mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
Merge remote-tracking branch 'origin/litellm_rust_admission_poc_messages' into litellm_perf_1k_rps
This commit is contained in:
commit
34d641239e
19 changed files with 1825 additions and 37 deletions
407
litellm-rust/Cargo.lock
generated
407
litellm-rust/Cargo.lock
generated
|
|
@ -2,6 +2,20 @@
|
|||
# It is not intended for manual editing.
|
||||
version = 4
|
||||
|
||||
[[package]]
|
||||
name = "ahash"
|
||||
version = "0.8.12"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "5a15f179cd60c4584b8a8c596927aadc462e27f2ca70c04e0071964a73ba7a75"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"getrandom 0.3.4",
|
||||
"once_cell",
|
||||
"serde",
|
||||
"version_check",
|
||||
"zerocopy",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "aho-corasick"
|
||||
version = "1.1.5"
|
||||
|
|
@ -418,7 +432,7 @@ checksum = "edca88bc138befd0323b20752846e6587272d3b03b0343c8ea28a6f819e6e71f"
|
|||
dependencies = [
|
||||
"async-trait",
|
||||
"axum-core",
|
||||
"base64",
|
||||
"base64 0.22.1",
|
||||
"bytes",
|
||||
"futures-util",
|
||||
"http 1.4.2",
|
||||
|
|
@ -468,6 +482,12 @@ dependencies = [
|
|||
"tracing",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "base64"
|
||||
version = "0.13.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "9e1b586273c5702936fe7b7d6896644d8be71e6314cfe09d3167c95f712589e8"
|
||||
|
||||
[[package]]
|
||||
name = "base64"
|
||||
version = "0.22.1"
|
||||
|
|
@ -542,6 +562,15 @@ version = "0.3.0"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "37b2a672a2cb129a2e41c10b1224bb368f9f37a2b16b612598138befd7b37eb5"
|
||||
|
||||
[[package]]
|
||||
name = "castaway"
|
||||
version = "0.2.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "dec551ab6e7578819132c713a93c022a05d60159dc86e7a7050223577484c55a"
|
||||
dependencies = [
|
||||
"rustversion",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "cc"
|
||||
version = "1.3.0"
|
||||
|
|
@ -644,6 +673,21 @@ version = "0.5.4"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "0c9ea0ac24bc397ab3c98583a3c9ba74fa56b09a4449bbe172b9b1ddb016027a"
|
||||
|
||||
[[package]]
|
||||
name = "compact_str"
|
||||
version = "0.9.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "9dfdd1c2274d9aa354115b09dc9a901d6c5576818cdf70d14cae2bdb47df00ab"
|
||||
dependencies = [
|
||||
"castaway",
|
||||
"cfg-if",
|
||||
"itoa",
|
||||
"rustversion",
|
||||
"ryu",
|
||||
"serde",
|
||||
"static_assertions",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "const-oid"
|
||||
version = "0.10.2"
|
||||
|
|
@ -696,7 +740,7 @@ dependencies = [
|
|||
"ciborium",
|
||||
"clap",
|
||||
"criterion-plot",
|
||||
"itertools",
|
||||
"itertools 0.13.0",
|
||||
"num-traits",
|
||||
"oorandom",
|
||||
"page_size",
|
||||
|
|
@ -716,7 +760,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
|||
checksum = "d8d80a2f4f5b554395e47b5d8305bc3d27813bacb73493eb1001e8f76dae29ea"
|
||||
dependencies = [
|
||||
"cast",
|
||||
"itertools",
|
||||
"itertools 0.13.0",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -778,6 +822,56 @@ dependencies = [
|
|||
"cmov",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "daachorse"
|
||||
version = "3.0.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "5614204febbc33cc07a2806aa6440b904ac012b68eecc37f4493ea4a76455a3d"
|
||||
|
||||
[[package]]
|
||||
name = "darling"
|
||||
version = "0.20.11"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "fc7f46116c46ff9ab3eb1597a45688b6715c6e628b5c133e288e709a29bcb4ee"
|
||||
dependencies = [
|
||||
"darling_core",
|
||||
"darling_macro",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "darling_core"
|
||||
version = "0.20.11"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "0d00b9596d185e565c2207a0b01f8bd1a135483d02d9b7b0a54b11da8d53412e"
|
||||
dependencies = [
|
||||
"fnv",
|
||||
"ident_case",
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"strsim",
|
||||
"syn 2.0.119",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "darling_macro"
|
||||
version = "0.20.11"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "fc34b93ccb385b40dc71c6fceac4b2ad23662c7eeb248cf10d529b7e055b6ead"
|
||||
dependencies = [
|
||||
"darling_core",
|
||||
"quote",
|
||||
"syn 2.0.119",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "dary_heap"
|
||||
version = "0.3.9"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "8b1e3a325bc115f096c8b77bbf027a7c2592230e70be2d985be950d3d5e60ebe"
|
||||
dependencies = [
|
||||
"serde",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "data-encoding"
|
||||
version = "2.11.0"
|
||||
|
|
@ -790,6 +884,37 @@ version = "0.5.8"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7cd812cc2bc1d69d4764bd80df88b4317eaef9e773c75226407d9bc0876b211c"
|
||||
|
||||
[[package]]
|
||||
name = "derive_builder"
|
||||
version = "0.20.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "507dfb09ea8b7fa618fcf76e953f4f5e192547945816d5358edffe39f6f94947"
|
||||
dependencies = [
|
||||
"derive_builder_macro",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "derive_builder_core"
|
||||
version = "0.20.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "2d5bcf7b024d6835cfb3d473887cd966994907effbe9227e8c8219824d06c4e8"
|
||||
dependencies = [
|
||||
"darling",
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 2.0.119",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "derive_builder_macro"
|
||||
version = "0.20.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "ab63b0e2bf4d5928aff72e83a7dace85d7bba5fe12dcc3c5a572d78caffd3f3c"
|
||||
dependencies = [
|
||||
"derive_builder_core",
|
||||
"syn 2.0.119",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "digest"
|
||||
version = "0.10.7"
|
||||
|
|
@ -841,6 +966,12 @@ version = "1.0.2"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f"
|
||||
|
||||
[[package]]
|
||||
name = "esaxx-rs"
|
||||
version = "0.1.10"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d817e038c30374a4bcb22f94d0a8a0e216958d4c3dcde369b1439fec4bdda6e6"
|
||||
|
||||
[[package]]
|
||||
name = "fastrand"
|
||||
version = "2.5.0"
|
||||
|
|
@ -964,6 +1095,18 @@ dependencies = [
|
|||
"wasm-bindgen",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "getrandom"
|
||||
version = "0.3.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "899def5c37c4fd7b2664648c28120ecec138e4d395b459e5ca34f9cce2dd77fd"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"libc",
|
||||
"r-efi 5.3.0",
|
||||
"wasip2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "getrandom"
|
||||
version = "0.4.3"
|
||||
|
|
@ -973,7 +1116,7 @@ dependencies = [
|
|||
"cfg-if",
|
||||
"js-sys",
|
||||
"libc",
|
||||
"r-efi",
|
||||
"r-efi 6.0.0",
|
||||
"rand_core 0.10.1",
|
||||
"wasm-bindgen",
|
||||
]
|
||||
|
|
@ -1220,7 +1363,7 @@ version = "0.1.20"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "96547c2556ec9d12fb1578c4eaf448b04993e7fb79cbaad930a656880a6bdfa0"
|
||||
dependencies = [
|
||||
"base64",
|
||||
"base64 0.22.1",
|
||||
"bytes",
|
||||
"futures-channel",
|
||||
"futures-util",
|
||||
|
|
@ -1319,6 +1462,12 @@ dependencies = [
|
|||
"zerovec",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ident_case"
|
||||
version = "1.0.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b9e0384b61958566e926dc50660321d12159025e767c18e043daf26b70104c39"
|
||||
|
||||
[[package]]
|
||||
name = "idna"
|
||||
version = "1.1.0"
|
||||
|
|
@ -1348,6 +1497,8 @@ checksum = "d466e9454f08e4a911e14806c24e16fba1b4c121d1ea474396f396069cf949d9"
|
|||
dependencies = [
|
||||
"equivalent",
|
||||
"hashbrown",
|
||||
"serde",
|
||||
"serde_core",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -1365,6 +1516,15 @@ dependencies = [
|
|||
"either",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "itertools"
|
||||
version = "0.14.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "2b192c782037fadd9cfa75548310488aabdbf3d2da73885b31bd0abd03351285"
|
||||
dependencies = [
|
||||
"either",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "itoa"
|
||||
version = "1.0.18"
|
||||
|
|
@ -1409,7 +1569,7 @@ name = "litellm-ai-gateway"
|
|||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"axum",
|
||||
"base64",
|
||||
"base64 0.22.1",
|
||||
"futures-channel",
|
||||
"futures-util",
|
||||
"litellm-config",
|
||||
|
|
@ -1447,7 +1607,8 @@ dependencies = [
|
|||
"aws-sigv4",
|
||||
"aws-smithy-runtime-api",
|
||||
"aws-types",
|
||||
"base64",
|
||||
"base64 0.22.1",
|
||||
"indexmap",
|
||||
"rand 0.8.7",
|
||||
"reqwest",
|
||||
"rstest",
|
||||
|
|
@ -1455,6 +1616,7 @@ dependencies = [
|
|||
"serde_json",
|
||||
"sha2 0.10.9",
|
||||
"thiserror 2.0.19",
|
||||
"tokenizers",
|
||||
"tokio",
|
||||
"tracing",
|
||||
"tracing-subscriber",
|
||||
|
|
@ -1507,6 +1669,22 @@ version = "0.1.2"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "112b39cec0b298b6c1999fee3e31427f74f676e4cb9879ed1a121b43661a4154"
|
||||
|
||||
[[package]]
|
||||
name = "macro_rules_attribute"
|
||||
version = "0.2.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b3ae8f6d608c795738406608304d30a2dfbdc8e58e44f7ba43236da5208ded3c"
|
||||
dependencies = [
|
||||
"macro_rules_attribute-proc_macro",
|
||||
"pastey",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "macro_rules_attribute-proc_macro"
|
||||
version = "0.2.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "fc04a4c58212d57930a24bf47d3fa87485264a3a054e9c10e042eb373573ad3c"
|
||||
|
||||
[[package]]
|
||||
name = "matchit"
|
||||
version = "0.7.3"
|
||||
|
|
@ -1535,6 +1713,12 @@ dependencies = [
|
|||
"unicase",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "minimal-lexical"
|
||||
version = "0.2.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "68354c5c6bd36d73ff3feceb05efa59b6acb7626617f4962be322a825e61f79a"
|
||||
|
||||
[[package]]
|
||||
name = "mio"
|
||||
version = "1.2.2"
|
||||
|
|
@ -1546,6 +1730,38 @@ dependencies = [
|
|||
"windows-sys 0.61.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "monostate"
|
||||
version = "0.1.18"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "3341a273f6c9d5bef1908f17b7267bbab0e95c9bf69a0d4dcf8e9e1b2c76ef67"
|
||||
dependencies = [
|
||||
"monostate-impl",
|
||||
"serde",
|
||||
"serde_core",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "monostate-impl"
|
||||
version = "0.1.18"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e4db6d5580af57bf992f59068d4ea26fd518574ff48d7639b255a36f9de6e7e9"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 2.0.119",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "nom"
|
||||
version = "7.1.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d273983c5a657a70a3e8f2a01329822f3b8c8172b73826411a55751e404a0a4a"
|
||||
dependencies = [
|
||||
"memchr",
|
||||
"minimal-lexical",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "num-conv"
|
||||
version = "0.2.2"
|
||||
|
|
@ -1576,6 +1792,28 @@ version = "1.21.4"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50"
|
||||
|
||||
[[package]]
|
||||
name = "onig"
|
||||
version = "6.5.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "0cc3cbf698f9438986c11a880c90a6d04b9de27575afd28bbf45b154b6c709e2"
|
||||
dependencies = [
|
||||
"bitflags",
|
||||
"libc",
|
||||
"once_cell",
|
||||
"onig_sys",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "onig_sys"
|
||||
version = "69.9.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "1e68317604e77e53b85896388e1a803c1d21b74c899ec9e5e1112db90735edd7"
|
||||
dependencies = [
|
||||
"cc",
|
||||
"pkg-config",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "oorandom"
|
||||
version = "11.1.5"
|
||||
|
|
@ -1604,6 +1842,18 @@ dependencies = [
|
|||
"winapi",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "paste"
|
||||
version = "1.0.15"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "57c0d7b74b563b49d38dae00a0c37d4d6de9b432382b2892f0574ddcae73fd0a"
|
||||
|
||||
[[package]]
|
||||
name = "pastey"
|
||||
version = "0.2.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "2ee67f1008b1ba2321834326597b8e186293b049a023cdef258527550b9935b4"
|
||||
|
||||
[[package]]
|
||||
name = "percent-encoding"
|
||||
version = "2.3.2"
|
||||
|
|
@ -1850,6 +2100,12 @@ dependencies = [
|
|||
"proc-macro2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "r-efi"
|
||||
version = "5.3.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "69cdb34c158ceb288df11e18b4bd39de994f6657d83847bdffdbd7f346754b0f"
|
||||
|
||||
[[package]]
|
||||
name = "r-efi"
|
||||
version = "6.0.0"
|
||||
|
|
@ -1863,10 +2119,20 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
|||
checksum = "22f6172bdec972074665ed81ed53b71da00bfc44b65a753cfde883ec4c702a1a"
|
||||
dependencies = [
|
||||
"libc",
|
||||
"rand_chacha",
|
||||
"rand_chacha 0.3.1",
|
||||
"rand_core 0.6.4",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rand"
|
||||
version = "0.9.5"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b9ef1d0d795eb7d84685bca4f72f3649f064e6641543d3a8c415898726a57b41"
|
||||
dependencies = [
|
||||
"rand_chacha 0.9.0",
|
||||
"rand_core 0.9.5",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rand"
|
||||
version = "0.10.2"
|
||||
|
|
@ -1888,6 +2154,16 @@ dependencies = [
|
|||
"rand_core 0.6.4",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rand_chacha"
|
||||
version = "0.9.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d3022b5f1df60f26e1ffddd6c66e8aa15de382ae63b3a0c1bfc0e4d3e3f325cb"
|
||||
dependencies = [
|
||||
"ppv-lite86",
|
||||
"rand_core 0.9.5",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rand_core"
|
||||
version = "0.6.4"
|
||||
|
|
@ -1897,6 +2173,15 @@ dependencies = [
|
|||
"getrandom 0.2.17",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rand_core"
|
||||
version = "0.9.5"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "76afc826de14238e6e8c374ddcc1fa19e374fd8dd986b0d2af0d02377261d83c"
|
||||
dependencies = [
|
||||
"getrandom 0.3.4",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rand_core"
|
||||
version = "0.10.1"
|
||||
|
|
@ -1922,6 +2207,17 @@ dependencies = [
|
|||
"rayon-core",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rayon-cond"
|
||||
version = "0.4.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "2964d0cf57a3e7a06e8183d14a8b527195c706b7983549cd5462d5aa3747438f"
|
||||
dependencies = [
|
||||
"either",
|
||||
"itertools 0.14.0",
|
||||
"rayon",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rayon-core"
|
||||
version = "1.13.0"
|
||||
|
|
@ -1979,7 +2275,7 @@ version = "0.12.28"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "eddd3ca559203180a307f12d114c268abf583f59b03cb906fd0b3ff8646c1147"
|
||||
dependencies = [
|
||||
"base64",
|
||||
"base64 0.22.1",
|
||||
"bytes",
|
||||
"futures-channel",
|
||||
"futures-core",
|
||||
|
|
@ -2361,12 +2657,36 @@ dependencies = [
|
|||
"windows-sys 0.61.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "spm_precompiled"
|
||||
version = "0.1.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "5851699c4033c63636f7ea4cf7b7c1f1bf06d0cc03cfb42e711de5a5c46cf326"
|
||||
dependencies = [
|
||||
"base64 0.13.1",
|
||||
"nom",
|
||||
"serde",
|
||||
"unicode-segmentation",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "stable_deref_trait"
|
||||
version = "1.2.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "6ce2be8dc25455e1f91df71bfa12ad37d7af1092ae736f3a6cd0e37bc7810596"
|
||||
|
||||
[[package]]
|
||||
name = "static_assertions"
|
||||
version = "1.1.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "a2eb9349b6444b326872e140eb1cf5e7c522154d69e7a0ffb0fb81c06b37543f"
|
||||
|
||||
[[package]]
|
||||
name = "strsim"
|
||||
version = "0.11.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7da8b5736845d9f2fcb837ea5d9e2628564b3b043a70948a3f0b778838c5fb4f"
|
||||
|
||||
[[package]]
|
||||
name = "subtle"
|
||||
version = "2.6.1"
|
||||
|
|
@ -2535,6 +2855,39 @@ version = "0.1.1"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20"
|
||||
|
||||
[[package]]
|
||||
name = "tokenizers"
|
||||
version = "0.23.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7afbf6e88718afcc138bad01d6ccc3051dbbc3b2ce9793d8b8a3aeb610969cfc"
|
||||
dependencies = [
|
||||
"ahash",
|
||||
"compact_str",
|
||||
"daachorse",
|
||||
"dary_heap",
|
||||
"derive_builder",
|
||||
"esaxx-rs",
|
||||
"getrandom 0.3.4",
|
||||
"itertools 0.14.0",
|
||||
"log",
|
||||
"macro_rules_attribute",
|
||||
"monostate",
|
||||
"onig",
|
||||
"paste",
|
||||
"rand 0.9.5",
|
||||
"rayon",
|
||||
"rayon-cond",
|
||||
"regex",
|
||||
"regex-syntax",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"spm_precompiled",
|
||||
"thiserror 2.0.19",
|
||||
"unicode-normalization-alignments",
|
||||
"unicode-segmentation",
|
||||
"unicode_categories",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tokio"
|
||||
version = "1.53.0"
|
||||
|
|
@ -2773,6 +3126,27 @@ version = "1.0.24"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75"
|
||||
|
||||
[[package]]
|
||||
name = "unicode-normalization-alignments"
|
||||
version = "0.1.12"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "43f613e4fa046e69818dd287fdc4bc78175ff20331479dab6e1b0f98d57062de"
|
||||
dependencies = [
|
||||
"smallvec",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "unicode-segmentation"
|
||||
version = "1.13.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c6f5d3c3b1bf09027a88a6bc961fc00497d651009560b5463668dc81b0fa87a8"
|
||||
|
||||
[[package]]
|
||||
name = "unicode_categories"
|
||||
version = "0.1.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "39ec24b3121d976906ece63c9daad25b85969647682eee313cb5779fdd69e14e"
|
||||
|
||||
[[package]]
|
||||
name = "untrusted"
|
||||
version = "0.9.0"
|
||||
|
|
@ -2856,6 +3230,15 @@ version = "0.11.1+wasi-snapshot-preview1"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b"
|
||||
|
||||
[[package]]
|
||||
name = "wasip2"
|
||||
version = "1.0.4+wasi-0.2.12"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b67efb37e106e55ce722a510d6b5f9c17f083e5fc79afc2badeb12cc313d9487"
|
||||
dependencies = [
|
||||
"wit-bindgen",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "wasm-bindgen"
|
||||
version = "0.2.126"
|
||||
|
|
@ -3081,6 +3464,12 @@ dependencies = [
|
|||
"memchr",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "wit-bindgen"
|
||||
version = "0.57.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "1ebf944e87a7c253233ad6766e082e3cd714b5d03812acc24c318f549614536e"
|
||||
|
||||
[[package]]
|
||||
name = "writeable"
|
||||
version = "0.6.3"
|
||||
|
|
|
|||
|
|
@ -15,6 +15,10 @@ thiserror.workspace = true
|
|||
tracing.workspace = true
|
||||
tracing-subscriber = { workspace = true, optional = true }
|
||||
sha2.workspace = true
|
||||
indexmap = { version = "2.14.0", features = ["serde"] }
|
||||
# HuggingFace tokenizer for input token counting; without the default features it
|
||||
# pulls no HTTP client or progress bars, only the `onig` regex backend.
|
||||
tokenizers = { version = "0.23.1", default-features = false, features = ["onig"] }
|
||||
aws-config = { version = "1.9.0", default-features = false, features = ["rustls", "rt-tokio"], optional = true }
|
||||
aws-credential-types = { version = "1.3.0", features = ["hardcoded-credentials"], optional = true }
|
||||
aws-sdk-sts = { version = "1.108.0", default-features = false, features = ["rustls", "rt-tokio"], optional = true }
|
||||
|
|
|
|||
|
|
@ -43,3 +43,13 @@ pub const EMPTY_TEXT_PLACEHOLDER: &str =
|
|||
"[System: Empty message content sanitised to satisfy protocol]";
|
||||
|
||||
pub const FUNCTION_TRACE_TARGET: &str = "litellm::function_trace";
|
||||
|
||||
/// Message accounting `litellm.token_counter` adds on top of the raw encoding
|
||||
/// for non-OpenAI models (`litellm/litellm_core_utils/token_counter.py`).
|
||||
pub(crate) const TOKENS_PER_MESSAGE: usize = 3;
|
||||
pub(crate) const TOKENS_PER_NAME: usize = 1;
|
||||
pub(crate) const REPLY_PRIMING_TOKENS: usize = 3;
|
||||
pub(crate) const TOOL_DEFINITIONS_TOKENS: usize = 9;
|
||||
pub(crate) const TOOLS_WITH_SYSTEM_MESSAGE_DISCOUNT: usize = 4;
|
||||
pub(crate) const TOOL_CHOICE_NONE_TOKENS: usize = 1;
|
||||
pub(crate) const NAMED_TOOL_CHOICE_TOKENS: usize = 7;
|
||||
|
|
|
|||
|
|
@ -14,5 +14,6 @@ pub mod realtime;
|
|||
pub mod responses;
|
||||
pub mod router;
|
||||
pub mod routing_utils;
|
||||
pub mod token_counter;
|
||||
|
||||
pub use error::Error;
|
||||
|
|
|
|||
207
litellm-rust/crates/core/src/token_counter/mod.rs
Normal file
207
litellm-rust/crates/core/src/token_counter/mod.rs
Normal file
|
|
@ -0,0 +1,207 @@
|
|||
//! Input token counting for a request body, mirroring `litellm.token_counter`
|
||||
//! for the shapes it can count exactly. Everything else is declined so the host
|
||||
//! keeps its own counter as the reference.
|
||||
|
||||
mod python_json;
|
||||
mod tools;
|
||||
pub mod types;
|
||||
|
||||
use serde::Serialize;
|
||||
use thiserror::Error as ThisError;
|
||||
|
||||
use crate::constants::{
|
||||
NAMED_TOOL_CHOICE_TOKENS, REPLY_PRIMING_TOKENS, TOKENS_PER_MESSAGE, TOKENS_PER_NAME,
|
||||
TOOL_CHOICE_NONE_TOKENS, TOOL_DEFINITIONS_TOKENS, TOOLS_WITH_SYSTEM_MESSAGE_DISCOUNT,
|
||||
};
|
||||
use tools::format_function_definitions;
|
||||
use types::{
|
||||
ContentBlock, ContentItem, CountableRequest, Message, MessageContent, TextValue, ToolChoice,
|
||||
};
|
||||
|
||||
#[derive(Debug, ThisError, PartialEq, Eq)]
|
||||
pub enum TokenCountError {
|
||||
#[error("failed to load tokenizer: {0}")]
|
||||
Load(String),
|
||||
/// The body is outside the shape this counter mirrors exactly. Hosts with a
|
||||
/// reference counter treat this as "fall back", not "fail".
|
||||
#[error("unsupported by the rust token counter: {0}")]
|
||||
Unsupported(String),
|
||||
#[error("tokenization failed: {0}")]
|
||||
Encode(String),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize)]
|
||||
pub struct InputTokenCount {
|
||||
pub model: Option<String>,
|
||||
pub input_tokens: usize,
|
||||
}
|
||||
|
||||
/// A loaded HuggingFace tokenizer plus the message accounting Python applies on
|
||||
/// top of it. Encoding is CPU-bound and synchronous; hosts run it off their
|
||||
/// event loop.
|
||||
pub struct TokenCounter {
|
||||
tokenizer: tokenizers::Tokenizer,
|
||||
}
|
||||
|
||||
impl TokenCounter {
|
||||
/// Load a HuggingFace `tokenizer.json` document. The host reads the file.
|
||||
pub fn from_json(tokenizer_json: &str) -> Result<Self, TokenCountError> {
|
||||
let tokenizer = tokenizer_json
|
||||
.parse::<tokenizers::Tokenizer>()
|
||||
.map_err(|error| TokenCountError::Load(error.to_string()))?;
|
||||
Ok(Self { tokenizer })
|
||||
}
|
||||
|
||||
pub fn count_text(&self, text: &str) -> Result<usize, TokenCountError> {
|
||||
self.tokenizer
|
||||
.encode_fast(text, true)
|
||||
.map(|encoding| encoding.len())
|
||||
.map_err(|error| TokenCountError::Encode(error.to_string()))
|
||||
}
|
||||
|
||||
/// Mirrors the host's key precedence: `messages`, then `prompt`, then
|
||||
/// `input`, then `query` plus `documents`.
|
||||
pub fn count_request(
|
||||
&self,
|
||||
request: &CountableRequest,
|
||||
) -> Result<InputTokenCount, TokenCountError> {
|
||||
let input_tokens = if let Some(messages) = &request.messages {
|
||||
self.count_messages(request, messages)?
|
||||
} else if let Some(prompt) = &request.prompt {
|
||||
self.count_text_value(prompt)?
|
||||
} else if let Some(input) = &request.input {
|
||||
self.count_text_value(input)?
|
||||
} else if request.query.is_some() || request.documents.is_some() {
|
||||
self.count_optional_text_value(request.query.as_ref())?
|
||||
+ self.count_optional_text_value(request.documents.as_ref())?
|
||||
} else {
|
||||
return Err(TokenCountError::Unsupported(
|
||||
"request has no countable input".to_string(),
|
||||
));
|
||||
};
|
||||
Ok(InputTokenCount {
|
||||
model: request.model.clone(),
|
||||
input_tokens,
|
||||
})
|
||||
}
|
||||
|
||||
fn count_messages(
|
||||
&self,
|
||||
request: &CountableRequest,
|
||||
messages: &[Message],
|
||||
) -> Result<usize, TokenCountError> {
|
||||
let message_tokens = messages
|
||||
.iter()
|
||||
.map(|message| self.count_message(message))
|
||||
.sum::<Result<usize, _>>()?;
|
||||
let includes_system_message = messages
|
||||
.iter()
|
||||
.any(|message| message.role.as_deref() == Some("system"));
|
||||
let extra_tokens = self.count_extra(
|
||||
request.tools.as_deref().unwrap_or_default(),
|
||||
request.tool_choice.as_ref(),
|
||||
includes_system_message,
|
||||
)?;
|
||||
Ok(message_tokens + extra_tokens)
|
||||
}
|
||||
|
||||
fn count_optional_text_value(
|
||||
&self,
|
||||
value: Option<&TextValue>,
|
||||
) -> Result<usize, TokenCountError> {
|
||||
value.map_or(Ok(0), |value| self.count_text_value(value))
|
||||
}
|
||||
|
||||
/// `str()` for scalars, `json.dumps()` for objects, lists flattened, nulls
|
||||
/// skipped. Floats are declined because Python's `repr` and Rust's float
|
||||
/// formatting disagree on exponents.
|
||||
fn count_text_value(&self, value: &TextValue) -> Result<usize, TokenCountError> {
|
||||
match value {
|
||||
TextValue::Null => Ok(0),
|
||||
TextValue::Bool(true) => self.count_text("True"),
|
||||
TextValue::Bool(false) => self.count_text("False"),
|
||||
TextValue::Integer(number) => self.count_text(&number.to_string()),
|
||||
TextValue::Float(_) => Err(TokenCountError::Unsupported(
|
||||
"float text values are counted by the python path".to_string(),
|
||||
)),
|
||||
TextValue::Text(text) => self.count_text(text),
|
||||
TextValue::List(items) => items
|
||||
.iter()
|
||||
.map(|item| self.count_text_value(item))
|
||||
.sum::<Result<usize, _>>(),
|
||||
TextValue::Object(_) => self.count_text(&python_json::dumps(value)?),
|
||||
}
|
||||
}
|
||||
|
||||
fn count_message(&self, message: &Message) -> Result<usize, TokenCountError> {
|
||||
let role_tokens = match &message.role {
|
||||
Some(role) => self.count_text(role)?,
|
||||
None => 0,
|
||||
};
|
||||
let name_tokens = match &message.name {
|
||||
Some(name) => self.count_text(name)? + TOKENS_PER_NAME,
|
||||
None => 0,
|
||||
};
|
||||
let content_tokens = match &message.content {
|
||||
Some(MessageContent::Text(text)) => self.count_text(text)?,
|
||||
Some(MessageContent::Blocks(items)) => items
|
||||
.iter()
|
||||
.map(|item| self.count_content_item(item))
|
||||
.sum::<Result<usize, _>>()?,
|
||||
None => 0,
|
||||
};
|
||||
Ok(TOKENS_PER_MESSAGE + role_tokens + name_tokens + content_tokens)
|
||||
}
|
||||
|
||||
fn count_content_item(&self, item: &ContentItem) -> Result<usize, TokenCountError> {
|
||||
match item {
|
||||
ContentItem::Text(text) => self.count_text(text),
|
||||
ContentItem::Block(ContentBlock::Text { text }) => self.count_text(text),
|
||||
ContentItem::Block(ContentBlock::Thinking { thinking }) => {
|
||||
if thinking.is_empty() {
|
||||
return Ok(0);
|
||||
}
|
||||
self.count_text(thinking)
|
||||
}
|
||||
ContentItem::Block(ContentBlock::ToolReference { tool_name }) => {
|
||||
match tool_name.as_deref().filter(|name| !name.is_empty()) {
|
||||
Some(name) => self.count_text(name),
|
||||
None => Ok(0),
|
||||
}
|
||||
}
|
||||
ContentItem::Block(ContentBlock::Unsupported) => Err(TokenCountError::Unsupported(
|
||||
"content block type is counted by the python path".to_string(),
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn count_extra(
|
||||
&self,
|
||||
tools: &[types::ToolDefinition],
|
||||
tool_choice: Option<&ToolChoice>,
|
||||
includes_system_message: bool,
|
||||
) -> Result<usize, TokenCountError> {
|
||||
let tool_tokens = if tools.is_empty() {
|
||||
0
|
||||
} else {
|
||||
let definitions = self.count_text(&format_function_definitions(tools)?)?;
|
||||
let discount = if includes_system_message {
|
||||
TOOLS_WITH_SYSTEM_MESSAGE_DISCOUNT
|
||||
} else {
|
||||
0
|
||||
};
|
||||
definitions + TOOL_DEFINITIONS_TOKENS - discount
|
||||
};
|
||||
let choice_tokens = match tool_choice {
|
||||
Some(ToolChoice::Mode(mode)) if mode == "none" => TOOL_CHOICE_NONE_TOKENS,
|
||||
Some(ToolChoice::Mode(_)) | None => 0,
|
||||
Some(ToolChoice::Named(named)) => {
|
||||
NAMED_TOOL_CHOICE_TOKENS + self.count_text(&named.function.name)?
|
||||
}
|
||||
};
|
||||
Ok(REPLY_PRIMING_TOKENS + tool_tokens + choice_tokens)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
79
litellm-rust/crates/core/src/token_counter/python_json.rs
Normal file
79
litellm-rust/crates/core/src/token_counter/python_json.rs
Normal file
|
|
@ -0,0 +1,79 @@
|
|||
//! `json.dumps(value)` with Python's default arguments: `", "` and `": "`
|
||||
//! separators, `ensure_ascii=True`, and keys in insertion order.
|
||||
|
||||
use std::fmt::Write;
|
||||
|
||||
use super::TokenCountError;
|
||||
use super::types::TextValue;
|
||||
|
||||
pub(super) fn dumps(value: &TextValue) -> Result<String, TokenCountError> {
|
||||
let mut out = String::new();
|
||||
write_value(&mut out, value)?;
|
||||
Ok(out)
|
||||
}
|
||||
|
||||
fn write_value(out: &mut String, value: &TextValue) -> Result<(), TokenCountError> {
|
||||
match value {
|
||||
TextValue::Null => out.push_str("null"),
|
||||
TextValue::Bool(true) => out.push_str("true"),
|
||||
TextValue::Bool(false) => out.push_str("false"),
|
||||
TextValue::Integer(number) => write_number(out, number),
|
||||
TextValue::Float(_) => {
|
||||
return Err(TokenCountError::Unsupported(
|
||||
"float repr is formatted by the python path".to_string(),
|
||||
));
|
||||
}
|
||||
TextValue::Text(text) => write_string(out, text),
|
||||
TextValue::List(items) => {
|
||||
out.push('[');
|
||||
for (index, item) in items.iter().enumerate() {
|
||||
if index > 0 {
|
||||
out.push_str(", ");
|
||||
}
|
||||
write_value(out, item)?;
|
||||
}
|
||||
out.push(']');
|
||||
}
|
||||
TextValue::Object(entries) => {
|
||||
out.push('{');
|
||||
for (index, (key, item)) in entries.iter().enumerate() {
|
||||
if index > 0 {
|
||||
out.push_str(", ");
|
||||
}
|
||||
write_string(out, key);
|
||||
out.push_str(": ");
|
||||
write_value(out, item)?;
|
||||
}
|
||||
out.push('}');
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn write_number(out: &mut String, number: &i64) {
|
||||
// Writing an integer into a String cannot fail.
|
||||
let _ = write!(out, "{number}");
|
||||
}
|
||||
|
||||
fn write_string(out: &mut String, text: &str) {
|
||||
out.push('"');
|
||||
for character in text.chars() {
|
||||
match character {
|
||||
'"' => out.push_str("\\\""),
|
||||
'\\' => out.push_str("\\\\"),
|
||||
'\n' => out.push_str("\\n"),
|
||||
'\r' => out.push_str("\\r"),
|
||||
'\t' => out.push_str("\\t"),
|
||||
'\u{08}' => out.push_str("\\b"),
|
||||
'\u{0c}' => out.push_str("\\f"),
|
||||
' '..='~' => out.push(character),
|
||||
_ => {
|
||||
let mut units = [0u16; 2];
|
||||
for unit in character.encode_utf16(&mut units) {
|
||||
let _ = write!(out, "\\u{unit:04x}");
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
out.push('"');
|
||||
}
|
||||
215
litellm-rust/crates/core/src/token_counter/tests.rs
Normal file
215
litellm-rust/crates/core/src/token_counter/tests.rs
Normal file
|
|
@ -0,0 +1,215 @@
|
|||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
/// Expected counts are pinned from `litellm.token_counter(model="claude-sonnet-4-5", ...)`
|
||||
/// so this test also guards Python parity.
|
||||
fn counter() -> 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");
|
||||
TokenCounter::from_json(&json).expect("anthropic tokenizer loads")
|
||||
}
|
||||
|
||||
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."},
|
||||
"plain string item",
|
||||
{"type":"thinking","thinking":"pondering"},
|
||||
{"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?"}],
|
||||
"tools":[
|
||||
{"type":"function","function":{"name":"get_weather","description":"Get weather","parameters":{
|
||||
"type":"object",
|
||||
"properties":{
|
||||
"location":{"type":"string","description":"City name"},
|
||||
"unit":{"type":"string","enum":["celsius","fahrenheit"]},
|
||||
"days":{"type":"integer"},
|
||||
"tags":{"type":"array","items":{"type":"string"}},
|
||||
"opts":{"type":"object","properties":{"verbose":{"type":"boolean"},"level":{"type":"integer","enum":[1,2]}},"required":["verbose"]},
|
||||
"anything":{}},
|
||||
"required":["location"]}}},
|
||||
{"type":"function","function":{"name":"noop"}}],
|
||||
"tool_choice":{"type":"function","function":{"name":"get_weather"}}}"#;
|
||||
|
||||
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_LIST: &str =
|
||||
r#"{"model":"claude-sonnet-4-5","prompt":["first prompt","second prompt"]}"#;
|
||||
|
||||
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 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"]}]}"#;
|
||||
|
||||
/// Expected counts are pinned from
|
||||
/// `litellm.proxy.spend_tracking.budget_reservation._count_input_tokens(body, "claude-sonnet-4-5")`.
|
||||
#[rstest]
|
||||
#[case::text_only(SIMPLE, 14)]
|
||||
#[case::content_blocks_name_and_system(BLOCKS_AND_SYSTEM, 45)]
|
||||
#[case::openai_tools_named_choice(TOOLS_OPENAI, 123)]
|
||||
#[case::anthropic_tools_system_discount_choice_none(TOOLS_ANTHROPIC_SYSTEM, 53)]
|
||||
#[case::completions_prompt(COMPLETIONS_PROMPT, 7)]
|
||||
#[case::completions_prompt_list(COMPLETIONS_PROMPT_LIST, 4)]
|
||||
#[case::responses_input_items(RESPONSES_INPUT, 62)]
|
||||
#[case::embeddings_token_ids(EMBEDDINGS_TOKEN_IDS, 5)]
|
||||
#[case::rerank_query_and_documents(RERANK, 41)]
|
||||
fn count_request_matches_python_token_counter(#[case] body: &str, #[case] expected: usize) {
|
||||
let request = CountableRequest::parse(body.as_bytes()).expect("fixture parses");
|
||||
let count = counter().count_request(&request).expect("fixture counts");
|
||||
assert_eq!(
|
||||
count,
|
||||
InputTokenCount {
|
||||
model: Some("claude-sonnet-4-5".to_string()),
|
||||
input_tokens: 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) {
|
||||
let request = CountableRequest::parse(body.as_bytes()).expect("fixture parses");
|
||||
let count = counter().count_request(&request).expect("fixture counts");
|
||||
assert_eq!(count.input_tokens, expected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn objects_dump_like_python_json_dumps() {
|
||||
let body = r#"{"model":"m","input":{"text":"caf\u00e9 \u2014 \ud83d\ude00 \"q\" \\ \n\t\u0001\u007f ~","n":-3,"ok":true,"no":false,"none":null,"list":[1,"a",{"z":[]}],"empty":{}}}"#;
|
||||
let request = CountableRequest::parse(body.as_bytes()).expect("fixture parses");
|
||||
let dumped = python_json::dumps(request.input.as_ref().expect("input is present"))
|
||||
.expect("fixture dumps");
|
||||
assert_eq!(
|
||||
dumped,
|
||||
r#"{"text": "caf\u00e9 \u2014 \ud83d\ude00 \"q\" \\ \n\t\u0001\u007f ~", "n": -3, "ok": true, "no": false, "none": null, "list": [1, "a", {"z": []}], "empty": {}}"#
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tool_definitions_render_like_python() {
|
||||
let request = CountableRequest::parse(TOOLS_OPENAI.as_bytes()).expect("fixture parses");
|
||||
let rendered = format_function_definitions(request.tools.as_deref().unwrap_or_default())
|
||||
.expect("fixture renders");
|
||||
let expected = "namespace functions {\n\n// Get weather\ntype get_weather = (_: {\n// City name\nlocation: string,\nunit?: \"celsius\" | \"fahrenheit\",\ndays?: number,\ntags?: string[],\nopts?: {\n verbose: boolean,\n level?: \"1\" | \"2\",\n},\nanything?: any,\n}) => any;\n\ntype noop = () => any;\n\n} // namespace functions";
|
||||
assert_eq!(rendered, expected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn union_types_and_anthropic_schema_render_like_python() {
|
||||
let request =
|
||||
CountableRequest::parse(TOOLS_ANTHROPIC_SYSTEM.as_bytes()).expect("fixture parses");
|
||||
let rendered = format_function_definitions(request.tools.as_deref().unwrap_or_default())
|
||||
.expect("fixture renders");
|
||||
assert_eq!(
|
||||
rendered,
|
||||
"namespace functions {\n\n// Get weather\ntype get_weather = (_: {\nlocation: any,\n}) => any;\n\n} // namespace functions"
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::not_json(b"not json" as &[u8])]
|
||||
#[case::messages_not_a_list(br#"{"model":"m","messages":"hi"}"#)]
|
||||
#[case::message_with_tool_calls(
|
||||
br#"{"model":"m","messages":[{"role":"assistant","tool_calls":[{"id":"1","type":"function","function":{"name":"f","arguments":"{}"}}]}]}"#
|
||||
)]
|
||||
#[case::dict_content(
|
||||
br#"{"model":"m","messages":[{"role":"user","content":{"type":"text","text":"x"}}]}"#
|
||||
)]
|
||||
#[case::float_enum(
|
||||
br#"{"model":"m","messages":[],"tools":[{"name":"f","input_schema":{"type":"object","properties":{"x":{"type":"number","enum":[1.5]}}}}]}"#
|
||||
)]
|
||||
#[case::anthropic_tool_choice_without_function(
|
||||
br#"{"model":"m","messages":[],"tool_choice":{"type":"auto"}}"#
|
||||
)]
|
||||
fn shapes_outside_the_mirror_are_declined_at_parse(#[case] body: &[u8]) {
|
||||
assert!(matches!(
|
||||
CountableRequest::parse(body),
|
||||
Err(TokenCountError::Unsupported(_))
|
||||
));
|
||||
}
|
||||
|
||||
#[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]) {
|
||||
let request = CountableRequest::parse(body).expect("shape parses");
|
||||
assert!(matches!(
|
||||
counter().count_request(&request),
|
||||
Err(TokenCountError::Unsupported(_))
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tool_choice_and_system_discount_change_the_count() {
|
||||
let counter = counter();
|
||||
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 + TOOL_CHOICE_NONE_TOKENS
|
||||
);
|
||||
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,
|
||||
TOOLS_WITH_SYSTEM_MESSAGE_DISCOUNT
|
||||
);
|
||||
assert_eq!(
|
||||
count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tools":[]}"#),
|
||||
base
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn loading_a_bad_tokenizer_is_a_load_error() {
|
||||
assert!(matches!(
|
||||
TokenCounter::from_json("{}"),
|
||||
Err(TokenCountError::Load(_))
|
||||
));
|
||||
}
|
||||
108
litellm-rust/crates/core/src/token_counter/tools.rs
Normal file
108
litellm-rust/crates/core/src/token_counter/tools.rs
Normal file
|
|
@ -0,0 +1,108 @@
|
|||
//! Renders tool definitions the way `litellm.token_counter` does before
|
||||
//! tokenizing them (the TypeScript-like namespace OpenAI appears to use).
|
||||
|
||||
use super::TokenCountError;
|
||||
use super::types::{EnumValue, FunctionDefinition, Schema, SchemaType, ToolDefinition};
|
||||
|
||||
pub(super) fn format_function_definitions(
|
||||
tools: &[ToolDefinition],
|
||||
) -> Result<String, TokenCountError> {
|
||||
let mut lines = vec!["namespace functions {".to_string(), String::new()];
|
||||
for tool in tools {
|
||||
let function = resolve_function(tool);
|
||||
let Some(name) = function.name.as_deref().filter(|name| !name.is_empty()) else {
|
||||
continue;
|
||||
};
|
||||
if let Some(description) = function.description.as_deref().filter(|d| !d.is_empty()) {
|
||||
lines.push(format!("// {description}"));
|
||||
}
|
||||
let parameters = function.parameters.unwrap_or_default();
|
||||
match ¶meters.properties {
|
||||
Some(properties) if !properties.is_empty() => {
|
||||
lines.push(format!("type {name} = (_: {{"));
|
||||
lines.push(format_object_parameters(¶meters, 0)?);
|
||||
lines.push("}) => any;".to_string());
|
||||
}
|
||||
_ => lines.push(format!("type {name} = () => any;")),
|
||||
}
|
||||
lines.push(String::new());
|
||||
}
|
||||
lines.push("} // namespace functions".to_string());
|
||||
Ok(lines.join("\n"))
|
||||
}
|
||||
|
||||
fn resolve_function(tool: &ToolDefinition) -> FunctionDefinition {
|
||||
match &tool.function {
|
||||
Some(function) => function.clone(),
|
||||
None => FunctionDefinition {
|
||||
name: tool.name.clone(),
|
||||
description: tool.description.clone(),
|
||||
parameters: tool
|
||||
.input_schema
|
||||
.clone()
|
||||
.or_else(|| tool.parameters.clone()),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
fn format_object_parameters(parameters: &Schema, indent: usize) -> Result<String, TokenCountError> {
|
||||
let Some(properties) = parameters.properties.as_ref().filter(|p| !p.is_empty()) else {
|
||||
return Ok(String::new());
|
||||
};
|
||||
let required = parameters.required.as_deref().unwrap_or_default();
|
||||
let mut lines = Vec::new();
|
||||
for (key, props) in properties {
|
||||
if let Some(description) = props.description.as_deref().filter(|d| !d.is_empty()) {
|
||||
lines.push(format!("// {description}"));
|
||||
}
|
||||
let question = if required.iter().any(|r| r == key) {
|
||||
""
|
||||
} else {
|
||||
"?"
|
||||
};
|
||||
lines.push(format!("{key}{question}: {},", format_type(props, indent)?));
|
||||
}
|
||||
let pad = " ".repeat(indent);
|
||||
Ok(lines
|
||||
.iter()
|
||||
.map(|line| format!("{pad}{line}"))
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n"))
|
||||
}
|
||||
|
||||
fn format_type(props: &Schema, indent: usize) -> Result<String, TokenCountError> {
|
||||
let Some(SchemaType::Name(schema_type)) = &props.schema_type else {
|
||||
return Ok("any".to_string());
|
||||
};
|
||||
match schema_type.as_str() {
|
||||
"string" | "integer" | "number" => Ok(match &props.enum_values {
|
||||
Some(values) => format_enum(values),
|
||||
None if schema_type == "string" => "string".to_string(),
|
||||
None => "number".to_string(),
|
||||
}),
|
||||
"array" => {
|
||||
let items = props.items.as_deref().ok_or(TokenCountError::Unsupported(
|
||||
"array parameter without items".to_string(),
|
||||
))?;
|
||||
Ok(format!("{}[]", format_type(items, indent)?))
|
||||
}
|
||||
"object" => Ok(format!(
|
||||
"{{\n{}\n}}",
|
||||
format_object_parameters(props, indent + 2)?
|
||||
)),
|
||||
"boolean" => Ok("boolean".to_string()),
|
||||
"null" => Ok("null".to_string()),
|
||||
_ => Ok("any".to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
fn format_enum(values: &[EnumValue]) -> String {
|
||||
values
|
||||
.iter()
|
||||
.map(|value| match value {
|
||||
EnumValue::Text(text) => format!("\"{text}\""),
|
||||
EnumValue::Integer(number) => format!("\"{number}\""),
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
.join(" | ")
|
||||
}
|
||||
161
litellm-rust/crates/core/src/token_counter/types.rs
Normal file
161
litellm-rust/crates/core/src/token_counter/types.rs
Normal file
|
|
@ -0,0 +1,161 @@
|
|||
use indexmap::IndexMap;
|
||||
use serde::{Deserialize, Deserializer};
|
||||
|
||||
use super::TokenCountError;
|
||||
|
||||
/// The parts of a request body the host's budget counter reads. Chat and
|
||||
/// Anthropic Messages bodies carry `messages`; completions carry `prompt`;
|
||||
/// Responses and embeddings carry `input`; rerank carries `query` and
|
||||
/// `documents`. The host checks key presence, not nullness, so an explicit
|
||||
/// `null` is kept distinct from an absent key. Anything outside this shape is
|
||||
/// declined so the host can fall back to its own counter instead of silently
|
||||
/// miscounting.
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq)]
|
||||
pub struct CountableRequest {
|
||||
pub model: Option<String>,
|
||||
#[serde(default, deserialize_with = "present_messages")]
|
||||
pub messages: Option<Vec<Message>>,
|
||||
pub tools: Option<Vec<ToolDefinition>>,
|
||||
pub tool_choice: Option<ToolChoice>,
|
||||
#[serde(default, deserialize_with = "present_text")]
|
||||
pub prompt: Option<TextValue>,
|
||||
#[serde(default, deserialize_with = "present_text")]
|
||||
pub input: Option<TextValue>,
|
||||
#[serde(default, deserialize_with = "present_text")]
|
||||
pub query: Option<TextValue>,
|
||||
#[serde(default, deserialize_with = "present_text")]
|
||||
pub documents: Option<TextValue>,
|
||||
}
|
||||
|
||||
impl CountableRequest {
|
||||
pub fn parse(body: &[u8]) -> Result<Self, TokenCountError> {
|
||||
serde_json::from_slice(body)
|
||||
.map_err(|error| TokenCountError::Unsupported(error.to_string()))
|
||||
}
|
||||
}
|
||||
|
||||
fn present_messages<'de, D: Deserializer<'de>>(
|
||||
deserializer: D,
|
||||
) -> Result<Option<Vec<Message>>, D::Error> {
|
||||
Option::<Vec<Message>>::deserialize(deserializer)
|
||||
.map(|messages| Some(messages.unwrap_or_default()))
|
||||
}
|
||||
|
||||
fn present_text<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Option<TextValue>, D::Error> {
|
||||
TextValue::deserialize(deserializer).map(Some)
|
||||
}
|
||||
|
||||
/// Free-form JSON the host counts as text: strings and integers via `str()`,
|
||||
/// objects via `json.dumps()`, lists flattened. Objects keep document order so
|
||||
/// the dumped text matches Python byte for byte.
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq)]
|
||||
#[serde(untagged)]
|
||||
pub enum TextValue {
|
||||
Null,
|
||||
Bool(bool),
|
||||
Integer(i64),
|
||||
Float(f64),
|
||||
Text(String),
|
||||
List(Vec<TextValue>),
|
||||
Object(IndexMap<String, TextValue>),
|
||||
}
|
||||
|
||||
/// Python counts every string-valued key of a message, so any key beyond these
|
||||
/// makes the shape unsupported rather than silently uncounted.
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct Message {
|
||||
pub role: Option<String>,
|
||||
pub name: Option<String>,
|
||||
pub content: Option<MessageContent>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq)]
|
||||
#[serde(untagged)]
|
||||
pub enum MessageContent {
|
||||
Text(String),
|
||||
Blocks(Vec<ContentItem>),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq)]
|
||||
#[serde(untagged)]
|
||||
pub enum ContentItem {
|
||||
Text(String),
|
||||
Block(ContentBlock),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq)]
|
||||
#[serde(tag = "type")]
|
||||
pub enum ContentBlock {
|
||||
#[serde(rename = "text")]
|
||||
Text { text: String },
|
||||
#[serde(rename = "thinking")]
|
||||
Thinking { thinking: String },
|
||||
#[serde(rename = "tool_reference")]
|
||||
ToolReference { tool_name: Option<String> },
|
||||
/// Images, documents, files and tool use/result blocks price through
|
||||
/// Python-only helpers, so they stay on the Python counter.
|
||||
#[serde(other)]
|
||||
Unsupported,
|
||||
}
|
||||
|
||||
/// Either the OpenAI `{"type": "function", "function": {...}}` shape or the
|
||||
/// Anthropic `{"name", "description", "input_schema"}` shape.
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq)]
|
||||
pub struct ToolDefinition {
|
||||
pub function: Option<FunctionDefinition>,
|
||||
pub name: Option<String>,
|
||||
pub description: Option<String>,
|
||||
pub input_schema: Option<Schema>,
|
||||
pub parameters: Option<Schema>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq)]
|
||||
pub struct FunctionDefinition {
|
||||
pub name: Option<String>,
|
||||
pub description: Option<String>,
|
||||
pub parameters: Option<Schema>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, Deserialize, PartialEq)]
|
||||
pub struct Schema {
|
||||
#[serde(rename = "type")]
|
||||
pub schema_type: Option<SchemaType>,
|
||||
pub description: Option<String>,
|
||||
#[serde(rename = "enum")]
|
||||
pub enum_values: Option<Vec<EnumValue>>,
|
||||
pub items: Option<Box<Schema>>,
|
||||
pub properties: Option<IndexMap<String, Schema>>,
|
||||
pub required: Option<Vec<String>>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq)]
|
||||
#[serde(untagged)]
|
||||
pub enum SchemaType {
|
||||
Name(String),
|
||||
Union(Vec<String>),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq)]
|
||||
#[serde(untagged)]
|
||||
pub enum EnumValue {
|
||||
Text(String),
|
||||
Integer(i64),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq)]
|
||||
#[serde(untagged)]
|
||||
pub enum ToolChoice {
|
||||
Mode(String),
|
||||
Named(NamedToolChoice),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq)]
|
||||
pub struct NamedToolChoice {
|
||||
pub function: NamedFunction,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq)]
|
||||
pub struct NamedFunction {
|
||||
pub name: String,
|
||||
}
|
||||
|
|
@ -3,7 +3,6 @@ use std::panic::AssertUnwindSafe;
|
|||
use std::time::Duration;
|
||||
|
||||
use futures_util::FutureExt;
|
||||
use litellm_core::error::Error;
|
||||
use litellm_python_interop::{Pythonized, panic_to_pyerr, release_gil};
|
||||
use pyo3::exceptions::PyRuntimeError;
|
||||
use pyo3::prelude::*;
|
||||
|
|
@ -11,14 +10,15 @@ use serde::Serialize;
|
|||
use tokio::runtime::{Handle, Runtime};
|
||||
use tokio::time::{self, MissedTickBehavior};
|
||||
|
||||
pub(crate) fn run_sync<T, F>(
|
||||
pub(crate) fn run_sync<T, E, F>(
|
||||
py: Python<'_>,
|
||||
future: F,
|
||||
map_error: fn(Error) -> PyErr,
|
||||
map_error: fn(E) -> PyErr,
|
||||
) -> PyResult<Py<PyAny>>
|
||||
where
|
||||
T: Serialize + Send + 'static,
|
||||
F: Future<Output = Result<T, Error>> + Send + 'static,
|
||||
E: Send + 'static,
|
||||
F: Future<Output = Result<T, E>> + Send + 'static,
|
||||
{
|
||||
run_sync_on(
|
||||
py,
|
||||
|
|
@ -28,15 +28,16 @@ where
|
|||
)
|
||||
}
|
||||
|
||||
fn run_sync_on<T, F>(
|
||||
fn run_sync_on<T, E, F>(
|
||||
py: Python<'_>,
|
||||
runtime: &Runtime,
|
||||
future: F,
|
||||
map_error: fn(Error) -> PyErr,
|
||||
map_error: fn(E) -> PyErr,
|
||||
) -> PyResult<Py<PyAny>>
|
||||
where
|
||||
T: Serialize + Send + 'static,
|
||||
F: Future<Output = Result<T, Error>> + Send + 'static,
|
||||
E: Send + 'static,
|
||||
F: Future<Output = Result<T, E>> + Send + 'static,
|
||||
{
|
||||
if Handle::try_current().is_ok() {
|
||||
return Err(PyRuntimeError::new_err(
|
||||
|
|
@ -49,14 +50,15 @@ where
|
|||
Pythonized(result).into_pyobject(py).map(Bound::unbind)
|
||||
}
|
||||
|
||||
pub(crate) fn run_async<T, F>(
|
||||
pub(crate) fn run_async<T, E, F>(
|
||||
py: Python<'_>,
|
||||
future: F,
|
||||
map_error: fn(Error) -> PyErr,
|
||||
map_error: fn(E) -> PyErr,
|
||||
) -> PyResult<Bound<'_, PyAny>>
|
||||
where
|
||||
T: Serialize + Send + 'static,
|
||||
F: Future<Output = Result<T, Error>> + Send + 'static,
|
||||
E: Send + 'static,
|
||||
F: Future<Output = Result<T, E>> + Send + 'static,
|
||||
{
|
||||
pyo3_async_runtimes::tokio::future_into_py(py, async move {
|
||||
let result = catch_future_panic(future).await?;
|
||||
|
|
@ -65,7 +67,7 @@ where
|
|||
})
|
||||
}
|
||||
|
||||
fn map_core_result<T>(result: Result<T, Error>, map_error: fn(Error) -> PyErr) -> PyResult<T> {
|
||||
fn map_core_result<T, E>(result: Result<T, E>, map_error: fn(E) -> PyErr) -> PyResult<T> {
|
||||
match result {
|
||||
Ok(value) => Ok(value),
|
||||
Err(error) => Err(
|
||||
|
|
@ -75,9 +77,9 @@ fn map_core_result<T>(result: Result<T, Error>, map_error: fn(Error) -> PyErr) -
|
|||
}
|
||||
}
|
||||
|
||||
async fn catch_future_panic<T, F>(future: F) -> PyResult<Result<T, Error>>
|
||||
async fn catch_future_panic<T, E, F>(future: F) -> PyResult<Result<T, E>>
|
||||
where
|
||||
F: Future<Output = Result<T, Error>>,
|
||||
F: Future<Output = Result<T, E>>,
|
||||
{
|
||||
AssertUnwindSafe(future)
|
||||
.catch_unwind()
|
||||
|
|
@ -85,9 +87,9 @@ where
|
|||
.map_err(panic_to_pyerr)
|
||||
}
|
||||
|
||||
async fn wait_for_sync_result<T, F>(future: F) -> PyResult<Result<T, Error>>
|
||||
async fn wait_for_sync_result<T, E, F>(future: F) -> PyResult<Result<T, E>>
|
||||
where
|
||||
F: Future<Output = Result<T, Error>>,
|
||||
F: Future<Output = Result<T, E>>,
|
||||
{
|
||||
let future = catch_future_panic(future);
|
||||
tokio::pin!(future);
|
||||
|
|
@ -114,6 +116,7 @@ mod tests {
|
|||
use std::thread;
|
||||
use std::time::Instant;
|
||||
|
||||
use litellm_core::error::Error;
|
||||
use pyo3::panic::PanicException;
|
||||
use pyo3::types::{PyDict, PyModule};
|
||||
use serde::Serializer;
|
||||
|
|
@ -237,7 +240,7 @@ mod tests {
|
|||
|
||||
let error = runtime.block_on(async {
|
||||
Python::attach(|py| {
|
||||
run_sync::<bool, _>(py, async { Ok(true) }, runtime_error)
|
||||
run_sync::<bool, Error, _>(py, async { Ok(true) }, runtime_error)
|
||||
.expect_err("sync route should reject a nested Tokio runtime")
|
||||
})
|
||||
});
|
||||
|
|
@ -273,7 +276,7 @@ mod tests {
|
|||
fn sync_runner_maps_a_panicked_future() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let error = run_sync::<bool, _>(
|
||||
let error = run_sync::<bool, Error, _>(
|
||||
py,
|
||||
poll_fn(|_| -> Poll<Result<bool, Error>> { panic!("route future panicked") }),
|
||||
runtime_error,
|
||||
|
|
@ -289,7 +292,7 @@ mod tests {
|
|||
fn sync_runner_maps_a_panicked_error_mapper() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let error = run_sync::<bool, _>(
|
||||
let error = run_sync::<bool, Error, _>(
|
||||
py,
|
||||
async { Err(Error::InvalidRequest("invalid".to_string())) },
|
||||
panicking_error_mapper,
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ mod execution;
|
|||
mod function_trace;
|
||||
mod marshal;
|
||||
mod routes;
|
||||
mod token_counter;
|
||||
|
||||
use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection;
|
||||
use pyo3::prelude::*;
|
||||
|
|
@ -71,6 +72,7 @@ mod _native {
|
|||
super::errors::register(module)?;
|
||||
super::routes::register(module)?;
|
||||
module.add_class::<super::ResponsesWebSocketConnection>()?;
|
||||
super::token_counter::register(module)?;
|
||||
super::diagnostics::register(module)
|
||||
}
|
||||
}
|
||||
|
|
@ -106,6 +108,7 @@ mod tests {
|
|||
"chat_completions",
|
||||
"achat_completions",
|
||||
"ResponsesWebSocketConnection",
|
||||
"TokenCounter",
|
||||
"gil_stats",
|
||||
];
|
||||
|
||||
|
|
|
|||
64
litellm-rust/crates/python-bridge/src/token_counter.rs
Normal file
64
litellm-rust/crates/python-bridge/src/token_counter.rs
Normal file
|
|
@ -0,0 +1,64 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use litellm_core::token_counter::types::CountableRequest;
|
||||
use litellm_core::token_counter::{
|
||||
InputTokenCount, TokenCountError, TokenCounter as CoreTokenCounter,
|
||||
};
|
||||
use litellm_python_interop::release_gil;
|
||||
use pyo3::exceptions::{PyRuntimeError, PyValueError};
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::PyAny;
|
||||
|
||||
use crate::errors::RustBridgeDeclined;
|
||||
use crate::execution::run_async;
|
||||
|
||||
/// Counts the input tokens of a raw request body off the Python event loop with
|
||||
/// the GIL released. Python owns which requests get here and what to do with
|
||||
/// the count.
|
||||
#[pyclass(frozen)]
|
||||
struct TokenCounter {
|
||||
inner: Arc<CoreTokenCounter>,
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl TokenCounter {
|
||||
#[new]
|
||||
fn new(py: Python<'_>, tokenizer_json: &str) -> PyResult<Self> {
|
||||
let inner = release_gil(py, || CoreTokenCounter::from_json(tokenizer_json))
|
||||
.map_err(token_count_error_to_pyerr)?;
|
||||
Ok(Self {
|
||||
inner: Arc::new(inner),
|
||||
})
|
||||
}
|
||||
|
||||
fn acount_request<'py>(&self, py: Python<'py>, body: &[u8]) -> PyResult<Bound<'py, PyAny>> {
|
||||
let counter = Arc::clone(&self.inner);
|
||||
let body = body.to_vec();
|
||||
run_async(
|
||||
py,
|
||||
async move {
|
||||
tokio::task::spawn_blocking(move || count_body(&counter, &body))
|
||||
.await
|
||||
.map_err(|error| TokenCountError::Encode(error.to_string()))?
|
||||
},
|
||||
token_count_error_to_pyerr,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
fn count_body(counter: &CoreTokenCounter, body: &[u8]) -> Result<InputTokenCount, TokenCountError> {
|
||||
let request = CountableRequest::parse(body)?;
|
||||
counter.count_request(&request)
|
||||
}
|
||||
|
||||
fn token_count_error_to_pyerr(error: TokenCountError) -> PyErr {
|
||||
match error {
|
||||
TokenCountError::Load(message) => PyValueError::new_err(message),
|
||||
TokenCountError::Unsupported(message) => RustBridgeDeclined::new_err(message),
|
||||
TokenCountError::Encode(message) => PyRuntimeError::new_err(message),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
module.add_class::<TokenCounter>()
|
||||
}
|
||||
|
|
@ -91,6 +91,7 @@ from litellm.proxy.common_utils.http_parsing_utils import (
|
|||
_safe_get_request_query_params,
|
||||
_safe_set_request_parsed_body,
|
||||
populate_request_with_path_params,
|
||||
read_raw_json_body,
|
||||
)
|
||||
from litellm.proxy.common_utils.realtime_utils import _realtime_request_body
|
||||
from litellm.proxy.common_utils.user_api_key_cache import (
|
||||
|
|
@ -2650,6 +2651,7 @@ async def _run_centralized_common_checks(
|
|||
|
||||
await _reserve_budget_after_common_checks(
|
||||
user_api_key_auth_obj=user_api_key_auth_obj,
|
||||
request=request,
|
||||
request_data=request_data,
|
||||
route=route,
|
||||
llm_router=llm_router,
|
||||
|
|
@ -2685,6 +2687,7 @@ async def _reserve_budget_after_common_checks(
|
|||
general_settings: dict,
|
||||
end_user_id: str | None = None,
|
||||
end_user_object: LiteLLM_EndUserTable | None = None,
|
||||
request: Request | None = None,
|
||||
) -> None:
|
||||
user_api_key_auth_obj.budget_reservation = None
|
||||
if skip_budget_checks:
|
||||
|
|
@ -2710,6 +2713,7 @@ async def _reserve_budget_after_common_checks(
|
|||
end_user_object=end_user_object,
|
||||
apply_user_budget_to_team_keys=general_settings.get("apply_user_budget_to_team_keys") is True,
|
||||
fail_closed_budget_enforcement=general_settings.get("fail_closed_budget_enforcement") is True,
|
||||
raw_body=await read_raw_json_body(request=request),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -213,6 +213,18 @@ async def _read_request_body(request: Request | None) -> dict:
|
|||
return {}
|
||||
|
||||
|
||||
async def read_raw_json_body(request: Request | None) -> bytes | None:
|
||||
if request is None or _safe_get_request_parsed_body(request=request) is None:
|
||||
return None
|
||||
content_type: Final = _safe_get_request_headers(request=request).get("content-type", "")
|
||||
if _is_form_content_type(content_type):
|
||||
return None
|
||||
try:
|
||||
return await request.body()
|
||||
except RuntimeError:
|
||||
return None
|
||||
|
||||
|
||||
def _safe_get_request_parsed_body(request: Request | None) -> dict | None:
|
||||
if request is None:
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -34,6 +34,7 @@ from litellm.proxy.common_utils.user_api_key_cache import (
|
|||
)
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging
|
||||
from litellm.router import Router
|
||||
from litellm.rust_bridge.token_counter import count_anthropic_input_tokens, uses_anthropic_tokenizer
|
||||
from litellm.types.proxy.model_access_group_budget import ModelAccessGroupBudget
|
||||
from litellm.types.router import DeploymentTypedDict
|
||||
|
||||
|
|
@ -210,6 +211,7 @@ async def reserve_budget_for_request(
|
|||
end_user_object: object = None,
|
||||
apply_user_budget_to_team_keys: bool = False,
|
||||
fail_closed_budget_enforcement: bool = False,
|
||||
raw_body: bytes | None = None,
|
||||
) -> dict | None:
|
||||
if valid_token is None or not RouteChecks.is_llm_api_route(route=route):
|
||||
return None
|
||||
|
|
@ -237,6 +239,7 @@ async def reserve_budget_for_request(
|
|||
request_body=request_body,
|
||||
route=route,
|
||||
llm_router=llm_router,
|
||||
raw_body=raw_body,
|
||||
)
|
||||
|
||||
current_spend_by_counter_key: Final[dict[str, float]] = {}
|
||||
|
|
@ -1355,24 +1358,46 @@ async def count_request_input_tokens(
|
|||
request_body: dict,
|
||||
route: str,
|
||||
llm_router: Router | None,
|
||||
raw_body: bytes | None = None,
|
||||
) -> Mapping[str, int]:
|
||||
"""Input-token count per candidate model, counted once per request.
|
||||
|
||||
Tokenizing is the reservation path's dominant CPU cost and is O(prompt), so
|
||||
counting a large prompt inline stalls every other request on the worker.
|
||||
Large prompts are counted in a worker thread, and the counts are reused by
|
||||
both the max-cost and the input-cost estimate.
|
||||
Models on the Anthropic tokenizer are counted from the raw body by the Rust
|
||||
bridge when it is enabled, which parses and tokenizes with the GIL released.
|
||||
Everything it declines is counted in Python, large prompts in a worker
|
||||
thread. The counts are reused by both the max-cost and the input-cost
|
||||
estimate.
|
||||
"""
|
||||
models: Final = _get_request_models(request_body=request_body, route=route, llm_router=llm_router)
|
||||
if not models:
|
||||
return MappingProxyType({})
|
||||
if _approximate_input_size(request_body) < TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS:
|
||||
return _count_input_tokens_for_models(request_body=request_body, models=models)
|
||||
return await asyncio.to_thread(
|
||||
_count_input_tokens_for_models,
|
||||
request_body=request_body,
|
||||
models=models,
|
||||
rust_count: Final = (
|
||||
await count_anthropic_input_tokens(raw_body)
|
||||
if raw_body is not None and any(uses_anthropic_tokenizer(model) for model in models)
|
||||
else None
|
||||
)
|
||||
rust_counts: Final = MappingProxyType(
|
||||
{
|
||||
model: rust_count.input_tokens
|
||||
for model in models
|
||||
if rust_count is not None and uses_anthropic_tokenizer(model)
|
||||
}
|
||||
)
|
||||
python_models: Final = tuple(model for model in models if model not in rust_counts)
|
||||
if not python_models:
|
||||
return rust_counts
|
||||
python_counts: Final = (
|
||||
_count_input_tokens_for_models(request_body=request_body, models=python_models)
|
||||
if _approximate_input_size(request_body) < TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS
|
||||
else await asyncio.to_thread(
|
||||
_count_input_tokens_for_models,
|
||||
request_body=request_body,
|
||||
models=python_models,
|
||||
)
|
||||
)
|
||||
return MappingProxyType({**rust_counts, **python_counts})
|
||||
|
||||
|
||||
def _count_input_tokens_for_models(
|
||||
|
|
|
|||
79
litellm/rust_bridge/token_counter.py
Normal file
79
litellm/rust_bridge/token_counter.py
Normal file
|
|
@ -0,0 +1,79 @@
|
|||
"""Thin Python wrapper for the native Rust input token counter."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable
|
||||
from dataclasses import dataclass
|
||||
from functools import lru_cache
|
||||
from typing import Final, Protocol, cast # noqa: TID251 # native extension exposes dynamically typed callables
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.rust_bridge.bindings import NativeBinding
|
||||
from litellm.rust_bridge.configuration import rust_enabled
|
||||
from litellm.rust_bridge.runtime import BridgeErrorContext, RustHandled, aattempt
|
||||
|
||||
|
||||
class RustTokenCounter(Protocol):
|
||||
def acount_request(self, body: bytes) -> Awaitable[object]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class RustTokenCounterFactory(Protocol):
|
||||
def __call__(self, tokenizer_json: str) -> RustTokenCounter:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class InputTokenCount:
|
||||
model: str | None
|
||||
input_tokens: int
|
||||
|
||||
|
||||
_INPUT_TOKEN_COUNT: Final = TypeAdapter(InputTokenCount)
|
||||
|
||||
|
||||
def _as_factory(value: object) -> RustTokenCounterFactory | None:
|
||||
return (
|
||||
cast( # cast-ok: native extension protocol is runtime-defined
|
||||
RustTokenCounterFactory, value
|
||||
)
|
||||
if callable(value)
|
||||
else None
|
||||
)
|
||||
|
||||
|
||||
TOKEN_COUNTER: Final = NativeBinding("TokenCounter", validate=_as_factory)
|
||||
|
||||
|
||||
def uses_anthropic_tokenizer(model: str) -> bool:
|
||||
if litellm.disable_token_counter is True or litellm.disable_hf_tokenizer_download is True:
|
||||
return False
|
||||
return model in litellm.anthropic_models and "claude-3" not in model
|
||||
|
||||
|
||||
@lru_cache(maxsize=4)
|
||||
def _anthropic_counter(factory: RustTokenCounterFactory) -> RustTokenCounter:
|
||||
from litellm.utils import claude_json_str
|
||||
|
||||
return factory(claude_json_str)
|
||||
|
||||
|
||||
async def count_anthropic_input_tokens(body: bytes) -> InputTokenCount | None:
|
||||
if not rust_enabled():
|
||||
return None
|
||||
factory: Final = TOKEN_COUNTER.load()
|
||||
if factory is None:
|
||||
return None
|
||||
try:
|
||||
attempt: Final = await aattempt(
|
||||
native_call=lambda: _anthropic_counter(factory).acount_request(body),
|
||||
adapt=_INPUT_TOKEN_COUNT.validate_python,
|
||||
context=BridgeErrorContext(route="token_counter", provider="anthropic", model=""),
|
||||
)
|
||||
except (RuntimeError, ValueError) as error:
|
||||
verbose_logger.debug("Rust token counter failed, counting in Python: %s", error)
|
||||
return None
|
||||
return attempt.value if isinstance(attempt, RustHandled) else None
|
||||
|
|
@ -26,9 +26,58 @@ from litellm.proxy.common_utils.http_parsing_utils import (
|
|||
get_tags_from_request_body,
|
||||
numeric_form_fields,
|
||||
populate_request_with_path_params,
|
||||
read_raw_json_body,
|
||||
)
|
||||
|
||||
|
||||
def _starlette_request(body: bytes, content_type: str) -> Request:
|
||||
scope = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/v1/messages",
|
||||
"headers": [(b"content-type", content_type.encode())],
|
||||
"query_string": b"",
|
||||
}
|
||||
chunks = iter((body,))
|
||||
|
||||
async def receive():
|
||||
return {"type": "http.request", "body": next(chunks, b""), "more_body": False}
|
||||
|
||||
return Request(scope, receive)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_read_raw_json_body_returns_the_bytes_the_parsed_body_came_from():
|
||||
body = b'{"model": "claude-sonnet-4-5", "messages": [{"role": "user", "content": "hi"}]}'
|
||||
request = _starlette_request(body, "application/json")
|
||||
|
||||
assert await _read_request_body(request) == orjson.loads(body)
|
||||
assert await read_raw_json_body(request) == body
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_read_raw_json_body_is_none_until_the_body_has_been_parsed():
|
||||
request = _starlette_request(b'{"model": "claude-sonnet-4-5"}', "application/json")
|
||||
|
||||
assert await read_raw_json_body(request) is None
|
||||
assert await read_raw_json_body(None) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_read_raw_json_body_is_none_for_form_bodies():
|
||||
request = _starlette_request(b"model=claude-sonnet-4-5", "application/x-www-form-urlencoded")
|
||||
|
||||
assert await _read_request_body(request) == {"model": "claude-sonnet-4-5"}
|
||||
assert await read_raw_json_body(request) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_read_raw_json_body_is_none_for_a_request_that_only_mocks_the_parsed_body_path():
|
||||
mock_request = MagicMock()
|
||||
|
||||
assert await read_raw_json_body(mock_request) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_request_body_caching():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,13 +1,21 @@
|
|||
import json
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
import litellm
|
||||
from litellm.caching import DualCache
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.spend_tracking.budget_reservation import estimate_request_max_cost, reserve_budget_for_request
|
||||
from litellm.proxy.spend_tracking.budget_reservation import (
|
||||
count_request_input_tokens,
|
||||
estimate_request_max_cost,
|
||||
reserve_budget_for_request,
|
||||
)
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.rust_bridge import bindings, configuration
|
||||
from litellm.rust_bridge import token_counter as rust_token_counter
|
||||
|
||||
TOKEN_COUNTING_ROUTES: Final = (
|
||||
"/responses/input_tokens",
|
||||
|
|
@ -139,3 +147,128 @@ def test_bedrock_converse_body_reserves_the_prompt_not_the_context_window():
|
|||
)
|
||||
assert converse_cost is not None and invoke_cost is not None
|
||||
assert invoke_cost < converse_cost < 2 * invoke_cost
|
||||
|
||||
|
||||
ANTHROPIC_TOKENIZER_MODEL: Final = "claude-sonnet-4-5-20250929"
|
||||
RUST_COUNTED_BODY: Final = {"model": ANTHROPIC_TOKENIZER_MODEL, "max_tokens": 16, "messages": ANTHROPIC_MESSAGES}
|
||||
RUST_INPUT_TOKENS: Final = 4_321
|
||||
|
||||
|
||||
class _FakeDeclined(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class _FakeUpstream(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class _FakeNative:
|
||||
RustBridgeDeclined = _FakeDeclined
|
||||
RustUpstreamError = _FakeUpstream
|
||||
|
||||
|
||||
class _RecordingCounter:
|
||||
bodies: Final[list[bytes]] = []
|
||||
|
||||
def __init__(self, tokenizer_json: str) -> None:
|
||||
pass
|
||||
|
||||
async def acount_request(self, body: bytes) -> object:
|
||||
self.bodies.append(body)
|
||||
return {"model": ANTHROPIC_TOKENIZER_MODEL, "input_tokens": RUST_INPUT_TOKENS}
|
||||
|
||||
|
||||
class _DecliningCounter:
|
||||
def __init__(self, tokenizer_json: str) -> None:
|
||||
pass
|
||||
|
||||
async def acount_request(self, body: bytes) -> object:
|
||||
raise _FakeDeclined("unsupported content block")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def rust_counter(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(bindings, "get_native_bridge", lambda: _FakeNative())
|
||||
rust_token_counter._anthropic_counter.cache_clear()
|
||||
configuration.reset_rust_configuration()
|
||||
_RecordingCounter.bodies.clear()
|
||||
yield
|
||||
rust_token_counter.TOKEN_COUNTER.reset()
|
||||
rust_token_counter._anthropic_counter.cache_clear()
|
||||
configuration.reset_rust_configuration()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("route", "request_body"),
|
||||
(
|
||||
("/v1/messages", RUST_COUNTED_BODY),
|
||||
("/v1/chat/completions", {"model": ANTHROPIC_TOKENIZER_MODEL, "messages": ANTHROPIC_MESSAGES}),
|
||||
("/v1/completions", {"model": ANTHROPIC_TOKENIZER_MODEL, "prompt": "hi"}),
|
||||
("/v1/responses", {"model": ANTHROPIC_TOKENIZER_MODEL, "input": "hi"}),
|
||||
("/v1/embeddings", {"model": ANTHROPIC_TOKENIZER_MODEL, "input": ["hi"]}),
|
||||
("/v1/rerank", {"model": ANTHROPIC_TOKENIZER_MODEL, "query": "hi", "documents": ["a"]}),
|
||||
),
|
||||
)
|
||||
async def test_rust_count_replaces_python_tokenizing_on_every_llm_route(
|
||||
rust_counter: None, route: str, request_body: dict
|
||||
) -> None:
|
||||
litellm.rust(True)
|
||||
rust_token_counter.TOKEN_COUNTER.override(_RecordingCounter)
|
||||
raw_body: Final = json.dumps(request_body).encode()
|
||||
|
||||
counts: Final = await count_request_input_tokens(
|
||||
request_body=request_body, route=route, llm_router=None, raw_body=raw_body
|
||||
)
|
||||
|
||||
assert dict(counts) == {ANTHROPIC_TOKENIZER_MODEL: RUST_INPUT_TOKENS}
|
||||
assert _RecordingCounter.bodies == [raw_body]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rust_decline_falls_back_to_python_count(rust_counter: None) -> None:
|
||||
litellm.rust(True)
|
||||
rust_token_counter.TOKEN_COUNTER.override(_DecliningCounter)
|
||||
python_counts: Final = await count_request_input_tokens(
|
||||
request_body=RUST_COUNTED_BODY, route="/v1/messages", llm_router=None
|
||||
)
|
||||
|
||||
counts: Final = await count_request_input_tokens(
|
||||
request_body=RUST_COUNTED_BODY,
|
||||
route="/v1/messages",
|
||||
llm_router=None,
|
||||
raw_body=json.dumps(RUST_COUNTED_BODY).encode(),
|
||||
)
|
||||
|
||||
assert dict(counts) == dict(python_counts)
|
||||
assert counts[ANTHROPIC_TOKENIZER_MODEL] != RUST_INPUT_TOKENS
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disabled_rust_never_sees_the_raw_body(rust_counter: None) -> None:
|
||||
litellm.rust(False)
|
||||
rust_token_counter.TOKEN_COUNTER.override(_RecordingCounter)
|
||||
|
||||
counts: Final = await count_request_input_tokens(
|
||||
request_body=RUST_COUNTED_BODY,
|
||||
route="/v1/messages",
|
||||
llm_router=None,
|
||||
raw_body=json.dumps(RUST_COUNTED_BODY).encode(),
|
||||
)
|
||||
|
||||
assert _RecordingCounter.bodies == []
|
||||
assert counts[ANTHROPIC_TOKENIZER_MODEL] != RUST_INPUT_TOKENS
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_anthropic_tokenizer_models_stay_in_python(rust_counter: None) -> None:
|
||||
litellm.rust(True)
|
||||
rust_token_counter.TOKEN_COUNTER.override(_RecordingCounter)
|
||||
body: Final = {"model": "gpt-4o", "messages": ANTHROPIC_MESSAGES}
|
||||
|
||||
counts: Final = await count_request_input_tokens(
|
||||
request_body=body, route="/v1/chat/completions", llm_router=None, raw_body=json.dumps(body).encode()
|
||||
)
|
||||
|
||||
assert _RecordingCounter.bodies == []
|
||||
assert counts["gpt-4o"] != RUST_INPUT_TOKENS
|
||||
|
|
|
|||
242
tests/test_litellm/rust_bridge/test_token_counter.py
Normal file
242
tests/test_litellm/rust_bridge/test_token_counter.py
Normal file
|
|
@ -0,0 +1,242 @@
|
|||
"""Tests for the Rust input token counter bridge.
|
||||
|
||||
The native factory is dependency-injected through ``TOKEN_COUNTER.override``
|
||||
so the fallback cases run without the compiled extension present. The parity
|
||||
cases need the extension and are skipped when it is not built.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.proxy.spend_tracking.budget_reservation import _count_input_tokens
|
||||
from litellm.rust_bridge import bindings, configuration
|
||||
from litellm.rust_bridge import token_counter as bridge
|
||||
|
||||
MODEL: Final = "claude-sonnet-4-5-20250929"
|
||||
BODY: Final = json.dumps({"model": MODEL, "messages": [{"role": "user", "content": "hello"}]}).encode()
|
||||
|
||||
|
||||
class _FakeDeclined(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class _FakeUpstream(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class _FakeNative:
|
||||
RustBridgeDeclined = _FakeDeclined
|
||||
RustUpstreamError = _FakeUpstream
|
||||
|
||||
|
||||
class _RecordingCounter:
|
||||
def __init__(self, tokenizer_json: str) -> None:
|
||||
self.tokenizer_json = tokenizer_json
|
||||
self.bodies: list[bytes] = []
|
||||
|
||||
async def acount_request(self, body: bytes) -> object:
|
||||
self.bodies.append(body)
|
||||
return {"model": MODEL, "input_tokens": 42}
|
||||
|
||||
|
||||
class _DecliningCounter:
|
||||
def __init__(self, tokenizer_json: str) -> None:
|
||||
pass
|
||||
|
||||
async def acount_request(self, body: bytes) -> object:
|
||||
raise _FakeDeclined("request has no messages")
|
||||
|
||||
|
||||
class _FailingCounter:
|
||||
def __init__(self, tokenizer_json: str) -> None:
|
||||
pass
|
||||
|
||||
async def acount_request(self, body: bytes) -> object:
|
||||
raise RuntimeError("encode failed")
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_bridge(monkeypatch: pytest.MonkeyPatch):
|
||||
bridge.TOKEN_COUNTER.reset()
|
||||
bridge._anthropic_counter.cache_clear()
|
||||
configuration.reset_rust_configuration()
|
||||
monkeypatch.setattr(bindings, "get_native_bridge", lambda: _FakeNative())
|
||||
yield
|
||||
bridge.TOKEN_COUNTER.reset()
|
||||
bridge._anthropic_counter.cache_clear()
|
||||
configuration.reset_rust_configuration()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disabled_bridge_never_constructs_a_counter() -> None:
|
||||
constructed: list[str] = []
|
||||
|
||||
def factory(tokenizer_json: str) -> _RecordingCounter:
|
||||
constructed.append(tokenizer_json)
|
||||
return _RecordingCounter(tokenizer_json)
|
||||
|
||||
litellm.rust(False)
|
||||
bridge.TOKEN_COUNTER.override(factory)
|
||||
|
||||
assert await bridge.count_anthropic_input_tokens(BODY) is None
|
||||
assert constructed == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_enabled_bridge_returns_typed_count_and_reuses_one_counter() -> None:
|
||||
counters: list[_RecordingCounter] = []
|
||||
|
||||
def factory(tokenizer_json: str) -> _RecordingCounter:
|
||||
counter = _RecordingCounter(tokenizer_json)
|
||||
counters.append(counter)
|
||||
return counter
|
||||
|
||||
litellm.rust(True)
|
||||
bridge.TOKEN_COUNTER.override(factory)
|
||||
|
||||
first: Final = await bridge.count_anthropic_input_tokens(BODY)
|
||||
second: Final = await bridge.count_anthropic_input_tokens(BODY)
|
||||
|
||||
assert first == bridge.InputTokenCount(model=MODEL, input_tokens=42)
|
||||
assert second == first
|
||||
assert len(counters) == 1
|
||||
assert counters[0].bodies == [BODY, BODY]
|
||||
assert json.loads(counters[0].tokenizer_json)["model"]["type"] == "BPE"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_native_module_falls_back(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
litellm.rust(True)
|
||||
monkeypatch.setattr(bindings, "get_native_bridge", lambda: None)
|
||||
|
||||
assert await bridge.count_anthropic_input_tokens(BODY) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_declined_request_falls_back() -> None:
|
||||
litellm.rust(True)
|
||||
bridge.TOKEN_COUNTER.override(_DecliningCounter)
|
||||
|
||||
assert await bridge.count_anthropic_input_tokens(BODY) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runtime_failure_falls_back() -> None:
|
||||
litellm.rust(True)
|
||||
bridge.TOKEN_COUNTER.override(_FailingCounter)
|
||||
|
||||
assert await bridge.count_anthropic_input_tokens(BODY) is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("model", "expected"),
|
||||
((MODEL, True), ("claude-3-5-sonnet-20241022", False), ("gpt-4o", False), ("my-router-alias", False)),
|
||||
)
|
||||
def test_uses_anthropic_tokenizer_mirrors_python_tokenizer_selection(model: str, expected: bool) -> None:
|
||||
assert bridge.uses_anthropic_tokenizer(model) is expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize("flag", ("disable_hf_tokenizer_download", "disable_token_counter"))
|
||||
def test_uses_anthropic_tokenizer_respects_python_opt_outs(monkeypatch: pytest.MonkeyPatch, flag: str) -> None:
|
||||
monkeypatch.setattr(litellm, flag, True)
|
||||
|
||||
assert bridge.uses_anthropic_tokenizer(MODEL) is False
|
||||
|
||||
|
||||
PARITY_REQUESTS: Final[tuple[dict[str, object], ...]] = (
|
||||
{"model": MODEL, "messages": [{"role": "user", "content": "Hello, how are you today?"}]},
|
||||
{
|
||||
"model": MODEL,
|
||||
"messages": [
|
||||
{"role": "system", "content": "You are terse."},
|
||||
{"role": "user", "name": "bob", "content": [{"type": "text", "text": "Summarize this."}]},
|
||||
{"role": "assistant", "content": "Sure."},
|
||||
],
|
||||
},
|
||||
{
|
||||
"model": MODEL,
|
||||
"messages": [{"role": "user", "content": "weather in sf?"}],
|
||||
"tools": [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"description": "Get weather",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"city": {"type": "string", "description": "City"},
|
||||
"unit": {"type": "string", "enum": ["c", "f"]},
|
||||
},
|
||||
"required": ["city"],
|
||||
},
|
||||
},
|
||||
}
|
||||
],
|
||||
"tool_choice": {"type": "function", "function": {"name": "get_weather"}},
|
||||
},
|
||||
{
|
||||
"model": MODEL,
|
||||
"messages": [{"role": "user", "content": "x " * 20_000}],
|
||||
},
|
||||
{"model": MODEL, "prompt": "Write a haiku about ships.", "max_tokens": 20},
|
||||
{"model": MODEL, "prompt": ["first prompt", "second prompt"]},
|
||||
{
|
||||
"model": MODEL,
|
||||
"instructions": "be terse",
|
||||
"input": [
|
||||
{"role": "user", "content": [{"type": "input_text", "text": "Summarise caf\u00e9 menus \u2014 \"ok\"?\n"}]},
|
||||
{"role": "assistant", "content": "Sure."},
|
||||
],
|
||||
},
|
||||
{"model": MODEL, "input": "a single embedding string"},
|
||||
{"model": MODEL, "input": [[101, 2023, 5], [7]], "encoding_format": "float"},
|
||||
{"model": MODEL, "query": "best harbour", "documents": ["doc one", {"text": "doc two", "title": "T", "n": 3}]},
|
||||
{"model": MODEL, "messages": None, "prompt": "messages key wins even when null"},
|
||||
{"prompt": "model comes from the route"},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("request_body", PARITY_REQUESTS)
|
||||
async def test_native_count_matches_python_budget_counter(
|
||||
monkeypatch: pytest.MonkeyPatch, request_body: dict[str, object]
|
||||
) -> None:
|
||||
native: Final = pytest.importorskip("litellm.rust_bridge._native")
|
||||
monkeypatch.setattr(bindings, "get_native_bridge", lambda: native)
|
||||
litellm.rust(True)
|
||||
|
||||
rust_count: Final = await bridge.count_anthropic_input_tokens(json.dumps(request_body).encode())
|
||||
python_count: Final = _count_input_tokens(request_body=request_body, model=MODEL)
|
||||
|
||||
assert rust_count is not None
|
||||
assert rust_count.model == request_body.get("model")
|
||||
assert rust_count.input_tokens == python_count
|
||||
|
||||
|
||||
DECLINED_REQUESTS: Final[tuple[dict[str, object], ...]] = (
|
||||
{
|
||||
"model": MODEL,
|
||||
"messages": [{"role": "user", "content": [{"type": "image_url", "image_url": {"url": "data:image/png;base64,AA"}}]}],
|
||||
},
|
||||
{"model": MODEL, "prompt": 1.5},
|
||||
{"model": MODEL, "documents": [{"score": 0.5}]},
|
||||
{"model": MODEL, "file": "audio.mp3"},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("request_body", DECLINED_REQUESTS)
|
||||
async def test_native_declines_shapes_python_prices_differently(
|
||||
monkeypatch: pytest.MonkeyPatch, request_body: dict[str, object]
|
||||
) -> None:
|
||||
native: Final = pytest.importorskip("litellm.rust_bridge._native")
|
||||
monkeypatch.setattr(bindings, "get_native_bridge", lambda: native)
|
||||
litellm.rust(True)
|
||||
|
||||
assert await bridge.count_anthropic_input_tokens(json.dumps(request_body).encode()) is None
|
||||
Loading…
Add table
Reference in a new issue