feat(rust): migrate high-level runtime paths

Move long-distance matching and high-level decompression from C shims into
Rust. The decoder now owns context, dictionary, parameter, one-shot, and
buffered streaming state while C retains allocation/configuration, legacy,
and trace leaves.

Move CLI parsing, safety policy, and dispatch into a separate Rust static
archive. Keeping it separate prevents library builds from retaining FIO
symbols, while C continues to own file opening, replacement, and I/O.
Program targets now select matching compression/decompression archives.

The remaining C boundary is intentional: high-level compression, optimal
parsing, dictionary building, legacy callbacks, and CLI file I/O still need
migration.

Test Plan:
- cargo test --all-targets (native and i686)
- cargo test --all-targets in rust/cli (native and i686)
- CLI crate compression-only and decompression-only feature tests
- native and i686 fuzzer/zstreamtest runs, plus legacy and dictionary tests
- ZSTD_C_PREDICT and ZSTD_HEAPMODE=0 fuzzer coverage
- library, dynamic-link, and program-target build/round-trip matrix

Refs: rust/README.md
This commit is contained in:
2026-07-11 09:03:41 +02:00
parent 031592da1e
commit fe7e24c770
14 changed files with 6886 additions and 4651 deletions
+4
View File
@@ -32,6 +32,8 @@ pub mod zstd_compress_superblock;
#[cfg(feature = "decompression")]
pub mod zstd_ddict;
#[cfg(feature = "decompression")]
pub mod zstd_decompress;
#[cfg(feature = "decompression")]
pub mod zstd_decompress_block;
#[cfg(feature = "compression")]
pub mod zstd_double_fast;
@@ -40,6 +42,8 @@ pub mod zstd_fast;
#[cfg(feature = "compression")]
pub mod zstd_lazy;
#[cfg(feature = "compression")]
pub mod zstd_ldm;
#[cfg(feature = "compression")]
pub mod zstd_opt_tree;
#[cfg(feature = "compression")]
pub mod zstd_presplit;
+1541
View File
@@ -0,0 +1,1541 @@
#![allow(non_camel_case_types)]
#![allow(non_snake_case)]
#![allow(clippy::missing_safety_doc)]
//! Rust command-line frontend for zstd.
//!
//! This is intentionally a parser and dispatch layer, not a second file I/O
//! implementation. It reuses the mature C `fileio` layer through its narrow
//! public-in-the-programs-tree ABI: file opening, safe replacement, sparse
//! writes, dictionary loading, streaming, and metadata preservation remain in
//! `programs/fileio.c` for this first migration step.
//!
//! Remaining C-only CLI boundaries are called out in `unsupported()` below:
//! benchmark execution, dictionary training, recursive/file-list expansion,
//! tracing, alternate-format selection, and the advanced directory modes.
use std::env;
use std::ffi::{CStr, CString, OsStr, OsString};
use std::fs;
use std::io::{self, IsTerminal, Write};
use std::os::raw::{c_char, c_int, c_uint};
use std::path::Path;
use std::ptr;
#[cfg(unix)]
use std::os::unix::ffi::{OsStrExt, OsStringExt};
#[cfg(unix)]
use std::os::unix::fs::FileTypeExt;
const DEFAULT_CLEVEL: i32 = 3;
#[cfg(feature = "compression")]
const DEFAULT_MAX_CLEVEL: i32 = 19;
const DEFAULT_MEM_LIMIT: u32 = 1 << 27;
const DEFAULT_LONG_WINDOW_LOG: u32 = 27;
const MAX_FAST_ACCELERATION: i32 = 128 << 10;
const STDIN_MARK: &str = "/*stdin*\\";
const STDOUT_MARK: &str = "/*stdout*\\";
#[cfg(windows)]
const NULL_MARK: &str = "NUL";
#[cfg(not(windows))]
const NULL_MARK: &str = "/dev/null";
#[cfg(feature = "compression")]
const ZSTD_SUFFIX: &[u8] = b".zst\0";
const FIO_ZSTD_COMPRESSION: c_int = 0;
const FIO_PS_AUTO: c_int = 0;
const FIO_PS_NEVER: c_int = 1;
const FIO_PS_ALWAYS: c_int = 2;
const ZSTD_PS_AUTO: c_int = 0;
const ZSTD_PS_ENABLE: c_int = 1;
const ZSTD_PS_DISABLE: c_int = 2;
#[repr(C)]
struct FIO_prefs_t {
_private: [u8; 0],
}
#[repr(C)]
struct FIO_ctx_t {
_private: [u8; 0],
}
#[repr(C)]
#[derive(Clone, Copy, Debug, Default)]
struct ZSTD_compressionParameters {
windowLog: u32,
chainLog: u32,
hashLog: u32,
searchLog: u32,
minMatch: u32,
targetLength: u32,
strategy: c_int,
}
unsafe extern "C" {
fn ZSTD_versionString() -> *const c_char;
fn ZSTD_rust_cli_expected_version() -> *const c_char;
static mut g_utilDisplayLevel: c_int;
#[cfg(feature = "compression")]
fn ZSTD_minCLevel() -> c_int;
#[cfg(feature = "compression")]
fn ZSTD_maxCLevel() -> c_int;
#[cfg(feature = "compression")]
fn UTIL_countPhysicalCores() -> c_int;
#[cfg(feature = "compression")]
fn UTIL_countLogicalCores() -> c_int;
fn FIO_createPreferences() -> *mut FIO_prefs_t;
fn FIO_freePreferences(prefs: *mut FIO_prefs_t);
fn FIO_createContext() -> *mut FIO_ctx_t;
fn FIO_freeContext(ctx: *mut FIO_ctx_t);
fn FIO_addAbortHandler();
fn FIO_setCompressionType(prefs: *mut FIO_prefs_t, compression_type: c_int);
fn FIO_overwriteMode(prefs: *mut FIO_prefs_t);
fn FIO_setAdaptiveMode(prefs: *mut FIO_prefs_t, adapt: c_int);
#[cfg(feature = "compression")]
fn FIO_setAdaptMin(prefs: *mut FIO_prefs_t, level: c_int);
#[cfg(feature = "compression")]
fn FIO_setAdaptMax(prefs: *mut FIO_prefs_t, level: c_int);
fn FIO_setUseRowMatchFinder(prefs: *mut FIO_prefs_t, mode: c_int);
fn FIO_setBlockSize(prefs: *mut FIO_prefs_t, block_size: c_int);
fn FIO_setChecksumFlag(prefs: *mut FIO_prefs_t, checksum: c_int);
fn FIO_setDictIDFlag(prefs: *mut FIO_prefs_t, dict_id: c_int);
fn FIO_setLdmBucketSizeLog(prefs: *mut FIO_prefs_t, value: c_int);
fn FIO_setLdmFlag(prefs: *mut FIO_prefs_t, value: c_uint);
fn FIO_setLdmHashRateLog(prefs: *mut FIO_prefs_t, value: c_int);
fn FIO_setLdmHashLog(prefs: *mut FIO_prefs_t, value: c_int);
fn FIO_setLdmMinMatch(prefs: *mut FIO_prefs_t, value: c_int);
fn FIO_setMemLimit(prefs: *mut FIO_prefs_t, limit: c_uint);
#[cfg(feature = "compression")]
fn FIO_setNbWorkers(prefs: *mut FIO_prefs_t, workers: c_int);
fn FIO_setOverlapLog(prefs: *mut FIO_prefs_t, value: c_int);
fn FIO_setRemoveSrcFile(prefs: *mut FIO_prefs_t, value: c_int);
fn FIO_setSparseWrite(prefs: *mut FIO_prefs_t, value: c_int);
fn FIO_setRsyncable(prefs: *mut FIO_prefs_t, value: c_int);
fn FIO_setStreamSrcSize(prefs: *mut FIO_prefs_t, value: usize);
fn FIO_setTargetCBlockSize(prefs: *mut FIO_prefs_t, value: usize);
fn FIO_setSrcSizeHint(prefs: *mut FIO_prefs_t, value: usize);
#[cfg(feature = "decompression")]
fn FIO_setTestMode(prefs: *mut FIO_prefs_t, value: c_int);
fn FIO_setLiteralCompressionMode(prefs: *mut FIO_prefs_t, value: c_int);
fn FIO_setProgressSetting(value: c_int);
fn FIO_setNotificationLevel(value: c_int);
fn FIO_setExcludeCompressedFile(prefs: *mut FIO_prefs_t, value: c_int);
fn FIO_setAllowBlockDevices(prefs: *mut FIO_prefs_t, value: c_int);
fn FIO_setContentSize(prefs: *mut FIO_prefs_t, value: c_int);
fn FIO_setAsyncIOFlag(prefs: *mut FIO_prefs_t, value: c_int);
fn FIO_setPassThroughFlag(prefs: *mut FIO_prefs_t, value: c_int);
fn FIO_setMMapDict(prefs: *mut FIO_prefs_t, value: c_int);
fn FIO_setNbFilesTotal(ctx: *mut FIO_ctx_t, value: c_int);
fn FIO_setHasStdinInput(ctx: *mut FIO_ctx_t, value: c_int);
fn FIO_setHasStdoutOutput(ctx: *mut FIO_ctx_t, value: c_int);
#[cfg(feature = "compression")]
fn FIO_compressFilename(
ctx: *mut FIO_ctx_t,
prefs: *mut FIO_prefs_t,
output: *const c_char,
input: *const c_char,
dict: *const c_char,
level: c_int,
params: ZSTD_compressionParameters,
) -> c_int;
#[cfg(feature = "decompression")]
fn FIO_decompressFilename(
ctx: *mut FIO_ctx_t,
prefs: *mut FIO_prefs_t,
output: *const c_char,
input: *const c_char,
dict: *const c_char,
) -> c_int;
#[cfg(feature = "compression")]
fn FIO_compressMultipleFilenames(
ctx: *mut FIO_ctx_t,
prefs: *mut FIO_prefs_t,
inputs: *const *const c_char,
output_mirror_dir: *const c_char,
output_dir: *const c_char,
output: *const c_char,
suffix: *const c_char,
dict: *const c_char,
level: c_int,
params: ZSTD_compressionParameters,
) -> c_int;
#[cfg(feature = "decompression")]
fn FIO_decompressMultipleFilenames(
ctx: *mut FIO_ctx_t,
prefs: *mut FIO_prefs_t,
inputs: *const *const c_char,
output_mirror_dir: *const c_char,
output_dir: *const c_char,
output: *const c_char,
dict: *const c_char,
) -> c_int;
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum Operation {
Compress,
Decompress,
Test,
}
#[derive(Debug)]
enum Action {
Run(Box<Cli>),
Help { advanced: bool },
Version { quiet: bool },
}
#[derive(Debug)]
struct Cli {
operation: Operation,
inputs: Vec<CString>,
output: Option<CString>,
dictionary: Option<CString>,
level: i32,
ultra: bool,
display_level: i32,
force: bool,
force_stdout: bool,
remove_source: bool,
checksum: Option<i32>,
sparse: Option<i32>,
pass_through: Option<i32>,
content_size: i32,
dict_id: Option<i32>,
async_io: Option<i32>,
mmap_dict: i32,
progress: i32,
workers: Option<i32>,
block_size: Option<usize>,
mem_limit: Option<u32>,
ldm: bool,
ldm_hash_log: Option<i32>,
ldm_min_match: Option<i32>,
ldm_bucket_size_log: Option<i32>,
ldm_hash_rate_log: Option<i32>,
overlap_log: Option<i32>,
adapt: bool,
adapt_min: Option<i32>,
adapt_max: Option<i32>,
rsyncable: bool,
stream_src_size: Option<usize>,
target_cblock_size: Option<usize>,
src_size_hint: Option<usize>,
literal_compression: Option<i32>,
row_match_finder: i32,
exclude_compressed: bool,
compression_params: ZSTD_compressionParameters,
unsupported_program: Option<String>,
}
impl Cli {
fn new(program_name: &str) -> Self {
let mut cli = Self {
operation: Operation::Compress,
inputs: Vec::new(),
output: None,
dictionary: None,
level: default_level(),
ultra: false,
display_level: 2,
force: false,
force_stdout: false,
remove_source: false,
checksum: None,
sparse: None,
pass_through: None,
content_size: 1,
dict_id: None,
async_io: None,
mmap_dict: ZSTD_PS_AUTO,
progress: FIO_PS_AUTO,
workers: None,
block_size: None,
mem_limit: None,
ldm: false,
ldm_hash_log: None,
ldm_min_match: None,
ldm_bucket_size_log: None,
ldm_hash_rate_log: None,
overlap_log: None,
adapt: false,
adapt_min: None,
adapt_max: None,
rsyncable: false,
stream_src_size: None,
target_cblock_size: None,
src_size_hint: None,
literal_compression: None,
row_match_finder: ZSTD_PS_AUTO,
exclude_compressed: false,
compression_params: ZSTD_compressionParameters::default(),
unsupported_program: None,
};
match program_name {
"unzstd" => cli.operation = Operation::Decompress,
"zstdmt" => cli.workers = Some(0),
"zstdcat" | "zcat" => {
cli.operation = Operation::Decompress;
cli.output = Some(cstring(STDOUT_MARK).expect("static stdout marker"));
cli.force = true;
cli.force_stdout = true;
cli.pass_through = Some(1);
cli.display_level = 1;
}
"gzip" | "gunzip" | "gzcat" | "lzma" | "unlzma" | "xz" | "unxz" | "lz4" | "unlz4" => {
cli.unsupported_program = Some(program_name.to_owned())
}
_ => {}
}
cli
}
}
fn cstring(value: &str) -> Result<CString, String> {
CString::new(value).map_err(|_| format!("argument contains an interior NUL: {value:?}"))
}
fn os_cstring(value: &OsStr) -> Result<CString, String> {
#[cfg(unix)]
{
CString::new(value.as_bytes()).map_err(|_| "file name contains an interior NUL".to_owned())
}
#[cfg(not(unix))]
{
cstring(&value.to_string_lossy())
}
}
fn default_level() -> i32 {
match env::var("ZSTD_CLEVEL") {
Ok(value) => value.parse::<i32>().unwrap_or(DEFAULT_CLEVEL),
Err(_) => DEFAULT_CLEVEL,
}
}
#[cfg(feature = "compression")]
unsafe fn default_worker_count() -> i32 {
if let Ok(value) = env::var("ZSTD_NBTHREADS") {
if let Ok(workers) = value.parse::<u32>() {
if let Ok(workers) = i32::try_from(workers) {
return workers;
}
}
}
let logical_cores = unsafe { UTIL_countLogicalCores() }.max(1);
(logical_cores / 4).clamp(1, 4)
}
#[cfg(feature = "compression")]
unsafe fn resolved_worker_count(workers: Option<i32>) -> i32 {
match workers {
Some(0) => unsafe { UTIL_countPhysicalCores() }.max(1),
Some(workers) => workers,
None => unsafe { default_worker_count() },
}
}
fn program_basename(value: &OsStr) -> String {
Path::new(value)
.file_name()
.unwrap_or(value)
.to_string_lossy()
.split('.')
.next()
.unwrap_or("zstd")
.to_owned()
}
fn usage(advanced: bool) {
let mut out = io::stdout().lock();
let _ = writeln!(
out,
"Compress or decompress INPUT file(s); reads stdin when INPUT is '-' or omitted."
);
let _ = writeln!(out, "\nUsage: zstd [OPTIONS...] [INPUT... | -] [-o OUTPUT]");
let _ = writeln!(out, "\nCore options:");
let _ = writeln!(
out,
" -o OUTPUT, -c, --stdout Select output file or stdout"
);
let _ = writeln!(out, " -d, --decompress Decompress");
let _ = writeln!(out, " -t, --test Test compressed input");
let _ = writeln!(
out,
" -# Compression level (default {DEFAULT_CLEVEL})"
);
let _ = writeln!(out, " -D DICT Use a dictionary");
let _ = writeln!(
out,
" -f, --force Overwrite output / allow stdio"
);
let _ = writeln!(
out,
" -k, --keep | --rm Preserve or remove source after success"
);
let _ = writeln!(out, " -q, --quiet | -v, --verbose Adjust display level");
let _ = writeln!(out, " -V, --version Print version");
let _ = writeln!(out, " -h | -H, --help Print help");
if advanced {
let _ = writeln!(out, "\nImplemented advanced compression controls:");
let _ = writeln!(
out,
" --fast[=#], --ultra, --long[=#], --threads=#, --block-size=#"
);
let _ = writeln!(
out,
" --zstd=wlog=#,clog=#,hlog=#,slog=#,mml=#,tlen=#,strat=#"
);
let _ = writeln!(
out,
" --[no-]check, --[no-]sparse, --[no-]progress, --[no-]asyncio"
);
let _ = writeln!(
out,
" --adapt[=min=#,max=#], --rsyncable, --[no-]row-match-finder"
);
let _ = writeln!(
out,
"\nNot yet migrated: benchmark, dictionary training, recursive/file-list expansion,"
);
let _ = writeln!(out, "trace, alternate formats, and output-directory modes.");
}
}
fn print_version(quiet: bool) {
let version = unsafe { CStr::from_ptr(ZSTD_versionString()) }
.to_string_lossy()
.into_owned();
if quiet {
println!("{version}");
} else {
println!(
"*** Zstandard CLI ({}-bit) v{version}, by Yann Collet ***",
usize::BITS
);
}
}
fn check_lib_version() -> Result<(), String> {
let expected = unsafe { CStr::from_ptr(ZSTD_rust_cli_expected_version()) };
let actual = unsafe { CStr::from_ptr(ZSTD_versionString()) };
if expected == actual {
return Ok(());
}
Err(format!(
"incorrect library version (expecting: {}; actual: {})",
expected.to_string_lossy(),
actual.to_string_lossy()
))
}
fn parse_size(value: &str) -> Result<usize, String> {
let split = value
.find(|character: char| !character.is_ascii_digit())
.unwrap_or(value.len());
let (digits, suffix) = value.split_at(split);
if digits.is_empty() {
return Err(format!("expected a numeric value, got {value:?}"));
}
let mut number = digits
.parse::<usize>()
.map_err(|_| format!("numeric value overflows size_t: {value:?}"))?;
let normalized = suffix.trim_end_matches('B').trim_end_matches('i');
let shift = match normalized {
"" => 0,
"K" | "k" => 10,
"M" | "m" => 20,
"G" | "g" => 30,
_ => return Err(format!("unsupported numeric suffix in {value:?}")),
};
number = number
.checked_shl(shift)
.ok_or_else(|| format!("numeric value overflows size_t: {value:?}"))?;
Ok(number)
}
fn parse_u32(value: &str, name: &str) -> Result<u32, String> {
let size = parse_size(value)?;
u32::try_from(size).map_err(|_| format!("{name} is too large: {value:?}"))
}
fn parse_i32(value: &str, name: &str) -> Result<i32, String> {
value
.parse::<i32>()
.map_err(|_| format!("invalid {name}: {value:?}"))
}
fn parse_worker_count(value: &str) -> Result<i32, String> {
let workers = parse_i32(value, "thread count")?;
if workers < 0 {
return Err(format!("thread count must not be negative: {value:?}"));
}
Ok(workers)
}
fn next_value(
attached: Option<&str>,
args: &[OsString],
index: &mut usize,
option: &str,
) -> Result<String, String> {
if let Some(value) = attached {
if !value.is_empty() {
return Ok(value.to_owned());
}
}
*index += 1;
let Some(value) = args.get(*index) else {
return Err(format!("missing argument for {option}"));
};
let rendered = value.to_string_lossy().into_owned();
if rendered.starts_with('-') {
return Err(format!(
"{option} cannot be separated from its argument by another option"
));
}
Ok(rendered)
}
fn next_os_value(args: &[OsString], index: &mut usize, option: &str) -> Result<OsString, String> {
*index += 1;
let Some(value) = args.get(*index) else {
return Err(format!("missing argument for {option}"));
};
if value.to_string_lossy().starts_with('-') {
return Err(format!(
"{option} cannot be separated from its argument by another option"
));
}
Ok(value.clone())
}
#[cfg(unix)]
fn short_attached_value(value: &OsStr, start: usize) -> Option<OsString> {
let bytes = &value.as_bytes()[start..];
if bytes.is_empty() {
return None;
}
let bytes = if bytes.first() == Some(&b'=') {
&bytes[1..]
} else {
bytes
};
Some(OsString::from_vec(bytes.to_vec()))
}
#[cfg(not(unix))]
fn short_attached_value(value: &OsStr, start: usize) -> Option<OsString> {
let rendered = value.to_string_lossy();
let attached = &rendered[start..];
if attached.is_empty() {
return None;
}
Some(OsString::from(
attached.strip_prefix('=').unwrap_or(attached),
))
}
fn parse_compression_parameters(value: &str, cli: &mut Cli) -> Result<(), String> {
for item in value.split(',') {
let Some((name, raw)) = item.split_once('=') else {
return Err(format!("invalid --zstd parameter {item:?}"));
};
let parsed = parse_u32(raw, name)?;
match name {
"windowLog" | "wlog" => cli.compression_params.windowLog = parsed,
"chainLog" | "clog" => cli.compression_params.chainLog = parsed,
"hashLog" | "hlog" => cli.compression_params.hashLog = parsed,
"searchLog" | "slog" => cli.compression_params.searchLog = parsed,
"minMatch" | "mml" => cli.compression_params.minMatch = parsed,
"targetLength" | "tlen" => cli.compression_params.targetLength = parsed,
"strategy" | "strat" => cli.compression_params.strategy = parsed as c_int,
"overlapLog" | "ovlog" => cli.overlap_log = Some(parsed as i32),
"ldmHashLog" | "lhlog" => cli.ldm_hash_log = Some(parsed as i32),
"ldmMinMatch" | "lmml" => cli.ldm_min_match = Some(parsed as i32),
"ldmBucketSizeLog" | "lblog" => cli.ldm_bucket_size_log = Some(parsed as i32),
"ldmHashRateLog" | "lhrlog" => cli.ldm_hash_rate_log = Some(parsed as i32),
_ => return Err(format!("unknown --zstd parameter {name:?}")),
}
}
Ok(())
}
fn parse_adapt(value: &str, cli: &mut Cli) -> Result<(), String> {
cli.adapt = true;
if value.is_empty() {
return Ok(());
}
for item in value.split(',') {
let Some((name, raw)) = item.split_once('=') else {
return Err(format!("invalid --adapt parameter {item:?}"));
};
match name {
"min" => cli.adapt_min = Some(parse_i32(raw, "adapt minimum")?),
"max" => cli.adapt_max = Some(parse_i32(raw, "adapt maximum")?),
_ => return Err(format!("unknown --adapt parameter {name:?}")),
}
}
if let (Some(minimum), Some(maximum)) = (cli.adapt_min, cli.adapt_max) {
if minimum > maximum {
return Err("--adapt minimum must not exceed its maximum".to_owned());
}
}
Ok(())
}
fn unsupported(option: &str) -> Result<(), String> {
Err(format!(
"{option} is not yet implemented by the Rust CLI frontend"
))
}
fn parse_long_option(
option: &str,
args: &[OsString],
index: &mut usize,
cli: &mut Cli,
) -> Result<Option<Action>, String> {
let (name, attached) = option
.split_once('=')
.map_or((option, None), |(name, value)| (name, Some(value)));
if attached.is_some()
&& matches!(
name,
"--compress"
| "--decompress"
| "--uncompress"
| "--test"
| "--force"
| "--keep"
| "--rm"
| "--stdout"
| "--version"
| "--help"
| "--verbose"
| "--quiet"
| "--check"
| "--no-check"
| "--sparse"
| "--no-sparse"
| "--pass-through"
| "--no-pass-through"
| "--content-size"
| "--no-content-size"
| "--no-dictID"
| "--asyncio"
| "--no-asyncio"
| "--mmap-dict"
| "--no-mmap-dict"
| "--progress"
| "--no-progress"
| "--ultra"
| "--no-row-match-finder"
| "--row-match-finder"
| "--rsyncable"
| "--compress-literals"
| "--no-compress-literals"
| "--exclude-compressed"
| "--no-name"
)
{
return Err(format!("{name} does not take an argument"));
}
match name {
"--" => Ok(None),
"--compress" => {
cli.operation = Operation::Compress;
Ok(None)
}
"--decompress" | "--uncompress" => {
cli.operation = Operation::Decompress;
Ok(None)
}
"--test" => {
cli.operation = Operation::Test;
Ok(None)
}
"--force" => {
cli.force = true;
Ok(None)
}
"--keep" => {
cli.remove_source = false;
Ok(None)
}
"--no-name" => Ok(None),
"--rm" => {
cli.remove_source = true;
Ok(None)
}
"--stdout" => {
cli.output = Some(cstring(STDOUT_MARK)?);
cli.force_stdout = true;
Ok(None)
}
"--version" => Ok(Some(Action::Version {
quiet: cli.display_level < 2,
})),
"--help" => Ok(Some(Action::Help { advanced: true })),
"--verbose" => {
cli.display_level += 1;
Ok(None)
}
"--quiet" => {
cli.display_level -= 1;
Ok(None)
}
"--check" => {
cli.checksum = Some(2);
Ok(None)
}
"--no-check" => {
cli.checksum = Some(0);
Ok(None)
}
"--sparse" => {
cli.sparse = Some(2);
Ok(None)
}
"--no-sparse" => {
cli.sparse = Some(0);
Ok(None)
}
"--pass-through" => {
cli.pass_through = Some(1);
Ok(None)
}
"--no-pass-through" => {
cli.pass_through = Some(0);
Ok(None)
}
"--content-size" => {
cli.content_size = 1;
Ok(None)
}
"--no-content-size" => {
cli.content_size = 0;
Ok(None)
}
"--no-dictID" => {
cli.dict_id = Some(0);
Ok(None)
}
"--asyncio" => {
cli.async_io = Some(1);
Ok(None)
}
"--no-asyncio" => {
cli.async_io = Some(0);
Ok(None)
}
"--mmap-dict" => {
cli.mmap_dict = ZSTD_PS_ENABLE;
Ok(None)
}
"--no-mmap-dict" => {
cli.mmap_dict = ZSTD_PS_DISABLE;
Ok(None)
}
"--progress" => {
cli.progress = FIO_PS_ALWAYS;
Ok(None)
}
"--no-progress" => {
cli.progress = FIO_PS_NEVER;
Ok(None)
}
"--ultra" => {
cli.ultra = true;
Ok(None)
}
"--fast" => {
let mut level = match attached {
Some(value) => parse_i32(value, "fast level")?,
None => 1,
};
if level <= 0 {
return Err("fast level must be positive".to_owned());
}
level = level.min(MAX_FAST_ACCELERATION);
cli.level = -level;
Ok(None)
}
"--long" => {
cli.ldm = true;
cli.ultra = true;
if let Some(value) = attached {
cli.compression_params.windowLog = parse_u32(value, "long window log")?;
} else if cli.compression_params.windowLog == 0 {
cli.compression_params.windowLog = DEFAULT_LONG_WINDOW_LOG;
}
Ok(None)
}
"--adapt" => {
parse_adapt(attached.unwrap_or(""), cli)?;
Ok(None)
}
"--no-row-match-finder" => {
cli.row_match_finder = ZSTD_PS_DISABLE;
Ok(None)
}
"--row-match-finder" => {
cli.row_match_finder = ZSTD_PS_ENABLE;
Ok(None)
}
"--rsyncable" => {
cli.rsyncable = true;
Ok(None)
}
"--compress-literals" => {
cli.literal_compression = Some(ZSTD_PS_ENABLE);
Ok(None)
}
"--no-compress-literals" => {
cli.literal_compression = Some(ZSTD_PS_DISABLE);
Ok(None)
}
"--exclude-compressed" => {
cli.exclude_compressed = true;
Ok(None)
}
"--threads" => {
let value = next_value(attached, args, index, "--threads")?;
cli.workers = Some(parse_worker_count(&value)?);
Ok(None)
}
"--memlimit" | "--memory" | "--memlimit-decompress" => {
let value = next_value(attached, args, index, name)?;
cli.mem_limit = Some(parse_u32(&value, "memory limit")?);
Ok(None)
}
"--block-size" => {
let value = next_value(attached, args, index, name)?;
cli.block_size = Some(parse_size(&value)?);
Ok(None)
}
"--stream-size" => {
let value = next_value(attached, args, index, name)?;
cli.stream_src_size = Some(parse_size(&value)?);
Ok(None)
}
"--target-compressed-block-size" => {
let value = next_value(attached, args, index, name)?;
cli.target_cblock_size = Some(parse_size(&value)?);
Ok(None)
}
"--size-hint" => {
let value = next_value(attached, args, index, name)?;
cli.src_size_hint = Some(parse_size(&value)?);
Ok(None)
}
"--zstd" => {
let value = next_value(attached, args, index, "--zstd")?;
parse_compression_parameters(&value, cli)?;
Ok(None)
}
"--list"
| "--train"
| "--train-cover"
| "--train-fastcover"
| "--train-legacy"
| "--max"
| "--maxdict"
| "--dictID"
| "--filelist"
| "--output-dir-flat"
| "--output-dir-mirror"
| "--patch-from"
| "--trace"
| "--format"
| "--priority"
| "--single-thread"
| "--auto-threads"
| "--fake-stdin-is-console"
| "--fake-stdout-is-console"
| "--fake-stderr-is-console"
| "--trace-file-stat"
| "--show-default-cparams" => {
unsupported(name)?;
Ok(None)
}
_ => Err(format!("unknown option {option:?}")),
}
}
fn parse_short_options(
value: &str,
raw_value: &OsStr,
args: &[OsString],
index: &mut usize,
cli: &mut Cli,
) -> Result<Option<Action>, String> {
let mut offset = 1usize;
let bytes = value.as_bytes();
while offset < bytes.len() {
let option = bytes[offset] as char;
match option {
'0'..='9' => {
let mut digits_end = offset;
while digits_end < bytes.len() && bytes[digits_end].is_ascii_digit() {
digits_end += 1;
}
cli.level = parse_i32(&value[offset..digits_end], "compression level")?;
offset = digits_end;
continue;
}
'd' => cli.operation = Operation::Decompress,
'z' => cli.operation = Operation::Compress,
't' => cli.operation = Operation::Test,
'c' => {
cli.output = Some(cstring(STDOUT_MARK)?);
cli.force_stdout = true;
}
'f' => cli.force = true,
'k' => cli.remove_source = false,
'n' => {}
'q' => cli.display_level -= 1,
'v' => cli.display_level += 1,
'C' => cli.checksum = Some(2),
'h' => return Ok(Some(Action::Help { advanced: false })),
'H' => return Ok(Some(Action::Help { advanced: true })),
'V' => {
return Ok(Some(Action::Version {
quiet: cli.display_level < 2,
}))
}
'o' | 'D' | 'T' | 'M' | 'B' => {
let attached = short_attached_value(raw_value, offset + 1);
let argument = attached
.map(Ok)
.unwrap_or_else(|| next_os_value(args, index, &format!("-{option}")))?;
match option {
'o' => cli.output = Some(os_cstring(&argument)?),
'D' => cli.dictionary = Some(os_cstring(&argument)?),
'T' => cli.workers = Some(parse_worker_count(&argument.to_string_lossy())?),
'M' => {
cli.mem_limit =
Some(parse_u32(&argument.to_string_lossy(), "memory limit")?)
}
'B' => cli.block_size = Some(parse_size(&argument.to_string_lossy())?),
_ => unreachable!(),
}
break;
}
'b' | 'e' | 'i' | 'l' | 'p' | 'P' | 'r' | 's' | 'S' => {
unsupported(&format!("-{option}"))?;
}
_ => return Err(format!("unknown option -{option}")),
}
offset += 1;
}
Ok(None)
}
fn parse_args(args: Vec<OsString>) -> Result<Action, String> {
let program_name = args
.first()
.map_or_else(|| "zstd".to_owned(), |value| program_basename(value));
let mut cli = Cli::new(&program_name);
let mut end_of_options = false;
let mut index = 1usize;
while index < args.len() {
let rendered = args[index].to_string_lossy().into_owned();
if end_of_options {
cli.inputs.push(os_cstring(&args[index])?);
} else if rendered == "--" {
end_of_options = true;
} else if rendered == "-" {
cli.inputs.push(cstring(STDIN_MARK)?);
} else if rendered.starts_with("--") {
if let Some(action) = parse_long_option(&rendered, &args, &mut index, &mut cli)? {
return Ok(action);
}
} else if rendered.starts_with('-') {
if let Some(action) =
parse_short_options(&rendered, &args[index], &args, &mut index, &mut cli)?
{
return Ok(action);
}
} else {
cli.inputs.push(os_cstring(&args[index])?);
}
index += 1;
}
Ok(Action::Run(Box::new(cli)))
}
unsafe fn apply_preferences(cli: &Cli, prefs: *mut FIO_prefs_t, ctx: *mut FIO_ctx_t) {
let display_level = if is_stdout(cli.output.as_ref()) && cli.display_level == 2 {
1
} else {
cli.display_level
};
unsafe {
g_utilDisplayLevel = display_level;
FIO_setCompressionType(prefs, FIO_ZSTD_COMPRESSION);
FIO_setNotificationLevel(display_level);
FIO_setProgressSetting(
if !io::stderr().is_terminal() && cli.progress != FIO_PS_ALWAYS {
FIO_PS_NEVER
} else {
cli.progress
},
);
FIO_setRemoveSrcFile(
prefs,
i32::from(
cli.remove_source
&& cli.operation != Operation::Test
&& !is_stdout(cli.output.as_ref()),
),
);
FIO_setAllowBlockDevices(prefs, i32::from(cli.force));
FIO_setMMapDict(prefs, cli.mmap_dict);
FIO_setUseRowMatchFinder(prefs, cli.row_match_finder);
FIO_setMemLimit(
prefs,
cli.mem_limit
.filter(|limit| *limit != 0)
.unwrap_or_else(|| {
if cli.compression_params.windowLog == 0 {
DEFAULT_MEM_LIMIT
} else {
1_u32 << (cli.compression_params.windowLog & 31)
}
}),
);
#[cfg(feature = "compression")]
FIO_setNbWorkers(prefs, resolved_worker_count(cli.workers));
FIO_setLdmFlag(prefs, u32::from(cli.ldm));
FIO_setAdaptiveMode(prefs, i32::from(cli.adapt));
FIO_setRsyncable(prefs, i32::from(cli.rsyncable));
FIO_setExcludeCompressedFile(prefs, i32::from(cli.exclude_compressed));
if cli.force {
FIO_overwriteMode(prefs);
}
if let Some(value) = cli.checksum {
FIO_setChecksumFlag(prefs, value);
}
if cli.operation == Operation::Compress {
FIO_setSparseWrite(prefs, 0);
} else if let Some(value) = cli.sparse {
FIO_setSparseWrite(prefs, value);
}
if let Some(value) = cli.pass_through {
FIO_setPassThroughFlag(prefs, value);
}
FIO_setContentSize(prefs, cli.content_size);
if let Some(value) = cli.dict_id {
FIO_setDictIDFlag(prefs, value);
}
if let Some(value) = cli.async_io {
FIO_setAsyncIOFlag(prefs, value);
}
if let Some(value) = cli.block_size {
FIO_setBlockSize(prefs, value as c_int);
}
if let Some(value) = cli.ldm_hash_log {
FIO_setLdmHashLog(prefs, value);
}
if let Some(value) = cli.ldm_min_match {
FIO_setLdmMinMatch(prefs, value);
}
if let Some(value) = cli.ldm_bucket_size_log {
FIO_setLdmBucketSizeLog(prefs, value);
}
if let Some(value) = cli.ldm_hash_rate_log {
FIO_setLdmHashRateLog(prefs, value);
}
if let Some(value) = cli.overlap_log {
FIO_setOverlapLog(prefs, value);
}
#[cfg(feature = "compression")]
{
FIO_setAdaptMin(prefs, cli.adapt_min.unwrap_or_else(|| ZSTD_minCLevel()));
FIO_setAdaptMax(prefs, cli.adapt_max.unwrap_or_else(|| ZSTD_maxCLevel()));
}
if let Some(value) = cli.stream_src_size {
FIO_setStreamSrcSize(prefs, value);
}
if let Some(value) = cli.target_cblock_size {
FIO_setTargetCBlockSize(prefs, value);
}
if let Some(value) = cli.src_size_hint {
FIO_setSrcSizeHint(prefs, value);
}
if let Some(value) = cli.literal_compression {
FIO_setLiteralCompressionMode(prefs, value);
}
FIO_setNbFilesTotal(ctx, cli.inputs.len() as c_int);
FIO_setHasStdinInput(ctx, i32::from(cli.inputs.iter().any(is_stdin)));
FIO_setHasStdoutOutput(ctx, i32::from(is_stdout(cli.output.as_ref())));
}
}
fn is_stdout(value: Option<&CString>) -> bool {
value.is_some_and(|value| value.as_bytes() == STDOUT_MARK.as_bytes())
}
fn is_stdin(value: &CString) -> bool {
value.as_bytes() == STDIN_MARK.as_bytes()
}
#[cfg(unix)]
fn is_non_fifo_symlink(input: &CString) -> bool {
let path = Path::new(OsStr::from_bytes(input.as_bytes()));
let Ok(metadata) = fs::symlink_metadata(path) else {
return false;
};
metadata.file_type().is_symlink()
&& !fs::metadata(path).is_ok_and(|target| target.file_type().is_fifo())
}
#[cfg(not(unix))]
fn is_non_fifo_symlink(input: &CString) -> bool {
let path = Path::new(&input.to_string_lossy().into_owned());
fs::symlink_metadata(path).is_ok_and(|metadata| metadata.file_type().is_symlink())
}
fn filter_symlink_inputs(cli: &mut Cli) {
if cli.force {
return;
}
cli.inputs.retain(|input| {
if is_stdin(input) || !is_non_fifo_symlink(input) {
return true;
}
if cli.display_level >= 2 {
eprintln!(
"zstd: Warning : {} is a symbolic link, ignoring",
input.to_string_lossy()
);
}
false
});
}
fn check_terminal_safety(cli: &Cli) -> Result<(), String> {
let has_stdin = cli.inputs.iter().any(is_stdin);
if has_stdin && !cli.force && io::stdin().is_terminal() {
return Err("stdin is a console, aborting".to_owned());
}
if has_stdin
&& is_stdout(cli.output.as_ref())
&& !cli.force
&& !cli.force_stdout
&& cli.operation != Operation::Decompress
&& io::stdout().is_terminal()
{
return Err("stdout is a console, aborting".to_owned());
}
Ok(())
}
#[cfg(feature = "compression")]
unsafe fn run_compress(
cli: &Cli,
ctx: *mut FIO_ctx_t,
prefs: *mut FIO_prefs_t,
inputs: &[*const c_char],
output: *const c_char,
dictionary: *const c_char,
) -> c_int {
if inputs.len() == 1 && !output.is_null() {
unsafe {
FIO_compressFilename(
ctx,
prefs,
output,
inputs[0],
dictionary,
cli.level,
cli.compression_params,
)
}
} else {
unsafe {
FIO_compressMultipleFilenames(
ctx,
prefs,
inputs.as_ptr(),
ptr::null(),
ptr::null(),
output,
ZSTD_SUFFIX.as_ptr().cast(),
dictionary,
cli.level,
cli.compression_params,
)
}
}
}
#[cfg(feature = "decompression")]
unsafe fn run_decompress(
operation: Operation,
ctx: *mut FIO_ctx_t,
prefs: *mut FIO_prefs_t,
inputs: &[*const c_char],
output: *const c_char,
dictionary: *const c_char,
) -> c_int {
match operation {
Operation::Test => {
let null_output = cstring(NULL_MARK).expect("static null marker");
unsafe {
FIO_setTestMode(prefs, 1);
FIO_decompressMultipleFilenames(
ctx,
prefs,
inputs.as_ptr(),
ptr::null(),
ptr::null(),
null_output.as_ptr(),
dictionary,
)
}
}
Operation::Decompress if inputs.len() == 1 && !output.is_null() => unsafe {
FIO_decompressFilename(ctx, prefs, output, inputs[0], dictionary)
},
Operation::Decompress => unsafe {
FIO_decompressMultipleFilenames(
ctx,
prefs,
inputs.as_ptr(),
ptr::null(),
ptr::null(),
output,
dictionary,
)
},
Operation::Compress => unreachable!("compression is dispatched separately"),
}
}
fn run_cli(mut cli: Cli) -> Result<i32, String> {
if let Some(program_name) = &cli.unsupported_program {
return Err(format!(
"{program_name} compatibility mode is not yet implemented by the Rust CLI frontend"
));
}
let explicit_input_count = cli.inputs.len();
filter_symlink_inputs(&mut cli);
if explicit_input_count > 0 && cli.inputs.is_empty() {
return Ok(1);
}
if cli.operation == Operation::Test {
cli.output = Some(cstring(NULL_MARK)?);
cli.remove_source = false;
}
if cli.inputs.is_empty() {
cli.inputs.push(cstring(STDIN_MARK)?);
if cli.output.is_none() {
cli.output = Some(cstring(STDOUT_MARK)?);
}
}
if cli.inputs.len() == 1
&& cli.inputs[0].as_bytes() == STDIN_MARK.as_bytes()
&& cli.output.is_none()
{
cli.output = Some(cstring(STDOUT_MARK)?);
}
check_terminal_safety(&cli)?;
if cli.operation == Operation::Compress {
#[cfg(not(feature = "compression"))]
return Err("Compression not supported".to_owned());
#[cfg(feature = "compression")]
{
let min_level = unsafe { ZSTD_minCLevel() };
let max_level = unsafe { ZSTD_maxCLevel() };
let ceiling = if cli.ultra {
max_level
} else {
DEFAULT_MAX_CLEVEL.min(max_level)
};
if cli.level > ceiling {
eprintln!("zstd: warning: compression level reduced to {ceiling}");
cli.level = ceiling;
}
if cli.level < min_level {
return Err(format!(
"compression level {} is below {min_level}",
cli.level
));
}
if let (Some(minimum), Some(maximum)) = (cli.adapt_min, cli.adapt_max) {
if minimum > maximum {
return Err("adaptation minimum exceeds maximum".to_owned());
}
cli.level = cli.level.clamp(minimum, maximum);
}
}
} else {
#[cfg(not(feature = "decompression"))]
return Err("Decompression not supported".to_owned());
}
let prefs = unsafe { FIO_createPreferences() };
let ctx = unsafe { FIO_createContext() };
if prefs.is_null() || ctx.is_null() {
unsafe {
if !prefs.is_null() {
FIO_freePreferences(prefs);
}
if !ctx.is_null() {
FIO_freeContext(ctx);
}
}
return Err("could not allocate C file-I/O state".to_owned());
}
let result = {
unsafe {
FIO_addAbortHandler();
apply_preferences(&cli, prefs, ctx);
}
let output = cli
.output
.as_ref()
.map_or(ptr::null(), |value| value.as_ptr());
let dictionary = cli
.dictionary
.as_ref()
.map_or(ptr::null(), |value| value.as_ptr());
let inputs: Vec<*const c_char> = cli.inputs.iter().map(|value| value.as_ptr()).collect();
match cli.operation {
Operation::Compress => {
#[cfg(feature = "compression")]
{
unsafe { run_compress(&cli, ctx, prefs, &inputs, output, dictionary) }
}
#[cfg(not(feature = "compression"))]
unreachable!("unsupported compression was rejected above")
}
Operation::Decompress | Operation::Test => {
#[cfg(feature = "decompression")]
{
unsafe {
run_decompress(cli.operation, ctx, prefs, &inputs, output, dictionary)
}
}
#[cfg(not(feature = "decompression"))]
unreachable!("unsupported decompression was rejected above")
}
}
};
unsafe {
FIO_freePreferences(prefs);
FIO_freeContext(ctx);
}
Ok(result)
}
fn run_from_args(args: Vec<OsString>) -> c_int {
match parse_args(args) {
Ok(Action::Help { advanced }) => {
usage(advanced);
0
}
Ok(Action::Version { quiet }) => {
print_version(quiet);
0
}
Ok(Action::Run(cli)) => match run_cli(*cli) {
Ok(result) => result,
Err(error) => {
eprintln!("zstd: {error}");
1
}
},
Err(error) => {
eprintln!("zstd: {error}\nTry `zstd --help` for usage.");
1
}
}
}
unsafe fn argv_to_os_strings(
arg_count: c_int,
argv: *const *const c_char,
) -> Result<Vec<OsString>, String> {
if arg_count <= 0 || argv.is_null() {
return Err("invalid argv supplied by C main".to_owned());
}
let count = arg_count as usize;
let mut args = Vec::with_capacity(count);
for index in 0..count {
let argument = unsafe { *argv.add(index) };
if argument.is_null() {
return Err(format!("argv[{index}] is null"));
}
let bytes = unsafe { CStr::from_ptr(argument) }.to_bytes();
#[cfg(unix)]
args.push(OsString::from_vec(bytes.to_vec()));
#[cfg(not(unix))]
args.push(OsString::from(String::from_utf8_lossy(bytes).into_owned()));
}
Ok(args)
}
/// C `main()` entry point retained by the small `programs/zstdcli.c` forwarder.
#[no_mangle]
pub unsafe extern "C" fn ZSTD_rust_cli_main(arg_count: c_int, argv: *const *const c_char) -> c_int {
if let Err(error) = check_lib_version() {
eprintln!("zstd: {error}");
return 1;
}
match unsafe { argv_to_os_strings(arg_count, argv) } {
Ok(args) => run_from_args(args),
Err(error) => {
eprintln!("zstd: {error}");
1
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn parse(values: &[&str]) -> Cli {
let args = values.iter().map(OsString::from).collect();
match parse_args(args).expect("arguments should parse") {
Action::Run(cli) => *cli,
Action::Help { .. } | Action::Version { .. } => panic!("expected a run action"),
}
}
#[test]
fn defaults_preserve_the_c_fileio_contract() {
let cli = parse(&["zstd", "input"]);
assert_eq!(cli.content_size, 1);
assert_eq!(cli.workers, None);
assert_eq!(cli.operation, Operation::Compress);
}
#[test]
fn long_mode_uses_the_cli_default_window_and_enables_ultra() {
let cli = parse(&["zstd", "--long", "input"]);
assert!(cli.ldm);
assert!(cli.ultra);
assert_eq!(cli.compression_params.windowLog, DEFAULT_LONG_WINDOW_LOG);
}
#[test]
fn long_mode_does_not_replace_an_explicit_window_log() {
let cli = parse(&["zstd", "--zstd=wlog=25", "--long", "input"]);
assert_eq!(cli.compression_params.windowLog, 25);
}
#[test]
fn stdio_and_dictionary_short_options_are_preserved() {
let cli = parse(&["zstd", "-dc", "-D", "dict", "-"]);
assert_eq!(cli.operation, Operation::Decompress);
assert!(is_stdout(cli.output.as_ref()));
assert_eq!(
cli.dictionary.as_deref().map(CStr::to_bytes),
Some(&b"dict"[..])
);
assert_eq!(
cli.inputs
.iter()
.map(|input| input.as_bytes())
.collect::<Vec<_>>(),
vec![STDIN_MARK.as_bytes()]
);
}
#[test]
fn stdout_selection_does_not_enable_force_or_pass_through() {
let cli = parse(&["zstd", "-c", "input"]);
assert!(cli.force_stdout);
assert!(!cli.force);
assert_eq!(cli.pass_through, None);
}
#[test]
fn short_level_can_be_combined_with_other_flags() {
let cli = parse(&["zstd", "-5q", "input"]);
assert_eq!(cli.level, 5);
assert_eq!(cli.display_level, 1);
}
#[test]
fn short_option_equals_form_is_accepted() {
let cli = parse(&["zstd", "-T=2", "-M=64M", "-B=1M", "input"]);
assert_eq!(cli.workers, Some(2));
assert_eq!(cli.mem_limit, Some(64 << 20));
assert_eq!(cli.block_size, Some(1 << 20));
}
#[test]
fn valueless_long_flags_reject_attached_values() {
let error = parse_args(vec![OsString::from("zstd"), OsString::from("--rm=0")])
.expect_err("an attached value must not activate --rm");
assert!(error.contains("does not take an argument"));
}
#[test]
fn zstdmt_uses_auto_threads() {
let cli = parse(&["zstdmt", "input"]);
assert_eq!(cli.workers, Some(0));
}
#[test]
fn alternate_format_aliases_fail_before_processing_files() {
let cli = Cli::new("gzip");
assert_eq!(cli.unsupported_program.as_deref(), Some("gzip"));
}
#[cfg(unix)]
#[test]
fn short_path_arguments_keep_non_utf8_bytes() {
let output = OsString::from_vec(vec![b'-', b'o', 0xff, b'.', b'z', b's', b't']);
let action = parse_args(vec![
OsString::from("zstd"),
output,
OsString::from("input"),
])
.expect("arguments should parse");
let Action::Run(cli) = action else {
panic!("expected a run action");
};
assert_eq!(
cli.output.as_deref().map(CStr::to_bytes),
Some(&b"\xff.zst"[..])
);
}
#[test]
fn unsupported_modes_fail_during_parsing() {
let error = parse_args(vec![OsString::from("zstd"), OsString::from("--train")])
.expect_err("training has not yet been migrated");
assert!(error.contains("not yet implemented"));
}
}
+3474
View File
@@ -0,0 +1,3474 @@
#![allow(non_camel_case_types)]
#![allow(non_snake_case)]
#![allow(clippy::missing_safety_doc)]
#![allow(clippy::too_many_arguments)]
#![allow(clippy::not_unsafe_ptr_arg_deref)]
//! Frame, context, and streaming decompression orchestration.
//!
//! `ZSTD_DCtx` deliberately remains C-owned. The companion C translation
//! unit projects its build-configuration-dependent leaves into
//! [`ZSTD_rustDctxView`]; this module owns the decoder state machine and
//! public ABI while never assumes offsets for the private C context.
use crate::entropy_common::FSE_readNCount;
use crate::errors::{ERR_isError, ZstdErrorCode, ERROR};
#[cfg(feature = "huf-force-decompress-x1")]
use crate::huf_decompress::HUF_readDTableX1_wksp;
#[cfg(not(feature = "huf-force-decompress-x1"))]
use crate::huf_decompress::HUF_readDTableX2_wksp;
use crate::mem::{MEM_32bits, MEM_readLE16, MEM_readLE32, MEM_readLE64};
use crate::xxhash::{XXH64_digest, XXH64_reset, XXH64_state_t, XXH64_update, XXH64};
use crate::zstd_ddict::{
ZSTD_DDict, ZSTD_DDict_dictContent, ZSTD_DDict_dictSize, ZSTD_copyDDictParameters,
ZSTD_freeDDict, ZSTD_getDictID_fromDDict,
};
use std::cmp::{max, min};
use std::ffi::c_void;
use std::mem::{size_of, MaybeUninit};
use std::os::raw::{c_int, c_uint};
use std::ptr;
const ZSTD_MAGICNUMBER: u32 = 0xFD2F_B528;
const ZSTD_MAGIC_DICTIONARY: u32 = 0xEC30_A437;
const ZSTD_MAGIC_SKIPPABLE_START: u32 = 0x184D_2A50;
const ZSTD_MAGIC_SKIPPABLE_MASK: u32 = 0xFFFF_FFF0;
const ZSTD_FRAMEIDSIZE: usize = 4;
const ZSTD_SKIPPABLEHEADERSIZE: usize = 8;
const ZSTD_BLOCKHEADERSIZE: usize = 3;
const ZSTD_BLOCKSIZE_MAX: usize = 128 << 10;
const ZSTD_BLOCKSIZE_MAX_MIN: usize = 1 << 10;
const ZSTD_WINDOWLOG_ABSOLUTEMIN: usize = 10;
const ZSTD_WINDOWLOG_LIMIT_DEFAULT: usize = 27;
const ZSTD_WINDOWLOG_MAX_32: usize = 30;
const ZSTD_WINDOWLOG_MAX_64: usize = 31;
const WILDCOPY_OVERLENGTH: usize = 32;
const ZSTD_WORKSPACETOOLARGE_FACTOR: usize = 3;
const ZSTD_WORKSPACETOOLARGE_MAXDURATION: usize = 128;
const ZSTD_HUFFDTABLE_CAPACITY_LOG: usize = 12;
const HUF_DTABLE_SIZE: usize = 1 + (1 << ZSTD_HUFFDTABLE_CAPACITY_LOG);
const ZSTD_BUILD_FSE_TABLE_WKSP_SIZE_U32: usize = 157;
const LL_FSE_LOG: usize = 9;
const OFF_FSE_LOG: usize = 8;
const ML_FSE_LOG: usize = 9;
const MAX_LL: usize = 35;
const MAX_ML: usize = 52;
const MAX_OFF: usize = 31;
const ZSTD_REP_NUM: usize = 3;
const ZSTD_CONTENTSIZE_UNKNOWN: u64 = u64::MAX;
const ZSTD_CONTENTSIZE_ERROR: u64 = u64::MAX - 1;
const ZSTD_F_ZSTD1: c_int = 0;
const ZSTD_F_ZSTD1_MAGICLESS: c_int = 1;
const ZSTD_FRAME: c_int = 0;
const ZSTD_SKIPPABLE_FRAME: c_int = 1;
const ZSTD_BM_BUFFERED: c_int = 0;
const ZSTD_BM_STABLE: c_int = 1;
const ZSTD_D_VALIDATE_CHECKSUM: c_int = 0;
const ZSTD_D_IGNORE_CHECKSUM: c_int = 1;
const ZSTD_RMD_REF_SINGLE_DDICT: c_int = 0;
const ZSTD_RMD_REF_MULTIPLE_DDICTS: c_int = 1;
const ZSTD_DLM_BY_COPY: c_int = 0;
const ZSTD_DLM_BY_REF: c_int = 1;
const ZSTD_DCT_AUTO: c_int = 0;
const ZSTD_DCT_RAW_CONTENT: c_int = 1;
const ZSTD_USE_INDEFINITELY: c_int = -1;
const ZSTD_DONT_USE: c_int = 0;
const ZSTD_USE_ONCE: c_int = 1;
const ZSTDDS_GET_FRAME_HEADER_SIZE: c_int = 0;
const ZSTDDS_DECODE_FRAME_HEADER: c_int = 1;
const ZSTDDS_DECODE_BLOCK_HEADER: c_int = 2;
const ZSTDDS_DECOMPRESS_BLOCK: c_int = 3;
const ZSTDDS_DECOMPRESS_LAST_BLOCK: c_int = 4;
const ZSTDDS_CHECK_CHECKSUM: c_int = 5;
const ZSTDDS_DECODE_SKIPPABLE_HEADER: c_int = 6;
const ZSTDDS_SKIP_FRAME: c_int = 7;
const ZDSS_INIT: c_int = 0;
const ZDSS_LOAD_HEADER: c_int = 1;
const ZDSS_READ: c_int = 2;
const ZDSS_LOAD: c_int = 3;
const ZDSS_FLUSH: c_int = 4;
const BT_RAW: c_int = 0;
const BT_RLE: c_int = 1;
const BT_COMPRESSED: c_int = 2;
const BT_RESERVED: c_int = 3;
const ZSTD_D_WINDOW_LOG_MAX: c_int = 100;
const ZSTD_D_FORMAT: c_int = 1000;
const ZSTD_D_STABLE_OUT_BUFFER: c_int = 1001;
const ZSTD_D_FORCE_IGNORE_CHECKSUM: c_int = 1002;
const ZSTD_D_REF_MULTIPLE_DDICTS: c_int = 1003;
const ZSTD_D_DISABLE_HUFFMAN_ASSEMBLY: c_int = 1004;
const ZSTD_D_MAX_BLOCK_SIZE: c_int = 1005;
const ZSTD_RESET_SESSION_ONLY: c_int = 1;
const ZSTD_RESET_PARAMETERS: c_int = 2;
const ZSTD_RESET_SESSION_AND_PARAMETERS: c_int = 3;
const ZSTD_NIT_FRAME_HEADER: c_int = 0;
const ZSTD_NIT_BLOCK_HEADER: c_int = 1;
const ZSTD_NIT_BLOCK: c_int = 2;
const ZSTD_NIT_LAST_BLOCK: c_int = 3;
const ZSTD_NIT_CHECKSUM: c_int = 4;
const ZSTD_NIT_SKIPPABLE_FRAME: c_int = 5;
const LL_BASE: [u32; MAX_LL + 1] = [
0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 18, 20, 22, 24, 28, 32, 40, 48, 64,
0x80, 0x100, 0x200, 0x400, 0x800, 0x1000, 0x2000, 0x4000, 0x8000, 0x10000,
];
const OF_BASE: [u32; MAX_OFF + 1] = [
0, 1, 1, 5, 0xD, 0x1D, 0x3D, 0x7D, 0xFD, 0x1FD, 0x3FD, 0x7FD, 0xFFD, 0x1FFD, 0x3FFD, 0x7FFD,
0xFFFD, 0x1FFFD, 0x3FFFD, 0x7FFFD, 0xFFFFD, 0x1FFFFD, 0x3FFFFD, 0x7FFFFD, 0xFFFFFD, 0x1FFFFFD,
0x3FFFFFD, 0x7FFFFFD, 0xFFFFFFD, 0x1FFFFFFD, 0x3FFFFFFD, 0x7FFFFFFD,
];
const OF_BITS: [u8; MAX_OFF + 1] = [
0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25,
26, 27, 28, 29, 30, 31,
];
const ML_BASE: [u32; MAX_ML + 1] = [
3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27,
28, 29, 30, 31, 32, 33, 34, 35, 37, 39, 41, 43, 47, 51, 59, 67, 83, 99, 0x83, 0x103, 0x203,
0x403, 0x803, 0x1003, 0x2003, 0x4003, 0x8003, 0x10003,
];
const LL_BITS: [u8; MAX_LL + 1] = [
0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 2, 2, 3, 3, 4, 6, 7, 8, 9, 10, 11,
12, 13, 14, 15, 16,
];
const ML_BITS: [u8; MAX_ML + 1] = [
0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
1, 1, 1, 1, 2, 2, 3, 3, 4, 4, 5, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16,
];
#[repr(C)]
pub struct ZSTD_DCtx {
_private: [u8; 0],
}
pub type ZSTD_DStream = ZSTD_DCtx;
#[repr(C)]
#[derive(Clone, Copy, Debug, Default)]
pub struct ZSTD_FrameHeader {
pub frame_content_size: u64,
pub window_size: u64,
pub block_size_max: c_uint,
pub frame_type: c_int,
pub header_size: c_uint,
pub dict_id: c_uint,
pub checksum_flag: c_uint,
pub reserved1: c_uint,
pub reserved2: c_uint,
}
#[repr(C)]
#[derive(Clone, Copy)]
pub struct ZSTD_inBuffer {
pub src: *const c_void,
pub size: usize,
pub pos: usize,
}
#[repr(C)]
#[derive(Clone, Copy)]
pub struct ZSTD_outBuffer {
pub dst: *mut c_void,
pub size: usize,
pub pos: usize,
}
#[repr(C)]
#[derive(Clone, Copy, Debug, Default)]
pub struct ZSTD_bounds {
pub error: usize,
pub lower_bound: c_int,
pub upper_bound: c_int,
}
type ZstdAllocFunction = unsafe extern "C" fn(*mut c_void, usize) -> *mut c_void;
type ZstdFreeFunction = unsafe extern "C" fn(*mut c_void, *mut c_void);
#[repr(C)]
#[derive(Clone, Copy)]
pub struct ZSTD_customMem {
custom_alloc: Option<ZstdAllocFunction>,
custom_free: Option<ZstdFreeFunction>,
opaque: *mut c_void,
}
#[repr(C)]
#[derive(Clone, Copy, Default)]
struct BlockProperties {
block_type: c_int,
last_block: u32,
orig_size: u32,
}
#[repr(C)]
pub struct ZSTD_entropyDTables_t {
ll_table: [crate::zstd_decompress_block::ZSTD_seqSymbol; 1 + (1 << LL_FSE_LOG)],
of_table: [crate::zstd_decompress_block::ZSTD_seqSymbol; 1 + (1 << OFF_FSE_LOG)],
ml_table: [crate::zstd_decompress_block::ZSTD_seqSymbol; 1 + (1 << ML_FSE_LOG)],
huf_table: [u32; HUF_DTABLE_SIZE],
rep: [u32; ZSTD_REP_NUM],
workspace: [u32; ZSTD_BUILD_FSE_TABLE_WKSP_SIZE_U32],
}
/// C-provided leaves of `ZSTD_DCtx_s`. Every pointer is produced under the
/// active C preprocessor configuration; Rust never hard-codes a private
/// decoder-context offset.
#[repr(C)]
#[derive(Clone, Copy)]
struct ZSTD_rustDctxView {
dctx: *mut c_void,
llt_ptr: *mut c_void,
mlt_ptr: *mut c_void,
oft_ptr: *mut c_void,
huf_ptr: *mut c_void,
entropy: *mut c_void,
workspace: *mut c_void,
workspace_size: usize,
previous_dst_end: *mut c_void,
prefix_start: *mut c_void,
virtual_start: *mut c_void,
dict_end: *mut c_void,
expected: *mut c_void,
f_params: *mut c_void,
processed_c_size: *mut c_void,
decoded_size: *mut c_void,
b_type: *mut c_void,
stage: *mut c_void,
lit_entropy: *mut c_void,
fse_entropy: *mut c_void,
xxh_state: *mut c_void,
header_size: *mut c_void,
format: *mut c_void,
force_ignore_checksum: *mut c_void,
validate_checksum: *mut c_void,
lit_ptr: *mut c_void,
custom_mem: *mut c_void,
lit_size: *mut c_void,
rle_size: *mut c_void,
static_size: *mut c_void,
is_frame_decompression: *mut c_void,
ddict_local: *mut c_void,
ddict: *mut c_void,
dict_id: *mut c_void,
ddict_is_cold: *mut c_void,
dict_uses: *mut c_void,
ddict_set: *mut c_void,
ref_multiple_ddicts: *mut c_void,
disable_huf_asm: *mut c_void,
max_block_size_param: *mut c_void,
stream_stage: *mut c_void,
in_buff: *mut c_void,
in_buff_size: *mut c_void,
in_pos: *mut c_void,
max_window_size: *mut c_void,
out_buff: *mut c_void,
out_buff_size: *mut c_void,
out_start: *mut c_void,
out_end: *mut c_void,
lh_size: *mut c_void,
legacy_context: *mut c_void,
previous_legacy_version: *mut c_void,
legacy_version: *mut c_void,
hostage_byte: *mut c_void,
no_forward_progress: *mut c_void,
out_buffer_mode: *mut c_void,
expected_out_buffer: *mut c_void,
lit_buffer: *mut c_void,
lit_buffer_end: *mut c_void,
lit_buffer_location: *mut c_void,
lit_extra_buffer: *mut c_void,
lit_extra_buffer_size: usize,
header_buffer: *mut c_void,
header_buffer_size: usize,
oversized_duration: *mut c_void,
fuzz_begin: *mut c_void,
fuzz_end: *mut c_void,
dctx_size: usize,
}
unsafe extern "C" {
fn ZSTD_rust_dctx_view(dctx: *mut ZSTD_DCtx, out: *mut ZSTD_rustDctxView);
fn ZSTD_rust_dctx_sizeof() -> usize;
fn ZSTD_rust_dctx_alloc(custom_mem: ZSTD_customMem) -> *mut ZSTD_DCtx;
fn ZSTD_rust_dctx_free_storage(dctx: *mut ZSTD_DCtx, custom_mem: ZSTD_customMem);
fn ZSTD_rust_dctx_init_platform(dctx: *mut ZSTD_DCtx);
fn ZSTD_rust_dctx_default_max_window_size() -> usize;
fn ZSTD_rust_no_forward_progress_max() -> c_int;
fn ZSTD_rust_heapmode() -> c_int;
fn ZSTD_rust_decompress_stack(
dst: *mut c_void,
dst_capacity: usize,
src: *const c_void,
src_size: usize,
) -> usize;
fn ZSTD_rust_custom_malloc(size: usize, custom_mem: ZSTD_customMem) -> *mut c_void;
fn ZSTD_rust_custom_calloc(size: usize, custom_mem: ZSTD_customMem) -> *mut c_void;
fn ZSTD_rust_custom_free(allocation: *mut c_void, custom_mem: ZSTD_customMem);
fn ZSTD_rust_create_ddict(
dict: *const c_void,
dict_size: usize,
dict_load_method: c_int,
dict_content_type: c_int,
custom_mem: ZSTD_customMem,
) -> *mut ZSTD_DDict;
fn ZSTD_rust_dctx_trace_end(
dctx: *mut ZSTD_DCtx,
uncompressed_size: u64,
compressed_size: u64,
streaming: c_int,
);
fn ZSTD_rust_dctx_trace_begin(dctx: *mut ZSTD_DCtx);
fn ZSTD_rust_dctx_copy_prefix(dst: *mut ZSTD_DCtx, src: *const ZSTD_DCtx);
fn ZSTD_rust_legacy_is(src: *const c_void, src_size: usize) -> c_uint;
fn ZSTD_rust_legacy_get_decompressed_size(src: *const c_void, src_size: usize) -> u64;
fn ZSTD_rust_legacy_find_compressed_size(src: *const c_void, src_size: usize) -> usize;
fn ZSTD_rust_legacy_frame_size_info(
src: *const c_void,
src_size: usize,
compressed_size: *mut usize,
decompressed_bound: *mut u64,
nb_blocks: *mut usize,
) -> usize;
fn ZSTD_rust_legacy_decompress(
dst: *mut c_void,
dst_capacity: usize,
src: *const c_void,
src_size: usize,
dict: *const c_void,
dict_size: usize,
) -> usize;
fn ZSTD_rust_legacy_decompress_stream(
dctx: *mut ZSTD_DCtx,
output: *mut ZSTD_outBuffer,
input: *mut ZSTD_inBuffer,
dict: *const c_void,
dict_size: usize,
) -> usize;
fn ZSTD_rust_legacy_free_stream(dctx: *mut ZSTD_DCtx);
fn ZSTD_decompressBlock_internal(
dctx: *mut ZSTD_DCtx,
dst: *mut c_void,
dst_capacity: usize,
src: *const c_void,
src_size: usize,
streaming: c_int,
) -> usize;
fn ZSTD_checkContinuity(dctx: *mut ZSTD_DCtx, dst: *const c_void, dst_size: usize);
}
#[inline]
unsafe fn field<T: Copy>(slot: *mut c_void) -> T {
unsafe { slot.cast::<T>().read() }
}
#[inline]
unsafe fn set_field<T>(slot: *mut c_void, value: T) {
unsafe { slot.cast::<T>().write(value) }
}
unsafe fn dctx_view(dctx: *mut ZSTD_DCtx) -> ZSTD_rustDctxView {
let mut view = MaybeUninit::<ZSTD_rustDctxView>::zeroed();
unsafe { ZSTD_rust_dctx_view(dctx, view.as_mut_ptr()) };
unsafe { view.assume_init() }
}
#[inline]
fn frame_header_prefix(format: c_int) -> usize {
if format == ZSTD_F_ZSTD1 {
5
} else {
1
}
}
#[inline]
fn frame_header_min(format: c_int) -> usize {
if format == ZSTD_F_ZSTD1 {
6
} else {
2
}
}
#[inline]
fn window_log_max() -> usize {
if MEM_32bits() {
ZSTD_WINDOWLOG_MAX_32
} else {
ZSTD_WINDOWLOG_MAX_64
}
}
#[inline]
unsafe fn const_ptr_add(ptr: *const u8, amount: usize) -> *const u8 {
if ptr.is_null() {
debug_assert_eq!(amount, 0);
ptr
} else {
unsafe { ptr.add(amount) }
}
}
#[inline]
unsafe fn ptr_distance(end: *const u8, start: *const u8) -> usize {
(end as usize).wrapping_sub(start as usize)
}
#[inline]
unsafe fn copy_bytes(dst: *mut u8, src: *const u8, len: usize) {
if len != 0 {
unsafe { ptr::copy(src, dst, len) };
}
}
#[inline]
unsafe fn limit_copy(dst: *mut u8, dst_capacity: usize, src: *const u8, src_size: usize) -> usize {
let len = min(dst_capacity, src_size);
if len != 0 {
unsafe { ptr::copy_nonoverlapping(src, dst, len) };
}
len
}
#[inline]
fn default_custom_mem() -> ZSTD_customMem {
ZSTD_customMem {
custom_alloc: None,
custom_free: None,
opaque: ptr::null_mut(),
}
}
#[inline]
fn custom_mem_valid(custom_mem: ZSTD_customMem) -> bool {
custom_mem.custom_alloc.is_some() == custom_mem.custom_free.is_some()
}
#[inline]
unsafe fn get_frame_header_ptr(view: &ZSTD_rustDctxView) -> *mut ZSTD_FrameHeader {
view.f_params.cast()
}
#[inline]
unsafe fn entropy_ptr(view: &ZSTD_rustDctxView) -> *mut ZSTD_entropyDTables_t {
view.entropy.cast()
}
#[inline]
unsafe fn dctx_custom_mem(view: &ZSTD_rustDctxView) -> ZSTD_customMem {
unsafe { field(view.custom_mem) }
}
#[inline]
unsafe fn dctx_ddict(view: &ZSTD_rustDctxView) -> *const ZSTD_DDict {
unsafe { field(view.ddict) }
}
#[inline]
unsafe fn set_dctx_ddict(view: &ZSTD_rustDctxView, ddict: *const ZSTD_DDict) {
unsafe { set_field(view.ddict, ddict) }
}
#[inline]
unsafe fn dctx_ddict_local(view: &ZSTD_rustDctxView) -> *mut ZSTD_DDict {
unsafe { field(view.ddict_local) }
}
#[inline]
unsafe fn set_dctx_ddict_local(view: &ZSTD_rustDctxView, ddict: *mut ZSTD_DDict) {
unsafe { set_field(view.ddict_local, ddict) }
}
#[inline]
unsafe fn get_pointer(slot: *mut c_void) -> *const u8 {
unsafe { field(slot) }
}
#[inline]
unsafe fn set_pointer(slot: *mut c_void, value: *const u8) {
unsafe { set_field(slot, value) }
}
#[inline]
unsafe fn get_mut_pointer(slot: *mut c_void) -> *mut u8 {
unsafe { field(slot) }
}
#[inline]
unsafe fn set_mut_pointer(slot: *mut c_void, value: *mut u8) {
unsafe { set_field(slot, value) }
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_isFrame(buffer: *const c_void, size: usize) -> c_uint {
if size < ZSTD_FRAMEIDSIZE || buffer.is_null() {
return 0;
}
let magic = unsafe { MEM_readLE32(buffer) };
if magic == ZSTD_MAGICNUMBER
|| (magic & ZSTD_MAGIC_SKIPPABLE_MASK) == ZSTD_MAGIC_SKIPPABLE_START
{
return 1;
}
unsafe { ZSTD_rust_legacy_is(buffer, size) }
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_isSkippableFrame(buffer: *const c_void, size: usize) -> c_uint {
if size < ZSTD_FRAMEIDSIZE || buffer.is_null() {
return 0;
}
u32::from(
unsafe { MEM_readLE32(buffer) } & ZSTD_MAGIC_SKIPPABLE_MASK == ZSTD_MAGIC_SKIPPABLE_START,
)
}
unsafe fn frame_header_size_internal(src: *const u8, src_size: usize, format: c_int) -> usize {
let min_input_size = frame_header_prefix(format);
if src_size < min_input_size {
return ERROR(ZstdErrorCode::SrcSizeWrong);
}
let fhd = unsafe { *src.add(min_input_size - 1) };
let dict_id = fhd & 3;
let single_segment = (fhd >> 5) & 1;
let fcs_id = fhd >> 6;
let did_size = [0usize, 1, 2, 4][dict_id as usize];
let fcs_size = [0usize, 2, 4, 8][fcs_id as usize];
min_input_size
+ usize::from(single_segment == 0)
+ did_size
+ fcs_size
+ usize::from(single_segment != 0 && fcs_id == 0)
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_frameHeaderSize(src: *const c_void, src_size: usize) -> usize {
unsafe { frame_header_size_internal(src.cast(), src_size, ZSTD_F_ZSTD1) }
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_getFrameHeader_advanced(
zfh: *mut ZSTD_FrameHeader,
src: *const c_void,
src_size: usize,
format: c_int,
) -> usize {
let min_input_size = frame_header_prefix(format);
if src_size != 0 && src.is_null() {
return ERROR(ZstdErrorCode::Generic);
}
if src_size < min_input_size {
if src_size != 0 && format != ZSTD_F_ZSTD1_MAGICLESS {
let mut header = ZSTD_MAGICNUMBER.to_le_bytes();
unsafe {
ptr::copy_nonoverlapping(src.cast::<u8>(), header.as_mut_ptr(), min(4, src_size))
};
if u32::from_le_bytes(header) != ZSTD_MAGICNUMBER {
header = ZSTD_MAGIC_SKIPPABLE_START.to_le_bytes();
unsafe {
ptr::copy_nonoverlapping(
src.cast::<u8>(),
header.as_mut_ptr(),
min(4, src_size),
)
};
if u32::from_le_bytes(header) & ZSTD_MAGIC_SKIPPABLE_MASK
!= ZSTD_MAGIC_SKIPPABLE_START
{
return ERROR(ZstdErrorCode::PrefixUnknown);
}
}
}
return min_input_size;
}
if zfh.is_null() {
return ERROR(ZstdErrorCode::Generic);
}
unsafe { zfh.write(ZSTD_FrameHeader::default()) };
let ip = src.cast::<u8>();
if format != ZSTD_F_ZSTD1_MAGICLESS && unsafe { MEM_readLE32(src) } != ZSTD_MAGICNUMBER {
let magic = unsafe { MEM_readLE32(src) };
if magic & ZSTD_MAGIC_SKIPPABLE_MASK == ZSTD_MAGIC_SKIPPABLE_START {
if src_size < ZSTD_SKIPPABLEHEADERSIZE {
return ZSTD_SKIPPABLEHEADERSIZE;
}
unsafe {
(*zfh).frame_type = ZSTD_SKIPPABLE_FRAME;
(*zfh).dict_id = magic - ZSTD_MAGIC_SKIPPABLE_START;
(*zfh).header_size = ZSTD_SKIPPABLEHEADERSIZE as c_uint;
(*zfh).frame_content_size = MEM_readLE32(ip.add(ZSTD_FRAMEIDSIZE).cast()) as u64;
}
return 0;
}
return ERROR(ZstdErrorCode::PrefixUnknown);
}
let fh_size = unsafe { frame_header_size_internal(ip, src_size, format) };
if ERR_isError(fh_size) {
return fh_size;
}
if src_size < fh_size {
return fh_size;
}
unsafe { (*zfh).header_size = fh_size as c_uint };
let fhd = unsafe { *ip.add(min_input_size - 1) };
if fhd & 0x08 != 0 {
return ERROR(ZstdErrorCode::FrameParameterUnsupported);
}
let dict_id_size_code = fhd & 3;
let checksum_flag = (fhd >> 2) & 1;
let single_segment = (fhd >> 5) & 1;
let fcs_id = fhd >> 6;
let mut pos = min_input_size;
let mut window_size = 0u64;
let mut dict_id = 0u32;
let mut frame_content_size = ZSTD_CONTENTSIZE_UNKNOWN;
if single_segment == 0 {
let wl = unsafe { *ip.add(pos) };
pos += 1;
let window_log = usize::from(wl >> 3) + ZSTD_WINDOWLOG_ABSOLUTEMIN;
if window_log > window_log_max() {
return ERROR(ZstdErrorCode::FrameParameterWindowTooLarge);
}
window_size = 1u64 << window_log;
window_size = window_size.wrapping_add((window_size >> 3) * u64::from(wl & 7));
}
match dict_id_size_code {
0 => {}
1 => {
dict_id = unsafe { *ip.add(pos) } as u32;
pos += 1;
}
2 => {
dict_id = unsafe { MEM_readLE16(ip.add(pos).cast()) as u32 };
pos += 2;
}
_ => {
dict_id = unsafe { MEM_readLE32(ip.add(pos).cast()) };
pos += 4;
}
}
match fcs_id {
0 => {
if single_segment != 0 {
frame_content_size = unsafe { *ip.add(pos) } as u64;
}
}
1 => frame_content_size = unsafe { MEM_readLE16(ip.add(pos).cast()) as u64 + 256 },
2 => frame_content_size = unsafe { MEM_readLE32(ip.add(pos).cast()) as u64 },
_ => frame_content_size = unsafe { MEM_readLE64(ip.add(pos).cast()) },
}
if single_segment != 0 {
window_size = frame_content_size;
}
unsafe {
(*zfh).frame_type = ZSTD_FRAME;
(*zfh).frame_content_size = frame_content_size;
(*zfh).window_size = window_size;
(*zfh).block_size_max = min(window_size, ZSTD_BLOCKSIZE_MAX as u64) as c_uint;
(*zfh).dict_id = dict_id;
(*zfh).checksum_flag = checksum_flag as c_uint;
}
0
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_getFrameHeader(
zfh: *mut ZSTD_FrameHeader,
src: *const c_void,
src_size: usize,
) -> usize {
unsafe { ZSTD_getFrameHeader_advanced(zfh, src, src_size, ZSTD_F_ZSTD1) }
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_getFrameContentSize(src: *const c_void, src_size: usize) -> u64 {
if unsafe { ZSTD_rust_legacy_is(src, src_size) } != 0 {
let size = unsafe { ZSTD_rust_legacy_get_decompressed_size(src, src_size) };
return if size == 0 {
ZSTD_CONTENTSIZE_UNKNOWN
} else {
size
};
}
let mut zfh = ZSTD_FrameHeader::default();
if unsafe { ZSTD_getFrameHeader(&mut zfh, src, src_size) } != 0 {
return ZSTD_CONTENTSIZE_ERROR;
}
if zfh.frame_type == ZSTD_SKIPPABLE_FRAME {
0
} else {
zfh.frame_content_size
}
}
unsafe fn read_skippable_frame_size(src: *const u8, src_size: usize) -> usize {
if src_size < ZSTD_SKIPPABLEHEADERSIZE {
return ERROR(ZstdErrorCode::SrcSizeWrong);
}
let size = unsafe { MEM_readLE32(src.add(ZSTD_FRAMEIDSIZE).cast()) };
if size.wrapping_add(ZSTD_SKIPPABLEHEADERSIZE as u32) < size {
return ERROR(ZstdErrorCode::FrameParameterUnsupported);
}
let frame_size = ZSTD_SKIPPABLEHEADERSIZE + size as usize;
if frame_size > src_size {
return ERROR(ZstdErrorCode::SrcSizeWrong);
}
frame_size
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_readSkippableFrame(
dst: *mut c_void,
dst_capacity: usize,
magic_variant: *mut c_uint,
src: *const c_void,
src_size: usize,
) -> usize {
if src_size < ZSTD_SKIPPABLEHEADERSIZE {
return ERROR(ZstdErrorCode::SrcSizeWrong);
}
let src_u8 = src.cast::<u8>();
let magic = unsafe { MEM_readLE32(src) };
let frame_size = unsafe { read_skippable_frame_size(src_u8, src_size) };
if ERR_isError(frame_size) {
return frame_size;
}
let content_size = frame_size - ZSTD_SKIPPABLEHEADERSIZE;
if unsafe { ZSTD_isSkippableFrame(src, src_size) } == 0 {
return ERROR(ZstdErrorCode::FrameParameterUnsupported);
}
if content_size > dst_capacity {
return ERROR(ZstdErrorCode::DstSizeTooSmall);
}
if content_size != 0 && !dst.is_null() {
unsafe {
copy_bytes(
dst.cast(),
src_u8.add(ZSTD_SKIPPABLEHEADERSIZE),
content_size,
)
};
}
if !magic_variant.is_null() {
unsafe { magic_variant.write(magic - ZSTD_MAGIC_SKIPPABLE_START) };
}
content_size
}
#[repr(C)]
#[derive(Clone, Copy, Default)]
struct FrameSizeInfo {
nb_blocks: usize,
compressed_size: usize,
decompressed_bound: u64,
}
#[inline]
fn error_frame_size_info(error: usize) -> FrameSizeInfo {
FrameSizeInfo {
nb_blocks: 0,
compressed_size: error,
decompressed_bound: ZSTD_CONTENTSIZE_ERROR,
}
}
unsafe fn find_frame_size_info(src: *const u8, src_size: usize, format: c_int) -> FrameSizeInfo {
if format == ZSTD_F_ZSTD1 && unsafe { ZSTD_rust_legacy_is(src.cast(), src_size) } != 0 {
let mut info = FrameSizeInfo::default();
let status = unsafe {
ZSTD_rust_legacy_frame_size_info(
src.cast(),
src_size,
&mut info.compressed_size,
&mut info.decompressed_bound,
&mut info.nb_blocks,
)
};
return if ERR_isError(status) {
error_frame_size_info(status)
} else {
info
};
}
if format == ZSTD_F_ZSTD1
&& src_size >= ZSTD_SKIPPABLEHEADERSIZE
&& unsafe { MEM_readLE32(src.cast()) } & ZSTD_MAGIC_SKIPPABLE_MASK
== ZSTD_MAGIC_SKIPPABLE_START
{
return FrameSizeInfo {
nb_blocks: 0,
compressed_size: unsafe { read_skippable_frame_size(src, src_size) },
decompressed_bound: 0,
};
}
let mut zfh = ZSTD_FrameHeader::default();
let header_result =
unsafe { ZSTD_getFrameHeader_advanced(&mut zfh, src.cast(), src_size, format) };
if ERR_isError(header_result) {
return error_frame_size_info(header_result);
}
if header_result != 0 {
return error_frame_size_info(ERROR(ZstdErrorCode::SrcSizeWrong));
}
let mut ip = unsafe { src.add(zfh.header_size as usize) };
let mut remaining = src_size - zfh.header_size as usize;
let mut nb_blocks = 0usize;
loop {
let mut block = BlockProperties::default();
let c_block_size = unsafe {
crate::zstd_decompress_block::ZSTD_getcBlockSize(
ip.cast(),
remaining,
(&mut block as *mut BlockProperties).cast(),
)
};
if ERR_isError(c_block_size) {
return error_frame_size_info(c_block_size);
}
let total_size = match ZSTD_BLOCKHEADERSIZE.checked_add(c_block_size) {
Some(size) => size,
None => return error_frame_size_info(ERROR(ZstdErrorCode::SrcSizeWrong)),
};
if total_size > remaining {
return error_frame_size_info(ERROR(ZstdErrorCode::SrcSizeWrong));
}
ip = unsafe { ip.add(total_size) };
remaining -= total_size;
nb_blocks += 1;
if block.last_block != 0 {
break;
}
}
if zfh.checksum_flag != 0 {
if remaining < 4 {
return error_frame_size_info(ERROR(ZstdErrorCode::SrcSizeWrong));
}
ip = unsafe { ip.add(4) };
}
FrameSizeInfo {
nb_blocks,
compressed_size: unsafe { ip.offset_from(src) as usize },
decompressed_bound: if zfh.frame_content_size != ZSTD_CONTENTSIZE_UNKNOWN {
zfh.frame_content_size
} else {
(nb_blocks as u64).wrapping_mul(zfh.block_size_max as u64)
},
}
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_findFrameCompressedSize(
src: *const c_void,
src_size: usize,
) -> usize {
unsafe { find_frame_size_info(src.cast(), src_size, ZSTD_F_ZSTD1).compressed_size }
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_findDecompressedSize(src: *const c_void, mut src_size: usize) -> u64 {
let mut input = src.cast::<u8>();
let mut total = 0u64;
while src_size >= frame_header_prefix(ZSTD_F_ZSTD1) {
if unsafe { MEM_readLE32(input.cast()) } & ZSTD_MAGIC_SKIPPABLE_MASK
== ZSTD_MAGIC_SKIPPABLE_START
{
let size = unsafe { read_skippable_frame_size(input, src_size) };
if ERR_isError(size) {
return ZSTD_CONTENTSIZE_ERROR;
}
input = unsafe { input.add(size) };
src_size -= size;
continue;
}
let frame_size = unsafe { ZSTD_getFrameContentSize(input.cast(), src_size) };
if frame_size >= ZSTD_CONTENTSIZE_ERROR {
return frame_size;
}
let next = total.wrapping_add(frame_size);
if next < total {
return ZSTD_CONTENTSIZE_ERROR;
}
total = next;
let compressed = unsafe { ZSTD_findFrameCompressedSize(input.cast(), src_size) };
if ERR_isError(compressed) || compressed > src_size {
return ZSTD_CONTENTSIZE_ERROR;
}
input = unsafe { input.add(compressed) };
src_size -= compressed;
}
if src_size != 0 {
ZSTD_CONTENTSIZE_ERROR
} else {
total
}
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_getDecompressedSize(src: *const c_void, src_size: usize) -> u64 {
let result = unsafe { ZSTD_getFrameContentSize(src, src_size) };
if result >= ZSTD_CONTENTSIZE_ERROR {
0
} else {
result
}
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_decompressBound(src: *const c_void, mut src_size: usize) -> u64 {
let mut input = src.cast::<u8>();
let mut bound = 0u64;
while src_size != 0 {
let info = unsafe { find_frame_size_info(input, src_size, ZSTD_F_ZSTD1) };
if ERR_isError(info.compressed_size) || info.decompressed_bound == ZSTD_CONTENTSIZE_ERROR {
return ZSTD_CONTENTSIZE_ERROR;
}
if info.compressed_size > src_size {
return ZSTD_CONTENTSIZE_ERROR;
}
bound = bound.wrapping_add(info.decompressed_bound);
input = unsafe { input.add(info.compressed_size) };
src_size -= info.compressed_size;
}
bound
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_decompressionMargin(
src: *const c_void,
mut src_size: usize,
) -> usize {
let mut input = src.cast::<u8>();
let mut margin = 0usize;
let mut max_block_size = 0usize;
while src_size != 0 {
let info = unsafe { find_frame_size_info(input, src_size, ZSTD_F_ZSTD1) };
let mut zfh = ZSTD_FrameHeader::default();
let header = unsafe { ZSTD_getFrameHeader(&mut zfh, input.cast(), src_size) };
if ERR_isError(header) {
return header;
}
if header != 0
|| ERR_isError(info.compressed_size)
|| info.decompressed_bound == ZSTD_CONTENTSIZE_ERROR
{
return ERROR(ZstdErrorCode::CorruptionDetected);
}
if zfh.frame_type == ZSTD_FRAME {
margin = margin
.wrapping_add(zfh.header_size as usize)
.wrapping_add(if zfh.checksum_flag != 0 { 4 } else { 0 })
.wrapping_add(ZSTD_BLOCKHEADERSIZE.wrapping_mul(info.nb_blocks));
max_block_size = max(max_block_size, zfh.block_size_max as usize);
} else {
margin = margin.wrapping_add(info.compressed_size);
}
if info.compressed_size > src_size {
return ERROR(ZstdErrorCode::CorruptionDetected);
}
input = unsafe { input.add(info.compressed_size) };
src_size -= info.compressed_size;
}
margin.wrapping_add(max_block_size)
}
unsafe fn ref_dict_content(view: &ZSTD_rustDctxView, dict: *const u8, dict_size: usize) -> usize {
let previous = unsafe { get_pointer(view.previous_dst_end) };
let prefix = unsafe { get_pointer(view.prefix_start) };
unsafe {
set_pointer(view.dict_end, previous);
/* Do the same address arithmetic as the C implementation without
* forming a Rust pointer outside of an allocation. These are virtual
* history addresses and are only compared/subtracted by the block
* decoder, never dereferenced until they again name live history. */
let history = (previous as usize).wrapping_sub(prefix as usize);
set_pointer(
view.virtual_start,
(dict as usize).wrapping_sub(history) as *const u8,
);
set_pointer(view.prefix_start, dict);
set_pointer(view.previous_dst_end, const_ptr_add(dict, dict_size));
if !view.fuzz_begin.is_null() {
set_pointer(view.fuzz_begin, dict);
set_pointer(view.fuzz_end, const_ptr_add(dict, dict_size));
}
}
0
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_loadDEntropy(
entropy: *mut ZSTD_entropyDTables_t,
dict: *const c_void,
dict_size: usize,
) -> usize {
if dict_size <= 8 || dict.is_null() {
return ERROR(ZstdErrorCode::DictionaryCorrupted);
}
let mut dict_ptr = unsafe { dict.cast::<u8>().add(8) };
let dict_end = unsafe { dict.cast::<u8>().add(dict_size) };
let workspace = entropy.cast::<c_void>();
let workspace_size = size_of::<crate::zstd_decompress_block::ZSTD_seqSymbol>()
* ((1 << LL_FSE_LOG) + (1 << OFF_FSE_LOG) + (1 << ML_FSE_LOG) + 3);
let huf_size = unsafe {
#[cfg(feature = "huf-force-decompress-x1")]
{
HUF_readDTableX1_wksp(
(*entropy).huf_table.as_mut_ptr(),
dict_ptr.cast(),
ptr_distance(dict_end, dict_ptr),
workspace,
workspace_size,
0,
)
}
#[cfg(not(feature = "huf-force-decompress-x1"))]
{
HUF_readDTableX2_wksp(
(*entropy).huf_table.as_mut_ptr(),
dict_ptr.cast(),
ptr_distance(dict_end, dict_ptr),
workspace,
workspace_size,
0,
)
}
};
if ERR_isError(huf_size) {
return ERROR(ZstdErrorCode::DictionaryCorrupted);
}
dict_ptr = unsafe { dict_ptr.add(huf_size) };
unsafe fn load_table(
table: *mut crate::zstd_decompress_block::ZSTD_seqSymbol,
max_symbol: usize,
max_log: usize,
base: *const u32,
bits: *const u8,
workspace: *mut u32,
dict_ptr: *const u8,
dict_end: *const u8,
) -> Result<usize, usize> {
let mut norm = [0i16; MAX_ML + 1];
let mut max = max_symbol as c_uint;
let mut log = 0u32;
let size = unsafe {
FSE_readNCount(
norm.as_mut_ptr(),
&mut max,
&mut log,
dict_ptr.cast(),
ptr_distance(dict_end, dict_ptr),
)
};
if ERR_isError(size) || max as usize > max_symbol || log as usize > max_log {
return Err(ERROR(ZstdErrorCode::DictionaryCorrupted));
}
unsafe {
crate::zstd_decompress_block::ZSTD_buildFSETable(
table,
norm.as_ptr(),
max,
base,
bits,
log,
workspace.cast(),
ZSTD_BUILD_FSE_TABLE_WKSP_SIZE_U32 * size_of::<u32>(),
0,
);
}
Ok(size)
}
let entropy_ref = unsafe { &mut *entropy };
let off_size = match unsafe {
load_table(
entropy_ref.of_table.as_mut_ptr(),
MAX_OFF,
OFF_FSE_LOG,
OF_BASE.as_ptr(),
OF_BITS.as_ptr(),
entropy_ref.workspace.as_mut_ptr(),
dict_ptr,
dict_end,
)
} {
Ok(size) => size,
Err(error) => return error,
};
dict_ptr = unsafe { dict_ptr.add(off_size) };
let ml_size = match unsafe {
load_table(
entropy_ref.ml_table.as_mut_ptr(),
MAX_ML,
ML_FSE_LOG,
ML_BASE.as_ptr(),
ML_BITS.as_ptr(),
entropy_ref.workspace.as_mut_ptr(),
dict_ptr,
dict_end,
)
} {
Ok(size) => size,
Err(error) => return error,
};
dict_ptr = unsafe { dict_ptr.add(ml_size) };
let ll_size = match unsafe {
load_table(
entropy_ref.ll_table.as_mut_ptr(),
MAX_LL,
LL_FSE_LOG,
LL_BASE.as_ptr(),
LL_BITS.as_ptr(),
entropy_ref.workspace.as_mut_ptr(),
dict_ptr,
dict_end,
)
} {
Ok(size) => size,
Err(error) => return error,
};
dict_ptr = unsafe { dict_ptr.add(ll_size) };
if unsafe { ptr_distance(dict_end, dict_ptr) } < 12 {
return ERROR(ZstdErrorCode::DictionaryCorrupted);
}
let content_size = unsafe { ptr_distance(dict_end, dict_ptr.add(12)) };
for rep in &mut entropy_ref.rep {
let value = unsafe { MEM_readLE32(dict_ptr.cast()) };
dict_ptr = unsafe { dict_ptr.add(4) };
if value == 0 || value as usize > content_size {
return ERROR(ZstdErrorCode::DictionaryCorrupted);
}
*rep = value;
}
unsafe { dict_ptr.offset_from(dict.cast()) as usize }
}
#[repr(C)]
struct DDictHashSet {
table: *mut *const ZSTD_DDict,
size: usize,
count: usize,
}
unsafe fn ddict_hash_index(set: *const DDictHashSet, dict_id: u32) -> usize {
let hash = unsafe { XXH64((&dict_id as *const u32).cast(), size_of::<u32>(), 0) };
hash as usize & (unsafe { (*set).size } - 1)
}
unsafe fn ddict_hashset_create(custom_mem: ZSTD_customMem) -> *mut DDictHashSet {
let set = unsafe { ZSTD_rust_custom_malloc(size_of::<DDictHashSet>(), custom_mem) }
.cast::<DDictHashSet>();
if set.is_null() {
return ptr::null_mut();
}
let table = unsafe { ZSTD_rust_custom_calloc(64 * size_of::<*const ZSTD_DDict>(), custom_mem) }
.cast::<*const ZSTD_DDict>();
if table.is_null() {
unsafe { ZSTD_rust_custom_free(set.cast(), custom_mem) };
return ptr::null_mut();
}
unsafe {
set.write(DDictHashSet {
table,
size: 64,
count: 0,
});
}
set
}
unsafe fn ddict_hashset_free(set: *mut DDictHashSet, custom_mem: ZSTD_customMem) {
if set.is_null() {
return;
}
unsafe {
ZSTD_rust_custom_free((*set).table.cast(), custom_mem);
ZSTD_rust_custom_free(set.cast(), custom_mem);
}
}
unsafe fn ddict_hashset_emplace(set: *mut DDictHashSet, ddict: *const ZSTD_DDict) -> usize {
let dict_id = unsafe { ZSTD_getDictID_fromDDict(ddict) };
let mut index = unsafe { ddict_hash_index(set, dict_id) };
let mask = unsafe { (*set).size } - 1;
if unsafe { (*set).count == (*set).size } {
return ERROR(ZstdErrorCode::Generic);
}
while !unsafe { *(*set).table.add(index) }.is_null() {
if unsafe { ZSTD_getDictID_fromDDict(*(*set).table.add(index)) } == dict_id {
unsafe { (*set).table.add(index).write(ddict) };
return 0;
}
index = (index + 1) & mask;
}
unsafe {
(*set).table.add(index).write(ddict);
(*set).count += 1;
}
0
}
unsafe fn ddict_hashset_expand(set: *mut DDictHashSet, custom_mem: ZSTD_customMem) -> usize {
let old_table = unsafe { (*set).table };
let old_size = unsafe { (*set).size };
let new_size = match old_size.checked_mul(2) {
Some(size) => size,
None => return ERROR(ZstdErrorCode::MemoryAllocation),
};
let new_table =
unsafe { ZSTD_rust_custom_calloc(new_size * size_of::<*const ZSTD_DDict>(), custom_mem) }
.cast::<*const ZSTD_DDict>();
if new_table.is_null() {
return ERROR(ZstdErrorCode::MemoryAllocation);
}
unsafe {
(*set).table = new_table;
(*set).size = new_size;
(*set).count = 0;
for index in 0..old_size {
let ddict = *old_table.add(index);
if !ddict.is_null() {
let result = ddict_hashset_emplace(set, ddict);
if ERR_isError(result) {
return result;
}
}
}
ZSTD_rust_custom_free(old_table.cast(), custom_mem);
}
0
}
unsafe fn ddict_hashset_add(
set: *mut DDictHashSet,
ddict: *const ZSTD_DDict,
custom_mem: ZSTD_customMem,
) -> usize {
let should_expand = unsafe { (*set).count }
.wrapping_mul(4)
.wrapping_div(unsafe { (*set).size })
.wrapping_mul(3)
!= 0;
if should_expand {
let result = unsafe { ddict_hashset_expand(set, custom_mem) };
if ERR_isError(result) {
return result;
}
}
unsafe { ddict_hashset_emplace(set, ddict) }
}
unsafe fn ddict_hashset_get(set: *const DDictHashSet, dict_id: u32) -> *const ZSTD_DDict {
let mut index = unsafe { ddict_hash_index(set, dict_id) };
let mask = unsafe { (*set).size } - 1;
loop {
let ddict = unsafe { *(*set).table.add(index) };
let current_id = unsafe { ZSTD_getDictID_fromDDict(ddict) };
if current_id == dict_id || current_id == 0 {
return ddict;
}
index = (index + 1) & mask;
}
}
unsafe fn reset_parameters(view: &ZSTD_rustDctxView) {
unsafe {
set_field(view.format, ZSTD_F_ZSTD1);
set_field(
view.max_window_size,
ZSTD_rust_dctx_default_max_window_size(),
);
set_field(view.out_buffer_mode, ZSTD_BM_BUFFERED);
set_field(view.force_ignore_checksum, ZSTD_D_VALIDATE_CHECKSUM);
set_field(view.ref_multiple_ddicts, ZSTD_RMD_REF_SINGLE_DDICT);
set_field(view.disable_huf_asm, 0 as c_int);
set_field(view.max_block_size_param, 0 as c_int);
}
}
unsafe fn init_dctx_internal(view: &ZSTD_rustDctxView) {
unsafe {
set_field(view.static_size, 0usize);
set_dctx_ddict(view, ptr::null());
set_dctx_ddict_local(view, ptr::null_mut());
set_pointer(view.dict_end, ptr::null());
set_field(view.ddict_is_cold, 0 as c_int);
set_field(view.dict_uses, ZSTD_DONT_USE);
set_mut_pointer(view.in_buff, ptr::null_mut());
set_field(view.in_buff_size, 0usize);
set_field(view.out_buff_size, 0usize);
set_field(view.stream_stage, ZDSS_INIT);
if !view.legacy_context.is_null() {
set_field(view.legacy_context, ptr::null_mut::<c_void>());
set_field(view.previous_legacy_version, 0u32);
set_field(view.legacy_version, 0u32);
}
set_field(view.no_forward_progress, 0 as c_int);
set_field(view.oversized_duration, 0usize);
set_field(view.is_frame_decompression, 1 as c_int);
set_field(view.ddict_set, ptr::null_mut::<DDictHashSet>());
reset_parameters(view);
if !view.fuzz_end.is_null() {
set_pointer(view.fuzz_end, ptr::null());
}
ZSTD_rust_dctx_init_platform(view.dctx.cast());
}
}
unsafe fn clear_dict(view: &ZSTD_rustDctxView) {
unsafe {
let local = dctx_ddict_local(view);
if !local.is_null() {
ZSTD_freeDDict(local);
}
set_dctx_ddict_local(view, ptr::null_mut());
set_dctx_ddict(view, ptr::null());
set_field(view.dict_uses, ZSTD_DONT_USE);
}
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_sizeof_DCtx(dctx: *const ZSTD_DCtx) -> usize {
if dctx.is_null() {
return 0;
}
let view = unsafe { dctx_view(dctx.cast_mut()) };
let local = unsafe { dctx_ddict_local(&view) };
let ddict_size = if local.is_null() {
0
} else {
unsafe { crate::zstd_ddict::ZSTD_sizeof_DDict(local) }
};
view.dctx_size
.wrapping_add(ddict_size)
.wrapping_add(unsafe { field::<usize>(view.in_buff_size) })
.wrapping_add(unsafe { field::<usize>(view.out_buff_size) })
}
#[no_mangle]
pub extern "C" fn ZSTD_estimateDCtxSize() -> usize {
unsafe { ZSTD_rust_dctx_sizeof() }
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_initStaticDCtx(
workspace: *mut c_void,
workspace_size: usize,
) -> *mut ZSTD_DCtx {
let dctx_size = unsafe { ZSTD_rust_dctx_sizeof() };
if workspace.is_null() || (workspace as usize & 7) != 0 || workspace_size < dctx_size {
return ptr::null_mut();
}
let dctx = workspace.cast::<ZSTD_DCtx>();
let view = unsafe { dctx_view(dctx) };
unsafe {
set_field(view.custom_mem, default_custom_mem());
init_dctx_internal(&view);
set_field(view.static_size, workspace_size);
set_mut_pointer(view.in_buff, workspace.cast::<u8>().add(dctx_size));
}
dctx
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_createDCtx_advanced(custom_mem: ZSTD_customMem) -> *mut ZSTD_DCtx {
if !custom_mem_valid(custom_mem) {
return ptr::null_mut();
}
let dctx = unsafe { ZSTD_rust_dctx_alloc(custom_mem) };
if dctx.is_null() {
return ptr::null_mut();
}
let view = unsafe { dctx_view(dctx) };
unsafe {
set_field(view.custom_mem, custom_mem);
init_dctx_internal(&view);
}
dctx
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_createDCtx() -> *mut ZSTD_DCtx {
unsafe { ZSTD_createDCtx_advanced(default_custom_mem()) }
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_freeDCtx(dctx: *mut ZSTD_DCtx) -> usize {
if dctx.is_null() {
return 0;
}
let view = unsafe { dctx_view(dctx) };
if unsafe { field::<usize>(view.static_size) } != 0 {
return ERROR(ZstdErrorCode::MemoryAllocation);
}
let custom_mem = unsafe { dctx_custom_mem(&view) };
unsafe {
clear_dict(&view);
let in_buff = get_mut_pointer(view.in_buff);
ZSTD_rust_custom_free(in_buff.cast(), custom_mem);
set_mut_pointer(view.in_buff, ptr::null_mut());
let set: *mut DDictHashSet = field(view.ddict_set);
ddict_hashset_free(set, custom_mem);
set_field(view.ddict_set, ptr::null_mut::<DDictHashSet>());
if !view.legacy_context.is_null() {
ZSTD_rust_legacy_free_stream(dctx);
}
ZSTD_rust_dctx_free_storage(dctx, custom_mem);
}
0
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_copyDCtx(dst: *mut ZSTD_DCtx, src: *const ZSTD_DCtx) {
unsafe { ZSTD_rust_dctx_copy_prefix(dst, src) }
}
unsafe fn select_frame_ddict(view: &ZSTD_rustDctxView) {
let set: *mut DDictHashSet = unsafe { field(view.ddict_set) };
if set.is_null() || unsafe { dctx_ddict(view) }.is_null() {
return;
}
let dict_id = unsafe { (*get_frame_header_ptr(view)).dict_id };
let ddict = unsafe { ddict_hashset_get(set, dict_id) };
if !ddict.is_null() {
unsafe {
clear_dict(view);
set_field(view.dict_id, dict_id);
set_dctx_ddict(view, ddict);
set_field(view.dict_uses, ZSTD_USE_INDEFINITELY);
}
}
}
unsafe fn decode_frame_header(
view: &ZSTD_rustDctxView,
src: *const u8,
header_size: usize,
) -> usize {
let format = unsafe { field::<c_int>(view.format) };
let result = unsafe {
ZSTD_getFrameHeader_advanced(get_frame_header_ptr(view), src.cast(), header_size, format)
};
if ERR_isError(result) {
return result;
}
if result != 0 {
return ERROR(ZstdErrorCode::SrcSizeWrong);
}
if unsafe { field::<c_int>(view.ref_multiple_ddicts) } == ZSTD_RMD_REF_MULTIPLE_DDICTS
&& !unsafe { field::<*mut DDictHashSet>(view.ddict_set) }.is_null()
{
unsafe { select_frame_ddict(view) };
}
if view.fuzz_begin.is_null()
&& unsafe { (*get_frame_header_ptr(view)).dict_id } != 0
&& unsafe { field::<u32>(view.dict_id) } != unsafe { (*get_frame_header_ptr(view)).dict_id }
{
return ERROR(ZstdErrorCode::DictionaryWrong);
}
let validate = u32::from(
unsafe { (*get_frame_header_ptr(view)).checksum_flag } != 0
&& unsafe { field::<c_int>(view.force_ignore_checksum) } == ZSTD_D_VALIDATE_CHECKSUM,
);
unsafe {
set_field(view.validate_checksum, validate);
if validate != 0 {
let _ = XXH64_reset(view.xxh_state.cast::<XXH64_state_t>(), 0);
}
let processed = field::<u64>(view.processed_c_size).wrapping_add(header_size as u64);
set_field(view.processed_c_size, processed);
}
0
}
unsafe fn decompress_begin(view: &ZSTD_rustDctxView) -> usize {
unsafe {
ZSTD_rust_dctx_trace_begin(view.dctx.cast());
let format = field::<c_int>(view.format);
set_field(view.expected, frame_header_prefix(format));
set_field(view.stage, ZSTDDS_GET_FRAME_HEADER_SIZE);
set_field(view.processed_c_size, 0u64);
set_field(view.decoded_size, 0u64);
set_pointer(view.previous_dst_end, ptr::null());
set_pointer(view.prefix_start, ptr::null());
set_pointer(view.virtual_start, ptr::null());
set_pointer(view.dict_end, ptr::null());
let entropy = &mut *entropy_ptr(view);
entropy.huf_table[0] = (ZSTD_HUFFDTABLE_CAPACITY_LOG as u32).wrapping_mul(0x0100_0001);
set_field(view.lit_entropy, 0u32);
set_field(view.fse_entropy, 0u32);
set_field(view.dict_id, 0u32);
set_field(view.b_type, BT_RESERVED);
set_field(view.is_frame_decompression, 1 as c_int);
entropy.rep = [1, 4, 8];
set_field(view.llt_ptr, entropy.ll_table.as_ptr());
set_field(view.mlt_ptr, entropy.ml_table.as_ptr());
set_field(view.oft_ptr, entropy.of_table.as_ptr());
set_field(view.huf_ptr, entropy.huf_table.as_ptr());
}
0
}
unsafe fn decompress_insert_dictionary(
view: &ZSTD_rustDctxView,
mut dict: *const u8,
mut dict_size: usize,
) -> usize {
if dict_size < 8 || unsafe { MEM_readLE32(dict.cast()) } != ZSTD_MAGIC_DICTIONARY {
return unsafe { ref_dict_content(view, dict, dict_size) };
}
unsafe {
set_field(
view.dict_id,
MEM_readLE32(dict.add(ZSTD_FRAMEIDSIZE).cast()),
)
};
let entropy_size = unsafe { ZSTD_loadDEntropy(entropy_ptr(view), dict.cast(), dict_size) };
if ERR_isError(entropy_size) {
return ERROR(ZstdErrorCode::DictionaryCorrupted);
}
dict = unsafe { dict.add(entropy_size) };
dict_size -= entropy_size;
unsafe {
set_field(view.lit_entropy, 1u32);
set_field(view.fse_entropy, 1u32);
ref_dict_content(view, dict, dict_size)
}
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_decompressBegin(dctx: *mut ZSTD_DCtx) -> usize {
if dctx.is_null() {
return ERROR(ZstdErrorCode::Generic);
}
unsafe { decompress_begin(&dctx_view(dctx)) }
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_decompressBegin_usingDict(
dctx: *mut ZSTD_DCtx,
dict: *const c_void,
dict_size: usize,
) -> usize {
let view = unsafe { dctx_view(dctx) };
let result = unsafe { decompress_begin(&view) };
if ERR_isError(result) {
return result;
}
if !dict.is_null() && dict_size != 0 {
let result = unsafe { decompress_insert_dictionary(&view, dict.cast(), dict_size) };
if ERR_isError(result) {
return ERROR(ZstdErrorCode::DictionaryCorrupted);
}
}
0
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_decompressBegin_usingDDict(
dctx: *mut ZSTD_DCtx,
ddict: *const ZSTD_DDict,
) -> usize {
let view = unsafe { dctx_view(dctx) };
if !ddict.is_null() {
let dict_start = unsafe { ZSTD_DDict_dictContent(ddict) };
let dict_size = unsafe { ZSTD_DDict_dictSize(ddict) };
let dict_end = unsafe { dict_start.cast::<u8>().add(dict_size).cast::<c_void>() };
unsafe {
set_field(
view.ddict_is_cold,
c_int::from(get_pointer(view.dict_end).cast::<c_void>() != dict_end),
)
};
}
let result = unsafe { decompress_begin(&view) };
if ERR_isError(result) {
return result;
}
if !ddict.is_null() {
unsafe { ZSTD_copyDDictParameters(dctx.cast(), ddict) };
}
0
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_getDictID_fromDict(dict: *const c_void, dict_size: usize) -> c_uint {
if dict.is_null() || dict_size < 8 || unsafe { MEM_readLE32(dict) } != ZSTD_MAGIC_DICTIONARY {
0
} else {
unsafe { MEM_readLE32(dict.cast::<u8>().add(ZSTD_FRAMEIDSIZE).cast()) }
}
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_getDictID_fromFrame(src: *const c_void, src_size: usize) -> c_uint {
let mut zfh = ZSTD_FrameHeader::default();
if ERR_isError(unsafe { ZSTD_getFrameHeader(&mut zfh, src, src_size) }) {
0
} else {
zfh.dict_id
}
}
/*-*************************************************************
* Frame decoding
***************************************************************/
#[inline]
unsafe fn copy_raw_block(
dst: *mut u8,
dst_capacity: usize,
src: *const u8,
src_size: usize,
) -> usize {
if src_size > dst_capacity {
return ERROR(ZstdErrorCode::DstSizeTooSmall);
}
if dst.is_null() {
return if src_size == 0 {
0
} else {
ERROR(ZstdErrorCode::DstBufferNull)
};
}
unsafe { copy_bytes(dst, src, src_size) };
src_size
}
#[inline]
unsafe fn set_rle_block(
dst: *mut u8,
dst_capacity: usize,
byte: u8,
regenerated_size: usize,
) -> usize {
if regenerated_size > dst_capacity {
return ERROR(ZstdErrorCode::DstSizeTooSmall);
}
if dst.is_null() {
return if regenerated_size == 0 {
0
} else {
ERROR(ZstdErrorCode::DstBufferNull)
};
}
if regenerated_size != 0 {
unsafe { dst.write_bytes(byte, regenerated_size) };
}
regenerated_size
}
/// Decode exactly one non-skippable modern frame and advance the caller's
/// input cursor. The state preparation deliberately stays in the public
/// `decompressBegin*()` calls, just as in the C implementation.
unsafe fn decompress_frame(
dctx: *mut ZSTD_DCtx,
view: &ZSTD_rustDctxView,
dst: *mut c_void,
dst_capacity: usize,
src_ptr: &mut *const u8,
src_size_ptr: &mut usize,
) -> usize {
let istart = *src_ptr;
let mut ip = istart;
let ostart = dst.cast::<u8>();
/* `dst == NULL, dstCapacity == 0` is supported for empty frames. A
* wrapping endpoint preserves C's address-only calculation until a
* block decoder reports the appropriate null/size error. */
let oend = ostart.wrapping_add(dst_capacity);
let mut op = ostart;
let mut remaining = *src_size_ptr;
let format = unsafe { field::<c_int>(view.format) };
if remaining < frame_header_min(format) + ZSTD_BLOCKHEADERSIZE {
return ERROR(ZstdErrorCode::SrcSizeWrong);
}
let header_size =
unsafe { frame_header_size_internal(ip, frame_header_prefix(format), format) };
if ERR_isError(header_size) {
return header_size;
}
if remaining < header_size + ZSTD_BLOCKHEADERSIZE {
return ERROR(ZstdErrorCode::SrcSizeWrong);
}
let result = unsafe { decode_frame_header(view, ip, header_size) };
if ERR_isError(result) {
return result;
}
ip = unsafe { ip.add(header_size) };
remaining -= header_size;
let max_block_size_param = unsafe { field::<c_int>(view.max_block_size_param) };
if max_block_size_param != 0 {
let mut params = unsafe { field::<ZSTD_FrameHeader>(view.f_params) };
params.block_size_max = min(params.block_size_max, max_block_size_param as c_uint);
unsafe { set_field(view.f_params, params) };
}
loop {
let mut block = BlockProperties::default();
let c_block_size = unsafe {
crate::zstd_decompress_block::ZSTD_getcBlockSize(
ip.cast(),
remaining,
(&mut block as *mut BlockProperties).cast(),
)
};
if ERR_isError(c_block_size) {
return c_block_size;
}
ip = unsafe { ip.add(ZSTD_BLOCKHEADERSIZE) };
remaining -= ZSTD_BLOCKHEADERSIZE;
if c_block_size > remaining {
return ERROR(ZstdErrorCode::SrcSizeWrong);
}
let mut block_end = oend;
if (ip as usize) >= (op as usize) && (ip as usize) < (block_end as usize) {
block_end = op.wrapping_add((ip as usize).wrapping_sub(op as usize));
}
let block_capacity = (block_end as usize).wrapping_sub(op as usize);
let decoded_size = match block.block_type {
BT_COMPRESSED => unsafe {
ZSTD_decompressBlock_internal(
dctx,
op.cast(),
block_capacity,
ip.cast(),
c_block_size,
0,
)
},
/* This deliberately uses `oend`, not `block_end`: memmove is
* overlap-safe for raw blocks. */
BT_RAW => unsafe {
copy_raw_block(
op,
(oend as usize).wrapping_sub(op as usize),
ip,
c_block_size,
)
},
BT_RLE => {
if c_block_size == 0 {
return ERROR(ZstdErrorCode::CorruptionDetected);
}
unsafe { set_rle_block(op, block_capacity, *ip, block.orig_size as usize) }
}
_ => ERROR(ZstdErrorCode::CorruptionDetected),
};
if ERR_isError(decoded_size) {
return decoded_size;
}
if unsafe { field::<u32>(view.validate_checksum) } != 0 {
let _ = unsafe {
XXH64_update(
view.xxh_state.cast::<XXH64_state_t>(),
op.cast(),
decoded_size,
)
};
}
if decoded_size != 0 {
op = unsafe { op.add(decoded_size) };
}
ip = unsafe { ip.add(c_block_size) };
remaining -= c_block_size;
if block.last_block != 0 {
break;
}
}
let params = unsafe { field::<ZSTD_FrameHeader>(view.f_params) };
let decoded = (op as usize).wrapping_sub(ostart as usize);
if params.frame_content_size != ZSTD_CONTENTSIZE_UNKNOWN
&& decoded as u64 != params.frame_content_size
{
return ERROR(ZstdErrorCode::CorruptionDetected);
}
if params.checksum_flag != 0 {
if remaining < 4 {
return ERROR(ZstdErrorCode::ChecksumWrong);
}
if unsafe { field::<c_int>(view.force_ignore_checksum) } == ZSTD_D_VALIDATE_CHECKSUM {
let calculated = unsafe { XXH64_digest(view.xxh_state.cast::<XXH64_state_t>()) } as u32;
let read = unsafe { MEM_readLE32(ip.cast()) };
if calculated != read {
return ERROR(ZstdErrorCode::ChecksumWrong);
}
}
ip = unsafe { ip.add(4) };
remaining -= 4;
}
unsafe {
ZSTD_rust_dctx_trace_end(
dctx,
decoded as u64,
(ip as usize).wrapping_sub(istart as usize) as u64,
0,
);
}
*src_ptr = ip;
*src_size_ptr = remaining;
decoded
}
unsafe fn decompress_multi_frame(
dctx: *mut ZSTD_DCtx,
dst: *mut c_void,
mut dst_capacity: usize,
src: *const c_void,
mut src_size: usize,
mut dict: *const c_void,
mut dict_size: usize,
ddict: *const ZSTD_DDict,
) -> usize {
if dctx.is_null() {
return ERROR(ZstdErrorCode::Generic);
}
let view = unsafe { dctx_view(dctx) };
if !ddict.is_null() {
dict = unsafe { ZSTD_DDict_dictContent(ddict) };
dict_size = unsafe { ZSTD_DDict_dictSize(ddict) };
}
let dst_start = dst.cast::<u8>();
let mut output = dst_start;
let mut input = src.cast::<u8>();
let mut more_than_one_frame = false;
let starting_input = frame_header_prefix(unsafe { field::<c_int>(view.format) });
while src_size >= starting_input {
if unsafe { field::<c_int>(view.format) } == ZSTD_F_ZSTD1
&& unsafe { ZSTD_rust_legacy_is(input.cast(), src_size) } != 0
{
let frame_size =
unsafe { ZSTD_rust_legacy_find_compressed_size(input.cast(), src_size) };
if ERR_isError(frame_size) {
return frame_size;
}
if unsafe { field::<usize>(view.static_size) } != 0 {
return ERROR(ZstdErrorCode::MemoryAllocation);
}
if frame_size > src_size {
return ERROR(ZstdErrorCode::SrcSizeWrong);
}
let decoded = unsafe {
ZSTD_rust_legacy_decompress(
output.cast(),
dst_capacity,
input.cast(),
frame_size,
dict,
dict_size,
)
};
if ERR_isError(decoded) {
return decoded;
}
let expected = unsafe { ZSTD_getFrameContentSize(input.cast(), src_size) };
if expected == ZSTD_CONTENTSIZE_ERROR
|| (expected != ZSTD_CONTENTSIZE_UNKNOWN && expected != decoded as u64)
{
return ERROR(ZstdErrorCode::CorruptionDetected);
}
if decoded > dst_capacity {
return ERROR(ZstdErrorCode::DstSizeTooSmall);
}
if decoded != 0 {
output = unsafe { output.add(decoded) };
}
dst_capacity -= decoded;
input = unsafe { input.add(frame_size) };
src_size -= frame_size;
continue;
}
if unsafe { field::<c_int>(view.format) } == ZSTD_F_ZSTD1 && src_size >= ZSTD_FRAMEIDSIZE {
let magic = unsafe { MEM_readLE32(input.cast()) };
if magic & ZSTD_MAGIC_SKIPPABLE_MASK == ZSTD_MAGIC_SKIPPABLE_START {
let size = unsafe { read_skippable_frame_size(input, src_size) };
if ERR_isError(size) {
return size;
}
input = unsafe { input.add(size) };
src_size -= size;
continue;
}
}
let init = if !ddict.is_null() {
unsafe { ZSTD_decompressBegin_usingDDict(dctx, ddict) }
} else {
unsafe { ZSTD_decompressBegin_usingDict(dctx, dict, dict_size) }
};
if ERR_isError(init) {
return init;
}
unsafe { ZSTD_checkContinuity(dctx, output.cast(), dst_capacity) };
let decoded = unsafe {
decompress_frame(
dctx,
&view,
output.cast(),
dst_capacity,
&mut input,
&mut src_size,
)
};
if decoded == ERROR(ZstdErrorCode::PrefixUnknown) && more_than_one_frame {
return ERROR(ZstdErrorCode::SrcSizeWrong);
}
if ERR_isError(decoded) {
return decoded;
}
if decoded > dst_capacity {
return ERROR(ZstdErrorCode::DstSizeTooSmall);
}
if decoded != 0 {
output = unsafe { output.add(decoded) };
}
dst_capacity -= decoded;
more_than_one_frame = true;
}
if src_size != 0 {
return ERROR(ZstdErrorCode::SrcSizeWrong);
}
(output as usize).wrapping_sub(dst_start as usize)
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_insertBlock(
dctx: *mut ZSTD_DCtx,
block_start: *const c_void,
block_size: usize,
) -> usize {
if dctx.is_null() {
return ERROR(ZstdErrorCode::Generic);
}
unsafe {
ZSTD_checkContinuity(dctx, block_start, block_size);
let view = dctx_view(dctx);
set_pointer(
view.previous_dst_end,
const_ptr_add(block_start.cast(), block_size),
);
}
block_size
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_decompress_usingDict(
dctx: *mut ZSTD_DCtx,
dst: *mut c_void,
dst_capacity: usize,
src: *const c_void,
src_size: usize,
dict: *const c_void,
dict_size: usize,
) -> usize {
unsafe {
decompress_multi_frame(
dctx,
dst,
dst_capacity,
src,
src_size,
dict,
dict_size,
ptr::null(),
)
}
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_decompress_usingDDict(
dctx: *mut ZSTD_DCtx,
dst: *mut c_void,
dst_capacity: usize,
src: *const c_void,
src_size: usize,
ddict: *const ZSTD_DDict,
) -> usize {
unsafe {
decompress_multi_frame(
dctx,
dst,
dst_capacity,
src,
src_size,
ptr::null(),
0,
ddict,
)
}
}
unsafe fn get_ddict(view: &ZSTD_rustDctxView) -> *const ZSTD_DDict {
match unsafe { field::<c_int>(view.dict_uses) } {
ZSTD_DONT_USE => {
unsafe { clear_dict(view) };
ptr::null()
}
ZSTD_USE_INDEFINITELY => unsafe { dctx_ddict(view) },
ZSTD_USE_ONCE => {
unsafe { set_field(view.dict_uses, ZSTD_DONT_USE) };
unsafe { dctx_ddict(view) }
}
_ => {
unsafe { clear_dict(view) };
ptr::null()
}
}
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_decompressDCtx(
dctx: *mut ZSTD_DCtx,
dst: *mut c_void,
dst_capacity: usize,
src: *const c_void,
src_size: usize,
) -> usize {
if dctx.is_null() {
return ERROR(ZstdErrorCode::Generic);
}
let view = unsafe { dctx_view(dctx) };
let ddict = unsafe { get_ddict(&view) };
unsafe { ZSTD_decompress_usingDDict(dctx, dst, dst_capacity, src, src_size, ddict) }
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_decompress(
dst: *mut c_void,
dst_capacity: usize,
src: *const c_void,
src_size: usize,
) -> usize {
if unsafe { ZSTD_rust_heapmode() } < 1 {
return unsafe { ZSTD_rust_decompress_stack(dst, dst_capacity, src, src_size) };
}
let dctx = unsafe { ZSTD_createDCtx() };
if dctx.is_null() {
return ERROR(ZstdErrorCode::MemoryAllocation);
}
let result = unsafe { ZSTD_decompressDCtx(dctx, dst, dst_capacity, src, src_size) };
let _ = unsafe { ZSTD_freeDCtx(dctx) };
result
}
/*-**************************************
* Advanced bufferless decompression
****************************************/
#[inline]
unsafe fn next_src_size_with_input_size(view: &ZSTD_rustDctxView, input_size: usize) -> usize {
let stage = unsafe { field::<c_int>(view.stage) };
if (stage == ZSTDDS_DECOMPRESS_BLOCK || stage == ZSTDDS_DECOMPRESS_LAST_BLOCK)
&& unsafe { field::<c_int>(view.b_type) } == BT_RAW
{
let expected = unsafe { field::<usize>(view.expected) };
return max(1, min(input_size, expected));
}
unsafe { field::<usize>(view.expected) }
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_nextSrcSizeToDecompress(dctx: *mut ZSTD_DCtx) -> usize {
if dctx.is_null() {
return 0;
}
let view = unsafe { dctx_view(dctx) };
unsafe { field(view.expected) }
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_nextInputType(dctx: *mut ZSTD_DCtx) -> c_int {
if dctx.is_null() {
return ZSTD_NIT_FRAME_HEADER;
}
let view = unsafe { dctx_view(dctx) };
match unsafe { field::<c_int>(view.stage) } {
ZSTDDS_GET_FRAME_HEADER_SIZE | ZSTDDS_DECODE_FRAME_HEADER => ZSTD_NIT_FRAME_HEADER,
ZSTDDS_DECODE_BLOCK_HEADER => ZSTD_NIT_BLOCK_HEADER,
ZSTDDS_DECOMPRESS_BLOCK => ZSTD_NIT_BLOCK,
ZSTDDS_DECOMPRESS_LAST_BLOCK => ZSTD_NIT_LAST_BLOCK,
ZSTDDS_CHECK_CHECKSUM => ZSTD_NIT_CHECKSUM,
ZSTDDS_DECODE_SKIPPABLE_HEADER | ZSTDDS_SKIP_FRAME => ZSTD_NIT_SKIPPABLE_FRAME,
_ => ZSTD_NIT_FRAME_HEADER,
}
}
#[inline]
unsafe fn is_skip_frame(view: &ZSTD_rustDctxView) -> bool {
(unsafe { field::<c_int>(view.stage) }) == ZSTDDS_SKIP_FRAME
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_decompressContinue(
dctx: *mut ZSTD_DCtx,
dst: *mut c_void,
dst_capacity: usize,
src: *const c_void,
src_size: usize,
) -> usize {
if dctx.is_null() {
return ERROR(ZstdErrorCode::Generic);
}
let view = unsafe { dctx_view(dctx) };
if src_size != unsafe { next_src_size_with_input_size(&view, src_size) } {
return ERROR(ZstdErrorCode::SrcSizeWrong);
}
unsafe { ZSTD_checkContinuity(dctx, dst.cast(), dst_capacity) };
let processed = unsafe { field::<u64>(view.processed_c_size) }.wrapping_add(src_size as u64);
unsafe { set_field(view.processed_c_size, processed) };
match unsafe { field::<c_int>(view.stage) } {
ZSTDDS_GET_FRAME_HEADER_SIZE => {
let format = unsafe { field::<c_int>(view.format) };
let header_buffer = view.header_buffer.cast::<u8>();
if format == ZSTD_F_ZSTD1
&& src_size >= ZSTD_FRAMEIDSIZE
&& unsafe { MEM_readLE32(src) } & ZSTD_MAGIC_SKIPPABLE_MASK
== ZSTD_MAGIC_SKIPPABLE_START
{
unsafe { copy_bytes(header_buffer, src.cast(), src_size) };
unsafe {
set_field(view.expected, ZSTD_SKIPPABLEHEADERSIZE - src_size);
set_field(view.stage, ZSTDDS_DECODE_SKIPPABLE_HEADER);
}
return 0;
}
let header_size = unsafe { frame_header_size_internal(src.cast(), src_size, format) };
if ERR_isError(header_size) {
return header_size;
}
unsafe {
copy_bytes(header_buffer, src.cast(), src_size);
set_field(view.header_size, header_size);
set_field(view.expected, header_size - src_size);
set_field(view.stage, ZSTDDS_DECODE_FRAME_HEADER);
}
0
}
ZSTDDS_DECODE_FRAME_HEADER => {
let header_size = unsafe { field::<usize>(view.header_size) };
let offset = header_size - src_size;
unsafe {
copy_bytes(
view.header_buffer.cast::<u8>().add(offset),
src.cast(),
src_size,
);
}
let result =
unsafe { decode_frame_header(&view, view.header_buffer.cast(), header_size) };
if ERR_isError(result) {
return result;
}
unsafe {
set_field(view.expected, ZSTD_BLOCKHEADERSIZE);
set_field(view.stage, ZSTDDS_DECODE_BLOCK_HEADER);
}
0
}
ZSTDDS_DECODE_BLOCK_HEADER => {
let mut block = BlockProperties::default();
let c_block_size = unsafe {
crate::zstd_decompress_block::ZSTD_getcBlockSize(
src,
ZSTD_BLOCKHEADERSIZE,
(&mut block as *mut BlockProperties).cast(),
)
};
if ERR_isError(c_block_size) {
return c_block_size;
}
let params = unsafe { field::<ZSTD_FrameHeader>(view.f_params) };
if c_block_size > params.block_size_max as usize {
return ERROR(ZstdErrorCode::CorruptionDetected);
}
unsafe {
set_field(view.expected, c_block_size);
set_field(view.b_type, block.block_type);
set_field(view.rle_size, block.orig_size as usize);
}
if c_block_size != 0 {
unsafe {
set_field(
view.stage,
if block.last_block != 0 {
ZSTDDS_DECOMPRESS_LAST_BLOCK
} else {
ZSTDDS_DECOMPRESS_BLOCK
},
);
}
return 0;
}
unsafe {
if block.last_block != 0 {
if params.checksum_flag != 0 {
set_field(view.expected, 4usize);
set_field(view.stage, ZSTDDS_CHECK_CHECKSUM);
} else {
set_field(view.expected, 0usize);
set_field(view.stage, ZSTDDS_GET_FRAME_HEADER_SIZE);
}
} else {
set_field(view.expected, ZSTD_BLOCKHEADERSIZE);
set_field(view.stage, ZSTDDS_DECODE_BLOCK_HEADER);
}
}
0
}
ZSTDDS_DECOMPRESS_BLOCK | ZSTDDS_DECOMPRESS_LAST_BLOCK => {
let stage = unsafe { field::<c_int>(view.stage) };
let b_type = unsafe { field::<c_int>(view.b_type) };
let decoded = match b_type {
BT_COMPRESSED => {
let result = unsafe {
ZSTD_decompressBlock_internal(dctx, dst, dst_capacity, src, src_size, 1)
};
unsafe { set_field(view.expected, 0usize) };
result
}
BT_RAW => {
let result =
unsafe { copy_raw_block(dst.cast(), dst_capacity, src.cast(), src_size) };
if ERR_isError(result) {
return result;
}
let expected = unsafe { field::<usize>(view.expected) } - result;
unsafe { set_field(view.expected, expected) };
result
}
BT_RLE => {
if src_size == 0 {
return ERROR(ZstdErrorCode::CorruptionDetected);
}
let result = unsafe {
set_rle_block(
dst.cast(),
dst_capacity,
*src.cast::<u8>(),
field::<usize>(view.rle_size),
)
};
unsafe { set_field(view.expected, 0usize) };
result
}
_ => return ERROR(ZstdErrorCode::CorruptionDetected),
};
if ERR_isError(decoded) {
return decoded;
}
let params = unsafe { field::<ZSTD_FrameHeader>(view.f_params) };
if decoded > params.block_size_max as usize {
return ERROR(ZstdErrorCode::CorruptionDetected);
}
let decoded_total =
unsafe { field::<u64>(view.decoded_size) }.wrapping_add(decoded as u64);
unsafe {
set_field(view.decoded_size, decoded_total);
if field::<u32>(view.validate_checksum) != 0 {
let _ = XXH64_update(view.xxh_state.cast::<XXH64_state_t>(), dst, decoded);
}
set_pointer(view.previous_dst_end, const_ptr_add(dst.cast(), decoded));
}
if unsafe { field::<usize>(view.expected) } != 0 {
return decoded;
}
if stage == ZSTDDS_DECOMPRESS_LAST_BLOCK {
if params.frame_content_size != ZSTD_CONTENTSIZE_UNKNOWN
&& decoded_total != params.frame_content_size
{
return ERROR(ZstdErrorCode::CorruptionDetected);
}
unsafe {
if params.checksum_flag != 0 {
set_field(view.expected, 4usize);
set_field(view.stage, ZSTDDS_CHECK_CHECKSUM);
} else {
ZSTD_rust_dctx_trace_end(
dctx,
decoded_total,
field::<u64>(view.processed_c_size),
1,
);
set_field(view.expected, 0usize);
set_field(view.stage, ZSTDDS_GET_FRAME_HEADER_SIZE);
}
}
} else {
unsafe {
set_field(view.stage, ZSTDDS_DECODE_BLOCK_HEADER);
set_field(view.expected, ZSTD_BLOCKHEADERSIZE);
}
}
decoded
}
ZSTDDS_CHECK_CHECKSUM => {
if src_size != 4 {
return ERROR(ZstdErrorCode::SrcSizeWrong);
}
if unsafe { field::<u32>(view.validate_checksum) } != 0 {
let calculated =
unsafe { XXH64_digest(view.xxh_state.cast::<XXH64_state_t>()) } as u32;
let read = unsafe { MEM_readLE32(src) };
if calculated != read {
return ERROR(ZstdErrorCode::ChecksumWrong);
}
}
unsafe {
ZSTD_rust_dctx_trace_end(
dctx,
field::<u64>(view.decoded_size),
field::<u64>(view.processed_c_size),
1,
);
set_field(view.expected, 0usize);
set_field(view.stage, ZSTDDS_GET_FRAME_HEADER_SIZE);
}
0
}
ZSTDDS_DECODE_SKIPPABLE_HEADER => {
if src_size > ZSTD_SKIPPABLEHEADERSIZE {
return ERROR(ZstdErrorCode::SrcSizeWrong);
}
unsafe {
let header = view.header_buffer.cast::<u8>();
copy_bytes(
header.add(ZSTD_SKIPPABLEHEADERSIZE - src_size),
src.cast(),
src_size,
);
set_field(
view.expected,
MEM_readLE32(header.add(ZSTD_FRAMEIDSIZE).cast()) as usize,
);
set_field(view.stage, ZSTDDS_SKIP_FRAME);
}
0
}
ZSTDDS_SKIP_FRAME => {
unsafe {
set_field(view.expected, 0usize);
set_field(view.stage, ZSTDDS_GET_FRAME_HEADER_SIZE);
}
0
}
_ => ERROR(ZstdErrorCode::Generic),
}
}
/*-**************************************
* Streaming context management
****************************************/
#[no_mangle]
pub unsafe extern "C" fn ZSTD_createDStream() -> *mut ZSTD_DStream {
unsafe { ZSTD_createDCtx() }
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_initStaticDStream(
workspace: *mut c_void,
workspace_size: usize,
) -> *mut ZSTD_DStream {
unsafe { ZSTD_initStaticDCtx(workspace, workspace_size) }
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_createDStream_advanced(
custom_mem: ZSTD_customMem,
) -> *mut ZSTD_DStream {
unsafe { ZSTD_createDCtx_advanced(custom_mem) }
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_freeDStream(zds: *mut ZSTD_DStream) -> usize {
unsafe { ZSTD_freeDCtx(zds) }
}
#[no_mangle]
pub extern "C" fn ZSTD_DStreamInSize() -> usize {
ZSTD_BLOCKSIZE_MAX + ZSTD_BLOCKHEADERSIZE
}
#[no_mangle]
pub extern "C" fn ZSTD_DStreamOutSize() -> usize {
ZSTD_BLOCKSIZE_MAX
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_DCtx_loadDictionary_advanced(
dctx: *mut ZSTD_DCtx,
dict: *const c_void,
dict_size: usize,
dict_load_method: c_int,
dict_content_type: c_int,
) -> usize {
if dctx.is_null() {
return ERROR(ZstdErrorCode::Generic);
}
let view = unsafe { dctx_view(dctx) };
if unsafe { field::<c_int>(view.stream_stage) } != ZDSS_INIT {
return ERROR(ZstdErrorCode::StageWrong);
}
unsafe { clear_dict(&view) };
if !dict.is_null() && dict_size != 0 {
let ddict = unsafe {
ZSTD_rust_create_ddict(
dict,
dict_size,
dict_load_method,
dict_content_type,
dctx_custom_mem(&view),
)
};
if ddict.is_null() {
return ERROR(ZstdErrorCode::MemoryAllocation);
}
unsafe {
set_dctx_ddict_local(&view, ddict);
set_dctx_ddict(&view, ddict);
set_field(view.dict_uses, ZSTD_USE_INDEFINITELY);
}
}
0
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_DCtx_loadDictionary_byReference(
dctx: *mut ZSTD_DCtx,
dict: *const c_void,
dict_size: usize,
) -> usize {
unsafe {
ZSTD_DCtx_loadDictionary_advanced(dctx, dict, dict_size, ZSTD_DLM_BY_REF, ZSTD_DCT_AUTO)
}
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_DCtx_loadDictionary(
dctx: *mut ZSTD_DCtx,
dict: *const c_void,
dict_size: usize,
) -> usize {
unsafe {
ZSTD_DCtx_loadDictionary_advanced(dctx, dict, dict_size, ZSTD_DLM_BY_COPY, ZSTD_DCT_AUTO)
}
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_DCtx_refPrefix_advanced(
dctx: *mut ZSTD_DCtx,
prefix: *const c_void,
prefix_size: usize,
dict_content_type: c_int,
) -> usize {
let result = unsafe {
ZSTD_DCtx_loadDictionary_advanced(
dctx,
prefix,
prefix_size,
ZSTD_DLM_BY_REF,
dict_content_type,
)
};
if ERR_isError(result) {
return result;
}
if dctx.is_null() {
return ERROR(ZstdErrorCode::Generic);
}
let view = unsafe { dctx_view(dctx) };
unsafe { set_field(view.dict_uses, ZSTD_USE_ONCE) };
0
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_DCtx_refPrefix(
dctx: *mut ZSTD_DCtx,
prefix: *const c_void,
prefix_size: usize,
) -> usize {
unsafe { ZSTD_DCtx_refPrefix_advanced(dctx, prefix, prefix_size, ZSTD_DCT_RAW_CONTENT) }
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_initDStream_usingDict(
zds: *mut ZSTD_DStream,
dict: *const c_void,
dict_size: usize,
) -> usize {
let result = unsafe { ZSTD_DCtx_reset(zds, ZSTD_RESET_SESSION_ONLY) };
if ERR_isError(result) {
return result;
}
let result = unsafe { ZSTD_DCtx_loadDictionary(zds, dict, dict_size) };
if ERR_isError(result) {
return result;
}
if zds.is_null() {
return ERROR(ZstdErrorCode::Generic);
}
let view = unsafe { dctx_view(zds) };
frame_header_prefix(unsafe { field::<c_int>(view.format) })
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_initDStream(zds: *mut ZSTD_DStream) -> usize {
let result = unsafe { ZSTD_DCtx_reset(zds, ZSTD_RESET_SESSION_ONLY) };
if ERR_isError(result) {
return result;
}
let result = unsafe { ZSTD_DCtx_refDDict(zds, ptr::null()) };
if ERR_isError(result) {
return result;
}
if zds.is_null() {
return ERROR(ZstdErrorCode::Generic);
}
let view = unsafe { dctx_view(zds) };
frame_header_prefix(unsafe { field::<c_int>(view.format) })
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_initDStream_usingDDict(
zds: *mut ZSTD_DStream,
ddict: *const ZSTD_DDict,
) -> usize {
let result = unsafe { ZSTD_DCtx_reset(zds, ZSTD_RESET_SESSION_ONLY) };
if ERR_isError(result) {
return result;
}
let result = unsafe { ZSTD_DCtx_refDDict(zds, ddict) };
if ERR_isError(result) {
return result;
}
if zds.is_null() {
return ERROR(ZstdErrorCode::Generic);
}
let view = unsafe { dctx_view(zds) };
frame_header_prefix(unsafe { field::<c_int>(view.format) })
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_resetDStream(zds: *mut ZSTD_DStream) -> usize {
let result = unsafe { ZSTD_DCtx_reset(zds, ZSTD_RESET_SESSION_ONLY) };
if ERR_isError(result) {
return result;
}
if zds.is_null() {
return ERROR(ZstdErrorCode::Generic);
}
let view = unsafe { dctx_view(zds) };
frame_header_prefix(unsafe { field::<c_int>(view.format) })
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_DCtx_refDDict(
dctx: *mut ZSTD_DCtx,
ddict: *const ZSTD_DDict,
) -> usize {
if dctx.is_null() {
return ERROR(ZstdErrorCode::Generic);
}
let view = unsafe { dctx_view(dctx) };
if unsafe { field::<c_int>(view.stream_stage) } != ZDSS_INIT {
return ERROR(ZstdErrorCode::StageWrong);
}
unsafe { clear_dict(&view) };
if ddict.is_null() {
return 0;
}
unsafe {
set_dctx_ddict(&view, ddict);
set_field(view.dict_uses, ZSTD_USE_INDEFINITELY);
}
if unsafe { field::<c_int>(view.ref_multiple_ddicts) } == ZSTD_RMD_REF_MULTIPLE_DDICTS {
let mut set: *mut DDictHashSet = unsafe { field(view.ddict_set) };
if set.is_null() {
if unsafe { field::<usize>(view.static_size) } != 0 {
return ERROR(ZstdErrorCode::ParameterUnsupported);
}
set = unsafe { ddict_hashset_create(dctx_custom_mem(&view)) };
if set.is_null() {
return ERROR(ZstdErrorCode::MemoryAllocation);
}
unsafe { set_field(view.ddict_set, set) };
}
let result = unsafe { ddict_hashset_add(set, ddict, dctx_custom_mem(&view)) };
if ERR_isError(result) {
return result;
}
}
0
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_DCtx_setMaxWindowSize(
dctx: *mut ZSTD_DCtx,
max_window_size: usize,
) -> usize {
if dctx.is_null() {
return ERROR(ZstdErrorCode::Generic);
}
let view = unsafe { dctx_view(dctx) };
let bounds = ZSTD_dParam_getBounds(ZSTD_D_WINDOW_LOG_MAX);
let minimum = 1usize << bounds.lower_bound;
let maximum = 1usize << bounds.upper_bound;
if unsafe { field::<c_int>(view.stream_stage) } != ZDSS_INIT {
return ERROR(ZstdErrorCode::StageWrong);
}
if max_window_size < minimum || max_window_size > maximum {
return ERROR(ZstdErrorCode::ParameterOutOfBound);
}
unsafe { set_field(view.max_window_size, max_window_size) };
0
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_DCtx_setFormat(dctx: *mut ZSTD_DCtx, format: c_int) -> usize {
unsafe { ZSTD_DCtx_setParameter(dctx, ZSTD_D_FORMAT, format) }
}
#[no_mangle]
pub extern "C" fn ZSTD_dParam_getBounds(param: c_int) -> ZSTD_bounds {
match param {
ZSTD_D_WINDOW_LOG_MAX => ZSTD_bounds {
error: 0,
lower_bound: ZSTD_WINDOWLOG_ABSOLUTEMIN as c_int,
upper_bound: window_log_max() as c_int,
},
ZSTD_D_FORMAT => ZSTD_bounds {
error: 0,
lower_bound: ZSTD_F_ZSTD1,
upper_bound: ZSTD_F_ZSTD1_MAGICLESS,
},
ZSTD_D_STABLE_OUT_BUFFER => ZSTD_bounds {
error: 0,
lower_bound: ZSTD_BM_BUFFERED,
upper_bound: ZSTD_BM_STABLE,
},
ZSTD_D_FORCE_IGNORE_CHECKSUM => ZSTD_bounds {
error: 0,
lower_bound: ZSTD_D_VALIDATE_CHECKSUM,
upper_bound: ZSTD_D_IGNORE_CHECKSUM,
},
ZSTD_D_REF_MULTIPLE_DDICTS => ZSTD_bounds {
error: 0,
lower_bound: ZSTD_RMD_REF_SINGLE_DDICT,
upper_bound: ZSTD_RMD_REF_MULTIPLE_DDICTS,
},
ZSTD_D_DISABLE_HUFFMAN_ASSEMBLY => ZSTD_bounds {
error: 0,
lower_bound: 0,
upper_bound: 1,
},
ZSTD_D_MAX_BLOCK_SIZE => ZSTD_bounds {
error: 0,
lower_bound: ZSTD_BLOCKSIZE_MAX_MIN as c_int,
upper_bound: ZSTD_BLOCKSIZE_MAX as c_int,
},
_ => ZSTD_bounds {
error: ERROR(ZstdErrorCode::ParameterUnsupported),
lower_bound: 0,
upper_bound: 0,
},
}
}
#[inline]
fn dparam_within_bounds(param: c_int, value: c_int) -> bool {
let bounds = ZSTD_dParam_getBounds(param);
!ERR_isError(bounds.error) && value >= bounds.lower_bound && value <= bounds.upper_bound
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_DCtx_getParameter(
dctx: *mut ZSTD_DCtx,
param: c_int,
value: *mut c_int,
) -> usize {
if dctx.is_null() || value.is_null() {
return ERROR(ZstdErrorCode::Generic);
}
let view = unsafe { dctx_view(dctx) };
let result = match param {
ZSTD_D_WINDOW_LOG_MAX => {
let window = unsafe { field::<usize>(view.max_window_size) } as u32;
(u32::BITS - 1 - window.leading_zeros()) as c_int
}
ZSTD_D_FORMAT => unsafe { field::<c_int>(view.format) },
ZSTD_D_STABLE_OUT_BUFFER => unsafe { field::<c_int>(view.out_buffer_mode) },
ZSTD_D_FORCE_IGNORE_CHECKSUM => unsafe { field::<c_int>(view.force_ignore_checksum) },
ZSTD_D_REF_MULTIPLE_DDICTS => unsafe { field::<c_int>(view.ref_multiple_ddicts) },
ZSTD_D_DISABLE_HUFFMAN_ASSEMBLY => unsafe { field::<c_int>(view.disable_huf_asm) },
ZSTD_D_MAX_BLOCK_SIZE => unsafe { field::<c_int>(view.max_block_size_param) },
_ => return ERROR(ZstdErrorCode::ParameterUnsupported),
};
unsafe { value.write(result) };
0
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_DCtx_setParameter(
dctx: *mut ZSTD_DCtx,
param: c_int,
mut value: c_int,
) -> usize {
if dctx.is_null() {
return ERROR(ZstdErrorCode::Generic);
}
let view = unsafe { dctx_view(dctx) };
if unsafe { field::<c_int>(view.stream_stage) } != ZDSS_INIT {
return ERROR(ZstdErrorCode::StageWrong);
}
match param {
ZSTD_D_WINDOW_LOG_MAX => {
if value == 0 {
value = ZSTD_WINDOWLOG_LIMIT_DEFAULT as c_int;
}
if !dparam_within_bounds(param, value) {
return ERROR(ZstdErrorCode::ParameterOutOfBound);
}
unsafe { set_field(view.max_window_size, 1usize << value) };
}
ZSTD_D_FORMAT => {
if !dparam_within_bounds(param, value) {
return ERROR(ZstdErrorCode::ParameterOutOfBound);
}
unsafe { set_field(view.format, value) };
}
ZSTD_D_STABLE_OUT_BUFFER => {
if !dparam_within_bounds(param, value) {
return ERROR(ZstdErrorCode::ParameterOutOfBound);
}
unsafe { set_field(view.out_buffer_mode, value) };
}
ZSTD_D_FORCE_IGNORE_CHECKSUM => {
if !dparam_within_bounds(param, value) {
return ERROR(ZstdErrorCode::ParameterOutOfBound);
}
unsafe { set_field(view.force_ignore_checksum, value) };
}
ZSTD_D_REF_MULTIPLE_DDICTS => {
if !dparam_within_bounds(param, value) {
return ERROR(ZstdErrorCode::ParameterOutOfBound);
}
if unsafe { field::<usize>(view.static_size) } != 0 {
return ERROR(ZstdErrorCode::ParameterUnsupported);
}
unsafe { set_field(view.ref_multiple_ddicts, value) };
}
ZSTD_D_DISABLE_HUFFMAN_ASSEMBLY => {
if !dparam_within_bounds(param, value) {
return ERROR(ZstdErrorCode::ParameterOutOfBound);
}
unsafe { set_field(view.disable_huf_asm, c_int::from(value != 0)) };
}
ZSTD_D_MAX_BLOCK_SIZE => {
if value != 0 && !dparam_within_bounds(param, value) {
return ERROR(ZstdErrorCode::ParameterOutOfBound);
}
unsafe { set_field(view.max_block_size_param, value) };
}
_ => return ERROR(ZstdErrorCode::ParameterUnsupported),
}
0
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_DCtx_reset(dctx: *mut ZSTD_DCtx, reset: c_int) -> usize {
if dctx.is_null() {
return ERROR(ZstdErrorCode::Generic);
}
let view = unsafe { dctx_view(dctx) };
if reset == ZSTD_RESET_SESSION_ONLY || reset == ZSTD_RESET_SESSION_AND_PARAMETERS {
unsafe {
set_field(view.stream_stage, ZDSS_INIT);
set_field(view.no_forward_progress, 0 as c_int);
set_field(view.is_frame_decompression, 1 as c_int);
}
}
if reset == ZSTD_RESET_PARAMETERS || reset == ZSTD_RESET_SESSION_AND_PARAMETERS {
if unsafe { field::<c_int>(view.stream_stage) } != ZDSS_INIT {
return ERROR(ZstdErrorCode::StageWrong);
}
unsafe {
clear_dict(&view);
reset_parameters(&view);
}
}
0
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_sizeof_DStream(dctx: *const ZSTD_DStream) -> usize {
unsafe { ZSTD_sizeof_DCtx(dctx) }
}
unsafe fn decoding_buffer_size_internal(
window_size: u64,
frame_content_size: u64,
block_size_max: usize,
) -> usize {
let block_size = min(
min(window_size, ZSTD_BLOCKSIZE_MAX as u64) as usize,
block_size_max,
);
let needed_ring = window_size
.wrapping_add((block_size as u64).wrapping_mul(2))
.wrapping_add((WILDCOPY_OVERLENGTH as u64).wrapping_mul(2));
let needed = min(frame_content_size, needed_ring);
let result = needed as usize;
if result as u64 != needed {
return ERROR(ZstdErrorCode::FrameParameterWindowTooLarge);
}
result
}
#[no_mangle]
pub extern "C" fn ZSTD_decodingBufferSize_min(window_size: u64, frame_content_size: u64) -> usize {
unsafe { decoding_buffer_size_internal(window_size, frame_content_size, ZSTD_BLOCKSIZE_MAX) }
}
#[no_mangle]
pub extern "C" fn ZSTD_estimateDStreamSize(window_size: usize) -> usize {
let block_size = min(window_size, ZSTD_BLOCKSIZE_MAX);
let out_size = ZSTD_decodingBufferSize_min(window_size as u64, ZSTD_CONTENTSIZE_UNKNOWN);
ZSTD_estimateDCtxSize()
.wrapping_add(block_size)
.wrapping_add(out_size)
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_estimateDStreamSize_fromFrame(
src: *const c_void,
src_size: usize,
) -> usize {
let mut zfh = ZSTD_FrameHeader::default();
let result = unsafe { ZSTD_getFrameHeader(&mut zfh, src, src_size) };
if ERR_isError(result) {
return result;
}
if result != 0 {
return ERROR(ZstdErrorCode::SrcSizeWrong);
}
if zfh.window_size > (1u64 << window_log_max()) {
return ERROR(ZstdErrorCode::FrameParameterWindowTooLarge);
}
ZSTD_estimateDStreamSize(zfh.window_size as usize)
}
#[inline]
unsafe fn dctx_is_overflow(
view: &ZSTD_rustDctxView,
needed_in_size: usize,
needed_out_size: usize,
) -> bool {
let current = unsafe { field::<usize>(view.in_buff_size) }
.wrapping_add(unsafe { field::<usize>(view.out_buff_size) });
let needed = needed_in_size
.wrapping_add(needed_out_size)
.wrapping_mul(ZSTD_WORKSPACETOOLARGE_FACTOR);
current >= needed
}
#[inline]
unsafe fn update_oversized_duration(
view: &ZSTD_rustDctxView,
needed_in_size: usize,
needed_out_size: usize,
) {
let duration = if unsafe { dctx_is_overflow(view, needed_in_size, needed_out_size) } {
unsafe { field::<usize>(view.oversized_duration) }.wrapping_add(1)
} else {
0
};
unsafe { set_field(view.oversized_duration, duration) };
}
#[inline]
unsafe fn oversized_too_long(view: &ZSTD_rustDctxView) -> bool {
(unsafe { field::<usize>(view.oversized_duration) }) >= ZSTD_WORKSPACETOOLARGE_MAXDURATION
}
unsafe fn check_out_buffer(view: &ZSTD_rustDctxView, output: &ZSTD_outBuffer) -> usize {
if unsafe { field::<c_int>(view.out_buffer_mode) } != ZSTD_BM_STABLE
|| unsafe { field::<c_int>(view.stream_stage) } == ZDSS_INIT
{
return 0;
}
let expected = unsafe { field::<ZSTD_outBuffer>(view.expected_out_buffer) };
if expected.dst == output.dst && expected.size == output.size && expected.pos == output.pos {
0
} else {
ERROR(ZstdErrorCode::DstBufferWrong)
}
}
/// Invoke the bufferless state machine from the streaming adapter and translate
/// its output into either the rolling internal buffer or the stable user buffer.
unsafe fn decompress_continue_stream(
dctx: *mut ZSTD_DCtx,
view: &ZSTD_rustDctxView,
op: &mut *mut u8,
oend: *mut u8,
src: *const u8,
src_size: usize,
) -> usize {
let skip = unsafe { is_skip_frame(view) };
if unsafe { field::<c_int>(view.out_buffer_mode) } == ZSTD_BM_BUFFERED {
let out_start = unsafe { field::<usize>(view.out_start) };
let out_size = unsafe { field::<usize>(view.out_buff_size) };
let dst_size = if skip {
0
} else {
out_size.wrapping_sub(out_start)
};
let out_buff = unsafe { get_mut_pointer(view.out_buff) };
let decoded = unsafe {
ZSTD_decompressContinue(
dctx,
out_buff.wrapping_add(out_start).cast(),
dst_size,
src.cast(),
src_size,
)
};
if ERR_isError(decoded) {
return decoded;
}
unsafe {
if decoded == 0 && !skip {
set_field(view.stream_stage, ZDSS_READ);
} else {
set_field(view.out_end, out_start.wrapping_add(decoded));
set_field(view.stream_stage, ZDSS_FLUSH);
}
}
} else {
let dst_size = (oend as usize).wrapping_sub(*op as usize);
let decoded =
unsafe { ZSTD_decompressContinue(dctx, (*op).cast(), dst_size, src.cast(), src_size) };
if ERR_isError(decoded) {
return decoded;
}
if decoded != 0 {
*op = unsafe { (*op).add(decoded) };
}
unsafe { set_field(view.stream_stage, ZDSS_READ) };
}
0
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_decompressStream(
zds: *mut ZSTD_DStream,
output: *mut ZSTD_outBuffer,
input: *mut ZSTD_inBuffer,
) -> usize {
if zds.is_null() || output.is_null() || input.is_null() {
return ERROR(ZstdErrorCode::Generic);
}
let output_ref = unsafe { &mut *output };
let input_ref = unsafe { &mut *input };
if input_ref.pos > input_ref.size {
return ERROR(ZstdErrorCode::SrcSizeWrong);
}
if output_ref.pos > output_ref.size {
return ERROR(ZstdErrorCode::DstSizeTooSmall);
}
let view = unsafe { dctx_view(zds) };
let src = input_ref.src.cast::<u8>();
let istart = src.wrapping_add(input_ref.pos);
let iend = src.wrapping_add(input_ref.size);
let mut ip = istart;
let dst = output_ref.dst.cast::<u8>();
let ostart = dst.wrapping_add(output_ref.pos);
let oend = dst.wrapping_add(output_ref.size);
let mut op = ostart;
let mut some_more_work = true;
let output_check = unsafe { check_out_buffer(&view, output_ref) };
if ERR_isError(output_check) {
return output_check;
}
while some_more_work {
match unsafe { field::<c_int>(view.stream_stage) } {
ZDSS_INIT => {
unsafe {
set_field(view.stream_stage, ZDSS_LOAD_HEADER);
set_field(view.lh_size, 0usize);
set_field(view.in_pos, 0usize);
set_field(view.out_start, 0usize);
set_field(view.out_end, 0usize);
if !view.legacy_version.is_null() {
set_field(view.legacy_version, 0u32);
}
set_field(view.hostage_byte, 0u32);
set_field(view.expected_out_buffer, *output_ref);
}
continue;
}
ZDSS_LOAD_HEADER => {
if !view.legacy_version.is_null()
&& unsafe { field::<u32>(view.legacy_version) } != 0
{
let ddict = unsafe { dctx_ddict(&view) };
let dict = if ddict.is_null() {
ptr::null()
} else {
unsafe { ZSTD_DDict_dictContent(ddict) }
};
let dict_size = if ddict.is_null() {
0
} else {
unsafe { ZSTD_DDict_dictSize(ddict) }
};
return unsafe {
ZSTD_rust_legacy_decompress_stream(zds, output, input, dict, dict_size)
};
}
let lh_size = unsafe { field::<usize>(view.lh_size) };
let header_result = unsafe {
ZSTD_getFrameHeader_advanced(
get_frame_header_ptr(&view),
view.header_buffer.cast(),
lh_size,
field::<c_int>(view.format),
)
};
if unsafe { field::<c_int>(view.ref_multiple_ddicts) }
== ZSTD_RMD_REF_MULTIPLE_DDICTS
&& !unsafe { field::<*mut DDictHashSet>(view.ddict_set) }.is_null()
{
unsafe { select_frame_ddict(&view) };
}
if ERR_isError(header_result) {
let available = (iend as usize).wrapping_sub(istart as usize);
if unsafe { ZSTD_rust_legacy_is(istart.cast(), available) } != 0 {
if unsafe { field::<usize>(view.static_size) } != 0 {
return ERROR(ZstdErrorCode::MemoryAllocation);
}
let ddict = unsafe { get_ddict(&view) };
let dict = if ddict.is_null() {
ptr::null()
} else {
unsafe { ZSTD_DDict_dictContent(ddict) }
};
let dict_size = if ddict.is_null() {
0
} else {
unsafe { ZSTD_DDict_dictSize(ddict) }
};
return unsafe {
ZSTD_rust_legacy_decompress_stream(zds, output, input, dict, dict_size)
};
}
return header_result;
}
if header_result != 0 {
let to_load = header_result - lh_size;
let remaining_input = (iend as usize).wrapping_sub(ip as usize);
if to_load > remaining_input {
if remaining_input != 0 {
unsafe {
copy_bytes(
view.header_buffer.cast::<u8>().add(lh_size),
ip,
remaining_input,
);
set_field(view.lh_size, lh_size + remaining_input);
}
}
input_ref.pos = input_ref.size;
let check = unsafe {
ZSTD_getFrameHeader_advanced(
get_frame_header_ptr(&view),
view.header_buffer.cast(),
field::<usize>(view.lh_size),
field::<c_int>(view.format),
)
};
if ERR_isError(check) {
return check;
}
let minimum = max(
frame_header_min(unsafe { field::<c_int>(view.format) }),
header_result,
);
return minimum
.wrapping_sub(unsafe { field::<usize>(view.lh_size) })
.wrapping_add(ZSTD_BLOCKHEADERSIZE);
}
unsafe {
copy_bytes(view.header_buffer.cast::<u8>().add(lh_size), ip, to_load);
set_field(view.lh_size, header_result);
}
ip = unsafe { ip.add(to_load) };
continue;
}
let params = unsafe { field::<ZSTD_FrameHeader>(view.f_params) };
let available = (iend as usize).wrapping_sub(istart as usize);
if params.frame_content_size != ZSTD_CONTENTSIZE_UNKNOWN
&& params.frame_type != ZSTD_SKIPPABLE_FRAME
&& (oend as usize).wrapping_sub(op as usize)
>= params.frame_content_size as usize
{
let frame_size = unsafe {
find_frame_size_info(istart, available, field::<c_int>(view.format))
.compressed_size
};
if frame_size <= available {
let ddict = unsafe { get_ddict(&view) };
let decoded = unsafe {
ZSTD_decompress_usingDDict(
zds,
op.cast(),
(oend as usize).wrapping_sub(op as usize),
istart.cast(),
frame_size,
ddict,
)
};
if ERR_isError(decoded) {
return decoded;
}
ip = unsafe { istart.add(frame_size) };
if decoded != 0 {
op = unsafe { op.add(decoded) };
}
unsafe {
set_field(view.expected, 0usize);
set_field(view.stream_stage, ZDSS_INIT);
}
some_more_work = false;
continue;
}
}
if unsafe { field::<c_int>(view.out_buffer_mode) } == ZSTD_BM_STABLE
&& params.frame_type != ZSTD_SKIPPABLE_FRAME
&& params.frame_content_size != ZSTD_CONTENTSIZE_UNKNOWN
&& (oend as usize).wrapping_sub(op as usize)
< params.frame_content_size as usize
{
return ERROR(ZstdErrorCode::DstSizeTooSmall);
}
let ddict = unsafe { get_ddict(&view) };
let begin = unsafe { ZSTD_decompressBegin_usingDDict(zds, ddict) };
if ERR_isError(begin) {
return begin;
}
let format = unsafe { field::<c_int>(view.format) };
if format == ZSTD_F_ZSTD1
&& unsafe { MEM_readLE32(view.header_buffer) } & ZSTD_MAGIC_SKIPPABLE_MASK
== ZSTD_MAGIC_SKIPPABLE_START
{
unsafe {
set_field(
view.expected,
MEM_readLE32(
view.header_buffer.cast::<u8>().add(ZSTD_FRAMEIDSIZE).cast(),
) as usize,
);
set_field(view.stage, ZSTDDS_SKIP_FRAME);
}
} else {
let result = unsafe {
decode_frame_header(
&view,
view.header_buffer.cast(),
field::<usize>(view.lh_size),
)
};
if ERR_isError(result) {
return result;
}
unsafe {
set_field(view.expected, ZSTD_BLOCKHEADERSIZE);
set_field(view.stage, ZSTDDS_DECODE_BLOCK_HEADER);
}
}
let mut frame_params = unsafe { field::<ZSTD_FrameHeader>(view.f_params) };
frame_params.window_size =
max(frame_params.window_size, 1u64 << ZSTD_WINDOWLOG_ABSOLUTEMIN);
if frame_params.window_size > unsafe { field::<usize>(view.max_window_size) } as u64
{
return ERROR(ZstdErrorCode::FrameParameterWindowTooLarge);
}
let max_block_size_param = unsafe { field::<c_int>(view.max_block_size_param) };
if max_block_size_param != 0 {
frame_params.block_size_max =
min(frame_params.block_size_max, max_block_size_param as c_uint);
}
unsafe { set_field(view.f_params, frame_params) };
let needed_in_size = max(frame_params.block_size_max as usize, 4);
let needed_out_size =
if unsafe { field::<c_int>(view.out_buffer_mode) } == ZSTD_BM_BUFFERED {
let size = unsafe {
decoding_buffer_size_internal(
frame_params.window_size,
frame_params.frame_content_size,
frame_params.block_size_max as usize,
)
};
if ERR_isError(size) {
return size;
}
size
} else {
0
};
unsafe { update_oversized_duration(&view, needed_in_size, needed_out_size) };
let too_small = unsafe { field::<usize>(view.in_buff_size) } < needed_in_size
|| unsafe { field::<usize>(view.out_buff_size) } < needed_out_size;
let too_large = unsafe { oversized_too_long(&view) };
if too_small || too_large {
let buffer_size = needed_in_size.wrapping_add(needed_out_size);
if unsafe { field::<usize>(view.static_size) } != 0 {
let static_size = unsafe { field::<usize>(view.static_size) };
if buffer_size > static_size.wrapping_sub(view.dctx_size) {
return ERROR(ZstdErrorCode::MemoryAllocation);
}
} else {
unsafe {
ZSTD_rust_custom_free(
get_mut_pointer(view.in_buff).cast(),
dctx_custom_mem(&view),
);
set_field(view.in_buff_size, 0usize);
set_field(view.out_buff_size, 0usize);
}
let allocation =
unsafe { ZSTD_rust_custom_malloc(buffer_size, dctx_custom_mem(&view)) };
if allocation.is_null() {
return ERROR(ZstdErrorCode::MemoryAllocation);
}
unsafe { set_mut_pointer(view.in_buff, allocation.cast()) };
}
let in_buff = unsafe { get_mut_pointer(view.in_buff) };
unsafe {
set_field(view.in_buff_size, needed_in_size);
set_mut_pointer(view.out_buff, in_buff.wrapping_add(needed_in_size));
set_field(view.out_buff_size, needed_out_size);
}
}
unsafe { set_field(view.stream_stage, ZDSS_READ) };
continue;
}
ZDSS_READ => {
let available = (iend as usize).wrapping_sub(ip as usize);
let needed = unsafe { next_src_size_with_input_size(&view, available) };
if needed == 0 {
unsafe { set_field(view.stream_stage, ZDSS_INIT) };
some_more_work = false;
continue;
}
if available >= needed {
let result = unsafe {
decompress_continue_stream(zds, &view, &mut op, oend, ip, needed)
};
if ERR_isError(result) {
return result;
}
ip = unsafe { ip.add(needed) };
continue;
}
if ip == iend {
some_more_work = false;
continue;
}
unsafe { set_field(view.stream_stage, ZDSS_LOAD) };
continue;
}
ZDSS_LOAD => {
let needed = unsafe { field::<usize>(view.expected) };
let in_pos = unsafe { field::<usize>(view.in_pos) };
let to_load = needed.wrapping_sub(in_pos);
let skip = unsafe { is_skip_frame(&view) };
let available = (iend as usize).wrapping_sub(ip as usize);
let loaded = if skip {
min(to_load, available)
} else {
let in_size = unsafe { field::<usize>(view.in_buff_size) };
if to_load > in_size.wrapping_sub(in_pos) {
return ERROR(ZstdErrorCode::CorruptionDetected);
}
unsafe {
limit_copy(
get_mut_pointer(view.in_buff).wrapping_add(in_pos),
to_load,
ip,
available,
)
}
};
if loaded != 0 {
ip = unsafe { ip.add(loaded) };
unsafe { set_field(view.in_pos, in_pos + loaded) };
}
if loaded < to_load {
some_more_work = false;
continue;
}
unsafe { set_field(view.in_pos, 0usize) };
let result = unsafe {
decompress_continue_stream(
zds,
&view,
&mut op,
oend,
get_mut_pointer(view.in_buff).cast(),
needed,
)
};
if ERR_isError(result) {
return result;
}
continue;
}
ZDSS_FLUSH => {
let out_start = unsafe { field::<usize>(view.out_start) };
let out_end = unsafe { field::<usize>(view.out_end) };
let to_flush = out_end.wrapping_sub(out_start);
let flushed = unsafe {
limit_copy(
op,
(oend as usize).wrapping_sub(op as usize),
get_mut_pointer(view.out_buff).wrapping_add(out_start),
to_flush,
)
};
if flushed != 0 {
op = unsafe { op.add(flushed) };
}
let new_out_start = out_start + flushed;
unsafe { set_field(view.out_start, new_out_start) };
if flushed == to_flush {
unsafe {
set_field(view.stream_stage, ZDSS_READ);
let frame_params = field::<ZSTD_FrameHeader>(view.f_params);
if field::<usize>(view.out_buff_size)
< frame_params.frame_content_size as usize
&& new_out_start + frame_params.block_size_max as usize
> field::<usize>(view.out_buff_size)
{
set_field(view.out_start, 0usize);
set_field(view.out_end, 0usize);
}
}
continue;
}
some_more_work = false;
continue;
}
_ => return ERROR(ZstdErrorCode::Generic),
}
}
input_ref.pos = (ip as usize).wrapping_sub(src as usize);
output_ref.pos = (op as usize).wrapping_sub(dst as usize);
unsafe { set_field(view.expected_out_buffer, *output_ref) };
if ip == istart && op == ostart {
let stalled = unsafe { field::<c_int>(view.no_forward_progress) } + 1;
unsafe { set_field(view.no_forward_progress, stalled) };
if stalled >= unsafe { ZSTD_rust_no_forward_progress_max() } {
if op == oend {
return ERROR(ZstdErrorCode::NoForwardProgressDestFull);
}
if ip == iend {
return ERROR(ZstdErrorCode::NoForwardProgressInputEmpty);
}
return ERROR(ZstdErrorCode::Generic);
}
} else {
unsafe { set_field(view.no_forward_progress, 0 as c_int) };
}
let mut hint = unsafe { field::<usize>(view.expected) };
if hint == 0 {
if unsafe { field::<usize>(view.out_end) } == unsafe { field::<usize>(view.out_start) } {
if unsafe { field::<u32>(view.hostage_byte) } != 0 {
if input_ref.pos >= input_ref.size {
unsafe { set_field(view.stream_stage, ZDSS_READ) };
return 1;
}
input_ref.pos += 1;
}
return 0;
}
if unsafe { field::<u32>(view.hostage_byte) } == 0 {
if input_ref.pos == 0 {
return ERROR(ZstdErrorCode::Generic);
}
input_ref.pos -= 1;
unsafe { set_field(view.hostage_byte, 1u32) };
}
return 1;
}
if unsafe { ZSTD_nextInputType(zds) } == ZSTD_NIT_BLOCK {
hint = hint.wrapping_add(ZSTD_BLOCKHEADERSIZE);
}
let in_pos = unsafe { field::<usize>(view.in_pos) };
if in_pos > hint {
return ERROR(ZstdErrorCode::CorruptionDetected);
}
hint - in_pos
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_decompressStream_simpleArgs(
dctx: *mut ZSTD_DCtx,
dst: *mut c_void,
dst_capacity: usize,
dst_pos: *mut usize,
src: *const c_void,
src_size: usize,
src_pos: *mut usize,
) -> usize {
if dst_pos.is_null() || src_pos.is_null() {
return ERROR(ZstdErrorCode::Generic);
}
let mut output = ZSTD_outBuffer {
dst,
size: dst_capacity,
pos: unsafe { dst_pos.read() },
};
let mut input = ZSTD_inBuffer {
src,
size: src_size,
pos: unsafe { src_pos.read() },
};
let result = unsafe { ZSTD_decompressStream(dctx, &mut output, &mut input) };
unsafe {
dst_pos.write(output.pos);
src_pos.write(input.pos);
}
result
}
+1062
View File
@@ -0,0 +1,1062 @@
#![allow(non_camel_case_types)]
#![allow(non_snake_case)]
#![allow(clippy::missing_safety_doc)]
#![allow(clippy::too_many_arguments)]
//! Long distance matching.
//!
//! The C translation unit owns opaque compression-context dispatch and exports
//! the immutable gear table. This module owns gear splitting, LDM table
//! maintenance, raw-sequence generation, and raw-sequence consumption.
use crate::errors::{ERR_isError, ZstdErrorCode, ERROR};
use crate::mem::{MEM_64bits, MEM_isLittleEndian, MEM_read16, MEM_read32, MEM_readST};
use crate::xxhash::XXH64;
use std::ffi::c_void;
use std::mem::size_of;
use std::os::raw::c_int;
const LDM_BATCH_SIZE: usize = 64;
const LDM_BUCKET_SIZE_LOG: u32 = 4;
const LDM_MIN_MATCH_LENGTH: u32 = 64;
const HASH_READ_SIZE: usize = 8;
const ZSTD_REP_NUM: usize = 3;
const ZSTD_WINDOW_START_INDEX: u32 = 2;
#[repr(C)]
#[derive(Clone, Copy)]
struct LdmEntry {
offset: u32,
checksum: u32,
}
#[repr(C)]
#[derive(Clone, Copy)]
struct RawSeq {
offset: u32,
lit_length: u32,
match_length: u32,
}
#[repr(C)]
struct RawSeqStore {
seq: *mut RawSeq,
pos: usize,
pos_in_sequence: usize,
size: usize,
capacity: usize,
}
#[repr(C)]
struct LdmParams {
enable_ldm: c_int,
hash_log: u32,
bucket_size_log: u32,
min_match_length: u32,
hash_rate_log: u32,
window_log: u32,
}
#[repr(C)]
struct LdmWindow {
next_src: *const u8,
base: *const u8,
dict_base: *const u8,
dict_limit: u32,
low_limit: u32,
nb_overflow_corrections: u32,
}
#[derive(Clone, Copy)]
struct RollingHashState {
rolling: u64,
stop_mask: u64,
}
#[derive(Clone, Copy)]
struct MatchCandidate {
split: *const u8,
hash: u32,
checksum: u32,
bucket: *mut LdmEntry,
}
const EMPTY_CANDIDATE: MatchCandidate = MatchCandidate {
split: std::ptr::null(),
hash: 0,
checksum: 0,
bucket: std::ptr::null_mut(),
};
unsafe extern "C" {
fn ZSTD_ldm_rust_gearTable() -> *const u64;
fn ZSTD_ldm_rust_prepareBlock(context: *mut c_void, anchor: *const c_void);
fn ZSTD_ldm_rust_compressLiterals(
context: *mut c_void,
seq_store: *mut c_void,
reps: *mut u32,
src: *const c_void,
src_size: usize,
) -> usize;
fn ZSTD_ldm_rust_storeSeq(
seq_store: *mut c_void,
lit_length: usize,
literals: *const c_void,
lit_limit: *const c_void,
off_base: u32,
match_length: usize,
);
fn ZSTD_ldm_rust_setLdmSeqStore(context: *mut c_void, raw_seq_store: *const c_void);
}
#[inline]
fn ptr_lt(left: *const u8, right: *const u8) -> bool {
(left as usize) < (right as usize)
}
#[inline]
fn ptr_gt(left: *const u8, right: *const u8) -> bool {
(left as usize) > (right as usize)
}
#[inline]
unsafe fn index_from(base: *const u8, ptr: *const u8) -> u32 {
unsafe { ptr.offset_from(base) as u32 }
}
#[inline]
fn common_bytes(word: usize) -> usize {
let zeros = if MEM_isLittleEndian() {
word.trailing_zeros()
} else {
word.leading_zeros()
};
(zeros / 8) as usize
}
unsafe fn count(mut input: *const u8, mut matched: *const u8, input_limit: *const u8) -> usize {
let input_start = input;
let word_size = size_of::<usize>();
while unsafe { input_limit.offset_from(input) as usize } >= word_size {
let diff =
unsafe { MEM_readST(matched.cast::<c_void>()) ^ MEM_readST(input.cast::<c_void>()) };
if diff != 0 {
return unsafe { input.offset_from(input_start) as usize } + common_bytes(diff);
}
input = input.wrapping_add(word_size);
matched = matched.wrapping_add(word_size);
}
if MEM_64bits()
&& unsafe { input_limit.offset_from(input) as usize } >= 4
&& unsafe { MEM_read32(matched.cast::<c_void>()) == MEM_read32(input.cast::<c_void>()) }
{
input = input.wrapping_add(4);
matched = matched.wrapping_add(4);
}
if unsafe { input_limit.offset_from(input) as usize } >= 2
&& unsafe { MEM_read16(matched.cast::<c_void>()) == MEM_read16(input.cast::<c_void>()) }
{
input = input.wrapping_add(2);
matched = matched.wrapping_add(2);
}
if ptr_lt(input, input_limit) && unsafe { *input == *matched } {
input = input.wrapping_add(1);
}
unsafe { input.offset_from(input_start) as usize }
}
unsafe fn count_2segments(
input: *const u8,
matched: *const u8,
input_end: *const u8,
match_end: *const u8,
input_start: *const u8,
) -> usize {
let match_remaining = unsafe { match_end.offset_from(matched) as usize };
let input_remaining = unsafe { input_end.offset_from(input) as usize };
let first_end = input.wrapping_add(match_remaining.min(input_remaining));
let first_count = unsafe { count(input, matched, first_end) };
if matched.wrapping_add(first_count) != match_end {
return first_count;
}
first_count + unsafe { count(input.wrapping_add(first_count), input_start, input_end) }
}
#[inline]
fn bounded(lower: u32, value: u32, upper: u32) -> u32 {
value.max(lower).min(upper)
}
#[inline]
unsafe fn ldm_bucket(hash_table: *mut LdmEntry, hash: u32, bucket_size_log: u32) -> *mut LdmEntry {
unsafe { hash_table.add((hash as usize) << bucket_size_log) }
}
unsafe fn ldm_insert_entry(
hash_table: *mut LdmEntry,
bucket_offsets: *mut u8,
hash: u32,
entry: LdmEntry,
bucket_size_log: u32,
) {
let offset = unsafe { *bucket_offsets.add(hash as usize) };
let bucket = unsafe { ldm_bucket(hash_table, hash, bucket_size_log) };
unsafe { *bucket.add(offset as usize) = entry };
unsafe {
*bucket_offsets.add(hash as usize) =
offset.wrapping_add(1) & ((1u32.wrapping_shl(bucket_size_log)).wrapping_sub(1) as u8)
};
}
fn gear_init(params: &LdmParams) -> RollingHashState {
let max_bits_in_mask = params.min_match_length.min(64);
let hash_rate_log = params.hash_rate_log;
let stop_mask = if hash_rate_log > 0 && hash_rate_log <= max_bits_in_mask {
((1u64 << hash_rate_log) - 1) << (max_bits_in_mask - hash_rate_log)
} else {
(1u64 << hash_rate_log) - 1
};
RollingHashState {
/* C assigns `~(U32)0`, which is a 32-bit all-ones value. */
rolling: u32::MAX as u64,
stop_mask,
}
}
/*
* This intentionally leaves `state.rolling` unchanged: the reference C
* routine computes a local hash but never writes it back to the state.
*/
unsafe fn gear_reset(_state: &mut RollingHashState, _data: *const u8, _min_match_length: usize) {}
unsafe fn gear_feed(
state: &mut RollingHashState,
gear_table: *const u64,
data: *const u8,
size: usize,
splits: &mut [usize; LDM_BATCH_SIZE],
num_splits: &mut usize,
) -> usize {
let mut hash = state.rolling;
let mut n = 0usize;
while n + 3 < size {
for _ in 0..4 {
hash = hash
.wrapping_shl(1)
.wrapping_add(unsafe { *gear_table.add(*data.add(n) as usize) });
n += 1;
if (hash & state.stop_mask) == 0 {
splits[*num_splits] = n;
*num_splits += 1;
if *num_splits == LDM_BATCH_SIZE {
state.rolling = hash;
return n;
}
}
}
}
while n < size {
hash = hash
.wrapping_shl(1)
.wrapping_add(unsafe { *gear_table.add(*data.add(n) as usize) });
n += 1;
if (hash & state.stop_mask) == 0 {
splits[*num_splits] = n;
*num_splits += 1;
if *num_splits == LDM_BATCH_SIZE {
state.rolling = hash;
return n;
}
}
}
state.rolling = hash;
n
}
unsafe fn count_backwards_match(
mut input: *const u8,
anchor: *const u8,
mut matched: *const u8,
match_base: *const u8,
) -> usize {
let mut match_length = 0usize;
while ptr_gt(input, anchor)
&& ptr_gt(matched, match_base)
&& unsafe { *input.wrapping_sub(1) == *matched.wrapping_sub(1) }
{
input = input.wrapping_sub(1);
matched = matched.wrapping_sub(1);
match_length += 1;
}
match_length
}
unsafe fn count_backwards_match_2segments(
input: *const u8,
anchor: *const u8,
matched: *const u8,
match_base: *const u8,
ext_dict_start: *const u8,
ext_dict_end: *const u8,
) -> usize {
let match_length = unsafe { count_backwards_match(input, anchor, matched, match_base) };
if matched.wrapping_sub(match_length) != match_base || match_base == ext_dict_start {
return match_length;
}
match_length
+ unsafe {
count_backwards_match(
input.wrapping_sub(match_length),
anchor,
ext_dict_end,
ext_dict_start,
)
}
}
fn window_has_ext_dict(window: &LdmWindow) -> bool {
window.low_limit < window.dict_limit
}
fn window_can_overflow_correct(
window: &LdmWindow,
cycle_log: u32,
max_dist: u32,
loaded_dict_end: u32,
src: *const u8,
) -> bool {
let cycle_size = 1u32.wrapping_shl(cycle_log);
let current = unsafe { index_from(window.base, src) };
let min_index = cycle_size
.wrapping_add(max_dist.max(cycle_size))
.wrapping_add(ZSTD_WINDOW_START_INDEX);
let adjustment = window.nb_overflow_corrections.wrapping_add(1);
let adjusted = min_index.wrapping_mul(adjustment).max(min_index);
let index_large_enough = current > adjusted;
let dictionary_invalidated = current > max_dist.wrapping_add(loaded_dict_end);
index_large_enough && dictionary_invalidated
}
fn window_needs_overflow_correction(
window: &LdmWindow,
cycle_log: u32,
max_dist: u32,
loaded_dict_end: u32,
src: *const u8,
src_end: *const u8,
overflow_correct_frequently: bool,
) -> bool {
if overflow_correct_frequently
&& window_can_overflow_correct(window, cycle_log, max_dist, loaded_dict_end, src)
{
return true;
}
let current = unsafe { index_from(window.base, src_end) };
let current_max = if size_of::<usize>() == 8 {
3500u32 * (1 << 20)
} else {
2000u32 * (1 << 20)
};
current > current_max
}
unsafe fn window_correct_overflow(
window: &mut LdmWindow,
cycle_log: u32,
max_dist: u32,
src: *const u8,
) -> u32 {
let cycle_size = 1u32.wrapping_shl(cycle_log);
let cycle_mask = cycle_size.wrapping_sub(1);
let current = unsafe { index_from(window.base, src) };
let current_cycle = current & cycle_mask;
let current_cycle_correction = if current_cycle < ZSTD_WINDOW_START_INDEX {
cycle_size.max(ZSTD_WINDOW_START_INDEX)
} else {
0
};
let new_current = current_cycle
.wrapping_add(current_cycle_correction)
.wrapping_add(max_dist.max(cycle_size));
let correction = current.wrapping_sub(new_current);
window.base = window.base.wrapping_add(correction as usize);
window.dict_base = window.dict_base.wrapping_add(correction as usize);
if window.low_limit < correction.wrapping_add(ZSTD_WINDOW_START_INDEX) {
window.low_limit = ZSTD_WINDOW_START_INDEX;
} else {
window.low_limit = window.low_limit.wrapping_sub(correction);
}
if window.dict_limit < correction.wrapping_add(ZSTD_WINDOW_START_INDEX) {
window.dict_limit = ZSTD_WINDOW_START_INDEX;
} else {
window.dict_limit = window.dict_limit.wrapping_sub(correction);
}
window.nb_overflow_corrections = window.nb_overflow_corrections.wrapping_add(1);
correction
}
fn window_enforce_max_dist(
window: &mut LdmWindow,
block_end: *const u8,
max_dist: u32,
loaded_dict_end: &mut u32,
) {
let block_end_index = unsafe { index_from(window.base, block_end) };
if block_end_index > max_dist.wrapping_add(*loaded_dict_end) {
let new_low_limit = block_end_index.wrapping_sub(max_dist);
if window.low_limit < new_low_limit {
window.low_limit = new_low_limit;
}
if window.dict_limit < window.low_limit {
window.dict_limit = window.low_limit;
}
*loaded_dict_end = 0;
}
}
unsafe fn reduce_table(table: *mut LdmEntry, size: u32, reducer_value: u32) {
for index in 0..size as usize {
let entry = unsafe { table.add(index) };
if unsafe { (*entry).offset < reducer_value } {
unsafe { (*entry).offset = 0 };
} else {
unsafe { (*entry).offset = (*entry).offset.wrapping_sub(reducer_value) };
}
}
}
unsafe fn generate_sequences_internal(
hash_table: *mut LdmEntry,
bucket_offsets: *mut u8,
window: &LdmWindow,
raw_seq_store: *mut RawSeqStore,
params: &LdmParams,
gear_table: *const u64,
src: *const u8,
src_size: usize,
) -> usize {
let ext_dict = window_has_ext_dict(window);
let min_match_length = params.min_match_length as usize;
let entries_per_bucket = 1usize << params.bucket_size_log;
let hbits = params.hash_log.wrapping_sub(params.bucket_size_log);
let dict_limit = window.dict_limit;
let lowest_index = if ext_dict {
window.low_limit
} else {
dict_limit
};
let base = window.base;
let dict_base = window.dict_base;
let dict_start = dict_base.wrapping_add(lowest_index as usize);
let dict_end = dict_base.wrapping_add(dict_limit as usize);
let low_prefix_ptr = base.wrapping_add(dict_limit as usize);
let iend = src.wrapping_add(src_size);
let ilimit = iend.wrapping_sub(HASH_READ_SIZE);
let mut anchor = src;
let mut ip = src;
if src_size < min_match_length {
return src_size;
}
let mut hash_state = gear_init(params);
unsafe { gear_reset(&mut hash_state, ip, min_match_length) };
ip = ip.wrapping_add(min_match_length);
while ptr_lt(ip, ilimit) {
let mut splits = [0usize; LDM_BATCH_SIZE];
let mut num_splits = 0usize;
let hashed = unsafe {
gear_feed(
&mut hash_state,
gear_table,
ip,
ilimit.offset_from(ip) as usize,
&mut splits,
&mut num_splits,
)
};
let mut candidates = [EMPTY_CANDIDATE; LDM_BATCH_SIZE];
for index in 0..num_splits {
let split = ip
.wrapping_add(splits[index])
.wrapping_sub(min_match_length);
let xxhash = unsafe { XXH64(split.cast::<c_void>(), min_match_length, 0) };
let hash = (xxhash as u32) & ((1u32 << hbits) - 1);
candidates[index] = MatchCandidate {
split,
hash,
checksum: (xxhash >> 32) as u32,
bucket: unsafe { ldm_bucket(hash_table, hash, params.bucket_size_log) },
};
}
for candidate in candidates.iter().take(num_splits) {
let split = candidate.split;
let new_entry = LdmEntry {
offset: unsafe { index_from(base, split) },
checksum: candidate.checksum,
};
if ptr_lt(split, anchor) {
unsafe {
ldm_insert_entry(
hash_table,
bucket_offsets,
candidate.hash,
new_entry,
params.bucket_size_log,
)
};
continue;
}
let mut forward_match_length = 0usize;
let mut backward_match_length = 0usize;
let mut best_match_length = 0usize;
let mut best_offset = None;
for entry_index in 0..entries_per_bucket {
let entry = unsafe { *candidate.bucket.add(entry_index) };
if entry.checksum != candidate.checksum || entry.offset <= lowest_index {
continue;
}
let (current_forward, current_backward) = if ext_dict {
let match_base = if entry.offset < dict_limit {
dict_base
} else {
base
};
let matched = match_base.wrapping_add(entry.offset as usize);
let match_end = if entry.offset < dict_limit {
dict_end
} else {
iend
};
let low_match = if entry.offset < dict_limit {
dict_start
} else {
low_prefix_ptr
};
let forward =
unsafe { count_2segments(split, matched, iend, match_end, low_prefix_ptr) };
if forward < min_match_length {
continue;
}
let backward = unsafe {
count_backwards_match_2segments(
split, anchor, matched, low_match, dict_start, dict_end,
)
};
(forward, backward)
} else {
let matched = base.wrapping_add(entry.offset as usize);
let forward = unsafe { count(split, matched, iend) };
if forward < min_match_length {
continue;
}
let backward =
unsafe { count_backwards_match(split, anchor, matched, low_prefix_ptr) };
(forward, backward)
};
let total = current_forward + current_backward;
if total > best_match_length {
best_match_length = total;
forward_match_length = current_forward;
backward_match_length = current_backward;
best_offset = Some(entry.offset);
}
}
let Some(best_offset) = best_offset else {
unsafe {
ldm_insert_entry(
hash_table,
bucket_offsets,
candidate.hash,
new_entry,
params.bucket_size_log,
)
};
continue;
};
let raw_seq_store = unsafe { &mut *raw_seq_store };
if raw_seq_store.size == raw_seq_store.capacity {
return ERROR(ZstdErrorCode::DstSizeTooSmall);
}
let sequence = unsafe { raw_seq_store.seq.add(raw_seq_store.size) };
unsafe {
(*sequence).lit_length = split
.wrapping_sub(backward_match_length)
.offset_from(anchor) as u32;
(*sequence).match_length = (forward_match_length + backward_match_length) as u32;
(*sequence).offset = index_from(base, split).wrapping_sub(best_offset);
}
raw_seq_store.size += 1;
unsafe {
ldm_insert_entry(
hash_table,
bucket_offsets,
candidate.hash,
new_entry,
params.bucket_size_log,
)
};
anchor = split.wrapping_add(forward_match_length);
if ptr_gt(anchor, ip.wrapping_add(hashed)) {
unsafe {
gear_reset(
&mut hash_state,
anchor.wrapping_sub(min_match_length),
min_match_length,
)
};
ip = anchor.wrapping_sub(hashed);
break;
}
}
ip = ip.wrapping_add(hashed);
}
unsafe { iend.offset_from(anchor) as usize }
}
/// Rust implementation called by the C ABI wrapper for parameter adjustment.
#[no_mangle]
pub unsafe extern "C" fn ZSTD_rust_ldm_adjustParameters(
params: *mut c_void,
window_log: u32,
strategy: c_int,
hash_log_max: u32,
bucket_size_log_max: u32,
btultra: c_int,
) {
let params = unsafe { &mut *params.cast::<LdmParams>() };
params.window_log = window_log;
if params.hash_rate_log == 0 {
if params.hash_log > 0 {
if params.window_log > params.hash_log {
params.hash_rate_log = params.window_log - params.hash_log;
}
} else {
params.hash_rate_log = 7u32.wrapping_sub((strategy / 3) as u32);
}
}
if params.hash_log == 0 {
params.hash_log = bounded(
6,
params.window_log.wrapping_sub(params.hash_rate_log),
hash_log_max,
);
}
if params.min_match_length == 0 {
params.min_match_length = LDM_MIN_MATCH_LENGTH;
if strategy >= btultra {
params.min_match_length /= 2;
}
}
if params.bucket_size_log == 0 {
params.bucket_size_log = bounded(LDM_BUCKET_SIZE_LOG, strategy as u32, bucket_size_log_max);
}
params.bucket_size_log = params.bucket_size_log.min(params.hash_log);
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_rust_ldm_getTableSize(
params: *const c_void,
enable_ldm: c_int,
redzone_size: usize,
) -> usize {
let params = unsafe { &*params.cast::<LdmParams>() };
let hash_size = 1usize << params.hash_log;
let bucket_log = params.bucket_size_log.min(params.hash_log);
let bucket_size = 1usize << (params.hash_log - bucket_log);
let alloc_size = |size: usize| {
if size == 0 {
0
} else {
size + 2 * redzone_size
}
};
if enable_ldm != 0 {
alloc_size(bucket_size) + alloc_size(hash_size * size_of::<LdmEntry>())
} else {
0
}
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_rust_ldm_getMaxNbSeq(
params: *const c_void,
enable_ldm: c_int,
max_chunk_size: usize,
) -> usize {
let params = unsafe { &*params.cast::<LdmParams>() };
if enable_ldm != 0 {
max_chunk_size / params.min_match_length as usize
} else {
0
}
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_rust_ldm_fillHashTable(
hash_table: *mut c_void,
bucket_offsets: *mut u8,
base: *const u8,
mut input: *const u8,
input_end: *const u8,
params: *const c_void,
) {
let hash_table = hash_table.cast::<LdmEntry>();
let params = unsafe { &*params.cast::<LdmParams>() };
let min_match_length = params.min_match_length as usize;
let hbits = params.hash_log.wrapping_sub(params.bucket_size_log);
let input_start = input;
let gear_table = unsafe { ZSTD_ldm_rust_gearTable() };
let mut hash_state = gear_init(params);
while ptr_lt(input, input_end) {
let mut splits = [0usize; LDM_BATCH_SIZE];
let mut num_splits = 0usize;
let hashed = unsafe {
gear_feed(
&mut hash_state,
gear_table,
input,
input_end.offset_from(input) as usize,
&mut splits,
&mut num_splits,
)
};
for split_index in splits.iter().take(num_splits) {
if input.wrapping_add(*split_index) >= input_start.wrapping_add(min_match_length) {
let split = input
.wrapping_add(*split_index)
.wrapping_sub(min_match_length);
let xxhash = unsafe { XXH64(split.cast::<c_void>(), min_match_length, 0) };
let hash = (xxhash as u32) & ((1u32 << hbits) - 1);
unsafe {
ldm_insert_entry(
hash_table,
bucket_offsets,
hash,
LdmEntry {
offset: index_from(base, split),
checksum: (xxhash >> 32) as u32,
},
params.bucket_size_log,
)
};
}
}
input = input.wrapping_add(hashed);
}
}
/// Rust implementation called by the C ABI wrapper for LDM sequence generation.
#[no_mangle]
pub unsafe extern "C" fn ZSTD_rust_ldm_generateSequences(
hash_table: *mut c_void,
bucket_offsets: *mut u8,
window: *mut c_void,
loaded_dict_end: *mut u32,
raw_seq_store: *mut c_void,
params: *const c_void,
src: *const c_void,
src_size: usize,
overflow_correct_frequently: c_int,
) -> usize {
let hash_table = hash_table.cast::<LdmEntry>();
let params = unsafe { &*params.cast::<LdmParams>() };
let window = unsafe { &mut *window.cast::<LdmWindow>() };
let raw_seq_store = raw_seq_store.cast::<RawSeqStore>();
let max_dist = 1u32.wrapping_shl(params.window_log);
let input = src.cast::<u8>();
let input_end = input.wrapping_add(src_size);
const MAX_CHUNK_SIZE: usize = 1 << 20;
let num_chunks =
src_size / MAX_CHUNK_SIZE + usize::from(!src_size.is_multiple_of(MAX_CHUNK_SIZE));
let mut leftover_size = 0usize;
let gear_table = unsafe { ZSTD_ldm_rust_gearTable() };
for chunk in 0..num_chunks {
let chunk_start = input.wrapping_add(chunk * MAX_CHUNK_SIZE);
let remaining = unsafe { input_end.offset_from(chunk_start) as usize };
let chunk_end = if remaining < MAX_CHUNK_SIZE {
input_end
} else {
chunk_start.wrapping_add(MAX_CHUNK_SIZE)
};
let chunk_size = unsafe { chunk_end.offset_from(chunk_start) as usize };
let raw_seq_store_ref = unsafe { &mut *raw_seq_store };
if raw_seq_store_ref.size >= raw_seq_store_ref.capacity {
break;
}
let previous_size = raw_seq_store_ref.size;
let loaded = unsafe { &mut *loaded_dict_end };
if window_needs_overflow_correction(
window,
0,
max_dist,
*loaded,
chunk_start,
chunk_end,
overflow_correct_frequently != 0,
) {
let hash_size = 1u32.wrapping_shl(params.hash_log);
let correction = unsafe { window_correct_overflow(window, 0, max_dist, chunk_start) };
unsafe { reduce_table(hash_table, hash_size, correction) };
*loaded = 0;
}
window_enforce_max_dist(window, chunk_end, max_dist, loaded);
let leftover = unsafe {
generate_sequences_internal(
hash_table,
bucket_offsets,
window,
raw_seq_store,
params,
gear_table,
chunk_start,
chunk_size,
)
};
if ERR_isError(leftover) {
return leftover;
}
let raw_seq_store_ref = unsafe { &mut *raw_seq_store };
if previous_size < raw_seq_store_ref.size {
unsafe {
(*raw_seq_store_ref.seq.add(previous_size)).lit_length =
(*raw_seq_store_ref.seq.add(previous_size))
.lit_length
.wrapping_add(leftover_size as u32)
};
leftover_size = leftover;
} else {
leftover_size += chunk_size;
}
}
0
}
unsafe fn skip_sequences(raw_seq_store: *mut RawSeqStore, mut src_size: usize, min_match: u32) {
let raw_seq_store = unsafe { &mut *raw_seq_store };
while src_size > 0 && raw_seq_store.pos < raw_seq_store.size {
let sequence = unsafe { raw_seq_store.seq.add(raw_seq_store.pos) };
if src_size <= unsafe { (*sequence).lit_length as usize } {
unsafe {
(*sequence).lit_length = (*sequence).lit_length.wrapping_sub(src_size as u32)
};
return;
}
src_size -= unsafe { (*sequence).lit_length as usize };
unsafe { (*sequence).lit_length = 0 };
if src_size < unsafe { (*sequence).match_length as usize } {
unsafe {
(*sequence).match_length = (*sequence).match_length.wrapping_sub(src_size as u32)
};
if unsafe { (*sequence).match_length < min_match } {
if raw_seq_store.pos + 1 < raw_seq_store.size {
unsafe {
(*sequence.add(1)).lit_length = (*sequence.add(1))
.lit_length
.wrapping_add((*sequence).match_length)
};
}
raw_seq_store.pos += 1;
}
return;
}
src_size -= unsafe { (*sequence).match_length as usize };
unsafe { (*sequence).match_length = 0 };
raw_seq_store.pos += 1;
}
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_rust_ldm_skipSequences(
raw_seq_store: *mut c_void,
src_size: usize,
min_match: u32,
) {
unsafe { skip_sequences(raw_seq_store.cast::<RawSeqStore>(), src_size, min_match) };
}
#[no_mangle]
pub unsafe extern "C" fn ZSTD_rust_ldm_skipRawSeqStoreBytes(
raw_seq_store: *mut c_void,
nb_bytes: usize,
) {
let raw_seq_store = unsafe { &mut *raw_seq_store.cast::<RawSeqStore>() };
let mut current_position = raw_seq_store.pos_in_sequence.wrapping_add(nb_bytes) as u32;
while current_position != 0 && raw_seq_store.pos < raw_seq_store.size {
let sequence = unsafe { *raw_seq_store.seq.add(raw_seq_store.pos) };
if current_position >= sequence.lit_length.wrapping_add(sequence.match_length) {
current_position = current_position
.wrapping_sub(sequence.lit_length)
.wrapping_sub(sequence.match_length);
raw_seq_store.pos += 1;
} else {
raw_seq_store.pos_in_sequence = current_position as usize;
break;
}
}
if current_position == 0 || raw_seq_store.pos == raw_seq_store.size {
raw_seq_store.pos_in_sequence = 0;
}
}
unsafe fn maybe_split_sequence(
raw_seq_store: *mut RawSeqStore,
remaining: u32,
min_match: u32,
) -> RawSeq {
let raw_seq_store_ref = unsafe { &mut *raw_seq_store };
let mut sequence = unsafe { *raw_seq_store_ref.seq.add(raw_seq_store_ref.pos) };
if remaining >= sequence.lit_length.wrapping_add(sequence.match_length) {
raw_seq_store_ref.pos += 1;
return sequence;
}
if remaining <= sequence.lit_length {
sequence.offset = 0;
} else {
sequence.match_length = remaining.wrapping_sub(sequence.lit_length);
if sequence.match_length < min_match {
sequence.offset = 0;
}
}
unsafe { skip_sequences(raw_seq_store, remaining as usize, min_match) };
sequence
}
/// Rust implementation called by the C ABI wrapper for LDM block integration.
#[no_mangle]
pub unsafe extern "C" fn ZSTD_rust_ldm_blockCompress(
raw_seq_store: *mut c_void,
block_context: *mut c_void,
seq_store: *mut c_void,
reps: *mut u32,
src: *const c_void,
src_size: usize,
min_match: u32,
use_optimal_parser: c_int,
) -> usize {
let raw_seq_store = raw_seq_store.cast::<RawSeqStore>();
let input = src.cast::<u8>();
let input_end = input.wrapping_add(src_size);
if use_optimal_parser != 0 {
unsafe { ZSTD_ldm_rust_setLdmSeqStore(block_context, raw_seq_store.cast::<c_void>()) };
let last_literals = unsafe {
ZSTD_ldm_rust_compressLiterals(block_context, seq_store, reps, src, src_size)
};
unsafe { ZSTD_rust_ldm_skipRawSeqStoreBytes(raw_seq_store.cast::<c_void>(), src_size) };
return last_literals;
}
let mut input_position = input;
while unsafe { (*raw_seq_store).pos < (*raw_seq_store).size }
&& ptr_lt(input_position, input_end)
{
let sequence = unsafe {
maybe_split_sequence(
raw_seq_store,
input_end.offset_from(input_position) as u32,
min_match,
)
};
if sequence.offset == 0 {
break;
}
unsafe { ZSTD_ldm_rust_prepareBlock(block_context, input_position.cast::<c_void>()) };
let new_lit_length = unsafe {
ZSTD_ldm_rust_compressLiterals(
block_context,
seq_store,
reps,
input_position.cast::<c_void>(),
sequence.lit_length as usize,
)
};
input_position = input_position.wrapping_add(sequence.lit_length as usize);
unsafe {
*reps.add(2) = *reps.add(1);
*reps.add(1) = *reps;
*reps = sequence.offset;
ZSTD_ldm_rust_storeSeq(
seq_store,
new_lit_length,
input_position.wrapping_sub(new_lit_length).cast::<c_void>(),
input_end.cast::<c_void>(),
sequence.offset.wrapping_add(ZSTD_REP_NUM as u32),
sequence.match_length as usize,
);
}
input_position = input_position.wrapping_add(sequence.match_length as usize);
}
unsafe { ZSTD_ldm_rust_prepareBlock(block_context, input_position.cast::<c_void>()) };
unsafe {
ZSTD_ldm_rust_compressLiterals(
block_context,
seq_store,
reps,
input_position.cast::<c_void>(),
input_end.offset_from(input_position) as usize,
)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parameter_defaults_follow_the_c_rules() {
let mut params = LdmParams {
enable_ldm: 1,
hash_log: 0,
bucket_size_log: 0,
min_match_length: 0,
hash_rate_log: 0,
window_log: 0,
};
unsafe {
ZSTD_rust_ldm_adjustParameters(
(&mut params as *mut LdmParams).cast::<c_void>(),
20,
3,
30,
8,
8,
)
};
assert_eq!(params.window_log, 20);
assert_eq!(params.hash_rate_log, 6);
assert_eq!(params.hash_log, 14);
assert_eq!(params.bucket_size_log, 4);
assert_eq!(params.min_match_length, 64);
}
#[test]
fn raw_sequence_skipping_merges_short_tail_matches() {
let mut sequences = [
RawSeq {
offset: 8,
lit_length: 2,
match_length: 10,
},
RawSeq {
offset: 9,
lit_length: 1,
match_length: 12,
},
];
let mut store = RawSeqStore {
seq: sequences.as_mut_ptr(),
pos: 0,
pos_in_sequence: 0,
size: sequences.len(),
capacity: sequences.len(),
};
unsafe { skip_sequences(&mut store, 10, 4) };
assert_eq!(store.pos, 1);
assert_eq!(sequences[1].lit_length, 3);
}
}