use serde::Deserialize;
use std::collections::HashMap;
use std::env;
use std::fs::{self, File};
use std::io::{BufWriter, Write};
use std::path::Path;
#[derive(Deserialize)]
struct Tokens {
single_byte: Vec<String>,
double_byte: Vec<Vec<String>>,
}
fn token_key(b: &[u8]) -> u32 {
match *b {
[a] => a as u32,
[a, b] => u32::from_le_bytes([a, b, 0, 0]),
[a, b, c] => u32::from_le_bytes([a, b, c, 0]),
[a, b, c, d] => u32::from_le_bytes([a, b, c, d]),
_ => {
let len = b.len();
let head = u32::from_le_bytes([b[0], b[1], b[2], b[3]]);
let tail = u32::from_le_bytes([b[len - 4], b[len - 3], b[len - 2], b[len - 1]]);
head.rotate_left(11) ^ tail ^ (len as u32).wrapping_mul(0x9E37_79B1)
}
}
}
fn build_table(entries: &[(u32, u16, u16)]) -> (u32, Vec<u32>, Vec<u16>, Vec<u16>) {
let mut k = usize::BITS - (entries.len() * 4 / 3).leading_zeros();
loop {
let size = 1usize << k;
let mut mul: u32 = 0x9E37_79B1;
for _ in 0..256 {
let mut keys = vec![0u32; size];
let mut meta = vec![0u16; size];
let mut off = vec![0u16; size];
'insert: {
for &(x, m, o) in entries {
let mut h = (x.wrapping_mul(mul) >> (32 - k)) as usize;
let mut disp = 0usize;
while keys[h] != 0 {
h = (h + 1) & (size - 1);
disp += 1;
if disp > 8 {
break 'insert;
}
}
keys[h] = x;
meta[h] = m;
off[h] = o;
}
return (mul, keys, meta, off);
}
mul = mul.wrapping_mul(0x85EB_CA6B) | 1;
}
k += 1;
}
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
println!("cargo:rerun-if-changed=src/tokens.json");
println!("cargo:rerun-if-changed=build.rs");
let path = Path::new(&env::var("OUT_DIR")?).join("token_maps.rs");
let mut file = BufWriter::new(File::create(&path)?);
let tokens_json = fs::read_to_string("src/tokens.json")?;
let tokens: Tokens = serde_json::from_str(&tokens_json)?;
let mut values: Vec<(String, u16)> = Vec::new();
let mut seen: HashMap<String, u16> = HashMap::new();
for (i, token) in tokens.single_byte.iter().enumerate() {
if !token.is_empty() {
let kind = u16::try_from(i)?;
assert!(kind < 256, "single-byte token index out of range");
if let Some(existing) = seen.get(token) {
panic!("duplicate token {:?}: {} vs {}", token, existing, kind);
}
seen.insert(token.clone(), kind);
values.push((token.clone(), kind));
}
}
for (dict_idx, dict) in tokens.double_byte.iter().enumerate() {
for (token_idx, token) in dict.iter().enumerate() {
if !token.is_empty() {
assert!(token_idx < 256, "double-byte token index out of range");
assert!(dict_idx + 1 < 256, "double-byte dict index out of range");
let kind = u16::try_from((dict_idx + 1) * 256 + token_idx)?;
if let Some(existing) = seen.get(token) {
panic!("duplicate token {:?}: {} vs {}", token, existing, kind);
}
seen.insert(token.clone(), kind);
values.push((token.clone(), kind));
}
}
}
let max_len = values.iter().map(|(t, _)| t.len()).max().unwrap_or(0);
assert!(max_len < 256, "length no longer fits blob prefix byte");
let mut blob: Vec<u8> = Vec::new();
let mut entries: Vec<(u32, u16, u16)> = Vec::new();
for (t, k) in &values {
let b = t.as_bytes();
assert!(!b.contains(&0), "token {t:?} contains NUL");
assert!(*k < (1 << 11), "kind overflows meta word");
let x = token_key(b);
assert!(x != 0, "token key collides with empty-slot sentinel");
let (tag, off) = if b.len() <= 4 {
(b.len() as u16 - 1, 0)
} else {
let o = u16::try_from(blob.len())?;
blob.push(b.len() as u8);
blob.extend_from_slice(b);
(31, o)
};
entries.push((x, *k | (tag << 11), off));
}
let (tok_mul, tok_keys, tok_meta, tok_off) = build_table(&entries);
writeln!(file, "const MAX_TOKEN_LEN: usize = {max_len};")?;
writeln!(file, "const TOK_MUL: u32 = {tok_mul};")?;
writeln!(
file,
"const TOK_SHIFT: u32 = {};",
32 - tok_keys.len().trailing_zeros()
)?;
writeln!(file, "const TOK_MASK: usize = {};", tok_keys.len() - 1)?;
writeln!(
file,
"static TOK_KEYS: [u32; {}] = {:?};",
tok_keys.len(),
tok_keys
)?;
writeln!(
file,
"static TOK_META: [u16; {}] = {:?};",
tok_meta.len(),
tok_meta
)?;
writeln!(
file,
"static TOK_OFF: [u16; {}] = {:?};",
tok_off.len(),
tok_off
)?;
writeln!(file, "static TOK_BLOB: [u8; {}] = {:?};", blob.len(), blob)?;
writeln!(file, "\nstatic SINGLE_BYTE_TOKENS: &[&str] = &[")?;
for token in &tokens.single_byte {
writeln!(file, " {:?},", token)?;
}
writeln!(file, "];")?;
writeln!(file, "\nstatic DOUBLE_BYTE_TOKENS: &[&[&str]] = &[")?;
for dict in &tokens.double_byte {
writeln!(file, " &[")?;
for token in dict {
writeln!(file, " {:?},", token)?;
}
writeln!(file, " ],")?;
}
writeln!(file, "];")?;
Ok(())
}