merge(litellm_internal_staging): resolve otel/team/org conflicts

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
shivam 2026-07-30 00:50:55 +00:00
commit 5ad0aeeca2
118 changed files with 5711 additions and 1243 deletions

1
.gitignore vendored
View file

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

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,10 +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";
/// 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,49 +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(response) => {
drop(response);
Err(litellm_core::CoreError::InvalidResponse(
"non-streaming messages execution returned a stream".to_string(),
))
}
}
}
pub(crate) enum MessagesResponse {
Json(Value),
Stream(reqwest::Response),
}
pub(crate) async fn execute_messages(
request: MessagesRequest<'_>,
stream: bool,
) -> CoreResult<MessagesResponse> {
let prepared = prepare_messages_call(request)?;
if stream {
execute_messages_provider_stream(prepared)
.await
.map(MessagesResponse::Stream)
} else {
execute_messages_provider_call(prepared)
.await
.map(MessagesResponse::Json)
}
}
#[cfg(test)]
mod tests;

View file

@ -1,24 +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) 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,12 +1,12 @@
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(reqwest::Response),
@ -52,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(upstream) => {
MessagesResponse::Stream(upstream)
}
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,11 +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 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 {

View file

@ -1,15 +1,13 @@
use litellm_core::CoreResult;
use litellm_core::error::CoreError;
use serde_json::Value;
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::ProviderMessagesRequest;
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
use super::types::{AnthropicMessagesResponse, ProviderMessagesRequest};
pub(super) async fn execute_messages_provider_call(
request: ProviderMessagesRequest,
) -> CoreResult<Value> {
) -> 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);
@ -39,12 +37,7 @@ pub(super) async fn execute_messages_provider_call(
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}"))
})
request.config.transform_response(&request.model, response)
}
pub(super) async fn execute_messages_provider_stream(

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,9 +1,8 @@
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 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};
pub(super) fn prepare_messages_call(

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

@ -30,15 +30,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
@ -133,6 +131,23 @@ def destination_resource_attrs(destination: "OtelDestination") -> Mapping[str, s
return dict(destination.resource_attributes)
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)
@ -200,8 +215,16 @@ def _otlp_metrics_endpoint(endpoint: str | None) -> str | None:
def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader":
"""Build a metric reader mirroring v1's exporter selection.
``console`` (and any unrecognized kind) exports to the console; ``otlp_http``/``otlp_grpc``
export over OTLP with the configured endpoint/headers, on a 5s period.
``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,
@ -213,18 +236,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,
@ -238,7 +255,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

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

@ -1173,7 +1173,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
@ -1190,6 +1189,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

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

@ -2008,17 +2008,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

@ -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( # noqa: C901 # single endpoint handling many optional
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( # noqa: C901 # single endpoint handling many optional
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,
@ -4856,7 +4907,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,
@ -5032,7 +5083,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:
@ -5105,7 +5156,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(
@ -5854,7 +5905,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

@ -1051,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,
@ -1089,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,

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

@ -144,7 +144,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)
@ -174,7 +174,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,
)
@ -513,7 +513,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
@ -775,7 +775,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,
@ -1470,9 +1470,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,
@ -1480,7 +1480,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,
)
@ -1716,7 +1716,7 @@ async def update_team(
)
# Verify caller has access to manage this team
team_for_auth = LiteLLM_TeamTable(**existing_team_row.model_dump())
team_for_auth = LiteLLM_TeamTable.model_validate(existing_team_row.model_dump())
await _verify_team_access(
team_obj=team_for_auth,
user_api_key_dict=user_api_key_dict,
@ -2017,7 +2017,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,
@ -2595,7 +2595,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,
@ -2640,10 +2640,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,
}
)
@ -2715,7 +2717,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
@ -2919,7 +2921,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
@ -3265,7 +3267,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(
@ -3389,12 +3391,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()
@ -3584,7 +3588,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 ##
@ -3619,9 +3623,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()
@ -3832,7 +3836,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,
)
@ -3881,7 +3885,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,
)
@ -3925,13 +3929,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
@ -4099,7 +4103,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):
@ -4714,7 +4718,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 (
@ -4814,7 +4818,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 (
@ -4882,7 +4886,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.
@ -4949,7 +4953,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.
@ -5210,7 +5214,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

@ -11214,11 +11214,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 +11296,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 +11831,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

@ -8636,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
@ -8670,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.
@ -8687,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
@ -8704,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]

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

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

@ -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).
@ -497,6 +521,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

@ -2334,9 +2334,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

@ -0,0 +1,133 @@
"""
The cache dashboard chart is fed by /global/activity/cache_hits. Aggregation
lives server-side: the SQL groups per call_type (splitting cache hits vs
successful vs failed requests; failed spend logs have call_type '' today and
must surface as 'Unknown'), and the endpoint returns chart-ready groups,
totals for the stat cards, and the filter options for the UI dropdowns.
"""
import json
from unittest.mock import AsyncMock, MagicMock
import pytest
from fastapi import HTTPException
from litellm.proxy.analytics_endpoints.analytics_endpoints import get_global_activity
from litellm.proxy.analytics_endpoints.cache_activity import (
GROUPS_SQL,
CacheActivityGroup,
compute_totals,
)
GROUP_ROWS = [
{
"call_type": "acompletion",
"api_requests": 1000,
"cache_hits": 300,
"failed_requests": 200,
"cached_completion_tokens": 12000,
"generated_completion_tokens": 48000,
},
{
"call_type": "Unknown",
"api_requests": 0,
"cache_hits": 0,
"failed_requests": 110,
"cached_completion_tokens": 0,
"generated_completion_tokens": 0,
},
]
KEY_ALIAS_ROWS = [{"key_alias": "Unnamed Key"}, {"key_alias": "my-key"}]
MODEL_ROWS = [{"model": "gpt-5.1"}]
def build_prisma(query_raw: AsyncMock) -> MagicMock:
prisma = MagicMock()
prisma.db.query_raw = query_raw
return prisma
def dispatching_query_raw() -> AsyncMock:
async def dispatch(sql: str, *params: object) -> list[dict[str, object]]:
if "GROUP BY" in sql:
return GROUP_ROWS
if "key_alias" in sql:
return KEY_ALIAS_ROWS
return MODEL_ROWS
return AsyncMock(side_effect=dispatch)
@pytest.fixture
def mock_prisma(monkeypatch: pytest.MonkeyPatch) -> MagicMock:
prisma = build_prisma(dispatching_query_raw())
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma)
return prisma
@pytest.mark.asyncio
async def test_returns_groups_totals_and_filter_options(mock_prisma: MagicMock):
response = await get_global_activity(
start_date="2026-07-01", end_date="2026-07-27", key_aliases=[], models=[]
)
assert [group.call_type for group in response.groups] == ["acompletion", "Unknown"]
assert response.groups[0].api_requests == 1000
assert response.groups[0].failed_requests == 200
assert response.totals.api_requests == 1000
assert response.totals.cache_hits == 300
assert response.totals.failed_requests == 310
assert response.totals.cached_completion_tokens == 12000
assert response.totals.cache_hit_ratio == pytest.approx((300 / 1610) * 100)
assert response.filter_options.key_aliases == ["Unnamed Key", "my-key"]
assert response.filter_options.models == ["gpt-5.1"]
@pytest.mark.asyncio
async def test_filters_are_passed_to_sql_as_json_arrays(mock_prisma: MagicMock):
await get_global_activity(
start_date="2026-07-01",
end_date="2026-07-27",
key_aliases=["my-key"],
models=["gpt-5.1", "claude-opus-4-8"],
)
groups_call = next(
call for call in mock_prisma.db.query_raw.call_args_list if "GROUP BY" in call.args[0]
)
assert groups_call.args[3] == json.dumps(["my-key"])
assert groups_call.args[4] == json.dumps(["gpt-5.1", "claude-opus-4-8"])
@pytest.mark.asyncio
async def test_rejects_malformed_dates_with_400(mock_prisma: MagicMock):
with pytest.raises(HTTPException) as exc_info:
await get_global_activity(start_date="07/01/2026", end_date="2026-07-27", key_aliases=[], models=[])
assert exc_info.value.status_code == 400
mock_prisma.db.query_raw.assert_not_called()
def test_totals_ratio_is_zero_without_requests():
totals = compute_totals([])
assert totals.cache_hit_ratio == 0.0
assert totals.api_requests == 0
def test_totals_denominator_includes_failed_requests():
group = CacheActivityGroup(
call_type="acompletion",
api_requests=60,
cache_hits=20,
failed_requests=20,
cached_completion_tokens=0,
generated_completion_tokens=0,
)
assert compute_totals([group]).cache_hit_ratio == pytest.approx(20.0)
def test_groups_sql_splits_failures_and_labels_empty_call_type_unknown():
assert "SUM(CASE WHEN sl.\"status\" = 'failure' THEN 1 ELSE 0 END)" in GROUPS_SQL
assert "CASE WHEN sl.\"call_type\" = '' THEN 'Unknown' ELSE sl.\"call_type\" END" in GROUPS_SQL

View file

@ -50,6 +50,7 @@ from litellm.proxy.auth.auth_checks import (
vector_store_access_check,
)
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.caching.redis_cache import RedisCache
from litellm.constants import DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
@ -4211,6 +4212,11 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias():
len(teams)==1 before populating the cache.
3. When team_alias is None, NO alias-key operation happens (no
delete of an empty-keyed entry, no spurious write).
4. DELETES the team_id-keyed entry from the internal usage cache
BEFORE the fresh write (LIT-4391). `_get_team_object_from_cache`
consults the internal usage cache first, so a leftover copy there
(backfilled from a Redis shared with `user_api_key_cache`) would
keep serving the pre-update team allowlist.
"""
from unittest.mock import AsyncMock, MagicMock
@ -4257,9 +4263,14 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias():
# (2) team_alias-keyed entry is deleted in BOTH the in-memory cache
# and the Redis dual cache (mirrors _delete_cache_key_object pattern).
cache.delete_cache.assert_called_once_with(key="team_alias:H-Capacity")
logging_obj.internal_usage_cache.dual_cache.async_delete_cache.assert_awaited_once_with(
key="team_alias:H-Capacity"
)
# (4) internal usage cache: team_id entry deleted BEFORE the fresh
# write, alias entry deleted as before.
internal_deleted_keys = [
(c.kwargs.get("key") or c.args[0])
for c in logging_obj.internal_usage_cache.dual_cache.async_delete_cache.await_args_list
]
assert internal_deleted_keys == ["team_id:team-1234", "team_alias:H-Capacity"]
# ===== team_alias is None: no alias-key operation =====
aliasless = LiteLLM_TeamTableCachedObj(**{**base_team_row, "team_alias": None})
@ -4277,7 +4288,9 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias():
)
cache2.delete_cache.assert_not_called()
logging_obj2.internal_usage_cache.dual_cache.async_delete_cache.assert_not_awaited()
logging_obj2.internal_usage_cache.dual_cache.async_delete_cache.assert_awaited_once_with(
key="team_id:team-no-alias"
)
written_keys_aliasless = [
(c.kwargs.get("key") or c.args[0])
for c in cache2.async_set_cache.await_args_list
@ -4285,6 +4298,145 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias():
assert written_keys_aliasless == ["team_id:team-no-alias"]
class _SharedFakeRedis(RedisCache):
"""Dict-backed stand-in for the single Redis that both
``user_api_key_cache`` (enable_redis_auth_cache) and
``proxy_logging_obj.internal_usage_cache.dual_cache`` share in the
LIT-4391 deployment topology. Only the methods DualCache calls are
implemented; ``super().__init__`` is skipped intentionally."""
def __init__(self):
self._store: dict = {}
async def async_set_cache(self, key, value, **kwargs):
self._store[key] = json.dumps(value)
async def async_get_cache(self, key, **kwargs):
raw = self._store.get(key)
return json.loads(raw) if raw is not None else None
async def async_delete_cache(self, key):
self._store.pop(key, None)
@pytest.mark.asyncio
async def test_team_update_not_shadowed_by_internal_usage_cache_lit_4391():
"""
Regression test for LIT-4391: keys with models=["all-team-models"] kept
getting 403 team_model_access_denied for models added via /team/update.
`_get_team_object_from_cache` consults the internal usage cache BEFORE
`user_api_key_cache`. When both share one Redis (enable_redis_auth_cache),
any team read backfills the internal cache's in-memory tier with the team
object. `_cache_team_object` (the /team/update refresh) only wrote
`user_api_key_cache`, so that backfilled copy kept shadowing the update
until its TTL expired — and the auth-time write-back then pushed the stale
copy back into the shared Redis, making the staleness self-sustaining.
Pins:
1. After `_cache_team_object` writes an updated team, `get_team_object`
returns the UPDATED model list even though the internal usage cache's
in-memory tier was backfilled with the pre-update team.
2. The shared Redis still holds the updated team afterwards — the
internal-cache invalidation must happen BEFORE the fresh write, or it
would wipe the value it just wrote.
"""
from litellm.caching.dual_cache import DualCache
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
from litellm.proxy.auth.auth_checks import _cache_team_object, get_team_object
team_id = "team-lit-4391"
shared_redis = _SharedFakeRedis()
user_api_key_cache = UserApiKeyCache(redis_cache=shared_redis)
proxy_logging_obj = MagicMock()
proxy_logging_obj.internal_usage_cache.dual_cache = DualCache(
redis_cache=shared_redis,
default_in_memory_ttl=300,
)
prisma_client = MagicMock()
await _cache_team_object(
team_id=team_id,
team_table=LiteLLM_TeamTableCachedObj(team_id=team_id, models=["model-a"]),
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
primed = await get_team_object(
team_id=team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
assert primed is not None and primed.models == ["model-a"]
await _cache_team_object(
team_id=team_id,
team_table=LiteLLM_TeamTableCachedObj(
team_id=team_id, models=["model-a", "model-b"]
),
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
refreshed = await get_team_object(
team_id=team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
assert refreshed is not None and refreshed.models == ["model-a", "model-b"], (
"get_team_object served a stale team allowlist after _cache_team_object "
f"refreshed it. Got models={refreshed.models if refreshed else None}"
)
redis_copy = await shared_redis.async_get_cache(f"team_id:{team_id}")
assert redis_copy is not None and redis_copy["models"] == ["model-a", "model-b"], (
"The shared Redis lost the refreshed team object — the internal-cache "
"invalidation must run BEFORE the fresh write, not after. "
f"Got: {redis_copy}"
)
@pytest.mark.asyncio
async def test_cache_team_object_tolerates_cache_invalidation_failures():
"""
Greptile review on the LIT-4391 fix: `_cache_team_object` runs after a
successful DB fetch (inside `get_team_object`) and after every team
mutation's DB write. A cache-backend error during the best-effort
invalidations must NOT fail those operations — otherwise a Redis blip
turns a healthy team lookup into a 404 and a committed /team/update into
a 500. The authoritative team_id-keyed write must still happen.
"""
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
from litellm.proxy.auth.auth_checks import _cache_team_object
cache = MagicMock()
cache.async_set_cache = AsyncMock()
cache.delete_cache = MagicMock(side_effect=Exception("redis down"))
logging_obj = MagicMock()
logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock(
side_effect=Exception("redis down")
)
await _cache_team_object(
team_id="team-cache-outage",
team_table=LiteLLM_TeamTableCachedObj(
team_id="team-cache-outage",
team_alias="cache-outage-alias",
models=["model-a"],
),
user_api_key_cache=cache,
proxy_logging_obj=logging_obj,
)
written_keys = [
(c.kwargs.get("key") or c.args[0])
for c in cache.async_set_cache.await_args_list
]
assert written_keys == ["team_id:team-cache-outage"]
MODEL_DISCOVERY_ROUTES = [
"/v1/models",
"/models",
@ -4762,3 +4914,254 @@ async def test_skip_user_budget_on_team_key_flag_restores_old_behavior():
request=MagicMock(spec=Request),
)
assert result is True
@pytest.mark.asyncio
async def test_get_default_end_user_budget_db_fetch_returns_validated_budget(monkeypatch):
from litellm.proxy.auth.auth_checks import get_default_end_user_budget
monkeypatch.setattr(litellm, "max_end_user_budget_id", "budget-default-1")
budget_row = MagicMock()
budget_row.dict = lambda: {"budget_id": "budget-default-1", "max_budget": 12.5, "tpm_limit": 100}
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row)
mock_cache = MagicMock()
mock_cache.async_get_cache = AsyncMock(return_value=None)
mock_cache.async_set_cache = AsyncMock()
result = await get_default_end_user_budget(
prisma_client=mock_prisma_client,
user_api_key_cache=mock_cache,
)
assert isinstance(result, LiteLLM_BudgetTable)
assert result.max_budget == 12.5
assert result.tpm_limit == 100
mock_cache.async_set_cache.assert_awaited_once()
assert mock_cache.async_set_cache.call_args.kwargs["value"] is result
@pytest.mark.asyncio
async def test_get_end_user_object_db_fetch_returns_validated_end_user():
from litellm.proxy.auth.auth_checks import get_end_user_object
end_user_row = MagicMock()
end_user_row.dict = lambda: {"user_id": "eu-1", "blocked": False, "spend": 3.0}
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_endusertable.find_unique = AsyncMock(return_value=end_user_row)
mock_cache = MagicMock()
mock_cache.async_get_cache = AsyncMock(return_value=None)
mock_cache.async_set_cache = AsyncMock()
result = await get_end_user_object(
end_user_id="eu-1",
prisma_client=mock_prisma_client,
user_api_key_cache=mock_cache,
)
assert isinstance(result, LiteLLM_EndUserTable)
assert result.user_id == "eu-1"
assert result.blocked is False
assert result.spend == 3.0
@pytest.mark.asyncio
async def test_get_team_membership_db_fetch_returns_validated_membership():
from litellm.proxy._types import LiteLLM_TeamMembership
from litellm.proxy.auth.auth_checks import get_team_membership
membership_row = MagicMock()
membership_row.dict = lambda: {"user_id": "u-1", "team_id": "t-1", "spend": 1.5}
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=membership_row)
mock_cache = MagicMock()
mock_cache.async_get_cache = AsyncMock(return_value=None)
mock_cache.async_set_cache = AsyncMock()
result = await get_team_membership(
user_id="u-1",
team_id="t-1",
prisma_client=mock_prisma_client,
user_api_key_cache=mock_cache,
)
assert isinstance(result, LiteLLM_TeamMembership)
assert result.user_id == "u-1"
assert result.team_id == "t-1"
assert result.spend == 1.5
@pytest.mark.asyncio
async def test_get_access_object_db_fetch_returns_validated_access_group():
from litellm.proxy._types import LiteLLM_AccessGroupTable
from litellm.proxy.auth.auth_checks import get_access_object
access_row = MagicMock()
access_row.dict = lambda: {
"access_group_id": "ag-1",
"access_group_name": "group one",
"access_model_names": ["gpt-4"],
}
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=access_row)
mock_cache = MagicMock()
mock_cache.async_get_cache = AsyncMock(return_value=None)
mock_cache.async_set_cache = AsyncMock()
result = await get_access_object(
access_group_id="ag-1",
prisma_client=mock_prisma_client,
user_api_key_cache=mock_cache,
proxy_logging_obj=None,
)
assert isinstance(result, LiteLLM_AccessGroupTable)
assert result.access_group_id == "ag-1"
assert result.access_model_names == ["gpt-4"]
@pytest.mark.asyncio
async def test_get_team_object_by_alias_db_fetch_returns_cached_obj():
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
from litellm.proxy.auth.auth_checks import get_team_object_by_alias
team_row = MagicMock()
team_row.model_dump = lambda: {"team_id": "t-9", "team_alias": "alias-9", "models": ["gpt-4"]}
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row])
mock_cache = MagicMock()
mock_cache.async_get_cache = AsyncMock(return_value=None)
mock_cache.async_set_cache = AsyncMock()
result = await get_team_object_by_alias(
team_alias="alias-9",
prisma_client=mock_prisma_client,
user_api_key_cache=mock_cache,
)
assert isinstance(result, LiteLLM_TeamTableCachedObj)
assert result.team_id == "t-9"
assert result.team_alias == "alias-9"
assert result.models == ["gpt-4"]
@pytest.mark.asyncio
async def test_get_org_object_by_alias_db_fetch_returns_validated_org():
from litellm.proxy._types import LiteLLM_OrganizationTable
from litellm.proxy.auth.auth_checks import get_org_object_by_alias
org_row = MagicMock()
org_row.model_dump = lambda: {
"organization_id": "org-1",
"organization_alias": "org-alias",
"budget_id": "b-1",
"created_by": "admin",
"updated_by": "admin",
"models": [],
}
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_organizationtable.find_many = AsyncMock(return_value=[org_row])
mock_cache = MagicMock()
mock_cache.async_get_cache = AsyncMock(return_value=None)
mock_cache.async_set_cache = AsyncMock()
result = await get_org_object_by_alias(
org_alias="org-alias",
prisma_client=mock_prisma_client,
user_api_key_cache=mock_cache,
)
assert isinstance(result, LiteLLM_OrganizationTable)
assert result.organization_id == "org-1"
assert result.budget_id == "b-1"
@pytest.mark.asyncio
async def test_get_object_permission_db_fetch_returns_validated_permission():
from litellm.proxy.auth.auth_checks import get_object_permission
perm_row = MagicMock()
perm_row.dict = lambda: {"object_permission_id": "op-1", "vector_stores": ["vs-1"]}
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=perm_row)
mock_cache = MagicMock()
mock_cache.async_get_cache = AsyncMock(return_value=None)
mock_cache.async_set_cache = AsyncMock()
result = await get_object_permission(
object_permission_id="op-1",
prisma_client=mock_prisma_client,
user_api_key_cache=mock_cache,
)
assert isinstance(result, LiteLLM_ObjectPermissionTable)
assert result.object_permission_id == "op-1"
assert result.vector_stores == ["vs-1"]
@pytest.mark.asyncio
async def test_get_managed_vector_store_rows_by_uuids_db_fetch_validates_rows():
from litellm.proxy._types import LiteLLM_ManagedVectorStoresTable
from litellm.proxy.auth.auth_checks import get_managed_vector_store_rows_by_uuids
vs_row = MagicMock()
vs_row.model_dump = lambda: {"vector_store_id": "vs-7", "custom_llm_provider": "openai"}
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_managedvectorstorestable.find_many = AsyncMock(return_value=[vs_row])
mock_cache = MagicMock()
mock_cache.async_get_cache = AsyncMock(return_value=None)
mock_cache.async_set_cache = AsyncMock()
result = await get_managed_vector_store_rows_by_uuids(
uuids=["vs-7"],
prisma_client=mock_prisma_client,
user_api_key_cache=mock_cache,
)
assert len(result) == 1
assert isinstance(result[0], LiteLLM_ManagedVectorStoresTable)
assert result[0].vector_store_id == "vs-7"
assert result[0].custom_llm_provider == "openai"
@pytest.mark.asyncio
async def test_get_project_object_db_fetch_returns_cached_obj():
from litellm.proxy._types import LiteLLM_ProjectTableCachedObj
from litellm.proxy.auth.auth_checks import get_project_object
project_row = MagicMock()
project_row.model_dump = lambda: {"project_id": "p-1", "project_alias": "proj", "team_id": "t-1"}
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_projecttable.find_unique = AsyncMock(return_value=project_row)
mock_cache = MagicMock()
mock_cache.async_get_cache = AsyncMock(return_value=None)
mock_cache.async_set_cache = AsyncMock()
result = await get_project_object(
project_id="p-1",
prisma_client=mock_prisma_client,
user_api_key_cache=mock_cache,
)
assert isinstance(result, LiteLLM_ProjectTableCachedObj)
assert result.project_id == "p-1"
assert result.project_alias == "proj"

View file

@ -2918,6 +2918,123 @@ async def test_team_metadata_refreshed_from_team_object_during_auth():
setattr(_proxy_server_mod, k, v)
@pytest.mark.asyncio
async def test_auth_flow_never_persists_fallback_team_object_lit_4391():
"""
Regression test for LIT-4391 (stale team allowlist poisoning).
When `get_team_object` fails at the "Check 6" team-auth step (cache miss
inside the DB-throttle window, DB blip, ...), the builder falls back to a
team object reconstructed from the CACHED token's team_* snapshot — which
can be arbitrarily stale (e.g. pre-/team/update models).
The builder used to write that team object back into `user_api_key_cache`
under "team_id:<id>" after Check 6. Writing a cache-read (or worse, a
token-snapshot) value back into the shared cache re-poisons it — with
enable_redis_auth_cache it clobbered the fresh team `/team/update` had
just written to Redis, making the stale allowlist self-sustaining across
requests. Only authoritative writers (`_cache_team_object` on DB reads and
team mutations) may populate the team cache.
Pins: the auth flow completes on the fallback path WITHOUT writing any
"team_id:*" cache entry.
"""
from starlette.datastructures import URL
from starlette.requests import Request
from fastapi import HTTPException
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder
api_key = "sk-test-lit-4391-no-team-writeback"
valid_token = UserAPIKeyAuth(
api_key=api_key,
token=api_key,
user_role=LitellmUserRoles.INTERNAL_USER,
team_id="team-lit-4391",
team_models=["model-a"],
models=["all-team-models"],
)
mock_cache = AsyncMock()
mock_cache.async_get_cache = AsyncMock(return_value=valid_token)
mock_cache.async_set_cache = AsyncMock(return_value=None)
mock_proxy_logging_obj = MagicMock()
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
import litellm.proxy.proxy_server as _proxy_server_mod
_attrs = {
"prisma_client": MagicMock(),
"user_api_key_cache": mock_cache,
"proxy_logging_obj": mock_proxy_logging_obj,
"master_key": "sk-master-key",
"general_settings": {},
"llm_model_list": [],
"llm_router": None,
"open_telemetry_logger": None,
"model_max_budget_limiter": MagicMock(),
"user_custom_auth": None,
"jwt_handler": None,
"litellm_proxy_admin_name": "admin",
}
_originals = {k: getattr(_proxy_server_mod, k, None) for k in _attrs}
try:
for k, v in _attrs.items():
setattr(_proxy_server_mod, k, v)
request = Request(scope={"type": "http"})
request._url = URL(url="/chat/completions")
with (
patch(
"litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key",
new_callable=AsyncMock,
return_value=valid_token,
),
patch(
"litellm.proxy.auth.user_api_key_auth.get_team_object",
new_callable=AsyncMock,
side_effect=HTTPException(
status_code=404,
detail={"error": "Team doesn't exist in db."},
),
),
):
result = await _user_api_key_auth_builder(
request=request,
api_key=f"Bearer {api_key}",
azure_api_key_header="",
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
request_data={},
)
assert result.team_id == "team-lit-4391"
team_cache_writes = [
key
for c in mock_cache.async_set_cache.await_args_list
if isinstance(key := (c.kwargs.get("key") if "key" in c.kwargs else c.args[0]), str)
and key.startswith("team_id:")
]
assert team_cache_writes == [], (
"The auth flow wrote a team object into the cache. Fallback/"
"cache-read team objects must never be persisted — only "
"_cache_team_object (DB reads and team mutations) may write "
f"'team_id:*' entries. Got writes: {team_cache_writes}"
)
finally:
for k, v in _originals.items():
setattr(_proxy_server_mod, k, v)
# ---------------------------------------------------------------------------
# _run_centralized_common_checks — centralized authz gate
@ -4576,101 +4693,6 @@ async def test_real_jwt_still_requires_license_when_jwt_auth_enabled(monkeypatch
assert "enterprise only feature" in message
@pytest.mark.asyncio
async def test_auth_path_caches_team_object_under_canonical_team_id_key():
"""Regression for LIT-4000: the auth builder must cache the team object under
the canonical ``team_id:{id}`` key that ``get_team_object`` and
``_update_team_cache`` read, never under the raw ``team_id`` (and never under
a ``None`` key, which Redis rejects with a NoneType key error). A raw or None
key is silently dropped by Redis / never served back, so every request
re-hits Postgres for the team object instead of the L2 cache.
Drives the real builder for a team-scoped key against a real in-memory
``UserApiKeyCache`` and reads the team object back. Mutating the cache key at
the write site to the raw ``valid_token.team_id`` (or ``None``) makes the
canonical-key read miss and fails this test.
"""
from fastapi import Request
from starlette.datastructures import URL
import litellm.proxy.proxy_server as _proxy_server_mod
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.proxy_server import hash_token
team_id = "team-lit-4000"
api_key = "sk-lit-4000-team-key"
cache = UserApiKeyCache()
team_token = UserAPIKeyAuth(token=hash_token(api_key), team_id=team_id)
team_obj = LiteLLM_TeamTableCachedObj(team_id=team_id)
proxy_logging_obj = MagicMock()
proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
attrs = {
"prisma_client": MagicMock(),
"user_api_key_cache": cache,
"proxy_logging_obj": proxy_logging_obj,
"master_key": "sk-test-master",
"general_settings": {"allow_requests_on_db_unavailable": False},
"llm_model_list": [],
"llm_router": None,
"open_telemetry_logger": None,
"model_max_budget_limiter": MagicMock(),
"user_custom_auth": None,
"jwt_handler": None,
"litellm_proxy_admin_name": "admin",
}
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
try:
for k, v in attrs.items():
setattr(_proxy_server_mod, k, v)
request = Request(scope={"type": "http"})
request._url = URL(url="/chat/completions")
with (
patch(
"litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key",
AsyncMock(return_value=team_token),
),
patch(
"litellm.proxy.auth.user_api_key_auth.get_team_object",
AsyncMock(return_value=team_obj),
),
patch(
"litellm.proxy.auth.user_api_key_auth._enforce_key_and_fallback_model_access",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.auth.user_api_key_auth._return_user_api_key_auth_obj",
new_callable=AsyncMock,
return_value=team_token,
),
patch(
"litellm.proxy.auth.auth_exception_handler.seed_request_identity",
),
):
await _user_api_key_auth_builder(
request=request,
api_key=f"Bearer {api_key}",
azure_api_key_header="",
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
request_data={},
)
finally:
for k, v in originals.items():
setattr(_proxy_server_mod, k, v)
served = cache.get_cache(
key=f"team_id:{team_id}", model_type=LiteLLM_TeamTableCachedObj
)
assert served is not None and served.team_id == team_id
assert cache.get_cache(key=team_id) is None
assert cache.get_cache(key=None) is None
@pytest.mark.asyncio
async def test_auth_does_not_rewrite_cached_key_object_back_into_cache():
"""A cache-hit auth must not write the token back into the cache.

View file

@ -51,7 +51,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
from litellm.proxy.utils import ProxyLogging
from litellm.router import Router
from litellm.types.llms.openai import BatchJobStatus
from litellm.types.utils import LiteLLMBatch
from litellm.types.utils import CredentialItem, LiteLLMBatch
from fastapi import Response
@ -2091,3 +2091,154 @@ async def test_retrieve__unified_no_router_500(retrieve_harness):
assert exc.value.code == "500"
retrieve_harness.router_aretrieve.assert_not_called()
retrieve_harness.litellm_aretrieve.assert_not_called()
# =========================================================================== #
# SCENARIO 3 + configured deployments: a provider-only call (custom-llm-provider
# header, no model anywhere) must resolve the gateway/team deployment's named
# credential for that provider and attach it to the provider call kwargs,
# instead of silently falling through to the host environment's default
# credentials (regression: vertex batch jobs landing in the hosting env's GCP
# project because litellm_credential_name never reached the call).
# =========================================================================== #
VERTEX_NAMED_CREDENTIAL = CredentialItem(
credential_name="vertex-named-cred",
credential_info={},
credential_values={
"vertex_project": "customer-project",
"vertex_location": "us-central1",
"vertex_credentials": "/creds/customer-sa.json",
},
)
def vertex_named_credential_router() -> Router:
return Router(
model_list=[
{
"model_name": "gemini-2.5-pro",
"litellm_params": {
"model": "vertex_ai/gemini-2.5-pro",
"litellm_credential_name": "vertex-named-cred",
},
}
]
)
@pytest.mark.asyncio
async def test_create__provider_only_resolves_named_vertex_credentials(harness):
"""Provider-only create must attach the configured named credential, and must
NOT turn the call into a model-routed one (no model kwarg injected)."""
set_body(
harness,
{
"input_file_id": "file-plain",
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
},
)
harness.provider_from_headers.return_value = "vertex_ai"
with patch.object(litellm, "credential_list", [VERTEX_NAMED_CREDENTIAL]):
with patch.object(proxy_server, "llm_router", vertex_named_credential_router()):
await call_create(harness)
assert harness.acreate_kwargs() == {
"custom_llm_provider": "vertex_ai",
"input_file_id": "file-plain",
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
"metadata": None,
"vertex_project": "customer-project",
"vertex_location": "us-central1",
"vertex_credentials": "/creds/customer-sa.json",
}
@pytest.mark.asyncio
async def test_create__provider_only_ignores_other_provider_deployments(harness):
"""A provider-only vertex call must not pick up credentials from deployments
of a different provider; with no vertex deployment the payload is exactly the
pre-fix env-var fallback."""
set_body(
harness,
{
"input_file_id": "file-plain",
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
},
)
harness.provider_from_headers.return_value = "vertex_ai"
openai_only_router = Router(
model_list=[
{
"model_name": "gpt-4o",
"litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-global-openai"},
}
]
)
with patch.object(proxy_server, "llm_router", openai_only_router):
await call_create(harness)
assert harness.acreate_kwargs() == {
"custom_llm_provider": "vertex_ai",
"input_file_id": "file-plain",
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
"metadata": None,
}
@pytest.mark.asyncio
async def test_retrieve__provider_only_resolves_named_vertex_credentials(retrieve_harness):
retrieve_harness.provider_from_headers.return_value = "vertex_ai"
with patch.object(litellm, "credential_list", [VERTEX_NAMED_CREDENTIAL]):
with patch.object(proxy_server, "llm_router", vertex_named_credential_router()):
await call_retrieve(retrieve_harness, "batch-raw-xyz")
assert retrieve_harness.aretrieve_kwargs() == {
"custom_llm_provider": "vertex_ai",
"batch_id": "batch-raw-xyz",
"vertex_project": "customer-project",
"vertex_location": "us-central1",
"vertex_credentials": "/creds/customer-sa.json",
}
@pytest.mark.asyncio
async def test_list__provider_only_resolves_named_vertex_credentials(list_harness):
list_harness.provider_from_headers.return_value = "vertex_ai"
with patch.object(litellm, "credential_list", [VERTEX_NAMED_CREDENTIAL]):
with patch.object(proxy_server, "llm_router", vertex_named_credential_router()):
await call_list(list_harness)
assert list_harness.alist_kwargs() == {
"custom_llm_provider": "vertex_ai",
"after": None,
"limit": None,
"vertex_project": "customer-project",
"vertex_location": "us-central1",
"vertex_credentials": "/creds/customer-sa.json",
}
@pytest.mark.asyncio
async def test_cancel__provider_only_resolves_named_vertex_credentials(cancel_harness):
cancel_harness.provider_from_headers.return_value = "vertex_ai"
with patch.object(litellm, "credential_list", [VERTEX_NAMED_CREDENTIAL]):
with patch.object(proxy_server, "llm_router", vertex_named_credential_router()):
await call_cancel(cancel_harness, "batch-raw-xyz")
assert cancel_harness.acancel_kwargs() == {
"custom_llm_provider": "vertex_ai",
"batch_id": "batch-raw-xyz",
"vertex_project": "customer-project",
"vertex_location": "us-central1",
"vertex_credentials": "/creds/customer-sa.json",
}

View file

@ -332,6 +332,21 @@ class TestScimTransformations:
assert scim_group.members[1].value == "test2@example.com"
assert scim_group.members[1].display == "test2@example.com"
@pytest.mark.asyncio
async def test_transform_team_marks_members_as_users(
self, mock_team, mock_prisma_client
):
"""A LiteLLM team only holds users, and stating the member type keeps the
response from emitting a null ``type`` now that SCIMMember carries one."""
mock_client, _ = mock_prisma_client
with patch("litellm.proxy.proxy_server.prisma_client", mock_client):
scim_group = await ScimTransformations.transform_litellm_team_to_scim_group(
mock_team
)
assert [member.type for member in scim_group.members] == ["User", "User"]
def test_get_scim_user_name(self, mock_user, mock_user_minimal):
# User with email
result = ScimTransformations._get_scim_user_name(mock_user)

View file

@ -108,6 +108,19 @@ class TestGetSchemas:
assert "displayName" in attr_names
assert "members" in attr_names
def test_group_schema_advertises_member_type(self):
"""IdPs read the schema to learn we understand ``members.type``, which is how
a nested group announces itself."""
schemas = _get_schemas()
group_schema = next(
s for s in schemas if s.id == "urn:ietf:params:scim:schemas:core:2.0:Group"
)
members = next(a for a in group_schema.attributes if a.name == "members")
member_type = next(a for a in members.subAttributes or [] if a.name == "type")
assert member_type.type == "string"
assert member_type.multiValued is False
assert "Group" in (member_type.description or "")
def test_schema_meta_fields(self):
schemas = _get_schemas()
user_schema = next(

View file

@ -20,8 +20,10 @@ from litellm.proxy.management_endpoints.scim.scim_v2 import (
_extract_group_member_ids,
_extract_ids_from_path_filter,
_handle_team_membership_changes,
_parse_member_entries,
_process_group_patch_operations,
_recompute_scim_member_roles,
_resolve_group_member_ids,
create_group,
create_user,
delete_group,
@ -36,6 +38,8 @@ from litellm.proxy.management_endpoints.scim.scim_v2 import (
)
from litellm.types.proxy.management_endpoints.scim_v2 import (
SCIM_ENTERPRISE_USER_SCHEMA,
SCIM_MANAGED_TEAM_METADATA_KEY,
SCIM_TEAM_DATA_METADATA_KEY,
SCIMGroup,
SCIMMember,
SCIMPatchOp,
@ -1611,7 +1615,10 @@ async def test_update_group_with_nonexistent_users_rejects(mocker, monkeypatch):
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
# Mock team operations
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team)
def mock_team_lookup(where):
return mock_existing_team if where["team_id"] == group_id else None
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(side_effect=mock_team_lookup)
# Mock updated team response
mock_updated_team = mocker.MagicMock()
@ -1775,6 +1782,7 @@ async def test_extract_group_member_ids_with_flag_true_creates_users(mocker, mon
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.db = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None)
# Mock user lookup - only existing-user exists initially
def mock_user_lookup(where):
@ -1842,6 +1850,7 @@ async def test_extract_group_member_ids_with_flag_false_rejects(mocker, monkeypa
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.db = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None)
# Mock user lookup - only existing-user exists
def mock_user_lookup(where):
@ -1902,6 +1911,7 @@ async def test_process_group_patch_operations_with_flag_true_creates_users(mocke
# Mock user lookup - new-user-1 doesn't exist
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None)
# Mock user creation
created_user = NewUserResponse(user_id="new-user-1", key="test-key-1")
@ -1956,6 +1966,7 @@ async def test_process_group_patch_operations_with_flag_false_rejects(mocker, mo
# Mock user lookup - new-user-1 doesn't exist
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None)
# Execute the function - should raise HTTPException
with pytest.raises(HTTPException) as exc_info:
@ -3519,3 +3530,874 @@ async def test_process_group_patch_replace_empty_value_does_not_use_path_filter(
)
assert final_members == set()
def _member_resolution_prisma(mocker, *, users: set, teams: set, unmanaged_teams: frozenset = frozenset()):
"""Prisma mock where only the given ids resolve to a user row / team row.
``teams`` are teams a SCIM group write created, so they carry provenance;
``unmanaged_teams`` resolve too but look like a team an admin created here.
"""
def team_row(team_id: str) -> LiteLLM_TeamTable | None:
if team_id in teams:
return LiteLLM_TeamTable(team_id=team_id, metadata={SCIM_MANAGED_TEAM_METADATA_KEY: True})
if team_id in unmanaged_teams:
return LiteLLM_TeamTable(team_id=team_id, metadata={})
return None
prisma_client = mocker.MagicMock()
prisma_client.db = mocker.MagicMock()
prisma_client.db.litellm_usertable = mocker.MagicMock()
prisma_client.db.litellm_usertable.find_unique = AsyncMock(
side_effect=lambda where: (
LiteLLM_UserTable(user_id=where["user_id"]) if where["user_id"] in users else None
)
)
prisma_client.db.litellm_teamtable = mocker.MagicMock()
prisma_client.db.litellm_teamtable.find_unique = AsyncMock(side_effect=lambda where: team_row(where["team_id"]))
return prisma_client
@pytest.fixture
def scim_upsert_user_enabled(monkeypatch):
from litellm.proxy.proxy_server import proxy_config
async def mock_get_config():
return {"litellm_settings": {"scim_upsert_user": True}}
monkeypatch.setattr(proxy_config, "get_config", mock_get_config)
@pytest.fixture
def scim_upsert_user_disabled(monkeypatch):
from litellm.proxy.proxy_server import proxy_config
async def mock_get_config():
return {"litellm_settings": {"scim_upsert_user": False}}
monkeypatch.setattr(proxy_config, "get_config", mock_get_config)
@pytest.mark.asyncio
async def test_create_group_ignores_nested_group_members(mocker, scim_upsert_user_enabled):
"""Entra sends nested groups as members with ``type: "Group"``. Treating that
GUID as a user id provisioned a phantom internal user per nested group."""
nested_group_id = "8f1e9d70-0000-4a0e-9a1e-nested"
scim_group = SCIMGroup(
schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"],
id="parent-group",
displayName="Parent Group",
members=[
SCIMMember(value="real-user", display="Real User", type="User"),
SCIMMember(value=nested_group_id, display="Nested Group", type="Group"),
],
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
AsyncMock(return_value=_member_resolution_prisma(mocker, users={"real-user"}, teams=set())),
)
create_user_mock = mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists",
AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")),
)
new_team_mock = mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2.new_team",
AsyncMock(return_value=mocker.MagicMock()),
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group",
AsyncMock(return_value=scim_group),
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._recompute_scim_member_roles",
AsyncMock(),
)
await create_group(group=scim_group)
create_user_mock.assert_not_called()
assert new_team_mock.call_args.kwargs["data"].members_with_roles == [Member(user_id="real-user", role="user")]
@pytest.mark.asyncio
async def test_update_group_ignores_nested_group_members(mocker, scim_upsert_user_enabled):
"""PUT /Groups must drop nested-group members too, so a full sync from the IdP
neither provisions nor enrolls the nested group's GUID."""
group_id = "parent-group"
nested_group_id = "8f1e9d70-0000-4a0e-9a1e-nested"
existing_team = LiteLLM_TeamTable(
team_id=group_id,
team_alias="Parent Group",
members=[],
members_with_roles=[],
metadata={},
)
scim_group = SCIMGroup(
schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"],
id=group_id,
displayName="Parent Group",
members=[
SCIMMember(value="real-user", display="Real User", type="User"),
SCIMMember(value=nested_group_id, display="Nested Group", type="Group"),
],
)
prisma_client = _member_resolution_prisma(mocker, users={"real-user"}, teams={group_id})
prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=existing_team)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
AsyncMock(return_value=prisma_client),
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._check_team_exists",
AsyncMock(return_value=existing_team),
)
create_user_mock = mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists",
AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")),
)
patch_membership_mock = mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership",
AsyncMock(),
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._recompute_scim_member_roles",
AsyncMock(),
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group",
AsyncMock(return_value=scim_group),
)
await update_group(group_id=group_id, group=scim_group)
create_user_mock.assert_not_called()
enrolled = {call.kwargs["user_id"] for call in patch_membership_mock.call_args_list}
assert enrolled == {"real-user"}
@pytest.mark.asyncio
async def test_process_group_patch_operations_ignores_nested_group_members(mocker, scim_upsert_user_enabled):
"""PATCH bodies bypass SCIMGroup parsing, so ``type`` must be read off the raw
member dicts; otherwise a nested group is indistinguishable from a user id."""
nested_group_id = "8f1e9d70-0000-4a0e-9a1e-nested"
patch_ops = SCIMPatchOp(
schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"],
Operations=[
SCIMPatchOperation(
op="add",
path="members",
value=[
{"value": "real-user", "display": "Real User", "type": "User"},
{"value": nested_group_id, "display": "Nested Group", "type": "Group"},
],
)
],
)
existing_team = LiteLLM_TeamTable(
team_id="parent-group",
team_alias="Parent Group",
members=[],
members_with_roles=[Member(user_id="incumbent", role="user")],
)
create_user_mock = mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists",
AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")),
)
_, final_members, _ = await _process_group_patch_operations(
patch_ops=patch_ops,
existing_team=existing_team,
prisma_client=_member_resolution_prisma(mocker, users={"real-user"}, teams=set()),
)
create_user_mock.assert_not_called()
assert final_members == {"incumbent", "real-user"}
@pytest.mark.asyncio
async def test_process_group_patch_operations_ignores_lowercase_group_type(mocker, scim_upsert_user_enabled):
"""The ``type`` comparison is case-insensitive; IdPs are not consistent about it."""
nested_group_id = "8f1e9d70-0000-4a0e-9a1e-nested"
patch_ops = SCIMPatchOp(
schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"],
Operations=[
SCIMPatchOperation(op="add", path="members", value=[{"value": nested_group_id, "type": "group"}])
],
)
existing_team = LiteLLM_TeamTable(
team_id="parent-group",
team_alias="Parent Group",
members=[],
members_with_roles=[],
)
create_user_mock = mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists",
AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")),
)
_, final_members, _ = await _process_group_patch_operations(
patch_ops=patch_ops,
existing_team=existing_team,
prisma_client=_member_resolution_prisma(mocker, users=set(), teams=set()),
)
create_user_mock.assert_not_called()
assert final_members == set()
@pytest.mark.asyncio
async def test_process_group_patch_operations_skips_member_matching_existing_team(mocker, scim_upsert_user_enabled):
"""Okta sends filtered paths and untyped ids, so a nested group arrives with no
``type`` at all; an id that names an existing team is still not a user."""
patch_ops = SCIMPatchOp(
schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"],
Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "child-team"}])],
)
existing_team = LiteLLM_TeamTable(
team_id="parent-group",
team_alias="Parent Group",
members=[],
members_with_roles=[Member(user_id="incumbent", role="user")],
)
create_user_mock = mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists",
AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")),
)
_, final_members, _ = await _process_group_patch_operations(
patch_ops=patch_ops,
existing_team=existing_team,
prisma_client=_member_resolution_prisma(mocker, users=set(), teams={"child-team", "parent-group"}),
)
create_user_mock.assert_not_called()
assert final_members == {"incumbent"}
@pytest.mark.asyncio
async def test_process_group_patch_operations_prefers_user_over_team_for_colliding_id(
mocker, scim_upsert_user_enabled
):
"""Nothing stops a user id from also being a team id, so the user lookup has to
win; ordering the team check first would silently stop syncing that user."""
patch_ops = SCIMPatchOp(
schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"],
Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "dual-id"}])],
)
existing_team = LiteLLM_TeamTable(
team_id="parent-group",
team_alias="Parent Group",
members=[],
members_with_roles=[],
)
_, final_members, _ = await _process_group_patch_operations(
patch_ops=patch_ops,
existing_team=existing_team,
prisma_client=_member_resolution_prisma(mocker, users={"dual-id"}, teams={"dual-id"}),
)
assert final_members == {"dual-id"}
@pytest.mark.asyncio
async def test_create_group_strict_mode_accepts_group_and_team_members(mocker, scim_upsert_user_disabled):
"""Strict mode (scim_upsert_user=False) rejects unknown *users*; a nested group
is not a user, so it must be dropped rather than 400 the whole sync."""
scim_group = SCIMGroup(
schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"],
id="parent-group",
displayName="Parent Group",
members=[
SCIMMember(value="real-user", type="User"),
SCIMMember(value="nested-group-guid", type="Group"),
SCIMMember(value="child-team"),
],
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
AsyncMock(return_value=_member_resolution_prisma(mocker, users={"real-user"}, teams={"child-team"})),
)
create_user_mock = mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists",
AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")),
)
new_team_mock = mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2.new_team",
AsyncMock(return_value=mocker.MagicMock()),
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group",
AsyncMock(return_value=scim_group),
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._recompute_scim_member_roles",
AsyncMock(),
)
await create_group(group=scim_group)
create_user_mock.assert_not_called()
assert new_team_mock.call_args.kwargs["data"].members_with_roles == [Member(user_id="real-user", role="user")]
@pytest.mark.asyncio
async def test_create_group_strict_mode_still_rejects_unknown_user(mocker, scim_upsert_user_disabled):
"""The strict-mode 400 must name the unknown *user* and stay quiet about the
nested group sharing the request."""
scim_group = SCIMGroup(
schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"],
id="parent-group",
displayName="Parent Group",
members=[
SCIMMember(value="nested-group-guid", type="Group"),
SCIMMember(value="unknown-user"),
],
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
AsyncMock(return_value=_member_resolution_prisma(mocker, users=set(), teams=set())),
)
with pytest.raises(ProxyException) as exc_info:
await create_group(group=scim_group)
assert int(exc_info.value.code) == 400
assert "unknown-user" in str(exc_info.value.message)
assert "nested-group-guid" not in str(exc_info.value.message)
@pytest.mark.asyncio
async def test_process_group_patch_remove_unknown_member_does_not_create_user(mocker, scim_upsert_user_enabled):
"""A ``remove`` of an id we don't know is an idempotent no-op. Upserting the id
first, only to drop it from the roster, made removals a phantom-user factory."""
patch_ops = SCIMPatchOp(
schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"],
Operations=[SCIMPatchOperation(op="remove", path="members", value=[{"value": "long-gone"}])],
)
existing_team = LiteLLM_TeamTable(
team_id="parent-group",
team_alias="Parent Group",
members=[],
members_with_roles=[Member(user_id="keep-user", role="user")],
)
create_user_mock = mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists",
AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")),
)
_, final_members, _ = await _process_group_patch_operations(
patch_ops=patch_ops,
existing_team=existing_team,
prisma_client=_member_resolution_prisma(mocker, users={"keep-user"}, teams=set()),
)
create_user_mock.assert_not_called()
assert final_members == {"keep-user"}
@pytest.mark.parametrize(
"operation",
[
SCIMPatchOperation(op="remove", path='members[value eq "long-gone"]', value=None),
SCIMPatchOperation(op="remove", path="members", value=[{"value": "long-gone"}]),
SCIMPatchOperation(op="remove", path="members", value=[{"value": "long-gone", "type": "Group"}]),
],
ids=["path-filter", "unknown-id", "nested-group"],
)
@pytest.mark.asyncio
async def test_process_group_patch_remove_unknown_member_does_not_reject_in_strict_mode(
mocker, scim_upsert_user_disabled, operation
):
"""Strict mode must not 400 a removal: refusing to drop an id the IdP already
forgot leaves the roster permanently out of sync."""
patch_ops = SCIMPatchOp(
schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"],
Operations=[operation],
)
existing_team = LiteLLM_TeamTable(
team_id="parent-group",
team_alias="Parent Group",
members=[],
members_with_roles=[Member(user_id="keep-user", role="user")],
)
_, final_members, _ = await _process_group_patch_operations(
patch_ops=patch_ops,
existing_team=existing_team,
prisma_client=_member_resolution_prisma(mocker, users={"keep-user"}, teams=set()),
)
assert final_members == {"keep-user"}
@pytest.mark.asyncio
async def test_process_group_patch_remove_drops_member_without_user_row(mocker, scim_upsert_user_enabled):
"""Phantom members already on a roster (their user row is gone) must still be
removable, so the removal id is honoured even though it resolves to nothing."""
patch_ops = SCIMPatchOp(
schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"],
Operations=[SCIMPatchOperation(op="remove", path="members", value=[{"value": "phantom"}])],
)
existing_team = LiteLLM_TeamTable(
team_id="parent-group",
team_alias="Parent Group",
members=[],
members_with_roles=[
Member(user_id="keep-user", role="user"),
Member(user_id="phantom", role="user"),
],
)
create_user_mock = mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists",
AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")),
)
_, final_members, _ = await _process_group_patch_operations(
patch_ops=patch_ops,
existing_team=existing_team,
prisma_client=_member_resolution_prisma(mocker, users={"keep-user"}, teams=set()),
)
create_user_mock.assert_not_called()
assert final_members == {"keep-user"}
_NESTED_GROUP_ID = "8f1e9d70-0000-4a0e-9a1e-nested"
@pytest.mark.parametrize(
"member_entry, user_rows, team_rows",
[
({"value": _NESTED_GROUP_ID, "type": "Group"}, {"keep-user", _NESTED_GROUP_ID}, set()),
({"value": _NESTED_GROUP_ID, "type": "Group"}, {"keep-user"}, set()),
({"value": _NESTED_GROUP_ID, "type": "Group"}, {"keep-user"}, {_NESTED_GROUP_ID}),
({"value": _NESTED_GROUP_ID}, {"keep-user"}, {_NESTED_GROUP_ID}),
],
ids=["phantom-user-row-exists", "user-row-already-deleted", "child-group-is-a-team", "untyped-team-id"],
)
@pytest.mark.asyncio
async def test_process_group_patch_remove_discards_non_user_member(
mocker, scim_upsert_user_enabled, member_entry, user_rows, team_rows
):
"""Rosters written before nested groups were understood still carry those ids,
and the IdP removes them exactly as it added them; a removal that resolved its
ids first would classify them as non-users and leave them stuck on the team."""
patch_ops = SCIMPatchOp(
schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"],
Operations=[SCIMPatchOperation(op="remove", path="members", value=[member_entry])],
)
existing_team = LiteLLM_TeamTable(
team_id="parent-group",
team_alias="Parent Group",
members=[],
members_with_roles=[
Member(user_id="keep-user", role="user"),
Member(user_id=_NESTED_GROUP_ID, role="user"),
],
)
create_user_mock = mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists",
AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")),
)
_, final_members, _ = await _process_group_patch_operations(
patch_ops=patch_ops,
existing_team=existing_team,
prisma_client=_member_resolution_prisma(mocker, users=user_rows, teams=team_rows),
)
create_user_mock.assert_not_called()
assert final_members == {"keep-user"}
@pytest.mark.asyncio
async def test_process_group_patch_add_keeps_member_typed_user_that_collides_with_team_id(
mocker, scim_upsert_user_enabled
):
"""The team lookup only exists to catch nested groups that arrive untyped. An id
the IdP calls a User is a user, and IdP ids collide with team ids easily."""
patch_ops = SCIMPatchOp(
schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"],
Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "123456", "type": "User"}])],
)
existing_team = LiteLLM_TeamTable(
team_id="parent-group",
team_alias="Parent Group",
members=[],
members_with_roles=[],
)
create_user_mock = mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists",
AsyncMock(return_value=NewUserResponse(user_id="123456", key="new-key")),
)
_, final_members, _ = await _process_group_patch_operations(
patch_ops=patch_ops,
existing_team=existing_team,
prisma_client=_member_resolution_prisma(mocker, users=set(), teams={"123456", "parent-group"}),
)
assert create_user_mock.call_args.kwargs["user_id"] == "123456"
assert final_members == {"123456"}
@pytest.mark.parametrize("member_type", ["Device", " group ", "Machine"])
@pytest.mark.asyncio
async def test_process_group_patch_operations_skips_non_user_member_types(
mocker, scim_upsert_user_enabled, member_type
):
"""A team holds users, so a member that declares itself to be anything else is
dropped; enumerating the types worth skipping would leave the next one to leak."""
patch_ops = SCIMPatchOp(
schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"],
Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "not-a-user", "type": member_type}])],
)
existing_team = LiteLLM_TeamTable(
team_id="parent-group",
team_alias="Parent Group",
members=[],
members_with_roles=[],
)
create_user_mock = mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists",
AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")),
)
_, final_members, _ = await _process_group_patch_operations(
patch_ops=patch_ops,
existing_team=existing_team,
prisma_client=_member_resolution_prisma(mocker, users=set(), teams=set()),
)
create_user_mock.assert_not_called()
assert final_members == set()
@pytest.mark.parametrize(
"team_metadata, expect_provisioned",
[
({SCIM_MANAGED_TEAM_METADATA_KEY: True}, False),
({SCIM_TEAM_DATA_METADATA_KEY: {"displayName": "Child.Apps"}}, False),
({}, True),
(None, True),
({SCIM_MANAGED_TEAM_METADATA_KEY: False}, True),
({SCIM_TEAM_DATA_METADATA_KEY: None}, True),
],
ids=[
"scim-managed",
"legacy-scim-data",
"admin-created",
"no-metadata",
"marker-unset",
"legacy-key-without-value",
],
)
@pytest.mark.asyncio
async def test_process_group_patch_team_match_needs_scim_provenance(
mocker, scim_upsert_user_enabled, team_metadata, expect_provisioned
):
"""A bare member id that names a team is only evidence of a nested group when the
identity provider is what wrote that team. Teams created here can share an id with
a real user, and skipping those members stops provisioning them entirely."""
patch_ops = SCIMPatchOp(
schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"],
Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "child-team"}])],
)
existing_team = LiteLLM_TeamTable(
team_id="parent-group",
team_alias="Parent Group",
members=[],
members_with_roles=[],
)
prisma_client = _member_resolution_prisma(mocker, users=set(), teams=set())
prisma_client.db.litellm_teamtable.find_unique = AsyncMock(
return_value=LiteLLM_TeamTable(team_id="child-team", metadata=team_metadata)
)
create_user_mock = mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists",
AsyncMock(return_value=NewUserResponse(user_id="child-team", key="new-key")),
)
_, final_members, _ = await _process_group_patch_operations(
patch_ops=patch_ops,
existing_team=existing_team,
prisma_client=prisma_client,
)
assert create_user_mock.called is expect_provisioned
assert final_members == ({"child-team"} if expect_provisioned else set())
@pytest.mark.asyncio
async def test_create_group_strict_mode_rejects_id_matching_admin_created_team(mocker, scim_upsert_user_disabled):
"""Strict mode drops nested groups but reports unknown users. A team an admin
created here says nothing about the member, so the member is an unknown user."""
scim_group = SCIMGroup(
schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"],
id="parent-group",
displayName="Parent Group",
members=[SCIMMember(value="admin-team")],
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
AsyncMock(
return_value=_member_resolution_prisma(
mocker, users=set(), teams=set(), unmanaged_teams=frozenset({"admin-team"})
)
),
)
with pytest.raises(ProxyException) as exc_info:
await create_group(group=scim_group)
assert int(exc_info.value.code) == 400
assert "admin-team" in str(exc_info.value.message)
@pytest.mark.asyncio
async def test_create_group_stamps_scim_provenance(mocker, scim_upsert_user_enabled):
"""The provenance the classifier reads only exists if the group writes stamp it;
a SCIM-created team that carries no mark looks admin-created forever after."""
scim_group = SCIMGroup(
schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"],
id="child-group",
displayName="Child.Apps",
members=[],
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
AsyncMock(return_value=_member_resolution_prisma(mocker, users=set(), teams=set())),
)
new_team_mock = mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2.new_team",
AsyncMock(return_value=mocker.MagicMock()),
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group",
AsyncMock(return_value=scim_group),
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._recompute_scim_member_roles",
AsyncMock(),
)
await create_group(group=scim_group)
assert new_team_mock.call_args.kwargs["data"].metadata == {SCIM_MANAGED_TEAM_METADATA_KEY: True}
@pytest.mark.asyncio
async def test_update_group_stamps_scim_provenance(mocker, scim_upsert_user_enabled):
"""A PUT full sync adopts a team the identity provider now owns, and the stamp has
to land alongside the existing metadata rather than replacing it."""
import json
group_id = "child-group"
existing_team = LiteLLM_TeamTable(
team_id=group_id,
team_alias="Child.Apps",
members=[],
members_with_roles=[],
metadata={"existing_key": "kept"},
)
scim_group = SCIMGroup(
schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"],
id=group_id,
displayName="Child.Apps",
members=[],
)
prisma_client = _member_resolution_prisma(mocker, users=set(), teams=set())
prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=existing_team)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
AsyncMock(return_value=prisma_client),
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._check_team_exists",
AsyncMock(return_value=existing_team),
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._recompute_scim_member_roles",
AsyncMock(),
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group",
AsyncMock(return_value=scim_group),
)
await update_group(group_id=group_id, group=scim_group)
written = json.loads(prisma_client.db.litellm_teamtable.update.call_args.kwargs["data"]["metadata"])
assert written[SCIM_MANAGED_TEAM_METADATA_KEY] is True
assert written["existing_key"] == "kept"
assert SCIM_TEAM_DATA_METADATA_KEY in written
@pytest.mark.asyncio
async def test_process_group_patch_stamps_scim_provenance(mocker, scim_upsert_user_enabled):
"""PATCH is how Okta adopts a group, so a membership-only patch has to stamp the
team too; otherwise the group it manages never gains provenance."""
patch_ops = SCIMPatchOp(
schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"],
Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "real-user"}])],
)
existing_team = LiteLLM_TeamTable(
team_id="parent-group",
team_alias="Parent Group",
members=[],
members_with_roles=[],
metadata={"existing_key": "kept"},
)
update_data, _, _ = await _process_group_patch_operations(
patch_ops=patch_ops,
existing_team=existing_team,
prisma_client=_member_resolution_prisma(mocker, users={"real-user"}, teams=set()),
)
assert update_data["metadata"][SCIM_MANAGED_TEAM_METADATA_KEY] is True
assert update_data["metadata"]["existing_key"] == "kept"
@pytest.mark.parametrize("member_type", ["direct", "Device"])
@pytest.mark.asyncio
async def test_process_group_patch_keeps_existing_user_with_unrecognized_type(
mocker, scim_upsert_user_enabled, member_type
):
"""Clients do stamp non-canonical types on real members (RFC 7643 defines
``direct`` for ``User.groups``). Dropping a member whose id is a live user would
revoke that user's team access on the next full sync, so the type is only
grounds for skipping once the user lookup has missed."""
patch_ops = SCIMPatchOp(
schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"],
Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "real-user", "type": member_type}])],
)
existing_team = LiteLLM_TeamTable(
team_id="parent-group",
team_alias="Parent Group",
members=[],
members_with_roles=[],
)
create_user_mock = mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists",
AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")),
)
_, final_members, _ = await _process_group_patch_operations(
patch_ops=patch_ops,
existing_team=existing_team,
prisma_client=_member_resolution_prisma(mocker, users={"real-user"}, teams=set()),
)
create_user_mock.assert_not_called()
assert final_members == {"real-user"}
@pytest.mark.parametrize(
"second_creation",
[None, NewUserResponse(user_id="dup-user", key="second-key")],
ids=["second-creation-fails", "both-creations-succeed"],
)
@pytest.mark.asyncio
async def test_resolve_group_member_ids_dedupes_repeated_member(mocker, scim_upsert_user_enabled, second_creation):
"""An id the request lists twice is one member. Admitting it twice writes a
duplicate members_with_roles row, and the second creation of the same id fails
against the real unique constraint even when the first one succeeded."""
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists",
AsyncMock(side_effect=[NewUserResponse(user_id="dup-user", key="first-key"), second_creation]),
)
result = await _resolve_group_member_ids(
members=[SCIMMember(value="dup-user"), SCIMMember(value="dup-user")],
created_via="scim_group_membership",
prisma_client=_member_resolution_prisma(mocker, users=set(), teams=set()),
)
assert result.all_member_ids == ["dup-user"]
@pytest.mark.parametrize(
"operation",
[
SCIMPatchOperation(op="add", path="members", value=[{"value": " "}]),
SCIMPatchOperation(op="remove", path="members", value=[{"value": " "}]),
SCIMPatchOperation(op="remove", path='members[value eq " "]', value=None),
],
ids=["add", "remove", "remove-path-filter"],
)
@pytest.mark.asyncio
async def test_process_group_patch_rejects_blank_member_id(mocker, scim_upsert_user_enabled, operation):
"""A blank id names nobody. The removal path stopped resolving its members, so it
has to keep rejecting one on its own."""
patch_ops = SCIMPatchOp(schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], Operations=[operation])
existing_team = LiteLLM_TeamTable(
team_id="parent-group",
team_alias="Parent Group",
members=[],
members_with_roles=[Member(user_id="keep-user", role="user")],
)
with pytest.raises(HTTPException) as exc_info:
await _process_group_patch_operations(
patch_ops=patch_ops,
existing_team=existing_team,
prisma_client=_member_resolution_prisma(mocker, users={"keep-user"}, teams=set()),
)
assert exc_info.value.status_code == 400
def test_scim_member_round_trips_type():
"""``type`` has to survive parsing; dropping it is what made a nested group
look like a user id."""
assert SCIMMember.model_validate({"value": "x", "type": "Group"}).type == "Group"
assert SCIMMember(value="x").type is None
@pytest.mark.parametrize("junk_type", [123, True, {}, [], 1.5])
def test_scim_member_treats_non_string_type_as_absent(junk_type):
"""Before ``type`` was a field, junk in it was parsed away; typing the field must
not start rejecting those requests, and both parsers have to agree it is typeless."""
assert SCIMMember.model_validate({"value": "x", "type": junk_type}).type is None
assert _parse_member_entries([{"value": "x", "type": junk_type}])[0].type is None
@pytest.mark.asyncio
async def test_get_groups_members_are_typed_as_users(mocker):
"""Group members we report back are always users, and saying so keeps the
response from emitting a null ``type``."""
team = LiteLLM_TeamTable(
team_id="team-1",
team_alias="Team One",
members=[],
members_with_roles=[Member(user_id="member-1", role="user")],
)
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.db = mocker.MagicMock()
mock_prisma_client.db.litellm_teamtable = mocker.MagicMock()
mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team])
mock_prisma_client.db.litellm_teamtable.count = AsyncMock(return_value=1)
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(
return_value=mocker.MagicMock(user_id="member-1", user_email="member-1@example.com")
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
AsyncMock(return_value=mock_prisma_client),
)
response = await get_groups(startIndex=1, count=10, filter=None)
assert [m.type for m in response.Resources[0].members] == ["User"]

View file

@ -650,14 +650,14 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys():
assert key_data.metrics.spend == 10.0
def _daily_user_spend_record(*, user_id, api_key, spend):
def _daily_user_spend_record(*, user_id, api_key, spend, model="gpt-4", model_group="gpt-4"):
"""A LiteLLM_DailyUserSpend row as the per-user breakdown reads it."""
return SimpleNamespace(
date="2024-01-01",
user_id=user_id,
api_key=api_key,
model="gpt-4",
model_group="gpt-4",
model=model,
model_group=model_group,
custom_llm_provider="openai",
mcp_namespaced_tool_name=None,
endpoint="/chat/completions",
@ -731,6 +731,64 @@ async def test_get_daily_activity_applies_resolve_entity_metadata_to_breakdown()
assert entities["user-no-email"].metadata == {}
@pytest.mark.asyncio
async def test_model_groups_breakdown_keys_by_public_name_with_model_fallback():
"""The usage UI labels model traffic with the model_groups breakdown.
Keys must be the requested public model name (model_group), and rows with a
NULL or empty model_group (pre-routing failures, rows written before the
column existed) must fall back to their model name instead of being dropped
from the breakdown. The models breakdown keeps the upstream litellm names.
"""
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
records = [
_daily_user_spend_record(
user_id="u1", api_key="key-1", spend=7.0, model="gpt-5.2", model_group="gpt-5.2-eu"
),
_daily_user_spend_record(
user_id="u1", api_key="key-1", spend=3.0, model="gpt-5.2", model_group=None
),
_daily_user_spend_record(
user_id="u1", api_key="key-1", spend=2.0, model="claude-x", model_group=""
),
]
mock_table = MagicMock()
mock_table.count = AsyncMock(return_value=len(records))
mock_table.find_many = AsyncMock(return_value=records)
mock_prisma.db.litellm_dailyuserspend = mock_table
mock_prisma.db.litellm_verificationtoken = MagicMock()
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
result = await get_daily_activity(
prisma_client=mock_prisma,
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id=None,
entity_metadata_field=None,
start_date="2024-01-01",
end_date="2024-01-01",
model=None,
api_key=None,
page=1,
page_size=1000,
)
breakdown = result.results[0].breakdown
assert set(breakdown.model_groups.keys()) == {"gpt-5.2-eu", "gpt-5.2", "claude-x"}
assert breakdown.model_groups["gpt-5.2-eu"].metrics.spend == 7.0
assert breakdown.model_groups["gpt-5.2"].metrics.spend == 3.0
assert breakdown.model_groups["claude-x"].metrics.spend == 2.0
assert breakdown.model_groups["gpt-5.2"].api_key_breakdown["key-1"].metrics.spend == 3.0
assert set(breakdown.models.keys()) == {"gpt-5.2", "claude-x"}
assert breakdown.models["gpt-5.2"].metrics.spend == 10.0
assert breakdown.models["claude-x"].metrics.spend == 2.0
class TestAdjustDatesForTimezone:
"""
Regression tests for the timezone double-counting bug.
@ -852,6 +910,38 @@ class TestBuildAggregatedSqlQuery:
assert "model = $4" in sql
assert "api_key = $5" in sql
def test_model_group_rollups_fall_back_to_model_name(self):
"""Aggregated model_groups rollups must fall back to model for group-less rows.
The (date, model_group) grouping level cannot recover the model column
after the fact (it is rolled up), so the fallback has to happen in SQL;
without it, group-less rows silently vanish from the model_groups
breakdown that the usage UI now renders by default. Group-less rows are
stored as empty strings, not NULL (spend_tracking_utils defaults
model_group to ""), so a plain COALESCE is not enough: the fallback must
be NULLIF-wrapped to catch both
"""
sql, _ = _build_aggregated_sql_query(
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id=None,
start_date="2026-07-01",
end_date="2026-07-01",
model=None,
api_key=None,
)
normalized = " ".join(sql.split())
fallback = "COALESCE(NULLIF(model_group, ''), model)"
assert f"{fallback} AS model_group" in normalized
assert (
f"GROUPING(date, api_key, model, {fallback}, "
"custom_llm_provider, mcp_namespaced_tool_name, endpoint) AS group_level" in normalized
)
assert f"(date, {fallback}), (date, {fallback}, api_key)," in normalized
assert "(date, model_group)" not in normalized
assert "COALESCE(model_group, model)" not in normalized
@pytest.mark.asyncio
async def test_get_daily_activity_aggregated_empty_result_set():

View file

@ -3667,3 +3667,38 @@ async def test_add_user_to_team_keeps_already_a_member_quiet(mocker, caplog):
)
assert [r.getMessage() for r in caplog.records if r.levelno >= logging.ERROR] == []
@pytest.mark.asyncio
async def test_get_user_info_for_proxy_admin_validates_keys_and_teams():
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.proxy._types import LiteLLM_TeamTable, UserAPIKeyAuth
from litellm.proxy.management_endpoints.internal_user_endpoints import (
_get_user_info_for_proxy_admin,
)
raw_rows = [
{
"teams": [
{"team_id": "team-b", "team_alias": "beta"},
{"team_id": "team-a", "team_alias": "alpha"},
],
"keys": [
{"token": "hashed-token-1", "team_id": "team-a", "models": None, "spend": 1.0},
],
}
]
mock_prisma_client = MagicMock()
mock_prisma_client.db.query_raw = AsyncMock(return_value=raw_rows)
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client):
result = await _get_user_info_for_proxy_admin(user_api_key_dict=UserAPIKeyAuth(user_id=None))
assert all(isinstance(team, LiteLLM_TeamTable) for team in result.teams)
assert [team.team_alias for team in result.teams] == ["alpha", "beta"]
assert len(result.keys) == 1
returned_key = result.keys[0]
assert returned_key["team_id"] == "team-a"
assert returned_key["models"] == []

View file

@ -2507,6 +2507,195 @@ async def test_update_key_nonexistent_key_returns_404(monkeypatch):
assert "Authentication Error" not in str(exc_info.value.message)
def _setup_update_key_mocks(monkeypatch, mock_prisma_client):
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", AsyncMock())
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock())
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
monkeypatch.setattr("litellm.store_audit_logs", False)
@pytest.mark.asyncio
async def test_update_key_by_alias_only(monkeypatch):
"""
/key/update identified by key_alias alone resolves the key row via
find_many on the alias and updates using the resolved token.
"""
from litellm.proxy.management_endpoints.key_management_endpoints import (
update_key_fn,
)
hashed_token = "0d62f396c1317066f55a96086517047c737087c61eb2bf016b72e6298927b15b"
key_in_db = LiteLLM_VerificationToken(
token=hashed_token,
key_alias="prod-alias",
user_id="test-user",
max_budget=200.0,
)
mock_prisma_client = AsyncMock()
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[key_in_db]
)
mock_prisma_client.db.litellm_verificationtoken.find_first = AsyncMock(
return_value=None
)
mock_prisma_client.update_data = AsyncMock(
return_value={"data": {"max_budget": 50.0}}
)
_setup_update_key_mocks(monkeypatch, mock_prisma_client)
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user"
)
request_data = UpdateKeyRequest(key_alias="prod-alias", max_budget=50.0)
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object"
) as mock_delete_cache:
mock_delete_cache.return_value = None
result = await update_key_fn(
request=MagicMock(),
data=request_data,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
mock_prisma_client.db.litellm_verificationtoken.find_many.assert_called_once_with(
where={"key_alias": "prod-alias"}, take=2
)
assert request_data.key == hashed_token
mock_prisma_client.db.litellm_verificationtoken.find_unique.assert_not_called()
mock_prisma_client.update_data.assert_awaited_once()
assert mock_prisma_client.update_data.call_args.kwargs["token"] == hashed_token
assert (
mock_prisma_client.update_data.call_args.kwargs["data"]["token"] == hashed_token
)
assert result["key"] == hashed_token
@pytest.mark.asyncio
async def test_update_key_by_alias_not_found_returns_404(monkeypatch):
"""
/key/update with a key_alias matching no key returns 404.
"""
from litellm.proxy.management_endpoints.key_management_endpoints import (
update_key_fn,
)
mock_prisma_client = AsyncMock()
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[]
)
_setup_update_key_mocks(monkeypatch, mock_prisma_client)
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user"
)
with pytest.raises(ProxyException) as exc_info:
await update_key_fn(
request=MagicMock(),
data=UpdateKeyRequest(key_alias="no-such-alias", max_budget=50.0),
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
assert str(exc_info.value.code) == "404"
assert "not found" in str(exc_info.value.message).lower()
mock_prisma_client.update_data.assert_not_called()
@pytest.mark.asyncio
async def test_update_key_by_duplicate_alias_returns_400(monkeypatch):
"""
/key/update with a key_alias shared by multiple keys returns 400
instead of silently updating one of them.
"""
from litellm.proxy.management_endpoints.key_management_endpoints import (
update_key_fn,
)
rows = [
LiteLLM_VerificationToken(token="hashed-token-1", key_alias="dup-alias"),
LiteLLM_VerificationToken(token="hashed-token-2", key_alias="dup-alias"),
]
mock_prisma_client = AsyncMock()
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=rows
)
_setup_update_key_mocks(monkeypatch, mock_prisma_client)
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user"
)
with pytest.raises(ProxyException) as exc_info:
await update_key_fn(
request=MagicMock(),
data=UpdateKeyRequest(key_alias="dup-alias", max_budget=50.0),
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
assert str(exc_info.value.code) == "400"
assert "multiple keys" in str(exc_info.value.message).lower()
mock_prisma_client.update_data.assert_not_called()
@pytest.mark.asyncio
async def test_update_key_with_key_and_alias_selects_by_key(monkeypatch):
"""
Regression: passing both key and key_alias keeps today's behavior. The key
identifies the row (find_unique, never find_many) and key_alias is the new
alias to set; the response echoes the caller-passed key.
"""
from litellm.proxy.management_endpoints.key_management_endpoints import (
update_key_fn,
)
hashed_token = "0d62f396c1317066f55a96086517047c737087c61eb2bf016b72e6298927b15b"
key_in_db = LiteLLM_VerificationToken(
token=hashed_token,
key_alias="old-name",
user_id="test-user",
)
mock_prisma_client = AsyncMock()
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=key_in_db
)
mock_prisma_client.db.litellm_verificationtoken.find_first = AsyncMock(
return_value=None
)
mock_prisma_client.update_data = AsyncMock(
return_value={"data": {"key_alias": "new-name"}}
)
_setup_update_key_mocks(monkeypatch, mock_prisma_client)
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user"
)
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object"
) as mock_delete_cache:
mock_delete_cache.return_value = None
result = await update_key_fn(
request=MagicMock(),
data=UpdateKeyRequest(key="sk-test-key", key_alias="new-name"),
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
mock_prisma_client.db.litellm_verificationtoken.find_many.assert_not_called()
mock_prisma_client.db.litellm_verificationtoken.find_unique.assert_called_once()
assert mock_prisma_client.update_data.call_args.kwargs["token"] == "sk-test-key"
assert result["key"] == "sk-test-key"
@pytest.mark.asyncio
async def test_block_key_existing_key_succeeds(monkeypatch):
"""

View file

@ -6578,6 +6578,7 @@ async def test_update_team_org_scoped_tpm_rpm_bypasses_user_limit():
return_value=mock_existing_team
)
mock_cache.async_set_cache = AsyncMock()
mock_logging.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
# Mock team update
mock_updated_team = MagicMock(spec=LiteLLM_TeamTable)
@ -10223,3 +10224,68 @@ def test_patch_team_route_publishes_its_request_body_schema():
assert schema == {"$ref": "#/components/schemas/PatchTeamRequest"}
properties = app.openapi()["components"]["schemas"]["PatchTeamRequest"]["properties"]
assert "tpm_limit" in properties and "metadata" in properties
@pytest.mark.asyncio
async def test_get_all_team_memberships_validates_rows():
from litellm.proxy._types import LiteLLM_TeamMembership
from litellm.proxy.management_endpoints.team_endpoints import (
get_all_team_memberships,
)
membership_row = MagicMock()
membership_row.model_dump = lambda: {
"user_id": "member-1",
"team_id": "team-1",
"spend": 2.5,
}
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_teammembership.find_many = AsyncMock(return_value=[membership_row])
result = await get_all_team_memberships(mock_prisma_client, ["team-1"], user_id="member-1")
assert len(result) == 1
assert isinstance(result[0], LiteLLM_TeamMembership)
assert result[0].user_id == "member-1"
assert result[0].team_id == "team-1"
assert result[0].spend == 2.5
find_many_kwargs = mock_prisma_client.db.litellm_teammembership.find_many.call_args.kwargs
assert find_many_kwargs["where"] == {"team_id": {"in": ["team-1"]}, "user_id": {"in": ["member-1"]}}
@pytest.mark.asyncio
async def test_list_available_teams_filters_joined_and_validates_rows(monkeypatch):
from fastapi import Request
import litellm
from litellm.proxy.management_endpoints.team_endpoints import list_available_teams
monkeypatch.setattr(
litellm,
"default_internal_user_params",
{"available_teams": ["team-open", "team-joined"]},
)
user_row = MagicMock()
user_row.model_dump = lambda: {"user_id": "u-1", "teams": ["team-joined"]}
open_team_row = MagicMock()
open_team_row.model_dump = lambda: {"team_id": "team-open", "team_alias": "open team"}
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row)
mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[open_team_row])
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client):
result = await list_available_teams(
http_request=MagicMock(spec=Request),
user_api_key_dict=UserAPIKeyAuth(user_id="u-1"),
)
assert len(result) == 1
assert isinstance(result[0], LiteLLM_TeamTable)
assert result[0].team_id == "team-open"
assert result[0].team_alias == "open team"
find_many_kwargs = mock_prisma_client.db.litellm_teamtable.find_many.call_args.kwargs
assert find_many_kwargs["where"] == {"team_id": {"in": ["team-open"]}}

View file

@ -2610,3 +2610,444 @@ def test_list_files_with_all_proxy_models_team_uses_openai_deployment(
assert captured_kwargs.get("api_key") == "team-openai-key"
assert captured_kwargs.get("custom_llm_provider") == "openai"
proxy_logging_obj.post_call_failure_hook.assert_not_called()
def _setup_vertex_named_credential_router(monkeypatch) -> Router:
from litellm.types.utils import CredentialItem
monkeypatch.setattr(
litellm,
"credential_list",
[
CredentialItem(
credential_name="vertex-named-cred",
credential_info={},
credential_values={
"vertex_project": "customer-project",
"vertex_location": "us-central1",
"vertex_credentials": "/creds/customer-sa.json",
},
)
],
)
return Router(
model_list=[
{
"model_name": "gemini-2.5-pro",
"litellm_params": {
"model": "vertex_ai/gemini-2.5-pro",
"litellm_credential_name": "vertex-named-cred",
},
}
]
)
def _assert_vertex_named_credentials_attached(captured_kwargs: dict) -> None:
assert captured_kwargs.get("custom_llm_provider") == "vertex_ai"
assert captured_kwargs.get("vertex_project") == "customer-project"
assert captured_kwargs.get("vertex_location") == "us-central1"
assert captured_kwargs.get("vertex_credentials") == "/creds/customer-sa.json"
assert captured_kwargs.get("model") is None
def test_create_file_provider_only_resolves_named_vertex_credentials(
mocker: MockerFixture, monkeypatch
):
"""
POST /v1/files with only a custom-llm-provider header (no model, no
target_model_names) must attach the configured named vertex credential to
the upstream call instead of falling through to google.auth.default(),
which uploads into the hosting environment's GCP project.
"""
import litellm.proxy.proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
router = _setup_vertex_named_credential_router(monkeypatch)
proxy_logging_obj = setup_proxy_logging_object(monkeypatch, router)
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router)
proxy_logging_obj.update_request_status = mocker.AsyncMock()
proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock()
captured_kwargs: dict = {}
async def _mock_acreate_file(**kwargs):
captured_kwargs.update(kwargs)
return OpenAIFileObject(
id="file-vertex-123",
object="file",
bytes=2,
created_at=1234567890,
filename="batch.jsonl",
purpose="batch",
status="uploaded",
)
monkeypatch.setattr(litellm, "acreate_file", _mock_acreate_file)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
api_key="test-key",
user_role=LitellmUserRoles.PROXY_ADMIN,
user_id="test-user",
)
try:
response = client.post(
"/v1/files",
files={"file": ("batch.jsonl", b"{}", "application/jsonl")},
data={"purpose": "batch"},
headers={
"Authorization": "Bearer test-key",
"custom-llm-provider": "vertex_ai",
},
)
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
assert response.status_code == 200, response.text
_assert_vertex_named_credentials_attached(captured_kwargs)
proxy_logging_obj.post_call_failure_hook.assert_not_called()
def test_get_file_provider_only_resolves_named_vertex_credentials(
mocker: MockerFixture, monkeypatch
):
import litellm.proxy.proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
router = _setup_vertex_named_credential_router(monkeypatch)
proxy_logging_obj = setup_proxy_logging_object(monkeypatch, router)
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router)
proxy_logging_obj.update_request_status = mocker.AsyncMock()
proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock()
captured_kwargs: dict = {}
async def _mock_afile_retrieve(**kwargs):
captured_kwargs.update(kwargs)
return OpenAIFileObject(
id="file-abc123",
object="file",
bytes=2,
created_at=1234567890,
filename="batch.jsonl",
purpose="batch",
status="uploaded",
)
monkeypatch.setattr(litellm, "afile_retrieve", _mock_afile_retrieve)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
api_key="test-key",
user_role=LitellmUserRoles.PROXY_ADMIN,
user_id="test-user",
)
try:
response = client.get(
"/v1/files/file-abc123",
headers={
"Authorization": "Bearer test-key",
"custom-llm-provider": "vertex_ai",
},
)
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
assert response.status_code == 200, response.text
assert captured_kwargs.get("file_id") == "file-abc123"
_assert_vertex_named_credentials_attached(captured_kwargs)
proxy_logging_obj.post_call_failure_hook.assert_not_called()
def test_get_file_content_provider_only_resolves_named_vertex_credentials(
mocker: MockerFixture, monkeypatch
):
import litellm.proxy.proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
router = _setup_vertex_named_credential_router(monkeypatch)
proxy_logging_obj = setup_proxy_logging_object(monkeypatch, router)
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router)
proxy_logging_obj.update_request_status = mocker.AsyncMock()
proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock()
captured_kwargs: dict = {}
async def _mock_afile_content(**kwargs):
captured_kwargs.update(kwargs)
return HttpxBinaryResponseContent(
response=httpx.Response(
status_code=200,
content=b"vertex-bytes",
headers={"content-type": "application/octet-stream"},
)
)
monkeypatch.setattr(litellm, "afile_content", _mock_afile_content)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
api_key="test-key",
user_role=LitellmUserRoles.PROXY_ADMIN,
user_id="test-user",
)
try:
response = client.get(
"/v1/files/file-abc123/content",
headers={
"Authorization": "Bearer test-key",
"custom-llm-provider": "vertex_ai",
},
)
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
assert response.status_code == 200, response.text
assert response.content == b"vertex-bytes"
assert captured_kwargs.get("file_id") == "file-abc123"
_assert_vertex_named_credentials_attached(captured_kwargs)
proxy_logging_obj.post_call_failure_hook.assert_not_called()
def test_delete_file_provider_only_resolves_named_vertex_credentials(
mocker: MockerFixture, monkeypatch
):
import litellm.proxy.proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
router = _setup_vertex_named_credential_router(monkeypatch)
proxy_logging_obj = setup_proxy_logging_object(monkeypatch, router)
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router)
proxy_logging_obj.update_request_status = mocker.AsyncMock()
proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock()
captured_kwargs: dict = {}
async def _mock_afile_delete(**kwargs):
captured_kwargs.update(kwargs)
return OpenAIFileObject(
id="file-abc123",
object="file",
bytes=2,
created_at=1234567890,
filename="batch.jsonl",
purpose="batch",
status="uploaded",
)
monkeypatch.setattr(litellm, "afile_delete", _mock_afile_delete)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
api_key="test-key",
user_role=LitellmUserRoles.PROXY_ADMIN,
user_id="test-user",
)
try:
response = client.delete(
"/v1/files/file-abc123",
headers={
"Authorization": "Bearer test-key",
"custom-llm-provider": "vertex_ai",
},
)
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
assert response.status_code == 200, response.text
assert captured_kwargs.get("file_id") == "file-abc123"
_assert_vertex_named_credentials_attached(captured_kwargs)
proxy_logging_obj.post_call_failure_hook.assert_not_called()
def test_create_file_provider_only_skips_other_team_vertex_deployment(
mocker: MockerFixture, monkeypatch
):
"""
Regression: with a team-scoped vertex deployment indexed before a global
one under the same model name, a provider-only upload from a different
team must use the global deployment's credentials, never the other
team's.
"""
import litellm.proxy.proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
router = Router(
model_list=[
{
"model_name": "gemini-2.5-pro",
"litellm_params": {
"model": "vertex_ai/gemini-2.5-pro",
"vertex_project": "team-b-project",
},
"model_info": {
"id": "team-b-vertex",
"team_id": "team-b",
"team_public_model_name": "gemini-2.5-pro",
},
},
{
"model_name": "gemini-2.5-pro",
"litellm_params": {
"model": "vertex_ai/gemini-2.5-pro",
"vertex_project": "shared-project",
},
},
]
)
proxy_logging_obj = setup_proxy_logging_object(monkeypatch, router)
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router)
proxy_logging_obj.update_request_status = mocker.AsyncMock()
proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock()
captured_kwargs: dict = {}
async def _mock_acreate_file(**kwargs):
captured_kwargs.update(kwargs)
return OpenAIFileObject(
id="file-vertex-456",
object="file",
bytes=2,
created_at=1234567890,
filename="batch.jsonl",
purpose="batch",
status="uploaded",
)
monkeypatch.setattr(litellm, "acreate_file", _mock_acreate_file)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
api_key="test-key",
user_role=LitellmUserRoles.INTERNAL_USER,
user_id="test-user",
team_id="team-a",
team_models=["gemini-2.5-pro"],
)
try:
response = client.post(
"/v1/files",
files={"file": ("batch.jsonl", b"{}", "application/jsonl")},
data={"purpose": "batch"},
headers={
"Authorization": "Bearer test-key",
"custom-llm-provider": "vertex_ai",
},
)
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
assert response.status_code == 200, response.text
assert captured_kwargs.get("vertex_project") == "shared-project"
proxy_logging_obj.post_call_failure_hook.assert_not_called()
def _team_openai_plus_global_anthropic_router() -> Router:
return Router(
model_list=[
{
"model_name": "team-gpt",
"litellm_params": {
"model": "openai/gpt-4o",
"api_key": "team-openai-key",
},
"model_info": {
"id": "team-a-openai",
"team_id": "team-a",
"team_public_model_name": "team-gpt",
},
},
{
"model_name": "claude-opus-4-6",
"litellm_params": {
"model": "anthropic/claude-opus-4-6",
"api_key": "anthropic-key",
},
},
]
)
def _list_files_captured_kwargs(
mocker: MockerFixture, monkeypatch, router: Router, key_models: list
) -> dict:
import litellm.proxy.proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
proxy_logging_obj = setup_proxy_logging_object(monkeypatch, router)
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router)
proxy_logging_obj.update_request_status = mocker.AsyncMock()
proxy_logging_obj.post_call_success_hook = mocker.AsyncMock(return_value=[])
proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock()
captured_kwargs: dict = {}
async def _mock_afile_list(**kwargs):
captured_kwargs.update(kwargs)
return []
monkeypatch.setattr(litellm, "afile_list", _mock_afile_list)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
api_key="test-key",
user_role=LitellmUserRoles.INTERNAL_USER,
user_id="test-user",
team_id="team-a",
team_models=["team-gpt", "claude-opus-4-6"],
models=key_models,
)
try:
response = client.get(
"/v1/files",
headers={
"Authorization": "Bearer test-key",
"custom-llm-provider": "openai",
},
)
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
assert response.status_code == 200, response.text
return captured_kwargs
def test_list_files_key_restricted_to_other_provider_does_not_leak_team_openai_credentials(
mocker: MockerFixture, monkeypatch
):
"""
Regression: a key restricted to an anthropic model on a team that also has
an openai deployment must not attach the team's openai credentials to a
provider-only openai files call; key-level model restrictions apply to
credential resolution, not just completions.
"""
captured_kwargs = _list_files_captured_kwargs(
mocker, monkeypatch, _team_openai_plus_global_anthropic_router(), ["claude-opus-4-6"]
)
assert captured_kwargs.get("api_key") != "team-openai-key"
def test_list_files_key_allowed_openai_model_still_resolves_team_credentials(
mocker: MockerFixture, monkeypatch
):
"""
A key whose allowlist includes the team's openai model keeps resolving that
deployment's credentials for provider-only openai files calls.
"""
captured_kwargs = _list_files_captured_kwargs(
mocker, monkeypatch, _team_openai_plus_global_anthropic_router(), ["team-gpt"]
)
assert captured_kwargs.get("api_key") == "team-openai-key"

View file

@ -1548,6 +1548,7 @@ async def test_ui_view_session_spend_logs_pagination(client, monkeypatch):
assert page_size == 1
assert skip == 1 # page=2, page_size=1
assert 'ORDER BY "startTime" DESC' in sql_query
assert '"user" = $4' not in sql_query
return [mock_spend_logs[0]]
class MockPrismaClient:
@ -1558,20 +1559,144 @@ async def test_ui_view_session_spend_logs_pagination(client, monkeypatch):
mock_prisma_client = MockPrismaClient()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
response = client.get(
"/spend/logs/session/ui",
params={"session_id": "session-123", "page": 2, "page_size": 1},
headers={"Authorization": "Bearer sk-test"},
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user"
)
assert response.status_code == 200
data = response.json()
assert data["total"] == 2
assert data["page"] == 2
assert data["page_size"] == 1
assert data["total_pages"] == 2
assert len(data["data"]) == 1
assert data["data"][0]["request_id"] == "req1"
try:
response = client.get(
"/spend/logs/session/ui",
params={"session_id": "session-123", "page": 2, "page_size": 1},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 200
data = response.json()
assert data["total"] == 2
assert data["page"] == 2
assert data["page_size"] == 1
assert data["total_pages"] == 2
assert len(data["data"]) == 1
assert data["data"][0]["request_id"] == "req1"
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_ui_view_session_spend_logs_scopes_non_admin_to_own_logs(client, monkeypatch):
own_log = {
"id": "log1",
"request_id": "req1",
"session_id": "session-123",
"user": "user-1",
"startTime": "2024-01-01T00:00:00Z",
}
class MockDB:
async def count(self, *args, **kwargs):
assert kwargs.get("where") == {"session_id": "session-123", "user": "user-1"}
return 1
async def query_raw(self, sql_query, session_id, page_size, skip, scoped_user):
assert session_id == "session-123"
assert scoped_user == "user-1"
assert '"user" = $4' in sql_query
return [own_log]
class MockPrismaClient:
def __init__(self):
self.db = MockDB()
self.db.litellm_spendlogs = self.db
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient())
async def no_permitted_teams(*args, **kwargs):
return []
monkeypatch.setattr(
"litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs",
no_permitted_teams,
)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER, user_id="user-1"
)
try:
response = client.get(
"/spend/logs/session/ui",
params={"session_id": "session-123", "page": 1, "page_size": 50},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 200
data = response.json()
assert data["total"] == 1
assert [row["request_id"] for row in data["data"]] == ["req1"]
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_ui_view_session_spend_logs_includes_permitted_team_logs(client, monkeypatch):
class MockDB:
async def count(self, *args, **kwargs):
assert kwargs.get("where") == {
"session_id": "session-123",
"OR": [
{"user": "user-1"},
{"team_id": {"in": ["team-9"]}},
],
}
return 1
async def query_raw(self, sql_query, session_id, page_size, skip, scoped_user, team_ids):
assert session_id == "session-123"
assert scoped_user == "user-1"
assert team_ids == ["team-9"]
assert '("user" = $4 OR team_id = ANY($5::text[]))' in sql_query
return [
{
"id": "log2",
"request_id": "req2",
"session_id": "session-123",
"team_id": "team-9",
"startTime": "2024-01-02T00:00:00Z",
}
]
class MockPrismaClient:
def __init__(self):
self.db = MockDB()
self.db.litellm_spendlogs = self.db
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient())
async def permitted_teams(*args, **kwargs):
return ["team-9"]
monkeypatch.setattr(
"litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs",
permitted_teams,
)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER, user_id="user-1"
)
try:
response = client.get(
"/spend/logs/session/ui",
params={"session_id": "session-123", "page": 1, "page_size": 50},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 200
data = response.json()
assert data["total"] == 1
assert [row["request_id"] for row in data["data"]] == ["req2"]
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio

View file

@ -1571,14 +1571,14 @@ async def test_get_all_team_models():
with patch("litellm.proxy.proxy_server.LiteLLM_TeamTable") as mock_team_table_class:
# Configure the mock class to return proper instances
def mock_team_table_constructor(**kwargs):
def mock_team_table_constructor(data):
mock_instance = MagicMock()
mock_instance.team_id = kwargs["team_id"]
mock_instance.models = kwargs["models"]
mock_instance.access_group_ids = kwargs.get("access_group_ids")
mock_instance.team_id = data["team_id"]
mock_instance.models = data["models"]
mock_instance.access_group_ids = data.get("access_group_ids")
return mock_instance
mock_team_table_class.side_effect = mock_team_table_constructor
mock_team_table_class.model_validate.side_effect = mock_team_table_constructor
result = await get_all_team_models(
user_teams="*",
@ -1607,7 +1607,7 @@ async def test_get_all_team_models():
mock_litellm_teamtable.find_many.return_value = [mock_team1]
with patch("litellm.proxy.proxy_server.LiteLLM_TeamTable") as mock_team_table_class:
mock_team_table_class.side_effect = mock_team_table_constructor
mock_team_table_class.model_validate.side_effect = mock_team_table_constructor
result = await get_all_team_models(
user_teams=["team1"],
@ -1658,7 +1658,7 @@ async def test_get_all_team_models():
mock_router.get_model_list.side_effect = mock_get_model_list_with_none
with patch("litellm.proxy.proxy_server.LiteLLM_TeamTable") as mock_team_table_class:
mock_team_table_class.side_effect = mock_team_table_constructor
mock_team_table_class.model_validate.side_effect = mock_team_table_constructor
result = await get_all_team_models(
user_teams=["team1"],
@ -2373,14 +2373,14 @@ async def test_get_all_team_models_with_access_groups():
with patch("litellm.proxy.proxy_server.LiteLLM_TeamTable") as mock_tt_class:
def mock_team_table_constructor(**kwargs):
def mock_team_table_constructor(data):
mock_instance = MagicMock()
mock_instance.team_id = kwargs["team_id"]
mock_instance.models = kwargs["models"]
mock_instance.access_group_ids = kwargs.get("access_group_ids")
mock_instance.team_id = data["team_id"]
mock_instance.models = data["models"]
mock_instance.access_group_ids = data.get("access_group_ids")
return mock_instance
mock_tt_class.side_effect = mock_team_table_constructor
mock_tt_class.model_validate.side_effect = mock_team_table_constructor
result = await get_all_team_models(
user_teams=["team1"],

View file

@ -139,3 +139,23 @@ def test_key_request_router_settings_keeps_enable_tag_filtering():
dumped = req.router_settings.model_dump(exclude_none=True)
assert dumped["enable_tag_filtering"] is True
assert dumped["num_retries"] == 2
def test_update_key_request_requires_key_or_key_alias():
"""``/key/update`` can be addressed by ``key`` or by ``key_alias``;
a request with neither has no way to identify the target key and must
fail validation before hitting the endpoint."""
import pydantic
from litellm.proxy._types import UpdateKeyRequest
with pytest.raises(pydantic.ValidationError, match="either key or key_alias must be provided"):
UpdateKeyRequest(max_budget=10.0)
by_key = UpdateKeyRequest(key="sk-1234")
assert by_key.key == "sk-1234"
assert by_key.key_alias is None
by_alias = UpdateKeyRequest(key_alias="my-alias")
assert by_alias.key is None
assert by_alias.key_alias == "my-alias"

View file

@ -1014,9 +1014,9 @@ def test_tiered_pricing_only_deployment_selects_router_model_id():
router = Router(
model_list=[
{
"model_name": "qwen-3.7-plus",
"model_name": "qwen-tier-only",
"litellm_params": {
"model": "dashscope/qwen3.7-plus",
"model": "dashscope/qwen-tier-only-test",
"api_key": "sk-fake",
},
"model_info": {
@ -1037,10 +1037,12 @@ def test_tiered_pricing_only_deployment_selects_router_model_id():
assert entry.get("input_cost_per_token") is None
assert entry.get("tiered_pricing") is not None
# The stripped shared alias must not carry tiered pricing.
assert litellm.model_cost["dashscope/qwen3.7-plus"].get("tiered_pricing") is None
assert (
litellm.model_cost["dashscope/qwen-tier-only-test"].get("tiered_pricing") is None
)
selected = _select_model_name_for_cost_calc(
model="dashscope/qwen3.7-plus",
model="dashscope/qwen-tier-only-test",
completion_response=None,
custom_pricing=True,
custom_llm_provider="dashscope",

View file

@ -3755,6 +3755,182 @@ def test_get_deployment_credentials_with_provider_team_wildcard_priority():
assert global_credentials["api_key"] == "global-key"
def test_get_deployment_credentials_with_provider_skips_other_team_deployment():
"""
Regression: a team-scoped deployment sharing a model_name with a global
deployment must never resolve for another team's (or an unscoped) caller,
even when it is indexed first; the shared global deployment wins instead.
"""
router = litellm.Router(
model_list=[
{
"model_name": "gemini-2.5-pro",
"litellm_params": {
"model": "vertex_ai/gemini-2.5-pro",
"vertex_project": "team-b-project",
},
"model_info": {
"id": "team-b-vertex",
"team_id": "team-b",
"team_public_model_name": "gemini-2.5-pro",
},
},
{
"model_name": "gemini-2.5-pro",
"litellm_params": {
"model": "vertex_ai/gemini-2.5-pro",
"vertex_project": "shared-project",
},
},
],
)
other_team_credentials = router.get_deployment_credentials_with_provider(
model_id="gemini-2.5-pro", team_id="team-a"
)
assert other_team_credentials is not None
assert other_team_credentials["vertex_project"] == "shared-project"
unscoped_credentials = router.get_deployment_credentials_with_provider(
model_id="gemini-2.5-pro"
)
assert unscoped_credentials is not None
assert unscoped_credentials["vertex_project"] == "shared-project"
owner_credentials = router.get_deployment_credentials_with_provider(
model_id="gemini-2.5-pro", team_id="team-b"
)
assert owner_credentials is not None
assert owner_credentials["vertex_project"] == "team-b-project"
def test_get_deployment_credentials_with_provider_no_fallback_to_other_team_only_name():
"""
When the only deployments under a model name belong to another team, other
callers must get None (env fallback) instead of that team's credentials.
"""
router = litellm.Router(
model_list=[
{
"model_name": "gemini-2.5-pro",
"litellm_params": {
"model": "vertex_ai/gemini-2.5-pro",
"vertex_project": "team-b-project",
},
"model_info": {
"id": "team-b-vertex",
"team_id": "team-b",
"team_public_model_name": "gemini-2.5-pro",
},
},
],
)
assert (
router.get_deployment_credentials_with_provider(
model_id="gemini-2.5-pro", team_id="team-a"
)
is None
)
assert (
router.get_deployment_credentials_with_provider(model_id="gemini-2.5-pro")
is None
)
def test_deployment_usable_by_team_helpers():
"""
Direct coverage of the team-ownership filter: a team-scoped deployment is
usable only by its owning team, shared deployments by anyone, and the
model-group picker returns the first usable deployment or None.
"""
router = litellm.Router(
model_list=[
{
"model_name": "gemini-2.5-pro",
"litellm_params": {
"model": "vertex_ai/gemini-2.5-pro",
"vertex_project": "team-b-project",
},
"model_info": {
"id": "team-b-vertex",
"team_id": "team-b",
"team_public_model_name": "gemini-2.5-pro",
},
},
{
"model_name": "gemini-2.5-pro",
"litellm_params": {
"model": "vertex_ai/gemini-2.5-pro",
"vertex_project": "shared-project",
},
},
],
)
team_owned, shared = router.model_list
assert router._deployment_usable_by_team(team_owned, "team-b") is True
assert router._deployment_usable_by_team(team_owned, "team-a") is False
assert router._deployment_usable_by_team(team_owned, None) is False
assert router._deployment_usable_by_team(shared, "team-a") is True
assert router._deployment_usable_by_team(shared, None) is True
picked = router._get_model_group_deployment_usable_by_team(
model_group_name="gemini-2.5-pro", team_id="team-a"
)
assert picked is not None
assert picked.litellm_params.vertex_project == "shared-project"
owner_picked = router._get_model_group_deployment_usable_by_team(
model_group_name="gemini-2.5-pro", team_id="team-b"
)
assert owner_picked is not None
assert owner_picked.litellm_params.vertex_project == "team-b-project"
assert (
router._get_model_group_deployment_usable_by_team(
model_group_name="unknown-model", team_id="team-a"
)
is None
)
def test_get_deployment_credentials_with_provider_skips_other_team_wildcard():
"""
Global wildcard resolution must skip a team-scoped wildcard deployment for
callers outside that team, falling through to the shared wildcard entry.
"""
router = litellm.Router(
model_list=[
{
"model_name": "openai/*",
"litellm_params": {"model": "openai/*", "api_key": "team-b-key"},
"model_info": {
"id": "team-b-wildcard",
"team_id": "team-b",
"team_public_model_name": "openai/*",
},
},
{
"model_name": "openai/*",
"litellm_params": {"model": "openai/*", "api_key": "global-key"},
},
],
)
other_team_credentials = router.get_deployment_credentials_with_provider(
model_id="openai/gpt-5.2", team_id="team-a"
)
assert other_team_credentials is not None
assert other_team_credentials["api_key"] == "global-key"
owner_credentials = router.get_deployment_credentials_with_provider(
model_id="openai/gpt-5.2", team_id="team-b"
)
assert owner_credentials is not None
assert owner_credentials["api_key"] == "team-b-key"
def test_team_wildcard_credentials_not_usable_after_delete_deployment():
"""
Regression: team_pattern_routers retained deleted deployments, so a team

View file

@ -152,7 +152,7 @@
"count": 1
},
"prefer-const": {
"count": 3
"count": 1
},
"react-hooks/purity": {
"count": 1

View file

@ -4,36 +4,50 @@ import { screen, waitFor, within } from "@testing-library/react";
import { renderWithProviders } from "../../../../../tests/test-utils";
import CacheDashboard from "./cache_dashboard";
const { adminGlobalCacheActivity, cachingHealthCheckCall } = vi.hoisted(() => ({
adminGlobalCacheActivity: vi.fn(),
const { useCacheActivity, cachingHealthCheckCall } = vi.hoisted(() => ({
useCacheActivity: vi.fn(),
cachingHealthCheckCall: vi.fn(),
}));
vi.mock("@/components/networking", () => ({
adminGlobalCacheActivity,
cachingHealthCheckCall,
}));
const cacheActivity = [
{
api_key: "sk-1",
model: "gpt-5.1",
call_type: "acompletion",
total_rows: 1500,
cache_hit_true_rows: 300,
cached_completion_tokens: 12000,
generated_completion_tokens: 48000,
vi.mock("@/app/(dashboard)/hooks/caching/useCacheActivity", () => ({
useCacheActivity,
}));
const cacheActivity = {
groups: [
{
call_type: "acompletion",
api_requests: 1000,
cache_hits: 300,
failed_requests: 200,
cached_completion_tokens: 12000,
generated_completion_tokens: 48000,
},
{
call_type: "aembedding",
api_requests: 550,
cache_hits: 100,
failed_requests: 50,
cached_completion_tokens: 2000,
generated_completion_tokens: 9000,
},
],
totals: {
api_requests: 1550,
cache_hits: 400,
failed_requests: 250,
cached_completion_tokens: 14000,
cache_hit_ratio: (400 / 2200) * 100,
},
{
api_key: "sk-2",
model: "text-embedding-3-large",
call_type: "aembedding",
total_rows: 700,
cache_hit_true_rows: 100,
cached_completion_tokens: 2000,
generated_completion_tokens: 9000,
filter_options: {
key_aliases: ["my-key", "Unnamed Key"],
models: ["gpt-5.1", "text-embedding-3-large"],
},
];
};
const renderDashboard = () =>
renderWithProviders(
@ -75,7 +89,7 @@ const legendFillByCategory = (card: HTMLElement) =>
describe("CacheDashboard cache analytics charts", () => {
beforeEach(() => {
vi.clearAllMocks();
adminGlobalCacheActivity.mockResolvedValue(cacheActivity);
useCacheActivity.mockReturnValue({ data: cacheActivity, refetch: vi.fn() });
});
it("renders both chart card titles", async () => {
@ -108,8 +122,13 @@ describe("CacheDashboard cache analytics charts", () => {
expect(legendFillByCategory(requestsCard)).toEqual({
"LLM API requests": "var(--color-sky-500, #0ea5e9)",
"Cache hit": "var(--color-teal-500, #14b8a6)",
"Failed requests": "var(--color-red-500, #ef4444)",
});
expect(barFills(requestsCard)).toEqual(["var(--color-sky-500, #0ea5e9)", "var(--color-teal-500, #14b8a6)"]);
expect(barFills(requestsCard)).toEqual([
"var(--color-sky-500, #0ea5e9)",
"var(--color-teal-500, #14b8a6)",
"var(--color-red-500, #ef4444)",
]);
});
it("renders the tokens chart with each category legend-bound to its fill and stacked in order", async () => {
@ -133,18 +152,39 @@ describe("CacheDashboard cache analytics charts", () => {
}
});
it("stacks the two categories into one column per call_type", async () => {
it("stacks all categories into one column per call_type", async () => {
renderDashboard();
const { requestsCard, tokensCard } = await findChartCards();
for (const card of [requestsCard, tokensCard]) {
const expectedRects = { requests: 6, tokens: 4 };
for (const [card, rectCount] of [
[requestsCard, expectedRects.requests],
[tokensCard, expectedRects.tokens],
] as const) {
const rects = Array.from(card.querySelectorAll("path.recharts-rectangle"));
expect(rects).toHaveLength(4);
expect(rects).toHaveLength(rectCount);
const xPositions = rects.map((rect) => rect.getAttribute("d")?.split(",")[0]);
expect(new Set(xPositions).size).toBe(2);
}
});
it("renders the server-computed cache hit ratio", async () => {
renderDashboard();
expect(await screen.findByText("18.18%")).toBeInTheDocument();
});
it("passes the date range and selected filters to the activity query", () => {
renderDashboard();
expect(useCacheActivity).toHaveBeenCalledWith({
startDate: expect.stringMatching(/^\d{4}-\d{2}-\d{2}$/),
endDate: expect.stringMatching(/^\d{4}-\d{2}-\d{2}$/),
keyAliases: [],
models: [],
});
});
it("formats y-axis ticks with compact notation", async () => {
renderDashboard();
const { requestsCard, tokensCard } = await findChartCards();

View file

@ -19,13 +19,29 @@ import {
import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs";
import { RefreshCw } from "lucide-react";
import { adminGlobalCacheActivity, cachingHealthCheckCall } from "@/components/networking";
import { cachingHealthCheckCall } from "@/components/networking";
import { useCacheActivity, type CacheActivityGroup } from "@/app/(dashboard)/hooks/caching/useCacheActivity";
// Import the new component
import { CacheHealthTab } from "./cache_health";
import CacheSettings from "./cache_settings";
import CoordinationRedisSettings from "./coordination_redis_settings";
const REQUEST_SERIES = {
apiRequests: "LLM API requests",
cacheHits: "Cache hit",
failed: "Failed requests",
} as const;
const toChartDatum = (group: CacheActivityGroup) => ({
name: group.call_type,
[REQUEST_SERIES.apiRequests]: group.api_requests,
[REQUEST_SERIES.cacheHits]: group.cache_hits,
[REQUEST_SERIES.failed]: group.failed_requests,
"Cached Completion Tokens": group.cached_completion_tokens,
"Generated Completion Tokens": group.generated_completion_tokens,
});
const formatDateWithoutTZ = (date: Date | undefined) => {
if (!date) return undefined;
return date.toISOString().split("T")[0];
@ -49,26 +65,6 @@ interface CachePageProps {
premiumUser: boolean;
}
interface cacheDataItem {
api_key: string;
model: string;
cache_hit_true_rows: number;
cached_completion_tokens: number;
total_rows: number;
generated_completion_tokens: number;
call_type: string;
// Add other properties as needed
}
type uiData = {
name: string;
"LLM API requests": number;
"Cache hit": number;
"Cached Completion Tokens": number;
"Generated Completion Tokens": number;
};
interface CacheHealthResponse {
status?: string;
cache_type?: string;
@ -97,13 +93,8 @@ const deepParse = (input: any) => {
};
const CacheDashboard: React.FC<CachePageProps> = ({ accessToken, token, userRole, userID, premiumUser }) => {
const [filteredData, setFilteredData] = useState<uiData[]>([]);
const [selectedApiKeys, setSelectedApiKeys] = useState<string[]>([]);
const [selectedModels, setSelectedModels] = useState<string[]>([]);
const [data, setData] = useState<cacheDataItem[]>([]);
const [cachedResponses, setCachedResponses] = useState("0");
const [cachedTokens, setCachedTokens] = useState("0");
const [cacheHitRatio, setCacheHitRatio] = useState("0");
const [dateValue, setDateValue] = useState<DateRangePickerValue>({
from: new Date(Date.now() - 7 * 24 * 60 * 60 * 1000),
@ -113,120 +104,24 @@ const CacheDashboard: React.FC<CachePageProps> = ({ accessToken, token, userRole
const [lastRefreshed, setLastRefreshed] = useState("");
const [healthCheckResponse, setHealthCheckResponse] = useState<any>("");
useEffect(() => {
if (!accessToken || !dateValue) {
return;
}
const fetchData = async () => {
const response = await adminGlobalCacheActivity(
accessToken,
formatDateWithoutTZ(dateValue.from),
formatDateWithoutTZ(dateValue.to),
);
setData(response);
};
fetchData();
const currentDate = new Date();
setLastRefreshed(currentDate.toLocaleString());
}, [accessToken]);
const uniqueApiKeys = Array.from(new Set(data.map((item) => item?.api_key ?? "")));
const uniqueModels = Array.from(new Set(data.map((item) => item?.model ?? "")));
const uniqueCallTypes = Array.from(new Set(data.map((item) => item?.call_type ?? "")));
const updateCachingData = async (startTime: Date | undefined, endTime: Date | undefined) => {
if (!startTime || !endTime || !accessToken) {
return;
}
let new_cache_data = await adminGlobalCacheActivity(
accessToken,
formatDateWithoutTZ(startTime),
formatDateWithoutTZ(endTime),
);
setData(new_cache_data);
};
const { data: activity, refetch } = useCacheActivity({
startDate: formatDateWithoutTZ(dateValue.from),
endDate: formatDateWithoutTZ(dateValue.to),
keyAliases: selectedApiKeys,
models: selectedModels,
});
useEffect(() => {
let newData: cacheDataItem[] = data;
if (selectedApiKeys.length > 0) {
newData = newData.filter((item) => selectedApiKeys.includes(item.api_key));
}
setLastRefreshed(new Date().toLocaleString());
}, []);
if (selectedModels.length > 0) {
newData = newData.filter((item) => selectedModels.includes(item.model));
}
/*
Data looks like this
[{"api_key":"sk-test-mock-key-001","call_type":"acompletion","model":"llama3-8b-8192","total_rows":13,"cache_hit_true_rows":0},
{"api_key":"sk-test-mock-key-002","call_type":"None","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0},
{"api_key":"sk-test-mock-key-123","call_type":"acompletion","model":"gpt-3.5-turbo","total_rows":19,"cache_hit_true_rows":0},
{"api_key":"sk-test-mock-key-123","call_type":"aimage_generation","model":"","total_rows":3,"cache_hit_true_rows":0},
{"api_key":"sk-test-mock-key-003","call_type":"None","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0},
{"api_key":"sk-test-mock-key-004","call_type":"","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0},
{"api_key":"sk-test-mock-key-005","call_type":"","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0},
*/
// What data we need for bar chat
// ui_data = [
// {
// name: "Call Type",
// Cache hit: 20,
// LLM API requests: 10,
// }
// ]
let llm_api_requests = 0;
let cache_hits = 0;
let cached_tokens = 0;
const processedData = newData.reduce((acc: uiData[], item) => {
if (!item.call_type) {
item.call_type = "Unknown";
}
llm_api_requests += (item.total_rows || 0) - (item.cache_hit_true_rows || 0);
cache_hits += item.cache_hit_true_rows || 0;
cached_tokens += item.cached_completion_tokens || 0;
const existingItem = acc.find((i) => i.name === item.call_type);
if (existingItem) {
existingItem["LLM API requests"] += (item.total_rows || 0) - (item.cache_hit_true_rows || 0);
existingItem["Cache hit"] += item.cache_hit_true_rows || 0;
existingItem["Cached Completion Tokens"] += item.cached_completion_tokens || 0;
existingItem["Generated Completion Tokens"] += item.generated_completion_tokens || 0;
} else {
acc.push({
name: item.call_type,
"LLM API requests": (item.total_rows || 0) - (item.cache_hit_true_rows || 0),
"Cache hit": item.cache_hit_true_rows || 0,
"Cached Completion Tokens": item.cached_completion_tokens || 0,
"Generated Completion Tokens": item.generated_completion_tokens || 0,
});
}
return acc;
}, []);
// set header cache statistics
setCachedResponses(valueFormatterNumbers(cache_hits));
setCachedTokens(valueFormatterNumbers(cached_tokens));
let allRequests = cache_hits + llm_api_requests;
if (allRequests > 0) {
let cache_hit_ratio = ((cache_hits / allRequests) * 100).toFixed(2);
setCacheHitRatio(cache_hit_ratio);
} else {
setCacheHitRatio("0");
}
setFilteredData(processedData);
}, [selectedApiKeys, selectedModels, dateValue, data]);
const uniqueApiKeys = activity?.filter_options.key_aliases ?? [];
const uniqueModels = activity?.filter_options.models ?? [];
const chartData = (activity?.groups ?? []).map(toChartDatum);
const handleRefreshClick = () => {
// Update the 'lastRefreshed' state to the current date and time
const currentDate = new Date();
setLastRefreshed(currentDate.toLocaleString());
refetch();
setLastRefreshed(new Date().toLocaleString());
};
const runCachingHealthCheck = async () => {
@ -257,10 +152,12 @@ const CacheDashboard: React.FC<CachePageProps> = ({ accessToken, token, userRole
}
};
const totals = activity?.totals;
const hasRequests = totals != null && totals.api_requests + totals.cache_hits + totals.failed_requests > 0;
const statCards = [
{ label: "Cache Hit Ratio", value: `${cacheHitRatio}%` },
{ label: "Cache Hits", value: cachedResponses },
{ label: "Cached Completion Tokens", value: cachedTokens },
{ label: "Cache Hit Ratio", value: `${hasRequests ? totals.cache_hit_ratio.toFixed(2) : "0"}%` },
{ label: "Cache Hits", value: valueFormatterNumbers(totals?.cache_hits ?? 0) },
{ label: "Cached Completion Tokens", value: valueFormatterNumbers(totals?.cached_completion_tokens ?? 0) },
];
return (
@ -380,7 +277,6 @@ const CacheDashboard: React.FC<CachePageProps> = ({ accessToken, token, userRole
value={dateValue}
onValueChange={(value) => {
setDateValue(value);
updateCachingData(value.from, value.to);
}}
/>
</div>
@ -404,12 +300,12 @@ const CacheDashboard: React.FC<CachePageProps> = ({ accessToken, token, userRole
</CardHeader>
<CardContent>
<BarChart
data={filteredData}
data={chartData}
stack={true}
index="name"
valueFormatter={valueFormatterNumbers}
categories={["LLM API requests", "Cache hit"]}
colors={["sky", "teal"]}
categories={[REQUEST_SERIES.apiRequests, REQUEST_SERIES.cacheHits, REQUEST_SERIES.failed]}
colors={["sky", "teal", "red"]}
yAxisWidth={48}
/>
</CardContent>
@ -423,7 +319,7 @@ const CacheDashboard: React.FC<CachePageProps> = ({ accessToken, token, userRole
</CardHeader>
<CardContent>
<BarChart
data={filteredData}
data={chartData}
stack={true}
index="name"
valueFormatter={valueFormatterNumbers}

View file

@ -0,0 +1,72 @@
import { renderHook } from "@testing-library/react";
import { beforeEach, describe, expect, it, vi } from "vitest";
import { useCacheActivity, type CacheActivityParams } from "./useCacheActivity";
const useQueryMock = vi.fn();
vi.mock("@/lib/http/api", () => ({
$api: { useQuery: (...args: unknown[]) => useQueryMock(...args) },
}));
const mockUseAuthorized = vi.fn();
vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({
default: () => mockUseAuthorized(),
}));
const params: CacheActivityParams = {
startDate: "2026-07-20",
endDate: "2026-07-27",
keyAliases: ["my-key"],
models: ["gpt-5.1"],
};
const lastCallOptions = (): { enabled: boolean } => {
const calls = useQueryMock.mock.calls;
return calls[calls.length - 1][3] as { enabled: boolean };
};
describe("useCacheActivity", () => {
beforeEach(() => {
vi.clearAllMocks();
useQueryMock.mockReturnValue({ data: undefined });
mockUseAuthorized.mockReturnValue({ accessToken: "test-access-token" });
});
it("queries GET /global/activity/cache_hits with dates and filters as query params", () => {
renderHook(() => useCacheActivity(params));
expect(useQueryMock).toHaveBeenCalledWith(
"get",
"/global/activity/cache_hits",
{
params: {
query: {
start_date: "2026-07-20",
end_date: "2026-07-27",
key_aliases: ["my-key"],
models: ["gpt-5.1"],
},
},
},
expect.any(Object),
);
});
it("enables the query when authorized and both dates are set", () => {
renderHook(() => useCacheActivity(params));
expect(lastCallOptions().enabled).toBe(true);
});
it("disables the query without an access token", () => {
mockUseAuthorized.mockReturnValue({ accessToken: null });
renderHook(() => useCacheActivity(params));
expect(lastCallOptions().enabled).toBe(false);
});
it("disables the query while the date range is incomplete", () => {
renderHook(() => useCacheActivity({ ...params, endDate: undefined }));
expect(lastCallOptions().enabled).toBe(false);
});
});

View file

@ -0,0 +1,32 @@
import { $api } from "@/lib/http/api";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import type { components } from "@/lib/http/schema";
export type CacheActivityResponse = components["schemas"]["CacheActivityResponse"];
export type CacheActivityGroup = components["schemas"]["CacheActivityGroup"];
export interface CacheActivityParams {
startDate: string | undefined;
endDate: string | undefined;
keyAliases: string[];
models: string[];
}
export const useCacheActivity = ({ startDate, endDate, keyAliases, models }: CacheActivityParams) => {
const { accessToken } = useAuthorized();
return $api.useQuery(
"get",
"/global/activity/cache_hits",
{
params: {
query: {
start_date: startDate ?? "",
end_date: endDate ?? "",
key_aliases: keyAliases,
models,
},
},
},
{ enabled: Boolean(accessToken && startDate && endDate) },
);
};

View file

@ -1,7 +1,11 @@
export function navigateWithParams(mutate: (params: URLSearchParams) => void): void {
export function navigateWithParams(mutate: (params: URLSearchParams) => void, mode: "push" | "replace" = "push"): void {
const params = new URLSearchParams(window.location.search);
mutate(params);
const qs = params.toString();
const url = qs ? `${window.location.pathname}?${qs}` : window.location.pathname;
window.history.pushState(null, "", url);
if (mode === "replace") {
window.history.replaceState(null, "", url);
} else {
window.history.pushState(null, "", url);
}
}

View file

@ -1,7 +1,9 @@
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { render, screen } from "@testing-library/react";
import { act, render, screen } from "@testing-library/react";
import React from "react";
import { describe, expect, it, vi } from "vitest";
import { beforeEach, describe, expect, it, vi } from "vitest";
import type OrganizationsTableComponent from "./OrganizationsTable";
import type OrganizationInfoViewComponent from "@/components/organization/organization_view";
vi.mock("@/components/vector_store_management/VectorStoreSelector", () => ({
__esModule: true,
@ -18,12 +20,50 @@ vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({
userRole: null,
}),
}));
type OrganizationsTableProps = React.ComponentProps<typeof OrganizationsTableComponent>;
type OrganizationInfoViewProps = React.ComponentProps<typeof OrganizationInfoViewComponent>;
let capturedTableProps: OrganizationsTableProps | null = null;
vi.mock("./OrganizationsTable", () => ({
__esModule: true,
default: (props: { isLoading: boolean }) => (
<div data-testid="organizations-table">isLoading:{String(props.isLoading)}</div>
),
default: (props: OrganizationsTableProps) => {
capturedTableProps = props;
return <div data-testid="organizations-table">isLoading:{String(props.isLoading)}</div>;
},
}));
const mockOrgInfoView = vi.fn<(props: OrganizationInfoViewProps) => void>();
vi.mock("@/components/organization/organization_view", () => ({
__esModule: true,
default: (props: OrganizationInfoViewProps) => {
mockOrgInfoView(props);
return <div data-testid="organization-info-view" />;
},
}));
// The selected org is URL-derived (?org=) via useOrgDetailRouting. Next's real useSearchParams
// re-renders subscribers on history.pushState/replaceState; mirror that so URL changes propagate.
vi.mock("next/navigation", async () => {
const { useSyncExternalStore } = await import("react");
const LOCATION_CHANGE_EVENT = "test-locationchange";
for (const method of ["pushState", "replaceState"] as const) {
const original = window.history[method].bind(window.history);
window.history[method] = (...args: Parameters<History["pushState"]>) => {
original(...args);
window.dispatchEvent(new Event(LOCATION_CHANGE_EVENT));
};
}
const subscribe = (onChange: () => void) => {
window.addEventListener(LOCATION_CHANGE_EVENT, onChange);
window.addEventListener("popstate", onChange);
return () => {
window.removeEventListener(LOCATION_CHANGE_EVENT, onChange);
window.removeEventListener("popstate", onChange);
};
};
return {
useSearchParams: () => new URLSearchParams(useSyncExternalStore(subscribe, () => window.location.search)),
};
});
import OrganizationsPanel from "./OrganizationsPanel";
@ -34,6 +74,12 @@ const renderWithQueryClient = (ui: React.ReactElement) => {
return render(<QueryClientProvider client={queryClient}>{ui}</QueryClientProvider>);
};
beforeEach(() => {
capturedTableProps = null;
mockOrgInfoView.mockClear();
window.history.replaceState(null, "", "/organizations/");
});
describe("OrganizationsPanel", () => {
it("gates non-premium users behind the enterprise notice", () => {
renderWithQueryClient(<OrganizationsPanel userRole="Admin" accessToken={null} premiumUser={false} />);
@ -55,3 +101,60 @@ describe("OrganizationsPanel", () => {
expect(screen.getByTestId("organizations-table")).toHaveTextContent("isLoading:false");
});
});
describe("OrganizationsPanel - org detail deep link (?org=)", () => {
it("clicking an organization pushes ?org= and opens the detail view", () => {
renderWithQueryClient(<OrganizationsPanel userRole="Admin" accessToken={null} premiumUser={true} />);
act(() => capturedTableProps?.onOrganizationClick("org-deep-link"));
expect(window.location.search).toContain("org=org-deep-link");
expect(mockOrgInfoView).toHaveBeenLastCalledWith(expect.objectContaining({ organizationId: "org-deep-link" }));
});
it("opens the org detail directly from a ?org= deep link", () => {
window.history.replaceState(null, "", "/organizations/?org=org-from-url");
renderWithQueryClient(<OrganizationsPanel userRole="Admin" accessToken={null} premiumUser={true} />);
expect(mockOrgInfoView).toHaveBeenLastCalledWith(
expect.objectContaining({ organizationId: "org-from-url", editOrg: false }),
);
expect(screen.queryByTestId("organizations-table")).not.toBeInTheDocument();
});
it("closing the org detail removes ?org= and returns to the list", () => {
window.history.replaceState(null, "", "/organizations/?org=org-from-url");
renderWithQueryClient(<OrganizationsPanel userRole="Admin" accessToken={null} premiumUser={true} />);
act(() => mockOrgInfoView.mock.calls.at(-1)?.[0].onClose());
expect(window.location.search).not.toContain("org=");
expect(screen.queryByTestId("organization-info-view")).not.toBeInTheDocument();
expect(screen.getByTestId("organizations-table")).toBeInTheDocument();
});
it("the edit action opens the detail in edit mode with ?org= set", () => {
renderWithQueryClient(<OrganizationsPanel userRole="Admin" accessToken={null} premiumUser={true} />);
act(() => capturedTableProps?.onEditClick("org-edit"));
expect(window.location.search).toContain("org=org-edit");
expect(mockOrgInfoView).toHaveBeenLastCalledWith(
expect.objectContaining({ organizationId: "org-edit", editOrg: true }),
);
});
it("a plain row click after leaving an edit view via browser history does not reopen in edit mode", () => {
renderWithQueryClient(<OrganizationsPanel userRole="Admin" accessToken={null} premiumUser={true} />);
act(() => capturedTableProps?.onEditClick("org-edit"));
expect(mockOrgInfoView).toHaveBeenLastCalledWith(expect.objectContaining({ editOrg: true }));
act(() => window.history.pushState(null, "", "/organizations/"));
act(() => capturedTableProps?.onOrganizationClick("org-plain"));
expect(mockOrgInfoView).toHaveBeenLastCalledWith(
expect.objectContaining({ organizationId: "org-plain", editOrg: false }),
);
});
});

View file

@ -1,5 +1,6 @@
import { organizationKeys, useOrganizations } from "@/app/(dashboard)/hooks/organizations/useOrganizations";
import { useUserModels } from "@/app/(dashboard)/hooks/models/useModels";
import { useOrgDetailRouting } from "@/app/(dashboard)/organizations/detailNavigation";
import OrganizationFilters, { FilterState } from "@/app/(dashboard)/organizations/OrganizationFilters";
import { useQueryClient } from "@tanstack/react-query";
import React, { useState } from "react";
@ -19,7 +20,7 @@ interface OrganizationsPanelProps {
}
const OrganizationsPanel: React.FC<OrganizationsPanelProps> = ({ userRole, accessToken, premiumUser }) => {
const [selectedOrgId, setSelectedOrgId] = useState<string | null>(null);
const { orgId: selectedOrgId, openOrg, close: closeOrgDetail } = useOrgDetailRouting();
const [editOrg, setEditOrg] = useState(false);
const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false);
const [orgToDelete, setOrgToDelete] = useState<string | null>(null);
@ -108,7 +109,7 @@ const OrganizationsPanel: React.FC<OrganizationsPanelProps> = ({ userRole, acces
<OrganizationInfoView
organizationId={selectedOrgId}
onClose={() => {
setSelectedOrgId(null);
closeOrgDetail();
setEditOrg(false);
}}
accessToken={accessToken}
@ -132,9 +133,12 @@ const OrganizationsPanel: React.FC<OrganizationsPanelProps> = ({ userRole, acces
isLoading={isLoading}
userRole={userRole}
searchActive={searchActive}
onOrganizationClick={setSelectedOrgId}
onOrganizationClick={(organizationId) => {
setEditOrg(false);
openOrg(organizationId);
}}
onEditClick={(organizationId) => {
setSelectedOrgId(organizationId);
openOrg(organizationId);
setEditOrg(true);
}}
onDeleteClick={handleDelete}

View file

@ -0,0 +1,53 @@
/* @vitest-environment jsdom */
import { act, renderHook } from "@testing-library/react";
import { beforeEach, describe, expect, it, vi } from "vitest";
import { useOrgDetailRouting } from "./detailNavigation";
vi.mock("next/navigation", () => ({ useSearchParams: () => new URLSearchParams(window.location.search) }));
describe("useOrgDetailRouting", () => {
beforeEach(() => {
window.history.pushState(null, "", "/organizations/");
});
it("openOrg sets ?org= via history.pushState (no full navigation)", () => {
const spy = vi.spyOn(window.history, "pushState");
const { result } = renderHook(() => useOrgDetailRouting());
act(() => result.current.openOrg("org-abc123"));
expect(spy).toHaveBeenCalledWith(null, "", expect.stringContaining("org=org-abc123"));
spy.mockRestore();
});
it("openOrg preserves unrelated query params", () => {
window.history.pushState(null, "", "/organizations/?foo=bar");
const spy = vi.spyOn(window.history, "pushState");
const { result } = renderHook(() => useOrgDetailRouting());
act(() => result.current.openOrg("org-abc123"));
const url = spy.mock.calls.at(-1)?.[2] as string;
expect(url).toContain("foo=bar");
expect(url).toContain("org=org-abc123");
spy.mockRestore();
});
it("close removes only the org param", () => {
window.history.pushState(null, "", "/organizations/?foo=bar&org=org-abc123");
const spy = vi.spyOn(window.history, "pushState");
const { result } = renderHook(() => useOrgDetailRouting());
act(() => result.current.close());
const url = spy.mock.calls.at(-1)?.[2] as string;
expect(url).toContain("foo=bar");
expect(url).not.toContain("org=");
spy.mockRestore();
});
it("exposes orgId from ?org=", () => {
window.history.pushState(null, "", "/organizations/?org=org-abc123");
const { result } = renderHook(() => useOrgDetailRouting());
expect(result.current.orgId).toBe("org-abc123");
});
it("orgId is null when no org param is present", () => {
const { result } = renderHook(() => useOrgDetailRouting());
expect(result.current.orgId).toBeNull();
});
});

View file

@ -0,0 +1,32 @@
import { useSearchParams } from "next/navigation";
import { useCallback } from "react";
import { navigateWithParams } from "../navigateWithParams";
export interface OrgDetailRouting {
orgId: string | null;
openOrg: (id: string) => void;
close: () => void;
}
export function useOrgDetailRouting(): OrgDetailRouting {
const searchParams = useSearchParams();
const openOrg = useCallback((id: string) => {
navigateWithParams((params) => {
params.set("org", id);
});
}, []);
const close = useCallback(() => {
navigateWithParams((params) => {
params.delete("org");
});
}, []);
return {
orgId: searchParams?.get("org") ?? null,
openOrg,
close,
};
}

View file

@ -0,0 +1,53 @@
/* @vitest-environment jsdom */
import { act, renderHook } from "@testing-library/react";
import { beforeEach, describe, expect, it, vi } from "vitest";
import { useTeamDetailRouting } from "./detailNavigation";
vi.mock("next/navigation", () => ({ useSearchParams: () => new URLSearchParams(window.location.search) }));
describe("useTeamDetailRouting", () => {
beforeEach(() => {
window.history.pushState(null, "", "/teams/");
});
it("openTeam sets ?team= via history.pushState (no full navigation)", () => {
const spy = vi.spyOn(window.history, "pushState");
const { result } = renderHook(() => useTeamDetailRouting());
act(() => result.current.openTeam("team-abc123"));
expect(spy).toHaveBeenCalledWith(null, "", expect.stringContaining("team=team-abc123"));
spy.mockRestore();
});
it("openTeam preserves unrelated query params", () => {
window.history.pushState(null, "", "/teams/?foo=bar");
const spy = vi.spyOn(window.history, "pushState");
const { result } = renderHook(() => useTeamDetailRouting());
act(() => result.current.openTeam("team-abc123"));
const url = spy.mock.calls.at(-1)?.[2] as string;
expect(url).toContain("foo=bar");
expect(url).toContain("team=team-abc123");
spy.mockRestore();
});
it("close removes only the team param", () => {
window.history.pushState(null, "", "/teams/?foo=bar&team=team-abc123");
const spy = vi.spyOn(window.history, "pushState");
const { result } = renderHook(() => useTeamDetailRouting());
act(() => result.current.close());
const url = spy.mock.calls.at(-1)?.[2] as string;
expect(url).toContain("foo=bar");
expect(url).not.toContain("team=");
spy.mockRestore();
});
it("exposes teamId from ?team=", () => {
window.history.pushState(null, "", "/teams/?team=team-abc123");
const { result } = renderHook(() => useTeamDetailRouting());
expect(result.current.teamId).toBe("team-abc123");
});
it("teamId is null when no team param is present", () => {
const { result } = renderHook(() => useTeamDetailRouting());
expect(result.current.teamId).toBeNull();
});
});

View file

@ -0,0 +1,32 @@
import { useSearchParams } from "next/navigation";
import { useCallback } from "react";
import { navigateWithParams } from "../navigateWithParams";
export interface TeamDetailRouting {
teamId: string | null;
openTeam: (id: string) => void;
close: () => void;
}
export function useTeamDetailRouting(): TeamDetailRouting {
const searchParams = useSearchParams();
const openTeam = useCallback((id: string) => {
navigateWithParams((params) => {
params.set("team", id);
});
}, []);
const close = useCallback(() => {
navigateWithParams((params) => {
params.delete("team");
});
}, []);
return {
teamId: searchParams?.get("team") ?? null,
openTeam,
close,
};
}

View file

@ -1,4 +1,4 @@
import { act, fireEvent, render, screen, waitFor } from "@testing-library/react";
import { act, fireEvent, render, screen, waitFor, within } from "@testing-library/react";
import { beforeAll, beforeEach, describe, expect, it, vi } from "vitest";
import * as networking from "@/components/networking";
import EntityUsage from "./EntityUsage";
@ -497,7 +497,7 @@ describe("EntityUsage", () => {
it.each([
["Cost", "Tag Spend Overview"],
["Model Activity", "metrics-source:models"],
["Model Activity", "metrics-source:model_groups"],
["Key Activity", "metrics-source:api_keys"],
["Endpoint Activity", "Endpoint Usage Panel"],
])("shows only the %s panel for a non-team entity type", async (tabLabel, marker) => {
@ -518,7 +518,7 @@ describe("EntityUsage", () => {
it.each([
["Cost", "Team Spend Overview"],
["Model Activity", "metrics-source:models"],
["Model Activity", "metrics-source:model_groups"],
["Agent Activity", "metrics-source:entities"],
["Key Activity", "metrics-source:api_keys"],
["Endpoint Activity", "Endpoint Usage Panel"],
@ -584,15 +584,41 @@ describe("EntityUsage", () => {
expect(screen.getByText("Request / Token Consumption")).toBeInTheDocument();
});
it("should display Top Models title for non-agent entity types", async () => {
it("should display Top Public Model Names title for non-agent entity types", async () => {
render(<EntityUsage {...defaultProps} entityType="tag" />);
await waitFor(() => {
expect(mockTagDailyActivityCall).toHaveBeenCalled();
});
const topModelsElements = screen.getAllByText("Top Models");
expect(topModelsElements.length).toBeGreaterThan(0);
expect(screen.getByText("Top Public Model Names")).toBeInTheDocument();
});
it("defaults Model Activity to public model names and toggles to litellm models", async () => {
const { container } = render(<EntityUsage {...defaultProps} />);
await waitFor(() => {
expect(mockTagDailyActivityCall).toHaveBeenCalled();
});
act(() => {
fireEvent.click(screen.getByText("Model Activity"));
});
const modelActivityPanel = () => selectedPanels(container)[0] as HTMLElement;
expect(modelActivityPanel().textContent).toContain("metrics-source:model_groups");
act(() => {
fireEvent.click(within(modelActivityPanel()).getByText("Litellm Model Name"));
});
expect(modelActivityPanel().textContent).toContain("metrics-source:models");
act(() => {
fireEvent.click(within(modelActivityPanel()).getByText("Public Model Name"));
});
expect(modelActivityPanel().textContent).toContain("metrics-source:model_groups");
});
it("should display Top Agents title for agent entity type", async () => {

View file

@ -49,6 +49,7 @@ import {
} from "@/components/UsagePage/types";
import { valueFormatterSpend } from "@/components/UsagePage/utils/value_formatters";
import EndpointUsage from "../EndpointUsage/EndpointUsage";
import ModelViewToggle, { ModelViewType } from "../ModelViewToggle";
import TopKeyView from "@/components/UsagePage/components/EntityUsage/TopKeyView";
import TopModelView from "./TopModelView";
@ -110,6 +111,7 @@ const ENTITY_FETCH_FNS: Record<EntityType, (...args: any[]) => Promise<any>> = {
const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, entityId, entityList, dateValue }) => {
const { teams } = useTeams();
const [selectedTags, setSelectedTags] = useState<string[]>([]);
const [modelViewType, setModelViewType] = useState<ModelViewType>("groups");
const [topKeysLimit, setTopKeysLimit] = useState<number>(5);
const [topModelsLimit, setTopModelsLimit] = useState<number>(5);
const [topAgentsLimit, setTopAgentsLimit] = useState<number>(5);
@ -153,14 +155,15 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
const agentSpendData = agentSpendDataRaw as unknown as EntitySpendData;
const modelMetrics = processActivityData(spendData, "models", teams || []);
const modelBreakdownKey = modelViewType === "groups" ? "model_groups" : "models";
const modelMetrics = processActivityData(spendData, modelBreakdownKey, teams || []);
const keyMetrics = processActivityData(spendData, "api_keys", teams || []);
const agentMetrics = entityType === "team" ? processActivityData(agentSpendData, "entities", teams || []) : {};
const getTopModels = () => {
const modelSpend: { [key: string]: any } = {};
spendData.results.forEach((day) => {
Object.entries(day.breakdown.models || {}).forEach(([model, metrics]) => {
Object.entries(day.breakdown[modelBreakdownKey] || {}).forEach(([model, metrics]) => {
if (!modelSpend[model]) {
modelSpend[model] = {
spend: 0,
@ -406,6 +409,8 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
const capitalizedEntityLabel = entityType.charAt(0).toUpperCase() + entityType.slice(1);
const modelViewTitle = modelViewType === "groups" ? "Top Public Model Names" : "Top Litellm Models";
const costPanel = (
<Grid numItems={2} className="gap-2 w-full">
{/* Total Spend Card */}
@ -604,7 +609,10 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
{/* Top Models */}
<Col numColSpan={1}>
<Card>
<Title>{entityType === "agent" ? "Top Agents" : "Top Models"}</Title>
<div className="flex justify-between items-center">
<Title>{entityType === "agent" ? "Top Agents" : modelViewTitle}</Title>
<ModelViewToggle value={modelViewType} onChange={setModelViewType} />
</div>
<TopModelView
topModels={getTopModels()}
topModelsLimit={topModelsLimit}
@ -691,7 +699,14 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
{
key: "models",
label: entityType === "agent" ? "Request / Token Consumption" : "Model Activity",
content: <ActivityMetrics modelMetrics={modelMetrics} hidePromptCachingMetrics={entityType === "agent"} />,
content: (
<>
<div className="flex justify-end mt-2 mb-4">
<ModelViewToggle value={modelViewType} onChange={setModelViewType} />
</div>
<ActivityMetrics modelMetrics={modelMetrics} hidePromptCachingMetrics={entityType === "agent"} />
</>
),
},
...(entityType === "team"
? [{ key: "agents", label: "Agent Activity", content: <ActivityMetrics modelMetrics={agentMetrics} /> }]

View file

@ -0,0 +1,29 @@
export type ModelViewType = "groups" | "individual";
const MODEL_VIEW_OPTIONS: readonly { value: ModelViewType; label: string }[] = [
{ value: "groups", label: "Public Model Name" },
{ value: "individual", label: "Litellm Model Name" },
];
interface ModelViewToggleProps {
value: ModelViewType;
onChange: (value: ModelViewType) => void;
}
export default function ModelViewToggle({ value, onChange }: ModelViewToggleProps) {
return (
<div className="flex bg-gray-100 rounded-lg p-1">
{MODEL_VIEW_OPTIONS.map((option) => (
<button
key={option.value}
className={`px-3 py-1 text-sm rounded-md transition-colors ${
value === option.value ? "bg-white shadow-xs text-gray-900" : "text-gray-600 hover:text-gray-900"
}`}
onClick={() => onChange(option.value)}
>
{option.label}
</button>
))}
</div>
);
}

View file

@ -30,8 +30,10 @@ vi.mock("@/components/networking", () => ({
// Mock child components to simplify testing
vi.mock("@/components/activity_metrics", () => ({
ActivityMetrics: () => <div>Activity Metrics</div>,
processActivityData: () => ({ data: [], metadata: {} }),
ActivityMetrics: ({ modelMetrics }: { modelMetrics?: { __source?: string } }) => (
<div>{`activity-source:${modelMetrics?.__source ?? "none"}`}</div>
),
processActivityData: (_data: unknown, key: string) => ({ __source: key }),
}));
vi.mock("@/components/view_user_spend", () => ({
@ -1043,8 +1045,8 @@ describe("UsagePage", () => {
// Default should be "groups" view showing "Top Public Model Names"
expect(screen.getByText("Top Public Model Names")).toBeInTheDocument();
expect(screen.getByText("Public Model Name")).toBeInTheDocument();
expect(screen.getByText("Litellm Model Name")).toBeInTheDocument();
expect(screen.getAllByText("Public Model Name").length).toBeGreaterThan(0);
expect(screen.getAllByText("Litellm Model Name").length).toBeGreaterThan(0);
});
it("should switch to Litellm Model Name view on toggle click", async () => {
@ -1055,7 +1057,7 @@ describe("UsagePage", () => {
});
// Click the "Litellm Model Name" toggle
const litellmToggle = screen.getByText("Litellm Model Name");
const litellmToggle = screen.getAllByText("Litellm Model Name")[0];
act(() => {
fireEvent.click(litellmToggle);
});
@ -1074,7 +1076,7 @@ describe("UsagePage", () => {
});
// Switch to individual first
const litellmToggle = screen.getByText("Litellm Model Name");
const litellmToggle = screen.getAllByText("Litellm Model Name")[0];
act(() => {
fireEvent.click(litellmToggle);
});
@ -1084,7 +1086,7 @@ describe("UsagePage", () => {
});
// Switch back to groups
const publicToggle = screen.getByText("Public Model Name");
const publicToggle = screen.getAllByText("Public Model Name")[0];
act(() => {
fireEvent.click(publicToggle);
});
@ -1093,6 +1095,34 @@ describe("UsagePage", () => {
expect(screen.getByText("Top Public Model Names")).toBeInTheDocument();
});
});
it("should feed the Model Activity tab from the model_groups breakdown by default", async () => {
renderWithProviders(<UsagePage {...defaultProps} />);
await waitFor(() => {
expect(mockUserDailyActivityAggregatedCall).toHaveBeenCalled();
});
expect(screen.getByText("activity-source:model_groups")).toBeInTheDocument();
expect(screen.queryByText("activity-source:models")).not.toBeInTheDocument();
});
it("should switch the Model Activity tab to the litellm models breakdown on toggle click", async () => {
renderWithProviders(<UsagePage {...defaultProps} />);
await waitFor(() => {
expect(mockUserDailyActivityAggregatedCall).toHaveBeenCalled();
});
act(() => {
fireEvent.click(screen.getAllByText("Litellm Model Name")[0]);
});
await waitFor(() => {
expect(screen.getByText("activity-source:models")).toBeInTheDocument();
});
expect(screen.queryByText("activity-source:model_groups")).not.toBeInTheDocument();
});
});
describe("customer usage banner", () => {

View file

@ -55,6 +55,7 @@ import { DailyData, KeyMetricWithMetadata, MetricWithMetadata } from "@/componen
import { valueFormatterSpend } from "@/components/UsagePage/utils/value_formatters";
import EndpointUsage from "./EndpointUsage/EndpointUsage";
import EntityUsage, { EntityList } from "./EntityUsage/EntityUsage";
import ModelViewToggle, { ModelViewType } from "./ModelViewToggle";
import SpendByProvider from "./EntityUsage/SpendByProvider";
import TopKeyView from "@/components/UsagePage/components/EntityUsage/TopKeyView";
import UsageAIChatPanel from "./UsageAIChatPanel";
@ -143,7 +144,7 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
// For admins: null means global view (all users), a string means filter by that user
// For non-admins: always set to their own user ID
const [selectedUserId, setSelectedUserId] = useState<string | null>(isAdmin ? null : userID || null);
const [modelViewType, setModelViewType] = useState<"groups" | "individual">("groups");
const [modelViewType, setModelViewType] = useState<ModelViewType>("groups");
const [isCloudZeroModalOpen, setIsCloudZeroModalOpen] = useState(false);
const [isGlobalExportModalOpen, setIsGlobalExportModalOpen] = useState(false);
const [isAiChatOpen, setIsAiChatOpen] = useState(false);
@ -438,7 +439,10 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
() => [...userSpendData.results].sort((a, b) => new Date(a.date).getTime() - new Date(b.date).getTime()),
[userSpendData.results],
);
const modelMetrics = useMemo(() => processActivityData(userSpendData, "models", teams), [userSpendData, teams]);
const modelMetrics = useMemo(
() => processActivityData(userSpendData, modelViewType === "groups" ? "model_groups" : "models", teams),
[userSpendData, modelViewType, teams],
);
const keyMetrics = useMemo(() => processActivityData(userSpendData, "api_keys", teams), [userSpendData, teams]);
const mcpServerMetrics = useMemo(
() => processActivityData(userSpendData, "mcp_servers", teams),
@ -753,28 +757,7 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
value={topModelsLimit}
onChange={(value) => setTopModelsLimit(value as number)}
/>
<div className="flex bg-gray-100 rounded-lg p-1">
<button
className={`px-3 py-1 text-sm rounded-md transition-colors ${
modelViewType === "groups"
? "bg-white shadow-xs text-gray-900"
: "text-gray-600 hover:text-gray-900"
}`}
onClick={() => setModelViewType("groups")}
>
Public Model Name
</button>
<button
className={`px-3 py-1 text-sm rounded-md transition-colors ${
modelViewType === "individual"
? "bg-white shadow-xs text-gray-900"
: "text-gray-600 hover:text-gray-900"
}`}
onClick={() => setModelViewType("individual")}
>
Litellm Model Name
</button>
</div>
<ModelViewToggle value={modelViewType} onChange={setModelViewType} />
</div>
{loading ? (
<ChartLoader isDateChanging={isDateChanging} />
@ -839,6 +822,9 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
{/* Activity Panel */}
<TabPanel>
<div className="flex justify-end mt-2 mb-4">
<ModelViewToggle value={modelViewType} onChange={setModelViewType} />
</div>
<ActivityMetrics modelMetrics={modelMetrics} />
</TabPanel>
<TabPanel>

View file

@ -72,6 +72,31 @@ vi.mock("@/components/team/TeamInfo", () => ({
},
}));
// The selected team is URL-derived (?team=) via useTeamDetailRouting. Next's real useSearchParams
// re-renders subscribers on history.pushState/replaceState; mirror that so URL changes propagate.
vi.mock("next/navigation", async () => {
const { useSyncExternalStore } = await import("react");
const LOCATION_CHANGE_EVENT = "test-locationchange";
for (const method of ["pushState", "replaceState"] as const) {
const original = window.history[method].bind(window.history);
window.history[method] = (...args: Parameters<History["pushState"]>) => {
original(...args);
window.dispatchEvent(new Event(LOCATION_CHANGE_EVENT));
};
}
const subscribe = (onChange: () => void) => {
window.addEventListener(LOCATION_CHANGE_EVENT, onChange);
window.addEventListener("popstate", onChange);
return () => {
window.removeEventListener(LOCATION_CHANGE_EVENT, onChange);
window.removeEventListener("popstate", onChange);
};
};
return {
useSearchParams: () => new URLSearchParams(useSyncExternalStore(subscribe, () => window.location.search)),
};
});
vi.mock("./ModelSelect/ModelSelect", () => {
const ModelSelect = React.forwardRef(({ value, onChange, dataTestId, id }: any, ref: any) => {
return (
@ -159,6 +184,7 @@ const renderWithQueryClient = (component: React.ReactElement) => {
// Re-establish safe defaults before every test (clearAllMocks keeps return values, so restore them here).
beforeEach(() => {
mockTeamsTableProps = null;
window.history.replaceState(null, "", "/teams/");
});
describe("Teams - handleCreate organization handling", () => {
@ -436,6 +462,47 @@ describe("Teams - premium props", () => {
});
});
describe("Teams - team detail deep link (?team=)", () => {
beforeEach(() => {
vi.clearAllMocks();
mockTeamInfoView.mockClear();
vi.mocked(fetchAvailableModelsForTeamOrKey).mockResolvedValue([]);
vi.mocked(fetchMCPAccessGroups).mockResolvedValue([]);
vi.mocked(getGuardrailsList).mockResolvedValue({ guardrails: [] });
mockUseOrganizations.mockReturnValue({ data: [] });
});
it("selecting a team pushes ?team= to the URL", async () => {
renderWithQueryClient(<Teams accessToken="test-token" userID="user-123" userRole="Admin" />);
await waitFor(() => expect(mockTeamsTableProps).not.toBeNull());
act(() => mockTeamsTableProps.onSelectTeam({ ...baseTableTeam, team_id: "team-deep-link" }));
expect(window.location.search).toContain("team=team-deep-link");
await waitFor(() => expect(mockTeamInfoView).toHaveBeenCalled());
expect(mockTeamInfoView).toHaveBeenLastCalledWith(expect.objectContaining({ teamId: "team-deep-link" }));
});
it("opens the team detail view directly from a ?team= deep link", async () => {
window.history.replaceState(null, "", "/teams/?team=team-from-url");
renderWithQueryClient(<Teams accessToken="test-token" userID="user-123" userRole="Admin" />);
await waitFor(() => expect(mockTeamInfoView).toHaveBeenCalled());
expect(mockTeamInfoView).toHaveBeenLastCalledWith(expect.objectContaining({ teamId: "team-from-url" }));
});
it("closing the team detail view removes ?team= from the URL", async () => {
window.history.replaceState(null, "", "/teams/?team=team-from-url");
renderWithQueryClient(<Teams accessToken="test-token" userID="user-123" userRole="Admin" />);
await waitFor(() => expect(mockTeamInfoView).toHaveBeenCalled());
act(() => mockTeamInfoView.mock.calls.at(-1)?.[0].onClose());
expect(window.location.search).not.toContain("team=");
await waitFor(() => expect(screen.queryByTestId("team-info-view")).not.toBeInTheDocument());
});
});
describe("Teams - Create Team CTA is grouped with the tabs on the left", () => {
beforeEach(() => {
vi.clearAllMocks();

View file

@ -12,6 +12,7 @@ import { useQueryClient } from "@tanstack/react-query";
import { PageHeader } from "@/components/shared/PageHeader";
import { Button as UIButton } from "@/components/ui/button";
import { teamsTableKeys } from "@/app/(dashboard)/hooks/teams/useTeams";
import { useTeamDetailRouting } from "@/app/(dashboard)/teams/detailNavigation";
import { TeamsTable } from "./TeamsPage/TeamsTable";
import AccessGroupSelector from "./common_components/AccessGroupSelector";
import PassThroughRoutesSelector from "./common_components/PassThroughRoutesSelector";
@ -135,7 +136,7 @@ const Teams: React.FC<TeamProps> = ({ accessToken, userID, userRole, premiumUser
const [editModalVisible, setEditModalVisible] = useState(false);
const [selectedTeam, setSelectedTeam] = useState<Team | null>(null);
const [selectedTeamId, setSelectedTeamId] = useState<string | null>(null);
const { teamId: selectedTeamId, openTeam, close: closeTeamDetail } = useTeamDetailRouting();
const [editTeam, setEditTeam] = useState<boolean>(false);
const [isTeamModalVisible, setIsTeamModalVisible] = useState(false);
@ -482,12 +483,12 @@ const Teams: React.FC<TeamProps> = ({ accessToken, userID, userRole, premiumUser
userID={userID}
onSelectTeam={(team) => {
setSelectedTeam(team);
setSelectedTeamId(team.team_id);
openTeam(team.team_id);
setEditTeam(false);
}}
onEditTeam={(team) => {
setSelectedTeam(team);
setSelectedTeamId(team.team_id);
openTeam(team.team_id);
setEditTeam(true);
}}
onDeleteTeam={handleDelete}
@ -547,11 +548,11 @@ const Teams: React.FC<TeamProps> = ({ accessToken, userID, userRole, premiumUser
}}
onClose={() => {
setSelectedTeam(null);
setSelectedTeamId(null);
closeTeamDetail();
setEditTeam(false);
}}
accessToken={accessToken}
is_team_admin={is_team_admin(selectedTeam)}
is_team_admin={is_team_admin(selectedTeam?.team_id === selectedTeamId ? selectedTeam : null)}
is_proxy_admin={userRole == "Admin"}
userModels={userModels}
editTeam={editTeam}

View file

@ -716,6 +716,60 @@ describe("processActivityData", () => {
expect(result["gpt-4"].total_spend).toBe(100.5);
});
it("should process model_groups data keyed by public model name including fallback entries", () => {
const upstreamModelMetrics = {
...EMPTY_SPEND_METRICS,
spend: 10,
api_requests: 10,
successful_requests: 10,
};
const dailyActivityWithModelGroups: { results: DailyData[] } = {
results: [
{
date: "2025-01-01",
metrics: upstreamModelMetrics,
breakdown: {
...EMPTY_BREAKDOWN,
models: {
"gpt-5.2": {
metrics: upstreamModelMetrics,
metadata: {},
api_key_breakdown: {},
},
},
model_groups: {
"gpt-5.2-eu": {
metrics: { ...EMPTY_SPEND_METRICS, spend: 7, api_requests: 7, successful_requests: 7 },
metadata: {},
api_key_breakdown: {
"key-1": {
metrics: { ...EMPTY_SPEND_METRICS, spend: 7, api_requests: 7, total_tokens: 700 },
metadata: { key_alias: "eu-key", team_id: "team1" },
},
},
},
"gpt-5.2": {
metrics: { ...EMPTY_SPEND_METRICS, spend: 3, api_requests: 3, successful_requests: 3 },
metadata: {},
api_key_breakdown: {},
},
},
},
},
],
};
const result = processActivityData(dailyActivityWithModelGroups, "model_groups");
expect(Object.keys(result).sort()).toEqual(["gpt-5.2", "gpt-5.2-eu"]);
expect(result["gpt-5.2-eu"].label).toBe("gpt-5.2-eu");
expect(result["gpt-5.2-eu"].total_spend).toBe(7);
expect(result["gpt-5.2-eu"].top_api_keys).toHaveLength(1);
expect(result["gpt-5.2-eu"].top_api_keys[0].key_alias).toBe("eu-key");
expect(result["gpt-5.2"].total_spend).toBe(3);
expect(result["gpt-5.2"].total_requests).toBe(3);
});
it("should process data for mcp_servers key", () => {
const dailyActivityWithMCP: { results: DailyData[] } = {
results: [

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