chore: merge upstream main into reasoning history support
|
|
@ -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 ;;
|
||||
|
|
|
|||
BIN
.github/assets/roi-calculator/00-original-setup.png
vendored
Normal file
|
After Width: | Height: | Size: 80 KiB |
BIN
.github/assets/roi-calculator/01-connect-github.png
vendored
Normal file
|
After Width: | Height: | Size: 58 KiB |
BIN
.github/assets/roi-calculator/02-repositories.png
vendored
Normal file
|
After Width: | Height: | Size: 63 KiB |
BIN
.github/assets/roi-calculator/03-estimator-schedule.png
vendored
Normal file
|
After Width: | Height: | Size: 70 KiB |
BIN
.github/assets/roi-calculator/04-backfill-progress.png
vendored
Normal file
|
After Width: | Height: | Size: 47 KiB |
BIN
.github/assets/roi-calculator/06-overview.png
vendored
Normal file
|
After Width: | Height: | Size: 76 KiB |
BIN
.github/assets/roi-calculator/07-people-unmatched.png
vendored
Normal file
|
After Width: | Height: | Size: 75 KiB |
BIN
.github/assets/roi-calculator/08-match-email.png
vendored
Normal file
|
After Width: | Height: | Size: 39 KiB |
BIN
.github/assets/roi-calculator/09-people-matched.png
vendored
Normal file
|
After Width: | Height: | Size: 70 KiB |
BIN
.github/assets/roi-calculator/10-pr-reasoning.png
vendored
Normal file
|
After Width: | Height: | Size: 93 KiB |
BIN
.github/assets/roi-calculator/11-settings.png
vendored
Normal file
|
After Width: | Height: | Size: 73 KiB |
BIN
.github/assets/roi-calculator/12-restart-setup.png
vendored
Normal file
|
After Width: | Height: | Size: 39 KiB |
BIN
.github/assets/roi-calculator/13-advanced-settings.png
vendored
Normal file
|
After Width: | Height: | Size: 81 KiB |
BIN
.github/assets/roi-calculator/14-overview-pulls.png
vendored
Normal file
|
After Width: | Height: | Size: 72 KiB |
BIN
.github/assets/roi-calculator/15-sample-preview.png
vendored
Normal file
|
After Width: | Height: | Size: 76 KiB |
BIN
.github/assets/roi-calculator/16-calculator-sidebar.png
vendored
Normal file
|
After Width: | Height: | Size: 50 KiB |
BIN
.github/assets/roi-calculator/19-matching-calculator-icons.png
vendored
Normal file
|
After Width: | Height: | Size: 59 KiB |
BIN
.github/assets/roi-calculator/20-partial-repository-report.png
vendored
Normal file
|
After Width: | Height: | Size: 57 KiB |
BIN
.github/assets/roi-calculator/21-empty-repository-preserved-report.png
vendored
Normal file
|
After Width: | Height: | Size: 56 KiB |
BIN
.github/assets/roi-calculator/22-partial-calculation-explanation.png
vendored
Normal file
|
After Width: | Height: | Size: 65 KiB |
BIN
.github/assets/roi-calculator/23-estimator-outage-preserved-report.png
vendored
Normal file
|
After Width: | Height: | Size: 55 KiB |
4
.github/workflows/test-unit.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -99,7 +99,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 44358
|
||||
"limit": 44802
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 109
|
||||
|
|
|
|||
86
litellm-rust/Cargo.lock
generated
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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" }
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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::{
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
}
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
})
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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::*;
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
||||
|
|
|
|||
137
litellm-rust/crates/python-bridge/src/routes/traces.rs
Normal 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, ¶meters).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)
|
||||
}
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
7
litellm-rust/crates/traces/AGENTS.md
Normal 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
|
||||
24
litellm-rust/crates/traces/Cargo.toml
Normal 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
|
||||
32
litellm-rust/crates/traces/config/reader.xml
Normal 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>
|
||||
47
litellm-rust/crates/traces/migrations/0001_otel_traces.sql
Normal 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
|
||||
25
litellm-rust/crates/traces/migrations/0002_agent_traces.sql
Normal 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
|
||||
|
|
@ -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
|
||||
42
litellm-rust/crates/traces/migrations/0004_spend_logs.sql
Normal 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)
|
||||
|
|
@ -0,0 +1 @@
|
|||
ALTER TABLE {database}.otel_traces MODIFY TTL toDateTime(Timestamp) + INTERVAL {trace_retention_days} DAY
|
||||
|
|
@ -0,0 +1 @@
|
|||
ALTER TABLE {database}.agent_traces_by_key MODIFY TTL toDateTime(StartTs) + INTERVAL {trace_retention_days} DAY
|
||||
|
|
@ -0,0 +1 @@
|
|||
ALTER TABLE {database}.spend_logs MODIFY TTL toDateTime(start_time) + INTERVAL {spend_log_retention_days} DAY
|
||||
35
litellm-rust/crates/traces/src/error.rs
Normal 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,
|
||||
}
|
||||
151
litellm-rust/crates/traces/src/insert.rs
Normal 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)
|
||||
));
|
||||
}
|
||||
}
|
||||
90
litellm-rust/crates/traces/src/lib.rs
Normal 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
|
||||
}
|
||||
}
|
||||
221
litellm-rust/crates/traces/src/otlp.rs
Normal 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(),
|
||||
}
|
||||
}
|
||||
87
litellm-rust/crates/traces/src/schema.rs
Normal 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(())
|
||||
}
|
||||
114
litellm-rust/crates/traces/src/sql.rs
Normal 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)
|
||||
}
|
||||
273
litellm-rust/crates/traces/tests/admin_sql.rs
Normal 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!("{}¶m_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",
|
||||
¶meters).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(())
|
||||
}
|
||||
40
litellm-rust/crates/traces/tests/insert.rs
Normal 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());
|
||||
}
|
||||
428
litellm-rust/crates/traces/tests/migrations.rs
Normal 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());
|
||||
}
|
||||
47
litellm-rust/crates/traces/tests/otlp.rs
Normal 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());
|
||||
}
|
||||
11
litellm-rust/crates/traces/tests/queries.rs
Normal 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);
|
||||
}
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
100
litellm/integrations/clickhouse/clickhouse_batch_logger.py
Normal 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
|
||||
11
litellm/integrations/clickhouse/schema.py
Normal 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)
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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({})
|
||||
),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 = []
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
]
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
)
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 []
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 ()
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
252
litellm/proxy/agent_endpoints/auth/managed_authorization.py
Normal 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
|
||||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||