From 85cc7173fa4f72ea0f197b19031431630a4d6eb0 Mon Sep 17 00:00:00 2001 From: yassin Date: Wed, 9 Sep 2026 06:06:04 +0000 Subject: [PATCH] feat(ai-gateway): rust admission layer for /v1/messages (parse once, parallel bounded checks) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm-rust/Cargo.lock | 405 +++++++++++++++++- litellm-rust/crates/ai-gateway/Cargo.toml | 6 +- .../ai-gateway/src/admission/extract.rs | 49 +++ .../ai-gateway/src/admission/identity.rs | 221 ++++++++++ .../crates/ai-gateway/src/admission/limits.rs | 171 ++++++++ .../crates/ai-gateway/src/admission/mod.rs | 389 +++++++++++++++++ .../ai-gateway/src/admission/tokenizer.rs | 111 +++++ .../crates/ai-gateway/src/auth/mod.rs | 19 +- .../crates/ai-gateway/src/constants.rs | 36 ++ litellm-rust/crates/ai-gateway/src/lib.rs | 2 + litellm-rust/crates/ai-gateway/src/main.rs | 15 + .../ai-gateway/src/routes/messages/mod.rs | 89 +++- .../ai-gateway/src/routes/responses/mod.rs | 8 + litellm-rust/crates/ai-gateway/src/state.rs | 5 +- .../crates/ai-gateway/src/trace_parity.rs | 8 + 15 files changed, 1505 insertions(+), 29 deletions(-) create mode 100644 litellm-rust/crates/ai-gateway/src/admission/extract.rs create mode 100644 litellm-rust/crates/ai-gateway/src/admission/identity.rs create mode 100644 litellm-rust/crates/ai-gateway/src/admission/limits.rs create mode 100644 litellm-rust/crates/ai-gateway/src/admission/mod.rs create mode 100644 litellm-rust/crates/ai-gateway/src/admission/tokenizer.rs diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 43b9ec1aac2..6578d873faa 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -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 = "1.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6f55d7153ba3b507595872a3874803f07a8a81d1e888abed8e5db7da0597d6e2" + +[[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" @@ -1365,6 +1514,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 +1567,7 @@ name = "litellm-ai-gateway" version = "0.1.0" dependencies = [ "axum", - "base64", + "base64 0.22.1", "futures-channel", "futures-util", "litellm-config", @@ -1421,6 +1579,8 @@ dependencies = [ "serde_json", "sha2 0.10.9", "subtle", + "thiserror 2.0.19", + "tokenizers", "tokio", "tokio-tungstenite", "tower", @@ -1447,7 +1607,7 @@ dependencies = [ "aws-sigv4", "aws-smithy-runtime-api", "aws-types", - "base64", + "base64 0.22.1", "rand 0.8.7", "reqwest", "rstest", @@ -1507,6 +1667,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 +1711,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 +1728,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 +1790,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 +1840,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 +2098,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 +2117,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 +2152,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 +2171,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 +2205,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 +2273,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 +2655,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 +2853,39 @@ version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" +[[package]] +name = "tokenizers" +version = "0.23.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "44e5bea67576e04b6ff8564c5d9e09c2ef0cf476502245f2f120e497769d3112" +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 +3124,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 +3228,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 +3462,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" diff --git a/litellm-rust/crates/ai-gateway/Cargo.toml b/litellm-rust/crates/ai-gateway/Cargo.toml index 74cf66e88a2..6df0a534bdb 100644 --- a/litellm-rust/crates/ai-gateway/Cargo.toml +++ b/litellm-rust/crates/ai-gateway/Cargo.toml @@ -32,15 +32,19 @@ serde_json.workspace = true base64.workspace = true axum = { workspace = true, features = ["ws"], optional = true } serde.workspace = true +thiserror.workspace = true subtle = { workspace = true, optional = true } # sha2 hashes the master key into user_api_key_hash (matches the proxy's # SHA-256 hash_token) so the plaintext credential never enters a log payload. sha2 = { workspace = true, optional = true } tower = { version = "0.5.3", features = ["util"], optional = true } +# HuggingFace tokenizer for the admission layer's input token count; without the +# default features it pulls no HTTP client or progress bars, only the `onig` regex. +tokenizers = { version = "0.23.1", default-features = false, features = ["onig"], optional = true } [features] default = [] -server = ["dep:axum", "dep:subtle", "dep:sha2"] +server = ["dep:axum", "dep:subtle", "dep:sha2", "dep:tokenizers"] # Build the gateway's config from the proxy YAML via an embedded Python # interpreter (links libpython; requires `litellm` importable at runtime). python-config = ["litellm-config/python"] diff --git a/litellm-rust/crates/ai-gateway/src/admission/extract.rs b/litellm-rust/crates/ai-gateway/src/admission/extract.rs new file mode 100644 index 00000000000..859024a5f54 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/admission/extract.rs @@ -0,0 +1,49 @@ +//! Axum extractor running [`Admission`] on the raw request body. + +use axum::body::to_bytes; +use axum::extract::{FromRequest, Request}; +use axum::http::StatusCode; +use axum::response::{IntoResponse, Response}; + +use crate::auth::bearer_token; +use crate::state::AppState; + +use super::{Admission, Admitted, Rejection}; + +/// Handler argument that yields the admitted request, or the rejection response. +pub struct Admit(pub Admitted); + +#[axum::async_trait] +impl FromRequest for Admit { + type Rejection = Response; + + async fn from_request(request: Request, state: &AppState) -> Result { + let (parts, body) = request.into_parts(); + let admission: &Admission = &state.admission; + let raw = to_bytes(body, admission.max_request_bytes()) + .await + .map_err(|error| (StatusCode::PAYLOAD_TOO_LARGE, error.to_string()).into_response())?; + admission + .admit(bearer_token(&parts.headers), &raw) + .await + .map(Admit) + .map_err(|rejection| reject(&rejection)) + } +} + +fn reject(rejection: &Rejection) -> Response { + let status = match rejection { + Rejection::Unauthorized => StatusCode::UNAUTHORIZED, + Rejection::IdentityUnavailable(_) => StatusCode::SERVICE_UNAVAILABLE, + Rejection::ModelNotAllowed(_) => StatusCode::FORBIDDEN, + Rejection::RequestTooLarge { .. } => StatusCode::PAYLOAD_TOO_LARGE, + Rejection::ContextTooLarge { .. } | Rejection::InvalidRequest(_) => StatusCode::BAD_REQUEST, + Rejection::LimitExceeded(_) => StatusCode::TOO_MANY_REQUESTS, + Rejection::Tokenizer(_) => StatusCode::INTERNAL_SERVER_ERROR, + }; + ( + status, + axum::Json(serde_json::json!({"error": {"message": rejection.to_string()}})), + ) + .into_response() +} diff --git a/litellm-rust/crates/ai-gateway/src/admission/identity.rs b/litellm-rust/crates/ai-gateway/src/admission/identity.rs new file mode 100644 index 00000000000..a570870d89c --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/admission/identity.rs @@ -0,0 +1,221 @@ +//! Who is calling: the master key, or a virtual key resolved through the Python proxy. +//! +//! Virtual keys are cached by their SHA-256 hash for the life of the process, so the network +//! round trip to `/key/info` happens once per key, not once per request. + +use std::collections::HashMap; +use std::sync::{Arc, RwLock}; +use std::time::Duration; + +use serde::Deserialize; +use subtle::ConstantTimeEq; + +use crate::auth::hash_token; +use crate::constants::{KEY_INFO_TIMEOUT_SECS, PROXY_KEY_INFO_PATH}; + +/// Limits attached to a virtual key. `None` means unlimited, as in the proxy. +#[derive(Clone, Debug, Default, Deserialize, PartialEq)] +pub struct KeyLimits { + #[serde(default)] + pub models: Vec, + #[serde(default)] + pub max_budget: Option, + #[serde(default)] + pub spend: f64, + #[serde(default)] + pub tpm_limit: Option, + #[serde(default)] + pub rpm_limit: Option, +} + +impl KeyLimits { + pub fn allows_model(&self, model: &str) -> bool { + self.models.is_empty() || self.models.iter().any(|allowed| allowed == model) + } +} + +#[derive(Clone, Debug, PartialEq)] +pub enum Identity { + Master, + VirtualKey { + key_hash: String, + limits: Arc, + }, +} + +impl Identity { + pub fn key_hash(&self) -> &str { + match self { + Identity::Master => "litellm_proxy_master_key", + Identity::VirtualKey { key_hash, .. } => key_hash, + } + } +} + +#[derive(Debug, thiserror::Error, PartialEq)] +pub enum IdentityError { + #[error("missing or invalid bearer token")] + Unauthorized, + #[error("key lookup failed: {0}")] + LookupFailed(String), +} + +#[derive(Deserialize)] +struct KeyInfoResponse { + info: KeyLimits, +} + +/// Resolves bearer tokens; misses go to the proxy's `/key/info`. +pub struct IdentityCache { + master_key: Option>, + proxy_base_url: String, + http: reqwest::Client, + cache: RwLock>>, +} + +impl IdentityCache { + pub fn new(master_key: Option>, proxy_base_url: String) -> Self { + Self { + master_key, + proxy_base_url: proxy_base_url.trim_end_matches('/').to_string(), + http: reqwest::Client::builder() + .timeout(Duration::from_secs(KEY_INFO_TIMEOUT_SECS)) + .build() + .unwrap_or_default(), + cache: RwLock::new(HashMap::new()), + } + } + + /// Seed the cache, so tests and offline hosts never call the proxy. + pub fn insert(&self, token: &str, limits: KeyLimits) { + if let Ok(mut cache) = self.cache.write() { + cache.insert(hash_token(token), Arc::new(limits)); + } + } + + pub async fn resolve(&self, token: Option<&str>) -> Result { + let Some(token) = token.map(str::trim).filter(|token| !token.is_empty()) else { + return Err(IdentityError::Unauthorized); + }; + if let Some(master) = self.master_key.as_deref() + && bool::from(token.as_bytes().ct_eq(master.as_bytes())) + { + return Ok(Identity::Master); + } + let key_hash = hash_token(token); + let cached = self + .cache + .read() + .ok() + .and_then(|cache| cache.get(&key_hash).cloned()); + let limits = match cached { + Some(limits) => limits, + None => { + let limits = Arc::new(self.fetch(token).await?); + if let Ok(mut cache) = self.cache.write() { + cache.insert(key_hash.clone(), Arc::clone(&limits)); + } + limits + } + }; + Ok(Identity::VirtualKey { key_hash, limits }) + } + + async fn fetch(&self, token: &str) -> Result { + let Some(master) = self.master_key.as_deref() else { + return Err(IdentityError::Unauthorized); + }; + let response = self + .http + .get(format!("{}{PROXY_KEY_INFO_PATH}", self.proxy_base_url)) + .query(&[("key", token)]) + .bearer_auth(master) + .send() + .await + .map_err(|error| IdentityError::LookupFailed(error.without_url().to_string()))?; + match response.status().as_u16() { + 200 => response + .json::() + .await + .map(|body| body.info) + .map_err(|error| IdentityError::LookupFailed(error.without_url().to_string())), + 400..=404 => Err(IdentityError::Unauthorized), + status => Err(IdentityError::LookupFailed(format!( + "proxy answered {status}" + ))), + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn cache() -> IdentityCache { + IdentityCache::new( + Some(Arc::from("sk-master")), + "http://127.0.0.1:1".to_string(), + ) + } + + #[tokio::test] + async fn master_key_is_unlimited_and_never_looked_up() { + assert_eq!( + cache().resolve(Some("sk-master")).await, + Ok(Identity::Master) + ); + } + + #[tokio::test] + async fn missing_token_is_unauthorized() { + assert_eq!( + cache().resolve(None).await, + Err(IdentityError::Unauthorized) + ); + assert_eq!( + cache().resolve(Some(" ")).await, + Err(IdentityError::Unauthorized) + ); + } + + #[tokio::test] + async fn seeded_virtual_key_resolves_from_cache_without_network() { + let cache = cache(); + let limits = KeyLimits { + models: vec!["claude".to_string()], + max_budget: Some(10.0), + spend: 1.5, + tpm_limit: Some(1000), + rpm_limit: None, + }; + cache.insert("sk-virtual", limits.clone()); + let identity = cache.resolve(Some("sk-virtual")).await.unwrap(); + match identity { + Identity::VirtualKey { + key_hash, + limits: resolved, + } => { + assert_eq!(key_hash, hash_token("sk-virtual")); + assert_eq!(*resolved, limits); + } + Identity::Master => panic!("virtual key resolved as master"), + } + } + + #[tokio::test] + async fn unknown_key_lookup_failure_is_reported_not_admitted() { + let error = cache().resolve(Some("sk-unknown")).await.unwrap_err(); + assert!(matches!(error, IdentityError::LookupFailed(_)), "{error:?}"); + } + + #[test] + fn empty_model_list_allows_every_model() { + assert!(KeyLimits::default().allows_model("anything")); + let limits = KeyLimits { + models: vec!["a".to_string()], + ..KeyLimits::default() + }; + assert!(limits.allows_model("a")); + assert!(!limits.allows_model("b")); + } +} diff --git a/litellm-rust/crates/ai-gateway/src/admission/limits.rs b/litellm-rust/crates/ai-gateway/src/admission/limits.rs new file mode 100644 index 00000000000..34a61a7364b --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/admission/limits.rs @@ -0,0 +1,171 @@ +//! Per-key budget and per-minute request/token windows, checked and reserved under one lock. +//! +//! Process-local: in a multi-replica deployment these counters would live in Redis, as the +//! proxy's do. The check itself is a couple of integer compares per request. + +use std::collections::HashMap; +use std::sync::Mutex; +use std::time::{Duration, Instant}; + +use crate::constants::DEFAULT_INPUT_COST_PER_TOKEN; + +use super::identity::{Identity, KeyLimits}; + +const WINDOW: Duration = Duration::from_secs(60); + +#[derive(Debug, PartialEq, Eq, Clone, Copy)] +pub enum LimitExceeded { + Budget, + TokensPerMinute, + RequestsPerMinute, +} + +#[derive(Debug)] +struct KeyWindow { + started: Instant, + tokens: u64, + requests: u64, + reserved_spend: f64, +} + +#[derive(Default)] +pub struct Limits { + windows: Mutex>, +} + +impl Limits { + /// Admit `input_tokens` for the identity, or say which limit it would cross. + pub fn reserve(&self, identity: &Identity, input_tokens: usize) -> Result<(), LimitExceeded> { + let Identity::VirtualKey { key_hash, limits } = identity else { + return Ok(()); + }; + let Ok(mut windows) = self.windows.lock() else { + return Ok(()); + }; + let now = Instant::now(); + let window = windows.entry(key_hash.clone()).or_insert(KeyWindow { + started: now, + tokens: 0, + requests: 0, + reserved_spend: 0.0, + }); + if now.duration_since(window.started) >= WINDOW { + window.started = now; + window.tokens = 0; + window.requests = 0; + } + let tokens = input_tokens as u64; + let cost = input_tokens as f64 * DEFAULT_INPUT_COST_PER_TOKEN; + check(limits, window, tokens, cost)?; + window.tokens += tokens; + window.requests += 1; + window.reserved_spend += cost; + Ok(()) + } +} + +fn check( + limits: &KeyLimits, + window: &KeyWindow, + tokens: u64, + cost: f64, +) -> Result<(), LimitExceeded> { + if let Some(max_budget) = limits.max_budget + && limits.spend + window.reserved_spend + cost > max_budget + { + return Err(LimitExceeded::Budget); + } + if let Some(tpm) = limits.tpm_limit + && window.tokens + tokens > tpm + { + return Err(LimitExceeded::TokensPerMinute); + } + if let Some(rpm) = limits.rpm_limit + && window.requests + 1 > rpm + { + return Err(LimitExceeded::RequestsPerMinute); + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use super::*; + + fn key(limits: KeyLimits) -> Identity { + Identity::VirtualKey { + key_hash: "hash".to_string(), + limits: Arc::new(limits), + } + } + + #[test] + fn master_key_is_never_limited() { + let limits = Limits::default(); + assert_eq!(limits.reserve(&Identity::Master, usize::MAX), Ok(())); + } + + #[test] + fn tpm_window_accumulates_and_rejects_on_overflow() { + let limits = Limits::default(); + let identity = key(KeyLimits { + tpm_limit: Some(100), + ..KeyLimits::default() + }); + assert_eq!(limits.reserve(&identity, 60), Ok(())); + assert_eq!( + limits.reserve(&identity, 50), + Err(LimitExceeded::TokensPerMinute) + ); + assert_eq!(limits.reserve(&identity, 40), Ok(())); + } + + #[test] + fn rpm_counts_requests() { + let limits = Limits::default(); + let identity = key(KeyLimits { + rpm_limit: Some(2), + ..KeyLimits::default() + }); + assert_eq!(limits.reserve(&identity, 1), Ok(())); + assert_eq!(limits.reserve(&identity, 1), Ok(())); + assert_eq!( + limits.reserve(&identity, 1), + Err(LimitExceeded::RequestsPerMinute) + ); + } + + #[test] + fn budget_includes_prior_spend_and_local_reservations() { + let limits = Limits::default(); + let identity = key(KeyLimits { + max_budget: Some(1.0), + spend: 0.5, + ..KeyLimits::default() + }); + let tokens_for_quarter_dollar = (0.25 / DEFAULT_INPUT_COST_PER_TOKEN) as usize; + assert_eq!(limits.reserve(&identity, tokens_for_quarter_dollar), Ok(())); + assert_eq!(limits.reserve(&identity, tokens_for_quarter_dollar), Ok(())); + assert_eq!( + limits.reserve(&identity, tokens_for_quarter_dollar), + Err(LimitExceeded::Budget) + ); + } + + #[test] + fn a_rejected_request_reserves_nothing() { + let limits = Limits::default(); + let identity = key(KeyLimits { + tpm_limit: Some(10), + rpm_limit: Some(5), + ..KeyLimits::default() + }); + assert_eq!( + limits.reserve(&identity, 11), + Err(LimitExceeded::TokensPerMinute) + ); + assert_eq!(limits.reserve(&identity, 10), Ok(())); + } +} diff --git a/litellm-rust/crates/ai-gateway/src/admission/mod.rs b/litellm-rust/crates/ai-gateway/src/admission/mod.rs new file mode 100644 index 00000000000..9fb785292ea --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/admission/mod.rs @@ -0,0 +1,389 @@ +//! Admission for `/v1/messages`: the checks the Python proxy runs before a request reaches +//! the provider (identity, model access, size, token count, budget, rate limits). +//! +//! The body is parsed once from raw bytes. Identity resolution and tokenization run +//! concurrently, tokenization on the blocking pool behind a semaphore, and the budget and +//! per-minute checks reuse that single token count. Everything after the first request for a +//! key is an in-memory read. +//! +//! Proof of concept, not at parity with `user_api_key_auth` and the proxy hooks: identity +//! comes from the proxy's `/key/info` once per key, and budgets and TPM/RPM windows are +//! process-local. + +pub mod extract; +pub mod identity; +pub mod limits; +pub mod tokenizer; + +use std::path::Path; +use std::sync::Arc; +use std::time::Instant; + +use serde_json::Value; + +use crate::constants::{ + DEFAULT_MAX_INPUT_TOKENS, DEFAULT_MAX_REQUEST_BYTES, DEFAULT_PROXY_BASE_URL, + DEFAULT_TOKENIZER_CONCURRENCY, +}; + +pub use extract::Admit; +pub use identity::{Identity, IdentityCache, IdentityError, KeyLimits}; +pub use limits::{LimitExceeded, Limits}; +pub use tokenizer::{TokenCounter, TokenizerError}; + +#[derive(Debug, thiserror::Error)] +pub enum Rejection { + #[error("missing or invalid bearer token")] + Unauthorized, + #[error("key lookup failed: {0}")] + IdentityUnavailable(String), + #[error("key is not allowed to call model '{0}'")] + ModelNotAllowed(String), + #[error("request body of {bytes} bytes exceeds the {max} byte limit")] + RequestTooLarge { bytes: usize, max: usize }, + #[error("input of {tokens} tokens exceeds the {max} token limit")] + ContextTooLarge { tokens: usize, max: usize }, + #[error("{0:?} limit exceeded")] + LimitExceeded(LimitExceeded), + #[error("invalid request: {0}")] + InvalidRequest(String), + #[error("tokenization failed: {0}")] + Tokenizer(String), +} + +impl From for Rejection { + fn from(error: IdentityError) -> Self { + match error { + IdentityError::Unauthorized => Rejection::Unauthorized, + IdentityError::LookupFailed(reason) => Rejection::IdentityUnavailable(reason), + } + } +} + +/// The request after the single parse: what admission needs plus the body to forward. +#[derive(Debug)] +pub struct ParsedRequest { + pub body: Value, + pub model: String, + /// Everything the tokenizer sees: system prompt, message content, tool schemas. + pub text: String, +} + +impl ParsedRequest { + pub fn parse(raw: &[u8]) -> Result { + let body: Value = serde_json::from_slice(raw) + .map_err(|error| Rejection::InvalidRequest(format!("body is not JSON: {error}")))?; + let Some(object) = body.as_object() else { + return Err(Rejection::InvalidRequest( + "body must be a JSON object".to_string(), + )); + }; + let model = object + .get("model") + .and_then(Value::as_str) + .map(str::trim) + .filter(|model| !model.is_empty()) + .ok_or_else(|| Rejection::InvalidRequest("body requires a model".to_string()))? + .to_string(); + let mut text = String::with_capacity(raw.len()); + if let Some(system) = object.get("system") { + push_content(system, &mut text); + } + for message in object + .get("messages") + .and_then(Value::as_array) + .into_iter() + .flatten() + { + if let Some(content) = message.get("content") { + push_content(content, &mut text); + } + } + if let Some(tools) = object.get("tools") { + text.push_str(&tools.to_string()); + } + Ok(Self { body, model, text }) + } +} + +fn push_content(content: &Value, text: &mut String) { + match content { + Value::String(value) => { + text.push_str(value); + text.push('\n'); + } + Value::Array(blocks) => { + for block in blocks { + match block.get("text").and_then(Value::as_str) { + Some(value) => { + text.push_str(value); + text.push('\n'); + } + None => text.push_str(&block.to_string()), + } + } + } + other => text.push_str(&other.to_string()), + } +} + +fn env_usize(name: &str, default: usize) -> usize { + std::env::var(name) + .ok() + .and_then(|value| value.trim().parse().ok()) + .unwrap_or(default) +} + +#[derive(Debug)] +pub struct Admitted { + pub body: Value, + pub identity: Identity, + pub input_tokens: usize, + pub elapsed_ms: f64, +} + +pub struct Admission { + identities: IdentityCache, + limits: Limits, + tokens: TokenCounter, + max_request_bytes: usize, + max_input_tokens: usize, +} + +impl Admission { + pub fn new(identities: IdentityCache, tokens: TokenCounter) -> Self { + Self { + identities, + limits: Limits::default(), + tokens, + max_request_bytes: DEFAULT_MAX_REQUEST_BYTES, + max_input_tokens: DEFAULT_MAX_INPUT_TOKENS, + } + } + + /// Build from `LITELLM_PROXY_BASE_URL`, `LITELLM_ANTHROPIC_TOKENIZER_PATH`, + /// `LITELLM_TOKENIZER_CONCURRENCY`, `LITELLM_MAX_REQUEST_BYTES` and `LITELLM_MAX_INPUT_TOKENS`. + /// Without a tokenizer path the token count is approximated from the input length. + pub fn from_env(master_key: Option>) -> Result { + let proxy_base_url = std::env::var("LITELLM_PROXY_BASE_URL") + .ok() + .map(|url| url.trim().to_string()) + .filter(|url| !url.is_empty()) + .unwrap_or_else(|| DEFAULT_PROXY_BASE_URL.to_string()); + let tokens = match std::env::var("LITELLM_ANTHROPIC_TOKENIZER_PATH") { + Ok(path) if !path.trim().is_empty() => TokenCounter::from_file( + Path::new(path.trim()), + env_usize( + "LITELLM_TOKENIZER_CONCURRENCY", + DEFAULT_TOKENIZER_CONCURRENCY, + ), + )?, + _ => TokenCounter::approximate(), + }; + Ok( + Self::new(IdentityCache::new(master_key, proxy_base_url), tokens).with_limits( + env_usize("LITELLM_MAX_REQUEST_BYTES", DEFAULT_MAX_REQUEST_BYTES), + env_usize("LITELLM_MAX_INPUT_TOKENS", DEFAULT_MAX_INPUT_TOKENS), + ), + ) + } + + pub fn with_limits(self, max_request_bytes: usize, max_input_tokens: usize) -> Self { + Self { + max_request_bytes, + max_input_tokens, + ..self + } + } + + pub fn identities(&self) -> &IdentityCache { + &self.identities + } + + pub fn tokens(&self) -> &TokenCounter { + &self.tokens + } + + pub fn max_request_bytes(&self) -> usize { + self.max_request_bytes + } + + pub async fn admit(&self, bearer: Option<&str>, raw: &[u8]) -> Result { + let started = Instant::now(); + let bearer = bearer + .map(str::trim) + .filter(|token| !token.is_empty()) + .ok_or(Rejection::Unauthorized)?; + if raw.len() > self.max_request_bytes { + return Err(Rejection::RequestTooLarge { + bytes: raw.len(), + max: self.max_request_bytes, + }); + } + let ParsedRequest { body, model, text } = ParsedRequest::parse(raw)?; + let (identity, input_tokens) = tokio::join!( + self.identities.resolve(Some(bearer)), + self.tokens.count(text) + ); + let identity = identity?; + let input_tokens = input_tokens + .map_err(|error: TokenizerError| Rejection::Tokenizer(error.to_string()))?; + if let Identity::VirtualKey { limits, .. } = &identity + && !limits.allows_model(&model) + { + return Err(Rejection::ModelNotAllowed(model)); + } + if input_tokens > self.max_input_tokens { + return Err(Rejection::ContextTooLarge { + tokens: input_tokens, + max: self.max_input_tokens, + }); + } + self.limits + .reserve(&identity, input_tokens) + .map_err(Rejection::LimitExceeded)?; + Ok(Admitted { + body, + identity, + input_tokens, + elapsed_ms: started.elapsed().as_secs_f64() * 1000.0, + }) + } +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use serde_json::json; + + use super::*; + + fn admission() -> Admission { + let identities = IdentityCache::new( + Some(Arc::from("sk-master")), + "http://127.0.0.1:1".to_string(), + ); + identities.insert( + "sk-limited", + KeyLimits { + models: vec!["claude".to_string()], + max_budget: None, + spend: 0.0, + tpm_limit: Some(100), + rpm_limit: None, + }, + ); + Admission::new(identities, TokenCounter::approximate()) + } + + fn body(model: &str, words: usize) -> Vec { + json!({ + "model": model, + "max_tokens": 16, + "messages": [{"role": "user", "content": "word ".repeat(words)}] + }) + .to_string() + .into_bytes() + } + + #[test] + fn parse_collects_system_messages_and_tools_once() { + let raw = json!({ + "model": "claude", + "system": [{"type": "text", "text": "be brief"}], + "messages": [ + {"role": "user", "content": "hi"}, + {"role": "assistant", "content": [{"type": "text", "text": "hello"}]} + ], + "tools": [{"name": "t", "input_schema": {"type": "object"}}] + }) + .to_string(); + let parsed = ParsedRequest::parse(raw.as_bytes()).unwrap(); + assert_eq!(parsed.model, "claude"); + assert!(parsed.text.contains("be brief\n")); + assert!(parsed.text.contains("hi\n")); + assert!(parsed.text.contains("hello\n")); + assert!(parsed.text.contains("input_schema")); + assert_eq!(parsed.body["messages"][0]["content"], "hi"); + } + + #[test] + fn parse_rejects_non_object_and_missing_model() { + assert!(matches!( + ParsedRequest::parse(b"[]"), + Err(Rejection::InvalidRequest(_)) + )); + assert!(matches!( + ParsedRequest::parse(br#"{"messages": []}"#), + Err(Rejection::InvalidRequest(_)) + )); + assert!(matches!( + ParsedRequest::parse(b"{not json"), + Err(Rejection::InvalidRequest(_)) + )); + } + + #[tokio::test] + async fn master_key_is_admitted_with_a_token_count() { + let admitted = admission() + .admit(Some("sk-master"), &body("claude", 40)) + .await + .unwrap(); + assert_eq!(admitted.identity, Identity::Master); + assert!(admitted.input_tokens > 0); + assert_eq!(admitted.body["model"], "claude"); + } + + #[tokio::test] + async fn missing_bearer_is_unauthorized() { + assert!(matches!( + admission().admit(None, &body("claude", 1)).await, + Err(Rejection::Unauthorized) + )); + } + + #[tokio::test] + async fn virtual_key_model_access_is_enforced() { + assert!(matches!( + admission().admit(Some("sk-limited"), &body("other", 1)).await, + Err(Rejection::ModelNotAllowed(model)) if model == "other" + )); + } + + #[tokio::test] + async fn one_token_count_feeds_the_tpm_window() { + let admission = admission(); + let first = admission + .admit(Some("sk-limited"), &body("claude", 40)) + .await + .unwrap(); + assert!(first.input_tokens > 40, "{}", first.input_tokens); + assert!(matches!( + admission + .admit(Some("sk-limited"), &body("claude", 40)) + .await, + Err(Rejection::LimitExceeded(LimitExceeded::TokensPerMinute)) + )); + } + + #[tokio::test] + async fn oversized_bodies_are_rejected_before_parsing() { + let admission = admission().with_limits(16, 1_000_000); + assert!(matches!( + admission + .admit(Some("sk-master"), &body("claude", 10)) + .await, + Err(Rejection::RequestTooLarge { .. }) + )); + } + + #[tokio::test] + async fn context_limit_uses_the_same_count() { + let admission = admission().with_limits(1 << 20, 10); + assert!(matches!( + admission.admit(Some("sk-master"), &body("claude", 40)).await, + Err(Rejection::ContextTooLarge { tokens, max: 10 }) if tokens > 10 + )); + } +} diff --git a/litellm-rust/crates/ai-gateway/src/admission/tokenizer.rs b/litellm-rust/crates/ai-gateway/src/admission/tokenizer.rs new file mode 100644 index 00000000000..f604ba2fbb7 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/admission/tokenizer.rs @@ -0,0 +1,111 @@ +//! Input token counting kept off the async worker threads. +//! +//! Large inputs are encoded on the blocking pool behind a semaphore, so a burst of 100K-token +//! requests can never stall the threads that accept and answer small requests. + +use std::path::Path; +use std::sync::Arc; + +use tokio::sync::Semaphore; + +use crate::constants::{APPROX_BYTES_PER_TOKEN, TOKENIZE_INLINE_MAX_BYTES}; + +#[derive(Debug, thiserror::Error)] +pub enum TokenizerError { + #[error("failed to load tokenizer: {0}")] + Load(String), + #[error("tokenization failed: {0}")] + Encode(String), +} + +#[derive(Clone)] +enum Backend { + HuggingFace(Arc), + Approximate, +} + +/// Counts input tokens with a bounded number of concurrent encodes. +#[derive(Clone)] +pub struct TokenCounter { + backend: Backend, + permits: Arc, +} + +impl TokenCounter { + /// Load a HuggingFace `tokenizer.json` (the proxy ships the Anthropic one under + /// `litellm/litellm_core_utils/tokenizers/anthropic_tokenizer.json`). + pub fn from_file(path: &Path, concurrency: usize) -> Result { + let tokenizer = tokenizers::Tokenizer::from_file(path) + .map_err(|error| TokenizerError::Load(error.to_string()))?; + Ok(Self { + backend: Backend::HuggingFace(Arc::new(tokenizer)), + permits: Arc::new(Semaphore::new(concurrency.max(1))), + }) + } + + /// `len / APPROX_BYTES_PER_TOKEN`, for hosts without a tokenizer file. + pub fn approximate() -> Self { + Self { + backend: Backend::Approximate, + permits: Arc::new(Semaphore::new(1)), + } + } + + pub fn is_exact(&self) -> bool { + matches!(self.backend, Backend::HuggingFace(_)) + } + + pub async fn count(&self, text: String) -> Result { + let tokenizer = match &self.backend { + Backend::Approximate => return Ok(text.len().div_ceil(APPROX_BYTES_PER_TOKEN)), + Backend::HuggingFace(tokenizer) => Arc::clone(tokenizer), + }; + if text.len() <= TOKENIZE_INLINE_MAX_BYTES { + return encode_len(&tokenizer, &text); + } + let _permit = self + .permits + .acquire() + .await + .map_err(|_| TokenizerError::Encode("tokenizer pool closed".to_string()))?; + tokio::task::spawn_blocking(move || encode_len(&tokenizer, &text)) + .await + .map_err(|error| TokenizerError::Encode(error.to_string()))? + } +} + +fn encode_len(tokenizer: &tokenizers::Tokenizer, text: &str) -> Result { + tokenizer + .encode_fast(text, false) + .map(|encoding| encoding.len()) + .map_err(|error| TokenizerError::Encode(error.to_string())) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn approximate_counter_rounds_up() { + let counter = TokenCounter::approximate(); + assert_eq!(counter.count("abcde".to_string()).await.unwrap(), 2); + assert_eq!(counter.count(String::new()).await.unwrap(), 0); + } + + #[tokio::test] + async fn loads_anthropic_tokenizer_and_counts_off_thread() { + let path = Path::new(env!("CARGO_MANIFEST_DIR")) + .join("../../../litellm/litellm_core_utils/tokenizers/anthropic_tokenizer.json"); + let counter = TokenCounter::from_file(&path, 1).expect("tokenizer loads"); + assert!(counter.is_exact()); + let small = counter.count("hello world".to_string()).await.unwrap(); + assert!((1..=4).contains(&small), "got {small}"); + let large = "the quick brown fox ".repeat(2000); + let large_len = large.len(); + let count = counter.count(large).await.unwrap(); + assert!( + count > large_len / 8 && count < large_len / 2, + "got {count}" + ); + } +} diff --git a/litellm-rust/crates/ai-gateway/src/auth/mod.rs b/litellm-rust/crates/ai-gateway/src/auth/mod.rs index b09d8285c3a..f0e89f3d901 100644 --- a/litellm-rust/crates/ai-gateway/src/auth/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/auth/mod.rs @@ -10,7 +10,7 @@ use axum::extract::FromRequestParts; use axum::http::StatusCode; -use axum::http::header::AUTHORIZATION; +use axum::http::header::{AUTHORIZATION, HeaderMap}; use axum::http::request::Parts; use sha2::{Digest, Sha256}; use subtle::ConstantTimeEq; @@ -35,6 +35,15 @@ pub fn hash_token(token: &str) -> String { hex } +/// The trimmed token after `Authorization: Bearer `, if the header carries one. +pub fn bearer_token(headers: &HeaderMap) -> Option<&str> { + headers + .get(AUTHORIZATION) + .and_then(|value| value.to_str().ok()) + .and_then(|value| value.strip_prefix("Bearer ")) + .map(str::trim) +} + /// Extractor that requires the configured master key as a bearer token. /// /// Rejections: `500` when no master key is configured (permanent @@ -56,13 +65,7 @@ impl FromRequestParts for RequireMasterKey { "gateway auth not configured (set LITELLM_MASTER_KEY)".to_string(), )); }; - let provided = parts - .headers - .get(AUTHORIZATION) - .and_then(|value| value.to_str().ok()) - .and_then(|value| value.strip_prefix("Bearer ")) - .map(str::trim); - match provided { + match bearer_token(&parts.headers) { Some(token) if bool::from(token.as_bytes().ct_eq(expected.as_bytes())) => Ok(Self), _ => Err(( StatusCode::UNAUTHORIZED, diff --git a/litellm-rust/crates/ai-gateway/src/constants.rs b/litellm-rust/crates/ai-gateway/src/constants.rs index 78af374bf70..c48c37c3cda 100644 --- a/litellm-rust/crates/ai-gateway/src/constants.rs +++ b/litellm-rust/crates/ai-gateway/src/constants.rs @@ -40,3 +40,39 @@ pub(crate) const MESSAGES_ROUTE_PATH: &str = "/v1/messages"; #[cfg(feature = "server")] pub(crate) const MESSAGES_HEADERS_NOT_FORWARDED: &[&str] = &["authorization", "connection", "content-length", "host"]; + +/// Response header carrying the wall time the gateway spent admitting a request. +#[cfg(feature = "server")] +pub(crate) const ADMISSION_DURATION_HEADER: &str = "x-litellm-admission-duration-ms"; + +/// Largest `/v1/messages` body accepted before parsing. Override: `LITELLM_MAX_REQUEST_BYTES`. +#[cfg(feature = "server")] +pub(crate) const DEFAULT_MAX_REQUEST_BYTES: usize = 32 * 1024 * 1024; + +/// Largest admitted input token count. Override: `LITELLM_MAX_INPUT_TOKENS`. +#[cfg(feature = "server")] +pub(crate) const DEFAULT_MAX_INPUT_TOKENS: usize = 1_000_000; + +/// Concurrent tokenizer runs on the blocking pool. Override: `LITELLM_TOKENIZER_CONCURRENCY`. +#[cfg(feature = "server")] +pub(crate) const DEFAULT_TOKENIZER_CONCURRENCY: usize = 2; + +/// Inputs at or under this size are tokenized inline; larger ones go to the blocking pool. +#[cfg(feature = "server")] +pub(crate) const TOKENIZE_INLINE_MAX_BYTES: usize = 16 * 1024; + +/// Bytes per token used when no tokenizer file is configured. +#[cfg(feature = "server")] +pub(crate) const APPROX_BYTES_PER_TOKEN: usize = 4; + +/// Input price used to reserve budget before the provider reports usage (USD per token). +#[cfg(feature = "server")] +pub(crate) const DEFAULT_INPUT_COST_PER_TOKEN: f64 = 3e-6; + +/// The Python proxy endpoint that resolves a virtual key to its limits. +#[cfg(feature = "server")] +pub(crate) const PROXY_KEY_INFO_PATH: &str = "/key/info"; + +/// Timeout for a virtual-key lookup against the Python proxy. +#[cfg(feature = "server")] +pub(crate) const KEY_INFO_TIMEOUT_SECS: u64 = 5; diff --git a/litellm-rust/crates/ai-gateway/src/lib.rs b/litellm-rust/crates/ai-gateway/src/lib.rs index 08fbde564ed..6c2178a8255 100644 --- a/litellm-rust/crates/ai-gateway/src/lib.rs +++ b/litellm-rust/crates/ai-gateway/src/lib.rs @@ -17,6 +17,8 @@ mod client; pub mod io; pub mod ocr; +#[cfg(feature = "server")] +pub mod admission; #[cfg(feature = "server")] pub mod auth; #[cfg(feature = "server")] diff --git a/litellm-rust/crates/ai-gateway/src/main.rs b/litellm-rust/crates/ai-gateway/src/main.rs index 88d7b1dbcf8..3e1d325cc72 100644 --- a/litellm-rust/crates/ai-gateway/src/main.rs +++ b/litellm-rust/crates/ai-gateway/src/main.rs @@ -11,6 +11,7 @@ use std::sync::Arc; +use litellm_ai_gateway::admission::Admission; use litellm_ai_gateway::io::realtime_pool::{PoolConfig, RealtimePool, upstream_key}; use litellm_ai_gateway::routes; use litellm_ai_gateway::state::AppState; @@ -49,6 +50,19 @@ async fn main() { let router = Arc::new(build_router()); + let admission = match Admission::from_env(master_key.clone()) { + Ok(admission) => Arc::new(admission), + Err(error) => { + eprintln!("admission setup failed: {error}"); + std::process::exit(1); + } + }; + if !admission.tokens().is_exact() { + eprintln!( + "warning: LITELLM_ANTHROPIC_TOKENIZER_PATH is not set; input tokens are approximated" + ); + } + // Build the pre-warmed realtime pool and register each deployment's upstream // so the background replenisher starts warming it. `REALTIME_POOL_SIZE=0` // yields a disabled pool → every connect fresh-dials (original behavior). @@ -70,6 +84,7 @@ async fn main() { let state = AppState { router, master_key, + admission, loggers: Arc::new(loggers), realtime_pool, }; diff --git a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs index bb9f3851a77..9d26b48b565 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs @@ -12,8 +12,10 @@ use axum::routing::post; use litellm_core::Error; use serde_json::{Map, Value}; -use crate::auth::RequireMasterKey; -use crate::constants::{MESSAGES_HEADERS_NOT_FORWARDED, MESSAGES_ROUTE_PATH}; +use crate::admission::{Admit, Admitted}; +use crate::constants::{ + ADMISSION_DURATION_HEADER, MESSAGES_HEADERS_NOT_FORWARDED, MESSAGES_ROUTE_PATH, +}; use crate::state::AppState; /// This route's contribution to the app router. @@ -28,19 +30,27 @@ pub fn router() -> Router { skip_all )] async fn handle( - _auth: RequireMasterKey, State(state): State, headers: HeaderMap, - Json(body): Json, + Admit(admitted): Admit, ) -> Result { + let Admitted { + body, elapsed_ms, .. + } = admitted; let extra_headers = forwarded_headers(&headers)?; - match service::run(&state.router, body, extra_headers) + let mut response = match service::run(&state.router, body, extra_headers) .await .map_err(MessagesRouteError::from)? { - service::MessagesResponse::Json(body) => Ok(Json(body).into_response()), - service::MessagesResponse::Stream(upstream) => stream_response(upstream), + service::MessagesResponse::Json(body) => Json(body).into_response(), + service::MessagesResponse::Stream(upstream) => stream_response(upstream)?, + }; + if let Ok(value) = HeaderValue::from_str(&format!("{elapsed_ms:.3}")) { + response + .headers_mut() + .insert(ADMISSION_DURATION_HEADER, value); } + Ok(response) } fn stream_response(upstream: reqwest::Response) -> Result { @@ -149,6 +159,8 @@ mod tests { use tower::ServiceExt; use super::super::app; + use crate::admission::{Admission, IdentityCache, KeyLimits, TokenCounter}; + use crate::constants::ADMISSION_DURATION_HEADER; use crate::io::realtime_pool::RealtimePool; use crate::state::AppState; @@ -172,6 +184,10 @@ mod tests { }, }])), master_key: master_key.map(Arc::from), + admission: Arc::new(Admission::new( + IdentityCache::new(master_key.map(Arc::from), "http://127.0.0.1:1".to_string()), + TokenCounter::approximate(), + )), loggers: Arc::new(Vec::new()), realtime_pool: RealtimePool::disabled(), } @@ -460,6 +476,54 @@ mod tests { server.await.expect("upstream task completes"); } + #[tokio::test] + async fn route_reports_admission_time_and_admits_virtual_keys_by_model() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); + let (api_base, server) = upstream(listener).await; + let state = state("claude-test", api_base, Some("master-key")); + state.admission.identities().insert( + "sk-virtual", + KeyLimits { + models: vec!["claude-test".to_string()], + ..KeyLimits::default() + }, + ); + let request = |model: &str| { + Request::builder() + .method("POST") + .uri("/v1/messages") + .header("authorization", "Bearer sk-virtual") + .header("content-type", "application/json") + .body(Body::from( + json!({ + "model": model, + "max_tokens": 16, + "messages": [{"role": "user", "content": "hello"}] + }) + .to_string(), + )) + .expect("request builds") + }; + let denied = app(state.clone()) + .oneshot(request("other-model")) + .await + .expect("route responds"); + assert_eq!(denied.status(), StatusCode::FORBIDDEN); + let admitted = app(state) + .oneshot(request("claude-test")) + .await + .expect("route responds"); + assert_eq!(admitted.status(), StatusCode::OK); + let elapsed: f64 = admitted + .headers() + .get(ADMISSION_DURATION_HEADER) + .and_then(|value| value.to_str().ok()) + .and_then(|value| value.parse().ok()) + .expect("admission duration header is a number"); + assert!(elapsed >= 0.0); + server.await.expect("upstream task completes"); + } + #[tokio::test] async fn route_rejects_missing_master_key() { let app = app(state( @@ -482,12 +546,17 @@ mod tests { } #[tokio::test] - async fn route_rejects_invalid_master_key() { + async fn route_rejects_unknown_key_when_identity_lookup_is_unavailable() { let app = app(state( "claude-test", "http://127.0.0.1:1".to_string(), Some("master-key"), )); + let body = json!({ + "model": "claude-test", + "max_tokens": 8, + "messages": [{"role": "user", "content": "hi"}] + }); let response = app .oneshot( Request::builder() @@ -495,12 +564,12 @@ mod tests { .uri("/v1/messages") .header("authorization", "Bearer wrong-key") .header("content-type", "application/json") - .body(Body::from("{}")) + .body(Body::from(body.to_string())) .expect("request builds"), ) .await .expect("route responds"); - assert_eq!(response.status(), StatusCode::UNAUTHORIZED); + assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); } #[tokio::test] diff --git a/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs index a94853e106d..5efd2bf94e0 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs @@ -239,6 +239,7 @@ async fn bridge( #[cfg(test)] mod tests { use super::*; + use crate::admission::{Admission, IdentityCache, TokenCounter}; use crate::io::realtime_pool::RealtimePool; use crate::state::AppState; use axum::body::Body; @@ -316,6 +317,13 @@ mod tests { AppState { router: Arc::new(ModelRouter::default()), master_key: Some(Arc::from("master-key")), + admission: Arc::new(Admission::new( + IdentityCache::new( + Some(Arc::from("master-key")), + "http://127.0.0.1:1".to_string(), + ), + TokenCounter::approximate(), + )), loggers: Arc::new(Vec::new()), realtime_pool: RealtimePool::disabled(), } diff --git a/litellm-rust/crates/ai-gateway/src/state.rs b/litellm-rust/crates/ai-gateway/src/state.rs index 3b61d8309ea..b7c0e72f26a 100644 --- a/litellm-rust/crates/ai-gateway/src/state.rs +++ b/litellm-rust/crates/ai-gateway/src/state.rs @@ -1,9 +1,10 @@ use std::sync::Arc; -use crate::io::realtime_pool::RealtimePool; use litellm_core::router::Router; +use crate::admission::Admission; use crate::integrations::custom_logger::CustomLogger; +use crate::io::realtime_pool::RealtimePool; /// Shared application state handed to every route handler. #[derive(Clone)] @@ -12,6 +13,8 @@ pub struct AppState { /// The gateway master key. Any caller presenting it as a bearer token may /// invoke the gateway. `None` → auth not configured (routes fail closed). pub master_key: Option>, + /// Per-request admission (identity, model access, size, tokens, limits) for `/v1/messages`. + pub admission: Arc, /// Logging callbacks fanned out at the end of each realtime session. pub loggers: Arc>>, /// Pre-warmed upstream realtime connection pool. Disabled diff --git a/litellm-rust/crates/ai-gateway/src/trace_parity.rs b/litellm-rust/crates/ai-gateway/src/trace_parity.rs index 614852c541d..b8c54d6e1f9 100644 --- a/litellm-rust/crates/ai-gateway/src/trace_parity.rs +++ b/litellm-rust/crates/ai-gateway/src/trace_parity.rs @@ -11,6 +11,7 @@ use serde::Serialize; use serde_json::Value; use tower::ServiceExt; +use crate::admission::{Admission, IdentityCache, TokenCounter}; use crate::io::realtime_pool::RealtimePool; use crate::routes; use crate::state::AppState; @@ -37,6 +38,13 @@ pub async fn messages_request( }, }])), master_key: Some(Arc::from("trace-master-key")), + admission: Arc::new(Admission::new( + IdentityCache::new( + Some(Arc::from("trace-master-key")), + "http://127.0.0.1:1".to_string(), + ), + TokenCounter::approximate(), + )), loggers: Arc::new(Vec::new()), realtime_pool: RealtimePool::disabled(), };