diff --git a/.gitignore b/.gitignore index b812d45e349..13f2202305d 100644 --- a/.gitignore +++ b/.gitignore @@ -141,3 +141,4 @@ crash.*.log .coverage ui/litellm-dashboard/out/ +litellm.log diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 28602fc235f..db3c2502e94 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -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 diff --git a/litellm-rust/ADDING_A_PROVIDER.md b/litellm-rust/ADDING_A_PROVIDER.md index 5f933ec4fa8..857a744e014 100644 --- a/litellm-rust/ADDING_A_PROVIDER.md +++ b/litellm-rust/ADDING_A_PROVIDER.md @@ -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//`; `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//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///transformation.rs`: implement that trait as a `const __CONFIG`, mirroring the Python provider tree. Add parity unit tests. -3. **HTTP / transport (the host)** — `crates/providers/src/.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 (request) -> CoreResult`, the Rust equivalent of `litellm.()`, plus a `_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///transformation.rs`: implement that trait as a `const __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`. diff --git a/litellm-rust/AGENTS.md b/litellm-rust/AGENTS.md index 398eec4685c..36a5ad5a8f4 100644 --- a/litellm-rust/AGENTS.md +++ b/litellm-rust/AGENTS.md @@ -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//`, 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. diff --git a/litellm-rust/CLAUDE.md b/litellm-rust/CLAUDE.md index 0659e63df39..fe6ceedbb86 100644 --- a/litellm-rust/CLAUDE.md +++ b/litellm-rust/CLAUDE.md @@ -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//` owns the route contract, shared types, and provider - template traits. For OCR, this means `core/src/ocr`. +- `core/src//` 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///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 `_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. diff --git a/litellm-rust/README.md b/litellm-rust/README.md index 1646c90ad76..bcccf93300b 100644 --- a/litellm-rust/README.md +++ b/litellm-rust/README.md @@ -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///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///transformation.rs`. The bridge exposes one +function per top-level route, mirroring the core entrypoints. ## Checks diff --git a/litellm-rust/crates/CODING_STANDARDS/PROVIDER_CODING_STANDARDS.md b/litellm-rust/crates/CODING_STANDARDS/PROVIDER_CODING_STANDARDS.md index ed44dc4c729..a1860d8a9c9 100644 --- a/litellm-rust/crates/CODING_STANDARDS/PROVIDER_CODING_STANDARDS.md +++ b/litellm-rust/crates/CODING_STANDARDS/PROVIDER_CODING_STANDARDS.md @@ -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///`; a route is a module, never a new crate. -9. Route entry point stays thin: `()` -> `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::::()` -> `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 diff --git a/litellm-rust/crates/ai-gateway/AGENTS.md b/litellm-rust/crates/ai-gateway/AGENTS.md index d9e6e1adde5..92567091cd3 100644 --- a/litellm-rust/crates/ai-gateway/AGENTS.md +++ b/litellm-rust/crates/ai-gateway/AGENTS.md @@ -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. diff --git a/litellm-rust/crates/ai-gateway/README.md b/litellm-rust/crates/ai-gateway/README.md index f913beff6d5..7a6c620ee84 100644 --- a/litellm-rust/crates/ai-gateway/README.md +++ b/litellm-rust/crates/ai-gateway/README.md @@ -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. diff --git a/litellm-rust/crates/ai-gateway/src/constants.rs b/litellm-rust/crates/ai-gateway/src/constants.rs index 74808cf1ce6..78af374bf70 100644 --- a/litellm-rust/crates/ai-gateway/src/constants.rs +++ b/litellm-rust/crates/ai-gateway/src/constants.rs @@ -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] = diff --git a/litellm-rust/crates/ai-gateway/src/io/messages.rs b/litellm-rust/crates/ai-gateway/src/io/messages.rs deleted file mode 100644 index 86170e45678..00000000000 --- a/litellm-rust/crates/ai-gateway/src/io/messages.rs +++ /dev/null @@ -1 +0,0 @@ -pub use crate::messages::{MessagesRequest, messages}; diff --git a/litellm-rust/crates/ai-gateway/src/io/mod.rs b/litellm-rust/crates/ai-gateway/src/io/mod.rs index 6129a808965..cce56dd2121 100644 --- a/litellm-rust/crates/ai-gateway/src/io/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/io/mod.rs @@ -1,5 +1,4 @@ pub mod audio_transcription; -pub mod messages; pub mod ocr; pub mod realtime; pub mod realtime_pool; diff --git a/litellm-rust/crates/ai-gateway/src/lib.rs b/litellm-rust/crates/ai-gateway/src/lib.rs index c44d661c29e..057db6457c4 100644 --- a/litellm-rust/crates/ai-gateway/src/lib.rs +++ b/litellm-rust/crates/ai-gateway/src/lib.rs @@ -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 diff --git a/litellm-rust/crates/ai-gateway/src/messages/mod.rs b/litellm-rust/crates/ai-gateway/src/messages/mod.rs deleted file mode 100644 index fd2dd546941..00000000000 --- a/litellm-rust/crates/ai-gateway/src/messages/mod.rs +++ /dev/null @@ -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 { - 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 { - 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; diff --git a/litellm-rust/crates/ai-gateway/src/messages/types.rs b/litellm-rust/crates/ai-gateway/src/messages/types.rs deleted file mode 100644 index 848fadb4b02..00000000000 --- a/litellm-rust/crates/ai-gateway/src/messages/types.rs +++ /dev/null @@ -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>, - pub timeout: Option, -} - -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, -} diff --git a/litellm-rust/crates/ai-gateway/src/routes/AGENTS.md b/litellm-rust/crates/ai-gateway/src/routes/AGENTS.md index 02c5f18c4f3..3eee43e7a2f 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/AGENTS.md +++ b/litellm-rust/crates/ai-gateway/src/routes/AGENTS.md @@ -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//`. - 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. diff --git a/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs b/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs index 75ed26e5be8..5f4c5fe8de4 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs @@ -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}")) }) } diff --git a/litellm-rust/crates/core/AGENTS.md b/litellm-rust/crates/core/AGENTS.md index 8740dccaf01..aee8b4937ef 100644 --- a/litellm-rust/crates/core/AGENTS.md +++ b/litellm-rust/crates/core/AGENTS.md @@ -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//` 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. diff --git a/litellm-rust/crates/core/CLAUDE.md b/litellm-rust/crates/core/CLAUDE.md index 20873878967..5d36305ded5 100644 --- a/litellm-rust/crates/core/CLAUDE.md +++ b/litellm-rust/crates/core/CLAUDE.md @@ -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 `_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 diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index 65c6db7412c..ab8050734f2 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -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"] } diff --git a/litellm-rust/crates/core/src/constants.rs b/litellm-rust/crates/core/src/constants.rs index 5826a5bc9c1..caada1d98b0 100644 --- a/litellm-rust/crates/core/src/constants.rs +++ b/litellm-rust/crates/core/src/constants.rs @@ -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"; diff --git a/litellm-rust/crates/ai-gateway/src/messages/client.rs b/litellm-rust/crates/core/src/messages/client.rs similarity index 100% rename from litellm-rust/crates/ai-gateway/src/messages/client.rs rename to litellm-rust/crates/core/src/messages/client.rs diff --git a/litellm-rust/crates/ai-gateway/src/messages/common_utils.rs b/litellm-rust/crates/core/src/messages/common_utils.rs similarity index 83% rename from litellm-rust/crates/ai-gateway/src/messages/common_utils.rs rename to litellm-rust/crates/core/src/messages/common_utils.rs index 68ecc3f17c1..9dcfcaa71e3 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/common_utils.rs +++ b/litellm-rust/crates/core/src/messages/common_utils.rs @@ -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 { diff --git a/litellm-rust/crates/ai-gateway/src/messages/handler.rs b/litellm-rust/crates/core/src/messages/handler.rs similarity index 84% rename from litellm-rust/crates/ai-gateway/src/messages/handler.rs rename to litellm-rust/crates/core/src/messages/handler.rs index 90c12367f50..1c895f66eba 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/handler.rs +++ b/litellm-rust/crates/core/src/messages/handler.rs @@ -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 { +) -> CoreResult { 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( diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index ec2fbb969a6..acb36d89daf 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -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 { + execute_messages_provider_call(prepare_messages_call(request)?).await +} + +pub async fn messages_stream(request: MessagesRequest<'_>) -> CoreResult { + execute_messages_provider_stream(prepare_messages_call(request)?).await +} + +#[cfg(test)] +mod tests; diff --git a/litellm-rust/crates/ai-gateway/src/messages/prepare.rs b/litellm-rust/crates/core/src/messages/prepare.rs similarity index 92% rename from litellm-rust/crates/ai-gateway/src/messages/prepare.rs rename to litellm-rust/crates/core/src/messages/prepare.rs index 9a027490eb6..94b5b1eaed7 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/prepare.rs +++ b/litellm-rust/crates/core/src/messages/prepare.rs @@ -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( diff --git a/litellm-rust/crates/ai-gateway/src/messages/tests.rs b/litellm-rust/crates/core/src/messages/tests.rs similarity index 97% rename from litellm-rust/crates/ai-gateway/src/messages/tests.rs rename to litellm-rust/crates/core/src/messages/tests.rs index 23a53e98045..9fc1763683b 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/tests.rs +++ b/litellm-rust/crates/core/src/messages/tests.rs @@ -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"); diff --git a/litellm-rust/crates/core/src/messages/types.rs b/litellm-rust/crates/core/src/messages/types.rs index 11fe17ea40f..b9f807c29fd 100644 --- a/litellm-rust/crates/core/src/messages/types.rs +++ b/litellm-rust/crates/core/src/messages/types.rs @@ -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>, + pub timeout: Option, +} + +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, +} + #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] #[serde(untagged)] pub enum SystemPrompt { diff --git a/litellm-rust/crates/python-bridge/AGENTS.md b/litellm-rust/crates/python-bridge/AGENTS.md index d6d3d90e6ab..ad3cddfa5fd 100644 --- a/litellm-rust/crates/python-bridge/AGENTS.md +++ b/litellm-rust/crates/python-bridge/AGENTS.md @@ -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. diff --git a/litellm-rust/crates/python-bridge/CLAUDE.md b/litellm-rust/crates/python-bridge/CLAUDE.md index e5d021ec25b..3ce8b8c639a 100644 --- a/litellm-rust/crates/python-bridge/CLAUDE.md +++ b/litellm-rust/crates/python-bridge/CLAUDE.md @@ -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 diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index ee9bdd0b81f..f0cc26a0cca 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -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> { Ok(json.call_method1("loads", (encoded,))?.unbind()) } +fn messages_response_to_py( + py: Python<'_>, + response: AnthropicMessagesResponse, +) -> PyResult> { + 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)) }) } diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index 12465377b51..11e7ab062b6 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -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) diff --git a/litellm/integrations/otel/model/semconv.py b/litellm/integrations/otel/model/semconv.py index 44b2f7e0488..1abe8ca33fa 100644 --- a/litellm/integrations/otel/model/semconv.py +++ b/litellm/integrations/otel/model/semconv.py @@ -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" diff --git a/litellm/integrations/otel/model/utils.py b/litellm/integrations/otel/model/utils.py index f37afc97879..ab54a558a9a 100644 --- a/litellm/integrations/otel/model/utils.py +++ b/litellm/integrations/otel/model/utils.py @@ -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 diff --git a/litellm/integrations/otel/plumbing/providers.py b/litellm/integrations/otel/plumbing/providers.py index 6bb067ed170..6c276184a0d 100644 --- a/litellm/integrations/otel/plumbing/providers.py +++ b/litellm/integrations/otel/plumbing/providers.py @@ -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`` (Grafana Cloud does, because a bare space + is not representable there) has to reach the exporter as ``Basic ``, + 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 ``) 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() diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 2cef600ea32..87d9b6afc18 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -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", diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 5c677fc0c91..8e4904b1128 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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 diff --git a/litellm/proxy/analytics_endpoints/analytics_endpoints.py b/litellm/proxy/analytics_endpoints/analytics_endpoints.py index 4c1ff31e5a1..cb22c468c1e 100644 --- a/litellm/proxy/analytics_endpoints/analytics_endpoints.py +++ b/litellm/proxy/analytics_endpoints/analytics_endpoints.py @@ -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 [], + ) diff --git a/litellm/proxy/analytics_endpoints/cache_activity.py b/litellm/proxy/analytics_endpoints/cache_activity.py new file mode 100644 index 00000000000..6e20382f56e --- /dev/null +++ b/litellm/proxy/analytics_endpoints/cache_activity.py @@ -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 [])], + ), + ) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 07dfdc4fb43..c46bc110ca8 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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() diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 3e4a454242e..0111be75b1a 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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: diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index a2dc5e1caf5..a91b29002e3 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -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 diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index 9b756d14815..8a5a31710cf 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -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), diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index a2c16e88839..1bd0a19bfb3 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -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 diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 34932d08751..44425734dca 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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: diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 1c0e7211493..b6422d7f5ae 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -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, diff --git a/litellm/proxy/management_endpoints/scim/scim_transformations.py b/litellm/proxy/management_endpoints/scim/scim_transformations.py index 65651752944..d572255fd32 100644 --- a/litellm/proxy/management_endpoints/scim/scim_transformations.py +++ b/litellm/proxy/management_endpoints/scim/scim_transformations.py @@ -175,6 +175,7 @@ class ScimTransformations: SCIMMember( value=ScimTransformations._get_scim_member_value(member), display=ScimTransformations._get_scim_member_display(member), + type="User", ) ) diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index 90eae5bbb21..e44299d018d 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -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 diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 769ae085d33..9f79496a010 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -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, diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index efd1d6b6cee..b2e36188681 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -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, diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index 0c9aa667751..f1bcfbafe58 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -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, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 4486cd7de59..70484eb1e4e 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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 diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index 0c525ee9466..42788227acc 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -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( diff --git a/litellm/router.py b/litellm/router.py index 759d6a0024c..69535d7c74a 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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] diff --git a/litellm/types/proxy/management_endpoints/scim_v2.py b/litellm/types/proxy/management_endpoints/scim_v2.py index 8c434481975..e7f5c85e6c6 100644 --- a/litellm/types/proxy/management_endpoints/scim_v2.py +++ b/litellm/types/proxy/management_endpoints/scim_v2.py @@ -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): diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index db28118d52b..0edd3bd5f30 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -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", diff --git a/tests/e2e/guardrails/guardrails_client.py b/tests/e2e/guardrails/guardrails_client.py index 5a54a4f0bbc..93861d19922 100644 --- a/tests/e2e/guardrails/guardrails_client.py +++ b/tests/e2e/guardrails/guardrails_client.py @@ -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 ) diff --git a/tests/e2e/guardrails/test_presidio_guardrail_e2e.py b/tests/e2e/guardrails/test_presidio_guardrail_e2e.py deleted file mode 100644 index 9742dfc6ae7..00000000000 --- a/tests/e2e/guardrails/test_presidio_guardrail_e2e.py +++ /dev/null @@ -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 `` placeholders (e.g. -``) 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 = "" - -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}" - ) diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_components.py b/tests/test_litellm/integrations/otel/test_otel_v2_components.py index bace586b13d..85122e29a7f 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_components.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_components.py @@ -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``. 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). diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_logger.py b/tests/test_litellm/integrations/otel/test_otel_v2_logger.py index 95d6b598d5e..9693950e834 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_logger.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_logger.py @@ -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", } diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py b/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py index 29067f91b5a..b2d89053ba3 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py @@ -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( diff --git a/tests/test_litellm/integrations/test_opentelemetry.py b/tests/test_litellm/integrations/test_opentelemetry.py index 7ffd09b931f..05205cb76f2 100644 --- a/tests/test_litellm/integrations/test_opentelemetry.py +++ b/tests/test_litellm/integrations/test_opentelemetry.py @@ -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 ) diff --git a/tests/test_litellm/proxy/analytics_endpoints/__init__.py b/tests/test_litellm/proxy/analytics_endpoints/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/proxy/analytics_endpoints/test_analytics_endpoints.py b/tests/test_litellm/proxy/analytics_endpoints/test_analytics_endpoints.py new file mode 100644 index 00000000000..c48b8cfd5a5 --- /dev/null +++ b/tests/test_litellm/proxy/analytics_endpoints/test_analytics_endpoints.py @@ -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 diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index ccb20976df9..34a353966bf 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -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" diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index dc9723c5eb9..8c131758cb3 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -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:" 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. diff --git a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py index 9573fddd435..6a185988c9b 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py @@ -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", + } diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py index 458c7c42eb6..6970e34f759 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py @@ -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) diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_discovery.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_discovery.py index 94ca0dc11f5..6ced5264267 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_discovery.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_discovery.py @@ -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( diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py index 7bb74285ac6..e333bf1e3fe 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py @@ -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"] diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py index 0769568d6cd..c2a0d34a915 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py @@ -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(): diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index 5cbc3e72d83..8de42ca89da 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -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"] == [] diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 51f72f91dc3..867ef759fb3 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -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): """ diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 5202c8cbfc0..a658a4b7353 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -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"]}} diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index 46ecb31e1c8..ac01c6ae1d1 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -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" diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index 67945436987..71206687b5c 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -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 diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 536d24d4b4e..62f5ced7a39 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -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"], diff --git a/tests/test_litellm/proxy/test_proxy_types.py b/tests/test_litellm/proxy/test_proxy_types.py index c8e0b3a730a..5354de182a0 100644 --- a/tests/test_litellm/proxy/test_proxy_types.py +++ b/tests/test_litellm/proxy/test_proxy_types.py @@ -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" diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 276ee96ed65..e6ae1f85cfd 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -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", diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 6db19d63a01..46b5ce65c3f 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -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 diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index 66e2055f579..1606c177a73 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -152,7 +152,7 @@ "count": 1 }, "prefer-const": { - "count": 3 + "count": 1 }, "react-hooks/purity": { "count": 1 diff --git a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.test.tsx index e1d02e9352d..fd8dd011b05 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.test.tsx @@ -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(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.tsx b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.tsx index 95b73d1aacb..47c266ceac0 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.tsx @@ -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 = ({ accessToken, token, userRole, userID, premiumUser }) => { - const [filteredData, setFilteredData] = useState([]); const [selectedApiKeys, setSelectedApiKeys] = useState([]); const [selectedModels, setSelectedModels] = useState([]); - const [data, setData] = useState([]); - const [cachedResponses, setCachedResponses] = useState("0"); - const [cachedTokens, setCachedTokens] = useState("0"); - const [cacheHitRatio, setCacheHitRatio] = useState("0"); const [dateValue, setDateValue] = useState({ from: new Date(Date.now() - 7 * 24 * 60 * 60 * 1000), @@ -113,120 +104,24 @@ const CacheDashboard: React.FC = ({ accessToken, token, userRole const [lastRefreshed, setLastRefreshed] = useState(""); const [healthCheckResponse, setHealthCheckResponse] = useState(""); - 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 = ({ 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 = ({ accessToken, token, userRole value={dateValue} onValueChange={(value) => { setDateValue(value); - updateCachingData(value.from, value.to); }} /> @@ -404,12 +300,12 @@ const CacheDashboard: React.FC = ({ accessToken, token, userRole @@ -423,7 +319,7 @@ const CacheDashboard: React.FC = ({ accessToken, token, userRole ({ + $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); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/caching/useCacheActivity.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/caching/useCacheActivity.ts new file mode 100644 index 00000000000..af4486ad33b --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/caching/useCacheActivity.ts @@ -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) }, + ); +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/navigateWithParams.ts b/ui/litellm-dashboard/src/app/(dashboard)/navigateWithParams.ts index c49c73c3578..5acf444a359 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/navigateWithParams.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/navigateWithParams.ts @@ -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); + } } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.test.tsx index d381e5e65ca..3f9de478069 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.test.tsx @@ -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; +type OrganizationInfoViewProps = React.ComponentProps; + +let capturedTableProps: OrganizationsTableProps | null = null; vi.mock("./OrganizationsTable", () => ({ __esModule: true, - default: (props: { isLoading: boolean }) => ( -
isLoading:{String(props.isLoading)}
- ), + default: (props: OrganizationsTableProps) => { + capturedTableProps = props; + return
isLoading:{String(props.isLoading)}
; + }, })); +const mockOrgInfoView = vi.fn<(props: OrganizationInfoViewProps) => void>(); +vi.mock("@/components/organization/organization_view", () => ({ + __esModule: true, + default: (props: OrganizationInfoViewProps) => { + mockOrgInfoView(props); + return
; + }, +})); + +// 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) => { + 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({ui}); }; +beforeEach(() => { + capturedTableProps = null; + mockOrgInfoView.mockClear(); + window.history.replaceState(null, "", "/organizations/"); +}); + describe("OrganizationsPanel", () => { it("gates non-premium users behind the enterprise notice", () => { renderWithQueryClient(); @@ -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(); + + 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(); + + 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(); + + 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(); + + 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(); + + 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 }), + ); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.tsx index b1c026d3904..a21c0669677 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.tsx @@ -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 = ({ userRole, accessToken, premiumUser }) => { - const [selectedOrgId, setSelectedOrgId] = useState(null); + const { orgId: selectedOrgId, openOrg, close: closeOrgDetail } = useOrgDetailRouting(); const [editOrg, setEditOrg] = useState(false); const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false); const [orgToDelete, setOrgToDelete] = useState(null); @@ -108,7 +109,7 @@ const OrganizationsPanel: React.FC = ({ userRole, acces { - setSelectedOrgId(null); + closeOrgDetail(); setEditOrg(false); }} accessToken={accessToken} @@ -132,9 +133,12 @@ const OrganizationsPanel: React.FC = ({ 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} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/detailNavigation.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/organizations/detailNavigation.test.ts new file mode 100644 index 00000000000..46b7c4313ea --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/detailNavigation.test.ts @@ -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(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/detailNavigation.ts b/ui/litellm-dashboard/src/app/(dashboard)/organizations/detailNavigation.ts new file mode 100644 index 00000000000..8c55c7b750c --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/detailNavigation.ts @@ -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, + }; +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/teams/detailNavigation.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/teams/detailNavigation.test.ts new file mode 100644 index 00000000000..e5d5b1a4073 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/teams/detailNavigation.test.ts @@ -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(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/teams/detailNavigation.ts b/ui/litellm-dashboard/src/app/(dashboard)/teams/detailNavigation.ts new file mode 100644 index 00000000000..d5208f094cb --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/teams/detailNavigation.ts @@ -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, + }; +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx index 89c38c6274f..82ca66b10c0 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx @@ -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(); 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(); + + 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 () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx index e330983b6f9..4d44791d1a9 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx @@ -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 Promise> = { const EntityUsage: React.FC = ({ accessToken, entityType, entityId, entityList, dateValue }) => { const { teams } = useTeams(); const [selectedTags, setSelectedTags] = useState([]); + const [modelViewType, setModelViewType] = useState("groups"); const [topKeysLimit, setTopKeysLimit] = useState(5); const [topModelsLimit, setTopModelsLimit] = useState(5); const [topAgentsLimit, setTopAgentsLimit] = useState(5); @@ -153,14 +155,15 @@ const EntityUsage: React.FC = ({ 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 = ({ 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 = ( {/* Total Spend Card */} @@ -604,7 +609,10 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti {/* Top Models */} - {entityType === "agent" ? "Top Agents" : "Top Models"} +
+ {entityType === "agent" ? "Top Agents" : modelViewTitle} + +
= ({ accessToken, entityType, enti { key: "models", label: entityType === "agent" ? "Request / Token Consumption" : "Model Activity", - content: , + content: ( + <> +
+ +
+ + + ), }, ...(entityType === "team" ? [{ key: "agents", label: "Agent Activity", content: }] diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/ModelViewToggle.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/ModelViewToggle.tsx new file mode 100644 index 00000000000..0ee6dd3b19c --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/ModelViewToggle.tsx @@ -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 ( +
+ {MODEL_VIEW_OPTIONS.map((option) => ( + + ))} +
+ ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.test.tsx index cf122137f91..98dae51fa37 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.test.tsx @@ -30,8 +30,10 @@ vi.mock("@/components/networking", () => ({ // Mock child components to simplify testing vi.mock("@/components/activity_metrics", () => ({ - ActivityMetrics: () =>
Activity Metrics
, - processActivityData: () => ({ data: [], metadata: {} }), + ActivityMetrics: ({ modelMetrics }: { modelMetrics?: { __source?: string } }) => ( +
{`activity-source:${modelMetrics?.__source ?? "none"}`}
+ ), + 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(); + + 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(); + + 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", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx index d2f75609d18..46a17017d39 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx @@ -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 = ({ 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(isAdmin ? null : userID || null); - const [modelViewType, setModelViewType] = useState<"groups" | "individual">("groups"); + const [modelViewType, setModelViewType] = useState("groups"); const [isCloudZeroModalOpen, setIsCloudZeroModalOpen] = useState(false); const [isGlobalExportModalOpen, setIsGlobalExportModalOpen] = useState(false); const [isAiChatOpen, setIsAiChatOpen] = useState(false); @@ -438,7 +439,10 @@ const UsagePage: React.FC = ({ 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 = ({ teams, organizations }) => { value={topModelsLimit} onChange={(value) => setTopModelsLimit(value as number)} /> -
- - -
+
{loading ? ( @@ -839,6 +822,9 @@ const UsagePage: React.FC = ({ teams, organizations }) => { {/* Activity Panel */} +
+ +
diff --git a/ui/litellm-dashboard/src/components/Teams.test.tsx b/ui/litellm-dashboard/src/components/Teams.test.tsx index 7065b1a5fb6..742a88d864c 100644 --- a/ui/litellm-dashboard/src/components/Teams.test.tsx +++ b/ui/litellm-dashboard/src/components/Teams.test.tsx @@ -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) => { + 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(); + + 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(); + + 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(); + + 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(); diff --git a/ui/litellm-dashboard/src/components/Teams.tsx b/ui/litellm-dashboard/src/components/Teams.tsx index 20e9e78e7e4..0a3c7fc736d 100644 --- a/ui/litellm-dashboard/src/components/Teams.tsx +++ b/ui/litellm-dashboard/src/components/Teams.tsx @@ -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 = ({ accessToken, userID, userRole, premiumUser const [editModalVisible, setEditModalVisible] = useState(false); const [selectedTeam, setSelectedTeam] = useState(null); - const [selectedTeamId, setSelectedTeamId] = useState(null); + const { teamId: selectedTeamId, openTeam, close: closeTeamDetail } = useTeamDetailRouting(); const [editTeam, setEditTeam] = useState(false); const [isTeamModalVisible, setIsTeamModalVisible] = useState(false); @@ -482,12 +483,12 @@ const Teams: React.FC = ({ 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 = ({ 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} diff --git a/ui/litellm-dashboard/src/components/activity_metrics.test.tsx b/ui/litellm-dashboard/src/components/activity_metrics.test.tsx index f37643f2366..21d66991bb5 100644 --- a/ui/litellm-dashboard/src/components/activity_metrics.test.tsx +++ b/ui/litellm-dashboard/src/components/activity_metrics.test.tsx @@ -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: [ diff --git a/ui/litellm-dashboard/src/components/activity_metrics.tsx b/ui/litellm-dashboard/src/components/activity_metrics.tsx index 54ac5ae0ee6..b8199ee29b6 100644 --- a/ui/litellm-dashboard/src/components/activity_metrics.tsx +++ b/ui/litellm-dashboard/src/components/activity_metrics.tsx @@ -362,7 +362,7 @@ export const formatKeyLabel = (modelData: KeyMetricWithMetadata, model: string, // Process data function export const processActivityData = ( dailyActivity: { results: DailyData[] }, - key: "models" | "api_keys" | "mcp_servers" | "entities", + key: "models" | "model_groups" | "api_keys" | "mcp_servers" | "entities", teams: Team[] = [], ): Record => { const modelMetrics: Record = {}; diff --git a/ui/litellm-dashboard/src/components/leftnav.test.tsx b/ui/litellm-dashboard/src/components/leftnav.test.tsx index 2c8fe52f97f..87c830e69de 100644 --- a/ui/litellm-dashboard/src/components/leftnav.test.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.test.tsx @@ -80,6 +80,12 @@ describe("Sidebar (leftnav)", () => { collapsed: false, }; + it("should link the logo to the UI home route rather than the proxy origin", () => { + renderWithProviders(); + + expect(screen.getByRole("link", { name: /litellm home/i })).toHaveAttribute("href", "/ui"); + }); + it("renders all top-level (non-nested) tabs for admin", () => { renderWithProviders(); diff --git a/ui/litellm-dashboard/src/components/leftnav.tsx b/ui/litellm-dashboard/src/components/leftnav.tsx index 4aa15f5fc28..af76fccf9eb 100644 --- a/ui/litellm-dashboard/src/components/leftnav.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.tsx @@ -582,7 +582,7 @@ const Sidebar_: React.FC = ({
- + LiteLLM { expect(screen.getByRole("button", { name: /open account menu/i })).toBeInTheDocument(); }); + it("should link the logo to the UI home route rather than the proxy origin", () => { + renderWithProviders(); + + expect(screen.getByRole("link", { name: /litellm brand/i })).toHaveAttribute("href", "/ui"); + }); + it("should display user information in dropdown", async () => { const user = userEvent.setup(); renderWithProviders(); diff --git a/ui/litellm-dashboard/src/components/navbar.tsx b/ui/litellm-dashboard/src/components/navbar.tsx index 40638e7b8ba..c999ee8035e 100644 --- a/ui/litellm-dashboard/src/components/navbar.tsx +++ b/ui/litellm-dashboard/src/components/navbar.tsx @@ -3,6 +3,7 @@ import { useDisableBouncingIcon } from "@/app/(dashboard)/hooks/useDisableBounci import { useDisableShowPrompts } from "@/app/(dashboard)/hooks/useDisableShowPrompts"; import { useWorker } from "@/hooks/useWorker"; import { getProxyBaseUrl } from "@/components/networking"; +import { migratedHref } from "@/utils/migratedPages"; import { useTheme } from "@/contexts/ThemeContext"; import { clearTokenCookies } from "@/utils/cookieUtils"; import { clearStoredReturnUrl, getLoginUrl } from "@/utils/returnUrlUtils"; @@ -75,7 +76,7 @@ const Navbar: React.FC = ({ )}
- +
{ - try { - let url = proxyBaseUrl ? `${proxyBaseUrl}/global/activity/cache_hits` : `/global/activity/cache_hits`; - - if (startTime && endTime) { - url += `?start_date=${startTime}&end_date=${endTime}`; - } - - const requestOptions = { - method: "GET", - headers: { - [globalLitellmHeaderName]: `Bearer ${accessToken}`, - }, - }; - - const response = await fetch(url, requestOptions); - - if (!response.ok) { - const errorData = await response.json(); - const errorMessage = deriveErrorMessage(errorData); - handleError(errorMessage); - throw new Error(errorMessage); - } - - const data = await response.json(); - return data; - } catch (error) { - console.error("Failed to fetch spend data:", error); - throw error; - } -}; - export const adminGlobalActivityPerModel = async ( accessToken: string, startTime: string | undefined, diff --git a/ui/litellm-dashboard/src/components/object_permissions_view.tsx b/ui/litellm-dashboard/src/components/object_permissions_view.tsx index 687d1a5a846..327e127e8fe 100644 --- a/ui/litellm-dashboard/src/components/object_permissions_view.tsx +++ b/ui/litellm-dashboard/src/components/object_permissions_view.tsx @@ -28,7 +28,7 @@ export function ObjectPermissionsView({ const searchTools = objectPermission?.search_tools || []; const content = ( -
+
+
Object Permissions diff --git a/ui/litellm-dashboard/src/components/organization/organization_view.test.tsx b/ui/litellm-dashboard/src/components/organization/organization_view.test.tsx index 7fed639c491..ec7baecad2d 100644 --- a/ui/litellm-dashboard/src/components/organization/organization_view.test.tsx +++ b/ui/litellm-dashboard/src/components/organization/organization_view.test.tsx @@ -6,7 +6,14 @@ import { renderWithProviders } from "../../../tests/test-utils"; import OrganizationInfoView from "./organization_view"; import { useOrganization } from "@/app/(dashboard)/hooks/organizations/useOrganizations"; -// Mock networking calls used by the component's mutation handlers +vi.mock("next/navigation", () => ({ + useRouter: () => ({ push: vi.fn() }), + usePathname: () => "/organizations", + useSearchParams: () => new URLSearchParams(window.location.search), +})); + +// Mock networking calls used by the component's mutation handlers. entityLinks -> migratedPages +// imports serverRootPath from the same module, so the mock must export it too. vi.mock("../networking", () => { return { __esModule: true, @@ -14,6 +21,7 @@ vi.mock("../networking", () => { organizationMemberUpdateCall: vi.fn(), organizationMemberDeleteCall: vi.fn(), organizationUpdateCall: vi.fn(), + serverRootPath: "", }; }); @@ -206,6 +214,58 @@ test("should display team ID as fallback when alias is not found", async () => { }); }); +test("links each team badge to that team's detail page", async () => { + const orgWithTeams = { + ...mockOrg, + teams: [{ team_id: "team_123" }, { team_id: "team_456" }], + }; + mockUseOrganization.mockReturnValue({ data: orgWithTeams, isLoading: false } as any); + + renderWithProviders( + {}} + accessToken="test-token" + is_org_admin={false} + is_proxy_admin={false} + userModels={[]} + editOrg={false} + />, + ); + + await waitFor(() => { + expect(screen.getByRole("link", { name: "Engineering Team" })).toHaveAttribute( + "href", + expect.stringContaining("/teams?team=team_123"), + ); + expect(screen.getByRole("link", { name: "Marketing Team" })).toHaveAttribute( + "href", + expect.stringContaining("/teams?team=team_456"), + ); + }); +}); + +test("model badges stay non-clickable", async () => { + mockUseOrganization.mockReturnValue({ data: mockOrg, isLoading: false } as any); + + renderWithProviders( + {}} + accessToken="test-token" + is_org_admin={false} + is_proxy_admin={false} + userModels={[]} + editOrg={false} + />, + ); + + await waitFor(() => { + expect(screen.getByText("gpt-4o-mini")).toBeInTheDocument(); + }); + expect(screen.queryByRole("link", { name: "gpt-4o-mini" })).not.toBeInTheDocument(); +}); + test("should keep unsaved settings edits when switching tabs and back", async () => { mockUseOrganization.mockReturnValue({ data: mockOrg, isLoading: false } as any); diff --git a/ui/litellm-dashboard/src/components/organization/organization_view.tsx b/ui/litellm-dashboard/src/components/organization/organization_view.tsx index e96296b1b39..e9c502fa14f 100644 --- a/ui/litellm-dashboard/src/components/organization/organization_view.tsx +++ b/ui/litellm-dashboard/src/components/organization/organization_view.tsx @@ -4,12 +4,13 @@ import { useQueryClient } from "@tanstack/react-query"; import { useVisitedTabs } from "@/hooks/useVisitedTabs"; import { MoneyCell } from "@/components/shared/table_cells"; import CopyButton from "@/components/shared/CopyButton"; -import { Badge } from "@/components/ui/badge"; import { Button } from "@/components/ui/button"; import { Card, CardContent } from "@/components/ui/card"; import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; import { formatNumberWithCommas } from "@/utils/dataUtils"; +import { teamDetailHref } from "@/utils/entityLinks"; import { createTeamAliasMap } from "@/utils/teamUtils"; +import { BadgeLink } from "@/components/shared/BadgeLink"; import type { ColumnsType } from "antd/es/table"; import { ArrowLeft } from "lucide-react"; import React, { useMemo, useState } from "react"; @@ -225,13 +226,9 @@ const OrganizationInfoView: React.FC = ({

Models

{orgData.models.length === 0 ? ( - All proxy models + All proxy models ) : ( - orgData.models.map((model, index) => ( - - {model} - - )) + orgData.models.map((model, index) => {model}) )}
@@ -242,9 +239,9 @@ const OrganizationInfoView: React.FC = ({

Teams

{orgData.teams?.map((team, index) => ( - + {teamAliasMap[team.team_id] || team.team_id} - + ))}
@@ -331,9 +328,7 @@ const OrganizationInfoView: React.FC = ({

Models

{orgData.models.map((model, index) => ( - - {model} - + {model} ))}
diff --git a/ui/litellm-dashboard/src/components/shared/BadgeLink.test.tsx b/ui/litellm-dashboard/src/components/shared/BadgeLink.test.tsx new file mode 100644 index 00000000000..10a192c8af0 --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/BadgeLink.test.tsx @@ -0,0 +1,42 @@ +/* @vitest-environment jsdom */ +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + +import { BadgeLink } from "./BadgeLink"; + +const push = vi.fn(); +vi.mock("next/navigation", () => ({ useRouter: () => ({ push }) })); + +describe("BadgeLink", () => { + beforeEach(() => { + push.mockClear(); + }); + + it("renders an anchor pointing at the target href", () => { + render(My Team); + expect(screen.getByRole("link", { name: "My Team" })).toHaveAttribute("href", "/ui/teams?team=t1"); + }); + + it("navigates client-side on plain click", async () => { + const user = userEvent.setup(); + render(My Team); + await user.click(screen.getByRole("link", { name: "My Team" })); + expect(push).toHaveBeenCalledWith("/ui/teams?team=t1"); + }); + + it("leaves modified clicks to the browser so new-tab shortcuts keep working", async () => { + const user = userEvent.setup(); + render(My Team); + await user.keyboard("{Meta>}"); + await user.click(screen.getByRole("link", { name: "My Team" })); + await user.keyboard("{/Meta}"); + expect(push).not.toHaveBeenCalled(); + }); + + it("renders a plain same-sized badge when no href is given", () => { + render(all-proxy-models); + expect(screen.getByText("all-proxy-models")).toBeInTheDocument(); + expect(screen.queryByRole("link", { name: "all-proxy-models" })).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/shared/BadgeLink.tsx b/ui/litellm-dashboard/src/components/shared/BadgeLink.tsx new file mode 100644 index 00000000000..444d2acdcfe --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/BadgeLink.tsx @@ -0,0 +1,46 @@ +"use client"; + +import { useRouter } from "next/navigation"; +import * as React from "react"; + +import { Badge } from "@/components/ui/badge"; +import { cn } from "@/lib/cva.config"; + +const ENTITY_BADGE_SIZE = "px-2.5 py-1 text-sm"; + +interface BadgeLinkProps { + href?: string; + variant?: React.ComponentProps["variant"]; + className?: string; + children: React.ReactNode; +} + +export function BadgeLink({ href, variant = "secondary", className, children }: BadgeLinkProps) { + const router = useRouter(); + + if (!href) { + return ( + + {children} + + ); + } + + const handleClick = (e: React.MouseEvent) => { + const hasModifierKey = e.metaKey || e.ctrlKey || e.shiftKey; + const isNativeNewTabClick = hasModifierKey || e.button === 1; + if (isNativeNewTabClick) return; + e.preventDefault(); + router.push(href); + }; + + return ( + } + > + {children} + + ); +} diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx index c41108c295e..d09735274d1 100644 --- a/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx @@ -445,6 +445,29 @@ describe("TeamInfoView", () => { }); }); + it("shows edit tabs when the fetched team data marks the session user as team admin, even without the is_team_admin prop", async () => { + vi.mocked(networking.teamInfoCall).mockResolvedValue( + createMockTeamData({ + members_with_roles: [ + { + user_id: "user-1", + user_email: "admin@test.com", + role: "admin", + spend: 0, + budget_id: "budget1", + }, + ], + }), + ); + + renderWithProviders(); + + await waitFor(() => { + expect(screen.getByRole("tab", { name: "Settings" })).toBeInTheDocument(); + }); + expect(screen.getByRole("tab", { name: "Members" })).toBeInTheDocument(); + }); + it("should navigate to settings tab when clicked", async () => { const user = userEvent.setup({ delay: null }); vi.mocked(networking.teamInfoCall).mockResolvedValue(createMockTeamData()); diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx index ff384f6e85d..1b23802a189 100644 --- a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx @@ -226,7 +226,15 @@ const TeamInfoView: React.FC = ({ return unfurlWildcardModelsInList(selected, userModels); }, [selectedModelsInForm, teamData, userModels]); - const canEditTeam = is_team_admin || is_proxy_admin || is_org_admin || isOrgAdminForTeam; + const isTeamAdminFromTeamData = useMemo( + () => + teamData?.team_info?.members_with_roles?.some( + (member) => member.user_id != null && member.user_id === userId && member.role === "admin", + ) ?? false, + [teamData, userId], + ); + + const canEditTeam = is_team_admin || is_proxy_admin || is_org_admin || isOrgAdminForTeam || isTeamAdminFromTeamData; // Destinations that will receive this team's traces, resolved server-side by // /team/info from credential_info.access. Names only, visible to every team viewer. diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.test.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.test.tsx index e8e2fae3f4d..49870e51c2b 100644 --- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.test.tsx @@ -1,7 +1,7 @@ import { screen, waitFor, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import moment from "moment"; -import { beforeEach, describe, expect, it, vi } from "vitest"; +import { afterAll, beforeAll, beforeEach, describe, expect, it, vi } from "vitest"; import { renderWithProviders, testQueryClient } from "../../../tests/test-utils"; import type { LogEntry } from "./columns"; @@ -22,11 +22,75 @@ vi.mock("@/components/key_team_helpers/filter_helpers", () => ({ })); vi.mock("./LogDetailsDrawer", () => ({ - LogDetailsDrawer: function LogDetailsDrawerMock({ open }: { open: boolean }) { - return
{open ? "open" : "closed"}
; + LogDetailsDrawer: function LogDetailsDrawerMock({ + open, + logEntry, + sessionId, + onClose, + allLogs = [], + onSelectLog, + }: { + open: boolean; + logEntry?: { request_id: string } | null; + sessionId?: string | null; + onClose: () => void; + allLogs?: { request_id: string }[]; + onSelectLog?: (log: { request_id: string }) => void; + }) { + const nextLog = allLogs.find((log) => log.request_id !== logEntry?.request_id); + return ( +
+ {open ? "open" : "closed"} + + +
+ ); }, })); +vi.mock("next/navigation", async (importOriginal) => { + const actual = await importOriginal(); + const { useSyncExternalStore } = await import("react"); + return { + ...actual, + useSearchParams: () => { + const search = useSyncExternalStore( + (onChange: () => void) => { + window.addEventListener("test-locationchange", onChange); + window.addEventListener("popstate", onChange); + return () => { + window.removeEventListener("test-locationchange", onChange); + window.removeEventListener("popstate", onChange); + }; + }, + () => window.location.search, + ); + return new URLSearchParams(search); + }, + }; +}); + +const originalPushState = window.history.pushState.bind(window.history); +const originalReplaceState = window.history.replaceState.bind(window.history); +beforeAll(() => { + window.history.pushState = (data, unused, url) => { + originalPushState(data, unused, url); + window.dispatchEvent(new Event("test-locationchange")); + }; + window.history.replaceState = (data, unused, url) => { + originalReplaceState(data, unused, url); + window.dispatchEvent(new Event("test-locationchange")); + }; +}); +afterAll(() => { + window.history.pushState = originalPushState; + window.history.replaceState = originalReplaceState; +}); + import { uiSpendLogsCall } from "../networking"; const logEntry = (overrides: Partial): LogEntry => ({ @@ -73,6 +137,7 @@ describe("RequestLogsPanel", () => { vi.clearAllMocks(); sessionStorage.clear(); testQueryClient.clear(); + window.history.replaceState(null, "", "/logs/"); respondWith([]); }); @@ -185,6 +250,161 @@ describe("RequestLogsPanel", () => { }); }); + describe("shareable log links (?log_id=)", () => { + const drawer = () => screen.getByTestId("log-details-drawer"); + + it("clicking a row writes ?log_id= to the URL and opens the drawer", async () => { + const user = userEvent.setup(); + respondWith([logEntry({ request_id: "req-1" })]); + renderWithProviders(); + + await waitFor(() => expect(row("req-1")).not.toBeNull()); + await user.click(row("req-1") as HTMLElement); + + expect(new URLSearchParams(window.location.search).get("log_id")).toBe("req-1"); + await waitFor(() => { + expect(drawer()).toHaveTextContent("open"); + expect(drawer()).toHaveAttribute("data-log-id", "req-1"); + }); + }); + + it("opens the drawer on load when ?log_id= matches a log in the loaded page", async () => { + window.history.replaceState(null, "", "/logs/?log_id=req-2"); + respondWith([logEntry({ request_id: "req-1" }), logEntry({ request_id: "req-2" })]); + renderWithProviders(); + + await waitFor(() => { + expect(drawer()).toHaveTextContent("open"); + expect(drawer()).toHaveAttribute("data-log-id", "req-2"); + }); + }); + + it("fetches the log by request_id and opens the drawer when it is not in the loaded page", async () => { + window.history.replaceState(null, "", "/logs/?log_id=req-old"); + vi.mocked(uiSpendLogsCall).mockImplementation(async ({ params }) => + params?.request_id === "req-old" + ? { data: [logEntry({ request_id: "req-old" })], total: 1, page: 1, page_size: 1, total_pages: 1 } + : { data: [], total: 0, page: 1, page_size: 50, total_pages: 0 }, + ); + renderWithProviders(); + + await waitFor(() => { + expect(drawer()).toHaveTextContent("open"); + expect(drawer()).toHaveAttribute("data-log-id", "req-old"); + }); + + const byIdCall = vi + .mocked(uiSpendLogsCall) + .mock.calls.find(([options]) => options.params?.request_id === "req-old")?.[0]; + if (!byIdCall) throw new Error("expected a by-id uiSpendLogsCall"); + expect(byIdCall.page).toBe(1); + expect(byIdCall.page_size).toBe(1); + }); + + it("closing the drawer removes ?log_id= from the URL and closes the drawer", async () => { + const user = userEvent.setup(); + respondWith([logEntry({ request_id: "req-1" })]); + renderWithProviders(); + + await waitFor(() => expect(row("req-1")).not.toBeNull()); + await user.click(row("req-1") as HTMLElement); + await waitFor(() => expect(drawer()).toHaveTextContent("open")); + + await user.click(screen.getByRole("button", { name: "close-drawer" })); + + expect(new URLSearchParams(window.location.search).get("log_id")).toBeNull(); + await waitFor(() => expect(drawer()).toHaveTextContent("closed")); + }); + + it("switching logs inside the drawer replaces the URL, so back closes the drawer in one step", async () => { + const user = userEvent.setup(); + respondWith([logEntry({ request_id: "req-1" }), logEntry({ request_id: "req-2" })]); + renderWithProviders(); + + await waitFor(() => expect(row("req-1")).not.toBeNull()); + await user.click(row("req-1") as HTMLElement); + await waitFor(() => expect(drawer()).toHaveAttribute("data-log-id", "req-1")); + + await user.click(screen.getByRole("button", { name: "select-next-log" })); + await waitFor(() => expect(drawer()).toHaveAttribute("data-log-id", "req-2")); + expect(new URLSearchParams(window.location.search).get("log_id")).toBe("req-2"); + + window.history.back(); + + await waitFor(() => expect(drawer()).toHaveTextContent("closed")); + expect(new URLSearchParams(window.location.search).get("log_id")).toBeNull(); + }); + + it("clicking a session id writes ?session_id= and ?log_id= and opens the session drawer", async () => { + const user = userEvent.setup(); + respondWith([logEntry({ request_id: "req-solo", session_id: "sess-solo", session_total_count: 1 })]); + renderWithProviders(); + + await waitFor(() => expect(row("req-solo")).not.toBeNull()); + await user.click(within(row("req-solo") as HTMLElement).getByText("sess-solo")); + + const params = new URLSearchParams(window.location.search); + expect(params.get("session_id")).toBe("sess-solo"); + expect(params.get("log_id")).toBe("req-solo"); + await waitFor(() => { + expect(drawer()).toHaveTextContent("open"); + expect(drawer()).toHaveAttribute("data-session-id", "sess-solo"); + }); + }); + + it("clicking a log row clears a lingering ?session_id= so the drawer shows the clicked log", async () => { + const user = userEvent.setup(); + respondWith([ + logEntry({ request_id: "req-a", session_id: "sess-a", session_total_count: 1 }), + logEntry({ request_id: "req-b" }), + ]); + renderWithProviders(); + + await waitFor(() => expect(row("req-a")).not.toBeNull()); + await user.click(within(row("req-a") as HTMLElement).getByText("sess-a")); + await waitFor(() => expect(new URLSearchParams(window.location.search).get("session_id")).toBe("sess-a")); + + await user.click(row("req-b") as HTMLElement); + + const params = new URLSearchParams(window.location.search); + expect(params.get("log_id")).toBe("req-b"); + expect(params.get("session_id")).toBeNull(); + await waitFor(() => { + expect(drawer()).toHaveAttribute("data-log-id", "req-b"); + expect(drawer()).toHaveAttribute("data-session-id", ""); + }); + }); + + it("browser back after opening via a session id closes the drawer", async () => { + const user = userEvent.setup(); + respondWith([logEntry({ request_id: "req-solo", session_id: "sess-solo", session_total_count: 1 })]); + renderWithProviders(); + + await waitFor(() => expect(row("req-solo")).not.toBeNull()); + await user.click(within(row("req-solo") as HTMLElement).getByText("sess-solo")); + await waitFor(() => expect(drawer()).toHaveTextContent("open")); + + window.history.back(); + + await waitFor(() => expect(drawer()).toHaveTextContent("closed")); + expect(new URLSearchParams(window.location.search).get("session_id")).toBeNull(); + }); + + it("opens a deep-linked multi-call session log in session mode", async () => { + window.history.replaceState(null, "", "/logs/?log_id=req-llm"); + respondWith([ + logEntry({ request_id: "req-llm", call_type: "acompletion", session_id: "sess-1", session_total_count: 3 }), + ]); + renderWithProviders(); + + await waitFor(() => { + expect(drawer()).toHaveTextContent("open"); + expect(drawer()).toHaveAttribute("data-log-id", "req-llm"); + expect(drawer()).toHaveAttribute("data-session-id", "sess-1"); + }); + }); + }); + describe("live tail", () => { it("shows the auto-refresh banner on the first page and hides it once stopped", async () => { const user = userEvent.setup(); diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.tsx index b3b1e8c0640..06c8ca26a7e 100644 --- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.tsx @@ -8,7 +8,7 @@ import { useCallback, useEffect, useMemo, useState } from "react"; import { AutoRouterModelGroupsProvider } from "@/components/shared/table_cells"; import { internalUserRoles } from "../../utils/roles"; import type { KeyResponse } from "../key_team_helpers/key_list"; -import { keyInfoV1Call } from "../networking"; +import { keyInfoV1Call, uiSpendLogsCall } from "../networking"; import KeyInfoView from "../templates/key_info_view"; import type { LogEntry } from "./columns"; import { AGENT_CALL_TYPES, MCP_CALL_TYPES } from "./constants"; @@ -17,8 +17,10 @@ import { formatLogsWindow, getLogsWindowEndBound, LOG_FILTER_IDS, + type PaginatedResponse, useLogFilterLogic, } from "./log_filter_logic"; +import { useLogDetailRouting } from "./logDetailRouting"; import { LogDetailsDrawer } from "./LogDetailsDrawer"; import { LiveTailBanner, LogsTableToolbar } from "./LogsTableToolbar"; import { RequestLogsTable } from "./RequestLogsTable"; @@ -52,8 +54,15 @@ export default function RequestLogsPanel({ accessToken, token, userRole, userID, const [selectedKeyIdInfoView, setSelectedKeyIdInfoView] = useState(null); const [selectedLog, setSelectedLog] = useState(null); - const [isDrawerOpen, setIsDrawerOpen] = useState(false); - const [selectedSessionId, setSelectedSessionId] = useState(null); + + const { + logId: urlLogId, + sessionId: urlSessionId, + openLog, + openSession, + selectLog, + close: closeUrlLog, + } = useLogDetailRouting(); const [isLiveTail, setIsLiveTail] = useState(() => { const storedValue = sessionStorage.getItem("isLiveTail"); @@ -106,6 +115,43 @@ export default function RequestLogsPanel({ accessToken, token, userRole, userID, const { data: selectedKeyInfo } = useQuery(keyInfoQueryOptions); + const urlLogQueryOptions: UseQueryOptions = { + queryKey: ["logs", "byId", urlLogId, accessToken], + queryFn: async () => { + if (urlLogId === null) return null; + const window = formatLogsWindow(startTime, endTime, isCustomDate); + const response: PaginatedResponse = await uiSpendLogsCall({ + accessToken, + start_date: window.start_date, + end_date: window.end_date, + page: 1, + page_size: 1, + params: { request_id: urlLogId }, + }); + return response.data.find((log) => log.request_id === urlLogId) ?? null; + }, + enabled: urlLogId !== null && selectedLog?.request_id !== urlLogId, + staleTime: Infinity, + }; + + const { data: urlLog } = useQuery(urlLogQueryOptions); + + const displayLog = useMemo(() => { + if (urlLogId === null) return null; + if (selectedLog?.request_id === urlLogId) return selectedLog; + return filteredLogs.data.find((log) => log.request_id === urlLogId) ?? urlLog ?? null; + }, [urlLogId, selectedLog, filteredLogs.data, urlLog]); + + const displaySessionId = useMemo(() => { + if (urlSessionId !== null) return urlSessionId; + if (displayLog?.session_id !== undefined && (displayLog.session_total_count || 1) > 1) { + return displayLog.session_id; + } + return null; + }, [urlSessionId, displayLog]); + + const isDrawerOpen = displayLog !== null || displaySessionId !== null; + const rows = useMemo(() => { const searchedLogs = filteredLogs.data; @@ -186,22 +232,30 @@ export default function RequestLogsPanel({ accessToken, token, userRole, userID, resetToFirstPage(); }, [resetToFirstPage]); - const handleRowClick = useCallback((log: LogEntry) => { - const isMultiCallSession = log.session_id !== undefined && (log.session_total_count || 1) > 1; - setSelectedSessionId(isMultiCallSession ? log.session_id ?? null : null); - setSelectedLog(log); - setIsDrawerOpen(true); - }, []); + const handleRowClick = useCallback( + (log: LogEntry) => { + setSelectedLog(log); + openLog(log.request_id); + }, + [openLog], + ); const handleSessionClick = useCallback( (sessionId: string) => { if (!sessionId) return; const log = rows.find((candidate) => candidate.session_id === sessionId) ?? null; - setSelectedSessionId(sessionId); setSelectedLog(log); - setIsDrawerOpen(true); + openSession(sessionId, log?.request_id ?? null); }, - [rows], + [rows, openSession], + ); + + const handleSelectLog = useCallback( + (log: LogEntry) => { + setSelectedLog(log); + selectLog(log.request_id); + }, + [selectLog], ); const handleKeyHashClick = useCallback((keyHash: string) => { @@ -267,15 +321,12 @@ export default function RequestLogsPanel({ accessToken, token, userRole, userID, { - setIsDrawerOpen(false); - setSelectedSessionId(null); - }} - logEntry={selectedLog} - sessionId={selectedSessionId} + onClose={closeUrlLog} + logEntry={displayLog} + sessionId={displaySessionId} accessToken={accessToken} allLogs={rows} - onSelectLog={setSelectedLog} + onSelectLog={handleSelectLog} startTime={moment(startTime).utc().format("YYYY-MM-DD HH:mm:ss")} /> diff --git a/ui/litellm-dashboard/src/components/view_logs/logDetailRouting.ts b/ui/litellm-dashboard/src/components/view_logs/logDetailRouting.ts new file mode 100644 index 00000000000..5b311c94627 --- /dev/null +++ b/ui/litellm-dashboard/src/components/view_logs/logDetailRouting.ts @@ -0,0 +1,60 @@ +import { useSearchParams } from "next/navigation"; +import { useCallback } from "react"; + +import { navigateWithParams } from "@/app/(dashboard)/navigateWithParams"; + +export const LOG_ID_QUERY_PARAM = "log_id"; +export const SESSION_ID_QUERY_PARAM = "session_id"; + +export interface LogDetailRouting { + logId: string | null; + sessionId: string | null; + openLog: (requestId: string) => void; + openSession: (sessionId: string, requestId: string | null) => void; + selectLog: (requestId: string) => void; + close: () => void; +} + +export function useLogDetailRouting(): LogDetailRouting { + const searchParams = useSearchParams(); + + const openLog = useCallback((requestId: string) => { + navigateWithParams((params) => { + params.set(LOG_ID_QUERY_PARAM, requestId); + params.delete(SESSION_ID_QUERY_PARAM); + }); + }, []); + + const openSession = useCallback((sessionId: string, requestId: string | null) => { + navigateWithParams((params) => { + params.set(SESSION_ID_QUERY_PARAM, sessionId); + if (requestId === null) { + params.delete(LOG_ID_QUERY_PARAM); + } else { + params.set(LOG_ID_QUERY_PARAM, requestId); + } + }); + }, []); + + const selectLog = useCallback((requestId: string) => { + navigateWithParams((params) => { + params.set(LOG_ID_QUERY_PARAM, requestId); + }, "replace"); + }, []); + + const close = useCallback(() => { + navigateWithParams((params) => { + params.delete(LOG_ID_QUERY_PARAM); + params.delete(SESSION_ID_QUERY_PARAM); + }); + }, []); + + return { + logId: searchParams?.get(LOG_ID_QUERY_PARAM) ?? null, + sessionId: searchParams?.get(SESSION_ID_QUERY_PARAM) ?? null, + openLog, + openSession, + selectLog, + close, + }; +} diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index c3b66d53883..ca9285ca27b 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -2484,11 +2484,6 @@ export interface paths { /** * Get Credentials * @description [BETA] endpoint. This might change unexpectedly. - * - * Proxy-admin only (a proxy-admin-viewer may read). Credentials, including - * admin-owned logging destinations, are managed exclusively by the proxy admin; - * tenants never read them over the API. Secret values are masked for both - * admin-tier readers, exactly as they were before this feature. */ get: operations["get_credentials_credentials_get"]; put?: never; @@ -2611,9 +2606,6 @@ export interface paths { /** * Update Credential * @description [BETA] endpoint. This might change unexpectedly. - * - * Proxy-admin only. Credentials, including admin-owned logging destinations and - * their ``access`` scoping, are managed exclusively by the proxy admin. */ patch: operations["update_credential_credentials__credential_name__patch"]; trace?: never; @@ -4364,25 +4356,9 @@ export interface paths { }; /** * Get Global Activity - * @description 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 - * } + * @description 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. */ get: operations["get_global_activity_global_activity_cache_hits_get"]; put?: never; @@ -6943,8 +6919,8 @@ export interface paths { * @description 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. @@ -21750,6 +21726,48 @@ export interface components { /** Total Requested */ total_requested: number; }; + /** CacheActivityFilterOptions */ + CacheActivityFilterOptions: { + /** Key Aliases */ + key_aliases: string[]; + /** Models */ + models: string[]; + }; + /** CacheActivityGroup */ + CacheActivityGroup: { + /** Api Requests */ + api_requests: number; + /** Cache Hits */ + cache_hits: number; + /** Cached Completion Tokens */ + cached_completion_tokens: number; + /** Call Type */ + call_type: string; + /** Failed Requests */ + failed_requests: number; + /** Generated Completion Tokens */ + generated_completion_tokens: number; + }; + /** CacheActivityResponse */ + CacheActivityResponse: { + filter_options: components["schemas"]["CacheActivityFilterOptions"]; + /** Groups */ + groups: components["schemas"]["CacheActivityGroup"][]; + totals: components["schemas"]["CacheActivityTotals"]; + }; + /** CacheActivityTotals */ + CacheActivityTotals: { + /** Api Requests */ + api_requests: number; + /** Cache Hit Ratio */ + cache_hit_ratio: number; + /** Cache Hits */ + cache_hits: number; + /** Cached Completion Tokens */ + cached_completion_tokens: number; + /** Failed Requests */ + failed_requests: number; + }; /** CachePingResponse */ CachePingResponse: { /** Cache Type */ @@ -32468,7 +32486,7 @@ export interface components { /** Guardrails */ guardrails?: string[] | null; /** Key */ - key: string; + key?: string | null; /** Key Alias */ key_alias?: string | null; /** Max Budget */ @@ -40644,11 +40662,15 @@ export interface operations { }; get_global_activity_global_activity_cache_hits_get: { parameters: { - query?: { + query: { /** @description Time from which to start viewing spend */ - start_date?: string | null; + start_date: string; /** @description Time till which to view spend */ - end_date?: string | null; + end_date: string; + /** @description Only include spend from these key aliases */ + key_aliases?: string[] | null; + /** @description Only include spend for these models */ + models?: string[] | null; }; header?: never; path?: never; @@ -40662,7 +40684,7 @@ export interface operations { [name: string]: unknown; }; content: { - "application/json": components["schemas"]["LiteLLM_SpendLogs"][]; + "application/json": components["schemas"]["CacheActivityResponse"]; }; }; /** @description Validation Error */ diff --git a/ui/litellm-dashboard/src/utils/entityLinks.ts b/ui/litellm-dashboard/src/utils/entityLinks.ts new file mode 100644 index 00000000000..2659a307866 --- /dev/null +++ b/ui/litellm-dashboard/src/utils/entityLinks.ts @@ -0,0 +1,5 @@ +import { migratedHref } from "@/utils/migratedPages"; + +export function teamDetailHref(teamId: string): string { + return `${migratedHref("teams")}?team=${encodeURIComponent(teamId)}`; +}