refactor(rust): rescope /v1/messages admission to a Python-called Rust token counter

Drop the gateway-side admission layer (identity cache, /key/info lookups,
process-local budget/TPM/RPM limits) and keep only the CPU work in Rust: a
typed litellm-core token counter that parses the raw body once and counts
input tokens with the GIL released, exposed as TokenCounter in the PyO3
bridge. The Python proxy's existing auth dependency passes the raw body
into budget reservation, which uses the Rust count for models on the
Anthropic tokenizer and falls back to Python for anything Rust declines,
when the bridge is disabled, or when the native module is unavailable.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yassin 2026-09-09 19:19:54 +00:00
parent 09badc21ab
commit 5605b6d0eb
32 changed files with 1208 additions and 1143 deletions

View file

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

View file

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

View file

@ -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<AppState> for Admit {
type Rejection = Response;
async fn from_request(request: Request, state: &AppState) -> Result<Self, Self::Rejection> {
let (parts, body) = request.into_parts();
let admission: &Admission = &state.admission;
let raw = to_bytes(body, admission.max_request_bytes())
.await
.map_err(|error| (StatusCode::PAYLOAD_TOO_LARGE, error.to_string()).into_response())?;
admission
.admit(bearer_token(&parts.headers), &raw)
.await
.map(Admit)
.map_err(|rejection| reject(&rejection))
}
}
fn reject(rejection: &Rejection) -> Response {
let status = match rejection {
Rejection::Unauthorized => StatusCode::UNAUTHORIZED,
Rejection::IdentityUnavailable(_) => StatusCode::SERVICE_UNAVAILABLE,
Rejection::ModelNotAllowed(_) => StatusCode::FORBIDDEN,
Rejection::RequestTooLarge { .. } => StatusCode::PAYLOAD_TOO_LARGE,
Rejection::ContextTooLarge { .. } | Rejection::InvalidRequest(_) => StatusCode::BAD_REQUEST,
Rejection::LimitExceeded(_) => StatusCode::TOO_MANY_REQUESTS,
Rejection::Tokenizer(_) => StatusCode::INTERNAL_SERVER_ERROR,
};
(
status,
axum::Json(serde_json::json!({"error": {"message": rejection.to_string()}})),
)
.into_response()
}

View file

@ -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<String>,
#[serde(default)]
pub max_budget: Option<f64>,
#[serde(default)]
pub spend: f64,
#[serde(default)]
pub tpm_limit: Option<u64>,
#[serde(default)]
pub rpm_limit: Option<u64>,
}
impl KeyLimits {
pub fn allows_model(&self, model: &str) -> bool {
self.models.is_empty() || self.models.iter().any(|allowed| allowed == model)
}
}
#[derive(Clone, Debug, PartialEq)]
pub enum Identity {
Master,
VirtualKey {
key_hash: String,
limits: Arc<KeyLimits>,
},
}
impl Identity {
pub fn key_hash(&self) -> &str {
match self {
Identity::Master => "litellm_proxy_master_key",
Identity::VirtualKey { key_hash, .. } => key_hash,
}
}
}
#[derive(Debug, thiserror::Error, PartialEq)]
pub enum IdentityError {
#[error("missing or invalid bearer token")]
Unauthorized,
#[error("key lookup failed: {0}")]
LookupFailed(String),
}
#[derive(Deserialize)]
struct KeyInfoResponse {
info: KeyLimits,
}
/// Resolves bearer tokens; misses go to the proxy's `/key/info`.
pub struct IdentityCache {
master_key: Option<Arc<str>>,
proxy_base_url: String,
http: reqwest::Client,
cache: RwLock<HashMap<String, Arc<KeyLimits>>>,
}
impl IdentityCache {
pub fn new(master_key: Option<Arc<str>>, proxy_base_url: String) -> Self {
Self {
master_key,
proxy_base_url: proxy_base_url.trim_end_matches('/').to_string(),
http: reqwest::Client::builder()
.timeout(Duration::from_secs(KEY_INFO_TIMEOUT_SECS))
.build()
.unwrap_or_default(),
cache: RwLock::new(HashMap::new()),
}
}
/// Seed the cache, so tests and offline hosts never call the proxy.
pub fn insert(&self, token: &str, limits: KeyLimits) {
if let Ok(mut cache) = self.cache.write() {
cache.insert(hash_token(token), Arc::new(limits));
}
}
pub async fn resolve(&self, token: Option<&str>) -> Result<Identity, IdentityError> {
let Some(token) = token.map(str::trim).filter(|token| !token.is_empty()) else {
return Err(IdentityError::Unauthorized);
};
if let Some(master) = self.master_key.as_deref()
&& bool::from(token.as_bytes().ct_eq(master.as_bytes()))
{
return Ok(Identity::Master);
}
let key_hash = hash_token(token);
let cached = self
.cache
.read()
.ok()
.and_then(|cache| cache.get(&key_hash).cloned());
let limits = match cached {
Some(limits) => limits,
None => {
let limits = Arc::new(self.fetch(token).await?);
if let Ok(mut cache) = self.cache.write() {
cache.insert(key_hash.clone(), Arc::clone(&limits));
}
limits
}
};
Ok(Identity::VirtualKey { key_hash, limits })
}
async fn fetch(&self, token: &str) -> Result<KeyLimits, IdentityError> {
let Some(master) = self.master_key.as_deref() else {
return Err(IdentityError::Unauthorized);
};
let response = self
.http
.get(format!("{}{PROXY_KEY_INFO_PATH}", self.proxy_base_url))
.query(&[("key", token)])
.bearer_auth(master)
.send()
.await
.map_err(|error| IdentityError::LookupFailed(error.without_url().to_string()))?;
match response.status().as_u16() {
200 => response
.json::<KeyInfoResponse>()
.await
.map(|body| body.info)
.map_err(|error| IdentityError::LookupFailed(error.without_url().to_string())),
400..=404 => Err(IdentityError::Unauthorized),
status => Err(IdentityError::LookupFailed(format!(
"proxy answered {status}"
))),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn cache() -> IdentityCache {
IdentityCache::new(
Some(Arc::from("sk-master")),
"http://127.0.0.1:1".to_string(),
)
}
#[tokio::test]
async fn master_key_is_unlimited_and_never_looked_up() {
assert_eq!(
cache().resolve(Some("sk-master")).await,
Ok(Identity::Master)
);
}
#[tokio::test]
async fn missing_token_is_unauthorized() {
assert_eq!(
cache().resolve(None).await,
Err(IdentityError::Unauthorized)
);
assert_eq!(
cache().resolve(Some(" ")).await,
Err(IdentityError::Unauthorized)
);
}
#[tokio::test]
async fn seeded_virtual_key_resolves_from_cache_without_network() {
let cache = cache();
let limits = KeyLimits {
models: vec!["claude".to_string()],
max_budget: Some(10.0),
spend: 1.5,
tpm_limit: Some(1000),
rpm_limit: None,
};
cache.insert("sk-virtual", limits.clone());
let identity = cache.resolve(Some("sk-virtual")).await.unwrap();
match identity {
Identity::VirtualKey {
key_hash,
limits: resolved,
} => {
assert_eq!(key_hash, hash_token("sk-virtual"));
assert_eq!(*resolved, limits);
}
Identity::Master => panic!("virtual key resolved as master"),
}
}
#[tokio::test]
async fn unknown_key_lookup_failure_is_reported_not_admitted() {
let error = cache().resolve(Some("sk-unknown")).await.unwrap_err();
assert!(matches!(error, IdentityError::LookupFailed(_)), "{error:?}");
}
#[test]
fn empty_model_list_allows_every_model() {
assert!(KeyLimits::default().allows_model("anything"));
let limits = KeyLimits {
models: vec!["a".to_string()],
..KeyLimits::default()
};
assert!(limits.allows_model("a"));
assert!(!limits.allows_model("b"));
}
}

View file

@ -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<HashMap<String, KeyWindow>>,
}
impl Limits {
/// Admit `input_tokens` for the identity, or say which limit it would cross.
pub fn reserve(&self, identity: &Identity, input_tokens: usize) -> Result<(), LimitExceeded> {
let Identity::VirtualKey { key_hash, limits } = identity else {
return Ok(());
};
let Ok(mut windows) = self.windows.lock() else {
return Ok(());
};
let now = Instant::now();
let window = windows.entry(key_hash.clone()).or_insert(KeyWindow {
started: now,
tokens: 0,
requests: 0,
reserved_spend: 0.0,
});
if now.duration_since(window.started) >= WINDOW {
window.started = now;
window.tokens = 0;
window.requests = 0;
}
let tokens = input_tokens as u64;
let cost = input_tokens as f64 * DEFAULT_INPUT_COST_PER_TOKEN;
check(limits, window, tokens, cost)?;
window.tokens += tokens;
window.requests += 1;
window.reserved_spend += cost;
Ok(())
}
}
fn check(
limits: &KeyLimits,
window: &KeyWindow,
tokens: u64,
cost: f64,
) -> Result<(), LimitExceeded> {
if let Some(max_budget) = limits.max_budget
&& limits.spend + window.reserved_spend + cost > max_budget
{
return Err(LimitExceeded::Budget);
}
if let Some(tpm) = limits.tpm_limit
&& window.tokens + tokens > tpm
{
return Err(LimitExceeded::TokensPerMinute);
}
if let Some(rpm) = limits.rpm_limit
&& window.requests + 1 > rpm
{
return Err(LimitExceeded::RequestsPerMinute);
}
Ok(())
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use super::*;
fn key(limits: KeyLimits) -> Identity {
Identity::VirtualKey {
key_hash: "hash".to_string(),
limits: Arc::new(limits),
}
}
#[test]
fn master_key_is_never_limited() {
let limits = Limits::default();
assert_eq!(limits.reserve(&Identity::Master, usize::MAX), Ok(()));
}
#[test]
fn tpm_window_accumulates_and_rejects_on_overflow() {
let limits = Limits::default();
let identity = key(KeyLimits {
tpm_limit: Some(100),
..KeyLimits::default()
});
assert_eq!(limits.reserve(&identity, 60), Ok(()));
assert_eq!(
limits.reserve(&identity, 50),
Err(LimitExceeded::TokensPerMinute)
);
assert_eq!(limits.reserve(&identity, 40), Ok(()));
}
#[test]
fn rpm_counts_requests() {
let limits = Limits::default();
let identity = key(KeyLimits {
rpm_limit: Some(2),
..KeyLimits::default()
});
assert_eq!(limits.reserve(&identity, 1), Ok(()));
assert_eq!(limits.reserve(&identity, 1), Ok(()));
assert_eq!(
limits.reserve(&identity, 1),
Err(LimitExceeded::RequestsPerMinute)
);
}
#[test]
fn budget_includes_prior_spend_and_local_reservations() {
let limits = Limits::default();
let identity = key(KeyLimits {
max_budget: Some(1.0),
spend: 0.5,
..KeyLimits::default()
});
let tokens_for_quarter_dollar = (0.25 / DEFAULT_INPUT_COST_PER_TOKEN) as usize;
assert_eq!(limits.reserve(&identity, tokens_for_quarter_dollar), Ok(()));
assert_eq!(limits.reserve(&identity, tokens_for_quarter_dollar), Ok(()));
assert_eq!(
limits.reserve(&identity, tokens_for_quarter_dollar),
Err(LimitExceeded::Budget)
);
}
#[test]
fn a_rejected_request_reserves_nothing() {
let limits = Limits::default();
let identity = key(KeyLimits {
tpm_limit: Some(10),
rpm_limit: Some(5),
..KeyLimits::default()
});
assert_eq!(
limits.reserve(&identity, 11),
Err(LimitExceeded::TokensPerMinute)
);
assert_eq!(limits.reserve(&identity, 10), Ok(()));
}
}

View file

@ -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<IdentityError> for Rejection {
fn from(error: IdentityError) -> Self {
match error {
IdentityError::Unauthorized => Rejection::Unauthorized,
IdentityError::LookupFailed(reason) => Rejection::IdentityUnavailable(reason),
}
}
}
/// The request after the single parse: what admission needs plus the body to forward.
#[derive(Debug)]
pub struct ParsedRequest {
pub body: Value,
pub model: String,
/// Everything the tokenizer sees: system prompt, message content, tool schemas.
pub text: String,
}
impl ParsedRequest {
pub fn parse(raw: &[u8]) -> Result<Self, Rejection> {
let body: Value = serde_json::from_slice(raw)
.map_err(|error| Rejection::InvalidRequest(format!("body is not JSON: {error}")))?;
let Some(object) = body.as_object() else {
return Err(Rejection::InvalidRequest(
"body must be a JSON object".to_string(),
));
};
let model = object
.get("model")
.and_then(Value::as_str)
.map(str::trim)
.filter(|model| !model.is_empty())
.ok_or_else(|| Rejection::InvalidRequest("body requires a model".to_string()))?
.to_string();
let mut text = String::with_capacity(raw.len());
if let Some(system) = object.get("system") {
push_content(system, &mut text);
}
for message in object
.get("messages")
.and_then(Value::as_array)
.into_iter()
.flatten()
{
if let Some(content) = message.get("content") {
push_content(content, &mut text);
}
}
if let Some(tools) = object.get("tools") {
text.push_str(&tools.to_string());
}
Ok(Self { body, model, text })
}
}
fn push_content(content: &Value, text: &mut String) {
match content {
Value::String(value) => {
text.push_str(value);
text.push('\n');
}
Value::Array(blocks) => {
for block in blocks {
match block.get("text").and_then(Value::as_str) {
Some(value) => {
text.push_str(value);
text.push('\n');
}
None => text.push_str(&block.to_string()),
}
}
}
other => text.push_str(&other.to_string()),
}
}
fn env_usize(name: &str, default: usize) -> usize {
std::env::var(name)
.ok()
.and_then(|value| value.trim().parse().ok())
.unwrap_or(default)
}
#[derive(Debug)]
pub struct Admitted {
pub body: Value,
pub identity: Identity,
pub input_tokens: usize,
pub elapsed_ms: f64,
}
pub struct Admission {
identities: IdentityCache,
limits: Limits,
tokens: TokenCounter,
max_request_bytes: usize,
max_input_tokens: usize,
}
impl Admission {
pub fn new(identities: IdentityCache, tokens: TokenCounter) -> Self {
Self {
identities,
limits: Limits::default(),
tokens,
max_request_bytes: DEFAULT_MAX_REQUEST_BYTES,
max_input_tokens: DEFAULT_MAX_INPUT_TOKENS,
}
}
/// Build from `LITELLM_PROXY_BASE_URL`, `LITELLM_ANTHROPIC_TOKENIZER_PATH`,
/// `LITELLM_TOKENIZER_CONCURRENCY`, `LITELLM_MAX_REQUEST_BYTES` and `LITELLM_MAX_INPUT_TOKENS`.
/// Without a tokenizer path the token count is approximated from the input length.
pub fn from_env(master_key: Option<Arc<str>>) -> Result<Self, TokenizerError> {
let proxy_base_url = std::env::var("LITELLM_PROXY_BASE_URL")
.ok()
.map(|url| url.trim().to_string())
.filter(|url| !url.is_empty())
.unwrap_or_else(|| DEFAULT_PROXY_BASE_URL.to_string());
let tokens = match std::env::var("LITELLM_ANTHROPIC_TOKENIZER_PATH") {
Ok(path) if !path.trim().is_empty() => TokenCounter::from_file(
Path::new(path.trim()),
env_usize(
"LITELLM_TOKENIZER_CONCURRENCY",
DEFAULT_TOKENIZER_CONCURRENCY,
),
)?,
_ => TokenCounter::approximate(),
};
Ok(
Self::new(IdentityCache::new(master_key, proxy_base_url), tokens).with_limits(
env_usize("LITELLM_MAX_REQUEST_BYTES", DEFAULT_MAX_REQUEST_BYTES),
env_usize("LITELLM_MAX_INPUT_TOKENS", DEFAULT_MAX_INPUT_TOKENS),
),
)
}
pub fn with_limits(self, max_request_bytes: usize, max_input_tokens: usize) -> Self {
Self {
max_request_bytes,
max_input_tokens,
..self
}
}
pub fn identities(&self) -> &IdentityCache {
&self.identities
}
pub fn tokens(&self) -> &TokenCounter {
&self.tokens
}
pub fn max_request_bytes(&self) -> usize {
self.max_request_bytes
}
pub async fn admit(&self, bearer: Option<&str>, raw: &[u8]) -> Result<Admitted, Rejection> {
let started = Instant::now();
let bearer = bearer
.map(str::trim)
.filter(|token| !token.is_empty())
.ok_or(Rejection::Unauthorized)?;
if raw.len() > self.max_request_bytes {
return Err(Rejection::RequestTooLarge {
bytes: raw.len(),
max: self.max_request_bytes,
});
}
let ParsedRequest { body, model, text } = ParsedRequest::parse(raw)?;
let (identity, input_tokens) = tokio::join!(
self.identities.resolve(Some(bearer)),
self.tokens.count(text)
);
let identity = identity?;
let input_tokens = input_tokens
.map_err(|error: TokenizerError| Rejection::Tokenizer(error.to_string()))?;
if let Identity::VirtualKey { limits, .. } = &identity
&& !limits.allows_model(&model)
{
return Err(Rejection::ModelNotAllowed(model));
}
if input_tokens > self.max_input_tokens {
return Err(Rejection::ContextTooLarge {
tokens: input_tokens,
max: self.max_input_tokens,
});
}
self.limits
.reserve(&identity, input_tokens)
.map_err(Rejection::LimitExceeded)?;
Ok(Admitted {
body,
identity,
input_tokens,
elapsed_ms: started.elapsed().as_secs_f64() * 1000.0,
})
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use serde_json::json;
use super::*;
fn admission() -> Admission {
let identities = IdentityCache::new(
Some(Arc::from("sk-master")),
"http://127.0.0.1:1".to_string(),
);
identities.insert(
"sk-limited",
KeyLimits {
models: vec!["claude".to_string()],
max_budget: None,
spend: 0.0,
tpm_limit: Some(100),
rpm_limit: None,
},
);
Admission::new(identities, TokenCounter::approximate())
}
fn body(model: &str, words: usize) -> Vec<u8> {
json!({
"model": model,
"max_tokens": 16,
"messages": [{"role": "user", "content": "word ".repeat(words)}]
})
.to_string()
.into_bytes()
}
#[test]
fn parse_collects_system_messages_and_tools_once() {
let raw = json!({
"model": "claude",
"system": [{"type": "text", "text": "be brief"}],
"messages": [
{"role": "user", "content": "hi"},
{"role": "assistant", "content": [{"type": "text", "text": "hello"}]}
],
"tools": [{"name": "t", "input_schema": {"type": "object"}}]
})
.to_string();
let parsed = ParsedRequest::parse(raw.as_bytes()).unwrap();
assert_eq!(parsed.model, "claude");
assert!(parsed.text.contains("be brief\n"));
assert!(parsed.text.contains("hi\n"));
assert!(parsed.text.contains("hello\n"));
assert!(parsed.text.contains("input_schema"));
assert_eq!(parsed.body["messages"][0]["content"], "hi");
}
#[test]
fn parse_rejects_non_object_and_missing_model() {
assert!(matches!(
ParsedRequest::parse(b"[]"),
Err(Rejection::InvalidRequest(_))
));
assert!(matches!(
ParsedRequest::parse(br#"{"messages": []}"#),
Err(Rejection::InvalidRequest(_))
));
assert!(matches!(
ParsedRequest::parse(b"{not json"),
Err(Rejection::InvalidRequest(_))
));
}
#[tokio::test]
async fn master_key_is_admitted_with_a_token_count() {
let admitted = admission()
.admit(Some("sk-master"), &body("claude", 40))
.await
.unwrap();
assert_eq!(admitted.identity, Identity::Master);
assert!(admitted.input_tokens > 0);
assert_eq!(admitted.body["model"], "claude");
}
#[tokio::test]
async fn missing_bearer_is_unauthorized() {
assert!(matches!(
admission().admit(None, &body("claude", 1)).await,
Err(Rejection::Unauthorized)
));
}
#[tokio::test]
async fn virtual_key_model_access_is_enforced() {
assert!(matches!(
admission().admit(Some("sk-limited"), &body("other", 1)).await,
Err(Rejection::ModelNotAllowed(model)) if model == "other"
));
}
#[tokio::test]
async fn one_token_count_feeds_the_tpm_window() {
let admission = admission();
let first = admission
.admit(Some("sk-limited"), &body("claude", 40))
.await
.unwrap();
assert!(first.input_tokens > 40, "{}", first.input_tokens);
assert!(matches!(
admission
.admit(Some("sk-limited"), &body("claude", 40))
.await,
Err(Rejection::LimitExceeded(LimitExceeded::TokensPerMinute))
));
}
#[tokio::test]
async fn oversized_bodies_are_rejected_before_parsing() {
let admission = admission().with_limits(16, 1_000_000);
assert!(matches!(
admission
.admit(Some("sk-master"), &body("claude", 10))
.await,
Err(Rejection::RequestTooLarge { .. })
));
}
#[tokio::test]
async fn context_limit_uses_the_same_count() {
let admission = admission().with_limits(1 << 20, 10);
assert!(matches!(
admission.admit(Some("sk-master"), &body("claude", 40)).await,
Err(Rejection::ContextTooLarge { tokens, max: 10 }) if tokens > 10
));
}
}

View file

@ -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<tokenizers::Tokenizer>),
Approximate,
}
/// Counts input tokens with a bounded number of concurrent encodes.
#[derive(Clone)]
pub struct TokenCounter {
backend: Backend,
permits: Arc<Semaphore>,
}
impl TokenCounter {
/// Load a HuggingFace `tokenizer.json` (the proxy ships the Anthropic one under
/// `litellm/litellm_core_utils/tokenizers/anthropic_tokenizer.json`).
pub fn from_file(path: &Path, concurrency: usize) -> Result<Self, TokenizerError> {
let tokenizer = tokenizers::Tokenizer::from_file(path)
.map_err(|error| TokenizerError::Load(error.to_string()))?;
Ok(Self {
backend: Backend::HuggingFace(Arc::new(tokenizer)),
permits: Arc::new(Semaphore::new(concurrency.max(1))),
})
}
/// `len / APPROX_BYTES_PER_TOKEN`, for hosts without a tokenizer file.
pub fn approximate() -> Self {
Self {
backend: Backend::Approximate,
permits: Arc::new(Semaphore::new(1)),
}
}
pub fn is_exact(&self) -> bool {
matches!(self.backend, Backend::HuggingFace(_))
}
pub async fn count(&self, text: String) -> Result<usize, TokenizerError> {
let tokenizer = match &self.backend {
Backend::Approximate => return Ok(text.len().div_ceil(APPROX_BYTES_PER_TOKEN)),
Backend::HuggingFace(tokenizer) => Arc::clone(tokenizer),
};
if text.len() <= TOKENIZE_INLINE_MAX_BYTES {
return encode_len(&tokenizer, &text);
}
let _permit = self
.permits
.acquire()
.await
.map_err(|_| TokenizerError::Encode("tokenizer pool closed".to_string()))?;
tokio::task::spawn_blocking(move || encode_len(&tokenizer, &text))
.await
.map_err(|error| TokenizerError::Encode(error.to_string()))?
}
}
fn encode_len(tokenizer: &tokenizers::Tokenizer, text: &str) -> Result<usize, TokenizerError> {
tokenizer
.encode_fast(text, false)
.map(|encoding| encoding.len())
.map_err(|error| TokenizerError::Encode(error.to_string()))
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn approximate_counter_rounds_up() {
let counter = TokenCounter::approximate();
assert_eq!(counter.count("abcde".to_string()).await.unwrap(), 2);
assert_eq!(counter.count(String::new()).await.unwrap(), 0);
}
#[tokio::test]
async fn loads_anthropic_tokenizer_and_counts_off_thread() {
let path = Path::new(env!("CARGO_MANIFEST_DIR"))
.join("../../../litellm/litellm_core_utils/tokenizers/anthropic_tokenizer.json");
let counter = TokenCounter::from_file(&path, 1).expect("tokenizer loads");
assert!(counter.is_exact());
let small = counter.count("hello world".to_string()).await.unwrap();
assert!((1..=4).contains(&small), "got {small}");
let large = "the quick brown fox ".repeat(2000);
let large_len = large.len();
let count = counter.count(large).await.unwrap();
assert!(
count > large_len / 8 && count < large_len / 2,
"got {count}"
);
}
}

View file

@ -10,7 +10,7 @@
use axum::extract::FromRequestParts;
use axum::http::StatusCode;
use axum::http::header::{AUTHORIZATION, 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<AppState> 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,

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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<Self, TokenCountError> {
let tokenizer = tokenizer_json
.parse::<tokenizers::Tokenizer>()
.map_err(|error| TokenCountError::Load(error.to_string()))?;
Ok(Self { tokenizer })
}
pub fn count_text(&self, text: &str) -> Result<usize, TokenCountError> {
self.tokenizer
.encode_fast(text, true)
.map(|encoding| encoding.len())
.map_err(|error| TokenCountError::Encode(error.to_string()))
}
pub fn count_request(
&self,
request: &CountableRequest,
) -> Result<InputTokenCount, TokenCountError> {
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::<Result<usize, _>>()?;
let includes_system_message = messages
.iter()
.any(|message| message.role.as_deref() == Some("system"));
let extra_tokens = self.count_extra(
request.tools.as_deref().unwrap_or_default(),
request.tool_choice.as_ref(),
includes_system_message,
)?;
Ok(InputTokenCount {
model: request.model.clone(),
input_tokens: message_tokens + extra_tokens,
})
}
fn count_message(&self, message: &Message) -> Result<usize, TokenCountError> {
let role_tokens = match &message.role {
Some(role) => self.count_text(role)?,
None => 0,
};
let name_tokens = match &message.name {
Some(name) => self.count_text(name)? + TOKENS_PER_NAME,
None => 0,
};
let content_tokens = match &message.content {
Some(MessageContent::Text(text)) => self.count_text(text)?,
Some(MessageContent::Blocks(items)) => items
.iter()
.map(|item| self.count_content_item(item))
.sum::<Result<usize, _>>()?,
None => 0,
};
Ok(TOKENS_PER_MESSAGE + role_tokens + name_tokens + content_tokens)
}
fn count_content_item(&self, item: &ContentItem) -> Result<usize, TokenCountError> {
match item {
ContentItem::Text(text) => self.count_text(text),
ContentItem::Block(ContentBlock::Text { text }) => self.count_text(text),
ContentItem::Block(ContentBlock::Thinking { thinking }) => {
if thinking.is_empty() {
return Ok(0);
}
self.count_text(thinking)
}
ContentItem::Block(ContentBlock::ToolReference { tool_name }) => {
match tool_name.as_deref().filter(|name| !name.is_empty()) {
Some(name) => self.count_text(name),
None => Ok(0),
}
}
ContentItem::Block(ContentBlock::Unsupported) => Err(TokenCountError::Unsupported(
"content block type is counted by the python path".to_string(),
)),
}
}
fn count_extra(
&self,
tools: &[types::ToolDefinition],
tool_choice: Option<&ToolChoice>,
includes_system_message: bool,
) -> Result<usize, TokenCountError> {
let tool_tokens = if tools.is_empty() {
0
} else {
let definitions = self.count_text(&format_function_definitions(tools)?)?;
let discount = if includes_system_message {
TOOLS_WITH_SYSTEM_MESSAGE_DISCOUNT
} else {
0
};
definitions + TOOL_DEFINITIONS_TOKENS - discount
};
let choice_tokens = match tool_choice {
Some(ToolChoice::Mode(mode)) if mode == "none" => TOOL_CHOICE_NONE_TOKENS,
Some(ToolChoice::Mode(_)) | None => 0,
Some(ToolChoice::Named(named)) => {
NAMED_TOOL_CHOICE_TOKENS + self.count_text(&named.function.name)?
}
};
Ok(REPLY_PRIMING_TOKENS + tool_tokens + choice_tokens)
}
}
#[cfg(test)]
mod tests;

View file

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

View file

@ -0,0 +1,108 @@
//! Renders tool definitions the way `litellm.token_counter` does before
//! tokenizing them (the TypeScript-like namespace OpenAI appears to use).
use super::TokenCountError;
use super::types::{EnumValue, FunctionDefinition, Schema, SchemaType, ToolDefinition};
pub(super) fn format_function_definitions(
tools: &[ToolDefinition],
) -> Result<String, TokenCountError> {
let mut lines = vec!["namespace functions {".to_string(), String::new()];
for tool in tools {
let function = resolve_function(tool);
let Some(name) = function.name.as_deref().filter(|name| !name.is_empty()) else {
continue;
};
if let Some(description) = function.description.as_deref().filter(|d| !d.is_empty()) {
lines.push(format!("// {description}"));
}
let parameters = function.parameters.unwrap_or_default();
match &parameters.properties {
Some(properties) if !properties.is_empty() => {
lines.push(format!("type {name} = (_: {{"));
lines.push(format_object_parameters(&parameters, 0)?);
lines.push("}) => any;".to_string());
}
_ => lines.push(format!("type {name} = () => any;")),
}
lines.push(String::new());
}
lines.push("} // namespace functions".to_string());
Ok(lines.join("\n"))
}
fn resolve_function(tool: &ToolDefinition) -> FunctionDefinition {
match &tool.function {
Some(function) => function.clone(),
None => FunctionDefinition {
name: tool.name.clone(),
description: tool.description.clone(),
parameters: tool
.input_schema
.clone()
.or_else(|| tool.parameters.clone()),
},
}
}
fn format_object_parameters(parameters: &Schema, indent: usize) -> Result<String, TokenCountError> {
let Some(properties) = parameters.properties.as_ref().filter(|p| !p.is_empty()) else {
return Ok(String::new());
};
let required = parameters.required.as_deref().unwrap_or_default();
let mut lines = Vec::new();
for (key, props) in properties {
if let Some(description) = props.description.as_deref().filter(|d| !d.is_empty()) {
lines.push(format!("// {description}"));
}
let question = if required.iter().any(|r| r == key) {
""
} else {
"?"
};
lines.push(format!("{key}{question}: {},", format_type(props, indent)?));
}
let pad = " ".repeat(indent);
Ok(lines
.iter()
.map(|line| format!("{pad}{line}"))
.collect::<Vec<_>>()
.join("\n"))
}
fn format_type(props: &Schema, indent: usize) -> Result<String, TokenCountError> {
let Some(SchemaType::Name(schema_type)) = &props.schema_type else {
return Ok("any".to_string());
};
match schema_type.as_str() {
"string" | "integer" | "number" => Ok(match &props.enum_values {
Some(values) => format_enum(values),
None if schema_type == "string" => "string".to_string(),
None => "number".to_string(),
}),
"array" => {
let items = props.items.as_deref().ok_or(TokenCountError::Unsupported(
"array parameter without items".to_string(),
))?;
Ok(format!("{}[]", format_type(items, indent)?))
}
"object" => Ok(format!(
"{{\n{}\n}}",
format_object_parameters(props, indent + 2)?
)),
"boolean" => Ok("boolean".to_string()),
"null" => Ok("null".to_string()),
_ => Ok("any".to_string()),
}
}
fn format_enum(values: &[EnumValue]) -> String {
values
.iter()
.map(|value| match value {
EnumValue::Text(text) => format!("\"{text}\""),
EnumValue::Integer(number) => format!("\"{number}\""),
})
.collect::<Vec<_>>()
.join(" | ")
}

View file

@ -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<Vec<Message>>,
pub tools: Option<Vec<ToolDefinition>>,
pub tool_choice: Option<ToolChoice>,
}
impl CountableRequest {
pub fn parse(body: &[u8]) -> Result<Self, TokenCountError> {
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<String>,
pub name: Option<String>,
pub content: Option<MessageContent>,
}
#[derive(Clone, Debug, Deserialize, PartialEq)]
#[serde(untagged)]
pub enum MessageContent {
Text(String),
Blocks(Vec<ContentItem>),
}
#[derive(Clone, Debug, Deserialize, PartialEq)]
#[serde(untagged)]
pub enum ContentItem {
Text(String),
Block(ContentBlock),
}
#[derive(Clone, Debug, Deserialize, PartialEq)]
#[serde(tag = "type")]
pub enum ContentBlock {
#[serde(rename = "text")]
Text { text: String },
#[serde(rename = "thinking")]
Thinking { thinking: String },
#[serde(rename = "tool_reference")]
ToolReference { tool_name: Option<String> },
/// Images, documents, files and tool use/result blocks price through
/// Python-only helpers, so they stay on the Python counter.
#[serde(other)]
Unsupported,
}
/// Either the OpenAI `{"type": "function", "function": {...}}` shape or the
/// Anthropic `{"name", "description", "input_schema"}` shape.
#[derive(Clone, Debug, Deserialize, PartialEq)]
pub struct ToolDefinition {
pub function: Option<FunctionDefinition>,
pub name: Option<String>,
pub description: Option<String>,
pub input_schema: Option<Schema>,
pub parameters: Option<Schema>,
}
#[derive(Clone, Debug, Deserialize, PartialEq)]
pub struct FunctionDefinition {
pub name: Option<String>,
pub description: Option<String>,
pub parameters: Option<Schema>,
}
#[derive(Clone, Debug, Default, Deserialize, PartialEq)]
pub struct Schema {
#[serde(rename = "type")]
pub schema_type: Option<SchemaType>,
pub description: Option<String>,
#[serde(rename = "enum")]
pub enum_values: Option<Vec<EnumValue>>,
pub items: Option<Box<Schema>>,
pub properties: Option<IndexMap<String, Schema>>,
pub required: Option<Vec<String>>,
}
#[derive(Clone, Debug, Deserialize, PartialEq)]
#[serde(untagged)]
pub enum SchemaType {
Name(String),
Union(Vec<String>),
}
#[derive(Clone, Debug, Deserialize, PartialEq)]
#[serde(untagged)]
pub enum EnumValue {
Text(String),
Integer(i64),
}
#[derive(Clone, Debug, Deserialize, PartialEq)]
#[serde(untagged)]
pub enum ToolChoice {
Mode(String),
Named(NamedToolChoice),
}
#[derive(Clone, Debug, Deserialize, PartialEq)]
pub struct NamedToolChoice {
pub function: NamedFunction,
}
#[derive(Clone, Debug, Deserialize, PartialEq)]
pub struct NamedFunction {
pub name: String,
}

View file

@ -3,7 +3,6 @@ use std::panic::AssertUnwindSafe;
use std::time::Duration;
use futures_util::FutureExt;
use litellm_core::error::Error;
use litellm_python_interop::{Pythonized, panic_to_pyerr, release_gil};
use pyo3::exceptions::PyRuntimeError;
use pyo3::prelude::*;
@ -11,14 +10,15 @@ use serde::Serialize;
use tokio::runtime::{Handle, Runtime};
use tokio::time::{self, MissedTickBehavior};
pub(crate) fn run_sync<T, F>(
pub(crate) fn run_sync<T, E, F>(
py: Python<'_>,
future: F,
map_error: fn(Error) -> PyErr,
map_error: fn(E) -> PyErr,
) -> PyResult<Py<PyAny>>
where
T: Serialize + Send + 'static,
F: Future<Output = Result<T, Error>> + Send + 'static,
E: Send + 'static,
F: Future<Output = Result<T, E>> + Send + 'static,
{
run_sync_on(
py,
@ -28,15 +28,16 @@ where
)
}
fn run_sync_on<T, F>(
fn run_sync_on<T, E, F>(
py: Python<'_>,
runtime: &Runtime,
future: F,
map_error: fn(Error) -> PyErr,
map_error: fn(E) -> PyErr,
) -> PyResult<Py<PyAny>>
where
T: Serialize + Send + 'static,
F: Future<Output = Result<T, Error>> + Send + 'static,
E: Send + 'static,
F: Future<Output = Result<T, E>> + Send + 'static,
{
if Handle::try_current().is_ok() {
return Err(PyRuntimeError::new_err(
@ -49,14 +50,15 @@ where
Pythonized(result).into_pyobject(py).map(Bound::unbind)
}
pub(crate) fn run_async<T, F>(
pub(crate) fn run_async<T, E, F>(
py: Python<'_>,
future: F,
map_error: fn(Error) -> PyErr,
map_error: fn(E) -> PyErr,
) -> PyResult<Bound<'_, PyAny>>
where
T: Serialize + Send + 'static,
F: Future<Output = Result<T, Error>> + Send + 'static,
E: Send + 'static,
F: Future<Output = Result<T, E>> + Send + 'static,
{
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let result = catch_future_panic(future).await?;
@ -65,7 +67,7 @@ where
})
}
fn map_core_result<T>(result: Result<T, Error>, map_error: fn(Error) -> PyErr) -> PyResult<T> {
fn map_core_result<T, E>(result: Result<T, E>, map_error: fn(E) -> PyErr) -> PyResult<T> {
match result {
Ok(value) => Ok(value),
Err(error) => Err(
@ -75,9 +77,9 @@ fn map_core_result<T>(result: Result<T, Error>, map_error: fn(Error) -> PyErr) -
}
}
async fn catch_future_panic<T, F>(future: F) -> PyResult<Result<T, Error>>
async fn catch_future_panic<T, E, F>(future: F) -> PyResult<Result<T, E>>
where
F: Future<Output = Result<T, Error>>,
F: Future<Output = Result<T, E>>,
{
AssertUnwindSafe(future)
.catch_unwind()
@ -85,9 +87,9 @@ where
.map_err(panic_to_pyerr)
}
async fn wait_for_sync_result<T, F>(future: F) -> PyResult<Result<T, Error>>
async fn wait_for_sync_result<T, E, F>(future: F) -> PyResult<Result<T, E>>
where
F: Future<Output = Result<T, Error>>,
F: Future<Output = Result<T, E>>,
{
let future = catch_future_panic(future);
tokio::pin!(future);
@ -114,6 +116,7 @@ mod tests {
use std::thread;
use std::time::Instant;
use litellm_core::error::Error;
use pyo3::panic::PanicException;
use pyo3::types::{PyDict, PyModule};
use serde::Serializer;
@ -237,7 +240,7 @@ mod tests {
let error = runtime.block_on(async {
Python::attach(|py| {
run_sync::<bool, _>(py, async { Ok(true) }, runtime_error)
run_sync::<bool, Error, _>(py, async { Ok(true) }, runtime_error)
.expect_err("sync route should reject a nested Tokio runtime")
})
});
@ -273,7 +276,7 @@ mod tests {
fn sync_runner_maps_a_panicked_future() {
Python::initialize();
Python::attach(|py| {
let error = run_sync::<bool, _>(
let error = run_sync::<bool, Error, _>(
py,
poll_fn(|_| -> Poll<Result<bool, Error>> { panic!("route future panicked") }),
runtime_error,
@ -289,7 +292,7 @@ mod tests {
fn sync_runner_maps_a_panicked_error_mapper() {
Python::initialize();
Python::attach(|py| {
let error = run_sync::<bool, _>(
let error = run_sync::<bool, Error, _>(
py,
async { Err(Error::InvalidRequest("invalid".to_string())) },
panicking_error_mapper,

View file

@ -5,6 +5,7 @@ mod execution;
mod function_trace;
mod marshal;
mod routes;
mod token_counter;
use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection;
use pyo3::prelude::*;
@ -71,6 +72,7 @@ mod _native {
super::errors::register(module)?;
super::routes::register(module)?;
module.add_class::<super::ResponsesWebSocketConnection>()?;
super::token_counter::register(module)?;
super::diagnostics::register(module)
}
}
@ -106,6 +108,7 @@ mod tests {
"chat_completions",
"achat_completions",
"ResponsesWebSocketConnection",
"TokenCounter",
"gil_stats",
];

View file

@ -0,0 +1,64 @@
use std::sync::Arc;
use litellm_core::token_counter::types::CountableRequest;
use litellm_core::token_counter::{
InputTokenCount, TokenCountError, TokenCounter as CoreTokenCounter,
};
use litellm_python_interop::release_gil;
use pyo3::exceptions::{PyRuntimeError, PyValueError};
use pyo3::prelude::*;
use pyo3::types::PyAny;
use crate::errors::RustBridgeDeclined;
use crate::execution::run_async;
/// Counts the input tokens of a raw request body off the Python event loop with
/// the GIL released. Python owns which requests get here and what to do with
/// the count.
#[pyclass(frozen)]
struct TokenCounter {
inner: Arc<CoreTokenCounter>,
}
#[pymethods]
impl TokenCounter {
#[new]
fn new(py: Python<'_>, tokenizer_json: &str) -> PyResult<Self> {
let inner = release_gil(py, || CoreTokenCounter::from_json(tokenizer_json))
.map_err(token_count_error_to_pyerr)?;
Ok(Self {
inner: Arc::new(inner),
})
}
fn acount_request<'py>(&self, py: Python<'py>, body: &[u8]) -> PyResult<Bound<'py, PyAny>> {
let counter = Arc::clone(&self.inner);
let body = body.to_vec();
run_async(
py,
async move {
tokio::task::spawn_blocking(move || count_body(&counter, &body))
.await
.map_err(|error| TokenCountError::Encode(error.to_string()))?
},
token_count_error_to_pyerr,
)
}
}
fn count_body(counter: &CoreTokenCounter, body: &[u8]) -> Result<InputTokenCount, TokenCountError> {
let request = CountableRequest::parse(body)?;
counter.count_request(&request)
}
fn token_count_error_to_pyerr(error: TokenCountError) -> PyErr {
match error {
TokenCountError::Load(message) => PyValueError::new_err(message),
TokenCountError::Unsupported(message) => RustBridgeDeclined::new_err(message),
TokenCountError::Encode(message) => PyRuntimeError::new_err(message),
}
}
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
module.add_class::<TokenCounter>()
}

View file

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

View file

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

View file

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

View file

@ -0,0 +1,79 @@
"""Thin Python wrapper for the native Rust input token counter."""
from __future__ import annotations
from collections.abc import Awaitable
from dataclasses import dataclass
from functools import lru_cache
from typing import Final, Protocol, cast # noqa: TID251 # native extension exposes dynamically typed callables
from pydantic import TypeAdapter
import litellm
from litellm._logging import verbose_logger
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.configuration import rust_enabled
from litellm.rust_bridge.runtime import BridgeErrorContext, RustHandled, aattempt
class RustTokenCounter(Protocol):
def acount_request(self, body: bytes) -> Awaitable[object]:
raise NotImplementedError
class RustTokenCounterFactory(Protocol):
def __call__(self, tokenizer_json: str) -> RustTokenCounter:
raise NotImplementedError
@dataclass(frozen=True, slots=True)
class InputTokenCount:
model: str
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

View file

@ -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():
"""

View file

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

View file

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