Skip to content
104 changes: 102 additions & 2 deletions lib/llm/src/backend.rs
Original file line number Diff line number Diff line change
Expand Up @@ -182,7 +182,20 @@ impl

let data = output.data.as_ref().unwrap();

let result = state.decoder.process_token_ids(&data.token_ids).unwrap();
let result = match state.decoder.process_token_ids(&data.token_ids) {
Ok(result) => result,
Err(e) => {
tracing::error!("Failed to process token_ids: {e}");
state.stream.context().stop_generating();
state.finished = true;
let mut output = output;
if let Some(data) = &mut output.data {
data.finish_reason =
Some(FinishReason::Error(format!("decode error: {e}")));
Comment thread
biswapanda marked this conversation as resolved.
}
return Some((output, state));
}
};

// NOTE: the `finish_reason` is computed from the generated `token_ids` alone.
// The `data` field can have a `finish_reason` set, coming from the underlying
Expand Down Expand Up @@ -583,16 +596,103 @@ impl Decoder {
}
}

#[cfg(test)]
mod tests {
use super::*;
use crate::tokenizers::traits;
use std::sync::Arc;

#[test]
fn test_char_boundary_drain() {
use super::Decoder;
let mut s = String::from("helloñworld"); // 12 bytes total ñ is 2 bytes
let max_bytes = 6; // 12 - 6 = 6 which is inside ñ
assert!(!s.is_char_boundary(s.len() - max_bytes)); // initially we are not on a char boundary
Decoder::maybe_drain_to_max_bytes(&mut s, max_bytes);
assert!(s.is_char_boundary(0)); // front of jail string on valid char boundary
assert_eq!(s, "ñworld");
}

/// A mock tokenizer that always returns Err from decode().
/// Used to test the error propagation path in Decoder::process_token_ids().
struct FailingDecoder;

impl traits::Encoder for FailingDecoder {
fn encode(&self, _input: &str) -> anyhow::Result<crate::tokenizers::Encoding> {
Ok(crate::tokenizers::Encoding::Sp(vec![]))
}
fn encode_batch(
&self,
_inputs: &[&str],
) -> anyhow::Result<Vec<crate::tokenizers::Encoding>> {
Ok(vec![])
}
}

impl traits::Decoder for FailingDecoder {
fn decode(
&self,
_token_ids: &[TokenIdType],
_skip_special_tokens: bool,
) -> anyhow::Result<String> {
Err(anyhow::anyhow!(
"Unable to decode into a valid UTF-8 string: incomplete utf-8 byte sequence from index 6"
))
}
}

impl traits::Tokenizer for FailingDecoder {}

/// When the tokenizer's decode() returns Err, Decoder::process_token_ids()
/// should propagate the error. In the backend unfold closure, this error
/// gets caught and converted to FinishReason::Error.
#[test]
fn test_decoder_process_token_ids_propagates_decode_error() {
let tokenizer: Arc<dyn traits::Tokenizer> = Arc::new(FailingDecoder);
let decode_stream = crate::tokenizers::DecodeStream::new(tokenizer, &[], false);
let stop_conditions = StopConditions::default();

let mut decoder = Decoder::new(decode_stream, stop_conditions, false, None);

let result = decoder.process_token_ids(&[42]);
assert!(
result.is_err(),
"process_token_ids should propagate decode errors"
);

let err_msg = result.err().unwrap().to_string();
assert!(
err_msg.contains("incomplete utf-8 byte sequence"),
"error should contain the original decode error message, got: {err_msg}"
);
}

/// Verify that the error message format matches what the backend unfold
/// closure would wrap into FinishReason::Error.
#[test]
fn test_decoder_error_message_format_for_finish_reason() {
let tokenizer: Arc<dyn traits::Tokenizer> = Arc::new(FailingDecoder);
let decode_stream = crate::tokenizers::DecodeStream::new(tokenizer, &[], false);
let stop_conditions = StopConditions::default();

let mut decoder = Decoder::new(decode_stream, stop_conditions, false, None);

let result = decoder.process_token_ids(&[42]);
let err = result.err().expect("should be Err");

// This is what the backend unfold closure does:
let finish_reason = FinishReason::Error(format!("decode error: {err}"));
match &finish_reason {
FinishReason::Error(msg) => {
assert!(
msg.starts_with("decode error:"),
"FinishReason::Error should have 'decode error:' prefix, got: {msg}"
);
assert!(
msg.contains("incomplete utf-8 byte sequence"),
"FinishReason::Error should contain original error, got: {msg}"
);
}
other => panic!("Expected FinishReason::Error, got: {:?}", other),
}
}
}
4 changes: 4 additions & 0 deletions lib/llm/src/tokenizers.rs
Original file line number Diff line number Diff line change
Expand Up @@ -60,6 +60,10 @@ pub mod traits {
fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>>;
}

/// Implementations **must** use lossy UTF-8 conversion (e.g. `String::from_utf8_lossy`)
/// so that partial multi-byte sequences produce U+FFFD (`�`) rather than returning `Err`.
/// `DecodeStream::step()` relies on the replacement character to detect incomplete
/// sequences and buffer tokens until the full character arrives.
pub trait Decoder: Send + Sync {
fn decode(&self, token_ids: &[TokenIdType], skip_special_tokens: bool) -> Result<String>;
}
Expand Down
204 changes: 201 additions & 3 deletions lib/llm/src/tokenizers/tiktoken.rs
Comment thread
biswapanda marked this conversation as resolved.
Original file line number Diff line number Diff line change
Expand Up @@ -100,9 +100,12 @@ impl Decoder for TikTokenTokenizer {
token_ids.to_vec()
};

self.bpe
.decode(ids)
.map_err(|err| Error::msg(format!("Error decoding tiktoken tokens: {err}")))
// Use lossy UTF-8 conversion so that partial multi-byte sequences become U+FFFD (�).
// This is critical for incremental detokenization: DecodeStream::step() relies on
// the replacement character to detect incomplete sequences and buffer tokens until
// a complete character arrives. CoreBPE::decode() would error on invalid UTF-8 instead.
let bytes: Vec<u8> = self.bpe._decode_native_and_split(ids).flatten().collect();
Ok(String::from_utf8_lossy(&bytes).into_owned())
}
}

Expand Down Expand Up @@ -236,7 +239,9 @@ fn load_special_tokens(directory: &Path, num_base_tokens: usize) -> Result<FxHas
#[cfg(test)]
mod tests {
use super::*;
use crate::tokenizers::DecodeStream;
use std::io::Write;
use std::sync::Arc;

fn create_test_tiktoken_file(dir: &Path) -> String {
let engine = base64::engine::general_purpose::STANDARD;
Expand Down Expand Up @@ -443,6 +448,199 @@ mod tests {
assert!(tokens.len() > 2);
}

/// Helper: create a tiktoken file that includes raw byte tokens (byte fallback tokens).
fn create_test_tiktoken_file_with_byte_tokens(dir: &Path) -> String {
Comment thread
biswapanda marked this conversation as resolved.
let engine = base64::engine::general_purpose::STANDARD;
let mut content = String::new();

let tokens: Vec<(&[u8], u32)> = vec![
(b"h", 0),
(b"e", 1),
(b"l", 2),
(b"o", 3),
(b" ", 4),
(b"hello", 5),
];

for (token, rank) in &tokens {
let encoded = engine.encode(token);
content.push_str(&format!("{encoded} {rank}\n"));
}

// Byte-fallback tokens: individual bytes that form CJK character "你" (U+4F60)
// UTF-8 encoding: 0xE4 0xBD 0xA0
let byte_tokens: Vec<(Vec<u8>, u32)> =
vec![(vec![0xE4], 100), (vec![0xBD], 101), (vec![0xA0], 102)];

for (token, rank) in &byte_tokens {
let encoded = engine.encode(token);
content.push_str(&format!("{encoded} {rank}\n"));
}

// Bytes for emoji "😀" (U+1F600) — 4-byte UTF-8: 0xF0 0x9F 0x98 0x80
let emoji_tokens: Vec<(Vec<u8>, u32)> = vec![
(vec![0xF0], 200),
(vec![0x9F], 201),
(vec![0x98], 202),
(vec![0x80], 203),
];

for (token, rank) in &emoji_tokens {
let encoded = engine.encode(token);
content.push_str(&format!("{encoded} {rank}\n"));
}

let file_path = dir.join("tiktoken.model");
let mut file = std::fs::File::create(&file_path).unwrap();
file.write_all(content.as_bytes()).unwrap();
file_path.to_str().unwrap().to_string()
}

fn create_byte_token_tokenizer(dir: &Path) -> TikTokenTokenizer {
let file_path = create_test_tiktoken_file_with_byte_tokens(dir);
let special_tokens = FxHashMap::default();
let pattern = r"[\w]+|[^\w\s]+|\s+";
TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap()
}

/// Reproduces the original panic: decoding a single byte-fallback token that is
/// part of a multi-byte UTF-8 character. Before the fix, CoreBPE::decode() would
/// call String::from_utf8() on [0xE4] and error with "incomplete utf-8 byte sequence".
#[test]
fn test_decode_single_incomplete_utf8_byte_does_not_error() {
let dir = tempfile::tempdir().unwrap();
let tokenizer = create_byte_token_tokenizer(dir.path());

let result = tokenizer.decode(&[100], false);
assert!(
result.is_ok(),
"decode() should not error on incomplete UTF-8 bytes"
);
let text = result.unwrap();
assert!(
text.contains('\u{FFFD}'),
"incomplete UTF-8 byte should produce replacement character, got: {:?}",
text
);
}

/// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
#[test]
fn test_decode_two_of_three_utf8_bytes_does_not_error() {
let dir = tempfile::tempdir().unwrap();
let tokenizer = create_byte_token_tokenizer(dir.path());

let result = tokenizer.decode(&[100, 101], false);
assert!(result.is_ok());
let text = result.unwrap();
assert!(
text.contains('\u{FFFD}'),
"incomplete 2-of-3 UTF-8 bytes should produce replacement character, got: {:?}",
text
);
}

/// When all bytes of a multi-byte character are present, the concatenated bytes form
/// valid UTF-8, so this test passes both before and after the fix. It serves as a
/// correctness check that the lossy conversion doesn't corrupt complete characters.
#[test]
fn test_decode_complete_multibyte_utf8_produces_correct_char() {
let dir = tempfile::tempdir().unwrap();
let tokenizer = create_byte_token_tokenizer(dir.path());

let result = tokenizer.decode(&[100, 101, 102], false);
assert!(result.is_ok());
assert_eq!(result.unwrap(), "你");
}

/// All 4 emoji bytes together form valid UTF-8, so this passes both before and after
/// the fix. Validates that lossy conversion doesn't alter complete multi-byte sequences.
#[test]
fn test_decode_complete_4byte_emoji_from_byte_tokens() {
let dir = tempfile::tempdir().unwrap();
let tokenizer = create_byte_token_tokenizer(dir.path());

let result = tokenizer.decode(&[200, 201, 202, 203], false);
assert!(result.is_ok());
assert_eq!(result.unwrap(), "😀");
}

/// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
#[test]
fn test_decode_partial_emoji_does_not_error() {
let dir = tempfile::tempdir().unwrap();
let tokenizer = create_byte_token_tokenizer(dir.path());

let result = tokenizer.decode(&[200], false);
assert!(result.is_ok());
assert!(result.unwrap().contains('\u{FFFD}'));
}

/// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
#[test]
fn test_decode_mixed_ascii_and_incomplete_bytes() {
let dir = tempfile::tempdir().unwrap();
let tokenizer = create_byte_token_tokenizer(dir.path());

let result = tokenizer.decode(&[5, 100], false);
assert!(result.is_ok());
let text = result.unwrap();
assert!(
text.starts_with("hello"),
"should start with 'hello', got: {:?}",
text
);
assert!(
text.contains('\u{FFFD}'),
"trailing incomplete byte should produce U+FFFD"
);
}

/// End-to-end incremental detokenization: DecodeStream buffers partial bytes,
/// emits the complete character once all bytes arrive.
/// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
#[test]
fn test_decode_stream_incremental_multibyte_reassembly() {
let dir = tempfile::tempdir().unwrap();
let tokenizer = create_byte_token_tokenizer(dir.path());
let tokenizer_arc: Arc<dyn crate::tokenizers::traits::Tokenizer> = Arc::new(tokenizer);

let mut stream = DecodeStream::new(tokenizer_arc, &[5], false);

let r1 = stream.step(100).unwrap();
assert_eq!(r1, None, "first byte of 3-byte char should be buffered");

let r2 = stream.step(101).unwrap();
assert_eq!(r2, None, "second byte of 3-byte char should be buffered");

let r3 = stream.step(102).unwrap();
assert!(r3.is_some(), "third byte should complete the character");
assert_eq!(r3.unwrap(), "你");
}

/// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
#[test]
fn test_decode_stream_incremental_emoji_reassembly() {
let dir = tempfile::tempdir().unwrap();
let tokenizer = create_byte_token_tokenizer(dir.path());
let tokenizer_arc: Arc<dyn crate::tokenizers::traits::Tokenizer> = Arc::new(tokenizer);

let mut stream = DecodeStream::new(tokenizer_arc, &[5], false);

let r1 = stream.step(200).unwrap();
assert_eq!(r1, None, "byte 1/4 of emoji should be buffered");

let r2 = stream.step(201).unwrap();
assert_eq!(r2, None, "byte 2/4 of emoji should be buffered");

let r3 = stream.step(202).unwrap();
assert_eq!(r3, None, "byte 3/4 of emoji should be buffered");

let r4 = stream.step(203).unwrap();
assert!(r4.is_some(), "byte 4/4 should complete the emoji");
assert_eq!(r4.unwrap(), "😀");
}

#[test]
fn test_tiktoken_encode_batch() {
let dir = tempfile::tempdir().unwrap();
Expand Down
Loading
Loading