// 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 huf_max_nb_bits = 12
///|
fn read_huffman_tree_description(
src : Bytes,
start : Int,
end_pos : Int,
) -> (Int, Int, Array[Int], Array[Int], Array[Int]) raise ZstdError {
if start >= end_pos {
raise CorruptionDetected
}
let header_byte = src[start].to_uint().reinterpret_as_int()
let weights : Array[Int] = if header_byte < 128 {
if start + 1 + header_byte > end_pos {
raise CorruptionDetected
}
read_huffman_weights_fse(src, start + 1, start + 1 + header_byte)
} else {
let number_of_weights = header_byte - 127
if number_of_weights <= 0 || number_of_weights > 128 {
raise CorruptionDetected
}
let packed_size = (number_of_weights + 1) >> 1
if start + 1 + packed_size > end_pos {
raise CorruptionDetected
}
let parsed_weights : Array[Int] = Array::new()
let mut i = 0
while i < number_of_weights {
parsed_weights.push(0)
i = i + 1
}
i = 0
while i < number_of_weights {
let b = src[start + 1 + (i >> 1)].to_uint().reinterpret_as_int()
if (i & 1) == 0 {
parsed_weights[i] = (b >> 4) & 0xF
} else {
parsed_weights[i] = b & 0xF
}
i = i + 1
}
parsed_weights
}
let consumed_size = if header_byte < 128 {
1 + header_byte
} else {
1 + ((weights.length() + 1) >> 1)
}
let mut weight_total = 0
let mut non_zero_weights = 0
let mut i = 0
while i < weights.length() {
let w = weights[i]
if w < 0 || w > huf_max_nb_bits {
raise CorruptionDetected
}
if w > 0 {
weight_total = weight_total + ((1 : Int) << (w - 1))
non_zero_weights = non_zero_weights + 1
}
i = i + 1
}
if weight_total <= 0 {
raise CorruptionDetected
}
let max_bits = high_bit_positive(weight_total) + 1
if max_bits <= 0 || max_bits > huf_max_nb_bits {
raise CorruptionDetected
}
let total_weight = (1 : Int) << max_bits
let rest = total_weight - weight_total
if rest <= 0 || !is_power_of_two(rest) {
raise CorruptionDetected
}
let last_weight = high_bit_positive(rest) + 1
if last_weight <= 0 || last_weight > huf_max_nb_bits {
raise CorruptionDetected
}
weights.push(last_weight)
non_zero_weights = non_zero_weights + 1
if non_zero_weights < 2 {
raise CorruptionDetected
}
let (left, right, symbol) = build_huffman_tree_from_weights(weights, max_bits)
(consumed_size, max_bits, left, right, symbol)
}
///|
fn build_huffman_tree_from_weights(
weights : Array[Int],
max_bits : Int,
) -> (Array[Int], Array[Int], Array[Int]) raise ZstdError {
let nb_bits : Array[Int] = Array::new()
let mut i = 0
while i < weights.length() {
let w = weights[i]
let n = if w > 0 { max_bits + 1 - w } else { 0 }
if n < 0 || n > max_bits {
raise CorruptionDetected
}
nb_bits.push(n)
i = i + 1
}
let symbols_by_weight : Array[Int] = Array::new()
let mut weight = 1
while weight <= max_bits {
i = 0
while i < weights.length() {
if weights[i] == weight {
symbols_by_weight.push(i)
}
i = i + 1
}
weight = weight + 1
}
if symbols_by_weight.length() < 2 {
raise CorruptionDetected
}
let codes : Array[Int] = Array::new()
i = 0
while i < nb_bits.length() {
codes.push(0)
i = i + 1
}
let first_symbol = symbols_by_weight[0]
let first_bits = nb_bits[first_symbol]
if first_bits <= 0 {
raise CorruptionDetected
}
let code_ref : Ref[Int] = { val: 0 }
let current_bits_ref : Ref[Int] = { val: first_bits }
i = 0
while i < symbols_by_weight.length() {
let s = symbols_by_weight[i]
let n = nb_bits[s]
if n <= 0 {
raise CorruptionDetected
}
if n > current_bits_ref.val {
raise CorruptionDetected
}
if n < current_bits_ref.val {
code_ref.val = code_ref.val >> (current_bits_ref.val - n)
current_bits_ref.val = n
}
codes[s] = code_ref.val
code_ref.val = code_ref.val + 1
i = i + 1
}
let left : Array[Int] = Array::new()
let right : Array[Int] = Array::new()
let symbol : Array[Int] = Array::new()
ignore(add_huffman_node(left, right, symbol))
i = 0
while i < nb_bits.length() {
let n = nb_bits[i]
if n > 0 {
insert_huffman_code(left, right, symbol, codes[i], n, i)
}
i = i + 1
}
(left, right, symbol)
}
///|
fn add_huffman_node(
left : Array[Int],
right : Array[Int],
symbol : Array[Int],
) -> Int {
left.push(-1)
right.push(-1)
symbol.push(-1)
left.length() - 1
}
///|
fn insert_huffman_code(
left : Array[Int],
right : Array[Int],
symbol : Array[Int],
code : Int,
nb_bits : Int,
value : Int,
) -> Unit raise ZstdError {
let mut node = 0
let mut bit_index = nb_bits - 1
while bit_index >= 0 {
let bit = (code >> bit_index) & 1
let last_bit = bit_index == 0
if bit == 0 {
let next = left[node]
if last_bit {
if next != -1 {
raise CorruptionDetected
}
let leaf = add_huffman_node(left, right, symbol)
symbol[leaf] = value
left[node] = leaf
} else if next == -1 {
let created = add_huffman_node(left, right, symbol)
left[node] = created
node = created
} else {
node = next
}
} else {
let next = right[node]
if last_bit {
if next != -1 {
raise CorruptionDetected
}
let leaf = add_huffman_node(left, right, symbol)
symbol[leaf] = value
right[node] = leaf
} else if next == -1 {
let created = add_huffman_node(left, right, symbol)
right[node] = created
node = created
} else {
node = next
}
}
bit_index = bit_index - 1
}
}
///|
fn decode_huffman_single_stream(
src : Bytes,
start : Int,
end_pos : Int,
output_size : Int,
max_bits : Int,
left : Array[Int],
right : Array[Int],
symbol : Array[Int],
) -> Bytes raise ZstdError {
ignore(max_bits)
if output_size < 0 || start >= end_pos {
raise CorruptionDetected
}
let out : Array[Byte] = Array::new()
let br_byte : Ref[Int] = { val: end_pos - 1 }
let br_bit : Ref[Int] = { val: -1 }
init_reverse_bit_reader(src, start, br_byte, br_bit)
let mut i = 0
while i < output_size {
let decoded = decode_one_huffman_symbol(
src, start, br_byte, br_bit, left, right, symbol,
)
out.push(decoded.reinterpret_as_uint().to_byte())
i = i + 1
}
if !reverse_bits_consumed(start, br_byte, br_bit) {
raise CorruptionDetected
}
Bytes::from_array(out)
}
///|
fn decode_huffman_four_streams(
src : Bytes,
start : Int,
end_pos : Int,
output_size : Int,
max_bits : Int,
left : Array[Int],
right : Array[Int],
symbol : Array[Int],
) -> Bytes raise ZstdError {
if output_size < 0 || start + 6 > end_pos {
raise CorruptionDetected
}
let stream1_size = src[start].to_uint().reinterpret_as_int() +
(src[start + 1].to_uint().reinterpret_as_int() << 8)
let stream2_size = src[start + 2].to_uint().reinterpret_as_int() +
(src[start + 3].to_uint().reinterpret_as_int() << 8)
let stream3_size = src[start + 4].to_uint().reinterpret_as_int() +
(src[start + 5].to_uint().reinterpret_as_int() << 8)
let streams_start = start + 6
let remaining = end_pos - streams_start
let stream4_size = remaining - stream1_size - stream2_size - stream3_size
if stream1_size <= 0 ||
stream2_size <= 0 ||
stream3_size <= 0 ||
stream4_size <= 0 {
raise CorruptionDetected
}
let stream1_start = streams_start
let stream1_end = stream1_start + stream1_size
let stream2_start = stream1_end
let stream2_end = stream2_start + stream2_size
let stream3_start = stream2_end
let stream3_end = stream3_start + stream3_size
let stream4_start = stream3_end
let stream4_end = stream4_start + stream4_size
if stream4_end != end_pos {
raise CorruptionDetected
}
let segment_size = (output_size + 3) / 4
let out1_size = segment_size
let out2_size = segment_size
let out3_size = segment_size
let out4_size = output_size - 3 * segment_size
if out4_size < 0 {
raise CorruptionDetected
}
let out1 = decode_huffman_single_stream(
src, stream1_start, stream1_end, out1_size, max_bits, left, right, symbol,
)
let out2 = decode_huffman_single_stream(
src, stream2_start, stream2_end, out2_size, max_bits, left, right, symbol,
)
let out3 = decode_huffman_single_stream(
src, stream3_start, stream3_end, out3_size, max_bits, left, right, symbol,
)
let out4 = decode_huffman_single_stream(
src, stream4_start, stream4_end, out4_size, max_bits, left, right, symbol,
)
let out : Array[Byte] = Array::new()
append_bytes(out, out1, 0, out1.length())
append_bytes(out, out2, 0, out2.length())
append_bytes(out, out3, 0, out3.length())
append_bytes(out, out4, 0, out4.length())
Bytes::from_array(out)
}
///|
fn decode_one_huffman_symbol(
src : Bytes,
start : Int,
br_byte : Ref[Int],
br_bit : Ref[Int],
left : Array[Int],
right : Array[Int],
symbol : Array[Int],
) -> Int raise ZstdError {
let mut node = 0
while true {
if node < 0 ||
node >= left.length() ||
node >= right.length() ||
node >= symbol.length() {
raise CorruptionDetected
}
let s = symbol[node]
if s >= 0 {
return s
}
let bit = read_reverse_bit(src, start, br_byte, br_bit)
node = if bit == 0 { left[node] } else { right[node] }
}
raise CorruptionDetected
}
///|
fn is_power_of_two(value : Int) -> Bool {
value > 0 && (value & (value - 1)) == 0
}