diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 6578d873faa..982c473ac1f 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -824,9 +824,9 @@ dependencies = [ [[package]] name = "daachorse" -version = "1.0.1" +version = "3.0.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6f55d7153ba3b507595872a3874803f07a8a81d1e888abed8e5db7da0597d6e2" +checksum = "5614204febbc33cc07a2806aa6440b904ac012b68eecc37f4493ea4a76455a3d" [[package]] name = "darling" @@ -1497,6 +1497,8 @@ checksum = "d466e9454f08e4a911e14806c24e16fba1b4c121d1ea474396f396069cf949d9" dependencies = [ "equivalent", "hashbrown", + "serde", + "serde_core", ] [[package]] @@ -1579,8 +1581,6 @@ dependencies = [ "serde_json", "sha2 0.10.9", "subtle", - "thiserror 2.0.19", - "tokenizers", "tokio", "tokio-tungstenite", "tower", @@ -1608,6 +1608,7 @@ dependencies = [ "aws-smithy-runtime-api", "aws-types", "base64 0.22.1", + "indexmap", "rand 0.8.7", "reqwest", "rstest", @@ -1615,6 +1616,7 @@ dependencies = [ "serde_json", "sha2 0.10.9", "thiserror 2.0.19", + "tokenizers", "tokio", "tracing", "tracing-subscriber", @@ -2855,9 +2857,9 @@ checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" [[package]] name = "tokenizers" -version = "0.23.1" +version = "0.23.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "44e5bea67576e04b6ff8564c5d9e09c2ef0cf476502245f2f120e497769d3112" +checksum = "7afbf6e88718afcc138bad01d6ccc3051dbbc3b2ce9793d8b8a3aeb610969cfc" dependencies = [ "ahash", "compact_str", diff --git a/litellm-rust/crates/ai-gateway/Cargo.toml b/litellm-rust/crates/ai-gateway/Cargo.toml index 6df0a534bdb..74cf66e88a2 100644 --- a/litellm-rust/crates/ai-gateway/Cargo.toml +++ b/litellm-rust/crates/ai-gateway/Cargo.toml @@ -32,19 +32,15 @@ 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", "dep:tokenizers"] +server = ["dep:axum", "dep:subtle", "dep:sha2"] # 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 deleted file mode 100644 index 859024a5f54..00000000000 --- a/litellm-rust/crates/ai-gateway/src/admission/extract.rs +++ /dev/null @@ -1,49 +0,0 @@ -//! 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 deleted file mode 100644 index a570870d89c..00000000000 --- a/litellm-rust/crates/ai-gateway/src/admission/identity.rs +++ /dev/null @@ -1,221 +0,0 @@ -//! 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 deleted file mode 100644 index 34a61a7364b..00000000000 --- a/litellm-rust/crates/ai-gateway/src/admission/limits.rs +++ /dev/null @@ -1,171 +0,0 @@ -//! 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 deleted file mode 100644 index 9fb785292ea..00000000000 --- a/litellm-rust/crates/ai-gateway/src/admission/mod.rs +++ /dev/null @@ -1,389 +0,0 @@ -//! 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 deleted file mode 100644 index f604ba2fbb7..00000000000 --- a/litellm-rust/crates/ai-gateway/src/admission/tokenizer.rs +++ /dev/null @@ -1,111 +0,0 @@ -//! 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 f0e89f3d901..b09d8285c3a 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, HeaderMap}; +use axum::http::header::AUTHORIZATION; use axum::http::request::Parts; use sha2::{Digest, Sha256}; use subtle::ConstantTimeEq; @@ -35,15 +35,6 @@ 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 @@ -65,7 +56,13 @@ impl FromRequestParts for RequireMasterKey { "gateway auth not configured (set LITELLM_MASTER_KEY)".to_string(), )); }; - match bearer_token(&parts.headers) { + 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 { 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 c48c37c3cda..78af374bf70 100644 --- a/litellm-rust/crates/ai-gateway/src/constants.rs +++ b/litellm-rust/crates/ai-gateway/src/constants.rs @@ -40,39 +40,3 @@ 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 6c2178a8255..08fbde564ed 100644 --- a/litellm-rust/crates/ai-gateway/src/lib.rs +++ b/litellm-rust/crates/ai-gateway/src/lib.rs @@ -17,8 +17,6 @@ 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 3e1d325cc72..88d7b1dbcf8 100644 --- a/litellm-rust/crates/ai-gateway/src/main.rs +++ b/litellm-rust/crates/ai-gateway/src/main.rs @@ -11,7 +11,6 @@ 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; @@ -50,19 +49,6 @@ 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). @@ -84,7 +70,6 @@ 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 9d26b48b565..bb9f3851a77 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs @@ -12,10 +12,8 @@ use axum::routing::post; use litellm_core::Error; use serde_json::{Map, Value}; -use crate::admission::{Admit, Admitted}; -use crate::constants::{ - ADMISSION_DURATION_HEADER, MESSAGES_HEADERS_NOT_FORWARDED, MESSAGES_ROUTE_PATH, -}; +use crate::auth::RequireMasterKey; +use crate::constants::{MESSAGES_HEADERS_NOT_FORWARDED, MESSAGES_ROUTE_PATH}; use crate::state::AppState; /// This route's contribution to the app router. @@ -30,27 +28,19 @@ pub fn router() -> Router { skip_all )] async fn handle( + _auth: RequireMasterKey, State(state): State, headers: HeaderMap, - Admit(admitted): Admit, + Json(body): Json, ) -> Result { - let Admitted { - body, elapsed_ms, .. - } = admitted; let extra_headers = forwarded_headers(&headers)?; - let mut response = match service::run(&state.router, body, extra_headers) + match service::run(&state.router, body, extra_headers) .await .map_err(MessagesRouteError::from)? { - 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); + service::MessagesResponse::Json(body) => Ok(Json(body).into_response()), + service::MessagesResponse::Stream(upstream) => stream_response(upstream), } - Ok(response) } fn stream_response(upstream: reqwest::Response) -> Result { @@ -159,8 +149,6 @@ 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; @@ -184,10 +172,6 @@ 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(), } @@ -476,54 +460,6 @@ 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( @@ -546,17 +482,12 @@ mod tests { } #[tokio::test] - async fn route_rejects_unknown_key_when_identity_lookup_is_unavailable() { + async fn route_rejects_invalid_master_key() { 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() @@ -564,12 +495,12 @@ mod tests { .uri("/v1/messages") .header("authorization", "Bearer wrong-key") .header("content-type", "application/json") - .body(Body::from(body.to_string())) + .body(Body::from("{}")) .expect("request builds"), ) .await .expect("route responds"); - assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); + assert_eq!(response.status(), StatusCode::UNAUTHORIZED); } #[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 5efd2bf94e0..a94853e106d 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs @@ -239,7 +239,6 @@ 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; @@ -317,13 +316,6 @@ 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 b7c0e72f26a..3b61d8309ea 100644 --- a/litellm-rust/crates/ai-gateway/src/state.rs +++ b/litellm-rust/crates/ai-gateway/src/state.rs @@ -1,10 +1,9 @@ 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)] @@ -13,8 +12,6 @@ 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 b8c54d6e1f9..614852c541d 100644 --- a/litellm-rust/crates/ai-gateway/src/trace_parity.rs +++ b/litellm-rust/crates/ai-gateway/src/trace_parity.rs @@ -11,7 +11,6 @@ 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; @@ -38,13 +37,6 @@ 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(), }; diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index c0de7ff3977..2f157b2c20d 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -15,6 +15,10 @@ thiserror.workspace = true tracing.workspace = true tracing-subscriber = { workspace = true, optional = true } sha2.workspace = true +indexmap = { version = "2.14.0", features = ["serde"] } +# HuggingFace tokenizer for input token counting; without the default features it +# pulls no HTTP client or progress bars, only the `onig` regex backend. +tokenizers = { version = "0.23.1", default-features = false, features = ["onig"] } aws-config = { version = "1.9.0", default-features = false, features = ["rustls", "rt-tokio"], optional = true } aws-credential-types = { version = "1.3.0", features = ["hardcoded-credentials"], optional = true } aws-sdk-sts = { version = "1.108.0", default-features = false, features = ["rustls", "rt-tokio"], optional = true } diff --git a/litellm-rust/crates/core/src/constants.rs b/litellm-rust/crates/core/src/constants.rs index fc81f4fa029..1ff8848b52d 100644 --- a/litellm-rust/crates/core/src/constants.rs +++ b/litellm-rust/crates/core/src/constants.rs @@ -43,3 +43,13 @@ pub const EMPTY_TEXT_PLACEHOLDER: &str = "[System: Empty message content sanitised to satisfy protocol]"; pub const FUNCTION_TRACE_TARGET: &str = "litellm::function_trace"; + +/// Message accounting `litellm.token_counter` adds on top of the raw encoding +/// for non-OpenAI models (`litellm/litellm_core_utils/token_counter.py`). +pub(crate) const TOKENS_PER_MESSAGE: usize = 3; +pub(crate) const TOKENS_PER_NAME: usize = 1; +pub(crate) const REPLY_PRIMING_TOKENS: usize = 3; +pub(crate) const TOOL_DEFINITIONS_TOKENS: usize = 9; +pub(crate) const TOOLS_WITH_SYSTEM_MESSAGE_DISCOUNT: usize = 4; +pub(crate) const TOOL_CHOICE_NONE_TOKENS: usize = 1; +pub(crate) const NAMED_TOOL_CHOICE_TOKENS: usize = 7; diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs index b93e084f57e..ec42a8301f5 100644 --- a/litellm-rust/crates/core/src/lib.rs +++ b/litellm-rust/crates/core/src/lib.rs @@ -14,5 +14,6 @@ pub mod realtime; pub mod responses; pub mod router; pub mod routing_utils; +pub mod token_counter; pub use error::Error; diff --git a/litellm-rust/crates/core/src/token_counter/mod.rs b/litellm-rust/crates/core/src/token_counter/mod.rs new file mode 100644 index 00000000000..637691407a4 --- /dev/null +++ b/litellm-rust/crates/core/src/token_counter/mod.rs @@ -0,0 +1,156 @@ +//! Input token counting for a request body, mirroring `litellm.token_counter` +//! for the shapes it can count exactly. Everything else is declined so the host +//! keeps its own counter as the reference. + +mod tools; +pub mod types; + +use serde::Serialize; +use thiserror::Error as ThisError; + +use crate::constants::{ + NAMED_TOOL_CHOICE_TOKENS, REPLY_PRIMING_TOKENS, TOKENS_PER_MESSAGE, TOKENS_PER_NAME, + TOOL_CHOICE_NONE_TOKENS, TOOL_DEFINITIONS_TOKENS, TOOLS_WITH_SYSTEM_MESSAGE_DISCOUNT, +}; +use tools::format_function_definitions; +use types::{ContentBlock, ContentItem, CountableRequest, Message, MessageContent, ToolChoice}; + +#[derive(Debug, ThisError, PartialEq, Eq)] +pub enum TokenCountError { + #[error("failed to load tokenizer: {0}")] + Load(String), + /// The body is outside the shape this counter mirrors exactly. Hosts with a + /// reference counter treat this as "fall back", not "fail". + #[error("unsupported by the rust token counter: {0}")] + Unsupported(String), + #[error("tokenization failed: {0}")] + Encode(String), +} + +#[derive(Clone, Debug, PartialEq, Eq, Serialize)] +pub struct InputTokenCount { + pub model: String, + pub input_tokens: usize, +} + +/// A loaded HuggingFace tokenizer plus the message accounting Python applies on +/// top of it. Encoding is CPU-bound and synchronous; hosts run it off their +/// event loop. +pub struct TokenCounter { + tokenizer: tokenizers::Tokenizer, +} + +impl TokenCounter { + /// Load a HuggingFace `tokenizer.json` document. The host reads the file. + pub fn from_json(tokenizer_json: &str) -> Result { + let tokenizer = tokenizer_json + .parse::() + .map_err(|error| TokenCountError::Load(error.to_string()))?; + Ok(Self { tokenizer }) + } + + pub fn count_text(&self, text: &str) -> Result { + self.tokenizer + .encode_fast(text, true) + .map(|encoding| encoding.len()) + .map_err(|error| TokenCountError::Encode(error.to_string())) + } + + pub fn count_request( + &self, + request: &CountableRequest, + ) -> Result { + let messages = request + .messages + .as_deref() + .ok_or_else(|| TokenCountError::Unsupported("request has no messages".to_string()))?; + let message_tokens = messages + .iter() + .map(|message| self.count_message(message)) + .sum::>()?; + let includes_system_message = messages + .iter() + .any(|message| message.role.as_deref() == Some("system")); + let extra_tokens = self.count_extra( + request.tools.as_deref().unwrap_or_default(), + request.tool_choice.as_ref(), + includes_system_message, + )?; + Ok(InputTokenCount { + model: request.model.clone(), + input_tokens: message_tokens + extra_tokens, + }) + } + + fn count_message(&self, message: &Message) -> Result { + let role_tokens = match &message.role { + Some(role) => self.count_text(role)?, + None => 0, + }; + let name_tokens = match &message.name { + Some(name) => self.count_text(name)? + TOKENS_PER_NAME, + None => 0, + }; + let content_tokens = match &message.content { + Some(MessageContent::Text(text)) => self.count_text(text)?, + Some(MessageContent::Blocks(items)) => items + .iter() + .map(|item| self.count_content_item(item)) + .sum::>()?, + None => 0, + }; + Ok(TOKENS_PER_MESSAGE + role_tokens + name_tokens + content_tokens) + } + + fn count_content_item(&self, item: &ContentItem) -> Result { + match item { + ContentItem::Text(text) => self.count_text(text), + ContentItem::Block(ContentBlock::Text { text }) => self.count_text(text), + ContentItem::Block(ContentBlock::Thinking { thinking }) => { + if thinking.is_empty() { + return Ok(0); + } + self.count_text(thinking) + } + ContentItem::Block(ContentBlock::ToolReference { tool_name }) => { + match tool_name.as_deref().filter(|name| !name.is_empty()) { + Some(name) => self.count_text(name), + None => Ok(0), + } + } + ContentItem::Block(ContentBlock::Unsupported) => Err(TokenCountError::Unsupported( + "content block type is counted by the python path".to_string(), + )), + } + } + + fn count_extra( + &self, + tools: &[types::ToolDefinition], + tool_choice: Option<&ToolChoice>, + includes_system_message: bool, + ) -> Result { + let tool_tokens = if tools.is_empty() { + 0 + } else { + let definitions = self.count_text(&format_function_definitions(tools)?)?; + let discount = if includes_system_message { + TOOLS_WITH_SYSTEM_MESSAGE_DISCOUNT + } else { + 0 + }; + definitions + TOOL_DEFINITIONS_TOKENS - discount + }; + let choice_tokens = match tool_choice { + Some(ToolChoice::Mode(mode)) if mode == "none" => TOOL_CHOICE_NONE_TOKENS, + Some(ToolChoice::Mode(_)) | None => 0, + Some(ToolChoice::Named(named)) => { + NAMED_TOOL_CHOICE_TOKENS + self.count_text(&named.function.name)? + } + }; + Ok(REPLY_PRIMING_TOKENS + tool_tokens + choice_tokens) + } +} + +#[cfg(test)] +mod tests; diff --git a/litellm-rust/crates/core/src/token_counter/tests.rs b/litellm-rust/crates/core/src/token_counter/tests.rs new file mode 100644 index 00000000000..c171bb32bdc --- /dev/null +++ b/litellm-rust/crates/core/src/token_counter/tests.rs @@ -0,0 +1,168 @@ +use rstest::rstest; + +use super::*; + +/// Expected counts are pinned from `litellm.token_counter(model="claude-sonnet-4-5", ...)` +/// so this test also guards Python parity. +fn counter() -> TokenCounter { + let path = concat!( + env!("CARGO_MANIFEST_DIR"), + "/../../../litellm/litellm_core_utils/tokenizers/anthropic_tokenizer.json" + ); + let json = std::fs::read_to_string(path).expect("anthropic tokenizer json is in the repo"); + TokenCounter::from_json(&json).expect("anthropic tokenizer loads") +} + +const SIMPLE: &str = r#"{"model":"claude-sonnet-4-5","messages":[{"role":"user","content":"Hello, how are you today?"}]}"#; + +const BLOCKS_AND_SYSTEM: &str = r#"{"model":"claude-sonnet-4-5","messages":[ + {"role":"system","content":"You are a terse assistant."}, + {"role":"user","name":"alice","content":[ + {"type":"text","text":"Summarise this paragraph about ships and harbours."}, + "plain string item", + {"type":"thinking","thinking":"pondering"}, + {"type":"tool_reference","tool_name":"get_weather"}]}, + {"role":"assistant","content":[{"type":"text","text":"Sure.","cache_control":{"type":"ephemeral"}}]}]}"#; + +const TOOLS_OPENAI: &str = r#"{"model":"claude-sonnet-4-5","messages":[{"role":"user","content":"weather?"}], + "tools":[ + {"type":"function","function":{"name":"get_weather","description":"Get weather","parameters":{ + "type":"object", + "properties":{ + "location":{"type":"string","description":"City name"}, + "unit":{"type":"string","enum":["celsius","fahrenheit"]}, + "days":{"type":"integer"}, + "tags":{"type":"array","items":{"type":"string"}}, + "opts":{"type":"object","properties":{"verbose":{"type":"boolean"},"level":{"type":"integer","enum":[1,2]}},"required":["verbose"]}, + "anything":{}}, + "required":["location"]}}}, + {"type":"function","function":{"name":"noop"}}], + "tool_choice":{"type":"function","function":{"name":"get_weather"}}}"#; + +const TOOLS_ANTHROPIC_SYSTEM: &str = r#"{"model":"claude-sonnet-4-5", + "messages":[{"role":"system","content":"sys"},{"role":"user","content":"weather?"}], + "tools":[{"name":"get_weather","description":"Get weather","input_schema":{ + "type":"object","properties":{"location":{"type":["string","null"]}},"required":["location"]}}], + "tool_choice":"none"}"#; + +#[rstest] +#[case::text_only(SIMPLE, 14)] +#[case::content_blocks_name_and_system(BLOCKS_AND_SYSTEM, 45)] +#[case::openai_tools_named_choice(TOOLS_OPENAI, 123)] +#[case::anthropic_tools_system_discount_choice_none(TOOLS_ANTHROPIC_SYSTEM, 53)] +fn count_request_matches_python_token_counter(#[case] body: &str, #[case] expected: usize) { + let request = CountableRequest::parse(body.as_bytes()).expect("fixture parses"); + let count = counter().count_request(&request).expect("fixture counts"); + assert_eq!( + count, + InputTokenCount { + model: "claude-sonnet-4-5".to_string(), + input_tokens: expected, + } + ); +} + +#[test] +fn tool_definitions_render_like_python() { + let request = CountableRequest::parse(TOOLS_OPENAI.as_bytes()).expect("fixture parses"); + let rendered = format_function_definitions(request.tools.as_deref().unwrap_or_default()) + .expect("fixture renders"); + let expected = "namespace functions {\n\n// Get weather\ntype get_weather = (_: {\n// City name\nlocation: string,\nunit?: \"celsius\" | \"fahrenheit\",\ndays?: number,\ntags?: string[],\nopts?: {\n verbose: boolean,\n level?: \"1\" | \"2\",\n},\nanything?: any,\n}) => any;\n\ntype noop = () => any;\n\n} // namespace functions"; + assert_eq!(rendered, expected); +} + +#[test] +fn union_types_and_anthropic_schema_render_like_python() { + let request = + CountableRequest::parse(TOOLS_ANTHROPIC_SYSTEM.as_bytes()).expect("fixture parses"); + let rendered = format_function_definitions(request.tools.as_deref().unwrap_or_default()) + .expect("fixture renders"); + assert_eq!( + rendered, + "namespace functions {\n\n// Get weather\ntype get_weather = (_: {\nlocation: any,\n}) => any;\n\n} // namespace functions" + ); +} + +#[rstest] +#[case::not_json(b"not json" as &[u8])] +#[case::missing_model(br#"{"messages":[]}"#)] +#[case::messages_not_a_list(br#"{"model":"m","messages":"hi"}"#)] +#[case::message_with_tool_calls( + br#"{"model":"m","messages":[{"role":"assistant","tool_calls":[{"id":"1","type":"function","function":{"name":"f","arguments":"{}"}}]}]}"# +)] +#[case::dict_content( + br#"{"model":"m","messages":[{"role":"user","content":{"type":"text","text":"x"}}]}"# +)] +#[case::float_enum( + br#"{"model":"m","messages":[],"tools":[{"name":"f","input_schema":{"type":"object","properties":{"x":{"type":"number","enum":[1.5]}}}}]}"# +)] +#[case::anthropic_tool_choice_without_function( + br#"{"model":"m","messages":[],"tool_choice":{"type":"auto"}}"# +)] +fn shapes_outside_the_mirror_are_declined_at_parse(#[case] body: &[u8]) { + assert!(matches!( + CountableRequest::parse(body), + Err(TokenCountError::Unsupported(_)) + )); +} + +#[rstest] +#[case::no_messages(br#"{"model":"m"}"# as &[u8])] +#[case::image_block( + br#"{"model":"m","messages":[{"role":"user","content":[{"type":"image","source":{"type":"base64","media_type":"image/png","data":"AA=="}}]}]}"# +)] +#[case::tool_result_block( + br#"{"model":"m","messages":[{"role":"user","content":[{"type":"tool_result","tool_use_id":"1","content":"ok"}]}]}"# +)] +#[case::array_without_items( + br#"{"model":"m","messages":[],"tools":[{"name":"f","input_schema":{"type":"object","properties":{"x":{"type":"array"}}}}]}"# +)] +fn shapes_outside_the_mirror_are_declined_at_count(#[case] body: &[u8]) { + let request = CountableRequest::parse(body).expect("shape parses"); + assert!(matches!( + counter().count_request(&request), + Err(TokenCountError::Unsupported(_)) + )); +} + +#[test] +fn tool_choice_and_system_discount_change_the_count() { + let counter = counter(); + let count = |body: &str| { + counter + .count_request(&CountableRequest::parse(body.as_bytes()).expect("parses")) + .expect("counts") + .input_tokens + }; + let base = count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}]}"#); + assert_eq!( + count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tool_choice":"none"}"#), + base + TOOL_CHOICE_NONE_TOKENS + ); + assert_eq!( + count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tool_choice":"auto"}"#), + base + ); + let with_tools = count( + r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tools":[{"name":"f"}]}"#, + ); + let with_tools_and_system = count( + r#"{"model":"m","messages":[{"role":"system","content":"hi"}],"tools":[{"name":"f"}]}"#, + ); + assert_eq!( + with_tools - with_tools_and_system, + TOOLS_WITH_SYSTEM_MESSAGE_DISCOUNT + ); + assert_eq!( + count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tools":[]}"#), + base + ); +} + +#[test] +fn loading_a_bad_tokenizer_is_a_load_error() { + assert!(matches!( + TokenCounter::from_json("{}"), + Err(TokenCountError::Load(_)) + )); +} diff --git a/litellm-rust/crates/core/src/token_counter/tools.rs b/litellm-rust/crates/core/src/token_counter/tools.rs new file mode 100644 index 00000000000..45d6ab5de69 --- /dev/null +++ b/litellm-rust/crates/core/src/token_counter/tools.rs @@ -0,0 +1,108 @@ +//! Renders tool definitions the way `litellm.token_counter` does before +//! tokenizing them (the TypeScript-like namespace OpenAI appears to use). + +use super::TokenCountError; +use super::types::{EnumValue, FunctionDefinition, Schema, SchemaType, ToolDefinition}; + +pub(super) fn format_function_definitions( + tools: &[ToolDefinition], +) -> Result { + let mut lines = vec!["namespace functions {".to_string(), String::new()]; + for tool in tools { + let function = resolve_function(tool); + let Some(name) = function.name.as_deref().filter(|name| !name.is_empty()) else { + continue; + }; + if let Some(description) = function.description.as_deref().filter(|d| !d.is_empty()) { + lines.push(format!("// {description}")); + } + let parameters = function.parameters.unwrap_or_default(); + match ¶meters.properties { + Some(properties) if !properties.is_empty() => { + lines.push(format!("type {name} = (_: {{")); + lines.push(format_object_parameters(¶meters, 0)?); + lines.push("}) => any;".to_string()); + } + _ => lines.push(format!("type {name} = () => any;")), + } + lines.push(String::new()); + } + lines.push("} // namespace functions".to_string()); + Ok(lines.join("\n")) +} + +fn resolve_function(tool: &ToolDefinition) -> FunctionDefinition { + match &tool.function { + Some(function) => function.clone(), + None => FunctionDefinition { + name: tool.name.clone(), + description: tool.description.clone(), + parameters: tool + .input_schema + .clone() + .or_else(|| tool.parameters.clone()), + }, + } +} + +fn format_object_parameters(parameters: &Schema, indent: usize) -> Result { + let Some(properties) = parameters.properties.as_ref().filter(|p| !p.is_empty()) else { + return Ok(String::new()); + }; + let required = parameters.required.as_deref().unwrap_or_default(); + let mut lines = Vec::new(); + for (key, props) in properties { + if let Some(description) = props.description.as_deref().filter(|d| !d.is_empty()) { + lines.push(format!("// {description}")); + } + let question = if required.iter().any(|r| r == key) { + "" + } else { + "?" + }; + lines.push(format!("{key}{question}: {},", format_type(props, indent)?)); + } + let pad = " ".repeat(indent); + Ok(lines + .iter() + .map(|line| format!("{pad}{line}")) + .collect::>() + .join("\n")) +} + +fn format_type(props: &Schema, indent: usize) -> Result { + let Some(SchemaType::Name(schema_type)) = &props.schema_type else { + return Ok("any".to_string()); + }; + match schema_type.as_str() { + "string" | "integer" | "number" => Ok(match &props.enum_values { + Some(values) => format_enum(values), + None if schema_type == "string" => "string".to_string(), + None => "number".to_string(), + }), + "array" => { + let items = props.items.as_deref().ok_or(TokenCountError::Unsupported( + "array parameter without items".to_string(), + ))?; + Ok(format!("{}[]", format_type(items, indent)?)) + } + "object" => Ok(format!( + "{{\n{}\n}}", + format_object_parameters(props, indent + 2)? + )), + "boolean" => Ok("boolean".to_string()), + "null" => Ok("null".to_string()), + _ => Ok("any".to_string()), + } +} + +fn format_enum(values: &[EnumValue]) -> String { + values + .iter() + .map(|value| match value { + EnumValue::Text(text) => format!("\"{text}\""), + EnumValue::Integer(number) => format!("\"{number}\""), + }) + .collect::>() + .join(" | ") +} diff --git a/litellm-rust/crates/core/src/token_counter/types.rs b/litellm-rust/crates/core/src/token_counter/types.rs new file mode 100644 index 00000000000..e3b58e9327d --- /dev/null +++ b/litellm-rust/crates/core/src/token_counter/types.rs @@ -0,0 +1,122 @@ +use indexmap::IndexMap; +use serde::Deserialize; + +use super::TokenCountError; + +/// The parts of a request body `litellm.token_counter` reads when a host counts +/// input tokens for budget checks. Anything outside this shape is declined so +/// the host can fall back to its own counter instead of silently miscounting. +#[derive(Clone, Debug, Deserialize, PartialEq)] +pub struct CountableRequest { + pub model: String, + pub messages: Option>, + pub tools: Option>, + pub tool_choice: Option, +} + +impl CountableRequest { + pub fn parse(body: &[u8]) -> Result { + serde_json::from_slice(body) + .map_err(|error| TokenCountError::Unsupported(error.to_string())) + } +} + +/// Python counts every string-valued key of a message, so any key beyond these +/// makes the shape unsupported rather than silently uncounted. +#[derive(Clone, Debug, Deserialize, PartialEq)] +#[serde(deny_unknown_fields)] +pub struct Message { + pub role: Option, + pub name: Option, + pub content: Option, +} + +#[derive(Clone, Debug, Deserialize, PartialEq)] +#[serde(untagged)] +pub enum MessageContent { + Text(String), + Blocks(Vec), +} + +#[derive(Clone, Debug, Deserialize, PartialEq)] +#[serde(untagged)] +pub enum ContentItem { + Text(String), + Block(ContentBlock), +} + +#[derive(Clone, Debug, Deserialize, PartialEq)] +#[serde(tag = "type")] +pub enum ContentBlock { + #[serde(rename = "text")] + Text { text: String }, + #[serde(rename = "thinking")] + Thinking { thinking: String }, + #[serde(rename = "tool_reference")] + ToolReference { tool_name: Option }, + /// Images, documents, files and tool use/result blocks price through + /// Python-only helpers, so they stay on the Python counter. + #[serde(other)] + Unsupported, +} + +/// Either the OpenAI `{"type": "function", "function": {...}}` shape or the +/// Anthropic `{"name", "description", "input_schema"}` shape. +#[derive(Clone, Debug, Deserialize, PartialEq)] +pub struct ToolDefinition { + pub function: Option, + pub name: Option, + pub description: Option, + pub input_schema: Option, + pub parameters: Option, +} + +#[derive(Clone, Debug, Deserialize, PartialEq)] +pub struct FunctionDefinition { + pub name: Option, + pub description: Option, + pub parameters: Option, +} + +#[derive(Clone, Debug, Default, Deserialize, PartialEq)] +pub struct Schema { + #[serde(rename = "type")] + pub schema_type: Option, + pub description: Option, + #[serde(rename = "enum")] + pub enum_values: Option>, + pub items: Option>, + pub properties: Option>, + pub required: Option>, +} + +#[derive(Clone, Debug, Deserialize, PartialEq)] +#[serde(untagged)] +pub enum SchemaType { + Name(String), + Union(Vec), +} + +#[derive(Clone, Debug, Deserialize, PartialEq)] +#[serde(untagged)] +pub enum EnumValue { + Text(String), + Integer(i64), +} + +#[derive(Clone, Debug, Deserialize, PartialEq)] +#[serde(untagged)] +pub enum ToolChoice { + Mode(String), + Named(NamedToolChoice), +} + +#[derive(Clone, Debug, Deserialize, PartialEq)] +pub struct NamedToolChoice { + pub function: NamedFunction, +} + +#[derive(Clone, Debug, Deserialize, PartialEq)] +pub struct NamedFunction { + pub name: String, +} diff --git a/litellm-rust/crates/python-bridge/src/execution.rs b/litellm-rust/crates/python-bridge/src/execution.rs index f3648158cf6..b57197b9ddf 100644 --- a/litellm-rust/crates/python-bridge/src/execution.rs +++ b/litellm-rust/crates/python-bridge/src/execution.rs @@ -3,7 +3,6 @@ use std::panic::AssertUnwindSafe; use std::time::Duration; use futures_util::FutureExt; -use litellm_core::error::Error; use litellm_python_interop::{Pythonized, panic_to_pyerr, release_gil}; use pyo3::exceptions::PyRuntimeError; use pyo3::prelude::*; @@ -11,14 +10,15 @@ use serde::Serialize; use tokio::runtime::{Handle, Runtime}; use tokio::time::{self, MissedTickBehavior}; -pub(crate) fn run_sync( +pub(crate) fn run_sync( py: Python<'_>, future: F, - map_error: fn(Error) -> PyErr, + map_error: fn(E) -> PyErr, ) -> PyResult> where T: Serialize + Send + 'static, - F: Future> + Send + 'static, + E: Send + 'static, + F: Future> + Send + 'static, { run_sync_on( py, @@ -28,15 +28,16 @@ where ) } -fn run_sync_on( +fn run_sync_on( py: Python<'_>, runtime: &Runtime, future: F, - map_error: fn(Error) -> PyErr, + map_error: fn(E) -> PyErr, ) -> PyResult> where T: Serialize + Send + 'static, - F: Future> + Send + 'static, + E: Send + 'static, + F: Future> + Send + 'static, { if Handle::try_current().is_ok() { return Err(PyRuntimeError::new_err( @@ -49,14 +50,15 @@ where Pythonized(result).into_pyobject(py).map(Bound::unbind) } -pub(crate) fn run_async( +pub(crate) fn run_async( py: Python<'_>, future: F, - map_error: fn(Error) -> PyErr, + map_error: fn(E) -> PyErr, ) -> PyResult> where T: Serialize + Send + 'static, - F: Future> + Send + 'static, + E: Send + 'static, + F: Future> + Send + 'static, { pyo3_async_runtimes::tokio::future_into_py(py, async move { let result = catch_future_panic(future).await?; @@ -65,7 +67,7 @@ where }) } -fn map_core_result(result: Result, map_error: fn(Error) -> PyErr) -> PyResult { +fn map_core_result(result: Result, map_error: fn(E) -> PyErr) -> PyResult { match result { Ok(value) => Ok(value), Err(error) => Err( @@ -75,9 +77,9 @@ fn map_core_result(result: Result, map_error: fn(Error) -> PyErr) - } } -async fn catch_future_panic(future: F) -> PyResult> +async fn catch_future_panic(future: F) -> PyResult> where - F: Future>, + F: Future>, { AssertUnwindSafe(future) .catch_unwind() @@ -85,9 +87,9 @@ where .map_err(panic_to_pyerr) } -async fn wait_for_sync_result(future: F) -> PyResult> +async fn wait_for_sync_result(future: F) -> PyResult> where - F: Future>, + F: Future>, { let future = catch_future_panic(future); tokio::pin!(future); @@ -114,6 +116,7 @@ mod tests { use std::thread; use std::time::Instant; + use litellm_core::error::Error; use pyo3::panic::PanicException; use pyo3::types::{PyDict, PyModule}; use serde::Serializer; @@ -237,7 +240,7 @@ mod tests { let error = runtime.block_on(async { Python::attach(|py| { - run_sync::(py, async { Ok(true) }, runtime_error) + run_sync::(py, async { Ok(true) }, runtime_error) .expect_err("sync route should reject a nested Tokio runtime") }) }); @@ -273,7 +276,7 @@ mod tests { fn sync_runner_maps_a_panicked_future() { Python::initialize(); Python::attach(|py| { - let error = run_sync::( + let error = run_sync::( py, poll_fn(|_| -> Poll> { panic!("route future panicked") }), runtime_error, @@ -289,7 +292,7 @@ mod tests { fn sync_runner_maps_a_panicked_error_mapper() { Python::initialize(); Python::attach(|py| { - let error = run_sync::( + let error = run_sync::( py, async { Err(Error::InvalidRequest("invalid".to_string())) }, panicking_error_mapper, diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 384f0be5a1b..2d8c93be52e 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -5,6 +5,7 @@ mod execution; mod function_trace; mod marshal; mod routes; +mod token_counter; use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection; use pyo3::prelude::*; @@ -71,6 +72,7 @@ mod _native { super::errors::register(module)?; super::routes::register(module)?; module.add_class::()?; + super::token_counter::register(module)?; super::diagnostics::register(module) } } @@ -106,6 +108,7 @@ mod tests { "chat_completions", "achat_completions", "ResponsesWebSocketConnection", + "TokenCounter", "gil_stats", ]; diff --git a/litellm-rust/crates/python-bridge/src/token_counter.rs b/litellm-rust/crates/python-bridge/src/token_counter.rs new file mode 100644 index 00000000000..2df39be0097 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/token_counter.rs @@ -0,0 +1,64 @@ +use std::sync::Arc; + +use litellm_core::token_counter::types::CountableRequest; +use litellm_core::token_counter::{ + InputTokenCount, TokenCountError, TokenCounter as CoreTokenCounter, +}; +use litellm_python_interop::release_gil; +use pyo3::exceptions::{PyRuntimeError, PyValueError}; +use pyo3::prelude::*; +use pyo3::types::PyAny; + +use crate::errors::RustBridgeDeclined; +use crate::execution::run_async; + +/// Counts the input tokens of a raw request body off the Python event loop with +/// the GIL released. Python owns which requests get here and what to do with +/// the count. +#[pyclass(frozen)] +struct TokenCounter { + inner: Arc, +} + +#[pymethods] +impl TokenCounter { + #[new] + fn new(py: Python<'_>, tokenizer_json: &str) -> PyResult { + let inner = release_gil(py, || CoreTokenCounter::from_json(tokenizer_json)) + .map_err(token_count_error_to_pyerr)?; + Ok(Self { + inner: Arc::new(inner), + }) + } + + fn acount_request<'py>(&self, py: Python<'py>, body: &[u8]) -> PyResult> { + let counter = Arc::clone(&self.inner); + let body = body.to_vec(); + run_async( + py, + async move { + tokio::task::spawn_blocking(move || count_body(&counter, &body)) + .await + .map_err(|error| TokenCountError::Encode(error.to_string()))? + }, + token_count_error_to_pyerr, + ) + } +} + +fn count_body(counter: &CoreTokenCounter, body: &[u8]) -> Result { + let request = CountableRequest::parse(body)?; + counter.count_request(&request) +} + +fn token_count_error_to_pyerr(error: TokenCountError) -> PyErr { + match error { + TokenCountError::Load(message) => PyValueError::new_err(message), + TokenCountError::Unsupported(message) => RustBridgeDeclined::new_err(message), + TokenCountError::Encode(message) => PyRuntimeError::new_err(message), + } +} + +pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { + module.add_class::() +} diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index bfccb703e76..0d7dae7891d 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -91,6 +91,7 @@ from litellm.proxy.common_utils.http_parsing_utils import ( _safe_get_request_query_params, _safe_set_request_parsed_body, populate_request_with_path_params, + read_raw_json_body, ) from litellm.proxy.common_utils.realtime_utils import _realtime_request_body from litellm.proxy.common_utils.user_api_key_cache import ( @@ -2650,6 +2651,7 @@ async def _run_centralized_common_checks( await _reserve_budget_after_common_checks( user_api_key_auth_obj=user_api_key_auth_obj, + request=request, request_data=request_data, route=route, llm_router=llm_router, @@ -2685,6 +2687,7 @@ async def _reserve_budget_after_common_checks( general_settings: dict, end_user_id: str | None = None, end_user_object: LiteLLM_EndUserTable | None = None, + request: Request | None = None, ) -> None: user_api_key_auth_obj.budget_reservation = None if skip_budget_checks: @@ -2710,6 +2713,7 @@ async def _reserve_budget_after_common_checks( end_user_object=end_user_object, apply_user_budget_to_team_keys=general_settings.get("apply_user_budget_to_team_keys") is True, fail_closed_budget_enforcement=general_settings.get("fail_closed_budget_enforcement") is True, + raw_body=await read_raw_json_body(request=request), ) diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index 54a0f18fd63..0bc35887f5b 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -213,6 +213,18 @@ async def _read_request_body(request: Request | None) -> dict: return {} +async def read_raw_json_body(request: Request | None) -> bytes | None: + if request is None or _safe_get_request_parsed_body(request=request) is None: + return None + content_type: Final = _safe_get_request_headers(request=request).get("content-type", "") + if _is_form_content_type(content_type): + return None + try: + return await request.body() + except RuntimeError: + return None + + def _safe_get_request_parsed_body(request: Request | None) -> dict | None: if request is None: return None diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index bf5eadcd85d..b469cd70d76 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -34,6 +34,7 @@ from litellm.proxy.common_utils.user_api_key_cache import ( ) from litellm.proxy.utils import PrismaClient, ProxyLogging from litellm.router import Router +from litellm.rust_bridge.token_counter import count_anthropic_input_tokens, uses_anthropic_tokenizer from litellm.types.proxy.model_access_group_budget import ModelAccessGroupBudget from litellm.types.router import DeploymentTypedDict @@ -210,6 +211,7 @@ async def reserve_budget_for_request( end_user_object: object = None, apply_user_budget_to_team_keys: bool = False, fail_closed_budget_enforcement: bool = False, + raw_body: bytes | None = None, ) -> dict | None: if valid_token is None or not RouteChecks.is_llm_api_route(route=route): return None @@ -237,6 +239,7 @@ async def reserve_budget_for_request( request_body=request_body, route=route, llm_router=llm_router, + raw_body=raw_body, ) current_spend_by_counter_key: Final[dict[str, float]] = {} @@ -1355,24 +1358,46 @@ async def count_request_input_tokens( request_body: dict, route: str, llm_router: Router | None, + raw_body: bytes | None = None, ) -> Mapping[str, int]: """Input-token count per candidate model, counted once per request. Tokenizing is the reservation path's dominant CPU cost and is O(prompt), so counting a large prompt inline stalls every other request on the worker. - Large prompts are counted in a worker thread, and the counts are reused by - both the max-cost and the input-cost estimate. + Models on the Anthropic tokenizer are counted from the raw body by the Rust + bridge when it is enabled, which parses and tokenizes with the GIL released. + Everything it declines is counted in Python, large prompts in a worker + thread. The counts are reused by both the max-cost and the input-cost + estimate. """ models: Final = _get_request_models(request_body=request_body, route=route, llm_router=llm_router) if not models: return MappingProxyType({}) - if _approximate_input_size(request_body) < TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS: - return _count_input_tokens_for_models(request_body=request_body, models=models) - return await asyncio.to_thread( - _count_input_tokens_for_models, - request_body=request_body, - models=models, + rust_count: Final = ( + await count_anthropic_input_tokens(raw_body) + if raw_body is not None and any(uses_anthropic_tokenizer(model) for model in models) + else None ) + rust_counts: Final = MappingProxyType( + { + model: rust_count.input_tokens + for model in models + if rust_count is not None and uses_anthropic_tokenizer(model) + } + ) + python_models: Final = tuple(model for model in models if model not in rust_counts) + if not python_models: + return rust_counts + python_counts: Final = ( + _count_input_tokens_for_models(request_body=request_body, models=python_models) + if _approximate_input_size(request_body) < TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS + else await asyncio.to_thread( + _count_input_tokens_for_models, + request_body=request_body, + models=python_models, + ) + ) + return MappingProxyType({**rust_counts, **python_counts}) def _count_input_tokens_for_models( diff --git a/litellm/rust_bridge/token_counter.py b/litellm/rust_bridge/token_counter.py new file mode 100644 index 00000000000..618ec7610ed --- /dev/null +++ b/litellm/rust_bridge/token_counter.py @@ -0,0 +1,79 @@ +"""Thin Python wrapper for the native Rust input token counter.""" + +from __future__ import annotations + +from collections.abc import Awaitable +from dataclasses import dataclass +from functools import lru_cache +from typing import Final, Protocol, cast # noqa: TID251 # native extension exposes dynamically typed callables + +from pydantic import TypeAdapter + +import litellm +from litellm._logging import verbose_logger +from litellm.rust_bridge.bindings import NativeBinding +from litellm.rust_bridge.configuration import rust_enabled +from litellm.rust_bridge.runtime import BridgeErrorContext, RustHandled, aattempt + + +class RustTokenCounter(Protocol): + def acount_request(self, body: bytes) -> Awaitable[object]: + raise NotImplementedError + + +class RustTokenCounterFactory(Protocol): + def __call__(self, tokenizer_json: str) -> RustTokenCounter: + raise NotImplementedError + + +@dataclass(frozen=True, slots=True) +class InputTokenCount: + model: str + input_tokens: int + + +_INPUT_TOKEN_COUNT: Final = TypeAdapter(InputTokenCount) + + +def _as_factory(value: object) -> RustTokenCounterFactory | None: + return ( + cast( # cast-ok: native extension protocol is runtime-defined + RustTokenCounterFactory, value + ) + if callable(value) + else None + ) + + +TOKEN_COUNTER: Final = NativeBinding("TokenCounter", validate=_as_factory) + + +def uses_anthropic_tokenizer(model: str) -> bool: + if litellm.disable_hf_tokenizer_download is True: + return False + return model in litellm.anthropic_models and "claude-3" not in model + + +@lru_cache(maxsize=4) +def _anthropic_counter(factory: RustTokenCounterFactory) -> RustTokenCounter: + from litellm.utils import claude_json_str + + return factory(claude_json_str) + + +async def count_anthropic_input_tokens(body: bytes) -> InputTokenCount | None: + if not rust_enabled(): + return None + factory: Final = TOKEN_COUNTER.load() + if factory is None: + return None + try: + attempt: Final = await aattempt( + native_call=lambda: _anthropic_counter(factory).acount_request(body), + adapt=_INPUT_TOKEN_COUNT.validate_python, + context=BridgeErrorContext(route="token_counter", provider="anthropic", model=""), + ) + except (RuntimeError, ValueError) as error: + verbose_logger.debug("Rust token counter failed, counting in Python: %s", error) + return None + return attempt.value if isinstance(attempt, RustHandled) else None diff --git a/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py b/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py index 011571a37e0..bc4e756eb65 100644 --- a/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py +++ b/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py @@ -26,9 +26,58 @@ from litellm.proxy.common_utils.http_parsing_utils import ( get_tags_from_request_body, numeric_form_fields, populate_request_with_path_params, + read_raw_json_body, ) +def _starlette_request(body: bytes, content_type: str) -> Request: + scope = { + "type": "http", + "method": "POST", + "path": "/v1/messages", + "headers": [(b"content-type", content_type.encode())], + "query_string": b"", + } + chunks = iter((body,)) + + async def receive(): + return {"type": "http.request", "body": next(chunks, b""), "more_body": False} + + return Request(scope, receive) + + +@pytest.mark.asyncio +async def test_read_raw_json_body_returns_the_bytes_the_parsed_body_came_from(): + body = b'{"model": "claude-sonnet-4-5", "messages": [{"role": "user", "content": "hi"}]}' + request = _starlette_request(body, "application/json") + + assert await _read_request_body(request) == orjson.loads(body) + assert await read_raw_json_body(request) == body + + +@pytest.mark.asyncio +async def test_read_raw_json_body_is_none_until_the_body_has_been_parsed(): + request = _starlette_request(b'{"model": "claude-sonnet-4-5"}', "application/json") + + assert await read_raw_json_body(request) is None + assert await read_raw_json_body(None) is None + + +@pytest.mark.asyncio +async def test_read_raw_json_body_is_none_for_form_bodies(): + request = _starlette_request(b"model=claude-sonnet-4-5", "application/x-www-form-urlencoded") + + assert await _read_request_body(request) == {"model": "claude-sonnet-4-5"} + assert await read_raw_json_body(request) is None + + +@pytest.mark.asyncio +async def test_read_raw_json_body_is_none_for_a_request_that_only_mocks_the_parsed_body_path(): + mock_request = MagicMock() + + assert await read_raw_json_body(mock_request) is None + + @pytest.mark.asyncio async def test_request_body_caching(): """ diff --git a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py index de6c7c2a40a..9918588a1b6 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py +++ b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py @@ -1,13 +1,21 @@ +import json from typing import Final import pytest -import litellm.proxy.proxy_server as proxy_server +import litellm from litellm.caching import DualCache +from litellm.proxy import proxy_server from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache -from litellm.proxy.spend_tracking.budget_reservation import estimate_request_max_cost, reserve_budget_for_request +from litellm.proxy.spend_tracking.budget_reservation import ( + count_request_input_tokens, + estimate_request_max_cost, + reserve_budget_for_request, +) from litellm.proxy.utils import ProxyLogging +from litellm.rust_bridge import bindings, configuration +from litellm.rust_bridge import token_counter as rust_token_counter TOKEN_COUNTING_ROUTES: Final = ( "/responses/input_tokens", @@ -139,3 +147,115 @@ def test_bedrock_converse_body_reserves_the_prompt_not_the_context_window(): ) assert converse_cost is not None and invoke_cost is not None assert invoke_cost < converse_cost < 2 * invoke_cost + + +ANTHROPIC_TOKENIZER_MODEL: Final = "claude-sonnet-4-5-20250929" +RUST_COUNTED_BODY: Final = {"model": ANTHROPIC_TOKENIZER_MODEL, "max_tokens": 16, "messages": ANTHROPIC_MESSAGES} +RUST_INPUT_TOKENS: Final = 4_321 + + +class _FakeDeclined(Exception): + pass + + +class _FakeUpstream(Exception): + pass + + +class _FakeNative: + RustBridgeDeclined = _FakeDeclined + RustUpstreamError = _FakeUpstream + + +class _RecordingCounter: + bodies: Final[list[bytes]] = [] + + def __init__(self, tokenizer_json: str) -> None: + pass + + async def acount_request(self, body: bytes) -> object: + self.bodies.append(body) + return {"model": ANTHROPIC_TOKENIZER_MODEL, "input_tokens": RUST_INPUT_TOKENS} + + +class _DecliningCounter: + def __init__(self, tokenizer_json: str) -> None: + pass + + async def acount_request(self, body: bytes) -> object: + raise _FakeDeclined("unsupported content block") + + +@pytest.fixture +def rust_counter(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(bindings, "get_native_bridge", lambda: _FakeNative()) + rust_token_counter._anthropic_counter.cache_clear() + configuration.reset_rust_configuration() + _RecordingCounter.bodies.clear() + yield + rust_token_counter.TOKEN_COUNTER.reset() + rust_token_counter._anthropic_counter.cache_clear() + configuration.reset_rust_configuration() + + +@pytest.mark.asyncio +async def test_rust_count_replaces_python_tokenizing_for_anthropic_models(rust_counter: None) -> None: + litellm.rust(True) + rust_token_counter.TOKEN_COUNTER.override(_RecordingCounter) + raw_body: Final = json.dumps(RUST_COUNTED_BODY).encode() + + counts: Final = await count_request_input_tokens( + request_body=RUST_COUNTED_BODY, route="/v1/messages", llm_router=None, raw_body=raw_body + ) + + assert dict(counts) == {ANTHROPIC_TOKENIZER_MODEL: RUST_INPUT_TOKENS} + assert _RecordingCounter.bodies == [raw_body] + + +@pytest.mark.asyncio +async def test_rust_decline_falls_back_to_python_count(rust_counter: None) -> None: + litellm.rust(True) + rust_token_counter.TOKEN_COUNTER.override(_DecliningCounter) + python_counts: Final = await count_request_input_tokens( + request_body=RUST_COUNTED_BODY, route="/v1/messages", llm_router=None + ) + + counts: Final = await count_request_input_tokens( + request_body=RUST_COUNTED_BODY, + route="/v1/messages", + llm_router=None, + raw_body=json.dumps(RUST_COUNTED_BODY).encode(), + ) + + assert dict(counts) == dict(python_counts) + assert counts[ANTHROPIC_TOKENIZER_MODEL] != RUST_INPUT_TOKENS + + +@pytest.mark.asyncio +async def test_disabled_rust_never_sees_the_raw_body(rust_counter: None) -> None: + litellm.rust(False) + rust_token_counter.TOKEN_COUNTER.override(_RecordingCounter) + + counts: Final = await count_request_input_tokens( + request_body=RUST_COUNTED_BODY, + route="/v1/messages", + llm_router=None, + raw_body=json.dumps(RUST_COUNTED_BODY).encode(), + ) + + assert _RecordingCounter.bodies == [] + assert counts[ANTHROPIC_TOKENIZER_MODEL] != RUST_INPUT_TOKENS + + +@pytest.mark.asyncio +async def test_non_anthropic_tokenizer_models_stay_in_python(rust_counter: None) -> None: + litellm.rust(True) + rust_token_counter.TOKEN_COUNTER.override(_RecordingCounter) + body: Final = {"model": "gpt-4o", "messages": ANTHROPIC_MESSAGES} + + counts: Final = await count_request_input_tokens( + request_body=body, route="/v1/chat/completions", llm_router=None, raw_body=json.dumps(body).encode() + ) + + assert _RecordingCounter.bodies == [] + assert counts["gpt-4o"] != RUST_INPUT_TOKENS diff --git a/tests/test_litellm/rust_bridge/test_token_counter.py b/tests/test_litellm/rust_bridge/test_token_counter.py new file mode 100644 index 00000000000..4d29251dca0 --- /dev/null +++ b/tests/test_litellm/rust_bridge/test_token_counter.py @@ -0,0 +1,224 @@ +"""Tests for the Rust input token counter bridge. + +The native factory is dependency-injected through ``TOKEN_COUNTER.override`` +so the fallback cases run without the compiled extension present. The parity +cases need the extension and are skipped when it is not built. +""" + +from __future__ import annotations + +import json +from typing import Final + +import pytest + +import litellm +from litellm.rust_bridge import bindings, configuration +from litellm.rust_bridge import token_counter as bridge + +MODEL: Final = "claude-sonnet-4-5-20250929" +BODY: Final = json.dumps({"model": MODEL, "messages": [{"role": "user", "content": "hello"}]}).encode() + + +class _FakeDeclined(Exception): + pass + + +class _FakeUpstream(Exception): + pass + + +class _FakeNative: + RustBridgeDeclined = _FakeDeclined + RustUpstreamError = _FakeUpstream + + +class _RecordingCounter: + def __init__(self, tokenizer_json: str) -> None: + self.tokenizer_json = tokenizer_json + self.bodies: list[bytes] = [] + + async def acount_request(self, body: bytes) -> object: + self.bodies.append(body) + return {"model": MODEL, "input_tokens": 42} + + +class _DecliningCounter: + def __init__(self, tokenizer_json: str) -> None: + pass + + async def acount_request(self, body: bytes) -> object: + raise _FakeDeclined("request has no messages") + + +class _FailingCounter: + def __init__(self, tokenizer_json: str) -> None: + pass + + async def acount_request(self, body: bytes) -> object: + raise RuntimeError("encode failed") + + +@pytest.fixture(autouse=True) +def _reset_bridge(monkeypatch: pytest.MonkeyPatch): + bridge.TOKEN_COUNTER.reset() + bridge._anthropic_counter.cache_clear() + configuration.reset_rust_configuration() + monkeypatch.setattr(bindings, "get_native_bridge", lambda: _FakeNative()) + yield + bridge.TOKEN_COUNTER.reset() + bridge._anthropic_counter.cache_clear() + configuration.reset_rust_configuration() + + +@pytest.mark.asyncio +async def test_disabled_bridge_never_constructs_a_counter() -> None: + constructed: list[str] = [] + + def factory(tokenizer_json: str) -> _RecordingCounter: + constructed.append(tokenizer_json) + return _RecordingCounter(tokenizer_json) + + litellm.rust(False) + bridge.TOKEN_COUNTER.override(factory) + + assert await bridge.count_anthropic_input_tokens(BODY) is None + assert constructed == [] + + +@pytest.mark.asyncio +async def test_enabled_bridge_returns_typed_count_and_reuses_one_counter() -> None: + counters: list[_RecordingCounter] = [] + + def factory(tokenizer_json: str) -> _RecordingCounter: + counter = _RecordingCounter(tokenizer_json) + counters.append(counter) + return counter + + litellm.rust(True) + bridge.TOKEN_COUNTER.override(factory) + + first: Final = await bridge.count_anthropic_input_tokens(BODY) + second: Final = await bridge.count_anthropic_input_tokens(BODY) + + assert first == bridge.InputTokenCount(model=MODEL, input_tokens=42) + assert second == first + assert len(counters) == 1 + assert counters[0].bodies == [BODY, BODY] + assert json.loads(counters[0].tokenizer_json)["model"]["type"] == "BPE" + + +@pytest.mark.asyncio +async def test_missing_native_module_falls_back(monkeypatch: pytest.MonkeyPatch) -> None: + litellm.rust(True) + monkeypatch.setattr(bindings, "get_native_bridge", lambda: None) + + assert await bridge.count_anthropic_input_tokens(BODY) is None + + +@pytest.mark.asyncio +async def test_declined_request_falls_back() -> None: + litellm.rust(True) + bridge.TOKEN_COUNTER.override(_DecliningCounter) + + assert await bridge.count_anthropic_input_tokens(BODY) is None + + +@pytest.mark.asyncio +async def test_runtime_failure_falls_back() -> None: + litellm.rust(True) + bridge.TOKEN_COUNTER.override(_FailingCounter) + + assert await bridge.count_anthropic_input_tokens(BODY) is None + + +@pytest.mark.parametrize( + ("model", "expected"), + ((MODEL, True), ("claude-3-5-sonnet-20241022", False), ("gpt-4o", False), ("my-router-alias", False)), +) +def test_uses_anthropic_tokenizer_mirrors_python_tokenizer_selection(model: str, expected: bool) -> None: + assert bridge.uses_anthropic_tokenizer(model) is expected + + +def test_uses_anthropic_tokenizer_respects_hf_download_opt_out(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "disable_hf_tokenizer_download", True) + + assert bridge.uses_anthropic_tokenizer(MODEL) is False + + +PARITY_REQUESTS: Final[tuple[dict[str, object], ...]] = ( + {"model": MODEL, "messages": [{"role": "user", "content": "Hello, how are you today?"}]}, + { + "model": MODEL, + "messages": [ + {"role": "system", "content": "You are terse."}, + {"role": "user", "name": "bob", "content": [{"type": "text", "text": "Summarize this."}]}, + {"role": "assistant", "content": "Sure."}, + ], + }, + { + "model": MODEL, + "messages": [{"role": "user", "content": "weather in sf?"}], + "tools": [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather", + "parameters": { + "type": "object", + "properties": { + "city": {"type": "string", "description": "City"}, + "unit": {"type": "string", "enum": ["c", "f"]}, + }, + "required": ["city"], + }, + }, + } + ], + "tool_choice": {"type": "function", "function": {"name": "get_weather"}}, + }, + { + "model": MODEL, + "messages": [{"role": "user", "content": "x " * 20_000}], + }, +) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("request_body", PARITY_REQUESTS) +async def test_native_count_matches_python_token_counter( + monkeypatch: pytest.MonkeyPatch, request_body: dict[str, object] +) -> None: + native: Final = pytest.importorskip("litellm.rust_bridge._native") + monkeypatch.setattr(bindings, "get_native_bridge", lambda: native) + litellm.rust(True) + + rust_count: Final = await bridge.count_anthropic_input_tokens(json.dumps(request_body).encode()) + python_count: Final = litellm.token_counter( + model=MODEL, + messages=request_body["messages"], + tools=request_body.get("tools"), + tool_choice=request_body.get("tool_choice"), + ) + + assert rust_count is not None + assert rust_count.model == MODEL + assert rust_count.input_tokens == python_count + + +@pytest.mark.asyncio +async def test_native_declines_image_content(monkeypatch: pytest.MonkeyPatch) -> None: + native: Final = pytest.importorskip("litellm.rust_bridge._native") + monkeypatch.setattr(bindings, "get_native_bridge", lambda: native) + litellm.rust(True) + body: Final = json.dumps( + { + "model": MODEL, + "messages": [ + {"role": "user", "content": [{"type": "image_url", "image_url": {"url": "data:image/png;base64,AA"}}]} + ], + } + ).encode() + + assert await bridge.count_anthropic_input_tokens(body) is None