mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
chore: merge internal staging into interview branch
Some checks are pending
LiteLLM Rust / rustfmt, clippy, test (push) Waiting to run
Some checks are pending
LiteLLM Rust / rustfmt, clippy, test (push) Waiting to run
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
commit
2008816278
172 changed files with 10261 additions and 2515 deletions
1
.gitignore
vendored
1
.gitignore
vendored
|
|
@ -141,3 +141,4 @@ crash.*.log
|
|||
.coverage
|
||||
|
||||
ui/litellm-dashboard/out/
|
||||
litellm.log
|
||||
|
|
|
|||
2
Makefile
2
Makefile
|
|
@ -176,6 +176,8 @@ lint-ruff-FULL-dev: install-dev
|
|||
if [ -n "$$files" ]; then echo "$$files" | xargs $(UV_RUN) ruff check; \
|
||||
else echo "No changed .py files to check."; fi
|
||||
|
||||
lint-basedpyright lint-basedpyright-budget-update: export NODE_OPTIONS := --max-old-space-size=12288
|
||||
|
||||
lint-basedpyright: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
|
||||
($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --base origin/litellm_internal_staging
|
||||
|
||||
|
|
|
|||
|
|
@ -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,12 +36,6 @@ pub(crate) const DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS: u64 = 300;
|
|||
#[cfg(feature = "server")]
|
||||
pub(crate) const MESSAGES_ROUTE_PATH: &str = "/v1/messages";
|
||||
|
||||
/// Provider name used by the Anthropic Messages route when a deployment's
|
||||
/// provider model does not carry an explicit provider prefix.
|
||||
pub(crate) const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic";
|
||||
pub(crate) const AZURE_ANTHROPIC_MESSAGES_PROVIDER: &str = "azure_ai";
|
||||
pub(crate) const BEDROCK_MESSAGES_PROVIDER: &str = "bedrock";
|
||||
|
||||
/// Request headers owned by the gateway and never forwarded upstream.
|
||||
#[cfg(feature = "server")]
|
||||
pub(crate) const MESSAGES_HEADERS_NOT_FORWARDED: &[&str] =
|
||||
|
|
|
|||
|
|
@ -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,160 +0,0 @@
|
|||
use std::collections::BTreeMap;
|
||||
use std::time::SystemTime;
|
||||
|
||||
use litellm_core::CoreResult;
|
||||
use litellm_core::error::CoreError;
|
||||
use litellm_core::providers::bedrock::aws_base::{
|
||||
AwsAuthConfig, resolve_credentials, sign_bedrock_post,
|
||||
};
|
||||
use serde_json::Value;
|
||||
|
||||
use super::client::http_client;
|
||||
use super::common_utils::truncate_error_body;
|
||||
use super::types::ProviderMessagesRequest;
|
||||
use crate::constants::{
|
||||
ANTHROPIC_MESSAGES_PROVIDER, AZURE_ANTHROPIC_MESSAGES_PROVIDER, BEDROCK_MESSAGES_PROVIDER,
|
||||
};
|
||||
|
||||
fn environment_lookup(key: &str) -> Option<String> {
|
||||
std::env::var(key).ok()
|
||||
}
|
||||
|
||||
async fn signed_request(
|
||||
request: &ProviderMessagesRequest,
|
||||
body: &[u8],
|
||||
) -> CoreResult<Vec<(String, String)>> {
|
||||
if request.provider != BEDROCK_MESSAGES_PROVIDER {
|
||||
return Ok(request.upstream_headers.clone());
|
||||
}
|
||||
if let Some(token) = &request.bearer_token {
|
||||
return Ok(request
|
||||
.upstream_headers
|
||||
.iter()
|
||||
.filter(|(name, _)| {
|
||||
!matches!(
|
||||
name.to_ascii_lowercase().as_str(),
|
||||
"authorization" | "x-api-key" | "anthropic-version"
|
||||
)
|
||||
})
|
||||
.cloned()
|
||||
.chain([
|
||||
("Authorization".to_string(), format!("Bearer {token}")),
|
||||
("content-type".to_string(), "application/json".to_string()),
|
||||
])
|
||||
.collect());
|
||||
}
|
||||
let headers = request
|
||||
.upstream_headers
|
||||
.iter()
|
||||
.filter(|(name, _)| {
|
||||
!matches!(
|
||||
name.to_ascii_lowercase().as_str(),
|
||||
"authorization" | "x-api-key" | "anthropic-version" | "host" | "content-length"
|
||||
)
|
||||
})
|
||||
.cloned()
|
||||
.chain(std::iter::once((
|
||||
"content-type".to_string(),
|
||||
"application/json".to_string(),
|
||||
)))
|
||||
.collect::<BTreeMap<_, _>>();
|
||||
let region = request.signing_region.as_deref().ok_or_else(|| {
|
||||
CoreError::InvalidRequest("Bedrock signing region was not resolved".to_string())
|
||||
})?;
|
||||
let credentials = resolve_credentials(AwsAuthConfig::default(), &environment_lookup).await?;
|
||||
let signed = sign_bedrock_post(
|
||||
&request.url,
|
||||
body,
|
||||
&headers,
|
||||
region,
|
||||
&credentials,
|
||||
SystemTime::now(),
|
||||
)?;
|
||||
Ok(signed.into_iter().collect())
|
||||
}
|
||||
|
||||
pub(super) async fn execute_messages_provider_call(
|
||||
request: ProviderMessagesRequest,
|
||||
) -> CoreResult<Value> {
|
||||
let body = serde_json::to_vec(&request.body).map_err(|error| {
|
||||
CoreError::InvalidRequest(format!("invalid messages request body: {error}"))
|
||||
})?;
|
||||
let headers = signed_request(&request, &body).await?;
|
||||
let mut request_builder = http_client().post(&request.url).body(body);
|
||||
for (key, value) in &headers {
|
||||
request_builder = request_builder.header(key, value);
|
||||
}
|
||||
if let Some(duration) = request.timeout {
|
||||
request_builder = request_builder.timeout(duration);
|
||||
}
|
||||
|
||||
let response = request_builder
|
||||
.send()
|
||||
.await
|
||||
.map_err(|err| CoreError::Network(err.to_string()))?;
|
||||
|
||||
let status = response.status();
|
||||
let text = response
|
||||
.text()
|
||||
.await
|
||||
.map_err(|err| CoreError::Network(err.to_string()))?;
|
||||
|
||||
if !status.is_success() {
|
||||
return Err(CoreError::Http {
|
||||
status: status.as_u16(),
|
||||
body: truncate_error_body(&text),
|
||||
});
|
||||
}
|
||||
|
||||
let response = serde_json::from_str(&text).map_err(|err| {
|
||||
CoreError::InvalidResponse(format!("invalid messages response JSON: {err}"))
|
||||
})?;
|
||||
let transformed = request
|
||||
.config
|
||||
.transform_response(&request.model, response)?;
|
||||
serde_json::to_value(transformed).map_err(|err| {
|
||||
CoreError::InvalidResponse(format!("failed to serialize messages response: {err}"))
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) async fn execute_messages_provider_stream(
|
||||
request: ProviderMessagesRequest,
|
||||
) -> CoreResult<reqwest::Response> {
|
||||
if !matches!(
|
||||
request.provider.as_str(),
|
||||
ANTHROPIC_MESSAGES_PROVIDER | AZURE_ANTHROPIC_MESSAGES_PROVIDER | BEDROCK_MESSAGES_PROVIDER
|
||||
) {
|
||||
return Err(CoreError::InvalidRequest(
|
||||
"streaming messages is not supported for this provider".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
let body = serde_json::to_vec(&request.body).map_err(|error| {
|
||||
CoreError::InvalidRequest(format!("invalid messages request body: {error}"))
|
||||
})?;
|
||||
let headers = signed_request(&request, &body).await?;
|
||||
let mut request_builder = http_client().post(&request.url).body(body);
|
||||
for (key, value) in &headers {
|
||||
request_builder = request_builder.header(key, value);
|
||||
}
|
||||
if let Some(duration) = request.timeout {
|
||||
request_builder = request_builder.timeout(duration);
|
||||
}
|
||||
|
||||
let response = request_builder
|
||||
.send()
|
||||
.await
|
||||
.map_err(|err| CoreError::Network(err.to_string()))?;
|
||||
let status = response.status();
|
||||
if !status.is_success() {
|
||||
let text = response
|
||||
.text()
|
||||
.await
|
||||
.map_err(|err| CoreError::Network(err.to_string()))?;
|
||||
return Err(CoreError::Http {
|
||||
status: status.as_u16(),
|
||||
body: truncate_error_body(&text),
|
||||
});
|
||||
}
|
||||
Ok(response)
|
||||
}
|
||||
|
|
@ -1,51 +0,0 @@
|
|||
use litellm_core::CoreResult;
|
||||
use serde_json::Value;
|
||||
|
||||
mod client;
|
||||
mod common_utils;
|
||||
mod handler;
|
||||
mod prepare;
|
||||
mod types;
|
||||
|
||||
pub use types::MessagesRequest;
|
||||
|
||||
use handler::{execute_messages_provider_call, execute_messages_provider_stream};
|
||||
use prepare::prepare_messages_call;
|
||||
|
||||
pub async fn messages(request: MessagesRequest<'_>) -> CoreResult<Value> {
|
||||
match execute_messages(request, false).await? {
|
||||
MessagesResponse::Json(body) => Ok(body),
|
||||
MessagesResponse::Stream { .. } => Err(litellm_core::CoreError::InvalidResponse(
|
||||
"non-streaming messages execution returned a stream".to_string(),
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) enum MessagesResponse {
|
||||
Json(Value),
|
||||
#[allow(dead_code)]
|
||||
Stream {
|
||||
provider: String,
|
||||
response: reqwest::Response,
|
||||
},
|
||||
}
|
||||
|
||||
pub(crate) async fn execute_messages(
|
||||
request: MessagesRequest<'_>,
|
||||
stream: bool,
|
||||
) -> CoreResult<MessagesResponse> {
|
||||
let prepared = prepare_messages_call(request)?;
|
||||
if stream {
|
||||
let provider = prepared.provider.clone();
|
||||
execute_messages_provider_stream(prepared)
|
||||
.await
|
||||
.map(|response| MessagesResponse::Stream { provider, response })
|
||||
} else {
|
||||
execute_messages_provider_call(prepared)
|
||||
.await
|
||||
.map(MessagesResponse::Json)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
|
@ -1,26 +0,0 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_core::messages::transformation::AnthropicMessagesProviderConfig;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
pub struct MessagesRequest<'a> {
|
||||
pub model: &'a str,
|
||||
pub body: Value,
|
||||
pub api_key: Option<&'a str>,
|
||||
pub api_base: Option<&'a str>,
|
||||
pub custom_llm_provider: Option<&'a str>,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
pub(crate) struct ProviderMessagesRequest {
|
||||
pub(crate) provider: String,
|
||||
pub(crate) model: String,
|
||||
pub(crate) config: &'static dyn AnthropicMessagesProviderConfig,
|
||||
pub(crate) url: String,
|
||||
pub(crate) body: Value,
|
||||
pub(crate) upstream_headers: Vec<(String, String)>,
|
||||
pub(crate) signing_region: Option<String>,
|
||||
pub(crate) bearer_token: Option<String>,
|
||||
pub(crate) timeout: Option<Duration>,
|
||||
}
|
||||
|
|
@ -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,18 +1,15 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use litellm_core::constants::ANTHROPIC_MESSAGES_PROVIDER;
|
||||
use litellm_core::messages::types::MessagesRequest;
|
||||
use litellm_core::messages::{messages, messages_stream};
|
||||
use litellm_core::router::Router;
|
||||
use litellm_core::{CoreError, CoreResult};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
|
||||
use crate::messages::{MessagesRequest, execute_messages};
|
||||
|
||||
pub(crate) enum MessagesResponse {
|
||||
Json(Value),
|
||||
Stream {
|
||||
provider: String,
|
||||
response: reqwest::Response,
|
||||
},
|
||||
Stream(reqwest::Response),
|
||||
}
|
||||
|
||||
pub async fn run(
|
||||
|
|
@ -55,13 +52,14 @@ pub async fn run(
|
|||
extra_headers,
|
||||
timeout: None,
|
||||
};
|
||||
let stream = request.body.get("stream").and_then(Value::as_bool) == Some(true);
|
||||
execute_messages(request, stream)
|
||||
.await
|
||||
.map(|response| match response {
|
||||
crate::messages::MessagesResponse::Json(body) => MessagesResponse::Json(body),
|
||||
crate::messages::MessagesResponse::Stream { provider, response } => {
|
||||
MessagesResponse::Stream { provider, response }
|
||||
}
|
||||
if request.body.get("stream").and_then(Value::as_bool) == Some(true) {
|
||||
return messages_stream(request).await.map(MessagesResponse::Stream);
|
||||
}
|
||||
|
||||
let response = messages(request).await?;
|
||||
serde_json::to_value(response)
|
||||
.map(MessagesResponse::Json)
|
||||
.map_err(|err| {
|
||||
CoreError::InvalidResponse(format!("failed to serialize messages response: {err}"))
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,12 +1,11 @@
|
|||
use litellm_core::CoreResult;
|
||||
use litellm_core::error::{CoreError, json_type_name};
|
||||
use litellm_core::messages::transformation::AnthropicMessagesProviderConfig;
|
||||
use litellm_core::providers::anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG;
|
||||
use litellm_core::providers::azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG;
|
||||
use litellm_core::providers::bedrock::messages::transformation::BEDROCK_MESSAGES_CONFIG;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::constants::MESSAGES_ERROR_BODY_MAX_CHARS;
|
||||
use crate::error::{CoreError, CoreResult, json_type_name};
|
||||
use crate::providers::anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG;
|
||||
use crate::providers::azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG;
|
||||
|
||||
use super::transformation::AnthropicMessagesProviderConfig;
|
||||
|
||||
pub(super) fn truncate_error_body(body: &str) -> String {
|
||||
if body.chars().count() <= MESSAGES_ERROR_BODY_MAX_CHARS {
|
||||
|
|
@ -22,7 +21,6 @@ pub(super) fn messages_provider_config(
|
|||
match provider {
|
||||
"anthropic" => Some(&ANTHROPIC_MESSAGES_CONFIG),
|
||||
"azure_ai" => Some(&AZURE_ANTHROPIC_MESSAGES_CONFIG),
|
||||
"bedrock" => Some(&BEDROCK_MESSAGES_CONFIG),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
76
litellm-rust/crates/core/src/messages/handler.rs
Normal file
76
litellm-rust/crates/core/src/messages/handler.rs
Normal file
|
|
@ -0,0 +1,76 @@
|
|||
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
|
||||
use crate::error::{CoreError, CoreResult};
|
||||
|
||||
use super::client::http_client;
|
||||
use super::common_utils::truncate_error_body;
|
||||
use super::types::{AnthropicMessagesResponse, ProviderMessagesRequest};
|
||||
|
||||
pub(super) async fn execute_messages_provider_call(
|
||||
request: ProviderMessagesRequest,
|
||||
) -> CoreResult<AnthropicMessagesResponse> {
|
||||
let mut request_builder = http_client().post(&request.url).json(&request.body);
|
||||
for (key, value) in &request.upstream_headers {
|
||||
request_builder = request_builder.header(key, value);
|
||||
}
|
||||
if let Some(duration) = request.timeout {
|
||||
request_builder = request_builder.timeout(duration);
|
||||
}
|
||||
|
||||
let response = request_builder
|
||||
.send()
|
||||
.await
|
||||
.map_err(|err| CoreError::Network(err.to_string()))?;
|
||||
|
||||
let status = response.status();
|
||||
let text = response
|
||||
.text()
|
||||
.await
|
||||
.map_err(|err| CoreError::Network(err.to_string()))?;
|
||||
|
||||
if !status.is_success() {
|
||||
return Err(CoreError::Http {
|
||||
status: status.as_u16(),
|
||||
body: truncate_error_body(&text),
|
||||
});
|
||||
}
|
||||
|
||||
let response = serde_json::from_str(&text).map_err(|err| {
|
||||
CoreError::InvalidResponse(format!("invalid messages response JSON: {err}"))
|
||||
})?;
|
||||
request.config.transform_response(&request.model, response)
|
||||
}
|
||||
|
||||
pub(super) async fn execute_messages_provider_stream(
|
||||
request: ProviderMessagesRequest,
|
||||
) -> CoreResult<reqwest::Response> {
|
||||
if request.provider != ANTHROPIC_MESSAGES_PROVIDER {
|
||||
return Err(CoreError::InvalidRequest(
|
||||
"streaming messages is not supported for this provider".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
let mut request_builder = http_client().post(&request.url).json(&request.body);
|
||||
for (key, value) in &request.upstream_headers {
|
||||
request_builder = request_builder.header(key, value);
|
||||
}
|
||||
if let Some(duration) = request.timeout {
|
||||
request_builder = request_builder.timeout(duration);
|
||||
}
|
||||
|
||||
let response = request_builder
|
||||
.send()
|
||||
.await
|
||||
.map_err(|err| CoreError::Network(err.to_string()))?;
|
||||
let status = response.status();
|
||||
if !status.is_success() {
|
||||
let text = response
|
||||
.text()
|
||||
.await
|
||||
.map_err(|err| CoreError::Network(err.to_string()))?;
|
||||
return Err(CoreError::Http {
|
||||
status: status.as_u16(),
|
||||
body: truncate_error_body(&text),
|
||||
});
|
||||
}
|
||||
Ok(response)
|
||||
}
|
||||
|
|
@ -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,12 +1,9 @@
|
|||
use litellm_core::CoreError;
|
||||
use litellm_core::CoreResult;
|
||||
use litellm_core::messages::transformation::MessagesAuthStrategy;
|
||||
use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
|
||||
use serde_json::Value;
|
||||
use crate::error::{CoreError, CoreResult};
|
||||
use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
|
||||
|
||||
use super::common_utils::{has_bearer_auth, has_header, messages_provider_config, string_headers};
|
||||
use super::transformation::MessagesAuthStrategy;
|
||||
use super::types::{MessagesRequest, ProviderMessagesRequest};
|
||||
use crate::constants::BEDROCK_MESSAGES_PROVIDER;
|
||||
|
||||
pub(super) fn prepare_messages_call(
|
||||
request: MessagesRequest<'_>,
|
||||
|
|
@ -34,45 +31,18 @@ pub(super) fn prepare_messages_call(
|
|||
|
||||
let mut headers = string_headers(request.extra_headers)?;
|
||||
|
||||
let is_bedrock = provider == BEDROCK_MESSAGES_PROVIDER;
|
||||
let auth_strategy = if is_bedrock {
|
||||
MessagesAuthStrategy::Header("authorization")
|
||||
} else {
|
||||
config.auth_strategy()
|
||||
};
|
||||
let bearer_token = if is_bedrock {
|
||||
request
|
||||
.api_key
|
||||
.map(str::to_string)
|
||||
.or_else(|| env_lookup("AWS_BEARER_TOKEN_BEDROCK"))
|
||||
.filter(|token| !token.trim().is_empty())
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let auth_header = match auth_strategy {
|
||||
MessagesAuthStrategy::Bearer
|
||||
if has_header(&headers, "authorization")
|
||||
|| (config.accepts_bearer_auth() && has_bearer_auth(&headers)) =>
|
||||
{
|
||||
None
|
||||
}
|
||||
MessagesAuthStrategy::Header(name)
|
||||
if has_header(&headers, name)
|
||||
|| (config.accepts_bearer_auth() && has_bearer_auth(&headers)) =>
|
||||
{
|
||||
None
|
||||
}
|
||||
MessagesAuthStrategy::Bearer => {
|
||||
let api_key = config.resolve_api_key(request.api_key, &env_lookup)?;
|
||||
Some(("authorization".to_string(), format!("Bearer {api_key}")))
|
||||
}
|
||||
MessagesAuthStrategy::Header(name) => {
|
||||
let api_key = config.resolve_api_key(request.api_key, &env_lookup)?;
|
||||
Some((name.to_string(), api_key))
|
||||
}
|
||||
};
|
||||
if let Some(header) = auth_header {
|
||||
headers.push(header);
|
||||
let auth_strategy = config.auth_strategy();
|
||||
let already_authorized = has_header(&headers, auth_strategy.header_name())
|
||||
|| (config.accepts_bearer_auth() && has_bearer_auth(&headers));
|
||||
if !already_authorized {
|
||||
let api_key = config.resolve_api_key(request.api_key, &env_lookup)?;
|
||||
let auth_header = match auth_strategy {
|
||||
MessagesAuthStrategy::Bearer => {
|
||||
("authorization".to_string(), format!("Bearer {api_key}"))
|
||||
}
|
||||
MessagesAuthStrategy::Header(name) => (name.to_string(), api_key),
|
||||
};
|
||||
headers.push(auth_header);
|
||||
}
|
||||
|
||||
for (name, value) in config.default_headers() {
|
||||
|
|
@ -81,9 +51,7 @@ pub(super) fn prepare_messages_call(
|
|||
}
|
||||
}
|
||||
|
||||
let _stream = request.body.get("stream").and_then(Value::as_bool) == Some(true);
|
||||
let url = config.complete_url(request.api_base, &model, &env_lookup)?;
|
||||
let signing_region = config.signing_region(request.api_base, &env_lookup);
|
||||
let typed_request = serde_json::from_value(request.body).map_err(|err| {
|
||||
CoreError::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
|
||||
})?;
|
||||
|
|
@ -101,8 +69,6 @@ pub(super) fn prepare_messages_call(
|
|||
url,
|
||||
body,
|
||||
upstream_headers: headers,
|
||||
signing_region,
|
||||
bearer_token,
|
||||
timeout: request.timeout,
|
||||
})
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -29,15 +29,13 @@ from opentelemetry.sdk.trace.export.in_memory_span_exporter import (
|
|||
InMemorySpanExporter,
|
||||
)
|
||||
from opentelemetry.trace import Span, SpanKind, Tracer
|
||||
from opentelemetry.util.re import parse_env_headers
|
||||
|
||||
from litellm._version import version as litellm_version
|
||||
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
|
||||
from litellm.integrations.otel.model.semconv import LiteLLM
|
||||
from litellm.integrations.otel.model.spans import LiteLLMSpanKind
|
||||
|
||||
# Re-exported so ``providers.parse_headers`` remains a stable entry point.
|
||||
from litellm.integrations.otel.model.utils import parse_headers as parse_headers
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.metrics import Meter
|
||||
from opentelemetry.sdk.metrics.export import MetricReader
|
||||
|
|
@ -119,6 +117,23 @@ def _otlp_traces_endpoint(endpoint: str | None) -> str | None:
|
|||
return endpoint + "/v1/traces"
|
||||
|
||||
|
||||
def parse_headers(raw: str | None) -> dict[str, str]:
|
||||
"""Parse an OTLP ``"k=v,k=v"`` header string into a dict.
|
||||
|
||||
``OTEL_EXPORTER_OTLP_HEADERS`` is W3C Baggage encoded per the OTLP spec, so
|
||||
values are percent-decoded: a vendor that documents
|
||||
``Authorization=Basic%20<token>`` (Grafana Cloud does, because a bare space
|
||||
is not representable there) has to reach the exporter as ``Basic <token>``,
|
||||
not with a literal ``%20`` that the backend rejects as malformed. The SDK's
|
||||
own parser is used so litellm decodes exactly what the OTLP exporters do
|
||||
when they read the env var themselves; ``liberal`` keeps values that are not
|
||||
percent-encoded (``Authorization=Bearer <token>``) working unchanged.
|
||||
"""
|
||||
if not raw:
|
||||
return {}
|
||||
return dict(parse_env_headers(raw, liberal=True))
|
||||
|
||||
|
||||
def _exporter_from_spec(spec: ExporterSpec) -> SpanExporter:
|
||||
kind = (spec.kind or "console").lower()
|
||||
factory = _EXPORTER_FACTORIES.get(kind)
|
||||
|
|
@ -191,6 +206,13 @@ def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader":
|
|||
``console`` (and any unrecognized kind) exports to the console; ``otlp_http``
|
||||
and ``otlp_grpc`` export over OTLP with the configured endpoint/headers. The
|
||||
reader exports on a 5s period, matching v1.
|
||||
|
||||
Histograms keep the SDK's default cumulative temporality. Prometheus-backed
|
||||
OTLP receivers (Grafana Cloud / Mimir, and the Prometheus OTLP endpoint)
|
||||
reject delta histograms outright with ``invalid temporality and type
|
||||
combination``, which drops the whole metric batch, while backends that
|
||||
prefer delta still accept cumulative. The enterprise billing exporter
|
||||
already relies on the same default.
|
||||
"""
|
||||
from opentelemetry.sdk.metrics.export import (
|
||||
ConsoleMetricExporter,
|
||||
|
|
@ -202,18 +224,12 @@ def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader":
|
|||
from opentelemetry.exporter.otlp.proto.http.metric_exporter import (
|
||||
OTLPMetricExporter as HTTPMetricExporter,
|
||||
)
|
||||
from opentelemetry.sdk.metrics import Histogram
|
||||
from opentelemetry.sdk.metrics.export import AggregationTemporality
|
||||
|
||||
exporter: Any = HTTPMetricExporter(
|
||||
endpoint=_otlp_metrics_endpoint(config.endpoint),
|
||||
headers=parse_headers(config.headers),
|
||||
preferred_temporality={Histogram: AggregationTemporality.DELTA},
|
||||
)
|
||||
elif kind in ("otlp_grpc", "grpc"):
|
||||
from opentelemetry.sdk.metrics import Histogram
|
||||
from opentelemetry.sdk.metrics.export import AggregationTemporality
|
||||
|
||||
try:
|
||||
from opentelemetry.exporter.otlp.proto.grpc.metric_exporter import (
|
||||
OTLPMetricExporter as GRPCMetricExporter,
|
||||
|
|
@ -227,7 +243,6 @@ def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader":
|
|||
exporter = GRPCMetricExporter(
|
||||
endpoint=config.endpoint,
|
||||
headers=parse_headers(config.headers),
|
||||
preferred_temporality={Histogram: AggregationTemporality.DELTA},
|
||||
)
|
||||
else:
|
||||
exporter = ConsoleMetricExporter()
|
||||
|
|
|
|||
|
|
@ -1267,7 +1267,7 @@ def _get_dummy_thought_signature() -> str:
|
|||
def convert_to_gemini_tool_call_invoke(
|
||||
message: ChatCompletionAssistantMessage,
|
||||
model: Optional[str] = None,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
forward_function_call_id: bool = False,
|
||||
) -> List[VertexPartType]:
|
||||
"""
|
||||
OpenAI tool invokes:
|
||||
|
|
@ -1317,16 +1317,12 @@ def convert_to_gemini_tool_call_invoke(
|
|||
VertexGeminiConfig,
|
||||
)
|
||||
|
||||
forward_tool_call_id = bool(
|
||||
model and VertexGeminiConfig._forward_gemini_function_call_id(model, custom_llm_provider)
|
||||
)
|
||||
|
||||
if tool_calls is not None:
|
||||
for idx, tool in enumerate(tool_calls):
|
||||
if "function" in tool:
|
||||
gemini_function_call: Optional[VertexFunctionCall] = _gemini_tool_call_invoke_helper(
|
||||
function_call_params=tool["function"],
|
||||
tool_call_id=(tool.get("id") if forward_tool_call_id else None),
|
||||
tool_call_id=(tool.get("id") if forward_function_call_id else None),
|
||||
)
|
||||
if gemini_function_call is not None:
|
||||
part_dict: VertexPartType = {"function_call": gemini_function_call}
|
||||
|
|
@ -1378,8 +1374,7 @@ def convert_to_gemini_tool_call_invoke(
|
|||
def convert_to_gemini_tool_call_result(
|
||||
message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage],
|
||||
last_message_with_tool_calls: Optional[dict],
|
||||
model: Optional[str] = None,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
forward_function_call_id: bool = False,
|
||||
) -> Union[VertexPartType, List[VertexPartType]]:
|
||||
"""
|
||||
OpenAI message with a tool result looks like:
|
||||
|
|
@ -1501,14 +1496,8 @@ def convert_to_gemini_tool_call_result(
|
|||
name = tool.get("function", {}).get("name", "")
|
||||
|
||||
# Echo the OpenAI tool_call_id on functionResponse (strip thought-signature suffix).
|
||||
# Only Google AI Studio Gemini 3+ accepts `id` on function_response parts.
|
||||
# Vertex AI and older Gemini models reject the field with HTTP 400.
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
VertexGeminiConfig,
|
||||
)
|
||||
|
||||
gemini_call_id: Optional[str] = None
|
||||
if model and VertexGeminiConfig._forward_gemini_function_call_id(model, custom_llm_provider):
|
||||
if forward_function_call_id:
|
||||
raw_tool_call_id = message.get("tool_call_id")
|
||||
if raw_tool_call_id and isinstance(raw_tool_call_id, str):
|
||||
stripped_id = raw_tool_call_id.split(THOUGHT_SIGNATURE_SEPARATOR, 1)[0]
|
||||
|
|
|
|||
|
|
@ -393,24 +393,25 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
if compaction_event is not None:
|
||||
return compaction_event
|
||||
|
||||
if self.sent_content_block_start is False:
|
||||
self.sent_content_block_start = True
|
||||
self.sent_content_block_finish = False
|
||||
self.chunk_queue.append(
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": self.current_content_block_index,
|
||||
"content_block": {"type": "text", "text": ""},
|
||||
}
|
||||
)
|
||||
return self.chunk_queue.popleft()
|
||||
|
||||
for chunk in self.completion_stream:
|
||||
if chunk == "None" or chunk is None:
|
||||
raise Exception
|
||||
|
||||
should_start_new_block = self._should_start_new_content_block(chunk)
|
||||
if should_start_new_block:
|
||||
is_opening_first_block = self.sent_content_block_start is False
|
||||
if is_opening_first_block and self._is_blank_delta(chunk):
|
||||
continue
|
||||
if is_opening_first_block:
|
||||
self.sent_content_block_start = True
|
||||
self.sent_content_block_finish = False
|
||||
self.chunk_queue.append(
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": self.current_content_block_index,
|
||||
"content_block": self.current_content_block_start,
|
||||
}
|
||||
)
|
||||
elif should_start_new_block:
|
||||
self._increment_content_block_index()
|
||||
|
||||
# applied_edits only needs to flow to the final message_delta
|
||||
|
|
@ -447,7 +448,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
# ``not self.queued_usage_chunk``.
|
||||
continue
|
||||
|
||||
if should_start_new_block and not self.sent_content_block_finish:
|
||||
if should_start_new_block and not is_opening_first_block and not self.sent_content_block_finish:
|
||||
# Queue the sequence: content_block_stop -> content_block_start
|
||||
# -> (optionally) the trigger chunk's delta.
|
||||
#
|
||||
|
|
@ -615,25 +616,25 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
if compaction_event is not None:
|
||||
return compaction_event
|
||||
|
||||
if self.sent_content_block_start is False:
|
||||
self.sent_content_block_start = True
|
||||
self.sent_content_block_finish = False
|
||||
self.chunk_queue.append(
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": self.current_content_block_index,
|
||||
"content_block": {"type": "text", "text": ""},
|
||||
}
|
||||
)
|
||||
return self.chunk_queue.popleft()
|
||||
|
||||
async for chunk in self.completion_stream:
|
||||
if chunk == "None" or chunk is None:
|
||||
raise Exception
|
||||
|
||||
# Check if we need to start a new content block
|
||||
should_start_new_block = self._should_start_new_content_block(chunk)
|
||||
if should_start_new_block:
|
||||
is_opening_first_block = self.sent_content_block_start is False
|
||||
if is_opening_first_block and self._is_blank_delta(chunk):
|
||||
continue
|
||||
if is_opening_first_block:
|
||||
self.sent_content_block_start = True
|
||||
self.sent_content_block_finish = False
|
||||
self.chunk_queue.append(
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": self.current_content_block_index,
|
||||
"content_block": self.current_content_block_start,
|
||||
}
|
||||
)
|
||||
elif should_start_new_block:
|
||||
self._increment_content_block_index()
|
||||
|
||||
# applied_edits only needs to flow to the final message_delta
|
||||
|
|
@ -664,7 +665,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
# Check if this processed chunk has a stop_reason - hold it for next chunk
|
||||
|
||||
if not self.queued_usage_chunk:
|
||||
if should_start_new_block and not self.sent_content_block_finish:
|
||||
if should_start_new_block and not is_opening_first_block and not self.sent_content_block_finish:
|
||||
# Queue the sequence: content_block_stop -> content_block_start
|
||||
# -> (optionally) the trigger chunk's delta.
|
||||
#
|
||||
|
|
@ -875,6 +876,22 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
return False
|
||||
return bool(delta.get(_delta_payload_field(delta_type)))
|
||||
|
||||
@staticmethod
|
||||
def _is_blank_delta(chunk: "ModelResponseStream") -> bool:
|
||||
choice = chunk.choices[0]
|
||||
if choice.finish_reason is not None:
|
||||
return False
|
||||
delta = choice.delta
|
||||
if getattr(delta, "tool_calls", None):
|
||||
return False
|
||||
if getattr(delta, "content", None):
|
||||
return False
|
||||
if getattr(delta, "reasoning_content", None):
|
||||
return False
|
||||
if getattr(delta, "thinking_blocks", None):
|
||||
return False
|
||||
return True
|
||||
|
||||
def _should_start_new_content_block(self, chunk: "ModelResponseStream") -> bool:
|
||||
"""
|
||||
Determine if we should start a new content block based on the processed chunk.
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ Why separate file? Make it easy to see how transformation works
|
|||
"""
|
||||
|
||||
import re
|
||||
from typing import List, Optional, Tuple, Literal
|
||||
from typing import List, Optional, Sequence, Tuple, Literal
|
||||
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.llms.vertex_ai import CachedContentRequestBody
|
||||
|
|
@ -152,6 +152,20 @@ def separate_cached_messages(
|
|||
return cached_messages, non_cached_messages
|
||||
|
||||
|
||||
def cached_messages_end_on_supported_turn(cached_messages: Sequence[AllMessageValues]) -> bool:
|
||||
"""
|
||||
The cachedContents API rejects contents ending on a model turn, which is how it
|
||||
classifies both assistant messages and tool results, with HTTP 400
|
||||
"Requests ending with a model turn are not supported". System messages are
|
||||
extracted into system_instruction before contents are built, so the terminal
|
||||
turn is the last non-system message.
|
||||
"""
|
||||
non_system_messages = tuple(message for message in cached_messages if message.get("role") != "system")
|
||||
if not non_system_messages:
|
||||
return bool(cached_messages)
|
||||
return non_system_messages[-1].get("role") not in ("assistant", "tool", "function")
|
||||
|
||||
|
||||
def transform_openai_messages_to_gemini_context_caching(
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
|
|
|
|||
|
|
@ -22,6 +22,7 @@ from litellm.types.llms.vertex_ai import (
|
|||
from ..common_utils import VertexAIError, get_vertex_base_url
|
||||
from ..vertex_llm_base import VertexBase
|
||||
from .transformation import (
|
||||
cached_messages_end_on_supported_turn,
|
||||
separate_cached_messages,
|
||||
transform_openai_messages_to_gemini_context_caching,
|
||||
)
|
||||
|
|
@ -308,6 +309,14 @@ class ContextCachingEndpoints(VertexBase):
|
|||
if len(cached_messages) == 0:
|
||||
return messages, optional_params, None
|
||||
|
||||
if not cached_messages_end_on_supported_turn(cached_messages):
|
||||
verbose_logger.debug(
|
||||
"Vertex AI context caching: cached message block ends on a model turn once "
|
||||
"system messages are extracted, which the cachedContents API rejects. "
|
||||
"Skipping context caching."
|
||||
)
|
||||
return messages, optional_params, None
|
||||
|
||||
# Gemini requires a minimum of 1024 tokens for context caching.
|
||||
# Skip caching if the cached content is too small to avoid API errors.
|
||||
if not is_prompt_caching_valid_prompt(
|
||||
|
|
@ -459,6 +468,14 @@ class ContextCachingEndpoints(VertexBase):
|
|||
if len(cached_messages) == 0:
|
||||
return messages, optional_params, None
|
||||
|
||||
if not cached_messages_end_on_supported_turn(cached_messages):
|
||||
verbose_logger.debug(
|
||||
"Vertex AI context caching: cached message block ends on a model turn once "
|
||||
"system messages are extracted, which the cachedContents API rejects. "
|
||||
"Skipping context caching."
|
||||
)
|
||||
return messages, optional_params, None
|
||||
|
||||
# Gemini requires a minimum of 1024 tokens for context caching.
|
||||
# Skip caching if the cached content is too small to avoid API errors.
|
||||
if not is_prompt_caching_valid_prompt(
|
||||
|
|
|
|||
|
|
@ -1,7 +1,9 @@
|
|||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from urllib.parse import unquote
|
||||
from typing import Any, Coroutine, Optional, Tuple, Union
|
||||
from typing import Any, Coroutine, Mapping, Optional, Tuple, Union
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -10,6 +12,7 @@ from litellm.integrations.gcs_bucket.gcs_bucket_base import (
|
|||
GCSBucketBase,
|
||||
GCSLoggingConfig,
|
||||
)
|
||||
from litellm.types.utils import StandardCallbackDynamicParams
|
||||
from litellm.litellm_core_utils.cloud_storage_security import (
|
||||
VERTEX_AI_MANAGED_GCS_PREFIX,
|
||||
should_allow_legacy_cloud_file_ids,
|
||||
|
|
@ -39,6 +42,35 @@ class VertexAIFilesHandler(GCSBucketBase):
|
|||
llm_provider=LlmProviders.VERTEX_AI,
|
||||
)
|
||||
|
||||
def _resolve_read_gcs_config(
|
||||
self,
|
||||
litellm_params: Mapping[str, object] | None,
|
||||
vertex_credentials: VERTEX_CREDENTIALS_TYPES | None,
|
||||
) -> tuple[str | None, str | None]:
|
||||
"""
|
||||
Resolve the GCS bucket and service-account credentials for the read/content path.
|
||||
|
||||
Sources them from the deployment's ``litellm_params`` (``gcs_bucket_name`` /
|
||||
``bucket_name`` and ``vertex_credentials``), mirroring the write path in
|
||||
``VertexAIFilesConfig._get_configured_bucket_name``, and falls back to the global
|
||||
``GCS_BUCKET_NAME`` / ``GCS_PATH_SERVICE_ACCOUNT`` env vars. This lets Vertex batch
|
||||
run entirely at the model-group level, so output written to a per-model bucket is
|
||||
readable without setting the global env vars.
|
||||
"""
|
||||
params: Mapping[str, object] = litellm_params or {}
|
||||
bucket_candidate = params.get("gcs_bucket_name") or params.get("bucket_name")
|
||||
configured_bucket_name = bucket_candidate if isinstance(bucket_candidate, str) else os.getenv("GCS_BUCKET_NAME")
|
||||
|
||||
credentials = params.get("vertex_credentials") or vertex_credentials
|
||||
if isinstance(credentials, dict):
|
||||
path_service_account: str | None = json.dumps(credentials)
|
||||
elif isinstance(credentials, str):
|
||||
path_service_account = credentials
|
||||
else:
|
||||
path_service_account = os.getenv("GCS_PATH_SERVICE_ACCOUNT")
|
||||
|
||||
return configured_bucket_name, path_service_account
|
||||
|
||||
def _extract_bucket_and_object_from_file_id(
|
||||
self,
|
||||
file_id: str,
|
||||
|
|
@ -91,7 +123,17 @@ class VertexAIFilesHandler(GCSBucketBase):
|
|||
if not file_id:
|
||||
raise ValueError("file_id is required in file_content_request")
|
||||
|
||||
gcs_logging_config: GCSLoggingConfig = await self.get_gcs_logging_config(kwargs={})
|
||||
configured_bucket_name, path_service_account = self._resolve_read_gcs_config(
|
||||
litellm_params=litellm_params,
|
||||
vertex_credentials=vertex_credentials,
|
||||
)
|
||||
dynamic_params = StandardCallbackDynamicParams(
|
||||
gcs_bucket_name=configured_bucket_name,
|
||||
gcs_path_service_account=path_service_account,
|
||||
)
|
||||
gcs_logging_config: GCSLoggingConfig = await self.get_gcs_logging_config(
|
||||
kwargs={"standard_callback_dynamic_params": dynamic_params}
|
||||
)
|
||||
bucket_name, object_path = self._extract_bucket_and_object_from_file_id(
|
||||
file_id=file_id,
|
||||
configured_bucket_name=gcs_logging_config["bucket_name"],
|
||||
|
|
|
|||
|
|
@ -661,6 +661,10 @@ def _gemini_convert_messages_with_history(
|
|||
vertex_project = litellm_params.get("vertex_project") or litellm_params.get("vertex_ai_project")
|
||||
vertex_credentials = litellm_params.get("vertex_credentials") or litellm_params.get("vertex_ai_credentials")
|
||||
|
||||
from .vertex_and_google_ai_studio_gemini import VertexGeminiConfig
|
||||
|
||||
forward_function_call_id = VertexGeminiConfig._forward_gemini_function_call_id(model or "")
|
||||
|
||||
try:
|
||||
while msg_i < len(messages):
|
||||
user_content: List[PartType] = []
|
||||
|
|
@ -910,7 +914,7 @@ def _gemini_convert_messages_with_history(
|
|||
gemini_tool_call_parts = convert_to_gemini_tool_call_invoke(
|
||||
assistant_msg,
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
forward_function_call_id=forward_function_call_id,
|
||||
)
|
||||
## check if gemini_tool_call already exists in assistant_content
|
||||
for gemini_tool_call_part in gemini_tool_call_parts:
|
||||
|
|
@ -973,8 +977,7 @@ def _gemini_convert_messages_with_history(
|
|||
_part = convert_to_gemini_tool_call_result(
|
||||
messages[msg_i], # type: ignore
|
||||
last_message_with_tool_calls, # type: ignore
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
forward_function_call_id=forward_function_call_id,
|
||||
)
|
||||
msg_i += 1
|
||||
# Handle both single part and list of parts (for Computer Use with images)
|
||||
|
|
|
|||
|
|
@ -289,15 +289,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
return False
|
||||
|
||||
@staticmethod
|
||||
def _forward_gemini_function_call_id(model: str, custom_llm_provider: Optional[str] = None) -> bool:
|
||||
def _forward_gemini_function_call_id(model: str) -> bool:
|
||||
"""
|
||||
Whether to include `id` on function_call / function_response parts.
|
||||
|
||||
Gemini 3+ on Google AI Studio accepts (and returns) `id` for strict
|
||||
tool-call matching. Vertex AI rejects the field with HTTP 400.
|
||||
Gemini 3+ accepts (and returns) `id` for strict tool-call matching, on Vertex AI and
|
||||
Google AI Studio alike. Older Gemini models reject the field with HTTP 400.
|
||||
"""
|
||||
if custom_llm_provider != "gemini":
|
||||
return False
|
||||
return VertexGeminiConfig._is_gemini_3_or_newer(model)
|
||||
|
||||
def _supports_penalty_parameters(self, model: str) -> bool:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -200,6 +200,13 @@ _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: tuple[MCPAuth, ...] = (
|
|||
)
|
||||
|
||||
|
||||
# OAuth discovery retry cooldown for servers whose endpoints stay unresolved. The base is one
|
||||
# reload cadence so a transient upstream failure recovers immediately; the cap bounds the request
|
||||
# amplification and log volume of a permanently broken configuration.
|
||||
_OAUTH_DISCOVERY_RETRY_BASE_SECONDS = 30.0
|
||||
_OAUTH_DISCOVERY_RETRY_MAX_SECONDS = 900.0
|
||||
|
||||
|
||||
def _blank_to_none(value: str | None) -> str | None:
|
||||
"""Collapse an absent, empty, or whitespace-only string to ``None``.
|
||||
|
||||
|
|
@ -247,6 +254,7 @@ def _endpoints_yield_to_issuer(
|
|||
authorization_url: str | None,
|
||||
token_url: str | None,
|
||||
registration_url: str | None,
|
||||
server_ref: str,
|
||||
) -> tuple[str | None, str | None, str | None]:
|
||||
"""The single rule that makes an admin-configured ``issuer`` the sole authoritative endpoint
|
||||
source (RFC 8414 §3.3): when it is set for a discovery auth type, the stored/manual
|
||||
|
|
@ -256,9 +264,29 @@ def _endpoints_yield_to_issuer(
|
|||
i.e. all ``None`` when issuer-anchored, else the inputs unchanged. Called at every resolution site
|
||||
so the invariant holds in one place instead of being re-derived per merge.
|
||||
"""
|
||||
if issuer is not None and is_discovery_auth_type:
|
||||
return None, None, None
|
||||
return authorization_url, token_url, registration_url
|
||||
if issuer is None or not is_discovery_auth_type:
|
||||
return authorization_url, token_url, registration_url
|
||||
discarded = sorted(
|
||||
label
|
||||
for label, value in (
|
||||
("authorization_url", authorization_url),
|
||||
("token_url", token_url),
|
||||
("registration_url", registration_url),
|
||||
)
|
||||
if value
|
||||
)
|
||||
if discarded:
|
||||
verbose_logger.warning(
|
||||
"MCP server %s has a pinned Issuer, so its stored %s %s not used: an anchored issuer is the "
|
||||
"sole endpoint source (RFC 8414 section 3.3) and a failed issuer fetch fails closed rather "
|
||||
"than falling back to them. To use manually configured endpoints instead, clear the Issuer "
|
||||
"field and re-enter the endpoint urls (clearing the Issuer also clears endpoints that may "
|
||||
"have been resolved under it), or clear the Issuer alone to re-discover from the server url.",
|
||||
server_ref,
|
||||
", ".join(discarded),
|
||||
"is" if len(discarded) == 1 else "are",
|
||||
)
|
||||
return None, None, None
|
||||
|
||||
|
||||
def _normalized_authorize_endpoint(url: str) -> str:
|
||||
|
|
@ -280,6 +308,68 @@ def _issuer_matches(claimed_issuer: object, configured_issuer: str) -> bool:
|
|||
return _normalized_authorize_endpoint(claimed_issuer) == _normalized_authorize_endpoint(configured_issuer)
|
||||
|
||||
|
||||
def _flow_endpoints_missing(
|
||||
auth_type: MCPAuthType | None,
|
||||
oauth2_flow: str | None,
|
||||
authorization_url: str | None,
|
||||
token_url: str | None,
|
||||
token_exchange_endpoint: str | None = None,
|
||||
) -> bool:
|
||||
"""Whether a built server is missing an endpoint its flow needs to run at all.
|
||||
|
||||
Used by the reload fast-path exemption: discovery runs at build time only, and the fast path
|
||||
reuses an unchanged row's registry entry verbatim, so a server whose discovery came back empty
|
||||
(transient upstream failure, rate limiting) would stay broken until some unrelated config write
|
||||
bumps ``updated_at``, serving its 400 the whole time. Rebuilding just these entries retries
|
||||
discovery on the normal reload cadence. It costs no extra fetch for servers that resolved, and
|
||||
none for those with no discovery source, since the build skips discovery for both.
|
||||
"""
|
||||
if auth_type == MCPAuth.oauth2_token_exchange:
|
||||
# A configured exchange endpoint replaces discovery entirely; only a server that must
|
||||
# discover its token endpoint and still has none is unresolved.
|
||||
return token_exchange_endpoint is None and token_url is None
|
||||
if auth_type not in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES:
|
||||
return False
|
||||
if oauth2_flow == "client_credentials":
|
||||
return token_url is None
|
||||
return authorization_url is None or token_url is None
|
||||
|
||||
|
||||
def _oauth_endpoints_unresolved(server: MCPServer) -> bool:
|
||||
"""``_flow_endpoints_missing`` over a built registry entry, for the reload fast-path check.
|
||||
|
||||
The flow comes from ``effective_oauth2_flow``, the one column-first, shape-fallback judge every
|
||||
flow decision uses, not from the raw column: a legacy row the startup backfill deliberately left
|
||||
unstamped (the ambiguous M2M shape) serves M2M at request time, and reading the bare column here
|
||||
would classify it as interactive-missing-endpoints and re-run discovery on every reload.
|
||||
"""
|
||||
if (
|
||||
server.auth_type == MCPAuth.oauth2_token_exchange
|
||||
and server.token_exchange_profile == "entra_obo"
|
||||
and not server.scopes
|
||||
):
|
||||
# entra_obo fails closed at exchange time without a scope (token_exchanger.py), and scopes
|
||||
# can come from resource discovery, so a server that resolved its endpoints but no scopes is
|
||||
# still unresolved for its flow.
|
||||
return True
|
||||
if server.is_dcr_bridge and not server.client_id and server.registration_url is None:
|
||||
# A DCR bridge with no admin-configured client can only register callers through the
|
||||
# upstream's registration endpoint, so a build that resolved the authorize and token
|
||||
# endpoints but not registration_endpoint (partial metadata) is still unresolved for its
|
||||
# flow and must keep retrying; without this it silently degrades to the short-circuit arm
|
||||
# until an unrelated config write. Scopes are deliberately NOT part of completeness: they
|
||||
# are a request hint the authorization server bounds at consent (RFC 6749 section 3.3),
|
||||
# and a server without them is fully functional.
|
||||
return True
|
||||
return _flow_endpoints_missing(
|
||||
server.auth_type,
|
||||
MCPServerManager.effective_oauth2_flow(server),
|
||||
server.authorization_url,
|
||||
server.token_url,
|
||||
server.token_exchange_endpoint,
|
||||
)
|
||||
|
||||
|
||||
def _endpoints_corroborate_authorization_url(
|
||||
source_authorization_url: str | None,
|
||||
trusted_authorization_url: str | None,
|
||||
|
|
@ -311,11 +401,10 @@ def _carry_forward_resolved_oauth_endpoints(new_server: MCPServer, previous_serv
|
|||
during re-discovery downgrades a working server (``authorization_url`` set) to a broken one
|
||||
(``None``, /authorize 400s) with no configuration change. Mirrors the ``short_prefix``
|
||||
carry-forward. Skipped when the server's ``url`` or ``auth_type`` changed, since the previous
|
||||
endpoints may then belong to a different upstream. ``registration_url`` IS carried even though
|
||||
``_persist_discovered_oauth_endpoints`` refuses to write it to the row: carrying only restores
|
||||
the same in-memory value the previous build already ran with, while persisting it would flip
|
||||
``_dcr_bridge_relays_client_registration`` (which keys off the stored column) for dcr_bridge
|
||||
servers that never had one configured.
|
||||
endpoints may then belong to a different upstream. Discovery results live only on the in-memory
|
||||
registry entry; the gateway never writes them to the row, whose OAuth columns carry admin intent
|
||||
alone, so this carry is the sole last-known-good mechanism and restores exactly the values the
|
||||
previous build already ran with.
|
||||
|
||||
Carry-forward is a non-manual endpoint source, so the same trust rule as discovery applies: the
|
||||
previous ``token_url``/``registration_url``/``scopes`` are carried only when the previous
|
||||
|
|
@ -1182,6 +1271,40 @@ class MCPServerManager:
|
|||
# empty result, or failure). Used to throttle re-probes for servers that do
|
||||
# not return instructions, and to apply a short cooldown after failures.
|
||||
self._upstream_initialize_instructions_probed_at: dict[str, float] = {}
|
||||
# Per-server (consecutive failures, monotonic timestamp) for OAuth discovery retries, so a
|
||||
# server whose endpoints never resolve backs off instead of re-running the full
|
||||
# RFC 9728 -> 8414 chain, and re-logging its warning, on every reload forever.
|
||||
self._oauth_discovery_retry_state: dict[
|
||||
str, tuple[int, float]
|
||||
] = {} # mutable-ok: retry cooldown cache, keyed per server and pruned on success
|
||||
|
||||
def _oauth_discovery_retry_due(self, server_id: str) -> bool:
|
||||
"""Whether an unresolved server is due for another discovery attempt.
|
||||
|
||||
The reload fast-path exemption is what retries a failed discovery, so without a cooldown a
|
||||
permanently unresolvable server re-runs the whole RFC 9728 -> RFC 8414 -> origin-fallback
|
||||
chain and re-emits its unresolved-endpoints warning on every reload, per server, forever.
|
||||
Delay doubles per consecutive failure from ``_OAUTH_DISCOVERY_RETRY_BASE_SECONDS`` up to
|
||||
``_OAUTH_DISCOVERY_RETRY_MAX_SECONDS``, so a transient outage still recovers on the next
|
||||
reload while a broken configuration settles to one attempt per cap.
|
||||
"""
|
||||
state = self._oauth_discovery_retry_state.get(server_id)
|
||||
if state is None:
|
||||
return True
|
||||
failures, attempted_at = state
|
||||
delay = min(
|
||||
_OAUTH_DISCOVERY_RETRY_BASE_SECONDS * (2 ** max(failures - 1, 0)),
|
||||
_OAUTH_DISCOVERY_RETRY_MAX_SECONDS,
|
||||
)
|
||||
return (time.monotonic() - attempted_at) >= delay
|
||||
|
||||
def _record_oauth_discovery_outcome(self, server: MCPServer) -> None:
|
||||
"""Advance or clear a server's retry cooldown after a rebuild resolved it or did not."""
|
||||
if not _oauth_endpoints_unresolved(server):
|
||||
self._oauth_discovery_retry_state.pop(server.server_id, None)
|
||||
return
|
||||
failures, _ = self._oauth_discovery_retry_state.get(server.server_id, (0, 0.0))
|
||||
self._oauth_discovery_retry_state[server.server_id] = (failures + 1, time.monotonic())
|
||||
|
||||
def _remember_upstream_initialize_instructions(self, server: MCPServer, client: MCPClient) -> None:
|
||||
raw = getattr(client, "_last_initialize_instructions", None)
|
||||
|
|
@ -1357,6 +1480,7 @@ class MCPServerManager:
|
|||
manual_authorization_url,
|
||||
manual_token_url,
|
||||
manual_registration_url,
|
||||
server_name or server_id,
|
||||
)
|
||||
should_discover = _has_oauth_discovery_source(server_url, use_issuer_anchor) and (
|
||||
is_discovery_auth_type or obo_needs_discovery
|
||||
|
|
@ -1834,7 +1958,6 @@ class MCPServerManager:
|
|||
*,
|
||||
credentials_are_encrypted: bool = True,
|
||||
env_vars_are_encrypted: Optional[bool] = None,
|
||||
persist_discovered_endpoints: bool = True,
|
||||
) -> MCPServer:
|
||||
_mcp_info: MCPInfo = mcp_server.mcp_info or {}
|
||||
env_dict = _deserialize_json_dict(getattr(mcp_server, "env", None))
|
||||
|
|
@ -1925,7 +2048,12 @@ class MCPServerManager:
|
|||
or self._obo_needs_endpoint_discovery(auth_type, token_exchange_endpoint, manual_token_url),
|
||||
)
|
||||
manual_authorization_url, manual_token_url, manual_registration_url = _endpoints_yield_to_issuer(
|
||||
manual_issuer, is_discovery_auth_type, manual_authorization_url, manual_token_url, manual_registration_url
|
||||
manual_issuer,
|
||||
is_discovery_auth_type,
|
||||
manual_authorization_url,
|
||||
manual_token_url,
|
||||
manual_registration_url,
|
||||
mcp_server.alias or mcp_server.server_name or mcp_server.server_id,
|
||||
)
|
||||
gated_oauth_metadata = await self._resolve_table_oauth_metadata(
|
||||
mcp_server=mcp_server,
|
||||
|
|
@ -2033,143 +2161,8 @@ class MCPServerManager:
|
|||
max_concurrent_requests=getattr(mcp_server, "max_concurrent_requests", None),
|
||||
)
|
||||
_warn_internal_delegate_pkce_if_applicable(new_server, source="database")
|
||||
if persist_discovered_endpoints:
|
||||
await self._persist_discovered_obo_token_url(
|
||||
server_id=mcp_server.server_id,
|
||||
auth_type=auth_type,
|
||||
existing_token_url=manual_token_url,
|
||||
discovered_token_url=new_server.token_url,
|
||||
)
|
||||
await self._persist_discovered_oauth_endpoints(
|
||||
server_id=mcp_server.server_id,
|
||||
auth_type=auth_type,
|
||||
existing_issuer=manual_issuer,
|
||||
existing_authorization_url=manual_authorization_url,
|
||||
existing_token_url=manual_token_url,
|
||||
existing_scopes=scopes,
|
||||
metadata=gated_oauth_metadata,
|
||||
is_issuer_anchored=use_issuer_anchor,
|
||||
)
|
||||
return new_server
|
||||
|
||||
async def _persist_discovered_obo_token_url(
|
||||
self,
|
||||
*,
|
||||
server_id: str,
|
||||
auth_type: Optional[MCPAuthType],
|
||||
existing_token_url: Optional[str],
|
||||
discovered_token_url: Optional[str],
|
||||
) -> None:
|
||||
"""Write a freshly discovered OBO token endpoint back onto the DB row.
|
||||
|
||||
``build_mcp_server_from_table`` resolves ``token_url`` via RFC 9728 -> RFC 8414 for an
|
||||
``oauth2_token_exchange`` server that has none configured, but that resolved value otherwise
|
||||
lives only on the returned in-memory object; the row keeps ``token_url=None`` so every rebuild
|
||||
re-runs discovery, and a transient upstream outage during a rebuild leaves the server with no
|
||||
endpoint until discovery next succeeds. Persisting it makes ``_obo_needs_endpoint_discovery``
|
||||
return False on the next build. Fires at most once per server (skipped once the row has a
|
||||
value), and is best-effort: a write failure just means discovery runs again next time.
|
||||
"""
|
||||
if auth_type != MCPAuth.oauth2_token_exchange:
|
||||
return
|
||||
if existing_token_url or not discovered_token_url:
|
||||
return
|
||||
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415
|
||||
|
||||
if prisma_client is None:
|
||||
return
|
||||
try:
|
||||
await MCPServerRepository(prisma_client).table.update(
|
||||
where={"server_id": server_id},
|
||||
data={"token_url": discovered_token_url},
|
||||
)
|
||||
verbose_logger.debug("Persisted discovered OBO token_url for MCP server %s", server_id)
|
||||
except Exception as exc: # noqa: BLE001 - best-effort; a failed write re-discovers next build
|
||||
verbose_logger.warning("Failed to persist discovered OBO token_url for MCP server %s: %s", server_id, exc)
|
||||
|
||||
async def _persist_discovered_oauth_endpoints(
|
||||
self,
|
||||
*,
|
||||
server_id: str,
|
||||
auth_type: MCPAuthType | None,
|
||||
existing_issuer: str | None,
|
||||
existing_authorization_url: str | None,
|
||||
existing_token_url: str | None,
|
||||
existing_scopes: list[str] | None,
|
||||
metadata: MCPOAuthMetadata | None,
|
||||
is_issuer_anchored: bool = False,
|
||||
) -> None:
|
||||
"""Write freshly discovered OAuth endpoints back onto the DB row.
|
||||
|
||||
Same rationale as ``_persist_discovered_obo_token_url`` but for the interactive oauth2
|
||||
family: discovered ``authorization_url``/``token_url``/``scopes`` otherwise live only on
|
||||
the in-memory registry entry, which is rebuilt on every client connect (the DCR reuse path
|
||||
calls ``update_server``) and on every post-write DB reload, so one failed re-discovery
|
||||
serves the 400 "authorization url is not configured" from /authorize until a later rebuild succeeds.
|
||||
Only fills row fields that are currently empty, never persists origin-fallback guesses
|
||||
(RFC 9728/8414-advertised metadata only), and deliberately skips ``registration_url``
|
||||
because ``_dcr_bridge_relays_client_registration`` keys off that column. Best-effort: a
|
||||
failed write re-discovers on the next build. Scopes go through ``update_mcp_server`` so
|
||||
they merge into the credentials blob without touching the stored client credentials.
|
||||
|
||||
For an issuer-anchored server (``is_issuer_anchored``) the endpoints are re-derived from the
|
||||
§3.3-validated issuer document on every build, so they are NOT persisted into the endpoint
|
||||
columns: persisting them would make the next build see populated endpoints and treat them as
|
||||
authoritative stored values, defeating the "endpoints come solely from the issuer" invariant.
|
||||
Only the resource-driven scopes are persisted for such servers.
|
||||
"""
|
||||
if auth_type not in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES:
|
||||
return
|
||||
if metadata is None or metadata.from_origin_fallback:
|
||||
return
|
||||
issuer_update = (
|
||||
{"issuer": metadata.discovered_issuer} if metadata.discovered_issuer and not existing_issuer else {}
|
||||
)
|
||||
authorization_url_update = (
|
||||
{"authorization_url": metadata.authorization_url}
|
||||
if metadata.authorization_url and not existing_authorization_url and not is_issuer_anchored
|
||||
else {}
|
||||
)
|
||||
token_url_update = (
|
||||
{"token_url": metadata.token_url}
|
||||
if metadata.token_url and not existing_token_url and not is_issuer_anchored
|
||||
else {}
|
||||
)
|
||||
scopes_update = {"credentials": {"scopes": metadata.scopes}} if metadata.scopes and not existing_scopes else {}
|
||||
updates: dict[str, object] = {
|
||||
**issuer_update,
|
||||
**authorization_url_update,
|
||||
**token_url_update,
|
||||
**scopes_update,
|
||||
}
|
||||
if not updates:
|
||||
return
|
||||
from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # db.py imports this module at load
|
||||
update_mcp_server,
|
||||
)
|
||||
from litellm.proxy._types import UpdateMCPServerRequest # noqa: PLC0415 # heavy module; import at call time
|
||||
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime value, set after startup
|
||||
|
||||
if prisma_client is None:
|
||||
return
|
||||
try:
|
||||
await update_mcp_server(
|
||||
prisma_client=prisma_client,
|
||||
data=UpdateMCPServerRequest.model_validate({"server_id": server_id, **updates}),
|
||||
touched_by="mcp_oauth_discovery",
|
||||
)
|
||||
verbose_logger.info(
|
||||
"Persisted discovered OAuth endpoints for MCP server %s: %s",
|
||||
server_id,
|
||||
sorted(updates),
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 - best-effort; a failed write re-discovers next build
|
||||
verbose_logger.warning(
|
||||
"Failed to persist discovered OAuth endpoints for MCP server %s: %s",
|
||||
server_id,
|
||||
exc,
|
||||
)
|
||||
|
||||
async def _maybe_register_openapi_tools(self, server: MCPServer, *, initialize_mapping: bool = True):
|
||||
"""Register OpenAPI tools if the server has a spec_path configured."""
|
||||
if server.spec_path:
|
||||
|
|
@ -5347,6 +5340,10 @@ class MCPServerManager:
|
|||
and existing_server.updated_at is not None
|
||||
and server.updated_at is not None
|
||||
and existing_server.updated_at == server.updated_at
|
||||
and not (
|
||||
_oauth_endpoints_unresolved(existing_server)
|
||||
and self._oauth_discovery_retry_due(server.server_id)
|
||||
)
|
||||
):
|
||||
# Re-use existing server instance to avoid re-running build_mcp_server_from_table()
|
||||
# which can perform network discovery for OAuth2 servers.
|
||||
|
|
@ -5364,6 +5361,7 @@ class MCPServerManager:
|
|||
# already-decrypted records add_server/update_server are handed.
|
||||
# Decrypt them while building the registry entry.
|
||||
new_server = await self.build_mcp_server_from_table(server, env_vars_are_encrypted=True)
|
||||
self._record_oauth_discovery_outcome(new_server)
|
||||
# Carry the cached short_prefix from the previous registry entry
|
||||
# (if any) so the prefix is stable across reloads.
|
||||
if existing_server is not None and existing_server.short_prefix:
|
||||
|
|
|
|||
|
|
@ -0,0 +1,148 @@
|
|||
"""One-time heal for MCP server rows whose ``issuer`` a released version wrote by itself.
|
||||
|
||||
Until the write was removed, OAuth discovery stamped the issuer it discovered onto the ``issuer``
|
||||
column trust-on-first-use. That column means "the admin pinned this trust anchor", so the next
|
||||
registry build read the gateway's own output back as admin intent: the server turned issuer-anchored
|
||||
(RFC 8414 section 3.3), its stored authorization/token/registration URLs stopped applying, and a
|
||||
failed issuer-document fetch left it with no authorize endpoint (GH #34985).
|
||||
|
||||
Deleting the write fixes every row created afterwards but cannot fix a row already stamped, which
|
||||
still reads as pinned. This heals those rows by clearing the stamp so their configured endpoints
|
||||
apply again.
|
||||
|
||||
The signal is a heuristic, and deliberately a narrow one. ``updated_by`` records only the most recent
|
||||
writer, and no audit trail says which field that writer touched, so "discovery wrote this issuer" is
|
||||
not directly knowable. Two independent clauses bound it, and each rules out a different way of
|
||||
destroying a pin an admin meant.
|
||||
|
||||
Configured endpoints must be present. A deliberately pinned row very often has none, both because the
|
||||
Issuer field is documented as overriding them and because ``update_mcp_server`` clears them when an
|
||||
issuer changes, so "issuer set, endpoints empty" is the canonical shape of a real pin and must never
|
||||
be cleared on this evidence. Skipping those rows costs little: with nothing configured to restore, the
|
||||
anchored and resource-rooted paths resolve from the same upstream document, and the row still gets the
|
||||
unresolved-endpoint retry and the anchored-discard warning.
|
||||
|
||||
The configured endpoints must also share the issuer's origin. A stamped issuer is by construction the
|
||||
one self-attested by the authorization-server document discovery reached from this very server, so
|
||||
endpoints typed alongside it address that same authority. An admin who pinned an issuer and typed
|
||||
endpoints for a different authority is expressing an intent that clearing the issuer would discard, so
|
||||
that row is warned about and never healed.
|
||||
|
||||
What survives both clauses is a row whose configured endpoints and stamped issuer share an origin,
|
||||
which is exactly the GH #34985 shape. An admin who pinned that same origin by hand lands here too, and
|
||||
for them the clear is close to a no-op: their typed endpoints keep serving and still anchor the
|
||||
RFC 9700 corroboration gate, with only the stricter section 3.3 anchoring lost. Every heal logs the
|
||||
cleared value so it can be restored, and the clear is recorded under this module's actor so the heal
|
||||
runs at most once per row.
|
||||
"""
|
||||
|
||||
from typing import Protocol
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import canonicalize_url_identity
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
# The actor the removed discovery write-back stamped rows with.
|
||||
_DISCOVERY_ACTOR = "mcp_oauth_discovery"
|
||||
|
||||
# The actor recorded on a healed row, which also makes the heal idempotent: once a row is cleared it
|
||||
# no longer matches ``updated_by == _DISCOVERY_ACTOR`` and is never reconsidered.
|
||||
_BACKFILL_ACTOR = "mcp_oauth_issuer_stamp_backfill"
|
||||
|
||||
_AUTH_TYPES_WITH_ISSUER_ANCHORING = ("oauth2", "true_passthrough", "oauth_delegate")
|
||||
|
||||
|
||||
def _origin(url: str) -> str | None:
|
||||
"""The scheme-and-authority identity of ``url``, or ``None`` when it has none.
|
||||
|
||||
Built on the shared URL canonicalizer so the lowercase-host and default-port rules match the
|
||||
RFC 8414 issuer comparison the resolution path uses, instead of being re-derived here.
|
||||
"""
|
||||
parsed = urlparse(canonicalize_url_identity(url))
|
||||
if not parsed.scheme or not parsed.netloc:
|
||||
return None
|
||||
return f"{parsed.scheme}://{parsed.netloc}"
|
||||
|
||||
|
||||
class _MCPServerRow(Protocol):
|
||||
"""The MCP server row fields this heal reads, so the untyped DB record is narrowed once here."""
|
||||
|
||||
server_id: str
|
||||
alias: str | None
|
||||
server_name: str | None
|
||||
auth_type: str | None
|
||||
issuer: str | None
|
||||
authorization_url: str | None
|
||||
token_url: str | None
|
||||
registration_url: str | None
|
||||
updated_by: str | None
|
||||
|
||||
|
||||
def _is_stamped_issuer_row(row: _MCPServerRow) -> bool:
|
||||
"""Whether this row carries the full signature of a gateway-written issuer stamp.
|
||||
|
||||
The whole rule lives here, including the writer check the query also filters on, so the decision
|
||||
to clear an admin-visible field is auditable in one place rather than split between a predicate
|
||||
and a query.
|
||||
"""
|
||||
if getattr(row, "updated_by", None) != _DISCOVERY_ACTOR:
|
||||
return False
|
||||
if not (getattr(row, "issuer", None) or "").strip():
|
||||
return False
|
||||
if getattr(row, "auth_type", None) not in _AUTH_TYPES_WITH_ISSUER_ANCHORING:
|
||||
return False
|
||||
configured = tuple(
|
||||
value.strip()
|
||||
for value in (row.authorization_url, row.token_url, row.registration_url)
|
||||
if value and value.strip()
|
||||
)
|
||||
if not configured:
|
||||
return False
|
||||
issuer_origin = _origin(row.issuer or "")
|
||||
return issuer_origin is not None and all(_origin(endpoint) == issuer_origin for endpoint in configured)
|
||||
|
||||
|
||||
async def backfill_discovery_stamped_issuers(prisma_client: PrismaClient) -> int:
|
||||
"""Clear gateway-written issuer stamps, returning the number of rows healed."""
|
||||
candidate_rows: list[_MCPServerRow] = await prisma_client.db.litellm_mcpservertable.find_many(
|
||||
where={
|
||||
"updated_by": _DISCOVERY_ACTOR,
|
||||
"auth_type": {"in": list(_AUTH_TYPES_WITH_ISSUER_ANCHORING)},
|
||||
},
|
||||
)
|
||||
stamped = tuple(row for row in candidate_rows if _is_stamped_issuer_row(row))
|
||||
if not stamped:
|
||||
return 0
|
||||
|
||||
healed = 0
|
||||
for row in stamped:
|
||||
try:
|
||||
await prisma_client.db.litellm_mcpservertable.update(
|
||||
where={"server_id": row.server_id},
|
||||
data={"issuer": None, "updated_by": _BACKFILL_ACTOR},
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 - per-row best effort; the next boot retries
|
||||
verbose_proxy_logger.warning(
|
||||
"MCP issuer stamp backfill: could not heal server_id=%s: %s", row.server_id, exc
|
||||
)
|
||||
continue
|
||||
healed += 1
|
||||
verbose_proxy_logger.warning(
|
||||
"MCP issuer stamp backfill: cleared issuer %r on server_id=%s (alias=%s). OAuth discovery "
|
||||
"had written that value onto the Issuer column, which made the server issuer-anchored and "
|
||||
"fail-closed, and its configured Authorization/Token/Registration URLs were being ignored "
|
||||
"as a result; those now apply again. If you pinned this issuer deliberately, set it again "
|
||||
"via the dashboard or PUT /v1/mcp/server to restore RFC 8414 section 3.3 anchoring.",
|
||||
row.issuer,
|
||||
row.server_id,
|
||||
row.alias or row.server_name,
|
||||
)
|
||||
|
||||
if healed:
|
||||
verbose_proxy_logger.warning(
|
||||
"MCP issuer stamp backfill: healed %d server(s) whose Issuer had been written by OAuth "
|
||||
"discovery rather than by an admin",
|
||||
healed,
|
||||
)
|
||||
return healed
|
||||
|
|
@ -1169,7 +1169,6 @@ class GenerateKeyResponse(KeyRequestBase):
|
|||
class UpdateKeyRequest(KeyRequestBase):
|
||||
# Note: the defaults of all Params here MUST BE NONE
|
||||
# else they will get overwritten
|
||||
key: str # type: ignore
|
||||
duration: Optional[str] = None
|
||||
spend: Optional[float] = None
|
||||
metadata: Optional[dict] = None
|
||||
|
|
@ -1186,6 +1185,12 @@ class UpdateKeyRequest(KeyRequestBase):
|
|||
raise ValueError("temp_budget_increase and temp_budget_expiry must be set together")
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_key_identifier(self) -> "UpdateKeyRequest":
|
||||
if self.key is None and self.key_alias is None:
|
||||
raise ValueError("either key or key_alias must be provided")
|
||||
return self
|
||||
|
||||
|
||||
class RegenerateKeyRequest(GenerateKeyRequest):
|
||||
# This needs to be different from UpdateKeyRequest, because "key" is optional for this
|
||||
|
|
@ -4330,7 +4335,13 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
|
|||
team_id_upsert: bool = False
|
||||
team_ids_jwt_field: Optional[str] = None
|
||||
upsert_sso_user_to_team: bool = False
|
||||
team_allowed_routes: List[str] = ["openai_routes", "info_routes", "mcp_routes"]
|
||||
team_allowed_routes: List[str] = [
|
||||
"openai_routes",
|
||||
"info_routes",
|
||||
"mcp_routes",
|
||||
"/v1/messages",
|
||||
"/v1/messages/count_tokens",
|
||||
]
|
||||
team_id_default: Optional[str] = Field(
|
||||
default=None,
|
||||
description="If no team_id given, default permissions/spend-tracking to this team.s",
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -1432,7 +1432,7 @@ def _extract_models_from_managed_resource_id(
|
|||
)
|
||||
_append_model_candidates(
|
||||
candidates=candidates,
|
||||
value=get_model_id_from_unified_batch_id(unified_file_id),
|
||||
value=_resolve_model_id_with_router(get_model_id_from_unified_batch_id(unified_file_id), llm_router),
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug("Unable to extract model from managed file/batch ID: %s", str(e))
|
||||
|
|
@ -1442,7 +1442,10 @@ def _extract_models_from_managed_resource_id(
|
|||
|
||||
parsed_id = parse_unified_id(resource_id)
|
||||
if parsed_id:
|
||||
_append_model_candidates(candidates=candidates, value=parsed_id.get("model_id"))
|
||||
_append_model_candidates(
|
||||
candidates=candidates,
|
||||
value=_resolve_model_id_with_router(parsed_id.get("model_id"), llm_router),
|
||||
)
|
||||
_append_model_candidates(candidates=candidates, value=parsed_id.get("target_model_names"))
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug("Unable to extract model from unified managed resource ID: %s", str(e))
|
||||
|
|
|
|||
|
|
@ -1960,17 +1960,6 @@ async def _user_api_key_auth_builder(
|
|||
else:
|
||||
valid_token.team_object_permission = None
|
||||
|
||||
# Cache under the canonical "team_id:{id}" key so get_team_object and
|
||||
# _update_team_cache serve this write from the L2 cache. The guard keeps a
|
||||
# non-team (personal) key, whose team_id is None, from reaching the cache
|
||||
# layer, which Redis rejects with a NoneType key error.
|
||||
if valid_token.team_id is not None and _team_obj is not None:
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=f"team_id:{valid_token.team_id}",
|
||||
value=_team_obj,
|
||||
model_type=LiteLLM_TeamTableCachedObj,
|
||||
)
|
||||
|
||||
# Fetch project object if key belongs to a project
|
||||
_project_obj = None
|
||||
if valid_token.project_id is not None:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -10,11 +10,32 @@ uv tool install 'litellm[proxy]'
|
|||
|
||||
## Configuration
|
||||
|
||||
The CLI can be configured using environment variables or command-line options:
|
||||
The CLI can be configured using environment variables, command-line options, or a persistent config file:
|
||||
|
||||
- `LITELLM_PROXY_URL`: Base URL of the LiteLLM proxy server (default: http://localhost:4000)
|
||||
- `LITELLM_PROXY_API_KEY`: API key for authentication
|
||||
|
||||
To stop exporting `LITELLM_PROXY_URL` in every shell session, store the proxy URL once in `~/.litellm/config.json`:
|
||||
|
||||
```bash
|
||||
lite config set base_url https://your-proxy.example.com
|
||||
```
|
||||
|
||||
Manage the stored config with:
|
||||
|
||||
```bash
|
||||
lite config get base_url # print the stored value
|
||||
lite config get # print all stored config
|
||||
lite config unset base_url # remove the stored value
|
||||
```
|
||||
|
||||
The base URL is resolved in this order of precedence:
|
||||
|
||||
1. `--base-url` command-line option
|
||||
2. `LITELLM_PROXY_URL` environment variable
|
||||
3. `base_url` from `~/.litellm/config.json`
|
||||
4. `http://localhost:4000`
|
||||
|
||||
## Global Options
|
||||
|
||||
- `--version`, `-v`: Print the LiteLLM Proxy client and server version and exit.
|
||||
|
|
@ -581,6 +602,8 @@ The CLI respects the following environment variables:
|
|||
- `LITELLM_PROXY_URL`: Base URL of the proxy server
|
||||
- `LITELLM_PROXY_API_KEY`: API key for authentication
|
||||
|
||||
`LITELLM_PROXY_URL` takes precedence over a `base_url` stored via `lite config set`, and the `--base-url` option overrides both. See the Configuration section for the full precedence order.
|
||||
|
||||
## Examples
|
||||
|
||||
1. List all models in table format:
|
||||
|
|
|
|||
|
|
@ -15,6 +15,8 @@ from rich.table import Table
|
|||
from litellm.constants import CLI_JWT_EXPIRATION_HOURS
|
||||
from litellm.litellm_core_utils.cli_token_utils import is_cli_token_fresh
|
||||
|
||||
from .private_json import write_private_json
|
||||
|
||||
|
||||
# Token storage utilities
|
||||
def get_token_file_path() -> str:
|
||||
|
|
@ -27,11 +29,7 @@ def get_token_file_path() -> str:
|
|||
|
||||
def save_token(token_data: Dict[str, Any]) -> None:
|
||||
"""Save token data to file"""
|
||||
token_file = get_token_file_path()
|
||||
with open(token_file, "w") as f:
|
||||
json.dump(token_data, f, indent=2)
|
||||
# Set file permissions to be readable only by owner
|
||||
os.chmod(token_file, 0o600)
|
||||
write_private_json(get_token_file_path(), token_data)
|
||||
|
||||
|
||||
def load_token() -> Optional[Dict[str, Any]]:
|
||||
|
|
|
|||
108
litellm/proxy/client/cli/commands/config.py
Normal file
108
litellm/proxy/client/cli/commands/config.py
Normal file
|
|
@ -0,0 +1,108 @@
|
|||
import json
|
||||
import os
|
||||
import sys
|
||||
from collections.abc import Mapping
|
||||
from pathlib import Path
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import click
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from .private_json import write_private_json
|
||||
|
||||
ALLOWED_CONFIG_KEYS: tuple[str, ...] = ("base_url",)
|
||||
|
||||
_config_adapter: TypeAdapter[Mapping[str, str]] = TypeAdapter(Mapping[str, str])
|
||||
|
||||
|
||||
def get_config_file_path() -> str:
|
||||
"""Get the path to the persistent CLI config file"""
|
||||
home_dir = Path.home()
|
||||
config_dir = home_dir / ".litellm"
|
||||
return str(config_dir / "config.json")
|
||||
|
||||
|
||||
def load_config() -> Mapping[str, str]:
|
||||
"""Load CLI config from file; returns {} if missing or unreadable"""
|
||||
try:
|
||||
config_file = get_config_file_path()
|
||||
except RuntimeError:
|
||||
return {}
|
||||
if not os.path.exists(config_file):
|
||||
return {}
|
||||
try:
|
||||
with open(config_file, "r") as f:
|
||||
return _config_adapter.validate_python(json.load(f))
|
||||
except (OSError, ValueError) as e:
|
||||
click.echo(f"Warning: ignoring invalid config file {config_file}: {e}", err=True)
|
||||
return {}
|
||||
|
||||
|
||||
def save_config(config: Mapping[str, str]) -> None:
|
||||
"""Save CLI config to file"""
|
||||
write_private_json(get_config_file_path(), config)
|
||||
|
||||
|
||||
def get_config_value(key: str) -> str | None:
|
||||
"""Get a single value from the persistent CLI config"""
|
||||
return load_config().get(key)
|
||||
|
||||
|
||||
@click.group(name="config")
|
||||
def config_commands() -> None:
|
||||
"""Manage persistent CLI configuration (~/.litellm/config.json)"""
|
||||
|
||||
|
||||
@config_commands.command(name="set")
|
||||
@click.argument("key")
|
||||
@click.argument("value")
|
||||
def set_config(key: str, value: str) -> None:
|
||||
"""Set a config KEY to VALUE (e.g. `lite config set base_url https://your-proxy.example.com`)"""
|
||||
if key not in ALLOWED_CONFIG_KEYS:
|
||||
raise click.UsageError(f"Unknown config key '{key}'. Allowed keys: {', '.join(ALLOWED_CONFIG_KEYS)}")
|
||||
|
||||
if key == "base_url":
|
||||
parsed = urlparse(value)
|
||||
if parsed.scheme not in ("http", "https") or not parsed.netloc:
|
||||
raise click.UsageError("base_url must be a full http:// or https:// URL including a host")
|
||||
if "?" in value or "#" in value:
|
||||
raise click.UsageError("base_url must not include a query string or fragment")
|
||||
|
||||
normalized_value = value.rstrip("/")
|
||||
save_config({**load_config(), key: normalized_value})
|
||||
click.echo(f"Set {key} = {normalized_value} in {get_config_file_path()}")
|
||||
|
||||
|
||||
@config_commands.command(name="get")
|
||||
@click.argument("key", required=False)
|
||||
def get_config(key: str | None) -> None:
|
||||
"""Print the value of KEY, or all stored config when KEY is omitted"""
|
||||
config = load_config()
|
||||
|
||||
if key is not None:
|
||||
value = config.get(key)
|
||||
if value is None:
|
||||
click.echo(f"{key} is not set", err=True)
|
||||
sys.exit(1)
|
||||
click.echo(value)
|
||||
return
|
||||
|
||||
if not config:
|
||||
click.echo("(no config set)")
|
||||
return
|
||||
|
||||
for entry_key, entry_value in config.items():
|
||||
click.echo(f"{entry_key} = {entry_value}")
|
||||
|
||||
|
||||
@config_commands.command(name="unset")
|
||||
@click.argument("key")
|
||||
def unset_config(key: str) -> None:
|
||||
"""Remove KEY from the config file"""
|
||||
config = load_config()
|
||||
if key not in config:
|
||||
click.echo(f"{key} was not set")
|
||||
return
|
||||
|
||||
save_config({k: v for k, v in config.items() if k != key})
|
||||
click.echo(f"Removed {key} from {get_config_file_path()}")
|
||||
20
litellm/proxy/client/cli/commands/private_json.py
Normal file
20
litellm/proxy/client/cli/commands/private_json.py
Normal file
|
|
@ -0,0 +1,20 @@
|
|||
import json
|
||||
import os
|
||||
import tempfile
|
||||
from collections.abc import Mapping
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def write_private_json(path: str, data: Mapping[str, object]) -> None:
|
||||
"""Atomically write JSON to path with owner-only permissions (0600)"""
|
||||
parent = Path(path).parent
|
||||
parent.mkdir(parents=True, exist_ok=True)
|
||||
fd, tmp_path = tempfile.mkstemp(dir=str(parent), prefix=".tmp-", suffix=".json")
|
||||
try:
|
||||
with os.fdopen(fd, "w") as f:
|
||||
json.dump(data, f, indent=2)
|
||||
f.flush()
|
||||
os.fsync(f.fileno())
|
||||
os.replace(tmp_path, path)
|
||||
finally:
|
||||
Path(tmp_path).unlink(missing_ok=True)
|
||||
|
|
@ -11,6 +11,7 @@ from .commands.agents import agent_commands
|
|||
from .commands.auth import auth_group, get_stored_api_key, login, logout, whoami
|
||||
from .commands.autoroute.commands import autoroute_group
|
||||
from .commands.chat import chat
|
||||
from .commands.config import config_commands, get_config_value
|
||||
from .commands.credentials import credentials
|
||||
from .commands.encryption import encryption
|
||||
from .commands.http import http
|
||||
|
|
@ -45,27 +46,16 @@ def print_version(base_url: str, api_key: Optional[str]):
|
|||
@click.option(
|
||||
"--version",
|
||||
"-v",
|
||||
"show_version",
|
||||
is_flag=True,
|
||||
is_eager=True,
|
||||
expose_value=False,
|
||||
help="Show the LiteLLM Proxy CLI and server version and exit.",
|
||||
callback=lambda ctx, param, value: (
|
||||
(
|
||||
print_version(
|
||||
ctx.params.get("base_url") or "http://localhost:4000",
|
||||
ctx.params.get("api_key"),
|
||||
)
|
||||
or ctx.exit()
|
||||
)
|
||||
if value and not ctx.resilient_parsing
|
||||
else None
|
||||
),
|
||||
)
|
||||
@click.option(
|
||||
"--base-url",
|
||||
envvar="LITELLM_PROXY_URL",
|
||||
show_envvar=True,
|
||||
default="http://localhost:4000",
|
||||
default=None,
|
||||
show_default="base_url from `lite config`, else http://localhost:4000",
|
||||
help="Base URL of the LiteLLM proxy server",
|
||||
)
|
||||
@click.option(
|
||||
|
|
@ -75,13 +65,16 @@ def print_version(base_url: str, api_key: Optional[str]):
|
|||
help="API key for authentication",
|
||||
)
|
||||
@click.pass_context
|
||||
def cli(ctx: click.Context, base_url: str, api_key: Optional[str]) -> None:
|
||||
def cli(ctx: click.Context, show_version: bool, base_url: str | None, api_key: Optional[str]) -> None:
|
||||
"""LiteLLM Proxy CLI - Manage your LiteLLM proxy server"""
|
||||
ctx.ensure_object(dict)
|
||||
|
||||
stored_base_url = get_config_value("base_url")
|
||||
base_url_provided = base_url is not None
|
||||
|
||||
# Normalize once here so every downstream command (login, agents, http, ...) can safely
|
||||
# do f"{base_url}/some/path" without producing a double slash.
|
||||
base_url = base_url.rstrip("/")
|
||||
base_url = ((stored_base_url or "http://localhost:4000") if base_url is None else base_url).rstrip("/")
|
||||
|
||||
# If no API key provided via flag or environment variable, try to load from saved token.
|
||||
# Pass base_url so we only use the stored key when it was issued for this server.
|
||||
|
|
@ -94,8 +87,13 @@ def cli(ctx: click.Context, base_url: str, api_key: Optional[str]) -> None:
|
|||
# apiKeyHelper is invoked bare (no flags) -- commands that must work
|
||||
# unattended (print-token) need to tell "user didn't say" apart from
|
||||
# "user said localhost:4000 on purpose" so they can fall back to
|
||||
# whatever server the stored token was actually issued for.
|
||||
ctx.obj["base_url_explicit"] = ctx.get_parameter_source("base_url") != click.core.ParameterSource.DEFAULT
|
||||
# whatever server the stored token was actually issued for. A base_url
|
||||
# saved via `lite config set` counts as the user saying it.
|
||||
ctx.obj["base_url_explicit"] = base_url_provided or bool(stored_base_url)
|
||||
|
||||
if show_version:
|
||||
print_version(base_url, api_key)
|
||||
ctx.exit()
|
||||
|
||||
# If no subcommand was invoked, start interactive mode
|
||||
if ctx.invoked_subcommand is None:
|
||||
|
|
@ -141,6 +139,7 @@ cli.add_command(down)
|
|||
cli.add_command(model_groups)
|
||||
# Add the autoroute command group (QA auto-routing against your real proxy)
|
||||
cli.add_command(autoroute_group, name="autoroute")
|
||||
cli.add_command(config_commands)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
|
|
|||
|
|
@ -1829,6 +1829,15 @@ async def add_litellm_data_to_request(
|
|||
return data
|
||||
|
||||
|
||||
def _warn_stale_team_alias_once(warning_key: str, message: str, *args: str) -> None:
|
||||
if warning_key in _STALE_TEAM_ALIAS_WARNING_KEYS:
|
||||
return
|
||||
_STALE_TEAM_ALIAS_WARNING_KEYS[warning_key] = None
|
||||
while len(_STALE_TEAM_ALIAS_WARNING_KEYS) > _MAX_STALE_ALIAS_WARNING_KEYS:
|
||||
_STALE_TEAM_ALIAS_WARNING_KEYS.popitem(last=False)
|
||||
verbose_proxy_logger.warning(message, *args)
|
||||
|
||||
|
||||
def _update_model_if_team_alias_exists(
|
||||
data: dict,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -1848,49 +1857,63 @@ def _update_model_if_team_alias_exists(
|
|||
Note: model_aliases for team models are deprecated. This function only applies
|
||||
to legacy non-team-scoped aliases. Team-scoped deployments use team_public_model_name
|
||||
and are resolved via map_team_model in route_llm_request.
|
||||
|
||||
An alias that targets a team-scoped internal name (``model_name_{team_id}_{uuid}``)
|
||||
with no live deployment behind it is never applied: the deployment was deleted, so
|
||||
the rewrite could only fail with an error naming a model the caller never sent.
|
||||
Keeping the requested model name lets it resolve against the deployments that still
|
||||
exist (e.g. a gateway-level model group shared with the team).
|
||||
"""
|
||||
_model = data.get("model")
|
||||
if _model and user_api_key_dict.team_model_aliases and _model in user_api_key_dict.team_model_aliases:
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
if not _model or not user_api_key_dict.team_model_aliases or _model not in user_api_key_dict.team_model_aliases:
|
||||
return
|
||||
|
||||
# Skip alias rewrite if this model resolves to team-specific deployments
|
||||
# (team models use team_public_model_name, not model_aliases)
|
||||
aliased_target = user_api_key_dict.team_model_aliases[_model]
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
|
||||
# Optional bypass for stale aliases from pre-PR deployments:
|
||||
# only enabled via feature flag to preserve backwards compatibility.
|
||||
# Cached at module level to avoid hot-path secret lookups on every request.
|
||||
global _ENABLE_TEAM_STALE_ALIAS_BYPASS
|
||||
if _ENABLE_TEAM_STALE_ALIAS_BYPASS is None:
|
||||
_ENABLE_TEAM_STALE_ALIAS_BYPASS = get_secret_bool("LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS", False)
|
||||
enable_stale_alias_bypass = _ENABLE_TEAM_STALE_ALIAS_BYPASS
|
||||
# Check if the alias points to a team-scoped UUID name
|
||||
# (format: "model_name_{team_id}_{uuid}")
|
||||
is_stale_team_alias = aliased_target.startswith(f"model_name_{user_api_key_dict.team_id}_")
|
||||
if is_stale_team_alias and llm_router:
|
||||
# This is a stale alias from pre-PR deployments.
|
||||
# Check if current team deployments exist for the public name.
|
||||
key = (user_api_key_dict.team_id, _model)
|
||||
if key in llm_router.team_model_to_deployment_indices:
|
||||
if enable_stale_alias_bypass:
|
||||
# Team deployments exist; skip stale alias
|
||||
return
|
||||
warning_key = f"{user_api_key_dict.team_id}:{_model}:{aliased_target}"
|
||||
if warning_key not in _STALE_TEAM_ALIAS_WARNING_KEYS:
|
||||
_STALE_TEAM_ALIAS_WARNING_KEYS[warning_key] = None
|
||||
while len(_STALE_TEAM_ALIAS_WARNING_KEYS) > _MAX_STALE_ALIAS_WARNING_KEYS:
|
||||
_STALE_TEAM_ALIAS_WARNING_KEYS.popitem(last=False)
|
||||
verbose_proxy_logger.warning(
|
||||
"Stale team model alias detected for model='%s', team_id='%s'. "
|
||||
"New sibling deployments may be unreachable. "
|
||||
"Set LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS=true to enable "
|
||||
"team-scoped sibling routing.",
|
||||
_sanitize_for_log(_model),
|
||||
user_api_key_dict.team_id,
|
||||
)
|
||||
# Skip alias rewrite if this model resolves to team-specific deployments
|
||||
# (team models use team_public_model_name, not model_aliases)
|
||||
aliased_target = user_api_key_dict.team_model_aliases[_model]
|
||||
|
||||
data["model"] = aliased_target
|
||||
return
|
||||
# Optional bypass for stale aliases from pre-PR deployments:
|
||||
# only enabled via feature flag to preserve backwards compatibility.
|
||||
# Cached at module level to avoid hot-path secret lookups on every request.
|
||||
global _ENABLE_TEAM_STALE_ALIAS_BYPASS
|
||||
if _ENABLE_TEAM_STALE_ALIAS_BYPASS is None:
|
||||
_ENABLE_TEAM_STALE_ALIAS_BYPASS = get_secret_bool("LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS", False)
|
||||
enable_stale_alias_bypass = _ENABLE_TEAM_STALE_ALIAS_BYPASS
|
||||
# Check if the alias points to a team-scoped UUID name
|
||||
# (format: "model_name_{team_id}_{uuid}")
|
||||
is_stale_team_alias = aliased_target.startswith(f"model_name_{user_api_key_dict.team_id}_")
|
||||
if is_stale_team_alias and llm_router:
|
||||
if aliased_target not in llm_router.model_name_to_deployment_indices:
|
||||
_warn_stale_team_alias_once(
|
||||
f"deleted:{user_api_key_dict.team_id}:{_model}:{aliased_target}",
|
||||
"Team model alias for model='%s', team_id='%s' targets '%s', which has no live "
|
||||
"deployment. Routing with the requested model name instead; remove the stale "
|
||||
"entry from the team's model_aliases to silence this warning.",
|
||||
_sanitize_for_log(_model),
|
||||
_sanitize_for_log(user_api_key_dict.team_id),
|
||||
_sanitize_for_log(aliased_target),
|
||||
)
|
||||
return
|
||||
# This is a stale alias from pre-PR deployments.
|
||||
# Check if current team deployments exist for the public name.
|
||||
key = (user_api_key_dict.team_id, _model)
|
||||
if key in llm_router.team_model_to_deployment_indices:
|
||||
if enable_stale_alias_bypass:
|
||||
# Team deployments exist; skip stale alias
|
||||
return
|
||||
_warn_stale_team_alias_once(
|
||||
f"{user_api_key_dict.team_id}:{_model}:{aliased_target}",
|
||||
"Stale team model alias detected for model='%s', team_id='%s'. "
|
||||
"New sibling deployments may be unreachable. "
|
||||
"Set LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS=true to enable "
|
||||
"team-scoped sibling routing.",
|
||||
_sanitize_for_log(_model),
|
||||
_sanitize_for_log(user_api_key_dict.team_id),
|
||||
)
|
||||
|
||||
data["model"] = aliased_target
|
||||
|
||||
|
||||
def _update_model_if_key_alias_exists(
|
||||
|
|
|
|||
|
|
@ -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(
|
|||
Update an existing API key's parameters.
|
||||
|
||||
Parameters:
|
||||
- key: str - The key to update
|
||||
- key_alias: Optional[str] - User-friendly key alias
|
||||
- key: Optional[str] - The key to update. Either key or key_alias must be provided.
|
||||
- key_alias: Optional[str] - User-friendly key alias. If key is omitted, also identifies the key to update (must match exactly one key, same as /key/delete's key_aliases)
|
||||
- user_id: Optional[str] - User ID associated with key
|
||||
- team_id: Optional[str] - Team ID associated with key
|
||||
- agent_id: Optional[str] - The agent id associated with the key.
|
||||
|
|
@ -2592,14 +2639,14 @@ async def update_key_fn(
|
|||
detail={"error": f"max_budget must be a non-negative finite number. Received: {data.max_budget}"},
|
||||
)
|
||||
|
||||
data_json: dict = data.model_dump(exclude_unset=True)
|
||||
key = data_json.pop("key")
|
||||
|
||||
# get the row from db
|
||||
existing_key_row = await _get_and_validate_existing_key(
|
||||
token=data.key,
|
||||
prisma_client=prisma_client,
|
||||
key_alias=data.key_alias,
|
||||
)
|
||||
key = _resolve_token_to_update(data=data, existing_key_row=existing_key_row)
|
||||
data.key = key
|
||||
|
||||
await _validate_update_key_data(
|
||||
data=data,
|
||||
|
|
@ -3047,10 +3094,12 @@ async def bulk_update_team_keys(
|
|||
)
|
||||
|
||||
# team_id from validated scope, never user payload — drives _check_team_key_limits.
|
||||
update_key_request = UpdateKeyRequest(
|
||||
key=token,
|
||||
team_id=data.team_id,
|
||||
**update_field_dict,
|
||||
update_key_request = UpdateKeyRequest.model_validate(
|
||||
{
|
||||
"key": token,
|
||||
"team_id": data.team_id,
|
||||
**update_field_dict,
|
||||
}
|
||||
)
|
||||
updated_key_info = await _process_single_key_update(
|
||||
update_key_request=update_key_request,
|
||||
|
|
@ -4048,12 +4097,14 @@ def _transform_verification_tokens_to_deleted_records(
|
|||
records = []
|
||||
for key in keys:
|
||||
key_payload = key.model_dump()
|
||||
deleted_record = LiteLLM_DeletedVerificationToken(
|
||||
**key_payload,
|
||||
deleted_at=deleted_at,
|
||||
deleted_by=user_api_key_dict.user_id,
|
||||
deleted_by_api_key=user_api_key_dict.api_key,
|
||||
litellm_changed_by=litellm_changed_by,
|
||||
deleted_record = LiteLLM_DeletedVerificationToken.model_validate(
|
||||
{
|
||||
**key_payload,
|
||||
"deleted_at": deleted_at,
|
||||
"deleted_by": user_api_key_dict.user_id,
|
||||
"deleted_by_api_key": user_api_key_dict.api_key,
|
||||
"litellm_changed_by": litellm_changed_by,
|
||||
}
|
||||
)
|
||||
record = deleted_record.model_dump()
|
||||
|
||||
|
|
@ -4535,7 +4586,7 @@ async def _execute_virtual_key_regeneration(
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
response = GenerateKeyResponse(**updated_token_dict)
|
||||
response = GenerateKeyResponse.model_validate(updated_token_dict)
|
||||
asyncio.create_task(
|
||||
KeyManagementEventHooks.async_key_rotated_hook(
|
||||
data=data,
|
||||
|
|
@ -4853,7 +4904,7 @@ async def _check_proxy_or_team_admin_for_key(
|
|||
)
|
||||
|
||||
|
||||
def _validate_reset_spend_value(reset_to: Any, key_in_db: LiteLLM_VerificationToken) -> float:
|
||||
def _validate_reset_spend_value(reset_to: object, key_in_db: LiteLLM_VerificationToken) -> float:
|
||||
if not isinstance(reset_to, (int, float)):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
|
|
@ -5029,7 +5080,7 @@ async def validate_key_list_check(
|
|||
code=status.HTTP_403_FORBIDDEN,
|
||||
)
|
||||
|
||||
complete_user_info = LiteLLM_UserTable(**complete_user_info_db_obj.model_dump())
|
||||
complete_user_info = LiteLLM_UserTable.model_validate(complete_user_info_db_obj.model_dump())
|
||||
|
||||
# internal user can only see their own keys
|
||||
if user_id:
|
||||
|
|
@ -5102,7 +5153,7 @@ async def _fetch_user_team_objects(
|
|||
if teams is None:
|
||||
return []
|
||||
|
||||
return [LiteLLM_TeamTable(**team.model_dump()) for team in teams]
|
||||
return [LiteLLM_TeamTable.model_validate(team.model_dump()) for team in teams]
|
||||
|
||||
|
||||
def _get_admin_team_ids_from_objects(
|
||||
|
|
@ -5851,7 +5902,7 @@ async def _list_key_helper(
|
|||
if return_full_object is True or (expand and "user" in expand):
|
||||
if use_deleted_table:
|
||||
# Use deleted key type to preserve deleted_at, deleted_by, etc.
|
||||
key_list.append(LiteLLM_DeletedVerificationToken(**key_dict))
|
||||
key_list.append(LiteLLM_DeletedVerificationToken.model_validate(key_dict))
|
||||
else:
|
||||
key_list.append(UserAPIKeyAuth(**key_dict)) # Return full key object
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -1526,7 +1526,6 @@ if MCP_AVAILABLE:
|
|||
temporary_server = await global_mcp_server_manager.build_mcp_server_from_table(
|
||||
temp_record,
|
||||
credentials_are_encrypted=False,
|
||||
persist_discovered_endpoints=False,
|
||||
)
|
||||
_cache_temporary_mcp_server(
|
||||
temporary_server,
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ Endpoints here:
|
|||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, Dict, List, Tuple
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
|
@ -16,6 +17,9 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
|||
|
||||
# Clear cache and reload models to pick up the access group changes
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
live_model_ids_snapshot,
|
||||
model_info_as_mapping,
|
||||
reload_serving_verdict,
|
||||
clear_cache,
|
||||
)
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
|
@ -72,11 +76,92 @@ def add_access_group_to_deployment(model_info: Dict[str, Any], access_group: str
|
|||
return model_info, True
|
||||
|
||||
|
||||
def _raise_http_if_reload_degraded_serving(
|
||||
before: frozenset[str],
|
||||
written_models: Sequence[tuple[str, object]],
|
||||
access_group: str,
|
||||
) -> None:
|
||||
"""Same verdict as the model-write endpoints, expressed through this file's
|
||||
HTTPException error convention, with the metadata-only obligation: these writes
|
||||
change group membership, not the models themselves, so a row that was already not
|
||||
serving before the reload is never blamed here; only a model this reload stopped
|
||||
serving is reported."""
|
||||
missing, collateral = reload_serving_verdict(before=before, written_models=written_models, written_must_serve=False)
|
||||
gone = tuple(dict.fromkeys((*missing, *collateral)))
|
||||
if not gone:
|
||||
return
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={
|
||||
"error": (
|
||||
f"Access group '{access_group}' was saved to the database, but model id(s) {list(gone)} that "
|
||||
"this pod was serving are no longer live after the reload it triggered. Other pods reload on "
|
||||
"their own interval. Check server logs for 'Error upserting deployment' for the cause."
|
||||
)
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
async def _tag_deployment_with_access_group(
|
||||
model_id: str,
|
||||
model_info: object,
|
||||
access_group: str,
|
||||
prisma_client: PrismaClient,
|
||||
) -> tuple[str, Mapping[str, object]] | None:
|
||||
"""Write `access_group` into one deployment's model_info; returns the
|
||||
(model_id, updated model_info) pair when a write happened, None when the
|
||||
deployment already carried the group."""
|
||||
updated_model_info, was_modified = add_access_group_to_deployment(
|
||||
model_info=dict(_readable_model_info_or_raise(model_id=model_id, model_info=model_info)),
|
||||
access_group=access_group,
|
||||
)
|
||||
if not was_modified:
|
||||
return None
|
||||
await ModelRepository(prisma_client).table.update(
|
||||
where={"model_id": model_id},
|
||||
data={"model_info": json.dumps(updated_model_info)},
|
||||
)
|
||||
verbose_proxy_logger.debug(f"Updated deployment {model_id} with access group: {access_group}")
|
||||
return (model_id, updated_model_info)
|
||||
|
||||
|
||||
def _readable_model_info_or_raise(model_id: str, model_info: object) -> Mapping[str, object]:
|
||||
"""These helpers rewrite the model_info column wholesale, so a present-but-unreadable
|
||||
value must refuse loudly rather than be silently replaced with a fresh object; an
|
||||
absent value stays a legitimate empty start."""
|
||||
parsed = model_info_as_mapping(model_info)
|
||||
if parsed is None and model_info is not None:
|
||||
raise ValueError(f"model_info for deployment {model_id} is not a readable JSON object; refusing to rewrite it")
|
||||
return parsed or {}
|
||||
|
||||
|
||||
async def _strip_access_group_from_deployment(
|
||||
model_id: str,
|
||||
model_info: object,
|
||||
access_group: str,
|
||||
prisma_client: PrismaClient,
|
||||
) -> tuple[str, Mapping[str, object]] | None:
|
||||
"""Remove `access_group` from one deployment's model_info; returns the
|
||||
(model_id, updated model_info) pair when a write happened, None when the
|
||||
deployment did not carry the group."""
|
||||
updated_model_info, was_modified = remove_access_group_from_deployment(
|
||||
model_info=dict(_readable_model_info_or_raise(model_id=model_id, model_info=model_info)),
|
||||
access_group=access_group,
|
||||
)
|
||||
if not was_modified:
|
||||
return None
|
||||
await ModelRepository(prisma_client).table.update(
|
||||
where={"model_id": model_id},
|
||||
data={"model_info": json.dumps(updated_model_info)},
|
||||
)
|
||||
return (model_id, updated_model_info)
|
||||
|
||||
|
||||
async def update_deployments_with_access_group(
|
||||
model_names: List[str],
|
||||
access_group: str,
|
||||
prisma_client: PrismaClient,
|
||||
) -> int:
|
||||
) -> tuple[tuple[str, Mapping[str, object]], ...]:
|
||||
"""
|
||||
Update all deployments for the given model names to include the access group.
|
||||
|
||||
|
|
@ -86,20 +171,15 @@ async def update_deployments_with_access_group(
|
|||
prisma_client: Database client
|
||||
|
||||
Returns:
|
||||
int: Number of deployments updated
|
||||
The (model_id, updated model_info) pair of every deployment actually written,
|
||||
so callers can verify each one survived the post-write reload
|
||||
"""
|
||||
models_updated = 0
|
||||
deployments = await ModelRepository(prisma_client).table.find_many(where={"model_name": {"in": model_names}})
|
||||
verbose_proxy_logger.debug(f"Found {len(deployments)} deployments for model_names: {model_names}")
|
||||
|
||||
found_names = {deployment.model_name for deployment in deployments}
|
||||
for model_name in model_names:
|
||||
verbose_proxy_logger.debug(f"Updating deployments for model_name: {model_name}")
|
||||
|
||||
# Get all deployments with this model_name
|
||||
deployments = await ModelRepository(prisma_client).table.find_many(where={"model_name": model_name})
|
||||
|
||||
verbose_proxy_logger.debug(f"Found {len(deployments)} deployments for model_name: {model_name}")
|
||||
|
||||
# If no deployments found, this is a config model (not in DB)
|
||||
if len(deployments) == 0:
|
||||
if model_name not in found_names:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
|
|
@ -107,65 +187,52 @@ async def update_deployments_with_access_group(
|
|||
},
|
||||
)
|
||||
|
||||
# Update each deployment
|
||||
for deployment in deployments:
|
||||
model_info = deployment.model_info or {}
|
||||
|
||||
# Add access group using helper
|
||||
updated_model_info, was_modified = add_access_group_to_deployment(
|
||||
model_info=model_info,
|
||||
access_group=access_group,
|
||||
)
|
||||
|
||||
# Only update in DB if modified
|
||||
if was_modified:
|
||||
await ModelRepository(prisma_client).table.update(
|
||||
where={"model_id": deployment.model_id},
|
||||
data={"model_info": json.dumps(updated_model_info)},
|
||||
)
|
||||
|
||||
models_updated += 1
|
||||
verbose_proxy_logger.debug(
|
||||
f"Updated deployment {deployment.model_id} with access group: {access_group}"
|
||||
)
|
||||
|
||||
return models_updated
|
||||
tagged = [
|
||||
await _tag_deployment_with_access_group(
|
||||
model_id=deployment.model_id,
|
||||
model_info=deployment.model_info,
|
||||
access_group=access_group,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
for deployment in deployments
|
||||
]
|
||||
return tuple(pair for pair in tagged if pair is not None)
|
||||
|
||||
|
||||
async def update_specific_deployments_with_access_group(
|
||||
model_ids: List[str],
|
||||
access_group: str,
|
||||
prisma_client: PrismaClient,
|
||||
) -> int:
|
||||
) -> tuple[tuple[str, Mapping[str, object]], ...]:
|
||||
"""
|
||||
Update specific deployments (by model_id) to include the access group.
|
||||
|
||||
Unlike update_deployments_with_access_group which tags ALL deployments sharing
|
||||
a model_name, this function only tags the specific deployments identified by
|
||||
their unique model_id.
|
||||
their unique model_id. Returns the (model_id, updated model_info) pair of every
|
||||
deployment actually written.
|
||||
"""
|
||||
models_updated = 0
|
||||
for model_id in model_ids:
|
||||
verbose_proxy_logger.debug(f"Updating specific deployment model_id: {model_id}")
|
||||
deployment = await ModelRepository(prisma_client).table.find_unique(where={"model_id": model_id})
|
||||
if deployment is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": f"Deployment with model_id '{model_id}' not found in Database."},
|
||||
)
|
||||
model_info = deployment.model_info or {}
|
||||
updated_model_info, was_modified = add_access_group_to_deployment(
|
||||
model_info=model_info,
|
||||
verbose_proxy_logger.debug(f"Updating specific deployment model_ids: {model_ids}")
|
||||
tagged = [
|
||||
await _tag_deployment_with_access_group(
|
||||
model_id=model_id,
|
||||
model_info=(await _find_deployment_or_400(model_id=model_id, prisma_client=prisma_client)),
|
||||
access_group=access_group,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
if was_modified:
|
||||
await ModelRepository(prisma_client).table.update(
|
||||
where={"model_id": model_id},
|
||||
data={"model_info": json.dumps(updated_model_info)},
|
||||
)
|
||||
models_updated += 1
|
||||
verbose_proxy_logger.debug(f"Updated deployment {model_id} with access group: {access_group}")
|
||||
return models_updated
|
||||
for model_id in model_ids
|
||||
]
|
||||
return tuple(pair for pair in tagged if pair is not None)
|
||||
|
||||
|
||||
async def _find_deployment_or_400(model_id: str, prisma_client: PrismaClient) -> Mapping[str, object] | None:
|
||||
deployment = await ModelRepository(prisma_client).table.find_unique(where={"model_id": model_id})
|
||||
if deployment is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": f"Deployment with model_id '{model_id}' not found in Database."},
|
||||
)
|
||||
return deployment.model_info
|
||||
|
||||
|
||||
def remove_access_group_from_deployment(model_info: Dict[str, Any], access_group: str) -> Tuple[Dict[str, Any], bool]:
|
||||
|
|
@ -335,20 +402,28 @@ async def create_model_group(
|
|||
# Update deployments using the appropriate method
|
||||
if use_model_ids:
|
||||
assert data.model_ids is not None
|
||||
models_updated = await update_specific_deployments_with_access_group(
|
||||
updated_pairs = await update_specific_deployments_with_access_group(
|
||||
model_ids=data.model_ids,
|
||||
access_group=data.access_group,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
else:
|
||||
assert data.model_names is not None
|
||||
models_updated = await update_deployments_with_access_group(
|
||||
updated_pairs = await update_deployments_with_access_group(
|
||||
model_names=data.model_names,
|
||||
access_group=data.access_group,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
models_updated = len(updated_pairs)
|
||||
|
||||
live_before_reload = live_model_ids_snapshot()
|
||||
|
||||
await clear_cache()
|
||||
_raise_http_if_reload_degraded_serving(
|
||||
before=live_before_reload,
|
||||
written_models=updated_pairs,
|
||||
access_group=data.access_group,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
f"Successfully created access group '{data.access_group}' with {models_updated} models updated"
|
||||
|
|
@ -573,38 +648,42 @@ async def update_access_group(
|
|||
# Step 1: Remove access group from ALL DB deployments (skip config models)
|
||||
all_deployments = await ModelRepository(prisma_client).table.find_many()
|
||||
|
||||
for deployment in all_deployments:
|
||||
model_info = deployment.model_info or {}
|
||||
|
||||
updated_model_info, was_modified = remove_access_group_from_deployment(
|
||||
model_info=model_info,
|
||||
stripped = [
|
||||
await _strip_access_group_from_deployment(
|
||||
model_id=deployment.model_id,
|
||||
model_info=deployment.model_info,
|
||||
access_group=access_group,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
if was_modified:
|
||||
await ModelRepository(prisma_client).table.update(
|
||||
where={"model_id": deployment.model_id},
|
||||
data={"model_info": json.dumps(updated_model_info)},
|
||||
)
|
||||
for deployment in all_deployments
|
||||
]
|
||||
stripped_pairs = tuple(pair for pair in stripped if pair is not None)
|
||||
|
||||
# Step 2: Add access group using the appropriate method
|
||||
if use_model_ids:
|
||||
assert data.model_ids is not None
|
||||
models_updated = await update_specific_deployments_with_access_group(
|
||||
updated_pairs = await update_specific_deployments_with_access_group(
|
||||
model_ids=data.model_ids,
|
||||
access_group=access_group,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
else:
|
||||
assert data.model_names is not None
|
||||
models_updated = await update_deployments_with_access_group(
|
||||
updated_pairs = await update_deployments_with_access_group(
|
||||
model_names=data.model_names,
|
||||
access_group=access_group,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
models_updated = len(updated_pairs)
|
||||
|
||||
# Clear cache and reload models to pick up the access group changes
|
||||
live_before_reload = live_model_ids_snapshot()
|
||||
await clear_cache()
|
||||
_raise_http_if_reload_degraded_serving(
|
||||
before=live_before_reload,
|
||||
written_models=list({**dict(stripped_pairs), **dict(updated_pairs)}.items()),
|
||||
access_group=access_group,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
f"Successfully updated access group '{access_group}' with {models_updated} models updated"
|
||||
|
|
@ -686,25 +765,27 @@ async def delete_access_group(
|
|||
try:
|
||||
# Remove access group from all DB deployments (skip config models)
|
||||
all_deployments = await ModelRepository(prisma_client).table.find_many()
|
||||
models_updated = 0
|
||||
|
||||
for deployment in all_deployments:
|
||||
model_info = deployment.model_info or {}
|
||||
|
||||
updated_model_info, was_modified = remove_access_group_from_deployment(
|
||||
model_info=model_info,
|
||||
removed = [
|
||||
await _strip_access_group_from_deployment(
|
||||
model_id=deployment.model_id,
|
||||
model_info=deployment.model_info,
|
||||
access_group=access_group,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
if was_modified:
|
||||
await ModelRepository(prisma_client).table.update(
|
||||
where={"model_id": deployment.model_id},
|
||||
data={"model_info": json.dumps(updated_model_info)},
|
||||
)
|
||||
models_updated += 1
|
||||
for deployment in all_deployments
|
||||
]
|
||||
removed_pairs = tuple(pair for pair in removed if pair is not None)
|
||||
models_updated = len(removed_pairs)
|
||||
|
||||
# Clear cache and reload models to pick up the access group changes
|
||||
live_before_reload = live_model_ids_snapshot()
|
||||
await clear_cache()
|
||||
_raise_http_if_reload_degraded_serving(
|
||||
before=live_before_reload,
|
||||
written_models=removed_pairs,
|
||||
access_group=access_group,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
f"Successfully deleted access group '{access_group}' from {models_updated} deployments"
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ model/{model_id}/update - PATCH endpoint for model update.
|
|||
import asyncio
|
||||
import datetime
|
||||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union, cast
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Header, Request, status
|
||||
|
|
@ -52,13 +53,19 @@ from litellm.proxy.utils import PrismaClient
|
|||
from litellm.repositories.model_repository import ModelRepository
|
||||
from litellm.repositories.table_repositories import ModelTableRepository
|
||||
from litellm.repositories.team_repository import TeamRepository
|
||||
from litellm.router import Router
|
||||
from litellm.types.proxy.management_endpoints.model_management_endpoints import (
|
||||
UpdateUsefulLinksRequest,
|
||||
)
|
||||
from litellm.router_utils.auto_router_model_naming import (
|
||||
STRATEGY_ROUTER_PARAM_FIELDS,
|
||||
validate_strategy_router_model_write,
|
||||
)
|
||||
from litellm.types.router import (
|
||||
SPECIAL_MODEL_INFO_PARAMS,
|
||||
Deployment,
|
||||
DeploymentTypedDict,
|
||||
GenericLiteLLMParams,
|
||||
LiteLLMParamsTypedDict,
|
||||
updateDeployment,
|
||||
)
|
||||
|
|
@ -96,6 +103,45 @@ async def get_db_model(model_id: str, prisma_client: PrismaClient) -> Optional[D
|
|||
return deployment_pydantic_obj
|
||||
|
||||
|
||||
def _strategy_router_write_violation(
|
||||
incoming_params: GenericLiteLLMParams | None,
|
||||
existing_params: GenericLiteLLMParams | None,
|
||||
) -> str | None:
|
||||
"""Reject writes that would corrupt a strategy router's pseudo-model.
|
||||
|
||||
An auto-router deployment's ``litellm_params.model`` (``auto_router/...``) is
|
||||
the discriminator the router loads it by; a write that mangles it makes the
|
||||
router drop the deployment silently under ``ignore_invalid_deployments``.
|
||||
Only writes that supply ``litellm_params.model`` are judged, against the
|
||||
merged (stored + incoming) params, so partial patches and restores of an
|
||||
already-corrupted row stay legal. Returns the violation, or None.
|
||||
"""
|
||||
if incoming_params is None or incoming_params.model is None:
|
||||
return None
|
||||
present_fields = frozenset(
|
||||
field
|
||||
for field in STRATEGY_ROUTER_PARAM_FIELDS
|
||||
for source in (incoming_params, existing_params)
|
||||
if source is not None and getattr(source, field, None) is not None
|
||||
)
|
||||
return validate_strategy_router_model_write(model=incoming_params.model, present_fields=present_fields)
|
||||
|
||||
|
||||
def _raise_on_strategy_router_write_violation(
|
||||
incoming_params: GenericLiteLLMParams | None,
|
||||
existing_params: GenericLiteLLMParams | None,
|
||||
) -> None:
|
||||
violation = _strategy_router_write_violation(incoming_params=incoming_params, existing_params=existing_params)
|
||||
if violation is None:
|
||||
return
|
||||
raise ProxyException(
|
||||
message=violation,
|
||||
type=ProxyErrorTypes.validation_error.value,
|
||||
code=status.HTTP_400_BAD_REQUEST,
|
||||
param="litellm_params.model",
|
||||
)
|
||||
|
||||
|
||||
def update_db_model(db_model: Deployment, updated_patch: updateDeployment) -> PrismaCompatibleUpdateDBModel:
|
||||
merged_deployment_dict = DeploymentTypedDict(
|
||||
model_name=db_model.model_name,
|
||||
|
|
@ -253,6 +299,11 @@ async def patch_model(
|
|||
param="blocked",
|
||||
)
|
||||
|
||||
_raise_on_strategy_router_write_violation(
|
||||
incoming_params=patch_data.litellm_params,
|
||||
existing_params=db_model.litellm_params,
|
||||
)
|
||||
|
||||
# Handle team model updates with proper alias management
|
||||
update_data = await _update_team_model_in_db(
|
||||
db_model=db_model,
|
||||
|
|
@ -272,6 +323,7 @@ async def patch_model(
|
|||
)
|
||||
|
||||
# Clear cache and reload models (uses config setting or defaults to preserving config models for DB updates)
|
||||
live_before_reload = live_model_ids_snapshot()
|
||||
await clear_cache()
|
||||
|
||||
## CREATE AUDIT LOG ##
|
||||
|
|
@ -288,6 +340,12 @@ async def patch_model(
|
|||
)
|
||||
)
|
||||
|
||||
raise_if_reload_degraded_serving(
|
||||
before=live_before_reload,
|
||||
written_models=[(model_id, getattr(updated_model, "model_info", None))],
|
||||
action="update",
|
||||
)
|
||||
|
||||
return updated_model
|
||||
|
||||
except Exception as e:
|
||||
|
|
@ -370,6 +428,7 @@ async def _set_model_blocked_status(
|
|||
},
|
||||
)
|
||||
|
||||
live_before_reload = live_model_ids_snapshot()
|
||||
await clear_cache()
|
||||
|
||||
asyncio.create_task(
|
||||
|
|
@ -387,6 +446,12 @@ async def _set_model_blocked_status(
|
|||
)
|
||||
)
|
||||
|
||||
raise_if_reload_degraded_serving(
|
||||
before=live_before_reload,
|
||||
written_models=[(data.model_id, getattr(updated_model, "model_info", None))],
|
||||
action=action,
|
||||
)
|
||||
|
||||
return updated_model
|
||||
|
||||
except Exception as e:
|
||||
|
|
@ -713,13 +778,8 @@ async def _get_team_deployments(
|
|||
# Confirm team_id in model_info (defensive check)
|
||||
result = []
|
||||
for row in response:
|
||||
model_info = row.model_info
|
||||
if isinstance(model_info, str):
|
||||
try:
|
||||
model_info = json.loads(model_info)
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
if isinstance(model_info, dict) and model_info.get("team_id") == team_id:
|
||||
model_info = model_info_as_mapping(row.model_info)
|
||||
if model_info is not None and model_info.get("team_id") == team_id:
|
||||
result.append(row)
|
||||
return result
|
||||
|
||||
|
|
@ -770,13 +830,8 @@ async def _get_team_public_model_names(
|
|||
deployments = await _get_team_deployments(team_id, prisma_client)
|
||||
public_names: Set[str] = set()
|
||||
for row in deployments:
|
||||
model_info = row.model_info
|
||||
if isinstance(model_info, str):
|
||||
try:
|
||||
model_info = json.loads(model_info)
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
if isinstance(model_info, dict):
|
||||
model_info = model_info_as_mapping(row.model_info)
|
||||
if model_info is not None:
|
||||
public_name = model_info.get("team_public_model_name")
|
||||
if public_name:
|
||||
public_names.add(public_name)
|
||||
|
|
@ -788,6 +843,7 @@ async def _remove_unbacked_team_models(
|
|||
prisma_client: PrismaClient,
|
||||
user_api_key_cache: Any,
|
||||
proxy_logging_obj: Any,
|
||||
llm_router: Router | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Strip a deleted team model's public name(s) from team.models and refresh the cache.
|
||||
|
|
@ -795,26 +851,50 @@ async def _remove_unbacked_team_models(
|
|||
Must be called after the deployment row is deleted: a public name is removed only
|
||||
when no remaining team deployment still backs it, so a load-balanced replica isn't
|
||||
revoked while siblings serve it, and concurrent deletes can't leave a ghost.
|
||||
|
||||
Legacy team models (created before team_public_model_name existed) store a
|
||||
``{public_name: "model_name_{team_id}_{uuid}"}`` entry in the team's model_aliases,
|
||||
so the alias scan runs for every team model; skipping it for internal-shaped names
|
||||
left stale aliases that rewrote requests to deployments that no longer exist.
|
||||
Aliases are scrubbed only when the deleted deployment's name no longer resolves in
|
||||
the router, so deleting one replica of a load-balanced group never breaks aliases
|
||||
that still route to the surviving replicas (in any team).
|
||||
|
||||
A public name that still resolves to a live router deployment (e.g. a gateway-level
|
||||
model group shared with the team) is kept in team.models, so deleting a per-team
|
||||
duplicate does not revoke the team's access to the shared deployment.
|
||||
"""
|
||||
team_id = model_params.model_info.team_id
|
||||
if team_id is None:
|
||||
return
|
||||
|
||||
# BYOK models carry an internal `model_name_{team_id}_{uuid}` name that can never
|
||||
# be a team alias value, so skip the full litellm_modeltable scan for them.
|
||||
removed_model_aliases: List[Tuple[str, str]] = []
|
||||
if not model_params.model_name.startswith(f"model_name_{team_id}_"):
|
||||
removed_model_aliases = await delete_team_model_alias(
|
||||
deleted_name_still_served = (
|
||||
llm_router is not None and model_params.model_name in llm_router.model_name_to_deployment_indices
|
||||
)
|
||||
removed_model_aliases: List[Tuple[str, str]] = (
|
||||
[]
|
||||
if deleted_name_still_served
|
||||
else await delete_team_model_alias(
|
||||
public_model_name=model_params.model_name,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
names_to_remove = {alias for alias_team_id, alias in removed_model_aliases if alias_team_id == team_id}
|
||||
if model_params.model_info.team_public_model_name is not None:
|
||||
names_to_remove.add(model_params.model_info.team_public_model_name)
|
||||
|
||||
if names_to_remove:
|
||||
names_to_remove -= await _get_team_public_model_names(team_id=team_id, prisma_client=prisma_client)
|
||||
)
|
||||
removed_alias_names = {alias for alias_team_id, alias in removed_model_aliases if alias_team_id == team_id}
|
||||
candidate_names = (
|
||||
removed_alias_names | {model_params.model_info.team_public_model_name}
|
||||
if model_params.model_info.team_public_model_name is not None
|
||||
else removed_alias_names
|
||||
)
|
||||
if not candidate_names:
|
||||
return
|
||||
|
||||
team_backed_names = await _get_team_public_model_names(team_id=team_id, prisma_client=prisma_client)
|
||||
router_served_names = (
|
||||
frozenset(name for name in candidate_names if name in llm_router.model_name_to_deployment_indices)
|
||||
if llm_router is not None
|
||||
else frozenset()
|
||||
)
|
||||
names_to_remove = candidate_names - team_backed_names - router_served_names
|
||||
if not names_to_remove:
|
||||
return
|
||||
|
||||
|
|
@ -853,18 +933,11 @@ async def _update_existing_team_model_assignment(
|
|||
def _get_team_public_model_name(
|
||||
model_info: Optional[Union[dict, str]],
|
||||
) -> Optional[str]:
|
||||
if isinstance(model_info, dict):
|
||||
value = model_info.get("team_public_model_name")
|
||||
return value if isinstance(value, str) else None
|
||||
if isinstance(model_info, str):
|
||||
try:
|
||||
parsed = json.loads(model_info)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
if isinstance(parsed, dict):
|
||||
value = parsed.get("team_public_model_name")
|
||||
return value if isinstance(value, str) else None
|
||||
return None
|
||||
parsed = model_info_as_mapping(model_info)
|
||||
if parsed is None:
|
||||
return None
|
||||
value = parsed.get("team_public_model_name")
|
||||
return value if isinstance(value, str) else None
|
||||
|
||||
old_public_name = db_model.model_info.team_public_model_name if db_model.model_info else None
|
||||
|
||||
|
|
@ -978,7 +1051,7 @@ class ModelManagementAuthChecks:
|
|||
status_code=400,
|
||||
detail={"error": "Team id={} does not exist in db".format(model_params.model_info.team_id)},
|
||||
)
|
||||
existing_team_row = LiteLLM_TeamTable(**_existing_team_row.model_dump())
|
||||
existing_team_row = LiteLLM_TeamTable.model_validate(_existing_team_row.model_dump())
|
||||
|
||||
ModelManagementAuthChecks.can_user_make_team_model_call(
|
||||
team_id=model_params.model_info.team_id,
|
||||
|
|
@ -1016,7 +1089,7 @@ class ModelManagementAuthChecks:
|
|||
status_code=400,
|
||||
detail={"error": "Team id={} does not exist in db".format(model_params.model_info.team_id)},
|
||||
)
|
||||
team_obj = LiteLLM_TeamTable(**team_obj_row.model_dump())
|
||||
team_obj = LiteLLM_TeamTable.model_validate(team_obj_row.model_dump())
|
||||
|
||||
return ModelManagementAuthChecks.can_user_make_team_model_call(
|
||||
team_id=model_params.model_info.team_id,
|
||||
|
|
@ -1120,6 +1193,7 @@ async def delete_model(
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
## CREATE AUDIT LOG ##
|
||||
|
|
@ -1267,6 +1341,11 @@ async def add_new_model(
|
|||
premium_user=premium_user,
|
||||
)
|
||||
|
||||
_raise_on_strategy_router_write_violation(
|
||||
incoming_params=model_params.litellm_params,
|
||||
existing_params=None,
|
||||
)
|
||||
|
||||
model_response: Optional[LiteLLM_ProxyModelTable] = None
|
||||
# update DB
|
||||
if store_model_in_db is True:
|
||||
|
|
@ -1275,6 +1354,7 @@ async def add_new_model(
|
|||
- store keys separately
|
||||
"""
|
||||
|
||||
live_before_reload = live_model_ids_snapshot()
|
||||
try:
|
||||
_original_litellm_model_name = model_params.model_name
|
||||
if model_params.model_info.team_id is None:
|
||||
|
|
@ -1330,6 +1410,12 @@ async def add_new_model(
|
|||
)
|
||||
)
|
||||
|
||||
raise_if_reload_degraded_serving(
|
||||
before=live_before_reload,
|
||||
written_models=[(model_response.model_id, getattr(model_response, "model_info", None))],
|
||||
action="create",
|
||||
)
|
||||
|
||||
return model_response
|
||||
|
||||
except Exception as e:
|
||||
|
|
@ -1414,6 +1500,11 @@ async def update_model(
|
|||
premium_user=premium_user,
|
||||
)
|
||||
|
||||
_raise_on_strategy_router_write_violation(
|
||||
incoming_params=model_params.litellm_params,
|
||||
existing_params=deployment.litellm_params,
|
||||
)
|
||||
|
||||
# update DB
|
||||
if store_model_in_db is True:
|
||||
_existing_litellm_params_dict = dict(_existing_litellm_params.litellm_params)
|
||||
|
|
@ -1450,8 +1541,8 @@ async def update_model(
|
|||
)
|
||||
|
||||
# Clear cache and reload models (uses config setting or defaults to preserving config models for DB updates)
|
||||
live_before_reload = live_model_ids_snapshot()
|
||||
await clear_cache()
|
||||
|
||||
## CREATE AUDIT LOG ##
|
||||
asyncio.create_task(
|
||||
create_object_audit_log(
|
||||
|
|
@ -1474,6 +1565,12 @@ async def update_model(
|
|||
)
|
||||
)
|
||||
|
||||
raise_if_reload_degraded_serving(
|
||||
before=live_before_reload,
|
||||
written_models=[(_model_id, getattr(model_response, "model_info", None))],
|
||||
action="update",
|
||||
)
|
||||
|
||||
return model_response
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
|
|
@ -1677,6 +1774,114 @@ def _deduplicate_litellm_router_models(models: List[Dict]) -> List[Dict]:
|
|||
return unique_models
|
||||
|
||||
|
||||
def model_info_as_mapping(model_info: object) -> Mapping[str, object] | None:
|
||||
"""A DB row's model_info column arrives as a dict or as its JSON string depending on
|
||||
the query path, and every consumer needs the mapping. Single owner of that parse:
|
||||
returns None when no usable mapping exists (None, an unparseable string, or JSON
|
||||
that is not an object), and callers choose what None means for them."""
|
||||
if isinstance(model_info, Mapping):
|
||||
return model_info
|
||||
if not isinstance(model_info, str):
|
||||
return None
|
||||
try:
|
||||
parsed = json.loads(model_info)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
return parsed if isinstance(parsed, Mapping) else None
|
||||
|
||||
|
||||
def _expects_liveness_on_this_pod(model_info: object) -> bool:
|
||||
from litellm.router import model_info_is_active_for_environment
|
||||
|
||||
try:
|
||||
return model_info_is_active_for_environment(model_info=model_info_as_mapping(model_info))
|
||||
except ValueError:
|
||||
return True
|
||||
|
||||
|
||||
def live_model_ids_snapshot() -> frozenset[str]:
|
||||
"""The ids this pod's router is currently serving, read fresh from the module global
|
||||
because a reload can rebind it. The empirical ground truth every verdict below is
|
||||
computed from; an absent router serves nothing."""
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
|
||||
if llm_router is None:
|
||||
return frozenset()
|
||||
return frozenset(llm_router.get_model_ids())
|
||||
|
||||
|
||||
def reload_serving_verdict(
|
||||
before: frozenset[str],
|
||||
written_models: Sequence[tuple[str, object]],
|
||||
written_must_serve: bool,
|
||||
) -> tuple[tuple[str, ...], tuple[str, ...]]:
|
||||
"""Judge a write-triggered reload by diffing the router's serving state instead of
|
||||
trusting any layer of the reload stack to report its own failure.
|
||||
|
||||
The full cell matrix, per id:
|
||||
- written, must-serve (the write's purpose is this model's serving state): live now
|
||||
is fine; not live is reported unless the row is deliberately inactive for this
|
||||
pod's LITELLM_ENVIRONMENT; a row whose model_info cannot be read counts as
|
||||
expecting to serve, so its drop is still reported
|
||||
- written, metadata-only (must_not_degrade): live before and gone now is reported;
|
||||
a row that was already not serving stays silent, because its deadness predates
|
||||
this write and blaming it would block unrelated metadata fixes
|
||||
- not written but live before and gone now: collateral degradation of this pod
|
||||
caused by the reload this request triggered (a wholesale re-add failure, or a
|
||||
newly introduced conflict), always reported
|
||||
|
||||
Returns (written ids violating their obligation, collateral ids no longer served).
|
||||
Best effort under concurrent admin writes: the snapshot spans only this request.
|
||||
"""
|
||||
now = live_model_ids_snapshot()
|
||||
written_ids = frozenset(model_id for model_id, _ in written_models)
|
||||
if written_must_serve:
|
||||
missing = tuple(
|
||||
model_id
|
||||
for model_id, model_info in written_models
|
||||
if model_id not in now and _expects_liveness_on_this_pod(model_info)
|
||||
)
|
||||
else:
|
||||
missing = tuple(model_id for model_id, _ in written_models if model_id in before and model_id not in now)
|
||||
collateral = tuple(sorted(before - now - written_ids))
|
||||
return (missing, collateral)
|
||||
|
||||
|
||||
def raise_if_reload_degraded_serving(
|
||||
before: frozenset[str],
|
||||
written_models: Sequence[tuple[str, object]],
|
||||
action: str,
|
||||
) -> None:
|
||||
"""The caller-visible error this pod's model-write endpoints owe their caller when
|
||||
the model they wrote is not being served after the reload they triggered. The DB
|
||||
write is durable either way and every other pod reloads on its own interval; this
|
||||
speaks only for the handling pod."""
|
||||
missing, collateral = reload_serving_verdict(before=before, written_models=written_models, written_must_serve=True)
|
||||
if not missing and not collateral:
|
||||
return
|
||||
missing_clause = (
|
||||
f"the model id(s) {list(missing)} are not live in this pod's router after the reload and are not "
|
||||
"being served by this pod."
|
||||
if missing
|
||||
else "the reload it triggered degraded this pod's serving state."
|
||||
)
|
||||
collateral_clause = (
|
||||
f" Previously served model id(s) {list(collateral)} are also no longer being served by this pod."
|
||||
if collateral
|
||||
else ""
|
||||
)
|
||||
raise ProxyException(
|
||||
message=(
|
||||
f"Model {action} was saved to the database, but {missing_clause}{collateral_clause} "
|
||||
"Other pods reload on their own interval. Check server logs for 'Error upserting deployment' or "
|
||||
"'Error creating deployment' for the cause."
|
||||
),
|
||||
type=ProxyErrorTypes.internal_server_error,
|
||||
code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
param=None,
|
||||
)
|
||||
|
||||
|
||||
async def clear_cache():
|
||||
"""
|
||||
Clear router caches and reload models.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -141,7 +141,7 @@ from litellm.types.proxy.management_endpoints.team_endpoints import (
|
|||
router = APIRouter()
|
||||
|
||||
|
||||
def _sanitize_for_log(value: Any) -> str:
|
||||
def _sanitize_for_log(value: object) -> str:
|
||||
"""Strip CR/LF from user-controlled values to prevent log injection."""
|
||||
try:
|
||||
text = str(value)
|
||||
|
|
@ -171,7 +171,7 @@ async def _refresh_cached_team(
|
|||
"""
|
||||
await _cache_team_object(
|
||||
team_id=team_row.team_id,
|
||||
team_table=LiteLLM_TeamTableCachedObj(**team_row.model_dump()),
|
||||
team_table=LiteLLM_TeamTableCachedObj.model_validate(team_row.model_dump()),
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
|
@ -510,7 +510,7 @@ async def get_all_team_memberships(
|
|||
|
||||
returned_tm: List[LiteLLM_TeamMembership] = []
|
||||
for tm in team_memberships:
|
||||
returned_tm.append(LiteLLM_TeamMembership(**tm.model_dump()))
|
||||
returned_tm.append(LiteLLM_TeamMembership.model_validate(tm.model_dump()))
|
||||
|
||||
return returned_tm
|
||||
|
||||
|
|
@ -772,7 +772,7 @@ async def _check_org_team_limits(
|
|||
# Convert teams to LiteLLM_TeamTable objects
|
||||
team_objs: List[LiteLLM_TeamTable] = []
|
||||
for team in teams:
|
||||
team_objs.append(LiteLLM_TeamTable(**team.model_dump()))
|
||||
team_objs.append(LiteLLM_TeamTable.model_validate(team.model_dump()))
|
||||
|
||||
check_org_team_model_specific_limits(
|
||||
teams=team_objs,
|
||||
|
|
@ -1467,9 +1467,9 @@ async def fetch_and_validate_organization(
|
|||
)
|
||||
|
||||
is_proxy_admin = user_api_key_dict is not None and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
organization = LiteLLM_OrganizationTableWithMembers(**organization_row.model_dump())
|
||||
organization = LiteLLM_OrganizationTableWithMembers.model_validate(organization_row.model_dump())
|
||||
validate_team_org_change(
|
||||
team=LiteLLM_TeamTable(**existing_team_row.model_dump()),
|
||||
team=LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()),
|
||||
organization=organization,
|
||||
llm_router=llm_router,
|
||||
is_proxy_admin=is_proxy_admin,
|
||||
|
|
@ -1477,7 +1477,7 @@ async def fetch_and_validate_organization(
|
|||
|
||||
if is_proxy_admin:
|
||||
await _auto_add_team_members_to_organization(
|
||||
team=LiteLLM_TeamTable(**existing_team_row.model_dump()),
|
||||
team=LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()),
|
||||
organization=organization,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
|
@ -1714,7 +1714,7 @@ async def update_team(
|
|||
|
||||
# Verify caller has access to manage this team
|
||||
await _verify_team_access(
|
||||
team_obj=LiteLLM_TeamTable(**existing_team_row.model_dump()),
|
||||
team_obj=LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
|
|
@ -2013,7 +2013,7 @@ async def patch_team(
|
|||
existing_metadata = existing_team_row.metadata if isinstance(existing_team_row.metadata, dict) else {}
|
||||
patch_fields["metadata"] = apply_json_merge_patch(existing_metadata, patch_fields["metadata"])
|
||||
|
||||
update_request = UpdateTeamRequest(team_id=team_id, **patch_fields)
|
||||
update_request = UpdateTeamRequest.model_validate({"team_id": team_id, **patch_fields})
|
||||
|
||||
result = await update_team(
|
||||
data=update_request,
|
||||
|
|
@ -2591,7 +2591,7 @@ async def team_member_add(
|
|||
detail={"error": f"Team not found for team_id={getattr(data, 'team_id', None)}"},
|
||||
)
|
||||
|
||||
complete_team_data = LiteLLM_TeamTable(**existing_team_row.model_dump())
|
||||
complete_team_data = LiteLLM_TeamTable.model_validate(existing_team_row.model_dump())
|
||||
|
||||
team_member_add_duplication_check(
|
||||
data=data,
|
||||
|
|
@ -2636,10 +2636,12 @@ async def team_member_add(
|
|||
|
||||
_emit_team_members_metric(complete_team_data)
|
||||
|
||||
return TeamAddMemberResponse(
|
||||
**updated_team.model_dump(),
|
||||
updated_users=updated_users,
|
||||
updated_team_memberships=updated_team_memberships,
|
||||
return TeamAddMemberResponse.model_validate(
|
||||
{
|
||||
**updated_team.model_dump(),
|
||||
"updated_users": updated_users,
|
||||
"updated_team_memberships": updated_team_memberships,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -2711,7 +2713,7 @@ async def team_member_delete(
|
|||
status_code=400,
|
||||
detail={"error": "Team id={} does not exist in db".format(data.team_id)},
|
||||
)
|
||||
existing_team_row = LiteLLM_TeamTable(**_existing_team_row.model_dump())
|
||||
existing_team_row = LiteLLM_TeamTable.model_validate(_existing_team_row.model_dump())
|
||||
|
||||
## CHECK IF USER IS PROXY ADMIN OR TEAM ADMIN OR ORG ADMIN
|
||||
|
||||
|
|
@ -2915,7 +2917,7 @@ async def team_member_update(
|
|||
status_code=400,
|
||||
detail={"error": "Team id={} does not exist in db".format(data.team_id)},
|
||||
)
|
||||
existing_team_row = LiteLLM_TeamTable(**_existing_team_row.model_dump())
|
||||
existing_team_row = LiteLLM_TeamTable.model_validate(_existing_team_row.model_dump())
|
||||
|
||||
## CHECK IF USER IS PROXY ADMIN OR TEAM ADMIN OR ORG ADMIN
|
||||
|
||||
|
|
@ -3261,7 +3263,7 @@ async def delete_team(
|
|||
status_code=404,
|
||||
detail={"error": f"Team not found, passed team_id={team_id}"},
|
||||
)
|
||||
team_row_pydantic = LiteLLM_TeamTable(**team_row_base.model_dump())
|
||||
team_row_pydantic = LiteLLM_TeamTable.model_validate(team_row_base.model_dump())
|
||||
|
||||
# Verify caller has access to manage this team
|
||||
await _verify_team_access(
|
||||
|
|
@ -3385,12 +3387,14 @@ def _transform_teams_to_deleted_records(
|
|||
records = []
|
||||
for team in teams:
|
||||
team_payload = team.model_dump()
|
||||
deleted_record = LiteLLM_DeletedTeamTable(
|
||||
**team_payload,
|
||||
deleted_at=deleted_at,
|
||||
deleted_by=user_api_key_dict.user_id,
|
||||
deleted_by_api_key=user_api_key_dict.api_key,
|
||||
litellm_changed_by=litellm_changed_by,
|
||||
deleted_record = LiteLLM_DeletedTeamTable.model_validate(
|
||||
{
|
||||
**team_payload,
|
||||
"deleted_at": deleted_at,
|
||||
"deleted_by": user_api_key_dict.user_id,
|
||||
"deleted_by_api_key": user_api_key_dict.api_key,
|
||||
"litellm_changed_by": litellm_changed_by,
|
||||
}
|
||||
)
|
||||
record = deleted_record.model_dump()
|
||||
|
||||
|
|
@ -3580,7 +3584,7 @@ async def team_info(
|
|||
)
|
||||
await validate_membership(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
team_table=LiteLLM_TeamTable(**team_info.model_dump()),
|
||||
team_table=LiteLLM_TeamTable.model_validate(team_info.model_dump()),
|
||||
)
|
||||
|
||||
## GET ALL KEYS ##
|
||||
|
|
@ -3615,9 +3619,9 @@ async def team_info(
|
|||
returned_tm = await get_all_team_memberships(prisma_client, [team_id], user_id=None)
|
||||
|
||||
if isinstance(team_info, dict):
|
||||
_team_info = TeamInfoResponseObjectTeamTable(**team_info)
|
||||
_team_info = TeamInfoResponseObjectTeamTable.model_validate(team_info)
|
||||
elif isinstance(team_info, BaseModel):
|
||||
_team_info = TeamInfoResponseObjectTeamTable(**team_info.model_dump())
|
||||
_team_info = TeamInfoResponseObjectTeamTable.model_validate(team_info.model_dump())
|
||||
else:
|
||||
_team_info = TeamInfoResponseObjectTeamTable()
|
||||
|
||||
|
|
@ -3823,7 +3827,7 @@ async def block_team(
|
|||
|
||||
# Verify caller has access to manage this team
|
||||
await _verify_team_access(
|
||||
team_obj=LiteLLM_TeamTable(**existing_team.model_dump()),
|
||||
team_obj=LiteLLM_TeamTable.model_validate(existing_team.model_dump()),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
|
|
@ -3872,7 +3876,7 @@ async def unblock_team(
|
|||
|
||||
# Verify caller has access to manage this team
|
||||
await _verify_team_access(
|
||||
team_obj=LiteLLM_TeamTable(**existing_team.model_dump()),
|
||||
team_obj=LiteLLM_TeamTable.model_validate(existing_team.model_dump()),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
|
|
@ -3916,13 +3920,13 @@ async def list_available_teams(
|
|||
status_code=404,
|
||||
detail={"error": "User not found"},
|
||||
)
|
||||
user_info_correct_type = LiteLLM_UserTable(**user_info.model_dump())
|
||||
user_info_correct_type = LiteLLM_UserTable.model_validate(user_info.model_dump())
|
||||
|
||||
available_teams = [team for team in available_teams if team not in user_info_correct_type.teams]
|
||||
|
||||
available_teams_db = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": available_teams}})
|
||||
|
||||
available_teams_correct_type = [LiteLLM_TeamTable(**team.model_dump()) for team in available_teams_db]
|
||||
available_teams_correct_type = [LiteLLM_TeamTable.model_validate(team.model_dump()) for team in available_teams_db]
|
||||
|
||||
return available_teams_correct_type
|
||||
|
||||
|
|
@ -4090,7 +4094,7 @@ def _convert_teams_to_response_models(
|
|||
team_dict = team.dict()
|
||||
|
||||
if use_deleted_table:
|
||||
team_list.append(LiteLLM_DeletedTeamTable(**team_dict))
|
||||
team_list.append(LiteLLM_DeletedTeamTable.model_validate(team_dict))
|
||||
else:
|
||||
members_with_roles = team_dict.get("members_with_roles")
|
||||
if not isinstance(members_with_roles, list):
|
||||
|
|
@ -4705,7 +4709,7 @@ async def team_model_add(
|
|||
detail={"error": f"Team not found, passed team_id={data.team_id}"},
|
||||
)
|
||||
|
||||
team_obj = LiteLLM_TeamTable(**team_row.model_dump())
|
||||
team_obj = LiteLLM_TeamTable.model_validate(team_row.model_dump())
|
||||
|
||||
# Authorization check - only proxy admin, team admin, or org admin can add models
|
||||
if (
|
||||
|
|
@ -4805,7 +4809,7 @@ async def team_model_delete(
|
|||
detail={"error": f"Team not found, passed team_id={data.team_id}"},
|
||||
)
|
||||
|
||||
team_obj = LiteLLM_TeamTable(**team_row.model_dump())
|
||||
team_obj = LiteLLM_TeamTable.model_validate(team_row.model_dump())
|
||||
|
||||
# Authorization check - only proxy admin, team admin, or org admin can remove models
|
||||
if (
|
||||
|
|
@ -4873,7 +4877,7 @@ async def team_member_permissions(
|
|||
check_db_only=True,
|
||||
)
|
||||
|
||||
complete_team_data = LiteLLM_TeamTable(**existing_team_row.model_dump())
|
||||
complete_team_data = LiteLLM_TeamTable.model_validate(existing_team_row.model_dump())
|
||||
|
||||
# Admin Viewer follows the read-parity rule: see team permissions like
|
||||
# a Proxy Admin would. Team / org admins keep their existing scope.
|
||||
|
|
@ -4940,7 +4944,7 @@ async def update_team_member_permissions(
|
|||
check_db_only=True,
|
||||
)
|
||||
|
||||
complete_team_data = LiteLLM_TeamTable(**existing_team_row.model_dump())
|
||||
complete_team_data = LiteLLM_TeamTable.model_validate(existing_team_row.model_dump())
|
||||
|
||||
# Available-team self-join must NOT grant write access to team-wide
|
||||
# permission policies; only proxy/team/org admins can update them.
|
||||
|
|
@ -5201,7 +5205,7 @@ async def get_team_daily_activity(
|
|||
if not _user_has_admin_view(user_api_key_dict) and team_ids_list and team_aliases:
|
||||
has_full_team_view = True
|
||||
for team_alias in team_aliases:
|
||||
team_obj = LiteLLM_TeamTable(**team_alias.model_dump())
|
||||
team_obj = LiteLLM_TeamTable.model_validate(team_alias.model_dump())
|
||||
is_admin = _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj)
|
||||
has_perm = _team_member_has_permission(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -6758,6 +6758,9 @@ class ProxyConfig:
|
|||
from litellm.proxy._experimental.mcp_server.oauth2_flow_backfill import (
|
||||
backfill_null_oauth2_flows,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.oauth_issuer_stamp_backfill import (
|
||||
backfill_discovery_stamped_issuers,
|
||||
)
|
||||
|
||||
try:
|
||||
if prisma_client is not None:
|
||||
|
|
@ -6767,6 +6770,16 @@ class ProxyConfig:
|
|||
"litellm.proxy.proxy_server.py::ProxyConfig:_init_mcp_servers_in_db backfill - {}".format(str(e))
|
||||
)
|
||||
|
||||
try:
|
||||
if prisma_client is not None:
|
||||
await backfill_discovery_stamped_issuers(prisma_client)
|
||||
except Exception as e: # noqa: BLE001
|
||||
verbose_proxy_logger.exception(
|
||||
"litellm.proxy.proxy_server.py::ProxyConfig:_init_mcp_servers_in_db issuer stamp backfill - {}".format(
|
||||
str(e)
|
||||
)
|
||||
)
|
||||
|
||||
try:
|
||||
await global_mcp_server_manager.reload_servers_from_database()
|
||||
except Exception as e:
|
||||
|
|
@ -6778,6 +6791,31 @@ class ProxyConfig:
|
|||
if self._should_load_db_object(object_type="mcp"):
|
||||
await self._init_mcp_servers_in_db()
|
||||
|
||||
async def reload_mcp_servers_from_db(self) -> None:
|
||||
"""Registry refresh only, for the periodic job in store_model_in_db-off deployments.
|
||||
|
||||
Deliberately narrower than ``init_mcp_servers_from_db``: the oauth2_flow backfill is a write
|
||||
path that only needs to run once at startup, so the cadence here is purely the read-side
|
||||
reload whose fast-path exemption retries failed OAuth discovery. Gated the same way, so an
|
||||
admin who excluded mcp from supported_db_objects opts out of this too.
|
||||
"""
|
||||
if not self._should_load_db_object(object_type="mcp"):
|
||||
return
|
||||
from litellm.proxy._experimental.mcp_server.utils import is_mcp_available
|
||||
|
||||
if not is_mcp_available():
|
||||
return
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
try:
|
||||
await global_mcp_server_manager.reload_servers_from_database()
|
||||
except Exception as e: # noqa: BLE001 # scheduled job: a reload failure must not kill the recurring retry
|
||||
verbose_proxy_logger.exception(
|
||||
"litellm.proxy.proxy_server.py::ProxyConfig:reload_mcp_servers_from_db - {}".format(str(e))
|
||||
)
|
||||
|
||||
async def _init_agents_in_db(self, prisma_client: PrismaClient):
|
||||
from litellm.proxy.agent_endpoints.agent_registry import (
|
||||
global_agent_registry as AGENT_REGISTRY,
|
||||
|
|
@ -8099,6 +8137,22 @@ class ProxyStartupEvent:
|
|||
|
||||
if store_model_in_db is not True:
|
||||
await proxy_config.init_mcp_servers_from_db()
|
||||
if prisma_client is not None:
|
||||
# DB-backed MCP servers are live objects in every mode, so the registry refresh that
|
||||
# store_model_in_db=True deployments get via the add_deployment job must run here
|
||||
# too; without it, a server whose OAuth discovery failed at startup is rebuilt only
|
||||
# by a management write, since the reload fast path is the retry's only driver.
|
||||
mcp_reload_interval_seconds = proxy_config_reload_interval_seconds
|
||||
if not isinstance(mcp_reload_interval_seconds, int) or mcp_reload_interval_seconds <= 0:
|
||||
mcp_reload_interval_seconds = 30
|
||||
scheduler.add_job(
|
||||
proxy_config.reload_mcp_servers_from_db,
|
||||
"interval",
|
||||
seconds=mcp_reload_interval_seconds,
|
||||
id="reload_mcp_servers_job",
|
||||
replace_existing=True,
|
||||
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
|
||||
)
|
||||
|
||||
await cls._initialize_slack_alerting_jobs(
|
||||
scheduler=scheduler,
|
||||
|
|
@ -11214,11 +11268,15 @@ async def get_all_team_models(
|
|||
|
||||
if user_teams == "*":
|
||||
team_db_objects = await TeamRepository(prisma_client).table.find_many()
|
||||
team_db_objects_typed = [LiteLLM_TeamTable(**team_db_object.model_dump()) for team_db_object in team_db_objects]
|
||||
team_db_objects_typed = [
|
||||
LiteLLM_TeamTable.model_validate(team_db_object.model_dump()) for team_db_object in team_db_objects
|
||||
]
|
||||
else:
|
||||
team_db_objects = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": user_teams}})
|
||||
|
||||
team_db_objects_typed = [LiteLLM_TeamTable(**team_db_object.model_dump()) for team_db_object in team_db_objects]
|
||||
team_db_objects_typed = [
|
||||
LiteLLM_TeamTable.model_validate(team_db_object.model_dump()) for team_db_object in team_db_objects
|
||||
]
|
||||
|
||||
team_models = _add_team_models_to_all_models(
|
||||
team_db_objects_typed=team_db_objects_typed,
|
||||
|
|
@ -11292,7 +11350,7 @@ async def _populate_team_access_on_models(
|
|||
where={"user_id": user_api_key_dict.user_id}
|
||||
)
|
||||
if user_db_object is not None:
|
||||
user_object = LiteLLM_UserTable(**user_db_object.model_dump())
|
||||
user_object = LiteLLM_UserTable.model_validate(user_db_object.model_dump())
|
||||
user_teams = user_object.teams or []
|
||||
direct_access_models = get_direct_access_models(
|
||||
user_db_object=user_object,
|
||||
|
|
@ -11827,7 +11885,7 @@ async def _load_team_object_for_model_filter(team_id: str, prisma_client: Prisma
|
|||
if team_db_object is None:
|
||||
verbose_proxy_logger.warning(f"Team {team_id} not found in database")
|
||||
return None
|
||||
return LiteLLM_TeamTable(**team_db_object.model_dump())
|
||||
return LiteLLM_TeamTable.model_validate(team_db_object.model_dump())
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error fetching team {team_id}: {str(e)}")
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ import threading
|
|||
import time
|
||||
import traceback
|
||||
from collections import defaultdict
|
||||
from collections.abc import Mapping
|
||||
from functools import lru_cache
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
|
|
@ -109,6 +110,9 @@ from litellm.router_utils.batch_utils import (
|
|||
replace_model_in_jsonl,
|
||||
should_replace_model_in_jsonl,
|
||||
)
|
||||
from litellm.router_utils.auto_router_model_naming import (
|
||||
classify_strategy_router_model,
|
||||
)
|
||||
from litellm.router_utils.client_initalization_utils import InitalizeCachedClient
|
||||
from litellm.router_utils.clientside_credential_handler import (
|
||||
get_dynamic_litellm_params,
|
||||
|
|
@ -270,6 +274,43 @@ def _cost_value_as_float(value: Union[str, int, float, None]) -> Optional[float]
|
|||
return None
|
||||
|
||||
|
||||
def model_info_is_active_for_environment(model_info: Mapping[str, object] | None) -> bool:
|
||||
"""Single owner of the environment-gating rule: a deployment whose model_info names
|
||||
`supported_environments` loads only on pods whose LITELLM_ENVIRONMENT is in that list.
|
||||
`Router.deployment_is_active_for_environment` delegates here, and the model-write
|
||||
endpoints consult the same rule to tell a deliberately inactive model from one that
|
||||
was dropped by a failed reload."""
|
||||
if model_info is None:
|
||||
return True
|
||||
supported_environments = model_info.get("supported_environments")
|
||||
if supported_environments is None:
|
||||
return True
|
||||
if not isinstance(supported_environments, (list, tuple)):
|
||||
raise ValueError(
|
||||
f"supported_environments must be a list of {VALID_LITELLM_ENVIRONMENTS}. "
|
||||
f"but set as: {supported_environments} for model_info: {model_info}"
|
||||
)
|
||||
litellm_environment = get_secret_str(secret_name="LITELLM_ENVIRONMENT")
|
||||
if litellm_environment is None:
|
||||
raise ValueError("Set 'supported_environments' for model but not 'LITELLM_ENVIRONMENT' set in .env")
|
||||
|
||||
if litellm_environment not in VALID_LITELLM_ENVIRONMENTS:
|
||||
raise ValueError(
|
||||
f"LITELLM_ENVIRONMENT must be one of {VALID_LITELLM_ENVIRONMENTS}. but set as: {litellm_environment}"
|
||||
)
|
||||
|
||||
for _env in supported_environments:
|
||||
if _env not in VALID_LITELLM_ENVIRONMENTS:
|
||||
raise ValueError(
|
||||
f"supported_environments must be one of {VALID_LITELLM_ENVIRONMENTS}. but set as: {_env} "
|
||||
f"for model_info: {model_info}"
|
||||
)
|
||||
|
||||
if litellm_environment in supported_environments:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
_PreRoutingStrategyT = TypeVar("_PreRoutingStrategyT")
|
||||
|
||||
|
||||
|
|
@ -7585,15 +7626,7 @@ class Router:
|
|||
but NOT "auto_router/complexity_router" or "auto_router/adaptive_router"
|
||||
(which use the complexity-router and adaptive-router strategies).
|
||||
"""
|
||||
if litellm_params.model.startswith("auto_router/complexity_router"):
|
||||
return False # This is handled by complexity_router
|
||||
if litellm_params.model.startswith("auto_router/adaptive_router"):
|
||||
return False # This is handled by adaptive_router
|
||||
if litellm_params.model.startswith("auto_router/quality_router"):
|
||||
return False # This is handled by quality_router
|
||||
if litellm_params.model.startswith("auto_router/"):
|
||||
return True
|
||||
return False
|
||||
return classify_strategy_router_model(litellm_params.model) == "semantic"
|
||||
|
||||
@staticmethod
|
||||
def _deployment_tags(deployment: Deployment) -> tuple[str, ...]:
|
||||
|
|
@ -7648,9 +7681,7 @@ class Router:
|
|||
|
||||
Returns True if the litellm_params model starts with "auto_router/complexity_router"
|
||||
"""
|
||||
if litellm_params.model.startswith("auto_router/complexity_router"):
|
||||
return True
|
||||
return False
|
||||
return classify_strategy_router_model(litellm_params.model) == "complexity"
|
||||
|
||||
def init_complexity_router_deployment(self, deployment: Deployment):
|
||||
"""
|
||||
|
|
@ -7700,7 +7731,7 @@ class Router:
|
|||
|
||||
def _is_adaptive_router_deployment(self, litellm_params: LiteLLM_Params) -> bool:
|
||||
"""True when this deployment opts in via the `auto_router/adaptive_router` model prefix."""
|
||||
return litellm_params.model.startswith("auto_router/adaptive_router")
|
||||
return classify_strategy_router_model(litellm_params.model) == "adaptive"
|
||||
|
||||
def _deployment_participates_in_adaptive_routing(self, litellm_params: LiteLLM_Params) -> bool:
|
||||
"""True when this deployment owns an `adaptive_routers` entry once finalized:
|
||||
|
|
@ -7926,9 +7957,7 @@ class Router:
|
|||
|
||||
Returns True if the litellm_params model starts with "auto_router/quality_router".
|
||||
"""
|
||||
if litellm_params.model.startswith("auto_router/quality_router"):
|
||||
return True
|
||||
return False
|
||||
return classify_strategy_router_model(litellm_params.model) == "quality"
|
||||
|
||||
def init_quality_router_deployment(self, deployment: Deployment):
|
||||
"""
|
||||
|
|
@ -7982,30 +8011,7 @@ class Router:
|
|||
- ValueError: If LITELLM_ENVIRONMENT is not set in .env or not one of the valid values
|
||||
- ValueError: If supported_environments is not set in model_info or not one of the valid values
|
||||
"""
|
||||
if (
|
||||
deployment.model_info is None
|
||||
or "supported_environments" not in deployment.model_info
|
||||
or deployment.model_info["supported_environments"] is None
|
||||
):
|
||||
return True
|
||||
litellm_environment = get_secret_str(secret_name="LITELLM_ENVIRONMENT")
|
||||
if litellm_environment is None:
|
||||
raise ValueError("Set 'supported_environments' for model but not 'LITELLM_ENVIRONMENT' set in .env")
|
||||
|
||||
if litellm_environment not in VALID_LITELLM_ENVIRONMENTS:
|
||||
raise ValueError(
|
||||
f"LITELLM_ENVIRONMENT must be one of {VALID_LITELLM_ENVIRONMENTS}. but set as: {litellm_environment}"
|
||||
)
|
||||
|
||||
for _env in deployment.model_info["supported_environments"]:
|
||||
if _env not in VALID_LITELLM_ENVIRONMENTS:
|
||||
raise ValueError(
|
||||
f"supported_environments must be one of {VALID_LITELLM_ENVIRONMENTS}. but set as: {_env} for deployment: {deployment}"
|
||||
)
|
||||
|
||||
if litellm_environment in deployment.model_info["supported_environments"]:
|
||||
return True
|
||||
return False
|
||||
return model_info_is_active_for_environment(model_info=deployment.model_info)
|
||||
|
||||
def set_model_list(self, model_list: list):
|
||||
original_model_list = copy.deepcopy(model_list)
|
||||
|
|
@ -8630,6 +8636,33 @@ class Router:
|
|||
raise Exception("Model Name invalid - {}".format(type(model)))
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _deployment_usable_by_team(model: Union[Mapping, Deployment], team_id: str | None) -> bool:
|
||||
"""
|
||||
A team-scoped deployment (``model_info.team_id`` set) is only usable by
|
||||
callers from that same team; deployments without a team owner are shared.
|
||||
"""
|
||||
model_info = model.get("model_info") if isinstance(model, dict) else model.model_info
|
||||
owner_team_id = model_info.get("team_id") if model_info is not None else None
|
||||
return owner_team_id is None or owner_team_id == team_id
|
||||
|
||||
def _get_model_group_deployment_usable_by_team(
|
||||
self, model_group_name: str, team_id: str | None
|
||||
) -> Deployment | None:
|
||||
"""
|
||||
Like ``get_deployment_by_model_group_name``, but skips deployments owned
|
||||
by other teams so a shared model name never resolves another team's
|
||||
credentials.
|
||||
"""
|
||||
indices = self.model_name_to_deployment_indices.get(model_group_name) or ()
|
||||
usable = (
|
||||
self.model_list[idx] for idx in indices if self._deployment_usable_by_team(self.model_list[idx], team_id)
|
||||
)
|
||||
first_usable = next(usable, None)
|
||||
if first_usable is None:
|
||||
return None
|
||||
return Deployment(**first_usable) if isinstance(first_usable, dict) else first_usable
|
||||
|
||||
def get_configured_token_limits(self, model_name: str) -> "tuple[int | None, int | None]":
|
||||
"""
|
||||
Return (max_input_tokens, max_output_tokens) explicitly configured in a concrete
|
||||
|
|
@ -8664,7 +8697,10 @@ class Router:
|
|||
model_id: Model ID or model name from model_list (e.g., "gpt-4o-litellm")
|
||||
team_id: Optional team id of the caller. When set, team-scoped
|
||||
deployments (indexed by team public model name, including team
|
||||
wildcard models like "openai/*") are also considered.
|
||||
wildcard models like "openai/*") are also considered. Name and
|
||||
wildcard lookups never resolve a deployment owned by a
|
||||
different team, so shared model names can't leak another
|
||||
team's credentials.
|
||||
|
||||
Returns:
|
||||
Dictionary containing api_key, api_base, custom_llm_provider, etc.
|
||||
|
|
@ -8681,7 +8717,7 @@ class Router:
|
|||
|
||||
# If not found, try by model_group_name
|
||||
if deployment is None:
|
||||
deployment = self.get_deployment_by_model_group_name(model_group_name=model_id)
|
||||
deployment = self._get_model_group_deployment_usable_by_team(model_group_name=model_id, team_id=team_id)
|
||||
|
||||
# If not found, check team-scoped deployments whose team public model
|
||||
# name exactly matches model_id (wildcard team names are matched via
|
||||
|
|
@ -8698,7 +8734,12 @@ class Router:
|
|||
if deployment is None:
|
||||
team_pattern_router = self.team_pattern_routers.get(team_id) if team_id is not None else None
|
||||
team_wildcard_models = (team_pattern_router.route(model_id) or []) if team_pattern_router else []
|
||||
potential_wildcard_models = team_wildcard_models or self.pattern_router.route(model_id) or []
|
||||
global_wildcard_models = [
|
||||
wildcard_model
|
||||
for wildcard_model in (self.pattern_router.route(model_id) or [])
|
||||
if self._deployment_usable_by_team(wildcard_model, team_id)
|
||||
]
|
||||
potential_wildcard_models = team_wildcard_models or global_wildcard_models
|
||||
if potential_wildcard_models:
|
||||
# Use the first matching wildcard deployment
|
||||
deployment_dict = potential_wildcard_models[0]
|
||||
|
|
@ -9519,7 +9560,12 @@ class Router:
|
|||
return None
|
||||
|
||||
# Strategy 1: Check if model_id directly matches a model_name or deployment ID
|
||||
if model_id in self.model_names or self.has_model_id(model_id):
|
||||
if model_id in self.model_names:
|
||||
return model_id
|
||||
if self.has_model_id(model_id):
|
||||
deployment = self.get_deployment(model_id=model_id)
|
||||
if deployment is not None and deployment.model_name:
|
||||
return deployment.model_name
|
||||
return model_id
|
||||
|
||||
# Strategy 2: Search through router's model_list to find by litellm_params.model
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ from __future__ import annotations
|
|||
import asyncio
|
||||
import random
|
||||
import re
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Any, Literal, Union, cast
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
|
@ -25,6 +26,7 @@ from pydantic import BaseModel
|
|||
from litellm._logging import verbose_router_logger
|
||||
from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.llms.base_llm.base_utils import type_to_response_format_param
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
from .config import (
|
||||
|
|
@ -112,6 +114,16 @@ def _classifier_call_metadata(metadata: dict[str, Any] | None) -> dict[str, Any]
|
|||
}
|
||||
|
||||
|
||||
def _effective_turn_off_message_logging(request_kwargs: Mapping[str, Any] | None) -> bool | None:
|
||||
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
|
||||
initialize_standard_callback_dynamic_params,
|
||||
)
|
||||
|
||||
return initialize_standard_callback_dynamic_params(dict(request_kwargs) if request_kwargs else {}).get(
|
||||
"turn_off_message_logging"
|
||||
)
|
||||
|
||||
|
||||
class DimensionScore:
|
||||
"""Represents a score for a single dimension with optional signal."""
|
||||
|
||||
|
|
@ -427,7 +439,17 @@ class ComplexityRouter(CustomLogger):
|
|||
# attributed to the calling key/team instead of being dropped. Excludes the
|
||||
# parent request's budget reservation, which the routed completion (not this
|
||||
# internal classifier call) is responsible for reconciling.
|
||||
metadata = _classifier_call_metadata((request_kwargs or {}).get("litellm_metadata"))
|
||||
request_metadata = (request_kwargs or {}).get("litellm_metadata") or (request_kwargs or {}).get("metadata")
|
||||
metadata = _classifier_call_metadata(request_metadata)
|
||||
turn_off_message_logging = _effective_turn_off_message_logging(request_kwargs)
|
||||
|
||||
proxy_server_request = {
|
||||
"body": {
|
||||
"model": llm_config.model,
|
||||
"messages": [{"role": "user", "content": classification_prompt}],
|
||||
"response_format": type_to_response_format_param(TierClassification),
|
||||
}
|
||||
}
|
||||
|
||||
response: ModelResponse = await self.litellm_router_instance.acompletion(
|
||||
model=llm_config.model,
|
||||
|
|
@ -435,6 +457,8 @@ class ComplexityRouter(CustomLogger):
|
|||
response_format=TierClassification,
|
||||
timeout=llm_config.timeout_ms / 1000,
|
||||
metadata=metadata,
|
||||
proxy_server_request=proxy_server_request,
|
||||
turn_off_message_logging=turn_off_message_logging,
|
||||
)
|
||||
content = response.choices[0].message.content
|
||||
if not content:
|
||||
|
|
@ -821,8 +845,16 @@ class ComplexityRouter(CustomLogger):
|
|||
# key/team budget. Key/team attribution fields are preserved for spend logging.
|
||||
metadata = _classifier_call_metadata(request_kwargs.get("metadata"))
|
||||
litellm_metadata = _classifier_call_metadata(request_kwargs.get("litellm_metadata"))
|
||||
turn_off_message_logging = _effective_turn_off_message_logging(request_kwargs)
|
||||
proxy_server_request = {"body": {"model": self.config.embedding_model, "input": [user_message]}}
|
||||
query_vector = (
|
||||
await encoder.aencode_queries([user_message], metadata=metadata, litellm_metadata=litellm_metadata)
|
||||
await encoder.aencode_queries(
|
||||
[user_message],
|
||||
metadata=metadata,
|
||||
litellm_metadata=litellm_metadata,
|
||||
proxy_server_request=proxy_server_request,
|
||||
turn_off_message_logging=turn_off_message_logging,
|
||||
)
|
||||
)[0]
|
||||
route_choice = await routelayer.acall(vector=query_vector)
|
||||
|
||||
|
|
|
|||
|
|
@ -73,13 +73,20 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
precise_minute = f"{current_date}-{current_hour}-{current_minute}"
|
||||
|
||||
response_ms = end_time - start_time
|
||||
if isinstance(response_ms, timedelta):
|
||||
# normalize to float seconds up-front: non-chat responses
|
||||
# (embeddings, speech, image) skip the ModelResponse branch
|
||||
# below, and a raw timedelta appended to the latency list
|
||||
# breaks JSON serialization when the router cache syncs to
|
||||
# Redis (issue #33169)
|
||||
response_ms = response_ms.total_seconds()
|
||||
time_to_first_token_response_time = None
|
||||
|
||||
if kwargs.get("stream", None) is not None and kwargs["stream"] is True:
|
||||
# only log ttft for streaming request
|
||||
time_to_first_token_response_time = kwargs.get("completion_start_time", end_time) - start_time
|
||||
|
||||
final_value: Union[float, timedelta] = response_ms
|
||||
final_value: float = response_ms
|
||||
time_to_first_token: Optional[float] = None
|
||||
total_tokens = 0
|
||||
|
||||
|
|
@ -89,15 +96,12 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
completion_tokens = _usage.completion_tokens
|
||||
total_tokens = _usage.total_tokens
|
||||
|
||||
# Handle both timedelta and float response times
|
||||
if isinstance(response_ms, timedelta):
|
||||
response_seconds = response_ms.total_seconds()
|
||||
else:
|
||||
response_seconds = response_ms
|
||||
# response_ms is already normalized to float seconds above
|
||||
response_seconds = response_ms
|
||||
|
||||
final_value = safe_divide_seconds(response_seconds, completion_tokens)
|
||||
if final_value is not None:
|
||||
final_value = float(final_value)
|
||||
normalized_value = safe_divide_seconds(response_seconds, completion_tokens)
|
||||
if normalized_value is not None:
|
||||
final_value = float(normalized_value)
|
||||
else:
|
||||
final_value = response_seconds
|
||||
|
||||
|
|
@ -262,12 +266,19 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
precise_minute = f"{current_date}-{current_hour}-{current_minute}"
|
||||
|
||||
response_ms = end_time - start_time
|
||||
if isinstance(response_ms, timedelta):
|
||||
# normalize to float seconds up-front: non-chat responses
|
||||
# (embeddings, speech, image) skip the ModelResponse branch
|
||||
# below, and a raw timedelta appended to the latency list
|
||||
# breaks JSON serialization when the router cache syncs to
|
||||
# Redis (issue #33169)
|
||||
response_ms = response_ms.total_seconds()
|
||||
time_to_first_token_response_time = None
|
||||
if kwargs.get("stream", None) is not None and kwargs["stream"] is True:
|
||||
# only log ttft for streaming request
|
||||
time_to_first_token_response_time = kwargs.get("completion_start_time", end_time) - start_time
|
||||
|
||||
final_value: Union[float, timedelta] = response_ms
|
||||
final_value: float = response_ms
|
||||
total_tokens = 0
|
||||
time_to_first_token: Optional[float] = None
|
||||
|
||||
|
|
@ -277,17 +288,14 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
completion_tokens = _usage.completion_tokens
|
||||
total_tokens = _usage.total_tokens
|
||||
|
||||
# Handle both timedelta and float response times
|
||||
if isinstance(response_ms, timedelta):
|
||||
response_seconds = response_ms.total_seconds()
|
||||
else:
|
||||
response_seconds = response_ms
|
||||
# response_ms is already normalized to float seconds above
|
||||
response_seconds = response_ms
|
||||
|
||||
final_value = safe_divide_seconds(response_seconds, completion_tokens)
|
||||
if final_value is not None:
|
||||
final_value = float(final_value)
|
||||
normalized_value = safe_divide_seconds(response_seconds, completion_tokens)
|
||||
if normalized_value is not None:
|
||||
final_value = float(normalized_value)
|
||||
else:
|
||||
final_value = response_ms
|
||||
final_value = response_seconds
|
||||
|
||||
if time_to_first_token_response_time is not None:
|
||||
if isinstance(time_to_first_token_response_time, timedelta):
|
||||
|
|
|
|||
101
litellm/router_utils/auto_router_model_naming.py
Normal file
101
litellm/router_utils/auto_router_model_naming.py
Normal file
|
|
@ -0,0 +1,101 @@
|
|||
"""Naming contract for strategy-router (auto-router) pseudo-models.
|
||||
|
||||
A deployment whose ``litellm_params.model`` starts with ``auto_router/`` does not
|
||||
name a provider model; the string is the discriminator that selects which
|
||||
pre-routing strategy owns the deployment. This module is the single source of
|
||||
truth for classifying that string (``Router._is_*_router_deployment`` delegates
|
||||
here) and for checking that a client-supplied write leaves the deployment
|
||||
coherent, so management endpoints can reject corruption with a 400 instead of
|
||||
the router silently dropping the deployment at load time under
|
||||
``ignore_invalid_deployments``.
|
||||
"""
|
||||
|
||||
from typing import Literal, Mapping
|
||||
|
||||
AUTO_ROUTER_MODEL_PREFIX = "auto_router/"
|
||||
|
||||
StrategyRouterKind = Literal["semantic", "complexity", "adaptive", "quality"]
|
||||
|
||||
STRATEGY_ROUTER_PARAM_FIELDS: frozenset[str] = frozenset(
|
||||
{
|
||||
"auto_router_config",
|
||||
"auto_router_config_path",
|
||||
"auto_router_default_model",
|
||||
"auto_router_embedding_model",
|
||||
"complexity_router_config",
|
||||
"complexity_router_default_model",
|
||||
"adaptive_router_config",
|
||||
"quality_router_config",
|
||||
"quality_router_default_model",
|
||||
}
|
||||
)
|
||||
|
||||
_REQUIRED_FIELD_GROUPS: Mapping[StrategyRouterKind, tuple[tuple[str, ...], ...]] = {
|
||||
"semantic": (
|
||||
("auto_router_config", "auto_router_config_path"),
|
||||
("auto_router_default_model",),
|
||||
("auto_router_embedding_model",),
|
||||
),
|
||||
"complexity": (("complexity_router_config", "complexity_router_default_model"),),
|
||||
"adaptive": (("adaptive_router_config",),),
|
||||
"quality": (("quality_router_config", "quality_router_default_model"),),
|
||||
}
|
||||
|
||||
|
||||
def classify_strategy_router_model(model: str) -> StrategyRouterKind | None:
|
||||
"""Classify a ``litellm_params.model`` string the way the Router does.
|
||||
|
||||
Returns None for regular provider models. Mirrors Router registration
|
||||
exactly: reserved names are matched by prefix, everything else under
|
||||
``auto_router/`` is a semantic router.
|
||||
"""
|
||||
if not model.startswith(AUTO_ROUTER_MODEL_PREFIX):
|
||||
return None
|
||||
remainder = model[len(AUTO_ROUTER_MODEL_PREFIX) :]
|
||||
if remainder.startswith("complexity_router"):
|
||||
return "complexity"
|
||||
if remainder.startswith("adaptive_router"):
|
||||
return "adaptive"
|
||||
if remainder.startswith("quality_router"):
|
||||
return "quality"
|
||||
return "semantic"
|
||||
|
||||
|
||||
def validate_strategy_router_model_write(model: str, present_fields: frozenset[str]) -> str | None:
|
||||
"""Check that writing ``model`` leaves a deployment the router can load.
|
||||
|
||||
``present_fields`` is the set of strategy-router param fields that are
|
||||
non-None on the deployment after the write (stored fields merged with the
|
||||
incoming ones). Returns a human-readable violation, or None when coherent.
|
||||
"""
|
||||
kind = classify_strategy_router_model(model)
|
||||
if kind is None:
|
||||
offending = sorted(present_fields & STRATEGY_ROUTER_PARAM_FIELDS)
|
||||
if offending:
|
||||
return (
|
||||
f"litellm_params.model='{model}' does not start with '{AUTO_ROUTER_MODEL_PREFIX}' but the "
|
||||
f"deployment carries auto-router settings ({', '.join(offending)}), so the router could not "
|
||||
f"load it. Keep the '{AUTO_ROUTER_MODEL_PREFIX}' prefix; to change the name clients call, "
|
||||
"edit the public model_name instead."
|
||||
)
|
||||
return None
|
||||
remainder = model[len(AUTO_ROUTER_MODEL_PREFIX) :]
|
||||
if remainder.startswith(AUTO_ROUTER_MODEL_PREFIX):
|
||||
return (
|
||||
f"litellm_params.model='{model}' repeats the '{AUTO_ROUTER_MODEL_PREFIX}' prefix, so the router "
|
||||
f"could not load it. Use '{remainder}'; to change the name clients call, edit the public "
|
||||
"model_name instead."
|
||||
)
|
||||
if not remainder:
|
||||
return (
|
||||
f"litellm_params.model='{model}' is missing the router name after the '{AUTO_ROUTER_MODEL_PREFIX}' prefix."
|
||||
)
|
||||
missing = tuple(
|
||||
" or ".join(group) for group in _REQUIRED_FIELD_GROUPS[kind] if not any(f in present_fields for f in group)
|
||||
)
|
||||
if missing:
|
||||
return (
|
||||
f"litellm_params.model='{model}' selects the {kind} router, which requires "
|
||||
f"{'; '.join(missing)} in litellm_params."
|
||||
)
|
||||
return None
|
||||
|
|
@ -16,7 +16,7 @@ GeminiEmbeddingInput = Union[EmbeddingInput, List[List[str]]]
|
|||
|
||||
class FunctionResponse(TypedDict, total=False):
|
||||
# `id` correlates this response with the originating `functionCall` part.
|
||||
# Supported on Google AI Studio Gemini 3.5+; Vertex AI rejects this field.
|
||||
# Supported on Gemini 3+; older Gemini models reject this field.
|
||||
id: str
|
||||
name: Required[str]
|
||||
response: Optional[dict]
|
||||
|
|
@ -24,8 +24,8 @@ class FunctionResponse(TypedDict, total=False):
|
|||
|
||||
|
||||
class FunctionCall(TypedDict, total=False):
|
||||
# `id` correlates the corresponding `functionResponse` on Google AI Studio
|
||||
# Gemini 3.5+. Vertex AI and older Gemini models omit/reject this field.
|
||||
# `id` correlates the corresponding `functionResponse` on Gemini 3+.
|
||||
# Older Gemini models omit/reject this field.
|
||||
id: str
|
||||
name: Required[str]
|
||||
args: Optional[dict]
|
||||
|
|
@ -58,8 +58,8 @@ class PartType(TypedDict, total=False):
|
|||
|
||||
|
||||
class HttpxFunctionCall(TypedDict, total=False):
|
||||
# `id` correlates the corresponding `functionResponse` on Google AI Studio
|
||||
# Gemini 3.5+. Vertex AI and older Gemini models omit/reject this field.
|
||||
# `id` correlates the corresponding `functionResponse` on Gemini 3+.
|
||||
# Older Gemini models omit/reject this field.
|
||||
id: str
|
||||
name: Required[str]
|
||||
args: dict
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -24,7 +24,7 @@
|
|||
"limit": 130
|
||||
},
|
||||
"ANN401": {
|
||||
"limit": 2015
|
||||
"limit": 2013
|
||||
},
|
||||
"ASYNC230": {
|
||||
"limit": 14
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
|
|
@ -11,13 +11,14 @@ request/response bodies are co-located here because only this suite speaks MCP.
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
import time
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, RootModel
|
||||
|
||||
from e2e_http import Headers, NoBody, Result, Success, unwrap
|
||||
from e2e_http import Headers, NoBody, Result, Success, UnknownApiError, unwrap
|
||||
from models import KeyGenerateBody, ObjectPermission
|
||||
from proxy_client import ProxyClient
|
||||
|
||||
|
|
@ -270,6 +271,60 @@ class McpClient:
|
|||
)
|
||||
time.sleep(self.proxy.poll_interval)
|
||||
|
||||
def await_call_tool(
|
||||
self,
|
||||
key: str,
|
||||
*,
|
||||
server_id: str,
|
||||
name: str,
|
||||
arguments: McpToolArguments,
|
||||
) -> McpCallToolResponse:
|
||||
"""Poll tools/call until the result is not a multi-worker registry miss.
|
||||
|
||||
Retries only on the gateway's own cold-worker 500 shapes (Tool <name>
|
||||
not found / server_not_found). Upstream tool errors and other 500s fail
|
||||
immediately so non-idempotent calls are not repeated.
|
||||
"""
|
||||
deadline = time.monotonic() + self.proxy.poll_timeout
|
||||
last: Result[McpCallToolResponse] | None = None
|
||||
while True:
|
||||
last = self.call_tool(key, server_id=server_id, name=name, arguments=arguments)
|
||||
if not _is_mcp_not_synced(last, tool_name=name):
|
||||
return unwrap(last)
|
||||
if time.monotonic() >= deadline:
|
||||
raise AssertionError(
|
||||
f"tools/call for {name!r} on server {server_id} still missing on the "
|
||||
f"data plane after {self.proxy.poll_timeout}s (multi-worker registry lag); "
|
||||
f"last result: {last}"
|
||||
)
|
||||
time.sleep(self.proxy.poll_interval)
|
||||
|
||||
def await_call_tool_denied(
|
||||
self,
|
||||
key: str,
|
||||
*,
|
||||
server_id: str,
|
||||
name: str,
|
||||
arguments: McpToolArguments,
|
||||
) -> UnknownApiError:
|
||||
"""Poll tools/call until a cold-worker miss clears and the call is 403 access_denied."""
|
||||
deadline = time.monotonic() + self.proxy.poll_timeout
|
||||
last: Result[McpCallToolResponse] | None = None
|
||||
while True:
|
||||
last = self.call_tool(key, server_id=server_id, name=name, arguments=arguments)
|
||||
if isinstance(last, UnknownApiError) and last.status_code == 403:
|
||||
return last
|
||||
if not _is_mcp_not_synced(last, tool_name=name):
|
||||
raise AssertionError(
|
||||
f"ungranted key's tools/call was not 403 access_denied: {last}"
|
||||
)
|
||||
if time.monotonic() >= deadline:
|
||||
raise AssertionError(
|
||||
f"ungranted key never got 403 for {name!r} within {self.proxy.poll_timeout}s; "
|
||||
f"last result: {last}"
|
||||
)
|
||||
time.sleep(self.proxy.poll_interval)
|
||||
|
||||
def register_mcp_content_filter(self, *, name: str, blocked_keyword: str) -> str:
|
||||
"""Register a default-on content-filter guardrail that runs on the MCP
|
||||
tool-call hook (pre_mcp_call) and blocks a single keyword. The keyword is
|
||||
|
|
@ -317,5 +372,39 @@ class McpClient:
|
|||
)
|
||||
|
||||
|
||||
def _is_mcp_not_synced(
|
||||
result: Result[McpCallToolResponse],
|
||||
*,
|
||||
tool_name: str | None = None,
|
||||
) -> bool:
|
||||
"""True only for gateway multi-worker registry misses, not upstream errors.
|
||||
|
||||
Matches the proxy's own shapes:
|
||||
- ValueError ``Tool <name> not found`` wrapped as HTTP 500 (cold tool map /
|
||||
unresolved server on this process)
|
||||
- REST ``server_not_found`` when this worker has not loaded the MCP server row
|
||||
|
||||
Does not treat arbitrary 500 bodies that merely mention "tool" and "not found"
|
||||
(e.g. upstream MCP payload text) as lag, so await_call_tool does not retry
|
||||
real failures or non-idempotent calls.
|
||||
"""
|
||||
if not isinstance(result, UnknownApiError) or result.status_code != 500:
|
||||
return False
|
||||
body = result.body
|
||||
body_l = body.lower()
|
||||
|
||||
if "server_not_found" in body_l:
|
||||
return True
|
||||
if re.search(r"mcp server ['\"][^'\"]+['\"] was not found", body_l):
|
||||
return True
|
||||
|
||||
# Gateway: "Tool search_datadog_logs not found" (optionally inside a longer message)
|
||||
if tool_name is not None:
|
||||
return (
|
||||
re.search(rf"\btool\s+{re.escape(tool_name)}\s+not found\b", body_l) is not None
|
||||
)
|
||||
return re.search(r"\btool\s+\S+\s+not found\b", body_l) is not None
|
||||
|
||||
|
||||
def build_client(proxy: ProxyClient) -> McpClient:
|
||||
return McpClient(proxy=proxy)
|
||||
|
|
|
|||
|
|
@ -29,6 +29,7 @@ class TestMcpAccessGroupToolSelection:
|
|||
) -> None:
|
||||
group = f"e2e-mcp-grp-{unique_marker()}"
|
||||
server_id = register_datadog_mcp(client, resources, mcp_access_groups=[group])
|
||||
client.await_registered(server_id)
|
||||
|
||||
granted = client.generate_key(
|
||||
user_id=f"e2e-mcp-ag-granted-{unique_marker()}",
|
||||
|
|
|
|||
|
|
@ -60,6 +60,7 @@ class TestDatadogMcpRoundTrip:
|
|||
_assert_datadog_logger_active(client.proxy)
|
||||
|
||||
server_id = register_datadog_mcp(client, resources)
|
||||
client.await_registered(server_id)
|
||||
marker = f"{MARKER_PREFIX}{unique_marker()}"
|
||||
|
||||
key = client.generate_key(
|
||||
|
|
@ -78,22 +79,19 @@ class TestDatadogMcpRoundTrip:
|
|||
)
|
||||
|
||||
tool_name = client.await_tool(key, server_id, SEARCH_LOGS_TOOL)
|
||||
|
||||
call = unwrap(
|
||||
client.call_tool(
|
||||
key,
|
||||
server_id=server_id,
|
||||
name=tool_name,
|
||||
arguments={
|
||||
"query": marker,
|
||||
"from": DD_SEARCH_FROM,
|
||||
"to": "now",
|
||||
"max_tokens": 5000,
|
||||
"telemetry": {
|
||||
"intent": "e2e assert seeded litellm completion log is searchable via MCP"
|
||||
},
|
||||
call = client.await_call_tool(
|
||||
key,
|
||||
server_id=server_id,
|
||||
name=tool_name,
|
||||
arguments={
|
||||
"query": marker,
|
||||
"from": DD_SEARCH_FROM,
|
||||
"to": "now",
|
||||
"max_tokens": 5000,
|
||||
"telemetry": {
|
||||
"intent": "e2e assert seeded litellm completion log is searchable via MCP"
|
||||
},
|
||||
)
|
||||
},
|
||||
)
|
||||
assert call.is_error is not True, f"search_datadog_logs errored: {call}"
|
||||
body = call.all_text
|
||||
|
|
|
|||
|
|
@ -16,7 +16,7 @@ import pytest
|
|||
|
||||
from datadog_mcp import SEARCH_LOGS_TOOL, register_datadog_mcp
|
||||
from e2e_config import DD_SEARCH_FROM, unique_marker
|
||||
from e2e_http import UnknownApiError, unwrap
|
||||
from e2e_http import unwrap
|
||||
from lifecycle import ResourceManager
|
||||
from mcp_client import McpClient
|
||||
|
||||
|
|
@ -72,13 +72,12 @@ class TestMcpKeyWithoutAccessIsDenied:
|
|||
"max_tokens": 1000,
|
||||
"telemetry": {"intent": "e2e control call proving granted key can invoke Datadog MCP"},
|
||||
}
|
||||
permitted_call = unwrap(
|
||||
client.call_tool(permitted_key, server_id=server_id, name=tool_name, arguments=search_args)
|
||||
permitted_call = client.await_call_tool(
|
||||
permitted_key, server_id=server_id, name=tool_name, arguments=search_args
|
||||
)
|
||||
assert permitted_call.is_error is not True, f"granted key's tool call errored: {permitted_call}"
|
||||
|
||||
match client.call_tool(denied_key, server_id=server_id, name=tool_name, arguments=search_args):
|
||||
case UnknownApiError(status_code=403, body=body):
|
||||
assert "access_denied" in body, f"403 was not an MCP access denial: {body}"
|
||||
case other:
|
||||
pytest.fail(f"ungranted key's tool call was not refused with 403 access_denied: {other}")
|
||||
denied = client.await_call_tool_denied(
|
||||
denied_key, server_id=server_id, name=tool_name, arguments=search_args
|
||||
)
|
||||
assert "access_denied" in denied.body, f"403 was not an MCP access denial: {denied.body}"
|
||||
|
|
|
|||
|
|
@ -2184,7 +2184,8 @@ def test_team_alias_stale_bypass_disabled_by_default(monkeypatch):
|
|||
pre_call_utils._ENABLE_TEAM_STALE_ALIAS_BYPASS = None
|
||||
|
||||
class _MockRouter:
|
||||
team_model_to_deployment_indices = {("team-1", "gpt-4o"): [0]}
|
||||
model_name_to_deployment_indices = {"model_name_team-1_legacy-uuid": [0]}
|
||||
team_model_to_deployment_indices = {("team-1", "gpt-4o"): [1]}
|
||||
|
||||
test_data = {"model": "gpt-4o"}
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
|
|
@ -2209,7 +2210,8 @@ def test_team_alias_stale_bypass_enabled_by_flag(monkeypatch):
|
|||
pre_call_utils._ENABLE_TEAM_STALE_ALIAS_BYPASS = None
|
||||
|
||||
class _MockRouter:
|
||||
team_model_to_deployment_indices = {("team-1", "gpt-4o"): [0]}
|
||||
model_name_to_deployment_indices = {"model_name_team-1_legacy-uuid": [0]}
|
||||
team_model_to_deployment_indices = {("team-1", "gpt-4o"): [1]}
|
||||
|
||||
test_data = {"model": "gpt-4o"}
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
|
|
|
|||
|
|
@ -2659,6 +2659,23 @@ def test_resolve_model_name_from_model_id():
|
|||
result = router.resolve_model_name_from_model_id("gpt-5-mini")
|
||||
assert result == "gpt-5-mini"
|
||||
|
||||
# Test case 10: model_id is a deployment ID (hash) that differs from the
|
||||
# public model_name. Regression for #32580: managed batch/file IDs embed the
|
||||
# deployment model_id, and it must resolve back to the public model_name so
|
||||
# team model-access checks compare against the model group, not the hash.
|
||||
model_list = [
|
||||
{
|
||||
"model_name": "bedrock-batch-model",
|
||||
"litellm_params": {
|
||||
"model": "bedrock/global.anthropic.claude-haiku-4-5-20251001-v1:0",
|
||||
},
|
||||
"model_info": {"id": "8d0eaa7e6c6f54a425dfd0062cb6b0dc"},
|
||||
},
|
||||
]
|
||||
router = Router(model_list=model_list)
|
||||
result = router.resolve_model_name_from_model_id("8d0eaa7e6c6f54a425dfd0062cb6b0dc")
|
||||
assert result == "bedrock-batch-model"
|
||||
|
||||
|
||||
def test_get_valid_args():
|
||||
"""Test get_valid_args static method returns valid Router.__init__ arguments"""
|
||||
|
|
|
|||
|
|
@ -5,6 +5,8 @@ Regression test for afile_retrieve called without credentials in
|
|||
async_post_call_success_hook when processing completed batch responses.
|
||||
"""
|
||||
|
||||
import json
|
||||
|
||||
import pytest
|
||||
from typing import Optional
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
|
@ -385,3 +387,59 @@ async def test_afile_content_error_reports_unified_id_not_provider_uri():
|
|||
message = str(exc_info.value)
|
||||
assert unified_file_id in message
|
||||
assert s3_uri not in message
|
||||
|
||||
|
||||
def _make_real_managed_files_instance():
|
||||
"""Create a _PROXY_LiteLLMManagedFiles with a real store_unified_file_id but
|
||||
an AsyncMock prisma client, so the DB write path itself can be asserted."""
|
||||
from litellm_enterprise.proxy.hooks.managed_files import (
|
||||
_PROXY_LiteLLMManagedFiles,
|
||||
)
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.async_set_cache = AsyncMock()
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_managedfiletable.upsert = AsyncMock()
|
||||
mock_prisma.db.litellm_managedfiletable.create = AsyncMock(
|
||||
side_effect=AssertionError(
|
||||
"store_unified_file_id must upsert, not create, on the retrieve path"
|
||||
)
|
||||
)
|
||||
|
||||
return (
|
||||
_PROXY_LiteLLMManagedFiles(
|
||||
internal_usage_cache=mock_cache,
|
||||
prisma_client=mock_prisma,
|
||||
),
|
||||
mock_prisma,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_store_unified_file_id_is_idempotent_via_upsert():
|
||||
"""Regression test for the managed-batch retrieve 500 (UniqueViolationError on
|
||||
unified_file_id): re-registering an already-stored output file id must upsert on
|
||||
unified_file_id, never do an unconditional create that raises on conflict."""
|
||||
managed_files, mock_prisma = _make_real_managed_files_instance()
|
||||
file_id = "litellm_proxy_unified_output_id_abc"
|
||||
model_mappings = {"model-deploy-xyz": "file-output-abc"}
|
||||
|
||||
for _ in range(2):
|
||||
await managed_files.store_unified_file_id(
|
||||
file_id=file_id,
|
||||
file_object=_make_file_object(),
|
||||
litellm_parent_otel_span=None,
|
||||
model_mappings=model_mappings,
|
||||
user_api_key_dict=_make_user_api_key_dict(),
|
||||
)
|
||||
|
||||
mock_prisma.db.litellm_managedfiletable.create.assert_not_awaited()
|
||||
upsert_mock = mock_prisma.db.litellm_managedfiletable.upsert
|
||||
assert upsert_mock.await_count == 2
|
||||
for upsert_call in upsert_mock.await_args_list:
|
||||
assert upsert_call.kwargs["where"] == {"unified_file_id": file_id}
|
||||
upsert_data = upsert_call.kwargs["data"]
|
||||
assert upsert_data["create"]["unified_file_id"] == file_id
|
||||
assert json.loads(upsert_data["create"]["model_mappings"]) == model_mappings
|
||||
assert json.loads(upsert_data["update"]["model_mappings"]) == model_mappings
|
||||
|
|
|
|||
|
|
@ -451,6 +451,30 @@ def test_parse_headers():
|
|||
assert providers.parse_headers("no-equals") == {}
|
||||
|
||||
|
||||
def test_parse_headers_percent_decodes_values():
|
||||
"""A percent-encoded OTLP header value reaches the exporter decoded.
|
||||
|
||||
``OTEL_EXPORTER_OTLP_HEADERS`` is W3C Baggage encoded, and Grafana Cloud
|
||||
documents ``Authorization=Basic%20<token>``. Forwarding the literal ``%20``
|
||||
makes the backend reject the export as a malformed credential.
|
||||
"""
|
||||
token = "MTMzNzc4MzpnbGNfZXlKdklqb2lNVEl6TkNJPQ=="
|
||||
assert providers.parse_headers(f"Authorization=Basic%20{token}") == {"authorization": f"Basic {token}"}
|
||||
assert providers.parse_headers("x-scope-orgid=team%20a") == {"x-scope-orgid": "team a"}
|
||||
|
||||
|
||||
def test_parse_headers_keeps_unencoded_values_working():
|
||||
"""Values that are not percent-encoded keep parsing unchanged.
|
||||
|
||||
Vendors that document a bare space, and litellm's own presets, must survive
|
||||
the switch to the spec-compliant parser. Base64 padding also means a value
|
||||
can contain ``=``, so only the first one may split the pair.
|
||||
"""
|
||||
assert providers.parse_headers("Authorization=Bearer sk-123") == {"authorization": "Bearer sk-123"}
|
||||
assert providers.parse_headers("api_key=abc,space_id=xyz") == {"api_key": "abc", "space_id": "xyz"}
|
||||
assert providers.parse_headers("api_key=YWJjZA==") == {"api_key": "YWJjZA=="}
|
||||
|
||||
|
||||
def test_otlp_traces_endpoint_normalization():
|
||||
norm = providers._otlp_traces_endpoint
|
||||
# A base endpoint gets the signal path appended (the common OTLP env shape).
|
||||
|
|
@ -487,6 +511,24 @@ def test_build_span_exporter_variants():
|
|||
assert "OTLPSpanExporter" in type(http_exporter).__name__
|
||||
|
||||
|
||||
def test_otlp_metric_exporter_uses_cumulative_histogram_temporality():
|
||||
"""Histograms must export as cumulative, not delta.
|
||||
|
||||
Prometheus-backed OTLP receivers (Grafana Cloud / Mimir) reject delta
|
||||
histograms with ``invalid temporality and type combination`` and drop the
|
||||
entire metric batch, so a delta default silently loses every GenAI metric.
|
||||
"""
|
||||
from opentelemetry.sdk.metrics import Histogram
|
||||
from opentelemetry.sdk.metrics.export import AggregationTemporality
|
||||
|
||||
reader = providers.build_metric_reader(
|
||||
OpenTelemetryV2Config(exporter="otlp_http", endpoint="http://h:4318")
|
||||
)
|
||||
temporality = reader._exporter._preferred_temporality # noqa: SLF001 # exporter exposes no public accessor
|
||||
|
||||
assert temporality[Histogram] is AggregationTemporality.CUMULATIVE
|
||||
|
||||
|
||||
def test_otlp_logs_endpoint_normalization():
|
||||
norm = providers._otlp_logs_endpoint
|
||||
# A base endpoint gets the signal path appended (the common OTLP env shape).
|
||||
|
|
|
|||
|
|
@ -2120,9 +2120,9 @@ def test_valid_metric_filter_records_six_metrics(monkeypatch):
|
|||
assert _emitted_metric_names(reader) == {
|
||||
"gen_ai.client.operation.duration",
|
||||
"gen_ai.client.token.usage",
|
||||
"gen_ai.client.token.cost",
|
||||
"gen_ai.client.response.time_to_first_token",
|
||||
"gen_ai.client.response.time_per_output_token",
|
||||
"gen_ai.usage.cost",
|
||||
"gen_ai.server.time_to_first_token",
|
||||
"gen_ai.server.time_per_output_token",
|
||||
"gen_ai.client.response.duration",
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -37,6 +37,13 @@ def _load_openapi_spec_dict() -> Dict[str, Any]:
|
|||
)
|
||||
|
||||
|
||||
def _declared_type_value(variant_schema: Dict[str, Any]) -> Any:
|
||||
"""The single `type` value a union variant pins, whether spelled as a const or a 1-item enum."""
|
||||
type_property = variant_schema.get("properties", {}).get("type", {})
|
||||
enum_values = type_property.get("enum") or []
|
||||
return type_property.get("const") or (enum_values[0] if len(enum_values) == 1 else None)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def spec_dict() -> Dict[str, Any]:
|
||||
"""Load raw spec dict for manual validation."""
|
||||
|
|
@ -105,26 +112,51 @@ class TestRequestCompliance:
|
|||
assert "string" in input_types, "Input should support string"
|
||||
assert "array" in input_types, "Input should support array"
|
||||
|
||||
def test_content_schema_uses_discriminator(self, spec_dict):
|
||||
"""Verify Content uses type discriminator."""
|
||||
def test_content_variants_are_identified_by_their_type_field(self, spec_dict):
|
||||
"""Verify a Content part can be told apart by its `type`, however the spec spells that.
|
||||
|
||||
Our transformation reads `type` off each content part to route it, so what has to hold is
|
||||
that every variant of the union pins a distinct `type` value and that text is one of them.
|
||||
A spec may express that with an OpenAPI `discriminator` on the union or with a `const` on
|
||||
each member's own `type`; both are equivalent for us, so accepting only the first makes
|
||||
this test fail on a stylistic change upstream that costs us nothing.
|
||||
"""
|
||||
content_schema = spec_dict["components"]["schemas"]["Content"]
|
||||
|
||||
assert "discriminator" in content_schema
|
||||
assert content_schema["discriminator"]["propertyName"] == "type"
|
||||
|
||||
# Check TextContent is an option (via mapping if present, or via oneOf refs)
|
||||
mapping = content_schema["discriminator"].get("mapping")
|
||||
if mapping:
|
||||
assert "text" in mapping
|
||||
print(f"Content type discriminator mapping: {list(mapping.keys())}")
|
||||
else:
|
||||
# Discriminator without explicit mapping — verify via oneOf
|
||||
one_of = content_schema.get("oneOf", [])
|
||||
ref_names = [opt["$ref"].split("/")[-1] for opt in one_of if "$ref" in opt]
|
||||
discriminator = content_schema.get("discriminator")
|
||||
if discriminator is not None:
|
||||
assert (
|
||||
"TextContent" in ref_names
|
||||
), f"TextContent not found in oneOf refs: {ref_names}"
|
||||
print(f"Content type discriminator (no mapping), oneOf refs: {ref_names}")
|
||||
discriminator.get("propertyName") == "type"
|
||||
), f"Content is discriminated on {discriminator.get('propertyName')!r}, not 'type'"
|
||||
|
||||
variant_names = [
|
||||
option["$ref"].split("/")[-1]
|
||||
for option in content_schema.get("oneOf", [])
|
||||
if "$ref" in option
|
||||
]
|
||||
assert variant_names, f"Content is not a union of named variants: {content_schema}"
|
||||
|
||||
mapping = (discriminator or {}).get("mapping") or {}
|
||||
type_values = {
|
||||
variant: mapping_value
|
||||
for mapping_value, ref in mapping.items()
|
||||
for variant in [ref.split("/")[-1]]
|
||||
} or {
|
||||
variant: _declared_type_value(spec_dict["components"]["schemas"].get(variant, {}))
|
||||
for variant in variant_names
|
||||
}
|
||||
|
||||
assert set(type_values) == set(variant_names) and all(type_values.values()), (
|
||||
f"every Content variant needs a discoverable type value, "
|
||||
f"got {type_values} for variants {sorted(variant_names)}"
|
||||
)
|
||||
assert len(set(type_values.values())) == len(type_values), (
|
||||
f"Content variants must pin DISTINCT type values, got {type_values}"
|
||||
)
|
||||
assert type_values.get("TextContent") == "text", (
|
||||
f"TextContent must be reachable as type 'text', got {type_values}"
|
||||
)
|
||||
print(f"Content variants by type: {type_values}")
|
||||
|
||||
def test_text_content_schema(self, spec_dict):
|
||||
"""Verify TextContent schema."""
|
||||
|
|
|
|||
|
|
@ -506,3 +506,162 @@ def test_empty_content_chunk_mid_text_block_is_suppressed_sync():
|
|||
|
||||
assert _text_deltas(events) == ["Hi", " there"]
|
||||
_assert_deltas_match_their_block_type(events)
|
||||
|
||||
|
||||
def _thinking_first_chunks() -> List[MagicMock]:
|
||||
return [
|
||||
_thinking_chunk("Let me think"),
|
||||
_thinking_chunk("about it."),
|
||||
_make_chunk(Delta(content="42")),
|
||||
_make_chunk(Delta(content=None), finish_reason="stop"),
|
||||
]
|
||||
|
||||
|
||||
def _assert_thinking_first_block_opens_at_index_zero(events: List[dict]) -> None:
|
||||
starts = [
|
||||
(e["index"], e["content_block"]["type"])
|
||||
for e in events
|
||||
if e.get("type") == "content_block_start"
|
||||
]
|
||||
assert starts == [(0, "thinking"), (1, "text")], starts
|
||||
assert "" not in _text_deltas(events)
|
||||
assert _thinking_deltas(events) == ["Let me think", "about it."]
|
||||
assert _text_deltas(events) == ["42"]
|
||||
_assert_deltas_match_their_block_type(events)
|
||||
|
||||
|
||||
def test_thinking_first_stream_opens_thinking_block_at_index_zero_sync():
|
||||
"""Bug A regression: when the model's first output is reasoning the adapter
|
||||
must open the first content block as ``thinking`` at index 0. The previous
|
||||
code pre-emitted a hardcoded empty ``text`` block at index 0 before
|
||||
inspecting any upstream chunk, then opened ``thinking`` at index 1; strict
|
||||
Anthropic SDK clients with thinking enabled reject that stream with
|
||||
"Content block is not a thinking block".
|
||||
"""
|
||||
wrapper = AnthropicStreamWrapper(
|
||||
completion_stream=iter(_thinking_first_chunks()),
|
||||
model="claude-x",
|
||||
)
|
||||
_assert_thinking_first_block_opens_at_index_zero(_drain_sync(wrapper))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_thinking_first_stream_opens_thinking_block_at_index_zero_async():
|
||||
"""Async twin of the Bug A regression; the proxy serves the async iterator,
|
||||
so the first block must be ``thinking`` at index 0 on this path too.
|
||||
"""
|
||||
wrapper = AnthropicStreamWrapper(
|
||||
completion_stream=_AsyncStream(_thinking_first_chunks()),
|
||||
model="claude-x",
|
||||
)
|
||||
_assert_thinking_first_block_opens_at_index_zero(await _drain_async(wrapper))
|
||||
|
||||
|
||||
def _reasoning_content_chunk(reasoning: str) -> MagicMock:
|
||||
return _make_chunk(Delta(content=None, reasoning_content=reasoning))
|
||||
|
||||
|
||||
def _reasoning_first_chunks() -> List[MagicMock]:
|
||||
return [
|
||||
_reasoning_content_chunk("Let me think"),
|
||||
_reasoning_content_chunk("about it."),
|
||||
_make_chunk(Delta(content="42")),
|
||||
_make_chunk(Delta(content=None), finish_reason="stop"),
|
||||
]
|
||||
|
||||
|
||||
def test_reasoning_content_first_stream_opens_thinking_block_at_index_zero_sync():
|
||||
"""The reported backend (hosted_vllm; vLLM and SGLang reasoning parsers)
|
||||
surfaces reasoning as OpenAI ``reasoning_content`` with no
|
||||
``thinking_blocks``. Such a stream must also open the first content block as
|
||||
``thinking`` at index 0, exercising the reasoning_content branch of the
|
||||
chunk translator rather than the thinking_blocks branch the other twins use.
|
||||
"""
|
||||
wrapper = AnthropicStreamWrapper(
|
||||
completion_stream=iter(_reasoning_first_chunks()),
|
||||
model="claude-x",
|
||||
)
|
||||
_assert_thinking_first_block_opens_at_index_zero(_drain_sync(wrapper))
|
||||
|
||||
|
||||
def _blank_lead_chunks() -> List[MagicMock]:
|
||||
return [
|
||||
_make_chunk(Delta(content=None)),
|
||||
_thinking_chunk("Let me think"),
|
||||
_thinking_chunk("about it."),
|
||||
_make_chunk(Delta(content="42")),
|
||||
_make_chunk(Delta(content=None), finish_reason="stop"),
|
||||
]
|
||||
|
||||
|
||||
def _role_only_reasoning_content_lead_chunks() -> List[MagicMock]:
|
||||
return [
|
||||
_make_chunk(Delta(role="assistant", content=None, tool_calls=[])),
|
||||
_reasoning_content_chunk("Let me think"),
|
||||
_reasoning_content_chunk("about it."),
|
||||
_make_chunk(Delta(content="42")),
|
||||
_make_chunk(Delta(content=None), finish_reason="stop"),
|
||||
]
|
||||
|
||||
|
||||
def test_contentless_lead_chunk_does_not_open_text_block_before_thinking_sync():
|
||||
"""OpenAI-compatible streaming backends open the response with a contentless
|
||||
priming chunk (an empty delta, e.g. the {role: assistant} lead-in) before the
|
||||
first real token. Such a lead chunk must NOT commit index 0 to an empty text
|
||||
block; the following thinking chunk must still open thinking at index 0, or
|
||||
strict Anthropic SDK clients reject the stream.
|
||||
"""
|
||||
wrapper = AnthropicStreamWrapper(
|
||||
completion_stream=iter(_blank_lead_chunks()),
|
||||
model="claude-x",
|
||||
)
|
||||
_assert_thinking_first_block_opens_at_index_zero(_drain_sync(wrapper))
|
||||
|
||||
|
||||
def test_role_only_lead_chunk_does_not_open_text_block_before_reasoning_content_sync():
|
||||
wrapper = AnthropicStreamWrapper(
|
||||
completion_stream=iter(_role_only_reasoning_content_lead_chunks()),
|
||||
model="claude-x",
|
||||
)
|
||||
_assert_thinking_first_block_opens_at_index_zero(_drain_sync(wrapper))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_contentless_lead_chunk_does_not_open_text_block_before_thinking_async():
|
||||
"""Async twin; the proxy serves the async iterator, so the contentless lead
|
||||
chunk must be skipped on this path too.
|
||||
"""
|
||||
wrapper = AnthropicStreamWrapper(
|
||||
completion_stream=_AsyncStream(_blank_lead_chunks()),
|
||||
model="claude-x",
|
||||
)
|
||||
_assert_thinking_first_block_opens_at_index_zero(await _drain_async(wrapper))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_role_only_lead_chunk_does_not_open_text_block_before_reasoning_content_async():
|
||||
wrapper = AnthropicStreamWrapper(
|
||||
completion_stream=_AsyncStream(_role_only_reasoning_content_lead_chunks()),
|
||||
model="claude-x",
|
||||
)
|
||||
_assert_thinking_first_block_opens_at_index_zero(await _drain_async(wrapper))
|
||||
|
||||
|
||||
def test_finish_first_chunk_is_not_deferred_sync():
|
||||
"""A stream whose first upstream chunk is already the finish event must not
|
||||
be skipped by the blank-delta deferral. ``_is_blank_delta`` returns False
|
||||
for a finish chunk so the message_delta still flows (with an empty text
|
||||
block opened and closed first); without that guard the deferral would drop
|
||||
the terminal event entirely.
|
||||
"""
|
||||
chunks = [_make_chunk(Delta(content=None), finish_reason="stop")]
|
||||
wrapper = AnthropicStreamWrapper(completion_stream=iter(chunks), model="claude-x")
|
||||
events = _drain_sync(wrapper)
|
||||
|
||||
assert [e["type"] for e in events] == [
|
||||
"message_start",
|
||||
"content_block_start",
|
||||
"content_block_stop",
|
||||
"message_delta",
|
||||
"message_stop",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -134,13 +134,6 @@ def test_anthropic_stream_wrapper_single_tool_call():
|
|||
# Verify the expected sequence of chunk types
|
||||
expected_types = [
|
||||
"message_start", # Initial message start
|
||||
# TODO: for future contributors: if the initial content_block_start
|
||||
# respects the upstream's starting chunk, the initial empty text block
|
||||
# should be removed (and this test should be updated accordingly)
|
||||
# ---------------------------------------------------------------------
|
||||
"content_block_start", # Initial empty text block start
|
||||
"content_block_stop", # End of empty text block
|
||||
# ---------------------------------------------------------------------
|
||||
"content_block_start", # Start of first tool_use content block
|
||||
"content_block_delta", # {"city":
|
||||
"content_block_delta", # "NY"}
|
||||
|
|
@ -196,13 +189,6 @@ def test_anthropic_stream_wrapper_back_to_back_tool_calls():
|
|||
# Verify the expected sequence of chunk types
|
||||
expected_types = [
|
||||
"message_start", # Initial message start
|
||||
# TODO: for future contributors: if the initial content_block_start
|
||||
# respects the upstream's starting chunk, the initial empty text block
|
||||
# should be removed (and this test should be updated accordingly)
|
||||
# ---------------------------------------------------------------------
|
||||
"content_block_start", # Initial empty text block start
|
||||
"content_block_stop", # End of empty text block
|
||||
# ---------------------------------------------------------------------
|
||||
"content_block_start", # Start of first tool_use content block
|
||||
"content_block_delta", # {"city":
|
||||
"content_block_delta", # "NY"}
|
||||
|
|
@ -267,13 +253,6 @@ def test_anthropic_stream_wrapper_interleaved_tool_calls_and_text():
|
|||
# Verify the expected sequence of chunk types
|
||||
expected_types = [
|
||||
"message_start", # Initial message start
|
||||
# TODO: for future contributors: if the initial content_block_start
|
||||
# respects the upstream's starting chunk, the initial empty text block
|
||||
# should be removed (and this test should be updated accordingly)
|
||||
# ---------------------------------------------------------------------
|
||||
"content_block_start", # Initial empty text block start
|
||||
"content_block_stop", # End of empty text block
|
||||
# ---------------------------------------------------------------------
|
||||
"content_block_start", # Start of first tool_use content block
|
||||
"content_block_delta", # {"city":
|
||||
"content_block_delta", # "NY"}
|
||||
|
|
|
|||
|
|
@ -1452,6 +1452,157 @@ class TestContextCachingEndpoints:
|
|||
# Restart the patcher so teardown_method can stop it cleanly
|
||||
self._token_check_patcher.start()
|
||||
|
||||
def _model_turn_final_messages(self, final_cached_role):
|
||||
tool_call = {
|
||||
"id": "call_abc123",
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "arguments": '{"location": "Boston"}'},
|
||||
}
|
||||
cached_tail = {
|
||||
"assistant": [],
|
||||
"tool": [
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_abc123",
|
||||
"content": "72F and sunny",
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
],
|
||||
"system": [
|
||||
{
|
||||
"role": "system",
|
||||
"content": "Tool results are authoritative.",
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
],
|
||||
}[final_cached_role]
|
||||
return [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Use the weather tool for every answer.",
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [tool_call],
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
},
|
||||
*cached_tail,
|
||||
{"role": "user", "content": "What is the weather in Boston?"},
|
||||
]
|
||||
|
||||
@pytest.mark.parametrize("final_cached_role", ["assistant", "tool", "system"])
|
||||
def test_check_and_create_cache_skips_when_cached_block_ends_on_model_turn(
|
||||
self, final_cached_role
|
||||
):
|
||||
"""The cachedContents API rejects contents ending on an assistant or tool turn
|
||||
with HTTP 400 "Requests ending with a model turn are not supported", so the
|
||||
request must proceed uncached instead of failing.
|
||||
"""
|
||||
all_messages = self._model_turn_final_messages(final_cached_role)
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
|
||||
result = self.context_caching.check_and_create_cache(
|
||||
messages=all_messages,
|
||||
optional_params=optional_params,
|
||||
api_key="test_key",
|
||||
api_base=None,
|
||||
model="gemini-3.6-flash",
|
||||
client=self.mock_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
cached_content=None,
|
||||
custom_llm_provider="vertex_ai",
|
||||
vertex_project="test_project",
|
||||
vertex_location="us-central1",
|
||||
vertex_auth_header="test_token",
|
||||
)
|
||||
|
||||
messages, returned_params, returned_cache = result
|
||||
assert messages == all_messages
|
||||
assert returned_cache is None
|
||||
assert "tools" in returned_params
|
||||
self.mock_client.get.assert_not_called()
|
||||
self.mock_client.post.assert_not_called()
|
||||
|
||||
@pytest.mark.parametrize("final_cached_role", ["assistant", "tool", "system"])
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_check_and_create_cache_skips_when_cached_block_ends_on_model_turn(
|
||||
self, final_cached_role
|
||||
):
|
||||
"""Async variant: an unsupported terminal turn skips caching instead of failing."""
|
||||
all_messages = self._model_turn_final_messages(final_cached_role)
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
|
||||
result = await self.context_caching.async_check_and_create_cache(
|
||||
messages=all_messages,
|
||||
optional_params=optional_params,
|
||||
api_key="test_key",
|
||||
api_base=None,
|
||||
model="gemini-3.6-flash",
|
||||
client=self.mock_async_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
cached_content=None,
|
||||
custom_llm_provider="vertex_ai",
|
||||
vertex_project="test_project",
|
||||
vertex_location="us-central1",
|
||||
vertex_auth_header="test_token",
|
||||
)
|
||||
|
||||
messages, returned_params, returned_cache = result
|
||||
assert messages == all_messages
|
||||
assert returned_cache is None
|
||||
assert "tools" in returned_params
|
||||
self.mock_async_client.get.assert_not_called()
|
||||
self.mock_async_client.post.assert_not_called()
|
||||
|
||||
|
||||
def test_cached_messages_end_on_supported_turn():
|
||||
from litellm.llms.vertex_ai.context_caching.transformation import (
|
||||
cached_messages_end_on_supported_turn,
|
||||
)
|
||||
|
||||
assert (
|
||||
cached_messages_end_on_supported_turn(
|
||||
[{"role": "assistant", "content": "hi"}, {"role": "user", "content": "hello"}]
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert cached_messages_end_on_supported_turn([{"role": "system", "content": "be brief"}]) is True
|
||||
assert cached_messages_end_on_supported_turn([{"role": "assistant", "content": "hi"}]) is False
|
||||
assert (
|
||||
cached_messages_end_on_supported_turn(
|
||||
[
|
||||
{"role": "user", "content": "hello"},
|
||||
{"role": "assistant", "content": "hi"},
|
||||
{"role": "system", "content": "be brief"},
|
||||
]
|
||||
)
|
||||
is False
|
||||
)
|
||||
assert (
|
||||
cached_messages_end_on_supported_turn(
|
||||
[{"role": "system", "content": "be brief"}, {"role": "user", "content": "hello"}]
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert (
|
||||
cached_messages_end_on_supported_turn([{"role": "tool", "tool_call_id": "x", "content": "y"}])
|
||||
is False
|
||||
)
|
||||
assert (
|
||||
cached_messages_end_on_supported_turn([{"role": "function", "name": "f", "content": "y"}])
|
||||
is False
|
||||
)
|
||||
assert cached_messages_end_on_supported_turn([]) is False
|
||||
|
||||
|
||||
class TestCheckCachePagination:
|
||||
"""Test pagination logic in check_cache and async_check_cache methods."""
|
||||
|
|
|
|||
|
|
@ -31,10 +31,7 @@ class TestVertexAIFilesHandler:
|
|||
def test_extract_bucket_and_object_from_file_id_standard_path(self):
|
||||
"""Test extraction of bucket and object from URL-encoded file_id with standard path"""
|
||||
# Sample file_id with nested folder structure
|
||||
file_id = (
|
||||
"gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files"
|
||||
"%2Ftest-folder%2Fsub-folder%2Ftest-file.txt"
|
||||
)
|
||||
file_id = "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files%2Ftest-folder%2Fsub-folder%2Ftest-file.txt"
|
||||
|
||||
bucket_name, object_path = self.handler._extract_bucket_and_object_from_file_id(
|
||||
file_id=file_id,
|
||||
|
|
@ -105,21 +102,14 @@ class TestVertexAIFilesHandler:
|
|||
async def test_afile_content_success(self):
|
||||
"""Test successful async file content retrieval"""
|
||||
# Setup test data
|
||||
file_id = (
|
||||
"gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files"
|
||||
"%2Fuploads%2Fabc-test-file.txt"
|
||||
)
|
||||
file_id = "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-test-file.txt"
|
||||
expected_content = b"test file content"
|
||||
|
||||
file_content_request = FileContentRequest(
|
||||
file_id=file_id, extra_headers=None, extra_body=None
|
||||
)
|
||||
file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
|
||||
|
||||
# Mock the download_gcs_object method
|
||||
with (
|
||||
patch.object(
|
||||
self.handler, "download_gcs_object", new_callable=AsyncMock
|
||||
) as mock_download,
|
||||
patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download,
|
||||
patch.object(
|
||||
self.handler,
|
||||
"get_gcs_logging_config",
|
||||
|
|
@ -148,15 +138,9 @@ class TestVertexAIFilesHandler:
|
|||
# Verify the download was called with correct parameters
|
||||
mock_download.assert_called_once()
|
||||
call_args = mock_download.call_args
|
||||
assert (
|
||||
call_args.kwargs["object_name"]
|
||||
== "litellm-vertex-files/uploads/abc-test-file.txt"
|
||||
)
|
||||
assert call_args.kwargs["object_name"] == "litellm-vertex-files/uploads/abc-test-file.txt"
|
||||
assert "standard_callback_dynamic_params" in call_args.kwargs
|
||||
assert (
|
||||
call_args.kwargs["standard_callback_dynamic_params"]["gcs_bucket_name"]
|
||||
== "test-bucket"
|
||||
)
|
||||
assert call_args.kwargs["standard_callback_dynamic_params"]["gcs_bucket_name"] == "test-bucket"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_afile_content_missing_file_id(self):
|
||||
|
|
@ -164,9 +148,7 @@ class TestVertexAIFilesHandler:
|
|||
file_content_request = FileContentRequest(extra_headers=None, extra_body=None)
|
||||
|
||||
# Should raise ValueError for missing file_id
|
||||
with pytest.raises(
|
||||
ValueError, match="file_id is required in file_content_request"
|
||||
):
|
||||
with pytest.raises(ValueError, match="file_id is required in file_content_request"):
|
||||
await self.handler.afile_content(
|
||||
file_content_request=file_content_request,
|
||||
vertex_credentials=None,
|
||||
|
|
@ -179,20 +161,13 @@ class TestVertexAIFilesHandler:
|
|||
@pytest.mark.asyncio
|
||||
async def test_afile_content_download_failure(self):
|
||||
"""Test async file content retrieval when download fails"""
|
||||
file_id = (
|
||||
"gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files"
|
||||
"%2Fuploads%2Fabc-test-file.txt"
|
||||
)
|
||||
file_id = "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-test-file.txt"
|
||||
|
||||
file_content_request = FileContentRequest(
|
||||
file_id=file_id, extra_headers=None, extra_body=None
|
||||
)
|
||||
file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
|
||||
|
||||
# Mock download to return None (failure)
|
||||
with (
|
||||
patch.object(
|
||||
self.handler, "download_gcs_object", new_callable=AsyncMock
|
||||
) as mock_download,
|
||||
patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download,
|
||||
patch.object(
|
||||
self.handler,
|
||||
"get_gcs_logging_config",
|
||||
|
|
@ -216,14 +191,130 @@ class TestVertexAIFilesHandler:
|
|||
max_retries=3,
|
||||
)
|
||||
|
||||
def test_resolve_read_gcs_config_prefers_per_model_bucket(self, monkeypatch):
|
||||
monkeypatch.setenv("GCS_BUCKET_NAME", "env-default-bucket")
|
||||
monkeypatch.setenv("GCS_PATH_SERVICE_ACCOUNT", "/env/sa.json")
|
||||
|
||||
bucket, service_account = self.handler._resolve_read_gcs_config(
|
||||
litellm_params={
|
||||
"gcs_bucket_name": "my-model-bucket",
|
||||
"vertex_credentials": "/model/sa.json",
|
||||
},
|
||||
vertex_credentials=None,
|
||||
)
|
||||
|
||||
assert bucket == "my-model-bucket"
|
||||
assert service_account == "/model/sa.json"
|
||||
|
||||
def test_resolve_read_gcs_config_falls_back_to_env(self, monkeypatch):
|
||||
monkeypatch.setenv("GCS_BUCKET_NAME", "env-default-bucket")
|
||||
monkeypatch.setenv("GCS_PATH_SERVICE_ACCOUNT", "/env/sa.json")
|
||||
|
||||
bucket, service_account = self.handler._resolve_read_gcs_config(litellm_params={}, vertex_credentials=None)
|
||||
|
||||
assert bucket == "env-default-bucket"
|
||||
assert service_account == "/env/sa.json"
|
||||
|
||||
def test_resolve_read_gcs_config_serializes_dict_credentials(self, monkeypatch):
|
||||
monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False)
|
||||
|
||||
_, service_account = self.handler._resolve_read_gcs_config(
|
||||
litellm_params={"gcs_bucket_name": "my-model-bucket"},
|
||||
vertex_credentials={"type": "service_account", "project_id": "p"},
|
||||
)
|
||||
|
||||
assert service_account == '{"type": "service_account", "project_id": "p"}'
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_afile_content_honors_per_model_bucket_over_env(self, monkeypatch):
|
||||
"""
|
||||
Regression for #32640: a batch output written to a per-model gcs_bucket_name must be
|
||||
readable even when the global GCS_BUCKET_NAME points at a different bucket. Before the
|
||||
fix the read path resolved the bucket from env only and raised
|
||||
"file_id bucket does not match the configured storage bucket".
|
||||
"""
|
||||
monkeypatch.setenv("GCS_BUCKET_NAME", "env-default-bucket")
|
||||
monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False)
|
||||
|
||||
file_id = "gs%3A%2F%2Fmy-model-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-batch-output.jsonl"
|
||||
file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
|
||||
|
||||
with (
|
||||
patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download,
|
||||
patch.object(
|
||||
self.handler,
|
||||
"get_or_create_vertex_instance",
|
||||
new_callable=AsyncMock,
|
||||
return_value=object(),
|
||||
),
|
||||
):
|
||||
mock_download.return_value = b"batch output"
|
||||
|
||||
result = await self.handler.afile_content(
|
||||
file_content_request=file_content_request,
|
||||
vertex_credentials="/model/sa.json",
|
||||
vertex_project="test-project",
|
||||
vertex_location="us-central1",
|
||||
timeout=60.0,
|
||||
max_retries=0,
|
||||
litellm_params={
|
||||
"gcs_bucket_name": "my-model-bucket",
|
||||
"vertex_credentials": "/model/sa.json",
|
||||
},
|
||||
)
|
||||
|
||||
assert isinstance(result, HttpxBinaryResponseContent)
|
||||
assert result.response.content == b"batch output"
|
||||
|
||||
dynamic_params = mock_download.call_args.kwargs["standard_callback_dynamic_params"]
|
||||
assert dynamic_params["gcs_bucket_name"] == "my-model-bucket"
|
||||
assert dynamic_params["gcs_path_service_account"] == "/model/sa.json"
|
||||
assert mock_download.call_args.kwargs["object_name"] == "litellm-vertex-files/uploads/abc-batch-output.jsonl"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_afile_content_reads_without_global_env_bucket(self, monkeypatch):
|
||||
"""
|
||||
Regression for #32640: with no global GCS_BUCKET_NAME set, a model-group-level
|
||||
deployment (per-model gcs_bucket_name) must still be readable. Before the fix the read
|
||||
path raised "GCS_BUCKET_NAME is not set in the environment".
|
||||
"""
|
||||
monkeypatch.delenv("GCS_BUCKET_NAME", raising=False)
|
||||
monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False)
|
||||
|
||||
file_id = "gs%3A%2F%2Fmy-model-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-batch-output.jsonl"
|
||||
file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
|
||||
|
||||
with (
|
||||
patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download,
|
||||
patch.object(
|
||||
self.handler,
|
||||
"get_or_create_vertex_instance",
|
||||
new_callable=AsyncMock,
|
||||
return_value=object(),
|
||||
),
|
||||
):
|
||||
mock_download.return_value = b"batch output"
|
||||
|
||||
result = await self.handler.afile_content(
|
||||
file_content_request=file_content_request,
|
||||
vertex_credentials="/model/sa.json",
|
||||
vertex_project="test-project",
|
||||
vertex_location="us-central1",
|
||||
timeout=60.0,
|
||||
max_retries=0,
|
||||
litellm_params={"gcs_bucket_name": "my-model-bucket"},
|
||||
)
|
||||
|
||||
assert isinstance(result, HttpxBinaryResponseContent)
|
||||
dynamic_params = mock_download.call_args.kwargs["standard_callback_dynamic_params"]
|
||||
assert dynamic_params["gcs_bucket_name"] == "my-model-bucket"
|
||||
|
||||
def test_file_content_sync_success(self):
|
||||
"""Test successful sync file content retrieval"""
|
||||
file_id = "gs%3A%2F%2Ftest-bucket%2Ftest-file.txt"
|
||||
expected_content = b"test file content"
|
||||
|
||||
file_content_request = FileContentRequest(
|
||||
file_id=file_id, extra_headers=None, extra_body=None
|
||||
)
|
||||
file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
|
||||
|
||||
# Create expected response
|
||||
mock_response = httpx.Response(
|
||||
|
|
@ -261,25 +352,17 @@ class TestVertexAIFilesHandler:
|
|||
file_id = "gs%3A%2F%2Ftest-bucket%2Ftest-file.txt"
|
||||
expected_content = b"test file content"
|
||||
|
||||
file_content_request = FileContentRequest(
|
||||
file_id=file_id, extra_headers=None, extra_body=None
|
||||
)
|
||||
file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
|
||||
|
||||
# Mock the afile_content method
|
||||
with patch.object(
|
||||
self.handler, "afile_content", new_callable=AsyncMock
|
||||
) as mock_afile_content:
|
||||
with patch.object(self.handler, "afile_content", new_callable=AsyncMock) as mock_afile_content:
|
||||
mock_response = httpx.Response(
|
||||
status_code=200,
|
||||
content=expected_content,
|
||||
headers={"content-type": "application/octet-stream"},
|
||||
request=httpx.Request(
|
||||
method="GET", url="gs://test-bucket/test-file.txt"
|
||||
),
|
||||
)
|
||||
mock_afile_content.return_value = HttpxBinaryResponseContent(
|
||||
response=mock_response
|
||||
request=httpx.Request(method="GET", url="gs://test-bucket/test-file.txt"),
|
||||
)
|
||||
mock_afile_content.return_value = HttpxBinaryResponseContent(response=mock_response)
|
||||
|
||||
# Call the method with _is_async=True
|
||||
result = self.handler.file_content(
|
||||
|
|
|
|||
|
|
@ -2276,82 +2276,8 @@ def test_is_gemini_3_or_newer():
|
|||
assert VertexGeminiConfig._is_gemini_3_or_newer("") == False
|
||||
|
||||
|
||||
def test_forward_gemini_function_call_id_vertex_vs_google_ai_studio():
|
||||
"""Vertex AI rejects `id` on function_call/function_response; Google AI Studio accepts it on Gemini 3.5+."""
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
VertexGeminiConfig,
|
||||
)
|
||||
|
||||
model = "gemini-3.5-flash"
|
||||
assert (
|
||||
VertexGeminiConfig._forward_gemini_function_call_id(model, "vertex_ai") is False
|
||||
)
|
||||
assert (
|
||||
VertexGeminiConfig._forward_gemini_function_call_id(model, "vertex_ai_beta")
|
||||
is False
|
||||
)
|
||||
assert VertexGeminiConfig._forward_gemini_function_call_id(model, "gemini") is True
|
||||
assert VertexGeminiConfig._forward_gemini_function_call_id(model, None) is False
|
||||
assert (
|
||||
VertexGeminiConfig._forward_gemini_function_call_id(
|
||||
"gemini-2.5-flash", "gemini"
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
|
||||
def test_vertex_ai_gemini_35_tool_calls_omit_function_call_id():
|
||||
"""Regression: Vertex must not send OpenAI tool_call id inside Gemini function_call parts."""
|
||||
from litellm.llms.vertex_ai.gemini.transformation import (
|
||||
_gemini_convert_messages_with_history,
|
||||
)
|
||||
|
||||
messages = [
|
||||
{"role": "user", "content": "Explore this directory"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_50e7e0fe0989464a89f188eda443",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "read",
|
||||
"arguments": '{"filePath": "/tmp"}',
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_50e7e0fe0989464a89f188eda443",
|
||||
"content": "ok",
|
||||
},
|
||||
]
|
||||
|
||||
contents = _gemini_convert_messages_with_history(
|
||||
messages=messages,
|
||||
model="gemini-3.5-flash",
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
|
||||
for content in contents:
|
||||
for part in content.get("parts", []):
|
||||
fc = part.get("function_call")
|
||||
if fc is not None:
|
||||
assert "id" not in fc, f"Vertex payload must not include id: {fc}"
|
||||
fr = part.get("function_response")
|
||||
if fr is not None:
|
||||
assert "id" not in fr, f"Vertex payload must not include id: {fr}"
|
||||
|
||||
|
||||
def test_google_ai_studio_gemini_35_tool_calls_include_function_call_id():
|
||||
from litellm.llms.vertex_ai.gemini.transformation import (
|
||||
_gemini_convert_messages_with_history,
|
||||
)
|
||||
|
||||
tool_call_id = "call_50e7e0fe0989464a89f188eda443"
|
||||
messages = [
|
||||
def _tool_call_messages(tool_call_id: str):
|
||||
return [
|
||||
{"role": "user", "content": "hi"},
|
||||
{
|
||||
"role": "assistant",
|
||||
|
|
@ -2374,12 +2300,8 @@ def test_google_ai_studio_gemini_35_tool_calls_include_function_call_id():
|
|||
},
|
||||
]
|
||||
|
||||
contents = _gemini_convert_messages_with_history(
|
||||
messages=messages,
|
||||
model="gemini-3.5-flash",
|
||||
custom_llm_provider="gemini",
|
||||
)
|
||||
|
||||
def _collect_function_call_ids(contents):
|
||||
function_call_ids = []
|
||||
function_response_ids = []
|
||||
for content in contents:
|
||||
|
|
@ -2390,9 +2312,120 @@ def test_google_ai_studio_gemini_35_tool_calls_include_function_call_id():
|
|||
fr = part.get("function_response")
|
||||
if fr is not None:
|
||||
function_response_ids.append(fr.get("id"))
|
||||
return function_call_ids, function_response_ids
|
||||
|
||||
assert function_call_ids == [tool_call_id]
|
||||
assert function_response_ids == [tool_call_id]
|
||||
|
||||
def test_forward_gemini_function_call_id_is_gated_on_model_version_only():
|
||||
"""Gemini 3+ takes `id` on Vertex AI and Google AI Studio alike; older models reject it."""
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
VertexGeminiConfig,
|
||||
)
|
||||
|
||||
assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-3.5-flash") is True
|
||||
assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-3-pro") is True
|
||||
assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-2.5-flash") is False
|
||||
assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-2.0-flash") is False
|
||||
|
||||
|
||||
@pytest.mark.parametrize("custom_llm_provider", ["vertex_ai", "vertex_ai_beta", "gemini"])
|
||||
def test_gemini_35_tool_calls_include_function_call_id(custom_llm_provider):
|
||||
"""Vertex AI accepts `id` on Gemini 3+, so it must be sent there and not just on AI Studio.
|
||||
|
||||
Both parts are asserted together: Vertex pairs a result to its call by id, so emitting one
|
||||
side without the other would break strict tool-call matching.
|
||||
"""
|
||||
from litellm.llms.vertex_ai.gemini.transformation import (
|
||||
_gemini_convert_messages_with_history,
|
||||
)
|
||||
|
||||
tool_call_id = "call_50e7e0fe0989464a89f188eda443"
|
||||
contents = _gemini_convert_messages_with_history(
|
||||
messages=_tool_call_messages(tool_call_id),
|
||||
model="gemini-3.5-flash",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
assert _collect_function_call_ids(contents) == ([tool_call_id], [tool_call_id])
|
||||
|
||||
|
||||
@pytest.mark.parametrize("custom_llm_provider", ["vertex_ai", "gemini"])
|
||||
def test_gemini_25_tool_calls_omit_function_call_id(custom_llm_provider):
|
||||
"""Regression: models older than Gemini 3 reject `id`, so the key must be absent entirely."""
|
||||
from litellm.llms.vertex_ai.gemini.transformation import (
|
||||
_gemini_convert_messages_with_history,
|
||||
)
|
||||
|
||||
contents = _gemini_convert_messages_with_history(
|
||||
messages=_tool_call_messages("call_50e7e0fe0989464a89f188eda443"),
|
||||
model="gemini-2.5-flash",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
for content in contents:
|
||||
for part in content.get("parts", []):
|
||||
fc = part.get("function_call")
|
||||
if fc is not None:
|
||||
assert "id" not in fc, f"gemini-2.5 payload must not include id: {fc}"
|
||||
fr = part.get("function_response")
|
||||
if fr is not None:
|
||||
assert "id" not in fr, f"gemini-2.5 payload must not include id: {fr}"
|
||||
|
||||
|
||||
def test_vertex_ai_forwarded_function_call_id_strips_thought_signature_suffix():
|
||||
"""The thought signature rides along on the OpenAI id but must not reach Vertex.
|
||||
|
||||
Vertex now sees this code path for the first time, so the suffix has to be stripped here too.
|
||||
"""
|
||||
from litellm.llms.vertex_ai.gemini.transformation import (
|
||||
_gemini_convert_messages_with_history,
|
||||
)
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import (
|
||||
THOUGHT_SIGNATURE_SEPARATOR,
|
||||
)
|
||||
|
||||
bare_id = "call_50e7e0fe0989464a89f188eda443"
|
||||
contents = _gemini_convert_messages_with_history(
|
||||
messages=_tool_call_messages(f"{bare_id}{THOUGHT_SIGNATURE_SEPARATOR}sig123"),
|
||||
model="gemini-3.5-flash",
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
|
||||
_, function_response_ids = _collect_function_call_ids(contents)
|
||||
assert function_response_ids == [bare_id]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", ["gemini-3.5-flash", "gemini-2.5-flash"])
|
||||
def test_tool_response_without_matching_tool_call_is_rejected(model):
|
||||
"""An unpairable tool result must raise, not ship a functionResponse with no matching call."""
|
||||
from litellm.llms.vertex_ai.gemini.transformation import (
|
||||
_gemini_convert_messages_with_history,
|
||||
)
|
||||
|
||||
messages = [
|
||||
{"role": "user", "content": "hi"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_50e7e0fe0989464a89f188eda443",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "read",
|
||||
"arguments": '{"filePath": "/tmp"}',
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "tool", "content": "ok"},
|
||||
]
|
||||
|
||||
with pytest.raises(Exception, match="Missing corresponding tool call"):
|
||||
_gemini_convert_messages_with_history(
|
||||
messages=messages,
|
||||
model=model,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
|
||||
|
||||
def test_reasoning_effort_maps_to_thinking_level_gemini_3():
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue