chore: merge upstream main into reasoning history support

This commit is contained in:
jibanez-staticduo 2026-09-30 23:53:45 +02:00
commit c7837bea6f
No known key found for this signature in database
334 changed files with 40680 additions and 1011 deletions

View file

@ -148,7 +148,10 @@ legacy_paths() {
echo tests/unit/proxy/test_proxy_server.py ;;
proxy-db-proxy-utils) echo tests/unit/proxy/test_proxy_utils.py ;;
proxy-extras) echo tests/unit/litellm_proxy_extras ;;
proxy-infra) echo tests/unit/gateway ;;
proxy-infra)
echo tests/unit/gateway
echo tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py
echo tests/unit/proxy/roi_calculator ;;
responses-caching-types)
find tests/unit/responses -name 'test_*.py' -not -path 'tests/unit/responses/mcp/*'
echo tests/unit/types ;;

Binary file not shown.

After

Width:  |  Height:  |  Size: 80 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 58 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 63 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 70 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 47 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 76 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 75 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 39 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 70 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 93 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 73 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 39 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 81 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 72 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 76 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 50 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 59 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 57 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 56 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 65 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 55 KiB

View file

@ -79,7 +79,9 @@ jobs:
- shard: integrations
artifact-name: integrations
test-path: ""
test-path: >-
tests/test_litellm/integrations
tests/test_litellm/tracing
unit-flag: integrations
workers: 2
reruns: 3

View file

@ -99,7 +99,7 @@
"limit": 0
},
"reportUnknownArgumentType": {
"limit": 44358
"limit": 44802
},
"reportUnknownLambdaType": {
"limit": 109

View file

@ -1274,6 +1274,18 @@ version = "0.4.33"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6e8ccc4ea9f6acc32d102c0f6d471d11d913ad15f20c04de743374861fa1d414"
[[package]]
name = "const-hex"
version = "1.19.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "0e59eef12462b0f9b0a3620219be5d639afd79fe39dff0a42c3997061f9298b4"
dependencies = [
"cfg-if",
"cpufeatures 0.2.17",
"proptest",
"serde_core",
]
[[package]]
name = "const-oid"
version = "0.9.6"
@ -2372,9 +2384,9 @@ dependencies = [
"http-body-util",
"hyper 1.10.1",
"lazy_static",
"opentelemetry",
"opentelemetry 0.32.0",
"opentelemetry-semantic-conventions",
"opentelemetry_sdk",
"opentelemetry_sdk 0.32.1",
"percent-encoding",
"pin-project",
"prost",
@ -4075,6 +4087,7 @@ dependencies = [
"litellm-secrets-aws",
"litellm-secrets-types",
"litellm-token-counter",
"litellm-traces",
"litellm-tracing",
"pyo3",
"pyo3-async-runtimes",
@ -4351,6 +4364,25 @@ dependencies = [
"tiktoken-rs",
]
[[package]]
name = "litellm-traces"
version = "0.1.0"
dependencies = [
"base64 0.22.1",
"flate2",
"litellm-http",
"opentelemetry-proto",
"prost",
"rstest",
"serde",
"serde_json",
"testcontainers-modules",
"thiserror 2.0.19",
"time",
"tokio",
"url",
]
[[package]]
name = "litellm-tracing"
version = "0.1.0"
@ -4760,6 +4792,33 @@ dependencies = [
"tracing",
]
[[package]]
name = "opentelemetry"
version = "0.33.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6cdb0b1b267eb9db3331b434ed9ddab10d50e280a9adf9d13e5233e2002b61b5"
dependencies = [
"futures-core",
"futures-sink",
"js-sys",
"pin-project-lite",
"thiserror 2.0.19",
]
[[package]]
name = "opentelemetry-proto"
version = "0.33.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "25da1ac11a0aeccf38d7f77ee0348715adaf8340f65ad46c94a02c6b20e2f65d"
dependencies = [
"base64 0.22.1",
"const-hex",
"opentelemetry 0.33.0",
"opentelemetry_sdk 0.33.0",
"prost",
"serde",
]
[[package]]
name = "opentelemetry-semantic-conventions"
version = "0.32.1"
@ -4775,7 +4834,23 @@ dependencies = [
"futures-channel",
"futures-executor",
"futures-util",
"opentelemetry",
"opentelemetry 0.32.0",
"percent-encoding",
"portable-atomic",
"rand 0.9.5",
"thiserror 2.0.19",
]
[[package]]
name = "opentelemetry_sdk"
version = "0.33.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "cb39533d9d1c912123efd7d41d7e0c29d16917b60ce15b4c8d87cb1af7f67520"
dependencies = [
"futures-channel",
"futures-executor",
"futures-util",
"opentelemetry 0.33.0",
"percent-encoding",
"portable-atomic",
"rand 0.9.5",
@ -5704,6 +5779,7 @@ checksum = "16a1cfa75cc186dd73d5818e510e042e40927bccc9c236b061cea97e1eb08029"
dependencies = [
"base64 0.23.1",
"bytes",
"encoding_rs",
"futures-core",
"futures-util",
"h2 0.4.15",
@ -5715,6 +5791,7 @@ dependencies = [
"hyper-util",
"js-sys",
"log",
"mime",
"percent-encoding",
"pin-project-lite",
"quinn",
@ -6945,6 +7022,7 @@ dependencies = [
"memchr",
"parse-display",
"pin-project-lite",
"reqwest 0.13.5",
"serde",
"serde_json",
"serde_with",
@ -7505,7 +7583,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "adbc64cba7137545b8044cb1fe9814f7aacf3c6b5f9b45be8bb5db538befdb26"
dependencies = [
"js-sys",
"opentelemetry",
"opentelemetry 0.32.0",
"tracing",
"tracing-core",
"tracing-subscriber",

View file

@ -12,6 +12,7 @@ repository = "https://github.com/BerriAI/litellm"
litellm-config = { path = "crates/config" }
litellm-router = { path = "crates/router" }
litellm-tracing = { path = "crates/tracing" }
litellm-traces = { path = "crates/traces" }
litellm-core = { path = "crates/core" }
litellm-gateway-mcp = { path = "crates/gateway-mcp" }
litellm-gateway = { path = "crates/gateway" }

View file

@ -203,6 +203,85 @@ fn threshold_tiers_and_boundaries() {
assert_eq!(calculate(&specification, &flex).unwrap().input(), 600.0);
}
#[rstest]
#[case::ultrafast_above_threshold(ServiceTier::Ultrafast, 300_000, 9_301_000.0, 37_000.0)]
#[case::ultrafast_at_threshold(ServiceTier::Ultrafast, 272_000, 544_500.0, 5_000.0)]
#[case::standard_above_threshold(ServiceTier::Standard, 300_000, 3_300_600.0, 13_000.0)]
#[case::priority_above_threshold(ServiceTier::Priority, 300_000, 5_701_000.0, 23_000.0)]
fn tiered_long_context_rates_are_selected_by_service_tier(
#[case] service_tier: ServiceTier,
#[case] prompt_tokens: u64,
#[case] expected_input: f64,
#[case] expected_output: f64,
) {
let standard = Rates {
cache_read: Rate::Value(3.0),
..rates(Rate::Value(1.0), Rate::Value(2.0))
};
let tiers = [
TierRates {
tier: ServiceTier::Priority,
rates: Rates {
cache_read: Rate::Value(5.0),
..rates(Rate::Value(3.0), Rate::Value(4.0))
},
},
TierRates {
tier: ServiceTier::Ultrafast,
rates: Rates {
cache_read: Rate::Value(7.0),
..rates(Rate::Value(2.0), Rate::Value(5.0))
},
},
];
let threshold_tiers = [
TierRates {
tier: ServiceTier::Priority,
rates: Rates {
cache_read: Rate::Value(29.0),
..rates(Rate::Value(19.0), Rate::Value(23.0))
},
},
TierRates {
tier: ServiceTier::Ultrafast,
rates: Rates {
cache_read: Rate::Value(41.0),
..rates(Rate::Value(31.0), Rate::Value(37.0))
},
},
];
let thresholds = [ThresholdRates {
above_prompt_tokens: 272_000,
standard: Rates {
cache_read: Rate::Value(17.0),
..rates(Rate::Value(11.0), Rate::Value(13.0))
},
tiers: &threshold_tiers,
}];
let pricing = Pricing {
standard,
tiers: &tiers,
thresholds: &thresholds,
off_peak: None,
};
let base = request();
let long_context_request = Request {
usage: Usage {
prompt_tokens,
completion_tokens: 1_000,
cache_read_tokens: 100,
cache_write_tokens: 0,
..base.usage
},
service_tier,
..base
};
let cost = calculate(&pricing, &long_context_request).unwrap();
assert_eq!(cost.input(), expected_input);
assert_eq!(cost.output(), expected_output);
}
#[test]
fn compile_rejects_ambiguous_rates() {
let duplicate = ThresholdRates {

View file

@ -21,6 +21,7 @@ tiktoken = ["litellm-token-counter/tiktoken"]
[dependencies]
fancy-regex.workspace = true
litellm-tracing.workspace = true
litellm-traces.workspace = true
litellm-host.workspace = true
bytes.workspace = true
futures-util.workspace = true

View file

@ -1,5 +1,5 @@
use crate::cache::cache_error;
use crate::logger::run_sync_value;
use crate::execution::run_sync_value;
use litellm_cache_gcs::{DEFAULT_ENDPOINT, GcsConfig};
use litellm_cache_redis_semantic::RedisSemanticConfig;
use litellm_host_python::release_gil;

View file

@ -470,7 +470,7 @@ impl NativeResponseCache {
match self {
Self::Exact(_) | Self::QdrantSemantic(_) => {
let service = self.clone();
crate::logger::run_async(
crate::execution::run_async(
py,
async move {
service
@ -495,7 +495,7 @@ impl NativeResponseCache {
match self {
Self::Exact(_) | Self::QdrantSemantic(_) => {
let service = self.clone();
crate::logger::run_async(
crate::execution::run_async(
py,
async move { service.async_lookup(&request, now()).await },
cache_error,
@ -550,7 +550,7 @@ impl NativeResponseCache {
match self {
Self::Exact(_) | Self::QdrantSemantic(_) => {
let service = self.clone();
crate::logger::run_async(
crate::execution::run_async(
py,
async move { service.async_store(&request, response, now()).await },
cache_error,
@ -619,7 +619,7 @@ impl NativeResponseCache {
match self {
Self::Exact(_) | Self::QdrantSemantic(_) => {
let service = self.clone();
crate::logger::run_async(
crate::execution::run_async(
py,
async move { service.async_store_batch(entries, now()).await },
cache_error,

View file

@ -1,5 +1,5 @@
use crate::cache::cache_error;
use crate::logger::run_async;
use crate::execution::run_async;
use std::{collections::VecDeque, time::Duration};
use litellm_cache::Error;

View file

@ -144,7 +144,7 @@ impl NativeCacheHandle {
self.check_process()?;
let request = request(key, None)?;
let backend = self.backend.clone();
crate::logger::run_async(
crate::execution::run_async(
py,
async move { backend.async_lookup(&request, super::request::now()).await },
cache_error,
@ -163,7 +163,7 @@ impl NativeCacheHandle {
let request = request(key, ttl)?;
let value: Value = from_py(value)?;
let backend = self.backend.clone();
crate::logger::run_async(
crate::execution::run_async(
py,
async move {
backend
@ -188,7 +188,7 @@ impl NativeCacheHandle {
.map(|(key, value)| Ok((request(key, ttl)?, value)))
.collect::<PyResult<Vec<_>>>()?;
let backend = self.backend.clone();
crate::logger::run_async(
crate::execution::run_async(
py,
async move {
backend
@ -202,19 +202,19 @@ impl NativeCacheHandle {
fn flush(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
self.check_process()?;
let backend = self.backend.clone();
crate::logger::run_sync(py, async move { backend.async_flush().await }, cache_error)
crate::execution::run_sync(py, async move { backend.async_flush().await }, cache_error)
}
fn async_flush<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
self.check_process()?;
let backend = self.backend.clone();
crate::logger::run_async(py, async move { backend.async_flush().await }, cache_error)
crate::execution::run_async(py, async move { backend.async_flush().await }, cache_error)
}
fn ping<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
self.check_process()?;
let storage = self.storage.clone();
crate::logger::run_async(
crate::execution::run_async(
py,
async move {
match storage {
@ -229,7 +229,7 @@ impl NativeCacheHandle {
fn disconnect<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
self.check_process()?;
let storage = self.storage.clone();
crate::logger::run_async(
crate::execution::run_async(
py,
async move {
match storage {
@ -244,7 +244,7 @@ impl NativeCacheHandle {
fn delete<'py>(&self, py: Python<'py>, keys: Vec<String>) -> PyResult<Bound<'py, PyAny>> {
self.check_process()?;
let storage = self.storage.clone();
crate::logger::run_async(
crate::execution::run_async(
py,
async move {
for key in keys {

View file

@ -1,4 +1,4 @@
use crate::logger::run_async;
use crate::execution::run_async;
use litellm_cache_response::PartialHits;
use litellm_host_python::{ExecutionStep, from_py, release_gil, to_py};
use pyo3::{

View file

@ -13,7 +13,7 @@ where
E: Send + 'static,
F: Future<Output = Result<T, E>> + Send + 'static,
{
litellm_host_python::run_sync(py, super::capture(py).instrument(future), map_error)
litellm_host_python::run_sync(py, crate::logger::capture(py).instrument(future), map_error)
}
pub(crate) fn run_async<T, E, F>(
@ -26,7 +26,7 @@ where
E: Send + 'static,
F: Future<Output = Result<T, E>> + Send + 'static,
{
litellm_host_python::run_async(py, super::capture(py).instrument(future), map_error)
litellm_host_python::run_async(py, crate::logger::capture(py).instrument(future), map_error)
}
pub(crate) fn run_sync_value<T, F>(py: Python<'_>, future: F) -> PyResult<T>
@ -34,7 +34,7 @@ where
T: Send + 'static,
F: Future<Output = PyResult<T>> + Send + 'static,
{
litellm_host_python::run_sync_value(py, super::capture(py).instrument(future))
litellm_host_python::run_sync_value(py, crate::logger::capture(py).instrument(future))
}
pub(crate) fn run_async_value<T, F>(py: Python<'_>, future: F) -> PyResult<Bound<'_, PyAny>>
@ -42,5 +42,5 @@ where
T: for<'py> IntoPyObject<'py> + Send + 'static,
F: Future<Output = PyResult<T>> + Send + 'static,
{
litellm_host_python::run_async_value(py, super::capture(py).instrument(future))
litellm_host_python::run_async_value(py, crate::logger::capture(py).instrument(future))
}

View file

@ -4,6 +4,7 @@ mod coercion;
mod credentials;
mod diagnostics;
mod errors;
mod execution;
mod http;
mod lifecycle;
mod logger;
@ -42,6 +43,8 @@ mod _native {
use crate::routes::responses::{ResponsesWebSocketConnection, aresponses, responses};
#[pymodule_export]
use crate::routes::token_counter::TokenCounter;
#[pymodule_export]
use crate::routes::traces::{NativeTraceStorage, trace_decode_otlp};
#[cfg(feature = "huggingface")]
#[pymodule_export]
use crate::tokenizer::HuggingFaceEncoding;
@ -106,6 +109,8 @@ mod tests {
"aresponses",
"ResponsesWebSocketConnection",
"NativeDiagnosticProcessor",
"NativeTraceStorage",
"trace_decode_otlp",
"TokenCounter",
"Tokenizer",
"gil_stats",

View file

@ -1,7 +1,5 @@
mod execution;
mod machine;
pub(crate) use execution::{run_async, run_async_value, run_sync, run_sync_value};
pub(crate) use machine::LoggedMachine;
use litellm_host_python::Pythonized;

View file

@ -76,7 +76,7 @@ async fn traced_operation(_secret: &str) -> PyResult<()> {
#[pyfunction]
fn span_warning(py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
super::run_async_value(py, traced_operation("private-key-sentinel"))
crate::execution::run_async_value(py, traced_operation("private-key-sentinel"))
}
#[pyfunction]
@ -93,7 +93,7 @@ fn levels(py: Python<'_>) {
#[pyfunction]
fn asynchronous_warning(py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
super::run_async_value(py, async {
crate::execution::run_async_value(py, async {
tokio::task::yield_now().await;
litellm_tracing::warn!("async warning");
Ok(())
@ -102,7 +102,7 @@ fn asynchronous_warning(py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
#[pyfunction]
fn synchronous_warning(py: Python<'_>) -> PyResult<()> {
super::run_sync_value(py, async {
crate::execution::run_sync_value(py, async {
tokio::task::yield_now().await;
litellm_tracing::warn!("sync warning");
Ok(())
@ -111,7 +111,7 @@ fn synchronous_warning(py: Python<'_>) -> PyResult<()> {
#[pyfunction]
fn synchronous_failure(py: Python<'_>) -> PyResult<()> {
super::run_sync_value(py, async {
crate::execution::run_sync_value(py, async {
litellm_tracing::warn!("failure diagnostic");
Err(pyo3::exceptions::PyValueError::new_err("request failed"))
})

View file

@ -1,4 +1,4 @@
use crate::logger::{run_async, run_sync};
use crate::execution::{run_async, run_sync};
use litellm_core::audio_transcription::{
AudioTranscriptionRoute, Error, types::AudioTranscriptionRequest,
};

View file

@ -2,7 +2,7 @@ mod host;
use pyo3::types::{PyDict, PyTuple};
use crate::logger::{run_async, run_sync};
use crate::execution::{run_async, run_sync};
use litellm_core::chat_completions::{ChatCompletionsRoute, Error, types::ChatCompletionsRequest};
use litellm_llms_types::formats::chat_completions::ChatCompletionsResponse;
use pyo3::prelude::*;

View file

@ -6,6 +6,7 @@ pub(crate) mod messages;
pub(crate) mod ocr;
pub(crate) mod responses;
pub(crate) mod token_counter;
pub(crate) mod traces;
use litellm_callbacks_legacy_python::LoggingOperation;
use litellm_callbacks_legacy_python::{LegacyLogging, PublicCall};

View file

@ -142,7 +142,7 @@ impl ResponsesWebSocketConnection {
) -> PyResult<Bound<'py, PyAny>> {
let headers = marshal_headers(headers)?;
let timeout = optional_timeout(timeout_seconds);
crate::logger::run_async_value(py, async move {
crate::execution::run_async_value(py, async move {
let inner = RustResponsesWebSocketConnection::connect_url(&url, &headers, timeout)
.await
.map_err(route_error_to_pyerr)?;
@ -152,21 +152,21 @@ impl ResponsesWebSocketConnection {
fn send_text<'py>(&self, py: Python<'py>, text: String) -> PyResult<Bound<'py, PyAny>> {
let inner = self.inner.clone();
crate::logger::run_async_value(py, async move {
crate::execution::run_async_value(py, async move {
inner.send_text(text).await.map_err(route_error_to_pyerr)
})
}
fn recv_text<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
let inner = self.inner.clone();
crate::logger::run_async_value(py, async move {
crate::execution::run_async_value(py, async move {
inner.recv_text().await.map_err(route_error_to_pyerr)
})
}
fn close<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
let inner = self.inner.clone();
crate::logger::run_async_value(py, async move {
crate::execution::run_async_value(py, async move {
inner.close().await.map_err(route_error_to_pyerr)
})
}

View file

@ -1,4 +1,4 @@
use crate::logger::run_async;
use crate::execution::run_async;
use std::sync::Arc;
use std::{num::NonZero, thread::available_parallelism};

View file

@ -0,0 +1,137 @@
use std::collections::BTreeMap;
use litellm_http::ClientVariant;
use litellm_traces::{Connection, Error, InsertTable, Parameter};
use pyo3::{
exceptions::{PyOverflowError, PyRuntimeError, PyValueError},
prelude::*,
};
fn map_error(error: Error) -> PyErr {
match error {
Error::InvalidRow | Error::InvalidTable | Error::InvalidSchema | Error::EmptySql => {
PyValueError::new_err(error.to_string())
}
Error::InsertTooLarge => PyOverflowError::new_err(error.to_string()),
Error::InvalidUrl
| Error::QueryFailed(_)
| Error::InsertFailed(_)
| Error::SchemaFailed(_)
| Error::ResponseTooLarge
| Error::InvalidResponse
| Error::Transport => PyRuntimeError::new_err(error.to_string()),
}
}
#[pyclass]
pub struct NativeTraceStorage {
database: String,
writer: Connection,
reader: Option<Connection>,
}
#[pymethods]
impl NativeTraceStorage {
#[new]
fn new(database: String, url: &str, reader_url: Option<&str>) -> PyResult<Self> {
litellm_traces::schema_statements(&database, 1, 1).map_err(map_error)?;
Ok(Self {
writer: Connection::writer(url).map_err(map_error)?,
reader: reader_url
.map(|value| Connection::reader(value, &database))
.transpose()
.map_err(map_error)?,
database,
})
}
fn ensure_schema<'py>(
&self,
py: Python<'py>,
trace_retention_days: u32,
spend_log_retention_days: u32,
) -> PyResult<Bound<'py, PyAny>> {
let client = crate::http::host_client(py, ClientVariant::NoRedirect)?;
let connection = self.writer.clone();
let database = self.database.clone();
crate::execution::run_async(
py,
async move {
litellm_traces::ensure_schema(
&client,
&connection,
&database,
trace_retention_days,
spend_log_retention_days,
)
.await
},
map_error,
)
}
fn insert_rows<'py>(
&self,
py: Python<'py>,
table: &str,
#[pyo3(from_py_with = litellm_host_python::from_py_argument)] rows: Vec<
BTreeMap<String, serde_json::Value>,
>,
) -> PyResult<Bound<'py, PyAny>> {
let table = InsertTable::parse(table).map_err(map_error)?;
let client = crate::http::host_client(py, ClientVariant::NoRedirect)?;
let connection = self.writer.clone();
let database = self.database.clone();
crate::execution::run_async(
py,
async move {
litellm_traces::insert_rows(&client, &connection, &database, table, rows).await
},
map_error,
)
}
fn query<'py>(
&self,
py: Python<'py>,
sql: String,
#[pyo3(from_py_with = litellm_host_python::from_py_argument)] parameters: BTreeMap<
String,
Parameter,
>,
) -> PyResult<Bound<'py, PyAny>> {
let connection = self.reader.clone().ok_or_else(|| {
PyRuntimeError::new_err("Trace reads require a separate ClickHouse reader URL")
})?;
let client = crate::http::host_client(py, ClientVariant::NoRedirect)?;
crate::execution::run_async(
py,
async move { litellm_traces::execute_read(&client, &connection, &sql, &parameters).await },
map_error,
)
}
}
#[pyfunction]
pub fn trace_decode_otlp<'py>(
py: Python<'py>,
body: &[u8],
content_type: Option<&str>,
content_encoding: Option<&str>,
max_decompressed_bytes: usize,
) -> PyResult<Bound<'py, PyAny>> {
let spans = py
.detach(|| {
litellm_traces::decode_otlp(
body,
content_type,
content_encoding,
max_decompressed_bytes,
)
})
.map_err(|error| match error {
litellm_traces::DecodeError::TooLarge => PyOverflowError::new_err(error.to_string()),
_ => PyValueError::new_err(error.to_string()),
})?;
litellm_host_python::Pythonized(spans).into_pyobject(py)
}

View file

@ -1,7 +1,7 @@
use std::{collections::BTreeMap, sync::Arc};
use litellm_core_utils::settings::{Lookup, ProcessEnvironment};
use litellm_host_python::{from_py, json_object_field, run_async_value, run_sync_value, to_py};
use litellm_host_python::{from_py, json_object_field, to_py};
use litellm_secrets::{
KeyManagementSettings, KeyManagementSystem, Secret, SecretManager, load_native_manager,
read_secret_from_python_manager,
@ -13,6 +13,8 @@ use pyo3::{
types::PyDict,
};
use crate::execution::{run_async_value, run_sync_value};
#[derive(Clone, PartialEq)]
struct Configuration {
system: KeyManagementSystem,

View file

@ -0,0 +1,7 @@
- Rust owns OTLP wire decoding, ClickHouse schema, row encoding, named reads, connection validation and transport
- Keep this crate independent of Python; PyO3 conversion and public Python exceptions belong in `python-bridge`
- Keep the SQL migrations here as the only ClickHouse schema definition
- Use typed query parameters and a dedicated SELECT-only reader with server-side limits
- Keep `config/reader.xml` grants on the database the schema is created in (CLICKHOUSE_DATABASE, default `litellm`)
- Bound insert time and encoded bytes; make retry deduplication behavior explicit for supported ClickHouse versions
- Test storage behavior through the crate's public API against ClickHouse

View file

@ -0,0 +1,24 @@
[package]
name = "litellm-traces"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
[dependencies]
base64.workspace = true
flate2.workspace = true
opentelemetry-proto = { version = "0.33.0", default-features = false, features = ["gen-tonic-messages", "trace", "with-serde"] }
prost = "0.14.4"
time = { workspace = true, features = ["formatting"] }
litellm-http.workspace = true
serde.workspace = true
serde_json.workspace = true
thiserror.workspace = true
url.workspace = true
[dev-dependencies]
litellm-http = { workspace = true, features = ["test-support"] }
rstest.workspace = true
testcontainers-modules = { version = "0.15.0", features = ["clickhouse"] }
tokio.workspace = true

View file

@ -0,0 +1,32 @@
<clickhouse>
<profiles>
<litellm_traces_reader>
<readonly>1</readonly>
<max_execution_time>10</max_execution_time>
<max_result_rows>1000</max_result_rows>
<max_result_bytes>4194304</max_result_bytes>
<result_overflow_mode>throw</result_overflow_mode>
<max_memory_usage>268435456</max_memory_usage>
<constraints>
<readonly><readonly/></readonly>
<max_execution_time><readonly/></max_execution_time>
<max_result_rows><readonly/></max_result_rows>
<max_result_bytes><readonly/></max_result_bytes>
<result_overflow_mode><readonly/></result_overflow_mode>
<max_memory_usage><readonly/></max_memory_usage>
</constraints>
</litellm_traces_reader>
</profiles>
<users>
<litellm_traces_reader>
<password from_env="LITELLM_TRACES_READER_PASSWORD"/>
<networks><ip>::/0</ip></networks>
<profile>litellm_traces_reader</profile>
<grants>
<query>GRANT SELECT ON litellm.otel_traces</query>
<query>GRANT SELECT ON litellm.agent_traces_by_key</query>
<query>GRANT SELECT ON litellm.spend_logs</query>
</grants>
</litellm_traces_reader>
</users>
</clickhouse>

View file

@ -0,0 +1,47 @@
CREATE TABLE IF NOT EXISTS {database}.otel_traces
(
Timestamp DateTime64(9) CODEC(Delta, ZSTD(1)),
TraceId String CODEC(ZSTD(1)),
SpanId String CODEC(ZSTD(1)),
ParentSpanId String CODEC(ZSTD(1)),
TraceState String CODEC(ZSTD(1)),
SpanName LowCardinality(String) CODEC(ZSTD(1)),
SpanKind LowCardinality(String) CODEC(ZSTD(1)),
ServiceName LowCardinality(String) CODEC(ZSTD(1)),
ResourceAttributes Map(LowCardinality(String), String) CODEC(ZSTD(1)),
ScopeName String CODEC(ZSTD(1)),
ScopeVersion String CODEC(ZSTD(1)),
SpanAttributes Map(LowCardinality(String), String) CODEC(ZSTD(1)),
Duration UInt64 CODEC(ZSTD(1)),
StatusCode LowCardinality(String) CODEC(ZSTD(1)),
StatusMessage String CODEC(ZSTD(1)),
`Events.Timestamp` Array(DateTime64(9)) CODEC(ZSTD(1)),
`Events.Name` Array(LowCardinality(String)) CODEC(ZSTD(1)),
`Events.Attributes` Array(Map(LowCardinality(String), String)) CODEC(ZSTD(1)),
`Links.TraceId` Array(String) CODEC(ZSTD(1)),
`Links.SpanId` Array(String) CODEC(ZSTD(1)),
`Links.TraceState` Array(String) CODEC(ZSTD(1)),
`Links.Attributes` Array(Map(LowCardinality(String), String)) CODEC(ZSTD(1)),
TeamId LowCardinality(String) DEFAULT ResourceAttributes['litellm.team_id'],
ApiKeyHash String DEFAULT ResourceAttributes['litellm.api_key_hash'],
ObservationType LowCardinality(String) DEFAULT multiIf(
ParentSpanId = '', 'agent',
SpanAttributes['gen_ai.operation.name'] = 'invoke_agent', 'agent',
SpanAttributes['gen_ai.operation.name'] IN ('chat', 'text_completion', 'generate_content'), 'llm',
SpanAttributes['gen_ai.operation.name'] = 'execute_tool', 'tool',
'chain'),
AgentName LowCardinality(String) DEFAULT SpanAttributes['gen_ai.agent.name'],
LiteLLMRequestId String DEFAULT SpanAttributes['gen_ai.response.id'],
Model LowCardinality(String) DEFAULT SpanAttributes['gen_ai.request.model'],
InputTokens UInt32 DEFAULT toUInt32OrZero(SpanAttributes['gen_ai.usage.input_tokens']),
OutputTokens UInt32 DEFAULT toUInt32OrZero(SpanAttributes['gen_ai.usage.output_tokens']),
Input String CODEC(ZSTD(3)),
Output String CODEC(ZSTD(3)),
InputPreview String DEFAULT substring(Input, 1, 240),
INDEX idx_trace_id TraceId TYPE bloom_filter(0.001) GRANULARITY 1,
INDEX idx_req_id LiteLLMRequestId TYPE bloom_filter(0.01) GRANULARITY 1
)
ENGINE = MergeTree
PARTITION BY toDate(Timestamp)
ORDER BY (TeamId, ServiceName, toDateTime(Timestamp), TraceId)
SETTINGS ttl_only_drop_parts = 1, non_replicated_deduplication_window = 1000

View file

@ -0,0 +1,25 @@
CREATE TABLE IF NOT EXISTS {database}.agent_traces_by_key
(
TeamId LowCardinality(String),
ApiKeyHash String,
TraceId String,
StartTs SimpleAggregateFunction(min, DateTime64(9)),
EndTs SimpleAggregateFunction(max, DateTime64(9)),
ServiceName SimpleAggregateFunction(any, LowCardinality(String)),
RootName SimpleAggregateFunction(anyLast, Nullable(String)),
RootInput SimpleAggregateFunction(anyLast, Nullable(String)),
RootStatus SimpleAggregateFunction(anyLast, Nullable(String)),
SpanCount SimpleAggregateFunction(sum, UInt64),
AgentCount SimpleAggregateFunction(sum, UInt64),
LlmCount SimpleAggregateFunction(sum, UInt64),
ToolCount SimpleAggregateFunction(sum, UInt64),
ErrorCount SimpleAggregateFunction(sum, UInt64),
InputTokens SimpleAggregateFunction(sum, UInt64),
OutputTokens SimpleAggregateFunction(sum, UInt64),
Models SimpleAggregateFunction(groupUniqArrayArray, Array(String)),
AgentNames SimpleAggregateFunction(groupUniqArrayArray, Array(String)),
RequestIds SimpleAggregateFunction(groupArrayArray, Array(String))
)
ENGINE = AggregatingMergeTree
ORDER BY (TeamId, ApiKeyHash, TraceId)
SETTINGS non_replicated_deduplication_window = 1000

View file

@ -0,0 +1,22 @@
CREATE MATERIALIZED VIEW IF NOT EXISTS {database}.agent_traces_by_key_mv
TO {database}.agent_traces_by_key AS
SELECT
TeamId, ApiKeyHash, TraceId,
min(Timestamp) AS StartTs,
max(Timestamp + toIntervalNanosecond(Duration)) AS EndTs,
any(ServiceName) AS ServiceName,
anyLastIf(toNullable(SpanName), ParentSpanId = '') AS RootName,
anyLastIf(toNullable(InputPreview), ParentSpanId = '') AS RootInput,
anyLastIf(toNullable(StatusCode), ParentSpanId = '') AS RootStatus,
count() AS SpanCount,
countIf(ObservationType = 'agent') AS AgentCount,
countIf(ObservationType = 'llm') AS LlmCount,
countIf(ObservationType = 'tool') AS ToolCount,
countIf(StatusCode = 'STATUS_CODE_ERROR') AS ErrorCount,
sum(InputTokens) AS InputTokens,
sum(OutputTokens) AS OutputTokens,
groupUniqArrayIf(toString(Model), Model != '') AS Models,
groupUniqArrayIf(SpanName, ObservationType = 'agent') AS AgentNames,
groupArrayIf(LiteLLMRequestId, LiteLLMRequestId != '') AS RequestIds
FROM {database}.otel_traces
GROUP BY TeamId, ApiKeyHash, TraceId

View file

@ -0,0 +1,42 @@
CREATE TABLE IF NOT EXISTS {database}.spend_logs
(
request_id String,
response_id String,
call_type LowCardinality(String),
api_key String,
key_alias String,
team_id LowCardinality(String),
team_alias String,
organization_id String,
user String,
end_user String,
model LowCardinality(String),
model_group LowCardinality(String),
model_id String,
custom_llm_provider LowCardinality(String),
api_base String,
spend Float64,
prompt_tokens UInt32,
completion_tokens UInt32,
total_tokens UInt32,
cache_read_tokens UInt32,
cache_write_tokens UInt32,
start_time DateTime64(3),
end_time DateTime64(3),
completion_start_time Nullable(DateTime64(3)),
status LowCardinality(String),
error_str String,
cache_hit Bool,
session_id String,
trace_id String,
span_id String,
request_tags Array(String),
metadata String CODEC(ZSTD(3)),
messages String CODEC(ZSTD(3)),
response String CODEC(ZSTD(3)),
INDEX idx_response_id response_id TYPE bloom_filter(0.001) GRANULARITY 1,
INDEX idx_trace_id trace_id TYPE bloom_filter(0.001) GRANULARITY 1
)
ENGINE = ReplacingMergeTree(end_time)
PARTITION BY toYYYYMM(start_time)
ORDER BY (team_id, start_time, request_id)

View file

@ -0,0 +1 @@
ALTER TABLE {database}.otel_traces MODIFY TTL toDateTime(Timestamp) + INTERVAL {trace_retention_days} DAY

View file

@ -0,0 +1 @@
ALTER TABLE {database}.agent_traces_by_key MODIFY TTL toDateTime(StartTs) + INTERVAL {trace_retention_days} DAY

View file

@ -0,0 +1 @@
ALTER TABLE {database}.spend_logs MODIFY TTL toDateTime(start_time) + INTERVAL {spend_log_retention_days} DAY

View file

@ -0,0 +1,35 @@
#[derive(Debug, thiserror::Error)]
pub enum Error {
#[error("invalid ClickHouse insert row")]
InvalidRow,
#[error("invalid ClickHouse insert table")]
InvalidTable,
#[error("invalid ClickHouse HTTP URL")]
InvalidUrl,
#[error("database must be a nonempty SQL identifier and retention must be positive")]
InvalidSchema,
#[error("SQL query must not be empty")]
EmptySql,
#[error("ClickHouse query failed with HTTP status {0}")]
QueryFailed(u16),
#[error("ClickHouse insert failed with HTTP status {0}")]
InsertFailed(u16),
#[error("ClickHouse insert exceeds the encoded size limit")]
InsertTooLarge,
#[error("ClickHouse schema setup failed with HTTP status {0}")]
SchemaFailed(u16),
#[error("ClickHouse query exceeded the response size limit")]
ResponseTooLarge,
#[error("ClickHouse returned an invalid or failed JSON query response")]
InvalidResponse,
#[error("ClickHouse query transport failed")]
Transport,
}
#[derive(Debug, thiserror::Error)]
pub enum DecodeError {
#[error("invalid OTLP trace payload")]
InvalidPayload,
#[error("OTLP trace payload exceeds the decompressed size limit")]
TooLarge,
}

View file

@ -0,0 +1,151 @@
use std::{collections::BTreeMap, io::Write, time::Duration};
use flate2::{Compression, write::GzEncoder};
use litellm_http::Client;
use serde_json::Value;
use time::{OffsetDateTime, format_description::well_known::Rfc3339};
use crate::{Connection, Error};
const MAX_INSERT_BYTES: usize = 64 * 1024 * 1024;
const INSERT_TIMEOUT: Duration = Duration::from_secs(30);
pub enum InsertTable {
OtelTraces,
SpendLogs,
}
impl InsertTable {
pub fn parse(value: &str) -> Result<Self, Error> {
match value {
"otel_traces" => Ok(Self::OtelTraces),
"spend_logs" => Ok(Self::SpendLogs),
_ => Err(Error::InvalidTable),
}
}
fn name(&self) -> &'static str {
match self {
Self::OtelTraces => "otel_traces",
Self::SpendLogs => "spend_logs",
}
}
}
pub async fn insert_rows(
client: &Client,
connection: &Connection,
database: &str,
table: InsertTable,
rows: Vec<BTreeMap<String, Value>>,
) -> Result<(), Error> {
if rows.is_empty() {
return Ok(());
}
let encoded = encode_rows_with_limit(rows, MAX_INSERT_BYTES)?;
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
encoder
.write_all(encoded.as_bytes())
.map_err(|_| Error::InvalidRow)?;
let body = encoder.finish().map_err(|_| Error::InvalidRow)?;
let mut url = connection.url().clone();
url.query_pairs_mut()
.append_pair(
"query",
&format!(
"INSERT INTO `{database}`.{} FORMAT JSONEachRow",
table.name()
),
)
.append_pair("async_insert", "1")
.append_pair("async_insert_deduplicate", "1")
.append_pair("wait_for_async_insert", "1")
.append_pair("date_time_input_format", "best_effort");
let response = client
.post(url)
.timeout(INSERT_TIMEOUT)
.header("Content-Encoding", "gzip")
.body(body)
.send()
.await
.map_err(|_| Error::Transport)?;
if !response.status().is_success() {
return Err(Error::InsertFailed(response.status().as_u16()));
}
Ok(())
}
pub fn encode_rows(rows: Vec<BTreeMap<String, Value>>) -> Result<String, Error> {
encode_rows_with_limit(rows, usize::MAX)
}
fn encode_rows_with_limit(
rows: Vec<BTreeMap<String, Value>>,
limit: usize,
) -> Result<String, Error> {
let mut body = Vec::new();
for row in rows {
let encoded = row
.into_iter()
.map(|(name, value)| insert_value(&name, value).map(|value| (name, value)))
.collect::<Result<BTreeMap<_, _>, _>>()?;
let record = serde_json::to_vec(&encoded).map_err(|_| Error::InvalidRow)?;
let size = body
.len()
.checked_add(record.len())
.and_then(|size| size.checked_add(usize::from(!body.is_empty())))
.ok_or(Error::InsertTooLarge)?;
if size > limit {
return Err(Error::InsertTooLarge);
}
if !body.is_empty() {
body.push(b'\n');
}
body.extend_from_slice(&record);
}
String::from_utf8(body).map_err(|_| Error::InvalidRow)
}
fn insert_value(name: &str, value: Value) -> Result<Value, Error> {
let multiplier = match name {
"Timestamp" => 1,
"start_time" | "end_time" | "completion_start_time" => 1_000_000,
_ => return Ok(value),
};
if name == "completion_start_time" && value.is_null() {
return Ok(value);
}
let timestamp = value.as_i64().ok_or(Error::InvalidRow)?;
let datetime = OffsetDateTime::from_unix_timestamp_nanos(i128::from(timestamp) * multiplier)
.map_err(|_| Error::InvalidRow)?;
datetime
.format(&Rfc3339)
.map(Value::String)
.map_err(|_| Error::InvalidRow)
}
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
use rstest::rstest;
use serde_json::json;
use super::encode_rows_with_limit;
use crate::Error;
#[rstest]
fn encoded_limit_counts_utf8_bytes_across_rows() {
let rows = vec![
BTreeMap::from([("Input".to_owned(), json!("雪"))]),
BTreeMap::from([("Input".to_owned(), json!("雪"))]),
];
let encoded = encode_rows_with_limit(rows.clone(), usize::MAX).expect("valid rows");
assert!(encode_rows_with_limit(rows.clone(), encoded.len()).is_ok());
assert!(matches!(
encode_rows_with_limit(rows, encoded.len() - 1),
Err(Error::InsertTooLarge)
));
}
}

View file

@ -0,0 +1,90 @@
mod error;
mod insert;
mod otlp;
mod schema;
mod sql;
pub use error::{DecodeError, Error};
pub use insert::{InsertTable, encode_rows, insert_rows};
pub use otlp::{DecodedSpan, decode_otlp};
pub use schema::{ensure_schema, schema_statements};
pub use sql::{Parameter, execute_read};
use url::Url;
#[derive(Clone)]
pub struct Connection {
url: Url,
}
impl Connection {
pub fn parse(value: &str) -> Result<Self, Error> {
let url = Url::parse(value).map_err(|_| Error::InvalidUrl)?;
if !matches!(url.scheme(), "http" | "https") || url.host().is_none() {
return Err(Error::InvalidUrl);
}
Ok(Self { url })
}
pub fn configured(
url: &str,
database: &str,
user: &str,
password: &str,
) -> Result<Self, Error> {
let mut connection = Self::parse(url)?;
connection
.url
.set_username(user)
.map_err(|_| Error::InvalidUrl)?;
connection
.url
.set_password(Some(password))
.map_err(|_| Error::InvalidUrl)?;
let pairs: Vec<_> = connection
.url
.query_pairs()
.filter(|(key, _)| !matches!(key.as_ref(), "database" | "user" | "password"))
.map(|(key, value)| (key.into_owned(), value.into_owned()))
.collect();
connection
.url
.query_pairs_mut()
.clear()
.extend_pairs(pairs)
.append_pair("database", database);
Ok(connection)
}
pub fn writer(url: &str) -> Result<Self, Error> {
let mut connection = Self::parse(url)?;
let pairs: Vec<_> = connection
.url
.query_pairs()
.filter(|(key, _)| !matches!(key.as_ref(), "database" | "readonly" | "query"))
.map(|(key, value)| (key.into_owned(), value.into_owned()))
.collect();
connection.url.query_pairs_mut().clear().extend_pairs(pairs);
Ok(connection)
}
pub fn reader(url: &str, database: &str) -> Result<Self, Error> {
let mut connection = Self::parse(url)?;
let pairs: Vec<_> = connection
.url
.query_pairs()
.filter(|(key, _)| key != "database")
.map(|(key, value)| (key.into_owned(), value.into_owned()))
.collect();
connection
.url
.query_pairs_mut()
.clear()
.extend_pairs(pairs)
.append_pair("database", database);
Ok(connection)
}
pub fn url(&self) -> &Url {
&self.url
}
}

View file

@ -0,0 +1,221 @@
use std::{collections::BTreeMap, io::Read};
use base64::Engine;
use flate2::read::GzDecoder;
use opentelemetry_proto::tonic::{
collector::trace::v1::ExportTraceServiceRequest,
common::v1::{AnyValue, KeyValue, any_value::Value as AttributeValue},
trace::v1::{Span, span::SpanKind, status::StatusCode},
};
use prost::Message;
use serde::Serialize;
use serde_json::Value;
use crate::DecodeError;
#[derive(Serialize)]
pub struct DecodedEvent {
pub name: String,
pub attributes: BTreeMap<String, String>,
}
#[derive(Serialize)]
pub struct DecodedSpan {
pub trace_id: String,
pub span_id: String,
pub parent_span_id: String,
pub trace_state: String,
pub name: String,
pub kind: String,
pub resource_attributes: BTreeMap<String, String>,
pub scope_name: String,
pub scope_version: String,
pub attributes: BTreeMap<String, String>,
pub start_ns: u64,
pub end_ns: u64,
pub status_code: String,
pub status_message: String,
pub events: Vec<DecodedEvent>,
}
pub fn decode_otlp(
body: &[u8],
content_type: Option<&str>,
content_encoding: Option<&str>,
max_decompressed_bytes: usize,
) -> Result<Vec<DecodedSpan>, DecodeError> {
let payload = if content_encoding == Some("gzip") || body.starts_with(&[0x1f, 0x8b]) {
let limit = u64::try_from(max_decompressed_bytes).map_err(|_| DecodeError::TooLarge)?;
let mut decoded = Vec::new();
GzDecoder::new(body)
.take(limit + 1)
.read_to_end(&mut decoded)
.map_err(|_| DecodeError::InvalidPayload)?;
decoded
} else {
body.to_vec()
};
if payload.len() > max_decompressed_bytes {
return Err(DecodeError::TooLarge);
}
let request = if content_type.is_some_and(|value| value.contains("json")) {
let value: Value =
serde_json::from_slice(&payload).map_err(|_| DecodeError::InvalidPayload)?;
serde_json::from_value(normalize_json_ids(value)?)
.map_err(|_| DecodeError::InvalidPayload)?
} else {
ExportTraceServiceRequest::decode(payload.as_slice())
.map_err(|_| DecodeError::InvalidPayload)?
};
Ok(request
.resource_spans
.into_iter()
.flat_map(|resource_spans| {
let resource_attributes = attributes(
resource_spans
.resource
.map(|resource| resource.attributes)
.unwrap_or_default(),
);
resource_spans
.scope_spans
.into_iter()
.flat_map(move |scope_spans| {
let scope = scope_spans.scope.unwrap_or_default();
let resource_attributes = resource_attributes.clone();
scope_spans.spans.into_iter().map(move |span| {
decoded_span(span, &resource_attributes, &scope.name, &scope.version)
})
})
})
.collect())
}
fn normalize_json_ids(value: Value) -> Result<Value, DecodeError> {
match value {
Value::Object(fields) => fields
.into_iter()
.map(|(name, value)| {
let normalized = if matches!(name.as_str(), "traceId" | "spanId" | "parentSpanId") {
let encoded = value.as_str().ok_or(DecodeError::InvalidPayload)?;
let bytes = base64::engine::general_purpose::STANDARD
.decode(encoded)
.map_err(|_| DecodeError::InvalidPayload)?;
Value::String(hex_bytes(&bytes))
} else if name == "kind" && value.is_string() {
let kind = SpanKind::from_str_name(value.as_str().unwrap_or_default())
.ok_or(DecodeError::InvalidPayload)?;
Value::from(kind as i32)
} else if name == "code" && value.is_string() {
let code = StatusCode::from_str_name(value.as_str().unwrap_or_default())
.ok_or(DecodeError::InvalidPayload)?;
Value::from(code as i32)
} else {
normalize_json_ids(value)?
};
Ok((name, normalized))
})
.collect::<Result<serde_json::Map<_, _>, _>>()
.map(Value::Object),
Value::Array(values) => values
.into_iter()
.map(normalize_json_ids)
.collect::<Result<Vec<_>, _>>()
.map(Value::Array),
value => Ok(value),
}
}
fn hex_bytes(bytes: &[u8]) -> String {
bytes.iter().map(|byte| format!("{byte:02x}")).collect()
}
fn decoded_span(
span: Span,
resource_attributes: &BTreeMap<String, String>,
scope_name: &str,
scope_version: &str,
) -> DecodedSpan {
let status = span.status.unwrap_or_default();
DecodedSpan {
trace_id: hex_bytes(&span.trace_id),
span_id: hex_bytes(&span.span_id),
parent_span_id: hex_bytes(&span.parent_span_id),
trace_state: span.trace_state,
name: span.name,
kind: SpanKind::try_from(span.kind)
.unwrap_or(SpanKind::Unspecified)
.as_str_name()
.to_owned(),
resource_attributes: resource_attributes.clone(),
scope_name: scope_name.to_owned(),
scope_version: scope_version.to_owned(),
attributes: attributes(span.attributes),
start_ns: span.start_time_unix_nano,
end_ns: span.end_time_unix_nano,
status_code: StatusCode::try_from(status.code)
.unwrap_or(StatusCode::Unset)
.as_str_name()
.to_owned(),
status_message: status.message,
events: span
.events
.into_iter()
.map(|event| DecodedEvent {
name: event.name,
attributes: attributes(event.attributes),
})
.collect(),
}
}
fn attributes(values: Vec<KeyValue>) -> BTreeMap<String, String> {
values
.into_iter()
.map(|entry| {
(
entry.key,
entry.value.as_ref().map(attribute_text).unwrap_or_default(),
)
})
.collect()
}
fn attribute_text(value: &AnyValue) -> String {
match value.value.as_ref() {
Some(AttributeValue::StringValue(value)) => value.clone(),
Some(AttributeValue::BoolValue(value)) => value.to_string(),
Some(AttributeValue::IntValue(value)) => value.to_string(),
Some(AttributeValue::DoubleValue(value)) => {
serde_json::to_string(value).unwrap_or_default()
}
Some(AttributeValue::BytesValue(value)) => String::from_utf8_lossy(value).into_owned(),
Some(AttributeValue::ArrayValue(value)) => format!(
"[{}]",
value
.values
.iter()
.map(|value| serde_json::to_string(&attribute_text(value)).unwrap_or_default())
.collect::<Vec<_>>()
.join(", ")
),
Some(AttributeValue::KvlistValue(value)) => format!(
"{{{}}}",
value
.values
.iter()
.map(|entry| format!(
"{}: {}",
serde_json::to_string(&entry.key).unwrap_or_default(),
serde_json::to_string(
&entry.value.as_ref().map(attribute_text).unwrap_or_default()
)
.unwrap_or_default()
))
.collect::<Vec<_>>()
.join(", ")
),
Some(AttributeValue::StringValueStrindex(value)) => value.to_string(),
None => String::new(),
}
}

View file

@ -0,0 +1,87 @@
use litellm_http::Client;
use std::time::Duration;
use crate::Connection;
use crate::Error;
const SCHEMA_REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
const MIGRATIONS: [&str; 7] = [
include_str!("../migrations/0001_otel_traces.sql"),
include_str!("../migrations/0002_agent_traces.sql"),
include_str!("../migrations/0003_agent_traces_mv.sql"),
include_str!("../migrations/0004_spend_logs.sql"),
include_str!("../migrations/0005_otel_traces_ttl.sql"),
include_str!("../migrations/0006_agent_traces_ttl.sql"),
include_str!("../migrations/0007_spend_logs_ttl.sql"),
];
pub fn schema_statements(
database: &str,
trace_retention_days: u32,
spend_log_retention_days: u32,
) -> Result<Vec<String>, Error> {
if database.is_empty()
|| !database
.bytes()
.all(|c| c.is_ascii_alphanumeric() || c == b'_')
|| trace_retention_days == 0
|| spend_log_retention_days == 0
{
return Err(Error::InvalidSchema);
}
let database = format!("`{database}`");
Ok(
std::iter::once(format!("CREATE DATABASE IF NOT EXISTS {database}"))
.chain(MIGRATIONS.iter().map(|sql| {
sql.replace("{database}", &database)
.replace("{trace_retention_days}", &trace_retention_days.to_string())
.replace(
"{spend_log_retention_days}",
&spend_log_retention_days.to_string(),
)
}))
.collect(),
)
}
pub async fn ensure_schema(
client: &Client,
connection: &Connection,
database: &str,
trace_retention_days: u32,
spend_log_retention_days: u32,
) -> Result<(), Error> {
ensure_schema_with_timeout(
client,
connection,
database,
trace_retention_days,
spend_log_retention_days,
SCHEMA_REQUEST_TIMEOUT,
)
.await
}
async fn ensure_schema_with_timeout(
client: &Client,
connection: &Connection,
database: &str,
trace_retention_days: u32,
spend_log_retention_days: u32,
request_timeout: Duration,
) -> Result<(), Error> {
for statement in schema_statements(database, trace_retention_days, spend_log_retention_days)? {
let response = client
.post(connection.url().clone())
.timeout(request_timeout)
.body(statement)
.send()
.await
.map_err(|_| Error::Transport)?;
if !response.status().is_success() {
return Err(Error::SchemaFailed(response.status().as_u16()));
}
}
Ok(())
}

View file

@ -0,0 +1,114 @@
use std::{collections::BTreeMap, time::Duration};
use serde::Deserialize;
use litellm_http::Client;
use crate::{Connection, Error};
const MAX_RESPONSE_BYTES: usize = 4 * 1024 * 1024;
#[derive(Debug, Deserialize)]
#[serde(untagged)]
pub enum Parameter {
Text(String),
Integer(i64),
Strings(Vec<String>),
}
impl Parameter {
fn encoded(&self) -> String {
match self {
Self::Text(value) => escaped(value),
Self::Integer(value) => value.to_string(),
Self::Strings(values) => format!(
"[{}]",
values
.iter()
.map(|value| format!("'{}'", escaped(value).replace('\'', "\\'")))
.collect::<Vec<_>>()
.join(",")
),
}
}
}
fn escaped(value: &str) -> String {
value
.replace('\\', "\\\\")
.replace('\t', "\\t")
.replace('\n', "\\n")
.replace('\r', "\\r")
.replace('\0', "\\0")
}
pub async fn execute_read(
client: &Client,
connection: &Connection,
sql: &str,
parameters: &BTreeMap<String, Parameter>,
) -> Result<String, Error> {
if sql.trim().is_empty() {
return Err(Error::EmptySql);
}
let mut url = connection.url().clone();
let existing_pairs: Vec<(String, String)> = url
.query_pairs()
.filter(|(key, _)| {
!key.starts_with("param_")
&& !matches!(
key.as_ref(),
"query"
| "readonly"
| "default_format"
| "max_result_rows"
| "result_overflow_mode"
| "max_execution_time"
| "wait_end_of_query"
)
})
.map(|(key, value)| (key.into_owned(), value.into_owned()))
.collect();
url.query_pairs_mut()
.clear()
.extend_pairs(existing_pairs)
.append_pair("readonly", "1")
.append_pair("max_result_rows", "1000")
.append_pair("result_overflow_mode", "throw")
.append_pair("max_execution_time", "10")
.append_pair("wait_end_of_query", "1")
.append_pair("default_format", "JSON");
url.query_pairs_mut().extend_pairs(
parameters
.iter()
.map(|(name, value)| (format!("param_{name}"), value.encoded())),
);
let request = client
.post(url)
.timeout(Duration::from_secs(15))
.body(sql.to_owned());
let mut response = request.send().await.map_err(|_| Error::Transport)?;
if !response.status().is_success() {
return Err(Error::QueryFailed(response.status().as_u16()));
}
let mut body = Vec::new();
while let Some(chunk) = response.chunk().await.map_err(|_| Error::Transport)? {
if body.len() + chunk.len() > MAX_RESPONSE_BYTES {
return Err(Error::ResponseTooLarge);
}
body.extend_from_slice(&chunk);
}
let json: serde_json::Value =
serde_json::from_slice(&body).map_err(|_| Error::InvalidResponse)?;
if json.get("exception").is_some() || !json.get("data").is_some_and(serde_json::Value::is_array)
{
return Err(Error::InvalidResponse);
}
String::from_utf8(body).map_err(|_| Error::InvalidResponse)
}

View file

@ -0,0 +1,273 @@
use litellm_http::Client;
use litellm_traces::{Connection, Error, Parameter, execute_read};
use rstest::{fixture, rstest};
use serde_json::Value;
use std::collections::BTreeMap;
use testcontainers_modules::{
clickhouse::ClickHouse,
testcontainers::{ContainerAsync, ImageExt, runners::AsyncRunner},
};
const CLICKHOUSE_TAG: &str =
"26.9.6.6@sha256:eb4870e7ca7ed70c259eebfcfbee6cf797017f6b5436c2926bbbfe3d4d28486e";
struct Database {
_container: ContainerAsync<ClickHouse>,
url: String,
admin_url: String,
client: Client,
}
#[fixture]
async fn database() -> Result<Database, Box<dyn std::error::Error>> {
let container = ClickHouse::default()
.with_tag(CLICKHOUSE_TAG)
.with_env_var("CLICKHOUSE_SKIP_USER_SETUP", "1")
.with_env_var("LITELLM_TRACES_READER_PASSWORD", "test_password")
.with_copy_to(
"/etc/clickhouse-server/users.d/litellm-traces-reader.xml",
include_bytes!("../config/reader.xml").to_vec(),
)
.start()
.await?;
let admin_url = format!(
"http://{}:{}",
container.get_host().await?,
container.get_host_port_ipv4(8123).await?,
);
let client = Client::no_redirect_for_test();
for sql in [
"CREATE DATABASE litellm",
"CREATE TABLE litellm.otel_traces (n UInt8) ENGINE = Memory",
"INSERT INTO litellm.otel_traces VALUES (1)",
"CREATE TABLE litellm.agent_traces_by_key (n UInt8) ENGINE = Memory",
"INSERT INTO litellm.agent_traces_by_key VALUES (4)",
"CREATE TABLE litellm.spend_logs (n UInt8) ENGINE = Memory",
"INSERT INTO litellm.spend_logs VALUES (3)",
"CREATE TABLE litellm.private_traces (n UInt8) ENGINE = Memory",
"CREATE TABLE private_traces (n UInt8) ENGINE = Memory",
] {
client
.post(&admin_url)
.body(sql)
.send()
.await?
.error_for_status()?;
}
let url = format!(
"{}?database=litellm",
admin_url.replacen("http://", "http://litellm_traces_reader:test_password@", 1)
);
Ok(Database {
_container: container,
url,
admin_url,
client,
})
}
#[rstest]
#[tokio::test]
async fn admin_sql_reads_rows_with_enforced_settings(
#[future(awt)] database: Result<Database, Box<dyn std::error::Error>>,
) -> Result<(), Box<dyn std::error::Error>> {
let database = database?;
let connection = Connection::parse(&format!(
"{}&readonly=0&default_format=TabSeparated&query=SELECT+2",
database.url,
))?;
let result = read(
&database.client,
&connection,
"SELECT n AS answer FROM otel_traces",
)
.await?;
let json: Value = serde_json::from_str(&result)?;
assert_eq!(json["data"][0]["answer"], 1);
let result = read(
&database.client,
&connection,
"SELECT n AS answer FROM agent_traces_by_key",
)
.await?;
let json: Value = serde_json::from_str(&result)?;
assert_eq!(json["data"][0]["answer"], 4);
Ok(())
}
#[rstest]
#[case::table("CREATE TABLE admin_sql_test (n UInt8) ENGINE = Memory")]
#[case::insert("INSERT INTO otel_traces VALUES (2)")]
#[case::drop("DROP TABLE otel_traces")]
#[case::named_collection("CREATE NAMED COLLECTION admin_sql_test AS host = 'localhost'")]
#[case::settings("SET readonly = 0")]
#[case::inline_settings("SELECT n FROM otel_traces SETTINGS readonly = 0")]
#[case::time_limit("SELECT n FROM otel_traces SETTINGS max_execution_time = 0")]
#[case::row_limit("SELECT n FROM otel_traces SETTINGS max_result_rows = 0")]
#[case::byte_limit("SELECT n FROM otel_traces SETTINGS max_result_bytes = 0")]
#[case::memory_limit("SELECT n FROM otel_traces SETTINGS max_memory_usage = 0")]
#[case::other_table("SELECT * FROM private_traces")]
#[tokio::test]
async fn reader_rejects_writes_and_privilege_escalation(
#[future(awt)] database: Result<Database, Box<dyn std::error::Error>>,
#[case] sql: &str,
) -> Result<(), Box<dyn std::error::Error>> {
let database = database?;
let connection = Connection::parse(&format!("{}&readonly=0", database.url))?;
let result = read(&database.client, &connection, sql).await;
assert!(matches!(result, Err(Error::QueryFailed(_))), "{result:?}");
let rows = read(&database.client, &connection, "SELECT n FROM otel_traces").await?;
let json: Value = serde_json::from_str(&rows)?;
assert_eq!(json["data"], serde_json::json!([{ "n": 1 }]));
Ok(())
}
#[rstest]
#[tokio::test]
async fn admin_sql_rejects_errors_after_output_starts(
#[future(awt)] database: Result<Database, Box<dyn std::error::Error>>,
) -> Result<(), Box<dyn std::error::Error>> {
let database = database?;
let connection = Connection::parse(&format!(
"{}?max_block_size=1&buffer_size=1&http_write_exception_in_output_format=1\
&send_progress_in_http_headers=1&http_headers_progress_interval_ms=0",
database.admin_url,
))?;
let result = read(
&database.client,
&connection,
"SELECT sleepEachRow(0.2), throwIf(number = 2) FROM numbers(5)",
)
.await;
assert!(
matches!(result, Err(Error::InvalidResponse)),
"expected an error embedded in a successful HTTP response: {result:?}"
);
Ok(())
}
#[rstest]
#[tokio::test]
async fn admin_sql_enforces_result_row_limit(
#[future(awt)] database: Result<Database, Box<dyn std::error::Error>>,
) -> Result<(), Box<dyn std::error::Error>> {
let database = database?;
let connection = Connection::parse(&format!(
"{}&max_result_rows=0&result_overflow_mode=throw&wait_end_of_query=1",
database.url,
))?;
let result = read(
&database.client,
&connection,
"SELECT number FROM numbers(1001)",
)
.await;
assert!(matches!(result, Err(Error::QueryFailed(_))), "{result:?}");
Ok(())
}
#[rstest]
#[tokio::test]
async fn admin_sql_enforces_response_byte_limit(
#[future(awt)] database: Result<Database, Box<dyn std::error::Error>>,
) -> Result<(), Box<dyn std::error::Error>> {
let database = database?;
let connection = Connection::parse(&database.admin_url)?;
let result = read(
&database.client,
&connection,
"SELECT repeat('x', 512 * 1024) AS payload FROM numbers(9)",
)
.await;
assert!(matches!(result, Err(Error::ResponseTooLarge)), "{result:?}");
Ok(())
}
#[rstest]
#[case::plain("test_password", "test_password")]
#[case::encoded("p@ss/word%", "p%40ss%2Fword%25")]
#[tokio::test]
async fn admin_sql_authenticates_url_credentials(
#[future(awt)] database: Result<Database, Box<dyn std::error::Error>>,
#[case] password: &str,
#[case] encoded_password: &str,
) -> Result<(), Box<dyn std::error::Error>> {
let database = database?;
database
.client
.post(&database.admin_url)
.body(format!(
"CREATE USER sql_reader IDENTIFIED WITH plaintext_password BY '{password}'"
))
.send()
.await?
.error_for_status()?;
let connection = Connection::parse(&database.admin_url.replacen(
"http://",
&format!("http://sql_reader:{encoded_password}@"),
1,
))?;
let result = read(
&database.client,
&connection,
"SELECT currentUser() AS username",
)
.await?;
let json: Value = serde_json::from_str(&result)?;
assert_eq!(json["data"][0]["username"], "sql_reader");
Ok(())
}
async fn read(client: &Client, connection: &Connection, sql: &str) -> Result<String, Error> {
execute_read(client, connection, sql, &BTreeMap::new()).await
}
#[rstest]
#[case::sql("'; DROP TABLE otel_traces; --")]
#[case::escapes("back\\slash\ttab\nline\0null")]
#[tokio::test]
async fn query_parameters_preserve_values_and_replace_url_parameters(
#[case] value: &str,
#[future(awt)] database: Result<Database, Box<dyn std::error::Error>>,
) -> Result<(), Box<dyn std::error::Error>> {
let database = database?;
let connection = Connection::parse(&format!("{}&param_value=wrong", database.url))?;
let values = vec![
"a'b".to_owned(),
"back\\slash".to_owned(),
"line\nbreak".to_owned(),
"雪".to_owned(),
];
let parameters = BTreeMap::from([
("value".to_owned(), Parameter::Text(value.into())),
("teams".to_owned(), Parameter::Strings(values.clone())),
("number".to_owned(), Parameter::Integer(-42)),
]);
let body = execute_read(&database.client, &connection,
"SELECT {value:String} AS value, {teams:Array(String)} AS teams, toInt32({number:Int64}) AS number",
&parameters).await?;
let json: Value = serde_json::from_str(&body)?;
assert_eq!(json["data"][0]["value"], value);
assert_eq!(json["data"][0]["teams"], serde_json::json!(values));
assert_eq!(json["data"][0]["number"], -42);
assert!(
read(&database.client, &connection, "SELECT n FROM otel_traces")
.await
.is_ok()
);
Ok(())
}

View file

@ -0,0 +1,40 @@
use std::collections::BTreeMap;
use litellm_traces::encode_rows;
use rstest::rstest;
use serde_json::{Value, json};
#[rstest]
#[case::span("Timestamp", json!(1_234_567_890), json!("1970-01-01T00:00:01.23456789Z"))]
#[case::start("start_time", json!(1_234), json!("1970-01-01T00:00:01.234Z"))]
#[case::end("end_time", json!(2_345), json!("1970-01-01T00:00:02.345Z"))]
#[case::completion("completion_start_time", json!(1_345), json!("1970-01-01T00:00:01.345Z"))]
#[case::absent_completion("completion_start_time", Value::Null, Value::Null)]
#[case::before_epoch("Timestamp", json!(-1), json!("1969-12-31T23:59:59.999999999Z"))]
fn insert_encoding_preserves_timestamp_precision_and_other_fields(
#[case] field: &str,
#[case] value: Value,
#[case] expected: Value,
) {
let rows = vec![BTreeMap::from([
(field.to_owned(), value),
("SpanAttributes".into(), json!({"message": "a\nb\\c\"雪"})),
("InputTokens".into(), json!(42)),
])];
let encoded = encode_rows(rows).expect("valid row");
let actual: Value = serde_json::from_str(&encoded).expect("JSONEachRow record");
assert_eq!(
actual,
json!({
field: expected, "SpanAttributes": {"message": "a\nb\\c\"雪"}, "InputTokens": 42
})
);
}
#[rstest]
#[case::fractional(json!(1.25))]
#[case::out_of_range(json!(u64::MAX))]
#[case::null(Value::Null)]
fn insert_encoding_rejects_invalid_span_timestamps(#[case] timestamp: Value) {
assert!(encode_rows(vec![BTreeMap::from([("Timestamp".into(), timestamp)])]).is_err());
}

View file

@ -0,0 +1,428 @@
use std::{collections::BTreeMap, time::Duration};
use litellm_http::Client;
use litellm_traces::{
Connection, Error, InsertTable, encode_rows, ensure_schema, execute_read, schema_statements,
};
use rstest::{fixture, rstest};
use testcontainers_modules::{
clickhouse::ClickHouse,
testcontainers::{ContainerAsync, ImageExt, runners::AsyncRunner},
};
const CLICKHOUSE_TAG: &str =
"26.9.6.6@sha256:eb4870e7ca7ed70c259eebfcfbee6cf797017f6b5436c2926bbbfe3d4d28486e";
type TestResult<T = ()> = Result<T, Box<dyn std::error::Error>>;
struct ClickHouseDatabase {
_container: ContainerAsync<ClickHouse>,
url: String,
client: Client,
}
#[fixture]
async fn database() -> TestResult<ClickHouseDatabase> {
let container = ClickHouse::default()
.with_tag(CLICKHOUSE_TAG)
.with_env_var("CLICKHOUSE_SKIP_USER_SETUP", "1")
.start()
.await?;
let url = format!(
"http://{}:{}",
container.get_host().await?,
container.get_host_port_ipv4(8123).await?
);
Ok(ClickHouseDatabase {
_container: container,
url,
client: Client::no_redirect_for_test(),
})
}
async fn insert_rows(
database: &ClickHouseDatabase,
table: &str,
rows: Vec<BTreeMap<String, serde_json::Value>>,
) -> TestResult {
database
.client
.post(&database.url)
.query(&[
(
"query",
format!("INSERT INTO trace_test.{table} FORMAT JSONEachRow"),
),
("date_time_input_format", "best_effort".into()),
])
.body(encode_rows(rows)?)
.send()
.await?
.error_for_status()?;
Ok(())
}
async fn execute_write(database: &ClickHouseDatabase, sql: &str) -> TestResult {
database
.client
.post(&database.url)
.body(sql.to_owned())
.send()
.await?
.error_for_status()?;
Ok(())
}
async fn read_json(database: &ClickHouseDatabase, sql: &str) -> TestResult<serde_json::Value> {
let connection = Connection::configured(&database.url, "trace_test", "default", "")?;
let body = execute_read(&database.client, &connection, sql, &BTreeMap::new()).await?;
Ok(serde_json::from_str(&body)?)
}
async fn table_rows(database: &ClickHouseDatabase, table: &str) -> TestResult<u64> {
let response = read_json(
database,
&format!("SELECT count() AS rows FROM trace_test.{table}"),
)
.await?;
Ok(response["data"][0]["rows"]
.as_u64()
.expect("ClickHouse returns row counts as unsigned integers"))
}
async fn mutation_rows(database: &ClickHouseDatabase) -> TestResult<u64> {
let response = read_json(
database,
"SELECT count() AS rows FROM system.mutations WHERE database = 'trace_test'",
)
.await?;
Ok(response["data"][0]["rows"]
.as_u64()
.expect("ClickHouse returns mutation counts as unsigned integers"))
}
#[rstest]
#[tokio::test]
async fn schema_supports_span_rollups_and_spend_joins(
#[future(awt)] database: TestResult<ClickHouseDatabase>,
) -> TestResult {
let database = database?;
let writer = Connection::writer(&database.url)?;
ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?;
ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?;
let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64;
let span = serde_json::from_value(serde_json::json!({
"Timestamp": timestamp, "TraceId": "trace-1", "SpanId": "span-1", "ParentSpanId": "",
"ServiceName": "proxy", "SpanName": "request", "Input": "hello world",
"ResourceAttributes": {"litellm.team_id": "team-1", "litellm.api_key_hash": "hash-1"},
"SpanAttributes": {"gen_ai.response.id": "response-1", "gen_ai.usage.input_tokens": "12"}
}))?;
let spend = serde_json::from_value(serde_json::json!({
"request_id": "request-1", "response_id": "response-1", "team_id": "team-1", "spend": 0.125,
"start_time": timestamp / 1_000_000, "end_time": timestamp / 1_000_000 + 100,
"completion_start_time": null
}))?;
insert_rows(&database, "otel_traces", vec![span]).await?;
insert_rows(&database, "spend_logs", vec![spend]).await?;
let body = read_json(
&database,
"SELECT o.TeamId, o.ApiKeyHash, o.ObservationType, o.InputPreview, s.spend, \
toString(toUnixTimestamp64Nano(o.Timestamp)) AS timestamp_ns, \
toString(toUnixTimestamp64Milli(s.start_time)) AS start_ms \
FROM trace_test.otel_traces o JOIN trace_test.spend_logs s \
ON o.LiteLLMRequestId = s.response_id AND o.TeamId = s.team_id",
)
.await?;
assert_eq!(
body["data"],
serde_json::json!([{
"TeamId": "team-1", "ApiKeyHash": "hash-1", "ObservationType": "agent",
"InputPreview": "hello world", "spend": 0.125,
"timestamp_ns": timestamp.to_string(), "start_ms": (timestamp / 1_000_000).to_string()
}])
);
let body = read_json(
&database,
"SELECT toUInt32(sum(SpanCount)) AS spans, toUInt32(sum(InputTokens)) AS tokens \
FROM trace_test.agent_traces_by_key WHERE TeamId = 'team-1' AND TraceId = 'trace-1'",
)
.await?;
assert_eq!(
body["data"],
serde_json::json!([{"spans": 1, "tokens": 12}])
);
Ok(())
}
#[rstest]
#[tokio::test]
async fn retried_trace_insert_does_not_inflate_rollup(
#[future(awt)] database: TestResult<ClickHouseDatabase>,
) -> TestResult {
let database = database?;
let writer = Connection::writer(&database.url)?;
ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?;
let row: BTreeMap<String, serde_json::Value> = serde_json::from_value(serde_json::json!({
"Timestamp": time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64,
"TraceId": "retried-trace", "SpanId": "span-1", "ParentSpanId": "",
"TeamId": "team-1", "ApiKeyHash": "key-1", "SpanName": "root", "InputTokens": 7
}))?;
for _ in 0..2 {
litellm_traces::insert_rows(
&database.client,
&writer,
"trace_test",
InsertTable::OtelTraces,
vec![row.clone()],
)
.await?;
}
let counts = read_json(
&database,
"SELECT toUInt32(sum(SpanCount)) AS spans, toUInt32(sum(InputTokens)) AS tokens \
FROM trace_test.agent_traces_by_key WHERE TraceId = 'retried-trace'",
)
.await?;
assert_eq!(table_rows(&database, "otel_traces").await?, 1);
assert_eq!(counts["data"][0]["spans"], 1);
assert_eq!(counts["data"][0]["tokens"], 7);
Ok(())
}
#[rstest]
#[tokio::test]
async fn keyed_rollup_keeps_same_trace_ids_separate_by_api_key(
#[future(awt)] database: TestResult<ClickHouseDatabase>,
) -> TestResult {
let database = database?;
let writer = Connection::writer(&database.url)?;
ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?;
let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64;
let rows = vec![
serde_json::from_value(serde_json::json!({
"Timestamp": timestamp, "TraceId": "shared-id", "SpanId": "root-one",
"ParentSpanId": "", "SpanName": "root-one", "Input": "private-one",
"ResourceAttributes": {"litellm.api_key_hash": "key-one"}
}))?,
serde_json::from_value(serde_json::json!({
"Timestamp": timestamp, "TraceId": "shared-id", "SpanId": "root-two",
"ParentSpanId": "", "SpanName": "root-two", "Input": "private-two",
"ResourceAttributes": {"litellm.api_key_hash": "key-two"}
}))?,
];
insert_rows(&database, "otel_traces", rows).await?;
execute_write(
&database,
"OPTIMIZE TABLE trace_test.agent_traces_by_key FINAL",
)
.await?;
let rows = read_json(
&database,
"SELECT ApiKeyHash, any(RootInput) AS RootInput \
FROM trace_test.agent_traces_by_key WHERE TraceId = 'shared-id' \
GROUP BY ApiKeyHash ORDER BY ApiKeyHash",
)
.await?;
assert_eq!(
rows["data"],
serde_json::json!([
{"ApiKeyHash": "key-one", "RootInput": "private-one"},
{"ApiKeyHash": "key-two", "RootInput": "private-two"}
])
);
Ok(())
}
#[rstest]
#[tokio::test]
async fn rollup_merges_spans_across_days_without_losing_root_fields(
#[future(awt)] database: TestResult<ClickHouseDatabase>,
) -> TestResult {
let database = database?;
let writer = Connection::writer(&database.url)?;
ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?;
let day_start = time::OffsetDateTime::now_utc()
.replace_time(time::Time::MIDNIGHT)
.unix_timestamp_nanos() as i64;
let root = serde_json::from_value(serde_json::json!({
"Timestamp": day_start - 1_000_000_000, "TraceId": "cross-day", "SpanId": "span-root",
"ParentSpanId": "", "ServiceName": "proxy", "SpanName": "root", "Input": "root input",
"StatusCode": "STATUS_CODE_ERROR",
"ResourceAttributes": {"litellm.team_id": "team-1"}
}))?;
insert_rows(&database, "otel_traces", vec![root]).await?;
let child = serde_json::from_value(serde_json::json!({
"Timestamp": day_start + 1_000_000_000, "TraceId": "cross-day", "SpanId": "span-child",
"ParentSpanId": "span-root", "ServiceName": "proxy", "SpanName": "child",
"StatusCode": "STATUS_CODE_UNSET",
"ResourceAttributes": {"litellm.team_id": "team-1"}
}))?;
insert_rows(&database, "otel_traces", vec![child]).await?;
execute_write(
&database,
"OPTIMIZE TABLE trace_test.agent_traces_by_key FINAL",
)
.await?;
let response = read_json(
&database,
"SELECT count() AS rows, any(RootName) AS RootName, any(RootInput) AS RootInput, \
any(RootStatus) AS RootStatus, sum(SpanCount) AS SpanCount \
FROM trace_test.agent_traces_by_key",
)
.await?;
assert_eq!(
response["data"],
serde_json::json!([{
"rows": 1, "RootName": "root", "RootInput": "root input",
"RootStatus": "STATUS_CODE_ERROR", "SpanCount": 2
}])
);
Ok(())
}
#[rstest]
#[tokio::test]
async fn spend_deduplication_preserves_subsecond_requests_and_retries(
#[future(awt)] database: TestResult<ClickHouseDatabase>,
) -> TestResult {
let database = database?;
let writer = Connection::writer(&database.url)?;
ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?;
let now_ms = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64 / 1_000_000;
let base_start_time = now_ms / 1000 * 1000;
let first_start_time = base_start_time + 100;
let second_start_time = base_start_time + 200;
let first = serde_json::from_value(serde_json::json!({
"request_id": "same-request", "team_id": "team-1", "spend": 1.0,
"start_time": first_start_time, "end_time": first_start_time + 1000
}))?;
let second = serde_json::from_value(serde_json::json!({
"request_id": "same-request", "team_id": "team-1", "spend": 2.0,
"start_time": second_start_time, "end_time": second_start_time + 1200
}))?;
let retry = serde_json::from_value(serde_json::json!({
"request_id": "same-request", "team_id": "team-1", "spend": 1.0,
"start_time": first_start_time, "end_time": first_start_time + 2000
}))?;
insert_rows(&database, "spend_logs", vec![first]).await?;
insert_rows(&database, "spend_logs", vec![second]).await?;
insert_rows(&database, "spend_logs", vec![retry]).await?;
execute_write(&database, "OPTIMIZE TABLE trace_test.spend_logs FINAL").await?;
let rows = read_json(
&database,
"SELECT toString(toUnixTimestamp64Milli(start_time)) AS start_time, \
toString(toUnixTimestamp64Milli(end_time)) AS end_time \
FROM trace_test.spend_logs ORDER BY start_time",
)
.await?;
assert_eq!(
rows["data"],
serde_json::json!([
{
"start_time": first_start_time.to_string(),
"end_time": (first_start_time + 2000).to_string()
},
{
"start_time": second_start_time.to_string(),
"end_time": (second_start_time + 1200).to_string()
}
])
);
assert_eq!(table_rows(&database, "spend_logs").await?, 2);
Ok(())
}
#[rstest]
#[tokio::test]
async fn retention_changes_materialize_existing_rows_and_remain_idempotent(
#[future(awt)] database: TestResult<ClickHouseDatabase>,
) -> TestResult {
let database = database?;
let writer = Connection::writer(&database.url)?;
ensure_schema(&database.client, &writer, "trace_test", 30, 30).await?;
let old_time = time::OffsetDateTime::now_utc() - time::Duration::days(20);
let old_timestamp_ns = old_time.unix_timestamp_nanos() as i64;
let old_timestamp_ms = old_timestamp_ns / 1_000_000;
let span = serde_json::from_value(serde_json::json!({
"Timestamp": old_timestamp_ns, "TraceId": "expired", "SpanId": "span-old",
"ParentSpanId": "", "ServiceName": "proxy", "SpanName": "old-root", "Input": "old input",
"ResourceAttributes": {"litellm.team_id": "team-1"}
}))?;
let spend = serde_json::from_value(serde_json::json!({
"request_id": "old-request", "team_id": "team-1", "spend": 1.0,
"start_time": old_timestamp_ms, "end_time": old_timestamp_ms + 1000
}))?;
insert_rows(&database, "otel_traces", vec![span]).await?;
insert_rows(&database, "spend_logs", vec![spend]).await?;
assert_eq!(table_rows(&database, "agent_traces_by_key").await?, 1);
ensure_schema(&database.client, &writer, "trace_test", 14, 14).await?;
let deadline = tokio::time::Instant::now() + Duration::from_secs(60);
loop {
let response = read_json(
&database,
"SELECT countIf(is_done = 0) AS pending \
FROM system.mutations WHERE database = 'trace_test'",
)
.await?;
let pending = response["data"][0]["pending"]
.as_u64()
.expect("ClickHouse returns pending mutation counts as unsigned integers");
if pending == 0 {
break;
}
assert!(
tokio::time::Instant::now() < deadline,
"ClickHouse TTL mutations did not finish before the deadline"
);
tokio::time::sleep(Duration::from_millis(100)).await;
}
execute_write(&database, "OPTIMIZE TABLE trace_test.otel_traces FINAL").await?;
execute_write(
&database,
"OPTIMIZE TABLE trace_test.agent_traces_by_key FINAL",
)
.await?;
execute_write(&database, "OPTIMIZE TABLE trace_test.spend_logs FINAL").await?;
assert_eq!(table_rows(&database, "otel_traces").await?, 0);
assert_eq!(table_rows(&database, "agent_traces_by_key").await?, 0);
assert_eq!(table_rows(&database, "spend_logs").await?, 0);
let mutation_count = mutation_rows(&database).await?;
ensure_schema(&database.client, &writer, "trace_test", 14, 14).await?;
assert_eq!(mutation_rows(&database).await?, mutation_count);
Ok(())
}
#[rstest]
#[tokio::test]
async fn schema_statement_timeout_maps_to_transport_error() -> TestResult {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
let address = listener.local_addr()?;
let server = tokio::spawn(async move {
let (_connection, _) = listener.accept().await.expect("accept schema request");
std::future::pending::<()>().await;
});
let client = Client::no_redirect_for_test();
let url = format!("http://{address}");
let writer = Connection::writer(&url)?;
let result = tokio::time::timeout(
Duration::from_secs(35),
ensure_schema(&client, &writer, "trace_test", 7, 14),
)
.await;
server.abort();
assert!(matches!(result, Ok(Err(Error::Transport))), "{result:?}");
Ok(())
}
#[rstest]
#[case::empty("", 7, 14)]
#[case::sql("db; DROP DATABASE default", 7, 14)]
#[case::trace_retention("traces", 0, 14)]
#[case::spend_retention("traces", 7, 0)]
fn schema_rejects_invalid_configuration(
#[case] database: &str,
#[case] traces: u32,
#[case] spend: u32,
) {
assert!(schema_statements(database, traces, spend).is_err());
}

View file

@ -0,0 +1,47 @@
use flate2::{Compression, write::GzEncoder};
use litellm_traces::decode_otlp;
use rstest::rstest;
use std::io::Write;
const FIXTURE: &[u8] = include_bytes!(
"../../../../tests/test_litellm/tracing/fixtures/langsmith_deep_agent_export.json"
);
#[rstest]
#[case::json(FIXTURE, Some("application/json"), None)]
#[case::gzip_json(FIXTURE, Some("application/json"), Some("gzip"))]
fn decodes_neutral_spans(
#[case] body: &[u8],
#[case] content_type: Option<&str>,
#[case] content_encoding: Option<&str>,
) {
let payload = if content_encoding == Some("gzip") {
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
encoder.write_all(body).expect("gzip input");
encoder.finish().expect("gzip payload")
} else {
body.to_vec()
};
let spans = decode_otlp(&payload, content_type, content_encoding, 8 * 1024 * 1024)
.expect("valid OTLP export");
assert_eq!(spans.len(), 6);
assert_eq!(spans[0].trace_id, "4bad42b84e9de3ba46fc870185f8f023");
assert_eq!(spans[0].resource_attributes["service.name"], "agent-demo");
assert_eq!(spans[0].scope_name, "langsmith");
assert!(
spans
.iter()
.any(|span| span.attributes.contains_key("gen_ai.prompt"))
);
}
#[rstest]
#[case::invalid(b"not protobuf", None, 8 * 1024 * 1024)]
#[case::too_large(FIXTURE, Some("application/json"), 1)]
fn rejects_invalid_or_oversized_payload(
#[case] body: &[u8],
#[case] content_type: Option<&str>,
#[case] limit: usize,
) {
assert!(decode_otlp(body, content_type, None, limit).is_err());
}

View file

@ -0,0 +1,11 @@
use litellm_traces::Connection;
use rstest::rstest;
#[rstest]
#[case::http("http://localhost:8123", true)]
#[case::https("https://localhost:8443", true)]
#[case::tcp("tcp://localhost:9000", false)]
#[case::missing_host("http://", false)]
fn accepts_only_clickhouse_http_urls(#[case] value: &str, #[case] expected: bool) {
assert_eq!(Connection::parse(value).is_ok(), expected);
}

View file

@ -34,7 +34,9 @@
"token-efficient-tools-2025-02-19": "token-efficient-tools-2025-02-19",
"web-fetch-2025-09-10": "web-fetch-2025-09-10",
"web-search-2025-03-05": "web-search-2025-03-05",
"mid-conversation-output-config-2026-07-01": "mid-conversation-output-config-2026-07-01"
"mid-conversation-output-config-2026-07-01": "mid-conversation-output-config-2026-07-01",
"thinking-display-updates-2026-08-18": "thinking-display-updates-2026-08-18",
"mid-conversation-tool-changes-2026-07-01": "mid-conversation-tool-changes-2026-07-01"
},
"azure_ai": {
"advisor-tool-2026-03-01": null,
@ -136,7 +138,9 @@
"tool-search-tool-2025-10-19": "tool-search-tool-2025-10-19",
"web-fetch-2025-09-10": null,
"web-search-2025-03-05": null,
"mid-conversation-output-config-2026-07-01": "mid-conversation-output-config-2026-07-01"
"mid-conversation-output-config-2026-07-01": "mid-conversation-output-config-2026-07-01",
"thinking-display-updates-2026-08-18": "thinking-display-updates-2026-08-18",
"mid-conversation-tool-changes-2026-07-01": "mid-conversation-tool-changes-2026-07-01"
},
"bedrock_mantle": {
"advanced-tool-use-2025-11-20": "tool-search-tool-2025-10-19",

View file

@ -405,8 +405,8 @@ class Cache:
forward_reasoning_content: Final = kwargs.get(
"forward_reasoning_content", nested_litellm_params.get("forward_reasoning_content")
)
if forward_reasoning_content is True:
cache_key += "forward_reasoning_content: True"
if forward_reasoning_content is False:
cache_key += "forward_reasoning_content: False"
reasoning_content_field: Final = kwargs.get(
"reasoning_content_field", nested_litellm_params.get("reasoning_content_field")
)

View file

@ -46,6 +46,19 @@ ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG: Final[frozenset[str]] = frozenset(
)
DEFAULT_BATCH_SIZE: Final = int(os.getenv("DEFAULT_BATCH_SIZE", 512))
DEFAULT_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_FLUSH_INTERVAL_SECONDS", 5))
# Agent tracing / ClickHouse
CLICKHOUSE_BATCH_SIZE: Final = get_env_int("CLICKHOUSE_BATCH_SIZE", 10_000)
CLICKHOUSE_FLUSH_INTERVAL_SECONDS: Final = float(os.getenv("CLICKHOUSE_FLUSH_INTERVAL_SECONDS", "1.0"))
CLICKHOUSE_MAX_BUFFERED_ROWS: Final = get_env_int("CLICKHOUSE_MAX_BUFFERED_ROWS", 200_000)
CLICKHOUSE_MAX_RETRIES: Final = get_env_int("CLICKHOUSE_MAX_RETRIES", 3)
AGENT_TRACING_RETENTION_DAYS: Final = get_env_int("AGENT_TRACING_RETENTION_DAYS", 30)
AGENT_TRACING_SPEND_LOG_RETENTION_DAYS: Final = get_env_int("AGENT_TRACING_SPEND_LOG_RETENTION_DAYS", 90)
OTLP_MAX_BODY_BYTES: Final = get_env_int("OTLP_MAX_BODY_BYTES", 8 * 1024 * 1024)
OTLP_MAX_ATTRIBUTE_VALUE_BYTES: Final = get_env_int("OTLP_MAX_ATTRIBUTE_VALUE_BYTES", 64 * 1024)
OTLP_RETRY_AFTER_SECONDS: Final = get_env_int("OTLP_RETRY_AFTER_SECONDS", 2)
OTLP_OFFLOAD_DECODE_BYTES: Final = get_env_int("OTLP_OFFLOAD_DECODE_BYTES", 256 * 1024)
AGENT_TRACING_INPUT_PREVIEW_CHARS: Final = get_env_int("AGENT_TRACING_INPUT_PREVIEW_CHARS", 240)
AGENT_TRACING_LIST_PAGE_SIZE: Final = get_env_int("AGENT_TRACING_LIST_PAGE_SIZE", 50)
DEFAULT_S3_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_S3_FLUSH_INTERVAL_SECONDS", 10))
DEFAULT_S3_BATCH_SIZE: Final = int(os.getenv("DEFAULT_S3_BATCH_SIZE", 512))
DEFAULT_S3_MAX_CONCURRENT_UPLOADS: Final = int(os.getenv("DEFAULT_S3_MAX_CONCURRENT_UPLOADS", "16"))

View file

@ -2008,7 +2008,6 @@ def response_cost_calculator(
else:
if isinstance(response_object, BaseModel):
if hasattr(response_object, "_hidden_params"):
response_object._hidden_params["optional_params"] = optional_params
provider_response_cost: Final = get_response_cost_from_hidden_params(response_object._hidden_params)
if provider_response_cost is not None:
return provider_response_cost

View file

@ -0,0 +1,100 @@
"""
Shared base for everything LiteLLM writes to ClickHouse.
Built on `CustomBatchLogger`: rows accumulate in `log_queue` and are flushed as one
gzip JSONEachRow insert, either every `CLICKHOUSE_FLUSH_INTERVAL_SECONDS` or as soon as
`batch_size` rows are queued. Subclasses only pick the table and build rows:
- `ClickHouseSpendLogger` -> spend_logs (LiteLLM requests, via the `clickhouse` callback)
"""
import asyncio
import os
from typing import Any, ClassVar
from litellm._logging import verbose_logger
from litellm.constants import (
CLICKHOUSE_BATCH_SIZE,
CLICKHOUSE_FLUSH_INTERVAL_SECONDS,
CLICKHOUSE_MAX_BUFFERED_ROWS,
CLICKHOUSE_MAX_RETRIES,
)
from litellm.integrations.custom_batch_logger import CustomBatchLogger
from litellm.rust_bridge.traces import TraceStorage
def clickhouse_storage_from_env() -> TraceStorage:
return TraceStorage(
database=os.getenv("CLICKHOUSE_DATABASE", "litellm"),
url=os.getenv("CLICKHOUSE_URL", ""),
)
class ClickHouseBatchLogger(CustomBatchLogger):
table: ClassVar[str]
def __init__(self, storage: TraceStorage | None = None) -> None:
self.storage = storage or clickhouse_storage_from_env()
self.rows_written = 0
self.rows_dropped = 0
self._failed_attempts = 0
super().__init__(
flush_lock=asyncio.Lock(),
batch_size=CLICKHOUSE_BATCH_SIZE,
flush_interval=CLICKHOUSE_FLUSH_INTERVAL_SECONDS,
)
try:
asyncio.get_running_loop().create_task(self.periodic_flush())
except RuntimeError: # no loop yet (e.g. sync config load); proxy startup calls start()
pass
def start(self) -> None:
asyncio.get_running_loop().create_task(self.periodic_flush())
def is_full(self) -> bool:
"""Backpressure signal: producers should reject (429) instead of enqueueing."""
return len(self.log_queue) >= CLICKHOUSE_MAX_BUFFERED_ROWS
def enqueue(self, rows: list[dict[str, Any]]) -> None:
"""Never awaits ClickHouse. Kicks off an early flush once a full batch is queued."""
self.log_queue.extend(rows)
if len(self.log_queue) >= self.batch_size:
asyncio.get_running_loop().create_task(self.flush_queue())
async def flush_queue(self) -> None:
# Swap the queue under the lock so rows enqueued during the insert are kept.
if self.flush_lock is None:
return
async with self.flush_lock:
while self.log_queue:
batch = self.log_queue[: self.batch_size]
self.log_queue = self.log_queue[len(batch) :]
if not await self._insert(batch):
break
async def async_send_batch(self) -> None:
await self.flush_queue()
async def _insert(self, batch: list[dict[str, Any]]) -> bool:
try:
await self.storage.insert_rows(self.table, batch)
self.rows_written += len(batch)
self._failed_attempts = 0
return True
except Exception as e:
self._failed_attempts += 1
if self._failed_attempts >= CLICKHOUSE_MAX_RETRIES:
self.rows_dropped += len(batch)
self._failed_attempts = 0
verbose_logger.error(
"ClickHouse: dropped %s rows for %s after %s attempts: %s",
len(batch),
self.table,
CLICKHOUSE_MAX_RETRIES,
e,
)
else:
# put it back; the next periodic flush retries it
self.log_queue = batch + self.log_queue
verbose_logger.warning("ClickHouse: insert into %s failed, will retry: %s", self.table, e)
return False

View file

@ -0,0 +1,11 @@
from typing import Final
from litellm.rust_bridge.traces import TraceStorage
OTEL_TRACES_TABLE: Final = "otel_traces"
AGENT_TRACES_BY_KEY_TABLE: Final = "agent_traces_by_key"
SPEND_LOGS_TABLE: Final = "spend_logs"
async def ensure_schema(storage: TraceStorage, trace_retention_days: int, spend_log_retention_days: int) -> None:
await storage.ensure_schema(trace_retention_days, spend_log_retention_days)

View file

@ -27,7 +27,7 @@ class CustomBatchLogger(CustomLogger):
self,
flush_lock: asyncio.Lock | None = None,
batch_size: int | None = None,
flush_interval: int | None = None,
flush_interval: float | None = None,
max_queue_size: int | None = None,
**kwargs,
) -> None:

View file

@ -213,6 +213,15 @@ nothing here imports outside it:
`config.yaml` — the latter reach the config through the logger's constructor
kwargs. `baggage_team_metadata_keys` is empty by default, so none of a team's
free-form metadata is promoted until each sub-key is explicitly allowlisted.
`excluded_services` withholds datastore spans from key/team `callback_vars`
destinations while the operator's own exporters keep them: set
`LITELLM_OTEL_EXCLUDED_SERVICES` (comma-separated) or `excluded_services`
(a YAML list) under `callback_settings.otel`, naming the datastore services
to withhold (`redis`, `postgres`, `batch_write_to_db`, `redis_*`, or their
`db.system.name` spellings `redis` / `postgresql`). Unknown names are logged
as an error and ignored. A span is withheld when its `db.system.name` /
`db.system` attribute is in the set, so request root, auth, guardrail and
model spans can never be excluded.
- [`baggage.py`](./model/baggage.py) — the single definition of which request-identity
values are promoted into Baggage (so child spans inherit them) and under which
attribute keys.

View file

@ -30,7 +30,7 @@ from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.otel.emitter import SpanEmitter, stamp_error
from litellm.integrations.otel.mappers import resolve_mappers
from litellm.integrations.otel.model.baggage import promoted_baggage
from litellm.integrations.otel.model.config import OpenTelemetryV2Config
from litellm.integrations.otel.model.config import OpenTelemetryV2Config, excluded_db_systems_from
from litellm.integrations.otel.model.metadata import (
LLMCallEvent,
RequestIdentity,
@ -898,12 +898,29 @@ def publish_global_otel_v2_provider(
"""
global _published_v2_provider
logger: Final = select_global_otel_v2_logger(in_memory_loggers, registered=registered)
attach_tenant_fan_out(logger.tracer_provider, *_v2_configs(in_memory_loggers, logger))
attach_tenant_fan_out(
logger.tracer_provider,
*_v2_configs(in_memory_loggers, logger),
excluded_db_systems=_excluded_db_systems(logger),
)
set_global_provider(logger.tracer_provider)
_published_v2_provider = logger.tracer_provider # rebind-ok: startup records the one provider carrying the fan-out
return logger
def _excluded_db_systems(logger: "OpenTelemetryV2") -> frozenset[str]:
"""The datastore services withheld from tenant destinations.
``callback_settings.otel.excluded_services`` wins over the env var whichever
logger got published: with ``callbacks: [langfuse_otel, otel]`` the ``otel``
callback folds into the preset, whose config is env-only.
"""
configured: Final = litellm.callback_settings.get("otel", {}).get("excluded_services")
if configured is None:
return logger.config.excluded_services
return excluded_db_systems_from(configured)
def _v2_configs(in_memory_loggers: Sequence[object], logger: "OpenTelemetryV2") -> tuple[OpenTelemetryV2Config, ...]:
"""Every v2 logger's config, the published logger's first.
@ -963,7 +980,11 @@ def fan_out_provider() -> ApiTracerProvider:
return published
logger: Final = _registered_v2_logger()
if logger is not None:
attach_tenant_fan_out(logger.tracer_provider, logger.config)
attach_tenant_fan_out(
logger.tracer_provider,
logger.config,
excluded_db_systems=_excluded_db_systems(logger),
)
return logger.tracer_provider
return get_tracer_provider()

View file

@ -4,14 +4,16 @@ from enum import Enum
from functools import lru_cache
from typing import Annotated, Any, Final
from pydantic import AliasChoices, BaseModel, Field, field_validator, model_validator
from pydantic import AliasChoices, BaseModel, Field, TypeAdapter, ValidationError, field_validator, model_validator
from pydantic_settings import BaseSettings, NoDecode, SettingsConfigDict
from litellm._logging import verbose_logger
from litellm.integrations.otel.model.baggage import (
BAGGAGE_PROMOTED_KEYS,
DEFAULT_BAGGAGE_METADATA_KEYS,
DEFAULT_BAGGAGE_TEAM_METADATA_KEYS,
)
from litellm.integrations.otel.model.spans import POSTGRESQL, db_system
from litellm.types.utils import OtelSpanScope
#: Master feature-flag env var. The logger is inert until this is truthy.
@ -174,6 +176,19 @@ class OpenTelemetryV2Config(BaseSettings):
"key/team destinations are not affected."
),
)
excluded_services: Annotated[frozenset[str], NoDecode] = Field(
default_factory=frozenset,
validation_alias=AliasChoices("excluded_services", "LITELLM_OTEL_EXCLUDED_SERVICES"),
description=(
"Datastore services whose spans are withheld from key/team ``callback_vars`` "
"OTel destinations (the operator's own exporters still receive them). Accepted "
"values are the datastore ``ServiceTypes`` names (``redis``, ``postgres``, "
"``batch_write_to_db``, ``redis_*``) or their ``db.system.name`` spellings "
"(``redis``, ``postgresql``); stored normalized to ``db.system.name`` values. "
"Configure via the ``LITELLM_OTEL_EXCLUDED_SERVICES`` env var (comma-separated) "
"or ``callback_settings.otel.excluded_services`` in config.yaml (a YAML list)."
),
)
# ----- explicit multi-destination / vocabulary configuration ------------ #
@ -284,6 +299,11 @@ class OpenTelemetryV2Config(BaseSettings):
return [item.strip() for item in value.split(",") if item.strip()]
return value
@field_validator("excluded_services", mode="before")
@classmethod
def _read_excluded_services(cls, value: object) -> frozenset[str]:
return excluded_service_names(value)
@model_validator(mode="after")
def _normalize(self) -> "OpenTelemetryV2Config":
# An endpoint with the default exporter kind implies OTLP/HTTP.
@ -316,6 +336,7 @@ class OpenTelemetryV2Config(BaseSettings):
if self.legacy_compat and "legacy" not in names:
names.append("legacy")
self.mapper_names = names
self.excluded_services = _normalize_excluded_services(self.excluded_services)
return self
@property
@ -334,3 +355,55 @@ class OpenTelemetryV2Config(BaseSettings):
@classmethod
def from_env(cls) -> "OpenTelemetryV2Config":
return cls()
_EXCLUDED_SERVICES_INPUT: Final[TypeAdapter[str | tuple[object, ...]]] = TypeAdapter(str | tuple[object, ...])
def excluded_db_systems_from(value: object) -> frozenset[str]:
"""Normalize a raw ``excluded_services`` value without building a settings model that rereads the env"""
return _normalize_excluded_services(excluded_service_names(value))
def excluded_service_names(value: object) -> frozenset[str]:
"""Read a YAML list or comma-separated string of service names, logging and dropping unusable input
so a malformed value cannot stop the OTel logger from being built"""
if value is None:
return frozenset()
try:
parsed: Final = _EXCLUDED_SERVICES_INPUT.validate_python(value)
except ValidationError:
verbose_logger.error("excluded_services must be a list or comma-separated string; %r ignored", value)
return frozenset()
items: Final = tuple(parsed.split(",")) if isinstance(parsed, str) else parsed
return frozenset(name for item in items if (name := _service_name(item)))
def _service_name(item: object) -> str:
if not isinstance(item, str):
verbose_logger.error("excluded_services must be a list of service names; %r ignored", item)
return ""
return item.strip().lower()
def _normalize_excluded_services(services: frozenset[str]) -> frozenset[str]:
"""Fold each accepted spelling to its ``db.system.name`` value.
``postgres`` and ``postgresql`` name the same system, as do every
``ServiceTypes`` member that ``db_system`` maps. Anything else means the
operator pointed the setting at a span family it cannot cover; those names
are logged and dropped so a typo cannot take the proxy down.
"""
resolved: Final = frozenset(
system for service in services if (system := _db_system_for_excluded_service(service)) is not None
)
return resolved
def _db_system_for_excluded_service(service: str) -> str | None:
resolved: Final = db_system(service) if service != POSTGRESQL else POSTGRESQL
if resolved is None:
verbose_logger.error(
"excluded_services: %r is not a datastore service; ignored. Allowed: postgres, redis", service
)
return resolved

View file

@ -418,6 +418,13 @@ def _is_database_span(attributes: Mapping[str, AttributeValue]) -> bool:
return any(key in attributes for key in _DB_SYSTEM_KEYS)
def _is_excluded_database_span(attributes: Mapping[str, AttributeValue], excluded: frozenset[str]) -> bool:
if not excluded:
return False
system: Final = attributes.get(DB.SYSTEM_NAME) or attributes.get(DB.SYSTEM_LEGACY)
return isinstance(system, str) and system in excluded
def _is_tenant_owned_span(attributes: Mapping[str, AttributeValue]) -> bool:
return any(key in attributes for key in _TENANT_OWNED_KEYS)
@ -549,10 +556,12 @@ class TenantFanOutSpanProcessor(SpanProcessor):
processor_factory: 'Callable[["OtelDestination"], SpanProcessor | None] | None' = None,
shutdown_drain_seconds: float = _SHUTDOWN_DRAIN_SECONDS,
operator_sinks: 'Mapping[_SinkKey, "OtelSpanScope"]' = MappingProxyType({}),
excluded_db_systems: frozenset[str] = frozenset(),
pending_drains: int = _MAX_PENDING_DRAINS,
drain_pool: _DrainPool | None = None,
) -> None:
self._operator_sinks: Final = operator_sinks
self._excluded_db_systems: Final = excluded_db_systems
self._drain_seconds: Final = shutdown_drain_seconds
self._lock: Final = threading.Condition()
self._closed = False # guarded by ``_lock``: an unlocked read races the teardown it gates
@ -567,9 +576,12 @@ class TenantFanOutSpanProcessor(SpanProcessor):
def on_end(self, span: ReadableSpan) -> None:
suppressed: Final = suppressed_backends()
attributes: Final = span.attributes or _NO_ATTRIBUTES
for destination in request_destinations():
if self._operator_already_writes(span, destination, suppressed) or not _in_scope(
span, destination.span_scope
if (
self._operator_already_writes(span, destination, suppressed)
or not _in_scope(span, destination.span_scope)
or _is_excluded_database_span(attributes, self._excluded_db_systems)
):
continue
processor = self._acquire(destination)
@ -1155,7 +1167,9 @@ def build_tracer_provider(
_FAN_OUT_ATTACH_LOCK: Final = threading.Lock()
def attach_tenant_fan_out(provider: TracerProvider, *configs: OpenTelemetryV2Config) -> None:
def attach_tenant_fan_out(
provider: TracerProvider, *configs: OpenTelemetryV2Config, excluded_db_systems: frozenset[str] = frozenset()
) -> None:
"""Give ``provider`` the fan-out that delivers spans to key/team destinations.
Called on the one provider published as the OTel global, and idempotent so a
@ -1164,12 +1178,18 @@ def attach_tenant_fan_out(provider: TracerProvider, *configs: OpenTelemetryV2Con
so exactly one fan-out lands. ``configs`` name the operator's own exporters, one
config per v2 logger since each keeps its own provider and still writes its
account, so an additive destination pointing at any of them is delivered once
rather than twice.
rather than twice. ``excluded_db_systems`` only filters what the fan-out
delivers, never the operator's own exporters.
"""
with _FAN_OUT_ATTACH_LOCK:
if any(isinstance(processor, TenantFanOutSpanProcessor) for processor in _attached_processors(provider)):
return
provider.add_span_processor(TenantFanOutSpanProcessor(operator_sinks=operator_sink_scopes(*configs)))
provider.add_span_processor(
TenantFanOutSpanProcessor(
operator_sinks=operator_sink_scopes(*configs),
excluded_db_systems=excluded_db_systems,
)
)
def deliverable_destinations(

View file

@ -2015,7 +2015,11 @@ def is_unsignable_thinking_block(block: object) -> bool:
return not (isinstance(thinking_text, str) and len(thinking_text.strip()) > 0)
def strip_encrypted_reasoning_from_messages(messages: object) -> None:
def strip_encrypted_reasoning_from_messages(
messages: object,
*,
should_strip: Callable[[Mapping[str, object]], bool] | None = None,
) -> None:
"""Drop the bridge-tagged reasoning blocks a routed deployment cannot decrypt from
Anthropic-shaped history.
@ -2030,7 +2034,7 @@ def strip_encrypted_reasoning_from_messages(messages: object) -> None:
if not isinstance(messages, list):
return
for content in anthropic_content_lists(cast(list[object], messages)): # cast-ok: untyped client json
_strip_encrypted_reasoning_from_blocks(content)
_strip_encrypted_reasoning_from_blocks(content, should_strip=should_strip)
def anthropic_content_lists(messages: Sequence[object]) -> Iterator[object]:
@ -2043,9 +2047,18 @@ def anthropic_content_lists(messages: Sequence[object]) -> Iterator[object]:
)
def _strip_encrypted_reasoning_from_blocks(content: object) -> None:
def _strip_encrypted_reasoning_from_blocks(
content: object,
*,
should_strip: Callable[[Mapping[str, object]], bool] | None = None,
) -> None:
blocks: Final = cast(list[object], content) # cast-ok: narrowed by the caller's isinstance
kept: Final = tuple(block for block in blocks if not is_encrypted_reasoning_block(block))
kept: Final = tuple(
block
for block in blocks
if not is_encrypted_reasoning_block(block)
or (should_strip is not None and not should_strip(cast(Mapping[str, object], block)))
)
blocks[:] = kept

View file

@ -23,7 +23,7 @@ def should_normalize_reasoning_content(field: object, *, model: str, provider: s
def normalize_reasoning_content(
messages: Sequence[AllMessageValues], *, forward: bool = True, normalize: bool = True
messages: Sequence[AllMessageValues], *, forward: bool = True, normalize: bool = True, strings_only: bool = False
) -> list[AllMessageValues]: # mutable-ok: provider request contract
def normalize_message(message: AllMessageValues) -> AllMessageValues:
if message["role"] != "assistant":
@ -40,7 +40,10 @@ def normalize_reasoning_content(
**MappingProxyType({key: value for key, value in history.items() if key not in removed_fields}),
**(
MappingProxyType({"reasoning": reasoning})
if normalize and forward and reasoning is not None
if normalize
and forward
and reasoning is not None
and (not strings_only or isinstance(reasoning, str))
else MappingProxyType({})
),
}

View file

@ -157,7 +157,8 @@ class HostedVLLMChatConfig(OpenAIGPTConfig):
) -> dict[str, object]: # mutable-ok: provider request contract
request_messages: Final = normalize_reasoning_content(
messages,
forward=litellm_params.get("forward_reasoning_content") is True,
forward=litellm_params.get("forward_reasoning_content") is not False,
strings_only=True,
normalize=should_normalize_reasoning_content(
litellm_params.get("reasoning_content_field"), model=model, provider="hosted_vllm"
),
@ -195,12 +196,14 @@ class HostedVLLMChatConfig(OpenAIGPTConfig):
"""
Support translating:
- video files from file_id or file_data to video_url
- thinking_blocks on assistant messages are removed,
and content lists are converted to strings for vLLM compatibility
- thinking_blocks and non-string reasoning_content on assistant messages
are removed, and content lists are converted to strings for vLLM compatibility
"""
for message in messages:
if message["role"] == "assistant":
message.pop("thinking_blocks", None)
if not isinstance(message.get("reasoning_content"), str):
message.pop("reasoning_content", None)
existing_content = message.get("content")
if isinstance(existing_content, list):
text_parts = []

View file

@ -53,6 +53,10 @@ from litellm.llms.base_llm.guardrail_translation.base_translation import (
)
from litellm.llms.base_llm.guardrail_translation.utils import (
blocked_responses_stream_usage,
effective_skip_system_message_for_guardrail,
merge_guardrailed_scoped_messages,
role_out_of_guardrail_scope,
scoped_structured_message_indices,
stream_item_field,
stream_item_fingerprint,
stream_item_items,
@ -376,6 +380,17 @@ class _RequestFields(NamedTuple):
class _ExtractedInputs(NamedTuple):
inputs: GenericGuardrailAPIInputs
task_mappings: tuple[tuple[int, int | None], ...]
instructions: str | None
def scannable_instructions(data: Mapping[str, object], *, skip_system: bool = False) -> str | None:
instructions: Final = data.get("instructions")
return instructions if isinstance(instructions, str) and instructions and not skip_system else None
def _input_item_role(item: object) -> str:
role: Final = item.get("role") if isinstance(item, Mapping) else None
return role.lower() if isinstance(role, str) else ""
def _patched_request_fields(
@ -494,7 +509,14 @@ class OpenAIResponsesHandler(BaseTranslation):
input_data: Final[str | ResponseInputParam | None] = data.get("input")
if not isinstance(input_data, (str, list)):
return data
skip_system: Final = effective_skip_system_message_for_guardrail(guardrail_to_apply)
structured_messages: Final = self.get_structured_messages(data)
scoped_indices: Final = scoped_structured_message_indices(
structured_messages or [], scan_only_tool_results=False, skip_system=skip_system, skip_tool=False
)
scoped_structured_messages: Final = (
[structured_messages[index] for index in scoped_indices] if structured_messages else None
)
raw_tools: Final = data.get("tools")
original_tools: Final[tuple[Mapping[str, object], ...]] = (
tuple(raw_tools) if isinstance(raw_tools, list) else ()
@ -502,11 +524,13 @@ class OpenAIResponsesHandler(BaseTranslation):
flattened_tool_groups: Final = tuple(
form.chat_tools for form in LiteLLMCompletionResponsesConfig.responses_tools_to_chat_forms(original_tools)
)
extracted: Final = self._extract_guardrail_inputs(data, input_data, flattened_tool_groups)
extracted: Final = self._extract_guardrail_inputs(
data, input_data, flattened_tool_groups, skip_system=skip_system
)
if not extracted.inputs.get("texts"):
return data
if structured_messages:
extracted.inputs["structured_messages"] = structured_messages
if scoped_structured_messages:
extracted.inputs["structured_messages"] = scoped_structured_messages
guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail(
inputs=extracted.inputs,
request_data=data,
@ -516,37 +540,63 @@ class OpenAIResponsesHandler(BaseTranslation):
self._apply_guardrailed_tools_to_data(
data, original_tools, flattened_tool_groups, guardrailed_inputs.get("tools")
)
written_back: Final = self._written_back_request_fields(data, structured_messages, guardrailed_inputs)
written_back: Final = self._written_back_request_fields(
data,
structured_messages or (),
scoped_indices,
scoped_structured_messages,
guardrail_to_apply,
guardrailed_inputs,
)
if written_back is not None:
data["input"] = list(written_back.input) # mutable-ok: JSON body
if written_back.instructions is None:
data.pop("instructions", None)
else:
data["instructions"] = written_back.instructions # rebind-ok: data is an out-param
elif isinstance(input_data, str):
guardrailed_texts: Final = guardrailed_inputs.get("texts") or ()
if len(guardrailed_texts) > 1:
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
data["input"] = guardrailed_texts[0] if guardrailed_texts else input_data # rebind-ok: data is an out-param
else:
rewritten_texts: Final = guardrailed_inputs.get("texts") or ()
if len(rewritten_texts) != len(extracted.task_mappings):
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
await self._apply_guardrail_responses_to_input(
messages=input_data,
responses=rewritten_texts,
task_mappings=extracted.task_mappings,
)
await self._apply_guardrailed_texts(data, input_data, extracted, guardrail_to_apply, guardrailed_inputs)
verbose_proxy_logger.debug("OpenAI Responses API: Processed input messages: %s", data.get("input"))
return data
async def _apply_guardrailed_texts(
self,
data: dict[str, object],
input_data: "str | ResponseInputParam",
extracted: _ExtractedInputs,
guardrail_to_apply: "CustomGuardrail",
guardrailed_inputs: GenericGuardrailAPIInputs,
) -> None:
returned_texts: Final = guardrailed_inputs.get("texts")
if not returned_texts:
return
rewritten_texts: Final = tuple(returned_texts)
offset: Final = 0 if extracted.instructions is None else 1
input_texts: Final = rewritten_texts[offset:]
expected: Final = 1 if isinstance(input_data, str) else len(extracted.task_mappings)
if len(rewritten_texts) != offset + expected:
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
if offset:
data["instructions"] = rewritten_texts[0] # rebind-ok: data is an out-param
if isinstance(input_data, str):
data["input"] = input_texts[0] # rebind-ok: data is an out-param
return
await self._apply_guardrail_responses_to_input(
messages=input_data,
responses=input_texts,
task_mappings=extracted.task_mappings,
)
def _extract_guardrail_inputs(
self,
data: Mapping[str, object],
input_data: "str | ResponseInputParam",
flattened_tool_groups: Sequence[Sequence[Mapping[str, object]]],
*,
skip_system: bool = False,
) -> _ExtractedInputs:
texts_to_check: Final[list[str]] = []
instructions: Final = scannable_instructions(data, skip_system=skip_system)
texts_to_check: Final[list[str]] = [] if instructions is None else [instructions]
images_to_check: Final[list[str]] = []
task_mappings: Final[list[tuple[int, int | None]]] = []
tools_to_check: Final[list[ChatCompletionToolParam]] = list( # mutable-ok: guardrail inputs want a list
@ -562,6 +612,10 @@ class OpenAIResponsesHandler(BaseTranslation):
texts_to_check.append(input_data)
else:
for msg_idx, message in enumerate(input_data):
if role_out_of_guardrail_scope(
_input_item_role(message), skip_system_message=skip_system, skip_tool_message=False
):
continue
self._extract_input_text_and_images(
message=message,
msg_idx=msg_idx,
@ -577,22 +631,32 @@ class OpenAIResponsesHandler(BaseTranslation):
model: Final = data.get("model")
if isinstance(model, str):
inputs["model"] = model
return _ExtractedInputs(inputs=inputs, task_mappings=tuple(task_mappings))
return _ExtractedInputs(inputs=inputs, task_mappings=tuple(task_mappings), instructions=instructions)
@staticmethod
def _written_back_request_fields(
data: Mapping[str, object],
structured_messages: Sequence[AllMessageValues] | None,
structured_messages: Sequence[AllMessageValues],
scoped_indices: Sequence[int],
scoped_structured_messages: Sequence[AllMessageValues] | None,
guardrail_to_apply: "CustomGuardrail",
guardrailed_inputs: GenericGuardrailAPIInputs,
) -> _RequestFields | None:
guardrailed: Final = guardrailed_inputs.get("structured_messages")
if guardrailed is None or guardrailed is structured_messages:
if guardrailed is None or guardrailed is scoped_structured_messages:
return None
covers_full_request: Final = len(scoped_indices) == len(structured_messages) or (
guardrail_to_apply.structured_messages_cover_full_request() and len(guardrailed) == len(structured_messages)
)
merged: Final = (
guardrailed
if covers_full_request
else merge_guardrailed_scoped_messages(
full_messages=structured_messages, scoped_indices=scoped_indices, guardrailed_scoped=guardrailed
)
)
return _patch_or_convert_request_fields(
data.get("input"),
data.get("instructions"),
structured_messages or (),
guardrailed,
data.get("input"), data.get("instructions"), structured_messages, merged
)
def extract_request_tool_names(self, data: dict) -> list[str]:

View file

@ -3358,7 +3358,7 @@
"supports_function_calling": true
},
"azure_ai/claude-haiku-4-5": {
"deprecation_date": "2026-10-19",
"deprecation_date": "2026-11-15",
"cache_creation_input_token_cost": 1.25e-06,
"cache_creation_input_token_cost_above_1hr": 2e-06,
"cache_read_input_token_cost": 1e-07,
@ -3378,10 +3378,11 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"prompt_cache_min_tokens": 4096
"prompt_cache_min_tokens": 4096,
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/model-retirement-schedule"
},
"azure_ai/claude-opus-4-5": {
"deprecation_date": "2026-10-19",
"deprecation_date": "2026-11-24",
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -3402,7 +3403,8 @@
"supports_tool_choice": true,
"supports_vision": true,
"supports_output_config": true,
"prompt_cache_min_tokens": 4096
"prompt_cache_min_tokens": 4096,
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/model-retirement-schedule"
},
"azure_ai/claude-opus-4-6": {
"deprecation_date": "2027-02-02",
@ -3640,7 +3642,7 @@
"prompt_cache_min_tokens": 1024
},
"azure_ai/claude-sonnet-4-5": {
"deprecation_date": "2026-10-19",
"deprecation_date": "2026-11-15",
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
@ -3660,7 +3662,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"prompt_cache_min_tokens": 1024
"prompt_cache_min_tokens": 1024,
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/model-retirement-schedule"
},
"azure_ai/claude-sonnet-5": {
"deprecation_date": "2027-06-30",
@ -30721,6 +30724,7 @@
"output_cost_per_image": 0.08
},
"gemini/veo-3.1-fast-generate-preview": {
"deprecation_date": "2026-10-22",
"litellm_provider": "gemini",
"max_input_tokens": 1024,
"max_tokens": 1024,
@ -30737,6 +30741,7 @@
]
},
"gemini/veo-3.1-generate-preview": {
"deprecation_date": "2026-10-22",
"litellm_provider": "gemini",
"max_input_tokens": 1024,
"max_tokens": 1024,
@ -30752,6 +30757,7 @@
]
},
"gemini/veo-3.1-lite-generate-preview": {
"deprecation_date": "2026-10-22",
"litellm_provider": "gemini",
"max_input_tokens": 1024,
"max_tokens": 1024,
@ -32912,10 +32918,13 @@
"gpt-image-2.5-flare": {
"cache_read_input_image_token_cost": 2e-06,
"cache_read_input_token_cost": 1.25e-06,
"cache_read_input_token_cost_batches": 6.25e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "openai",
"mode": "image_generation",
"input_cost_per_image_token": 8e-06,
"input_cost_per_image_token_batches": 4e-06,
"input_cost_per_token_batches": 2.5e-06,
"output_cost_per_image_token": 3e-05,
"supported_endpoints": [
"/v1/images/generations",
@ -32944,10 +32953,13 @@
"gpt-image-2.5-sunburst": {
"cache_read_input_image_token_cost": 2e-06,
"cache_read_input_token_cost": 1.25e-06,
"cache_read_input_token_cost_batches": 6.25e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "openai",
"mode": "image_generation",
"input_cost_per_image_token": 8e-06,
"input_cost_per_image_token_batches": 4e-06,
"input_cost_per_token_batches": 2.5e-06,
"output_cost_per_image_token": 3e-05,
"supported_endpoints": [
"/v1/images/generations",
@ -38611,6 +38623,7 @@
},
"mistral/zai-glm-5-2": {
"cache_read_input_token_cost": 1.4e-07,
"deprecation_date": "2026-10-31",
"input_cost_per_token": 1.4e-06,
"litellm_provider": "mistral",
"max_input_tokens": 1048576,
@ -38741,6 +38754,7 @@
"source": "https://mistral.ai/pricing#api-pricing"
},
"mistral/mistral-ocr-4-0": {
"deprecation_date": "2026-09-30",
"litellm_provider": "mistral",
"ocr_cost_per_page": 0.004,
"ocr_cost_per_page_batches": 0.002,
@ -60616,6 +60630,7 @@
"supports_vision": true
},
"mistral/labs-leanstral-1-5": {
"deprecation_date": "2026-09-30",
"input_cost_per_token": 0.0,
"litellm_provider": "mistral",
"max_input_tokens": 262144,
@ -61337,13 +61352,16 @@
},
"fireworks_ai/nemotron-lightning-3p5-30b-a3b": {
"cache_read_input_token_cost": 1e-08,
"cache_read_input_token_cost_priority": 1.25e-08,
"input_cost_per_token": 5e-08,
"input_cost_per_token_priority": 6.25e-08,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 262144,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 2e-07,
"output_cost_per_token_priority": 2.5e-07,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -61353,13 +61371,16 @@
},
"fireworks_ai/nemotron-3-ultra-nvfp4": {
"cache_read_input_token_cost": 1.2e-07,
"cache_read_input_token_cost_priority": 1.5e-07,
"input_cost_per_token": 6e-07,
"input_cost_per_token_priority": 7.5e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 262144,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 2.4e-06,
"output_cost_per_token_priority": 3e-06,
"source": "https://api.fireworks.ai/v1/serverless/models",
"supports_function_calling": true,
"supports_reasoning": true,
@ -61389,13 +61410,16 @@
},
"fireworks_ai/accounts/fireworks/models/nemotron-lightning-3p5-30b-a3b": {
"cache_read_input_token_cost": 1e-08,
"cache_read_input_token_cost_priority": 1.25e-08,
"input_cost_per_token": 5e-08,
"input_cost_per_token_priority": 6.25e-08,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 262144,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 2e-07,
"output_cost_per_token_priority": 2.5e-07,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -61405,13 +61429,16 @@
},
"fireworks_ai/accounts/fireworks/models/nemotron-3-ultra-nvfp4": {
"cache_read_input_token_cost": 1.2e-07,
"cache_read_input_token_cost_priority": 1.5e-07,
"input_cost_per_token": 6e-07,
"input_cost_per_token_priority": 7.5e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 262144,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 2.4e-06,
"output_cost_per_token_priority": 3e-06,
"source": "https://api.fireworks.ai/v1/serverless/models",
"supports_function_calling": true,
"supports_reasoning": true,
@ -64321,13 +64348,16 @@
},
"fireworks_ai/accounts/fireworks/routers/glm-5p3-us": {
"cache_read_input_token_cost": 3.9e-07,
"cache_read_input_token_cost_priority": 4.875e-07,
"input_cost_per_token": 2.1e-06,
"input_cost_per_token_priority": 2.625e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 6.6e-06,
"output_cost_per_token_priority": 8.25e-06,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -64356,13 +64386,16 @@
},
"fireworks_ai/glm-5p3-us": {
"cache_read_input_token_cost": 3.9e-07,
"cache_read_input_token_cost_priority": 4.875e-07,
"input_cost_per_token": 2.1e-06,
"input_cost_per_token_priority": 2.625e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 6.6e-06,
"output_cost_per_token_priority": 8.25e-06,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -64446,12 +64479,15 @@
},
"fireworks_ai/accounts/fireworks/routers/glm-5p3-flash-us": {
"cache_read_input_token_cost": 4.5e-08,
"cache_read_input_token_cost_priority": 5.625e-08,
"input_cost_per_token": 2.25e-07,
"input_cost_per_token_priority": 2.8125e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_tokens": 1048576,
"mode": "chat",
"output_cost_per_token": 7.5e-07,
"output_cost_per_token_priority": 9.375e-07,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_response_schema": true,
@ -64477,12 +64513,15 @@
},
"fireworks_ai/glm-5p3-flash-us": {
"cache_read_input_token_cost": 4.5e-08,
"cache_read_input_token_cost_priority": 5.625e-08,
"input_cost_per_token": 2.25e-07,
"input_cost_per_token_priority": 2.8125e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_tokens": 1048576,
"mode": "chat",
"output_cost_per_token": 7.5e-07,
"output_cost_per_token_priority": 9.375e-07,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_response_schema": true,
@ -77478,11 +77517,14 @@
},
"fireworks_ai/accounts/fireworks/models/ember-1": {
"cache_read_input_token_cost": 3e-07,
"cache_read_input_token_cost_priority": 3.75e-07,
"input_cost_per_token": 3e-06,
"input_cost_per_token_priority": 3.75e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"output_cost_per_token_priority": 1.875e-05,
"source": "https://api.fireworks.ai/v1/serverless/models",
"supports_function_calling": true,
"supports_reasoning": true,
@ -79350,5 +79392,33 @@
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true
},
"vertex_ai/gemini-3.8-flash-tts": {
"input_cost_per_token": 5e-07,
"litellm_provider": "vertex_ai",
"max_input_tokens": 8192,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "audio_speech",
"output_cost_per_audio_token": 9e-06,
"output_cost_per_token": 9e-06,
"source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing",
"supported_endpoints": [
"/v1/audio/speech"
]
},
"vertex_ai/gemini-3.8-flash-lite-tts": {
"input_cost_per_token": 5e-07,
"litellm_provider": "vertex_ai",
"max_input_tokens": 8192,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "audio_speech",
"output_cost_per_audio_token": 6e-06,
"output_cost_per_token": 6e-06,
"source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing",
"supported_endpoints": [
"/v1/audio/speech"
]
}
}

View file

@ -0,0 +1,74 @@
from types import MappingProxyType
from typing import Final
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_managed_agent_ceilings
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
from litellm.types.proxy.agent_identity import AgentIdentityFailure
async def _delegated_resource_subject(user_id: str) -> UserAPIKeyAuth:
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
human: Final = await MCPRequestHandler.reload_admitted_user(user_id, requires_fresh_policy=True)
return human.model_copy(update=MappingProxyType({"mcp_explicit_grants_only": True}))
async def managed_agent_servers(auth: UserAPIKeyAuth) -> tuple[str, ...]:
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
agent: Final = auth.managed_agent_policy
if agent is None:
return ()
try:
base: Final = frozenset(await MCPRequestHandler.get_allowed_mcp_servers_for_agent(auth))
ceilings: Final = await resolve_managed_agent_ceilings(agent)
expanded: Final = tuple(
frozenset(global_mcp_server_manager.expand_permission_list(sorted(ceiling.mcp_server_ids)))
for ceiling in ceilings
)
grouped: Final = frozenset(server for server in base if all(server in ceiling for ceiling in expanded))
caller_capped, _ = await MCPRequestHandler.apply_agent_caller_ceiling(sorted(grouped), auth)
own: Final = frozenset(caller_capped)
context: Final = auth.managed_agent_context
if context is None or context.mode == "autonomous":
return tuple(sorted(own))
if context.user_id is None:
return ()
human: Final = await _delegated_resource_subject(context.user_id)
allowed: Final = await MCPRequestHandler.resolve_admitted_subject_servers(
human, allowed_team_ids=frozenset((auth.team_id,)) if auth.team_id else frozenset()
)
return tuple(sorted(own.intersection(allowed)))
except Exception: # noqa: BLE001 # Authorization boundary: every unresolved policy must deny access
raise_identity_failure(
AgentIdentityFailure(code="policy_unavailable", message="Agent MCP policy is unavailable")
)
async def managed_agent_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None:
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
if server_id not in await managed_agent_servers(auth):
return []
try:
granted: Final = await MCPRequestHandler.get_agent_tool_permissions_for_server(server_id, auth)
own: Final = await MCPRequestHandler.apply_agent_caller_tool_ceiling(granted, server_id, auth)
context: Final = auth.managed_agent_context
if context is None or context.mode == "autonomous":
return None if own is None else sorted(own)
if context.user_id is None:
return []
human: Final = await _delegated_resource_subject(context.user_id)
human_tools: Final = await MCPRequestHandler.resolve_admitted_subject_tools(
server_id, human, allowed_team_ids=frozenset((auth.team_id,)) if auth.team_id else frozenset()
)
if own is None:
return human_tools
return sorted(own) if human_tools is None else sorted(frozenset(own).intersection(human_tools))
except Exception: # noqa: BLE001 # Authorization boundary: every unresolved policy must deny access
raise_identity_failure(
AgentIdentityFailure(code="policy_unavailable", message="Agent tool policy is unavailable")
)

View file

@ -51,6 +51,7 @@ from litellm.proxy.agent_endpoints.auth.agent_access_groups import (
resolve_agent_access_group_ceiling,
)
from litellm.proxy.agent_endpoints.auth.agent_caller import agent_caller_auth
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_policy
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
from litellm.proxy.auth.user_api_key_auth import (
_get_bearer_token_or_received_api_key, # pyright: ignore[reportPrivateUsage] # shared x-litellm-api-key parser lives with user_api_key_auth
@ -67,7 +68,6 @@ from litellm.repositories.table_repositories import (
AgentsRepository,
MCPServerRepository,
)
from litellm.repositories.user_repository import UserRepository
from litellm.types.mcp_server.mcp_server_manager import MCPServer
if TYPE_CHECKING:
@ -1086,7 +1086,7 @@ class MCPRequestHandler:
assert_never(identity.subject_type)
@staticmethod
async def reload_admitted_user(user_id: str) -> UserAPIKeyAuth:
async def reload_admitted_user(user_id: str, *, requires_fresh_policy: bool = False) -> UserAPIKeyAuth:
"""Reload the live user an interactively-minted envelope references and admit them as themselves.
The user's own object permission and ``org_id`` ride on the returned ``UserAPIKeyAuth``, and the
@ -1111,6 +1111,7 @@ class MCPRequestHandler:
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
check_db_only=requires_fresh_policy,
)
# Resolve the user's own MCP object permission (get_user_object does not load it) so the shared
# get_allowed_mcp_servers can grant the user their litellm-granted servers. Reuses the same
@ -1119,6 +1120,7 @@ class MCPRequestHandler:
if user_object is not None and object_permission is None and user_object.object_permission_id:
object_permission = await get_object_permission(
object_permission_id=user_object.object_permission_id,
check_db_only=requires_fresh_policy,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
@ -1147,6 +1149,7 @@ class MCPRequestHandler:
# Server-only marker, set AFTER construction: the before-validator strips it from any validated
# input, so caller-supplied data (key metadata, JWT claims) can never forge it.
admitted.mcp_admitted_user_subject = True
admitted.requires_fresh_policy = requires_fresh_policy
# Carry each granting team's per-server mcp_rpm_limit: this subject reaches servers through
# several teams under its own identity, so without this a cross-team user outruns every team's
# limit. Resolved from the same roster-checked sources as the grant union, so a team throttles
@ -1202,7 +1205,7 @@ class MCPRequestHandler:
return None
@staticmethod
async def _reload_admitted_key(key_hash: str) -> UserAPIKeyAuth:
async def _reload_admitted_key(key_hash: str, *, check_db_only: bool = False) -> UserAPIKeyAuth:
"""Reload the live key record an admitted envelope references and re-check live policy.
Resolving the current ``UserAPIKeyAuth`` (cache first, then DB) is what stops the
@ -1234,6 +1237,7 @@ class MCPRequestHandler:
hashed_token=key_hash,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
check_db_only=check_db_only,
)
except (ProxyException, HTTPException):
raise HTTPException(status_code=401, detail="Invalid or expired credential") from None
@ -1597,6 +1601,11 @@ class MCPRequestHandler:
"""
from litellm.proxy.proxy_server import general_settings
if managed_agent_policy(user_api_key_auth) is not None:
from litellm.proxy._experimental.mcp_server.auth.managed_agent_access import managed_agent_servers
return MCPServerAccess(server_ids=await managed_agent_servers(user_api_key_auth), scope="scoped")
key_object_permission: Final = MCPRequestHandler._get_key_object_permission(user_api_key_auth)
try:
@ -1606,7 +1615,7 @@ class MCPRequestHandler:
# independent; an opt-out silences only its own source, inside the recursive call).
if _is_mcp_admitted_user_subject(user_api_key_auth) and user_api_key_auth is not None:
return MCPServerAccess(
server_ids=tuple(await MCPRequestHandler._resolve_admitted_subject_servers(user_api_key_auth)),
server_ids=tuple(await MCPRequestHandler.resolve_admitted_subject_servers(user_api_key_auth)),
)
# Get allowed servers from key and team
@ -1703,7 +1712,7 @@ class MCPRequestHandler:
if user_api_key_auth and user_api_key_auth.agent_id:
agent_capped: Final = _agent_capped_servers(
allowed_mcp_servers,
await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth),
await MCPRequestHandler.get_allowed_mcp_servers_for_agent(user_api_key_auth),
await MCPRequestHandler._get_agent_access_group_server_ceiling(user_api_key_auth),
)
if agent_capped is not None:
@ -1716,7 +1725,7 @@ class MCPRequestHandler:
#########################################################
# Cap an agent key at what the user and team that invoked the agent may reach
#########################################################
caller_capped, caller_restricts = await MCPRequestHandler._apply_agent_caller_ceiling(
caller_capped, caller_restricts = await MCPRequestHandler.apply_agent_caller_ceiling(
allowed_mcp_servers, user_api_key_auth
)
@ -1829,10 +1838,14 @@ class MCPRequestHandler:
scoped.object_permission = auth.object_permission
scoped.object_permission_id = auth.object_permission_id
scoped.access_group_ids = auth.access_group_ids
scoped.requires_fresh_policy = auth.requires_fresh_policy
scoped.mcp_explicit_grants_only = auth.mcp_explicit_grants_only
return scoped
@staticmethod
async def _admitted_subject_sources(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]:
async def admitted_subject_sources(
auth: UserAPIKeyAuth, *, allowed_team_ids: frozenset[str] | None = None
) -> list[UserAPIKeyAuth]:
"""The independent sources a keyless admitted subject reaches MCP servers through: their own
direct grants, plus every team they are a live roster member of.
@ -1849,6 +1862,8 @@ class MCPRequestHandler:
if not auth.user_id or prisma_client is None:
return sources
for team_id in await MCPRequestHandler._resolve_user_team_ids(auth.user_id, auth):
if allowed_team_ids is not None and team_id not in allowed_team_ids:
continue
team_obj = await MCPRequestHandler._roster_team_object(team_id, auth)
if team_obj is None:
continue
@ -1886,6 +1901,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(auth and auth.requires_fresh_policy),
)
except Exception as e: # noqa: BLE001 # per-source isolation: one team's blip must not deny the others
# Fault isolation is per SOURCE: an unresolvable team contributes nothing (fail closed for
@ -1932,7 +1948,9 @@ class MCPRequestHandler:
return team_obj
@staticmethod
async def admitted_source_grants(auth: UserAPIKeyAuth) -> list[tuple[UserAPIKeyAuth, set[str]]]:
async def admitted_source_grants(
auth: UserAPIKeyAuth, *, allowed_team_ids: frozenset[str] | None = None
) -> list[tuple[UserAPIKeyAuth, set[str]]]:
"""``(source, the servers that source grants)`` for every source of an admitted subject.
THE owner of "which source reaches which server". The reachable union, the per-team throttle
@ -1941,15 +1959,17 @@ class MCPRequestHandler:
roster instead of by grant charged unrelated teams' buckets)."""
return [
(source, set(await MCPRequestHandler.get_allowed_mcp_servers(source, keyless_source=True)))
for source in await MCPRequestHandler._admitted_subject_sources(auth)
for source in await MCPRequestHandler.admitted_subject_sources(auth, allowed_team_ids=allowed_team_ids)
]
@staticmethod
async def _resolve_admitted_subject_servers(auth: UserAPIKeyAuth) -> list[str]:
async def resolve_admitted_subject_servers(
auth: UserAPIKeyAuth, *, allowed_team_ids: frozenset[str] | None = None
) -> list[str]:
"""Union of what each of the admitted subject's sources reaches, each answered by the
canonical resolver so no rule is reimplemented for this caller shape."""
reachable: Final[set[str]] = set()
for _source, granted in await MCPRequestHandler.admitted_source_grants(auth):
for _source, granted in await MCPRequestHandler.admitted_source_grants(auth, allowed_team_ids=allowed_team_ids):
reachable.update(granted)
return list(reachable)
@ -2007,7 +2027,9 @@ class MCPRequestHandler:
return min((source for source, _ in granting), key=lambda s: s.team_id or "")
@staticmethod
async def _resolve_admitted_subject_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None:
async def resolve_admitted_subject_tools(
server_id: str, auth: UserAPIKeyAuth, *, allowed_team_ids: frozenset[str] | None = None
) -> list[str] | None:
"""Effective tool allowlist on ``server_id`` for an admitted subject, as the union over the
sources that actually grant that server.
@ -2029,7 +2051,7 @@ class MCPRequestHandler:
) or await MCPRequestHandler.admin_view_unscoped(auth)
allowed: Final[set[str]] = set()
for source, granted in await MCPRequestHandler.admitted_source_grants(auth):
for source, granted in await MCPRequestHandler.admitted_source_grants(auth, allowed_team_ids=allowed_team_ids):
# The open channel is evaluated against the user's OWN source (team_id is None), so that
# source's restrictions apply to it; a team's rules never ride an open-channel server.
if server_id not in granted and not (reachable_via_open_channel and source.team_id is None):
@ -2088,6 +2110,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if not team_obj:
@ -2098,6 +2121,8 @@ class MCPRequestHandler:
@staticmethod
async def _toolset_tool_permissions(
object_permission: LiteLLM_ObjectPermissionTable | None,
*,
requires_fresh_policy: bool = False,
) -> Mapping[str, Sequence[str]]:
"""The ``server_id -> tool names`` grants of this permission row's toolsets, empty when it
declares none. The shared resolver for the team, org, and internal-user levels, so a toolset
@ -2114,7 +2139,8 @@ class MCPRequestHandler:
if object_permission is None or not object_permission.mcp_toolsets:
return _EMPTY_TOOLSET_GRANTS
resolved: Final = await global_mcp_server_manager.resolve_toolset_tool_permissions(
toolset_ids=object_permission.mcp_toolsets
toolset_ids=object_permission.mcp_toolsets,
requires_fresh_policy=requires_fresh_policy,
)
if not resolved:
raise UnloadableEntitlementError(
@ -2126,10 +2152,15 @@ class MCPRequestHandler:
async def _toolset_tools_for_server(
object_permission: LiteLLM_ObjectPermissionTable | None,
server_id: str,
*,
requires_fresh_policy: bool = False,
) -> Sequence[str] | None:
"""Tool names this row's toolsets grant on ``server_id``, ``None`` when its toolsets place
no restriction on that server (it declares no toolsets, or none of them name it)."""
return (await MCPRequestHandler._toolset_tool_permissions(object_permission)).get(server_id)
grants: Final = await MCPRequestHandler._toolset_tool_permissions(
object_permission, requires_fresh_policy=requires_fresh_policy
)
return grants.get(server_id)
@staticmethod
def _union_tool_grants(
@ -2171,6 +2202,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
@staticmethod
@ -2219,12 +2251,17 @@ class MCPRequestHandler:
if not user_api_key_auth:
return None
if managed_agent_policy(user_api_key_auth) is not None:
from litellm.proxy._experimental.mcp_server.auth.managed_agent_access import managed_agent_tools
return await managed_agent_tools(server_id, user_api_key_auth)
try:
# FIRST statement, mirroring get_allowed_mcp_servers: a keyless admitted subject resolves per
# source and shares nothing with the single-credential prelude below. Ordering is the invariant:
# sat after the prelude, a fault in a lookup the subject never uses denied tools its teams grant.
if _is_mcp_admitted_user_subject(user_api_key_auth):
return await MCPRequestHandler._resolve_admitted_subject_tools(server_id, user_api_key_auth)
return await MCPRequestHandler.resolve_admitted_subject_tools(server_id, user_api_key_auth)
# Get key and team object permissions (already loaded in main auth flow)
key_obj_perm: Final = MCPRequestHandler._get_key_object_permission(user_api_key_auth)
@ -2249,9 +2286,12 @@ class MCPRequestHandler:
# tool-level check sees the key's full effective tool scope
key_toolset_ids: Final = (key_obj_perm.mcp_toolsets or []) if key_obj_perm else []
key_toolset_tools: Final = (
(await global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=key_toolset_ids)).get(
server_id
)
(
await global_mcp_server_manager.resolve_toolset_tool_permissions(
toolset_ids=key_toolset_ids,
requires_fresh_policy=user_api_key_auth.requires_fresh_policy,
)
).get(server_id)
if key_toolset_ids
else None
)
@ -2265,7 +2305,9 @@ class MCPRequestHandler:
# Tools granted through the team's toolsets restrict this server exactly
# as the team's direct tool permissions do, mirroring the key path above
team_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(team_obj_perm, server_id)
team_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
team_obj_perm, server_id, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
)
team_tools: Final = MCPRequestHandler._union_tool_grants(team_direct_tools, team_toolset_tools)
# Apply same inheritance logic as get_allowed_mcp_servers
@ -2291,7 +2333,7 @@ class MCPRequestHandler:
)
allowed_tools = _as_list(
await MCPRequestHandler._apply_agent_caller_tool_ceiling(allowed_tools, server_id, user_api_key_auth)
await MCPRequestHandler.apply_agent_caller_tool_ceiling(allowed_tools, server_id, user_api_key_auth)
)
return await MCPRequestHandler._apply_agent_and_org_tool_ceilings(
@ -2334,7 +2376,7 @@ class MCPRequestHandler:
if user_api_key_auth.agent_id:
# Pre-fetch agent object_permission once to avoid a duplicate DB query.
agent_obj_perm: Final = await MCPRequestHandler._get_agent_object_permission(user_api_key_auth)
agent_tools: Final = await MCPRequestHandler._get_agent_tool_permissions_for_server(
agent_tools: Final = await MCPRequestHandler.get_agent_tool_permissions_for_server(
server_id=server_id,
user_api_key_auth=user_api_key_auth,
agent_object_permission=agent_obj_perm,
@ -2365,7 +2407,9 @@ class MCPRequestHandler:
if org_obj_perm and org_obj_perm.mcp_tool_permissions
else None
)
org_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(org_obj_perm, server_id)
org_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
org_obj_perm, server_id, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
)
org_tools: Final = MCPRequestHandler._union_tool_grants(org_direct_tools, org_toolset_tools)
if org_tools is not None:
allowed_tools = (
@ -2456,6 +2500,7 @@ class MCPRequestHandler:
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if not raw_server_ids:
return []
@ -2502,6 +2547,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if key_object_permission is None:
return []
@ -2518,7 +2564,8 @@ class MCPRequestHandler:
# Get MCP servers from access groups
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
key_object_permission.mcp_access_groups or []
key_object_permission.mcp_access_groups or [],
requires_fresh_policy=user_api_key_auth.requires_fresh_policy,
)
# servers referenced in tool permissions should also be accessible
@ -2531,7 +2578,14 @@ class MCPRequestHandler:
# ceilings as any other key-level grant
toolset_ids: Final = key_object_permission.mcp_toolsets or []
toolset_servers: Final = (
list((await global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=toolset_ids)).keys())
list(
(
await global_mcp_server_manager.resolve_toolset_tool_permissions(
toolset_ids=toolset_ids,
requires_fresh_policy=user_api_key_auth.requires_fresh_policy,
)
).keys()
)
if toolset_ids
else []
)
@ -2550,7 +2604,7 @@ class MCPRequestHandler:
"""Get allowed MCP servers a caller inherits from the team it is pinned to.
Exactly one team, or none. A subject that reaches servers through SEVERAL teams does not
fan out here: it is resolved one source per team in ``_resolve_admitted_subject_servers``,
fan out here: it is resolved one source per team in ``resolve_admitted_subject_servers``,
and each of those sources pins a single ``team_id`` before reaching this point. Keeping the
fan-out here as well would be a second multi-team path to drift from that one.
"""
@ -2568,7 +2622,7 @@ class MCPRequestHandler:
which must NOT silently gain the union across every team the user belongs to), and it covers
each single-source auth an admitted subject fans out into — those pin a team_id, so they land
on the first branch. The admitted subject itself never reaches here: it resolves per source
in ``_resolve_admitted_subject_servers`` before this point. The ``UI_TEAM_ID`` sentinel
in ``resolve_admitted_subject_servers`` before this point. The ``UI_TEAM_ID`` sentinel
resolves to no teams exactly as before."""
if user_api_key_auth is None or not user_api_key_auth.team_id:
return []
@ -2596,6 +2650,7 @@ class MCPRequestHandler:
user_id_upsert=False,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
except Exception as e: # noqa: BLE001 # a team-resolution blip narrows access, never raises
verbose_logger.warning("Failed to resolve user teams for MCP grant: %s", e)
@ -2605,7 +2660,12 @@ class MCPRequestHandler:
return list(dict.fromkeys(t for t in user_object.teams if t and t != UI_TEAM_ID))
@staticmethod
async def _team_granted_servers(team_obj: LiteLLM_TeamTable, team_access_group_servers: list[str]) -> set[str]:
async def _team_granted_servers(
team_obj: LiteLLM_TeamTable,
team_access_group_servers: list[str],
*,
requires_fresh_policy: bool = False,
) -> set[str]:
"""The raw MCP-server set a team grants (before any org ceiling): its object_permission (direct
``mcp_servers``, the ``all_proxy_servers`` sentinel → the full registry, legacy access groups,
tool-perm-referenced servers, toolset-referenced servers) unioned with its unified
@ -2620,13 +2680,17 @@ class MCPRequestHandler:
if SpecialMCPServerName.all_proxy_servers.value in (object_permissions.mcp_servers or []):
return set(global_mcp_server_manager.get_registry().keys())
legacy_access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
object_permissions.mcp_access_groups or []
object_permissions.mcp_access_groups or [],
requires_fresh_policy=requires_fresh_policy,
)
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(
object_permissions, requires_fresh_policy=requires_fresh_policy
)
return (
set(global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or []))
| set(legacy_access_group_servers)
| set(global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys())
| (await MCPRequestHandler._toolset_tool_permissions(object_permissions)).keys()
| toolset_grants.keys()
| set(team_access_group_servers)
)
@ -2667,6 +2731,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if team_obj is None:
return []
@ -2680,12 +2745,19 @@ class MCPRequestHandler:
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
servers: Final = await MCPRequestHandler._team_granted_servers(team_obj, team_access_group_servers)
servers: Final = await MCPRequestHandler._team_granted_servers(
team_obj,
team_access_group_servers,
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
return list(servers)
except Exception as e:
if isinstance(e, UnloadableEntitlementError):
if isinstance(e, UnloadableEntitlementError) or (
user_api_key_auth is not None and user_api_key_auth.requires_fresh_policy
):
raise
verbose_logger.warning("Failed to get allowed MCP servers for team: %s", e)
return []
@ -2716,6 +2788,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
except Exception as e: # noqa: BLE001 # a named entitlement we cannot read denies, whatever the read failed with
raise unloadable from e
@ -2811,7 +2884,8 @@ class MCPRequestHandler:
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or [])
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
object_permissions.mcp_access_groups or []
object_permissions.mcp_access_groups or [],
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
tool_perm_servers: Final = list(
@ -2820,7 +2894,10 @@ class MCPRequestHandler:
# servers referenced by the org's toolset grants are part of the org ceiling,
# exactly as servers referenced by its inline tool permissions are
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(object_permissions)
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(
object_permissions,
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
all_servers: Final = tuple(
{*direct_mcp_servers, *access_group_servers, *tool_perm_servers, *toolset_grants}
@ -2912,7 +2989,8 @@ class MCPRequestHandler:
# Get MCP servers from access groups
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
object_permission.mcp_access_groups or []
object_permission.mcp_access_groups or [],
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
# servers referenced in tool permissions should also be accessible
@ -2961,7 +3039,9 @@ class MCPRequestHandler:
return None
user_id: Final = user_api_key_auth.user_id
object_permission_id: Final = await MCPRequestHandler._user_object_permission_id(user_id, prisma_client)
object_permission_id: Final = await MCPRequestHandler._user_object_permission_id(
user_id, prisma_client, check_db_only=user_api_key_auth.requires_fresh_policy
)
if object_permission_id is None:
return None
@ -2971,6 +3051,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if object_permission is None:
raise ValueError(
@ -2979,7 +3060,9 @@ class MCPRequestHandler:
return object_permission
@staticmethod
async def _user_object_permission_id(user_id: str, prisma_client: "PrismaClient") -> str | None:
async def _user_object_permission_id(
user_id: str, prisma_client: "PrismaClient", *, check_db_only: bool = False
) -> str | None:
"""The permission row this human's user row links to, or None when they link none.
Caches the link (with a sentinel for "links none") so a human without an entitlement costs no
@ -2988,16 +3071,23 @@ class MCPRequestHandler:
whether someone is entitled is the state that existed before this level, so it places no
ceiling. Only a link we DID resolve can make the caller deny.
"""
from litellm.proxy.auth.auth_checks import get_user_object
from litellm.proxy.proxy_server import user_api_key_cache
cache_key: Final = user_object_permission_id_cache_key(user_id)
try:
cached: Final[object] = await user_api_key_cache.async_get_cache(key=cache_key)
cached: Final[object] = None if check_db_only else await user_api_key_cache.async_get_cache(key=cache_key)
if cached == USER_NO_MCP_PERMISSION_SENTINEL:
return None
if isinstance(cached, str) and cached:
return cached
user_row: Final = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id})
user_row: Final = await get_user_object(
user_id=user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
check_db_only=check_db_only,
)
linked: Final[object] = getattr(user_row, "object_permission_id", None) if user_row is not None else None
object_permission_id: Final = linked if isinstance(linked, str) and linked else None
await user_api_key_cache.async_set_cache(
@ -3006,7 +3096,9 @@ class MCPRequestHandler:
ttl=get_management_object_ttl(user_api_key_cache),
)
return object_permission_id
except Exception as e: # noqa: BLE001 # unknown whether entitled at all: no ceiling, as before
except Exception as e: # noqa: BLE001 # Legacy callers retain their existing optional user-ceiling behavior
if check_db_only:
raise HTTPException(503, "User policy is unavailable") from e
verbose_logger.warning("MCP user entitlement: link for %r unresolved, no ceiling: %s", user_id, e)
return None
@ -3031,13 +3123,17 @@ class MCPRequestHandler:
return []
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or [])
fresh: Final = bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy)
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
object_permissions.mcp_access_groups or []
object_permissions.mcp_access_groups or [],
requires_fresh_policy=fresh,
)
tool_perm_servers: Final = list(
global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys()
)
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(object_permissions)
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(
object_permissions, requires_fresh_policy=fresh
)
return tuple({*direct_mcp_servers, *access_group_servers, *tool_perm_servers, *toolset_grants})
except Exception as e: # noqa: BLE001 # any resolution fault is an unresolved ceiling, never "no ceiling"
verbose_logger.warning("Failed to get allowed MCP servers for user: %s", e)
@ -3075,7 +3171,7 @@ class MCPRequestHandler:
return capped, True
@staticmethod
async def _apply_agent_caller_ceiling(
async def apply_agent_caller_ceiling(
allowed_mcp_servers: Sequence[str],
user_api_key_auth: UserAPIKeyAuth | None = None,
) -> tuple[tuple[str, ...], bool]:
@ -3119,9 +3215,13 @@ class MCPRequestHandler:
(any non-empty entitlement, or an unresolved one, disqualifies), exactly as
``operator_open_server_ids`` reads the same row. The one owner of this predicate: the
server-axis registry resolution in ``get_allowed_mcp_servers`` and the tools-axis open
channel in ``_resolve_admitted_subject_tools`` both consult it, so the two axes cannot
channel in ``resolve_admitted_subject_tools`` both consult it, so the two axes cannot
disagree."""
if user_api_key_auth is None or not user_api_key_has_admin_view(user_api_key_auth):
if (
user_api_key_auth is None
or user_api_key_auth.mcp_explicit_grants_only
or not user_api_key_has_admin_view(user_api_key_auth)
):
return False
object_permission: Final = user_api_key_auth.object_permission
credential_scoped: Final = (
@ -3167,7 +3267,11 @@ class MCPRequestHandler:
user_direct_tools: Final = global_mcp_server_manager.expand_tool_permissions(
object_permissions.mcp_tool_permissions
).get(server_id)
user_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(object_permissions, server_id)
user_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
object_permissions,
server_id,
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
user_tools: Final = MCPRequestHandler._union_tool_grants(user_direct_tools, user_toolset_tools)
if user_tools is None:
return allowed_tools
@ -3176,7 +3280,7 @@ class MCPRequestHandler:
return list(set(allowed_tools) & set(user_tools))
@staticmethod
async def _apply_agent_caller_tool_ceiling(
async def apply_agent_caller_tool_ceiling(
allowed_tools: Sequence[str] | None,
server_id: str,
user_api_key_auth: UserAPIKeyAuth | None = None,
@ -3184,7 +3288,7 @@ class MCPRequestHandler:
"""Narrow an agent key's tools on ``server_id`` to those the invoking user and team (echoed back
by the agent as ``x-litellm-user-id`` / ``x-litellm-team-id``) may call: the echoed team's tool
grants when it names any on this server, then the echoed user's own tool entitlement. The tools
axis twin of ``_apply_agent_caller_ceiling``, so the headers only ever narrow. Denies every tool
axis twin of ``apply_agent_caller_ceiling``, so the headers only ever narrow. Denies every tool
on the server when the caller's team cannot be loaded, since a caller we cannot resolve must not
read as unrestricted."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
@ -3196,7 +3300,9 @@ class MCPRequestHandler:
return allowed_tools
try:
team_obj_perm: Final = await MCPRequestHandler._get_team_object_permission(caller_auth)
team_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(team_obj_perm, server_id)
team_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
team_obj_perm, server_id, requires_fresh_policy=caller_auth.requires_fresh_policy
)
except Exception as e: # noqa: BLE001 # an unresolved caller team must deny, not widen
verbose_logger.warning(
"MCP agent caller team tool ceiling unresolvable, denying tools on %r: %s", server_id, e
@ -3241,7 +3347,11 @@ class MCPRequestHandler:
end_user_direct_tools: Final = global_mcp_server_manager.expand_tool_permissions(
object_permissions.mcp_tool_permissions
).get(server_id)
end_user_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(object_permissions, server_id)
end_user_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
object_permissions,
server_id,
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
end_user_tools: Final = MCPRequestHandler._union_tool_grants(end_user_direct_tools, end_user_toolset_tools)
if end_user_tools is None:
return allowed_tools
@ -3302,6 +3412,11 @@ class MCPRequestHandler:
if not user_api_key_auth or not user_api_key_auth.agent_id:
return None
managed: Final = managed_agent_policy(user_api_key_auth)
if managed is not None:
permission: Final = managed.object_permission
return LiteLLM_ObjectPermissionTable.model_validate(permission) if permission is not None else None
if prisma_client is None:
verbose_logger.debug("prisma_client is None")
return None
@ -3319,7 +3434,7 @@ class MCPRequestHandler:
)
@staticmethod
async def _get_allowed_mcp_servers_for_agent(
async def get_allowed_mcp_servers_for_agent(
user_api_key_auth: UserAPIKeyAuth | None = None,
agent_object_permission: LiteLLM_ObjectPermissionTable | None = None,
) -> list[str]:
@ -3358,12 +3473,16 @@ class MCPRequestHandler:
obj_perm.mcp_servers or []
)
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
obj_perm.mcp_access_groups or []
obj_perm.mcp_access_groups or [],
requires_fresh_policy=user_api_key_auth.requires_fresh_policy,
)
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(obj_perm)
return list({*expanded_direct_servers, *access_group_servers, *toolset_grants})
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(
obj_perm, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
)
inline_tools: Final = global_mcp_server_manager.expand_tool_permissions(obj_perm.mcp_tool_permissions)
return list({*expanded_direct_servers, *access_group_servers, *toolset_grants, *inline_tools})
except Exception as e:
if isinstance(e, UnloadableEntitlementError):
if managed_agent_policy(user_api_key_auth) is not None or isinstance(e, UnloadableEntitlementError):
raise
verbose_logger.warning("Failed to get allowed MCP servers for agent: %s", e)
return []
@ -3390,7 +3509,7 @@ class MCPRequestHandler:
return frozenset(global_mcp_server_manager.expand_permission_list(sorted(ceiling.mcp_server_ids)))
@staticmethod
async def _get_agent_tool_permissions_for_server(
async def get_agent_tool_permissions_for_server(
server_id: str,
user_api_key_auth: UserAPIKeyAuth | None = None,
agent_object_permission: LiteLLM_ObjectPermissionTable | None = None,
@ -3430,11 +3549,13 @@ class MCPRequestHandler:
if obj_perm.mcp_tool_permissions
else None
)
toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(obj_perm, server_id)
toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
obj_perm, server_id, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
)
agent_tools: Final = MCPRequestHandler._union_tool_grants(direct_tools, toolset_tools)
return list(agent_tools) if agent_tools else None
return list(agent_tools) if agent_tools is not None else None
except Exception as e:
if isinstance(e, UnloadableEntitlementError):
if managed_agent_policy(user_api_key_auth) is not None or isinstance(e, UnloadableEntitlementError):
raise
verbose_logger.warning("Failed to get agent tool permissions for server: %s", e)
return None
@ -3452,28 +3573,38 @@ class MCPRequestHandler:
return server_ids
@staticmethod
async def _get_db_server_ids_for_access_groups(prisma_client, access_groups: list[str]) -> set[str]:
async def _get_db_server_ids_for_access_groups(
prisma_client,
access_groups: list[str],
*,
use_writer: bool = False,
) -> set[str]:
"""
Helper to get server_ids from DB servers that match any of the given access groups.
"""
server_ids: Final[set[str]] = set()
if access_groups and prisma_client is not None:
try:
mcp_servers: Final = await MCPServerRepository(prisma_client).table.find_many(
mcp_servers: Final = await MCPServerRepository(prisma_client, use_writer=use_writer).table.find_many(
where={"mcp_access_groups": {"hasSome": access_groups}}
)
for server in mcp_servers:
server_ids.add(server.server_id)
except Exception as e:
if use_writer:
raise
verbose_logger.debug("Error getting MCP servers from access groups: %s", e)
return server_ids
@staticmethod
async def _get_mcp_servers_from_access_groups(
access_groups: list[str],
*,
requires_fresh_policy: bool = False,
) -> list[str]:
"""
Resolve MCP access groups to server IDs by querying BOTH the MCP server table (DB) AND config-loaded servers
Resolve MCP access groups to server IDs by querying BOTH the MCP server table (DB) AND config-loaded servers.
``requires_fresh_policy`` reads the writer and propagates a read fault instead of resolving to no servers.
"""
from litellm.proxy.proxy_server import prisma_client
@ -3489,11 +3620,15 @@ class MCPRequestHandler:
)
# Use the new helper for DB servers
db_server_ids = await MCPRequestHandler._get_db_server_ids_for_access_groups(prisma_client, access_groups)
db_server_ids = await MCPRequestHandler._get_db_server_ids_for_access_groups(
prisma_client, access_groups, use_writer=requires_fresh_policy
)
server_ids.update(db_server_ids)
return list(server_ids)
except Exception as e:
if requires_fresh_policy:
raise
verbose_logger.warning("Failed to get MCP servers from access groups: %s", e)
return []
@ -3548,6 +3683,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if key_object_permission is None:
return []
@ -3591,6 +3727,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if team_obj is None:
verbose_logger.debug("team_obj is None")

View file

@ -170,6 +170,11 @@ async def identity_from_subject_token(
return _refusal_for(denied, denied.message)
except Exception as denied: # noqa: BLE001 # auth_jwt raises a plain Exception on signature and claim failures
return _refusal_for(denied, denied)
if result.get("agent_id") is not None:
return SubjectTokenRefusal(
error="invalid_request",
description="Agent tokens require direct JWT authentication; this exchange supports users only",
)
user_id: Final = result["user_id"]
if user_id is None:
return SubjectTokenRefusal(error="invalid_request", description="subject_token names no user the gateway knows")

View file

@ -181,6 +181,7 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
is_per_server_oauth_discovery_eligible,
)
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_policy
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
from litellm.proxy.common_utils.user_api_key_cache import get_management_object_ttl
@ -3428,7 +3429,9 @@ class MCPServerManager:
``allow_all_server_ids`` / ``submitted_server_ids`` are injectable so the server union,
which precomputes both for its fallback path, does not compute them twice."""
if user_api_key_auth is not None and user_api_key_auth.mcp_toolset_id is not None:
if user_api_key_auth is not None and (
user_api_key_auth.mcp_toolset_id is not None or user_api_key_auth.mcp_explicit_grants_only
):
return set()
if allow_all_server_ids is None:
allow_all_server_ids = self.get_allow_all_keys_server_ids()
@ -3477,9 +3480,14 @@ class MCPServerManager:
2. If admin and no object_permission, return all servers
3. Otherwise, use standard permission checks
"""
if managed_agent_policy(user_api_key_auth) is not None:
managed: Final = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth)
return managed if access is None else [server for server in managed if server in access.server_ids]
from litellm.proxy.proxy_server import general_settings as proxy_general_settings
resolved_general_settings: Final = proxy_general_settings if general_settings is None else general_settings
explicit_grants_only: Final = bool(user_api_key_auth and user_api_key_auth.mcp_explicit_grants_only)
allow_all_server_ids: Final = self.get_allow_all_keys_server_ids()
# A keyless admitted subject is resolved per grant source, and channel decisions that are
@ -3511,7 +3519,7 @@ class MCPServerManager:
# only keys without their own mcp_servers list get submitted servers unioned in.
submitted_server_ids: Final = (
[]
if has_explicit_object_permission
if has_explicit_object_permission or explicit_grants_only
else await self._get_active_submitted_mcp_server_ids_for_user(user_api_key_auth)
)
@ -3580,12 +3588,14 @@ class MCPServerManager:
return [
server_id
for server_id in dict.fromkeys(allow_all_server_ids + submitted_server_ids)
if scope is None or server_id == scope
if not explicit_grants_only and (scope is None or server_id == scope)
]
async def resolve_toolset_tool_permissions(
self,
toolset_ids: list[str],
*,
requires_fresh_policy: bool = False,
) -> dict[str, list[str]]:
"""
Resolve a list of toolset IDs into a mcp_tool_permissions dict.
@ -3595,6 +3605,10 @@ class MCPServerManager:
Redis-backed ``DualCache`` in production) so that cache entries are
shared across workers and cold-cache DB hits are minimised.
``requires_fresh_policy`` bypasses the cache and reads the writer so a
revocation is honoured on the very next request; a read fault then
propagates instead of resolving to no grants.
A row names a tool on the server identified by ``server_id``, so the
stored name is the tool's own name and is used as written. It is never
reduced by the server's wire prefix: that prefix is added on the way out
@ -3609,12 +3623,16 @@ class MCPServerManager:
return {}
cache_key: Final = "toolset_perms:" + ",".join(sorted(toolset_ids))
cached: Final[dict[str, list[str]] | None] = await user_api_key_cache.async_get_cache(key=cache_key)
cached: Final[dict[str, list[str]] | None] = (
None if requires_fresh_policy else await user_api_key_cache.async_get_cache(key=cache_key)
)
if cached is not None:
return cached
try:
toolsets: Final = await list_mcp_toolsets(prisma_client, toolset_ids=toolset_ids)
toolsets: Final = await list_mcp_toolsets(
prisma_client, toolset_ids=toolset_ids, use_writer=requires_fresh_policy
)
tool_permissions: Final[dict[str, list[str]]] = {}
for toolset in toolsets:
for tool in toolset.tools:
@ -3628,6 +3646,8 @@ class MCPServerManager:
)
return tool_permissions
except Exception as e:
if requires_fresh_policy:
raise
verbose_logger.warning("Failed to resolve toolset permissions: %s", e)
return {}

View file

@ -5,7 +5,7 @@ from dataclasses import dataclass
from datetime import datetime
from traceback import walk_tb
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal
from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict
from uuid import uuid4
import anyio
@ -14,6 +14,7 @@ import httpx2
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
from pydantic import ValidationError
from starlette.datastructures import Headers
from typing_extensions import ReadOnly
from litellm._logging import verbose_logger
from litellm.constants import MCP_CLIENT_TIMEOUT, MCP_TOOL_LISTING_TIMEOUT
@ -63,7 +64,27 @@ if TYPE_CHECKING:
from litellm.proxy.utils import ProxyLogging
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
from litellm.types.mcp import MCPAuth
from litellm.types.utils import CallTypes
from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall
class _MCPModelMetadata(TypedDict):
model_group: ReadOnly[str]
def _stamp_mcp_tool_metadata(logging_obj: "LiteLLMLoggingObj | None", server_id: str, tool_name: str) -> None:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
if logging_obj is None:
return
server: Final = global_mcp_server_manager.get_mcp_server_by_id(
server_id
) or global_mcp_server_manager.get_mcp_server_by_name(server_id)
metadata: Final[StandardLoggingMCPToolCall] = {
"name": tool_name,
"mcp_server_name": server.name if server is not None else server_id,
}
logging_obj.model_call_details["mcp_tool_call_metadata"] = metadata
MCP_AVAILABLE: bool = True
try:
@ -1193,6 +1214,12 @@ if MCP_AVAILABLE:
},
)
data["model"] = f"MCP: {tool_name}"
model_metadata: Final[_MCPModelMetadata] = {
**(data.get("metadata") or MappingProxyType({})),
"model_group": f"MCP: {tool_name}",
}
data["metadata"] = model_metadata
proxy_base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data)
_request_start_time: Final = datetime.now() # noqa: DTZ005 # naive to match the tool start time below
try:
@ -1226,6 +1253,8 @@ if MCP_AVAILABLE:
if "metadata" in data and "user_api_key_auth" in data["metadata"]:
data["user_api_key_auth"] = data["metadata"]["user_api_key_auth"]
_stamp_mcp_tool_metadata(logging_obj, server_id, tool_name)
# Resolve allowed MCP servers with IP filtering
(
allowed_mcp_servers,

View file

@ -65,9 +65,9 @@ class MCPToolsetTable(Protocol):
async def delete(self, where: Mapping[str, object]) -> MCPToolsetRow: ...
def _toolset_table(prisma_client: PrismaClient) -> MCPToolsetTable:
def _toolset_table(prisma_client: PrismaClient, *, use_writer: bool = False) -> MCPToolsetTable:
"""The toolset table actions of the prisma client."""
return MCPToolsetRepository(prisma_client).table
return MCPToolsetRepository(prisma_client, use_writer=use_writer).table
def _toolset_from_row(row: MCPToolsetRow) -> MCPToolset:
@ -107,12 +107,16 @@ async def get_mcp_toolset(
async def list_mcp_toolsets(
prisma_client: PrismaClient,
toolset_ids: Sequence[str] | None = None,
*,
use_writer: bool = False,
) -> Sequence[MCPToolset]:
try:
where: Final[Mapping[str, object]] = {} if toolset_ids is None else {"toolset_id": {"in": toolset_ids}}
rows: Final = await _toolset_table(prisma_client).find_many(where=where)
rows: Final = await _toolset_table(prisma_client, use_writer=use_writer).find_many(where=where)
return [_toolset_from_row(r) for r in rows]
except Exception as e:
if use_writer:
raise
verbose_proxy_logger.warning("litellm.proxy._experimental.mcp_server.toolset_db::list_mcp_toolsets - %s", e)
return []

View file

@ -58,6 +58,7 @@ async def resolve_ui_session_team_ids(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
check_db_only=user_api_key_auth.requires_fresh_policy,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
@ -92,7 +93,9 @@ async def admitted_user_context(user_api_key_auth: UserAPIKeyAuth) -> UserAPIKey
)
try:
admitted: Final = await MCPRequestHandler.reload_admitted_user(user_id)
admitted: Final = await MCPRequestHandler.reload_admitted_user(
user_id, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
)
except HTTPException as e:
verbose_logger.warning("MCP dashboard session: admitted-subject reload failed for %s: %s", user_id, e.detail)
return None

View file

@ -133,6 +133,11 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = (
module_path="litellm.proxy.management_endpoints.model_insights_endpoints",
path_prefixes=("/model-insights",),
),
LazyFeature(
name="roi_calculator",
module_path="litellm.proxy.management_endpoints.roi_calculator_endpoints",
path_prefixes=("/roi-calculator",),
),
LazyFeature(
name="search_tools",
module_path="litellm.proxy.search_endpoints.search_tool_management",

File diff suppressed because it is too large Load diff

View file

@ -1,7 +1,7 @@
import enum
import json
import os
from collections.abc import Callable, Mapping
from collections.abc import Callable, Mapping, Sequence
from datetime import datetime
from types import MappingProxyType
from typing import TYPE_CHECKING, Annotated, Any, Final, Literal, NamedTuple, TypeAlias
@ -15,6 +15,7 @@ from pydantic import (
Json,
JsonValue,
PositiveInt,
PrivateAttr,
field_validator,
model_validator,
)
@ -519,6 +520,10 @@ class LiteLLMRoutes(enum.Enum):
"/v1/rag/ingest",
"/rag/query",
"/v1/rag/query",
# agent tracing: OTLP ingest + reads (scoped to the caller's team in the handler)
"/v1/traces",
"/v1/traces/{trace_id}",
"/v1/traces/{trace_id}/spans/{span_id}",
]
anthropic_routes = [
@ -2240,6 +2245,12 @@ class ResetTeamBudgetRequest(LiteLLMPydanticObjectBase):
class DeleteTeamRequest(LiteLLMPydanticObjectBase):
team_ids: list[str] # required
@field_validator("team_ids")
@classmethod
def distinct_team_ids(cls, team_ids: Sequence[str]) -> list[str]:
"""One delete per team: a repeated id would otherwise write its tombstone and audit row twice."""
return list(dict.fromkeys(team_ids))
class BlockTeamRequest(LiteLLMPydanticObjectBase):
team_id: str # required
@ -3319,6 +3330,7 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
# single-owner so its meaning stays trustworthy.
mcp_session_resource_server_id: str | None = Field(default=None, exclude=True)
mcp_toolset_id: str | None = Field(default=None, exclude=True)
authenticated_by_custom_auth: bool = Field(default=False, exclude=True)
via_virtual_key: bool = Field(
default=False,
exclude=True,
@ -3334,6 +3346,7 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
invoked_agent_policy: AgentResponse | None = Field(default=None, exclude=True)
agent_invocation_cost: float | None = Field(default=None, exclude=True)
billing_agent_policy: AgentResponse | None = Field(default=None, exclude=True)
_managed_delegation_verified: bool = PrivateAttr(default=False)
managed_agent_policy: AgentResponse | None = Field(default=None, exclude=True)
managed_agent_context: ManagedAgentContext | None = Field(default=None, exclude=True)
agent_caller: AgentCaller | None = Field(
@ -3379,6 +3392,7 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
values.pop("mcp_session_resource_server_id", None)
values.pop("mcp_toolset_id", None)
values.pop("via_virtual_key", None)
values.pop("authenticated_by_custom_auth", None)
values.pop("agent_caller", None)
values.pop("managed_agent_context", None)
values.pop("managed_agent_policy", None)

View file

@ -597,7 +597,6 @@ async def get_agent_card(
if agent is None:
raise HTTPException(status_code=404, detail=f"Agent '{agent_id}' not found")
# Check agent permission (skip for admin users)
is_allowed: Final = await AgentRequestHandler.is_agent_allowed(
agent_id=agent.agent_id,
user_api_key_auth=user_api_key_dict,
@ -723,6 +722,8 @@ async def invoke_agent_a2a(
detail=f"Agent '{agent_id}' is not allowed for your key/team. Contact proxy admin for access.",
)
user_api_key_dict.invoked_agent_id = agent.agent_id
_enforce_inbound_trace_id(agent, request)
# Get backend URL and agent name
@ -760,6 +761,10 @@ async def invoke_agent_a2a(
if "metadata" not in body:
body["metadata"] = {}
body["metadata"]["agent_id"] = agent.agent_id
body["metadata"]["model_group"] = f"a2a_agent/{agent_name}"
body["metadata"]["model_info"] = { # mutable-ok: request hooks mutate metadata before JSON logging
"id": agent.agent_id
}
body["agent_id"] = agent.agent_id
body.update(
@ -863,6 +868,7 @@ async def invoke_agent_a2a(
# results written by the unified_guardrail hook are captured.
logging_obj._defer_async_logging = True
response = await asend_message(
model=f"a2a_agent/{agent_name}",
request=a2a_request,
api_base=agent_url,
litellm_params=litellm_params,

View file

@ -57,7 +57,7 @@ async def route_a2a_agent_request(
user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
)
if not is_admin:
if not is_admin or agent.identity_managed:
is_allowed: Final = await AgentRequestHandler.is_agent_allowed(
agent_id=agent.agent_id,
user_api_key_auth=user_api_key_dict,

View file

@ -1,13 +1,16 @@
import asyncio
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from typing import Final, TypeAlias
from typing import TYPE_CHECKING, Final, TypeAlias
from fastapi import HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import LiteLLM_AccessGroupTable
if TYPE_CHECKING:
from litellm.types.agents import AgentResponse
AccessGroupIds: TypeAlias = tuple[str, ...]
AccessGroupIdsLoader: TypeAlias = Callable[[str], Awaitable[AccessGroupIds]] # mutable-ok: Callable params
LoadedAccessGroup: TypeAlias = LiteLLM_AccessGroupTable | None
@ -34,7 +37,7 @@ async def _registry_access_group_ids(agent_id: str) -> AccessGroupIds:
return tuple(agent.access_group_ids or ()) if agent is not None else ()
async def _load_access_group(access_group_id: str) -> LoadedAccessGroup:
async def _load_access_group(access_group_id: str, *, check_db_only: bool = False) -> LoadedAccessGroup:
from litellm.proxy.auth.auth_checks import get_access_object
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
@ -47,8 +50,11 @@ async def _load_access_group(access_group_id: str) -> LoadedAccessGroup:
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
check_db_only=check_db_only,
)
except HTTPException as e:
if check_db_only:
raise
verbose_proxy_logger.warning(
"Agent access group %s could not be loaded, treating it as empty: %s", access_group_id, e.detail
)
@ -59,13 +65,20 @@ async def resolve_agent_access_group_ceiling(
agent_id: str,
load_access_group_ids: AccessGroupIdsLoader = _registry_access_group_ids,
load_access_group: AccessGroupLoader = _load_access_group,
*,
check_db_only: bool = False,
) -> AgentAccessGroupCeiling | None:
"""``None`` when the agent has no access groups attached, so nothing is capped."""
access_group_ids: Final = await load_access_group_ids(agent_id)
if not access_group_ids:
return None
loaded: Final = await asyncio.gather(*(load_access_group(group_id) for group_id in access_group_ids))
loaded: Final = await asyncio.gather(
*(
_load_access_group(group_id, check_db_only=True) if check_db_only else load_access_group(group_id)
for group_id in access_group_ids
)
)
groups: Final = tuple(group for group in loaded if group is not None)
return AgentAccessGroupCeiling(
access_group_ids=access_group_ids,
@ -73,3 +86,16 @@ async def resolve_agent_access_group_ceiling(
mcp_server_ids=frozenset(server_id for group in groups for server_id in group.access_mcp_server_ids),
agent_ids=frozenset(target_id for group in groups for target_id in group.access_agent_ids),
)
async def resolve_managed_agent_ceilings(agent: "AgentResponse") -> tuple[AgentAccessGroupCeiling, ...]:
async def authoritative_group(group_id: str) -> LoadedAccessGroup:
return await _load_access_group(group_id, check_db_only=True)
async def manual_ids(_agent_id: str) -> AccessGroupIds:
return tuple(agent.access_group_ids or ())
manual: Final = await resolve_agent_access_group_ceiling(
agent.agent_id, load_access_group_ids=manual_ids, load_access_group=authoritative_group
)
return (manual,) if manual is not None else ()

View file

@ -8,6 +8,7 @@ can only narrow access and need no trust.
"""
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final
from litellm._logging import verbose_proxy_logger
@ -45,7 +46,7 @@ def agent_caller_auth(user_api_key_auth: UserAPIKeyAuth) -> UserAPIKeyAuth | Non
user_id=caller.user_id,
team_id=caller.team_id,
parent_otel_span=user_api_key_auth.parent_otel_span,
)
).model_copy(update=MappingProxyType({"requires_fresh_policy": user_api_key_auth.requires_fresh_policy}))
async def load_agent_caller_team(user_api_key_auth: UserAPIKeyAuth) -> LiteLLM_TeamTable | None:

View file

@ -8,8 +8,11 @@ Follows the same pattern as MCP permission handling.
import asyncio
from collections.abc import Awaitable, Callable, Sequence
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final, TypeAlias
from fastapi import HTTPException
from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.ui_session_utils import build_effective_auth_contexts
from litellm.proxy._types import (
@ -24,6 +27,7 @@ from litellm.proxy.agent_endpoints.auth.agent_access_groups import (
resolve_agent_access_group_ceiling,
)
from litellm.proxy.agent_endpoints.auth.agent_caller import agent_caller_auth
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_policy
from litellm.repositories.table_repositories import AgentsRepository
from litellm.types.agents import AgentResponse
@ -83,13 +87,23 @@ class AgentRequestHandler:
async def resolve_agent_access(
user_api_key_auth: UserAPIKeyAuth | None = None,
resolve_ceiling: CeilingResolver = resolve_agent_access_group_ceiling,
*,
strict: bool = False,
) -> AgentAccess:
"""Agents the key may reach: key and team grants, intersected with the agent's access group ceiling
and, for an agent key acting on behalf of an invoking user, with that user's team grants."""
key_team_access: Final = await AgentRequestHandler._resolve_key_team_agent_access(user_api_key_auth)
caller_access: Final = await AgentRequestHandler._agent_caller_access(user_api_key_auth)
if managed_agent_policy(user_api_key_auth) is not None:
return await _managed_actor_agent_access(user_api_key_auth)
key_team_access: Final = await AgentRequestHandler.resolve_key_team_agent_access(
user_api_key_auth, strict=strict
)
if strict and isinstance(key_team_access, UnrestrictedAgentAccess):
return RestrictedAgentAccess(frozenset())
caller_access: Final = await AgentRequestHandler.agent_caller_access(user_api_key_auth, strict=strict)
own_access: Final = _intersect_agent_access(key_team_access, caller_access)
agent_ceiling: Final = await AgentRequestHandler._agent_access_group_ceiling(user_api_key_auth, resolve_ceiling)
agent_ceiling: Final = await AgentRequestHandler._agent_access_group_ceiling(
user_api_key_auth, resolve_ceiling, strict=strict
)
if agent_ceiling is None:
return own_access
if isinstance(own_access, UnrestrictedAgentAccess):
@ -97,20 +111,26 @@ class AgentRequestHandler:
return RestrictedAgentAccess(own_access.agent_ids & agent_ceiling)
@staticmethod
async def _agent_caller_access(user_api_key_auth: UserAPIKeyAuth | None) -> AgentAccess:
async def agent_caller_access(user_api_key_auth: UserAPIKeyAuth | None, *, strict: bool = False) -> AgentAccess:
caller_auth: Final = agent_caller_auth(user_api_key_auth) if user_api_key_auth else None
if caller_auth is None:
return UnrestrictedAgentAccess()
return await AgentRequestHandler._get_allowed_agents_for_team(caller_auth)
return await AgentRequestHandler._get_allowed_agents_for_team(caller_auth, strict=strict)
@staticmethod
async def _resolve_key_team_agent_access(
async def resolve_key_team_agent_access(
user_api_key_auth: UserAPIKeyAuth | None,
*,
strict: bool = False,
) -> AgentAccess:
try:
key_access: Final = await AgentRequestHandler._get_allowed_agents_for_key(user_api_key_auth)
team_access: Final = await AgentRequestHandler._get_allowed_agents_for_team(user_api_key_auth)
key_access: Final = await AgentRequestHandler.get_allowed_agents_for_key(user_api_key_auth, strict=strict)
team_access: Final = await AgentRequestHandler._get_allowed_agents_for_team(
user_api_key_auth, strict=strict
)
except Exception as e:
if strict:
raise HTTPException(503, "Agent invocation policy is unavailable") from e
verbose_logger.warning("Failed to get allowed agents: %s", e)
return UnrestrictedAgentAccess()
return _intersect_agent_access(key_access, team_access)
@ -119,10 +139,16 @@ class AgentRequestHandler:
async def _agent_access_group_ceiling(
user_api_key_auth: UserAPIKeyAuth | None,
resolve_ceiling: CeilingResolver,
*,
strict: bool = False,
) -> frozenset[str] | None:
if user_api_key_auth is None or not user_api_key_auth.agent_id:
return None
ceiling: Final = await resolve_ceiling(user_api_key_auth.agent_id)
ceiling: Final = (
await resolve_agent_access_group_ceiling(user_api_key_auth.agent_id, check_db_only=True)
if strict
else await resolve_ceiling(user_api_key_auth.agent_id)
)
if ceiling is None:
return None
return _to_stable_ids(ceiling.agent_ids)
@ -144,6 +170,49 @@ class AgentRequestHandler:
bool: True if agent is allowed, False otherwise
"""
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
from litellm.proxy.proxy_server import prisma_client
from litellm.types.proxy.agent_identity import AgentIdentityFailure
registered: Final = global_agent_registry.get_agent_by_id(agent_id)
registry_managed: Final = isinstance(registered, AgentResponse) and registered.identity_managed
if registry_managed or prisma_client is not None:
target: Final = await AgentIdentityStore.from_client(prisma_client).agent(agent_id)
if isinstance(target, AgentIdentityFailure):
raise_identity_failure(target)
elif target is None and registry_managed:
return False
elif isinstance(target, AgentResponse) and target.identity_managed:
if (
not target.enabled
or target.identity is None
or not target.identity.active
or user_api_key_auth is None
):
return False
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
key_hash: Final = user_api_key_auth.api_key or user_api_key_auth.token
authority: Final = (
await MCPRequestHandler._reload_admitted_key(key_hash, check_db_only=True) # pyright: ignore[reportPrivateUsage] # the authoritative key reload has no public seam
if key_hash
and managed_agent_policy(user_api_key_auth) is None
and not user_api_key_auth.is_session_token
and not user_api_key_auth.authenticated_by_custom_auth
else user_api_key_auth
)
fresh_auth: Final = authority.model_copy(
update=MappingProxyType(
{"requires_fresh_policy": True, "agent_caller": user_api_key_auth.agent_caller}
)
)
explicit: Final = await _granted_agent_ids(
fresh_auth,
_strict_agent_access,
build_effective_auth_contexts,
)
return target.agent_id in explicit
match await AgentRequestHandler.resolve_agent_access(user_api_key_auth, resolve_ceiling):
case UnrestrictedAgentAccess():
@ -202,8 +271,10 @@ class AgentRequestHandler:
return team_obj.object_permission
@staticmethod
async def _get_allowed_agents_for_key(
async def get_allowed_agents_for_key(
user_api_key_auth: UserAPIKeyAuth | None = None,
*,
strict: bool = False,
) -> AgentAccess:
"""
Get allowed agents for a key.
@ -237,24 +308,36 @@ class AgentRequestHandler:
return UnrestrictedAgentAccess()
access_group_agents: Final = (
tuple(await AgentRequestHandler._get_agents_from_access_groups(list(declared_access_groups)))
tuple(
await AgentRequestHandler._get_agents_from_access_groups(
declared_access_groups, check_db_only=strict
)
)
if declared_access_groups
else ()
)
unified_agents: Final = (
tuple(await AgentRequestHandler._get_unified_access_group_agents(list(key_access_group_ids)))
tuple(
await AgentRequestHandler._get_unified_access_group_agents(
key_access_group_ids, check_db_only=strict
)
)
if key_access_group_ids
else ()
)
return RestrictedAgentAccess(frozenset(direct_agents + access_group_agents + unified_agents))
except Exception as e:
if strict:
raise HTTPException(503, "Agent invocation policy is unavailable") from e
verbose_logger.warning("Failed to get allowed agents for key: %s", e)
return UnrestrictedAgentAccess()
@staticmethod
async def _get_allowed_agents_for_team(
user_api_key_auth: UserAPIKeyAuth | None = None,
*,
strict: bool = False,
) -> AgentAccess:
"""
Get allowed agents for a team.
@ -263,7 +346,7 @@ class AgentRequestHandler:
2. Also includes agents from team's access_group_ids (unified access groups)
Fetches the team object once and reuses it for both permission sources.
Declared-but-empty grants stay restricted; see `_get_allowed_agents_for_key`.
Declared-but-empty grants stay restricted; see `get_allowed_agents_for_key`.
"""
if user_api_key_auth is None:
return UnrestrictedAgentAccess()
@ -280,7 +363,7 @@ class AgentRequestHandler:
)
if not prisma_client:
return UnrestrictedAgentAccess()
return RestrictedAgentAccess(frozenset()) if strict else UnrestrictedAgentAccess()
# Fetch the team object once for both permission sources
team_obj: Final = await get_team_object(
@ -289,10 +372,11 @@ class AgentRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=strict,
)
if team_obj is None:
return UnrestrictedAgentAccess()
return RestrictedAgentAccess(frozenset()) if strict else UnrestrictedAgentAccess()
# 1. Get agents from object_permission (native permissions)
object_permissions: Final = team_obj.object_permission
@ -307,18 +391,28 @@ class AgentRequestHandler:
return UnrestrictedAgentAccess()
access_group_agents: Final = (
tuple(await AgentRequestHandler._get_agents_from_access_groups(list(declared_access_groups)))
tuple(
await AgentRequestHandler._get_agents_from_access_groups(
declared_access_groups, check_db_only=strict
)
)
if declared_access_groups
else ()
)
unified_agents: Final = (
tuple(await AgentRequestHandler._get_unified_access_group_agents(list(team_access_group_ids)))
tuple(
await AgentRequestHandler._get_unified_access_group_agents(
team_access_group_ids, check_db_only=strict
)
)
if team_access_group_ids
else ()
)
return RestrictedAgentAccess(frozenset(direct_agents + access_group_agents + unified_agents))
except Exception as e:
if strict:
raise HTTPException(503, "Agent invocation policy is unavailable") from e
# litellm-dashboard is the default UI team and will never have agents;
# skip noisy warnings for it.
if user_api_key_auth.team_id != UI_TEAM_ID:
@ -326,7 +420,9 @@ class AgentRequestHandler:
return UnrestrictedAgentAccess()
@staticmethod
def _get_config_agent_ids_for_access_groups(config_agents: list, access_groups: list[str]) -> set[str]:
def _get_config_agent_ids_for_access_groups(
config_agents: Sequence[AgentResponse], access_groups: Sequence[str]
) -> set[str]:
"""
Helper to get agent_ids from config-loaded agents that match any of the given access groups.
"""
@ -339,7 +435,9 @@ class AgentRequestHandler:
return server_ids
@staticmethod
async def _get_db_agent_ids_for_access_groups(prisma_client, access_groups: list[str]) -> set[str]:
async def _get_db_agent_ids_for_access_groups(
prisma_client, access_groups: Sequence[str], *, check_db_only: bool = False
) -> set[str]:
"""
Helper to get agent_ids from DB agents that match any of the given access groups.
@ -349,23 +447,27 @@ class AgentRequestHandler:
if not access_groups or prisma_client is None:
return set()
agents: Final = await AgentsRepository(prisma_client).table.find_many(
agents: Final = await AgentsRepository(prisma_client, use_writer=check_db_only).table.find_many(
where={"agent_access_groups": {"hasSome": access_groups}}
)
return {agent.agent_id for agent in agents}
@staticmethod
async def _get_unified_access_group_agents(access_group_ids: list[str]) -> list[str]:
async def _get_unified_access_group_agents(
access_group_ids: Sequence[str], *, check_db_only: bool = False
) -> list[str]:
"""
Resolve unified access group ids to agent IDs.
"""
from litellm.proxy.auth.auth_checks import _get_agent_ids_from_access_groups
return await _get_agent_ids_from_access_groups(access_group_ids=access_group_ids)
return await _get_agent_ids_from_access_groups(access_group_ids=access_group_ids, check_db_only=check_db_only)
@staticmethod
async def _get_agents_from_access_groups(
access_groups: list[str],
access_groups: Sequence[str],
*,
check_db_only: bool = False,
) -> list[str]:
"""
Resolve agent access groups to agent IDs by querying BOTH the agent table (DB) AND config-loaded agents.
@ -373,14 +475,13 @@ class AgentRequestHandler:
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
from litellm.proxy.proxy_server import prisma_client
# Use the helper for config-loaded agents
config_agent_ids: Final = AgentRequestHandler._get_config_agent_ids_for_access_groups(
global_agent_registry.agent_list, access_groups
)
# Use the helper for DB agents
db_agent_ids: Final = await AgentRequestHandler._get_db_agent_ids_for_access_groups(
prisma_client, access_groups
prisma_client, access_groups, check_db_only=check_db_only
)
return list(config_agent_ids | db_agent_ids)
@ -531,4 +632,90 @@ async def accessible_agents(
AgentRequestHandler.resolve_agent_access if resolve_access is None else resolve_access,
effective_contexts,
)
return tuple(agent for agent in agents if agent.agent_id in allowed_agent_ids)
allowed: Final = await asyncio.gather(
*(
AgentRequestHandler.is_agent_allowed(agent.agent_id, user_api_key_auth)
for agent in agents
if agent.identity_managed
)
)
managed_ids: Final = frozenset(
agent.agent_id
for agent, permitted in zip((agent for agent in agents if agent.identity_managed), allowed)
if permitted
)
return tuple(
agent
for agent in agents
if (agent.agent_id in managed_ids if agent.identity_managed else agent.agent_id in allowed_agent_ids)
)
async def _strict_agent_access(auth: UserAPIKeyAuth) -> AgentAccess:
return await AgentRequestHandler.resolve_agent_access(auth, strict=True)
async def _managed_actor_agent_access(auth: UserAPIKeyAuth) -> AgentAccess:
agent: Final = managed_agent_policy(auth)
if agent is None or not agent.object_permission:
return RestrictedAgentAccess(frozenset())
permission: Final = LiteLLM_ObjectPermissionTable.model_validate(agent.object_permission or MappingProxyType({}))
own_auth: Final = UserAPIKeyAuth(object_permission=permission)
own: Final = _granted_ids(await AgentRequestHandler.get_allowed_agents_for_key(own_auth, strict=True))
from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_managed_agent_ceilings
ceilings: Final = await resolve_managed_agent_ceilings(agent)
grouped: Final = frozenset(target for target in own if all(target in ceiling.agent_ids for ceiling in ceilings))
caller: Final = await AgentRequestHandler.agent_caller_access(auth, strict=True)
capped: Final = grouped if isinstance(caller, UnrestrictedAgentAccess) else grouped & caller.agent_ids
context: Final = auth.managed_agent_context
if context is None or context.mode == "autonomous":
return RestrictedAgentAccess(capped)
if context.user_id is None:
return RestrictedAgentAccess(frozenset())
human_ids: Final = await verified_human_agent_grants(context.user_id, auth.team_id)
return RestrictedAgentAccess(capped.intersection(human_ids))
async def _verified_human_agent_sources(
user_id: str | None, *, allowed_team_ids: frozenset[str] | None = None
) -> tuple[tuple[str | None, frozenset[str]], ...]:
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
if user_id is None:
return ()
human: Final = await MCPRequestHandler.reload_admitted_user(user_id, requires_fresh_policy=True)
sources: Final = await MCPRequestHandler.admitted_subject_sources(human, allowed_team_ids=allowed_team_ids)
access: Final = await asyncio.gather(*(_strict_agent_access(source) for source in sources))
return tuple((source.team_id, _granted_ids(grant)) for source, grant in zip(sources, access, strict=True))
async def verified_human_agent_grants(user_id: str | None, team_id: str | None = None) -> frozenset[str]:
sources: Final = await _verified_human_agent_sources(
user_id, allowed_team_ids=frozenset((team_id,)) if team_id else frozenset()
)
return frozenset().union(*(grants for source, grants in sources if source is None or source == team_id))
async def resolve_delegated_agent_team(
user_id: str | None,
agent_id: str,
team_id: str | None,
*,
explicit_team: bool,
allowed_team_ids: frozenset[str] | None = None,
) -> str | None:
sources: Final = await _verified_human_agent_sources(user_id)
if any(source is None and agent_id in grants for source, grants in sources):
return team_id
granting_teams: Final = frozenset(
source
for source, grants in sources
if source is not None and agent_id in grants and (allowed_team_ids is None or source in allowed_team_ids)
)
if team_id in granting_teams:
return team_id
if not explicit_team and granting_teams:
return min(granting_teams)
raise HTTPException(403, "Select a team that grants access to this agent using x-litellm-team-id")

View file

@ -0,0 +1,252 @@
from collections.abc import Mapping
from itertools import product
from types import MappingProxyType
from typing import Annotated, Final, Literal
from pydantic import Field, TypeAdapter, ValidationError
from litellm.proxy._types import LiteLLMRoutes, UserAPIKeyAuth
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
from litellm.types.agents import AgentResponse
from litellm.types.proxy.agent_identity import AgentIdentityFailure, ManagedAgentContext
_MANAGED_REALTIME_ROUTES: Final = frozenset(("/realtime", "/v1/realtime", "/openai/v1/realtime"))
_MANAGED_MODEL_ROUTES: Final = frozenset(
f"{prefix}/{operation}"
for prefix, operation in product(
("", "/v1"),
(
"chat/completions",
"completions",
"embeddings",
"responses",
"messages",
"messages/count_tokens",
"images/generations",
"images/edits",
"audio/transcriptions",
"audio/speech",
"moderations",
"rerank",
"ocr",
),
)
) | frozenset(
(
"/openai/v1/responses",
"/v2/rerank",
"/claude_code_gateway/v1/messages",
"/claude_code_gateway/v1/messages/count_tokens",
"/cursor/chat/completions",
)
)
_MANAGED_MODEL_PATHS: Final = (
"/engines/{model:path}/chat/completions",
"/engines/{model:path}/completions",
"/engines/{model:path}/embeddings",
"/openai/deployments/{model:path}/chat/completions",
"/openai/deployments/{model:path}/completions",
"/openai/deployments/{model:path}/embeddings",
"/openai/deployments/{model:path}/images/generations",
"/openai/deployments/{model:path}/images/edits",
"/v1beta/models/{model_name:path}:countTokens",
"/v1beta/models/{model_name:path}:generateContent",
"/v1beta/models/{model_name:path}:streamGenerateContent",
"/models/{model_name:path}:countTokens",
"/models/{model_name:path}:generateContent",
"/models/{model_name:path}:streamGenerateContent",
)
_MANAGED_MCP_ROUTES: Final = tuple(
route for route in LiteLLMRoutes.mcp_inference_routes.value if route not in ("/token", "/introspect")
)
_MODEL_ROUTE_KINDS: Final[
Mapping[str, Literal["image_generation", "image_edit", "moderation", "speech", "body", "path"]]
] = MappingProxyType(
{
"/images/generations": "image_generation",
"/images/edits": "image_edit",
"/moderations": "moderation",
"/audio/transcriptions": "moderation",
"/audio/speech": "speech",
"/rerank": "body",
"/messages/count_tokens": "body",
":countTokens": "path",
}
)
def managed_agent_route_allowed(route: str, method: str | None) -> bool:
from litellm.proxy.auth.route_checks import RouteChecks
if route in ("/agents", "/v1/agents"):
return method in (None, "GET", "HEAD")
if route in _MANAGED_REALTIME_ROUTES:
return method in (None, "GET")
if route in _MANAGED_MODEL_ROUTES or RouteChecks.check_route_access(route, _MANAGED_MODEL_PATHS):
return method in (None, "POST")
return RouteChecks.check_route_access(route, _MANAGED_MCP_ROUTES) or RouteChecks.check_route_access(
route, LiteLLMRoutes.agent_inference_routes.value
)
def managed_inference_request(
route: str,
body: Mapping[str, object],
settings: Mapping[str, object],
cli_model: str | None,
path_model: object = None,
query_model: object = None,
) -> dict[str, object]:
from litellm.proxy.auth.route_checks import RouteChecks
if route in _MANAGED_REALTIME_ROUTES:
model: Final = query_model or body.get("model")
if not isinstance(model, str) or not model:
raise_identity_failure(
AgentIdentityFailure(message="Managed inference requires an explicit or configured model")
)
return {**body, "model": model} # mutable-ok: centralized auth hooks add request tags and budget metadata
if route not in _MANAGED_MODEL_ROUTES and not RouteChecks.check_route_access(route, _MANAGED_MODEL_PATHS):
return dict(body) # mutable-ok: centralized auth hooks add request tags and budget metadata
from litellm.proxy.common_utils.http_parsing_utils import resolve_inference_model
kind: Final = next((kind for suffix, kind in _MODEL_ROUTE_KINDS.items() if route.endswith(suffix)), "completion")
endpoint_model: Final = path_model or (
query_model if route.endswith(("/completions", "/embeddings", "/images/generations", "/images/edits")) else None
)
effective: Final = resolve_inference_model(body.get("model"), settings, cli_model, endpoint_model, kind=kind)
if not isinstance(effective, str) or not effective:
raise_identity_failure(
AgentIdentityFailure(message="Managed inference requires an explicit or configured model")
)
return {**body, "model": effective} # mutable-ok: centralized auth hooks add request tags and budget metadata
def managed_agent_policy(auth: "UserAPIKeyAuth | None") -> AgentResponse | None:
"""The admitted managed policy, or ``None`` when the subject was never admitted as a managed agent.
``admit_managed_actor`` only assigns ``managed_agent_policy`` after ``actor_admission_failure``
has verified the bound context, so an ``AgentResponse`` here means admission succeeded.
"""
policy: Final = auth.managed_agent_policy if auth is not None else None
return policy if isinstance(policy, AgentResponse) else None
async def admit_managed_actor(auth: UserAPIKeyAuth, store: AgentIdentityStore | None) -> None:
delegation_verified: Final = auth._managed_delegation_verified # pyright: ignore[reportPrivateUsage] # the one-shot delegation marker is a PrivateAttr by design
auth._managed_delegation_verified = False # pyright: ignore[reportPrivateUsage] # consumed here so a replayed token cannot reuse it
if auth.agent_id is None:
return
if store is None:
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
registered: Final = global_agent_registry.get_agent_by_id(auth.agent_id)
if auth.managed_agent_context is not None or (
registered is not None and (registered.identity_managed or registered.identity is not None)
):
raise_identity_failure(
AgentIdentityFailure(code="policy_unavailable", message="Managed agent policy requires a database")
)
return
agent: Final = await store.agent(auth.agent_id)
if isinstance(agent, AgentIdentityFailure):
raise_identity_failure(agent)
if agent is None:
retired: Final = await store.retired_agent(auth.agent_id)
if isinstance(retired, AgentIdentityFailure):
raise_identity_failure(retired)
if auth.managed_agent_context is not None or retired:
raise_identity_failure(AgentIdentityFailure(message="Agent no longer exists"))
return
if not agent.identity_managed:
return
if auth.jwt_claims and auth.managed_agent_context is None:
raise_identity_failure(AgentIdentityFailure(message="A managed agent requires a matching verified identity"))
failure: Final = actor_admission_failure(agent, auth.managed_agent_context)
if failure is not None:
raise_identity_failure(failure)
auth.managed_agent_policy = agent
auth.billing_agent_policy = agent
auth.requires_fresh_policy = True
if (
auth.managed_agent_context is not None
and auth.managed_agent_context.mode == "delegated"
and not delegation_verified
):
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import verified_human_agent_grants
grants: Final = await verified_human_agent_grants(auth.managed_agent_context.user_id, auth.team_id)
if agent.agent_id not in grants:
raise_identity_failure(
AgentIdentityFailure(message="The delegated user is not permitted to invoke this agent")
)
def actor_admission_failure(
agent: AgentResponse,
context: ManagedAgentContext | None,
) -> AgentIdentityFailure | None:
if not agent.enabled or agent.identity is None or not agent.identity.active:
return AgentIdentityFailure(message="Agent execution is disabled")
if context is None:
return AgentIdentityFailure(message="This agent requires its bound identity provider token")
if context.agent_id != agent.agent_id or context.binding_revision != agent.identity.revision:
return AgentIdentityFailure(message="Agent identity changed during authentication; retry")
if agent.execution_mode not in (context.mode, "both"):
return AgentIdentityFailure(message="Agent is not enabled for this execution mode")
if context.mode == "delegated" and not context.user_id:
return AgentIdentityFailure(message="A verified human subject is required")
return None
_INVOCATION_COST: Final = TypeAdapter(Annotated[float, Field(ge=0, allow_inf_nan=False)])
def invocation_target(route: str, body: Mapping[str, object]) -> str | None:
components: Final = tuple(route.strip("/").split("/"))
path: Final = components[1:] if components and components[0] == "v1" else components
if len(path) >= 2 and path[0] == "a2a":
return path[1] or None
model: Final = body.get("model")
return model.removeprefix("a2a/") or None if isinstance(model, str) and model.startswith("a2a/") else None
async def prepare_agent_invocation(
auth: UserAPIKeyAuth, target_name: str, store: AgentIdentityStore | None, *, billable: bool = True
) -> None:
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import AgentRequestHandler
from litellm.proxy.common_utils.registry_read_through import get_agent_with_read_through
registered: Final = await get_agent_with_read_through(target_name)
if registered is None:
return
registered_managed: Final = registered.identity_managed or registered.identity is not None
if store is None and registered_managed:
raise_identity_failure(
AgentIdentityFailure(code="policy_unavailable", message="Managed agent policy requires a database")
)
target: Final = await store.agent(registered.agent_id) if store is not None else None
if isinstance(target, AgentIdentityFailure):
raise_identity_failure(target)
if target is None and registered_managed:
raise_identity_failure(AgentIdentityFailure(message="Invoked agent no longer exists"))
effective: Final = target if target is not None else registered
if not effective.identity_managed and auth.managed_agent_policy is None:
return
if not await AgentRequestHandler.is_agent_allowed(effective.agent_id, auth):
raise_identity_failure(AgentIdentityFailure(message="The caller is not permitted to invoke this agent"))
auth.invoked_agent_id = effective.agent_id
auth.invoked_agent_policy = effective
if auth.agent_id is None and effective.identity_managed:
auth.billing_agent_policy = effective
raw_fee: Final = (effective.litellm_params or MappingProxyType({})).get("cost_per_query", 0.0) if billable else 0.0
try:
fee: Final = _INVOCATION_COST.validate_python(raw_fee)
except ValidationError:
raise_identity_failure(
AgentIdentityFailure(code="policy_unavailable", message="Agent invocation price is invalid")
)
auth.agent_invocation_cost = fee

View file

@ -28,6 +28,7 @@ if TYPE_CHECKING:
LiteLLM_AgentIdentityWhereUniqueInput,
LiteLLM_AgentsTableInclude,
LiteLLM_AgentsTableWhereUniqueInput,
LiteLLM_RetiredAgentWhereUniqueInput,
LiteLLM_VerifiedSubjectCreateInput,
LiteLLM_VerifiedSubjectUpsertInput,
LiteLLM_VerifiedSubjectWhereUniqueInput,
@ -183,7 +184,8 @@ class AgentIdentityStore:
if self.retired_agents is None:
return AgentIdentityFailure(code="policy_unavailable", message="Agent history is unavailable")
try:
return await self.retired_agents.table.find_unique(where={"original_agent_id": agent_id}) is not None
where: Final[LiteLLM_RetiredAgentWhereUniqueInput] = {"original_agent_id": agent_id}
return await self.retired_agents.table.find_unique(where=where) is not None
except Exception:
return AgentIdentityFailure(code="policy_unavailable", message="Agent history is unavailable")

View file

@ -14,6 +14,7 @@ import math
import re
import time
from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence
from functools import partial
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Generic, Literal, Optional, Protocol, TypeAlias
@ -77,6 +78,7 @@ from litellm.proxy.agent_endpoints.auth.agent_caller import (
load_agent_caller_team,
load_agent_caller_user,
)
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_policy
from litellm.proxy.auth.budget_throttle import (
budget_throttle_percentage,
should_throttle_budget_exceeded,
@ -1057,6 +1059,20 @@ async def common_checks(
code=status.HTTP_400_BAD_REQUEST,
)
managed_policy: Final = managed_agent_policy(valid_token)
if _model and valid_token is not None and managed_policy is not None:
managed_models: Final = (managed_policy.object_permission or MappingProxyType({})).get("models", ())
if not isinstance(managed_models, (list, tuple)) or not managed_models:
raise HTTPException(403, "This agent has no model grants")
_can_object_call_model(
model=_resolve_team_alias(_model, valid_token.team_model_aliases, valid_token.team_id, llm_router),
llm_router=llm_router,
models=list(managed_models),
team_id=valid_token.team_id,
object_type="agent",
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
)
await _check_agent_access_group_model_access(model=_model, valid_token=valid_token, llm_router=llm_router)
await _check_agent_caller_model_access(
model=_model,
@ -2642,7 +2658,7 @@ async def get_user_object(
)
if should_check_db:
response = await _user_table(UserRepository(prisma_client)).find_unique(
response = await _user_table(UserRepository(prisma_client, use_writer=bool(check_db_only))).find_unique(
where={"user_id": user_id}, include={"organization_memberships": True}
)
@ -2680,7 +2696,7 @@ async def get_user_object(
budget_duration=new_user_params["budget_duration"]
)
response = await _user_table(UserRepository(prisma_client)).create(
response = await _user_table(UserRepository(prisma_client, use_writer=bool(check_db_only))).create(
data=new_user_params,
include={"organization_memberships": True},
)
@ -3126,9 +3142,9 @@ class TeamNotFoundError(HTTPException):
@log_db_metrics
async def _get_team_db_check(
team_id: str, prisma_client: PrismaClient, team_id_upsert: bool | None = None
team_id: str, prisma_client: PrismaClient, team_id_upsert: bool | None = None, *, use_writer: bool = False
) -> "_PrismaTeamRow | None":
response = await _team_table(TeamRepository(prisma_client)).find_unique(
response = await _team_table(TeamRepository(prisma_client, use_writer=use_writer)).find_unique(
where={"team_id": team_id}, include=_TEAM_GRANT_RELATIONS
)
@ -3162,6 +3178,7 @@ async def _get_team_object_from_user_api_key_cache(
proxy_logging_obj: ProxyLogging | None,
key: str,
team_id_upsert: bool | None = None,
use_writer: bool = False,
) -> LiteLLM_TeamTableCachedObj:
db_access_time_key: Final = key
should_check_db: Final = _should_check_db(
@ -3170,7 +3187,9 @@ async def _get_team_object_from_user_api_key_cache(
db_cache_expiry=db_cache_expiry,
)
if should_check_db:
response = await _get_team_db_check(team_id=team_id, prisma_client=prisma_client, team_id_upsert=team_id_upsert)
response = await _get_team_db_check(
team_id=team_id, prisma_client=prisma_client, team_id_upsert=team_id_upsert, use_writer=use_writer
)
# The database answered and the row is not there. Distinct from every
# other failure here, which leaves the team's grant unknown.
if response is None:
@ -3192,8 +3211,11 @@ async def _get_team_object_from_user_api_key_cache(
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=proxy_logging_obj,
check_db_only=use_writer,
)
except Exception as e:
if use_writer:
raise
verbose_proxy_logger.debug(
"Failed to load object_permission for team %s with object_permission_id=%s: %s",
team_id,
@ -3283,6 +3305,7 @@ async def get_team_object(
db_cache_expiry=db_cache_expiry,
key=key,
team_id_upsert=team_id_upsert,
use_writer=bool(check_db_only),
)
except TeamNotFoundError:
raise
@ -3328,16 +3351,15 @@ async def get_access_object(
prisma_client: DatabaseClient | None,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging | None = None,
*,
check_db_only: bool = False,
) -> LiteLLM_AccessGroupTable:
"""
- Check if access_group_id in proxy AccessGroupTable
- Always checks cache first, then DB only when not found in cache
- Checks cache first unless authoritative writer admission is requested
- if valid, return LiteLLM_AccessGroupTable object
- if not, then raise an error
Unlike get_team_object, this has no check_cache_only or check_db_only flags;
it always follows cache-first-then-db semantics.
Raises:
- HTTPException: If access group doesn't exist in db or cache (status_code=404)
"""
@ -3346,18 +3368,19 @@ async def get_access_object(
key: Final = f"access_group_id:{access_group_id}"
cached_access_obj: Final = await user_api_key_cache.async_get_cache(
key=key,
model_type=LiteLLM_AccessGroupTable,
cached_access_obj: Final = (
None
if check_db_only
else await user_api_key_cache.async_get_cache(key=key, model_type=LiteLLM_AccessGroupTable)
)
if cached_access_obj is not None:
return cached_access_obj
# Not in cache - fetch from DB
try:
response: Final = await _dictable_table(AccessGroupRepository(prisma_client), "access_group").find_unique(
where={"access_group_id": access_group_id}
)
response: Final = await _dictable_table(
AccessGroupRepository(prisma_client, use_writer=check_db_only), "access_group"
).find_unique(where={"access_group_id": access_group_id})
if response is None:
raise HTTPException(
@ -3384,8 +3407,12 @@ async def get_access_object(
access_group_id,
)
raise HTTPException(
status_code=404,
detail={"error": f"Access group doesn't exist in db. Access group={access_group_id}. Error: {e}"},
status_code=503 if check_db_only else 404,
detail=(
"Access group policy is unavailable"
if check_db_only
else {"error": f"Access group doesn't exist in db. Access group={access_group_id}. Error: {e}"}
),
)
@ -3719,6 +3746,8 @@ async def _fetch_key_object_from_db_with_reconnect(
parent_otel_span: Span | None,
proxy_logging_obj: ProxyLogging | None,
deadline_seconds: float | None = None,
*,
check_db_only: bool = False,
) -> BaseModel | None:
"""
Fetch key object from DB and retry once if a DB connection error can be healed.
@ -3732,6 +3761,7 @@ async def _fetch_key_object_from_db_with_reconnect(
prisma_client=prisma_client,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=check_db_only,
),
name="key",
deadline_seconds=deadline_seconds,
@ -3743,10 +3773,13 @@ async def _fetch_key_object_from_db_unbounded(
prisma_client: PrismaClient,
parent_otel_span: Span | None,
proxy_logging_obj: ProxyLogging | None,
*,
check_db_only: bool = False,
) -> BaseModel | None:
fetch: Final = partial(prisma_client.get_data, use_writer=True) if check_db_only else prisma_client.get_data
async with db_lookup_gate.current():
try:
return await prisma_client.get_data(
return await fetch(
token=hashed_token,
table_name="combined_view",
parent_otel_span=parent_otel_span,
@ -3768,7 +3801,7 @@ async def _fetch_key_object_from_db_unbounded(
lock_timeout_seconds=auth_reconnect_lock_timeout,
)
if did_reconnect:
return await prisma_client.get_data(
return await fetch(
token=hashed_token,
table_name="combined_view",
parent_otel_span=parent_otel_span,
@ -3856,6 +3889,8 @@ async def get_key_object(
parent_otel_span: Span | None = None,
proxy_logging_obj: ProxyLogging | None = None,
check_cache_only: bool | None = None,
*,
check_db_only: bool = False,
) -> UserAPIKeyAuth:
"""
- Check if team id in proxy Team Table
@ -3870,9 +3905,8 @@ async def get_key_object(
# Same flow as before: use cache only when we have a hit we can turn into UserAPIKeyAuth
# (dict from Redis / model_dump, or UserAPIKeyAuth from in-memory). Otherwise fall through to DB.
user_api_key_auth: Final = await user_api_key_cache.async_get_cache(
key=key,
model_type=UserAPIKeyAuth,
user_api_key_auth: Final = (
None if check_db_only else await user_api_key_cache.async_get_cache(key=key, model_type=UserAPIKeyAuth)
)
if user_api_key_auth is not None:
return _copy_user_api_key_auth_for_cache(user_api_key_obj=user_api_key_auth)
@ -3886,6 +3920,7 @@ async def get_key_object(
prisma_client=prisma_client,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=check_db_only,
)
if _valid_token is None:
@ -3899,7 +3934,7 @@ async def get_key_object(
_response: Final = UserAPIKeyAuth.model_validate(_valid_token.model_dump(exclude_none=True))
# Load object_permission if object_permission_id exists but object_permission is not loaded
if _response.object_permission_id and not _response.object_permission:
if _response.object_permission_id and (check_db_only or not _response.object_permission):
try:
_response.object_permission = await get_object_permission(
object_permission_id=_response.object_permission_id,
@ -3907,14 +3942,20 @@ async def get_key_object(
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=check_db_only,
)
except Exception as e:
if check_db_only:
raise
verbose_proxy_logger.debug(
"Failed to load object_permission for key with object_permission_id=%s: %s",
_response.object_permission_id,
e,
)
if check_db_only:
return _response
# save the key object to cache
await _cache_key_object(
hashed_token=hashed_token,
@ -3944,6 +3985,7 @@ async def get_object_permission(
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Span | None = None,
proxy_logging_obj: ProxyLogging | None = None,
check_db_only: bool = False,
) -> LiteLLM_ObjectPermissionTable | None:
"""
- Check if object permission id in proxy ObjectPermissionTable
@ -3955,9 +3997,13 @@ async def get_object_permission(
# check if in cache
key: Final = object_permission_cache_key(object_permission_id)
deserialized_perm: Final = await user_api_key_cache.async_get_cache(
key=key,
model_type=LiteLLM_ObjectPermissionTable,
deserialized_perm: Final = (
None
if check_db_only
else await user_api_key_cache.async_get_cache(
key=key,
model_type=LiteLLM_ObjectPermissionTable,
)
)
if deserialized_perm is not None:
return deserialized_perm
@ -3965,10 +4011,12 @@ async def get_object_permission(
# else, check db
try:
response: Final = await _dictable_table(
ObjectPermissionRepository(prisma_client), "object_permission"
ObjectPermissionRepository(prisma_client, use_writer=check_db_only), "object_permission"
).find_unique(where={"object_permission_id": object_permission_id})
if response is None:
if check_db_only:
raise HTTPException(status_code=403, detail="Referenced object permission does not exist")
return None
_perm_obj: Final = LiteLLM_ObjectPermissionTable.model_validate(response.dict())
@ -3981,6 +4029,8 @@ async def get_object_permission(
return _perm_obj
except Exception:
if check_db_only:
raise
return None
@ -4190,6 +4240,7 @@ async def _get_resources_from_access_groups(
prisma_client: DatabaseClient | None = None,
user_api_key_cache: UserApiKeyCache | None = None,
proxy_logging_obj: ProxyLogging | None = None,
check_db_only: bool = False,
) -> list[str]:
"""
Fetch access groups by their IDs (from cache or DB) and collect
@ -4232,9 +4283,12 @@ async def _get_resources_from_access_groups(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
check_db_only=check_db_only,
)
resources.extend(getattr(ag, resource_field, []))
except Exception:
if check_db_only:
raise
verbose_proxy_logger.debug(
"Could not fetch access group %s for resource field %s",
ag_id,
@ -4267,6 +4321,7 @@ async def _get_mcp_server_ids_from_access_groups(
prisma_client: PrismaClient | None = None,
user_api_key_cache: UserApiKeyCache | None = None,
proxy_logging_obj: ProxyLogging | None = None,
check_db_only: bool = False,
) -> list[str]:
"""
Collect MCP server IDs from unified access groups.
@ -4278,6 +4333,7 @@ async def _get_mcp_server_ids_from_access_groups(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
check_db_only=check_db_only,
)
@ -4286,6 +4342,7 @@ async def _get_agent_ids_from_access_groups(
prisma_client: PrismaClient | None = None,
user_api_key_cache: UserApiKeyCache | None = None,
proxy_logging_obj: ProxyLogging | None = None,
check_db_only: bool = False,
) -> list[str]:
"""
Collect agent IDs from unified access groups.
@ -4297,6 +4354,7 @@ async def _get_agent_ids_from_access_groups(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
check_db_only=check_db_only,
)
@ -4496,26 +4554,37 @@ async def _check_agent_access_group_model_access(
"""Attached groups naming no model deny every model; the empty allowlist in ``_can_object_call_model`` allows."""
if not model or valid_token is None or not valid_token.agent_id:
return True
ceiling: Final = await resolve_ceiling(valid_token.agent_id)
if ceiling is None:
return True
if not ceiling.models:
raise ModelAccessDeniedProxyException(
message=model_access_denied_client_message(model=model),
internal_message=f"agent {valid_token.agent_id} access groups {ceiling.access_group_ids} grant no models",
type=ProxyErrorTypes.agent_model_access_denied,
param="model",
code=status.HTTP_403_FORBIDDEN,
)
dispatched: Final = _resolve_team_alias(model, valid_token.team_model_aliases, valid_token.team_id, llm_router)
return _can_object_call_model(
model=dispatched,
llm_router=llm_router,
models=sorted(ceiling.models),
team_id=valid_token.team_id,
object_type="agent",
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_managed_agent_ceilings
managed: Final = managed_agent_policy(valid_token)
unmanaged: Final = await resolve_ceiling(valid_token.agent_id) if managed is None else None
ceilings: Final = (
await resolve_managed_agent_ceilings(managed)
if managed is not None
else (unmanaged,)
if unmanaged is not None
else ()
)
dispatched: Final = _resolve_team_alias(model, valid_token.team_model_aliases, valid_token.team_id, llm_router)
for ceiling in ceilings:
if not ceiling.models:
raise ModelAccessDeniedProxyException(
message=model_access_denied_client_message(model=model),
internal_message=f"agent {valid_token.agent_id} access groups grant no models",
type=ProxyErrorTypes.agent_model_access_denied,
param="model",
code=status.HTTP_403_FORBIDDEN,
)
_can_object_call_model(
model=dispatched,
llm_router=llm_router,
models=sorted(ceiling.models),
team_id=valid_token.team_id,
object_type="agent",
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
)
return True
LoadedCallerTeam: TypeAlias = LiteLLM_TeamTable | None

View file

@ -52,6 +52,10 @@ from litellm.proxy._types import (
TeamMemberAddRequest,
UserAPIKeyAuth,
)
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_route_allowed
from litellm.proxy.agent_endpoints.identity import has_legacy_identity
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore, resolve_managed_agent
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
from litellm.proxy.auth.auth_checks import can_team_access_model
from litellm.proxy.auth.model_access_denied import (
ModelAccessDeniedHTTPException,
@ -67,6 +71,7 @@ from litellm.proxy.common_utils.user_api_key_cache import (
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.repositories.user_repository import UserRepository
from litellm.types.agents import AgentResponse
from litellm.types.proxy.agent_identity import AgentIdentityFailure
from litellm.types.proxy.auth.auth_checks import UserNotFoundError
from .auth_checks import (
@ -157,6 +162,8 @@ class HeaderTeam:
class AgentLookup(Protocol):
"""The registered-agent lookups a JWT agent claim is matched against."""
def get_agent_list(self) -> Sequence[AgentResponse]: ...
def get_agent_by_id(self, agent_id: str) -> AgentResponse | None:
"""The agent registered under ``agent_id``, if any."""
@ -167,6 +174,9 @@ class AgentLookup(Protocol):
class _NoRegisteredAgents:
"""The lookup in force until the proxy binds its agent registry: no agent is registered, so no claim matches."""
def get_agent_list(self) -> tuple[AgentResponse, ...]:
return ()
def get_agent_by_id(self, agent_id: str) -> None:
return None
@ -398,7 +408,7 @@ class JWTHandler:
return []
def get_all_jwt_team_ids(self, token: dict) -> list[str]:
def get_all_jwt_team_ids(self, token: dict[str, object]) -> list[str]:
"""
Return team IDs from both the plural ``team_ids_jwt_field`` and the
singular ``team_id_jwt_field`` claim (string or list of strings), as a
@ -522,7 +532,7 @@ class JWTHandler:
team_id = default_value
return team_id
def get_team_alias(self, token: dict, default_value: str | None) -> str | None:
def get_team_alias(self, token: dict[str, object], default_value: str | None) -> str | None:
"""
Extract team name/alias from JWT token using the configured team_alias_jwt_field.
@ -1096,6 +1106,15 @@ class JWTHandler:
"options": options or None,
}
def managed_issuer_is_trusted(self, issuer: object) -> bool:
if not isinstance(issuer, str):
return False
configured: Final = self.litellm_jwtauth.issuers or ()
for item in configured:
if item.issuer == issuer:
return bool(item.audience) and not item.disable_audience_validation
return issuer == os.getenv("JWT_ISSUER") and bool(os.getenv("JWT_AUDIENCE"))
def _get_configured_issuer(self, token: str) -> JWTIssuerConfig | None:
litellm_jwtauth: Final[_JWTAuthSettings | None] = getattr(self, "litellm_jwtauth", None)
if litellm_jwtauth is None:
@ -1488,7 +1507,12 @@ class JWTAuthManager:
agent: Final = agent_registry.get_agent_by_id(agent_id=agent_claim) or agent_registry.get_agent_by_name(
agent_name=agent_claim
)
if agent is None:
if (
agent is None
or agent.identity_managed
or agent.identity is not None
or has_legacy_identity(agent.litellm_params)
):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=f"No registered agent matches JWT claim {jwt_handler.litellm_jwtauth.agent_id_jwt_field}={agent_claim}",
@ -2159,7 +2183,7 @@ class JWTAuthManager:
parent_otel_span: Span | None,
proxy_logging_obj: ProxyLogging,
team_id_upsert: bool | None,
) -> tuple:
) -> tuple[str | None, LiteLLM_TeamTable | None, LiteLLM_TeamMembership | None]:
"""
If JWT did not resolve team_id, but the user belongs to exactly one team
in LiteLLM, load that team (and membership when user_id is set) so that
@ -2478,12 +2502,39 @@ class JWTAuthManager:
"""Resolve and authorize JWT context; only normal admission supplies provisioning."""
handler: Final = jwt_handler
jwt_valid_token: Final = await JWTAuthManager.authenticate_jwt(api_key, handler)
managed: Final = await resolve_managed_agent(jwt_valid_token, prisma_client, cache=user_api_key_cache)
if managed is not None:
if not handler.managed_issuer_is_trusted(jwt_valid_token.get("iss")):
raise HTTPException(403, "Managed agents require trusted JWT issuer and audience validation")
if not managed_agent_route_allowed(route, request_method):
raise HTTPException(403, "Agent identities can only access inference and agent discovery routes")
evidence: Final = await AgentIdentityStore.from_client(prisma_client).record_authentication(managed)
if isinstance(evidence, AgentIdentityFailure):
raise_identity_failure(evidence)
if managed.mode == "autonomous":
return JWTAuthBuilderResult(
is_proxy_admin=False,
team_id=None,
team_object=None,
user_id=None,
user_email=None,
user_object=None,
org_id=None,
org_object=None,
end_user_id=None,
end_user_object=None,
token=api_key,
team_membership=None,
jwt_claims=jwt_valid_token,
agent_id=managed.agent_id,
managed_agent_context=managed,
)
team_id_upsert: Final = provisioning.team_id_upsert if provisioning is not None else False
model: Final = request_data.get("model")
requested_model: Final = model if isinstance(model, str) else None
# Check RBAC
rbac_role: Final = handler.get_rbac_role(token=jwt_valid_token)
rbac_role: Final = handler.get_rbac_role(token=jwt_valid_token) if managed is None else None
await JWTAuthManager.check_rbac_role(handler, jwt_valid_token, general_settings, request_data, route, rbac_role)
# Check Scope Based Access
@ -2499,7 +2550,11 @@ class JWTAuthManager:
object_id = handler.get_object_id(token=jwt_valid_token, default_value=None)
# Get basic user info
user_id, user_email, valid_user_email = await JWTAuthManager.get_user_info(handler, jwt_valid_token)
user_id, user_email, valid_user_email = (
(managed.user_id, None, None)
if managed is not None
else await JWTAuthManager.get_user_info(handler, jwt_valid_token)
)
# Get IDs
org_id: Final = handler.get_org_id(token=jwt_valid_token, default_value=None)
@ -2514,23 +2569,31 @@ class JWTAuthManager:
elif rbac_role == LitellmUserRoles.INTERNAL_USER:
user_id = object_id
agent_id: Final = JWTAuthManager.resolve_agent_id(
jwt_handler=handler,
jwt_valid_token=jwt_valid_token,
agent_registry=handler.agent_lookup,
agent_id: Final = (
managed.agent_id
if managed is not None
else JWTAuthManager.resolve_agent_id(
jwt_handler=handler,
jwt_valid_token=jwt_valid_token,
agent_registry=handler.agent_lookup,
)
)
# Check admin access
admin_result: Final = await JWTAuthManager.check_admin_access(
handler,
scopes,
route,
user_id,
org_id,
api_key,
jwt_valid_token,
user_email=user_email,
agent_id=agent_id,
admin_result: Final = (
None
if managed is not None
else await JWTAuthManager.check_admin_access(
handler,
scopes,
route,
user_id,
org_id,
api_key,
jwt_valid_token,
user_email=user_email,
agent_id=agent_id,
)
)
if admin_result:
await JWTAuthManager._attach_team_from_header_for_admin(
@ -2673,8 +2736,47 @@ class JWTAuthManager:
team_id_upsert=team_id_upsert,
)
if team_id and not JWTAuthManager._team_has_passthrough_route_access(
team_object=team_object,
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import resolve_delegated_agent_team
claimed_teams: Final[frozenset[str]] = (
frozenset(handler.get_all_jwt_team_ids(jwt_valid_token)) if managed is not None else frozenset()
)
scoped_teams: Final[frozenset[str] | None] = claimed_teams or (
frozenset((team_id,))
if managed is not None and team_id and handler.get_team_alias(jwt_valid_token, default_value=None)
else None
)
granting_team: Final = (
await resolve_delegated_agent_team(
managed.user_id,
managed.agent_id,
team_id,
explicit_team=header_team is not None,
allowed_team_ids=None if handler.litellm_jwtauth.fallback_to_db_teams else scoped_teams,
)
if managed is not None
else team_id
)
if granting_team is not None and granting_team != team_id:
if not JWTAuthManager._is_team_route_allowed(route, request_method, handler):
raise HTTPException(403, "The granting team is not allowed to access this route")
selected_team_id: Final[str | None] = granting_team if granting_team is not None else team_id
selected_team_object: Final[LiteLLM_TeamTable | None] = (
await get_team_object(
team_id=selected_team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=True,
)
if selected_team_id is not None and selected_team_id != team_id
else team_object
)
if selected_team_id and not JWTAuthManager._team_has_passthrough_route_access(
team_object=selected_team_object,
route=route,
request_method=request_method,
team_allowed_routes=handler.litellm_jwtauth.team_allowed_routes,
@ -2696,7 +2798,7 @@ class JWTAuthManager:
user_email=user_email,
org_id=org_id,
end_user_id=end_user_id,
team_id=team_id,
team_id=selected_team_id,
valid_user_email=valid_user_email,
jwt_handler=handler,
prisma_client=prisma_client,
@ -2705,13 +2807,13 @@ class JWTAuthManager:
proxy_logging_obj=proxy_logging_obj,
route=route,
org_alias=org_alias,
user_id_upsert=provisioning.user_id_upsert if provisioning is not None else False,
user_id_upsert=provisioning.user_id_upsert if provisioning is not None and managed is None else False,
)
# Derive org_id from org_object if resolved by alias
resolved_org_id: Final = org_object.organization_id if org_object else org_id
if provisioning is not None:
if provisioning is not None and managed is None:
await JWTAuthManager.sync_user_role_and_teams(
jwt_handler=handler,
jwt_valid_token=jwt_valid_token,
@ -2721,7 +2823,7 @@ class JWTAuthManager:
)
# If JWT did not resolve team_id, attempt a team fallback.
if team_id is None and db_team_fallback:
if selected_team_id is None and db_team_fallback:
(
team_id,
team_object,
@ -2750,7 +2852,7 @@ class JWTAuthManager:
team_allowed_routes=handler.litellm_jwtauth.team_allowed_routes,
):
JWTAuthManager._raise_team_passthrough_route_denial(route=route)
elif team_id is None:
elif selected_team_id is None:
(
team_id,
team_object,
@ -2764,9 +2866,9 @@ class JWTAuthManager:
proxy_logging_obj=proxy_logging_obj,
team_id_upsert=team_id_upsert,
)
elif provisional_header_team is not None and team_id == provisional_header_team.team_id:
elif provisional_header_team is not None and selected_team_id == provisional_header_team.team_id:
JWTAuthManager._validate_header_team_in_db_membership(
team_id=team_id,
team_id=selected_team_id,
user_object=user_object,
header_value=provisional_header_team.header_value,
)
@ -2783,28 +2885,35 @@ class JWTAuthManager:
),
)
authorized_team_id: Final[str | None] = selected_team_id if selected_team_id is not None else team_id
authorized_team_object: Final[LiteLLM_TeamTable | None] = (
selected_team_object if selected_team_id is not None else team_object
)
## MAP USER TO TEAMS
if provisioning is not None:
if provisioning is not None and managed is None:
await JWTAuthManager.map_user_to_teams(
user_object=user_object,
team_object=team_object,
team_object=authorized_team_object,
)
# Validate that a valid rbac id is returned for spend tracking
JWTAuthManager.validate_object_id(
user_id=user_id,
team_id=team_id,
team_id=authorized_team_id,
enforce_rbac=bool(general_settings.get("enforce_rbac", False)),
is_proxy_admin=False,
)
# check if user is proxy admin
is_proxy_admin: Final = bool(user_object and user_object.user_role == LitellmUserRoles.PROXY_ADMIN)
is_proxy_admin: Final = managed is None and bool(
user_object and user_object.user_role == LitellmUserRoles.PROXY_ADMIN
)
return JWTAuthBuilderResult(
is_proxy_admin=is_proxy_admin,
team_id=team_id,
team_object=team_object,
team_id=authorized_team_id,
team_object=authorized_team_object,
user_id=user_id,
user_email=(user_object.user_email if user_object is not None and user_object.user_email else user_email),
user_object=user_object,
@ -2816,6 +2925,7 @@ class JWTAuthManager:
team_membership=team_membership_object,
jwt_claims=jwt_valid_token,
agent_id=agent_id,
managed_agent_context=managed,
)
@staticmethod
@ -2826,11 +2936,13 @@ class JWTAuthManager:
"""Keep JWT identity and permission attribution identical across consumers."""
user: Final = result["user_object"]
admin: Final = result["is_proxy_admin"]
return UserAPIKeyAuth(
auth: Final = UserAPIKeyAuth(
api_key=None,
user_role=(
LitellmUserRoles.PROXY_ADMIN
if admin
else LitellmUserRoles.INTERNAL_USER
if result.get("managed_agent_context") is not None
else LitellmUserRoles(user.user_role)
if user is not None and user.user_role is not None
else LitellmUserRoles.INTERNAL_USER
@ -2852,3 +2964,8 @@ class JWTAuthManager:
user_id=result["user_id"],
),
)
auth.managed_agent_context = result.get("managed_agent_context")
auth._managed_delegation_verified = ( # pyright: ignore[reportPrivateUsage] # JWT admission produces the one-shot proof consumed by managed authorization
auth.managed_agent_context is not None and auth.managed_agent_context.mode == "delegated"
)
return auth

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