///|
/// B/M/E/S states used by Chinese unknown-word recognition.
pub(all) enum ChineseHmmState {
  Begin
  Middle
  End
  Single
} derive(Eq, Compare, @debug.Debug)

///|
fn chinese_hmm_state_name(state : ChineseHmmState) -> String {
  match state {
    Begin => "B"
    Middle => "M"
    End => "E"
    Single => "S"
  }
}

///|
fn chinese_hmm_state_index(state : ChineseHmmState) -> Int {
  match state {
    Begin => 0
    Middle => 1
    End => 2
    Single => 3
  }
}

///|
fn chinese_hmm_state_from_index(index : Int) -> ChineseHmmState {
  match index {
    0 => Begin
    1 => Middle
    2 => End
    _ => Single
  }
}

///|
fn chinese_hmm_impossible_score() -> Double {
  -3.14e100
}

///|
fn chinese_hmm_safe_score(score : Double) -> Double {
  if score.is_nan() {
    chinese_hmm_impossible_score()
  } else {
    score
  }
}

///|
fn chinese_hmm_legal_start(state : ChineseHmmState) -> Bool {
  state == Begin || state == Single
}

///|
fn chinese_hmm_legal_transition(
  from_state : ChineseHmmState,
  to_state : ChineseHmmState,
) -> Bool {
  match from_state {
    Begin => to_state == Middle || to_state == End
    Middle => to_state == Middle || to_state == End
    End => to_state == Begin || to_state == Single
    Single => to_state == Begin || to_state == Single
  }
}

///|
/// Open HMM score provider. All values are natural-log scores.
pub(open) trait ChineseHmmModel {
  fn start_score(Self, ChineseHmmState) -> Double
  fn transition_score(Self, ChineseHmmState, ChineseHmmState) -> Double
  fn emission_score(Self, ChineseHmmState, Char) -> Double
}

///|
/// Structural errors raised while snapshotting one table-backed HMM model.
pub(all) suberror ChineseHmmModelError {
  IllegalStart(ChineseHmmState)
  IllegalTransition(ChineseHmmState, ChineseHmmState)
  DuplicateStart(ChineseHmmState)
  DuplicateTransition(ChineseHmmState, ChineseHmmState)
  DuplicateEmission(ChineseHmmState, Char)
  MissingStart(ChineseHmmState)
  MissingTransition(ChineseHmmState, ChineseHmmState)
  InvalidStartScore(ChineseHmmState)
  InvalidTransitionScore(ChineseHmmState, ChineseHmmState)
  InvalidEmissionScore(ChineseHmmState, Char)
  InvalidUnknownEmissionScore
} derive(Eq, @debug.Debug)

///|
pub impl Show for ChineseHmmModelError with fn output(self, logger) {
  match self {
    IllegalStart(state) =>
      logger.write_string(
        "Illegal Chinese HMM start state: \{chinese_hmm_state_name(state)}",
      )
    IllegalTransition(from_state, to_state) =>
      logger.write_string(
        "Illegal Chinese HMM transition: \{chinese_hmm_state_name(from_state)}->\{chinese_hmm_state_name(to_state)}",
      )
    DuplicateStart(state) =>
      logger.write_string(
        "Duplicate Chinese HMM start score: \{chinese_hmm_state_name(state)}",
      )
    DuplicateTransition(from_state, to_state) =>
      logger.write_string(
        "Duplicate Chinese HMM transition score: \{chinese_hmm_state_name(from_state)}->\{chinese_hmm_state_name(to_state)}",
      )
    DuplicateEmission(state, character) =>
      logger.write_string(
        "Duplicate Chinese HMM emission score: \{chinese_hmm_state_name(state)}/\{character}",
      )
    MissingStart(state) =>
      logger.write_string(
        "Missing Chinese HMM start score: \{chinese_hmm_state_name(state)}",
      )
    MissingTransition(from_state, to_state) =>
      logger.write_string(
        "Missing Chinese HMM transition score: \{chinese_hmm_state_name(from_state)}->\{chinese_hmm_state_name(to_state)}",
      )
    InvalidStartScore(state) =>
      logger.write_string(
        "Chinese HMM start score is NaN: \{chinese_hmm_state_name(state)}",
      )
    InvalidTransitionScore(from_state, to_state) =>
      logger.write_string(
        "Chinese HMM transition score is NaN: \{chinese_hmm_state_name(from_state)}->\{chinese_hmm_state_name(to_state)}",
      )
    InvalidEmissionScore(state, character) =>
      logger.write_string(
        "Chinese HMM emission score is NaN: \{chinese_hmm_state_name(state)}/\{character}",
      )
    InvalidUnknownEmissionScore =>
      logger.write_string("Chinese HMM unknown-emission score is NaN")
  }
}

///|
/// Immutable public facade over snapshotted HMM score tables.
pub struct TableChineseHmmModel {
  start_scores : ReadOnlyArray[Double]
  transition_scores : ReadOnlyArray[Double]
  emission_scores : ReadOnlyArray[@hashmap.HashMap[Char, Double]]
  unknown_emission_score : Double
}

///|
/// Builds a strict table model.
///
/// Both legal start states and all eight legal transitions must be supplied
/// exactly once. Emissions are sparse and fall back to
/// `unknown_emission_score`.
pub fn TableChineseHmmModel::new(
  start_scores : Array[(ChineseHmmState, Double)],
  transition_scores : Array[(ChineseHmmState, ChineseHmmState, Double)],
  emission_scores : Array[(ChineseHmmState, Char, Double)],
  unknown_emission_score : Double,
) -> TableChineseHmmModel raise ChineseHmmModelError {
  let starts = Array::make(4, chinese_hmm_impossible_score())
  let seen_starts = Array::make(4, false)
  for entry in start_scores {
    let (state, score) = entry
    guard chinese_hmm_legal_start(state) else {
      raise ChineseHmmModelError::IllegalStart(state)
    }
    let index = chinese_hmm_state_index(state)
    guard !seen_starts[index] else {
      raise ChineseHmmModelError::DuplicateStart(state)
    }
    guard !score.is_nan() else {
      raise ChineseHmmModelError::InvalidStartScore(state)
    }
    seen_starts[index] = true
    starts[index] = score
  }
  for state in [ChineseHmmState::Begin, ChineseHmmState::Single] {
    guard seen_starts[chinese_hmm_state_index(state)] else {
      raise ChineseHmmModelError::MissingStart(state)
    }
  }
  let transitions = Array::make(16, chinese_hmm_impossible_score())
  let seen_transitions = Array::make(16, false)
  for entry in transition_scores {
    let (from_state, to_state, score) = entry
    guard chinese_hmm_legal_transition(from_state, to_state) else {
      raise ChineseHmmModelError::IllegalTransition(from_state, to_state)
    }
    let index = chinese_hmm_state_index(from_state) * 4 +
      chinese_hmm_state_index(to_state)
    guard !seen_transitions[index] else {
      raise ChineseHmmModelError::DuplicateTransition(from_state, to_state)
    }
    guard !score.is_nan() else {
      raise ChineseHmmModelError::InvalidTransitionScore(from_state, to_state)
    }
    seen_transitions[index] = true
    transitions[index] = score
  }
  for
    from_state in [
      ChineseHmmState::Begin,
      ChineseHmmState::Middle,
      ChineseHmmState::End,
      ChineseHmmState::Single,
    ] {
    for
      to_state in [
        ChineseHmmState::Begin,
        ChineseHmmState::Middle,
        ChineseHmmState::End,
        ChineseHmmState::Single,
      ] {
      if chinese_hmm_legal_transition(from_state, to_state) {
        let index = chinese_hmm_state_index(from_state) * 4 +
          chinese_hmm_state_index(to_state)
        guard seen_transitions[index] else {
          raise ChineseHmmModelError::MissingTransition(from_state, to_state)
        }
      }
    }
  }
  let emissions : Array[@hashmap.HashMap[Char, Double]] = [
    @hashmap.HashMap([]),
    @hashmap.HashMap([]),
    @hashmap.HashMap([]),
    @hashmap.HashMap([]),
  ]
  for entry in emission_scores {
    let (state, character, score) = entry
    let index = chinese_hmm_state_index(state)
    guard !emissions[index].contains(character) else {
      raise ChineseHmmModelError::DuplicateEmission(state, character)
    }
    guard !score.is_nan() else {
      raise ChineseHmmModelError::InvalidEmissionScore(state, character)
    }
    emissions[index].set(character, score)
  }
  guard !unknown_emission_score.is_nan() else {
    raise ChineseHmmModelError::InvalidUnknownEmissionScore
  }
  {
    start_scores: ReadOnlyArray::from_array(starts),
    transition_scores: ReadOnlyArray::from_array(transitions),
    emission_scores: ReadOnlyArray::from_array(emissions),
    unknown_emission_score,
  }
}

///|
pub impl ChineseHmmModel for TableChineseHmmModel with fn start_score(
  self,
  state,
) {
  self.start_scores[chinese_hmm_state_index(state)]
}

///|
pub impl ChineseHmmModel for TableChineseHmmModel with fn transition_score(
  self,
  from_state,
  to_state,
) {
  self.transition_scores[chinese_hmm_state_index(from_state) * 4 +
  chinese_hmm_state_index(to_state)]
}

///|
pub impl ChineseHmmModel for TableChineseHmmModel with fn emission_score(
  self,
  state,
  character,
) {
  match self.emission_scores[chinese_hmm_state_index(state)].get(character) {
    Some(score) => score
    None => self.unknown_emission_score
  }
}

///|
fn chinese_hmm_previous_pair(state_index : Int) -> (Int, Int) {
  match state_index {
    0 => (2, 3)
    1 => (0, 1)
    2 => (0, 1)
    _ => (2, 3)
  }
}

///|
/// Runs constrained B/M/E/S Viterbi and returns token lengths covering the
/// complete `[start, end)` character range.
fn chinese_hmm_segment_lengths(
  model : &ChineseHmmModel,
  characters : ReadOnlyArray[Char],
  start : Int,
  end : Int,
) -> Array[Int] {
  let character_count = end - start
  if character_count <= 0 {
    return []
  }
  let backpointers : Array[Array[Int]] = [Array::make(4, -1)]
  let mut previous_scores = Array::make(4, chinese_hmm_impossible_score())
  let mut previous_reachable = Array::make(4, false)
  for state in [ChineseHmmState::Begin, ChineseHmmState::Single] {
    let state_index = chinese_hmm_state_index(state)
    previous_reachable[state_index] = true
    previous_scores[state_index] = chinese_hmm_safe_score(
        model.start_score(state),
      ) +
      chinese_hmm_safe_score(model.emission_score(state, characters[start]))
  }
  for offset in 1.. first_score) {
        current_reachable[state_index] = true
        current_scores[state_index] = second_score
        current_backpointers[state_index] = second_previous
      } else if previous_reachable[first_previous] {
        current_reachable[state_index] = true
        current_scores[state_index] = first_score
        current_backpointers[state_index] = first_previous
      }
    }
    backpointers.push(current_backpointers)
    previous_scores = current_scores
    previous_reachable = current_reachable
  }
  let mut final_state = 2
  if previous_reachable[3] &&
    (!previous_reachable[2] || previous_scores[3] >= previous_scores[2]) {
    final_state = 3
  }
  let state_path = Array::make(character_count, final_state)
  let mut offset = character_count - 1
  while offset > 0 {
    state_path[offset - 1] = backpointers[offset][state_path[offset]]
    offset -= 1
  }
  let lengths : Array[Int] = []
  let mut word_start = 0
  for index in 0.. word_start = index
      Middle => ()
      End => {
        lengths.push(index - word_start + 1)
        word_start = index + 1
      }
      Single => {
        lengths.push(1)
        word_start = index + 1
      }
    }
  }
  lengths
}

///|
fn chinese_dictionary_has_exact_match(
  dictionary : &ChineseDictionary,
  characters : ReadOnlyArray[Char],
  start : Int,
  end : Int,
) -> Bool {
  let expected_length = end - start
  for candidate in dictionary.matches_at(characters, start, end) {
    if candidate.length == expected_length && candidate.frequency > 0 {
      return true
    }
  }
  false
}