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>
This commit is contained in:
yassin 2026-09-09 06:06:04 +00:00
parent ee7c7e14f3
commit 85cc7173fa
15 changed files with 1505 additions and 29 deletions

405
litellm-rust/Cargo.lock generated
View file

@ -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"

View file

@ -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"]

View file

@ -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<AppState> for Admit {
type Rejection = Response;
async fn from_request(request: Request, state: &AppState) -> Result<Self, Self::Rejection> {
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()
}

View file

@ -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<String>,
#[serde(default)]
pub max_budget: Option<f64>,
#[serde(default)]
pub spend: f64,
#[serde(default)]
pub tpm_limit: Option<u64>,
#[serde(default)]
pub rpm_limit: Option<u64>,
}
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<KeyLimits>,
},
}
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<Arc<str>>,
proxy_base_url: String,
http: reqwest::Client,
cache: RwLock<HashMap<String, Arc<KeyLimits>>>,
}
impl IdentityCache {
pub fn new(master_key: Option<Arc<str>>, 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<Identity, IdentityError> {
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<KeyLimits, IdentityError> {
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::<KeyInfoResponse>()
.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"));
}
}

View file

@ -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<HashMap<String, KeyWindow>>,
}
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(()));
}
}

View file

@ -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<IdentityError> 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<Self, Rejection> {
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<Arc<str>>) -> Result<Self, TokenizerError> {
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<Admitted, Rejection> {
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<u8> {
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
));
}
}

View file

@ -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<tokenizers::Tokenizer>),
Approximate,
}
/// Counts input tokens with a bounded number of concurrent encodes.
#[derive(Clone)]
pub struct TokenCounter {
backend: Backend,
permits: Arc<Semaphore>,
}
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<Self, TokenizerError> {
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<usize, TokenizerError> {
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<usize, TokenizerError> {
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}"
);
}
}

View file

@ -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<AppState> 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,

View file

@ -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;

View file

@ -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")]

View file

@ -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,
};

View file

@ -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<AppState> {
skip_all
)]
async fn handle(
_auth: RequireMasterKey,
State(state): State<AppState>,
headers: HeaderMap,
Json(body): Json<Value>,
Admit(admitted): Admit,
) -> Result<Response, MessagesRouteError> {
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<Response, MessagesRouteError> {
@ -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]

View file

@ -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(),
}

View file

@ -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<Arc<str>>,
/// Per-request admission (identity, model access, size, tokens, limits) for `/v1/messages`.
pub admission: Arc<Admission>,
/// Logging callbacks fanned out at the end of each realtime session.
pub loggers: Arc<Vec<Arc<dyn CustomLogger>>>,
/// Pre-warmed upstream realtime connection pool. Disabled

View file

@ -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(),
};