// Copyright 2025 International Digital Economy Academy
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
///|
let sequence_table_kind_rle = 0
///|
let sequence_table_kind_predefined = 1
///|
let sequence_table_kind_compressed = 2
///|
fn decode_compressed_block_minimal(
src : Bytes,
block_start : Int,
block_size : Int,
out : Array[Byte],
frame_out_start : Int,
history : Bytes,
rep1 : Int,
rep2 : Int,
rep3 : Int,
prev_huf_valid : Ref[Bool],
prev_huf_max_bits : Ref[Int],
prev_huf_left : Ref[Array[Int]],
prev_huf_right : Ref[Array[Int]],
prev_huf_symbol : Ref[Array[Int]],
prev_ll_valid : Ref[Bool],
prev_ll_kind : Ref[Int],
prev_ll_code : Ref[UInt],
prev_ll_table_log : Ref[Int],
prev_ll_table_next_state : Ref[Array[Int]],
prev_ll_table_nb_add_bits : Ref[Array[Int]],
prev_ll_table_nb_bits : Ref[Array[Int]],
prev_ll_table_base_values : Ref[Array[Int]],
prev_off_valid : Ref[Bool],
prev_off_kind : Ref[Int],
prev_off_code : Ref[UInt],
prev_off_table_log : Ref[Int],
prev_off_table_next_state : Ref[Array[Int]],
prev_off_table_nb_add_bits : Ref[Array[Int]],
prev_off_table_nb_bits : Ref[Array[Int]],
prev_off_table_base_values : Ref[Array[Int]],
prev_ml_valid : Ref[Bool],
prev_ml_kind : Ref[Int],
prev_ml_code : Ref[UInt],
prev_ml_table_log : Ref[Int],
prev_ml_table_next_state : Ref[Array[Int]],
prev_ml_table_nb_add_bits : Ref[Array[Int]],
prev_ml_table_nb_bits : Ref[Array[Int]],
prev_ml_table_base_values : Ref[Array[Int]],
window_size? : Int = 0,
) -> (Int, UInt64, Int, Int, Int) raise ZstdError {
let src_len = src.length()
if frame_out_start < 0 || frame_out_start > out.length() {
raise CorruptionDetected
}
ensure_range(src_len, block_start, block_size)
let block_end = block_start + block_size
let (seq_pos0, literals) = decode_literals_section_minimal(
src, block_start, block_end, prev_huf_valid, prev_huf_max_bits, prev_huf_left,
prev_huf_right, prev_huf_symbol,
)
let (number_of_sequences, seq_pos0_end) = parse_sequence_count(
src, seq_pos0, block_end,
)
if number_of_sequences == 0 {
if seq_pos0_end != block_end {
raise CorruptionDetected
}
append_bytes(out, literals, 0, literals.length())
return (seq_pos0_end, literals.length().to_uint64(), rep1, rep2, rep3)
}
ensure_range(src_len, seq_pos0_end, 1)
let modes = src[seq_pos0_end].to_uint()
if (modes & 0x3) != 0 {
raise CorruptionDetected
}
let ll_mode = (modes >> 6) & 0x3
let off_mode = (modes >> 4) & 0x3
let ml_mode = (modes >> 2) & 0x3
let seq_pos_ref : Ref[Int] = { val: seq_pos0_end + 1 }
let (ll_kind, ll_code) = decode_sequence_source(
src, block_end, seq_pos_ref, ll_mode, prev_ll_valid, prev_ll_kind, prev_ll_code,
prev_ll_table_log, prev_ll_table_next_state, prev_ll_table_nb_add_bits, prev_ll_table_nb_bits,
prev_ll_table_base_values, sequence_symbol_literal_length,
)
let (off_kind, off_code) = decode_sequence_source(
src, block_end, seq_pos_ref, off_mode, prev_off_valid, prev_off_kind, prev_off_code,
prev_off_table_log, prev_off_table_next_state, prev_off_table_nb_add_bits, prev_off_table_nb_bits,
prev_off_table_base_values, sequence_symbol_offset,
)
let (ml_kind, ml_code) = decode_sequence_source(
src, block_end, seq_pos_ref, ml_mode, prev_ml_valid, prev_ml_kind, prev_ml_code,
prev_ml_table_log, prev_ml_table_next_state, prev_ml_table_nb_add_bits, prev_ml_table_nb_bits,
prev_ml_table_base_values, sequence_symbol_match_length,
)
let seq_pos = seq_pos_ref.val
let has_sequence_bitstream = seq_pos < block_end
if (
ll_kind != sequence_table_kind_rle ||
off_kind != sequence_table_kind_rle ||
ml_kind != sequence_table_kind_rle
) &&
!has_sequence_bitstream {
raise CorruptionDetected
}
let br_start = seq_pos
let br_byte : Ref[Int] = { val: block_end - 1 }
let br_bit : Ref[Int] = { val: -1 }
if has_sequence_bitstream {
init_reverse_bit_reader(src, br_start, br_byte, br_bit)
}
let ll_state : Ref[Int] = { val: 0 }
let off_state : Ref[Int] = { val: 0 }
let ml_state : Ref[Int] = { val: 0 }
if ll_kind == sequence_table_kind_predefined {
ll_state.val = init_fse_state_reverse(
src, br_start, br_byte, br_bit, ll_predefined_table_log,
)
} else if ll_kind == sequence_table_kind_compressed {
if prev_ll_table_log.val <= 0 {
raise CorruptionDetected
}
ll_state.val = init_fse_state_reverse(
src,
br_start,
br_byte,
br_bit,
prev_ll_table_log.val,
)
}
if off_kind == sequence_table_kind_predefined {
off_state.val = init_fse_state_reverse(
src, br_start, br_byte, br_bit, of_predefined_table_log,
)
} else if off_kind == sequence_table_kind_compressed {
if prev_off_table_log.val <= 0 {
raise CorruptionDetected
}
off_state.val = init_fse_state_reverse(
src,
br_start,
br_byte,
br_bit,
prev_off_table_log.val,
)
}
if ml_kind == sequence_table_kind_predefined {
ml_state.val = init_fse_state_reverse(
src, br_start, br_byte, br_bit, ml_predefined_table_log,
)
} else if ml_kind == sequence_table_kind_compressed {
if prev_ml_table_log.val <= 0 {
raise CorruptionDetected
}
ml_state.val = init_fse_state_reverse(
src,
br_start,
br_byte,
br_bit,
prev_ml_table_log.val,
)
}
let mut literal_pos = 0
let mut r1 = rep1
let mut r2 = rep2
let mut r3 = rep3
let mut produced : UInt64 = 0
let seq_count = number_of_sequences.reinterpret_as_int()
let mut seq_idx = 0
while seq_idx < seq_count {
let is_last_sequence = seq_idx + 1 == seq_count
let (ll_next_state, ll_nb_add_bits, ll_nb_state_bits, ll_base) = sequence_symbol_entry(
ll_kind,
ll_code,
ll_state.val,
prev_ll_table_next_state.val,
prev_ll_table_nb_add_bits.val,
prev_ll_table_nb_bits.val,
prev_ll_table_base_values.val,
sequence_symbol_literal_length,
)
let (ml_next_state, ml_nb_add_bits, ml_nb_state_bits, ml_base) = sequence_symbol_entry(
ml_kind,
ml_code,
ml_state.val,
prev_ml_table_next_state.val,
prev_ml_table_nb_add_bits.val,
prev_ml_table_nb_bits.val,
prev_ml_table_base_values.val,
sequence_symbol_match_length,
)
let (off_next_state, off_nb_add_bits, off_nb_state_bits, off_base) = sequence_symbol_entry(
off_kind,
off_code,
off_state.val,
prev_off_table_next_state.val,
prev_off_table_nb_add_bits.val,
prev_off_table_nb_bits.val,
prev_off_table_base_values.val,
sequence_symbol_offset,
)
let ll0 = ll_base == 0
let offset = if off_nb_add_bits > 1 {
if !has_sequence_bitstream {
raise CorruptionDetected
}
let off_extra = read_reverse_bits(
src, br_start, br_byte, br_bit, off_nb_add_bits,
)
let value = off_base + off_extra.reinterpret_as_int()
if value <= 0 {
raise CorruptionDetected
}
r3 = r2
r2 = r1
r1 = value
value
} else if off_nb_add_bits == 0 {
let value = if ll0 { r2 } else { r1 }
if ll0 {
r2 = r1
}
r1 = value
value
} else {
if !has_sequence_bitstream {
raise CorruptionDetected
}
let low = read_reverse_bits(src, br_start, br_byte, br_bit, 1).reinterpret_as_int()
let offset_code = off_base + (if ll0 { 1 } else { 0 }) + low
let value = if offset_code == 1 {
r2
} else if offset_code == 2 {
r3
} else if offset_code == 3 {
r1 - 1
} else {
raise CorruptionDetected
}
if value <= 0 {
raise CorruptionDetected
}
if offset_code != 1 {
r3 = r2
}
r2 = r1
r1 = value
value
}
let ml_extra = if ml_nb_add_bits == 0 {
(0 : UInt)
} else if has_sequence_bitstream {
read_reverse_bits(src, br_start, br_byte, br_bit, ml_nb_add_bits)
} else {
raise CorruptionDetected
}
let ll_extra = if ll_nb_add_bits == 0 {
(0 : UInt)
} else if has_sequence_bitstream {
read_reverse_bits(src, br_start, br_byte, br_bit, ll_nb_add_bits)
} else {
raise CorruptionDetected
}
let lit_len = ll_base + ll_extra.reinterpret_as_int()
let match_len = ml_base + ml_extra.reinterpret_as_int()
if literal_pos + lit_len > literals.length() {
raise CorruptionDetected
}
append_bytes(out, literals, literal_pos, lit_len)
literal_pos = literal_pos + lit_len
produced = produced + lit_len.to_uint64()
let mut match_i = 0
while match_i < match_len {
let value = read_match_byte(
out,
frame_out_start,
history,
offset,
window_size~,
)
out.push(value)
match_i = match_i + 1
}
produced = produced + match_len.to_uint64()
if !is_last_sequence {
if ll_kind != sequence_table_kind_rle {
ll_state.val = update_fse_state_reverse(
src, br_start, br_byte, br_bit, ll_next_state, ll_nb_state_bits,
)
}
if ml_kind != sequence_table_kind_rle {
ml_state.val = update_fse_state_reverse(
src, br_start, br_byte, br_bit, ml_next_state, ml_nb_state_bits,
)
}
if off_kind != sequence_table_kind_rle {
off_state.val = update_fse_state_reverse(
src, br_start, br_byte, br_bit, off_next_state, off_nb_state_bits,
)
}
}
seq_idx = seq_idx + 1
}
let tail_len = literals.length() - literal_pos
append_bytes(out, literals, literal_pos, tail_len)
produced = produced + tail_len.to_uint64()
if has_sequence_bitstream && !reverse_bits_consumed(br_start, br_byte, br_bit) {
raise CorruptionDetected
}
(block_end, produced, r1, r2, r3)
}
///|
fn read_match_byte(
out : Array[Byte],
frame_out_start : Int,
history : Bytes,
offset : Int,
window_size? : Int = 0,
) -> Byte raise ZstdError {
if offset <= 0 || frame_out_start < 0 || frame_out_start > out.length() {
raise CorruptionDetected
}
if window_size > 0 && offset > window_size {
raise CorruptionDetected
}
let produced = out.length() - frame_out_start
let history_len = history.length()
if offset > produced + history_len {
raise CorruptionDetected
}
let from_tail = offset - 1
if from_tail < produced {
return out[out.length() - 1 - from_tail]
}
let history_from_tail = from_tail - produced
history[history_len - 1 - history_from_tail]
}
///|
fn sequence_symbol_entry(
kind : Int,
code : UInt,
state : Int,
table_next_state : Array[Int],
table_nb_add_bits : Array[Int],
table_nb_bits : Array[Int],
table_base_values : Array[Int],
symbol_type : Int,
) -> (Int, Int, Int, Int) raise ZstdError {
if kind == sequence_table_kind_rle {
let (base, nb_add_bits) = sequence_symbol_base_additional_bits(
symbol_type,
code.reinterpret_as_int(),
)
(0, nb_add_bits, 0, base)
} else if kind == sequence_table_kind_predefined {
if symbol_type == sequence_symbol_literal_length {
ll_predefined_entry(state)
} else if symbol_type == sequence_symbol_offset {
of_predefined_entry(state)
} else if symbol_type == sequence_symbol_match_length {
ml_predefined_entry(state)
} else {
raise CorruptionDetected
}
} else if kind == sequence_table_kind_compressed {
fse_table_entry(
table_next_state, table_nb_add_bits, table_nb_bits, table_base_values, state,
)
} else {
raise CorruptionDetected
}
}
///|
fn fse_table_entry(
table_next_state : Array[Int],
table_nb_add_bits : Array[Int],
table_nb_bits : Array[Int],
table_base_values : Array[Int],
state : Int,
) -> (Int, Int, Int, Int) raise ZstdError {
if state < 0 ||
state >= table_next_state.length() ||
state >= table_nb_add_bits.length() ||
state >= table_nb_bits.length() ||
state >= table_base_values.length() {
raise CorruptionDetected
}
(
table_next_state[state],
table_nb_add_bits[state],
table_nb_bits[state],
table_base_values[state],
)
}
///|
fn decode_sequence_source(
src : Bytes,
block_end : Int,
seq_pos_ref : Ref[Int],
mode : UInt,
prev_valid : Ref[Bool],
prev_kind : Ref[Int],
prev_code : Ref[UInt],
prev_table_log : Ref[Int],
prev_table_next_state : Ref[Array[Int]],
prev_table_nb_add_bits : Ref[Array[Int]],
prev_table_nb_bits : Ref[Array[Int]],
prev_table_base_values : Ref[Array[Int]],
symbol_type : Int,
) -> (Int, UInt) raise ZstdError {
if mode == 1 {
if seq_pos_ref.val >= block_end {
raise CorruptionDetected
}
ensure_range(src.length(), seq_pos_ref.val, 1)
let code = src[seq_pos_ref.val].to_uint()
seq_pos_ref.val = seq_pos_ref.val + 1
prev_valid.val = true
prev_kind.val = sequence_table_kind_rle
prev_code.val = code
(sequence_table_kind_rle, code)
} else if mode == 0 {
prev_valid.val = true
prev_kind.val = sequence_table_kind_predefined
prev_code.val = 0
(sequence_table_kind_predefined, (0 : UInt))
} else if mode == 2 {
let (
table_header_size,
table_log,
table_next_state,
table_nb_add_bits,
table_nb_bits,
table_base_values,
) = build_sequence_fse_table_from_header(
src,
seq_pos_ref.val,
block_end,
symbol_type,
)
seq_pos_ref.val = seq_pos_ref.val + table_header_size
prev_valid.val = true
prev_kind.val = sequence_table_kind_compressed
prev_code.val = 0
prev_table_log.val = table_log
prev_table_next_state.val = table_next_state
prev_table_nb_add_bits.val = table_nb_add_bits
prev_table_nb_bits.val = table_nb_bits
prev_table_base_values.val = table_base_values
(sequence_table_kind_compressed, (0 : UInt))
} else if mode == 3 {
if !prev_valid.val {
raise CorruptionDetected
}
if prev_kind.val == sequence_table_kind_compressed &&
prev_table_log.val <= 0 {
raise CorruptionDetected
}
(prev_kind.val, prev_code.val)
} else {
raise CorruptionDetected
}
}
///|
fn init_fse_state_reverse(
src : Bytes,
br_start : Int,
br_byte : Ref[Int],
br_bit : Ref[Int],
table_log : Int,
) -> Int raise ZstdError {
read_reverse_bits(src, br_start, br_byte, br_bit, table_log).reinterpret_as_int()
}
///|
fn update_fse_state_reverse(
src : Bytes,
br_start : Int,
br_byte : Ref[Int],
br_bit : Ref[Int],
next_state : Int,
nb_bits : Int,
) -> Int raise ZstdError {
next_state +
read_reverse_bits(src, br_start, br_byte, br_bit, nb_bits).reinterpret_as_int()
}
///|
fn offset_base_from_code(code : UInt) -> Int raise ZstdError {
if code == 0 {
0
} else if code == 1 {
1
} else if code <= 31 {
(((1 : UInt64) << code.reinterpret_as_int()) - (3 : UInt64)).to_int()
} else {
raise CorruptionDetected
}
}
///|
fn parse_sequence_count(
src : Bytes,
seq_start : Int,
block_end : Int,
) -> (UInt, Int) raise ZstdError {
if seq_start >= block_end {
raise CorruptionDetected
}
let byte0 = src[seq_start].to_uint()
let mut seq_pos = seq_start + 1
let count : UInt = if byte0 < 128 {
byte0
} else if byte0 < 255 {
ensure_range(src.length(), seq_pos, 1)
let byte1 = src[seq_pos].to_uint()
seq_pos = seq_pos + 1
((byte0 - 0x80) << 8) + byte1
} else {
ensure_range(src.length(), seq_pos, 2)
let byte1 = src[seq_pos].to_uint()
let byte2 = src[seq_pos + 1].to_uint()
seq_pos = seq_pos + 2
byte1 + (byte2 << 8) + 0x7F00
}
(count, seq_pos)
}