///|
const TOKENIZER_KIND_CHAR : String = "char"

///|
const TOKENIZER_KIND_WORD : String = "word"

///|
const TOKENIZER_KIND_BPE : String = "bpe"

///|
pub const DEFAULT_BPE_VOCAB_SIZE : Int = 512

///|
pub(all) enum TokenizerConfig {
  CharacterLevel
  WordLevel
  BpeLevel(Int)
} derive(Eq, Debug)

///|
pub(all) suberror TokenizerError {
  UnknownCharacter(Char)
  UnknownToken(String)
  InvalidTokenId(Int)
} derive(Eq, Debug)

///|
pub struct CharTokenizer {
  priv chars : Array[Char]
  priv char_to_id : Map[Char, Int]
}

///|
pub struct WordTokenizer {
  priv tokens : Array[String]
  priv token_to_id : Map[String, Int]
}

///|
pub struct BpeMerge {
  priv left : String
  priv right : String
  priv merged : String
}

///|
pub fn BpeMerge::BpeMerge(
  left : String,
  right : String,
  merged : String,
) -> BpeMerge {
  { left, right, merged }
}

///|
pub fn BpeMerge::left(self : BpeMerge) -> String {
  self.left
}

///|
pub fn BpeMerge::right(self : BpeMerge) -> String {
  self.right
}

///|
pub fn BpeMerge::merged(self : BpeMerge) -> String {
  self.merged
}

///|
pub struct BpeTokenizer {
  priv tokens : Array[String]
  priv token_to_id : Map[String, Int]
  priv merges : Array[BpeMerge]
}

///|
pub(all) enum Tokenizer {
  Character(CharTokenizer)
  Word(WordTokenizer)
  Bpe(BpeTokenizer)
}

///|
pub struct TokenDataset {
  priv tokenizer : Tokenizer
  priv train_ids : Array[Int]
  priv val_ids : Array[Int]
}

///|
pub fn TokenDataset::tokenizer(self : TokenDataset) -> Tokenizer {
  self.tokenizer
}

///|
pub fn TokenDataset::train_ids(self : TokenDataset) -> Array[Int] {
  self.train_ids.copy()
}

///|
pub fn TokenDataset::val_ids(self : TokenDataset) -> Array[Int] {
  self.val_ids.copy()
}

///|
fn collect_sorted_chars(text : String) -> Array[Char] {
  let seen : Map[Char, Unit] = {}
  for ch in text {
    seen[ch] = ()
  }
  let chars : Array[Char] = []
  for ch, _ in seen {
    chars.push(ch)
  }
  chars.sort()
  chars
}

///|
fn char_map(chars : Array[Char]) -> Map[Char, Int] {
  let char_to_id : Map[Char, Int] = {}
  for i in 0.. Array[String] {
  let tokens : Array[String] = []
  let buf = StringBuilder()
  fn flush() -> Unit {
    if !buf.is_empty() {
      tokens.push(buf.to_string())
      buf.reset()
    }
  }
  for ch in text {
    match ch {
      '\n' => {
        flush()
        tokens.push("\n")
      }
      ' ' | '\t' | '\r' => flush()
      _ => buf.write_char(ch)
    }
  }
  flush()
  tokens
}

///|
fn collect_sorted_tokens(tokens : Array[String]) -> Array[String] {
  let seen : Map[String, Unit] = {}
  for token in tokens {
    seen[token] = ()
  }
  let vocabulary : Array[String] = []
  for token, _ in seen {
    vocabulary.push(token)
  }
  vocabulary.sort()
  vocabulary
}

///|
fn token_map(tokens : Array[String]) -> Map[String, Int] {
  let token_to_id : Map[String, Int] = {}
  for i in 0.. Unit {
  if tokens.length() == 0 {
    abort("tokenizer vocabulary must not be empty")
  }
  let seen : Map[String, Unit] = {}
  for token in tokens {
    if token == "" {
      abort("tokenizer vocabulary must not contain empty tokens")
    }
    if seen.contains(token) {
      abort("tokenizer vocabulary must not contain duplicates")
    }
    seen[token] = ()
  }
}

///|
fn split_char_tokens(text : String) -> Array[String] {
  let tokens : Array[String] = []
  for ch in text {
    tokens.push(ch.to_string())
  }
  tokens
}

///|
fn join_word_tokens(tokens : Array[String]) -> String {
  let buf = StringBuilder()
  let mut needs_space = false
  for token in tokens {
    if token == "\n" {
      buf.write_char('\n')
      needs_space = false
    } else {
      if needs_space {
        buf.write_char(' ')
      }
      buf.write_string(token)
      needs_space = true
    }
  }
  buf.to_string()
}

///|
fn join_tokens(tokens : Array[String]) -> String {
  let buf = StringBuilder()
  for token in tokens {
    buf.write_string(token)
  }
  buf.to_string()
}

///|
fn pair_key(left : String, right : String) -> String {
  left + "\u{1f}" + right
}

///|
priv struct BpePair {
  left : String
  right : String
  merged : String
}

///|
fn best_bpe_pair(
  tokens : Array[String],
  token_to_id : Map[String, Int],
) -> BpePair? {
  if tokens.length() < 2 {
    return None
  }
  let counts : Map[String, Int] = {}
  let lefts : Map[String, String] = {}
  let rights : Map[String, String] = {}
  for i in 0..<(tokens.length() - 1) {
    let left = tokens[i]
    let right = tokens[i + 1]
    let merged = left + right
    if !token_to_id.contains(merged) {
      let key = pair_key(left, right)
      counts[key] = counts.get_or_default(key, 0) + 1
      lefts[key] = left
      rights[key] = right
    }
  }
  let mut best_key = ""
  let mut best_count = 0
  for key, count in counts {
    if count > best_count ||
      (count == best_count && best_count > 0 && key.compare(best_key) < 0) {
      best_key = key
      best_count = count
    }
  }
  if best_count < 2 {
    return None
  }
  let left = lefts[best_key]
  let right = rights[best_key]
  Some({ left, right, merged: left + right })
}

///|
fn apply_bpe_merge(
  tokens : Array[String],
  left : String,
  right : String,
  merged : String,
) -> Array[String] {
  let out : Array[String] = []
  let mut i = 0
  while i < tokens.length() {
    if i + 1 < tokens.length() && tokens[i] == left && tokens[i + 1] == right {
      out.push(merged)
      i += 2
    } else {
      out.push(tokens[i])
      i += 1
    }
  }
  out
}

///|
pub fn CharTokenizer::from_chars(chars : Array[Char]) -> CharTokenizer {
  if chars.length() == 0 {
    abort("character tokenizer vocabulary must not be empty")
  }
  let sorted = chars.copy()
  sorted.sort()
  for i in 1.. (CharTokenizer, Array[Int]) {
  if text.char_length() == 0 {
    abort("character tokenizer training text must not be empty")
  }
  let tokenizer = CharTokenizer::from_chars(collect_sorted_chars(text))
  let ids = tokenizer.encode(text) catch {
    UnknownCharacter(_) => abort("trained tokenizer rejected training text")
    UnknownToken(_) => abort("unexpected word tokenizer error")
    InvalidTokenId(_) => abort("unexpected token id error while encoding text")
  }
  (tokenizer, ids)
}

///|
pub fn BpeTokenizer::from_vocabulary_and_merges(
  vocabulary : Array[String],
  merges : Array[BpeMerge],
) -> BpeTokenizer {
  validate_ordered_tokens(vocabulary)
  let token_to_id = token_map(vocabulary)
  for merge in merges {
    if !token_to_id.contains(merge.left) ||
      !token_to_id.contains(merge.right) ||
      !token_to_id.contains(merge.merged) {
      abort("BPE merge rule must reference vocabulary tokens")
    }
  }
  { tokens: vocabulary.copy(), token_to_id, merges: merges.copy() }
}

///|
pub fn BpeTokenizer::train(
  text : String,
  vocab_size : Int,
) -> (BpeTokenizer, Array[Int]) {
  if vocab_size <= 0 {
    abort("BPE vocabulary size must be positive")
  }
  let mut current = split_char_tokens(text)
  if current.length() == 0 {
    abort("BPE tokenizer training text must not be empty")
  }
  let vocabulary : Array[String] = []
  for ch in collect_sorted_chars(text) {
    vocabulary.push(ch.to_string())
  }
  let token_to_id = token_map(vocabulary)
  let merges : Array[BpeMerge] = []
  while vocabulary.length() < vocab_size {
    match best_bpe_pair(current, token_to_id) {
      Some(pair) => {
        current = apply_bpe_merge(current, pair.left, pair.right, pair.merged)
        token_to_id[pair.merged] = vocabulary.length()
        vocabulary.push(pair.merged)
        merges.push(BpeMerge(pair.left, pair.right, pair.merged))
      }
      None => break
    }
  }
  let tokenizer = BpeTokenizer::from_vocabulary_and_merges(vocabulary, merges)
  let ids = tokenizer.encode(text) catch {
    UnknownToken(_) => abort("trained tokenizer rejected training text")
    UnknownCharacter(_) => abort("trained tokenizer rejected training text")
    InvalidTokenId(_) => abort("unexpected token id error while encoding text")
  }
  (tokenizer, ids)
}

///|
pub fn BpeTokenizer::kind_name(_self : BpeTokenizer) -> String {
  TOKENIZER_KIND_BPE
}

///|
pub fn BpeTokenizer::vocab_size(self : BpeTokenizer) -> Int {
  self.tokens.length()
}

///|
pub fn BpeTokenizer::vocabulary(self : BpeTokenizer) -> Array[String] {
  self.tokens.copy()
}

///|
pub fn BpeTokenizer::merges(self : BpeTokenizer) -> Array[BpeMerge] {
  self.merges.copy()
}

///|
pub fn BpeTokenizer::encode(
  self : BpeTokenizer,
  text : String,
) -> Array[Int] raise TokenizerError {
  let mut tokens : Array[String] = []
  for ch in text {
    let token = ch.to_string()
    if !self.token_to_id.contains(token) {
      raise UnknownCharacter(ch)
    }
    tokens.push(token)
  }
  for merge in self.merges {
    tokens = apply_bpe_merge(tokens, merge.left, merge.right, merge.merged)
  }
  let ids : Array[Int] = []
  for token in tokens {
    match self.token_to_id.get(token) {
      Some(id) => ids.push(id)
      None => raise UnknownToken(token)
    }
  }
  ids
}

///|
pub fn BpeTokenizer::decode(
  self : BpeTokenizer,
  ids : Array[Int],
) -> String raise TokenizerError {
  let tokens : Array[String] = []
  for id in ids {
    match self.tokens.get(id) {
      Some(token) => tokens.push(token)
      None => raise InvalidTokenId(id)
    }
  }
  join_tokens(tokens)
}

///|
pub fn WordTokenizer::from_tokens(tokens : Array[String]) -> WordTokenizer {
  if tokens.length() == 0 {
    abort("word tokenizer vocabulary must not be empty")
  }
  let sorted = tokens.copy()
  sorted.sort()
  for i in 0.. 0 && sorted[i - 1] == sorted[i] {
      abort("word tokenizer vocabulary must not contain duplicates")
    }
  }
  { tokens: sorted, token_to_id: token_map(sorted) }
}

///|
pub fn WordTokenizer::train(text : String) -> (WordTokenizer, Array[Int]) {
  let tokens = split_word_tokens(text)
  if tokens.length() == 0 {
    abort("word tokenizer training text must not be empty")
  }
  let tokenizer = WordTokenizer::from_tokens(collect_sorted_tokens(tokens))
  let ids = tokenizer.encode(text) catch {
    UnknownToken(_) => abort("trained tokenizer rejected training text")
    UnknownCharacter(_) => abort("unexpected character tokenizer error")
    InvalidTokenId(_) => abort("unexpected token id error while encoding text")
  }
  (tokenizer, ids)
}

///|
pub fn WordTokenizer::kind_name(_self : WordTokenizer) -> String {
  TOKENIZER_KIND_WORD
}

///|
pub fn WordTokenizer::vocab_size(self : WordTokenizer) -> Int {
  self.tokens.length()
}

///|
pub fn WordTokenizer::vocabulary(self : WordTokenizer) -> Array[String] {
  self.tokens.copy()
}

///|
pub fn WordTokenizer::encode(
  self : WordTokenizer,
  text : String,
) -> Array[Int] raise TokenizerError {
  let ids : Array[Int] = []
  for token in split_word_tokens(text) {
    match self.token_to_id.get(token) {
      Some(id) => ids.push(id)
      None => raise UnknownToken(token)
    }
  }
  ids
}

///|
pub fn WordTokenizer::decode(
  self : WordTokenizer,
  ids : Array[Int],
) -> String raise TokenizerError {
  let tokens : Array[String] = []
  for id in ids {
    match self.tokens.get(id) {
      Some(token) => tokens.push(token)
      None => raise InvalidTokenId(id)
    }
  }
  join_word_tokens(tokens)
}

///|
pub fn Tokenizer::from_char(tokenizer : CharTokenizer) -> Tokenizer {
  Character(tokenizer)
}

///|
pub fn Tokenizer::from_word(tokenizer : WordTokenizer) -> Tokenizer {
  Word(tokenizer)
}

///|
pub fn Tokenizer::from_bpe(tokenizer : BpeTokenizer) -> Tokenizer {
  Bpe(tokenizer)
}

///|
pub fn Tokenizer::train_char(text : String) -> (Tokenizer, Array[Int]) {
  let (tokenizer, ids) = CharTokenizer::train(text)
  (Tokenizer::from_char(tokenizer), ids)
}

///|
pub fn Tokenizer::train_word(text : String) -> (Tokenizer, Array[Int]) {
  let (tokenizer, ids) = WordTokenizer::train(text)
  (Tokenizer::from_word(tokenizer), ids)
}

///|
pub fn Tokenizer::train_bpe(
  text : String,
  vocab_size : Int,
) -> (Tokenizer, Array[Int]) {
  let (tokenizer, ids) = BpeTokenizer::train(text, vocab_size)
  (Tokenizer::from_bpe(tokenizer), ids)
}

///|
pub fn Tokenizer::train(
  text : String,
  config : TokenizerConfig,
) -> (Tokenizer, Array[Int]) {
  match config {
    CharacterLevel => Tokenizer::train_char(text)
    WordLevel => Tokenizer::train_word(text)
    BpeLevel(vocab_size) => Tokenizer::train_bpe(text, vocab_size)
  }
}

///|
pub fn Tokenizer::kind_name(self : Tokenizer) -> String {
  match self {
    Character(tokenizer) => tokenizer.kind_name()
    Word(tokenizer) => tokenizer.kind_name()
    Bpe(tokenizer) => tokenizer.kind_name()
  }
}

///|
pub fn Tokenizer::vocab_size(self : Tokenizer) -> Int {
  match self {
    Character(tokenizer) => tokenizer.vocab_size()
    Word(tokenizer) => tokenizer.vocab_size()
    Bpe(tokenizer) => tokenizer.vocab_size()
  }
}

///|
pub fn Tokenizer::vocabulary(self : Tokenizer) -> Array[String] {
  match self {
    Character(tokenizer) => {
      let vocabulary : Array[String] = []
      for ch in tokenizer.vocabulary() {
        vocabulary.push(ch.to_string())
      }
      vocabulary
    }
    Word(tokenizer) => tokenizer.vocabulary()
    Bpe(tokenizer) => tokenizer.vocabulary()
  }
}

///|
pub fn Tokenizer::bpe_merges(self : Tokenizer) -> Array[BpeMerge] {
  match self {
    Bpe(tokenizer) => tokenizer.merges()
    _ => []
  }
}

///|
pub fn Tokenizer::encode(
  self : Tokenizer,
  text : String,
) -> Array[Int] raise TokenizerError {
  match self {
    Character(tokenizer) => tokenizer.encode(text)
    Word(tokenizer) => tokenizer.encode(text)
    Bpe(tokenizer) => tokenizer.encode(text)
  }
}

///|
pub fn Tokenizer::decode(
  self : Tokenizer,
  ids : Array[Int],
) -> String raise TokenizerError {
  match self {
    Character(tokenizer) => tokenizer.decode(ids)
    Word(tokenizer) => tokenizer.decode(ids)
    Bpe(tokenizer) => tokenizer.decode(ids)
  }
}

///|
pub fn CharTokenizer::kind_name(_self : CharTokenizer) -> String {
  TOKENIZER_KIND_CHAR
}

///|
pub fn CharTokenizer::vocab_size(self : CharTokenizer) -> Int {
  self.chars.length()
}

///|
pub fn CharTokenizer::vocabulary(self : CharTokenizer) -> Array[Char] {
  self.chars.copy()
}

///|
pub fn CharTokenizer::encode(
  self : CharTokenizer,
  text : String,
) -> Array[Int] raise TokenizerError {
  let ids : Array[Int] = []
  for ch in text {
    match self.char_to_id.get(ch) {
      Some(id) => ids.push(id)
      None => raise UnknownCharacter(ch)
    }
  }
  ids
}

///|
pub fn CharTokenizer::decode(
  self : CharTokenizer,
  ids : Array[Int],
) -> String raise TokenizerError {
  let chars : Array[Char] = []
  for id in ids {
    match self.chars.get(id) {
      Some(ch) => chars.push(ch)
      None => raise InvalidTokenId(id)
    }
  }
  String::from_array(chars)
}

///|
pub fn prepare_token_dataset(
  text : String,
  config : TokenizerConfig,
) -> TokenDataset {
  let (tokenizer, ids) = Tokenizer::train(text, config)
  let split = ids.length() * 9 / 10
  {
    tokenizer,
    train_ids: ids[0:split].to_owned(),
    val_ids: ids[split:ids.length()].to_owned(),
  }
}