mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
merge(litellm_internal_staging): resolve otel/team/org conflicts
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
commit
5ad0aeeca2
118 changed files with 5711 additions and 1243 deletions
1
.gitignore
vendored
1
.gitignore
vendored
|
|
@ -141,3 +141,4 @@ crash.*.log
|
|||
.coverage
|
||||
|
||||
ui/litellm-dashboard/out/
|
||||
litellm.log
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1,10 +1,11 @@
|
|||
# Adding a provider / route to litellm-rust
|
||||
|
||||
Three layers, same for every route (see `ocr` and `realtime` as references):
|
||||
Everything for a route lives in `crates/core/src/<route>/`; `crates/core/src/messages` is the reference. A host (the axum gateway, the Python bridge) only calls the route's entrypoint.
|
||||
|
||||
1. **Transform contract (pure)** — `crates/core/src/<route>/transformation.rs`: a `…ProviderConfig` trait (URL build + request/response transforms) + types in `types.rs`. No network, env, or auth.
|
||||
2. **Provider config (pure)** — `crates/providers/src/<provider>/<route>/transformation.rs`: implement that trait as a `const <PROVIDER>_<ROUTE>_CONFIG`, mirroring the Python provider tree. Add parity unit tests.
|
||||
3. **HTTP / transport (the host)** — `crates/providers/src/<route>.rs` (e.g. `ocr.rs`, `realtime.rs`): the callable fn (`run_ocr`, `realtime`). It resolves the key, builds the auth header, builds URL + transforms via the config, then does the network call. This is the only layer allowed to do I/O.
|
||||
1. **Entrypoint** — `mod.rs`: `pub async fn <route>(request) -> CoreResult<Response>`, the Rust equivalent of `litellm.<route>()`, plus a `<route>_stream` variant when the route streams. It is the only thing a host touches.
|
||||
2. **Transform contract** — `transformation.rs`: a `…ProviderConfig` trait (URL build + request/response transforms) with types in `types.rs`.
|
||||
3. **Provider config** — `crates/core/src/providers/<provider>/<route>/transformation.rs`: implement that trait as a `const <PROVIDER>_<ROUTE>_CONFIG`, mirroring the Python provider tree. Add parity unit tests.
|
||||
4. **Prepare + handler** — `prepare.rs` resolves provider/model, credentials, auth headers, and URL, then transforms the request; `handler.rs` performs the provider call through the shared client in `client.rs` and transforms the response.
|
||||
|
||||
## Coding standards
|
||||
|
||||
|
|
@ -25,4 +26,4 @@ variants of it. The test for a good abstraction is that adding the next provider
|
|||
is a few declarative lines, not a new file of duplicated flow. Only diverge from
|
||||
the base when behavior is genuinely different, and say so explicitly in the PR.
|
||||
|
||||
**Calling:** the host invokes the route fn — the Python bridge calls `run_ocr`; the `ai-gateway` server calls `realtime`. Register new modules in `lib.rs` / `mod.rs`, then run `cargo fmt && cargo clippy --workspace -- -D warnings && cargo test --workspace`.
|
||||
**Calling:** hosts invoke the core entrypoint — the Python bridge and the `ai-gateway` route service both call `litellm_core::messages::messages`. Never add a provider handler to `ai-gateway`. Register new modules in `lib.rs` / `mod.rs`, then run `cargo fmt && cargo clippy --workspace -- -D warnings && cargo test --workspace`.
|
||||
|
|
|
|||
|
|
@ -4,14 +4,30 @@ litellm-rust has exactly THREE crates. A crate is a LAYER, not a route. Routes (
|
|||
|
||||
## Crates
|
||||
|
||||
| Crate | Role | Pure / I/O |
|
||||
|-------|------|------------|
|
||||
| litellm-core | Translation layer — types, route contracts (traits), provider transforms (modules under providers/), and the router. Builds requests/responses; no network. | Pure |
|
||||
| litellm-ai-gateway | Routes + host — the only crate that touches the network. HTTP/WebSocket I/O (modules under io/) plus the axum server binary (behind the `server` feature). | I/O |
|
||||
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — a thin adapter over litellm-ai-gateway's I/O. | Binding |
|
||||
| Crate | Role |
|
||||
|-------|------|
|
||||
| litellm-core | The LiteLLM SDK in Rust. One public entrypoint per top-level call (`messages::messages()`), owning types, transforms, provider resolution, auth, and the provider HTTP call. Call it, get a typed response. |
|
||||
| litellm-ai-gateway | The axum server (behind the `server` feature) plus the WebSocket hosts. Translates HTTP/WS to core entrypoints; owns no provider logic and no handlers. |
|
||||
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — marshals Python objects and calls core entrypoints. |
|
||||
|
||||
Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge.
|
||||
|
||||
## Where a route lives
|
||||
|
||||
A top-level LiteLLM call is a module under `crates/core/src/<route>/`, shaped like `messages`:
|
||||
|
||||
```
|
||||
core/src/messages/
|
||||
mod.rs # pub async fn messages(..) -> CoreResult<..> (+ messages_stream for SSE)
|
||||
types.rs # request/response types, MessagesRequest
|
||||
transformation.rs # the provider template trait
|
||||
prepare.rs # provider resolution, auth headers, URL
|
||||
handler.rs # the provider call
|
||||
client.rs # the shared reqwest client
|
||||
```
|
||||
|
||||
Handlers never live in `ai-gateway`. `ocr`, `audio_transcription`, and `realtime` are still hosted there from before this rule; they move to `core` as they are touched.
|
||||
|
||||
Adding a crate: default to a MODULE. New crate ONLY on a real trigger — separate artifact (binary/cdylib), proc-macro, shared foundation, or publishable standalone. A new provider or route is none of these.
|
||||
|
||||
Adding a crate fails crates/core/tests/workspace_crate_allowlist.rs until you update its allowlist and this file — intentional.
|
||||
|
|
|
|||
|
|
@ -23,21 +23,34 @@ the base when behavior is genuinely different, and say so explicitly in the PR.
|
|||
|
||||
## Crates (exactly three — see AGENTS.md)
|
||||
|
||||
`litellm-core` describes work; `litellm-ai-gateway` executes it; `litellm-python-bridge`
|
||||
exposes it to the Python SDK. A crate is a **layer**, not a route — add modules, not crates.
|
||||
`litellm-core` **is** the LiteLLM SDK in Rust: it makes the LLM call.
|
||||
`litellm-ai-gateway` is an HTTP/WebSocket server in front of it, and
|
||||
`litellm-python-bridge` exposes it to the Python SDK. A crate is a **layer**, not
|
||||
a route — add modules, not crates.
|
||||
|
||||
## Core Boundary
|
||||
|
||||
`litellm-core` is the pure translation layer; the `litellm-ai-gateway` host executes work.
|
||||
`litellm-core` owns the whole call. The Rust equivalent of `litellm.messages()`
|
||||
is `litellm_core::messages::messages(request).await`: you call it, it does the
|
||||
provider call, and you get a typed non-streaming response back.
|
||||
|
||||
Route-level Rust structure mirrors LiteLLM's Python responsibilities:
|
||||
- `core/src/<route>/` owns the route contract, shared types, and provider
|
||||
template traits. For OCR, this means `core/src/ocr`.
|
||||
- `core/src/<route>/` owns the route end to end: the public entrypoint fn named
|
||||
after the route in `mod.rs`, the request/response types (`types.rs`), the
|
||||
provider template trait (`transformation.rs`), the provider/auth/URL
|
||||
resolution (`prepare.rs`), the HTTP client (`client.rs`), and the handler that
|
||||
performs the call (`handler.rs`). `core/src/messages` is the reference.
|
||||
- `core/src/providers/<provider>/<route>/transformation.rs` owns the
|
||||
provider-specific transform. For Mistral OCR, this means
|
||||
`core/src/providers/mistral/ocr/transformation.rs`.
|
||||
- Network execution lives in the host crate `ai-gateway` (`ai-gateway/src/io/`),
|
||||
never inside `core`.
|
||||
provider-specific transform. For Anthropic Messages, this means
|
||||
`core/src/providers/anthropic/messages/transformation.rs`.
|
||||
- Handlers live in `core`, never in a host. `ai-gateway` must not contain a
|
||||
route handler that talks to a provider; its axum route reads the HTTP request,
|
||||
picks a deployment, and calls the `core` entrypoint. `python-bridge` marshals
|
||||
Python objects and calls the same entrypoint.
|
||||
|
||||
Streaming keeps the same shape: the route entrypoint has a `<route>_stream`
|
||||
variant in `core` that returns the upstream response so a host can splice it to
|
||||
its own caller; the host still owns no provider logic.
|
||||
|
||||
Call-hook and lifecycle instrumentation, including phase timing, usage
|
||||
accumulation, and callback payload construction, always lives in `core`.
|
||||
|
|
@ -45,21 +58,31 @@ Hosts feed observed events into core and dispatch the completed payloads through
|
|||
their I/O logger; hosts must not own callback orchestration.
|
||||
|
||||
Allowed in `core`:
|
||||
- Pure request transforms
|
||||
- Pure response transforms
|
||||
- Pure stream chunk normalization
|
||||
- The public entrypoint for a top-level LiteLLM call
|
||||
- Request/response transforms and stream chunk normalization
|
||||
- Provider resolution, auth header construction, and URL building
|
||||
- The provider HTTP call itself, through a shared reused client with connect and
|
||||
request timeouts
|
||||
- Shared data types and validation errors
|
||||
- Deterministic token/cost helper logic
|
||||
|
||||
Not allowed in `core`:
|
||||
- Network calls
|
||||
- Environment variable or secret reads
|
||||
- Serving HTTP: axum routes, extractors, and transport concerns stay in the host
|
||||
- Filesystem access
|
||||
- Database or cache access
|
||||
- Provider SDK signing or auth flows
|
||||
- Database access
|
||||
- Config file reading and rollout state
|
||||
- Logging callbacks, spend writes, or custom callbacks
|
||||
- Global mutable runtime state
|
||||
|
||||
Env reads in `core` are limited to credential fallback inside a route's
|
||||
`prepare.rs` (the `env_lookup` closure), mirroring what the Python SDK does when
|
||||
no key is passed. Everything else config-shaped is resolved by the host and
|
||||
passed in.
|
||||
|
||||
Routes still hosted in `ai-gateway` (`ocr`, `audio_transcription`, `realtime`)
|
||||
predate this rule and are being moved into `core` route modules; do not add new
|
||||
ones there, and prefer moving one when you touch it.
|
||||
|
||||
Python owns rollout state and fallback while Rust is being introduced. Rust
|
||||
paths must be off by default until parity tests prove equivalence with Python.
|
||||
A new provider/route may instead be implemented rust-only with no Python
|
||||
|
|
@ -93,10 +116,10 @@ the first PR:
|
|||
- Preserve Python output shape intentionally. If a field is always serialized as
|
||||
`null` for Python parity, leave a short comment explaining that parity choice.
|
||||
|
||||
## Host I/O Rules
|
||||
## Network I/O Rules
|
||||
|
||||
These rules apply when adding future crates or modules that execute network I/O,
|
||||
such as `ai-gateway`, router hosts, or standalone servers:
|
||||
These rules apply to every module that executes network I/O, whether it is a
|
||||
`core` route handler or a host such as `ai-gateway`:
|
||||
|
||||
- Set connect and full-request timeouts. No unbounded waits.
|
||||
- Reuse HTTP clients; do not construct clients per request.
|
||||
|
|
|
|||
|
|
@ -2,18 +2,31 @@
|
|||
|
||||
This workspace contains the staged Rust implementation for LiteLLM.
|
||||
|
||||
Rust starts as a pure transform core used by the existing Python host. Python
|
||||
continues to own auth, configuration, network I/O, retries, routing, logging,
|
||||
`litellm-core` is the LiteLLM SDK in Rust: one entrypoint per top-level call
|
||||
that makes the LLM call and hands back a typed response, the same shape as
|
||||
`litellm.messages()` in Python.
|
||||
|
||||
```rust
|
||||
let response = litellm_core::messages::messages(MessagesRequest {
|
||||
model: "claude-sonnet-4-5",
|
||||
body,
|
||||
api_key: Some(key),
|
||||
..
|
||||
})
|
||||
.await?;
|
||||
```
|
||||
|
||||
Python continues to own configuration, retries, routing policy, logging,
|
||||
callbacks, spend tracking, and customer plugins until each Rust path has parity
|
||||
coverage and production evidence.
|
||||
|
||||
## Crates
|
||||
|
||||
| Crate | Role | Pure / I/O |
|
||||
|-------|------|------------|
|
||||
| litellm-core | Translation layer — types, route contracts (traits), provider transforms (modules under providers/), and the router. Builds requests/responses; no network. | Pure |
|
||||
| litellm-ai-gateway | Routes + host — the only crate that touches the network. HTTP/WebSocket I/O (modules under io/) plus the axum server binary (behind the `server` feature). | I/O |
|
||||
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — a thin adapter over litellm-ai-gateway's I/O. | Binding |
|
||||
| Crate | Role |
|
||||
|-------|------|
|
||||
| litellm-core | The SDK. Per-route entrypoints (`messages::messages()`), types, provider transforms (modules under `providers/`), provider resolution, auth, the provider HTTP call, and the router. |
|
||||
| litellm-ai-gateway | The axum server (behind the `server` feature) and WebSocket hosts. Translates HTTP/WS to core entrypoints; no provider handlers. |
|
||||
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — marshals Python objects and calls core entrypoints. |
|
||||
|
||||
Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge.
|
||||
|
||||
|
|
@ -21,16 +34,16 @@ Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-
|
|||
|
||||
```text
|
||||
crates/
|
||||
core/ Route contracts, shared pure types, errors, and templates.
|
||||
src/ocr/
|
||||
providers/ Provider-specific pure transforms.
|
||||
src/mistral/ocr/transformation.rs
|
||||
core/ The SDK: route modules + provider transforms.
|
||||
src/messages/ mod.rs (entrypoint), types, transformation, prepare, handler, client
|
||||
src/providers/anthropic/messages/transformation.rs
|
||||
ai-gateway/ Axum server + WebSocket hosts; calls core entrypoints.
|
||||
python-bridge/ PyO3 bridge for Python LiteLLM.
|
||||
```
|
||||
|
||||
The folder shape should follow the Python provider tree:
|
||||
`providers/src/<provider>/<route>/transformation.rs`. The bridge should expose
|
||||
one function per top-level route, starting with `ocr(payload)`.
|
||||
The folder shape follows the Python provider tree:
|
||||
`core/src/providers/<provider>/<route>/transformation.rs`. The bridge exposes one
|
||||
function per top-level route, mirroring the core entrypoints.
|
||||
|
||||
## Checks
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
# Provider coding standards (litellm-rust)
|
||||
|
||||
Rules for adding or changing an LLM provider/route in `litellm-rust`. OCR (`MISTRAL_OCR_CONFIG`) is the reference; `messages` (`ANTHROPIC_MESSAGES_CONFIG`) is the next port.
|
||||
Rules for adding or changing an LLM provider/route in `litellm-rust`. `messages` (`core/src/messages`, `ANTHROPIC_MESSAGES_CONFIG`) is the reference: a route is a `core` module with a public entrypoint that makes the call and returns a typed response.
|
||||
|
||||
## Provider resolution
|
||||
|
||||
|
|
@ -16,10 +16,10 @@ Rules for adding or changing an LLM provider/route in `litellm-rust`. OCR (`MIST
|
|||
|
||||
## Boundaries
|
||||
|
||||
7. Layers never cross: `core` = pure transforms/types (no network, env, secrets, auth, logging, global mutable state); `ai-gateway` = all I/O, auth headers, HTTP/SSE, lifecycle hooks; `python-bridge` = thin PyO3 adapter.
|
||||
7. Layers never cross: `core` = the call itself (entrypoint, types, transforms, provider resolution, auth headers, provider HTTP, lifecycle hooks); `ai-gateway` = serving HTTP/WS (routing, extractors, auth of *our* callers, streaming to the client); `python-bridge` = thin PyO3 adapter. Hosts call the core entrypoint; they never build a provider request.
|
||||
8. Generic/route files contain zero provider-specific branches. A provider is one module under `core/src/providers/<provider>/<route>/`; a route is a module, never a new crate.
|
||||
9. Route entry point stays thin: `<route>()` -> `prepare_*` -> `CallLifecycle::run_request`, which owns the pre_call -> during_call -> provider call -> success/failure order and phase timing. Handlers validate and delegate; no business logic in them.
|
||||
10. Constants (URLs, env-var names, API versions, error messages) live in a crate `constants.rs`, never inline. Env reads happen only at the host/config layer, with the `DEFAULT_*` fallback defined in `constants.rs`.
|
||||
9. Route entry point stays thin: `core::<route>::<route>()` -> `prepare_*` -> handler (or `CallLifecycle::run_request`, which owns the pre_call -> during_call -> provider call -> success/failure order and phase timing). Axum handlers validate and delegate to a service that calls the entrypoint; no business logic in them.
|
||||
10. Constants (URLs, env-var names, API versions, error messages) live in a crate `constants.rs`, never inline. Config-shaped env reads happen at the host/config layer with the `DEFAULT_*` fallback defined in `constants.rs`; the only env read in `core` is the credential fallback in a route's `prepare.rs`.
|
||||
|
||||
## Types and errors
|
||||
|
||||
|
|
@ -33,7 +33,7 @@ Rules for adding or changing an LLM provider/route in `litellm-rust`. OCR (`MIST
|
|||
|
||||
16. Never log request/response bodies, base64 payloads, document contents, or secrets. Truncate and bound any upstream body before it crosses a host boundary.
|
||||
17. Treat empty/whitespace credentials, URLs, and config values as absent at the host resolution layer.
|
||||
18. Host I/O sets connect + request timeouts (no unbounded waits), reuses a shared HTTP client, and prefers rustls TLS.
|
||||
18. Network I/O sets connect + request timeouts (no unbounded waits), reuses a shared HTTP client, and prefers rustls TLS.
|
||||
|
||||
## Tests and rollout
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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] =
|
||||
|
|
|
|||
|
|
@ -1 +0,0 @@
|
|||
pub use crate::messages::{MessagesRequest, messages};
|
||||
|
|
@ -1,5 +1,4 @@
|
|||
pub mod audio_transcription;
|
||||
pub mod messages;
|
||||
pub mod ocr;
|
||||
pub mod realtime;
|
||||
pub mod realtime_pool;
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1,49 +0,0 @@
|
|||
use litellm_core::CoreResult;
|
||||
use serde_json::Value;
|
||||
|
||||
mod client;
|
||||
mod common_utils;
|
||||
mod handler;
|
||||
mod prepare;
|
||||
mod types;
|
||||
|
||||
pub use types::MessagesRequest;
|
||||
|
||||
use handler::{execute_messages_provider_call, execute_messages_provider_stream};
|
||||
use prepare::prepare_messages_call;
|
||||
|
||||
pub async fn messages(request: MessagesRequest<'_>) -> CoreResult<Value> {
|
||||
match execute_messages(request, false).await? {
|
||||
MessagesResponse::Json(body) => Ok(body),
|
||||
MessagesResponse::Stream(response) => {
|
||||
drop(response);
|
||||
Err(litellm_core::CoreError::InvalidResponse(
|
||||
"non-streaming messages execution returned a stream".to_string(),
|
||||
))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) enum MessagesResponse {
|
||||
Json(Value),
|
||||
Stream(reqwest::Response),
|
||||
}
|
||||
|
||||
pub(crate) async fn execute_messages(
|
||||
request: MessagesRequest<'_>,
|
||||
stream: bool,
|
||||
) -> CoreResult<MessagesResponse> {
|
||||
let prepared = prepare_messages_call(request)?;
|
||||
if stream {
|
||||
execute_messages_provider_stream(prepared)
|
||||
.await
|
||||
.map(MessagesResponse::Stream)
|
||||
} else {
|
||||
execute_messages_provider_call(prepared)
|
||||
.await
|
||||
.map(MessagesResponse::Json)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
|
@ -1,24 +0,0 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_core::messages::transformation::AnthropicMessagesProviderConfig;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
pub struct MessagesRequest<'a> {
|
||||
pub model: &'a str,
|
||||
pub body: Value,
|
||||
pub api_key: Option<&'a str>,
|
||||
pub api_base: Option<&'a str>,
|
||||
pub custom_llm_provider: Option<&'a str>,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
pub(crate) struct ProviderMessagesRequest {
|
||||
pub(crate) provider: String,
|
||||
pub(crate) model: String,
|
||||
pub(crate) config: &'static dyn AnthropicMessagesProviderConfig,
|
||||
pub(crate) url: String,
|
||||
pub(crate) body: Value,
|
||||
pub(crate) upstream_headers: Vec<(String, String)>,
|
||||
pub(crate) timeout: Option<Duration>,
|
||||
}
|
||||
|
|
@ -19,7 +19,10 @@ async fn handle(...) -> impl IntoResponse { ... }
|
|||
When a route has business logic worth testing without axum, put it in a sibling
|
||||
`service` (a file, or a folder if the route grows). The route file stays the
|
||||
**axum surface** (router + handler + any socket/SSE adapter); `service` is plain
|
||||
Rust with **no axum types**. `realtime/` is the example:
|
||||
Rust with **no axum types**, and its job is to pick the deployment and call the
|
||||
`core` route entrypoint (see `messages/service.rs` calling
|
||||
`litellm_core::messages::messages`). Never build a provider request, resolve a
|
||||
key, or perform the provider call here. `realtime/` is the older example:
|
||||
```
|
||||
realtime/
|
||||
mod.rs # axum surface: router() + handler + the WS<->events adapter
|
||||
|
|
@ -33,6 +36,8 @@ genuinely gets hard to read.
|
|||
`crate::auth::RequireMasterKey` to its arguments; it runs during extraction.
|
||||
Never re-implement the check per route.
|
||||
- **Handlers contain no business logic; `service` contains no axum types.**
|
||||
- **No provider handlers in this crate.** Transforms, auth headers, and the
|
||||
provider HTTP call live in `core/src/<route>/`.
|
||||
- A route owns its paths in its own `router()`; `mod.rs` only merges.
|
||||
- Cross-cutting concerns (logging, CORS, timeouts) → Tower layers in `mod.rs`,
|
||||
not duplicated in handlers.
|
||||
|
|
|
|||
|
|
@ -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}"))
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,3 +1,7 @@
|
|||
litellm-core is the PURE translation layer — types, route contracts (traits), provider transforms (modules under `providers/`), and the router. No network, no I/O, no env reads.
|
||||
litellm-core is the LiteLLM SDK in Rust — it makes the LLM call. Each top-level call is a module under `src/<route>/` exposing a public entrypoint named after the route (`messages::messages()`, the Rust equivalent of `litellm.messages()`): you call it and get a typed non-streaming response back.
|
||||
|
||||
Routes (ocr, realtime) and providers (mistral, openai) are modules, not crates.
|
||||
A route module owns everything the call needs: types, the provider template trait, provider transforms (under `providers/`), provider/auth/URL resolution, and the handler that performs the HTTP call. Handlers belong here, never in a host crate.
|
||||
|
||||
Not here: serving HTTP (axum routes, extractors), config file reading, rollout state, databases, or callback dispatch. Env reads are limited to credential fallback in a route's `prepare.rs`.
|
||||
|
||||
Routes (messages, ocr, realtime) and providers (anthropic, mistral, openai) are modules, not crates.
|
||||
|
|
|
|||
|
|
@ -4,20 +4,28 @@ Rules for `litellm-rust/crates/core`.
|
|||
|
||||
## Responsibility
|
||||
|
||||
`core` owns shared data types, typed errors, and deterministic helper contracts.
|
||||
It must stay pure and host-independent.
|
||||
`core` is the LiteLLM SDK in Rust: it makes the LLM call. Every top-level
|
||||
LiteLLM call has a public entrypoint here, named after the route
|
||||
(`messages::messages()` is the Rust equivalent of `litellm.messages()`), and
|
||||
calling it returns a typed non-streaming response.
|
||||
|
||||
Allowed:
|
||||
- The public entrypoint for a route, plus its `<route>_stream` variant when the
|
||||
route supports streaming.
|
||||
- Provider resolution, auth header construction, URL building, and the provider
|
||||
HTTP call (shared reused client, connect + request timeouts).
|
||||
- Shared request/response structs.
|
||||
- Typed errors with stable, non-sensitive messages.
|
||||
- Deterministic validation helpers.
|
||||
- Serialization helpers that intentionally mirror Python output shape.
|
||||
- Route templates that match Python base config responsibilities, such as
|
||||
`ocr::transformation::OcrProviderConfig`.
|
||||
`messages::transformation::AnthropicMessagesProviderConfig`.
|
||||
|
||||
Not allowed:
|
||||
- Network, filesystem, database, cache, or environment access.
|
||||
- Secret reads or auth/header construction.
|
||||
- Serving HTTP: axum routers, extractors, and other transport concerns.
|
||||
- Filesystem, database, or cache access.
|
||||
- Config file reading or rollout state; the host resolves those and passes them
|
||||
in. Env reads are limited to credential fallback in a route's `prepare.rs`.
|
||||
- Logging callbacks, tracing spans, spend writes, or customer callbacks.
|
||||
- Provider-specific branching that belongs in `providers`.
|
||||
- Panics for user/provider-controlled input.
|
||||
|
|
@ -33,10 +41,21 @@ typed field on a struct, not a raw string threaded through the API.
|
|||
|
||||
## Structure
|
||||
|
||||
Use route names directly under `src/`: `ocr`, future `messages`,
|
||||
Use route names directly under `src/`: `messages`, `ocr`, future
|
||||
`chat_completions`, `embeddings`, and similar top-level LiteLLM calls. Do not
|
||||
invent broad names like `engine` for route contracts.
|
||||
|
||||
`src/messages` is the reference shape for a route module:
|
||||
|
||||
```
|
||||
mod.rs pub async fn messages(..) (+ messages_stream)
|
||||
types.rs request/response types
|
||||
transformation.rs the provider template trait
|
||||
prepare.rs provider resolution, auth headers, URL
|
||||
handler.rs the provider call
|
||||
client.rs the shared reqwest client
|
||||
```
|
||||
|
||||
## Parity Rules
|
||||
|
||||
- Every shared type used by a provider transform needs unit tests for
|
||||
|
|
|
|||
|
|
@ -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"] }
|
||||
|
|
|
|||
|
|
@ -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";
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
@ -1,15 +1,13 @@
|
|||
use litellm_core::CoreResult;
|
||||
use litellm_core::error::CoreError;
|
||||
use serde_json::Value;
|
||||
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
|
||||
use crate::error::{CoreError, CoreResult};
|
||||
|
||||
use super::client::http_client;
|
||||
use super::common_utils::truncate_error_body;
|
||||
use super::types::ProviderMessagesRequest;
|
||||
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
|
||||
use super::types::{AnthropicMessagesResponse, ProviderMessagesRequest};
|
||||
|
||||
pub(super) async fn execute_messages_provider_call(
|
||||
request: ProviderMessagesRequest,
|
||||
) -> CoreResult<Value> {
|
||||
) -> CoreResult<AnthropicMessagesResponse> {
|
||||
let mut request_builder = http_client().post(&request.url).json(&request.body);
|
||||
for (key, value) in &request.upstream_headers {
|
||||
request_builder = request_builder.header(key, value);
|
||||
|
|
@ -39,12 +37,7 @@ pub(super) async fn execute_messages_provider_call(
|
|||
let response = serde_json::from_str(&text).map_err(|err| {
|
||||
CoreError::InvalidResponse(format!("invalid messages response JSON: {err}"))
|
||||
})?;
|
||||
let transformed = request
|
||||
.config
|
||||
.transform_response(&request.model, response)?;
|
||||
serde_json::to_value(transformed).map_err(|err| {
|
||||
CoreError::InvalidResponse(format!("failed to serialize messages response: {err}"))
|
||||
})
|
||||
request.config.transform_response(&request.model, response)
|
||||
}
|
||||
|
||||
pub(super) async fn execute_messages_provider_stream(
|
||||
|
|
@ -1,2 +1,32 @@
|
|||
//! The Anthropic Messages call, the Rust equivalent of Python's
|
||||
//! `litellm.messages()`.
|
||||
//!
|
||||
//! [`messages`] is the top-level entrypoint: give it a model, a body, and
|
||||
//! credentials, and it resolves the provider, transforms the request, calls the
|
||||
//! provider, and returns a typed non-streaming response. [`messages_stream`]
|
||||
//! is the streaming variant; it hands the raw upstream response back so a host
|
||||
//! can splice the event stream to its own caller.
|
||||
|
||||
mod client;
|
||||
mod common_utils;
|
||||
mod handler;
|
||||
mod prepare;
|
||||
pub mod transformation;
|
||||
pub mod types;
|
||||
|
||||
use crate::error::CoreResult;
|
||||
|
||||
use handler::{execute_messages_provider_call, execute_messages_provider_stream};
|
||||
use prepare::prepare_messages_call;
|
||||
use types::{AnthropicMessagesResponse, MessagesRequest};
|
||||
|
||||
pub async fn messages(request: MessagesRequest<'_>) -> CoreResult<AnthropicMessagesResponse> {
|
||||
execute_messages_provider_call(prepare_messages_call(request)?).await
|
||||
}
|
||||
|
||||
pub async fn messages_stream(request: MessagesRequest<'_>) -> CoreResult<reqwest::Response> {
|
||||
execute_messages_provider_stream(prepare_messages_call(request)?).await
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
@ -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");
|
||||
|
|
@ -1,6 +1,30 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::transformation::AnthropicMessagesProviderConfig;
|
||||
|
||||
pub struct MessagesRequest<'a> {
|
||||
pub model: &'a str,
|
||||
pub body: Value,
|
||||
pub api_key: Option<&'a str>,
|
||||
pub api_base: Option<&'a str>,
|
||||
pub custom_llm_provider: Option<&'a str>,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
pub(super) struct ProviderMessagesRequest {
|
||||
pub(super) provider: String,
|
||||
pub(super) model: String,
|
||||
pub(super) config: &'static dyn AnthropicMessagesProviderConfig,
|
||||
pub(super) url: String,
|
||||
pub(super) body: Value,
|
||||
pub(super) upstream_headers: Vec<(String, String)>,
|
||||
pub(super) timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum SystemPrompt {
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -4,10 +4,11 @@ use std::time::Duration;
|
|||
use litellm_ai_gateway::io::audio_transcription::{
|
||||
AudioTranscriptionRequest, audio_transcription as run_audio_transcription,
|
||||
};
|
||||
use litellm_ai_gateway::io::messages::{MessagesRequest, messages as run_messages};
|
||||
use litellm_ai_gateway::io::ocr::{OcrRequest, ocr as run_ocr};
|
||||
use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection;
|
||||
use litellm_core::error::CoreError;
|
||||
use litellm_core::messages::messages as run_messages;
|
||||
use litellm_core::messages::types::{AnthropicMessagesResponse, MessagesRequest};
|
||||
use pyo3::exceptions::{PyRuntimeError, PyValueError};
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::{PyAny, PyDict};
|
||||
|
|
@ -35,6 +36,15 @@ fn json_to_py(py: Python<'_>, value: Value) -> PyResult<Py<PyAny>> {
|
|||
Ok(json.call_method1("loads", (encoded,))?.unbind())
|
||||
}
|
||||
|
||||
fn messages_response_to_py(
|
||||
py: Python<'_>,
|
||||
response: AnthropicMessagesResponse,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
let value =
|
||||
serde_json::to_value(response).map_err(|err| PyValueError::new_err(err.to_string()))?;
|
||||
json_to_py(py, value)
|
||||
}
|
||||
|
||||
fn core_error_to_pyerr(err: CoreError) -> PyErr {
|
||||
match err {
|
||||
CoreError::Auth(message) => PyValueError::new_err(message),
|
||||
|
|
@ -382,7 +392,7 @@ fn messages(
|
|||
});
|
||||
|
||||
match result {
|
||||
Ok(value) => json_to_py(py, value),
|
||||
Ok(response) => messages_response_to_py(py, response),
|
||||
Err(err) => Err(core_error_to_pyerr(err)),
|
||||
}
|
||||
}
|
||||
|
|
@ -404,7 +414,7 @@ fn amessages(
|
|||
marshal_messages_inputs(py, body, extra_headers, timeout_seconds)?;
|
||||
|
||||
pyo3_async_runtimes::tokio::future_into_py(py, async move {
|
||||
let value = run_messages(MessagesRequest {
|
||||
let response = run_messages(MessagesRequest {
|
||||
model: &model,
|
||||
body,
|
||||
api_key: api_key.as_deref(),
|
||||
|
|
@ -416,7 +426,7 @@ fn amessages(
|
|||
.await
|
||||
.map_err(core_error_to_pyerr)?;
|
||||
|
||||
Python::attach(|py| json_to_py(py, value))
|
||||
Python::attach(|py| messages_response_to_py(py, response))
|
||||
})
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -30,15 +30,13 @@ from opentelemetry.sdk.trace.export.in_memory_span_exporter import (
|
|||
InMemorySpanExporter,
|
||||
)
|
||||
from opentelemetry.trace import Span, SpanKind, Tracer
|
||||
from opentelemetry.util.re import parse_env_headers
|
||||
|
||||
from litellm._version import version as litellm_version
|
||||
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
|
||||
from litellm.integrations.otel.model.semconv import LiteLLM
|
||||
from litellm.integrations.otel.model.spans import LiteLLMSpanKind
|
||||
|
||||
# Re-exported so ``providers.parse_headers`` remains a stable entry point.
|
||||
from litellm.integrations.otel.model.utils import parse_headers as parse_headers
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.metrics import Meter
|
||||
from opentelemetry.sdk.metrics.export import MetricReader
|
||||
|
|
@ -133,6 +131,23 @@ def destination_resource_attrs(destination: "OtelDestination") -> Mapping[str, s
|
|||
return dict(destination.resource_attributes)
|
||||
|
||||
|
||||
def parse_headers(raw: str | None) -> dict[str, str]:
|
||||
"""Parse an OTLP ``"k=v,k=v"`` header string into a dict.
|
||||
|
||||
``OTEL_EXPORTER_OTLP_HEADERS`` is W3C Baggage encoded per the OTLP spec, so
|
||||
values are percent-decoded: a vendor that documents
|
||||
``Authorization=Basic%20<token>`` (Grafana Cloud does, because a bare space
|
||||
is not representable there) has to reach the exporter as ``Basic <token>``,
|
||||
not with a literal ``%20`` that the backend rejects as malformed. The SDK's
|
||||
own parser is used so litellm decodes exactly what the OTLP exporters do
|
||||
when they read the env var themselves; ``liberal`` keeps values that are not
|
||||
percent-encoded (``Authorization=Bearer <token>``) working unchanged.
|
||||
"""
|
||||
if not raw:
|
||||
return {}
|
||||
return dict(parse_env_headers(raw, liberal=True))
|
||||
|
||||
|
||||
def _exporter_from_spec(spec: ExporterSpec) -> SpanExporter:
|
||||
kind = (spec.kind or "console").lower()
|
||||
factory = _EXPORTER_FACTORIES.get(kind)
|
||||
|
|
@ -200,8 +215,16 @@ def _otlp_metrics_endpoint(endpoint: str | None) -> str | None:
|
|||
def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader":
|
||||
"""Build a metric reader mirroring v1's exporter selection.
|
||||
|
||||
``console`` (and any unrecognized kind) exports to the console; ``otlp_http``/``otlp_grpc``
|
||||
export over OTLP with the configured endpoint/headers, on a 5s period.
|
||||
``console`` (and any unrecognized kind) exports to the console; ``otlp_http``
|
||||
and ``otlp_grpc`` export over OTLP with the configured endpoint/headers. The
|
||||
reader exports on a 5s period, matching v1.
|
||||
|
||||
Histograms keep the SDK's default cumulative temporality. Prometheus-backed
|
||||
OTLP receivers (Grafana Cloud / Mimir, and the Prometheus OTLP endpoint)
|
||||
reject delta histograms outright with ``invalid temporality and type
|
||||
combination``, which drops the whole metric batch, while backends that
|
||||
prefer delta still accept cumulative. The enterprise billing exporter
|
||||
already relies on the same default.
|
||||
"""
|
||||
from opentelemetry.sdk.metrics.export import (
|
||||
ConsoleMetricExporter,
|
||||
|
|
@ -213,18 +236,12 @@ def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader":
|
|||
from opentelemetry.exporter.otlp.proto.http.metric_exporter import (
|
||||
OTLPMetricExporter as HTTPMetricExporter,
|
||||
)
|
||||
from opentelemetry.sdk.metrics import Histogram
|
||||
from opentelemetry.sdk.metrics.export import AggregationTemporality
|
||||
|
||||
exporter: Any = HTTPMetricExporter(
|
||||
endpoint=_otlp_metrics_endpoint(config.endpoint),
|
||||
headers=parse_headers(config.headers),
|
||||
preferred_temporality={Histogram: AggregationTemporality.DELTA},
|
||||
)
|
||||
elif kind in ("otlp_grpc", "grpc"):
|
||||
from opentelemetry.sdk.metrics import Histogram
|
||||
from opentelemetry.sdk.metrics.export import AggregationTemporality
|
||||
|
||||
try:
|
||||
from opentelemetry.exporter.otlp.proto.grpc.metric_exporter import (
|
||||
OTLPMetricExporter as GRPCMetricExporter,
|
||||
|
|
@ -238,7 +255,6 @@ def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader":
|
|||
exporter = GRPCMetricExporter(
|
||||
endpoint=config.endpoint,
|
||||
headers=parse_headers(config.headers),
|
||||
preferred_temporality={Histogram: AggregationTemporality.DELTA},
|
||||
)
|
||||
else:
|
||||
exporter = ConsoleMetricExporter()
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 [],
|
||||
)
|
||||
|
|
|
|||
137
litellm/proxy/analytics_endpoints/cache_activity.py
Normal file
137
litellm/proxy/analytics_endpoints/cache_activity.py
Normal file
|
|
@ -0,0 +1,137 @@
|
|||
import asyncio
|
||||
import json
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Sequence
|
||||
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
UNKNOWN_CALL_TYPE = "Unknown"
|
||||
|
||||
|
||||
class CacheActivityGroup(BaseModel):
|
||||
call_type: str
|
||||
api_requests: int
|
||||
cache_hits: int
|
||||
failed_requests: int
|
||||
cached_completion_tokens: int
|
||||
generated_completion_tokens: int
|
||||
|
||||
|
||||
class CacheActivityTotals(BaseModel):
|
||||
api_requests: int
|
||||
cache_hits: int
|
||||
failed_requests: int
|
||||
cached_completion_tokens: int
|
||||
cache_hit_ratio: float
|
||||
|
||||
|
||||
class CacheActivityFilterOptions(BaseModel):
|
||||
key_aliases: list[str]
|
||||
models: list[str]
|
||||
|
||||
|
||||
class CacheActivityResponse(BaseModel):
|
||||
groups: list[CacheActivityGroup]
|
||||
totals: CacheActivityTotals
|
||||
filter_options: CacheActivityFilterOptions
|
||||
|
||||
|
||||
GROUPS_SQL = """
|
||||
SELECT
|
||||
CASE WHEN sl."call_type" = '' THEN 'Unknown' ELSE sl."call_type" END AS call_type,
|
||||
(COUNT(*)
|
||||
- SUM(CASE WHEN COALESCE(sl."cache_hit", '') = 'True' THEN 1 ELSE 0 END)
|
||||
- SUM(CASE WHEN sl."status" = 'failure' THEN 1 ELSE 0 END))::int AS api_requests,
|
||||
SUM(CASE WHEN COALESCE(sl."cache_hit", '') = 'True' THEN 1 ELSE 0 END)::int AS cache_hits,
|
||||
SUM(CASE WHEN sl."status" = 'failure' THEN 1 ELSE 0 END)::int AS failed_requests,
|
||||
SUM(CASE WHEN COALESCE(sl."cache_hit", '') = 'True' THEN sl."completion_tokens" ELSE 0 END)::int
|
||||
AS cached_completion_tokens,
|
||||
SUM(CASE WHEN COALESCE(sl."cache_hit", '') != 'True' THEN sl."completion_tokens" ELSE 0 END)::int
|
||||
AS generated_completion_tokens
|
||||
FROM "LiteLLM_SpendLogs" sl
|
||||
LEFT JOIN "LiteLLM_VerificationToken" vt ON sl."api_key" = vt."token"
|
||||
WHERE
|
||||
sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC')
|
||||
AND sl."startTime" < (($2::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC')
|
||||
AND ($3::jsonb = '[]'::jsonb
|
||||
OR COALESCE(vt."key_alias", 'Unnamed Key') IN (SELECT jsonb_array_elements_text($3::jsonb)))
|
||||
AND ($4::jsonb = '[]'::jsonb
|
||||
OR sl."model" IN (SELECT jsonb_array_elements_text($4::jsonb)))
|
||||
GROUP BY 1
|
||||
ORDER BY (COUNT(*)) DESC
|
||||
"""
|
||||
|
||||
KEY_ALIAS_OPTIONS_SQL = """
|
||||
SELECT DISTINCT COALESCE(vt."key_alias", 'Unnamed Key') AS key_alias
|
||||
FROM "LiteLLM_SpendLogs" sl
|
||||
LEFT JOIN "LiteLLM_VerificationToken" vt ON sl."api_key" = vt."token"
|
||||
WHERE
|
||||
sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC')
|
||||
AND sl."startTime" < (($2::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC')
|
||||
ORDER BY 1
|
||||
"""
|
||||
|
||||
MODEL_OPTIONS_SQL = """
|
||||
SELECT DISTINCT sl."model" AS model
|
||||
FROM "LiteLLM_SpendLogs" sl
|
||||
WHERE
|
||||
sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC')
|
||||
AND sl."startTime" < (($2::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC')
|
||||
AND sl."model" != ''
|
||||
ORDER BY 1
|
||||
"""
|
||||
|
||||
|
||||
class _KeyAliasRow(BaseModel):
|
||||
key_alias: str
|
||||
|
||||
|
||||
class _ModelRow(BaseModel):
|
||||
model: str
|
||||
|
||||
|
||||
_groups_adapter = TypeAdapter(list[CacheActivityGroup])
|
||||
_key_alias_rows_adapter = TypeAdapter(list[_KeyAliasRow])
|
||||
_model_rows_adapter = TypeAdapter(list[_ModelRow])
|
||||
|
||||
|
||||
def compute_totals(groups: Sequence[CacheActivityGroup]) -> CacheActivityTotals:
|
||||
api_requests = sum(group.api_requests for group in groups)
|
||||
cache_hits = sum(group.cache_hits for group in groups)
|
||||
failed_requests = sum(group.failed_requests for group in groups)
|
||||
all_requests = api_requests + cache_hits + failed_requests
|
||||
return CacheActivityTotals(
|
||||
api_requests=api_requests,
|
||||
cache_hits=cache_hits,
|
||||
failed_requests=failed_requests,
|
||||
cached_completion_tokens=sum(group.cached_completion_tokens for group in groups),
|
||||
cache_hit_ratio=(cache_hits / all_requests) * 100 if all_requests > 0 else 0.0,
|
||||
)
|
||||
|
||||
|
||||
async def get_cache_activity(
|
||||
prisma_client: "PrismaClient",
|
||||
start_date: datetime,
|
||||
end_date: datetime,
|
||||
key_aliases: Sequence[str],
|
||||
models: Sequence[str],
|
||||
) -> CacheActivityResponse:
|
||||
group_rows, key_alias_rows, model_rows = await asyncio.gather(
|
||||
prisma_client.db.query_raw(
|
||||
GROUPS_SQL, start_date, end_date, json.dumps(list(key_aliases)), json.dumps(list(models))
|
||||
),
|
||||
prisma_client.db.query_raw(KEY_ALIAS_OPTIONS_SQL, start_date, end_date),
|
||||
prisma_client.db.query_raw(MODEL_OPTIONS_SQL, start_date, end_date),
|
||||
)
|
||||
groups = _groups_adapter.validate_python(group_rows or [])
|
||||
return CacheActivityResponse(
|
||||
groups=groups,
|
||||
totals=compute_totals(groups),
|
||||
filter_options=CacheActivityFilterOptions(
|
||||
key_aliases=[row.key_alias for row in _key_alias_rows_adapter.validate_python(key_alias_rows or [])],
|
||||
models=[row.model for row in _model_rows_adapter.validate_python(model_rows or [])],
|
||||
),
|
||||
)
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -175,6 +175,7 @@ class ScimTransformations:
|
|||
SCIMMember(
|
||||
value=ScimTransformations._get_scim_member_value(member),
|
||||
display=ScimTransformations._get_scim_member_display(member),
|
||||
type="User",
|
||||
)
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -1,141 +0,0 @@
|
|||
"""Live e2e: the built-in Presidio PII guardrail masks PII on the request and on
|
||||
the model output.
|
||||
|
||||
Presidio replaces detected PII with `<ENTITY_TYPE>` placeholders (e.g.
|
||||
`<EMAIL_ADDRESS>`) via a real analyzer + anonymizer. Two modes are checked
|
||||
independently, each opted into per request (default_on=False) so it never touches
|
||||
unrelated traffic:
|
||||
|
||||
- pre_call: the prompt is anonymized before it reaches the model, so a
|
||||
repeat-verbatim request comes back with the placeholder, never the raw email
|
||||
- post_call (apply_to_output): PII the model itself emits is masked on the way
|
||||
out, so the caller never receives the raw value the model produced
|
||||
|
||||
A third mode, logging_only, is not covered here: the raw email stayed in the OTEL
|
||||
span's `gen_ai.input.messages` on every attempt over a full poll deadline while
|
||||
these two modes masked correctly, so that cell is tracked in LIT-4841 rather than
|
||||
asserted against known-failing behavior.
|
||||
|
||||
Analyzer/anonymizer bases come from PRESIDIO_ANALYZER_API_BASE /
|
||||
PRESIDIO_ANONYMIZER_API_BASE (compose provides the in-network hosts; point them at
|
||||
locally published container ports for a host run). The chat backend is a gemini
|
||||
deployment created for the test.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
|
||||
import pytest
|
||||
|
||||
from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, unique_marker
|
||||
from e2e_http import unwrap
|
||||
from guardrails_client import GuardrailMode, GuardrailsClient, PresidioParamsBody
|
||||
from lifecycle import ResourceManager
|
||||
from models import ChatResponse
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
RAW_EMAIL = "alice.example.person@example.com"
|
||||
PLACEHOLDER = "<EMAIL_ADDRESS>"
|
||||
|
||||
ECHO_REQUEST = f"Repeat the following text back exactly, verbatim, with no changes: My email is {RAW_EMAIL}"
|
||||
EMIT_REQUEST = f"Output exactly this one line and nothing else: Please contact {RAW_EMAIL} today"
|
||||
|
||||
|
||||
def _content(response: ChatResponse) -> str:
|
||||
if not response.choices:
|
||||
return ""
|
||||
message = response.choices[0].message
|
||||
return (message.content if message else None) or ""
|
||||
|
||||
|
||||
def _presidio_params(
|
||||
mode: GuardrailMode, *, apply_to_output: bool = False, logging_only: bool = False
|
||||
) -> PresidioParamsBody:
|
||||
analyzer = os.environ["PRESIDIO_ANALYZER_API_BASE"]
|
||||
anonymizer = os.environ["PRESIDIO_ANONYMIZER_API_BASE"]
|
||||
return PresidioParamsBody(
|
||||
mode=mode,
|
||||
default_on=False,
|
||||
presidio_analyzer_api_base=analyzer,
|
||||
presidio_anonymizer_api_base=anonymizer,
|
||||
apply_to_output=apply_to_output,
|
||||
logging_only=logging_only,
|
||||
)
|
||||
|
||||
|
||||
def _poll_until_masked(call: Callable[[], str]) -> str:
|
||||
"""Retry a call until the guardrail masks its PII, returning the last content.
|
||||
|
||||
Registering a guardrail is a control-plane write; the data-plane worker that
|
||||
serves /chat/completions only picks it up on its next periodic DB sync (~30s
|
||||
in proxy_server.py), so a call issued the instant after the create runs
|
||||
against a worker that has no guardrail yet and passes the raw value through.
|
||||
That is in-flight propagation, not a masking failure. Polling to the deadline
|
||||
waits it out, so the assertions that follow judge the synced state; if the
|
||||
mask never lands the last unmasked content is returned and they still fail.
|
||||
"""
|
||||
deadline = time.monotonic() + POLL_TIMEOUT
|
||||
last = call()
|
||||
while time.monotonic() < deadline:
|
||||
if PLACEHOLDER in last and RAW_EMAIL not in last:
|
||||
return last
|
||||
time.sleep(POLL_INTERVAL)
|
||||
last = call()
|
||||
return last
|
||||
|
||||
|
||||
class TestPresidioGuardrail:
|
||||
@pytest.mark.covers(
|
||||
"guardrail.presidio.pre_call.masks",
|
||||
exercised_on=["chat_completions"],
|
||||
)
|
||||
def test_pre_call_masks_pii_before_the_model_sees_it(
|
||||
self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str
|
||||
) -> None:
|
||||
model = client.create_backend_model(resources, prefix="e2e-presidio-pre")
|
||||
name = f"e2e-presidio-pre-{unique_marker()}"
|
||||
guardrail_id = client.register(name, _presidio_params("pre_call"))
|
||||
resources.defer(lambda: client.delete_guardrail(guardrail_id))
|
||||
|
||||
echoed = _poll_until_masked(
|
||||
lambda: _content(
|
||||
unwrap(client.chat(scoped_key, model, ECHO_REQUEST, guardrails=[name], max_tokens=128))
|
||||
)
|
||||
)
|
||||
assert RAW_EMAIL not in echoed, (
|
||||
"pre_call masking must strip the raw email before the model sees it, but the "
|
||||
f"model echoed it back: {echoed[:300]!r}"
|
||||
)
|
||||
assert PLACEHOLDER in echoed, (
|
||||
"the model should have echoed the masked placeholder the guardrail substituted, "
|
||||
f"got: {echoed[:300]!r}"
|
||||
)
|
||||
|
||||
@pytest.mark.covers(
|
||||
"guardrail.presidio.post_call.masks",
|
||||
exercised_on=["chat_completions"],
|
||||
)
|
||||
def test_post_call_masks_pii_in_model_output(
|
||||
self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str
|
||||
) -> None:
|
||||
model = client.create_backend_model(resources, prefix="e2e-presidio-post")
|
||||
name = f"e2e-presidio-post-{unique_marker()}"
|
||||
guardrail_id = client.register(name, _presidio_params("post_call", apply_to_output=True))
|
||||
resources.defer(lambda: client.delete_guardrail(guardrail_id))
|
||||
|
||||
out = _poll_until_masked(
|
||||
lambda: _content(
|
||||
unwrap(client.chat(scoped_key, model, EMIT_REQUEST, guardrails=[name], max_tokens=128))
|
||||
)
|
||||
)
|
||||
assert RAW_EMAIL not in out, (
|
||||
"post_call masking must strip PII the model emitted, but the raw email reached the "
|
||||
f"caller: {out[:300]!r}"
|
||||
)
|
||||
assert PLACEHOLDER in out, (
|
||||
f"the masked placeholder should replace the model's PII output, got: {out[:300]!r}"
|
||||
)
|
||||
|
|
@ -451,6 +451,30 @@ def test_parse_headers():
|
|||
assert providers.parse_headers("no-equals") == {}
|
||||
|
||||
|
||||
def test_parse_headers_percent_decodes_values():
|
||||
"""A percent-encoded OTLP header value reaches the exporter decoded.
|
||||
|
||||
``OTEL_EXPORTER_OTLP_HEADERS`` is W3C Baggage encoded, and Grafana Cloud
|
||||
documents ``Authorization=Basic%20<token>``. Forwarding the literal ``%20``
|
||||
makes the backend reject the export as a malformed credential.
|
||||
"""
|
||||
token = "MTMzNzc4MzpnbGNfZXlKdklqb2lNVEl6TkNJPQ=="
|
||||
assert providers.parse_headers(f"Authorization=Basic%20{token}") == {"authorization": f"Basic {token}"}
|
||||
assert providers.parse_headers("x-scope-orgid=team%20a") == {"x-scope-orgid": "team a"}
|
||||
|
||||
|
||||
def test_parse_headers_keeps_unencoded_values_working():
|
||||
"""Values that are not percent-encoded keep parsing unchanged.
|
||||
|
||||
Vendors that document a bare space, and litellm's own presets, must survive
|
||||
the switch to the spec-compliant parser. Base64 padding also means a value
|
||||
can contain ``=``, so only the first one may split the pair.
|
||||
"""
|
||||
assert providers.parse_headers("Authorization=Bearer sk-123") == {"authorization": "Bearer sk-123"}
|
||||
assert providers.parse_headers("api_key=abc,space_id=xyz") == {"api_key": "abc", "space_id": "xyz"}
|
||||
assert providers.parse_headers("api_key=YWJjZA==") == {"api_key": "YWJjZA=="}
|
||||
|
||||
|
||||
def test_otlp_traces_endpoint_normalization():
|
||||
norm = providers._otlp_traces_endpoint
|
||||
# A base endpoint gets the signal path appended (the common OTLP env shape).
|
||||
|
|
@ -497,6 +521,24 @@ def test_build_span_exporter_variants():
|
|||
assert "OTLPSpanExporter" in type(http_exporter).__name__
|
||||
|
||||
|
||||
def test_otlp_metric_exporter_uses_cumulative_histogram_temporality():
|
||||
"""Histograms must export as cumulative, not delta.
|
||||
|
||||
Prometheus-backed OTLP receivers (Grafana Cloud / Mimir) reject delta
|
||||
histograms with ``invalid temporality and type combination`` and drop the
|
||||
entire metric batch, so a delta default silently loses every GenAI metric.
|
||||
"""
|
||||
from opentelemetry.sdk.metrics import Histogram
|
||||
from opentelemetry.sdk.metrics.export import AggregationTemporality
|
||||
|
||||
reader = providers.build_metric_reader(
|
||||
OpenTelemetryV2Config(exporter="otlp_http", endpoint="http://h:4318")
|
||||
)
|
||||
temporality = reader._exporter._preferred_temporality # noqa: SLF001 # exporter exposes no public accessor
|
||||
|
||||
assert temporality[Histogram] is AggregationTemporality.CUMULATIVE
|
||||
|
||||
|
||||
def test_otlp_logs_endpoint_normalization():
|
||||
norm = providers._otlp_logs_endpoint
|
||||
# A base endpoint gets the signal path appended (the common OTLP env shape).
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
0
tests/test_litellm/proxy/analytics_endpoints/__init__.py
Normal file
0
tests/test_litellm/proxy/analytics_endpoints/__init__.py
Normal file
|
|
@ -0,0 +1,133 @@
|
|||
"""
|
||||
The cache dashboard chart is fed by /global/activity/cache_hits. Aggregation
|
||||
lives server-side: the SQL groups per call_type (splitting cache hits vs
|
||||
successful vs failed requests; failed spend logs have call_type '' today and
|
||||
must surface as 'Unknown'), and the endpoint returns chart-ready groups,
|
||||
totals for the stat cards, and the filter options for the UI dropdowns.
|
||||
"""
|
||||
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy.analytics_endpoints.analytics_endpoints import get_global_activity
|
||||
from litellm.proxy.analytics_endpoints.cache_activity import (
|
||||
GROUPS_SQL,
|
||||
CacheActivityGroup,
|
||||
compute_totals,
|
||||
)
|
||||
|
||||
GROUP_ROWS = [
|
||||
{
|
||||
"call_type": "acompletion",
|
||||
"api_requests": 1000,
|
||||
"cache_hits": 300,
|
||||
"failed_requests": 200,
|
||||
"cached_completion_tokens": 12000,
|
||||
"generated_completion_tokens": 48000,
|
||||
},
|
||||
{
|
||||
"call_type": "Unknown",
|
||||
"api_requests": 0,
|
||||
"cache_hits": 0,
|
||||
"failed_requests": 110,
|
||||
"cached_completion_tokens": 0,
|
||||
"generated_completion_tokens": 0,
|
||||
},
|
||||
]
|
||||
KEY_ALIAS_ROWS = [{"key_alias": "Unnamed Key"}, {"key_alias": "my-key"}]
|
||||
MODEL_ROWS = [{"model": "gpt-5.1"}]
|
||||
|
||||
|
||||
def build_prisma(query_raw: AsyncMock) -> MagicMock:
|
||||
prisma = MagicMock()
|
||||
prisma.db.query_raw = query_raw
|
||||
return prisma
|
||||
|
||||
|
||||
def dispatching_query_raw() -> AsyncMock:
|
||||
async def dispatch(sql: str, *params: object) -> list[dict[str, object]]:
|
||||
if "GROUP BY" in sql:
|
||||
return GROUP_ROWS
|
||||
if "key_alias" in sql:
|
||||
return KEY_ALIAS_ROWS
|
||||
return MODEL_ROWS
|
||||
|
||||
return AsyncMock(side_effect=dispatch)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_prisma(monkeypatch: pytest.MonkeyPatch) -> MagicMock:
|
||||
prisma = build_prisma(dispatching_query_raw())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma)
|
||||
return prisma
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_groups_totals_and_filter_options(mock_prisma: MagicMock):
|
||||
response = await get_global_activity(
|
||||
start_date="2026-07-01", end_date="2026-07-27", key_aliases=[], models=[]
|
||||
)
|
||||
|
||||
assert [group.call_type for group in response.groups] == ["acompletion", "Unknown"]
|
||||
assert response.groups[0].api_requests == 1000
|
||||
assert response.groups[0].failed_requests == 200
|
||||
assert response.totals.api_requests == 1000
|
||||
assert response.totals.cache_hits == 300
|
||||
assert response.totals.failed_requests == 310
|
||||
assert response.totals.cached_completion_tokens == 12000
|
||||
assert response.totals.cache_hit_ratio == pytest.approx((300 / 1610) * 100)
|
||||
assert response.filter_options.key_aliases == ["Unnamed Key", "my-key"]
|
||||
assert response.filter_options.models == ["gpt-5.1"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_filters_are_passed_to_sql_as_json_arrays(mock_prisma: MagicMock):
|
||||
await get_global_activity(
|
||||
start_date="2026-07-01",
|
||||
end_date="2026-07-27",
|
||||
key_aliases=["my-key"],
|
||||
models=["gpt-5.1", "claude-opus-4-8"],
|
||||
)
|
||||
|
||||
groups_call = next(
|
||||
call for call in mock_prisma.db.query_raw.call_args_list if "GROUP BY" in call.args[0]
|
||||
)
|
||||
assert groups_call.args[3] == json.dumps(["my-key"])
|
||||
assert groups_call.args[4] == json.dumps(["gpt-5.1", "claude-opus-4-8"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rejects_malformed_dates_with_400(mock_prisma: MagicMock):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await get_global_activity(start_date="07/01/2026", end_date="2026-07-27", key_aliases=[], models=[])
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
mock_prisma.db.query_raw.assert_not_called()
|
||||
|
||||
|
||||
def test_totals_ratio_is_zero_without_requests():
|
||||
totals = compute_totals([])
|
||||
|
||||
assert totals.cache_hit_ratio == 0.0
|
||||
assert totals.api_requests == 0
|
||||
|
||||
|
||||
def test_totals_denominator_includes_failed_requests():
|
||||
group = CacheActivityGroup(
|
||||
call_type="acompletion",
|
||||
api_requests=60,
|
||||
cache_hits=20,
|
||||
failed_requests=20,
|
||||
cached_completion_tokens=0,
|
||||
generated_completion_tokens=0,
|
||||
)
|
||||
|
||||
assert compute_totals([group]).cache_hit_ratio == pytest.approx(20.0)
|
||||
|
||||
|
||||
def test_groups_sql_splits_failures_and_labels_empty_call_type_unknown():
|
||||
assert "SUM(CASE WHEN sl.\"status\" = 'failure' THEN 1 ELSE 0 END)" in GROUPS_SQL
|
||||
assert "CASE WHEN sl.\"call_type\" = '' THEN 'Unknown' ELSE sl.\"call_type\" END" in GROUPS_SQL
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -2918,6 +2918,123 @@ async def test_team_metadata_refreshed_from_team_object_during_auth():
|
|||
setattr(_proxy_server_mod, k, v)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auth_flow_never_persists_fallback_team_object_lit_4391():
|
||||
"""
|
||||
Regression test for LIT-4391 (stale team allowlist poisoning).
|
||||
|
||||
When `get_team_object` fails at the "Check 6" team-auth step (cache miss
|
||||
inside the DB-throttle window, DB blip, ...), the builder falls back to a
|
||||
team object reconstructed from the CACHED token's team_* snapshot — which
|
||||
can be arbitrarily stale (e.g. pre-/team/update models).
|
||||
|
||||
The builder used to write that team object back into `user_api_key_cache`
|
||||
under "team_id:<id>" after Check 6. Writing a cache-read (or worse, a
|
||||
token-snapshot) value back into the shared cache re-poisons it — with
|
||||
enable_redis_auth_cache it clobbered the fresh team `/team/update` had
|
||||
just written to Redis, making the stale allowlist self-sustaining across
|
||||
requests. Only authoritative writers (`_cache_team_object` on DB reads and
|
||||
team mutations) may populate the team cache.
|
||||
|
||||
Pins: the auth flow completes on the fallback path WITHOUT writing any
|
||||
"team_id:*" cache entry.
|
||||
"""
|
||||
from starlette.datastructures import URL
|
||||
from starlette.requests import Request
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder
|
||||
|
||||
api_key = "sk-test-lit-4391-no-team-writeback"
|
||||
valid_token = UserAPIKeyAuth(
|
||||
api_key=api_key,
|
||||
token=api_key,
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
team_id="team-lit-4391",
|
||||
team_models=["model-a"],
|
||||
models=["all-team-models"],
|
||||
)
|
||||
|
||||
mock_cache = AsyncMock()
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=valid_token)
|
||||
mock_cache.async_set_cache = AsyncMock(return_value=None)
|
||||
|
||||
mock_proxy_logging_obj = MagicMock()
|
||||
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
|
||||
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
|
||||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
||||
|
||||
import litellm.proxy.proxy_server as _proxy_server_mod
|
||||
|
||||
_attrs = {
|
||||
"prisma_client": MagicMock(),
|
||||
"user_api_key_cache": mock_cache,
|
||||
"proxy_logging_obj": mock_proxy_logging_obj,
|
||||
"master_key": "sk-master-key",
|
||||
"general_settings": {},
|
||||
"llm_model_list": [],
|
||||
"llm_router": None,
|
||||
"open_telemetry_logger": None,
|
||||
"model_max_budget_limiter": MagicMock(),
|
||||
"user_custom_auth": None,
|
||||
"jwt_handler": None,
|
||||
"litellm_proxy_admin_name": "admin",
|
||||
}
|
||||
_originals = {k: getattr(_proxy_server_mod, k, None) for k in _attrs}
|
||||
|
||||
try:
|
||||
for k, v in _attrs.items():
|
||||
setattr(_proxy_server_mod, k, v)
|
||||
|
||||
request = Request(scope={"type": "http"})
|
||||
request._url = URL(url="/chat/completions")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key",
|
||||
new_callable=AsyncMock,
|
||||
return_value=valid_token,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
side_effect=HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": "Team doesn't exist in db."},
|
||||
),
|
||||
),
|
||||
):
|
||||
result = await _user_api_key_auth_builder(
|
||||
request=request,
|
||||
api_key=f"Bearer {api_key}",
|
||||
azure_api_key_header="",
|
||||
anthropic_api_key_header=None,
|
||||
google_ai_studio_api_key_header=None,
|
||||
azure_apim_header=None,
|
||||
request_data={},
|
||||
)
|
||||
|
||||
assert result.team_id == "team-lit-4391"
|
||||
|
||||
team_cache_writes = [
|
||||
key
|
||||
for c in mock_cache.async_set_cache.await_args_list
|
||||
if isinstance(key := (c.kwargs.get("key") if "key" in c.kwargs else c.args[0]), str)
|
||||
and key.startswith("team_id:")
|
||||
]
|
||||
assert team_cache_writes == [], (
|
||||
"The auth flow wrote a team object into the cache. Fallback/"
|
||||
"cache-read team objects must never be persisted — only "
|
||||
"_cache_team_object (DB reads and team mutations) may write "
|
||||
f"'team_id:*' entries. Got writes: {team_cache_writes}"
|
||||
)
|
||||
|
||||
finally:
|
||||
for k, v in _originals.items():
|
||||
setattr(_proxy_server_mod, k, v)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# _run_centralized_common_checks — centralized authz gate
|
||||
|
|
@ -4576,101 +4693,6 @@ async def test_real_jwt_still_requires_license_when_jwt_auth_enabled(monkeypatch
|
|||
assert "enterprise only feature" in message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auth_path_caches_team_object_under_canonical_team_id_key():
|
||||
"""Regression for LIT-4000: the auth builder must cache the team object under
|
||||
the canonical ``team_id:{id}`` key that ``get_team_object`` and
|
||||
``_update_team_cache`` read, never under the raw ``team_id`` (and never under
|
||||
a ``None`` key, which Redis rejects with a NoneType key error). A raw or None
|
||||
key is silently dropped by Redis / never served back, so every request
|
||||
re-hits Postgres for the team object instead of the L2 cache.
|
||||
|
||||
Drives the real builder for a team-scoped key against a real in-memory
|
||||
``UserApiKeyCache`` and reads the team object back. Mutating the cache key at
|
||||
the write site to the raw ``valid_token.team_id`` (or ``None``) makes the
|
||||
canonical-key read miss and fails this test.
|
||||
"""
|
||||
from fastapi import Request
|
||||
from starlette.datastructures import URL
|
||||
|
||||
import litellm.proxy.proxy_server as _proxy_server_mod
|
||||
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
|
||||
from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.proxy_server import hash_token
|
||||
|
||||
team_id = "team-lit-4000"
|
||||
api_key = "sk-lit-4000-team-key"
|
||||
cache = UserApiKeyCache()
|
||||
|
||||
team_token = UserAPIKeyAuth(token=hash_token(api_key), team_id=team_id)
|
||||
team_obj = LiteLLM_TeamTableCachedObj(team_id=team_id)
|
||||
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
||||
attrs = {
|
||||
"prisma_client": MagicMock(),
|
||||
"user_api_key_cache": cache,
|
||||
"proxy_logging_obj": proxy_logging_obj,
|
||||
"master_key": "sk-test-master",
|
||||
"general_settings": {"allow_requests_on_db_unavailable": False},
|
||||
"llm_model_list": [],
|
||||
"llm_router": None,
|
||||
"open_telemetry_logger": None,
|
||||
"model_max_budget_limiter": MagicMock(),
|
||||
"user_custom_auth": None,
|
||||
"jwt_handler": None,
|
||||
"litellm_proxy_admin_name": "admin",
|
||||
}
|
||||
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
|
||||
try:
|
||||
for k, v in attrs.items():
|
||||
setattr(_proxy_server_mod, k, v)
|
||||
request = Request(scope={"type": "http"})
|
||||
request._url = URL(url="/chat/completions")
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key",
|
||||
AsyncMock(return_value=team_token),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.get_team_object",
|
||||
AsyncMock(return_value=team_obj),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth._enforce_key_and_fallback_model_access",
|
||||
new_callable=AsyncMock,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth._return_user_api_key_auth_obj",
|
||||
new_callable=AsyncMock,
|
||||
return_value=team_token,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.auth.auth_exception_handler.seed_request_identity",
|
||||
),
|
||||
):
|
||||
await _user_api_key_auth_builder(
|
||||
request=request,
|
||||
api_key=f"Bearer {api_key}",
|
||||
azure_api_key_header="",
|
||||
anthropic_api_key_header=None,
|
||||
google_ai_studio_api_key_header=None,
|
||||
azure_apim_header=None,
|
||||
request_data={},
|
||||
)
|
||||
finally:
|
||||
for k, v in originals.items():
|
||||
setattr(_proxy_server_mod, k, v)
|
||||
|
||||
served = cache.get_cache(
|
||||
key=f"team_id:{team_id}", model_type=LiteLLM_TeamTableCachedObj
|
||||
)
|
||||
assert served is not None and served.team_id == team_id
|
||||
assert cache.get_cache(key=team_id) is None
|
||||
assert cache.get_cache(key=None) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auth_does_not_rewrite_cached_key_object_back_into_cache():
|
||||
"""A cache-hit auth must not write the token back into the cache.
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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"] == []
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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"]}}
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -152,7 +152,7 @@
|
|||
"count": 1
|
||||
},
|
||||
"prefer-const": {
|
||||
"count": 3
|
||||
"count": 1
|
||||
},
|
||||
"react-hooks/purity": {
|
||||
"count": 1
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -19,13 +19,29 @@ import {
|
|||
import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs";
|
||||
|
||||
import { RefreshCw } from "lucide-react";
|
||||
import { adminGlobalCacheActivity, cachingHealthCheckCall } from "@/components/networking";
|
||||
import { cachingHealthCheckCall } from "@/components/networking";
|
||||
import { useCacheActivity, type CacheActivityGroup } from "@/app/(dashboard)/hooks/caching/useCacheActivity";
|
||||
|
||||
// Import the new component
|
||||
import { CacheHealthTab } from "./cache_health";
|
||||
import CacheSettings from "./cache_settings";
|
||||
import CoordinationRedisSettings from "./coordination_redis_settings";
|
||||
|
||||
const REQUEST_SERIES = {
|
||||
apiRequests: "LLM API requests",
|
||||
cacheHits: "Cache hit",
|
||||
failed: "Failed requests",
|
||||
} as const;
|
||||
|
||||
const toChartDatum = (group: CacheActivityGroup) => ({
|
||||
name: group.call_type,
|
||||
[REQUEST_SERIES.apiRequests]: group.api_requests,
|
||||
[REQUEST_SERIES.cacheHits]: group.cache_hits,
|
||||
[REQUEST_SERIES.failed]: group.failed_requests,
|
||||
"Cached Completion Tokens": group.cached_completion_tokens,
|
||||
"Generated Completion Tokens": group.generated_completion_tokens,
|
||||
});
|
||||
|
||||
const formatDateWithoutTZ = (date: Date | undefined) => {
|
||||
if (!date) return undefined;
|
||||
return date.toISOString().split("T")[0];
|
||||
|
|
@ -49,26 +65,6 @@ interface CachePageProps {
|
|||
premiumUser: boolean;
|
||||
}
|
||||
|
||||
interface cacheDataItem {
|
||||
api_key: string;
|
||||
model: string;
|
||||
cache_hit_true_rows: number;
|
||||
cached_completion_tokens: number;
|
||||
total_rows: number;
|
||||
generated_completion_tokens: number;
|
||||
call_type: string;
|
||||
|
||||
// Add other properties as needed
|
||||
}
|
||||
|
||||
type uiData = {
|
||||
name: string;
|
||||
"LLM API requests": number;
|
||||
"Cache hit": number;
|
||||
"Cached Completion Tokens": number;
|
||||
"Generated Completion Tokens": number;
|
||||
};
|
||||
|
||||
interface CacheHealthResponse {
|
||||
status?: string;
|
||||
cache_type?: string;
|
||||
|
|
@ -97,13 +93,8 @@ const deepParse = (input: any) => {
|
|||
};
|
||||
|
||||
const CacheDashboard: React.FC<CachePageProps> = ({ accessToken, token, userRole, userID, premiumUser }) => {
|
||||
const [filteredData, setFilteredData] = useState<uiData[]>([]);
|
||||
const [selectedApiKeys, setSelectedApiKeys] = useState<string[]>([]);
|
||||
const [selectedModels, setSelectedModels] = useState<string[]>([]);
|
||||
const [data, setData] = useState<cacheDataItem[]>([]);
|
||||
const [cachedResponses, setCachedResponses] = useState("0");
|
||||
const [cachedTokens, setCachedTokens] = useState("0");
|
||||
const [cacheHitRatio, setCacheHitRatio] = useState("0");
|
||||
|
||||
const [dateValue, setDateValue] = useState<DateRangePickerValue>({
|
||||
from: new Date(Date.now() - 7 * 24 * 60 * 60 * 1000),
|
||||
|
|
@ -113,120 +104,24 @@ const CacheDashboard: React.FC<CachePageProps> = ({ accessToken, token, userRole
|
|||
const [lastRefreshed, setLastRefreshed] = useState("");
|
||||
const [healthCheckResponse, setHealthCheckResponse] = useState<any>("");
|
||||
|
||||
useEffect(() => {
|
||||
if (!accessToken || !dateValue) {
|
||||
return;
|
||||
}
|
||||
const fetchData = async () => {
|
||||
const response = await adminGlobalCacheActivity(
|
||||
accessToken,
|
||||
formatDateWithoutTZ(dateValue.from),
|
||||
formatDateWithoutTZ(dateValue.to),
|
||||
);
|
||||
setData(response);
|
||||
};
|
||||
fetchData();
|
||||
|
||||
const currentDate = new Date();
|
||||
setLastRefreshed(currentDate.toLocaleString());
|
||||
}, [accessToken]);
|
||||
|
||||
const uniqueApiKeys = Array.from(new Set(data.map((item) => item?.api_key ?? "")));
|
||||
const uniqueModels = Array.from(new Set(data.map((item) => item?.model ?? "")));
|
||||
const uniqueCallTypes = Array.from(new Set(data.map((item) => item?.call_type ?? "")));
|
||||
|
||||
const updateCachingData = async (startTime: Date | undefined, endTime: Date | undefined) => {
|
||||
if (!startTime || !endTime || !accessToken) {
|
||||
return;
|
||||
}
|
||||
|
||||
let new_cache_data = await adminGlobalCacheActivity(
|
||||
accessToken,
|
||||
formatDateWithoutTZ(startTime),
|
||||
formatDateWithoutTZ(endTime),
|
||||
);
|
||||
|
||||
setData(new_cache_data);
|
||||
};
|
||||
const { data: activity, refetch } = useCacheActivity({
|
||||
startDate: formatDateWithoutTZ(dateValue.from),
|
||||
endDate: formatDateWithoutTZ(dateValue.to),
|
||||
keyAliases: selectedApiKeys,
|
||||
models: selectedModels,
|
||||
});
|
||||
|
||||
useEffect(() => {
|
||||
let newData: cacheDataItem[] = data;
|
||||
if (selectedApiKeys.length > 0) {
|
||||
newData = newData.filter((item) => selectedApiKeys.includes(item.api_key));
|
||||
}
|
||||
setLastRefreshed(new Date().toLocaleString());
|
||||
}, []);
|
||||
|
||||
if (selectedModels.length > 0) {
|
||||
newData = newData.filter((item) => selectedModels.includes(item.model));
|
||||
}
|
||||
|
||||
/*
|
||||
Data looks like this
|
||||
[{"api_key":"sk-test-mock-key-001","call_type":"acompletion","model":"llama3-8b-8192","total_rows":13,"cache_hit_true_rows":0},
|
||||
{"api_key":"sk-test-mock-key-002","call_type":"None","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0},
|
||||
{"api_key":"sk-test-mock-key-123","call_type":"acompletion","model":"gpt-3.5-turbo","total_rows":19,"cache_hit_true_rows":0},
|
||||
{"api_key":"sk-test-mock-key-123","call_type":"aimage_generation","model":"","total_rows":3,"cache_hit_true_rows":0},
|
||||
{"api_key":"sk-test-mock-key-003","call_type":"None","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0},
|
||||
{"api_key":"sk-test-mock-key-004","call_type":"","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0},
|
||||
{"api_key":"sk-test-mock-key-005","call_type":"","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0},
|
||||
*/
|
||||
|
||||
// What data we need for bar chat
|
||||
// ui_data = [
|
||||
// {
|
||||
// name: "Call Type",
|
||||
// Cache hit: 20,
|
||||
// LLM API requests: 10,
|
||||
// }
|
||||
// ]
|
||||
|
||||
let llm_api_requests = 0;
|
||||
let cache_hits = 0;
|
||||
let cached_tokens = 0;
|
||||
const processedData = newData.reduce((acc: uiData[], item) => {
|
||||
if (!item.call_type) {
|
||||
item.call_type = "Unknown";
|
||||
}
|
||||
|
||||
llm_api_requests += (item.total_rows || 0) - (item.cache_hit_true_rows || 0);
|
||||
cache_hits += item.cache_hit_true_rows || 0;
|
||||
cached_tokens += item.cached_completion_tokens || 0;
|
||||
|
||||
const existingItem = acc.find((i) => i.name === item.call_type);
|
||||
if (existingItem) {
|
||||
existingItem["LLM API requests"] += (item.total_rows || 0) - (item.cache_hit_true_rows || 0);
|
||||
existingItem["Cache hit"] += item.cache_hit_true_rows || 0;
|
||||
existingItem["Cached Completion Tokens"] += item.cached_completion_tokens || 0;
|
||||
existingItem["Generated Completion Tokens"] += item.generated_completion_tokens || 0;
|
||||
} else {
|
||||
acc.push({
|
||||
name: item.call_type,
|
||||
"LLM API requests": (item.total_rows || 0) - (item.cache_hit_true_rows || 0),
|
||||
"Cache hit": item.cache_hit_true_rows || 0,
|
||||
"Cached Completion Tokens": item.cached_completion_tokens || 0,
|
||||
"Generated Completion Tokens": item.generated_completion_tokens || 0,
|
||||
});
|
||||
}
|
||||
return acc;
|
||||
}, []);
|
||||
|
||||
// set header cache statistics
|
||||
setCachedResponses(valueFormatterNumbers(cache_hits));
|
||||
setCachedTokens(valueFormatterNumbers(cached_tokens));
|
||||
let allRequests = cache_hits + llm_api_requests;
|
||||
if (allRequests > 0) {
|
||||
let cache_hit_ratio = ((cache_hits / allRequests) * 100).toFixed(2);
|
||||
setCacheHitRatio(cache_hit_ratio);
|
||||
} else {
|
||||
setCacheHitRatio("0");
|
||||
}
|
||||
|
||||
setFilteredData(processedData);
|
||||
}, [selectedApiKeys, selectedModels, dateValue, data]);
|
||||
const uniqueApiKeys = activity?.filter_options.key_aliases ?? [];
|
||||
const uniqueModels = activity?.filter_options.models ?? [];
|
||||
const chartData = (activity?.groups ?? []).map(toChartDatum);
|
||||
|
||||
const handleRefreshClick = () => {
|
||||
// Update the 'lastRefreshed' state to the current date and time
|
||||
const currentDate = new Date();
|
||||
setLastRefreshed(currentDate.toLocaleString());
|
||||
refetch();
|
||||
setLastRefreshed(new Date().toLocaleString());
|
||||
};
|
||||
|
||||
const runCachingHealthCheck = async () => {
|
||||
|
|
@ -257,10 +152,12 @@ const CacheDashboard: React.FC<CachePageProps> = ({ accessToken, token, userRole
|
|||
}
|
||||
};
|
||||
|
||||
const totals = activity?.totals;
|
||||
const hasRequests = totals != null && totals.api_requests + totals.cache_hits + totals.failed_requests > 0;
|
||||
const statCards = [
|
||||
{ label: "Cache Hit Ratio", value: `${cacheHitRatio}%` },
|
||||
{ label: "Cache Hits", value: cachedResponses },
|
||||
{ label: "Cached Completion Tokens", value: cachedTokens },
|
||||
{ label: "Cache Hit Ratio", value: `${hasRequests ? totals.cache_hit_ratio.toFixed(2) : "0"}%` },
|
||||
{ label: "Cache Hits", value: valueFormatterNumbers(totals?.cache_hits ?? 0) },
|
||||
{ label: "Cached Completion Tokens", value: valueFormatterNumbers(totals?.cached_completion_tokens ?? 0) },
|
||||
];
|
||||
|
||||
return (
|
||||
|
|
@ -380,7 +277,6 @@ const CacheDashboard: React.FC<CachePageProps> = ({ accessToken, token, userRole
|
|||
value={dateValue}
|
||||
onValueChange={(value) => {
|
||||
setDateValue(value);
|
||||
updateCachingData(value.from, value.to);
|
||||
}}
|
||||
/>
|
||||
</div>
|
||||
|
|
@ -404,12 +300,12 @@ const CacheDashboard: React.FC<CachePageProps> = ({ accessToken, token, userRole
|
|||
</CardHeader>
|
||||
<CardContent>
|
||||
<BarChart
|
||||
data={filteredData}
|
||||
data={chartData}
|
||||
stack={true}
|
||||
index="name"
|
||||
valueFormatter={valueFormatterNumbers}
|
||||
categories={["LLM API requests", "Cache hit"]}
|
||||
colors={["sky", "teal"]}
|
||||
categories={[REQUEST_SERIES.apiRequests, REQUEST_SERIES.cacheHits, REQUEST_SERIES.failed]}
|
||||
colors={["sky", "teal", "red"]}
|
||||
yAxisWidth={48}
|
||||
/>
|
||||
</CardContent>
|
||||
|
|
@ -423,7 +319,7 @@ const CacheDashboard: React.FC<CachePageProps> = ({ accessToken, token, userRole
|
|||
</CardHeader>
|
||||
<CardContent>
|
||||
<BarChart
|
||||
data={filteredData}
|
||||
data={chartData}
|
||||
stack={true}
|
||||
index="name"
|
||||
valueFormatter={valueFormatterNumbers}
|
||||
|
|
|
|||
|
|
@ -0,0 +1,72 @@
|
|||
import { renderHook } from "@testing-library/react";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import { useCacheActivity, type CacheActivityParams } from "./useCacheActivity";
|
||||
|
||||
const useQueryMock = vi.fn();
|
||||
vi.mock("@/lib/http/api", () => ({
|
||||
$api: { useQuery: (...args: unknown[]) => useQueryMock(...args) },
|
||||
}));
|
||||
|
||||
const mockUseAuthorized = vi.fn();
|
||||
vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({
|
||||
default: () => mockUseAuthorized(),
|
||||
}));
|
||||
|
||||
const params: CacheActivityParams = {
|
||||
startDate: "2026-07-20",
|
||||
endDate: "2026-07-27",
|
||||
keyAliases: ["my-key"],
|
||||
models: ["gpt-5.1"],
|
||||
};
|
||||
|
||||
const lastCallOptions = (): { enabled: boolean } => {
|
||||
const calls = useQueryMock.mock.calls;
|
||||
return calls[calls.length - 1][3] as { enabled: boolean };
|
||||
};
|
||||
|
||||
describe("useCacheActivity", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
useQueryMock.mockReturnValue({ data: undefined });
|
||||
mockUseAuthorized.mockReturnValue({ accessToken: "test-access-token" });
|
||||
});
|
||||
|
||||
it("queries GET /global/activity/cache_hits with dates and filters as query params", () => {
|
||||
renderHook(() => useCacheActivity(params));
|
||||
|
||||
expect(useQueryMock).toHaveBeenCalledWith(
|
||||
"get",
|
||||
"/global/activity/cache_hits",
|
||||
{
|
||||
params: {
|
||||
query: {
|
||||
start_date: "2026-07-20",
|
||||
end_date: "2026-07-27",
|
||||
key_aliases: ["my-key"],
|
||||
models: ["gpt-5.1"],
|
||||
},
|
||||
},
|
||||
},
|
||||
expect.any(Object),
|
||||
);
|
||||
});
|
||||
|
||||
it("enables the query when authorized and both dates are set", () => {
|
||||
renderHook(() => useCacheActivity(params));
|
||||
|
||||
expect(lastCallOptions().enabled).toBe(true);
|
||||
});
|
||||
|
||||
it("disables the query without an access token", () => {
|
||||
mockUseAuthorized.mockReturnValue({ accessToken: null });
|
||||
renderHook(() => useCacheActivity(params));
|
||||
|
||||
expect(lastCallOptions().enabled).toBe(false);
|
||||
});
|
||||
|
||||
it("disables the query while the date range is incomplete", () => {
|
||||
renderHook(() => useCacheActivity({ ...params, endDate: undefined }));
|
||||
|
||||
expect(lastCallOptions().enabled).toBe(false);
|
||||
});
|
||||
});
|
||||
|
|
@ -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) },
|
||||
);
|
||||
};
|
||||
|
|
@ -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);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,7 +1,9 @@
|
|||
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
|
||||
import { render, screen } from "@testing-library/react";
|
||||
import { act, render, screen } from "@testing-library/react";
|
||||
import React from "react";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import type OrganizationsTableComponent from "./OrganizationsTable";
|
||||
import type OrganizationInfoViewComponent from "@/components/organization/organization_view";
|
||||
|
||||
vi.mock("@/components/vector_store_management/VectorStoreSelector", () => ({
|
||||
__esModule: true,
|
||||
|
|
@ -18,12 +20,50 @@ vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({
|
|||
userRole: null,
|
||||
}),
|
||||
}));
|
||||
type OrganizationsTableProps = React.ComponentProps<typeof OrganizationsTableComponent>;
|
||||
type OrganizationInfoViewProps = React.ComponentProps<typeof OrganizationInfoViewComponent>;
|
||||
|
||||
let capturedTableProps: OrganizationsTableProps | null = null;
|
||||
vi.mock("./OrganizationsTable", () => ({
|
||||
__esModule: true,
|
||||
default: (props: { isLoading: boolean }) => (
|
||||
<div data-testid="organizations-table">isLoading:{String(props.isLoading)}</div>
|
||||
),
|
||||
default: (props: OrganizationsTableProps) => {
|
||||
capturedTableProps = props;
|
||||
return <div data-testid="organizations-table">isLoading:{String(props.isLoading)}</div>;
|
||||
},
|
||||
}));
|
||||
const mockOrgInfoView = vi.fn<(props: OrganizationInfoViewProps) => void>();
|
||||
vi.mock("@/components/organization/organization_view", () => ({
|
||||
__esModule: true,
|
||||
default: (props: OrganizationInfoViewProps) => {
|
||||
mockOrgInfoView(props);
|
||||
return <div data-testid="organization-info-view" />;
|
||||
},
|
||||
}));
|
||||
|
||||
// The selected org is URL-derived (?org=) via useOrgDetailRouting. Next's real useSearchParams
|
||||
// re-renders subscribers on history.pushState/replaceState; mirror that so URL changes propagate.
|
||||
vi.mock("next/navigation", async () => {
|
||||
const { useSyncExternalStore } = await import("react");
|
||||
const LOCATION_CHANGE_EVENT = "test-locationchange";
|
||||
for (const method of ["pushState", "replaceState"] as const) {
|
||||
const original = window.history[method].bind(window.history);
|
||||
window.history[method] = (...args: Parameters<History["pushState"]>) => {
|
||||
original(...args);
|
||||
window.dispatchEvent(new Event(LOCATION_CHANGE_EVENT));
|
||||
};
|
||||
}
|
||||
const subscribe = (onChange: () => void) => {
|
||||
window.addEventListener(LOCATION_CHANGE_EVENT, onChange);
|
||||
window.addEventListener("popstate", onChange);
|
||||
return () => {
|
||||
window.removeEventListener(LOCATION_CHANGE_EVENT, onChange);
|
||||
window.removeEventListener("popstate", onChange);
|
||||
};
|
||||
};
|
||||
return {
|
||||
useSearchParams: () => new URLSearchParams(useSyncExternalStore(subscribe, () => window.location.search)),
|
||||
};
|
||||
});
|
||||
|
||||
import OrganizationsPanel from "./OrganizationsPanel";
|
||||
|
||||
|
|
@ -34,6 +74,12 @@ const renderWithQueryClient = (ui: React.ReactElement) => {
|
|||
return render(<QueryClientProvider client={queryClient}>{ui}</QueryClientProvider>);
|
||||
};
|
||||
|
||||
beforeEach(() => {
|
||||
capturedTableProps = null;
|
||||
mockOrgInfoView.mockClear();
|
||||
window.history.replaceState(null, "", "/organizations/");
|
||||
});
|
||||
|
||||
describe("OrganizationsPanel", () => {
|
||||
it("gates non-premium users behind the enterprise notice", () => {
|
||||
renderWithQueryClient(<OrganizationsPanel userRole="Admin" accessToken={null} premiumUser={false} />);
|
||||
|
|
@ -55,3 +101,60 @@ describe("OrganizationsPanel", () => {
|
|||
expect(screen.getByTestId("organizations-table")).toHaveTextContent("isLoading:false");
|
||||
});
|
||||
});
|
||||
|
||||
describe("OrganizationsPanel - org detail deep link (?org=)", () => {
|
||||
it("clicking an organization pushes ?org= and opens the detail view", () => {
|
||||
renderWithQueryClient(<OrganizationsPanel userRole="Admin" accessToken={null} premiumUser={true} />);
|
||||
|
||||
act(() => capturedTableProps?.onOrganizationClick("org-deep-link"));
|
||||
|
||||
expect(window.location.search).toContain("org=org-deep-link");
|
||||
expect(mockOrgInfoView).toHaveBeenLastCalledWith(expect.objectContaining({ organizationId: "org-deep-link" }));
|
||||
});
|
||||
|
||||
it("opens the org detail directly from a ?org= deep link", () => {
|
||||
window.history.replaceState(null, "", "/organizations/?org=org-from-url");
|
||||
renderWithQueryClient(<OrganizationsPanel userRole="Admin" accessToken={null} premiumUser={true} />);
|
||||
|
||||
expect(mockOrgInfoView).toHaveBeenLastCalledWith(
|
||||
expect.objectContaining({ organizationId: "org-from-url", editOrg: false }),
|
||||
);
|
||||
expect(screen.queryByTestId("organizations-table")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("closing the org detail removes ?org= and returns to the list", () => {
|
||||
window.history.replaceState(null, "", "/organizations/?org=org-from-url");
|
||||
renderWithQueryClient(<OrganizationsPanel userRole="Admin" accessToken={null} premiumUser={true} />);
|
||||
|
||||
act(() => mockOrgInfoView.mock.calls.at(-1)?.[0].onClose());
|
||||
|
||||
expect(window.location.search).not.toContain("org=");
|
||||
expect(screen.queryByTestId("organization-info-view")).not.toBeInTheDocument();
|
||||
expect(screen.getByTestId("organizations-table")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("the edit action opens the detail in edit mode with ?org= set", () => {
|
||||
renderWithQueryClient(<OrganizationsPanel userRole="Admin" accessToken={null} premiumUser={true} />);
|
||||
|
||||
act(() => capturedTableProps?.onEditClick("org-edit"));
|
||||
|
||||
expect(window.location.search).toContain("org=org-edit");
|
||||
expect(mockOrgInfoView).toHaveBeenLastCalledWith(
|
||||
expect.objectContaining({ organizationId: "org-edit", editOrg: true }),
|
||||
);
|
||||
});
|
||||
|
||||
it("a plain row click after leaving an edit view via browser history does not reopen in edit mode", () => {
|
||||
renderWithQueryClient(<OrganizationsPanel userRole="Admin" accessToken={null} premiumUser={true} />);
|
||||
|
||||
act(() => capturedTableProps?.onEditClick("org-edit"));
|
||||
expect(mockOrgInfoView).toHaveBeenLastCalledWith(expect.objectContaining({ editOrg: true }));
|
||||
|
||||
act(() => window.history.pushState(null, "", "/organizations/"));
|
||||
act(() => capturedTableProps?.onOrganizationClick("org-plain"));
|
||||
|
||||
expect(mockOrgInfoView).toHaveBeenLastCalledWith(
|
||||
expect.objectContaining({ organizationId: "org-plain", editOrg: false }),
|
||||
);
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import { organizationKeys, useOrganizations } from "@/app/(dashboard)/hooks/organizations/useOrganizations";
|
||||
import { useUserModels } from "@/app/(dashboard)/hooks/models/useModels";
|
||||
import { useOrgDetailRouting } from "@/app/(dashboard)/organizations/detailNavigation";
|
||||
import OrganizationFilters, { FilterState } from "@/app/(dashboard)/organizations/OrganizationFilters";
|
||||
import { useQueryClient } from "@tanstack/react-query";
|
||||
import React, { useState } from "react";
|
||||
|
|
@ -19,7 +20,7 @@ interface OrganizationsPanelProps {
|
|||
}
|
||||
|
||||
const OrganizationsPanel: React.FC<OrganizationsPanelProps> = ({ userRole, accessToken, premiumUser }) => {
|
||||
const [selectedOrgId, setSelectedOrgId] = useState<string | null>(null);
|
||||
const { orgId: selectedOrgId, openOrg, close: closeOrgDetail } = useOrgDetailRouting();
|
||||
const [editOrg, setEditOrg] = useState(false);
|
||||
const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false);
|
||||
const [orgToDelete, setOrgToDelete] = useState<string | null>(null);
|
||||
|
|
@ -108,7 +109,7 @@ const OrganizationsPanel: React.FC<OrganizationsPanelProps> = ({ userRole, acces
|
|||
<OrganizationInfoView
|
||||
organizationId={selectedOrgId}
|
||||
onClose={() => {
|
||||
setSelectedOrgId(null);
|
||||
closeOrgDetail();
|
||||
setEditOrg(false);
|
||||
}}
|
||||
accessToken={accessToken}
|
||||
|
|
@ -132,9 +133,12 @@ const OrganizationsPanel: React.FC<OrganizationsPanelProps> = ({ userRole, acces
|
|||
isLoading={isLoading}
|
||||
userRole={userRole}
|
||||
searchActive={searchActive}
|
||||
onOrganizationClick={setSelectedOrgId}
|
||||
onOrganizationClick={(organizationId) => {
|
||||
setEditOrg(false);
|
||||
openOrg(organizationId);
|
||||
}}
|
||||
onEditClick={(organizationId) => {
|
||||
setSelectedOrgId(organizationId);
|
||||
openOrg(organizationId);
|
||||
setEditOrg(true);
|
||||
}}
|
||||
onDeleteClick={handleDelete}
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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,
|
||||
};
|
||||
}
|
||||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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,
|
||||
};
|
||||
}
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
import { act, fireEvent, render, screen, waitFor } from "@testing-library/react";
|
||||
import { act, fireEvent, render, screen, waitFor, within } from "@testing-library/react";
|
||||
import { beforeAll, beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import * as networking from "@/components/networking";
|
||||
import EntityUsage from "./EntityUsage";
|
||||
|
|
@ -497,7 +497,7 @@ describe("EntityUsage", () => {
|
|||
|
||||
it.each([
|
||||
["Cost", "Tag Spend Overview"],
|
||||
["Model Activity", "metrics-source:models"],
|
||||
["Model Activity", "metrics-source:model_groups"],
|
||||
["Key Activity", "metrics-source:api_keys"],
|
||||
["Endpoint Activity", "Endpoint Usage Panel"],
|
||||
])("shows only the %s panel for a non-team entity type", async (tabLabel, marker) => {
|
||||
|
|
@ -518,7 +518,7 @@ describe("EntityUsage", () => {
|
|||
|
||||
it.each([
|
||||
["Cost", "Team Spend Overview"],
|
||||
["Model Activity", "metrics-source:models"],
|
||||
["Model Activity", "metrics-source:model_groups"],
|
||||
["Agent Activity", "metrics-source:entities"],
|
||||
["Key Activity", "metrics-source:api_keys"],
|
||||
["Endpoint Activity", "Endpoint Usage Panel"],
|
||||
|
|
@ -584,15 +584,41 @@ describe("EntityUsage", () => {
|
|||
expect(screen.getByText("Request / Token Consumption")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display Top Models title for non-agent entity types", async () => {
|
||||
it("should display Top Public Model Names title for non-agent entity types", async () => {
|
||||
render(<EntityUsage {...defaultProps} entityType="tag" />);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(mockTagDailyActivityCall).toHaveBeenCalled();
|
||||
});
|
||||
|
||||
const topModelsElements = screen.getAllByText("Top Models");
|
||||
expect(topModelsElements.length).toBeGreaterThan(0);
|
||||
expect(screen.getByText("Top Public Model Names")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("defaults Model Activity to public model names and toggles to litellm models", async () => {
|
||||
const { container } = render(<EntityUsage {...defaultProps} />);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(mockTagDailyActivityCall).toHaveBeenCalled();
|
||||
});
|
||||
|
||||
act(() => {
|
||||
fireEvent.click(screen.getByText("Model Activity"));
|
||||
});
|
||||
|
||||
const modelActivityPanel = () => selectedPanels(container)[0] as HTMLElement;
|
||||
expect(modelActivityPanel().textContent).toContain("metrics-source:model_groups");
|
||||
|
||||
act(() => {
|
||||
fireEvent.click(within(modelActivityPanel()).getByText("Litellm Model Name"));
|
||||
});
|
||||
|
||||
expect(modelActivityPanel().textContent).toContain("metrics-source:models");
|
||||
|
||||
act(() => {
|
||||
fireEvent.click(within(modelActivityPanel()).getByText("Public Model Name"));
|
||||
});
|
||||
|
||||
expect(modelActivityPanel().textContent).toContain("metrics-source:model_groups");
|
||||
});
|
||||
|
||||
it("should display Top Agents title for agent entity type", async () => {
|
||||
|
|
|
|||
|
|
@ -49,6 +49,7 @@ import {
|
|||
} from "@/components/UsagePage/types";
|
||||
import { valueFormatterSpend } from "@/components/UsagePage/utils/value_formatters";
|
||||
import EndpointUsage from "../EndpointUsage/EndpointUsage";
|
||||
import ModelViewToggle, { ModelViewType } from "../ModelViewToggle";
|
||||
import TopKeyView from "@/components/UsagePage/components/EntityUsage/TopKeyView";
|
||||
import TopModelView from "./TopModelView";
|
||||
|
||||
|
|
@ -110,6 +111,7 @@ const ENTITY_FETCH_FNS: Record<EntityType, (...args: any[]) => Promise<any>> = {
|
|||
const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, entityId, entityList, dateValue }) => {
|
||||
const { teams } = useTeams();
|
||||
const [selectedTags, setSelectedTags] = useState<string[]>([]);
|
||||
const [modelViewType, setModelViewType] = useState<ModelViewType>("groups");
|
||||
const [topKeysLimit, setTopKeysLimit] = useState<number>(5);
|
||||
const [topModelsLimit, setTopModelsLimit] = useState<number>(5);
|
||||
const [topAgentsLimit, setTopAgentsLimit] = useState<number>(5);
|
||||
|
|
@ -153,14 +155,15 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
|
|||
|
||||
const agentSpendData = agentSpendDataRaw as unknown as EntitySpendData;
|
||||
|
||||
const modelMetrics = processActivityData(spendData, "models", teams || []);
|
||||
const modelBreakdownKey = modelViewType === "groups" ? "model_groups" : "models";
|
||||
const modelMetrics = processActivityData(spendData, modelBreakdownKey, teams || []);
|
||||
const keyMetrics = processActivityData(spendData, "api_keys", teams || []);
|
||||
const agentMetrics = entityType === "team" ? processActivityData(agentSpendData, "entities", teams || []) : {};
|
||||
|
||||
const getTopModels = () => {
|
||||
const modelSpend: { [key: string]: any } = {};
|
||||
spendData.results.forEach((day) => {
|
||||
Object.entries(day.breakdown.models || {}).forEach(([model, metrics]) => {
|
||||
Object.entries(day.breakdown[modelBreakdownKey] || {}).forEach(([model, metrics]) => {
|
||||
if (!modelSpend[model]) {
|
||||
modelSpend[model] = {
|
||||
spend: 0,
|
||||
|
|
@ -406,6 +409,8 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
|
|||
|
||||
const capitalizedEntityLabel = entityType.charAt(0).toUpperCase() + entityType.slice(1);
|
||||
|
||||
const modelViewTitle = modelViewType === "groups" ? "Top Public Model Names" : "Top Litellm Models";
|
||||
|
||||
const costPanel = (
|
||||
<Grid numItems={2} className="gap-2 w-full">
|
||||
{/* Total Spend Card */}
|
||||
|
|
@ -604,7 +609,10 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
|
|||
{/* Top Models */}
|
||||
<Col numColSpan={1}>
|
||||
<Card>
|
||||
<Title>{entityType === "agent" ? "Top Agents" : "Top Models"}</Title>
|
||||
<div className="flex justify-between items-center">
|
||||
<Title>{entityType === "agent" ? "Top Agents" : modelViewTitle}</Title>
|
||||
<ModelViewToggle value={modelViewType} onChange={setModelViewType} />
|
||||
</div>
|
||||
<TopModelView
|
||||
topModels={getTopModels()}
|
||||
topModelsLimit={topModelsLimit}
|
||||
|
|
@ -691,7 +699,14 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
|
|||
{
|
||||
key: "models",
|
||||
label: entityType === "agent" ? "Request / Token Consumption" : "Model Activity",
|
||||
content: <ActivityMetrics modelMetrics={modelMetrics} hidePromptCachingMetrics={entityType === "agent"} />,
|
||||
content: (
|
||||
<>
|
||||
<div className="flex justify-end mt-2 mb-4">
|
||||
<ModelViewToggle value={modelViewType} onChange={setModelViewType} />
|
||||
</div>
|
||||
<ActivityMetrics modelMetrics={modelMetrics} hidePromptCachingMetrics={entityType === "agent"} />
|
||||
</>
|
||||
),
|
||||
},
|
||||
...(entityType === "team"
|
||||
? [{ key: "agents", label: "Agent Activity", content: <ActivityMetrics modelMetrics={agentMetrics} /> }]
|
||||
|
|
|
|||
|
|
@ -0,0 +1,29 @@
|
|||
export type ModelViewType = "groups" | "individual";
|
||||
|
||||
const MODEL_VIEW_OPTIONS: readonly { value: ModelViewType; label: string }[] = [
|
||||
{ value: "groups", label: "Public Model Name" },
|
||||
{ value: "individual", label: "Litellm Model Name" },
|
||||
];
|
||||
|
||||
interface ModelViewToggleProps {
|
||||
value: ModelViewType;
|
||||
onChange: (value: ModelViewType) => void;
|
||||
}
|
||||
|
||||
export default function ModelViewToggle({ value, onChange }: ModelViewToggleProps) {
|
||||
return (
|
||||
<div className="flex bg-gray-100 rounded-lg p-1">
|
||||
{MODEL_VIEW_OPTIONS.map((option) => (
|
||||
<button
|
||||
key={option.value}
|
||||
className={`px-3 py-1 text-sm rounded-md transition-colors ${
|
||||
value === option.value ? "bg-white shadow-xs text-gray-900" : "text-gray-600 hover:text-gray-900"
|
||||
}`}
|
||||
onClick={() => onChange(option.value)}
|
||||
>
|
||||
{option.label}
|
||||
</button>
|
||||
))}
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
|
@ -30,8 +30,10 @@ vi.mock("@/components/networking", () => ({
|
|||
|
||||
// Mock child components to simplify testing
|
||||
vi.mock("@/components/activity_metrics", () => ({
|
||||
ActivityMetrics: () => <div>Activity Metrics</div>,
|
||||
processActivityData: () => ({ data: [], metadata: {} }),
|
||||
ActivityMetrics: ({ modelMetrics }: { modelMetrics?: { __source?: string } }) => (
|
||||
<div>{`activity-source:${modelMetrics?.__source ?? "none"}`}</div>
|
||||
),
|
||||
processActivityData: (_data: unknown, key: string) => ({ __source: key }),
|
||||
}));
|
||||
|
||||
vi.mock("@/components/view_user_spend", () => ({
|
||||
|
|
@ -1043,8 +1045,8 @@ describe("UsagePage", () => {
|
|||
|
||||
// Default should be "groups" view showing "Top Public Model Names"
|
||||
expect(screen.getByText("Top Public Model Names")).toBeInTheDocument();
|
||||
expect(screen.getByText("Public Model Name")).toBeInTheDocument();
|
||||
expect(screen.getByText("Litellm Model Name")).toBeInTheDocument();
|
||||
expect(screen.getAllByText("Public Model Name").length).toBeGreaterThan(0);
|
||||
expect(screen.getAllByText("Litellm Model Name").length).toBeGreaterThan(0);
|
||||
});
|
||||
|
||||
it("should switch to Litellm Model Name view on toggle click", async () => {
|
||||
|
|
@ -1055,7 +1057,7 @@ describe("UsagePage", () => {
|
|||
});
|
||||
|
||||
// Click the "Litellm Model Name" toggle
|
||||
const litellmToggle = screen.getByText("Litellm Model Name");
|
||||
const litellmToggle = screen.getAllByText("Litellm Model Name")[0];
|
||||
act(() => {
|
||||
fireEvent.click(litellmToggle);
|
||||
});
|
||||
|
|
@ -1074,7 +1076,7 @@ describe("UsagePage", () => {
|
|||
});
|
||||
|
||||
// Switch to individual first
|
||||
const litellmToggle = screen.getByText("Litellm Model Name");
|
||||
const litellmToggle = screen.getAllByText("Litellm Model Name")[0];
|
||||
act(() => {
|
||||
fireEvent.click(litellmToggle);
|
||||
});
|
||||
|
|
@ -1084,7 +1086,7 @@ describe("UsagePage", () => {
|
|||
});
|
||||
|
||||
// Switch back to groups
|
||||
const publicToggle = screen.getByText("Public Model Name");
|
||||
const publicToggle = screen.getAllByText("Public Model Name")[0];
|
||||
act(() => {
|
||||
fireEvent.click(publicToggle);
|
||||
});
|
||||
|
|
@ -1093,6 +1095,34 @@ describe("UsagePage", () => {
|
|||
expect(screen.getByText("Top Public Model Names")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
it("should feed the Model Activity tab from the model_groups breakdown by default", async () => {
|
||||
renderWithProviders(<UsagePage {...defaultProps} />);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(mockUserDailyActivityAggregatedCall).toHaveBeenCalled();
|
||||
});
|
||||
|
||||
expect(screen.getByText("activity-source:model_groups")).toBeInTheDocument();
|
||||
expect(screen.queryByText("activity-source:models")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should switch the Model Activity tab to the litellm models breakdown on toggle click", async () => {
|
||||
renderWithProviders(<UsagePage {...defaultProps} />);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(mockUserDailyActivityAggregatedCall).toHaveBeenCalled();
|
||||
});
|
||||
|
||||
act(() => {
|
||||
fireEvent.click(screen.getAllByText("Litellm Model Name")[0]);
|
||||
});
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("activity-source:models")).toBeInTheDocument();
|
||||
});
|
||||
expect(screen.queryByText("activity-source:model_groups")).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
describe("customer usage banner", () => {
|
||||
|
|
|
|||
|
|
@ -55,6 +55,7 @@ import { DailyData, KeyMetricWithMetadata, MetricWithMetadata } from "@/componen
|
|||
import { valueFormatterSpend } from "@/components/UsagePage/utils/value_formatters";
|
||||
import EndpointUsage from "./EndpointUsage/EndpointUsage";
|
||||
import EntityUsage, { EntityList } from "./EntityUsage/EntityUsage";
|
||||
import ModelViewToggle, { ModelViewType } from "./ModelViewToggle";
|
||||
import SpendByProvider from "./EntityUsage/SpendByProvider";
|
||||
import TopKeyView from "@/components/UsagePage/components/EntityUsage/TopKeyView";
|
||||
import UsageAIChatPanel from "./UsageAIChatPanel";
|
||||
|
|
@ -143,7 +144,7 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
|
|||
// For admins: null means global view (all users), a string means filter by that user
|
||||
// For non-admins: always set to their own user ID
|
||||
const [selectedUserId, setSelectedUserId] = useState<string | null>(isAdmin ? null : userID || null);
|
||||
const [modelViewType, setModelViewType] = useState<"groups" | "individual">("groups");
|
||||
const [modelViewType, setModelViewType] = useState<ModelViewType>("groups");
|
||||
const [isCloudZeroModalOpen, setIsCloudZeroModalOpen] = useState(false);
|
||||
const [isGlobalExportModalOpen, setIsGlobalExportModalOpen] = useState(false);
|
||||
const [isAiChatOpen, setIsAiChatOpen] = useState(false);
|
||||
|
|
@ -438,7 +439,10 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
|
|||
() => [...userSpendData.results].sort((a, b) => new Date(a.date).getTime() - new Date(b.date).getTime()),
|
||||
[userSpendData.results],
|
||||
);
|
||||
const modelMetrics = useMemo(() => processActivityData(userSpendData, "models", teams), [userSpendData, teams]);
|
||||
const modelMetrics = useMemo(
|
||||
() => processActivityData(userSpendData, modelViewType === "groups" ? "model_groups" : "models", teams),
|
||||
[userSpendData, modelViewType, teams],
|
||||
);
|
||||
const keyMetrics = useMemo(() => processActivityData(userSpendData, "api_keys", teams), [userSpendData, teams]);
|
||||
const mcpServerMetrics = useMemo(
|
||||
() => processActivityData(userSpendData, "mcp_servers", teams),
|
||||
|
|
@ -753,28 +757,7 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
|
|||
value={topModelsLimit}
|
||||
onChange={(value) => setTopModelsLimit(value as number)}
|
||||
/>
|
||||
<div className="flex bg-gray-100 rounded-lg p-1">
|
||||
<button
|
||||
className={`px-3 py-1 text-sm rounded-md transition-colors ${
|
||||
modelViewType === "groups"
|
||||
? "bg-white shadow-xs text-gray-900"
|
||||
: "text-gray-600 hover:text-gray-900"
|
||||
}`}
|
||||
onClick={() => setModelViewType("groups")}
|
||||
>
|
||||
Public Model Name
|
||||
</button>
|
||||
<button
|
||||
className={`px-3 py-1 text-sm rounded-md transition-colors ${
|
||||
modelViewType === "individual"
|
||||
? "bg-white shadow-xs text-gray-900"
|
||||
: "text-gray-600 hover:text-gray-900"
|
||||
}`}
|
||||
onClick={() => setModelViewType("individual")}
|
||||
>
|
||||
Litellm Model Name
|
||||
</button>
|
||||
</div>
|
||||
<ModelViewToggle value={modelViewType} onChange={setModelViewType} />
|
||||
</div>
|
||||
{loading ? (
|
||||
<ChartLoader isDateChanging={isDateChanging} />
|
||||
|
|
@ -839,6 +822,9 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
|
|||
|
||||
{/* Activity Panel */}
|
||||
<TabPanel>
|
||||
<div className="flex justify-end mt-2 mb-4">
|
||||
<ModelViewToggle value={modelViewType} onChange={setModelViewType} />
|
||||
</div>
|
||||
<ActivityMetrics modelMetrics={modelMetrics} />
|
||||
</TabPanel>
|
||||
<TabPanel>
|
||||
|
|
|
|||
|
|
@ -72,6 +72,31 @@ vi.mock("@/components/team/TeamInfo", () => ({
|
|||
},
|
||||
}));
|
||||
|
||||
// The selected team is URL-derived (?team=) via useTeamDetailRouting. Next's real useSearchParams
|
||||
// re-renders subscribers on history.pushState/replaceState; mirror that so URL changes propagate.
|
||||
vi.mock("next/navigation", async () => {
|
||||
const { useSyncExternalStore } = await import("react");
|
||||
const LOCATION_CHANGE_EVENT = "test-locationchange";
|
||||
for (const method of ["pushState", "replaceState"] as const) {
|
||||
const original = window.history[method].bind(window.history);
|
||||
window.history[method] = (...args: Parameters<History["pushState"]>) => {
|
||||
original(...args);
|
||||
window.dispatchEvent(new Event(LOCATION_CHANGE_EVENT));
|
||||
};
|
||||
}
|
||||
const subscribe = (onChange: () => void) => {
|
||||
window.addEventListener(LOCATION_CHANGE_EVENT, onChange);
|
||||
window.addEventListener("popstate", onChange);
|
||||
return () => {
|
||||
window.removeEventListener(LOCATION_CHANGE_EVENT, onChange);
|
||||
window.removeEventListener("popstate", onChange);
|
||||
};
|
||||
};
|
||||
return {
|
||||
useSearchParams: () => new URLSearchParams(useSyncExternalStore(subscribe, () => window.location.search)),
|
||||
};
|
||||
});
|
||||
|
||||
vi.mock("./ModelSelect/ModelSelect", () => {
|
||||
const ModelSelect = React.forwardRef(({ value, onChange, dataTestId, id }: any, ref: any) => {
|
||||
return (
|
||||
|
|
@ -159,6 +184,7 @@ const renderWithQueryClient = (component: React.ReactElement) => {
|
|||
// Re-establish safe defaults before every test (clearAllMocks keeps return values, so restore them here).
|
||||
beforeEach(() => {
|
||||
mockTeamsTableProps = null;
|
||||
window.history.replaceState(null, "", "/teams/");
|
||||
});
|
||||
|
||||
describe("Teams - handleCreate organization handling", () => {
|
||||
|
|
@ -436,6 +462,47 @@ describe("Teams - premium props", () => {
|
|||
});
|
||||
});
|
||||
|
||||
describe("Teams - team detail deep link (?team=)", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
mockTeamInfoView.mockClear();
|
||||
vi.mocked(fetchAvailableModelsForTeamOrKey).mockResolvedValue([]);
|
||||
vi.mocked(fetchMCPAccessGroups).mockResolvedValue([]);
|
||||
vi.mocked(getGuardrailsList).mockResolvedValue({ guardrails: [] });
|
||||
mockUseOrganizations.mockReturnValue({ data: [] });
|
||||
});
|
||||
|
||||
it("selecting a team pushes ?team= to the URL", async () => {
|
||||
renderWithQueryClient(<Teams accessToken="test-token" userID="user-123" userRole="Admin" />);
|
||||
|
||||
await waitFor(() => expect(mockTeamsTableProps).not.toBeNull());
|
||||
act(() => mockTeamsTableProps.onSelectTeam({ ...baseTableTeam, team_id: "team-deep-link" }));
|
||||
|
||||
expect(window.location.search).toContain("team=team-deep-link");
|
||||
await waitFor(() => expect(mockTeamInfoView).toHaveBeenCalled());
|
||||
expect(mockTeamInfoView).toHaveBeenLastCalledWith(expect.objectContaining({ teamId: "team-deep-link" }));
|
||||
});
|
||||
|
||||
it("opens the team detail view directly from a ?team= deep link", async () => {
|
||||
window.history.replaceState(null, "", "/teams/?team=team-from-url");
|
||||
renderWithQueryClient(<Teams accessToken="test-token" userID="user-123" userRole="Admin" />);
|
||||
|
||||
await waitFor(() => expect(mockTeamInfoView).toHaveBeenCalled());
|
||||
expect(mockTeamInfoView).toHaveBeenLastCalledWith(expect.objectContaining({ teamId: "team-from-url" }));
|
||||
});
|
||||
|
||||
it("closing the team detail view removes ?team= from the URL", async () => {
|
||||
window.history.replaceState(null, "", "/teams/?team=team-from-url");
|
||||
renderWithQueryClient(<Teams accessToken="test-token" userID="user-123" userRole="Admin" />);
|
||||
|
||||
await waitFor(() => expect(mockTeamInfoView).toHaveBeenCalled());
|
||||
act(() => mockTeamInfoView.mock.calls.at(-1)?.[0].onClose());
|
||||
|
||||
expect(window.location.search).not.toContain("team=");
|
||||
await waitFor(() => expect(screen.queryByTestId("team-info-view")).not.toBeInTheDocument());
|
||||
});
|
||||
});
|
||||
|
||||
describe("Teams - Create Team CTA is grouped with the tabs on the left", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ import { useQueryClient } from "@tanstack/react-query";
|
|||
import { PageHeader } from "@/components/shared/PageHeader";
|
||||
import { Button as UIButton } from "@/components/ui/button";
|
||||
import { teamsTableKeys } from "@/app/(dashboard)/hooks/teams/useTeams";
|
||||
import { useTeamDetailRouting } from "@/app/(dashboard)/teams/detailNavigation";
|
||||
import { TeamsTable } from "./TeamsPage/TeamsTable";
|
||||
import AccessGroupSelector from "./common_components/AccessGroupSelector";
|
||||
import PassThroughRoutesSelector from "./common_components/PassThroughRoutesSelector";
|
||||
|
|
@ -135,7 +136,7 @@ const Teams: React.FC<TeamProps> = ({ accessToken, userID, userRole, premiumUser
|
|||
const [editModalVisible, setEditModalVisible] = useState(false);
|
||||
|
||||
const [selectedTeam, setSelectedTeam] = useState<Team | null>(null);
|
||||
const [selectedTeamId, setSelectedTeamId] = useState<string | null>(null);
|
||||
const { teamId: selectedTeamId, openTeam, close: closeTeamDetail } = useTeamDetailRouting();
|
||||
const [editTeam, setEditTeam] = useState<boolean>(false);
|
||||
|
||||
const [isTeamModalVisible, setIsTeamModalVisible] = useState(false);
|
||||
|
|
@ -482,12 +483,12 @@ const Teams: React.FC<TeamProps> = ({ accessToken, userID, userRole, premiumUser
|
|||
userID={userID}
|
||||
onSelectTeam={(team) => {
|
||||
setSelectedTeam(team);
|
||||
setSelectedTeamId(team.team_id);
|
||||
openTeam(team.team_id);
|
||||
setEditTeam(false);
|
||||
}}
|
||||
onEditTeam={(team) => {
|
||||
setSelectedTeam(team);
|
||||
setSelectedTeamId(team.team_id);
|
||||
openTeam(team.team_id);
|
||||
setEditTeam(true);
|
||||
}}
|
||||
onDeleteTeam={handleDelete}
|
||||
|
|
@ -547,11 +548,11 @@ const Teams: React.FC<TeamProps> = ({ accessToken, userID, userRole, premiumUser
|
|||
}}
|
||||
onClose={() => {
|
||||
setSelectedTeam(null);
|
||||
setSelectedTeamId(null);
|
||||
closeTeamDetail();
|
||||
setEditTeam(false);
|
||||
}}
|
||||
accessToken={accessToken}
|
||||
is_team_admin={is_team_admin(selectedTeam)}
|
||||
is_team_admin={is_team_admin(selectedTeam?.team_id === selectedTeamId ? selectedTeam : null)}
|
||||
is_proxy_admin={userRole == "Admin"}
|
||||
userModels={userModels}
|
||||
editTeam={editTeam}
|
||||
|
|
|
|||
|
|
@ -716,6 +716,60 @@ describe("processActivityData", () => {
|
|||
expect(result["gpt-4"].total_spend).toBe(100.5);
|
||||
});
|
||||
|
||||
it("should process model_groups data keyed by public model name including fallback entries", () => {
|
||||
const upstreamModelMetrics = {
|
||||
...EMPTY_SPEND_METRICS,
|
||||
spend: 10,
|
||||
api_requests: 10,
|
||||
successful_requests: 10,
|
||||
};
|
||||
const dailyActivityWithModelGroups: { results: DailyData[] } = {
|
||||
results: [
|
||||
{
|
||||
date: "2025-01-01",
|
||||
metrics: upstreamModelMetrics,
|
||||
breakdown: {
|
||||
...EMPTY_BREAKDOWN,
|
||||
models: {
|
||||
"gpt-5.2": {
|
||||
metrics: upstreamModelMetrics,
|
||||
metadata: {},
|
||||
api_key_breakdown: {},
|
||||
},
|
||||
},
|
||||
model_groups: {
|
||||
"gpt-5.2-eu": {
|
||||
metrics: { ...EMPTY_SPEND_METRICS, spend: 7, api_requests: 7, successful_requests: 7 },
|
||||
metadata: {},
|
||||
api_key_breakdown: {
|
||||
"key-1": {
|
||||
metrics: { ...EMPTY_SPEND_METRICS, spend: 7, api_requests: 7, total_tokens: 700 },
|
||||
metadata: { key_alias: "eu-key", team_id: "team1" },
|
||||
},
|
||||
},
|
||||
},
|
||||
"gpt-5.2": {
|
||||
metrics: { ...EMPTY_SPEND_METRICS, spend: 3, api_requests: 3, successful_requests: 3 },
|
||||
metadata: {},
|
||||
api_key_breakdown: {},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
],
|
||||
};
|
||||
|
||||
const result = processActivityData(dailyActivityWithModelGroups, "model_groups");
|
||||
|
||||
expect(Object.keys(result).sort()).toEqual(["gpt-5.2", "gpt-5.2-eu"]);
|
||||
expect(result["gpt-5.2-eu"].label).toBe("gpt-5.2-eu");
|
||||
expect(result["gpt-5.2-eu"].total_spend).toBe(7);
|
||||
expect(result["gpt-5.2-eu"].top_api_keys).toHaveLength(1);
|
||||
expect(result["gpt-5.2-eu"].top_api_keys[0].key_alias).toBe("eu-key");
|
||||
expect(result["gpt-5.2"].total_spend).toBe(3);
|
||||
expect(result["gpt-5.2"].total_requests).toBe(3);
|
||||
});
|
||||
|
||||
it("should process data for mcp_servers key", () => {
|
||||
const dailyActivityWithMCP: { results: DailyData[] } = {
|
||||
results: [
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue