Merge remote-tracking branch 'refs/remotes/upstream/litellm_internal_staging' into codex/rust-auth-control-plane

# Conflicts:
#	litellm-rust/Cargo.lock
#	litellm-rust/Cargo.toml
#	litellm-rust/crates/ai-gateway/Cargo.toml
#	litellm-rust/crates/ai-gateway/src/routes/realtime/mod.rs
#	litellm/proxy/proxy_server.py
This commit is contained in:
Ishaan Jaff 2026-06-24 15:38:24 -07:00
commit c5c0d5fce6
No known key found for this signature in database
102 changed files with 8986 additions and 800 deletions

View file

@ -29,6 +29,7 @@ jobs:
tests/test_litellm/proxy/_experimental
tests/test_litellm/proxy/experimental
tests/test_litellm/proxy/common_utils
tests/test_litellm/proxy/logging_endpoints
tests/test_litellm/proxy/test_*.py
workers: 2
reruns: 2

3
.gitignore vendored
View file

@ -123,3 +123,6 @@ crash.*.log
# and should be committed.
.vscode
.pin_list.txt
# pytest coverage data
.coverage

View file

@ -84,6 +84,8 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = (
"/active/callbacks",
"/callbacks",
"/team_callback",
# Rust data-plane gateway → proxy control-plane API (logging today, auth later)
"/v1/rust_control_plane/",
# Alerting / email / IP allowlist
"/alerting/",
"/email/",

View file

@ -69,7 +69,7 @@
},
"reportMatchNotExhaustive": {
"baseline": 1,
"slack": 3
"slack": 0
},
"reportMissingParameterType": {
"baseline": 3933,

View file

@ -13,7 +13,9 @@ from litellm import Router, verbose_logger
from litellm._uuid import uuid
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
from litellm.litellm_core_utils.prompt_templates.common_utils import (
extract_file_metadata,
)
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
from litellm.llms.base_llm.managed_resources.isolation import (
build_list_page,
@ -981,9 +983,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
target_model_names_list: List[str],
) -> OpenAIFileObject:
## GET THE FILE TYPE FROM THE CREATE FILE REQUEST
file_data = extract_file_data(create_file_request["file"])
file_type = file_data["content_type"]
_, file_type = extract_file_metadata(create_file_request["file"])
output_file_id = file_objects[0].id
model_id = file_objects[0]._hidden_params.get("model_id")

View file

@ -77,6 +77,20 @@ such as `ai-gateway`, router hosts, or standalone servers:
- Avoid `expect`/`unwrap` in server startup and request paths unless the panic is
impossible by construction and documented.
## Constants
Magic numbers and fixed strings go in a crate-level `constants.rs`, never
hardcoded inline — the Rust mirror of Python's `litellm/constants.py`.
- Each crate that needs them has `src/constants.rs` (declared `mod constants;`);
import from it (`use crate::constants::...`). Don't scatter `const` values at
the top of feature modules.
- An env-overridable tunable still lives in `constants.rs` as its `DEFAULT_*`
value; the env read (with fallback to that default) happens at the host/config
resolution layer, not in `core`/`providers`.
- Exception: a value that is purely local to one function and has no meaning
elsewhere may stay inline, but prefer `constants.rs` when in doubt.
## Checks
Run these before pushing Rust changes. The same checks run in GitHub Actions

View file

@ -0,0 +1,12 @@
# ai-gateway architecture
The Rust ai-gateway does LLM inference (realtime WebSocket). Spend tracking is an
API callback: it POSTs each finished session to the LiteLLM proxy, which records
spend and runs the usual callbacks.
```mermaid
flowchart LR
C[client] <--> G[Rust ai-gateway<br/>LLM inference]
G <--> O[OpenAI realtime]
G -. spend tracking callback .-> P[litellm proxy]
```

View file

@ -15,8 +15,11 @@ required-features = ["server"]
[dependencies]
litellm-core.workspace = true
# reqwest (rustls + json) is used by io/ocr and ships realtime logs to the
# Python proxy callbacks API.
reqwest.workspace = true
tokio = { workspace = true, features = ["rt-multi-thread", "macros", "net", "time"] }
# `sync` powers the bounded mpsc channel the realtime logger drains.
tokio = { workspace = true, features = ["rt-multi-thread", "macros", "net", "time", "sync"] }
tokio-tungstenite.workspace = true
futures-util.workspace = true
serde_json.workspace = true

View file

@ -17,8 +17,9 @@ dials OpenAI upstream, and splices the two sockets frame-by-frame.
Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge.
- **Client endpoint:** `wss://<host>/v1/realtime?model=<model>` (WebSocket)
- **Auth:** `Authorization: Bearer <key>` — proxy-admin only for realtime serving today. Non-admin virtual keys are verified but rejected until realtime usage is reported back to the LiteLLM proxy. Fails closed.
- **Auth:** `Authorization: Bearer <key>` — the master key (admin) or a virtual key validated by the LiteLLM proxy (see below). Fails closed.
- **Health:** `GET /health/readiness`, `GET /health/liveness`, `GET /health/gil`
- **Request logs:** POSTed to a LiteLLM proxy at `/v1/rust_control_plane/logs` (see [Request logging](#request-logging))
> **Realtime serving is pure Rust.** Python is used at **load time only** — to
> read the config once at boot. The realtime hot path never touches Python.
@ -42,9 +43,9 @@ route/model permissions are enforced in exactly one place. A client presents
load. A bounded ≤200-entry cache is available for future high-RPS *per-request*
routes via `LITELLM_AUTH_CACHE_TTL_SECS > 0`, which trades budget/rate-limit
freshness (bounded by the TTL) for fewer control-plane calls — like the proxy's
own ~60s auth cache. Non-admin virtual-key identities are rejected on the
realtime route until realtime usage is reported back to the proxy, so callers
cannot consume upstream usage without key/team spend attribution.
own ~60s auth cache. Realtime session usage is reported back to the proxy with
the resolved key/user/team metadata, so key and team spend logs keep their
normal attribution.
Data plane → control plane is itself authenticated with a **dedicated data-plane key**
(NOT the master key — least privilege): the gateway sends
@ -74,13 +75,11 @@ export LITELLM_AUTH_VERIFY_URL=https://<proxy-host>/v1/rust_control_plane/authen
./litellm-ai-gateway
```
Fails closed: a missing/wrong data-plane key, an unreachable proxy, or a rejected key
all yield `401`; a valid non-admin virtual key currently yields `403` on
`/v1/realtime` until realtime spend reporting is available. Revocation and budget
changes take effect within the cache TTL (if caching is enabled via
`LITELLM_AUTH_CACHE_TTL_SECS`; off by default → every connection re-verifies).
Keep the proxy on a private network — the verify endpoint is internal-only and
excluded from the public OpenAPI spec.
Fails closed: a missing/wrong data-plane key, an unreachable proxy, or a rejected
key all yield `401`. Revocation and budget changes take effect within the cache
TTL (if caching is enabled via `LITELLM_AUTH_CACHE_TTL_SECS`; off by default →
every connection re-verifies). Keep the proxy on a private network — the verify
endpoint is internal-only and excluded from the public OpenAPI spec.
## Configuration (config.yaml)
@ -122,11 +121,12 @@ overridden at deploy time (e.g. a Render secret file mounted at the same path).
| `LITELLM_CONFIG_PATH` | yes (config mode) | — | Path to the config.yaml the gateway loads its `model_list` from. The Docker image defaults this to `/app/config.yaml`. |
| `LITELLM_MASTER_KEY` | yes | — | Admin bearer token (checked locally, no proxy call). Unset ⇒ the master-key path is disabled. |
| `LITELLM_DATA_PLANE_KEY` | for virtual keys | — | Dedicated secret the gateway sends as `X-LiteLLM-Data-Plane-Key` to authenticate itself to the proxy's verify endpoint. **Must match the proxy's `LITELLM_DATA_PLANE_KEY`.** Not the master key. |
| `LITELLM_AUTH_VERIFY_URL` | for virtual keys | `http://localhost:4000/v1/rust_control_plane/authentication` | The proxy's verify endpoint the gateway delegates virtual-key auth to. Non-admin virtual keys are still rejected by `/v1/realtime` until realtime spend reporting is wired. |
| `LITELLM_AUTH_VERIFY_URL` | for virtual keys | `http://localhost:4000/v1/rust_control_plane/authentication` | The proxy's verify endpoint the gateway delegates virtual-key auth to. |
| `LITELLM_AUTH_CACHE_TTL_SECS` | no | `0` | Verified-key cache TTL. **`0` = off (default)** → every connection re-verifies (budget/rate-limit enforced each time); cheap for realtime since auth is per-connection. Set `> 0` only for high-RPS per-request routes, trading budget/rate-limit freshness for fewer proxy calls. |
| `OPENAI_API_KEY` | yes | — | Upstream OpenAI key. Referenced by config.yaml as `os.environ/OPENAI_API_KEY` for the gateway→OpenAI dial. |
| `HOST` | no | `127.0.0.1` | **Set to `0.0.0.0` in any container/deploy** or external traffic is refused. |
| `PORT` | no | `4001` | Listen port. Render and most PaaS inject this automatically. |
| `LITELLM_PROXY_BASE_URL` | no | `http://localhost:4000` | LiteLLM proxy that request logs are POSTed to. See [Request logging](#request-logging). |
> Secrets (`LITELLM_MASTER_KEY`, `LITELLM_DATA_PLANE_KEY`, `OPENAI_API_KEY`) are never
> baked into the image or `render.yaml` — inject them at deploy time only.
@ -145,6 +145,18 @@ This mode links no libpython and needs no config file, but it only supports one
hard-coded OpenAI deployment. **config.yaml is the recommended path** — use the
stand-in only for the leanest possible build.
## Request logging
The gateway runs no spend logic. When a session ends it builds one
`StandardLoggingPayload` and POSTs it to `{LITELLM_PROXY_BASE_URL}/v1/rust_control_plane/logs`
(admin-only, bearer = `LITELLM_MASTER_KEY`), and the proxy replays it through its
normal callbacks (spend logs, Langfuse, etc.). The POST is non-blocking: a bounded
channel drained by a background worker, dropping with a counter if the proxy is
down. It sends one payload per session. Both env vars are in the table above.
Worker tuning, rarely needed: `LITELLM_LOG_CHANNEL_CAPACITY` (4096),
`LITELLM_LOG_BATCH_SIZE` (256), `LITELLM_LOG_FLUSH_INTERVAL_MS` (500).
## Build & run with Docker
The image is built `--features python-config` and installs litellm **from this

View file

@ -2,12 +2,11 @@
//! keeps handlers clean and auth testable).
//!
//! For now this supports both the local **master key** and virtual keys resolved
//! through the Python control plane. Individual routes can still narrow the
//! accepted identities. For example, realtime rejects non-admin virtual keys
//! until realtime usage is reported back to the proxy for spend attribution.
//! through the Python control plane. Routes receive the resolved identity and can
//! pass it to logging/billing paths for spend attribution.
//!
//! A handler opts in by adding [`RequireMasterKey`] to its arguments; auth then
//! runs during extraction, before the handler body. Routes never re-implement it.
//! A handler opts in by adding an auth extractor to its arguments; auth then runs
//! during extraction, before the handler body. Routes never re-implement it.
pub mod cache;
pub mod client;
@ -21,10 +20,29 @@ use axum::extract::FromRequestParts;
use axum::http::header::AUTHORIZATION;
use axum::http::request::Parts;
use axum::http::StatusCode;
use sha2::{Digest, Sha256};
use subtle::ConstantTimeEq;
use crate::state::AppState;
/// SHA-256 hex digest of a token — the exact transform the Python proxy applies
/// (`litellm.proxy.utils.hash_token`).
///
/// STRICT REQUIREMENT: a raw key (`LITELLM_MASTER_KEY`, a virtual key, …) must
/// **never** leave this gateway in a log payload. Spend logs and every callback
/// integration receive `user_api_key_hash`, so that field must be this hash, not
/// the credential. Hashing here also means the value matches the key's hash in
/// `LiteLLM_SpendLogs.api_key`, so realtime spend joins with the rest of LiteLLM.
pub fn hash_token(token: &str) -> String {
let digest = Sha256::digest(token.as_bytes());
let mut hex = String::with_capacity(digest.len() * 2);
for byte in digest {
use std::fmt::Write;
let _ = write!(hex, "{byte:02x}");
}
hex
}
/// Extractor that requires the configured master key as a bearer token.
///
/// Rejections: `500` when no master key is configured (permanent
@ -61,3 +79,23 @@ impl FromRequestParts<AppState> for RequireMasterKey {
}
}
}
#[cfg(test)]
mod tests {
use super::hash_token;
#[test]
fn hash_token_matches_python_sha256_hexdigest() {
// Must equal hashlib.sha256("sk-1234".encode()).hexdigest() — the value
// the proxy stores in LiteLLM_SpendLogs.api_key.
assert_eq!(
hash_token("sk-1234"),
"88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b"
);
// 64 lowercase hex chars, and never the raw input.
let h = hash_token("sk-secret");
assert_eq!(h.len(), 64);
assert!(h.chars().all(|c| c.is_ascii_hexdigit()));
assert_ne!(h, "sk-secret");
}
}

View file

@ -204,6 +204,7 @@ mod tests {
AppState {
router: Arc::new(Router::new(vec![])),
master_key: master_key.map(Arc::from),
loggers: Arc::new(Vec::new()),
realtime_pool: RealtimePool::disabled(),
authenticator,
key_cache: Arc::new(KeyCache::new()),

View file

@ -0,0 +1,29 @@
//! Crate-level constants for the ai-gateway.
//!
//! Per `litellm-rust/CLAUDE.md`, magic numbers and fixed strings live here
//! (the Rust mirror of Python's `litellm/constants.py`), not inline in feature
//! modules. Env-overridable tunables keep their `DEFAULT_*` value here; the env
//! read + fallback happens at the host/config layer.
/// Default LiteLLM control-plane base URL for request-log egress when
/// `LITELLM_PROXY_BASE_URL` is unset.
pub(crate) const DEFAULT_PROXY_BASE_URL: &str = "http://localhost:4000";
/// The logs ingest path appended to the proxy base. Not a tunable; it is the
/// proxy's API contract (the rust-control-plane router on the Python proxy).
pub(crate) const RUST_CONTROL_PLANE_LOGS_PATH: &str = "/v1/rust_control_plane/logs";
/// Default bounded channel depth for the log-egress worker.
/// Override: `LITELLM_LOG_CHANNEL_CAPACITY`.
pub(crate) const DEFAULT_CHANNEL_CAPACITY: usize = 4096;
/// Default max records POSTed per request to the control plane.
/// Override: `LITELLM_LOG_BATCH_SIZE`.
pub(crate) const DEFAULT_MAX_BATCH_SIZE: usize = 256;
/// Default partial-batch flush cadence, in ms.
/// Override: `LITELLM_LOG_FLUSH_INTERVAL_MS`.
pub(crate) const DEFAULT_FLUSH_INTERVAL_MS: u64 = 500;
/// Provider attributed to realtime sessions in the logging payload.
pub(crate) const DEFAULT_PROVIDER: &str = "openai";

View file

@ -0,0 +1,24 @@
//! The `CustomLogger` trait — the Rust mirror of Python
//! `litellm/integrations/custom_logger.py::CustomLogger`.
//!
//! Synchronous (no `async_trait`): callbacks are O(1) enqueue-and-return so the
//! realtime splice never blocks on a logger. Default bodies are no-ops so a
//! logger can implement only the events it cares about.
use crate::integrations::types::{LogError, LoggingError, StandardLoggingPayload};
pub trait CustomLogger: Send + Sync {
/// Record a successful call. Default: no-op.
fn log_success_event(&self, _payload: &StandardLoggingPayload) -> Result<(), LogError> {
Ok(())
}
/// Record a failed call. Default: no-op.
fn log_failure_event(
&self,
_payload: &StandardLoggingPayload,
_error: &LoggingError,
) -> Result<(), LogError> {
Ok(())
}
}

View file

@ -0,0 +1,209 @@
//! A `CustomLogger` that ships finished events to the LiteLLM Python proxy's
//! `/v1/rust_control_plane/logs` endpoint.
//!
//! The callback path is non-blocking: `log_success_event` / `log_failure_event`
//! build a `LogRecord` and `try_send` it onto a bounded channel, returning a
//! `LogError` (never panicking, never awaiting) if the channel is full or the
//! worker has gone away. A spawned background worker drains the channel, batches
//! records into `{"records":[...]}`, and POSTs them to the proxy with a pooled
//! `reqwest::Client`.
use std::sync::Arc;
use std::time::Duration;
use reqwest::Client;
use tokio::sync::mpsc::{self, Receiver, Sender};
use tokio::time::interval;
use crate::constants::{
DEFAULT_CHANNEL_CAPACITY, DEFAULT_FLUSH_INTERVAL_MS, DEFAULT_MAX_BATCH_SIZE,
DEFAULT_PROXY_BASE_URL, RUST_CONTROL_PLANE_LOGS_PATH,
};
use crate::integrations::custom_logger::CustomLogger;
use crate::integrations::types::{
CallbackLogsRequest, LogError, LogRecord, LoggingError, StandardLoggingPayload,
};
/// Egress worker tunables. Each field defaults to the matching `DEFAULT_*` const
/// in `crate::constants` and is overridable via an env var (read once at logger
/// construction).
struct EgressTunables {
channel_capacity: usize,
max_batch_size: usize,
flush_interval: Duration,
}
impl EgressTunables {
fn from_env() -> Self {
Self {
channel_capacity: env_positive(
"LITELLM_LOG_CHANNEL_CAPACITY",
DEFAULT_CHANNEL_CAPACITY,
),
max_batch_size: env_positive("LITELLM_LOG_BATCH_SIZE", DEFAULT_MAX_BATCH_SIZE),
flush_interval: Duration::from_millis(env_positive(
"LITELLM_LOG_FLUSH_INTERVAL_MS",
DEFAULT_FLUSH_INTERVAL_MS,
)),
}
}
}
/// Parse a positive integer env var, falling back to `default` on missing,
/// unparseable, or non-positive values. Generic over the integer type so one
/// helper serves both the `usize` capacities and the `u64` interval.
fn env_positive<T>(name: &str, default: T) -> T
where
T: std::str::FromStr + PartialOrd + From<u8>,
{
let zero = T::from(0u8);
std::env::var(name)
.ok()
.and_then(|value| value.trim().parse::<T>().ok())
.filter(|n| *n > zero)
.unwrap_or(default)
}
/// Ships realtime logging events to the LiteLLM Python proxy.
pub struct LiteLLMPythonProxyAPILogger {
sink: Sender<LogRecord>,
}
impl LiteLLMPythonProxyAPILogger {
/// Spawn the background worker and return a logger handle. `base` is the
/// proxy base URL (no trailing path); `master_key` is sent as a bearer token.
pub fn start(base: String, master_key: String) -> Arc<Self> {
let tunables = EgressTunables::from_env();
let (sink, receiver) = mpsc::channel::<LogRecord>(tunables.channel_capacity);
let url = format!(
"{}{}",
base.trim_end_matches('/'),
RUST_CONTROL_PLANE_LOGS_PATH
);
let client = Client::new();
tokio::spawn(worker_loop(
receiver,
client,
url,
master_key,
tunables.max_batch_size,
tunables.flush_interval,
));
Arc::new(Self { sink })
}
/// Build a logger from the environment: `LITELLM_PROXY_BASE_URL` (default
/// `http://localhost:4000`) and `LITELLM_MASTER_KEY`.
///
/// `LITELLM_PROXY_BASE_URL` is treated as the full base and the route is
/// appended verbatim, so if the proxy runs under a `SERVER_ROOT_PATH`
/// (e.g. served at `https://host/litellm`), include it in the base
/// (`LITELLM_PROXY_BASE_URL=https://host/litellm`) and the POST lands at
/// `https://host/litellm/v1/rust_control_plane/logs`.
pub fn from_env() -> Arc<Self> {
let base = std::env::var("LITELLM_PROXY_BASE_URL")
.ok()
.filter(|value| !value.trim().is_empty())
.unwrap_or_else(|| DEFAULT_PROXY_BASE_URL.to_string());
let key = std::env::var("LITELLM_MASTER_KEY").unwrap_or_default();
Self::start(base, key)
}
fn enqueue(&self, record: LogRecord) -> Result<(), LogError> {
self.sink.try_send(record).map_err(|err| match err {
mpsc::error::TrySendError::Full(_) => LogError::channel_full(),
mpsc::error::TrySendError::Closed(_) => LogError::channel_closed(),
})
}
}
impl CustomLogger for LiteLLMPythonProxyAPILogger {
fn log_success_event(&self, payload: &StandardLoggingPayload) -> Result<(), LogError> {
self.enqueue(LogRecord {
status: "success".to_string(),
payload: payload.clone(),
error: None,
})
}
fn log_failure_event(
&self,
payload: &StandardLoggingPayload,
error: &LoggingError,
) -> Result<(), LogError> {
self.enqueue(LogRecord {
status: "failure".to_string(),
payload: payload.clone(),
error: Some(format!("{}: {}", error.kind, error.message)),
})
}
}
/// Drain the channel, batching records and POSTing them to the proxy. Exits when
/// the channel is closed (all senders dropped) and drained.
async fn worker_loop(
mut receiver: Receiver<LogRecord>,
client: Client,
url: String,
master_key: String,
max_batch_size: usize,
flush_interval: Duration,
) {
let mut ticker = interval(flush_interval);
let mut batch: Vec<LogRecord> = Vec::with_capacity(max_batch_size);
loop {
tokio::select! {
maybe_record = receiver.recv() => {
match maybe_record {
Some(record) => {
batch.push(record);
if batch.len() >= max_batch_size {
flush(&client, &url, &master_key, &mut batch).await;
}
}
None => {
// Channel closed: flush remaining and exit.
flush(&client, &url, &master_key, &mut batch).await;
break;
}
}
}
_ = ticker.tick() => {
flush(&client, &url, &master_key, &mut batch).await;
}
}
}
}
/// POST the current batch (if any), clearing it. Errors are logged, not fatal.
async fn flush(client: &Client, url: &str, master_key: &str, batch: &mut Vec<LogRecord>) {
if batch.is_empty() {
return;
}
let records = std::mem::take(batch)
.into_iter()
.map(LogRecord::into_callback_record)
.collect();
let body = CallbackLogsRequest { records };
let response = client
.post(url)
.bearer_auth(master_key)
.json(&body)
.send()
.await;
match response {
Ok(resp) if resp.status().is_success() => {}
Ok(resp) => {
eprintln!(
"litellm-ai-gateway: callback logs POST returned {} to {url}",
resp.status()
);
}
Err(err) => {
eprintln!("litellm-ai-gateway: callback logs POST failed to {url}: {err}");
}
}
}

View file

@ -0,0 +1,10 @@
//! Pure-Rust logging integrations. Names map 1:1 to Python
//! `litellm/integrations/`:
//! - [`custom_logger::CustomLogger`] — the callback trait
//! - [`litellm_python_proxy_api::LiteLLMPythonProxyAPILogger`] — ships events
//! to the Python proxy's `/v1/callbacks/logs` endpoint
//! - [`types`] — the typed `StandardLoggingPayload` wire contract
pub mod custom_logger;
pub mod litellm_python_proxy_api;
pub mod types;

View file

@ -0,0 +1,164 @@
//! Typed payloads for the LiteLLM `/v1/callbacks/logs` realtime-logging contract.
//!
//! Field names below are the EXACT JSON keys the Python replay path + spend-logs
//! builder read. Note the deliberate mix:
//! - `startTime` / `endTime` are camelCase (epoch f64 seconds)
//! - `response_cost` / `prompt_tokens` / etc. are snake_case
//!
//! Mirrors Python `litellm/integrations/` + the proxy `CallbackLogsRequest`
//! contract 1:1.
use serde::Serialize;
use serde_json::Value;
use std::collections::HashMap;
/// Cumulative token usage for a realtime session.
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct Usage {
pub prompt_tokens: u64,
pub completion_tokens: u64,
pub total_tokens: u64,
}
/// Cost-attribution metadata threaded from the authenticated request.
#[derive(Clone, Debug, Default)]
pub struct RequestMetadata {
pub user_api_key_hash: Option<String>,
pub user_api_key_user_id: Option<String>,
pub user_api_key_team_id: Option<String>,
}
/// A logging-callback failure (e.g. a custom logger raised). Mirrors the Python
/// failure-event shape: a message plus an exception kind/class name.
#[derive(Clone, Debug)]
pub struct LoggingError {
pub message: String,
pub kind: String,
}
/// A non-fatal error returned by a `CustomLogger` when it cannot enqueue an
/// event (channel full or the background worker has shut down).
#[derive(Clone, Debug)]
pub struct LogError {
pub message: String,
pub kind: String,
}
impl LogError {
pub fn channel_full() -> Self {
Self {
message: "logging channel is full; dropping record".to_string(),
kind: "ChannelFull".to_string(),
}
}
pub fn channel_closed() -> Self {
Self {
message: "logging channel is closed; worker has shut down".to_string(),
kind: "ChannelClosed".to_string(),
}
}
}
impl std::fmt::Display for LogError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}: {}", self.kind, self.message)
}
}
impl std::error::Error for LogError {}
/// Batch wrapper — the top-level request body.
/// Matches Python `CallbackLogsRequest { records: list[CallbackLogRecord] }`.
#[derive(Serialize)]
pub struct CallbackLogsRequest {
pub records: Vec<CallbackLogRecord>,
}
/// One finished logging event.
/// Matches `CallbackLogRecord { status, standard_logging_payload, error? }`.
#[derive(Serialize)]
pub struct CallbackLogRecord {
/// "success" | "failure". On "failure", `error` (or payload.error_str)
/// becomes the replayed exception string.
pub status: String,
pub standard_logging_payload: StandardLoggingPayload,
/// Only meaningful when status == "failure". Omitted on success.
#[serde(skip_serializing_if = "Option::is_none")]
pub error: Option<String>,
}
/// The self-describing payload. Field names are the EXACT JSON keys the Python
/// replay path + spend-logs builder read.
#[derive(Clone, Debug, Serialize)]
pub struct StandardLoggingPayload {
pub id: String,
pub litellm_call_id: String,
/// e.g. "realtime", "acompletion". Falls back to "acompletion" if absent.
pub call_type: String,
pub model: String,
pub custom_llm_provider: String,
/// Spend ($) written to LiteLLM_SpendLogs.spend.
pub response_cost: f64,
pub prompt_tokens: u64,
pub completion_tokens: u64,
pub total_tokens: u64,
/// EPOCH SECONDS as float — camelCase keys, NOT snake_case.
#[serde(rename = "startTime")]
pub start_time: f64,
#[serde(rename = "endTime")]
pub end_time: f64,
pub stream: bool,
pub metadata: StandardLoggingMetadata,
/// Optional; stored as request input on the spend log row.
#[serde(skip_serializing_if = "Option::is_none")]
pub messages: Option<Value>,
}
/// Cost-attribution keys. The replayer maps these into litellm_params.metadata,
/// which the spend-logs builder reads to set user / team_id / organization_id.
#[derive(Clone, Debug, Serialize, Default)]
pub struct StandardLoggingMetadata {
pub user_api_key_hash: Option<String>, // -> SpendLogs.api_key
pub user_api_key_user_id: Option<String>, // -> SpendLogs.user
pub user_api_key_team_id: Option<String>, // -> SpendLogs.team_id
// Optional but read by the builder; include when known:
#[serde(skip_serializing_if = "Option::is_none")]
pub user_api_key_alias: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub user_api_key_org_id: Option<String>, // -> SpendLogs.organization_id
#[serde(skip_serializing_if = "Option::is_none")]
pub user_api_key_end_user_id: Option<String>, // -> SpendLogs.end_user
#[serde(skip_serializing_if = "Option::is_none")]
pub spend_logs_metadata: Option<HashMap<String, Value>>,
}
/// The unit handed to a `CustomLogger` sink: a finished payload plus its status
/// and (on failure) the replayed error string.
#[derive(Clone, Debug)]
pub struct LogRecord {
pub status: String,
pub payload: StandardLoggingPayload,
pub error: Option<String>,
}
impl LogRecord {
pub fn into_callback_record(self) -> CallbackLogRecord {
CallbackLogRecord {
status: self.status,
standard_logging_payload: self.payload,
error: self.error,
}
}
}

View file

@ -126,6 +126,9 @@ pub(crate) async fn read_event(upstream_rx: &mut UpstreamRx) -> CoreResult<Realt
/// `session.created` here; the fresh-dial path passes `None` and lets the upstream
/// deliver it). Then a single select loop forwards both directions through the
/// transforms until either side closes or the idle timeout fires.
/// `observe` is invoked on **upstream→client** events only (the trusted side that
/// carries `session.created` and `response.done` usage) — never on client events,
/// so a client cannot fabricate usage into its own logs.
#[allow(clippy::too_many_arguments)]
pub(crate) async fn splice<In, Out>(
model: &str,
@ -133,6 +136,7 @@ pub(crate) async fn splice<In, Out>(
mut upstream_rx: UpstreamRx,
prelude: Option<RealtimeEvent>,
idle_timeout: Option<Duration>,
mut observe: impl FnMut(&RealtimeEvent) + Send,
mut client_in: In,
mut client_out: Out,
) -> CoreResult<()>
@ -165,6 +169,10 @@ where
// client -> upstream
client_event = client_in.next() => {
let Some(event) = client_event else { break }; // client disconnected
// NOTE: do NOT observe client events. session.created / response.done
// (carrying usage) are server→client events; observing the client arm
// would let an authenticated client POST a fabricated response.done and
// inflate its own spend log. Logging observes upstream events only.
for outbound in config.transform_realtime_request(&event, model)?.events {
let payload = serde_json::to_string(&outbound)
.map_err(|err| CoreError::InvalidResponse(err.to_string()))?;
@ -181,6 +189,7 @@ where
Message::Text(text) => {
let event: RealtimeEvent = serde_json::from_str(&text)
.map_err(|err| CoreError::InvalidResponse(err.to_string()))?;
observe(&event);
for outbound in config.transform_realtime_response(&event, model)?.events {
client_out
.send(outbound)
@ -207,11 +216,13 @@ where
/// framework-agnostic; the gateway adapts its axum socket to these. This is the
/// fresh-dial path: dial, then splice. The pool's warm-handoff path skips the dial
/// and calls [`splice`] directly with a buffered `session.created`.
#[allow(clippy::too_many_arguments)]
pub async fn realtime<In, Out>(
model: &str,
api_key: Option<&str>,
api_base: Option<&str>,
idle_timeout: Option<Duration>,
observe: impl FnMut(&RealtimeEvent) + Send,
client_in: In,
client_out: Out,
) -> CoreResult<()>
@ -229,6 +240,7 @@ where
upstream_rx,
None,
idle_timeout,
observe,
client_in,
client_out,
)
@ -238,10 +250,12 @@ where
/// Splice a pre-warmed upstream (taken from [`crate::io::realtime_pool`]) to the
/// client. Relays the buffered `session.created` first, then splices exactly like
/// the fresh-dial path — so a warm session is indistinguishable from a fresh one.
#[allow(clippy::too_many_arguments)]
pub async fn realtime_warm<In, Out>(
model: &str,
handoff: crate::io::realtime_pool::WarmHandoff,
idle_timeout: Option<Duration>,
observe: impl FnMut(&RealtimeEvent) + Send,
client_in: In,
client_out: Out,
) -> CoreResult<()>
@ -256,6 +270,7 @@ where
handoff.rx,
Some(handoff.session_created),
idle_timeout,
observe,
client_in,
client_out,
)
@ -303,6 +318,7 @@ mod tests {
Some(&key_owned),
None,
None,
|_| {},
client_in,
client_out,
)

View file

@ -24,5 +24,15 @@ pub mod routes;
#[cfg(feature = "server")]
pub mod state;
// Realtime request logging. Only the server serves realtime, so these are
// `server`-gated; `io::realtime` exposes the generic `observe` hook while the
// collector and callback fan-out live here.
#[cfg(feature = "server")]
mod constants;
#[cfg(feature = "server")]
pub mod integrations;
#[cfg(feature = "server")]
mod realtime;
#[cfg(feature = "python-config")]
pub mod python;

View file

@ -18,6 +18,8 @@ use litellm_ai_gateway::routes;
use litellm_ai_gateway::state::AppState;
use litellm_core::router::{Deployment, LiteLLMParams, Router};
use litellm_ai_gateway::integrations::custom_logger::CustomLogger;
use litellm_ai_gateway::integrations::litellm_python_proxy_api::LiteLLMPythonProxyAPILogger;
#[cfg(feature = "python-config")]
use litellm_ai_gateway::python;
@ -44,6 +46,12 @@ async fn main() {
);
}
// Spawn the realtime-logging worker (drains a channel → POSTs batches to the
// Python proxy's /v1/callbacks/logs). Built here so the spawn lands on the
// tokio runtime. `from_env` reads LITELLM_PROXY_BASE_URL + LITELLM_MASTER_KEY.
let proxy_logger = LiteLLMPythonProxyAPILogger::from_env();
let loggers: Vec<Arc<dyn CustomLogger>> = vec![proxy_logger];
let router = Arc::new(build_router());
// Build the pre-warmed realtime pool and register each deployment's upstream
@ -85,6 +93,7 @@ async fn main() {
let state = AppState {
router,
master_key,
loggers: Arc::new(loggers),
realtime_pool,
authenticator,
key_cache,

View file

@ -0,0 +1,4 @@
//! Realtime logging collector. Observes the realtime event stream and emits a
//! `StandardLoggingPayload` to the registered callbacks on session close.
pub mod streaming;

View file

@ -0,0 +1,352 @@
//! `RealTimeStreaming` — the realtime logging collector.
//!
//! Mirrors Python `litellm.realtime_api.main.RealTimeStreaming`: it observes the
//! event stream in O(1) (never buffering frames), accumulating just the fields
//! the spend log needs (model, id, cumulative usage), then on session close
//! builds a `StandardLoggingPayload` and fans it out to every registered
//! `CustomLogger`.
use std::sync::Arc;
use std::time::{SystemTime, UNIX_EPOCH};
use litellm_core::realtime::types::RealtimeEvent;
use serde_json::Value;
use crate::constants::DEFAULT_PROVIDER;
use crate::integrations::custom_logger::CustomLogger;
use crate::integrations::types::{
RequestMetadata, StandardLoggingMetadata, StandardLoggingPayload, Usage,
};
/// Current wall-clock time as epoch seconds (float), matching the Python
/// `startTime`/`endTime` contract.
fn epoch_seconds() -> f64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_secs_f64())
.unwrap_or(0.0)
}
/// Status of a finished realtime session, mapped to the callback record status.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum SessionStatus {
Success,
Failure,
}
/// Accumulates realtime session state and emits a logging payload on close.
pub struct RealTimeStreaming {
callbacks: Vec<Arc<dyn CustomLogger>>,
/// REQUEST-ID RULE: the SpendLogs `request_id` == the OpenAI realtime session
/// id (`sess_…`), captured from `session.created`. Both `id` and
/// `litellm_call_id` are set to that value so the Python writer logs the same
/// id regardless of which field it reads. The gateway-generated `rt-…` id
/// (the constructor seed) is only a fallback for sessions that fail before
/// `session.created` arrives.
litellm_call_id: String,
/// See the request-id rule above — mirrors `litellm_call_id`.
id: String,
model: String,
custom_llm_provider: String,
usage: Usage,
response_cost: f64,
start_time: f64,
end_time: f64,
metadata: RequestMetadata,
/// Count of logging callbacks that failed to enqueue (non-fatal).
dropped: u64,
}
impl RealTimeStreaming {
/// Create a collector for one session. `litellm_call_id` is the gateway's
/// per-connection id; `model` is the requested model (a sane default until
/// `session.created` reports the upstream model).
pub fn new(
callbacks: Vec<Arc<dyn CustomLogger>>,
litellm_call_id: String,
model: String,
metadata: RequestMetadata,
) -> Self {
let now = epoch_seconds();
Self {
callbacks,
id: litellm_call_id.clone(),
litellm_call_id,
model,
custom_llm_provider: DEFAULT_PROVIDER.to_string(),
usage: Usage::default(),
response_cost: 0.0,
start_time: now,
end_time: now,
metadata,
dropped: 0,
}
}
/// Number of logging callbacks that failed to enqueue so far (test/observ.).
#[allow(dead_code)]
pub fn dropped(&self) -> u64 {
self.dropped
}
/// Observe one realtime event. O(1): updates accumulated state only; never
/// buffers frames. Safe to call on every event in either direction.
pub fn observe(&mut self, event: &RealtimeEvent) {
match event.event_type.as_str() {
"session.created" | "session.updated" => self.on_session(event),
"response.done" => self.on_response_done(event),
_ => {}
}
}
/// `session.created` / `session.updated` → capture upstream id + model.
/// Per the request-id rule, the OpenAI session id becomes BOTH `id` and
/// `litellm_call_id`, replacing the gateway-generated fallback.
fn on_session(&mut self, event: &RealtimeEvent) {
let session = event.data.get("session").and_then(Value::as_object);
if let Some(id) = session.and_then(|s| s.get("id")).and_then(Value::as_str) {
if !id.is_empty() {
self.id = id.to_string();
self.litellm_call_id = id.to_string();
}
}
if let Some(model) = session.and_then(|s| s.get("model")).and_then(Value::as_str) {
if !model.is_empty() {
self.model = model.to_string();
}
}
}
/// `response.done` → add this response's usage to the cumulative totals.
fn on_response_done(&mut self, event: &RealtimeEvent) {
let usage = event
.data
.get("response")
.and_then(Value::as_object)
.and_then(|r| r.get("usage"))
.and_then(Value::as_object);
let Some(usage) = usage else { return };
let input = usage.get("input_tokens").and_then(Value::as_u64);
let output = usage.get("output_tokens").and_then(Value::as_u64);
let total = usage.get("total_tokens").and_then(Value::as_u64);
if let Some(input) = input {
self.usage.prompt_tokens += input;
}
if let Some(output) = output {
self.usage.completion_tokens += output;
}
// Prefer the upstream-reported total; otherwise derive it.
match total {
Some(total) => self.usage.total_tokens += total,
None => {
self.usage.total_tokens += input.unwrap_or(0) + output.unwrap_or(0);
}
}
}
/// Set the per-session response cost ($). Cost computation is Python-side in
/// the proxy; the gateway forwards 0.0 by default and lets the proxy price.
/// Public API (exercised in tests) for the future path where the gateway
/// prices realtime sessions itself.
#[allow(dead_code)]
pub fn set_response_cost(&mut self, cost: f64) {
self.response_cost = cost;
}
/// Build the `StandardLoggingPayload` from accumulated state.
pub fn build_payload(&self) -> StandardLoggingPayload {
StandardLoggingPayload {
id: self.id.clone(),
litellm_call_id: self.litellm_call_id.clone(),
call_type: "realtime".to_string(),
model: self.model.clone(),
custom_llm_provider: self.custom_llm_provider.clone(),
response_cost: self.response_cost,
prompt_tokens: self.usage.prompt_tokens,
completion_tokens: self.usage.completion_tokens,
total_tokens: self.usage.total_tokens,
start_time: self.start_time,
end_time: self.end_time,
stream: true,
metadata: StandardLoggingMetadata {
user_api_key_hash: self.metadata.user_api_key_hash.clone(),
user_api_key_user_id: self.metadata.user_api_key_user_id.clone(),
user_api_key_team_id: self.metadata.user_api_key_team_id.clone(),
..Default::default()
},
messages: None,
}
}
/// Finish the session: stamp the end time and fan the payload out to every
/// callback. On a logger enqueue error we bump a non-fatal counter (the
/// realtime session has already ended; a dropped log must never propagate).
pub fn log_messages(&mut self, status: SessionStatus) {
self.end_time = epoch_seconds();
let payload = self.build_payload();
match status {
SessionStatus::Success => {
for callback in &self.callbacks {
if let Err(err) = callback.log_success_event(&payload) {
self.dropped += 1;
eprintln!("litellm-ai-gateway: log_success_event dropped: {err}");
}
}
}
SessionStatus::Failure => {
let error = crate::integrations::types::LoggingError {
message: "realtime session ended in failure".to_string(),
kind: "RealtimeSessionError".to_string(),
};
for callback in &self.callbacks {
if let Err(err) = callback.log_failure_event(&payload, &error) {
self.dropped += 1;
eprintln!("litellm-ai-gateway: log_failure_event dropped: {err}");
}
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::integrations::types::{LogError, LoggingError};
use std::sync::atomic::{AtomicU64, Ordering};
fn event(raw: &str) -> RealtimeEvent {
serde_json::from_str(raw).expect("valid event json")
}
/// A test logger that records the last payload it saw.
#[derive(Default)]
struct CapturingLogger {
calls: AtomicU64,
last_model: std::sync::Mutex<Option<String>>,
last_total_tokens: AtomicU64,
}
impl CustomLogger for CapturingLogger {
fn log_success_event(&self, payload: &StandardLoggingPayload) -> Result<(), LogError> {
self.calls.fetch_add(1, Ordering::SeqCst);
*self.last_model.lock().unwrap() = Some(payload.model.clone());
self.last_total_tokens
.store(payload.total_tokens, Ordering::SeqCst);
Ok(())
}
}
#[test]
fn observe_accumulates_model_and_tokens_then_logs() {
let logger = Arc::new(CapturingLogger::default());
let callbacks: Vec<Arc<dyn CustomLogger>> = vec![logger.clone()];
let mut streaming = RealTimeStreaming::new(
callbacks,
"call_abc".to_string(),
"gpt-realtime".to_string(),
RequestMetadata {
user_api_key_hash: Some("hash123".to_string()),
user_api_key_user_id: Some("user-1".to_string()),
user_api_key_team_id: Some("team-1".to_string()),
},
);
streaming.observe(&event(
r#"{"type":"session.created","session":{"id":"sess_001","model":"gpt-realtime-2025"}}"#,
));
streaming.observe(&event(
r#"{"type":"response.done","response":{"usage":{"input_tokens":10,"output_tokens":5,"total_tokens":15}}}"#,
));
// A second response.done accumulates.
streaming.observe(&event(
r#"{"type":"response.done","response":{"usage":{"input_tokens":3,"output_tokens":2,"total_tokens":5}}}"#,
));
let payload = streaming.build_payload();
assert_eq!(payload.model, "gpt-realtime-2025");
// Request-id rule: session.created's id becomes BOTH id and
// litellm_call_id (replacing the "call_abc" gateway fallback), so the
// SpendLogs request_id is always the OpenAI session id.
assert_eq!(payload.id, "sess_001");
assert_eq!(payload.litellm_call_id, "sess_001");
assert_eq!(payload.prompt_tokens, 13);
assert_eq!(payload.completion_tokens, 7);
assert_eq!(payload.total_tokens, 20);
assert_eq!(payload.response_cost, 0.0);
assert_eq!(payload.call_type, "realtime");
assert_eq!(payload.custom_llm_provider, "openai");
assert_eq!(
payload.metadata.user_api_key_hash.as_deref(),
Some("hash123")
);
streaming.log_messages(SessionStatus::Success);
assert_eq!(logger.calls.load(Ordering::SeqCst), 1);
assert_eq!(
logger.last_model.lock().unwrap().as_deref(),
Some("gpt-realtime-2025")
);
assert_eq!(logger.last_total_tokens.load(Ordering::SeqCst), 20);
assert_eq!(streaming.dropped(), 0);
}
#[test]
fn payload_serializes_with_camelcase_times_and_realtime_call_type() {
let mut streaming = RealTimeStreaming::new(
Vec::new(),
"call_xyz".to_string(),
"gpt-realtime".to_string(),
RequestMetadata::default(),
);
streaming.observe(&event(
r#"{"type":"response.done","response":{"usage":{"input_tokens":1,"output_tokens":1,"total_tokens":2}}}"#,
));
streaming.set_response_cost(0.0042);
let payload = streaming.build_payload();
let json = serde_json::to_string(&payload).expect("serialize payload");
assert!(json.contains("\"startTime\""), "missing startTime: {json}");
assert!(json.contains("\"endTime\""), "missing endTime: {json}");
assert!(
json.contains("\"call_type\":\"realtime\""),
"missing call_type realtime: {json}"
);
assert!(
json.contains("\"response_cost\""),
"missing response_cost: {json}"
);
assert_eq!(payload.response_cost, 0.0042);
}
/// A logger whose enqueue always fails should bump the dropped counter, not
/// panic or propagate.
#[test]
fn failing_logger_bumps_dropped_counter() {
struct FailingLogger;
impl CustomLogger for FailingLogger {
fn log_success_event(&self, _p: &StandardLoggingPayload) -> Result<(), LogError> {
Err(LogError::channel_full())
}
fn log_failure_event(
&self,
_p: &StandardLoggingPayload,
_e: &LoggingError,
) -> Result<(), LogError> {
Err(LogError::channel_closed())
}
}
let callbacks: Vec<Arc<dyn CustomLogger>> = vec![Arc::new(FailingLogger)];
let mut streaming = RealTimeStreaming::new(
callbacks,
"call_1".to_string(),
"gpt-realtime".to_string(),
RequestMetadata::default(),
);
streaming.log_messages(SessionStatus::Success);
assert_eq!(streaming.dropped(), 1);
}
}

View file

@ -5,7 +5,9 @@
mod service;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::time::{SystemTime, UNIX_EPOCH};
use crate::io::realtime_pool::RealtimePool;
use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
@ -20,8 +22,26 @@ use litellm_core::router::Router as ModelRouter;
use serde::Deserialize;
use crate::auth::UserApiKeyAuth;
use crate::integrations::custom_logger::CustomLogger;
use crate::integrations::types::RequestMetadata;
use crate::realtime::streaming::{RealTimeStreaming, SessionStatus};
use crate::state::AppState;
/// Process-local monotonic counter, mixed into the per-session call id so two
/// sessions opened in the same nanosecond still get distinct ids.
static CALL_SEQ: AtomicU64 = AtomicU64::new(0);
/// Generate a per-connection `litellm_call_id`. No external uuid dep: epoch
/// nanos + a process-local sequence is unique enough for log correlation.
fn new_call_id() -> String {
let nanos = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_nanos())
.unwrap_or(0);
let seq = CALL_SEQ.fetch_add(1, Ordering::Relaxed);
format!("rt-{nanos:x}-{seq:x}")
}
/// This route's contribution to the app router.
pub fn router() -> Router<AppState> {
Router::new().route("/v1/realtime", get(handle))
@ -32,31 +52,30 @@ struct RealtimeQuery {
model: String,
}
fn require_realtime_billing_safe_auth(auth: &UserApiKeyAuth) -> Result<(), (StatusCode, String)> {
if auth.is_proxy_admin() {
return Ok(());
fn request_metadata_for_auth(auth: &UserApiKeyAuth, master_key: Option<&str>) -> RequestMetadata {
let user_api_key_hash = auth.api_key.clone().or_else(|| {
auth.is_proxy_admin()
.then(|| master_key.map(crate::auth::hash_token))
.flatten()
});
RequestMetadata {
user_api_key_hash,
user_api_key_user_id: auth.user_id.clone(),
user_api_key_team_id: auth.team_id.clone(),
}
Err((
StatusCode::FORBIDDEN,
"non-admin virtual-key realtime is disabled until realtime usage is reported to the control plane".to_string(),
))
}
/// Auth runs via the `UserApiKeyAuth` extractor (master key → admin, otherwise
/// cache then the swappable authenticator). Non-admin virtual keys are rejected
/// until realtime usage is reported back to the proxy; otherwise callers could
/// consume upstream realtime spend without that spend being charged to their key
/// or team. We validate the model BEFORE the upgrade so failures are clean HTTP
/// (400/404), not a socket that opens then closes, then hand the socket to
/// `bridge`.
/// cache then the swappable authenticator). We validate the model BEFORE the
/// upgrade so failures are clean HTTP (400/404), not a socket that opens then
/// closes, then hand the socket to `bridge`.
async fn handle(
auth: UserApiKeyAuth,
ws: WebSocketUpgrade,
State(state): State<AppState>,
Query(query): Query<RealtimeQuery>,
) -> Result<Response, (StatusCode, String)> {
require_realtime_billing_safe_auth(&auth)?;
if query.model.trim().is_empty() {
return Err((
StatusCode::BAD_REQUEST,
@ -72,26 +91,58 @@ async fn handle(
let router = state.router.clone();
let pool = state.realtime_pool.clone();
let loggers = state.loggers.clone();
let master_key = state.master_key.clone();
let model = query.model;
Ok(ws.on_upgrade(move |socket| bridge(socket, router, pool, model)))
Ok(ws.on_upgrade(move |socket| bridge(socket, router, pool, loggers, auth, master_key, model)))
}
/// Adapt the axum socket (text frames) to the typed-event `Stream`/`Sink` the
/// service wants, keeping axum types out of `service`.
///
/// This is also the realtime-logging seam: every upstream→client event (the
/// direction carrying `session.created` and `response.done` with usage) is fed
/// to a [`RealTimeStreaming`] collector via the splice's `observe` callback. The
/// observe is O(1) and never buffers frames. When the splice returns (any of the
/// three break paths — client disconnect, upstream close, idle timeout), we flush
/// one logging payload to the registered callbacks.
async fn bridge(
socket: WebSocket,
router: Arc<ModelRouter>,
pool: Arc<RealtimePool>,
loggers: Arc<Vec<Arc<dyn CustomLogger>>>,
auth: UserApiKeyAuth,
master_key: Option<Arc<str>>,
model: String,
) {
let (ws_sink, ws_stream) = socket.split();
// Attribute the spend log to the key that authenticated this session. For a
// virtual key, the Python verifier returns the already-hashed key + user/team
// attribution. For the local master-key fast path, hash the master key here.
// A non-null user_api_key_hash is required for the Python spend logger to
// write a SpendLogs row.
let metadata = request_metadata_for_auth(&auth, master_key.as_deref());
// Owned by THIS task only. The splice observes it via a synchronous `&mut`
// callback (below), so there is no Arc/Mutex/atomic on the per-frame hot
// path — just a monomorphized FnMut mutating stack-local fields. This is
// what lets observe scale: 10K concurrent sessions = 10K independent
// collectors, zero cross-task synchronization.
let mut collector = RealTimeStreaming::new(
loggers.as_ref().clone(),
new_call_id(),
model.clone(),
metadata,
);
let client_in = ws_stream.filter_map(|message| async move {
match message {
Ok(Message::Text(text)) => serde_json::from_str::<RealtimeEvent>(&text).ok(),
_ => None,
}
});
// Plain forwarding sink — no observe here anymore.
let client_out = ws_sink.with(|event: RealtimeEvent| async move {
Ok::<Message, axum::Error>(Message::Text(
serde_json::to_string(&event).unwrap_or_default(),
@ -99,7 +150,28 @@ async fn bridge(
});
futures_util::pin_mut!(client_in, client_out);
let _ = service::run(&router, &pool, &model, None, client_in, client_out).await;
// The observe closure borrows `&mut collector` for the duration of the
// splice; the borrow ends when `run` returns, freeing the collector for the
// single post-session `log_messages` flush. `run` picks a pooled (warm) or
// fresh upstream — observe fires on the upstream arm either way.
let result = service::run(
&router,
&pool,
&model,
None,
|event: &RealtimeEvent| collector.observe(event),
client_in,
client_out,
)
.await;
let status = if result.is_ok() {
SessionStatus::Success
} else {
SessionStatus::Failure
};
collector.log_messages(status);
}
#[cfg(test)]
@ -107,19 +179,30 @@ mod tests {
use super::*;
#[test]
fn realtime_allows_proxy_admin_identity() {
assert!(require_realtime_billing_safe_auth(&UserApiKeyAuth::admin()).is_ok());
fn request_metadata_uses_virtual_key_identity() {
let metadata = request_metadata_for_auth(
&UserApiKeyAuth {
api_key: Some("hashed-key".to_string()),
user_id: Some("user-1".to_string()),
team_id: Some("team-1".to_string()),
..UserApiKeyAuth::default()
},
Some("sk-master"),
);
assert_eq!(metadata.user_api_key_hash.as_deref(), Some("hashed-key"));
assert_eq!(metadata.user_api_key_user_id.as_deref(), Some("user-1"));
assert_eq!(metadata.user_api_key_team_id.as_deref(), Some("team-1"));
}
#[test]
fn realtime_rejects_non_admin_virtual_keys_until_spend_is_reported() {
let err = require_realtime_billing_safe_auth(&UserApiKeyAuth {
user_role: Some("internal_user".to_string()),
..UserApiKeyAuth::default()
})
.unwrap_err();
fn request_metadata_hashes_master_key_for_local_admin_fast_path() {
let metadata = request_metadata_for_auth(&UserApiKeyAuth::admin(), Some("sk-master"));
let expected_hash = crate::auth::hash_token("sk-master");
assert_eq!(err.0, StatusCode::FORBIDDEN);
assert!(err.1.contains("usage is reported"));
assert_eq!(
metadata.user_api_key_hash.as_deref(),
Some(expected_hash.as_str())
);
}
}

View file

@ -26,6 +26,7 @@ pub async fn run<In, Out>(
pool: &RealtimePool,
model: &str,
idle_timeout: Option<Duration>,
observe: impl FnMut(&RealtimeEvent) + Send,
client_in: In,
client_out: Out,
) -> CoreResult<()>
@ -56,6 +57,7 @@ where
provider_model,
handoff,
idle_timeout,
observe,
client_in,
client_out,
)
@ -69,6 +71,7 @@ where
params.api_key.as_deref(),
params.api_base.as_deref(),
idle_timeout,
observe,
client_in,
client_out,
)

View file

@ -3,6 +3,8 @@ use std::sync::Arc;
use crate::io::realtime_pool::RealtimePool;
use litellm_core::router::Router;
use crate::integrations::custom_logger::CustomLogger;
/// Shared application state handed to every route handler.
#[derive(Clone)]
pub struct AppState {
@ -10,6 +12,8 @@ pub struct AppState {
/// The gateway master key. Any caller presenting it as a bearer token may
/// invoke the gateway. `None` → auth not configured (routes fail closed).
pub master_key: Option<Arc<str>>,
/// 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
/// (`RealtimePool::disabled()`) when `REALTIME_POOL_SIZE=0`, in which case
/// every realtime connect fresh-dials exactly as before.

View file

@ -1,5 +1,5 @@
import json
from typing import Any, List, Literal, Optional, Tuple
from typing import Any, Iterator, List, Literal, Optional, Tuple
import litellm
from litellm._logging import verbose_logger
@ -314,6 +314,70 @@ def _get_file_content_as_dictionary(file_content: bytes) -> List[dict]:
raise e
def _iter_batch_input_lines(file_content: bytes) -> Iterator[bytes]:
"""
Yield non-empty JSONL lines (unparsed) one at a time, so a caller can parse
each row in its own try/except and a single malformed line cannot abort the
whole pass. Peak memory stays bounded for large batch files.
"""
start, length, newline = 0, len(file_content), ord("\n")
while start < length:
idx = file_content.find(newline, start)
if idx == -1:
chunk, start = file_content[start:], length
else:
chunk, start = file_content[start:idx], idx + 1
line = chunk.strip()
if line:
yield line
def _iter_batch_input_entries(file_content: bytes) -> Iterator[dict]:
"""
Yield parsed batch input JSONL entries one at a time without materializing the
whole file as a list, so peak memory stays bounded. Raises on a malformed line;
callers that must survive bad rows should iterate ``_iter_batch_input_lines``
and parse per-row instead.
"""
for line in _iter_batch_input_lines(file_content):
yield json.loads(line)
# A batch request's input tokens scale roughly with its serialized size, so this
# is a conservative per-row fallback when the token counter cannot measure a row.
_BATCH_TOKEN_ESTIMATE_BYTES_PER_TOKEN = 4
def _estimate_batch_entry_tokens(raw_line: bytes) -> int:
"""Conservative token estimate for a batch row the token counter cannot measure
(or that cannot be parsed). Keeps the batch token total non-zero so a crafted
row cannot evade the TPM limit, without hard-rejecting a legitimate batch."""
return max(1, len(raw_line) // _BATCH_TOKEN_ESTIMATE_BYTES_PER_TOKEN)
def _count_entry_tokens(
entry: dict,
model_name: Optional[str] = None,
) -> int:
"""Token-count a single batch input entry's body (chat / text / embedding)."""
body = entry.get("body", {}) or {}
model = body.get("model", model_name or "")
messages = body.get("messages")
if messages:
return token_counter(model=model, messages=messages)
prompt = body.get("prompt")
if prompt:
return _count_prompt_or_input_tokens(model=model, value=prompt)
input_data = body.get("input")
if input_data:
return _count_prompt_or_input_tokens(model=model, value=input_data)
return 0
def _get_batch_job_cost_from_file_content(
file_content_dictionary: List[dict],
custom_llm_provider: Literal[
@ -396,70 +460,6 @@ def _get_batch_job_total_usage_from_file_content(
)
def _get_models_from_batch_input_file_content(
file_content_dictionary: List[dict],
) -> List[str]:
"""Extract the distinct ``body.model`` values from a batch *input* file.
Used by the proxy's batch pre-call hook to enforce that the caller is
authorized for every model named inside the JSONL — not just the one
on the outer request — so the proxy's per-key model allowlist isn't
bypassed by smuggling expensive models into the batch file.
"""
models: List[str] = []
seen: set = set()
for _item in file_content_dictionary:
body = _item.get("body") or {}
model = body.get("model")
if model and model not in seen:
seen.add(model)
models.append(model)
return models
def _get_batch_job_input_file_usage(
file_content_dictionary: List[dict],
custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai",
model_name: Optional[str] = None,
) -> Usage:
"""
Count the number of tokens in the input file
Used for batch rate limiting to count the number of tokens in the input file
"""
prompt_tokens: int = 0
completion_tokens: int = 0
for _item in file_content_dictionary:
body = _item.get("body", {})
model = body.get("model", model_name or "")
# Chat completion payloads.
messages = body.get("messages")
if messages:
prompt_tokens += token_counter(model=model, messages=messages)
continue
# Text completion payloads (`prompt`).
prompt = body.get("prompt")
if prompt:
prompt_tokens += _count_prompt_or_input_tokens(model=model, value=prompt)
continue
# Embedding payloads (`input`).
input_data = body.get("input")
if input_data:
prompt_tokens += _count_prompt_or_input_tokens(
model=model, value=input_data
)
return Usage(
total_tokens=prompt_tokens + completion_tokens,
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
)
def _count_prompt_or_input_tokens(model: str, value: Any) -> int:
"""Token-count a ``prompt`` / ``input`` field that the OpenAI batch
schema allows in four shapes:

View file

@ -30,6 +30,9 @@ DEFAULT_SQS_BATCH_SIZE = int(os.getenv("DEFAULT_SQS_BATCH_SIZE", 512))
SQS_SEND_MESSAGE_ACTION = "SendMessage"
SQS_API_VERSION = "2012-11-05"
DEFAULT_MAX_RETRIES = int(os.getenv("DEFAULT_MAX_RETRIES", 2))
# Max records accepted in one POST /v1/callbacks/logs batch. Bounds the blast
# radius: each record fans out to spend logs + every callback integration.
MAX_CALLBACK_LOG_RECORDS = 1000
DEFAULT_MAX_RECURSE_DEPTH = int(os.getenv("DEFAULT_MAX_RECURSE_DEPTH", 100))
DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER = int(
os.getenv("DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER", 10)

View file

@ -224,6 +224,7 @@ class MCPClient:
extra_headers: Optional[Dict[str, str]] = None,
ssl_verify: Optional[VerifyTypes] = None,
aws_auth: Optional[httpx.Auth] = None,
resolved_auth: Optional[httpx.Auth] = None,
sampling_callback: Optional[Callable] = None,
elicitation_callback: Optional[Callable] = None,
logging_callback: Optional[Callable] = None,
@ -237,6 +238,9 @@ class MCPClient:
self.extra_headers: Optional[Dict[str, str]] = extra_headers
self.ssl_verify: Optional[VerifyTypes] = ssl_verify
self._aws_auth: Optional[httpx.Auth] = aws_auth
# A pre-resolved httpx.Auth (e.g. from the v2 credential resolver) attached to the
# upstream client's auth= slot, taking precedence over the SigV4 aws_auth.
self._resolved_auth: Optional[httpx.Auth] = resolved_auth
self._last_initialize_instructions: Optional[str] = None
self._sampling_callback: Optional[Callable] = sampling_callback
self._elicitation_callback: Optional[Callable] = elicitation_callback
@ -482,11 +486,15 @@ class MCPClient:
verbose_logger.debug(
f"MCP client using SSL configuration: {type(ssl_config).__name__}"
)
# Use SigV4 auth if configured and no explicit auth provided.
# The MCP SDK's sse_client and streamable_http_client call this
# factory without passing auth=, so self._aws_auth is used.
# For non-SigV4 clients, self._aws_auth is None — no behavior change.
effective_auth = auth if auth is not None else self._aws_auth
# The MCP SDK's sse_client and streamable_http_client call this factory without
# passing auth=, so the fallback is used: a v2-resolved auth if present, else the
# SigV4 aws_auth. Both are None for the common case — no behavior change.
fallback_auth = (
self._resolved_auth
if self._resolved_auth is not None
else self._aws_auth
)
effective_auth = auth if auth is not None else fallback_auth
return httpx.AsyncClient(
headers=headers,
timeout=timeout,

View file

@ -3,6 +3,22 @@ from typing import Optional
from litellm.types.llms.openai import CreateFileRequest
from litellm.types.utils import ExtractedFileData
# MIME types a .jsonl batch upload is plausibly labeled with. Clients are
# inconsistent (text/plain, application/json, octet-stream, ndjson, ...), so a
# batch file must not silently bypass the streaming path just because of its
# declared type. ``purpose == "batch"`` is the authoritative signal; non-JSONL
# content still fails loudly when the rows are parsed.
_BATCH_JSONL_CONTENT_TYPES = frozenset(
{
"application/jsonl",
"application/json",
"application/octet-stream",
"application/x-ndjson",
"application/x-jsonlines",
"text/plain",
}
)
class FilesAPIUtils:
"""
@ -24,9 +40,24 @@ class FilesAPIUtils:
and extracted_file_data.get("content") is not None
)
@staticmethod
def is_batch_jsonl_request(
create_file_data: CreateFileRequest, content_type: Optional[str]
) -> bool:
"""
Batch-jsonl check from metadata only, so the body can stay a streamable
Path/handle instead of being read into memory.
"""
return (
create_file_data.get("purpose") == "batch"
and FilesAPIUtils.valid_content_type(content_type)
and create_file_data.get("file") is not None
)
@staticmethod
def valid_content_type(content_type: Optional[str]) -> bool:
"""
Check if the content type is valid
Whether the upload's MIME type is one a batch JSONL file is plausibly
sent as (see ``_BATCH_JSONL_CONTENT_TYPES``).
"""
return content_type in set(["application/jsonl", "application/octet-stream"])
return content_type in _BATCH_JSONL_CONTENT_TYPES

View file

@ -14,6 +14,7 @@ OPTIONAL_KWARGS_KEYS = frozenset(
"azure_password",
"azure_scope",
"timeout",
"gcs_bucket_name",
"bucket_name",
"vertex_credentials",
"vertex_project",

View file

@ -757,6 +757,46 @@ def update_responses_tools_with_model_file_ids(
return updated_tools
def extract_file_metadata(file_data: FileTypes) -> Tuple[Optional[str], Optional[str]]:
"""
Resolve (filename, content_type) without reading the file body.
Mirrors extract_file_data's metadata resolution but never calls .read(), so
it stays O(1) on large uploads. Use this when only metadata is needed (batch
detection, GCS object naming) and the body must remain a streamable Path/handle.
"""
filename: Optional[str] = None
content_type: Optional[str] = None
file_content: Any = None
if isinstance(file_data, tuple):
if len(file_data) == 2:
filename, file_content = file_data
elif len(file_data) == 3:
filename, file_content, content_type = file_data
elif len(file_data) == 4:
filename, file_content, content_type, _ = file_data
elif isinstance(file_data, InMemoryFile):
filename = file_data.name
content_type = file_data.content_type
else:
file_content = file_data
if filename is None:
if isinstance(file_content, PathLike):
filename = Path(file_content).name
elif isinstance(file_content, io.IOBase):
name_attr = getattr(file_content, "name", None)
if isinstance(name_attr, str):
filename = Path(name_attr).name
if not content_type:
guessed = mimetypes.guess_type(filename)[0] if filename else None
content_type = guessed or "application/octet-stream"
return filename, content_type
def extract_file_data(file_data: FileTypes) -> ExtractedFileData:
"""
Extracts and processes file data from various input formats.

View file

@ -1,5 +1,5 @@
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
from typing import TYPE_CHECKING, Any, Dict, Iterator, List, Optional, Union
import httpx
from openai.types.file_deleted import FileDeleted
@ -32,6 +32,22 @@ else:
Router = Any
class BaseFileUploadStream(ABC):
"""Re-iterable request body that yields an upload's bytes lazily.
A provider returns one of these (inside the upload config from
``transform_create_file_request``) when the upload body can be produced
incrementally; the HTTP handler then sends it in bounded chunks instead of
buffering the whole payload, which is what exhausts memory on large uploads.
``iter_bytes`` must return a fresh iterator each call so the body can be
replayed if the upload is retried.
"""
@abstractmethod
def iter_bytes(self) -> Iterator[bytes]: ...
class BaseFilesConfig(BaseConfig):
@property
@abstractmethod

View file

@ -1,13 +1,14 @@
import asyncio
import json
import ssl
from functools import lru_cache
from urllib.parse import parse_qs, urlencode, urlparse, urlunparse
from typing import (
TYPE_CHECKING,
Any,
AsyncIterator,
Coroutine,
Dict,
Iterator,
List,
Literal,
Optional,
@ -16,6 +17,7 @@ from typing import (
cast,
get_type_hints,
)
from urllib.parse import parse_qs, urlencode, urlparse, urlunparse
import httpx # type: ignore
from openai.types.file_deleted import FileDeleted
@ -27,8 +29,8 @@ import litellm.types.utils
from litellm._logging import _redact_string, verbose_logger
from litellm.anthropic_beta_headers_manager import update_headers_with_filtered_beta
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming
from litellm.litellm_core_utils.asyncify import run_async_function
from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
from litellm.llms.base_llm.anthropic_messages.transformation import (
BaseAnthropicMessagesConfig,
@ -107,6 +109,7 @@ from litellm.types.llms.openai import (
ResponsesAPIOptionalRequestParams,
ResponsesAPIResponse,
)
from litellm.types.realtime import RealtimeQueryParams
from litellm.types.rerank import RerankResponse
from litellm.types.responses.main import DeleteResponseResult
from litellm.types.router import GenericLiteLLMParams
@ -132,7 +135,6 @@ from litellm.types.vector_stores import (
VectorStoreSearchOptionalRequestParams,
VectorStoreSearchResponse,
)
from litellm.types.realtime import RealtimeQueryParams
from litellm.types.videos.main import VideoObject
from litellm.utils import (
CustomStreamWrapper,
@ -3438,6 +3440,23 @@ class BaseLLMHTTPHandler:
data=presigned_request["data"],
timeout=timeout,
)
elif (
isinstance(transformed_request, dict)
and "resumable_chunked_upload" in transformed_request
):
try:
upload_response = self._resumable_chunked_upload(
client=sync_httpx_client,
initiate_url=api_base,
base_headers=headers,
config=cast(Dict[str, Any], transformed_request)[
"resumable_chunked_upload"
],
timeout=timeout,
)
except Exception as e:
verbose_logger.exception(f"Error creating file: {e}")
raise self._handle_error(e=e, provider_config=provider_config)
elif isinstance(transformed_request, str) or isinstance(
transformed_request, bytes
):
@ -3519,7 +3538,15 @@ class BaseLLMHTTPHandler:
input="",
api_key="",
additional_args={
"complete_input_dict": transformed_request,
# A resumable upload config holds a reference to the (potentially
# huge) upload payload; logging deep-copies additional_args, so log
# a placeholder instead of re-materializing the payload.
"complete_input_dict": (
"<resumable chunked upload>"
if isinstance(transformed_request, dict)
and "resumable_chunked_upload" in transformed_request
else transformed_request
),
"api_base": api_base,
"headers": headers,
},
@ -3596,6 +3623,23 @@ class BaseLLMHTTPHandler:
data=presigned_request["data"],
timeout=timeout,
)
elif (
isinstance(transformed_request, dict)
and "resumable_chunked_upload" in transformed_request
):
try:
upload_response = await self._aresumable_chunked_upload(
client=async_httpx_client,
initiate_url=api_base,
base_headers=headers,
config=cast(Dict[str, Any], transformed_request)[
"resumable_chunked_upload"
],
timeout=timeout,
)
except Exception as e:
verbose_logger.exception(f"Error creating file: {e}")
raise self._handle_error(e=e, provider_config=provider_config)
elif isinstance(transformed_request, str) or isinstance(
transformed_request, bytes
):
@ -3639,6 +3683,224 @@ class BaseLLMHTTPHandler:
litellm_params=litellm_params,
)
# 8 MiB; a 256 KiB multiple, which GCS requires for every non-final chunk.
_RESUMABLE_CHUNK_SIZE = 8 * 1024 * 1024
@staticmethod
def _iter_resumable_chunks(
byte_iter: Iterator[bytes], chunk_size: int
) -> Iterator[bytes]:
"""Regroup a byte stream into ``chunk_size`` pieces, yielding a final
partial piece only when it is non-empty. Every full piece is exactly
``chunk_size`` bytes (kept a 256 KiB multiple for GCS) and never more than
one chunk is buffered. An exactly chunk-aligned stream yields only full
chunks, so the upload finalizes on its last data chunk instead of making
an extra empty request; a 0-byte stream yields nothing and the caller
finalizes with a single empty request.
"""
buf = bytearray()
for piece in byte_iter:
buf.extend(piece)
while len(buf) >= chunk_size:
yield bytes(buf[:chunk_size])
del buf[:chunk_size]
if buf:
yield bytes(buf)
@staticmethod
def _resumable_content_range(offset: int, data_len: int, is_final: bool) -> str:
if not is_final:
return f"bytes {offset}-{offset + data_len - 1}/*"
total = offset + data_len
if data_len == 0:
return f"bytes */{total}"
return f"bytes {offset}-{total - 1}/{total}"
@staticmethod
def _resumable_request_kwargs(
headers: dict,
content: bytes,
timeout: Optional[Union[float, httpx.Timeout]],
) -> dict:
kwargs: Dict[str, Any] = {"headers": headers, "content": content}
if timeout is not None:
kwargs["timeout"] = timeout
return kwargs
def _resumable_chunked_upload(
self,
*,
client: HTTPHandler,
initiate_url: str,
base_headers: dict,
config: dict,
timeout: Optional[Union[float, httpx.Timeout]],
) -> httpx.Response:
"""Open a GCS resumable session, then PUT the body in bounded chunks so a
large upload is never held in memory in full."""
stream = config["body_stream"]
chunk_size = config.get("chunk_size", self._RESUMABLE_CHUNK_SIZE)
session_url_header = config.get("session_url_header", "location")
httpx_client = client.client
init_headers = {**base_headers, **config.get("initiate_headers", {})}
init_req = httpx_client.build_request(
"POST",
initiate_url,
**self._resumable_request_kwargs(init_headers, b"", timeout),
)
init_resp = httpx_client.send(init_req, follow_redirects=False)
init_resp.read()
if init_resp.status_code not in (200, 201):
init_resp.raise_for_status()
session_url = init_resp.headers.get(session_url_header)
if not session_url:
raise ValueError(
f"resumable upload: no session URL in '{session_url_header}' header"
)
offset = 0
pending: Optional[bytes] = None
for chunk in self._iter_resumable_chunks(stream.iter_bytes(), chunk_size):
if pending is not None:
self._send_resumable_chunk(
httpx_client,
session_url,
base_headers,
pending,
offset,
is_final=False,
timeout=timeout,
)
offset += len(pending)
pending = chunk
return self._send_resumable_chunk(
httpx_client,
session_url,
base_headers,
pending or b"",
offset,
is_final=True,
timeout=timeout,
)
def _send_resumable_chunk(
self,
httpx_client: httpx.Client,
url: str,
base_headers: dict,
data: bytes,
offset: int,
*,
is_final: bool,
timeout: Optional[Union[float, httpx.Timeout]],
) -> httpx.Response:
headers = {
**base_headers,
"Content-Range": self._resumable_content_range(offset, len(data), is_final),
}
req = httpx_client.build_request(
"PUT", url, **self._resumable_request_kwargs(headers, data, timeout)
)
resp = httpx_client.send(req, follow_redirects=False)
resp.read()
if resp.status_code not in ((200, 201) if is_final else (308,)):
# 4xx/5xx raise here; the ValueError catches an unexpected success
# status (e.g. a 200 where the protocol expects a 308 between chunks).
resp.raise_for_status()
raise ValueError(f"resumable upload: unexpected status {resp.status_code}")
return resp
async def _aresumable_chunked_upload(
self,
*,
client: AsyncHTTPHandler,
initiate_url: str,
base_headers: dict,
config: dict,
timeout: Optional[Union[float, httpx.Timeout]],
) -> httpx.Response:
stream = config["body_stream"]
chunk_size = config.get("chunk_size", self._RESUMABLE_CHUNK_SIZE)
session_url_header = config.get("session_url_header", "location")
httpx_client = client.client
init_headers = {**base_headers, **config.get("initiate_headers", {})}
init_req = httpx_client.build_request(
"POST",
initiate_url,
**self._resumable_request_kwargs(init_headers, b"", timeout),
)
init_resp = await httpx_client.send(init_req, follow_redirects=False)
await init_resp.aread()
if init_resp.status_code not in (200, 201):
init_resp.raise_for_status()
session_url = init_resp.headers.get(session_url_header)
if not session_url:
raise ValueError(
f"resumable upload: no session URL in '{session_url_header}' header"
)
offset = 0
pending: Optional[bytes] = None
# Producing each chunk runs the synchronous per-row transform for that
# chunk's worth of rows. Pull it off the event loop thread so a large
# upload does not block other concurrent requests between PUTs.
chunk_iter = self._iter_resumable_chunks(stream.iter_bytes(), chunk_size)
done = object()
while True:
chunk = await asyncio.to_thread(next, chunk_iter, done)
if chunk is done:
break
if pending is not None:
await self._asend_resumable_chunk(
httpx_client,
session_url,
base_headers,
pending,
offset,
is_final=False,
timeout=timeout,
)
offset += len(pending)
pending = chunk
return await self._asend_resumable_chunk(
httpx_client,
session_url,
base_headers,
pending or b"",
offset,
is_final=True,
timeout=timeout,
)
async def _asend_resumable_chunk(
self,
httpx_client: httpx.AsyncClient,
url: str,
base_headers: dict,
data: bytes,
offset: int,
*,
is_final: bool,
timeout: Optional[Union[float, httpx.Timeout]],
) -> httpx.Response:
headers = {
**base_headers,
"Content-Range": self._resumable_content_range(offset, len(data), is_final),
}
req = httpx_client.build_request(
"PUT", url, **self._resumable_request_kwargs(headers, data, timeout)
)
resp = await httpx_client.send(req, follow_redirects=False)
await resp.aread()
if resp.status_code not in ((200, 201) if is_final else (308,)):
# 4xx/5xx raise here; the ValueError catches an unexpected success
# status (e.g. a 200 where the protocol expects a 308 between chunks).
resp.raise_for_status()
raise ValueError(f"resumable upload: unexpected status {resp.status_code}")
return resp
def create_batch(
self,
create_batch_data: "CreateBatchRequest",

View file

@ -17,17 +17,13 @@ from litellm.litellm_core_utils.cloud_storage_security import (
)
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.openai import (
CreateFileRequest,
FileContentRequest,
HttpxBinaryResponseContent,
OpenAIFileObject,
)
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.types.llms.vertex_ai import VERTEX_CREDENTIALS_TYPES
from .transformation import VertexAIFilesConfig, VertexAIJsonlFilesTransformation
vertex_ai_files_transformation = VertexAIJsonlFilesTransformation()
from .transformation import VertexAIFilesConfig
class VertexAIFilesHandler(GCSBucketBase):
@ -43,82 +39,6 @@ class VertexAIFilesHandler(GCSBucketBase):
llm_provider=LlmProviders.VERTEX_AI,
)
async def async_create_file(
self,
create_file_data: CreateFileRequest,
api_base: Optional[str],
vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES],
vertex_project: Optional[str],
vertex_location: Optional[str],
timeout: Union[float, httpx.Timeout],
max_retries: Optional[int],
) -> OpenAIFileObject:
gcs_logging_config: GCSLoggingConfig = await self.get_gcs_logging_config(
kwargs={}
)
headers = await self.construct_request_headers(
vertex_instance=gcs_logging_config["vertex_instance"],
service_account_json=gcs_logging_config["path_service_account"],
)
bucket_name = gcs_logging_config["bucket_name"]
(
logging_payload,
object_name,
) = vertex_ai_files_transformation.transform_openai_file_content_to_vertex_ai_file_content(
openai_file_content=create_file_data.get("file")
)
gcs_upload_response = await self._log_json_data_on_gcs(
headers=headers,
bucket_name=bucket_name,
object_name=object_name,
logging_payload=logging_payload,
)
return vertex_ai_files_transformation.transform_gcs_bucket_response_to_openai_file_object(
create_file_data=create_file_data,
gcs_upload_response=gcs_upload_response,
)
def create_file(
self,
_is_async: bool,
create_file_data: CreateFileRequest,
api_base: Optional[str],
vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES],
vertex_project: Optional[str],
vertex_location: Optional[str],
timeout: Union[float, httpx.Timeout],
max_retries: Optional[int],
) -> Union[OpenAIFileObject, Coroutine[Any, Any, OpenAIFileObject]]:
"""
Creates a file on VertexAI GCS Bucket
Only supported for Async litellm.acreate_file
"""
if _is_async:
return self.async_create_file(
create_file_data=create_file_data,
api_base=api_base,
vertex_credentials=vertex_credentials,
vertex_project=vertex_project,
vertex_location=vertex_location,
timeout=timeout,
max_retries=max_retries,
)
else:
return asyncio.run(
self.async_create_file(
create_file_data=create_file_data,
api_base=api_base,
vertex_credentials=vertex_credentials,
vertex_project=vertex_project,
vertex_location=vertex_location,
timeout=timeout,
max_retries=max_retries,
)
)
def _extract_bucket_and_object_from_file_id(
self,
file_id: str,

View file

@ -1,9 +1,21 @@
import base64
import io
import itertools
import json
import os
import re
import time
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
from typing import (
Any,
Callable,
Dict,
Iterable,
Iterator,
List,
Optional,
Tuple,
Union,
)
import httpx
from httpx import Headers, Response
@ -22,9 +34,13 @@ from litellm.litellm_core_utils.cloud_storage_security import (
validate_managed_cloud_file_id,
)
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
from litellm.litellm_core_utils.prompt_templates.common_utils import (
extract_file_data,
extract_file_metadata,
)
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.files.transformation import (
BaseFileUploadStream,
BaseFilesConfig,
LiteLLMLoggingObj,
)
@ -44,8 +60,9 @@ from litellm.types.llms.openai import (
OpenAIFileObject,
PathLike,
)
from litellm.types.files import ResumableChunkedUploadConfig
from litellm.types.llms.vertex_ai import GcsBucketResponse
from litellm.types.utils import ExtractedFileData, LlmProviders, ModelResponse
from litellm.types.utils import LlmProviders, ModelResponse
from ..common_utils import VertexAIError
from ..vertex_llm_base import VertexBase
@ -137,42 +154,140 @@ def _get_litellm_batch_custom_id_from_labels(labels: Dict[str, Any]) -> str:
return str(labels.get("litellm_custom_id", "unknown"))
def _openai_batch_jsonl_entries_to_vertex_wrapped_requests(
openai_jsonl_content: List[Dict[str, Any]],
def _openai_batch_jsonl_entry_to_vertex_wrapped_request(
openai_entry: Dict[str, Any],
map_openai_to_vertex_params: Callable[[Dict[str, Any]], Dict[str, Any]],
) -> List[Dict[str, Any]]:
) -> Dict[str, Any]:
"""
Transforms OpenAI JSONL batch entries to Vertex AI JSONL lines.
Transforms a single OpenAI JSONL batch entry into its Vertex wrapped request.
jsonl body for vertex is {"request": <request_body>}
Example Vertex jsonl
{"request":{"contents": [{"role": "user", "parts": [{"text": "What is the relation between the following video and image samples?"}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/video/animals.mp4", "mimeType": "video/mp4"}}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/image/cricket.jpeg", "mimeType": "image/jpeg"}}]}]}}
{"request":{"contents": [{"role": "user", "parts": [{"text": "Describe what is happening in this video."}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/video/another_video.mov", "mimeType": "video/mov"}}]}]}}
"""
openai_request_body = openai_entry.get("body") or {}
vertex_request_body = _transform_request_body(
messages=openai_request_body.get("messages", []),
model=openai_request_body.get("model", ""),
optional_params=map_openai_to_vertex_params(openai_request_body),
custom_llm_provider="vertex_ai",
litellm_params={},
cached_content=None,
)
custom_id = openai_entry.get("custom_id")
if custom_id is not None:
if "labels" not in vertex_request_body:
vertex_request_body["labels"] = {}
_set_litellm_batch_custom_id_labels(vertex_request_body["labels"], custom_id)
return {"request": vertex_request_body}
def _iter_stripped_lines(raw_lines: Iterable[Union[str, bytes]]) -> Iterator[str]:
"""Decode (when needed), strip, and drop blank lines from an iterable of lines."""
for raw in raw_lines:
line = raw.decode("utf-8") if isinstance(raw, (bytes, bytearray)) else raw
line = line.strip()
if line:
yield line
def _iter_openai_jsonl_lines(openai_file_content: FileTypes) -> Iterator[str]:
"""
Yield non-empty JSONL lines one at a time without materializing the whole
payload, so peak memory stays bounded regardless of payload size. Mirrors
``str.splitlines()`` + ``line.strip()`` for ``\\n`` / ``\\r\\n`` delimited
JSONL.
"""
content: Any = openai_file_content
if isinstance(content, tuple):
content = content[1]
if isinstance(content, (bytes, bytearray)):
# Scan for newlines in place so a large in-memory payload is not copied
# into a BytesIO just to iterate it line by line.
newline = ord("\n")
start, length = 0, len(content)
while start < length:
idx = content.find(newline, start)
if idx == -1:
chunk, start = content[start:], length
else:
chunk, start = content[start:idx], idx + 1
line = chunk.decode("utf-8").strip()
if line:
yield line
return
if isinstance(content, str):
yield from _iter_stripped_lines(io.StringIO(content))
return
if isinstance(content, PathLike):
with open(str(content), "rb") as handle:
yield from _iter_stripped_lines(handle)
return
if hasattr(content, "read"):
# The handle is read twice per upload (first-row probe for the GCS
# object name, then the body stream), so it must rewind to 0. A
# non-seekable handle would silently resume mid-stream and drop the
# already-consumed first row, so reject it loudly instead.
seek = getattr(content, "seek", None)
if seek is None:
raise ValueError(
"Batch upload file handle must be seekable; got a non-seekable "
"stream. Pass bytes, a path, or a seekable handle."
)
try:
seek(0)
except (OSError, ValueError) as e:
raise ValueError(
"Batch upload file handle must be seekable so it can be re-read "
"for the GCS object name and the upload body."
) from e
yield from _iter_stripped_lines(content)
return
raise ValueError("Unsupported file content type")
def _iter_openai_jsonl_entries(
openai_file_content: FileTypes,
) -> Iterator[Dict[str, Any]]:
for line in _iter_openai_jsonl_lines(openai_file_content):
yield json.loads(line)
class _OpenAIToVertexBatchUploadStream(BaseFileUploadStream):
"""Streams an OpenAI batch JSONL upload as Vertex-wrapped JSONL one row at a
time, so the transformed payload is never held in full.
The transform runs lazily as the HTTP client pulls each chunk, which keeps
peak memory at one row regardless of how large the batch file is.
"""
vertex_jsonl_content = []
for _openai_jsonl_content in openai_jsonl_content:
openai_request_body = _openai_jsonl_content.get("body") or {}
vertex_request_body = _transform_request_body(
messages=openai_request_body.get("messages", []),
model=openai_request_body.get("model", ""),
optional_params=map_openai_to_vertex_params(openai_request_body),
custom_llm_provider="vertex_ai",
litellm_params={},
cached_content=None,
)
def __init__(
self,
openai_file_content: FileTypes,
map_openai_to_vertex_params: Callable[[Dict[str, Any]], Dict[str, Any]],
) -> None:
self._openai_file_content = openai_file_content
self._map_openai_to_vertex_params = map_openai_to_vertex_params
# Add custom_id as a label for correlation in batch outputs
custom_id = _openai_jsonl_content.get("custom_id")
if custom_id is not None:
if "labels" not in vertex_request_body:
vertex_request_body["labels"] = {}
_set_litellm_batch_custom_id_labels(
vertex_request_body["labels"], custom_id
def _iter_vertex_jsonl_chunks(self) -> Iterator[bytes]:
first = True
for entry in _iter_openai_jsonl_entries(self._openai_file_content):
wrapped = _openai_batch_jsonl_entry_to_vertex_wrapped_request(
entry, self._map_openai_to_vertex_params
)
prefix = b"" if first else b"\n"
first = False
yield prefix + json.dumps(wrapped).encode("utf-8")
vertex_jsonl_content.append({"request": vertex_request_body})
return vertex_jsonl_content
def iter_bytes(self) -> Iterator[bytes]:
return self._iter_vertex_jsonl_chunks()
class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
@ -181,7 +296,6 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
"""
def __init__(self):
self.jsonl_transformation = VertexAIJsonlFilesTransformation()
super().__init__()
@property
@ -208,43 +322,6 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
headers["Authorization"] = f"Bearer {api_key}"
return headers
def _get_content_from_openai_file(self, openai_file_content: FileTypes) -> str:
"""
Helper to extract content from various OpenAI file types and return as string.
Handles:
- Direct content (str, bytes, IO[bytes])
- Tuple formats: (filename, content, [content_type], [headers])
- PathLike objects
"""
content: Union[str, bytes] = b""
# Extract file content from tuple if necessary
if isinstance(openai_file_content, tuple):
# Take the second element which is always the file content
file_content = openai_file_content[1]
else:
file_content = openai_file_content
# Handle different file content types
if isinstance(file_content, str):
# String content can be used directly
content = file_content
elif isinstance(file_content, bytes):
# Bytes content can be decoded
content = file_content
elif isinstance(file_content, PathLike): # PathLike
with open(str(file_content), "rb") as f:
content = f.read()
elif hasattr(file_content, "read"): # IO[bytes]
# File-like objects need to be read
content = file_content.read()
# Ensure content is string
if isinstance(content, bytes):
content = content.decode("utf-8")
return content
def _get_gcs_object_name_from_batch_jsonl(
self,
openai_jsonl_content: List[Dict[str, Any]],
@ -261,32 +338,21 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
object_name = f"{VERTEX_AI_MANAGED_GCS_PREFIX}{safe_model_path}/{uuid.uuid4()}"
return object_name
def get_object_name(
self, extracted_file_data: ExtractedFileData, purpose: str
) -> str:
def get_object_name(self, file_data: FileTypes, purpose: str) -> str:
"""
Get the object name for the request
Get the object name for the request.
Reads only the first JSONL entry (streamed) for batch files, so a large
upload is never materialized just to derive the GCS object name.
"""
extracted_file_data_content = extracted_file_data.get("content")
if extracted_file_data_content is None:
raise ValueError("file content is required")
if purpose == "batch":
## 1. If jsonl, check if there's a model name
file_content = self._get_content_from_openai_file(
extracted_file_data_content
)
# Split into lines and parse each line as JSON
openai_jsonl_content = [
json.loads(line) for line in file_content.splitlines() if line.strip()
]
if len(openai_jsonl_content) > 0:
return self._get_gcs_object_name_from_batch_jsonl(openai_jsonl_content)
## 1. If jsonl, derive the object name from the first entry's model
first_entry = next(_iter_openai_jsonl_entries(file_data), None)
if first_entry is not None:
return self._get_gcs_object_name_from_batch_jsonl([first_entry])
## 2. If not jsonl, store under a server-generated managed object name
filename = extracted_file_data.get("filename")
filename, _ = extract_file_metadata(file_data)
return build_managed_cloud_object_name(
prefix=f"{VERTEX_AI_MANAGED_GCS_PREFIX}uploads/",
filename=filename,
@ -294,7 +360,11 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
)
def _get_configured_bucket_name(self, litellm_params: Dict) -> str:
bucket_name = litellm_params.get("bucket_name") or os.getenv("GCS_BUCKET_NAME")
bucket_name = (
litellm_params.get("gcs_bucket_name")
or litellm_params.get("bucket_name")
or os.getenv("GCS_BUCKET_NAME")
)
if not bucket_name:
raise ValueError("GCS bucket_name is required")
return bucket_name
@ -319,12 +389,21 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
raise ValueError("file is required")
if purpose is None:
raise ValueError("purpose is required")
extracted_file_data = extract_file_data(file_data)
object_name = self.get_object_name(extracted_file_data, purpose)
_, content_type = extract_file_metadata(file_data)
object_name = self.get_object_name(file_data, purpose)
if object_prefix:
object_name = f"{object_prefix}/{object_name}"
encoded_object_name = encode_gcs_object_name_for_url(object_name)
endpoint = f"upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={encoded_object_name}"
# Batch jsonl is streamed via a resumable session (bounded memory on
# large uploads); everything else is a single simple-media upload.
upload_type = (
"resumable"
if FilesAPIUtils.is_batch_jsonl_request(
create_file_data=data, content_type=content_type
)
else "media"
)
endpoint = f"upload/storage/v1/b/{bucket_name}/o?uploadType={upload_type}&name={encoded_object_name}"
api_base = api_base or "https://storage.googleapis.com"
if not api_base:
raise ValueError("api_base is required")
@ -366,14 +445,6 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
)
return vertex_params
def _transform_openai_jsonl_content_to_vertex_ai_jsonl_content(
self, openai_jsonl_content: List[Dict[str, Any]]
) -> List[Dict[str, Any]]:
return _openai_batch_jsonl_entries_to_vertex_wrapped_requests(
openai_jsonl_content=openai_jsonl_content,
map_openai_to_vertex_params=self._map_openai_to_vertex_params,
)
def transform_create_file_request(
self,
model: str,
@ -384,40 +455,34 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
"""
2 Cases:
1. Handle basic file upload
2. Handle batch file upload (.jsonl)
2. Handle batch file upload (.jsonl), streamed to a GCS resumable
session so large uploads stay memory-bounded.
"""
file_data = create_file_data.get("file")
if file_data is None:
raise ValueError("file is required")
extracted_file_data = extract_file_data(file_data)
extracted_file_data_content = extracted_file_data.get("content")
if extracted_file_data_content is None:
raise ValueError("file content is required")
if FilesAPIUtils.is_batch_jsonl_file(
_, content_type = extract_file_metadata(file_data)
if FilesAPIUtils.is_batch_jsonl_request(
create_file_data=create_file_data,
extracted_file_data=extracted_file_data,
content_type=content_type,
):
## 1. If jsonl, check if there's a model name
file_content = self._get_content_from_openai_file(
extracted_file_data_content
)
# Split into lines and parse each line as JSON
openai_jsonl_content = [
json.loads(line) for line in file_content.splitlines() if line.strip()
]
vertex_jsonl_content = (
self._transform_openai_jsonl_content_to_vertex_ai_jsonl_content(
openai_jsonl_content
return {
"resumable_chunked_upload": ResumableChunkedUploadConfig(
body_stream=_OpenAIToVertexBatchUploadStream(
file_data,
self._map_openai_to_vertex_params,
),
initiate_headers={
"X-Upload-Content-Type": "application/json",
},
)
)
return "\n".join(json.dumps(item) for item in vertex_jsonl_content)
elif isinstance(extracted_file_data_content, bytes):
}
extracted_file_data_content = extract_file_data(file_data).get("content")
if isinstance(extracted_file_data_content, bytes):
return extracted_file_data_content
else:
raise ValueError("Unsupported file content type")
raise ValueError("Unsupported file content type")
def transform_create_file_response(
self,
@ -642,39 +707,38 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
}
"""
try:
# Decode content
content_str = content.decode("utf-8")
# Check if it's JSONL (multiple lines)
lines = content_str.strip().split("\n")
if not lines:
# Read the result file one row at a time. Batch output files can be
# as large as the (multi-GB) input, so splitting into a list of rows
# and building a second list of transformed rows peaks at several full
# copies and OOMs on retrieval.
lines = _iter_openai_jsonl_lines(content)
try:
first_line = next(lines)
except StopIteration:
return content
# Try to parse the first line to see if it's Vertex AI batch output
first_line = json.loads(lines[0])
# Check if it has Vertex AI batch output structure with discriminating fields
# Must have request, response, and processed_time
# Plus either candidates (success) or status (error)
has_base_structure = (
"response" in first_line
and "request" in first_line
and "processed_time" in first_line
# Identify a Vertex AI batch output from the first row's
# discriminating fields. Anything else (e.g. a binary file whose
# first line is not valid UTF-8/JSON) raises and falls through to the
# passthrough below, leaving the content untouched.
first_row = json.loads(first_line)
is_vertex_batch_output = (
"request" in first_row
and "response" in first_row
and "processed_time" in first_row
and (
"candidates" in first_row.get("response", {})
or "promptFeedback" in first_row.get("response", {})
or bool(first_row.get("status"))
)
)
has_success_or_error = (
"candidates" in first_line.get("response", {})
or "promptFeedback" in first_line.get("response", {})
or bool(first_line.get("status"))
)
if not (has_base_structure and has_success_or_error):
# Not a Vertex AI batch output, return as-is
if not is_vertex_batch_output:
return content
vertex_gemini_config = VertexGeminiConfig()
# Always use a fresh local Logging object for the per-line transformation
# so we never mutate the caller's logging_obj (which already went through
# pre_call and has its own model/start_time/optional_params set).
# Use a fresh Logging object for the per-row transform so we never
# mutate the caller's (which already ran pre_call with its own
# model/start_time/optional_params).
batch_transform_logging_obj = Logging(
model="",
messages=[],
@ -691,29 +755,27 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
request=httpx.Request(method="POST", url="https://example.com"),
)
# Transform all lines
transformed_lines = []
for line in lines:
if not line.strip():
continue
# Transform each row straight into the output buffer, so peak memory
# stays at ~one row plus the output. If any row fails, return the
# original content unchanged.
output = bytearray()
for line in itertools.chain([first_line], lines):
try:
vertex_output = json.loads(line)
openai_output = (
self._transform_single_vertex_batch_output_to_openai(
vertex_output=vertex_output,
vertex_output=json.loads(line),
vertex_gemini_config=vertex_gemini_config,
logging_obj=batch_transform_logging_obj,
mock_httpx_response=mock_httpx_response,
)
)
transformed_lines.append(json.dumps(openai_output))
except Exception:
# If any line fails, return original content
return content
if output:
output += b"\n"
output += json.dumps(openai_output).encode("utf-8")
# Return transformed content
return "\n".join(transformed_lines).encode("utf-8")
return bytes(output)
except Exception:
# If anything fails, return original content
@ -795,137 +857,3 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
"message": f"Failed to transform response: {str(e)}",
},
}
class VertexAIJsonlFilesTransformation(VertexGeminiConfig):
"""
Transforms OpenAI /v1/files/* requests to VertexAI /v1/files/* requests
"""
def transform_openai_file_content_to_vertex_ai_file_content(
self, openai_file_content: Optional[FileTypes] = None
) -> Tuple[str, str]:
"""
Transforms OpenAI FileContentRequest to VertexAI FileContentRequest
"""
if openai_file_content is None:
raise ValueError("contents of file are None")
# Read the content of the file
file_content = self._get_content_from_openai_file(openai_file_content)
# Split into lines and parse each line as JSON
openai_jsonl_content = [
json.loads(line) for line in file_content.splitlines() if line.strip()
]
vertex_jsonl_content = (
self._transform_openai_jsonl_content_to_vertex_ai_jsonl_content(
openai_jsonl_content
)
)
vertex_jsonl_string = "\n".join(
json.dumps(item) for item in vertex_jsonl_content
)
object_name = self._get_gcs_object_name(
openai_jsonl_content=openai_jsonl_content
)
return vertex_jsonl_string, object_name
def _transform_openai_jsonl_content_to_vertex_ai_jsonl_content(
self, openai_jsonl_content: List[Dict[str, Any]]
) -> List[Dict[str, Any]]:
return _openai_batch_jsonl_entries_to_vertex_wrapped_requests(
openai_jsonl_content=openai_jsonl_content,
map_openai_to_vertex_params=self._map_openai_to_vertex_params,
)
def _get_gcs_object_name(
self,
openai_jsonl_content: List[Dict[str, Any]],
) -> str:
"""
Gets a unique GCS object name for the VertexAI batch prediction job
named as: litellm-vertex-{model}-{uuid}
"""
_model = openai_jsonl_content[0].get("body", {}).get("model", "")
if "publishers/google/models" not in _model:
_model = f"publishers/google/models/{_model}"
safe_model_path = sanitize_cloud_object_path(_model, fallback="model")
object_name = f"{VERTEX_AI_MANAGED_GCS_PREFIX}{safe_model_path}/{uuid.uuid4()}"
return object_name
def _map_openai_to_vertex_params(
self,
openai_request_body: Dict[str, Any],
) -> Dict[str, Any]:
"""
wrapper to call VertexGeminiConfig.map_openai_params
"""
_model = openai_request_body.get("model", "")
vertex_params = self.map_openai_params(
model=_model,
non_default_params=openai_request_body,
optional_params={},
drop_params=False,
)
return vertex_params
def _get_content_from_openai_file(self, openai_file_content: FileTypes) -> str:
"""
Helper to extract content from various OpenAI file types and return as string.
Handles:
- Direct content (str, bytes, IO[bytes])
- Tuple formats: (filename, content, [content_type], [headers])
- PathLike objects
"""
content: Union[str, bytes] = b""
# Extract file content from tuple if necessary
if isinstance(openai_file_content, tuple):
# Take the second element which is always the file content
file_content = openai_file_content[1]
else:
file_content = openai_file_content
# Handle different file content types
if isinstance(file_content, str):
# String content can be used directly
content = file_content
elif isinstance(file_content, bytes):
# Bytes content can be decoded
content = file_content
elif isinstance(file_content, PathLike): # PathLike
with open(str(file_content), "rb") as f:
content = f.read()
elif hasattr(file_content, "read"): # IO[bytes]
# File-like objects need to be read
content = file_content.read()
# Ensure content is string
if isinstance(content, bytes):
content = content.decode("utf-8")
return content
def transform_gcs_bucket_response_to_openai_file_object(
self, create_file_data: CreateFileRequest, gcs_upload_response: Dict[str, Any]
) -> OpenAIFileObject:
"""
Transforms GCS Bucket upload file response to OpenAI FileObject
"""
gcs_id = gcs_upload_response.get("id", "")
# Remove the last numeric ID from the path
gcs_id = "/".join(gcs_id.split("/")[:-1]) if gcs_id else ""
return OpenAIFileObject(
purpose=create_file_data.get("purpose", "batch"),
id=f"gs://{gcs_id}",
filename=gcs_upload_response.get("name", ""),
created_at=_convert_vertex_datetime_to_openai_datetime(
vertex_datetime=gcs_upload_response.get("timeCreated", "")
),
status="uploaded",
bytes=gcs_upload_response.get("size", 0),
object="file",
)

View file

@ -56,6 +56,16 @@ from litellm.proxy._experimental.mcp_server.sampling_handler import (
MCP_SAMPLING_AVAILABLE,
)
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import resolve_mcp_auth
from litellm.proxy._experimental.mcp_server.outbound_credentials import (
Error,
Ok,
UpstreamCredentialProvider,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
raise_public,
to_server_spec,
to_subject,
)
from litellm.proxy._experimental.mcp_server.utils import (
MCP_TOOL_PREFIX_SEPARATOR,
MCPMissingUserEnvVarsError,
@ -511,7 +521,8 @@ class MCPServerManager:
return "client_credentials"
return None
def __init__(self):
def __init__(self, cred_provider: Optional[UpstreamCredentialProvider] = None):
self._cred_provider = cred_provider or UpstreamCredentialProvider()
self.registry: Dict[str, MCPServer] = {}
self.config_mcp_servers: Dict[str, MCPServer] = {}
"""
@ -1942,11 +1953,19 @@ class MCPServerManager:
Returns:
Configured MCP client instance.
"""
auth_value = await resolve_mcp_auth(
server, mcp_auth_header, subject_token=subject_token
)
transport = server.transport or MCPTransport.sse
spec = None if transport == MCPTransport.stdio else to_server_spec(server)
# A per-request override is the caller-supplied credential v1 turns into the upstream
# auth, so it must win; defer those to v1 (this defer falls away once the per-user modes
# stop writing mcp_auth_header). An inbound header already in extra_headers is handled on
# the v2 path below, not here.
if spec is not None and mcp_auth_header:
spec = None
auth_value = (
await resolve_mcp_auth(server, mcp_auth_header, subject_token=subject_token)
if spec is None
else None
)
# Create sampling and elicitation callbacks for this client
sampling_cb = (
@ -2017,6 +2036,43 @@ class MCPServerManager:
# For HTTP/SSE transports
server_url = server.url or ""
if spec is not None:
match await self._cred_provider.resolve_credentials(
to_subject(user_api_key_auth, subject_token), spec
):
case Ok(auth):
resolved_auth = auth
# Do not override an Authorization already supplied via extra_headers
# (a guardrail hook such as the JWT signer, static_headers, or a
# forwarded caller header): v1 applies those last, so they win. NoOpAuth
# has no header_name and so never skips.
header_name = getattr(resolved_auth, "header_name", None)
if (
header_name
and extra_headers
and any(
key.lower() == header_name.lower()
for key in extra_headers
)
):
resolved_auth = None
case Error(err):
raise_public(err)
return MCPClient(
server_url=server_url,
transport_type=transport,
auth_type=server.auth_type,
timeout=(
server.timeout
if server.timeout is not None
else MCP_CLIENT_TIMEOUT
),
extra_headers=extra_headers,
resolved_auth=resolved_auth,
sampling_callback=sampling_cb,
elicitation_callback=elicitation_cb,
)
# Create SigV4 auth if configured
aws_auth = None
if server.auth_type == MCPAuth.aws_sigv4:

View file

@ -0,0 +1,140 @@
"""The v1 <-> v2 bridge for the credential resolver.
These edge functions translate v1's request objects into the resolver's typed inputs and map
its typed errors onto the proxy's public exception contract. They import v1 and live outside the
package's public surface so the resolver core (``resolver.py`` / ``types.py``) stays v1-free.
Nothing wires them into ``_create_mcp_client`` yet.
``to_server_spec`` maps only the modes the resolver has gone live for, returning ``None`` for
every other mode so the caller defers to v1 (parity-safe); it grows one branch per migrated mode.
"""
from __future__ import annotations
import base64
from typing import TYPE_CHECKING, NoReturn, Optional
from fastapi import HTTPException
from pydantic import SecretStr
from typing_extensions import assert_never
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
ApiKeyConfig,
CredError,
NoneConfig,
ServerSpec,
SharedKey,
Subject,
)
from litellm.types.mcp import MCPAuth
if TYPE_CHECKING:
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
def to_subject(
user_api_key_auth: Optional[UserAPIKeyAuth], subject_token: Optional[str]
) -> Subject:
"""Map v1's authenticated principal onto the resolver's Subject.
tenant_id / subject_id are empty for an unauthenticated caller; the per-user arms must reject
an empty subject rather than share one credential slot across callers.
"""
inbound = SecretStr(subject_token) if subject_token else None
if user_api_key_auth is None:
return Subject(tenant_id="", subject_id="", inbound_token=inbound)
return Subject(
tenant_id=user_api_key_auth.org_id or user_api_key_auth.team_id or "",
subject_id=user_api_key_auth.user_id or "",
inbound_token=inbound,
)
def to_server_spec(server: MCPServer) -> Optional[ServerSpec]:
"""Map a v1 server onto a ServerSpec for a migrated mode, or None to defer to v1.
BYOK is the per-user source of the ``api_key`` mode; its scheme rides on ``auth_type`` just
like a shared key, but the value is per-user and not migrated yet, so a BYOK server defers
to v1 regardless of ``auth_type`` (this guard is the seam the BYOK arm replaces later).
Dispatches on the declared ``auth_type``. The match is exhaustive over ``MCPAuthType`` with
an ``assert_never`` tail, so a newly added auth mode fails the type gate here until it is
explicitly mapped or explicitly deferred, rather than silently falling through to v1. Live
modes: ``none`` and the static-header family (``api_key`` plus the Authorization schemes),
all shared-key; every other mode returns None and stays on v1.
"""
if server.is_byok:
return (
None # per-user BYOK source not migrated yet -> defer to v1 (any auth_type)
)
resource = server.url or server.server_id
auth_type = server.auth_type
match auth_type:
case None | MCPAuth.none:
if server.is_oauth_passthrough:
return None # passthrough is not migrated yet -> defer to v1
return ServerSpec(
server_id=server.server_id, resource=resource, config=NoneConfig()
)
case MCPAuth.api_key:
return _shared_key_spec(server, resource, "X-API-Key", "")
case MCPAuth.bearer_token:
return _shared_key_spec(server, resource, "Authorization", "Bearer")
case MCPAuth.token:
return _shared_key_spec(server, resource, "Authorization", "token")
case MCPAuth.authorization:
return _shared_key_spec(server, resource, "Authorization", "")
case MCPAuth.basic:
return _shared_key_spec(
server, resource, "Authorization", "Basic", encode=True
)
case MCPAuth.oauth2 | MCPAuth.oauth2_token_exchange | MCPAuth.aws_sigv4:
return None # OAuth grants and SigV4 are not migrated yet -> defer to v1
assert_never(auth_type)
def _shared_key_spec(
server: MCPServer,
resource: str,
header_name: str,
value_prefix: str,
*,
encode: bool = False,
) -> Optional[ServerSpec]:
"""Build an api_key spec from the server's static token, or defer (None) if it is absent.
Covers the whole shared-key static-header family: ``api_key`` on ``X-API-Key`` and the
Authorization schemes (bearer / token / authorization sent verbatim, basic base64-encoded).
"""
token = server.authentication_token
if not token:
return None # no key configured -> defer to v1 (parity-safe)
value = base64.b64encode(token.encode("utf-8")).decode() if encode else token
return ServerSpec(
server_id=server.server_id,
resource=resource,
config=ApiKeyConfig(
header_name=header_name,
value_prefix=value_prefix,
key_source=SharedKey(value=SecretStr(value)),
),
)
def raise_public(error: CredError) -> NoReturn:
"""Map a resolver CredError onto the proxy's public HTTP contract. The one edge that raises."""
match error.tag:
case "unauthorized":
raise HTTPException(status_code=401, detail=error.summary)
case "misconfigured":
raise HTTPException(status_code=500, detail=error.summary)
case "upstream_unavailable":
raise HTTPException(status_code=503, detail=error.summary)
case "unsupported_mode":
raise HTTPException(status_code=500, detail=error.summary)
case "precondition_required":
raise HTTPException(status_code=412, detail=error.summary)
case "not_implemented":
raise HTTPException(status_code=501, detail=error.summary)
assert_never(error.tag)

View file

@ -7,9 +7,9 @@ no precedence cascade. It is wildcard-free with an `assert_never` tail, so addin
an arm fails the type gate (basedpyright `reportMatchNotExhaustive`); a bypassed gate fails loudly
at runtime instead of returning `None`.
This skeleton ships every arm as a `not_implemented` stub. Each mode's real body, with its
injected seam, lands in its own follow-up PR; until then the arm returns a typed error rather
than silently producing no credential. Pure v2: no imports from v1.
`none` and `api_key` (shared-key source) are live; the remaining arms are `not_implemented`
stubs that each land in a follow-up PR with their injected seam. The self-contained arms read
straight from the config and need no collaborator. Pure v2: no imports from v1.
"""
from __future__ import annotations
@ -17,8 +17,13 @@ from __future__ import annotations
import httpx
from typing_extensions import assert_never
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
NoOpAuth,
StaticHeaderAuth,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
Error,
Ok,
Result,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
@ -26,11 +31,13 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
AuthorizationCodeConfig,
AuthSpecKind,
AwsSigV4Config,
Byok,
ClientCredentialsConfig,
CredError,
NoneConfig,
PassthroughConfig,
ServerSpec,
SharedKey,
Subject,
TokenExchangeConfig,
)
@ -40,7 +47,7 @@ class UpstreamCredentialProvider:
"""Produces the one `httpx.Auth` for a `(subject, upstream)` pair, per declared mode.
Collaborators (the per-mode credential stores and token fetchers) are injected as each arm
is built; the skeleton needs none, since every arm is a stub.
is built; the live `none` and `api_key`-shared arms read from the config and need none.
"""
async def resolve_credentials(
@ -48,9 +55,9 @@ class UpstreamCredentialProvider:
) -> Result[httpx.Auth, CredError]:
match server.config:
case NoneConfig():
return _not_implemented(AuthSpecKind.none)
case ApiKeyConfig():
return _not_implemented(AuthSpecKind.api_key)
return Ok(NoOpAuth())
case ApiKeyConfig() as config:
return self._api_key(config)
case PassthroughConfig():
return _not_implemented(AuthSpecKind.passthrough)
case ClientCredentialsConfig():
@ -63,6 +70,22 @@ class UpstreamCredentialProvider:
return _not_implemented(AuthSpecKind.aws_sigv4)
assert_never(server.config)
def _api_key(self, config: ApiKeyConfig) -> Result[httpx.Auth, CredError]:
match config.key_source:
case SharedKey() as source:
header_name, header_value = config.header(
source.value.get_secret_value()
)
return Ok(StaticHeaderAuth(header_value, header_name=header_name))
case Byok():
# Per-user key pulled from the credential store; lands with that seam.
return Error(
CredError.of_not_implemented(
"api_key BYOK source not implemented yet"
)
)
assert_never(config.key_source)
def _not_implemented(kind: AuthSpecKind) -> Result[httpx.Auth, CredError]:
return Error(

View file

@ -21,6 +21,7 @@ from typing import (
TYPE_CHECKING,
Any,
Dict,
Iterable,
List,
Literal,
NoReturn,
@ -32,13 +33,15 @@ from typing import (
from fastapi import HTTPException
from pydantic import BaseModel
import json
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.batches.batch_utils import (
_count_entry_tokens,
_estimate_batch_entry_tokens,
_extract_file_access_credentials,
_get_batch_job_input_file_usage,
_get_file_content_as_dictionary,
_get_models_from_batch_input_file_content,
_iter_batch_input_lines,
)
from litellm.exceptions import RateLimitErrorCategory
from litellm.integrations.custom_logger import CustomLogger
@ -537,6 +540,8 @@ class _PROXY_BatchRateLimiter(CustomLogger):
# Managed files require bypassing the HTTP endpoint (which runs access-check hooks)
# and calling the managed files hook directly with the user's credentials.
is_managed_file = _is_base64_encoded_unified_file_id(file_id)
# For managed files the unified file id encodes the proxy model
# alias(es) the file was uploaded for; auth validates against those.
target_model_names = (
get_models_from_unified_file_id(is_managed_file)
if is_managed_file
@ -568,7 +573,38 @@ class _PROXY_BatchRateLimiter(CustomLogger):
f"Expected bytes content from file retrieval for {file_id}, "
f"got {type(file_content_bytes)}"
)
file_content_as_dict = _get_file_content_as_dictionary(file_content_bytes)
# Single streaming pass over the JSONL lines, accounting each row
# independently. One bad row can never abort the pass: a malformed
# line is skipped (its request can't run upstream anyway) and a row
# the token counter can't measure falls back to a conservative
# size-based estimate. This guarantees two things a restricted caller
# must not be able to break by crafting a row that raises:
# 1. The allowlist check below always sees every parseable
# ``body.model`` (the loop never stops early), so models can't be
# smuggled in after a bad row.
# 2. The token total is never silently zeroed, so the TPM limit
# can't be evaded by sending uncountable rows.
# Counting stays best-effort, so a legitimate (e.g. multimodal) row
# the counter can't measure is estimated, not hard-rejected.
models: set = set()
total_tokens = 0
request_count = 0
for raw_line in _iter_batch_input_lines(file_content_bytes):
request_count += 1
try:
entry = json.loads(raw_line)
except Exception:
total_tokens += _estimate_batch_entry_tokens(raw_line)
continue
if isinstance(entry, dict):
model = (entry.get("body") or {}).get("model")
if model:
models.add(model)
try:
total_tokens += _count_entry_tokens(entry)
except Exception:
total_tokens += _estimate_batch_entry_tokens(raw_line)
# Validate every model named in the batch JSONL against the
# caller's per-key model allowlist. Without this, a caller
@ -578,17 +614,12 @@ class _PROXY_BatchRateLimiter(CustomLogger):
if user_api_key_dict is not None:
await self._enforce_batch_file_model_access(
user_api_key_dict=user_api_key_dict,
file_content_as_dict=file_content_as_dict,
models=models,
target_model_names=target_model_names or None,
)
input_file_usage = _get_batch_job_input_file_usage(
file_content_dictionary=file_content_as_dict,
custom_llm_provider=custom_llm_provider,
)
request_count = len(file_content_as_dict)
return BatchFileUsage(
total_tokens=input_file_usage.total_tokens,
total_tokens=total_tokens,
request_count=request_count,
)
@ -614,14 +645,15 @@ class _PROXY_BatchRateLimiter(CustomLogger):
async def _enforce_batch_file_model_access(
self,
user_api_key_dict: UserAPIKeyAuth,
file_content_as_dict: List[dict],
models: Optional[Iterable[str]] = None,
target_model_names: Optional[List[str]] = None,
) -> None:
"""Reject the batch if the caller is not authorized for the upload target.
For managed files, ``target_model_names`` (from the unified file id) is
the proxy alias the file was uploaded for and is used directly for auth.
For legacy/non-managed files, falls back to ``body.model`` values in the JSONL.
the proxy alias the file was uploaded for and is checked directly.
Otherwise the ``body.model`` values collected from the JSONL (``models``)
are checked.
Reuses standard auth helpers so the same model access rules the proxy
enforces on `/chat/completions` apply here.
@ -640,10 +672,9 @@ class _PROXY_BatchRateLimiter(CustomLogger):
if target_model_names:
models = target_model_names
else:
models = _get_models_from_batch_input_file_content(file_content_as_dict)
if not models:
return
if not models:
return
team_object = None
if (

View file

@ -0,0 +1,208 @@
"""
Ingest pre-built logging payloads from external producers and replay them
through LiteLLM's standard success/failure callback fan-out.
This exists for hosts that own a request outside the Python process — e.g. the
`litellm-rust` gateway proxying realtime websockets. Those hosts can't use the
in-process logging object, so they POST a finished `StandardLoggingPayload` here
and Python replays it through the exact same path a normal completion uses:
`Logging.async_success_handler` / `async_failure_handler`. Every registered
callback (spend logs, Langfuse, Datadog, ...) fires unchanged — there is no
spend-logs-specific or callback-specific code here, only the replay.
The endpoint is generic: realtime is the first producer, but the contract is the
self-describing `StandardLoggingPayload`, so completions/responses can use it too.
"""
import uuid
from datetime import datetime, timezone
from typing import Any
from fastapi import APIRouter, Depends, HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.types.proxy.callback_logs_endpoints import (
CallbackLogFailure,
CallbackLogRecord,
CallbackLogsRequest,
CallbackLogsResponse,
)
# Routes the Python proxy exposes for the Rust data-plane gateway to call into
# (logging today; auth/budgets later). Namespaced under /v1/rust_control_plane so
# they're clearly distinct from the proxy's own control-plane/management routes.
rust_control_plane_router = APIRouter(
prefix="/v1/rust_control_plane", tags=["rust control plane"]
)
class CallbackLogsReplayer:
"""
Replays finished logging payloads through LiteLLM's callback fan-out.
Each helper is small and pure so the replay path is easy to read and test:
rebuild a `Logging` object from the payload, seed `model_call_details` with
exactly what the callbacks read, then dispatch to the success/failure
handler. No spend/callback logic lives here — only the replay.
"""
@staticmethod
def _epoch_to_datetime(value: Any) -> datetime:
"""`StandardLoggingPayload` stores startTime/endTime as float epoch seconds."""
if isinstance(value, (int, float)):
return datetime.fromtimestamp(float(value), tz=timezone.utc)
if isinstance(value, datetime):
return value
return datetime.now(tz=timezone.utc)
@staticmethod
def _build_logging_obj(payload: dict[str, Any]) -> LiteLLMLogging:
"""
Reconstruct a `Logging` object from a finished payload and seed
`model_call_details` with exactly what the success/failure callbacks
read: the prebuilt `standard_logging_object`, the resolved
`response_cost`, and the `litellm_params.metadata` keys used for cost
attribution. Setting `standard_logging_object` up front makes the handler
skip rebuilding it.
"""
model = payload.get("model") or ""
call_type = payload.get("call_type") or "acompletion"
start_time = CallbackLogsReplayer._epoch_to_datetime(payload.get("startTime"))
call_id = (
payload.get("litellm_call_id") or payload.get("id") or str(uuid.uuid4())
)
logging_obj = LiteLLMLogging(
model=model,
messages=payload.get("messages") or [],
# A replayed payload is always a *terminal*, fully-aggregated event —
# the producer (e.g. the rust gateway) already collected the whole
# session before POSTing. Never mark it streaming: a streaming
# Logging object makes async_success_handler wait for a
# complete_streaming_response that will never arrive, so the spend
# log is never written.
stream=False,
call_type=call_type,
start_time=start_time,
litellm_call_id=call_id,
function_id="",
)
metadata: dict[str, Any] = payload.get("metadata") or {}
litellm_metadata: dict[str, Any] = {
"user_api_key": metadata.get("user_api_key_hash"),
"user_api_key_alias": metadata.get("user_api_key_alias"),
"user_api_key_user_id": metadata.get("user_api_key_user_id"),
"user_api_key_team_id": metadata.get("user_api_key_team_id"),
"user_api_key_org_id": metadata.get("user_api_key_org_id"),
"user_api_key_end_user_id": metadata.get("user_api_key_end_user_id"),
"spend_logs_metadata": metadata.get("spend_logs_metadata"),
}
logging_obj.model_call_details.update(
{
"model": model,
"call_type": call_type,
"custom_llm_provider": payload.get("custom_llm_provider"),
"response_cost": payload.get("response_cost") or 0.0,
"standard_logging_object": payload,
"litellm_params": {"metadata": litellm_metadata},
"cache_hit": payload.get("cache_hit") or False,
}
)
return logging_obj
@staticmethod
def _response_obj_from_payload(payload: dict[str, Any]) -> dict[str, Any]:
"""Minimal response object so usage-derived spend-log fields resolve."""
return {
"id": payload.get("id"),
"usage": {
"prompt_tokens": payload.get("prompt_tokens", 0),
"completion_tokens": payload.get("completion_tokens", 0),
"total_tokens": payload.get("total_tokens", 0),
},
}
async def replay(self, record: CallbackLogRecord) -> None:
"""Replay one record through the matching success/failure handler."""
payload = record.standard_logging_payload
verbose_proxy_logger.debug(
"CallbackLogsReplayer: replaying %s record id=%s model=%s call_type=%s",
record.status,
payload.get("id"),
payload.get("model"),
payload.get("call_type"),
)
logging_obj = self._build_logging_obj(payload)
start_time = self._epoch_to_datetime(payload.get("startTime"))
end_time = self._epoch_to_datetime(payload.get("endTime"))
if record.status == "success":
await logging_obj.async_success_handler(
result=self._response_obj_from_payload(payload),
start_time=start_time,
end_time=end_time,
)
else:
error_str = record.error or payload.get("error_str") or "replayed failure"
await logging_obj.async_failure_handler(
Exception(error_str),
traceback_exception="",
start_time=start_time,
end_time=end_time,
)
async def replay_batch(
self, records: list[CallbackLogRecord]
) -> CallbackLogsResponse:
"""Replay a batch; a single bad record never sinks the rest. Each failure
is reported back with its batch index so the caller can retry/triage it."""
processed = 0
failures: list[CallbackLogFailure] = []
for index, record in enumerate(records):
try:
await self.replay(record)
processed += 1
except Exception as e:
failures.append(CallbackLogFailure(index=index, error=str(e)))
verbose_proxy_logger.exception(
"CallbackLogsReplayer: failed to replay record %s: %s",
index,
str(e),
)
verbose_proxy_logger.debug(
"CallbackLogsReplayer: batch done processed=%s failed=%s",
processed,
len(failures),
)
return CallbackLogsResponse(
processed=processed, failed=len(failures), failures=failures
)
@rust_control_plane_router.post(
"/logs",
dependencies=[Depends(user_api_key_auth)],
)
async def ingest_callback_logs(
body: CallbackLogsRequest,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> CallbackLogsResponse:
"""
Replay a batch of finished logging payloads through the callback fan-out.
Admin-only: the payloads write spend logs and trigger every callback, so this
is a trusted internal route, not a public surface.
"""
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
raise HTTPException(
status_code=403,
detail="/v1/rust_control_plane/logs is admin-only (proxy admin key required).",
)
return await CallbackLogsReplayer().replay_batch(body.records)

View file

@ -874,6 +874,8 @@ async def _common_key_generation_helper(
object_permission=data_json.get("object_permission"),
team_obj=team_table,
prisma_client=prisma_client,
is_proxy_admin=user_api_key_dict.user_role
== LitellmUserRoles.PROXY_ADMIN.value,
)
if normalized_object_permission is not None:
data_json["object_permission"] = normalized_object_permission
@ -2164,6 +2166,7 @@ async def _validate_mcp_servers_for_key_update(
existing_key_row: Any,
prisma_client: Any,
user_api_key_cache: Any,
is_proxy_admin: bool,
) -> Optional[dict]:
"""Validate MCP servers in object_permission against the effective team."""
effective_team_obj = team_obj
@ -2186,6 +2189,7 @@ async def _validate_mcp_servers_for_key_update(
object_permission=object_permission_dict,
team_obj=effective_team_obj,
prisma_client=prisma_client,
is_proxy_admin=is_proxy_admin,
)
await validate_key_search_tools_against_team(
object_permission=object_permission_dict,
@ -2422,6 +2426,7 @@ async def _validate_update_key_data(
existing_key_row=existing_key_row,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
is_proxy_admin=_is_proxy_admin,
)
if normalized_object_permission is not None:
data.object_permission = LiteLLM_ObjectPermissionBase(

View file

@ -469,6 +469,7 @@ async def validate_key_mcp_servers_against_team(
object_permission: Optional[dict],
team_obj: Optional["LiteLLM_TeamTableCachedObj"],
prisma_client: Optional[PrismaClient] = None,
is_proxy_admin: bool = False,
) -> Optional[dict]:
"""
Validate that MCP servers requested on a key are within the allowed scope.
@ -476,12 +477,17 @@ async def validate_key_mcp_servers_against_team(
Rules:
- If key is in a team: key's mcp_servers must be a subset of
(team's allowed servers + allow_all_keys servers)
- If key is NOT in a team: key's mcp_servers must only contain
allow_all_keys servers
- If key is NOT in a team and the caller is a proxy admin: any server or
access group may be assigned. A proxy admin can already reach every MCP
server, and runtime access is granted directly from the key's own
object_permission, so the key is scoped to exactly what the admin selected
- If key is NOT in a team and the caller is not a proxy admin: key's
mcp_servers must only contain allow_all_keys servers
- If team has no MCP config: key can only use allow_all_keys servers
Raises HTTPException(403) if validation fails.
"""
teamless_admin_assignment = team_obj is None and is_proxy_admin
requested_servers = _extract_requested_mcp_server_ids(object_permission)
requested_access_groups = _extract_requested_mcp_access_groups(object_permission)
@ -526,7 +532,11 @@ async def validate_key_mcp_servers_against_team(
identifier_to_server_ids
)
disallowed_servers = active_requested_servers - all_allowed_servers
allowed_servers = all_allowed_servers
if teamless_admin_assignment:
allowed_servers = all_allowed_servers | active_requested_servers
disallowed_servers = active_requested_servers - allowed_servers
if disallowed_servers:
if team_obj is not None:
team_id = team_obj.team_id
@ -557,7 +567,11 @@ async def validate_key_mcp_servers_against_team(
):
team_access_groups = set(team_obj.object_permission.mcp_access_groups)
disallowed_groups = requested_access_groups - team_access_groups
allowed_access_groups = team_access_groups
if teamless_admin_assignment:
allowed_access_groups = team_access_groups | requested_access_groups
disallowed_groups = requested_access_groups - allowed_access_groups
if disallowed_groups:
if team_obj is not None:
team_id = team_obj.team_id

View file

@ -7,7 +7,7 @@
import asyncio
import traceback
from typing import Any, Optional, cast, get_args
from typing import Any, BinaryIO, Optional, Union, cast, get_args
import httpx
from fastapi import (
@ -97,16 +97,18 @@ def get_files_provider_config(
return None
def get_first_json_object(file_content_bytes: bytes) -> Optional[dict]:
def get_first_json_object(file_source: Union[bytes, BinaryIO]) -> Optional[dict]:
try:
# Decode the bytes to a string and split into lines
file_content = file_content_bytes.decode("utf-8")
first_line = file_content.splitlines()[0].strip()
# Parse the JSON object from the first line
json_object = json.loads(first_line)
return json_object
except (json.JSONDecodeError, UnicodeDecodeError):
if isinstance(file_source, (bytes, bytearray)):
newline = file_source.find(b"\n")
raw = file_source if newline == -1 else file_source[:newline]
first_line = raw.decode("utf-8")
else:
file_source.seek(0)
first_line = file_source.readline().decode("utf-8")
file_source.seek(0)
return json.loads(first_line.strip())
except (json.JSONDecodeError, UnicodeDecodeError, OSError, ValueError):
return None
@ -327,9 +329,15 @@ async def create_file(
data: Dict = {}
try:
# Use orjson to parse JSON data, orjson speeds up requests significantly
# Read the file content
file_content = await file.read()
# Batch uploads can be gigabytes. Starlette has already spooled the upload
# to disk, so stream from that handle instead of reading it into memory.
# Other uploads are small and stay in-memory bytes.
file_source: Union[bytes, BinaryIO]
if purpose == "batch":
await file.seek(0)
file_source = file.file
else:
file_source = await file.read()
custom_llm_provider = (
provider
or get_custom_llm_provider_from_request_headers(request=request)
@ -454,13 +462,13 @@ async def create_file(
)
# Prepare the file data according to FileTypes
file_data = (file.filename, file_content, file.content_type)
file_data = (file.filename, file_source, file.content_type)
## check if model is a loadbalanced model
router_model: Optional[str] = None
is_router_model = False
if litellm.enable_loadbalancing_on_batch_endpoints is True:
json_obj = get_first_json_object(file_content_bytes=file_content)
json_obj = get_first_json_object(file_source)
if json_obj:
router_model = get_model_from_json_obj(json_object=json_obj)
is_router_model = is_known_model(

View file

@ -346,6 +346,9 @@ from litellm.proxy.hooks.prompt_injection_detection import (
from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger
from litellm.proxy.image_endpoints.endpoints import router as image_router
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
from litellm.proxy.logging_endpoints.callback_logs_endpoints import (
rust_control_plane_router,
)
from litellm.proxy.rust_control_plane.auth_endpoints import (
router as rust_control_plane_auth_router,
)
@ -16642,6 +16645,7 @@ app.include_router(analytics_router)
app.include_router(callback_management_endpoints_router)
app.include_router(debugging_endpoints_router)
app.include_router(rust_control_plane_auth_router)
app.include_router(rust_control_plane_router)
app.include_router(ui_crud_endpoints_router)
app.include_router(openai_files_router)
app.include_router(team_callback_router)

View file

@ -82,43 +82,83 @@ def replace_model_in_jsonl(file_content: FileTypes, new_model_name: str) -> File
if isinstance(file_content, PathLike):
return file_content
# Decode the bytes to a string and split into lines
# If file_content is a file-like object, read the bytes
if hasattr(file_content, "read"):
file_content_bytes = file_content.read() # type: ignore
elif isinstance(file_content, tuple):
file_content_bytes = file_content[1]
else:
file_content_bytes = file_content
# Decode the bytes to a string and split into lines
if isinstance(file_content_bytes, bytes):
file_content_str = file_content_bytes.decode("utf-8")
elif isinstance(file_content_bytes, str):
file_content_str = file_content_bytes
# Iterate the source line-by-line WITHOUT reading it all into memory. A
# spooled upload handle (managed batches stream from it) is read straight
# off its backing; bytes/str are wrapped so they iterate line-by-line.
source = file_content[1] if isinstance(file_content, tuple) else file_content
if hasattr(source, "read"):
if hasattr(source, "seek"):
try:
source.seek(0) # type: ignore[attr-defined]
except (OSError, ValueError):
pass
line_iter: object = source
elif isinstance(source, (bytes, bytearray)):
line_iter = io.BytesIO(bytes(source))
elif isinstance(source, str):
line_iter = io.StringIO(source)
else:
return file_content
# Parse JSONL properly, handling potential multiline JSON objects
json_objects = parse_jsonl_with_embedded_newlines(file_content_str)
# Rewrite one row at a time, writing straight into the output buffer
# instead of holding every parsed row in a list. Peak memory stays at
# ~one row plus the output rather than several full copies of the file,
# which the managed-files path depends on (it re-runs this rewrite once
# per target model). Lines are accumulated so JSON objects that span
# multiple physical lines still parse. Streaming the handle also means
# the model rewrite is actually applied to tuple-wrapped upload handles;
# otherwise a restricted body.model would survive and bypass the batch
# model allowlist (which validates the upload target alias).
output = InMemoryFile(
b"", name="modified_file.jsonl", content_type="application/jsonl"
)
wrote_any = False
buffer = ""
for raw_line in line_iter: # type: ignore[attr-defined]
buffer += (
raw_line.decode("utf-8")
if isinstance(raw_line, (bytes, bytearray))
else raw_line
)
stripped = buffer.strip()
if not stripped:
buffer = ""
continue
try:
json_object = json.loads(stripped)
except json.JSONDecodeError:
continue # object not complete yet; keep accumulating
if isinstance(json_object, dict) and isinstance(
json_object.get("body"), dict
):
json_object["body"]["model"] = new_model_name
output.write(
(("\n" if wrote_any else "") + json.dumps(json_object)).encode("utf-8")
)
wrote_any = True
buffer = ""
if buffer.strip():
# A row never parsed (truncated/malformed, or it swallowed the rows
# that followed it). Returning the partial `output` would silently
# drop those rows; return the unchanged original so the provider
# rejects the batch loudly instead of accepting a truncated one.
verbose_logger.error(
f"error parsing trailing batch content: {buffer[:100]}..."
)
if hasattr(source, "seek"):
try:
source.seek(0) # type: ignore[attr-defined]
except (OSError, ValueError):
pass
return file_content
# If no valid JSON objects were found, return the original content
if len(json_objects) == 0:
if not wrote_any:
return file_content
modified_lines = []
for json_object in json_objects:
# Replace the model name if it exists
if "body" in json_object:
json_object["body"]["model"] = new_model_name
# Convert the modified JSON object back to a string
modified_lines.append(json.dumps(json_object))
# Reassemble the modified lines and return as bytes
modified_file_content = "\n".join(modified_lines).encode("utf-8")
return InMemoryFile(modified_file_content, name="modified_file.jsonl", content_type="application/jsonl") # type: ignore
output.seek(0)
return output # type: ignore
except (json.JSONDecodeError, UnicodeDecodeError, TypeError):
# return the original file content if there is an error replacing the model name

View file

@ -321,3 +321,21 @@ class TwoStepFileUploadConfig(TypedDict, total=False):
upload_request: Required[TwoStepFileUploadRequest]
upload_url_location: Required[Literal["headers", "body"]]
upload_url_key: str
class ResumableChunkedUploadConfig(TypedDict, total=False):
"""Drives a memory-bounded resumable upload (GCS JSON API).
The handler POSTs to the upload URL to open a session, reads the session URI
from ``session_url_header``, then PUTs ``body_stream`` to that URI in
``chunk_size``-byte chunks (a 256 KiB multiple) using Content-Range, so the
payload is never buffered in full and the transfer is resumable.
``body_stream`` is a ``BaseFileUploadStream``; it is typed ``Any`` here to
avoid importing the llms layer into types.
"""
body_stream: Required[Any]
chunk_size: int
session_url_header: str
initiate_headers: Dict[str, str]

View file

@ -0,0 +1,44 @@
"""
Types for the callback-logs ingest endpoint (POST /v1/callbacks/logs).
External producers (e.g. the litellm-rust gateway) POST finished logging
payloads here; the proxy replays them through the standard callback fan-out.
"""
from typing import Any, Literal, Optional
from pydantic import BaseModel, Field
from litellm.constants import MAX_CALLBACK_LOG_RECORDS
class CallbackLogRecord(BaseModel):
"""A single finished logging event to replay through the callbacks."""
status: Literal["success", "failure"]
standard_logging_payload: dict[str, Any]
error: Optional[str] = None
class CallbackLogsRequest(BaseModel):
"""A batch of logging events posted by an external producer."""
# Bounded so one POST can't trigger an unbounded callback/DB fan-out (each
# record fires every registered integration). Over the cap → 422.
records: list[CallbackLogRecord] = Field(..., max_length=MAX_CALLBACK_LOG_RECORDS)
class CallbackLogFailure(BaseModel):
"""A record that failed to replay, identified by its index in the batch."""
index: int
error: str
class CallbackLogsResponse(BaseModel):
"""Per-batch result: counts plus per-record failure detail so the caller can
distinguish a transient callback error from a structurally bad payload."""
processed: int
failed: int
failures: list[CallbackLogFailure] = Field(default_factory=list)

View file

@ -180,6 +180,9 @@ class CredentialLiteLLMParams(BaseModel):
## UNIFIED PROJECT/REGION ##
region_name: Optional[str] = None
## OBJECT STORAGE (files / batches) ##
gcs_bucket_name: Optional[str] = None
## AWS BEDROCK / SAGEMAKER ##
aws_access_key_id: Optional[str] = None
aws_secret_access_key: Optional[str] = None

View file

@ -27,7 +27,7 @@ from litellm.integrations.custom_logger import CustomLogger
from litellm.types.utils import StandardLoggingPayload
import socket
import httpx
from unittest.mock import patch, MagicMock
from unittest.mock import patch, MagicMock, AsyncMock
def _can_resolve_openai():
@ -513,10 +513,26 @@ async def test_avertex_batch_prediction(monkeypatch):
mock_response.status_code = 200
return mock_response
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
side_effect=mock_side_effect,
) as mock_global_post:
# Batch jsonl file creation now streams to a GCS resumable session via
# _aresumable_chunked_upload (httpx send), not AsyncHTTPHandler.post, so mock
# that entry point to return the GCS object response. The resumable protocol
# itself is covered in test_vertex_ai_files_streaming.py.
mock_upload_response = httpx.Response(
200,
json=mock_file_response,
request=httpx.Request("PUT", "https://storage.googleapis.com/upload"),
)
with (
patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
side_effect=mock_side_effect,
) as mock_global_post,
patch(
"litellm.llms.custom_httpx.llm_http_handler.BaseLLMHTTPHandler._aresumable_chunked_upload",
new_callable=AsyncMock,
return_value=mock_upload_response,
),
):
litellm.set_verbose = True
litellm._turn_on_debug()
file_name = "vertex_batch_completions.jsonl"

View file

@ -0,0 +1,93 @@
# Budget Code Matrix
What LiteLLM actually implements for budgets: every entity that can carry a dollar
budget, how the limit is enforced, and where in the code it happens. This is the
"what we support" reference; the companion `BUDGET_TEST_COVERAGE_MATRIX.md` maps
each row to its tests and the e2e gaps.
Over-budget surfaces as a `budget_exceeded` error (the live suite
`tests/otel_tests/test_e2e_budgeting.py` asserts `type == "budget_exceeded"`,
`code == "429"`); the underlying `BudgetExceededError` is defined in
`litellm/exceptions.py` (`status_code=400`). Enforcement runs in `common_checks()`
/ `auth_checks.py` at auth time, plus pre-call reservation in
`budget_reservation.py`.
Legend for "Enforced": **block** = request rejected; **filter** = router skips the
deployment; **alert** = notify only, request proceeds.
---
## 1. Per-entity dollar budgets
| Entity | Budget stored | Hard `max_budget` | Soft budget | Per-window | Model budget | Reset by `budget_duration` |
|--------|---------------|-------------------|-------------|------------|--------------|----------------------------|
| API key | `LiteLLM_VerificationToken` (direct cols + `budget_id` FK) | block (`_virtual_key_max_budget_check`) | alert (`_virtual_key_soft_budget_check`) + 80% alert | block (`_virtual_key_multi_budget_check`) | block (`model_max_budget_limiter.is_key_within_model_budget`) | keys reset job |
| Internal user | `LiteLLM_UserTable` (direct cols) | block (`common_checks`, only when not on a team) | - | - | via `model_max_budget` json | users reset job |
| Team | `LiteLLM_TeamTable` (direct cols) | block (`_team_max_budget_check`) | alert (`_team_soft_budget_check`) | block (`_team_multi_budget_check`) | via `model_max_budget` | teams reset job |
| Team member | `LiteLLM_TeamMembership` -> `LiteLLM_BudgetTable` | block (`_check_team_member_budget`) | - | - | - | budget-table reset job |
| End-user / customer | `LiteLLM_EndUserTable` -> `LiteLLM_BudgetTable` | block (`_check_end_user_budget`) | - | - | block (`is_end_user_within_model_budget`) | budget-table reset job |
| Organization | `LiteLLM_OrganizationTable` -> `LiteLLM_BudgetTable` | block (`_organization_max_budget_check`) | - | - | via budget-table | budget-table reset job |
| Tag | `LiteLLM_TagTable` -> `LiteLLM_BudgetTable` | block (`_tag_max_budget_check`) | - | - | via budget-table | budget-table reset job |
| Project | `LiteLLM_ProjectTable` -> `LiteLLM_BudgetTable` | block (`_project_max_budget_check`) | alert (`_project_soft_budget_check`) | - | - | budget-table reset job |
| Provider (router) | config `provider_budget_config` (in-memory) | filter (`router_strategy/budget_limiter`) | - | yes (time window) | - | window TTL |
| Global proxy | `litellm.max_budget` (config) | block (`_global_proxy_budget_check`) | - | - | - | - |
Notes / flags from the code:
- **User budget only enforced off-team**: `common_checks` skips the personal-user
budget when the key belongs to a team (team budget governs instead).
- **Comparison operators are inconsistent**: key/user use `>=`, team/end-user main
budget use `>`. Spend exactly at `max_budget` blocks a key but not a team.
- **Provider budgets are filter-only**: an over-budget provider is removed from
routing; if all are over budget the router raises
`no_deployments_with_provider_budget_routing` (not a per-entity block).
- **Enforcement timing differs by entity**: key / user / org / team-member / tag /
model enforce off real-time reservation counters (block within ~2 calls);
**end-user** enforcement reads `EndUserTable.spend`, which only updates on the
`proxy_batch_write_at` flush, so it lags by that interval (verified live).
## 2. Budget mechanisms
| Mechanism | What it does | Code |
|-----------|--------------|------|
| Pre-call reservation | Estimates max request cost, atomically reserves against redis spend counters for key/team/user/end_user/tag/team_member/org before the call; blocks if a counter would exceed | `spend_tracking/budget_reservation.py` |
| Post-call reconciliation | Adjusts the reservation to the actual cost once known | `reconcile_budget_reservation` |
| Read-time enforcement | Auth-time check of current spend vs `max_budget` | `auth_checks.common_checks` + per-entity `_*_max_budget_check` |
| Soft budget / alerts | At `soft_budget` (or 80% of max) fire Slack/email alert, do not block | `_virtual_key_soft_budget_check`, `_team_soft_budget_check`, `budget_alerts` |
| Multi-window budgets | `budget_limits` list of `{budget_duration, max_budget}`; each window enforced + reset independently | `_virtual_key_multi_budget_check`, `reset_budget_windows` |
| Model-level budgets | `model_max_budget` dict (per model: `budget_limit` + `time_period`) on key/user/end_user | `hooks/model_max_budget_limiter.py` |
| Reset by duration | Job zeros `spend`, recomputes `budget_reset_at = now + duration_in_seconds(budget_duration)`, invalidates redis counters | `common_utils/reset_budget_job.py`, `duration_parser.duration_in_seconds` |
| Zero-cost bypass | Models with no configured price bypass budget reservation | `budget_reservation` zero-cost path |
## 3. Budget management surface (endpoints)
| Action | Endpoint | Handler |
|--------|----------|---------|
| Create budget | `POST /budget/new` | `new_budget` |
| Update budget | `POST /budget/update` | `update_budget` |
| Budget info | `POST /budget/info` (`{"budgets": [id]}`) | `info_budget` |
| Budget settings | `GET /budget/settings` | `budget_settings` |
| List budgets | `GET /budget/list` | `list_budget` |
| Delete budget | `POST /budget/delete` (`{"id": id}`) | `delete_budget` |
| Set on key | `POST /key/generate`, `/key/update` (`max_budget`, `soft_budget`, `budget_duration`, `model_max_budget`, `budget_id`) | key mgmt |
| Set on user | `POST /user/new` (`max_budget`, `budget_duration`) | internal user |
| Set on team | `POST /team/new` (`max_budget`, `soft_budget`, `team_member_budget`) | team |
| Set on team member | `POST /team/member_add` (`max_budget_in_team`) | team |
| Set on org | `POST /organization/new` (`max_budget`, `soft_budget`, `model_max_budget`) | org |
| Set on customer | `POST /customer/new`, `/customer/update` (`max_budget`, `budget_id`) | customer |
| Set on tag | `POST /tag/new`, `/tag/update` (`max_budget`) | tag mgmt |
| Read budget+spend | `/key/info`, `/user/info`, `/team/info`, `/organization/info`, `/customer/info`, `/budget/info` | per-entity info |
Endpoint method/shape gotchas verified live: `/organization/delete` is **DELETE**
with `{"organization_ids": [id]}`; `/budget/info` takes `{"budgets": [id]}`;
`model_max_budget` entries use `{"budget_limit", "time_period"}`.
## 4. Config knobs
| Setting | Effect |
|---------|--------|
| `litellm.max_budget` | proxy-wide hard cap (global proxy budget) |
| `max_internal_user_budget` / `default_max_internal_user_budget` | default `max_budget` for internal users |
| `internal_user_budget_duration` | default reset duration for internal users |
| `max_end_user_budget` / `max_end_user_budget_id` | default budget for end-users |
| `default_team_params` | default `max_budget` / `budget_duration` / limits for teams |
| `provider_budget_config` (router) | per-provider spend caps + windows |

View file

@ -0,0 +1,78 @@
# Budget Test Coverage Matrix
Maps every row of `BUDGET_CODE_MATRIX.md` (what LiteLLM implements) to its tests
and level, then marks the live e2e coverage this suite adds.
Levels: `unit` mocked (`AsyncMock` on `get_current_spend`/prisma); `router` live
router with fake deployments; `live-e2e` real proxy, real key/team, real requests
until blocked. Status: `covered` / `partial` / `gap`.
Pre-existing live coverage outside this suite:
- `tests/otel_tests/test_e2e_budgeting.py` - key + team enforcement, budget update.
- `tests/local_testing/test_router_budget_limiter.py` - provider / tag / deployment
budgets at the router.
This suite (`tests/e2e/budgets/`) adds the missing live coverage and runs
on the shared lifecycle (every entity it creates is deleted on teardown).
---
## Per-entity enforcement
| Entity | Unit | Pre-existing live | This suite (live) | Status |
|--------|------|-------------------|-------------------|--------|
| API key | `test_budget_reservation.py`, `test_max_budget_limiter.py` | `otel_tests` | `test_budget_enforcement_e2e::test_key_budget_blocks` | **covered** |
| Team | `test_team_budget_limits.py` | `otel_tests` | (org test builds a team) | **covered** |
| Internal user | auth unit tests | - | `test_internal_user_budget_blocks` | **covered (new)** |
| Team member | `test_team_member_budget.py` | - | `test_team_member_budget_blocks` | **covered (new)** |
| End-user / customer | `test_custom_auth_end_user_budget.py` | - | `test_end_user_budget_blocks` | **covered (new)** |
| Organization | `test_organization_budget_enforcement.py` (flagged weak) | - | `test_organization_budget_blocks` | **covered (new)** |
| Tag (proxy-level) | - | router only | `test_tag_budget_e2e::test_tag_budget_blocks_tagged_requests` | **covered (new)** |
| Model-level (`model_max_budget`) | `test_unit_test_max_model_budget_limiter.py` | - | `test_model_max_budget_e2e::test_model_max_budget_isolates_per_model` | **covered (new)** |
| Provider (router) | `test_budget_limiter_hotpath.py` | `test_router_budget_limiter.py` | - | **covered** (router) |
| Global proxy (`litellm.max_budget`) | unit | - | - | **gap** (needs a config-level cap; not key-settable) |
## Budget mechanisms
| Mechanism | Unit | This suite (live) | Status |
|-----------|------|-------------------|--------|
| Pre-call reservation | `test_budget_reservation.py` | exercised by every enforcement test | **partial** |
| Soft budget / alerts | `SlackAlerting/test_budget_alert_types.py` | `test_soft_budget_e2e::test_soft_budget_does_not_block` | **covered (new)** (block-vs-alert; the alert side-effect itself stays unit) |
| Budget CRUD | `test_budget_endpoints.py` | `test_budget_crud_e2e` (roundtrip + delete) | **covered (new)** |
| Reset scheduling | `test_proxy_budget_reset.py` | `test_budget_crud_e2e::test_budget_duration_schedules_reset_on_key` | **covered (new)** (scheduling; actual zeroing is time-dependent -> unit) |
| Multi-window budgets | `test_multi_budget_windows.py` | - | **gap** (window setup is fiddly; left to unit for now) |
| Read budget+spend | `test_spend_management_endpoints.py` | `/key/info` asserted in CRUD + enforcement | **partial** |
## Remaining gaps (intentionally not live-tested)
- **Global proxy budget** (`litellm.max_budget`): set via proxy config, not a
per-key API, so it needs a dedicated proxy boot with that config rather than a
runtime-created entity. Out of scope for the per-entity suite.
- **Multi-window budgets**: the `budget_limits` list shape and per-window reset are
covered by `test_multi_budget_windows.py` (unit); a live version would need to
wait out a short window to see the reset, which is time-dependent.
- **Soft-budget alert delivery**: whether the Slack/email actually fires is not
observable from the proxy API; unit tests own that. The live test pins the
load-bearing behavior (soft does not block).
- **Reset zeroing after the window elapses**: time-dependent; unit tests own the
reset-job logic. The live test pins that `budget_reset_at` is scheduled.
## This suite's files
| File | Covers |
|------|--------|
| `test_budget_enforcement_e2e.py` | key / internal-user / end-user / organization / team-member hard enforcement |
| `test_model_max_budget_e2e.py` | per-model caps isolate by model |
| `test_soft_budget_e2e.py` | soft budget alerts but does not block |
| `test_tag_budget_e2e.py` | proxy-level tag budget blocks tagged requests, spares others |
| `test_budget_crud_e2e.py` | `/budget/*` CRUD roundtrip + delete + `budget_reset_at` scheduling |
## Pattern + timing
Create the entity with a tiny `max_budget`, drive spend until a `budget_exceeded`
block. The enforcement helper is two-phase: a fast warmup (key/user/org/member/tag/
model block within ~2 calls off real-time counters), then a poll across the ~60s
batch-write window (end-user enforcement reads table spend that lags). Skip on a
non-budget error (provider down / key missing); fail if the budget is never
enforced. Chat tests use `gpt-5.5` (the model with a working key on the reference
proxy); swap the literal if your proxy differs.

View file

@ -0,0 +1,418 @@
"""Client for budget e2e tests: the shared Gateway plus budget-bearing entity
management (user / team / team-member / org / customer / tag / budget-table) and
info reads.
Over-budget surfaces as a ``budget_exceeded`` error; ``is_budget_block`` detects it
on a chat outcome. Create methods return the new id and raise on failure; tests
register the matching delete with ``resources.defer(...)`` for cleanup. The request
and response models are co-located here because only this suite uses them.
"""
from __future__ import annotations
from dataclasses import dataclass
from pydantic import AliasPath, BaseModel, Field, RootModel
from e2e_gateway import Gateway, build_gateway
from e2e_http import NoBody, StreamingResponse, Success, unwrap
from models import (
BudgetWindow,
ChatBody,
ChatMessage,
ChatMetadata,
KeyGenerateBody,
ModelBudgetEntry,
)
class UserNewBody(BaseModel):
max_budget: float
class UserNewResponse(BaseModel):
user_id: str
class UserDeleteBody(BaseModel):
user_ids: list[str]
class CustomerNewBody(BaseModel):
user_id: str
max_budget: float
class OrgNewBody(BaseModel):
organization_alias: str
max_budget: float
class OrgNewResponse(BaseModel):
organization_id: str
class OrgDeleteBody(BaseModel):
organization_ids: list[str]
class TeamMember(BaseModel):
role: str
user_id: str
class TeamNewBody(BaseModel):
team_alias: str
max_budget: float | None = None
organization_id: str | None = None
budget_limits: list[BudgetWindow] | None = None
class TeamNewResponse(BaseModel):
team_id: str
class TeamDeleteBody(BaseModel):
team_ids: list[str]
class TeamMemberAddBody(BaseModel):
team_id: str
member: TeamMember
max_budget_in_team: float | None = None
class TeamMemberUpdateBody(BaseModel):
team_id: str
user_id: str
max_budget_in_team: float | None = None
budget_duration: str | None = None
class TeamMembershipRow(BaseModel):
user_id: str | None = None
budget_reset_at: str | None = Field(
default=None,
validation_alias=AliasPath("litellm_budget_table", "budget_reset_at"),
)
class TeamInfoParams(BaseModel):
team_id: str
class TeamInfoResponse(BaseModel):
team_memberships: list[TeamMembershipRow] = []
class TagNewBody(BaseModel):
name: str
max_budget: float
class TagDeleteBody(BaseModel):
name: str
class BudgetNewBody(BaseModel):
max_budget: float
soft_budget: float | None = None
budget_duration: str | None = None
class BudgetNewResponse(BaseModel):
budget_id: str
class BudgetDeleteBody(BaseModel):
id: str
class BudgetInfoBody(BaseModel):
budgets: list[str]
class BudgetRow(BaseModel):
budget_id: str | None = None
max_budget: float | None = None
soft_budget: float | None = None
budget_duration: str | None = None
budget_reset_at: str | None = None
class BudgetInfoResponse(RootModel[list[BudgetRow]]):
pass
def is_budget_block(result: StreamingResponse) -> bool:
"""True if the call was rejected for being over budget (vs a provider error)."""
return not result.ok and "budget_exceeded" in result.body
def model_budget(model: str, limit: float, period: str = "30d") -> dict[str, ModelBudgetEntry]:
"""A model_max_budget entry: per-model cap with a reset window."""
return {model: ModelBudgetEntry(budget_limit=limit, time_period=period)}
@dataclass(frozen=True, slots=True)
class BudgetClient:
gateway: Gateway
# ---- generic key ops (delegate to the shared Gateway) ---------------
def generate_key(
self,
*,
models: list[str] | None = None,
max_budget: float | None = None,
soft_budget: float | None = None,
budget_duration: str | None = None,
budget_id: str | None = None,
user_id: str | None = None,
team_id: str | None = None,
model_max_budget: dict[str, ModelBudgetEntry] | None = None,
budget_limits: list[BudgetWindow] | None = None,
) -> str:
return self.gateway.generate_key(
KeyGenerateBody(
models=models or [],
max_budget=max_budget,
soft_budget=soft_budget,
budget_duration=budget_duration,
budget_id=budget_id,
user_id=user_id,
team_id=team_id,
model_max_budget=model_max_budget,
budget_limits=budget_limits,
)
)
def delete_key(self, key: str) -> None:
self.gateway.delete_key(key)
def delete_customers(self, user_ids: list[str]) -> None:
self.gateway.delete_customers(user_ids)
# ---- chat (raw HTTP outcome: a budget block surfaces as a non-2xx) --
def chat(
self,
key: str,
model: str,
content: str,
*,
max_tokens: int | None = None,
user: str | None = None,
tags: list[str] | None = None,
) -> StreamingResponse:
return self.gateway.transport.send(
"/chat/completions",
headers=self.gateway.transport.bearer(key),
json=ChatBody(
model=model,
messages=[ChatMessage(role="user", content=content)],
max_tokens=max_tokens,
user=user,
metadata=ChatMetadata(tags=tags) if tags else None,
),
)
# ---- internal user --------------------------------------------------
def create_user(self, *, max_budget: float) -> str:
return unwrap(
self.gateway.transport.post(
"/user/new",
headers=self.gateway.transport.master,
json=UserNewBody(max_budget=max_budget),
response_type=UserNewResponse,
)
).user_id
def delete_user(self, user_id: str) -> None:
_ = self.gateway.transport.post(
"/user/delete",
headers=self.gateway.transport.master,
json=UserDeleteBody(user_ids=[user_id]),
response_type=NoBody,
)
# ---- customer / end-user -------------------------------------------
def create_customer(self, customer_id: str, *, max_budget: float) -> str:
resp = self.gateway.transport.send(
"/customer/new",
headers=self.gateway.transport.master,
json=CustomerNewBody(user_id=customer_id, max_budget=max_budget),
)
assert resp.ok, resp.body
return customer_id
# ---- organization ---------------------------------------------------
def create_org(self, *, max_budget: float, alias: str) -> str:
return unwrap(
self.gateway.transport.post(
"/organization/new",
headers=self.gateway.transport.master,
json=OrgNewBody(organization_alias=alias, max_budget=max_budget),
response_type=OrgNewResponse,
)
).organization_id
def delete_org(self, org_id: str) -> None:
_ = self.gateway.transport.delete(
"/organization/delete",
headers=self.gateway.transport.master,
json=OrgDeleteBody(organization_ids=[org_id]),
response_type=NoBody,
)
# ---- team -----------------------------------------------------------
def create_team(
self,
*,
alias: str,
max_budget: float | None = None,
organization_id: str | None = None,
budget_limits: list[BudgetWindow] | None = None,
) -> str:
return unwrap(
self.gateway.transport.post(
"/team/new",
headers=self.gateway.transport.master,
json=TeamNewBody(
team_alias=alias,
max_budget=max_budget,
organization_id=organization_id,
budget_limits=budget_limits,
),
response_type=TeamNewResponse,
)
).team_id
def delete_team(self, team_id: str) -> None:
_ = self.gateway.transport.post(
"/team/delete",
headers=self.gateway.transport.master,
json=TeamDeleteBody(team_ids=[team_id]),
response_type=NoBody,
)
def add_team_member(self, team_id: str, user_id: str, *, max_budget_in_team: float | None = None) -> None:
resp = self.gateway.transport.send(
"/team/member_add",
headers=self.gateway.transport.master,
json=TeamMemberAddBody(
team_id=team_id,
member=TeamMember(role="user", user_id=user_id),
max_budget_in_team=max_budget_in_team,
),
)
assert resp.ok, resp.body
def update_team_member(
self,
team_id: str,
user_id: str,
*,
max_budget_in_team: float | None = None,
budget_duration: str | None = None,
) -> None:
resp = self.gateway.transport.send(
"/team/member_update",
headers=self.gateway.transport.master,
json=TeamMemberUpdateBody(
team_id=team_id,
user_id=user_id,
max_budget_in_team=max_budget_in_team,
budget_duration=budget_duration,
),
)
assert resp.ok, resp.body
def member_budget_reset_at(self, team_id: str, user_id: str) -> str | None:
"""The member's per-team budget_reset_at as /team/info reports it, or None if
no reset is scheduled. The reset job advances this each time the window
elapses; a job that skips the row leaves it pinned forever."""
result = self.gateway.transport.get(
"/team/info",
headers=self.gateway.transport.master,
params=TeamInfoParams(team_id=team_id),
response_type=TeamInfoResponse,
)
match result:
case Success(data=data):
return next(
(row.budget_reset_at for row in data.team_memberships if row.user_id == user_id),
None,
)
case _:
return None
# ---- tag ------------------------------------------------------------
def create_tag(self, name: str, *, max_budget: float) -> str:
resp = self.gateway.transport.send(
"/tag/new",
headers=self.gateway.transport.master,
json=TagNewBody(name=name, max_budget=max_budget),
)
assert resp.ok, resp.body
return name
def delete_tag(self, name: str) -> None:
_ = self.gateway.transport.post(
"/tag/delete",
headers=self.gateway.transport.master,
json=TagDeleteBody(name=name),
response_type=NoBody,
)
# ---- budget table ---------------------------------------------------
def create_budget(
self,
*,
max_budget: float,
soft_budget: float | None = None,
budget_duration: str | None = None,
) -> str:
return unwrap(
self.gateway.transport.post(
"/budget/new",
headers=self.gateway.transport.master,
json=BudgetNewBody(
max_budget=max_budget,
soft_budget=soft_budget,
budget_duration=budget_duration,
),
response_type=BudgetNewResponse,
)
).budget_id
def delete_budget(self, budget_id: str) -> None:
_ = self.gateway.transport.post(
"/budget/delete",
headers=self.gateway.transport.master,
json=BudgetDeleteBody(id=budget_id),
response_type=NoBody,
)
def budget_info(self, budget_id: str) -> tuple[BudgetRow, ...]:
result = self.gateway.transport.post(
"/budget/info",
headers=self.gateway.transport.master,
json=BudgetInfoBody(budgets=[budget_id]),
response_type=BudgetInfoResponse,
)
match result:
case Success(data=data):
return tuple(data.root)
case _:
return ()
def build_client() -> BudgetClient:
return BudgetClient(gateway=build_gateway())

View file

@ -0,0 +1,16 @@
"""Budgets suite's `client` fixture.
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
live in the parent tests/e2e/conftest.py. BudgetClient holds the shared Gateway,
so the `resources` fixture cleans up keys through it; tests register entity deletes
via `resources.defer(...)`.
"""
import pytest
from budget_client import BudgetClient, build_client
@pytest.fixture(scope="session")
def client() -> BudgetClient:
return build_client()

View file

@ -0,0 +1,60 @@
"""Live e2e for the budget management surface (no LLM calls, fast).
Covers the budget-table CRUD round-trip and that `budget_duration` schedules a
`budget_reset_at`. The actual zeroing after the window is time-dependent, so we
assert the reset is *scheduled* (now + duration), not waited out.
"""
from datetime import datetime, timezone
import pytest
from budget_client import BudgetClient
from lifecycle import ResourceManager
pytestmark = pytest.mark.e2e
def test_budget_crud_roundtrip(client: BudgetClient, resources: ResourceManager) -> None:
budget_id = client.create_budget(max_budget=12.5, soft_budget=10.0, budget_duration="30d")
resources.defer(lambda: client.delete_budget(budget_id))
rows = client.budget_info(budget_id)
assert rows, f"/budget/info returned nothing for {budget_id}"
row = rows[0]
assert row.max_budget == 12.5
assert row.soft_budget == 10.0
assert row.budget_reset_at, "budget_duration did not schedule a reset"
# Attach the budget to a key and confirm the key reflects it.
key = client.generate_key(budget_id=budget_id)
resources.defer(lambda: client.delete_key(key))
info = client.gateway.key_info(key)
linked = info.litellm_budget_table
assert info.budget_id == budget_id or (linked is not None and linked.max_budget == 12.5), (
f"key does not reflect attached budget: {info.budget_id}, {linked}"
)
def test_budget_delete_removes_it(client: BudgetClient, resources: ResourceManager) -> None:
budget_id = client.create_budget(max_budget=1.0)
resources.defer(lambda: client.delete_budget(budget_id))
client.delete_budget(budget_id)
assert not client.budget_info(budget_id), "budget still present after delete"
def test_budget_duration_schedules_reset_on_key(client: BudgetClient, resources: ResourceManager) -> None:
key = client.generate_key(max_budget=10.0, budget_duration="30d")
resources.defer(lambda: client.delete_key(key))
reset_at = client.gateway.key_info(key).budget_reset_at
assert reset_at, "budget_duration did not set budget_reset_at on the key"
# budget_duration schedules a FUTURE reset. Don't assume now+30d exactly: the
# proxy may align the reset to a calendar boundary (e.g. start of next month),
# so "30d" can land ~12 days out mid-month. Assert it's scheduled ahead.
# get current time -> assert budget from days_left - budget_duration == days_left
reset_dt = datetime.fromisoformat(str(reset_at).replace("Z", "+00:00"))
days_out = (reset_dt - datetime.now(timezone.utc)).total_seconds() / 86400
assert 0 < days_out < 40, f"reset should be scheduled ahead, got {days_out:.1f}d out"

View file

@ -0,0 +1,148 @@
"""Live e2e: a tiny max_budget on an entity actually blocks requests.
Each entity is an E2ECase (lifecycle.E2ECase) driven by run_case: init() creates
the budgeted entity + a key, run() drives spend until a `budget_exceeded` block,
teardown() deletes everything init() created (always runs, even on failure/skip).
Covers the entities with no prior live coverage - internal user, end-user,
organization, team member. See BUDGET_TEST_COVERAGE_MATRIX.md.
A non-budget error fails hard (never a skip); if calls never get blocked, budget
enforcement is broken -> fail.
"""
import time
from dataclasses import dataclass, field
from typing import Callable, List, Type
import pytest
from budget_client import BudgetClient, is_budget_block
from e2e_config import unique_marker
from e2e_http import require_successful_call
from lifecycle import run_case
pytestmark = pytest.mark.e2e
def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") -> None:
"""Send paid calls until the entity's budget blocks one. Key/user/org/member
block within a couple calls off real-time reservation counters; the end-user
budget enforces off table spend that lands on the batch write, so it takes a
few more. A non-budget error fails hard (never a skip)."""
for _ in range(40):
result = client.chat(
key,
"claude-haiku-4-5",
f"spend {unique_marker()}",
max_tokens=16,
user=user or None,
)
if is_budget_block(result):
return
require_successful_call(result)
time.sleep(2)
pytest.fail("budget never enforced within the call budget")
@dataclass
class _BudgetCase:
"""Base E2ECase: a key under some budgeted entity must get blocked.
Subclasses set up the budgeted entity in init() and register every created id
in `_undo` (run LIFO in teardown so a key is deleted before its team/org).
"""
client: BudgetClient
key: str = ""
_undo: List[Callable[[], None]] = field(
default_factory=list
) # mutable-ok: per-case teardown registry
def init(self) -> None:
raise NotImplementedError
def run(self) -> None:
_assert_budget_blocks(self.client, self.key)
def teardown(self) -> None:
for undo in reversed(self._undo):
undo()
class KeyBudgetCase(_BudgetCase):
def init(self) -> None:
self.key = self.client.generate_key(max_budget=3e-6)
self._undo.append(lambda: self.client.delete_key(self.key))
class InternalUserBudgetCase(_BudgetCase):
def init(self) -> None:
user_id = self.client.create_user(max_budget=3e-6)
self._undo.append(lambda: self.client.delete_user(user_id))
# personal key (no team) -> the user budget governs
self.key = self.client.generate_key(user_id=user_id)
self._undo.append(lambda: self.client.delete_key(self.key))
class EndUserBudgetCase(_BudgetCase):
def init(self) -> None:
customer = f"e2e-budget-cust-{unique_marker()}"
self.client.create_customer(customer, max_budget=3e-6)
self._undo.append(lambda: self.client.delete_customers([customer]))
self.key = self.client.generate_key(models=["claude-haiku-4-5"])
self._undo.append(lambda: self.client.delete_key(self.key))
self._customer = customer
def run(self) -> None:
_assert_budget_blocks(self.client, self.key, user=self._customer)
class OrganizationBudgetCase(_BudgetCase):
def init(self) -> None:
# Org carries the tiny budget; the team under it has none, so a block here
# is org-level enforcement (the historically weak link).
org_id = self.client.create_org(
max_budget=3e-6, alias=f"e2e-budget-org-{unique_marker()}"
)
self._undo.append(lambda: self.client.delete_org(org_id))
team_id = self.client.create_team(
alias=f"e2e-budget-team-{unique_marker()}", organization_id=org_id
)
self._undo.append(lambda: self.client.delete_team(team_id))
self.key = self.client.generate_key(team_id=team_id)
self._undo.append(lambda: self.client.delete_key(self.key))
class TeamMemberBudgetCase(_BudgetCase):
def init(self) -> None:
# Member's per-team budget is tiny while the team has a large budget, so a
# block proves member-level (not team-level) enforcement.
team_id = self.client.create_team(
alias=f"e2e-budget-team-{unique_marker()}", max_budget=100.0
)
self._undo.append(lambda: self.client.delete_team(team_id))
user_id = self.client.create_user(max_budget=100.0)
self._undo.append(lambda: self.client.delete_user(user_id))
self.client.add_team_member(team_id, user_id, max_budget_in_team=3e-6)
self.key = self.client.generate_key(team_id=team_id, user_id=user_id)
self._undo.append(lambda: self.client.delete_key(self.key))
def _case_id(case_cls: Type[_BudgetCase]) -> str:
return case_cls.__name__
@pytest.mark.parametrize(
"case_cls",
[
KeyBudgetCase,
InternalUserBudgetCase,
EndUserBudgetCase,
OrganizationBudgetCase,
TeamMemberBudgetCase,
],
ids=_case_id,
)
def test_budget_enforcement(
client: BudgetClient, case_cls: Type[_BudgetCase]
) -> None:
run_case(case_cls(client))

View file

@ -0,0 +1,59 @@
"""Live e2e: a key budget resets (zeroes spend) after its budget_duration.
Short budget_duration (30s) + the fast-rescheduled reset job: a key blocked for
exceeding its max_budget starts succeeding again once the duration elapses and the
reset job zeroes key.spend. Closes the reset-zeroing gap in
BUDGET_TEST_COVERAGE_MATRIX.md (reset_budget_for_litellm_keys), which the unit
suite covers but no live test did - distinct from the per-window reset in
test_multi_window_budget_e2e.py.
"""
import time
import pytest
from budget_client import BudgetClient, is_budget_block
from e2e_config import unique_marker
from e2e_http import require_successful_call
from lifecycle import ResourceManager
pytestmark = pytest.mark.e2e
def _call(client: BudgetClient, key: str):
return client.chat(
key, "claude-haiku-4-5", f"reset {unique_marker()}", max_tokens=16
)
def test_key_budget_resets_after_duration(
client: BudgetClient, resources: ResourceManager
) -> None:
key = client.generate_key(max_budget=3e-6, budget_duration="30s")
resources.defer(lambda: client.delete_key(key))
# 1. exceed the budget -> litellm returns budget_exceeded
blocked = False
for _ in range(20):
result = _call(client, key)
if is_budget_block(result):
blocked = True
break
require_successful_call(result)
time.sleep(2)
assert blocked, "key budget never enforced"
# 2. once the 30s duration elapses + the reset job runs, key.spend zeroes and
# calls flow again. The window is wall-clock-aligned, so the reset lands up to
# a window later, then the rescheduler (~15-20s) zeroes the spend; allow
# generous headroom over that. A stuck rescheduler is caught by the wait-loop
# timeout, not this elapsed bound.
start = time.monotonic()
while time.monotonic() < start + 150:
time.sleep(5)
result = _call(client, key)
if result.ok:
assert time.monotonic() - start < 120, "reset too slow for a 30s budget"
return
assert is_budget_block(result), f"non-budget error: {result.body[:200]}"
pytest.fail("key budget never reset within 150s")

View file

@ -0,0 +1,57 @@
"""Live e2e: per-model budgets (`model_max_budget`) isolate by model.
A key caps one model tiny and leaves another generous. Exhausting the capped
model must block *that* model while the other still works - proving the per-model
cap is enforced independently, not as a key-wide budget. Closes the
model_max_budget gap in BUDGET_TEST_COVERAGE_MATRIX.md.
"""
import time
import pytest
from budget_client import BudgetClient, is_budget_block, model_budget
from e2e_config import unique_marker
from e2e_http import require_successful_call
from lifecycle import ResourceManager
pytestmark = pytest.mark.e2e
CAPPED_MODEL = "claude-haiku-4-5"
FREE_MODEL = "gemini-2.5-flash"
def _call(client: BudgetClient, key: str, model: str):
result = client.chat(key, model, f"hi {unique_marker()}", max_tokens=16)
if not result.ok and not is_budget_block(result):
require_successful_call(result)
return result
def test_model_max_budget_isolates_per_model(
client: BudgetClient, resources: ResourceManager
) -> None:
key = client.generate_key(
model_max_budget={
**model_budget(CAPPED_MODEL, 1e-6),
**model_budget(FREE_MODEL, 1000.0),
}
)
resources.defer(lambda: client.delete_key(key))
# Exhaust the capped model.
blocked = False
deadline = time.monotonic() + 60
while time.monotonic() < deadline:
if is_budget_block(_call(client, key, CAPPED_MODEL)):
blocked = True
break
time.sleep(1)
assert blocked, f"{CAPPED_MODEL} per-model budget never enforced"
# The other model shares the key but has its own (large) cap -> still works.
other = _call(client, key, FREE_MODEL)
assert not is_budget_block(other), (
f"{FREE_MODEL} was blocked by {CAPPED_MODEL}'s budget; per-model caps not isolated"
)
require_successful_call(other)

View file

@ -0,0 +1,70 @@
"""Live e2e: multi-window budgets (budget_limits) enforce AND reset per window.
Short windows make the time limit reachable inside a test: a tight 30s window and
a roomy 1m window. The 30s window blocks once its tiny cap is exceeded, then - once
its 30s elapses and the reset job runs (rescheduled fast via
PROXY_BUDGET_RESCHEDULER_* in docker-compose) - the window resets and calls flow
again. Closes the multi-window gap (enforcement + per-window reset) in
BUDGET_TEST_COVERAGE_MATRIX.md, which the unit suite covered but no live test did.
"""
import time
import pytest
from budget_client import BudgetClient, is_budget_block
from e2e_config import unique_marker
from e2e_http import require_successful_call
from lifecycle import ResourceManager
from models import BudgetWindow
pytestmark = pytest.mark.e2e
WINDOW_SECONDS = 30 # the tight window; calls succeed again only after it elapses
def _call(client: BudgetClient, key: str):
return client.chat(
key, "claude-haiku-4-5", f"window {unique_marker()}", max_tokens=16
)
def test_short_window_blocks_then_resets(
client: BudgetClient, resources: ResourceManager
) -> None:
key = client.generate_key(
budget_limits=[
BudgetWindow(budget_duration=f"{WINDOW_SECONDS}s", max_budget=3e-6),
BudgetWindow(budget_duration="1m", max_budget=1.0), # roomy: never blocks
]
)
resources.defer(lambda: client.delete_key(key))
# 1. exhaust the tight window -> litellm returns budget_exceeded
start = time.monotonic()
blocked = False
for _ in range(20):
result = _call(client, key)
if is_budget_block(result):
blocked = True
break
require_successful_call(result)
time.sleep(2)
assert blocked, f"{WINDOW_SECONDS}s window never enforced"
# 2. the window resets at the next wall-clock-aligned boundary (up to a window
# after start), then the reset job (~15-20s rescheduler) zeroes the spend.
# Allow generous headroom for that alignment + rescheduler latency; a stuck
# rescheduler is caught by the wait-loop timeout, not this elapsed bound.
deadline = time.monotonic() + 150
while time.monotonic() < deadline:
time.sleep(5)
result = _call(client, key)
if result.ok:
elapsed = time.monotonic() - start
assert elapsed < WINDOW_SECONDS + 90, (
f"reset took {elapsed:.0f}s - too long for a {WINDOW_SECONDS}s window"
)
return
assert is_budget_block(result), f"non-budget error during reset wait: {result.body[:200]}"
pytest.fail(f"{WINDOW_SECONDS}s window never reset within 150s")

View file

@ -0,0 +1,35 @@
"""Live e2e: soft_budget alerts but does NOT block.
A key with a tiny `soft_budget` well under a large `max_budget`: spend crosses the
soft threshold within a couple calls, but requests keep succeeding (soft budget is
advisory). Closes the soft_budget gap in BUDGET_TEST_COVERAGE_MATRIX.md. The alert
side-effect (Slack/email) is not observable from the proxy API, so we assert the
load-bearing behavior: soft != block.
"""
import pytest
from budget_client import BudgetClient, is_budget_block
from e2e_config import unique_marker
from e2e_http import require_successful_call
from lifecycle import ResourceManager
pytestmark = pytest.mark.e2e
def test_soft_budget_does_not_block(
client: BudgetClient, resources: ResourceManager
) -> None:
# soft far below max: spend crosses soft immediately, stays under max.
key = client.generate_key(max_budget=1000.0, soft_budget=1e-9)
resources.defer(lambda: client.delete_key(key))
for _ in range(3):
result = client.chat(
key, "claude-haiku-4-5", f"hi {unique_marker()}", max_tokens=16
)
assert not is_budget_block(result), (
"soft_budget blocked a request; it must alert only, not block "
f"(body={result.body[:200]})"
)
require_successful_call(result) # any other non-2xx (e.g. provider down) is a hard fail

View file

@ -0,0 +1,144 @@
"""Live e2e: concurrent cold-counter reseeds keep the spend counter equal to DB spend (#26829).
Regression for the cross-pod spend-counter multiplication. Real requests build a key's
DB spend through the spend writer; the Redis spend counter then expires (the e2e proxy
sets a short default_redis_ttl) and goes cold. The proxy runs several workers sharing one
Redis, so a concurrent burst makes more than one worker reseed the same cold counter at
once. The fix seeds with SET NX - one worker initializes the counter at the DB spend and
the rest read it back - so the counter still equals the DB spend (plus the burst's own
small cost). The pre-#26829 additive reseed stacked the DB spend once per worker, leaving
the counter at ~N x the real spend.
The test reads the shared counter straight from Redis and asserts it equals the DB spend,
not a multiple. It also asserts the counter actually went cold before the burst, so a proxy
that never expires the counter (no short TTL) fails loudly instead of passing vacuously.
Skipped when the e2e Redis is not reachable.
"""
import hashlib
import os
import time
from concurrent.futures import ThreadPoolExecutor
from threading import Barrier
from typing import TYPE_CHECKING
import pytest
from budget_client import BudgetClient
from e2e_config import unique_marker
from e2e_http import StreamingResponse
from lifecycle import ResourceManager
if TYPE_CHECKING:
import redis
from redis.cluster import RedisCluster
pytestmark = pytest.mark.e2e
MODEL = "claude-haiku-4-5"
ACCUMULATE_CALLS = 24
BURST = 6
# proxy_batch_write_at (60s) flushes the spend to the DB and default_redis_ttl (20s)
# expires the counter; this waits out both.
COLD_WAIT_SECONDS = 80
def _redis() -> "redis.Redis[str] | RedisCluster[str]":
"""The proxy's Redis. The deployed runner sets REDIS_HOST to the serverless
ElastiCache, which is always TLS + cluster-mode; without it, fall back to a
local standalone redis for docker-compose runs."""
import redis
host = os.getenv("REDIS_HOST")
if not host:
return redis.Redis(host="localhost", port=6380, decode_responses=True, socket_connect_timeout=2)
from redis.cluster import RedisCluster
return RedisCluster(
host=host,
port=int(os.getenv("REDIS_PORT", "6379")),
ssl=True,
decode_responses=True,
socket_connect_timeout=2,
)
def _spend_counter(rds: "redis.Redis[str] | RedisCluster[str]", key: str) -> float | None:
"""The shared spend counter for `key`, or None if it is cold. A cluster client
can't run a keyspace SCAN that spans shards, so read the key directly - the stage
gateway sets no cache namespace, so the key is the bare ``spend:key:{sha256(key)}``.
A standalone client matches by suffix, so the local cache namespace (litellm.caching)
need not be hard-coded here."""
from redis.cluster import RedisCluster
digest = hashlib.sha256(key.encode()).hexdigest()
suffix = f"spend:key:{digest}"
if isinstance(rds, RedisCluster):
raw = rds.get(suffix)
return float(raw) if raw is not None else None
matches = list(rds.scan_iter(match=f"*{suffix}"))
if not matches:
return None
raw = rds.get(matches[0])
return float(raw) if raw is not None else None
def _chat(client: BudgetClient, key: str) -> StreamingResponse:
return client.chat(key, MODEL, f"reseed {unique_marker()}", max_tokens=16)
def _accumulate(client: BudgetClient, key: str, count: int) -> None:
def one(_: int) -> StreamingResponse:
return _chat(client, key)
with ThreadPoolExecutor(max_workers=8) as pool:
list(pool.map(one, range(count)))
def _burst(client: BudgetClient, key: str, count: int) -> None:
"""Fire `count` requests that start together, so multiple workers reseed the cold
counter concurrently rather than one warming it before the others arrive."""
barrier = Barrier(count)
def one(_: int) -> StreamingResponse:
barrier.wait()
return _chat(client, key)
with ThreadPoolExecutor(max_workers=count) as pool:
list(pool.map(one, range(count)))
def test_cold_counter_reseed_keeps_counter_equal_to_db_spend(
client: BudgetClient, resources: ResourceManager
) -> None:
try:
rds = _redis()
rds.ping()
except Exception as exc: # noqa: BLE001 - any connect failure means skip
pytest.skip(f"e2e redis not reachable (set REDIS_HOST/REDIS_PORT): {exc}")
key = client.generate_key(max_budget=1.0, models=[MODEL])
resources.defer(lambda: client.delete_key(key))
_accumulate(client, key, ACCUMULATE_CALLS)
time.sleep(COLD_WAIT_SECONDS)
assert _spend_counter(rds, key) is None, (
"the spend counter never went cold; default_redis_ttl must be short enough for it "
"to expire, otherwise the burst reads a warm counter and the reseed is never exercised"
)
db_spend = client.gateway.key_info(key).spend or 0.0
assert db_spend > 0, f"no DB spend accumulated from real calls: {db_spend}"
_burst(client, key, BURST)
time.sleep(3)
counter = _spend_counter(rds, key)
assert counter is not None, "the burst did not reseed the cold counter"
assert db_spend * 0.95 <= counter < db_spend * 1.7, (
f"redis spend counter {counter} does not equal DB spend {db_spend} (expected ~equal "
f"plus the burst's small cost); a near-multiple means the cold-counter reseed stacked "
f"the DB spend once per worker instead of seeding it once (#26829)"
)

View file

@ -0,0 +1,59 @@
"""Live e2e: proxy-level tag budgets block tagged requests.
A tag with a tiny budget: requests carrying that tag get blocked once the tag's
spend is exceeded, while a request with a different tag (no budget) still works.
Closes the proxy-level tag-budget gap in BUDGET_TEST_COVERAGE_MATRIX.md (today
only router-level tag budgets are tested).
"""
import time
import pytest
from budget_client import BudgetClient, is_budget_block
from e2e_config import unique_marker
from e2e_http import require_successful_call
from lifecycle import ResourceManager
pytestmark = pytest.mark.e2e
TINY_BUDGET = 1e-6
def _tagged_call(client: BudgetClient, key: str, tag: str):
result = client.chat(
key,
"claude-haiku-4-5",
f"hi {unique_marker()}",
tags=[tag],
max_tokens=16,
)
if not result.ok and not is_budget_block(result):
require_successful_call(result)
return result
def test_tag_budget_blocks_tagged_requests(
client: BudgetClient, scoped_key: str, resources: ResourceManager
) -> None:
budgeted_tag = f"e2e-budget-tag-{unique_marker()}"
client.create_tag(budgeted_tag, max_budget=TINY_BUDGET)
resources.defer(lambda: client.delete_tag(budgeted_tag))
# Requests under the budgeted tag get blocked once its spend is exceeded.
blocked = False
deadline = time.monotonic() + 60
while time.monotonic() < deadline:
if is_budget_block(_tagged_call(client, scoped_key, budgeted_tag)):
blocked = True
break
time.sleep(1)
assert blocked, f"tag budget for {budgeted_tag!r} never enforced"
# A request with an unbudgeted tag on the same key is unaffected.
free_tag = f"e2e-free-tag-{unique_marker()}"
other = _tagged_call(client, scoped_key, free_tag)
assert not is_budget_block(other), (
f"unbudgeted tag {free_tag!r} was blocked by {budgeted_tag!r}'s budget"
)
require_successful_call(other)

View file

@ -0,0 +1,107 @@
"""Live e2e: a team member's per-team budget attributes spend and enforces a cap.
The team carries a large budget while the one enrolled member is capped at a tiny
per-team budget, so any block is member-level, not team-level. Two scenarios share
that single member:
- attribution: the member's calls land in the spend logs tagged with both the team_id
and the member's user_id, so per-member spend can be billed back
- enforcement: once the member's spend passes the per-team budget, calls are blocked
with budget_exceeded while the team's own budget is nowhere near exhausted
Per-member budgets enforce off batch-written spend (~60s), so a quick burst all goes
through; the block only lands once that spend flushes.
"""
import time
from collections.abc import Iterator
from dataclasses import dataclass
import pytest
from budget_client import BudgetClient, is_budget_block
from e2e_config import unique_marker
from e2e_http import Success, require_successful_call
from lifecycle import ResourceManager
from models import ChatBody, ChatMessage
pytestmark = pytest.mark.e2e
MODEL = "claude-haiku-4-5"
TEAM_BUDGET = 100.0
MEMBER_BUDGET = 3e-6
BURST = 6
@dataclass(frozen=True, slots=True)
class _Member:
team_id: str
user_id: str
key: str
@pytest.fixture(scope="class")
def member(client: BudgetClient) -> Iterator[_Member]:
"""A team with a large budget plus one member capped at a tiny per-team budget,
and that member's key. Shared across the class; torn down when it finishes.
Cleanups register progressively and run LIFO best-effort through ResourceManager,
so a partial-setup failure still releases what came before and one failed delete
never strands the rest on the shared proxy."""
resources = ResourceManager(client=client.gateway)
try:
marker = unique_marker()
team_id = client.create_team(alias=f"e2e-team-member-{marker}", max_budget=TEAM_BUDGET)
resources.defer(lambda: client.delete_team(team_id))
user_id = client.create_user(max_budget=TEAM_BUDGET)
resources.defer(lambda: client.delete_user(user_id))
client.add_team_member(team_id, user_id, max_budget_in_team=MEMBER_BUDGET)
key = client.generate_key(team_id=team_id, user_id=user_id)
resources.defer(lambda: client.delete_key(key))
yield _Member(team_id=team_id, user_id=user_id, key=key)
finally:
resources.teardown()
def _send(client: BudgetClient, key: str) -> str | None:
"""One member call; its response id (== the spend-log request_id) if it went
through, else None."""
match client.gateway.chat(
key,
ChatBody(
model=MODEL,
messages=[ChatMessage(role="user", content=f"hi {unique_marker()}")],
max_tokens=16,
),
):
case Success(data=response):
return response.id
case _:
return None
class TestTeamMemberBudget:
def test_member_spend_attributed_to_team_and_user(self, client: BudgetClient, member: _Member) -> None:
sent = frozenset(rid for rid in (_send(client, member.key) for _ in range(BURST)) if rid)
assert sent, "no member call went through; cannot check attribution"
rows = client.gateway.poll_logs_for_key(
member.key, predicate=lambda rs: bool(sent & {r.request_id for r in rs})
)
logged = [row for row in rows if row.request_id in sent]
assert logged, f"none of the member's {len(sent)} calls reached the spend logs"
for row in logged:
assert row.team_id == member.team_id, (
f"call {row.request_id} logged under team {row.team_id}, not the member's team {member.team_id}"
)
assert row.user == member.user_id, (
f"call {row.request_id} logged under user {row.user}, not member {member.user_id}"
)
def test_member_spend_over_budget_is_blocked(self, client: BudgetClient, member: _Member) -> None:
for _ in range(40):
result = client.chat(member.key, MODEL, f"spend {unique_marker()}", max_tokens=16)
if is_budget_block(result):
return
require_successful_call(result)
time.sleep(2)
pytest.fail("per-member budget never enforced within the call budget")

View file

@ -0,0 +1,47 @@
import time
from datetime import datetime
import pytest
from budget_client import BudgetClient
from e2e_config import unique_marker
from e2e_http import require_successful_call
from lifecycle import ResourceManager
pytestmark = pytest.mark.e2e
MEMBER_BUDGET = 1.0 # default member budget is $50, we're testing with a smaller value
def _as_datetime(value: str) -> datetime:
return datetime.fromisoformat(value.replace("Z", "+00:00"))
def test_team_member_budget_reset_keeps_advancing(client: BudgetClient, resources: ResourceManager) -> None:
team_id = client.create_team(alias=f"e2e-member-reset-{unique_marker()}", max_budget=100.0)
resources.defer(lambda: client.delete_team(team_id))
user_id = client.create_user(max_budget=100.0)
resources.defer(lambda: client.delete_user(user_id))
# add the member, then update them onto a short per-team budget window
client.add_team_member(team_id, user_id, max_budget_in_team=MEMBER_BUDGET)
client.update_team_member(team_id, user_id, max_budget_in_team=MEMBER_BUDGET, budget_duration="30s")
scheduled = client.member_budget_reset_at(team_id, user_id)
assert scheduled, "updating the member with a budget_duration set no budget_reset_at"
first_reset = _as_datetime(scheduled)
# the member can spend within the team while the window is live
key = client.generate_key(team_id=team_id, user_id=user_id)
resources.defer(lambda: client.delete_key(key))
require_successful_call(client.chat(key, "claude-haiku-4-5", f"reset {unique_marker()}", max_tokens=16))
# once the window elapses the reset job must move budget_reset_at forward; a job
# that skips the member's budget row (the #25109 regression) leaves it pinned at
# first_reset forever
deadline = time.monotonic() + 150
while time.monotonic() < deadline:
time.sleep(5)
current = client.member_budget_reset_at(team_id, user_id)
if current and _as_datetime(current) > first_reset:
return
pytest.fail(f"member budget_reset_at never advanced past {first_reset.isoformat()} in 150s")

View file

@ -0,0 +1,79 @@
"""Live e2e: a team's multi-window budgets (budget_limits) enforce AND reset per window.
The team analog of test_multi_window_budget_e2e.py (which covers keys). A team is
created with a tight 30s window and a roomy 1m window; a key on that team blocks once
the tight window's cap is exceeded, then - once the 30s elapses and the reset job runs
(rescheduled fast via PROXY_BUDGET_RESCHEDULER_* in docker-compose) - the window resets
and calls flow again. This exercises the reset_budget_windows TEAM branch (raw SQL over
LiteLLM_TeamTable.budget_limits, the literal #25109 path), which had no live coverage.
Fails at team creation today: /team/new writes the raw window list straight to the
Json? column, where Prisma rejects it (500), unlike the key path and /team/update which
json.dumps it first. Marked xfail(strict=True) so the suite stays green while the bug
persists and flips to a failure the moment the write is fixed and the marker should be
removed.
"""
import time
import pytest
from budget_client import BudgetClient, is_budget_block
from e2e_config import unique_marker
from e2e_http import require_successful_call
from lifecycle import ResourceManager
from models import BudgetWindow
pytestmark = pytest.mark.e2e
WINDOW_SECONDS = 30
def _call(client: BudgetClient, key: str):
return client.chat(key, "claude-haiku-4-5", f"team-window {unique_marker()}", max_tokens=16)
@pytest.mark.xfail(
strict=True,
reason="known proxy bug: /team/new writes budget_limits straight to the Json? "
"column and Prisma rejects it (500), unlike the key path and /team/update which "
"json.dumps first; remove this marker once that write is fixed",
)
def test_team_short_window_blocks_then_resets(client: BudgetClient, resources: ResourceManager) -> None:
team_id = client.create_team(
alias=f"e2e-team-window-{unique_marker()}",
budget_limits=[
BudgetWindow(budget_duration=f"{WINDOW_SECONDS}s", max_budget=3e-6),
BudgetWindow(budget_duration="1m", max_budget=1.0), # roomy: never blocks
],
)
resources.defer(lambda: client.delete_team(team_id))
key = client.generate_key(team_id=team_id)
resources.defer(lambda: client.delete_key(key))
# 1. exhaust the tight window -> litellm returns budget_exceeded
start = time.monotonic()
blocked = False
for _ in range(20):
result = _call(client, key)
if is_budget_block(result):
blocked = True
break
require_successful_call(result)
time.sleep(2)
assert blocked, f"team {WINDOW_SECONDS}s window never enforced"
# 2. the window resets at the next wall-clock-aligned boundary (up to a window
# after start), then the reset job (~15-20s rescheduler) zeroes the spend.
# Allow generous headroom for that alignment + rescheduler latency; a stuck
# rescheduler is caught by the wait-loop timeout, not this elapsed bound.
deadline = time.monotonic() + 150
while time.monotonic() < deadline:
time.sleep(5)
result = _call(client, key)
if result.ok:
elapsed = time.monotonic() - start
assert elapsed < WINDOW_SECONDS + 90, f"reset took {elapsed:.0f}s - too long for a {WINDOW_SECONDS}s window"
return
assert is_budget_block(result), f"non-budget error during reset wait: {result.body[:200]}"
pytest.fail(f"team {WINDOW_SECONDS}s window never reset within 150s")

120
tests/e2e/conftest.py Normal file
View file

@ -0,0 +1,120 @@
"""Shared fixtures for all live e2e suites under tests/e2e/.
Design rule: skip on environment, fail on behavior. Live tests (marked `e2e`)
skip when no proxy answers; once a request reaches the proxy, behavior is
asserted. Pure unit coverage of the harness itself carries no `e2e` marker and
runs regardless of whether a proxy is up.
Lifecycle: the `resources` fixture maps the init -> run -> teardown contract
(lifecycle.E2ECase) onto pytest - setup is init(), the test body is run(), and
teardown deletes every resource the test created on the long-lived proxy.
Each suite provides its own `client` fixture (a lifecycle.ResourceClient); these
shared fixtures build on it.
"""
import functools
import sys
from pathlib import Path
from typing import Iterator
import pytest
import requests
from e2e_config import CONTROL_PLANE_BASE_URL, PROXY_BASE_URL
from lifecycle import GatewayProvider, ResourceManager
_E2E_TEST_RAN = pytest.StashKey[bool]()
def pytest_configure(config: pytest.Config) -> None:
config.addinivalue_line(
"markers",
"e2e: live test that requires a running proxy and real provider keys",
)
def _liveness_reason(label: str, base_url: str) -> str | None:
"""None if `base_url` answers its liveness probe, else a skip reason."""
try:
resp = requests.get(f"{base_url}/health/liveliness", timeout=5)
except requests.RequestException as exc:
return f"No live {label} at {base_url}: {exc}"
if resp.status_code >= 500:
return f"{label} at {base_url} returned {resp.status_code}"
return None
@functools.lru_cache(maxsize=1)
def _proxy_skip_reason() -> str | None:
"""Probe the proxy once per session. None if it answers, else a skip reason. In
a split deployment the management/admin control plane is a separate service, so
require it too (when it differs) - else its tests would fail rather than skip."""
reason = _liveness_reason("proxy", PROXY_BASE_URL)
if reason is not None:
return reason
if CONTROL_PLANE_BASE_URL != PROXY_BASE_URL:
return _liveness_reason("control plane", CONTROL_PLANE_BASE_URL)
return None
def pytest_runtest_setup(item: pytest.Item) -> None:
"""Skip `e2e`-marked tests unless a proxy answers its liveness probe. Unmarked
tests (unit coverage of the harness) don't touch the proxy, so they run even
when none is up."""
if item.get_closest_marker("e2e") is None:
return
reason = _proxy_skip_reason()
if reason is not None:
pytest.skip(reason)
def pytest_runtest_call(item: pytest.Item) -> None:
"""Mark that an e2e test body actually ran (not skipped at setup). Skipped
sessions never reach this hook, so the session-finish cleanup can use it as a
guard before truncating the spend-log DB. Tests under `tests/e2e/` without the
`e2e` marker (pure unit coverage for the harness itself) never hit the proxy,
so they must not arm the destructive DB truncate."""
if item.get_closest_marker("e2e") is None:
return
item.session.stash[_E2E_TEST_RAN] = True
def pytest_sessionfinish(session: pytest.Session, exitstatus: int) -> None:
"""Once the whole e2e session is done (all suites), truncate the spend logs so
the DB doesn't accumulate test rows. Skipped sessions (no live proxy, no test
actually executed) leave the DB alone so a `DATABASE_URL` pointing at a shared
instance is never wiped without an e2e run. Best-effort: a cleanup failure (no
DB reachable) must not fail the run. The spend_tracking dir goes on sys.path
only for this import and is removed after, so a broader `pytest tests/` run is
not left with a mutated path."""
if not session.stash.get(_E2E_TEST_RAN, False):
return
spend_dir = str(Path(__file__).parent / "spend_tracking")
sys.path.insert(0, spend_dir)
try:
from spend_e2e_client import reset_spend_logs # pyright: ignore
reset_spend_logs()
except Exception as exc: # noqa: BLE001 - cleanup is best-effort
print(f"spend-log cleanup skipped: {exc}")
finally:
if spend_dir in sys.path:
sys.path.remove(spend_dir)
@pytest.fixture
def resources(client: GatewayProvider) -> Iterator[ResourceManager]:
"""init -> run -> teardown: create a manager, run the test, release resources.
Cleanup goes through the shared Gateway, whatever the suite's client adds."""
manager = ResourceManager(client=client.gateway)
manager.init()
yield manager
manager.teardown()
@pytest.fixture
def scoped_key(resources: ResourceManager) -> str:
"""A fresh all-models key per test, auto-deleted by the resources teardown."""
return resources.key()

34
tests/e2e/e2e_config.py Normal file
View file

@ -0,0 +1,34 @@
"""Generic configuration for live e2e tests against a running LiteLLM proxy.
Shared by every e2e suite under tests/e2e/. Values come from the
environment so the same tests run against localhost or a deployed proxy.
"""
import os
import uuid
PROXY_BASE_URL = os.environ.get("LITELLM_PROXY_URL", "http://localhost:4000").rstrip("/")
MASTER_KEY = os.environ.get("LITELLM_MASTER_KEY", "sk-1234")
# Control-plane (management/admin) base URL. In a split control-plane/data-plane
# deployment the LLM data plane (PROXY_BASE_URL: /chat, /embeddings, native
# passthrough) and the management API (keys, users, teams, orgs, budgets, spend,
# model info, /openapi.json) are served by *different* services. The suite drives
# both through one Transport that routes by path (see transport.SplitTransport).
# Defaults to PROXY_BASE_URL so a monolithic proxy serving everything on one URL
# behaves exactly as before.
CONTROL_PLANE_BASE_URL = os.environ.get(
"LITELLM_CONTROL_PLANE_URL", PROXY_BASE_URL
).rstrip("/")
# Writes on the proxy are eventually consistent (e.g. spend rows flush on
# proxy_batch_write_at, ~60s). Read-backs poll to this deadline, never sleep-once.
POLL_TIMEOUT = float(os.environ.get("E2E_POLL_TIMEOUT", "120"))
POLL_INTERVAL = float(os.environ.get("E2E_POLL_INTERVAL", "5"))
REQUEST_TIMEOUT = float(os.environ.get("E2E_REQUEST_TIMEOUT", "60"))
def unique_marker() -> str:
"""A short unique token per call/run, so concurrent runs and the shared
response cache never collide on prompts, tags, or customer ids."""
return uuid.uuid4().hex[:12]

211
tests/e2e/e2e_gateway.py Normal file
View file

@ -0,0 +1,211 @@
"""Gateway: the shared proxy operations, DI'd into every client (composition).
A frozen-slots dataclass holding a Transport plus poll config. Clients hold a
Gateway and add their own route methods; the lifecycle ResourceManager uses the
Gateway's key/customer methods for cleanup. Read-backs are eventually consistent
(proxy_batch_write_at ~60s) so they poll to a deadline.
"""
from __future__ import annotations
import time
from collections.abc import Callable
from dataclasses import dataclass
from e2e_http import (
NoBody,
ProbeResult,
Result,
StreamingResponse,
Success,
unwrap,
)
from models import (
ChatBody,
ChatResponse,
CustomerDeleteBody,
EmbedBody,
EmbedResponse,
KeyDeleteBody,
KeyGenerateBody,
KeyGenerateResponse,
KeyInfo,
KeyInfoParams,
KeyInfoResponse,
ModelInfoEntry,
ModelInfoResponse,
SpendLogRow,
SpendLogs,
SpendLogsParams,
)
from e2e_config import (
CONTROL_PLANE_BASE_URL,
MASTER_KEY,
POLL_INTERVAL,
POLL_TIMEOUT,
PROXY_BASE_URL,
REQUEST_TIMEOUT,
)
from transport import HttpTransport, SplitTransport, Transport
RowsPredicate = Callable[[list[SpendLogRow]], bool]
@dataclass(frozen=True, slots=True)
class Gateway:
transport: Transport
poll_timeout: float = 120.0
poll_interval: float = 5.0
# ---- keys / customers (satisfies lifecycle.ResourceClient) ----------
def generate_key(self, body: KeyGenerateBody) -> str:
return unwrap(
self.transport.post(
"/key/generate",
headers=self.transport.master,
json=body,
response_type=KeyGenerateResponse,
)
).key
def delete_key(self, key: str) -> None:
_ = self.transport.post(
"/key/delete",
headers=self.transport.master,
json=KeyDeleteBody(keys=[key]),
response_type=NoBody,
)
def delete_customers(self, user_ids: list[str]) -> None:
if not user_ids:
return
_ = self.transport.post(
"/customer/delete",
headers=self.transport.master,
json=CustomerDeleteBody(user_ids=user_ids),
response_type=NoBody,
)
def key_info(self, key: str) -> KeyInfo:
return unwrap(
self.transport.get(
"/key/info",
headers=self.transport.master,
params=KeyInfoParams(key=key),
response_type=KeyInfoResponse,
)
).info
def model_info(self) -> list[ModelInfoEntry]:
"""Every configured deployment with the price the proxy resolved for it
(config override merged over cost-map defaults)."""
return unwrap(
self.transport.get(
"/model/info",
headers=self.transport.master,
params=NoBody(),
response_type=ModelInfoResponse,
)
).data
# ---- LLM calls ------------------------------------------------------
def chat(self, key: str, body: ChatBody) -> Result[ChatResponse]:
return self.transport.post(
"/chat/completions",
headers=self.transport.bearer(key),
json=body,
response_type=ChatResponse,
)
def chat_stream(self, key: str, body: ChatBody) -> StreamingResponse:
return self.transport.stream(
"/chat/completions", headers=self.transport.bearer(key), json=body
)
def embed(self, key: str, body: EmbedBody) -> Result[EmbedResponse]:
return self.transport.post(
"/embeddings",
headers=self.transport.bearer(key),
json=body,
response_type=EmbedResponse,
)
# ---- spend read-back ------------------------------------------------
def spend_logs(self, params: SpendLogsParams) -> list[SpendLogRow]:
result = self.transport.get(
"/spend/logs",
headers=self.transport.master,
params=params,
response_type=SpendLogs,
)
match result:
case Success(data=logs):
return logs.root
case _:
return []
def poll_logs_for_key(
self, key: str, *, min_rows: int = 1, predicate: RowsPredicate | None = None
) -> list[SpendLogRow]:
return self._poll(
lambda: self.spend_logs(SpendLogsParams(api_key=key)), min_rows, predicate
)
def poll_logs_for_request_id(
self,
request_id: str,
*,
min_rows: int = 1,
predicate: RowsPredicate | None = None,
) -> list[SpendLogRow]:
return self._poll(
lambda: self.spend_logs(SpendLogsParams(request_id=request_id)),
min_rows,
predicate,
)
def _poll(
self,
fetch: Callable[[], list[SpendLogRow]],
min_rows: int,
predicate: RowsPredicate | None,
) -> list[SpendLogRow]:
deadline = time.monotonic() + self.poll_timeout
rows: list[SpendLogRow] = []
while time.monotonic() < deadline:
rows = fetch()
if len(rows) >= min_rows and (predicate is None or predicate(rows)):
return rows
time.sleep(self.poll_interval)
return rows
# ---- route probe ----------------------------------------------------
def probe(self, path: str, *, params: NoBody) -> ProbeResult:
return self.transport.probe(path, params=params)
def build_gateway() -> Gateway:
"""The Gateway every suite's client is built from: a SplitTransport that routes
LLM calls to the data plane (PROXY_BASE_URL) and management/admin calls to the
control plane (CONTROL_PLANE_BASE_URL), with the shared poll budget. The two
base URLs are the same for a monolithic proxy, so routing is then a no-op."""
return Gateway(
transport=SplitTransport(
data=HttpTransport(
base_url=PROXY_BASE_URL,
master_key=MASTER_KEY,
request_timeout=REQUEST_TIMEOUT,
),
control=HttpTransport(
base_url=CONTROL_PLANE_BASE_URL,
master_key=MASTER_KEY,
request_timeout=REQUEST_TIMEOUT,
),
),
poll_timeout=POLL_TIMEOUT,
poll_interval=POLL_INTERVAL,
)

306
tests/e2e/e2e_http.py Normal file
View file

@ -0,0 +1,306 @@
"""The ONLY module permitted to call ``requests.*``.
Enforced by tests/code_coverage_tests/check_e2e_no_raw_requests.py. Every request
body / query / header / response is a pydantic model; outcomes are a tagged union
(``Result[R]``) so callers ``match`` on them instead of catching exceptions.
Named e2e_http (not http) so it does not shadow the stdlib ``http`` package that
requests itself imports.
"""
from __future__ import annotations
from typing import Generic, Iterator, Literal, NewType, TypeVar, cast
import pytest
import requests
from pydantic import BaseModel, ConfigDict, Field
URL = NewType("URL", str)
class Headers(BaseModel):
"""Base for header models. Subclasses may alias to hyphenated header names
(e.g. ``x-litellm-api-key``); serialization uses by_alias."""
model_config = ConfigDict(populate_by_name=True)
class AuthHeaders(Headers):
# litellm accepts either; set whichever the call needs, leave the other None.
authorization: str | None = None
x_litellm_api_key: str | None = Field(default=None, alias="x-litellm-api-key")
class NoBody(BaseModel):
"""Empty body/query for routes that take none."""
# ---------- Result types ----------
R = TypeVar("R", bound=BaseModel)
class Success(BaseModel, Generic[R]):
kind: Literal["success"] = "success"
data: R
class NetworkError(BaseModel):
kind: Literal["network"] = "network"
message: str
class UnauthorizedError(BaseModel):
kind: Literal["unauthorized"] = "unauthorized"
class RateLimitedError(BaseModel):
kind: Literal["rate_limited"] = "rate_limited"
retry_after_seconds: int | None = None
# litellm overloads 429 for budget_exceeded too, so keep the body to tell them apart.
body: str = ""
class ValidationError(BaseModel):
kind: Literal["validation"] = "validation"
message: str
class UnknownApiError(BaseModel):
kind: Literal["unknown"] = "unknown"
status_code: int
body: str
type Result[R: BaseModel] = (
Success[R]
| NetworkError
| UnauthorizedError
| RateLimitedError
| ValidationError
| UnknownApiError
)
class ProbeResult(BaseModel):
"""A route's reachability: status + body, no schema validation. Healthy ==
route exists (not 404) and the handler did not crash (not 5xx)."""
status_code: int
body: str
@property
def healthy(self) -> bool:
return 200 <= self.status_code < 500 and self.status_code != 404
class StreamingResponse(BaseModel):
"""Raw outcome for calls whose body is provider-native or streamed: status, the
x-litellm-call-id header (== SpendLogs.request_id), the content-type (which
tells streaming `text/event-stream` from non-streaming `application/json`), and
the body. Used by passthrough and streaming, where one validated JSON model
does not fit."""
status_code: int
call_id: str | None = None # x-litellm-call-id header
content_type: str | None = None
body: str
chunks: int = 0 # streamed events (0 for non-streaming)
@property
def ok(self) -> bool:
return 200 <= self.status_code < 300
@property
def is_streaming(self) -> bool:
return "text/event-stream" in (self.content_type or "")
def _hdr(resp: requests.Response, name: str) -> str | None:
value = resp.headers.get(name)
return value if isinstance(value, str) else None
def unwrap[R: BaseModel](result: Result[R]) -> R:
match result:
case Success(data=data):
return data
case _:
raise AssertionError(result)
def is_ok[R: BaseModel](result: Result[R]) -> bool:
match result:
case Success():
return True
case _:
return False
def require_successful_call(result: StreamingResponse) -> None:
"""A call that should have succeeded but didn't is a hard failure, never a skip:
if the proxy can't make a call it's expected to, the test must fail."""
if result.ok:
return
pytest.fail(
f"upstream call failed (status {result.status_code}); body={result.body[:300]}"
)
def _headers(headers: BaseModel) -> dict[str, str]:
dumped: dict[str, object] = headers.model_dump(by_alias=True, exclude_none=True)
return {key: str(value) for key, value in dumped.items()}
def _params(params: BaseModel | None) -> dict[str, str]:
if params is None:
return {}
dumped: dict[str, object] = params.model_dump(by_alias=True, exclude_none=True)
return {key: str(value) for key, value in dumped.items()}
def _classify[R: BaseModel](
resp: requests.Response, response_type: type[R]
) -> Result[R]:
if resp.status_code == 401:
return UnauthorizedError()
if resp.status_code == 429:
return RateLimitedError(body=resp.text)
if not resp.ok:
return UnknownApiError(status_code=resp.status_code, body=resp.text)
try:
return Success(data=response_type.model_validate(resp.json()))
except Exception as exc: # noqa: BLE001 - any parse/validation failure is a value
return ValidationError(message=str(exc))
def post[R: BaseModel](
url: URL,
*,
headers: BaseModel,
json: BaseModel,
response_type: type[R],
timeout: float = 30.0,
) -> Result[R]:
try:
resp = requests.post(
str(url),
headers=_headers(headers),
json=json.model_dump(by_alias=True, exclude_none=True),
timeout=timeout,
)
except requests.RequestException as exc:
return NetworkError(message=str(exc))
return _classify(resp, response_type)
def get[R: BaseModel](
url: URL,
*,
headers: BaseModel,
params: BaseModel,
response_type: type[R],
timeout: float = 30.0,
) -> Result[R]:
try:
resp = requests.get(
str(url),
headers=_headers(headers),
params=params.model_dump(by_alias=True, exclude_none=True),
timeout=timeout,
)
except requests.RequestException as exc:
return NetworkError(message=str(exc))
return _classify(resp, response_type)
def delete[R: BaseModel](
url: URL,
*,
headers: BaseModel,
json: BaseModel,
response_type: type[R],
timeout: float = 30.0,
) -> Result[R]:
try:
resp = requests.delete(
str(url),
headers=_headers(headers),
json=json.model_dump(by_alias=True, exclude_none=True),
timeout=timeout,
)
except requests.RequestException as exc:
return NetworkError(message=str(exc))
return _classify(resp, response_type)
def probe(
url: URL, *, headers: BaseModel, params: BaseModel, timeout: float = 30.0
) -> ProbeResult:
try:
resp = requests.get(
str(url),
headers=_headers(headers),
params=params.model_dump(by_alias=True, exclude_none=True),
timeout=timeout,
)
except requests.RequestException as exc:
return ProbeResult(status_code=-1, body=str(exc))
return ProbeResult(status_code=resp.status_code, body=resp.text)
def _streaming_outcome(resp: requests.Response, stream: bool) -> StreamingResponse:
call_id = _hdr(resp, "x-litellm-call-id")
content_type = _hdr(resp, "content-type")
if not stream or not (200 <= resp.status_code < 300):
return StreamingResponse(
status_code=resp.status_code,
call_id=call_id,
content_type=content_type,
body=resp.text,
)
lines = cast("Iterator[bytes]", resp.iter_lines())
chunks = sum(1 for line in lines if line)
return StreamingResponse(
status_code=resp.status_code,
call_id=call_id,
content_type=content_type,
body="<streamed>",
chunks=chunks,
)
def send(
url: URL,
*,
headers: BaseModel,
json: BaseModel,
params: BaseModel | None = None,
stream: bool = False,
timeout: float = 60.0,
) -> StreamingResponse:
"""Raw POST returning the unparsed HTTP outcome: status, full body, and the
x-litellm-call-id header. For native/passthrough bodies and for calls judged by
status rather than a typed JSON model (e.g. a budget block is a non-2xx). With
``stream=True`` the SSE body is consumed and its events counted instead."""
try:
resp = requests.post(
str(url),
headers=_headers(headers),
params=_params(params),
json=json.model_dump(by_alias=True, exclude_none=True),
stream=stream,
timeout=timeout,
)
except requests.RequestException as exc:
return StreamingResponse(status_code=-1, body=str(exc))
return _streaming_outcome(resp, stream)
def stream(
url: URL, *, headers: BaseModel, json: BaseModel, timeout: float = 60.0
) -> StreamingResponse:
"""Streaming (SSE) call: consumes the stream counting events, and captures the
x-litellm-call-id + content-type headers. Body is elided."""
return send(url, headers=headers, json=json, stream=True, timeout=timeout)

View file

@ -0,0 +1,168 @@
# This default config file aims to support most popular model providers out of the box
#In general, the model name used by the client will be the same as the ones from the provider (For example, you will use "anthropic.claude-3-5-sonnet-20240620-v1:0" when you're calling LiteLLM just like you would when calling Amazon Bedrock directly)
#In the case where there are model name conflicts, a prefix will be used (For example, the Azure and the openAI model names conflict, so when you are using Azure, you will use "azure/gpt-4o-realtime-preview-2024-10-01")
#Some model providers require additional user-specific configuration (such as Azure which requires you to specify your own api_base with your resource name, and your api_version).
#In this case, the provider is commented out, and you should uncomment it and provide your specific info
#For more detailed information about each provider, refer to the docs: https://docs.litellm.ai/docs/providers
#If you are not interested in a particular provider, just remove it from your config.yaml, and redeploy, and it will no longer show up in your LiteLLM deployment
#If a particular provider is not working, double check your .env file, and make sure you have provided a valid api key for that provider, and then redeploy
#Full details on guardrails here: https://docs.litellm.ai/docs/proxy/guardrails/bedrock
general_settings:
store_prompts_in_spend_logs: true
master_key: os.environ/LITELLM_MASTER_KEY
proxy_batch_write_at: 60
database_connection_pool_limit: 10
# disable_error_logs: True
forward_client_headers_to_llm_api: false
maximum_spend_logs_retention_period: "60d" # GSE-13389: Cleanup logs older than 60 days
maximum_spend_logs_cleanup_cron: "0 1 * * *" # 01:00 UTC daily = 18:00 PDT
database_url: os.environ/DATABASE_URL
control_plane_url: os.environ/CONTROL_PLANE_URL
alerts: ["email"]
proxy_budget_rescheduler_min_time: 15
proxy_budget_rescheduler_max_time: 20
# fallbacks: [{"gpt-4": ["anthropic.claude-3-5-sonnet-20240620-v1:0"]}] #Configure fallbacks for context window exeeded errors (In this example, we will fall back to Claude Sonnet if over 8000 tokens, which is gpt-4's limit)
# default_fallbacks: ["anthropic.claude-3-haiku-20240307-v1:0"] #Configure fallbacks for any error for every model (the above fallback configurations override this one)
# environment_variables:
# STORE_MODEL_IN_DB: 'True'
# LITELLM_LOG: "DEBUG"
litellm_settings:
drop_params: True
# Spend counters inherit this as their Redis TTL, so an idle counter goes cold and
# the next request reseeds it from the DB; kept short to exercise the cross-pod
# reseed path in test_spend_counter_reseed_e2e. Response-cache writes pass their own
# ttl and are unaffected.
default_redis_ttl: 20
request_timeout: 600
num_retries: 3
json_logs: true
store_audit_logs: True
cache: true
cache_params:
type: redis
host: redis
port: 6379
password: os.environ/REDIS_PASSWORD
namespace: litellm.caching
ttl: 16600
# max_budget: 1000000000.0 # (float) sets max budget in dollars across the entire proxy across all API keys. Note, the budget does not apply to the master key. That is the only exception.
# budget_duration: 1mo # (str) frequency of budget reset - You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"), months ("1mo").
# max_internal_user_budget: 1000000000.0 # (float) sets default budget in dollars for each internal user. (Doesn't apply to Admins. Doesn't apply to Teams. Doesn't apply to master key)
# internal_user_budget_duration: "1mo" # (str) frequency of budget reset - You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"), months ("1mo").
# success_callback: ["s3_v2"]
# failure_callback: ["s3_v2"]
# service_callback: ["datadog"]
callbacks: ["arize_phoenix", "datadog", "smtp_email", "prometheus", "otel"]
require_auth_for_metrics_endpoint: false
#type: redis-semantic
#similarity_threshold: 0.8 # similarity threshold for semantic cache
#redis_semantic_cache_embedding_model: text-embedding-ada-002 # only works with text-embedding-ada-002 for now... https://github.com/BerriAI/litellm/issues/4001
router_settings:
routing_strategy: simple-shuffle
num_retries: 3
allowed_fails: 5
cooldown_time: 30
# When gemini deployments are exhausted (provider 429 / auth), cross over to
# working models. Exercised by tests/e2e/router/test_rate_limiter.py.
fallbacks:
- gemini-2.5-flash: ["gpt-5.5", "claude-haiku-4-5"]
#ttl: Optional[float]
#default_in_memory_ttl: Optional[float]
#default_in_redis_ttl: Optional[float]
model_list:
- model_name: gpt-5.5
litellm_params:
model: openai/gpt-5.5
api_key: os.environ/OPENAI_API_KEY
- model_name: claude-haiku-4-5
litellm_params:
model: anthropic/claude-haiku-4-5
api_key: os.environ/ANTHROPIC_API_KEY
# Same underlying model via Vertex AI — distinct routing/auth path
# # (service-account JSON), so it gets its own model_name.
- model_name: gemini-2.5-flash-vertex
litellm_params:
model: vertex_ai/gemini-2.5-flash
vertex_project: os.environ/VERTEXAI_PROJECT
vertex_location: us-central1
vertex_credentials: os.environ/VERTEXAI_CREDENTIALS
- model_name: gemini-2.5-flash
litellm_params:
model: gemini/gemini-2.5-flash
api_key: os.environ/GEMINI_API_KEY
# load balancing to a different deployment, if gemini gets rate limited.
- model_name: gemini-2.5-flash
litellm_params:
model: gemini/gemini-2.5-flash
api_key: os.environ/GEMINI_API_KEY
# Custom per-token pricing exercised by llm_translation/test_custom_pricing_e2e.py.
# Rates deliberately exceed canonical gemini-2.5-flash (input 3e-7 / output 2.5e-6)
# so an override that is ignored or under-applied reports spend at the base rate
# and fails that test. The test reads these same rates back from this file.
- model_name: custom-priced-flash
litellm_params:
model: gemini/gemini-2.5-flash
api_key: os.environ/GEMINI_API_KEY
input_cost_per_token: 0.00005
output_cost_per_token: 0.0001
# embedding models
- model_name: openai-text-embedding-3-small
litellm_params:
model: openai/text-embedding-3-small
api_key: os.environ/OPENAI_API_KEY
- model_name: gemini-2-embedding
litellm_params:
model: gemini/gemini-2-embedding
api_key: os.environ/GEMINI_API_KEY
# realtime models
- model_name: openai-realtime
litellm_params:
model: openai/realtime-2
api_key: os.environ/OPENAI_API_KEY
model_info:
mode: realtime
mcp_servers:
deepwiki_mcp:
url: "https://mcp.deepwiki.com/mcp"
auth_type: none
description: "just a test"
atlassian:
url: "https://mcp.atlassian.com/v1/mcp"
auth_type: oauth2
authorization_url: https://auth.atlassian.com/authorize
guardrails:
- guardrail_name: "presidio-pii"
litellm_params:
guardrail: presidio
mode: pre_call
presidio_analyzer_api_base: os.environ/PRESIDIO_ANALYZER_API_BASE
presidio_anonymizer_api_base: os.environ/PRESIDIO_ANONYMIZER_API_BASE
default_on: false
pii_entities_config:
EMAIL_ADDRESS: BLOCK
CREDIT_CARD: BLOCK
US_SSN: BLOCK
PHONE_NUMBER: BLOCK

117
tests/e2e/lifecycle.py Normal file
View file

@ -0,0 +1,117 @@
"""Lifecycle contract and resource cleanup for stateful e2e tests.
Shared by every e2e suite under tests/e2e/. The proxy under test is
long-lived and never reset between tests, so anything a test creates (keys,
customers, teams, orgs, users, guardrails, budgets, ...) persists unless
explicitly deleted. Every check follows an init -> run -> teardown lifecycle;
teardown releases each resource init() created, even when run() raises.
In pytest terms (see conftest.py): the `resources` fixture's setup is init(),
the test body is run(), and the fixture's teardown is teardown().
"""
from dataclasses import dataclass, field
from typing import Callable, List, Protocol, runtime_checkable
from e2e_gateway import Gateway
from models import KeyGenerateBody
@runtime_checkable
class E2ECase(Protocol):
"""A stateful e2e check run against a long-lived proxy.
init() acquires resources, run() exercises behaviour and asserts, teardown()
releases everything init() created. teardown() must run even if init() fails
partway or run() raises.
"""
def init(self) -> None: ...
def run(self) -> None: ...
def teardown(self) -> None: ...
def run_case(case: E2ECase) -> None:
"""Drive a case through its lifecycle: init -> run -> teardown.
teardown always runs - even when init() fails partway or run() raises (or
skips) - so resources the case already registered on the long-lived proxy are
released. init() is inside the try because cases register cleanups
progressively (e.g. create team, then user, then key), and a failure after
the first creation must still release what came before.
"""
try:
case.init()
case.run()
finally:
case.teardown()
@runtime_checkable
class ResourceClient(Protocol):
"""Proxy operations the convenience creators use. Resource types without a
creator here are handled generically via ResourceManager.defer(). The Gateway
satisfies this."""
def generate_key(self, body: KeyGenerateBody) -> str: ...
def delete_key(self, key: str) -> None: ...
def delete_customers(self, user_ids: List[str]) -> None: ...
@runtime_checkable
class GatewayProvider(Protocol):
"""Every suite's client exposes the shared Gateway, which the resources fixture
uses for cleanup. The client adds its own route methods on top."""
@property
def gateway(self) -> Gateway: ...
@dataclass
class ResourceManager:
"""Registry of teardown actions for resources a test creates on the stateful
proxy.
Not limited to any resource type: register a cleanup with ``defer()`` for a
key, customer, team, org, user, guardrail, budget, MCP server - anything with
a delete. The two most common resources have sugar (``key``, ``customer``);
everything else is ``resources.defer(lambda: client.delete_team(team_id))``.
Cleanups run LIFO (so a resource is removed before whatever it depends on) and
best-effort (one failing cleanup never blocks the rest).
"""
client: ResourceClient
_cleanups: List[Callable[[], None]] = field(
default_factory=list
) # mutable-ok: append-only teardown registry
def init(self) -> None:
"""No global setup needed today; present for lifecycle symmetry."""
return None
def defer(self, cleanup: Callable[[], None]) -> None:
"""Register a teardown action for any resource the test just created."""
self._cleanups.append(cleanup)
def key(self) -> str:
"""Create an all-models virtual key; delete it on teardown."""
key = self.client.generate_key(KeyGenerateBody(models=[]))
self.defer(lambda: self.client.delete_key(key))
return key
def customer(self, customer_id: str) -> str:
"""Track an end-user id (from the `user` param); delete it on teardown."""
self.defer(lambda: self.client.delete_customers([customer_id]))
return customer_id
def teardown(self) -> None:
for cleanup in reversed(self._cleanups):
try:
cleanup()
except Exception:
pass # best-effort: a failed cleanup must not block the rest

View file

@ -0,0 +1,84 @@
# LLM Translation Test Coverage Matrix
Scope: the proxy's two translation surfaces, end to end against a live proxy.
1. **Passthrough** - the client speaks the provider's NATIVE API (Gemini
`generateContent`, Anthropic `/v1/messages`); the proxy forwards it and still
logs a costed `SpendLogs` row (`call_type="pass_through_endpoint"`). Routes:
`/gemini`, `/anthropic`, `/vertex_ai`, `/openai`, `/bedrock`, `/cohere`,
`/mistral`, `/vllm`.
2. **Non-passthrough** - the client speaks OpenAI format
(`/chat/completions`, `/embeddings`); litellm translates to/from the provider.
The two axes that must work in production for each: **passthrough vs
non-passthrough** and **streaming vs non-streaming**, with **cost logged** and
**tool calls** working in every cell.
Companion: live suite `test_passthrough_e2e.py` (this directory). The
non-passthrough chat/embedding cells are exercised by `../spend_tracking/`.
Levels: `live` real provider + proxy + SpendLogs row; `unit` mocked.
Status: `covered` / `partial` / `gap`.
---
## Passthrough endpoints (native provider format)
| Provider | Non-streaming | Streaming | Tool calls | Cost logged | Status |
|----------|---------------|-----------|------------|-------------|--------|
| Gemini (`/gemini/v1beta/models/{m}:generateContent` / `:streamGenerateContent`) | live | live | live | live | **covered** |
| Anthropic (`/anthropic/v1/messages`) | live | live | live | live | **covered** |
| Vertex AI (`/vertex_ai/...`) | - | - | - | - | gap (gcloud auth) |
| OpenAI / Bedrock / Cohere / Mistral / VLLM | - | - | - | - | gap |
Each covered cell asserts: `call_type == "pass_through_endpoint"`, `spend > 0`,
`status == "success"`, correct `custom_llm_provider`/`model`, row correlated by the
`x-litellm-call-id` header. Gemini non-streaming also pins `request_tags`
propagation; streaming pins `chunks > 0` then a costed row; tool tests assert the
provider emitted a tool call (`functionCall` / `tool_use`) and it was costed.
Cost on passthrough is computed in the success handler by transforming the native
response to a `ModelResponse` and calling `litellm.completion_cost()`; for
streaming, chunks are buffered and costed after the stream ends. This is the path
most likely to silently break and the one a mock can't prove works.
## Non-passthrough endpoints (OpenAI-compatible translation)
| Modality | Non-streaming | Streaming | Tool calls | Cost logged | Status |
|----------|---------------|-----------|------------|-------------|--------|
| Chat | live (spend suite) | live (spend suite) | gap | live | partial |
| Embeddings | live (spend suite) | n/a | n/a | live | covered |
| Responses / image / audio / rerank / realtime | - | - | - | - | gap |
## This suite's files
| Test | Cell |
|------|------|
| `test_gemini_passthrough_nonstreaming_logs_cost` | gemini native, non-stream, cost + tags |
| `test_gemini_passthrough_streaming_logs_cost` | gemini native, stream, cost |
| `test_gemini_passthrough_tool_call_logs_cost` | gemini native, tool call, cost |
| `test_anthropic_passthrough_nonstreaming_logs_cost` | anthropic native, non-stream, cost |
| `test_anthropic_passthrough_streaming_logs_cost` | anthropic native, stream, cost |
| `test_anthropic_passthrough_tool_call_logs_cost` | anthropic native, tool call, cost |
## Gaps
- Vertex / OpenAI / Bedrock / Cohere passthrough (same shape; add once the
provider credential is configured; Vertex is closest - route exists, auth stale).
- Non-passthrough tool calls over `/chat/completions` end to end with cost.
- Image / audio / rerank / responses / realtime translation + cost.
- Streaming cost-injection (`include_cost_in_streaming_usage`); passthrough on
client disconnect (partial-usage logging).
## Adding a provider/modality
Extend `PassthroughClient` with the native call (it inherits keys, cleanup, and
SpendLogs polling from `ProxyClient`), then add a test that calls it,
`require_successful_call(result)`, and `_costed_row(...)`.
## Timing
Passthrough spend is logged asynchronously after the response and lands on the
`proxy_batch_write_at` (~60s) cycle, so cost assertions poll
`/spend/logs?request_id=<x-litellm-call-id>` to a deadline. Streaming cost is only
known after the stream is fully consumed.

View file

@ -0,0 +1,15 @@
"""LLM-translation suite's `client` fixture.
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
live in the parent tests/e2e/conftest.py. PassthroughClient holds the shared
Gateway, so the `resources` fixture cleans up keys this suite creates.
"""
import pytest
from passthrough_client import PassthroughClient, build_client
@pytest.fixture(scope="session")
def client() -> PassthroughClient:
return build_client()

View file

@ -0,0 +1,163 @@
"""Client for LLM-translation e2e tests over the proxy's passthrough endpoints.
A passthrough request is sent in the PROVIDER's native format (Gemini
generateContent, Anthropic /v1/messages) to the proxy, which forwards it to the
provider and still logs a SpendLogs row (call_type="pass_through_endpoint"). The
litellm virtual key is passed as the provider key; the proxy swaps in the real env
credential. SpendLogs.request_id == the x-litellm-call-id response header. The
native request models are co-located here because only this suite uses them.
"""
from __future__ import annotations
from dataclasses import dataclass
from pydantic import BaseModel, Field
from e2e_gateway import Gateway, build_gateway
from e2e_http import Headers, StreamingResponse
from models import ChatMessage
class JsonSchemaProperty(BaseModel):
type: str
class JsonSchema(BaseModel):
type: str
properties: dict[str, JsonSchemaProperty]
required: list[str]
class GeminiHeaders(Headers):
x_goog_api_key: str = Field(serialization_alias="x-goog-api-key")
content_type: str = Field(
default="application/json", serialization_alias="Content-Type"
)
tags: str | None = None
class AnthropicHeaders(Headers):
x_api_key: str = Field(serialization_alias="x-api-key")
anthropic_version: str = Field(
default="2023-06-01", serialization_alias="anthropic-version"
)
content_type: str = Field(
default="application/json", serialization_alias="Content-Type"
)
tags: str | None = None
class AltSseParams(BaseModel):
alt: str = "sse"
class GeminiPart(BaseModel):
text: str
class GeminiContent(BaseModel):
role: str = "user"
parts: list[GeminiPart]
class GeminiFunctionDeclaration(BaseModel):
name: str
description: str
parameters: JsonSchema
class GeminiTool(BaseModel):
function_declarations: list[GeminiFunctionDeclaration] = Field(
serialization_alias="functionDeclarations"
)
class GeminiGenerateBody(BaseModel):
contents: list[GeminiContent]
tools: list[GeminiTool] | None = None
class AnthropicTool(BaseModel):
name: str
description: str
input_schema: JsonSchema
class AnthropicMessageBody(BaseModel):
model: str
max_tokens: int
messages: list[ChatMessage]
tools: list[AnthropicTool] | None = None
stream: bool = False
def _tags_header(tags: list[str] | None) -> str | None:
return ",".join(tags) if tags else None
@dataclass(frozen=True, slots=True)
class PassthroughClient:
gateway: Gateway
# ---- Gemini native passthrough (/gemini/v1beta/...) -----------------
def gemini_generate(
self,
key: str,
model: str,
text: str,
*,
tools: list[GeminiTool] | None = None,
tags: list[str] | None = None,
) -> StreamingResponse:
return self.gateway.transport.send(
f"/gemini/v1beta/models/{model}:generateContent",
headers=GeminiHeaders(x_goog_api_key=key, tags=_tags_header(tags)),
json=GeminiGenerateBody(
contents=[GeminiContent(parts=[GeminiPart(text=text)])], tools=tools
),
)
def gemini_stream(
self, key: str, model: str, text: str, *, tags: list[str] | None = None
) -> StreamingResponse:
return self.gateway.transport.send(
f"/gemini/v1beta/models/{model}:streamGenerateContent",
headers=GeminiHeaders(x_goog_api_key=key, tags=_tags_header(tags)),
json=GeminiGenerateBody(
contents=[GeminiContent(parts=[GeminiPart(text=text)])]
),
params=AltSseParams(),
stream=True,
)
# ---- Anthropic native passthrough (/anthropic/v1/messages) ----------
def anthropic_message(
self,
key: str,
model: str,
text: str,
*,
max_tokens: int = 64,
tools: list[AnthropicTool] | None = None,
stream: bool = False,
tags: list[str] | None = None,
) -> StreamingResponse:
return self.gateway.transport.send(
"/anthropic/v1/messages",
headers=AnthropicHeaders(x_api_key=key, tags=_tags_header(tags)),
json=AnthropicMessageBody(
model=model,
max_tokens=max_tokens,
messages=[ChatMessage(role="user", content=text)],
tools=tools,
stream=stream,
),
stream=stream,
)
def build_client() -> PassthroughClient:
return PassthroughClient(gateway=build_gateway())

View file

@ -0,0 +1,219 @@
"""Live e2e: a model's custom per-token pricing is loaded, billed, and isolated.
The gateway config declares ``custom-priced-flash`` (gemini-2.5-flash underneath)
with input/output rates deliberately far above the canonical gemini price, read
back here from the same config file. Three behaviors are checked independently:
- billing: a real call's logged cost breakdown charges input and output tokens at
the custom rates, each component checked separately (a base-rate bill lands
~100x lower; a swapped input/output rate passes a total-only check but not this)
- reporting: /model/info surfaces those rates for the model
- isolation: gemini-2.5-flash shares the same underlying gemini/gemini-2.5-flash
but sets no override, so it must keep its own price; an override that leaks into
the shared cost map misprices it. This fails on a real proxy gap today, so it is
marked xfail(strict=True): the suite stays green while the leak persists and
flips to a failure the moment isolation is fixed and the marker should be removed.
"""
import time
from dataclasses import dataclass
from pathlib import Path
import pytest
import yaml
from pydantic import BaseModel, RootModel
from e2e_config import unique_marker
from e2e_http import Success, unwrap
from models import ChatBody, ChatMessage, CustomPricing, ModelInfoEntry, SpendLogsParams
from passthrough_client import PassthroughClient
pytestmark = pytest.mark.e2e
CUSTOM_MODEL = "custom-priced-flash"
BASE_MODEL = "gemini-2.5-flash"
CONFIG_PATH = Path(__file__).resolve().parents[1] / "gateway" / "litellm-config.yml"
@dataclass(frozen=True, slots=True)
class _Rates:
input_per_token: float
output_per_token: float
class _ConfiguredModel(BaseModel):
model_name: str
litellm_params: CustomPricing
class _GatewayConfig(BaseModel):
model_list: list[_ConfiguredModel]
class _CostBreakdown(BaseModel):
input_cost: float | None = None
output_cost: float | None = None
class _RowMetadata(BaseModel):
cost_breakdown: _CostBreakdown | None = None
class _SpendRow(BaseModel):
request_id: str | None = None
prompt_tokens: int | None = None
completion_tokens: int | None = None
metadata: _RowMetadata | None = None
class _SpendRows(RootModel[list[_SpendRow]]):
pass
def _approx_equal(actual: float, expected: float) -> bool:
"""Within 1% or 1e-9 absolute - spend math, not exact float identity."""
return abs(actual - expected) <= max(1e-9, abs(expected) * 1e-2)
def _configured_pricing(model_name: str) -> _Rates:
"""The custom rates declared for `model_name` in the gateway config the proxy
runs with - the source of truth the billed and reported prices are checked
against."""
config = _GatewayConfig.model_validate(yaml.safe_load(CONFIG_PATH.read_text()))
for entry in config.model_list:
if entry.model_name == model_name:
pricing = entry.litellm_params
assert pricing.input_cost_per_token and pricing.output_cost_per_token, (
f"{model_name} declares no custom per-token rates in {CONFIG_PATH.name}"
)
return _Rates(pricing.input_cost_per_token, pricing.output_cost_per_token)
pytest.fail(f"{model_name} not found in {CONFIG_PATH.name}")
def _model_info_entry(
entries: list[ModelInfoEntry], model_name: str
) -> ModelInfoEntry:
for entry in entries:
if entry.model_name == model_name:
return entry
pytest.fail(f"{model_name} absent from /model/info; the override did not load")
def _poll_breakdown_row(
client: PassthroughClient, key: str, response_id: str | None
) -> _SpendRow:
"""Poll /spend/logs until the call's row lands with a cost breakdown (rows
flush ~60s behind the call via proxy_batch_write_at)."""
deadline = time.monotonic() + client.gateway.poll_timeout
while time.monotonic() < deadline:
result = client.gateway.transport.get(
"/spend/logs",
headers=client.gateway.transport.master,
params=SpendLogsParams(api_key=key),
response_type=_SpendRows,
)
match result:
case Success(data=data):
rows = data.root
case _:
rows = []
priced = [
row
for row in rows
if row.metadata
and row.metadata.cost_breakdown
and row.metadata.cost_breakdown.input_cost is not None
]
for row in priced:
if response_id and row.request_id == response_id:
return row
if priced and response_id is None:
return priced[0]
time.sleep(client.gateway.poll_interval)
pytest.fail("no spend row with a cost breakdown landed before the deadline")
def test_custom_pricing_is_billed_at_configured_rate(
client: PassthroughClient, scoped_key: str
) -> None:
rates = _configured_pricing(CUSTOM_MODEL)
chat = unwrap(
client.gateway.chat(
scoped_key,
ChatBody(
model=CUSTOM_MODEL,
messages=[
ChatMessage(
role="user", content=f"reply with one word {unique_marker()}"
)
],
max_tokens=16,
),
)
)
row = _poll_breakdown_row(client, scoped_key, chat.id)
assert row.metadata and row.metadata.cost_breakdown # guaranteed by the poll
breakdown = row.metadata.cost_breakdown
prompt = row.prompt_tokens or 0
completion = row.completion_tokens or 0
assert prompt > 0 and completion > 0, f"call tokens not logged on the row: {row}"
input_cost = breakdown.input_cost
output_cost = breakdown.output_cost
assert input_cost is not None and output_cost is not None, (
f"row cost breakdown missing input/output cost: {breakdown}"
)
assert _approx_equal(input_cost, prompt * rates.input_per_token), (
f"input_cost {input_cost} != {prompt} tokens * {rates.input_per_token} "
f"= {prompt * rates.input_per_token}"
)
assert _approx_equal(output_cost, completion * rates.output_per_token), (
f"output_cost {output_cost} != {completion} tokens * {rates.output_per_token} "
f"= {completion * rates.output_per_token}"
)
def test_model_info_reports_custom_pricing(client: PassthroughClient) -> None:
rates = _configured_pricing(CUSTOM_MODEL)
entry = _model_info_entry(client.gateway.model_info(), CUSTOM_MODEL)
assert entry.litellm_params.input_cost_per_token == rates.input_per_token, (
f"/model/info litellm_params input rate "
f"{entry.litellm_params.input_cost_per_token} != configured "
f"{rates.input_per_token}"
)
assert entry.litellm_params.output_cost_per_token == rates.output_per_token, (
f"/model/info litellm_params output rate "
f"{entry.litellm_params.output_cost_per_token} != configured "
f"{rates.output_per_token}"
)
@pytest.mark.xfail(
strict=True,
reason="known proxy bug: a deployment's custom per-token pricing leaks into the "
"shared cost map for sibling deployments of the same underlying model; remove "
"this marker once isolation is fixed",
)
def test_custom_pricing_is_isolated_from_sibling_deployment(
client: PassthroughClient,
) -> None:
entries = {entry.model_name: entry for entry in client.gateway.model_info()}
custom = entries.get(CUSTOM_MODEL)
base = entries.get(BASE_MODEL)
assert custom is not None, f"{CUSTOM_MODEL} absent from /model/info"
assert base is not None, f"{BASE_MODEL} absent from /model/info"
# custom-priced-flash overrides pricing; gemini-2.5-flash shares the same
# underlying gemini/gemini-2.5-flash but sets no override, so it must keep its
# own price. Equal rates mean the override leaked into the shared cost map.
assert (
base.model_info.input_cost_per_token != custom.model_info.input_cost_per_token
), (
f"{BASE_MODEL} input rate {base.model_info.input_cost_per_token} matches "
f"{CUSTOM_MODEL}'s override {custom.model_info.input_cost_per_token}; "
f"per-deployment custom pricing is not isolated"
)

View file

@ -0,0 +1,159 @@
"""Live e2e for LLM-translation passthrough endpoints.
Each test sends a NATIVE provider request through the proxy's passthrough route
and verifies the proxy still logged a costed SpendLogs row
(call_type="pass_through_endpoint"), correlated by the x-litellm-call-id header.
Covered: gemini ("gemini-2.5-flash") + anthropic ("claude-haiku-4-5"), streaming +
non-streaming, plus native tool calls. See LLM_TRANSLATION_COVERAGE_MATRIX.md.
A passthrough call returning non-2xx fails hard (never a skip); once it returns
2xx, a missing or zero-cost SpendLogs row fails too.
"""
import pytest
from e2e_config import unique_marker
from e2e_http import StreamingResponse, require_successful_call
from models import SpendLogRow
from passthrough_client import (
AnthropicTool,
GeminiFunctionDeclaration,
GeminiTool,
JsonSchema,
JsonSchemaProperty,
PassthroughClient,
)
pytestmark = pytest.mark.e2e
def _fetch_cost_breakdown(client: PassthroughClient, result: StreamingResponse) -> SpendLogRow:
"""The passthrough call's logged row, polled until it carries a cost.
Asserts (not skips) that a 2xx passthrough call produced a costed row - the
whole point of passthrough spend tracking.
"""
assert result.call_id, "passthrough response had no x-litellm-call-id header"
rows = client.gateway.poll_logs_for_request_id(
result.call_id,
predicate=lambda rs: (rs[0].spend or 0) > 0,
)
assert rows, f"no SpendLogs row for passthrough call_id {result.call_id}"
row = rows[0]
assert row.call_type == "pass_through_endpoint"
assert (row.spend or 0) > 0, f"passthrough call was not costed: {row}"
assert row.status == "success"
return row
# ---- Gemini passthrough ------------------------------------------------
def test_gemini_passthrough_nonstreaming_logs_cost(
client: PassthroughClient, scoped_key: str
) -> None:
tag = f"e2e-passthrough-{unique_marker()}"
result = client.gemini_generate(
scoped_key, "gemini-2.5-flash", "Say hello in one word", tags=[tag, "gemini"]
)
require_successful_call(result)
row = _fetch_cost_breakdown(client, result)
assert row.custom_llm_provider == "gemini"
assert "gemini" in (row.model or "")
assert tag in (row.request_tags or []), f"tags not logged: {row.request_tags}"
def test_gemini_passthrough_streaming_logs_cost(
client: PassthroughClient, scoped_key: str
) -> None:
result = client.gemini_stream(scoped_key, "gemini-2.5-flash", "Count to five")
require_successful_call(result)
assert result.chunks > 0, "streaming passthrough produced no events"
row = _fetch_cost_breakdown(client, result)
assert row.custom_llm_provider == "gemini"
def test_gemini_passthrough_tool_call_logs_cost(
client: PassthroughClient, scoped_key: str
) -> None:
result = client.gemini_generate(
scoped_key,
"gemini-2.5-flash",
"What is the weather in Paris? Use the get_weather tool.",
tools=[
GeminiTool(
function_declarations=[
GeminiFunctionDeclaration(
name="get_weather",
description="Get the weather for a city",
parameters=JsonSchema(
type="object",
properties={"city": JsonSchemaProperty(type="string")},
required=["city"],
),
)
]
)
],
)
require_successful_call(result)
assert "functionCall" in result.body, "gemini did not emit a tool call"
row = _fetch_cost_breakdown(client, result)
assert row.custom_llm_provider == "gemini"
# ---- Anthropic passthrough ---------------------------------------------
def test_anthropic_passthrough_nonstreaming_logs_cost(
client: PassthroughClient, scoped_key: str
) -> None:
result = client.anthropic_message(scoped_key, "claude-haiku-4-5", "Say hello")
require_successful_call(result)
row = _fetch_cost_breakdown(client, result)
assert row.custom_llm_provider == "anthropic"
assert "claude" in (row.model or "")
def test_anthropic_passthrough_streaming_logs_cost(
client: PassthroughClient, scoped_key: str
) -> None:
result = client.anthropic_message(
scoped_key, "claude-haiku-4-5", "Count to five", stream=True
)
require_successful_call(result)
assert result.chunks > 0, "streaming passthrough produced no events"
row = _fetch_cost_breakdown(client, result)
assert row.custom_llm_provider == "anthropic"
def test_anthropic_passthrough_tool_call_logs_cost(
client: PassthroughClient, scoped_key: str
) -> None:
result = client.anthropic_message(
scoped_key,
"claude-haiku-4-5",
"What is the weather in Paris? Use the get_weather tool.",
tools=[
AnthropicTool(
name="get_weather",
description="Get the weather for a city",
input_schema=JsonSchema(
type="object",
properties={"city": JsonSchemaProperty(type="string")},
required=["city"],
),
)
],
)
require_successful_call(result)
assert "tool_use" in result.body, "anthropic did not emit a tool call"
row = _fetch_cost_breakdown(client, result)
assert row.custom_llm_provider == "anthropic"

240
tests/e2e/models.py Normal file
View file

@ -0,0 +1,240 @@
"""Shared pydantic request/response models for the e2e gateway.
Only the fields the tests read are modelled; pydantic ignores the rest, so a
response validates without mirroring every proxy field. No untyped dicts.
"""
from __future__ import annotations
from pydantic import BaseModel, ConfigDict, RootModel
# ---------- keys ----------
class ModelBudgetEntry(BaseModel):
budget_limit: float
time_period: str
class BudgetWindow(BaseModel):
budget_duration: str
max_budget: float
class KeyGenerateBody(BaseModel):
models: list[str] = []
duration: str | None = None
max_budget: float | None = None
soft_budget: float | None = None
budget_duration: str | None = None
user_id: str | None = None
team_id: str | None = None
budget_id: str | None = None
model_max_budget: dict[str, ModelBudgetEntry] | None = None
budget_limits: list[BudgetWindow] | None = None
tpm_limit: int | None = None
rpm_limit: int | None = None
class KeyGenerateResponse(BaseModel):
key: str
class KeyDeleteBody(BaseModel):
keys: list[str]
class KeyInfoParams(BaseModel):
key: str
class LiteLLMBudgetTable(BaseModel):
max_budget: float | None = None
soft_budget: float | None = None
budget_duration: str | None = None
budget_reset_at: str | None = None
class KeyInfo(BaseModel):
spend: float | None = None
max_budget: float | None = None
budget_reset_at: str | None = None
budget_id: str | None = None
litellm_budget_table: LiteLLMBudgetTable | None = None
class KeyInfoResponse(BaseModel):
info: KeyInfo
# ---------- customers ----------
class CustomerDeleteBody(BaseModel):
user_ids: list[str]
# ---------- chat / embeddings ----------
class ChatMetadata(BaseModel):
tags: list[str] | None = None
class ChatMessage(BaseModel):
role: str
content: str
class ChatBody(BaseModel):
model: str
messages: list[ChatMessage]
stream: bool = False
max_tokens: int | None = None
user: str | None = None
metadata: ChatMetadata | None = None
class OutMessage(BaseModel):
content: str | None = None
class ChatChoice(BaseModel):
message: OutMessage | None = None
class Usage(BaseModel):
prompt_tokens: int | None = None
completion_tokens: int | None = None
total_tokens: int | None = None
class ChatResponse(BaseModel):
id: str | None = None
model: str | None = None
choices: list[ChatChoice] = []
usage: Usage | None = None
class EmbedBody(BaseModel):
model: str
input: str
class EmbedResponse(BaseModel):
model: str | None = None
# ---------- spend logs ----------
class SpendLogRow(BaseModel):
request_id: str | None = None
model: str | None = None
spend: float | None = None
status: str | None = None
cache_hit: str | None = None
call_type: str | None = None
custom_llm_provider: str | None = None
team_id: str | None = None
user: str | None = None
end_user: str | None = None
prompt_tokens: int | None = None
completion_tokens: int | None = None
total_tokens: int | None = None
request_tags: list[str] | None = None
class SpendLogs(RootModel[list[SpendLogRow]]):
pass
class SpendLogsParams(BaseModel):
request_id: str | None = None
api_key: str | None = None
# ---------- spend calculate ----------
class SpendCalculateBody(BaseModel):
model: str
messages: list[ChatMessage]
class SpendCalculateResponse(BaseModel):
cost: float
# ---------- route probing ----------
class DateRangeParams(BaseModel):
start_date: str
end_date: str
class RouteSpec(RootModel[dict[str, object]]):
"""One /openapi.json path entry: a map of HTTP method -> operation. Only the
method names are read, so the operation specs stay opaque."""
@property
def methods(self) -> frozenset[str]:
return frozenset(method.lower() for method in self.root)
class OpenAPISchema(BaseModel):
paths: dict[str, RouteSpec] = {}
# ---------- model info / custom pricing ----------
class CustomPricing(BaseModel):
"""The per-token custom-pricing fields a deployment can override in
litellm_params - the token-cost subset of litellm's CustomPricingLiteLLMParams
the proxy applies to chat spend. All optional: a config sets only what it
overrides, and /model/info echoes the rates the proxy resolved."""
model_config = ConfigDict(extra="ignore")
input_cost_per_token: float | None = None
output_cost_per_token: float | None = None
cache_read_input_token_cost: float | None = None
cache_creation_input_token_cost: float | None = None
def overrides(self) -> dict[str, float]:
"""The rates actually declared (non-null) - e.g. those a config.yml sets."""
declared = {
"input_cost_per_token": self.input_cost_per_token,
"output_cost_per_token": self.output_cost_per_token,
"cache_read_input_token_cost": self.cache_read_input_token_cost,
"cache_creation_input_token_cost": self.cache_creation_input_token_cost,
}
return {field: rate for field, rate in declared.items() if rate is not None}
def token_cost(self, prompt_tokens: int, completion_tokens: int) -> float:
"""Spend for a fresh (uncached) call under these rates: the proxy's
custom-pricing formula (prompt * input + completion * output)."""
assert (
self.input_cost_per_token is not None
and self.output_cost_per_token is not None
), "custom pricing has no per-token rates"
return (
prompt_tokens * self.input_cost_per_token
+ completion_tokens * self.output_cost_per_token
)
class ModelInfoEntry(BaseModel):
"""One /model/info row. `litellm_params` is the configured deployment (carries
any custom-pricing override); `model_info` is the price the proxy resolved for
it - the override merged over the cost-map defaults."""
model_config = ConfigDict(protected_namespaces=())
model_name: str
litellm_params: CustomPricing = CustomPricing()
model_info: CustomPricing = CustomPricing()
class ModelInfoResponse(BaseModel):
data: list[ModelInfoEntry] = []

7
tests/e2e/pytest.ini Normal file
View file

@ -0,0 +1,7 @@
[pytest]
# Config when any e2e suite under tests/e2e/ is run directly, e.g.
# uv run pytest tests/e2e/spend_tracking/ -v
# The e2e marker is also registered in conftest.py for runs rooted elsewhere.
addopts = --strict-markers --strict-config
markers =
e2e: live test that requires a running proxy and real provider keys

View file

@ -0,0 +1,78 @@
# Spend Tracking Test Coverage Matrix
Scope: every distinct spend-tracking code path, mapped to the test that exercises
it and the level it runs at. Highlights where a live e2e check is the only thing
that would catch a regression.
Companion: live suite `test_spend_tracking_e2e.py` + route breadth
`test_spend_routes.py` (this directory). Offline regression suite:
`tests/test_litellm/proxy/spend_tracking/`. Reference PR: BerriAI/litellm#29956.
Levels: `unit` mocked; `integration` real DB/cost-map; `live` real provider +
proxy + SpendLogs rows. Status: `covered` / `partial` / `gap`.
---
## SpendLogs row construction (`spend_tracking_utils.get_logging_payload`)
| Path | Existing | Level | Status | Live e2e |
|------|----------|-------|--------|----------|
| `_get_status_for_spend_log` | `test_spend_tracking_utils.py` | unit | covered | yes (status read off the row) |
| cache-hit `request_id` suffix | `test_spend_tracking_utils.py` | unit | covered | yes (`test_cache_hit_is_zero_cost_and_suffixed`) |
| failure status + zero spend | `test_spend_tracking_utils.py` | unit | covered | no (live failure logging is non-deterministic across providers) |
| per-model / per-provider attribution | `test_spend_tracking_utils.py` | unit | covered | yes (`test_each_model_on_a_shared_key_gets_its_own_row`) |
| field population (model/tokens/api_key/team/org) | `test_spend_tracking_utils.py` | unit | partial | yes (asserts real values) |
| `request_tags` propagation | `test_db_spend_update_writer.py` | unit | partial | yes (`test_request_tags_round_trip`) |
| `end_user` attribution | unit | unit | partial | yes (`test_end_user_spend_attributed_on_row`) |
## Cost calculation by modality
| Modality | Existing | Status | Live e2e |
|----------|----------|--------|----------|
| Chat (non-stream) | `test_cost_calculator.py`, `local_testing/test_completion_cost.py` | covered | yes (`test_chat_completion_writes_nonzero_spend_row`) |
| Chat (streaming) | `test_streaming_interrupt_spend_tracking.py` | partial | yes (`test_streaming_chat_completion_tracks_spend`) |
| Embedding | `test_cost_calculator.py` (#29956) | partial | yes (`test_embedding_writes_nonzero_spend_row`) |
| Pass-through (gemini/anthropic) | `pass_through_tests/*.test.js` + `llm_translation/` suite | covered | yes (llm_translation suite) |
| Image / audio / rerank / responses / realtime | per-provider unit cost tests | partial/gap | gap |
## Entity spend aggregation
| Entity | Existing | Status | Live e2e |
|--------|----------|--------|----------|
| API key | `test_db_spend_update_writer.py`, `test_spend_counters.py` | covered | yes (`test_key_spend_equals_sum_of_logs`) |
| Tag | `test_update_daily_tag_spend.py` | partial | yes (`test_request_tags_round_trip`, propagation only) |
| End-user | `test_proxy_update_spend.py` | covered | yes |
| Spend == sum(logs) consistency | none | gap | yes (key aggregate == sum of rows) |
## Spend read endpoints (verification surface)
| Endpoint | Existing | Status | Live e2e |
|----------|----------|--------|----------|
| `/spend/logs` (request_id / api_key) | `test_spend_management_endpoints.py` | covered | yes (primary read path; `test_spend_logs_endpoint_returns_spend` asserts 200 + spend, never 5xx) |
| `/spend/calculate` | `local_testing/test_spend_calculate_endpoint.py` | covered | yes (`test_spend_calculate_returns_nonzero_cost`) |
| `/spend/tags` | `test_spend_management_endpoints.py` | partial | yes (`test_spend_routes.py` route probe) |
| whole spend GET surface (22 routes) | unit per-handler | partial | yes (`test_spend_routes.py` probes each for 404/5xx) |
## What this suite pins
| Test | Invariant |
|------|-----------|
| `test_chat_completion_writes_nonzero_spend_row` | nonzero cost, token arithmetic, status, row findable by `response.id` |
| `test_streaming_chat_completion_tracks_spend` | streamed responses still costed |
| `test_embedding_writes_nonzero_spend_row` | embedding cost, `completion_tokens == 0` |
| `test_cache_hit_is_zero_cost_and_suffixed` | cache hits not double-charged; `_cache_hit` suffix |
| `test_key_spend_equals_sum_of_logs` | key aggregate == sum of rows |
| `test_request_tags_round_trip` | tags persist onto the row |
| `test_end_user_spend_attributed_on_row` | `end_user` attributed + costed |
| `test_each_model_on_a_shared_key_gets_its_own_row` | per-model/provider rows, correct model + cost, distinct request_ids matching response id |
| `test_spend_calculate_returns_nonzero_cost` | cost-map smoke (no batch wait) |
| `test_spend_logs_endpoint_returns_spend` | `/spend/logs` returns 200 + the key's spend, never a 5xx (intermittent-500 regression) |
| `test_spend_routes.py` (23) | no spend route 404s or 5xxs |
## Design + timing
`proxy_batch_write_at` (~60s) means rows land late; every read polls to a deadline.
Fresh scoped key per test (isolation, xdist-safe, cleaned up). Assert invariants
(`spend > 0`, `total == prompt + completion`, aggregate == sum), not literal
$/token values, so pricing drift is not a failure. Skip on environment (no proxy /
no provider key), fail on behavior (a real 2xx call with a wrong/missing row).

View file

@ -0,0 +1,16 @@
"""Spend-tracking suite's `client` fixture.
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
live in the parent tests/e2e/conftest.py. SpendClient exposes the shared Gateway
(GatewayProvider), so the `resources` fixture cleans up keys and customers this
suite creates.
"""
import pytest
from spend_e2e_client import SpendClient, build_client
@pytest.fixture(scope="session")
def client() -> SpendClient:
return build_client()

View file

@ -0,0 +1,167 @@
"""Spend-tracking e2e client: a Gateway plus the spend-specific read endpoints.
Generic proxy operations (keys, customers, chat/embed, route probing, SpendLogs
polling) come from the shared Gateway, DI'd in (composition, not inheritance).
This client adds only the spend surface: /spend/calculate, key-spend
polling, and the route probes the breadth test uses.
Re-exports unwrap / is_ok / unique_marker / SpendLogRow so the tests import their
helpers from one place.
"""
from __future__ import annotations
import os
import time
from collections.abc import Callable
from dataclasses import dataclass
from e2e_config import unique_marker
from e2e_http import (
NoBody,
ProbeResult,
Result,
StreamingResponse,
is_ok,
unwrap,
)
from e2e_gateway import Gateway, build_gateway
from models import (
ChatBody,
ChatMessage,
ChatMetadata,
ChatResponse,
DateRangeParams,
EmbedBody,
EmbedResponse,
OpenAPISchema,
SpendCalculateBody,
SpendCalculateResponse,
SpendLogRow,
)
__all__ = [
"SpendClient",
"build_client",
"reset_spend_logs",
"unique_marker",
"unwrap",
"is_ok",
"SpendLogRow",
"ProbeResult",
]
def reset_spend_logs() -> None:
"""Truncate LiteLLM_SpendLogs for a clean slate. No proxy endpoint deletes
spend logs (/global/spend/reset keeps them), so go to the DB directly. Uses
DATABASE_URL (default: the local docker postgres on its mapped host port; note
the in-container `@db` host isn't resolvable from the host, so default to
localhost).
"""
import psycopg
url = os.environ.get(
"DATABASE_URL",
"postgresql://llmproxy:dbpassword9090@localhost:5432/litellm",
)
with psycopg.connect(url) as conn:
_ = conn.execute('TRUNCATE TABLE "LiteLLM_SpendLogs"')
def _chat_body(
model: str,
content: str,
*,
max_tokens: int | None = None,
tags: list[str] | None = None,
user: str | None = None,
stream: bool = False,
) -> ChatBody:
return ChatBody(
model=model,
messages=[ChatMessage(role="user", content=content)],
max_tokens=max_tokens,
stream=stream,
user=user,
metadata=ChatMetadata(tags=tags) if tags else None,
)
@dataclass(frozen=True, slots=True)
class SpendClient:
gateway: Gateway
def chat(
self,
key: str,
model: str,
content: str,
*,
max_tokens: int | None = None,
tags: list[str] | None = None,
user: str | None = None,
) -> Result[ChatResponse]:
return self.gateway.chat(
key, _chat_body(model, content, max_tokens=max_tokens, tags=tags, user=user)
)
def chat_stream(
self, key: str, model: str, content: str, *, max_tokens: int | None = None
) -> StreamingResponse:
return self.gateway.chat_stream(
key, _chat_body(model, content, max_tokens=max_tokens, stream=True)
)
def embed(self, key: str, model: str, content: str) -> Result[EmbedResponse]:
return self.gateway.embed(key, EmbedBody(model=model, input=content))
def poll_logs_for_key(
self,
key: str,
*,
min_rows: int = 1,
predicate: Callable[[list[SpendLogRow]], bool] | None = None,
) -> list[SpendLogRow]:
return self.gateway.poll_logs_for_key(
key, min_rows=min_rows, predicate=predicate
)
def calculate_spend(self, model: str, content: str) -> float:
return unwrap(
self.gateway.transport.post(
"/spend/calculate",
headers=self.gateway.transport.master,
json=SpendCalculateBody(
model=model, messages=[ChatMessage(role="user", content=content)]
),
response_type=SpendCalculateResponse,
)
).cost
def poll_key_spend(self, key: str, *, minimum: float = 0.0) -> float:
deadline = time.monotonic() + self.gateway.poll_timeout
spend = 0.0
while time.monotonic() < deadline:
spend = self.gateway.key_info(key).spend or 0.0
if spend > minimum:
return spend
time.sleep(self.gateway.poll_interval)
return spend
def probe(self, path: str, *, params: DateRangeParams) -> ProbeResult:
return self.gateway.transport.probe(path, params=params)
def openapi(self) -> OpenAPISchema:
return unwrap(
self.gateway.transport.get(
"/openapi.json",
headers=self.gateway.transport.master,
params=NoBody(),
response_type=OpenAPISchema,
)
)
def build_client() -> SpendClient:
return SpendClient(gateway=build_gateway())

View file

@ -0,0 +1,96 @@
"""Breadth check: query every route on the spend read surface and show what it
returns.
Spend tracking sprawls across many routes (model-cost / key / user / team / org /
customer aggregation, tags, and activity reports). Most are served with
`include_in_schema=False`, so they do NOT appear in `/openapi.json` - discovery
from the schema alone misses ~70% of the surface. So we probe a curated, verified
list directly, plus any spend route the schema does list (to auto-catch new ones).
Each probe captures status AND body, so a failure shows the proxy's actual error
(a 500 traceback, a 404 meaning the route was removed) rather than a bare code.
Run with `-rA` (or `-s`) to print every route's response, not just failures.
Healthy == route exists (not 404) and handler did not crash (not 5xx). A 4xx
(missing params / auth nuance) still means the route is wired and ran. Cheap and
fast: no batch-write wait, no provider calls.
"""
from datetime import datetime, timedelta, timezone
import pytest
from models import DateRangeParams
from spend_e2e_client import SpendClient
pytestmark = pytest.mark.e2e
# Verified present and responsive on a live proxy. One per row of the spend
# surface: key / user / team / org / customer aggregation, model-cost, tags,
# activity.
SPEND_ROUTES = (
"/spend/keys",
"/spend/users",
"/spend/tags",
"/spend/logs",
"/spend/logs/ui",
"/global/spend",
"/global/spend/keys",
"/global/spend/teams",
"/global/spend/models",
"/global/spend/provider",
"/global/spend/report",
"/global/spend/tags",
"/global/spend/logs",
"/global/spend/all_tag_names",
"/global/activity",
"/global/activity/model",
"/global/activity/exceptions",
"/key/list",
"/user/list",
"/team/list",
"/organization/list",
"/customer/list",
)
_SPEND_PREFIXES = ("/spend", "/global/spend", "/global/activity")
def _date_range() -> DateRangeParams:
# Satisfies date-required endpoints (report/activity/provider); ignored elsewhere.
end = datetime.now(timezone.utc).date()
start = end - timedelta(days=1)
return DateRangeParams(start_date=start.isoformat(), end_date=end.isoformat())
@pytest.mark.parametrize("route", SPEND_ROUTES)
def test_spend_route_responsive(client: SpendClient, route: str) -> None:
result = client.probe(route, params=_date_range())
print(f"{route} -> {result.status_code}\n{result.body[:600]}")
assert result.healthy, f"{route} -> {result.status_code}\n{result.body[:600]}"
def test_schema_listed_spend_routes_are_responsive(client: SpendClient) -> None:
"""Probe any spend GET route the schema lists that isn't in SPEND_ROUTES."""
schema = client.openapi()
assert schema.paths, "/openapi.json had no paths"
discovered = [
path
for path, spec in schema.paths.items()
if "get" in spec.methods
and "{" not in path
and any(path.startswith(prefix) for prefix in _SPEND_PREFIXES)
]
extras = [path for path in discovered if path not in SPEND_ROUTES]
params = _date_range()
results = [(path, client.probe(path, params=params)) for path in extras]
for path, result in results:
print(f"{path} -> {result.status_code}")
offenders = [
f"{path} -> {result.status_code}\n{result.body[:600]}"
for path, result in results
if not result.healthy
]
assert not offenders, "non-responsive schema spend routes:\n" + "\n".join(offenders)

View file

@ -0,0 +1,328 @@
"""Live end-to-end spend-tracking tests against a running proxy.
Run against a proxy started with the gateway config. Coverage rationale:
SPEND_TRACKING_COVERAGE_MATRIX.md.
Model names are literals from that config: chat tests hit "gemini-2.5-flash",
embedding tests hit "openai-text-embedding-3-small".
Every test: fresh scoped key (isolation) -> real provider call -> unwrap (hard
fail if the proxy couldn't make a call it should) -> poll /spend/logs to a
deadline (rows land ~60s later via proxy_batch_write_at) -> assert invariants on
the real row (spend, token arithmetic, status, cache).
Assertions target invariants, not literals: a regression in the spend pipeline
fails the test; a pricing or token-count drift does not.
"""
import time
from collections.abc import Callable
import pytest
from e2e_http import Success
from lifecycle import ResourceManager
from models import SpendLogs, SpendLogsParams
from spend_e2e_client import SpendClient, SpendLogRow, unique_marker, unwrap
pytestmark = pytest.mark.e2e
def _approx_equal(actual: float, expected: float) -> bool:
"""Within 1% or 1e-9 absolute - spend math, not exact float identity."""
return abs(actual - expected) <= max(1e-9, abs(expected) * 1e-2)
def _summarize(rows: list[SpendLogRow]) -> list[dict[str, object]]:
fields = {
"request_id",
"model",
"spend",
"status",
"cache_hit",
"prompt_tokens",
"completion_tokens",
"total_tokens",
}
return [row.model_dump(include=fields) for row in rows]
def _require_row(
rows: list[SpendLogRow], predicate: Callable[[SpendLogRow], bool], what: str
) -> SpendLogRow:
matches = [r for r in rows if predicate(r)]
assert matches, (
f"no SpendLogs row {what} after polling; saw {len(rows)} row(s): "
f"{_summarize(rows)}"
)
return matches[0]
def test_chat_completion_writes_nonzero_spend_row(
client: SpendClient, scoped_key: str
) -> None:
chat = unwrap(
client.chat(
scoped_key,
"gemini-2.5-flash",
f"reply with one word {unique_marker()}",
max_tokens=16,
)
)
rows = client.poll_logs_for_key(
scoped_key, predicate=lambda rs: any(r.status == "success" for r in rs)
)
row = _require_row(rows, lambda r: r.status == "success", "for the chat call")
assert (row.spend or 0) > 0, f"chat row should cost > 0: {_summarize(rows)}"
assert row.status == "success"
assert row.cache_hit != "True", "fresh call must not be a cache hit"
assert "gemini-2.5-flash" in (row.model or "")
prompt = row.prompt_tokens or 0
completion = row.completion_tokens or 0
total = row.total_tokens or 0
assert prompt > 0 and completion > 0
assert total == prompt + completion, f"token arithmetic broken: {_summarize(rows)}"
if chat.id:
assert any(
r.request_id == chat.id for r in rows
), f"row request_id != client response.id ({chat.id})"
def test_streaming_chat_completion_tracks_spend(
client: SpendClient, scoped_key: str
) -> None:
result = client.chat_stream(
scoped_key,
"gemini-2.5-flash",
f"count to three {unique_marker()}",
max_tokens=64,
)
assert (
result.ok
), f"stream failed (status {result.status_code}): {result.body[:300]}"
rows = client.poll_logs_for_key(
scoped_key, predicate=lambda rs: any((r.spend or 0) > 0 for r in rs)
)
row = _require_row(
rows, lambda r: (r.spend or 0) > 0, "with nonzero spend for the stream"
)
prompt = row.prompt_tokens or 0
completion = row.completion_tokens or 0
assert (
prompt > 0 and completion > 0
), f"streaming tokens not tracked: {_summarize(rows)}"
assert (row.total_tokens or 0) == prompt + completion
def test_embedding_writes_nonzero_spend_row(
client: SpendClient, scoped_key: str
) -> None:
_ = unwrap(
client.embed(
scoped_key,
"openai-text-embedding-3-small",
f"vectorize this sentence {unique_marker()}",
)
)
rows = client.poll_logs_for_key(
scoped_key, predicate=lambda rs: any((r.spend or 0) > 0 for r in rs)
)
row = _require_row(
rows, lambda r: (r.spend or 0) > 0, "with nonzero spend for the embedding"
)
assert (row.prompt_tokens or 0) > 0
assert (row.completion_tokens or 0) == 0, "embeddings have no completion tokens"
assert "text-embedding-3-small" in (row.model or "")
def test_cache_hit_is_zero_cost_and_suffixed(
client: SpendClient, scoped_key: str
) -> None:
# Unique marker shared by both calls: call 1 is a guaranteed cache MISS (fresh
# content, paid), call 2 repeats the identical request and HITS the cache just
# populated. The marker keeps each run isolated - a fixed prompt would persist
# in the shared response cache across runs and make both calls hit (flaky).
prompt = f"What is the capital of France? Answer in one word. {unique_marker()}"
_ = unwrap(client.chat(scoped_key, "gemini-2.5-flash", prompt, max_tokens=16))
_ = unwrap(client.chat(scoped_key, "gemini-2.5-flash", prompt, max_tokens=16))
rows = client.poll_logs_for_key(
scoped_key, predicate=lambda rs: any(r.cache_hit == "True" for r in rs)
)
cache_rows = [r for r in rows if r.cache_hit == "True"]
if not cache_rows:
pytest.skip(
"no cache-hit row observed; caching may be disabled on this proxy. "
f"rows seen: {_summarize(rows)}"
)
cache_row = cache_rows[0]
assert (
cache_row.spend or 0
) == 0.0, f"cache hit was charged (double-charge regression): {_summarize(rows)}"
assert "_cache_hit" in (cache_row.request_id or ""), (
"cache-hit row missing the _cache_hit request_id suffix; "
"duplicate-key collisions will silently drop rows"
)
paid_rows = [r for r in rows if r.cache_hit != "True"]
assert any(
(r.spend or 0) > 0 for r in paid_rows
), f"the non-cached call should still be charged: {_summarize(rows)}"
def test_key_spend_equals_sum_of_logs(client: SpendClient, scoped_key: str) -> None:
for _ in range(2):
_ = unwrap(
client.chat(
scoped_key,
"gemini-2.5-flash",
f"say hi {unique_marker()}",
max_tokens=16,
)
)
rows = client.poll_logs_for_key(
scoped_key,
min_rows=2,
predicate=lambda rs: sum((r.spend or 0) for r in rs) > 0,
)
assert len(rows) >= 2, f"expected >=2 rows for the key, saw {_summarize(rows)}"
logs_total = sum((r.spend or 0) for r in rows)
assert logs_total > 0
key_spend = client.poll_key_spend(scoped_key, minimum=logs_total * 0.999)
assert _approx_equal(
key_spend, logs_total
), f"key aggregate {key_spend} != sum of logs {logs_total}; rows: {_summarize(rows)}"
def test_request_tags_round_trip(client: SpendClient, scoped_key: str) -> None:
tag = f"e2e-spend-{unique_marker()}"
_ = unwrap(
client.chat(
scoped_key, "gemini-2.5-flash", "tagged request", tags=[tag], max_tokens=16
)
)
rows = client.poll_logs_for_key(
scoped_key, predicate=lambda rs: any(tag in (r.request_tags or []) for r in rs)
)
_require_row(
rows, lambda r: tag in (r.request_tags or []), f"carrying request tag {tag!r}"
)
def test_end_user_spend_attributed_on_row(
client: SpendClient, scoped_key: str, resources: ResourceManager
) -> None:
customer = resources.customer(f"e2e-cust-{unique_marker()}")
_ = unwrap(
client.chat(scoped_key, "gemini-2.5-flash", "hi", user=customer, max_tokens=16)
)
rows = client.poll_logs_for_key(
scoped_key, predicate=lambda rs: any(r.end_user == customer for r in rs)
)
row = _require_row(
rows, lambda r: r.end_user == customer, f"attributed to end_user {customer!r}"
)
assert (row.spend or 0) > 0, f"end-user row should cost > 0: {_summarize(rows)}"
def test_each_model_on_a_shared_key_gets_its_own_row(
client: SpendClient, scoped_key: str
) -> None:
"""One key calling two different models, on two providers, gets one spend row per
call - each carrying its own model and a nonzero cost, under distinct request_ids
that match the call's response id. Pins per-model/per-provider attribution: a
regression that stamps the wrong model on the row, bills a call's cost to the
sibling deployment, or collapses both calls onto one request_id fails here."""
gemini = unwrap(
client.chat(
scoped_key, "gemini-2.5-flash", f"one word {unique_marker()}", max_tokens=16
)
)
claude = unwrap(
client.chat(
scoped_key, "claude-haiku-4-5", f"one word {unique_marker()}", max_tokens=16
)
)
def both_models_costed(rows: list[SpendLogRow]) -> bool:
costed = [r.model or "" for r in rows if (r.spend or 0) > 0]
return any("gemini-2.5-flash" in m for m in costed) and any(
"claude-haiku-4-5" in m for m in costed
)
rows = client.poll_logs_for_key(scoped_key, min_rows=2, predicate=both_models_costed)
gemini_row = _require_row(
rows, lambda r: "gemini-2.5-flash" in (r.model or ""), "for the gemini call"
)
claude_row = _require_row(
rows, lambda r: "claude-haiku-4-5" in (r.model or ""), "for the claude call"
)
assert (gemini_row.spend or 0) > 0, f"gemini row should cost > 0: {_summarize(rows)}"
assert (claude_row.spend or 0) > 0, f"claude row should cost > 0: {_summarize(rows)}"
assert (
gemini_row.request_id != claude_row.request_id
), f"two distinct calls collapsed onto one request_id: {_summarize(rows)}"
if gemini.id:
assert (
gemini_row.request_id == gemini.id
), f"gemini row request_id {gemini_row.request_id} != response id {gemini.id}"
if claude.id:
assert (
claude_row.request_id == claude.id
), f"claude row request_id {claude_row.request_id} != response id {claude.id}"
def test_spend_calculate_returns_nonzero_cost(client: SpendClient) -> None:
cost = client.calculate_spend(
"gemini-2.5-flash", "estimate the cost of this request"
)
assert cost > 0, (
"/spend/calculate returned 0 for gemini-2.5-flash; "
"cost map may be missing this model"
)
def test_spend_logs_endpoint_returns_spend(
client: SpendClient, scoped_key: str
) -> None:
"""The /spend/logs read endpoint returns a 200 carrying the key's spend, never a
5xx. Regression for intermittent 500s (DB query / serialization errors under load)
on this endpoint: every poll asserts a success response, not just a truthy row
list, so a 500 fails loudly instead of being swallowed as 'no rows yet'; the
call's nonzero spend must surface before the deadline."""
unwrap(
client.chat(
scoped_key, "gemini-2.5-flash", f"spend logs {unique_marker()}", max_tokens=16
)
)
gateway = client.gateway
deadline = time.monotonic() + gateway.poll_timeout
while True:
result = gateway.transport.get(
"/spend/logs",
headers=gateway.transport.master,
params=SpendLogsParams(api_key=scoped_key),
response_type=SpendLogs,
)
assert isinstance(result, Success), f"/spend/logs did not return 200 OK: {result}"
rows = result.data.root
if sum((r.spend or 0) for r in rows) > 0:
return
if time.monotonic() >= deadline:
pytest.fail(
f"/spend/logs never surfaced the key's spend before the deadline; "
f"saw {_summarize(rows)}"
)
time.sleep(gateway.poll_interval)

View file

@ -0,0 +1,46 @@
"""Unit coverage for the lifecycle harness (lifecycle.run_case).
Cases register cleanups progressively during init() (create team, then user, then
key), so a failure partway through init() must still release whatever was already
created on the long-lived shared proxy. This guards that contract.
"""
from dataclasses import dataclass, field
from typing import Callable, List
import pytest
from lifecycle import run_case
@dataclass
class _PartialInitCase:
"""init() registers a cleanup, then raises before finishing - mirroring a real
case that creates a resource, registers its delete, then fails on the next
step."""
released: List[str] = field(default_factory=list)
_undo: List[Callable[[], None]] = field(default_factory=list)
def init(self) -> None:
self._undo.append(lambda: self.released.append("first"))
raise RuntimeError("init failed after registering the first resource")
def run(self) -> None:
raise AssertionError("run() must not execute when init() failed")
def teardown(self) -> None:
for undo in reversed(self._undo):
undo()
def test_run_case_releases_resources_when_init_fails_partway() -> None:
case = _PartialInitCase()
with pytest.raises(RuntimeError, match="init failed"):
run_case(case)
assert case.released == ["first"], (
"a resource registered before init() failed must still be released, or it "
"leaks on the long-lived shared proxy"
)

244
tests/e2e/transport.py Normal file
View file

@ -0,0 +1,244 @@
"""Transport: the typed request primitives clients use, behind a Protocol.
`Transport` is what each client depends on (composition + DI); `HttpTransport` is
the concrete frozen-slots dataclass that fulfils it via the e2e_http wrapper. No
client touches requests.* or builds raw dicts; they pass pydantic models here.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Protocol
from pydantic import BaseModel
import e2e_http
from e2e_http import URL, AuthHeaders, ProbeResult, Result, StreamingResponse
class Transport(Protocol):
def post[R: BaseModel](
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
) -> Result[R]: ...
def stream(
self, path: str, *, headers: BaseModel, json: BaseModel
) -> StreamingResponse: ...
def send(
self,
path: str,
*,
headers: BaseModel,
json: BaseModel,
params: BaseModel | None = None,
stream: bool = False,
) -> StreamingResponse: ...
def get[R: BaseModel](
self,
path: str,
*,
headers: BaseModel,
params: BaseModel,
response_type: type[R],
) -> Result[R]: ...
def delete[R: BaseModel](
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
) -> Result[R]: ...
def probe(self, path: str, *, params: BaseModel) -> ProbeResult: ...
def bearer(self, key: str) -> AuthHeaders: ...
@property
def master(self) -> AuthHeaders: ...
@dataclass(frozen=True, slots=True)
class HttpTransport:
base_url: str
master_key: str
request_timeout: float = 60.0
def _url(self, path: str) -> URL:
return URL(f"{self.base_url.rstrip('/')}{path}")
def bearer(self, key: str) -> AuthHeaders:
return AuthHeaders(authorization=f"Bearer {key}")
@property
def master(self) -> AuthHeaders:
return self.bearer(self.master_key)
def post[R: BaseModel](
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
) -> Result[R]:
return e2e_http.post(
self._url(path),
headers=headers,
json=json,
response_type=response_type,
timeout=self.request_timeout,
)
def get[R: BaseModel](
self,
path: str,
*,
headers: BaseModel,
params: BaseModel,
response_type: type[R],
) -> Result[R]:
return e2e_http.get(
self._url(path),
headers=headers,
params=params,
response_type=response_type,
timeout=self.request_timeout,
)
def delete[R: BaseModel](
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
) -> Result[R]:
return e2e_http.delete(
self._url(path),
headers=headers,
json=json,
response_type=response_type,
timeout=self.request_timeout,
)
def stream(
self, path: str, *, headers: BaseModel, json: BaseModel
) -> StreamingResponse:
return e2e_http.stream(
self._url(path), headers=headers, json=json, timeout=self.request_timeout
)
def send(
self,
path: str,
*,
headers: BaseModel,
json: BaseModel,
params: BaseModel | None = None,
stream: bool = False,
) -> StreamingResponse:
return e2e_http.send(
self._url(path),
headers=headers,
json=json,
params=params,
stream=stream,
timeout=self.request_timeout,
)
def probe(self, path: str, *, params: BaseModel) -> ProbeResult:
return e2e_http.probe(
self._url(path),
headers=self.master,
params=params,
timeout=self.request_timeout,
)
# Top-level management/admin route groups. In a split deployment these are served
# by the control plane (a different service from the LLM data plane). LLM routes
# (/chat, /embeddings, and native passthrough like /gemini, /anthropic) are NOT
# here and fall through to the data plane. Matched as path prefixes.
CONTROL_PLANE_PREFIXES: tuple[str, ...] = (
"/key",
"/user",
"/team",
"/organization",
"/customer",
"/tag",
"/budget",
"/model/info",
"/spend",
"/global",
"/openapi.json",
)
def is_control_plane_path(path: str) -> bool:
"""True if `path` is a management/admin route (served by the control plane in a
split deployment), false for LLM data-plane routes."""
return path.startswith(CONTROL_PLANE_PREFIXES)
@dataclass(frozen=True, slots=True)
class SplitTransport:
"""A Transport that dispatches each call by path to one of two backends: the
management/admin control plane or the LLM data plane.
Litellm can run as a split control-plane/data-plane deployment where the two
surfaces live on different services. Clients here stay plane-agnostic — they
keep calling ``transport.post("/budget/new", ...)`` or
``transport.send("/chat/completions", ...)`` — and routing happens in one place
by path (see ``CONTROL_PLANE_PREFIXES``). When ``control`` and ``data`` share a
base URL (the monolithic default), routing is a no-op. ``bearer``/``master``
are plane-agnostic (same master key both planes), so they come from ``data``.
"""
data: HttpTransport
control: HttpTransport
def _route(self, path: str) -> HttpTransport:
return self.control if is_control_plane_path(path) else self.data
def bearer(self, key: str) -> AuthHeaders:
return self.data.bearer(key)
@property
def master(self) -> AuthHeaders:
return self.data.master
def post[R: BaseModel](
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
) -> Result[R]:
return self._route(path).post(
path, headers=headers, json=json, response_type=response_type
)
def get[R: BaseModel](
self,
path: str,
*,
headers: BaseModel,
params: BaseModel,
response_type: type[R],
) -> Result[R]:
return self._route(path).get(
path, headers=headers, params=params, response_type=response_type
)
def delete[R: BaseModel](
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
) -> Result[R]:
return self._route(path).delete(
path, headers=headers, json=json, response_type=response_type
)
def stream(
self, path: str, *, headers: BaseModel, json: BaseModel
) -> StreamingResponse:
return self._route(path).stream(path, headers=headers, json=json)
def send(
self,
path: str,
*,
headers: BaseModel,
json: BaseModel,
params: BaseModel | None = None,
stream: bool = False,
) -> StreamingResponse:
return self._route(path).send(
path, headers=headers, json=json, params=params, stream=stream
)
def probe(self, path: str, *, params: BaseModel) -> ProbeResult:
return self._route(path).probe(path, params=params)

View file

@ -45,7 +45,18 @@ async def test_mcp_server_works_without_config_auth_value():
@pytest.mark.parametrize("token_key", ["authentication_token", "auth_value"])
async def test_mcp_server_config_auth_value_header_used(token_key):
"""Ensure auth header is sent when auth token configured in config"""
"""Ensure the configured auth token is emitted as the upstream Authorization header.
The token is resolved through the v2 credential resolver and rides on the client's
httpx.Auth, so assert the header it writes onto the request rather than the (now
credential-free) _get_auth_headers() dict.
"""
import httpx
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
StaticHeaderAuth,
)
config = {
"test_server": {
"url": "https://api.example.com/mcp",
@ -60,7 +71,8 @@ async def test_mcp_server_config_auth_value_header_used(token_key):
server = next(iter(manager.config_mcp_servers.values()))
client = await manager._create_mcp_client(server)
headers = client._get_auth_headers()
assert headers["Authorization"] == "Bearer example_token"
assert isinstance(client._resolved_auth, StaticHeaderAuth)
emitted = next(client._resolved_auth.auth_flow(httpx.Request("POST", server.url)))
assert emitted.headers["Authorization"] == "Bearer example_token"
assert client.auth_type == MCPAuth.bearer_token

View file

@ -1089,7 +1089,9 @@ async def test_list_tools_only_returns_allowed_servers(monkeypatch):
mock_client_constructor,
):
# Call list_tools
tools = await test_manager.list_tools(user_api_key_auth=MagicMock())
from litellm.proxy._types import UserAPIKeyAuth
tools = await test_manager.list_tools(user_api_key_auth=UserAPIKeyAuth())
# Should only return tools from server_a
assert len(tools) == 1
# The server should use the server_name as prefix since no alias is provided

11
tests/pyrightconfig.json Normal file
View file

@ -0,0 +1,11 @@
{
"include": ["e2e"],
"exclude": ["**/node_modules", "**/__pycache__"],
"pythonVersion": "3.12",
"typeCheckingMode": "strict",
"enableTypeIgnoreComments": false,
"reportMissingImports": false,
"reportPrivateImportUsage": false,
"reportExplicitAny": "error",
"reportAny": "error"
}

View file

@ -1,17 +1,10 @@
import sys
import os
import traceback
from dotenv import load_dotenv
from fastapi import Request
from datetime import datetime
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
from litellm import Router
import pytest
import litellm
from unittest.mock import patch, MagicMock, AsyncMock
import json
from io import BytesIO
@ -76,6 +69,29 @@ def test_tuple_input(sample_jsonl_bytes):
assert result.content_type == "application/jsonl"
def test_tuple_with_file_handle_rewrites_model(sample_jsonl_bytes):
"""Security regression: when the tuple's content element is a file handle
(batch uploads stream from the spooled upload handle), the model must still
be rewritten. Otherwise a restricted body.model survives unmodified and
bypasses the batch model allowlist, which only checks the upload target."""
new_model = "approved-target-model"
handle = BytesIO(sample_jsonl_bytes)
test_tuple = ("test.jsonl", handle, "application/json")
result = replace_model_in_jsonl(test_tuple, new_model)
assert isinstance(result, InMemoryFile)
rows = [
json.loads(line)
for line in result.getvalue().decode("utf-8").splitlines()
if line.strip()
]
assert rows, "rewrite must produce rows"
# every row now carries the rewritten target, not the original (restricted) model
assert all(row["body"]["model"] == new_model for row in rows)
assert all(row["body"]["model"] != "gpt-5.5" for row in rows)
def test_file_like_object(sample_file_like):
"""Test with file-like object input"""
new_model = "claude-3"
@ -129,9 +145,9 @@ def test_should_replace_model_in_jsonl():
"""Test that should_replace_model_in_jsonl returns the correct value"""
from litellm.router_utils.batch_utils import should_replace_model_in_jsonl
assert should_replace_model_in_jsonl(purpose="batch") == True
assert should_replace_model_in_jsonl(purpose="test") == False
assert should_replace_model_in_jsonl(purpose="user_data") == False
assert should_replace_model_in_jsonl(purpose="batch") is True
assert should_replace_model_in_jsonl(purpose="test") is False
assert should_replace_model_in_jsonl(purpose="user_data") is False
def test_parse_jsonl_with_embedded_newlines_simple():
@ -217,6 +233,63 @@ def test_parse_jsonl_with_embedded_newlines_whitespace_only():
assert len(result) == 0
def test_replace_model_in_jsonl_malformed_middle_row_returns_original():
"""Regression: a malformed/truncated middle row must not silently drop the
rows that follow it. The streaming rewrite accumulates physical lines into a
buffer; a row that never parses poisons the buffer so every later valid row
is concatenated into it and dropped. Returning that partial rewrite would
ship a truncated batch with no error to the caller. Instead the original
content is returned unchanged so the provider rejects the bad batch loudly."""
content = (
b'{"custom_id":"a","body":{"model":"x"}}\n'
b'{"custom_id":"b","body":{"model":\n' # truncated, never completes
b'{"custom_id":"c","body":{"model":"x"}}\n'
)
result = replace_model_in_jsonl(content, "new-model")
assert (
result == content
), "must return the original unchanged, not a partial rewrite"
def test_replace_model_in_jsonl_malformed_row_seekable_handle_rewound():
"""When the source is a seekable handle that gets consumed during the failed
rewrite, it must be rewound to 0 so the caller can re-read the full original."""
content = (
b'{"custom_id":"a","body":{"model":"x"}}\n'
b'{"custom_id":"b","body":{"model":\n'
b'{"custom_id":"c","body":{"model":"x"}}\n'
)
handle = BytesIO(content)
result = replace_model_in_jsonl(handle, "new-model")
assert result is handle
assert handle.read() == content, "handle must be rewound for the caller to re-read"
def test_replace_model_in_jsonl_multi_row_rewrites_every_model():
"""Happy path: a well-formed multi-row file gets every row's model rewritten
and no row is dropped."""
content = (
b'{"custom_id":"a","body":{"model":"old1"}}\n'
b'{"custom_id":"b","body":{"model":"old2"}}\n'
b'{"custom_id":"c","body":{"model":"old3"}}\n'
)
result = replace_model_in_jsonl(content, "new-model")
assert isinstance(result, InMemoryFile)
rows = [
json.loads(line)
for line in result.getvalue().decode("utf-8").splitlines()
if line.strip()
]
assert [row["custom_id"] for row in rows] == ["a", "b", "c"]
assert all(row["body"]["model"] == "new-model" for row in rows)
def test_replace_model_in_jsonl_with_embedded_newlines():
"""Test that replace_model_in_jsonl works correctly with embedded newlines in content"""
# Create a JSONL with embedded newlines in the message content

View file

@ -1,6 +1,5 @@
import asyncio
import os
import ssl
import sys
from unittest.mock import AsyncMock, MagicMock, patch
@ -543,5 +542,45 @@ class TestExecuteSessionOperationSurfacesTransportError:
assert result == "done"
class TestMCPClientResolvedAuth:
"""A pre-resolved httpx.Auth is attached to the upstream client's auth= slot."""
@pytest.mark.asyncio
async def test_resolved_auth_feeds_the_auth_slot(self):
resolved = httpx.Auth()
client = MCPClient(
server_url="https://upstream.example.com", resolved_auth=resolved
)
http_client = client._create_httpx_client_factory()()
try:
assert http_client.auth is resolved
finally:
await http_client.aclose()
@pytest.mark.asyncio
async def test_resolved_auth_takes_precedence_over_aws_auth(self):
resolved = httpx.Auth()
client = MCPClient(
server_url="https://upstream.example.com",
resolved_auth=resolved,
aws_auth=httpx.Auth(),
)
http_client = client._create_httpx_client_factory()()
try:
assert http_client.auth is resolved
finally:
await http_client.aclose()
@pytest.mark.asyncio
async def test_without_resolved_auth_falls_back_to_aws_auth(self):
aws = httpx.Auth()
client = MCPClient(server_url="https://upstream.example.com", aws_auth=aws)
http_client = client._create_httpx_client_factory()()
try:
assert http_client.auth is aws
finally:
await http_client.aclose()
if __name__ == "__main__":
pytest.main([__file__])

View file

@ -8,8 +8,8 @@ Regression test for: UTF-8 codec error when uploading binary files
"""
import io
import json
import pytest
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
@ -137,11 +137,11 @@ class TestVertexAIBinaryFileUpload:
), "Binary file data should remain as bytes"
@pytest.mark.asyncio
async def test_jsonl_file_upload_returns_string(self):
async def test_jsonl_file_upload_returns_resumable_stream(self):
"""
Test that JSONL files (text) are correctly transformed to strings.
This ensures we handle both binary and text files correctly.
Test that JSONL batch files are transformed into a resumable-upload config
carrying a streaming body (not a buffered bytes payload), so the handler
can stream the upload to GCS in bounded chunks.
"""
# Create mock JSONL content
mock_jsonl_content = (
@ -164,10 +164,16 @@ class TestVertexAIBinaryFileUpload:
litellm_params={},
)
# JSONL files should be transformed to string
assert isinstance(
transformed_request, str
), f"Expected string for JSONL file, got {type(transformed_request)}"
assert (
isinstance(transformed_request, dict)
and "resumable_chunked_upload" in transformed_request
), f"Expected a resumable upload config for JSONL, got {type(transformed_request)}"
stream = transformed_request["resumable_chunked_upload"]["body_stream"]
decoded = json.loads(b"".join(stream.iter_bytes()).decode("utf-8"))
assert (
"request" in decoded
), "JSONL transform must wrap each row in {'request': ...}"
@pytest.mark.asyncio
async def test_mixed_file_types_in_sequence(self):
@ -208,7 +214,7 @@ class TestVertexAIBinaryFileUpload:
optional_params={},
litellm_params={},
)
assert isinstance(result2, str)
assert isinstance(result2, dict) and "resumable_chunked_upload" in result2
# Test 3: Upload another binary file
binary_content2 = b"\xc4\xe5\xf2\xe5\xeb"
@ -251,7 +257,7 @@ class TestVertexAIBinaryFileUpload:
},
"text_files": {
"input_type": "str or bytes",
"output_type": "str",
"output_type": "bytes",
"examples": ["JSONL", "CSV", "TXT"],
"http_method": "POST",
"encoding": "UTF-8",

View file

@ -0,0 +1,696 @@
"""
Tests for the streaming OpenAI -> Vertex JSONL batch transform.
The transform converts batch uploads entry-by-entry rather than materializing
the payload in full intermediate lists (decoded str, parsed dicts, transformed
dicts, joined output), which keeps peak memory bounded on large uploads.
These tests lock in the behaviour that would regress if the streaming path were
replaced by a list-based pipeline:
1. Byte-for-byte output parity with a list pipeline (wire format).
2. The streaming transform peaks at a clear fraction of a list pipeline on the
same input (relative differential, robust to GC noise).
3. ``get_object_name`` only parses the first JSONL row, so a payload whose
later rows are not valid JSON does not raise.
4. A tuple-wrapped file handle uploaded through the real create_file ordering
keeps every row, including entry 0 (no partial upload from a consumed
cursor).
"""
import gc
import io
import json
import time
import tracemalloc
import httpx
import pytest
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.base_llm.files.transformation import BaseFileUploadStream
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.llms.vertex_ai.files.transformation import (
VertexAIFilesConfig,
_OpenAIToVertexBatchUploadStream,
_get_litellm_batch_custom_id_from_labels,
_iter_openai_jsonl_entries,
_iter_openai_jsonl_lines,
_openai_batch_jsonl_entry_to_vertex_wrapped_request,
)
from litellm.types.llms.openai import CreateFileRequest
def _resumable_stream(transformed) -> BaseFileUploadStream:
"""Pull the streaming body out of a resumable-upload transform result."""
return transformed["resumable_chunked_upload"]["body_stream"]
def _join_upload_body(transformed) -> bytes:
"""Materialize a transform result's upload body for byte-level assertions."""
if isinstance(transformed, dict) and "resumable_chunked_upload" in transformed:
return b"".join(_resumable_stream(transformed).iter_bytes())
if isinstance(transformed, BaseFileUploadStream):
return b"".join(transformed.iter_bytes())
if isinstance(transformed, str):
return transformed.encode("utf-8")
return transformed
def _make_openai_jsonl_bytes(n_rows: int, padding: int = 400) -> bytes:
pad = "x" * padding
rows = []
for i in range(n_rows):
rows.append(
json.dumps(
{
"custom_id": f"request-{i}",
"method": "POST",
"url": "/v1/chat/completions",
"body": {
"model": "gemini-2.5-flash",
"messages": [{"role": "user", "content": f"{pad} {i}"}],
"max_tokens": 4,
},
}
)
)
return ("\n".join(rows)).encode("utf-8")
def _reference_vertex_jsonl_string(cfg: VertexAIFilesConfig, content: str) -> str:
"""Row-by-row reference output built eagerly from the live single-entry
transform, so the streaming path can be checked against it for parity."""
entries = [json.loads(line) for line in content.splitlines() if line.strip()]
return "\n".join(
json.dumps(
_openai_batch_jsonl_entry_to_vertex_wrapped_request(
entry, cfg._map_openai_to_vertex_params
)
)
for entry in entries
)
class TestStreamingOutputParity:
def test_transform_create_file_request_returns_resumable_stream_parity(self):
cfg = VertexAIFilesConfig()
raw = _make_openai_jsonl_bytes(300)
request: CreateFileRequest = {
"file": ("batch.jsonl", raw, "application/jsonl"),
"purpose": "batch",
}
out = cfg.transform_create_file_request(
model="", create_file_data=request, optional_params={}, litellm_params={}
)
# A batch upload must be a resumable-upload config carrying a streaming
# body, so the handler can chunk it; a buffered bytes/str return would
# defeat the OOM fix.
assert isinstance(out, dict) and "resumable_chunked_upload" in out
assert isinstance(_resumable_stream(out), BaseFileUploadStream)
assert _join_upload_body(out).decode("utf-8") == _reference_vertex_jsonl_string(
cfg, raw.decode("utf-8")
)
class TestFileLikeInputNotPartiallyConsumed:
"""
In ``llm_http_handler.create_file`` the object-name step
(get_complete_file_url -> get_object_name) runs before
transform_create_file_request, and both read the same create_file_data
source. When the file is a tuple-wrapped open handle, the streaming reader
must still emit every row including entry 0: ``_iter_openai_jsonl_lines``
rewinds a seekable source (seek(0)) before each pass, so the object-name
step's partial read of the cursor does not consume the upload. A partial
upload missing the first request would be silent and hard to catch, so this
locks the full-payload invariant in.
"""
def test_filehandle_create_file_keeps_first_entry(self):
cfg = VertexAIFilesConfig()
n_rows = 25
raw = _make_openai_jsonl_bytes(n_rows)
create_file_data: CreateFileRequest = {
"file": ("batch.jsonl", io.BytesIO(raw), "application/jsonl"),
"purpose": "batch",
}
# Object-name step first (as the handler does), then the transform, both
# reading the same live BytesIO handle.
cfg.get_complete_file_url(
api_base=None,
api_key=None,
model="",
optional_params={},
litellm_params={"gcs_bucket_name": "test-bucket"},
data=create_file_data,
)
out = cfg.transform_create_file_request(
model="",
create_file_data=create_file_data,
optional_params={},
litellm_params={},
)
lines = _join_upload_body(out).decode("utf-8").splitlines()
assert len(lines) == n_rows, "no batch row may be dropped from the upload"
first_labels = json.loads(lines[0])["request"]["labels"]
assert _get_litellm_batch_custom_id_from_labels(first_labels) == "request-0"
class TestStreamingLineIterator:
def test_skips_blank_and_whitespace_lines(self):
content = b'{"a": 1}\n\n \n{"b": 2}\n'
assert list(_iter_openai_jsonl_lines(content)) == ['{"a": 1}', '{"b": 2}']
def test_handles_crlf_and_missing_trailing_newline(self):
content = b'{"a": 1}\r\n{"b": 2}'
assert [json.loads(line) for line in _iter_openai_jsonl_lines(content)] == [
{"a": 1},
{"b": 2},
]
def test_accepts_str_bytes_tuple_and_filelike(self):
expected = [{"a": 1}, {"b": 2}]
text = '{"a": 1}\n{"b": 2}\n'
for source in (
text,
text.encode("utf-8"),
("name.jsonl", text.encode("utf-8"), "application/jsonl"),
io.BytesIO(text.encode("utf-8")),
):
assert list(_iter_openai_jsonl_entries(source)) == expected
def test_str_input_without_trailing_newline(self):
assert list(_iter_openai_jsonl_lines('{"a": 1}\n{"b": 2}')) == [
'{"a": 1}',
'{"b": 2}',
]
def test_pathlike_input_is_read_line_by_line(self, tmp_path):
path = tmp_path / "batch.jsonl"
path.write_bytes(b'{"a": 1}\n{"b": 2}\n')
assert list(_iter_openai_jsonl_entries(path)) == [{"a": 1}, {"b": 2}]
def test_unsupported_content_type_raises(self):
with pytest.raises(ValueError, match="Unsupported file content type"):
list(_iter_openai_jsonl_lines(12345)) # type: ignore[arg-type]
def test_non_seekable_handle_raises_instead_of_dropping_first_row(self):
# The handle is read twice (object-name probe, then body). A non-seekable
# handle can't rewind, so it must fail loudly rather than silently resume
# mid-stream and omit the opening batch request.
class _NonSeekable:
def __init__(self, raw: bytes):
self._buf = io.BytesIO(raw)
def read(self, *args):
return self._buf.read(*args)
def __iter__(self):
return iter(self._buf)
def seek(self, *args):
raise io.UnsupportedOperation("not seekable")
handle = _NonSeekable(
b'{"custom_id": "request-0"}\n{"custom_id": "request-1"}\n'
)
with pytest.raises(ValueError, match="seekable"):
list(_iter_openai_jsonl_lines(handle))
def test_is_lazy_does_not_parse_past_first_entry(self):
# Second row is invalid JSON; pulling only the first entry must not raise.
content = b'{"custom_id": "first"}\nnot-json-at-all\n'
gen = _iter_openai_jsonl_entries(content)
assert next(gen)["custom_id"] == "first"
with pytest.raises(json.JSONDecodeError):
next(gen)
class TestGetObjectNameLazyParse:
def test_only_parses_first_row_for_model(self):
cfg = VertexAIFilesConfig()
# Tail rows are deliberately not valid JSON. Parsing the whole payload
# would raise here; a first-row-only parse must not.
raw = (
b'{"custom_id": "r-0", "body": {"model": "gemini-2.5-flash"}}\n'
b"garbage line that is not json\n"
)
object_name = cfg.get_object_name(
("batch.jsonl", raw, "application/jsonl"), purpose="batch"
)
assert "gemini-2.5-flash" in object_name
class TestStreamingPeakMemory:
"""
Differential guard: the streaming transform must stay well under the peak
that a list pipeline incurs on the same input. If the hot path builds full
intermediate lists, the streaming assertion fails.
The assertion that matters is the *relative* one: ``streaming_peak`` must be
a clear fraction of ``list_peak`` on the identical input. Absolute
``tracemalloc`` ratios drift with GC timing and the live set carried in from
earlier tests, so they make poor CI gates; the relative comparison cancels
that shared noise and is exactly what regresses (toward 1.0) when the hot
path builds full intermediate lists. ``gc.collect()`` before each
measurement removes any garbage the previous run left behind.
"""
def _measure(self, fn):
gc.collect()
tracemalloc.start()
try:
fn()
_, peak = tracemalloc.get_traced_memory()
finally:
tracemalloc.stop()
return peak
def test_streaming_peak_well_below_list_pipeline(self):
cfg = VertexAIFilesConfig()
raw = _make_openai_jsonl_bytes(8000)
content_str = raw.decode("utf-8")
def drain_stream():
# Consume the upload body one row at a time, as the chunked uploader
# does, without accumulating it.
for _ in _OpenAIToVertexBatchUploadStream(
raw, cfg._map_openai_to_vertex_params
).iter_bytes():
pass
streaming_peak = self._measure(drain_stream)
list_peak = self._measure(
lambda: _reference_vertex_jsonl_string(cfg, content_str)
)
# Core guard: the lazily consumed streaming body peaks well under a list
# pipeline that materializes every transformed row. Building full
# intermediate lists in the hot path pushes this ratio back toward 1.0.
assert streaming_peak < list_peak * 0.6, (
f"streaming peak {streaming_peak} not a clear win over list pipeline "
f"{list_peak} (ratio {streaming_peak / list_peak:.2f})"
)
def test_get_object_name_does_not_scale_with_payload(self):
cfg = VertexAIFilesConfig()
raw = _make_openai_jsonl_bytes(8000)
file_data = ("batch.jsonl", raw, "application/jsonl")
# The payload bytes already exist before measurement starts, so a lazy
# first-row parse should allocate only a small fraction of the payload;
# parsing every row would blow past this bound.
peak = self._measure(lambda: cfg.get_object_name(file_data, purpose="batch"))
assert (
peak / len(raw) < 2.0
), "get_object_name should not copy the whole payload"
class TestPathSourcedStreaming:
"""
The proxy spools large batch uploads to a temp file and passes a pathlib.Path
as the file content instead of pre-reading bytes, so the transform streams
from disk. These lock in that a Path source yields identical output, keeps
every row, stays memory-bounded, and is re-iterable (multi-model uploads).
"""
def _write_jsonl(self, tmp_path, n_rows, padding=400):
raw = _make_openai_jsonl_bytes(n_rows, padding=padding)
path = tmp_path / "batch.jsonl"
path.write_bytes(raw)
return path, raw
def _batch_request(self, path) -> CreateFileRequest:
return {"file": ("batch.jsonl", path, "application/jsonl"), "purpose": "batch"}
def test_transform_from_path_matches_legacy_and_keeps_all_rows(self, tmp_path):
cfg = VertexAIFilesConfig()
n_rows = 200
path, raw = self._write_jsonl(tmp_path, n_rows)
data = self._batch_request(path)
url = cfg.get_complete_file_url(
api_base=None,
api_key=None,
model="",
optional_params={},
litellm_params={"gcs_bucket_name": "test-bucket"},
data=data,
)
assert "uploadType=resumable" in url
out = cfg.transform_create_file_request(
model="", create_file_data=data, optional_params={}, litellm_params={}
)
assert isinstance(out, dict) and "resumable_chunked_upload" in out
body = _join_upload_body(out).decode("utf-8")
assert body == _reference_vertex_jsonl_string(cfg, raw.decode("utf-8"))
lines = body.splitlines()
assert len(lines) == n_rows, "no batch row may be dropped from a Path source"
first_labels = json.loads(lines[0])["request"]["labels"]
assert _get_litellm_batch_custom_id_from_labels(first_labels) == "request-0"
def test_path_source_peak_stays_below_payload(self, tmp_path):
cfg = VertexAIFilesConfig()
path, raw = self._write_jsonl(tmp_path, 8000)
data = self._batch_request(path)
def run():
cfg.get_complete_file_url(
api_base=None,
api_key=None,
model="",
optional_params={},
litellm_params={"gcs_bucket_name": "test-bucket"},
data=data,
)
out = cfg.transform_create_file_request(
model="", create_file_data=data, optional_params={}, litellm_params={}
)
for _ in _resumable_stream(out).iter_bytes():
pass # drain without accumulating
gc.collect()
tracemalloc.start()
try:
run()
_, peak = tracemalloc.get_traced_memory()
finally:
tracemalloc.stop()
# Streaming from disk must not materialize the payload. Reading the whole
# file into bytes (the pre-fix path) would push peak past the file size.
assert peak < len(raw) * 0.3, (
f"peak {peak} not bounded vs payload {len(raw)} "
f"(ratio {peak / len(raw):.2f})"
)
def test_path_source_stream_is_reiterable(self, tmp_path):
cfg = VertexAIFilesConfig()
path, _ = self._write_jsonl(tmp_path, 50)
data = self._batch_request(path)
out = cfg.transform_create_file_request(
model="", create_file_data=data, optional_params={}, litellm_params={}
)
stream = _resumable_stream(out)
first = b"".join(stream.iter_bytes())
second = b"".join(stream.iter_bytes())
assert first == second and len(first) > 0
_GCS_OBJECT_JSON = {
"id": "test-bucket/litellm-vertex-files/x/123",
"name": "litellm-vertex-files/x",
"size": "0",
"timeCreated": "2026-01-01T00:00:00.000000Z",
"purpose": "batch",
}
class _FixedBytesStream(BaseFileUploadStream):
"""Streaming body of exact, controllable bytes for protocol-edge tests."""
def __init__(self, data: bytes, piece: int = 64):
self._data = data
self._piece = piece
def iter_bytes(self):
for i in range(0, len(self._data), self._piece):
yield self._data[i : i + self._piece]
def _logging_obj() -> Logging:
return Logging(
model="",
messages=[],
stream=False,
call_type="acreate_file",
start_time=time.time(),
litellm_call_id="test",
function_id="",
)
def _gcs_resumable_mock(session_url: str, final_status: int = 200):
"""A fake GCS resumable endpoint: POST opens a session (URI in Location),
each PUT appends and returns 308 until the final chunk returns 200/201."""
state = {"received": bytearray(), "ranges": [], "methods": [], "urls": []}
async def handler(request: httpx.Request) -> httpx.Response:
state["methods"].append(request.method)
state["urls"].append(str(request.url))
if request.method == "POST":
return httpx.Response(200, headers={"location": session_url})
body = await request.aread()
content_range = request.headers["content-range"]
state["ranges"].append(content_range)
state["received"].extend(body)
if content_range.rsplit("/", 1)[-1] == "*":
return httpx.Response(
308, headers={"range": f"bytes=0-{len(state['received']) - 1}"}
)
return httpx.Response(final_status, json=_GCS_OBJECT_JSON)
return handler, state
def _async_handler_with(mock) -> AsyncHTTPHandler:
handler = AsyncHTTPHandler()
handler.client = httpx.AsyncClient(transport=httpx.MockTransport(mock))
return handler
class TestResumableUploadUrl:
def test_batch_jsonl_uses_resumable_upload_type(self):
cfg = VertexAIFilesConfig()
request: CreateFileRequest = {
"file": ("batch.jsonl", _make_openai_jsonl_bytes(3), "application/jsonl"),
"purpose": "batch",
}
url = cfg.get_complete_file_url(
api_base=None,
api_key=None,
model="",
optional_params={},
litellm_params={"gcs_bucket_name": "test-bucket"},
data=request,
)
assert "uploadType=resumable" in url
assert "uploadType=media" not in url
def test_batch_text_plain_uses_resumable_upload_type(self):
# Clients often label a .jsonl batch upload as text/plain; it must still
# take the streaming/resumable path, not the buffered media path.
cfg = VertexAIFilesConfig()
request: CreateFileRequest = {
"file": ("batch.jsonl", _make_openai_jsonl_bytes(3), "text/plain"),
"purpose": "batch",
}
url = cfg.get_complete_file_url(
api_base=None,
api_key=None,
model="",
optional_params={},
litellm_params={"gcs_bucket_name": "test-bucket"},
data=request,
)
assert "uploadType=resumable" in url
assert "uploadType=media" not in url
def test_binary_upload_stays_simple_media(self):
cfg = VertexAIFilesConfig()
request: CreateFileRequest = {
"file": ("doc.pdf", b"%PDF-1.4 binary", "application/pdf"),
"purpose": "user_data",
}
url = cfg.get_complete_file_url(
api_base=None,
api_key=None,
model="",
optional_params={},
litellm_params={"gcs_bucket_name": "test-bucket"},
data=request,
)
assert "uploadType=media" in url
assert "uploadType=resumable" not in url
class TestResumableStreamBody:
def test_stream_matches_legacy_pipeline(self):
cfg = VertexAIFilesConfig()
raw = _make_openai_jsonl_bytes(120)
stream = _OpenAIToVertexBatchUploadStream(raw, cfg._map_openai_to_vertex_params)
assert b"".join(stream.iter_bytes()).decode(
"utf-8"
) == _reference_vertex_jsonl_string(cfg, raw.decode("utf-8"))
def test_stream_is_reiterable_for_retries(self):
# A one-shot generator would make a transport retry upload an empty body;
# iter_bytes() must yield the full payload every call.
cfg = VertexAIFilesConfig()
raw = _make_openai_jsonl_bytes(40)
stream = _OpenAIToVertexBatchUploadStream(raw, cfg._map_openai_to_vertex_params)
first = b"".join(stream.iter_bytes())
second = b"".join(stream.iter_bytes())
assert first == second and len(first) > 0
def test_stream_is_reiterable_for_seekable_file_like_input(self):
# A seekable handle (BytesIO, temp file) must be rewound between calls;
# otherwise the first iter_bytes() exhausts it and a retry would upload
# an empty body silently.
cfg = VertexAIFilesConfig()
raw = _make_openai_jsonl_bytes(40)
stream = _OpenAIToVertexBatchUploadStream(
io.BytesIO(raw), cfg._map_openai_to_vertex_params
)
first = b"".join(stream.iter_bytes())
second = b"".join(stream.iter_bytes())
assert first == second and len(first) > 0
class TestResumableChunking:
def test_intermediate_chunks_are_exactly_chunk_size(self):
pieces = list(BaseLLMHTTPHandler._iter_resumable_chunks(iter([b"x" * 10]), 4))
assert pieces == [b"xxxx", b"xxxx", b"xx"]
def test_exact_multiple_yields_no_trailing_empty(self):
# An exactly chunk-aligned stream yields only full chunks; the upload
# finalizes on the last data chunk instead of an extra empty request.
pieces = list(BaseLLMHTTPHandler._iter_resumable_chunks(iter([b"x" * 8]), 4))
assert pieces == [b"xxxx", b"xxxx"]
def test_empty_stream_yields_nothing(self):
# A 0-byte stream yields no chunks; the caller finalizes with one empty
# request (bytes */0).
assert list(BaseLLMHTTPHandler._iter_resumable_chunks(iter([]), 4)) == []
def test_default_chunk_size_is_256kib_multiple(self):
assert BaseLLMHTTPHandler._RESUMABLE_CHUNK_SIZE % (256 * 1024) == 0
def test_content_range_intermediate_uses_star_total(self):
assert (
BaseLLMHTTPHandler._resumable_content_range(0, 4096, is_final=False)
== "bytes 0-4095/*"
)
def test_content_range_final_uses_real_total(self):
assert (
BaseLLMHTTPHandler._resumable_content_range(8192, 100, is_final=True)
== "bytes 8192-8291/8292"
)
def test_content_range_empty_finalize(self):
assert (
BaseLLMHTTPHandler._resumable_content_range(8192, 0, is_final=True)
== "bytes */8192"
)
@pytest.mark.asyncio
class TestResumableUploadProtocol:
"""End-to-end against a faked GCS resumable endpoint. These are the tests
that fail if the handler buffers the whole body, drops bytes, mislabels a
Content-Range, follows the 308 instead of continuing, or skips finalize."""
async def _run(self, raw: bytes, chunk_size: int, final_status: int = 200):
cfg = VertexAIFilesConfig()
request: CreateFileRequest = {
"file": ("batch.jsonl", raw, "application/jsonl"),
"purpose": "batch",
}
api_base = cfg.get_complete_file_url(
api_base=None,
api_key=None,
model="",
optional_params={},
litellm_params={"gcs_bucket_name": "test-bucket"},
data=request,
)
transformed = cfg.transform_create_file_request(
model="", create_file_data=request, optional_params={}, litellm_params={}
)
transformed["resumable_chunked_upload"]["chunk_size"] = chunk_size
expected = _join_upload_body(transformed)
session_url = "https://storage.googleapis.com/upload/sess?upload_id=SID"
mock, state = _gcs_resumable_mock(session_url, final_status=final_status)
response = await BaseLLMHTTPHandler().async_create_file(
transformed_request=transformed,
litellm_params={},
provider_config=cfg,
headers={"Authorization": "Bearer x"},
api_base=api_base,
logging_obj=_logging_obj(),
client=_async_handler_with(mock),
timeout=None,
)
return expected, state, response, session_url, api_base
async def test_streams_in_chunks_and_reassembles(self):
raw = _make_openai_jsonl_bytes(300)
chunk_size = 4096
expected, state, response, session_url, api_base = await self._run(
raw, chunk_size
)
# One session-open POST, then a sequence of chunk PUTs.
assert state["methods"][0] == "POST"
assert set(state["methods"][1:]) == {"PUT"}
assert state["methods"].count("PUT") >= 2, "payload must span multiple chunks"
# POST opens a resumable session; every chunk goes to the session URI.
assert "uploadType=resumable" in state["urls"][0]
assert all(u == session_url for u in state["urls"][1:])
# Every non-final chunk is exactly chunk_size with an unknown-total range;
# the final chunk carries the real total.
intermediate = state["ranges"][:-1]
for index, content_range in enumerate(intermediate):
assert (
content_range
== f"bytes {index * chunk_size}-{(index + 1) * chunk_size - 1}/*"
)
total = len(expected)
last_offset = len(intermediate) * chunk_size
if last_offset == total: # payload landed on a chunk boundary
assert state["ranges"][-1] == f"bytes */{total}"
else:
assert state["ranges"][-1] == f"bytes {last_offset}-{total - 1}/{total}"
# The bytes GCS received are exactly the transformed batch payload.
assert bytes(state["received"]) == expected
assert response.object == "file"
async def test_exact_multiple_finalizes_on_last_data_chunk(self):
# A body that is an exact multiple of the chunk size finalizes on its
# last data chunk (bytes (TOTAL-chunk)-(TOTAL-1)/TOTAL), with no extra
# empty finalize request.
chunk_size = 256
total = chunk_size * 3
stream = _FixedBytesStream(b"a" * total)
config = {"body_stream": stream, "chunk_size": chunk_size}
session_url = "https://storage.googleapis.com/upload/sess?upload_id=SID"
mock, state = _gcs_resumable_mock(session_url)
response = await BaseLLMHTTPHandler()._aresumable_chunked_upload(
client=_async_handler_with(mock),
initiate_url="https://storage.googleapis.com/upload?uploadType=resumable",
base_headers={"Authorization": "Bearer x"},
config=config,
timeout=None,
)
assert state["ranges"][-1] == f"bytes {total - chunk_size}-{total - 1}/{total}"
assert "*" not in state["ranges"][-1]
assert bytes(state["received"]) == b"a" * total
assert response.status_code == 200
async def test_failed_chunk_raises(self):
raw = _make_openai_jsonl_bytes(80)
with pytest.raises(Exception):
await self._run(raw, chunk_size=4096, final_status=403)

View file

@ -14,8 +14,8 @@ from unittest.mock import MagicMock
from litellm.llms.vertex_ai.files.transformation import (
VertexAIFilesConfig,
VertexAIJsonlFilesTransformation,
_get_litellm_batch_custom_id_from_labels,
_openai_batch_jsonl_entry_to_vertex_wrapped_request,
_sanitize_gcp_label_value,
)
from litellm.types.llms.openai import OpenAIFileObject, HttpxBinaryResponseContent
@ -33,7 +33,7 @@ class TestParseGcsUri:
def test_should_parse_standard_gs_uri(self, config):
file_id = "gs://my-bucket/litellm-vertex-files/path/to/object.jsonl"
bucket, encoded = config._parse_gcs_uri(
file_id, litellm_params={"bucket_name": "my-bucket"}
file_id, litellm_params={"gcs_bucket_name": "my-bucket"}
)
assert bucket == "my-bucket"
assert encoded == urllib.parse.quote(
@ -43,7 +43,7 @@ class TestParseGcsUri:
def test_should_parse_uri_with_nested_publisher_path(self, config):
uri = "gs://litellm-local/litellm-vertex-files/publishers/google/models/gemini-2.0-flash-001/abc-123"
bucket, encoded = config._parse_gcs_uri(
uri, litellm_params={"bucket_name": "litellm-local"}
uri, litellm_params={"gcs_bucket_name": "litellm-local"}
)
assert bucket == "litellm-local"
expected_path = (
@ -56,7 +56,7 @@ class TestParseGcsUri:
"gs://my-bucket/litellm-vertex-files/some/path", safe=""
)
bucket, encoded = config._parse_gcs_uri(
encoded_uri, litellm_params={"bucket_name": "my-bucket"}
encoded_uri, litellm_params={"gcs_bucket_name": "my-bucket"}
)
assert bucket == "my-bucket"
assert encoded == urllib.parse.quote("litellm-vertex-files/some/path", safe="")
@ -64,21 +64,21 @@ class TestParseGcsUri:
def test_should_reject_bucket_only(self, config):
with pytest.raises(ValueError, match="object name"):
config._parse_gcs_uri(
"gs://my-bucket", litellm_params={"bucket_name": "my-bucket"}
"gs://my-bucket", litellm_params={"gcs_bucket_name": "my-bucket"}
)
def test_should_reject_no_gs_prefix(self, config):
with pytest.raises(ValueError, match="gs://"):
config._parse_gcs_uri(
"my-bucket/litellm-vertex-files/object.txt",
litellm_params={"bucket_name": "my-bucket"},
litellm_params={"gcs_bucket_name": "my-bucket"},
)
def test_should_reject_unmanaged_object_path(self, config):
with pytest.raises(ValueError, match="LiteLLM-managed"):
config._parse_gcs_uri(
"gs://my-bucket/private/object.txt",
litellm_params={"bucket_name": "my-bucket"},
litellm_params={"gcs_bucket_name": "my-bucket"},
)
def test_should_reject_request_supplied_legacy_flag(self, config):
@ -86,7 +86,7 @@ class TestParseGcsUri:
config._parse_gcs_uri(
"gs://my-bucket/private/object.txt",
litellm_params={
"bucket_name": "my-bucket",
"gcs_bucket_name": "my-bucket",
"allow_legacy_cloud_file_ids": True,
},
)
@ -96,7 +96,7 @@ class TestParseGcsUri:
bucket, encoded = config._parse_gcs_uri(
"gs://my-bucket/private/object.txt",
litellm_params={
"bucket_name": "my-bucket",
"gcs_bucket_name": "my-bucket",
"_litellm_internal_model_credentials": trusted_credentials,
},
)
@ -109,7 +109,7 @@ class TestParseGcsUri:
config._parse_gcs_uri(
"gs://my-bucket/private/object.txt",
litellm_params={
"bucket_name": "my-bucket",
"gcs_bucket_name": "my-bucket",
"_litellm_internal_model_credentials": {
"allow_legacy_cloud_file_ids": True
},
@ -121,7 +121,7 @@ class TestParseGcsUri:
bucket, encoded = config._parse_gcs_uri(
"gs://my-bucket/team-a/private/object.txt",
litellm_params={
"bucket_name": "my-bucket/team-a",
"gcs_bucket_name": "my-bucket/team-a",
"_litellm_internal_model_credentials": trusted_credentials,
},
)
@ -135,7 +135,7 @@ class TestParseGcsUri:
config._parse_gcs_uri(
"gs://my-bucket/team-b/private/object.txt",
litellm_params={
"bucket_name": "my-bucket/team-a",
"gcs_bucket_name": "my-bucket/team-a",
"_litellm_internal_model_credentials": trusted_credentials,
},
)
@ -144,7 +144,7 @@ class TestParseGcsUri:
with pytest.raises(ValueError, match="configured storage bucket"):
config._parse_gcs_uri(
"gs://other-bucket/litellm-vertex-files/object.txt",
litellm_params={"bucket_name": "my-bucket"},
litellm_params={"gcs_bucket_name": "my-bucket"},
)
@ -156,7 +156,7 @@ class TestCreateFileUrl:
model="",
optional_params={},
litellm_params={
"bucket_name": "safe-bucket",
"gcs_bucket_name": "safe-bucket",
"litellm_metadata": {"gcs_bucket_name": "attacker-bucket"},
},
data={
@ -182,7 +182,7 @@ class TestTransformRetrieveFile:
url, params = config.transform_retrieve_file_request(
file_id=file_id,
optional_params={},
litellm_params={"bucket_name": "my-bucket"},
litellm_params={"gcs_bucket_name": "my-bucket"},
)
expected_encoded = urllib.parse.quote(
"litellm-vertex-files/path/to/file.jsonl", safe=""
@ -243,7 +243,7 @@ class TestTransformFileContent:
url, params = config.transform_file_content_request(
file_content_request={"file_id": file_id},
optional_params={},
litellm_params={"bucket_name": "my-bucket"},
litellm_params={"gcs_bucket_name": "my-bucket"},
)
encoded = urllib.parse.quote("litellm-vertex-files/path/to/file.jsonl", safe="")
assert (
@ -378,7 +378,7 @@ class TestTransformDeleteFile:
url, params = config.transform_delete_file_request(
file_id=file_id,
optional_params={},
litellm_params={"bucket_name": "my-bucket"},
litellm_params={"gcs_bucket_name": "my-bucket"},
)
encoded = urllib.parse.quote("litellm-vertex-files/path/to/file.jsonl", safe="")
assert (
@ -854,6 +854,106 @@ class TestVertexBatchOutputTransformation:
)
assert transformed_content == invalid_content
def test_binary_content_passthrough(self, config):
"""A binary file (PDF/video) whose first bytes are not valid UTF-8 must be
returned unchanged. The row-by-row transform only engages for a JSONL
batch output and must never line-parse or corrupt binary content."""
binary = b"%PDF-1.4\n%\xc4\xe5\xf2\xe5\xeb\xa7\n" + b"\x00\x01\x02\xff\xfe" * 64
assert config._try_transform_vertex_batch_output_to_openai(binary) == binary
def test_streaming_transform_peaks_below_list_pipeline(self, config):
"""The output transform must stream row-by-row, not build a list of every
parsed row and a second list of transformed rows. This guards against a
regression to the list pipeline, which peaks at several full copies and
OOMs on large result files. The relative comparison cancels shared noise
(per-row transform cost, GC timing) and only the list overhead differs.
"""
import gc
import tracemalloc
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
def vertex_row(index: int) -> dict:
return {
"status": "",
"processed_time": "2024-11-01T18:13:16.826+00:00",
"request": {
"contents": [{"role": "user", "parts": [{"text": "hi"}]}],
"labels": {"litellm_custom_id": f"r-{index}"},
},
"response": {
"candidates": [
{
"content": {
"parts": [{"text": "hello " * 20}],
"role": "model",
},
"finishReason": "STOP",
}
],
"modelVersion": "gemini-2.0-flash-001",
"usageMetadata": {
"promptTokenCount": 10,
"candidatesTokenCount": 20,
"totalTokenCount": 30,
},
},
}
content = ("\n".join(json.dumps(vertex_row(i)) for i in range(4000))).encode(
"utf-8"
)
def list_pipeline() -> bytes:
gemini_config = VertexGeminiConfig()
logging_obj = Logging(
model="",
messages=[],
stream=False,
call_type="batch_transform",
start_time=0.1,
litellm_call_id="",
function_id="",
)
logging_obj.optional_params = {}
mock_response = httpx.Response(
status_code=200,
headers={"content-type": "application/json"},
request=httpx.Request("POST", "https://example.com"),
)
rows = content.decode("utf-8").strip().split("\n")
transformed = [
json.dumps(
config._transform_single_vertex_batch_output_to_openai(
json.loads(row), gemini_config, logging_obj, mock_response
)
)
for row in rows
]
return "\n".join(transformed).encode("utf-8")
def peak_of(fn) -> int:
gc.collect()
tracemalloc.start()
try:
fn()
return tracemalloc.get_traced_memory()[1]
finally:
tracemalloc.stop()
streaming_peak = peak_of(
lambda: config._try_transform_vertex_batch_output_to_openai(content)
)
list_peak = peak_of(list_pipeline)
assert streaming_peak < list_peak * 0.75, (
f"streaming peak {streaming_peak} is not a clear win over the list "
f"pipeline {list_peak} (ratio {streaming_peak / list_peak:.2f})"
)
class TestTryTransformDoesNotMutateCallerLoggingObj:
"""Regression tests: _try_transform_vertex_batch_output_to_openai must not mutate
@ -953,12 +1053,23 @@ class TestTryTransformDoesNotMutateCallerLoggingObj:
assert transformed["response"]["status_code"] == 200
def _wrap_entries(openai_jsonl_content):
"""Vertex-wrapped requests for a list of OpenAI batch entries, built via the
live single-entry transform that the streaming upload path uses."""
cfg = VertexAIFilesConfig()
return [
_openai_batch_jsonl_entry_to_vertex_wrapped_request(
entry, cfg._map_openai_to_vertex_params
)
for entry in openai_jsonl_content
]
class TestVertexBatchCustomIdLabels:
"""Test custom_id handling in batch transformations"""
def test_custom_id_added_to_labels_in_vertex_request(self):
"""Test that custom_id from OpenAI format is added as a label in Vertex AI format"""
transformation = VertexAIJsonlFilesTransformation()
openai_jsonl_content = [
{
@ -973,11 +1084,7 @@ class TestVertexBatchCustomIdLabels:
}
]
vertex_jsonl_content = (
transformation._transform_openai_jsonl_content_to_vertex_ai_jsonl_content(
openai_jsonl_content
)
)
vertex_jsonl_content = _wrap_entries(openai_jsonl_content)
assert len(vertex_jsonl_content) == 1
vertex_request = vertex_jsonl_content[0]
@ -992,7 +1099,6 @@ class TestVertexBatchCustomIdLabels:
def test_long_custom_id_round_trips_across_raw_label_chunks(self):
"""Test that long custom_ids are not truncated in raw labels."""
transformation = VertexAIJsonlFilesTransformation()
custom_id_a = "shared-prefix-that-is-longer-than-thirty-six-bytes-A"
custom_id_b = "shared-prefix-that-is-longer-than-thirty-six-bytes-B"
@ -1009,11 +1115,7 @@ class TestVertexBatchCustomIdLabels:
for custom_id in (custom_id_a, custom_id_b)
]
vertex_jsonl_content = (
transformation._transform_openai_jsonl_content_to_vertex_ai_jsonl_content(
openai_jsonl_content
)
)
vertex_jsonl_content = _wrap_entries(openai_jsonl_content)
labels_a = vertex_jsonl_content[0]["request"]["labels"]
labels_b = vertex_jsonl_content[1]["request"]["labels"]
@ -1028,7 +1130,6 @@ class TestVertexBatchCustomIdLabels:
def test_multiple_requests_each_get_their_own_label(self):
"""Test that multiple requests each get their own custom_id label"""
transformation = VertexAIJsonlFilesTransformation()
openai_jsonl_content = [
{
@ -1043,11 +1144,7 @@ class TestVertexBatchCustomIdLabels:
for i in range(3)
]
vertex_jsonl_content = (
transformation._transform_openai_jsonl_content_to_vertex_ai_jsonl_content(
openai_jsonl_content
)
)
vertex_jsonl_content = _wrap_entries(openai_jsonl_content)
assert len(vertex_jsonl_content) == 3
@ -1063,7 +1160,6 @@ class TestVertexBatchCustomIdLabels:
def test_request_without_custom_id_has_no_label(self):
"""Test that requests without custom_id don't get a label"""
transformation = VertexAIJsonlFilesTransformation()
openai_jsonl_content = [
{
@ -1076,11 +1172,7 @@ class TestVertexBatchCustomIdLabels:
}
]
vertex_jsonl_content = (
transformation._transform_openai_jsonl_content_to_vertex_ai_jsonl_content(
openai_jsonl_content
)
)
vertex_jsonl_content = _wrap_entries(openai_jsonl_content)
# Should not have labels if no custom_id was provided
assert "labels" not in vertex_jsonl_content[0]["request"]
@ -1090,7 +1182,6 @@ class TestVertexBatchCustomIdLabels:
Test the full round trip: OpenAI format -> Vertex AI format -> Vertex AI output -> OpenAI output
Verify that custom_id is preserved through the entire flow.
"""
transformation = VertexAIJsonlFilesTransformation()
config = VertexAIFilesConfig()
# Step 1: Transform OpenAI input to Vertex AI format (mixed case exercises raw label)
@ -1106,11 +1197,7 @@ class TestVertexBatchCustomIdLabels:
}
]
vertex_input = (
transformation._transform_openai_jsonl_content_to_vertex_ai_jsonl_content(
openai_input
)
)
vertex_input = _wrap_entries(openai_input)
# Verify both labels are GCP-safe and encoded raw preserves round-trip.
assert (
@ -1154,7 +1241,6 @@ class TestVertexBatchCustomIdLabels:
def test_custom_id_label_sanitization(self):
"""Test that custom_id values are sanitized to meet GCP label constraints"""
transformation = VertexAIJsonlFilesTransformation()
# Test sanitization function
assert _sanitize_gcp_label_value("MyRequest-1") == "myrequest-1"
@ -1179,11 +1265,7 @@ class TestVertexBatchCustomIdLabels:
}
]
vertex_input = (
transformation._transform_openai_jsonl_content_to_vertex_ai_jsonl_content(
openai_input
)
)
vertex_input = _wrap_entries(openai_input)
# Verify both labels are safe for GCP labels.
assert (
@ -1192,3 +1274,47 @@ class TestVertexBatchCustomIdLabels:
raw_label = vertex_input[0]["request"]["labels"]["litellm_custom_id_raw"]
assert raw_label != "MyRequest-1"
assert _sanitize_gcp_label_value(raw_label) == raw_label
class TestConfiguredBucketNameResolution:
def test_should_resolve_new_gcs_bucket_name_key(self, config, monkeypatch):
monkeypatch.delenv("GCS_BUCKET_NAME", raising=False)
assert (
config._get_configured_bucket_name({"gcs_bucket_name": "my-new-bucket"})
== "my-new-bucket"
)
def test_should_resolve_legacy_bucket_name_key(self, config, monkeypatch):
monkeypatch.delenv("GCS_BUCKET_NAME", raising=False)
assert (
config._get_configured_bucket_name({"bucket_name": "my-legacy-bucket"})
== "my-legacy-bucket"
)
def test_should_prefer_new_key_over_legacy(self, config, monkeypatch):
monkeypatch.delenv("GCS_BUCKET_NAME", raising=False)
assert (
config._get_configured_bucket_name(
{"gcs_bucket_name": "new", "bucket_name": "legacy"}
)
== "new"
)
def test_should_fall_back_to_env(self, config, monkeypatch):
monkeypatch.setenv("GCS_BUCKET_NAME", "env-bucket")
assert config._get_configured_bucket_name({}) == "env-bucket"
def test_should_raise_when_no_bucket_anywhere(self, config, monkeypatch):
monkeypatch.delenv("GCS_BUCKET_NAME", raising=False)
with pytest.raises(ValueError, match="GCS bucket_name is required"):
config._get_configured_bucket_name({})
def test_legacy_kwarg_survives_get_litellm_params(self):
from litellm.litellm_core_utils.get_litellm_params import (
OPTIONAL_KWARGS_KEYS,
get_litellm_params,
)
assert "bucket_name" in OPTIONAL_KWARGS_KEYS
params = get_litellm_params(bucket_name="my-legacy-bucket")
assert params.get("bucket_name") == "my-legacy-bucket"

View file

@ -0,0 +1,141 @@
"""Tests for the v1 -> v2 bridge.
`to_server_spec` maps the migrated modes (none + the static-header family, shared-key) and
defers everything else to v1 by returning None; `to_subject` maps the principal; `raise_public`
maps each CredError onto its HTTP status. These pin the parity-critical mapping before the graft.
"""
import base64
from types import SimpleNamespace
import pytest
from fastapi import HTTPException
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
raise_public,
to_server_spec,
to_subject,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
ApiKeyConfig,
CredError,
NoneConfig,
SharedKey,
)
from litellm.types.mcp import MCPAuth, MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
def _server(**kwargs) -> MCPServer:
return MCPServer(server_id="s", name="n", transport=MCPTransport.http, **kwargs)
def test_none_maps_to_none_config():
spec = to_server_spec(_server(auth_type=None))
assert spec is not None
assert isinstance(spec.config, NoneConfig)
def test_api_key_maps_to_x_api_key_shared():
spec = to_server_spec(_server(auth_type=MCPAuth.api_key, authentication_token="k"))
assert spec is not None and isinstance(spec.config, ApiKeyConfig)
assert spec.config.header_name == "X-API-Key"
assert spec.config.value_prefix == ""
assert isinstance(spec.config.key_source, SharedKey)
assert spec.config.key_source.value.get_secret_value() == "k"
@pytest.mark.parametrize(
"auth_type, prefix",
[
(MCPAuth.bearer_token, "Bearer"),
(MCPAuth.token, "token"),
(MCPAuth.authorization, ""),
],
)
def test_authorization_schemes_map_with_their_prefix(auth_type, prefix):
spec = to_server_spec(_server(auth_type=auth_type, authentication_token="t"))
assert spec is not None and isinstance(spec.config, ApiKeyConfig)
assert spec.config.header_name == "Authorization"
assert spec.config.value_prefix == prefix
assert spec.config.key_source.value.get_secret_value() == "t"
def test_basic_scheme_base64_encodes_the_token():
spec = to_server_spec(
_server(auth_type=MCPAuth.basic, authentication_token="user:pass")
)
assert spec is not None and isinstance(spec.config, ApiKeyConfig)
assert spec.config.value_prefix == "Basic"
expected = base64.b64encode(b"user:pass").decode()
assert spec.config.key_source.value.get_secret_value() == expected
@pytest.mark.parametrize(
"server",
[
_server(auth_type=MCPAuth.api_key), # no token configured
_server(auth_type=MCPAuth.bearer_token), # no token configured
_server(auth_type=MCPAuth.oauth2),
_server(auth_type=MCPAuth.oauth2_token_exchange),
_server(auth_type=MCPAuth.aws_sigv4),
_server(
auth_type=None, oauth_passthrough=True, extra_headers=["Authorization"]
),
],
)
def test_unmigrated_modes_defer_to_v1(server):
# A None spec is the defer signal; the caller falls back to v1.
assert to_server_spec(server) is None
@pytest.mark.parametrize(
"server",
[
_server(auth_type=MCPAuth.api_key, is_byok=True),
# BYOK rides on auth_type, so it must defer for every scheme, not just api_key. A stray
# static token must not route a BYOK server to a v2 shared-key spec with the wrong value.
_server(auth_type=MCPAuth.bearer_token, is_byok=True, authentication_token="x"),
_server(auth_type=MCPAuth.basic, is_byok=True, authentication_token="x"),
_server(
auth_type=MCPAuth.authorization, is_byok=True, authentication_token="x"
),
_server(auth_type=MCPAuth.token, is_byok=True, authentication_token="x"),
_server(auth_type=None, is_byok=True),
],
)
def test_byok_defers_regardless_of_auth_type(server):
assert to_server_spec(server) is None
def test_to_subject_unauthenticated_is_empty_with_inbound_token():
subject = to_subject(None, "inbound-jwt")
assert subject.tenant_id == ""
assert subject.subject_id == ""
assert subject.inbound_token is not None
assert subject.inbound_token.get_secret_value() == "inbound-jwt"
def test_to_subject_maps_principal_fields():
principal = SimpleNamespace(org_id="org1", team_id="team1", user_id="user1")
subject = to_subject(principal, None)
assert subject.tenant_id == "org1"
assert subject.subject_id == "user1"
assert subject.inbound_token is None
@pytest.mark.parametrize(
"error, status",
[
(CredError.of_unauthorized("x"), 401),
(CredError.of_misconfigured("x"), 500),
(CredError.of_upstream_unavailable("x"), 503),
(CredError.of_unsupported_mode("x"), 500),
(CredError.of_precondition_required("x"), 412),
(CredError.of_not_implemented("x"), 501),
],
)
def test_raise_public_maps_each_error_to_its_status(error, status):
with pytest.raises(HTTPException) as exc_info:
raise_public(error)
assert exc_info.value.status_code == status

View file

@ -1,57 +1,104 @@
"""Tests for the resolver dispatch skeleton.
"""Tests for the resolver dispatch: live arms produce auth, stubbed arms fail closed.
Every mode must reach its own arm and, until that arm is built, return a typed
`not_implemented` CredError rather than silently producing no credential. Parametrizing over
one config per mode also guards reachability: if a `case` were dropped, that mode would fall to
the `assert_never` tail and raise here instead of returning the stub.
`none` and `api_key` (shared-key source) are implemented; every other arm, plus the `api_key`
BYOK source, returns a typed `not_implemented` error until its mode lands. Parametrizing the
stubs over one config each also guards reachability: a dropped `case` would hit `assert_never`
and raise instead of returning the stub.
"""
import httpx
import pytest
from pydantic import SecretStr
from litellm.proxy._experimental.mcp_server.outbound_credentials import (
ApiKeyConfig,
AuthorizationCodeConfig,
AuthSpecKind,
AwsSigV4Config,
Byok,
ClientCredentialsConfig,
Error,
NoneConfig,
NoOpAuth,
Ok,
PassthroughConfig,
ServerSpec,
SharedKey,
StaticHeaderAuth,
Subject,
TokenExchangeConfig,
UpstreamCredentialProvider,
)
_ONE_CONFIG_PER_MODE = [
(AuthSpecKind.none, NoneConfig()),
(AuthSpecKind.api_key, ApiKeyConfig(key_source=SharedKey(value=SecretStr("k")))),
(AuthSpecKind.passthrough, PassthroughConfig()),
(AuthSpecKind.client_credentials, ClientCredentialsConfig()),
(AuthSpecKind.token_exchange, TokenExchangeConfig()),
(AuthSpecKind.authorization_code, AuthorizationCodeConfig()),
(AuthSpecKind.aws_sigv4, AwsSigV4Config(region="us-east-1")),
_SUBJECT = Subject(tenant_id="", subject_id="")
def _spec(config):
return ServerSpec(
server_id="s", resource="https://upstream.example.com", config=config
)
def _emitted(auth: httpx.Auth) -> httpx.Headers:
request = httpx.Request("GET", "https://upstream.example.com/mcp")
flow = auth.auth_flow(request)
next(flow)
flow.close()
return request.headers
@pytest.mark.asyncio
async def test_none_mode_yields_a_no_op_auth():
result = await UpstreamCredentialProvider().resolve_credentials(
_SUBJECT, _spec(NoneConfig())
)
assert isinstance(result, Ok)
assert isinstance(result.ok, NoOpAuth)
@pytest.mark.asyncio
async def test_api_key_shared_emits_the_configured_header():
config = ApiKeyConfig(
header_name="X-API-Key",
value_prefix="",
key_source=SharedKey(value=SecretStr("secret-key")),
)
result = await UpstreamCredentialProvider().resolve_credentials(
_SUBJECT, _spec(config)
)
assert isinstance(result, Ok)
assert isinstance(result.ok, StaticHeaderAuth)
assert _emitted(result.ok)["X-API-Key"] == "secret-key"
@pytest.mark.asyncio
async def test_api_key_shared_honors_authorization_scheme():
config = ApiKeyConfig(
header_name="Authorization",
value_prefix="Bearer",
key_source=SharedKey(value=SecretStr("tok")),
)
result = await UpstreamCredentialProvider().resolve_credentials(
_SUBJECT, _spec(config)
)
assert isinstance(result, Ok)
assert _emitted(result.ok)["Authorization"] == "Bearer tok"
_STUBBED = [
("api_key_byok", ApiKeyConfig(key_source=Byok())),
("passthrough", PassthroughConfig()),
("client_credentials", ClientCredentialsConfig()),
("token_exchange", TokenExchangeConfig()),
("authorization_code", AuthorizationCodeConfig()),
("aws_sigv4", AwsSigV4Config(region="us-east-1")),
]
@pytest.mark.asyncio
@pytest.mark.parametrize("kind, config", _ONE_CONFIG_PER_MODE)
async def test_every_mode_reaches_its_arm_and_returns_not_implemented(kind, config):
spec = ServerSpec(
server_id="s", resource="https://upstream.example.com", config=config
@pytest.mark.parametrize("label, config", _STUBBED)
async def test_unbuilt_arms_fail_closed_with_not_implemented(label, config):
result = await UpstreamCredentialProvider().resolve_credentials(
_SUBJECT, _spec(config)
)
subject = Subject(tenant_id="", subject_id="")
result = await UpstreamCredentialProvider().resolve_credentials(subject, spec)
assert isinstance(result, Error)
assert result.error.tag == "not_implemented"
assert kind.value in result.error.summary
def test_all_seven_modes_are_covered():
# Guards that the parametrization (and therefore the dispatch) spans every AuthSpecKind, so a
# newly added mode without a test row is caught here rather than slipping through.
assert {kind for kind, _ in _ONE_CONFIG_PER_MODE} == set(AuthSpecKind)

View file

@ -4568,5 +4568,173 @@ class TestGetPublicMCPServersLegacyMode:
assert sorted(s.server_id for s in result) == ["a", "b"]
class TestCreateMcpClientV2Graft:
"""The PR4 v2-resolver graft in ``_create_mcp_client``.
Migrated HTTP/SSE modes (``none`` plus the static ``api_key`` family) resolve through the
injected ``UpstreamCredentialProvider`` into the ``resolved_auth`` slot; every other mode,
and every stdio server, defers to v1's ``auth_value`` path unchanged.
"""
def _http_server(self, **overrides: Any) -> MCPServer:
base: Dict[str, Any] = dict(
server_id="http-graft",
name="graft_server",
url="https://upstream.example.com/mcp",
transport=MCPTransport.http,
)
base.update(overrides)
return MCPServer(**base)
async def test_none_mode_resolves_to_noop_auth(self):
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
NoOpAuth,
)
client = await MCPServerManager()._create_mcp_client(
self._http_server(auth_type=None)
)
assert isinstance(client._resolved_auth, NoOpAuth)
assert client._mcp_auth_value is None
@pytest.mark.parametrize(
"auth_type, token, expected_name, expected_value",
[
(MCPAuth.api_key, "k-123", "X-API-Key", "k-123"),
(MCPAuth.bearer_token, "b-123", "Authorization", "Bearer b-123"),
(MCPAuth.token, "t-123", "Authorization", "token t-123"),
(MCPAuth.authorization, "raw-123", "Authorization", "raw-123"),
],
)
async def test_static_family_emits_expected_header(
self, auth_type, token, expected_name, expected_value
):
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
StaticHeaderAuth,
)
client = await MCPServerManager()._create_mcp_client(
self._http_server(auth_type=auth_type, authentication_token=token)
)
assert isinstance(client._resolved_auth, StaticHeaderAuth)
assert client._resolved_auth.header_name == expected_name
assert client._resolved_auth._header_value.get_secret_value() == expected_value
assert client._mcp_auth_value is None
async def test_basic_mode_base64_encodes(self):
import base64
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
StaticHeaderAuth,
)
client = await MCPServerManager()._create_mcp_client(
self._http_server(auth_type=MCPAuth.basic, authentication_token="user:pass")
)
encoded = base64.b64encode(b"user:pass").decode()
assert isinstance(client._resolved_auth, StaticHeaderAuth)
assert client._resolved_auth.header_name == "Authorization"
assert (
client._resolved_auth._header_value.get_secret_value() == f"Basic {encoded}"
)
async def test_deferred_mode_uses_v1_auth_value(self):
client = await MCPServerManager()._create_mcp_client(
self._http_server(
auth_type=MCPAuth.oauth2, authentication_token="legacy-token"
)
)
assert client._resolved_auth is None
assert client._mcp_auth_value == "legacy-token"
async def test_static_token_missing_defers_to_v1(self):
client = await MCPServerManager()._create_mcp_client(
self._http_server(auth_type=MCPAuth.api_key, authentication_token=None)
)
assert client._resolved_auth is None
async def test_stdio_migrated_auth_type_still_defers_to_v1(self):
client = await MCPServerManager()._create_mcp_client(
MCPServer(
server_id="stdio-graft",
name="stdio_graft",
transport=MCPTransport.stdio,
command="node",
args=["server.js"],
auth_type=MCPAuth.api_key,
authentication_token="k-stdio",
)
)
assert client.transport_type == MCPTransport.stdio
assert client._resolved_auth is None
assert client._mcp_auth_value == "k-stdio"
async def test_resolver_error_maps_to_http_exception(self):
from litellm.proxy._experimental.mcp_server.outbound_credentials import Error
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
CredError,
)
class _UnauthorizedProvider:
async def resolve_credentials(self, subject, server):
return Error(CredError.of_unauthorized("denied"))
manager = MCPServerManager(cred_provider=_UnauthorizedProvider())
with pytest.raises(HTTPException) as exc:
await manager._create_mcp_client(self._http_server(auth_type=None))
assert exc.value.status_code == 401
async def test_per_request_override_defers_to_v1(self):
# A per-request override (mcp_auth_header) must win over the shared static token,
# exactly as v1 did, so a migrated static server defers to v1 when one is present.
client = await MCPServerManager()._create_mcp_client(
self._http_server(
auth_type=MCPAuth.bearer_token, authentication_token="shared-tok"
),
mcp_auth_header="caller-override",
)
assert client._resolved_auth is None
assert client._mcp_auth_value == "caller-override"
async def test_conflicting_extra_header_skips_resolved_auth_on_v2(self):
# An Authorization already supplied via extra_headers (guardrail hook like the JWT
# signer, static_headers, or a forwarded caller header) must win. The server stays on
# the v2 path but skips resolved_auth, so nothing overwrites the inbound header.
client = await MCPServerManager()._create_mcp_client(
self._http_server(
auth_type=MCPAuth.bearer_token, authentication_token="shared-tok"
),
extra_headers={"Authorization": "Bearer hook-jwt"},
)
assert client._resolved_auth is None
assert client._mcp_auth_value is None
assert client._get_auth_headers()["Authorization"] == "Bearer hook-jwt"
async def test_none_with_extra_header_stays_v2_without_clobbering(self):
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
NoOpAuth,
)
# none resolves to NoOpAuth, which writes no header, so it cannot clobber an inbound
# Authorization; it stays on the v2 path and the inbound header is preserved verbatim.
client = await MCPServerManager()._create_mcp_client(
self._http_server(auth_type=None),
extra_headers={"Authorization": "Bearer hook-jwt"},
)
assert isinstance(client._resolved_auth, NoOpAuth)
assert client._get_auth_headers()["Authorization"] == "Bearer hook-jwt"
if __name__ == "__main__":
pytest.main([__file__])

View file

@ -14,155 +14,131 @@ from fastapi import HTTPException
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
def _models(file_content_as_dict):
"""Distinct body.model values, mirroring how the rate limiter collects the
models from a streamed batch file before the access check."""
return [
entry["body"]["model"]
for entry in file_content_as_dict
if (entry.get("body") or {}).get("model")
]
# ---------------------------------------------------------------------------
# Token counter — covers all three batch payload shapes
# ---------------------------------------------------------------------------
def test_token_counter_counts_chat_messages():
from litellm.batches.batch_utils import _get_batch_job_input_file_usage
from litellm.batches.batch_utils import _count_entry_tokens
usage = _get_batch_job_input_file_usage(
file_content_dictionary=[
{
"body": {
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "hello"}],
}
tokens = _count_entry_tokens(
{
"body": {
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "hello"}],
}
]
}
)
assert usage.prompt_tokens > 0
assert tokens > 0
def test_token_counter_counts_text_completion_prompt():
"""Pre-fix this returned 0 tokens (the function only inspected
"""Pre-fix this returned 0 tokens (the counter only inspected
`messages`), letting `prompt`-style batches slip past TPM limits."""
from litellm.batches.batch_utils import _get_batch_job_input_file_usage
from litellm.batches.batch_utils import _count_entry_tokens
usage = _get_batch_job_input_file_usage(
file_content_dictionary=[
{"body": {"model": "gpt-3.5-turbo-instruct", "prompt": "hello world"}}
]
tokens = _count_entry_tokens(
{"body": {"model": "gpt-3.5-turbo-instruct", "prompt": "hello world"}}
)
assert usage.prompt_tokens > 0
assert tokens > 0
def test_token_counter_counts_embedding_input_string():
from litellm.batches.batch_utils import _get_batch_job_input_file_usage
from litellm.batches.batch_utils import _count_entry_tokens
usage = _get_batch_job_input_file_usage(
file_content_dictionary=[
{"body": {"model": "text-embedding-3-small", "input": "hello world"}}
]
tokens = _count_entry_tokens(
{"body": {"model": "text-embedding-3-small", "input": "hello world"}}
)
assert usage.prompt_tokens > 0
assert tokens > 0
def test_token_counter_counts_embedding_input_list():
from litellm.batches.batch_utils import _get_batch_job_input_file_usage
from litellm.batches.batch_utils import _count_entry_tokens
usage = _get_batch_job_input_file_usage(
file_content_dictionary=[
{
"body": {
"model": "text-embedding-3-small",
"input": ["hello", "world"],
}
tokens = _count_entry_tokens(
{
"body": {
"model": "text-embedding-3-small",
"input": ["hello", "world"],
}
]
}
)
assert usage.prompt_tokens > 0
assert tokens > 0
def test_token_counter_counts_text_completion_prompt_list():
from litellm.batches.batch_utils import _get_batch_job_input_file_usage
from litellm.batches.batch_utils import _count_entry_tokens
usage = _get_batch_job_input_file_usage(
file_content_dictionary=[
{
"body": {
"model": "gpt-3.5-turbo-instruct",
"prompt": ["alpha", "beta"],
}
tokens = _count_entry_tokens(
{
"body": {
"model": "gpt-3.5-turbo-instruct",
"prompt": ["alpha", "beta"],
}
]
}
)
assert usage.prompt_tokens > 0
assert tokens > 0
def test_token_counter_counts_pre_tokenized_prompt_int_list():
"""OpenAI's text-completion API accepts a single pre-tokenized prompt as
a list of ints. Each int is one token; pre-fix this shape was silently
counted as zero, leaving a TPM bypass."""
from litellm.batches.batch_utils import _get_batch_job_input_file_usage
from litellm.batches.batch_utils import _count_entry_tokens
usage = _get_batch_job_input_file_usage(
file_content_dictionary=[
{
"body": {
"model": "gpt-3.5-turbo-instruct",
"prompt": [1, 2, 3, 4, 5],
}
tokens = _count_entry_tokens(
{
"body": {
"model": "gpt-3.5-turbo-instruct",
"prompt": [1, 2, 3, 4, 5],
}
]
}
)
assert usage.prompt_tokens == 5
assert tokens == 5
def test_token_counter_counts_pre_tokenized_prompt_list_of_int_lists():
"""Multiple pre-tokenized prompts (`list[list[int]]`) — the most
important bypass shape. A 1000-token batch must report 1000 tokens,
not zero."""
from litellm.batches.batch_utils import _get_batch_job_input_file_usage
from litellm.batches.batch_utils import _count_entry_tokens
usage = _get_batch_job_input_file_usage(
file_content_dictionary=[
{
"body": {
"model": "gpt-3.5-turbo-instruct",
"prompt": [[1] * 250, [2] * 250, [3] * 500],
}
tokens = _count_entry_tokens(
{
"body": {
"model": "gpt-3.5-turbo-instruct",
"prompt": [[1] * 250, [2] * 250, [3] * 500],
}
]
}
)
assert usage.prompt_tokens == 1000
assert tokens == 1000
def test_token_counter_counts_pre_tokenized_input_for_embeddings():
"""Same shape applies to embeddings (`input`)."""
from litellm.batches.batch_utils import _get_batch_job_input_file_usage
from litellm.batches.batch_utils import _count_entry_tokens
usage = _get_batch_job_input_file_usage(
file_content_dictionary=[
{
"body": {
"model": "text-embedding-3-small",
"input": [[1, 2, 3], [4, 5, 6]],
}
tokens = _count_entry_tokens(
{
"body": {
"model": "text-embedding-3-small",
"input": [[1, 2, 3], [4, 5, 6]],
}
]
}
)
assert usage.prompt_tokens == 6
# ---------------------------------------------------------------------------
# Model extractor
# ---------------------------------------------------------------------------
def test_model_extractor_returns_distinct_models():
from litellm.batches.batch_utils import _get_models_from_batch_input_file_content
models = _get_models_from_batch_input_file_content(
[
{"body": {"model": "gpt-4o", "messages": []}},
{"body": {"model": "gpt-4o", "messages": []}}, # duplicate
{"body": {"model": "gpt-4o-mini", "messages": []}},
{"body": {}}, # missing model
]
)
assert models == ["gpt-4o", "gpt-4o-mini"]
assert tokens == 6
# ---------------------------------------------------------------------------
@ -211,7 +187,7 @@ async def test_pre_call_rejects_unauthorized_model_in_batch_file():
with pytest.raises(HTTPException) as exc:
await rate_limiter._enforce_batch_file_model_access(
user_api_key_dict=user,
file_content_as_dict=file_dict,
models=_models(file_dict),
)
assert exc.value.status_code == 403
@ -250,7 +226,7 @@ async def test_pre_call_allows_all_team_models_key_when_model_in_team_allowlist(
with patch("litellm.proxy.proxy_server.llm_router", None):
await rate_limiter._enforce_batch_file_model_access(
user_api_key_dict=user,
file_content_as_dict=file_dict,
models=_models(file_dict),
)
@ -297,7 +273,7 @@ async def test_pre_call_uses_current_team_allowlist_for_all_team_models_key():
):
await rate_limiter._enforce_batch_file_model_access(
user_api_key_dict=user,
file_content_as_dict=file_dict,
models=_models(file_dict),
)
assert exc_info.value.status_code == 403
@ -358,7 +334,7 @@ async def test_pre_call_allows_all_team_models_key_via_current_team_object():
):
await rate_limiter._enforce_batch_file_model_access(
user_api_key_dict=user,
file_content_as_dict=file_dict,
models=_models(file_dict),
)
mock_get_team_object.assert_awaited_once()
@ -421,7 +397,7 @@ async def test_pre_call_denies_all_team_models_key_via_member_scope():
):
await rate_limiter._enforce_batch_file_model_access(
user_api_key_dict=user,
file_content_as_dict=file_dict,
models=_models(file_dict),
)
assert exc_info.value.status_code == 403
@ -479,7 +455,7 @@ async def test_pre_call_fails_closed_when_current_team_fetch_fails_for_all_team_
):
await rate_limiter._enforce_batch_file_model_access(
user_api_key_dict=user,
file_content_as_dict=file_dict,
models=_models(file_dict),
)
assert exc_info.value.status_code == expected_status
@ -524,7 +500,7 @@ async def test_pre_call_allows_authorized_model_in_batch_file():
# Should not raise
await rate_limiter._enforce_batch_file_model_access(
user_api_key_dict=user,
file_content_as_dict=file_dict,
models=_models(file_dict),
)
@ -744,7 +720,7 @@ async def test_pre_call_allows_stripped_provider_model_when_key_has_proxy_alias(
):
await rate_limiter._enforce_batch_file_model_access(
user_api_key_dict=user,
file_content_as_dict=file_dict,
models=_models(file_dict),
target_model_names=[proxy_alias],
)
@ -837,7 +813,7 @@ async def test_pre_call_uses_target_model_names_not_stripped_reverse_lookup(
):
await rate_limiter._enforce_batch_file_model_access(
user_api_key_dict=user,
file_content_as_dict=file_dict,
models=_models(file_dict),
target_model_names=[batch_alias],
)
@ -863,11 +839,11 @@ async def test_pre_call_skips_check_when_no_models_present():
# entirely.
await rate_limiter._enforce_batch_file_model_access(
user_api_key_dict=user,
file_content_as_dict=[],
models=_models([]),
)
await rate_limiter._enforce_batch_file_model_access(
user_api_key_dict=user,
file_content_as_dict=[{"body": {}}],
models=_models([{"body": {}}]),
)
@ -1390,3 +1366,272 @@ async def test_count_input_file_usage_raises_on_non_bytes_content():
user_api_key_dict=UserAPIKeyAuth(api_key="sk", models=["*"]),
data={},
)
# Streaming input counting — peak memory must not scale with a full dict list
# ---------------------------------------------------------------------------
def _make_batch_input_bytes(n_rows: int, padding: int = 200) -> bytes:
import json as _json
pad = "x" * padding
rows = []
for i in range(n_rows):
rows.append(
_json.dumps(
{
"custom_id": f"request-{i}",
"method": "POST",
"url": "/v1/chat/completions",
"body": {
"model": "gpt-4o" if i % 2 else "gpt-3.5-turbo",
"messages": [{"role": "user", "content": f"{pad} {i}"}],
},
}
)
)
return ("\n".join(rows)).encode("utf-8")
def test_iter_batch_input_entries_matches_dict_list():
from litellm.batches.batch_utils import (
_get_file_content_as_dictionary,
_iter_batch_input_entries,
)
raw = _make_batch_input_bytes(50)
streamed = list(_iter_batch_input_entries(raw))
assert streamed == _get_file_content_as_dictionary(raw)
assert streamed[0]["custom_id"] == "request-0"
# tolerant of blank lines and a missing trailing newline
assert list(_iter_batch_input_entries(raw + b"\n\n")) == streamed
def test_streaming_count_peak_below_dict_list():
import gc
import tracemalloc
from litellm.batches.batch_utils import (
_get_file_content_as_dictionary,
_iter_batch_input_entries,
)
raw = _make_batch_input_bytes(8000)
def _measure(fn):
gc.collect()
tracemalloc.start()
try:
fn()
_, peak = tracemalloc.get_traced_memory()
finally:
tracemalloc.stop()
return peak
def _stream():
count = 0
models: set = set()
for entry in _iter_batch_input_entries(raw):
count += 1
model = (entry.get("body") or {}).get("model")
if model:
models.add(model)
return count
def _build_list():
return len(_get_file_content_as_dictionary(raw))
stream_peak = _measure(_stream)
list_peak = _measure(_build_list)
assert stream_peak < list_peak * 0.5, (
f"streaming count peak {stream_peak} is not a clear win over the dict "
f"list {list_peak} (ratio {stream_peak / list_peak:.2f})"
)
@pytest.mark.asyncio
async def test_count_input_file_usage_streams_without_building_list():
"""count_input_file_usage must count requests/tokens in one streaming pass.
Mocks the download; asserts the count is correct and that the dict-list
helper is never called (a revert to the list approach would call it)."""
from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter
rate_limiter = _PROXY_BatchRateLimiter(
internal_usage_cache=MagicMock(),
parallel_request_limiter=MagicMock(),
)
raw = _make_batch_input_bytes(10)
fake_content = MagicMock()
fake_content.content = raw
with (
patch("litellm.afile_content", new=AsyncMock(return_value=fake_content)),
patch(
"litellm.batches.batch_utils._get_file_content_as_dictionary"
) as mock_dict_list,
):
usage = await rate_limiter.count_input_file_usage(
file_id="file-not-managed",
custom_llm_provider="openai",
user_api_key_dict=None,
)
assert usage.request_count == 10
assert usage.total_tokens > 0
mock_dict_list.assert_not_called()
def _one_row_batch_bytes(model: str) -> bytes:
import json as _json
return (
_json.dumps(
{
"custom_id": "r0",
"method": "POST",
"url": "/v1/chat/completions",
"body": {
"model": model,
"messages": [{"role": "user", "content": "x"}],
},
}
)
+ "\n"
).encode("utf-8")
@pytest.mark.asyncio
async def test_count_input_file_usage_enforces_models_when_token_counting_fails():
"""Security regression: a row whose content makes token counting raise must
NOT skip the model allowlist check. async_pre_call_hook swallows non-HTTP
exceptions and submits the batch, so a raised counting error would otherwise
fail open. The access check must still run and deny the restricted model."""
from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter
rate_limiter = _PROXY_BatchRateLimiter(
internal_usage_cache=MagicMock(),
parallel_request_limiter=MagicMock(),
)
fake_content = MagicMock()
fake_content.content = _one_row_batch_bytes("restricted-model")
user = UserAPIKeyAuth(
api_key="sk-x",
user_id="bob",
models=["only-allowed"],
user_role=LitellmUserRoles.INTERNAL_USER.value,
)
def _boom(*args, **kwargs):
raise ValueError("unsupported content part: input_audio")
deny = AsyncMock(side_effect=Exception("model not in allowlist"))
with (
patch("litellm.afile_content", new=AsyncMock(return_value=fake_content)),
patch("litellm.proxy.hooks.batch_rate_limiter._count_entry_tokens", new=_boom),
patch("litellm.proxy.auth.auth_checks.can_key_call_model", new=deny),
patch("litellm.proxy.proxy_server.llm_router", MagicMock(model_list=[])),
):
with pytest.raises(HTTPException) as exc:
await rate_limiter.count_input_file_usage(
file_id="file-not-managed",
custom_llm_provider="openai",
user_api_key_dict=user,
)
# The access check ran despite token counting failing, and denied the model.
deny.assert_awaited()
assert exc.value.status_code == 403
@pytest.mark.asyncio
async def test_count_input_file_usage_estimates_tokens_when_counting_fails_for_allowed_model():
"""A token-counting failure for an allowed model must not hard-block the batch
(the pre-streaming behavior let such batches through), but it also must not
zero the token total, which would let a caller evade the TPM limit by sending
rows the counter cannot measure. The row falls back to a conservative
size-based estimate so the batch proceeds with a non-zero count."""
from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter
rate_limiter = _PROXY_BatchRateLimiter(
internal_usage_cache=MagicMock(),
parallel_request_limiter=MagicMock(),
)
fake_content = MagicMock()
fake_content.content = _one_row_batch_bytes("allowed-model")
user = UserAPIKeyAuth(
api_key="sk-x",
user_id="bob",
models=["allowed-model"],
user_role=LitellmUserRoles.INTERNAL_USER.value,
)
def _boom(*args, **kwargs):
raise ValueError("unsupported content part: file")
allow = AsyncMock(return_value=True)
with (
patch("litellm.afile_content", new=AsyncMock(return_value=fake_content)),
patch("litellm.proxy.hooks.batch_rate_limiter._count_entry_tokens", new=_boom),
patch("litellm.proxy.auth.auth_checks.can_key_call_model", new=allow),
patch("litellm.proxy.proxy_server.llm_router", MagicMock(model_list=[])),
):
usage = await rate_limiter.count_input_file_usage(
file_id="file-not-managed",
custom_llm_provider="openai",
user_api_key_dict=user,
)
allow.assert_awaited()
assert usage.request_count == 1
# Estimated, not zeroed: a crafted uncountable row can't evade the TPM limit.
assert usage.total_tokens > 0
@pytest.mark.asyncio
async def test_count_input_file_usage_collects_models_after_malformed_line():
"""A malformed JSONL line must not abort model collection. A restricted model
named on a row AFTER a malformed line must still be collected and denied by the
allowlist check, otherwise a caller could hide a restricted model behind a bad
row."""
from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter
rate_limiter = _PROXY_BatchRateLimiter(
internal_usage_cache=MagicMock(),
parallel_request_limiter=MagicMock(),
)
fake_content = MagicMock()
fake_content.content = (
_one_row_batch_bytes("only-allowed")
+ b"{ this is not valid json\n"
+ _one_row_batch_bytes("restricted-model")
)
user = UserAPIKeyAuth(
api_key="sk-x",
user_id="bob",
models=["only-allowed"],
user_role=LitellmUserRoles.INTERNAL_USER.value,
)
async def _deny_restricted(model, **kwargs):
if model == "restricted-model":
raise Exception("model not in allowlist")
return True
deny = AsyncMock(side_effect=_deny_restricted)
with (
patch("litellm.afile_content", new=AsyncMock(return_value=fake_content)),
patch("litellm.proxy.auth.auth_checks.can_key_call_model", new=deny),
patch("litellm.proxy.proxy_server.llm_router", MagicMock(model_list=[])),
):
with pytest.raises(HTTPException) as exc:
await rate_limiter.count_input_file_usage(
file_id="file-not-managed",
custom_llm_provider="openai",
user_api_key_dict=user,
)
assert exc.value.status_code == 403

View file

@ -0,0 +1,195 @@
"""Unit tests for POST /v1/callbacks/logs (replay logging payloads → callbacks)."""
import time
import pytest
from fastapi import HTTPException
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.logging_endpoints.callback_logs_endpoints import (
CallbackLogsReplayer,
ingest_callback_logs,
)
from litellm.types.proxy.callback_logs_endpoints import (
CallbackLogRecord,
CallbackLogsRequest,
)
REQ_ID = "cb-logs-unit-test-1"
def _sample_payload(**overrides):
payload = {
"id": REQ_ID,
"litellm_call_id": REQ_ID,
"call_type": "acompletion",
"stream": False,
"response_cost": 0.0123,
"custom_llm_provider": "openai",
"total_tokens": 42,
"prompt_tokens": 30,
"completion_tokens": 12,
"startTime": time.time() - 2,
"endTime": time.time(),
"model": "gpt-4o-mini",
"metadata": {
"user_api_key_hash": "rust-gateway-test-key",
"user_api_key_user_id": "user-cb-logs-test",
"user_api_key_team_id": "team-cb-logs-test",
},
"messages": [{"role": "user", "content": "hi"}],
}
payload.update(overrides)
return payload
def test_epoch_to_datetime_handles_float_and_fallback():
dt = CallbackLogsReplayer._epoch_to_datetime(1_700_000_000.5)
assert dt.year == 2023
# Non-numeric input must not raise — falls back to "now".
assert CallbackLogsReplayer._epoch_to_datetime(None) is not None
def test_build_logging_obj_seeds_model_call_details():
obj = CallbackLogsReplayer._build_logging_obj(_sample_payload())
details = obj.model_call_details
# Prebuilt payload is set so the handler skips rebuilding it.
assert details["standard_logging_object"]["id"] == REQ_ID
assert details["response_cost"] == 0.0123
assert details["call_type"] == "acompletion"
# Metadata is mapped to the keys the cost-tracking callback reads.
md = details["litellm_params"]["metadata"]
assert md["user_api_key"] == "rust-gateway-test-key"
assert md["user_api_key_user_id"] == "user-cb-logs-test"
assert md["user_api_key_team_id"] == "team-cb-logs-test"
def test_response_obj_carries_usage():
obj = CallbackLogsReplayer._response_obj_from_payload(_sample_payload())
assert obj["usage"]["total_tokens"] == 42
assert obj["usage"]["prompt_tokens"] == 30
assert obj["usage"]["completion_tokens"] == 12
@pytest.mark.asyncio
async def test_success_record_invokes_success_handler(monkeypatch):
captured = {}
async def fake_success(self, result=None, start_time=None, end_time=None, **kwargs):
captured["standard_logging_object"] = self.model_call_details.get(
"standard_logging_object"
)
captured["result"] = result
monkeypatch.setattr(LiteLLMLogging, "async_success_handler", fake_success)
body = CallbackLogsRequest(
records=[
CallbackLogRecord(
status="success", standard_logging_payload=_sample_payload()
)
]
)
resp = await ingest_callback_logs(
body, user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
)
assert resp.processed == 1 and resp.failed == 0
assert captured["standard_logging_object"]["id"] == REQ_ID
assert captured["result"]["usage"]["total_tokens"] == 42
@pytest.mark.asyncio
async def test_failure_record_invokes_failure_handler(monkeypatch):
captured = {}
async def fake_failure(
self, exception, traceback_exception, start_time=None, end_time=None
):
captured["exception"] = str(exception)
monkeypatch.setattr(LiteLLMLogging, "async_failure_handler", fake_failure)
body = CallbackLogsRequest(
records=[
CallbackLogRecord(
status="failure",
standard_logging_payload=_sample_payload(),
error="upstream exploded",
)
]
)
resp = await ingest_callback_logs(
body, user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
)
assert resp.processed == 1 and resp.failed == 0
assert captured["exception"] == "upstream exploded"
@pytest.mark.asyncio
async def test_non_admin_is_rejected(monkeypatch):
async def fake_success(self, **kwargs):
return None
monkeypatch.setattr(LiteLLMLogging, "async_success_handler", fake_success)
body = CallbackLogsRequest(
records=[
CallbackLogRecord(
status="success", standard_logging_payload=_sample_payload()
)
]
)
with pytest.raises(HTTPException) as exc_info:
await ingest_callback_logs(
body,
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER),
)
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
async def test_one_bad_record_does_not_sink_the_batch(monkeypatch):
calls = {"n": 0}
async def flaky_success(
self, result=None, start_time=None, end_time=None, **kwargs
):
calls["n"] += 1
if calls["n"] == 1:
raise ValueError("boom on first record")
monkeypatch.setattr(LiteLLMLogging, "async_success_handler", flaky_success)
body = CallbackLogsRequest(
records=[
CallbackLogRecord(
status="success", standard_logging_payload=_sample_payload()
),
CallbackLogRecord(
status="success", standard_logging_payload=_sample_payload()
),
]
)
resp = await ingest_callback_logs(
body, user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
)
assert resp.processed == 1 and resp.failed == 1
# The failed record is reported back by index + error, not silently dropped.
assert len(resp.failures) == 1
assert resp.failures[0].index == 0
assert "boom on first record" in resp.failures[0].error
def test_batch_over_limit_is_rejected():
from litellm.constants import MAX_CALLBACK_LOG_RECORDS
from pydantic import ValidationError
# One over the cap must fail validation (422 at the API boundary), bounding
# the callback/DB fan-out a single POST can trigger.
too_many = [
CallbackLogRecord(status="success", standard_logging_payload=_sample_payload())
for _ in range(MAX_CALLBACK_LOG_RECORDS + 1)
]
with pytest.raises(ValidationError):
CallbackLogsRequest(records=too_many)

View file

@ -351,6 +351,110 @@ async def test_validate_no_team_non_global_server_raises(
assert "not in a team" in str(exc_info.value.detail)
@pytest.mark.asyncio
@patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
new=_make_mock_mcp_manager("private-server"),
)
@patch(
"litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids",
return_value=set(),
)
@patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_no_team_proxy_admin_can_assign_private_server(
mock_access_groups, mock_allow_all
):
"""Proxy admin assigning a non-global server to a teamless key — should pass (LIT-3815)."""
result = await validate_key_mcp_servers_against_team(
object_permission={"mcp_servers": ["private-server"]},
team_obj=None,
is_proxy_admin=True,
)
assert result["mcp_servers"] == ["private-server"]
@pytest.mark.asyncio
@patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
new=_make_mock_mcp_manager("private-server"),
)
@patch(
"litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids",
return_value=set(),
)
@patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_no_team_non_admin_private_server_still_raises(
mock_access_groups, mock_allow_all
):
"""The teamless override is gated on proxy admin — a non-admin still gets 403."""
with pytest.raises(HTTPException) as exc_info:
await validate_key_mcp_servers_against_team(
object_permission={"mcp_servers": ["private-server"]},
team_obj=None,
is_proxy_admin=False,
)
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
@patch(
"litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids",
return_value=set(),
)
@patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_no_team_proxy_admin_can_assign_access_group(
mock_access_groups, mock_allow_all
):
"""Proxy admin assigning an access group to a teamless key — should pass (LIT-3815)."""
result = await validate_key_mcp_servers_against_team(
object_permission={"mcp_access_groups": ["group-1"]},
team_obj=None,
is_proxy_admin=True,
)
assert result["mcp_access_groups"] == ["group-1"]
@pytest.mark.asyncio
@patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
new=_make_mock_mcp_manager("server-1", "server-outside"),
)
@patch(
"litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids",
return_value=set(),
)
@patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_proxy_admin_still_bounded_by_team_scope(
mock_access_groups, mock_allow_all
):
"""The override is scoped to teamless keys — an admin assigning beyond a team's scope still raises."""
team_obj = _make_team_obj(mcp_servers=["server-1"])
with pytest.raises(HTTPException) as exc_info:
await validate_key_mcp_servers_against_team(
object_permission={"mcp_servers": ["server-1", "server-outside"]},
team_obj=team_obj,
is_proxy_admin=True,
)
assert exc_info.value.status_code == 403
assert "server-outside" in str(exc_info.value.detail)
@pytest.mark.asyncio
@patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",

View file

@ -398,6 +398,83 @@ def test_mock_create_audio_file(mocker: MockerFixture, monkeypatch, llm_router:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
def test_create_file_batch_streams_from_upload_spool(monkeypatch, llm_router: Router):
"""
Batch uploads must be passed downstream as the upload's streamable file handle
(Starlette's already-spooled file), not read into an in-memory bytes object, so
the proxy never buffers the whole payload. Non-batch uploads keep the bytes path.
"""
import litellm.proxy.proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.openai_files_endpoints import files_endpoints as fe
from litellm.types.llms.openai import OpenAIFileObject
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", llm_router)
setup_proxy_logging_object(monkeypatch, llm_router)
captured: dict = {}
async def fake_route_create_file(*, _create_file_request, **kwargs):
file_elem = _create_file_request["file"][1]
captured["file_elem"] = file_elem
if hasattr(file_elem, "read") and hasattr(file_elem, "seek"):
file_elem.seek(0)
captured["streamed_content"] = file_elem.read()
return OpenAIFileObject(
id="dummy-id",
object="file",
bytes=0,
created_at=1234567890,
filename="batch.jsonl",
purpose="batch",
status="uploaded",
)
monkeypatch.setattr(fe, "route_create_file", fake_route_create_file)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test-user"
)
content = (
b'{"custom_id":"r-0","method":"POST","url":"/v1/chat/completions",'
b'"body":{"model":"gpt-3.5-turbo","messages":[{"role":"user","content":"hi"}]}}\n'
)
try:
resp = client.post(
"/v1/files",
files={"file": ("batch.jsonl", content, "application/jsonl")},
data={"purpose": "batch"},
headers={"Authorization": "Bearer test-key"},
)
assert resp.status_code == 200, resp.text
file_elem = captured["file_elem"]
assert not isinstance(
file_elem, (bytes, bytearray)
), "batch upload must be a streamable handle, not in-memory bytes"
assert hasattr(file_elem, "read") and hasattr(
file_elem, "seek"
), "batch upload must be a seekable file handle"
assert (
captured["streamed_content"] == content
), "the handle must stream the uploaded bytes"
captured.clear()
resp = client.post(
"/v1/files",
files={"file": ("data.jsonl", content, "application/jsonl")},
data={"purpose": "user_data"},
headers={"Authorization": "Bearer test-key"},
)
assert resp.status_code == 200, resp.text
assert isinstance(
captured["file_elem"], (bytes, bytearray)
), "non-batch upload must stay in-memory bytes"
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.flaky(retries=3, delay=2)
def test_target_storage_invokes_storage_backend(
mocker: MockerFixture, monkeypatch, llm_router: Router

Some files were not shown because too many files have changed in this diff Show more