diff --git a/.github/workflows/test-rust.yml b/.github/workflows/test-rust.yml index 551f783d4f9..70d828b5c8b 100644 --- a/.github/workflows/test-rust.yml +++ b/.github/workflows/test-rust.yml @@ -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 diff --git a/litellm-rust/crates/token-counter-fast/Cargo.toml b/litellm-rust/crates/token-counter-fast/Cargo.toml index 4950277f5b0..5127dc17f22 100644 --- a/litellm-rust/crates/token-counter-fast/Cargo.toml +++ b/litellm-rust/crates/token-counter-fast/Cargo.toml @@ -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] diff --git a/litellm-rust/crates/token-counter-fast/src/tiktoken.rs b/litellm-rust/crates/token-counter-fast/src/tiktoken.rs index 7a9e71ed587..16172b7a688 100644 --- a/litellm-rust/crates/token-counter-fast/src/tiktoken.rs +++ b/litellm-rust/crates/token-counter-fast/src/tiktoken.rs @@ -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::(); + 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()); diff --git a/litellm-rust/crates/token-counter-huggingface/Cargo.toml b/litellm-rust/crates/token-counter-huggingface/Cargo.toml index ad15d9f19d0..6d8cb85e524 100644 --- a/litellm-rust/crates/token-counter-huggingface/Cargo.toml +++ b/litellm-rust/crates/token-counter-huggingface/Cargo.toml @@ -7,4 +7,4 @@ repository.workspace = true [dependencies] thiserror.workspace = true -tokenizers = { version = "0.23.1", default-features = false, features = ["onig"] } +tokenizers.workspace = true diff --git a/litellm-rust/crates/token-counter-huggingface/src/error.rs b/litellm-rust/crates/token-counter-huggingface/src/error.rs new file mode 100644 index 00000000000..adc4551886f --- /dev/null +++ b/litellm-rust/crates/token-counter-huggingface/src/error.rs @@ -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), +} diff --git a/litellm-rust/crates/token-counter-huggingface/src/lib.rs b/litellm-rust/crates/token-counter-huggingface/src/lib.rs index 184956b5f39..8e05c2cca46 100644 --- a/litellm-rust/crates/token-counter-huggingface/src/lib.rs +++ b/litellm-rust/crates/token-counter-huggingface/src/lib.rs @@ -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); diff --git a/litellm-rust/crates/token-counter-tiktoken/src/error.rs b/litellm-rust/crates/token-counter-tiktoken/src/error.rs new file mode 100644 index 00000000000..e28cbbfb620 --- /dev/null +++ b/litellm-rust/crates/token-counter-tiktoken/src/error.rs @@ -0,0 +1,5 @@ +use thiserror::Error as ThisError; + +#[derive(Debug, ThisError)] +#[error("unsupported tokenizer: {0}")] +pub struct UnsupportedTokenizer(pub String); diff --git a/litellm-rust/crates/token-counter-tiktoken/src/lib.rs b/litellm-rust/crates/token-counter-tiktoken/src/lib.rs index b35c4876d00..ecdb3946eee 100644 --- a/litellm-rust/crates/token-counter-tiktoken/src/lib.rs +++ b/litellm-rust/crates/token-counter-tiktoken/src/lib.rs @@ -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"); } } diff --git a/litellm-rust/crates/token-counter/tests/token_counter.rs b/litellm-rust/crates/token-counter/tests/token_counter.rs index ebbe2f00ffc..542bd4a1fc4 100644 --- a/litellm-rust/crates/token-counter/tests/token_counter.rs +++ b/litellm-rust/crates/token-counter/tests/token_counter.rs @@ -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; +#[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; -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, 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 = 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 = 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" + )); +}