use crate::core::encoder::Encoder;
use std::cmp::Reverse;
use std::collections::BinaryHeap;
use super::encode::{Piece, Seed};
use super::nodes::Node;
use super::ranks::RankLookup;
pub(super) const SCAN_SYMBOL_LIMIT: usize = 64;
const GATE_ASCII: u8 = 64;
const GATE_MULTI: u8 = 16;
const fn build_byte_gate() -> [u8; 256] {
let mut gate = [GATE_MULTI; 256];
let mut b = 0;
while b < 0x80 {
gate[b] = GATE_ASCII;
b += 1;
}
gate[b' ' as usize] = GATE_MULTI;
gate[b'\t' as usize] = GATE_MULTI;
gate[b'\n' as usize] = GATE_MULTI;
gate[b'\r' as usize] = GATE_MULTI;
gate
}
static BYTE_GATE: [u8; 256] = build_byte_gate();
#[inline]
fn content_start(piece: &[u8]) -> usize {
match piece {
[0xE2, 0x96, 0x81, rest @ ..] if !rest.is_empty() => 3,
[0xC4, 0xA0, rest @ ..] if !rest.is_empty() => 2,
[ws, rest @ ..] if ws.is_ascii_whitespace() && !rest.is_empty() => 1,
_ => 0,
}
}
#[inline]
pub(super) fn prefers_scan(piece: &[u8], symbols: usize) -> bool {
if symbols > SCAN_SYMBOL_LIMIT {
return false;
}
if piece.len() <= GATE_MULTI as usize {
return true;
}
piece.len() <= BYTE_GATE[piece[content_start(piece)] as usize] as usize
}
#[inline]
fn rank_of(
piece: &[u8],
nodes: &[Node],
left: usize,
right: usize,
merge_ranks: RankLookup<'_>,
) -> (u32, u32) {
if left == usize::MAX || right == usize::MAX {
return (u32::MAX, u32::MAX);
}
if let Some(table) = merge_ranks.by_id() {
return merge_ranks.pair(table, nodes[left].id, nodes[right].id);
}
let (l, r) = (&nodes[left], &nodes[right]);
let len = l.len + r.len;
(merge_ranks.get(&piece[l.start..l.start + len]), u32::MAX)
}
#[inline]
fn absorb(nodes: &mut [Node], left: usize, right: usize, merged: u32) -> usize {
nodes[left].len += nodes[right].len;
nodes[left].id = merged;
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];
let mut merged = [u32::MAX; SCAN_SYMBOL_LIMIT];
for i in 0..count.saturating_sub(1) {
(ranks[i], merged[i]) = 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, merged[best]);
live -= 1;
ranks[right] = u32::MAX;
let prev = nodes[best].prev;
if prev != usize::MAX {
(ranks[prev], merged[prev]) = rank_of(piece, nodes, prev, best, merge_ranks);
}
(ranks[best], merged[best]) = rank_of(piece, nodes, best, new_next, merge_ranks);
}
}
#[inline]
fn key(rank: u32, left: usize) -> u64 {
((rank as u64) << 32) | (left as u64 & 0xFFFF_FFFF)
}
#[inline]
fn key_rank(key: u64) -> u32 {
(key >> 32) as u32
}
#[inline]
fn key_left(key: u64) -> usize {
(key & 0xFFFF_FFFF) as usize
}
#[derive(Default)]
pub(super) struct QueueScratch {
pub(super) ranks: Vec<u32>,
pub(super) merged: Vec<u32>,
pub(super) cold: Vec<u64>,
pub(super) hot: BinaryHeap<Reverse<u64>>,
}
impl QueueScratch {
pub(super) fn clear(&mut self) {
self.ranks.clear();
self.merged.clear();
self.cold.clear();
self.hot.clear();
}
#[cfg(test)]
pub(super) fn is_empty(&self) -> bool {
self.ranks.is_empty()
&& self.merged.is_empty()
&& self.cold.is_empty()
&& self.hot.is_empty()
}
}
fn merge_by_queue(
piece: &[u8],
nodes: &mut [Node],
merge_ranks: RankLookup<'_>,
q: &mut QueueScratch,
) -> usize {
let count = nodes.len();
q.ranks.clear();
q.ranks.reserve(count);
q.merged.clear();
q.merged.reserve(count);
q.cold.clear();
q.hot.clear();
for i in 0..count.saturating_sub(1) {
let (rank, merged) = rank_of(piece, nodes, i, i + 1, merge_ranks);
q.ranks.push(rank);
q.merged.push(merged);
if rank != u32::MAX {
q.cold.push(key(rank, i));
}
}
q.ranks.push(u32::MAX);
q.merged.push(u32::MAX);
q.cold.sort_unstable();
let mut cursor = 0usize;
let mut live = count;
loop {
while let Some(&k) = q.cold.get(cursor) {
if q.ranks[key_left(k)] == key_rank(k) {
break;
}
cursor += 1;
}
while let Some(&Reverse(k)) = q.hot.peek() {
if q.ranks[key_left(k)] == key_rank(k) {
break;
}
q.hot.pop();
}
let next = match (q.cold.get(cursor).copied(), q.hot.peek().map(|r| r.0)) {
(None, None) => return live,
(Some(cold), None) => {
cursor += 1;
cold
}
(None, Some(hot)) => {
q.hot.pop();
hot
}
(Some(cold), Some(hot)) => {
if cold < hot {
cursor += 1;
cold
} else {
q.hot.pop();
hot
}
}
};
let left = key_left(next);
let right = nodes[left].next;
let new_next = absorb(nodes, left, right, q.merged[left]);
live -= 1;
q.ranks[right] = u32::MAX;
let prev = nodes[left].prev;
if prev != usize::MAX {
let (rank, merged) = rank_of(piece, nodes, prev, left, merge_ranks);
q.ranks[prev] = rank;
q.merged[prev] = merged;
if rank != u32::MAX {
q.hot.push(Reverse(key(rank, prev)));
}
}
let (rank, merged) = rank_of(piece, nodes, left, new_next, merge_ranks);
q.ranks[left] = rank;
q.merged[left] = merged;
if rank != u32::MAX {
q.hot.push(Reverse(key(rank, left)));
}
}
}
fn link_and_merge(
piece: &[u8],
nodes: &mut [Node],
merge_ranks: RankLookup<'_>,
queue: &mut QueueScratch,
) -> usize {
let count = nodes.len();
debug_assert!(
nodes.first().is_none_or(|node| node.prev == usize::MAX)
&& nodes.last().is_none_or(|node| node.next == usize::MAX),
"nodes must arrive linked — see the seeding functions in `super::encode`"
);
if prefers_scan(piece, count) {
merge_by_scan(piece, nodes, merge_ranks)
} else {
merge_by_queue(piece, nodes, merge_ranks, queue)
}
}
#[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))
}
pub(super) fn merge_and_collect(
piece: &[u8],
nodes: &mut [Node],
merge_ranks: RankLookup<'_>,
id_encoder: &Encoder,
seeds: Option<&[Seed]>,
queue: &mut QueueScratch,
) -> Vec<Piece> {
link_and_merge(piece, nodes, merge_ranks, queue);
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_count(
piece: &[u8],
nodes: &mut [Node],
merge_ranks: RankLookup<'_>,
queue: &mut QueueScratch,
) -> usize {
link_and_merge(piece, nodes, merge_ranks, queue)
}
pub(super) fn merge_and_collect_ids_into(
piece: &[u8],
nodes: &mut [Node],
merge_ranks: RankLookup<'_>,
id_encoder: &Encoder,
out: &mut Vec<u32>,
queue: &mut QueueScratch,
) {
let live = link_and_merge(piece, nodes, merge_ranks, queue);
out.reserve(live);
let mut curr = 0;
if merge_ranks.by_id().is_some() {
while curr != usize::MAX {
out.push(nodes[curr].id);
curr = nodes[curr].next;
}
return;
}
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;
}
}