fix(rust): validate tokenizer ranks and cover backend features

This commit is contained in:
Yujong Lee 2026-09-20 15:13:59 -07:00
parent 56c5d31e73
commit 661da87c91
9 changed files with 326 additions and 204 deletions

View file

@ -120,6 +120,13 @@ jobs:
- run: cargo test --workspace --doc --locked
- name: Test token counter feature combinations
run: |
for features in '' fast huggingface tiktoken fast,huggingface fast,tiktoken huggingface,tiktoken fast,huggingface,tiktoken; do
cargo test -p litellm-token-counter --locked --no-default-features --features "$features"
cargo check -p litellm-python-bridge --locked --no-default-features --features "abi3${features:+,$features}"
done
rust-wheel:
runs-on: ubuntu-latest
timeout-minutes: 30

View file

@ -9,7 +9,7 @@ repository.workspace = true
base64.workspace = true
rustc-hash = "2.1.3"
thiserror.workspace = true
tokenizers = { version = "0.23.1", default-features = false, features = ["onig"] }
tokenizers.workspace = true
unicode-normalization-alignments = "0.1.12"
[dev-dependencies]

View file

@ -81,6 +81,9 @@ fn parse_line(line: &str) -> Result<(Box<[u8]>, Rank), Error> {
let rank = rank
.parse()
.map_err(|error| Error::Ranks(format!("rank is not an integer: {error}")))?;
if rank == NO_RANK {
return Err(Error::Ranks(format!("rank {rank} is reserved")));
}
Ok((bytes.into_boxed_slice(), rank))
}
@ -212,6 +215,21 @@ mod tests {
);
}
#[test]
fn reserved_merge_rank_is_rejected() {
let bytes = (0..=u8::MAX)
.map(|byte| format!("{} {byte}\n", STANDARD.encode([byte])))
.collect::<String>();
let rank_file = format!("{bytes}{} {NO_RANK}\n", STANDARD.encode(b"ab"));
assert!(matches!(
MergeRanks::parse(&rank_file),
Err(Error::Ranks(_))
));
let valid_rank_file = format!("{bytes}{} {}\n", STANDARD.encode(b"ab"), NO_RANK - 1);
let ranks = MergeRanks::parse(&valid_rank_file).unwrap();
assert_eq!(ranks.count_piece(b"aab", &mut MergeScratch::default()), 2);
}
#[test]
fn malformed_rank_files_are_rejected() {
assert!(MergeRanks::parse("IQ==").is_err());

View file

@ -7,4 +7,4 @@ repository.workspace = true
[dependencies]
thiserror.workspace = true
tokenizers = { version = "0.23.1", default-features = false, features = ["onig"] }
tokenizers.workspace = true

View file

@ -0,0 +1,9 @@
use thiserror::Error as ThisError;
#[derive(Debug, ThisError)]
pub enum Error {
#[error("failed to load tokenizer: {0}")]
Load(#[source] tokenizers::Error),
#[error("tokenization failed: {0}")]
Encode(#[source] tokenizers::Error),
}

View file

@ -1,14 +1,8 @@
#![forbid(unsafe_code)]
use thiserror::Error as ThisError;
mod error;
#[derive(Debug, ThisError)]
pub enum Error {
#[error("failed to load tokenizer: {0}")]
Load(#[source] tokenizers::Error),
#[error("tokenization failed: {0}")]
Encode(#[source] tokenizers::Error),
}
pub use error::Error;
pub struct HuggingFaceTokenizer(Box<tokenizers::Tokenizer>);

View file

@ -0,0 +1,5 @@
use thiserror::Error as ThisError;
#[derive(Debug, ThisError)]
#[error("unsupported tokenizer: {0}")]
pub struct UnsupportedTokenizer(pub String);

View file

@ -1,10 +1,8 @@
#![forbid(unsafe_code)]
use thiserror::Error as ThisError;
mod error;
#[derive(Debug, ThisError)]
#[error("unsupported tokenizer: {0}")]
pub struct UnsupportedTokenizer(pub String);
pub use error::UnsupportedTokenizer;
pub struct TiktokenTokenizer(&'static tiktoken_rs::CoreBPE);
@ -32,23 +30,41 @@ mod tests {
use super::*;
#[test]
fn special_tokens_are_counted_as_ordinary_text() {
let counter = TiktokenTokenizer::from_name("cl100k_base").unwrap();
assert!(counter.count_tokens("<|endoftext|>") > 1);
fn named_encodings_match_their_reference_counts() {
let encodings = [
("cl100k_base", tiktoken_rs::cl100k_base_singleton()),
("o200k_base", tiktoken_rs::o200k_base_singleton()),
("o200k_harmony", tiktoken_rs::o200k_harmony_singleton()),
("p50k_base", tiktoken_rs::p50k_base_singleton()),
("p50k_edit", tiktoken_rs::p50k_edit_singleton()),
("r50k_base", tiktoken_rs::r50k_base_singleton()),
("gpt2", tiktoken_rs::r50k_base_singleton()),
];
let texts = [
"",
"Hello, how are you today?",
"é e\u{301} 漢字 ع ३ 🙂 AfiⅣ",
" def function():\n return 123456789\r\n",
"<|endoftext|><|fim_prefix|><|start|>assistant<|message|>",
];
for (name, reference) in encodings {
let counter = TiktokenTokenizer::from_name(name).unwrap();
for text in texts {
assert_eq!(
counter.count_tokens(text),
reference.encode_ordinary(text).len(),
"{name}: {text:?}",
);
}
}
}
#[test]
fn all_python_tiktoken_encodings_are_available() {
for name in [
"cl100k_base",
"o200k_base",
"o200k_harmony",
"p50k_base",
"p50k_edit",
"r50k_base",
"gpt2",
] {
assert!(TiktokenTokenizer::from_name(name).is_ok(), "{name}");
}
fn unsupported_encoding_preserves_its_name() {
let Err(UnsupportedTokenizer(name)) = TiktokenTokenizer::from_name("unknown-encoding")
else {
panic!("unknown encoding must be rejected");
};
assert_eq!(name, "unknown-encoding");
}
}

View file

@ -1,25 +1,30 @@
#![cfg(any(feature = "fast", feature = "huggingface"))]
use rstest::rstest;
use litellm_token_counter::{CountableRequest, Error, InputTokenCount, TokenCounter};
#[cfg(any(feature = "fast", feature = "huggingface", feature = "tiktoken"))]
use litellm_token_counter::TokenCounter;
use litellm_token_counter::{CountableRequest, Error};
/// Expected counts are pinned from `litellm.token_counter(model="claude-sonnet-4-5", ...)`
/// so this test also guards Python parity.
type JsonLoader = fn(&str) -> Result<TokenCounter, Error>;
#[cfg(any(feature = "fast", feature = "huggingface"))]
mod json {
use super::*;
use litellm_token_counter::InputTokenCount;
fn counter(load: JsonLoader) -> TokenCounter {
let path = concat!(
env!("CARGO_MANIFEST_DIR"),
"/../../../litellm/litellm_core_utils/tokenizers/anthropic_tokenizer.json"
);
let json = std::fs::read_to_string(path).expect("anthropic tokenizer json is in the repo");
load(&json).expect("anthropic tokenizer loads")
}
/// Expected counts are pinned from `litellm.token_counter(model="claude-sonnet-4-5", ...)`
/// so this test also guards Python parity.
type JsonLoader = fn(&str) -> Result<TokenCounter, Error>;
const SIMPLE: &str = r#"{"model":"claude-sonnet-4-5","messages":[{"role":"user","content":"Hello, how are you today?"}]}"#;
fn counter(load: JsonLoader) -> TokenCounter {
let path = concat!(
env!("CARGO_MANIFEST_DIR"),
"/../../../litellm/litellm_core_utils/tokenizers/anthropic_tokenizer.json"
);
let json = std::fs::read_to_string(path).expect("anthropic tokenizer json is in the repo");
load(&json).expect("anthropic tokenizer loads")
}
const BLOCKS_AND_SYSTEM: &str = r#"{"model":"claude-sonnet-4-5","messages":[
const SIMPLE: &str = r#"{"model":"claude-sonnet-4-5","messages":[{"role":"user","content":"Hello, how are you today?"}]}"#;
const BLOCKS_AND_SYSTEM: &str = r#"{"model":"claude-sonnet-4-5","messages":[
{"role":"system","content":"You are a terse assistant."},
{"role":"user","name":"alice","content":[
{"type":"text","text":"Summarise this paragraph about ships and harbours."},
@ -28,7 +33,7 @@ const BLOCKS_AND_SYSTEM: &str = r#"{"model":"claude-sonnet-4-5","messages":[
{"type":"tool_reference","tool_name":"get_weather"}]},
{"role":"assistant","content":[{"type":"text","text":"Sure.","cache_control":{"type":"ephemeral"}}]}]}"#;
const TOOLS_OPENAI: &str = r#"{"model":"claude-sonnet-4-5","messages":[{"role":"user","content":"weather?"}],
const TOOLS_OPENAI: &str = r#"{"model":"claude-sonnet-4-5","messages":[{"role":"user","content":"weather?"}],
"tools":[
{"type":"function","function":{"name":"get_weather","description":"Get weather","parameters":{
"type":"object",
@ -43,154 +48,188 @@ const TOOLS_OPENAI: &str = r#"{"model":"claude-sonnet-4-5","messages":[{"role":"
{"type":"function","function":{"name":"noop"}}],
"tool_choice":{"type":"function","function":{"name":"get_weather"}}}"#;
const TOOLS_ANTHROPIC_SYSTEM: &str = r#"{"model":"claude-sonnet-4-5",
const TOOLS_ANTHROPIC_SYSTEM: &str = r#"{"model":"claude-sonnet-4-5",
"messages":[{"role":"system","content":"sys"},{"role":"user","content":"weather?"}],
"tools":[{"name":"get_weather","description":"Get weather","input_schema":{
"type":"object","properties":{"location":{"type":["string","null"]}},"required":["location"]}}],
"tool_choice":"none"}"#;
const COMPLETIONS_PROMPT: &str =
r#"{"model":"claude-sonnet-4-5","prompt":"Write a haiku about ships."}"#;
const COMPLETIONS_PROMPT: &str =
r#"{"model":"claude-sonnet-4-5","prompt":"Write a haiku about ships."}"#;
const COMPLETIONS_PROMPT_LIST: &str =
r#"{"model":"claude-sonnet-4-5","prompt":["first prompt","second prompt"]}"#;
const COMPLETIONS_PROMPT_LIST: &str =
r#"{"model":"claude-sonnet-4-5","prompt":["first prompt","second prompt"]}"#;
const RESPONSES_INPUT: &str = r#"{"model":"claude-sonnet-4-5","input":[
const RESPONSES_INPUT: &str = r#"{"model":"claude-sonnet-4-5","input":[
{"role":"user","content":[{"type":"input_text","text":"Summarise caf\u00e9 menus, na\u00efve \u2014 ok? \"quoted\"\n"}]},
{"role":"assistant","content":"Sure."}],"instructions":"be terse"}"#;
const EMBEDDINGS_TOKEN_IDS: &str =
r#"{"model":"claude-sonnet-4-5","input":[[101,2023,5],[7]],"encoding_format":"float"}"#;
const EMBEDDINGS_TOKEN_IDS: &str =
r#"{"model":"claude-sonnet-4-5","input":[[101,2023,5],[7]],"encoding_format":"float"}"#;
const RERANK: &str = r#"{"model":"claude-sonnet-4-5","query":"best harbour",
const RERANK: &str = r#"{"model":"claude-sonnet-4-5","query":"best harbour",
"documents":["doc one",{"text":"doc two","title":"T","n":3,"ok":true,"none":null,"tags":["a","b"]}]}"#;
fn assert_count_request_matches_python_token_counter(
load: JsonLoader,
body: &str,
expected: usize,
) {
let request = CountableRequest::parse(body.as_bytes()).expect("fixture parses");
let count = counter(load)
.count_request(&request)
.expect("fixture counts");
assert_eq!(
count,
InputTokenCount {
model: Some("claude-sonnet-4-5".to_string()),
input_tokens: expected,
}
);
}
fn assert_count_request_matches_python_token_counter(
load: JsonLoader,
body: &str,
expected: usize,
) {
let request = CountableRequest::parse(body.as_bytes()).expect("fixture parses");
let count = counter(load)
.count_request(&request)
.expect("fixture counts");
assert_eq!(
count,
InputTokenCount {
model: Some("claude-sonnet-4-5".to_string()),
input_tokens: expected,
}
);
}
fn assert_key_presence_follows_python(load: JsonLoader, body: &str, expected: usize) {
let request = CountableRequest::parse(body.as_bytes()).expect("fixture parses");
let count = counter(load)
.count_request(&request)
.expect("fixture counts");
assert_eq!(count.input_tokens, expected);
}
fn assert_key_presence_follows_python(load: JsonLoader, body: &str, expected: usize) {
let request = CountableRequest::parse(body.as_bytes()).expect("fixture parses");
let count = counter(load)
.count_request(&request)
.expect("fixture counts");
assert_eq!(count.input_tokens, expected);
}
fn assert_shapes_outside_the_mirror_are_declined_at_count(load: JsonLoader, body: &[u8]) {
let request = CountableRequest::parse(body).expect("shape parses");
assert!(matches!(
counter(load).count_request(&request),
Err(Error::MissingInput | Error::FloatText | Error::ContentBlock | Error::ArrayItems)
));
}
fn assert_shapes_outside_the_mirror_are_declined_at_count(load: JsonLoader, body: &[u8]) {
let request = CountableRequest::parse(body).expect("shape parses");
assert!(matches!(
counter(load).count_request(&request),
Err(Error::MissingInput | Error::FloatText | Error::ContentBlock | Error::ArrayItems)
));
}
fn assert_tool_choice_and_system_discount_change_the_count(load: JsonLoader) {
let counter = counter(load);
let count = |body: &str| {
counter
.count_request(&CountableRequest::parse(body.as_bytes()).expect("parses"))
.expect("counts")
.input_tokens
};
let base = count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}]}"#);
assert_eq!(
count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tool_choice":"none"}"#),
base + 1
);
assert_eq!(
count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tool_choice":"auto"}"#),
base
);
let with_tools = count(
r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tools":[{"name":"f"}]}"#,
);
let with_tools_and_system = count(
r#"{"model":"m","messages":[{"role":"system","content":"hi"}],"tools":[{"name":"f"}]}"#,
);
assert_eq!(with_tools - with_tools_and_system, 4);
assert_eq!(
count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tools":[]}"#),
base
);
}
fn assert_tool_choice_and_system_discount_change_the_count(load: JsonLoader) {
let counter = counter(load);
let count = |body: &str| {
counter
.count_request(&CountableRequest::parse(body.as_bytes()).expect("parses"))
.expect("counts")
.input_tokens
};
let base = count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}]}"#);
assert_eq!(
count(
r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tool_choice":"none"}"#
),
base + 1
);
assert_eq!(
count(
r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tool_choice":"auto"}"#
),
base
);
let with_tools = count(
r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tools":[{"name":"f"}]}"#,
);
let with_tools_and_system = count(
r#"{"model":"m","messages":[{"role":"system","content":"hi"}],"tools":[{"name":"f"}]}"#,
);
assert_eq!(with_tools - with_tools_and_system, 4);
assert_eq!(
count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tools":[]}"#),
base
);
}
fn assert_loading_a_bad_tokenizer_is_a_load_error(load: JsonLoader) {
assert!(matches!(load("{}"), Err(Error::Load(_))));
}
fn assert_loading_a_bad_tokenizer_is_a_load_error(load: JsonLoader) {
assert!(matches!(load("{}"), Err(Error::Load(_))));
}
macro_rules! json_backend_tests {
($loader:path) => {
#[rstest]
#[case::text_only(super::SIMPLE, 14)]
#[case::content_blocks_name_and_system(super::BLOCKS_AND_SYSTEM, 45)]
#[case::openai_tools_named_choice(super::TOOLS_OPENAI, 123)]
#[case::anthropic_tools_system_discount_choice_none(super::TOOLS_ANTHROPIC_SYSTEM, 53)]
#[case::completions_prompt(super::COMPLETIONS_PROMPT, 7)]
#[case::completions_prompt_list(super::COMPLETIONS_PROMPT_LIST, 4)]
#[case::responses_input_items(super::RESPONSES_INPUT, 62)]
#[case::embeddings_token_ids(super::EMBEDDINGS_TOKEN_IDS, 5)]
#[case::rerank_query_and_documents(super::RERANK, 41)]
fn count_request_matches_python_token_counter(
#[case] body: &str,
#[case] expected: usize,
) {
super::assert_count_request_matches_python_token_counter($loader, body, expected);
}
macro_rules! json_backend_tests {
($loader:path) => {
#[rstest]
#[case::text_only(super::SIMPLE, 14)]
#[case::content_blocks_name_and_system(super::BLOCKS_AND_SYSTEM, 45)]
#[case::openai_tools_named_choice(super::TOOLS_OPENAI, 123)]
#[case::anthropic_tools_system_discount_choice_none(super::TOOLS_ANTHROPIC_SYSTEM, 53)]
#[case::completions_prompt(super::COMPLETIONS_PROMPT, 7)]
#[case::completions_prompt_list(super::COMPLETIONS_PROMPT_LIST, 4)]
#[case::responses_input_items(super::RESPONSES_INPUT, 62)]
#[case::embeddings_token_ids(super::EMBEDDINGS_TOKEN_IDS, 5)]
#[case::rerank_query_and_documents(super::RERANK, 41)]
fn count_request_matches_python_token_counter(
#[case] body: &str,
#[case] expected: usize,
) {
super::assert_count_request_matches_python_token_counter($loader, body, expected);
}
#[rstest]
#[case::null_messages_win_over_prompt(
r#"{"model":"m","messages":null,"prompt":"ignored"}"#,
3
)]
#[case::model_from_route(r#"{"prompt":"hi"}"#, 1)]
#[case::bools_and_ints_use_python_str(r#"{"model":"m","prompt":[true,false,42]}"#, 3)]
#[case::null_prompt_counts_zero(r#"{"model":"m","prompt":null}"#, 0)]
fn key_presence_follows_python(#[case] body: &str, #[case] expected: usize) {
super::assert_key_presence_follows_python($loader, body, expected);
}
#[rstest]
#[case::null_messages_win_over_prompt(
r#"{"model":"m","messages":null,"prompt":"ignored"}"#,
3
)]
#[case::model_from_route(r#"{"prompt":"hi"}"#, 1)]
#[case::bools_and_ints_use_python_str(r#"{"model":"m","prompt":[true,false,42]}"#, 3)]
#[case::null_prompt_counts_zero(r#"{"model":"m","prompt":null}"#, 0)]
fn key_presence_follows_python(#[case] body: &str, #[case] expected: usize) {
super::assert_key_presence_follows_python($loader, body, expected);
}
#[rstest]
#[case::no_countable_input(br#"{"model":"m","instructions":"hi"}"# as &[u8])]
#[case::float_prompt(br#"{"model":"m","prompt":1.5}"#)]
#[case::float_inside_document(br#"{"model":"m","documents":[{"score":0.5}]}"#)]
#[case::image_block(
br#"{"model":"m","messages":[{"role":"user","content":[{"type":"image","source":{"type":"base64","media_type":"image/png","data":"AA=="}}]}]}"#
)]
#[case::tool_result_block(
br#"{"model":"m","messages":[{"role":"user","content":[{"type":"tool_result","tool_use_id":"1","content":"ok"}]}]}"#
)]
#[case::array_without_items(
br#"{"model":"m","messages":[],"tools":[{"name":"f","input_schema":{"type":"object","properties":{"x":{"type":"array"}}}}]}"#
)]
fn shapes_outside_the_mirror_are_declined_at_count(#[case] body: &[u8]) {
super::assert_shapes_outside_the_mirror_are_declined_at_count($loader, body);
}
#[rstest]
#[case::no_countable_input(br#"{"model":"m","instructions":"hi"}"# as &[u8])]
#[case::float_prompt(br#"{"model":"m","prompt":1.5}"#)]
#[case::float_inside_document(br#"{"model":"m","documents":[{"score":0.5}]}"#)]
#[case::image_block(
br#"{"model":"m","messages":[{"role":"user","content":[{"type":"image","source":{"type":"base64","media_type":"image/png","data":"AA=="}}]}]}"#
)]
#[case::tool_result_block(
br#"{"model":"m","messages":[{"role":"user","content":[{"type":"tool_result","tool_use_id":"1","content":"ok"}]}]}"#
)]
#[case::array_without_items(
br#"{"model":"m","messages":[],"tools":[{"name":"f","input_schema":{"type":"object","properties":{"x":{"type":"array"}}}}]}"#
)]
fn shapes_outside_the_mirror_are_declined_at_count(#[case] body: &[u8]) {
super::assert_shapes_outside_the_mirror_are_declined_at_count($loader, body);
}
#[test]
fn tool_choice_and_system_discount_change_the_count() {
super::assert_tool_choice_and_system_discount_change_the_count($loader);
}
#[test]
fn tool_choice_and_system_discount_change_the_count() {
super::assert_tool_choice_and_system_discount_change_the_count($loader);
}
#[test]
fn loading_a_bad_tokenizer_is_a_load_error() {
super::assert_loading_a_bad_tokenizer_is_a_load_error($loader);
}
};
#[test]
fn encoding_errors_preserve_the_backend_source() {
use std::error::Error as _;
let tokenizer = tokenizers::Tokenizer::new(
tokenizers::models::wordpiece::WordPiece::default(),
);
let expected = tokenizer.encode_fast("hello", true).unwrap_err();
let counter = $loader(&tokenizer.to_string(false).unwrap()).unwrap();
let request = CountableRequest::parse(br#"{"prompt":"hello"}"#).unwrap();
let error = counter.count_request(&request).unwrap_err();
assert!(matches!(error, Error::Encode(_)));
assert_eq!(error.source().unwrap().to_string(), expected.to_string());
}
#[test]
fn loading_a_bad_tokenizer_is_a_load_error() {
super::assert_loading_a_bad_tokenizer_is_a_load_error($loader);
}
};
}
#[cfg(feature = "fast")]
mod fast_json {
use super::*;
json_backend_tests!(TokenCounter::from_json_fast);
}
#[cfg(feature = "huggingface")]
mod huggingface_json {
use super::*;
json_backend_tests!(TokenCounter::from_json);
}
}
#[rstest]
@ -215,22 +254,8 @@ fn shapes_outside_the_mirror_are_declined_at_parse(#[case] body: &[u8]) {
));
}
#[cfg(feature = "fast")]
mod fast_json {
use super::*;
json_backend_tests!(TokenCounter::from_json_fast);
}
#[cfg(feature = "huggingface")]
mod huggingface_json {
use super::*;
json_backend_tests!(TokenCounter::from_json);
}
#[cfg(feature = "fast")]
mod fast {
#[cfg(any(feature = "fast", feature = "tiktoken"))]
mod tiktoken {
use super::*;
use serde::Deserialize;
@ -239,33 +264,65 @@ mod fast {
#[derive(Clone, Copy)]
struct TiktokenEncoding {
fixtures: &'static str,
rank_file: &'static str,
source: TokenizerSource,
load: fn(&str) -> Result<TokenCounter, Error>,
model: &'static str,
}
#[cfg(feature = "fast")]
const CL100K: TiktokenEncoding = TiktokenEncoding {
fixtures: "cl100k",
rank_file: "9b5ad71b2ce5302211f9c61530b329a4922fc6a4",
source: TokenizerSource::RankFile("9b5ad71b2ce5302211f9c61530b329a4922fc6a4"),
load: TokenCounter::from_cl100k_ranks,
model: "gpt-4",
};
#[cfg(feature = "fast")]
const O200K: TiktokenEncoding = TiktokenEncoding {
fixtures: "o200k",
rank_file: "fb374d419588a4632f3f557e76b4b70aebbca790",
source: TokenizerSource::RankFile("fb374d419588a4632f3f557e76b4b70aebbca790"),
load: TokenCounter::from_o200k_ranks,
model: "gpt-4o",
};
#[cfg(feature = "tiktoken")]
const TIKTOKEN_CL100K: TiktokenEncoding = TiktokenEncoding {
fixtures: "cl100k",
source: TokenizerSource::Name("cl100k_base"),
load: TokenCounter::from_tiktoken,
model: "gpt-4",
};
#[cfg(feature = "tiktoken")]
const TIKTOKEN_O200K: TiktokenEncoding = TiktokenEncoding {
fixtures: "o200k",
source: TokenizerSource::Name("o200k_base"),
load: TokenCounter::from_tiktoken,
model: "gpt-4o",
};
#[derive(Clone, Copy)]
enum TokenizerSource {
#[cfg(feature = "fast")]
RankFile(&'static str),
#[cfg(feature = "tiktoken")]
Name(&'static str),
}
fn tiktoken_counter(encoding: TiktokenEncoding) -> TokenCounter {
let path = format!(
"{}/../../../litellm/litellm_core_utils/tokenizers/{}",
env!("CARGO_MANIFEST_DIR"),
encoding.rank_file
);
let ranks = std::fs::read_to_string(&path).expect("rank file is in the repo");
(encoding.load)(&ranks).expect("ranks load")
match encoding.source {
#[cfg(feature = "tiktoken")]
TokenizerSource::Name(name) => (encoding.load)(name).expect("encoding loads"),
#[cfg(feature = "fast")]
TokenizerSource::RankFile(file) => {
let path = format!(
"{}/../../../litellm/litellm_core_utils/tokenizers/{file}",
env!("CARGO_MANIFEST_DIR"),
);
let ranks = std::fs::read_to_string(path).expect("rank file is in the repo");
(encoding.load)(&ranks).expect("ranks load")
}
}
}
fn tiktoken_fixture(encoding: TiktokenEncoding, name: &str) -> String {
@ -293,8 +350,10 @@ mod fast {
/// Reference counts come from `tiktoken.get_encoding(name)`; see
/// `token-counter-fast/tests/fixtures/generate.py`.
#[rstest]
#[case::cl100k(CL100K)]
#[case::o200k(O200K)]
#[cfg_attr(feature = "fast", case::fast_cl100k(CL100K))]
#[cfg_attr(feature = "fast", case::fast_o200k(O200K))]
#[cfg_attr(feature = "tiktoken", case::tiktoken_cl100k(TIKTOKEN_CL100K))]
#[cfg_attr(feature = "tiktoken", case::tiktoken_o200k(TIKTOKEN_O200K))]
fn tiktoken_text_counts_match_tiktoken(#[case] encoding: TiktokenEncoding) {
let counter = tiktoken_counter(encoding);
let fixtures: Vec<TextFixture> = tiktoken_fixture(encoding, "texts.jsonl")
@ -319,8 +378,10 @@ mod fast {
/// (`_count_input_tokens(body, model)`), so this pins the shared message,
/// tool and reply-priming accounting on the tiktoken paths as well.
#[rstest]
#[case::cl100k(CL100K)]
#[case::o200k(O200K)]
#[cfg_attr(feature = "fast", case::fast_cl100k(CL100K))]
#[cfg_attr(feature = "fast", case::fast_o200k(O200K))]
#[cfg_attr(feature = "tiktoken", case::tiktoken_cl100k(TIKTOKEN_CL100K))]
#[cfg_attr(feature = "tiktoken", case::tiktoken_o200k(TIKTOKEN_O200K))]
fn tiktoken_request_counts_match_python_admission_counter(#[case] encoding: TiktokenEncoding) {
let counter = tiktoken_counter(encoding);
let fixtures: Vec<RequestFixture> = tiktoken_fixture(encoding, "requests.jsonl")
@ -342,8 +403,10 @@ mod fast {
}
#[rstest]
#[case::cl100k(CL100K)]
#[case::o200k(O200K)]
#[cfg_attr(feature = "fast", case::fast_cl100k(CL100K))]
#[cfg_attr(feature = "fast", case::fast_o200k(O200K))]
#[cfg_attr(feature = "tiktoken", case::tiktoken_cl100k(TIKTOKEN_CL100K))]
#[cfg_attr(feature = "tiktoken", case::tiktoken_o200k(TIKTOKEN_O200K))]
fn tiktoken_shares_the_message_accounting_with_the_anthropic_path(
#[case] encoding: TiktokenEncoding,
) {
@ -369,6 +432,7 @@ mod fast {
);
}
#[cfg(feature = "fast")]
#[rstest]
#[case::empty("")]
#[case::not_base64("!!!! 0")]
@ -382,3 +446,12 @@ mod fast {
assert!(matches!((encoding.load)(rank_file), Err(Error::Ranks(_))));
}
}
#[cfg(feature = "tiktoken")]
#[test]
fn unsupported_encoding_reaches_the_counter_caller() {
assert!(matches!(
TokenCounter::from_tiktoken("unknown-encoding"),
Err(Error::UnsupportedTokenizer(name)) if name == "unknown-encoding"
));
}