use std::fmt;
const BASE: i64 = 36;
const TMIN: i64 = 1;
const TMAX: i64 = 26;
const INITIAL_BIAS: i64 = 72;
const DAMP: i64 = 700;
const SKEW: i64 = 38;
const DIGITS: &[u8; 36] = b"abcdefghijklmnopqrstuvwxyz0123456789";
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum PunycodeError {
NotAscii,
Incomplete,
InvalidDigit,
InvalidCodePoint,
}
impl fmt::Display for PunycodeError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let msg = match self {
PunycodeError::NotAscii => "input is not ASCII",
PunycodeError::Incomplete => "incomplete punycode string",
PunycodeError::InvalidDigit => "invalid extended code point",
PunycodeError::InvalidCodePoint => "invalid decoded character",
};
f.write_str(msg)
}
}
impl std::error::Error for PunycodeError {}
fn threshold(j: i64, bias: i64) -> i64 {
(BASE * (j + 1) - bias).clamp(TMIN, TMAX)
}
fn adapt(mut delta: i64, first: bool, numchars: i64) -> i64 {
delta /= if first { DAMP } else { 2 };
delta += delta / numchars;
let mut divisions = 0;
while delta > ((BASE - TMIN) * TMAX) / 2 {
delta /= BASE - TMIN;
divisions += BASE;
}
divisions + (BASE * delta) / (delta + SKEW)
}
fn segregate(text: &[char]) -> (Vec<u8>, Vec<char>) {
let mut base = Vec::new();
let mut extended = std::collections::BTreeSet::new();
for &c in text {
if (c as u32) < 128 {
base.push(c as u8);
} else {
extended.insert(c);
}
}
(base, extended.into_iter().collect())
}
fn selective_len(text: &[char], max: u32) -> i64 {
text.iter().filter(|&&c| (c as u32) < max).count() as i64
}
fn selective_find(text: &[char], target: char, mut index: i64, mut pos: i64) -> (i64, i64) {
let len = text.len() as i64;
loop {
pos += 1;
if pos == len {
return (-1, -1);
}
let c = text[pos as usize];
if c == target {
return (index + 1, pos);
} else if c < target {
index += 1;
}
}
}
fn insertion_unsort(text: &[char], extended: &[char]) -> Vec<i64> {
let mut oldchar: i64 = 0x80;
let mut oldindex: i64 = -1;
let mut result = Vec::new();
for &c in extended {
let mut index: i64 = -1;
let mut pos: i64 = -1;
let ch = c as i64;
let curlen = selective_len(text, c as u32);
let mut delta = (curlen + 1) * (ch - oldchar);
loop {
let (i, p) = selective_find(text, c, index, pos);
index = i;
pos = p;
if index == -1 {
break;
}
delta += index - oldindex;
result.push(delta - 1);
oldindex = index;
delta = 0;
}
oldchar = ch;
}
result
}
fn generate_generalized_integer(mut n: i64, bias: i64) -> Vec<u8> {
let mut result = Vec::new();
let mut j = 0;
loop {
let t = threshold(j, bias);
if n < t {
result.push(DIGITS[n as usize]);
return result;
}
result.push(DIGITS[(t + (n - t) % (BASE - t)) as usize]);
n = (n - t) / (BASE - t);
j += 1;
}
}
fn generate_integers(baselen: i64, deltas: &[i64]) -> Vec<u8> {
let mut result = Vec::new();
let mut bias = INITIAL_BIAS;
for (points, &delta) in deltas.iter().enumerate() {
result.extend(generate_generalized_integer(delta, bias));
bias = adapt(delta, points == 0, baselen + points as i64 + 1);
}
result
}
pub fn encode(text: &str) -> String {
let chars: Vec<char> = text.chars().collect();
let (mut base, extended) = segregate(&chars);
let deltas = insertion_unsort(&chars, &extended);
let encoded = generate_integers(base.len() as i64, &deltas);
if !base.is_empty() {
base.push(b'-');
}
base.extend(encoded);
String::from_utf8(base).expect("punycode output is ASCII")
}
fn decode_generalized_number(
extended: &[u8],
mut extpos: usize,
bias: i64,
) -> Result<(usize, i64), PunycodeError> {
let mut result = 0;
let mut w = 1;
let mut j = 0;
loop {
let ch = *extended.get(extpos).ok_or(PunycodeError::Incomplete)?;
extpos += 1;
let digit = if ch.is_ascii_uppercase() {
(ch - b'A') as i64
} else if ch.is_ascii_digit() {
(ch - b'0') as i64 + 26
} else {
return Err(PunycodeError::InvalidDigit);
};
let t = threshold(j, bias);
result += digit * w;
if digit < t {
return Ok((extpos, result));
}
w *= BASE - t;
j += 1;
}
}
fn insertion_sort(mut base: Vec<char>, extended: &[u8]) -> Result<Vec<char>, PunycodeError> {
let mut char_code: i64 = 0x80;
let mut pos: i64 = -1;
let mut bias = INITIAL_BIAS;
let mut extpos = 0;
while extpos < extended.len() {
let (newpos, delta) = decode_generalized_number(extended, extpos, bias)?;
pos += delta + 1;
char_code += pos / (base.len() as i64 + 1);
if char_code > 0x10FFFF {
return Err(PunycodeError::InvalidCodePoint);
}
pos %= base.len() as i64 + 1;
let ch = char::from_u32(char_code as u32).ok_or(PunycodeError::InvalidCodePoint)?;
base.insert(pos as usize, ch);
bias = adapt(delta, extpos == 0, base.len() as i64);
extpos = newpos;
}
Ok(base)
}
pub fn decode(text: &str) -> Result<String, PunycodeError> {
if !text.is_ascii() {
return Err(PunycodeError::NotAscii);
}
let bytes = text.as_bytes();
let (base, extended) = match bytes.iter().rposition(|&b| b == b'-') {
Some(pos) => (
bytes[..pos].iter().map(|&b| b as char).collect(),
bytes[pos + 1..].to_ascii_uppercase(),
),
None => (Vec::new(), bytes.to_ascii_uppercase()),
};
Ok(insertion_sort(base, &extended)?.into_iter().collect())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn encodes_known_values() {
assert_eq!(encode("münchen"), "mnchen-3ya");
assert_eq!(encode("abc"), "abc-"); assert_eq!(encode(""), "");
assert_eq!(encode("ü"), "tda");
}
#[test]
fn decodes_known_values() {
assert_eq!(decode("mnchen-3ya").unwrap(), "münchen");
assert_eq!(decode("abc-").unwrap(), "abc");
assert_eq!(decode("").unwrap(), "");
assert_eq!(decode("tda").unwrap(), "ü");
}
#[test]
fn round_trips() {
for s in [
"münchen",
"café",
"naïve",
"日本語",
"Ελληνικά",
"abc",
"Hello-World",
"",
"ñ",
"a1b2c3",
"emoji-💡-here",
"mixed café 日本 test",
] {
let encoded = encode(s);
assert_eq!(decode(&encoded).unwrap(), s, "round-trip failed for {s:?}");
}
}
#[test]
fn decode_rejects_bad_input() {
assert_eq!(decode("café"), Err(PunycodeError::NotAscii)); assert!(decode("-!").is_err()); }
}