// 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.

///|
struct Encoding {
  encoder : Map[BytesView, Int]
  decoder : FixedArray[Bytes]
  special_encoder : Map[String, Int]
  special_decoder : Map[Int, Bytes]
  regex : @pcre2.Code
}

///|
pub fn Encoding::new(
  mergeable_ranks~ : Map[Bytes, Int],
  special_tokens~ : Map[String, Int],
  pat_str~ : String,
) -> Encoding raise {
  let encoder : Map[BytesView, Int] = {}
  let decoder = FixedArray::make(mergeable_ranks.length(), b"")
  for piece, rank in mergeable_ranks {
    encoder[piece] = rank
    decoder[rank] = piece
  }
  let special_encoder = {}
  let special_decoder = {}
  for token, rank in special_tokens {
    special_encoder[token] = rank
    special_decoder[rank] = @encoding/utf8.encode(token)
  }
  let regex = @pcre2.compile(pat_str)
  regex.jit_compile(complete=true)
  Encoding::{ encoder, decoder, special_encoder, special_decoder, regex }
}

///|
pub fn Encoding::special_tokens(self : Encoding, piece : String) -> Int? {
  self.special_encoder.get(piece)
}

///|
fn arg_min(vec : Array[Int]) -> Int {
  let mut value = @int.max_value
  let mut index = -1
  for i = 0; i < vec.length(); i = i + 1 {
    if vec[i] < value {
      value = vec[i]
      index = i
    }
  }
  index
}

///|
fn Encoding::lookup(
  self : Encoding,
  piece : Bytes,
  tokens : Array[Int],
) -> Unit {
  for i in 0.. index
      None => @int.max_value
    }
    tokens.push(token)
  }
}

///|
fn Encoding::rank(self : Encoding, start : Int, tokens : Array[Int]) -> Int {
  let ranks = []
  for i = start; i < tokens.length() - 1; i = i + 1 {
    let l = tokens[i]
    let r = tokens[i + 1]
    let l = self.decoder[l]
    let r = self.decoder[r]
    let rank = match self.encoder.get(l + r) {
      Some(merge) => merge
      None => @int.max_value
    }
    ranks.push(rank)
  }
  arg_min(ranks)
}

///|
fn Encoding::merge(self : Encoding, start : Int, tokens : Array[Int]) -> Unit {
  for index = self.rank(start, tokens)
      index != -1
      index = self.rank(start, tokens) {
    let index = start + index
    let l = tokens[index]
    let r = tokens[index + 1]
    let l = self.decoder[l]
    let r = self.decoder[r]
    let merge = l + r
    let merge = self.encoder[merge]
    tokens[index] = merge
    tokens.remove(index + 1) |> ignore()
  }
}

///|
pub fn Encoding::encode(
  self : Encoding,
  piece : StringView,
) -> Array[Int] raise {
  let tokens = []
  let matches = self.regex.matches(piece)
  while matches.next() is Some(matched) {
    let piece = @encoding/utf8.encode(matched[0])
    match self.encoder.get(piece) {
      Some(token) => tokens.push(token)
      None => {
        let start = tokens.length()
        self.lookup(piece, tokens)
        self.merge(start, tokens)
      }
    }
  }
  tokens
}

///|
suberror DecodingError {
  InvalidToken(Int)
  MalformedUtf8(Bytes)
} derive(Show)

///|
pub fn Encoding::decode(
  self : Encoding,
  tokens : ArrayView[Int],
) -> String raise DecodingError {
  let buffer = @buffer.new()
  for token in tokens {
    if self.decoder.get(token) is Some(piece) {
      buffer.write_bytes(piece)
      continue
    }
    if self.special_decoder.get(token) is Some(piece) {
      buffer.write_bytes(piece)
      continue
    }
    raise DecodingError::InvalidToken(token)
  }
  let contents = buffer.contents()
  @encoding/utf8.decode(contents) catch {
    _ => raise DecodingError::MalformedUtf8(contents)
  }
}