diff --git a/rust/src/zstd_cli.rs b/rust/src/zstd_cli.rs index 41c653a5d..6095cabe5 100644 --- a/rust/src/zstd_cli.rs +++ b/rust/src/zstd_cli.rs @@ -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, output_dir_flat: Option, output_dir_mirror: Option, + 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, } @@ -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 { + parse_u32(value, &format!("{option} {field}")) +} + +fn parse_cover_parameters(value: &str) -> Result { + 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 { + 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 { + 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 { 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 { + #[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 { if let Some(program_name) = &cli.unsupported_program { return Err(format!( @@ -1664,6 +1968,9 @@ fn run_cli(mut cli: Cli) -> Result { 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 { #[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]