diff --git a/.github/workflows/test-unit-proxy-infra.yml b/.github/workflows/test-unit-proxy-infra.yml index 336e53ee3d7..2681400e8b7 100644 --- a/.github/workflows/test-unit-proxy-infra.yml +++ b/.github/workflows/test-unit-proxy-infra.yml @@ -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 diff --git a/.gitignore b/.gitignore index fda3311fe02..3563d7c8c2d 100644 --- a/.gitignore +++ b/.gitignore @@ -123,3 +123,6 @@ crash.*.log # and should be committed. .vscode .pin_list.txt + +# pytest coverage data +.coverage diff --git a/backend/routes/allowlist.py b/backend/routes/allowlist.py index 2f65f99c292..b67f7d42127 100644 --- a/backend/routes/allowlist.py +++ b/backend/routes/allowlist.py @@ -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/", diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index f5b0a9aaf81..1af0148e452 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -69,7 +69,7 @@ }, "reportMatchNotExhaustive": { "baseline": 1, - "slack": 3 + "slack": 0 }, "reportMissingParameterType": { "baseline": 3933, diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index 8486e37384e..af4870bb1a5 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -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") diff --git a/litellm-rust/CLAUDE.md b/litellm-rust/CLAUDE.md index 0a985b833f3..7c723e570ef 100644 --- a/litellm-rust/CLAUDE.md +++ b/litellm-rust/CLAUDE.md @@ -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 diff --git a/litellm-rust/crates/ai-gateway/ARCHITECTURE.md b/litellm-rust/crates/ai-gateway/ARCHITECTURE.md new file mode 100644 index 00000000000..733953bbdb3 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/ARCHITECTURE.md @@ -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
LLM inference] + G <--> O[OpenAI realtime] + G -. spend tracking callback .-> P[litellm proxy] +``` diff --git a/litellm-rust/crates/ai-gateway/Cargo.toml b/litellm-rust/crates/ai-gateway/Cargo.toml index 5032259826b..567a881fe19 100644 --- a/litellm-rust/crates/ai-gateway/Cargo.toml +++ b/litellm-rust/crates/ai-gateway/Cargo.toml @@ -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 diff --git a/litellm-rust/crates/ai-gateway/README.md b/litellm-rust/crates/ai-gateway/README.md index 9e946c6e2ef..50d4c5249ac 100644 --- a/litellm-rust/crates/ai-gateway/README.md +++ b/litellm-rust/crates/ai-gateway/README.md @@ -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:///v1/realtime?model=` (WebSocket) -- **Auth:** `Authorization: Bearer ` — 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 ` — 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:///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 diff --git a/litellm-rust/crates/ai-gateway/src/auth/mod.rs b/litellm-rust/crates/ai-gateway/src/auth/mod.rs index 1920ff35b44..09ab1759f98 100644 --- a/litellm-rust/crates/ai-gateway/src/auth/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/auth/mod.rs @@ -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 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"); + } +} diff --git a/litellm-rust/crates/ai-gateway/src/auth/user_api_key.rs b/litellm-rust/crates/ai-gateway/src/auth/user_api_key.rs index 534e72bd2aa..15f2b16e91f 100644 --- a/litellm-rust/crates/ai-gateway/src/auth/user_api_key.rs +++ b/litellm-rust/crates/ai-gateway/src/auth/user_api_key.rs @@ -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()), diff --git a/litellm-rust/crates/ai-gateway/src/constants.rs b/litellm-rust/crates/ai-gateway/src/constants.rs new file mode 100644 index 00000000000..3116a4c9932 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/constants.rs @@ -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"; diff --git a/litellm-rust/crates/ai-gateway/src/integrations/custom_logger.rs b/litellm-rust/crates/ai-gateway/src/integrations/custom_logger.rs new file mode 100644 index 00000000000..53b599d8c98 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/integrations/custom_logger.rs @@ -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(()) + } +} diff --git a/litellm-rust/crates/ai-gateway/src/integrations/litellm_python_proxy_api.rs b/litellm-rust/crates/ai-gateway/src/integrations/litellm_python_proxy_api.rs new file mode 100644 index 00000000000..165a90d8dbe --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/integrations/litellm_python_proxy_api.rs @@ -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(name: &str, default: T) -> T +where + T: std::str::FromStr + PartialOrd + From, +{ + let zero = T::from(0u8); + std::env::var(name) + .ok() + .and_then(|value| value.trim().parse::().ok()) + .filter(|n| *n > zero) + .unwrap_or(default) +} + +/// Ships realtime logging events to the LiteLLM Python proxy. +pub struct LiteLLMPythonProxyAPILogger { + sink: Sender, +} + +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 { + let tunables = EgressTunables::from_env(); + let (sink, receiver) = mpsc::channel::(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 { + 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, + client: Client, + url: String, + master_key: String, + max_batch_size: usize, + flush_interval: Duration, +) { + let mut ticker = interval(flush_interval); + let mut batch: Vec = 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) { + 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}"); + } + } +} diff --git a/litellm-rust/crates/ai-gateway/src/integrations/mod.rs b/litellm-rust/crates/ai-gateway/src/integrations/mod.rs new file mode 100644 index 00000000000..8799be0c040 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/integrations/mod.rs @@ -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; diff --git a/litellm-rust/crates/ai-gateway/src/integrations/types.rs b/litellm-rust/crates/ai-gateway/src/integrations/types.rs new file mode 100644 index 00000000000..d61a1f816a7 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/integrations/types.rs @@ -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, + pub user_api_key_user_id: Option, + pub user_api_key_team_id: Option, +} + +/// 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, +} + +/// 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, +} + +/// 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, +} + +/// 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, // -> SpendLogs.api_key + pub user_api_key_user_id: Option, // -> SpendLogs.user + pub user_api_key_team_id: Option, // -> 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, + #[serde(skip_serializing_if = "Option::is_none")] + pub user_api_key_org_id: Option, // -> SpendLogs.organization_id + #[serde(skip_serializing_if = "Option::is_none")] + pub user_api_key_end_user_id: Option, // -> SpendLogs.end_user + #[serde(skip_serializing_if = "Option::is_none")] + pub spend_logs_metadata: Option>, +} + +/// 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, +} + +impl LogRecord { + pub fn into_callback_record(self) -> CallbackLogRecord { + CallbackLogRecord { + status: self.status, + standard_logging_payload: self.payload, + error: self.error, + } + } +} diff --git a/litellm-rust/crates/ai-gateway/src/io/realtime.rs b/litellm-rust/crates/ai-gateway/src/io/realtime.rs index 5d538d95fa4..4047de5cb26 100644 --- a/litellm-rust/crates/ai-gateway/src/io/realtime.rs +++ b/litellm-rust/crates/ai-gateway/src/io/realtime.rs @@ -126,6 +126,9 @@ pub(crate) async fn read_event(upstream_rx: &mut UpstreamRx) -> CoreResult( model: &str, @@ -133,6 +136,7 @@ pub(crate) async fn splice( mut upstream_rx: UpstreamRx, prelude: Option, idle_timeout: Option, + 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( model: &str, api_key: Option<&str>, api_base: Option<&str>, idle_timeout: Option, + 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( model: &str, handoff: crate::io::realtime_pool::WarmHandoff, idle_timeout: Option, + 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, ) diff --git a/litellm-rust/crates/ai-gateway/src/lib.rs b/litellm-rust/crates/ai-gateway/src/lib.rs index a2c228aeb30..6c04fbb7626 100644 --- a/litellm-rust/crates/ai-gateway/src/lib.rs +++ b/litellm-rust/crates/ai-gateway/src/lib.rs @@ -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; diff --git a/litellm-rust/crates/ai-gateway/src/main.rs b/litellm-rust/crates/ai-gateway/src/main.rs index c1c721510b7..9363c1cf623 100644 --- a/litellm-rust/crates/ai-gateway/src/main.rs +++ b/litellm-rust/crates/ai-gateway/src/main.rs @@ -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> = 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, diff --git a/litellm-rust/crates/ai-gateway/src/realtime/mod.rs b/litellm-rust/crates/ai-gateway/src/realtime/mod.rs new file mode 100644 index 00000000000..82be596ba86 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/realtime/mod.rs @@ -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; diff --git a/litellm-rust/crates/ai-gateway/src/realtime/streaming.rs b/litellm-rust/crates/ai-gateway/src/realtime/streaming.rs new file mode 100644 index 00000000000..34c82897808 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/realtime/streaming.rs @@ -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>, + /// 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>, + 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>, + 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> = 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> = 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); + } +} diff --git a/litellm-rust/crates/ai-gateway/src/routes/realtime/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/realtime/mod.rs index 91f106b20c5..cf484274268 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/realtime/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/realtime/mod.rs @@ -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 { 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, Query(query): Query, ) -> Result { - 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, pool: Arc, + loggers: Arc>>, + auth: UserApiKeyAuth, + master_key: Option>, 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::(&text).ok(), _ => None, } }); + // Plain forwarding sink — no observe here anymore. let client_out = ws_sink.with(|event: RealtimeEvent| async move { Ok::(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()) + ); } } diff --git a/litellm-rust/crates/ai-gateway/src/routes/realtime/service.rs b/litellm-rust/crates/ai-gateway/src/routes/realtime/service.rs index c78ca8df446..d6c31edd454 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/realtime/service.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/realtime/service.rs @@ -26,6 +26,7 @@ pub async fn run( pool: &RealtimePool, model: &str, idle_timeout: Option, + 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, ) diff --git a/litellm-rust/crates/ai-gateway/src/state.rs b/litellm-rust/crates/ai-gateway/src/state.rs index 95102d30694..cc35bd51b97 100644 --- a/litellm-rust/crates/ai-gateway/src/state.rs +++ b/litellm-rust/crates/ai-gateway/src/state.rs @@ -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>, + /// Logging callbacks fanned out at the end of each realtime session. + pub loggers: Arc>>, /// 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. diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index 74e753b09ea..aeec58f1dfc 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -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: diff --git a/litellm/constants.py b/litellm/constants.py index 212d34357f8..09235106c63 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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) diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index c6d427e7f09..5baa7cbc9c5 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -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, diff --git a/litellm/files/utils.py b/litellm/files/utils.py index a2b9a42c154..a0df7a89b0f 100644 --- a/litellm/files/utils.py +++ b/litellm/files/utils.py @@ -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 diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index fc3c25e0d95..c88f8b77dc2 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -14,6 +14,7 @@ OPTIONAL_KWARGS_KEYS = frozenset( "azure_password", "azure_scope", "timeout", + "gcs_bucket_name", "bucket_name", "vertex_credentials", "vertex_project", diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index fe34731759f..bf9ce3b0acb 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -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. diff --git a/litellm/llms/base_llm/files/transformation.py b/litellm/llms/base_llm/files/transformation.py index c3abfafc552..85016c7a5c4 100644 --- a/litellm/llms/base_llm/files/transformation.py +++ b/litellm/llms/base_llm/files/transformation.py @@ -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 diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 948c90f9f99..d33ec295e94 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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": ( + "" + 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", diff --git a/litellm/llms/vertex_ai/files/handler.py b/litellm/llms/vertex_ai/files/handler.py index c31bfde69e7..176cfe98411 100644 --- a/litellm/llms/vertex_ai/files/handler.py +++ b/litellm/llms/vertex_ai/files/handler.py @@ -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, diff --git a/litellm/llms/vertex_ai/files/transformation.py b/litellm/llms/vertex_ai/files/transformation.py index f30518bc7ca..d5164d8c1c2 100644 --- a/litellm/llms/vertex_ai/files/transformation.py +++ b/litellm/llms/vertex_ai/files/transformation.py @@ -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": } 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", - ) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 5e704b889ae..4e8e04245c2 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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: diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py new file mode 100644 index 00000000000..39db2314aee --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py @@ -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) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py index 7bcdb3e6529..969bbf01ec8 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py @@ -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( diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index 5b691beccbf..91c604e6204 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -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 ( diff --git a/litellm/proxy/logging_endpoints/__init__.py b/litellm/proxy/logging_endpoints/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/proxy/logging_endpoints/callback_logs_endpoints.py b/litellm/proxy/logging_endpoints/callback_logs_endpoints.py new file mode 100644 index 00000000000..a96a5431294 --- /dev/null +++ b/litellm/proxy/logging_endpoints/callback_logs_endpoints.py @@ -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) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 2d49297c8e9..cd40f3155ee 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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( diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index 07c355f2cd9..c29d1ac178d 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -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 diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index 944423632ef..7c19804dde3 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -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( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 97bd598c406..b134414f65e 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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) diff --git a/litellm/router_utils/batch_utils.py b/litellm/router_utils/batch_utils.py index 5e58479825b..ddec753d362 100644 --- a/litellm/router_utils/batch_utils.py +++ b/litellm/router_utils/batch_utils.py @@ -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 diff --git a/litellm/types/files.py b/litellm/types/files.py index bf56894329c..1b2d7e30f1f 100644 --- a/litellm/types/files.py +++ b/litellm/types/files.py @@ -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] diff --git a/litellm/types/proxy/callback_logs_endpoints.py b/litellm/types/proxy/callback_logs_endpoints.py new file mode 100644 index 00000000000..ef148274ca7 --- /dev/null +++ b/litellm/types/proxy/callback_logs_endpoints.py @@ -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) diff --git a/litellm/types/router.py b/litellm/types/router.py index 607bfd584fd..b5285f11f8b 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -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 diff --git a/tests/batches_tests/test_openai_batches_and_files.py b/tests/batches_tests/test_openai_batches_and_files.py index bccb5eaaacb..8a2d5f33805 100644 --- a/tests/batches_tests/test_openai_batches_and_files.py +++ b/tests/batches_tests/test_openai_batches_and_files.py @@ -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" diff --git a/tests/e2e/budgets/BUDGET_CODE_MATRIX.md b/tests/e2e/budgets/BUDGET_CODE_MATRIX.md new file mode 100644 index 00000000000..63cb41228d2 --- /dev/null +++ b/tests/e2e/budgets/BUDGET_CODE_MATRIX.md @@ -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 | diff --git a/tests/e2e/budgets/BUDGET_TEST_COVERAGE_MATRIX.md b/tests/e2e/budgets/BUDGET_TEST_COVERAGE_MATRIX.md new file mode 100644 index 00000000000..62bfc1fdd41 --- /dev/null +++ b/tests/e2e/budgets/BUDGET_TEST_COVERAGE_MATRIX.md @@ -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. diff --git a/tests/e2e/budgets/budget_client.py b/tests/e2e/budgets/budget_client.py new file mode 100644 index 00000000000..af8021f9b93 --- /dev/null +++ b/tests/e2e/budgets/budget_client.py @@ -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()) diff --git a/tests/e2e/budgets/conftest.py b/tests/e2e/budgets/conftest.py new file mode 100644 index 00000000000..236822f4309 --- /dev/null +++ b/tests/e2e/budgets/conftest.py @@ -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() diff --git a/tests/e2e/budgets/test_budget_crud_e2e.py b/tests/e2e/budgets/test_budget_crud_e2e.py new file mode 100644 index 00000000000..e697eca0051 --- /dev/null +++ b/tests/e2e/budgets/test_budget_crud_e2e.py @@ -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" diff --git a/tests/e2e/budgets/test_budget_enforcement_e2e.py b/tests/e2e/budgets/test_budget_enforcement_e2e.py new file mode 100644 index 00000000000..d03288b637a --- /dev/null +++ b/tests/e2e/budgets/test_budget_enforcement_e2e.py @@ -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)) diff --git a/tests/e2e/budgets/test_budget_reset_e2e.py b/tests/e2e/budgets/test_budget_reset_e2e.py new file mode 100644 index 00000000000..dcf776db9a2 --- /dev/null +++ b/tests/e2e/budgets/test_budget_reset_e2e.py @@ -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") diff --git a/tests/e2e/budgets/test_model_max_budget_e2e.py b/tests/e2e/budgets/test_model_max_budget_e2e.py new file mode 100644 index 00000000000..44e6a333ef0 --- /dev/null +++ b/tests/e2e/budgets/test_model_max_budget_e2e.py @@ -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) diff --git a/tests/e2e/budgets/test_multi_window_budget_e2e.py b/tests/e2e/budgets/test_multi_window_budget_e2e.py new file mode 100644 index 00000000000..553ad1ce701 --- /dev/null +++ b/tests/e2e/budgets/test_multi_window_budget_e2e.py @@ -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") diff --git a/tests/e2e/budgets/test_soft_budget_e2e.py b/tests/e2e/budgets/test_soft_budget_e2e.py new file mode 100644 index 00000000000..407de7ae467 --- /dev/null +++ b/tests/e2e/budgets/test_soft_budget_e2e.py @@ -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 diff --git a/tests/e2e/budgets/test_spend_counter_reseed_e2e.py b/tests/e2e/budgets/test_spend_counter_reseed_e2e.py new file mode 100644 index 00000000000..a6860aeef43 --- /dev/null +++ b/tests/e2e/budgets/test_spend_counter_reseed_e2e.py @@ -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)" + ) diff --git a/tests/e2e/budgets/test_tag_budget_e2e.py b/tests/e2e/budgets/test_tag_budget_e2e.py new file mode 100644 index 00000000000..7cec5bc96c1 --- /dev/null +++ b/tests/e2e/budgets/test_tag_budget_e2e.py @@ -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) diff --git a/tests/e2e/budgets/test_team_member_budget_e2e.py b/tests/e2e/budgets/test_team_member_budget_e2e.py new file mode 100644 index 00000000000..301617bfdca --- /dev/null +++ b/tests/e2e/budgets/test_team_member_budget_e2e.py @@ -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") diff --git a/tests/e2e/budgets/test_team_member_budget_reset_e2e.py b/tests/e2e/budgets/test_team_member_budget_reset_e2e.py new file mode 100644 index 00000000000..2749f16a26e --- /dev/null +++ b/tests/e2e/budgets/test_team_member_budget_reset_e2e.py @@ -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") diff --git a/tests/e2e/budgets/test_team_multi_window_budget_e2e.py b/tests/e2e/budgets/test_team_multi_window_budget_e2e.py new file mode 100644 index 00000000000..946ab8ee1f2 --- /dev/null +++ b/tests/e2e/budgets/test_team_multi_window_budget_e2e.py @@ -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") diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py new file mode 100644 index 00000000000..cc95c7538dd --- /dev/null +++ b/tests/e2e/conftest.py @@ -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() diff --git a/tests/e2e/e2e_config.py b/tests/e2e/e2e_config.py new file mode 100644 index 00000000000..3865804b08f --- /dev/null +++ b/tests/e2e/e2e_config.py @@ -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] diff --git a/tests/e2e/e2e_gateway.py b/tests/e2e/e2e_gateway.py new file mode 100644 index 00000000000..b700145d434 --- /dev/null +++ b/tests/e2e/e2e_gateway.py @@ -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, + ) diff --git a/tests/e2e/e2e_http.py b/tests/e2e/e2e_http.py new file mode 100644 index 00000000000..7458f316852 --- /dev/null +++ b/tests/e2e/e2e_http.py @@ -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="", + 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) diff --git a/tests/e2e/gateway/litellm-config.yml b/tests/e2e/gateway/litellm-config.yml new file mode 100644 index 00000000000..56c209ee750 --- /dev/null +++ b/tests/e2e/gateway/litellm-config.yml @@ -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 diff --git a/tests/e2e/lifecycle.py b/tests/e2e/lifecycle.py new file mode 100644 index 00000000000..fdf2137584e --- /dev/null +++ b/tests/e2e/lifecycle.py @@ -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 diff --git a/tests/e2e/llm_translation/LLM_TRANSLATION_COVERAGE_MATRIX.md b/tests/e2e/llm_translation/LLM_TRANSLATION_COVERAGE_MATRIX.md new file mode 100644 index 00000000000..5e4a448857f --- /dev/null +++ b/tests/e2e/llm_translation/LLM_TRANSLATION_COVERAGE_MATRIX.md @@ -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=` to a deadline. Streaming cost is only +known after the stream is fully consumed. diff --git a/tests/e2e/llm_translation/conftest.py b/tests/e2e/llm_translation/conftest.py new file mode 100644 index 00000000000..fbf008cf085 --- /dev/null +++ b/tests/e2e/llm_translation/conftest.py @@ -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() diff --git a/tests/e2e/llm_translation/passthrough_client.py b/tests/e2e/llm_translation/passthrough_client.py new file mode 100644 index 00000000000..fff4064a328 --- /dev/null +++ b/tests/e2e/llm_translation/passthrough_client.py @@ -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()) diff --git a/tests/e2e/llm_translation/test_custom_pricing_e2e.py b/tests/e2e/llm_translation/test_custom_pricing_e2e.py new file mode 100644 index 00000000000..4b3e87b78e1 --- /dev/null +++ b/tests/e2e/llm_translation/test_custom_pricing_e2e.py @@ -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" + ) diff --git a/tests/e2e/llm_translation/test_passthrough_e2e.py b/tests/e2e/llm_translation/test_passthrough_e2e.py new file mode 100644 index 00000000000..37d55c665b3 --- /dev/null +++ b/tests/e2e/llm_translation/test_passthrough_e2e.py @@ -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" diff --git a/tests/e2e/models.py b/tests/e2e/models.py new file mode 100644 index 00000000000..fbeb3d44fa5 --- /dev/null +++ b/tests/e2e/models.py @@ -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] = [] diff --git a/tests/e2e/pytest.ini b/tests/e2e/pytest.ini new file mode 100644 index 00000000000..7799f6b16a2 --- /dev/null +++ b/tests/e2e/pytest.ini @@ -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 diff --git a/tests/e2e/spend_tracking/SPEND_TRACKING_COVERAGE_MATRIX.md b/tests/e2e/spend_tracking/SPEND_TRACKING_COVERAGE_MATRIX.md new file mode 100644 index 00000000000..53c4d4ace83 --- /dev/null +++ b/tests/e2e/spend_tracking/SPEND_TRACKING_COVERAGE_MATRIX.md @@ -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). diff --git a/tests/e2e/spend_tracking/conftest.py b/tests/e2e/spend_tracking/conftest.py new file mode 100644 index 00000000000..1d01ab3d17a --- /dev/null +++ b/tests/e2e/spend_tracking/conftest.py @@ -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() diff --git a/tests/e2e/spend_tracking/spend_e2e_client.py b/tests/e2e/spend_tracking/spend_e2e_client.py new file mode 100644 index 00000000000..d749d69f1a4 --- /dev/null +++ b/tests/e2e/spend_tracking/spend_e2e_client.py @@ -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()) diff --git a/tests/e2e/spend_tracking/test_spend_routes.py b/tests/e2e/spend_tracking/test_spend_routes.py new file mode 100644 index 00000000000..e3c96a4d578 --- /dev/null +++ b/tests/e2e/spend_tracking/test_spend_routes.py @@ -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) diff --git a/tests/e2e/spend_tracking/test_spend_tracking_e2e.py b/tests/e2e/spend_tracking/test_spend_tracking_e2e.py new file mode 100644 index 00000000000..8c9e913b10f --- /dev/null +++ b/tests/e2e/spend_tracking/test_spend_tracking_e2e.py @@ -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) diff --git a/tests/e2e/test_lifecycle.py b/tests/e2e/test_lifecycle.py new file mode 100644 index 00000000000..d3c559dd2ed --- /dev/null +++ b/tests/e2e/test_lifecycle.py @@ -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" + ) diff --git a/tests/e2e/transport.py b/tests/e2e/transport.py new file mode 100644 index 00000000000..37412fc0cf5 --- /dev/null +++ b/tests/e2e/transport.py @@ -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) diff --git a/tests/mcp_tests/test_mcp_auth_priority.py b/tests/mcp_tests/test_mcp_auth_priority.py index ad6e9438edd..7ae0f59afe5 100644 --- a/tests/mcp_tests/test_mcp_auth_priority.py +++ b/tests/mcp_tests/test_mcp_auth_priority.py @@ -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 diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index eea2f2721ab..5f8fcbf835e 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -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 diff --git a/tests/pyrightconfig.json b/tests/pyrightconfig.json new file mode 100644 index 00000000000..5757c97f812 --- /dev/null +++ b/tests/pyrightconfig.json @@ -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" +} \ No newline at end of file diff --git a/tests/router_unit_tests/test_router_batch_utils.py b/tests/router_unit_tests/test_router_batch_utils.py index 1b8f713a437..b8760906645 100644 --- a/tests/router_unit_tests/test_router_batch_utils.py +++ b/tests/router_unit_tests/test_router_batch_utils.py @@ -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 diff --git a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py index c9e500b4a5b..704dd7f92f1 100644 --- a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py +++ b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py @@ -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__]) diff --git a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_binary_file_upload.py b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_binary_file_upload.py index d4586134b13..122518d4acb 100644 --- a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_binary_file_upload.py +++ b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_binary_file_upload.py @@ -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", diff --git a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_streaming.py b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_streaming.py new file mode 100644 index 00000000000..cd556c48b6b --- /dev/null +++ b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_streaming.py @@ -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) diff --git a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_transformation.py b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_transformation.py index 7c063c72607..8c5305ee67b 100644 --- a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_transformation.py @@ -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" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py new file mode 100644 index 00000000000..5481d60a22a --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py index 7885617aa46..75be6dfc157 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py @@ -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) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 8dbee1daa36..35a67391315 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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__]) diff --git a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py index a6f6e651487..1d4d39ec140 100644 --- a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py +++ b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py @@ -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 diff --git a/tests/test_litellm/proxy/logging_endpoints/__init__.py b/tests/test_litellm/proxy/logging_endpoints/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/proxy/logging_endpoints/test_callback_logs_endpoints.py b/tests/test_litellm/proxy/logging_endpoints/test_callback_logs_endpoints.py new file mode 100644 index 00000000000..40e89329b8d --- /dev/null +++ b/tests/test_litellm/proxy/logging_endpoints/test_callback_logs_endpoints.py @@ -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) diff --git a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py index c77ac11ffc1..2b38d732e9d 100644 --- a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py @@ -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", diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index f42639cee8a..46ecb31e1c8 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -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 diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 3e2150848a7..470d38caf10 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -3062,6 +3062,36 @@ def test_get_deployment_credentials_with_provider_aws_bedrock_runtime_endpoint() assert credentials["custom_llm_provider"] == "bedrock" +def test_get_deployment_credentials_with_provider_includes_bucket_name(): + """ + Regression: bucket_name must survive the CredentialLiteLLMParams filter so + managed-files batch retrieval can resolve the GCS/S3 bucket. Previously it was + dropped, causing "GCS bucket_name is required" when fetching batch output files. + """ + router = litellm.Router( + model_list=[ + { + "model_name": "vertex-gemini", + "litellm_params": { + "model": "vertex_ai/gemini-3.5-flash", + "vertex_project": "my-project", + "vertex_location": "global", + "gcs_bucket_name": "my-batch-bucket", + }, + } + ], + ) + + credentials = router.get_deployment_credentials_with_provider( + model_id="vertex-gemini" + ) + + assert credentials is not None + assert credentials["gcs_bucket_name"] == "my-batch-bucket" + assert credentials["vertex_project"] == "my-project" + assert credentials["custom_llm_provider"] == "vertex_ai" + + def test_get_deployment_credentials_with_provider_resolves_credential_name(): """ Test that get_deployment_credentials_with_provider correctly resolves diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 6fae14ee6ec..3669acae67e 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -17236,6 +17236,29 @@ export interface paths { patch?: never; trace?: never; }; + "/v1/rust_control_plane/logs": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** + * Ingest Callback Logs + * @description 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. + */ + post: operations["ingest_callback_logs_v1_rust_control_plane_logs_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/v1/search": { parameters: { query?: never; @@ -21512,6 +21535,54 @@ export interface components { /** Callback Name */ callback_name: string; }; + /** + * CallbackLogFailure + * @description A record that failed to replay, identified by its index in the batch. + */ + CallbackLogFailure: { + /** Error */ + error: string; + /** Index */ + index: number; + }; + /** + * CallbackLogRecord + * @description A single finished logging event to replay through the callbacks. + */ + CallbackLogRecord: { + /** Error */ + error?: string | null; + /** Standard Logging Payload */ + standard_logging_payload: { + [key: string]: unknown; + }; + /** + * Status + * @enum {string} + */ + status: "success" | "failure"; + }; + /** + * CallbackLogsRequest + * @description A batch of logging events posted by an external producer. + */ + CallbackLogsRequest: { + /** Records */ + records: components["schemas"]["CallbackLogRecord"][]; + }; + /** + * CallbackLogsResponse + * @description Per-batch result: counts plus per-record failure detail so the caller can + * distinguish a transient callback error from a structurally bad payload. + */ + CallbackLogsResponse: { + /** Failed */ + failed: number; + /** Failures */ + failures?: components["schemas"]["CallbackLogFailure"][]; + /** Processed */ + processed: number; + }; /** CallbacksByType */ CallbacksByType: { /** Failure */ @@ -54686,6 +54757,39 @@ export interface operations { }; }; }; + ingest_callback_logs_v1_rust_control_plane_logs_post: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody: { + content: { + "application/json": components["schemas"]["CallbackLogsRequest"]; + }; + }; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["CallbackLogsResponse"]; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; search_v1_search_post: { parameters: { query?: {