chore: merge internal staging into interview branch
Some checks are pending
LiteLLM Rust / rustfmt, clippy, test (push) Waiting to run

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Ishaan Jaff 2026-07-29 20:32:52 -07:00
commit 2008816278
172 changed files with 10261 additions and 2515 deletions

1
.gitignore vendored
View file

@ -141,3 +141,4 @@ crash.*.log
.coverage
ui/litellm-dashboard/out/
litellm.log

View file

@ -176,6 +176,8 @@ lint-ruff-FULL-dev: install-dev
if [ -n "$$files" ]; then echo "$$files" | xargs $(UV_RUN) ruff check; \
else echo "No changed .py files to check."; fi
lint-basedpyright lint-basedpyright-budget-update: export NODE_OPTIONS := --max-old-space-size=12288
lint-basedpyright: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --base origin/litellm_internal_staging

View file

@ -1,9 +1,9 @@
{
"reportAny": {
"limit": 34906
"limit": 33216
},
"reportArgumentType": {
"limit": 2701
"limit": 2648
},
"reportAssignmentType": {
"limit": 330
@ -24,7 +24,7 @@
"limit": 42
},
"reportExplicitAny": {
"limit": 10230
"limit": 10228
},
"reportFunctionMemberAccess": {
"limit": 11
@ -99,7 +99,7 @@
"limit": 0
},
"reportUnknownArgumentType": {
"limit": 45870
"limit": 45567
},
"reportUnknownLambdaType": {
"limit": 113

View file

@ -1,10 +1,11 @@
# Adding a provider / route to litellm-rust
Three layers, same for every route (see `ocr` and `realtime` as references):
Everything for a route lives in `crates/core/src/<route>/`; `crates/core/src/messages` is the reference. A host (the axum gateway, the Python bridge) only calls the route's entrypoint.
1. **Transform contract (pure)** — `crates/core/src/<route>/transformation.rs`: a `…ProviderConfig` trait (URL build + request/response transforms) + types in `types.rs`. No network, env, or auth.
2. **Provider config (pure)** — `crates/providers/src/<provider>/<route>/transformation.rs`: implement that trait as a `const <PROVIDER>_<ROUTE>_CONFIG`, mirroring the Python provider tree. Add parity unit tests.
3. **HTTP / transport (the host)** — `crates/providers/src/<route>.rs` (e.g. `ocr.rs`, `realtime.rs`): the callable fn (`run_ocr`, `realtime`). It resolves the key, builds the auth header, builds URL + transforms via the config, then does the network call. This is the only layer allowed to do I/O.
1. **Entrypoint** — `mod.rs`: `pub async fn <route>(request) -> CoreResult<Response>`, the Rust equivalent of `litellm.<route>()`, plus a `<route>_stream` variant when the route streams. It is the only thing a host touches.
2. **Transform contract** — `transformation.rs`: a `…ProviderConfig` trait (URL build + request/response transforms) with types in `types.rs`.
3. **Provider config** — `crates/core/src/providers/<provider>/<route>/transformation.rs`: implement that trait as a `const <PROVIDER>_<ROUTE>_CONFIG`, mirroring the Python provider tree. Add parity unit tests.
4. **Prepare + handler** — `prepare.rs` resolves provider/model, credentials, auth headers, and URL, then transforms the request; `handler.rs` performs the provider call through the shared client in `client.rs` and transforms the response.
## Coding standards
@ -25,4 +26,4 @@ variants of it. The test for a good abstraction is that adding the next provider
is a few declarative lines, not a new file of duplicated flow. Only diverge from
the base when behavior is genuinely different, and say so explicitly in the PR.
**Calling:** the host invokes the route fn — the Python bridge calls `run_ocr`; the `ai-gateway` server calls `realtime`. Register new modules in `lib.rs` / `mod.rs`, then run `cargo fmt && cargo clippy --workspace -- -D warnings && cargo test --workspace`.
**Calling:** hosts invoke the core entrypoint — the Python bridge and the `ai-gateway` route service both call `litellm_core::messages::messages`. Never add a provider handler to `ai-gateway`. Register new modules in `lib.rs` / `mod.rs`, then run `cargo fmt && cargo clippy --workspace -- -D warnings && cargo test --workspace`.

View file

@ -4,14 +4,30 @@ litellm-rust has exactly THREE crates. A crate is a LAYER, not a route. Routes (
## Crates
| Crate | Role | Pure / I/O |
|-------|------|------------|
| litellm-core | Translation layer — types, route contracts (traits), provider transforms (modules under providers/), and the router. Builds requests/responses; no network. | Pure |
| litellm-ai-gateway | Routes + host — the only crate that touches the network. HTTP/WebSocket I/O (modules under io/) plus the axum server binary (behind the `server` feature). | I/O |
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — a thin adapter over litellm-ai-gateway's I/O. | Binding |
| Crate | Role |
|-------|------|
| litellm-core | The LiteLLM SDK in Rust. One public entrypoint per top-level call (`messages::messages()`), owning types, transforms, provider resolution, auth, and the provider HTTP call. Call it, get a typed response. |
| litellm-ai-gateway | The axum server (behind the `server` feature) plus the WebSocket hosts. Translates HTTP/WS to core entrypoints; owns no provider logic and no handlers. |
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — marshals Python objects and calls core entrypoints. |
Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge.
## Where a route lives
A top-level LiteLLM call is a module under `crates/core/src/<route>/`, shaped like `messages`:
```
core/src/messages/
mod.rs # pub async fn messages(..) -> CoreResult<..> (+ messages_stream for SSE)
types.rs # request/response types, MessagesRequest
transformation.rs # the provider template trait
prepare.rs # provider resolution, auth headers, URL
handler.rs # the provider call
client.rs # the shared reqwest client
```
Handlers never live in `ai-gateway`. `ocr`, `audio_transcription`, and `realtime` are still hosted there from before this rule; they move to `core` as they are touched.
Adding a crate: default to a MODULE. New crate ONLY on a real trigger — separate artifact (binary/cdylib), proc-macro, shared foundation, or publishable standalone. A new provider or route is none of these.
Adding a crate fails crates/core/tests/workspace_crate_allowlist.rs until you update its allowlist and this file — intentional.

View file

@ -23,21 +23,34 @@ the base when behavior is genuinely different, and say so explicitly in the PR.
## Crates (exactly three — see AGENTS.md)
`litellm-core` describes work; `litellm-ai-gateway` executes it; `litellm-python-bridge`
exposes it to the Python SDK. A crate is a **layer**, not a route — add modules, not crates.
`litellm-core` **is** the LiteLLM SDK in Rust: it makes the LLM call.
`litellm-ai-gateway` is an HTTP/WebSocket server in front of it, and
`litellm-python-bridge` exposes it to the Python SDK. A crate is a **layer**, not
a route — add modules, not crates.
## Core Boundary
`litellm-core` is the pure translation layer; the `litellm-ai-gateway` host executes work.
`litellm-core` owns the whole call. The Rust equivalent of `litellm.messages()`
is `litellm_core::messages::messages(request).await`: you call it, it does the
provider call, and you get a typed non-streaming response back.
Route-level Rust structure mirrors LiteLLM's Python responsibilities:
- `core/src/<route>/` owns the route contract, shared types, and provider
template traits. For OCR, this means `core/src/ocr`.
- `core/src/<route>/` owns the route end to end: the public entrypoint fn named
after the route in `mod.rs`, the request/response types (`types.rs`), the
provider template trait (`transformation.rs`), the provider/auth/URL
resolution (`prepare.rs`), the HTTP client (`client.rs`), and the handler that
performs the call (`handler.rs`). `core/src/messages` is the reference.
- `core/src/providers/<provider>/<route>/transformation.rs` owns the
provider-specific transform. For Mistral OCR, this means
`core/src/providers/mistral/ocr/transformation.rs`.
- Network execution lives in the host crate `ai-gateway` (`ai-gateway/src/io/`),
never inside `core`.
provider-specific transform. For Anthropic Messages, this means
`core/src/providers/anthropic/messages/transformation.rs`.
- Handlers live in `core`, never in a host. `ai-gateway` must not contain a
route handler that talks to a provider; its axum route reads the HTTP request,
picks a deployment, and calls the `core` entrypoint. `python-bridge` marshals
Python objects and calls the same entrypoint.
Streaming keeps the same shape: the route entrypoint has a `<route>_stream`
variant in `core` that returns the upstream response so a host can splice it to
its own caller; the host still owns no provider logic.
Call-hook and lifecycle instrumentation, including phase timing, usage
accumulation, and callback payload construction, always lives in `core`.
@ -45,21 +58,31 @@ Hosts feed observed events into core and dispatch the completed payloads through
their I/O logger; hosts must not own callback orchestration.
Allowed in `core`:
- Pure request transforms
- Pure response transforms
- Pure stream chunk normalization
- The public entrypoint for a top-level LiteLLM call
- Request/response transforms and stream chunk normalization
- Provider resolution, auth header construction, and URL building
- The provider HTTP call itself, through a shared reused client with connect and
request timeouts
- Shared data types and validation errors
- Deterministic token/cost helper logic
Not allowed in `core`:
- Network calls
- Environment variable or secret reads
- Serving HTTP: axum routes, extractors, and transport concerns stay in the host
- Filesystem access
- Database or cache access
- Provider SDK signing or auth flows
- Database access
- Config file reading and rollout state
- Logging callbacks, spend writes, or custom callbacks
- Global mutable runtime state
Env reads in `core` are limited to credential fallback inside a route's
`prepare.rs` (the `env_lookup` closure), mirroring what the Python SDK does when
no key is passed. Everything else config-shaped is resolved by the host and
passed in.
Routes still hosted in `ai-gateway` (`ocr`, `audio_transcription`, `realtime`)
predate this rule and are being moved into `core` route modules; do not add new
ones there, and prefer moving one when you touch it.
Python owns rollout state and fallback while Rust is being introduced. Rust
paths must be off by default until parity tests prove equivalence with Python.
A new provider/route may instead be implemented rust-only with no Python
@ -93,10 +116,10 @@ the first PR:
- Preserve Python output shape intentionally. If a field is always serialized as
`null` for Python parity, leave a short comment explaining that parity choice.
## Host I/O Rules
## Network I/O Rules
These rules apply when adding future crates or modules that execute network I/O,
such as `ai-gateway`, router hosts, or standalone servers:
These rules apply to every module that executes network I/O, whether it is a
`core` route handler or a host such as `ai-gateway`:
- Set connect and full-request timeouts. No unbounded waits.
- Reuse HTTP clients; do not construct clients per request.

View file

@ -2,18 +2,31 @@
This workspace contains the staged Rust implementation for LiteLLM.
Rust starts as a pure transform core used by the existing Python host. Python
continues to own auth, configuration, network I/O, retries, routing, logging,
`litellm-core` is the LiteLLM SDK in Rust: one entrypoint per top-level call
that makes the LLM call and hands back a typed response, the same shape as
`litellm.messages()` in Python.
```rust
let response = litellm_core::messages::messages(MessagesRequest {
model: "claude-sonnet-4-5",
body,
api_key: Some(key),
..
})
.await?;
```
Python continues to own configuration, retries, routing policy, logging,
callbacks, spend tracking, and customer plugins until each Rust path has parity
coverage and production evidence.
## Crates
| Crate | Role | Pure / I/O |
|-------|------|------------|
| litellm-core | Translation layer — types, route contracts (traits), provider transforms (modules under providers/), and the router. Builds requests/responses; no network. | Pure |
| litellm-ai-gateway | Routes + host — the only crate that touches the network. HTTP/WebSocket I/O (modules under io/) plus the axum server binary (behind the `server` feature). | I/O |
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — a thin adapter over litellm-ai-gateway's I/O. | Binding |
| Crate | Role |
|-------|------|
| litellm-core | The SDK. Per-route entrypoints (`messages::messages()`), types, provider transforms (modules under `providers/`), provider resolution, auth, the provider HTTP call, and the router. |
| litellm-ai-gateway | The axum server (behind the `server` feature) and WebSocket hosts. Translates HTTP/WS to core entrypoints; no provider handlers. |
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — marshals Python objects and calls core entrypoints. |
Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge.
@ -21,16 +34,16 @@ Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-
```text
crates/
core/ Route contracts, shared pure types, errors, and templates.
src/ocr/
providers/ Provider-specific pure transforms.
src/mistral/ocr/transformation.rs
core/ The SDK: route modules + provider transforms.
src/messages/ mod.rs (entrypoint), types, transformation, prepare, handler, client
src/providers/anthropic/messages/transformation.rs
ai-gateway/ Axum server + WebSocket hosts; calls core entrypoints.
python-bridge/ PyO3 bridge for Python LiteLLM.
```
The folder shape should follow the Python provider tree:
`providers/src/<provider>/<route>/transformation.rs`. The bridge should expose
one function per top-level route, starting with `ocr(payload)`.
The folder shape follows the Python provider tree:
`core/src/providers/<provider>/<route>/transformation.rs`. The bridge exposes one
function per top-level route, mirroring the core entrypoints.
## Checks

View file

@ -1,6 +1,6 @@
# Provider coding standards (litellm-rust)
Rules for adding or changing an LLM provider/route in `litellm-rust`. OCR (`MISTRAL_OCR_CONFIG`) is the reference; `messages` (`ANTHROPIC_MESSAGES_CONFIG`) is the next port.
Rules for adding or changing an LLM provider/route in `litellm-rust`. `messages` (`core/src/messages`, `ANTHROPIC_MESSAGES_CONFIG`) is the reference: a route is a `core` module with a public entrypoint that makes the call and returns a typed response.
## Provider resolution
@ -16,10 +16,10 @@ Rules for adding or changing an LLM provider/route in `litellm-rust`. OCR (`MIST
## Boundaries
7. Layers never cross: `core` = pure transforms/types (no network, env, secrets, auth, logging, global mutable state); `ai-gateway` = all I/O, auth headers, HTTP/SSE, lifecycle hooks; `python-bridge` = thin PyO3 adapter.
7. Layers never cross: `core` = the call itself (entrypoint, types, transforms, provider resolution, auth headers, provider HTTP, lifecycle hooks); `ai-gateway` = serving HTTP/WS (routing, extractors, auth of *our* callers, streaming to the client); `python-bridge` = thin PyO3 adapter. Hosts call the core entrypoint; they never build a provider request.
8. Generic/route files contain zero provider-specific branches. A provider is one module under `core/src/providers/<provider>/<route>/`; a route is a module, never a new crate.
9. Route entry point stays thin: `<route>()` -> `prepare_*` -> `CallLifecycle::run_request`, which owns the pre_call -> during_call -> provider call -> success/failure order and phase timing. Handlers validate and delegate; no business logic in them.
10. Constants (URLs, env-var names, API versions, error messages) live in a crate `constants.rs`, never inline. Env reads happen only at the host/config layer, with the `DEFAULT_*` fallback defined in `constants.rs`.
9. Route entry point stays thin: `core::<route>::<route>()` -> `prepare_*` -> handler (or `CallLifecycle::run_request`, which owns the pre_call -> during_call -> provider call -> success/failure order and phase timing). Axum handlers validate and delegate to a service that calls the entrypoint; no business logic in them.
10. Constants (URLs, env-var names, API versions, error messages) live in a crate `constants.rs`, never inline. Config-shaped env reads happen at the host/config layer with the `DEFAULT_*` fallback defined in `constants.rs`; the only env read in `core` is the credential fallback in a route's `prepare.rs`.
## Types and errors
@ -33,7 +33,7 @@ Rules for adding or changing an LLM provider/route in `litellm-rust`. OCR (`MIST
16. Never log request/response bodies, base64 payloads, document contents, or secrets. Truncate and bound any upstream body before it crosses a host boundary.
17. Treat empty/whitespace credentials, URLs, and config values as absent at the host resolution layer.
18. Host I/O sets connect + request timeouts (no unbounded waits), reuses a shared HTTP client, and prefers rustls TLS.
18. Network I/O sets connect + request timeouts (no unbounded waits), reuses a shared HTTP client, and prefers rustls TLS.
## Tests and rollout

View file

@ -1,7 +1,9 @@
# ai-gateway — folder architecture
The Axum server that fronts the Rust gateway. It owns transport + config + auth
only; deployment selection lives in `core::router`, transforms in `core`/`providers`.
only; deployment selection lives in `core::router`, and the LLM call itself
(transforms, auth headers, provider HTTP) lives behind a `core` route entrypoint
such as `litellm_core::messages::messages`. No provider handler lives here.
```
src/
@ -32,6 +34,11 @@ src/
args; it runs during extraction. Never re-implement the check per route.
- **Handlers are thin.** A handler validates and delegates to its `service`. No
business logic, no provider calls, no transforms in handlers.
- **Services call `core`, they don't reimplement it.** A `service` picks the
deployment and calls the `core` route entrypoint. Provider resolution, auth
headers, URL building, and the HTTP call are `core`'s job; a service that
builds a provider request itself is a bug (`routes/messages/service.rs` is
the reference).
- **State is shared and cheap to clone.** Long-lived handles live behind `Arc` in
`state.rs`; read env/config only in `main.rs` when building state.

View file

@ -8,11 +8,11 @@ dials OpenAI upstream, and splices the two sockets frame-by-frame.
`litellm-rust` is exactly three crates (a crate is a **layer**, not a route):
| Crate | Role | Pure / I/O |
|-------|------|------------|
| litellm-core | Translation layer — types, route contracts (traits), provider transforms (modules under `providers/`), and the router. Builds requests/responses; no network. | Pure |
| litellm-ai-gateway | Routes + host — the only crate that touches the network. HTTP/WebSocket I/O (modules under `io/`) plus the Axum server binary (behind the `server` feature). | I/O |
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — a thin adapter over litellm-ai-gateway's I/O. | Binding |
| Crate | Role |
|-------|------|
| litellm-core | The LiteLLM SDK in Rust — per-route entrypoints (`messages::messages()`) that resolve the provider, transform, and make the call; plus types, provider transforms, and the router. |
| litellm-ai-gateway | The Axum server (behind the `server` feature) and WebSocket hosts. Translates HTTP/WS to core entrypoints; no provider handlers. |
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — marshals Python objects and calls core entrypoints. |
Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge.

View file

@ -29,18 +29,6 @@ pub(crate) const DEFAULT_FLUSH_INTERVAL_MS: u64 = 500;
#[cfg(feature = "server")]
pub(crate) const DEFAULT_PROVIDER: &str = "openai";
/// Full-request timeout ceiling for Anthropic Messages provider calls, in
/// seconds. Mirrors the Python Anthropic Messages default. The per-request
/// timeout from `litellm_params` still overrides this on the request builder.
pub(crate) const MESSAGES_TIMEOUT_SECS: u64 = 600;
/// Connect timeout for Anthropic Messages provider calls, in seconds.
pub(crate) const MESSAGES_CONNECT_TIMEOUT_SECS: u64 = 10;
/// Max characters of an upstream error body echoed across the host boundary
/// before truncation, so provider bodies are bounded and data-minimized.
pub(crate) const MESSAGES_ERROR_BODY_MAX_CHARS: usize = 256;
pub(crate) const DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS: u64 = 10;
pub(crate) const DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS: u64 = 300;
@ -48,12 +36,6 @@ pub(crate) const DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS: u64 = 300;
#[cfg(feature = "server")]
pub(crate) const MESSAGES_ROUTE_PATH: &str = "/v1/messages";
/// Provider name used by the Anthropic Messages route when a deployment's
/// provider model does not carry an explicit provider prefix.
pub(crate) const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic";
pub(crate) const AZURE_ANTHROPIC_MESSAGES_PROVIDER: &str = "azure_ai";
pub(crate) const BEDROCK_MESSAGES_PROVIDER: &str = "bedrock";
/// Request headers owned by the gateway and never forwarded upstream.
#[cfg(feature = "server")]
pub(crate) const MESSAGES_HEADERS_NOT_FORWARDED: &[&str] =

View file

@ -1 +0,0 @@
pub use crate::messages::{MessagesRequest, messages};

View file

@ -1,5 +1,4 @@
pub mod audio_transcription;
pub mod messages;
pub mod ocr;
pub mod realtime;
pub mod realtime_pool;

View file

@ -4,7 +4,9 @@
//! without pulling in the HTTP server:
//!
//! - Call-type modules such as [`ocr`]: provider transforms, lifecycle hooks,
//! and provider I/O. Always available — no feature required.
//! and provider I/O. Always available — no feature required. These predate the
//! rule that a route's entrypoint and handler live in `litellm-core` (see
//! `litellm_core::messages`) and move there as they are touched.
//! - [`io`]: compatibility exports and realtime WebSocket splice helpers.
//! - The server modules ([`auth`], [`routes`], [`state`]) and anything pulling
//! `axum` are gated behind the `server` feature, which the `litellm-ai-gateway`
@ -14,7 +16,6 @@
pub mod audio_transcription;
mod client;
pub mod io;
pub mod messages;
pub mod ocr;
/// GIL-activity tracking. Pure (atomics only); shared by the `server` routes and

View file

@ -1,160 +0,0 @@
use std::collections::BTreeMap;
use std::time::SystemTime;
use litellm_core::CoreResult;
use litellm_core::error::CoreError;
use litellm_core::providers::bedrock::aws_base::{
AwsAuthConfig, resolve_credentials, sign_bedrock_post,
};
use serde_json::Value;
use super::client::http_client;
use super::common_utils::truncate_error_body;
use super::types::ProviderMessagesRequest;
use crate::constants::{
ANTHROPIC_MESSAGES_PROVIDER, AZURE_ANTHROPIC_MESSAGES_PROVIDER, BEDROCK_MESSAGES_PROVIDER,
};
fn environment_lookup(key: &str) -> Option<String> {
std::env::var(key).ok()
}
async fn signed_request(
request: &ProviderMessagesRequest,
body: &[u8],
) -> CoreResult<Vec<(String, String)>> {
if request.provider != BEDROCK_MESSAGES_PROVIDER {
return Ok(request.upstream_headers.clone());
}
if let Some(token) = &request.bearer_token {
return Ok(request
.upstream_headers
.iter()
.filter(|(name, _)| {
!matches!(
name.to_ascii_lowercase().as_str(),
"authorization" | "x-api-key" | "anthropic-version"
)
})
.cloned()
.chain([
("Authorization".to_string(), format!("Bearer {token}")),
("content-type".to_string(), "application/json".to_string()),
])
.collect());
}
let headers = request
.upstream_headers
.iter()
.filter(|(name, _)| {
!matches!(
name.to_ascii_lowercase().as_str(),
"authorization" | "x-api-key" | "anthropic-version" | "host" | "content-length"
)
})
.cloned()
.chain(std::iter::once((
"content-type".to_string(),
"application/json".to_string(),
)))
.collect::<BTreeMap<_, _>>();
let region = request.signing_region.as_deref().ok_or_else(|| {
CoreError::InvalidRequest("Bedrock signing region was not resolved".to_string())
})?;
let credentials = resolve_credentials(AwsAuthConfig::default(), &environment_lookup).await?;
let signed = sign_bedrock_post(
&request.url,
body,
&headers,
region,
&credentials,
SystemTime::now(),
)?;
Ok(signed.into_iter().collect())
}
pub(super) async fn execute_messages_provider_call(
request: ProviderMessagesRequest,
) -> CoreResult<Value> {
let body = serde_json::to_vec(&request.body).map_err(|error| {
CoreError::InvalidRequest(format!("invalid messages request body: {error}"))
})?;
let headers = signed_request(&request, &body).await?;
let mut request_builder = http_client().post(&request.url).body(body);
for (key, value) in &headers {
request_builder = request_builder.header(key, value);
}
if let Some(duration) = request.timeout {
request_builder = request_builder.timeout(duration);
}
let response = request_builder
.send()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
let status = response.status();
let text = response
.text()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
if !status.is_success() {
return Err(CoreError::Http {
status: status.as_u16(),
body: truncate_error_body(&text),
});
}
let response = serde_json::from_str(&text).map_err(|err| {
CoreError::InvalidResponse(format!("invalid messages response JSON: {err}"))
})?;
let transformed = request
.config
.transform_response(&request.model, response)?;
serde_json::to_value(transformed).map_err(|err| {
CoreError::InvalidResponse(format!("failed to serialize messages response: {err}"))
})
}
pub(super) async fn execute_messages_provider_stream(
request: ProviderMessagesRequest,
) -> CoreResult<reqwest::Response> {
if !matches!(
request.provider.as_str(),
ANTHROPIC_MESSAGES_PROVIDER | AZURE_ANTHROPIC_MESSAGES_PROVIDER | BEDROCK_MESSAGES_PROVIDER
) {
return Err(CoreError::InvalidRequest(
"streaming messages is not supported for this provider".to_string(),
));
}
let body = serde_json::to_vec(&request.body).map_err(|error| {
CoreError::InvalidRequest(format!("invalid messages request body: {error}"))
})?;
let headers = signed_request(&request, &body).await?;
let mut request_builder = http_client().post(&request.url).body(body);
for (key, value) in &headers {
request_builder = request_builder.header(key, value);
}
if let Some(duration) = request.timeout {
request_builder = request_builder.timeout(duration);
}
let response = request_builder
.send()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
let status = response.status();
if !status.is_success() {
let text = response
.text()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
return Err(CoreError::Http {
status: status.as_u16(),
body: truncate_error_body(&text),
});
}
Ok(response)
}

View file

@ -1,51 +0,0 @@
use litellm_core::CoreResult;
use serde_json::Value;
mod client;
mod common_utils;
mod handler;
mod prepare;
mod types;
pub use types::MessagesRequest;
use handler::{execute_messages_provider_call, execute_messages_provider_stream};
use prepare::prepare_messages_call;
pub async fn messages(request: MessagesRequest<'_>) -> CoreResult<Value> {
match execute_messages(request, false).await? {
MessagesResponse::Json(body) => Ok(body),
MessagesResponse::Stream { .. } => Err(litellm_core::CoreError::InvalidResponse(
"non-streaming messages execution returned a stream".to_string(),
)),
}
}
pub(crate) enum MessagesResponse {
Json(Value),
#[allow(dead_code)]
Stream {
provider: String,
response: reqwest::Response,
},
}
pub(crate) async fn execute_messages(
request: MessagesRequest<'_>,
stream: bool,
) -> CoreResult<MessagesResponse> {
let prepared = prepare_messages_call(request)?;
if stream {
let provider = prepared.provider.clone();
execute_messages_provider_stream(prepared)
.await
.map(|response| MessagesResponse::Stream { provider, response })
} else {
execute_messages_provider_call(prepared)
.await
.map(MessagesResponse::Json)
}
}
#[cfg(test)]
mod tests;

View file

@ -1,26 +0,0 @@
use std::time::Duration;
use litellm_core::messages::transformation::AnthropicMessagesProviderConfig;
use serde_json::{Map, Value};
pub struct MessagesRequest<'a> {
pub model: &'a str,
pub body: Value,
pub api_key: Option<&'a str>,
pub api_base: Option<&'a str>,
pub custom_llm_provider: Option<&'a str>,
pub extra_headers: Option<Map<String, Value>>,
pub timeout: Option<Duration>,
}
pub(crate) struct ProviderMessagesRequest {
pub(crate) provider: String,
pub(crate) model: String,
pub(crate) config: &'static dyn AnthropicMessagesProviderConfig,
pub(crate) url: String,
pub(crate) body: Value,
pub(crate) upstream_headers: Vec<(String, String)>,
pub(crate) signing_region: Option<String>,
pub(crate) bearer_token: Option<String>,
pub(crate) timeout: Option<Duration>,
}

View file

@ -19,7 +19,10 @@ async fn handle(...) -> impl IntoResponse { ... }
When a route has business logic worth testing without axum, put it in a sibling
`service` (a file, or a folder if the route grows). The route file stays the
**axum surface** (router + handler + any socket/SSE adapter); `service` is plain
Rust with **no axum types**. `realtime/` is the example:
Rust with **no axum types**, and its job is to pick the deployment and call the
`core` route entrypoint (see `messages/service.rs` calling
`litellm_core::messages::messages`). Never build a provider request, resolve a
key, or perform the provider call here. `realtime/` is the older example:
```
realtime/
mod.rs # axum surface: router() + handler + the WS<->events adapter
@ -33,6 +36,8 @@ genuinely gets hard to read.
`crate::auth::RequireMasterKey` to its arguments; it runs during extraction.
Never re-implement the check per route.
- **Handlers contain no business logic; `service` contains no axum types.**
- **No provider handlers in this crate.** Transforms, auth headers, and the
provider HTTP call live in `core/src/<route>/`.
- A route owns its paths in its own `router()`; `mod.rs` only merges.
- Cross-cutting concerns (logging, CORS, timeouts) → Tower layers in `mod.rs`,
not duplicated in handlers.

View file

@ -1,18 +1,15 @@
use std::sync::Arc;
use litellm_core::constants::ANTHROPIC_MESSAGES_PROVIDER;
use litellm_core::messages::types::MessagesRequest;
use litellm_core::messages::{messages, messages_stream};
use litellm_core::router::Router;
use litellm_core::{CoreError, CoreResult};
use serde_json::{Map, Value};
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
use crate::messages::{MessagesRequest, execute_messages};
pub(crate) enum MessagesResponse {
Json(Value),
Stream {
provider: String,
response: reqwest::Response,
},
Stream(reqwest::Response),
}
pub async fn run(
@ -55,13 +52,14 @@ pub async fn run(
extra_headers,
timeout: None,
};
let stream = request.body.get("stream").and_then(Value::as_bool) == Some(true);
execute_messages(request, stream)
.await
.map(|response| match response {
crate::messages::MessagesResponse::Json(body) => MessagesResponse::Json(body),
crate::messages::MessagesResponse::Stream { provider, response } => {
MessagesResponse::Stream { provider, response }
}
if request.body.get("stream").and_then(Value::as_bool) == Some(true) {
return messages_stream(request).await.map(MessagesResponse::Stream);
}
let response = messages(request).await?;
serde_json::to_value(response)
.map(MessagesResponse::Json)
.map_err(|err| {
CoreError::InvalidResponse(format!("failed to serialize messages response: {err}"))
})
}

View file

@ -1,3 +1,7 @@
litellm-core is the PURE translation layer — types, route contracts (traits), provider transforms (modules under `providers/`), and the router. No network, no I/O, no env reads.
litellm-core is the LiteLLM SDK in Rust — it makes the LLM call. Each top-level call is a module under `src/<route>/` exposing a public entrypoint named after the route (`messages::messages()`, the Rust equivalent of `litellm.messages()`): you call it and get a typed non-streaming response back.
Routes (ocr, realtime) and providers (mistral, openai) are modules, not crates.
A route module owns everything the call needs: types, the provider template trait, provider transforms (under `providers/`), provider/auth/URL resolution, and the handler that performs the HTTP call. Handlers belong here, never in a host crate.
Not here: serving HTTP (axum routes, extractors), config file reading, rollout state, databases, or callback dispatch. Env reads are limited to credential fallback in a route's `prepare.rs`.
Routes (messages, ocr, realtime) and providers (anthropic, mistral, openai) are modules, not crates.

View file

@ -4,20 +4,28 @@ Rules for `litellm-rust/crates/core`.
## Responsibility
`core` owns shared data types, typed errors, and deterministic helper contracts.
It must stay pure and host-independent.
`core` is the LiteLLM SDK in Rust: it makes the LLM call. Every top-level
LiteLLM call has a public entrypoint here, named after the route
(`messages::messages()` is the Rust equivalent of `litellm.messages()`), and
calling it returns a typed non-streaming response.
Allowed:
- The public entrypoint for a route, plus its `<route>_stream` variant when the
route supports streaming.
- Provider resolution, auth header construction, URL building, and the provider
HTTP call (shared reused client, connect + request timeouts).
- Shared request/response structs.
- Typed errors with stable, non-sensitive messages.
- Deterministic validation helpers.
- Serialization helpers that intentionally mirror Python output shape.
- Route templates that match Python base config responsibilities, such as
`ocr::transformation::OcrProviderConfig`.
`messages::transformation::AnthropicMessagesProviderConfig`.
Not allowed:
- Network, filesystem, database, cache, or environment access.
- Secret reads or auth/header construction.
- Serving HTTP: axum routers, extractors, and other transport concerns.
- Filesystem, database, or cache access.
- Config file reading or rollout state; the host resolves those and passes them
in. Env reads are limited to credential fallback in a route's `prepare.rs`.
- Logging callbacks, tracing spans, spend writes, or customer callbacks.
- Provider-specific branching that belongs in `providers`.
- Panics for user/provider-controlled input.
@ -33,10 +41,21 @@ typed field on a struct, not a raw string threaded through the API.
## Structure
Use route names directly under `src/`: `ocr`, future `messages`,
Use route names directly under `src/`: `messages`, `ocr`, future
`chat_completions`, `embeddings`, and similar top-level LiteLLM calls. Do not
invent broad names like `engine` for route contracts.
`src/messages` is the reference shape for a route module:
```
mod.rs pub async fn messages(..) (+ messages_stream)
types.rs request/response types
transformation.rs the provider template trait
prepare.rs provider resolution, auth headers, URL
handler.rs the provider call
client.rs the shared reqwest client
```
## Parity Rules
- Every shared type used by a provider transform needs unit tests for

View file

@ -7,6 +7,7 @@ repository.workspace = true
[dependencies]
rand.workspace = true
reqwest.workspace = true
serde.workspace = true
serde_json.workspace = true
thiserror.workspace = true
@ -30,5 +31,4 @@ bedrock-auth = [
]
[dev-dependencies]
reqwest.workspace = true
tokio = { workspace = true, features = ["macros", "rt-multi-thread"] }

View file

@ -1,3 +1,19 @@
pub const OPENAI_DEFAULT_API_BASE: &str = "https://api.openai.com";
pub const OPENAI_RESPONSES_DEFAULT_API_BASE: &str = "https://api.openai.com/v1";
pub const OPENAI_RESPONSES_PATH: &str = "/responses";
/// Full-request timeout ceiling for Anthropic Messages provider calls, in
/// seconds. Mirrors the Python Anthropic Messages default. The per-request
/// timeout from the caller still overrides this on the request builder.
pub(crate) const MESSAGES_TIMEOUT_SECS: u64 = 600;
/// Connect timeout for Anthropic Messages provider calls, in seconds.
pub(crate) const MESSAGES_CONNECT_TIMEOUT_SECS: u64 = 10;
/// Max characters of an upstream error body echoed across the call boundary
/// before truncation, so provider bodies are bounded and data-minimized.
pub(crate) const MESSAGES_ERROR_BODY_MAX_CHARS: usize = 256;
/// Provider name used for Anthropic Messages when a deployment's provider model
/// does not carry an explicit provider prefix.
pub const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic";

View file

@ -1,12 +1,11 @@
use litellm_core::CoreResult;
use litellm_core::error::{CoreError, json_type_name};
use litellm_core::messages::transformation::AnthropicMessagesProviderConfig;
use litellm_core::providers::anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG;
use litellm_core::providers::azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG;
use litellm_core::providers::bedrock::messages::transformation::BEDROCK_MESSAGES_CONFIG;
use serde_json::{Map, Value};
use crate::constants::MESSAGES_ERROR_BODY_MAX_CHARS;
use crate::error::{CoreError, CoreResult, json_type_name};
use crate::providers::anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG;
use crate::providers::azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG;
use super::transformation::AnthropicMessagesProviderConfig;
pub(super) fn truncate_error_body(body: &str) -> String {
if body.chars().count() <= MESSAGES_ERROR_BODY_MAX_CHARS {
@ -22,7 +21,6 @@ pub(super) fn messages_provider_config(
match provider {
"anthropic" => Some(&ANTHROPIC_MESSAGES_CONFIG),
"azure_ai" => Some(&AZURE_ANTHROPIC_MESSAGES_CONFIG),
"bedrock" => Some(&BEDROCK_MESSAGES_CONFIG),
_ => None,
}
}

View file

@ -0,0 +1,76 @@
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
use crate::error::{CoreError, CoreResult};
use super::client::http_client;
use super::common_utils::truncate_error_body;
use super::types::{AnthropicMessagesResponse, ProviderMessagesRequest};
pub(super) async fn execute_messages_provider_call(
request: ProviderMessagesRequest,
) -> CoreResult<AnthropicMessagesResponse> {
let mut request_builder = http_client().post(&request.url).json(&request.body);
for (key, value) in &request.upstream_headers {
request_builder = request_builder.header(key, value);
}
if let Some(duration) = request.timeout {
request_builder = request_builder.timeout(duration);
}
let response = request_builder
.send()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
let status = response.status();
let text = response
.text()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
if !status.is_success() {
return Err(CoreError::Http {
status: status.as_u16(),
body: truncate_error_body(&text),
});
}
let response = serde_json::from_str(&text).map_err(|err| {
CoreError::InvalidResponse(format!("invalid messages response JSON: {err}"))
})?;
request.config.transform_response(&request.model, response)
}
pub(super) async fn execute_messages_provider_stream(
request: ProviderMessagesRequest,
) -> CoreResult<reqwest::Response> {
if request.provider != ANTHROPIC_MESSAGES_PROVIDER {
return Err(CoreError::InvalidRequest(
"streaming messages is not supported for this provider".to_string(),
));
}
let mut request_builder = http_client().post(&request.url).json(&request.body);
for (key, value) in &request.upstream_headers {
request_builder = request_builder.header(key, value);
}
if let Some(duration) = request.timeout {
request_builder = request_builder.timeout(duration);
}
let response = request_builder
.send()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
let status = response.status();
if !status.is_success() {
let text = response
.text()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
return Err(CoreError::Http {
status: status.as_u16(),
body: truncate_error_body(&text),
});
}
Ok(response)
}

View file

@ -1,2 +1,32 @@
//! The Anthropic Messages call, the Rust equivalent of Python's
//! `litellm.messages()`.
//!
//! [`messages`] is the top-level entrypoint: give it a model, a body, and
//! credentials, and it resolves the provider, transforms the request, calls the
//! provider, and returns a typed non-streaming response. [`messages_stream`]
//! is the streaming variant; it hands the raw upstream response back so a host
//! can splice the event stream to its own caller.
mod client;
mod common_utils;
mod handler;
mod prepare;
pub mod transformation;
pub mod types;
use crate::error::CoreResult;
use handler::{execute_messages_provider_call, execute_messages_provider_stream};
use prepare::prepare_messages_call;
use types::{AnthropicMessagesResponse, MessagesRequest};
pub async fn messages(request: MessagesRequest<'_>) -> CoreResult<AnthropicMessagesResponse> {
execute_messages_provider_call(prepare_messages_call(request)?).await
}
pub async fn messages_stream(request: MessagesRequest<'_>) -> CoreResult<reqwest::Response> {
execute_messages_provider_stream(prepare_messages_call(request)?).await
}
#[cfg(test)]
mod tests;

View file

@ -1,12 +1,9 @@
use litellm_core::CoreError;
use litellm_core::CoreResult;
use litellm_core::messages::transformation::MessagesAuthStrategy;
use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
use serde_json::Value;
use crate::error::{CoreError, CoreResult};
use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
use super::common_utils::{has_bearer_auth, has_header, messages_provider_config, string_headers};
use super::transformation::MessagesAuthStrategy;
use super::types::{MessagesRequest, ProviderMessagesRequest};
use crate::constants::BEDROCK_MESSAGES_PROVIDER;
pub(super) fn prepare_messages_call(
request: MessagesRequest<'_>,
@ -34,45 +31,18 @@ pub(super) fn prepare_messages_call(
let mut headers = string_headers(request.extra_headers)?;
let is_bedrock = provider == BEDROCK_MESSAGES_PROVIDER;
let auth_strategy = if is_bedrock {
MessagesAuthStrategy::Header("authorization")
} else {
config.auth_strategy()
};
let bearer_token = if is_bedrock {
request
.api_key
.map(str::to_string)
.or_else(|| env_lookup("AWS_BEARER_TOKEN_BEDROCK"))
.filter(|token| !token.trim().is_empty())
} else {
None
};
let auth_header = match auth_strategy {
MessagesAuthStrategy::Bearer
if has_header(&headers, "authorization")
|| (config.accepts_bearer_auth() && has_bearer_auth(&headers)) =>
{
None
}
MessagesAuthStrategy::Header(name)
if has_header(&headers, name)
|| (config.accepts_bearer_auth() && has_bearer_auth(&headers)) =>
{
None
}
MessagesAuthStrategy::Bearer => {
let api_key = config.resolve_api_key(request.api_key, &env_lookup)?;
Some(("authorization".to_string(), format!("Bearer {api_key}")))
}
MessagesAuthStrategy::Header(name) => {
let api_key = config.resolve_api_key(request.api_key, &env_lookup)?;
Some((name.to_string(), api_key))
}
};
if let Some(header) = auth_header {
headers.push(header);
let auth_strategy = config.auth_strategy();
let already_authorized = has_header(&headers, auth_strategy.header_name())
|| (config.accepts_bearer_auth() && has_bearer_auth(&headers));
if !already_authorized {
let api_key = config.resolve_api_key(request.api_key, &env_lookup)?;
let auth_header = match auth_strategy {
MessagesAuthStrategy::Bearer => {
("authorization".to_string(), format!("Bearer {api_key}"))
}
MessagesAuthStrategy::Header(name) => (name.to_string(), api_key),
};
headers.push(auth_header);
}
for (name, value) in config.default_headers() {
@ -81,9 +51,7 @@ pub(super) fn prepare_messages_call(
}
}
let _stream = request.body.get("stream").and_then(Value::as_bool) == Some(true);
let url = config.complete_url(request.api_base, &model, &env_lookup)?;
let signing_region = config.signing_region(request.api_base, &env_lookup);
let typed_request = serde_json::from_value(request.body).map_err(|err| {
CoreError::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
})?;
@ -101,8 +69,6 @@ pub(super) fn prepare_messages_call(
url,
body,
upstream_headers: headers,
signing_region,
bearer_token,
timeout: request.timeout,
})
}

View file

@ -1,14 +1,16 @@
use std::time::Duration;
use litellm_core::error::CoreError;
use serde_json::{Map, Value, json};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use crate::error::CoreError;
use super::common_utils::{
has_bearer_auth, has_header, messages_provider_config, string_headers, truncate_error_body,
};
use super::{MessagesRequest, messages};
use super::messages;
use super::types::MessagesRequest;
async fn read_http_request(socket: &mut TcpStream) -> String {
let mut request = Vec::new();
@ -152,8 +154,8 @@ async fn messages_round_trip_builds_azure_request_and_passes_response_through()
.await
.expect("messages request succeeds");
assert_eq!(response["content"][0]["text"], "hi");
assert_eq!(response["stop_reason"], "end_turn");
assert_eq!(response.content[0]["text"], "hi");
assert_eq!(response.stop_reason.as_deref(), Some("end_turn"));
let request = server.await.expect("server task completes");
let (head, body) = request.split_once("\r\n\r\n").expect("has body");
@ -208,8 +210,8 @@ async fn messages_round_trip_builds_native_anthropic_request() {
.await
.expect("messages request succeeds");
assert_eq!(response["content"][0]["text"], "hi");
assert_eq!(response["stop_reason"], "end_turn");
assert_eq!(response.content[0]["text"], "hi");
assert_eq!(response.stop_reason.as_deref(), Some("end_turn"));
let request = server.await.expect("server task completes");
let (head, _) = request.split_once("\r\n\r\n").expect("has body");

View file

@ -1,6 +1,30 @@
use std::time::Duration;
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use super::transformation::AnthropicMessagesProviderConfig;
pub struct MessagesRequest<'a> {
pub model: &'a str,
pub body: Value,
pub api_key: Option<&'a str>,
pub api_base: Option<&'a str>,
pub custom_llm_provider: Option<&'a str>,
pub extra_headers: Option<Map<String, Value>>,
pub timeout: Option<Duration>,
}
pub(super) struct ProviderMessagesRequest {
pub(super) provider: String,
pub(super) model: String,
pub(super) config: &'static dyn AnthropicMessagesProviderConfig,
pub(super) url: String,
pub(super) body: Value,
pub(super) upstream_headers: Vec<(String, String)>,
pub(super) timeout: Option<Duration>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub enum SystemPrompt {

View file

@ -1,3 +1,3 @@
litellm-python-bridge is the PyO3 cdylib that exposes Rust to the litellm Python SDK — a thin adapter (Python objects → Rust calls → Python results) over litellm-ai-gateway.
litellm-python-bridge is the PyO3 cdylib that exposes Rust to the litellm Python SDK — a thin adapter (Python objects → Rust calls → Python results) over the litellm-core route entrypoints (e.g. `litellm_core::messages::messages`).
Keep it thin: no business logic, no transforms, no I/O orchestration — just marshal in/out and call into litellm-ai-gateway.
Keep it thin: no business logic, no transforms, no I/O orchestration — just marshal in/out and call the core entrypoint.

View file

@ -11,11 +11,11 @@ Python-compatible dictionaries.
## Bridge Shape
- Prefer one stable method per top-level LiteLLM route, for example
`ocr(payload)`.
`messages(...)`, calling the matching `litellm-core` entrypoint.
- Do not add one exported PyO3 function per provider helper unless there is a
measured reason.
- Provider dispatch belongs in Rust route modules such as
`litellm_providers::ocr`, not in this PyO3 crate.
- Provider dispatch belongs in the `litellm-core` route module (e.g.
`litellm_core::messages`), not in this PyO3 crate.
- Python owns rollout state and fallback. Rust should return errors; Python
decides whether to raise or fall back. For a rust-only provider/route (no
Python reference), the Python side is a thin dispatch that calls Rust and

View file

@ -4,10 +4,11 @@ use std::time::Duration;
use litellm_ai_gateway::io::audio_transcription::{
AudioTranscriptionRequest, audio_transcription as run_audio_transcription,
};
use litellm_ai_gateway::io::messages::{MessagesRequest, messages as run_messages};
use litellm_ai_gateway::io::ocr::{OcrRequest, ocr as run_ocr};
use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection;
use litellm_core::error::CoreError;
use litellm_core::messages::messages as run_messages;
use litellm_core::messages::types::{AnthropicMessagesResponse, MessagesRequest};
use pyo3::exceptions::{PyRuntimeError, PyValueError};
use pyo3::prelude::*;
use pyo3::types::{PyAny, PyDict};
@ -35,6 +36,15 @@ fn json_to_py(py: Python<'_>, value: Value) -> PyResult<Py<PyAny>> {
Ok(json.call_method1("loads", (encoded,))?.unbind())
}
fn messages_response_to_py(
py: Python<'_>,
response: AnthropicMessagesResponse,
) -> PyResult<Py<PyAny>> {
let value =
serde_json::to_value(response).map_err(|err| PyValueError::new_err(err.to_string()))?;
json_to_py(py, value)
}
fn core_error_to_pyerr(err: CoreError) -> PyErr {
match err {
CoreError::Auth(message) => PyValueError::new_err(message),
@ -382,7 +392,7 @@ fn messages(
});
match result {
Ok(value) => json_to_py(py, value),
Ok(response) => messages_response_to_py(py, response),
Err(err) => Err(core_error_to_pyerr(err)),
}
}
@ -404,7 +414,7 @@ fn amessages(
marshal_messages_inputs(py, body, extra_headers, timeout_seconds)?;
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let value = run_messages(MessagesRequest {
let response = run_messages(MessagesRequest {
model: &model,
body,
api_key: api_key.as_deref(),
@ -416,7 +426,7 @@ fn amessages(
.await
.map_err(core_error_to_pyerr)?;
Python::attach(|py| json_to_py(py, value))
Python::attach(|py| messages_response_to_py(py, response))
})
}

View file

@ -27,6 +27,7 @@ from litellm.integrations.opentelemetry_utils.gen_ai_semconv import (
OTELSemconvCategory,
parse_semconv_opt_in,
)
from litellm.integrations.otel.model.semconv import Metric
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.secret_redaction import redact_string
from litellm.secret_managers.main import get_secret_bool, str_to_bool
@ -597,32 +598,32 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
meter = meter_provider.get_meter(__name__)
self._operation_duration_histogram = meter.create_histogram(
name="gen_ai.client.operation.duration", # Replace with semconv constant in otel 1.38
name=Metric.OPERATION_DURATION,
description="GenAI operation duration",
unit="s",
)
self._token_usage_histogram = meter.create_histogram(
name="gen_ai.client.token.usage", # Replace with semconv constant in otel 1.38
name=Metric.TOKEN_USAGE,
description="GenAI token usage",
unit="{token}",
)
self._cost_histogram = meter.create_histogram(
name="gen_ai.client.token.cost",
name=Metric.TOKEN_COST,
description="GenAI request cost",
unit="USD",
)
self._time_to_first_token_histogram = meter.create_histogram(
name="gen_ai.client.response.time_to_first_token",
name=Metric.TIME_TO_FIRST_TOKEN,
description="Time to first token for streaming requests",
unit="s",
)
self._time_per_output_token_histogram = meter.create_histogram(
name="gen_ai.client.response.time_per_output_token",
name=Metric.TIME_PER_OUTPUT_TOKEN,
description="Average time per output token (generation time / completion tokens)",
unit="s",
)
self._response_duration_histogram = meter.create_histogram(
name="gen_ai.client.response.duration",
name=Metric.RESPONSE_DURATION,
description="Total LLM API generation time (excludes LiteLLM overhead)",
unit="s",
)
@ -2980,10 +2981,12 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
def _get_metric_reader(self):
"""
Get the appropriate metric reader based on the configuration.
Histograms keep the SDK's default cumulative temporality: Prometheus-backed
OTLP receivers reject delta histograms and drop the whole batch, while
backends that prefer delta still accept cumulative.
"""
from opentelemetry.sdk.metrics import Histogram
from opentelemetry.sdk.metrics.export import (
AggregationTemporality,
ConsoleMetricExporter,
PeriodicExportingMetricReader,
)
@ -3014,7 +3017,6 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
exporter = OTLPMetricExporter(
endpoint=normalized_endpoint,
headers=_split_otel_headers,
preferred_temporality={Histogram: AggregationTemporality.DELTA},
)
return PeriodicExportingMetricReader(exporter, export_interval_millis=5000)
@ -3032,7 +3034,6 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
exporter = OTLPMetricExporter(
endpoint=normalized_endpoint,
headers=_split_otel_headers,
preferred_temporality={Histogram: AggregationTemporality.DELTA},
)
return PeriodicExportingMetricReader(exporter, export_interval_millis=5000)

View file

@ -257,13 +257,27 @@ class LiteLLM:
class Metric:
"""GenAI metric instrument names."""
"""GenAI metric instrument names.
Every name here that a convention or a backend defines uses that name, so a
consumer charting GenAI telemetry finds litellm's series where it looks for
them. ``TOKEN_USAGE``, ``OPERATION_DURATION``, ``TIME_TO_FIRST_TOKEN`` and
``TIME_PER_OUTPUT_TOKEN`` are semconv instruments, defined in the GenAI
conventions; the ``gen_ai.client.response.*`` spellings litellm used for the
latter two are not conventions at all, so nothing downstream could chart
them. Cost has no semconv instrument, so it takes ``gen_ai.usage.cost``, the
name backends already query for spend.
``RESPONSE_DURATION`` keeps its vendor spelling deliberately: the closest
convention, ``gen_ai.server.request.duration``, would collide in meaning with
``OPERATION_DURATION``, which litellm already emits for the whole operation.
"""
TOKEN_USAGE: Final = "gen_ai.client.token.usage"
OPERATION_DURATION: Final = "gen_ai.client.operation.duration"
TOKEN_COST: Final = "gen_ai.client.token.cost"
TIME_TO_FIRST_TOKEN: Final = "gen_ai.client.response.time_to_first_token"
TIME_PER_OUTPUT_TOKEN: Final = "gen_ai.client.response.time_per_output_token"
TOKEN_COST: Final = "gen_ai.usage.cost"
TIME_TO_FIRST_TOKEN: Final = "gen_ai.server.time_to_first_token"
TIME_PER_OUTPUT_TOKEN: Final = "gen_ai.server.time_per_output_token"
RESPONSE_DURATION: Final = "gen_ai.client.response.duration"

View file

@ -1,9 +1,11 @@
"""Shared, OpenTelemetry-free helpers for the otel integration.
Generic value coercion (for reading heterogeneous logging-payload dicts), time
conversion, and header parsing — pulled out of the individual modules so they
live in one place. Deliberately free of any ``opentelemetry`` import so the
OTel-free sources of truth (payloads, semconv, spans, config) can use it too.
Generic value coercion (for reading heterogeneous logging-payload dicts) and
time conversion — pulled out of the individual modules so they live in one
place. Deliberately free of any ``opentelemetry`` import so the OTel-free
sources of truth (payloads, semconv, spans, config) can use it too. OTLP header
parsing lives in :mod:`litellm.integrations.otel.plumbing.providers` instead,
because it delegates to the OTel SDK's own W3C Baggage parser.
"""
from datetime import datetime
@ -89,15 +91,3 @@ def to_seconds(value: datetime | float | int | str | None) -> float | None:
except ValueError:
continue
return None
def parse_headers(raw: str | None) -> dict[str, str]:
"""Parse an OTLP ``"k=v,k=v"`` header string into a dict."""
headers: dict[str, str] = {}
if not raw:
return headers
for pair in raw.split(","):
if "=" in pair:
key, _, value = pair.partition("=")
headers[key.strip()] = value.strip()
return headers

View file

@ -29,15 +29,13 @@ from opentelemetry.sdk.trace.export.in_memory_span_exporter import (
InMemorySpanExporter,
)
from opentelemetry.trace import Span, SpanKind, Tracer
from opentelemetry.util.re import parse_env_headers
from litellm._version import version as litellm_version
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
from litellm.integrations.otel.model.semconv import LiteLLM
from litellm.integrations.otel.model.spans import LiteLLMSpanKind
# Re-exported so ``providers.parse_headers`` remains a stable entry point.
from litellm.integrations.otel.model.utils import parse_headers as parse_headers
if TYPE_CHECKING:
from opentelemetry.metrics import Meter
from opentelemetry.sdk.metrics.export import MetricReader
@ -119,6 +117,23 @@ def _otlp_traces_endpoint(endpoint: str | None) -> str | None:
return endpoint + "/v1/traces"
def parse_headers(raw: str | None) -> dict[str, str]:
"""Parse an OTLP ``"k=v,k=v"`` header string into a dict.
``OTEL_EXPORTER_OTLP_HEADERS`` is W3C Baggage encoded per the OTLP spec, so
values are percent-decoded: a vendor that documents
``Authorization=Basic%20<token>`` (Grafana Cloud does, because a bare space
is not representable there) has to reach the exporter as ``Basic <token>``,
not with a literal ``%20`` that the backend rejects as malformed. The SDK's
own parser is used so litellm decodes exactly what the OTLP exporters do
when they read the env var themselves; ``liberal`` keeps values that are not
percent-encoded (``Authorization=Bearer <token>``) working unchanged.
"""
if not raw:
return {}
return dict(parse_env_headers(raw, liberal=True))
def _exporter_from_spec(spec: ExporterSpec) -> SpanExporter:
kind = (spec.kind or "console").lower()
factory = _EXPORTER_FACTORIES.get(kind)
@ -191,6 +206,13 @@ def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader":
``console`` (and any unrecognized kind) exports to the console; ``otlp_http``
and ``otlp_grpc`` export over OTLP with the configured endpoint/headers. The
reader exports on a 5s period, matching v1.
Histograms keep the SDK's default cumulative temporality. Prometheus-backed
OTLP receivers (Grafana Cloud / Mimir, and the Prometheus OTLP endpoint)
reject delta histograms outright with ``invalid temporality and type
combination``, which drops the whole metric batch, while backends that
prefer delta still accept cumulative. The enterprise billing exporter
already relies on the same default.
"""
from opentelemetry.sdk.metrics.export import (
ConsoleMetricExporter,
@ -202,18 +224,12 @@ def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader":
from opentelemetry.exporter.otlp.proto.http.metric_exporter import (
OTLPMetricExporter as HTTPMetricExporter,
)
from opentelemetry.sdk.metrics import Histogram
from opentelemetry.sdk.metrics.export import AggregationTemporality
exporter: Any = HTTPMetricExporter(
endpoint=_otlp_metrics_endpoint(config.endpoint),
headers=parse_headers(config.headers),
preferred_temporality={Histogram: AggregationTemporality.DELTA},
)
elif kind in ("otlp_grpc", "grpc"):
from opentelemetry.sdk.metrics import Histogram
from opentelemetry.sdk.metrics.export import AggregationTemporality
try:
from opentelemetry.exporter.otlp.proto.grpc.metric_exporter import (
OTLPMetricExporter as GRPCMetricExporter,
@ -227,7 +243,6 @@ def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader":
exporter = GRPCMetricExporter(
endpoint=config.endpoint,
headers=parse_headers(config.headers),
preferred_temporality={Histogram: AggregationTemporality.DELTA},
)
else:
exporter = ConsoleMetricExporter()

View file

@ -1267,7 +1267,7 @@ def _get_dummy_thought_signature() -> str:
def convert_to_gemini_tool_call_invoke(
message: ChatCompletionAssistantMessage,
model: Optional[str] = None,
custom_llm_provider: Optional[str] = None,
forward_function_call_id: bool = False,
) -> List[VertexPartType]:
"""
OpenAI tool invokes:
@ -1317,16 +1317,12 @@ def convert_to_gemini_tool_call_invoke(
VertexGeminiConfig,
)
forward_tool_call_id = bool(
model and VertexGeminiConfig._forward_gemini_function_call_id(model, custom_llm_provider)
)
if tool_calls is not None:
for idx, tool in enumerate(tool_calls):
if "function" in tool:
gemini_function_call: Optional[VertexFunctionCall] = _gemini_tool_call_invoke_helper(
function_call_params=tool["function"],
tool_call_id=(tool.get("id") if forward_tool_call_id else None),
tool_call_id=(tool.get("id") if forward_function_call_id else None),
)
if gemini_function_call is not None:
part_dict: VertexPartType = {"function_call": gemini_function_call}
@ -1378,8 +1374,7 @@ def convert_to_gemini_tool_call_invoke(
def convert_to_gemini_tool_call_result(
message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage],
last_message_with_tool_calls: Optional[dict],
model: Optional[str] = None,
custom_llm_provider: Optional[str] = None,
forward_function_call_id: bool = False,
) -> Union[VertexPartType, List[VertexPartType]]:
"""
OpenAI message with a tool result looks like:
@ -1501,14 +1496,8 @@ def convert_to_gemini_tool_call_result(
name = tool.get("function", {}).get("name", "")
# Echo the OpenAI tool_call_id on functionResponse (strip thought-signature suffix).
# Only Google AI Studio Gemini 3+ accepts `id` on function_response parts.
# Vertex AI and older Gemini models reject the field with HTTP 400.
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
gemini_call_id: Optional[str] = None
if model and VertexGeminiConfig._forward_gemini_function_call_id(model, custom_llm_provider):
if forward_function_call_id:
raw_tool_call_id = message.get("tool_call_id")
if raw_tool_call_id and isinstance(raw_tool_call_id, str):
stripped_id = raw_tool_call_id.split(THOUGHT_SIGNATURE_SEPARATOR, 1)[0]

View file

@ -393,24 +393,25 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
if compaction_event is not None:
return compaction_event
if self.sent_content_block_start is False:
self.sent_content_block_start = True
self.sent_content_block_finish = False
self.chunk_queue.append(
{
"type": "content_block_start",
"index": self.current_content_block_index,
"content_block": {"type": "text", "text": ""},
}
)
return self.chunk_queue.popleft()
for chunk in self.completion_stream:
if chunk == "None" or chunk is None:
raise Exception
should_start_new_block = self._should_start_new_content_block(chunk)
if should_start_new_block:
is_opening_first_block = self.sent_content_block_start is False
if is_opening_first_block and self._is_blank_delta(chunk):
continue
if is_opening_first_block:
self.sent_content_block_start = True
self.sent_content_block_finish = False
self.chunk_queue.append(
{
"type": "content_block_start",
"index": self.current_content_block_index,
"content_block": self.current_content_block_start,
}
)
elif should_start_new_block:
self._increment_content_block_index()
# applied_edits only needs to flow to the final message_delta
@ -447,7 +448,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
# ``not self.queued_usage_chunk``.
continue
if should_start_new_block and not self.sent_content_block_finish:
if should_start_new_block and not is_opening_first_block and not self.sent_content_block_finish:
# Queue the sequence: content_block_stop -> content_block_start
# -> (optionally) the trigger chunk's delta.
#
@ -615,25 +616,25 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
if compaction_event is not None:
return compaction_event
if self.sent_content_block_start is False:
self.sent_content_block_start = True
self.sent_content_block_finish = False
self.chunk_queue.append(
{
"type": "content_block_start",
"index": self.current_content_block_index,
"content_block": {"type": "text", "text": ""},
}
)
return self.chunk_queue.popleft()
async for chunk in self.completion_stream:
if chunk == "None" or chunk is None:
raise Exception
# Check if we need to start a new content block
should_start_new_block = self._should_start_new_content_block(chunk)
if should_start_new_block:
is_opening_first_block = self.sent_content_block_start is False
if is_opening_first_block and self._is_blank_delta(chunk):
continue
if is_opening_first_block:
self.sent_content_block_start = True
self.sent_content_block_finish = False
self.chunk_queue.append(
{
"type": "content_block_start",
"index": self.current_content_block_index,
"content_block": self.current_content_block_start,
}
)
elif should_start_new_block:
self._increment_content_block_index()
# applied_edits only needs to flow to the final message_delta
@ -664,7 +665,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
# Check if this processed chunk has a stop_reason - hold it for next chunk
if not self.queued_usage_chunk:
if should_start_new_block and not self.sent_content_block_finish:
if should_start_new_block and not is_opening_first_block and not self.sent_content_block_finish:
# Queue the sequence: content_block_stop -> content_block_start
# -> (optionally) the trigger chunk's delta.
#
@ -875,6 +876,22 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
return False
return bool(delta.get(_delta_payload_field(delta_type)))
@staticmethod
def _is_blank_delta(chunk: "ModelResponseStream") -> bool:
choice = chunk.choices[0]
if choice.finish_reason is not None:
return False
delta = choice.delta
if getattr(delta, "tool_calls", None):
return False
if getattr(delta, "content", None):
return False
if getattr(delta, "reasoning_content", None):
return False
if getattr(delta, "thinking_blocks", None):
return False
return True
def _should_start_new_content_block(self, chunk: "ModelResponseStream") -> bool:
"""
Determine if we should start a new content block based on the processed chunk.

View file

@ -5,7 +5,7 @@ Why separate file? Make it easy to see how transformation works
"""
import re
from typing import List, Optional, Tuple, Literal
from typing import List, Optional, Sequence, Tuple, Literal
from litellm.types.llms.openai import AllMessageValues
from litellm.types.llms.vertex_ai import CachedContentRequestBody
@ -152,6 +152,20 @@ def separate_cached_messages(
return cached_messages, non_cached_messages
def cached_messages_end_on_supported_turn(cached_messages: Sequence[AllMessageValues]) -> bool:
"""
The cachedContents API rejects contents ending on a model turn, which is how it
classifies both assistant messages and tool results, with HTTP 400
"Requests ending with a model turn are not supported". System messages are
extracted into system_instruction before contents are built, so the terminal
turn is the last non-system message.
"""
non_system_messages = tuple(message for message in cached_messages if message.get("role") != "system")
if not non_system_messages:
return bool(cached_messages)
return non_system_messages[-1].get("role") not in ("assistant", "tool", "function")
def transform_openai_messages_to_gemini_context_caching(
model: str,
messages: List[AllMessageValues],

View file

@ -22,6 +22,7 @@ from litellm.types.llms.vertex_ai import (
from ..common_utils import VertexAIError, get_vertex_base_url
from ..vertex_llm_base import VertexBase
from .transformation import (
cached_messages_end_on_supported_turn,
separate_cached_messages,
transform_openai_messages_to_gemini_context_caching,
)
@ -308,6 +309,14 @@ class ContextCachingEndpoints(VertexBase):
if len(cached_messages) == 0:
return messages, optional_params, None
if not cached_messages_end_on_supported_turn(cached_messages):
verbose_logger.debug(
"Vertex AI context caching: cached message block ends on a model turn once "
"system messages are extracted, which the cachedContents API rejects. "
"Skipping context caching."
)
return messages, optional_params, None
# Gemini requires a minimum of 1024 tokens for context caching.
# Skip caching if the cached content is too small to avoid API errors.
if not is_prompt_caching_valid_prompt(
@ -459,6 +468,14 @@ class ContextCachingEndpoints(VertexBase):
if len(cached_messages) == 0:
return messages, optional_params, None
if not cached_messages_end_on_supported_turn(cached_messages):
verbose_logger.debug(
"Vertex AI context caching: cached message block ends on a model turn once "
"system messages are extracted, which the cachedContents API rejects. "
"Skipping context caching."
)
return messages, optional_params, None
# Gemini requires a minimum of 1024 tokens for context caching.
# Skip caching if the cached content is too small to avoid API errors.
if not is_prompt_caching_valid_prompt(

View file

@ -1,7 +1,9 @@
import asyncio
import json
import os
import time
from urllib.parse import unquote
from typing import Any, Coroutine, Optional, Tuple, Union
from typing import Any, Coroutine, Mapping, Optional, Tuple, Union
import httpx
@ -10,6 +12,7 @@ from litellm.integrations.gcs_bucket.gcs_bucket_base import (
GCSBucketBase,
GCSLoggingConfig,
)
from litellm.types.utils import StandardCallbackDynamicParams
from litellm.litellm_core_utils.cloud_storage_security import (
VERTEX_AI_MANAGED_GCS_PREFIX,
should_allow_legacy_cloud_file_ids,
@ -39,6 +42,35 @@ class VertexAIFilesHandler(GCSBucketBase):
llm_provider=LlmProviders.VERTEX_AI,
)
def _resolve_read_gcs_config(
self,
litellm_params: Mapping[str, object] | None,
vertex_credentials: VERTEX_CREDENTIALS_TYPES | None,
) -> tuple[str | None, str | None]:
"""
Resolve the GCS bucket and service-account credentials for the read/content path.
Sources them from the deployment's ``litellm_params`` (``gcs_bucket_name`` /
``bucket_name`` and ``vertex_credentials``), mirroring the write path in
``VertexAIFilesConfig._get_configured_bucket_name``, and falls back to the global
``GCS_BUCKET_NAME`` / ``GCS_PATH_SERVICE_ACCOUNT`` env vars. This lets Vertex batch
run entirely at the model-group level, so output written to a per-model bucket is
readable without setting the global env vars.
"""
params: Mapping[str, object] = litellm_params or {}
bucket_candidate = params.get("gcs_bucket_name") or params.get("bucket_name")
configured_bucket_name = bucket_candidate if isinstance(bucket_candidate, str) else os.getenv("GCS_BUCKET_NAME")
credentials = params.get("vertex_credentials") or vertex_credentials
if isinstance(credentials, dict):
path_service_account: str | None = json.dumps(credentials)
elif isinstance(credentials, str):
path_service_account = credentials
else:
path_service_account = os.getenv("GCS_PATH_SERVICE_ACCOUNT")
return configured_bucket_name, path_service_account
def _extract_bucket_and_object_from_file_id(
self,
file_id: str,
@ -91,7 +123,17 @@ class VertexAIFilesHandler(GCSBucketBase):
if not file_id:
raise ValueError("file_id is required in file_content_request")
gcs_logging_config: GCSLoggingConfig = await self.get_gcs_logging_config(kwargs={})
configured_bucket_name, path_service_account = self._resolve_read_gcs_config(
litellm_params=litellm_params,
vertex_credentials=vertex_credentials,
)
dynamic_params = StandardCallbackDynamicParams(
gcs_bucket_name=configured_bucket_name,
gcs_path_service_account=path_service_account,
)
gcs_logging_config: GCSLoggingConfig = await self.get_gcs_logging_config(
kwargs={"standard_callback_dynamic_params": dynamic_params}
)
bucket_name, object_path = self._extract_bucket_and_object_from_file_id(
file_id=file_id,
configured_bucket_name=gcs_logging_config["bucket_name"],

View file

@ -661,6 +661,10 @@ def _gemini_convert_messages_with_history(
vertex_project = litellm_params.get("vertex_project") or litellm_params.get("vertex_ai_project")
vertex_credentials = litellm_params.get("vertex_credentials") or litellm_params.get("vertex_ai_credentials")
from .vertex_and_google_ai_studio_gemini import VertexGeminiConfig
forward_function_call_id = VertexGeminiConfig._forward_gemini_function_call_id(model or "")
try:
while msg_i < len(messages):
user_content: List[PartType] = []
@ -910,7 +914,7 @@ def _gemini_convert_messages_with_history(
gemini_tool_call_parts = convert_to_gemini_tool_call_invoke(
assistant_msg,
model=model,
custom_llm_provider=custom_llm_provider,
forward_function_call_id=forward_function_call_id,
)
## check if gemini_tool_call already exists in assistant_content
for gemini_tool_call_part in gemini_tool_call_parts:
@ -973,8 +977,7 @@ def _gemini_convert_messages_with_history(
_part = convert_to_gemini_tool_call_result(
messages[msg_i], # type: ignore
last_message_with_tool_calls, # type: ignore
model=model,
custom_llm_provider=custom_llm_provider,
forward_function_call_id=forward_function_call_id,
)
msg_i += 1
# Handle both single part and list of parts (for Computer Use with images)

View file

@ -289,15 +289,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return False
@staticmethod
def _forward_gemini_function_call_id(model: str, custom_llm_provider: Optional[str] = None) -> bool:
def _forward_gemini_function_call_id(model: str) -> bool:
"""
Whether to include `id` on function_call / function_response parts.
Gemini 3+ on Google AI Studio accepts (and returns) `id` for strict
tool-call matching. Vertex AI rejects the field with HTTP 400.
Gemini 3+ accepts (and returns) `id` for strict tool-call matching, on Vertex AI and
Google AI Studio alike. Older Gemini models reject the field with HTTP 400.
"""
if custom_llm_provider != "gemini":
return False
return VertexGeminiConfig._is_gemini_3_or_newer(model)
def _supports_penalty_parameters(self, model: str) -> bool:

View file

@ -13567,6 +13567,56 @@
}
]
},
"dashscope/qwen3.7-max": {
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 2.5e-06,
"litellm_provider": "dashscope",
"max_input_tokens": 991808,
"max_output_tokens": 65536,
"max_tokens": 65536,
"mode": "chat",
"output_cost_per_token": 7.5e-06,
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"dashscope/qwen3.7-plus": {
"litellm_provider": "dashscope",
"max_input_tokens": 991808,
"max_output_tokens": 65536,
"max_tokens": 65536,
"mode": "chat",
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"tiered_pricing": [
{
"cache_read_input_token_cost": 8e-08,
"input_cost_per_token": 4e-07,
"output_cost_per_token": 1.6e-06,
"range": [
0,
256000.0
]
},
{
"cache_read_input_token_cost": 2.4e-07,
"input_cost_per_token": 1.2e-06,
"output_cost_per_token": 4.8e-06,
"range": [
256000.0,
1000000.0
]
}
]
},
"dashscope/qwq-plus": {
"input_cost_per_token": 8e-07,
"litellm_provider": "dashscope",

View file

@ -200,6 +200,13 @@ _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: tuple[MCPAuth, ...] = (
)
# OAuth discovery retry cooldown for servers whose endpoints stay unresolved. The base is one
# reload cadence so a transient upstream failure recovers immediately; the cap bounds the request
# amplification and log volume of a permanently broken configuration.
_OAUTH_DISCOVERY_RETRY_BASE_SECONDS = 30.0
_OAUTH_DISCOVERY_RETRY_MAX_SECONDS = 900.0
def _blank_to_none(value: str | None) -> str | None:
"""Collapse an absent, empty, or whitespace-only string to ``None``.
@ -247,6 +254,7 @@ def _endpoints_yield_to_issuer(
authorization_url: str | None,
token_url: str | None,
registration_url: str | None,
server_ref: str,
) -> tuple[str | None, str | None, str | None]:
"""The single rule that makes an admin-configured ``issuer`` the sole authoritative endpoint
source (RFC 8414 §3.3): when it is set for a discovery auth type, the stored/manual
@ -256,9 +264,29 @@ def _endpoints_yield_to_issuer(
i.e. all ``None`` when issuer-anchored, else the inputs unchanged. Called at every resolution site
so the invariant holds in one place instead of being re-derived per merge.
"""
if issuer is not None and is_discovery_auth_type:
return None, None, None
return authorization_url, token_url, registration_url
if issuer is None or not is_discovery_auth_type:
return authorization_url, token_url, registration_url
discarded = sorted(
label
for label, value in (
("authorization_url", authorization_url),
("token_url", token_url),
("registration_url", registration_url),
)
if value
)
if discarded:
verbose_logger.warning(
"MCP server %s has a pinned Issuer, so its stored %s %s not used: an anchored issuer is the "
"sole endpoint source (RFC 8414 section 3.3) and a failed issuer fetch fails closed rather "
"than falling back to them. To use manually configured endpoints instead, clear the Issuer "
"field and re-enter the endpoint urls (clearing the Issuer also clears endpoints that may "
"have been resolved under it), or clear the Issuer alone to re-discover from the server url.",
server_ref,
", ".join(discarded),
"is" if len(discarded) == 1 else "are",
)
return None, None, None
def _normalized_authorize_endpoint(url: str) -> str:
@ -280,6 +308,68 @@ def _issuer_matches(claimed_issuer: object, configured_issuer: str) -> bool:
return _normalized_authorize_endpoint(claimed_issuer) == _normalized_authorize_endpoint(configured_issuer)
def _flow_endpoints_missing(
auth_type: MCPAuthType | None,
oauth2_flow: str | None,
authorization_url: str | None,
token_url: str | None,
token_exchange_endpoint: str | None = None,
) -> bool:
"""Whether a built server is missing an endpoint its flow needs to run at all.
Used by the reload fast-path exemption: discovery runs at build time only, and the fast path
reuses an unchanged row's registry entry verbatim, so a server whose discovery came back empty
(transient upstream failure, rate limiting) would stay broken until some unrelated config write
bumps ``updated_at``, serving its 400 the whole time. Rebuilding just these entries retries
discovery on the normal reload cadence. It costs no extra fetch for servers that resolved, and
none for those with no discovery source, since the build skips discovery for both.
"""
if auth_type == MCPAuth.oauth2_token_exchange:
# A configured exchange endpoint replaces discovery entirely; only a server that must
# discover its token endpoint and still has none is unresolved.
return token_exchange_endpoint is None and token_url is None
if auth_type not in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES:
return False
if oauth2_flow == "client_credentials":
return token_url is None
return authorization_url is None or token_url is None
def _oauth_endpoints_unresolved(server: MCPServer) -> bool:
"""``_flow_endpoints_missing`` over a built registry entry, for the reload fast-path check.
The flow comes from ``effective_oauth2_flow``, the one column-first, shape-fallback judge every
flow decision uses, not from the raw column: a legacy row the startup backfill deliberately left
unstamped (the ambiguous M2M shape) serves M2M at request time, and reading the bare column here
would classify it as interactive-missing-endpoints and re-run discovery on every reload.
"""
if (
server.auth_type == MCPAuth.oauth2_token_exchange
and server.token_exchange_profile == "entra_obo"
and not server.scopes
):
# entra_obo fails closed at exchange time without a scope (token_exchanger.py), and scopes
# can come from resource discovery, so a server that resolved its endpoints but no scopes is
# still unresolved for its flow.
return True
if server.is_dcr_bridge and not server.client_id and server.registration_url is None:
# A DCR bridge with no admin-configured client can only register callers through the
# upstream's registration endpoint, so a build that resolved the authorize and token
# endpoints but not registration_endpoint (partial metadata) is still unresolved for its
# flow and must keep retrying; without this it silently degrades to the short-circuit arm
# until an unrelated config write. Scopes are deliberately NOT part of completeness: they
# are a request hint the authorization server bounds at consent (RFC 6749 section 3.3),
# and a server without them is fully functional.
return True
return _flow_endpoints_missing(
server.auth_type,
MCPServerManager.effective_oauth2_flow(server),
server.authorization_url,
server.token_url,
server.token_exchange_endpoint,
)
def _endpoints_corroborate_authorization_url(
source_authorization_url: str | None,
trusted_authorization_url: str | None,
@ -311,11 +401,10 @@ def _carry_forward_resolved_oauth_endpoints(new_server: MCPServer, previous_serv
during re-discovery downgrades a working server (``authorization_url`` set) to a broken one
(``None``, /authorize 400s) with no configuration change. Mirrors the ``short_prefix``
carry-forward. Skipped when the server's ``url`` or ``auth_type`` changed, since the previous
endpoints may then belong to a different upstream. ``registration_url`` IS carried even though
``_persist_discovered_oauth_endpoints`` refuses to write it to the row: carrying only restores
the same in-memory value the previous build already ran with, while persisting it would flip
``_dcr_bridge_relays_client_registration`` (which keys off the stored column) for dcr_bridge
servers that never had one configured.
endpoints may then belong to a different upstream. Discovery results live only on the in-memory
registry entry; the gateway never writes them to the row, whose OAuth columns carry admin intent
alone, so this carry is the sole last-known-good mechanism and restores exactly the values the
previous build already ran with.
Carry-forward is a non-manual endpoint source, so the same trust rule as discovery applies: the
previous ``token_url``/``registration_url``/``scopes`` are carried only when the previous
@ -1182,6 +1271,40 @@ class MCPServerManager:
# empty result, or failure). Used to throttle re-probes for servers that do
# not return instructions, and to apply a short cooldown after failures.
self._upstream_initialize_instructions_probed_at: dict[str, float] = {}
# Per-server (consecutive failures, monotonic timestamp) for OAuth discovery retries, so a
# server whose endpoints never resolve backs off instead of re-running the full
# RFC 9728 -> 8414 chain, and re-logging its warning, on every reload forever.
self._oauth_discovery_retry_state: dict[
str, tuple[int, float]
] = {} # mutable-ok: retry cooldown cache, keyed per server and pruned on success
def _oauth_discovery_retry_due(self, server_id: str) -> bool:
"""Whether an unresolved server is due for another discovery attempt.
The reload fast-path exemption is what retries a failed discovery, so without a cooldown a
permanently unresolvable server re-runs the whole RFC 9728 -> RFC 8414 -> origin-fallback
chain and re-emits its unresolved-endpoints warning on every reload, per server, forever.
Delay doubles per consecutive failure from ``_OAUTH_DISCOVERY_RETRY_BASE_SECONDS`` up to
``_OAUTH_DISCOVERY_RETRY_MAX_SECONDS``, so a transient outage still recovers on the next
reload while a broken configuration settles to one attempt per cap.
"""
state = self._oauth_discovery_retry_state.get(server_id)
if state is None:
return True
failures, attempted_at = state
delay = min(
_OAUTH_DISCOVERY_RETRY_BASE_SECONDS * (2 ** max(failures - 1, 0)),
_OAUTH_DISCOVERY_RETRY_MAX_SECONDS,
)
return (time.monotonic() - attempted_at) >= delay
def _record_oauth_discovery_outcome(self, server: MCPServer) -> None:
"""Advance or clear a server's retry cooldown after a rebuild resolved it or did not."""
if not _oauth_endpoints_unresolved(server):
self._oauth_discovery_retry_state.pop(server.server_id, None)
return
failures, _ = self._oauth_discovery_retry_state.get(server.server_id, (0, 0.0))
self._oauth_discovery_retry_state[server.server_id] = (failures + 1, time.monotonic())
def _remember_upstream_initialize_instructions(self, server: MCPServer, client: MCPClient) -> None:
raw = getattr(client, "_last_initialize_instructions", None)
@ -1357,6 +1480,7 @@ class MCPServerManager:
manual_authorization_url,
manual_token_url,
manual_registration_url,
server_name or server_id,
)
should_discover = _has_oauth_discovery_source(server_url, use_issuer_anchor) and (
is_discovery_auth_type or obo_needs_discovery
@ -1834,7 +1958,6 @@ class MCPServerManager:
*,
credentials_are_encrypted: bool = True,
env_vars_are_encrypted: Optional[bool] = None,
persist_discovered_endpoints: bool = True,
) -> MCPServer:
_mcp_info: MCPInfo = mcp_server.mcp_info or {}
env_dict = _deserialize_json_dict(getattr(mcp_server, "env", None))
@ -1925,7 +2048,12 @@ class MCPServerManager:
or self._obo_needs_endpoint_discovery(auth_type, token_exchange_endpoint, manual_token_url),
)
manual_authorization_url, manual_token_url, manual_registration_url = _endpoints_yield_to_issuer(
manual_issuer, is_discovery_auth_type, manual_authorization_url, manual_token_url, manual_registration_url
manual_issuer,
is_discovery_auth_type,
manual_authorization_url,
manual_token_url,
manual_registration_url,
mcp_server.alias or mcp_server.server_name or mcp_server.server_id,
)
gated_oauth_metadata = await self._resolve_table_oauth_metadata(
mcp_server=mcp_server,
@ -2033,143 +2161,8 @@ class MCPServerManager:
max_concurrent_requests=getattr(mcp_server, "max_concurrent_requests", None),
)
_warn_internal_delegate_pkce_if_applicable(new_server, source="database")
if persist_discovered_endpoints:
await self._persist_discovered_obo_token_url(
server_id=mcp_server.server_id,
auth_type=auth_type,
existing_token_url=manual_token_url,
discovered_token_url=new_server.token_url,
)
await self._persist_discovered_oauth_endpoints(
server_id=mcp_server.server_id,
auth_type=auth_type,
existing_issuer=manual_issuer,
existing_authorization_url=manual_authorization_url,
existing_token_url=manual_token_url,
existing_scopes=scopes,
metadata=gated_oauth_metadata,
is_issuer_anchored=use_issuer_anchor,
)
return new_server
async def _persist_discovered_obo_token_url(
self,
*,
server_id: str,
auth_type: Optional[MCPAuthType],
existing_token_url: Optional[str],
discovered_token_url: Optional[str],
) -> None:
"""Write a freshly discovered OBO token endpoint back onto the DB row.
``build_mcp_server_from_table`` resolves ``token_url`` via RFC 9728 -> RFC 8414 for an
``oauth2_token_exchange`` server that has none configured, but that resolved value otherwise
lives only on the returned in-memory object; the row keeps ``token_url=None`` so every rebuild
re-runs discovery, and a transient upstream outage during a rebuild leaves the server with no
endpoint until discovery next succeeds. Persisting it makes ``_obo_needs_endpoint_discovery``
return False on the next build. Fires at most once per server (skipped once the row has a
value), and is best-effort: a write failure just means discovery runs again next time.
"""
if auth_type != MCPAuth.oauth2_token_exchange:
return
if existing_token_url or not discovered_token_url:
return
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415
if prisma_client is None:
return
try:
await MCPServerRepository(prisma_client).table.update(
where={"server_id": server_id},
data={"token_url": discovered_token_url},
)
verbose_logger.debug("Persisted discovered OBO token_url for MCP server %s", server_id)
except Exception as exc: # noqa: BLE001 - best-effort; a failed write re-discovers next build
verbose_logger.warning("Failed to persist discovered OBO token_url for MCP server %s: %s", server_id, exc)
async def _persist_discovered_oauth_endpoints(
self,
*,
server_id: str,
auth_type: MCPAuthType | None,
existing_issuer: str | None,
existing_authorization_url: str | None,
existing_token_url: str | None,
existing_scopes: list[str] | None,
metadata: MCPOAuthMetadata | None,
is_issuer_anchored: bool = False,
) -> None:
"""Write freshly discovered OAuth endpoints back onto the DB row.
Same rationale as ``_persist_discovered_obo_token_url`` but for the interactive oauth2
family: discovered ``authorization_url``/``token_url``/``scopes`` otherwise live only on
the in-memory registry entry, which is rebuilt on every client connect (the DCR reuse path
calls ``update_server``) and on every post-write DB reload, so one failed re-discovery
serves the 400 "authorization url is not configured" from /authorize until a later rebuild succeeds.
Only fills row fields that are currently empty, never persists origin-fallback guesses
(RFC 9728/8414-advertised metadata only), and deliberately skips ``registration_url``
because ``_dcr_bridge_relays_client_registration`` keys off that column. Best-effort: a
failed write re-discovers on the next build. Scopes go through ``update_mcp_server`` so
they merge into the credentials blob without touching the stored client credentials.
For an issuer-anchored server (``is_issuer_anchored``) the endpoints are re-derived from the
§3.3-validated issuer document on every build, so they are NOT persisted into the endpoint
columns: persisting them would make the next build see populated endpoints and treat them as
authoritative stored values, defeating the "endpoints come solely from the issuer" invariant.
Only the resource-driven scopes are persisted for such servers.
"""
if auth_type not in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES:
return
if metadata is None or metadata.from_origin_fallback:
return
issuer_update = (
{"issuer": metadata.discovered_issuer} if metadata.discovered_issuer and not existing_issuer else {}
)
authorization_url_update = (
{"authorization_url": metadata.authorization_url}
if metadata.authorization_url and not existing_authorization_url and not is_issuer_anchored
else {}
)
token_url_update = (
{"token_url": metadata.token_url}
if metadata.token_url and not existing_token_url and not is_issuer_anchored
else {}
)
scopes_update = {"credentials": {"scopes": metadata.scopes}} if metadata.scopes and not existing_scopes else {}
updates: dict[str, object] = {
**issuer_update,
**authorization_url_update,
**token_url_update,
**scopes_update,
}
if not updates:
return
from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # db.py imports this module at load
update_mcp_server,
)
from litellm.proxy._types import UpdateMCPServerRequest # noqa: PLC0415 # heavy module; import at call time
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime value, set after startup
if prisma_client is None:
return
try:
await update_mcp_server(
prisma_client=prisma_client,
data=UpdateMCPServerRequest.model_validate({"server_id": server_id, **updates}),
touched_by="mcp_oauth_discovery",
)
verbose_logger.info(
"Persisted discovered OAuth endpoints for MCP server %s: %s",
server_id,
sorted(updates),
)
except Exception as exc: # noqa: BLE001 - best-effort; a failed write re-discovers next build
verbose_logger.warning(
"Failed to persist discovered OAuth endpoints for MCP server %s: %s",
server_id,
exc,
)
async def _maybe_register_openapi_tools(self, server: MCPServer, *, initialize_mapping: bool = True):
"""Register OpenAPI tools if the server has a spec_path configured."""
if server.spec_path:
@ -5347,6 +5340,10 @@ class MCPServerManager:
and existing_server.updated_at is not None
and server.updated_at is not None
and existing_server.updated_at == server.updated_at
and not (
_oauth_endpoints_unresolved(existing_server)
and self._oauth_discovery_retry_due(server.server_id)
)
):
# Re-use existing server instance to avoid re-running build_mcp_server_from_table()
# which can perform network discovery for OAuth2 servers.
@ -5364,6 +5361,7 @@ class MCPServerManager:
# already-decrypted records add_server/update_server are handed.
# Decrypt them while building the registry entry.
new_server = await self.build_mcp_server_from_table(server, env_vars_are_encrypted=True)
self._record_oauth_discovery_outcome(new_server)
# Carry the cached short_prefix from the previous registry entry
# (if any) so the prefix is stable across reloads.
if existing_server is not None and existing_server.short_prefix:

View file

@ -0,0 +1,148 @@
"""One-time heal for MCP server rows whose ``issuer`` a released version wrote by itself.
Until the write was removed, OAuth discovery stamped the issuer it discovered onto the ``issuer``
column trust-on-first-use. That column means "the admin pinned this trust anchor", so the next
registry build read the gateway's own output back as admin intent: the server turned issuer-anchored
(RFC 8414 section 3.3), its stored authorization/token/registration URLs stopped applying, and a
failed issuer-document fetch left it with no authorize endpoint (GH #34985).
Deleting the write fixes every row created afterwards but cannot fix a row already stamped, which
still reads as pinned. This heals those rows by clearing the stamp so their configured endpoints
apply again.
The signal is a heuristic, and deliberately a narrow one. ``updated_by`` records only the most recent
writer, and no audit trail says which field that writer touched, so "discovery wrote this issuer" is
not directly knowable. Two independent clauses bound it, and each rules out a different way of
destroying a pin an admin meant.
Configured endpoints must be present. A deliberately pinned row very often has none, both because the
Issuer field is documented as overriding them and because ``update_mcp_server`` clears them when an
issuer changes, so "issuer set, endpoints empty" is the canonical shape of a real pin and must never
be cleared on this evidence. Skipping those rows costs little: with nothing configured to restore, the
anchored and resource-rooted paths resolve from the same upstream document, and the row still gets the
unresolved-endpoint retry and the anchored-discard warning.
The configured endpoints must also share the issuer's origin. A stamped issuer is by construction the
one self-attested by the authorization-server document discovery reached from this very server, so
endpoints typed alongside it address that same authority. An admin who pinned an issuer and typed
endpoints for a different authority is expressing an intent that clearing the issuer would discard, so
that row is warned about and never healed.
What survives both clauses is a row whose configured endpoints and stamped issuer share an origin,
which is exactly the GH #34985 shape. An admin who pinned that same origin by hand lands here too, and
for them the clear is close to a no-op: their typed endpoints keep serving and still anchor the
RFC 9700 corroboration gate, with only the stricter section 3.3 anchoring lost. Every heal logs the
cleared value so it can be restored, and the clear is recorded under this module's actor so the heal
runs at most once per row.
"""
from typing import Protocol
from urllib.parse import urlparse
from litellm._logging import verbose_proxy_logger
from litellm.proxy._experimental.mcp_server.oauth_utils import canonicalize_url_identity
from litellm.proxy.utils import PrismaClient
# The actor the removed discovery write-back stamped rows with.
_DISCOVERY_ACTOR = "mcp_oauth_discovery"
# The actor recorded on a healed row, which also makes the heal idempotent: once a row is cleared it
# no longer matches ``updated_by == _DISCOVERY_ACTOR`` and is never reconsidered.
_BACKFILL_ACTOR = "mcp_oauth_issuer_stamp_backfill"
_AUTH_TYPES_WITH_ISSUER_ANCHORING = ("oauth2", "true_passthrough", "oauth_delegate")
def _origin(url: str) -> str | None:
"""The scheme-and-authority identity of ``url``, or ``None`` when it has none.
Built on the shared URL canonicalizer so the lowercase-host and default-port rules match the
RFC 8414 issuer comparison the resolution path uses, instead of being re-derived here.
"""
parsed = urlparse(canonicalize_url_identity(url))
if not parsed.scheme or not parsed.netloc:
return None
return f"{parsed.scheme}://{parsed.netloc}"
class _MCPServerRow(Protocol):
"""The MCP server row fields this heal reads, so the untyped DB record is narrowed once here."""
server_id: str
alias: str | None
server_name: str | None
auth_type: str | None
issuer: str | None
authorization_url: str | None
token_url: str | None
registration_url: str | None
updated_by: str | None
def _is_stamped_issuer_row(row: _MCPServerRow) -> bool:
"""Whether this row carries the full signature of a gateway-written issuer stamp.
The whole rule lives here, including the writer check the query also filters on, so the decision
to clear an admin-visible field is auditable in one place rather than split between a predicate
and a query.
"""
if getattr(row, "updated_by", None) != _DISCOVERY_ACTOR:
return False
if not (getattr(row, "issuer", None) or "").strip():
return False
if getattr(row, "auth_type", None) not in _AUTH_TYPES_WITH_ISSUER_ANCHORING:
return False
configured = tuple(
value.strip()
for value in (row.authorization_url, row.token_url, row.registration_url)
if value and value.strip()
)
if not configured:
return False
issuer_origin = _origin(row.issuer or "")
return issuer_origin is not None and all(_origin(endpoint) == issuer_origin for endpoint in configured)
async def backfill_discovery_stamped_issuers(prisma_client: PrismaClient) -> int:
"""Clear gateway-written issuer stamps, returning the number of rows healed."""
candidate_rows: list[_MCPServerRow] = await prisma_client.db.litellm_mcpservertable.find_many(
where={
"updated_by": _DISCOVERY_ACTOR,
"auth_type": {"in": list(_AUTH_TYPES_WITH_ISSUER_ANCHORING)},
},
)
stamped = tuple(row for row in candidate_rows if _is_stamped_issuer_row(row))
if not stamped:
return 0
healed = 0
for row in stamped:
try:
await prisma_client.db.litellm_mcpservertable.update(
where={"server_id": row.server_id},
data={"issuer": None, "updated_by": _BACKFILL_ACTOR},
)
except Exception as exc: # noqa: BLE001 - per-row best effort; the next boot retries
verbose_proxy_logger.warning(
"MCP issuer stamp backfill: could not heal server_id=%s: %s", row.server_id, exc
)
continue
healed += 1
verbose_proxy_logger.warning(
"MCP issuer stamp backfill: cleared issuer %r on server_id=%s (alias=%s). OAuth discovery "
"had written that value onto the Issuer column, which made the server issuer-anchored and "
"fail-closed, and its configured Authorization/Token/Registration URLs were being ignored "
"as a result; those now apply again. If you pinned this issuer deliberately, set it again "
"via the dashboard or PUT /v1/mcp/server to restore RFC 8414 section 3.3 anchoring.",
row.issuer,
row.server_id,
row.alias or row.server_name,
)
if healed:
verbose_proxy_logger.warning(
"MCP issuer stamp backfill: healed %d server(s) whose Issuer had been written by OAuth "
"discovery rather than by an admin",
healed,
)
return healed

View file

@ -1169,7 +1169,6 @@ class GenerateKeyResponse(KeyRequestBase):
class UpdateKeyRequest(KeyRequestBase):
# Note: the defaults of all Params here MUST BE NONE
# else they will get overwritten
key: str # type: ignore
duration: Optional[str] = None
spend: Optional[float] = None
metadata: Optional[dict] = None
@ -1186,6 +1185,12 @@ class UpdateKeyRequest(KeyRequestBase):
raise ValueError("temp_budget_increase and temp_budget_expiry must be set together")
return self
@model_validator(mode="after")
def validate_key_identifier(self) -> "UpdateKeyRequest":
if self.key is None and self.key_alias is None:
raise ValueError("either key or key_alias must be provided")
return self
class RegenerateKeyRequest(GenerateKeyRequest):
# This needs to be different from UpdateKeyRequest, because "key" is optional for this
@ -4330,7 +4335,13 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
team_id_upsert: bool = False
team_ids_jwt_field: Optional[str] = None
upsert_sso_user_to_team: bool = False
team_allowed_routes: List[str] = ["openai_routes", "info_routes", "mcp_routes"]
team_allowed_routes: List[str] = [
"openai_routes",
"info_routes",
"mcp_routes",
"/v1/messages",
"/v1/messages/count_tokens",
]
team_id_default: Optional[str] = Field(
default=None,
description="If no team_id given, default permissions/spend-tracking to this team.s",

View file

@ -1,105 +1,61 @@
#### Analytics Endpoints #####
from datetime import datetime, timezone
from typing import List, Optional
from typing import Annotated
import fastapi
from fastapi import APIRouter, Depends, HTTPException, status
from litellm.proxy._types import *
from litellm.proxy.analytics_endpoints.cache_activity import CacheActivityResponse, get_cache_activity
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
router = APIRouter()
def _parse_date(value: str, param_name: str) -> datetime:
try:
return datetime.strptime(value, "%Y-%m-%d").replace(tzinfo=timezone.utc)
except ValueError:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": f"{param_name} must be a YYYY-MM-DD date, got {value!r}"},
)
@router.get(
"/global/activity/cache_hits",
tags=["Budget & Spend Tracking"],
dependencies=[Depends(user_api_key_auth)],
responses={
200: {"model": List[LiteLLM_SpendLogs]},
},
response_model=CacheActivityResponse,
include_in_schema=False,
)
async def get_global_activity(
start_date: Optional[str] = fastapi.Query(
default=None,
description="Time from which to start viewing spend",
),
end_date: Optional[str] = fastapi.Query(
default=None,
description="Time till which to view spend",
),
):
start_date: Annotated[str, fastapi.Query(description="Time from which to start viewing spend")],
end_date: Annotated[str, fastapi.Query(description="Time till which to view spend")],
key_aliases: Annotated[
list[str] | None, fastapi.Query(description="Only include spend from these key aliases")
] = None,
models: Annotated[list[str] | None, fastapi.Query(description="Only include spend for these models")] = None,
) -> CacheActivityResponse:
"""
Get number of cache hits, vs misses
{
"daily_data": [
const chartdata = [
{
date: 'Jan 22',
cache_hits: 10,
llm_api_calls: 2000
},
{
date: 'Jan 23',
cache_hits: 10,
llm_api_calls: 12
},
],
"sum_cache_hits": 20,
"sum_llm_api_calls": 2012
}
Cache activity for the Admin UI cache dashboard, aggregated per call_type:
cache hits vs successful LLM API requests vs failed requests, plus totals
for the stat cards and the available key-alias/model filter options.
"""
if start_date is None or end_date is None:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": "Please provide start_date and end_date"},
)
start_date_obj = datetime.strptime(start_date, "%Y-%m-%d").replace(tzinfo=timezone.utc)
end_date_obj = datetime.strptime(end_date, "%Y-%m-%d").replace(tzinfo=timezone.utc)
from litellm.proxy.proxy_server import prisma_client
try:
if prisma_client is None:
raise ValueError(
"Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys"
)
sql_query = """
SELECT
CASE
WHEN vt."key_alias" IS NOT NULL THEN vt."key_alias"
ELSE 'Unnamed Key'
END AS api_key,
sl."call_type",
sl."model",
COUNT(*) AS total_rows,
SUM(CASE WHEN sl."cache_hit" = 'True' THEN 1 ELSE 0 END) AS cache_hit_true_rows,
SUM(CASE WHEN sl."cache_hit" = 'True' THEN sl."completion_tokens" ELSE 0 END) AS cached_completion_tokens,
SUM(CASE WHEN sl."cache_hit" != 'True' THEN sl."completion_tokens" ELSE 0 END) AS generated_completion_tokens
FROM "LiteLLM_SpendLogs" sl
LEFT JOIN "LiteLLM_VerificationToken" vt ON sl."api_key" = vt."token"
WHERE
sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC')
AND sl."startTime" < (($2::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC')
GROUP BY
vt."key_alias",
sl."call_type",
sl."model"
"""
db_response = await prisma_client.db.query_raw(sql_query, start_date_obj, end_date_obj)
if db_response is None:
return []
return db_response
except Exception as e:
if prisma_client is None:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": str(e)},
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail={
"error": "Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys"
},
)
return await get_cache_activity(
prisma_client=prisma_client,
start_date=_parse_date(start_date, "start_date"),
end_date=_parse_date(end_date, "end_date"),
key_aliases=key_aliases or [],
models=models or [],
)

View file

@ -0,0 +1,137 @@
import asyncio
import json
from datetime import datetime
from typing import TYPE_CHECKING, Sequence
from pydantic import BaseModel, TypeAdapter
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient
UNKNOWN_CALL_TYPE = "Unknown"
class CacheActivityGroup(BaseModel):
call_type: str
api_requests: int
cache_hits: int
failed_requests: int
cached_completion_tokens: int
generated_completion_tokens: int
class CacheActivityTotals(BaseModel):
api_requests: int
cache_hits: int
failed_requests: int
cached_completion_tokens: int
cache_hit_ratio: float
class CacheActivityFilterOptions(BaseModel):
key_aliases: list[str]
models: list[str]
class CacheActivityResponse(BaseModel):
groups: list[CacheActivityGroup]
totals: CacheActivityTotals
filter_options: CacheActivityFilterOptions
GROUPS_SQL = """
SELECT
CASE WHEN sl."call_type" = '' THEN 'Unknown' ELSE sl."call_type" END AS call_type,
(COUNT(*)
- SUM(CASE WHEN COALESCE(sl."cache_hit", '') = 'True' THEN 1 ELSE 0 END)
- SUM(CASE WHEN sl."status" = 'failure' THEN 1 ELSE 0 END))::int AS api_requests,
SUM(CASE WHEN COALESCE(sl."cache_hit", '') = 'True' THEN 1 ELSE 0 END)::int AS cache_hits,
SUM(CASE WHEN sl."status" = 'failure' THEN 1 ELSE 0 END)::int AS failed_requests,
SUM(CASE WHEN COALESCE(sl."cache_hit", '') = 'True' THEN sl."completion_tokens" ELSE 0 END)::int
AS cached_completion_tokens,
SUM(CASE WHEN COALESCE(sl."cache_hit", '') != 'True' THEN sl."completion_tokens" ELSE 0 END)::int
AS generated_completion_tokens
FROM "LiteLLM_SpendLogs" sl
LEFT JOIN "LiteLLM_VerificationToken" vt ON sl."api_key" = vt."token"
WHERE
sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC')
AND sl."startTime" < (($2::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC')
AND ($3::jsonb = '[]'::jsonb
OR COALESCE(vt."key_alias", 'Unnamed Key') IN (SELECT jsonb_array_elements_text($3::jsonb)))
AND ($4::jsonb = '[]'::jsonb
OR sl."model" IN (SELECT jsonb_array_elements_text($4::jsonb)))
GROUP BY 1
ORDER BY (COUNT(*)) DESC
"""
KEY_ALIAS_OPTIONS_SQL = """
SELECT DISTINCT COALESCE(vt."key_alias", 'Unnamed Key') AS key_alias
FROM "LiteLLM_SpendLogs" sl
LEFT JOIN "LiteLLM_VerificationToken" vt ON sl."api_key" = vt."token"
WHERE
sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC')
AND sl."startTime" < (($2::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC')
ORDER BY 1
"""
MODEL_OPTIONS_SQL = """
SELECT DISTINCT sl."model" AS model
FROM "LiteLLM_SpendLogs" sl
WHERE
sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC')
AND sl."startTime" < (($2::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC')
AND sl."model" != ''
ORDER BY 1
"""
class _KeyAliasRow(BaseModel):
key_alias: str
class _ModelRow(BaseModel):
model: str
_groups_adapter = TypeAdapter(list[CacheActivityGroup])
_key_alias_rows_adapter = TypeAdapter(list[_KeyAliasRow])
_model_rows_adapter = TypeAdapter(list[_ModelRow])
def compute_totals(groups: Sequence[CacheActivityGroup]) -> CacheActivityTotals:
api_requests = sum(group.api_requests for group in groups)
cache_hits = sum(group.cache_hits for group in groups)
failed_requests = sum(group.failed_requests for group in groups)
all_requests = api_requests + cache_hits + failed_requests
return CacheActivityTotals(
api_requests=api_requests,
cache_hits=cache_hits,
failed_requests=failed_requests,
cached_completion_tokens=sum(group.cached_completion_tokens for group in groups),
cache_hit_ratio=(cache_hits / all_requests) * 100 if all_requests > 0 else 0.0,
)
async def get_cache_activity(
prisma_client: "PrismaClient",
start_date: datetime,
end_date: datetime,
key_aliases: Sequence[str],
models: Sequence[str],
) -> CacheActivityResponse:
group_rows, key_alias_rows, model_rows = await asyncio.gather(
prisma_client.db.query_raw(
GROUPS_SQL, start_date, end_date, json.dumps(list(key_aliases)), json.dumps(list(models))
),
prisma_client.db.query_raw(KEY_ALIAS_OPTIONS_SQL, start_date, end_date),
prisma_client.db.query_raw(MODEL_OPTIONS_SQL, start_date, end_date),
)
groups = _groups_adapter.validate_python(group_rows or [])
return CacheActivityResponse(
groups=groups,
totals=compute_totals(groups),
filter_options=CacheActivityFilterOptions(
key_aliases=[row.key_alias for row in _key_alias_rows_adapter.validate_python(key_alias_rows or [])],
models=[row.model for row in _model_rows_adapter.validate_python(model_rows or [])],
),
)

View file

@ -953,7 +953,7 @@ async def get_default_end_user_budget(
)
return None
_budget_obj = LiteLLM_BudgetTable(**budget_record.dict())
_budget_obj = LiteLLM_BudgetTable.model_validate(budget_record.dict())
# Cache the budget for 60 seconds
await user_api_key_cache.async_set_cache(
key=cache_key,
@ -999,7 +999,7 @@ async def get_team_member_default_budget(
if isinstance(cached_budget, LiteLLM_BudgetTable):
return cached_budget
if isinstance(cached_budget, dict):
return LiteLLM_BudgetTable(**cached_budget)
return LiteLLM_BudgetTable.model_validate(cached_budget)
try:
budget_record = await BudgetRepository(prisma_client).table.find_unique(where={"budget_id": budget_id})
@ -1014,7 +1014,7 @@ async def get_team_member_default_budget(
ttl=get_management_object_ttl(user_api_key_cache),
)
return LiteLLM_BudgetTable(**budget_record.dict())
return LiteLLM_BudgetTable.model_validate(budget_record.dict())
except Exception:
verbose_proxy_logger.exception(f"Error fetching team-default member budget {budget_id}")
@ -1168,7 +1168,7 @@ async def get_end_user_object(
raise Exception
# Convert to LiteLLM_EndUserTable object
_response = LiteLLM_EndUserTable(**response.dict())
_response = LiteLLM_EndUserTable.model_validate(response.dict())
# Apply default budget if needed
_response = await _apply_default_budget_to_end_user(
@ -1360,7 +1360,7 @@ async def get_tag_objects_batch(
for db_tag in db_tags:
tag_name = db_tag.tag_name
cache_key = f"tag:{tag_name}"
_tag_obj = LiteLLM_TagTable(**db_tag.dict())
_tag_obj = LiteLLM_TagTable.model_validate(db_tag.dict())
await user_api_key_cache.async_set_cache(
key=cache_key,
value=_tag_obj,
@ -1453,7 +1453,7 @@ async def get_team_membership(
if response is None:
return None
_response = LiteLLM_TeamMembership(**response.dict())
_response = LiteLLM_TeamMembership.model_validate(response.dict())
await user_api_key_cache.async_set_cache(
key=_key,
value=_response,
@ -1719,13 +1719,13 @@ async def get_user_object(
if response.organization_memberships is not None and len(response.organization_memberships) > 0:
# dump each organization membership to type LiteLLM_OrganizationMembershipTable
_dumped_memberships = [
LiteLLM_OrganizationMembershipTable(**membership.model_dump())
LiteLLM_OrganizationMembershipTable.model_validate(membership.model_dump())
for membership in response.organization_memberships
if membership is not None
]
response.organization_memberships = _dumped_memberships
_response = LiteLLM_UserTable(**dict(response))
_response = LiteLLM_UserTable.model_validate(dict(response))
response_dict = _response.model_dump()
# save the user object to cache
@ -1781,9 +1781,22 @@ async def _cache_team_object(
## CACHE REFRESH TIME!
team_table.last_refreshed_at = time.time()
key = "team_id:{}".format(team_id)
if proxy_logging_obj is not None:
try:
await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=key)
except Exception as e: # noqa: BLE001 # best-effort invalidation: any cache backend error must not fail the write
verbose_proxy_logger.warning(
"Failed to invalidate internal usage cache entry %s; "
"a stale team object may be served until its TTL expires: %s",
key,
e,
)
# team_id is the table primary key — guaranteed unique, safe to write.
await _cache_management_object(
key="team_id:{}".format(team_id),
key=key,
value=team_table,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
@ -1805,9 +1818,17 @@ async def _cache_team_object(
# the cache from a verified single row.
if team_table.team_alias:
alias_key = "team_alias:{}".format(team_table.team_alias)
user_api_key_cache.delete_cache(key=alias_key)
if proxy_logging_obj is not None:
await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=alias_key)
try:
user_api_key_cache.delete_cache(key=alias_key)
if proxy_logging_obj is not None:
await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=alias_key)
except Exception as e: # noqa: BLE001 # best-effort invalidation: any cache backend error must not fail the mutation
verbose_proxy_logger.warning(
"Failed to invalidate cached team alias entry %s; "
"a stale team object may be served until its TTL expires: %s",
alias_key,
e,
)
async def _cache_key_object(
@ -1862,7 +1883,7 @@ async def _get_team_db_check(team_id: str, prisma_client: PrismaClient, team_id_
http_request=mock_request,
user_api_key_dict=system_admin_user,
)
response = LiteLLM_TeamTable(**created_team_dict)
response = LiteLLM_TeamTable.model_validate(created_team_dict)
return response
@ -1894,7 +1915,7 @@ async def _get_team_object_from_user_api_key_cache(
if response is None:
raise Exception
_response = LiteLLM_TeamTableCachedObj(**response.dict())
_response = LiteLLM_TeamTableCachedObj.model_validate(response.dict())
# Load object_permission if object_permission_id exists but object_permission is not loaded
if _response.object_permission_id and not _response.object_permission:
@ -2085,7 +2106,7 @@ async def get_access_object(
detail={"error": f"Access group doesn't exist in db. Access group={access_group_id}."},
)
_response = LiteLLM_AccessGroupTable(**response.dict())
_response = LiteLLM_AccessGroupTable.model_validate(response.dict())
# Save to cache
await _cache_access_object(
@ -2170,7 +2191,7 @@ async def get_team_object_by_alias(
)
team = teams[0]
team_obj = LiteLLM_TeamTableCachedObj(**team.model_dump())
team_obj = LiteLLM_TeamTableCachedObj.model_validate(team.model_dump())
# Load object_permission if object_permission_id exists but object_permission is not loaded
if team_obj.object_permission_id and not team_obj.object_permission:
@ -2272,7 +2293,7 @@ async def get_org_object_by_alias(
)
org = orgs[0]
org_obj = LiteLLM_OrganizationTable(**org.model_dump())
org_obj = LiteLLM_OrganizationTable.model_validate(org.model_dump())
# Cache the result
await user_api_key_cache.async_set_cache(
@ -2605,7 +2626,7 @@ async def get_object_permission(
if response is None:
return None
_perm_obj = LiteLLM_ObjectPermissionTable(**response.dict())
_perm_obj = LiteLLM_ObjectPermissionTable.model_validate(response.dict())
await user_api_key_cache.async_set_cache(
key=key,
value=_perm_obj,
@ -2665,7 +2686,7 @@ async def get_managed_vector_store_rows_by_uuids(
row_dict = dict(row) if hasattr(row, "__dict__") else {}
if not row_dict:
continue
cached_obj = LiteLLM_ManagedVectorStoresTable(**row_dict)
cached_obj = LiteLLM_ManagedVectorStoresTable.model_validate(row_dict)
key = "managed_vector_store_id:{}".format(cached_obj.vector_store_id)
await user_api_key_cache.async_set_cache(
key=key,
@ -2746,7 +2767,7 @@ async def get_org_object(
f"Organization doesn't exist in db. Organization={org_id}. Create organization via `/organization/new` call."
)
_org_obj = LiteLLM_OrganizationTable(**response.model_dump())
_org_obj = LiteLLM_OrganizationTable.model_validate(response.model_dump())
# Cache the result
await user_api_key_cache.async_set_cache(
key=cache_key,
@ -4221,7 +4242,7 @@ async def get_project_object(
if project_row is None:
return None
project_obj = LiteLLM_ProjectTableCachedObj(**project_row.model_dump())
project_obj = LiteLLM_ProjectTableCachedObj.model_validate(project_row.model_dump())
# Cache with TTL following _cache_management_object pattern
project_obj.last_refreshed_at = time.time()

View file

@ -1432,7 +1432,7 @@ def _extract_models_from_managed_resource_id(
)
_append_model_candidates(
candidates=candidates,
value=get_model_id_from_unified_batch_id(unified_file_id),
value=_resolve_model_id_with_router(get_model_id_from_unified_batch_id(unified_file_id), llm_router),
)
except Exception as e:
verbose_proxy_logger.debug("Unable to extract model from managed file/batch ID: %s", str(e))
@ -1442,7 +1442,10 @@ def _extract_models_from_managed_resource_id(
parsed_id = parse_unified_id(resource_id)
if parsed_id:
_append_model_candidates(candidates=candidates, value=parsed_id.get("model_id"))
_append_model_candidates(
candidates=candidates,
value=_resolve_model_id_with_router(parsed_id.get("model_id"), llm_router),
)
_append_model_candidates(candidates=candidates, value=parsed_id.get("target_model_names"))
except Exception as e:
verbose_proxy_logger.debug("Unable to extract model from unified managed resource ID: %s", str(e))

View file

@ -1960,17 +1960,6 @@ async def _user_api_key_auth_builder(
else:
valid_token.team_object_permission = None
# Cache under the canonical "team_id:{id}" key so get_team_object and
# _update_team_cache serve this write from the L2 cache. The guard keeps a
# non-team (personal) key, whose team_id is None, from reaching the cache
# layer, which Redis rejects with a NoneType key error.
if valid_token.team_id is not None and _team_obj is not None:
await user_api_key_cache.async_set_cache(
key=f"team_id:{valid_token.team_id}",
value=_team_obj,
model_type=LiteLLM_TeamTableCachedObj,
)
# Fetch project object if key belongs to a project
_project_obj = None
if valid_token.project_id is not None:

View file

@ -23,6 +23,7 @@ from litellm.proxy.common_utils.openai_endpoint_utils import (
)
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
apply_team_provider_credentials,
decode_model_from_file_id,
encode_batch_response_ids,
encode_file_id_with_model,
@ -295,6 +296,12 @@ async def create_batch(
verbose_proxy_logger.debug(f"Created batch using model: {model_param}")
else:
# SCENARIO 3: Fallback to custom_llm_provider (uses env variables)
apply_team_provider_credentials(
data=cast(dict, _create_batch_data), # cast-ok: TypedDict is a dict at runtime
llm_router=llm_router,
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
response = await litellm.acreate_batch(
custom_llm_provider=custom_llm_provider,
**_create_batch_data, # type: ignore
@ -525,6 +532,12 @@ async def retrieve_batch(
or get_custom_llm_provider_from_request_query(request=request)
or "openai"
)
apply_team_provider_credentials(
data=data,
llm_router=llm_router,
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
response = await litellm.aretrieve_batch(
custom_llm_provider=custom_llm_provider,
**data, # type: ignore
@ -718,6 +731,12 @@ async def list_batches(
or get_custom_llm_provider_from_request_query(request=request)
or "openai"
)
apply_team_provider_credentials(
data=data,
llm_router=llm_router,
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
response = await litellm.alist_batches(
custom_llm_provider=custom_llm_provider, # type: ignore
after=after,
@ -908,6 +927,12 @@ async def cancel_batch(
# Extract batch_id from data to avoid "multiple values for keyword argument" error
# data was cast from CancelBatchRequest which already contains batch_id
data.pop("batch_id", None)
apply_team_provider_credentials(
data=data,
llm_router=llm_router,
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
_cancel_batch_data = CancelBatchRequest(batch_id=batch_id, **data)
response = await litellm.acancel_batch(
custom_llm_provider=custom_llm_provider, # type: ignore

View file

@ -10,11 +10,32 @@ uv tool install 'litellm[proxy]'
## Configuration
The CLI can be configured using environment variables or command-line options:
The CLI can be configured using environment variables, command-line options, or a persistent config file:
- `LITELLM_PROXY_URL`: Base URL of the LiteLLM proxy server (default: http://localhost:4000)
- `LITELLM_PROXY_API_KEY`: API key for authentication
To stop exporting `LITELLM_PROXY_URL` in every shell session, store the proxy URL once in `~/.litellm/config.json`:
```bash
lite config set base_url https://your-proxy.example.com
```
Manage the stored config with:
```bash
lite config get base_url # print the stored value
lite config get # print all stored config
lite config unset base_url # remove the stored value
```
The base URL is resolved in this order of precedence:
1. `--base-url` command-line option
2. `LITELLM_PROXY_URL` environment variable
3. `base_url` from `~/.litellm/config.json`
4. `http://localhost:4000`
## Global Options
- `--version`, `-v`: Print the LiteLLM Proxy client and server version and exit.
@ -581,6 +602,8 @@ The CLI respects the following environment variables:
- `LITELLM_PROXY_URL`: Base URL of the proxy server
- `LITELLM_PROXY_API_KEY`: API key for authentication
`LITELLM_PROXY_URL` takes precedence over a `base_url` stored via `lite config set`, and the `--base-url` option overrides both. See the Configuration section for the full precedence order.
## Examples
1. List all models in table format:

View file

@ -15,6 +15,8 @@ from rich.table import Table
from litellm.constants import CLI_JWT_EXPIRATION_HOURS
from litellm.litellm_core_utils.cli_token_utils import is_cli_token_fresh
from .private_json import write_private_json
# Token storage utilities
def get_token_file_path() -> str:
@ -27,11 +29,7 @@ def get_token_file_path() -> str:
def save_token(token_data: Dict[str, Any]) -> None:
"""Save token data to file"""
token_file = get_token_file_path()
with open(token_file, "w") as f:
json.dump(token_data, f, indent=2)
# Set file permissions to be readable only by owner
os.chmod(token_file, 0o600)
write_private_json(get_token_file_path(), token_data)
def load_token() -> Optional[Dict[str, Any]]:

View file

@ -0,0 +1,108 @@
import json
import os
import sys
from collections.abc import Mapping
from pathlib import Path
from urllib.parse import urlparse
import click
from pydantic import TypeAdapter
from .private_json import write_private_json
ALLOWED_CONFIG_KEYS: tuple[str, ...] = ("base_url",)
_config_adapter: TypeAdapter[Mapping[str, str]] = TypeAdapter(Mapping[str, str])
def get_config_file_path() -> str:
"""Get the path to the persistent CLI config file"""
home_dir = Path.home()
config_dir = home_dir / ".litellm"
return str(config_dir / "config.json")
def load_config() -> Mapping[str, str]:
"""Load CLI config from file; returns {} if missing or unreadable"""
try:
config_file = get_config_file_path()
except RuntimeError:
return {}
if not os.path.exists(config_file):
return {}
try:
with open(config_file, "r") as f:
return _config_adapter.validate_python(json.load(f))
except (OSError, ValueError) as e:
click.echo(f"Warning: ignoring invalid config file {config_file}: {e}", err=True)
return {}
def save_config(config: Mapping[str, str]) -> None:
"""Save CLI config to file"""
write_private_json(get_config_file_path(), config)
def get_config_value(key: str) -> str | None:
"""Get a single value from the persistent CLI config"""
return load_config().get(key)
@click.group(name="config")
def config_commands() -> None:
"""Manage persistent CLI configuration (~/.litellm/config.json)"""
@config_commands.command(name="set")
@click.argument("key")
@click.argument("value")
def set_config(key: str, value: str) -> None:
"""Set a config KEY to VALUE (e.g. `lite config set base_url https://your-proxy.example.com`)"""
if key not in ALLOWED_CONFIG_KEYS:
raise click.UsageError(f"Unknown config key '{key}'. Allowed keys: {', '.join(ALLOWED_CONFIG_KEYS)}")
if key == "base_url":
parsed = urlparse(value)
if parsed.scheme not in ("http", "https") or not parsed.netloc:
raise click.UsageError("base_url must be a full http:// or https:// URL including a host")
if "?" in value or "#" in value:
raise click.UsageError("base_url must not include a query string or fragment")
normalized_value = value.rstrip("/")
save_config({**load_config(), key: normalized_value})
click.echo(f"Set {key} = {normalized_value} in {get_config_file_path()}")
@config_commands.command(name="get")
@click.argument("key", required=False)
def get_config(key: str | None) -> None:
"""Print the value of KEY, or all stored config when KEY is omitted"""
config = load_config()
if key is not None:
value = config.get(key)
if value is None:
click.echo(f"{key} is not set", err=True)
sys.exit(1)
click.echo(value)
return
if not config:
click.echo("(no config set)")
return
for entry_key, entry_value in config.items():
click.echo(f"{entry_key} = {entry_value}")
@config_commands.command(name="unset")
@click.argument("key")
def unset_config(key: str) -> None:
"""Remove KEY from the config file"""
config = load_config()
if key not in config:
click.echo(f"{key} was not set")
return
save_config({k: v for k, v in config.items() if k != key})
click.echo(f"Removed {key} from {get_config_file_path()}")

View file

@ -0,0 +1,20 @@
import json
import os
import tempfile
from collections.abc import Mapping
from pathlib import Path
def write_private_json(path: str, data: Mapping[str, object]) -> None:
"""Atomically write JSON to path with owner-only permissions (0600)"""
parent = Path(path).parent
parent.mkdir(parents=True, exist_ok=True)
fd, tmp_path = tempfile.mkstemp(dir=str(parent), prefix=".tmp-", suffix=".json")
try:
with os.fdopen(fd, "w") as f:
json.dump(data, f, indent=2)
f.flush()
os.fsync(f.fileno())
os.replace(tmp_path, path)
finally:
Path(tmp_path).unlink(missing_ok=True)

View file

@ -11,6 +11,7 @@ from .commands.agents import agent_commands
from .commands.auth import auth_group, get_stored_api_key, login, logout, whoami
from .commands.autoroute.commands import autoroute_group
from .commands.chat import chat
from .commands.config import config_commands, get_config_value
from .commands.credentials import credentials
from .commands.encryption import encryption
from .commands.http import http
@ -45,27 +46,16 @@ def print_version(base_url: str, api_key: Optional[str]):
@click.option(
"--version",
"-v",
"show_version",
is_flag=True,
is_eager=True,
expose_value=False,
help="Show the LiteLLM Proxy CLI and server version and exit.",
callback=lambda ctx, param, value: (
(
print_version(
ctx.params.get("base_url") or "http://localhost:4000",
ctx.params.get("api_key"),
)
or ctx.exit()
)
if value and not ctx.resilient_parsing
else None
),
)
@click.option(
"--base-url",
envvar="LITELLM_PROXY_URL",
show_envvar=True,
default="http://localhost:4000",
default=None,
show_default="base_url from `lite config`, else http://localhost:4000",
help="Base URL of the LiteLLM proxy server",
)
@click.option(
@ -75,13 +65,16 @@ def print_version(base_url: str, api_key: Optional[str]):
help="API key for authentication",
)
@click.pass_context
def cli(ctx: click.Context, base_url: str, api_key: Optional[str]) -> None:
def cli(ctx: click.Context, show_version: bool, base_url: str | None, api_key: Optional[str]) -> None:
"""LiteLLM Proxy CLI - Manage your LiteLLM proxy server"""
ctx.ensure_object(dict)
stored_base_url = get_config_value("base_url")
base_url_provided = base_url is not None
# Normalize once here so every downstream command (login, agents, http, ...) can safely
# do f"{base_url}/some/path" without producing a double slash.
base_url = base_url.rstrip("/")
base_url = ((stored_base_url or "http://localhost:4000") if base_url is None else base_url).rstrip("/")
# If no API key provided via flag or environment variable, try to load from saved token.
# Pass base_url so we only use the stored key when it was issued for this server.
@ -94,8 +87,13 @@ def cli(ctx: click.Context, base_url: str, api_key: Optional[str]) -> None:
# apiKeyHelper is invoked bare (no flags) -- commands that must work
# unattended (print-token) need to tell "user didn't say" apart from
# "user said localhost:4000 on purpose" so they can fall back to
# whatever server the stored token was actually issued for.
ctx.obj["base_url_explicit"] = ctx.get_parameter_source("base_url") != click.core.ParameterSource.DEFAULT
# whatever server the stored token was actually issued for. A base_url
# saved via `lite config set` counts as the user saying it.
ctx.obj["base_url_explicit"] = base_url_provided or bool(stored_base_url)
if show_version:
print_version(base_url, api_key)
ctx.exit()
# If no subcommand was invoked, start interactive mode
if ctx.invoked_subcommand is None:
@ -141,6 +139,7 @@ cli.add_command(down)
cli.add_command(model_groups)
# Add the autoroute command group (QA auto-routing against your real proxy)
cli.add_command(autoroute_group, name="autoroute")
cli.add_command(config_commands)
if __name__ == "__main__":

View file

@ -1829,6 +1829,15 @@ async def add_litellm_data_to_request(
return data
def _warn_stale_team_alias_once(warning_key: str, message: str, *args: str) -> None:
if warning_key in _STALE_TEAM_ALIAS_WARNING_KEYS:
return
_STALE_TEAM_ALIAS_WARNING_KEYS[warning_key] = None
while len(_STALE_TEAM_ALIAS_WARNING_KEYS) > _MAX_STALE_ALIAS_WARNING_KEYS:
_STALE_TEAM_ALIAS_WARNING_KEYS.popitem(last=False)
verbose_proxy_logger.warning(message, *args)
def _update_model_if_team_alias_exists(
data: dict,
user_api_key_dict: UserAPIKeyAuth,
@ -1848,49 +1857,63 @@ def _update_model_if_team_alias_exists(
Note: model_aliases for team models are deprecated. This function only applies
to legacy non-team-scoped aliases. Team-scoped deployments use team_public_model_name
and are resolved via map_team_model in route_llm_request.
An alias that targets a team-scoped internal name (``model_name_{team_id}_{uuid}``)
with no live deployment behind it is never applied: the deployment was deleted, so
the rewrite could only fail with an error naming a model the caller never sent.
Keeping the requested model name lets it resolve against the deployments that still
exist (e.g. a gateway-level model group shared with the team).
"""
_model = data.get("model")
if _model and user_api_key_dict.team_model_aliases and _model in user_api_key_dict.team_model_aliases:
from litellm.proxy.proxy_server import llm_router
if not _model or not user_api_key_dict.team_model_aliases or _model not in user_api_key_dict.team_model_aliases:
return
# Skip alias rewrite if this model resolves to team-specific deployments
# (team models use team_public_model_name, not model_aliases)
aliased_target = user_api_key_dict.team_model_aliases[_model]
from litellm.proxy.proxy_server import llm_router
# Optional bypass for stale aliases from pre-PR deployments:
# only enabled via feature flag to preserve backwards compatibility.
# Cached at module level to avoid hot-path secret lookups on every request.
global _ENABLE_TEAM_STALE_ALIAS_BYPASS
if _ENABLE_TEAM_STALE_ALIAS_BYPASS is None:
_ENABLE_TEAM_STALE_ALIAS_BYPASS = get_secret_bool("LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS", False)
enable_stale_alias_bypass = _ENABLE_TEAM_STALE_ALIAS_BYPASS
# Check if the alias points to a team-scoped UUID name
# (format: "model_name_{team_id}_{uuid}")
is_stale_team_alias = aliased_target.startswith(f"model_name_{user_api_key_dict.team_id}_")
if is_stale_team_alias and llm_router:
# This is a stale alias from pre-PR deployments.
# Check if current team deployments exist for the public name.
key = (user_api_key_dict.team_id, _model)
if key in llm_router.team_model_to_deployment_indices:
if enable_stale_alias_bypass:
# Team deployments exist; skip stale alias
return
warning_key = f"{user_api_key_dict.team_id}:{_model}:{aliased_target}"
if warning_key not in _STALE_TEAM_ALIAS_WARNING_KEYS:
_STALE_TEAM_ALIAS_WARNING_KEYS[warning_key] = None
while len(_STALE_TEAM_ALIAS_WARNING_KEYS) > _MAX_STALE_ALIAS_WARNING_KEYS:
_STALE_TEAM_ALIAS_WARNING_KEYS.popitem(last=False)
verbose_proxy_logger.warning(
"Stale team model alias detected for model='%s', team_id='%s'. "
"New sibling deployments may be unreachable. "
"Set LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS=true to enable "
"team-scoped sibling routing.",
_sanitize_for_log(_model),
user_api_key_dict.team_id,
)
# Skip alias rewrite if this model resolves to team-specific deployments
# (team models use team_public_model_name, not model_aliases)
aliased_target = user_api_key_dict.team_model_aliases[_model]
data["model"] = aliased_target
return
# Optional bypass for stale aliases from pre-PR deployments:
# only enabled via feature flag to preserve backwards compatibility.
# Cached at module level to avoid hot-path secret lookups on every request.
global _ENABLE_TEAM_STALE_ALIAS_BYPASS
if _ENABLE_TEAM_STALE_ALIAS_BYPASS is None:
_ENABLE_TEAM_STALE_ALIAS_BYPASS = get_secret_bool("LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS", False)
enable_stale_alias_bypass = _ENABLE_TEAM_STALE_ALIAS_BYPASS
# Check if the alias points to a team-scoped UUID name
# (format: "model_name_{team_id}_{uuid}")
is_stale_team_alias = aliased_target.startswith(f"model_name_{user_api_key_dict.team_id}_")
if is_stale_team_alias and llm_router:
if aliased_target not in llm_router.model_name_to_deployment_indices:
_warn_stale_team_alias_once(
f"deleted:{user_api_key_dict.team_id}:{_model}:{aliased_target}",
"Team model alias for model='%s', team_id='%s' targets '%s', which has no live "
"deployment. Routing with the requested model name instead; remove the stale "
"entry from the team's model_aliases to silence this warning.",
_sanitize_for_log(_model),
_sanitize_for_log(user_api_key_dict.team_id),
_sanitize_for_log(aliased_target),
)
return
# This is a stale alias from pre-PR deployments.
# Check if current team deployments exist for the public name.
key = (user_api_key_dict.team_id, _model)
if key in llm_router.team_model_to_deployment_indices:
if enable_stale_alias_bypass:
# Team deployments exist; skip stale alias
return
_warn_stale_team_alias_once(
f"{user_api_key_dict.team_id}:{_model}:{aliased_target}",
"Stale team model alias detected for model='%s', team_id='%s'. "
"New sibling deployments may be unreachable. "
"Set LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS=true to enable "
"team-scoped sibling routing.",
_sanitize_for_log(_model),
_sanitize_for_log(user_api_key_dict.team_id),
)
data["model"] = aliased_target
def _update_model_if_key_alias_exists(

View file

@ -233,27 +233,28 @@ def update_breakdown_metrics(
)
# Update model group breakdown
if record.model_group and record.model_group not in breakdown.model_groups:
breakdown.model_groups[record.model_group] = MetricWithMetadata(
model_group_key = record.model_group or record.model
if model_group_key and model_group_key not in breakdown.model_groups:
breakdown.model_groups[model_group_key] = MetricWithMetadata(
metrics=SpendMetrics(),
metadata=model_metadata.get(record.model_group, {}),
metadata=model_metadata.get(model_group_key, {}),
)
if record.model_group:
breakdown.model_groups[record.model_group].metrics = update_metrics(
breakdown.model_groups[record.model_group].metrics, record
if model_group_key:
breakdown.model_groups[model_group_key].metrics = update_metrics(
breakdown.model_groups[model_group_key].metrics, record
)
# Update API key breakdown for this model
if record.api_key not in breakdown.model_groups[record.model_group].api_key_breakdown:
breakdown.model_groups[record.model_group].api_key_breakdown[record.api_key] = KeyMetricWithMetadata(
if record.api_key not in breakdown.model_groups[model_group_key].api_key_breakdown:
breakdown.model_groups[model_group_key].api_key_breakdown[record.api_key] = KeyMetricWithMetadata(
metrics=SpendMetrics(),
metadata=KeyMetadata(
key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None),
team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None),
),
)
breakdown.model_groups[record.model_group].api_key_breakdown[record.api_key].metrics = update_metrics(
breakdown.model_groups[record.model_group].api_key_breakdown[record.api_key].metrics,
breakdown.model_groups[model_group_key].api_key_breakdown[record.api_key].metrics = update_metrics(
breakdown.model_groups[model_group_key].api_key_breakdown[record.api_key].metrics,
record,
)
@ -574,11 +575,11 @@ def _build_aggregated_sql_query(
date,
api_key,
model,
model_group,
COALESCE(NULLIF(model_group, ''), model) AS model_group,
custom_llm_provider,
mcp_namespaced_tool_name,
endpoint,
GROUPING(date, api_key, model, model_group,
GROUPING(date, api_key, model, COALESCE(NULLIF(model_group, ''), model),
custom_llm_provider, mcp_namespaced_tool_name,
endpoint) AS group_level,
SUM(spend)::float AS spend,
@ -599,8 +600,8 @@ def _build_aggregated_sql_query(
(date, api_key),
(date, model),
(date, model, api_key),
(date, model_group),
(date, model_group, api_key),
(date, COALESCE(NULLIF(model_group, ''), model)),
(date, COALESCE(NULLIF(model_group, ''), model), api_key),
(date, custom_llm_provider),
(date, custom_llm_provider, api_key),
(date, mcp_namespaced_tool_name),

View file

@ -513,7 +513,7 @@ async def new_user(
response_dict["key"] = response.get("token", "")
new_user_response = NewUserResponse(**response_dict)
new_user_response = NewUserResponse.model_validate(response_dict)
#########################################################
########## USER CREATED HOOK ################
@ -879,7 +879,7 @@ async def _check_user_info_v2_access(
# Get all teams the caller belongs to
teams = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": caller_user.teams}})
for team in teams:
team_obj = LiteLLM_TeamTable(**team.model_dump())
team_obj = LiteLLM_TeamTable.model_validate(team.model_dump())
if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj):
# Check if target user is in this team
if team.team_id in (target_user.teams or []):
@ -1013,11 +1013,11 @@ async def _get_user_info_for_proxy_admin(user_api_key_dict: UserAPIKeyAuth):
for key in _keys_in_db:
if key.get("models") is None:
key["models"] = []
keys_in_db.append(LiteLLM_VerificationToken(**key))
keys_in_db.append(LiteLLM_VerificationToken.model_validate(key))
# cast all teams to LiteLLM_TeamTable
_teams_in_db: list = results[0]["teams"] or []
_teams_in_db = [LiteLLM_TeamTable(**team) for team in _teams_in_db]
_teams_in_db = [LiteLLM_TeamTable.model_validate(team) for team in _teams_in_db]
_teams_in_db.sort(key=lambda x: getattr(x, "team_alias", "") or "")
returned_keys = _process_keys_for_user_info(keys=keys_in_db, all_teams=_teams_in_db)
@ -1146,7 +1146,7 @@ async def _schedule_user_update_audit_log(
try:
updated_user_row = await UserRepository(prisma_client).table.find_first(where={"user_id": response["user_id"]})
if updated_user_row:
user_row_typed = LiteLLM_UserTable(**updated_user_row.model_dump(exclude_none=True))
user_row_typed = LiteLLM_UserTable.model_validate(updated_user_row.model_dump(exclude_none=True))
asyncio.create_task(
UserManagementEventHooks.create_internal_user_audit_log(
user_id=user_row_typed.user_id,
@ -1172,7 +1172,7 @@ def _check_user_update_authz(
raise HTTPException(status_code=403, detail="Only proxy admins can modify user roles.")
if existing_user_row is not None:
typed_row = LiteLLM_UserTable(**existing_user_row.model_dump(exclude_none=True))
typed_row = LiteLLM_UserTable.model_validate(existing_user_row.model_dump(exclude_none=True))
if not can_user_call_user_update(user_api_key_dict=user_api_key_dict, user_info=typed_row):
raise HTTPException(
status_code=403,
@ -1248,7 +1248,7 @@ async def _update_single_user_helper(
_check_user_update_authz(user_request, user_api_key_dict, existing_user_row)
if existing_user_row is not None:
existing_user_row = LiteLLM_UserTable(**existing_user_row.model_dump(exclude_none=True))
existing_user_row = LiteLLM_UserTable.model_validate(existing_user_row.model_dump(exclude_none=True))
# Prevent budget self-escalation (GHSA-wvg4-6222-3q4r): non-admin callers
# must not be able to raise their own budget/spend fields.
@ -1998,7 +1998,11 @@ async def get_users(
for user in users:
user_dump = user.model_dump()
user_dump["metadata"] = _redact_scim_enterprise_metadata(user_dump.get("metadata"))
user_list.append(LiteLLM_UserTableWithKeyCount(**user_dump, key_count=user_key_counts.get(user.user_id, 0)))
user_list.append(
LiteLLM_UserTableWithKeyCount.model_validate(
{**user_dump, "key_count": user_key_counts.get(user.user_id, 0)}
)
)
else:
user_list = []
@ -2157,7 +2161,7 @@ async def delete_user(
teams_to_update = []
for team in fetch_all_teams:
is_member_in_team, new_team_members = _cleanup_members_with_roles(
existing_team_row=LiteLLM_TeamTable(**team.model_dump()),
existing_team_row=LiteLLM_TeamTable.model_validate(team.model_dump()),
data=TeamMemberDeleteRequest(
team_id=team.team_id,
user_id=user_row.user_id,
@ -2438,7 +2442,7 @@ async def ui_view_users(
if not users:
return []
return [LiteLLM_UserTableFiltered(**user.model_dump()) for user in users]
return [LiteLLM_UserTableFiltered.model_validate(user.model_dump()) for user in users]
except HTTPException:
raise

View file

@ -1078,7 +1078,7 @@ async def _common_key_generation_helper(
response["soft_budget"] = data.soft_budget # include the user-input soft budget in the response
response = GenerateKeyResponse(**response)
response = GenerateKeyResponse.model_validate(response)
response.token = response.token_id # remap token to use the hash, and leave the key in the `key` field [TODO]: clean up generate_key_helper_fn to do this
@ -2023,7 +2023,7 @@ def _validate_max_budget(max_budget: Optional[float]) -> None:
async def _get_and_validate_existing_key(
token: str, prisma_client: Optional[PrismaClient]
token: str | None, prisma_client: Optional[PrismaClient], key_alias: str | None = None
) -> LiteLLM_VerificationToken:
"""
Get existing key from database and validate it exists.
@ -2031,12 +2031,13 @@ async def _get_and_validate_existing_key(
Args:
token: The key token to look up
prisma_client: Prisma client instance
key_alias: Alias to look the key up by when token is not provided
Returns:
LiteLLM_VerificationToken: The existing key row
Raises:
ProxyException: 404 if key is not found
ProxyException: 404 if key is not found, 400 if the alias matches multiple keys
"""
if prisma_client is None:
raise HTTPException(
@ -2044,19 +2045,65 @@ async def _get_and_validate_existing_key(
detail={"error": "Database not connected"},
)
hashed_token = _hash_token_if_needed(token=token)
if token is not None:
hashed_token = _hash_token_if_needed(token=token)
existing_key_row = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": hashed_token})
existing_key_row: LiteLLM_VerificationToken | None = await VerificationTokenRepository(
prisma_client
).table.find_unique(where={"token": hashed_token})
if existing_key_row is None:
if existing_key_row is None:
raise ProxyException(
message="Key not found.",
type=ProxyErrorTypes.not_found_error,
param="key",
code=status.HTTP_404_NOT_FOUND,
)
return existing_key_row
if key_alias is None:
raise ProxyException(
message="either key or key_alias must be provided",
type=ProxyErrorTypes.bad_request_error,
param="key",
code=status.HTTP_400_BAD_REQUEST,
)
rows: list[LiteLLM_VerificationToken] = await VerificationTokenRepository(prisma_client).table.find_many(
where={"key_alias": key_alias}, take=2
)
if len(rows) == 0:
raise ProxyException(
message=f"Key not found. No key with key_alias='{key_alias}'.",
type=ProxyErrorTypes.not_found_error,
param="key_alias",
code=status.HTTP_404_NOT_FOUND,
)
if len(rows) > 1:
raise ProxyException(
message=f"Multiple keys share key_alias='{key_alias}', so it cannot be used as an identifier.",
type=ProxyErrorTypes.bad_request_error,
param="key_alias",
code=status.HTTP_400_BAD_REQUEST,
)
return rows[0]
def _resolve_token_to_update(data: UpdateKeyRequest, existing_key_row: LiteLLM_VerificationToken) -> str:
if data.key is not None:
return data.key
if existing_key_row.token is None:
raise ProxyException(
message="Key not found.",
type=ProxyErrorTypes.not_found_error,
param="key",
code=status.HTTP_404_NOT_FOUND,
)
return existing_key_row
return existing_key_row.token
async def _process_single_key_update(
@ -2508,8 +2555,8 @@ async def update_key_fn(
Update an existing API key's parameters.
Parameters:
- key: str - The key to update
- key_alias: Optional[str] - User-friendly key alias
- key: Optional[str] - The key to update. Either key or key_alias must be provided.
- key_alias: Optional[str] - User-friendly key alias. If key is omitted, also identifies the key to update (must match exactly one key, same as /key/delete's key_aliases)
- user_id: Optional[str] - User ID associated with key
- team_id: Optional[str] - Team ID associated with key
- agent_id: Optional[str] - The agent id associated with the key.
@ -2592,14 +2639,14 @@ async def update_key_fn(
detail={"error": f"max_budget must be a non-negative finite number. Received: {data.max_budget}"},
)
data_json: dict = data.model_dump(exclude_unset=True)
key = data_json.pop("key")
# get the row from db
existing_key_row = await _get_and_validate_existing_key(
token=data.key,
prisma_client=prisma_client,
key_alias=data.key_alias,
)
key = _resolve_token_to_update(data=data, existing_key_row=existing_key_row)
data.key = key
await _validate_update_key_data(
data=data,
@ -3047,10 +3094,12 @@ async def bulk_update_team_keys(
)
# team_id from validated scope, never user payload — drives _check_team_key_limits.
update_key_request = UpdateKeyRequest(
key=token,
team_id=data.team_id,
**update_field_dict,
update_key_request = UpdateKeyRequest.model_validate(
{
"key": token,
"team_id": data.team_id,
**update_field_dict,
}
)
updated_key_info = await _process_single_key_update(
update_key_request=update_key_request,
@ -4048,12 +4097,14 @@ def _transform_verification_tokens_to_deleted_records(
records = []
for key in keys:
key_payload = key.model_dump()
deleted_record = LiteLLM_DeletedVerificationToken(
**key_payload,
deleted_at=deleted_at,
deleted_by=user_api_key_dict.user_id,
deleted_by_api_key=user_api_key_dict.api_key,
litellm_changed_by=litellm_changed_by,
deleted_record = LiteLLM_DeletedVerificationToken.model_validate(
{
**key_payload,
"deleted_at": deleted_at,
"deleted_by": user_api_key_dict.user_id,
"deleted_by_api_key": user_api_key_dict.api_key,
"litellm_changed_by": litellm_changed_by,
}
)
record = deleted_record.model_dump()
@ -4535,7 +4586,7 @@ async def _execute_virtual_key_regeneration(
proxy_logging_obj=proxy_logging_obj,
)
response = GenerateKeyResponse(**updated_token_dict)
response = GenerateKeyResponse.model_validate(updated_token_dict)
asyncio.create_task(
KeyManagementEventHooks.async_key_rotated_hook(
data=data,
@ -4853,7 +4904,7 @@ async def _check_proxy_or_team_admin_for_key(
)
def _validate_reset_spend_value(reset_to: Any, key_in_db: LiteLLM_VerificationToken) -> float:
def _validate_reset_spend_value(reset_to: object, key_in_db: LiteLLM_VerificationToken) -> float:
if not isinstance(reset_to, (int, float)):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@ -5029,7 +5080,7 @@ async def validate_key_list_check(
code=status.HTTP_403_FORBIDDEN,
)
complete_user_info = LiteLLM_UserTable(**complete_user_info_db_obj.model_dump())
complete_user_info = LiteLLM_UserTable.model_validate(complete_user_info_db_obj.model_dump())
# internal user can only see their own keys
if user_id:
@ -5102,7 +5153,7 @@ async def _fetch_user_team_objects(
if teams is None:
return []
return [LiteLLM_TeamTable(**team.model_dump()) for team in teams]
return [LiteLLM_TeamTable.model_validate(team.model_dump()) for team in teams]
def _get_admin_team_ids_from_objects(
@ -5851,7 +5902,7 @@ async def _list_key_helper(
if return_full_object is True or (expand and "user" in expand):
if use_deleted_table:
# Use deleted key type to preserve deleted_at, deleted_by, etc.
key_list.append(LiteLLM_DeletedVerificationToken(**key_dict))
key_list.append(LiteLLM_DeletedVerificationToken.model_validate(key_dict))
else:
key_list.append(UserAPIKeyAuth(**key_dict)) # Return full key object
else:

View file

@ -1526,7 +1526,6 @@ if MCP_AVAILABLE:
temporary_server = await global_mcp_server_manager.build_mcp_server_from_table(
temp_record,
credentials_are_encrypted=False,
persist_discovered_endpoints=False,
)
_cache_temporary_mcp_server(
temporary_server,

View file

@ -6,6 +6,7 @@ Endpoints here:
"""
import json
from collections.abc import Mapping, Sequence
from typing import Any, Dict, List, Tuple
from fastapi import APIRouter, Depends, HTTPException
@ -16,6 +17,9 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
# Clear cache and reload models to pick up the access group changes
from litellm.proxy.management_endpoints.model_management_endpoints import (
live_model_ids_snapshot,
model_info_as_mapping,
reload_serving_verdict,
clear_cache,
)
from litellm.proxy.utils import PrismaClient
@ -72,11 +76,92 @@ def add_access_group_to_deployment(model_info: Dict[str, Any], access_group: str
return model_info, True
def _raise_http_if_reload_degraded_serving(
before: frozenset[str],
written_models: Sequence[tuple[str, object]],
access_group: str,
) -> None:
"""Same verdict as the model-write endpoints, expressed through this file's
HTTPException error convention, with the metadata-only obligation: these writes
change group membership, not the models themselves, so a row that was already not
serving before the reload is never blamed here; only a model this reload stopped
serving is reported."""
missing, collateral = reload_serving_verdict(before=before, written_models=written_models, written_must_serve=False)
gone = tuple(dict.fromkeys((*missing, *collateral)))
if not gone:
return
raise HTTPException(
status_code=500,
detail={
"error": (
f"Access group '{access_group}' was saved to the database, but model id(s) {list(gone)} that "
"this pod was serving are no longer live after the reload it triggered. Other pods reload on "
"their own interval. Check server logs for 'Error upserting deployment' for the cause."
)
},
)
async def _tag_deployment_with_access_group(
model_id: str,
model_info: object,
access_group: str,
prisma_client: PrismaClient,
) -> tuple[str, Mapping[str, object]] | None:
"""Write `access_group` into one deployment's model_info; returns the
(model_id, updated model_info) pair when a write happened, None when the
deployment already carried the group."""
updated_model_info, was_modified = add_access_group_to_deployment(
model_info=dict(_readable_model_info_or_raise(model_id=model_id, model_info=model_info)),
access_group=access_group,
)
if not was_modified:
return None
await ModelRepository(prisma_client).table.update(
where={"model_id": model_id},
data={"model_info": json.dumps(updated_model_info)},
)
verbose_proxy_logger.debug(f"Updated deployment {model_id} with access group: {access_group}")
return (model_id, updated_model_info)
def _readable_model_info_or_raise(model_id: str, model_info: object) -> Mapping[str, object]:
"""These helpers rewrite the model_info column wholesale, so a present-but-unreadable
value must refuse loudly rather than be silently replaced with a fresh object; an
absent value stays a legitimate empty start."""
parsed = model_info_as_mapping(model_info)
if parsed is None and model_info is not None:
raise ValueError(f"model_info for deployment {model_id} is not a readable JSON object; refusing to rewrite it")
return parsed or {}
async def _strip_access_group_from_deployment(
model_id: str,
model_info: object,
access_group: str,
prisma_client: PrismaClient,
) -> tuple[str, Mapping[str, object]] | None:
"""Remove `access_group` from one deployment's model_info; returns the
(model_id, updated model_info) pair when a write happened, None when the
deployment did not carry the group."""
updated_model_info, was_modified = remove_access_group_from_deployment(
model_info=dict(_readable_model_info_or_raise(model_id=model_id, model_info=model_info)),
access_group=access_group,
)
if not was_modified:
return None
await ModelRepository(prisma_client).table.update(
where={"model_id": model_id},
data={"model_info": json.dumps(updated_model_info)},
)
return (model_id, updated_model_info)
async def update_deployments_with_access_group(
model_names: List[str],
access_group: str,
prisma_client: PrismaClient,
) -> int:
) -> tuple[tuple[str, Mapping[str, object]], ...]:
"""
Update all deployments for the given model names to include the access group.
@ -86,20 +171,15 @@ async def update_deployments_with_access_group(
prisma_client: Database client
Returns:
int: Number of deployments updated
The (model_id, updated model_info) pair of every deployment actually written,
so callers can verify each one survived the post-write reload
"""
models_updated = 0
deployments = await ModelRepository(prisma_client).table.find_many(where={"model_name": {"in": model_names}})
verbose_proxy_logger.debug(f"Found {len(deployments)} deployments for model_names: {model_names}")
found_names = {deployment.model_name for deployment in deployments}
for model_name in model_names:
verbose_proxy_logger.debug(f"Updating deployments for model_name: {model_name}")
# Get all deployments with this model_name
deployments = await ModelRepository(prisma_client).table.find_many(where={"model_name": model_name})
verbose_proxy_logger.debug(f"Found {len(deployments)} deployments for model_name: {model_name}")
# If no deployments found, this is a config model (not in DB)
if len(deployments) == 0:
if model_name not in found_names:
raise HTTPException(
status_code=400,
detail={
@ -107,65 +187,52 @@ async def update_deployments_with_access_group(
},
)
# Update each deployment
for deployment in deployments:
model_info = deployment.model_info or {}
# Add access group using helper
updated_model_info, was_modified = add_access_group_to_deployment(
model_info=model_info,
access_group=access_group,
)
# Only update in DB if modified
if was_modified:
await ModelRepository(prisma_client).table.update(
where={"model_id": deployment.model_id},
data={"model_info": json.dumps(updated_model_info)},
)
models_updated += 1
verbose_proxy_logger.debug(
f"Updated deployment {deployment.model_id} with access group: {access_group}"
)
return models_updated
tagged = [
await _tag_deployment_with_access_group(
model_id=deployment.model_id,
model_info=deployment.model_info,
access_group=access_group,
prisma_client=prisma_client,
)
for deployment in deployments
]
return tuple(pair for pair in tagged if pair is not None)
async def update_specific_deployments_with_access_group(
model_ids: List[str],
access_group: str,
prisma_client: PrismaClient,
) -> int:
) -> tuple[tuple[str, Mapping[str, object]], ...]:
"""
Update specific deployments (by model_id) to include the access group.
Unlike update_deployments_with_access_group which tags ALL deployments sharing
a model_name, this function only tags the specific deployments identified by
their unique model_id.
their unique model_id. Returns the (model_id, updated model_info) pair of every
deployment actually written.
"""
models_updated = 0
for model_id in model_ids:
verbose_proxy_logger.debug(f"Updating specific deployment model_id: {model_id}")
deployment = await ModelRepository(prisma_client).table.find_unique(where={"model_id": model_id})
if deployment is None:
raise HTTPException(
status_code=400,
detail={"error": f"Deployment with model_id '{model_id}' not found in Database."},
)
model_info = deployment.model_info or {}
updated_model_info, was_modified = add_access_group_to_deployment(
model_info=model_info,
verbose_proxy_logger.debug(f"Updating specific deployment model_ids: {model_ids}")
tagged = [
await _tag_deployment_with_access_group(
model_id=model_id,
model_info=(await _find_deployment_or_400(model_id=model_id, prisma_client=prisma_client)),
access_group=access_group,
prisma_client=prisma_client,
)
if was_modified:
await ModelRepository(prisma_client).table.update(
where={"model_id": model_id},
data={"model_info": json.dumps(updated_model_info)},
)
models_updated += 1
verbose_proxy_logger.debug(f"Updated deployment {model_id} with access group: {access_group}")
return models_updated
for model_id in model_ids
]
return tuple(pair for pair in tagged if pair is not None)
async def _find_deployment_or_400(model_id: str, prisma_client: PrismaClient) -> Mapping[str, object] | None:
deployment = await ModelRepository(prisma_client).table.find_unique(where={"model_id": model_id})
if deployment is None:
raise HTTPException(
status_code=400,
detail={"error": f"Deployment with model_id '{model_id}' not found in Database."},
)
return deployment.model_info
def remove_access_group_from_deployment(model_info: Dict[str, Any], access_group: str) -> Tuple[Dict[str, Any], bool]:
@ -335,20 +402,28 @@ async def create_model_group(
# Update deployments using the appropriate method
if use_model_ids:
assert data.model_ids is not None
models_updated = await update_specific_deployments_with_access_group(
updated_pairs = await update_specific_deployments_with_access_group(
model_ids=data.model_ids,
access_group=data.access_group,
prisma_client=prisma_client,
)
else:
assert data.model_names is not None
models_updated = await update_deployments_with_access_group(
updated_pairs = await update_deployments_with_access_group(
model_names=data.model_names,
access_group=data.access_group,
prisma_client=prisma_client,
)
models_updated = len(updated_pairs)
live_before_reload = live_model_ids_snapshot()
await clear_cache()
_raise_http_if_reload_degraded_serving(
before=live_before_reload,
written_models=updated_pairs,
access_group=data.access_group,
)
verbose_proxy_logger.info(
f"Successfully created access group '{data.access_group}' with {models_updated} models updated"
@ -573,38 +648,42 @@ async def update_access_group(
# Step 1: Remove access group from ALL DB deployments (skip config models)
all_deployments = await ModelRepository(prisma_client).table.find_many()
for deployment in all_deployments:
model_info = deployment.model_info or {}
updated_model_info, was_modified = remove_access_group_from_deployment(
model_info=model_info,
stripped = [
await _strip_access_group_from_deployment(
model_id=deployment.model_id,
model_info=deployment.model_info,
access_group=access_group,
prisma_client=prisma_client,
)
if was_modified:
await ModelRepository(prisma_client).table.update(
where={"model_id": deployment.model_id},
data={"model_info": json.dumps(updated_model_info)},
)
for deployment in all_deployments
]
stripped_pairs = tuple(pair for pair in stripped if pair is not None)
# Step 2: Add access group using the appropriate method
if use_model_ids:
assert data.model_ids is not None
models_updated = await update_specific_deployments_with_access_group(
updated_pairs = await update_specific_deployments_with_access_group(
model_ids=data.model_ids,
access_group=access_group,
prisma_client=prisma_client,
)
else:
assert data.model_names is not None
models_updated = await update_deployments_with_access_group(
updated_pairs = await update_deployments_with_access_group(
model_names=data.model_names,
access_group=access_group,
prisma_client=prisma_client,
)
models_updated = len(updated_pairs)
# Clear cache and reload models to pick up the access group changes
live_before_reload = live_model_ids_snapshot()
await clear_cache()
_raise_http_if_reload_degraded_serving(
before=live_before_reload,
written_models=list({**dict(stripped_pairs), **dict(updated_pairs)}.items()),
access_group=access_group,
)
verbose_proxy_logger.info(
f"Successfully updated access group '{access_group}' with {models_updated} models updated"
@ -686,25 +765,27 @@ async def delete_access_group(
try:
# Remove access group from all DB deployments (skip config models)
all_deployments = await ModelRepository(prisma_client).table.find_many()
models_updated = 0
for deployment in all_deployments:
model_info = deployment.model_info or {}
updated_model_info, was_modified = remove_access_group_from_deployment(
model_info=model_info,
removed = [
await _strip_access_group_from_deployment(
model_id=deployment.model_id,
model_info=deployment.model_info,
access_group=access_group,
prisma_client=prisma_client,
)
if was_modified:
await ModelRepository(prisma_client).table.update(
where={"model_id": deployment.model_id},
data={"model_info": json.dumps(updated_model_info)},
)
models_updated += 1
for deployment in all_deployments
]
removed_pairs = tuple(pair for pair in removed if pair is not None)
models_updated = len(removed_pairs)
# Clear cache and reload models to pick up the access group changes
live_before_reload = live_model_ids_snapshot()
await clear_cache()
_raise_http_if_reload_degraded_serving(
before=live_before_reload,
written_models=removed_pairs,
access_group=access_group,
)
verbose_proxy_logger.info(
f"Successfully deleted access group '{access_group}' from {models_updated} deployments"

View file

@ -13,6 +13,7 @@ model/{model_id}/update - PATCH endpoint for model update.
import asyncio
import datetime
import json
from collections.abc import Mapping, Sequence
from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union, cast
from fastapi import APIRouter, Depends, HTTPException, Header, Request, status
@ -52,13 +53,19 @@ from litellm.proxy.utils import PrismaClient
from litellm.repositories.model_repository import ModelRepository
from litellm.repositories.table_repositories import ModelTableRepository
from litellm.repositories.team_repository import TeamRepository
from litellm.router import Router
from litellm.types.proxy.management_endpoints.model_management_endpoints import (
UpdateUsefulLinksRequest,
)
from litellm.router_utils.auto_router_model_naming import (
STRATEGY_ROUTER_PARAM_FIELDS,
validate_strategy_router_model_write,
)
from litellm.types.router import (
SPECIAL_MODEL_INFO_PARAMS,
Deployment,
DeploymentTypedDict,
GenericLiteLLMParams,
LiteLLMParamsTypedDict,
updateDeployment,
)
@ -96,6 +103,45 @@ async def get_db_model(model_id: str, prisma_client: PrismaClient) -> Optional[D
return deployment_pydantic_obj
def _strategy_router_write_violation(
incoming_params: GenericLiteLLMParams | None,
existing_params: GenericLiteLLMParams | None,
) -> str | None:
"""Reject writes that would corrupt a strategy router's pseudo-model.
An auto-router deployment's ``litellm_params.model`` (``auto_router/...``) is
the discriminator the router loads it by; a write that mangles it makes the
router drop the deployment silently under ``ignore_invalid_deployments``.
Only writes that supply ``litellm_params.model`` are judged, against the
merged (stored + incoming) params, so partial patches and restores of an
already-corrupted row stay legal. Returns the violation, or None.
"""
if incoming_params is None or incoming_params.model is None:
return None
present_fields = frozenset(
field
for field in STRATEGY_ROUTER_PARAM_FIELDS
for source in (incoming_params, existing_params)
if source is not None and getattr(source, field, None) is not None
)
return validate_strategy_router_model_write(model=incoming_params.model, present_fields=present_fields)
def _raise_on_strategy_router_write_violation(
incoming_params: GenericLiteLLMParams | None,
existing_params: GenericLiteLLMParams | None,
) -> None:
violation = _strategy_router_write_violation(incoming_params=incoming_params, existing_params=existing_params)
if violation is None:
return
raise ProxyException(
message=violation,
type=ProxyErrorTypes.validation_error.value,
code=status.HTTP_400_BAD_REQUEST,
param="litellm_params.model",
)
def update_db_model(db_model: Deployment, updated_patch: updateDeployment) -> PrismaCompatibleUpdateDBModel:
merged_deployment_dict = DeploymentTypedDict(
model_name=db_model.model_name,
@ -253,6 +299,11 @@ async def patch_model(
param="blocked",
)
_raise_on_strategy_router_write_violation(
incoming_params=patch_data.litellm_params,
existing_params=db_model.litellm_params,
)
# Handle team model updates with proper alias management
update_data = await _update_team_model_in_db(
db_model=db_model,
@ -272,6 +323,7 @@ async def patch_model(
)
# Clear cache and reload models (uses config setting or defaults to preserving config models for DB updates)
live_before_reload = live_model_ids_snapshot()
await clear_cache()
## CREATE AUDIT LOG ##
@ -288,6 +340,12 @@ async def patch_model(
)
)
raise_if_reload_degraded_serving(
before=live_before_reload,
written_models=[(model_id, getattr(updated_model, "model_info", None))],
action="update",
)
return updated_model
except Exception as e:
@ -370,6 +428,7 @@ async def _set_model_blocked_status(
},
)
live_before_reload = live_model_ids_snapshot()
await clear_cache()
asyncio.create_task(
@ -387,6 +446,12 @@ async def _set_model_blocked_status(
)
)
raise_if_reload_degraded_serving(
before=live_before_reload,
written_models=[(data.model_id, getattr(updated_model, "model_info", None))],
action=action,
)
return updated_model
except Exception as e:
@ -713,13 +778,8 @@ async def _get_team_deployments(
# Confirm team_id in model_info (defensive check)
result = []
for row in response:
model_info = row.model_info
if isinstance(model_info, str):
try:
model_info = json.loads(model_info)
except (TypeError, ValueError):
continue
if isinstance(model_info, dict) and model_info.get("team_id") == team_id:
model_info = model_info_as_mapping(row.model_info)
if model_info is not None and model_info.get("team_id") == team_id:
result.append(row)
return result
@ -770,13 +830,8 @@ async def _get_team_public_model_names(
deployments = await _get_team_deployments(team_id, prisma_client)
public_names: Set[str] = set()
for row in deployments:
model_info = row.model_info
if isinstance(model_info, str):
try:
model_info = json.loads(model_info)
except (TypeError, ValueError):
continue
if isinstance(model_info, dict):
model_info = model_info_as_mapping(row.model_info)
if model_info is not None:
public_name = model_info.get("team_public_model_name")
if public_name:
public_names.add(public_name)
@ -788,6 +843,7 @@ async def _remove_unbacked_team_models(
prisma_client: PrismaClient,
user_api_key_cache: Any,
proxy_logging_obj: Any,
llm_router: Router | None = None,
) -> None:
"""
Strip a deleted team model's public name(s) from team.models and refresh the cache.
@ -795,26 +851,50 @@ async def _remove_unbacked_team_models(
Must be called after the deployment row is deleted: a public name is removed only
when no remaining team deployment still backs it, so a load-balanced replica isn't
revoked while siblings serve it, and concurrent deletes can't leave a ghost.
Legacy team models (created before team_public_model_name existed) store a
``{public_name: "model_name_{team_id}_{uuid}"}`` entry in the team's model_aliases,
so the alias scan runs for every team model; skipping it for internal-shaped names
left stale aliases that rewrote requests to deployments that no longer exist.
Aliases are scrubbed only when the deleted deployment's name no longer resolves in
the router, so deleting one replica of a load-balanced group never breaks aliases
that still route to the surviving replicas (in any team).
A public name that still resolves to a live router deployment (e.g. a gateway-level
model group shared with the team) is kept in team.models, so deleting a per-team
duplicate does not revoke the team's access to the shared deployment.
"""
team_id = model_params.model_info.team_id
if team_id is None:
return
# BYOK models carry an internal `model_name_{team_id}_{uuid}` name that can never
# be a team alias value, so skip the full litellm_modeltable scan for them.
removed_model_aliases: List[Tuple[str, str]] = []
if not model_params.model_name.startswith(f"model_name_{team_id}_"):
removed_model_aliases = await delete_team_model_alias(
deleted_name_still_served = (
llm_router is not None and model_params.model_name in llm_router.model_name_to_deployment_indices
)
removed_model_aliases: List[Tuple[str, str]] = (
[]
if deleted_name_still_served
else await delete_team_model_alias(
public_model_name=model_params.model_name,
prisma_client=prisma_client,
)
names_to_remove = {alias for alias_team_id, alias in removed_model_aliases if alias_team_id == team_id}
if model_params.model_info.team_public_model_name is not None:
names_to_remove.add(model_params.model_info.team_public_model_name)
if names_to_remove:
names_to_remove -= await _get_team_public_model_names(team_id=team_id, prisma_client=prisma_client)
)
removed_alias_names = {alias for alias_team_id, alias in removed_model_aliases if alias_team_id == team_id}
candidate_names = (
removed_alias_names | {model_params.model_info.team_public_model_name}
if model_params.model_info.team_public_model_name is not None
else removed_alias_names
)
if not candidate_names:
return
team_backed_names = await _get_team_public_model_names(team_id=team_id, prisma_client=prisma_client)
router_served_names = (
frozenset(name for name in candidate_names if name in llm_router.model_name_to_deployment_indices)
if llm_router is not None
else frozenset()
)
names_to_remove = candidate_names - team_backed_names - router_served_names
if not names_to_remove:
return
@ -853,18 +933,11 @@ async def _update_existing_team_model_assignment(
def _get_team_public_model_name(
model_info: Optional[Union[dict, str]],
) -> Optional[str]:
if isinstance(model_info, dict):
value = model_info.get("team_public_model_name")
return value if isinstance(value, str) else None
if isinstance(model_info, str):
try:
parsed = json.loads(model_info)
except (TypeError, ValueError):
return None
if isinstance(parsed, dict):
value = parsed.get("team_public_model_name")
return value if isinstance(value, str) else None
return None
parsed = model_info_as_mapping(model_info)
if parsed is None:
return None
value = parsed.get("team_public_model_name")
return value if isinstance(value, str) else None
old_public_name = db_model.model_info.team_public_model_name if db_model.model_info else None
@ -978,7 +1051,7 @@ class ModelManagementAuthChecks:
status_code=400,
detail={"error": "Team id={} does not exist in db".format(model_params.model_info.team_id)},
)
existing_team_row = LiteLLM_TeamTable(**_existing_team_row.model_dump())
existing_team_row = LiteLLM_TeamTable.model_validate(_existing_team_row.model_dump())
ModelManagementAuthChecks.can_user_make_team_model_call(
team_id=model_params.model_info.team_id,
@ -1016,7 +1089,7 @@ class ModelManagementAuthChecks:
status_code=400,
detail={"error": "Team id={} does not exist in db".format(model_params.model_info.team_id)},
)
team_obj = LiteLLM_TeamTable(**team_obj_row.model_dump())
team_obj = LiteLLM_TeamTable.model_validate(team_obj_row.model_dump())
return ModelManagementAuthChecks.can_user_make_team_model_call(
team_id=model_params.model_info.team_id,
@ -1120,6 +1193,7 @@ async def delete_model(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
llm_router=llm_router,
)
## CREATE AUDIT LOG ##
@ -1267,6 +1341,11 @@ async def add_new_model(
premium_user=premium_user,
)
_raise_on_strategy_router_write_violation(
incoming_params=model_params.litellm_params,
existing_params=None,
)
model_response: Optional[LiteLLM_ProxyModelTable] = None
# update DB
if store_model_in_db is True:
@ -1275,6 +1354,7 @@ async def add_new_model(
- store keys separately
"""
live_before_reload = live_model_ids_snapshot()
try:
_original_litellm_model_name = model_params.model_name
if model_params.model_info.team_id is None:
@ -1330,6 +1410,12 @@ async def add_new_model(
)
)
raise_if_reload_degraded_serving(
before=live_before_reload,
written_models=[(model_response.model_id, getattr(model_response, "model_info", None))],
action="create",
)
return model_response
except Exception as e:
@ -1414,6 +1500,11 @@ async def update_model(
premium_user=premium_user,
)
_raise_on_strategy_router_write_violation(
incoming_params=model_params.litellm_params,
existing_params=deployment.litellm_params,
)
# update DB
if store_model_in_db is True:
_existing_litellm_params_dict = dict(_existing_litellm_params.litellm_params)
@ -1450,8 +1541,8 @@ async def update_model(
)
# Clear cache and reload models (uses config setting or defaults to preserving config models for DB updates)
live_before_reload = live_model_ids_snapshot()
await clear_cache()
## CREATE AUDIT LOG ##
asyncio.create_task(
create_object_audit_log(
@ -1474,6 +1565,12 @@ async def update_model(
)
)
raise_if_reload_degraded_serving(
before=live_before_reload,
written_models=[(_model_id, getattr(model_response, "model_info", None))],
action="update",
)
return model_response
except Exception as e:
verbose_proxy_logger.exception(
@ -1677,6 +1774,114 @@ def _deduplicate_litellm_router_models(models: List[Dict]) -> List[Dict]:
return unique_models
def model_info_as_mapping(model_info: object) -> Mapping[str, object] | None:
"""A DB row's model_info column arrives as a dict or as its JSON string depending on
the query path, and every consumer needs the mapping. Single owner of that parse:
returns None when no usable mapping exists (None, an unparseable string, or JSON
that is not an object), and callers choose what None means for them."""
if isinstance(model_info, Mapping):
return model_info
if not isinstance(model_info, str):
return None
try:
parsed = json.loads(model_info)
except (TypeError, ValueError):
return None
return parsed if isinstance(parsed, Mapping) else None
def _expects_liveness_on_this_pod(model_info: object) -> bool:
from litellm.router import model_info_is_active_for_environment
try:
return model_info_is_active_for_environment(model_info=model_info_as_mapping(model_info))
except ValueError:
return True
def live_model_ids_snapshot() -> frozenset[str]:
"""The ids this pod's router is currently serving, read fresh from the module global
because a reload can rebind it. The empirical ground truth every verdict below is
computed from; an absent router serves nothing."""
from litellm.proxy.proxy_server import llm_router
if llm_router is None:
return frozenset()
return frozenset(llm_router.get_model_ids())
def reload_serving_verdict(
before: frozenset[str],
written_models: Sequence[tuple[str, object]],
written_must_serve: bool,
) -> tuple[tuple[str, ...], tuple[str, ...]]:
"""Judge a write-triggered reload by diffing the router's serving state instead of
trusting any layer of the reload stack to report its own failure.
The full cell matrix, per id:
- written, must-serve (the write's purpose is this model's serving state): live now
is fine; not live is reported unless the row is deliberately inactive for this
pod's LITELLM_ENVIRONMENT; a row whose model_info cannot be read counts as
expecting to serve, so its drop is still reported
- written, metadata-only (must_not_degrade): live before and gone now is reported;
a row that was already not serving stays silent, because its deadness predates
this write and blaming it would block unrelated metadata fixes
- not written but live before and gone now: collateral degradation of this pod
caused by the reload this request triggered (a wholesale re-add failure, or a
newly introduced conflict), always reported
Returns (written ids violating their obligation, collateral ids no longer served).
Best effort under concurrent admin writes: the snapshot spans only this request.
"""
now = live_model_ids_snapshot()
written_ids = frozenset(model_id for model_id, _ in written_models)
if written_must_serve:
missing = tuple(
model_id
for model_id, model_info in written_models
if model_id not in now and _expects_liveness_on_this_pod(model_info)
)
else:
missing = tuple(model_id for model_id, _ in written_models if model_id in before and model_id not in now)
collateral = tuple(sorted(before - now - written_ids))
return (missing, collateral)
def raise_if_reload_degraded_serving(
before: frozenset[str],
written_models: Sequence[tuple[str, object]],
action: str,
) -> None:
"""The caller-visible error this pod's model-write endpoints owe their caller when
the model they wrote is not being served after the reload they triggered. The DB
write is durable either way and every other pod reloads on its own interval; this
speaks only for the handling pod."""
missing, collateral = reload_serving_verdict(before=before, written_models=written_models, written_must_serve=True)
if not missing and not collateral:
return
missing_clause = (
f"the model id(s) {list(missing)} are not live in this pod's router after the reload and are not "
"being served by this pod."
if missing
else "the reload it triggered degraded this pod's serving state."
)
collateral_clause = (
f" Previously served model id(s) {list(collateral)} are also no longer being served by this pod."
if collateral
else ""
)
raise ProxyException(
message=(
f"Model {action} was saved to the database, but {missing_clause}{collateral_clause} "
"Other pods reload on their own interval. Check server logs for 'Error upserting deployment' or "
"'Error creating deployment' for the cause."
),
type=ProxyErrorTypes.internal_server_error,
code=status.HTTP_500_INTERNAL_SERVER_ERROR,
param=None,
)
async def clear_cache():
"""
Clear router caches and reload models.

View file

@ -175,6 +175,7 @@ class ScimTransformations:
SCIMMember(
value=ScimTransformations._get_scim_member_value(member),
display=ScimTransformations._get_scim_member_display(member),
type="User",
)
)

View file

@ -5,7 +5,9 @@ This is an enterprise feature and requires a premium license.
"""
import re
from typing import Any, Dict, Iterable, List, Optional, Set, Tuple
from collections.abc import Mapping, Sequence
from itertools import chain
from typing import Any, Dict, Iterable, List, NamedTuple, Optional, Set, Tuple
from fastapi import (
APIRouter,
@ -17,8 +19,8 @@ from fastapi import (
Request,
Response,
)
from pydantic import BaseModel, ValidationError
from typing_extensions import TypedDict
from pydantic import BaseModel, TypeAdapter, ValidationError
from typing_extensions import TypedDict, assert_never
import litellm
from litellm._logging import verbose_proxy_logger
@ -50,7 +52,11 @@ from litellm.proxy.management_endpoints.team_endpoints import (
team_member_add,
team_member_delete,
)
from litellm.proxy.utils import _premium_user_check, handle_exception_on_proxy
from litellm.proxy.utils import (
PrismaClient,
_premium_user_check,
handle_exception_on_proxy,
)
from litellm.repositories.table_repositories import (
InvitationLinkRepository,
OrganizationMembershipRepository,
@ -143,7 +149,11 @@ class ScimUserData(TypedDict):
class GroupMemberExtractionResult(BaseModel):
"""Result of extracting and processing group members."""
"""Result of extracting and processing group members.
``all_member_ids`` is deduped order-preserving; ``existing_member_ids`` is not,
so a repeated resolved id appears once in the former and twice in the latter.
"""
existing_member_ids: List[str]
created_users: List[NewUserResponse]
@ -371,6 +381,216 @@ async def _recompute_scim_member_roles(prisma_client: Any, user_ids: Iterable[st
)
class _ResolvedUserMember(NamedTuple):
user_id: str
class _SkippedGroupMember(NamedTuple):
value: str
reason: Literal["nested_group", "non_user_type", "existing_team"]
class _UnknownMember(NamedTuple):
value: str
_ClassifiedGroupMember = Union[_ResolvedUserMember, _SkippedGroupMember, _UnknownMember]
class _PartitionedMembers(NamedTuple):
resolved_ids: tuple[str, ...]
skipped: tuple[_SkippedGroupMember, ...]
unknown_ids: tuple[str, ...]
def _member_value(member: SCIMMember) -> str:
"""A member id is opaque to us but has to be there; an empty one is a client error."""
if not member.value or not member.value.strip():
raise HTTPException(
status_code=400,
detail={"error": "Invalid member: user ID cannot be empty."},
)
return member.value
def _normalized_member_type(member: SCIMMember) -> str | None:
"""The canonical ``type`` a member declares, lowercased; blank or absent means none."""
normalized = (member.type or "").strip().lower()
return normalized or None
_JSON_OBJECT_ADAPTER = TypeAdapter(Dict[str, object])
def _json_object_fields(raw: object) -> Mapping[str, object] | None:
"""A typed, read-only view of a JSON object, or None when it is not one."""
try:
return _JSON_OBJECT_ADAPTER.validate_python(raw)
except ValidationError:
return None
def _team_metadata_has_scim_provenance(team_metadata: object) -> bool:
"""Whether a group write from the identity provider left its mark on this team.
``SCIM_TEAM_DATA_METADATA_KEY`` counts because PUT has been writing it since
long before the explicit marker, so a team the identity provider already
syncs is recognized without waiting to be written again.
"""
fields = _json_object_fields(team_metadata)
if fields is None:
return False
return bool(fields.get(SCIM_MANAGED_TEAM_METADATA_KEY)) or fields.get(SCIM_TEAM_DATA_METADATA_KEY) is not None
async def _classify_group_member(member: SCIMMember, prisma_client: PrismaClient) -> _ClassifiedGroupMember:
"""
Decide what a single SCIM group member refers to.
A LiteLLM team only holds users, so a member is dropped when it declares a type
other than ``User`` or when its id names an existing team. Both of those checks
are placed around the user lookup rather than before it, because the id of a
real user is the one thing that outranks them:
- ``"type": "Group"`` (what Entra sends for a nested group) is dropped without
a lookup. This bug provisioned nested group GUIDs as users, so those rows
exist in the wild and would otherwise resolve as members all over again.
- any other unrecognized type is dropped only after the user lookup misses.
Clients do send non-canonical types on real members (RFC 7643 defines
``direct`` for ``User.groups``), and dropping a live user over one would
revoke that user's team access on the next full sync.
- an id that names an existing team is dropped only when the member arrives
untyped, which is how Okta sends nested groups, and only when that team is
one the identity provider writes. An id the IdP called a User is a user
even if some team happens to share the id, and a team created here rather
than through SCIM is not evidence of anything about the member.
"""
value = _member_value(member)
member_type = _normalized_member_type(member)
if member_type == "group":
return _SkippedGroupMember(value=value, reason="nested_group")
user = await UserRepository(prisma_client).table.find_unique(where={"user_id": value})
if user is not None:
return _ResolvedUserMember(user_id=value)
if member_type is not None and member_type != "user":
return _SkippedGroupMember(value=value, reason="non_user_type")
if member_type is None:
team = await TeamRepository(prisma_client).table.find_unique(where={"team_id": value})
if team is not None and _team_metadata_has_scim_provenance(team.metadata):
return _SkippedGroupMember(value=value, reason="existing_team")
return _UnknownMember(value=value)
def _bucketed_member(entry: _ClassifiedGroupMember) -> _PartitionedMembers:
"""The single-member partition one classified entry contributes."""
match entry:
case _ResolvedUserMember(user_id=user_id):
return _PartitionedMembers(resolved_ids=(user_id,), skipped=(), unknown_ids=())
case _SkippedGroupMember():
return _PartitionedMembers(resolved_ids=(), skipped=(entry,), unknown_ids=())
case _UnknownMember(value=value):
return _PartitionedMembers(resolved_ids=(), skipped=(), unknown_ids=(value,))
case _:
assert_never(entry)
def _partition_classified_members(classified: Iterable[_ClassifiedGroupMember]) -> _PartitionedMembers:
"""Split classified members into the buckets the resolver acts on, keeping request order."""
bucketed = tuple(_bucketed_member(entry) for entry in classified)
return _PartitionedMembers(
resolved_ids=tuple(chain.from_iterable(bucket.resolved_ids for bucket in bucketed)),
skipped=tuple(chain.from_iterable(bucket.skipped for bucket in bucketed)),
unknown_ids=tuple(chain.from_iterable(bucket.unknown_ids for bucket in bucketed)),
)
def _admitted_member_id(entry: _ClassifiedGroupMember, created_ids: frozenset[str]) -> str | None:
match entry:
case _ResolvedUserMember(user_id=user_id):
return user_id
case _UnknownMember(value=value):
return value if value in created_ids else None
case _SkippedGroupMember():
return None
case _:
assert_never(entry)
def _admitted_member_ids(classified: Iterable[_ClassifiedGroupMember], created_ids: frozenset[str]) -> tuple[str, ...]:
"""Member ids that survive resolution, in the order the request listed them.
An id the request repeats is one member: the roster these ids are written to
holds one row per member, and a second creation attempt for the same id fails
against the real unique constraint even though the first one succeeded.
"""
return tuple(
dict.fromkeys(
member_id for entry in classified if (member_id := _admitted_member_id(entry, created_ids)) is not None
)
)
async def _resolve_group_member_ids(
members: Sequence[SCIMMember],
created_via: str,
prisma_client: PrismaClient,
) -> GroupMemberExtractionResult:
"""
Resolve SCIM group members to LiteLLM user ids, dropping members that are not users.
Only the operations that put ids onto a roster resolve their members: an id
that resolves to nothing is created when litellm_settings.scim_upsert_user is
True (default) and rejected per SCIM 2.0 otherwise. Removals do not come
through here; dropping an id is idempotent, so it needs neither a lookup nor a
user to drop.
Raises:
HTTPException: 400 when a member id is empty, or when scim_upsert_user is
False and a member id is neither an existing user, an existing team, nor a
member declared to be something other than a user.
"""
classified = tuple([await _classify_group_member(member, prisma_client) for member in members])
partition = _partition_classified_members(classified)
for skipped in partition.skipped:
verbose_proxy_logger.info(
"SCIM: ignoring non-user group member '%s' (%s); LiteLLM teams contain users only",
skipped.value,
skipped.reason,
)
if partition.unknown_ids and not await _get_scim_upsert_user_setting():
raise HTTPException(
status_code=400,
detail={
"error": f"User with ID '{partition.unknown_ids[0]}' does not exist. "
"Please create the user first via POST /Users before adding to group."
},
)
creations = tuple(
[
(user_id, await _create_user_if_not_exists(user_id=user_id, created_via=created_via))
for user_id in partition.unknown_ids
]
)
created_users = tuple(created for _, created in creations if created is not None)
return GroupMemberExtractionResult(
existing_member_ids=partition.resolved_ids,
created_users=created_users,
all_member_ids=_admitted_member_ids(
classified,
frozenset(user_id for user_id, created in creations if created is not None),
),
)
async def _extract_group_member_ids(group: SCIMGroup) -> GroupMemberExtractionResult:
"""
Extract member IDs from SCIMGroup, validating that all users exist.
@ -386,56 +606,10 @@ async def _extract_group_member_ids(group: SCIMGroup) -> GroupMemberExtractionRe
HTTPException: If scim_upsert_user is False and any member user does not exist (400 Bad Request)
"""
prisma_client = await _get_prisma_client_or_raise_exception()
existing_member_ids = []
created_users = []
all_member_ids = []
# Check the feature flag
scim_upsert_user = await _get_scim_upsert_user_setting()
if group.members:
for member in group.members:
user_id = member.value
# Validate user_id is not empty
if not user_id or not user_id.strip():
raise HTTPException(
status_code=400,
detail={"error": "Invalid member: user ID cannot be empty."},
)
# Check if user exists
user = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id})
if user:
existing_member_ids.append(user_id)
all_member_ids.append(user_id)
else:
if scim_upsert_user:
# Create the user if they don't exist (backward compatible behavior)
created_user = await _create_user_if_not_exists(
user_id=user_id, created_via="scim_group_membership"
)
if created_user:
created_users.append(created_user)
all_member_ids.append(user_id)
# If creation failed, user is skipped (logged in helper)
else:
# User doesn't exist - reject per SCIM 2.0 protocol
# This prevents security issues where users not assigned to app
# get provisioned via group membership
raise HTTPException(
status_code=400,
detail={
"error": f"User with ID '{user_id}' does not exist. "
"Please create the user first via POST /Users before adding to group."
},
)
return GroupMemberExtractionResult(
existing_member_ids=existing_member_ids,
created_users=created_users,
all_member_ids=all_member_ids,
return await _resolve_group_member_ids(
members=group.members or [],
created_via="scim_group_membership",
prisma_client=prisma_client,
)
@ -448,7 +622,7 @@ async def _get_team_members_display(member_ids: List[str]) -> List[SCIMMember]:
user = await UserRepository(prisma_client).table.find_unique(where={"user_id": member_id})
if user:
display_name = user.user_email or user.user_id
members.append(SCIMMember(value=user.user_id, display=display_name))
members.append(SCIMMember(value=user.user_id, display=display_name, type="User"))
return members
@ -863,6 +1037,14 @@ def _get_schemas() -> list:
type="string",
description="Member display name.",
),
SCIMSchemaAttribute(
name="type",
type="string",
description=(
'The type of member; canonical values are "User" and "Group". '
"Only members of type User are honored, LiteLLM teams contain users only."
),
),
],
),
],
@ -1317,7 +1499,7 @@ async def delete_user(
where={"team_id": team.team_id}, data={"members": new_members}
)
team_row = LiteLLM_TeamTable(**team.model_dump())
team_row = LiteLLM_TeamTable.model_validate(team.model_dump())
if any(member.user_id == user_id for member in team_row.members_with_roles or []):
await team_member_delete(
data=TeamMemberDeleteRequest(team_id=team_row.team_id, user_id=user_id),
@ -1336,21 +1518,42 @@ async def delete_user(
raise handle_exception_on_proxy(e)
def _extract_group_values(value: Any) -> List[str]:
def _parse_member_entry(entry: object) -> SCIMMember | None:
"""Parse one entry of a SCIM patch value, or None when it carries no id."""
if isinstance(entry, str):
return SCIMMember(value=entry)
fields = _json_object_fields(entry)
if fields is None:
return None
entry_value = fields.get("value")
if not entry_value:
return None
entry_display = fields.get("display")
entry_type = fields.get("type")
return SCIMMember(
value=str(entry_value),
display=str(entry_display) if entry_display is not None else None,
type=entry_type if isinstance(entry_type, str) else None,
)
def _parse_member_entries(value: object) -> tuple[SCIMMember, ...]:
"""Parse a SCIM patch value into members, keeping each entry's ``type``.
PATCH bodies bypass SCIMGroup parsing (SCIMPatchOperation.value is untyped),
so member objects arrive as raw dicts and the ``type`` that marks a nested
group would otherwise be lost.
"""
entries: tuple[object, ...] = tuple(value) if isinstance(value, list) else (value,)
return tuple(member for member in (_parse_member_entry(entry) for entry in entries) if member is not None)
def _extract_group_values(value: object) -> List[str]:
"""Return group ids from a SCIM patch value."""
group_values: List[str] = []
if isinstance(value, list):
for v in value:
if isinstance(v, dict) and v.get("value"):
group_values.append(str(v.get("value")))
elif isinstance(v, str):
group_values.append(v)
elif isinstance(value, dict):
if value.get("value"):
group_values.append(str(value.get("value")))
elif isinstance(value, str):
group_values.append(value)
return group_values
return [member.value for member in _parse_member_entries(value)]
def _extract_ids_from_path_filter(path: str | None, attribute: str) -> List[str]:
@ -1833,6 +2036,7 @@ async def create_group(
team_id=team_id,
team_alias=group.displayName,
members_with_roles=members_with_roles,
metadata={SCIM_MANAGED_TEAM_METADATA_KEY: True},
),
http_request=Request(scope={"type": "http", "path": "/scim/v2/Groups"}),
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
@ -1875,7 +2079,11 @@ async def update_group(
# Prepare update data
existing_metadata = existing_team.metadata if existing_team.metadata else {}
updated_metadata = {**existing_metadata, "scim_data": group.model_dump()}
updated_metadata = {
**existing_metadata,
SCIM_TEAM_DATA_METADATA_KEY: group.model_dump(),
SCIM_MANAGED_TEAM_METADATA_KEY: True,
}
update_data = {
"team_alias": group.displayName,
@ -1968,12 +2176,17 @@ async def _process_group_patch_operations(
is absolute: it declares the roster is exactly this set, so the caller must
reconcile against it as a set-to-target rather than rebasing it onto a
concurrently-mutated roster.
A ``remove`` drops the ids it names without resolving them first. Removal is
idempotent and cannot put anything on a roster, while resolving would make it
conditional on what the id turns out to be and leave members we should never
have admitted - the phantom users this endpoint used to create for nested
groups - impossible to clean up.
"""
update_data: Dict[str, Any] = {}
# Create a fresh copy of existing metadata to avoid Prisma issues
existing_metadata = existing_team.metadata or {}
metadata = dict(existing_metadata) if existing_metadata else {}
metadata = {**(existing_team.metadata or {}), SCIM_MANAGED_TEAM_METADATA_KEY: True}
# Track member changes. members_with_roles is the source of truth for team
# membership; the legacy `members` column is not populated by team creation
@ -2001,50 +2214,26 @@ async def _process_group_patch_operations(
metadata["externalId"] = str(value)
elif path.startswith("members"):
# Handle member operations
member_values = _extract_group_values(value)
if not member_values and value is None:
member_values = _extract_ids_from_path_filter(op.path, "members")
# Check the feature flag
scim_upsert_user = await _get_scim_upsert_user_setting()
# Validate all users exist or create them based on feature flag
valid_members = []
for member_id in member_values:
# Validate member_id is not empty
if not member_id or not member_id.strip():
raise HTTPException(
status_code=400,
detail={"error": "Invalid member: user ID cannot be empty."},
)
patched_members = (
_parse_member_entries(value)
if value is not None
else tuple(
SCIMMember(value=member_id) for member_id in _extract_ids_from_path_filter(op.path, "members")
)
)
user = await UserRepository(prisma_client).table.find_unique(where={"user_id": member_id})
if user:
valid_members.append(member_id)
else:
if scim_upsert_user:
# Create the user if they don't exist (backward compatible behavior)
created_user = await _create_user_if_not_exists(
user_id=member_id, created_via="scim_group_patch"
)
if created_user:
valid_members.append(member_id)
# If creation failed, user is skipped (logged in helper)
else:
# User doesn't exist - reject per SCIM 2.0 protocol
raise HTTPException(
status_code=400,
detail={
"error": f"User with ID '{member_id}' does not exist. "
"Please create the user first via POST /Users before adding to group."
},
)
if op_type == "replace":
final_members = set(valid_members)
elif op_type == "add":
final_members.update(valid_members)
elif op_type == "remove":
for member_id in valid_members:
final_members.discard(member_id)
if op_type == "remove":
final_members = final_members - {_member_value(member) for member in patched_members}
else:
member_result = await _resolve_group_member_ids(
members=patched_members,
created_via="scim_group_patch",
prisma_client=prisma_client,
)
if op_type == "replace":
final_members = set(member_result.all_member_ids)
elif op_type == "add":
final_members = final_members | set(member_result.all_member_ids)
else:
# Handle other generic metadata
if op_type == "remove":
@ -2052,9 +2241,7 @@ async def _process_group_patch_operations(
else:
metadata[path] = value
# Include metadata in update data if it exists
if metadata:
update_data["metadata"] = metadata
update_data["metadata"] = metadata
member_replace_present = any(
op.op == "replace" and (op.path or "").lower().startswith("members") for op in patch_ops.Operations
@ -2145,7 +2332,9 @@ async def patch_group(
refreshed_team = await TeamRepository(prisma_client).table.find_unique(where={"team_id": group_id})
refreshed_current = (
set(await _get_team_member_user_ids_from_team(LiteLLM_TeamTable(**refreshed_team.model_dump())))
set(
await _get_team_member_user_ids_from_team(LiteLLM_TeamTable.model_validate(refreshed_team.model_dump()))
)
if refreshed_team
else snapshot_members
)
@ -2173,7 +2362,7 @@ async def patch_group(
# Convert to SCIM format and return
scim_group = await ScimTransformations.transform_litellm_team_to_scim_group(
LiteLLM_TeamTable(**updated_team.model_dump())
LiteLLM_TeamTable.model_validate(updated_team.model_dump())
)
return scim_group

View file

@ -141,7 +141,7 @@ from litellm.types.proxy.management_endpoints.team_endpoints import (
router = APIRouter()
def _sanitize_for_log(value: Any) -> str:
def _sanitize_for_log(value: object) -> str:
"""Strip CR/LF from user-controlled values to prevent log injection."""
try:
text = str(value)
@ -171,7 +171,7 @@ async def _refresh_cached_team(
"""
await _cache_team_object(
team_id=team_row.team_id,
team_table=LiteLLM_TeamTableCachedObj(**team_row.model_dump()),
team_table=LiteLLM_TeamTableCachedObj.model_validate(team_row.model_dump()),
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
@ -510,7 +510,7 @@ async def get_all_team_memberships(
returned_tm: List[LiteLLM_TeamMembership] = []
for tm in team_memberships:
returned_tm.append(LiteLLM_TeamMembership(**tm.model_dump()))
returned_tm.append(LiteLLM_TeamMembership.model_validate(tm.model_dump()))
return returned_tm
@ -772,7 +772,7 @@ async def _check_org_team_limits(
# Convert teams to LiteLLM_TeamTable objects
team_objs: List[LiteLLM_TeamTable] = []
for team in teams:
team_objs.append(LiteLLM_TeamTable(**team.model_dump()))
team_objs.append(LiteLLM_TeamTable.model_validate(team.model_dump()))
check_org_team_model_specific_limits(
teams=team_objs,
@ -1467,9 +1467,9 @@ async def fetch_and_validate_organization(
)
is_proxy_admin = user_api_key_dict is not None and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
organization = LiteLLM_OrganizationTableWithMembers(**organization_row.model_dump())
organization = LiteLLM_OrganizationTableWithMembers.model_validate(organization_row.model_dump())
validate_team_org_change(
team=LiteLLM_TeamTable(**existing_team_row.model_dump()),
team=LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()),
organization=organization,
llm_router=llm_router,
is_proxy_admin=is_proxy_admin,
@ -1477,7 +1477,7 @@ async def fetch_and_validate_organization(
if is_proxy_admin:
await _auto_add_team_members_to_organization(
team=LiteLLM_TeamTable(**existing_team_row.model_dump()),
team=LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()),
organization=organization,
prisma_client=prisma_client,
)
@ -1714,7 +1714,7 @@ async def update_team(
# Verify caller has access to manage this team
await _verify_team_access(
team_obj=LiteLLM_TeamTable(**existing_team_row.model_dump()),
team_obj=LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()),
user_api_key_dict=user_api_key_dict,
)
@ -2013,7 +2013,7 @@ async def patch_team(
existing_metadata = existing_team_row.metadata if isinstance(existing_team_row.metadata, dict) else {}
patch_fields["metadata"] = apply_json_merge_patch(existing_metadata, patch_fields["metadata"])
update_request = UpdateTeamRequest(team_id=team_id, **patch_fields)
update_request = UpdateTeamRequest.model_validate({"team_id": team_id, **patch_fields})
result = await update_team(
data=update_request,
@ -2591,7 +2591,7 @@ async def team_member_add(
detail={"error": f"Team not found for team_id={getattr(data, 'team_id', None)}"},
)
complete_team_data = LiteLLM_TeamTable(**existing_team_row.model_dump())
complete_team_data = LiteLLM_TeamTable.model_validate(existing_team_row.model_dump())
team_member_add_duplication_check(
data=data,
@ -2636,10 +2636,12 @@ async def team_member_add(
_emit_team_members_metric(complete_team_data)
return TeamAddMemberResponse(
**updated_team.model_dump(),
updated_users=updated_users,
updated_team_memberships=updated_team_memberships,
return TeamAddMemberResponse.model_validate(
{
**updated_team.model_dump(),
"updated_users": updated_users,
"updated_team_memberships": updated_team_memberships,
}
)
@ -2711,7 +2713,7 @@ async def team_member_delete(
status_code=400,
detail={"error": "Team id={} does not exist in db".format(data.team_id)},
)
existing_team_row = LiteLLM_TeamTable(**_existing_team_row.model_dump())
existing_team_row = LiteLLM_TeamTable.model_validate(_existing_team_row.model_dump())
## CHECK IF USER IS PROXY ADMIN OR TEAM ADMIN OR ORG ADMIN
@ -2915,7 +2917,7 @@ async def team_member_update(
status_code=400,
detail={"error": "Team id={} does not exist in db".format(data.team_id)},
)
existing_team_row = LiteLLM_TeamTable(**_existing_team_row.model_dump())
existing_team_row = LiteLLM_TeamTable.model_validate(_existing_team_row.model_dump())
## CHECK IF USER IS PROXY ADMIN OR TEAM ADMIN OR ORG ADMIN
@ -3261,7 +3263,7 @@ async def delete_team(
status_code=404,
detail={"error": f"Team not found, passed team_id={team_id}"},
)
team_row_pydantic = LiteLLM_TeamTable(**team_row_base.model_dump())
team_row_pydantic = LiteLLM_TeamTable.model_validate(team_row_base.model_dump())
# Verify caller has access to manage this team
await _verify_team_access(
@ -3385,12 +3387,14 @@ def _transform_teams_to_deleted_records(
records = []
for team in teams:
team_payload = team.model_dump()
deleted_record = LiteLLM_DeletedTeamTable(
**team_payload,
deleted_at=deleted_at,
deleted_by=user_api_key_dict.user_id,
deleted_by_api_key=user_api_key_dict.api_key,
litellm_changed_by=litellm_changed_by,
deleted_record = LiteLLM_DeletedTeamTable.model_validate(
{
**team_payload,
"deleted_at": deleted_at,
"deleted_by": user_api_key_dict.user_id,
"deleted_by_api_key": user_api_key_dict.api_key,
"litellm_changed_by": litellm_changed_by,
}
)
record = deleted_record.model_dump()
@ -3580,7 +3584,7 @@ async def team_info(
)
await validate_membership(
user_api_key_dict=user_api_key_dict,
team_table=LiteLLM_TeamTable(**team_info.model_dump()),
team_table=LiteLLM_TeamTable.model_validate(team_info.model_dump()),
)
## GET ALL KEYS ##
@ -3615,9 +3619,9 @@ async def team_info(
returned_tm = await get_all_team_memberships(prisma_client, [team_id], user_id=None)
if isinstance(team_info, dict):
_team_info = TeamInfoResponseObjectTeamTable(**team_info)
_team_info = TeamInfoResponseObjectTeamTable.model_validate(team_info)
elif isinstance(team_info, BaseModel):
_team_info = TeamInfoResponseObjectTeamTable(**team_info.model_dump())
_team_info = TeamInfoResponseObjectTeamTable.model_validate(team_info.model_dump())
else:
_team_info = TeamInfoResponseObjectTeamTable()
@ -3823,7 +3827,7 @@ async def block_team(
# Verify caller has access to manage this team
await _verify_team_access(
team_obj=LiteLLM_TeamTable(**existing_team.model_dump()),
team_obj=LiteLLM_TeamTable.model_validate(existing_team.model_dump()),
user_api_key_dict=user_api_key_dict,
)
@ -3872,7 +3876,7 @@ async def unblock_team(
# Verify caller has access to manage this team
await _verify_team_access(
team_obj=LiteLLM_TeamTable(**existing_team.model_dump()),
team_obj=LiteLLM_TeamTable.model_validate(existing_team.model_dump()),
user_api_key_dict=user_api_key_dict,
)
@ -3916,13 +3920,13 @@ async def list_available_teams(
status_code=404,
detail={"error": "User not found"},
)
user_info_correct_type = LiteLLM_UserTable(**user_info.model_dump())
user_info_correct_type = LiteLLM_UserTable.model_validate(user_info.model_dump())
available_teams = [team for team in available_teams if team not in user_info_correct_type.teams]
available_teams_db = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": available_teams}})
available_teams_correct_type = [LiteLLM_TeamTable(**team.model_dump()) for team in available_teams_db]
available_teams_correct_type = [LiteLLM_TeamTable.model_validate(team.model_dump()) for team in available_teams_db]
return available_teams_correct_type
@ -4090,7 +4094,7 @@ def _convert_teams_to_response_models(
team_dict = team.dict()
if use_deleted_table:
team_list.append(LiteLLM_DeletedTeamTable(**team_dict))
team_list.append(LiteLLM_DeletedTeamTable.model_validate(team_dict))
else:
members_with_roles = team_dict.get("members_with_roles")
if not isinstance(members_with_roles, list):
@ -4705,7 +4709,7 @@ async def team_model_add(
detail={"error": f"Team not found, passed team_id={data.team_id}"},
)
team_obj = LiteLLM_TeamTable(**team_row.model_dump())
team_obj = LiteLLM_TeamTable.model_validate(team_row.model_dump())
# Authorization check - only proxy admin, team admin, or org admin can add models
if (
@ -4805,7 +4809,7 @@ async def team_model_delete(
detail={"error": f"Team not found, passed team_id={data.team_id}"},
)
team_obj = LiteLLM_TeamTable(**team_row.model_dump())
team_obj = LiteLLM_TeamTable.model_validate(team_row.model_dump())
# Authorization check - only proxy admin, team admin, or org admin can remove models
if (
@ -4873,7 +4877,7 @@ async def team_member_permissions(
check_db_only=True,
)
complete_team_data = LiteLLM_TeamTable(**existing_team_row.model_dump())
complete_team_data = LiteLLM_TeamTable.model_validate(existing_team_row.model_dump())
# Admin Viewer follows the read-parity rule: see team permissions like
# a Proxy Admin would. Team / org admins keep their existing scope.
@ -4940,7 +4944,7 @@ async def update_team_member_permissions(
check_db_only=True,
)
complete_team_data = LiteLLM_TeamTable(**existing_team_row.model_dump())
complete_team_data = LiteLLM_TeamTable.model_validate(existing_team_row.model_dump())
# Available-team self-join must NOT grant write access to team-wide
# permission policies; only proxy/team/org admins can update them.
@ -5201,7 +5205,7 @@ async def get_team_daily_activity(
if not _user_has_admin_view(user_api_key_dict) and team_ids_list and team_aliases:
has_full_team_view = True
for team_alias in team_aliases:
team_obj = LiteLLM_TeamTable(**team_alias.model_dump())
team_obj = LiteLLM_TeamTable.model_validate(team_alias.model_dump())
is_admin = _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj)
has_perm = _team_member_has_permission(
user_api_key_dict=user_api_key_dict,

View file

@ -14,6 +14,7 @@ from litellm.types.utils import SpecialEnums
if TYPE_CHECKING:
from fastapi import Request
from litellm.proxy._types import UserAPIKeyAuth
from litellm.router import Router
@ -294,9 +295,8 @@ def get_credentials_for_model(
def get_team_provider_credentials(
llm_router: Optional["Router"],
team_models: List[str],
user_api_key_dict: "UserAPIKeyAuth",
custom_llm_provider: str,
team_id: Optional[str] = None,
) -> Optional[dict]:
"""
Resolve upstream credentials for a provider-scoped file operation
@ -304,21 +304,61 @@ def get_team_provider_credentials(
Priority:
1. The team's own (BYOK) deployment for this provider — a deployment whose
``model_info.team_id`` matches ``team_id``. This keeps team-scoped listings
on the team's own provider account/key instead of a shared global one.
2. Fallback: any deployment the team is granted access to for this provider,
expanding wildcard routes and the all-proxy-models sentinel.
``model_info.team_id`` matches the caller's team. This keeps team-scoped
listings on the team's own provider account/key instead of a shared
global one.
2. Fallback: any deployment the caller is granted access to for this
provider, expanding wildcard routes and the all-proxy-models sentinel.
Credential lookup is always scoped to the team's allowlist, so a team can
never resolve a provider key for a deployment it isn't authorized to use.
Credential lookup is scoped to both the team's allowlist and the key's own
model allowlist (``user_api_key_dict.models``), so neither a team nor a
restricted key within a team can resolve a provider key for a deployment
it isn't authorized to use. A key restricted to an explicit model list
only narrows the team scope; sentinel-bearing keys (all-proxy-models /
all-team-models) defer to the team scope instead of widening past it.
Returns None when the router is unavailable or no authorized deployment
matches, so the caller can fall back to default credential resolution.
"""
if llm_router is None:
return None
from litellm.proxy._types import SpecialModelNames
from litellm.proxy.auth.model_checks import get_complete_model_list, get_key_models
team_id = user_api_key_dict.team_id
team_models = user_api_key_dict.team_models or []
proxy_model_list = llm_router.get_model_names(team_id=team_id)
model_access_groups = llm_router.get_model_access_groups()
raw_key_models = user_api_key_dict.models or []
sentinel_values = {
SpecialModelNames.all_proxy_models.value,
SpecialModelNames.all_team_models.value,
}
key_is_restricted = bool(raw_key_models) and not (set(raw_key_models) & sentinel_values)
key_model_allowlist = (
tuple(
dict.fromkeys(
get_key_models(
user_api_key_dict=user_api_key_dict,
proxy_model_list=proxy_model_list,
model_access_groups=model_access_groups,
)
)
)
if key_is_restricted
else ()
)
key_model_allowlist_set = frozenset(key_model_allowlist)
def _key_may_use(public_model_name: Optional[str]) -> bool:
if not key_model_allowlist_set:
return True
return public_model_name is not None and public_model_name in key_model_allowlist_set
def _provider_credentials(model_id: str) -> Optional[dict]:
credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_id)
credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_id, team_id=team_id)
if credentials is not None and credentials.get("custom_llm_provider") == custom_llm_provider:
return credentials
return None
@ -332,27 +372,27 @@ def get_team_provider_credentials(
deployment_id = model_info.get("id")
if deployment_id is None:
continue
if not _key_may_use(model_info.get("team_public_model_name") or deployment.get("model_name")):
continue
credentials = _provider_credentials(deployment_id)
if credentials is not None:
return credentials
# 2. Fall back to deployments the team is allowed to access. The
# all-proxy-models sentinel isn't expanded by get_complete_model_list, so
# normalize it to an empty allowlist, which defers to the team-scoped
# proxy model list. A team with a restricted allowlist (e.g. anthropic
# only) therefore never resolves another provider's key.
from litellm.proxy._types import SpecialModelNames
from litellm.proxy.auth.model_checks import get_complete_model_list
# 2. Fall back to deployments the caller is allowed to access. The key's
# effective allowlist (sentinels and access groups already expanded by
# get_key_models) wins when set; otherwise the team's allowlist applies.
# The all-proxy-models sentinel isn't expanded by
# get_complete_model_list, so normalize it to an empty allowlist, which
# defers to the team-scoped proxy model list. A team or key with a
# restricted allowlist (e.g. anthropic only) therefore never resolves
# another provider's key.
grants_all_models = SpecialModelNames.all_proxy_models.value in team_models
effective_team_models = [] if grants_all_models else team_models
proxy_model_list = llm_router.get_model_names(team_id=team_id)
model_access_groups = llm_router.get_model_access_groups()
models_to_try = list(
dict.fromkeys(
get_complete_model_list(
key_models=[],
key_models=list(key_model_allowlist),
team_models=effective_team_models,
proxy_model_list=proxy_model_list,
user_model=None,
@ -373,6 +413,28 @@ def get_team_provider_credentials(
return None
def apply_team_provider_credentials(
data: dict, # mutable-ok: credentials are merged into the request payload in place, same contract as prepare_data_with_credentials
llm_router: Optional["Router"],
user_api_key_dict: "UserAPIKeyAuth",
custom_llm_provider: str,
) -> None:
"""
Resolve credentials for a provider-only request (no model pinned) via
``get_team_provider_credentials`` and merge them into ``data`` in-place.
Leaves ``data`` untouched when no authorized deployment matches, so the
caller falls back to environment-variable credentials exactly as before.
"""
credentials = get_team_provider_credentials(
llm_router=llm_router,
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
if credentials is None:
return
prepare_data_with_credentials(data=data, credentials=credentials)
def prepare_data_with_credentials(
data: dict,
credentials: dict,

View file

@ -43,10 +43,10 @@ from litellm.litellm_core_utils.cloud_storage_security import (
)
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
apply_team_provider_credentials,
encode_file_id_with_model,
extract_file_creation_params,
get_credentials_for_model,
get_team_provider_credentials,
handle_model_based_routing,
prepare_data_with_credentials,
validate_managed_files_requirement,
@ -253,6 +253,12 @@ async def route_create_file(
_create_file_request=_create_file_request,
)
else:
apply_team_provider_credentials(
data=cast(dict, _create_file_request), # cast-ok: TypedDict is a plain dict at runtime; merged in place
llm_router=llm_router,
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
# get configs for custom_llm_provider
llm_provider_config = get_files_provider_config(custom_llm_provider=custom_llm_provider)
if llm_provider_config is not None:
@ -735,6 +741,14 @@ async def get_file_content(
check_file_id_encoding=True,
)
if not should_route:
apply_team_provider_credentials(
data=data,
llm_router=llm_router,
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
from litellm.proxy.openai_files_endpoints.file_content_streaming_handler import (
FileContentStreamingHandler,
)
@ -983,6 +997,12 @@ async def get_file(
# Remove file_id from data to avoid "multiple values for keyword argument" error
# data was initialized with {"file_id": file_id}
data.pop("file_id", None)
apply_team_provider_credentials(
data=data,
llm_router=llm_router,
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
response = await litellm.afile_retrieve(
custom_llm_provider=custom_llm_provider,
file_id=file_id,
@ -1183,6 +1203,12 @@ async def delete_file(
)
else:
data.pop("file_id", None)
apply_team_provider_credentials(
data=data,
llm_router=llm_router,
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
response = await litellm.afile_delete(
custom_llm_provider=custom_llm_provider,
file_id=file_id,
@ -1354,14 +1380,12 @@ async def list_files(
# No model/target_model_names pinned: resolve upstream credentials from
# the team's deployment for this provider so the call is authenticated
# against the team's own account (e.g. the team's openai deployment).
team_credentials = get_team_provider_credentials(
apply_team_provider_credentials(
data=data,
llm_router=llm_router,
team_models=user_api_key_dict.team_models or [],
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
team_id=user_api_key_dict.team_id,
)
if team_credentials is not None:
prepare_data_with_credentials(data=data, credentials=team_credentials)
response = await litellm.afile_list(
custom_llm_provider=custom_llm_provider,

View file

@ -6758,6 +6758,9 @@ class ProxyConfig:
from litellm.proxy._experimental.mcp_server.oauth2_flow_backfill import (
backfill_null_oauth2_flows,
)
from litellm.proxy._experimental.mcp_server.oauth_issuer_stamp_backfill import (
backfill_discovery_stamped_issuers,
)
try:
if prisma_client is not None:
@ -6767,6 +6770,16 @@ class ProxyConfig:
"litellm.proxy.proxy_server.py::ProxyConfig:_init_mcp_servers_in_db backfill - {}".format(str(e))
)
try:
if prisma_client is not None:
await backfill_discovery_stamped_issuers(prisma_client)
except Exception as e: # noqa: BLE001
verbose_proxy_logger.exception(
"litellm.proxy.proxy_server.py::ProxyConfig:_init_mcp_servers_in_db issuer stamp backfill - {}".format(
str(e)
)
)
try:
await global_mcp_server_manager.reload_servers_from_database()
except Exception as e:
@ -6778,6 +6791,31 @@ class ProxyConfig:
if self._should_load_db_object(object_type="mcp"):
await self._init_mcp_servers_in_db()
async def reload_mcp_servers_from_db(self) -> None:
"""Registry refresh only, for the periodic job in store_model_in_db-off deployments.
Deliberately narrower than ``init_mcp_servers_from_db``: the oauth2_flow backfill is a write
path that only needs to run once at startup, so the cadence here is purely the read-side
reload whose fast-path exemption retries failed OAuth discovery. Gated the same way, so an
admin who excluded mcp from supported_db_objects opts out of this too.
"""
if not self._should_load_db_object(object_type="mcp"):
return
from litellm.proxy._experimental.mcp_server.utils import is_mcp_available
if not is_mcp_available():
return
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
try:
await global_mcp_server_manager.reload_servers_from_database()
except Exception as e: # noqa: BLE001 # scheduled job: a reload failure must not kill the recurring retry
verbose_proxy_logger.exception(
"litellm.proxy.proxy_server.py::ProxyConfig:reload_mcp_servers_from_db - {}".format(str(e))
)
async def _init_agents_in_db(self, prisma_client: PrismaClient):
from litellm.proxy.agent_endpoints.agent_registry import (
global_agent_registry as AGENT_REGISTRY,
@ -8099,6 +8137,22 @@ class ProxyStartupEvent:
if store_model_in_db is not True:
await proxy_config.init_mcp_servers_from_db()
if prisma_client is not None:
# DB-backed MCP servers are live objects in every mode, so the registry refresh that
# store_model_in_db=True deployments get via the add_deployment job must run here
# too; without it, a server whose OAuth discovery failed at startup is rebuilt only
# by a management write, since the reload fast path is the retry's only driver.
mcp_reload_interval_seconds = proxy_config_reload_interval_seconds
if not isinstance(mcp_reload_interval_seconds, int) or mcp_reload_interval_seconds <= 0:
mcp_reload_interval_seconds = 30
scheduler.add_job(
proxy_config.reload_mcp_servers_from_db,
"interval",
seconds=mcp_reload_interval_seconds,
id="reload_mcp_servers_job",
replace_existing=True,
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
)
await cls._initialize_slack_alerting_jobs(
scheduler=scheduler,
@ -11214,11 +11268,15 @@ async def get_all_team_models(
if user_teams == "*":
team_db_objects = await TeamRepository(prisma_client).table.find_many()
team_db_objects_typed = [LiteLLM_TeamTable(**team_db_object.model_dump()) for team_db_object in team_db_objects]
team_db_objects_typed = [
LiteLLM_TeamTable.model_validate(team_db_object.model_dump()) for team_db_object in team_db_objects
]
else:
team_db_objects = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": user_teams}})
team_db_objects_typed = [LiteLLM_TeamTable(**team_db_object.model_dump()) for team_db_object in team_db_objects]
team_db_objects_typed = [
LiteLLM_TeamTable.model_validate(team_db_object.model_dump()) for team_db_object in team_db_objects
]
team_models = _add_team_models_to_all_models(
team_db_objects_typed=team_db_objects_typed,
@ -11292,7 +11350,7 @@ async def _populate_team_access_on_models(
where={"user_id": user_api_key_dict.user_id}
)
if user_db_object is not None:
user_object = LiteLLM_UserTable(**user_db_object.model_dump())
user_object = LiteLLM_UserTable.model_validate(user_db_object.model_dump())
user_teams = user_object.teams or []
direct_access_models = get_direct_access_models(
user_db_object=user_object,
@ -11827,7 +11885,7 @@ async def _load_team_object_for_model_filter(team_id: str, prisma_client: Prisma
if team_db_object is None:
verbose_proxy_logger.warning(f"Team {team_id} not found in database")
return None
return LiteLLM_TeamTable(**team_db_object.model_dump())
return LiteLLM_TeamTable.model_validate(team_db_object.model_dump())
except Exception as e:
verbose_proxy_logger.exception(f"Error fetching team {team_id}: {str(e)}")
return None

View file

@ -3308,8 +3308,36 @@ async def ui_view_session_spend_logs(
detail="Database not connected",
)
# Build query conditions
where_conditions = {"session_id": session_id}
if _is_admin_view_safe(user_api_key_dict=user_api_key_dict):
scope_sql = ""
scope_params = ()
where_conditions = {"session_id": session_id}
else:
try:
permitted_team_ids = (
await _get_permitted_team_ids_for_spend_logs(
prisma_client=prisma_client,
user_api_key_dict=user_api_key_dict,
)
if _can_user_view_spend_log(user_api_key_dict=user_api_key_dict)
else []
)
except Exception: # noqa: BLE001 # mirror /spend/logs/ui: failed team lookup falls back to own-logs-only scope
permitted_team_ids = []
if permitted_team_ids:
scope_sql = ' AND ("user" = $4 OR team_id = ANY($5::text[]))'
scope_params = (user_api_key_dict.user_id, permitted_team_ids)
where_conditions = {
"session_id": session_id,
"OR": [
{"user": user_api_key_dict.user_id},
{"team_id": {"in": permitted_team_ids}},
],
}
else:
scope_sql = ' AND "user" = $4'
scope_params = (user_api_key_dict.user_id,)
where_conditions = {"session_id": session_id, "user": user_api_key_dict.user_id}
# Calculate pagination offsets
skip = (page - 1) * page_size
@ -3318,7 +3346,7 @@ async def ui_view_session_spend_logs(
total_records = await SpendLogsRepository(prisma_client).table.count(where=where_conditions)
# Query with raw SQL to exclude heavy columns (messages, response, proxy_server_request)
sql_query = """
sql_query = f"""
SELECT
request_id, call_type, api_key, spend, total_tokens,
prompt_tokens, completion_tokens, "startTime", "endTime",
@ -3328,11 +3356,11 @@ async def ui_view_session_spend_logs(
organization_id, end_user, requester_ip_address,
session_id, status, mcp_namespaced_tool_name, agent_id
FROM "LiteLLM_SpendLogs"
WHERE session_id = $1
WHERE session_id = $1{scope_sql}
ORDER BY "startTime" DESC
LIMIT $2 OFFSET $3
"""
result = await prisma_client.db.query_raw(sql_query, session_id, page_size, skip)
result = await prisma_client.db.query_raw(sql_query, session_id, page_size, skip, *scope_params)
total_pages = (total_records + page_size - 1) // page_size
@ -3548,7 +3576,7 @@ async def _can_team_member_view_log(
team_row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id})
if team_row is None:
return False
team_obj = LiteLLM_TeamTable(**team_row.model_dump())
team_obj = LiteLLM_TeamTable.model_validate(team_row.model_dump())
if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj):
return True
return _team_member_has_permission(
@ -3640,7 +3668,7 @@ async def _get_permitted_team_ids_for_spend_logs(
permitted: List[str] = []
for team_row in team_rows:
team_obj = LiteLLM_TeamTable(**team_row.model_dump())
team_obj = LiteLLM_TeamTable.model_validate(team_row.model_dump())
if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj):
permitted.append(team_obj.team_id)
elif _team_member_has_permission(

View file

@ -19,6 +19,7 @@ import threading
import time
import traceback
from collections import defaultdict
from collections.abc import Mapping
from functools import lru_cache
from typing import (
TYPE_CHECKING,
@ -109,6 +110,9 @@ from litellm.router_utils.batch_utils import (
replace_model_in_jsonl,
should_replace_model_in_jsonl,
)
from litellm.router_utils.auto_router_model_naming import (
classify_strategy_router_model,
)
from litellm.router_utils.client_initalization_utils import InitalizeCachedClient
from litellm.router_utils.clientside_credential_handler import (
get_dynamic_litellm_params,
@ -270,6 +274,43 @@ def _cost_value_as_float(value: Union[str, int, float, None]) -> Optional[float]
return None
def model_info_is_active_for_environment(model_info: Mapping[str, object] | None) -> bool:
"""Single owner of the environment-gating rule: a deployment whose model_info names
`supported_environments` loads only on pods whose LITELLM_ENVIRONMENT is in that list.
`Router.deployment_is_active_for_environment` delegates here, and the model-write
endpoints consult the same rule to tell a deliberately inactive model from one that
was dropped by a failed reload."""
if model_info is None:
return True
supported_environments = model_info.get("supported_environments")
if supported_environments is None:
return True
if not isinstance(supported_environments, (list, tuple)):
raise ValueError(
f"supported_environments must be a list of {VALID_LITELLM_ENVIRONMENTS}. "
f"but set as: {supported_environments} for model_info: {model_info}"
)
litellm_environment = get_secret_str(secret_name="LITELLM_ENVIRONMENT")
if litellm_environment is None:
raise ValueError("Set 'supported_environments' for model but not 'LITELLM_ENVIRONMENT' set in .env")
if litellm_environment not in VALID_LITELLM_ENVIRONMENTS:
raise ValueError(
f"LITELLM_ENVIRONMENT must be one of {VALID_LITELLM_ENVIRONMENTS}. but set as: {litellm_environment}"
)
for _env in supported_environments:
if _env not in VALID_LITELLM_ENVIRONMENTS:
raise ValueError(
f"supported_environments must be one of {VALID_LITELLM_ENVIRONMENTS}. but set as: {_env} "
f"for model_info: {model_info}"
)
if litellm_environment in supported_environments:
return True
return False
_PreRoutingStrategyT = TypeVar("_PreRoutingStrategyT")
@ -7585,15 +7626,7 @@ class Router:
but NOT "auto_router/complexity_router" or "auto_router/adaptive_router"
(which use the complexity-router and adaptive-router strategies).
"""
if litellm_params.model.startswith("auto_router/complexity_router"):
return False # This is handled by complexity_router
if litellm_params.model.startswith("auto_router/adaptive_router"):
return False # This is handled by adaptive_router
if litellm_params.model.startswith("auto_router/quality_router"):
return False # This is handled by quality_router
if litellm_params.model.startswith("auto_router/"):
return True
return False
return classify_strategy_router_model(litellm_params.model) == "semantic"
@staticmethod
def _deployment_tags(deployment: Deployment) -> tuple[str, ...]:
@ -7648,9 +7681,7 @@ class Router:
Returns True if the litellm_params model starts with "auto_router/complexity_router"
"""
if litellm_params.model.startswith("auto_router/complexity_router"):
return True
return False
return classify_strategy_router_model(litellm_params.model) == "complexity"
def init_complexity_router_deployment(self, deployment: Deployment):
"""
@ -7700,7 +7731,7 @@ class Router:
def _is_adaptive_router_deployment(self, litellm_params: LiteLLM_Params) -> bool:
"""True when this deployment opts in via the `auto_router/adaptive_router` model prefix."""
return litellm_params.model.startswith("auto_router/adaptive_router")
return classify_strategy_router_model(litellm_params.model) == "adaptive"
def _deployment_participates_in_adaptive_routing(self, litellm_params: LiteLLM_Params) -> bool:
"""True when this deployment owns an `adaptive_routers` entry once finalized:
@ -7926,9 +7957,7 @@ class Router:
Returns True if the litellm_params model starts with "auto_router/quality_router".
"""
if litellm_params.model.startswith("auto_router/quality_router"):
return True
return False
return classify_strategy_router_model(litellm_params.model) == "quality"
def init_quality_router_deployment(self, deployment: Deployment):
"""
@ -7982,30 +8011,7 @@ class Router:
- ValueError: If LITELLM_ENVIRONMENT is not set in .env or not one of the valid values
- ValueError: If supported_environments is not set in model_info or not one of the valid values
"""
if (
deployment.model_info is None
or "supported_environments" not in deployment.model_info
or deployment.model_info["supported_environments"] is None
):
return True
litellm_environment = get_secret_str(secret_name="LITELLM_ENVIRONMENT")
if litellm_environment is None:
raise ValueError("Set 'supported_environments' for model but not 'LITELLM_ENVIRONMENT' set in .env")
if litellm_environment not in VALID_LITELLM_ENVIRONMENTS:
raise ValueError(
f"LITELLM_ENVIRONMENT must be one of {VALID_LITELLM_ENVIRONMENTS}. but set as: {litellm_environment}"
)
for _env in deployment.model_info["supported_environments"]:
if _env not in VALID_LITELLM_ENVIRONMENTS:
raise ValueError(
f"supported_environments must be one of {VALID_LITELLM_ENVIRONMENTS}. but set as: {_env} for deployment: {deployment}"
)
if litellm_environment in deployment.model_info["supported_environments"]:
return True
return False
return model_info_is_active_for_environment(model_info=deployment.model_info)
def set_model_list(self, model_list: list):
original_model_list = copy.deepcopy(model_list)
@ -8630,6 +8636,33 @@ class Router:
raise Exception("Model Name invalid - {}".format(type(model)))
return None
@staticmethod
def _deployment_usable_by_team(model: Union[Mapping, Deployment], team_id: str | None) -> bool:
"""
A team-scoped deployment (``model_info.team_id`` set) is only usable by
callers from that same team; deployments without a team owner are shared.
"""
model_info = model.get("model_info") if isinstance(model, dict) else model.model_info
owner_team_id = model_info.get("team_id") if model_info is not None else None
return owner_team_id is None or owner_team_id == team_id
def _get_model_group_deployment_usable_by_team(
self, model_group_name: str, team_id: str | None
) -> Deployment | None:
"""
Like ``get_deployment_by_model_group_name``, but skips deployments owned
by other teams so a shared model name never resolves another team's
credentials.
"""
indices = self.model_name_to_deployment_indices.get(model_group_name) or ()
usable = (
self.model_list[idx] for idx in indices if self._deployment_usable_by_team(self.model_list[idx], team_id)
)
first_usable = next(usable, None)
if first_usable is None:
return None
return Deployment(**first_usable) if isinstance(first_usable, dict) else first_usable
def get_configured_token_limits(self, model_name: str) -> "tuple[int | None, int | None]":
"""
Return (max_input_tokens, max_output_tokens) explicitly configured in a concrete
@ -8664,7 +8697,10 @@ class Router:
model_id: Model ID or model name from model_list (e.g., "gpt-4o-litellm")
team_id: Optional team id of the caller. When set, team-scoped
deployments (indexed by team public model name, including team
wildcard models like "openai/*") are also considered.
wildcard models like "openai/*") are also considered. Name and
wildcard lookups never resolve a deployment owned by a
different team, so shared model names can't leak another
team's credentials.
Returns:
Dictionary containing api_key, api_base, custom_llm_provider, etc.
@ -8681,7 +8717,7 @@ class Router:
# If not found, try by model_group_name
if deployment is None:
deployment = self.get_deployment_by_model_group_name(model_group_name=model_id)
deployment = self._get_model_group_deployment_usable_by_team(model_group_name=model_id, team_id=team_id)
# If not found, check team-scoped deployments whose team public model
# name exactly matches model_id (wildcard team names are matched via
@ -8698,7 +8734,12 @@ class Router:
if deployment is None:
team_pattern_router = self.team_pattern_routers.get(team_id) if team_id is not None else None
team_wildcard_models = (team_pattern_router.route(model_id) or []) if team_pattern_router else []
potential_wildcard_models = team_wildcard_models or self.pattern_router.route(model_id) or []
global_wildcard_models = [
wildcard_model
for wildcard_model in (self.pattern_router.route(model_id) or [])
if self._deployment_usable_by_team(wildcard_model, team_id)
]
potential_wildcard_models = team_wildcard_models or global_wildcard_models
if potential_wildcard_models:
# Use the first matching wildcard deployment
deployment_dict = potential_wildcard_models[0]
@ -9519,7 +9560,12 @@ class Router:
return None
# Strategy 1: Check if model_id directly matches a model_name or deployment ID
if model_id in self.model_names or self.has_model_id(model_id):
if model_id in self.model_names:
return model_id
if self.has_model_id(model_id):
deployment = self.get_deployment(model_id=model_id)
if deployment is not None and deployment.model_name:
return deployment.model_name
return model_id
# Strategy 2: Search through router's model_list to find by litellm_params.model

View file

@ -18,6 +18,7 @@ from __future__ import annotations
import asyncio
import random
import re
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Literal, Union, cast
from pydantic import BaseModel
@ -25,6 +26,7 @@ from pydantic import BaseModel
from litellm._logging import verbose_router_logger
from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY
from litellm.integrations.custom_logger import CustomLogger
from litellm.llms.base_llm.base_utils import type_to_response_format_param
from litellm.types.utils import ModelResponse
from .config import (
@ -112,6 +114,16 @@ def _classifier_call_metadata(metadata: dict[str, Any] | None) -> dict[str, Any]
}
def _effective_turn_off_message_logging(request_kwargs: Mapping[str, Any] | None) -> bool | None:
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
initialize_standard_callback_dynamic_params,
)
return initialize_standard_callback_dynamic_params(dict(request_kwargs) if request_kwargs else {}).get(
"turn_off_message_logging"
)
class DimensionScore:
"""Represents a score for a single dimension with optional signal."""
@ -427,7 +439,17 @@ class ComplexityRouter(CustomLogger):
# attributed to the calling key/team instead of being dropped. Excludes the
# parent request's budget reservation, which the routed completion (not this
# internal classifier call) is responsible for reconciling.
metadata = _classifier_call_metadata((request_kwargs or {}).get("litellm_metadata"))
request_metadata = (request_kwargs or {}).get("litellm_metadata") or (request_kwargs or {}).get("metadata")
metadata = _classifier_call_metadata(request_metadata)
turn_off_message_logging = _effective_turn_off_message_logging(request_kwargs)
proxy_server_request = {
"body": {
"model": llm_config.model,
"messages": [{"role": "user", "content": classification_prompt}],
"response_format": type_to_response_format_param(TierClassification),
}
}
response: ModelResponse = await self.litellm_router_instance.acompletion(
model=llm_config.model,
@ -435,6 +457,8 @@ class ComplexityRouter(CustomLogger):
response_format=TierClassification,
timeout=llm_config.timeout_ms / 1000,
metadata=metadata,
proxy_server_request=proxy_server_request,
turn_off_message_logging=turn_off_message_logging,
)
content = response.choices[0].message.content
if not content:
@ -821,8 +845,16 @@ class ComplexityRouter(CustomLogger):
# key/team budget. Key/team attribution fields are preserved for spend logging.
metadata = _classifier_call_metadata(request_kwargs.get("metadata"))
litellm_metadata = _classifier_call_metadata(request_kwargs.get("litellm_metadata"))
turn_off_message_logging = _effective_turn_off_message_logging(request_kwargs)
proxy_server_request = {"body": {"model": self.config.embedding_model, "input": [user_message]}}
query_vector = (
await encoder.aencode_queries([user_message], metadata=metadata, litellm_metadata=litellm_metadata)
await encoder.aencode_queries(
[user_message],
metadata=metadata,
litellm_metadata=litellm_metadata,
proxy_server_request=proxy_server_request,
turn_off_message_logging=turn_off_message_logging,
)
)[0]
route_choice = await routelayer.acall(vector=query_vector)

View file

@ -73,13 +73,20 @@ class LowestLatencyLoggingHandler(CustomLogger):
precise_minute = f"{current_date}-{current_hour}-{current_minute}"
response_ms = end_time - start_time
if isinstance(response_ms, timedelta):
# normalize to float seconds up-front: non-chat responses
# (embeddings, speech, image) skip the ModelResponse branch
# below, and a raw timedelta appended to the latency list
# breaks JSON serialization when the router cache syncs to
# Redis (issue #33169)
response_ms = response_ms.total_seconds()
time_to_first_token_response_time = None
if kwargs.get("stream", None) is not None and kwargs["stream"] is True:
# only log ttft for streaming request
time_to_first_token_response_time = kwargs.get("completion_start_time", end_time) - start_time
final_value: Union[float, timedelta] = response_ms
final_value: float = response_ms
time_to_first_token: Optional[float] = None
total_tokens = 0
@ -89,15 +96,12 @@ class LowestLatencyLoggingHandler(CustomLogger):
completion_tokens = _usage.completion_tokens
total_tokens = _usage.total_tokens
# Handle both timedelta and float response times
if isinstance(response_ms, timedelta):
response_seconds = response_ms.total_seconds()
else:
response_seconds = response_ms
# response_ms is already normalized to float seconds above
response_seconds = response_ms
final_value = safe_divide_seconds(response_seconds, completion_tokens)
if final_value is not None:
final_value = float(final_value)
normalized_value = safe_divide_seconds(response_seconds, completion_tokens)
if normalized_value is not None:
final_value = float(normalized_value)
else:
final_value = response_seconds
@ -262,12 +266,19 @@ class LowestLatencyLoggingHandler(CustomLogger):
precise_minute = f"{current_date}-{current_hour}-{current_minute}"
response_ms = end_time - start_time
if isinstance(response_ms, timedelta):
# normalize to float seconds up-front: non-chat responses
# (embeddings, speech, image) skip the ModelResponse branch
# below, and a raw timedelta appended to the latency list
# breaks JSON serialization when the router cache syncs to
# Redis (issue #33169)
response_ms = response_ms.total_seconds()
time_to_first_token_response_time = None
if kwargs.get("stream", None) is not None and kwargs["stream"] is True:
# only log ttft for streaming request
time_to_first_token_response_time = kwargs.get("completion_start_time", end_time) - start_time
final_value: Union[float, timedelta] = response_ms
final_value: float = response_ms
total_tokens = 0
time_to_first_token: Optional[float] = None
@ -277,17 +288,14 @@ class LowestLatencyLoggingHandler(CustomLogger):
completion_tokens = _usage.completion_tokens
total_tokens = _usage.total_tokens
# Handle both timedelta and float response times
if isinstance(response_ms, timedelta):
response_seconds = response_ms.total_seconds()
else:
response_seconds = response_ms
# response_ms is already normalized to float seconds above
response_seconds = response_ms
final_value = safe_divide_seconds(response_seconds, completion_tokens)
if final_value is not None:
final_value = float(final_value)
normalized_value = safe_divide_seconds(response_seconds, completion_tokens)
if normalized_value is not None:
final_value = float(normalized_value)
else:
final_value = response_ms
final_value = response_seconds
if time_to_first_token_response_time is not None:
if isinstance(time_to_first_token_response_time, timedelta):

View file

@ -0,0 +1,101 @@
"""Naming contract for strategy-router (auto-router) pseudo-models.
A deployment whose ``litellm_params.model`` starts with ``auto_router/`` does not
name a provider model; the string is the discriminator that selects which
pre-routing strategy owns the deployment. This module is the single source of
truth for classifying that string (``Router._is_*_router_deployment`` delegates
here) and for checking that a client-supplied write leaves the deployment
coherent, so management endpoints can reject corruption with a 400 instead of
the router silently dropping the deployment at load time under
``ignore_invalid_deployments``.
"""
from typing import Literal, Mapping
AUTO_ROUTER_MODEL_PREFIX = "auto_router/"
StrategyRouterKind = Literal["semantic", "complexity", "adaptive", "quality"]
STRATEGY_ROUTER_PARAM_FIELDS: frozenset[str] = frozenset(
{
"auto_router_config",
"auto_router_config_path",
"auto_router_default_model",
"auto_router_embedding_model",
"complexity_router_config",
"complexity_router_default_model",
"adaptive_router_config",
"quality_router_config",
"quality_router_default_model",
}
)
_REQUIRED_FIELD_GROUPS: Mapping[StrategyRouterKind, tuple[tuple[str, ...], ...]] = {
"semantic": (
("auto_router_config", "auto_router_config_path"),
("auto_router_default_model",),
("auto_router_embedding_model",),
),
"complexity": (("complexity_router_config", "complexity_router_default_model"),),
"adaptive": (("adaptive_router_config",),),
"quality": (("quality_router_config", "quality_router_default_model"),),
}
def classify_strategy_router_model(model: str) -> StrategyRouterKind | None:
"""Classify a ``litellm_params.model`` string the way the Router does.
Returns None for regular provider models. Mirrors Router registration
exactly: reserved names are matched by prefix, everything else under
``auto_router/`` is a semantic router.
"""
if not model.startswith(AUTO_ROUTER_MODEL_PREFIX):
return None
remainder = model[len(AUTO_ROUTER_MODEL_PREFIX) :]
if remainder.startswith("complexity_router"):
return "complexity"
if remainder.startswith("adaptive_router"):
return "adaptive"
if remainder.startswith("quality_router"):
return "quality"
return "semantic"
def validate_strategy_router_model_write(model: str, present_fields: frozenset[str]) -> str | None:
"""Check that writing ``model`` leaves a deployment the router can load.
``present_fields`` is the set of strategy-router param fields that are
non-None on the deployment after the write (stored fields merged with the
incoming ones). Returns a human-readable violation, or None when coherent.
"""
kind = classify_strategy_router_model(model)
if kind is None:
offending = sorted(present_fields & STRATEGY_ROUTER_PARAM_FIELDS)
if offending:
return (
f"litellm_params.model='{model}' does not start with '{AUTO_ROUTER_MODEL_PREFIX}' but the "
f"deployment carries auto-router settings ({', '.join(offending)}), so the router could not "
f"load it. Keep the '{AUTO_ROUTER_MODEL_PREFIX}' prefix; to change the name clients call, "
"edit the public model_name instead."
)
return None
remainder = model[len(AUTO_ROUTER_MODEL_PREFIX) :]
if remainder.startswith(AUTO_ROUTER_MODEL_PREFIX):
return (
f"litellm_params.model='{model}' repeats the '{AUTO_ROUTER_MODEL_PREFIX}' prefix, so the router "
f"could not load it. Use '{remainder}'; to change the name clients call, edit the public "
"model_name instead."
)
if not remainder:
return (
f"litellm_params.model='{model}' is missing the router name after the '{AUTO_ROUTER_MODEL_PREFIX}' prefix."
)
missing = tuple(
" or ".join(group) for group in _REQUIRED_FIELD_GROUPS[kind] if not any(f in present_fields for f in group)
)
if missing:
return (
f"litellm_params.model='{model}' selects the {kind} router, which requires "
f"{'; '.join(missing)} in litellm_params."
)
return None

View file

@ -16,7 +16,7 @@ GeminiEmbeddingInput = Union[EmbeddingInput, List[List[str]]]
class FunctionResponse(TypedDict, total=False):
# `id` correlates this response with the originating `functionCall` part.
# Supported on Google AI Studio Gemini 3.5+; Vertex AI rejects this field.
# Supported on Gemini 3+; older Gemini models reject this field.
id: str
name: Required[str]
response: Optional[dict]
@ -24,8 +24,8 @@ class FunctionResponse(TypedDict, total=False):
class FunctionCall(TypedDict, total=False):
# `id` correlates the corresponding `functionResponse` on Google AI Studio
# Gemini 3.5+. Vertex AI and older Gemini models omit/reject this field.
# `id` correlates the corresponding `functionResponse` on Gemini 3+.
# Older Gemini models omit/reject this field.
id: str
name: Required[str]
args: Optional[dict]
@ -58,8 +58,8 @@ class PartType(TypedDict, total=False):
class HttpxFunctionCall(TypedDict, total=False):
# `id` correlates the corresponding `functionResponse` on Google AI Studio
# Gemini 3.5+. Vertex AI and older Gemini models omit/reject this field.
# `id` correlates the corresponding `functionResponse` on Gemini 3+.
# Older Gemini models omit/reject this field.
id: str
name: Required[str]
args: dict

View file

@ -18,6 +18,9 @@ SCIM_ENTERPRISE_METADATA_KEY = "scim_enterprise"
SCIM_ENTITLEMENTS_METADATA_KEY = "scim_entitlements"
SCIM_ROLES_METADATA_KEY = "scim_roles"
SCIM_MANAGED_TEAM_METADATA_KEY = "scim_managed"
SCIM_TEAM_DATA_METADATA_KEY = "scim_data"
class LiteLLM_UserScimMetadata(BaseModel):
"""
@ -131,6 +134,15 @@ class SCIMUser(SCIMResource):
class SCIMMember(BaseModel):
value: str # User ID
display: Optional[str] = None # Username or email
type: str | None = None
@field_validator("type", mode="before")
@classmethod
def normalize_type(cls, v: object) -> str | None:
"""Anything that is not a string carries no canonical type, and rejecting the
request over it would be a regression: before this field existed the value was
parsed away silently."""
return v if isinstance(v, str) else None
class SCIMGroup(SCIMResource):

View file

@ -13567,6 +13567,56 @@
}
]
},
"dashscope/qwen3.7-max": {
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 2.5e-06,
"litellm_provider": "dashscope",
"max_input_tokens": 991808,
"max_output_tokens": 65536,
"max_tokens": 65536,
"mode": "chat",
"output_cost_per_token": 7.5e-06,
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"dashscope/qwen3.7-plus": {
"litellm_provider": "dashscope",
"max_input_tokens": 991808,
"max_output_tokens": 65536,
"max_tokens": 65536,
"mode": "chat",
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"tiered_pricing": [
{
"cache_read_input_token_cost": 8e-08,
"input_cost_per_token": 4e-07,
"output_cost_per_token": 1.6e-06,
"range": [
0,
256000.0
]
},
{
"cache_read_input_token_cost": 2.4e-07,
"input_cost_per_token": 1.2e-06,
"output_cost_per_token": 4.8e-06,
"range": [
256000.0,
1000000.0
]
}
]
},
"dashscope/qwq-plus": {
"input_cost_per_token": 8e-07,
"litellm_provider": "dashscope",

View file

@ -24,7 +24,7 @@
"limit": 130
},
"ANN401": {
"limit": 2015
"limit": 2013
},
"ASYNC230": {
"limit": 14

View file

@ -63,16 +63,6 @@ class OpenAIModerationParamsBody(GuardrailParamsBase):
model: str | None = None
class PresidioParamsBody(GuardrailParamsBase):
guardrail: Literal["presidio"] = "presidio"
presidio_analyzer_api_base: str | None = None
presidio_anonymizer_api_base: str | None = None
# apply_to_output masks PII the model itself emitted, which also makes the
# guardrail run post_call. logging_only masks what the proxy logs.
apply_to_output: bool | None = None
logging_only: bool | None = None
class BlockCodeExecutionParamsBody(GuardrailParamsBase):
guardrail: Literal["block_code_execution"] = "block_code_execution"
@ -81,7 +71,6 @@ GuardrailParamsBody = (
ContentFilterParamsBody
| BedrockGuardrailParamsBody
| OpenAIModerationParamsBody
| PresidioParamsBody
| BlockCodeExecutionParamsBody
)

View file

@ -1,141 +0,0 @@
"""Live e2e: the built-in Presidio PII guardrail masks PII on the request and on
the model output.
Presidio replaces detected PII with `<ENTITY_TYPE>` placeholders (e.g.
`<EMAIL_ADDRESS>`) via a real analyzer + anonymizer. Two modes are checked
independently, each opted into per request (default_on=False) so it never touches
unrelated traffic:
- pre_call: the prompt is anonymized before it reaches the model, so a
repeat-verbatim request comes back with the placeholder, never the raw email
- post_call (apply_to_output): PII the model itself emits is masked on the way
out, so the caller never receives the raw value the model produced
A third mode, logging_only, is not covered here: the raw email stayed in the OTEL
span's `gen_ai.input.messages` on every attempt over a full poll deadline while
these two modes masked correctly, so that cell is tracked in LIT-4841 rather than
asserted against known-failing behavior.
Analyzer/anonymizer bases come from PRESIDIO_ANALYZER_API_BASE /
PRESIDIO_ANONYMIZER_API_BASE (compose provides the in-network hosts; point them at
locally published container ports for a host run). The chat backend is a gemini
deployment created for the test.
"""
from __future__ import annotations
import os
import time
from collections.abc import Callable
import pytest
from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, unique_marker
from e2e_http import unwrap
from guardrails_client import GuardrailMode, GuardrailsClient, PresidioParamsBody
from lifecycle import ResourceManager
from models import ChatResponse
pytestmark = pytest.mark.e2e
RAW_EMAIL = "alice.example.person@example.com"
PLACEHOLDER = "<EMAIL_ADDRESS>"
ECHO_REQUEST = f"Repeat the following text back exactly, verbatim, with no changes: My email is {RAW_EMAIL}"
EMIT_REQUEST = f"Output exactly this one line and nothing else: Please contact {RAW_EMAIL} today"
def _content(response: ChatResponse) -> str:
if not response.choices:
return ""
message = response.choices[0].message
return (message.content if message else None) or ""
def _presidio_params(
mode: GuardrailMode, *, apply_to_output: bool = False, logging_only: bool = False
) -> PresidioParamsBody:
analyzer = os.environ["PRESIDIO_ANALYZER_API_BASE"]
anonymizer = os.environ["PRESIDIO_ANONYMIZER_API_BASE"]
return PresidioParamsBody(
mode=mode,
default_on=False,
presidio_analyzer_api_base=analyzer,
presidio_anonymizer_api_base=anonymizer,
apply_to_output=apply_to_output,
logging_only=logging_only,
)
def _poll_until_masked(call: Callable[[], str]) -> str:
"""Retry a call until the guardrail masks its PII, returning the last content.
Registering a guardrail is a control-plane write; the data-plane worker that
serves /chat/completions only picks it up on its next periodic DB sync (~30s
in proxy_server.py), so a call issued the instant after the create runs
against a worker that has no guardrail yet and passes the raw value through.
That is in-flight propagation, not a masking failure. Polling to the deadline
waits it out, so the assertions that follow judge the synced state; if the
mask never lands the last unmasked content is returned and they still fail.
"""
deadline = time.monotonic() + POLL_TIMEOUT
last = call()
while time.monotonic() < deadline:
if PLACEHOLDER in last and RAW_EMAIL not in last:
return last
time.sleep(POLL_INTERVAL)
last = call()
return last
class TestPresidioGuardrail:
@pytest.mark.covers(
"guardrail.presidio.pre_call.masks",
exercised_on=["chat_completions"],
)
def test_pre_call_masks_pii_before_the_model_sees_it(
self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str
) -> None:
model = client.create_backend_model(resources, prefix="e2e-presidio-pre")
name = f"e2e-presidio-pre-{unique_marker()}"
guardrail_id = client.register(name, _presidio_params("pre_call"))
resources.defer(lambda: client.delete_guardrail(guardrail_id))
echoed = _poll_until_masked(
lambda: _content(
unwrap(client.chat(scoped_key, model, ECHO_REQUEST, guardrails=[name], max_tokens=128))
)
)
assert RAW_EMAIL not in echoed, (
"pre_call masking must strip the raw email before the model sees it, but the "
f"model echoed it back: {echoed[:300]!r}"
)
assert PLACEHOLDER in echoed, (
"the model should have echoed the masked placeholder the guardrail substituted, "
f"got: {echoed[:300]!r}"
)
@pytest.mark.covers(
"guardrail.presidio.post_call.masks",
exercised_on=["chat_completions"],
)
def test_post_call_masks_pii_in_model_output(
self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str
) -> None:
model = client.create_backend_model(resources, prefix="e2e-presidio-post")
name = f"e2e-presidio-post-{unique_marker()}"
guardrail_id = client.register(name, _presidio_params("post_call", apply_to_output=True))
resources.defer(lambda: client.delete_guardrail(guardrail_id))
out = _poll_until_masked(
lambda: _content(
unwrap(client.chat(scoped_key, model, EMIT_REQUEST, guardrails=[name], max_tokens=128))
)
)
assert RAW_EMAIL not in out, (
"post_call masking must strip PII the model emitted, but the raw email reached the "
f"caller: {out[:300]!r}"
)
assert PLACEHOLDER in out, (
f"the masked placeholder should replace the model's PII output, got: {out[:300]!r}"
)

View file

@ -11,13 +11,14 @@ request/response bodies are co-located here because only this suite speaks MCP.
from __future__ import annotations
import re
import time
from collections.abc import Mapping
from dataclasses import dataclass
from pydantic import BaseModel, ConfigDict, Field, RootModel
from e2e_http import Headers, NoBody, Result, Success, unwrap
from e2e_http import Headers, NoBody, Result, Success, UnknownApiError, unwrap
from models import KeyGenerateBody, ObjectPermission
from proxy_client import ProxyClient
@ -270,6 +271,60 @@ class McpClient:
)
time.sleep(self.proxy.poll_interval)
def await_call_tool(
self,
key: str,
*,
server_id: str,
name: str,
arguments: McpToolArguments,
) -> McpCallToolResponse:
"""Poll tools/call until the result is not a multi-worker registry miss.
Retries only on the gateway's own cold-worker 500 shapes (Tool <name>
not found / server_not_found). Upstream tool errors and other 500s fail
immediately so non-idempotent calls are not repeated.
"""
deadline = time.monotonic() + self.proxy.poll_timeout
last: Result[McpCallToolResponse] | None = None
while True:
last = self.call_tool(key, server_id=server_id, name=name, arguments=arguments)
if not _is_mcp_not_synced(last, tool_name=name):
return unwrap(last)
if time.monotonic() >= deadline:
raise AssertionError(
f"tools/call for {name!r} on server {server_id} still missing on the "
f"data plane after {self.proxy.poll_timeout}s (multi-worker registry lag); "
f"last result: {last}"
)
time.sleep(self.proxy.poll_interval)
def await_call_tool_denied(
self,
key: str,
*,
server_id: str,
name: str,
arguments: McpToolArguments,
) -> UnknownApiError:
"""Poll tools/call until a cold-worker miss clears and the call is 403 access_denied."""
deadline = time.monotonic() + self.proxy.poll_timeout
last: Result[McpCallToolResponse] | None = None
while True:
last = self.call_tool(key, server_id=server_id, name=name, arguments=arguments)
if isinstance(last, UnknownApiError) and last.status_code == 403:
return last
if not _is_mcp_not_synced(last, tool_name=name):
raise AssertionError(
f"ungranted key's tools/call was not 403 access_denied: {last}"
)
if time.monotonic() >= deadline:
raise AssertionError(
f"ungranted key never got 403 for {name!r} within {self.proxy.poll_timeout}s; "
f"last result: {last}"
)
time.sleep(self.proxy.poll_interval)
def register_mcp_content_filter(self, *, name: str, blocked_keyword: str) -> str:
"""Register a default-on content-filter guardrail that runs on the MCP
tool-call hook (pre_mcp_call) and blocks a single keyword. The keyword is
@ -317,5 +372,39 @@ class McpClient:
)
def _is_mcp_not_synced(
result: Result[McpCallToolResponse],
*,
tool_name: str | None = None,
) -> bool:
"""True only for gateway multi-worker registry misses, not upstream errors.
Matches the proxy's own shapes:
- ValueError ``Tool <name> not found`` wrapped as HTTP 500 (cold tool map /
unresolved server on this process)
- REST ``server_not_found`` when this worker has not loaded the MCP server row
Does not treat arbitrary 500 bodies that merely mention "tool" and "not found"
(e.g. upstream MCP payload text) as lag, so await_call_tool does not retry
real failures or non-idempotent calls.
"""
if not isinstance(result, UnknownApiError) or result.status_code != 500:
return False
body = result.body
body_l = body.lower()
if "server_not_found" in body_l:
return True
if re.search(r"mcp server ['\"][^'\"]+['\"] was not found", body_l):
return True
# Gateway: "Tool search_datadog_logs not found" (optionally inside a longer message)
if tool_name is not None:
return (
re.search(rf"\btool\s+{re.escape(tool_name)}\s+not found\b", body_l) is not None
)
return re.search(r"\btool\s+\S+\s+not found\b", body_l) is not None
def build_client(proxy: ProxyClient) -> McpClient:
return McpClient(proxy=proxy)

View file

@ -29,6 +29,7 @@ class TestMcpAccessGroupToolSelection:
) -> None:
group = f"e2e-mcp-grp-{unique_marker()}"
server_id = register_datadog_mcp(client, resources, mcp_access_groups=[group])
client.await_registered(server_id)
granted = client.generate_key(
user_id=f"e2e-mcp-ag-granted-{unique_marker()}",

View file

@ -60,6 +60,7 @@ class TestDatadogMcpRoundTrip:
_assert_datadog_logger_active(client.proxy)
server_id = register_datadog_mcp(client, resources)
client.await_registered(server_id)
marker = f"{MARKER_PREFIX}{unique_marker()}"
key = client.generate_key(
@ -78,22 +79,19 @@ class TestDatadogMcpRoundTrip:
)
tool_name = client.await_tool(key, server_id, SEARCH_LOGS_TOOL)
call = unwrap(
client.call_tool(
key,
server_id=server_id,
name=tool_name,
arguments={
"query": marker,
"from": DD_SEARCH_FROM,
"to": "now",
"max_tokens": 5000,
"telemetry": {
"intent": "e2e assert seeded litellm completion log is searchable via MCP"
},
call = client.await_call_tool(
key,
server_id=server_id,
name=tool_name,
arguments={
"query": marker,
"from": DD_SEARCH_FROM,
"to": "now",
"max_tokens": 5000,
"telemetry": {
"intent": "e2e assert seeded litellm completion log is searchable via MCP"
},
)
},
)
assert call.is_error is not True, f"search_datadog_logs errored: {call}"
body = call.all_text

View file

@ -16,7 +16,7 @@ import pytest
from datadog_mcp import SEARCH_LOGS_TOOL, register_datadog_mcp
from e2e_config import DD_SEARCH_FROM, unique_marker
from e2e_http import UnknownApiError, unwrap
from e2e_http import unwrap
from lifecycle import ResourceManager
from mcp_client import McpClient
@ -72,13 +72,12 @@ class TestMcpKeyWithoutAccessIsDenied:
"max_tokens": 1000,
"telemetry": {"intent": "e2e control call proving granted key can invoke Datadog MCP"},
}
permitted_call = unwrap(
client.call_tool(permitted_key, server_id=server_id, name=tool_name, arguments=search_args)
permitted_call = client.await_call_tool(
permitted_key, server_id=server_id, name=tool_name, arguments=search_args
)
assert permitted_call.is_error is not True, f"granted key's tool call errored: {permitted_call}"
match client.call_tool(denied_key, server_id=server_id, name=tool_name, arguments=search_args):
case UnknownApiError(status_code=403, body=body):
assert "access_denied" in body, f"403 was not an MCP access denial: {body}"
case other:
pytest.fail(f"ungranted key's tool call was not refused with 403 access_denied: {other}")
denied = client.await_call_tool_denied(
denied_key, server_id=server_id, name=tool_name, arguments=search_args
)
assert "access_denied" in denied.body, f"403 was not an MCP access denial: {denied.body}"

View file

@ -2184,7 +2184,8 @@ def test_team_alias_stale_bypass_disabled_by_default(monkeypatch):
pre_call_utils._ENABLE_TEAM_STALE_ALIAS_BYPASS = None
class _MockRouter:
team_model_to_deployment_indices = {("team-1", "gpt-4o"): [0]}
model_name_to_deployment_indices = {"model_name_team-1_legacy-uuid": [0]}
team_model_to_deployment_indices = {("team-1", "gpt-4o"): [1]}
test_data = {"model": "gpt-4o"}
user_api_key_dict = UserAPIKeyAuth(
@ -2209,7 +2210,8 @@ def test_team_alias_stale_bypass_enabled_by_flag(monkeypatch):
pre_call_utils._ENABLE_TEAM_STALE_ALIAS_BYPASS = None
class _MockRouter:
team_model_to_deployment_indices = {("team-1", "gpt-4o"): [0]}
model_name_to_deployment_indices = {"model_name_team-1_legacy-uuid": [0]}
team_model_to_deployment_indices = {("team-1", "gpt-4o"): [1]}
test_data = {"model": "gpt-4o"}
user_api_key_dict = UserAPIKeyAuth(

View file

@ -2659,6 +2659,23 @@ def test_resolve_model_name_from_model_id():
result = router.resolve_model_name_from_model_id("gpt-5-mini")
assert result == "gpt-5-mini"
# Test case 10: model_id is a deployment ID (hash) that differs from the
# public model_name. Regression for #32580: managed batch/file IDs embed the
# deployment model_id, and it must resolve back to the public model_name so
# team model-access checks compare against the model group, not the hash.
model_list = [
{
"model_name": "bedrock-batch-model",
"litellm_params": {
"model": "bedrock/global.anthropic.claude-haiku-4-5-20251001-v1:0",
},
"model_info": {"id": "8d0eaa7e6c6f54a425dfd0062cb6b0dc"},
},
]
router = Router(model_list=model_list)
result = router.resolve_model_name_from_model_id("8d0eaa7e6c6f54a425dfd0062cb6b0dc")
assert result == "bedrock-batch-model"
def test_get_valid_args():
"""Test get_valid_args static method returns valid Router.__init__ arguments"""

View file

@ -5,6 +5,8 @@ Regression test for afile_retrieve called without credentials in
async_post_call_success_hook when processing completed batch responses.
"""
import json
import pytest
from typing import Optional
from unittest.mock import AsyncMock, MagicMock, patch
@ -385,3 +387,59 @@ async def test_afile_content_error_reports_unified_id_not_provider_uri():
message = str(exc_info.value)
assert unified_file_id in message
assert s3_uri not in message
def _make_real_managed_files_instance():
"""Create a _PROXY_LiteLLMManagedFiles with a real store_unified_file_id but
an AsyncMock prisma client, so the DB write path itself can be asserted."""
from litellm_enterprise.proxy.hooks.managed_files import (
_PROXY_LiteLLMManagedFiles,
)
mock_cache = MagicMock()
mock_cache.async_set_cache = AsyncMock()
mock_prisma = MagicMock()
mock_prisma.db.litellm_managedfiletable.upsert = AsyncMock()
mock_prisma.db.litellm_managedfiletable.create = AsyncMock(
side_effect=AssertionError(
"store_unified_file_id must upsert, not create, on the retrieve path"
)
)
return (
_PROXY_LiteLLMManagedFiles(
internal_usage_cache=mock_cache,
prisma_client=mock_prisma,
),
mock_prisma,
)
@pytest.mark.asyncio
async def test_store_unified_file_id_is_idempotent_via_upsert():
"""Regression test for the managed-batch retrieve 500 (UniqueViolationError on
unified_file_id): re-registering an already-stored output file id must upsert on
unified_file_id, never do an unconditional create that raises on conflict."""
managed_files, mock_prisma = _make_real_managed_files_instance()
file_id = "litellm_proxy_unified_output_id_abc"
model_mappings = {"model-deploy-xyz": "file-output-abc"}
for _ in range(2):
await managed_files.store_unified_file_id(
file_id=file_id,
file_object=_make_file_object(),
litellm_parent_otel_span=None,
model_mappings=model_mappings,
user_api_key_dict=_make_user_api_key_dict(),
)
mock_prisma.db.litellm_managedfiletable.create.assert_not_awaited()
upsert_mock = mock_prisma.db.litellm_managedfiletable.upsert
assert upsert_mock.await_count == 2
for upsert_call in upsert_mock.await_args_list:
assert upsert_call.kwargs["where"] == {"unified_file_id": file_id}
upsert_data = upsert_call.kwargs["data"]
assert upsert_data["create"]["unified_file_id"] == file_id
assert json.loads(upsert_data["create"]["model_mappings"]) == model_mappings
assert json.loads(upsert_data["update"]["model_mappings"]) == model_mappings

View file

@ -451,6 +451,30 @@ def test_parse_headers():
assert providers.parse_headers("no-equals") == {}
def test_parse_headers_percent_decodes_values():
"""A percent-encoded OTLP header value reaches the exporter decoded.
``OTEL_EXPORTER_OTLP_HEADERS`` is W3C Baggage encoded, and Grafana Cloud
documents ``Authorization=Basic%20<token>``. Forwarding the literal ``%20``
makes the backend reject the export as a malformed credential.
"""
token = "MTMzNzc4MzpnbGNfZXlKdklqb2lNVEl6TkNJPQ=="
assert providers.parse_headers(f"Authorization=Basic%20{token}") == {"authorization": f"Basic {token}"}
assert providers.parse_headers("x-scope-orgid=team%20a") == {"x-scope-orgid": "team a"}
def test_parse_headers_keeps_unencoded_values_working():
"""Values that are not percent-encoded keep parsing unchanged.
Vendors that document a bare space, and litellm's own presets, must survive
the switch to the spec-compliant parser. Base64 padding also means a value
can contain ``=``, so only the first one may split the pair.
"""
assert providers.parse_headers("Authorization=Bearer sk-123") == {"authorization": "Bearer sk-123"}
assert providers.parse_headers("api_key=abc,space_id=xyz") == {"api_key": "abc", "space_id": "xyz"}
assert providers.parse_headers("api_key=YWJjZA==") == {"api_key": "YWJjZA=="}
def test_otlp_traces_endpoint_normalization():
norm = providers._otlp_traces_endpoint
# A base endpoint gets the signal path appended (the common OTLP env shape).
@ -487,6 +511,24 @@ def test_build_span_exporter_variants():
assert "OTLPSpanExporter" in type(http_exporter).__name__
def test_otlp_metric_exporter_uses_cumulative_histogram_temporality():
"""Histograms must export as cumulative, not delta.
Prometheus-backed OTLP receivers (Grafana Cloud / Mimir) reject delta
histograms with ``invalid temporality and type combination`` and drop the
entire metric batch, so a delta default silently loses every GenAI metric.
"""
from opentelemetry.sdk.metrics import Histogram
from opentelemetry.sdk.metrics.export import AggregationTemporality
reader = providers.build_metric_reader(
OpenTelemetryV2Config(exporter="otlp_http", endpoint="http://h:4318")
)
temporality = reader._exporter._preferred_temporality # noqa: SLF001 # exporter exposes no public accessor
assert temporality[Histogram] is AggregationTemporality.CUMULATIVE
def test_otlp_logs_endpoint_normalization():
norm = providers._otlp_logs_endpoint
# A base endpoint gets the signal path appended (the common OTLP env shape).

View file

@ -2120,9 +2120,9 @@ def test_valid_metric_filter_records_six_metrics(monkeypatch):
assert _emitted_metric_names(reader) == {
"gen_ai.client.operation.duration",
"gen_ai.client.token.usage",
"gen_ai.client.token.cost",
"gen_ai.client.response.time_to_first_token",
"gen_ai.client.response.time_per_output_token",
"gen_ai.usage.cost",
"gen_ai.server.time_to_first_token",
"gen_ai.server.time_per_output_token",
"gen_ai.client.response.duration",
}

View file

@ -39,9 +39,9 @@ from litellm.integrations.otel.plumbing.providers import ( # noqa: E402
OPERATION_DURATION = "gen_ai.client.operation.duration"
TOKEN_USAGE = "gen_ai.client.token.usage"
TOKEN_COST = "gen_ai.client.token.cost"
TIME_TO_FIRST_TOKEN = "gen_ai.client.response.time_to_first_token"
TIME_PER_OUTPUT_TOKEN = "gen_ai.client.response.time_per_output_token"
TOKEN_COST = "gen_ai.usage.cost"
TIME_TO_FIRST_TOKEN = "gen_ai.server.time_to_first_token"
TIME_PER_OUTPUT_TOKEN = "gen_ai.server.time_per_output_token"
RESPONSE_DURATION = "gen_ai.client.response.duration"
ALL_METRICS = frozenset(

View file

@ -412,6 +412,40 @@ class TestOpenTelemetryProviderInitialization(unittest.TestCase):
current_provider is existing_provider
), "Existing TracerProvider should be respected and not overridden"
@patch.dict(
os.environ, {"LITELLM_OTEL_INTEGRATION_ENABLE_METRICS": "true"}, clear=True
)
def test_init_metrics_creates_instruments_under_their_published_names(self):
"""
The v1 engine's instrument names are a public contract.
Every name here is what a backend queries: four are GenAI semantic
conventions and gen_ai.usage.cost is the name backends query for spend.
A rename is breaking for anyone charting them, so it has to be a
deliberate edit to the shared Metric constants and to this list, never
a silent drift between the v1 and v2 engines.
"""
from opentelemetry import metrics
metrics.set_meter_provider(MeterProvider(metric_readers=[InMemoryMetricReader()]))
otel_integration = OpenTelemetry(config=OpenTelemetryConfig.from_env())
assert {
otel_integration._operation_duration_histogram.name,
otel_integration._token_usage_histogram.name,
otel_integration._cost_histogram.name,
otel_integration._time_to_first_token_histogram.name,
otel_integration._time_per_output_token_histogram.name,
otel_integration._response_duration_histogram.name,
} == {
"gen_ai.client.operation.duration",
"gen_ai.client.token.usage",
"gen_ai.usage.cost",
"gen_ai.server.time_to_first_token",
"gen_ai.server.time_per_output_token",
"gen_ai.client.response.duration",
}
@patch.dict(
os.environ, {"LITELLM_OTEL_INTEGRATION_ENABLE_METRICS": "true"}, clear=True
)

View file

@ -37,6 +37,13 @@ def _load_openapi_spec_dict() -> Dict[str, Any]:
)
def _declared_type_value(variant_schema: Dict[str, Any]) -> Any:
"""The single `type` value a union variant pins, whether spelled as a const or a 1-item enum."""
type_property = variant_schema.get("properties", {}).get("type", {})
enum_values = type_property.get("enum") or []
return type_property.get("const") or (enum_values[0] if len(enum_values) == 1 else None)
@pytest.fixture(scope="module")
def spec_dict() -> Dict[str, Any]:
"""Load raw spec dict for manual validation."""
@ -105,26 +112,51 @@ class TestRequestCompliance:
assert "string" in input_types, "Input should support string"
assert "array" in input_types, "Input should support array"
def test_content_schema_uses_discriminator(self, spec_dict):
"""Verify Content uses type discriminator."""
def test_content_variants_are_identified_by_their_type_field(self, spec_dict):
"""Verify a Content part can be told apart by its `type`, however the spec spells that.
Our transformation reads `type` off each content part to route it, so what has to hold is
that every variant of the union pins a distinct `type` value and that text is one of them.
A spec may express that with an OpenAPI `discriminator` on the union or with a `const` on
each member's own `type`; both are equivalent for us, so accepting only the first makes
this test fail on a stylistic change upstream that costs us nothing.
"""
content_schema = spec_dict["components"]["schemas"]["Content"]
assert "discriminator" in content_schema
assert content_schema["discriminator"]["propertyName"] == "type"
# Check TextContent is an option (via mapping if present, or via oneOf refs)
mapping = content_schema["discriminator"].get("mapping")
if mapping:
assert "text" in mapping
print(f"Content type discriminator mapping: {list(mapping.keys())}")
else:
# Discriminator without explicit mapping — verify via oneOf
one_of = content_schema.get("oneOf", [])
ref_names = [opt["$ref"].split("/")[-1] for opt in one_of if "$ref" in opt]
discriminator = content_schema.get("discriminator")
if discriminator is not None:
assert (
"TextContent" in ref_names
), f"TextContent not found in oneOf refs: {ref_names}"
print(f"Content type discriminator (no mapping), oneOf refs: {ref_names}")
discriminator.get("propertyName") == "type"
), f"Content is discriminated on {discriminator.get('propertyName')!r}, not 'type'"
variant_names = [
option["$ref"].split("/")[-1]
for option in content_schema.get("oneOf", [])
if "$ref" in option
]
assert variant_names, f"Content is not a union of named variants: {content_schema}"
mapping = (discriminator or {}).get("mapping") or {}
type_values = {
variant: mapping_value
for mapping_value, ref in mapping.items()
for variant in [ref.split("/")[-1]]
} or {
variant: _declared_type_value(spec_dict["components"]["schemas"].get(variant, {}))
for variant in variant_names
}
assert set(type_values) == set(variant_names) and all(type_values.values()), (
f"every Content variant needs a discoverable type value, "
f"got {type_values} for variants {sorted(variant_names)}"
)
assert len(set(type_values.values())) == len(type_values), (
f"Content variants must pin DISTINCT type values, got {type_values}"
)
assert type_values.get("TextContent") == "text", (
f"TextContent must be reachable as type 'text', got {type_values}"
)
print(f"Content variants by type: {type_values}")
def test_text_content_schema(self, spec_dict):
"""Verify TextContent schema."""

View file

@ -506,3 +506,162 @@ def test_empty_content_chunk_mid_text_block_is_suppressed_sync():
assert _text_deltas(events) == ["Hi", " there"]
_assert_deltas_match_their_block_type(events)
def _thinking_first_chunks() -> List[MagicMock]:
return [
_thinking_chunk("Let me think"),
_thinking_chunk("about it."),
_make_chunk(Delta(content="42")),
_make_chunk(Delta(content=None), finish_reason="stop"),
]
def _assert_thinking_first_block_opens_at_index_zero(events: List[dict]) -> None:
starts = [
(e["index"], e["content_block"]["type"])
for e in events
if e.get("type") == "content_block_start"
]
assert starts == [(0, "thinking"), (1, "text")], starts
assert "" not in _text_deltas(events)
assert _thinking_deltas(events) == ["Let me think", "about it."]
assert _text_deltas(events) == ["42"]
_assert_deltas_match_their_block_type(events)
def test_thinking_first_stream_opens_thinking_block_at_index_zero_sync():
"""Bug A regression: when the model's first output is reasoning the adapter
must open the first content block as ``thinking`` at index 0. The previous
code pre-emitted a hardcoded empty ``text`` block at index 0 before
inspecting any upstream chunk, then opened ``thinking`` at index 1; strict
Anthropic SDK clients with thinking enabled reject that stream with
"Content block is not a thinking block".
"""
wrapper = AnthropicStreamWrapper(
completion_stream=iter(_thinking_first_chunks()),
model="claude-x",
)
_assert_thinking_first_block_opens_at_index_zero(_drain_sync(wrapper))
@pytest.mark.asyncio
async def test_thinking_first_stream_opens_thinking_block_at_index_zero_async():
"""Async twin of the Bug A regression; the proxy serves the async iterator,
so the first block must be ``thinking`` at index 0 on this path too.
"""
wrapper = AnthropicStreamWrapper(
completion_stream=_AsyncStream(_thinking_first_chunks()),
model="claude-x",
)
_assert_thinking_first_block_opens_at_index_zero(await _drain_async(wrapper))
def _reasoning_content_chunk(reasoning: str) -> MagicMock:
return _make_chunk(Delta(content=None, reasoning_content=reasoning))
def _reasoning_first_chunks() -> List[MagicMock]:
return [
_reasoning_content_chunk("Let me think"),
_reasoning_content_chunk("about it."),
_make_chunk(Delta(content="42")),
_make_chunk(Delta(content=None), finish_reason="stop"),
]
def test_reasoning_content_first_stream_opens_thinking_block_at_index_zero_sync():
"""The reported backend (hosted_vllm; vLLM and SGLang reasoning parsers)
surfaces reasoning as OpenAI ``reasoning_content`` with no
``thinking_blocks``. Such a stream must also open the first content block as
``thinking`` at index 0, exercising the reasoning_content branch of the
chunk translator rather than the thinking_blocks branch the other twins use.
"""
wrapper = AnthropicStreamWrapper(
completion_stream=iter(_reasoning_first_chunks()),
model="claude-x",
)
_assert_thinking_first_block_opens_at_index_zero(_drain_sync(wrapper))
def _blank_lead_chunks() -> List[MagicMock]:
return [
_make_chunk(Delta(content=None)),
_thinking_chunk("Let me think"),
_thinking_chunk("about it."),
_make_chunk(Delta(content="42")),
_make_chunk(Delta(content=None), finish_reason="stop"),
]
def _role_only_reasoning_content_lead_chunks() -> List[MagicMock]:
return [
_make_chunk(Delta(role="assistant", content=None, tool_calls=[])),
_reasoning_content_chunk("Let me think"),
_reasoning_content_chunk("about it."),
_make_chunk(Delta(content="42")),
_make_chunk(Delta(content=None), finish_reason="stop"),
]
def test_contentless_lead_chunk_does_not_open_text_block_before_thinking_sync():
"""OpenAI-compatible streaming backends open the response with a contentless
priming chunk (an empty delta, e.g. the {role: assistant} lead-in) before the
first real token. Such a lead chunk must NOT commit index 0 to an empty text
block; the following thinking chunk must still open thinking at index 0, or
strict Anthropic SDK clients reject the stream.
"""
wrapper = AnthropicStreamWrapper(
completion_stream=iter(_blank_lead_chunks()),
model="claude-x",
)
_assert_thinking_first_block_opens_at_index_zero(_drain_sync(wrapper))
def test_role_only_lead_chunk_does_not_open_text_block_before_reasoning_content_sync():
wrapper = AnthropicStreamWrapper(
completion_stream=iter(_role_only_reasoning_content_lead_chunks()),
model="claude-x",
)
_assert_thinking_first_block_opens_at_index_zero(_drain_sync(wrapper))
@pytest.mark.asyncio
async def test_contentless_lead_chunk_does_not_open_text_block_before_thinking_async():
"""Async twin; the proxy serves the async iterator, so the contentless lead
chunk must be skipped on this path too.
"""
wrapper = AnthropicStreamWrapper(
completion_stream=_AsyncStream(_blank_lead_chunks()),
model="claude-x",
)
_assert_thinking_first_block_opens_at_index_zero(await _drain_async(wrapper))
@pytest.mark.asyncio
async def test_role_only_lead_chunk_does_not_open_text_block_before_reasoning_content_async():
wrapper = AnthropicStreamWrapper(
completion_stream=_AsyncStream(_role_only_reasoning_content_lead_chunks()),
model="claude-x",
)
_assert_thinking_first_block_opens_at_index_zero(await _drain_async(wrapper))
def test_finish_first_chunk_is_not_deferred_sync():
"""A stream whose first upstream chunk is already the finish event must not
be skipped by the blank-delta deferral. ``_is_blank_delta`` returns False
for a finish chunk so the message_delta still flows (with an empty text
block opened and closed first); without that guard the deferral would drop
the terminal event entirely.
"""
chunks = [_make_chunk(Delta(content=None), finish_reason="stop")]
wrapper = AnthropicStreamWrapper(completion_stream=iter(chunks), model="claude-x")
events = _drain_sync(wrapper)
assert [e["type"] for e in events] == [
"message_start",
"content_block_start",
"content_block_stop",
"message_delta",
"message_stop",
]

View file

@ -134,13 +134,6 @@ def test_anthropic_stream_wrapper_single_tool_call():
# Verify the expected sequence of chunk types
expected_types = [
"message_start", # Initial message start
# TODO: for future contributors: if the initial content_block_start
# respects the upstream's starting chunk, the initial empty text block
# should be removed (and this test should be updated accordingly)
# ---------------------------------------------------------------------
"content_block_start", # Initial empty text block start
"content_block_stop", # End of empty text block
# ---------------------------------------------------------------------
"content_block_start", # Start of first tool_use content block
"content_block_delta", # {"city":
"content_block_delta", # "NY"}
@ -196,13 +189,6 @@ def test_anthropic_stream_wrapper_back_to_back_tool_calls():
# Verify the expected sequence of chunk types
expected_types = [
"message_start", # Initial message start
# TODO: for future contributors: if the initial content_block_start
# respects the upstream's starting chunk, the initial empty text block
# should be removed (and this test should be updated accordingly)
# ---------------------------------------------------------------------
"content_block_start", # Initial empty text block start
"content_block_stop", # End of empty text block
# ---------------------------------------------------------------------
"content_block_start", # Start of first tool_use content block
"content_block_delta", # {"city":
"content_block_delta", # "NY"}
@ -267,13 +253,6 @@ def test_anthropic_stream_wrapper_interleaved_tool_calls_and_text():
# Verify the expected sequence of chunk types
expected_types = [
"message_start", # Initial message start
# TODO: for future contributors: if the initial content_block_start
# respects the upstream's starting chunk, the initial empty text block
# should be removed (and this test should be updated accordingly)
# ---------------------------------------------------------------------
"content_block_start", # Initial empty text block start
"content_block_stop", # End of empty text block
# ---------------------------------------------------------------------
"content_block_start", # Start of first tool_use content block
"content_block_delta", # {"city":
"content_block_delta", # "NY"}

View file

@ -1452,6 +1452,157 @@ class TestContextCachingEndpoints:
# Restart the patcher so teardown_method can stop it cleanly
self._token_check_patcher.start()
def _model_turn_final_messages(self, final_cached_role):
tool_call = {
"id": "call_abc123",
"type": "function",
"function": {"name": "get_weather", "arguments": '{"location": "Boston"}'},
}
cached_tail = {
"assistant": [],
"tool": [
{
"role": "tool",
"tool_call_id": "call_abc123",
"content": "72F and sunny",
"cache_control": {"type": "ephemeral"},
}
],
"system": [
{
"role": "system",
"content": "Tool results are authoritative.",
"cache_control": {"type": "ephemeral"},
}
],
}[final_cached_role]
return [
{
"role": "user",
"content": [
{
"type": "text",
"text": "Use the weather tool for every answer.",
"cache_control": {"type": "ephemeral"},
}
],
},
{
"role": "assistant",
"content": "",
"tool_calls": [tool_call],
"cache_control": {"type": "ephemeral"},
},
*cached_tail,
{"role": "user", "content": "What is the weather in Boston?"},
]
@pytest.mark.parametrize("final_cached_role", ["assistant", "tool", "system"])
def test_check_and_create_cache_skips_when_cached_block_ends_on_model_turn(
self, final_cached_role
):
"""The cachedContents API rejects contents ending on an assistant or tool turn
with HTTP 400 "Requests ending with a model turn are not supported", so the
request must proceed uncached instead of failing.
"""
all_messages = self._model_turn_final_messages(final_cached_role)
optional_params = self.sample_optional_params.copy()
result = self.context_caching.check_and_create_cache(
messages=all_messages,
optional_params=optional_params,
api_key="test_key",
api_base=None,
model="gemini-3.6-flash",
client=self.mock_client,
timeout=30.0,
logging_obj=self.mock_logging,
cached_content=None,
custom_llm_provider="vertex_ai",
vertex_project="test_project",
vertex_location="us-central1",
vertex_auth_header="test_token",
)
messages, returned_params, returned_cache = result
assert messages == all_messages
assert returned_cache is None
assert "tools" in returned_params
self.mock_client.get.assert_not_called()
self.mock_client.post.assert_not_called()
@pytest.mark.parametrize("final_cached_role", ["assistant", "tool", "system"])
@pytest.mark.asyncio
async def test_async_check_and_create_cache_skips_when_cached_block_ends_on_model_turn(
self, final_cached_role
):
"""Async variant: an unsupported terminal turn skips caching instead of failing."""
all_messages = self._model_turn_final_messages(final_cached_role)
optional_params = self.sample_optional_params.copy()
result = await self.context_caching.async_check_and_create_cache(
messages=all_messages,
optional_params=optional_params,
api_key="test_key",
api_base=None,
model="gemini-3.6-flash",
client=self.mock_async_client,
timeout=30.0,
logging_obj=self.mock_logging,
cached_content=None,
custom_llm_provider="vertex_ai",
vertex_project="test_project",
vertex_location="us-central1",
vertex_auth_header="test_token",
)
messages, returned_params, returned_cache = result
assert messages == all_messages
assert returned_cache is None
assert "tools" in returned_params
self.mock_async_client.get.assert_not_called()
self.mock_async_client.post.assert_not_called()
def test_cached_messages_end_on_supported_turn():
from litellm.llms.vertex_ai.context_caching.transformation import (
cached_messages_end_on_supported_turn,
)
assert (
cached_messages_end_on_supported_turn(
[{"role": "assistant", "content": "hi"}, {"role": "user", "content": "hello"}]
)
is True
)
assert cached_messages_end_on_supported_turn([{"role": "system", "content": "be brief"}]) is True
assert cached_messages_end_on_supported_turn([{"role": "assistant", "content": "hi"}]) is False
assert (
cached_messages_end_on_supported_turn(
[
{"role": "user", "content": "hello"},
{"role": "assistant", "content": "hi"},
{"role": "system", "content": "be brief"},
]
)
is False
)
assert (
cached_messages_end_on_supported_turn(
[{"role": "system", "content": "be brief"}, {"role": "user", "content": "hello"}]
)
is True
)
assert (
cached_messages_end_on_supported_turn([{"role": "tool", "tool_call_id": "x", "content": "y"}])
is False
)
assert (
cached_messages_end_on_supported_turn([{"role": "function", "name": "f", "content": "y"}])
is False
)
assert cached_messages_end_on_supported_turn([]) is False
class TestCheckCachePagination:
"""Test pagination logic in check_cache and async_check_cache methods."""

View file

@ -31,10 +31,7 @@ class TestVertexAIFilesHandler:
def test_extract_bucket_and_object_from_file_id_standard_path(self):
"""Test extraction of bucket and object from URL-encoded file_id with standard path"""
# Sample file_id with nested folder structure
file_id = (
"gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files"
"%2Ftest-folder%2Fsub-folder%2Ftest-file.txt"
)
file_id = "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files%2Ftest-folder%2Fsub-folder%2Ftest-file.txt"
bucket_name, object_path = self.handler._extract_bucket_and_object_from_file_id(
file_id=file_id,
@ -105,21 +102,14 @@ class TestVertexAIFilesHandler:
async def test_afile_content_success(self):
"""Test successful async file content retrieval"""
# Setup test data
file_id = (
"gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files"
"%2Fuploads%2Fabc-test-file.txt"
)
file_id = "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-test-file.txt"
expected_content = b"test file content"
file_content_request = FileContentRequest(
file_id=file_id, extra_headers=None, extra_body=None
)
file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
# Mock the download_gcs_object method
with (
patch.object(
self.handler, "download_gcs_object", new_callable=AsyncMock
) as mock_download,
patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download,
patch.object(
self.handler,
"get_gcs_logging_config",
@ -148,15 +138,9 @@ class TestVertexAIFilesHandler:
# Verify the download was called with correct parameters
mock_download.assert_called_once()
call_args = mock_download.call_args
assert (
call_args.kwargs["object_name"]
== "litellm-vertex-files/uploads/abc-test-file.txt"
)
assert call_args.kwargs["object_name"] == "litellm-vertex-files/uploads/abc-test-file.txt"
assert "standard_callback_dynamic_params" in call_args.kwargs
assert (
call_args.kwargs["standard_callback_dynamic_params"]["gcs_bucket_name"]
== "test-bucket"
)
assert call_args.kwargs["standard_callback_dynamic_params"]["gcs_bucket_name"] == "test-bucket"
@pytest.mark.asyncio
async def test_afile_content_missing_file_id(self):
@ -164,9 +148,7 @@ class TestVertexAIFilesHandler:
file_content_request = FileContentRequest(extra_headers=None, extra_body=None)
# Should raise ValueError for missing file_id
with pytest.raises(
ValueError, match="file_id is required in file_content_request"
):
with pytest.raises(ValueError, match="file_id is required in file_content_request"):
await self.handler.afile_content(
file_content_request=file_content_request,
vertex_credentials=None,
@ -179,20 +161,13 @@ class TestVertexAIFilesHandler:
@pytest.mark.asyncio
async def test_afile_content_download_failure(self):
"""Test async file content retrieval when download fails"""
file_id = (
"gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files"
"%2Fuploads%2Fabc-test-file.txt"
)
file_id = "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-test-file.txt"
file_content_request = FileContentRequest(
file_id=file_id, extra_headers=None, extra_body=None
)
file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
# Mock download to return None (failure)
with (
patch.object(
self.handler, "download_gcs_object", new_callable=AsyncMock
) as mock_download,
patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download,
patch.object(
self.handler,
"get_gcs_logging_config",
@ -216,14 +191,130 @@ class TestVertexAIFilesHandler:
max_retries=3,
)
def test_resolve_read_gcs_config_prefers_per_model_bucket(self, monkeypatch):
monkeypatch.setenv("GCS_BUCKET_NAME", "env-default-bucket")
monkeypatch.setenv("GCS_PATH_SERVICE_ACCOUNT", "/env/sa.json")
bucket, service_account = self.handler._resolve_read_gcs_config(
litellm_params={
"gcs_bucket_name": "my-model-bucket",
"vertex_credentials": "/model/sa.json",
},
vertex_credentials=None,
)
assert bucket == "my-model-bucket"
assert service_account == "/model/sa.json"
def test_resolve_read_gcs_config_falls_back_to_env(self, monkeypatch):
monkeypatch.setenv("GCS_BUCKET_NAME", "env-default-bucket")
monkeypatch.setenv("GCS_PATH_SERVICE_ACCOUNT", "/env/sa.json")
bucket, service_account = self.handler._resolve_read_gcs_config(litellm_params={}, vertex_credentials=None)
assert bucket == "env-default-bucket"
assert service_account == "/env/sa.json"
def test_resolve_read_gcs_config_serializes_dict_credentials(self, monkeypatch):
monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False)
_, service_account = self.handler._resolve_read_gcs_config(
litellm_params={"gcs_bucket_name": "my-model-bucket"},
vertex_credentials={"type": "service_account", "project_id": "p"},
)
assert service_account == '{"type": "service_account", "project_id": "p"}'
@pytest.mark.asyncio
async def test_afile_content_honors_per_model_bucket_over_env(self, monkeypatch):
"""
Regression for #32640: a batch output written to a per-model gcs_bucket_name must be
readable even when the global GCS_BUCKET_NAME points at a different bucket. Before the
fix the read path resolved the bucket from env only and raised
"file_id bucket does not match the configured storage bucket".
"""
monkeypatch.setenv("GCS_BUCKET_NAME", "env-default-bucket")
monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False)
file_id = "gs%3A%2F%2Fmy-model-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-batch-output.jsonl"
file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
with (
patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download,
patch.object(
self.handler,
"get_or_create_vertex_instance",
new_callable=AsyncMock,
return_value=object(),
),
):
mock_download.return_value = b"batch output"
result = await self.handler.afile_content(
file_content_request=file_content_request,
vertex_credentials="/model/sa.json",
vertex_project="test-project",
vertex_location="us-central1",
timeout=60.0,
max_retries=0,
litellm_params={
"gcs_bucket_name": "my-model-bucket",
"vertex_credentials": "/model/sa.json",
},
)
assert isinstance(result, HttpxBinaryResponseContent)
assert result.response.content == b"batch output"
dynamic_params = mock_download.call_args.kwargs["standard_callback_dynamic_params"]
assert dynamic_params["gcs_bucket_name"] == "my-model-bucket"
assert dynamic_params["gcs_path_service_account"] == "/model/sa.json"
assert mock_download.call_args.kwargs["object_name"] == "litellm-vertex-files/uploads/abc-batch-output.jsonl"
@pytest.mark.asyncio
async def test_afile_content_reads_without_global_env_bucket(self, monkeypatch):
"""
Regression for #32640: with no global GCS_BUCKET_NAME set, a model-group-level
deployment (per-model gcs_bucket_name) must still be readable. Before the fix the read
path raised "GCS_BUCKET_NAME is not set in the environment".
"""
monkeypatch.delenv("GCS_BUCKET_NAME", raising=False)
monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False)
file_id = "gs%3A%2F%2Fmy-model-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-batch-output.jsonl"
file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
with (
patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download,
patch.object(
self.handler,
"get_or_create_vertex_instance",
new_callable=AsyncMock,
return_value=object(),
),
):
mock_download.return_value = b"batch output"
result = await self.handler.afile_content(
file_content_request=file_content_request,
vertex_credentials="/model/sa.json",
vertex_project="test-project",
vertex_location="us-central1",
timeout=60.0,
max_retries=0,
litellm_params={"gcs_bucket_name": "my-model-bucket"},
)
assert isinstance(result, HttpxBinaryResponseContent)
dynamic_params = mock_download.call_args.kwargs["standard_callback_dynamic_params"]
assert dynamic_params["gcs_bucket_name"] == "my-model-bucket"
def test_file_content_sync_success(self):
"""Test successful sync file content retrieval"""
file_id = "gs%3A%2F%2Ftest-bucket%2Ftest-file.txt"
expected_content = b"test file content"
file_content_request = FileContentRequest(
file_id=file_id, extra_headers=None, extra_body=None
)
file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
# Create expected response
mock_response = httpx.Response(
@ -261,25 +352,17 @@ class TestVertexAIFilesHandler:
file_id = "gs%3A%2F%2Ftest-bucket%2Ftest-file.txt"
expected_content = b"test file content"
file_content_request = FileContentRequest(
file_id=file_id, extra_headers=None, extra_body=None
)
file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
# Mock the afile_content method
with patch.object(
self.handler, "afile_content", new_callable=AsyncMock
) as mock_afile_content:
with patch.object(self.handler, "afile_content", new_callable=AsyncMock) as mock_afile_content:
mock_response = httpx.Response(
status_code=200,
content=expected_content,
headers={"content-type": "application/octet-stream"},
request=httpx.Request(
method="GET", url="gs://test-bucket/test-file.txt"
),
)
mock_afile_content.return_value = HttpxBinaryResponseContent(
response=mock_response
request=httpx.Request(method="GET", url="gs://test-bucket/test-file.txt"),
)
mock_afile_content.return_value = HttpxBinaryResponseContent(response=mock_response)
# Call the method with _is_async=True
result = self.handler.file_content(

View file

@ -2276,82 +2276,8 @@ def test_is_gemini_3_or_newer():
assert VertexGeminiConfig._is_gemini_3_or_newer("") == False
def test_forward_gemini_function_call_id_vertex_vs_google_ai_studio():
"""Vertex AI rejects `id` on function_call/function_response; Google AI Studio accepts it on Gemini 3.5+."""
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
model = "gemini-3.5-flash"
assert (
VertexGeminiConfig._forward_gemini_function_call_id(model, "vertex_ai") is False
)
assert (
VertexGeminiConfig._forward_gemini_function_call_id(model, "vertex_ai_beta")
is False
)
assert VertexGeminiConfig._forward_gemini_function_call_id(model, "gemini") is True
assert VertexGeminiConfig._forward_gemini_function_call_id(model, None) is False
assert (
VertexGeminiConfig._forward_gemini_function_call_id(
"gemini-2.5-flash", "gemini"
)
is False
)
def test_vertex_ai_gemini_35_tool_calls_omit_function_call_id():
"""Regression: Vertex must not send OpenAI tool_call id inside Gemini function_call parts."""
from litellm.llms.vertex_ai.gemini.transformation import (
_gemini_convert_messages_with_history,
)
messages = [
{"role": "user", "content": "Explore this directory"},
{
"role": "assistant",
"content": "",
"tool_calls": [
{
"id": "call_50e7e0fe0989464a89f188eda443",
"type": "function",
"function": {
"name": "read",
"arguments": '{"filePath": "/tmp"}',
},
}
],
},
{
"role": "tool",
"tool_call_id": "call_50e7e0fe0989464a89f188eda443",
"content": "ok",
},
]
contents = _gemini_convert_messages_with_history(
messages=messages,
model="gemini-3.5-flash",
custom_llm_provider="vertex_ai",
)
for content in contents:
for part in content.get("parts", []):
fc = part.get("function_call")
if fc is not None:
assert "id" not in fc, f"Vertex payload must not include id: {fc}"
fr = part.get("function_response")
if fr is not None:
assert "id" not in fr, f"Vertex payload must not include id: {fr}"
def test_google_ai_studio_gemini_35_tool_calls_include_function_call_id():
from litellm.llms.vertex_ai.gemini.transformation import (
_gemini_convert_messages_with_history,
)
tool_call_id = "call_50e7e0fe0989464a89f188eda443"
messages = [
def _tool_call_messages(tool_call_id: str):
return [
{"role": "user", "content": "hi"},
{
"role": "assistant",
@ -2374,12 +2300,8 @@ def test_google_ai_studio_gemini_35_tool_calls_include_function_call_id():
},
]
contents = _gemini_convert_messages_with_history(
messages=messages,
model="gemini-3.5-flash",
custom_llm_provider="gemini",
)
def _collect_function_call_ids(contents):
function_call_ids = []
function_response_ids = []
for content in contents:
@ -2390,9 +2312,120 @@ def test_google_ai_studio_gemini_35_tool_calls_include_function_call_id():
fr = part.get("function_response")
if fr is not None:
function_response_ids.append(fr.get("id"))
return function_call_ids, function_response_ids
assert function_call_ids == [tool_call_id]
assert function_response_ids == [tool_call_id]
def test_forward_gemini_function_call_id_is_gated_on_model_version_only():
"""Gemini 3+ takes `id` on Vertex AI and Google AI Studio alike; older models reject it."""
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-3.5-flash") is True
assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-3-pro") is True
assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-2.5-flash") is False
assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-2.0-flash") is False
@pytest.mark.parametrize("custom_llm_provider", ["vertex_ai", "vertex_ai_beta", "gemini"])
def test_gemini_35_tool_calls_include_function_call_id(custom_llm_provider):
"""Vertex AI accepts `id` on Gemini 3+, so it must be sent there and not just on AI Studio.
Both parts are asserted together: Vertex pairs a result to its call by id, so emitting one
side without the other would break strict tool-call matching.
"""
from litellm.llms.vertex_ai.gemini.transformation import (
_gemini_convert_messages_with_history,
)
tool_call_id = "call_50e7e0fe0989464a89f188eda443"
contents = _gemini_convert_messages_with_history(
messages=_tool_call_messages(tool_call_id),
model="gemini-3.5-flash",
custom_llm_provider=custom_llm_provider,
)
assert _collect_function_call_ids(contents) == ([tool_call_id], [tool_call_id])
@pytest.mark.parametrize("custom_llm_provider", ["vertex_ai", "gemini"])
def test_gemini_25_tool_calls_omit_function_call_id(custom_llm_provider):
"""Regression: models older than Gemini 3 reject `id`, so the key must be absent entirely."""
from litellm.llms.vertex_ai.gemini.transformation import (
_gemini_convert_messages_with_history,
)
contents = _gemini_convert_messages_with_history(
messages=_tool_call_messages("call_50e7e0fe0989464a89f188eda443"),
model="gemini-2.5-flash",
custom_llm_provider=custom_llm_provider,
)
for content in contents:
for part in content.get("parts", []):
fc = part.get("function_call")
if fc is not None:
assert "id" not in fc, f"gemini-2.5 payload must not include id: {fc}"
fr = part.get("function_response")
if fr is not None:
assert "id" not in fr, f"gemini-2.5 payload must not include id: {fr}"
def test_vertex_ai_forwarded_function_call_id_strips_thought_signature_suffix():
"""The thought signature rides along on the OpenAI id but must not reach Vertex.
Vertex now sees this code path for the first time, so the suffix has to be stripped here too.
"""
from litellm.llms.vertex_ai.gemini.transformation import (
_gemini_convert_messages_with_history,
)
from litellm.litellm_core_utils.prompt_templates.factory import (
THOUGHT_SIGNATURE_SEPARATOR,
)
bare_id = "call_50e7e0fe0989464a89f188eda443"
contents = _gemini_convert_messages_with_history(
messages=_tool_call_messages(f"{bare_id}{THOUGHT_SIGNATURE_SEPARATOR}sig123"),
model="gemini-3.5-flash",
custom_llm_provider="vertex_ai",
)
_, function_response_ids = _collect_function_call_ids(contents)
assert function_response_ids == [bare_id]
@pytest.mark.parametrize("model", ["gemini-3.5-flash", "gemini-2.5-flash"])
def test_tool_response_without_matching_tool_call_is_rejected(model):
"""An unpairable tool result must raise, not ship a functionResponse with no matching call."""
from litellm.llms.vertex_ai.gemini.transformation import (
_gemini_convert_messages_with_history,
)
messages = [
{"role": "user", "content": "hi"},
{
"role": "assistant",
"content": "",
"tool_calls": [
{
"id": "call_50e7e0fe0989464a89f188eda443",
"type": "function",
"function": {
"name": "read",
"arguments": '{"filePath": "/tmp"}',
},
}
],
},
{"role": "tool", "content": "ok"},
]
with pytest.raises(Exception, match="Missing corresponding tool call"):
_gemini_convert_messages_with_history(
messages=messages,
model=model,
custom_llm_provider="vertex_ai",
)
def test_reasoning_effort_maps_to_thinking_level_gemini_3():

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