mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge remote-tracking branch 'origin/main' into litellm_propagate_4xx_missing_params
This commit is contained in:
commit
484e96f534
285 changed files with 27304 additions and 921 deletions
|
|
@ -99,7 +99,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 44358
|
||||
"limit": 44802
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 109
|
||||
|
|
|
|||
|
|
@ -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==",
|
||||
|
|
|
|||
|
|
@ -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 $$;
|
||||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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==",
|
||||
|
|
|
|||
19
litellm-rust/Cargo.lock
generated
19
litellm-rust/Cargo.lock
generated
|
|
@ -4075,6 +4075,7 @@ dependencies = [
|
|||
"litellm-secrets-aws",
|
||||
"litellm-secrets-types",
|
||||
"litellm-token-counter",
|
||||
"litellm-traces",
|
||||
"litellm-tracing",
|
||||
"pyo3",
|
||||
"pyo3-async-runtimes",
|
||||
|
|
@ -4351,6 +4352,21 @@ dependencies = [
|
|||
"tiktoken-rs",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-traces"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"litellm-http",
|
||||
"rstest",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"testcontainers-modules",
|
||||
"thiserror 2.0.19",
|
||||
"time",
|
||||
"tokio",
|
||||
"url",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-tracing"
|
||||
version = "0.1.0"
|
||||
|
|
@ -5704,6 +5720,7 @@ checksum = "16a1cfa75cc186dd73d5818e510e042e40927bccc9c236b061cea97e1eb08029"
|
|||
dependencies = [
|
||||
"base64 0.23.1",
|
||||
"bytes",
|
||||
"encoding_rs",
|
||||
"futures-core",
|
||||
"futures-util",
|
||||
"h2 0.4.15",
|
||||
|
|
@ -5715,6 +5732,7 @@ dependencies = [
|
|||
"hyper-util",
|
||||
"js-sys",
|
||||
"log",
|
||||
"mime",
|
||||
"percent-encoding",
|
||||
"pin-project-lite",
|
||||
"quinn",
|
||||
|
|
@ -6945,6 +6963,7 @@ dependencies = [
|
|||
"memchr",
|
||||
"parse-display",
|
||||
"pin-project-lite",
|
||||
"reqwest 0.13.5",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"serde_with",
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ repository = "https://github.com/BerriAI/litellm"
|
|||
litellm-config = { path = "crates/config" }
|
||||
litellm-router = { path = "crates/router" }
|
||||
litellm-tracing = { path = "crates/tracing" }
|
||||
litellm-traces = { path = "crates/traces" }
|
||||
litellm-core = { path = "crates/core" }
|
||||
litellm-gateway-mcp = { path = "crates/gateway-mcp" }
|
||||
litellm-gateway = { path = "crates/gateway" }
|
||||
|
|
|
|||
|
|
@ -203,6 +203,85 @@ fn threshold_tiers_and_boundaries() {
|
|||
assert_eq!(calculate(&specification, &flex).unwrap().input(), 600.0);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::ultrafast_above_threshold(ServiceTier::Ultrafast, 300_000, 9_301_000.0, 37_000.0)]
|
||||
#[case::ultrafast_at_threshold(ServiceTier::Ultrafast, 272_000, 544_500.0, 5_000.0)]
|
||||
#[case::standard_above_threshold(ServiceTier::Standard, 300_000, 3_300_600.0, 13_000.0)]
|
||||
#[case::priority_above_threshold(ServiceTier::Priority, 300_000, 5_701_000.0, 23_000.0)]
|
||||
fn tiered_long_context_rates_are_selected_by_service_tier(
|
||||
#[case] service_tier: ServiceTier,
|
||||
#[case] prompt_tokens: u64,
|
||||
#[case] expected_input: f64,
|
||||
#[case] expected_output: f64,
|
||||
) {
|
||||
let standard = Rates {
|
||||
cache_read: Rate::Value(3.0),
|
||||
..rates(Rate::Value(1.0), Rate::Value(2.0))
|
||||
};
|
||||
let tiers = [
|
||||
TierRates {
|
||||
tier: ServiceTier::Priority,
|
||||
rates: Rates {
|
||||
cache_read: Rate::Value(5.0),
|
||||
..rates(Rate::Value(3.0), Rate::Value(4.0))
|
||||
},
|
||||
},
|
||||
TierRates {
|
||||
tier: ServiceTier::Ultrafast,
|
||||
rates: Rates {
|
||||
cache_read: Rate::Value(7.0),
|
||||
..rates(Rate::Value(2.0), Rate::Value(5.0))
|
||||
},
|
||||
},
|
||||
];
|
||||
let threshold_tiers = [
|
||||
TierRates {
|
||||
tier: ServiceTier::Priority,
|
||||
rates: Rates {
|
||||
cache_read: Rate::Value(29.0),
|
||||
..rates(Rate::Value(19.0), Rate::Value(23.0))
|
||||
},
|
||||
},
|
||||
TierRates {
|
||||
tier: ServiceTier::Ultrafast,
|
||||
rates: Rates {
|
||||
cache_read: Rate::Value(41.0),
|
||||
..rates(Rate::Value(31.0), Rate::Value(37.0))
|
||||
},
|
||||
},
|
||||
];
|
||||
let thresholds = [ThresholdRates {
|
||||
above_prompt_tokens: 272_000,
|
||||
standard: Rates {
|
||||
cache_read: Rate::Value(17.0),
|
||||
..rates(Rate::Value(11.0), Rate::Value(13.0))
|
||||
},
|
||||
tiers: &threshold_tiers,
|
||||
}];
|
||||
let pricing = Pricing {
|
||||
standard,
|
||||
tiers: &tiers,
|
||||
thresholds: &thresholds,
|
||||
off_peak: None,
|
||||
};
|
||||
let base = request();
|
||||
let long_context_request = Request {
|
||||
usage: Usage {
|
||||
prompt_tokens,
|
||||
completion_tokens: 1_000,
|
||||
cache_read_tokens: 100,
|
||||
cache_write_tokens: 0,
|
||||
..base.usage
|
||||
},
|
||||
service_tier,
|
||||
..base
|
||||
};
|
||||
let cost = calculate(&pricing, &long_context_request).unwrap();
|
||||
|
||||
assert_eq!(cost.input(), expected_input);
|
||||
assert_eq!(cost.output(), expected_output);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn compile_rejects_ambiguous_rates() {
|
||||
let duplicate = ThresholdRates {
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ tiktoken = ["litellm-token-counter/tiktoken"]
|
|||
[dependencies]
|
||||
fancy-regex.workspace = true
|
||||
litellm-tracing.workspace = true
|
||||
litellm-traces.workspace = true
|
||||
litellm-host.workspace = true
|
||||
bytes.workspace = true
|
||||
futures-util.workspace = true
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
use crate::cache::cache_error;
|
||||
use crate::logger::run_sync_value;
|
||||
use crate::execution::run_sync_value;
|
||||
use litellm_cache_gcs::{DEFAULT_ENDPOINT, GcsConfig};
|
||||
use litellm_cache_redis_semantic::RedisSemanticConfig;
|
||||
use litellm_host_python::release_gil;
|
||||
|
|
|
|||
|
|
@ -470,7 +470,7 @@ impl NativeResponseCache {
|
|||
match self {
|
||||
Self::Exact(_) | Self::QdrantSemantic(_) => {
|
||||
let service = self.clone();
|
||||
crate::logger::run_async(
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move {
|
||||
service
|
||||
|
|
@ -495,7 +495,7 @@ impl NativeResponseCache {
|
|||
match self {
|
||||
Self::Exact(_) | Self::QdrantSemantic(_) => {
|
||||
let service = self.clone();
|
||||
crate::logger::run_async(
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move { service.async_lookup(&request, now()).await },
|
||||
cache_error,
|
||||
|
|
@ -550,7 +550,7 @@ impl NativeResponseCache {
|
|||
match self {
|
||||
Self::Exact(_) | Self::QdrantSemantic(_) => {
|
||||
let service = self.clone();
|
||||
crate::logger::run_async(
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move { service.async_store(&request, response, now()).await },
|
||||
cache_error,
|
||||
|
|
@ -619,7 +619,7 @@ impl NativeResponseCache {
|
|||
match self {
|
||||
Self::Exact(_) | Self::QdrantSemantic(_) => {
|
||||
let service = self.clone();
|
||||
crate::logger::run_async(
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move { service.async_store_batch(entries, now()).await },
|
||||
cache_error,
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
use crate::cache::cache_error;
|
||||
use crate::logger::run_async;
|
||||
use crate::execution::run_async;
|
||||
use std::{collections::VecDeque, time::Duration};
|
||||
|
||||
use litellm_cache::Error;
|
||||
|
|
|
|||
|
|
@ -144,7 +144,7 @@ impl NativeCacheHandle {
|
|||
self.check_process()?;
|
||||
let request = request(key, None)?;
|
||||
let backend = self.backend.clone();
|
||||
crate::logger::run_async(
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move { backend.async_lookup(&request, super::request::now()).await },
|
||||
cache_error,
|
||||
|
|
@ -163,7 +163,7 @@ impl NativeCacheHandle {
|
|||
let request = request(key, ttl)?;
|
||||
let value: Value = from_py(value)?;
|
||||
let backend = self.backend.clone();
|
||||
crate::logger::run_async(
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move {
|
||||
backend
|
||||
|
|
@ -188,7 +188,7 @@ impl NativeCacheHandle {
|
|||
.map(|(key, value)| Ok((request(key, ttl)?, value)))
|
||||
.collect::<PyResult<Vec<_>>>()?;
|
||||
let backend = self.backend.clone();
|
||||
crate::logger::run_async(
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move {
|
||||
backend
|
||||
|
|
@ -202,19 +202,19 @@ impl NativeCacheHandle {
|
|||
fn flush(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
|
||||
self.check_process()?;
|
||||
let backend = self.backend.clone();
|
||||
crate::logger::run_sync(py, async move { backend.async_flush().await }, cache_error)
|
||||
crate::execution::run_sync(py, async move { backend.async_flush().await }, cache_error)
|
||||
}
|
||||
|
||||
fn async_flush<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
|
||||
self.check_process()?;
|
||||
let backend = self.backend.clone();
|
||||
crate::logger::run_async(py, async move { backend.async_flush().await }, cache_error)
|
||||
crate::execution::run_async(py, async move { backend.async_flush().await }, cache_error)
|
||||
}
|
||||
|
||||
fn ping<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
|
||||
self.check_process()?;
|
||||
let storage = self.storage.clone();
|
||||
crate::logger::run_async(
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move {
|
||||
match storage {
|
||||
|
|
@ -229,7 +229,7 @@ impl NativeCacheHandle {
|
|||
fn disconnect<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
|
||||
self.check_process()?;
|
||||
let storage = self.storage.clone();
|
||||
crate::logger::run_async(
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move {
|
||||
match storage {
|
||||
|
|
@ -244,7 +244,7 @@ impl NativeCacheHandle {
|
|||
fn delete<'py>(&self, py: Python<'py>, keys: Vec<String>) -> PyResult<Bound<'py, PyAny>> {
|
||||
self.check_process()?;
|
||||
let storage = self.storage.clone();
|
||||
crate::logger::run_async(
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move {
|
||||
for key in keys {
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
use crate::logger::run_async;
|
||||
use crate::execution::run_async;
|
||||
use litellm_cache_response::PartialHits;
|
||||
use litellm_host_python::{ExecutionStep, from_py, release_gil, to_py};
|
||||
use pyo3::{
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ where
|
|||
E: Send + 'static,
|
||||
F: Future<Output = Result<T, E>> + Send + 'static,
|
||||
{
|
||||
litellm_host_python::run_sync(py, super::capture(py).instrument(future), map_error)
|
||||
litellm_host_python::run_sync(py, crate::logger::capture(py).instrument(future), map_error)
|
||||
}
|
||||
|
||||
pub(crate) fn run_async<T, E, F>(
|
||||
|
|
@ -26,7 +26,7 @@ where
|
|||
E: Send + 'static,
|
||||
F: Future<Output = Result<T, E>> + Send + 'static,
|
||||
{
|
||||
litellm_host_python::run_async(py, super::capture(py).instrument(future), map_error)
|
||||
litellm_host_python::run_async(py, crate::logger::capture(py).instrument(future), map_error)
|
||||
}
|
||||
|
||||
pub(crate) fn run_sync_value<T, F>(py: Python<'_>, future: F) -> PyResult<T>
|
||||
|
|
@ -34,7 +34,7 @@ where
|
|||
T: Send + 'static,
|
||||
F: Future<Output = PyResult<T>> + Send + 'static,
|
||||
{
|
||||
litellm_host_python::run_sync_value(py, super::capture(py).instrument(future))
|
||||
litellm_host_python::run_sync_value(py, crate::logger::capture(py).instrument(future))
|
||||
}
|
||||
|
||||
pub(crate) fn run_async_value<T, F>(py: Python<'_>, future: F) -> PyResult<Bound<'_, PyAny>>
|
||||
|
|
@ -42,5 +42,5 @@ where
|
|||
T: for<'py> IntoPyObject<'py> + Send + 'static,
|
||||
F: Future<Output = PyResult<T>> + Send + 'static,
|
||||
{
|
||||
litellm_host_python::run_async_value(py, super::capture(py).instrument(future))
|
||||
litellm_host_python::run_async_value(py, crate::logger::capture(py).instrument(future))
|
||||
}
|
||||
|
|
@ -4,6 +4,7 @@ mod coercion;
|
|||
mod credentials;
|
||||
mod diagnostics;
|
||||
mod errors;
|
||||
mod execution;
|
||||
mod http;
|
||||
mod lifecycle;
|
||||
mod logger;
|
||||
|
|
@ -42,6 +43,8 @@ mod _native {
|
|||
use crate::routes::responses::{ResponsesWebSocketConnection, aresponses, responses};
|
||||
#[pymodule_export]
|
||||
use crate::routes::token_counter::TokenCounter;
|
||||
#[pymodule_export]
|
||||
use crate::routes::traces::{trace_encode_rows, trace_ensure_schema, trace_query};
|
||||
#[cfg(feature = "huggingface")]
|
||||
#[pymodule_export]
|
||||
use crate::tokenizer::HuggingFaceEncoding;
|
||||
|
|
@ -106,6 +109,9 @@ mod tests {
|
|||
"aresponses",
|
||||
"ResponsesWebSocketConnection",
|
||||
"NativeDiagnosticProcessor",
|
||||
"trace_encode_rows",
|
||||
"trace_ensure_schema",
|
||||
"trace_query",
|
||||
"TokenCounter",
|
||||
"Tokenizer",
|
||||
"gil_stats",
|
||||
|
|
|
|||
|
|
@ -1,7 +1,5 @@
|
|||
mod execution;
|
||||
mod machine;
|
||||
|
||||
pub(crate) use execution::{run_async, run_async_value, run_sync, run_sync_value};
|
||||
pub(crate) use machine::LoggedMachine;
|
||||
|
||||
use litellm_host_python::Pythonized;
|
||||
|
|
|
|||
|
|
@ -76,7 +76,7 @@ async fn traced_operation(_secret: &str) -> PyResult<()> {
|
|||
|
||||
#[pyfunction]
|
||||
fn span_warning(py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
|
||||
super::run_async_value(py, traced_operation("private-key-sentinel"))
|
||||
crate::execution::run_async_value(py, traced_operation("private-key-sentinel"))
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
|
|
@ -93,7 +93,7 @@ fn levels(py: Python<'_>) {
|
|||
|
||||
#[pyfunction]
|
||||
fn asynchronous_warning(py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
|
||||
super::run_async_value(py, async {
|
||||
crate::execution::run_async_value(py, async {
|
||||
tokio::task::yield_now().await;
|
||||
litellm_tracing::warn!("async warning");
|
||||
Ok(())
|
||||
|
|
@ -102,7 +102,7 @@ fn asynchronous_warning(py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
|
|||
|
||||
#[pyfunction]
|
||||
fn synchronous_warning(py: Python<'_>) -> PyResult<()> {
|
||||
super::run_sync_value(py, async {
|
||||
crate::execution::run_sync_value(py, async {
|
||||
tokio::task::yield_now().await;
|
||||
litellm_tracing::warn!("sync warning");
|
||||
Ok(())
|
||||
|
|
@ -111,7 +111,7 @@ fn synchronous_warning(py: Python<'_>) -> PyResult<()> {
|
|||
|
||||
#[pyfunction]
|
||||
fn synchronous_failure(py: Python<'_>) -> PyResult<()> {
|
||||
super::run_sync_value(py, async {
|
||||
crate::execution::run_sync_value(py, async {
|
||||
litellm_tracing::warn!("failure diagnostic");
|
||||
Err(pyo3::exceptions::PyValueError::new_err("request failed"))
|
||||
})
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
use crate::logger::{run_async, run_sync};
|
||||
use crate::execution::{run_async, run_sync};
|
||||
use litellm_core::audio_transcription::{
|
||||
AudioTranscriptionRoute, Error, types::AudioTranscriptionRequest,
|
||||
};
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ mod host;
|
|||
|
||||
use pyo3::types::{PyDict, PyTuple};
|
||||
|
||||
use crate::logger::{run_async, run_sync};
|
||||
use crate::execution::{run_async, run_sync};
|
||||
use litellm_core::chat_completions::{ChatCompletionsRoute, Error, types::ChatCompletionsRequest};
|
||||
use litellm_llms_types::formats::chat_completions::ChatCompletionsResponse;
|
||||
use pyo3::prelude::*;
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ pub(crate) mod messages;
|
|||
pub(crate) mod ocr;
|
||||
pub(crate) mod responses;
|
||||
pub(crate) mod token_counter;
|
||||
pub(crate) mod traces;
|
||||
|
||||
use litellm_callbacks_legacy_python::LoggingOperation;
|
||||
use litellm_callbacks_legacy_python::{LegacyLogging, PublicCall};
|
||||
|
|
|
|||
|
|
@ -142,7 +142,7 @@ impl ResponsesWebSocketConnection {
|
|||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let headers = marshal_headers(headers)?;
|
||||
let timeout = optional_timeout(timeout_seconds);
|
||||
crate::logger::run_async_value(py, async move {
|
||||
crate::execution::run_async_value(py, async move {
|
||||
let inner = RustResponsesWebSocketConnection::connect_url(&url, &headers, timeout)
|
||||
.await
|
||||
.map_err(route_error_to_pyerr)?;
|
||||
|
|
@ -152,21 +152,21 @@ impl ResponsesWebSocketConnection {
|
|||
|
||||
fn send_text<'py>(&self, py: Python<'py>, text: String) -> PyResult<Bound<'py, PyAny>> {
|
||||
let inner = self.inner.clone();
|
||||
crate::logger::run_async_value(py, async move {
|
||||
crate::execution::run_async_value(py, async move {
|
||||
inner.send_text(text).await.map_err(route_error_to_pyerr)
|
||||
})
|
||||
}
|
||||
|
||||
fn recv_text<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
|
||||
let inner = self.inner.clone();
|
||||
crate::logger::run_async_value(py, async move {
|
||||
crate::execution::run_async_value(py, async move {
|
||||
inner.recv_text().await.map_err(route_error_to_pyerr)
|
||||
})
|
||||
}
|
||||
|
||||
fn close<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
|
||||
let inner = self.inner.clone();
|
||||
crate::logger::run_async_value(py, async move {
|
||||
crate::execution::run_async_value(py, async move {
|
||||
inner.close().await.map_err(route_error_to_pyerr)
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
use crate::logger::run_async;
|
||||
use crate::execution::run_async;
|
||||
use std::sync::Arc;
|
||||
use std::{num::NonZero, thread::available_parallelism};
|
||||
|
||||
|
|
|
|||
83
litellm-rust/crates/python-bridge/src/routes/traces.rs
Normal file
83
litellm-rust/crates/python-bridge/src/routes/traces.rs
Normal file
|
|
@ -0,0 +1,83 @@
|
|||
use std::collections::BTreeMap;
|
||||
|
||||
use litellm_http::ClientVariant;
|
||||
use litellm_traces::{Connection, Error, Parameter};
|
||||
use pyo3::{
|
||||
exceptions::{PyRuntimeError, PyValueError},
|
||||
prelude::*,
|
||||
};
|
||||
|
||||
fn map_error(error: Error) -> PyErr {
|
||||
match error {
|
||||
Error::InvalidRow | Error::InvalidSchema | Error::EmptySql => {
|
||||
PyValueError::new_err(error.to_string())
|
||||
}
|
||||
Error::InvalidUrl
|
||||
| Error::QueryFailed(_)
|
||||
| Error::SchemaFailed(_)
|
||||
| Error::ResponseTooLarge
|
||||
| Error::InvalidResponse
|
||||
| Error::Transport => PyRuntimeError::new_err(error.to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
pub fn trace_ensure_schema<'py>(
|
||||
py: Python<'py>,
|
||||
url: &str,
|
||||
database: String,
|
||||
user: &str,
|
||||
password: &str,
|
||||
trace_retention_days: u32,
|
||||
spend_log_retention_days: u32,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let connection = Connection::writer(url, user, password).map_err(map_error)?;
|
||||
let client = crate::http::host_client(py, ClientVariant::NoRedirect)?;
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move {
|
||||
litellm_traces::ensure_schema(
|
||||
&client,
|
||||
&connection,
|
||||
&database,
|
||||
trace_retention_days,
|
||||
spend_log_retention_days,
|
||||
)
|
||||
.await
|
||||
},
|
||||
map_error,
|
||||
)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
pub fn trace_query<'py>(
|
||||
py: Python<'py>,
|
||||
url: &str,
|
||||
database: &str,
|
||||
user: &str,
|
||||
password: &str,
|
||||
sql: String,
|
||||
#[pyo3(from_py_with = litellm_host_python::from_py_argument)] parameters: BTreeMap<
|
||||
String,
|
||||
Parameter,
|
||||
>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let connection = Connection::configured(url, database, user, password).map_err(map_error)?;
|
||||
let client = crate::http::host_client(py, ClientVariant::NoRedirect)?;
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move { litellm_traces::execute_read(&client, &connection, &sql, ¶meters).await },
|
||||
map_error,
|
||||
)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
pub fn trace_encode_rows(
|
||||
py: Python<'_>,
|
||||
#[pyo3(from_py_with = litellm_host_python::from_py_argument)] rows: Vec<
|
||||
BTreeMap<String, serde_json::Value>,
|
||||
>,
|
||||
) -> PyResult<String> {
|
||||
py.detach(|| litellm_traces::encode_rows(rows))
|
||||
.map_err(map_error)
|
||||
}
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
6
litellm-rust/crates/traces/AGENTS.md
Normal file
6
litellm-rust/crates/traces/AGENTS.md
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
- Rust owns OTLP wire decoding, ClickHouse schema, row encoding, named reads, connection validation and transport
|
||||
- Keep this crate independent of Python; PyO3 conversion and public Python exceptions belong in `python-bridge`
|
||||
- Keep the SQL migrations here as the only ClickHouse schema definition
|
||||
- Use typed query parameters and a dedicated SELECT-only reader with server-side limits
|
||||
- Bound insert time and encoded bytes; make retry deduplication behavior explicit for supported ClickHouse versions
|
||||
- Test storage behavior through the crate's public API against ClickHouse
|
||||
20
litellm-rust/crates/traces/Cargo.toml
Normal file
20
litellm-rust/crates/traces/Cargo.toml
Normal file
|
|
@ -0,0 +1,20 @@
|
|||
[package]
|
||||
name = "litellm-traces"
|
||||
version = "0.1.0"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
time = { workspace = true, features = ["formatting"] }
|
||||
litellm-http.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
thiserror.workspace = true
|
||||
url.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
rstest.workspace = true
|
||||
testcontainers-modules = { version = "0.15.0", features = ["clickhouse"] }
|
||||
tokio.workspace = true
|
||||
32
litellm-rust/crates/traces/config/reader.xml
Normal file
32
litellm-rust/crates/traces/config/reader.xml
Normal file
|
|
@ -0,0 +1,32 @@
|
|||
<clickhouse>
|
||||
<profiles>
|
||||
<litellm_traces_reader>
|
||||
<readonly>1</readonly>
|
||||
<max_execution_time>10</max_execution_time>
|
||||
<max_result_rows>1000</max_result_rows>
|
||||
<max_result_bytes>4194304</max_result_bytes>
|
||||
<result_overflow_mode>throw</result_overflow_mode>
|
||||
<max_memory_usage>268435456</max_memory_usage>
|
||||
<constraints>
|
||||
<readonly><readonly/></readonly>
|
||||
<max_execution_time><readonly/></max_execution_time>
|
||||
<max_result_rows><readonly/></max_result_rows>
|
||||
<max_result_bytes><readonly/></max_result_bytes>
|
||||
<result_overflow_mode><readonly/></result_overflow_mode>
|
||||
<max_memory_usage><readonly/></max_memory_usage>
|
||||
</constraints>
|
||||
</litellm_traces_reader>
|
||||
</profiles>
|
||||
<users>
|
||||
<litellm_traces_reader>
|
||||
<password from_env="LITELLM_TRACES_READER_PASSWORD"/>
|
||||
<networks><ip>::/0</ip></networks>
|
||||
<profile>litellm_traces_reader</profile>
|
||||
<grants>
|
||||
<query>GRANT SELECT ON litellm.otel_traces</query>
|
||||
<query>GRANT SELECT ON litellm.agent_traces</query>
|
||||
<query>GRANT SELECT ON litellm.spend_logs</query>
|
||||
</grants>
|
||||
</litellm_traces_reader>
|
||||
</users>
|
||||
</clickhouse>
|
||||
47
litellm-rust/crates/traces/migrations/0001_otel_traces.sql
Normal file
47
litellm-rust/crates/traces/migrations/0001_otel_traces.sql
Normal file
|
|
@ -0,0 +1,47 @@
|
|||
CREATE TABLE IF NOT EXISTS {database}.otel_traces
|
||||
(
|
||||
Timestamp DateTime64(9) CODEC(Delta, ZSTD(1)),
|
||||
TraceId String CODEC(ZSTD(1)),
|
||||
SpanId String CODEC(ZSTD(1)),
|
||||
ParentSpanId String CODEC(ZSTD(1)),
|
||||
TraceState String CODEC(ZSTD(1)),
|
||||
SpanName LowCardinality(String) CODEC(ZSTD(1)),
|
||||
SpanKind LowCardinality(String) CODEC(ZSTD(1)),
|
||||
ServiceName LowCardinality(String) CODEC(ZSTD(1)),
|
||||
ResourceAttributes Map(LowCardinality(String), String) CODEC(ZSTD(1)),
|
||||
ScopeName String CODEC(ZSTD(1)),
|
||||
ScopeVersion String CODEC(ZSTD(1)),
|
||||
SpanAttributes Map(LowCardinality(String), String) CODEC(ZSTD(1)),
|
||||
Duration UInt64 CODEC(ZSTD(1)),
|
||||
StatusCode LowCardinality(String) CODEC(ZSTD(1)),
|
||||
StatusMessage String CODEC(ZSTD(1)),
|
||||
`Events.Timestamp` Array(DateTime64(9)) CODEC(ZSTD(1)),
|
||||
`Events.Name` Array(LowCardinality(String)) CODEC(ZSTD(1)),
|
||||
`Events.Attributes` Array(Map(LowCardinality(String), String)) CODEC(ZSTD(1)),
|
||||
`Links.TraceId` Array(String) CODEC(ZSTD(1)),
|
||||
`Links.SpanId` Array(String) CODEC(ZSTD(1)),
|
||||
`Links.TraceState` Array(String) CODEC(ZSTD(1)),
|
||||
`Links.Attributes` Array(Map(LowCardinality(String), String)) CODEC(ZSTD(1)),
|
||||
TeamId LowCardinality(String) DEFAULT ResourceAttributes['litellm.team_id'],
|
||||
ApiKeyHash String DEFAULT ResourceAttributes['litellm.api_key_hash'],
|
||||
ObservationType LowCardinality(String) DEFAULT multiIf(
|
||||
ParentSpanId = '', 'agent',
|
||||
SpanAttributes['gen_ai.operation.name'] = 'invoke_agent', 'agent',
|
||||
SpanAttributes['gen_ai.operation.name'] IN ('chat', 'text_completion', 'generate_content'), 'llm',
|
||||
SpanAttributes['gen_ai.operation.name'] = 'execute_tool', 'tool',
|
||||
'chain'),
|
||||
AgentName LowCardinality(String) DEFAULT SpanAttributes['gen_ai.agent.name'],
|
||||
LiteLLMRequestId String DEFAULT SpanAttributes['gen_ai.response.id'],
|
||||
Model LowCardinality(String) DEFAULT SpanAttributes['gen_ai.request.model'],
|
||||
InputTokens UInt32 DEFAULT toUInt32OrZero(SpanAttributes['gen_ai.usage.input_tokens']),
|
||||
OutputTokens UInt32 DEFAULT toUInt32OrZero(SpanAttributes['gen_ai.usage.output_tokens']),
|
||||
Input String CODEC(ZSTD(3)),
|
||||
Output String CODEC(ZSTD(3)),
|
||||
InputPreview String DEFAULT substring(Input, 1, 240),
|
||||
INDEX idx_trace_id TraceId TYPE bloom_filter(0.001) GRANULARITY 1,
|
||||
INDEX idx_req_id LiteLLMRequestId TYPE bloom_filter(0.01) GRANULARITY 1
|
||||
)
|
||||
ENGINE = MergeTree
|
||||
PARTITION BY toDate(Timestamp)
|
||||
ORDER BY (TeamId, ServiceName, toDateTime(Timestamp), TraceId)
|
||||
SETTINGS ttl_only_drop_parts = 1
|
||||
23
litellm-rust/crates/traces/migrations/0002_agent_traces.sql
Normal file
23
litellm-rust/crates/traces/migrations/0002_agent_traces.sql
Normal file
|
|
@ -0,0 +1,23 @@
|
|||
CREATE TABLE IF NOT EXISTS {database}.agent_traces
|
||||
(
|
||||
TeamId LowCardinality(String),
|
||||
TraceId String,
|
||||
StartTs SimpleAggregateFunction(min, DateTime64(9)),
|
||||
EndTs SimpleAggregateFunction(max, DateTime64(9)),
|
||||
ServiceName SimpleAggregateFunction(any, LowCardinality(String)),
|
||||
RootName SimpleAggregateFunction(anyLast, Nullable(String)),
|
||||
RootInput SimpleAggregateFunction(anyLast, Nullable(String)),
|
||||
RootStatus SimpleAggregateFunction(anyLast, Nullable(String)),
|
||||
SpanCount SimpleAggregateFunction(sum, UInt64),
|
||||
AgentCount SimpleAggregateFunction(sum, UInt64),
|
||||
LlmCount SimpleAggregateFunction(sum, UInt64),
|
||||
ToolCount SimpleAggregateFunction(sum, UInt64),
|
||||
ErrorCount SimpleAggregateFunction(sum, UInt64),
|
||||
InputTokens SimpleAggregateFunction(sum, UInt64),
|
||||
OutputTokens SimpleAggregateFunction(sum, UInt64),
|
||||
Models SimpleAggregateFunction(groupUniqArrayArray, Array(String)),
|
||||
AgentNames SimpleAggregateFunction(groupUniqArrayArray, Array(String)),
|
||||
RequestIds SimpleAggregateFunction(groupArrayArray, Array(String))
|
||||
)
|
||||
ENGINE = AggregatingMergeTree
|
||||
ORDER BY (TeamId, TraceId)
|
||||
|
|
@ -0,0 +1,21 @@
|
|||
CREATE MATERIALIZED VIEW IF NOT EXISTS {database}.agent_traces_mv TO {database}.agent_traces AS
|
||||
SELECT
|
||||
TeamId, TraceId,
|
||||
min(Timestamp) AS StartTs,
|
||||
max(Timestamp + toIntervalNanosecond(Duration)) AS EndTs,
|
||||
any(ServiceName) AS ServiceName,
|
||||
anyLastIf(toNullable(SpanName), ParentSpanId = '') AS RootName,
|
||||
anyLastIf(toNullable(InputPreview), ParentSpanId = '') AS RootInput,
|
||||
anyLastIf(toNullable(StatusCode), ParentSpanId = '') AS RootStatus,
|
||||
count() AS SpanCount,
|
||||
countIf(ObservationType = 'agent') AS AgentCount,
|
||||
countIf(ObservationType = 'llm') AS LlmCount,
|
||||
countIf(ObservationType = 'tool') AS ToolCount,
|
||||
countIf(StatusCode = 'STATUS_CODE_ERROR') AS ErrorCount,
|
||||
sum(InputTokens) AS InputTokens,
|
||||
sum(OutputTokens) AS OutputTokens,
|
||||
groupUniqArrayIf(toString(Model), Model != '') AS Models,
|
||||
groupUniqArrayIf(SpanName, ObservationType = 'agent') AS AgentNames,
|
||||
groupArrayIf(LiteLLMRequestId, LiteLLMRequestId != '') AS RequestIds
|
||||
FROM {database}.otel_traces
|
||||
GROUP BY TeamId, TraceId
|
||||
42
litellm-rust/crates/traces/migrations/0004_spend_logs.sql
Normal file
42
litellm-rust/crates/traces/migrations/0004_spend_logs.sql
Normal file
|
|
@ -0,0 +1,42 @@
|
|||
CREATE TABLE IF NOT EXISTS {database}.spend_logs
|
||||
(
|
||||
request_id String,
|
||||
response_id String,
|
||||
call_type LowCardinality(String),
|
||||
api_key String,
|
||||
key_alias String,
|
||||
team_id LowCardinality(String),
|
||||
team_alias String,
|
||||
organization_id String,
|
||||
user String,
|
||||
end_user String,
|
||||
model LowCardinality(String),
|
||||
model_group LowCardinality(String),
|
||||
model_id String,
|
||||
custom_llm_provider LowCardinality(String),
|
||||
api_base String,
|
||||
spend Float64,
|
||||
prompt_tokens UInt32,
|
||||
completion_tokens UInt32,
|
||||
total_tokens UInt32,
|
||||
cache_read_tokens UInt32,
|
||||
cache_write_tokens UInt32,
|
||||
start_time DateTime64(3),
|
||||
end_time DateTime64(3),
|
||||
completion_start_time Nullable(DateTime64(3)),
|
||||
status LowCardinality(String),
|
||||
error_str String,
|
||||
cache_hit Bool,
|
||||
session_id String,
|
||||
trace_id String,
|
||||
span_id String,
|
||||
request_tags Array(String),
|
||||
metadata String CODEC(ZSTD(3)),
|
||||
messages String CODEC(ZSTD(3)),
|
||||
response String CODEC(ZSTD(3)),
|
||||
INDEX idx_response_id response_id TYPE bloom_filter(0.001) GRANULARITY 1,
|
||||
INDEX idx_trace_id trace_id TYPE bloom_filter(0.001) GRANULARITY 1
|
||||
)
|
||||
ENGINE = ReplacingMergeTree(end_time)
|
||||
PARTITION BY toYYYYMM(start_time)
|
||||
ORDER BY (team_id, start_time, request_id)
|
||||
|
|
@ -0,0 +1 @@
|
|||
ALTER TABLE {database}.otel_traces MODIFY TTL toDateTime(Timestamp) + INTERVAL {trace_retention_days} DAY
|
||||
|
|
@ -0,0 +1 @@
|
|||
ALTER TABLE {database}.agent_traces MODIFY TTL toDateTime(StartTs) + INTERVAL {trace_retention_days} DAY
|
||||
|
|
@ -0,0 +1 @@
|
|||
ALTER TABLE {database}.spend_logs MODIFY TTL toDateTime(start_time) + INTERVAL {spend_log_retention_days} DAY
|
||||
21
litellm-rust/crates/traces/src/error.rs
Normal file
21
litellm-rust/crates/traces/src/error.rs
Normal file
|
|
@ -0,0 +1,21 @@
|
|||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum Error {
|
||||
#[error("invalid ClickHouse insert row")]
|
||||
InvalidRow,
|
||||
#[error("invalid ClickHouse HTTP URL")]
|
||||
InvalidUrl,
|
||||
#[error("database must be a nonempty SQL identifier and retention must be positive")]
|
||||
InvalidSchema,
|
||||
#[error("SQL query must not be empty")]
|
||||
EmptySql,
|
||||
#[error("ClickHouse query failed with HTTP status {0}")]
|
||||
QueryFailed(u16),
|
||||
#[error("ClickHouse schema setup failed with HTTP status {0}")]
|
||||
SchemaFailed(u16),
|
||||
#[error("ClickHouse query exceeded the response size limit")]
|
||||
ResponseTooLarge,
|
||||
#[error("ClickHouse returned an invalid or failed JSON query response")]
|
||||
InvalidResponse,
|
||||
#[error("ClickHouse query transport failed")]
|
||||
Transport,
|
||||
}
|
||||
37
litellm-rust/crates/traces/src/insert.rs
Normal file
37
litellm-rust/crates/traces/src/insert.rs
Normal file
|
|
@ -0,0 +1,37 @@
|
|||
use std::collections::BTreeMap;
|
||||
|
||||
use serde_json::Value;
|
||||
use time::{OffsetDateTime, format_description::well_known::Rfc3339};
|
||||
|
||||
use crate::Error;
|
||||
|
||||
pub fn encode_rows(rows: Vec<BTreeMap<String, Value>>) -> Result<String, Error> {
|
||||
rows.into_iter()
|
||||
.map(|row| {
|
||||
let encoded = row
|
||||
.into_iter()
|
||||
.map(|(name, value)| insert_value(&name, value).map(|value| (name, value)))
|
||||
.collect::<Result<BTreeMap<_, _>, _>>()?;
|
||||
serde_json::to_string(&encoded).map_err(|_| Error::InvalidRow)
|
||||
})
|
||||
.collect::<Result<Vec<_>, _>>()
|
||||
.map(|rows| rows.join("\n"))
|
||||
}
|
||||
|
||||
fn insert_value(name: &str, value: Value) -> Result<Value, Error> {
|
||||
let multiplier = match name {
|
||||
"Timestamp" => 1,
|
||||
"start_time" | "end_time" | "completion_start_time" => 1_000_000,
|
||||
_ => return Ok(value),
|
||||
};
|
||||
if name == "completion_start_time" && value.is_null() {
|
||||
return Ok(value);
|
||||
}
|
||||
let timestamp = value.as_i64().ok_or(Error::InvalidRow)?;
|
||||
let datetime = OffsetDateTime::from_unix_timestamp_nanos(i128::from(timestamp) * multiplier)
|
||||
.map_err(|_| Error::InvalidRow)?;
|
||||
datetime
|
||||
.format(&Rfc3339)
|
||||
.map(Value::String)
|
||||
.map_err(|_| Error::InvalidRow)
|
||||
}
|
||||
73
litellm-rust/crates/traces/src/lib.rs
Normal file
73
litellm-rust/crates/traces/src/lib.rs
Normal file
|
|
@ -0,0 +1,73 @@
|
|||
mod error;
|
||||
mod insert;
|
||||
mod schema;
|
||||
mod sql;
|
||||
|
||||
pub use error::Error;
|
||||
pub use insert::encode_rows;
|
||||
pub use schema::{ensure_schema, schema_statements};
|
||||
pub use sql::{Parameter, execute_read};
|
||||
use url::Url;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct Connection {
|
||||
url: Url,
|
||||
}
|
||||
|
||||
impl Connection {
|
||||
pub fn parse(value: &str) -> Result<Self, Error> {
|
||||
let url = Url::parse(value).map_err(|_| Error::InvalidUrl)?;
|
||||
if !matches!(url.scheme(), "http" | "https") || url.host().is_none() {
|
||||
return Err(Error::InvalidUrl);
|
||||
}
|
||||
Ok(Self { url })
|
||||
}
|
||||
|
||||
pub fn configured(
|
||||
url: &str,
|
||||
database: &str,
|
||||
user: &str,
|
||||
password: &str,
|
||||
) -> Result<Self, Error> {
|
||||
let mut connection = Self::parse(url)?;
|
||||
connection
|
||||
.url
|
||||
.set_username(user)
|
||||
.map_err(|_| Error::InvalidUrl)?;
|
||||
connection
|
||||
.url
|
||||
.set_password(Some(password))
|
||||
.map_err(|_| Error::InvalidUrl)?;
|
||||
let pairs: Vec<_> = connection
|
||||
.url
|
||||
.query_pairs()
|
||||
.filter(|(key, _)| !matches!(key.as_ref(), "database" | "user" | "password"))
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect();
|
||||
connection
|
||||
.url
|
||||
.query_pairs_mut()
|
||||
.clear()
|
||||
.extend_pairs(pairs)
|
||||
.append_pair("database", database);
|
||||
Ok(connection)
|
||||
}
|
||||
|
||||
pub fn writer(url: &str, user: &str, password: &str) -> Result<Self, Error> {
|
||||
let mut connection = Self::parse(url)?;
|
||||
connection
|
||||
.url
|
||||
.set_username(user)
|
||||
.map_err(|_| Error::InvalidUrl)?;
|
||||
connection
|
||||
.url
|
||||
.set_password(Some(password))
|
||||
.map_err(|_| Error::InvalidUrl)?;
|
||||
connection.url.set_query(None);
|
||||
Ok(connection)
|
||||
}
|
||||
|
||||
pub fn url(&self) -> &Url {
|
||||
&self.url
|
||||
}
|
||||
}
|
||||
66
litellm-rust/crates/traces/src/schema.rs
Normal file
66
litellm-rust/crates/traces/src/schema.rs
Normal file
|
|
@ -0,0 +1,66 @@
|
|||
use litellm_http::Client;
|
||||
use std::time::Duration;
|
||||
|
||||
use crate::Connection;
|
||||
use crate::Error;
|
||||
|
||||
const MIGRATIONS: [&str; 7] = [
|
||||
include_str!("../migrations/0001_otel_traces.sql"),
|
||||
include_str!("../migrations/0002_agent_traces.sql"),
|
||||
include_str!("../migrations/0003_agent_traces_mv.sql"),
|
||||
include_str!("../migrations/0004_spend_logs.sql"),
|
||||
include_str!("../migrations/0005_otel_traces_ttl.sql"),
|
||||
include_str!("../migrations/0006_agent_traces_ttl.sql"),
|
||||
include_str!("../migrations/0007_spend_logs_ttl.sql"),
|
||||
];
|
||||
|
||||
pub fn schema_statements(
|
||||
database: &str,
|
||||
trace_retention_days: u32,
|
||||
spend_log_retention_days: u32,
|
||||
) -> Result<Vec<String>, Error> {
|
||||
if database.is_empty()
|
||||
|| !database
|
||||
.bytes()
|
||||
.all(|c| c.is_ascii_alphanumeric() || c == b'_')
|
||||
|| trace_retention_days == 0
|
||||
|| spend_log_retention_days == 0
|
||||
{
|
||||
return Err(Error::InvalidSchema);
|
||||
}
|
||||
let database = format!("`{database}`");
|
||||
Ok(
|
||||
std::iter::once(format!("CREATE DATABASE IF NOT EXISTS {database}"))
|
||||
.chain(MIGRATIONS.iter().map(|sql| {
|
||||
sql.replace("{database}", &database)
|
||||
.replace("{trace_retention_days}", &trace_retention_days.to_string())
|
||||
.replace(
|
||||
"{spend_log_retention_days}",
|
||||
&spend_log_retention_days.to_string(),
|
||||
)
|
||||
}))
|
||||
.collect(),
|
||||
)
|
||||
}
|
||||
|
||||
pub async fn ensure_schema(
|
||||
client: &Client,
|
||||
connection: &Connection,
|
||||
database: &str,
|
||||
trace_retention_days: u32,
|
||||
spend_log_retention_days: u32,
|
||||
) -> Result<(), Error> {
|
||||
for statement in schema_statements(database, trace_retention_days, spend_log_retention_days)? {
|
||||
let response = client
|
||||
.post(connection.url().clone())
|
||||
.timeout(Duration::from_secs(10))
|
||||
.body(statement)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|_| Error::Transport)?;
|
||||
if !response.status().is_success() {
|
||||
return Err(Error::SchemaFailed(response.status().as_u16()));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
114
litellm-rust/crates/traces/src/sql.rs
Normal file
114
litellm-rust/crates/traces/src/sql.rs
Normal file
|
|
@ -0,0 +1,114 @@
|
|||
use std::{collections::BTreeMap, time::Duration};
|
||||
|
||||
use serde::Deserialize;
|
||||
|
||||
use litellm_http::Client;
|
||||
|
||||
use crate::{Connection, Error};
|
||||
|
||||
const MAX_RESPONSE_BYTES: usize = 4 * 1024 * 1024;
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum Parameter {
|
||||
Text(String),
|
||||
Integer(i64),
|
||||
Strings(Vec<String>),
|
||||
}
|
||||
|
||||
impl Parameter {
|
||||
fn encoded(&self) -> String {
|
||||
match self {
|
||||
Self::Text(value) => escaped(value),
|
||||
Self::Integer(value) => value.to_string(),
|
||||
Self::Strings(values) => format!(
|
||||
"[{}]",
|
||||
values
|
||||
.iter()
|
||||
.map(|value| format!("'{}'", escaped(value).replace('\'', "\\'")))
|
||||
.collect::<Vec<_>>()
|
||||
.join(",")
|
||||
),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn escaped(value: &str) -> String {
|
||||
value
|
||||
.replace('\\', "\\\\")
|
||||
.replace('\t', "\\t")
|
||||
.replace('\n', "\\n")
|
||||
.replace('\r', "\\r")
|
||||
.replace('\0', "\\0")
|
||||
}
|
||||
|
||||
pub async fn execute_read(
|
||||
client: &Client,
|
||||
connection: &Connection,
|
||||
sql: &str,
|
||||
parameters: &BTreeMap<String, Parameter>,
|
||||
) -> Result<String, Error> {
|
||||
if sql.trim().is_empty() {
|
||||
return Err(Error::EmptySql);
|
||||
}
|
||||
|
||||
let mut url = connection.url().clone();
|
||||
|
||||
let existing_pairs: Vec<(String, String)> = url
|
||||
.query_pairs()
|
||||
.filter(|(key, _)| {
|
||||
!key.starts_with("param_")
|
||||
&& !matches!(
|
||||
key.as_ref(),
|
||||
"query"
|
||||
| "readonly"
|
||||
| "default_format"
|
||||
| "max_result_rows"
|
||||
| "result_overflow_mode"
|
||||
| "max_execution_time"
|
||||
| "wait_end_of_query"
|
||||
)
|
||||
})
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect();
|
||||
url.query_pairs_mut()
|
||||
.clear()
|
||||
.extend_pairs(existing_pairs)
|
||||
.append_pair("readonly", "1")
|
||||
.append_pair("max_result_rows", "1000")
|
||||
.append_pair("result_overflow_mode", "throw")
|
||||
.append_pair("max_execution_time", "10")
|
||||
.append_pair("wait_end_of_query", "1")
|
||||
.append_pair("default_format", "JSON");
|
||||
|
||||
url.query_pairs_mut().extend_pairs(
|
||||
parameters
|
||||
.iter()
|
||||
.map(|(name, value)| (format!("param_{name}"), value.encoded())),
|
||||
);
|
||||
|
||||
let request = client
|
||||
.post(url)
|
||||
.timeout(Duration::from_secs(15))
|
||||
.body(sql.to_owned());
|
||||
let mut response = request.send().await.map_err(|_| Error::Transport)?;
|
||||
if !response.status().is_success() {
|
||||
return Err(Error::QueryFailed(response.status().as_u16()));
|
||||
}
|
||||
|
||||
let mut body = Vec::new();
|
||||
while let Some(chunk) = response.chunk().await.map_err(|_| Error::Transport)? {
|
||||
if body.len() + chunk.len() > MAX_RESPONSE_BYTES {
|
||||
return Err(Error::ResponseTooLarge);
|
||||
}
|
||||
body.extend_from_slice(&chunk);
|
||||
}
|
||||
|
||||
let json: serde_json::Value =
|
||||
serde_json::from_slice(&body).map_err(|_| Error::InvalidResponse)?;
|
||||
if json.get("exception").is_some() || !json.get("data").is_some_and(serde_json::Value::is_array)
|
||||
{
|
||||
return Err(Error::InvalidResponse);
|
||||
}
|
||||
String::from_utf8(body).map_err(|_| Error::InvalidResponse)
|
||||
}
|
||||
259
litellm-rust/crates/traces/tests/admin_sql.rs
Normal file
259
litellm-rust/crates/traces/tests/admin_sql.rs
Normal file
|
|
@ -0,0 +1,259 @@
|
|||
use litellm_http::Client;
|
||||
use litellm_traces::{Connection, Error, Parameter, execute_read};
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::Value;
|
||||
use std::collections::BTreeMap;
|
||||
use testcontainers_modules::{
|
||||
clickhouse::ClickHouse,
|
||||
testcontainers::{ContainerAsync, ImageExt, runners::AsyncRunner},
|
||||
};
|
||||
|
||||
const CLICKHOUSE_TAG: &str =
|
||||
"26.9.6.6@sha256:eb4870e7ca7ed70c259eebfcfbee6cf797017f6b5436c2926bbbfe3d4d28486e";
|
||||
|
||||
struct Database {
|
||||
_container: ContainerAsync<ClickHouse>,
|
||||
url: String,
|
||||
admin_url: String,
|
||||
client: Client,
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
async fn database() -> Result<Database, Box<dyn std::error::Error>> {
|
||||
let container = ClickHouse::default()
|
||||
.with_tag(CLICKHOUSE_TAG)
|
||||
.with_env_var("CLICKHOUSE_SKIP_USER_SETUP", "1")
|
||||
.with_env_var("LITELLM_TRACES_READER_PASSWORD", "test_password")
|
||||
.with_copy_to(
|
||||
"/etc/clickhouse-server/users.d/litellm-traces-reader.xml",
|
||||
include_bytes!("../config/reader.xml").to_vec(),
|
||||
)
|
||||
.start()
|
||||
.await?;
|
||||
let admin_url = format!(
|
||||
"http://{}:{}",
|
||||
container.get_host().await?,
|
||||
container.get_host_port_ipv4(8123).await?,
|
||||
);
|
||||
let client = Client::no_redirect_for_test();
|
||||
for sql in [
|
||||
"CREATE DATABASE litellm",
|
||||
"CREATE TABLE litellm.otel_traces (n UInt8) ENGINE = Memory",
|
||||
"INSERT INTO litellm.otel_traces VALUES (1)",
|
||||
"CREATE TABLE litellm.private_traces (n UInt8) ENGINE = Memory",
|
||||
] {
|
||||
client
|
||||
.post(&admin_url)
|
||||
.body(sql)
|
||||
.send()
|
||||
.await?
|
||||
.error_for_status()?;
|
||||
}
|
||||
let url = format!(
|
||||
"{}?database=litellm",
|
||||
admin_url.replacen("http://", "http://litellm_traces_reader:test_password@", 1)
|
||||
);
|
||||
Ok(Database {
|
||||
_container: container,
|
||||
url,
|
||||
admin_url,
|
||||
client,
|
||||
})
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn admin_sql_reads_rows_with_enforced_settings(
|
||||
#[future(awt)] database: Result<Database, Box<dyn std::error::Error>>,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let database = database?;
|
||||
let connection = Connection::parse(&format!(
|
||||
"{}&readonly=0&default_format=TabSeparated&query=SELECT+2",
|
||||
database.url,
|
||||
))?;
|
||||
|
||||
let result = read(
|
||||
&database.client,
|
||||
&connection,
|
||||
"SELECT n AS answer FROM otel_traces",
|
||||
)
|
||||
.await?;
|
||||
let json: Value = serde_json::from_str(&result)?;
|
||||
assert_eq!(json["data"][0]["answer"], 1);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::table("CREATE TABLE admin_sql_test (n UInt8) ENGINE = Memory")]
|
||||
#[case::insert("INSERT INTO otel_traces VALUES (2)")]
|
||||
#[case::drop("DROP TABLE otel_traces")]
|
||||
#[case::named_collection("CREATE NAMED COLLECTION admin_sql_test AS host = 'localhost'")]
|
||||
#[case::settings("SET readonly = 0")]
|
||||
#[case::inline_settings("SELECT n FROM otel_traces SETTINGS readonly = 0")]
|
||||
#[case::time_limit("SELECT n FROM otel_traces SETTINGS max_execution_time = 0")]
|
||||
#[case::row_limit("SELECT n FROM otel_traces SETTINGS max_result_rows = 0")]
|
||||
#[case::byte_limit("SELECT n FROM otel_traces SETTINGS max_result_bytes = 0")]
|
||||
#[case::memory_limit("SELECT n FROM otel_traces SETTINGS max_memory_usage = 0")]
|
||||
#[case::other_table("SELECT * FROM private_traces")]
|
||||
#[tokio::test]
|
||||
async fn reader_rejects_writes_and_privilege_escalation(
|
||||
#[future(awt)] database: Result<Database, Box<dyn std::error::Error>>,
|
||||
#[case] sql: &str,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let database = database?;
|
||||
let connection = Connection::parse(&format!("{}&readonly=0", database.url))?;
|
||||
|
||||
let result = read(&database.client, &connection, sql).await;
|
||||
|
||||
assert!(matches!(result, Err(Error::QueryFailed(_))), "{result:?}");
|
||||
let rows = read(&database.client, &connection, "SELECT n FROM otel_traces").await?;
|
||||
let json: Value = serde_json::from_str(&rows)?;
|
||||
assert_eq!(json["data"], serde_json::json!([{ "n": 1 }]));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn admin_sql_rejects_errors_after_output_starts(
|
||||
#[future(awt)] database: Result<Database, Box<dyn std::error::Error>>,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let database = database?;
|
||||
let connection = Connection::parse(&format!(
|
||||
"{}?max_block_size=1&buffer_size=1&http_write_exception_in_output_format=1\
|
||||
&send_progress_in_http_headers=1&http_headers_progress_interval_ms=0",
|
||||
database.admin_url,
|
||||
))?;
|
||||
|
||||
let result = read(
|
||||
&database.client,
|
||||
&connection,
|
||||
"SELECT sleepEachRow(0.2), throwIf(number = 2) FROM numbers(5)",
|
||||
)
|
||||
.await;
|
||||
|
||||
assert!(
|
||||
matches!(result, Err(Error::InvalidResponse)),
|
||||
"expected an error embedded in a successful HTTP response: {result:?}"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn admin_sql_enforces_result_row_limit(
|
||||
#[future(awt)] database: Result<Database, Box<dyn std::error::Error>>,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let database = database?;
|
||||
let connection = Connection::parse(&format!(
|
||||
"{}&max_result_rows=0&result_overflow_mode=throw&wait_end_of_query=1",
|
||||
database.url,
|
||||
))?;
|
||||
|
||||
let result = read(
|
||||
&database.client,
|
||||
&connection,
|
||||
"SELECT number FROM numbers(1001)",
|
||||
)
|
||||
.await;
|
||||
|
||||
assert!(matches!(result, Err(Error::QueryFailed(_))), "{result:?}");
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn admin_sql_enforces_response_byte_limit(
|
||||
#[future(awt)] database: Result<Database, Box<dyn std::error::Error>>,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let database = database?;
|
||||
let connection = Connection::parse(&database.admin_url)?;
|
||||
|
||||
let result = read(
|
||||
&database.client,
|
||||
&connection,
|
||||
"SELECT repeat('x', 512 * 1024) AS payload FROM numbers(9)",
|
||||
)
|
||||
.await;
|
||||
|
||||
assert!(matches!(result, Err(Error::ResponseTooLarge)), "{result:?}");
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::plain("test_password", "test_password")]
|
||||
#[case::encoded("p@ss/word%", "p%40ss%2Fword%25")]
|
||||
#[tokio::test]
|
||||
async fn admin_sql_authenticates_url_credentials(
|
||||
#[future(awt)] database: Result<Database, Box<dyn std::error::Error>>,
|
||||
#[case] password: &str,
|
||||
#[case] encoded_password: &str,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let database = database?;
|
||||
database
|
||||
.client
|
||||
.post(&database.admin_url)
|
||||
.body(format!(
|
||||
"CREATE USER sql_reader IDENTIFIED WITH plaintext_password BY '{password}'"
|
||||
))
|
||||
.send()
|
||||
.await?
|
||||
.error_for_status()?;
|
||||
let connection = Connection::parse(&database.admin_url.replacen(
|
||||
"http://",
|
||||
&format!("http://sql_reader:{encoded_password}@"),
|
||||
1,
|
||||
))?;
|
||||
|
||||
let result = read(
|
||||
&database.client,
|
||||
&connection,
|
||||
"SELECT currentUser() AS username",
|
||||
)
|
||||
.await?;
|
||||
let json: Value = serde_json::from_str(&result)?;
|
||||
|
||||
assert_eq!(json["data"][0]["username"], "sql_reader");
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn read(client: &Client, connection: &Connection, sql: &str) -> Result<String, Error> {
|
||||
execute_read(client, connection, sql, &BTreeMap::new()).await
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::sql("'; DROP TABLE otel_traces; --")]
|
||||
#[case::escapes("back\\slash\ttab\nline\0null")]
|
||||
#[tokio::test]
|
||||
async fn query_parameters_preserve_values_and_replace_url_parameters(
|
||||
#[case] value: &str,
|
||||
#[future(awt)] database: Result<Database, Box<dyn std::error::Error>>,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let database = database?;
|
||||
let connection = Connection::parse(&format!("{}¶m_value=wrong", database.url))?;
|
||||
let values = vec![
|
||||
"a'b".to_owned(),
|
||||
"back\\slash".to_owned(),
|
||||
"line\nbreak".to_owned(),
|
||||
"雪".to_owned(),
|
||||
];
|
||||
let parameters = BTreeMap::from([
|
||||
("value".to_owned(), Parameter::Text(value.into())),
|
||||
("teams".to_owned(), Parameter::Strings(values.clone())),
|
||||
("number".to_owned(), Parameter::Integer(-42)),
|
||||
]);
|
||||
let body = execute_read(&database.client, &connection,
|
||||
"SELECT {value:String} AS value, {teams:Array(String)} AS teams, toInt32({number:Int64}) AS number",
|
||||
¶meters).await?;
|
||||
let json: Value = serde_json::from_str(&body)?;
|
||||
assert_eq!(json["data"][0]["value"], value);
|
||||
assert_eq!(json["data"][0]["teams"], serde_json::json!(values));
|
||||
assert_eq!(json["data"][0]["number"], -42);
|
||||
assert!(
|
||||
read(&database.client, &connection, "SELECT n FROM otel_traces")
|
||||
.await
|
||||
.is_ok()
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
40
litellm-rust/crates/traces/tests/insert.rs
Normal file
40
litellm-rust/crates/traces/tests/insert.rs
Normal file
|
|
@ -0,0 +1,40 @@
|
|||
use std::collections::BTreeMap;
|
||||
|
||||
use litellm_traces::encode_rows;
|
||||
use rstest::rstest;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
#[rstest]
|
||||
#[case::span("Timestamp", json!(1_234_567_890), json!("1970-01-01T00:00:01.23456789Z"))]
|
||||
#[case::start("start_time", json!(1_234), json!("1970-01-01T00:00:01.234Z"))]
|
||||
#[case::end("end_time", json!(2_345), json!("1970-01-01T00:00:02.345Z"))]
|
||||
#[case::completion("completion_start_time", json!(1_345), json!("1970-01-01T00:00:01.345Z"))]
|
||||
#[case::absent_completion("completion_start_time", Value::Null, Value::Null)]
|
||||
#[case::before_epoch("Timestamp", json!(-1), json!("1969-12-31T23:59:59.999999999Z"))]
|
||||
fn insert_encoding_preserves_timestamp_precision_and_other_fields(
|
||||
#[case] field: &str,
|
||||
#[case] value: Value,
|
||||
#[case] expected: Value,
|
||||
) {
|
||||
let rows = vec![BTreeMap::from([
|
||||
(field.to_owned(), value),
|
||||
("SpanAttributes".into(), json!({"message": "a\nb\\c\"雪"})),
|
||||
("InputTokens".into(), json!(42)),
|
||||
])];
|
||||
let encoded = encode_rows(rows).expect("valid row");
|
||||
let actual: Value = serde_json::from_str(&encoded).expect("JSONEachRow record");
|
||||
assert_eq!(
|
||||
actual,
|
||||
json!({
|
||||
field: expected, "SpanAttributes": {"message": "a\nb\\c\"雪"}, "InputTokens": 42
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::fractional(json!(1.25))]
|
||||
#[case::out_of_range(json!(u64::MAX))]
|
||||
#[case::null(Value::Null)]
|
||||
fn insert_encoding_rejects_invalid_span_timestamps(#[case] timestamp: Value) {
|
||||
assert!(encode_rows(vec![BTreeMap::from([("Timestamp".into(), timestamp)])]).is_err());
|
||||
}
|
||||
340
litellm-rust/crates/traces/tests/migrations.rs
Normal file
340
litellm-rust/crates/traces/tests/migrations.rs
Normal file
|
|
@ -0,0 +1,340 @@
|
|||
use std::{collections::BTreeMap, time::Duration};
|
||||
|
||||
use litellm_http::Client;
|
||||
use litellm_traces::{
|
||||
Connection, Error, encode_rows, ensure_schema, execute_read, schema_statements,
|
||||
};
|
||||
use rstest::{fixture, rstest};
|
||||
use testcontainers_modules::{
|
||||
clickhouse::ClickHouse,
|
||||
testcontainers::{ContainerAsync, ImageExt, runners::AsyncRunner},
|
||||
};
|
||||
|
||||
const CLICKHOUSE_TAG: &str =
|
||||
"26.9.6.6@sha256:eb4870e7ca7ed70c259eebfcfbee6cf797017f6b5436c2926bbbfe3d4d28486e";
|
||||
|
||||
type TestResult<T = ()> = Result<T, Box<dyn std::error::Error>>;
|
||||
|
||||
struct ClickHouseDatabase {
|
||||
_container: ContainerAsync<ClickHouse>,
|
||||
url: String,
|
||||
client: Client,
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
async fn database() -> TestResult<ClickHouseDatabase> {
|
||||
let container = ClickHouse::default()
|
||||
.with_tag(CLICKHOUSE_TAG)
|
||||
.with_env_var("CLICKHOUSE_SKIP_USER_SETUP", "1")
|
||||
.start()
|
||||
.await?;
|
||||
let url = format!(
|
||||
"http://{}:{}",
|
||||
container.get_host().await?,
|
||||
container.get_host_port_ipv4(8123).await?
|
||||
);
|
||||
Ok(ClickHouseDatabase {
|
||||
_container: container,
|
||||
url,
|
||||
client: Client::no_redirect_for_test(),
|
||||
})
|
||||
}
|
||||
|
||||
async fn insert_rows(
|
||||
database: &ClickHouseDatabase,
|
||||
table: &str,
|
||||
rows: Vec<BTreeMap<String, serde_json::Value>>,
|
||||
) -> TestResult {
|
||||
database
|
||||
.client
|
||||
.post(&database.url)
|
||||
.query(&[
|
||||
(
|
||||
"query",
|
||||
format!("INSERT INTO trace_test.{table} FORMAT JSONEachRow"),
|
||||
),
|
||||
("date_time_input_format", "best_effort".into()),
|
||||
])
|
||||
.body(encode_rows(rows)?)
|
||||
.send()
|
||||
.await?
|
||||
.error_for_status()?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn execute_write(database: &ClickHouseDatabase, sql: &str) -> TestResult {
|
||||
database
|
||||
.client
|
||||
.post(&database.url)
|
||||
.body(sql.to_owned())
|
||||
.send()
|
||||
.await?
|
||||
.error_for_status()?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn read_json(database: &ClickHouseDatabase, sql: &str) -> TestResult<serde_json::Value> {
|
||||
let connection = Connection::configured(&database.url, "trace_test", "default", "")?;
|
||||
let body = execute_read(&database.client, &connection, sql, &BTreeMap::new()).await?;
|
||||
Ok(serde_json::from_str(&body)?)
|
||||
}
|
||||
|
||||
async fn table_rows(database: &ClickHouseDatabase, table: &str) -> TestResult<u64> {
|
||||
let response = read_json(
|
||||
database,
|
||||
&format!("SELECT count() AS rows FROM trace_test.{table}"),
|
||||
)
|
||||
.await?;
|
||||
Ok(response["data"][0]["rows"]
|
||||
.as_u64()
|
||||
.expect("ClickHouse returns row counts as unsigned integers"))
|
||||
}
|
||||
|
||||
async fn mutation_rows(database: &ClickHouseDatabase) -> TestResult<u64> {
|
||||
let response = read_json(
|
||||
database,
|
||||
"SELECT count() AS rows FROM system.mutations WHERE database = 'trace_test'",
|
||||
)
|
||||
.await?;
|
||||
Ok(response["data"][0]["rows"]
|
||||
.as_u64()
|
||||
.expect("ClickHouse returns mutation counts as unsigned integers"))
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn schema_supports_span_rollups_and_spend_joins(
|
||||
#[future(awt)] database: TestResult<ClickHouseDatabase>,
|
||||
) -> TestResult {
|
||||
let database = database?;
|
||||
let writer = Connection::writer(&database.url, "default", "")?;
|
||||
ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?;
|
||||
ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?;
|
||||
let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64;
|
||||
let span = serde_json::from_value(serde_json::json!({
|
||||
"Timestamp": timestamp, "TraceId": "trace-1", "SpanId": "span-1", "ParentSpanId": "",
|
||||
"ServiceName": "proxy", "SpanName": "request", "Input": "hello world",
|
||||
"ResourceAttributes": {"litellm.team_id": "team-1", "litellm.api_key_hash": "hash-1"},
|
||||
"SpanAttributes": {"gen_ai.response.id": "response-1", "gen_ai.usage.input_tokens": "12"}
|
||||
}))?;
|
||||
let spend = serde_json::from_value(serde_json::json!({
|
||||
"request_id": "request-1", "response_id": "response-1", "team_id": "team-1", "spend": 0.125,
|
||||
"start_time": timestamp / 1_000_000, "end_time": timestamp / 1_000_000 + 100,
|
||||
"completion_start_time": null
|
||||
}))?;
|
||||
insert_rows(&database, "otel_traces", vec![span]).await?;
|
||||
insert_rows(&database, "spend_logs", vec![spend]).await?;
|
||||
let body = read_json(
|
||||
&database,
|
||||
"SELECT o.TeamId, o.ApiKeyHash, o.ObservationType, o.InputPreview, s.spend, \
|
||||
toString(toUnixTimestamp64Nano(o.Timestamp)) AS timestamp_ns, \
|
||||
toString(toUnixTimestamp64Milli(s.start_time)) AS start_ms \
|
||||
FROM trace_test.otel_traces o JOIN trace_test.spend_logs s \
|
||||
ON o.LiteLLMRequestId = s.response_id AND o.TeamId = s.team_id",
|
||||
)
|
||||
.await?;
|
||||
assert_eq!(
|
||||
body["data"],
|
||||
serde_json::json!([{
|
||||
"TeamId": "team-1", "ApiKeyHash": "hash-1", "ObservationType": "agent",
|
||||
"InputPreview": "hello world", "spend": 0.125,
|
||||
"timestamp_ns": timestamp.to_string(), "start_ms": (timestamp / 1_000_000).to_string()
|
||||
}])
|
||||
);
|
||||
let body = read_json(
|
||||
&database,
|
||||
"SELECT toUInt32(sum(SpanCount)) AS spans, toUInt32(sum(InputTokens)) AS tokens \
|
||||
FROM trace_test.agent_traces WHERE TeamId = 'team-1' AND TraceId = 'trace-1'",
|
||||
)
|
||||
.await?;
|
||||
assert_eq!(
|
||||
body["data"],
|
||||
serde_json::json!([{"spans": 1, "tokens": 12}])
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn rollup_merges_spans_across_days_without_losing_root_fields(
|
||||
#[future(awt)] database: TestResult<ClickHouseDatabase>,
|
||||
) -> TestResult {
|
||||
let database = database?;
|
||||
let writer = Connection::writer(&database.url, "default", "")?;
|
||||
ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?;
|
||||
let day_start = time::OffsetDateTime::now_utc()
|
||||
.replace_time(time::Time::MIDNIGHT)
|
||||
.unix_timestamp_nanos() as i64;
|
||||
let root = serde_json::from_value(serde_json::json!({
|
||||
"Timestamp": day_start - 1_000_000_000, "TraceId": "cross-day", "SpanId": "span-root",
|
||||
"ParentSpanId": "", "ServiceName": "proxy", "SpanName": "root", "Input": "root input",
|
||||
"StatusCode": "STATUS_CODE_ERROR",
|
||||
"ResourceAttributes": {"litellm.team_id": "team-1"}
|
||||
}))?;
|
||||
insert_rows(&database, "otel_traces", vec![root]).await?;
|
||||
let child = serde_json::from_value(serde_json::json!({
|
||||
"Timestamp": day_start + 1_000_000_000, "TraceId": "cross-day", "SpanId": "span-child",
|
||||
"ParentSpanId": "span-root", "ServiceName": "proxy", "SpanName": "child",
|
||||
"StatusCode": "STATUS_CODE_UNSET",
|
||||
"ResourceAttributes": {"litellm.team_id": "team-1"}
|
||||
}))?;
|
||||
insert_rows(&database, "otel_traces", vec![child]).await?;
|
||||
execute_write(&database, "OPTIMIZE TABLE trace_test.agent_traces FINAL").await?;
|
||||
let response = read_json(
|
||||
&database,
|
||||
"SELECT count() AS rows, any(RootName) AS RootName, any(RootInput) AS RootInput, \
|
||||
any(RootStatus) AS RootStatus, sum(SpanCount) AS SpanCount \
|
||||
FROM trace_test.agent_traces",
|
||||
)
|
||||
.await?;
|
||||
assert_eq!(
|
||||
response["data"],
|
||||
serde_json::json!([{
|
||||
"rows": 1, "RootName": "root", "RootInput": "root input",
|
||||
"RootStatus": "STATUS_CODE_ERROR", "SpanCount": 2
|
||||
}])
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn spend_deduplication_preserves_subsecond_requests_and_retries(
|
||||
#[future(awt)] database: TestResult<ClickHouseDatabase>,
|
||||
) -> TestResult {
|
||||
let database = database?;
|
||||
let writer = Connection::writer(&database.url, "default", "")?;
|
||||
ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?;
|
||||
let now_ms = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64 / 1_000_000;
|
||||
let base_start_time = now_ms / 1000 * 1000;
|
||||
let first_start_time = base_start_time + 100;
|
||||
let second_start_time = base_start_time + 200;
|
||||
let first = serde_json::from_value(serde_json::json!({
|
||||
"request_id": "same-request", "team_id": "team-1", "spend": 1.0,
|
||||
"start_time": first_start_time, "end_time": first_start_time + 1000
|
||||
}))?;
|
||||
let second = serde_json::from_value(serde_json::json!({
|
||||
"request_id": "same-request", "team_id": "team-1", "spend": 2.0,
|
||||
"start_time": second_start_time, "end_time": second_start_time + 1200
|
||||
}))?;
|
||||
let retry = serde_json::from_value(serde_json::json!({
|
||||
"request_id": "same-request", "team_id": "team-1", "spend": 1.0,
|
||||
"start_time": first_start_time, "end_time": first_start_time + 2000
|
||||
}))?;
|
||||
insert_rows(&database, "spend_logs", vec![first]).await?;
|
||||
insert_rows(&database, "spend_logs", vec![second]).await?;
|
||||
insert_rows(&database, "spend_logs", vec![retry]).await?;
|
||||
execute_write(&database, "OPTIMIZE TABLE trace_test.spend_logs FINAL").await?;
|
||||
let rows = read_json(
|
||||
&database,
|
||||
"SELECT toString(toUnixTimestamp64Milli(start_time)) AS start_time, \
|
||||
toString(toUnixTimestamp64Milli(end_time)) AS end_time \
|
||||
FROM trace_test.spend_logs ORDER BY start_time",
|
||||
)
|
||||
.await?;
|
||||
assert_eq!(
|
||||
rows["data"],
|
||||
serde_json::json!([
|
||||
{
|
||||
"start_time": first_start_time.to_string(),
|
||||
"end_time": (first_start_time + 2000).to_string()
|
||||
},
|
||||
{
|
||||
"start_time": second_start_time.to_string(),
|
||||
"end_time": (second_start_time + 1200).to_string()
|
||||
}
|
||||
])
|
||||
);
|
||||
assert_eq!(table_rows(&database, "spend_logs").await?, 2);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn retention_changes_materialize_existing_rows_and_remain_idempotent(
|
||||
#[future(awt)] database: TestResult<ClickHouseDatabase>,
|
||||
) -> TestResult {
|
||||
let database = database?;
|
||||
let writer = Connection::writer(&database.url, "default", "")?;
|
||||
ensure_schema(&database.client, &writer, "trace_test", 30, 30).await?;
|
||||
let old_time = time::OffsetDateTime::now_utc() - time::Duration::days(20);
|
||||
let old_timestamp_ns = old_time.unix_timestamp_nanos() as i64;
|
||||
let old_timestamp_ms = old_timestamp_ns / 1_000_000;
|
||||
let span = serde_json::from_value(serde_json::json!({
|
||||
"Timestamp": old_timestamp_ns, "TraceId": "expired", "SpanId": "span-old",
|
||||
"ParentSpanId": "", "ServiceName": "proxy", "SpanName": "old-root", "Input": "old input",
|
||||
"ResourceAttributes": {"litellm.team_id": "team-1"}
|
||||
}))?;
|
||||
let spend = serde_json::from_value(serde_json::json!({
|
||||
"request_id": "old-request", "team_id": "team-1", "spend": 1.0,
|
||||
"start_time": old_timestamp_ms, "end_time": old_timestamp_ms + 1000
|
||||
}))?;
|
||||
insert_rows(&database, "otel_traces", vec![span]).await?;
|
||||
insert_rows(&database, "spend_logs", vec![spend]).await?;
|
||||
ensure_schema(&database.client, &writer, "trace_test", 14, 14).await?;
|
||||
let deadline = tokio::time::Instant::now() + Duration::from_secs(60);
|
||||
loop {
|
||||
let response = read_json(
|
||||
&database,
|
||||
"SELECT countIf(is_done = 0) AS pending \
|
||||
FROM system.mutations WHERE database = 'trace_test'",
|
||||
)
|
||||
.await?;
|
||||
let pending = response["data"][0]["pending"]
|
||||
.as_u64()
|
||||
.expect("ClickHouse returns pending mutation counts as unsigned integers");
|
||||
if pending == 0 {
|
||||
break;
|
||||
}
|
||||
assert!(
|
||||
tokio::time::Instant::now() < deadline,
|
||||
"ClickHouse TTL mutations did not finish before the deadline"
|
||||
);
|
||||
tokio::time::sleep(Duration::from_millis(100)).await;
|
||||
}
|
||||
execute_write(&database, "OPTIMIZE TABLE trace_test.otel_traces FINAL").await?;
|
||||
execute_write(&database, "OPTIMIZE TABLE trace_test.agent_traces FINAL").await?;
|
||||
execute_write(&database, "OPTIMIZE TABLE trace_test.spend_logs FINAL").await?;
|
||||
assert_eq!(table_rows(&database, "otel_traces").await?, 0);
|
||||
assert_eq!(table_rows(&database, "agent_traces").await?, 0);
|
||||
assert_eq!(table_rows(&database, "spend_logs").await?, 0);
|
||||
let mutation_count = mutation_rows(&database).await?;
|
||||
ensure_schema(&database.client, &writer, "trace_test", 14, 14).await?;
|
||||
assert_eq!(mutation_rows(&database).await?, mutation_count);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn schema_statement_timeout_maps_to_transport_error() -> TestResult {
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
|
||||
let address = listener.local_addr()?;
|
||||
let server = tokio::spawn(async move {
|
||||
let (_connection, _) = listener.accept().await.expect("accept schema request");
|
||||
std::future::pending::<()>().await;
|
||||
});
|
||||
let client = Client::no_redirect_for_test();
|
||||
let url = format!("http://{address}");
|
||||
let writer = Connection::writer(&url, "default", "")?;
|
||||
let result = tokio::time::timeout(
|
||||
Duration::from_secs(12),
|
||||
ensure_schema(&client, &writer, "trace_test", 7, 14),
|
||||
)
|
||||
.await;
|
||||
server.abort();
|
||||
assert!(matches!(result, Ok(Err(Error::Transport))), "{result:?}");
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::empty("", 7, 14)]
|
||||
#[case::sql("db; DROP DATABASE default", 7, 14)]
|
||||
#[case::trace_retention("traces", 0, 14)]
|
||||
#[case::spend_retention("traces", 7, 0)]
|
||||
fn schema_rejects_invalid_configuration(
|
||||
#[case] database: &str,
|
||||
#[case] traces: u32,
|
||||
#[case] spend: u32,
|
||||
) {
|
||||
assert!(schema_statements(database, traces, spend).is_err());
|
||||
}
|
||||
11
litellm-rust/crates/traces/tests/queries.rs
Normal file
11
litellm-rust/crates/traces/tests/queries.rs
Normal file
|
|
@ -0,0 +1,11 @@
|
|||
use litellm_traces::Connection;
|
||||
use rstest::rstest;
|
||||
|
||||
#[rstest]
|
||||
#[case::http("http://localhost:8123", true)]
|
||||
#[case::https("https://localhost:8443", true)]
|
||||
#[case::tcp("tcp://localhost:9000", false)]
|
||||
#[case::missing_host("http://", false)]
|
||||
fn accepts_only_clickhouse_http_urls(#[case] value: &str, #[case] expected: bool) {
|
||||
assert_eq!(Connection::parse(value).is_ok(), expected);
|
||||
}
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -213,6 +213,15 @@ nothing here imports outside it:
|
|||
`config.yaml` — the latter reach the config through the logger's constructor
|
||||
kwargs. `baggage_team_metadata_keys` is empty by default, so none of a team's
|
||||
free-form metadata is promoted until each sub-key is explicitly allowlisted.
|
||||
`excluded_services` withholds datastore spans from key/team `callback_vars`
|
||||
destinations while the operator's own exporters keep them: set
|
||||
`LITELLM_OTEL_EXCLUDED_SERVICES` (comma-separated) or `excluded_services`
|
||||
(a YAML list) under `callback_settings.otel`, naming the datastore services
|
||||
to withhold (`redis`, `postgres`, `batch_write_to_db`, `redis_*`, or their
|
||||
`db.system.name` spellings `redis` / `postgresql`). Unknown names are logged
|
||||
as an error and ignored. A span is withheld when its `db.system.name` /
|
||||
`db.system` attribute is in the set, so request root, auth, guardrail and
|
||||
model spans can never be excluded.
|
||||
- [`baggage.py`](./model/baggage.py) — the single definition of which request-identity
|
||||
values are promoted into Baggage (so child spans inherit them) and under which
|
||||
attribute keys.
|
||||
|
|
|
|||
|
|
@ -30,7 +30,7 @@ from litellm.integrations.custom_logger import CustomLogger
|
|||
from litellm.integrations.otel.emitter import SpanEmitter, stamp_error
|
||||
from litellm.integrations.otel.mappers import resolve_mappers
|
||||
from litellm.integrations.otel.model.baggage import promoted_baggage
|
||||
from litellm.integrations.otel.model.config import OpenTelemetryV2Config
|
||||
from litellm.integrations.otel.model.config import OpenTelemetryV2Config, excluded_db_systems_from
|
||||
from litellm.integrations.otel.model.metadata import (
|
||||
LLMCallEvent,
|
||||
RequestIdentity,
|
||||
|
|
@ -898,12 +898,29 @@ def publish_global_otel_v2_provider(
|
|||
"""
|
||||
global _published_v2_provider
|
||||
logger: Final = select_global_otel_v2_logger(in_memory_loggers, registered=registered)
|
||||
attach_tenant_fan_out(logger.tracer_provider, *_v2_configs(in_memory_loggers, logger))
|
||||
attach_tenant_fan_out(
|
||||
logger.tracer_provider,
|
||||
*_v2_configs(in_memory_loggers, logger),
|
||||
excluded_db_systems=_excluded_db_systems(logger),
|
||||
)
|
||||
set_global_provider(logger.tracer_provider)
|
||||
_published_v2_provider = logger.tracer_provider # rebind-ok: startup records the one provider carrying the fan-out
|
||||
return logger
|
||||
|
||||
|
||||
def _excluded_db_systems(logger: "OpenTelemetryV2") -> frozenset[str]:
|
||||
"""The datastore services withheld from tenant destinations.
|
||||
|
||||
``callback_settings.otel.excluded_services`` wins over the env var whichever
|
||||
logger got published: with ``callbacks: [langfuse_otel, otel]`` the ``otel``
|
||||
callback folds into the preset, whose config is env-only.
|
||||
"""
|
||||
configured: Final = litellm.callback_settings.get("otel", {}).get("excluded_services")
|
||||
if configured is None:
|
||||
return logger.config.excluded_services
|
||||
return excluded_db_systems_from(configured)
|
||||
|
||||
|
||||
def _v2_configs(in_memory_loggers: Sequence[object], logger: "OpenTelemetryV2") -> tuple[OpenTelemetryV2Config, ...]:
|
||||
"""Every v2 logger's config, the published logger's first.
|
||||
|
||||
|
|
@ -963,7 +980,11 @@ def fan_out_provider() -> ApiTracerProvider:
|
|||
return published
|
||||
logger: Final = _registered_v2_logger()
|
||||
if logger is not None:
|
||||
attach_tenant_fan_out(logger.tracer_provider, logger.config)
|
||||
attach_tenant_fan_out(
|
||||
logger.tracer_provider,
|
||||
logger.config,
|
||||
excluded_db_systems=_excluded_db_systems(logger),
|
||||
)
|
||||
return logger.tracer_provider
|
||||
return get_tracer_provider()
|
||||
|
||||
|
|
|
|||
|
|
@ -4,14 +4,16 @@ from enum import Enum
|
|||
from functools import lru_cache
|
||||
from typing import Annotated, Any, Final
|
||||
|
||||
from pydantic import AliasChoices, BaseModel, Field, field_validator, model_validator
|
||||
from pydantic import AliasChoices, BaseModel, Field, TypeAdapter, ValidationError, field_validator, model_validator
|
||||
from pydantic_settings import BaseSettings, NoDecode, SettingsConfigDict
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.otel.model.baggage import (
|
||||
BAGGAGE_PROMOTED_KEYS,
|
||||
DEFAULT_BAGGAGE_METADATA_KEYS,
|
||||
DEFAULT_BAGGAGE_TEAM_METADATA_KEYS,
|
||||
)
|
||||
from litellm.integrations.otel.model.spans import POSTGRESQL, db_system
|
||||
from litellm.types.utils import OtelSpanScope
|
||||
|
||||
#: Master feature-flag env var. The logger is inert until this is truthy.
|
||||
|
|
@ -174,6 +176,19 @@ class OpenTelemetryV2Config(BaseSettings):
|
|||
"key/team destinations are not affected."
|
||||
),
|
||||
)
|
||||
excluded_services: Annotated[frozenset[str], NoDecode] = Field(
|
||||
default_factory=frozenset,
|
||||
validation_alias=AliasChoices("excluded_services", "LITELLM_OTEL_EXCLUDED_SERVICES"),
|
||||
description=(
|
||||
"Datastore services whose spans are withheld from key/team ``callback_vars`` "
|
||||
"OTel destinations (the operator's own exporters still receive them). Accepted "
|
||||
"values are the datastore ``ServiceTypes`` names (``redis``, ``postgres``, "
|
||||
"``batch_write_to_db``, ``redis_*``) or their ``db.system.name`` spellings "
|
||||
"(``redis``, ``postgresql``); stored normalized to ``db.system.name`` values. "
|
||||
"Configure via the ``LITELLM_OTEL_EXCLUDED_SERVICES`` env var (comma-separated) "
|
||||
"or ``callback_settings.otel.excluded_services`` in config.yaml (a YAML list)."
|
||||
),
|
||||
)
|
||||
|
||||
# ----- explicit multi-destination / vocabulary configuration ------------ #
|
||||
|
||||
|
|
@ -284,6 +299,11 @@ class OpenTelemetryV2Config(BaseSettings):
|
|||
return [item.strip() for item in value.split(",") if item.strip()]
|
||||
return value
|
||||
|
||||
@field_validator("excluded_services", mode="before")
|
||||
@classmethod
|
||||
def _read_excluded_services(cls, value: object) -> frozenset[str]:
|
||||
return excluded_service_names(value)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _normalize(self) -> "OpenTelemetryV2Config":
|
||||
# An endpoint with the default exporter kind implies OTLP/HTTP.
|
||||
|
|
@ -316,6 +336,7 @@ class OpenTelemetryV2Config(BaseSettings):
|
|||
if self.legacy_compat and "legacy" not in names:
|
||||
names.append("legacy")
|
||||
self.mapper_names = names
|
||||
self.excluded_services = _normalize_excluded_services(self.excluded_services)
|
||||
return self
|
||||
|
||||
@property
|
||||
|
|
@ -334,3 +355,55 @@ class OpenTelemetryV2Config(BaseSettings):
|
|||
@classmethod
|
||||
def from_env(cls) -> "OpenTelemetryV2Config":
|
||||
return cls()
|
||||
|
||||
|
||||
_EXCLUDED_SERVICES_INPUT: Final[TypeAdapter[str | tuple[object, ...]]] = TypeAdapter(str | tuple[object, ...])
|
||||
|
||||
|
||||
def excluded_db_systems_from(value: object) -> frozenset[str]:
|
||||
"""Normalize a raw ``excluded_services`` value without building a settings model that rereads the env"""
|
||||
return _normalize_excluded_services(excluded_service_names(value))
|
||||
|
||||
|
||||
def excluded_service_names(value: object) -> frozenset[str]:
|
||||
"""Read a YAML list or comma-separated string of service names, logging and dropping unusable input
|
||||
so a malformed value cannot stop the OTel logger from being built"""
|
||||
if value is None:
|
||||
return frozenset()
|
||||
try:
|
||||
parsed: Final = _EXCLUDED_SERVICES_INPUT.validate_python(value)
|
||||
except ValidationError:
|
||||
verbose_logger.error("excluded_services must be a list or comma-separated string; %r ignored", value)
|
||||
return frozenset()
|
||||
items: Final = tuple(parsed.split(",")) if isinstance(parsed, str) else parsed
|
||||
return frozenset(name for item in items if (name := _service_name(item)))
|
||||
|
||||
|
||||
def _service_name(item: object) -> str:
|
||||
if not isinstance(item, str):
|
||||
verbose_logger.error("excluded_services must be a list of service names; %r ignored", item)
|
||||
return ""
|
||||
return item.strip().lower()
|
||||
|
||||
|
||||
def _normalize_excluded_services(services: frozenset[str]) -> frozenset[str]:
|
||||
"""Fold each accepted spelling to its ``db.system.name`` value.
|
||||
|
||||
``postgres`` and ``postgresql`` name the same system, as do every
|
||||
``ServiceTypes`` member that ``db_system`` maps. Anything else means the
|
||||
operator pointed the setting at a span family it cannot cover; those names
|
||||
are logged and dropped so a typo cannot take the proxy down.
|
||||
"""
|
||||
resolved: Final = frozenset(
|
||||
system for service in services if (system := _db_system_for_excluded_service(service)) is not None
|
||||
)
|
||||
return resolved
|
||||
|
||||
|
||||
def _db_system_for_excluded_service(service: str) -> str | None:
|
||||
resolved: Final = db_system(service) if service != POSTGRESQL else POSTGRESQL
|
||||
if resolved is None:
|
||||
verbose_logger.error(
|
||||
"excluded_services: %r is not a datastore service; ignored. Allowed: postgres, redis", service
|
||||
)
|
||||
return resolved
|
||||
|
|
|
|||
|
|
@ -418,6 +418,13 @@ def _is_database_span(attributes: Mapping[str, AttributeValue]) -> bool:
|
|||
return any(key in attributes for key in _DB_SYSTEM_KEYS)
|
||||
|
||||
|
||||
def _is_excluded_database_span(attributes: Mapping[str, AttributeValue], excluded: frozenset[str]) -> bool:
|
||||
if not excluded:
|
||||
return False
|
||||
system: Final = attributes.get(DB.SYSTEM_NAME) or attributes.get(DB.SYSTEM_LEGACY)
|
||||
return isinstance(system, str) and system in excluded
|
||||
|
||||
|
||||
def _is_tenant_owned_span(attributes: Mapping[str, AttributeValue]) -> bool:
|
||||
return any(key in attributes for key in _TENANT_OWNED_KEYS)
|
||||
|
||||
|
|
@ -549,10 +556,12 @@ class TenantFanOutSpanProcessor(SpanProcessor):
|
|||
processor_factory: 'Callable[["OtelDestination"], SpanProcessor | None] | None' = None,
|
||||
shutdown_drain_seconds: float = _SHUTDOWN_DRAIN_SECONDS,
|
||||
operator_sinks: 'Mapping[_SinkKey, "OtelSpanScope"]' = MappingProxyType({}),
|
||||
excluded_db_systems: frozenset[str] = frozenset(),
|
||||
pending_drains: int = _MAX_PENDING_DRAINS,
|
||||
drain_pool: _DrainPool | None = None,
|
||||
) -> None:
|
||||
self._operator_sinks: Final = operator_sinks
|
||||
self._excluded_db_systems: Final = excluded_db_systems
|
||||
self._drain_seconds: Final = shutdown_drain_seconds
|
||||
self._lock: Final = threading.Condition()
|
||||
self._closed = False # guarded by ``_lock``: an unlocked read races the teardown it gates
|
||||
|
|
@ -567,9 +576,12 @@ class TenantFanOutSpanProcessor(SpanProcessor):
|
|||
|
||||
def on_end(self, span: ReadableSpan) -> None:
|
||||
suppressed: Final = suppressed_backends()
|
||||
attributes: Final = span.attributes or _NO_ATTRIBUTES
|
||||
for destination in request_destinations():
|
||||
if self._operator_already_writes(span, destination, suppressed) or not _in_scope(
|
||||
span, destination.span_scope
|
||||
if (
|
||||
self._operator_already_writes(span, destination, suppressed)
|
||||
or not _in_scope(span, destination.span_scope)
|
||||
or _is_excluded_database_span(attributes, self._excluded_db_systems)
|
||||
):
|
||||
continue
|
||||
processor = self._acquire(destination)
|
||||
|
|
@ -1155,7 +1167,9 @@ def build_tracer_provider(
|
|||
_FAN_OUT_ATTACH_LOCK: Final = threading.Lock()
|
||||
|
||||
|
||||
def attach_tenant_fan_out(provider: TracerProvider, *configs: OpenTelemetryV2Config) -> None:
|
||||
def attach_tenant_fan_out(
|
||||
provider: TracerProvider, *configs: OpenTelemetryV2Config, excluded_db_systems: frozenset[str] = frozenset()
|
||||
) -> None:
|
||||
"""Give ``provider`` the fan-out that delivers spans to key/team destinations.
|
||||
|
||||
Called on the one provider published as the OTel global, and idempotent so a
|
||||
|
|
@ -1164,12 +1178,18 @@ def attach_tenant_fan_out(provider: TracerProvider, *configs: OpenTelemetryV2Con
|
|||
so exactly one fan-out lands. ``configs`` name the operator's own exporters, one
|
||||
config per v2 logger since each keeps its own provider and still writes its
|
||||
account, so an additive destination pointing at any of them is delivered once
|
||||
rather than twice.
|
||||
rather than twice. ``excluded_db_systems`` only filters what the fan-out
|
||||
delivers, never the operator's own exporters.
|
||||
"""
|
||||
with _FAN_OUT_ATTACH_LOCK:
|
||||
if any(isinstance(processor, TenantFanOutSpanProcessor) for processor in _attached_processors(provider)):
|
||||
return
|
||||
provider.add_span_processor(TenantFanOutSpanProcessor(operator_sinks=operator_sink_scopes(*configs)))
|
||||
provider.add_span_processor(
|
||||
TenantFanOutSpanProcessor(
|
||||
operator_sinks=operator_sink_scopes(*configs),
|
||||
excluded_db_systems=excluded_db_systems,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def deliverable_destinations(
|
||||
|
|
|
|||
|
|
@ -2015,7 +2015,11 @@ def is_unsignable_thinking_block(block: object) -> bool:
|
|||
return not (isinstance(thinking_text, str) and len(thinking_text.strip()) > 0)
|
||||
|
||||
|
||||
def strip_encrypted_reasoning_from_messages(messages: object) -> None:
|
||||
def strip_encrypted_reasoning_from_messages(
|
||||
messages: object,
|
||||
*,
|
||||
should_strip: Callable[[Mapping[str, object]], bool] | None = None,
|
||||
) -> None:
|
||||
"""Drop the bridge-tagged reasoning blocks a routed deployment cannot decrypt from
|
||||
Anthropic-shaped history.
|
||||
|
||||
|
|
@ -2030,7 +2034,7 @@ def strip_encrypted_reasoning_from_messages(messages: object) -> None:
|
|||
if not isinstance(messages, list):
|
||||
return
|
||||
for content in anthropic_content_lists(cast(list[object], messages)): # cast-ok: untyped client json
|
||||
_strip_encrypted_reasoning_from_blocks(content)
|
||||
_strip_encrypted_reasoning_from_blocks(content, should_strip=should_strip)
|
||||
|
||||
|
||||
def anthropic_content_lists(messages: Sequence[object]) -> Iterator[object]:
|
||||
|
|
@ -2043,9 +2047,18 @@ def anthropic_content_lists(messages: Sequence[object]) -> Iterator[object]:
|
|||
)
|
||||
|
||||
|
||||
def _strip_encrypted_reasoning_from_blocks(content: object) -> None:
|
||||
def _strip_encrypted_reasoning_from_blocks(
|
||||
content: object,
|
||||
*,
|
||||
should_strip: Callable[[Mapping[str, object]], bool] | None = None,
|
||||
) -> None:
|
||||
blocks: Final = cast(list[object], content) # cast-ok: narrowed by the caller's isinstance
|
||||
kept: Final = tuple(block for block in blocks if not is_encrypted_reasoning_block(block))
|
||||
kept: Final = tuple(
|
||||
block
|
||||
for block in blocks
|
||||
if not is_encrypted_reasoning_block(block)
|
||||
or (should_strip is not None and not should_strip(cast(Mapping[str, object], block)))
|
||||
)
|
||||
blocks[:] = kept
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -53,6 +53,10 @@ from litellm.llms.base_llm.guardrail_translation.base_translation import (
|
|||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
blocked_responses_stream_usage,
|
||||
effective_skip_system_message_for_guardrail,
|
||||
merge_guardrailed_scoped_messages,
|
||||
role_out_of_guardrail_scope,
|
||||
scoped_structured_message_indices,
|
||||
stream_item_field,
|
||||
stream_item_fingerprint,
|
||||
stream_item_items,
|
||||
|
|
@ -376,6 +380,17 @@ class _RequestFields(NamedTuple):
|
|||
class _ExtractedInputs(NamedTuple):
|
||||
inputs: GenericGuardrailAPIInputs
|
||||
task_mappings: tuple[tuple[int, int | None], ...]
|
||||
instructions: str | None
|
||||
|
||||
|
||||
def scannable_instructions(data: Mapping[str, object], *, skip_system: bool = False) -> str | None:
|
||||
instructions: Final = data.get("instructions")
|
||||
return instructions if isinstance(instructions, str) and instructions and not skip_system else None
|
||||
|
||||
|
||||
def _input_item_role(item: object) -> str:
|
||||
role: Final = item.get("role") if isinstance(item, Mapping) else None
|
||||
return role.lower() if isinstance(role, str) else ""
|
||||
|
||||
|
||||
def _patched_request_fields(
|
||||
|
|
@ -494,7 +509,14 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
input_data: Final[str | ResponseInputParam | None] = data.get("input")
|
||||
if not isinstance(input_data, (str, list)):
|
||||
return data
|
||||
skip_system: Final = effective_skip_system_message_for_guardrail(guardrail_to_apply)
|
||||
structured_messages: Final = self.get_structured_messages(data)
|
||||
scoped_indices: Final = scoped_structured_message_indices(
|
||||
structured_messages or [], scan_only_tool_results=False, skip_system=skip_system, skip_tool=False
|
||||
)
|
||||
scoped_structured_messages: Final = (
|
||||
[structured_messages[index] for index in scoped_indices] if structured_messages else None
|
||||
)
|
||||
raw_tools: Final = data.get("tools")
|
||||
original_tools: Final[tuple[Mapping[str, object], ...]] = (
|
||||
tuple(raw_tools) if isinstance(raw_tools, list) else ()
|
||||
|
|
@ -502,11 +524,13 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
flattened_tool_groups: Final = tuple(
|
||||
form.chat_tools for form in LiteLLMCompletionResponsesConfig.responses_tools_to_chat_forms(original_tools)
|
||||
)
|
||||
extracted: Final = self._extract_guardrail_inputs(data, input_data, flattened_tool_groups)
|
||||
extracted: Final = self._extract_guardrail_inputs(
|
||||
data, input_data, flattened_tool_groups, skip_system=skip_system
|
||||
)
|
||||
if not extracted.inputs.get("texts"):
|
||||
return data
|
||||
if structured_messages:
|
||||
extracted.inputs["structured_messages"] = structured_messages
|
||||
if scoped_structured_messages:
|
||||
extracted.inputs["structured_messages"] = scoped_structured_messages
|
||||
guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail(
|
||||
inputs=extracted.inputs,
|
||||
request_data=data,
|
||||
|
|
@ -516,37 +540,63 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
self._apply_guardrailed_tools_to_data(
|
||||
data, original_tools, flattened_tool_groups, guardrailed_inputs.get("tools")
|
||||
)
|
||||
written_back: Final = self._written_back_request_fields(data, structured_messages, guardrailed_inputs)
|
||||
written_back: Final = self._written_back_request_fields(
|
||||
data,
|
||||
structured_messages or (),
|
||||
scoped_indices,
|
||||
scoped_structured_messages,
|
||||
guardrail_to_apply,
|
||||
guardrailed_inputs,
|
||||
)
|
||||
if written_back is not None:
|
||||
data["input"] = list(written_back.input) # mutable-ok: JSON body
|
||||
if written_back.instructions is None:
|
||||
data.pop("instructions", None)
|
||||
else:
|
||||
data["instructions"] = written_back.instructions # rebind-ok: data is an out-param
|
||||
elif isinstance(input_data, str):
|
||||
guardrailed_texts: Final = guardrailed_inputs.get("texts") or ()
|
||||
if len(guardrailed_texts) > 1:
|
||||
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
|
||||
data["input"] = guardrailed_texts[0] if guardrailed_texts else input_data # rebind-ok: data is an out-param
|
||||
else:
|
||||
rewritten_texts: Final = guardrailed_inputs.get("texts") or ()
|
||||
if len(rewritten_texts) != len(extracted.task_mappings):
|
||||
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
|
||||
await self._apply_guardrail_responses_to_input(
|
||||
messages=input_data,
|
||||
responses=rewritten_texts,
|
||||
task_mappings=extracted.task_mappings,
|
||||
)
|
||||
await self._apply_guardrailed_texts(data, input_data, extracted, guardrail_to_apply, guardrailed_inputs)
|
||||
verbose_proxy_logger.debug("OpenAI Responses API: Processed input messages: %s", data.get("input"))
|
||||
return data
|
||||
|
||||
async def _apply_guardrailed_texts(
|
||||
self,
|
||||
data: dict[str, object],
|
||||
input_data: "str | ResponseInputParam",
|
||||
extracted: _ExtractedInputs,
|
||||
guardrail_to_apply: "CustomGuardrail",
|
||||
guardrailed_inputs: GenericGuardrailAPIInputs,
|
||||
) -> None:
|
||||
returned_texts: Final = guardrailed_inputs.get("texts")
|
||||
if not returned_texts:
|
||||
return
|
||||
rewritten_texts: Final = tuple(returned_texts)
|
||||
offset: Final = 0 if extracted.instructions is None else 1
|
||||
input_texts: Final = rewritten_texts[offset:]
|
||||
expected: Final = 1 if isinstance(input_data, str) else len(extracted.task_mappings)
|
||||
if len(rewritten_texts) != offset + expected:
|
||||
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
|
||||
if offset:
|
||||
data["instructions"] = rewritten_texts[0] # rebind-ok: data is an out-param
|
||||
if isinstance(input_data, str):
|
||||
data["input"] = input_texts[0] # rebind-ok: data is an out-param
|
||||
return
|
||||
await self._apply_guardrail_responses_to_input(
|
||||
messages=input_data,
|
||||
responses=input_texts,
|
||||
task_mappings=extracted.task_mappings,
|
||||
)
|
||||
|
||||
def _extract_guardrail_inputs(
|
||||
self,
|
||||
data: Mapping[str, object],
|
||||
input_data: "str | ResponseInputParam",
|
||||
flattened_tool_groups: Sequence[Sequence[Mapping[str, object]]],
|
||||
*,
|
||||
skip_system: bool = False,
|
||||
) -> _ExtractedInputs:
|
||||
texts_to_check: Final[list[str]] = []
|
||||
instructions: Final = scannable_instructions(data, skip_system=skip_system)
|
||||
texts_to_check: Final[list[str]] = [] if instructions is None else [instructions]
|
||||
images_to_check: Final[list[str]] = []
|
||||
task_mappings: Final[list[tuple[int, int | None]]] = []
|
||||
tools_to_check: Final[list[ChatCompletionToolParam]] = list( # mutable-ok: guardrail inputs want a list
|
||||
|
|
@ -562,6 +612,10 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
texts_to_check.append(input_data)
|
||||
else:
|
||||
for msg_idx, message in enumerate(input_data):
|
||||
if role_out_of_guardrail_scope(
|
||||
_input_item_role(message), skip_system_message=skip_system, skip_tool_message=False
|
||||
):
|
||||
continue
|
||||
self._extract_input_text_and_images(
|
||||
message=message,
|
||||
msg_idx=msg_idx,
|
||||
|
|
@ -577,22 +631,32 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
model: Final = data.get("model")
|
||||
if isinstance(model, str):
|
||||
inputs["model"] = model
|
||||
return _ExtractedInputs(inputs=inputs, task_mappings=tuple(task_mappings))
|
||||
return _ExtractedInputs(inputs=inputs, task_mappings=tuple(task_mappings), instructions=instructions)
|
||||
|
||||
@staticmethod
|
||||
def _written_back_request_fields(
|
||||
data: Mapping[str, object],
|
||||
structured_messages: Sequence[AllMessageValues] | None,
|
||||
structured_messages: Sequence[AllMessageValues],
|
||||
scoped_indices: Sequence[int],
|
||||
scoped_structured_messages: Sequence[AllMessageValues] | None,
|
||||
guardrail_to_apply: "CustomGuardrail",
|
||||
guardrailed_inputs: GenericGuardrailAPIInputs,
|
||||
) -> _RequestFields | None:
|
||||
guardrailed: Final = guardrailed_inputs.get("structured_messages")
|
||||
if guardrailed is None or guardrailed is structured_messages:
|
||||
if guardrailed is None or guardrailed is scoped_structured_messages:
|
||||
return None
|
||||
covers_full_request: Final = len(scoped_indices) == len(structured_messages) or (
|
||||
guardrail_to_apply.structured_messages_cover_full_request() and len(guardrailed) == len(structured_messages)
|
||||
)
|
||||
merged: Final = (
|
||||
guardrailed
|
||||
if covers_full_request
|
||||
else merge_guardrailed_scoped_messages(
|
||||
full_messages=structured_messages, scoped_indices=scoped_indices, guardrailed_scoped=guardrailed
|
||||
)
|
||||
)
|
||||
return _patch_or_convert_request_fields(
|
||||
data.get("input"),
|
||||
data.get("instructions"),
|
||||
structured_messages or (),
|
||||
guardrailed,
|
||||
data.get("input"), data.get("instructions"), structured_messages, merged
|
||||
)
|
||||
|
||||
def extract_request_tool_names(self, data: dict) -> list[str]:
|
||||
|
|
|
|||
|
|
@ -3358,7 +3358,7 @@
|
|||
"supports_function_calling": true
|
||||
},
|
||||
"azure_ai/claude-haiku-4-5": {
|
||||
"deprecation_date": "2026-10-19",
|
||||
"deprecation_date": "2026-11-15",
|
||||
"cache_creation_input_token_cost": 1.25e-06,
|
||||
"cache_creation_input_token_cost_above_1hr": 2e-06,
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
|
|
@ -3378,10 +3378,11 @@
|
|||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"prompt_cache_min_tokens": 4096
|
||||
"prompt_cache_min_tokens": 4096,
|
||||
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/model-retirement-schedule"
|
||||
},
|
||||
"azure_ai/claude-opus-4-5": {
|
||||
"deprecation_date": "2026-10-19",
|
||||
"deprecation_date": "2026-11-24",
|
||||
"cache_creation_input_token_cost": 6.25e-06,
|
||||
"cache_creation_input_token_cost_above_1hr": 1e-05,
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
|
|
@ -3402,7 +3403,8 @@
|
|||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_output_config": true,
|
||||
"prompt_cache_min_tokens": 4096
|
||||
"prompt_cache_min_tokens": 4096,
|
||||
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/model-retirement-schedule"
|
||||
},
|
||||
"azure_ai/claude-opus-4-6": {
|
||||
"deprecation_date": "2027-02-02",
|
||||
|
|
@ -3640,7 +3642,7 @@
|
|||
"prompt_cache_min_tokens": 1024
|
||||
},
|
||||
"azure_ai/claude-sonnet-4-5": {
|
||||
"deprecation_date": "2026-10-19",
|
||||
"deprecation_date": "2026-11-15",
|
||||
"cache_creation_input_token_cost": 3.75e-06,
|
||||
"cache_creation_input_token_cost_above_1hr": 6e-06,
|
||||
"cache_read_input_token_cost": 3e-07,
|
||||
|
|
@ -3660,7 +3662,8 @@
|
|||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"prompt_cache_min_tokens": 1024
|
||||
"prompt_cache_min_tokens": 1024,
|
||||
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/model-retirement-schedule"
|
||||
},
|
||||
"azure_ai/claude-sonnet-5": {
|
||||
"deprecation_date": "2027-06-30",
|
||||
|
|
@ -30721,6 +30724,7 @@
|
|||
"output_cost_per_image": 0.08
|
||||
},
|
||||
"gemini/veo-3.1-fast-generate-preview": {
|
||||
"deprecation_date": "2026-10-22",
|
||||
"litellm_provider": "gemini",
|
||||
"max_input_tokens": 1024,
|
||||
"max_tokens": 1024,
|
||||
|
|
@ -30737,6 +30741,7 @@
|
|||
]
|
||||
},
|
||||
"gemini/veo-3.1-generate-preview": {
|
||||
"deprecation_date": "2026-10-22",
|
||||
"litellm_provider": "gemini",
|
||||
"max_input_tokens": 1024,
|
||||
"max_tokens": 1024,
|
||||
|
|
@ -30752,6 +30757,7 @@
|
|||
]
|
||||
},
|
||||
"gemini/veo-3.1-lite-generate-preview": {
|
||||
"deprecation_date": "2026-10-22",
|
||||
"litellm_provider": "gemini",
|
||||
"max_input_tokens": 1024,
|
||||
"max_tokens": 1024,
|
||||
|
|
@ -32912,10 +32918,13 @@
|
|||
"gpt-image-2.5-flare": {
|
||||
"cache_read_input_image_token_cost": 2e-06,
|
||||
"cache_read_input_token_cost": 1.25e-06,
|
||||
"cache_read_input_token_cost_batches": 6.25e-07,
|
||||
"input_cost_per_token": 5e-06,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "image_generation",
|
||||
"input_cost_per_image_token": 8e-06,
|
||||
"input_cost_per_image_token_batches": 4e-06,
|
||||
"input_cost_per_token_batches": 2.5e-06,
|
||||
"output_cost_per_image_token": 3e-05,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations",
|
||||
|
|
@ -32944,10 +32953,13 @@
|
|||
"gpt-image-2.5-sunburst": {
|
||||
"cache_read_input_image_token_cost": 2e-06,
|
||||
"cache_read_input_token_cost": 1.25e-06,
|
||||
"cache_read_input_token_cost_batches": 6.25e-07,
|
||||
"input_cost_per_token": 5e-06,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "image_generation",
|
||||
"input_cost_per_image_token": 8e-06,
|
||||
"input_cost_per_image_token_batches": 4e-06,
|
||||
"input_cost_per_token_batches": 2.5e-06,
|
||||
"output_cost_per_image_token": 3e-05,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations",
|
||||
|
|
@ -38611,6 +38623,7 @@
|
|||
},
|
||||
"mistral/zai-glm-5-2": {
|
||||
"cache_read_input_token_cost": 1.4e-07,
|
||||
"deprecation_date": "2026-10-31",
|
||||
"input_cost_per_token": 1.4e-06,
|
||||
"litellm_provider": "mistral",
|
||||
"max_input_tokens": 1048576,
|
||||
|
|
@ -38741,6 +38754,7 @@
|
|||
"source": "https://mistral.ai/pricing#api-pricing"
|
||||
},
|
||||
"mistral/mistral-ocr-4-0": {
|
||||
"deprecation_date": "2026-09-30",
|
||||
"litellm_provider": "mistral",
|
||||
"ocr_cost_per_page": 0.004,
|
||||
"ocr_cost_per_page_batches": 0.002,
|
||||
|
|
@ -60616,6 +60630,7 @@
|
|||
"supports_vision": true
|
||||
},
|
||||
"mistral/labs-leanstral-1-5": {
|
||||
"deprecation_date": "2026-09-30",
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "mistral",
|
||||
"max_input_tokens": 262144,
|
||||
|
|
@ -61337,13 +61352,16 @@
|
|||
},
|
||||
"fireworks_ai/nemotron-lightning-3p5-30b-a3b": {
|
||||
"cache_read_input_token_cost": 1e-08,
|
||||
"cache_read_input_token_cost_priority": 1.25e-08,
|
||||
"input_cost_per_token": 5e-08,
|
||||
"input_cost_per_token_priority": 6.25e-08,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 262144,
|
||||
"max_output_tokens": 32768,
|
||||
"max_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2e-07,
|
||||
"output_cost_per_token_priority": 2.5e-07,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -61353,13 +61371,16 @@
|
|||
},
|
||||
"fireworks_ai/nemotron-3-ultra-nvfp4": {
|
||||
"cache_read_input_token_cost": 1.2e-07,
|
||||
"cache_read_input_token_cost_priority": 1.5e-07,
|
||||
"input_cost_per_token": 6e-07,
|
||||
"input_cost_per_token_priority": 7.5e-07,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 262144,
|
||||
"max_output_tokens": 32768,
|
||||
"max_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"output_cost_per_token_priority": 3e-06,
|
||||
"source": "https://api.fireworks.ai/v1/serverless/models",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -61389,13 +61410,16 @@
|
|||
},
|
||||
"fireworks_ai/accounts/fireworks/models/nemotron-lightning-3p5-30b-a3b": {
|
||||
"cache_read_input_token_cost": 1e-08,
|
||||
"cache_read_input_token_cost_priority": 1.25e-08,
|
||||
"input_cost_per_token": 5e-08,
|
||||
"input_cost_per_token_priority": 6.25e-08,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 262144,
|
||||
"max_output_tokens": 32768,
|
||||
"max_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2e-07,
|
||||
"output_cost_per_token_priority": 2.5e-07,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -61405,13 +61429,16 @@
|
|||
},
|
||||
"fireworks_ai/accounts/fireworks/models/nemotron-3-ultra-nvfp4": {
|
||||
"cache_read_input_token_cost": 1.2e-07,
|
||||
"cache_read_input_token_cost_priority": 1.5e-07,
|
||||
"input_cost_per_token": 6e-07,
|
||||
"input_cost_per_token_priority": 7.5e-07,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 262144,
|
||||
"max_output_tokens": 32768,
|
||||
"max_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"output_cost_per_token_priority": 3e-06,
|
||||
"source": "https://api.fireworks.ai/v1/serverless/models",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -64321,13 +64348,16 @@
|
|||
},
|
||||
"fireworks_ai/accounts/fireworks/routers/glm-5p3-us": {
|
||||
"cache_read_input_token_cost": 3.9e-07,
|
||||
"cache_read_input_token_cost_priority": 4.875e-07,
|
||||
"input_cost_per_token": 2.1e-06,
|
||||
"input_cost_per_token_priority": 2.625e-06,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 6.6e-06,
|
||||
"output_cost_per_token_priority": 8.25e-06,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -64356,13 +64386,16 @@
|
|||
},
|
||||
"fireworks_ai/glm-5p3-us": {
|
||||
"cache_read_input_token_cost": 3.9e-07,
|
||||
"cache_read_input_token_cost_priority": 4.875e-07,
|
||||
"input_cost_per_token": 2.1e-06,
|
||||
"input_cost_per_token_priority": 2.625e-06,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 6.6e-06,
|
||||
"output_cost_per_token_priority": 8.25e-06,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -64446,12 +64479,15 @@
|
|||
},
|
||||
"fireworks_ai/accounts/fireworks/routers/glm-5p3-flash-us": {
|
||||
"cache_read_input_token_cost": 4.5e-08,
|
||||
"cache_read_input_token_cost_priority": 5.625e-08,
|
||||
"input_cost_per_token": 2.25e-07,
|
||||
"input_cost_per_token_priority": 2.8125e-07,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_tokens": 1048576,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 7.5e-07,
|
||||
"output_cost_per_token_priority": 9.375e-07,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_response_schema": true,
|
||||
|
|
@ -64477,12 +64513,15 @@
|
|||
},
|
||||
"fireworks_ai/glm-5p3-flash-us": {
|
||||
"cache_read_input_token_cost": 4.5e-08,
|
||||
"cache_read_input_token_cost_priority": 5.625e-08,
|
||||
"input_cost_per_token": 2.25e-07,
|
||||
"input_cost_per_token_priority": 2.8125e-07,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_tokens": 1048576,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 7.5e-07,
|
||||
"output_cost_per_token_priority": 9.375e-07,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_response_schema": true,
|
||||
|
|
@ -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"
|
||||
]
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -0,0 +1,74 @@
|
|||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_managed_agent_ceilings
|
||||
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityFailure
|
||||
|
||||
|
||||
async def _delegated_resource_subject(user_id: str) -> UserAPIKeyAuth:
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
|
||||
|
||||
human: Final = await MCPRequestHandler.reload_admitted_user(user_id, requires_fresh_policy=True)
|
||||
return human.model_copy(update=MappingProxyType({"mcp_explicit_grants_only": True}))
|
||||
|
||||
|
||||
async def managed_agent_servers(auth: UserAPIKeyAuth) -> tuple[str, ...]:
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
agent: Final = auth.managed_agent_policy
|
||||
if agent is None:
|
||||
return ()
|
||||
|
||||
try:
|
||||
base: Final = frozenset(await MCPRequestHandler.get_allowed_mcp_servers_for_agent(auth))
|
||||
ceilings: Final = await resolve_managed_agent_ceilings(agent)
|
||||
expanded: Final = tuple(
|
||||
frozenset(global_mcp_server_manager.expand_permission_list(sorted(ceiling.mcp_server_ids)))
|
||||
for ceiling in ceilings
|
||||
)
|
||||
grouped: Final = frozenset(server for server in base if all(server in ceiling for ceiling in expanded))
|
||||
caller_capped, _ = await MCPRequestHandler.apply_agent_caller_ceiling(sorted(grouped), auth)
|
||||
own: Final = frozenset(caller_capped)
|
||||
context: Final = auth.managed_agent_context
|
||||
if context is None or context.mode == "autonomous":
|
||||
return tuple(sorted(own))
|
||||
if context.user_id is None:
|
||||
return ()
|
||||
human: Final = await _delegated_resource_subject(context.user_id)
|
||||
allowed: Final = await MCPRequestHandler.resolve_admitted_subject_servers(
|
||||
human, allowed_team_ids=frozenset((auth.team_id,)) if auth.team_id else frozenset()
|
||||
)
|
||||
return tuple(sorted(own.intersection(allowed)))
|
||||
except Exception: # noqa: BLE001 # Authorization boundary: every unresolved policy must deny access
|
||||
raise_identity_failure(
|
||||
AgentIdentityFailure(code="policy_unavailable", message="Agent MCP policy is unavailable")
|
||||
)
|
||||
|
||||
|
||||
async def managed_agent_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None:
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
|
||||
|
||||
if server_id not in await managed_agent_servers(auth):
|
||||
return []
|
||||
try:
|
||||
granted: Final = await MCPRequestHandler.get_agent_tool_permissions_for_server(server_id, auth)
|
||||
own: Final = await MCPRequestHandler.apply_agent_caller_tool_ceiling(granted, server_id, auth)
|
||||
context: Final = auth.managed_agent_context
|
||||
if context is None or context.mode == "autonomous":
|
||||
return None if own is None else sorted(own)
|
||||
if context.user_id is None:
|
||||
return []
|
||||
human: Final = await _delegated_resource_subject(context.user_id)
|
||||
human_tools: Final = await MCPRequestHandler.resolve_admitted_subject_tools(
|
||||
server_id, human, allowed_team_ids=frozenset((auth.team_id,)) if auth.team_id else frozenset()
|
||||
)
|
||||
if own is None:
|
||||
return human_tools
|
||||
return sorted(own) if human_tools is None else sorted(frozenset(own).intersection(human_tools))
|
||||
except Exception: # noqa: BLE001 # Authorization boundary: every unresolved policy must deny access
|
||||
raise_identity_failure(
|
||||
AgentIdentityFailure(code="policy_unavailable", message="Agent tool policy is unavailable")
|
||||
)
|
||||
|
|
@ -51,6 +51,7 @@ from litellm.proxy.agent_endpoints.auth.agent_access_groups import (
|
|||
resolve_agent_access_group_ceiling,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.auth.agent_caller import agent_caller_auth
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_policy
|
||||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
from litellm.proxy.auth.user_api_key_auth import (
|
||||
_get_bearer_token_or_received_api_key, # pyright: ignore[reportPrivateUsage] # shared x-litellm-api-key parser lives with user_api_key_auth
|
||||
|
|
@ -67,7 +68,6 @@ from litellm.repositories.table_repositories import (
|
|||
AgentsRepository,
|
||||
MCPServerRepository,
|
||||
)
|
||||
from litellm.repositories.user_repository import UserRepository
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -1086,7 +1086,7 @@ class MCPRequestHandler:
|
|||
assert_never(identity.subject_type)
|
||||
|
||||
@staticmethod
|
||||
async def reload_admitted_user(user_id: str) -> UserAPIKeyAuth:
|
||||
async def reload_admitted_user(user_id: str, *, requires_fresh_policy: bool = False) -> UserAPIKeyAuth:
|
||||
"""Reload the live user an interactively-minted envelope references and admit them as themselves.
|
||||
|
||||
The user's own object permission and ``org_id`` ride on the returned ``UserAPIKeyAuth``, and the
|
||||
|
|
@ -1111,6 +1111,7 @@ class MCPRequestHandler:
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
check_db_only=requires_fresh_policy,
|
||||
)
|
||||
# Resolve the user's own MCP object permission (get_user_object does not load it) so the shared
|
||||
# get_allowed_mcp_servers can grant the user their litellm-granted servers. Reuses the same
|
||||
|
|
@ -1119,6 +1120,7 @@ class MCPRequestHandler:
|
|||
if user_object is not None and object_permission is None and user_object.object_permission_id:
|
||||
object_permission = await get_object_permission(
|
||||
object_permission_id=user_object.object_permission_id,
|
||||
check_db_only=requires_fresh_policy,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
|
|
@ -1147,6 +1149,7 @@ class MCPRequestHandler:
|
|||
# Server-only marker, set AFTER construction: the before-validator strips it from any validated
|
||||
# input, so caller-supplied data (key metadata, JWT claims) can never forge it.
|
||||
admitted.mcp_admitted_user_subject = True
|
||||
admitted.requires_fresh_policy = requires_fresh_policy
|
||||
# Carry each granting team's per-server mcp_rpm_limit: this subject reaches servers through
|
||||
# several teams under its own identity, so without this a cross-team user outruns every team's
|
||||
# limit. Resolved from the same roster-checked sources as the grant union, so a team throttles
|
||||
|
|
@ -1202,7 +1205,7 @@ class MCPRequestHandler:
|
|||
return None
|
||||
|
||||
@staticmethod
|
||||
async def _reload_admitted_key(key_hash: str) -> UserAPIKeyAuth:
|
||||
async def _reload_admitted_key(key_hash: str, *, check_db_only: bool = False) -> UserAPIKeyAuth:
|
||||
"""Reload the live key record an admitted envelope references and re-check live policy.
|
||||
|
||||
Resolving the current ``UserAPIKeyAuth`` (cache first, then DB) is what stops the
|
||||
|
|
@ -1234,6 +1237,7 @@ class MCPRequestHandler:
|
|||
hashed_token=key_hash,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
check_db_only=check_db_only,
|
||||
)
|
||||
except (ProxyException, HTTPException):
|
||||
raise HTTPException(status_code=401, detail="Invalid or expired credential") from None
|
||||
|
|
@ -1597,6 +1601,11 @@ class MCPRequestHandler:
|
|||
"""
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
if managed_agent_policy(user_api_key_auth) is not None:
|
||||
from litellm.proxy._experimental.mcp_server.auth.managed_agent_access import managed_agent_servers
|
||||
|
||||
return MCPServerAccess(server_ids=await managed_agent_servers(user_api_key_auth), scope="scoped")
|
||||
|
||||
key_object_permission: Final = MCPRequestHandler._get_key_object_permission(user_api_key_auth)
|
||||
|
||||
try:
|
||||
|
|
@ -1606,7 +1615,7 @@ class MCPRequestHandler:
|
|||
# independent; an opt-out silences only its own source, inside the recursive call).
|
||||
if _is_mcp_admitted_user_subject(user_api_key_auth) and user_api_key_auth is not None:
|
||||
return MCPServerAccess(
|
||||
server_ids=tuple(await MCPRequestHandler._resolve_admitted_subject_servers(user_api_key_auth)),
|
||||
server_ids=tuple(await MCPRequestHandler.resolve_admitted_subject_servers(user_api_key_auth)),
|
||||
)
|
||||
|
||||
# Get allowed servers from key and team
|
||||
|
|
@ -1703,7 +1712,7 @@ class MCPRequestHandler:
|
|||
if user_api_key_auth and user_api_key_auth.agent_id:
|
||||
agent_capped: Final = _agent_capped_servers(
|
||||
allowed_mcp_servers,
|
||||
await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth),
|
||||
await MCPRequestHandler.get_allowed_mcp_servers_for_agent(user_api_key_auth),
|
||||
await MCPRequestHandler._get_agent_access_group_server_ceiling(user_api_key_auth),
|
||||
)
|
||||
if agent_capped is not None:
|
||||
|
|
@ -1716,7 +1725,7 @@ class MCPRequestHandler:
|
|||
#########################################################
|
||||
# Cap an agent key at what the user and team that invoked the agent may reach
|
||||
#########################################################
|
||||
caller_capped, caller_restricts = await MCPRequestHandler._apply_agent_caller_ceiling(
|
||||
caller_capped, caller_restricts = await MCPRequestHandler.apply_agent_caller_ceiling(
|
||||
allowed_mcp_servers, user_api_key_auth
|
||||
)
|
||||
|
||||
|
|
@ -1829,10 +1838,14 @@ class MCPRequestHandler:
|
|||
scoped.object_permission = auth.object_permission
|
||||
scoped.object_permission_id = auth.object_permission_id
|
||||
scoped.access_group_ids = auth.access_group_ids
|
||||
scoped.requires_fresh_policy = auth.requires_fresh_policy
|
||||
scoped.mcp_explicit_grants_only = auth.mcp_explicit_grants_only
|
||||
return scoped
|
||||
|
||||
@staticmethod
|
||||
async def _admitted_subject_sources(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]:
|
||||
async def admitted_subject_sources(
|
||||
auth: UserAPIKeyAuth, *, allowed_team_ids: frozenset[str] | None = None
|
||||
) -> list[UserAPIKeyAuth]:
|
||||
"""The independent sources a keyless admitted subject reaches MCP servers through: their own
|
||||
direct grants, plus every team they are a live roster member of.
|
||||
|
||||
|
|
@ -1849,6 +1862,8 @@ class MCPRequestHandler:
|
|||
if not auth.user_id or prisma_client is None:
|
||||
return sources
|
||||
for team_id in await MCPRequestHandler._resolve_user_team_ids(auth.user_id, auth):
|
||||
if allowed_team_ids is not None and team_id not in allowed_team_ids:
|
||||
continue
|
||||
team_obj = await MCPRequestHandler._roster_team_object(team_id, auth)
|
||||
if team_obj is None:
|
||||
continue
|
||||
|
|
@ -1886,6 +1901,7 @@ class MCPRequestHandler:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(auth and auth.requires_fresh_policy),
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # per-source isolation: one team's blip must not deny the others
|
||||
# Fault isolation is per SOURCE: an unresolvable team contributes nothing (fail closed for
|
||||
|
|
@ -1932,7 +1948,9 @@ class MCPRequestHandler:
|
|||
return team_obj
|
||||
|
||||
@staticmethod
|
||||
async def admitted_source_grants(auth: UserAPIKeyAuth) -> list[tuple[UserAPIKeyAuth, set[str]]]:
|
||||
async def admitted_source_grants(
|
||||
auth: UserAPIKeyAuth, *, allowed_team_ids: frozenset[str] | None = None
|
||||
) -> list[tuple[UserAPIKeyAuth, set[str]]]:
|
||||
"""``(source, the servers that source grants)`` for every source of an admitted subject.
|
||||
|
||||
THE owner of "which source reaches which server". The reachable union, the per-team throttle
|
||||
|
|
@ -1941,15 +1959,17 @@ class MCPRequestHandler:
|
|||
roster instead of by grant charged unrelated teams' buckets)."""
|
||||
return [
|
||||
(source, set(await MCPRequestHandler.get_allowed_mcp_servers(source, keyless_source=True)))
|
||||
for source in await MCPRequestHandler._admitted_subject_sources(auth)
|
||||
for source in await MCPRequestHandler.admitted_subject_sources(auth, allowed_team_ids=allowed_team_ids)
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
async def _resolve_admitted_subject_servers(auth: UserAPIKeyAuth) -> list[str]:
|
||||
async def resolve_admitted_subject_servers(
|
||||
auth: UserAPIKeyAuth, *, allowed_team_ids: frozenset[str] | None = None
|
||||
) -> list[str]:
|
||||
"""Union of what each of the admitted subject's sources reaches, each answered by the
|
||||
canonical resolver so no rule is reimplemented for this caller shape."""
|
||||
reachable: Final[set[str]] = set()
|
||||
for _source, granted in await MCPRequestHandler.admitted_source_grants(auth):
|
||||
for _source, granted in await MCPRequestHandler.admitted_source_grants(auth, allowed_team_ids=allowed_team_ids):
|
||||
reachable.update(granted)
|
||||
return list(reachable)
|
||||
|
||||
|
|
@ -2007,7 +2027,9 @@ class MCPRequestHandler:
|
|||
return min((source for source, _ in granting), key=lambda s: s.team_id or "")
|
||||
|
||||
@staticmethod
|
||||
async def _resolve_admitted_subject_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None:
|
||||
async def resolve_admitted_subject_tools(
|
||||
server_id: str, auth: UserAPIKeyAuth, *, allowed_team_ids: frozenset[str] | None = None
|
||||
) -> list[str] | None:
|
||||
"""Effective tool allowlist on ``server_id`` for an admitted subject, as the union over the
|
||||
sources that actually grant that server.
|
||||
|
||||
|
|
@ -2029,7 +2051,7 @@ class MCPRequestHandler:
|
|||
) or await MCPRequestHandler.admin_view_unscoped(auth)
|
||||
|
||||
allowed: Final[set[str]] = set()
|
||||
for source, granted in await MCPRequestHandler.admitted_source_grants(auth):
|
||||
for source, granted in await MCPRequestHandler.admitted_source_grants(auth, allowed_team_ids=allowed_team_ids):
|
||||
# The open channel is evaluated against the user's OWN source (team_id is None), so that
|
||||
# source's restrictions apply to it; a team's rules never ride an open-channel server.
|
||||
if server_id not in granted and not (reachable_via_open_channel and source.team_id is None):
|
||||
|
|
@ -2088,6 +2110,7 @@ class MCPRequestHandler:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
|
||||
if not team_obj:
|
||||
|
|
@ -2098,6 +2121,8 @@ class MCPRequestHandler:
|
|||
@staticmethod
|
||||
async def _toolset_tool_permissions(
|
||||
object_permission: LiteLLM_ObjectPermissionTable | None,
|
||||
*,
|
||||
requires_fresh_policy: bool = False,
|
||||
) -> Mapping[str, Sequence[str]]:
|
||||
"""The ``server_id -> tool names`` grants of this permission row's toolsets, empty when it
|
||||
declares none. The shared resolver for the team, org, and internal-user levels, so a toolset
|
||||
|
|
@ -2114,7 +2139,8 @@ class MCPRequestHandler:
|
|||
if object_permission is None or not object_permission.mcp_toolsets:
|
||||
return _EMPTY_TOOLSET_GRANTS
|
||||
resolved: Final = await global_mcp_server_manager.resolve_toolset_tool_permissions(
|
||||
toolset_ids=object_permission.mcp_toolsets
|
||||
toolset_ids=object_permission.mcp_toolsets,
|
||||
requires_fresh_policy=requires_fresh_policy,
|
||||
)
|
||||
if not resolved:
|
||||
raise UnloadableEntitlementError(
|
||||
|
|
@ -2126,10 +2152,15 @@ class MCPRequestHandler:
|
|||
async def _toolset_tools_for_server(
|
||||
object_permission: LiteLLM_ObjectPermissionTable | None,
|
||||
server_id: str,
|
||||
*,
|
||||
requires_fresh_policy: bool = False,
|
||||
) -> Sequence[str] | None:
|
||||
"""Tool names this row's toolsets grant on ``server_id``, ``None`` when its toolsets place
|
||||
no restriction on that server (it declares no toolsets, or none of them name it)."""
|
||||
return (await MCPRequestHandler._toolset_tool_permissions(object_permission)).get(server_id)
|
||||
grants: Final = await MCPRequestHandler._toolset_tool_permissions(
|
||||
object_permission, requires_fresh_policy=requires_fresh_policy
|
||||
)
|
||||
return grants.get(server_id)
|
||||
|
||||
@staticmethod
|
||||
def _union_tool_grants(
|
||||
|
|
@ -2171,6 +2202,7 @@ class MCPRequestHandler:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -2219,12 +2251,17 @@ class MCPRequestHandler:
|
|||
if not user_api_key_auth:
|
||||
return None
|
||||
|
||||
if managed_agent_policy(user_api_key_auth) is not None:
|
||||
from litellm.proxy._experimental.mcp_server.auth.managed_agent_access import managed_agent_tools
|
||||
|
||||
return await managed_agent_tools(server_id, user_api_key_auth)
|
||||
|
||||
try:
|
||||
# FIRST statement, mirroring get_allowed_mcp_servers: a keyless admitted subject resolves per
|
||||
# source and shares nothing with the single-credential prelude below. Ordering is the invariant:
|
||||
# sat after the prelude, a fault in a lookup the subject never uses denied tools its teams grant.
|
||||
if _is_mcp_admitted_user_subject(user_api_key_auth):
|
||||
return await MCPRequestHandler._resolve_admitted_subject_tools(server_id, user_api_key_auth)
|
||||
return await MCPRequestHandler.resolve_admitted_subject_tools(server_id, user_api_key_auth)
|
||||
|
||||
# Get key and team object permissions (already loaded in main auth flow)
|
||||
key_obj_perm: Final = MCPRequestHandler._get_key_object_permission(user_api_key_auth)
|
||||
|
|
@ -2249,9 +2286,12 @@ class MCPRequestHandler:
|
|||
# tool-level check sees the key's full effective tool scope
|
||||
key_toolset_ids: Final = (key_obj_perm.mcp_toolsets or []) if key_obj_perm else []
|
||||
key_toolset_tools: Final = (
|
||||
(await global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=key_toolset_ids)).get(
|
||||
server_id
|
||||
)
|
||||
(
|
||||
await global_mcp_server_manager.resolve_toolset_tool_permissions(
|
||||
toolset_ids=key_toolset_ids,
|
||||
requires_fresh_policy=user_api_key_auth.requires_fresh_policy,
|
||||
)
|
||||
).get(server_id)
|
||||
if key_toolset_ids
|
||||
else None
|
||||
)
|
||||
|
|
@ -2265,7 +2305,9 @@ class MCPRequestHandler:
|
|||
|
||||
# Tools granted through the team's toolsets restrict this server exactly
|
||||
# as the team's direct tool permissions do, mirroring the key path above
|
||||
team_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(team_obj_perm, server_id)
|
||||
team_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
|
||||
team_obj_perm, server_id, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
|
||||
)
|
||||
team_tools: Final = MCPRequestHandler._union_tool_grants(team_direct_tools, team_toolset_tools)
|
||||
|
||||
# Apply same inheritance logic as get_allowed_mcp_servers
|
||||
|
|
@ -2291,7 +2333,7 @@ class MCPRequestHandler:
|
|||
)
|
||||
|
||||
allowed_tools = _as_list(
|
||||
await MCPRequestHandler._apply_agent_caller_tool_ceiling(allowed_tools, server_id, user_api_key_auth)
|
||||
await MCPRequestHandler.apply_agent_caller_tool_ceiling(allowed_tools, server_id, user_api_key_auth)
|
||||
)
|
||||
|
||||
return await MCPRequestHandler._apply_agent_and_org_tool_ceilings(
|
||||
|
|
@ -2334,7 +2376,7 @@ class MCPRequestHandler:
|
|||
if user_api_key_auth.agent_id:
|
||||
# Pre-fetch agent object_permission once to avoid a duplicate DB query.
|
||||
agent_obj_perm: Final = await MCPRequestHandler._get_agent_object_permission(user_api_key_auth)
|
||||
agent_tools: Final = await MCPRequestHandler._get_agent_tool_permissions_for_server(
|
||||
agent_tools: Final = await MCPRequestHandler.get_agent_tool_permissions_for_server(
|
||||
server_id=server_id,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
agent_object_permission=agent_obj_perm,
|
||||
|
|
@ -2365,7 +2407,9 @@ class MCPRequestHandler:
|
|||
if org_obj_perm and org_obj_perm.mcp_tool_permissions
|
||||
else None
|
||||
)
|
||||
org_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(org_obj_perm, server_id)
|
||||
org_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
|
||||
org_obj_perm, server_id, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
|
||||
)
|
||||
org_tools: Final = MCPRequestHandler._union_tool_grants(org_direct_tools, org_toolset_tools)
|
||||
if org_tools is not None:
|
||||
allowed_tools = (
|
||||
|
|
@ -2456,6 +2500,7 @@ class MCPRequestHandler:
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
if not raw_server_ids:
|
||||
return []
|
||||
|
|
@ -2502,6 +2547,7 @@ class MCPRequestHandler:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
if key_object_permission is None:
|
||||
return []
|
||||
|
|
@ -2518,7 +2564,8 @@ class MCPRequestHandler:
|
|||
|
||||
# Get MCP servers from access groups
|
||||
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
|
||||
key_object_permission.mcp_access_groups or []
|
||||
key_object_permission.mcp_access_groups or [],
|
||||
requires_fresh_policy=user_api_key_auth.requires_fresh_policy,
|
||||
)
|
||||
|
||||
# servers referenced in tool permissions should also be accessible
|
||||
|
|
@ -2531,7 +2578,14 @@ class MCPRequestHandler:
|
|||
# ceilings as any other key-level grant
|
||||
toolset_ids: Final = key_object_permission.mcp_toolsets or []
|
||||
toolset_servers: Final = (
|
||||
list((await global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=toolset_ids)).keys())
|
||||
list(
|
||||
(
|
||||
await global_mcp_server_manager.resolve_toolset_tool_permissions(
|
||||
toolset_ids=toolset_ids,
|
||||
requires_fresh_policy=user_api_key_auth.requires_fresh_policy,
|
||||
)
|
||||
).keys()
|
||||
)
|
||||
if toolset_ids
|
||||
else []
|
||||
)
|
||||
|
|
@ -2550,7 +2604,7 @@ class MCPRequestHandler:
|
|||
"""Get allowed MCP servers a caller inherits from the team it is pinned to.
|
||||
|
||||
Exactly one team, or none. A subject that reaches servers through SEVERAL teams does not
|
||||
fan out here: it is resolved one source per team in ``_resolve_admitted_subject_servers``,
|
||||
fan out here: it is resolved one source per team in ``resolve_admitted_subject_servers``,
|
||||
and each of those sources pins a single ``team_id`` before reaching this point. Keeping the
|
||||
fan-out here as well would be a second multi-team path to drift from that one.
|
||||
"""
|
||||
|
|
@ -2568,7 +2622,7 @@ class MCPRequestHandler:
|
|||
which must NOT silently gain the union across every team the user belongs to), and it covers
|
||||
each single-source auth an admitted subject fans out into — those pin a team_id, so they land
|
||||
on the first branch. The admitted subject itself never reaches here: it resolves per source
|
||||
in ``_resolve_admitted_subject_servers`` before this point. The ``UI_TEAM_ID`` sentinel
|
||||
in ``resolve_admitted_subject_servers`` before this point. The ``UI_TEAM_ID`` sentinel
|
||||
resolves to no teams exactly as before."""
|
||||
if user_api_key_auth is None or not user_api_key_auth.team_id:
|
||||
return []
|
||||
|
|
@ -2596,6 +2650,7 @@ class MCPRequestHandler:
|
|||
user_id_upsert=False,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # a team-resolution blip narrows access, never raises
|
||||
verbose_logger.warning("Failed to resolve user teams for MCP grant: %s", e)
|
||||
|
|
@ -2605,7 +2660,12 @@ class MCPRequestHandler:
|
|||
return list(dict.fromkeys(t for t in user_object.teams if t and t != UI_TEAM_ID))
|
||||
|
||||
@staticmethod
|
||||
async def _team_granted_servers(team_obj: LiteLLM_TeamTable, team_access_group_servers: list[str]) -> set[str]:
|
||||
async def _team_granted_servers(
|
||||
team_obj: LiteLLM_TeamTable,
|
||||
team_access_group_servers: list[str],
|
||||
*,
|
||||
requires_fresh_policy: bool = False,
|
||||
) -> set[str]:
|
||||
"""The raw MCP-server set a team grants (before any org ceiling): its object_permission (direct
|
||||
``mcp_servers``, the ``all_proxy_servers`` sentinel → the full registry, legacy access groups,
|
||||
tool-perm-referenced servers, toolset-referenced servers) unioned with its unified
|
||||
|
|
@ -2620,13 +2680,17 @@ class MCPRequestHandler:
|
|||
if SpecialMCPServerName.all_proxy_servers.value in (object_permissions.mcp_servers or []):
|
||||
return set(global_mcp_server_manager.get_registry().keys())
|
||||
legacy_access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
|
||||
object_permissions.mcp_access_groups or []
|
||||
object_permissions.mcp_access_groups or [],
|
||||
requires_fresh_policy=requires_fresh_policy,
|
||||
)
|
||||
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(
|
||||
object_permissions, requires_fresh_policy=requires_fresh_policy
|
||||
)
|
||||
return (
|
||||
set(global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or []))
|
||||
| set(legacy_access_group_servers)
|
||||
| set(global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys())
|
||||
| (await MCPRequestHandler._toolset_tool_permissions(object_permissions)).keys()
|
||||
| toolset_grants.keys()
|
||||
| set(team_access_group_servers)
|
||||
)
|
||||
|
||||
|
|
@ -2667,6 +2731,7 @@ class MCPRequestHandler:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
if team_obj is None:
|
||||
return []
|
||||
|
|
@ -2680,12 +2745,19 @@ class MCPRequestHandler:
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
|
||||
servers: Final = await MCPRequestHandler._team_granted_servers(team_obj, team_access_group_servers)
|
||||
servers: Final = await MCPRequestHandler._team_granted_servers(
|
||||
team_obj,
|
||||
team_access_group_servers,
|
||||
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
return list(servers)
|
||||
except Exception as e:
|
||||
if isinstance(e, UnloadableEntitlementError):
|
||||
if isinstance(e, UnloadableEntitlementError) or (
|
||||
user_api_key_auth is not None and user_api_key_auth.requires_fresh_policy
|
||||
):
|
||||
raise
|
||||
verbose_logger.warning("Failed to get allowed MCP servers for team: %s", e)
|
||||
return []
|
||||
|
|
@ -2716,6 +2788,7 @@ class MCPRequestHandler:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # a named entitlement we cannot read denies, whatever the read failed with
|
||||
raise unloadable from e
|
||||
|
|
@ -2811,7 +2884,8 @@ class MCPRequestHandler:
|
|||
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or [])
|
||||
|
||||
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
|
||||
object_permissions.mcp_access_groups or []
|
||||
object_permissions.mcp_access_groups or [],
|
||||
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
|
||||
tool_perm_servers: Final = list(
|
||||
|
|
@ -2820,7 +2894,10 @@ class MCPRequestHandler:
|
|||
|
||||
# servers referenced by the org's toolset grants are part of the org ceiling,
|
||||
# exactly as servers referenced by its inline tool permissions are
|
||||
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(object_permissions)
|
||||
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(
|
||||
object_permissions,
|
||||
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
|
||||
all_servers: Final = tuple(
|
||||
{*direct_mcp_servers, *access_group_servers, *tool_perm_servers, *toolset_grants}
|
||||
|
|
@ -2912,7 +2989,8 @@ class MCPRequestHandler:
|
|||
|
||||
# Get MCP servers from access groups
|
||||
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
|
||||
object_permission.mcp_access_groups or []
|
||||
object_permission.mcp_access_groups or [],
|
||||
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
|
||||
# servers referenced in tool permissions should also be accessible
|
||||
|
|
@ -2961,7 +3039,9 @@ class MCPRequestHandler:
|
|||
return None
|
||||
|
||||
user_id: Final = user_api_key_auth.user_id
|
||||
object_permission_id: Final = await MCPRequestHandler._user_object_permission_id(user_id, prisma_client)
|
||||
object_permission_id: Final = await MCPRequestHandler._user_object_permission_id(
|
||||
user_id, prisma_client, check_db_only=user_api_key_auth.requires_fresh_policy
|
||||
)
|
||||
if object_permission_id is None:
|
||||
return None
|
||||
|
||||
|
|
@ -2971,6 +3051,7 @@ class MCPRequestHandler:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
if object_permission is None:
|
||||
raise ValueError(
|
||||
|
|
@ -2979,7 +3060,9 @@ class MCPRequestHandler:
|
|||
return object_permission
|
||||
|
||||
@staticmethod
|
||||
async def _user_object_permission_id(user_id: str, prisma_client: "PrismaClient") -> str | None:
|
||||
async def _user_object_permission_id(
|
||||
user_id: str, prisma_client: "PrismaClient", *, check_db_only: bool = False
|
||||
) -> str | None:
|
||||
"""The permission row this human's user row links to, or None when they link none.
|
||||
|
||||
Caches the link (with a sentinel for "links none") so a human without an entitlement costs no
|
||||
|
|
@ -2988,16 +3071,23 @@ class MCPRequestHandler:
|
|||
whether someone is entitled is the state that existed before this level, so it places no
|
||||
ceiling. Only a link we DID resolve can make the caller deny.
|
||||
"""
|
||||
from litellm.proxy.auth.auth_checks import get_user_object
|
||||
from litellm.proxy.proxy_server import user_api_key_cache
|
||||
|
||||
cache_key: Final = user_object_permission_id_cache_key(user_id)
|
||||
try:
|
||||
cached: Final[object] = await user_api_key_cache.async_get_cache(key=cache_key)
|
||||
cached: Final[object] = None if check_db_only else await user_api_key_cache.async_get_cache(key=cache_key)
|
||||
if cached == USER_NO_MCP_PERMISSION_SENTINEL:
|
||||
return None
|
||||
if isinstance(cached, str) and cached:
|
||||
return cached
|
||||
user_row: Final = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id})
|
||||
user_row: Final = await get_user_object(
|
||||
user_id=user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
check_db_only=check_db_only,
|
||||
)
|
||||
linked: Final[object] = getattr(user_row, "object_permission_id", None) if user_row is not None else None
|
||||
object_permission_id: Final = linked if isinstance(linked, str) and linked else None
|
||||
await user_api_key_cache.async_set_cache(
|
||||
|
|
@ -3006,7 +3096,9 @@ class MCPRequestHandler:
|
|||
ttl=get_management_object_ttl(user_api_key_cache),
|
||||
)
|
||||
return object_permission_id
|
||||
except Exception as e: # noqa: BLE001 # unknown whether entitled at all: no ceiling, as before
|
||||
except Exception as e: # noqa: BLE001 # Legacy callers retain their existing optional user-ceiling behavior
|
||||
if check_db_only:
|
||||
raise HTTPException(503, "User policy is unavailable") from e
|
||||
verbose_logger.warning("MCP user entitlement: link for %r unresolved, no ceiling: %s", user_id, e)
|
||||
return None
|
||||
|
||||
|
|
@ -3031,13 +3123,17 @@ class MCPRequestHandler:
|
|||
return []
|
||||
|
||||
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or [])
|
||||
fresh: Final = bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy)
|
||||
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
|
||||
object_permissions.mcp_access_groups or []
|
||||
object_permissions.mcp_access_groups or [],
|
||||
requires_fresh_policy=fresh,
|
||||
)
|
||||
tool_perm_servers: Final = list(
|
||||
global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys()
|
||||
)
|
||||
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(object_permissions)
|
||||
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(
|
||||
object_permissions, requires_fresh_policy=fresh
|
||||
)
|
||||
return tuple({*direct_mcp_servers, *access_group_servers, *tool_perm_servers, *toolset_grants})
|
||||
except Exception as e: # noqa: BLE001 # any resolution fault is an unresolved ceiling, never "no ceiling"
|
||||
verbose_logger.warning("Failed to get allowed MCP servers for user: %s", e)
|
||||
|
|
@ -3075,7 +3171,7 @@ class MCPRequestHandler:
|
|||
return capped, True
|
||||
|
||||
@staticmethod
|
||||
async def _apply_agent_caller_ceiling(
|
||||
async def apply_agent_caller_ceiling(
|
||||
allowed_mcp_servers: Sequence[str],
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
) -> tuple[tuple[str, ...], bool]:
|
||||
|
|
@ -3119,9 +3215,13 @@ class MCPRequestHandler:
|
|||
(any non-empty entitlement, or an unresolved one, disqualifies), exactly as
|
||||
``operator_open_server_ids`` reads the same row. The one owner of this predicate: the
|
||||
server-axis registry resolution in ``get_allowed_mcp_servers`` and the tools-axis open
|
||||
channel in ``_resolve_admitted_subject_tools`` both consult it, so the two axes cannot
|
||||
channel in ``resolve_admitted_subject_tools`` both consult it, so the two axes cannot
|
||||
disagree."""
|
||||
if user_api_key_auth is None or not user_api_key_has_admin_view(user_api_key_auth):
|
||||
if (
|
||||
user_api_key_auth is None
|
||||
or user_api_key_auth.mcp_explicit_grants_only
|
||||
or not user_api_key_has_admin_view(user_api_key_auth)
|
||||
):
|
||||
return False
|
||||
object_permission: Final = user_api_key_auth.object_permission
|
||||
credential_scoped: Final = (
|
||||
|
|
@ -3167,7 +3267,11 @@ class MCPRequestHandler:
|
|||
user_direct_tools: Final = global_mcp_server_manager.expand_tool_permissions(
|
||||
object_permissions.mcp_tool_permissions
|
||||
).get(server_id)
|
||||
user_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(object_permissions, server_id)
|
||||
user_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
|
||||
object_permissions,
|
||||
server_id,
|
||||
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
user_tools: Final = MCPRequestHandler._union_tool_grants(user_direct_tools, user_toolset_tools)
|
||||
if user_tools is None:
|
||||
return allowed_tools
|
||||
|
|
@ -3176,7 +3280,7 @@ class MCPRequestHandler:
|
|||
return list(set(allowed_tools) & set(user_tools))
|
||||
|
||||
@staticmethod
|
||||
async def _apply_agent_caller_tool_ceiling(
|
||||
async def apply_agent_caller_tool_ceiling(
|
||||
allowed_tools: Sequence[str] | None,
|
||||
server_id: str,
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
|
|
@ -3184,7 +3288,7 @@ class MCPRequestHandler:
|
|||
"""Narrow an agent key's tools on ``server_id`` to those the invoking user and team (echoed back
|
||||
by the agent as ``x-litellm-user-id`` / ``x-litellm-team-id``) may call: the echoed team's tool
|
||||
grants when it names any on this server, then the echoed user's own tool entitlement. The tools
|
||||
axis twin of ``_apply_agent_caller_ceiling``, so the headers only ever narrow. Denies every tool
|
||||
axis twin of ``apply_agent_caller_ceiling``, so the headers only ever narrow. Denies every tool
|
||||
on the server when the caller's team cannot be loaded, since a caller we cannot resolve must not
|
||||
read as unrestricted."""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
|
|
@ -3196,7 +3300,9 @@ class MCPRequestHandler:
|
|||
return allowed_tools
|
||||
try:
|
||||
team_obj_perm: Final = await MCPRequestHandler._get_team_object_permission(caller_auth)
|
||||
team_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(team_obj_perm, server_id)
|
||||
team_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
|
||||
team_obj_perm, server_id, requires_fresh_policy=caller_auth.requires_fresh_policy
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # an unresolved caller team must deny, not widen
|
||||
verbose_logger.warning(
|
||||
"MCP agent caller team tool ceiling unresolvable, denying tools on %r: %s", server_id, e
|
||||
|
|
@ -3241,7 +3347,11 @@ class MCPRequestHandler:
|
|||
end_user_direct_tools: Final = global_mcp_server_manager.expand_tool_permissions(
|
||||
object_permissions.mcp_tool_permissions
|
||||
).get(server_id)
|
||||
end_user_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(object_permissions, server_id)
|
||||
end_user_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
|
||||
object_permissions,
|
||||
server_id,
|
||||
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
end_user_tools: Final = MCPRequestHandler._union_tool_grants(end_user_direct_tools, end_user_toolset_tools)
|
||||
if end_user_tools is None:
|
||||
return allowed_tools
|
||||
|
|
@ -3302,6 +3412,11 @@ class MCPRequestHandler:
|
|||
if not user_api_key_auth or not user_api_key_auth.agent_id:
|
||||
return None
|
||||
|
||||
managed: Final = managed_agent_policy(user_api_key_auth)
|
||||
if managed is not None:
|
||||
permission: Final = managed.object_permission
|
||||
return LiteLLM_ObjectPermissionTable.model_validate(permission) if permission is not None else None
|
||||
|
||||
if prisma_client is None:
|
||||
verbose_logger.debug("prisma_client is None")
|
||||
return None
|
||||
|
|
@ -3319,7 +3434,7 @@ class MCPRequestHandler:
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
async def _get_allowed_mcp_servers_for_agent(
|
||||
async def get_allowed_mcp_servers_for_agent(
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
agent_object_permission: LiteLLM_ObjectPermissionTable | None = None,
|
||||
) -> list[str]:
|
||||
|
|
@ -3358,12 +3473,16 @@ class MCPRequestHandler:
|
|||
obj_perm.mcp_servers or []
|
||||
)
|
||||
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
|
||||
obj_perm.mcp_access_groups or []
|
||||
obj_perm.mcp_access_groups or [],
|
||||
requires_fresh_policy=user_api_key_auth.requires_fresh_policy,
|
||||
)
|
||||
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(obj_perm)
|
||||
return list({*expanded_direct_servers, *access_group_servers, *toolset_grants})
|
||||
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(
|
||||
obj_perm, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
|
||||
)
|
||||
inline_tools: Final = global_mcp_server_manager.expand_tool_permissions(obj_perm.mcp_tool_permissions)
|
||||
return list({*expanded_direct_servers, *access_group_servers, *toolset_grants, *inline_tools})
|
||||
except Exception as e:
|
||||
if isinstance(e, UnloadableEntitlementError):
|
||||
if managed_agent_policy(user_api_key_auth) is not None or isinstance(e, UnloadableEntitlementError):
|
||||
raise
|
||||
verbose_logger.warning("Failed to get allowed MCP servers for agent: %s", e)
|
||||
return []
|
||||
|
|
@ -3390,7 +3509,7 @@ class MCPRequestHandler:
|
|||
return frozenset(global_mcp_server_manager.expand_permission_list(sorted(ceiling.mcp_server_ids)))
|
||||
|
||||
@staticmethod
|
||||
async def _get_agent_tool_permissions_for_server(
|
||||
async def get_agent_tool_permissions_for_server(
|
||||
server_id: str,
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
agent_object_permission: LiteLLM_ObjectPermissionTable | None = None,
|
||||
|
|
@ -3430,11 +3549,13 @@ class MCPRequestHandler:
|
|||
if obj_perm.mcp_tool_permissions
|
||||
else None
|
||||
)
|
||||
toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(obj_perm, server_id)
|
||||
toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
|
||||
obj_perm, server_id, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
|
||||
)
|
||||
agent_tools: Final = MCPRequestHandler._union_tool_grants(direct_tools, toolset_tools)
|
||||
return list(agent_tools) if agent_tools else None
|
||||
return list(agent_tools) if agent_tools is not None else None
|
||||
except Exception as e:
|
||||
if isinstance(e, UnloadableEntitlementError):
|
||||
if managed_agent_policy(user_api_key_auth) is not None or isinstance(e, UnloadableEntitlementError):
|
||||
raise
|
||||
verbose_logger.warning("Failed to get agent tool permissions for server: %s", e)
|
||||
return None
|
||||
|
|
@ -3452,28 +3573,38 @@ class MCPRequestHandler:
|
|||
return server_ids
|
||||
|
||||
@staticmethod
|
||||
async def _get_db_server_ids_for_access_groups(prisma_client, access_groups: list[str]) -> set[str]:
|
||||
async def _get_db_server_ids_for_access_groups(
|
||||
prisma_client,
|
||||
access_groups: list[str],
|
||||
*,
|
||||
use_writer: bool = False,
|
||||
) -> set[str]:
|
||||
"""
|
||||
Helper to get server_ids from DB servers that match any of the given access groups.
|
||||
"""
|
||||
server_ids: Final[set[str]] = set()
|
||||
if access_groups and prisma_client is not None:
|
||||
try:
|
||||
mcp_servers: Final = await MCPServerRepository(prisma_client).table.find_many(
|
||||
mcp_servers: Final = await MCPServerRepository(prisma_client, use_writer=use_writer).table.find_many(
|
||||
where={"mcp_access_groups": {"hasSome": access_groups}}
|
||||
)
|
||||
for server in mcp_servers:
|
||||
server_ids.add(server.server_id)
|
||||
except Exception as e:
|
||||
if use_writer:
|
||||
raise
|
||||
verbose_logger.debug("Error getting MCP servers from access groups: %s", e)
|
||||
return server_ids
|
||||
|
||||
@staticmethod
|
||||
async def _get_mcp_servers_from_access_groups(
|
||||
access_groups: list[str],
|
||||
*,
|
||||
requires_fresh_policy: bool = False,
|
||||
) -> list[str]:
|
||||
"""
|
||||
Resolve MCP access groups to server IDs by querying BOTH the MCP server table (DB) AND config-loaded servers
|
||||
Resolve MCP access groups to server IDs by querying BOTH the MCP server table (DB) AND config-loaded servers.
|
||||
``requires_fresh_policy`` reads the writer and propagates a read fault instead of resolving to no servers.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
|
|
@ -3489,11 +3620,15 @@ class MCPRequestHandler:
|
|||
)
|
||||
|
||||
# Use the new helper for DB servers
|
||||
db_server_ids = await MCPRequestHandler._get_db_server_ids_for_access_groups(prisma_client, access_groups)
|
||||
db_server_ids = await MCPRequestHandler._get_db_server_ids_for_access_groups(
|
||||
prisma_client, access_groups, use_writer=requires_fresh_policy
|
||||
)
|
||||
server_ids.update(db_server_ids)
|
||||
|
||||
return list(server_ids)
|
||||
except Exception as e:
|
||||
if requires_fresh_policy:
|
||||
raise
|
||||
verbose_logger.warning("Failed to get MCP servers from access groups: %s", e)
|
||||
return []
|
||||
|
||||
|
|
@ -3548,6 +3683,7 @@ class MCPRequestHandler:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
if key_object_permission is None:
|
||||
return []
|
||||
|
|
@ -3591,6 +3727,7 @@ class MCPRequestHandler:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
if team_obj is None:
|
||||
verbose_logger.debug("team_obj is None")
|
||||
|
|
|
|||
|
|
@ -181,6 +181,7 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
is_per_server_oauth_discovery_eligible,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_policy
|
||||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
|
||||
from litellm.proxy.common_utils.user_api_key_cache import get_management_object_ttl
|
||||
|
|
@ -3428,7 +3429,9 @@ class MCPServerManager:
|
|||
|
||||
``allow_all_server_ids`` / ``submitted_server_ids`` are injectable so the server union,
|
||||
which precomputes both for its fallback path, does not compute them twice."""
|
||||
if user_api_key_auth is not None and user_api_key_auth.mcp_toolset_id is not None:
|
||||
if user_api_key_auth is not None and (
|
||||
user_api_key_auth.mcp_toolset_id is not None or user_api_key_auth.mcp_explicit_grants_only
|
||||
):
|
||||
return set()
|
||||
if allow_all_server_ids is None:
|
||||
allow_all_server_ids = self.get_allow_all_keys_server_ids()
|
||||
|
|
@ -3477,9 +3480,14 @@ class MCPServerManager:
|
|||
2. If admin and no object_permission, return all servers
|
||||
3. Otherwise, use standard permission checks
|
||||
"""
|
||||
if managed_agent_policy(user_api_key_auth) is not None:
|
||||
managed: Final = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth)
|
||||
return managed if access is None else [server for server in managed if server in access.server_ids]
|
||||
|
||||
from litellm.proxy.proxy_server import general_settings as proxy_general_settings
|
||||
|
||||
resolved_general_settings: Final = proxy_general_settings if general_settings is None else general_settings
|
||||
explicit_grants_only: Final = bool(user_api_key_auth and user_api_key_auth.mcp_explicit_grants_only)
|
||||
allow_all_server_ids: Final = self.get_allow_all_keys_server_ids()
|
||||
|
||||
# A keyless admitted subject is resolved per grant source, and channel decisions that are
|
||||
|
|
@ -3511,7 +3519,7 @@ class MCPServerManager:
|
|||
# only keys without their own mcp_servers list get submitted servers unioned in.
|
||||
submitted_server_ids: Final = (
|
||||
[]
|
||||
if has_explicit_object_permission
|
||||
if has_explicit_object_permission or explicit_grants_only
|
||||
else await self._get_active_submitted_mcp_server_ids_for_user(user_api_key_auth)
|
||||
)
|
||||
|
||||
|
|
@ -3580,12 +3588,14 @@ class MCPServerManager:
|
|||
return [
|
||||
server_id
|
||||
for server_id in dict.fromkeys(allow_all_server_ids + submitted_server_ids)
|
||||
if scope is None or server_id == scope
|
||||
if not explicit_grants_only and (scope is None or server_id == scope)
|
||||
]
|
||||
|
||||
async def resolve_toolset_tool_permissions(
|
||||
self,
|
||||
toolset_ids: list[str],
|
||||
*,
|
||||
requires_fresh_policy: bool = False,
|
||||
) -> dict[str, list[str]]:
|
||||
"""
|
||||
Resolve a list of toolset IDs into a mcp_tool_permissions dict.
|
||||
|
|
@ -3595,6 +3605,10 @@ class MCPServerManager:
|
|||
Redis-backed ``DualCache`` in production) so that cache entries are
|
||||
shared across workers and cold-cache DB hits are minimised.
|
||||
|
||||
``requires_fresh_policy`` bypasses the cache and reads the writer so a
|
||||
revocation is honoured on the very next request; a read fault then
|
||||
propagates instead of resolving to no grants.
|
||||
|
||||
A row names a tool on the server identified by ``server_id``, so the
|
||||
stored name is the tool's own name and is used as written. It is never
|
||||
reduced by the server's wire prefix: that prefix is added on the way out
|
||||
|
|
@ -3609,12 +3623,16 @@ class MCPServerManager:
|
|||
return {}
|
||||
|
||||
cache_key: Final = "toolset_perms:" + ",".join(sorted(toolset_ids))
|
||||
cached: Final[dict[str, list[str]] | None] = await user_api_key_cache.async_get_cache(key=cache_key)
|
||||
cached: Final[dict[str, list[str]] | None] = (
|
||||
None if requires_fresh_policy else await user_api_key_cache.async_get_cache(key=cache_key)
|
||||
)
|
||||
if cached is not None:
|
||||
return cached
|
||||
|
||||
try:
|
||||
toolsets: Final = await list_mcp_toolsets(prisma_client, toolset_ids=toolset_ids)
|
||||
toolsets: Final = await list_mcp_toolsets(
|
||||
prisma_client, toolset_ids=toolset_ids, use_writer=requires_fresh_policy
|
||||
)
|
||||
tool_permissions: Final[dict[str, list[str]]] = {}
|
||||
for toolset in toolsets:
|
||||
for tool in toolset.tools:
|
||||
|
|
@ -3628,6 +3646,8 @@ class MCPServerManager:
|
|||
)
|
||||
return tool_permissions
|
||||
except Exception as e:
|
||||
if requires_fresh_policy:
|
||||
raise
|
||||
verbose_logger.warning("Failed to resolve toolset permissions: %s", e)
|
||||
return {}
|
||||
|
||||
|
|
|
|||
|
|
@ -65,9 +65,9 @@ class MCPToolsetTable(Protocol):
|
|||
async def delete(self, where: Mapping[str, object]) -> MCPToolsetRow: ...
|
||||
|
||||
|
||||
def _toolset_table(prisma_client: PrismaClient) -> MCPToolsetTable:
|
||||
def _toolset_table(prisma_client: PrismaClient, *, use_writer: bool = False) -> MCPToolsetTable:
|
||||
"""The toolset table actions of the prisma client."""
|
||||
return MCPToolsetRepository(prisma_client).table
|
||||
return MCPToolsetRepository(prisma_client, use_writer=use_writer).table
|
||||
|
||||
|
||||
def _toolset_from_row(row: MCPToolsetRow) -> MCPToolset:
|
||||
|
|
@ -107,12 +107,16 @@ async def get_mcp_toolset(
|
|||
async def list_mcp_toolsets(
|
||||
prisma_client: PrismaClient,
|
||||
toolset_ids: Sequence[str] | None = None,
|
||||
*,
|
||||
use_writer: bool = False,
|
||||
) -> Sequence[MCPToolset]:
|
||||
try:
|
||||
where: Final[Mapping[str, object]] = {} if toolset_ids is None else {"toolset_id": {"in": toolset_ids}}
|
||||
rows: Final = await _toolset_table(prisma_client).find_many(where=where)
|
||||
rows: Final = await _toolset_table(prisma_client, use_writer=use_writer).find_many(where=where)
|
||||
return [_toolset_from_row(r) for r in rows]
|
||||
except Exception as e:
|
||||
if use_writer:
|
||||
raise
|
||||
verbose_proxy_logger.warning("litellm.proxy._experimental.mcp_server.toolset_db::list_mcp_toolsets - %s", e)
|
||||
return []
|
||||
|
||||
|
|
|
|||
|
|
@ -58,6 +58,7 @@ async def resolve_ui_session_team_ids(
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
check_db_only=user_api_key_auth.requires_fresh_policy,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
|
@ -92,7 +93,9 @@ async def admitted_user_context(user_api_key_auth: UserAPIKeyAuth) -> UserAPIKey
|
|||
)
|
||||
|
||||
try:
|
||||
admitted: Final = await MCPRequestHandler.reload_admitted_user(user_id)
|
||||
admitted: Final = await MCPRequestHandler.reload_admitted_user(
|
||||
user_id, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
|
||||
)
|
||||
except HTTPException as e:
|
||||
verbose_logger.warning("MCP dashboard session: admitted-subject reload failed for %s: %s", user_id, e.detail)
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -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": [
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -1,13 +1,16 @@
|
|||
import asyncio
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Final, TypeAlias
|
||||
from typing import TYPE_CHECKING, Final, TypeAlias
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import LiteLLM_AccessGroupTable
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
AccessGroupIds: TypeAlias = tuple[str, ...]
|
||||
AccessGroupIdsLoader: TypeAlias = Callable[[str], Awaitable[AccessGroupIds]] # mutable-ok: Callable params
|
||||
LoadedAccessGroup: TypeAlias = LiteLLM_AccessGroupTable | None
|
||||
|
|
@ -34,7 +37,7 @@ async def _registry_access_group_ids(agent_id: str) -> AccessGroupIds:
|
|||
return tuple(agent.access_group_ids or ()) if agent is not None else ()
|
||||
|
||||
|
||||
async def _load_access_group(access_group_id: str) -> LoadedAccessGroup:
|
||||
async def _load_access_group(access_group_id: str, *, check_db_only: bool = False) -> LoadedAccessGroup:
|
||||
from litellm.proxy.auth.auth_checks import get_access_object
|
||||
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
|
||||
|
||||
|
|
@ -47,8 +50,11 @@ async def _load_access_group(access_group_id: str) -> LoadedAccessGroup:
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=check_db_only,
|
||||
)
|
||||
except HTTPException as e:
|
||||
if check_db_only:
|
||||
raise
|
||||
verbose_proxy_logger.warning(
|
||||
"Agent access group %s could not be loaded, treating it as empty: %s", access_group_id, e.detail
|
||||
)
|
||||
|
|
@ -59,13 +65,20 @@ async def resolve_agent_access_group_ceiling(
|
|||
agent_id: str,
|
||||
load_access_group_ids: AccessGroupIdsLoader = _registry_access_group_ids,
|
||||
load_access_group: AccessGroupLoader = _load_access_group,
|
||||
*,
|
||||
check_db_only: bool = False,
|
||||
) -> AgentAccessGroupCeiling | None:
|
||||
"""``None`` when the agent has no access groups attached, so nothing is capped."""
|
||||
access_group_ids: Final = await load_access_group_ids(agent_id)
|
||||
if not access_group_ids:
|
||||
return None
|
||||
|
||||
loaded: Final = await asyncio.gather(*(load_access_group(group_id) for group_id in access_group_ids))
|
||||
loaded: Final = await asyncio.gather(
|
||||
*(
|
||||
_load_access_group(group_id, check_db_only=True) if check_db_only else load_access_group(group_id)
|
||||
for group_id in access_group_ids
|
||||
)
|
||||
)
|
||||
groups: Final = tuple(group for group in loaded if group is not None)
|
||||
return AgentAccessGroupCeiling(
|
||||
access_group_ids=access_group_ids,
|
||||
|
|
@ -73,3 +86,16 @@ async def resolve_agent_access_group_ceiling(
|
|||
mcp_server_ids=frozenset(server_id for group in groups for server_id in group.access_mcp_server_ids),
|
||||
agent_ids=frozenset(target_id for group in groups for target_id in group.access_agent_ids),
|
||||
)
|
||||
|
||||
|
||||
async def resolve_managed_agent_ceilings(agent: "AgentResponse") -> tuple[AgentAccessGroupCeiling, ...]:
|
||||
async def authoritative_group(group_id: str) -> LoadedAccessGroup:
|
||||
return await _load_access_group(group_id, check_db_only=True)
|
||||
|
||||
async def manual_ids(_agent_id: str) -> AccessGroupIds:
|
||||
return tuple(agent.access_group_ids or ())
|
||||
|
||||
manual: Final = await resolve_agent_access_group_ceiling(
|
||||
agent.agent_id, load_access_group_ids=manual_ids, load_access_group=authoritative_group
|
||||
)
|
||||
return (manual,) if manual is not None else ()
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ can only narrow access and need no trust.
|
|||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -45,7 +46,7 @@ def agent_caller_auth(user_api_key_auth: UserAPIKeyAuth) -> UserAPIKeyAuth | Non
|
|||
user_id=caller.user_id,
|
||||
team_id=caller.team_id,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
)
|
||||
).model_copy(update=MappingProxyType({"requires_fresh_policy": user_api_key_auth.requires_fresh_policy}))
|
||||
|
||||
|
||||
async def load_agent_caller_team(user_api_key_auth: UserAPIKeyAuth) -> LiteLLM_TeamTable | None:
|
||||
|
|
|
|||
|
|
@ -8,8 +8,11 @@ Follows the same pattern as MCP permission handling.
|
|||
import asyncio
|
||||
from collections.abc import Awaitable, Callable, Sequence
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final, TypeAlias
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy._experimental.mcp_server.ui_session_utils import build_effective_auth_contexts
|
||||
from litellm.proxy._types import (
|
||||
|
|
@ -24,6 +27,7 @@ from litellm.proxy.agent_endpoints.auth.agent_access_groups import (
|
|||
resolve_agent_access_group_ceiling,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.auth.agent_caller import agent_caller_auth
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_policy
|
||||
from litellm.repositories.table_repositories import AgentsRepository
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
|
|
@ -83,13 +87,23 @@ class AgentRequestHandler:
|
|||
async def resolve_agent_access(
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
resolve_ceiling: CeilingResolver = resolve_agent_access_group_ceiling,
|
||||
*,
|
||||
strict: bool = False,
|
||||
) -> AgentAccess:
|
||||
"""Agents the key may reach: key and team grants, intersected with the agent's access group ceiling
|
||||
and, for an agent key acting on behalf of an invoking user, with that user's team grants."""
|
||||
key_team_access: Final = await AgentRequestHandler._resolve_key_team_agent_access(user_api_key_auth)
|
||||
caller_access: Final = await AgentRequestHandler._agent_caller_access(user_api_key_auth)
|
||||
if managed_agent_policy(user_api_key_auth) is not None:
|
||||
return await _managed_actor_agent_access(user_api_key_auth)
|
||||
key_team_access: Final = await AgentRequestHandler.resolve_key_team_agent_access(
|
||||
user_api_key_auth, strict=strict
|
||||
)
|
||||
if strict and isinstance(key_team_access, UnrestrictedAgentAccess):
|
||||
return RestrictedAgentAccess(frozenset())
|
||||
caller_access: Final = await AgentRequestHandler.agent_caller_access(user_api_key_auth, strict=strict)
|
||||
own_access: Final = _intersect_agent_access(key_team_access, caller_access)
|
||||
agent_ceiling: Final = await AgentRequestHandler._agent_access_group_ceiling(user_api_key_auth, resolve_ceiling)
|
||||
agent_ceiling: Final = await AgentRequestHandler._agent_access_group_ceiling(
|
||||
user_api_key_auth, resolve_ceiling, strict=strict
|
||||
)
|
||||
if agent_ceiling is None:
|
||||
return own_access
|
||||
if isinstance(own_access, UnrestrictedAgentAccess):
|
||||
|
|
@ -97,20 +111,26 @@ class AgentRequestHandler:
|
|||
return RestrictedAgentAccess(own_access.agent_ids & agent_ceiling)
|
||||
|
||||
@staticmethod
|
||||
async def _agent_caller_access(user_api_key_auth: UserAPIKeyAuth | None) -> AgentAccess:
|
||||
async def agent_caller_access(user_api_key_auth: UserAPIKeyAuth | None, *, strict: bool = False) -> AgentAccess:
|
||||
caller_auth: Final = agent_caller_auth(user_api_key_auth) if user_api_key_auth else None
|
||||
if caller_auth is None:
|
||||
return UnrestrictedAgentAccess()
|
||||
return await AgentRequestHandler._get_allowed_agents_for_team(caller_auth)
|
||||
return await AgentRequestHandler._get_allowed_agents_for_team(caller_auth, strict=strict)
|
||||
|
||||
@staticmethod
|
||||
async def _resolve_key_team_agent_access(
|
||||
async def resolve_key_team_agent_access(
|
||||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
*,
|
||||
strict: bool = False,
|
||||
) -> AgentAccess:
|
||||
try:
|
||||
key_access: Final = await AgentRequestHandler._get_allowed_agents_for_key(user_api_key_auth)
|
||||
team_access: Final = await AgentRequestHandler._get_allowed_agents_for_team(user_api_key_auth)
|
||||
key_access: Final = await AgentRequestHandler.get_allowed_agents_for_key(user_api_key_auth, strict=strict)
|
||||
team_access: Final = await AgentRequestHandler._get_allowed_agents_for_team(
|
||||
user_api_key_auth, strict=strict
|
||||
)
|
||||
except Exception as e:
|
||||
if strict:
|
||||
raise HTTPException(503, "Agent invocation policy is unavailable") from e
|
||||
verbose_logger.warning("Failed to get allowed agents: %s", e)
|
||||
return UnrestrictedAgentAccess()
|
||||
return _intersect_agent_access(key_access, team_access)
|
||||
|
|
@ -119,10 +139,16 @@ class AgentRequestHandler:
|
|||
async def _agent_access_group_ceiling(
|
||||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
resolve_ceiling: CeilingResolver,
|
||||
*,
|
||||
strict: bool = False,
|
||||
) -> frozenset[str] | None:
|
||||
if user_api_key_auth is None or not user_api_key_auth.agent_id:
|
||||
return None
|
||||
ceiling: Final = await resolve_ceiling(user_api_key_auth.agent_id)
|
||||
ceiling: Final = (
|
||||
await resolve_agent_access_group_ceiling(user_api_key_auth.agent_id, check_db_only=True)
|
||||
if strict
|
||||
else await resolve_ceiling(user_api_key_auth.agent_id)
|
||||
)
|
||||
if ceiling is None:
|
||||
return None
|
||||
return _to_stable_ids(ceiling.agent_ids)
|
||||
|
|
@ -144,6 +170,49 @@ class AgentRequestHandler:
|
|||
bool: True if agent is allowed, False otherwise
|
||||
"""
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
|
||||
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityFailure
|
||||
|
||||
registered: Final = global_agent_registry.get_agent_by_id(agent_id)
|
||||
registry_managed: Final = isinstance(registered, AgentResponse) and registered.identity_managed
|
||||
if registry_managed or (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))
|
||||
|
|
|
|||
84
litellm/proxy/agent_endpoints/auth/managed_authorization.py
Normal file
84
litellm/proxy/agent_endpoints/auth/managed_authorization.py
Normal 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
|
||||
17
litellm/proxy/agent_endpoints/identity.py
Normal file
17
litellm/proxy/agent_endpoints/identity.py
Normal 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)
|
||||
252
litellm/proxy/agent_endpoints/identity_store.py
Normal file
252
litellm/proxy/agent_endpoints/identity_store.py
Normal 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
|
||||
220
litellm/proxy/agent_endpoints/managed_identity.py
Normal file
220
litellm/proxy/agent_endpoints/managed_identity.py
Normal 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")
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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``.
|
||||
|
|
|
|||
|
|
@ -28,12 +28,15 @@ from litellm.integrations.custom_guardrail import (
|
|||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
effective_scan_only_tool_results_for_guardrail,
|
||||
effective_skip_system_message_for_guardrail,
|
||||
role_out_of_guardrail_scope,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.llms.openai.responses.guardrail_translation.handler import scannable_instructions
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
add_guardrail_scan_id,
|
||||
|
|
@ -105,9 +108,14 @@ class _ResponsesInputItem(BaseModel):
|
|||
model_config = ConfigDict(extra="ignore")
|
||||
|
||||
type: str | None = None
|
||||
role: str | None = None
|
||||
content: str | tuple[_ResponsesContentPart, ...] | None = None
|
||||
|
||||
def text_count(self) -> int:
|
||||
def text_count(self, *, skip_system: bool) -> int:
|
||||
if role_out_of_guardrail_scope(
|
||||
(self.role or "").lower(), skip_system_message=skip_system, skip_tool_message=False
|
||||
):
|
||||
return 0
|
||||
if isinstance(self.content, str):
|
||||
return 1
|
||||
if self.content is None:
|
||||
|
|
@ -1636,10 +1644,10 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
|
||||
A message's texts are consumed only when they sit at the running position of
|
||||
``texts``; messages the translation handler added without a counterpart in
|
||||
``texts`` (Responses ``instructions``, ``function_call_output``, ``reasoning``)
|
||||
are skipped. The walk runs front-to-back and back-to-front and both must agree,
|
||||
so an added message whose text happens to equal a neighbouring real message's
|
||||
text cannot steal that text's attribution. Returns None otherwise.
|
||||
``texts`` (Responses ``function_call_output``, ``reasoning``) are skipped. The walk
|
||||
runs front-to-back and back-to-front and both must agree, so an added message whose
|
||||
text happens to equal a neighbouring real message's text cannot steal that text's
|
||||
attribution. Returns None otherwise.
|
||||
"""
|
||||
runs: Final = tuple(cls._message_texts(message) for message in messages)
|
||||
|
||||
|
|
@ -1660,17 +1668,19 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
)
|
||||
return forward if len(forward) == len(texts) and forward == backward else None
|
||||
|
||||
@classmethod
|
||||
@staticmethod
|
||||
def _reasoning_item_text_indices(
|
||||
cls,
|
||||
texts: Sequence[str],
|
||||
request_data: Mapping[str, object],
|
||||
*,
|
||||
skip_system: bool,
|
||||
) -> frozenset[int] | None:
|
||||
"""Return the ``texts`` indices flattened from Responses ``reasoning`` input items.
|
||||
|
||||
The Responses translation handler gives those model-authored items the default
|
||||
``user`` role, so the latest-turn selection must not mistake one for a human turn.
|
||||
Empty for requests without a Responses ``input`` item list; None when the raw items
|
||||
(after the leading ``instructions`` text, both minus whatever ``skip_system`` drops)
|
||||
do not account for every entry of ``texts``.
|
||||
"""
|
||||
try:
|
||||
|
|
@ -1679,10 +1689,11 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
return None
|
||||
if not isinstance(raw_input, tuple):
|
||||
return frozenset()
|
||||
counts: Final = tuple(item.text_count() for item in raw_input)
|
||||
if sum(counts) != len(texts):
|
||||
offset: Final = 0 if scannable_instructions(request_data, skip_system=skip_system) is None else 1
|
||||
counts: Final = tuple(item.text_count(skip_system=skip_system) for item in raw_input)
|
||||
if offset + sum(counts) != len(texts):
|
||||
return None
|
||||
starts: Final = itertools.accumulate(counts, initial=0)
|
||||
starts: Final = itertools.accumulate(counts, initial=offset)
|
||||
return frozenset(
|
||||
text_idx
|
||||
for item, count, start in zip(raw_input, counts, starts)
|
||||
|
|
@ -1690,9 +1701,8 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
for text_idx in range(start, start + count)
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _get_latest_user_text_indices(
|
||||
cls,
|
||||
self,
|
||||
texts: Sequence[str],
|
||||
messages: Sequence[AllMessageValues],
|
||||
request_data: Mapping[str, object],
|
||||
|
|
@ -1706,10 +1716,12 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
user/developer message exists, or the latest one carries text that never reached
|
||||
``texts`` (safety fallback to the role-filter scan).
|
||||
"""
|
||||
sources: Final = cls._text_source_message_indices(texts, messages)
|
||||
sources: Final = self._text_source_message_indices(texts, messages)
|
||||
if sources is None:
|
||||
return None
|
||||
reasoning: Final = cls._reasoning_item_text_indices(texts, request_data)
|
||||
reasoning: Final = self._reasoning_item_text_indices(
|
||||
texts, request_data, skip_system=effective_skip_system_message_for_guardrail(self)
|
||||
)
|
||||
if reasoning is None:
|
||||
return None
|
||||
reasoning_messages: Final = frozenset(sources[text_idx] for text_idx in reasoning)
|
||||
|
|
@ -1723,7 +1735,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
)
|
||||
if latest_human is None:
|
||||
return None
|
||||
if latest_human not in sources and cls._message_texts(messages[latest_human]):
|
||||
if latest_human not in sources and self._message_texts(messages[latest_human]):
|
||||
return None
|
||||
return frozenset(text_idx for text_idx, source in enumerate(sources) if source == latest_human)
|
||||
|
||||
|
|
|
|||
|
|
@ -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({})
|
||||
|
|
|
|||
|
|
@ -2290,7 +2290,7 @@ def _upstream_close_to_relay(task_results: Iterable[object]) -> Close | None:
|
|||
return upstream_close
|
||||
|
||||
|
||||
_WEBSOCKET_FORWARDED_HEADERS: Final = frozenset(("authorization", "x-api-key", "x-goog-user-project"))
|
||||
_WEBSOCKET_FORWARDED_HEADERS: Final = frozenset(("x-goog-user-project",))
|
||||
|
||||
|
||||
def _with_trace_context(headers: Mapping[str, str], parent_span: object) -> dict[str, str]:
|
||||
|
|
|
|||
|
|
@ -143,6 +143,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",
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -4194,6 +4194,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.
|
||||
|
|
@ -4205,7 +4207,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:
|
||||
|
|
@ -4873,6 +4875,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()
|
||||
|
|
@ -5171,12 +5174,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
|
||||
|
|
@ -5188,6 +5199,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")
|
||||
|
|
@ -5439,9 +5451,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 = {}
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"]:
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -20,6 +20,19 @@ class RustUpstreamError(Exception): ...
|
|||
class ForkedAfterNativeRuntimeStarted(RuntimeError): ...
|
||||
class ProcessReservedForForking(RuntimeError): ...
|
||||
|
||||
def trace_encode_rows(rows: Sequence[Mapping[str, JsonValue]]) -> str: ...
|
||||
def trace_ensure_schema(
|
||||
url: str, database: str, user: str, password: str, trace_retention_days: int, spend_log_retention_days: int
|
||||
) -> Future[None]: ...
|
||||
def trace_query(
|
||||
url: str,
|
||||
database: str,
|
||||
user: str,
|
||||
password: str,
|
||||
sql: str,
|
||||
parameters: Mapping[str, str | int | Sequence[str]],
|
||||
) -> Future[str]: ...
|
||||
|
||||
@final
|
||||
class NativeDiagnosticProcessor:
|
||||
def __new__(cls, minimum_custom_key_length: int) -> NativeDiagnosticProcessor: ...
|
||||
|
|
@ -338,6 +351,9 @@ __all__ = [
|
|||
"process_state_started",
|
||||
"reserve_process_for_forking",
|
||||
"responses",
|
||||
"trace_encode_rows",
|
||||
"trace_ensure_schema",
|
||||
"trace_query",
|
||||
"transcription",
|
||||
]
|
||||
|
||||
|
|
|
|||
69
litellm/rust_bridge/traces.py
Normal file
69
litellm/rust_bridge/traces.py
Normal file
|
|
@ -0,0 +1,69 @@
|
|||
from collections.abc import Awaitable, Mapping, Sequence
|
||||
from typing import Final, Protocol, cast
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, JsonValue
|
||||
|
||||
from litellm.rust_bridge.loader import get_native_bridge
|
||||
|
||||
|
||||
class NativeTraces(Protocol):
|
||||
def trace_encode_rows(self, rows: Sequence[Mapping[str, JsonValue]]) -> str: ...
|
||||
|
||||
def trace_ensure_schema(
|
||||
self,
|
||||
url: str,
|
||||
database: str,
|
||||
user: str,
|
||||
password: str,
|
||||
trace_retention_days: int,
|
||||
spend_log_retention_days: int,
|
||||
) -> Awaitable[None]: ...
|
||||
|
||||
def trace_query(
|
||||
self,
|
||||
url: str,
|
||||
database: str,
|
||||
user: str,
|
||||
password: str,
|
||||
sql: str,
|
||||
parameters: Mapping[str, str | int | Sequence[str]],
|
||||
) -> Awaitable[str]: ...
|
||||
|
||||
|
||||
class QueryResponse(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
data: list[dict[str, JsonValue]]
|
||||
|
||||
|
||||
def _native() -> NativeTraces:
|
||||
native: Final = get_native_bridge()
|
||||
if native is None:
|
||||
raise RuntimeError("Agent tracing requires the Rust extension")
|
||||
return cast(NativeTraces, native) # cast-ok: the native extension is validated against this protocol at call sites
|
||||
|
||||
|
||||
async def ensure_schema(
|
||||
url: str,
|
||||
database: str,
|
||||
user: str,
|
||||
password: str,
|
||||
trace_retention_days: int,
|
||||
spend_log_retention_days: int,
|
||||
) -> None:
|
||||
await _native().trace_ensure_schema(url, database, user, password, trace_retention_days, spend_log_retention_days)
|
||||
|
||||
|
||||
async def query(
|
||||
url: str,
|
||||
database: str,
|
||||
user: str,
|
||||
password: str,
|
||||
sql: str,
|
||||
parameters: Mapping[str, str | int | Sequence[str]],
|
||||
) -> list[dict[str, JsonValue]]:
|
||||
result: Final = await _native().trace_query(url, database, user, password, sql, parameters)
|
||||
return QueryResponse.model_validate_json(result).data
|
||||
|
||||
|
||||
def encode_rows(rows: Sequence[Mapping[str, JsonValue]]) -> bytes:
|
||||
return _native().trace_encode_rows(rows).encode("utf-8")
|
||||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue