feat(cli): port dictionary-training dispatch to Rust

Add Rust parsing and dispatch for the default, Cover, FastCover, and legacy
dictionary-training modes, including their parameter validation, output
defaults, and help text. Keep file loading behind the existing narrow dibio
bridge while the dictionary algorithms are supplied by the Rust library.

Test Plan:
- rustfmt +nightly --check --edition 2021 rust/src/zstd_cli.rs
- RUSTC_WRAPPER= CARGO_BUILD_RUSTC_WRAPPER= cargo test --manifest-path rust/cli/Cargo.toml
- RUSTC_WRAPPER= CARGO_BUILD_RUSTC_WRAPPER= cargo clippy --manifest-path rust/cli/Cargo.toml --all-targets -- -D warnings
- git diff --cached --check
This commit is contained in:
2026-07-12 10:44:57 +02:00
parent 393b2bec45
commit 49454dada9
+410 -15
View File
@@ -22,7 +22,8 @@
//! byte-identical with the C CLI.
//!
//! Remaining C-only CLI boundaries are called out in `unsupported()` below:
//! dictionary training, tracing, and alternate-format selection.
//! tracing and alternate-format selection. Dictionary training uses the
//! temporary `DiB_trainFromFiles` bridge while its algorithms are Rust-owned.
use std::env;
use std::ffi::{CStr, CString, OsStr, OsString};
@@ -43,6 +44,11 @@ const DEFAULT_MAX_CLEVEL: i32 = 19;
const DEFAULT_BENCH_NB_SECONDS: u32 = 3;
const DEFAULT_MEM_LIMIT: u32 = 1 << 27;
const DEFAULT_LONG_WINDOW_LOG: u32 = 27;
const DEFAULT_DICT_NAME: &str = "dictionary";
const DEFAULT_MAX_DICT_SIZE: usize = 110 << 10;
const DEFAULT_DICT_SELECTIVITY: u32 = 9;
const DEFAULT_SHRINK_DICT_REGRESSION: u32 = 1;
const DEFAULT_FASTCOVER_ACCEL: u32 = 1;
const MAX_FAST_ACCELERATION: i32 = 128 << 10;
const STDIN_MARK: &str = "/*stdin*\\";
const STDOUT_MARK: &str = "/*stdout*\\";
@@ -83,6 +89,49 @@ struct ZSTD_compressionParameters {
strategy: c_int,
}
#[repr(C)]
#[derive(Clone, Copy, Debug, Default, PartialEq)]
struct ZDICT_params_t {
compressionLevel: c_int,
notificationLevel: c_uint,
dictID: c_uint,
}
#[repr(C)]
#[derive(Clone, Copy, Debug, Default, PartialEq)]
struct ZDICT_cover_params_t {
k: c_uint,
d: c_uint,
steps: c_uint,
nbThreads: c_uint,
splitPoint: f64,
shrinkDict: c_uint,
shrinkDictMaxRegression: c_uint,
zParams: ZDICT_params_t,
}
#[repr(C)]
#[derive(Clone, Copy, Debug, Default, PartialEq)]
struct ZDICT_fastCover_params_t {
k: c_uint,
d: c_uint,
f: c_uint,
steps: c_uint,
nbThreads: c_uint,
splitPoint: f64,
accel: c_uint,
shrinkDict: c_uint,
shrinkDictMaxRegression: c_uint,
zParams: ZDICT_params_t,
}
#[repr(C)]
#[derive(Clone, Copy, Debug, Default, PartialEq)]
struct ZDICT_legacy_params_t {
selectivityLevel: c_uint,
zParams: ZDICT_params_t,
}
/// Mirror of the `FileNamesTable` in `programs/util.h`; the tables returned
/// by the `UTIL_*FNT` helpers are only read and released here, never resized.
#[repr(C)]
@@ -227,6 +276,21 @@ unsafe extern "C" {
block_size: usize,
nb_workers: c_int,
) -> c_int;
/// Narrow bridge to `programs/dibio.c`. The file loader and dictionary
/// algorithms remain on the C/Rust library side of this boundary.
#[cfg(feature = "compression")]
fn DiB_trainFromFiles(
dict_file_name: *const c_char,
max_dict_size: usize,
file_names: *const *const c_char,
nb_files: c_int,
chunk_size: usize,
legacy_params: *mut ZDICT_legacy_params_t,
cover_params: *mut ZDICT_cover_params_t,
fast_cover_params: *mut ZDICT_fastCover_params_t,
optimize: c_int,
mem_limit: c_uint,
) -> c_int;
#[cfg(feature = "compression")]
fn ZSTD_getCParams(
compression_level: c_int,
@@ -255,6 +319,15 @@ enum Operation {
Test,
Bench,
List,
Train,
}
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
enum TrainingAlgorithm {
Cover,
#[default]
FastCover,
Legacy,
}
#[derive(Debug)]
@@ -312,6 +385,12 @@ struct Cli {
file_lists: Vec<CString>,
output_dir_flat: Option<CString>,
output_dir_mirror: Option<CString>,
training_algorithm: TrainingAlgorithm,
max_dict_size: usize,
dictionary_id: u32,
dictionary_selectivity: u32,
cover_params: ZDICT_cover_params_t,
fast_cover_params: ZDICT_fastCover_params_t,
unsupported_program: Option<String>,
}
@@ -364,6 +443,12 @@ impl Cli {
file_lists: Vec::new(),
output_dir_flat: None,
output_dir_mirror: None,
training_algorithm: TrainingAlgorithm::default(),
max_dict_size: DEFAULT_MAX_DICT_SIZE,
dictionary_id: 0,
dictionary_selectivity: DEFAULT_DICT_SELECTIVITY,
cover_params: ZDICT_cover_params_t::default(),
fast_cover_params: default_fast_cover_params(),
unsupported_program: None,
};
@@ -530,8 +615,9 @@ fn usage(advanced: bool) {
);
let _ = writeln!(
out,
"\nNot yet migrated: dictionary training, trace, and alternate formats."
"\nDictionary builder:\n --train Create a dictionary from training files\n --train-cover[=k=#,d=#,steps=#,split=#,shrink[=#]]\n --train-fastcover[=k=#,d=#,f=#,steps=#,split=#,accel=#,shrink[=#]]\n --train-legacy[=s=#] Use the legacy algorithm\n --maxdict=# Maximum dictionary size (default {DEFAULT_MAX_DICT_SIZE})\n --dictID=# Force the dictionary ID (default: random)"
);
let _ = writeln!(out, "\nNot yet migrated: trace and alternate formats.");
}
}
@@ -751,6 +837,101 @@ fn parse_adapt(value: &str, cli: &mut Cli) -> Result<(), String> {
Ok(())
}
fn default_fast_cover_params() -> ZDICT_fastCover_params_t {
ZDICT_fastCover_params_t {
d: 8,
f: 20,
steps: 4,
splitPoint: 0.75,
accel: DEFAULT_FASTCOVER_ACCEL,
shrinkDictMaxRegression: DEFAULT_SHRINK_DICT_REGRESSION,
..ZDICT_fastCover_params_t::default()
}
}
fn parse_training_u32(value: &str, option: &str, field: &str) -> Result<u32, String> {
parse_u32(value, &format!("{option} {field}"))
}
fn parse_cover_parameters(value: &str) -> Result<ZDICT_cover_params_t, String> {
let mut params = ZDICT_cover_params_t::default();
for item in value.split(',') {
if item == "shrink" {
params.shrinkDict = 1;
params.shrinkDictMaxRegression = DEFAULT_SHRINK_DICT_REGRESSION;
continue;
}
if let Some(raw) = item.strip_prefix("shrink=") {
params.shrinkDict = 1;
params.shrinkDictMaxRegression =
parse_training_u32(raw, "--train-cover", "shrink regression")?;
continue;
}
let Some((field, raw)) = item.split_once('=') else {
return Err(format!("invalid --train-cover parameter {item:?}"));
};
let parsed = parse_training_u32(raw, "--train-cover", field)?;
match field {
"k" => params.k = parsed,
"d" => params.d = parsed,
"steps" => params.steps = parsed,
"split" => params.splitPoint = f64::from(parsed) / 100.0,
_ => return Err(format!("unknown --train-cover parameter {field:?}")),
}
}
Ok(params)
}
fn parse_fast_cover_parameters(value: &str) -> Result<ZDICT_fastCover_params_t, String> {
let mut params = ZDICT_fastCover_params_t::default();
for item in value.split(',') {
if item == "shrink" {
params.shrinkDict = 1;
params.shrinkDictMaxRegression = DEFAULT_SHRINK_DICT_REGRESSION;
continue;
}
if let Some(raw) = item.strip_prefix("shrink=") {
params.shrinkDict = 1;
params.shrinkDictMaxRegression =
parse_training_u32(raw, "--train-fastcover", "shrink regression")?;
continue;
}
let Some((field, raw)) = item.split_once('=') else {
return Err(format!("invalid --train-fastcover parameter {item:?}"));
};
let parsed = parse_training_u32(raw, "--train-fastcover", field)?;
match field {
"k" => params.k = parsed,
"d" => params.d = parsed,
"f" => params.f = parsed,
"steps" => params.steps = parsed,
"split" => params.splitPoint = f64::from(parsed) / 100.0,
"accel" => params.accel = parsed,
_ => return Err(format!("unknown --train-fastcover parameter {field:?}")),
}
}
Ok(params)
}
fn parse_legacy_parameters(value: &str) -> Result<u32, String> {
let Some(raw) = value
.strip_prefix("s=")
.or_else(|| value.strip_prefix("selectivity="))
else {
return Err(format!("invalid --train-legacy parameter {value:?}"));
};
parse_training_u32(raw, "--train-legacy", "selectivity")
}
fn select_training(cli: &mut Cli, algorithm: TrainingAlgorithm) -> Result<(), String> {
cli.operation = Operation::Train;
cli.training_algorithm = algorithm;
if cli.output.is_none() {
cli.output = Some(cstring(DEFAULT_DICT_NAME)?);
}
Ok(())
}
fn unsupported(option: &str) -> Result<(), String> {
Err(format!(
"{option} is not yet implemented by the Rust CLI frontend"
@@ -807,6 +988,7 @@ fn parse_long_option(
| "--no-name"
| "--list"
| "--show-default-cparams"
| "--train"
)
{
return Err(format!("{name} does not take an argument"));
@@ -1020,6 +1202,45 @@ fn parse_long_option(
cli.show_default_cparams = true;
Ok(None)
}
"--train" => {
select_training(cli, cli.training_algorithm)?;
Ok(None)
}
"--train-cover" => {
let params = attached.map_or_else(
|| Ok(ZDICT_cover_params_t::default()),
parse_cover_parameters,
)?;
select_training(cli, TrainingAlgorithm::Cover)?;
cli.cover_params = params;
Ok(None)
}
"--train-fastcover" => {
let params = attached.map_or_else(
|| Ok(ZDICT_fastCover_params_t::default()),
parse_fast_cover_parameters,
)?;
select_training(cli, TrainingAlgorithm::FastCover)?;
cli.fast_cover_params = params;
Ok(None)
}
"--train-legacy" => {
let selectivity =
attached.map_or_else(|| Ok(cli.dictionary_selectivity), parse_legacy_parameters)?;
select_training(cli, TrainingAlgorithm::Legacy)?;
cli.dictionary_selectivity = selectivity;
Ok(None)
}
"--maxdict" => {
let value = next_field(attached, args, index, name)?;
cli.max_dict_size = parse_u32(&value, "maximum dictionary size")? as usize;
Ok(None)
}
"--dictID" => {
let value = next_field(attached, args, index, name)?;
cli.dictionary_id = parse_u32(&value, "dictionary ID")?;
Ok(None)
}
"--filelist" => {
let value = next_field(attached, args, index, name)?;
cli.file_lists.push(cstring(&value)?);
@@ -1050,13 +1271,7 @@ fn parse_long_option(
cli.output_dir_mirror = Some(cstring(&value)?);
Ok(None)
}
"--train"
| "--train-cover"
| "--train-fastcover"
| "--train-legacy"
| "--max"
| "--maxdict"
| "--dictID"
"--max"
| "--patch-from"
| "--trace"
| "--format"
@@ -1539,7 +1754,7 @@ unsafe fn run_decompress(
dictionary,
)
},
Operation::Compress | Operation::Bench | Operation::List => {
Operation::Compress | Operation::Bench | Operation::List | Operation::Train => {
unreachable!("compression, benchmark, and list are dispatched separately")
}
}
@@ -1632,6 +1847,95 @@ fn run_bench(cli: &Cli) -> Result<i32, String> {
Ok(result)
}
/// Runs dictionary training through the narrow `programs/dibio.c` ABI. The
/// file loading and output path stay in that bridge for now; the selected
/// dictionary builder symbols are supplied by the Rust library archive.
fn run_train(cli: &Cli) -> Result<i32, String> {
#[cfg(not(feature = "compression"))]
{
let _ = cli;
return Err("training mode not available".to_owned());
}
#[cfg(feature = "compression")]
{
let default_output = cstring(DEFAULT_DICT_NAME).expect("static dictionary name");
let output = cli.output.as_ref().unwrap_or(&default_output);
let inputs: Vec<*const c_char> = cli.inputs.iter().map(|value| value.as_ptr()).collect();
let workers = unsafe { resolved_worker_count(cli.workers, cli.single_thread) } as c_uint;
let z_params = ZDICT_params_t {
compressionLevel: cli.level,
notificationLevel: cli.display_level as c_uint,
dictID: cli.dictionary_id,
};
let chunk_size = cli.block_size.unwrap_or(0);
let mem_limit = cli.mem_limit.unwrap_or(0);
let result = match cli.training_algorithm {
TrainingAlgorithm::Cover => {
let mut params = cli.cover_params;
params.nbThreads = workers;
params.zParams = z_params;
let optimize = c_int::from(params.k == 0 || params.d == 0);
unsafe {
DiB_trainFromFiles(
output.as_ptr(),
cli.max_dict_size,
inputs.as_ptr(),
inputs.len() as c_int,
chunk_size,
ptr::null_mut(),
&mut params,
ptr::null_mut(),
optimize,
mem_limit,
)
}
}
TrainingAlgorithm::FastCover => {
let mut params = cli.fast_cover_params;
params.nbThreads = workers;
params.zParams = z_params;
let optimize = c_int::from(params.k == 0 || params.d == 0);
unsafe {
DiB_trainFromFiles(
output.as_ptr(),
cli.max_dict_size,
inputs.as_ptr(),
inputs.len() as c_int,
chunk_size,
ptr::null_mut(),
ptr::null_mut(),
&mut params,
optimize,
mem_limit,
)
}
}
TrainingAlgorithm::Legacy => {
let mut params = ZDICT_legacy_params_t {
selectivityLevel: cli.dictionary_selectivity,
zParams: z_params,
};
unsafe {
DiB_trainFromFiles(
output.as_ptr(),
cli.max_dict_size,
inputs.as_ptr(),
inputs.len() as c_int,
chunk_size,
&mut params,
ptr::null_mut(),
ptr::null_mut(),
0,
mem_limit,
)
}
}
};
Ok(result)
}
}
fn run_cli(mut cli: Cli) -> Result<i32, String> {
if let Some(program_name) = &cli.unsupported_program {
return Err(format!(
@@ -1664,6 +1968,9 @@ fn run_cli(mut cli: Cli) -> Result<i32, String> {
if cli.operation == Operation::Bench {
return run_bench(&cli);
}
if cli.operation == Operation::Train {
return run_train(&cli);
}
if cli.operation == Operation::Test {
cli.output = Some(cstring(NULL_MARK)?);
cli.remove_source = false;
@@ -1784,7 +2091,7 @@ fn run_cli(mut cli: Cli) -> Result<i32, String> {
#[cfg(not(feature = "decompression"))]
unreachable!("unsupported decompression was rejected above")
}
Operation::Bench | Operation::List => {
Operation::Bench | Operation::List | Operation::Train => {
unreachable!("benchmark and list modes were dispatched earlier")
}
}
@@ -2026,11 +2333,99 @@ mod tests {
}
#[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");
fn train_defaults_to_fastcover_and_dictionary_output() {
let cli = parse(&["zstd", "--train", "sample"]);
assert!(error.contains("not yet implemented"));
assert_eq!(cli.operation, Operation::Train);
assert_eq!(cli.training_algorithm, TrainingAlgorithm::FastCover);
assert_eq!(
cli.output.as_deref().map(CStr::to_bytes),
Some(&b"dictionary"[..])
);
assert_eq!(cli.max_dict_size, DEFAULT_MAX_DICT_SIZE);
assert_eq!(cli.fast_cover_params.d, 8);
assert_eq!(cli.fast_cover_params.f, 20);
assert_eq!(cli.fast_cover_params.steps, 4);
assert_eq!(cli.fast_cover_params.splitPoint, 0.75);
assert_eq!(cli.fast_cover_params.accel, DEFAULT_FASTCOVER_ACCEL);
}
#[test]
fn train_cover_parameters_select_algorithm_and_parse_optional_fields() {
let cli = parse(&[
"zstd",
"--train-cover=k=48,d=8,steps=32,split=80,shrink=3",
"sample",
]);
assert_eq!(cli.operation, Operation::Train);
assert_eq!(cli.training_algorithm, TrainingAlgorithm::Cover);
assert_eq!(cli.cover_params.k, 48);
assert_eq!(cli.cover_params.d, 8);
assert_eq!(cli.cover_params.steps, 32);
assert_eq!(cli.cover_params.splitPoint, 0.8);
assert_eq!(cli.cover_params.shrinkDict, 1);
assert_eq!(cli.cover_params.shrinkDictMaxRegression, 3);
}
#[test]
fn train_fastcover_and_legacy_parameter_forms_share_training_options() {
let fast = parse(&[
"zstd",
"--train-fastcover=k=64,d=8,f=21,steps=12,split=75,accel=2,shrink",
"--maxdict",
"128K",
"--dictID=42",
"-B4K",
"-T2",
"-o",
"custom.dict",
"sample",
]);
assert_eq!(fast.training_algorithm, TrainingAlgorithm::FastCover);
assert_eq!(fast.fast_cover_params.k, 64);
assert_eq!(fast.fast_cover_params.d, 8);
assert_eq!(fast.fast_cover_params.f, 21);
assert_eq!(fast.fast_cover_params.steps, 12);
assert_eq!(fast.fast_cover_params.splitPoint, 0.75);
assert_eq!(fast.fast_cover_params.accel, 2);
assert_eq!(fast.fast_cover_params.shrinkDict, 1);
assert_eq!(fast.fast_cover_params.shrinkDictMaxRegression, 1);
assert_eq!(fast.max_dict_size, 128 << 10);
assert_eq!(fast.dictionary_id, 42);
assert_eq!(fast.block_size, Some(4 << 10));
assert_eq!(fast.workers, Some(2));
assert_eq!(
fast.output.as_deref().map(CStr::to_bytes),
Some(&b"custom.dict"[..])
);
let legacy = parse(&["zstd", "--train-legacy=selectivity=12", "sample"]);
assert_eq!(legacy.operation, Operation::Train);
assert_eq!(legacy.training_algorithm, TrainingAlgorithm::Legacy);
assert_eq!(legacy.dictionary_selectivity, 12);
}
#[test]
fn training_options_preserve_upstream_validation_boundaries() {
let attached = parse_args(vec![
OsString::from("zstd"),
OsString::from("--train=value"),
])
.expect_err("--train does not accept an attached value");
assert!(attached.contains("does not take an argument"));
let unknown = parse_args(vec![
OsString::from("zstd"),
OsString::from("--train-cover=unknown=1"),
])
.expect_err("unknown cover parameters must be rejected");
assert!(unknown.contains("unknown --train-cover parameter"));
let missing = parse_args(vec![OsString::from("zstd"), OsString::from("--dictID")])
.expect_err("--dictID requires a value");
assert!(missing.contains("missing argument for --dictID"));
}
#[test]