use crate::core::token_bytes::Encoder;
use std::collections::BinaryHeap;
use super::encode::{Piece, Seed};
use super::nodes::{Merge, Node};
use super::ranks::RankLookup;
pub(super) const SCAN_SYMBOL_LIMIT: usize = 64;
#[inline]
fn rank_of(
piece: &[u8],
nodes: &[Node],
left: usize,
right: usize,
merge_ranks: RankLookup<'_>,
) -> u32 {
if left == usize::MAX || right == usize::MAX {
return u32::MAX;
}
let (l, r) = (&nodes[left], &nodes[right]);
let len = l.len + r.len;
merge_ranks.get(&piece[l.start..l.start + len])
}
#[inline]
fn absorb(nodes: &mut [Node], left: usize, right: usize) -> usize {
nodes[left].len += nodes[right].len;
nodes[right].len = 0;
let new_next = nodes[right].next;
nodes[left].next = new_next;
if new_next != usize::MAX {
nodes[new_next].prev = left;
}
new_next
}
fn merge_by_scan(piece: &[u8], nodes: &mut [Node], merge_ranks: RankLookup<'_>) -> usize {
let count = nodes.len();
let mut ranks = [u32::MAX; SCAN_SYMBOL_LIMIT];
for (i, rank) in ranks.iter_mut().enumerate().take(count.saturating_sub(1)) {
*rank = rank_of(piece, nodes, i, i + 1, merge_ranks);
}
let mut live = count;
loop {
let mut best_rank = u32::MAX;
let mut best = usize::MAX;
for (i, &rank) in ranks.iter().enumerate().take(count) {
if rank < best_rank {
best_rank = rank;
best = i;
}
}
if best_rank == u32::MAX {
return live;
}
let right = nodes[best].next;
let new_next = absorb(nodes, best, right);
live -= 1;
ranks[right] = u32::MAX;
let prev = nodes[best].prev;
if prev != usize::MAX {
ranks[prev] = rank_of(piece, nodes, prev, best, merge_ranks);
}
ranks[best] = rank_of(piece, nodes, best, new_next, merge_ranks);
}
}
fn merge_by_heap(piece: &[u8], nodes: &mut [Node], merge_ranks: RankLookup<'_>) -> usize {
let count = nodes.len();
let push = |queue: &mut BinaryHeap<Merge>, left: usize, right: usize, nodes: &[Node]| {
let rank = rank_of(piece, nodes, left, right, merge_ranks);
if rank != u32::MAX {
queue.push(Merge {
left,
right,
rank,
len: nodes[left].len + nodes[right].len,
});
}
};
let mut queue: BinaryHeap<Merge> = BinaryHeap::with_capacity(count * 2);
for i in 0..count.saturating_sub(1) {
push(&mut queue, i, i + 1, nodes);
}
let mut live = count;
while let Some(candidate) = queue.pop() {
let (li, ri) = (candidate.left, candidate.right);
let (left, right) = (nodes[li], nodes[ri]);
if left.len == 0 || left.next != ri || left.len + right.len != candidate.len {
continue;
}
let new_next = absorb(nodes, li, ri);
live -= 1;
let prev = nodes[li].prev;
if prev != usize::MAX {
push(&mut queue, prev, li, nodes);
}
if new_next != usize::MAX {
push(&mut queue, li, new_next, nodes);
}
}
live
}
fn link_and_merge(piece: &[u8], nodes: &mut [Node], merge_ranks: RankLookup<'_>) -> usize {
let count = nodes.len();
for (i, node) in nodes.iter_mut().enumerate() {
node.prev = if i == 0 { usize::MAX } else { i - 1 };
node.next = if i + 1 == count { usize::MAX } else { i + 1 };
}
if count <= SCAN_SYMBOL_LIMIT {
merge_by_scan(piece, nodes, merge_ranks)
} else {
merge_by_heap(piece, nodes, merge_ranks)
}
}
#[inline]
fn resolve(
slice: &[u8],
node: &Node,
index: usize,
id_encoder: &Encoder,
seeds: Option<&[Seed]>,
) -> Option<u32> {
seeds
.and_then(|seeds| seeds.get(index))
.filter(|seed| seed.len == node.len)
.and_then(|seed| seed.id)
.or_else(|| id_encoder.get(slice).copied())
}
pub(super) fn merge_and_collect(
piece: &[u8],
mut nodes: Vec<Node>,
merge_ranks: RankLookup<'_>,
id_encoder: &Encoder,
seeds: Option<&[Seed]>,
) -> Vec<Piece> {
link_and_merge(piece, &mut nodes, merge_ranks);
let mut result: Vec<Piece> = Vec::with_capacity(piece.len());
let mut curr = 0;
let mut unresolved_run: Option<(usize, usize)> = None;
while curr != usize::MAX {
let node = &nodes[curr];
let slice = &piece[node.start..node.start + node.len];
if let Some(id) = resolve(slice, node, curr, id_encoder, seeds) {
if let Some((start, len)) = unresolved_run.take() {
result.push(Piece::Unresolved { start, len });
}
result.push(Piece::Token(id));
} else {
for (offset, &byte) in slice.iter().enumerate() {
if let Some(&id) = id_encoder.get(&[byte][..]) {
if let Some((start, len)) = unresolved_run.take() {
result.push(Piece::Unresolved { start, len });
}
result.push(Piece::Token(id));
} else {
let byte_start = node.start + offset;
match &mut unresolved_run {
Some((start, len)) if *start + *len == byte_start => *len += 1,
_ => {
if let Some((start, len)) = unresolved_run.take() {
result.push(Piece::Unresolved { start, len });
}
unresolved_run = Some((byte_start, 1));
}
}
}
}
}
curr = nodes[curr].next;
}
if let Some((start, len)) = unresolved_run.take() {
result.push(Piece::Unresolved { start, len });
}
result
}
pub(super) fn merge_and_collect_ids_into(
piece: &[u8],
nodes: &mut [Node],
merge_ranks: RankLookup<'_>,
id_encoder: &Encoder,
out: &mut Vec<u32>,
) {
let live = link_and_merge(piece, nodes, merge_ranks);
out.reserve(live);
let mut curr = 0;
while curr != usize::MAX {
let node = &nodes[curr];
let slice = &piece[node.start..node.start + node.len];
match id_encoder.get(slice) {
Some(&id) => out.push(id),
None => out.extend(slice.iter().filter_map(|b| id_encoder.get(&[*b][..]))),
}
curr = nodes[curr].next;
}
}