chore(mcp): merge main into listed-tool metadata branch

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-30 18:15:26 +00:00
commit c8870f0080
191 changed files with 9554 additions and 679 deletions

View file

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

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-enterprise"
version = "0.1.71"
version = "0.1.72"
description = "Package for LiteLLM Enterprise features"
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.1.71"
version = "0.1.72"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-enterprise==",

View file

@ -0,0 +1,97 @@
-- AlterTable
ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN IF NOT EXISTS "enabled" BOOLEAN NOT NULL DEFAULT true,
ADD COLUMN IF NOT EXISTS "execution_mode" TEXT NOT NULL DEFAULT 'autonomous',
ADD COLUMN IF NOT EXISTS "identity_managed" BOOLEAN NOT NULL DEFAULT false;
-- AlterTable
ALTER TABLE "LiteLLM_SpendLogs" ADD COLUMN IF NOT EXISTS "billing_agent_id" TEXT;
-- CreateTable
CREATE TABLE IF NOT EXISTS "LiteLLM_AgentIdentity" (
"agent_id" TEXT NOT NULL,
"active" BOOLEAN NOT NULL DEFAULT true,
"provider" TEXT NOT NULL,
"issuer" TEXT NOT NULL,
"tenant_id" TEXT NOT NULL,
"client_id" TEXT NOT NULL,
"service_principal_id" TEXT,
"required_roles" TEXT[] DEFAULT ARRAY[]::TEXT[],
"required_scopes" TEXT[] DEFAULT ARRAY['user_impersonation']::TEXT[],
"revision" TEXT NOT NULL,
"last_authenticated_at" TIMESTAMP(3),
CONSTRAINT "LiteLLM_AgentIdentity_pkey" PRIMARY KEY ("agent_id")
);
-- CreateTable
CREATE TABLE IF NOT EXISTS "LiteLLM_RetiredAgentIdentity" (
"binding_id" TEXT NOT NULL,
"agent_id" TEXT,
"provider" TEXT NOT NULL,
"issuer" TEXT NOT NULL,
"tenant_id" TEXT NOT NULL,
"client_id" TEXT NOT NULL,
CONSTRAINT "LiteLLM_RetiredAgentIdentity_pkey" PRIMARY KEY ("binding_id")
);
-- CreateTable
CREATE TABLE IF NOT EXISTS "LiteLLM_RetiredAgent" (
"original_agent_id" TEXT NOT NULL,
"retired_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
CONSTRAINT "LiteLLM_RetiredAgent_pkey" PRIMARY KEY ("original_agent_id")
);
-- CreateTable
CREATE TABLE IF NOT EXISTS "LiteLLM_VerifiedSubject" (
"subject_id" TEXT NOT NULL,
"issuer" TEXT NOT NULL,
"tenant_id" TEXT NOT NULL,
"oid" TEXT NOT NULL,
"kind" TEXT NOT NULL DEFAULT 'human',
"user_id" TEXT,
"verified_via" TEXT NOT NULL DEFAULT 'sso_interactive',
"verified_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
CONSTRAINT "LiteLLM_VerifiedSubject_pkey" PRIMARY KEY ("subject_id")
);
-- CreateIndex
CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_AgentIdentity_provider_tenant_id_client_id_key" ON "LiteLLM_AgentIdentity"("provider", "tenant_id", "client_id");
-- CreateIndex
CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_AgentIdentity_issuer_service_principal_id_key" ON "LiteLLM_AgentIdentity"("issuer", "service_principal_id");
-- CreateIndex
CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_RetiredAgentIdentity_provider_tenant_id_client_id_key" ON "LiteLLM_RetiredAgentIdentity"("provider", "tenant_id", "client_id");
-- CreateIndex
CREATE INDEX IF NOT EXISTS "LiteLLM_VerifiedSubject_user_id_idx" ON "LiteLLM_VerifiedSubject"("user_id");
-- CreateIndex
CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_VerifiedSubject_issuer_tenant_id_oid_key" ON "LiteLLM_VerifiedSubject"("issuer", "tenant_id", "oid");
-- AddForeignKey
DO $$
BEGIN
IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_AgentIdentity_agent_id_fkey') THEN
ALTER TABLE "LiteLLM_AgentIdentity" ADD CONSTRAINT "LiteLLM_AgentIdentity_agent_id_fkey" FOREIGN KEY ("agent_id") REFERENCES "LiteLLM_AgentsTable"("agent_id") ON DELETE CASCADE ON UPDATE CASCADE;
END IF;
END $$;
-- AddForeignKey
DO $$
BEGIN
IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_RetiredAgentIdentity_agent_id_fkey') THEN
ALTER TABLE "LiteLLM_RetiredAgentIdentity" ADD CONSTRAINT "LiteLLM_RetiredAgentIdentity_agent_id_fkey" FOREIGN KEY ("agent_id") REFERENCES "LiteLLM_AgentsTable"("agent_id") ON DELETE SET NULL ON UPDATE CASCADE;
END IF;
END $$;
-- AddForeignKey
DO $$
BEGIN
IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_VerifiedSubject_user_id_fkey') THEN
ALTER TABLE "LiteLLM_VerifiedSubject" ADD CONSTRAINT "LiteLLM_VerifiedSubject_user_id_fkey" FOREIGN KEY ("user_id") REFERENCES "LiteLLM_UserTable"("user_id") ON DELETE CASCADE ON UPDATE CASCADE;
END IF;
END $$;

View file

@ -78,6 +78,11 @@ model LiteLLM_AgentsTable {
object_permission_id String?
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
spend Float @default(0.0)
identity_managed Boolean @default(false)
enabled Boolean @default(true)
execution_mode String @default("autonomous")
identity LiteLLM_AgentIdentity?
retired_identities LiteLLM_RetiredAgentIdentity[]
tpm_limit Int?
rpm_limit Int?
session_tpm_limit Int?
@ -88,6 +93,56 @@ model LiteLLM_AgentsTable {
updated_by String
}
model LiteLLM_AgentIdentity {
agent_id String @id
active Boolean @default(true)
agent LiteLLM_AgentsTable @relation(fields: [agent_id], references: [agent_id], onDelete: Cascade)
provider String
issuer String
tenant_id String
client_id String
service_principal_id String?
required_roles String[] @default([])
required_scopes String[] @default(["user_impersonation"])
revision String @default(uuid())
last_authenticated_at DateTime?
@@unique([provider, tenant_id, client_id])
@@unique([issuer, service_principal_id])
}
model LiteLLM_RetiredAgentIdentity {
binding_id String @id @default(uuid())
agent_id String?
agent LiteLLM_AgentsTable? @relation(fields: [agent_id], references: [agent_id], onDelete: SetNull)
provider String
issuer String
tenant_id String
client_id String
@@unique([provider, tenant_id, client_id])
}
model LiteLLM_RetiredAgent {
original_agent_id String @id
retired_at DateTime @default(now())
}
model LiteLLM_VerifiedSubject {
subject_id String @id @default(uuid())
issuer String
tenant_id String
oid String
kind String @default("human")
user_id String?
user LiteLLM_UserTable? @relation(fields: [user_id], references: [user_id], onDelete: Cascade)
verified_via String @default("sso_interactive")
verified_at DateTime @default(now())
@@unique([issuer, tenant_id, oid])
@@index([user_id])
}
model LiteLLM_OrganizationTable {
organization_id String @id @default(uuid())
organization_alias String
@ -241,6 +296,7 @@ model LiteLLM_DeletedTeamTable {
// Track spend, rate limit, budget Users
model LiteLLM_UserTable {
verified_subjects LiteLLM_VerifiedSubject[]
user_id String @id
user_alias String?
team_id String?
@ -675,6 +731,7 @@ model LiteLLM_SpendLogs {
session_id String?
status String?
mcp_namespaced_tool_name String?
billing_agent_id String?
agent_id String?
proxy_server_request Json? @default("{}")
litellm_call_id String?

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-proxy-extras"
version = "0.4.102"
version = "0.4.103"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.4.102"
version = "0.4.103"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-proxy-extras==",

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -4,6 +4,7 @@ mod coercion;
mod credentials;
mod diagnostics;
mod errors;
mod execution;
mod http;
mod lifecycle;
mod logger;

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -33,7 +33,8 @@
"thinking-binding-controls-2026-08-01": "thinking-binding-controls-2026-08-01",
"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"
"web-search-2025-03-05": "web-search-2025-03-05",
"mid-conversation-output-config-2026-07-01": "mid-conversation-output-config-2026-07-01"
},
"azure_ai": {
"advisor-tool-2026-03-01": null,
@ -134,7 +135,8 @@
"token-efficient-tools-2025-02-19": null,
"tool-search-tool-2025-10-19": "tool-search-tool-2025-10-19",
"web-fetch-2025-09-10": null,
"web-search-2025-03-05": null
"web-search-2025-03-05": null,
"mid-conversation-output-config-2026-07-01": "mid-conversation-output-config-2026-07-01"
},
"bedrock_mantle": {
"advanced-tool-use-2025-11-20": "tool-search-tool-2025-10-19",

View file

@ -24,7 +24,7 @@ from litellm.types.caching import RedisPipelineIncrementOperation
from .base_cache import BaseCache
from .in_memory_cache import DEFAULT_MAX_SIZE_IN_MEMORY, InMemoryCache
from .redis_batch import BatchResult, RedisBatch, active_post_call_redis_batch
from .redis_batch import BatchResult, RedisBatch, active_post_call_redis_batch, active_request_redis_batch
from .redis_cache import RedisCache, RedisCircuitBreakerOpenError, log_redis_failure
if TYPE_CHECKING:
@ -279,6 +279,9 @@ class DualCache(BaseCache):
result = in_memory_result
if result is None and self.redis_cache is not None and local_only is False:
request_batch: Final = active_request_redis_batch(self.redis_cache)
if request_batch is not None and request_batch.read_as_missing(key):
return None
# If not found in in-memory cache, try fetching from Redis
redis_result: Final = await self.redis_cache.async_get_cache(key, parent_otel_span=parent_otel_span)
@ -323,6 +326,20 @@ class DualCache(BaseCache):
return sublist_keys, previous_access_times
def reserve_redis_batch_reads(self, keys: Sequence[str]) -> tuple[list[str], dict[str, float | None]]:
"""Reserve the memory-missed keys whose throttled Redis reads are due, as a batch read would."""
if self.redis_cache is None:
return [], {} # mutable-ok: API contract returns an empty list and dictionary
key_list: Final = list(keys) # mutable-ok: batch_get_cache takes a list
memory: Final = self.in_memory_cache
in_memory_result: Final = (
None
if memory is None # pyright: ignore[reportUnnecessaryComparison] # handle an absent in-memory tier
else memory.batch_get_cache(key_list)
)
result: Final = in_memory_result if in_memory_result is not None else tuple(None for _ in key_list)
return self._reserve_redis_batch_keys(time.time(), key_list, result)
def _rollback_redis_batch_key_reservations(self, previous_access_times: dict[str, float | None]) -> None:
with self._last_redis_batch_access_time_lock:
for key, previous_time in previous_access_times.items():
@ -502,12 +519,29 @@ class DualCache(BaseCache):
verbose_logger, logging.ERROR, "LiteLLM Cache: exception in async add_cache", e, with_traceback=True
)
async def async_set_cache_pre_call(self, key: str, value: object, ttl: float | None) -> BatchResult[None] | None:
"""Memory now, the Redis SET on the request's pipeline, sent with the next read any caller awaits; None
when no pipeline is open, so the caller takes its direct path."""
batch: Final = None if self.redis_cache is None else active_request_redis_batch(self.redis_cache)
return None if batch is None else await self._set_on_batch(batch, key, value, ttl)
async def async_set_cache_post_call(self, key: str, value: object, ttl: float | None) -> BatchResult[None] | None:
"""Memory now, the Redis SET on the request's post-call pipeline; None when no pipeline is open, so the
caller takes its direct path."""
batch: Final = None if self.redis_cache is None else active_post_call_redis_batch(self.redis_cache)
return None if batch is None else await self._set_on_batch(batch, key, value, ttl)
async def async_delete_cache_pre_call(self, key: str) -> BatchResult[None] | None:
"""Memory now, the Redis DEL on the request's pipeline; None when no pipeline is open, so the caller
takes its direct path."""
batch: Final = None if self.redis_cache is None else active_request_redis_batch(self.redis_cache)
if batch is None:
return None
if self.in_memory_cache is not None:
self.in_memory_cache.delete_cache(key)
return batch.delete(key)
async def _set_on_batch(self, batch: RedisBatch, key: str, value: object, ttl: float | None) -> BatchResult[None]:
effective_ttl: Final = self.default_in_memory_ttl if ttl is None else ttl
if self.in_memory_cache is not None:
await self.in_memory_cache.async_set_cache(key, value, ttl=effective_ttl)

View file

@ -47,6 +47,7 @@ class _RedisPipeline(Protocol):
def incrbyfloat(self, name: str, amount: float) -> object: ...
def expire(self, name: str, time: timedelta) -> object: ...
def set(self, name: str, value: str, ex: timedelta | None = None) -> object: ...
def delete(self, *names: str) -> object: ...
async def execute(self, raise_on_error: bool = True) -> list[object]: ...
@ -233,6 +234,27 @@ class _Set(_Op[None]):
await self._redis_cache.async_set_cache_pipeline_with_ttls(((self._key, self._value, self._ttl),))
class _Delete(_Op[None]):
"""DEL of one key, the pipelined twin of ``async_delete_cache``."""
__slots__ = ("_key", "_redis_cache")
def __init__(self, redis_cache: RedisCache, key: str) -> None:
super().__init__()
self._redis_cache: Final = redis_cache
self._key: Final = key
def enqueue(self, pipe: _RedisPipeline) -> int:
pipe.delete(self._redis_cache.check_and_fix_namespace(key=self._key))
return 1
def resolve(self, replies: Sequence[object]) -> None:
return None
async def run_alone(self) -> None:
await self._redis_cache.async_delete_cache(self._key)
class BatchResult(Generic[_T]):
"""Awaitable handle for one declared operation; awaiting it flushes the batch it belongs to."""
@ -269,10 +291,23 @@ class RedisBatch:
_pending: list[_Op[object]] = field(default_factory=list) # mutable-ok: drained by flush
_flush_hooks: list[Callable[[], None]] = field(default_factory=list) # mutable-ok: append-only registry
_lock: asyncio.Lock = field(default_factory=asyncio.Lock)
_misses: set[str] = field(default_factory=set) # mutable-ok: keys an MGET of this request read as absent
flushes: int = 0
def mget(self, keys: Sequence[str]) -> BatchResult[Mapping[str, object]]:
return self._declare(_MGet(self.redis_cache, keys))
op: Final = _MGet(self.redis_cache, keys)
op.future.add_done_callback(self._note_misses)
return self._declare(op)
def _note_misses(self, future: asyncio.Future[Mapping[str, object]]) -> None:
if future.cancelled() or future.exception() is not None:
return
self._misses.update(key for key, value in future.result().items() if value is None)
def read_as_missing(self, key: str) -> bool:
"""True when an MGET on this batch already found no value under ``key`` and nothing has set it since,
so a per-key GET later in the same request can be answered without another round trip."""
return key in self._misses
def script(
self, source: str, run: RegisteredScript, keys: Sequence[str], args: Sequence[_ScriptArg]
@ -283,8 +318,13 @@ class RedisBatch:
return self._declare(_Increment(self.redis_cache, key, value, ttl))
def set(self, key: str, value: object, ttl: float | None = None) -> BatchResult[None]:
self._misses.discard(key)
return self._declare(_Set(self.redis_cache, key, value, ttl))
def delete(self, key: str) -> BatchResult[None]:
self._misses.add(key)
return self._declare(_Delete(self.redis_cache, key))
def add_flush_hook(self, hook: Callable[[], None]) -> None:
"""Called at the start of every flush so lazily bound readers can declare their keys into the same trip."""
self._flush_hooks.append(hook)
@ -449,10 +489,6 @@ class RequestRedisBatches:
def batches(self) -> tuple[RedisBatch, ...]:
return tuple(self._batches.values())
@property
def post_call_batches(self) -> tuple[RedisBatch, ...]:
return tuple(self._post_call.values())
_active_request_batches: Final[ContextVar[RequestRedisBatches | None]] = ContextVar(
"request_redis_batches", default=None

View file

@ -1908,6 +1908,7 @@ SPEND_LOG_KEY_METADATA_CACHE_TTL: Final = 600
SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL: Final = 30
SPEND_LOG_KEY_METADATA_CACHE_MAX_ITEMS: Final = 10000
SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS: Final = 5000
SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE: Final = 100
# Short TTL for negative MCP access-group existence lookups. Keeps unauthenticated
# callers from forcing a DB query per request for unknown names, while bounding
# staleness so a transient DB error (which surfaces as an empty list) cannot

View file

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

View file

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

View file

@ -31,6 +31,7 @@ from litellm.llms.base_llm.base_utils import BaseLLMModelInfo, BaseTokenCounter
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.types.llms.anthropic import (
ANTHROPIC_HOSTED_TOOLS,
ANTHROPIC_MID_CONVERSATION_OUTPUT_CONFIG_BETA_HEADER,
ANTHROPIC_OAUTH_BETA_HEADER,
ANTHROPIC_OAUTH_TOKEN_PREFIX,
AllAnthropicToolsValues,
@ -326,6 +327,12 @@ class AnthropicModelInfo(BaseLLMModelInfo):
file_ids: Final = get_file_ids_from_messages(messages)
return len(file_ids) > 0
def is_mid_conversation_output_config_used(self, messages: list[AllMessageValues]) -> bool:
"""
Return if "output_config" is in a message
"""
return any("output_config" in message for message in messages)
def is_mcp_server_used(self, mcp_servers: list[AnthropicMcpServerTool] | None) -> bool:
if mcp_servers is None:
return False
@ -851,6 +858,7 @@ class AnthropicModelInfo(BaseLLMModelInfo):
mcp_server_used: bool = False,
*,
custom_llm_provider: str,
is_mid_conversation_output_config_used: bool = False,
) -> list[str]:
"""
Get list of common beta headers based on the features that are active.
@ -883,6 +891,9 @@ class AnthropicModelInfo(BaseLLMModelInfo):
if mcp_server_used:
betas.append("mcp-client-2025-04-04")
if is_mid_conversation_output_config_used:
betas.append(ANTHROPIC_MID_CONVERSATION_OUTPUT_CONFIG_BETA_HEADER)
return list(set(betas))
@staticmethod
@ -915,6 +926,7 @@ class AnthropicModelInfo(BaseLLMModelInfo):
container_with_skills_used: bool = False,
api_base: str | None = None,
use_bearer_for_custom_base: bool = False,
is_mid_conversation_output_config_used: bool = False,
) -> dict:
betas: Final = set()
# Anthropic no longer requires the prompt-caching beta header
@ -950,6 +962,9 @@ class AnthropicModelInfo(BaseLLMModelInfo):
if container_with_skills_used:
betas.add("skills-2025-10-02")
if is_mid_conversation_output_config_used:
betas.add(ANTHROPIC_MID_CONVERSATION_OUTPUT_CONFIG_BETA_HEADER)
_is_oauth: Final = api_key and api_key.startswith(ANTHROPIC_OAUTH_TOKEN_PREFIX)
headers: Final = {
"anthropic-version": anthropic_version or "2023-06-01",
@ -1015,6 +1030,7 @@ class AnthropicModelInfo(BaseLLMModelInfo):
mcp_server_used: Final = self.is_mcp_server_used(mcp_servers=optional_params.get("mcp_servers"))
pdf_used: Final = self.is_pdf_used(messages=messages)
file_id_used: Final = self.is_file_id_used(messages=messages)
is_mid_conversation_output_config_used: Final = self.is_mid_conversation_output_config_used(messages=messages)
web_search_tool_used: Final = self.is_web_search_tool_used(tools=tools)
tool_search_used: Final = self.is_tool_search_used(tools=tools)
programmatic_tool_calling_used: Final = self.is_programmatic_tool_calling_used(tools=tools)
@ -1032,6 +1048,7 @@ class AnthropicModelInfo(BaseLLMModelInfo):
api_key=api_key,
auth_token=auth_token,
file_id_used=file_id_used,
is_mid_conversation_output_config_used=is_mid_conversation_output_config_used,
web_search_tool_used=web_search_tool_used,
is_vertex_request=optional_params.get("is_vertex_request", False),
user_anthropic_beta_headers=user_anthropic_beta_headers,

View file

@ -255,6 +255,7 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
tool_search_used: Final = self.is_tool_search_used(tools)
programmatic_tool_calling_used: Final = self.is_programmatic_tool_calling_used(tools)
input_examples_used: Final = self.is_input_examples_used(tools)
is_mid_conversation_output_config_used: Final = self.is_mid_conversation_output_config_used(messages)
user_beta_set: Final = set(get_anthropic_beta_from_headers(headers))
beta_set: Final = set(user_beta_set)
@ -266,6 +267,7 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
file_id_used=self.is_file_id_used(messages),
mcp_server_used=self.is_mcp_server_used(optional_params.get("mcp_servers")),
custom_llm_provider="bedrock",
is_mid_conversation_output_config_used=is_mid_conversation_output_config_used,
)
beta_set.update(auto_betas)

View file

@ -515,7 +515,13 @@ class AmazonAnthropicClaudeMessagesConfig(
tool_search_used: Final = anthropic_model_info.is_tool_search_used(tools)
programmatic_tool_calling_used: Final = anthropic_model_info.is_programmatic_tool_calling_used(tools)
input_examples_used: Final = anthropic_model_info.is_input_examples_used(tools)
outgoing_messages_typed: Final = cast(
list[AllMessageValues],
anthropic_messages_request["messages"],
)
is_mid_conversation_output_config_used: Final = anthropic_model_info.is_mid_conversation_output_config_used(
outgoing_messages_typed
)
user_beta_set: Final = set(get_anthropic_beta_from_headers(headers))
beta_set: Final = set(user_beta_set)
auto_betas: Final = anthropic_model_info.get_anthropic_beta_list(
@ -528,6 +534,7 @@ class AmazonAnthropicClaudeMessagesConfig(
anthropic_messages_optional_request_params.get("mcp_servers")
),
custom_llm_provider="bedrock",
is_mid_conversation_output_config_used=is_mid_conversation_output_config_used,
)
beta_set.update(auto_betas)

View file

@ -3358,7 +3358,7 @@
"supports_function_calling": true
},
"azure_ai/claude-haiku-4-5": {
"deprecation_date": "2026-10-19",
"deprecation_date": "2026-11-15",
"cache_creation_input_token_cost": 1.25e-06,
"cache_creation_input_token_cost_above_1hr": 2e-06,
"cache_read_input_token_cost": 1e-07,
@ -3378,10 +3378,11 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"prompt_cache_min_tokens": 4096
"prompt_cache_min_tokens": 4096,
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/model-retirement-schedule"
},
"azure_ai/claude-opus-4-5": {
"deprecation_date": "2026-10-19",
"deprecation_date": "2026-11-24",
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -3402,7 +3403,8 @@
"supports_tool_choice": true,
"supports_vision": true,
"supports_output_config": true,
"prompt_cache_min_tokens": 4096
"prompt_cache_min_tokens": 4096,
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/model-retirement-schedule"
},
"azure_ai/claude-opus-4-6": {
"deprecation_date": "2027-02-02",
@ -3640,7 +3642,7 @@
"prompt_cache_min_tokens": 1024
},
"azure_ai/claude-sonnet-4-5": {
"deprecation_date": "2026-10-19",
"deprecation_date": "2026-11-15",
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
@ -3660,7 +3662,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"prompt_cache_min_tokens": 1024
"prompt_cache_min_tokens": 1024,
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/model-retirement-schedule"
},
"azure_ai/claude-sonnet-5": {
"deprecation_date": "2027-06-30",
@ -30721,6 +30724,7 @@
"output_cost_per_image": 0.08
},
"gemini/veo-3.1-fast-generate-preview": {
"deprecation_date": "2026-10-22",
"litellm_provider": "gemini",
"max_input_tokens": 1024,
"max_tokens": 1024,
@ -30737,6 +30741,7 @@
]
},
"gemini/veo-3.1-generate-preview": {
"deprecation_date": "2026-10-22",
"litellm_provider": "gemini",
"max_input_tokens": 1024,
"max_tokens": 1024,
@ -30752,6 +30757,7 @@
]
},
"gemini/veo-3.1-lite-generate-preview": {
"deprecation_date": "2026-10-22",
"litellm_provider": "gemini",
"max_input_tokens": 1024,
"max_tokens": 1024,
@ -32912,10 +32918,13 @@
"gpt-image-2.5-flare": {
"cache_read_input_image_token_cost": 2e-06,
"cache_read_input_token_cost": 1.25e-06,
"cache_read_input_token_cost_batches": 6.25e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "openai",
"mode": "image_generation",
"input_cost_per_image_token": 8e-06,
"input_cost_per_image_token_batches": 4e-06,
"input_cost_per_token_batches": 2.5e-06,
"output_cost_per_image_token": 3e-05,
"supported_endpoints": [
"/v1/images/generations",
@ -32944,10 +32953,13 @@
"gpt-image-2.5-sunburst": {
"cache_read_input_image_token_cost": 2e-06,
"cache_read_input_token_cost": 1.25e-06,
"cache_read_input_token_cost_batches": 6.25e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "openai",
"mode": "image_generation",
"input_cost_per_image_token": 8e-06,
"input_cost_per_image_token_batches": 4e-06,
"input_cost_per_token_batches": 2.5e-06,
"output_cost_per_image_token": 3e-05,
"supported_endpoints": [
"/v1/images/generations",
@ -38611,6 +38623,7 @@
},
"mistral/zai-glm-5-2": {
"cache_read_input_token_cost": 1.4e-07,
"deprecation_date": "2026-10-31",
"input_cost_per_token": 1.4e-06,
"litellm_provider": "mistral",
"max_input_tokens": 1048576,
@ -38741,6 +38754,7 @@
"source": "https://mistral.ai/pricing#api-pricing"
},
"mistral/mistral-ocr-4-0": {
"deprecation_date": "2026-09-30",
"litellm_provider": "mistral",
"ocr_cost_per_page": 0.004,
"ocr_cost_per_page_batches": 0.002,
@ -60616,6 +60630,7 @@
"supports_vision": true
},
"mistral/labs-leanstral-1-5": {
"deprecation_date": "2026-09-30",
"input_cost_per_token": 0.0,
"litellm_provider": "mistral",
"max_input_tokens": 262144,
@ -61337,13 +61352,16 @@
},
"fireworks_ai/nemotron-lightning-3p5-30b-a3b": {
"cache_read_input_token_cost": 1e-08,
"cache_read_input_token_cost_priority": 1.25e-08,
"input_cost_per_token": 5e-08,
"input_cost_per_token_priority": 6.25e-08,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 262144,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 2e-07,
"output_cost_per_token_priority": 2.5e-07,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -61353,13 +61371,16 @@
},
"fireworks_ai/nemotron-3-ultra-nvfp4": {
"cache_read_input_token_cost": 1.2e-07,
"cache_read_input_token_cost_priority": 1.5e-07,
"input_cost_per_token": 6e-07,
"input_cost_per_token_priority": 7.5e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 262144,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 2.4e-06,
"output_cost_per_token_priority": 3e-06,
"source": "https://api.fireworks.ai/v1/serverless/models",
"supports_function_calling": true,
"supports_reasoning": true,
@ -61389,13 +61410,16 @@
},
"fireworks_ai/accounts/fireworks/models/nemotron-lightning-3p5-30b-a3b": {
"cache_read_input_token_cost": 1e-08,
"cache_read_input_token_cost_priority": 1.25e-08,
"input_cost_per_token": 5e-08,
"input_cost_per_token_priority": 6.25e-08,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 262144,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 2e-07,
"output_cost_per_token_priority": 2.5e-07,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -61405,13 +61429,16 @@
},
"fireworks_ai/accounts/fireworks/models/nemotron-3-ultra-nvfp4": {
"cache_read_input_token_cost": 1.2e-07,
"cache_read_input_token_cost_priority": 1.5e-07,
"input_cost_per_token": 6e-07,
"input_cost_per_token_priority": 7.5e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 262144,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 2.4e-06,
"output_cost_per_token_priority": 3e-06,
"source": "https://api.fireworks.ai/v1/serverless/models",
"supports_function_calling": true,
"supports_reasoning": true,
@ -64321,13 +64348,16 @@
},
"fireworks_ai/accounts/fireworks/routers/glm-5p3-us": {
"cache_read_input_token_cost": 3.9e-07,
"cache_read_input_token_cost_priority": 4.875e-07,
"input_cost_per_token": 2.1e-06,
"input_cost_per_token_priority": 2.625e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 6.6e-06,
"output_cost_per_token_priority": 8.25e-06,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -64356,13 +64386,16 @@
},
"fireworks_ai/glm-5p3-us": {
"cache_read_input_token_cost": 3.9e-07,
"cache_read_input_token_cost_priority": 4.875e-07,
"input_cost_per_token": 2.1e-06,
"input_cost_per_token_priority": 2.625e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 6.6e-06,
"output_cost_per_token_priority": 8.25e-06,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -64446,12 +64479,15 @@
},
"fireworks_ai/accounts/fireworks/routers/glm-5p3-flash-us": {
"cache_read_input_token_cost": 4.5e-08,
"cache_read_input_token_cost_priority": 5.625e-08,
"input_cost_per_token": 2.25e-07,
"input_cost_per_token_priority": 2.8125e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_tokens": 1048576,
"mode": "chat",
"output_cost_per_token": 7.5e-07,
"output_cost_per_token_priority": 9.375e-07,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_response_schema": true,
@ -64477,12 +64513,15 @@
},
"fireworks_ai/glm-5p3-flash-us": {
"cache_read_input_token_cost": 4.5e-08,
"cache_read_input_token_cost_priority": 5.625e-08,
"input_cost_per_token": 2.25e-07,
"input_cost_per_token_priority": 2.8125e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_tokens": 1048576,
"mode": "chat",
"output_cost_per_token": 7.5e-07,
"output_cost_per_token_priority": 9.375e-07,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_response_schema": true,
@ -64571,6 +64610,7 @@
"source": "https://api.together.ai/v1/models"
},
"together_ai/nvidia/NVIDIA-Nemotron-Nano-9B-v2": {
"deprecation_date": "2026-02-25",
"input_cost_per_token": 6e-08,
"output_cost_per_token": 2.5e-07,
"litellm_provider": "together_ai",
@ -70451,6 +70491,7 @@
},
"together_ai/nvidia/nemotron-3-ultra-550b-a55b": {
"cache_read_input_token_cost": 2e-07,
"deprecation_date": "2026-08-27",
"input_cost_per_token": 6e-07,
"litellm_provider": "together_ai",
"max_input_tokens": 512288,
@ -77476,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,
@ -79218,12 +79262,12 @@
"input_cost_per_token_above_272k_tokens": 4e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_above_272k_tokens": 1.5e-05,
"source": "https://developers.openai.com/api/docs/pricing",
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-1-sol.html",
"supported_endpoints": [
"/v1/responses"
],
@ -79253,12 +79297,12 @@
"input_cost_per_token_above_272k_tokens": 4e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_above_272k_tokens": 1.5e-05,
"source": "https://developers.openai.com/api/docs/pricing",
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-1-sol.html",
"supported_modalities": [
"text",
"image"
@ -79285,12 +79329,12 @@
"input_cost_per_token_above_272k_tokens": 4.4e-06,
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "responses",
"output_cost_per_token": 1.1e-05,
"output_cost_per_token_above_272k_tokens": 1.65e-05,
"source": "https://developers.openai.com/api/docs/models/gpt-6.1-sol",
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-1-sol.html",
"supported_endpoints": [
"/v1/chat/completions",
"/v1/responses"
@ -79323,12 +79367,12 @@
"input_cost_per_token_above_272k_tokens": 4.4e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1.1e-05,
"output_cost_per_token_above_272k_tokens": 1.65e-05,
"source": "https://developers.openai.com/api/docs/pricing",
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-1-sol.html",
"supported_endpoints": [
"/v1/responses"
],
@ -79348,5 +79392,33 @@
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true
},
"vertex_ai/gemini-3.8-flash-tts": {
"input_cost_per_token": 5e-07,
"litellm_provider": "vertex_ai",
"max_input_tokens": 8192,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "audio_speech",
"output_cost_per_audio_token": 9e-06,
"output_cost_per_token": 9e-06,
"source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing",
"supported_endpoints": [
"/v1/audio/speech"
]
},
"vertex_ai/gemini-3.8-flash-lite-tts": {
"input_cost_per_token": 5e-07,
"litellm_provider": "vertex_ai",
"max_input_tokens": 8192,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "audio_speech",
"output_cost_per_audio_token": 6e-06,
"output_cost_per_token": 6e-06,
"source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing",
"supported_endpoints": [
"/v1/audio/speech"
]
}
}

View file

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

View file

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

View file

@ -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
@ -3494,7 +3495,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()
@ -3543,9 +3546,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
@ -3577,7 +3585,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)
)
@ -3646,12 +3654,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.
@ -3661,6 +3671,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
@ -3675,12 +3689,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:
@ -3694,6 +3712,8 @@ class MCPServerManager:
)
return tool_permissions
except Exception as e:
if requires_fresh_policy:
raise
verbose_logger.warning("Failed to resolve toolset permissions: %s", e)
return {}

View file

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

View file

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

View file

@ -2521,6 +2521,91 @@
"title": "AgentExtension",
"type": "object"
},
"AgentIdentityBinding": {
"properties": {
"active": {
"default": true,
"title": "Active",
"type": "boolean"
},
"agent_id": {
"title": "Agent Id",
"type": "string"
},
"client_id": {
"title": "Client Id",
"type": "string"
},
"issuer": {
"title": "Issuer",
"type": "string"
},
"last_authenticated_at": {
"anyOf": [
{
"format": "date-time",
"type": "string"
},
{
"type": "null"
}
],
"title": "Last Authenticated At"
},
"provider": {
"const": "microsoft_entra",
"title": "Provider",
"type": "string"
},
"required_roles": {
"default": [],
"items": {
"type": "string"
},
"title": "Required Roles",
"type": "array"
},
"required_scopes": {
"default": [
"user_impersonation"
],
"items": {
"type": "string"
},
"title": "Required Scopes",
"type": "array"
},
"revision": {
"title": "Revision",
"type": "string"
},
"service_principal_id": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "Service Principal Id"
},
"tenant_id": {
"title": "Tenant Id",
"type": "string"
}
},
"required": [
"agent_id",
"provider",
"tenant_id",
"client_id",
"issuer",
"revision"
],
"title": "AgentIdentityBinding",
"type": "object"
},
"AgentInterface": {
"description": "Declares a combination of a target URL and a transport protocol.",
"properties": {
@ -2972,6 +3057,21 @@
],
"title": "Created By"
},
"enabled": {
"default": true,
"title": "Enabled",
"type": "boolean"
},
"execution_mode": {
"default": "autonomous",
"enum": [
"autonomous",
"delegated",
"both"
],
"title": "Execution Mode",
"type": "string"
},
"extra_headers": {
"anyOf": [
{
@ -2986,6 +3086,26 @@
],
"title": "Extra Headers"
},
"identity": {
"anyOf": [
{
"$ref": "#/components/schemas/AgentIdentityBinding"
},
{
"type": "null"
}
]
},
"identity_managed": {
"default": false,
"title": "Identity Managed",
"type": "boolean"
},
"jwt_auth_configured": {
"default": false,
"title": "Jwt Auth Configured",
"type": "boolean"
},
"keys": {
"anyOf": [
{

View file

@ -15,6 +15,7 @@ from pydantic import (
Json,
JsonValue,
PositiveInt,
PrivateAttr,
field_validator,
model_validator,
)
@ -27,7 +28,7 @@ from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
validate_langfuse_span_scope_value,
validate_no_callback_env_reference,
)
from litellm.types.agents import AgentCaller
from litellm.types.agents import AgentCaller, AgentResponse
from litellm.types.integrations.compression_interception import (
CompressionSavingsMetadata,
)
@ -46,6 +47,7 @@ from litellm.types.mcp import (
MCPTransportType,
)
from litellm.types.mcp_server.mcp_server_manager import MCPInfo
from litellm.types.proxy.agent_identity import ManagedAgentContext
from litellm.types.proxy.carried_budget_state import (
OrgBudgetSnapshot,
TeamBudgetSnapshot,
@ -567,6 +569,7 @@ class LiteLLMRoutes(enum.Enum):
"/agents",
"/a2a/{agent_id}",
"/a2a/{agent_id}/message/send",
"/v1/a2a/{agent_id}/message/send",
"/a2a/{agent_id}/message/stream",
"/a2a/{agent_id}/.well-known/agent-card.json",
)
@ -3302,6 +3305,8 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
# metadata or JWT claims, so it cannot be forged to gain the team-inherited MCP grant union
# or to escape the caller-Authorization egress scrub. exclude=True keeps it out of serialization.
mcp_admitted_user_subject: bool = Field(default=False, exclude=True)
requires_fresh_policy: bool = Field(default=False, exclude=True)
mcp_explicit_grants_only: bool = Field(default=False, exclude=True)
# team_id -> that team's mcp_rpm_limit map, for a keyless admitted subject that reaches MCP
# servers through several teams at once and therefore has no single team_id for the limiter to
# key off. Server-only and stripped from validated input for the same reason as the marker
@ -3326,6 +3331,13 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
"user id."
),
)
invoked_agent_id: str | None = Field(default=None, exclude=True)
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(
default=None,
exclude=True,
@ -3363,11 +3375,19 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
# path via post-construction assignment. Strip it from any validated input (constructor
# kwargs, model_validate, a JWT/key claim splat) so it can never be forged from caller data.
values.pop("mcp_admitted_user_subject", None)
values.pop("requires_fresh_policy", None)
values.pop("mcp_explicit_grants_only", None)
values.pop("mcp_source_team_rpm_limits", None)
values.pop("mcp_session_resource_server_id", None)
values.pop("mcp_toolset_id", None)
values.pop("via_virtual_key", None)
values.pop("agent_caller", None)
values.pop("managed_agent_context", None)
values.pop("managed_agent_policy", None)
values.pop("invoked_agent_id", None)
values.pop("invoked_agent_policy", None)
values.pop("agent_invocation_cost", None)
values.pop("billing_agent_policy", None)
if values.get("api_key") is not None:
values.update({"token": cls._safe_hash_litellm_api_key(values.get("api_key"))})
if isinstance(values.get("api_key"), str):
@ -4063,6 +4083,11 @@ class SpendLogsRouterMetadata(TypedDict):
class SpendLogsMetadata(TypedDict):
actor_agent_id: ReadOnly[NotRequired[str | None]]
target_agent_id: ReadOnly[NotRequired[str | None]]
billing_agent_id: ReadOnly[NotRequired[str | None]]
agent_execution_mode: ReadOnly[NotRequired[str | None]]
verified_human_user_id: ReadOnly[NotRequired[str | None]]
autorouter_baseline_observation: ReadOnly[str | None]
"""
Specific metadata k,v pairs logged to spendlogs for easier cost tracking
@ -4126,6 +4151,7 @@ class SpendLogsPayload(TypedDict):
model_id: str | None
model_group: str | None
mcp_namespaced_tool_name: str | None
billing_agent_id: ReadOnly[NotRequired[str | None]]
agent_id: str | None
api_base: str
user: str
@ -5048,6 +5074,7 @@ class JWTAuthBuilderResult(TypedDict):
org_id: str | None
team_membership: LiteLLM_TeamMembership | None
jwt_claims: dict # Decoded JWT token claims (avoids re-decoding)
managed_agent_context: ReadOnly[NotRequired[ManagedAgentContext | None]]
agent_id: ReadOnly[str | None]

View file

@ -597,7 +597,6 @@ async def get_agent_card(
if agent is None:
raise HTTPException(status_code=404, detail=f"Agent '{agent_id}' not found")
# Check agent permission (skip for admin users)
is_allowed: Final = await AgentRequestHandler.is_agent_allowed(
agent_id=agent.agent_id,
user_api_key_auth=user_api_key_dict,

View file

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

View file

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

View file

@ -8,8 +8,11 @@ Follows the same pattern as MCP permission handling.
import asyncio
from collections.abc import Awaitable, Callable, Sequence
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final, TypeAlias
from fastapi import HTTPException
from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.ui_session_utils import build_effective_auth_contexts
from litellm.proxy._types import (
@ -24,6 +27,7 @@ from litellm.proxy.agent_endpoints.auth.agent_access_groups import (
resolve_agent_access_group_ceiling,
)
from litellm.proxy.agent_endpoints.auth.agent_caller import agent_caller_auth
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_policy
from litellm.repositories.table_repositories import AgentsRepository
from litellm.types.agents import AgentResponse
@ -83,13 +87,23 @@ class AgentRequestHandler:
async def resolve_agent_access(
user_api_key_auth: UserAPIKeyAuth | None = None,
resolve_ceiling: CeilingResolver = resolve_agent_access_group_ceiling,
*,
strict: bool = False,
) -> AgentAccess:
"""Agents the key may reach: key and team grants, intersected with the agent's access group ceiling
and, for an agent key acting on behalf of an invoking user, with that user's team grants."""
key_team_access: Final = await AgentRequestHandler._resolve_key_team_agent_access(user_api_key_auth)
caller_access: Final = await AgentRequestHandler._agent_caller_access(user_api_key_auth)
if managed_agent_policy(user_api_key_auth) is not None:
return await _managed_actor_agent_access(user_api_key_auth)
key_team_access: Final = await AgentRequestHandler.resolve_key_team_agent_access(
user_api_key_auth, strict=strict
)
if strict and isinstance(key_team_access, UnrestrictedAgentAccess):
return RestrictedAgentAccess(frozenset())
caller_access: Final = await AgentRequestHandler.agent_caller_access(user_api_key_auth, strict=strict)
own_access: Final = _intersect_agent_access(key_team_access, caller_access)
agent_ceiling: Final = await AgentRequestHandler._agent_access_group_ceiling(user_api_key_auth, resolve_ceiling)
agent_ceiling: Final = await AgentRequestHandler._agent_access_group_ceiling(
user_api_key_auth, resolve_ceiling, strict=strict
)
if agent_ceiling is None:
return own_access
if isinstance(own_access, UnrestrictedAgentAccess):
@ -97,20 +111,26 @@ class AgentRequestHandler:
return RestrictedAgentAccess(own_access.agent_ids & agent_ceiling)
@staticmethod
async def _agent_caller_access(user_api_key_auth: UserAPIKeyAuth | None) -> AgentAccess:
async def agent_caller_access(user_api_key_auth: UserAPIKeyAuth | None, *, strict: bool = False) -> AgentAccess:
caller_auth: Final = agent_caller_auth(user_api_key_auth) if user_api_key_auth else None
if caller_auth is None:
return UnrestrictedAgentAccess()
return await AgentRequestHandler._get_allowed_agents_for_team(caller_auth)
return await AgentRequestHandler._get_allowed_agents_for_team(caller_auth, strict=strict)
@staticmethod
async def _resolve_key_team_agent_access(
async def resolve_key_team_agent_access(
user_api_key_auth: UserAPIKeyAuth | None,
*,
strict: bool = False,
) -> AgentAccess:
try:
key_access: Final = await AgentRequestHandler._get_allowed_agents_for_key(user_api_key_auth)
team_access: Final = await AgentRequestHandler._get_allowed_agents_for_team(user_api_key_auth)
key_access: Final = await AgentRequestHandler.get_allowed_agents_for_key(user_api_key_auth, strict=strict)
team_access: Final = await AgentRequestHandler._get_allowed_agents_for_team(
user_api_key_auth, strict=strict
)
except Exception as e:
if strict:
raise HTTPException(503, "Agent invocation policy is unavailable") from e
verbose_logger.warning("Failed to get allowed agents: %s", e)
return UnrestrictedAgentAccess()
return _intersect_agent_access(key_access, team_access)
@ -119,10 +139,16 @@ class AgentRequestHandler:
async def _agent_access_group_ceiling(
user_api_key_auth: UserAPIKeyAuth | None,
resolve_ceiling: CeilingResolver,
*,
strict: bool = False,
) -> frozenset[str] | None:
if user_api_key_auth is None or not user_api_key_auth.agent_id:
return None
ceiling: Final = await resolve_ceiling(user_api_key_auth.agent_id)
ceiling: Final = (
await resolve_agent_access_group_ceiling(user_api_key_auth.agent_id, check_db_only=True)
if strict
else await resolve_ceiling(user_api_key_auth.agent_id)
)
if ceiling is None:
return None
return _to_stable_ids(ceiling.agent_ids)
@ -144,6 +170,49 @@ class AgentRequestHandler:
bool: True if agent is allowed, False otherwise
"""
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
from litellm.proxy.proxy_server import prisma_client
from litellm.types.proxy.agent_identity import AgentIdentityFailure
registered: Final = global_agent_registry.get_agent_by_id(agent_id)
registry_managed: Final = isinstance(registered, AgentResponse) and registered.identity_managed
if registry_managed or (registered is None and prisma_client is not None):
target: Final = await AgentIdentityStore.from_client(prisma_client).agent(agent_id)
if isinstance(target, AgentIdentityFailure):
if registry_managed:
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
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,60 @@ 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_grants(user_id: str | None, team_id: str | None = None) -> frozenset[str]:
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
if user_id is None:
return frozenset()
human: Final = await MCPRequestHandler.reload_admitted_user(user_id, requires_fresh_policy=True)
sources: Final = await MCPRequestHandler.admitted_subject_sources(
human, allowed_team_ids=frozenset((team_id,)) if team_id else frozenset()
)
human_access: Final = await asyncio.gather(*(_strict_agent_access(source) for source in sources))
return frozenset().union(*(_granted_ids(access) for access in human_access))

View file

@ -0,0 +1,84 @@
from typing import Final
from litellm.proxy._types import 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
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

View file

@ -0,0 +1,17 @@
from collections.abc import Mapping
from typing import Final
from fastapi import HTTPException
LEGACY_IDENTITY_MESSAGE: Final = (
"litellm_params.identity is not supported: bind an Entra application through the top-level identity field"
)
def has_legacy_identity(params: Mapping[str, object] | None) -> bool:
return params is not None and "identity" in params
def reject_legacy_identity(params: Mapping[str, object] | None) -> None:
if has_legacy_identity(params):
raise HTTPException(400, LEGACY_IDENTITY_MESSAGE)

View file

@ -0,0 +1,252 @@
import json
from collections.abc import Mapping
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Final
from litellm.proxy.agent_endpoints.managed_identity import classify_agent_subject
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, get_management_object_ttl
from litellm.repositories.table_repositories import (
AgentIdentityRepository,
AgentsRepository,
RetiredAgentIdentityRepository,
RetiredAgentRepository,
VerifiedSubjectRepository,
)
from litellm.types.agents import AgentResponse
from litellm.types.proxy.agent_identity import (
AgentIdentityFailure,
ManagedAgentContext,
MicrosoftInteractiveSubject,
VerifiedHumanSubject,
)
if TYPE_CHECKING:
from prisma.models import LiteLLM_VerifiedSubject
from prisma.types import (
LiteLLM_AgentIdentityUpdateManyMutationInput,
LiteLLM_AgentIdentityWhereInput,
LiteLLM_AgentIdentityWhereUniqueInput,
LiteLLM_AgentsTableInclude,
LiteLLM_AgentsTableWhereUniqueInput,
LiteLLM_RetiredAgentWhereUniqueInput,
LiteLLM_VerifiedSubjectCreateInput,
LiteLLM_VerifiedSubjectUpsertInput,
LiteLLM_VerifiedSubjectWhereUniqueInput,
)
class AgentIdentityStore:
@classmethod
def from_client(cls, client: object, *, cache: UserApiKeyCache | None = None) -> "AgentIdentityStore":
return cls(
AgentsRepository(client, use_writer=True),
AgentIdentityRepository(client, use_writer=True),
VerifiedSubjectRepository(client, use_writer=True),
RetiredAgentIdentityRepository(client, use_writer=True),
RetiredAgentRepository(client, use_writer=True),
cache=cache,
)
def __init__(
self,
agents: AgentsRepository,
identities: AgentIdentityRepository,
humans: VerifiedSubjectRepository,
retired: RetiredAgentIdentityRepository | None = None,
retired_agents: RetiredAgentRepository | None = None,
*,
cache: UserApiKeyCache | None = None,
) -> None:
self.agents = agents
self.identities = identities
self.humans = humans
self.retired = retired
self.retired_agents = retired_agents
self.cache = cache
async def agent(self, agent_id: str) -> AgentResponse | AgentIdentityFailure | None:
try:
where: Final[LiteLLM_AgentsTableWhereUniqueInput] = {"agent_id": agent_id}
include: Final[LiteLLM_AgentsTableInclude] = {
"identity": True,
"object_permission": True,
}
row: Final = await self.agents.table.find_unique(where=where, include=include)
if row is None:
return None
return AgentResponse.model_validate(row.model_dump())
except Exception:
return AgentIdentityFailure(code="policy_unavailable", message="Agent policy could not be loaded")
async def unbound_client(self, where: "LiteLLM_AgentIdentityWhereUniqueInput") -> AgentIdentityFailure | None:
if self.retired is not None:
try:
retired: Final = await self.retired.table.find_unique(where=where)
except Exception:
return AgentIdentityFailure(
code="policy_unavailable", message="Retired agent identity could not be checked"
)
if retired is not None:
return AgentIdentityFailure(message="This agent identity binding has been retired")
return None
async def _bound_agent_id(self, tenant_id: str, client_id: str) -> str | AgentIdentityFailure | None:
cache_key: Final = f"agent_identity:{json.dumps((tenant_id, client_id))}"
cached: Final[object] = await self.cache.async_get_cache(key=cache_key) if self.cache is not None else None
if isinstance(cached, str):
return cached
where: Final[LiteLLM_AgentIdentityWhereUniqueInput] = {
"provider_tenant_id_client_id": {
"provider": "microsoft_entra",
"tenant_id": tenant_id,
"client_id": client_id,
}
}
try:
row: Final = await self.identities.table.find_unique(where=where)
except Exception:
return AgentIdentityFailure(code="policy_unavailable", message="Agent identity could not be loaded")
if row is None:
return await self.unbound_client(where)
if self.cache is not None:
await self.cache.async_set_cache(
key=cache_key, value=row.agent_id, ttl=get_management_object_ttl(self.cache)
)
return row.agent_id
async def resolve_verified_claims(
self, claims: Mapping[str, object]
) -> ManagedAgentContext | AgentIdentityFailure | None:
issuer: Final = claims.get("iss")
tenant: Final = claims.get("tid")
client: Final = claims.get("azp")
if not isinstance(issuer, str) or not isinstance(tenant, str) or not isinstance(client, str):
return None
agent_id: Final = await self._bound_agent_id(tenant, client)
if agent_id is None or isinstance(agent_id, AgentIdentityFailure):
return agent_id
agent: Final = await self.agent(agent_id)
if isinstance(agent, AgentIdentityFailure):
return agent
if (
agent is None
or not agent.identity_managed
or not agent.enabled
or agent.identity is None
or not agent.identity.active
):
return AgentIdentityFailure(message="Agent is disabled or no longer bound to an identity")
subject: Final = classify_agent_subject(agent.identity, claims, agent.execution_mode)
if isinstance(subject, AgentIdentityFailure):
return subject
if subject.kind == "application":
return ManagedAgentContext(
agent_id=agent.agent_id,
binding_revision=agent.identity.revision,
mode=subject.mode,
subject_oid=subject.oid,
)
proven: Final = await self.subject(issuer, tenant, claims.get("oid"))
if isinstance(proven, AgentIdentityFailure):
return proven
human: Final = (
VerifiedHumanSubject.model_validate(proven.model_dump())
if proven is not None
and proven.kind == "human"
and proven.verified_via == "sso_interactive"
and proven.user_id is not None
else None
)
if human is None:
return AgentIdentityFailure(message="The delegated user must first sign in through trusted Microsoft SSO")
return ManagedAgentContext(
agent_id=agent.agent_id,
binding_revision=agent.identity.revision,
mode=subject.mode,
user_id=human.user_id,
subject_oid=subject.oid,
)
async def subject(
self, issuer: str, tenant_id: str, oid: object
) -> "LiteLLM_VerifiedSubject | AgentIdentityFailure | None":
if not isinstance(oid, str):
return None
try:
where: Final[LiteLLM_VerifiedSubjectWhereUniqueInput] = {
"issuer_tenant_id_oid": {"issuer": issuer, "tenant_id": tenant_id, "oid": oid}
}
return await self.humans.table.find_unique(where=where)
except Exception:
return AgentIdentityFailure(code="policy_unavailable", message="Subject classification is unavailable")
async def retired_agent(self, agent_id: str) -> bool | AgentIdentityFailure:
if self.retired_agents is None:
return AgentIdentityFailure(code="policy_unavailable", message="Agent history is unavailable")
try:
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")
async def record_authentication(self, context: ManagedAgentContext) -> AgentIdentityFailure | None:
try:
if context.binding_revision is None:
return AgentIdentityFailure(message="Agent authentication requires a binding revision")
where: Final[LiteLLM_AgentIdentityWhereInput] = {
"agent_id": context.agent_id,
"revision": context.binding_revision,
"active": True,
"agent": {"is": {"enabled": True, "identity_managed": True}},
}
data: Final[LiteLLM_AgentIdentityUpdateManyMutationInput] = {
"last_authenticated_at": datetime.now(timezone.utc)
}
count: Final = await self.identities.table.update_many(where=where, data=data)
if count != 1:
return AgentIdentityFailure(message="Agent identity changed during authentication; retry")
return None
except Exception:
return AgentIdentityFailure(code="policy_unavailable", message="Agent authentication could not be recorded")
async def enroll_interactive_human(
self,
subject: MicrosoftInteractiveSubject,
user_id: str,
) -> AgentIdentityFailure | None:
try:
where: Final[LiteLLM_VerifiedSubjectWhereUniqueInput] = {
"issuer_tenant_id_oid": {"issuer": subject.issuer, "tenant_id": subject.tenant_id, "oid": subject.oid}
}
create_data: Final[LiteLLM_VerifiedSubjectCreateInput] = {
"issuer": subject.issuer,
"tenant_id": subject.tenant_id,
"oid": subject.oid,
"user_id": user_id,
"verified_via": "sso_interactive",
}
data: Final[LiteLLM_VerifiedSubjectUpsertInput] = {"create": create_data, "update": {}}
row: Final = await self.humans.table.upsert(where=where, data=data)
if row.kind != "human" or row.user_id != user_id or row.verified_via != "sso_interactive":
return AgentIdentityFailure(message="Microsoft subject is already bound to another local identity")
return None
except Exception:
return AgentIdentityFailure(
code="policy_unavailable", message="Microsoft subject enrollment is unavailable"
)
async def resolve_managed_agent(
claims: Mapping[str, object],
client: object,
*,
cache: UserApiKeyCache | None = None,
) -> ManagedAgentContext | None:
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
if client is None:
return None
result: Final = await AgentIdentityStore.from_client(client, cache=cache).resolve_verified_claims(claims)
if isinstance(result, AgentIdentityFailure):
raise_identity_failure(result)
return result

View file

@ -0,0 +1,220 @@
from collections.abc import Mapping
from datetime import datetime
from typing import Final, NoReturn, TypedDict
from uuid import uuid4
from fastapi import HTTPException
from pydantic import TypeAdapter, ValidationError
from typing_extensions import ReadOnly
from litellm.types.agents import AgentResponse
from litellm.types.proxy.agent_identity import (
AgentExecutionMode,
AgentIdentityBinding,
AgentIdentityFailure,
AgentSubject,
EntraIdentityConfig,
)
_MODE: Final = TypeAdapter(AgentExecutionMode)
class IdentityFields(TypedDict, total=False):
provider: ReadOnly[str]
tenant_id: ReadOnly[str]
client_id: ReadOnly[str]
issuer: ReadOnly[str]
service_principal_id: ReadOnly[str | None]
required_roles: ReadOnly[tuple[str, ...]]
required_scopes: ReadOnly[tuple[str, ...]]
active: ReadOnly[bool]
revision: ReadOnly[str]
last_authenticated_at: ReadOnly[datetime | None]
class IdentityUpsert(TypedDict):
create: ReadOnly[IdentityFields]
update: ReadOnly[IdentityFields]
class IdentityRelationWrite(TypedDict, total=False):
create: ReadOnly[IdentityFields]
update: ReadOnly[IdentityFields]
upsert: ReadOnly[IdentityUpsert]
class IdentityHistoryKey(TypedDict):
provider: ReadOnly[str]
tenant_id: ReadOnly[str]
client_id: ReadOnly[str]
class IdentityHistoryWhere(TypedDict):
provider_tenant_id_client_id: ReadOnly[IdentityHistoryKey]
class IdentityHistoryEntry(IdentityHistoryKey):
issuer: ReadOnly[str]
class IdentityHistoryConnect(TypedDict):
where: ReadOnly[IdentityHistoryWhere]
create: ReadOnly[IdentityHistoryEntry]
class IdentityHistoryWrite(TypedDict):
connectOrCreate: ReadOnly[IdentityHistoryConnect]
class ManagedWriteFields(TypedDict, total=False):
enabled: ReadOnly[bool]
execution_mode: ReadOnly[AgentExecutionMode]
identity_managed: ReadOnly[bool]
identity: ReadOnly[IdentityRelationWrite]
retired_identities: ReadOnly[IdentityHistoryWrite]
def raise_identity_failure(failure: AgentIdentityFailure, status_code: int = 403) -> NoReturn:
raise HTTPException(503 if failure.code == "policy_unavailable" else status_code, failure.message)
def _configuration_failure(
identity: EntraIdentityConfig | AgentIdentityBinding | None,
mode: AgentExecutionMode,
enabling_without_binding: bool,
) -> AgentIdentityFailure | None:
if identity is not None and mode != "delegated" and not identity.service_principal_id:
return AgentIdentityFailure(
message="Autonomous mode requires the Enterprise application service-principal object ID"
)
if enabling_without_binding and (
identity is None or isinstance(identity, AgentIdentityBinding) and not identity.active
):
return AgentIdentityFailure(message="Bind an identity before enabling this managed agent")
return None
def managed_write_fields(
incoming: Mapping[str, object],
existing: AgentResponse | None,
updated_by: str,
) -> ManagedWriteFields | AgentIdentityFailure:
try:
identity: Final = (
EntraIdentityConfig.model_validate(incoming["identity"]) if incoming.get("identity") is not None else None
)
mode: Final = _MODE.validate_python(
incoming.get("execution_mode", existing.execution_mode if existing else "autonomous")
)
current_identity: Final = identity if "identity" in incoming else existing.identity if existing else None
failure: Final = _configuration_failure(
current_identity,
mode,
incoming.get("enabled") is True
and "identity" not in incoming
and bool(existing and existing.identity_managed),
)
if failure is not None:
return failure
empty: Final[ManagedWriteFields] = {}
identity_fields: Final = _identity_write(identity, existing) if "identity" in incoming else empty
result: Final[ManagedWriteFields] = {
**({"enabled": incoming["enabled"] is True} if "enabled" in incoming else {}),
**({"execution_mode": mode} if "execution_mode" in incoming else {}),
**identity_fields,
}
return result
except (ValidationError, ValueError) as exc:
return AgentIdentityFailure(message=f"Invalid agent identity configuration: {exc}")
def _identity_write(identity: EntraIdentityConfig | None, existing: AgentResponse | None) -> ManagedWriteFields:
if identity is None:
unbind: Final[ManagedWriteFields] = {
**(
{"identity": {"update": {"active": False, "revision": str(uuid4()), "last_authenticated_at": None}}}
if existing and existing.identity
else {}
),
**({"identity_managed": True, "enabled": False} if existing and existing.identity_managed else {}),
}
return unbind
if (
existing
and existing.identity
and existing.identity.active
and all(getattr(existing.identity, name) == value for name, value in identity.model_dump().items())
):
unchanged: Final[ManagedWriteFields] = {}
return unchanged
binding: Final[IdentityFields] = {
"provider": identity.provider,
"tenant_id": identity.tenant_id,
"client_id": identity.client_id,
"service_principal_id": identity.service_principal_id,
"required_roles": identity.required_roles,
"required_scopes": identity.required_scopes,
"issuer": identity.issuer,
"active": True,
"revision": str(uuid4()),
"last_authenticated_at": None,
}
result: Final[ManagedWriteFields] = {
"retired_identities": {
"connectOrCreate": {
"where": {
"provider_tenant_id_client_id": {
"provider": identity.provider,
"tenant_id": identity.tenant_id,
"client_id": identity.client_id,
}
},
"create": {
"provider": identity.provider,
"issuer": identity.issuer,
"tenant_id": identity.tenant_id,
"client_id": identity.client_id,
},
}
},
"identity_managed": True,
"identity": {"upsert": {"create": binding, "update": binding}} if existing else {"create": binding},
}
return result
def classify_agent_subject(
binding: AgentIdentityBinding,
claims: Mapping[str, object],
allowed_mode: AgentExecutionMode,
) -> AgentSubject | AgentIdentityFailure:
if (claims.get("iss"), claims.get("tid"), claims.get("azp")) != (
binding.issuer,
binding.tenant_id,
binding.client_id,
):
return AgentIdentityFailure(message="Token does not match the registered Entra application")
oid: Final = claims.get("oid")
if not isinstance(oid, str) or not oid:
return AgentIdentityFailure(message="Entra token must identify its object subject")
scope: Final = claims.get("scp")
facets: Final = claims.get("xms_sub_fct")
if facets is not None and (not isinstance(facets, str) or "13" in facets.split()):
return AgentIdentityFailure(message="Native agent-user authentication is not supported by this binding")
if scope is not None and not isinstance(scope, str):
return AgentIdentityFailure(message="Invalid delegated scope claim")
if isinstance(scope, str) and scope:
if allowed_mode == "autonomous" or oid == binding.service_principal_id or claims.get("idtyp") == "app":
return AgentIdentityFailure(message="Delegated token contradicts the configured agent identity or mode")
granted_scopes: Final = frozenset(scope.split())
if not granted_scopes or not frozenset(binding.required_scopes).issubset(granted_scopes):
return AgentIdentityFailure(message="Token lacks the required delegated scopes")
return AgentSubject(kind="delegated_subject", oid=oid, mode="delegated")
if allowed_mode == "delegated" or oid != binding.service_principal_id or claims.get("idtyp") == "user":
return AgentIdentityFailure(message="Application token contradicts the configured agent identity or mode")
roles: Final = claims.get("roles", ())
if not isinstance(roles, (list, tuple)) or any(not isinstance(role, str) for role in roles):
return AgentIdentityFailure(message="Invalid application roles claim")
if not frozenset(binding.required_roles).issubset(roles):
return AgentIdentityFailure(message="Token lacks the required application roles")
return AgentSubject(kind="application", oid=oid, mode="autonomous")

View file

@ -14,6 +14,7 @@ import math
import re
import time
from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence
from functools import partial
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Generic, Literal, Optional, Protocol, TypeAlias
@ -23,7 +24,7 @@ from typing_extensions import NotRequired, ReadOnly, Required, TypedDict, Unpack
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.caching.dual_cache import LimitedSizeOrderedDict
from litellm.caching.dual_cache import DualCache, LimitedSizeOrderedDict
from litellm.constants import (
CLI_JWT_EXPIRATION_HOURS,
CLI_SESSION_KEY_PREFIX,
@ -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,
@ -1784,11 +1800,12 @@ async def _load_bounded_registry(
if not isinstance(cached, _RegistryNotCached):
return cached
waited_for_another_load: Final = load_lock.locked()
async with load_lock:
# The request that held the lock has since cached an answer for everyone waiting on it.
cached_after_wait: Final = await _cached_registry(cache_key, overflow_sentinel, user_api_key_cache)
if not isinstance(cached_after_wait, _RegistryNotCached):
return cached_after_wait
if waited_for_another_load:
cached_after_wait: Final = await _cached_registry(cache_key, overflow_sentinel, user_api_key_cache)
if not isinstance(cached_after_wait, _RegistryNotCached):
return cached_after_wait
return await _fetch_and_cache_registry(
cache_key=cache_key,
@ -2641,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}
)
@ -2679,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},
)
@ -2782,17 +2799,12 @@ async def _cache_team_object(
team_table.last_refreshed_at = time.time()
key: Final = f"team_id:{team_id}"
usage_cache: Final = None if proxy_logging_obj is None else proxy_logging_obj.internal_usage_cache.dual_cache
# On a shared Redis the write below replaces the team entry and the alias DEL below removes the alias entry
# for both caches, so the usage cache only has its own memory to clear.
redis_shared: Final = usage_cache is not None and usage_cache.redis_cache is user_api_key_cache.redis_cache
if proxy_logging_obj is not None:
try:
await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=key)
except Exception as e: # noqa: BLE001 # best-effort invalidation: any cache backend error must not fail the write
verbose_proxy_logger.warning(
"Failed to invalidate internal usage cache entry %s; "
"a stale team object may be served until its TTL expires: %s",
key,
e,
)
await _invalidate_usage_cache_entry(usage_cache, key, redis_shared=redis_shared, stale="team object")
# team_id is the table primary key — guaranteed unique, safe to write.
await _cache_management_object(
@ -2819,9 +2831,11 @@ async def _cache_team_object(
if team_table.team_alias:
alias_key: Final = f"team_alias:{team_table.team_alias}"
try:
user_api_key_cache.delete_cache(key=alias_key)
if proxy_logging_obj is not None:
await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=alias_key)
pipelined_delete: Final = await user_api_key_cache.async_delete_cache_pre_call(alias_key)
if pipelined_delete is None:
await user_api_key_cache.async_delete_cache(key=alias_key)
else:
await pipelined_delete
except Exception as e: # noqa: BLE001 # best-effort invalidation: any cache backend error must not fail the mutation
verbose_proxy_logger.warning(
"Failed to invalidate cached team alias entry %s; "
@ -2829,6 +2843,30 @@ async def _cache_team_object(
alias_key,
e,
)
await _invalidate_usage_cache_entry(usage_cache, alias_key, redis_shared=redis_shared, stale="team alias")
async def _invalidate_usage_cache_entry(
usage_cache: DualCache | None,
key: str,
*,
redis_shared: bool,
stale: str,
) -> None:
if usage_cache is None:
return
try:
if redis_shared:
usage_cache.in_memory_cache.delete_cache(key)
else:
await usage_cache.async_delete_cache(key=key)
except Exception as e: # noqa: BLE001 # best-effort invalidation: any cache backend error must not fail the write
verbose_proxy_logger.warning(
"Failed to invalidate internal usage cache entry %s; a stale %s may be served until its TTL expires: %s",
key.replace("\r", "").replace("\n", ""),
stale,
e,
)
async def invalidate_team_member_spend_state(
@ -3104,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
)
@ -3140,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(
@ -3148,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:
@ -3170,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,
@ -3261,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
@ -3306,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)
"""
@ -3324,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(
@ -3362,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}"}
),
)
@ -3565,13 +3614,16 @@ async def get_org_object_by_alias(
)
LITELLM_SESSION_TOKEN_PREFIX: Final = "litellm_login_"
class ExperimentalUIJWTToken:
@staticmethod
def get_experimental_ui_login_jwt_auth_token(user_info: LiteLLM_UserTable) -> str:
from datetime import timedelta
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
encrypt_value_helper,
encrypt_bearer_token,
)
if user_info.user_role is None:
@ -3597,7 +3649,7 @@ class ExperimentalUIJWTToken:
user_role=LitellmUserRoles(user_info.user_role),
)
return encrypt_value_helper(valid_token.model_dump_json(exclude_none=True))
return encrypt_bearer_token(valid_token.model_dump_json(exclude_none=True), prefix=LITELLM_SESSION_TOKEN_PREFIX)
@staticmethod
def get_cli_jwt_auth_token(
@ -3628,7 +3680,7 @@ class ExperimentalUIJWTToken:
from datetime import timedelta
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
encrypt_value_helper,
encrypt_bearer_token,
)
if user_info.user_role is None:
@ -3666,7 +3718,7 @@ class ExperimentalUIJWTToken:
is_session_token=True,
)
return encrypt_value_helper(valid_token.model_dump_json(exclude_none=True))
return encrypt_bearer_token(valid_token.model_dump_json(exclude_none=True), prefix=LITELLM_SESSION_TOKEN_PREFIX)
@staticmethod
def get_key_object_from_ui_hash_key(
@ -3676,10 +3728,10 @@ class ExperimentalUIJWTToken:
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
decrypt_bearer_token,
)
decrypted_token: Final = decrypt_value_helper(hashed_token, key="ui_hash_key", exception_type="debug")
decrypted_token: Final = decrypt_bearer_token(hashed_token, prefix=LITELLM_SESSION_TOKEN_PREFIX)
if decrypted_token is None:
return None
try:
@ -3694,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.
@ -3707,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,
@ -3718,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,
@ -3743,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,
@ -3831,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
@ -3845,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)
@ -3861,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:
@ -3874,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,
@ -3882,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,
@ -3919,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
@ -3930,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
@ -3940,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())
@ -3956,6 +4029,8 @@ async def get_object_permission(
return _perm_obj
except Exception:
if check_db_only:
raise
return None
@ -4165,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
@ -4207,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,
@ -4242,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.
@ -4253,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,
)
@ -4261,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.
@ -4272,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,
)
@ -4471,26 +4554,37 @@ async def _check_agent_access_group_model_access(
"""Attached groups naming no model deny every model; the empty allowlist in ``_can_object_call_model`` allows."""
if not model or valid_token is None or not valid_token.agent_id:
return True
ceiling: Final = await resolve_ceiling(valid_token.agent_id)
if ceiling is None:
return True
if not ceiling.models:
raise ModelAccessDeniedProxyException(
message=model_access_denied_client_message(model=model),
internal_message=f"agent {valid_token.agent_id} access groups {ceiling.access_group_ids} grant no models",
type=ProxyErrorTypes.agent_model_access_denied,
param="model",
code=status.HTTP_403_FORBIDDEN,
)
dispatched: Final = _resolve_team_alias(model, valid_token.team_model_aliases, valid_token.team_id, llm_router)
return _can_object_call_model(
model=dispatched,
llm_router=llm_router,
models=sorted(ceiling.models),
team_id=valid_token.team_id,
object_type="agent",
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_managed_agent_ceilings
managed: Final = managed_agent_policy(valid_token)
unmanaged: Final = await resolve_ceiling(valid_token.agent_id) if managed is None else None
ceilings: Final = (
await resolve_managed_agent_ceilings(managed)
if managed is not None
else (unmanaged,)
if unmanaged is not None
else ()
)
dispatched: Final = _resolve_team_alias(model, valid_token.team_model_aliases, valid_token.team_id, llm_router)
for ceiling in ceilings:
if not ceiling.models:
raise ModelAccessDeniedProxyException(
message=model_access_denied_client_message(model=model),
internal_message=f"agent {valid_token.agent_id} access groups grant no models",
type=ProxyErrorTypes.agent_model_access_denied,
param="model",
code=status.HTTP_403_FORBIDDEN,
)
_can_object_call_model(
model=dispatched,
llm_router=llm_router,
models=sorted(ceiling.models),
team_id=valid_token.team_id,
object_type="agent",
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
)
return True
LoadedCallerTeam: TypeAlias = LiteLLM_TeamTable | None

View file

@ -15,7 +15,7 @@ from pydantic import BaseModel, TypeAdapter, ValidationError
from litellm._logging import verbose_proxy_logger
from litellm.caching.redis_batch import active_request_redis_batch
from litellm.caching.redis_cache import RedisCache
from litellm.constants import DEFAULT_IN_MEMORY_TTL
from litellm.constants import DEFAULT_IN_MEMORY_TTL, REGISTRY_ERROR_NEGATIVE_CACHE_TTL
from litellm.models.organization import LiteLLM_OrganizationTable
from litellm.models.team import LiteLLM_TeamTableCachedObj
from litellm.models.team_membership import LiteLLM_TeamMembership
@ -324,3 +324,35 @@ async def prefetch_auth_objects(
await _fill_from_db(refs, _missing_in_memory(missing, memory), user_api_key_cache, prisma_client)
except Exception as e: # noqa: BLE001 # warm-up only; the getters enforce and fail closed on their own
verbose_proxy_logger.warning("auth prefetch skipped, falling back to per-object lookups: %s", e)
def _identity_memory_ttl(value: object, management_ttl: float) -> float:
"""A registry stored as a string is a sentinel, written with the shorter of the two registry TTLs."""
return min(REGISTRY_ERROR_NEGATIVE_CACHE_TTL, management_ttl) if isinstance(value, str) else management_ttl
async def prefetch_identity_keys(cache_keys: Sequence[str], user_api_key_cache: UserApiKeyCache) -> None:
"""Warm the entries auth reads before it knows the key's owners (the key object, the end user and the two
registries) in one MGET on the request pipeline. Keys the MGET finds absent stay noted on the pipeline, so the
per-key getters that follow go to the database without a GET of their own. Best effort, like the
owner prefetch: the getters read and enforce on their own."""
try:
redis_cache: Final = user_api_key_cache.redis_cache
if redis_cache is None:
return
missing: Final = tuple(
key
for key in dict.fromkeys(cache_keys)
if user_api_key_cache.in_memory_cache_for(key).get_cache(key=key) is None
)
if not missing:
return
found: Final = _RowValues.validate_python(await _read_redis_rows(sorted(missing), redis_cache))
management_ttl: Final = get_management_object_ttl(user_api_key_cache)
except Exception as e: # noqa: BLE001 # warm-up only; the getters read Redis and the database on their own
verbose_proxy_logger.warning("auth identity prefetch skipped, falling back to per-key lookups: %s", e)
return
for key, value in ((key, found.get(key)) for key in missing):
if value is not None:
memory: _InMemoryCache = user_api_key_cache.in_memory_cache_for(key)
_set_in_memory(memory, key, value, _identity_memory_ttl(value, management_ttl))

View file

@ -73,7 +73,7 @@ from litellm.proxy.auth.auth_checks import (
)
from litellm.proxy.auth.auth_exception_handler import UserAPIKeyAuthExceptionHandler
from litellm.proxy.auth.auth_method import AuthMethod
from litellm.proxy.auth.auth_object_prefetch import AuthObjectRefs, prefetch_auth_objects
from litellm.proxy.auth.auth_object_prefetch import AuthObjectRefs, prefetch_auth_objects, prefetch_identity_keys
from litellm.proxy.auth.auth_utils import (
abbreviate_api_key,
get_end_user_id_from_request_body,
@ -120,6 +120,9 @@ from litellm.proxy.common_utils.model_listing_utils import claude_code_requested
from litellm.proxy.common_utils.realtime_utils import _realtime_request_body
from litellm.proxy.common_utils.user_api_key_cache import (
UserApiKeyCache,
end_user_cache_key,
end_user_restricted_registry_cache_key,
model_access_group_registry_cache_key,
team_membership_auth_cache_key,
)
from litellm.proxy.db.db_lookup_gate import bounded_db_lookup
@ -1892,6 +1895,11 @@ async def _user_api_key_auth_builder(
proxy_logging_obj=proxy_logging_obj,
route=route,
)
if prisma_client is not None:
await prefetch_identity_keys(
_identity_cache_keys(api_key, end_user_id=end_user_id, key_is_resolved=valid_token is not None),
user_api_key_cache=user_api_key_cache,
)
if end_user_id:
try:
end_user_params["end_user_id"] = end_user_id
@ -3196,6 +3204,15 @@ async def _authorize_authenticated_request(
# admin-only-route / model-access / budget checks) surface as
# ProxyException consistently with pre-refactor behavior.
try:
from litellm.proxy.agent_endpoints.auth.managed_authorization import admit_managed_actor
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
from litellm.proxy.proxy_server import prisma_client
if user_api_key_auth_obj.agent_id is not None:
await admit_managed_actor(
user_api_key_auth_obj,
AgentIdentityStore.from_client(prisma_client) if prisma_client is not None else None,
)
await _run_centralized_common_checks(
user_api_key_auth_obj=user_api_key_auth_obj,
request=request,
@ -3248,6 +3265,21 @@ def _spend_counter_redis_cache() -> RedisCache | None:
return spend_counter_cache.redis_cache
def _identity_cache_keys(api_key: str, *, end_user_id: str | None, key_is_resolved: bool) -> tuple[str, ...]:
"""Cache keys auth reads before it knows the key's owners, all known from the request alone. A key object is
cached under the hash of the bearer, so the bearer itself never reaches Redis."""
return tuple(
key
for key in (
None if key_is_resolved else hash_token(api_key),
None if not end_user_id else end_user_cache_key(end_user_id),
None if not end_user_id else end_user_restricted_registry_cache_key(),
model_access_group_registry_cache_key(),
)
if key is not None
)
async def _prefetch_referenced_auth_objects(
valid_token: UserAPIKeyAuth,
end_user_id: str | None,

View file

@ -72,26 +72,55 @@ def _derive_key(signing_key: str) -> bytes:
return hashlib.sha256(signing_key.encode()).digest()
def _encrypt_aes_gcm(value: str, signing_key: str) -> str:
"""Encrypt under AES-256-GCM and return the versioned ``v2:gcm:`` string."""
def _seal_aes_gcm(value: str, signing_key: str, aad: bytes | None) -> bytes:
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
nonce: Final = os.urandom(12)
# AESGCM.encrypt returns ciphertext || tag(16); wire format is nonce || that.
blob: Final = AESGCM(_derive_key(signing_key)).encrypt(nonce, value.encode("utf-8"), None)
return _V2_GCM_PREFIX + base64.urlsafe_b64encode(nonce + blob).decode("utf-8")
return nonce + AESGCM(_derive_key(signing_key)).encrypt(nonce, value.encode("utf-8"), aad)
def _open_aes_gcm(sealed: bytes, signing_key: str, aad: bytes | None) -> str:
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
# An empty plaintext still serializes to nonce(12) || tag(16) = 28 bytes, so a
# short/empty buffer here is a corrupt value: let AESGCM.decrypt raise and be
# swallowed by the caller (returns None/original), same as legacy.
return AESGCM(_derive_key(signing_key)).decrypt(sealed[:12], sealed[12:], aad).decode("utf-8")
def _encrypt_aes_gcm(value: str, signing_key: str) -> str:
"""Encrypt under AES-256-GCM and return the versioned ``v2:gcm:`` string."""
sealed: Final = _seal_aes_gcm(value=value, signing_key=signing_key, aad=None)
return _V2_GCM_PREFIX + base64.urlsafe_b64encode(sealed).decode("utf-8")
def _decrypt_aes_gcm(value: str, signing_key: str) -> str:
"""Decrypt a versioned ``v2:gcm:`` string produced by :func:`_encrypt_aes_gcm`."""
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
sealed: Final = base64.urlsafe_b64decode(value[len(_V2_GCM_PREFIX) :])
return _open_aes_gcm(sealed=sealed, signing_key=signing_key, aad=None)
raw: Final = base64.urlsafe_b64decode(value[len(_V2_GCM_PREFIX) :])
# An empty plaintext still serializes to nonce(12) || tag(16) = 28 bytes, so a
# short/empty buffer here is a corrupt value: let AESGCM.decrypt raise and be
# swallowed by decrypt_value_helper (returns None/original), same as legacy.
nonce, blob = raw[:12], raw[12:]
return AESGCM(_derive_key(signing_key)).decrypt(nonce, blob, None).decode("utf-8")
def encrypt_bearer_token(value: str, prefix: str) -> str:
"""AES-256-GCM as unpadded base64url behind ``prefix``, which is also the AAD so a token can't change kind."""
salt_key: Final = _get_salt_key()
if not isinstance(salt_key, str):
raise ValueError("Set LITELLM_SALT_KEY or a master key to mint bearer tokens")
sealed: Final = _seal_aes_gcm(value=value, signing_key=salt_key, aad=prefix.encode("utf-8"))
return prefix + base64.urlsafe_b64encode(sealed).decode("ascii").rstrip("=")
def decrypt_bearer_token(token: str, prefix: str) -> str | None:
"""None unless ``token`` came from :func:`encrypt_bearer_token` with the same ``prefix``."""
salt_key: Final = _get_salt_key()
if not isinstance(salt_key, str) or not token.startswith(prefix):
return None
encoded: Final = token.removeprefix(prefix)
try:
sealed: Final = base64.b64decode(encoded + "=" * (-len(encoded) % 4), altchars=b"-_", validate=True)
return _open_aes_gcm(sealed=sealed, signing_key=salt_key, aad=prefix.encode("utf-8"))
except Exception: # noqa: BLE001 # base64 and AES-GCM each raise their own "not a token" type
return None
def encrypt_value_helper(value: str, new_encryption_key: str | None = None):

View file

@ -17,6 +17,8 @@ from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec
if TYPE_CHECKING:
from opentelemetry.trace import Span
from litellm.caching.redis_batch import BatchResult
T = TypeVar("T", bound=BaseModel)
_HASHED_TOKEN_CACHE_KEY: Final = re.compile(r"[0-9a-f]{64}")
@ -27,6 +29,9 @@ def is_user_key_cache_key(key: str) -> bool:
return _HASHED_TOKEN_CACHE_KEY.fullmatch(key) is not None
_PIPELINED_SET_OPTIONS: Final = frozenset(("ttl",))
class UserApiKeyCache(DualCache):
"""
DualCache wrapper for UserAPIKeyAuth-like payloads.
@ -208,10 +213,23 @@ class UserApiKeyCache(DualCache):
return super().set_cache(key=key, value=payload, local_only=local_only, **kwargs)
async def async_set_cache(self, key: str | None, value: object, local_only: bool = False, **kwargs: object):
"""Inside a request the Redis SET rides the request's pipeline (memory is written at once); anywhere
else, or with options the pipeline does not carry, it goes to Redis directly as before."""
model_type: Final = cast(type[BaseModel] | None, kwargs.pop("model_type", None))
payload: Final[object] = CacheCodec.serialize(value, model_type=model_type)
ttl: Final = kwargs.get("ttl")
pipelined: Final = (
key is not None
and not local_only
and kwargs.keys() <= _PIPELINED_SET_OPTIONS
and (ttl is None or isinstance(ttl, (int, float)))
)
if key is not None and is_user_key_cache_key(key):
if pipelined and await self.key_object_cache.async_set_cache_pre_call(key, payload, ttl) is not None:
return None
return await self.key_object_cache.async_set_cache(key=key, value=payload, local_only=local_only, **kwargs)
if pipelined and await super().async_set_cache_pre_call(key, payload, ttl) is not None:
return None
return await super().async_set_cache(key=key, value=payload, local_only=local_only, **kwargs)
def delete_cache(self, key: str) -> None:
@ -226,6 +244,11 @@ class UserApiKeyCache(DualCache):
return
await super().async_delete_cache(key)
async def async_delete_cache_pre_call(self, key: str) -> BatchResult[None] | None:
if is_user_key_cache_key(key):
return await self.key_object_cache.async_delete_cache_pre_call(key)
return await super().async_delete_cache_pre_call(key)
async def async_delete_cache_keys(self, keys: Sequence[str]) -> None:
"""Batch twin of ``async_delete_cache``, partitioned like
``async_set_cache_pipeline``.

View file

@ -2442,11 +2442,13 @@ async def _update_key_row_with_soft_budget(
existing_key_row=existing_key_row,
changed_by=changed_by,
)
include_object_permission: Final[prisma.types.LiteLLM_VerificationTokenInclude] = {"object_permission": True}
updated_row: Final = await tx.litellm_verificationtoken.update(
where=key_where,
data=with_settings_updated_at(
prisma_client.jsonify_object(MappingProxyType({**update_values, "token": hashed_token}))
),
include=include_object_permission,
)
updated_data: Final[Mapping[str, object]] = (
updated_row.model_dump() if updated_row is not None else MappingProxyType({})

View file

@ -142,6 +142,7 @@ ROUTE_ENDPOINT_MAPPING: Final = {
"acancel_run": "/evals/{eval_id}/runs/{run_id}/cancel",
"adelete_run": "/evals/{eval_id}/runs/{run_id}",
"acreate_batch": "/batches",
"aretrieve_batch": "/batches",
}

View file

@ -78,6 +78,11 @@ model LiteLLM_AgentsTable {
object_permission_id String?
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
spend Float @default(0.0)
identity_managed Boolean @default(false)
enabled Boolean @default(true)
execution_mode String @default("autonomous")
identity LiteLLM_AgentIdentity?
retired_identities LiteLLM_RetiredAgentIdentity[]
tpm_limit Int?
rpm_limit Int?
session_tpm_limit Int?
@ -88,6 +93,56 @@ model LiteLLM_AgentsTable {
updated_by String
}
model LiteLLM_AgentIdentity {
agent_id String @id
active Boolean @default(true)
agent LiteLLM_AgentsTable @relation(fields: [agent_id], references: [agent_id], onDelete: Cascade)
provider String
issuer String
tenant_id String
client_id String
service_principal_id String?
required_roles String[] @default([])
required_scopes String[] @default(["user_impersonation"])
revision String @default(uuid())
last_authenticated_at DateTime?
@@unique([provider, tenant_id, client_id])
@@unique([issuer, service_principal_id])
}
model LiteLLM_RetiredAgentIdentity {
binding_id String @id @default(uuid())
agent_id String?
agent LiteLLM_AgentsTable? @relation(fields: [agent_id], references: [agent_id], onDelete: SetNull)
provider String
issuer String
tenant_id String
client_id String
@@unique([provider, tenant_id, client_id])
}
model LiteLLM_RetiredAgent {
original_agent_id String @id
retired_at DateTime @default(now())
}
model LiteLLM_VerifiedSubject {
subject_id String @id @default(uuid())
issuer String
tenant_id String
oid String
kind String @default("human")
user_id String?
user LiteLLM_UserTable? @relation(fields: [user_id], references: [user_id], onDelete: Cascade)
verified_via String @default("sso_interactive")
verified_at DateTime @default(now())
@@unique([issuer, tenant_id, oid])
@@index([user_id])
}
model LiteLLM_OrganizationTable {
organization_id String @id @default(uuid())
organization_alias String
@ -241,6 +296,7 @@ model LiteLLM_DeletedTeamTable {
// Track spend, rate limit, budget Users
model LiteLLM_UserTable {
verified_subjects LiteLLM_VerifiedSubject[]
user_id String @id
user_alias String?
team_id String?
@ -675,6 +731,7 @@ model LiteLLM_SpendLogs {
session_id String?
status String?
mcp_namespaced_tool_name String?
billing_agent_id String?
agent_id String?
proxy_server_request Json? @default("{}")
litellm_call_id String?

View file

@ -4,7 +4,7 @@ from collections.abc import Set as AbstractSet
from dataclasses import dataclass
from datetime import datetime, timedelta
from types import MappingProxyType
from typing import Final, TypeVar
from typing import Final, Literal, TypeVar
from pydantic import BaseModel, TypeAdapter
from typing_extensions import ReadOnly, TypedDict
@ -17,6 +17,7 @@ from litellm.constants import (
SPEND_LOG_KEY_METADATA_CACHE_TTL,
SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL,
SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS,
SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE,
)
from litellm.litellm_core_utils.litellm_logging import is_valid_sha256_hash
from litellm.proxy.utils import PrismaClient
@ -39,26 +40,58 @@ WHERE encode(sha256(convert_to(token, 'UTF8')), 'hex') = ANY($1::text[])
ORDER BY token, deleted_at DESC
"""
_SPEND_LOG_ALIAS_SQL: Final = """
SELECT api_key AS digest,
MIN(key_alias) AS first_alias,
MAX(key_alias) AS last_alias,
MIN(team_id) AS first_team,
MAX(team_id) AS last_team,
MIN(user_id) AS first_owner,
MAX(user_id) AS last_owner
FROM (
SELECT api_key,
NULLIF(metadata->>'user_api_key_alias', '') AS key_alias,
COALESCE(NULLIF(team_id, ''), NULLIF(metadata->>'user_api_key_team_id', '')) AS team_id,
COALESCE(NULLIF("user", ''), NULLIF(metadata->>'user_api_key_user_id', '')) AS user_id
FROM "LiteLLM_SpendLogs"
WHERE api_key = ANY($1::text[])
AND "startTime" >= $2::timestamp
AND "startTime" < $3::timestamp
) named
WHERE COALESCE(key_alias, user_id, team_id) IS NOT NULL
GROUP BY api_key
def _named_spend_log_edge_row_sql(
direction: Literal["ASC", "DESC"], since: Literal["$2::timestamp", "oldest_probe.stopped_at"]
) -> str:
return f"""
SELECT "startTime", key_alias, team_id, user_id
FROM (
SELECT "startTime",
NULLIF(metadata->>'user_api_key_alias', '') AS key_alias,
COALESCE(NULLIF(team_id, ''), NULLIF(metadata->>'user_api_key_team_id', '')) AS team_id,
COALESCE(NULLIF("user", ''), NULLIF(metadata->>'user_api_key_user_id', '')) AS user_id
FROM (
SELECT "startTime", metadata, team_id, "user"
FROM "LiteLLM_SpendLogs"
WHERE api_key = keys.digest
AND "startTime" >= {since}
AND "startTime" < $3::timestamp
ORDER BY "startTime" {direction}
LIMIT {SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE}
) edge
) named
WHERE COALESCE(key_alias, user_id, team_id) IS NOT NULL
ORDER BY "startTime" {direction}
LIMIT 1
"""
_OLDEST_PROBE_STOPPED_AT_SQL: Final = f"""
SELECT COALESCE(first_row."startTime", (
SELECT "startTime"
FROM "LiteLLM_SpendLogs"
WHERE api_key = keys.digest
AND "startTime" >= $2::timestamp
AND "startTime" < $3::timestamp
ORDER BY "startTime" ASC
OFFSET {SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE - 1}
LIMIT 1
)) AS stopped_at
"""
_SPEND_LOG_ALIAS_SQL: Final = f"""
SELECT keys.digest,
first_row.key_alias AS first_alias,
last_row.key_alias AS last_alias,
first_row.team_id AS first_team,
last_row.team_id AS last_team,
first_row.user_id AS first_owner,
last_row.user_id AS last_owner
FROM unnest($1::text[]) AS keys(digest)
LEFT JOIN LATERAL ({_named_spend_log_edge_row_sql("ASC", "$2::timestamp")}) first_row ON true
LEFT JOIN LATERAL ({_OLDEST_PROBE_STOPPED_AT_SQL}) oldest_probe ON true
LEFT JOIN LATERAL ({_named_spend_log_edge_row_sql("DESC", "oldest_probe.stopped_at")}) last_row ON true
"""
_DAILY_USER_SPEND_OWNER_SQL: Final = """
@ -69,6 +102,7 @@ GROUP BY api_key
"""
_SPEND_LOG_STATEMENT_TIMEOUT_SQL: Final = f"SET LOCAL statement_timeout = {SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS}"
_SPEND_LOG_NO_BITMAP_SCAN_SQL: Final = "SET LOCAL enable_bitmapscan = off"
_SPEND_LOG_TRANSACTION_TIMEOUT: Final = timedelta(milliseconds=2 * SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS)
_HASHED_JWT_PREFIX: Final = "hashed-jwt-"
@ -91,7 +125,9 @@ class _TokenDigestRow(BaseModel):
def _unanimous(first: str | None, last: str | None) -> str | None:
return first if first == last else None
if first is None:
return last
return first if last is None or first == last else None
class _SpendLogDigestRow(BaseModel):
@ -148,9 +184,12 @@ async def _rows_within_the_statement_timeout(
prisma_client: PrismaClient,
sql: str,
*params: object,
planner_settings: tuple[str, ...] = (),
) -> Sequence[Mapping[str, object]]:
async with prisma_client.db.tx(timeout=_SPEND_LOG_TRANSACTION_TIMEOUT) as transaction:
await transaction.execute_raw(_SPEND_LOG_STATEMENT_TIMEOUT_SQL)
for setting in planner_settings:
await transaction.execute_raw(setting)
return await transaction.query_raw(sql, *params)
@ -364,7 +403,14 @@ async def _query_spend_log_metadata(
) -> Mapping[str, KeyMetadataDict] | None:
start, end = window
rows: Final = await _db_or_empty(
lambda: _rows_within_the_statement_timeout(prisma_client, _SPEND_LOG_ALIAS_SQL, sorted(digests), start, end),
lambda: _rows_within_the_statement_timeout(
prisma_client,
_SPEND_LOG_ALIAS_SQL,
sorted(digests),
start,
end,
planner_settings=(_SPEND_LOG_NO_BITMAP_SCAN_SQL,),
),
"Failed spend-log alias recovery for %d missing keys: %s",
len(digests),
)

View file

@ -2444,7 +2444,7 @@ def _build_spend_log_search_condition(
f"(request_id = {raw} OR ("
f"\"startTime\" >= ({window_start}::timestamptz AT TIME ZONE 'UTC') "
f"AND \"startTime\" <= ({window_end}::timestamptz AT TIME ZONE 'UTC') "
f'AND (api_key = {raw} OR team_id = {raw} OR "user" = {raw} OR end_user = {raw} '
f'AND (litellm_call_id = {raw} OR api_key = {raw} OR team_id = {raw} OR "user" = {raw} OR end_user = {raw} '
f"OR session_id = {raw} OR model_id = {raw})))"
)
return _SpendLogSearchCondition(sql=sql, params=(search, start_date, end_date))
@ -2557,7 +2557,7 @@ async def ui_view_spend_logs(
search: str | None = fastapi.Query(
default=None,
description=(
"Match a log whose request_id, api_key (hash), team_id, user, end_user, "
"Match a log whose request_id, litellm_call_id, api_key (hash), team_id, user, end_user, "
"session_id, or model_id equals this value. request_id matches across all time; the other columns "
"match inside start_date/end_date, which stay required"
),

View file

@ -44,6 +44,7 @@ from litellm.litellm_core_utils.litellm_logging import (
)
from litellm.litellm_core_utils.ptu_pricing import azure_spillover
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps, strip_null_bytes
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload, SpendLogsRouterMetadata
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error
@ -1083,6 +1084,11 @@ def _get_messages_for_spend_logs_payload(
_SENSITIVE_REQUEST_BODY_KEYS: Final = frozenset({"secret_fields"})
_REQUEST_BODY_CREDENTIAL_MASKER: Final = SensitiveDataMasker(extra_sensitive_patterns=frozenset({"apikey"}))
def _is_request_body_credential(key: str, value: object) -> bool:
return isinstance(value, str) and _REQUEST_BODY_CREDENTIAL_MASKER.is_sensitive_key(key)
def _sanitize_request_body_for_spend_logs_payload(
@ -1094,8 +1100,9 @@ def _sanitize_request_body_for_spend_logs_payload(
Recursively sanitize request body to prevent logging large base64 strings or other large values.
Truncates strings longer than MAX_STRING_LENGTH_PROMPT_IN_DB characters and handles nested dictionaries.
Also strips keys listed in _SENSITIVE_REQUEST_BODY_KEYS (e.g. secret_fields
which contains raw HTTP headers including Authorization tokens).
At every nesting level, also strips keys listed in _SENSITIVE_REQUEST_BODY_KEYS (e.g. secret_fields,
which holds raw HTTP headers including Authorization tokens), and replaces string values under keys
SensitiveDataMasker classifies as credentials with REDACTED_BY_LITELM_STRING.
"""
from litellm.constants import (
LITELLM_TRUNCATED_PAYLOAD_FIELD,
@ -1152,7 +1159,11 @@ def _sanitize_request_body_for_spend_logs_payload(
return value
return value
return {k: _sanitize_value(v) for k, v in request_body.items() if k not in _SENSITIVE_REQUEST_BODY_KEYS}
return {
k: REDACTED_BY_LITELM_STRING if _is_request_body_credential(k, v) else _sanitize_value(v)
for k, v in request_body.items()
if k not in _SENSITIVE_REQUEST_BODY_KEYS
}
# Quoted-key form: ``"input"`` / ``'messages'`` / ``"prompt"`` followed by

View file

@ -258,7 +258,7 @@ if TYPE_CHECKING:
from prisma.actions import LiteLLM_DeprecatedVerificationTokenActions
from prisma.client import TransactionManager
from prisma.models import LiteLLM_DeprecatedVerificationToken
from prisma.types import HttpConfig
from prisma.types import HttpConfig, LiteLLM_VerificationTokenInclude
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
@ -4200,6 +4200,8 @@ _PRISMA_DEFAULT_TX_TIMEOUT: Final = timedelta(seconds=5)
async def _lookup_deprecated_key(
db: PrismaWrapper | RoutingPrismaWrapper,
hashed_token: str,
*,
check_db_only: bool = False,
) -> str | None:
"""
Check if a token exists in the deprecated keys table and is still within its grace period.
@ -4211,7 +4213,7 @@ async def _lookup_deprecated_key(
now_ts: Final = now.timestamp()
# Check cache first
cached: Final = _deprecated_key_cache.get(hashed_token)
cached: Final = None if check_db_only else _deprecated_key_cache.get(hashed_token)
if cached is not None:
active_token_id, cache_expires_at_ts, revoke_at_ts = cached
if now_ts < cache_expires_at_ts and now_ts < revoke_at_ts:
@ -4879,6 +4881,7 @@ class PrismaClient:
proxy_logging_obj: ProxyLogging | None = None,
budget_id_list: list[str] | None = None,
check_deprecated: bool = True,
use_writer: bool = False,
):
args_passed_in: Final = locals()
start_time: Final = time.time()
@ -5177,12 +5180,20 @@ class PrismaClient:
WHERE v.token = $1
"""
response = await self._query_first_with_cached_plan_fallback(sql_query, hashed_token)
response = (
await self.writer_db.query_first(sql_query, hashed_token)
if use_writer
else await self._query_first_with_cached_plan_fallback(sql_query, hashed_token)
)
# If not found in main table, check deprecated keys (grace period)
# check_deprecated=False on the recursive call prevents unbounded chaining
if response is None and hashed_token is not None and check_deprecated:
active_token_id: Final = await _lookup_deprecated_key(db=self.db, hashed_token=hashed_token)
active_token_id: Final = await _lookup_deprecated_key(
db=self.writer_db if use_writer else self.db,
hashed_token=hashed_token,
check_db_only=use_writer,
)
if active_token_id:
# The recursive call returns a finished
# LiteLLM_VerificationTokenView; the dict
@ -5194,6 +5205,7 @@ class PrismaClient:
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_deprecated=False,
use_writer=use_writer,
)
if deprecated_response is not None:
verbose_proxy_logger.debug("Deprecated key used during grace period")
@ -5445,9 +5457,11 @@ class PrismaClient:
# check if plain text or hash
token = _hash_token_if_needed(token=token)
db_data["token"] = token
include_object_permission: Final[LiteLLM_VerificationTokenInclude] = {"object_permission": True}
response: Final = await VerificationTokenRepository(self).table.update(
where={"token": token},
data=with_settings_updated_at(db_data),
include=include_object_permission,
)
verbose_proxy_logger.debug("\033[91m" + f"DB Token Table update succeeded {response}" + "\033[0m")
_data: dict = {}

View file

@ -15,9 +15,14 @@ if TYPE_CHECKING:
class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]):
"""Repository for object permission database operations."""
def __init__(self, prisma_client: object, *, use_writer: bool = False) -> None:
super().__init__(prisma_client)
self._use_writer = use_writer
@property
def table(self) -> TableActions["prisma_models.LiteLLM_ObjectPermissionTable"]:
return self.prisma_client.db.litellm_objectpermissiontable
database: Final = self.prisma_client.writer_db if self._use_writer else self.prisma_client.db
return database.litellm_objectpermissiontable
@property
def model_class(self) -> type[LiteLLM_ObjectPermissionTable]:

View file

@ -21,8 +21,9 @@ class PrismaTableRepository(Generic[RowT_co]):
table_name: str
def __init__(self, prisma_client: object):
def __init__(self, prisma_client: object, *, use_writer: bool = False) -> None:
self._prisma_client = prisma_client
self._use_writer = use_writer
@property
def prisma_client(self) -> Any:
@ -32,7 +33,9 @@ class PrismaTableRepository(Generic[RowT_co]):
@property
def table(self) -> TableActions[RowT_co]:
actions: Final[TableActions[RowT_co]] = getattr(self.prisma_client.db, self.table_name)
actions: Final[TableActions[RowT_co]] = getattr(
self.prisma_client.writer_db if self._use_writer else self.prisma_client.db, self.table_name
)
return wrap_table_actions_for_config_sync(actions=actions, table_name=self.table_name)
@ -44,6 +47,18 @@ class AgentsRepository(PrismaTableRepository["prisma_models.LiteLLM_AgentsTable"
table_name = "litellm_agentstable"
class AgentIdentityRepository(PrismaTableRepository["prisma_models.LiteLLM_AgentIdentity"]):
table_name = "litellm_agentidentity"
class RetiredAgentIdentityRepository(PrismaTableRepository["prisma_models.LiteLLM_RetiredAgentIdentity"]):
table_name = "litellm_retiredagentidentity"
class VerifiedSubjectRepository(PrismaTableRepository["prisma_models.LiteLLM_VerifiedSubject"]):
table_name = "litellm_verifiedsubject"
class ObjectPermissionRepository(PrismaTableRepository["prisma_models.LiteLLM_ObjectPermissionTable"]):
table_name = "litellm_objectpermissiontable"
@ -250,3 +265,7 @@ class AuditLogRepository(PrismaTableRepository["prisma_models.LiteLLM_AuditLog"]
class AdaptiveRouterSessionRepository(PrismaTableRepository["prisma_models.LiteLLM_AdaptiveRouterSession"]):
table_name = "litellm_adaptiveroutersession"
class RetiredAgentRepository(PrismaTableRepository["prisma_models.LiteLLM_RetiredAgent"]):
table_name = "litellm_retiredagent"

View file

@ -70,6 +70,9 @@ class _PrismaClientView(Protocol):
@property
def db(self) -> _PrismaTeamDb: ...
@property
def writer_db(self) -> _PrismaTeamDb: ...
_MEMBERS_WITH_ROLES_ADAPTER: Final = TypeAdapter(list[Member])
_JSON_ENCODED_TEAM_FIELDS: Final = (
@ -85,10 +88,14 @@ _JSON_ENCODED_TEAM_FIELDS: Final = (
class TeamRepository(BaseRepository[LiteLLM_TeamTable]):
"""Repository for team database operations."""
def __init__(self, prisma_client: object, *, use_writer: bool = False) -> None:
super().__init__(prisma_client)
self._use_writer = use_writer
@property
def _db(self) -> _PrismaTeamDb:
client: Final[_PrismaClientView] = self.prisma_client
return client.db
return client.writer_db if self._use_writer else client.db
@property
def table(self) -> TableActions["prisma_models.LiteLLM_TeamTable"]:

View file

@ -38,9 +38,14 @@ _PLACEHOLDER_ROWS_ADAPTER: Final = TypeAdapter(tuple[SCIMPlaceholder, ...])
class UserRepository(BaseRepository[LiteLLM_UserTable]):
"""Repository for user database operations."""
def __init__(self, prisma_client: object, *, use_writer: bool = False) -> None:
super().__init__(prisma_client)
self._use_writer = use_writer
@property
def table(self) -> TableActions["prisma_models.LiteLLM_UserTable"]:
return self.prisma_client.db.litellm_usertable
database: Final = self.prisma_client.writer_db if self._use_writer else self.prisma_client.db
return database.litellm_usertable
@property
def model_class(self) -> type[LiteLLM_UserTable]:

View file

@ -601,7 +601,7 @@ class BaseResponsesAPIStreamingIterator:
raw_headers: Final[Mapping[str, object]] = raw if isinstance(raw, Mapping) else EMPTY_MAPPING
# rebuild by value and let existing keys win: sharing the source dicts would alias what the proxy
# splats into the client's HTTP headers, and copying non-header keys would carry response_cost
target._hidden_params = { # mutable-ok: the cost calculator writes optional_params into _hidden_params
target._hidden_params = { # mutable-ok: logging aliases _hidden_params into request metadata and writes into it
"additional_headers": {**headers}, # mutable-ok: fresh copy, logging callbacks may mutate it
"headers": {**raw_headers}, # mutable-ok: fresh copy, logging callbacks may mutate it
**existing,

View file

@ -1,6 +1,6 @@
import base64
import re
from collections.abc import Iterable, Mapping, Sequence
from collections.abc import Callable, Iterable, Mapping, Sequence
from functools import reduce
from typing import Any, Final, Optional, TypeVar, Union, cast, get_type_hints, overload
@ -556,7 +556,11 @@ class ResponsesAPIRequestUtils:
return request_input
@staticmethod
def strip_encrypted_reasoning_from_input(request_input: object) -> None:
def strip_encrypted_reasoning_from_input(
request_input: object,
*,
should_strip: Callable[[Mapping[str, object]], bool] | None = None,
) -> None:
"""Drop reasoning items the routed deployment cannot decrypt, keeping their readable summary.
Mutates ``request_input`` in place: the router's fallback snapshot shares this
@ -565,7 +569,12 @@ class ResponsesAPIRequestUtils:
if not isinstance(request_input, list):
return
items: Final = cast(list[object], request_input) # cast-ok: untyped client json
stripped: Final = tuple(ResponsesAPIRequestUtils._without_encrypted_reasoning(item) for item in items)
stripped: Final = tuple(
ResponsesAPIRequestUtils._without_encrypted_reasoning(item)
if should_strip is None or (isinstance(item, Mapping) and should_strip(cast(Mapping[str, object], item)))
else item
for item in items
)
items[:] = (item for item in stripped if item is not None)
@staticmethod

View file

@ -1813,7 +1813,6 @@ class Router:
messages: list[dict[str, str]] | None,
input: str | list | None,
request_kwargs: dict | None,
prefetched_usage: PrefetchedUsage | None = None,
) -> Any | None:
"""
Asks the strategy selector for a deployment. Caller handles
@ -1839,14 +1838,6 @@ class Router:
messages=messages,
input=input,
)
case "usage-based-routing-v2" if isinstance(selector, LowestTPMLoggingHandler_v2):
return await selector.async_get_available_deployments(
model_group=model,
healthy_deployments=healthy_deployments,
messages=messages,
input=input,
prefetched_usage=prefetched_usage,
)
case "usage-based-routing-v2" | "cost-based-routing":
return await selector.async_get_available_deployments(
model_group=model,
@ -12958,7 +12949,6 @@ class Router:
specific_deployment: bool | None = False,
parent_otel_span: Span | None = None,
health_check_probe: bool = False,
routing_read_batch: RoutingReadBatch | None = None,
) -> list[dict] | dict:
"""
Get the healthy deployments for a model.
@ -13011,6 +13001,7 @@ class Router:
health_check_probe=health_check_probe,
)
routing_read_batch: Final = RoutingReadBatch.active()
cooldown_deployments: Final = (
await _async_get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=parent_otel_span)
if routing_read_batch is None
@ -13298,15 +13289,15 @@ class Router:
strategy, strategy_selector = self._get_routing_context(model, request_kwargs)
routing_read_batch: Final = RoutingReadBatch.for_strategy(strategy, strategy_selector)
healthy_deployments: Final = await self.async_get_healthy_deployments(
model=model,
request_kwargs=request_kwargs,
messages=messages,
input=input,
specific_deployment=specific_deployment,
parent_otel_span=parent_otel_span,
routing_read_batch=routing_read_batch,
)
with RoutingReadBatch.scoped(routing_read_batch):
healthy_deployments: Final = await self.async_get_healthy_deployments(
model=model,
request_kwargs=request_kwargs,
messages=messages,
input=input,
specific_deployment=specific_deployment,
parent_otel_span=parent_otel_span,
)
if isinstance(healthy_deployments, dict):
await self._async_override_selector_pre_call_check(
strategy, strategy_selector, healthy_deployments, parent_otel_span
@ -13328,16 +13319,18 @@ class Router:
model=model,
request_kwargs=request_kwargs,
)
deployment: Final = await self._select_deployment_async(
strategy=strategy,
selector=strategy_selector,
model=model,
healthy_deployments=healthy_deployments,
messages=messages,
input=input,
request_kwargs=request_kwargs,
prefetched_usage=routing_read_batch.prefetched_usage if routing_read_batch is not None else None,
)
with PrefetchedUsage.scoped(
routing_read_batch.prefetched_usage if routing_read_batch is not None else None
):
deployment: Final = await self._select_deployment_async(
strategy=strategy,
selector=strategy_selector,
model=model,
healthy_deployments=healthy_deployments,
messages=messages,
input=input,
request_kwargs=request_kwargs,
)
if deployment is None:
exception: Final = await async_raise_no_deployment_exception(
litellm_router_instance=self,

View file

@ -1,7 +1,9 @@
#### What this does ####
# identifies lowest tpm deployment
import random
from collections.abc import Mapping, Sequence
from collections.abc import Iterator, Mapping, Sequence
from contextlib import contextmanager
from contextvars import ContextVar
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Final
@ -32,6 +34,9 @@ class RoutingArgs(LiteLLMPydanticObjectBase):
ttl: int = 1 * 60 # 1min (RPM/TPM expire key)
_active_prefetched_usage: Final[ContextVar["PrefetchedUsage | None"]] = ContextVar("prefetched_usage", default=None)
@dataclass(frozen=True)
class PrefetchedUsage:
"""
@ -51,6 +56,19 @@ class PrefetchedUsage:
return None
return [self.values.get(key) for key in keys]
@staticmethod
@contextmanager
def scoped(usage: "PrefetchedUsage | None") -> Iterator[None]:
token: Final = _active_prefetched_usage.set(usage)
try:
yield
finally:
_active_prefetched_usage.reset(token)
@staticmethod
def active() -> "PrefetchedUsage | None":
return _active_prefetched_usage.get()
class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
"""
@ -436,18 +454,15 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
def usage_counter_keys(self, healthy_deployments: list) -> tuple[list[str], list[str]]:
"""The `<id>:<model>:tpm:<HH-MM>` and `<id>:<model>:rpm:<HH-MM>` counter keys selection reads."""
current_minute: Final = get_utc_datetime().strftime("%H-%M")
tpm_keys: Final[list[str]] = []
rpm_keys: Final[list[str]] = []
for m in healthy_deployments:
if isinstance(m, dict):
id = m.get("model_info", {}).get(
"id"
) # a deployment should always have an 'id'. this is set in router.py
deployment_name = m.get("litellm_params", {}).get("model")
tpm_keys.append(f"{id}:{deployment_name}:tpm:{current_minute}")
rpm_keys.append(f"{id}:{deployment_name}:rpm:{current_minute}")
return tpm_keys, rpm_keys
prefixes: Final = tuple(
f"{m.get('model_info', {}).get('id')}:{m.get('litellm_params', {}).get('model')}"
for m in healthy_deployments
if isinstance(m, dict)
)
return (
[f"{prefix}:tpm:{current_minute}" for prefix in prefixes],
[f"{prefix}:rpm:{current_minute}" for prefix in prefixes],
)
async def async_get_available_deployments(
self,
@ -455,13 +470,13 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
healthy_deployments: list,
messages: list[dict[str, str]] | None = None,
input: str | list | None = None,
prefetched_usage: PrefetchedUsage | None = None,
):
"""
Async implementation of get deployments.
Reduces time to retrieve the tpm/rpm values from cache. `prefetched_usage` skips the cache
read when it already holds this request's counters (see `RoutingReadBatch`).
Reduces time to retrieve the tpm/rpm values from cache. A `PrefetchedUsage` scoped
to this request skips the cache read when it already holds its counters (see
`RoutingReadBatch`).
"""
# get list of potential deployments
verbose_router_logger.debug(
@ -473,6 +488,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
tpm_keys, rpm_keys = self.usage_counter_keys(healthy_deployments)
combined_tpm_rpm_keys: Final = tpm_keys + rpm_keys
prefetched_usage: Final = PrefetchedUsage.active()
if prefetched_usage is not None and prefetched_usage.covers(combined_tpm_rpm_keys):
combined_tpm_rpm_values = prefetched_usage.values_for(combined_tpm_rpm_keys)
else:

View file

@ -36,7 +36,8 @@ Safe to enable globally:
- No cache required.
"""
from collections.abc import Iterator, Mapping
from collections.abc import Iterator, Mapping, Sequence
from functools import cache
from typing import TYPE_CHECKING, Final, Optional, cast
from litellm._logging import verbose_router_logger
@ -114,23 +115,31 @@ class EncryptedContentAffinityCheck(CustomLogger):
if not isinstance(request_input, list):
return None
for item in request_input:
if not isinstance(item, dict):
continue
return next(
(
model_id
for item in request_input
if (model_id := EncryptedContentAffinityCheck._model_id_of_input_item(item)) is not None
),
None,
)
# First, try to decode from item ID (if present)
item_id = item.get("id")
if item_id and isinstance(item_id, str):
decoded = ResponsesAPIRequestUtils._decode_encrypted_item_id(item_id)
if decoded:
return decoded.get("model_id")
@staticmethod
def _model_id_of_input_item(item: object) -> str | None:
if not isinstance(item, dict):
return None
# If no encoded ID, check if encrypted_content itself is wrapped
encrypted_content = item.get("encrypted_content")
if encrypted_content and isinstance(encrypted_content, str):
model_id = EncryptedContentAffinityCheck._model_id_from_wrapped_encrypted_content(encrypted_content)
if model_id:
return model_id
item_id: Final = item.get("id")
if item_id and isinstance(item_id, str):
decoded: Final = ResponsesAPIRequestUtils._decode_encrypted_item_id(item_id)
if decoded:
return decoded.get("model_id")
encrypted_content: Final = item.get("encrypted_content")
if encrypted_content and isinstance(encrypted_content, str):
model_id: Final = EncryptedContentAffinityCheck._model_id_from_wrapped_encrypted_content(encrypted_content)
if model_id:
return model_id
return None
@ -150,19 +159,20 @@ class EncryptedContentAffinityCheck(CustomLogger):
model_id, _ = ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(encrypted_content)
return model_id or None
@staticmethod
def _model_id_of_anthropic_block(block: Mapping[str, object]) -> str | None:
encrypted_content: Final = encrypted_content_of_block(block)
if encrypted_content is None:
return None
return EncryptedContentAffinityCheck._model_id_from_wrapped_encrypted_content(encrypted_content)
@staticmethod
def _extract_model_id_from_anthropic_messages(messages: object) -> str | None:
return next(
(
model_id
for block in EncryptedContentAffinityCheck._anthropic_content_blocks(messages)
if (encrypted_content := encrypted_content_of_block(block)) is not None
if (
model_id := EncryptedContentAffinityCheck._model_id_from_wrapped_encrypted_content(
encrypted_content
)
)
is not None
if (model_id := EncryptedContentAffinityCheck._model_id_of_anthropic_block(block)) is not None
),
None,
)
@ -243,6 +253,50 @@ class EncryptedContentAffinityCheck(CustomLogger):
]
return matches, originating
def _strip_reasoning_the_target_cannot_decrypt(
self,
request_input: object,
anthropic_messages: object,
target_deployments: Sequence[Mapping[str, object]],
) -> None:
target_ids: Final = frozenset(
str(model_info["id"])
for target in target_deployments
if isinstance((model_info := target.get("model_info")), Mapping) and model_info.get("id") is not None
)
target_boundaries: Final = frozenset(
boundary
for target in target_deployments
if (boundary := self._encryption_boundary_key(target.get("litellm_params"))) is not None
)
@cache
def target_can_decrypt(origin_model_id: str) -> bool:
if origin_model_id in target_ids:
return True
if self.router is None:
return False
origin: Final = self.router.get_deployment(model_id=origin_model_id)
origin_boundary: Final = (
self._encryption_boundary_key(origin.litellm_params.model_dump(exclude_none=True))
if origin is not None
else None
)
return origin_boundary is not None and origin_boundary in target_boundaries
def should_strip_input_item(item: Mapping[str, object]) -> bool:
origin_model_id: Final = self._model_id_of_input_item(item)
return origin_model_id is not None and not target_can_decrypt(origin_model_id)
def should_strip_anthropic_block(block: Mapping[str, object]) -> bool:
origin_model_id: Final = self._model_id_of_anthropic_block(block)
return origin_model_id is not None and not target_can_decrypt(origin_model_id)
ResponsesAPIRequestUtils.strip_encrypted_reasoning_from_input(
request_input, should_strip=should_strip_input_item
)
strip_encrypted_reasoning_from_messages(anthropic_messages, should_strip=should_strip_anthropic_block)
# ------------------------------------------------------------------
# Request routing (pre-call filter)
# ------------------------------------------------------------------
@ -303,6 +357,7 @@ class EncryptedContentAffinityCheck(CustomLogger):
model_id,
)
request_kwargs["_encrypted_content_affinity_pinned"] = True
self._strip_reasoning_the_target_cannot_decrypt(request_input, anthropic_messages, (deployment,))
return [deployment]
# Follow-up switched model_name (LIT-2531): pin by Azure resource instead.
@ -318,6 +373,7 @@ class EncryptedContentAffinityCheck(CustomLogger):
len(boundary_matches),
)
request_kwargs["_encrypted_content_affinity_pinned"] = True
self._strip_reasoning_the_target_cannot_decrypt(request_input, anthropic_messages, boundary_matches)
return boundary_matches
# The origin cannot serve this turn and no peer shares its encryption boundary, so its

View file

@ -8,11 +8,14 @@ different objects. `RoutingReadBatch` fetches both key sets in one
the usage slice to the strategy, so selection does not read again.
"""
import asyncio
import itertools
from collections.abc import Mapping, Sequence
from collections.abc import Iterator, Mapping, Sequence
from contextlib import contextmanager
from contextvars import ContextVar
from dataclasses import dataclass
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final
from typing import TYPE_CHECKING, Final
from litellm._logging import verbose_router_logger
from litellm.caching.dual_cache import DualCache
@ -21,31 +24,66 @@ from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2
from litellm.router_utils.cooldown_cache import CooldownCache
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
from opentelemetry.trace import Span
from litellm.router import Router as _Router
LitellmRouter = _Router
Span = _Span
else:
LitellmRouter = Any
Span = Any
from litellm.router import Router
_PREFETCH_SLOT: Final = "routing_read"
async def _backfill_prefetched_cache(
cache: DualCache,
due_keys: tuple[str, ...],
values: Mapping[str, object],
) -> None:
cache_keys: Final = list(due_keys) # mutable-ok: _prepare_batch_get takes a list
prepare_batch_get: Final = cache._prepare_batch_get # pyright: ignore[reportPrivateUsage] # memory backfill
pending: Final = await prepare_batch_get(cache_keys, local_only=True)
redis_values: Final = { # mutable-ok: _apply_batch_get accepts a dictionary
key: values[key]
for key, local in zip(due_keys, pending.result)
if local is None and values.get(key) is not None
}
apply_batch_get: Final = cache._apply_batch_get # pyright: ignore[reportPrivateUsage] # cache backfill
await apply_batch_get(pending, redis_values)
@dataclass(frozen=True, slots=True)
class RoutingPrefetch:
"""The cooldown and usage keys of a model group, declared on the request's Redis batch before admission
flushes it, so the routing read rides the same round trip as the rate limiter's Lua calls."""
keys: frozenset[str]
fetched: frozenset[str]
result: BatchResult[Mapping[str, object]]
reservations: tuple[tuple[DualCache, tuple[str, ...], dict[str, float | None]], ...]
def release(self) -> None:
for cache, _, previous_access_times in self.reservations:
cache._rollback_redis_batch_key_reservations( # pyright: ignore[reportPrivateUsage] # rollback
previous_access_times
)
async def _settle(self, future: asyncio.Future[Mapping[str, object]]) -> None:
if future.cancelled():
self.release()
return
if future.exception() is not None:
self.release()
return
values: Final = future.result()
try:
for cache, due_keys, _ in self.reservations:
await _backfill_prefetched_cache(cache, due_keys, values)
except Exception:
self.release()
raise
@staticmethod
def arm(
litellm_router_instance: LitellmRouter,
litellm_router_instance: "Router",
usage_selector: LowestTPMLoggingHandler_v2 | None,
deployments: list,
) -> None:
@ -60,9 +98,28 @@ class RoutingPrefetch:
() if usage_selector is None else tuple(itertools.chain(*usage_selector.usage_counter_keys(deployments)))
)
keys: Final = (*cooldown_keys, *usage_keys)
request.prefetched[_PREFETCH_SLOT] = RoutingPrefetch(
keys=frozenset(keys), result=request.batch(redis_cache).mget(keys)
cooldown_store: Final = litellm_router_instance.cooldown_cache.cooldown_store
cooldown_due, cooldown_previous = cooldown_store.reserve_redis_batch_reads(cooldown_keys)
usage_cache: Final = None if usage_selector is None else usage_selector.router_cache
usage_reservation: Final = None if usage_cache is None else usage_cache.reserve_redis_batch_reads(usage_keys)
usage_due: Final = () if usage_reservation is None else tuple(usage_reservation[0])
due: Final = (*cooldown_due, *usage_due)
reservations: Final = (
(cooldown_store, tuple(cooldown_due), cooldown_previous),
*(
()
if usage_cache is None or usage_reservation is None
else ((usage_cache, usage_due, usage_reservation[1]),)
),
)
if not due:
return
result: Final = request.batch(redis_cache).mget(due)
prefetch: Final = RoutingPrefetch(
keys=frozenset(keys), fetched=frozenset(due), result=result, reservations=reservations
)
result.on_settled(prefetch._settle)
request.prefetched[_PREFETCH_SLOT] = prefetch
@staticmethod
def armed() -> bool:
@ -78,14 +135,34 @@ class RoutingPrefetch:
armed: Final = request.prefetched.pop(_PREFETCH_SLOT, None)
if isinstance(armed, RoutingPrefetch) and armed.keys.issuperset(needed):
return armed
if isinstance(armed, RoutingPrefetch):
armed.release()
return None
_active_routing_read_batch: Final[ContextVar["RoutingReadBatch | None"]] = ContextVar(
"routing_read_batch", default=None
)
class RoutingReadBatch:
def __init__(self, usage_selector: LowestTPMLoggingHandler_v2 | None) -> None:
self.usage_selector: Final = usage_selector
self.prefetched_usage: PrefetchedUsage | None = None
@staticmethod
@contextmanager
def scoped(batch: "RoutingReadBatch | None") -> Iterator[None]:
token: Final = _active_routing_read_batch.set(batch)
try:
yield
finally:
_active_routing_read_batch.reset(token)
@staticmethod
def active() -> "RoutingReadBatch | None":
return _active_routing_read_batch.get()
@staticmethod
def for_strategy(strategy: str | None, selector: object) -> "RoutingReadBatch | None":
"""Usage-based routing reads its counters with the cooldown state; every other strategy reads only the
@ -97,9 +174,9 @@ class RoutingReadBatch:
async def async_get_cooldown_deployments(
self,
litellm_router_instance: LitellmRouter,
litellm_router_instance: "Router",
healthy_deployments: list,
parent_otel_span: Span | None,
parent_otel_span: "Span | None",
) -> list[str]:
"""
`_async_get_cooldown_deployments`, with the strategy's tpm/rpm counters for
@ -107,19 +184,23 @@ class RoutingReadBatch:
"""
model_ids: Final = litellm_router_instance.get_model_ids()
cooldown_keys: Final = [CooldownCache.get_cooldown_cache_key(model_id) for model_id in model_ids]
reads: Final[list[tuple[DualCache, list[str]]]] = [ # mutable-ok: the usage read is appended below
(litellm_router_instance.cooldown_cache.cooldown_store, cooldown_keys)
]
usage_keys: list[str] = [] # mutable-ok: DualCache batch reads take a list
if self.usage_selector is not None:
tpm_keys, rpm_keys = self.usage_selector.usage_counter_keys(healthy_deployments)
usage_keys = tpm_keys + rpm_keys
reads.append((self.usage_selector.router_cache, usage_keys))
selector: Final = self.usage_selector
usage_keys: Final = (
() if selector is None else tuple(itertools.chain(*selector.usage_counter_keys(healthy_deployments)))
)
reads: Final = (
(litellm_router_instance.cooldown_cache.cooldown_store, cooldown_keys),
*(
()
if selector is None
else ((selector.router_cache, list(usage_keys)),) # mutable-ok: DualCache batch reads take a list
),
)
results: Final = await self._read_prefetched(reads) or await DualCache.async_batch_get_cache_shared(
reads, parent_otel_span=parent_otel_span
)
cooldown_results: Final = results[0]
if self.usage_selector is not None:
if selector is not None:
usage_values: Final = results[1]
self.prefetched_usage = PrefetchedUsage(
keys=frozenset(usage_keys),
@ -134,7 +215,7 @@ class RoutingReadBatch:
@staticmethod
async def _read_prefetched(
reads: list[tuple[DualCache, list[str]]],
reads: Sequence[tuple[DualCache, list[str]]],
) -> list[list[object | None] | None] | None:
"""Serve the reads from the request's armed `RoutingPrefetch`, backfilling each cache's memory tier as
its own batch read would. None when nothing usable was armed or the prefetch failed."""
@ -149,6 +230,10 @@ class RoutingReadBatch:
results: Final[list[list[object | None] | None]] = [] # mutable-ok: filled per read below
for cache, keys in reads:
pending = await cache._prepare_batch_get(keys, local_only=True) # pyright: ignore[reportPrivateUsage] # same two-step read as async_batch_get_cache_shared
if any(
key not in prefetch.fetched for key, local_value in zip(keys, pending.result) if local_value is None
):
return None
missed = { # mutable-ok: _apply_batch_get takes a dict
key: values.get(key) for key, local in zip(keys, pending.result) if local is None
}

View file

@ -7,6 +7,10 @@ from pydantic import BaseModel, ConfigDict, Field, PrivateAttr, StrictInt, field
from typing_extensions import ReadOnly, Required, TypedDict
from litellm.types.llms.base import LiteLLMPydanticObjectBase
from litellm.types.proxy.agent_identity import (
AgentExecutionMode,
AgentIdentityBinding,
)
if TYPE_CHECKING:
from a2a.types import SendMessageResponse
@ -301,6 +305,11 @@ class AgentKeySummary(BaseModel):
class AgentResponse(BaseModel):
identity: AgentIdentityBinding | None = None
identity_managed: bool = False
enabled: bool = True
execution_mode: AgentExecutionMode = "autonomous"
jwt_auth_configured: bool = False
agent_id: str
agent_name: str
litellm_params: dict[str, object] | None = None

View file

@ -199,6 +199,7 @@ class ObservabilityOptions:
logger_fn: Callable[[Mapping[str, object]], None] | None = None
verbose: bool | None = None
no_log: bool | None = field(default=None, metadata=wire("no-log"))
log_client_error_tracebacks: bool | None = None
@dataclass(frozen=True, slots=True, kw_only=True)

View file

@ -774,6 +774,8 @@ ANTHROPIC_TOOL_SEARCH_TOOL_TYPES: Final = frozenset(
# Effort beta header constant
ANTHROPIC_EFFORT_BETA_HEADER: Final = "effort-2025-11-24"
ANTHROPIC_MID_CONVERSATION_OUTPUT_CONFIG_BETA_HEADER: Final = "mid-conversation-output-config-2026-07-01"
ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER: Final = "fine-grained-tool-streaming-2025-05-14"
# OAuth constants

View file

@ -0,0 +1,96 @@
from datetime import datetime
from typing import Literal, TypeAlias
from uuid import UUID
from pydantic import BaseModel, ConfigDict, Field, field_validator
AgentExecutionMode: TypeAlias = Literal["autonomous", "delegated", "both"]
class EntraIdentityConfig(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
provider: Literal["microsoft_entra"]
tenant_id: str
client_id: str
service_principal_id: str | None = None
required_roles: tuple[str, ...] = ()
required_scopes: tuple[str, ...] = Field(
default=("user_impersonation",),
description="Required delegated scopes. An empty list accepts any nonempty scope granted for this gateway.",
)
@field_validator("tenant_id", "client_id", "service_principal_id")
@classmethod
def normalize_identifier(cls, value: str | None) -> str | None:
return str(UUID(value)) if value is not None else None
@property
def issuer(self) -> str:
return f"https://login.microsoftonline.com/{self.tenant_id}/v2.0"
class AgentIdentityBinding(BaseModel):
model_config = ConfigDict(frozen=True)
agent_id: str
active: bool = True
provider: Literal["microsoft_entra"]
tenant_id: str
client_id: str
service_principal_id: str | None = None
issuer: str
required_roles: tuple[str, ...] = ()
required_scopes: tuple[str, ...] = ("user_impersonation",)
revision: str
last_authenticated_at: datetime | None = None
class AgentSubject(BaseModel):
model_config = ConfigDict(frozen=True)
kind: Literal["application", "delegated_subject"]
oid: str
mode: Literal["autonomous", "delegated"]
class AgentIdentityFailure(BaseModel):
model_config = ConfigDict(frozen=True)
code: Literal["identity_denied", "policy_unavailable"] = "identity_denied"
message: str
class ManagedAgentContext(BaseModel):
model_config = ConfigDict(frozen=True)
agent_id: str
binding_revision: str | None = None
mode: Literal["autonomous", "delegated"]
user_id: str | None = None
subject_oid: str | None = None
class VerifiedHumanSubject(BaseModel):
model_config = ConfigDict(frozen=True)
issuer: str
tenant_id: str
oid: str
user_id: str
class MicrosoftInteractiveSubject(BaseModel):
model_config = ConfigDict(frozen=True)
issuer: str
tenant_id: str
oid: str
class ManagedAgentIdentityStatus(BaseModel):
identity: AgentIdentityBinding | None = None
identity_managed: bool = False
enabled: bool = True
execution_mode: AgentExecutionMode = "autonomous"
last_authenticated_at: datetime | None = None

View file

@ -284,6 +284,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
cache_creation_input_token_cost_above_272k_tokens: float | None
cache_creation_input_token_cost_above_272k_tokens_priority: float | None
cache_creation_input_token_cost_above_272k_tokens_flex: float | None
cache_creation_input_token_cost_above_272k_tokens_ultrafast: ReadOnly[float | None]
cache_creation_input_token_cost_above_1hr: float | None
cache_creation_input_token_cost_flex: float | None # OpenAI flex service tier pricing
cache_creation_input_token_cost_priority: float | None # OpenAI priority service tier pricing
@ -300,6 +301,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
cache_read_input_token_cost_above_272k_tokens: float | None
cache_read_input_token_cost_above_272k_tokens_priority: float | None
cache_read_input_token_cost_above_272k_tokens_flex: float | None
cache_read_input_token_cost_above_272k_tokens_ultrafast: ReadOnly[float | None]
cache_read_input_token_cost_above_512k_tokens: float | None
cache_read_input_token_cost_batches: ReadOnly[float | None]
cache_read_input_token_cost_above_200k_tokens_batches: ReadOnly[float | None]
@ -319,6 +321,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
input_cost_per_token_above_272k_tokens: float | None # GPT-5.4/5.4-pro: prompts >272K priced at 2x input
input_cost_per_token_above_272k_tokens_priority: float | None
input_cost_per_token_above_272k_tokens_flex: float | None
input_cost_per_token_above_272k_tokens_ultrafast: ReadOnly[float | None]
input_cost_per_token_above_512k_tokens: float | None # MiniMax-M3: prompts >512K priced at 2x input
input_cost_per_character_above_128k_tokens: float | None # only for vertex ai models
input_cost_per_query: float | None # per-request pricing: rerank, search, and Bedrock Marengo embeddings
@ -360,6 +363,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
output_cost_per_token_above_272k_tokens: float | None # GPT-5.4/5.4-pro: prompts >272K priced at 1.5x output
output_cost_per_token_above_272k_tokens_priority: float | None
output_cost_per_token_above_272k_tokens_flex: float | None
output_cost_per_token_above_272k_tokens_ultrafast: ReadOnly[float | None]
output_cost_per_token_above_512k_tokens: float | None # MiniMax-M3: prompts >512K priced at 2x output
output_cost_per_character_above_128k_tokens: float | None # only for vertex ai models
output_cost_per_image: float | None
@ -3737,6 +3741,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
cache_creation_input_token_cost_above_272k_tokens: float | None = None
cache_creation_input_token_cost_above_272k_tokens_priority: float | None = None
cache_creation_input_token_cost_above_272k_tokens_flex: float | None = None
cache_creation_input_token_cost_above_272k_tokens_ultrafast: float | None = None
cache_creation_input_token_cost_flex: float | None = None
cache_creation_input_token_cost_priority: float | None = None
cache_creation_input_token_cost_ultrafast: float | None = None
@ -3749,6 +3754,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
cache_read_input_token_cost_above_200k_tokens_priority: float | None = None
cache_read_input_token_cost_above_272k_tokens_priority: float | None = None
cache_read_input_token_cost_above_272k_tokens_flex: float | None = None
cache_read_input_token_cost_above_272k_tokens_ultrafast: float | None = None
cache_read_input_token_cost_batches: float | None = None
cache_read_input_token_cost_above_200k_tokens_batches: float | None = None
cache_read_input_token_cost_above_272k_tokens_batches: float | None = None
@ -3765,6 +3771,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
input_cost_per_token_above_200k_tokens_priority: float | None = None
input_cost_per_token_above_272k_tokens_priority: float | None = None
input_cost_per_token_above_272k_tokens_flex: float | None = None
input_cost_per_token_above_272k_tokens_ultrafast: float | None = None
input_cost_per_token_above_200k_tokens_batches: float | None = None
input_cost_per_token_above_272k_tokens_batches: float | None = None
input_cost_per_query: float | None = None
@ -3791,6 +3798,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
output_cost_per_token_above_200k_tokens_priority: float | None = None
output_cost_per_token_above_272k_tokens_priority: float | None = None
output_cost_per_token_above_272k_tokens_flex: float | None = None
output_cost_per_token_above_272k_tokens_ultrafast: float | None = None
output_cost_per_token_above_200k_tokens_batches: float | None = None
output_cost_per_token_above_272k_tokens_batches: float | None = None
output_cost_per_character_above_128k_tokens: float | None = None

View file

@ -6104,6 +6104,9 @@ def _get_model_info_helper(
cache_creation_input_token_cost_above_272k_tokens_flex=_model_info.get(
"cache_creation_input_token_cost_above_272k_tokens_flex", None
),
cache_creation_input_token_cost_above_272k_tokens_ultrafast=_model_info.get(
"cache_creation_input_token_cost_above_272k_tokens_ultrafast", None
),
cache_creation_input_token_cost_flex=_model_info.get("cache_creation_input_token_cost_flex", None),
cache_creation_input_token_cost_priority=_model_info.get(
"cache_creation_input_token_cost_priority", None
@ -6129,6 +6132,9 @@ def _get_model_info_helper(
cache_read_input_token_cost_above_272k_tokens_flex=_model_info.get(
"cache_read_input_token_cost_above_272k_tokens_flex", None
),
cache_read_input_token_cost_above_272k_tokens_ultrafast=_model_info.get(
"cache_read_input_token_cost_above_272k_tokens_ultrafast", None
),
cache_read_input_token_cost_above_512k_tokens=_model_info.get(
"cache_read_input_token_cost_above_512k_tokens", None
),
@ -6167,6 +6173,9 @@ def _get_model_info_helper(
input_cost_per_token_above_272k_tokens_flex=_model_info.get(
"input_cost_per_token_above_272k_tokens_flex", None
),
input_cost_per_token_above_272k_tokens_ultrafast=_model_info.get(
"input_cost_per_token_above_272k_tokens_ultrafast", None
),
input_cost_per_token_above_512k_tokens=_model_info.get("input_cost_per_token_above_512k_tokens", None),
input_cost_per_query=_model_info.get("input_cost_per_query", None),
cost_per_second=_model_info.get("cost_per_second", None),
@ -6234,6 +6243,9 @@ def _get_model_info_helper(
output_cost_per_token_above_272k_tokens_flex=_model_info.get(
"output_cost_per_token_above_272k_tokens_flex", None
),
output_cost_per_token_above_272k_tokens_ultrafast=_model_info.get(
"output_cost_per_token_above_272k_tokens_ultrafast", None
),
output_cost_per_token_above_512k_tokens=_model_info.get(
"output_cost_per_token_above_512k_tokens", None
),

View file

@ -3358,7 +3358,7 @@
"supports_function_calling": true
},
"azure_ai/claude-haiku-4-5": {
"deprecation_date": "2026-10-19",
"deprecation_date": "2026-11-15",
"cache_creation_input_token_cost": 1.25e-06,
"cache_creation_input_token_cost_above_1hr": 2e-06,
"cache_read_input_token_cost": 1e-07,
@ -3378,10 +3378,11 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"prompt_cache_min_tokens": 4096
"prompt_cache_min_tokens": 4096,
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/model-retirement-schedule"
},
"azure_ai/claude-opus-4-5": {
"deprecation_date": "2026-10-19",
"deprecation_date": "2026-11-24",
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -3402,7 +3403,8 @@
"supports_tool_choice": true,
"supports_vision": true,
"supports_output_config": true,
"prompt_cache_min_tokens": 4096
"prompt_cache_min_tokens": 4096,
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/model-retirement-schedule"
},
"azure_ai/claude-opus-4-6": {
"deprecation_date": "2027-02-02",
@ -3640,7 +3642,7 @@
"prompt_cache_min_tokens": 1024
},
"azure_ai/claude-sonnet-4-5": {
"deprecation_date": "2026-10-19",
"deprecation_date": "2026-11-15",
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
@ -3660,7 +3662,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"prompt_cache_min_tokens": 1024
"prompt_cache_min_tokens": 1024,
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/model-retirement-schedule"
},
"azure_ai/claude-sonnet-5": {
"deprecation_date": "2027-06-30",
@ -30721,6 +30724,7 @@
"output_cost_per_image": 0.08
},
"gemini/veo-3.1-fast-generate-preview": {
"deprecation_date": "2026-10-22",
"litellm_provider": "gemini",
"max_input_tokens": 1024,
"max_tokens": 1024,
@ -30737,6 +30741,7 @@
]
},
"gemini/veo-3.1-generate-preview": {
"deprecation_date": "2026-10-22",
"litellm_provider": "gemini",
"max_input_tokens": 1024,
"max_tokens": 1024,
@ -30752,6 +30757,7 @@
]
},
"gemini/veo-3.1-lite-generate-preview": {
"deprecation_date": "2026-10-22",
"litellm_provider": "gemini",
"max_input_tokens": 1024,
"max_tokens": 1024,
@ -32912,10 +32918,13 @@
"gpt-image-2.5-flare": {
"cache_read_input_image_token_cost": 2e-06,
"cache_read_input_token_cost": 1.25e-06,
"cache_read_input_token_cost_batches": 6.25e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "openai",
"mode": "image_generation",
"input_cost_per_image_token": 8e-06,
"input_cost_per_image_token_batches": 4e-06,
"input_cost_per_token_batches": 2.5e-06,
"output_cost_per_image_token": 3e-05,
"supported_endpoints": [
"/v1/images/generations",
@ -32944,10 +32953,13 @@
"gpt-image-2.5-sunburst": {
"cache_read_input_image_token_cost": 2e-06,
"cache_read_input_token_cost": 1.25e-06,
"cache_read_input_token_cost_batches": 6.25e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "openai",
"mode": "image_generation",
"input_cost_per_image_token": 8e-06,
"input_cost_per_image_token_batches": 4e-06,
"input_cost_per_token_batches": 2.5e-06,
"output_cost_per_image_token": 3e-05,
"supported_endpoints": [
"/v1/images/generations",
@ -38611,6 +38623,7 @@
},
"mistral/zai-glm-5-2": {
"cache_read_input_token_cost": 1.4e-07,
"deprecation_date": "2026-10-31",
"input_cost_per_token": 1.4e-06,
"litellm_provider": "mistral",
"max_input_tokens": 1048576,
@ -38741,6 +38754,7 @@
"source": "https://mistral.ai/pricing#api-pricing"
},
"mistral/mistral-ocr-4-0": {
"deprecation_date": "2026-09-30",
"litellm_provider": "mistral",
"ocr_cost_per_page": 0.004,
"ocr_cost_per_page_batches": 0.002,
@ -60616,6 +60630,7 @@
"supports_vision": true
},
"mistral/labs-leanstral-1-5": {
"deprecation_date": "2026-09-30",
"input_cost_per_token": 0.0,
"litellm_provider": "mistral",
"max_input_tokens": 262144,
@ -61337,13 +61352,16 @@
},
"fireworks_ai/nemotron-lightning-3p5-30b-a3b": {
"cache_read_input_token_cost": 1e-08,
"cache_read_input_token_cost_priority": 1.25e-08,
"input_cost_per_token": 5e-08,
"input_cost_per_token_priority": 6.25e-08,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 262144,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 2e-07,
"output_cost_per_token_priority": 2.5e-07,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -61353,13 +61371,16 @@
},
"fireworks_ai/nemotron-3-ultra-nvfp4": {
"cache_read_input_token_cost": 1.2e-07,
"cache_read_input_token_cost_priority": 1.5e-07,
"input_cost_per_token": 6e-07,
"input_cost_per_token_priority": 7.5e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 262144,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 2.4e-06,
"output_cost_per_token_priority": 3e-06,
"source": "https://api.fireworks.ai/v1/serverless/models",
"supports_function_calling": true,
"supports_reasoning": true,
@ -61389,13 +61410,16 @@
},
"fireworks_ai/accounts/fireworks/models/nemotron-lightning-3p5-30b-a3b": {
"cache_read_input_token_cost": 1e-08,
"cache_read_input_token_cost_priority": 1.25e-08,
"input_cost_per_token": 5e-08,
"input_cost_per_token_priority": 6.25e-08,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 262144,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 2e-07,
"output_cost_per_token_priority": 2.5e-07,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -61405,13 +61429,16 @@
},
"fireworks_ai/accounts/fireworks/models/nemotron-3-ultra-nvfp4": {
"cache_read_input_token_cost": 1.2e-07,
"cache_read_input_token_cost_priority": 1.5e-07,
"input_cost_per_token": 6e-07,
"input_cost_per_token_priority": 7.5e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 262144,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 2.4e-06,
"output_cost_per_token_priority": 3e-06,
"source": "https://api.fireworks.ai/v1/serverless/models",
"supports_function_calling": true,
"supports_reasoning": true,
@ -64321,13 +64348,16 @@
},
"fireworks_ai/accounts/fireworks/routers/glm-5p3-us": {
"cache_read_input_token_cost": 3.9e-07,
"cache_read_input_token_cost_priority": 4.875e-07,
"input_cost_per_token": 2.1e-06,
"input_cost_per_token_priority": 2.625e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 6.6e-06,
"output_cost_per_token_priority": 8.25e-06,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -64356,13 +64386,16 @@
},
"fireworks_ai/glm-5p3-us": {
"cache_read_input_token_cost": 3.9e-07,
"cache_read_input_token_cost_priority": 4.875e-07,
"input_cost_per_token": 2.1e-06,
"input_cost_per_token_priority": 2.625e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 6.6e-06,
"output_cost_per_token_priority": 8.25e-06,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -64446,12 +64479,15 @@
},
"fireworks_ai/accounts/fireworks/routers/glm-5p3-flash-us": {
"cache_read_input_token_cost": 4.5e-08,
"cache_read_input_token_cost_priority": 5.625e-08,
"input_cost_per_token": 2.25e-07,
"input_cost_per_token_priority": 2.8125e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_tokens": 1048576,
"mode": "chat",
"output_cost_per_token": 7.5e-07,
"output_cost_per_token_priority": 9.375e-07,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_response_schema": true,
@ -64477,12 +64513,15 @@
},
"fireworks_ai/glm-5p3-flash-us": {
"cache_read_input_token_cost": 4.5e-08,
"cache_read_input_token_cost_priority": 5.625e-08,
"input_cost_per_token": 2.25e-07,
"input_cost_per_token_priority": 2.8125e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_tokens": 1048576,
"mode": "chat",
"output_cost_per_token": 7.5e-07,
"output_cost_per_token_priority": 9.375e-07,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_response_schema": true,
@ -64571,6 +64610,7 @@
"source": "https://api.together.ai/v1/models"
},
"together_ai/nvidia/NVIDIA-Nemotron-Nano-9B-v2": {
"deprecation_date": "2026-02-25",
"input_cost_per_token": 6e-08,
"output_cost_per_token": 2.5e-07,
"litellm_provider": "together_ai",
@ -70451,6 +70491,7 @@
},
"together_ai/nvidia/nemotron-3-ultra-550b-a55b": {
"cache_read_input_token_cost": 2e-07,
"deprecation_date": "2026-08-27",
"input_cost_per_token": 6e-07,
"litellm_provider": "together_ai",
"max_input_tokens": 512288,
@ -77476,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,
@ -79218,12 +79262,12 @@
"input_cost_per_token_above_272k_tokens": 4e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_above_272k_tokens": 1.5e-05,
"source": "https://developers.openai.com/api/docs/pricing",
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-1-sol.html",
"supported_endpoints": [
"/v1/responses"
],
@ -79253,12 +79297,12 @@
"input_cost_per_token_above_272k_tokens": 4e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_above_272k_tokens": 1.5e-05,
"source": "https://developers.openai.com/api/docs/pricing",
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-1-sol.html",
"supported_modalities": [
"text",
"image"
@ -79285,12 +79329,12 @@
"input_cost_per_token_above_272k_tokens": 4.4e-06,
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "responses",
"output_cost_per_token": 1.1e-05,
"output_cost_per_token_above_272k_tokens": 1.65e-05,
"source": "https://developers.openai.com/api/docs/models/gpt-6.1-sol",
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-1-sol.html",
"supported_endpoints": [
"/v1/chat/completions",
"/v1/responses"
@ -79323,12 +79367,12 @@
"input_cost_per_token_above_272k_tokens": 4.4e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1.1e-05,
"output_cost_per_token_above_272k_tokens": 1.65e-05,
"source": "https://developers.openai.com/api/docs/pricing",
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-1-sol.html",
"supported_endpoints": [
"/v1/responses"
],
@ -79348,5 +79392,33 @@
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true
},
"vertex_ai/gemini-3.8-flash-tts": {
"input_cost_per_token": 5e-07,
"litellm_provider": "vertex_ai",
"max_input_tokens": 8192,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "audio_speech",
"output_cost_per_audio_token": 9e-06,
"output_cost_per_token": 9e-06,
"source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing",
"supported_endpoints": [
"/v1/audio/speech"
]
},
"vertex_ai/gemini-3.8-flash-lite-tts": {
"input_cost_per_token": 5e-07,
"litellm_provider": "vertex_ai",
"max_input_tokens": 8192,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "audio_speech",
"output_cost_per_audio_token": 6e-06,
"output_cost_per_token": 6e-06,
"source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing",
"supported_endpoints": [
"/v1/audio/speech"
]
}
}

View file

@ -7,3 +7,13 @@ reason = "diskcache has no fixed release published; remove this entry once one e
id = "GHSA-h7x2-h6g9-p789"
ignoreUntil = 2026-10-14
reason = "mlflow has no fixed release published (3.16.0, 2026-09-04, and master still store gateway secret api_base unvalidated); remove this entry once one exists"
[[IgnoredVulns]]
id = "GHSA-hj66-6f7g-4r5v"
ignoreUntil = 2026-10-02
reason = "oauthlib 4.0.0 (the only fixed release, 2026-09-28) is inside the 3-day uv exclude-newer cooldown; bump oauthlib and remove this entry once it clears"
[[IgnoredVulns]]
id = "GHSA-xpv3-w29h-x7cv"
ignoreUntil = 2026-10-02
reason = "oauthlib 4.0.0 (the only fixed release, 2026-09-28) is inside the 3-day uv exclude-newer cooldown; bump oauthlib and remove this entry once it clears"

View file

@ -1,6 +1,6 @@
[project]
name = "litellm"
version = "1.104.0"
version = "1.105.0"
description = "Library to easily interface with LLM API providers"
readme = "README.md"
requires-python = ">=3.10, <3.15"
@ -75,8 +75,8 @@ proxy = [
"mcp>=2.2.0,<3",
"httpx2>=2.5.0,<3",
"pydantic>=2.12.0,<3",
"litellm-proxy-extras==0.4.102",
"litellm-enterprise==0.1.71",
"litellm-proxy-extras==0.4.103",
"litellm-enterprise==0.1.72",
"RestrictedPython>=8.5,<9.0",
"rich>=13.9.4,<14.0",
"InquirerPy>=0.3.4,<1.0",
@ -357,7 +357,7 @@ litellm-enterprise = { workspace = true }
members = ["enterprise", "litellm-proxy-extras"]
[tool.commitizen]
version = "1.104.0"
version = "1.105.0"
version_files = [
"pyproject.toml:^version",
]

View file

@ -78,6 +78,11 @@ model LiteLLM_AgentsTable {
object_permission_id String?
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
spend Float @default(0.0)
identity_managed Boolean @default(false)
enabled Boolean @default(true)
execution_mode String @default("autonomous")
identity LiteLLM_AgentIdentity?
retired_identities LiteLLM_RetiredAgentIdentity[]
tpm_limit Int?
rpm_limit Int?
session_tpm_limit Int?
@ -88,6 +93,56 @@ model LiteLLM_AgentsTable {
updated_by String
}
model LiteLLM_AgentIdentity {
agent_id String @id
active Boolean @default(true)
agent LiteLLM_AgentsTable @relation(fields: [agent_id], references: [agent_id], onDelete: Cascade)
provider String
issuer String
tenant_id String
client_id String
service_principal_id String?
required_roles String[] @default([])
required_scopes String[] @default(["user_impersonation"])
revision String @default(uuid())
last_authenticated_at DateTime?
@@unique([provider, tenant_id, client_id])
@@unique([issuer, service_principal_id])
}
model LiteLLM_RetiredAgentIdentity {
binding_id String @id @default(uuid())
agent_id String?
agent LiteLLM_AgentsTable? @relation(fields: [agent_id], references: [agent_id], onDelete: SetNull)
provider String
issuer String
tenant_id String
client_id String
@@unique([provider, tenant_id, client_id])
}
model LiteLLM_RetiredAgent {
original_agent_id String @id
retired_at DateTime @default(now())
}
model LiteLLM_VerifiedSubject {
subject_id String @id @default(uuid())
issuer String
tenant_id String
oid String
kind String @default("human")
user_id String?
user LiteLLM_UserTable? @relation(fields: [user_id], references: [user_id], onDelete: Cascade)
verified_via String @default("sso_interactive")
verified_at DateTime @default(now())
@@unique([issuer, tenant_id, oid])
@@index([user_id])
}
model LiteLLM_OrganizationTable {
organization_id String @id @default(uuid())
organization_alias String
@ -241,6 +296,7 @@ model LiteLLM_DeletedTeamTable {
// Track spend, rate limit, budget Users
model LiteLLM_UserTable {
verified_subjects LiteLLM_VerifiedSubject[]
user_id String @id
user_alias String?
team_id String?
@ -675,6 +731,7 @@ model LiteLLM_SpendLogs {
session_id String?
status String?
mcp_namespaced_tool_name String?
billing_agent_id String?
agent_id String?
proxy_server_request Json? @default("{}")
litellm_call_id String?

View file

@ -64,7 +64,7 @@ The suites run against a live proxy, so bring one up first by running the litell
For the opt-in browser profile, start the existing IdP first, then run `.github/e2e-stack/oidc-profile.sh "$PROXY_BASE_URL" <server-command>`. The wrapper creates a confidential client with an exact `/sso/callback` redirect and S256 PKCE, passes the client secret only through the child process environment, and removes the client on exit. It uses the existing generic OIDC handler with `GENERIC_USER_ID_ATTRIBUTE=sub`. Preserve the IdP's PostgreSQL data across restarts
`tests/e2e/ui/playwright.oidc.config.ts` uses an already running OIDC stack and separate storage/output files. Supply `E2E_OIDC_UI_URL`, `JWT_ISSUER`, `E2E_OIDC_USERNAME` and `E2E_OIDC_PASSWORD` for a seeded actor. Its setup follows the real login and callback path. The current Python canary qualifies browser-client configuration and token/userinfo identity mapping; browser journey specs under `ui/oidc/` are a separate coverage step
`tests/e2e/ui/playwright.oidc.config.ts` uses an already running OIDC stack and separate storage/output files. Supply `E2E_OIDC_UI_URL`, `JWT_ISSUER`, `E2E_OIDC_USERNAME` and `E2E_OIDC_PASSWORD` for a seeded actor. Its setup follows the real login and callback path. The current Python canary qualifies browser-client configuration and token/userinfo identity mapping. The specs under `ui/oidc/` drive a real dashboard SSO login and a real `lite login`, so start the proxy with `EXPERIMENTAL_UI_LOGIN=true` and at least one model it can actually serve. The CLI spec runs `lite` from `PATH` unless `E2E_LITE_CLI` names another executable, and it gives the CLI a temporary `HOME` with the keyring disabled so your own login is never touched. The main `playwright.config.ts` ignores `oidc/`
Every successful IdP create immediately registers cleanup, including partial setup failures. Cleanup failures emit warnings. Tokens are minted on demand, and the expiration test waits relative to the token's actual `exp` with a bounded clock-drift check. To check first-attempt behavior locally, run both files with `--reruns 0`:

View file

@ -60,3 +60,6 @@
- {id: other.auth.jwt.wrong_issuer_denied, module: other, tier: P0, area: auth, assertions: [wrong_issuer_denied], source: "auth/handle_jwt.py", rationale: "A signed token with the correct audience and an unexpected issuer is rejected"}
- {id: other.auth.jwt.wrong_audience_denied, module: other, tier: P0, area: auth, assertions: [wrong_audience_denied], source: "auth/handle_jwt.py", rationale: "A signed token from the trusted issuer intended for another app is rejected"}
- {id: other.auth.session_token.valid_allows, module: other, tier: P0, area: auth, assertions: [valid_allows], source: "auth/auth_checks.py ExperimentalUIJWTToken", rationale: "An unexpired LiteLLM-minted session token authenticates with the role it carries"}
- {id: other.auth.session_token.expired_denied, module: other, tier: P0, area: auth, assertions: [expired_denied], source: "auth/user_api_key_auth.py expiry check", rationale: "An expired session token is rejected with the expired-key error"}
- {id: other.auth.session_token.encrypted_value_denied, module: other, tier: P0, area: auth, assertions: [encrypted_value_denied], source: "auth/auth_checks.py ExperimentalUIJWTToken", rationale: "An encrypted value read back from a management route is not accepted as a bearer token"}

View file

@ -0,0 +1,91 @@
"""Live e2e: UI/CLI session tokens are accepted only while valid and only when minted as session tokens.
The runner mints its own session tokens under the proxy's salt key, so the valid and expired cases run in
seconds instead of waiting out a real login's expiry.
"""
from __future__ import annotations
import base64
import hashlib
import json
import os
from datetime import datetime, timedelta, timezone
from typing import Final
import pytest
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
from e2e_config import MASTER_KEY, unique_marker
from e2e_http import UnauthorizedError, unwrap
from lifecycle import ResourceManager
from models import KeyGenerateBody, KeyLoggingCallback, KeyLoggingCallbackVars, KeyMetadata
from other_client import OtherClient
pytestmark = pytest.mark.e2e
SALT_KEY: Final = os.environ.get("LITELLM_SALT_KEY") or MASTER_KEY
SESSION_TOKEN_PREFIX: Final = "litellm_login_"
ENCRYPTED_PREFIX: Final = "litellm_enc::"
def _admin_session_token(expires_at: datetime) -> str:
claims: Final = json.dumps(
{
"token": f"ui-token-{unique_marker()}",
"user_id": f"e2e-session-{unique_marker()}",
"user_role": "proxy_admin",
"team_id": "litellm-dashboard",
"expires": expires_at.isoformat(),
}
)
nonce: Final = os.urandom(12)
sealed: Final = AESGCM(hashlib.sha256(SALT_KEY.encode()).digest()).encrypt(
nonce, claims.encode(), SESSION_TOKEN_PREFIX.encode()
)
return SESSION_TOKEN_PREFIX + base64.urlsafe_b64encode(nonce + sealed).decode().rstrip("=")
class TestSessionToken:
@pytest.mark.covers("other.auth.session_token.valid_allows")
def test_unexpired_session_token_reaches_admin_route(self, client: OtherClient) -> None:
token: Final = _admin_session_token(datetime.now(timezone.utc) + timedelta(minutes=10))
listing: Final = unwrap(client.list_users_as(token))
assert listing.total >= 0, f"an unexpired admin session token did not reach /user/list: {listing}"
@pytest.mark.covers("other.auth.session_token.expired_denied")
def test_expired_session_token_is_denied(self, client: OtherClient) -> None:
token: Final = _admin_session_token(datetime.now(timezone.utc) - timedelta(minutes=1))
result: Final = client.list_users_as(token)
assert isinstance(result, UnauthorizedError), f"an expired session token must get 401, got {result}"
assert "expired" in result.body.lower(), f"expected the expired-key error, got {result.body[:300]}"
@pytest.mark.covers("other.auth.session_token.encrypted_value_denied")
def test_encrypted_stored_value_is_not_a_bearer_token(
self, client: OtherClient, resources: ResourceManager
) -> None:
stored_value: Final = f'{{"token": "{unique_marker()}", "user_role": "proxy_admin"}}'
key: Final = client.proxy.generate_key(
KeyGenerateBody(
key_alias=f"e2e-session-{unique_marker()}",
metadata=KeyMetadata(
logging=[
KeyLoggingCallback(
callback_name="langfuse",
callback_vars=KeyLoggingCallbackVars(langfuse_secret_key=stored_value),
)
]
),
)
)
resources.defer(lambda: client.proxy.delete_key(key))
metadata: Final = client.proxy.key_info(key).metadata
assert metadata is not None and metadata.logging, f"/key/info dropped the logging metadata: {metadata}"
encrypted: Final = metadata.logging[0].callback_vars.langfuse_secret_key
assert encrypted is not None and encrypted.startswith(ENCRYPTED_PREFIX), (
f"expected /key/info to return the stored secret encrypted, got {encrypted!r}"
)
for bearer in (encrypted.removeprefix(ENCRYPTED_PREFIX), encrypted):
result = client.list_users_as(bearer)
assert isinstance(result, UnauthorizedError), f"an encrypted stored value must get 401, got {result}"

View file

@ -51,6 +51,21 @@ export async function sendChatCompletion(request: APIRequestContext, opts: ChatO
return body.id as string;
}
export interface ServedChat {
requestId: string;
callId: string;
}
export async function sendChatCompletionWithCallId(request: APIRequestContext, opts: ChatOptions): Promise<ServedChat> {
const res = await postChatCompletion(request, opts);
expect(res.ok(), `chat completion for ${opts.model} failed (${res.status()}): ${await res.text()}`).toBe(true);
const callId = res.headers()["x-litellm-call-id"];
expect(callId, "proxy did not return an x-litellm-call-id header").toBeTruthy();
const body = await res.json();
expect(body.choices?.[0]?.message?.content).toContain(MOCK_RESPONSE_TEXT);
return { requestId: body.id as string, callId };
}
export interface ChatAttempt {
status: number;
body: string;
@ -124,7 +139,7 @@ export async function waitForSpendLog(
lastStatus = res.status();
if (res.ok()) {
const body = await res.json();
const rows = Array.isArray(body) ? body : (body?.data ?? []);
const rows = Array.isArray(body) ? body : body?.data ?? [];
if (rows.length > 0) {
return;
}

View file

@ -0,0 +1,87 @@
import { expect, test } from "@playwright/test";
import { execFile, spawn } from "node:child_process";
import * as fs from "node:fs";
import * as os from "node:os";
import * as path from "node:path";
import { promisify } from "node:util";
const LITE_CLI = process.env.E2E_LITE_CLI ?? "lite";
const SKIP_TEAM_SELECTION = "skip\n";
const execFileAsync = promisify(execFile);
function requiredEnv(name: string): string {
const value = process.env[name];
if (!value) throw new Error(`${name} must be set for the OIDC suite`);
return value;
}
test("CLI SSO login stores a session that lists models and completes a chat request", async ({ browser, baseURL }) => {
test.setTimeout(180_000);
const issuer = requiredEnv("JWT_ISSUER");
const home = fs.mkdtempSync(path.join(os.tmpdir(), "lite-cli-login-"));
const browserUrlFile = path.join(home, "browser-url");
const browserCommand = path.join(home, "browser.sh");
fs.writeFileSync(browserCommand, `#!/bin/sh\nprintf '%s' "$1" > '${browserUrlFile}'\n`, { mode: 0o700 });
const env = {
...process.env,
HOME: home,
LITELLM_CLI_DISABLE_KEYRING: "1",
BROWSER: browserCommand,
PYTHONUNBUFFERED: "1",
FORCE_COLOR: undefined,
NO_COLOR: "1",
LITELLM_PROXY_URL: baseURL,
LITELLM_PROXY_API_KEY: undefined,
};
const login = spawn(LITE_CLI, ["login"], { env });
let loginOutput = "";
login.stdout.on("data", (chunk: Buffer) => (loginOutput += chunk.toString()));
login.stderr.on("data", (chunk: Buffer) => (loginOutput += chunk.toString()));
const loginExit = new Promise<number | null>((resolve) => login.on("close", resolve));
login.stdin.end(SKIP_TEAM_SELECTION);
try {
await expect.poll(() => fs.existsSync(browserUrlFile), { timeout: 30_000 }).toBe(true);
await expect.poll(() => loginOutput).toMatch(/Verification code: \S+/);
const userCode = /Verification code: (\S+)/.exec(loginOutput)?.[1] ?? "";
const context = await browser.newContext({ storageState: { cookies: [], origins: [] } });
try {
const page = await context.newPage();
await page.goto(fs.readFileSync(browserUrlFile, "utf8"));
await expect(page).toHaveURL((url) => url.href.startsWith(`${issuer}/`));
await page.getByLabel("Username or email").fill(requiredEnv("E2E_OIDC_USERNAME"));
await page.getByLabel("Password", { exact: true }).fill(requiredEnv("E2E_OIDC_PASSWORD"));
await page.getByRole("button", { name: "Sign In", exact: true }).click();
await page.getByLabel("Verification code").fill(userCode);
await page.getByRole("button", { name: "Continue", exact: true }).click();
await expect(page.getByRole("heading", { name: "Authentication Successful!" })).toBeVisible();
} finally {
await context.close();
}
expect(await loginExit, loginOutput).toBe(0);
expect(loginOutput).toContain("Login successful!");
const stored: { key?: unknown } = JSON.parse(fs.readFileSync(path.join(home, ".litellm", "token.json"), "utf8"));
expect(typeof stored.key).toBe("string");
expect(stored.key, "CLI login issues a session token, not a virtual key").not.toMatch(/^sk-/);
const { stdout: modelsJson } = await execFileAsync(LITE_CLI, ["models", "list", "--format", "json"], { env });
const models: { id: string }[] = JSON.parse(modelsJson);
expect(models.length, "the stack serves at least one model").toBeGreaterThan(0);
const chatRequest = JSON.stringify({
model: models[0].id,
messages: [{ role: "user", content: "Reply with the single word: ok" }],
});
const { stdout: completionJson } = await execFileAsync(
LITE_CLI,
["http", "request", "POST", "/chat/completions", "-j", chatRequest],
{ env },
);
const completion: { choices: { message: { content: string | null } }[] } = JSON.parse(completionJson);
expect(completion.choices[0]?.message.content).toBeTruthy();
} finally {
login.kill();
fs.rmSync(home, { recursive: true, force: true });
}
});

View file

@ -0,0 +1,35 @@
import { expect, test, type Page as PlaywrightPage, type Response } from "@playwright/test";
import { Page } from "../fixtures/pages";
import { navigateToPage } from "../helpers/navigation";
function sessionKey(tokenCookie: string): string {
const claims: unknown = JSON.parse(Buffer.from(tokenCookie.split(".")[1] ?? "", "base64url").toString("utf8"));
const key = claims !== null && typeof claims === "object" && "key" in claims ? claims.key : undefined;
if (typeof key !== "string") throw new Error("The dashboard token cookie carries no key claim");
return key;
}
async function openPageAndCapture(page: PlaywrightPage, target: Page, apiPath: string): Promise<Response> {
const response = page.waitForResponse((r) => new URL(r.url()).pathname === apiPath);
await navigateToPage(page, target);
return response;
}
test("SSO login issues a session that authorizes dashboard data requests", async ({ page, context, baseURL }) => {
const tokenCookie = (await context.cookies(baseURL)).find((cookie) => cookie.name === "token");
expect(tokenCookie, "SSO login sets the dashboard token cookie").toBeDefined();
const key = sessionKey(tokenCookie?.value ?? "");
expect(key, "SSO login issues a session token, not a virtual key").not.toMatch(/^sk-/);
const keyList = await openPageAndCapture(page, Page.ApiKeys, "/key/list");
expect(keyList.request().headers()["authorization"]).toBe(`Bearer ${key}`);
expect(keyList.status()).toBe(200);
expect(Array.isArray((await keyList.json()).keys)).toBe(true);
const modelInfo = await openPageAndCapture(page, Page.Models, "/v2/model/info");
expect(modelInfo.request().headers()["authorization"]).toBe(`Bearer ${key}`);
expect(modelInfo.status()).toBe(200);
const models: { model_name: string }[] = (await modelInfo.json()).data;
expect(models.length, "the stack serves at least one model").toBeGreaterThan(0);
await expect(page.getByText(models[0].model_name, { exact: true }).first()).toBeVisible();
});

View file

@ -8,7 +8,7 @@ import { ARTIFACT_DIR, UI_BASE_URL } from "./constants";
export default defineConfig({
testDir: ".",
testMatch: ["**/*.spec.ts", "**/*.setup.ts"],
testIgnore: ["**/*.test.*", "**/integrationCritical/**"],
testIgnore: ["**/*.test.*", "**/integrationCritical/**", "oidc/**"],
/* Run tests in files in parallel */
fullyParallel: true,
/* Fail the build on CI if you accidentally left test.only in the source code. */

View file

@ -6,6 +6,7 @@ import {
CHAT_MODEL_A,
MOCK_RESPONSE_TEXT,
sendChatCompletion,
sendChatCompletionWithCallId,
waitForSpendLog,
waitForSpendLogByPrompt,
} from "../../helpers/traffic";
@ -95,6 +96,50 @@ test.describe("Logs page", () => {
await expect(drawer.getByText(MOCK_RESPONSE_TEXT, { exact: false }).first()).toBeVisible({ timeout: 20_000 });
});
test("a served request's Logs row and drawer show its x-litellm-call-id", async ({ page, request }) => {
const prompt = `logs-call-id-prompt-${uniqueSuffix()}`;
const { requestId, callId } = await sendChatCompletionWithCallId(request, {
model: CHAT_MODEL_A,
prompt,
});
expect(callId, "call id must differ from the provider response id for this check to mean anything").not.toBe(
requestId,
);
await waitForSpendLog(request, requestId);
await navigateToPage(page, Page.Logs);
await dismissFeedbackPopup(page);
const search = visibleTestId(page, "datatable-search");
await expect(search).toBeVisible({ timeout: 20_000 });
await search.fill(callId);
const row = requestLogsRows(page).filter({ hasText: requestId });
await expect(row, `no logs row for call id ${callId}`).toHaveCount(1, { timeout: 30_000 });
await expect(row, "the row itself shows only the request id").not.toContainText(callId);
await row.getByText(requestId).hover();
const tooltip = page.locator("[data-slot='tooltip-content']");
await expect(tooltip, "hovering the Request ID cell does not list the x-litellm-call-id").toContainText(
`x-litellm-call-id: ${callId}`,
{ timeout: 10_000 },
);
await tooltip.getByRole("button", { name: "Copy x-litellm-call-id" }).click();
if (await page.evaluate(() => window.isSecureContext)) {
await expect.poll(() => page.evaluate(() => navigator.clipboard.readText())).toBe(callId);
}
await row.click();
const drawer = page.getByRole("dialog").first();
await expect(drawer.getByText("Request & Response")).toBeVisible({ timeout: 20_000 });
await expect(drawer.getByText("x-litellm-call-id:"), "drawer header lacks the x-litellm-call-id line").toBeVisible({
timeout: 10_000,
});
await expect(
drawer.getByText(callId, { exact: false }).first(),
`drawer does not show x-litellm-call-id ${callId}`,
).toBeVisible({ timeout: 10_000 });
});
// Split out because only the copy path needs a secure context; folding it in would
// take the drawer-rendering coverage down with it.
test("the drawer copies the request and the response to the clipboard", async ({ page, request }) => {

View file

@ -3,6 +3,8 @@ import uuid
from collections.abc import Iterator, Mapping, Sequence
from contextlib import contextmanager
from dataclasses import dataclass
from datetime import datetime, timedelta
from hashlib import sha256
from itertools import chain
from typing import Final
@ -34,7 +36,16 @@ INSERT_SPEND_LOG: Final = (
" VALUES (%s, 'acompletion', %s, %s::timestamp, %s::timestamp, %s)"
)
DELETE_SPEND_LOG: Final = 'DELETE FROM "LiteLLM_SpendLogs" WHERE request_id = %s'
INSERT_SPEND_LOG_ROW: Final = (
'INSERT INTO "LiteLLM_SpendLogs" (request_id, call_type, api_key, "startTime", "endTime", metadata, team_id, "user")'
" VALUES (%s, 'acompletion', %s, %s::timestamp, %s::timestamp, %s, %s, %s)"
)
DELETE_SPEND_LOG_ROWS: Final = 'DELETE FROM "LiteLLM_SpendLogs" WHERE request_id = ANY(%s)'
DELETE_KEY_ROW: Final = 'DELETE FROM "LiteLLM_VerificationToken" WHERE token = %s'
DELETE_ARCHIVED_KEY_ROW: Final = 'DELETE FROM "LiteLLM_DeletedVerificationToken" WHERE token = %s'
LOCK_TABLE: Final = sql.SQL("LOCK TABLE {table} IN ACCESS EXCLUSIVE MODE")
SPEND_LOGS_TABLE: Final = "LiteLLM_SpendLogs"
FIRST_SPEND_LOG_AT: Final = datetime(2026, 2, 3, 12, 0, 0)
@dataclass(frozen=True, slots=True)
@ -67,6 +78,10 @@ def key_no_key_table_holds() -> str:
return f"integration-ownerless-{uuid.uuid4().hex}"
def digest_no_key_table_holds() -> str:
return sha256(uuid.uuid4().bytes).hexdigest()
def activity_of_key(
gateway: Gateway, path: str, api_key: str, *, reader: str | None = None, **filters: str
) -> httpx.Response:
@ -148,6 +163,56 @@ def spend_log_naming_only_an_alias(request_id: str, api_key: str, started: str,
connection.execute(DELETE_SPEND_LOG, (request_id,))
@dataclass(frozen=True, slots=True)
class SpendLogRow:
started: str
metadata: JsonValue = None
team_id: str | None = None
user: str | None = None
def started_at(index: int) -> str:
return (FIRST_SPEND_LOG_AT + timedelta(seconds=index)).strftime("%Y-%m-%d %H:%M:%S")
def nameless_rows(count: int, first_index: int = 0) -> tuple[SpendLogRow, ...]:
return tuple(SpendLogRow(started_at(first_index + offset), {}) for offset in range(count))
def named_row(index: int, alias: str) -> SpendLogRow:
return SpendLogRow(started_at(index), {"user_api_key_alias": alias})
@contextmanager
def spend_logs_of_key(
api_key: str, rows: Sequence[SpendLogRow], *, database_url: str | None = None
) -> Iterator[tuple[str, ...]]:
request_ids: Final = tuple(f"integration-{uuid.uuid4().hex}" for _ in rows)
with psycopg.connect(database_url or os.environ["DATABASE_URL"]) as connection:
connection.cursor().executemany(
INSERT_SPEND_LOG_ROW,
tuple(
(request_id, api_key, row.started, row.started, Jsonb(row.metadata), row.team_id, row.user)
for request_id, row in zip(request_ids, rows, strict=True)
),
)
try:
yield request_ids
finally:
delete_spend_logs(request_ids, database_url=database_url)
def delete_spend_logs(request_ids: Sequence[str], *, database_url: str | None = None) -> None:
with psycopg.connect(database_url or os.environ["DATABASE_URL"]) as connection:
connection.execute(DELETE_SPEND_LOG_ROWS, (list(request_ids),))
def purge_key_from_the_key_tables(digest: str, *, database_url: str | None = None) -> None:
with psycopg.connect(database_url or os.environ["DATABASE_URL"]) as connection:
connection.execute(DELETE_KEY_ROW, (digest,))
connection.execute(DELETE_ARCHIVED_KEY_ROW, (digest,))
@contextmanager
def locked_table(table: str, *, database_url: str | None = None) -> Iterator[None]:
with psycopg.connect(database_url or os.environ["DATABASE_URL"]) as connection:

View file

@ -324,7 +324,7 @@ ROUTES: Final[tuple[Route, ...]] = (
lambda s: Call("POST", "/team/update", {"team_id": s.team_id, "max_budget": 5}),
team_admin=403, others=403, org_admin=200),
Route("team_update_budget_permitted",
lambda s: Call("POST", "/team/update", {"team_id": s.team_id, "max_budget": 7}),
lambda s: Call("POST", "/team/update", {"team_id": s.team_id, "max_budget": 4}),
team_admin=200, others=403, org_admin=200, permission="max_budget"),
Route("project_new",
lambda s: Call("POST", "/project/new", {"team_id": s.team_id, "project_alias": f"matrix-{uuid.uuid4().hex}"}),
@ -414,6 +414,8 @@ def test_status_code(shared: TeamScenario, org_team: TeamScenario, route: Route,
team: Final = org_team if caller in ORG_CALLERS else shared
with team.gateway.scenario() as scenario:
s: Final = replace(team, scenario=scenario)
if route.name == "team_update_budget_permitted":
s.gateway.post("/team/update", {"team_id": s.team_id, "max_budget": 5})
if route.permission:
scenario.cleanups.enter_context(team_admin_permissions(s.gateway, (route.permission,)))
call: Final = route.call(s)
@ -421,5 +423,9 @@ def test_status_code(shared: TeamScenario, org_team: TeamScenario, route: Route,
assert response.status_code == route.expected(caller), (
f"{caller} {call.method} {call.path}: {response.status_code} {response.text}"
)
if route.name == "team_update_budget_permitted":
assert read_rows(
'SELECT max_budget FROM "LiteLLM_TeamTable" WHERE team_id = %s', (s.team_id,)
) == [{"max_budget": 4.0 if response.status_code == 200 else 5.0}]
if response.status_code == 200 and route.cleanup is not None:
route.cleanup(s, object_value(response.json()))

View file

@ -1,11 +1,15 @@
import json
from typing import Final
import uuid
from typing import Final, Literal
import httpx
import pytest
from pydantic import JsonValue
from tests.integration._support.client import JSON_OBJECT, Gateway, eventually, object_value, string_value
from tests.integration._support.database import read_rows
from tests.integration._support.upstream import delete_scenario, register_scenario
from tests.integration.cost_calculation.cost_tracking_case import JsonResponse
STANDARD_INPUT_RATE: Final = 0.001
STANDARD_OUTPUT_RATE: Final = 0.002
@ -69,3 +73,190 @@ def test_ultrafast_service_tier_bills_ultrafast_rates_and_keeps_pricing_off_the_
)
assert_chat_bills_rates(gateway, model, "ultrafast", ULTRAFAST_INPUT_RATE, ULTRAFAST_OUTPUT_RATE)
assert_chat_bills_rates(gateway, model, None, STANDARD_INPUT_RATE, STANDARD_OUTPUT_RATE)
LONG_CONTEXT_PRICING: Final[dict[str, JsonValue]] = {
"input_cost_per_token": 1e-06,
"output_cost_per_token": 2e-06,
"cache_read_input_token_cost": 1e-07,
"input_cost_per_token_above_272k_tokens": 3e-06,
"output_cost_per_token_above_272k_tokens": 4e-06,
"cache_read_input_token_cost_above_272k_tokens": 3e-07,
"input_cost_per_token_ultrafast": 1e-05,
"output_cost_per_token_ultrafast": 2e-05,
"cache_read_input_token_cost_ultrafast": 1e-06,
"input_cost_per_token_above_272k_tokens_ultrafast": 5e-05,
"output_cost_per_token_above_272k_tokens_ultrafast": 6e-05,
"cache_read_input_token_cost_above_272k_tokens_ultrafast": 5e-06,
"cache_creation_input_token_cost_above_272k_tokens_ultrafast": 6e-06,
}
LONG_PROMPT_TOKENS: Final = 300_000
SHORT_PROMPT_TOKENS: Final = 1_000
CACHED_TOKENS: Final = 400
COMPLETION_TOKENS: Final = 1_000
def _chat_response(service_tier: str | None, prompt_tokens: int) -> JsonResponse:
return JsonResponse(
content_type="application/json",
body={
"id": "chatcmpl-$UNIQUE_ID",
"object": "chat.completion",
"created": 1,
"model": "integration-ultrafast-long-context",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "long context answer"},
"finish_reason": "stop",
}
],
"usage": {
"prompt_tokens": prompt_tokens,
"completion_tokens": COMPLETION_TOKENS,
"total_tokens": prompt_tokens + COMPLETION_TOKENS,
"prompt_tokens_details": {"cached_tokens": CACHED_TOKENS},
},
**({} if service_tier is None else {"service_tier": service_tier}),
},
)
def _responses_response(service_tier: str | None, prompt_tokens: int) -> JsonResponse:
return JsonResponse(
content_type="application/json",
body={
"id": "resp_$UNIQUE_ID",
"object": "response",
"created_at": 1,
"status": "completed",
"model": "integration-ultrafast-long-context",
"output": [
{
"type": "message",
"id": "msg_$UNIQUE_ID",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "long context answer", "annotations": []}],
}
],
"usage": {
"input_tokens": prompt_tokens,
"output_tokens": COMPLETION_TOKENS,
"total_tokens": prompt_tokens + COMPLETION_TOKENS,
"input_tokens_details": {"cached_tokens": CACHED_TOKENS},
"output_tokens_details": {"reasoning_tokens": 0},
},
**({} if service_tier is None else {"service_tier": service_tier}),
},
)
def _surface_response(
surface: Literal["chat", "responses"], service_tier: str | None, prompt_tokens: int
) -> JsonResponse:
match surface:
case "chat":
return _chat_response(service_tier, prompt_tokens)
case "responses":
return _responses_response(service_tier, prompt_tokens)
def _surface_request(
surface: Literal["chat", "responses"], scenario_id: str, model: str, service_tier: str | None
) -> tuple[str, dict[str, JsonValue], str]:
match surface:
case "chat":
return (
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": "long context ultrafast control"}],
**({} if service_tier is None else {"service_tier": service_tier}),
},
f"/{scenario_id}/chat/completions",
)
case "responses":
return (
"/v1/responses",
{
"model": model,
"input": "long context ultrafast control",
**({} if service_tier is None else {"service_tier": service_tier}),
},
f"/{scenario_id}/responses",
)
@pytest.mark.parametrize(
("service_tier", "prompt_tokens", "input_rate", "cache_read_rate", "output_rate"),
(
("ultrafast", LONG_PROMPT_TOKENS, 5e-05, 5e-06, 6e-05),
("ultrafast", SHORT_PROMPT_TOKENS, 1e-05, 1e-06, 2e-05),
(None, LONG_PROMPT_TOKENS, 3e-06, 3e-07, 4e-06),
),
ids=("ultrafast_above_272k", "ultrafast_below_272k", "standard_above_272k"),
)
@pytest.mark.parametrize("surface", ("chat", "responses"), ids=("chat", "responses"))
def test_ultrafast_long_context_prompt_bills_ultrafast_long_context_rates(
gateway: Gateway,
surface: Literal["chat", "responses"],
service_tier: str | None,
prompt_tokens: int,
input_rate: float,
cache_read_rate: float,
output_rate: float,
) -> None:
with gateway.scenario() as scenario:
scenario_id: Final = f"ultrafast-long-context-{uuid.uuid4().hex}"
handle: Final = register_scenario(
scenario_id, _surface_response(surface, service_tier, prompt_tokens)
)
scenario.cleanups.callback(delete_scenario, handle)
key: Final = scenario.key()
model: Final = scenario.model(
model=f"openai/integration-ultrafast-long-context-{uuid.uuid4().hex}",
api_key=scenario_id,
api_base=handle.api_base(),
**LONG_CONTEXT_PRICING,
)
request_path, request_body, expected_upstream_path = _surface_request(surface, scenario_id, model, service_tier)
with httpx.Client(base_url=gateway.upstream_url, trust_env=False) as upstream:
upstream.get("/__observations").raise_for_status()
response: Final = gateway.request(
"POST",
request_path,
request_body,
key=key,
)
observations: Final = JSON_OBJECT.validate_json(upstream.get("/__observations").content)["requests"]
assert response.status_code == 200, response.text
expected_input: Final = (prompt_tokens - CACHED_TOKENS) * input_rate + CACHED_TOKENS * cache_read_rate
expected_output: Final = COMPLETION_TOKENS * output_rate
expected: Final = expected_input + expected_output
assert float(response.headers["x-litellm-response-cost"]) == pytest.approx(expected, rel=1e-6), response.text
request_id: Final = string_value(object_value(response.json())["id"])
rows: Final = eventually(
lambda: read_rows(
'SELECT spend, metadata, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" WHERE request_id = %s',
(request_id,),
),
lambda values: len(values) == 1,
seconds=70,
)
assert rows[0]["prompt_tokens"] == prompt_tokens
assert rows[0]["completion_tokens"] == COMPLETION_TOKENS
assert float(rows[0]["spend"]) == pytest.approx(expected, rel=1e-6)
metadata: Final = rows[0]["metadata"]
parsed: Final = json.loads(metadata) if isinstance(metadata, str) else object_value(metadata)
breakdown: Final = object_value(parsed["cost_breakdown"])
assert float(breakdown["input_cost"]) == pytest.approx(expected_input, rel=1e-6)
assert float(breakdown["output_cost"]) == pytest.approx(expected_output, rel=1e-6)
assert isinstance(observations, list)
assert len(observations) == 1
observation: Final = object_value(observations[0])
upstream_path: Final = string_value(observation["path"])
assert upstream_path == expected_upstream_path, upstream_path
body: Final = object_value(observation["body"])
assert body.get("service_tier") == service_tier, body
assert not set(LONG_CONTEXT_PRICING).intersection(body), body

View file

@ -12,6 +12,8 @@ def test_chat_over_responses_deployment_returns_length_when_output_tokens_run_ou
identity: Final = "responses-incomplete-" + uuid.uuid4().hex
def respond(request: Request) -> Reply:
if request.method == "GET" and request.target == "/v1/models":
return Reply(body=b'{"object":"list","data":[]}')
assert request.method == "POST" and request.target == "/responses", request.target
assert request.headers["authorization"] == "Bearer synthetic-openai-key"
body: Final = json.loads(request.body)
@ -56,7 +58,7 @@ def test_chat_over_responses_deployment_returns_length_when_output_tokens_run_ou
)
assert response.status_code == 200, response.text
body: Final = response.json()
assert len(wire.drain()) == 1
assert len(tuple(request for request in wire.drain() if request.method == "POST")) == 1
assert [choice["finish_reason"] for choice in body["choices"]] == ["length"], response.text
assert body["choices"][0]["message"]["content"] == "", response.text
assert body["choices"][0]["message"]["role"] == "assistant", response.text
@ -69,6 +71,8 @@ def test_messages_over_responses_deployment_with_max_tokens_1_is_clamped_to_16_i
identity: Final = "responses-clamp-" + uuid.uuid4().hex
def respond(request: Request) -> Reply:
if request.method == "GET" and request.target == "/v1/models":
return Reply(body=b'{"object":"list","data":[]}')
assert request.method == "POST" and request.target == "/responses", request.target
assert request.headers["authorization"] == "Bearer synthetic-openai-key"
body: Final = json.loads(request.body)
@ -132,7 +136,7 @@ def test_messages_over_responses_deployment_with_max_tokens_1_is_clamped_to_16_i
)
assert response.status_code == 200, response.text
body: Final = response.json()
assert len(wire.drain()) == 1
assert len(tuple(request for request in wire.drain() if request.method == "POST")) == 1
assert body["role"] == "assistant", response.text
assert body["content"] == [{"type": "text", "text": "ok"}], response.text
assert body["stop_reason"] == "end_turn", response.text
@ -143,6 +147,8 @@ def test_messages_over_responses_deployment_with_max_tokens_one_reaches_openai_a
identity: Final = "responses-min-tokens-" + uuid.uuid4().hex
def respond(request: Request) -> Reply:
if request.method == "GET" and request.target == "/v1/models":
return Reply(body=b'{"object":"list","data":[]}')
assert request.method == "POST" and request.target == "/responses", request.target
assert request.headers["authorization"] == "Bearer synthetic-openai-key"
body: Final = json.loads(request.body)
@ -185,6 +191,6 @@ def test_messages_over_responses_deployment_with_max_tokens_one_reaches_openai_a
)
assert response.status_code == 200, response.text
body: Final = response.json()
assert len(wire.drain()) == 1
assert len(tuple(request for request in wire.drain() if request.method == "POST")) == 1
assert body["content"] == [{"type": "text", "text": "ok"}], response.text
assert body["usage"]["input_tokens"] == 9 and body["usage"]["output_tokens"] == 1, response.text

View file

@ -0,0 +1,262 @@
from __future__ import annotations
import json
import shlex
import threading
import uuid
from collections.abc import Iterator
from contextlib import contextmanager
from datetime import UTC, datetime
from pathlib import Path
from queue import SimpleQueue
from types import MappingProxyType
from typing import Final
import httpx
import pytest
import yaml
from integration._support.client import Gateway, eventually
from integration._support.process import owned_proxy
from integration._support.redis_process import owned_redis
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue, TypeAdapter
from redis import Redis
from redis.exceptions import TimeoutError as RedisTimeoutError
JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
MONITOR_COMMAND: Final = TypeAdapter(dict[str, JsonValue])
OPENAI_MODEL: Final = "gpt-4o-mini"
MASTER_KEY: Final = "sk-integration-usage-routing-redis-reads"
API_KEY: Final = "synthetic-usage-routing-key"
ENDPOINT_PATHS: Final = MappingProxyType(
{
"/v1/chat/completions": ("/v1/chat/completions", "/v1/chat/completions"),
"/v1/messages": ("/v1/responses", "/v1/responses"),
"/v1/responses": ("/v1/responses", "/v1/responses"),
}
)
CHAT_RESPONSE: Final = json.dumps(
{
"id": "chatcmpl_usage_routing_redis_reads",
"object": "chat.completion",
"created": 1700000000,
"model": OPENAI_MODEL,
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "redis read contract"},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 7, "completion_tokens": 4, "total_tokens": 11},
}
).encode()
RESPONSES_RESPONSE: Final = json.dumps(
{
"id": "resp_usage_routing_redis_reads",
"object": "response",
"created_at": 1700000000,
"status": "completed",
"model": OPENAI_MODEL,
"output": [
{
"id": "msg_usage_routing_redis_reads",
"type": "message",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "redis read contract", "annotations": []}],
}
],
"usage": {"input_tokens": 7, "output_tokens": 4, "total_tokens": 11},
}
).encode()
def _request_object(body: bytes) -> dict[str, JsonValue]:
return JSON_OBJECT.validate_json(body)
def _deployment_list(
model_name: str, api_base: str, deployment_ids: tuple[str, str]
) -> list[dict[str, JsonValue]]:
return [
{
"model_name": model_name,
"litellm_params": {
"model": f"openai/{OPENAI_MODEL}",
"api_base": api_base,
"api_key": API_KEY,
"rpm": 1,
},
"model_info": {"id": deployment_id},
}
for deployment_id in deployment_ids
]
def _request_payload(endpoint: str, model_name: str, marker: str) -> dict[str, JsonValue]:
if endpoint == "/v1/responses":
return {"model": model_name, "input": marker, "max_output_tokens": 16, "store": False}
return {"model": model_name, "messages": [{"role": "user", "content": marker}], "max_tokens": 16}
def _expected_wire_body(endpoint: str, marker: str) -> dict[str, JsonValue]:
if endpoint == "/v1/messages":
return {
"model": OPENAI_MODEL,
"input": [{"type": "message", "role": "user", "content": [{"type": "input_text", "text": marker}]}],
"include": ["reasoning.encrypted_content"],
"max_output_tokens": 16,
}
if endpoint == "/v1/responses":
return {"model": OPENAI_MODEL, "input": marker, "max_output_tokens": 16, "store": False}
return {"model": OPENAI_MODEL, "messages": [{"role": "user", "content": marker}], "max_tokens": 16}
def _reply(request: Request) -> Reply:
if request.target == "/v1/models":
return Reply(body=json.dumps({"object": "list", "data": [{"id": OPENAI_MODEL, "object": "model"}]}).encode())
return Reply(body=RESPONSES_RESPONSE if request.target == "/v1/responses" else CHAT_RESPONSE)
@contextmanager
def _capture_redis_commands(host: str, port: int) -> Iterator[SimpleQueue[str]]:
commands: Final = SimpleQueue[str]()
started: Final = threading.Event()
armed: Final = threading.Event()
stopped: Final = threading.Event()
ready_marker: Final = f"monitor-ready-{uuid.uuid4().hex}"
stop_marker: Final = f"monitor-stop-{uuid.uuid4().hex}"
def capture() -> None:
with Redis(host=host, port=port, socket_timeout=1, decode_responses=True) as client:
with client.monitor() as monitor:
started.set()
stream: Final = iter(monitor.listen())
while not stopped.is_set():
try:
record: Final = MONITOR_COMMAND.validate_python(next(stream))
except RedisTimeoutError:
continue
command: Final = record.get("command")
if not isinstance(command, str):
continue
commands.put(command)
if ready_marker in command:
armed.set()
thread: Final = threading.Thread(target=capture, daemon=True)
thread.start()
try:
assert started.wait(timeout=5), "Redis MONITOR did not start"
with Redis(host=host, port=port, socket_timeout=1, decode_responses=True) as client:
client.set(ready_marker, "ready", ex=1)
assert armed.wait(timeout=5), "Redis MONITOR did not capture its readiness command"
yield commands
finally:
stopped.set()
with Redis(host=host, port=port, socket_timeout=1, decode_responses=True) as client:
client.set(stop_marker, "stop", ex=1)
thread.join(timeout=5)
assert not thread.is_alive(), "Redis MONITOR thread survived cleanup"
def _drain_mgets(commands: SimpleQueue[str]) -> tuple[tuple[str, tuple[str, ...]], ...]:
captured: Final = tuple(commands.get_nowait() for _ in range(commands.qsize()))
parsed: Final = tuple((line, tuple(shlex.split(line))) for line in captured)
return tuple((line, arguments) for line, arguments in parsed if arguments and arguments[0] == "MGET")
@pytest.mark.parametrize(
"endpoint",
("/v1/chat/completions", "/v1/messages", "/v1/responses"),
ids=("chat-completions", "messages", "responses"),
)
def test_proxy_usage_routing_reads_cooldown_tpm_then_rpm_from_redis(
endpoint: str, tmp_path: Path
) -> None:
with owned_redis(tmp_path) as cache, wire_server(_reply) as wire:
run_id: Final = uuid.uuid4().hex
model_name: Final = f"usage-redis-{run_id}"
deployment_ids: Final = (f"dep-a-{run_id[:8]}", f"dep-b-{run_id[:8]}")
configuration: Final = JSON_OBJECT.validate_python(
yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
)
config: Final = {
**configuration,
"model_list": _deployment_list(model_name, f"{wire.url}/v1", deployment_ids),
"router_settings": {
"routing_strategy": "usage-based-routing-v2",
"redis_host": cache.host,
"redis_port": cache.port,
},
}
config_path: Final = tmp_path / "usage-routing.yaml"
config_path.write_text(yaml.safe_dump(config))
with httpx.Client(base_url=wire.url, timeout=15, trust_env=False) as bootstrap_client:
bootstrap: Final = Gateway(bootstrap_client, MASTER_KEY, wire.url)
with owned_proxy(
bootstrap,
tmp_path,
{"REDIS_HOST": cache.host, "REDIS_PORT": str(cache.port)},
config=config_path,
) as candidate:
eventually(
lambda: wire.received.qsize(),
lambda received: received >= len(deployment_ids),
seconds=15,
)
wire.drain()
eventually(
lambda: datetime.now(UTC),
lambda current: current.second < 40,
seconds=65,
)
minute: Final = datetime.now(UTC).strftime("%H-%M")
markers: Final = tuple(f"{run_id}-{index}" for index in range(3))
payloads: Final = tuple(_request_payload(endpoint, model_name, marker) for marker in markers)
request_headers: Final = (
{"anthropic-version": "2023-06-01"} if endpoint == "/v1/messages" else {}
)
with _capture_redis_commands(cache.host, cache.port) as commands:
responses: Final = tuple(
candidate.request("POST", endpoint, payload, headers=request_headers) for payload in payloads
)
assert tuple(response.status_code for response in responses) == (200, 200, 429), [
response.text for response in responses
]
assert "No deployments available" in responses[2].text
served_ids: Final = tuple(response.headers["x-litellm-model-id"] for response in responses[:2])
assert set(served_ids) == set(deployment_ids), served_ids
rpm_keys: Final = tuple(
f"{deployment_id}:openai/{OPENAI_MODEL}:rpm:{minute}" for deployment_id in deployment_ids
)
with Redis(host=cache.host, port=cache.port, decode_responses=True) as redis_client:
rpm_values: Final = eventually(
lambda: tuple(redis_client.get(key) for key in rpm_keys),
lambda values: values == ("1", "1"),
seconds=15,
)
assert rpm_values == ("1", "1")
received: Final = wire.drain()
assert len(received) == 2
assert tuple(request.method for request in received) == ("POST", "POST")
assert tuple(request.target for request in received) == ENDPOINT_PATHS[endpoint]
observed_bodies: Final = tuple(_request_object(request.body) for request in received)
expected_bodies: Final = tuple(_expected_wire_body(endpoint, marker) for marker in markers[:2])
assert observed_bodies == expected_bodies, observed_bodies
expected_mget: Final = (
"MGET",
f"deployment:{deployment_ids[0]}:cooldown",
f"deployment:{deployment_ids[1]}:cooldown",
f"{deployment_ids[0]}:openai/{OPENAI_MODEL}:tpm:{minute}",
f"{deployment_ids[1]}:openai/{OPENAI_MODEL}:tpm:{minute}",
*(
f"{deployment_id}:openai/{OPENAI_MODEL}:rpm:{minute}"
for deployment_id in deployment_ids
),
)
mgets: Final = _drain_mgets(commands)
assert any(arguments == expected_mget for _, arguments in mgets), mgets
raw_mgets: Final = tuple(line for line, _ in mgets)
print(f"proxy {endpoint} MGETs: {raw_mgets}")

View file

@ -0,0 +1,191 @@
from __future__ import annotations
import asyncio
import json
import shlex
import threading
import uuid
from collections.abc import Iterator
from contextlib import contextmanager
from datetime import UTC, datetime
from pathlib import Path
from queue import SimpleQueue
from typing import Final
import litellm
import pytest
from integration._support.client import eventually
from integration._support.redis_process import owned_redis
from integration._support.wire import Reply, Request, wire_server
from litellm import Router
from pydantic import JsonValue, TypeAdapter
from redis import Redis
from redis.exceptions import TimeoutError as RedisTimeoutError
JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
MONITOR_COMMAND: Final = TypeAdapter(dict[str, JsonValue])
OPENAI_MODEL: Final = "gpt-4o-mini"
API_KEY: Final = "synthetic-usage-routing-key"
CHAT_RESPONSE: Final = json.dumps(
{
"id": "chatcmpl_usage_routing_sdk_redis_reads",
"object": "chat.completion",
"created": 1700000000,
"model": OPENAI_MODEL,
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "redis read contract"},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 7, "completion_tokens": 4, "total_tokens": 11},
}
).encode()
def _request_object(body: bytes) -> dict[str, JsonValue]:
return JSON_OBJECT.validate_json(body)
def _deployment_list(
model_name: str, api_base: str, deployment_ids: tuple[str, str]
) -> list[dict[str, JsonValue]]:
return [
{
"model_name": model_name,
"litellm_params": {
"model": f"openai/{OPENAI_MODEL}",
"api_base": api_base,
"api_key": API_KEY,
"rpm": 1,
},
"model_info": {"id": deployment_id},
}
for deployment_id in deployment_ids
]
def _reply(request: Request) -> Reply:
return Reply(body=CHAT_RESPONSE)
@contextmanager
def _capture_redis_commands(host: str, port: int) -> Iterator[SimpleQueue[str]]:
commands: Final = SimpleQueue[str]()
started: Final = threading.Event()
armed: Final = threading.Event()
stopped: Final = threading.Event()
ready_marker: Final = f"monitor-ready-{uuid.uuid4().hex}"
stop_marker: Final = f"monitor-stop-{uuid.uuid4().hex}"
def capture() -> None:
with Redis(host=host, port=port, socket_timeout=1, decode_responses=True) as client:
with client.monitor() as monitor:
started.set()
stream: Final = iter(monitor.listen())
while not stopped.is_set():
try:
record: Final = MONITOR_COMMAND.validate_python(next(stream))
except RedisTimeoutError:
continue
command: Final = record.get("command")
if not isinstance(command, str):
continue
commands.put(command)
if ready_marker in command:
armed.set()
thread: Final = threading.Thread(target=capture, daemon=True)
thread.start()
try:
assert started.wait(timeout=5), "Redis MONITOR did not start"
with Redis(host=host, port=port, socket_timeout=1, decode_responses=True) as client:
client.set(ready_marker, "ready", ex=1)
assert armed.wait(timeout=5), "Redis MONITOR did not capture its readiness command"
yield commands
finally:
stopped.set()
with Redis(host=host, port=port, socket_timeout=1, decode_responses=True) as client:
client.set(stop_marker, "stop", ex=1)
thread.join(timeout=5)
assert not thread.is_alive(), "Redis MONITOR thread survived cleanup"
def _drain_mgets(commands: SimpleQueue[str]) -> tuple[tuple[str, tuple[str, ...]], ...]:
captured: Final = tuple(commands.get_nowait() for _ in range(commands.qsize()))
parsed: Final = tuple((line, tuple(shlex.split(line))) for line in captured)
return tuple((line, arguments) for line, arguments in parsed if arguments and arguments[0] == "MGET")
def _model_id(response: object) -> str:
response_params: Final = getattr(response, "_hidden_params")
hidden_params: Final = JSON_OBJECT.validate_python(response_params)
model_id: Final = hidden_params.get("model_id")
assert isinstance(model_id, str), hidden_params
return model_id
async def _exercise_router(router: Router, model_name: str, markers: tuple[str, str, str]) -> tuple[str, str]:
first: Final = await router.acompletion(
model=model_name, messages=[{"role": "user", "content": markers[0]}], max_tokens=8
)
second: Final = await router.acompletion(
model=model_name, messages=[{"role": "user", "content": markers[1]}], max_tokens=8
)
with pytest.raises(litellm.RateLimitError, match="No deployments available"):
await router.acompletion(
model=model_name, messages=[{"role": "user", "content": markers[2]}], max_tokens=8
)
return _model_id(first), _model_id(second)
def test_sdk_usage_routing_reads_tpm_then_rpm_from_redis(tmp_path: Path) -> None:
with owned_redis(tmp_path) as cache, wire_server(_reply) as wire:
run_id: Final = uuid.uuid4().hex
model_name: Final = f"usage-redis-{run_id}"
deployment_ids: Final = (f"dep-a-{run_id[:8]}", f"dep-b-{run_id[:8]}")
router: Final = Router(
model_list=_deployment_list(model_name, f"{wire.url}/v1", deployment_ids),
routing_strategy="usage-based-routing-v2",
redis_host=cache.host,
redis_port=cache.port,
)
try:
eventually(
lambda: datetime.now(UTC),
lambda current: current.second < 40,
seconds=65,
)
minute: Final = datetime.now(UTC).strftime("%H-%M")
markers: Final = tuple(f"{run_id}-{index}" for index in range(3))
with _capture_redis_commands(cache.host, cache.port) as commands:
served_ids: Final = asyncio.run(_exercise_router(router, model_name, markers))
assert set(served_ids) == set(deployment_ids), served_ids
received: Final = wire.drain()
assert len(received) == 2
assert tuple(request.method for request in received) == ("POST", "POST")
assert tuple(request.target for request in received) == ("/v1/chat/completions",) * 2
observed_bodies: Final = tuple(_request_object(request.body) for request in received)
expected_bodies: Final = tuple(
{
"model": OPENAI_MODEL,
"messages": [{"role": "user", "content": marker}],
"max_tokens": 8,
}
for marker in markers[:2]
)
assert observed_bodies == expected_bodies, observed_bodies
expected_mget: Final = (
"MGET",
f"deployment:{deployment_ids[0]}:cooldown",
f"deployment:{deployment_ids[1]}:cooldown",
*(f"{deployment_id}:openai/{OPENAI_MODEL}:tpm:{minute}" for deployment_id in deployment_ids),
*(f"{deployment_id}:openai/{OPENAI_MODEL}:rpm:{minute}" for deployment_id in deployment_ids),
)
mgets: Final = _drain_mgets(commands)
assert any(arguments == expected_mget for _, arguments in mgets), mgets
raw_mgets: Final = tuple(line for line, _ in mgets)
print(f"sdk MGETs: {raw_mgets}")
finally:
router.reset()

View file

@ -2,11 +2,12 @@ from __future__ import annotations
import json
import uuid
from datetime import datetime, timedelta, timezone
from hashlib import sha256
from typing import Final
import pytest
from integration._support.client import JSON_OBJECT, Gateway, eventually, string_value
from integration._support.client import JSON_OBJECT, Gateway, eventually, object_value, string_value
from integration._support.database import read_rows
from integration._support.upstream import delete_scenario, register_scenario
from integration.cost_calculation.cost_tracking_case import JsonResponse, RoutedResponse, TextResponse
@ -116,6 +117,28 @@ def _batch_routes(model: str) -> RoutedResponse:
)
def _team_day_endpoints(gateway: Gateway, team: str, start_date: str, end_date: str) -> dict[str, object] | None:
response: Final = gateway.request(
"GET",
"/team/daily/activity",
params={"team_ids": team, "start_date": start_date, "end_date": end_date},
)
if response.status_code != 200:
return None
days: Final = response.json()["results"]
if not days:
return None
return object_value(object_value(object_value(days[0])["breakdown"])["endpoints"])
def _batches_total_tokens(endpoints: dict[str, object] | None) -> int | None:
if endpoints is None or "/batches" not in endpoints:
return None
metrics: Final = object_value(object_value(endpoints["/batches"])["metrics"])
total_tokens: Final = metrics["total_tokens"]
return int(total_tokens) if isinstance(total_tokens, (int, float, str)) else None
def _input_file(model: str) -> bytes:
return (
"\n".join(
@ -200,3 +223,79 @@ def test_completed_batch_spend_row_records_reasoning_tokens_and_error_file_failu
"reasoning_tokens": reasoning_tokens,
"text_tokens": completion_tokens - reasoning_tokens,
}, json.dumps(metadata)
INPUT_COST_PER_TOKEN: Final = 0.001
OUTPUT_COST_PER_TOKEN: Final = 0.002
BATCH_PROMPT_TOKENS: Final = FIRST_LINE["prompt_tokens"] + SECOND_LINE["prompt_tokens"]
BATCH_COMPLETION_TOKENS: Final = FIRST_LINE["completion_tokens"] + SECOND_LINE["completion_tokens"]
BATCH_SPEND: Final = (BATCH_PROMPT_TOKENS * INPUT_COST_PER_TOKEN + BATCH_COMPLETION_TOKENS * OUTPUT_COST_PER_TOKEN) / 2
def test_completed_batch_spend_lands_under_batches_in_team_endpoint_activity(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
scenario_id: Final = f"batch-endpoint-{uuid.uuid4().hex[:12]}"
handle: Final = register_scenario(scenario_id, _batch_routes("gpt-4o-mini"))
scenario.cleanups.callback(delete_scenario, handle)
model: Final = scenario.model(
api_base=handle.api_base(),
input_cost_per_token=INPUT_COST_PER_TOKEN,
output_cost_per_token=OUTPUT_COST_PER_TOKEN,
)
team: Final = scenario.team(models=[model])
key: Final = scenario.key(team_id=team, models=[model])
file_response: Final = gateway.request_multipart(
"/v1/files",
{"purpose": "batch", "model": model},
{"file": ("in.jsonl", _input_file(model), "application/jsonl")},
key=key,
)
assert file_response.status_code == 200, file_response.text
batch_response: Final = gateway.request(
"POST",
"/v1/batches",
{
"input_file_id": string_value(JSON_OBJECT.validate_json(file_response.content)["id"]),
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
"model": model,
},
key=key,
)
assert batch_response.status_code == 200, batch_response.text
batch_id: Final = string_value(JSON_OBJECT.validate_json(batch_response.content)["id"])
retrieval: Final = gateway.request("GET", f"/v1/batches/{batch_id}", key=key)
assert retrieval.status_code == 200, retrieval.text
assert retrieval.json()["status"] == "completed", retrieval.text
rows: Final = eventually(
lambda: read_rows(
'SELECT spend, prompt_tokens, completion_tokens, total_tokens FROM "LiteLLM_SpendLogs" '
"WHERE api_key=%s AND call_type='aretrieve_batch'",
(sha256(key.encode()).hexdigest(),),
),
lambda values: len(values) == 1,
seconds=70,
)
row: Final = rows[0]
assert float(row["spend"]) == pytest.approx(BATCH_SPEND), dict(row)
assert (row["prompt_tokens"], row["completion_tokens"]) == (
BATCH_PROMPT_TOKENS,
BATCH_COMPLETION_TOKENS,
), dict(row)
today: Final = datetime.now(timezone.utc)
endpoints: Final = eventually(
lambda: _team_day_endpoints(
gateway,
team,
(today - timedelta(days=1)).strftime("%Y-%m-%d"),
(today + timedelta(days=1)).strftime("%Y-%m-%d"),
),
lambda value: _batches_total_tokens(value) == BATCH_PROMPT_TOKENS + BATCH_COMPLETION_TOKENS,
seconds=70,
return_last_on_timeout=True,
)
assert endpoints is not None, "team daily activity returned no endpoint breakdown for the day"
assert set(endpoints) == {"/batches"}, endpoints
endpoint_metrics: Final = object_value(object_value(endpoints["/batches"])["metrics"])
assert float(endpoint_metrics["spend"]) == pytest.approx(BATCH_SPEND), endpoints
assert endpoint_metrics["total_tokens"] == BATCH_PROMPT_TOKENS + BATCH_COMPLETION_TOKENS, endpoints

View file

@ -0,0 +1,490 @@
import time
import uuid
from collections.abc import Callable, Iterator
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
from hashlib import sha256
from pathlib import Path
from typing import Final
import httpx
import pytest
from integration._support.client import Gateway, eventually, object_value
from integration._support.daily_activity import (
AGGREGATED_USER_ACTIVITY,
DAY,
ROUTES,
SPEND_LOGS_TABLE,
USER_SPEND,
Route,
SpendLogRow,
activity_of_key,
assert_key_reported,
daily_rows,
digest_no_key_table_holds,
key_metadata,
locked_table,
named_row,
nameless_rows,
records_of_key,
seeded_metrics,
seeded_row,
spend_logs_of_key,
started_at,
user_row,
user_with_an_email,
)
from integration._support.database import read_rows, scratch_database
from integration._support.process import OwnedProxy, owned_proxy_process
from pydantic import JsonValue
DAY_OUTSIDE_THE_WINDOW: Final = "2026-02-10"
GIVES_UP_WITHIN_SECONDS: Final = 10
CONCURRENT_READS: Final = 20
CACHED_MISS_CLEARS_WITHIN_SECONDS: Final = 45
ALIAS_OF_ONE_SPEND_LOG: Final = (
"SELECT metadata->>'user_api_key_alias' AS alias FROM \"LiteLLM_SpendLogs\" WHERE request_id = %s"
)
def _alias() -> str:
return f"integration-alias-{uuid.uuid4().hex}"
def _named_between_fifty_and_fifty(alias: str) -> tuple[SpendLogRow, ...]:
return (*nameless_rows(50), named_row(50, alias), *nameless_rows(50, 51))
def _oldest_named(alias: str) -> tuple[SpendLogRow, ...]:
return (named_row(0, alias), *nameless_rows(150, 1))
def _newest_named(alias: str) -> tuple[SpendLogRow, ...]:
return (*nameless_rows(150), named_row(150, alias))
def _both_edges_named(alias: str) -> tuple[SpendLogRow, ...]:
return (named_row(0, alias), *nameless_rows(150, 1), named_row(151, alias))
def _named_after_one_hundred(alias: str) -> tuple[SpendLogRow, ...]:
return (*nameless_rows(100), named_row(100, alias), *nameless_rows(99, 101))
def _named_after_ninety_nine(alias: str) -> tuple[SpendLogRow, ...]:
return (*nameless_rows(99), named_row(99, alias), *nameless_rows(100, 100))
def _named_only_in_the_middle(alias: str) -> tuple[SpendLogRow, ...]:
return (*nameless_rows(100), named_row(100, alias), *nameless_rows(100, 101))
def _renamed_and_renamed_back(alias: str, other: str) -> tuple[SpendLogRow, ...]:
return (
named_row(0, alias),
*nameless_rows(100, 1),
named_row(101, other),
*nameless_rows(100, 102),
named_row(202, alias),
)
def _team_in_the_column(team: str) -> SpendLogRow:
return SpendLogRow(started_at(0), {}, team_id=team)
def _team_in_the_metadata(team: str) -> SpendLogRow:
return SpendLogRow(started_at(0), {"user_api_key_team_id": team})
def _user_in_the_column(user: str) -> SpendLogRow:
return SpendLogRow(started_at(0), {}, user=user)
def _user_in_the_metadata(user: str) -> SpendLogRow:
return SpendLogRow(started_at(0), {"user_api_key_user_id": user})
def _activity_on_route(gateway: Gateway, route: Route, api_key: str, entity: str) -> httpx.Response:
filters: Final = {} if route.entity_filter is None else {route.entity_filter: entity}
return activity_of_key(gateway, route.path, api_key, **filters)
def _reported_aliases(response: httpx.Response, api_key: str) -> tuple[JsonValue, ...]:
if response.status_code != 200:
return ()
return tuple(
object_value(object_value(record)["metadata"])["key_alias"]
for record in records_of_key(object_value(response.json()), api_key)
)
def _names_the_key(api_key: str, alias: str) -> Callable[[httpx.Response], bool]:
def names(response: httpx.Response) -> bool:
reported: Final = _reported_aliases(response, api_key)
return bool(reported) and frozenset(reported) == frozenset((alias,))
return names
@contextmanager
def _proxy_on(gateway: Gateway, directory: Path, database_url: str, *, workers: int = 1) -> Iterator[OwnedProxy]:
with owned_proxy_process(
gateway,
directory,
{"DATABASE_URL": database_url},
remove_environment=("DATABASE_URL_READ_REPLICA",),
workers=workers,
) as owned:
yield owned
def _owner_on(candidate: Gateway) -> tuple[str, str]:
owner: Final = f"integration-{uuid.uuid4().hex}"
email: Final = f"{owner}@example.com"
candidate.post("/user/new", {"user_id": owner, "user_email": email, "auto_create_key": False})
return owner, email
@pytest.mark.parametrize("route", ROUTES, ids=lambda route: route.path.strip("/").replace("/", "_"))
def test_alias_named_only_by_a_spend_log_is_reported_on_every_daily_activity_route(
gateway: Gateway, route: Route
) -> None:
api_key: Final = digest_no_key_table_holds()
alias: Final = _alias()
entity: Final = f"integration-entity-{uuid.uuid4().hex}"
with gateway.scenario() as scenario:
owner, email = user_with_an_email(scenario)
entity_rows: Final = (
() if route.table == USER_SPEND else (seeded_row(route.table, route.entity_column, entity, api_key, DAY),)
)
filters: Final = {} if route.entity_filter is None else {route.entity_filter: entity}
with (
daily_rows((user_row(owner, api_key, DAY), *entity_rows)),
spend_logs_of_key(api_key, (named_row(0, alias),)),
):
assert_key_reported(
activity_of_key(gateway, route.path, api_key, **filters),
api_key,
DAY,
key_metadata(alias=alias, user=owner, email=email),
seeded_metrics(1),
)
@pytest.mark.parametrize(
"layout",
(
pytest.param(_named_between_fifty_and_fifty, id="named_between_50_and_50_nameless"),
pytest.param(_oldest_named, id="oldest_named_150_nameless_newer"),
pytest.param(_newest_named, id="newest_named_150_nameless_older"),
pytest.param(_both_edges_named, id="both_edges_named_150_nameless_between"),
pytest.param(_named_after_one_hundred, id="100_nameless_named_99_nameless"),
pytest.param(_named_after_ninety_nine, id="99_nameless_named_100_nameless"),
),
)
def test_alias_on_an_edge_of_the_window_is_reported_whatever_surrounds_it(
gateway: Gateway, layout: Callable[[str], tuple[SpendLogRow, ...]]
) -> None:
api_key: Final = digest_no_key_table_holds()
alias: Final = _alias()
with gateway.scenario() as scenario:
owner, email = user_with_an_email(scenario)
with daily_rows((user_row(owner, api_key, DAY),)), spend_logs_of_key(api_key, layout(alias)):
assert_key_reported(
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
api_key,
DAY,
key_metadata(alias=alias, user=owner, email=email),
seeded_metrics(1),
)
def test_alias_named_only_in_the_middle_of_two_hundred_nameless_rows_is_not_picked_up(gateway: Gateway) -> None:
api_key: Final = digest_no_key_table_holds()
with gateway.scenario() as scenario:
owner, email = user_with_an_email(scenario)
with (
daily_rows((user_row(owner, api_key, DAY),)),
spend_logs_of_key(api_key, _named_only_in_the_middle(_alias())),
):
assert_key_reported(
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
api_key,
DAY,
key_metadata(user=owner, email=email),
seeded_metrics(1),
)
def test_key_renamed_and_renamed_back_is_reported_with_the_alias_on_both_edges(gateway: Gateway) -> None:
api_key: Final = digest_no_key_table_holds()
alias: Final = _alias()
with gateway.scenario() as scenario:
owner, email = user_with_an_email(scenario)
rows: Final = _renamed_and_renamed_back(alias, _alias())
with daily_rows((user_row(owner, api_key, DAY),)), spend_logs_of_key(api_key, rows):
assert_key_reported(
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
api_key,
DAY,
key_metadata(alias=alias, user=owner, email=email),
seeded_metrics(1),
)
@pytest.mark.parametrize(
"spend_log_of_team",
(
pytest.param(_team_in_the_column, id="team_id_column"),
pytest.param(_team_in_the_metadata, id="team_id_in_metadata"),
),
)
def test_team_named_only_by_a_spend_log_is_reported_next_to_the_daily_owner(
gateway: Gateway, spend_log_of_team: Callable[[str], SpendLogRow]
) -> None:
api_key: Final = digest_no_key_table_holds()
team: Final = f"integration-team-{uuid.uuid4().hex}"
with gateway.scenario() as scenario:
owner, email = user_with_an_email(scenario)
with daily_rows((user_row(owner, api_key, DAY),)), spend_logs_of_key(api_key, (spend_log_of_team(team),)):
assert_key_reported(
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
api_key,
DAY,
key_metadata(team=team, user=owner, email=email),
seeded_metrics(1),
)
@pytest.mark.parametrize(
"spend_log_of_user",
(
pytest.param(_user_in_the_column, id="user_column"),
pytest.param(_user_in_the_metadata, id="user_id_in_metadata"),
),
)
def test_user_named_by_a_spend_log_beats_the_owner_the_daily_rows_name(
gateway: Gateway, spend_log_of_user: Callable[[str], SpendLogRow]
) -> None:
api_key: Final = digest_no_key_table_holds()
with gateway.scenario() as scenario:
daily_owner, _ = user_with_an_email(scenario)
log_user, log_email = user_with_an_email(scenario)
with (
daily_rows((user_row(daily_owner, api_key, DAY),)),
spend_logs_of_key(api_key, (spend_log_of_user(log_user),)),
):
assert_key_reported(
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
api_key,
DAY,
key_metadata(user=log_user, email=log_email),
seeded_metrics(1),
)
def test_hashed_jwt_digest_is_named_by_its_spend_log(gateway: Gateway) -> None:
api_key: Final = f"hashed-jwt-{sha256(uuid.uuid4().bytes).hexdigest()}"
alias: Final = _alias()
with gateway.scenario() as scenario:
owner, email = user_with_an_email(scenario)
with daily_rows((user_row(owner, api_key, DAY),)), spend_logs_of_key(api_key, (named_row(0, alias),)):
assert_key_reported(
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
api_key,
DAY,
key_metadata(alias=alias, user=owner, email=email),
seeded_metrics(1),
)
@pytest.mark.parametrize(
("started", "inside_the_window"),
(
pytest.param("2026-02-01 23:59:59", False, id="second_before_the_window"),
pytest.param("2026-02-02 00:00:00", True, id="first_second_of_the_window"),
pytest.param("2026-02-04 23:59:59", True, id="last_second_of_the_window"),
pytest.param("2026-02-05 00:00:00", False, id="first_second_after_the_window"),
),
)
def test_spend_log_names_the_key_only_from_one_day_before_to_two_days_after_the_read(
gateway: Gateway, started: str, inside_the_window: bool
) -> None:
api_key: Final = digest_no_key_table_holds()
alias: Final = _alias()
row: Final = SpendLogRow(started, {"user_api_key_alias": alias})
with gateway.scenario() as scenario:
owner, email = user_with_an_email(scenario)
with daily_rows((user_row(owner, api_key, DAY),)), spend_logs_of_key(api_key, (row,)):
assert_key_reported(
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
api_key,
DAY,
key_metadata(alias=alias if inside_the_window else None, user=owner, email=email),
seeded_metrics(1),
)
def test_two_aliases_on_the_two_edges_leave_the_key_unnamed(gateway: Gateway) -> None:
api_key: Final = digest_no_key_table_holds()
rows: Final = (named_row(0, _alias()), *nameless_rows(150, 1), named_row(151, _alias()))
with gateway.scenario() as scenario:
owner, email = user_with_an_email(scenario)
with daily_rows((user_row(owner, api_key, DAY),)), spend_logs_of_key(api_key, rows):
assert_key_reported(
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
api_key,
DAY,
key_metadata(user=owner, email=email),
seeded_metrics(1),
)
@pytest.mark.parametrize(
"unnamed_rows",
(
pytest.param((SpendLogRow(started_at(0), {"user_api_key_alias": ""}),), id="empty_string_alias"),
pytest.param(
(SpendLogRow(started_at(0), ["x"]), SpendLogRow(started_at(1), "x")), id="array_then_string_metadata"
),
),
)
def test_rows_without_a_usable_alias_do_not_hide_the_named_row_after_them(
gateway: Gateway, unnamed_rows: tuple[SpendLogRow, ...]
) -> None:
api_key: Final = digest_no_key_table_holds()
alias: Final = _alias()
rows: Final = (*unnamed_rows, named_row(len(unnamed_rows), alias))
with gateway.scenario() as scenario:
owner, email = user_with_an_email(scenario)
with daily_rows((user_row(owner, api_key, DAY),)), spend_logs_of_key(api_key, rows):
assert_key_reported(
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
api_key,
DAY,
key_metadata(alias=alias, user=owner, email=email),
seeded_metrics(1),
)
@pytest.mark.parametrize(
"stored_alias",
(
pytest.param(123, id="json_int"),
pytest.param(["a"], id="json_list"),
pytest.param("a" * 5000, id="five_kb_string"),
),
)
def test_alias_of_an_unexpected_shape_is_reported_as_postgres_renders_it(
gateway: Gateway, stored_alias: JsonValue
) -> None:
api_key: Final = digest_no_key_table_holds()
row: Final = SpendLogRow(started_at(0), {"user_api_key_alias": stored_alias})
with gateway.scenario() as scenario:
owner, email = user_with_an_email(scenario)
with daily_rows((user_row(owner, api_key, DAY),)), spend_logs_of_key(api_key, (row,)) as request_ids:
rendered: Final = read_rows(ALIAS_OF_ONE_SPEND_LOG, (request_ids[0],))[0]["alias"]
assert isinstance(rendered, str) and rendered, rendered
assert_key_reported(
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
api_key,
DAY,
key_metadata(alias=rendered, user=owner, email=email),
seeded_metrics(1),
)
@pytest.mark.timeout(300)
def test_alias_found_once_is_served_from_the_cache_for_the_same_window_only(gateway: Gateway, tmp_path: Path) -> None:
api_key: Final = digest_no_key_table_holds()
alias: Final = _alias()
with scratch_database() as database_url, _proxy_on(gateway, tmp_path, database_url) as owned:
owner, email = _owner_on(owned.gateway)
rows: Final = (user_row(owner, api_key, DAY), user_row(owner, api_key, DAY_OUTSIDE_THE_WINDOW))
with daily_rows(rows, database_url=database_url):
with spend_logs_of_key(api_key, (named_row(0, alias),), database_url=database_url):
first: Final = activity_of_key(owned.gateway, AGGREGATED_USER_ACTIVITY, api_key)
cached: Final = activity_of_key(owned.gateway, AGGREGATED_USER_ACTIVITY, api_key)
other_window: Final = owned.gateway.request(
"GET",
AGGREGATED_USER_ACTIVITY,
params={"start_date": DAY_OUTSIDE_THE_WINDOW, "end_date": DAY_OUTSIDE_THE_WINDOW, "api_key": api_key},
)
named: Final = key_metadata(alias=alias, user=owner, email=email)
assert_key_reported(first, api_key, DAY, named, seeded_metrics(1))
assert_key_reported(cached, api_key, DAY, named, seeded_metrics(1))
assert_key_reported(
other_window, api_key, DAY_OUTSIDE_THE_WINDOW, key_metadata(user=owner, email=email), seeded_metrics(1)
)
@pytest.mark.timeout(300)
def test_alias_logged_after_a_cached_miss_shows_once_the_miss_expires(gateway: Gateway, tmp_path: Path) -> None:
api_key: Final = digest_no_key_table_holds()
alias: Final = _alias()
with scratch_database() as database_url, _proxy_on(gateway, tmp_path, database_url) as owned:
owner, email = _owner_on(owned.gateway)
with daily_rows((user_row(owner, api_key, DAY),), database_url=database_url):
missed: Final = activity_of_key(owned.gateway, AGGREGATED_USER_ACTIVITY, api_key)
with spend_logs_of_key(api_key, (named_row(0, alias),), database_url=database_url):
named: Final = eventually(
lambda: activity_of_key(owned.gateway, AGGREGATED_USER_ACTIVITY, api_key),
_names_the_key(api_key, alias),
seconds=CACHED_MISS_CLEARS_WITHIN_SECONDS,
)
assert_key_reported(missed, api_key, DAY, key_metadata(user=owner, email=email), seeded_metrics(1))
assert_key_reported(named, api_key, DAY, key_metadata(alias=alias, user=owner, email=email), seeded_metrics(1))
@pytest.mark.timeout(300)
def test_alias_lookup_gives_up_while_spend_logs_are_locked_and_answers_once_they_are_not(
gateway: Gateway, tmp_path: Path
) -> None:
api_key: Final = digest_no_key_table_holds()
alias: Final = _alias()
with scratch_database() as database_url, _proxy_on(gateway, tmp_path, database_url, workers=2) as owned:
owner, email = _owner_on(owned.gateway)
with (
daily_rows((user_row(owner, api_key, DAY),), database_url=database_url),
spend_logs_of_key(api_key, (named_row(0, alias),), database_url=database_url),
):
with locked_table(SPEND_LOGS_TABLE, database_url=database_url):
started: Final = time.monotonic()
locked: Final = activity_of_key(owned.gateway, AGGREGATED_USER_ACTIVITY, api_key)
waited: Final = time.monotonic() - started
unlocked: Final = eventually(
lambda: activity_of_key(owned.gateway, AGGREGATED_USER_ACTIVITY, api_key),
_names_the_key(api_key, alias),
seconds=CACHED_MISS_CLEARS_WITHIN_SECONDS,
)
assert waited < GIVES_UP_WITHIN_SECONDS, waited
assert_key_reported(locked, api_key, DAY, key_metadata(user=owner, email=email), seeded_metrics(1))
assert_key_reported(unlocked, api_key, DAY, key_metadata(alias=alias, user=owner, email=email), seeded_metrics(1))
def test_concurrent_reads_over_every_route_all_name_a_fresh_key(gateway: Gateway) -> None:
api_key: Final = digest_no_key_table_holds()
alias: Final = _alias()
entity: Final = f"integration-entity-{uuid.uuid4().hex}"
entity_columns: Final = {route.table: route.entity_column for route in ROUTES if route.table != USER_SPEND}
with gateway.scenario() as scenario:
owner, email = user_with_an_email(scenario)
rows: Final = (
user_row(owner, api_key, DAY),
*(seeded_row(table, column, entity, api_key, DAY) for table, column in entity_columns.items()),
)
with (
daily_rows(rows),
spend_logs_of_key(api_key, (named_row(0, alias),)),
ThreadPoolExecutor(CONCURRENT_READS) as pool,
):
reads: Final = tuple(
pool.submit(_activity_on_route, gateway, ROUTES[index % len(ROUTES)], api_key, entity)
for index in range(CONCURRENT_READS)
)
responses: Final = tuple(read.result() for read in reads)
for response in responses:
assert_key_reported(
response, api_key, DAY, key_metadata(alias=alias, user=owner, email=email), seeded_metrics(1)
)

View file

@ -12,7 +12,7 @@ from typing import Final
import httpx
import pytest
from integration._support.client import Gateway, Scenario, eventually
from integration._support.client import Gateway, Scenario, eventually, string_value
from integration._support.daily_activity import (
AGGREGATED_USER_ACTIVITY,
DAY,
@ -25,6 +25,7 @@ from integration._support.daily_activity import (
daily_rows,
key_metadata,
key_no_key_table_holds,
purge_key_from_the_key_tables,
seeded_metrics,
seeded_row,
user_row,
@ -42,6 +43,10 @@ REQUESTS_OF_KEY: Final = (
'SELECT COALESCE(SUM(api_requests), 0)::int AS requests FROM "LiteLLM_DailyUserSpend" '
"WHERE api_key=%s AND user_id=%s"
)
NAMED_SPEND_LOGS_OF_KEY: Final = (
'SELECT COUNT(*)::int AS named FROM "LiteLLM_SpendLogs" '
"WHERE api_key=%s AND NULLIF(metadata->>'user_api_key_alias', '') IS NOT NULL"
)
UNIFIED_ENDPOINTS: Final = ("/v1/chat/completions", "/v1/messages", "/v1/responses")
REQUESTS_OF_A_BURST: Final = 21
READS_DURING_A_BURST: Final = 30
@ -210,6 +215,14 @@ def _wait_for_requests(api_key: str, user: str, requests: int) -> None:
)
def _wait_for_named_spend_logs(api_key: str, requests: int) -> None:
eventually(
lambda: read_rows(NAMED_SPEND_LOGS_OF_KEY, (api_key,)),
lambda rows: rows[0]["named"] == requests,
seconds=70,
)
def _cli_session_token(user: str, team: str) -> str:
cli_user: Final = LiteLLM_UserTable(user_id=user, user_role="internal_user", teams=[team], models=[])
return ExperimentalUIJWTToken.get_cli_jwt_auth_token(user_info=cli_user, team_id=team, team_alias="cli-team")
@ -244,6 +257,39 @@ def test_key_used_on_every_unified_endpoint_is_reported_with_its_own_alias_and_u
)
def test_key_purged_from_the_key_tables_is_reported_with_the_alias_its_spend_logs_name(gateway: Gateway) -> None:
prompts: Final = (_prompt(), _prompt(), _prompt())
with wire_server(_provider) as wire, gateway.scenario() as scenario:
model: Final = _priced_model(scenario, wire.url)
owner, email = user_with_an_email(scenario)
alias: Final = f"integration-alias-{uuid.uuid4().hex}"
generated: Final = gateway.post("/key/generate", {"user_id": owner, "key_alias": alias, "models": [model]})
key: Final = string_value(generated["key"])
stored: Final = sha256(key.encode()).hexdigest()
try:
answers: Final = tuple(
gateway.request("POST", endpoint, _request_body(endpoint, model, prompt), key=key)
for endpoint, prompt in zip(UNIFIED_ENDPOINTS, prompts, strict=True)
)
assert [answer.status_code for answer in answers] == [200, 200, 200], [answer.text for answer in answers]
received: Final = _sent_for_callers(wire.drain())
assert [request.target for request in received] == [
"/v1/chat/completions",
"/v1/responses",
"/v1/responses",
]
_wait_for_requests(stored, owner, 3)
_wait_for_named_spend_logs(stored, 3)
finally:
purge_key_from_the_key_tables(stored)
assert_key_owner_and_totals(
_activity_around_today(gateway, stored),
stored,
key_metadata(alias=alias, user=owner, email=email, exists=False),
_totals_of_requests(3),
)
def test_cli_session_spend_is_reported_with_the_user_and_team_of_the_session(
gateway: Gateway, monkeypatch: pytest.MonkeyPatch
) -> None:

View file

@ -690,7 +690,7 @@ def test_langfuse_logging_tool_calling():
]
response = litellm.completion(
model="gpt-3.5-turbo-1106",
model="gpt-6-luna",
messages=messages,
tools=tools,
tool_choice="auto", # auto is default, but we'll be explicit
@ -698,6 +698,8 @@ def test_langfuse_logging_tool_calling():
print("\nLLM Response1:\n", response)
response_message = response.choices[0].message
tool_calls = response.choices[0].message.tool_calls
assert response.choices[0].message.tool_calls
assert all(call.function.name == "get_current_weather" for call in response.choices[0].message.tool_calls)
# test_langfuse_logging_tool_calling()

View file

@ -39,7 +39,7 @@ def get_current_weather(location, unit="fahrenheit"):
@pytest.mark.parametrize(
"model",
[
"gpt-3.5-turbo-1106",
"gpt-6-luna",
"mistral/mistral-large-latest",
"claude-haiku-4-5-20251001",
"gemini/gemini-2.5-flash-lite",
@ -386,7 +386,7 @@ def test_parallel_function_call_stream():
}
]
response = litellm.completion(
model="gpt-3.5-turbo-1106",
model="gpt-6-luna",
messages=messages,
tools=tools,
stream=True,
@ -435,7 +435,7 @@ def test_parallel_function_call_stream():
) # extend conversation with function response
print(f"messages: {messages}")
second_response = litellm.completion(
model="gpt-3.5-turbo-1106", messages=messages, temperature=0.2, seed=22
model="gpt-6-luna", messages=messages, temperature=0.2, seed=22, reasoning_effort="none"
) # get a new response from the model where it can see the function response
print("second response\n", second_response)
return second_response

View file

@ -1,14 +1,18 @@
# What is this?
## Unit testing for the 'get_model_info()' function
import os
import re
from collections.abc import Collection, Mapping
from typing import List, Dict, Any
from typing import List, Dict, Any, Final, Literal
import pytest
import litellm
from litellm import get_model_info
from litellm.llms.bedrock.common_utils import BedrockModelInfo
from litellm.types.utils import ModelInfoBase
from litellm.utils import _invalidate_model_cost_lowercase_map
from unittest.mock import MagicMock, patch
@ -116,26 +120,31 @@ def test_get_model_info_ft_model_with_provider_prefix():
def _enforce_bedrock_converse_models(
model_cost: List[Dict[str, Any]], whitelist_models: List[str]
):
model_cost: Mapping[str, ModelInfoBase], whitelist_models: Collection[str]
) -> None:
"""
Assert all new bedrock chat models are added as `bedrock_converse` unless explicitly whitelisted.
Assert unlisted Bedrock chat models declare or inherit Converse routing.
"""
# Check for unwhitelisted models
for model, info in litellm.model_cost.items():
for model, info in model_cost.items():
if (
info["litellm_provider"] == "bedrock"
and info["mode"] == "chat"
and model not in whitelist_models
and not (
(base_model := BedrockModelInfo.get_base_model(model)) != model
and model_cost.get(base_model, {}).get("litellm_provider") == "bedrock_converse"
and BedrockModelInfo.get_bedrock_route(model) == "converse"
)
):
raise AssertionError(
f"New bedrock chat model detected: {model}. Please set `litellm_provider='bedrock_converse'` for this model."
f"Unlisted Bedrock chat model does not route to Converse: {model}"
)
def test_model_info_bedrock_converse(monkeypatch):
"""
Assert all new bedrock chat models are added as `bedrock_converse` unless explicitly whitelisted.
Assert unlisted Bedrock chat models declare or inherit Converse routing.
This ensures they are automatically routed to the converse endpoint.
"""
@ -173,7 +182,7 @@ def test_model_info_bedrock_converse_enforcement(monkeypatch):
whitelist_models = [line.strip() for line in file.readlines()]
# Check for unwhitelisted models
with pytest.raises(AssertionError):
with pytest.raises(AssertionError, match=r"fake\.bedrock-chat-model"):
_enforce_bedrock_converse_models(
model_cost=litellm.model_cost, whitelist_models=whitelist_models
)
@ -181,6 +190,27 @@ def test_model_info_bedrock_converse_enforcement(monkeypatch):
pytest.skip("whitelisted_bedrock_models.txt not found")
@pytest.mark.parametrize("region", ("us-gov-east-1", "us-gov-west-1"))
@pytest.mark.parametrize("base_provider", ("bedrock_converse", "bedrock"))
def test_regional_bedrock_alias_requires_canonical_converse_metadata(
region: str, base_provider: Literal["bedrock_converse", "bedrock"]
) -> None:
base_model: Final = next(
model for model in sorted(litellm.bedrock_converse_models) if BedrockModelInfo.get_base_model(model) == model
)
model: Final = f"bedrock/{region}/{base_model}"
model_cost: Final[Mapping[str, ModelInfoBase]] = {
model: {"litellm_provider": "bedrock", "mode": "chat"},
base_model: {"litellm_provider": base_provider, "mode": "chat"},
}
assert BedrockModelInfo.get_bedrock_route(model) == "converse"
if base_provider == "bedrock":
with pytest.raises(AssertionError, match=re.escape(model)):
_enforce_bedrock_converse_models(model_cost, ())
return
_enforce_bedrock_converse_models(model_cost, ())
def test_get_model_info_custom_provider():
# Custom provider example copied from https://docs.litellm.ai/docs/providers/custom_llm_server:
import litellm

View file

@ -83,13 +83,15 @@ def test_lunary_with_tools():
]
response = litellm.completion(
model="gpt-3.5-turbo-1106",
model="gpt-6-luna",
messages=messages,
tools=tools,
tool_choice="auto", # auto is default, but we'll be explicit
)
response_message = response.choices[0].message
assert response.choices[0].message.tool_calls
assert all(call.function.name == "get_current_weather" for call in response.choices[0].message.tool_calls)
print("\nLLM Response:\n", response.choices[0].message)

View file

@ -157,6 +157,7 @@ async def test_create_mcp_server_direct():
# Mock server manager
mock_manager.add_server = mock.AsyncMock()
mock_manager.reload_servers_from_database = mock.AsyncMock()
mock_manager.get_mcp_server_by_id.return_value = None
# Set up test data
server_id = str(uuid.uuid4())

View file

@ -0,0 +1,559 @@
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import pytest
from fastapi import HTTPException
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import mcp_server_manager
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_UserTable, UserAPIKeyAuth
from litellm.proxy.auth import auth_checks
from litellm.types.agents import AgentResponse
from litellm.types.proxy.agent_identity import ManagedAgentContext
def actor(tools: tuple[str, ...] | None, *, delegated: bool = False) -> UserAPIKeyAuth:
permission: Final = LiteLLM_ObjectPermissionTable(
object_permission_id="agent-permissions",
mcp_servers=["slack", "linear"],
mcp_tool_permissions={"slack": list(tools)} if tools is not None else None,
)
agent: Final = AgentResponse(
agent_id="publisher",
agent_name="Publisher",
agent_card_params={},
object_permission=permission.model_dump(),
identity_managed=True,
)
auth: Final = UserAPIKeyAuth(agent_id=agent.agent_id)
auth.managed_agent_policy = agent
auth.managed_agent_context = ManagedAgentContext(
agent_id=agent.agent_id,
mode="delegated" if delegated else "autonomous",
user_id="human" if delegated else None,
)
return auth
@pytest.fixture(autouse=True)
def isolated_manager(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", mcp_server_manager.MCPServerManager())
monkeypatch.setattr(proxy_server, "prisma_client", MagicMock())
@pytest.mark.asyncio
@pytest.mark.parametrize("tools", (None, (), ("read",), ("read", "write")))
async def test_autonomous_agent_uses_only_its_own_tool_grants(tools: tuple[str, ...] | None) -> None:
auth: Final = actor(tools)
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"slack", "linear"}
actual: Final = await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)
assert (frozenset(actual) if actual is not None else None) == (frozenset(tools) if tools is not None else None)
assert await MCPRequestHandler.get_allowed_tools_for_server("ungranted-server", auth) == []
@pytest.mark.asyncio
@pytest.mark.parametrize(
"agent_tools,user_tools,expected",
(
(None, ("read",), ("read",)),
(("read",), None, ("read",)),
(("read", "write"), ("read",), ("read",)),
(("read",), ("write",), ()),
((), None, ()),
),
)
async def test_delegated_server_and_tool_intersections(
monkeypatch: pytest.MonkeyPatch,
agent_tools: tuple[str, ...] | None,
user_tools: tuple[str, ...] | None,
expected: tuple[str, ...],
) -> None:
permission: Final = LiteLLM_ObjectPermissionTable(
object_permission_id="user-permissions",
mcp_servers=["slack", "user-only"],
mcp_tool_permissions={"slack": list(user_tools)} if user_tools is not None else None,
)
user: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=permission)
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=user))
auth: Final = actor(agent_tools, delegated=True)
assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == ["slack"]
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == list(expected)
assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == []
assert await MCPRequestHandler.get_allowed_tools_for_server("user-only", auth) == []
@pytest.mark.asyncio
async def test_unavailable_delegated_user_never_leaves_agent_permissions_unrestricted(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(side_effect=RuntimeError("DB unavailable")))
with pytest.raises(HTTPException) as failure:
await MCPRequestHandler.get_allowed_tools_for_server("slack", actor(None, delegated=True))
assert failure.value.status_code == 503
@pytest.mark.asyncio
@pytest.mark.parametrize("servers,expected", (((), ()), (("slack",), ("slack",)), (("user-only",), ())))
async def test_access_groups_cap_agent_servers_without_granting_new_ones(
monkeypatch: pytest.MonkeyPatch,
servers: tuple[str, ...],
expected: tuple[str, ...],
) -> None:
from litellm.proxy._types import LiteLLM_AccessGroupTable
group: Final = LiteLLM_AccessGroupTable(
access_group_id="group", access_group_name="Restricted", access_mcp_server_ids=list(servers)
)
monkeypatch.setattr(auth_checks, "get_access_object", AsyncMock(return_value=group))
auth: Final = actor(None)
assert auth.managed_agent_policy is not None
auth.managed_agent_policy = auth.managed_agent_policy.model_copy(update={"access_group_ids": ["group"]})
assert tuple(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == expected
if "slack" not in expected:
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == []
@pytest.mark.asyncio
@pytest.mark.parametrize("change", ["tools", "servers", "disabled", "outage"])
async def test_delegated_mcp_revokes_warm_human_policy_before_tool_execution(
monkeypatch: pytest.MonkeyPatch, change: str
) -> None:
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key
permission: Final = LiteLLM_ObjectPermissionTable(
object_permission_id="user-grant", mcp_servers=["slack"], mcp_tool_permissions={"slack": ["read", "write"]}
)
user: Final = LiteLLM_UserTable(
user_id="human", teams=[], organization_memberships=[], object_permission_id="user-grant"
)
cache: Final = UserApiKeyCache()
cache.set_cache("human", user)
cache.set_cache(object_permission_cache_key("user-grant"), permission)
client: Final = MagicMock()
client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=user)
client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission)
monkeypatch.setattr(proxy_server, "prisma_client", client)
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
auth: Final = actor(("read", "write"), delegated=True)
assert set(await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)) == {"read", "write"}
if change == "disabled":
client.writer_db.litellm_usertable.find_unique.return_value = user.model_copy(
update={"metadata": {"scim_active": False}}
)
elif change == "outage":
client.writer_db.litellm_usertable.find_unique.side_effect = RuntimeError("writer unavailable")
elif change == "servers":
client.writer_db.litellm_objectpermissiontable.find_unique.return_value = permission.model_copy(
update={"mcp_servers": [], "mcp_tool_permissions": {}}
)
else:
client.writer_db.litellm_objectpermissiontable.find_unique.return_value = permission.model_copy(
update={"mcp_tool_permissions": {"slack": ["read"]}}
)
if change in ("disabled", "outage"):
with pytest.raises(HTTPException):
await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)
else:
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == (
["read"] if change == "tools" else []
)
client.db.litellm_usertable.find_unique.assert_not_called()
client.db.litellm_objectpermissiontable.find_unique.assert_not_called()
def _server_row(server_id: str, access_groups: tuple[str, ...]) -> MagicMock:
row: Final = MagicMock()
row.server_id = server_id
row.mcp_access_groups = list(access_groups)
return row
def _toolset_row(server_id: str, tool_name: str) -> MagicMock:
row: Final = MagicMock()
row.tools = [{"server_id": server_id, "tool_name": tool_name}]
return row
@pytest.mark.asyncio
@pytest.mark.parametrize("change", ["tool", "server", "outage"])
async def test_autonomous_agent_toolset_and_access_group_revocations_bind_on_the_next_request(
monkeypatch: pytest.MonkeyPatch, change: str
) -> None:
"""The agent's entitlements are read through the shared toolset and access-group resolvers. Once the
writer revokes a tool or drops the server from the group, the next managed request must be denied
even though the legacy cache still holds the warm grant and the replica still shows the old rows"""
from litellm.caching.caching import DualCache
from litellm.proxy._experimental.mcp_server import toolset_db
warm_toolset: Final = _toolset_row("slack", "read")
list_toolsets: Final = AsyncMock(return_value=[warm_toolset])
monkeypatch.setattr(toolset_db, "list_mcp_toolsets", list_toolsets)
client: Final = MagicMock()
client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[_server_row("linear", ("grp",))])
client.writer_db.litellm_mcpservertable.find_many = AsyncMock(return_value=[_server_row("linear", ("grp",))])
monkeypatch.setattr(proxy_server, "prisma_client", client)
monkeypatch.setattr(proxy_server, "user_api_key_cache", DualCache())
permission: Final = LiteLLM_ObjectPermissionTable(
object_permission_id="agent-permissions", mcp_toolsets=["ts"], mcp_access_groups=["grp"]
)
auth: Final = actor(None)
assert auth.managed_agent_policy is not None
auth.managed_agent_policy = auth.managed_agent_policy.model_copy(
update={"object_permission": permission.model_dump()}
)
auth.requires_fresh_policy = True
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"slack", "linear"}
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == ["read"]
if change == "tool":
list_toolsets.return_value = [_toolset_row("slack", "other")]
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == ["other"]
elif change == "server":
client.writer_db.litellm_mcpservertable.find_many.return_value = []
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"slack"}
assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == []
else:
list_toolsets.side_effect = RuntimeError("writer unavailable")
with pytest.raises(HTTPException) as failure:
await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)
assert failure.value.status_code == 503
for call in list_toolsets.await_args_list:
assert call.kwargs["use_writer"] is True, "managed agent toolsets must be read from the writer"
client.db.litellm_mcpservertable.find_many.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.parametrize("role", ["proxy_admin", "proxy_admin_viewer", "internal_user"])
@pytest.mark.parametrize("open_channel", ["none", "operator", "submitted"])
@pytest.mark.parametrize("has_grant", [True, False])
@pytest.mark.parametrize("agent_tools", [("read", "write"), None])
async def test_delegated_mcp_uses_explicit_team_grants_even_for_dashboard_admins(
monkeypatch: pytest.MonkeyPatch,
role: str,
open_channel: str,
has_grant: bool,
agent_tools: tuple[str, ...] | None,
) -> None:
from litellm.proxy._types import LiteLLM_TeamTable
from litellm.types.mcp_server.mcp_server_manager import MCPServer
manager: Final = mcp_server_manager.global_mcp_server_manager
manager.registry = {
name: MCPServer(
server_id=name,
name=name,
transport="http",
url="https://example.com/mcp",
allow_all_keys=open_channel == "operator",
)
for name in ("slack", "linear")
}
from litellm.proxy._experimental.mcp_server import db
monkeypatch.setattr(
db,
"get_active_submitted_mcp_server_ids_for_user",
AsyncMock(return_value=["slack", "linear"] if open_channel == "submitted" else []),
)
permission: Final = LiteLLM_ObjectPermissionTable(
object_permission_id="team-grant", mcp_servers=["slack"], mcp_tool_permissions={"slack": ["read"]}
)
user: Final = LiteLLM_UserTable(
user_id="human", user_role=role, teams=["team"] if has_grant else [], organization_memberships=[]
)
team: Final = LiteLLM_TeamTable(
team_id="team",
models=[],
members_with_roles=[{"user_id": "human", "role": "user"}],
object_permission_id="team-grant",
)
client: Final = MagicMock()
client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=user)
client.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=team)
client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission)
monkeypatch.setattr(proxy_server, "prisma_client", client)
auth: Final = actor(agent_tools, delegated=True)
auth.team_id = "team"
assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == (["slack"] if has_grant else [])
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == (["read"] if has_grant else [])
assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == []
admitted: Final = await MCPRequestHandler.reload_admitted_user("human", requires_fresh_policy=True)
assert admitted.user_role == role
@pytest.mark.asyncio
async def test_explicit_grants_never_fall_back_to_open_servers_on_resolution_failure(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from litellm.proxy._experimental.mcp_server import db
from litellm.types.mcp_server.mcp_server_manager import MCPServer
manager: Final = mcp_server_manager.global_mcp_server_manager
manager.registry = {"slack": MCPServer(server_id="slack", name="slack", transport="http", allow_all_keys=True)}
monkeypatch.setattr(db, "get_active_submitted_mcp_server_ids_for_user", AsyncMock(return_value=["slack"]))
auth: Final = UserAPIKeyAuth(user_id="human")
auth.mcp_explicit_grants_only = True
with pytest.MonkeyPatch.context() as patcher:
patcher.setattr(MCPRequestHandler, "get_mcp_server_access", AsyncMock(side_effect=RuntimeError("unavailable")))
assert await manager.get_allowed_mcp_servers(auth) == []
auth.mcp_explicit_grants_only = False
assert await manager.get_allowed_mcp_servers(auth) == ["slack"]
@pytest.mark.asyncio
async def test_absent_agent_policy_and_missing_delegated_subject_grant_no_servers() -> None:
from litellm.proxy._experimental.mcp_server.auth.managed_agent_access import managed_agent_servers
assert await managed_agent_servers(UserAPIKeyAuth()) == ()
auth: Final = actor(None, delegated=True)
assert auth.managed_agent_context is not None
auth.managed_agent_context = auth.managed_agent_context.model_copy(update={"user_id": None})
assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == []
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == []
@pytest.mark.asyncio
async def test_tool_policy_outage_after_server_admission_fails_closed(monkeypatch: pytest.MonkeyPatch) -> None:
permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="human-grant", mcp_servers=["slack"])
user: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=permission)
monkeypatch.setattr(
auth_checks, "get_user_object", AsyncMock(side_effect=[user, RuntimeError("tool lookup unavailable")])
)
with pytest.raises(HTTPException) as failure:
await MCPRequestHandler.get_allowed_tools_for_server("slack", actor(None, delegated=True))
assert failure.value.status_code == 503
@pytest.mark.asyncio
@pytest.mark.parametrize("role", (None, "proxy_admin", "internal_user"))
@pytest.mark.parametrize("scoped", (False, True))
async def test_manager_preserves_managed_server_grants_across_open_channels(
monkeypatch: pytest.MonkeyPatch, role: str | None, scoped: bool
) -> None:
from litellm.proxy._experimental.mcp_server import db
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPServerAccess
from litellm.types.mcp_server.mcp_server_manager import MCPServer
manager: Final = mcp_server_manager.global_mcp_server_manager
manager.registry = {
"open": MCPServer(server_id="open", name="open", transport="http", allow_all_keys=True),
"submitted": MCPServer(server_id="submitted", name="submitted", transport="http"),
"passthrough": MCPServer(
server_id="passthrough", name="passthrough", transport="http", auth_type="true_passthrough"
),
}
monkeypatch.setattr(db, "get_active_submitted_mcp_server_ids_for_user", AsyncMock(return_value=["submitted"]))
auth: Final = actor(None)
auth.user_role = role
assert not auth.mcp_explicit_grants_only
access: Final = MCPServerAccess(server_ids=("slack", "open")) if scoped else None
assert set(await manager.get_allowed_mcp_servers(auth, access=access)) == (
{"slack"} if scoped else {"slack", "linear"}
)
@pytest.mark.asyncio
async def test_manager_does_not_replace_managed_policy_failure_with_open_servers(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from litellm.types.mcp_server.mcp_server_manager import MCPServer
manager: Final = mcp_server_manager.global_mcp_server_manager
manager.registry = {"open": MCPServer(server_id="open", name="open", transport="http", allow_all_keys=True)}
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(side_effect=RuntimeError("writer unavailable")))
with pytest.raises(HTTPException) as failure:
await manager.get_allowed_mcp_servers(actor(None, delegated=True))
assert failure.value.status_code == 503
@pytest.mark.asyncio
async def test_inline_tool_grant_admits_its_server_without_widening_tools() -> None:
auth: Final = actor(("read",))
assert auth.managed_agent_policy is not None
auth.managed_agent_policy = auth.managed_agent_policy.model_copy(
update={"object_permission": {"object_permission_id": "tools", "mcp_tool_permissions": {"slack": ["read"]}}}
)
assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == ["slack"]
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == ["read"]
assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == []
@pytest.mark.asyncio
@pytest.mark.parametrize("selected_team", (None, "selected"))
@pytest.mark.parametrize("selected_grant", (False, True))
async def test_delegation_never_borrows_another_teams_server_or_tools(
monkeypatch: pytest.MonkeyPatch, selected_team: str | None, selected_grant: bool
) -> None:
from litellm.proxy._types import LiteLLM_TeamTable
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
user: Final = LiteLLM_UserTable(user_id="human", teams=["selected", "other"], organization_memberships=[])
permission: Final = LiteLLM_ObjectPermissionTable(
object_permission_id="selected-grant",
mcp_servers=["slack"] if selected_grant else [],
mcp_tool_permissions={"slack": ["read"]} if selected_grant else {},
)
teams: Final = {
name: LiteLLM_TeamTable(
team_id=name,
models=[],
members_with_roles=[{"user_id": "human", "role": "user"}],
object_permission=permission if name == "selected" else LiteLLM_ObjectPermissionTable(
object_permission_id="other-grant", mcp_servers=["slack", "linear"]
),
)
for name in ("selected", "other")
}
async def get_team(team_id: str, **kwargs: object) -> LiteLLM_TeamTable:
return teams[team_id]
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=user))
monkeypatch.setattr(auth_checks, "get_team_object", get_team)
monkeypatch.setattr(proxy_server, "user_api_key_cache", UserApiKeyCache())
auth: Final = actor(None, delegated=True)
auth.team_id = selected_team
expected: Final = ["slack"] if selected_team and selected_grant else []
assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == expected
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == (["read"] if expected else [])
assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == []
ordinary: Final = await MCPRequestHandler.reload_admitted_user("human", requires_fresh_policy=True)
assert set(await MCPRequestHandler.resolve_admitted_subject_servers(ordinary)) == {"slack", "linear"}
@pytest.mark.asyncio
@pytest.mark.parametrize("entitlement", ("group", "toolset"))
async def test_managed_mcp_rejects_unavailable_authoritative_entitlements(
monkeypatch: pytest.MonkeyPatch, entitlement: str
) -> None:
client: Final = MagicMock()
client.writer_db.litellm_mcpservertable.find_many = AsyncMock(side_effect=RuntimeError("writer unavailable"))
client.writer_db.litellm_mcptoolsettable.find_many = AsyncMock(side_effect=RuntimeError("writer unavailable"))
monkeypatch.setattr(proxy_server, "prisma_client", client)
permission: Final = LiteLLM_ObjectPermissionTable(
object_permission_id="entitlements",
mcp_access_groups=["group"] if entitlement == "group" else [],
mcp_toolsets=["toolset"] if entitlement == "toolset" else [],
)
auth: Final = actor(None)
assert auth.managed_agent_policy is not None
auth.managed_agent_policy = auth.managed_agent_policy.model_copy(update={"object_permission": permission.model_dump()})
auth.requires_fresh_policy = True
with pytest.raises(HTTPException) as failure:
await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)
assert failure.value.status_code == 503
client.db.litellm_mcpservertable.find_many.assert_not_called()
client.db.litellm_mcptoolsettable.find_many.assert_not_called()
@pytest.mark.asyncio
async def test_managed_agent_mcp_access_is_capped_at_the_invoking_callers_grants(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""The managed MCP path must honour the agent_caller ceiling the same way the unmanaged path does:
the agent's own policy grants slack and linear, but the team echoed back on the request reaches
only slack, so the agent may use slack alone."""
from litellm.proxy._types import AgentCaller
monkeypatch.setattr(
MCPRequestHandler,
"_get_allowed_mcp_servers_for_team",
AsyncMock(return_value=["slack"]),
)
monkeypatch.setattr(
MCPRequestHandler,
"_apply_user_server_ceiling",
AsyncMock(side_effect=lambda servers, _auth: (tuple(servers), False)),
)
monkeypatch.setattr(
MCPRequestHandler,
"_get_team_object_permission",
AsyncMock(
return_value=LiteLLM_ObjectPermissionTable(
object_permission_id="caller-team-permissions",
mcp_servers=["slack"],
mcp_tool_permissions={"slack": ["read"]},
)
),
)
monkeypatch.setattr(
MCPRequestHandler,
"_apply_user_tool_ceiling",
AsyncMock(side_effect=lambda tools, _server_id, _auth: tools),
)
auth: Final = actor(("read", "write"))
auth.agent_caller = AgentCaller(user_id="alice", team_id="callers")
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"slack"}
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == ["read"]
@pytest.mark.asyncio
@pytest.mark.parametrize("fresh", [False, True])
@pytest.mark.parametrize("caller_kind", ["team", "user"])
async def test_caller_mcp_revocation_uses_fresh_policy(
monkeypatch: pytest.MonkeyPatch, fresh: bool, caller_kind: str,
) -> None:
from litellm.proxy._types import LiteLLM_TeamTable
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key
from litellm.types.agents import AgentCaller
cached_permission: Final = LiteLLM_ObjectPermissionTable(
object_permission_id="caller-permission", mcp_servers=["slack", "linear"],
mcp_tool_permissions={"slack": ["read", "write"]},
)
current_permission: Final = LiteLLM_ObjectPermissionTable(
object_permission_id="caller-permission", mcp_servers=["slack"],
mcp_tool_permissions={"slack": ["read"]},
)
team: Final = LiteLLM_TeamTable(
team_id="caller", object_permission_id="caller-permission", object_permission=current_permission,
)
user: Final = LiteLLM_UserTable(
user_id="caller", teams=[], object_permission_id="caller-permission", object_permission=current_permission,
)
database: Final = MagicMock()
database.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=team)
database.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=user)
database.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=current_permission)
cache: Final = UserApiKeyCache()
cache.set_cache("team_id:caller", team.model_copy(update={"object_permission": cached_permission}))
cache.set_cache("caller", user.model_copy(update={"object_permission": cached_permission}))
cache.set_cache(object_permission_cache_key("caller-permission"), cached_permission)
monkeypatch.setattr(proxy_server, "prisma_client", database)
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
auth: Final = actor(("read", "write"))
auth.requires_fresh_policy = fresh
auth.agent_caller = AgentCaller(team_id="caller") if caller_kind == "team" else AgentCaller(user_id="caller")
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == ({"slack"} if fresh else {"slack", "linear"})
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == (["read"] if fresh else ["read", "write"])
@pytest.mark.asyncio
@pytest.mark.parametrize("fresh", [False, True])
async def test_caller_team_outage_cannot_remove_authoritative_server_ceiling(
monkeypatch: pytest.MonkeyPatch, fresh: bool,
) -> None:
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.types.agents import AgentCaller
database: Final = MagicMock()
database.writer_db.litellm_teamtable.find_unique = AsyncMock(side_effect=RuntimeError("writer unavailable"))
database.db.litellm_teamtable.find_unique = AsyncMock(side_effect=RuntimeError("reader unavailable"))
monkeypatch.setattr(proxy_server, "prisma_client", database)
monkeypatch.setattr(proxy_server, "user_api_key_cache", UserApiKeyCache())
auth: Final = actor(("read",))
auth.agent_caller = AgentCaller(team_id="caller")
auth.requires_fresh_policy = fresh
if fresh:
with pytest.raises(HTTPException) as failure:
await MCPRequestHandler.get_allowed_mcp_servers(auth)
assert failure.value.status_code == 503
else:
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"slack", "linear"}

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