use super::state::{DecodeCursor, DecodeState};
use crate::core::tokenize::TokenizeError;
use std::convert::Infallible;
use std::sync::Arc;
pub struct StreamingDecoder {
cursor: DecodeCursor<Arc<DecodeState>>,
}
impl StreamingDecoder {
pub(crate) fn new(state: Arc<DecodeState>) -> Self {
Self {
cursor: DecodeCursor::new(state),
}
}
pub fn add_token(&mut self, id: u32) -> Result<Option<String>, TokenizeError> {
self.add_tokens(&[id])
}
pub fn add_tokens(&mut self, ids: &[u32]) -> Result<Option<String>, TokenizeError> {
self.cursor
.feed(ids, |id| Err(TokenizeError::InvalidTokenId(id)))
}
pub fn add_token_lossy(&mut self, id: u32) -> Option<String> {
self.add_tokens_lossy(&[id])
}
pub fn add_tokens_lossy(&mut self, ids: &[u32]) -> Option<String> {
match self.cursor.feed(ids, |_| Ok::<(), Infallible>(())) {
Ok(text) => text,
Err(never) => match never {},
}
}
pub fn flush(&mut self) -> String {
self.cursor.flush()
}
pub fn reset(&mut self) {
self.cursor.reset();
}
pub fn has_pending(&self) -> bool {
self.cursor.has_pending()
}
pub fn pending_bytes(&self) -> usize {
self.cursor.pending_bytes()
}
}
#[cfg(test)]
mod tests {
use crate::core::any_tokenizer::Backend;
use crate::core::byte_level::byte_level_encode;
use crate::core::pretrained::from_pretrained;
use crate::core::tokenize::TokenizeError;
use crate::core::tokenizer::Tokenizer;
use proptest::prelude::*;
use rustc_hash::FxHashMap;
use std::sync::OnceLock;
fn make_test_tokenizer() -> Tokenizer {
let mut encoder = FxHashMap::default();
for b in 0u8..=255 {
encoder.insert(vec![b], b as u32);
}
encoder.insert("Hello".as_bytes().to_vec(), 256);
encoder.insert("世界".as_bytes().to_vec(), 257);
let special_tokens = FxHashMap::default();
let pattern = r".";
Tokenizer::new(encoder, special_tokens, pattern).unwrap()
}
fn make_byte_level_tokenizer() -> Tokenizer {
let mut encoder = FxHashMap::default();
encoder.insert(byte_level_encode(b"Hello").into_bytes(), 100);
encoder.insert(byte_level_encode(b" world").into_bytes(), 101);
encoder.insert(byte_level_encode("你好".as_bytes()).into_bytes(), 102);
let ni_bytes = "你".as_bytes();
for (i, &b) in ni_bytes.iter().enumerate() {
let byte_level = byte_level_encode(&[b]);
encoder.insert(byte_level.into_bytes(), 200 + i as u32);
}
let mut special_tokens = FxHashMap::default();
special_tokens.insert("<|think|>".to_string(), 1000);
let pattern = r".";
Tokenizer::new_byte_level(encoder, special_tokens, pattern).unwrap()
}
fn make_metaspace_tokenizer() -> Tokenizer {
let mut encoder = FxHashMap::default();
encoder.insert("\u{2581}Hello".as_bytes().to_vec(), 10);
encoder.insert("\u{2581}world".as_bytes().to_vec(), 11);
encoder.insert(vec![0xE2, 0x96], 12);
encoder.insert(vec![0x81, b'x'], 13);
Tokenizer::new_with_metaspace_decoder(encoder, FxHashMap::default(), r".").unwrap()
}
fn make_byte_fallback_tokenizer() -> Tokenizer {
let mut encoder = FxHashMap::default();
encoder.insert(b"a".to_vec(), 1);
encoder.insert(b"c".to_vec(), 2);
for (i, b) in [0xF0u8, 0x90, 0x8D, 0x88].into_iter().enumerate() {
encoder.insert(format!("<0x{b:02X}>").into_bytes(), 10 + i as u32);
}
let byte_fallback = Tokenizer::byte_fallback_from_encoder(&encoder, None, true);
Tokenizer::new(encoder, FxHashMap::default(), r"\S+|\s+")
.expect("the test pattern compiles")
.with_byte_fallback(byte_fallback)
}
fn pretrained_bpe(name: &str) -> Tokenizer {
let any = from_pretrained(name).expect("bundled vocabulary loads");
match any.into_backend() {
Backend::Bpe(tokenizer) => tokenizer,
_ => panic!("{name} is a BPE vocabulary"),
}
}
fn cl100k_base() -> &'static Tokenizer {
static TOKENIZER: OnceLock<Tokenizer> = OnceLock::new();
TOKENIZER.get_or_init(|| pretrained_bpe("cl100k_base"))
}
fn deepseek_v3() -> &'static Tokenizer {
static TOKENIZER: OnceLock<Tokenizer> = OnceLock::new();
TOKENIZER.get_or_init(|| pretrained_bpe("deepseek_v3"))
}
const AGREEMENT_TEXTS: &[&str] = &[
"",
"Hello, world!",
"The quick brown fox jumps over the lazy dog. 1234567890",
"こんにちは世界、これはテストです。",
"Привет, мир! Здравствуйте.",
"🎉🚀 emoji 👨👩👧👦 family, and é combining e\u{0301}.",
"混合 mixed 텍스트 with\ttabs\nand spaces ",
"def f(x):\n return x ** 2 # code",
];
fn drive_strict(tokenizer: &Tokenizer, ids: &[u32], chunk: usize) -> String {
let mut decoder = tokenizer.streaming_decoder();
let mut out = String::new();
for group in ids.chunks(chunk.max(1)) {
if let Some(text) = decoder.add_tokens(group).expect("ids are all known") {
out.push_str(&text);
}
}
out.push_str(&decoder.flush());
out
}
fn drive_lossy(tokenizer: &Tokenizer, ids: &[u32]) -> String {
let mut decoder = tokenizer.streaming_decoder();
let mut out = String::new();
for &id in ids {
if let Some(text) = decoder.add_token_lossy(id) {
out.push_str(&text);
}
}
out.push_str(&decoder.flush());
out
}
#[test]
fn test_simple_ascii() {
let tokenizer = make_test_tokenizer();
let mut decoder = tokenizer.streaming_decoder();
assert_eq!(
decoder.add_token(b'H' as u32).unwrap(),
Some("H".to_string())
);
assert_eq!(
decoder.add_token(b'i' as u32).unwrap(),
Some("i".to_string())
);
assert!(!decoder.has_pending());
}
#[test]
fn test_multi_byte_complete() {
let tokenizer = make_test_tokenizer();
let mut decoder = tokenizer.streaming_decoder();
assert_eq!(decoder.add_token(257).unwrap(), Some("世界".to_string()));
assert!(!decoder.has_pending());
}
#[test]
fn test_multi_byte_split() {
let tokenizer = make_test_tokenizer();
let mut decoder = tokenizer.streaming_decoder();
assert_eq!(decoder.add_token(0xE4).unwrap(), None); assert!(decoder.has_pending());
assert_eq!(decoder.pending_bytes(), 1);
assert_eq!(decoder.add_token(0xB8).unwrap(), None); assert_eq!(decoder.pending_bytes(), 2);
assert_eq!(decoder.add_token(0x96).unwrap(), Some("世".to_string()));
assert!(!decoder.has_pending());
}
#[test]
fn test_flush_incomplete() {
let tokenizer = make_test_tokenizer();
let mut decoder = tokenizer.streaming_decoder();
decoder.add_token(0xE4).unwrap(); decoder.add_token(0xB8).unwrap();
let flushed = decoder.flush();
assert!(flushed.contains('\u{FFFD}')); assert!(!decoder.has_pending());
}
#[test]
fn test_reset() {
let tokenizer = make_test_tokenizer();
let mut decoder = tokenizer.streaming_decoder();
decoder.add_token(0xE4).unwrap();
assert!(decoder.has_pending());
decoder.reset();
assert!(!decoder.has_pending());
}
#[test]
fn test_mixed_complete_incomplete() {
let tokenizer = make_test_tokenizer();
let mut decoder = tokenizer.streaming_decoder();
let result1 = decoder.add_token(b'H' as u32).unwrap();
assert_eq!(result1, Some("H".to_string()));
assert!(!decoder.has_pending());
let result2 = decoder.add_token(0xE4).unwrap(); assert_eq!(result2, None);
assert!(decoder.has_pending());
}
#[test]
fn test_add_tokens_batch() {
let tokenizer = make_test_tokenizer();
let mut decoder = tokenizer.streaming_decoder();
let result = decoder
.add_tokens(&[b'H' as u32, b'i' as u32, b'!' as u32])
.unwrap();
assert_eq!(result, Some("Hi!".to_string()));
}
#[test]
fn test_byte_level_simple_ascii() {
let tokenizer = make_byte_level_tokenizer();
let mut decoder = tokenizer.streaming_decoder();
let result = decoder.add_token(100).unwrap();
assert_eq!(result, Some("Hello".to_string()));
assert!(!decoder.has_pending());
}
#[test]
fn test_byte_level_with_space() {
let tokenizer = make_byte_level_tokenizer();
let mut decoder = tokenizer.streaming_decoder();
let result = decoder.add_token(101).unwrap();
assert_eq!(result, Some(" world".to_string()));
}
#[test]
fn test_byte_level_chinese() {
let tokenizer = make_byte_level_tokenizer();
let mut decoder = tokenizer.streaming_decoder();
let result = decoder.add_token(102).unwrap();
assert_eq!(result, Some("你好".to_string()));
}
#[test]
fn test_byte_level_split_chinese() {
let tokenizer = make_byte_level_tokenizer();
let mut decoder = tokenizer.streaming_decoder();
let result1 = decoder.add_token(200).unwrap();
assert_eq!(result1, None);
assert!(decoder.has_pending());
let result2 = decoder.add_token(201).unwrap();
assert_eq!(result2, None);
assert!(decoder.has_pending());
let result3 = decoder.add_token(202).unwrap();
assert_eq!(result3, Some("你".to_string()));
assert!(!decoder.has_pending());
}
#[test]
fn test_byte_level_special_token() {
let tokenizer = make_byte_level_tokenizer();
let mut decoder = tokenizer.streaming_decoder();
let result = decoder.add_token(1000).unwrap();
assert_eq!(result, Some("<|think|>".to_string()));
}
#[test]
fn test_byte_level_mixed() {
let tokenizer = make_byte_level_tokenizer();
let mut decoder = tokenizer.streaming_decoder();
let result = decoder.add_tokens(&[100, 1000, 101]).unwrap();
assert_eq!(result, Some("Hello<|think|> world".to_string()));
}
#[test]
fn test_byte_level_flush() {
let tokenizer = make_byte_level_tokenizer();
let mut decoder = tokenizer.streaming_decoder();
decoder.add_token(200).unwrap();
decoder.add_token(201).unwrap();
assert!(decoder.has_pending());
let flushed = decoder.flush();
assert!(flushed.contains('\u{FFFD}'));
assert!(!decoder.has_pending());
}
#[test]
fn test_byte_level_reset() {
let tokenizer = make_byte_level_tokenizer();
let mut decoder = tokenizer.streaming_decoder();
decoder.add_token(200).unwrap();
assert!(decoder.has_pending());
decoder.reset();
assert!(!decoder.has_pending());
}
#[test]
fn test_special_decode_ids_are_skipped_like_decode() {
let skipped: rustc_hash::FxHashSet<u32> = [1000u32].into_iter().collect();
let tokenizer = make_byte_level_tokenizer().with_special_decode_ids(skipped);
let ids = [100, 1000, 101];
let mut decoder = tokenizer.streaming_decoder();
let streamed = decoder.add_tokens(&ids).unwrap().unwrap_or_default() + &decoder.flush();
assert_eq!(streamed, "Hello world");
assert_eq!(streamed, tokenizer.decode(&ids).unwrap());
}
#[test]
fn test_metaspace_decoder_applies_to_the_stream() {
let tokenizer = make_metaspace_tokenizer();
for ids in [vec![10, 11], vec![12, 13]] {
let expected = tokenizer.decode(&ids).unwrap();
assert_eq!(drive_strict(&tokenizer, &ids, 1), expected);
}
assert_eq!(drive_strict(&tokenizer, &[10, 11], 1), " Hello world");
assert_eq!(drive_strict(&tokenizer, &[12, 13], 1), " x");
}
#[test]
fn test_byte_fallback_char_reassembles_across_add_token_calls() {
let tokenizer = make_byte_fallback_tokenizer();
let ids = tokenizer.encode("a𐍈c");
assert_eq!(ids, vec![1, 10, 11, 12, 13, 2]);
let mut decoder = tokenizer.streaming_decoder();
assert_eq!(decoder.add_token(1).unwrap(), Some("a".to_string()));
for id in [10, 11, 12] {
assert_eq!(decoder.add_token(id).unwrap(), None);
assert!(decoder.has_pending());
}
assert_eq!(decoder.add_token(13).unwrap(), Some("𐍈".to_string()));
assert!(!decoder.has_pending());
assert_eq!(decoder.add_token(2).unwrap(), Some("c".to_string()));
assert_eq!(decoder.flush(), "");
let expected = tokenizer.decode(&ids).expect("real ids decode");
assert_eq!(expected, "a𐍈c");
for chunk in 1..=ids.len() {
assert_eq!(drive_strict(&tokenizer, &ids, chunk), expected);
}
assert_eq!(drive_lossy(&tokenizer, &ids), tokenizer.decode_lossy(&ids));
}
#[test]
fn test_unknown_id_strict_errors_and_lossy_skips() {
let tokenizer = make_byte_level_tokenizer();
let unknown = 999_999;
let mut strict = tokenizer.streaming_decoder();
assert!(matches!(
strict.add_token(unknown),
Err(TokenizeError::InvalidTokenId(id)) if id == unknown
));
assert!(tokenizer.decode(&[unknown]).is_err());
let mut lossy = tokenizer.streaming_decoder();
let emitted = lossy
.add_tokens_lossy(&[100, unknown, 101])
.unwrap_or_default();
let text = emitted + &lossy.flush();
assert_eq!(text, "Hello world");
assert_eq!(text, tokenizer.decode_lossy(&[100, unknown, 101]));
}
#[test]
fn test_decoder_is_owned_and_outlives_its_tokenizer() {
let mut decoder = {
let tokenizer = make_test_tokenizer();
tokenizer.streaming_decoder()
};
assert_eq!(
decoder.add_token(b'H' as u32).unwrap(),
Some("H".to_string())
);
}
#[test]
fn test_stream_matches_decode_cl100k_base() {
let tokenizer = cl100k_base();
for text in AGREEMENT_TEXTS {
let ids = tokenizer.encode(text);
let expected = tokenizer.decode(&ids).expect("real ids decode");
assert_eq!(drive_strict(tokenizer, &ids, 1), expected, "text: {text:?}");
assert_eq!(drive_lossy(tokenizer, &ids), tokenizer.decode_lossy(&ids));
}
}
#[test]
fn test_stream_matches_decode_deepseek_v3() {
let tokenizer = deepseek_v3();
for text in AGREEMENT_TEXTS {
let ids = tokenizer.encode(text);
let expected = tokenizer.decode(&ids).expect("real ids decode");
assert_eq!(drive_strict(tokenizer, &ids, 1), expected, "text: {text:?}");
assert_eq!(drive_lossy(tokenizer, &ids), tokenizer.decode_lossy(&ids));
}
}
proptest! {
#[test]
fn prop_chunking_matches_decode_cl100k_base(
text in ".{0,120}",
chunk in 1usize..8,
) {
let tokenizer = cl100k_base();
let ids = tokenizer.encode(&text);
let expected = tokenizer.decode(&ids).expect("real ids decode");
prop_assert_eq!(drive_strict(tokenizer, &ids, 1), expected.clone());
prop_assert_eq!(drive_strict(tokenizer, &ids, chunk), expected);
}
#[test]
fn prop_chunking_matches_decode_deepseek_v3(
text in ".{0,120}",
chunk in 1usize..8,
) {
let tokenizer = deepseek_v3();
let ids = tokenizer.encode(&text);
let expected = tokenizer.decode(&ids).expect("real ids decode");
prop_assert_eq!(drive_strict(tokenizer, &ids, 1), expected.clone());
prop_assert_eq!(drive_strict(tokenizer, &ids, chunk), expected);
}
#[test]
fn prop_arbitrary_ids_match_decode_lossy(
ids in prop::collection::vec(0u32..300, 0..48),
) {
let tokenizer = make_test_tokenizer();
prop_assert_eq!(drive_lossy(&tokenizer, &ids), tokenizer.decode_lossy(&ids));
}
#[test]
fn prop_reset_matches_a_fresh_decoder(
dirty in prop::collection::vec(0u32..300, 0..16),
ids in prop::collection::vec(0u32..300, 0..32),
) {
let tokenizer = make_test_tokenizer();
let mut reused = tokenizer.streaming_decoder();
reused.add_tokens_lossy(&dirty);
reused.reset();
prop_assert!(!reused.has_pending());
prop_assert_eq!(reused.pending_bytes(), 0);
let mut fresh = tokenizer.streaming_decoder();
let mut from_reused = String::new();
let mut from_fresh = String::new();
for &id in &ids {
let a = reused.add_token_lossy(id);
let b = fresh.add_token_lossy(id);
prop_assert_eq!(&a, &b);
prop_assert_eq!(reused.pending_bytes(), fresh.pending_bytes());
from_reused.push_str(&a.unwrap_or_default());
from_fresh.push_str(&b.unwrap_or_default());
}
from_reused.push_str(&reused.flush());
from_fresh.push_str(&fresh.flush());
prop_assert_eq!(from_reused, from_fresh);
}
}
}