mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
Merge remote-tracking branch 'refs/remotes/upstream/litellm_internal_staging' into codex/rust-auth-control-plane
# Conflicts: # litellm-rust/Cargo.lock # litellm-rust/Cargo.toml # litellm-rust/crates/ai-gateway/Cargo.toml # litellm-rust/crates/ai-gateway/src/routes/realtime/mod.rs # litellm/proxy/proxy_server.py
This commit is contained in:
commit
c5c0d5fce6
102 changed files with 8986 additions and 800 deletions
1
.github/workflows/test-unit-proxy-infra.yml
vendored
1
.github/workflows/test-unit-proxy-infra.yml
vendored
|
|
@ -29,6 +29,7 @@ jobs:
|
|||
tests/test_litellm/proxy/_experimental
|
||||
tests/test_litellm/proxy/experimental
|
||||
tests/test_litellm/proxy/common_utils
|
||||
tests/test_litellm/proxy/logging_endpoints
|
||||
tests/test_litellm/proxy/test_*.py
|
||||
workers: 2
|
||||
reruns: 2
|
||||
|
|
|
|||
3
.gitignore
vendored
3
.gitignore
vendored
|
|
@ -123,3 +123,6 @@ crash.*.log
|
|||
# and should be committed.
|
||||
.vscode
|
||||
.pin_list.txt
|
||||
|
||||
# pytest coverage data
|
||||
.coverage
|
||||
|
|
|
|||
|
|
@ -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/",
|
||||
|
|
|
|||
|
|
@ -69,7 +69,7 @@
|
|||
},
|
||||
"reportMatchNotExhaustive": {
|
||||
"baseline": 1,
|
||||
"slack": 3
|
||||
"slack": 0
|
||||
},
|
||||
"reportMissingParameterType": {
|
||||
"baseline": 3933,
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
12
litellm-rust/crates/ai-gateway/ARCHITECTURE.md
Normal file
12
litellm-rust/crates/ai-gateway/ARCHITECTURE.md
Normal file
|
|
@ -0,0 +1,12 @@
|
|||
# ai-gateway architecture
|
||||
|
||||
The Rust ai-gateway does LLM inference (realtime WebSocket). Spend tracking is an
|
||||
API callback: it POSTs each finished session to the LiteLLM proxy, which records
|
||||
spend and runs the usual callbacks.
|
||||
|
||||
```mermaid
|
||||
flowchart LR
|
||||
C[client] <--> G[Rust ai-gateway<br/>LLM inference]
|
||||
G <--> O[OpenAI realtime]
|
||||
G -. spend tracking callback .-> P[litellm proxy]
|
||||
```
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -17,8 +17,9 @@ dials OpenAI upstream, and splices the two sockets frame-by-frame.
|
|||
Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge.
|
||||
|
||||
- **Client endpoint:** `wss://<host>/v1/realtime?model=<model>` (WebSocket)
|
||||
- **Auth:** `Authorization: Bearer <key>` — proxy-admin only for realtime serving today. Non-admin virtual keys are verified but rejected until realtime usage is reported back to the LiteLLM proxy. Fails closed.
|
||||
- **Auth:** `Authorization: Bearer <key>` — the master key (admin) or a virtual key validated by the LiteLLM proxy (see below). Fails closed.
|
||||
- **Health:** `GET /health/readiness`, `GET /health/liveness`, `GET /health/gil`
|
||||
- **Request logs:** POSTed to a LiteLLM proxy at `/v1/rust_control_plane/logs` (see [Request logging](#request-logging))
|
||||
|
||||
> **Realtime serving is pure Rust.** Python is used at **load time only** — to
|
||||
> read the config once at boot. The realtime hot path never touches Python.
|
||||
|
|
@ -42,9 +43,9 @@ route/model permissions are enforced in exactly one place. A client presents
|
|||
load. A bounded ≤200-entry cache is available for future high-RPS *per-request*
|
||||
routes via `LITELLM_AUTH_CACHE_TTL_SECS > 0`, which trades budget/rate-limit
|
||||
freshness (bounded by the TTL) for fewer control-plane calls — like the proxy's
|
||||
own ~60s auth cache. Non-admin virtual-key identities are rejected on the
|
||||
realtime route until realtime usage is reported back to the proxy, so callers
|
||||
cannot consume upstream usage without key/team spend attribution.
|
||||
own ~60s auth cache. Realtime session usage is reported back to the proxy with
|
||||
the resolved key/user/team metadata, so key and team spend logs keep their
|
||||
normal attribution.
|
||||
|
||||
Data plane → control plane is itself authenticated with a **dedicated data-plane key**
|
||||
(NOT the master key — least privilege): the gateway sends
|
||||
|
|
@ -74,13 +75,11 @@ export LITELLM_AUTH_VERIFY_URL=https://<proxy-host>/v1/rust_control_plane/authen
|
|||
./litellm-ai-gateway
|
||||
```
|
||||
|
||||
Fails closed: a missing/wrong data-plane key, an unreachable proxy, or a rejected key
|
||||
all yield `401`; a valid non-admin virtual key currently yields `403` on
|
||||
`/v1/realtime` until realtime spend reporting is available. Revocation and budget
|
||||
changes take effect within the cache TTL (if caching is enabled via
|
||||
`LITELLM_AUTH_CACHE_TTL_SECS`; off by default → every connection re-verifies).
|
||||
Keep the proxy on a private network — the verify endpoint is internal-only and
|
||||
excluded from the public OpenAPI spec.
|
||||
Fails closed: a missing/wrong data-plane key, an unreachable proxy, or a rejected
|
||||
key all yield `401`. Revocation and budget changes take effect within the cache
|
||||
TTL (if caching is enabled via `LITELLM_AUTH_CACHE_TTL_SECS`; off by default →
|
||||
every connection re-verifies). Keep the proxy on a private network — the verify
|
||||
endpoint is internal-only and excluded from the public OpenAPI spec.
|
||||
|
||||
## Configuration (config.yaml)
|
||||
|
||||
|
|
@ -122,11 +121,12 @@ overridden at deploy time (e.g. a Render secret file mounted at the same path).
|
|||
| `LITELLM_CONFIG_PATH` | yes (config mode) | — | Path to the config.yaml the gateway loads its `model_list` from. The Docker image defaults this to `/app/config.yaml`. |
|
||||
| `LITELLM_MASTER_KEY` | yes | — | Admin bearer token (checked locally, no proxy call). Unset ⇒ the master-key path is disabled. |
|
||||
| `LITELLM_DATA_PLANE_KEY` | for virtual keys | — | Dedicated secret the gateway sends as `X-LiteLLM-Data-Plane-Key` to authenticate itself to the proxy's verify endpoint. **Must match the proxy's `LITELLM_DATA_PLANE_KEY`.** Not the master key. |
|
||||
| `LITELLM_AUTH_VERIFY_URL` | for virtual keys | `http://localhost:4000/v1/rust_control_plane/authentication` | The proxy's verify endpoint the gateway delegates virtual-key auth to. Non-admin virtual keys are still rejected by `/v1/realtime` until realtime spend reporting is wired. |
|
||||
| `LITELLM_AUTH_VERIFY_URL` | for virtual keys | `http://localhost:4000/v1/rust_control_plane/authentication` | The proxy's verify endpoint the gateway delegates virtual-key auth to. |
|
||||
| `LITELLM_AUTH_CACHE_TTL_SECS` | no | `0` | Verified-key cache TTL. **`0` = off (default)** → every connection re-verifies (budget/rate-limit enforced each time); cheap for realtime since auth is per-connection. Set `> 0` only for high-RPS per-request routes, trading budget/rate-limit freshness for fewer proxy calls. |
|
||||
| `OPENAI_API_KEY` | yes | — | Upstream OpenAI key. Referenced by config.yaml as `os.environ/OPENAI_API_KEY` for the gateway→OpenAI dial. |
|
||||
| `HOST` | no | `127.0.0.1` | **Set to `0.0.0.0` in any container/deploy** or external traffic is refused. |
|
||||
| `PORT` | no | `4001` | Listen port. Render and most PaaS inject this automatically. |
|
||||
| `LITELLM_PROXY_BASE_URL` | no | `http://localhost:4000` | LiteLLM proxy that request logs are POSTed to. See [Request logging](#request-logging). |
|
||||
|
||||
> Secrets (`LITELLM_MASTER_KEY`, `LITELLM_DATA_PLANE_KEY`, `OPENAI_API_KEY`) are never
|
||||
> baked into the image or `render.yaml` — inject them at deploy time only.
|
||||
|
|
@ -145,6 +145,18 @@ This mode links no libpython and needs no config file, but it only supports one
|
|||
hard-coded OpenAI deployment. **config.yaml is the recommended path** — use the
|
||||
stand-in only for the leanest possible build.
|
||||
|
||||
## Request logging
|
||||
|
||||
The gateway runs no spend logic. When a session ends it builds one
|
||||
`StandardLoggingPayload` and POSTs it to `{LITELLM_PROXY_BASE_URL}/v1/rust_control_plane/logs`
|
||||
(admin-only, bearer = `LITELLM_MASTER_KEY`), and the proxy replays it through its
|
||||
normal callbacks (spend logs, Langfuse, etc.). The POST is non-blocking: a bounded
|
||||
channel drained by a background worker, dropping with a counter if the proxy is
|
||||
down. It sends one payload per session. Both env vars are in the table above.
|
||||
|
||||
Worker tuning, rarely needed: `LITELLM_LOG_CHANNEL_CAPACITY` (4096),
|
||||
`LITELLM_LOG_BATCH_SIZE` (256), `LITELLM_LOG_FLUSH_INTERVAL_MS` (500).
|
||||
|
||||
## Build & run with Docker
|
||||
|
||||
The image is built `--features python-config` and installs litellm **from this
|
||||
|
|
|
|||
|
|
@ -2,12 +2,11 @@
|
|||
//! keeps handlers clean and auth testable).
|
||||
//!
|
||||
//! For now this supports both the local **master key** and virtual keys resolved
|
||||
//! through the Python control plane. Individual routes can still narrow the
|
||||
//! accepted identities. For example, realtime rejects non-admin virtual keys
|
||||
//! until realtime usage is reported back to the proxy for spend attribution.
|
||||
//! through the Python control plane. Routes receive the resolved identity and can
|
||||
//! pass it to logging/billing paths for spend attribution.
|
||||
//!
|
||||
//! A handler opts in by adding [`RequireMasterKey`] to its arguments; auth then
|
||||
//! runs during extraction, before the handler body. Routes never re-implement it.
|
||||
//! A handler opts in by adding an auth extractor to its arguments; auth then runs
|
||||
//! during extraction, before the handler body. Routes never re-implement it.
|
||||
|
||||
pub mod cache;
|
||||
pub mod client;
|
||||
|
|
@ -21,10 +20,29 @@ use axum::extract::FromRequestParts;
|
|||
use axum::http::header::AUTHORIZATION;
|
||||
use axum::http::request::Parts;
|
||||
use axum::http::StatusCode;
|
||||
use sha2::{Digest, Sha256};
|
||||
use subtle::ConstantTimeEq;
|
||||
|
||||
use crate::state::AppState;
|
||||
|
||||
/// SHA-256 hex digest of a token — the exact transform the Python proxy applies
|
||||
/// (`litellm.proxy.utils.hash_token`).
|
||||
///
|
||||
/// STRICT REQUIREMENT: a raw key (`LITELLM_MASTER_KEY`, a virtual key, …) must
|
||||
/// **never** leave this gateway in a log payload. Spend logs and every callback
|
||||
/// integration receive `user_api_key_hash`, so that field must be this hash, not
|
||||
/// the credential. Hashing here also means the value matches the key's hash in
|
||||
/// `LiteLLM_SpendLogs.api_key`, so realtime spend joins with the rest of LiteLLM.
|
||||
pub fn hash_token(token: &str) -> String {
|
||||
let digest = Sha256::digest(token.as_bytes());
|
||||
let mut hex = String::with_capacity(digest.len() * 2);
|
||||
for byte in digest {
|
||||
use std::fmt::Write;
|
||||
let _ = write!(hex, "{byte:02x}");
|
||||
}
|
||||
hex
|
||||
}
|
||||
|
||||
/// Extractor that requires the configured master key as a bearer token.
|
||||
///
|
||||
/// Rejections: `500` when no master key is configured (permanent
|
||||
|
|
@ -61,3 +79,23 @@ impl FromRequestParts<AppState> for RequireMasterKey {
|
|||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::hash_token;
|
||||
|
||||
#[test]
|
||||
fn hash_token_matches_python_sha256_hexdigest() {
|
||||
// Must equal hashlib.sha256("sk-1234".encode()).hexdigest() — the value
|
||||
// the proxy stores in LiteLLM_SpendLogs.api_key.
|
||||
assert_eq!(
|
||||
hash_token("sk-1234"),
|
||||
"88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b"
|
||||
);
|
||||
// 64 lowercase hex chars, and never the raw input.
|
||||
let h = hash_token("sk-secret");
|
||||
assert_eq!(h.len(), 64);
|
||||
assert!(h.chars().all(|c| c.is_ascii_hexdigit()));
|
||||
assert_ne!(h, "sk-secret");
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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()),
|
||||
|
|
|
|||
29
litellm-rust/crates/ai-gateway/src/constants.rs
Normal file
29
litellm-rust/crates/ai-gateway/src/constants.rs
Normal file
|
|
@ -0,0 +1,29 @@
|
|||
//! Crate-level constants for the ai-gateway.
|
||||
//!
|
||||
//! Per `litellm-rust/CLAUDE.md`, magic numbers and fixed strings live here
|
||||
//! (the Rust mirror of Python's `litellm/constants.py`), not inline in feature
|
||||
//! modules. Env-overridable tunables keep their `DEFAULT_*` value here; the env
|
||||
//! read + fallback happens at the host/config layer.
|
||||
|
||||
/// Default LiteLLM control-plane base URL for request-log egress when
|
||||
/// `LITELLM_PROXY_BASE_URL` is unset.
|
||||
pub(crate) const DEFAULT_PROXY_BASE_URL: &str = "http://localhost:4000";
|
||||
|
||||
/// The logs ingest path appended to the proxy base. Not a tunable; it is the
|
||||
/// proxy's API contract (the rust-control-plane router on the Python proxy).
|
||||
pub(crate) const RUST_CONTROL_PLANE_LOGS_PATH: &str = "/v1/rust_control_plane/logs";
|
||||
|
||||
/// Default bounded channel depth for the log-egress worker.
|
||||
/// Override: `LITELLM_LOG_CHANNEL_CAPACITY`.
|
||||
pub(crate) const DEFAULT_CHANNEL_CAPACITY: usize = 4096;
|
||||
|
||||
/// Default max records POSTed per request to the control plane.
|
||||
/// Override: `LITELLM_LOG_BATCH_SIZE`.
|
||||
pub(crate) const DEFAULT_MAX_BATCH_SIZE: usize = 256;
|
||||
|
||||
/// Default partial-batch flush cadence, in ms.
|
||||
/// Override: `LITELLM_LOG_FLUSH_INTERVAL_MS`.
|
||||
pub(crate) const DEFAULT_FLUSH_INTERVAL_MS: u64 = 500;
|
||||
|
||||
/// Provider attributed to realtime sessions in the logging payload.
|
||||
pub(crate) const DEFAULT_PROVIDER: &str = "openai";
|
||||
|
|
@ -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(())
|
||||
}
|
||||
}
|
||||
|
|
@ -0,0 +1,209 @@
|
|||
//! A `CustomLogger` that ships finished events to the LiteLLM Python proxy's
|
||||
//! `/v1/rust_control_plane/logs` endpoint.
|
||||
//!
|
||||
//! The callback path is non-blocking: `log_success_event` / `log_failure_event`
|
||||
//! build a `LogRecord` and `try_send` it onto a bounded channel, returning a
|
||||
//! `LogError` (never panicking, never awaiting) if the channel is full or the
|
||||
//! worker has gone away. A spawned background worker drains the channel, batches
|
||||
//! records into `{"records":[...]}`, and POSTs them to the proxy with a pooled
|
||||
//! `reqwest::Client`.
|
||||
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use reqwest::Client;
|
||||
use tokio::sync::mpsc::{self, Receiver, Sender};
|
||||
use tokio::time::interval;
|
||||
|
||||
use crate::constants::{
|
||||
DEFAULT_CHANNEL_CAPACITY, DEFAULT_FLUSH_INTERVAL_MS, DEFAULT_MAX_BATCH_SIZE,
|
||||
DEFAULT_PROXY_BASE_URL, RUST_CONTROL_PLANE_LOGS_PATH,
|
||||
};
|
||||
use crate::integrations::custom_logger::CustomLogger;
|
||||
use crate::integrations::types::{
|
||||
CallbackLogsRequest, LogError, LogRecord, LoggingError, StandardLoggingPayload,
|
||||
};
|
||||
|
||||
/// Egress worker tunables. Each field defaults to the matching `DEFAULT_*` const
|
||||
/// in `crate::constants` and is overridable via an env var (read once at logger
|
||||
/// construction).
|
||||
struct EgressTunables {
|
||||
channel_capacity: usize,
|
||||
max_batch_size: usize,
|
||||
flush_interval: Duration,
|
||||
}
|
||||
|
||||
impl EgressTunables {
|
||||
fn from_env() -> Self {
|
||||
Self {
|
||||
channel_capacity: env_positive(
|
||||
"LITELLM_LOG_CHANNEL_CAPACITY",
|
||||
DEFAULT_CHANNEL_CAPACITY,
|
||||
),
|
||||
max_batch_size: env_positive("LITELLM_LOG_BATCH_SIZE", DEFAULT_MAX_BATCH_SIZE),
|
||||
flush_interval: Duration::from_millis(env_positive(
|
||||
"LITELLM_LOG_FLUSH_INTERVAL_MS",
|
||||
DEFAULT_FLUSH_INTERVAL_MS,
|
||||
)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Parse a positive integer env var, falling back to `default` on missing,
|
||||
/// unparseable, or non-positive values. Generic over the integer type so one
|
||||
/// helper serves both the `usize` capacities and the `u64` interval.
|
||||
fn env_positive<T>(name: &str, default: T) -> T
|
||||
where
|
||||
T: std::str::FromStr + PartialOrd + From<u8>,
|
||||
{
|
||||
let zero = T::from(0u8);
|
||||
std::env::var(name)
|
||||
.ok()
|
||||
.and_then(|value| value.trim().parse::<T>().ok())
|
||||
.filter(|n| *n > zero)
|
||||
.unwrap_or(default)
|
||||
}
|
||||
|
||||
/// Ships realtime logging events to the LiteLLM Python proxy.
|
||||
pub struct LiteLLMPythonProxyAPILogger {
|
||||
sink: Sender<LogRecord>,
|
||||
}
|
||||
|
||||
impl LiteLLMPythonProxyAPILogger {
|
||||
/// Spawn the background worker and return a logger handle. `base` is the
|
||||
/// proxy base URL (no trailing path); `master_key` is sent as a bearer token.
|
||||
pub fn start(base: String, master_key: String) -> Arc<Self> {
|
||||
let tunables = EgressTunables::from_env();
|
||||
let (sink, receiver) = mpsc::channel::<LogRecord>(tunables.channel_capacity);
|
||||
let url = format!(
|
||||
"{}{}",
|
||||
base.trim_end_matches('/'),
|
||||
RUST_CONTROL_PLANE_LOGS_PATH
|
||||
);
|
||||
let client = Client::new();
|
||||
tokio::spawn(worker_loop(
|
||||
receiver,
|
||||
client,
|
||||
url,
|
||||
master_key,
|
||||
tunables.max_batch_size,
|
||||
tunables.flush_interval,
|
||||
));
|
||||
Arc::new(Self { sink })
|
||||
}
|
||||
|
||||
/// Build a logger from the environment: `LITELLM_PROXY_BASE_URL` (default
|
||||
/// `http://localhost:4000`) and `LITELLM_MASTER_KEY`.
|
||||
///
|
||||
/// `LITELLM_PROXY_BASE_URL` is treated as the full base and the route is
|
||||
/// appended verbatim, so if the proxy runs under a `SERVER_ROOT_PATH`
|
||||
/// (e.g. served at `https://host/litellm`), include it in the base
|
||||
/// (`LITELLM_PROXY_BASE_URL=https://host/litellm`) and the POST lands at
|
||||
/// `https://host/litellm/v1/rust_control_plane/logs`.
|
||||
pub fn from_env() -> Arc<Self> {
|
||||
let base = std::env::var("LITELLM_PROXY_BASE_URL")
|
||||
.ok()
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
.unwrap_or_else(|| DEFAULT_PROXY_BASE_URL.to_string());
|
||||
let key = std::env::var("LITELLM_MASTER_KEY").unwrap_or_default();
|
||||
Self::start(base, key)
|
||||
}
|
||||
|
||||
fn enqueue(&self, record: LogRecord) -> Result<(), LogError> {
|
||||
self.sink.try_send(record).map_err(|err| match err {
|
||||
mpsc::error::TrySendError::Full(_) => LogError::channel_full(),
|
||||
mpsc::error::TrySendError::Closed(_) => LogError::channel_closed(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl CustomLogger for LiteLLMPythonProxyAPILogger {
|
||||
fn log_success_event(&self, payload: &StandardLoggingPayload) -> Result<(), LogError> {
|
||||
self.enqueue(LogRecord {
|
||||
status: "success".to_string(),
|
||||
payload: payload.clone(),
|
||||
error: None,
|
||||
})
|
||||
}
|
||||
|
||||
fn log_failure_event(
|
||||
&self,
|
||||
payload: &StandardLoggingPayload,
|
||||
error: &LoggingError,
|
||||
) -> Result<(), LogError> {
|
||||
self.enqueue(LogRecord {
|
||||
status: "failure".to_string(),
|
||||
payload: payload.clone(),
|
||||
error: Some(format!("{}: {}", error.kind, error.message)),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// Drain the channel, batching records and POSTing them to the proxy. Exits when
|
||||
/// the channel is closed (all senders dropped) and drained.
|
||||
async fn worker_loop(
|
||||
mut receiver: Receiver<LogRecord>,
|
||||
client: Client,
|
||||
url: String,
|
||||
master_key: String,
|
||||
max_batch_size: usize,
|
||||
flush_interval: Duration,
|
||||
) {
|
||||
let mut ticker = interval(flush_interval);
|
||||
let mut batch: Vec<LogRecord> = Vec::with_capacity(max_batch_size);
|
||||
|
||||
loop {
|
||||
tokio::select! {
|
||||
maybe_record = receiver.recv() => {
|
||||
match maybe_record {
|
||||
Some(record) => {
|
||||
batch.push(record);
|
||||
if batch.len() >= max_batch_size {
|
||||
flush(&client, &url, &master_key, &mut batch).await;
|
||||
}
|
||||
}
|
||||
None => {
|
||||
// Channel closed: flush remaining and exit.
|
||||
flush(&client, &url, &master_key, &mut batch).await;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
_ = ticker.tick() => {
|
||||
flush(&client, &url, &master_key, &mut batch).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// POST the current batch (if any), clearing it. Errors are logged, not fatal.
|
||||
async fn flush(client: &Client, url: &str, master_key: &str, batch: &mut Vec<LogRecord>) {
|
||||
if batch.is_empty() {
|
||||
return;
|
||||
}
|
||||
let records = std::mem::take(batch)
|
||||
.into_iter()
|
||||
.map(LogRecord::into_callback_record)
|
||||
.collect();
|
||||
let body = CallbackLogsRequest { records };
|
||||
|
||||
let response = client
|
||||
.post(url)
|
||||
.bearer_auth(master_key)
|
||||
.json(&body)
|
||||
.send()
|
||||
.await;
|
||||
|
||||
match response {
|
||||
Ok(resp) if resp.status().is_success() => {}
|
||||
Ok(resp) => {
|
||||
eprintln!(
|
||||
"litellm-ai-gateway: callback logs POST returned {} to {url}",
|
||||
resp.status()
|
||||
);
|
||||
}
|
||||
Err(err) => {
|
||||
eprintln!("litellm-ai-gateway: callback logs POST failed to {url}: {err}");
|
||||
}
|
||||
}
|
||||
}
|
||||
10
litellm-rust/crates/ai-gateway/src/integrations/mod.rs
Normal file
10
litellm-rust/crates/ai-gateway/src/integrations/mod.rs
Normal file
|
|
@ -0,0 +1,10 @@
|
|||
//! Pure-Rust logging integrations. Names map 1:1 to Python
|
||||
//! `litellm/integrations/`:
|
||||
//! - [`custom_logger::CustomLogger`] — the callback trait
|
||||
//! - [`litellm_python_proxy_api::LiteLLMPythonProxyAPILogger`] — ships events
|
||||
//! to the Python proxy's `/v1/callbacks/logs` endpoint
|
||||
//! - [`types`] — the typed `StandardLoggingPayload` wire contract
|
||||
|
||||
pub mod custom_logger;
|
||||
pub mod litellm_python_proxy_api;
|
||||
pub mod types;
|
||||
164
litellm-rust/crates/ai-gateway/src/integrations/types.rs
Normal file
164
litellm-rust/crates/ai-gateway/src/integrations/types.rs
Normal file
|
|
@ -0,0 +1,164 @@
|
|||
//! Typed payloads for the LiteLLM `/v1/callbacks/logs` realtime-logging contract.
|
||||
//!
|
||||
//! Field names below are the EXACT JSON keys the Python replay path + spend-logs
|
||||
//! builder read. Note the deliberate mix:
|
||||
//! - `startTime` / `endTime` are camelCase (epoch f64 seconds)
|
||||
//! - `response_cost` / `prompt_tokens` / etc. are snake_case
|
||||
//!
|
||||
//! Mirrors Python `litellm/integrations/` + the proxy `CallbackLogsRequest`
|
||||
//! contract 1:1.
|
||||
|
||||
use serde::Serialize;
|
||||
use serde_json::Value;
|
||||
use std::collections::HashMap;
|
||||
|
||||
/// Cumulative token usage for a realtime session.
|
||||
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
|
||||
pub struct Usage {
|
||||
pub prompt_tokens: u64,
|
||||
pub completion_tokens: u64,
|
||||
pub total_tokens: u64,
|
||||
}
|
||||
|
||||
/// Cost-attribution metadata threaded from the authenticated request.
|
||||
#[derive(Clone, Debug, Default)]
|
||||
pub struct RequestMetadata {
|
||||
pub user_api_key_hash: Option<String>,
|
||||
pub user_api_key_user_id: Option<String>,
|
||||
pub user_api_key_team_id: Option<String>,
|
||||
}
|
||||
|
||||
/// A logging-callback failure (e.g. a custom logger raised). Mirrors the Python
|
||||
/// failure-event shape: a message plus an exception kind/class name.
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct LoggingError {
|
||||
pub message: String,
|
||||
pub kind: String,
|
||||
}
|
||||
|
||||
/// A non-fatal error returned by a `CustomLogger` when it cannot enqueue an
|
||||
/// event (channel full or the background worker has shut down).
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct LogError {
|
||||
pub message: String,
|
||||
pub kind: String,
|
||||
}
|
||||
|
||||
impl LogError {
|
||||
pub fn channel_full() -> Self {
|
||||
Self {
|
||||
message: "logging channel is full; dropping record".to_string(),
|
||||
kind: "ChannelFull".to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn channel_closed() -> Self {
|
||||
Self {
|
||||
message: "logging channel is closed; worker has shut down".to_string(),
|
||||
kind: "ChannelClosed".to_string(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Display for LogError {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
write!(f, "{}: {}", self.kind, self.message)
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for LogError {}
|
||||
|
||||
/// Batch wrapper — the top-level request body.
|
||||
/// Matches Python `CallbackLogsRequest { records: list[CallbackLogRecord] }`.
|
||||
#[derive(Serialize)]
|
||||
pub struct CallbackLogsRequest {
|
||||
pub records: Vec<CallbackLogRecord>,
|
||||
}
|
||||
|
||||
/// One finished logging event.
|
||||
/// Matches `CallbackLogRecord { status, standard_logging_payload, error? }`.
|
||||
#[derive(Serialize)]
|
||||
pub struct CallbackLogRecord {
|
||||
/// "success" | "failure". On "failure", `error` (or payload.error_str)
|
||||
/// becomes the replayed exception string.
|
||||
pub status: String,
|
||||
|
||||
pub standard_logging_payload: StandardLoggingPayload,
|
||||
|
||||
/// Only meaningful when status == "failure". Omitted on success.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub error: Option<String>,
|
||||
}
|
||||
|
||||
/// The self-describing payload. Field names are the EXACT JSON keys the Python
|
||||
/// replay path + spend-logs builder read.
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
pub struct StandardLoggingPayload {
|
||||
pub id: String,
|
||||
pub litellm_call_id: String,
|
||||
|
||||
/// e.g. "realtime", "acompletion". Falls back to "acompletion" if absent.
|
||||
pub call_type: String,
|
||||
|
||||
pub model: String,
|
||||
pub custom_llm_provider: String,
|
||||
|
||||
/// Spend ($) written to LiteLLM_SpendLogs.spend.
|
||||
pub response_cost: f64,
|
||||
|
||||
pub prompt_tokens: u64,
|
||||
pub completion_tokens: u64,
|
||||
pub total_tokens: u64,
|
||||
|
||||
/// EPOCH SECONDS as float — camelCase keys, NOT snake_case.
|
||||
#[serde(rename = "startTime")]
|
||||
pub start_time: f64,
|
||||
#[serde(rename = "endTime")]
|
||||
pub end_time: f64,
|
||||
|
||||
pub stream: bool,
|
||||
|
||||
pub metadata: StandardLoggingMetadata,
|
||||
|
||||
/// Optional; stored as request input on the spend log row.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub messages: Option<Value>,
|
||||
}
|
||||
|
||||
/// Cost-attribution keys. The replayer maps these into litellm_params.metadata,
|
||||
/// which the spend-logs builder reads to set user / team_id / organization_id.
|
||||
#[derive(Clone, Debug, Serialize, Default)]
|
||||
pub struct StandardLoggingMetadata {
|
||||
pub user_api_key_hash: Option<String>, // -> SpendLogs.api_key
|
||||
pub user_api_key_user_id: Option<String>, // -> SpendLogs.user
|
||||
pub user_api_key_team_id: Option<String>, // -> SpendLogs.team_id
|
||||
|
||||
// Optional but read by the builder; include when known:
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub user_api_key_alias: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub user_api_key_org_id: Option<String>, // -> SpendLogs.organization_id
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub user_api_key_end_user_id: Option<String>, // -> SpendLogs.end_user
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub spend_logs_metadata: Option<HashMap<String, Value>>,
|
||||
}
|
||||
|
||||
/// The unit handed to a `CustomLogger` sink: a finished payload plus its status
|
||||
/// and (on failure) the replayed error string.
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct LogRecord {
|
||||
pub status: String,
|
||||
pub payload: StandardLoggingPayload,
|
||||
pub error: Option<String>,
|
||||
}
|
||||
|
||||
impl LogRecord {
|
||||
pub fn into_callback_record(self) -> CallbackLogRecord {
|
||||
CallbackLogRecord {
|
||||
status: self.status,
|
||||
standard_logging_payload: self.payload,
|
||||
error: self.error,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -126,6 +126,9 @@ pub(crate) async fn read_event(upstream_rx: &mut UpstreamRx) -> CoreResult<Realt
|
|||
/// `session.created` here; the fresh-dial path passes `None` and lets the upstream
|
||||
/// deliver it). Then a single select loop forwards both directions through the
|
||||
/// transforms until either side closes or the idle timeout fires.
|
||||
/// `observe` is invoked on **upstream→client** events only (the trusted side that
|
||||
/// carries `session.created` and `response.done` usage) — never on client events,
|
||||
/// so a client cannot fabricate usage into its own logs.
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub(crate) async fn splice<In, Out>(
|
||||
model: &str,
|
||||
|
|
@ -133,6 +136,7 @@ pub(crate) async fn splice<In, Out>(
|
|||
mut upstream_rx: UpstreamRx,
|
||||
prelude: Option<RealtimeEvent>,
|
||||
idle_timeout: Option<Duration>,
|
||||
mut observe: impl FnMut(&RealtimeEvent) + Send,
|
||||
mut client_in: In,
|
||||
mut client_out: Out,
|
||||
) -> CoreResult<()>
|
||||
|
|
@ -165,6 +169,10 @@ where
|
|||
// client -> upstream
|
||||
client_event = client_in.next() => {
|
||||
let Some(event) = client_event else { break }; // client disconnected
|
||||
// NOTE: do NOT observe client events. session.created / response.done
|
||||
// (carrying usage) are server→client events; observing the client arm
|
||||
// would let an authenticated client POST a fabricated response.done and
|
||||
// inflate its own spend log. Logging observes upstream events only.
|
||||
for outbound in config.transform_realtime_request(&event, model)?.events {
|
||||
let payload = serde_json::to_string(&outbound)
|
||||
.map_err(|err| CoreError::InvalidResponse(err.to_string()))?;
|
||||
|
|
@ -181,6 +189,7 @@ where
|
|||
Message::Text(text) => {
|
||||
let event: RealtimeEvent = serde_json::from_str(&text)
|
||||
.map_err(|err| CoreError::InvalidResponse(err.to_string()))?;
|
||||
observe(&event);
|
||||
for outbound in config.transform_realtime_response(&event, model)?.events {
|
||||
client_out
|
||||
.send(outbound)
|
||||
|
|
@ -207,11 +216,13 @@ where
|
|||
/// framework-agnostic; the gateway adapts its axum socket to these. This is the
|
||||
/// fresh-dial path: dial, then splice. The pool's warm-handoff path skips the dial
|
||||
/// and calls [`splice`] directly with a buffered `session.created`.
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub async fn realtime<In, Out>(
|
||||
model: &str,
|
||||
api_key: Option<&str>,
|
||||
api_base: Option<&str>,
|
||||
idle_timeout: Option<Duration>,
|
||||
observe: impl FnMut(&RealtimeEvent) + Send,
|
||||
client_in: In,
|
||||
client_out: Out,
|
||||
) -> CoreResult<()>
|
||||
|
|
@ -229,6 +240,7 @@ where
|
|||
upstream_rx,
|
||||
None,
|
||||
idle_timeout,
|
||||
observe,
|
||||
client_in,
|
||||
client_out,
|
||||
)
|
||||
|
|
@ -238,10 +250,12 @@ where
|
|||
/// Splice a pre-warmed upstream (taken from [`crate::io::realtime_pool`]) to the
|
||||
/// client. Relays the buffered `session.created` first, then splices exactly like
|
||||
/// the fresh-dial path — so a warm session is indistinguishable from a fresh one.
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub async fn realtime_warm<In, Out>(
|
||||
model: &str,
|
||||
handoff: crate::io::realtime_pool::WarmHandoff,
|
||||
idle_timeout: Option<Duration>,
|
||||
observe: impl FnMut(&RealtimeEvent) + Send,
|
||||
client_in: In,
|
||||
client_out: Out,
|
||||
) -> CoreResult<()>
|
||||
|
|
@ -256,6 +270,7 @@ where
|
|||
handoff.rx,
|
||||
Some(handoff.session_created),
|
||||
idle_timeout,
|
||||
observe,
|
||||
client_in,
|
||||
client_out,
|
||||
)
|
||||
|
|
@ -303,6 +318,7 @@ mod tests {
|
|||
Some(&key_owned),
|
||||
None,
|
||||
None,
|
||||
|_| {},
|
||||
client_in,
|
||||
client_out,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -18,6 +18,8 @@ use litellm_ai_gateway::routes;
|
|||
use litellm_ai_gateway::state::AppState;
|
||||
use litellm_core::router::{Deployment, LiteLLMParams, Router};
|
||||
|
||||
use litellm_ai_gateway::integrations::custom_logger::CustomLogger;
|
||||
use litellm_ai_gateway::integrations::litellm_python_proxy_api::LiteLLMPythonProxyAPILogger;
|
||||
#[cfg(feature = "python-config")]
|
||||
use litellm_ai_gateway::python;
|
||||
|
||||
|
|
@ -44,6 +46,12 @@ async fn main() {
|
|||
);
|
||||
}
|
||||
|
||||
// Spawn the realtime-logging worker (drains a channel → POSTs batches to the
|
||||
// Python proxy's /v1/callbacks/logs). Built here so the spawn lands on the
|
||||
// tokio runtime. `from_env` reads LITELLM_PROXY_BASE_URL + LITELLM_MASTER_KEY.
|
||||
let proxy_logger = LiteLLMPythonProxyAPILogger::from_env();
|
||||
let loggers: Vec<Arc<dyn CustomLogger>> = vec![proxy_logger];
|
||||
|
||||
let router = Arc::new(build_router());
|
||||
|
||||
// Build the pre-warmed realtime pool and register each deployment's upstream
|
||||
|
|
@ -85,6 +93,7 @@ async fn main() {
|
|||
let state = AppState {
|
||||
router,
|
||||
master_key,
|
||||
loggers: Arc::new(loggers),
|
||||
realtime_pool,
|
||||
authenticator,
|
||||
key_cache,
|
||||
|
|
|
|||
4
litellm-rust/crates/ai-gateway/src/realtime/mod.rs
Normal file
4
litellm-rust/crates/ai-gateway/src/realtime/mod.rs
Normal file
|
|
@ -0,0 +1,4 @@
|
|||
//! Realtime logging collector. Observes the realtime event stream and emits a
|
||||
//! `StandardLoggingPayload` to the registered callbacks on session close.
|
||||
|
||||
pub mod streaming;
|
||||
352
litellm-rust/crates/ai-gateway/src/realtime/streaming.rs
Normal file
352
litellm-rust/crates/ai-gateway/src/realtime/streaming.rs
Normal file
|
|
@ -0,0 +1,352 @@
|
|||
//! `RealTimeStreaming` — the realtime logging collector.
|
||||
//!
|
||||
//! Mirrors Python `litellm.realtime_api.main.RealTimeStreaming`: it observes the
|
||||
//! event stream in O(1) (never buffering frames), accumulating just the fields
|
||||
//! the spend log needs (model, id, cumulative usage), then on session close
|
||||
//! builds a `StandardLoggingPayload` and fans it out to every registered
|
||||
//! `CustomLogger`.
|
||||
|
||||
use std::sync::Arc;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use litellm_core::realtime::types::RealtimeEvent;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::constants::DEFAULT_PROVIDER;
|
||||
use crate::integrations::custom_logger::CustomLogger;
|
||||
use crate::integrations::types::{
|
||||
RequestMetadata, StandardLoggingMetadata, StandardLoggingPayload, Usage,
|
||||
};
|
||||
|
||||
/// Current wall-clock time as epoch seconds (float), matching the Python
|
||||
/// `startTime`/`endTime` contract.
|
||||
fn epoch_seconds() -> f64 {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map(|d| d.as_secs_f64())
|
||||
.unwrap_or(0.0)
|
||||
}
|
||||
|
||||
/// Status of a finished realtime session, mapped to the callback record status.
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum SessionStatus {
|
||||
Success,
|
||||
Failure,
|
||||
}
|
||||
|
||||
/// Accumulates realtime session state and emits a logging payload on close.
|
||||
pub struct RealTimeStreaming {
|
||||
callbacks: Vec<Arc<dyn CustomLogger>>,
|
||||
/// REQUEST-ID RULE: the SpendLogs `request_id` == the OpenAI realtime session
|
||||
/// id (`sess_…`), captured from `session.created`. Both `id` and
|
||||
/// `litellm_call_id` are set to that value so the Python writer logs the same
|
||||
/// id regardless of which field it reads. The gateway-generated `rt-…` id
|
||||
/// (the constructor seed) is only a fallback for sessions that fail before
|
||||
/// `session.created` arrives.
|
||||
litellm_call_id: String,
|
||||
/// See the request-id rule above — mirrors `litellm_call_id`.
|
||||
id: String,
|
||||
model: String,
|
||||
custom_llm_provider: String,
|
||||
usage: Usage,
|
||||
response_cost: f64,
|
||||
start_time: f64,
|
||||
end_time: f64,
|
||||
metadata: RequestMetadata,
|
||||
/// Count of logging callbacks that failed to enqueue (non-fatal).
|
||||
dropped: u64,
|
||||
}
|
||||
|
||||
impl RealTimeStreaming {
|
||||
/// Create a collector for one session. `litellm_call_id` is the gateway's
|
||||
/// per-connection id; `model` is the requested model (a sane default until
|
||||
/// `session.created` reports the upstream model).
|
||||
pub fn new(
|
||||
callbacks: Vec<Arc<dyn CustomLogger>>,
|
||||
litellm_call_id: String,
|
||||
model: String,
|
||||
metadata: RequestMetadata,
|
||||
) -> Self {
|
||||
let now = epoch_seconds();
|
||||
Self {
|
||||
callbacks,
|
||||
id: litellm_call_id.clone(),
|
||||
litellm_call_id,
|
||||
model,
|
||||
custom_llm_provider: DEFAULT_PROVIDER.to_string(),
|
||||
usage: Usage::default(),
|
||||
response_cost: 0.0,
|
||||
start_time: now,
|
||||
end_time: now,
|
||||
metadata,
|
||||
dropped: 0,
|
||||
}
|
||||
}
|
||||
|
||||
/// Number of logging callbacks that failed to enqueue so far (test/observ.).
|
||||
#[allow(dead_code)]
|
||||
pub fn dropped(&self) -> u64 {
|
||||
self.dropped
|
||||
}
|
||||
|
||||
/// Observe one realtime event. O(1): updates accumulated state only; never
|
||||
/// buffers frames. Safe to call on every event in either direction.
|
||||
pub fn observe(&mut self, event: &RealtimeEvent) {
|
||||
match event.event_type.as_str() {
|
||||
"session.created" | "session.updated" => self.on_session(event),
|
||||
"response.done" => self.on_response_done(event),
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
/// `session.created` / `session.updated` → capture upstream id + model.
|
||||
/// Per the request-id rule, the OpenAI session id becomes BOTH `id` and
|
||||
/// `litellm_call_id`, replacing the gateway-generated fallback.
|
||||
fn on_session(&mut self, event: &RealtimeEvent) {
|
||||
let session = event.data.get("session").and_then(Value::as_object);
|
||||
if let Some(id) = session.and_then(|s| s.get("id")).and_then(Value::as_str) {
|
||||
if !id.is_empty() {
|
||||
self.id = id.to_string();
|
||||
self.litellm_call_id = id.to_string();
|
||||
}
|
||||
}
|
||||
if let Some(model) = session.and_then(|s| s.get("model")).and_then(Value::as_str) {
|
||||
if !model.is_empty() {
|
||||
self.model = model.to_string();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// `response.done` → add this response's usage to the cumulative totals.
|
||||
fn on_response_done(&mut self, event: &RealtimeEvent) {
|
||||
let usage = event
|
||||
.data
|
||||
.get("response")
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|r| r.get("usage"))
|
||||
.and_then(Value::as_object);
|
||||
let Some(usage) = usage else { return };
|
||||
|
||||
let input = usage.get("input_tokens").and_then(Value::as_u64);
|
||||
let output = usage.get("output_tokens").and_then(Value::as_u64);
|
||||
let total = usage.get("total_tokens").and_then(Value::as_u64);
|
||||
|
||||
if let Some(input) = input {
|
||||
self.usage.prompt_tokens += input;
|
||||
}
|
||||
if let Some(output) = output {
|
||||
self.usage.completion_tokens += output;
|
||||
}
|
||||
// Prefer the upstream-reported total; otherwise derive it.
|
||||
match total {
|
||||
Some(total) => self.usage.total_tokens += total,
|
||||
None => {
|
||||
self.usage.total_tokens += input.unwrap_or(0) + output.unwrap_or(0);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Set the per-session response cost ($). Cost computation is Python-side in
|
||||
/// the proxy; the gateway forwards 0.0 by default and lets the proxy price.
|
||||
/// Public API (exercised in tests) for the future path where the gateway
|
||||
/// prices realtime sessions itself.
|
||||
#[allow(dead_code)]
|
||||
pub fn set_response_cost(&mut self, cost: f64) {
|
||||
self.response_cost = cost;
|
||||
}
|
||||
|
||||
/// Build the `StandardLoggingPayload` from accumulated state.
|
||||
pub fn build_payload(&self) -> StandardLoggingPayload {
|
||||
StandardLoggingPayload {
|
||||
id: self.id.clone(),
|
||||
litellm_call_id: self.litellm_call_id.clone(),
|
||||
call_type: "realtime".to_string(),
|
||||
model: self.model.clone(),
|
||||
custom_llm_provider: self.custom_llm_provider.clone(),
|
||||
response_cost: self.response_cost,
|
||||
prompt_tokens: self.usage.prompt_tokens,
|
||||
completion_tokens: self.usage.completion_tokens,
|
||||
total_tokens: self.usage.total_tokens,
|
||||
start_time: self.start_time,
|
||||
end_time: self.end_time,
|
||||
stream: true,
|
||||
metadata: StandardLoggingMetadata {
|
||||
user_api_key_hash: self.metadata.user_api_key_hash.clone(),
|
||||
user_api_key_user_id: self.metadata.user_api_key_user_id.clone(),
|
||||
user_api_key_team_id: self.metadata.user_api_key_team_id.clone(),
|
||||
..Default::default()
|
||||
},
|
||||
messages: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Finish the session: stamp the end time and fan the payload out to every
|
||||
/// callback. On a logger enqueue error we bump a non-fatal counter (the
|
||||
/// realtime session has already ended; a dropped log must never propagate).
|
||||
pub fn log_messages(&mut self, status: SessionStatus) {
|
||||
self.end_time = epoch_seconds();
|
||||
let payload = self.build_payload();
|
||||
|
||||
match status {
|
||||
SessionStatus::Success => {
|
||||
for callback in &self.callbacks {
|
||||
if let Err(err) = callback.log_success_event(&payload) {
|
||||
self.dropped += 1;
|
||||
eprintln!("litellm-ai-gateway: log_success_event dropped: {err}");
|
||||
}
|
||||
}
|
||||
}
|
||||
SessionStatus::Failure => {
|
||||
let error = crate::integrations::types::LoggingError {
|
||||
message: "realtime session ended in failure".to_string(),
|
||||
kind: "RealtimeSessionError".to_string(),
|
||||
};
|
||||
for callback in &self.callbacks {
|
||||
if let Err(err) = callback.log_failure_event(&payload, &error) {
|
||||
self.dropped += 1;
|
||||
eprintln!("litellm-ai-gateway: log_failure_event dropped: {err}");
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::integrations::types::{LogError, LoggingError};
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
|
||||
fn event(raw: &str) -> RealtimeEvent {
|
||||
serde_json::from_str(raw).expect("valid event json")
|
||||
}
|
||||
|
||||
/// A test logger that records the last payload it saw.
|
||||
#[derive(Default)]
|
||||
struct CapturingLogger {
|
||||
calls: AtomicU64,
|
||||
last_model: std::sync::Mutex<Option<String>>,
|
||||
last_total_tokens: AtomicU64,
|
||||
}
|
||||
|
||||
impl CustomLogger for CapturingLogger {
|
||||
fn log_success_event(&self, payload: &StandardLoggingPayload) -> Result<(), LogError> {
|
||||
self.calls.fetch_add(1, Ordering::SeqCst);
|
||||
*self.last_model.lock().unwrap() = Some(payload.model.clone());
|
||||
self.last_total_tokens
|
||||
.store(payload.total_tokens, Ordering::SeqCst);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn observe_accumulates_model_and_tokens_then_logs() {
|
||||
let logger = Arc::new(CapturingLogger::default());
|
||||
let callbacks: Vec<Arc<dyn CustomLogger>> = vec![logger.clone()];
|
||||
let mut streaming = RealTimeStreaming::new(
|
||||
callbacks,
|
||||
"call_abc".to_string(),
|
||||
"gpt-realtime".to_string(),
|
||||
RequestMetadata {
|
||||
user_api_key_hash: Some("hash123".to_string()),
|
||||
user_api_key_user_id: Some("user-1".to_string()),
|
||||
user_api_key_team_id: Some("team-1".to_string()),
|
||||
},
|
||||
);
|
||||
|
||||
streaming.observe(&event(
|
||||
r#"{"type":"session.created","session":{"id":"sess_001","model":"gpt-realtime-2025"}}"#,
|
||||
));
|
||||
streaming.observe(&event(
|
||||
r#"{"type":"response.done","response":{"usage":{"input_tokens":10,"output_tokens":5,"total_tokens":15}}}"#,
|
||||
));
|
||||
// A second response.done accumulates.
|
||||
streaming.observe(&event(
|
||||
r#"{"type":"response.done","response":{"usage":{"input_tokens":3,"output_tokens":2,"total_tokens":5}}}"#,
|
||||
));
|
||||
|
||||
let payload = streaming.build_payload();
|
||||
assert_eq!(payload.model, "gpt-realtime-2025");
|
||||
// Request-id rule: session.created's id becomes BOTH id and
|
||||
// litellm_call_id (replacing the "call_abc" gateway fallback), so the
|
||||
// SpendLogs request_id is always the OpenAI session id.
|
||||
assert_eq!(payload.id, "sess_001");
|
||||
assert_eq!(payload.litellm_call_id, "sess_001");
|
||||
assert_eq!(payload.prompt_tokens, 13);
|
||||
assert_eq!(payload.completion_tokens, 7);
|
||||
assert_eq!(payload.total_tokens, 20);
|
||||
assert_eq!(payload.response_cost, 0.0);
|
||||
assert_eq!(payload.call_type, "realtime");
|
||||
assert_eq!(payload.custom_llm_provider, "openai");
|
||||
assert_eq!(
|
||||
payload.metadata.user_api_key_hash.as_deref(),
|
||||
Some("hash123")
|
||||
);
|
||||
|
||||
streaming.log_messages(SessionStatus::Success);
|
||||
assert_eq!(logger.calls.load(Ordering::SeqCst), 1);
|
||||
assert_eq!(
|
||||
logger.last_model.lock().unwrap().as_deref(),
|
||||
Some("gpt-realtime-2025")
|
||||
);
|
||||
assert_eq!(logger.last_total_tokens.load(Ordering::SeqCst), 20);
|
||||
assert_eq!(streaming.dropped(), 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn payload_serializes_with_camelcase_times_and_realtime_call_type() {
|
||||
let mut streaming = RealTimeStreaming::new(
|
||||
Vec::new(),
|
||||
"call_xyz".to_string(),
|
||||
"gpt-realtime".to_string(),
|
||||
RequestMetadata::default(),
|
||||
);
|
||||
streaming.observe(&event(
|
||||
r#"{"type":"response.done","response":{"usage":{"input_tokens":1,"output_tokens":1,"total_tokens":2}}}"#,
|
||||
));
|
||||
streaming.set_response_cost(0.0042);
|
||||
let payload = streaming.build_payload();
|
||||
let json = serde_json::to_string(&payload).expect("serialize payload");
|
||||
|
||||
assert!(json.contains("\"startTime\""), "missing startTime: {json}");
|
||||
assert!(json.contains("\"endTime\""), "missing endTime: {json}");
|
||||
assert!(
|
||||
json.contains("\"call_type\":\"realtime\""),
|
||||
"missing call_type realtime: {json}"
|
||||
);
|
||||
assert!(
|
||||
json.contains("\"response_cost\""),
|
||||
"missing response_cost: {json}"
|
||||
);
|
||||
assert_eq!(payload.response_cost, 0.0042);
|
||||
}
|
||||
|
||||
/// A logger whose enqueue always fails should bump the dropped counter, not
|
||||
/// panic or propagate.
|
||||
#[test]
|
||||
fn failing_logger_bumps_dropped_counter() {
|
||||
struct FailingLogger;
|
||||
impl CustomLogger for FailingLogger {
|
||||
fn log_success_event(&self, _p: &StandardLoggingPayload) -> Result<(), LogError> {
|
||||
Err(LogError::channel_full())
|
||||
}
|
||||
fn log_failure_event(
|
||||
&self,
|
||||
_p: &StandardLoggingPayload,
|
||||
_e: &LoggingError,
|
||||
) -> Result<(), LogError> {
|
||||
Err(LogError::channel_closed())
|
||||
}
|
||||
}
|
||||
let callbacks: Vec<Arc<dyn CustomLogger>> = vec![Arc::new(FailingLogger)];
|
||||
let mut streaming = RealTimeStreaming::new(
|
||||
callbacks,
|
||||
"call_1".to_string(),
|
||||
"gpt-realtime".to_string(),
|
||||
RequestMetadata::default(),
|
||||
);
|
||||
streaming.log_messages(SessionStatus::Success);
|
||||
assert_eq!(streaming.dropped(), 1);
|
||||
}
|
||||
}
|
||||
|
|
@ -5,7 +5,9 @@
|
|||
|
||||
mod service;
|
||||
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use std::sync::Arc;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use crate::io::realtime_pool::RealtimePool;
|
||||
use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
|
||||
|
|
@ -20,8 +22,26 @@ use litellm_core::router::Router as ModelRouter;
|
|||
use serde::Deserialize;
|
||||
|
||||
use crate::auth::UserApiKeyAuth;
|
||||
use crate::integrations::custom_logger::CustomLogger;
|
||||
use crate::integrations::types::RequestMetadata;
|
||||
use crate::realtime::streaming::{RealTimeStreaming, SessionStatus};
|
||||
use crate::state::AppState;
|
||||
|
||||
/// Process-local monotonic counter, mixed into the per-session call id so two
|
||||
/// sessions opened in the same nanosecond still get distinct ids.
|
||||
static CALL_SEQ: AtomicU64 = AtomicU64::new(0);
|
||||
|
||||
/// Generate a per-connection `litellm_call_id`. No external uuid dep: epoch
|
||||
/// nanos + a process-local sequence is unique enough for log correlation.
|
||||
fn new_call_id() -> String {
|
||||
let nanos = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map(|d| d.as_nanos())
|
||||
.unwrap_or(0);
|
||||
let seq = CALL_SEQ.fetch_add(1, Ordering::Relaxed);
|
||||
format!("rt-{nanos:x}-{seq:x}")
|
||||
}
|
||||
|
||||
/// This route's contribution to the app router.
|
||||
pub fn router() -> Router<AppState> {
|
||||
Router::new().route("/v1/realtime", get(handle))
|
||||
|
|
@ -32,31 +52,30 @@ struct RealtimeQuery {
|
|||
model: String,
|
||||
}
|
||||
|
||||
fn require_realtime_billing_safe_auth(auth: &UserApiKeyAuth) -> Result<(), (StatusCode, String)> {
|
||||
if auth.is_proxy_admin() {
|
||||
return Ok(());
|
||||
fn request_metadata_for_auth(auth: &UserApiKeyAuth, master_key: Option<&str>) -> RequestMetadata {
|
||||
let user_api_key_hash = auth.api_key.clone().or_else(|| {
|
||||
auth.is_proxy_admin()
|
||||
.then(|| master_key.map(crate::auth::hash_token))
|
||||
.flatten()
|
||||
});
|
||||
|
||||
RequestMetadata {
|
||||
user_api_key_hash,
|
||||
user_api_key_user_id: auth.user_id.clone(),
|
||||
user_api_key_team_id: auth.team_id.clone(),
|
||||
}
|
||||
Err((
|
||||
StatusCode::FORBIDDEN,
|
||||
"non-admin virtual-key realtime is disabled until realtime usage is reported to the control plane".to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
/// Auth runs via the `UserApiKeyAuth` extractor (master key → admin, otherwise
|
||||
/// cache then the swappable authenticator). Non-admin virtual keys are rejected
|
||||
/// until realtime usage is reported back to the proxy; otherwise callers could
|
||||
/// consume upstream realtime spend without that spend being charged to their key
|
||||
/// or team. We validate the model BEFORE the upgrade so failures are clean HTTP
|
||||
/// (400/404), not a socket that opens then closes, then hand the socket to
|
||||
/// `bridge`.
|
||||
/// cache then the swappable authenticator). We validate the model BEFORE the
|
||||
/// upgrade so failures are clean HTTP (400/404), not a socket that opens then
|
||||
/// closes, then hand the socket to `bridge`.
|
||||
async fn handle(
|
||||
auth: UserApiKeyAuth,
|
||||
ws: WebSocketUpgrade,
|
||||
State(state): State<AppState>,
|
||||
Query(query): Query<RealtimeQuery>,
|
||||
) -> Result<Response, (StatusCode, String)> {
|
||||
require_realtime_billing_safe_auth(&auth)?;
|
||||
|
||||
if query.model.trim().is_empty() {
|
||||
return Err((
|
||||
StatusCode::BAD_REQUEST,
|
||||
|
|
@ -72,26 +91,58 @@ async fn handle(
|
|||
|
||||
let router = state.router.clone();
|
||||
let pool = state.realtime_pool.clone();
|
||||
let loggers = state.loggers.clone();
|
||||
let master_key = state.master_key.clone();
|
||||
let model = query.model;
|
||||
Ok(ws.on_upgrade(move |socket| bridge(socket, router, pool, model)))
|
||||
Ok(ws.on_upgrade(move |socket| bridge(socket, router, pool, loggers, auth, master_key, model)))
|
||||
}
|
||||
|
||||
/// Adapt the axum socket (text frames) to the typed-event `Stream`/`Sink` the
|
||||
/// service wants, keeping axum types out of `service`.
|
||||
///
|
||||
/// This is also the realtime-logging seam: every upstream→client event (the
|
||||
/// direction carrying `session.created` and `response.done` with usage) is fed
|
||||
/// to a [`RealTimeStreaming`] collector via the splice's `observe` callback. The
|
||||
/// observe is O(1) and never buffers frames. When the splice returns (any of the
|
||||
/// three break paths — client disconnect, upstream close, idle timeout), we flush
|
||||
/// one logging payload to the registered callbacks.
|
||||
async fn bridge(
|
||||
socket: WebSocket,
|
||||
router: Arc<ModelRouter>,
|
||||
pool: Arc<RealtimePool>,
|
||||
loggers: Arc<Vec<Arc<dyn CustomLogger>>>,
|
||||
auth: UserApiKeyAuth,
|
||||
master_key: Option<Arc<str>>,
|
||||
model: String,
|
||||
) {
|
||||
let (ws_sink, ws_stream) = socket.split();
|
||||
|
||||
// Attribute the spend log to the key that authenticated this session. For a
|
||||
// virtual key, the Python verifier returns the already-hashed key + user/team
|
||||
// attribution. For the local master-key fast path, hash the master key here.
|
||||
// A non-null user_api_key_hash is required for the Python spend logger to
|
||||
// write a SpendLogs row.
|
||||
let metadata = request_metadata_for_auth(&auth, master_key.as_deref());
|
||||
|
||||
// Owned by THIS task only. The splice observes it via a synchronous `&mut`
|
||||
// callback (below), so there is no Arc/Mutex/atomic on the per-frame hot
|
||||
// path — just a monomorphized FnMut mutating stack-local fields. This is
|
||||
// what lets observe scale: 10K concurrent sessions = 10K independent
|
||||
// collectors, zero cross-task synchronization.
|
||||
let mut collector = RealTimeStreaming::new(
|
||||
loggers.as_ref().clone(),
|
||||
new_call_id(),
|
||||
model.clone(),
|
||||
metadata,
|
||||
);
|
||||
|
||||
let client_in = ws_stream.filter_map(|message| async move {
|
||||
match message {
|
||||
Ok(Message::Text(text)) => serde_json::from_str::<RealtimeEvent>(&text).ok(),
|
||||
_ => None,
|
||||
}
|
||||
});
|
||||
// Plain forwarding sink — no observe here anymore.
|
||||
let client_out = ws_sink.with(|event: RealtimeEvent| async move {
|
||||
Ok::<Message, axum::Error>(Message::Text(
|
||||
serde_json::to_string(&event).unwrap_or_default(),
|
||||
|
|
@ -99,7 +150,28 @@ async fn bridge(
|
|||
});
|
||||
|
||||
futures_util::pin_mut!(client_in, client_out);
|
||||
let _ = service::run(&router, &pool, &model, None, client_in, client_out).await;
|
||||
|
||||
// The observe closure borrows `&mut collector` for the duration of the
|
||||
// splice; the borrow ends when `run` returns, freeing the collector for the
|
||||
// single post-session `log_messages` flush. `run` picks a pooled (warm) or
|
||||
// fresh upstream — observe fires on the upstream arm either way.
|
||||
let result = service::run(
|
||||
&router,
|
||||
&pool,
|
||||
&model,
|
||||
None,
|
||||
|event: &RealtimeEvent| collector.observe(event),
|
||||
client_in,
|
||||
client_out,
|
||||
)
|
||||
.await;
|
||||
|
||||
let status = if result.is_ok() {
|
||||
SessionStatus::Success
|
||||
} else {
|
||||
SessionStatus::Failure
|
||||
};
|
||||
collector.log_messages(status);
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
|
@ -107,19 +179,30 @@ mod tests {
|
|||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn realtime_allows_proxy_admin_identity() {
|
||||
assert!(require_realtime_billing_safe_auth(&UserApiKeyAuth::admin()).is_ok());
|
||||
fn request_metadata_uses_virtual_key_identity() {
|
||||
let metadata = request_metadata_for_auth(
|
||||
&UserApiKeyAuth {
|
||||
api_key: Some("hashed-key".to_string()),
|
||||
user_id: Some("user-1".to_string()),
|
||||
team_id: Some("team-1".to_string()),
|
||||
..UserApiKeyAuth::default()
|
||||
},
|
||||
Some("sk-master"),
|
||||
);
|
||||
|
||||
assert_eq!(metadata.user_api_key_hash.as_deref(), Some("hashed-key"));
|
||||
assert_eq!(metadata.user_api_key_user_id.as_deref(), Some("user-1"));
|
||||
assert_eq!(metadata.user_api_key_team_id.as_deref(), Some("team-1"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn realtime_rejects_non_admin_virtual_keys_until_spend_is_reported() {
|
||||
let err = require_realtime_billing_safe_auth(&UserApiKeyAuth {
|
||||
user_role: Some("internal_user".to_string()),
|
||||
..UserApiKeyAuth::default()
|
||||
})
|
||||
.unwrap_err();
|
||||
fn request_metadata_hashes_master_key_for_local_admin_fast_path() {
|
||||
let metadata = request_metadata_for_auth(&UserApiKeyAuth::admin(), Some("sk-master"));
|
||||
let expected_hash = crate::auth::hash_token("sk-master");
|
||||
|
||||
assert_eq!(err.0, StatusCode::FORBIDDEN);
|
||||
assert!(err.1.contains("usage is reported"));
|
||||
assert_eq!(
|
||||
metadata.user_api_key_hash.as_deref(),
|
||||
Some(expected_hash.as_str())
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ pub async fn run<In, Out>(
|
|||
pool: &RealtimePool,
|
||||
model: &str,
|
||||
idle_timeout: Option<Duration>,
|
||||
observe: impl FnMut(&RealtimeEvent) + Send,
|
||||
client_in: In,
|
||||
client_out: Out,
|
||||
) -> CoreResult<()>
|
||||
|
|
@ -56,6 +57,7 @@ where
|
|||
provider_model,
|
||||
handoff,
|
||||
idle_timeout,
|
||||
observe,
|
||||
client_in,
|
||||
client_out,
|
||||
)
|
||||
|
|
@ -69,6 +71,7 @@ where
|
|||
params.api_key.as_deref(),
|
||||
params.api_base.as_deref(),
|
||||
idle_timeout,
|
||||
observe,
|
||||
client_in,
|
||||
client_out,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -3,6 +3,8 @@ use std::sync::Arc;
|
|||
use crate::io::realtime_pool::RealtimePool;
|
||||
use litellm_core::router::Router;
|
||||
|
||||
use crate::integrations::custom_logger::CustomLogger;
|
||||
|
||||
/// Shared application state handed to every route handler.
|
||||
#[derive(Clone)]
|
||||
pub struct AppState {
|
||||
|
|
@ -10,6 +12,8 @@ pub struct AppState {
|
|||
/// The gateway master key. Any caller presenting it as a bearer token may
|
||||
/// invoke the gateway. `None` → auth not configured (routes fail closed).
|
||||
pub master_key: Option<Arc<str>>,
|
||||
/// Logging callbacks fanned out at the end of each realtime session.
|
||||
pub loggers: Arc<Vec<Arc<dyn CustomLogger>>>,
|
||||
/// Pre-warmed upstream realtime connection pool. Disabled
|
||||
/// (`RealtimePool::disabled()`) when `REALTIME_POOL_SIZE=0`, in which case
|
||||
/// every realtime connect fresh-dials exactly as before.
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ OPTIONAL_KWARGS_KEYS = frozenset(
|
|||
"azure_password",
|
||||
"azure_scope",
|
||||
"timeout",
|
||||
"gcs_bucket_name",
|
||||
"bucket_name",
|
||||
"vertex_credentials",
|
||||
"vertex_project",
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1,13 +1,14 @@
|
|||
import asyncio
|
||||
import json
|
||||
import ssl
|
||||
from functools import lru_cache
|
||||
from urllib.parse import parse_qs, urlencode, urlparse, urlunparse
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
AsyncIterator,
|
||||
Coroutine,
|
||||
Dict,
|
||||
Iterator,
|
||||
List,
|
||||
Literal,
|
||||
Optional,
|
||||
|
|
@ -16,6 +17,7 @@ from typing import (
|
|||
cast,
|
||||
get_type_hints,
|
||||
)
|
||||
from urllib.parse import parse_qs, urlencode, urlparse, urlunparse
|
||||
|
||||
import httpx # type: ignore
|
||||
from openai.types.file_deleted import FileDeleted
|
||||
|
|
@ -27,8 +29,8 @@ import litellm.types.utils
|
|||
from litellm._logging import _redact_string, verbose_logger
|
||||
from litellm.anthropic_beta_headers_manager import update_headers_with_filtered_beta
|
||||
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
|
||||
from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming
|
||||
from litellm.litellm_core_utils.asyncify import run_async_function
|
||||
from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
from litellm.llms.base_llm.anthropic_messages.transformation import (
|
||||
BaseAnthropicMessagesConfig,
|
||||
|
|
@ -107,6 +109,7 @@ from litellm.types.llms.openai import (
|
|||
ResponsesAPIOptionalRequestParams,
|
||||
ResponsesAPIResponse,
|
||||
)
|
||||
from litellm.types.realtime import RealtimeQueryParams
|
||||
from litellm.types.rerank import RerankResponse
|
||||
from litellm.types.responses.main import DeleteResponseResult
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
|
@ -132,7 +135,6 @@ from litellm.types.vector_stores import (
|
|||
VectorStoreSearchOptionalRequestParams,
|
||||
VectorStoreSearchResponse,
|
||||
)
|
||||
from litellm.types.realtime import RealtimeQueryParams
|
||||
from litellm.types.videos.main import VideoObject
|
||||
from litellm.utils import (
|
||||
CustomStreamWrapper,
|
||||
|
|
@ -3438,6 +3440,23 @@ class BaseLLMHTTPHandler:
|
|||
data=presigned_request["data"],
|
||||
timeout=timeout,
|
||||
)
|
||||
elif (
|
||||
isinstance(transformed_request, dict)
|
||||
and "resumable_chunked_upload" in transformed_request
|
||||
):
|
||||
try:
|
||||
upload_response = self._resumable_chunked_upload(
|
||||
client=sync_httpx_client,
|
||||
initiate_url=api_base,
|
||||
base_headers=headers,
|
||||
config=cast(Dict[str, Any], transformed_request)[
|
||||
"resumable_chunked_upload"
|
||||
],
|
||||
timeout=timeout,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"Error creating file: {e}")
|
||||
raise self._handle_error(e=e, provider_config=provider_config)
|
||||
elif isinstance(transformed_request, str) or isinstance(
|
||||
transformed_request, bytes
|
||||
):
|
||||
|
|
@ -3519,7 +3538,15 @@ class BaseLLMHTTPHandler:
|
|||
input="",
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": transformed_request,
|
||||
# A resumable upload config holds a reference to the (potentially
|
||||
# huge) upload payload; logging deep-copies additional_args, so log
|
||||
# a placeholder instead of re-materializing the payload.
|
||||
"complete_input_dict": (
|
||||
"<resumable chunked upload>"
|
||||
if isinstance(transformed_request, dict)
|
||||
and "resumable_chunked_upload" in transformed_request
|
||||
else transformed_request
|
||||
),
|
||||
"api_base": api_base,
|
||||
"headers": headers,
|
||||
},
|
||||
|
|
@ -3596,6 +3623,23 @@ class BaseLLMHTTPHandler:
|
|||
data=presigned_request["data"],
|
||||
timeout=timeout,
|
||||
)
|
||||
elif (
|
||||
isinstance(transformed_request, dict)
|
||||
and "resumable_chunked_upload" in transformed_request
|
||||
):
|
||||
try:
|
||||
upload_response = await self._aresumable_chunked_upload(
|
||||
client=async_httpx_client,
|
||||
initiate_url=api_base,
|
||||
base_headers=headers,
|
||||
config=cast(Dict[str, Any], transformed_request)[
|
||||
"resumable_chunked_upload"
|
||||
],
|
||||
timeout=timeout,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"Error creating file: {e}")
|
||||
raise self._handle_error(e=e, provider_config=provider_config)
|
||||
elif isinstance(transformed_request, str) or isinstance(
|
||||
transformed_request, bytes
|
||||
):
|
||||
|
|
@ -3639,6 +3683,224 @@ class BaseLLMHTTPHandler:
|
|||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
# 8 MiB; a 256 KiB multiple, which GCS requires for every non-final chunk.
|
||||
_RESUMABLE_CHUNK_SIZE = 8 * 1024 * 1024
|
||||
|
||||
@staticmethod
|
||||
def _iter_resumable_chunks(
|
||||
byte_iter: Iterator[bytes], chunk_size: int
|
||||
) -> Iterator[bytes]:
|
||||
"""Regroup a byte stream into ``chunk_size`` pieces, yielding a final
|
||||
partial piece only when it is non-empty. Every full piece is exactly
|
||||
``chunk_size`` bytes (kept a 256 KiB multiple for GCS) and never more than
|
||||
one chunk is buffered. An exactly chunk-aligned stream yields only full
|
||||
chunks, so the upload finalizes on its last data chunk instead of making
|
||||
an extra empty request; a 0-byte stream yields nothing and the caller
|
||||
finalizes with a single empty request.
|
||||
"""
|
||||
buf = bytearray()
|
||||
for piece in byte_iter:
|
||||
buf.extend(piece)
|
||||
while len(buf) >= chunk_size:
|
||||
yield bytes(buf[:chunk_size])
|
||||
del buf[:chunk_size]
|
||||
if buf:
|
||||
yield bytes(buf)
|
||||
|
||||
@staticmethod
|
||||
def _resumable_content_range(offset: int, data_len: int, is_final: bool) -> str:
|
||||
if not is_final:
|
||||
return f"bytes {offset}-{offset + data_len - 1}/*"
|
||||
total = offset + data_len
|
||||
if data_len == 0:
|
||||
return f"bytes */{total}"
|
||||
return f"bytes {offset}-{total - 1}/{total}"
|
||||
|
||||
@staticmethod
|
||||
def _resumable_request_kwargs(
|
||||
headers: dict,
|
||||
content: bytes,
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
) -> dict:
|
||||
kwargs: Dict[str, Any] = {"headers": headers, "content": content}
|
||||
if timeout is not None:
|
||||
kwargs["timeout"] = timeout
|
||||
return kwargs
|
||||
|
||||
def _resumable_chunked_upload(
|
||||
self,
|
||||
*,
|
||||
client: HTTPHandler,
|
||||
initiate_url: str,
|
||||
base_headers: dict,
|
||||
config: dict,
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
) -> httpx.Response:
|
||||
"""Open a GCS resumable session, then PUT the body in bounded chunks so a
|
||||
large upload is never held in memory in full."""
|
||||
stream = config["body_stream"]
|
||||
chunk_size = config.get("chunk_size", self._RESUMABLE_CHUNK_SIZE)
|
||||
session_url_header = config.get("session_url_header", "location")
|
||||
httpx_client = client.client
|
||||
|
||||
init_headers = {**base_headers, **config.get("initiate_headers", {})}
|
||||
init_req = httpx_client.build_request(
|
||||
"POST",
|
||||
initiate_url,
|
||||
**self._resumable_request_kwargs(init_headers, b"", timeout),
|
||||
)
|
||||
init_resp = httpx_client.send(init_req, follow_redirects=False)
|
||||
init_resp.read()
|
||||
if init_resp.status_code not in (200, 201):
|
||||
init_resp.raise_for_status()
|
||||
session_url = init_resp.headers.get(session_url_header)
|
||||
if not session_url:
|
||||
raise ValueError(
|
||||
f"resumable upload: no session URL in '{session_url_header}' header"
|
||||
)
|
||||
|
||||
offset = 0
|
||||
pending: Optional[bytes] = None
|
||||
for chunk in self._iter_resumable_chunks(stream.iter_bytes(), chunk_size):
|
||||
if pending is not None:
|
||||
self._send_resumable_chunk(
|
||||
httpx_client,
|
||||
session_url,
|
||||
base_headers,
|
||||
pending,
|
||||
offset,
|
||||
is_final=False,
|
||||
timeout=timeout,
|
||||
)
|
||||
offset += len(pending)
|
||||
pending = chunk
|
||||
return self._send_resumable_chunk(
|
||||
httpx_client,
|
||||
session_url,
|
||||
base_headers,
|
||||
pending or b"",
|
||||
offset,
|
||||
is_final=True,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
def _send_resumable_chunk(
|
||||
self,
|
||||
httpx_client: httpx.Client,
|
||||
url: str,
|
||||
base_headers: dict,
|
||||
data: bytes,
|
||||
offset: int,
|
||||
*,
|
||||
is_final: bool,
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
) -> httpx.Response:
|
||||
headers = {
|
||||
**base_headers,
|
||||
"Content-Range": self._resumable_content_range(offset, len(data), is_final),
|
||||
}
|
||||
req = httpx_client.build_request(
|
||||
"PUT", url, **self._resumable_request_kwargs(headers, data, timeout)
|
||||
)
|
||||
resp = httpx_client.send(req, follow_redirects=False)
|
||||
resp.read()
|
||||
if resp.status_code not in ((200, 201) if is_final else (308,)):
|
||||
# 4xx/5xx raise here; the ValueError catches an unexpected success
|
||||
# status (e.g. a 200 where the protocol expects a 308 between chunks).
|
||||
resp.raise_for_status()
|
||||
raise ValueError(f"resumable upload: unexpected status {resp.status_code}")
|
||||
return resp
|
||||
|
||||
async def _aresumable_chunked_upload(
|
||||
self,
|
||||
*,
|
||||
client: AsyncHTTPHandler,
|
||||
initiate_url: str,
|
||||
base_headers: dict,
|
||||
config: dict,
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
) -> httpx.Response:
|
||||
stream = config["body_stream"]
|
||||
chunk_size = config.get("chunk_size", self._RESUMABLE_CHUNK_SIZE)
|
||||
session_url_header = config.get("session_url_header", "location")
|
||||
httpx_client = client.client
|
||||
|
||||
init_headers = {**base_headers, **config.get("initiate_headers", {})}
|
||||
init_req = httpx_client.build_request(
|
||||
"POST",
|
||||
initiate_url,
|
||||
**self._resumable_request_kwargs(init_headers, b"", timeout),
|
||||
)
|
||||
init_resp = await httpx_client.send(init_req, follow_redirects=False)
|
||||
await init_resp.aread()
|
||||
if init_resp.status_code not in (200, 201):
|
||||
init_resp.raise_for_status()
|
||||
session_url = init_resp.headers.get(session_url_header)
|
||||
if not session_url:
|
||||
raise ValueError(
|
||||
f"resumable upload: no session URL in '{session_url_header}' header"
|
||||
)
|
||||
|
||||
offset = 0
|
||||
pending: Optional[bytes] = None
|
||||
# Producing each chunk runs the synchronous per-row transform for that
|
||||
# chunk's worth of rows. Pull it off the event loop thread so a large
|
||||
# upload does not block other concurrent requests between PUTs.
|
||||
chunk_iter = self._iter_resumable_chunks(stream.iter_bytes(), chunk_size)
|
||||
done = object()
|
||||
while True:
|
||||
chunk = await asyncio.to_thread(next, chunk_iter, done)
|
||||
if chunk is done:
|
||||
break
|
||||
if pending is not None:
|
||||
await self._asend_resumable_chunk(
|
||||
httpx_client,
|
||||
session_url,
|
||||
base_headers,
|
||||
pending,
|
||||
offset,
|
||||
is_final=False,
|
||||
timeout=timeout,
|
||||
)
|
||||
offset += len(pending)
|
||||
pending = chunk
|
||||
return await self._asend_resumable_chunk(
|
||||
httpx_client,
|
||||
session_url,
|
||||
base_headers,
|
||||
pending or b"",
|
||||
offset,
|
||||
is_final=True,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
async def _asend_resumable_chunk(
|
||||
self,
|
||||
httpx_client: httpx.AsyncClient,
|
||||
url: str,
|
||||
base_headers: dict,
|
||||
data: bytes,
|
||||
offset: int,
|
||||
*,
|
||||
is_final: bool,
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
) -> httpx.Response:
|
||||
headers = {
|
||||
**base_headers,
|
||||
"Content-Range": self._resumable_content_range(offset, len(data), is_final),
|
||||
}
|
||||
req = httpx_client.build_request(
|
||||
"PUT", url, **self._resumable_request_kwargs(headers, data, timeout)
|
||||
)
|
||||
resp = await httpx_client.send(req, follow_redirects=False)
|
||||
await resp.aread()
|
||||
if resp.status_code not in ((200, 201) if is_final else (308,)):
|
||||
# 4xx/5xx raise here; the ValueError catches an unexpected success
|
||||
# status (e.g. a 200 where the protocol expects a 308 between chunks).
|
||||
resp.raise_for_status()
|
||||
raise ValueError(f"resumable upload: unexpected status {resp.status_code}")
|
||||
return resp
|
||||
|
||||
def create_batch(
|
||||
self,
|
||||
create_batch_data: "CreateBatchRequest",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -1,9 +1,21 @@
|
|||
import base64
|
||||
import io
|
||||
import itertools
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
|
||||
from typing import (
|
||||
Any,
|
||||
Callable,
|
||||
Dict,
|
||||
Iterable,
|
||||
Iterator,
|
||||
List,
|
||||
Optional,
|
||||
Tuple,
|
||||
Union,
|
||||
)
|
||||
|
||||
import httpx
|
||||
from httpx import Headers, Response
|
||||
|
|
@ -22,9 +34,13 @@ from litellm.litellm_core_utils.cloud_storage_security import (
|
|||
validate_managed_cloud_file_id,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
extract_file_data,
|
||||
extract_file_metadata,
|
||||
)
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.base_llm.files.transformation import (
|
||||
BaseFileUploadStream,
|
||||
BaseFilesConfig,
|
||||
LiteLLMLoggingObj,
|
||||
)
|
||||
|
|
@ -44,8 +60,9 @@ from litellm.types.llms.openai import (
|
|||
OpenAIFileObject,
|
||||
PathLike,
|
||||
)
|
||||
from litellm.types.files import ResumableChunkedUploadConfig
|
||||
from litellm.types.llms.vertex_ai import GcsBucketResponse
|
||||
from litellm.types.utils import ExtractedFileData, LlmProviders, ModelResponse
|
||||
from litellm.types.utils import LlmProviders, ModelResponse
|
||||
|
||||
from ..common_utils import VertexAIError
|
||||
from ..vertex_llm_base import VertexBase
|
||||
|
|
@ -137,42 +154,140 @@ def _get_litellm_batch_custom_id_from_labels(labels: Dict[str, Any]) -> str:
|
|||
return str(labels.get("litellm_custom_id", "unknown"))
|
||||
|
||||
|
||||
def _openai_batch_jsonl_entries_to_vertex_wrapped_requests(
|
||||
openai_jsonl_content: List[Dict[str, Any]],
|
||||
def _openai_batch_jsonl_entry_to_vertex_wrapped_request(
|
||||
openai_entry: Dict[str, Any],
|
||||
map_openai_to_vertex_params: Callable[[Dict[str, Any]], Dict[str, Any]],
|
||||
) -> List[Dict[str, Any]]:
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Transforms OpenAI JSONL batch entries to Vertex AI JSONL lines.
|
||||
Transforms a single OpenAI JSONL batch entry into its Vertex wrapped request.
|
||||
|
||||
jsonl body for vertex is {"request": <request_body>}
|
||||
Example Vertex jsonl
|
||||
{"request":{"contents": [{"role": "user", "parts": [{"text": "What is the relation between the following video and image samples?"}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/video/animals.mp4", "mimeType": "video/mp4"}}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/image/cricket.jpeg", "mimeType": "image/jpeg"}}]}]}}
|
||||
{"request":{"contents": [{"role": "user", "parts": [{"text": "Describe what is happening in this video."}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/video/another_video.mov", "mimeType": "video/mov"}}]}]}}
|
||||
"""
|
||||
openai_request_body = openai_entry.get("body") or {}
|
||||
vertex_request_body = _transform_request_body(
|
||||
messages=openai_request_body.get("messages", []),
|
||||
model=openai_request_body.get("model", ""),
|
||||
optional_params=map_openai_to_vertex_params(openai_request_body),
|
||||
custom_llm_provider="vertex_ai",
|
||||
litellm_params={},
|
||||
cached_content=None,
|
||||
)
|
||||
|
||||
custom_id = openai_entry.get("custom_id")
|
||||
if custom_id is not None:
|
||||
if "labels" not in vertex_request_body:
|
||||
vertex_request_body["labels"] = {}
|
||||
_set_litellm_batch_custom_id_labels(vertex_request_body["labels"], custom_id)
|
||||
|
||||
return {"request": vertex_request_body}
|
||||
|
||||
|
||||
def _iter_stripped_lines(raw_lines: Iterable[Union[str, bytes]]) -> Iterator[str]:
|
||||
"""Decode (when needed), strip, and drop blank lines from an iterable of lines."""
|
||||
for raw in raw_lines:
|
||||
line = raw.decode("utf-8") if isinstance(raw, (bytes, bytearray)) else raw
|
||||
line = line.strip()
|
||||
if line:
|
||||
yield line
|
||||
|
||||
|
||||
def _iter_openai_jsonl_lines(openai_file_content: FileTypes) -> Iterator[str]:
|
||||
"""
|
||||
Yield non-empty JSONL lines one at a time without materializing the whole
|
||||
payload, so peak memory stays bounded regardless of payload size. Mirrors
|
||||
``str.splitlines()`` + ``line.strip()`` for ``\\n`` / ``\\r\\n`` delimited
|
||||
JSONL.
|
||||
"""
|
||||
content: Any = openai_file_content
|
||||
if isinstance(content, tuple):
|
||||
content = content[1]
|
||||
|
||||
if isinstance(content, (bytes, bytearray)):
|
||||
# Scan for newlines in place so a large in-memory payload is not copied
|
||||
# into a BytesIO just to iterate it line by line.
|
||||
newline = ord("\n")
|
||||
start, length = 0, len(content)
|
||||
while start < length:
|
||||
idx = content.find(newline, start)
|
||||
if idx == -1:
|
||||
chunk, start = content[start:], length
|
||||
else:
|
||||
chunk, start = content[start:idx], idx + 1
|
||||
line = chunk.decode("utf-8").strip()
|
||||
if line:
|
||||
yield line
|
||||
return
|
||||
|
||||
if isinstance(content, str):
|
||||
yield from _iter_stripped_lines(io.StringIO(content))
|
||||
return
|
||||
|
||||
if isinstance(content, PathLike):
|
||||
with open(str(content), "rb") as handle:
|
||||
yield from _iter_stripped_lines(handle)
|
||||
return
|
||||
|
||||
if hasattr(content, "read"):
|
||||
# The handle is read twice per upload (first-row probe for the GCS
|
||||
# object name, then the body stream), so it must rewind to 0. A
|
||||
# non-seekable handle would silently resume mid-stream and drop the
|
||||
# already-consumed first row, so reject it loudly instead.
|
||||
seek = getattr(content, "seek", None)
|
||||
if seek is None:
|
||||
raise ValueError(
|
||||
"Batch upload file handle must be seekable; got a non-seekable "
|
||||
"stream. Pass bytes, a path, or a seekable handle."
|
||||
)
|
||||
try:
|
||||
seek(0)
|
||||
except (OSError, ValueError) as e:
|
||||
raise ValueError(
|
||||
"Batch upload file handle must be seekable so it can be re-read "
|
||||
"for the GCS object name and the upload body."
|
||||
) from e
|
||||
yield from _iter_stripped_lines(content)
|
||||
return
|
||||
|
||||
raise ValueError("Unsupported file content type")
|
||||
|
||||
|
||||
def _iter_openai_jsonl_entries(
|
||||
openai_file_content: FileTypes,
|
||||
) -> Iterator[Dict[str, Any]]:
|
||||
for line in _iter_openai_jsonl_lines(openai_file_content):
|
||||
yield json.loads(line)
|
||||
|
||||
|
||||
class _OpenAIToVertexBatchUploadStream(BaseFileUploadStream):
|
||||
"""Streams an OpenAI batch JSONL upload as Vertex-wrapped JSONL one row at a
|
||||
time, so the transformed payload is never held in full.
|
||||
|
||||
The transform runs lazily as the HTTP client pulls each chunk, which keeps
|
||||
peak memory at one row regardless of how large the batch file is.
|
||||
"""
|
||||
|
||||
vertex_jsonl_content = []
|
||||
for _openai_jsonl_content in openai_jsonl_content:
|
||||
openai_request_body = _openai_jsonl_content.get("body") or {}
|
||||
vertex_request_body = _transform_request_body(
|
||||
messages=openai_request_body.get("messages", []),
|
||||
model=openai_request_body.get("model", ""),
|
||||
optional_params=map_openai_to_vertex_params(openai_request_body),
|
||||
custom_llm_provider="vertex_ai",
|
||||
litellm_params={},
|
||||
cached_content=None,
|
||||
)
|
||||
def __init__(
|
||||
self,
|
||||
openai_file_content: FileTypes,
|
||||
map_openai_to_vertex_params: Callable[[Dict[str, Any]], Dict[str, Any]],
|
||||
) -> None:
|
||||
self._openai_file_content = openai_file_content
|
||||
self._map_openai_to_vertex_params = map_openai_to_vertex_params
|
||||
|
||||
# Add custom_id as a label for correlation in batch outputs
|
||||
custom_id = _openai_jsonl_content.get("custom_id")
|
||||
if custom_id is not None:
|
||||
if "labels" not in vertex_request_body:
|
||||
vertex_request_body["labels"] = {}
|
||||
_set_litellm_batch_custom_id_labels(
|
||||
vertex_request_body["labels"], custom_id
|
||||
def _iter_vertex_jsonl_chunks(self) -> Iterator[bytes]:
|
||||
first = True
|
||||
for entry in _iter_openai_jsonl_entries(self._openai_file_content):
|
||||
wrapped = _openai_batch_jsonl_entry_to_vertex_wrapped_request(
|
||||
entry, self._map_openai_to_vertex_params
|
||||
)
|
||||
prefix = b"" if first else b"\n"
|
||||
first = False
|
||||
yield prefix + json.dumps(wrapped).encode("utf-8")
|
||||
|
||||
vertex_jsonl_content.append({"request": vertex_request_body})
|
||||
return vertex_jsonl_content
|
||||
def iter_bytes(self) -> Iterator[bytes]:
|
||||
return self._iter_vertex_jsonl_chunks()
|
||||
|
||||
|
||||
class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
||||
|
|
@ -181,7 +296,6 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.jsonl_transformation = VertexAIJsonlFilesTransformation()
|
||||
super().__init__()
|
||||
|
||||
@property
|
||||
|
|
@ -208,43 +322,6 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
return headers
|
||||
|
||||
def _get_content_from_openai_file(self, openai_file_content: FileTypes) -> str:
|
||||
"""
|
||||
Helper to extract content from various OpenAI file types and return as string.
|
||||
|
||||
Handles:
|
||||
- Direct content (str, bytes, IO[bytes])
|
||||
- Tuple formats: (filename, content, [content_type], [headers])
|
||||
- PathLike objects
|
||||
"""
|
||||
content: Union[str, bytes] = b""
|
||||
# Extract file content from tuple if necessary
|
||||
if isinstance(openai_file_content, tuple):
|
||||
# Take the second element which is always the file content
|
||||
file_content = openai_file_content[1]
|
||||
else:
|
||||
file_content = openai_file_content
|
||||
|
||||
# Handle different file content types
|
||||
if isinstance(file_content, str):
|
||||
# String content can be used directly
|
||||
content = file_content
|
||||
elif isinstance(file_content, bytes):
|
||||
# Bytes content can be decoded
|
||||
content = file_content
|
||||
elif isinstance(file_content, PathLike): # PathLike
|
||||
with open(str(file_content), "rb") as f:
|
||||
content = f.read()
|
||||
elif hasattr(file_content, "read"): # IO[bytes]
|
||||
# File-like objects need to be read
|
||||
content = file_content.read()
|
||||
|
||||
# Ensure content is string
|
||||
if isinstance(content, bytes):
|
||||
content = content.decode("utf-8")
|
||||
|
||||
return content
|
||||
|
||||
def _get_gcs_object_name_from_batch_jsonl(
|
||||
self,
|
||||
openai_jsonl_content: List[Dict[str, Any]],
|
||||
|
|
@ -261,32 +338,21 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
object_name = f"{VERTEX_AI_MANAGED_GCS_PREFIX}{safe_model_path}/{uuid.uuid4()}"
|
||||
return object_name
|
||||
|
||||
def get_object_name(
|
||||
self, extracted_file_data: ExtractedFileData, purpose: str
|
||||
) -> str:
|
||||
def get_object_name(self, file_data: FileTypes, purpose: str) -> str:
|
||||
"""
|
||||
Get the object name for the request
|
||||
Get the object name for the request.
|
||||
|
||||
Reads only the first JSONL entry (streamed) for batch files, so a large
|
||||
upload is never materialized just to derive the GCS object name.
|
||||
"""
|
||||
extracted_file_data_content = extracted_file_data.get("content")
|
||||
|
||||
if extracted_file_data_content is None:
|
||||
raise ValueError("file content is required")
|
||||
|
||||
if purpose == "batch":
|
||||
## 1. If jsonl, check if there's a model name
|
||||
file_content = self._get_content_from_openai_file(
|
||||
extracted_file_data_content
|
||||
)
|
||||
|
||||
# Split into lines and parse each line as JSON
|
||||
openai_jsonl_content = [
|
||||
json.loads(line) for line in file_content.splitlines() if line.strip()
|
||||
]
|
||||
if len(openai_jsonl_content) > 0:
|
||||
return self._get_gcs_object_name_from_batch_jsonl(openai_jsonl_content)
|
||||
## 1. If jsonl, derive the object name from the first entry's model
|
||||
first_entry = next(_iter_openai_jsonl_entries(file_data), None)
|
||||
if first_entry is not None:
|
||||
return self._get_gcs_object_name_from_batch_jsonl([first_entry])
|
||||
|
||||
## 2. If not jsonl, store under a server-generated managed object name
|
||||
filename = extracted_file_data.get("filename")
|
||||
filename, _ = extract_file_metadata(file_data)
|
||||
return build_managed_cloud_object_name(
|
||||
prefix=f"{VERTEX_AI_MANAGED_GCS_PREFIX}uploads/",
|
||||
filename=filename,
|
||||
|
|
@ -294,7 +360,11 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
)
|
||||
|
||||
def _get_configured_bucket_name(self, litellm_params: Dict) -> str:
|
||||
bucket_name = litellm_params.get("bucket_name") or os.getenv("GCS_BUCKET_NAME")
|
||||
bucket_name = (
|
||||
litellm_params.get("gcs_bucket_name")
|
||||
or litellm_params.get("bucket_name")
|
||||
or os.getenv("GCS_BUCKET_NAME")
|
||||
)
|
||||
if not bucket_name:
|
||||
raise ValueError("GCS bucket_name is required")
|
||||
return bucket_name
|
||||
|
|
@ -319,12 +389,21 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
raise ValueError("file is required")
|
||||
if purpose is None:
|
||||
raise ValueError("purpose is required")
|
||||
extracted_file_data = extract_file_data(file_data)
|
||||
object_name = self.get_object_name(extracted_file_data, purpose)
|
||||
_, content_type = extract_file_metadata(file_data)
|
||||
object_name = self.get_object_name(file_data, purpose)
|
||||
if object_prefix:
|
||||
object_name = f"{object_prefix}/{object_name}"
|
||||
encoded_object_name = encode_gcs_object_name_for_url(object_name)
|
||||
endpoint = f"upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={encoded_object_name}"
|
||||
# Batch jsonl is streamed via a resumable session (bounded memory on
|
||||
# large uploads); everything else is a single simple-media upload.
|
||||
upload_type = (
|
||||
"resumable"
|
||||
if FilesAPIUtils.is_batch_jsonl_request(
|
||||
create_file_data=data, content_type=content_type
|
||||
)
|
||||
else "media"
|
||||
)
|
||||
endpoint = f"upload/storage/v1/b/{bucket_name}/o?uploadType={upload_type}&name={encoded_object_name}"
|
||||
api_base = api_base or "https://storage.googleapis.com"
|
||||
if not api_base:
|
||||
raise ValueError("api_base is required")
|
||||
|
|
@ -366,14 +445,6 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
)
|
||||
return vertex_params
|
||||
|
||||
def _transform_openai_jsonl_content_to_vertex_ai_jsonl_content(
|
||||
self, openai_jsonl_content: List[Dict[str, Any]]
|
||||
) -> List[Dict[str, Any]]:
|
||||
return _openai_batch_jsonl_entries_to_vertex_wrapped_requests(
|
||||
openai_jsonl_content=openai_jsonl_content,
|
||||
map_openai_to_vertex_params=self._map_openai_to_vertex_params,
|
||||
)
|
||||
|
||||
def transform_create_file_request(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -384,40 +455,34 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
"""
|
||||
2 Cases:
|
||||
1. Handle basic file upload
|
||||
2. Handle batch file upload (.jsonl)
|
||||
2. Handle batch file upload (.jsonl), streamed to a GCS resumable
|
||||
session so large uploads stay memory-bounded.
|
||||
"""
|
||||
file_data = create_file_data.get("file")
|
||||
if file_data is None:
|
||||
raise ValueError("file is required")
|
||||
extracted_file_data = extract_file_data(file_data)
|
||||
extracted_file_data_content = extracted_file_data.get("content")
|
||||
|
||||
if extracted_file_data_content is None:
|
||||
raise ValueError("file content is required")
|
||||
|
||||
if FilesAPIUtils.is_batch_jsonl_file(
|
||||
_, content_type = extract_file_metadata(file_data)
|
||||
if FilesAPIUtils.is_batch_jsonl_request(
|
||||
create_file_data=create_file_data,
|
||||
extracted_file_data=extracted_file_data,
|
||||
content_type=content_type,
|
||||
):
|
||||
## 1. If jsonl, check if there's a model name
|
||||
file_content = self._get_content_from_openai_file(
|
||||
extracted_file_data_content
|
||||
)
|
||||
|
||||
# Split into lines and parse each line as JSON
|
||||
openai_jsonl_content = [
|
||||
json.loads(line) for line in file_content.splitlines() if line.strip()
|
||||
]
|
||||
vertex_jsonl_content = (
|
||||
self._transform_openai_jsonl_content_to_vertex_ai_jsonl_content(
|
||||
openai_jsonl_content
|
||||
return {
|
||||
"resumable_chunked_upload": ResumableChunkedUploadConfig(
|
||||
body_stream=_OpenAIToVertexBatchUploadStream(
|
||||
file_data,
|
||||
self._map_openai_to_vertex_params,
|
||||
),
|
||||
initiate_headers={
|
||||
"X-Upload-Content-Type": "application/json",
|
||||
},
|
||||
)
|
||||
)
|
||||
return "\n".join(json.dumps(item) for item in vertex_jsonl_content)
|
||||
elif isinstance(extracted_file_data_content, bytes):
|
||||
}
|
||||
|
||||
extracted_file_data_content = extract_file_data(file_data).get("content")
|
||||
if isinstance(extracted_file_data_content, bytes):
|
||||
return extracted_file_data_content
|
||||
else:
|
||||
raise ValueError("Unsupported file content type")
|
||||
raise ValueError("Unsupported file content type")
|
||||
|
||||
def transform_create_file_response(
|
||||
self,
|
||||
|
|
@ -642,39 +707,38 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
}
|
||||
"""
|
||||
try:
|
||||
# Decode content
|
||||
content_str = content.decode("utf-8")
|
||||
|
||||
# Check if it's JSONL (multiple lines)
|
||||
lines = content_str.strip().split("\n")
|
||||
if not lines:
|
||||
# Read the result file one row at a time. Batch output files can be
|
||||
# as large as the (multi-GB) input, so splitting into a list of rows
|
||||
# and building a second list of transformed rows peaks at several full
|
||||
# copies and OOMs on retrieval.
|
||||
lines = _iter_openai_jsonl_lines(content)
|
||||
try:
|
||||
first_line = next(lines)
|
||||
except StopIteration:
|
||||
return content
|
||||
|
||||
# Try to parse the first line to see if it's Vertex AI batch output
|
||||
first_line = json.loads(lines[0])
|
||||
|
||||
# Check if it has Vertex AI batch output structure with discriminating fields
|
||||
# Must have request, response, and processed_time
|
||||
# Plus either candidates (success) or status (error)
|
||||
has_base_structure = (
|
||||
"response" in first_line
|
||||
and "request" in first_line
|
||||
and "processed_time" in first_line
|
||||
# Identify a Vertex AI batch output from the first row's
|
||||
# discriminating fields. Anything else (e.g. a binary file whose
|
||||
# first line is not valid UTF-8/JSON) raises and falls through to the
|
||||
# passthrough below, leaving the content untouched.
|
||||
first_row = json.loads(first_line)
|
||||
is_vertex_batch_output = (
|
||||
"request" in first_row
|
||||
and "response" in first_row
|
||||
and "processed_time" in first_row
|
||||
and (
|
||||
"candidates" in first_row.get("response", {})
|
||||
or "promptFeedback" in first_row.get("response", {})
|
||||
or bool(first_row.get("status"))
|
||||
)
|
||||
)
|
||||
has_success_or_error = (
|
||||
"candidates" in first_line.get("response", {})
|
||||
or "promptFeedback" in first_line.get("response", {})
|
||||
or bool(first_line.get("status"))
|
||||
)
|
||||
|
||||
if not (has_base_structure and has_success_or_error):
|
||||
# Not a Vertex AI batch output, return as-is
|
||||
if not is_vertex_batch_output:
|
||||
return content
|
||||
|
||||
vertex_gemini_config = VertexGeminiConfig()
|
||||
# Always use a fresh local Logging object for the per-line transformation
|
||||
# so we never mutate the caller's logging_obj (which already went through
|
||||
# pre_call and has its own model/start_time/optional_params set).
|
||||
# Use a fresh Logging object for the per-row transform so we never
|
||||
# mutate the caller's (which already ran pre_call with its own
|
||||
# model/start_time/optional_params).
|
||||
batch_transform_logging_obj = Logging(
|
||||
model="",
|
||||
messages=[],
|
||||
|
|
@ -691,29 +755,27 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
request=httpx.Request(method="POST", url="https://example.com"),
|
||||
)
|
||||
|
||||
# Transform all lines
|
||||
transformed_lines = []
|
||||
for line in lines:
|
||||
if not line.strip():
|
||||
continue
|
||||
|
||||
# Transform each row straight into the output buffer, so peak memory
|
||||
# stays at ~one row plus the output. If any row fails, return the
|
||||
# original content unchanged.
|
||||
output = bytearray()
|
||||
for line in itertools.chain([first_line], lines):
|
||||
try:
|
||||
vertex_output = json.loads(line)
|
||||
openai_output = (
|
||||
self._transform_single_vertex_batch_output_to_openai(
|
||||
vertex_output=vertex_output,
|
||||
vertex_output=json.loads(line),
|
||||
vertex_gemini_config=vertex_gemini_config,
|
||||
logging_obj=batch_transform_logging_obj,
|
||||
mock_httpx_response=mock_httpx_response,
|
||||
)
|
||||
)
|
||||
transformed_lines.append(json.dumps(openai_output))
|
||||
except Exception:
|
||||
# If any line fails, return original content
|
||||
return content
|
||||
if output:
|
||||
output += b"\n"
|
||||
output += json.dumps(openai_output).encode("utf-8")
|
||||
|
||||
# Return transformed content
|
||||
return "\n".join(transformed_lines).encode("utf-8")
|
||||
return bytes(output)
|
||||
|
||||
except Exception:
|
||||
# If anything fails, return original content
|
||||
|
|
@ -795,137 +857,3 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
"message": f"Failed to transform response: {str(e)}",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
class VertexAIJsonlFilesTransformation(VertexGeminiConfig):
|
||||
"""
|
||||
Transforms OpenAI /v1/files/* requests to VertexAI /v1/files/* requests
|
||||
"""
|
||||
|
||||
def transform_openai_file_content_to_vertex_ai_file_content(
|
||||
self, openai_file_content: Optional[FileTypes] = None
|
||||
) -> Tuple[str, str]:
|
||||
"""
|
||||
Transforms OpenAI FileContentRequest to VertexAI FileContentRequest
|
||||
"""
|
||||
|
||||
if openai_file_content is None:
|
||||
raise ValueError("contents of file are None")
|
||||
# Read the content of the file
|
||||
file_content = self._get_content_from_openai_file(openai_file_content)
|
||||
|
||||
# Split into lines and parse each line as JSON
|
||||
openai_jsonl_content = [
|
||||
json.loads(line) for line in file_content.splitlines() if line.strip()
|
||||
]
|
||||
vertex_jsonl_content = (
|
||||
self._transform_openai_jsonl_content_to_vertex_ai_jsonl_content(
|
||||
openai_jsonl_content
|
||||
)
|
||||
)
|
||||
vertex_jsonl_string = "\n".join(
|
||||
json.dumps(item) for item in vertex_jsonl_content
|
||||
)
|
||||
object_name = self._get_gcs_object_name(
|
||||
openai_jsonl_content=openai_jsonl_content
|
||||
)
|
||||
return vertex_jsonl_string, object_name
|
||||
|
||||
def _transform_openai_jsonl_content_to_vertex_ai_jsonl_content(
|
||||
self, openai_jsonl_content: List[Dict[str, Any]]
|
||||
) -> List[Dict[str, Any]]:
|
||||
return _openai_batch_jsonl_entries_to_vertex_wrapped_requests(
|
||||
openai_jsonl_content=openai_jsonl_content,
|
||||
map_openai_to_vertex_params=self._map_openai_to_vertex_params,
|
||||
)
|
||||
|
||||
def _get_gcs_object_name(
|
||||
self,
|
||||
openai_jsonl_content: List[Dict[str, Any]],
|
||||
) -> str:
|
||||
"""
|
||||
Gets a unique GCS object name for the VertexAI batch prediction job
|
||||
|
||||
named as: litellm-vertex-{model}-{uuid}
|
||||
"""
|
||||
_model = openai_jsonl_content[0].get("body", {}).get("model", "")
|
||||
if "publishers/google/models" not in _model:
|
||||
_model = f"publishers/google/models/{_model}"
|
||||
safe_model_path = sanitize_cloud_object_path(_model, fallback="model")
|
||||
object_name = f"{VERTEX_AI_MANAGED_GCS_PREFIX}{safe_model_path}/{uuid.uuid4()}"
|
||||
return object_name
|
||||
|
||||
def _map_openai_to_vertex_params(
|
||||
self,
|
||||
openai_request_body: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
wrapper to call VertexGeminiConfig.map_openai_params
|
||||
"""
|
||||
_model = openai_request_body.get("model", "")
|
||||
vertex_params = self.map_openai_params(
|
||||
model=_model,
|
||||
non_default_params=openai_request_body,
|
||||
optional_params={},
|
||||
drop_params=False,
|
||||
)
|
||||
return vertex_params
|
||||
|
||||
def _get_content_from_openai_file(self, openai_file_content: FileTypes) -> str:
|
||||
"""
|
||||
Helper to extract content from various OpenAI file types and return as string.
|
||||
|
||||
Handles:
|
||||
- Direct content (str, bytes, IO[bytes])
|
||||
- Tuple formats: (filename, content, [content_type], [headers])
|
||||
- PathLike objects
|
||||
"""
|
||||
content: Union[str, bytes] = b""
|
||||
# Extract file content from tuple if necessary
|
||||
if isinstance(openai_file_content, tuple):
|
||||
# Take the second element which is always the file content
|
||||
file_content = openai_file_content[1]
|
||||
else:
|
||||
file_content = openai_file_content
|
||||
|
||||
# Handle different file content types
|
||||
if isinstance(file_content, str):
|
||||
# String content can be used directly
|
||||
content = file_content
|
||||
elif isinstance(file_content, bytes):
|
||||
# Bytes content can be decoded
|
||||
content = file_content
|
||||
elif isinstance(file_content, PathLike): # PathLike
|
||||
with open(str(file_content), "rb") as f:
|
||||
content = f.read()
|
||||
elif hasattr(file_content, "read"): # IO[bytes]
|
||||
# File-like objects need to be read
|
||||
content = file_content.read()
|
||||
|
||||
# Ensure content is string
|
||||
if isinstance(content, bytes):
|
||||
content = content.decode("utf-8")
|
||||
|
||||
return content
|
||||
|
||||
def transform_gcs_bucket_response_to_openai_file_object(
|
||||
self, create_file_data: CreateFileRequest, gcs_upload_response: Dict[str, Any]
|
||||
) -> OpenAIFileObject:
|
||||
"""
|
||||
Transforms GCS Bucket upload file response to OpenAI FileObject
|
||||
"""
|
||||
gcs_id = gcs_upload_response.get("id", "")
|
||||
# Remove the last numeric ID from the path
|
||||
gcs_id = "/".join(gcs_id.split("/")[:-1]) if gcs_id else ""
|
||||
|
||||
return OpenAIFileObject(
|
||||
purpose=create_file_data.get("purpose", "batch"),
|
||||
id=f"gs://{gcs_id}",
|
||||
filename=gcs_upload_response.get("name", ""),
|
||||
created_at=_convert_vertex_datetime_to_openai_datetime(
|
||||
vertex_datetime=gcs_upload_response.get("timeCreated", "")
|
||||
),
|
||||
status="uploaded",
|
||||
bytes=gcs_upload_response.get("size", 0),
|
||||
object="file",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
0
litellm/proxy/logging_endpoints/__init__.py
Normal file
0
litellm/proxy/logging_endpoints/__init__.py
Normal file
208
litellm/proxy/logging_endpoints/callback_logs_endpoints.py
Normal file
208
litellm/proxy/logging_endpoints/callback_logs_endpoints.py
Normal file
|
|
@ -0,0 +1,208 @@
|
|||
"""
|
||||
Ingest pre-built logging payloads from external producers and replay them
|
||||
through LiteLLM's standard success/failure callback fan-out.
|
||||
|
||||
This exists for hosts that own a request outside the Python process — e.g. the
|
||||
`litellm-rust` gateway proxying realtime websockets. Those hosts can't use the
|
||||
in-process logging object, so they POST a finished `StandardLoggingPayload` here
|
||||
and Python replays it through the exact same path a normal completion uses:
|
||||
`Logging.async_success_handler` / `async_failure_handler`. Every registered
|
||||
callback (spend logs, Langfuse, Datadog, ...) fires unchanged — there is no
|
||||
spend-logs-specific or callback-specific code here, only the replay.
|
||||
|
||||
The endpoint is generic: realtime is the first producer, but the contract is the
|
||||
self-describing `StandardLoggingPayload`, so completions/responses can use it too.
|
||||
"""
|
||||
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.types.proxy.callback_logs_endpoints import (
|
||||
CallbackLogFailure,
|
||||
CallbackLogRecord,
|
||||
CallbackLogsRequest,
|
||||
CallbackLogsResponse,
|
||||
)
|
||||
|
||||
# Routes the Python proxy exposes for the Rust data-plane gateway to call into
|
||||
# (logging today; auth/budgets later). Namespaced under /v1/rust_control_plane so
|
||||
# they're clearly distinct from the proxy's own control-plane/management routes.
|
||||
rust_control_plane_router = APIRouter(
|
||||
prefix="/v1/rust_control_plane", tags=["rust control plane"]
|
||||
)
|
||||
|
||||
|
||||
class CallbackLogsReplayer:
|
||||
"""
|
||||
Replays finished logging payloads through LiteLLM's callback fan-out.
|
||||
|
||||
Each helper is small and pure so the replay path is easy to read and test:
|
||||
rebuild a `Logging` object from the payload, seed `model_call_details` with
|
||||
exactly what the callbacks read, then dispatch to the success/failure
|
||||
handler. No spend/callback logic lives here — only the replay.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _epoch_to_datetime(value: Any) -> datetime:
|
||||
"""`StandardLoggingPayload` stores startTime/endTime as float epoch seconds."""
|
||||
if isinstance(value, (int, float)):
|
||||
return datetime.fromtimestamp(float(value), tz=timezone.utc)
|
||||
if isinstance(value, datetime):
|
||||
return value
|
||||
return datetime.now(tz=timezone.utc)
|
||||
|
||||
@staticmethod
|
||||
def _build_logging_obj(payload: dict[str, Any]) -> LiteLLMLogging:
|
||||
"""
|
||||
Reconstruct a `Logging` object from a finished payload and seed
|
||||
`model_call_details` with exactly what the success/failure callbacks
|
||||
read: the prebuilt `standard_logging_object`, the resolved
|
||||
`response_cost`, and the `litellm_params.metadata` keys used for cost
|
||||
attribution. Setting `standard_logging_object` up front makes the handler
|
||||
skip rebuilding it.
|
||||
"""
|
||||
model = payload.get("model") or ""
|
||||
call_type = payload.get("call_type") or "acompletion"
|
||||
start_time = CallbackLogsReplayer._epoch_to_datetime(payload.get("startTime"))
|
||||
call_id = (
|
||||
payload.get("litellm_call_id") or payload.get("id") or str(uuid.uuid4())
|
||||
)
|
||||
|
||||
logging_obj = LiteLLMLogging(
|
||||
model=model,
|
||||
messages=payload.get("messages") or [],
|
||||
# A replayed payload is always a *terminal*, fully-aggregated event —
|
||||
# the producer (e.g. the rust gateway) already collected the whole
|
||||
# session before POSTing. Never mark it streaming: a streaming
|
||||
# Logging object makes async_success_handler wait for a
|
||||
# complete_streaming_response that will never arrive, so the spend
|
||||
# log is never written.
|
||||
stream=False,
|
||||
call_type=call_type,
|
||||
start_time=start_time,
|
||||
litellm_call_id=call_id,
|
||||
function_id="",
|
||||
)
|
||||
|
||||
metadata: dict[str, Any] = payload.get("metadata") or {}
|
||||
litellm_metadata: dict[str, Any] = {
|
||||
"user_api_key": metadata.get("user_api_key_hash"),
|
||||
"user_api_key_alias": metadata.get("user_api_key_alias"),
|
||||
"user_api_key_user_id": metadata.get("user_api_key_user_id"),
|
||||
"user_api_key_team_id": metadata.get("user_api_key_team_id"),
|
||||
"user_api_key_org_id": metadata.get("user_api_key_org_id"),
|
||||
"user_api_key_end_user_id": metadata.get("user_api_key_end_user_id"),
|
||||
"spend_logs_metadata": metadata.get("spend_logs_metadata"),
|
||||
}
|
||||
|
||||
logging_obj.model_call_details.update(
|
||||
{
|
||||
"model": model,
|
||||
"call_type": call_type,
|
||||
"custom_llm_provider": payload.get("custom_llm_provider"),
|
||||
"response_cost": payload.get("response_cost") or 0.0,
|
||||
"standard_logging_object": payload,
|
||||
"litellm_params": {"metadata": litellm_metadata},
|
||||
"cache_hit": payload.get("cache_hit") or False,
|
||||
}
|
||||
)
|
||||
return logging_obj
|
||||
|
||||
@staticmethod
|
||||
def _response_obj_from_payload(payload: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Minimal response object so usage-derived spend-log fields resolve."""
|
||||
return {
|
||||
"id": payload.get("id"),
|
||||
"usage": {
|
||||
"prompt_tokens": payload.get("prompt_tokens", 0),
|
||||
"completion_tokens": payload.get("completion_tokens", 0),
|
||||
"total_tokens": payload.get("total_tokens", 0),
|
||||
},
|
||||
}
|
||||
|
||||
async def replay(self, record: CallbackLogRecord) -> None:
|
||||
"""Replay one record through the matching success/failure handler."""
|
||||
payload = record.standard_logging_payload
|
||||
verbose_proxy_logger.debug(
|
||||
"CallbackLogsReplayer: replaying %s record id=%s model=%s call_type=%s",
|
||||
record.status,
|
||||
payload.get("id"),
|
||||
payload.get("model"),
|
||||
payload.get("call_type"),
|
||||
)
|
||||
|
||||
logging_obj = self._build_logging_obj(payload)
|
||||
start_time = self._epoch_to_datetime(payload.get("startTime"))
|
||||
end_time = self._epoch_to_datetime(payload.get("endTime"))
|
||||
|
||||
if record.status == "success":
|
||||
await logging_obj.async_success_handler(
|
||||
result=self._response_obj_from_payload(payload),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
else:
|
||||
error_str = record.error or payload.get("error_str") or "replayed failure"
|
||||
await logging_obj.async_failure_handler(
|
||||
Exception(error_str),
|
||||
traceback_exception="",
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
||||
async def replay_batch(
|
||||
self, records: list[CallbackLogRecord]
|
||||
) -> CallbackLogsResponse:
|
||||
"""Replay a batch; a single bad record never sinks the rest. Each failure
|
||||
is reported back with its batch index so the caller can retry/triage it."""
|
||||
processed = 0
|
||||
failures: list[CallbackLogFailure] = []
|
||||
for index, record in enumerate(records):
|
||||
try:
|
||||
await self.replay(record)
|
||||
processed += 1
|
||||
except Exception as e:
|
||||
failures.append(CallbackLogFailure(index=index, error=str(e)))
|
||||
verbose_proxy_logger.exception(
|
||||
"CallbackLogsReplayer: failed to replay record %s: %s",
|
||||
index,
|
||||
str(e),
|
||||
)
|
||||
verbose_proxy_logger.debug(
|
||||
"CallbackLogsReplayer: batch done processed=%s failed=%s",
|
||||
processed,
|
||||
len(failures),
|
||||
)
|
||||
return CallbackLogsResponse(
|
||||
processed=processed, failed=len(failures), failures=failures
|
||||
)
|
||||
|
||||
|
||||
@rust_control_plane_router.post(
|
||||
"/logs",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def ingest_callback_logs(
|
||||
body: CallbackLogsRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
) -> CallbackLogsResponse:
|
||||
"""
|
||||
Replay a batch of finished logging payloads through the callback fan-out.
|
||||
|
||||
Admin-only: the payloads write spend logs and trigger every callback, so this
|
||||
is a trusted internal route, not a public surface.
|
||||
"""
|
||||
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail="/v1/rust_control_plane/logs is admin-only (proxy admin key required).",
|
||||
)
|
||||
|
||||
return await CallbackLogsReplayer().replay_batch(body.records)
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
44
litellm/types/proxy/callback_logs_endpoints.py
Normal file
44
litellm/types/proxy/callback_logs_endpoints.py
Normal file
|
|
@ -0,0 +1,44 @@
|
|||
"""
|
||||
Types for the callback-logs ingest endpoint (POST /v1/callbacks/logs).
|
||||
|
||||
External producers (e.g. the litellm-rust gateway) POST finished logging
|
||||
payloads here; the proxy replays them through the standard callback fan-out.
|
||||
"""
|
||||
|
||||
from typing import Any, Literal, Optional
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from litellm.constants import MAX_CALLBACK_LOG_RECORDS
|
||||
|
||||
|
||||
class CallbackLogRecord(BaseModel):
|
||||
"""A single finished logging event to replay through the callbacks."""
|
||||
|
||||
status: Literal["success", "failure"]
|
||||
standard_logging_payload: dict[str, Any]
|
||||
error: Optional[str] = None
|
||||
|
||||
|
||||
class CallbackLogsRequest(BaseModel):
|
||||
"""A batch of logging events posted by an external producer."""
|
||||
|
||||
# Bounded so one POST can't trigger an unbounded callback/DB fan-out (each
|
||||
# record fires every registered integration). Over the cap → 422.
|
||||
records: list[CallbackLogRecord] = Field(..., max_length=MAX_CALLBACK_LOG_RECORDS)
|
||||
|
||||
|
||||
class CallbackLogFailure(BaseModel):
|
||||
"""A record that failed to replay, identified by its index in the batch."""
|
||||
|
||||
index: int
|
||||
error: str
|
||||
|
||||
|
||||
class CallbackLogsResponse(BaseModel):
|
||||
"""Per-batch result: counts plus per-record failure detail so the caller can
|
||||
distinguish a transient callback error from a structurally bad payload."""
|
||||
|
||||
processed: int
|
||||
failed: int
|
||||
failures: list[CallbackLogFailure] = Field(default_factory=list)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
93
tests/e2e/budgets/BUDGET_CODE_MATRIX.md
Normal file
93
tests/e2e/budgets/BUDGET_CODE_MATRIX.md
Normal file
|
|
@ -0,0 +1,93 @@
|
|||
# Budget Code Matrix
|
||||
|
||||
What LiteLLM actually implements for budgets: every entity that can carry a dollar
|
||||
budget, how the limit is enforced, and where in the code it happens. This is the
|
||||
"what we support" reference; the companion `BUDGET_TEST_COVERAGE_MATRIX.md` maps
|
||||
each row to its tests and the e2e gaps.
|
||||
|
||||
Over-budget surfaces as a `budget_exceeded` error (the live suite
|
||||
`tests/otel_tests/test_e2e_budgeting.py` asserts `type == "budget_exceeded"`,
|
||||
`code == "429"`); the underlying `BudgetExceededError` is defined in
|
||||
`litellm/exceptions.py` (`status_code=400`). Enforcement runs in `common_checks()`
|
||||
/ `auth_checks.py` at auth time, plus pre-call reservation in
|
||||
`budget_reservation.py`.
|
||||
|
||||
Legend for "Enforced": **block** = request rejected; **filter** = router skips the
|
||||
deployment; **alert** = notify only, request proceeds.
|
||||
|
||||
---
|
||||
|
||||
## 1. Per-entity dollar budgets
|
||||
|
||||
| Entity | Budget stored | Hard `max_budget` | Soft budget | Per-window | Model budget | Reset by `budget_duration` |
|
||||
|--------|---------------|-------------------|-------------|------------|--------------|----------------------------|
|
||||
| API key | `LiteLLM_VerificationToken` (direct cols + `budget_id` FK) | block (`_virtual_key_max_budget_check`) | alert (`_virtual_key_soft_budget_check`) + 80% alert | block (`_virtual_key_multi_budget_check`) | block (`model_max_budget_limiter.is_key_within_model_budget`) | keys reset job |
|
||||
| Internal user | `LiteLLM_UserTable` (direct cols) | block (`common_checks`, only when not on a team) | - | - | via `model_max_budget` json | users reset job |
|
||||
| Team | `LiteLLM_TeamTable` (direct cols) | block (`_team_max_budget_check`) | alert (`_team_soft_budget_check`) | block (`_team_multi_budget_check`) | via `model_max_budget` | teams reset job |
|
||||
| Team member | `LiteLLM_TeamMembership` -> `LiteLLM_BudgetTable` | block (`_check_team_member_budget`) | - | - | - | budget-table reset job |
|
||||
| End-user / customer | `LiteLLM_EndUserTable` -> `LiteLLM_BudgetTable` | block (`_check_end_user_budget`) | - | - | block (`is_end_user_within_model_budget`) | budget-table reset job |
|
||||
| Organization | `LiteLLM_OrganizationTable` -> `LiteLLM_BudgetTable` | block (`_organization_max_budget_check`) | - | - | via budget-table | budget-table reset job |
|
||||
| Tag | `LiteLLM_TagTable` -> `LiteLLM_BudgetTable` | block (`_tag_max_budget_check`) | - | - | via budget-table | budget-table reset job |
|
||||
| Project | `LiteLLM_ProjectTable` -> `LiteLLM_BudgetTable` | block (`_project_max_budget_check`) | alert (`_project_soft_budget_check`) | - | - | budget-table reset job |
|
||||
| Provider (router) | config `provider_budget_config` (in-memory) | filter (`router_strategy/budget_limiter`) | - | yes (time window) | - | window TTL |
|
||||
| Global proxy | `litellm.max_budget` (config) | block (`_global_proxy_budget_check`) | - | - | - | - |
|
||||
|
||||
Notes / flags from the code:
|
||||
- **User budget only enforced off-team**: `common_checks` skips the personal-user
|
||||
budget when the key belongs to a team (team budget governs instead).
|
||||
- **Comparison operators are inconsistent**: key/user use `>=`, team/end-user main
|
||||
budget use `>`. Spend exactly at `max_budget` blocks a key but not a team.
|
||||
- **Provider budgets are filter-only**: an over-budget provider is removed from
|
||||
routing; if all are over budget the router raises
|
||||
`no_deployments_with_provider_budget_routing` (not a per-entity block).
|
||||
- **Enforcement timing differs by entity**: key / user / org / team-member / tag /
|
||||
model enforce off real-time reservation counters (block within ~2 calls);
|
||||
**end-user** enforcement reads `EndUserTable.spend`, which only updates on the
|
||||
`proxy_batch_write_at` flush, so it lags by that interval (verified live).
|
||||
|
||||
## 2. Budget mechanisms
|
||||
|
||||
| Mechanism | What it does | Code |
|
||||
|-----------|--------------|------|
|
||||
| Pre-call reservation | Estimates max request cost, atomically reserves against redis spend counters for key/team/user/end_user/tag/team_member/org before the call; blocks if a counter would exceed | `spend_tracking/budget_reservation.py` |
|
||||
| Post-call reconciliation | Adjusts the reservation to the actual cost once known | `reconcile_budget_reservation` |
|
||||
| Read-time enforcement | Auth-time check of current spend vs `max_budget` | `auth_checks.common_checks` + per-entity `_*_max_budget_check` |
|
||||
| Soft budget / alerts | At `soft_budget` (or 80% of max) fire Slack/email alert, do not block | `_virtual_key_soft_budget_check`, `_team_soft_budget_check`, `budget_alerts` |
|
||||
| Multi-window budgets | `budget_limits` list of `{budget_duration, max_budget}`; each window enforced + reset independently | `_virtual_key_multi_budget_check`, `reset_budget_windows` |
|
||||
| Model-level budgets | `model_max_budget` dict (per model: `budget_limit` + `time_period`) on key/user/end_user | `hooks/model_max_budget_limiter.py` |
|
||||
| Reset by duration | Job zeros `spend`, recomputes `budget_reset_at = now + duration_in_seconds(budget_duration)`, invalidates redis counters | `common_utils/reset_budget_job.py`, `duration_parser.duration_in_seconds` |
|
||||
| Zero-cost bypass | Models with no configured price bypass budget reservation | `budget_reservation` zero-cost path |
|
||||
|
||||
## 3. Budget management surface (endpoints)
|
||||
|
||||
| Action | Endpoint | Handler |
|
||||
|--------|----------|---------|
|
||||
| Create budget | `POST /budget/new` | `new_budget` |
|
||||
| Update budget | `POST /budget/update` | `update_budget` |
|
||||
| Budget info | `POST /budget/info` (`{"budgets": [id]}`) | `info_budget` |
|
||||
| Budget settings | `GET /budget/settings` | `budget_settings` |
|
||||
| List budgets | `GET /budget/list` | `list_budget` |
|
||||
| Delete budget | `POST /budget/delete` (`{"id": id}`) | `delete_budget` |
|
||||
| Set on key | `POST /key/generate`, `/key/update` (`max_budget`, `soft_budget`, `budget_duration`, `model_max_budget`, `budget_id`) | key mgmt |
|
||||
| Set on user | `POST /user/new` (`max_budget`, `budget_duration`) | internal user |
|
||||
| Set on team | `POST /team/new` (`max_budget`, `soft_budget`, `team_member_budget`) | team |
|
||||
| Set on team member | `POST /team/member_add` (`max_budget_in_team`) | team |
|
||||
| Set on org | `POST /organization/new` (`max_budget`, `soft_budget`, `model_max_budget`) | org |
|
||||
| Set on customer | `POST /customer/new`, `/customer/update` (`max_budget`, `budget_id`) | customer |
|
||||
| Set on tag | `POST /tag/new`, `/tag/update` (`max_budget`) | tag mgmt |
|
||||
| Read budget+spend | `/key/info`, `/user/info`, `/team/info`, `/organization/info`, `/customer/info`, `/budget/info` | per-entity info |
|
||||
|
||||
Endpoint method/shape gotchas verified live: `/organization/delete` is **DELETE**
|
||||
with `{"organization_ids": [id]}`; `/budget/info` takes `{"budgets": [id]}`;
|
||||
`model_max_budget` entries use `{"budget_limit", "time_period"}`.
|
||||
|
||||
## 4. Config knobs
|
||||
|
||||
| Setting | Effect |
|
||||
|---------|--------|
|
||||
| `litellm.max_budget` | proxy-wide hard cap (global proxy budget) |
|
||||
| `max_internal_user_budget` / `default_max_internal_user_budget` | default `max_budget` for internal users |
|
||||
| `internal_user_budget_duration` | default reset duration for internal users |
|
||||
| `max_end_user_budget` / `max_end_user_budget_id` | default budget for end-users |
|
||||
| `default_team_params` | default `max_budget` / `budget_duration` / limits for teams |
|
||||
| `provider_budget_config` (router) | per-provider spend caps + windows |
|
||||
78
tests/e2e/budgets/BUDGET_TEST_COVERAGE_MATRIX.md
Normal file
78
tests/e2e/budgets/BUDGET_TEST_COVERAGE_MATRIX.md
Normal file
|
|
@ -0,0 +1,78 @@
|
|||
# Budget Test Coverage Matrix
|
||||
|
||||
Maps every row of `BUDGET_CODE_MATRIX.md` (what LiteLLM implements) to its tests
|
||||
and level, then marks the live e2e coverage this suite adds.
|
||||
|
||||
Levels: `unit` mocked (`AsyncMock` on `get_current_spend`/prisma); `router` live
|
||||
router with fake deployments; `live-e2e` real proxy, real key/team, real requests
|
||||
until blocked. Status: `covered` / `partial` / `gap`.
|
||||
|
||||
Pre-existing live coverage outside this suite:
|
||||
- `tests/otel_tests/test_e2e_budgeting.py` - key + team enforcement, budget update.
|
||||
- `tests/local_testing/test_router_budget_limiter.py` - provider / tag / deployment
|
||||
budgets at the router.
|
||||
|
||||
This suite (`tests/e2e/budgets/`) adds the missing live coverage and runs
|
||||
on the shared lifecycle (every entity it creates is deleted on teardown).
|
||||
|
||||
---
|
||||
|
||||
## Per-entity enforcement
|
||||
|
||||
| Entity | Unit | Pre-existing live | This suite (live) | Status |
|
||||
|--------|------|-------------------|-------------------|--------|
|
||||
| API key | `test_budget_reservation.py`, `test_max_budget_limiter.py` | `otel_tests` | `test_budget_enforcement_e2e::test_key_budget_blocks` | **covered** |
|
||||
| Team | `test_team_budget_limits.py` | `otel_tests` | (org test builds a team) | **covered** |
|
||||
| Internal user | auth unit tests | - | `test_internal_user_budget_blocks` | **covered (new)** |
|
||||
| Team member | `test_team_member_budget.py` | - | `test_team_member_budget_blocks` | **covered (new)** |
|
||||
| End-user / customer | `test_custom_auth_end_user_budget.py` | - | `test_end_user_budget_blocks` | **covered (new)** |
|
||||
| Organization | `test_organization_budget_enforcement.py` (flagged weak) | - | `test_organization_budget_blocks` | **covered (new)** |
|
||||
| Tag (proxy-level) | - | router only | `test_tag_budget_e2e::test_tag_budget_blocks_tagged_requests` | **covered (new)** |
|
||||
| Model-level (`model_max_budget`) | `test_unit_test_max_model_budget_limiter.py` | - | `test_model_max_budget_e2e::test_model_max_budget_isolates_per_model` | **covered (new)** |
|
||||
| Provider (router) | `test_budget_limiter_hotpath.py` | `test_router_budget_limiter.py` | - | **covered** (router) |
|
||||
| Global proxy (`litellm.max_budget`) | unit | - | - | **gap** (needs a config-level cap; not key-settable) |
|
||||
|
||||
## Budget mechanisms
|
||||
|
||||
| Mechanism | Unit | This suite (live) | Status |
|
||||
|-----------|------|-------------------|--------|
|
||||
| Pre-call reservation | `test_budget_reservation.py` | exercised by every enforcement test | **partial** |
|
||||
| Soft budget / alerts | `SlackAlerting/test_budget_alert_types.py` | `test_soft_budget_e2e::test_soft_budget_does_not_block` | **covered (new)** (block-vs-alert; the alert side-effect itself stays unit) |
|
||||
| Budget CRUD | `test_budget_endpoints.py` | `test_budget_crud_e2e` (roundtrip + delete) | **covered (new)** |
|
||||
| Reset scheduling | `test_proxy_budget_reset.py` | `test_budget_crud_e2e::test_budget_duration_schedules_reset_on_key` | **covered (new)** (scheduling; actual zeroing is time-dependent -> unit) |
|
||||
| Multi-window budgets | `test_multi_budget_windows.py` | - | **gap** (window setup is fiddly; left to unit for now) |
|
||||
| Read budget+spend | `test_spend_management_endpoints.py` | `/key/info` asserted in CRUD + enforcement | **partial** |
|
||||
|
||||
## Remaining gaps (intentionally not live-tested)
|
||||
|
||||
- **Global proxy budget** (`litellm.max_budget`): set via proxy config, not a
|
||||
per-key API, so it needs a dedicated proxy boot with that config rather than a
|
||||
runtime-created entity. Out of scope for the per-entity suite.
|
||||
- **Multi-window budgets**: the `budget_limits` list shape and per-window reset are
|
||||
covered by `test_multi_budget_windows.py` (unit); a live version would need to
|
||||
wait out a short window to see the reset, which is time-dependent.
|
||||
- **Soft-budget alert delivery**: whether the Slack/email actually fires is not
|
||||
observable from the proxy API; unit tests own that. The live test pins the
|
||||
load-bearing behavior (soft does not block).
|
||||
- **Reset zeroing after the window elapses**: time-dependent; unit tests own the
|
||||
reset-job logic. The live test pins that `budget_reset_at` is scheduled.
|
||||
|
||||
## This suite's files
|
||||
|
||||
| File | Covers |
|
||||
|------|--------|
|
||||
| `test_budget_enforcement_e2e.py` | key / internal-user / end-user / organization / team-member hard enforcement |
|
||||
| `test_model_max_budget_e2e.py` | per-model caps isolate by model |
|
||||
| `test_soft_budget_e2e.py` | soft budget alerts but does not block |
|
||||
| `test_tag_budget_e2e.py` | proxy-level tag budget blocks tagged requests, spares others |
|
||||
| `test_budget_crud_e2e.py` | `/budget/*` CRUD roundtrip + delete + `budget_reset_at` scheduling |
|
||||
|
||||
## Pattern + timing
|
||||
|
||||
Create the entity with a tiny `max_budget`, drive spend until a `budget_exceeded`
|
||||
block. The enforcement helper is two-phase: a fast warmup (key/user/org/member/tag/
|
||||
model block within ~2 calls off real-time counters), then a poll across the ~60s
|
||||
batch-write window (end-user enforcement reads table spend that lags). Skip on a
|
||||
non-budget error (provider down / key missing); fail if the budget is never
|
||||
enforced. Chat tests use `gpt-5.5` (the model with a working key on the reference
|
||||
proxy); swap the literal if your proxy differs.
|
||||
418
tests/e2e/budgets/budget_client.py
Normal file
418
tests/e2e/budgets/budget_client.py
Normal file
|
|
@ -0,0 +1,418 @@
|
|||
"""Client for budget e2e tests: the shared Gateway plus budget-bearing entity
|
||||
management (user / team / team-member / org / customer / tag / budget-table) and
|
||||
info reads.
|
||||
|
||||
Over-budget surfaces as a ``budget_exceeded`` error; ``is_budget_block`` detects it
|
||||
on a chat outcome. Create methods return the new id and raise on failure; tests
|
||||
register the matching delete with ``resources.defer(...)`` for cleanup. The request
|
||||
and response models are co-located here because only this suite uses them.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from pydantic import AliasPath, BaseModel, Field, RootModel
|
||||
|
||||
from e2e_gateway import Gateway, build_gateway
|
||||
from e2e_http import NoBody, StreamingResponse, Success, unwrap
|
||||
from models import (
|
||||
BudgetWindow,
|
||||
ChatBody,
|
||||
ChatMessage,
|
||||
ChatMetadata,
|
||||
KeyGenerateBody,
|
||||
ModelBudgetEntry,
|
||||
)
|
||||
|
||||
|
||||
class UserNewBody(BaseModel):
|
||||
max_budget: float
|
||||
|
||||
|
||||
class UserNewResponse(BaseModel):
|
||||
user_id: str
|
||||
|
||||
|
||||
class UserDeleteBody(BaseModel):
|
||||
user_ids: list[str]
|
||||
|
||||
|
||||
class CustomerNewBody(BaseModel):
|
||||
user_id: str
|
||||
max_budget: float
|
||||
|
||||
|
||||
class OrgNewBody(BaseModel):
|
||||
organization_alias: str
|
||||
max_budget: float
|
||||
|
||||
|
||||
class OrgNewResponse(BaseModel):
|
||||
organization_id: str
|
||||
|
||||
|
||||
class OrgDeleteBody(BaseModel):
|
||||
organization_ids: list[str]
|
||||
|
||||
|
||||
class TeamMember(BaseModel):
|
||||
role: str
|
||||
user_id: str
|
||||
|
||||
|
||||
class TeamNewBody(BaseModel):
|
||||
team_alias: str
|
||||
max_budget: float | None = None
|
||||
organization_id: str | None = None
|
||||
budget_limits: list[BudgetWindow] | None = None
|
||||
|
||||
|
||||
class TeamNewResponse(BaseModel):
|
||||
team_id: str
|
||||
|
||||
|
||||
class TeamDeleteBody(BaseModel):
|
||||
team_ids: list[str]
|
||||
|
||||
|
||||
class TeamMemberAddBody(BaseModel):
|
||||
team_id: str
|
||||
member: TeamMember
|
||||
max_budget_in_team: float | None = None
|
||||
|
||||
|
||||
class TeamMemberUpdateBody(BaseModel):
|
||||
team_id: str
|
||||
user_id: str
|
||||
max_budget_in_team: float | None = None
|
||||
budget_duration: str | None = None
|
||||
|
||||
|
||||
class TeamMembershipRow(BaseModel):
|
||||
user_id: str | None = None
|
||||
budget_reset_at: str | None = Field(
|
||||
default=None,
|
||||
validation_alias=AliasPath("litellm_budget_table", "budget_reset_at"),
|
||||
)
|
||||
|
||||
|
||||
class TeamInfoParams(BaseModel):
|
||||
team_id: str
|
||||
|
||||
|
||||
class TeamInfoResponse(BaseModel):
|
||||
team_memberships: list[TeamMembershipRow] = []
|
||||
|
||||
|
||||
class TagNewBody(BaseModel):
|
||||
name: str
|
||||
max_budget: float
|
||||
|
||||
|
||||
class TagDeleteBody(BaseModel):
|
||||
name: str
|
||||
|
||||
|
||||
class BudgetNewBody(BaseModel):
|
||||
max_budget: float
|
||||
soft_budget: float | None = None
|
||||
budget_duration: str | None = None
|
||||
|
||||
|
||||
class BudgetNewResponse(BaseModel):
|
||||
budget_id: str
|
||||
|
||||
|
||||
class BudgetDeleteBody(BaseModel):
|
||||
id: str
|
||||
|
||||
|
||||
class BudgetInfoBody(BaseModel):
|
||||
budgets: list[str]
|
||||
|
||||
|
||||
class BudgetRow(BaseModel):
|
||||
budget_id: str | None = None
|
||||
max_budget: float | None = None
|
||||
soft_budget: float | None = None
|
||||
budget_duration: str | None = None
|
||||
budget_reset_at: str | None = None
|
||||
|
||||
|
||||
class BudgetInfoResponse(RootModel[list[BudgetRow]]):
|
||||
pass
|
||||
|
||||
|
||||
def is_budget_block(result: StreamingResponse) -> bool:
|
||||
"""True if the call was rejected for being over budget (vs a provider error)."""
|
||||
return not result.ok and "budget_exceeded" in result.body
|
||||
|
||||
|
||||
def model_budget(model: str, limit: float, period: str = "30d") -> dict[str, ModelBudgetEntry]:
|
||||
"""A model_max_budget entry: per-model cap with a reset window."""
|
||||
return {model: ModelBudgetEntry(budget_limit=limit, time_period=period)}
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BudgetClient:
|
||||
gateway: Gateway
|
||||
|
||||
# ---- generic key ops (delegate to the shared Gateway) ---------------
|
||||
|
||||
def generate_key(
|
||||
self,
|
||||
*,
|
||||
models: list[str] | None = None,
|
||||
max_budget: float | None = None,
|
||||
soft_budget: float | None = None,
|
||||
budget_duration: str | None = None,
|
||||
budget_id: str | None = None,
|
||||
user_id: str | None = None,
|
||||
team_id: str | None = None,
|
||||
model_max_budget: dict[str, ModelBudgetEntry] | None = None,
|
||||
budget_limits: list[BudgetWindow] | None = None,
|
||||
) -> str:
|
||||
return self.gateway.generate_key(
|
||||
KeyGenerateBody(
|
||||
models=models or [],
|
||||
max_budget=max_budget,
|
||||
soft_budget=soft_budget,
|
||||
budget_duration=budget_duration,
|
||||
budget_id=budget_id,
|
||||
user_id=user_id,
|
||||
team_id=team_id,
|
||||
model_max_budget=model_max_budget,
|
||||
budget_limits=budget_limits,
|
||||
)
|
||||
)
|
||||
|
||||
def delete_key(self, key: str) -> None:
|
||||
self.gateway.delete_key(key)
|
||||
|
||||
def delete_customers(self, user_ids: list[str]) -> None:
|
||||
self.gateway.delete_customers(user_ids)
|
||||
|
||||
# ---- chat (raw HTTP outcome: a budget block surfaces as a non-2xx) --
|
||||
|
||||
def chat(
|
||||
self,
|
||||
key: str,
|
||||
model: str,
|
||||
content: str,
|
||||
*,
|
||||
max_tokens: int | None = None,
|
||||
user: str | None = None,
|
||||
tags: list[str] | None = None,
|
||||
) -> StreamingResponse:
|
||||
return self.gateway.transport.send(
|
||||
"/chat/completions",
|
||||
headers=self.gateway.transport.bearer(key),
|
||||
json=ChatBody(
|
||||
model=model,
|
||||
messages=[ChatMessage(role="user", content=content)],
|
||||
max_tokens=max_tokens,
|
||||
user=user,
|
||||
metadata=ChatMetadata(tags=tags) if tags else None,
|
||||
),
|
||||
)
|
||||
|
||||
# ---- internal user --------------------------------------------------
|
||||
|
||||
def create_user(self, *, max_budget: float) -> str:
|
||||
return unwrap(
|
||||
self.gateway.transport.post(
|
||||
"/user/new",
|
||||
headers=self.gateway.transport.master,
|
||||
json=UserNewBody(max_budget=max_budget),
|
||||
response_type=UserNewResponse,
|
||||
)
|
||||
).user_id
|
||||
|
||||
def delete_user(self, user_id: str) -> None:
|
||||
_ = self.gateway.transport.post(
|
||||
"/user/delete",
|
||||
headers=self.gateway.transport.master,
|
||||
json=UserDeleteBody(user_ids=[user_id]),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
||||
# ---- customer / end-user -------------------------------------------
|
||||
|
||||
def create_customer(self, customer_id: str, *, max_budget: float) -> str:
|
||||
resp = self.gateway.transport.send(
|
||||
"/customer/new",
|
||||
headers=self.gateway.transport.master,
|
||||
json=CustomerNewBody(user_id=customer_id, max_budget=max_budget),
|
||||
)
|
||||
assert resp.ok, resp.body
|
||||
return customer_id
|
||||
|
||||
# ---- organization ---------------------------------------------------
|
||||
|
||||
def create_org(self, *, max_budget: float, alias: str) -> str:
|
||||
return unwrap(
|
||||
self.gateway.transport.post(
|
||||
"/organization/new",
|
||||
headers=self.gateway.transport.master,
|
||||
json=OrgNewBody(organization_alias=alias, max_budget=max_budget),
|
||||
response_type=OrgNewResponse,
|
||||
)
|
||||
).organization_id
|
||||
|
||||
def delete_org(self, org_id: str) -> None:
|
||||
_ = self.gateway.transport.delete(
|
||||
"/organization/delete",
|
||||
headers=self.gateway.transport.master,
|
||||
json=OrgDeleteBody(organization_ids=[org_id]),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
||||
# ---- team -----------------------------------------------------------
|
||||
|
||||
def create_team(
|
||||
self,
|
||||
*,
|
||||
alias: str,
|
||||
max_budget: float | None = None,
|
||||
organization_id: str | None = None,
|
||||
budget_limits: list[BudgetWindow] | None = None,
|
||||
) -> str:
|
||||
return unwrap(
|
||||
self.gateway.transport.post(
|
||||
"/team/new",
|
||||
headers=self.gateway.transport.master,
|
||||
json=TeamNewBody(
|
||||
team_alias=alias,
|
||||
max_budget=max_budget,
|
||||
organization_id=organization_id,
|
||||
budget_limits=budget_limits,
|
||||
),
|
||||
response_type=TeamNewResponse,
|
||||
)
|
||||
).team_id
|
||||
|
||||
def delete_team(self, team_id: str) -> None:
|
||||
_ = self.gateway.transport.post(
|
||||
"/team/delete",
|
||||
headers=self.gateway.transport.master,
|
||||
json=TeamDeleteBody(team_ids=[team_id]),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
||||
def add_team_member(self, team_id: str, user_id: str, *, max_budget_in_team: float | None = None) -> None:
|
||||
resp = self.gateway.transport.send(
|
||||
"/team/member_add",
|
||||
headers=self.gateway.transport.master,
|
||||
json=TeamMemberAddBody(
|
||||
team_id=team_id,
|
||||
member=TeamMember(role="user", user_id=user_id),
|
||||
max_budget_in_team=max_budget_in_team,
|
||||
),
|
||||
)
|
||||
assert resp.ok, resp.body
|
||||
|
||||
def update_team_member(
|
||||
self,
|
||||
team_id: str,
|
||||
user_id: str,
|
||||
*,
|
||||
max_budget_in_team: float | None = None,
|
||||
budget_duration: str | None = None,
|
||||
) -> None:
|
||||
resp = self.gateway.transport.send(
|
||||
"/team/member_update",
|
||||
headers=self.gateway.transport.master,
|
||||
json=TeamMemberUpdateBody(
|
||||
team_id=team_id,
|
||||
user_id=user_id,
|
||||
max_budget_in_team=max_budget_in_team,
|
||||
budget_duration=budget_duration,
|
||||
),
|
||||
)
|
||||
assert resp.ok, resp.body
|
||||
|
||||
def member_budget_reset_at(self, team_id: str, user_id: str) -> str | None:
|
||||
"""The member's per-team budget_reset_at as /team/info reports it, or None if
|
||||
no reset is scheduled. The reset job advances this each time the window
|
||||
elapses; a job that skips the row leaves it pinned forever."""
|
||||
result = self.gateway.transport.get(
|
||||
"/team/info",
|
||||
headers=self.gateway.transport.master,
|
||||
params=TeamInfoParams(team_id=team_id),
|
||||
response_type=TeamInfoResponse,
|
||||
)
|
||||
match result:
|
||||
case Success(data=data):
|
||||
return next(
|
||||
(row.budget_reset_at for row in data.team_memberships if row.user_id == user_id),
|
||||
None,
|
||||
)
|
||||
case _:
|
||||
return None
|
||||
|
||||
# ---- tag ------------------------------------------------------------
|
||||
|
||||
def create_tag(self, name: str, *, max_budget: float) -> str:
|
||||
resp = self.gateway.transport.send(
|
||||
"/tag/new",
|
||||
headers=self.gateway.transport.master,
|
||||
json=TagNewBody(name=name, max_budget=max_budget),
|
||||
)
|
||||
assert resp.ok, resp.body
|
||||
return name
|
||||
|
||||
def delete_tag(self, name: str) -> None:
|
||||
_ = self.gateway.transport.post(
|
||||
"/tag/delete",
|
||||
headers=self.gateway.transport.master,
|
||||
json=TagDeleteBody(name=name),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
||||
# ---- budget table ---------------------------------------------------
|
||||
|
||||
def create_budget(
|
||||
self,
|
||||
*,
|
||||
max_budget: float,
|
||||
soft_budget: float | None = None,
|
||||
budget_duration: str | None = None,
|
||||
) -> str:
|
||||
return unwrap(
|
||||
self.gateway.transport.post(
|
||||
"/budget/new",
|
||||
headers=self.gateway.transport.master,
|
||||
json=BudgetNewBody(
|
||||
max_budget=max_budget,
|
||||
soft_budget=soft_budget,
|
||||
budget_duration=budget_duration,
|
||||
),
|
||||
response_type=BudgetNewResponse,
|
||||
)
|
||||
).budget_id
|
||||
|
||||
def delete_budget(self, budget_id: str) -> None:
|
||||
_ = self.gateway.transport.post(
|
||||
"/budget/delete",
|
||||
headers=self.gateway.transport.master,
|
||||
json=BudgetDeleteBody(id=budget_id),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
||||
def budget_info(self, budget_id: str) -> tuple[BudgetRow, ...]:
|
||||
result = self.gateway.transport.post(
|
||||
"/budget/info",
|
||||
headers=self.gateway.transport.master,
|
||||
json=BudgetInfoBody(budgets=[budget_id]),
|
||||
response_type=BudgetInfoResponse,
|
||||
)
|
||||
match result:
|
||||
case Success(data=data):
|
||||
return tuple(data.root)
|
||||
case _:
|
||||
return ()
|
||||
|
||||
|
||||
def build_client() -> BudgetClient:
|
||||
return BudgetClient(gateway=build_gateway())
|
||||
16
tests/e2e/budgets/conftest.py
Normal file
16
tests/e2e/budgets/conftest.py
Normal file
|
|
@ -0,0 +1,16 @@
|
|||
"""Budgets suite's `client` fixture.
|
||||
|
||||
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
|
||||
live in the parent tests/e2e/conftest.py. BudgetClient holds the shared Gateway,
|
||||
so the `resources` fixture cleans up keys through it; tests register entity deletes
|
||||
via `resources.defer(...)`.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from budget_client import BudgetClient, build_client
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def client() -> BudgetClient:
|
||||
return build_client()
|
||||
60
tests/e2e/budgets/test_budget_crud_e2e.py
Normal file
60
tests/e2e/budgets/test_budget_crud_e2e.py
Normal file
|
|
@ -0,0 +1,60 @@
|
|||
"""Live e2e for the budget management surface (no LLM calls, fast).
|
||||
|
||||
Covers the budget-table CRUD round-trip and that `budget_duration` schedules a
|
||||
`budget_reset_at`. The actual zeroing after the window is time-dependent, so we
|
||||
assert the reset is *scheduled* (now + duration), not waited out.
|
||||
"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import pytest
|
||||
|
||||
from budget_client import BudgetClient
|
||||
from lifecycle import ResourceManager
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
|
||||
def test_budget_crud_roundtrip(client: BudgetClient, resources: ResourceManager) -> None:
|
||||
budget_id = client.create_budget(max_budget=12.5, soft_budget=10.0, budget_duration="30d")
|
||||
resources.defer(lambda: client.delete_budget(budget_id))
|
||||
|
||||
rows = client.budget_info(budget_id)
|
||||
assert rows, f"/budget/info returned nothing for {budget_id}"
|
||||
row = rows[0]
|
||||
assert row.max_budget == 12.5
|
||||
assert row.soft_budget == 10.0
|
||||
assert row.budget_reset_at, "budget_duration did not schedule a reset"
|
||||
|
||||
# Attach the budget to a key and confirm the key reflects it.
|
||||
key = client.generate_key(budget_id=budget_id)
|
||||
resources.defer(lambda: client.delete_key(key))
|
||||
info = client.gateway.key_info(key)
|
||||
linked = info.litellm_budget_table
|
||||
assert info.budget_id == budget_id or (linked is not None and linked.max_budget == 12.5), (
|
||||
f"key does not reflect attached budget: {info.budget_id}, {linked}"
|
||||
)
|
||||
|
||||
|
||||
def test_budget_delete_removes_it(client: BudgetClient, resources: ResourceManager) -> None:
|
||||
budget_id = client.create_budget(max_budget=1.0)
|
||||
resources.defer(lambda: client.delete_budget(budget_id))
|
||||
client.delete_budget(budget_id)
|
||||
assert not client.budget_info(budget_id), "budget still present after delete"
|
||||
|
||||
|
||||
def test_budget_duration_schedules_reset_on_key(client: BudgetClient, resources: ResourceManager) -> None:
|
||||
key = client.generate_key(max_budget=10.0, budget_duration="30d")
|
||||
resources.defer(lambda: client.delete_key(key))
|
||||
|
||||
reset_at = client.gateway.key_info(key).budget_reset_at
|
||||
assert reset_at, "budget_duration did not set budget_reset_at on the key"
|
||||
|
||||
# budget_duration schedules a FUTURE reset. Don't assume now+30d exactly: the
|
||||
# proxy may align the reset to a calendar boundary (e.g. start of next month),
|
||||
# so "30d" can land ~12 days out mid-month. Assert it's scheduled ahead.
|
||||
|
||||
# get current time -> assert budget from days_left - budget_duration == days_left
|
||||
reset_dt = datetime.fromisoformat(str(reset_at).replace("Z", "+00:00"))
|
||||
days_out = (reset_dt - datetime.now(timezone.utc)).total_seconds() / 86400
|
||||
assert 0 < days_out < 40, f"reset should be scheduled ahead, got {days_out:.1f}d out"
|
||||
148
tests/e2e/budgets/test_budget_enforcement_e2e.py
Normal file
148
tests/e2e/budgets/test_budget_enforcement_e2e.py
Normal file
|
|
@ -0,0 +1,148 @@
|
|||
"""Live e2e: a tiny max_budget on an entity actually blocks requests.
|
||||
|
||||
Each entity is an E2ECase (lifecycle.E2ECase) driven by run_case: init() creates
|
||||
the budgeted entity + a key, run() drives spend until a `budget_exceeded` block,
|
||||
teardown() deletes everything init() created (always runs, even on failure/skip).
|
||||
Covers the entities with no prior live coverage - internal user, end-user,
|
||||
organization, team member. See BUDGET_TEST_COVERAGE_MATRIX.md.
|
||||
|
||||
A non-budget error fails hard (never a skip); if calls never get blocked, budget
|
||||
enforcement is broken -> fail.
|
||||
"""
|
||||
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable, List, Type
|
||||
|
||||
import pytest
|
||||
|
||||
from budget_client import BudgetClient, is_budget_block
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import require_successful_call
|
||||
from lifecycle import run_case
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") -> None:
|
||||
"""Send paid calls until the entity's budget blocks one. Key/user/org/member
|
||||
block within a couple calls off real-time reservation counters; the end-user
|
||||
budget enforces off table spend that lands on the batch write, so it takes a
|
||||
few more. A non-budget error fails hard (never a skip)."""
|
||||
for _ in range(40):
|
||||
result = client.chat(
|
||||
key,
|
||||
"claude-haiku-4-5",
|
||||
f"spend {unique_marker()}",
|
||||
max_tokens=16,
|
||||
user=user or None,
|
||||
)
|
||||
if is_budget_block(result):
|
||||
return
|
||||
require_successful_call(result)
|
||||
time.sleep(2)
|
||||
pytest.fail("budget never enforced within the call budget")
|
||||
|
||||
|
||||
@dataclass
|
||||
class _BudgetCase:
|
||||
"""Base E2ECase: a key under some budgeted entity must get blocked.
|
||||
|
||||
Subclasses set up the budgeted entity in init() and register every created id
|
||||
in `_undo` (run LIFO in teardown so a key is deleted before its team/org).
|
||||
"""
|
||||
|
||||
client: BudgetClient
|
||||
key: str = ""
|
||||
_undo: List[Callable[[], None]] = field(
|
||||
default_factory=list
|
||||
) # mutable-ok: per-case teardown registry
|
||||
|
||||
def init(self) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
def run(self) -> None:
|
||||
_assert_budget_blocks(self.client, self.key)
|
||||
|
||||
def teardown(self) -> None:
|
||||
for undo in reversed(self._undo):
|
||||
undo()
|
||||
|
||||
|
||||
class KeyBudgetCase(_BudgetCase):
|
||||
def init(self) -> None:
|
||||
self.key = self.client.generate_key(max_budget=3e-6)
|
||||
self._undo.append(lambda: self.client.delete_key(self.key))
|
||||
|
||||
|
||||
class InternalUserBudgetCase(_BudgetCase):
|
||||
def init(self) -> None:
|
||||
user_id = self.client.create_user(max_budget=3e-6)
|
||||
self._undo.append(lambda: self.client.delete_user(user_id))
|
||||
# personal key (no team) -> the user budget governs
|
||||
self.key = self.client.generate_key(user_id=user_id)
|
||||
self._undo.append(lambda: self.client.delete_key(self.key))
|
||||
|
||||
|
||||
class EndUserBudgetCase(_BudgetCase):
|
||||
def init(self) -> None:
|
||||
customer = f"e2e-budget-cust-{unique_marker()}"
|
||||
self.client.create_customer(customer, max_budget=3e-6)
|
||||
self._undo.append(lambda: self.client.delete_customers([customer]))
|
||||
self.key = self.client.generate_key(models=["claude-haiku-4-5"])
|
||||
self._undo.append(lambda: self.client.delete_key(self.key))
|
||||
self._customer = customer
|
||||
|
||||
def run(self) -> None:
|
||||
_assert_budget_blocks(self.client, self.key, user=self._customer)
|
||||
|
||||
|
||||
class OrganizationBudgetCase(_BudgetCase):
|
||||
def init(self) -> None:
|
||||
# Org carries the tiny budget; the team under it has none, so a block here
|
||||
# is org-level enforcement (the historically weak link).
|
||||
org_id = self.client.create_org(
|
||||
max_budget=3e-6, alias=f"e2e-budget-org-{unique_marker()}"
|
||||
)
|
||||
self._undo.append(lambda: self.client.delete_org(org_id))
|
||||
team_id = self.client.create_team(
|
||||
alias=f"e2e-budget-team-{unique_marker()}", organization_id=org_id
|
||||
)
|
||||
self._undo.append(lambda: self.client.delete_team(team_id))
|
||||
self.key = self.client.generate_key(team_id=team_id)
|
||||
self._undo.append(lambda: self.client.delete_key(self.key))
|
||||
|
||||
|
||||
class TeamMemberBudgetCase(_BudgetCase):
|
||||
def init(self) -> None:
|
||||
# Member's per-team budget is tiny while the team has a large budget, so a
|
||||
# block proves member-level (not team-level) enforcement.
|
||||
team_id = self.client.create_team(
|
||||
alias=f"e2e-budget-team-{unique_marker()}", max_budget=100.0
|
||||
)
|
||||
self._undo.append(lambda: self.client.delete_team(team_id))
|
||||
user_id = self.client.create_user(max_budget=100.0)
|
||||
self._undo.append(lambda: self.client.delete_user(user_id))
|
||||
self.client.add_team_member(team_id, user_id, max_budget_in_team=3e-6)
|
||||
self.key = self.client.generate_key(team_id=team_id, user_id=user_id)
|
||||
self._undo.append(lambda: self.client.delete_key(self.key))
|
||||
|
||||
|
||||
def _case_id(case_cls: Type[_BudgetCase]) -> str:
|
||||
return case_cls.__name__
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"case_cls",
|
||||
[
|
||||
KeyBudgetCase,
|
||||
InternalUserBudgetCase,
|
||||
EndUserBudgetCase,
|
||||
OrganizationBudgetCase,
|
||||
TeamMemberBudgetCase,
|
||||
],
|
||||
ids=_case_id,
|
||||
)
|
||||
def test_budget_enforcement(
|
||||
client: BudgetClient, case_cls: Type[_BudgetCase]
|
||||
) -> None:
|
||||
run_case(case_cls(client))
|
||||
59
tests/e2e/budgets/test_budget_reset_e2e.py
Normal file
59
tests/e2e/budgets/test_budget_reset_e2e.py
Normal file
|
|
@ -0,0 +1,59 @@
|
|||
"""Live e2e: a key budget resets (zeroes spend) after its budget_duration.
|
||||
|
||||
Short budget_duration (30s) + the fast-rescheduled reset job: a key blocked for
|
||||
exceeding its max_budget starts succeeding again once the duration elapses and the
|
||||
reset job zeroes key.spend. Closes the reset-zeroing gap in
|
||||
BUDGET_TEST_COVERAGE_MATRIX.md (reset_budget_for_litellm_keys), which the unit
|
||||
suite covers but no live test did - distinct from the per-window reset in
|
||||
test_multi_window_budget_e2e.py.
|
||||
"""
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from budget_client import BudgetClient, is_budget_block
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import require_successful_call
|
||||
from lifecycle import ResourceManager
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
|
||||
def _call(client: BudgetClient, key: str):
|
||||
return client.chat(
|
||||
key, "claude-haiku-4-5", f"reset {unique_marker()}", max_tokens=16
|
||||
)
|
||||
|
||||
|
||||
def test_key_budget_resets_after_duration(
|
||||
client: BudgetClient, resources: ResourceManager
|
||||
) -> None:
|
||||
key = client.generate_key(max_budget=3e-6, budget_duration="30s")
|
||||
resources.defer(lambda: client.delete_key(key))
|
||||
|
||||
# 1. exceed the budget -> litellm returns budget_exceeded
|
||||
blocked = False
|
||||
for _ in range(20):
|
||||
result = _call(client, key)
|
||||
if is_budget_block(result):
|
||||
blocked = True
|
||||
break
|
||||
require_successful_call(result)
|
||||
time.sleep(2)
|
||||
assert blocked, "key budget never enforced"
|
||||
|
||||
# 2. once the 30s duration elapses + the reset job runs, key.spend zeroes and
|
||||
# calls flow again. The window is wall-clock-aligned, so the reset lands up to
|
||||
# a window later, then the rescheduler (~15-20s) zeroes the spend; allow
|
||||
# generous headroom over that. A stuck rescheduler is caught by the wait-loop
|
||||
# timeout, not this elapsed bound.
|
||||
start = time.monotonic()
|
||||
while time.monotonic() < start + 150:
|
||||
time.sleep(5)
|
||||
result = _call(client, key)
|
||||
if result.ok:
|
||||
assert time.monotonic() - start < 120, "reset too slow for a 30s budget"
|
||||
return
|
||||
assert is_budget_block(result), f"non-budget error: {result.body[:200]}"
|
||||
pytest.fail("key budget never reset within 150s")
|
||||
57
tests/e2e/budgets/test_model_max_budget_e2e.py
Normal file
57
tests/e2e/budgets/test_model_max_budget_e2e.py
Normal file
|
|
@ -0,0 +1,57 @@
|
|||
"""Live e2e: per-model budgets (`model_max_budget`) isolate by model.
|
||||
|
||||
A key caps one model tiny and leaves another generous. Exhausting the capped
|
||||
model must block *that* model while the other still works - proving the per-model
|
||||
cap is enforced independently, not as a key-wide budget. Closes the
|
||||
model_max_budget gap in BUDGET_TEST_COVERAGE_MATRIX.md.
|
||||
"""
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from budget_client import BudgetClient, is_budget_block, model_budget
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import require_successful_call
|
||||
from lifecycle import ResourceManager
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
CAPPED_MODEL = "claude-haiku-4-5"
|
||||
FREE_MODEL = "gemini-2.5-flash"
|
||||
|
||||
|
||||
def _call(client: BudgetClient, key: str, model: str):
|
||||
result = client.chat(key, model, f"hi {unique_marker()}", max_tokens=16)
|
||||
if not result.ok and not is_budget_block(result):
|
||||
require_successful_call(result)
|
||||
return result
|
||||
|
||||
|
||||
def test_model_max_budget_isolates_per_model(
|
||||
client: BudgetClient, resources: ResourceManager
|
||||
) -> None:
|
||||
key = client.generate_key(
|
||||
model_max_budget={
|
||||
**model_budget(CAPPED_MODEL, 1e-6),
|
||||
**model_budget(FREE_MODEL, 1000.0),
|
||||
}
|
||||
)
|
||||
resources.defer(lambda: client.delete_key(key))
|
||||
|
||||
# Exhaust the capped model.
|
||||
blocked = False
|
||||
deadline = time.monotonic() + 60
|
||||
while time.monotonic() < deadline:
|
||||
if is_budget_block(_call(client, key, CAPPED_MODEL)):
|
||||
blocked = True
|
||||
break
|
||||
time.sleep(1)
|
||||
assert blocked, f"{CAPPED_MODEL} per-model budget never enforced"
|
||||
|
||||
# The other model shares the key but has its own (large) cap -> still works.
|
||||
other = _call(client, key, FREE_MODEL)
|
||||
assert not is_budget_block(other), (
|
||||
f"{FREE_MODEL} was blocked by {CAPPED_MODEL}'s budget; per-model caps not isolated"
|
||||
)
|
||||
require_successful_call(other)
|
||||
70
tests/e2e/budgets/test_multi_window_budget_e2e.py
Normal file
70
tests/e2e/budgets/test_multi_window_budget_e2e.py
Normal file
|
|
@ -0,0 +1,70 @@
|
|||
"""Live e2e: multi-window budgets (budget_limits) enforce AND reset per window.
|
||||
|
||||
Short windows make the time limit reachable inside a test: a tight 30s window and
|
||||
a roomy 1m window. The 30s window blocks once its tiny cap is exceeded, then - once
|
||||
its 30s elapses and the reset job runs (rescheduled fast via
|
||||
PROXY_BUDGET_RESCHEDULER_* in docker-compose) - the window resets and calls flow
|
||||
again. Closes the multi-window gap (enforcement + per-window reset) in
|
||||
BUDGET_TEST_COVERAGE_MATRIX.md, which the unit suite covered but no live test did.
|
||||
"""
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from budget_client import BudgetClient, is_budget_block
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import require_successful_call
|
||||
from lifecycle import ResourceManager
|
||||
from models import BudgetWindow
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
WINDOW_SECONDS = 30 # the tight window; calls succeed again only after it elapses
|
||||
|
||||
|
||||
def _call(client: BudgetClient, key: str):
|
||||
return client.chat(
|
||||
key, "claude-haiku-4-5", f"window {unique_marker()}", max_tokens=16
|
||||
)
|
||||
|
||||
|
||||
def test_short_window_blocks_then_resets(
|
||||
client: BudgetClient, resources: ResourceManager
|
||||
) -> None:
|
||||
key = client.generate_key(
|
||||
budget_limits=[
|
||||
BudgetWindow(budget_duration=f"{WINDOW_SECONDS}s", max_budget=3e-6),
|
||||
BudgetWindow(budget_duration="1m", max_budget=1.0), # roomy: never blocks
|
||||
]
|
||||
)
|
||||
resources.defer(lambda: client.delete_key(key))
|
||||
|
||||
# 1. exhaust the tight window -> litellm returns budget_exceeded
|
||||
start = time.monotonic()
|
||||
blocked = False
|
||||
for _ in range(20):
|
||||
result = _call(client, key)
|
||||
if is_budget_block(result):
|
||||
blocked = True
|
||||
break
|
||||
require_successful_call(result)
|
||||
time.sleep(2)
|
||||
assert blocked, f"{WINDOW_SECONDS}s window never enforced"
|
||||
|
||||
# 2. the window resets at the next wall-clock-aligned boundary (up to a window
|
||||
# after start), then the reset job (~15-20s rescheduler) zeroes the spend.
|
||||
# Allow generous headroom for that alignment + rescheduler latency; a stuck
|
||||
# rescheduler is caught by the wait-loop timeout, not this elapsed bound.
|
||||
deadline = time.monotonic() + 150
|
||||
while time.monotonic() < deadline:
|
||||
time.sleep(5)
|
||||
result = _call(client, key)
|
||||
if result.ok:
|
||||
elapsed = time.monotonic() - start
|
||||
assert elapsed < WINDOW_SECONDS + 90, (
|
||||
f"reset took {elapsed:.0f}s - too long for a {WINDOW_SECONDS}s window"
|
||||
)
|
||||
return
|
||||
assert is_budget_block(result), f"non-budget error during reset wait: {result.body[:200]}"
|
||||
pytest.fail(f"{WINDOW_SECONDS}s window never reset within 150s")
|
||||
35
tests/e2e/budgets/test_soft_budget_e2e.py
Normal file
35
tests/e2e/budgets/test_soft_budget_e2e.py
Normal file
|
|
@ -0,0 +1,35 @@
|
|||
"""Live e2e: soft_budget alerts but does NOT block.
|
||||
|
||||
A key with a tiny `soft_budget` well under a large `max_budget`: spend crosses the
|
||||
soft threshold within a couple calls, but requests keep succeeding (soft budget is
|
||||
advisory). Closes the soft_budget gap in BUDGET_TEST_COVERAGE_MATRIX.md. The alert
|
||||
side-effect (Slack/email) is not observable from the proxy API, so we assert the
|
||||
load-bearing behavior: soft != block.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from budget_client import BudgetClient, is_budget_block
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import require_successful_call
|
||||
from lifecycle import ResourceManager
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
|
||||
def test_soft_budget_does_not_block(
|
||||
client: BudgetClient, resources: ResourceManager
|
||||
) -> None:
|
||||
# soft far below max: spend crosses soft immediately, stays under max.
|
||||
key = client.generate_key(max_budget=1000.0, soft_budget=1e-9)
|
||||
resources.defer(lambda: client.delete_key(key))
|
||||
|
||||
for _ in range(3):
|
||||
result = client.chat(
|
||||
key, "claude-haiku-4-5", f"hi {unique_marker()}", max_tokens=16
|
||||
)
|
||||
assert not is_budget_block(result), (
|
||||
"soft_budget blocked a request; it must alert only, not block "
|
||||
f"(body={result.body[:200]})"
|
||||
)
|
||||
require_successful_call(result) # any other non-2xx (e.g. provider down) is a hard fail
|
||||
144
tests/e2e/budgets/test_spend_counter_reseed_e2e.py
Normal file
144
tests/e2e/budgets/test_spend_counter_reseed_e2e.py
Normal file
|
|
@ -0,0 +1,144 @@
|
|||
"""Live e2e: concurrent cold-counter reseeds keep the spend counter equal to DB spend (#26829).
|
||||
|
||||
Regression for the cross-pod spend-counter multiplication. Real requests build a key's
|
||||
DB spend through the spend writer; the Redis spend counter then expires (the e2e proxy
|
||||
sets a short default_redis_ttl) and goes cold. The proxy runs several workers sharing one
|
||||
Redis, so a concurrent burst makes more than one worker reseed the same cold counter at
|
||||
once. The fix seeds with SET NX - one worker initializes the counter at the DB spend and
|
||||
the rest read it back - so the counter still equals the DB spend (plus the burst's own
|
||||
small cost). The pre-#26829 additive reseed stacked the DB spend once per worker, leaving
|
||||
the counter at ~N x the real spend.
|
||||
|
||||
The test reads the shared counter straight from Redis and asserts it equals the DB spend,
|
||||
not a multiple. It also asserts the counter actually went cold before the burst, so a proxy
|
||||
that never expires the counter (no short TTL) fails loudly instead of passing vacuously.
|
||||
Skipped when the e2e Redis is not reachable.
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
import os
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from threading import Barrier
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import pytest
|
||||
|
||||
from budget_client import BudgetClient
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import StreamingResponse
|
||||
from lifecycle import ResourceManager
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import redis
|
||||
from redis.cluster import RedisCluster
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
MODEL = "claude-haiku-4-5"
|
||||
ACCUMULATE_CALLS = 24
|
||||
BURST = 6
|
||||
# proxy_batch_write_at (60s) flushes the spend to the DB and default_redis_ttl (20s)
|
||||
# expires the counter; this waits out both.
|
||||
COLD_WAIT_SECONDS = 80
|
||||
|
||||
|
||||
def _redis() -> "redis.Redis[str] | RedisCluster[str]":
|
||||
"""The proxy's Redis. The deployed runner sets REDIS_HOST to the serverless
|
||||
ElastiCache, which is always TLS + cluster-mode; without it, fall back to a
|
||||
local standalone redis for docker-compose runs."""
|
||||
import redis
|
||||
|
||||
host = os.getenv("REDIS_HOST")
|
||||
if not host:
|
||||
return redis.Redis(host="localhost", port=6380, decode_responses=True, socket_connect_timeout=2)
|
||||
|
||||
from redis.cluster import RedisCluster
|
||||
|
||||
return RedisCluster(
|
||||
host=host,
|
||||
port=int(os.getenv("REDIS_PORT", "6379")),
|
||||
ssl=True,
|
||||
decode_responses=True,
|
||||
socket_connect_timeout=2,
|
||||
)
|
||||
|
||||
|
||||
def _spend_counter(rds: "redis.Redis[str] | RedisCluster[str]", key: str) -> float | None:
|
||||
"""The shared spend counter for `key`, or None if it is cold. A cluster client
|
||||
can't run a keyspace SCAN that spans shards, so read the key directly - the stage
|
||||
gateway sets no cache namespace, so the key is the bare ``spend:key:{sha256(key)}``.
|
||||
A standalone client matches by suffix, so the local cache namespace (litellm.caching)
|
||||
need not be hard-coded here."""
|
||||
from redis.cluster import RedisCluster
|
||||
|
||||
digest = hashlib.sha256(key.encode()).hexdigest()
|
||||
suffix = f"spend:key:{digest}"
|
||||
if isinstance(rds, RedisCluster):
|
||||
raw = rds.get(suffix)
|
||||
return float(raw) if raw is not None else None
|
||||
|
||||
matches = list(rds.scan_iter(match=f"*{suffix}"))
|
||||
if not matches:
|
||||
return None
|
||||
raw = rds.get(matches[0])
|
||||
return float(raw) if raw is not None else None
|
||||
|
||||
|
||||
def _chat(client: BudgetClient, key: str) -> StreamingResponse:
|
||||
return client.chat(key, MODEL, f"reseed {unique_marker()}", max_tokens=16)
|
||||
|
||||
|
||||
def _accumulate(client: BudgetClient, key: str, count: int) -> None:
|
||||
def one(_: int) -> StreamingResponse:
|
||||
return _chat(client, key)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=8) as pool:
|
||||
list(pool.map(one, range(count)))
|
||||
|
||||
|
||||
def _burst(client: BudgetClient, key: str, count: int) -> None:
|
||||
"""Fire `count` requests that start together, so multiple workers reseed the cold
|
||||
counter concurrently rather than one warming it before the others arrive."""
|
||||
barrier = Barrier(count)
|
||||
|
||||
def one(_: int) -> StreamingResponse:
|
||||
barrier.wait()
|
||||
return _chat(client, key)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=count) as pool:
|
||||
list(pool.map(one, range(count)))
|
||||
|
||||
|
||||
def test_cold_counter_reseed_keeps_counter_equal_to_db_spend(
|
||||
client: BudgetClient, resources: ResourceManager
|
||||
) -> None:
|
||||
try:
|
||||
rds = _redis()
|
||||
rds.ping()
|
||||
except Exception as exc: # noqa: BLE001 - any connect failure means skip
|
||||
pytest.skip(f"e2e redis not reachable (set REDIS_HOST/REDIS_PORT): {exc}")
|
||||
|
||||
key = client.generate_key(max_budget=1.0, models=[MODEL])
|
||||
resources.defer(lambda: client.delete_key(key))
|
||||
|
||||
_accumulate(client, key, ACCUMULATE_CALLS)
|
||||
time.sleep(COLD_WAIT_SECONDS)
|
||||
|
||||
assert _spend_counter(rds, key) is None, (
|
||||
"the spend counter never went cold; default_redis_ttl must be short enough for it "
|
||||
"to expire, otherwise the burst reads a warm counter and the reseed is never exercised"
|
||||
)
|
||||
db_spend = client.gateway.key_info(key).spend or 0.0
|
||||
assert db_spend > 0, f"no DB spend accumulated from real calls: {db_spend}"
|
||||
|
||||
_burst(client, key, BURST)
|
||||
time.sleep(3)
|
||||
|
||||
counter = _spend_counter(rds, key)
|
||||
assert counter is not None, "the burst did not reseed the cold counter"
|
||||
assert db_spend * 0.95 <= counter < db_spend * 1.7, (
|
||||
f"redis spend counter {counter} does not equal DB spend {db_spend} (expected ~equal "
|
||||
f"plus the burst's small cost); a near-multiple means the cold-counter reseed stacked "
|
||||
f"the DB spend once per worker instead of seeding it once (#26829)"
|
||||
)
|
||||
59
tests/e2e/budgets/test_tag_budget_e2e.py
Normal file
59
tests/e2e/budgets/test_tag_budget_e2e.py
Normal file
|
|
@ -0,0 +1,59 @@
|
|||
"""Live e2e: proxy-level tag budgets block tagged requests.
|
||||
|
||||
A tag with a tiny budget: requests carrying that tag get blocked once the tag's
|
||||
spend is exceeded, while a request with a different tag (no budget) still works.
|
||||
Closes the proxy-level tag-budget gap in BUDGET_TEST_COVERAGE_MATRIX.md (today
|
||||
only router-level tag budgets are tested).
|
||||
"""
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from budget_client import BudgetClient, is_budget_block
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import require_successful_call
|
||||
from lifecycle import ResourceManager
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
TINY_BUDGET = 1e-6
|
||||
|
||||
|
||||
def _tagged_call(client: BudgetClient, key: str, tag: str):
|
||||
result = client.chat(
|
||||
key,
|
||||
"claude-haiku-4-5",
|
||||
f"hi {unique_marker()}",
|
||||
tags=[tag],
|
||||
max_tokens=16,
|
||||
)
|
||||
if not result.ok and not is_budget_block(result):
|
||||
require_successful_call(result)
|
||||
return result
|
||||
|
||||
|
||||
def test_tag_budget_blocks_tagged_requests(
|
||||
client: BudgetClient, scoped_key: str, resources: ResourceManager
|
||||
) -> None:
|
||||
budgeted_tag = f"e2e-budget-tag-{unique_marker()}"
|
||||
client.create_tag(budgeted_tag, max_budget=TINY_BUDGET)
|
||||
resources.defer(lambda: client.delete_tag(budgeted_tag))
|
||||
|
||||
# Requests under the budgeted tag get blocked once its spend is exceeded.
|
||||
blocked = False
|
||||
deadline = time.monotonic() + 60
|
||||
while time.monotonic() < deadline:
|
||||
if is_budget_block(_tagged_call(client, scoped_key, budgeted_tag)):
|
||||
blocked = True
|
||||
break
|
||||
time.sleep(1)
|
||||
assert blocked, f"tag budget for {budgeted_tag!r} never enforced"
|
||||
|
||||
# A request with an unbudgeted tag on the same key is unaffected.
|
||||
free_tag = f"e2e-free-tag-{unique_marker()}"
|
||||
other = _tagged_call(client, scoped_key, free_tag)
|
||||
assert not is_budget_block(other), (
|
||||
f"unbudgeted tag {free_tag!r} was blocked by {budgeted_tag!r}'s budget"
|
||||
)
|
||||
require_successful_call(other)
|
||||
107
tests/e2e/budgets/test_team_member_budget_e2e.py
Normal file
107
tests/e2e/budgets/test_team_member_budget_e2e.py
Normal file
|
|
@ -0,0 +1,107 @@
|
|||
"""Live e2e: a team member's per-team budget attributes spend and enforces a cap.
|
||||
|
||||
The team carries a large budget while the one enrolled member is capped at a tiny
|
||||
per-team budget, so any block is member-level, not team-level. Two scenarios share
|
||||
that single member:
|
||||
- attribution: the member's calls land in the spend logs tagged with both the team_id
|
||||
and the member's user_id, so per-member spend can be billed back
|
||||
- enforcement: once the member's spend passes the per-team budget, calls are blocked
|
||||
with budget_exceeded while the team's own budget is nowhere near exhausted
|
||||
|
||||
Per-member budgets enforce off batch-written spend (~60s), so a quick burst all goes
|
||||
through; the block only lands once that spend flushes.
|
||||
"""
|
||||
|
||||
import time
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass
|
||||
|
||||
import pytest
|
||||
|
||||
from budget_client import BudgetClient, is_budget_block
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import Success, require_successful_call
|
||||
from lifecycle import ResourceManager
|
||||
from models import ChatBody, ChatMessage
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
MODEL = "claude-haiku-4-5"
|
||||
TEAM_BUDGET = 100.0
|
||||
MEMBER_BUDGET = 3e-6
|
||||
BURST = 6
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Member:
|
||||
team_id: str
|
||||
user_id: str
|
||||
key: str
|
||||
|
||||
|
||||
@pytest.fixture(scope="class")
|
||||
def member(client: BudgetClient) -> Iterator[_Member]:
|
||||
"""A team with a large budget plus one member capped at a tiny per-team budget,
|
||||
and that member's key. Shared across the class; torn down when it finishes.
|
||||
Cleanups register progressively and run LIFO best-effort through ResourceManager,
|
||||
so a partial-setup failure still releases what came before and one failed delete
|
||||
never strands the rest on the shared proxy."""
|
||||
resources = ResourceManager(client=client.gateway)
|
||||
try:
|
||||
marker = unique_marker()
|
||||
team_id = client.create_team(alias=f"e2e-team-member-{marker}", max_budget=TEAM_BUDGET)
|
||||
resources.defer(lambda: client.delete_team(team_id))
|
||||
user_id = client.create_user(max_budget=TEAM_BUDGET)
|
||||
resources.defer(lambda: client.delete_user(user_id))
|
||||
client.add_team_member(team_id, user_id, max_budget_in_team=MEMBER_BUDGET)
|
||||
key = client.generate_key(team_id=team_id, user_id=user_id)
|
||||
resources.defer(lambda: client.delete_key(key))
|
||||
yield _Member(team_id=team_id, user_id=user_id, key=key)
|
||||
finally:
|
||||
resources.teardown()
|
||||
|
||||
|
||||
def _send(client: BudgetClient, key: str) -> str | None:
|
||||
"""One member call; its response id (== the spend-log request_id) if it went
|
||||
through, else None."""
|
||||
match client.gateway.chat(
|
||||
key,
|
||||
ChatBody(
|
||||
model=MODEL,
|
||||
messages=[ChatMessage(role="user", content=f"hi {unique_marker()}")],
|
||||
max_tokens=16,
|
||||
),
|
||||
):
|
||||
case Success(data=response):
|
||||
return response.id
|
||||
case _:
|
||||
return None
|
||||
|
||||
|
||||
class TestTeamMemberBudget:
|
||||
def test_member_spend_attributed_to_team_and_user(self, client: BudgetClient, member: _Member) -> None:
|
||||
sent = frozenset(rid for rid in (_send(client, member.key) for _ in range(BURST)) if rid)
|
||||
assert sent, "no member call went through; cannot check attribution"
|
||||
|
||||
rows = client.gateway.poll_logs_for_key(
|
||||
member.key, predicate=lambda rs: bool(sent & {r.request_id for r in rs})
|
||||
)
|
||||
logged = [row for row in rows if row.request_id in sent]
|
||||
assert logged, f"none of the member's {len(sent)} calls reached the spend logs"
|
||||
|
||||
for row in logged:
|
||||
assert row.team_id == member.team_id, (
|
||||
f"call {row.request_id} logged under team {row.team_id}, not the member's team {member.team_id}"
|
||||
)
|
||||
assert row.user == member.user_id, (
|
||||
f"call {row.request_id} logged under user {row.user}, not member {member.user_id}"
|
||||
)
|
||||
|
||||
def test_member_spend_over_budget_is_blocked(self, client: BudgetClient, member: _Member) -> None:
|
||||
for _ in range(40):
|
||||
result = client.chat(member.key, MODEL, f"spend {unique_marker()}", max_tokens=16)
|
||||
if is_budget_block(result):
|
||||
return
|
||||
require_successful_call(result)
|
||||
time.sleep(2)
|
||||
pytest.fail("per-member budget never enforced within the call budget")
|
||||
47
tests/e2e/budgets/test_team_member_budget_reset_e2e.py
Normal file
47
tests/e2e/budgets/test_team_member_budget_reset_e2e.py
Normal file
|
|
@ -0,0 +1,47 @@
|
|||
import time
|
||||
from datetime import datetime
|
||||
|
||||
import pytest
|
||||
|
||||
from budget_client import BudgetClient
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import require_successful_call
|
||||
from lifecycle import ResourceManager
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
MEMBER_BUDGET = 1.0 # default member budget is $50, we're testing with a smaller value
|
||||
|
||||
def _as_datetime(value: str) -> datetime:
|
||||
return datetime.fromisoformat(value.replace("Z", "+00:00"))
|
||||
|
||||
|
||||
def test_team_member_budget_reset_keeps_advancing(client: BudgetClient, resources: ResourceManager) -> None:
|
||||
team_id = client.create_team(alias=f"e2e-member-reset-{unique_marker()}", max_budget=100.0)
|
||||
resources.defer(lambda: client.delete_team(team_id))
|
||||
user_id = client.create_user(max_budget=100.0)
|
||||
resources.defer(lambda: client.delete_user(user_id))
|
||||
|
||||
# add the member, then update them onto a short per-team budget window
|
||||
client.add_team_member(team_id, user_id, max_budget_in_team=MEMBER_BUDGET)
|
||||
client.update_team_member(team_id, user_id, max_budget_in_team=MEMBER_BUDGET, budget_duration="30s")
|
||||
|
||||
scheduled = client.member_budget_reset_at(team_id, user_id)
|
||||
assert scheduled, "updating the member with a budget_duration set no budget_reset_at"
|
||||
first_reset = _as_datetime(scheduled)
|
||||
|
||||
# the member can spend within the team while the window is live
|
||||
key = client.generate_key(team_id=team_id, user_id=user_id)
|
||||
resources.defer(lambda: client.delete_key(key))
|
||||
require_successful_call(client.chat(key, "claude-haiku-4-5", f"reset {unique_marker()}", max_tokens=16))
|
||||
|
||||
# once the window elapses the reset job must move budget_reset_at forward; a job
|
||||
# that skips the member's budget row (the #25109 regression) leaves it pinned at
|
||||
# first_reset forever
|
||||
deadline = time.monotonic() + 150
|
||||
while time.monotonic() < deadline:
|
||||
time.sleep(5)
|
||||
current = client.member_budget_reset_at(team_id, user_id)
|
||||
if current and _as_datetime(current) > first_reset:
|
||||
return
|
||||
pytest.fail(f"member budget_reset_at never advanced past {first_reset.isoformat()} in 150s")
|
||||
79
tests/e2e/budgets/test_team_multi_window_budget_e2e.py
Normal file
79
tests/e2e/budgets/test_team_multi_window_budget_e2e.py
Normal file
|
|
@ -0,0 +1,79 @@
|
|||
"""Live e2e: a team's multi-window budgets (budget_limits) enforce AND reset per window.
|
||||
|
||||
The team analog of test_multi_window_budget_e2e.py (which covers keys). A team is
|
||||
created with a tight 30s window and a roomy 1m window; a key on that team blocks once
|
||||
the tight window's cap is exceeded, then - once the 30s elapses and the reset job runs
|
||||
(rescheduled fast via PROXY_BUDGET_RESCHEDULER_* in docker-compose) - the window resets
|
||||
and calls flow again. This exercises the reset_budget_windows TEAM branch (raw SQL over
|
||||
LiteLLM_TeamTable.budget_limits, the literal #25109 path), which had no live coverage.
|
||||
|
||||
Fails at team creation today: /team/new writes the raw window list straight to the
|
||||
Json? column, where Prisma rejects it (500), unlike the key path and /team/update which
|
||||
json.dumps it first. Marked xfail(strict=True) so the suite stays green while the bug
|
||||
persists and flips to a failure the moment the write is fixed and the marker should be
|
||||
removed.
|
||||
"""
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from budget_client import BudgetClient, is_budget_block
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import require_successful_call
|
||||
from lifecycle import ResourceManager
|
||||
from models import BudgetWindow
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
WINDOW_SECONDS = 30
|
||||
|
||||
|
||||
def _call(client: BudgetClient, key: str):
|
||||
return client.chat(key, "claude-haiku-4-5", f"team-window {unique_marker()}", max_tokens=16)
|
||||
|
||||
|
||||
@pytest.mark.xfail(
|
||||
strict=True,
|
||||
reason="known proxy bug: /team/new writes budget_limits straight to the Json? "
|
||||
"column and Prisma rejects it (500), unlike the key path and /team/update which "
|
||||
"json.dumps first; remove this marker once that write is fixed",
|
||||
)
|
||||
def test_team_short_window_blocks_then_resets(client: BudgetClient, resources: ResourceManager) -> None:
|
||||
team_id = client.create_team(
|
||||
alias=f"e2e-team-window-{unique_marker()}",
|
||||
budget_limits=[
|
||||
BudgetWindow(budget_duration=f"{WINDOW_SECONDS}s", max_budget=3e-6),
|
||||
BudgetWindow(budget_duration="1m", max_budget=1.0), # roomy: never blocks
|
||||
],
|
||||
)
|
||||
resources.defer(lambda: client.delete_team(team_id))
|
||||
key = client.generate_key(team_id=team_id)
|
||||
resources.defer(lambda: client.delete_key(key))
|
||||
|
||||
# 1. exhaust the tight window -> litellm returns budget_exceeded
|
||||
start = time.monotonic()
|
||||
blocked = False
|
||||
for _ in range(20):
|
||||
result = _call(client, key)
|
||||
if is_budget_block(result):
|
||||
blocked = True
|
||||
break
|
||||
require_successful_call(result)
|
||||
time.sleep(2)
|
||||
assert blocked, f"team {WINDOW_SECONDS}s window never enforced"
|
||||
|
||||
# 2. the window resets at the next wall-clock-aligned boundary (up to a window
|
||||
# after start), then the reset job (~15-20s rescheduler) zeroes the spend.
|
||||
# Allow generous headroom for that alignment + rescheduler latency; a stuck
|
||||
# rescheduler is caught by the wait-loop timeout, not this elapsed bound.
|
||||
deadline = time.monotonic() + 150
|
||||
while time.monotonic() < deadline:
|
||||
time.sleep(5)
|
||||
result = _call(client, key)
|
||||
if result.ok:
|
||||
elapsed = time.monotonic() - start
|
||||
assert elapsed < WINDOW_SECONDS + 90, f"reset took {elapsed:.0f}s - too long for a {WINDOW_SECONDS}s window"
|
||||
return
|
||||
assert is_budget_block(result), f"non-budget error during reset wait: {result.body[:200]}"
|
||||
pytest.fail(f"team {WINDOW_SECONDS}s window never reset within 150s")
|
||||
120
tests/e2e/conftest.py
Normal file
120
tests/e2e/conftest.py
Normal file
|
|
@ -0,0 +1,120 @@
|
|||
"""Shared fixtures for all live e2e suites under tests/e2e/.
|
||||
|
||||
Design rule: skip on environment, fail on behavior. Live tests (marked `e2e`)
|
||||
skip when no proxy answers; once a request reaches the proxy, behavior is
|
||||
asserted. Pure unit coverage of the harness itself carries no `e2e` marker and
|
||||
runs regardless of whether a proxy is up.
|
||||
|
||||
Lifecycle: the `resources` fixture maps the init -> run -> teardown contract
|
||||
(lifecycle.E2ECase) onto pytest - setup is init(), the test body is run(), and
|
||||
teardown deletes every resource the test created on the long-lived proxy.
|
||||
|
||||
Each suite provides its own `client` fixture (a lifecycle.ResourceClient); these
|
||||
shared fixtures build on it.
|
||||
"""
|
||||
|
||||
import functools
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Iterator
|
||||
|
||||
import pytest
|
||||
import requests
|
||||
|
||||
from e2e_config import CONTROL_PLANE_BASE_URL, PROXY_BASE_URL
|
||||
from lifecycle import GatewayProvider, ResourceManager
|
||||
|
||||
|
||||
_E2E_TEST_RAN = pytest.StashKey[bool]()
|
||||
|
||||
|
||||
def pytest_configure(config: pytest.Config) -> None:
|
||||
config.addinivalue_line(
|
||||
"markers",
|
||||
"e2e: live test that requires a running proxy and real provider keys",
|
||||
)
|
||||
|
||||
|
||||
def _liveness_reason(label: str, base_url: str) -> str | None:
|
||||
"""None if `base_url` answers its liveness probe, else a skip reason."""
|
||||
try:
|
||||
resp = requests.get(f"{base_url}/health/liveliness", timeout=5)
|
||||
except requests.RequestException as exc:
|
||||
return f"No live {label} at {base_url}: {exc}"
|
||||
if resp.status_code >= 500:
|
||||
return f"{label} at {base_url} returned {resp.status_code}"
|
||||
return None
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=1)
|
||||
def _proxy_skip_reason() -> str | None:
|
||||
"""Probe the proxy once per session. None if it answers, else a skip reason. In
|
||||
a split deployment the management/admin control plane is a separate service, so
|
||||
require it too (when it differs) - else its tests would fail rather than skip."""
|
||||
reason = _liveness_reason("proxy", PROXY_BASE_URL)
|
||||
if reason is not None:
|
||||
return reason
|
||||
if CONTROL_PLANE_BASE_URL != PROXY_BASE_URL:
|
||||
return _liveness_reason("control plane", CONTROL_PLANE_BASE_URL)
|
||||
return None
|
||||
|
||||
|
||||
def pytest_runtest_setup(item: pytest.Item) -> None:
|
||||
"""Skip `e2e`-marked tests unless a proxy answers its liveness probe. Unmarked
|
||||
tests (unit coverage of the harness) don't touch the proxy, so they run even
|
||||
when none is up."""
|
||||
if item.get_closest_marker("e2e") is None:
|
||||
return
|
||||
reason = _proxy_skip_reason()
|
||||
if reason is not None:
|
||||
pytest.skip(reason)
|
||||
|
||||
|
||||
def pytest_runtest_call(item: pytest.Item) -> None:
|
||||
"""Mark that an e2e test body actually ran (not skipped at setup). Skipped
|
||||
sessions never reach this hook, so the session-finish cleanup can use it as a
|
||||
guard before truncating the spend-log DB. Tests under `tests/e2e/` without the
|
||||
`e2e` marker (pure unit coverage for the harness itself) never hit the proxy,
|
||||
so they must not arm the destructive DB truncate."""
|
||||
if item.get_closest_marker("e2e") is None:
|
||||
return
|
||||
item.session.stash[_E2E_TEST_RAN] = True
|
||||
|
||||
|
||||
def pytest_sessionfinish(session: pytest.Session, exitstatus: int) -> None:
|
||||
"""Once the whole e2e session is done (all suites), truncate the spend logs so
|
||||
the DB doesn't accumulate test rows. Skipped sessions (no live proxy, no test
|
||||
actually executed) leave the DB alone so a `DATABASE_URL` pointing at a shared
|
||||
instance is never wiped without an e2e run. Best-effort: a cleanup failure (no
|
||||
DB reachable) must not fail the run. The spend_tracking dir goes on sys.path
|
||||
only for this import and is removed after, so a broader `pytest tests/` run is
|
||||
not left with a mutated path."""
|
||||
if not session.stash.get(_E2E_TEST_RAN, False):
|
||||
return
|
||||
spend_dir = str(Path(__file__).parent / "spend_tracking")
|
||||
sys.path.insert(0, spend_dir)
|
||||
try:
|
||||
from spend_e2e_client import reset_spend_logs # pyright: ignore
|
||||
|
||||
reset_spend_logs()
|
||||
except Exception as exc: # noqa: BLE001 - cleanup is best-effort
|
||||
print(f"spend-log cleanup skipped: {exc}")
|
||||
finally:
|
||||
if spend_dir in sys.path:
|
||||
sys.path.remove(spend_dir)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def resources(client: GatewayProvider) -> Iterator[ResourceManager]:
|
||||
"""init -> run -> teardown: create a manager, run the test, release resources.
|
||||
Cleanup goes through the shared Gateway, whatever the suite's client adds."""
|
||||
manager = ResourceManager(client=client.gateway)
|
||||
manager.init()
|
||||
yield manager
|
||||
manager.teardown()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def scoped_key(resources: ResourceManager) -> str:
|
||||
"""A fresh all-models key per test, auto-deleted by the resources teardown."""
|
||||
return resources.key()
|
||||
34
tests/e2e/e2e_config.py
Normal file
34
tests/e2e/e2e_config.py
Normal file
|
|
@ -0,0 +1,34 @@
|
|||
"""Generic configuration for live e2e tests against a running LiteLLM proxy.
|
||||
|
||||
Shared by every e2e suite under tests/e2e/. Values come from the
|
||||
environment so the same tests run against localhost or a deployed proxy.
|
||||
"""
|
||||
|
||||
import os
|
||||
import uuid
|
||||
|
||||
PROXY_BASE_URL = os.environ.get("LITELLM_PROXY_URL", "http://localhost:4000").rstrip("/")
|
||||
MASTER_KEY = os.environ.get("LITELLM_MASTER_KEY", "sk-1234")
|
||||
|
||||
# Control-plane (management/admin) base URL. In a split control-plane/data-plane
|
||||
# deployment the LLM data plane (PROXY_BASE_URL: /chat, /embeddings, native
|
||||
# passthrough) and the management API (keys, users, teams, orgs, budgets, spend,
|
||||
# model info, /openapi.json) are served by *different* services. The suite drives
|
||||
# both through one Transport that routes by path (see transport.SplitTransport).
|
||||
# Defaults to PROXY_BASE_URL so a monolithic proxy serving everything on one URL
|
||||
# behaves exactly as before.
|
||||
CONTROL_PLANE_BASE_URL = os.environ.get(
|
||||
"LITELLM_CONTROL_PLANE_URL", PROXY_BASE_URL
|
||||
).rstrip("/")
|
||||
|
||||
# Writes on the proxy are eventually consistent (e.g. spend rows flush on
|
||||
# proxy_batch_write_at, ~60s). Read-backs poll to this deadline, never sleep-once.
|
||||
POLL_TIMEOUT = float(os.environ.get("E2E_POLL_TIMEOUT", "120"))
|
||||
POLL_INTERVAL = float(os.environ.get("E2E_POLL_INTERVAL", "5"))
|
||||
REQUEST_TIMEOUT = float(os.environ.get("E2E_REQUEST_TIMEOUT", "60"))
|
||||
|
||||
|
||||
def unique_marker() -> str:
|
||||
"""A short unique token per call/run, so concurrent runs and the shared
|
||||
response cache never collide on prompts, tags, or customer ids."""
|
||||
return uuid.uuid4().hex[:12]
|
||||
211
tests/e2e/e2e_gateway.py
Normal file
211
tests/e2e/e2e_gateway.py
Normal file
|
|
@ -0,0 +1,211 @@
|
|||
"""Gateway: the shared proxy operations, DI'd into every client (composition).
|
||||
|
||||
A frozen-slots dataclass holding a Transport plus poll config. Clients hold a
|
||||
Gateway and add their own route methods; the lifecycle ResourceManager uses the
|
||||
Gateway's key/customer methods for cleanup. Read-backs are eventually consistent
|
||||
(proxy_batch_write_at ~60s) so they poll to a deadline.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
|
||||
from e2e_http import (
|
||||
NoBody,
|
||||
ProbeResult,
|
||||
Result,
|
||||
StreamingResponse,
|
||||
Success,
|
||||
unwrap,
|
||||
)
|
||||
from models import (
|
||||
ChatBody,
|
||||
ChatResponse,
|
||||
CustomerDeleteBody,
|
||||
EmbedBody,
|
||||
EmbedResponse,
|
||||
KeyDeleteBody,
|
||||
KeyGenerateBody,
|
||||
KeyGenerateResponse,
|
||||
KeyInfo,
|
||||
KeyInfoParams,
|
||||
KeyInfoResponse,
|
||||
ModelInfoEntry,
|
||||
ModelInfoResponse,
|
||||
SpendLogRow,
|
||||
SpendLogs,
|
||||
SpendLogsParams,
|
||||
)
|
||||
from e2e_config import (
|
||||
CONTROL_PLANE_BASE_URL,
|
||||
MASTER_KEY,
|
||||
POLL_INTERVAL,
|
||||
POLL_TIMEOUT,
|
||||
PROXY_BASE_URL,
|
||||
REQUEST_TIMEOUT,
|
||||
)
|
||||
from transport import HttpTransport, SplitTransport, Transport
|
||||
|
||||
RowsPredicate = Callable[[list[SpendLogRow]], bool]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Gateway:
|
||||
transport: Transport
|
||||
poll_timeout: float = 120.0
|
||||
poll_interval: float = 5.0
|
||||
|
||||
# ---- keys / customers (satisfies lifecycle.ResourceClient) ----------
|
||||
|
||||
def generate_key(self, body: KeyGenerateBody) -> str:
|
||||
return unwrap(
|
||||
self.transport.post(
|
||||
"/key/generate",
|
||||
headers=self.transport.master,
|
||||
json=body,
|
||||
response_type=KeyGenerateResponse,
|
||||
)
|
||||
).key
|
||||
|
||||
def delete_key(self, key: str) -> None:
|
||||
_ = self.transport.post(
|
||||
"/key/delete",
|
||||
headers=self.transport.master,
|
||||
json=KeyDeleteBody(keys=[key]),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
||||
def delete_customers(self, user_ids: list[str]) -> None:
|
||||
if not user_ids:
|
||||
return
|
||||
_ = self.transport.post(
|
||||
"/customer/delete",
|
||||
headers=self.transport.master,
|
||||
json=CustomerDeleteBody(user_ids=user_ids),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
||||
def key_info(self, key: str) -> KeyInfo:
|
||||
return unwrap(
|
||||
self.transport.get(
|
||||
"/key/info",
|
||||
headers=self.transport.master,
|
||||
params=KeyInfoParams(key=key),
|
||||
response_type=KeyInfoResponse,
|
||||
)
|
||||
).info
|
||||
|
||||
def model_info(self) -> list[ModelInfoEntry]:
|
||||
"""Every configured deployment with the price the proxy resolved for it
|
||||
(config override merged over cost-map defaults)."""
|
||||
return unwrap(
|
||||
self.transport.get(
|
||||
"/model/info",
|
||||
headers=self.transport.master,
|
||||
params=NoBody(),
|
||||
response_type=ModelInfoResponse,
|
||||
)
|
||||
).data
|
||||
|
||||
# ---- LLM calls ------------------------------------------------------
|
||||
|
||||
def chat(self, key: str, body: ChatBody) -> Result[ChatResponse]:
|
||||
return self.transport.post(
|
||||
"/chat/completions",
|
||||
headers=self.transport.bearer(key),
|
||||
json=body,
|
||||
response_type=ChatResponse,
|
||||
)
|
||||
|
||||
def chat_stream(self, key: str, body: ChatBody) -> StreamingResponse:
|
||||
return self.transport.stream(
|
||||
"/chat/completions", headers=self.transport.bearer(key), json=body
|
||||
)
|
||||
|
||||
def embed(self, key: str, body: EmbedBody) -> Result[EmbedResponse]:
|
||||
return self.transport.post(
|
||||
"/embeddings",
|
||||
headers=self.transport.bearer(key),
|
||||
json=body,
|
||||
response_type=EmbedResponse,
|
||||
)
|
||||
|
||||
# ---- spend read-back ------------------------------------------------
|
||||
|
||||
def spend_logs(self, params: SpendLogsParams) -> list[SpendLogRow]:
|
||||
result = self.transport.get(
|
||||
"/spend/logs",
|
||||
headers=self.transport.master,
|
||||
params=params,
|
||||
response_type=SpendLogs,
|
||||
)
|
||||
match result:
|
||||
case Success(data=logs):
|
||||
return logs.root
|
||||
case _:
|
||||
return []
|
||||
|
||||
def poll_logs_for_key(
|
||||
self, key: str, *, min_rows: int = 1, predicate: RowsPredicate | None = None
|
||||
) -> list[SpendLogRow]:
|
||||
return self._poll(
|
||||
lambda: self.spend_logs(SpendLogsParams(api_key=key)), min_rows, predicate
|
||||
)
|
||||
|
||||
def poll_logs_for_request_id(
|
||||
self,
|
||||
request_id: str,
|
||||
*,
|
||||
min_rows: int = 1,
|
||||
predicate: RowsPredicate | None = None,
|
||||
) -> list[SpendLogRow]:
|
||||
return self._poll(
|
||||
lambda: self.spend_logs(SpendLogsParams(request_id=request_id)),
|
||||
min_rows,
|
||||
predicate,
|
||||
)
|
||||
|
||||
def _poll(
|
||||
self,
|
||||
fetch: Callable[[], list[SpendLogRow]],
|
||||
min_rows: int,
|
||||
predicate: RowsPredicate | None,
|
||||
) -> list[SpendLogRow]:
|
||||
deadline = time.monotonic() + self.poll_timeout
|
||||
rows: list[SpendLogRow] = []
|
||||
while time.monotonic() < deadline:
|
||||
rows = fetch()
|
||||
if len(rows) >= min_rows and (predicate is None or predicate(rows)):
|
||||
return rows
|
||||
time.sleep(self.poll_interval)
|
||||
return rows
|
||||
|
||||
# ---- route probe ----------------------------------------------------
|
||||
|
||||
def probe(self, path: str, *, params: NoBody) -> ProbeResult:
|
||||
return self.transport.probe(path, params=params)
|
||||
|
||||
|
||||
def build_gateway() -> Gateway:
|
||||
"""The Gateway every suite's client is built from: a SplitTransport that routes
|
||||
LLM calls to the data plane (PROXY_BASE_URL) and management/admin calls to the
|
||||
control plane (CONTROL_PLANE_BASE_URL), with the shared poll budget. The two
|
||||
base URLs are the same for a monolithic proxy, so routing is then a no-op."""
|
||||
return Gateway(
|
||||
transport=SplitTransport(
|
||||
data=HttpTransport(
|
||||
base_url=PROXY_BASE_URL,
|
||||
master_key=MASTER_KEY,
|
||||
request_timeout=REQUEST_TIMEOUT,
|
||||
),
|
||||
control=HttpTransport(
|
||||
base_url=CONTROL_PLANE_BASE_URL,
|
||||
master_key=MASTER_KEY,
|
||||
request_timeout=REQUEST_TIMEOUT,
|
||||
),
|
||||
),
|
||||
poll_timeout=POLL_TIMEOUT,
|
||||
poll_interval=POLL_INTERVAL,
|
||||
)
|
||||
306
tests/e2e/e2e_http.py
Normal file
306
tests/e2e/e2e_http.py
Normal file
|
|
@ -0,0 +1,306 @@
|
|||
"""The ONLY module permitted to call ``requests.*``.
|
||||
|
||||
Enforced by tests/code_coverage_tests/check_e2e_no_raw_requests.py. Every request
|
||||
body / query / header / response is a pydantic model; outcomes are a tagged union
|
||||
(``Result[R]``) so callers ``match`` on them instead of catching exceptions.
|
||||
|
||||
Named e2e_http (not http) so it does not shadow the stdlib ``http`` package that
|
||||
requests itself imports.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Generic, Iterator, Literal, NewType, TypeVar, cast
|
||||
|
||||
import pytest
|
||||
import requests
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
URL = NewType("URL", str)
|
||||
|
||||
|
||||
class Headers(BaseModel):
|
||||
"""Base for header models. Subclasses may alias to hyphenated header names
|
||||
(e.g. ``x-litellm-api-key``); serialization uses by_alias."""
|
||||
|
||||
model_config = ConfigDict(populate_by_name=True)
|
||||
|
||||
|
||||
class AuthHeaders(Headers):
|
||||
# litellm accepts either; set whichever the call needs, leave the other None.
|
||||
authorization: str | None = None
|
||||
x_litellm_api_key: str | None = Field(default=None, alias="x-litellm-api-key")
|
||||
|
||||
|
||||
class NoBody(BaseModel):
|
||||
"""Empty body/query for routes that take none."""
|
||||
|
||||
|
||||
# ---------- Result types ----------
|
||||
|
||||
R = TypeVar("R", bound=BaseModel)
|
||||
|
||||
|
||||
class Success(BaseModel, Generic[R]):
|
||||
kind: Literal["success"] = "success"
|
||||
data: R
|
||||
|
||||
|
||||
class NetworkError(BaseModel):
|
||||
kind: Literal["network"] = "network"
|
||||
message: str
|
||||
|
||||
|
||||
class UnauthorizedError(BaseModel):
|
||||
kind: Literal["unauthorized"] = "unauthorized"
|
||||
|
||||
|
||||
class RateLimitedError(BaseModel):
|
||||
kind: Literal["rate_limited"] = "rate_limited"
|
||||
retry_after_seconds: int | None = None
|
||||
# litellm overloads 429 for budget_exceeded too, so keep the body to tell them apart.
|
||||
body: str = ""
|
||||
|
||||
|
||||
class ValidationError(BaseModel):
|
||||
kind: Literal["validation"] = "validation"
|
||||
message: str
|
||||
|
||||
|
||||
class UnknownApiError(BaseModel):
|
||||
kind: Literal["unknown"] = "unknown"
|
||||
status_code: int
|
||||
body: str
|
||||
|
||||
|
||||
type Result[R: BaseModel] = (
|
||||
Success[R]
|
||||
| NetworkError
|
||||
| UnauthorizedError
|
||||
| RateLimitedError
|
||||
| ValidationError
|
||||
| UnknownApiError
|
||||
)
|
||||
|
||||
|
||||
class ProbeResult(BaseModel):
|
||||
"""A route's reachability: status + body, no schema validation. Healthy ==
|
||||
route exists (not 404) and the handler did not crash (not 5xx)."""
|
||||
|
||||
status_code: int
|
||||
body: str
|
||||
|
||||
@property
|
||||
def healthy(self) -> bool:
|
||||
return 200 <= self.status_code < 500 and self.status_code != 404
|
||||
|
||||
|
||||
class StreamingResponse(BaseModel):
|
||||
"""Raw outcome for calls whose body is provider-native or streamed: status, the
|
||||
x-litellm-call-id header (== SpendLogs.request_id), the content-type (which
|
||||
tells streaming `text/event-stream` from non-streaming `application/json`), and
|
||||
the body. Used by passthrough and streaming, where one validated JSON model
|
||||
does not fit."""
|
||||
|
||||
status_code: int
|
||||
call_id: str | None = None # x-litellm-call-id header
|
||||
content_type: str | None = None
|
||||
body: str
|
||||
chunks: int = 0 # streamed events (0 for non-streaming)
|
||||
|
||||
@property
|
||||
def ok(self) -> bool:
|
||||
return 200 <= self.status_code < 300
|
||||
|
||||
@property
|
||||
def is_streaming(self) -> bool:
|
||||
return "text/event-stream" in (self.content_type or "")
|
||||
|
||||
|
||||
def _hdr(resp: requests.Response, name: str) -> str | None:
|
||||
value = resp.headers.get(name)
|
||||
return value if isinstance(value, str) else None
|
||||
|
||||
|
||||
def unwrap[R: BaseModel](result: Result[R]) -> R:
|
||||
match result:
|
||||
case Success(data=data):
|
||||
return data
|
||||
case _:
|
||||
raise AssertionError(result)
|
||||
|
||||
|
||||
def is_ok[R: BaseModel](result: Result[R]) -> bool:
|
||||
match result:
|
||||
case Success():
|
||||
return True
|
||||
case _:
|
||||
return False
|
||||
|
||||
|
||||
def require_successful_call(result: StreamingResponse) -> None:
|
||||
"""A call that should have succeeded but didn't is a hard failure, never a skip:
|
||||
if the proxy can't make a call it's expected to, the test must fail."""
|
||||
if result.ok:
|
||||
return
|
||||
pytest.fail(
|
||||
f"upstream call failed (status {result.status_code}); body={result.body[:300]}"
|
||||
)
|
||||
|
||||
|
||||
def _headers(headers: BaseModel) -> dict[str, str]:
|
||||
dumped: dict[str, object] = headers.model_dump(by_alias=True, exclude_none=True)
|
||||
return {key: str(value) for key, value in dumped.items()}
|
||||
|
||||
|
||||
def _params(params: BaseModel | None) -> dict[str, str]:
|
||||
if params is None:
|
||||
return {}
|
||||
dumped: dict[str, object] = params.model_dump(by_alias=True, exclude_none=True)
|
||||
return {key: str(value) for key, value in dumped.items()}
|
||||
|
||||
|
||||
def _classify[R: BaseModel](
|
||||
resp: requests.Response, response_type: type[R]
|
||||
) -> Result[R]:
|
||||
if resp.status_code == 401:
|
||||
return UnauthorizedError()
|
||||
if resp.status_code == 429:
|
||||
return RateLimitedError(body=resp.text)
|
||||
if not resp.ok:
|
||||
return UnknownApiError(status_code=resp.status_code, body=resp.text)
|
||||
try:
|
||||
return Success(data=response_type.model_validate(resp.json()))
|
||||
except Exception as exc: # noqa: BLE001 - any parse/validation failure is a value
|
||||
return ValidationError(message=str(exc))
|
||||
|
||||
|
||||
def post[R: BaseModel](
|
||||
url: URL,
|
||||
*,
|
||||
headers: BaseModel,
|
||||
json: BaseModel,
|
||||
response_type: type[R],
|
||||
timeout: float = 30.0,
|
||||
) -> Result[R]:
|
||||
try:
|
||||
resp = requests.post(
|
||||
str(url),
|
||||
headers=_headers(headers),
|
||||
json=json.model_dump(by_alias=True, exclude_none=True),
|
||||
timeout=timeout,
|
||||
)
|
||||
except requests.RequestException as exc:
|
||||
return NetworkError(message=str(exc))
|
||||
return _classify(resp, response_type)
|
||||
|
||||
|
||||
def get[R: BaseModel](
|
||||
url: URL,
|
||||
*,
|
||||
headers: BaseModel,
|
||||
params: BaseModel,
|
||||
response_type: type[R],
|
||||
timeout: float = 30.0,
|
||||
) -> Result[R]:
|
||||
try:
|
||||
resp = requests.get(
|
||||
str(url),
|
||||
headers=_headers(headers),
|
||||
params=params.model_dump(by_alias=True, exclude_none=True),
|
||||
timeout=timeout,
|
||||
)
|
||||
except requests.RequestException as exc:
|
||||
return NetworkError(message=str(exc))
|
||||
return _classify(resp, response_type)
|
||||
|
||||
|
||||
def delete[R: BaseModel](
|
||||
url: URL,
|
||||
*,
|
||||
headers: BaseModel,
|
||||
json: BaseModel,
|
||||
response_type: type[R],
|
||||
timeout: float = 30.0,
|
||||
) -> Result[R]:
|
||||
try:
|
||||
resp = requests.delete(
|
||||
str(url),
|
||||
headers=_headers(headers),
|
||||
json=json.model_dump(by_alias=True, exclude_none=True),
|
||||
timeout=timeout,
|
||||
)
|
||||
except requests.RequestException as exc:
|
||||
return NetworkError(message=str(exc))
|
||||
return _classify(resp, response_type)
|
||||
|
||||
|
||||
def probe(
|
||||
url: URL, *, headers: BaseModel, params: BaseModel, timeout: float = 30.0
|
||||
) -> ProbeResult:
|
||||
try:
|
||||
resp = requests.get(
|
||||
str(url),
|
||||
headers=_headers(headers),
|
||||
params=params.model_dump(by_alias=True, exclude_none=True),
|
||||
timeout=timeout,
|
||||
)
|
||||
except requests.RequestException as exc:
|
||||
return ProbeResult(status_code=-1, body=str(exc))
|
||||
return ProbeResult(status_code=resp.status_code, body=resp.text)
|
||||
|
||||
|
||||
def _streaming_outcome(resp: requests.Response, stream: bool) -> StreamingResponse:
|
||||
call_id = _hdr(resp, "x-litellm-call-id")
|
||||
content_type = _hdr(resp, "content-type")
|
||||
if not stream or not (200 <= resp.status_code < 300):
|
||||
return StreamingResponse(
|
||||
status_code=resp.status_code,
|
||||
call_id=call_id,
|
||||
content_type=content_type,
|
||||
body=resp.text,
|
||||
)
|
||||
lines = cast("Iterator[bytes]", resp.iter_lines())
|
||||
chunks = sum(1 for line in lines if line)
|
||||
return StreamingResponse(
|
||||
status_code=resp.status_code,
|
||||
call_id=call_id,
|
||||
content_type=content_type,
|
||||
body="<streamed>",
|
||||
chunks=chunks,
|
||||
)
|
||||
|
||||
|
||||
def send(
|
||||
url: URL,
|
||||
*,
|
||||
headers: BaseModel,
|
||||
json: BaseModel,
|
||||
params: BaseModel | None = None,
|
||||
stream: bool = False,
|
||||
timeout: float = 60.0,
|
||||
) -> StreamingResponse:
|
||||
"""Raw POST returning the unparsed HTTP outcome: status, full body, and the
|
||||
x-litellm-call-id header. For native/passthrough bodies and for calls judged by
|
||||
status rather than a typed JSON model (e.g. a budget block is a non-2xx). With
|
||||
``stream=True`` the SSE body is consumed and its events counted instead."""
|
||||
try:
|
||||
resp = requests.post(
|
||||
str(url),
|
||||
headers=_headers(headers),
|
||||
params=_params(params),
|
||||
json=json.model_dump(by_alias=True, exclude_none=True),
|
||||
stream=stream,
|
||||
timeout=timeout,
|
||||
)
|
||||
except requests.RequestException as exc:
|
||||
return StreamingResponse(status_code=-1, body=str(exc))
|
||||
return _streaming_outcome(resp, stream)
|
||||
|
||||
|
||||
def stream(
|
||||
url: URL, *, headers: BaseModel, json: BaseModel, timeout: float = 60.0
|
||||
) -> StreamingResponse:
|
||||
"""Streaming (SSE) call: consumes the stream counting events, and captures the
|
||||
x-litellm-call-id + content-type headers. Body is elided."""
|
||||
return send(url, headers=headers, json=json, stream=True, timeout=timeout)
|
||||
168
tests/e2e/gateway/litellm-config.yml
Normal file
168
tests/e2e/gateway/litellm-config.yml
Normal file
|
|
@ -0,0 +1,168 @@
|
|||
# This default config file aims to support most popular model providers out of the box
|
||||
|
||||
#In general, the model name used by the client will be the same as the ones from the provider (For example, you will use "anthropic.claude-3-5-sonnet-20240620-v1:0" when you're calling LiteLLM just like you would when calling Amazon Bedrock directly)
|
||||
#In the case where there are model name conflicts, a prefix will be used (For example, the Azure and the openAI model names conflict, so when you are using Azure, you will use "azure/gpt-4o-realtime-preview-2024-10-01")
|
||||
|
||||
#Some model providers require additional user-specific configuration (such as Azure which requires you to specify your own api_base with your resource name, and your api_version).
|
||||
#In this case, the provider is commented out, and you should uncomment it and provide your specific info
|
||||
|
||||
#For more detailed information about each provider, refer to the docs: https://docs.litellm.ai/docs/providers
|
||||
|
||||
#If you are not interested in a particular provider, just remove it from your config.yaml, and redeploy, and it will no longer show up in your LiteLLM deployment
|
||||
|
||||
#If a particular provider is not working, double check your .env file, and make sure you have provided a valid api key for that provider, and then redeploy
|
||||
|
||||
#Full details on guardrails here: https://docs.litellm.ai/docs/proxy/guardrails/bedrock
|
||||
general_settings:
|
||||
store_prompts_in_spend_logs: true
|
||||
master_key: os.environ/LITELLM_MASTER_KEY
|
||||
proxy_batch_write_at: 60
|
||||
database_connection_pool_limit: 10
|
||||
# disable_error_logs: True
|
||||
forward_client_headers_to_llm_api: false
|
||||
maximum_spend_logs_retention_period: "60d" # GSE-13389: Cleanup logs older than 60 days
|
||||
maximum_spend_logs_cleanup_cron: "0 1 * * *" # 01:00 UTC daily = 18:00 PDT
|
||||
database_url: os.environ/DATABASE_URL
|
||||
control_plane_url: os.environ/CONTROL_PLANE_URL
|
||||
alerts: ["email"]
|
||||
proxy_budget_rescheduler_min_time: 15
|
||||
proxy_budget_rescheduler_max_time: 20
|
||||
|
||||
# fallbacks: [{"gpt-4": ["anthropic.claude-3-5-sonnet-20240620-v1:0"]}] #Configure fallbacks for context window exeeded errors (In this example, we will fall back to Claude Sonnet if over 8000 tokens, which is gpt-4's limit)
|
||||
# default_fallbacks: ["anthropic.claude-3-haiku-20240307-v1:0"] #Configure fallbacks for any error for every model (the above fallback configurations override this one)
|
||||
# environment_variables:
|
||||
# STORE_MODEL_IN_DB: 'True'
|
||||
# LITELLM_LOG: "DEBUG"
|
||||
litellm_settings:
|
||||
drop_params: True
|
||||
# Spend counters inherit this as their Redis TTL, so an idle counter goes cold and
|
||||
# the next request reseeds it from the DB; kept short to exercise the cross-pod
|
||||
# reseed path in test_spend_counter_reseed_e2e. Response-cache writes pass their own
|
||||
# ttl and are unaffected.
|
||||
default_redis_ttl: 20
|
||||
request_timeout: 600
|
||||
num_retries: 3
|
||||
json_logs: true
|
||||
store_audit_logs: True
|
||||
cache: true
|
||||
cache_params:
|
||||
type: redis
|
||||
host: redis
|
||||
port: 6379
|
||||
password: os.environ/REDIS_PASSWORD
|
||||
namespace: litellm.caching
|
||||
ttl: 16600
|
||||
# max_budget: 1000000000.0 # (float) sets max budget in dollars across the entire proxy across all API keys. Note, the budget does not apply to the master key. That is the only exception.
|
||||
# budget_duration: 1mo # (str) frequency of budget reset - You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"), months ("1mo").
|
||||
# max_internal_user_budget: 1000000000.0 # (float) sets default budget in dollars for each internal user. (Doesn't apply to Admins. Doesn't apply to Teams. Doesn't apply to master key)
|
||||
# internal_user_budget_duration: "1mo" # (str) frequency of budget reset - You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"), months ("1mo").
|
||||
# success_callback: ["s3_v2"]
|
||||
# failure_callback: ["s3_v2"]
|
||||
# service_callback: ["datadog"]
|
||||
callbacks: ["arize_phoenix", "datadog", "smtp_email", "prometheus", "otel"]
|
||||
require_auth_for_metrics_endpoint: false
|
||||
#type: redis-semantic
|
||||
#similarity_threshold: 0.8 # similarity threshold for semantic cache
|
||||
#redis_semantic_cache_embedding_model: text-embedding-ada-002 # only works with text-embedding-ada-002 for now... https://github.com/BerriAI/litellm/issues/4001
|
||||
|
||||
router_settings:
|
||||
routing_strategy: simple-shuffle
|
||||
num_retries: 3
|
||||
allowed_fails: 5
|
||||
cooldown_time: 30
|
||||
# When gemini deployments are exhausted (provider 429 / auth), cross over to
|
||||
# working models. Exercised by tests/e2e/router/test_rate_limiter.py.
|
||||
fallbacks:
|
||||
- gemini-2.5-flash: ["gpt-5.5", "claude-haiku-4-5"]
|
||||
|
||||
#ttl: Optional[float]
|
||||
#default_in_memory_ttl: Optional[float]
|
||||
#default_in_redis_ttl: Optional[float]
|
||||
|
||||
model_list:
|
||||
- model_name: gpt-5.5
|
||||
litellm_params:
|
||||
model: openai/gpt-5.5
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
|
||||
- model_name: claude-haiku-4-5
|
||||
litellm_params:
|
||||
model: anthropic/claude-haiku-4-5
|
||||
api_key: os.environ/ANTHROPIC_API_KEY
|
||||
|
||||
# Same underlying model via Vertex AI — distinct routing/auth path
|
||||
# # (service-account JSON), so it gets its own model_name.
|
||||
- model_name: gemini-2.5-flash-vertex
|
||||
litellm_params:
|
||||
model: vertex_ai/gemini-2.5-flash
|
||||
vertex_project: os.environ/VERTEXAI_PROJECT
|
||||
vertex_location: us-central1
|
||||
vertex_credentials: os.environ/VERTEXAI_CREDENTIALS
|
||||
|
||||
- model_name: gemini-2.5-flash
|
||||
litellm_params:
|
||||
model: gemini/gemini-2.5-flash
|
||||
api_key: os.environ/GEMINI_API_KEY
|
||||
|
||||
# load balancing to a different deployment, if gemini gets rate limited.
|
||||
- model_name: gemini-2.5-flash
|
||||
litellm_params:
|
||||
model: gemini/gemini-2.5-flash
|
||||
api_key: os.environ/GEMINI_API_KEY
|
||||
|
||||
# Custom per-token pricing exercised by llm_translation/test_custom_pricing_e2e.py.
|
||||
# Rates deliberately exceed canonical gemini-2.5-flash (input 3e-7 / output 2.5e-6)
|
||||
# so an override that is ignored or under-applied reports spend at the base rate
|
||||
# and fails that test. The test reads these same rates back from this file.
|
||||
- model_name: custom-priced-flash
|
||||
litellm_params:
|
||||
model: gemini/gemini-2.5-flash
|
||||
api_key: os.environ/GEMINI_API_KEY
|
||||
input_cost_per_token: 0.00005
|
||||
output_cost_per_token: 0.0001
|
||||
|
||||
# embedding models
|
||||
- model_name: openai-text-embedding-3-small
|
||||
litellm_params:
|
||||
model: openai/text-embedding-3-small
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
|
||||
- model_name: gemini-2-embedding
|
||||
litellm_params:
|
||||
model: gemini/gemini-2-embedding
|
||||
api_key: os.environ/GEMINI_API_KEY
|
||||
|
||||
# realtime models
|
||||
- model_name: openai-realtime
|
||||
litellm_params:
|
||||
model: openai/realtime-2
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
model_info:
|
||||
mode: realtime
|
||||
|
||||
|
||||
mcp_servers:
|
||||
deepwiki_mcp:
|
||||
url: "https://mcp.deepwiki.com/mcp"
|
||||
auth_type: none
|
||||
description: "just a test"
|
||||
|
||||
atlassian:
|
||||
url: "https://mcp.atlassian.com/v1/mcp"
|
||||
auth_type: oauth2
|
||||
authorization_url: https://auth.atlassian.com/authorize
|
||||
|
||||
|
||||
guardrails:
|
||||
- guardrail_name: "presidio-pii"
|
||||
litellm_params:
|
||||
guardrail: presidio
|
||||
mode: pre_call
|
||||
presidio_analyzer_api_base: os.environ/PRESIDIO_ANALYZER_API_BASE
|
||||
presidio_anonymizer_api_base: os.environ/PRESIDIO_ANONYMIZER_API_BASE
|
||||
default_on: false
|
||||
pii_entities_config:
|
||||
EMAIL_ADDRESS: BLOCK
|
||||
CREDIT_CARD: BLOCK
|
||||
US_SSN: BLOCK
|
||||
PHONE_NUMBER: BLOCK
|
||||
117
tests/e2e/lifecycle.py
Normal file
117
tests/e2e/lifecycle.py
Normal file
|
|
@ -0,0 +1,117 @@
|
|||
"""Lifecycle contract and resource cleanup for stateful e2e tests.
|
||||
|
||||
Shared by every e2e suite under tests/e2e/. The proxy under test is
|
||||
long-lived and never reset between tests, so anything a test creates (keys,
|
||||
customers, teams, orgs, users, guardrails, budgets, ...) persists unless
|
||||
explicitly deleted. Every check follows an init -> run -> teardown lifecycle;
|
||||
teardown releases each resource init() created, even when run() raises.
|
||||
|
||||
In pytest terms (see conftest.py): the `resources` fixture's setup is init(),
|
||||
the test body is run(), and the fixture's teardown is teardown().
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable, List, Protocol, runtime_checkable
|
||||
|
||||
from e2e_gateway import Gateway
|
||||
from models import KeyGenerateBody
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class E2ECase(Protocol):
|
||||
"""A stateful e2e check run against a long-lived proxy.
|
||||
|
||||
init() acquires resources, run() exercises behaviour and asserts, teardown()
|
||||
releases everything init() created. teardown() must run even if init() fails
|
||||
partway or run() raises.
|
||||
"""
|
||||
|
||||
def init(self) -> None: ...
|
||||
|
||||
def run(self) -> None: ...
|
||||
|
||||
def teardown(self) -> None: ...
|
||||
|
||||
|
||||
def run_case(case: E2ECase) -> None:
|
||||
"""Drive a case through its lifecycle: init -> run -> teardown.
|
||||
|
||||
teardown always runs - even when init() fails partway or run() raises (or
|
||||
skips) - so resources the case already registered on the long-lived proxy are
|
||||
released. init() is inside the try because cases register cleanups
|
||||
progressively (e.g. create team, then user, then key), and a failure after
|
||||
the first creation must still release what came before.
|
||||
"""
|
||||
try:
|
||||
case.init()
|
||||
case.run()
|
||||
finally:
|
||||
case.teardown()
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class ResourceClient(Protocol):
|
||||
"""Proxy operations the convenience creators use. Resource types without a
|
||||
creator here are handled generically via ResourceManager.defer(). The Gateway
|
||||
satisfies this."""
|
||||
|
||||
def generate_key(self, body: KeyGenerateBody) -> str: ...
|
||||
|
||||
def delete_key(self, key: str) -> None: ...
|
||||
|
||||
def delete_customers(self, user_ids: List[str]) -> None: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class GatewayProvider(Protocol):
|
||||
"""Every suite's client exposes the shared Gateway, which the resources fixture
|
||||
uses for cleanup. The client adds its own route methods on top."""
|
||||
|
||||
@property
|
||||
def gateway(self) -> Gateway: ...
|
||||
|
||||
|
||||
@dataclass
|
||||
class ResourceManager:
|
||||
"""Registry of teardown actions for resources a test creates on the stateful
|
||||
proxy.
|
||||
|
||||
Not limited to any resource type: register a cleanup with ``defer()`` for a
|
||||
key, customer, team, org, user, guardrail, budget, MCP server - anything with
|
||||
a delete. The two most common resources have sugar (``key``, ``customer``);
|
||||
everything else is ``resources.defer(lambda: client.delete_team(team_id))``.
|
||||
|
||||
Cleanups run LIFO (so a resource is removed before whatever it depends on) and
|
||||
best-effort (one failing cleanup never blocks the rest).
|
||||
"""
|
||||
|
||||
client: ResourceClient
|
||||
_cleanups: List[Callable[[], None]] = field(
|
||||
default_factory=list
|
||||
) # mutable-ok: append-only teardown registry
|
||||
|
||||
def init(self) -> None:
|
||||
"""No global setup needed today; present for lifecycle symmetry."""
|
||||
return None
|
||||
|
||||
def defer(self, cleanup: Callable[[], None]) -> None:
|
||||
"""Register a teardown action for any resource the test just created."""
|
||||
self._cleanups.append(cleanup)
|
||||
|
||||
def key(self) -> str:
|
||||
"""Create an all-models virtual key; delete it on teardown."""
|
||||
key = self.client.generate_key(KeyGenerateBody(models=[]))
|
||||
self.defer(lambda: self.client.delete_key(key))
|
||||
return key
|
||||
|
||||
def customer(self, customer_id: str) -> str:
|
||||
"""Track an end-user id (from the `user` param); delete it on teardown."""
|
||||
self.defer(lambda: self.client.delete_customers([customer_id]))
|
||||
return customer_id
|
||||
|
||||
def teardown(self) -> None:
|
||||
for cleanup in reversed(self._cleanups):
|
||||
try:
|
||||
cleanup()
|
||||
except Exception:
|
||||
pass # best-effort: a failed cleanup must not block the rest
|
||||
84
tests/e2e/llm_translation/LLM_TRANSLATION_COVERAGE_MATRIX.md
Normal file
84
tests/e2e/llm_translation/LLM_TRANSLATION_COVERAGE_MATRIX.md
Normal file
|
|
@ -0,0 +1,84 @@
|
|||
# LLM Translation Test Coverage Matrix
|
||||
|
||||
Scope: the proxy's two translation surfaces, end to end against a live proxy.
|
||||
|
||||
1. **Passthrough** - the client speaks the provider's NATIVE API (Gemini
|
||||
`generateContent`, Anthropic `/v1/messages`); the proxy forwards it and still
|
||||
logs a costed `SpendLogs` row (`call_type="pass_through_endpoint"`). Routes:
|
||||
`/gemini`, `/anthropic`, `/vertex_ai`, `/openai`, `/bedrock`, `/cohere`,
|
||||
`/mistral`, `/vllm`.
|
||||
2. **Non-passthrough** - the client speaks OpenAI format
|
||||
(`/chat/completions`, `/embeddings`); litellm translates to/from the provider.
|
||||
|
||||
The two axes that must work in production for each: **passthrough vs
|
||||
non-passthrough** and **streaming vs non-streaming**, with **cost logged** and
|
||||
**tool calls** working in every cell.
|
||||
|
||||
Companion: live suite `test_passthrough_e2e.py` (this directory). The
|
||||
non-passthrough chat/embedding cells are exercised by `../spend_tracking/`.
|
||||
|
||||
Levels: `live` real provider + proxy + SpendLogs row; `unit` mocked.
|
||||
Status: `covered` / `partial` / `gap`.
|
||||
|
||||
---
|
||||
|
||||
## Passthrough endpoints (native provider format)
|
||||
|
||||
| Provider | Non-streaming | Streaming | Tool calls | Cost logged | Status |
|
||||
|----------|---------------|-----------|------------|-------------|--------|
|
||||
| Gemini (`/gemini/v1beta/models/{m}:generateContent` / `:streamGenerateContent`) | live | live | live | live | **covered** |
|
||||
| Anthropic (`/anthropic/v1/messages`) | live | live | live | live | **covered** |
|
||||
| Vertex AI (`/vertex_ai/...`) | - | - | - | - | gap (gcloud auth) |
|
||||
| OpenAI / Bedrock / Cohere / Mistral / VLLM | - | - | - | - | gap |
|
||||
|
||||
Each covered cell asserts: `call_type == "pass_through_endpoint"`, `spend > 0`,
|
||||
`status == "success"`, correct `custom_llm_provider`/`model`, row correlated by the
|
||||
`x-litellm-call-id` header. Gemini non-streaming also pins `request_tags`
|
||||
propagation; streaming pins `chunks > 0` then a costed row; tool tests assert the
|
||||
provider emitted a tool call (`functionCall` / `tool_use`) and it was costed.
|
||||
|
||||
Cost on passthrough is computed in the success handler by transforming the native
|
||||
response to a `ModelResponse` and calling `litellm.completion_cost()`; for
|
||||
streaming, chunks are buffered and costed after the stream ends. This is the path
|
||||
most likely to silently break and the one a mock can't prove works.
|
||||
|
||||
## Non-passthrough endpoints (OpenAI-compatible translation)
|
||||
|
||||
| Modality | Non-streaming | Streaming | Tool calls | Cost logged | Status |
|
||||
|----------|---------------|-----------|------------|-------------|--------|
|
||||
| Chat | live (spend suite) | live (spend suite) | gap | live | partial |
|
||||
| Embeddings | live (spend suite) | n/a | n/a | live | covered |
|
||||
| Responses / image / audio / rerank / realtime | - | - | - | - | gap |
|
||||
|
||||
## This suite's files
|
||||
|
||||
| Test | Cell |
|
||||
|------|------|
|
||||
| `test_gemini_passthrough_nonstreaming_logs_cost` | gemini native, non-stream, cost + tags |
|
||||
| `test_gemini_passthrough_streaming_logs_cost` | gemini native, stream, cost |
|
||||
| `test_gemini_passthrough_tool_call_logs_cost` | gemini native, tool call, cost |
|
||||
| `test_anthropic_passthrough_nonstreaming_logs_cost` | anthropic native, non-stream, cost |
|
||||
| `test_anthropic_passthrough_streaming_logs_cost` | anthropic native, stream, cost |
|
||||
| `test_anthropic_passthrough_tool_call_logs_cost` | anthropic native, tool call, cost |
|
||||
|
||||
## Gaps
|
||||
|
||||
- Vertex / OpenAI / Bedrock / Cohere passthrough (same shape; add once the
|
||||
provider credential is configured; Vertex is closest - route exists, auth stale).
|
||||
- Non-passthrough tool calls over `/chat/completions` end to end with cost.
|
||||
- Image / audio / rerank / responses / realtime translation + cost.
|
||||
- Streaming cost-injection (`include_cost_in_streaming_usage`); passthrough on
|
||||
client disconnect (partial-usage logging).
|
||||
|
||||
## Adding a provider/modality
|
||||
|
||||
Extend `PassthroughClient` with the native call (it inherits keys, cleanup, and
|
||||
SpendLogs polling from `ProxyClient`), then add a test that calls it,
|
||||
`require_successful_call(result)`, and `_costed_row(...)`.
|
||||
|
||||
## Timing
|
||||
|
||||
Passthrough spend is logged asynchronously after the response and lands on the
|
||||
`proxy_batch_write_at` (~60s) cycle, so cost assertions poll
|
||||
`/spend/logs?request_id=<x-litellm-call-id>` to a deadline. Streaming cost is only
|
||||
known after the stream is fully consumed.
|
||||
15
tests/e2e/llm_translation/conftest.py
Normal file
15
tests/e2e/llm_translation/conftest.py
Normal file
|
|
@ -0,0 +1,15 @@
|
|||
"""LLM-translation suite's `client` fixture.
|
||||
|
||||
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
|
||||
live in the parent tests/e2e/conftest.py. PassthroughClient holds the shared
|
||||
Gateway, so the `resources` fixture cleans up keys this suite creates.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from passthrough_client import PassthroughClient, build_client
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def client() -> PassthroughClient:
|
||||
return build_client()
|
||||
163
tests/e2e/llm_translation/passthrough_client.py
Normal file
163
tests/e2e/llm_translation/passthrough_client.py
Normal file
|
|
@ -0,0 +1,163 @@
|
|||
"""Client for LLM-translation e2e tests over the proxy's passthrough endpoints.
|
||||
|
||||
A passthrough request is sent in the PROVIDER's native format (Gemini
|
||||
generateContent, Anthropic /v1/messages) to the proxy, which forwards it to the
|
||||
provider and still logs a SpendLogs row (call_type="pass_through_endpoint"). The
|
||||
litellm virtual key is passed as the provider key; the proxy swaps in the real env
|
||||
credential. SpendLogs.request_id == the x-litellm-call-id response header. The
|
||||
native request models are co-located here because only this suite uses them.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from e2e_gateway import Gateway, build_gateway
|
||||
from e2e_http import Headers, StreamingResponse
|
||||
from models import ChatMessage
|
||||
|
||||
|
||||
class JsonSchemaProperty(BaseModel):
|
||||
type: str
|
||||
|
||||
|
||||
class JsonSchema(BaseModel):
|
||||
type: str
|
||||
properties: dict[str, JsonSchemaProperty]
|
||||
required: list[str]
|
||||
|
||||
|
||||
class GeminiHeaders(Headers):
|
||||
x_goog_api_key: str = Field(serialization_alias="x-goog-api-key")
|
||||
content_type: str = Field(
|
||||
default="application/json", serialization_alias="Content-Type"
|
||||
)
|
||||
tags: str | None = None
|
||||
|
||||
|
||||
class AnthropicHeaders(Headers):
|
||||
x_api_key: str = Field(serialization_alias="x-api-key")
|
||||
anthropic_version: str = Field(
|
||||
default="2023-06-01", serialization_alias="anthropic-version"
|
||||
)
|
||||
content_type: str = Field(
|
||||
default="application/json", serialization_alias="Content-Type"
|
||||
)
|
||||
tags: str | None = None
|
||||
|
||||
|
||||
class AltSseParams(BaseModel):
|
||||
alt: str = "sse"
|
||||
|
||||
|
||||
class GeminiPart(BaseModel):
|
||||
text: str
|
||||
|
||||
|
||||
class GeminiContent(BaseModel):
|
||||
role: str = "user"
|
||||
parts: list[GeminiPart]
|
||||
|
||||
|
||||
class GeminiFunctionDeclaration(BaseModel):
|
||||
name: str
|
||||
description: str
|
||||
parameters: JsonSchema
|
||||
|
||||
|
||||
class GeminiTool(BaseModel):
|
||||
function_declarations: list[GeminiFunctionDeclaration] = Field(
|
||||
serialization_alias="functionDeclarations"
|
||||
)
|
||||
|
||||
|
||||
class GeminiGenerateBody(BaseModel):
|
||||
contents: list[GeminiContent]
|
||||
tools: list[GeminiTool] | None = None
|
||||
|
||||
|
||||
class AnthropicTool(BaseModel):
|
||||
name: str
|
||||
description: str
|
||||
input_schema: JsonSchema
|
||||
|
||||
|
||||
class AnthropicMessageBody(BaseModel):
|
||||
model: str
|
||||
max_tokens: int
|
||||
messages: list[ChatMessage]
|
||||
tools: list[AnthropicTool] | None = None
|
||||
stream: bool = False
|
||||
|
||||
|
||||
def _tags_header(tags: list[str] | None) -> str | None:
|
||||
return ",".join(tags) if tags else None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PassthroughClient:
|
||||
gateway: Gateway
|
||||
|
||||
# ---- Gemini native passthrough (/gemini/v1beta/...) -----------------
|
||||
|
||||
def gemini_generate(
|
||||
self,
|
||||
key: str,
|
||||
model: str,
|
||||
text: str,
|
||||
*,
|
||||
tools: list[GeminiTool] | None = None,
|
||||
tags: list[str] | None = None,
|
||||
) -> StreamingResponse:
|
||||
return self.gateway.transport.send(
|
||||
f"/gemini/v1beta/models/{model}:generateContent",
|
||||
headers=GeminiHeaders(x_goog_api_key=key, tags=_tags_header(tags)),
|
||||
json=GeminiGenerateBody(
|
||||
contents=[GeminiContent(parts=[GeminiPart(text=text)])], tools=tools
|
||||
),
|
||||
)
|
||||
|
||||
def gemini_stream(
|
||||
self, key: str, model: str, text: str, *, tags: list[str] | None = None
|
||||
) -> StreamingResponse:
|
||||
return self.gateway.transport.send(
|
||||
f"/gemini/v1beta/models/{model}:streamGenerateContent",
|
||||
headers=GeminiHeaders(x_goog_api_key=key, tags=_tags_header(tags)),
|
||||
json=GeminiGenerateBody(
|
||||
contents=[GeminiContent(parts=[GeminiPart(text=text)])]
|
||||
),
|
||||
params=AltSseParams(),
|
||||
stream=True,
|
||||
)
|
||||
|
||||
# ---- Anthropic native passthrough (/anthropic/v1/messages) ----------
|
||||
|
||||
def anthropic_message(
|
||||
self,
|
||||
key: str,
|
||||
model: str,
|
||||
text: str,
|
||||
*,
|
||||
max_tokens: int = 64,
|
||||
tools: list[AnthropicTool] | None = None,
|
||||
stream: bool = False,
|
||||
tags: list[str] | None = None,
|
||||
) -> StreamingResponse:
|
||||
return self.gateway.transport.send(
|
||||
"/anthropic/v1/messages",
|
||||
headers=AnthropicHeaders(x_api_key=key, tags=_tags_header(tags)),
|
||||
json=AnthropicMessageBody(
|
||||
model=model,
|
||||
max_tokens=max_tokens,
|
||||
messages=[ChatMessage(role="user", content=text)],
|
||||
tools=tools,
|
||||
stream=stream,
|
||||
),
|
||||
stream=stream,
|
||||
)
|
||||
|
||||
|
||||
def build_client() -> PassthroughClient:
|
||||
return PassthroughClient(gateway=build_gateway())
|
||||
219
tests/e2e/llm_translation/test_custom_pricing_e2e.py
Normal file
219
tests/e2e/llm_translation/test_custom_pricing_e2e.py
Normal file
|
|
@ -0,0 +1,219 @@
|
|||
"""Live e2e: a model's custom per-token pricing is loaded, billed, and isolated.
|
||||
|
||||
The gateway config declares ``custom-priced-flash`` (gemini-2.5-flash underneath)
|
||||
with input/output rates deliberately far above the canonical gemini price, read
|
||||
back here from the same config file. Three behaviors are checked independently:
|
||||
|
||||
- billing: a real call's logged cost breakdown charges input and output tokens at
|
||||
the custom rates, each component checked separately (a base-rate bill lands
|
||||
~100x lower; a swapped input/output rate passes a total-only check but not this)
|
||||
- reporting: /model/info surfaces those rates for the model
|
||||
- isolation: gemini-2.5-flash shares the same underlying gemini/gemini-2.5-flash
|
||||
but sets no override, so it must keep its own price; an override that leaks into
|
||||
the shared cost map misprices it. This fails on a real proxy gap today, so it is
|
||||
marked xfail(strict=True): the suite stays green while the leak persists and
|
||||
flips to a failure the moment isolation is fixed and the marker should be removed.
|
||||
"""
|
||||
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import yaml
|
||||
from pydantic import BaseModel, RootModel
|
||||
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import Success, unwrap
|
||||
from models import ChatBody, ChatMessage, CustomPricing, ModelInfoEntry, SpendLogsParams
|
||||
from passthrough_client import PassthroughClient
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
CUSTOM_MODEL = "custom-priced-flash"
|
||||
BASE_MODEL = "gemini-2.5-flash"
|
||||
CONFIG_PATH = Path(__file__).resolve().parents[1] / "gateway" / "litellm-config.yml"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Rates:
|
||||
input_per_token: float
|
||||
output_per_token: float
|
||||
|
||||
|
||||
class _ConfiguredModel(BaseModel):
|
||||
model_name: str
|
||||
litellm_params: CustomPricing
|
||||
|
||||
|
||||
class _GatewayConfig(BaseModel):
|
||||
model_list: list[_ConfiguredModel]
|
||||
|
||||
|
||||
class _CostBreakdown(BaseModel):
|
||||
input_cost: float | None = None
|
||||
output_cost: float | None = None
|
||||
|
||||
|
||||
class _RowMetadata(BaseModel):
|
||||
cost_breakdown: _CostBreakdown | None = None
|
||||
|
||||
|
||||
class _SpendRow(BaseModel):
|
||||
request_id: str | None = None
|
||||
prompt_tokens: int | None = None
|
||||
completion_tokens: int | None = None
|
||||
metadata: _RowMetadata | None = None
|
||||
|
||||
|
||||
class _SpendRows(RootModel[list[_SpendRow]]):
|
||||
pass
|
||||
|
||||
|
||||
def _approx_equal(actual: float, expected: float) -> bool:
|
||||
"""Within 1% or 1e-9 absolute - spend math, not exact float identity."""
|
||||
return abs(actual - expected) <= max(1e-9, abs(expected) * 1e-2)
|
||||
|
||||
|
||||
def _configured_pricing(model_name: str) -> _Rates:
|
||||
"""The custom rates declared for `model_name` in the gateway config the proxy
|
||||
runs with - the source of truth the billed and reported prices are checked
|
||||
against."""
|
||||
config = _GatewayConfig.model_validate(yaml.safe_load(CONFIG_PATH.read_text()))
|
||||
for entry in config.model_list:
|
||||
if entry.model_name == model_name:
|
||||
pricing = entry.litellm_params
|
||||
assert pricing.input_cost_per_token and pricing.output_cost_per_token, (
|
||||
f"{model_name} declares no custom per-token rates in {CONFIG_PATH.name}"
|
||||
)
|
||||
return _Rates(pricing.input_cost_per_token, pricing.output_cost_per_token)
|
||||
pytest.fail(f"{model_name} not found in {CONFIG_PATH.name}")
|
||||
|
||||
|
||||
def _model_info_entry(
|
||||
entries: list[ModelInfoEntry], model_name: str
|
||||
) -> ModelInfoEntry:
|
||||
for entry in entries:
|
||||
if entry.model_name == model_name:
|
||||
return entry
|
||||
pytest.fail(f"{model_name} absent from /model/info; the override did not load")
|
||||
|
||||
|
||||
def _poll_breakdown_row(
|
||||
client: PassthroughClient, key: str, response_id: str | None
|
||||
) -> _SpendRow:
|
||||
"""Poll /spend/logs until the call's row lands with a cost breakdown (rows
|
||||
flush ~60s behind the call via proxy_batch_write_at)."""
|
||||
deadline = time.monotonic() + client.gateway.poll_timeout
|
||||
while time.monotonic() < deadline:
|
||||
result = client.gateway.transport.get(
|
||||
"/spend/logs",
|
||||
headers=client.gateway.transport.master,
|
||||
params=SpendLogsParams(api_key=key),
|
||||
response_type=_SpendRows,
|
||||
)
|
||||
match result:
|
||||
case Success(data=data):
|
||||
rows = data.root
|
||||
case _:
|
||||
rows = []
|
||||
priced = [
|
||||
row
|
||||
for row in rows
|
||||
if row.metadata
|
||||
and row.metadata.cost_breakdown
|
||||
and row.metadata.cost_breakdown.input_cost is not None
|
||||
]
|
||||
for row in priced:
|
||||
if response_id and row.request_id == response_id:
|
||||
return row
|
||||
if priced and response_id is None:
|
||||
return priced[0]
|
||||
time.sleep(client.gateway.poll_interval)
|
||||
pytest.fail("no spend row with a cost breakdown landed before the deadline")
|
||||
|
||||
|
||||
def test_custom_pricing_is_billed_at_configured_rate(
|
||||
client: PassthroughClient, scoped_key: str
|
||||
) -> None:
|
||||
rates = _configured_pricing(CUSTOM_MODEL)
|
||||
|
||||
chat = unwrap(
|
||||
client.gateway.chat(
|
||||
scoped_key,
|
||||
ChatBody(
|
||||
model=CUSTOM_MODEL,
|
||||
messages=[
|
||||
ChatMessage(
|
||||
role="user", content=f"reply with one word {unique_marker()}"
|
||||
)
|
||||
],
|
||||
max_tokens=16,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
row = _poll_breakdown_row(client, scoped_key, chat.id)
|
||||
assert row.metadata and row.metadata.cost_breakdown # guaranteed by the poll
|
||||
breakdown = row.metadata.cost_breakdown
|
||||
|
||||
prompt = row.prompt_tokens or 0
|
||||
completion = row.completion_tokens or 0
|
||||
assert prompt > 0 and completion > 0, f"call tokens not logged on the row: {row}"
|
||||
|
||||
input_cost = breakdown.input_cost
|
||||
output_cost = breakdown.output_cost
|
||||
assert input_cost is not None and output_cost is not None, (
|
||||
f"row cost breakdown missing input/output cost: {breakdown}"
|
||||
)
|
||||
assert _approx_equal(input_cost, prompt * rates.input_per_token), (
|
||||
f"input_cost {input_cost} != {prompt} tokens * {rates.input_per_token} "
|
||||
f"= {prompt * rates.input_per_token}"
|
||||
)
|
||||
assert _approx_equal(output_cost, completion * rates.output_per_token), (
|
||||
f"output_cost {output_cost} != {completion} tokens * {rates.output_per_token} "
|
||||
f"= {completion * rates.output_per_token}"
|
||||
)
|
||||
|
||||
|
||||
def test_model_info_reports_custom_pricing(client: PassthroughClient) -> None:
|
||||
rates = _configured_pricing(CUSTOM_MODEL)
|
||||
entry = _model_info_entry(client.gateway.model_info(), CUSTOM_MODEL)
|
||||
|
||||
assert entry.litellm_params.input_cost_per_token == rates.input_per_token, (
|
||||
f"/model/info litellm_params input rate "
|
||||
f"{entry.litellm_params.input_cost_per_token} != configured "
|
||||
f"{rates.input_per_token}"
|
||||
)
|
||||
assert entry.litellm_params.output_cost_per_token == rates.output_per_token, (
|
||||
f"/model/info litellm_params output rate "
|
||||
f"{entry.litellm_params.output_cost_per_token} != configured "
|
||||
f"{rates.output_per_token}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.xfail(
|
||||
strict=True,
|
||||
reason="known proxy bug: a deployment's custom per-token pricing leaks into the "
|
||||
"shared cost map for sibling deployments of the same underlying model; remove "
|
||||
"this marker once isolation is fixed",
|
||||
)
|
||||
def test_custom_pricing_is_isolated_from_sibling_deployment(
|
||||
client: PassthroughClient,
|
||||
) -> None:
|
||||
entries = {entry.model_name: entry for entry in client.gateway.model_info()}
|
||||
custom = entries.get(CUSTOM_MODEL)
|
||||
base = entries.get(BASE_MODEL)
|
||||
assert custom is not None, f"{CUSTOM_MODEL} absent from /model/info"
|
||||
assert base is not None, f"{BASE_MODEL} absent from /model/info"
|
||||
|
||||
# custom-priced-flash overrides pricing; gemini-2.5-flash shares the same
|
||||
# underlying gemini/gemini-2.5-flash but sets no override, so it must keep its
|
||||
# own price. Equal rates mean the override leaked into the shared cost map.
|
||||
assert (
|
||||
base.model_info.input_cost_per_token != custom.model_info.input_cost_per_token
|
||||
), (
|
||||
f"{BASE_MODEL} input rate {base.model_info.input_cost_per_token} matches "
|
||||
f"{CUSTOM_MODEL}'s override {custom.model_info.input_cost_per_token}; "
|
||||
f"per-deployment custom pricing is not isolated"
|
||||
)
|
||||
159
tests/e2e/llm_translation/test_passthrough_e2e.py
Normal file
159
tests/e2e/llm_translation/test_passthrough_e2e.py
Normal file
|
|
@ -0,0 +1,159 @@
|
|||
"""Live e2e for LLM-translation passthrough endpoints.
|
||||
|
||||
Each test sends a NATIVE provider request through the proxy's passthrough route
|
||||
and verifies the proxy still logged a costed SpendLogs row
|
||||
(call_type="pass_through_endpoint"), correlated by the x-litellm-call-id header.
|
||||
|
||||
Covered: gemini ("gemini-2.5-flash") + anthropic ("claude-haiku-4-5"), streaming +
|
||||
non-streaming, plus native tool calls. See LLM_TRANSLATION_COVERAGE_MATRIX.md.
|
||||
|
||||
A passthrough call returning non-2xx fails hard (never a skip); once it returns
|
||||
2xx, a missing or zero-cost SpendLogs row fails too.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import StreamingResponse, require_successful_call
|
||||
from models import SpendLogRow
|
||||
from passthrough_client import (
|
||||
AnthropicTool,
|
||||
GeminiFunctionDeclaration,
|
||||
GeminiTool,
|
||||
JsonSchema,
|
||||
JsonSchemaProperty,
|
||||
PassthroughClient,
|
||||
)
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
|
||||
def _fetch_cost_breakdown(client: PassthroughClient, result: StreamingResponse) -> SpendLogRow:
|
||||
"""The passthrough call's logged row, polled until it carries a cost.
|
||||
|
||||
Asserts (not skips) that a 2xx passthrough call produced a costed row - the
|
||||
whole point of passthrough spend tracking.
|
||||
"""
|
||||
assert result.call_id, "passthrough response had no x-litellm-call-id header"
|
||||
rows = client.gateway.poll_logs_for_request_id(
|
||||
result.call_id,
|
||||
predicate=lambda rs: (rs[0].spend or 0) > 0,
|
||||
)
|
||||
assert rows, f"no SpendLogs row for passthrough call_id {result.call_id}"
|
||||
row = rows[0]
|
||||
assert row.call_type == "pass_through_endpoint"
|
||||
assert (row.spend or 0) > 0, f"passthrough call was not costed: {row}"
|
||||
assert row.status == "success"
|
||||
return row
|
||||
|
||||
|
||||
# ---- Gemini passthrough ------------------------------------------------
|
||||
|
||||
|
||||
def test_gemini_passthrough_nonstreaming_logs_cost(
|
||||
client: PassthroughClient, scoped_key: str
|
||||
) -> None:
|
||||
tag = f"e2e-passthrough-{unique_marker()}"
|
||||
result = client.gemini_generate(
|
||||
scoped_key, "gemini-2.5-flash", "Say hello in one word", tags=[tag, "gemini"]
|
||||
)
|
||||
require_successful_call(result)
|
||||
|
||||
row = _fetch_cost_breakdown(client, result)
|
||||
assert row.custom_llm_provider == "gemini"
|
||||
assert "gemini" in (row.model or "")
|
||||
assert tag in (row.request_tags or []), f"tags not logged: {row.request_tags}"
|
||||
|
||||
|
||||
def test_gemini_passthrough_streaming_logs_cost(
|
||||
client: PassthroughClient, scoped_key: str
|
||||
) -> None:
|
||||
result = client.gemini_stream(scoped_key, "gemini-2.5-flash", "Count to five")
|
||||
require_successful_call(result)
|
||||
assert result.chunks > 0, "streaming passthrough produced no events"
|
||||
|
||||
row = _fetch_cost_breakdown(client, result)
|
||||
assert row.custom_llm_provider == "gemini"
|
||||
|
||||
|
||||
def test_gemini_passthrough_tool_call_logs_cost(
|
||||
client: PassthroughClient, scoped_key: str
|
||||
) -> None:
|
||||
result = client.gemini_generate(
|
||||
scoped_key,
|
||||
"gemini-2.5-flash",
|
||||
"What is the weather in Paris? Use the get_weather tool.",
|
||||
tools=[
|
||||
GeminiTool(
|
||||
function_declarations=[
|
||||
GeminiFunctionDeclaration(
|
||||
name="get_weather",
|
||||
description="Get the weather for a city",
|
||||
parameters=JsonSchema(
|
||||
type="object",
|
||||
properties={"city": JsonSchemaProperty(type="string")},
|
||||
required=["city"],
|
||||
),
|
||||
)
|
||||
]
|
||||
)
|
||||
],
|
||||
)
|
||||
require_successful_call(result)
|
||||
assert "functionCall" in result.body, "gemini did not emit a tool call"
|
||||
|
||||
row = _fetch_cost_breakdown(client, result)
|
||||
assert row.custom_llm_provider == "gemini"
|
||||
|
||||
|
||||
# ---- Anthropic passthrough ---------------------------------------------
|
||||
|
||||
|
||||
def test_anthropic_passthrough_nonstreaming_logs_cost(
|
||||
client: PassthroughClient, scoped_key: str
|
||||
) -> None:
|
||||
result = client.anthropic_message(scoped_key, "claude-haiku-4-5", "Say hello")
|
||||
require_successful_call(result)
|
||||
|
||||
row = _fetch_cost_breakdown(client, result)
|
||||
assert row.custom_llm_provider == "anthropic"
|
||||
assert "claude" in (row.model or "")
|
||||
|
||||
|
||||
def test_anthropic_passthrough_streaming_logs_cost(
|
||||
client: PassthroughClient, scoped_key: str
|
||||
) -> None:
|
||||
result = client.anthropic_message(
|
||||
scoped_key, "claude-haiku-4-5", "Count to five", stream=True
|
||||
)
|
||||
require_successful_call(result)
|
||||
assert result.chunks > 0, "streaming passthrough produced no events"
|
||||
|
||||
row = _fetch_cost_breakdown(client, result)
|
||||
assert row.custom_llm_provider == "anthropic"
|
||||
|
||||
|
||||
def test_anthropic_passthrough_tool_call_logs_cost(
|
||||
client: PassthroughClient, scoped_key: str
|
||||
) -> None:
|
||||
result = client.anthropic_message(
|
||||
scoped_key,
|
||||
"claude-haiku-4-5",
|
||||
"What is the weather in Paris? Use the get_weather tool.",
|
||||
tools=[
|
||||
AnthropicTool(
|
||||
name="get_weather",
|
||||
description="Get the weather for a city",
|
||||
input_schema=JsonSchema(
|
||||
type="object",
|
||||
properties={"city": JsonSchemaProperty(type="string")},
|
||||
required=["city"],
|
||||
),
|
||||
)
|
||||
],
|
||||
)
|
||||
require_successful_call(result)
|
||||
assert "tool_use" in result.body, "anthropic did not emit a tool call"
|
||||
|
||||
row = _fetch_cost_breakdown(client, result)
|
||||
assert row.custom_llm_provider == "anthropic"
|
||||
240
tests/e2e/models.py
Normal file
240
tests/e2e/models.py
Normal file
|
|
@ -0,0 +1,240 @@
|
|||
"""Shared pydantic request/response models for the e2e gateway.
|
||||
|
||||
Only the fields the tests read are modelled; pydantic ignores the rest, so a
|
||||
response validates without mirroring every proxy field. No untyped dicts.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, RootModel
|
||||
|
||||
# ---------- keys ----------
|
||||
|
||||
|
||||
class ModelBudgetEntry(BaseModel):
|
||||
budget_limit: float
|
||||
time_period: str
|
||||
|
||||
|
||||
class BudgetWindow(BaseModel):
|
||||
budget_duration: str
|
||||
max_budget: float
|
||||
|
||||
|
||||
class KeyGenerateBody(BaseModel):
|
||||
models: list[str] = []
|
||||
duration: str | None = None
|
||||
max_budget: float | None = None
|
||||
soft_budget: float | None = None
|
||||
budget_duration: str | None = None
|
||||
user_id: str | None = None
|
||||
team_id: str | None = None
|
||||
budget_id: str | None = None
|
||||
model_max_budget: dict[str, ModelBudgetEntry] | None = None
|
||||
budget_limits: list[BudgetWindow] | None = None
|
||||
tpm_limit: int | None = None
|
||||
rpm_limit: int | None = None
|
||||
|
||||
|
||||
class KeyGenerateResponse(BaseModel):
|
||||
key: str
|
||||
|
||||
|
||||
class KeyDeleteBody(BaseModel):
|
||||
keys: list[str]
|
||||
|
||||
|
||||
class KeyInfoParams(BaseModel):
|
||||
key: str
|
||||
|
||||
|
||||
class LiteLLMBudgetTable(BaseModel):
|
||||
max_budget: float | None = None
|
||||
soft_budget: float | None = None
|
||||
budget_duration: str | None = None
|
||||
budget_reset_at: str | None = None
|
||||
|
||||
|
||||
class KeyInfo(BaseModel):
|
||||
spend: float | None = None
|
||||
max_budget: float | None = None
|
||||
budget_reset_at: str | None = None
|
||||
budget_id: str | None = None
|
||||
litellm_budget_table: LiteLLMBudgetTable | None = None
|
||||
|
||||
|
||||
class KeyInfoResponse(BaseModel):
|
||||
info: KeyInfo
|
||||
|
||||
|
||||
# ---------- customers ----------
|
||||
|
||||
|
||||
class CustomerDeleteBody(BaseModel):
|
||||
user_ids: list[str]
|
||||
|
||||
|
||||
# ---------- chat / embeddings ----------
|
||||
|
||||
|
||||
class ChatMetadata(BaseModel):
|
||||
tags: list[str] | None = None
|
||||
|
||||
|
||||
class ChatMessage(BaseModel):
|
||||
role: str
|
||||
content: str
|
||||
|
||||
|
||||
class ChatBody(BaseModel):
|
||||
model: str
|
||||
messages: list[ChatMessage]
|
||||
stream: bool = False
|
||||
max_tokens: int | None = None
|
||||
user: str | None = None
|
||||
metadata: ChatMetadata | None = None
|
||||
|
||||
|
||||
class OutMessage(BaseModel):
|
||||
content: str | None = None
|
||||
|
||||
|
||||
class ChatChoice(BaseModel):
|
||||
message: OutMessage | None = None
|
||||
|
||||
|
||||
class Usage(BaseModel):
|
||||
prompt_tokens: int | None = None
|
||||
completion_tokens: int | None = None
|
||||
total_tokens: int | None = None
|
||||
|
||||
|
||||
class ChatResponse(BaseModel):
|
||||
id: str | None = None
|
||||
model: str | None = None
|
||||
choices: list[ChatChoice] = []
|
||||
usage: Usage | None = None
|
||||
|
||||
|
||||
class EmbedBody(BaseModel):
|
||||
model: str
|
||||
input: str
|
||||
|
||||
|
||||
class EmbedResponse(BaseModel):
|
||||
model: str | None = None
|
||||
|
||||
|
||||
# ---------- spend logs ----------
|
||||
|
||||
|
||||
class SpendLogRow(BaseModel):
|
||||
request_id: str | None = None
|
||||
model: str | None = None
|
||||
spend: float | None = None
|
||||
status: str | None = None
|
||||
cache_hit: str | None = None
|
||||
call_type: str | None = None
|
||||
custom_llm_provider: str | None = None
|
||||
team_id: str | None = None
|
||||
user: str | None = None
|
||||
end_user: str | None = None
|
||||
prompt_tokens: int | None = None
|
||||
completion_tokens: int | None = None
|
||||
total_tokens: int | None = None
|
||||
request_tags: list[str] | None = None
|
||||
|
||||
|
||||
class SpendLogs(RootModel[list[SpendLogRow]]):
|
||||
pass
|
||||
|
||||
|
||||
class SpendLogsParams(BaseModel):
|
||||
request_id: str | None = None
|
||||
api_key: str | None = None
|
||||
|
||||
|
||||
# ---------- spend calculate ----------
|
||||
|
||||
|
||||
class SpendCalculateBody(BaseModel):
|
||||
model: str
|
||||
messages: list[ChatMessage]
|
||||
|
||||
|
||||
class SpendCalculateResponse(BaseModel):
|
||||
cost: float
|
||||
|
||||
|
||||
# ---------- route probing ----------
|
||||
|
||||
|
||||
class DateRangeParams(BaseModel):
|
||||
start_date: str
|
||||
end_date: str
|
||||
|
||||
|
||||
class RouteSpec(RootModel[dict[str, object]]):
|
||||
"""One /openapi.json path entry: a map of HTTP method -> operation. Only the
|
||||
method names are read, so the operation specs stay opaque."""
|
||||
|
||||
@property
|
||||
def methods(self) -> frozenset[str]:
|
||||
return frozenset(method.lower() for method in self.root)
|
||||
|
||||
|
||||
class OpenAPISchema(BaseModel):
|
||||
paths: dict[str, RouteSpec] = {}
|
||||
|
||||
|
||||
# ---------- model info / custom pricing ----------
|
||||
|
||||
|
||||
class CustomPricing(BaseModel):
|
||||
"""The per-token custom-pricing fields a deployment can override in
|
||||
litellm_params - the token-cost subset of litellm's CustomPricingLiteLLMParams
|
||||
the proxy applies to chat spend. All optional: a config sets only what it
|
||||
overrides, and /model/info echoes the rates the proxy resolved."""
|
||||
|
||||
model_config = ConfigDict(extra="ignore")
|
||||
input_cost_per_token: float | None = None
|
||||
output_cost_per_token: float | None = None
|
||||
cache_read_input_token_cost: float | None = None
|
||||
cache_creation_input_token_cost: float | None = None
|
||||
|
||||
def overrides(self) -> dict[str, float]:
|
||||
"""The rates actually declared (non-null) - e.g. those a config.yml sets."""
|
||||
declared = {
|
||||
"input_cost_per_token": self.input_cost_per_token,
|
||||
"output_cost_per_token": self.output_cost_per_token,
|
||||
"cache_read_input_token_cost": self.cache_read_input_token_cost,
|
||||
"cache_creation_input_token_cost": self.cache_creation_input_token_cost,
|
||||
}
|
||||
return {field: rate for field, rate in declared.items() if rate is not None}
|
||||
|
||||
def token_cost(self, prompt_tokens: int, completion_tokens: int) -> float:
|
||||
"""Spend for a fresh (uncached) call under these rates: the proxy's
|
||||
custom-pricing formula (prompt * input + completion * output)."""
|
||||
assert (
|
||||
self.input_cost_per_token is not None
|
||||
and self.output_cost_per_token is not None
|
||||
), "custom pricing has no per-token rates"
|
||||
return (
|
||||
prompt_tokens * self.input_cost_per_token
|
||||
+ completion_tokens * self.output_cost_per_token
|
||||
)
|
||||
|
||||
|
||||
class ModelInfoEntry(BaseModel):
|
||||
"""One /model/info row. `litellm_params` is the configured deployment (carries
|
||||
any custom-pricing override); `model_info` is the price the proxy resolved for
|
||||
it - the override merged over the cost-map defaults."""
|
||||
|
||||
model_config = ConfigDict(protected_namespaces=())
|
||||
model_name: str
|
||||
litellm_params: CustomPricing = CustomPricing()
|
||||
model_info: CustomPricing = CustomPricing()
|
||||
|
||||
|
||||
class ModelInfoResponse(BaseModel):
|
||||
data: list[ModelInfoEntry] = []
|
||||
7
tests/e2e/pytest.ini
Normal file
7
tests/e2e/pytest.ini
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
[pytest]
|
||||
# Config when any e2e suite under tests/e2e/ is run directly, e.g.
|
||||
# uv run pytest tests/e2e/spend_tracking/ -v
|
||||
# The e2e marker is also registered in conftest.py for runs rooted elsewhere.
|
||||
addopts = --strict-markers --strict-config
|
||||
markers =
|
||||
e2e: live test that requires a running proxy and real provider keys
|
||||
78
tests/e2e/spend_tracking/SPEND_TRACKING_COVERAGE_MATRIX.md
Normal file
78
tests/e2e/spend_tracking/SPEND_TRACKING_COVERAGE_MATRIX.md
Normal file
|
|
@ -0,0 +1,78 @@
|
|||
# Spend Tracking Test Coverage Matrix
|
||||
|
||||
Scope: every distinct spend-tracking code path, mapped to the test that exercises
|
||||
it and the level it runs at. Highlights where a live e2e check is the only thing
|
||||
that would catch a regression.
|
||||
|
||||
Companion: live suite `test_spend_tracking_e2e.py` + route breadth
|
||||
`test_spend_routes.py` (this directory). Offline regression suite:
|
||||
`tests/test_litellm/proxy/spend_tracking/`. Reference PR: BerriAI/litellm#29956.
|
||||
|
||||
Levels: `unit` mocked; `integration` real DB/cost-map; `live` real provider +
|
||||
proxy + SpendLogs rows. Status: `covered` / `partial` / `gap`.
|
||||
|
||||
---
|
||||
|
||||
## SpendLogs row construction (`spend_tracking_utils.get_logging_payload`)
|
||||
|
||||
| Path | Existing | Level | Status | Live e2e |
|
||||
|------|----------|-------|--------|----------|
|
||||
| `_get_status_for_spend_log` | `test_spend_tracking_utils.py` | unit | covered | yes (status read off the row) |
|
||||
| cache-hit `request_id` suffix | `test_spend_tracking_utils.py` | unit | covered | yes (`test_cache_hit_is_zero_cost_and_suffixed`) |
|
||||
| failure status + zero spend | `test_spend_tracking_utils.py` | unit | covered | no (live failure logging is non-deterministic across providers) |
|
||||
| per-model / per-provider attribution | `test_spend_tracking_utils.py` | unit | covered | yes (`test_each_model_on_a_shared_key_gets_its_own_row`) |
|
||||
| field population (model/tokens/api_key/team/org) | `test_spend_tracking_utils.py` | unit | partial | yes (asserts real values) |
|
||||
| `request_tags` propagation | `test_db_spend_update_writer.py` | unit | partial | yes (`test_request_tags_round_trip`) |
|
||||
| `end_user` attribution | unit | unit | partial | yes (`test_end_user_spend_attributed_on_row`) |
|
||||
|
||||
## Cost calculation by modality
|
||||
|
||||
| Modality | Existing | Status | Live e2e |
|
||||
|----------|----------|--------|----------|
|
||||
| Chat (non-stream) | `test_cost_calculator.py`, `local_testing/test_completion_cost.py` | covered | yes (`test_chat_completion_writes_nonzero_spend_row`) |
|
||||
| Chat (streaming) | `test_streaming_interrupt_spend_tracking.py` | partial | yes (`test_streaming_chat_completion_tracks_spend`) |
|
||||
| Embedding | `test_cost_calculator.py` (#29956) | partial | yes (`test_embedding_writes_nonzero_spend_row`) |
|
||||
| Pass-through (gemini/anthropic) | `pass_through_tests/*.test.js` + `llm_translation/` suite | covered | yes (llm_translation suite) |
|
||||
| Image / audio / rerank / responses / realtime | per-provider unit cost tests | partial/gap | gap |
|
||||
|
||||
## Entity spend aggregation
|
||||
|
||||
| Entity | Existing | Status | Live e2e |
|
||||
|--------|----------|--------|----------|
|
||||
| API key | `test_db_spend_update_writer.py`, `test_spend_counters.py` | covered | yes (`test_key_spend_equals_sum_of_logs`) |
|
||||
| Tag | `test_update_daily_tag_spend.py` | partial | yes (`test_request_tags_round_trip`, propagation only) |
|
||||
| End-user | `test_proxy_update_spend.py` | covered | yes |
|
||||
| Spend == sum(logs) consistency | none | gap | yes (key aggregate == sum of rows) |
|
||||
|
||||
## Spend read endpoints (verification surface)
|
||||
|
||||
| Endpoint | Existing | Status | Live e2e |
|
||||
|----------|----------|--------|----------|
|
||||
| `/spend/logs` (request_id / api_key) | `test_spend_management_endpoints.py` | covered | yes (primary read path; `test_spend_logs_endpoint_returns_spend` asserts 200 + spend, never 5xx) |
|
||||
| `/spend/calculate` | `local_testing/test_spend_calculate_endpoint.py` | covered | yes (`test_spend_calculate_returns_nonzero_cost`) |
|
||||
| `/spend/tags` | `test_spend_management_endpoints.py` | partial | yes (`test_spend_routes.py` route probe) |
|
||||
| whole spend GET surface (22 routes) | unit per-handler | partial | yes (`test_spend_routes.py` probes each for 404/5xx) |
|
||||
|
||||
## What this suite pins
|
||||
|
||||
| Test | Invariant |
|
||||
|------|-----------|
|
||||
| `test_chat_completion_writes_nonzero_spend_row` | nonzero cost, token arithmetic, status, row findable by `response.id` |
|
||||
| `test_streaming_chat_completion_tracks_spend` | streamed responses still costed |
|
||||
| `test_embedding_writes_nonzero_spend_row` | embedding cost, `completion_tokens == 0` |
|
||||
| `test_cache_hit_is_zero_cost_and_suffixed` | cache hits not double-charged; `_cache_hit` suffix |
|
||||
| `test_key_spend_equals_sum_of_logs` | key aggregate == sum of rows |
|
||||
| `test_request_tags_round_trip` | tags persist onto the row |
|
||||
| `test_end_user_spend_attributed_on_row` | `end_user` attributed + costed |
|
||||
| `test_each_model_on_a_shared_key_gets_its_own_row` | per-model/provider rows, correct model + cost, distinct request_ids matching response id |
|
||||
| `test_spend_calculate_returns_nonzero_cost` | cost-map smoke (no batch wait) |
|
||||
| `test_spend_logs_endpoint_returns_spend` | `/spend/logs` returns 200 + the key's spend, never a 5xx (intermittent-500 regression) |
|
||||
| `test_spend_routes.py` (23) | no spend route 404s or 5xxs |
|
||||
|
||||
## Design + timing
|
||||
|
||||
`proxy_batch_write_at` (~60s) means rows land late; every read polls to a deadline.
|
||||
Fresh scoped key per test (isolation, xdist-safe, cleaned up). Assert invariants
|
||||
(`spend > 0`, `total == prompt + completion`, aggregate == sum), not literal
|
||||
$/token values, so pricing drift is not a failure. Skip on environment (no proxy /
|
||||
no provider key), fail on behavior (a real 2xx call with a wrong/missing row).
|
||||
16
tests/e2e/spend_tracking/conftest.py
Normal file
16
tests/e2e/spend_tracking/conftest.py
Normal file
|
|
@ -0,0 +1,16 @@
|
|||
"""Spend-tracking suite's `client` fixture.
|
||||
|
||||
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
|
||||
live in the parent tests/e2e/conftest.py. SpendClient exposes the shared Gateway
|
||||
(GatewayProvider), so the `resources` fixture cleans up keys and customers this
|
||||
suite creates.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from spend_e2e_client import SpendClient, build_client
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def client() -> SpendClient:
|
||||
return build_client()
|
||||
167
tests/e2e/spend_tracking/spend_e2e_client.py
Normal file
167
tests/e2e/spend_tracking/spend_e2e_client.py
Normal file
|
|
@ -0,0 +1,167 @@
|
|||
"""Spend-tracking e2e client: a Gateway plus the spend-specific read endpoints.
|
||||
|
||||
Generic proxy operations (keys, customers, chat/embed, route probing, SpendLogs
|
||||
polling) come from the shared Gateway, DI'd in (composition, not inheritance).
|
||||
This client adds only the spend surface: /spend/calculate, key-spend
|
||||
polling, and the route probes the breadth test uses.
|
||||
|
||||
Re-exports unwrap / is_ok / unique_marker / SpendLogRow so the tests import their
|
||||
helpers from one place.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import (
|
||||
NoBody,
|
||||
ProbeResult,
|
||||
Result,
|
||||
StreamingResponse,
|
||||
is_ok,
|
||||
unwrap,
|
||||
)
|
||||
from e2e_gateway import Gateway, build_gateway
|
||||
from models import (
|
||||
ChatBody,
|
||||
ChatMessage,
|
||||
ChatMetadata,
|
||||
ChatResponse,
|
||||
DateRangeParams,
|
||||
EmbedBody,
|
||||
EmbedResponse,
|
||||
OpenAPISchema,
|
||||
SpendCalculateBody,
|
||||
SpendCalculateResponse,
|
||||
SpendLogRow,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"SpendClient",
|
||||
"build_client",
|
||||
"reset_spend_logs",
|
||||
"unique_marker",
|
||||
"unwrap",
|
||||
"is_ok",
|
||||
"SpendLogRow",
|
||||
"ProbeResult",
|
||||
]
|
||||
|
||||
|
||||
def reset_spend_logs() -> None:
|
||||
"""Truncate LiteLLM_SpendLogs for a clean slate. No proxy endpoint deletes
|
||||
spend logs (/global/spend/reset keeps them), so go to the DB directly. Uses
|
||||
DATABASE_URL (default: the local docker postgres on its mapped host port; note
|
||||
the in-container `@db` host isn't resolvable from the host, so default to
|
||||
localhost).
|
||||
"""
|
||||
import psycopg
|
||||
|
||||
url = os.environ.get(
|
||||
"DATABASE_URL",
|
||||
"postgresql://llmproxy:dbpassword9090@localhost:5432/litellm",
|
||||
)
|
||||
with psycopg.connect(url) as conn:
|
||||
_ = conn.execute('TRUNCATE TABLE "LiteLLM_SpendLogs"')
|
||||
|
||||
|
||||
def _chat_body(
|
||||
model: str,
|
||||
content: str,
|
||||
*,
|
||||
max_tokens: int | None = None,
|
||||
tags: list[str] | None = None,
|
||||
user: str | None = None,
|
||||
stream: bool = False,
|
||||
) -> ChatBody:
|
||||
return ChatBody(
|
||||
model=model,
|
||||
messages=[ChatMessage(role="user", content=content)],
|
||||
max_tokens=max_tokens,
|
||||
stream=stream,
|
||||
user=user,
|
||||
metadata=ChatMetadata(tags=tags) if tags else None,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SpendClient:
|
||||
gateway: Gateway
|
||||
|
||||
def chat(
|
||||
self,
|
||||
key: str,
|
||||
model: str,
|
||||
content: str,
|
||||
*,
|
||||
max_tokens: int | None = None,
|
||||
tags: list[str] | None = None,
|
||||
user: str | None = None,
|
||||
) -> Result[ChatResponse]:
|
||||
return self.gateway.chat(
|
||||
key, _chat_body(model, content, max_tokens=max_tokens, tags=tags, user=user)
|
||||
)
|
||||
|
||||
def chat_stream(
|
||||
self, key: str, model: str, content: str, *, max_tokens: int | None = None
|
||||
) -> StreamingResponse:
|
||||
return self.gateway.chat_stream(
|
||||
key, _chat_body(model, content, max_tokens=max_tokens, stream=True)
|
||||
)
|
||||
|
||||
def embed(self, key: str, model: str, content: str) -> Result[EmbedResponse]:
|
||||
return self.gateway.embed(key, EmbedBody(model=model, input=content))
|
||||
|
||||
def poll_logs_for_key(
|
||||
self,
|
||||
key: str,
|
||||
*,
|
||||
min_rows: int = 1,
|
||||
predicate: Callable[[list[SpendLogRow]], bool] | None = None,
|
||||
) -> list[SpendLogRow]:
|
||||
return self.gateway.poll_logs_for_key(
|
||||
key, min_rows=min_rows, predicate=predicate
|
||||
)
|
||||
|
||||
def calculate_spend(self, model: str, content: str) -> float:
|
||||
return unwrap(
|
||||
self.gateway.transport.post(
|
||||
"/spend/calculate",
|
||||
headers=self.gateway.transport.master,
|
||||
json=SpendCalculateBody(
|
||||
model=model, messages=[ChatMessage(role="user", content=content)]
|
||||
),
|
||||
response_type=SpendCalculateResponse,
|
||||
)
|
||||
).cost
|
||||
|
||||
def poll_key_spend(self, key: str, *, minimum: float = 0.0) -> float:
|
||||
deadline = time.monotonic() + self.gateway.poll_timeout
|
||||
spend = 0.0
|
||||
while time.monotonic() < deadline:
|
||||
spend = self.gateway.key_info(key).spend or 0.0
|
||||
if spend > minimum:
|
||||
return spend
|
||||
time.sleep(self.gateway.poll_interval)
|
||||
return spend
|
||||
|
||||
def probe(self, path: str, *, params: DateRangeParams) -> ProbeResult:
|
||||
return self.gateway.transport.probe(path, params=params)
|
||||
|
||||
def openapi(self) -> OpenAPISchema:
|
||||
return unwrap(
|
||||
self.gateway.transport.get(
|
||||
"/openapi.json",
|
||||
headers=self.gateway.transport.master,
|
||||
params=NoBody(),
|
||||
response_type=OpenAPISchema,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def build_client() -> SpendClient:
|
||||
return SpendClient(gateway=build_gateway())
|
||||
96
tests/e2e/spend_tracking/test_spend_routes.py
Normal file
96
tests/e2e/spend_tracking/test_spend_routes.py
Normal file
|
|
@ -0,0 +1,96 @@
|
|||
"""Breadth check: query every route on the spend read surface and show what it
|
||||
returns.
|
||||
|
||||
Spend tracking sprawls across many routes (model-cost / key / user / team / org /
|
||||
customer aggregation, tags, and activity reports). Most are served with
|
||||
`include_in_schema=False`, so they do NOT appear in `/openapi.json` - discovery
|
||||
from the schema alone misses ~70% of the surface. So we probe a curated, verified
|
||||
list directly, plus any spend route the schema does list (to auto-catch new ones).
|
||||
|
||||
Each probe captures status AND body, so a failure shows the proxy's actual error
|
||||
(a 500 traceback, a 404 meaning the route was removed) rather than a bare code.
|
||||
Run with `-rA` (or `-s`) to print every route's response, not just failures.
|
||||
|
||||
Healthy == route exists (not 404) and handler did not crash (not 5xx). A 4xx
|
||||
(missing params / auth nuance) still means the route is wired and ran. Cheap and
|
||||
fast: no batch-write wait, no provider calls.
|
||||
"""
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
import pytest
|
||||
|
||||
from models import DateRangeParams
|
||||
from spend_e2e_client import SpendClient
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
# Verified present and responsive on a live proxy. One per row of the spend
|
||||
# surface: key / user / team / org / customer aggregation, model-cost, tags,
|
||||
# activity.
|
||||
SPEND_ROUTES = (
|
||||
"/spend/keys",
|
||||
"/spend/users",
|
||||
"/spend/tags",
|
||||
"/spend/logs",
|
||||
"/spend/logs/ui",
|
||||
"/global/spend",
|
||||
"/global/spend/keys",
|
||||
"/global/spend/teams",
|
||||
"/global/spend/models",
|
||||
"/global/spend/provider",
|
||||
"/global/spend/report",
|
||||
"/global/spend/tags",
|
||||
"/global/spend/logs",
|
||||
"/global/spend/all_tag_names",
|
||||
"/global/activity",
|
||||
"/global/activity/model",
|
||||
"/global/activity/exceptions",
|
||||
"/key/list",
|
||||
"/user/list",
|
||||
"/team/list",
|
||||
"/organization/list",
|
||||
"/customer/list",
|
||||
)
|
||||
|
||||
_SPEND_PREFIXES = ("/spend", "/global/spend", "/global/activity")
|
||||
|
||||
|
||||
def _date_range() -> DateRangeParams:
|
||||
# Satisfies date-required endpoints (report/activity/provider); ignored elsewhere.
|
||||
end = datetime.now(timezone.utc).date()
|
||||
start = end - timedelta(days=1)
|
||||
return DateRangeParams(start_date=start.isoformat(), end_date=end.isoformat())
|
||||
|
||||
|
||||
@pytest.mark.parametrize("route", SPEND_ROUTES)
|
||||
def test_spend_route_responsive(client: SpendClient, route: str) -> None:
|
||||
result = client.probe(route, params=_date_range())
|
||||
print(f"{route} -> {result.status_code}\n{result.body[:600]}")
|
||||
assert result.healthy, f"{route} -> {result.status_code}\n{result.body[:600]}"
|
||||
|
||||
|
||||
def test_schema_listed_spend_routes_are_responsive(client: SpendClient) -> None:
|
||||
"""Probe any spend GET route the schema lists that isn't in SPEND_ROUTES."""
|
||||
schema = client.openapi()
|
||||
assert schema.paths, "/openapi.json had no paths"
|
||||
|
||||
discovered = [
|
||||
path
|
||||
for path, spec in schema.paths.items()
|
||||
if "get" in spec.methods
|
||||
and "{" not in path
|
||||
and any(path.startswith(prefix) for prefix in _SPEND_PREFIXES)
|
||||
]
|
||||
extras = [path for path in discovered if path not in SPEND_ROUTES]
|
||||
|
||||
params = _date_range()
|
||||
results = [(path, client.probe(path, params=params)) for path in extras]
|
||||
for path, result in results:
|
||||
print(f"{path} -> {result.status_code}")
|
||||
offenders = [
|
||||
f"{path} -> {result.status_code}\n{result.body[:600]}"
|
||||
for path, result in results
|
||||
if not result.healthy
|
||||
]
|
||||
assert not offenders, "non-responsive schema spend routes:\n" + "\n".join(offenders)
|
||||
328
tests/e2e/spend_tracking/test_spend_tracking_e2e.py
Normal file
328
tests/e2e/spend_tracking/test_spend_tracking_e2e.py
Normal file
|
|
@ -0,0 +1,328 @@
|
|||
"""Live end-to-end spend-tracking tests against a running proxy.
|
||||
|
||||
Run against a proxy started with the gateway config. Coverage rationale:
|
||||
SPEND_TRACKING_COVERAGE_MATRIX.md.
|
||||
|
||||
Model names are literals from that config: chat tests hit "gemini-2.5-flash",
|
||||
embedding tests hit "openai-text-embedding-3-small".
|
||||
|
||||
Every test: fresh scoped key (isolation) -> real provider call -> unwrap (hard
|
||||
fail if the proxy couldn't make a call it should) -> poll /spend/logs to a
|
||||
deadline (rows land ~60s later via proxy_batch_write_at) -> assert invariants on
|
||||
the real row (spend, token arithmetic, status, cache).
|
||||
|
||||
Assertions target invariants, not literals: a regression in the spend pipeline
|
||||
fails the test; a pricing or token-count drift does not.
|
||||
"""
|
||||
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
|
||||
import pytest
|
||||
|
||||
from e2e_http import Success
|
||||
from lifecycle import ResourceManager
|
||||
from models import SpendLogs, SpendLogsParams
|
||||
from spend_e2e_client import SpendClient, SpendLogRow, unique_marker, unwrap
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
|
||||
def _approx_equal(actual: float, expected: float) -> bool:
|
||||
"""Within 1% or 1e-9 absolute - spend math, not exact float identity."""
|
||||
return abs(actual - expected) <= max(1e-9, abs(expected) * 1e-2)
|
||||
|
||||
|
||||
def _summarize(rows: list[SpendLogRow]) -> list[dict[str, object]]:
|
||||
fields = {
|
||||
"request_id",
|
||||
"model",
|
||||
"spend",
|
||||
"status",
|
||||
"cache_hit",
|
||||
"prompt_tokens",
|
||||
"completion_tokens",
|
||||
"total_tokens",
|
||||
}
|
||||
return [row.model_dump(include=fields) for row in rows]
|
||||
|
||||
|
||||
def _require_row(
|
||||
rows: list[SpendLogRow], predicate: Callable[[SpendLogRow], bool], what: str
|
||||
) -> SpendLogRow:
|
||||
matches = [r for r in rows if predicate(r)]
|
||||
assert matches, (
|
||||
f"no SpendLogs row {what} after polling; saw {len(rows)} row(s): "
|
||||
f"{_summarize(rows)}"
|
||||
)
|
||||
return matches[0]
|
||||
|
||||
|
||||
def test_chat_completion_writes_nonzero_spend_row(
|
||||
client: SpendClient, scoped_key: str
|
||||
) -> None:
|
||||
chat = unwrap(
|
||||
client.chat(
|
||||
scoped_key,
|
||||
"gemini-2.5-flash",
|
||||
f"reply with one word {unique_marker()}",
|
||||
max_tokens=16,
|
||||
)
|
||||
)
|
||||
|
||||
rows = client.poll_logs_for_key(
|
||||
scoped_key, predicate=lambda rs: any(r.status == "success" for r in rs)
|
||||
)
|
||||
row = _require_row(rows, lambda r: r.status == "success", "for the chat call")
|
||||
|
||||
assert (row.spend or 0) > 0, f"chat row should cost > 0: {_summarize(rows)}"
|
||||
assert row.status == "success"
|
||||
assert row.cache_hit != "True", "fresh call must not be a cache hit"
|
||||
assert "gemini-2.5-flash" in (row.model or "")
|
||||
|
||||
prompt = row.prompt_tokens or 0
|
||||
completion = row.completion_tokens or 0
|
||||
total = row.total_tokens or 0
|
||||
assert prompt > 0 and completion > 0
|
||||
assert total == prompt + completion, f"token arithmetic broken: {_summarize(rows)}"
|
||||
|
||||
if chat.id:
|
||||
assert any(
|
||||
r.request_id == chat.id for r in rows
|
||||
), f"row request_id != client response.id ({chat.id})"
|
||||
|
||||
|
||||
def test_streaming_chat_completion_tracks_spend(
|
||||
client: SpendClient, scoped_key: str
|
||||
) -> None:
|
||||
result = client.chat_stream(
|
||||
scoped_key,
|
||||
"gemini-2.5-flash",
|
||||
f"count to three {unique_marker()}",
|
||||
max_tokens=64,
|
||||
)
|
||||
assert (
|
||||
result.ok
|
||||
), f"stream failed (status {result.status_code}): {result.body[:300]}"
|
||||
|
||||
rows = client.poll_logs_for_key(
|
||||
scoped_key, predicate=lambda rs: any((r.spend or 0) > 0 for r in rs)
|
||||
)
|
||||
row = _require_row(
|
||||
rows, lambda r: (r.spend or 0) > 0, "with nonzero spend for the stream"
|
||||
)
|
||||
prompt = row.prompt_tokens or 0
|
||||
completion = row.completion_tokens or 0
|
||||
assert (
|
||||
prompt > 0 and completion > 0
|
||||
), f"streaming tokens not tracked: {_summarize(rows)}"
|
||||
assert (row.total_tokens or 0) == prompt + completion
|
||||
|
||||
|
||||
def test_embedding_writes_nonzero_spend_row(
|
||||
client: SpendClient, scoped_key: str
|
||||
) -> None:
|
||||
_ = unwrap(
|
||||
client.embed(
|
||||
scoped_key,
|
||||
"openai-text-embedding-3-small",
|
||||
f"vectorize this sentence {unique_marker()}",
|
||||
)
|
||||
)
|
||||
|
||||
rows = client.poll_logs_for_key(
|
||||
scoped_key, predicate=lambda rs: any((r.spend or 0) > 0 for r in rs)
|
||||
)
|
||||
row = _require_row(
|
||||
rows, lambda r: (r.spend or 0) > 0, "with nonzero spend for the embedding"
|
||||
)
|
||||
assert (row.prompt_tokens or 0) > 0
|
||||
assert (row.completion_tokens or 0) == 0, "embeddings have no completion tokens"
|
||||
assert "text-embedding-3-small" in (row.model or "")
|
||||
|
||||
|
||||
def test_cache_hit_is_zero_cost_and_suffixed(
|
||||
client: SpendClient, scoped_key: str
|
||||
) -> None:
|
||||
# Unique marker shared by both calls: call 1 is a guaranteed cache MISS (fresh
|
||||
# content, paid), call 2 repeats the identical request and HITS the cache just
|
||||
# populated. The marker keeps each run isolated - a fixed prompt would persist
|
||||
# in the shared response cache across runs and make both calls hit (flaky).
|
||||
prompt = f"What is the capital of France? Answer in one word. {unique_marker()}"
|
||||
_ = unwrap(client.chat(scoped_key, "gemini-2.5-flash", prompt, max_tokens=16))
|
||||
_ = unwrap(client.chat(scoped_key, "gemini-2.5-flash", prompt, max_tokens=16))
|
||||
|
||||
rows = client.poll_logs_for_key(
|
||||
scoped_key, predicate=lambda rs: any(r.cache_hit == "True" for r in rs)
|
||||
)
|
||||
cache_rows = [r for r in rows if r.cache_hit == "True"]
|
||||
if not cache_rows:
|
||||
pytest.skip(
|
||||
"no cache-hit row observed; caching may be disabled on this proxy. "
|
||||
f"rows seen: {_summarize(rows)}"
|
||||
)
|
||||
|
||||
cache_row = cache_rows[0]
|
||||
assert (
|
||||
cache_row.spend or 0
|
||||
) == 0.0, f"cache hit was charged (double-charge regression): {_summarize(rows)}"
|
||||
assert "_cache_hit" in (cache_row.request_id or ""), (
|
||||
"cache-hit row missing the _cache_hit request_id suffix; "
|
||||
"duplicate-key collisions will silently drop rows"
|
||||
)
|
||||
paid_rows = [r for r in rows if r.cache_hit != "True"]
|
||||
assert any(
|
||||
(r.spend or 0) > 0 for r in paid_rows
|
||||
), f"the non-cached call should still be charged: {_summarize(rows)}"
|
||||
|
||||
|
||||
def test_key_spend_equals_sum_of_logs(client: SpendClient, scoped_key: str) -> None:
|
||||
for _ in range(2):
|
||||
_ = unwrap(
|
||||
client.chat(
|
||||
scoped_key,
|
||||
"gemini-2.5-flash",
|
||||
f"say hi {unique_marker()}",
|
||||
max_tokens=16,
|
||||
)
|
||||
)
|
||||
|
||||
rows = client.poll_logs_for_key(
|
||||
scoped_key,
|
||||
min_rows=2,
|
||||
predicate=lambda rs: sum((r.spend or 0) for r in rs) > 0,
|
||||
)
|
||||
assert len(rows) >= 2, f"expected >=2 rows for the key, saw {_summarize(rows)}"
|
||||
logs_total = sum((r.spend or 0) for r in rows)
|
||||
assert logs_total > 0
|
||||
|
||||
key_spend = client.poll_key_spend(scoped_key, minimum=logs_total * 0.999)
|
||||
assert _approx_equal(
|
||||
key_spend, logs_total
|
||||
), f"key aggregate {key_spend} != sum of logs {logs_total}; rows: {_summarize(rows)}"
|
||||
|
||||
|
||||
def test_request_tags_round_trip(client: SpendClient, scoped_key: str) -> None:
|
||||
tag = f"e2e-spend-{unique_marker()}"
|
||||
_ = unwrap(
|
||||
client.chat(
|
||||
scoped_key, "gemini-2.5-flash", "tagged request", tags=[tag], max_tokens=16
|
||||
)
|
||||
)
|
||||
|
||||
rows = client.poll_logs_for_key(
|
||||
scoped_key, predicate=lambda rs: any(tag in (r.request_tags or []) for r in rs)
|
||||
)
|
||||
_require_row(
|
||||
rows, lambda r: tag in (r.request_tags or []), f"carrying request tag {tag!r}"
|
||||
)
|
||||
|
||||
|
||||
def test_end_user_spend_attributed_on_row(
|
||||
client: SpendClient, scoped_key: str, resources: ResourceManager
|
||||
) -> None:
|
||||
customer = resources.customer(f"e2e-cust-{unique_marker()}")
|
||||
_ = unwrap(
|
||||
client.chat(scoped_key, "gemini-2.5-flash", "hi", user=customer, max_tokens=16)
|
||||
)
|
||||
|
||||
rows = client.poll_logs_for_key(
|
||||
scoped_key, predicate=lambda rs: any(r.end_user == customer for r in rs)
|
||||
)
|
||||
row = _require_row(
|
||||
rows, lambda r: r.end_user == customer, f"attributed to end_user {customer!r}"
|
||||
)
|
||||
assert (row.spend or 0) > 0, f"end-user row should cost > 0: {_summarize(rows)}"
|
||||
|
||||
|
||||
def test_each_model_on_a_shared_key_gets_its_own_row(
|
||||
client: SpendClient, scoped_key: str
|
||||
) -> None:
|
||||
"""One key calling two different models, on two providers, gets one spend row per
|
||||
call - each carrying its own model and a nonzero cost, under distinct request_ids
|
||||
that match the call's response id. Pins per-model/per-provider attribution: a
|
||||
regression that stamps the wrong model on the row, bills a call's cost to the
|
||||
sibling deployment, or collapses both calls onto one request_id fails here."""
|
||||
gemini = unwrap(
|
||||
client.chat(
|
||||
scoped_key, "gemini-2.5-flash", f"one word {unique_marker()}", max_tokens=16
|
||||
)
|
||||
)
|
||||
claude = unwrap(
|
||||
client.chat(
|
||||
scoped_key, "claude-haiku-4-5", f"one word {unique_marker()}", max_tokens=16
|
||||
)
|
||||
)
|
||||
|
||||
def both_models_costed(rows: list[SpendLogRow]) -> bool:
|
||||
costed = [r.model or "" for r in rows if (r.spend or 0) > 0]
|
||||
return any("gemini-2.5-flash" in m for m in costed) and any(
|
||||
"claude-haiku-4-5" in m for m in costed
|
||||
)
|
||||
|
||||
rows = client.poll_logs_for_key(scoped_key, min_rows=2, predicate=both_models_costed)
|
||||
gemini_row = _require_row(
|
||||
rows, lambda r: "gemini-2.5-flash" in (r.model or ""), "for the gemini call"
|
||||
)
|
||||
claude_row = _require_row(
|
||||
rows, lambda r: "claude-haiku-4-5" in (r.model or ""), "for the claude call"
|
||||
)
|
||||
|
||||
assert (gemini_row.spend or 0) > 0, f"gemini row should cost > 0: {_summarize(rows)}"
|
||||
assert (claude_row.spend or 0) > 0, f"claude row should cost > 0: {_summarize(rows)}"
|
||||
assert (
|
||||
gemini_row.request_id != claude_row.request_id
|
||||
), f"two distinct calls collapsed onto one request_id: {_summarize(rows)}"
|
||||
if gemini.id:
|
||||
assert (
|
||||
gemini_row.request_id == gemini.id
|
||||
), f"gemini row request_id {gemini_row.request_id} != response id {gemini.id}"
|
||||
if claude.id:
|
||||
assert (
|
||||
claude_row.request_id == claude.id
|
||||
), f"claude row request_id {claude_row.request_id} != response id {claude.id}"
|
||||
|
||||
|
||||
def test_spend_calculate_returns_nonzero_cost(client: SpendClient) -> None:
|
||||
cost = client.calculate_spend(
|
||||
"gemini-2.5-flash", "estimate the cost of this request"
|
||||
)
|
||||
assert cost > 0, (
|
||||
"/spend/calculate returned 0 for gemini-2.5-flash; "
|
||||
"cost map may be missing this model"
|
||||
)
|
||||
|
||||
|
||||
def test_spend_logs_endpoint_returns_spend(
|
||||
client: SpendClient, scoped_key: str
|
||||
) -> None:
|
||||
"""The /spend/logs read endpoint returns a 200 carrying the key's spend, never a
|
||||
5xx. Regression for intermittent 500s (DB query / serialization errors under load)
|
||||
on this endpoint: every poll asserts a success response, not just a truthy row
|
||||
list, so a 500 fails loudly instead of being swallowed as 'no rows yet'; the
|
||||
call's nonzero spend must surface before the deadline."""
|
||||
unwrap(
|
||||
client.chat(
|
||||
scoped_key, "gemini-2.5-flash", f"spend logs {unique_marker()}", max_tokens=16
|
||||
)
|
||||
)
|
||||
|
||||
gateway = client.gateway
|
||||
deadline = time.monotonic() + gateway.poll_timeout
|
||||
while True:
|
||||
result = gateway.transport.get(
|
||||
"/spend/logs",
|
||||
headers=gateway.transport.master,
|
||||
params=SpendLogsParams(api_key=scoped_key),
|
||||
response_type=SpendLogs,
|
||||
)
|
||||
assert isinstance(result, Success), f"/spend/logs did not return 200 OK: {result}"
|
||||
rows = result.data.root
|
||||
if sum((r.spend or 0) for r in rows) > 0:
|
||||
return
|
||||
if time.monotonic() >= deadline:
|
||||
pytest.fail(
|
||||
f"/spend/logs never surfaced the key's spend before the deadline; "
|
||||
f"saw {_summarize(rows)}"
|
||||
)
|
||||
time.sleep(gateway.poll_interval)
|
||||
46
tests/e2e/test_lifecycle.py
Normal file
46
tests/e2e/test_lifecycle.py
Normal file
|
|
@ -0,0 +1,46 @@
|
|||
"""Unit coverage for the lifecycle harness (lifecycle.run_case).
|
||||
|
||||
Cases register cleanups progressively during init() (create team, then user, then
|
||||
key), so a failure partway through init() must still release whatever was already
|
||||
created on the long-lived shared proxy. This guards that contract.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable, List
|
||||
|
||||
import pytest
|
||||
|
||||
from lifecycle import run_case
|
||||
|
||||
|
||||
@dataclass
|
||||
class _PartialInitCase:
|
||||
"""init() registers a cleanup, then raises before finishing - mirroring a real
|
||||
case that creates a resource, registers its delete, then fails on the next
|
||||
step."""
|
||||
|
||||
released: List[str] = field(default_factory=list)
|
||||
_undo: List[Callable[[], None]] = field(default_factory=list)
|
||||
|
||||
def init(self) -> None:
|
||||
self._undo.append(lambda: self.released.append("first"))
|
||||
raise RuntimeError("init failed after registering the first resource")
|
||||
|
||||
def run(self) -> None:
|
||||
raise AssertionError("run() must not execute when init() failed")
|
||||
|
||||
def teardown(self) -> None:
|
||||
for undo in reversed(self._undo):
|
||||
undo()
|
||||
|
||||
|
||||
def test_run_case_releases_resources_when_init_fails_partway() -> None:
|
||||
case = _PartialInitCase()
|
||||
|
||||
with pytest.raises(RuntimeError, match="init failed"):
|
||||
run_case(case)
|
||||
|
||||
assert case.released == ["first"], (
|
||||
"a resource registered before init() failed must still be released, or it "
|
||||
"leaks on the long-lived shared proxy"
|
||||
)
|
||||
244
tests/e2e/transport.py
Normal file
244
tests/e2e/transport.py
Normal file
|
|
@ -0,0 +1,244 @@
|
|||
"""Transport: the typed request primitives clients use, behind a Protocol.
|
||||
|
||||
`Transport` is what each client depends on (composition + DI); `HttpTransport` is
|
||||
the concrete frozen-slots dataclass that fulfils it via the e2e_http wrapper. No
|
||||
client touches requests.* or builds raw dicts; they pass pydantic models here.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Protocol
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
import e2e_http
|
||||
from e2e_http import URL, AuthHeaders, ProbeResult, Result, StreamingResponse
|
||||
|
||||
|
||||
class Transport(Protocol):
|
||||
def post[R: BaseModel](
|
||||
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
|
||||
) -> Result[R]: ...
|
||||
|
||||
def stream(
|
||||
self, path: str, *, headers: BaseModel, json: BaseModel
|
||||
) -> StreamingResponse: ...
|
||||
|
||||
def send(
|
||||
self,
|
||||
path: str,
|
||||
*,
|
||||
headers: BaseModel,
|
||||
json: BaseModel,
|
||||
params: BaseModel | None = None,
|
||||
stream: bool = False,
|
||||
) -> StreamingResponse: ...
|
||||
|
||||
def get[R: BaseModel](
|
||||
self,
|
||||
path: str,
|
||||
*,
|
||||
headers: BaseModel,
|
||||
params: BaseModel,
|
||||
response_type: type[R],
|
||||
) -> Result[R]: ...
|
||||
|
||||
def delete[R: BaseModel](
|
||||
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
|
||||
) -> Result[R]: ...
|
||||
|
||||
def probe(self, path: str, *, params: BaseModel) -> ProbeResult: ...
|
||||
|
||||
def bearer(self, key: str) -> AuthHeaders: ...
|
||||
|
||||
@property
|
||||
def master(self) -> AuthHeaders: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class HttpTransport:
|
||||
base_url: str
|
||||
master_key: str
|
||||
request_timeout: float = 60.0
|
||||
|
||||
def _url(self, path: str) -> URL:
|
||||
return URL(f"{self.base_url.rstrip('/')}{path}")
|
||||
|
||||
def bearer(self, key: str) -> AuthHeaders:
|
||||
return AuthHeaders(authorization=f"Bearer {key}")
|
||||
|
||||
@property
|
||||
def master(self) -> AuthHeaders:
|
||||
return self.bearer(self.master_key)
|
||||
|
||||
def post[R: BaseModel](
|
||||
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
|
||||
) -> Result[R]:
|
||||
return e2e_http.post(
|
||||
self._url(path),
|
||||
headers=headers,
|
||||
json=json,
|
||||
response_type=response_type,
|
||||
timeout=self.request_timeout,
|
||||
)
|
||||
|
||||
def get[R: BaseModel](
|
||||
self,
|
||||
path: str,
|
||||
*,
|
||||
headers: BaseModel,
|
||||
params: BaseModel,
|
||||
response_type: type[R],
|
||||
) -> Result[R]:
|
||||
return e2e_http.get(
|
||||
self._url(path),
|
||||
headers=headers,
|
||||
params=params,
|
||||
response_type=response_type,
|
||||
timeout=self.request_timeout,
|
||||
)
|
||||
|
||||
def delete[R: BaseModel](
|
||||
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
|
||||
) -> Result[R]:
|
||||
return e2e_http.delete(
|
||||
self._url(path),
|
||||
headers=headers,
|
||||
json=json,
|
||||
response_type=response_type,
|
||||
timeout=self.request_timeout,
|
||||
)
|
||||
|
||||
def stream(
|
||||
self, path: str, *, headers: BaseModel, json: BaseModel
|
||||
) -> StreamingResponse:
|
||||
return e2e_http.stream(
|
||||
self._url(path), headers=headers, json=json, timeout=self.request_timeout
|
||||
)
|
||||
|
||||
def send(
|
||||
self,
|
||||
path: str,
|
||||
*,
|
||||
headers: BaseModel,
|
||||
json: BaseModel,
|
||||
params: BaseModel | None = None,
|
||||
stream: bool = False,
|
||||
) -> StreamingResponse:
|
||||
return e2e_http.send(
|
||||
self._url(path),
|
||||
headers=headers,
|
||||
json=json,
|
||||
params=params,
|
||||
stream=stream,
|
||||
timeout=self.request_timeout,
|
||||
)
|
||||
|
||||
def probe(self, path: str, *, params: BaseModel) -> ProbeResult:
|
||||
return e2e_http.probe(
|
||||
self._url(path),
|
||||
headers=self.master,
|
||||
params=params,
|
||||
timeout=self.request_timeout,
|
||||
)
|
||||
|
||||
|
||||
# Top-level management/admin route groups. In a split deployment these are served
|
||||
# by the control plane (a different service from the LLM data plane). LLM routes
|
||||
# (/chat, /embeddings, and native passthrough like /gemini, /anthropic) are NOT
|
||||
# here and fall through to the data plane. Matched as path prefixes.
|
||||
CONTROL_PLANE_PREFIXES: tuple[str, ...] = (
|
||||
"/key",
|
||||
"/user",
|
||||
"/team",
|
||||
"/organization",
|
||||
"/customer",
|
||||
"/tag",
|
||||
"/budget",
|
||||
"/model/info",
|
||||
"/spend",
|
||||
"/global",
|
||||
"/openapi.json",
|
||||
)
|
||||
|
||||
|
||||
def is_control_plane_path(path: str) -> bool:
|
||||
"""True if `path` is a management/admin route (served by the control plane in a
|
||||
split deployment), false for LLM data-plane routes."""
|
||||
return path.startswith(CONTROL_PLANE_PREFIXES)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SplitTransport:
|
||||
"""A Transport that dispatches each call by path to one of two backends: the
|
||||
management/admin control plane or the LLM data plane.
|
||||
|
||||
Litellm can run as a split control-plane/data-plane deployment where the two
|
||||
surfaces live on different services. Clients here stay plane-agnostic — they
|
||||
keep calling ``transport.post("/budget/new", ...)`` or
|
||||
``transport.send("/chat/completions", ...)`` — and routing happens in one place
|
||||
by path (see ``CONTROL_PLANE_PREFIXES``). When ``control`` and ``data`` share a
|
||||
base URL (the monolithic default), routing is a no-op. ``bearer``/``master``
|
||||
are plane-agnostic (same master key both planes), so they come from ``data``.
|
||||
"""
|
||||
|
||||
data: HttpTransport
|
||||
control: HttpTransport
|
||||
|
||||
def _route(self, path: str) -> HttpTransport:
|
||||
return self.control if is_control_plane_path(path) else self.data
|
||||
|
||||
def bearer(self, key: str) -> AuthHeaders:
|
||||
return self.data.bearer(key)
|
||||
|
||||
@property
|
||||
def master(self) -> AuthHeaders:
|
||||
return self.data.master
|
||||
|
||||
def post[R: BaseModel](
|
||||
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
|
||||
) -> Result[R]:
|
||||
return self._route(path).post(
|
||||
path, headers=headers, json=json, response_type=response_type
|
||||
)
|
||||
|
||||
def get[R: BaseModel](
|
||||
self,
|
||||
path: str,
|
||||
*,
|
||||
headers: BaseModel,
|
||||
params: BaseModel,
|
||||
response_type: type[R],
|
||||
) -> Result[R]:
|
||||
return self._route(path).get(
|
||||
path, headers=headers, params=params, response_type=response_type
|
||||
)
|
||||
|
||||
def delete[R: BaseModel](
|
||||
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
|
||||
) -> Result[R]:
|
||||
return self._route(path).delete(
|
||||
path, headers=headers, json=json, response_type=response_type
|
||||
)
|
||||
|
||||
def stream(
|
||||
self, path: str, *, headers: BaseModel, json: BaseModel
|
||||
) -> StreamingResponse:
|
||||
return self._route(path).stream(path, headers=headers, json=json)
|
||||
|
||||
def send(
|
||||
self,
|
||||
path: str,
|
||||
*,
|
||||
headers: BaseModel,
|
||||
json: BaseModel,
|
||||
params: BaseModel | None = None,
|
||||
stream: bool = False,
|
||||
) -> StreamingResponse:
|
||||
return self._route(path).send(
|
||||
path, headers=headers, json=json, params=params, stream=stream
|
||||
)
|
||||
|
||||
def probe(self, path: str, *, params: BaseModel) -> ProbeResult:
|
||||
return self._route(path).probe(path, params=params)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1089,7 +1089,9 @@ async def test_list_tools_only_returns_allowed_servers(monkeypatch):
|
|||
mock_client_constructor,
|
||||
):
|
||||
# Call list_tools
|
||||
tools = await test_manager.list_tools(user_api_key_auth=MagicMock())
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
tools = await test_manager.list_tools(user_api_key_auth=UserAPIKeyAuth())
|
||||
# Should only return tools from server_a
|
||||
assert len(tools) == 1
|
||||
# The server should use the server_name as prefix since no alias is provided
|
||||
|
|
|
|||
11
tests/pyrightconfig.json
Normal file
11
tests/pyrightconfig.json
Normal file
|
|
@ -0,0 +1,11 @@
|
|||
{
|
||||
"include": ["e2e"],
|
||||
"exclude": ["**/node_modules", "**/__pycache__"],
|
||||
"pythonVersion": "3.12",
|
||||
"typeCheckingMode": "strict",
|
||||
"enableTypeIgnoreComments": false,
|
||||
"reportMissingImports": false,
|
||||
"reportPrivateImportUsage": false,
|
||||
"reportExplicitAny": "error",
|
||||
"reportAny": "error"
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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__])
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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__])
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
0
tests/test_litellm/proxy/logging_endpoints/__init__.py
Normal file
0
tests/test_litellm/proxy/logging_endpoints/__init__.py
Normal file
|
|
@ -0,0 +1,195 @@
|
|||
"""Unit tests for POST /v1/callbacks/logs (replay logging payloads → callbacks)."""
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.logging_endpoints.callback_logs_endpoints import (
|
||||
CallbackLogsReplayer,
|
||||
ingest_callback_logs,
|
||||
)
|
||||
from litellm.types.proxy.callback_logs_endpoints import (
|
||||
CallbackLogRecord,
|
||||
CallbackLogsRequest,
|
||||
)
|
||||
|
||||
REQ_ID = "cb-logs-unit-test-1"
|
||||
|
||||
|
||||
def _sample_payload(**overrides):
|
||||
payload = {
|
||||
"id": REQ_ID,
|
||||
"litellm_call_id": REQ_ID,
|
||||
"call_type": "acompletion",
|
||||
"stream": False,
|
||||
"response_cost": 0.0123,
|
||||
"custom_llm_provider": "openai",
|
||||
"total_tokens": 42,
|
||||
"prompt_tokens": 30,
|
||||
"completion_tokens": 12,
|
||||
"startTime": time.time() - 2,
|
||||
"endTime": time.time(),
|
||||
"model": "gpt-4o-mini",
|
||||
"metadata": {
|
||||
"user_api_key_hash": "rust-gateway-test-key",
|
||||
"user_api_key_user_id": "user-cb-logs-test",
|
||||
"user_api_key_team_id": "team-cb-logs-test",
|
||||
},
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
}
|
||||
payload.update(overrides)
|
||||
return payload
|
||||
|
||||
|
||||
def test_epoch_to_datetime_handles_float_and_fallback():
|
||||
dt = CallbackLogsReplayer._epoch_to_datetime(1_700_000_000.5)
|
||||
assert dt.year == 2023
|
||||
# Non-numeric input must not raise — falls back to "now".
|
||||
assert CallbackLogsReplayer._epoch_to_datetime(None) is not None
|
||||
|
||||
|
||||
def test_build_logging_obj_seeds_model_call_details():
|
||||
obj = CallbackLogsReplayer._build_logging_obj(_sample_payload())
|
||||
details = obj.model_call_details
|
||||
# Prebuilt payload is set so the handler skips rebuilding it.
|
||||
assert details["standard_logging_object"]["id"] == REQ_ID
|
||||
assert details["response_cost"] == 0.0123
|
||||
assert details["call_type"] == "acompletion"
|
||||
# Metadata is mapped to the keys the cost-tracking callback reads.
|
||||
md = details["litellm_params"]["metadata"]
|
||||
assert md["user_api_key"] == "rust-gateway-test-key"
|
||||
assert md["user_api_key_user_id"] == "user-cb-logs-test"
|
||||
assert md["user_api_key_team_id"] == "team-cb-logs-test"
|
||||
|
||||
|
||||
def test_response_obj_carries_usage():
|
||||
obj = CallbackLogsReplayer._response_obj_from_payload(_sample_payload())
|
||||
assert obj["usage"]["total_tokens"] == 42
|
||||
assert obj["usage"]["prompt_tokens"] == 30
|
||||
assert obj["usage"]["completion_tokens"] == 12
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_success_record_invokes_success_handler(monkeypatch):
|
||||
captured = {}
|
||||
|
||||
async def fake_success(self, result=None, start_time=None, end_time=None, **kwargs):
|
||||
captured["standard_logging_object"] = self.model_call_details.get(
|
||||
"standard_logging_object"
|
||||
)
|
||||
captured["result"] = result
|
||||
|
||||
monkeypatch.setattr(LiteLLMLogging, "async_success_handler", fake_success)
|
||||
|
||||
body = CallbackLogsRequest(
|
||||
records=[
|
||||
CallbackLogRecord(
|
||||
status="success", standard_logging_payload=_sample_payload()
|
||||
)
|
||||
]
|
||||
)
|
||||
resp = await ingest_callback_logs(
|
||||
body, user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
)
|
||||
assert resp.processed == 1 and resp.failed == 0
|
||||
assert captured["standard_logging_object"]["id"] == REQ_ID
|
||||
assert captured["result"]["usage"]["total_tokens"] == 42
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failure_record_invokes_failure_handler(monkeypatch):
|
||||
captured = {}
|
||||
|
||||
async def fake_failure(
|
||||
self, exception, traceback_exception, start_time=None, end_time=None
|
||||
):
|
||||
captured["exception"] = str(exception)
|
||||
|
||||
monkeypatch.setattr(LiteLLMLogging, "async_failure_handler", fake_failure)
|
||||
|
||||
body = CallbackLogsRequest(
|
||||
records=[
|
||||
CallbackLogRecord(
|
||||
status="failure",
|
||||
standard_logging_payload=_sample_payload(),
|
||||
error="upstream exploded",
|
||||
)
|
||||
]
|
||||
)
|
||||
resp = await ingest_callback_logs(
|
||||
body, user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
)
|
||||
assert resp.processed == 1 and resp.failed == 0
|
||||
assert captured["exception"] == "upstream exploded"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_admin_is_rejected(monkeypatch):
|
||||
async def fake_success(self, **kwargs):
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(LiteLLMLogging, "async_success_handler", fake_success)
|
||||
|
||||
body = CallbackLogsRequest(
|
||||
records=[
|
||||
CallbackLogRecord(
|
||||
status="success", standard_logging_payload=_sample_payload()
|
||||
)
|
||||
]
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await ingest_callback_logs(
|
||||
body,
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER),
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_one_bad_record_does_not_sink_the_batch(monkeypatch):
|
||||
calls = {"n": 0}
|
||||
|
||||
async def flaky_success(
|
||||
self, result=None, start_time=None, end_time=None, **kwargs
|
||||
):
|
||||
calls["n"] += 1
|
||||
if calls["n"] == 1:
|
||||
raise ValueError("boom on first record")
|
||||
|
||||
monkeypatch.setattr(LiteLLMLogging, "async_success_handler", flaky_success)
|
||||
|
||||
body = CallbackLogsRequest(
|
||||
records=[
|
||||
CallbackLogRecord(
|
||||
status="success", standard_logging_payload=_sample_payload()
|
||||
),
|
||||
CallbackLogRecord(
|
||||
status="success", standard_logging_payload=_sample_payload()
|
||||
),
|
||||
]
|
||||
)
|
||||
resp = await ingest_callback_logs(
|
||||
body, user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
)
|
||||
assert resp.processed == 1 and resp.failed == 1
|
||||
# The failed record is reported back by index + error, not silently dropped.
|
||||
assert len(resp.failures) == 1
|
||||
assert resp.failures[0].index == 0
|
||||
assert "boom on first record" in resp.failures[0].error
|
||||
|
||||
|
||||
def test_batch_over_limit_is_rejected():
|
||||
from litellm.constants import MAX_CALLBACK_LOG_RECORDS
|
||||
from pydantic import ValidationError
|
||||
|
||||
# One over the cap must fail validation (422 at the API boundary), bounding
|
||||
# the callback/DB fan-out a single POST can trigger.
|
||||
too_many = [
|
||||
CallbackLogRecord(status="success", standard_logging_payload=_sample_payload())
|
||||
for _ in range(MAX_CALLBACK_LOG_RECORDS + 1)
|
||||
]
|
||||
with pytest.raises(ValidationError):
|
||||
CallbackLogsRequest(records=too_many)
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -398,6 +398,83 @@ def test_mock_create_audio_file(mocker: MockerFixture, monkeypatch, llm_router:
|
|||
app.dependency_overrides.pop(ps.user_api_key_auth, None)
|
||||
|
||||
|
||||
def test_create_file_batch_streams_from_upload_spool(monkeypatch, llm_router: Router):
|
||||
"""
|
||||
Batch uploads must be passed downstream as the upload's streamable file handle
|
||||
(Starlette's already-spooled file), not read into an in-memory bytes object, so
|
||||
the proxy never buffers the whole payload. Non-batch uploads keep the bytes path.
|
||||
"""
|
||||
import litellm.proxy.proxy_server as ps
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
from litellm.proxy.openai_files_endpoints import files_endpoints as fe
|
||||
from litellm.types.llms.openai import OpenAIFileObject
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", llm_router)
|
||||
setup_proxy_logging_object(monkeypatch, llm_router)
|
||||
|
||||
captured: dict = {}
|
||||
|
||||
async def fake_route_create_file(*, _create_file_request, **kwargs):
|
||||
file_elem = _create_file_request["file"][1]
|
||||
captured["file_elem"] = file_elem
|
||||
if hasattr(file_elem, "read") and hasattr(file_elem, "seek"):
|
||||
file_elem.seek(0)
|
||||
captured["streamed_content"] = file_elem.read()
|
||||
return OpenAIFileObject(
|
||||
id="dummy-id",
|
||||
object="file",
|
||||
bytes=0,
|
||||
created_at=1234567890,
|
||||
filename="batch.jsonl",
|
||||
purpose="batch",
|
||||
status="uploaded",
|
||||
)
|
||||
|
||||
monkeypatch.setattr(fe, "route_create_file", fake_route_create_file)
|
||||
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test-user"
|
||||
)
|
||||
|
||||
content = (
|
||||
b'{"custom_id":"r-0","method":"POST","url":"/v1/chat/completions",'
|
||||
b'"body":{"model":"gpt-3.5-turbo","messages":[{"role":"user","content":"hi"}]}}\n'
|
||||
)
|
||||
try:
|
||||
resp = client.post(
|
||||
"/v1/files",
|
||||
files={"file": ("batch.jsonl", content, "application/jsonl")},
|
||||
data={"purpose": "batch"},
|
||||
headers={"Authorization": "Bearer test-key"},
|
||||
)
|
||||
assert resp.status_code == 200, resp.text
|
||||
file_elem = captured["file_elem"]
|
||||
assert not isinstance(
|
||||
file_elem, (bytes, bytearray)
|
||||
), "batch upload must be a streamable handle, not in-memory bytes"
|
||||
assert hasattr(file_elem, "read") and hasattr(
|
||||
file_elem, "seek"
|
||||
), "batch upload must be a seekable file handle"
|
||||
assert (
|
||||
captured["streamed_content"] == content
|
||||
), "the handle must stream the uploaded bytes"
|
||||
|
||||
captured.clear()
|
||||
resp = client.post(
|
||||
"/v1/files",
|
||||
files={"file": ("data.jsonl", content, "application/jsonl")},
|
||||
data={"purpose": "user_data"},
|
||||
headers={"Authorization": "Bearer test-key"},
|
||||
)
|
||||
assert resp.status_code == 200, resp.text
|
||||
assert isinstance(
|
||||
captured["file_elem"], (bytes, bytearray)
|
||||
), "non-batch upload must stay in-memory bytes"
|
||||
finally:
|
||||
app.dependency_overrides.pop(ps.user_api_key_auth, None)
|
||||
|
||||
|
||||
@pytest.mark.flaky(retries=3, delay=2)
|
||||
def test_target_storage_invokes_storage_backend(
|
||||
mocker: MockerFixture, monkeypatch, llm_router: Router
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue