Skip to main content

lz4_flex/block/
compress.rs

1//! The compression algorithm.
2//!
3//! We make use of hash tables to find duplicates. This gives a reasonable compression ratio with a
4//! high performance. It has fixed memory usage, which contrary to other approaches, makes it less
5//! memory hungry.
6
7use crate::block::hashtable::HashTable;
8use crate::block::END_OFFSET;
9use crate::block::LZ4_MIN_LENGTH;
10use crate::block::MAX_DISTANCE;
11use crate::block::MFLIMIT;
12use crate::block::MINMATCH;
13#[cfg(all(feature = "alloc", not(feature = "safe-encode")))]
14use crate::sink::PtrSink;
15use crate::sink::Sink;
16use crate::sink::SliceSink;
17#[cfg(feature = "alloc")]
18#[allow(unused_imports)]
19use alloc::vec;
20
21#[cfg(feature = "alloc")]
22#[allow(unused_imports)]
23use alloc::vec::Vec;
24
25pub(crate) use super::hashtable::HashTable4K;
26pub(crate) use super::hashtable::HashTable4KU16;
27use super::{CompressError, WINDOW_SIZE};
28
29/// Increase step size after 1<<INCREASE_STEPSIZE_BITSHIFT non matches
30const INCREASE_STEPSIZE_BITSHIFT: usize = 5;
31
32/// Read a 4-byte "batch" from some position.
33///
34/// This will read a native-endian 4-byte integer from some position.
35#[inline]
36#[cfg(not(feature = "safe-encode"))]
37pub(super) fn get_batch(input: &[u8], n: usize) -> u32 {
38    unsafe { read_u32_ptr(input.as_ptr().add(n)) }
39}
40
41#[inline]
42#[cfg(feature = "safe-encode")]
43pub(super) fn get_batch(input: &[u8], n: usize) -> u32 {
44    u32::from_ne_bytes(input[n..n + 4].try_into().unwrap())
45}
46
47/// Read an usize sized "batch" from some position.
48///
49/// This will read a native-endian usize from some position.
50#[inline]
51#[allow(dead_code)]
52#[cfg(not(feature = "safe-encode"))]
53pub(super) fn get_batch_arch(input: &[u8], n: usize) -> usize {
54    unsafe { read_usize_ptr(input.as_ptr().add(n)) }
55}
56
57#[inline]
58#[allow(dead_code)]
59#[cfg(feature = "safe-encode")]
60pub(super) fn get_batch_arch(input: &[u8], n: usize) -> usize {
61    const USIZE_SIZE: usize = core::mem::size_of::<usize>();
62    let arr: &[u8; USIZE_SIZE] = input[n..n + USIZE_SIZE].try_into().unwrap();
63    usize::from_ne_bytes(*arr)
64}
65
66#[inline]
67fn token_from_literal(lit_len: usize) -> u8 {
68    if lit_len < 0xF {
69        // Since we can fit the literals length into it, there is no need for saturation.
70        (lit_len as u8) << 4
71    } else {
72        // We were unable to fit the literals into it, so we saturate to 0xF. We will later
73        // write the extensional value.
74        0xF0
75    }
76}
77
78#[inline]
79fn token_from_literal_and_match_length(lit_len: usize, duplicate_length: usize) -> u8 {
80    let mut token = if lit_len < 0xF {
81        // Since we can fit the literals length into it, there is no need for saturation.
82        (lit_len as u8) << 4
83    } else {
84        // We were unable to fit the literals into it, so we saturate to 0xF. We will later
85        // write the extensional value.
86        0xF0
87    };
88
89    token |= if duplicate_length < 0xF {
90        // We could fit it in.
91        duplicate_length as u8
92    } else {
93        // We were unable to fit it in, so we default to 0xF, which will later be extended.
94        0xF
95    };
96
97    token
98}
99
100/// Counts the number of same bytes in two byte streams.
101/// `input` is the complete input
102/// `cur` is the current position in the input. it will be incremented by the number of matched
103/// bytes `source` either the same as input or an external slice
104/// `candidate` is the candidate position in `source`
105///
106/// The function ignores the last END_OFFSET bytes in input as those should be literals.
107#[inline]
108#[cfg(feature = "safe-encode")]
109fn count_same_bytes(input: &[u8], cur: &mut usize, source: &[u8], candidate: usize) -> usize {
110    const USIZE_SIZE: usize = core::mem::size_of::<usize>();
111    let cur_slice = &input[*cur..input.len() - END_OFFSET];
112    let cand_slice = &source[candidate..];
113
114    let mut num = 0;
115    for (block1, block2) in cur_slice
116        .chunks_exact(USIZE_SIZE)
117        .zip(cand_slice.chunks_exact(USIZE_SIZE))
118    {
119        let input_block = usize::from_ne_bytes(block1.try_into().unwrap());
120        let match_block = usize::from_ne_bytes(block2.try_into().unwrap());
121
122        if input_block == match_block {
123            num += USIZE_SIZE;
124        } else {
125            let diff = input_block ^ match_block;
126            num += (diff.to_le().trailing_zeros() / 8) as usize;
127            *cur += num;
128            return num;
129        }
130    }
131
132    // If we're here we may have 1 to 7 bytes left to check close to the end of input
133    // or source slices. Since this is rare occurrence we mark it cold to get better
134    // ~5% better performance.
135    #[cold]
136    fn count_same_bytes_tail(a: &[u8], b: &[u8], offset: usize) -> usize {
137        a.iter()
138            .zip(b)
139            .skip(offset)
140            .take_while(|(a, b)| a == b)
141            .count()
142    }
143    num += count_same_bytes_tail(cur_slice, cand_slice, num);
144
145    *cur += num;
146    num
147}
148
149/// Counts the number of same bytes in two byte streams.
150/// `input` is the complete input
151/// `cur` is the current position in the input. it will be incremented by the number of matched
152/// bytes `source` either the same as input OR an external slice
153/// `candidate` is the candidate position in `source`
154///
155/// The function ignores the last END_OFFSET bytes in input as those should be literals.
156#[inline]
157#[cfg(not(feature = "safe-encode"))]
158fn count_same_bytes(input: &[u8], cur: &mut usize, source: &[u8], candidate: usize) -> usize {
159    let max_input_match = input.len().saturating_sub(*cur + END_OFFSET);
160    let max_candidate_match = source.len() - candidate;
161    // Considering both limits calc how far we may match in input.
162    let input_end = *cur + max_input_match.min(max_candidate_match);
163
164    let start = *cur;
165    let mut source_ptr = unsafe { source.as_ptr().add(candidate) };
166
167    // compare 4/8 bytes blocks depending on the arch
168    const STEP_SIZE: usize = core::mem::size_of::<usize>();
169    while *cur + STEP_SIZE <= input_end {
170        let diff = read_usize_ptr(unsafe { input.as_ptr().add(*cur) }) ^ read_usize_ptr(source_ptr);
171
172        if diff == 0 {
173            *cur += STEP_SIZE;
174            unsafe {
175                source_ptr = source_ptr.add(STEP_SIZE);
176            }
177        } else {
178            *cur += (diff.to_le().trailing_zeros() / 8) as usize;
179            return *cur - start;
180        }
181    }
182
183    // compare 4 bytes block
184    #[cfg(target_pointer_width = "64")]
185    {
186        if input_end - *cur >= 4 {
187            let diff = read_u32_ptr(unsafe { input.as_ptr().add(*cur) }) ^ read_u32_ptr(source_ptr);
188
189            if diff == 0 {
190                *cur += 4;
191                unsafe {
192                    source_ptr = source_ptr.add(4);
193                }
194            } else {
195                *cur += (diff.to_le().trailing_zeros() / 8) as usize;
196                return *cur - start;
197            }
198        }
199    }
200
201    // compare 2 bytes block
202    if input_end - *cur >= 2
203        && unsafe { read_u16_ptr(input.as_ptr().add(*cur)) == read_u16_ptr(source_ptr) }
204    {
205        *cur += 2;
206        unsafe {
207            source_ptr = source_ptr.add(2);
208        }
209    }
210
211    if *cur < input_end
212        && unsafe { input.as_ptr().add(*cur).read() } == unsafe { source_ptr.read() }
213    {
214        *cur += 1;
215    }
216
217    *cur - start
218}
219
220/// Write an integer to the output.
221///
222/// Each additional byte then represent a value from 0 to 255, which is added to the previous value
223/// to produce a total length. When the byte value is 255, another byte must read and added, and so
224/// on. There can be any number of bytes of value "255" following token
225#[inline]
226pub(super) fn write_integer(output: &mut impl Sink, mut n: usize) {
227    // Note: Since `n` is usually < 0xFF and writing multiple bytes to the output
228    // requires 2 branches of bound check (due to the possibility of add overflows)
229    // the simple byte at a time implementation below is faster in most cases.
230    while n >= 0xFF {
231        n -= 0xFF;
232        push_byte(output, 0xFF);
233    }
234    push_byte(output, n as u8);
235}
236
237/// Handle the last bytes from the input as literals
238#[cold]
239fn handle_last_literals(output: &mut impl Sink, input: &[u8], start: usize) {
240    let lit_len = input.len() - start;
241
242    let token = token_from_literal(lit_len);
243    push_byte(output, token);
244    if lit_len >= 0xF {
245        write_integer(output, lit_len - 0xF);
246    }
247    // Now, write the actual literals.
248    output.extend_from_slice(&input[start..]);
249}
250
251/// Moves the cursors back as long as the bytes match, to find additional bytes in a duplicate
252#[inline]
253#[cfg(feature = "safe-encode")]
254fn backtrack_match(
255    input: &[u8],
256    cur: &mut usize,
257    literal_start: usize,
258    source: &[u8],
259    candidate: &mut usize,
260) {
261    // Note: Even if iterator version of this loop has less branches inside the loop it has more
262    // branches before the loop. That in practice seems to make it slower than the while version
263    // bellow. TODO: It should be possible remove all bounds checks, since we are walking
264    // backwards
265    while *candidate > 0 && *cur > literal_start && input[*cur - 1] == source[*candidate - 1] {
266        *cur -= 1;
267        *candidate -= 1;
268    }
269}
270
271/// Moves the cursors back as long as the bytes match, to find additional bytes in a duplicate
272#[inline]
273#[cfg(not(feature = "safe-encode"))]
274fn backtrack_match(
275    input: &[u8],
276    cur: &mut usize,
277    literal_start: usize,
278    source: &[u8],
279    candidate: &mut usize,
280) {
281    while unsafe {
282        *candidate > 0
283            && *cur > literal_start
284            && input.get_unchecked(*cur - 1) == source.get_unchecked(*candidate - 1)
285    } {
286        *cur -= 1;
287        *candidate -= 1;
288    }
289}
290
291/// Compress all bytes of `input[input_pos..]` into `output`.
292///
293/// Bytes in `input[..input_pos]` are treated as a preamble and can be used for lookback.
294/// This part is known as the compressor "prefix".
295/// Bytes in `ext_dict` logically precede the bytes in `input` and can also be used for lookback.
296///
297/// `input_stream_offset` is the logical position of the first byte of `input`. This allows same
298/// `dict` to be used for many calls to `compress_internal` as we can "readdress" the first byte of
299/// `input` to be something other than 0.
300///
301/// `dict` is the dictionary of previously encoded sequences.
302///
303/// This is used to find duplicates in the stream so they are not written multiple times.
304///
305/// Every four bytes are hashed, and in the resulting slot their position in the input buffer
306/// is placed in the dict. This way we can easily look up a candidate to back references.
307///
308/// Returns the number of bytes written (compressed) into `output`.
309///
310/// # Const parameters
311/// `USE_DICT`: Disables usage of ext_dict (it'll panic if a non-empty slice is used).
312/// In other words, this generates more optimized code when an external dictionary isn't used.
313///
314/// A similar const argument could be used to disable the Prefix mode (eg. USE_PREFIX),
315/// which would impose `input_pos == 0 && input_stream_offset == 0`. Experiments didn't
316/// show significant improvement though.
317// Intentionally avoid inlining.
318// Empirical tests revealed it to be rarely better but often significantly detrimental.
319#[inline(never)]
320pub(crate) fn compress_internal<T: HashTable, const USE_DICT: bool, S: Sink>(
321    input: &[u8],
322    input_pos: usize,
323    output: &mut S,
324    dict: &mut T,
325    ext_dict: &[u8],
326    input_stream_offset: usize,
327) -> Result<usize, CompressError> {
328    assert!(input_pos <= input.len());
329    if USE_DICT {
330        assert!(ext_dict.len() <= super::WINDOW_SIZE);
331        assert!(ext_dict.len() <= input_stream_offset);
332        // Check for overflow hazard when using ext_dict
333        assert!(input_stream_offset
334            .checked_add(input.len())
335            .and_then(|i| i.checked_add(ext_dict.len()))
336            .is_some_and(|i| i <= isize::MAX as usize));
337    } else {
338        assert!(ext_dict.is_empty());
339    }
340    if output.capacity() - output.pos() < get_maximum_output_size(input.len() - input_pos) {
341        return Err(CompressError::OutputTooSmall);
342    }
343
344    let output_start_pos = output.pos();
345    if input.len() - input_pos < LZ4_MIN_LENGTH {
346        handle_last_literals(output, input, input_pos);
347        return Ok(output.pos() - output_start_pos);
348    }
349
350    let ext_dict_stream_offset = input_stream_offset - ext_dict.len();
351    let end_pos_check = input.len() - MFLIMIT;
352    let mut literal_start = input_pos;
353    let mut cur = input_pos;
354
355    if cur == 0 && input_stream_offset == 0 {
356        // According to the spec we can't start with a match,
357        // except when referencing another block.
358        let hash = T::get_hash_at(input, 0);
359        dict.put_at(hash, 0);
360        cur = 1;
361    }
362
363    loop {
364        // Read the next block into two sections, the literals and the duplicates.
365        let mut step_size;
366        let mut candidate;
367        let mut candidate_source;
368        let mut offset;
369        let mut non_match_count = 1 << INCREASE_STEPSIZE_BITSHIFT;
370        // The number of bytes before our cursor, where the duplicate starts.
371        let mut next_cur = cur;
372
373        // In this loop we search for duplicates via the hashtable. 4bytes or 8bytes are hashed and
374        // compared.
375        loop {
376            step_size = non_match_count >> INCREASE_STEPSIZE_BITSHIFT;
377            non_match_count += 1;
378
379            cur = next_cur;
380            next_cur += step_size;
381
382            // Same as cur + MFLIMIT > input.len()
383            if cur > end_pos_check {
384                handle_last_literals(output, input, literal_start);
385                return Ok(output.pos() - output_start_pos);
386            }
387            // Find a candidate in the dictionary with the hash of the current four bytes.
388            // Unchecked is safe as long as the values from the hash function don't exceed the size
389            // of the table. This is ensured by right shifting the hash values
390            // (`dict_bitshift`) to fit them in the table
391
392            // [Bounds Check]: Can be elided due to `end_pos_check` above
393            let hash = T::get_hash_at(input, cur);
394            candidate = dict.get_at(hash);
395            dict.put_at(hash, cur + input_stream_offset);
396
397            // Sanity check: Matches can't be ahead of `cur`.
398            debug_assert!(candidate <= input_stream_offset + cur);
399
400            // Two requirements to the candidate exists:
401            // - We should not return a position which is merely a hash collision, so that the
402            //   candidate actually matches what we search for.
403            // - We can address up to 16-bit offset, hence we are only able to address the candidate
404            //   if its offset is less than or equals to 0xFFFF.
405            if input_stream_offset + cur - candidate > MAX_DISTANCE {
406                continue;
407            }
408
409            if candidate >= input_stream_offset {
410                // match within input
411                offset = (input_stream_offset + cur - candidate) as u16;
412                candidate -= input_stream_offset;
413                candidate_source = input;
414            } else if USE_DICT {
415                // Sanity check, which may fail if we lost history beyond MAX_DISTANCE
416                debug_assert!(
417                    candidate >= ext_dict_stream_offset,
418                    "Lost history in ext dict mode"
419                );
420                // match within ext dict
421                offset = (input_stream_offset + cur - candidate) as u16;
422                candidate -= ext_dict_stream_offset;
423                candidate_source = ext_dict;
424            } else {
425                // Match is not reachable anymore
426                // eg. compressing an independent block frame w/o clearing
427                // the matches tables, only increasing input_stream_offset.
428                // Sanity check
429                debug_assert!(input_pos == 0, "Lost history in prefix mode");
430                continue;
431            }
432            // [Bounds Check]: Candidate is coming from the Hashmap. It can't be out of bounds, but
433            // impossible to prove for the compiler and remove the bounds checks.
434            let cand_bytes: u32 = get_batch(candidate_source, candidate);
435            // [Bounds Check]: Should be able to be elided due to `end_pos_check`.
436            let curr_bytes: u32 = get_batch(input, cur);
437
438            if cand_bytes == curr_bytes {
439                break;
440            }
441        }
442
443        // Extend the match backwards if we can
444        backtrack_match(
445            input,
446            &mut cur,
447            literal_start,
448            candidate_source,
449            &mut candidate,
450        );
451
452        // The length (in bytes) of the literals section.
453        let lit_len = cur - literal_start;
454
455        // Generate the higher half of the token.
456        cur += MINMATCH;
457        candidate += MINMATCH;
458        let duplicate_length = count_same_bytes(input, &mut cur, candidate_source, candidate);
459
460        // Note: The `- 2` offset was copied from the reference implementation, it could be
461        // arbitrary.
462        let hash = T::get_hash_at(input, cur - 2);
463        dict.put_at(hash, cur - 2 + input_stream_offset);
464
465        let token = token_from_literal_and_match_length(lit_len, duplicate_length);
466
467        // Push the token to the output stream.
468        push_byte(output, token);
469        // If we were unable to fit the literals length into the token, write the extensional
470        // part.
471        if lit_len >= 0xF {
472            write_integer(output, lit_len - 0xF);
473        }
474
475        // Now, write the actual literals.
476        //
477        // The unsafe version copies blocks of 8bytes, and therefore may copy up to 7bytes more than
478        // needed. This is safe, because the last 12 bytes (MF_LIMIT) are handled in
479        // handle_last_literals.
480        copy_literals_wild(output, input, literal_start, lit_len);
481        // write the offset in little endian.
482        push_u16(output, offset);
483
484        // If we were unable to fit the duplicates length into the token, write the
485        // extensional part.
486        if duplicate_length >= 0xF {
487            write_integer(output, duplicate_length - 0xF);
488        }
489        literal_start = cur;
490    }
491}
492
493#[inline]
494#[cfg(feature = "safe-encode")]
495fn push_byte(output: &mut impl Sink, el: u8) {
496    output.push(el);
497}
498
499#[inline]
500#[cfg(not(feature = "safe-encode"))]
501fn push_byte(output: &mut impl Sink, el: u8) {
502    unsafe {
503        core::ptr::write(output.pos_mut_ptr(), el);
504        output.set_pos(output.pos() + 1);
505    }
506}
507
508#[inline]
509#[cfg(feature = "safe-encode")]
510fn push_u16(output: &mut impl Sink, el: u16) {
511    output.extend_from_slice(&el.to_le_bytes());
512}
513
514#[inline]
515#[cfg(not(feature = "safe-encode"))]
516fn push_u16(output: &mut impl Sink, el: u16) {
517    unsafe {
518        core::ptr::copy_nonoverlapping(el.to_le_bytes().as_ptr(), output.pos_mut_ptr(), 2);
519        output.set_pos(output.pos() + 2);
520    }
521}
522
523#[inline(always)] // (always) necessary otherwise compiler fails to inline it
524#[cfg(feature = "safe-encode")]
525fn copy_literals_wild(output: &mut impl Sink, input: &[u8], input_start: usize, len: usize) {
526    output.extend_from_slice_wild(&input[input_start..input_start + len], len)
527}
528
529#[inline]
530#[cfg(not(feature = "safe-encode"))]
531fn copy_literals_wild(output: &mut impl Sink, input: &[u8], input_start: usize, len: usize) {
532    debug_assert!(input_start + len / 8 * 8 + ((len % 8) != 0) as usize * 8 <= input.len());
533    debug_assert!(output.pos() + len / 8 * 8 + ((len % 8) != 0) as usize * 8 <= output.capacity());
534    unsafe {
535        // Note: This used to be a wild copy loop of 8 bytes, but the compiler consistently
536        // transformed it into a call to memcopy, which hurts performance significantly for
537        // small copies, which are common.
538        let start_ptr = input.as_ptr().add(input_start);
539        match len {
540            0..=8 => core::ptr::copy_nonoverlapping(start_ptr, output.pos_mut_ptr(), 8),
541            9..=16 => core::ptr::copy_nonoverlapping(start_ptr, output.pos_mut_ptr(), 16),
542            17..=24 => core::ptr::copy_nonoverlapping(start_ptr, output.pos_mut_ptr(), 24),
543            _ => core::ptr::copy_nonoverlapping(start_ptr, output.pos_mut_ptr(), len),
544        }
545        output.set_pos(output.pos() + len);
546    }
547}
548
549/// Compress all bytes of `input` into `output`.
550/// The method chooses an appropriate hashtable to lookup duplicates.
551/// output should be preallocated with a size of
552/// `get_maximum_output_size`.
553///
554/// Returns the number of bytes written (compressed) into `output`.
555#[inline]
556pub(crate) fn compress_into_sink_with_dict<const USE_DICT: bool>(
557    input: &[u8],
558    output: &mut impl Sink,
559    mut dict_data: &[u8],
560) -> Result<usize, CompressError> {
561    if USE_DICT && dict_data.len() < MINMATCH {
562        return compress_into_sink_without_dict(input, output);
563    }
564
565    if dict_data.len() + input.len() < u16::MAX as usize {
566        let mut dict = HashTable4KU16::new();
567        init_dict(&mut dict, &mut dict_data);
568        compress_internal::<_, USE_DICT, _>(input, 0, output, &mut dict, dict_data, dict_data.len())
569    } else {
570        let mut dict = HashTable4K::new();
571        init_dict(&mut dict, &mut dict_data);
572        compress_internal::<_, USE_DICT, _>(input, 0, output, &mut dict, dict_data, dict_data.len())
573    }
574}
575
576/// Slow fallback for when the dictionary is too small to be useful. This avoids the overhead of
577/// inlining.
578#[cold]
579#[inline(never)]
580fn compress_into_sink_without_dict(
581    input: &[u8],
582    output: &mut impl Sink,
583) -> Result<usize, CompressError> {
584    compress_into_sink_with_dict::<false>(input, output, b"")
585}
586
587#[inline]
588fn init_dict<T: HashTable>(dict: &mut T, dict_data: &mut &[u8]) {
589    if dict_data.len() > WINDOW_SIZE {
590        *dict_data = &dict_data[dict_data.len() - WINDOW_SIZE..];
591    }
592    let mut i = 0usize;
593    while i + core::mem::size_of::<usize>() <= dict_data.len() {
594        let hash = T::get_hash_at(dict_data, i);
595        dict.put_at(hash, i);
596        // Note: The 3 byte step was copied from the reference implementation, it could be
597        // arbitrary.
598        i += 3;
599    }
600}
601
602/// Returns the maximum output size of the compressed data.
603/// Can be used to preallocate capacity on the output vector
604#[inline]
605pub const fn get_maximum_output_size(input_len: usize) -> usize {
606    16 + 4 + (input_len as u64 * 110 / 100) as usize
607}
608
609/// Compress all bytes of `input` into `output`.
610/// The method chooses an appropriate hashtable to lookup duplicates.
611/// output should be preallocated with a size of
612/// `get_maximum_output_size`.
613///
614/// Returns the number of bytes written (compressed) into `output`.
615#[inline]
616pub fn compress_into(input: &[u8], output: &mut [u8]) -> Result<usize, CompressError> {
617    compress_into_sink_with_dict::<false>(input, &mut SliceSink::new(output, 0), b"")
618}
619
620/// Compress all bytes of `input` into `output`.
621/// The method chooses an appropriate hashtable to lookup duplicates.
622/// output should be preallocated with a size of
623/// `get_maximum_output_size`.
624///
625/// Returns the number of bytes written (compressed) into `output`.
626#[inline]
627pub fn compress_into_with_dict(
628    input: &[u8],
629    output: &mut [u8],
630    dict_data: &[u8],
631) -> Result<usize, CompressError> {
632    compress_into_sink_with_dict::<true>(input, &mut SliceSink::new(output, 0), dict_data)
633}
634
635#[cfg(feature = "alloc")]
636#[inline]
637fn compress_into_vec_with_dict<const USE_DICT: bool>(
638    input: &[u8],
639    prepend_size: bool,
640    dict_data: &[u8],
641) -> Vec<u8> {
642    let prepend_size_num_bytes = if prepend_size { 4 } else { 0 };
643    let max_compressed_size = get_maximum_output_size(input.len()) + prepend_size_num_bytes;
644    if USE_DICT && dict_data.len() < MINMATCH {
645        return compress_into_vec_without_dict(input, prepend_size);
646    }
647    #[cfg(feature = "safe-encode")]
648    let mut compressed = {
649        let mut compressed: Vec<u8> = vec![0u8; max_compressed_size];
650        let out = if prepend_size {
651            compressed[..4].copy_from_slice(&(input.len() as u32).to_le_bytes());
652            &mut compressed[4..]
653        } else {
654            &mut compressed
655        };
656        let compressed_len =
657            compress_into_sink_with_dict::<USE_DICT>(input, &mut SliceSink::new(out, 0), dict_data)
658                .unwrap();
659
660        compressed.truncate(prepend_size_num_bytes + compressed_len);
661        compressed
662    };
663    #[cfg(not(feature = "safe-encode"))]
664    let mut compressed = {
665        let mut vec = Vec::with_capacity(max_compressed_size);
666        let start_pos = if prepend_size {
667            vec.extend_from_slice(&(input.len() as u32).to_le_bytes());
668            4
669        } else {
670            0
671        };
672        let compressed_len = compress_into_sink_with_dict::<USE_DICT>(
673            input,
674            &mut PtrSink::from_vec(&mut vec, start_pos),
675            dict_data,
676        )
677        .unwrap();
678        unsafe {
679            vec.set_len(prepend_size_num_bytes + compressed_len);
680        }
681        vec
682    };
683
684    compressed.shrink_to_fit();
685    compressed
686}
687
688#[cfg(feature = "alloc")]
689#[cold]
690#[inline(never)]
691fn compress_into_vec_without_dict(input: &[u8], prepend_size: bool) -> Vec<u8> {
692    compress_into_vec_with_dict::<false>(input, prepend_size, b"")
693}
694
695/// Compress all bytes of `input` into `output`. The uncompressed size will be prepended as a little
696/// endian u32. Can be used in conjunction with `decompress_size_prepended`
697#[cfg(feature = "alloc")]
698#[cfg_attr(docsrs, doc(cfg(feature = "alloc")))]
699#[inline]
700pub fn compress_prepend_size(input: &[u8]) -> Vec<u8> {
701    compress_into_vec_with_dict::<false>(input, true, b"")
702}
703
704/// Compress all bytes of `input`.
705#[cfg(feature = "alloc")]
706#[cfg_attr(docsrs, doc(cfg(feature = "alloc")))]
707#[inline]
708pub fn compress(input: &[u8]) -> Vec<u8> {
709    compress_into_vec_with_dict::<false>(input, false, b"")
710}
711
712/// Compress all bytes of `input` with an external dictionary.
713#[cfg(feature = "alloc")]
714#[cfg_attr(docsrs, doc(cfg(feature = "alloc")))]
715#[inline]
716pub fn compress_with_dict(input: &[u8], ext_dict: &[u8]) -> Vec<u8> {
717    compress_into_vec_with_dict::<true>(input, false, ext_dict)
718}
719
720/// Compress all bytes of `input` into `output`. The uncompressed size will be prepended as a little
721/// endian u32. Can be used in conjunction with `decompress_size_prepended_with_dict`
722#[cfg(feature = "alloc")]
723#[cfg_attr(docsrs, doc(cfg(feature = "alloc")))]
724#[inline]
725pub fn compress_prepend_size_with_dict(input: &[u8], ext_dict: &[u8]) -> Vec<u8> {
726    compress_into_vec_with_dict::<true>(input, true, ext_dict)
727}
728
729/// A reusable compression table that avoids re-allocating the internal hash table on every call.
730///
731/// This is useful when compressing many small inputs in a loop. Create one table and pass it
732/// to [`compress_into_with_table`] repeatedly.
733///
734/// # Example
735/// ```
736/// use lz4_flex::block::{compress_into_with_table, get_maximum_output_size, CompressTable};
737///
738/// let mut table = CompressTable::default();
739/// let input = b"hello world, hello world, hello!";
740/// let mut output = vec![0u8; get_maximum_output_size(input.len())];
741/// let compressed_len = compress_into_with_table(input, &mut output, &mut table).unwrap();
742/// ```
743pub enum CompressTable {
744    /// Table using 16-bit entries, suitable for inputs where `input.len() < u16::MAX`.
745    Small(HashTable4KU16),
746    /// Table using 32-bit entries, suitable for any input size.
747    Large(HashTable4K),
748}
749
750impl Default for CompressTable {
751    fn default() -> Self {
752        CompressTable::Small(HashTable4KU16::new())
753    }
754}
755
756impl CompressTable {
757    /// Create a small table (16-bit entries). More memory efficient, but only usable when the
758    /// total input size is less than 65535 bytes.
759    #[cfg(feature = "alloc")]
760    pub fn small() -> Self {
761        CompressTable::Small(HashTable4KU16::new())
762    }
763
764    /// Create a small table (16-bit entries). More memory efficient, but only usable when the
765    /// total input size is less than 65535 bytes.
766    #[cfg(not(feature = "alloc"))]
767    pub const fn small() -> Self {
768        CompressTable::Small(HashTable4KU16::new())
769    }
770
771    /// Create a large table (32-bit entries). Works for any input size.
772    #[cfg(feature = "alloc")]
773    pub fn large() -> Self {
774        CompressTable::Large(HashTable4K::new())
775    }
776
777    /// Create a large table (32-bit entries). Works for any input size.
778    #[cfg(not(feature = "alloc"))]
779    pub const fn large() -> Self {
780        CompressTable::Large(HashTable4K::new())
781    }
782}
783
784/// Compress all bytes of `input` into `output`, reusing a [`CompressTable`] to avoid
785/// re-allocating the internal hash table.
786///
787/// `output` should be preallocated with a size of [`get_maximum_output_size`].
788///
789/// Returns the number of bytes written (compressed) into `output`.
790///
791/// **Note:** If the table variant doesn't match the input size (e.g. a `Small` table is used
792/// with input >= 64KB), the table will be transparently upgraded. However, it won't be
793/// downgraded automatically. Without the `alloc` feature the upgrade constructs the new table
794/// on the stack, use [`CompressTable::large`] upfront to avoid this.
795#[inline]
796pub fn compress_into_with_table(
797    input: &[u8],
798    output: &mut [u8],
799    table: &mut CompressTable,
800) -> Result<usize, CompressError> {
801    if input.len() >= u16::MAX as usize && matches!(table, CompressTable::Small(_)) {
802        *table = CompressTable::Large(HashTable4K::new());
803    }
804
805    match table {
806        CompressTable::Small(dict) => {
807            dict.clear();
808            compress_internal::<_, false, _>(input, 0, &mut SliceSink::new(output, 0), dict, b"", 0)
809        }
810        CompressTable::Large(dict) => {
811            dict.clear();
812            compress_internal::<_, false, _>(input, 0, &mut SliceSink::new(output, 0), dict, b"", 0)
813        }
814    }
815}
816
817#[inline]
818#[cfg(not(feature = "safe-encode"))]
819fn read_u16_ptr(input: *const u8) -> u16 {
820    let mut num: u16 = 0;
821    unsafe {
822        core::ptr::copy_nonoverlapping(input, &mut num as *mut u16 as *mut u8, 2);
823    }
824    num
825}
826
827#[inline]
828#[cfg(not(feature = "safe-encode"))]
829fn read_u32_ptr(input: *const u8) -> u32 {
830    let mut num: u32 = 0;
831    unsafe {
832        core::ptr::copy_nonoverlapping(input, &mut num as *mut u32 as *mut u8, 4);
833    }
834    num
835}
836
837#[inline]
838#[cfg(not(feature = "safe-encode"))]
839fn read_usize_ptr(input: *const u8) -> usize {
840    let mut num: usize = 0;
841    unsafe {
842        core::ptr::copy_nonoverlapping(
843            input,
844            &mut num as *mut usize as *mut u8,
845            core::mem::size_of::<usize>(),
846        );
847    }
848    num
849}
850
851#[cfg(test)]
852mod tests {
853    use super::*;
854
855    #[test]
856    fn test_count_same_bytes() {
857        // 8byte aligned block, zeros and ones are added because the end/offset
858        let first: &[u8] = &[
859            1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
860        ];
861        let second: &[u8] = &[
862            1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1,
863        ];
864        assert_eq!(count_same_bytes(first, &mut 0, second, 0), 16);
865
866        // 4byte aligned block
867        let first: &[u8] = &[
868            1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 0, 0, 0, 0, 0, 0, 0, 0, 0,
869            0, 0, 0,
870        ];
871        let second: &[u8] = &[
872            1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 1, 1, 1, 1, 1, 1, 1, 1,
873            1, 1, 1,
874        ];
875        assert_eq!(count_same_bytes(first, &mut 0, second, 0), 20);
876
877        // 2byte aligned block
878        let first: &[u8] = &[
879            1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 3, 4, 0, 0, 0, 0, 0, 0, 0,
880            0, 0, 0, 0, 0,
881        ];
882        let second: &[u8] = &[
883            1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 3, 4, 1, 1, 1, 1, 1, 1, 1,
884            1, 1, 1, 1, 1,
885        ];
886        assert_eq!(count_same_bytes(first, &mut 0, second, 0), 22);
887
888        // 1byte aligned block
889        let first: &[u8] = &[
890            1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 3, 4, 5, 0, 0, 0, 0, 0, 0,
891            0, 0, 0, 0, 0, 0,
892        ];
893        let second: &[u8] = &[
894            1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 3, 4, 5, 1, 1, 1, 1, 1, 1,
895            1, 1, 1, 1, 1, 1,
896        ];
897        assert_eq!(count_same_bytes(first, &mut 0, second, 0), 23);
898
899        // 1byte aligned block - last byte different
900        let first: &[u8] = &[
901            1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 3, 4, 5, 0, 0, 0, 0, 0, 0,
902            0, 0, 0, 0, 0, 0,
903        ];
904        let second: &[u8] = &[
905            1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 3, 4, 6, 1, 1, 1, 1, 1, 1,
906            1, 1, 1, 1, 1, 1,
907        ];
908        assert_eq!(count_same_bytes(first, &mut 0, second, 0), 22);
909
910        // 1byte aligned block
911        let first: &[u8] = &[
912            1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 3, 9, 5, 0, 0, 0, 0, 0, 0,
913            0, 0, 0, 0, 0, 0,
914        ];
915        let second: &[u8] = &[
916            1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 1, 2, 3, 4, 3, 4, 6, 1, 1, 1, 1, 1, 1,
917            1, 1, 1, 1, 1, 1,
918        ];
919        assert_eq!(count_same_bytes(first, &mut 0, second, 0), 21);
920
921        for diff_idx in 8..100 {
922            let first: Vec<u8> = (0u8..255).cycle().take(100 + 12).collect();
923            let mut second = first.clone();
924            second[diff_idx] = 255;
925            for start in 0..=diff_idx {
926                let same_bytes = count_same_bytes(&first, &mut start.clone(), &second, start);
927                assert_eq!(same_bytes, diff_idx - start);
928            }
929        }
930    }
931
932    #[test]
933    fn test_bug() {
934        let input: &[u8] = &[
935            10, 12, 14, 16, 18, 10, 12, 14, 16, 18, 10, 12, 14, 16, 18, 10, 12, 14, 16, 18,
936        ];
937        let _out = compress(input);
938    }
939
940    #[test]
941    fn test_dict() {
942        let input: &[u8] = &[
943            10, 12, 14, 16, 18, 10, 12, 14, 16, 18, 10, 12, 14, 16, 18, 10, 12, 14, 16, 18,
944        ];
945        let dict = input;
946        let compressed = compress_with_dict(input, dict);
947        assert_lt!(compressed.len(), compress(input).len());
948
949        assert!(compressed.len() < compress(input).len());
950        let mut uncompressed = vec![0u8; input.len()];
951        let uncomp_size = crate::block::decompress::decompress_into_with_dict(
952            &compressed,
953            &mut uncompressed,
954            dict,
955        )
956        .unwrap();
957        uncompressed.truncate(uncomp_size);
958        assert_eq!(input, uncompressed);
959    }
960
961    #[test]
962    fn test_dict_no_panic() {
963        let input: &[u8] = &[
964            10, 12, 14, 16, 18, 10, 12, 14, 16, 18, 10, 12, 14, 16, 18, 10, 12, 14, 16, 18,
965        ];
966        let dict = &[10, 12, 14];
967        let _compressed = compress_with_dict(input, dict);
968    }
969
970    #[test]
971    fn compress_into_with_short_dict_does_not_panic() {
972        let input = [0u8; 13];
973
974        for dict_len in 0..MINMATCH {
975            let dict = vec![0u8; dict_len];
976            let mut output = vec![0u8; get_maximum_output_size(input.len())];
977            let compressed_len = compress_into_with_dict(&input, &mut output, &dict).unwrap();
978
979            let mut uncompressed = vec![0u8; input.len()];
980            let uncompressed_len = crate::block::decompress::decompress_into_with_dict(
981                &output[..compressed_len],
982                &mut uncompressed,
983                &dict,
984            )
985            .unwrap();
986            uncompressed.truncate(uncompressed_len);
987            assert_eq!(uncompressed, input);
988        }
989    }
990
991    #[test]
992    #[cfg(all(miri, not(feature = "safe-encode")))]
993    fn miri_compress_into_with_short_dict_reads_past_dict() {
994        let input = [0u8; 13];
995        let dict = [0u8; 1];
996        let mut output = vec![0u8; get_maximum_output_size(input.len())];
997
998        let _ = compress_into_with_dict(&input, &mut output, &dict);
999    }
1000
1001    #[test]
1002    fn test_dict_match_crossing() {
1003        let input: &[u8] = &[
1004            10, 12, 14, 16, 18, 10, 12, 14, 16, 18, 10, 12, 14, 16, 18, 10, 12, 14, 16, 18,
1005        ];
1006        let dict = input;
1007        let compressed = compress_with_dict(input, dict);
1008        assert_lt!(compressed.len(), compress(input).len());
1009
1010        let mut uncompressed = vec![0u8; input.len() * 2];
1011        // copy first half of the input into output
1012        let dict_cutoff = dict.len() / 2;
1013        let output_start = dict.len() - dict_cutoff;
1014        uncompressed[..output_start].copy_from_slice(&dict[dict_cutoff..]);
1015        let uncomp_len = {
1016            let mut sink = SliceSink::new(&mut uncompressed[..], output_start);
1017            crate::block::decompress::decompress_internal::<true, _>(
1018                &compressed,
1019                &mut sink,
1020                &dict[..dict_cutoff],
1021            )
1022            .unwrap()
1023        };
1024        assert_eq!(input.len(), uncomp_len);
1025        assert_eq!(
1026            input,
1027            &uncompressed[output_start..output_start + uncomp_len]
1028        );
1029    }
1030
1031    #[test]
1032    fn test_conformant_last_block() {
1033        // From the spec:
1034        // The last match must start at least 12 bytes before the end of block.
1035        // The last match is part of the penultimate sequence. It is followed by the last sequence,
1036        // which contains only literals. Note that, as a consequence, an independent block <
1037        // 13 bytes cannot be compressed, because the match must copy "something",
1038        // so it needs at least one prior byte.
1039        // When a block can reference data from another block, it can start immediately with a match
1040        // and no literal, so a block of 12 bytes can be compressed.
1041        let aaas: &[u8] = b"aaaaaaaaaaaaaaa";
1042
1043        // incompressible
1044        let out = compress(&aaas[..12]);
1045        assert_gt!(out.len(), 12);
1046        // compressible
1047        let out = compress(&aaas[..13]);
1048        assert_le!(out.len(), 13);
1049        let out = compress(&aaas[..14]);
1050        assert_le!(out.len(), 14);
1051        let out = compress(&aaas[..15]);
1052        assert_le!(out.len(), 15);
1053
1054        // dict incompressible
1055        let out = compress_with_dict(&aaas[..11], aaas);
1056        assert_gt!(out.len(), 11);
1057        // compressible
1058        let out = compress_with_dict(&aaas[..12], aaas);
1059        // According to the spec this _could_ compress, but it doesn't in this lib
1060        // as it aborts compression for any input len < LZ4_MIN_LENGTH
1061        assert_gt!(out.len(), 12);
1062        let out = compress_with_dict(&aaas[..13], aaas);
1063        assert_le!(out.len(), 13);
1064        let out = compress_with_dict(&aaas[..14], aaas);
1065        assert_le!(out.len(), 14);
1066        let out = compress_with_dict(&aaas[..15], aaas);
1067        assert_le!(out.len(), 15);
1068    }
1069
1070    #[test]
1071    fn test_dict_size() {
1072        let dict = vec![b'a'; 1024 * 1024];
1073        let input = &b"aaaaaaaaaaaaaaaaaaaaaaaaaaaaa"[..];
1074        let compressed = compress_prepend_size_with_dict(input, &dict);
1075        let decompressed =
1076            crate::block::decompress_size_prepended_with_dict(&compressed, &dict).unwrap();
1077        assert_eq!(decompressed, input);
1078    }
1079}