// 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_weight_fse_table_log_max = 6

///|
fn read_huffman_weights_fse(
  src : Bytes,
  start : Int,
  end_pos : Int,
) -> Array[Int] raise ZstdError {
  if start >= end_pos {
    raise CorruptionDetected
  }

  let (header_size, table_log, max_symbol, normalized_counter) = read_fse_ncount_header(
    src, start, end_pos, huf_max_nb_bits, huf_weight_fse_table_log_max,
  )

  let (next_state, nb_bits, symbols) = build_fse_symbol_decode_table(
    normalized_counter, max_symbol, table_log,
  )

  let payload_start = start + header_size
  if payload_start >= end_pos {
    raise CorruptionDetected
  }

  decode_fse_symbol_stream_reverse(
    src, payload_start, end_pos, table_log, next_state, nb_bits, symbols, 255,
  )
}

///|
fn decode_fse_symbol_stream_reverse(
  src : Bytes,
  start : Int,
  end_pos : Int,
  table_log : Int,
  next_state : Array[Int],
  nb_bits : Array[Int],
  symbols : Array[Int],
  max_output_size : Int,
) -> Array[Int] raise ZstdError {
  if table_log <= 0 ||
    start >= end_pos ||
    max_output_size <= 0 ||
    next_state.length() == 0 ||
    next_state.length() != nb_bits.length() ||
    next_state.length() != symbols.length() {
    raise CorruptionDetected
  }

  let table_size = (1 : Int) << table_log
  if table_size != next_state.length() {
    raise CorruptionDetected
  }

  let ds_ptr : Ref[Int] = { val: start }
  let ds_limit : Ref[Int] = { val: start + 8 }
  let ds_bits_consumed : Ref[Int] = { val: 0 }
  let ds_bit_container : Ref[UInt64] = { val: (0 : UInt64) }
  init_fse_dstream(
    src, start, end_pos, ds_ptr, ds_limit, ds_bits_consumed, ds_bit_container,
  )

  let state1_ref : Ref[Int] = {
    val: fse_read_bits(ds_bits_consumed, ds_bit_container, table_log).to_int(),
  }
  let state2_ref : Ref[Int] = {
    val: fse_read_bits(ds_bits_consumed, ds_bit_container, table_log).to_int(),
  }
  if state1_ref.val < 0 ||
    state1_ref.val >= table_size ||
    state2_ref.val < 0 ||
    state2_ref.val >= table_size {
    raise CorruptionDetected
  }

  if fse_reload_dstream(
      src, start, ds_ptr, ds_limit, ds_bits_consumed, ds_bit_container,
    ) ==
    fse_dstream_overflow {
    raise CorruptionDetected
  }

  let out : Array[Int] = Array::new()

  while true {
    let (symbol1, next1) = decode_fse_symbol_update(
      ds_bits_consumed,
      ds_bit_container,
      state1_ref.val,
      table_size,
      next_state,
      nb_bits,
      symbols,
    )
    out.push(symbol1)
    if out.length() > max_output_size {
      raise CorruptionDetected
    }
    state1_ref.val = next1

    if fse_reload_dstream(
        src, start, ds_ptr, ds_limit, ds_bits_consumed, ds_bit_container,
      ) ==
      fse_dstream_overflow {
      out.push(read_fse_table_symbol(state2_ref.val, table_size, symbols))
      break
    }

    let (symbol2, next2) = decode_fse_symbol_update(
      ds_bits_consumed,
      ds_bit_container,
      state2_ref.val,
      table_size,
      next_state,
      nb_bits,
      symbols,
    )
    out.push(symbol2)
    if out.length() > max_output_size {
      raise CorruptionDetected
    }
    state2_ref.val = next2

    if fse_reload_dstream(
        src, start, ds_ptr, ds_limit, ds_bits_consumed, ds_bit_container,
      ) ==
      fse_dstream_overflow {
      out.push(read_fse_table_symbol(state1_ref.val, table_size, symbols))
      break
    }
  }

  if out.length() <= 0 || out.length() > max_output_size {
    raise CorruptionDetected
  }
  out
}

///|
fn decode_fse_symbol_update(
  bits_consumed : Ref[Int],
  bit_container : Ref[UInt64],
  state : Int,
  table_size : Int,
  next_state : Array[Int],
  nb_bits : Array[Int],
  symbols : Array[Int],
) -> (Int, Int) raise ZstdError {
  let symbol = read_fse_table_symbol(state, table_size, symbols)

  let bit_count = nb_bits[state]
  if bit_count < 0 || bit_count > 24 {
    raise CorruptionDetected
  }
  let add = fse_read_bits(bits_consumed, bit_container, bit_count).to_int()

  let next = next_state[state] + add
  if next < 0 || next >= table_size {
    raise CorruptionDetected
  }
  (symbol, next)
}

///|
fn read_fse_table_symbol(
  state : Int,
  table_size : Int,
  symbols : Array[Int],
) -> Int raise ZstdError {
  if state < 0 || state >= table_size || state >= symbols.length() {
    raise CorruptionDetected
  }
  symbols[state]
}

///|
let fse_dstream_unfinished = 0

///|
let fse_dstream_end_of_buffer = 1

///|
let fse_dstream_completed = 2

///|
let fse_dstream_overflow = 3

///|
fn init_fse_dstream(
  src : Bytes,
  start : Int,
  end_pos : Int,
  ptr : Ref[Int],
  limit : Ref[Int],
  bits_consumed : Ref[Int],
  bit_container : Ref[UInt64],
) -> Unit raise ZstdError {
  if start < 0 || end_pos <= start || end_pos > src.length() {
    raise CorruptionDetected
  }
  let src_size = end_pos - start
  let last_byte = src[end_pos - 1].to_uint().reinterpret_as_int()
  if last_byte == 0 {
    raise CorruptionDetected
  }

  limit.val = start + 8
  if src_size >= 8 {
    ptr.val = end_pos - 8
    bit_container.val = read_u64_le(src, ptr.val)
    bits_consumed.val = 8 - high_bit_positive(last_byte)
  } else {
    ptr.val = start
    bit_container.val = read_u64_le_padded(src, start, end_pos)
    bits_consumed.val = 8 - high_bit_positive(last_byte) + (8 - src_size) * 8
  }
}

///|
fn read_u64_le_padded(src : Bytes, start : Int, end_pos : Int) -> UInt64 {
  let mut value : UInt64 = 0
  let mut i = 0
  while start + i < end_pos && i < 8 {
    value = value + (src[start + i].to_uint().to_uint64() << (i * 8))
    i = i + 1
  }
  value
}

///|
fn fse_read_bits(
  bits_consumed : Ref[Int],
  bit_container : Ref[UInt64],
  count : Int,
) -> UInt64 raise ZstdError {
  if count < 0 || count > 24 {
    raise CorruptionDetected
  }
  if count == 0 {
    return (0 : UInt64)
  }
  let start = (64 - bits_consumed.val - count) & 63
  let mask = ((1 : UInt64) << count) - (1 : UInt64)
  let value = (bit_container.val >> start) & mask
  bits_consumed.val = bits_consumed.val + count
  value
}

///|
fn fse_reload_dstream(
  src : Bytes,
  start : Int,
  ptr : Ref[Int],
  limit : Ref[Int],
  bits_consumed : Ref[Int],
  bit_container : Ref[UInt64],
) -> Int raise ZstdError {
  if bits_consumed.val > 64 {
    return fse_dstream_overflow
  }
  if ptr.val < start {
    raise CorruptionDetected
  }

  if ptr.val >= limit.val {
    let moved = bits_consumed.val >> 3
    ptr.val = ptr.val - moved
    if ptr.val < start {
      raise CorruptionDetected
    }
    bits_consumed.val = bits_consumed.val & 7
    bit_container.val = read_u64_le(src, ptr.val)
    return fse_dstream_unfinished
  }

  if ptr.val == start {
    if bits_consumed.val < 64 {
      return fse_dstream_end_of_buffer
    }
    return fse_dstream_completed
  }

  let mut nb_bytes = bits_consumed.val >> 3
  let mut result = fse_dstream_unfinished
  if ptr.val - nb_bytes < start {
    nb_bytes = ptr.val - start
    result = fse_dstream_end_of_buffer
  }
  ptr.val = ptr.val - nb_bytes
  bits_consumed.val = bits_consumed.val - nb_bytes * 8
  bit_container.val = read_u64_le(src, ptr.val)
  result
}