mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-24 00:52:24 +00:00
fix(rust): validate tokenizer ranks and cover backend features
This commit is contained in:
parent
56c5d31e73
commit
661da87c91
9 changed files with 326 additions and 204 deletions
7
.github/workflows/test-rust.yml
vendored
7
.github/workflows/test-rust.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
|
|
|
|||
|
|
@ -7,4 +7,4 @@ repository.workspace = true
|
|||
|
||||
[dependencies]
|
||||
thiserror.workspace = true
|
||||
tokenizers = { version = "0.23.1", default-features = false, features = ["onig"] }
|
||||
tokenizers.workspace = true
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
}
|
||||
|
|
@ -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>);
|
||||
|
||||
|
|
|
|||
5
litellm-rust/crates/token-counter-tiktoken/src/error.rs
Normal file
5
litellm-rust/crates/token-counter-tiktoken/src/error.rs
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
use thiserror::Error as ThisError;
|
||||
|
||||
#[derive(Debug, ThisError)]
|
||||
#[error("unsupported tokenizer: {0}")]
|
||||
pub struct UnsupportedTokenizer(pub String);
|
||||
|
|
@ -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");
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
));
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue