mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
chore(mcp): merge main into listed-tool metadata branch
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
commit
c8870f0080
191 changed files with 9554 additions and 679 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==",
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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::*;
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -3494,7 +3495,9 @@ class MCPServerManager:
|
|||
|
||||
``allow_all_server_ids`` / ``submitted_server_ids`` are injectable so the server union,
|
||||
which precomputes both for its fallback path, does not compute them twice."""
|
||||
if user_api_key_auth is not None and user_api_key_auth.mcp_toolset_id is not None:
|
||||
if user_api_key_auth is not None and (
|
||||
user_api_key_auth.mcp_toolset_id is not None or user_api_key_auth.mcp_explicit_grants_only
|
||||
):
|
||||
return set()
|
||||
if allow_all_server_ids is None:
|
||||
allow_all_server_ids = self.get_allow_all_keys_server_ids()
|
||||
|
|
@ -3543,9 +3546,14 @@ class MCPServerManager:
|
|||
2. If admin and no object_permission, return all servers
|
||||
3. Otherwise, use standard permission checks
|
||||
"""
|
||||
if managed_agent_policy(user_api_key_auth) is not None:
|
||||
managed: Final = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth)
|
||||
return managed if access is None else [server for server in managed if server in access.server_ids]
|
||||
|
||||
from litellm.proxy.proxy_server import general_settings as proxy_general_settings
|
||||
|
||||
resolved_general_settings: Final = proxy_general_settings if general_settings is None else general_settings
|
||||
explicit_grants_only: Final = bool(user_api_key_auth and user_api_key_auth.mcp_explicit_grants_only)
|
||||
allow_all_server_ids: Final = self.get_allow_all_keys_server_ids()
|
||||
|
||||
# A keyless admitted subject is resolved per grant source, and channel decisions that are
|
||||
|
|
@ -3577,7 +3585,7 @@ class MCPServerManager:
|
|||
# only keys without their own mcp_servers list get submitted servers unioned in.
|
||||
submitted_server_ids: Final = (
|
||||
[]
|
||||
if has_explicit_object_permission
|
||||
if has_explicit_object_permission or explicit_grants_only
|
||||
else await self._get_active_submitted_mcp_server_ids_for_user(user_api_key_auth)
|
||||
)
|
||||
|
||||
|
|
@ -3646,12 +3654,14 @@ class MCPServerManager:
|
|||
return [
|
||||
server_id
|
||||
for server_id in dict.fromkeys(allow_all_server_ids + submitted_server_ids)
|
||||
if scope is None or server_id == scope
|
||||
if not explicit_grants_only and (scope is None or server_id == scope)
|
||||
]
|
||||
|
||||
async def resolve_toolset_tool_permissions(
|
||||
self,
|
||||
toolset_ids: list[str],
|
||||
*,
|
||||
requires_fresh_policy: bool = False,
|
||||
) -> dict[str, list[str]]:
|
||||
"""
|
||||
Resolve a list of toolset IDs into a mcp_tool_permissions dict.
|
||||
|
|
@ -3661,6 +3671,10 @@ class MCPServerManager:
|
|||
Redis-backed ``DualCache`` in production) so that cache entries are
|
||||
shared across workers and cold-cache DB hits are minimised.
|
||||
|
||||
``requires_fresh_policy`` bypasses the cache and reads the writer so a
|
||||
revocation is honoured on the very next request; a read fault then
|
||||
propagates instead of resolving to no grants.
|
||||
|
||||
A row names a tool on the server identified by ``server_id``, so the
|
||||
stored name is the tool's own name and is used as written. It is never
|
||||
reduced by the server's wire prefix: that prefix is added on the way out
|
||||
|
|
@ -3675,12 +3689,16 @@ class MCPServerManager:
|
|||
return {}
|
||||
|
||||
cache_key: Final = "toolset_perms:" + ",".join(sorted(toolset_ids))
|
||||
cached: Final[dict[str, list[str]] | None] = await user_api_key_cache.async_get_cache(key=cache_key)
|
||||
cached: Final[dict[str, list[str]] | None] = (
|
||||
None if requires_fresh_policy else await user_api_key_cache.async_get_cache(key=cache_key)
|
||||
)
|
||||
if cached is not None:
|
||||
return cached
|
||||
|
||||
try:
|
||||
toolsets: Final = await list_mcp_toolsets(prisma_client, toolset_ids=toolset_ids)
|
||||
toolsets: Final = await list_mcp_toolsets(
|
||||
prisma_client, toolset_ids=toolset_ids, use_writer=requires_fresh_policy
|
||||
)
|
||||
tool_permissions: Final[dict[str, list[str]]] = {}
|
||||
for toolset in toolsets:
|
||||
for tool in toolset.tools:
|
||||
|
|
@ -3694,6 +3712,8 @@ class MCPServerManager:
|
|||
)
|
||||
return tool_permissions
|
||||
except Exception as e:
|
||||
if requires_fresh_policy:
|
||||
raise
|
||||
verbose_logger.warning("Failed to resolve toolset permissions: %s", e)
|
||||
return {}
|
||||
|
||||
|
|
|
|||
|
|
@ -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``.
|
||||
|
|
|
|||
|
|
@ -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({})
|
||||
|
|
|
|||
|
|
@ -142,6 +142,7 @@ ROUTE_ENDPOINT_MAPPING: Final = {
|
|||
"acancel_run": "/evals/{eval_id}/runs/{run_id}/cancel",
|
||||
"adelete_run": "/evals/{eval_id}/runs/{run_id}",
|
||||
"acreate_batch": "/batches",
|
||||
"aretrieve_batch": "/batches",
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -4200,6 +4200,8 @@ _PRISMA_DEFAULT_TX_TIMEOUT: Final = timedelta(seconds=5)
|
|||
async def _lookup_deprecated_key(
|
||||
db: PrismaWrapper | RoutingPrismaWrapper,
|
||||
hashed_token: str,
|
||||
*,
|
||||
check_db_only: bool = False,
|
||||
) -> str | None:
|
||||
"""
|
||||
Check if a token exists in the deprecated keys table and is still within its grace period.
|
||||
|
|
@ -4211,7 +4213,7 @@ async def _lookup_deprecated_key(
|
|||
now_ts: Final = now.timestamp()
|
||||
|
||||
# Check cache first
|
||||
cached: Final = _deprecated_key_cache.get(hashed_token)
|
||||
cached: Final = None if check_db_only else _deprecated_key_cache.get(hashed_token)
|
||||
if cached is not None:
|
||||
active_token_id, cache_expires_at_ts, revoke_at_ts = cached
|
||||
if now_ts < cache_expires_at_ts and now_ts < revoke_at_ts:
|
||||
|
|
@ -4879,6 +4881,7 @@ class PrismaClient:
|
|||
proxy_logging_obj: ProxyLogging | None = None,
|
||||
budget_id_list: list[str] | None = None,
|
||||
check_deprecated: bool = True,
|
||||
use_writer: bool = False,
|
||||
):
|
||||
args_passed_in: Final = locals()
|
||||
start_time: Final = time.time()
|
||||
|
|
@ -5177,12 +5180,20 @@ class PrismaClient:
|
|||
WHERE v.token = $1
|
||||
"""
|
||||
|
||||
response = await self._query_first_with_cached_plan_fallback(sql_query, hashed_token)
|
||||
response = (
|
||||
await self.writer_db.query_first(sql_query, hashed_token)
|
||||
if use_writer
|
||||
else await self._query_first_with_cached_plan_fallback(sql_query, hashed_token)
|
||||
)
|
||||
|
||||
# If not found in main table, check deprecated keys (grace period)
|
||||
# check_deprecated=False on the recursive call prevents unbounded chaining
|
||||
if response is None and hashed_token is not None and check_deprecated:
|
||||
active_token_id: Final = await _lookup_deprecated_key(db=self.db, hashed_token=hashed_token)
|
||||
active_token_id: Final = await _lookup_deprecated_key(
|
||||
db=self.writer_db if use_writer else self.db,
|
||||
hashed_token=hashed_token,
|
||||
check_db_only=use_writer,
|
||||
)
|
||||
if active_token_id:
|
||||
# The recursive call returns a finished
|
||||
# LiteLLM_VerificationTokenView; the dict
|
||||
|
|
@ -5194,6 +5205,7 @@ class PrismaClient:
|
|||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_deprecated=False,
|
||||
use_writer=use_writer,
|
||||
)
|
||||
if deprecated_response is not None:
|
||||
verbose_proxy_logger.debug("Deprecated key used during grace period")
|
||||
|
|
@ -5445,9 +5457,11 @@ class PrismaClient:
|
|||
# check if plain text or hash
|
||||
token = _hash_token_if_needed(token=token)
|
||||
db_data["token"] = token
|
||||
include_object_permission: Final[LiteLLM_VerificationTokenInclude] = {"object_permission": True}
|
||||
response: Final = await VerificationTokenRepository(self).table.update(
|
||||
where={"token": token},
|
||||
data=with_settings_updated_at(db_data),
|
||||
include=include_object_permission,
|
||||
)
|
||||
verbose_proxy_logger.debug("\033[91m" + f"DB Token Table update succeeded {response}" + "\033[0m")
|
||||
_data: dict = {}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -7,6 +7,10 @@ from pydantic import BaseModel, ConfigDict, Field, PrivateAttr, StrictInt, field
|
|||
from typing_extensions import ReadOnly, Required, TypedDict
|
||||
|
||||
from litellm.types.llms.base import LiteLLMPydanticObjectBase
|
||||
from litellm.types.proxy.agent_identity import (
|
||||
AgentExecutionMode,
|
||||
AgentIdentityBinding,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from a2a.types import SendMessageResponse
|
||||
|
|
@ -301,6 +305,11 @@ class AgentKeySummary(BaseModel):
|
|||
|
||||
|
||||
class AgentResponse(BaseModel):
|
||||
identity: AgentIdentityBinding | None = None
|
||||
identity_managed: bool = False
|
||||
enabled: bool = True
|
||||
execution_mode: AgentExecutionMode = "autonomous"
|
||||
jwt_auth_configured: bool = False
|
||||
agent_id: str
|
||||
agent_name: str
|
||||
litellm_params: dict[str, object] | None = None
|
||||
|
|
|
|||
|
|
@ -199,6 +199,7 @@ class ObservabilityOptions:
|
|||
logger_fn: Callable[[Mapping[str, object]], None] | None = None
|
||||
verbose: bool | None = None
|
||||
no_log: bool | None = field(default=None, metadata=wire("no-log"))
|
||||
log_client_error_tracebacks: bool | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True, kw_only=True)
|
||||
|
|
|
|||
|
|
@ -774,6 +774,8 @@ ANTHROPIC_TOOL_SEARCH_TOOL_TYPES: Final = frozenset(
|
|||
# Effort beta header constant
|
||||
ANTHROPIC_EFFORT_BETA_HEADER: Final = "effort-2025-11-24"
|
||||
|
||||
ANTHROPIC_MID_CONVERSATION_OUTPUT_CONFIG_BETA_HEADER: Final = "mid-conversation-output-config-2026-07-01"
|
||||
|
||||
ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER: Final = "fine-grained-tool-streaming-2025-05-14"
|
||||
|
||||
# OAuth constants
|
||||
|
|
|
|||
96
litellm/types/proxy/agent_identity.py
Normal file
96
litellm/types/proxy/agent_identity.py
Normal file
|
|
@ -0,0 +1,96 @@
|
|||
from datetime import datetime
|
||||
from typing import Literal, TypeAlias
|
||||
from uuid import UUID
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator
|
||||
|
||||
AgentExecutionMode: TypeAlias = Literal["autonomous", "delegated", "both"]
|
||||
|
||||
|
||||
class EntraIdentityConfig(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="forbid")
|
||||
|
||||
provider: Literal["microsoft_entra"]
|
||||
tenant_id: str
|
||||
client_id: str
|
||||
service_principal_id: str | None = None
|
||||
required_roles: tuple[str, ...] = ()
|
||||
required_scopes: tuple[str, ...] = Field(
|
||||
default=("user_impersonation",),
|
||||
description="Required delegated scopes. An empty list accepts any nonempty scope granted for this gateway.",
|
||||
)
|
||||
|
||||
@field_validator("tenant_id", "client_id", "service_principal_id")
|
||||
@classmethod
|
||||
def normalize_identifier(cls, value: str | None) -> str | None:
|
||||
return str(UUID(value)) if value is not None else None
|
||||
|
||||
@property
|
||||
def issuer(self) -> str:
|
||||
return f"https://login.microsoftonline.com/{self.tenant_id}/v2.0"
|
||||
|
||||
|
||||
class AgentIdentityBinding(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
agent_id: str
|
||||
active: bool = True
|
||||
provider: Literal["microsoft_entra"]
|
||||
tenant_id: str
|
||||
client_id: str
|
||||
service_principal_id: str | None = None
|
||||
issuer: str
|
||||
required_roles: tuple[str, ...] = ()
|
||||
required_scopes: tuple[str, ...] = ("user_impersonation",)
|
||||
revision: str
|
||||
last_authenticated_at: datetime | None = None
|
||||
|
||||
|
||||
class AgentSubject(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
kind: Literal["application", "delegated_subject"]
|
||||
oid: str
|
||||
mode: Literal["autonomous", "delegated"]
|
||||
|
||||
|
||||
class AgentIdentityFailure(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
code: Literal["identity_denied", "policy_unavailable"] = "identity_denied"
|
||||
message: str
|
||||
|
||||
|
||||
class ManagedAgentContext(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
agent_id: str
|
||||
binding_revision: str | None = None
|
||||
mode: Literal["autonomous", "delegated"]
|
||||
user_id: str | None = None
|
||||
subject_oid: str | None = None
|
||||
|
||||
|
||||
class VerifiedHumanSubject(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
issuer: str
|
||||
tenant_id: str
|
||||
oid: str
|
||||
user_id: str
|
||||
|
||||
|
||||
class MicrosoftInteractiveSubject(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
issuer: str
|
||||
tenant_id: str
|
||||
oid: str
|
||||
|
||||
|
||||
class ManagedAgentIdentityStatus(BaseModel):
|
||||
identity: AgentIdentityBinding | None = None
|
||||
identity_managed: bool = False
|
||||
enabled: bool = True
|
||||
execution_mode: AgentExecutionMode = "autonomous"
|
||||
last_authenticated_at: datetime | None = None
|
||||
|
|
@ -284,6 +284,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
|
|||
cache_creation_input_token_cost_above_272k_tokens: float | None
|
||||
cache_creation_input_token_cost_above_272k_tokens_priority: float | None
|
||||
cache_creation_input_token_cost_above_272k_tokens_flex: float | None
|
||||
cache_creation_input_token_cost_above_272k_tokens_ultrafast: ReadOnly[float | None]
|
||||
cache_creation_input_token_cost_above_1hr: float | None
|
||||
cache_creation_input_token_cost_flex: float | None # OpenAI flex service tier pricing
|
||||
cache_creation_input_token_cost_priority: float | None # OpenAI priority service tier pricing
|
||||
|
|
@ -300,6 +301,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
|
|||
cache_read_input_token_cost_above_272k_tokens: float | None
|
||||
cache_read_input_token_cost_above_272k_tokens_priority: float | None
|
||||
cache_read_input_token_cost_above_272k_tokens_flex: float | None
|
||||
cache_read_input_token_cost_above_272k_tokens_ultrafast: ReadOnly[float | None]
|
||||
cache_read_input_token_cost_above_512k_tokens: float | None
|
||||
cache_read_input_token_cost_batches: ReadOnly[float | None]
|
||||
cache_read_input_token_cost_above_200k_tokens_batches: ReadOnly[float | None]
|
||||
|
|
@ -319,6 +321,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
|
|||
input_cost_per_token_above_272k_tokens: float | None # GPT-5.4/5.4-pro: prompts >272K priced at 2x input
|
||||
input_cost_per_token_above_272k_tokens_priority: float | None
|
||||
input_cost_per_token_above_272k_tokens_flex: float | None
|
||||
input_cost_per_token_above_272k_tokens_ultrafast: ReadOnly[float | None]
|
||||
input_cost_per_token_above_512k_tokens: float | None # MiniMax-M3: prompts >512K priced at 2x input
|
||||
input_cost_per_character_above_128k_tokens: float | None # only for vertex ai models
|
||||
input_cost_per_query: float | None # per-request pricing: rerank, search, and Bedrock Marengo embeddings
|
||||
|
|
@ -360,6 +363,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
|
|||
output_cost_per_token_above_272k_tokens: float | None # GPT-5.4/5.4-pro: prompts >272K priced at 1.5x output
|
||||
output_cost_per_token_above_272k_tokens_priority: float | None
|
||||
output_cost_per_token_above_272k_tokens_flex: float | None
|
||||
output_cost_per_token_above_272k_tokens_ultrafast: ReadOnly[float | None]
|
||||
output_cost_per_token_above_512k_tokens: float | None # MiniMax-M3: prompts >512K priced at 2x output
|
||||
output_cost_per_character_above_128k_tokens: float | None # only for vertex ai models
|
||||
output_cost_per_image: float | None
|
||||
|
|
@ -3737,6 +3741,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
|
|||
cache_creation_input_token_cost_above_272k_tokens: float | None = None
|
||||
cache_creation_input_token_cost_above_272k_tokens_priority: float | None = None
|
||||
cache_creation_input_token_cost_above_272k_tokens_flex: float | None = None
|
||||
cache_creation_input_token_cost_above_272k_tokens_ultrafast: float | None = None
|
||||
cache_creation_input_token_cost_flex: float | None = None
|
||||
cache_creation_input_token_cost_priority: float | None = None
|
||||
cache_creation_input_token_cost_ultrafast: float | None = None
|
||||
|
|
@ -3749,6 +3754,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
|
|||
cache_read_input_token_cost_above_200k_tokens_priority: float | None = None
|
||||
cache_read_input_token_cost_above_272k_tokens_priority: float | None = None
|
||||
cache_read_input_token_cost_above_272k_tokens_flex: float | None = None
|
||||
cache_read_input_token_cost_above_272k_tokens_ultrafast: float | None = None
|
||||
cache_read_input_token_cost_batches: float | None = None
|
||||
cache_read_input_token_cost_above_200k_tokens_batches: float | None = None
|
||||
cache_read_input_token_cost_above_272k_tokens_batches: float | None = None
|
||||
|
|
@ -3765,6 +3771,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
|
|||
input_cost_per_token_above_200k_tokens_priority: float | None = None
|
||||
input_cost_per_token_above_272k_tokens_priority: float | None = None
|
||||
input_cost_per_token_above_272k_tokens_flex: float | None = None
|
||||
input_cost_per_token_above_272k_tokens_ultrafast: float | None = None
|
||||
input_cost_per_token_above_200k_tokens_batches: float | None = None
|
||||
input_cost_per_token_above_272k_tokens_batches: float | None = None
|
||||
input_cost_per_query: float | None = None
|
||||
|
|
@ -3791,6 +3798,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
|
|||
output_cost_per_token_above_200k_tokens_priority: float | None = None
|
||||
output_cost_per_token_above_272k_tokens_priority: float | None = None
|
||||
output_cost_per_token_above_272k_tokens_flex: float | None = None
|
||||
output_cost_per_token_above_272k_tokens_ultrafast: float | None = None
|
||||
output_cost_per_token_above_200k_tokens_batches: float | None = None
|
||||
output_cost_per_token_above_272k_tokens_batches: float | None = None
|
||||
output_cost_per_character_above_128k_tokens: float | None = None
|
||||
|
|
|
|||
|
|
@ -6104,6 +6104,9 @@ def _get_model_info_helper(
|
|||
cache_creation_input_token_cost_above_272k_tokens_flex=_model_info.get(
|
||||
"cache_creation_input_token_cost_above_272k_tokens_flex", None
|
||||
),
|
||||
cache_creation_input_token_cost_above_272k_tokens_ultrafast=_model_info.get(
|
||||
"cache_creation_input_token_cost_above_272k_tokens_ultrafast", None
|
||||
),
|
||||
cache_creation_input_token_cost_flex=_model_info.get("cache_creation_input_token_cost_flex", None),
|
||||
cache_creation_input_token_cost_priority=_model_info.get(
|
||||
"cache_creation_input_token_cost_priority", None
|
||||
|
|
@ -6129,6 +6132,9 @@ def _get_model_info_helper(
|
|||
cache_read_input_token_cost_above_272k_tokens_flex=_model_info.get(
|
||||
"cache_read_input_token_cost_above_272k_tokens_flex", None
|
||||
),
|
||||
cache_read_input_token_cost_above_272k_tokens_ultrafast=_model_info.get(
|
||||
"cache_read_input_token_cost_above_272k_tokens_ultrafast", None
|
||||
),
|
||||
cache_read_input_token_cost_above_512k_tokens=_model_info.get(
|
||||
"cache_read_input_token_cost_above_512k_tokens", None
|
||||
),
|
||||
|
|
@ -6167,6 +6173,9 @@ def _get_model_info_helper(
|
|||
input_cost_per_token_above_272k_tokens_flex=_model_info.get(
|
||||
"input_cost_per_token_above_272k_tokens_flex", None
|
||||
),
|
||||
input_cost_per_token_above_272k_tokens_ultrafast=_model_info.get(
|
||||
"input_cost_per_token_above_272k_tokens_ultrafast", None
|
||||
),
|
||||
input_cost_per_token_above_512k_tokens=_model_info.get("input_cost_per_token_above_512k_tokens", None),
|
||||
input_cost_per_query=_model_info.get("input_cost_per_query", None),
|
||||
cost_per_second=_model_info.get("cost_per_second", None),
|
||||
|
|
@ -6234,6 +6243,9 @@ def _get_model_info_helper(
|
|||
output_cost_per_token_above_272k_tokens_flex=_model_info.get(
|
||||
"output_cost_per_token_above_272k_tokens_flex", None
|
||||
),
|
||||
output_cost_per_token_above_272k_tokens_ultrafast=_model_info.get(
|
||||
"output_cost_per_token_above_272k_tokens_ultrafast", None
|
||||
),
|
||||
output_cost_per_token_above_512k_tokens=_model_info.get(
|
||||
"output_cost_per_token_above_512k_tokens", None
|
||||
),
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
]
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -7,3 +7,13 @@ reason = "diskcache has no fixed release published; remove this entry once one e
|
|||
id = "GHSA-h7x2-h6g9-p789"
|
||||
ignoreUntil = 2026-10-14
|
||||
reason = "mlflow has no fixed release published (3.16.0, 2026-09-04, and master still store gateway secret api_base unvalidated); remove this entry once one exists"
|
||||
|
||||
[[IgnoredVulns]]
|
||||
id = "GHSA-hj66-6f7g-4r5v"
|
||||
ignoreUntil = 2026-10-02
|
||||
reason = "oauthlib 4.0.0 (the only fixed release, 2026-09-28) is inside the 3-day uv exclude-newer cooldown; bump oauthlib and remove this entry once it clears"
|
||||
|
||||
[[IgnoredVulns]]
|
||||
id = "GHSA-xpv3-w29h-x7cv"
|
||||
ignoreUntil = 2026-10-02
|
||||
reason = "oauthlib 4.0.0 (the only fixed release, 2026-09-28) is inside the 3-day uv exclude-newer cooldown; bump oauthlib and remove this entry once it clears"
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm"
|
||||
version = "1.104.0"
|
||||
version = "1.105.0"
|
||||
description = "Library to easily interface with LLM API providers"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10, <3.15"
|
||||
|
|
@ -75,8 +75,8 @@ proxy = [
|
|||
"mcp>=2.2.0,<3",
|
||||
"httpx2>=2.5.0,<3",
|
||||
"pydantic>=2.12.0,<3",
|
||||
"litellm-proxy-extras==0.4.102",
|
||||
"litellm-enterprise==0.1.71",
|
||||
"litellm-proxy-extras==0.4.103",
|
||||
"litellm-enterprise==0.1.72",
|
||||
"RestrictedPython>=8.5,<9.0",
|
||||
"rich>=13.9.4,<14.0",
|
||||
"InquirerPy>=0.3.4,<1.0",
|
||||
|
|
@ -357,7 +357,7 @@ litellm-enterprise = { workspace = true }
|
|||
members = ["enterprise", "litellm-proxy-extras"]
|
||||
|
||||
[tool.commitizen]
|
||||
version = "1.104.0"
|
||||
version = "1.105.0"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -64,7 +64,7 @@ The suites run against a live proxy, so bring one up first by running the litell
|
|||
|
||||
For the opt-in browser profile, start the existing IdP first, then run `.github/e2e-stack/oidc-profile.sh "$PROXY_BASE_URL" <server-command>`. The wrapper creates a confidential client with an exact `/sso/callback` redirect and S256 PKCE, passes the client secret only through the child process environment, and removes the client on exit. It uses the existing generic OIDC handler with `GENERIC_USER_ID_ATTRIBUTE=sub`. Preserve the IdP's PostgreSQL data across restarts
|
||||
|
||||
`tests/e2e/ui/playwright.oidc.config.ts` uses an already running OIDC stack and separate storage/output files. Supply `E2E_OIDC_UI_URL`, `JWT_ISSUER`, `E2E_OIDC_USERNAME` and `E2E_OIDC_PASSWORD` for a seeded actor. Its setup follows the real login and callback path. The current Python canary qualifies browser-client configuration and token/userinfo identity mapping; browser journey specs under `ui/oidc/` are a separate coverage step
|
||||
`tests/e2e/ui/playwright.oidc.config.ts` uses an already running OIDC stack and separate storage/output files. Supply `E2E_OIDC_UI_URL`, `JWT_ISSUER`, `E2E_OIDC_USERNAME` and `E2E_OIDC_PASSWORD` for a seeded actor. Its setup follows the real login and callback path. The current Python canary qualifies browser-client configuration and token/userinfo identity mapping. The specs under `ui/oidc/` drive a real dashboard SSO login and a real `lite login`, so start the proxy with `EXPERIMENTAL_UI_LOGIN=true` and at least one model it can actually serve. The CLI spec runs `lite` from `PATH` unless `E2E_LITE_CLI` names another executable, and it gives the CLI a temporary `HOME` with the keyring disabled so your own login is never touched. The main `playwright.config.ts` ignores `oidc/`
|
||||
|
||||
Every successful IdP create immediately registers cleanup, including partial setup failures. Cleanup failures emit warnings. Tokens are minted on demand, and the expiration test waits relative to the token's actual `exp` with a bounded clock-drift check. To check first-attempt behavior locally, run both files with `--reruns 0`:
|
||||
|
||||
|
|
|
|||
|
|
@ -60,3 +60,6 @@
|
|||
|
||||
- {id: other.auth.jwt.wrong_issuer_denied, module: other, tier: P0, area: auth, assertions: [wrong_issuer_denied], source: "auth/handle_jwt.py", rationale: "A signed token with the correct audience and an unexpected issuer is rejected"}
|
||||
- {id: other.auth.jwt.wrong_audience_denied, module: other, tier: P0, area: auth, assertions: [wrong_audience_denied], source: "auth/handle_jwt.py", rationale: "A signed token from the trusted issuer intended for another app is rejected"}
|
||||
- {id: other.auth.session_token.valid_allows, module: other, tier: P0, area: auth, assertions: [valid_allows], source: "auth/auth_checks.py ExperimentalUIJWTToken", rationale: "An unexpired LiteLLM-minted session token authenticates with the role it carries"}
|
||||
- {id: other.auth.session_token.expired_denied, module: other, tier: P0, area: auth, assertions: [expired_denied], source: "auth/user_api_key_auth.py expiry check", rationale: "An expired session token is rejected with the expired-key error"}
|
||||
- {id: other.auth.session_token.encrypted_value_denied, module: other, tier: P0, area: auth, assertions: [encrypted_value_denied], source: "auth/auth_checks.py ExperimentalUIJWTToken", rationale: "An encrypted value read back from a management route is not accepted as a bearer token"}
|
||||
|
|
|
|||
91
tests/e2e/other/test_session_token_e2e.py
Normal file
91
tests/e2e/other/test_session_token_e2e.py
Normal file
|
|
@ -0,0 +1,91 @@
|
|||
"""Live e2e: UI/CLI session tokens are accepted only while valid and only when minted as session tokens.
|
||||
|
||||
The runner mints its own session tokens under the proxy's salt key, so the valid and expired cases run in
|
||||
seconds instead of waiting out a real login's expiry.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
|
||||
from e2e_config import MASTER_KEY, unique_marker
|
||||
from e2e_http import UnauthorizedError, unwrap
|
||||
from lifecycle import ResourceManager
|
||||
from models import KeyGenerateBody, KeyLoggingCallback, KeyLoggingCallbackVars, KeyMetadata
|
||||
from other_client import OtherClient
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
SALT_KEY: Final = os.environ.get("LITELLM_SALT_KEY") or MASTER_KEY
|
||||
SESSION_TOKEN_PREFIX: Final = "litellm_login_"
|
||||
ENCRYPTED_PREFIX: Final = "litellm_enc::"
|
||||
|
||||
|
||||
def _admin_session_token(expires_at: datetime) -> str:
|
||||
claims: Final = json.dumps(
|
||||
{
|
||||
"token": f"ui-token-{unique_marker()}",
|
||||
"user_id": f"e2e-session-{unique_marker()}",
|
||||
"user_role": "proxy_admin",
|
||||
"team_id": "litellm-dashboard",
|
||||
"expires": expires_at.isoformat(),
|
||||
}
|
||||
)
|
||||
nonce: Final = os.urandom(12)
|
||||
sealed: Final = AESGCM(hashlib.sha256(SALT_KEY.encode()).digest()).encrypt(
|
||||
nonce, claims.encode(), SESSION_TOKEN_PREFIX.encode()
|
||||
)
|
||||
return SESSION_TOKEN_PREFIX + base64.urlsafe_b64encode(nonce + sealed).decode().rstrip("=")
|
||||
|
||||
|
||||
class TestSessionToken:
|
||||
@pytest.mark.covers("other.auth.session_token.valid_allows")
|
||||
def test_unexpired_session_token_reaches_admin_route(self, client: OtherClient) -> None:
|
||||
token: Final = _admin_session_token(datetime.now(timezone.utc) + timedelta(minutes=10))
|
||||
listing: Final = unwrap(client.list_users_as(token))
|
||||
assert listing.total >= 0, f"an unexpired admin session token did not reach /user/list: {listing}"
|
||||
|
||||
@pytest.mark.covers("other.auth.session_token.expired_denied")
|
||||
def test_expired_session_token_is_denied(self, client: OtherClient) -> None:
|
||||
token: Final = _admin_session_token(datetime.now(timezone.utc) - timedelta(minutes=1))
|
||||
result: Final = client.list_users_as(token)
|
||||
assert isinstance(result, UnauthorizedError), f"an expired session token must get 401, got {result}"
|
||||
assert "expired" in result.body.lower(), f"expected the expired-key error, got {result.body[:300]}"
|
||||
|
||||
@pytest.mark.covers("other.auth.session_token.encrypted_value_denied")
|
||||
def test_encrypted_stored_value_is_not_a_bearer_token(
|
||||
self, client: OtherClient, resources: ResourceManager
|
||||
) -> None:
|
||||
stored_value: Final = f'{{"token": "{unique_marker()}", "user_role": "proxy_admin"}}'
|
||||
key: Final = client.proxy.generate_key(
|
||||
KeyGenerateBody(
|
||||
key_alias=f"e2e-session-{unique_marker()}",
|
||||
metadata=KeyMetadata(
|
||||
logging=[
|
||||
KeyLoggingCallback(
|
||||
callback_name="langfuse",
|
||||
callback_vars=KeyLoggingCallbackVars(langfuse_secret_key=stored_value),
|
||||
)
|
||||
]
|
||||
),
|
||||
)
|
||||
)
|
||||
resources.defer(lambda: client.proxy.delete_key(key))
|
||||
|
||||
metadata: Final = client.proxy.key_info(key).metadata
|
||||
assert metadata is not None and metadata.logging, f"/key/info dropped the logging metadata: {metadata}"
|
||||
encrypted: Final = metadata.logging[0].callback_vars.langfuse_secret_key
|
||||
assert encrypted is not None and encrypted.startswith(ENCRYPTED_PREFIX), (
|
||||
f"expected /key/info to return the stored secret encrypted, got {encrypted!r}"
|
||||
)
|
||||
|
||||
for bearer in (encrypted.removeprefix(ENCRYPTED_PREFIX), encrypted):
|
||||
result = client.list_users_as(bearer)
|
||||
assert isinstance(result, UnauthorizedError), f"an encrypted stored value must get 401, got {result}"
|
||||
|
|
@ -51,6 +51,21 @@ export async function sendChatCompletion(request: APIRequestContext, opts: ChatO
|
|||
return body.id as string;
|
||||
}
|
||||
|
||||
export interface ServedChat {
|
||||
requestId: string;
|
||||
callId: string;
|
||||
}
|
||||
|
||||
export async function sendChatCompletionWithCallId(request: APIRequestContext, opts: ChatOptions): Promise<ServedChat> {
|
||||
const res = await postChatCompletion(request, opts);
|
||||
expect(res.ok(), `chat completion for ${opts.model} failed (${res.status()}): ${await res.text()}`).toBe(true);
|
||||
const callId = res.headers()["x-litellm-call-id"];
|
||||
expect(callId, "proxy did not return an x-litellm-call-id header").toBeTruthy();
|
||||
const body = await res.json();
|
||||
expect(body.choices?.[0]?.message?.content).toContain(MOCK_RESPONSE_TEXT);
|
||||
return { requestId: body.id as string, callId };
|
||||
}
|
||||
|
||||
export interface ChatAttempt {
|
||||
status: number;
|
||||
body: string;
|
||||
|
|
@ -124,7 +139,7 @@ export async function waitForSpendLog(
|
|||
lastStatus = res.status();
|
||||
if (res.ok()) {
|
||||
const body = await res.json();
|
||||
const rows = Array.isArray(body) ? body : (body?.data ?? []);
|
||||
const rows = Array.isArray(body) ? body : body?.data ?? [];
|
||||
if (rows.length > 0) {
|
||||
return;
|
||||
}
|
||||
|
|
|
|||
87
tests/e2e/ui/oidc/cliLogin.spec.ts
Normal file
87
tests/e2e/ui/oidc/cliLogin.spec.ts
Normal file
|
|
@ -0,0 +1,87 @@
|
|||
import { expect, test } from "@playwright/test";
|
||||
import { execFile, spawn } from "node:child_process";
|
||||
import * as fs from "node:fs";
|
||||
import * as os from "node:os";
|
||||
import * as path from "node:path";
|
||||
import { promisify } from "node:util";
|
||||
|
||||
const LITE_CLI = process.env.E2E_LITE_CLI ?? "lite";
|
||||
const SKIP_TEAM_SELECTION = "skip\n";
|
||||
const execFileAsync = promisify(execFile);
|
||||
|
||||
function requiredEnv(name: string): string {
|
||||
const value = process.env[name];
|
||||
if (!value) throw new Error(`${name} must be set for the OIDC suite`);
|
||||
return value;
|
||||
}
|
||||
|
||||
test("CLI SSO login stores a session that lists models and completes a chat request", async ({ browser, baseURL }) => {
|
||||
test.setTimeout(180_000);
|
||||
const issuer = requiredEnv("JWT_ISSUER");
|
||||
const home = fs.mkdtempSync(path.join(os.tmpdir(), "lite-cli-login-"));
|
||||
const browserUrlFile = path.join(home, "browser-url");
|
||||
const browserCommand = path.join(home, "browser.sh");
|
||||
fs.writeFileSync(browserCommand, `#!/bin/sh\nprintf '%s' "$1" > '${browserUrlFile}'\n`, { mode: 0o700 });
|
||||
const env = {
|
||||
...process.env,
|
||||
HOME: home,
|
||||
LITELLM_CLI_DISABLE_KEYRING: "1",
|
||||
BROWSER: browserCommand,
|
||||
PYTHONUNBUFFERED: "1",
|
||||
FORCE_COLOR: undefined,
|
||||
NO_COLOR: "1",
|
||||
LITELLM_PROXY_URL: baseURL,
|
||||
LITELLM_PROXY_API_KEY: undefined,
|
||||
};
|
||||
const login = spawn(LITE_CLI, ["login"], { env });
|
||||
let loginOutput = "";
|
||||
login.stdout.on("data", (chunk: Buffer) => (loginOutput += chunk.toString()));
|
||||
login.stderr.on("data", (chunk: Buffer) => (loginOutput += chunk.toString()));
|
||||
const loginExit = new Promise<number | null>((resolve) => login.on("close", resolve));
|
||||
login.stdin.end(SKIP_TEAM_SELECTION);
|
||||
try {
|
||||
await expect.poll(() => fs.existsSync(browserUrlFile), { timeout: 30_000 }).toBe(true);
|
||||
await expect.poll(() => loginOutput).toMatch(/Verification code: \S+/);
|
||||
const userCode = /Verification code: (\S+)/.exec(loginOutput)?.[1] ?? "";
|
||||
|
||||
const context = await browser.newContext({ storageState: { cookies: [], origins: [] } });
|
||||
try {
|
||||
const page = await context.newPage();
|
||||
await page.goto(fs.readFileSync(browserUrlFile, "utf8"));
|
||||
await expect(page).toHaveURL((url) => url.href.startsWith(`${issuer}/`));
|
||||
await page.getByLabel("Username or email").fill(requiredEnv("E2E_OIDC_USERNAME"));
|
||||
await page.getByLabel("Password", { exact: true }).fill(requiredEnv("E2E_OIDC_PASSWORD"));
|
||||
await page.getByRole("button", { name: "Sign In", exact: true }).click();
|
||||
await page.getByLabel("Verification code").fill(userCode);
|
||||
await page.getByRole("button", { name: "Continue", exact: true }).click();
|
||||
await expect(page.getByRole("heading", { name: "Authentication Successful!" })).toBeVisible();
|
||||
} finally {
|
||||
await context.close();
|
||||
}
|
||||
|
||||
expect(await loginExit, loginOutput).toBe(0);
|
||||
expect(loginOutput).toContain("Login successful!");
|
||||
const stored: { key?: unknown } = JSON.parse(fs.readFileSync(path.join(home, ".litellm", "token.json"), "utf8"));
|
||||
expect(typeof stored.key).toBe("string");
|
||||
expect(stored.key, "CLI login issues a session token, not a virtual key").not.toMatch(/^sk-/);
|
||||
|
||||
const { stdout: modelsJson } = await execFileAsync(LITE_CLI, ["models", "list", "--format", "json"], { env });
|
||||
const models: { id: string }[] = JSON.parse(modelsJson);
|
||||
expect(models.length, "the stack serves at least one model").toBeGreaterThan(0);
|
||||
|
||||
const chatRequest = JSON.stringify({
|
||||
model: models[0].id,
|
||||
messages: [{ role: "user", content: "Reply with the single word: ok" }],
|
||||
});
|
||||
const { stdout: completionJson } = await execFileAsync(
|
||||
LITE_CLI,
|
||||
["http", "request", "POST", "/chat/completions", "-j", chatRequest],
|
||||
{ env },
|
||||
);
|
||||
const completion: { choices: { message: { content: string | null } }[] } = JSON.parse(completionJson);
|
||||
expect(completion.choices[0]?.message.content).toBeTruthy();
|
||||
} finally {
|
||||
login.kill();
|
||||
fs.rmSync(home, { recursive: true, force: true });
|
||||
}
|
||||
});
|
||||
35
tests/e2e/ui/oidc/dashboardLogin.spec.ts
Normal file
35
tests/e2e/ui/oidc/dashboardLogin.spec.ts
Normal file
|
|
@ -0,0 +1,35 @@
|
|||
import { expect, test, type Page as PlaywrightPage, type Response } from "@playwright/test";
|
||||
import { Page } from "../fixtures/pages";
|
||||
import { navigateToPage } from "../helpers/navigation";
|
||||
|
||||
function sessionKey(tokenCookie: string): string {
|
||||
const claims: unknown = JSON.parse(Buffer.from(tokenCookie.split(".")[1] ?? "", "base64url").toString("utf8"));
|
||||
const key = claims !== null && typeof claims === "object" && "key" in claims ? claims.key : undefined;
|
||||
if (typeof key !== "string") throw new Error("The dashboard token cookie carries no key claim");
|
||||
return key;
|
||||
}
|
||||
|
||||
async function openPageAndCapture(page: PlaywrightPage, target: Page, apiPath: string): Promise<Response> {
|
||||
const response = page.waitForResponse((r) => new URL(r.url()).pathname === apiPath);
|
||||
await navigateToPage(page, target);
|
||||
return response;
|
||||
}
|
||||
|
||||
test("SSO login issues a session that authorizes dashboard data requests", async ({ page, context, baseURL }) => {
|
||||
const tokenCookie = (await context.cookies(baseURL)).find((cookie) => cookie.name === "token");
|
||||
expect(tokenCookie, "SSO login sets the dashboard token cookie").toBeDefined();
|
||||
const key = sessionKey(tokenCookie?.value ?? "");
|
||||
expect(key, "SSO login issues a session token, not a virtual key").not.toMatch(/^sk-/);
|
||||
|
||||
const keyList = await openPageAndCapture(page, Page.ApiKeys, "/key/list");
|
||||
expect(keyList.request().headers()["authorization"]).toBe(`Bearer ${key}`);
|
||||
expect(keyList.status()).toBe(200);
|
||||
expect(Array.isArray((await keyList.json()).keys)).toBe(true);
|
||||
|
||||
const modelInfo = await openPageAndCapture(page, Page.Models, "/v2/model/info");
|
||||
expect(modelInfo.request().headers()["authorization"]).toBe(`Bearer ${key}`);
|
||||
expect(modelInfo.status()).toBe(200);
|
||||
const models: { model_name: string }[] = (await modelInfo.json()).data;
|
||||
expect(models.length, "the stack serves at least one model").toBeGreaterThan(0);
|
||||
await expect(page.getByText(models[0].model_name, { exact: true }).first()).toBeVisible();
|
||||
});
|
||||
|
|
@ -8,7 +8,7 @@ import { ARTIFACT_DIR, UI_BASE_URL } from "./constants";
|
|||
export default defineConfig({
|
||||
testDir: ".",
|
||||
testMatch: ["**/*.spec.ts", "**/*.setup.ts"],
|
||||
testIgnore: ["**/*.test.*", "**/integrationCritical/**"],
|
||||
testIgnore: ["**/*.test.*", "**/integrationCritical/**", "oidc/**"],
|
||||
/* Run tests in files in parallel */
|
||||
fullyParallel: true,
|
||||
/* Fail the build on CI if you accidentally left test.only in the source code. */
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ import {
|
|||
CHAT_MODEL_A,
|
||||
MOCK_RESPONSE_TEXT,
|
||||
sendChatCompletion,
|
||||
sendChatCompletionWithCallId,
|
||||
waitForSpendLog,
|
||||
waitForSpendLogByPrompt,
|
||||
} from "../../helpers/traffic";
|
||||
|
|
@ -95,6 +96,50 @@ test.describe("Logs page", () => {
|
|||
await expect(drawer.getByText(MOCK_RESPONSE_TEXT, { exact: false }).first()).toBeVisible({ timeout: 20_000 });
|
||||
});
|
||||
|
||||
test("a served request's Logs row and drawer show its x-litellm-call-id", async ({ page, request }) => {
|
||||
const prompt = `logs-call-id-prompt-${uniqueSuffix()}`;
|
||||
const { requestId, callId } = await sendChatCompletionWithCallId(request, {
|
||||
model: CHAT_MODEL_A,
|
||||
prompt,
|
||||
});
|
||||
expect(callId, "call id must differ from the provider response id for this check to mean anything").not.toBe(
|
||||
requestId,
|
||||
);
|
||||
await waitForSpendLog(request, requestId);
|
||||
|
||||
await navigateToPage(page, Page.Logs);
|
||||
await dismissFeedbackPopup(page);
|
||||
const search = visibleTestId(page, "datatable-search");
|
||||
await expect(search).toBeVisible({ timeout: 20_000 });
|
||||
await search.fill(callId);
|
||||
|
||||
const row = requestLogsRows(page).filter({ hasText: requestId });
|
||||
await expect(row, `no logs row for call id ${callId}`).toHaveCount(1, { timeout: 30_000 });
|
||||
await expect(row, "the row itself shows only the request id").not.toContainText(callId);
|
||||
|
||||
await row.getByText(requestId).hover();
|
||||
const tooltip = page.locator("[data-slot='tooltip-content']");
|
||||
await expect(tooltip, "hovering the Request ID cell does not list the x-litellm-call-id").toContainText(
|
||||
`x-litellm-call-id: ${callId}`,
|
||||
{ timeout: 10_000 },
|
||||
);
|
||||
await tooltip.getByRole("button", { name: "Copy x-litellm-call-id" }).click();
|
||||
if (await page.evaluate(() => window.isSecureContext)) {
|
||||
await expect.poll(() => page.evaluate(() => navigator.clipboard.readText())).toBe(callId);
|
||||
}
|
||||
|
||||
await row.click();
|
||||
const drawer = page.getByRole("dialog").first();
|
||||
await expect(drawer.getByText("Request & Response")).toBeVisible({ timeout: 20_000 });
|
||||
await expect(drawer.getByText("x-litellm-call-id:"), "drawer header lacks the x-litellm-call-id line").toBeVisible({
|
||||
timeout: 10_000,
|
||||
});
|
||||
await expect(
|
||||
drawer.getByText(callId, { exact: false }).first(),
|
||||
`drawer does not show x-litellm-call-id ${callId}`,
|
||||
).toBeVisible({ timeout: 10_000 });
|
||||
});
|
||||
|
||||
// Split out because only the copy path needs a secure context; folding it in would
|
||||
// take the drawer-rendering coverage down with it.
|
||||
test("the drawer copies the request and the response to the clipboard", async ({ page, request }) => {
|
||||
|
|
|
|||
|
|
@ -3,6 +3,8 @@ import uuid
|
|||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta
|
||||
from hashlib import sha256
|
||||
from itertools import chain
|
||||
from typing import Final
|
||||
|
||||
|
|
@ -34,7 +36,16 @@ INSERT_SPEND_LOG: Final = (
|
|||
" VALUES (%s, 'acompletion', %s, %s::timestamp, %s::timestamp, %s)"
|
||||
)
|
||||
DELETE_SPEND_LOG: Final = 'DELETE FROM "LiteLLM_SpendLogs" WHERE request_id = %s'
|
||||
INSERT_SPEND_LOG_ROW: Final = (
|
||||
'INSERT INTO "LiteLLM_SpendLogs" (request_id, call_type, api_key, "startTime", "endTime", metadata, team_id, "user")'
|
||||
" VALUES (%s, 'acompletion', %s, %s::timestamp, %s::timestamp, %s, %s, %s)"
|
||||
)
|
||||
DELETE_SPEND_LOG_ROWS: Final = 'DELETE FROM "LiteLLM_SpendLogs" WHERE request_id = ANY(%s)'
|
||||
DELETE_KEY_ROW: Final = 'DELETE FROM "LiteLLM_VerificationToken" WHERE token = %s'
|
||||
DELETE_ARCHIVED_KEY_ROW: Final = 'DELETE FROM "LiteLLM_DeletedVerificationToken" WHERE token = %s'
|
||||
LOCK_TABLE: Final = sql.SQL("LOCK TABLE {table} IN ACCESS EXCLUSIVE MODE")
|
||||
SPEND_LOGS_TABLE: Final = "LiteLLM_SpendLogs"
|
||||
FIRST_SPEND_LOG_AT: Final = datetime(2026, 2, 3, 12, 0, 0)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -67,6 +78,10 @@ def key_no_key_table_holds() -> str:
|
|||
return f"integration-ownerless-{uuid.uuid4().hex}"
|
||||
|
||||
|
||||
def digest_no_key_table_holds() -> str:
|
||||
return sha256(uuid.uuid4().bytes).hexdigest()
|
||||
|
||||
|
||||
def activity_of_key(
|
||||
gateway: Gateway, path: str, api_key: str, *, reader: str | None = None, **filters: str
|
||||
) -> httpx.Response:
|
||||
|
|
@ -148,6 +163,56 @@ def spend_log_naming_only_an_alias(request_id: str, api_key: str, started: str,
|
|||
connection.execute(DELETE_SPEND_LOG, (request_id,))
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SpendLogRow:
|
||||
started: str
|
||||
metadata: JsonValue = None
|
||||
team_id: str | None = None
|
||||
user: str | None = None
|
||||
|
||||
|
||||
def started_at(index: int) -> str:
|
||||
return (FIRST_SPEND_LOG_AT + timedelta(seconds=index)).strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
|
||||
def nameless_rows(count: int, first_index: int = 0) -> tuple[SpendLogRow, ...]:
|
||||
return tuple(SpendLogRow(started_at(first_index + offset), {}) for offset in range(count))
|
||||
|
||||
|
||||
def named_row(index: int, alias: str) -> SpendLogRow:
|
||||
return SpendLogRow(started_at(index), {"user_api_key_alias": alias})
|
||||
|
||||
|
||||
@contextmanager
|
||||
def spend_logs_of_key(
|
||||
api_key: str, rows: Sequence[SpendLogRow], *, database_url: str | None = None
|
||||
) -> Iterator[tuple[str, ...]]:
|
||||
request_ids: Final = tuple(f"integration-{uuid.uuid4().hex}" for _ in rows)
|
||||
with psycopg.connect(database_url or os.environ["DATABASE_URL"]) as connection:
|
||||
connection.cursor().executemany(
|
||||
INSERT_SPEND_LOG_ROW,
|
||||
tuple(
|
||||
(request_id, api_key, row.started, row.started, Jsonb(row.metadata), row.team_id, row.user)
|
||||
for request_id, row in zip(request_ids, rows, strict=True)
|
||||
),
|
||||
)
|
||||
try:
|
||||
yield request_ids
|
||||
finally:
|
||||
delete_spend_logs(request_ids, database_url=database_url)
|
||||
|
||||
|
||||
def delete_spend_logs(request_ids: Sequence[str], *, database_url: str | None = None) -> None:
|
||||
with psycopg.connect(database_url or os.environ["DATABASE_URL"]) as connection:
|
||||
connection.execute(DELETE_SPEND_LOG_ROWS, (list(request_ids),))
|
||||
|
||||
|
||||
def purge_key_from_the_key_tables(digest: str, *, database_url: str | None = None) -> None:
|
||||
with psycopg.connect(database_url or os.environ["DATABASE_URL"]) as connection:
|
||||
connection.execute(DELETE_KEY_ROW, (digest,))
|
||||
connection.execute(DELETE_ARCHIVED_KEY_ROW, (digest,))
|
||||
|
||||
|
||||
@contextmanager
|
||||
def locked_table(table: str, *, database_url: str | None = None) -> Iterator[None]:
|
||||
with psycopg.connect(database_url or os.environ["DATABASE_URL"]) as connection:
|
||||
|
|
|
|||
|
|
@ -324,7 +324,7 @@ ROUTES: Final[tuple[Route, ...]] = (
|
|||
lambda s: Call("POST", "/team/update", {"team_id": s.team_id, "max_budget": 5}),
|
||||
team_admin=403, others=403, org_admin=200),
|
||||
Route("team_update_budget_permitted",
|
||||
lambda s: Call("POST", "/team/update", {"team_id": s.team_id, "max_budget": 7}),
|
||||
lambda s: Call("POST", "/team/update", {"team_id": s.team_id, "max_budget": 4}),
|
||||
team_admin=200, others=403, org_admin=200, permission="max_budget"),
|
||||
Route("project_new",
|
||||
lambda s: Call("POST", "/project/new", {"team_id": s.team_id, "project_alias": f"matrix-{uuid.uuid4().hex}"}),
|
||||
|
|
@ -414,6 +414,8 @@ def test_status_code(shared: TeamScenario, org_team: TeamScenario, route: Route,
|
|||
team: Final = org_team if caller in ORG_CALLERS else shared
|
||||
with team.gateway.scenario() as scenario:
|
||||
s: Final = replace(team, scenario=scenario)
|
||||
if route.name == "team_update_budget_permitted":
|
||||
s.gateway.post("/team/update", {"team_id": s.team_id, "max_budget": 5})
|
||||
if route.permission:
|
||||
scenario.cleanups.enter_context(team_admin_permissions(s.gateway, (route.permission,)))
|
||||
call: Final = route.call(s)
|
||||
|
|
@ -421,5 +423,9 @@ def test_status_code(shared: TeamScenario, org_team: TeamScenario, route: Route,
|
|||
assert response.status_code == route.expected(caller), (
|
||||
f"{caller} {call.method} {call.path}: {response.status_code} {response.text}"
|
||||
)
|
||||
if route.name == "team_update_budget_permitted":
|
||||
assert read_rows(
|
||||
'SELECT max_budget FROM "LiteLLM_TeamTable" WHERE team_id = %s', (s.team_id,)
|
||||
) == [{"max_budget": 4.0 if response.status_code == 200 else 5.0}]
|
||||
if response.status_code == 200 and route.cleanup is not None:
|
||||
route.cleanup(s, object_value(response.json()))
|
||||
|
|
|
|||
|
|
@ -1,11 +1,15 @@
|
|||
import json
|
||||
from typing import Final
|
||||
import uuid
|
||||
from typing import Final, Literal
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from pydantic import JsonValue
|
||||
|
||||
from tests.integration._support.client import JSON_OBJECT, Gateway, eventually, object_value, string_value
|
||||
from tests.integration._support.database import read_rows
|
||||
from tests.integration._support.upstream import delete_scenario, register_scenario
|
||||
from tests.integration.cost_calculation.cost_tracking_case import JsonResponse
|
||||
|
||||
STANDARD_INPUT_RATE: Final = 0.001
|
||||
STANDARD_OUTPUT_RATE: Final = 0.002
|
||||
|
|
@ -69,3 +73,190 @@ def test_ultrafast_service_tier_bills_ultrafast_rates_and_keeps_pricing_off_the_
|
|||
)
|
||||
assert_chat_bills_rates(gateway, model, "ultrafast", ULTRAFAST_INPUT_RATE, ULTRAFAST_OUTPUT_RATE)
|
||||
assert_chat_bills_rates(gateway, model, None, STANDARD_INPUT_RATE, STANDARD_OUTPUT_RATE)
|
||||
|
||||
|
||||
LONG_CONTEXT_PRICING: Final[dict[str, JsonValue]] = {
|
||||
"input_cost_per_token": 1e-06,
|
||||
"output_cost_per_token": 2e-06,
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
"input_cost_per_token_above_272k_tokens": 3e-06,
|
||||
"output_cost_per_token_above_272k_tokens": 4e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 3e-07,
|
||||
"input_cost_per_token_ultrafast": 1e-05,
|
||||
"output_cost_per_token_ultrafast": 2e-05,
|
||||
"cache_read_input_token_cost_ultrafast": 1e-06,
|
||||
"input_cost_per_token_above_272k_tokens_ultrafast": 5e-05,
|
||||
"output_cost_per_token_above_272k_tokens_ultrafast": 6e-05,
|
||||
"cache_read_input_token_cost_above_272k_tokens_ultrafast": 5e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens_ultrafast": 6e-06,
|
||||
}
|
||||
LONG_PROMPT_TOKENS: Final = 300_000
|
||||
SHORT_PROMPT_TOKENS: Final = 1_000
|
||||
CACHED_TOKENS: Final = 400
|
||||
COMPLETION_TOKENS: Final = 1_000
|
||||
|
||||
|
||||
def _chat_response(service_tier: str | None, prompt_tokens: int) -> JsonResponse:
|
||||
return JsonResponse(
|
||||
content_type="application/json",
|
||||
body={
|
||||
"id": "chatcmpl-$UNIQUE_ID",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": "integration-ultrafast-long-context",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "long context answer"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": prompt_tokens,
|
||||
"completion_tokens": COMPLETION_TOKENS,
|
||||
"total_tokens": prompt_tokens + COMPLETION_TOKENS,
|
||||
"prompt_tokens_details": {"cached_tokens": CACHED_TOKENS},
|
||||
},
|
||||
**({} if service_tier is None else {"service_tier": service_tier}),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _responses_response(service_tier: str | None, prompt_tokens: int) -> JsonResponse:
|
||||
return JsonResponse(
|
||||
content_type="application/json",
|
||||
body={
|
||||
"id": "resp_$UNIQUE_ID",
|
||||
"object": "response",
|
||||
"created_at": 1,
|
||||
"status": "completed",
|
||||
"model": "integration-ultrafast-long-context",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_$UNIQUE_ID",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "long context answer", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"input_tokens": prompt_tokens,
|
||||
"output_tokens": COMPLETION_TOKENS,
|
||||
"total_tokens": prompt_tokens + COMPLETION_TOKENS,
|
||||
"input_tokens_details": {"cached_tokens": CACHED_TOKENS},
|
||||
"output_tokens_details": {"reasoning_tokens": 0},
|
||||
},
|
||||
**({} if service_tier is None else {"service_tier": service_tier}),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _surface_response(
|
||||
surface: Literal["chat", "responses"], service_tier: str | None, prompt_tokens: int
|
||||
) -> JsonResponse:
|
||||
match surface:
|
||||
case "chat":
|
||||
return _chat_response(service_tier, prompt_tokens)
|
||||
case "responses":
|
||||
return _responses_response(service_tier, prompt_tokens)
|
||||
|
||||
|
||||
def _surface_request(
|
||||
surface: Literal["chat", "responses"], scenario_id: str, model: str, service_tier: str | None
|
||||
) -> tuple[str, dict[str, JsonValue], str]:
|
||||
match surface:
|
||||
case "chat":
|
||||
return (
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": "long context ultrafast control"}],
|
||||
**({} if service_tier is None else {"service_tier": service_tier}),
|
||||
},
|
||||
f"/{scenario_id}/chat/completions",
|
||||
)
|
||||
case "responses":
|
||||
return (
|
||||
"/v1/responses",
|
||||
{
|
||||
"model": model,
|
||||
"input": "long context ultrafast control",
|
||||
**({} if service_tier is None else {"service_tier": service_tier}),
|
||||
},
|
||||
f"/{scenario_id}/responses",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("service_tier", "prompt_tokens", "input_rate", "cache_read_rate", "output_rate"),
|
||||
(
|
||||
("ultrafast", LONG_PROMPT_TOKENS, 5e-05, 5e-06, 6e-05),
|
||||
("ultrafast", SHORT_PROMPT_TOKENS, 1e-05, 1e-06, 2e-05),
|
||||
(None, LONG_PROMPT_TOKENS, 3e-06, 3e-07, 4e-06),
|
||||
),
|
||||
ids=("ultrafast_above_272k", "ultrafast_below_272k", "standard_above_272k"),
|
||||
)
|
||||
@pytest.mark.parametrize("surface", ("chat", "responses"), ids=("chat", "responses"))
|
||||
def test_ultrafast_long_context_prompt_bills_ultrafast_long_context_rates(
|
||||
gateway: Gateway,
|
||||
surface: Literal["chat", "responses"],
|
||||
service_tier: str | None,
|
||||
prompt_tokens: int,
|
||||
input_rate: float,
|
||||
cache_read_rate: float,
|
||||
output_rate: float,
|
||||
) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
scenario_id: Final = f"ultrafast-long-context-{uuid.uuid4().hex}"
|
||||
handle: Final = register_scenario(
|
||||
scenario_id, _surface_response(surface, service_tier, prompt_tokens)
|
||||
)
|
||||
scenario.cleanups.callback(delete_scenario, handle)
|
||||
key: Final = scenario.key()
|
||||
model: Final = scenario.model(
|
||||
model=f"openai/integration-ultrafast-long-context-{uuid.uuid4().hex}",
|
||||
api_key=scenario_id,
|
||||
api_base=handle.api_base(),
|
||||
**LONG_CONTEXT_PRICING,
|
||||
)
|
||||
request_path, request_body, expected_upstream_path = _surface_request(surface, scenario_id, model, service_tier)
|
||||
with httpx.Client(base_url=gateway.upstream_url, trust_env=False) as upstream:
|
||||
upstream.get("/__observations").raise_for_status()
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
request_path,
|
||||
request_body,
|
||||
key=key,
|
||||
)
|
||||
observations: Final = JSON_OBJECT.validate_json(upstream.get("/__observations").content)["requests"]
|
||||
assert response.status_code == 200, response.text
|
||||
expected_input: Final = (prompt_tokens - CACHED_TOKENS) * input_rate + CACHED_TOKENS * cache_read_rate
|
||||
expected_output: Final = COMPLETION_TOKENS * output_rate
|
||||
expected: Final = expected_input + expected_output
|
||||
assert float(response.headers["x-litellm-response-cost"]) == pytest.approx(expected, rel=1e-6), response.text
|
||||
request_id: Final = string_value(object_value(response.json())["id"])
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT spend, metadata, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" WHERE request_id = %s',
|
||||
(request_id,),
|
||||
),
|
||||
lambda values: len(values) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
assert rows[0]["prompt_tokens"] == prompt_tokens
|
||||
assert rows[0]["completion_tokens"] == COMPLETION_TOKENS
|
||||
assert float(rows[0]["spend"]) == pytest.approx(expected, rel=1e-6)
|
||||
metadata: Final = rows[0]["metadata"]
|
||||
parsed: Final = json.loads(metadata) if isinstance(metadata, str) else object_value(metadata)
|
||||
breakdown: Final = object_value(parsed["cost_breakdown"])
|
||||
assert float(breakdown["input_cost"]) == pytest.approx(expected_input, rel=1e-6)
|
||||
assert float(breakdown["output_cost"]) == pytest.approx(expected_output, rel=1e-6)
|
||||
assert isinstance(observations, list)
|
||||
assert len(observations) == 1
|
||||
observation: Final = object_value(observations[0])
|
||||
upstream_path: Final = string_value(observation["path"])
|
||||
assert upstream_path == expected_upstream_path, upstream_path
|
||||
body: Final = object_value(observation["body"])
|
||||
assert body.get("service_tier") == service_tier, body
|
||||
assert not set(LONG_CONTEXT_PRICING).intersection(body), body
|
||||
|
|
|
|||
|
|
@ -12,6 +12,8 @@ def test_chat_over_responses_deployment_returns_length_when_output_tokens_run_ou
|
|||
identity: Final = "responses-incomplete-" + uuid.uuid4().hex
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
if request.method == "GET" and request.target == "/v1/models":
|
||||
return Reply(body=b'{"object":"list","data":[]}')
|
||||
assert request.method == "POST" and request.target == "/responses", request.target
|
||||
assert request.headers["authorization"] == "Bearer synthetic-openai-key"
|
||||
body: Final = json.loads(request.body)
|
||||
|
|
@ -56,7 +58,7 @@ def test_chat_over_responses_deployment_returns_length_when_output_tokens_run_ou
|
|||
)
|
||||
assert response.status_code == 200, response.text
|
||||
body: Final = response.json()
|
||||
assert len(wire.drain()) == 1
|
||||
assert len(tuple(request for request in wire.drain() if request.method == "POST")) == 1
|
||||
assert [choice["finish_reason"] for choice in body["choices"]] == ["length"], response.text
|
||||
assert body["choices"][0]["message"]["content"] == "", response.text
|
||||
assert body["choices"][0]["message"]["role"] == "assistant", response.text
|
||||
|
|
@ -69,6 +71,8 @@ def test_messages_over_responses_deployment_with_max_tokens_1_is_clamped_to_16_i
|
|||
identity: Final = "responses-clamp-" + uuid.uuid4().hex
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
if request.method == "GET" and request.target == "/v1/models":
|
||||
return Reply(body=b'{"object":"list","data":[]}')
|
||||
assert request.method == "POST" and request.target == "/responses", request.target
|
||||
assert request.headers["authorization"] == "Bearer synthetic-openai-key"
|
||||
body: Final = json.loads(request.body)
|
||||
|
|
@ -132,7 +136,7 @@ def test_messages_over_responses_deployment_with_max_tokens_1_is_clamped_to_16_i
|
|||
)
|
||||
assert response.status_code == 200, response.text
|
||||
body: Final = response.json()
|
||||
assert len(wire.drain()) == 1
|
||||
assert len(tuple(request for request in wire.drain() if request.method == "POST")) == 1
|
||||
assert body["role"] == "assistant", response.text
|
||||
assert body["content"] == [{"type": "text", "text": "ok"}], response.text
|
||||
assert body["stop_reason"] == "end_turn", response.text
|
||||
|
|
@ -143,6 +147,8 @@ def test_messages_over_responses_deployment_with_max_tokens_one_reaches_openai_a
|
|||
identity: Final = "responses-min-tokens-" + uuid.uuid4().hex
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
if request.method == "GET" and request.target == "/v1/models":
|
||||
return Reply(body=b'{"object":"list","data":[]}')
|
||||
assert request.method == "POST" and request.target == "/responses", request.target
|
||||
assert request.headers["authorization"] == "Bearer synthetic-openai-key"
|
||||
body: Final = json.loads(request.body)
|
||||
|
|
@ -185,6 +191,6 @@ def test_messages_over_responses_deployment_with_max_tokens_one_reaches_openai_a
|
|||
)
|
||||
assert response.status_code == 200, response.text
|
||||
body: Final = response.json()
|
||||
assert len(wire.drain()) == 1
|
||||
assert len(tuple(request for request in wire.drain() if request.method == "POST")) == 1
|
||||
assert body["content"] == [{"type": "text", "text": "ok"}], response.text
|
||||
assert body["usage"]["input_tokens"] == 9 and body["usage"]["output_tokens"] == 1, response.text
|
||||
|
|
|
|||
|
|
@ -0,0 +1,262 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import shlex
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from contextlib import contextmanager
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from queue import SimpleQueue
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.client import Gateway, eventually
|
||||
from integration._support.process import owned_proxy
|
||||
from integration._support.redis_process import owned_redis
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
from redis import Redis
|
||||
from redis.exceptions import TimeoutError as RedisTimeoutError
|
||||
|
||||
JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
|
||||
MONITOR_COMMAND: Final = TypeAdapter(dict[str, JsonValue])
|
||||
OPENAI_MODEL: Final = "gpt-4o-mini"
|
||||
MASTER_KEY: Final = "sk-integration-usage-routing-redis-reads"
|
||||
API_KEY: Final = "synthetic-usage-routing-key"
|
||||
ENDPOINT_PATHS: Final = MappingProxyType(
|
||||
{
|
||||
"/v1/chat/completions": ("/v1/chat/completions", "/v1/chat/completions"),
|
||||
"/v1/messages": ("/v1/responses", "/v1/responses"),
|
||||
"/v1/responses": ("/v1/responses", "/v1/responses"),
|
||||
}
|
||||
)
|
||||
CHAT_RESPONSE: Final = json.dumps(
|
||||
{
|
||||
"id": "chatcmpl_usage_routing_redis_reads",
|
||||
"object": "chat.completion",
|
||||
"created": 1700000000,
|
||||
"model": OPENAI_MODEL,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "redis read contract"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 7, "completion_tokens": 4, "total_tokens": 11},
|
||||
}
|
||||
).encode()
|
||||
RESPONSES_RESPONSE: Final = json.dumps(
|
||||
{
|
||||
"id": "resp_usage_routing_redis_reads",
|
||||
"object": "response",
|
||||
"created_at": 1700000000,
|
||||
"status": "completed",
|
||||
"model": OPENAI_MODEL,
|
||||
"output": [
|
||||
{
|
||||
"id": "msg_usage_routing_redis_reads",
|
||||
"type": "message",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "redis read contract", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 7, "output_tokens": 4, "total_tokens": 11},
|
||||
}
|
||||
).encode()
|
||||
|
||||
|
||||
def _request_object(body: bytes) -> dict[str, JsonValue]:
|
||||
return JSON_OBJECT.validate_json(body)
|
||||
|
||||
|
||||
def _deployment_list(
|
||||
model_name: str, api_base: str, deployment_ids: tuple[str, str]
|
||||
) -> list[dict[str, JsonValue]]:
|
||||
return [
|
||||
{
|
||||
"model_name": model_name,
|
||||
"litellm_params": {
|
||||
"model": f"openai/{OPENAI_MODEL}",
|
||||
"api_base": api_base,
|
||||
"api_key": API_KEY,
|
||||
"rpm": 1,
|
||||
},
|
||||
"model_info": {"id": deployment_id},
|
||||
}
|
||||
for deployment_id in deployment_ids
|
||||
]
|
||||
|
||||
|
||||
def _request_payload(endpoint: str, model_name: str, marker: str) -> dict[str, JsonValue]:
|
||||
if endpoint == "/v1/responses":
|
||||
return {"model": model_name, "input": marker, "max_output_tokens": 16, "store": False}
|
||||
return {"model": model_name, "messages": [{"role": "user", "content": marker}], "max_tokens": 16}
|
||||
|
||||
|
||||
def _expected_wire_body(endpoint: str, marker: str) -> dict[str, JsonValue]:
|
||||
if endpoint == "/v1/messages":
|
||||
return {
|
||||
"model": OPENAI_MODEL,
|
||||
"input": [{"type": "message", "role": "user", "content": [{"type": "input_text", "text": marker}]}],
|
||||
"include": ["reasoning.encrypted_content"],
|
||||
"max_output_tokens": 16,
|
||||
}
|
||||
if endpoint == "/v1/responses":
|
||||
return {"model": OPENAI_MODEL, "input": marker, "max_output_tokens": 16, "store": False}
|
||||
return {"model": OPENAI_MODEL, "messages": [{"role": "user", "content": marker}], "max_tokens": 16}
|
||||
|
||||
|
||||
def _reply(request: Request) -> Reply:
|
||||
if request.target == "/v1/models":
|
||||
return Reply(body=json.dumps({"object": "list", "data": [{"id": OPENAI_MODEL, "object": "model"}]}).encode())
|
||||
return Reply(body=RESPONSES_RESPONSE if request.target == "/v1/responses" else CHAT_RESPONSE)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _capture_redis_commands(host: str, port: int) -> Iterator[SimpleQueue[str]]:
|
||||
commands: Final = SimpleQueue[str]()
|
||||
started: Final = threading.Event()
|
||||
armed: Final = threading.Event()
|
||||
stopped: Final = threading.Event()
|
||||
ready_marker: Final = f"monitor-ready-{uuid.uuid4().hex}"
|
||||
stop_marker: Final = f"monitor-stop-{uuid.uuid4().hex}"
|
||||
|
||||
def capture() -> None:
|
||||
with Redis(host=host, port=port, socket_timeout=1, decode_responses=True) as client:
|
||||
with client.monitor() as monitor:
|
||||
started.set()
|
||||
stream: Final = iter(monitor.listen())
|
||||
while not stopped.is_set():
|
||||
try:
|
||||
record: Final = MONITOR_COMMAND.validate_python(next(stream))
|
||||
except RedisTimeoutError:
|
||||
continue
|
||||
command: Final = record.get("command")
|
||||
if not isinstance(command, str):
|
||||
continue
|
||||
commands.put(command)
|
||||
if ready_marker in command:
|
||||
armed.set()
|
||||
|
||||
thread: Final = threading.Thread(target=capture, daemon=True)
|
||||
thread.start()
|
||||
try:
|
||||
assert started.wait(timeout=5), "Redis MONITOR did not start"
|
||||
with Redis(host=host, port=port, socket_timeout=1, decode_responses=True) as client:
|
||||
client.set(ready_marker, "ready", ex=1)
|
||||
assert armed.wait(timeout=5), "Redis MONITOR did not capture its readiness command"
|
||||
yield commands
|
||||
finally:
|
||||
stopped.set()
|
||||
with Redis(host=host, port=port, socket_timeout=1, decode_responses=True) as client:
|
||||
client.set(stop_marker, "stop", ex=1)
|
||||
thread.join(timeout=5)
|
||||
assert not thread.is_alive(), "Redis MONITOR thread survived cleanup"
|
||||
|
||||
|
||||
def _drain_mgets(commands: SimpleQueue[str]) -> tuple[tuple[str, tuple[str, ...]], ...]:
|
||||
captured: Final = tuple(commands.get_nowait() for _ in range(commands.qsize()))
|
||||
parsed: Final = tuple((line, tuple(shlex.split(line))) for line in captured)
|
||||
return tuple((line, arguments) for line, arguments in parsed if arguments and arguments[0] == "MGET")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"endpoint",
|
||||
("/v1/chat/completions", "/v1/messages", "/v1/responses"),
|
||||
ids=("chat-completions", "messages", "responses"),
|
||||
)
|
||||
def test_proxy_usage_routing_reads_cooldown_tpm_then_rpm_from_redis(
|
||||
endpoint: str, tmp_path: Path
|
||||
) -> None:
|
||||
with owned_redis(tmp_path) as cache, wire_server(_reply) as wire:
|
||||
run_id: Final = uuid.uuid4().hex
|
||||
model_name: Final = f"usage-redis-{run_id}"
|
||||
deployment_ids: Final = (f"dep-a-{run_id[:8]}", f"dep-b-{run_id[:8]}")
|
||||
configuration: Final = JSON_OBJECT.validate_python(
|
||||
yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
)
|
||||
config: Final = {
|
||||
**configuration,
|
||||
"model_list": _deployment_list(model_name, f"{wire.url}/v1", deployment_ids),
|
||||
"router_settings": {
|
||||
"routing_strategy": "usage-based-routing-v2",
|
||||
"redis_host": cache.host,
|
||||
"redis_port": cache.port,
|
||||
},
|
||||
}
|
||||
config_path: Final = tmp_path / "usage-routing.yaml"
|
||||
config_path.write_text(yaml.safe_dump(config))
|
||||
with httpx.Client(base_url=wire.url, timeout=15, trust_env=False) as bootstrap_client:
|
||||
bootstrap: Final = Gateway(bootstrap_client, MASTER_KEY, wire.url)
|
||||
with owned_proxy(
|
||||
bootstrap,
|
||||
tmp_path,
|
||||
{"REDIS_HOST": cache.host, "REDIS_PORT": str(cache.port)},
|
||||
config=config_path,
|
||||
) as candidate:
|
||||
eventually(
|
||||
lambda: wire.received.qsize(),
|
||||
lambda received: received >= len(deployment_ids),
|
||||
seconds=15,
|
||||
)
|
||||
wire.drain()
|
||||
eventually(
|
||||
lambda: datetime.now(UTC),
|
||||
lambda current: current.second < 40,
|
||||
seconds=65,
|
||||
)
|
||||
minute: Final = datetime.now(UTC).strftime("%H-%M")
|
||||
markers: Final = tuple(f"{run_id}-{index}" for index in range(3))
|
||||
payloads: Final = tuple(_request_payload(endpoint, model_name, marker) for marker in markers)
|
||||
request_headers: Final = (
|
||||
{"anthropic-version": "2023-06-01"} if endpoint == "/v1/messages" else {}
|
||||
)
|
||||
with _capture_redis_commands(cache.host, cache.port) as commands:
|
||||
responses: Final = tuple(
|
||||
candidate.request("POST", endpoint, payload, headers=request_headers) for payload in payloads
|
||||
)
|
||||
assert tuple(response.status_code for response in responses) == (200, 200, 429), [
|
||||
response.text for response in responses
|
||||
]
|
||||
assert "No deployments available" in responses[2].text
|
||||
served_ids: Final = tuple(response.headers["x-litellm-model-id"] for response in responses[:2])
|
||||
assert set(served_ids) == set(deployment_ids), served_ids
|
||||
rpm_keys: Final = tuple(
|
||||
f"{deployment_id}:openai/{OPENAI_MODEL}:rpm:{minute}" for deployment_id in deployment_ids
|
||||
)
|
||||
with Redis(host=cache.host, port=cache.port, decode_responses=True) as redis_client:
|
||||
rpm_values: Final = eventually(
|
||||
lambda: tuple(redis_client.get(key) for key in rpm_keys),
|
||||
lambda values: values == ("1", "1"),
|
||||
seconds=15,
|
||||
)
|
||||
assert rpm_values == ("1", "1")
|
||||
received: Final = wire.drain()
|
||||
assert len(received) == 2
|
||||
assert tuple(request.method for request in received) == ("POST", "POST")
|
||||
assert tuple(request.target for request in received) == ENDPOINT_PATHS[endpoint]
|
||||
observed_bodies: Final = tuple(_request_object(request.body) for request in received)
|
||||
expected_bodies: Final = tuple(_expected_wire_body(endpoint, marker) for marker in markers[:2])
|
||||
assert observed_bodies == expected_bodies, observed_bodies
|
||||
expected_mget: Final = (
|
||||
"MGET",
|
||||
f"deployment:{deployment_ids[0]}:cooldown",
|
||||
f"deployment:{deployment_ids[1]}:cooldown",
|
||||
f"{deployment_ids[0]}:openai/{OPENAI_MODEL}:tpm:{minute}",
|
||||
f"{deployment_ids[1]}:openai/{OPENAI_MODEL}:tpm:{minute}",
|
||||
*(
|
||||
f"{deployment_id}:openai/{OPENAI_MODEL}:rpm:{minute}"
|
||||
for deployment_id in deployment_ids
|
||||
),
|
||||
)
|
||||
mgets: Final = _drain_mgets(commands)
|
||||
assert any(arguments == expected_mget for _, arguments in mgets), mgets
|
||||
raw_mgets: Final = tuple(line for line, _ in mgets)
|
||||
print(f"proxy {endpoint} MGETs: {raw_mgets}")
|
||||
|
|
@ -0,0 +1,191 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import shlex
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from contextlib import contextmanager
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from queue import SimpleQueue
|
||||
from typing import Final
|
||||
|
||||
import litellm
|
||||
import pytest
|
||||
from integration._support.client import eventually
|
||||
from integration._support.redis_process import owned_redis
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from litellm import Router
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
from redis import Redis
|
||||
from redis.exceptions import TimeoutError as RedisTimeoutError
|
||||
|
||||
JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
|
||||
MONITOR_COMMAND: Final = TypeAdapter(dict[str, JsonValue])
|
||||
OPENAI_MODEL: Final = "gpt-4o-mini"
|
||||
API_KEY: Final = "synthetic-usage-routing-key"
|
||||
CHAT_RESPONSE: Final = json.dumps(
|
||||
{
|
||||
"id": "chatcmpl_usage_routing_sdk_redis_reads",
|
||||
"object": "chat.completion",
|
||||
"created": 1700000000,
|
||||
"model": OPENAI_MODEL,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "redis read contract"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 7, "completion_tokens": 4, "total_tokens": 11},
|
||||
}
|
||||
).encode()
|
||||
|
||||
|
||||
def _request_object(body: bytes) -> dict[str, JsonValue]:
|
||||
return JSON_OBJECT.validate_json(body)
|
||||
|
||||
|
||||
def _deployment_list(
|
||||
model_name: str, api_base: str, deployment_ids: tuple[str, str]
|
||||
) -> list[dict[str, JsonValue]]:
|
||||
return [
|
||||
{
|
||||
"model_name": model_name,
|
||||
"litellm_params": {
|
||||
"model": f"openai/{OPENAI_MODEL}",
|
||||
"api_base": api_base,
|
||||
"api_key": API_KEY,
|
||||
"rpm": 1,
|
||||
},
|
||||
"model_info": {"id": deployment_id},
|
||||
}
|
||||
for deployment_id in deployment_ids
|
||||
]
|
||||
|
||||
|
||||
def _reply(request: Request) -> Reply:
|
||||
return Reply(body=CHAT_RESPONSE)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _capture_redis_commands(host: str, port: int) -> Iterator[SimpleQueue[str]]:
|
||||
commands: Final = SimpleQueue[str]()
|
||||
started: Final = threading.Event()
|
||||
armed: Final = threading.Event()
|
||||
stopped: Final = threading.Event()
|
||||
ready_marker: Final = f"monitor-ready-{uuid.uuid4().hex}"
|
||||
stop_marker: Final = f"monitor-stop-{uuid.uuid4().hex}"
|
||||
|
||||
def capture() -> None:
|
||||
with Redis(host=host, port=port, socket_timeout=1, decode_responses=True) as client:
|
||||
with client.monitor() as monitor:
|
||||
started.set()
|
||||
stream: Final = iter(monitor.listen())
|
||||
while not stopped.is_set():
|
||||
try:
|
||||
record: Final = MONITOR_COMMAND.validate_python(next(stream))
|
||||
except RedisTimeoutError:
|
||||
continue
|
||||
command: Final = record.get("command")
|
||||
if not isinstance(command, str):
|
||||
continue
|
||||
commands.put(command)
|
||||
if ready_marker in command:
|
||||
armed.set()
|
||||
|
||||
thread: Final = threading.Thread(target=capture, daemon=True)
|
||||
thread.start()
|
||||
try:
|
||||
assert started.wait(timeout=5), "Redis MONITOR did not start"
|
||||
with Redis(host=host, port=port, socket_timeout=1, decode_responses=True) as client:
|
||||
client.set(ready_marker, "ready", ex=1)
|
||||
assert armed.wait(timeout=5), "Redis MONITOR did not capture its readiness command"
|
||||
yield commands
|
||||
finally:
|
||||
stopped.set()
|
||||
with Redis(host=host, port=port, socket_timeout=1, decode_responses=True) as client:
|
||||
client.set(stop_marker, "stop", ex=1)
|
||||
thread.join(timeout=5)
|
||||
assert not thread.is_alive(), "Redis MONITOR thread survived cleanup"
|
||||
|
||||
|
||||
def _drain_mgets(commands: SimpleQueue[str]) -> tuple[tuple[str, tuple[str, ...]], ...]:
|
||||
captured: Final = tuple(commands.get_nowait() for _ in range(commands.qsize()))
|
||||
parsed: Final = tuple((line, tuple(shlex.split(line))) for line in captured)
|
||||
return tuple((line, arguments) for line, arguments in parsed if arguments and arguments[0] == "MGET")
|
||||
|
||||
|
||||
def _model_id(response: object) -> str:
|
||||
response_params: Final = getattr(response, "_hidden_params")
|
||||
hidden_params: Final = JSON_OBJECT.validate_python(response_params)
|
||||
model_id: Final = hidden_params.get("model_id")
|
||||
assert isinstance(model_id, str), hidden_params
|
||||
return model_id
|
||||
|
||||
|
||||
async def _exercise_router(router: Router, model_name: str, markers: tuple[str, str, str]) -> tuple[str, str]:
|
||||
first: Final = await router.acompletion(
|
||||
model=model_name, messages=[{"role": "user", "content": markers[0]}], max_tokens=8
|
||||
)
|
||||
second: Final = await router.acompletion(
|
||||
model=model_name, messages=[{"role": "user", "content": markers[1]}], max_tokens=8
|
||||
)
|
||||
with pytest.raises(litellm.RateLimitError, match="No deployments available"):
|
||||
await router.acompletion(
|
||||
model=model_name, messages=[{"role": "user", "content": markers[2]}], max_tokens=8
|
||||
)
|
||||
return _model_id(first), _model_id(second)
|
||||
|
||||
|
||||
def test_sdk_usage_routing_reads_tpm_then_rpm_from_redis(tmp_path: Path) -> None:
|
||||
with owned_redis(tmp_path) as cache, wire_server(_reply) as wire:
|
||||
run_id: Final = uuid.uuid4().hex
|
||||
model_name: Final = f"usage-redis-{run_id}"
|
||||
deployment_ids: Final = (f"dep-a-{run_id[:8]}", f"dep-b-{run_id[:8]}")
|
||||
router: Final = Router(
|
||||
model_list=_deployment_list(model_name, f"{wire.url}/v1", deployment_ids),
|
||||
routing_strategy="usage-based-routing-v2",
|
||||
redis_host=cache.host,
|
||||
redis_port=cache.port,
|
||||
)
|
||||
try:
|
||||
eventually(
|
||||
lambda: datetime.now(UTC),
|
||||
lambda current: current.second < 40,
|
||||
seconds=65,
|
||||
)
|
||||
minute: Final = datetime.now(UTC).strftime("%H-%M")
|
||||
markers: Final = tuple(f"{run_id}-{index}" for index in range(3))
|
||||
with _capture_redis_commands(cache.host, cache.port) as commands:
|
||||
served_ids: Final = asyncio.run(_exercise_router(router, model_name, markers))
|
||||
assert set(served_ids) == set(deployment_ids), served_ids
|
||||
received: Final = wire.drain()
|
||||
assert len(received) == 2
|
||||
assert tuple(request.method for request in received) == ("POST", "POST")
|
||||
assert tuple(request.target for request in received) == ("/v1/chat/completions",) * 2
|
||||
observed_bodies: Final = tuple(_request_object(request.body) for request in received)
|
||||
expected_bodies: Final = tuple(
|
||||
{
|
||||
"model": OPENAI_MODEL,
|
||||
"messages": [{"role": "user", "content": marker}],
|
||||
"max_tokens": 8,
|
||||
}
|
||||
for marker in markers[:2]
|
||||
)
|
||||
assert observed_bodies == expected_bodies, observed_bodies
|
||||
expected_mget: Final = (
|
||||
"MGET",
|
||||
f"deployment:{deployment_ids[0]}:cooldown",
|
||||
f"deployment:{deployment_ids[1]}:cooldown",
|
||||
*(f"{deployment_id}:openai/{OPENAI_MODEL}:tpm:{minute}" for deployment_id in deployment_ids),
|
||||
*(f"{deployment_id}:openai/{OPENAI_MODEL}:rpm:{minute}" for deployment_id in deployment_ids),
|
||||
)
|
||||
mgets: Final = _drain_mgets(commands)
|
||||
assert any(arguments == expected_mget for _, arguments in mgets), mgets
|
||||
raw_mgets: Final = tuple(line for line, _ in mgets)
|
||||
print(f"sdk MGETs: {raw_mgets}")
|
||||
finally:
|
||||
router.reset()
|
||||
|
|
@ -2,11 +2,12 @@ from __future__ import annotations
|
|||
|
||||
import json
|
||||
import uuid
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from hashlib import sha256
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from integration._support.client import JSON_OBJECT, Gateway, eventually, string_value
|
||||
from integration._support.client import JSON_OBJECT, Gateway, eventually, object_value, string_value
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.upstream import delete_scenario, register_scenario
|
||||
from integration.cost_calculation.cost_tracking_case import JsonResponse, RoutedResponse, TextResponse
|
||||
|
|
@ -116,6 +117,28 @@ def _batch_routes(model: str) -> RoutedResponse:
|
|||
)
|
||||
|
||||
|
||||
def _team_day_endpoints(gateway: Gateway, team: str, start_date: str, end_date: str) -> dict[str, object] | None:
|
||||
response: Final = gateway.request(
|
||||
"GET",
|
||||
"/team/daily/activity",
|
||||
params={"team_ids": team, "start_date": start_date, "end_date": end_date},
|
||||
)
|
||||
if response.status_code != 200:
|
||||
return None
|
||||
days: Final = response.json()["results"]
|
||||
if not days:
|
||||
return None
|
||||
return object_value(object_value(object_value(days[0])["breakdown"])["endpoints"])
|
||||
|
||||
|
||||
def _batches_total_tokens(endpoints: dict[str, object] | None) -> int | None:
|
||||
if endpoints is None or "/batches" not in endpoints:
|
||||
return None
|
||||
metrics: Final = object_value(object_value(endpoints["/batches"])["metrics"])
|
||||
total_tokens: Final = metrics["total_tokens"]
|
||||
return int(total_tokens) if isinstance(total_tokens, (int, float, str)) else None
|
||||
|
||||
|
||||
def _input_file(model: str) -> bytes:
|
||||
return (
|
||||
"\n".join(
|
||||
|
|
@ -200,3 +223,79 @@ def test_completed_batch_spend_row_records_reasoning_tokens_and_error_file_failu
|
|||
"reasoning_tokens": reasoning_tokens,
|
||||
"text_tokens": completion_tokens - reasoning_tokens,
|
||||
}, json.dumps(metadata)
|
||||
|
||||
|
||||
INPUT_COST_PER_TOKEN: Final = 0.001
|
||||
OUTPUT_COST_PER_TOKEN: Final = 0.002
|
||||
BATCH_PROMPT_TOKENS: Final = FIRST_LINE["prompt_tokens"] + SECOND_LINE["prompt_tokens"]
|
||||
BATCH_COMPLETION_TOKENS: Final = FIRST_LINE["completion_tokens"] + SECOND_LINE["completion_tokens"]
|
||||
BATCH_SPEND: Final = (BATCH_PROMPT_TOKENS * INPUT_COST_PER_TOKEN + BATCH_COMPLETION_TOKENS * OUTPUT_COST_PER_TOKEN) / 2
|
||||
|
||||
|
||||
def test_completed_batch_spend_lands_under_batches_in_team_endpoint_activity(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
scenario_id: Final = f"batch-endpoint-{uuid.uuid4().hex[:12]}"
|
||||
handle: Final = register_scenario(scenario_id, _batch_routes("gpt-4o-mini"))
|
||||
scenario.cleanups.callback(delete_scenario, handle)
|
||||
model: Final = scenario.model(
|
||||
api_base=handle.api_base(),
|
||||
input_cost_per_token=INPUT_COST_PER_TOKEN,
|
||||
output_cost_per_token=OUTPUT_COST_PER_TOKEN,
|
||||
)
|
||||
team: Final = scenario.team(models=[model])
|
||||
key: Final = scenario.key(team_id=team, models=[model])
|
||||
file_response: Final = gateway.request_multipart(
|
||||
"/v1/files",
|
||||
{"purpose": "batch", "model": model},
|
||||
{"file": ("in.jsonl", _input_file(model), "application/jsonl")},
|
||||
key=key,
|
||||
)
|
||||
assert file_response.status_code == 200, file_response.text
|
||||
batch_response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/batches",
|
||||
{
|
||||
"input_file_id": string_value(JSON_OBJECT.validate_json(file_response.content)["id"]),
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"completion_window": "24h",
|
||||
"model": model,
|
||||
},
|
||||
key=key,
|
||||
)
|
||||
assert batch_response.status_code == 200, batch_response.text
|
||||
batch_id: Final = string_value(JSON_OBJECT.validate_json(batch_response.content)["id"])
|
||||
retrieval: Final = gateway.request("GET", f"/v1/batches/{batch_id}", key=key)
|
||||
assert retrieval.status_code == 200, retrieval.text
|
||||
assert retrieval.json()["status"] == "completed", retrieval.text
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT spend, prompt_tokens, completion_tokens, total_tokens FROM "LiteLLM_SpendLogs" '
|
||||
"WHERE api_key=%s AND call_type='aretrieve_batch'",
|
||||
(sha256(key.encode()).hexdigest(),),
|
||||
),
|
||||
lambda values: len(values) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
row: Final = rows[0]
|
||||
assert float(row["spend"]) == pytest.approx(BATCH_SPEND), dict(row)
|
||||
assert (row["prompt_tokens"], row["completion_tokens"]) == (
|
||||
BATCH_PROMPT_TOKENS,
|
||||
BATCH_COMPLETION_TOKENS,
|
||||
), dict(row)
|
||||
today: Final = datetime.now(timezone.utc)
|
||||
endpoints: Final = eventually(
|
||||
lambda: _team_day_endpoints(
|
||||
gateway,
|
||||
team,
|
||||
(today - timedelta(days=1)).strftime("%Y-%m-%d"),
|
||||
(today + timedelta(days=1)).strftime("%Y-%m-%d"),
|
||||
),
|
||||
lambda value: _batches_total_tokens(value) == BATCH_PROMPT_TOKENS + BATCH_COMPLETION_TOKENS,
|
||||
seconds=70,
|
||||
return_last_on_timeout=True,
|
||||
)
|
||||
assert endpoints is not None, "team daily activity returned no endpoint breakdown for the day"
|
||||
assert set(endpoints) == {"/batches"}, endpoints
|
||||
endpoint_metrics: Final = object_value(object_value(endpoints["/batches"])["metrics"])
|
||||
assert float(endpoint_metrics["spend"]) == pytest.approx(BATCH_SPEND), endpoints
|
||||
assert endpoint_metrics["total_tokens"] == BATCH_PROMPT_TOKENS + BATCH_COMPLETION_TOKENS, endpoints
|
||||
|
|
|
|||
490
tests/integration/spend/test_daily_activity_key_alias_probes.py
Normal file
490
tests/integration/spend/test_daily_activity_key_alias_probes.py
Normal file
|
|
@ -0,0 +1,490 @@
|
|||
import time
|
||||
import uuid
|
||||
from collections.abc import Callable, Iterator
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from contextlib import contextmanager
|
||||
from hashlib import sha256
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from integration._support.client import Gateway, eventually, object_value
|
||||
from integration._support.daily_activity import (
|
||||
AGGREGATED_USER_ACTIVITY,
|
||||
DAY,
|
||||
ROUTES,
|
||||
SPEND_LOGS_TABLE,
|
||||
USER_SPEND,
|
||||
Route,
|
||||
SpendLogRow,
|
||||
activity_of_key,
|
||||
assert_key_reported,
|
||||
daily_rows,
|
||||
digest_no_key_table_holds,
|
||||
key_metadata,
|
||||
locked_table,
|
||||
named_row,
|
||||
nameless_rows,
|
||||
records_of_key,
|
||||
seeded_metrics,
|
||||
seeded_row,
|
||||
spend_logs_of_key,
|
||||
started_at,
|
||||
user_row,
|
||||
user_with_an_email,
|
||||
)
|
||||
from integration._support.database import read_rows, scratch_database
|
||||
from integration._support.process import OwnedProxy, owned_proxy_process
|
||||
from pydantic import JsonValue
|
||||
|
||||
DAY_OUTSIDE_THE_WINDOW: Final = "2026-02-10"
|
||||
GIVES_UP_WITHIN_SECONDS: Final = 10
|
||||
CONCURRENT_READS: Final = 20
|
||||
CACHED_MISS_CLEARS_WITHIN_SECONDS: Final = 45
|
||||
ALIAS_OF_ONE_SPEND_LOG: Final = (
|
||||
"SELECT metadata->>'user_api_key_alias' AS alias FROM \"LiteLLM_SpendLogs\" WHERE request_id = %s"
|
||||
)
|
||||
|
||||
|
||||
def _alias() -> str:
|
||||
return f"integration-alias-{uuid.uuid4().hex}"
|
||||
|
||||
|
||||
def _named_between_fifty_and_fifty(alias: str) -> tuple[SpendLogRow, ...]:
|
||||
return (*nameless_rows(50), named_row(50, alias), *nameless_rows(50, 51))
|
||||
|
||||
|
||||
def _oldest_named(alias: str) -> tuple[SpendLogRow, ...]:
|
||||
return (named_row(0, alias), *nameless_rows(150, 1))
|
||||
|
||||
|
||||
def _newest_named(alias: str) -> tuple[SpendLogRow, ...]:
|
||||
return (*nameless_rows(150), named_row(150, alias))
|
||||
|
||||
|
||||
def _both_edges_named(alias: str) -> tuple[SpendLogRow, ...]:
|
||||
return (named_row(0, alias), *nameless_rows(150, 1), named_row(151, alias))
|
||||
|
||||
|
||||
def _named_after_one_hundred(alias: str) -> tuple[SpendLogRow, ...]:
|
||||
return (*nameless_rows(100), named_row(100, alias), *nameless_rows(99, 101))
|
||||
|
||||
|
||||
def _named_after_ninety_nine(alias: str) -> tuple[SpendLogRow, ...]:
|
||||
return (*nameless_rows(99), named_row(99, alias), *nameless_rows(100, 100))
|
||||
|
||||
|
||||
def _named_only_in_the_middle(alias: str) -> tuple[SpendLogRow, ...]:
|
||||
return (*nameless_rows(100), named_row(100, alias), *nameless_rows(100, 101))
|
||||
|
||||
|
||||
def _renamed_and_renamed_back(alias: str, other: str) -> tuple[SpendLogRow, ...]:
|
||||
return (
|
||||
named_row(0, alias),
|
||||
*nameless_rows(100, 1),
|
||||
named_row(101, other),
|
||||
*nameless_rows(100, 102),
|
||||
named_row(202, alias),
|
||||
)
|
||||
|
||||
|
||||
def _team_in_the_column(team: str) -> SpendLogRow:
|
||||
return SpendLogRow(started_at(0), {}, team_id=team)
|
||||
|
||||
|
||||
def _team_in_the_metadata(team: str) -> SpendLogRow:
|
||||
return SpendLogRow(started_at(0), {"user_api_key_team_id": team})
|
||||
|
||||
|
||||
def _user_in_the_column(user: str) -> SpendLogRow:
|
||||
return SpendLogRow(started_at(0), {}, user=user)
|
||||
|
||||
|
||||
def _user_in_the_metadata(user: str) -> SpendLogRow:
|
||||
return SpendLogRow(started_at(0), {"user_api_key_user_id": user})
|
||||
|
||||
|
||||
def _activity_on_route(gateway: Gateway, route: Route, api_key: str, entity: str) -> httpx.Response:
|
||||
filters: Final = {} if route.entity_filter is None else {route.entity_filter: entity}
|
||||
return activity_of_key(gateway, route.path, api_key, **filters)
|
||||
|
||||
|
||||
def _reported_aliases(response: httpx.Response, api_key: str) -> tuple[JsonValue, ...]:
|
||||
if response.status_code != 200:
|
||||
return ()
|
||||
return tuple(
|
||||
object_value(object_value(record)["metadata"])["key_alias"]
|
||||
for record in records_of_key(object_value(response.json()), api_key)
|
||||
)
|
||||
|
||||
|
||||
def _names_the_key(api_key: str, alias: str) -> Callable[[httpx.Response], bool]:
|
||||
def names(response: httpx.Response) -> bool:
|
||||
reported: Final = _reported_aliases(response, api_key)
|
||||
return bool(reported) and frozenset(reported) == frozenset((alias,))
|
||||
|
||||
return names
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _proxy_on(gateway: Gateway, directory: Path, database_url: str, *, workers: int = 1) -> Iterator[OwnedProxy]:
|
||||
with owned_proxy_process(
|
||||
gateway,
|
||||
directory,
|
||||
{"DATABASE_URL": database_url},
|
||||
remove_environment=("DATABASE_URL_READ_REPLICA",),
|
||||
workers=workers,
|
||||
) as owned:
|
||||
yield owned
|
||||
|
||||
|
||||
def _owner_on(candidate: Gateway) -> tuple[str, str]:
|
||||
owner: Final = f"integration-{uuid.uuid4().hex}"
|
||||
email: Final = f"{owner}@example.com"
|
||||
candidate.post("/user/new", {"user_id": owner, "user_email": email, "auto_create_key": False})
|
||||
return owner, email
|
||||
|
||||
|
||||
@pytest.mark.parametrize("route", ROUTES, ids=lambda route: route.path.strip("/").replace("/", "_"))
|
||||
def test_alias_named_only_by_a_spend_log_is_reported_on_every_daily_activity_route(
|
||||
gateway: Gateway, route: Route
|
||||
) -> None:
|
||||
api_key: Final = digest_no_key_table_holds()
|
||||
alias: Final = _alias()
|
||||
entity: Final = f"integration-entity-{uuid.uuid4().hex}"
|
||||
with gateway.scenario() as scenario:
|
||||
owner, email = user_with_an_email(scenario)
|
||||
entity_rows: Final = (
|
||||
() if route.table == USER_SPEND else (seeded_row(route.table, route.entity_column, entity, api_key, DAY),)
|
||||
)
|
||||
filters: Final = {} if route.entity_filter is None else {route.entity_filter: entity}
|
||||
with (
|
||||
daily_rows((user_row(owner, api_key, DAY), *entity_rows)),
|
||||
spend_logs_of_key(api_key, (named_row(0, alias),)),
|
||||
):
|
||||
assert_key_reported(
|
||||
activity_of_key(gateway, route.path, api_key, **filters),
|
||||
api_key,
|
||||
DAY,
|
||||
key_metadata(alias=alias, user=owner, email=email),
|
||||
seeded_metrics(1),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"layout",
|
||||
(
|
||||
pytest.param(_named_between_fifty_and_fifty, id="named_between_50_and_50_nameless"),
|
||||
pytest.param(_oldest_named, id="oldest_named_150_nameless_newer"),
|
||||
pytest.param(_newest_named, id="newest_named_150_nameless_older"),
|
||||
pytest.param(_both_edges_named, id="both_edges_named_150_nameless_between"),
|
||||
pytest.param(_named_after_one_hundred, id="100_nameless_named_99_nameless"),
|
||||
pytest.param(_named_after_ninety_nine, id="99_nameless_named_100_nameless"),
|
||||
),
|
||||
)
|
||||
def test_alias_on_an_edge_of_the_window_is_reported_whatever_surrounds_it(
|
||||
gateway: Gateway, layout: Callable[[str], tuple[SpendLogRow, ...]]
|
||||
) -> None:
|
||||
api_key: Final = digest_no_key_table_holds()
|
||||
alias: Final = _alias()
|
||||
with gateway.scenario() as scenario:
|
||||
owner, email = user_with_an_email(scenario)
|
||||
with daily_rows((user_row(owner, api_key, DAY),)), spend_logs_of_key(api_key, layout(alias)):
|
||||
assert_key_reported(
|
||||
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
|
||||
api_key,
|
||||
DAY,
|
||||
key_metadata(alias=alias, user=owner, email=email),
|
||||
seeded_metrics(1),
|
||||
)
|
||||
|
||||
|
||||
def test_alias_named_only_in_the_middle_of_two_hundred_nameless_rows_is_not_picked_up(gateway: Gateway) -> None:
|
||||
api_key: Final = digest_no_key_table_holds()
|
||||
with gateway.scenario() as scenario:
|
||||
owner, email = user_with_an_email(scenario)
|
||||
with (
|
||||
daily_rows((user_row(owner, api_key, DAY),)),
|
||||
spend_logs_of_key(api_key, _named_only_in_the_middle(_alias())),
|
||||
):
|
||||
assert_key_reported(
|
||||
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
|
||||
api_key,
|
||||
DAY,
|
||||
key_metadata(user=owner, email=email),
|
||||
seeded_metrics(1),
|
||||
)
|
||||
|
||||
|
||||
def test_key_renamed_and_renamed_back_is_reported_with_the_alias_on_both_edges(gateway: Gateway) -> None:
|
||||
api_key: Final = digest_no_key_table_holds()
|
||||
alias: Final = _alias()
|
||||
with gateway.scenario() as scenario:
|
||||
owner, email = user_with_an_email(scenario)
|
||||
rows: Final = _renamed_and_renamed_back(alias, _alias())
|
||||
with daily_rows((user_row(owner, api_key, DAY),)), spend_logs_of_key(api_key, rows):
|
||||
assert_key_reported(
|
||||
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
|
||||
api_key,
|
||||
DAY,
|
||||
key_metadata(alias=alias, user=owner, email=email),
|
||||
seeded_metrics(1),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"spend_log_of_team",
|
||||
(
|
||||
pytest.param(_team_in_the_column, id="team_id_column"),
|
||||
pytest.param(_team_in_the_metadata, id="team_id_in_metadata"),
|
||||
),
|
||||
)
|
||||
def test_team_named_only_by_a_spend_log_is_reported_next_to_the_daily_owner(
|
||||
gateway: Gateway, spend_log_of_team: Callable[[str], SpendLogRow]
|
||||
) -> None:
|
||||
api_key: Final = digest_no_key_table_holds()
|
||||
team: Final = f"integration-team-{uuid.uuid4().hex}"
|
||||
with gateway.scenario() as scenario:
|
||||
owner, email = user_with_an_email(scenario)
|
||||
with daily_rows((user_row(owner, api_key, DAY),)), spend_logs_of_key(api_key, (spend_log_of_team(team),)):
|
||||
assert_key_reported(
|
||||
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
|
||||
api_key,
|
||||
DAY,
|
||||
key_metadata(team=team, user=owner, email=email),
|
||||
seeded_metrics(1),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"spend_log_of_user",
|
||||
(
|
||||
pytest.param(_user_in_the_column, id="user_column"),
|
||||
pytest.param(_user_in_the_metadata, id="user_id_in_metadata"),
|
||||
),
|
||||
)
|
||||
def test_user_named_by_a_spend_log_beats_the_owner_the_daily_rows_name(
|
||||
gateway: Gateway, spend_log_of_user: Callable[[str], SpendLogRow]
|
||||
) -> None:
|
||||
api_key: Final = digest_no_key_table_holds()
|
||||
with gateway.scenario() as scenario:
|
||||
daily_owner, _ = user_with_an_email(scenario)
|
||||
log_user, log_email = user_with_an_email(scenario)
|
||||
with (
|
||||
daily_rows((user_row(daily_owner, api_key, DAY),)),
|
||||
spend_logs_of_key(api_key, (spend_log_of_user(log_user),)),
|
||||
):
|
||||
assert_key_reported(
|
||||
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
|
||||
api_key,
|
||||
DAY,
|
||||
key_metadata(user=log_user, email=log_email),
|
||||
seeded_metrics(1),
|
||||
)
|
||||
|
||||
|
||||
def test_hashed_jwt_digest_is_named_by_its_spend_log(gateway: Gateway) -> None:
|
||||
api_key: Final = f"hashed-jwt-{sha256(uuid.uuid4().bytes).hexdigest()}"
|
||||
alias: Final = _alias()
|
||||
with gateway.scenario() as scenario:
|
||||
owner, email = user_with_an_email(scenario)
|
||||
with daily_rows((user_row(owner, api_key, DAY),)), spend_logs_of_key(api_key, (named_row(0, alias),)):
|
||||
assert_key_reported(
|
||||
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
|
||||
api_key,
|
||||
DAY,
|
||||
key_metadata(alias=alias, user=owner, email=email),
|
||||
seeded_metrics(1),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("started", "inside_the_window"),
|
||||
(
|
||||
pytest.param("2026-02-01 23:59:59", False, id="second_before_the_window"),
|
||||
pytest.param("2026-02-02 00:00:00", True, id="first_second_of_the_window"),
|
||||
pytest.param("2026-02-04 23:59:59", True, id="last_second_of_the_window"),
|
||||
pytest.param("2026-02-05 00:00:00", False, id="first_second_after_the_window"),
|
||||
),
|
||||
)
|
||||
def test_spend_log_names_the_key_only_from_one_day_before_to_two_days_after_the_read(
|
||||
gateway: Gateway, started: str, inside_the_window: bool
|
||||
) -> None:
|
||||
api_key: Final = digest_no_key_table_holds()
|
||||
alias: Final = _alias()
|
||||
row: Final = SpendLogRow(started, {"user_api_key_alias": alias})
|
||||
with gateway.scenario() as scenario:
|
||||
owner, email = user_with_an_email(scenario)
|
||||
with daily_rows((user_row(owner, api_key, DAY),)), spend_logs_of_key(api_key, (row,)):
|
||||
assert_key_reported(
|
||||
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
|
||||
api_key,
|
||||
DAY,
|
||||
key_metadata(alias=alias if inside_the_window else None, user=owner, email=email),
|
||||
seeded_metrics(1),
|
||||
)
|
||||
|
||||
|
||||
def test_two_aliases_on_the_two_edges_leave_the_key_unnamed(gateway: Gateway) -> None:
|
||||
api_key: Final = digest_no_key_table_holds()
|
||||
rows: Final = (named_row(0, _alias()), *nameless_rows(150, 1), named_row(151, _alias()))
|
||||
with gateway.scenario() as scenario:
|
||||
owner, email = user_with_an_email(scenario)
|
||||
with daily_rows((user_row(owner, api_key, DAY),)), spend_logs_of_key(api_key, rows):
|
||||
assert_key_reported(
|
||||
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
|
||||
api_key,
|
||||
DAY,
|
||||
key_metadata(user=owner, email=email),
|
||||
seeded_metrics(1),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"unnamed_rows",
|
||||
(
|
||||
pytest.param((SpendLogRow(started_at(0), {"user_api_key_alias": ""}),), id="empty_string_alias"),
|
||||
pytest.param(
|
||||
(SpendLogRow(started_at(0), ["x"]), SpendLogRow(started_at(1), "x")), id="array_then_string_metadata"
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_rows_without_a_usable_alias_do_not_hide_the_named_row_after_them(
|
||||
gateway: Gateway, unnamed_rows: tuple[SpendLogRow, ...]
|
||||
) -> None:
|
||||
api_key: Final = digest_no_key_table_holds()
|
||||
alias: Final = _alias()
|
||||
rows: Final = (*unnamed_rows, named_row(len(unnamed_rows), alias))
|
||||
with gateway.scenario() as scenario:
|
||||
owner, email = user_with_an_email(scenario)
|
||||
with daily_rows((user_row(owner, api_key, DAY),)), spend_logs_of_key(api_key, rows):
|
||||
assert_key_reported(
|
||||
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
|
||||
api_key,
|
||||
DAY,
|
||||
key_metadata(alias=alias, user=owner, email=email),
|
||||
seeded_metrics(1),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"stored_alias",
|
||||
(
|
||||
pytest.param(123, id="json_int"),
|
||||
pytest.param(["a"], id="json_list"),
|
||||
pytest.param("a" * 5000, id="five_kb_string"),
|
||||
),
|
||||
)
|
||||
def test_alias_of_an_unexpected_shape_is_reported_as_postgres_renders_it(
|
||||
gateway: Gateway, stored_alias: JsonValue
|
||||
) -> None:
|
||||
api_key: Final = digest_no_key_table_holds()
|
||||
row: Final = SpendLogRow(started_at(0), {"user_api_key_alias": stored_alias})
|
||||
with gateway.scenario() as scenario:
|
||||
owner, email = user_with_an_email(scenario)
|
||||
with daily_rows((user_row(owner, api_key, DAY),)), spend_logs_of_key(api_key, (row,)) as request_ids:
|
||||
rendered: Final = read_rows(ALIAS_OF_ONE_SPEND_LOG, (request_ids[0],))[0]["alias"]
|
||||
assert isinstance(rendered, str) and rendered, rendered
|
||||
assert_key_reported(
|
||||
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
|
||||
api_key,
|
||||
DAY,
|
||||
key_metadata(alias=rendered, user=owner, email=email),
|
||||
seeded_metrics(1),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.timeout(300)
|
||||
def test_alias_found_once_is_served_from_the_cache_for_the_same_window_only(gateway: Gateway, tmp_path: Path) -> None:
|
||||
api_key: Final = digest_no_key_table_holds()
|
||||
alias: Final = _alias()
|
||||
with scratch_database() as database_url, _proxy_on(gateway, tmp_path, database_url) as owned:
|
||||
owner, email = _owner_on(owned.gateway)
|
||||
rows: Final = (user_row(owner, api_key, DAY), user_row(owner, api_key, DAY_OUTSIDE_THE_WINDOW))
|
||||
with daily_rows(rows, database_url=database_url):
|
||||
with spend_logs_of_key(api_key, (named_row(0, alias),), database_url=database_url):
|
||||
first: Final = activity_of_key(owned.gateway, AGGREGATED_USER_ACTIVITY, api_key)
|
||||
cached: Final = activity_of_key(owned.gateway, AGGREGATED_USER_ACTIVITY, api_key)
|
||||
other_window: Final = owned.gateway.request(
|
||||
"GET",
|
||||
AGGREGATED_USER_ACTIVITY,
|
||||
params={"start_date": DAY_OUTSIDE_THE_WINDOW, "end_date": DAY_OUTSIDE_THE_WINDOW, "api_key": api_key},
|
||||
)
|
||||
named: Final = key_metadata(alias=alias, user=owner, email=email)
|
||||
assert_key_reported(first, api_key, DAY, named, seeded_metrics(1))
|
||||
assert_key_reported(cached, api_key, DAY, named, seeded_metrics(1))
|
||||
assert_key_reported(
|
||||
other_window, api_key, DAY_OUTSIDE_THE_WINDOW, key_metadata(user=owner, email=email), seeded_metrics(1)
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.timeout(300)
|
||||
def test_alias_logged_after_a_cached_miss_shows_once_the_miss_expires(gateway: Gateway, tmp_path: Path) -> None:
|
||||
api_key: Final = digest_no_key_table_holds()
|
||||
alias: Final = _alias()
|
||||
with scratch_database() as database_url, _proxy_on(gateway, tmp_path, database_url) as owned:
|
||||
owner, email = _owner_on(owned.gateway)
|
||||
with daily_rows((user_row(owner, api_key, DAY),), database_url=database_url):
|
||||
missed: Final = activity_of_key(owned.gateway, AGGREGATED_USER_ACTIVITY, api_key)
|
||||
with spend_logs_of_key(api_key, (named_row(0, alias),), database_url=database_url):
|
||||
named: Final = eventually(
|
||||
lambda: activity_of_key(owned.gateway, AGGREGATED_USER_ACTIVITY, api_key),
|
||||
_names_the_key(api_key, alias),
|
||||
seconds=CACHED_MISS_CLEARS_WITHIN_SECONDS,
|
||||
)
|
||||
assert_key_reported(missed, api_key, DAY, key_metadata(user=owner, email=email), seeded_metrics(1))
|
||||
assert_key_reported(named, api_key, DAY, key_metadata(alias=alias, user=owner, email=email), seeded_metrics(1))
|
||||
|
||||
|
||||
@pytest.mark.timeout(300)
|
||||
def test_alias_lookup_gives_up_while_spend_logs_are_locked_and_answers_once_they_are_not(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
api_key: Final = digest_no_key_table_holds()
|
||||
alias: Final = _alias()
|
||||
with scratch_database() as database_url, _proxy_on(gateway, tmp_path, database_url, workers=2) as owned:
|
||||
owner, email = _owner_on(owned.gateway)
|
||||
with (
|
||||
daily_rows((user_row(owner, api_key, DAY),), database_url=database_url),
|
||||
spend_logs_of_key(api_key, (named_row(0, alias),), database_url=database_url),
|
||||
):
|
||||
with locked_table(SPEND_LOGS_TABLE, database_url=database_url):
|
||||
started: Final = time.monotonic()
|
||||
locked: Final = activity_of_key(owned.gateway, AGGREGATED_USER_ACTIVITY, api_key)
|
||||
waited: Final = time.monotonic() - started
|
||||
unlocked: Final = eventually(
|
||||
lambda: activity_of_key(owned.gateway, AGGREGATED_USER_ACTIVITY, api_key),
|
||||
_names_the_key(api_key, alias),
|
||||
seconds=CACHED_MISS_CLEARS_WITHIN_SECONDS,
|
||||
)
|
||||
assert waited < GIVES_UP_WITHIN_SECONDS, waited
|
||||
assert_key_reported(locked, api_key, DAY, key_metadata(user=owner, email=email), seeded_metrics(1))
|
||||
assert_key_reported(unlocked, api_key, DAY, key_metadata(alias=alias, user=owner, email=email), seeded_metrics(1))
|
||||
|
||||
|
||||
def test_concurrent_reads_over_every_route_all_name_a_fresh_key(gateway: Gateway) -> None:
|
||||
api_key: Final = digest_no_key_table_holds()
|
||||
alias: Final = _alias()
|
||||
entity: Final = f"integration-entity-{uuid.uuid4().hex}"
|
||||
entity_columns: Final = {route.table: route.entity_column for route in ROUTES if route.table != USER_SPEND}
|
||||
with gateway.scenario() as scenario:
|
||||
owner, email = user_with_an_email(scenario)
|
||||
rows: Final = (
|
||||
user_row(owner, api_key, DAY),
|
||||
*(seeded_row(table, column, entity, api_key, DAY) for table, column in entity_columns.items()),
|
||||
)
|
||||
with (
|
||||
daily_rows(rows),
|
||||
spend_logs_of_key(api_key, (named_row(0, alias),)),
|
||||
ThreadPoolExecutor(CONCURRENT_READS) as pool,
|
||||
):
|
||||
reads: Final = tuple(
|
||||
pool.submit(_activity_on_route, gateway, ROUTES[index % len(ROUTES)], api_key, entity)
|
||||
for index in range(CONCURRENT_READS)
|
||||
)
|
||||
responses: Final = tuple(read.result() for read in reads)
|
||||
for response in responses:
|
||||
assert_key_reported(
|
||||
response, api_key, DAY, key_metadata(alias=alias, user=owner, email=email), seeded_metrics(1)
|
||||
)
|
||||
|
|
@ -12,7 +12,7 @@ from typing import Final
|
|||
|
||||
import httpx
|
||||
import pytest
|
||||
from integration._support.client import Gateway, Scenario, eventually
|
||||
from integration._support.client import Gateway, Scenario, eventually, string_value
|
||||
from integration._support.daily_activity import (
|
||||
AGGREGATED_USER_ACTIVITY,
|
||||
DAY,
|
||||
|
|
@ -25,6 +25,7 @@ from integration._support.daily_activity import (
|
|||
daily_rows,
|
||||
key_metadata,
|
||||
key_no_key_table_holds,
|
||||
purge_key_from_the_key_tables,
|
||||
seeded_metrics,
|
||||
seeded_row,
|
||||
user_row,
|
||||
|
|
@ -42,6 +43,10 @@ REQUESTS_OF_KEY: Final = (
|
|||
'SELECT COALESCE(SUM(api_requests), 0)::int AS requests FROM "LiteLLM_DailyUserSpend" '
|
||||
"WHERE api_key=%s AND user_id=%s"
|
||||
)
|
||||
NAMED_SPEND_LOGS_OF_KEY: Final = (
|
||||
'SELECT COUNT(*)::int AS named FROM "LiteLLM_SpendLogs" '
|
||||
"WHERE api_key=%s AND NULLIF(metadata->>'user_api_key_alias', '') IS NOT NULL"
|
||||
)
|
||||
UNIFIED_ENDPOINTS: Final = ("/v1/chat/completions", "/v1/messages", "/v1/responses")
|
||||
REQUESTS_OF_A_BURST: Final = 21
|
||||
READS_DURING_A_BURST: Final = 30
|
||||
|
|
@ -210,6 +215,14 @@ def _wait_for_requests(api_key: str, user: str, requests: int) -> None:
|
|||
)
|
||||
|
||||
|
||||
def _wait_for_named_spend_logs(api_key: str, requests: int) -> None:
|
||||
eventually(
|
||||
lambda: read_rows(NAMED_SPEND_LOGS_OF_KEY, (api_key,)),
|
||||
lambda rows: rows[0]["named"] == requests,
|
||||
seconds=70,
|
||||
)
|
||||
|
||||
|
||||
def _cli_session_token(user: str, team: str) -> str:
|
||||
cli_user: Final = LiteLLM_UserTable(user_id=user, user_role="internal_user", teams=[team], models=[])
|
||||
return ExperimentalUIJWTToken.get_cli_jwt_auth_token(user_info=cli_user, team_id=team, team_alias="cli-team")
|
||||
|
|
@ -244,6 +257,39 @@ def test_key_used_on_every_unified_endpoint_is_reported_with_its_own_alias_and_u
|
|||
)
|
||||
|
||||
|
||||
def test_key_purged_from_the_key_tables_is_reported_with_the_alias_its_spend_logs_name(gateway: Gateway) -> None:
|
||||
prompts: Final = (_prompt(), _prompt(), _prompt())
|
||||
with wire_server(_provider) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _priced_model(scenario, wire.url)
|
||||
owner, email = user_with_an_email(scenario)
|
||||
alias: Final = f"integration-alias-{uuid.uuid4().hex}"
|
||||
generated: Final = gateway.post("/key/generate", {"user_id": owner, "key_alias": alias, "models": [model]})
|
||||
key: Final = string_value(generated["key"])
|
||||
stored: Final = sha256(key.encode()).hexdigest()
|
||||
try:
|
||||
answers: Final = tuple(
|
||||
gateway.request("POST", endpoint, _request_body(endpoint, model, prompt), key=key)
|
||||
for endpoint, prompt in zip(UNIFIED_ENDPOINTS, prompts, strict=True)
|
||||
)
|
||||
assert [answer.status_code for answer in answers] == [200, 200, 200], [answer.text for answer in answers]
|
||||
received: Final = _sent_for_callers(wire.drain())
|
||||
assert [request.target for request in received] == [
|
||||
"/v1/chat/completions",
|
||||
"/v1/responses",
|
||||
"/v1/responses",
|
||||
]
|
||||
_wait_for_requests(stored, owner, 3)
|
||||
_wait_for_named_spend_logs(stored, 3)
|
||||
finally:
|
||||
purge_key_from_the_key_tables(stored)
|
||||
assert_key_owner_and_totals(
|
||||
_activity_around_today(gateway, stored),
|
||||
stored,
|
||||
key_metadata(alias=alias, user=owner, email=email, exists=False),
|
||||
_totals_of_requests(3),
|
||||
)
|
||||
|
||||
|
||||
def test_cli_session_spend_is_reported_with_the_user_and_team_of_the_session(
|
||||
gateway: Gateway, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
|
|
|
|||
|
|
@ -690,7 +690,7 @@ def test_langfuse_logging_tool_calling():
|
|||
]
|
||||
|
||||
response = litellm.completion(
|
||||
model="gpt-3.5-turbo-1106",
|
||||
model="gpt-6-luna",
|
||||
messages=messages,
|
||||
tools=tools,
|
||||
tool_choice="auto", # auto is default, but we'll be explicit
|
||||
|
|
@ -698,6 +698,8 @@ def test_langfuse_logging_tool_calling():
|
|||
print("\nLLM Response1:\n", response)
|
||||
response_message = response.choices[0].message
|
||||
tool_calls = response.choices[0].message.tool_calls
|
||||
assert response.choices[0].message.tool_calls
|
||||
assert all(call.function.name == "get_current_weather" for call in response.choices[0].message.tool_calls)
|
||||
|
||||
|
||||
# test_langfuse_logging_tool_calling()
|
||||
|
|
|
|||
|
|
@ -39,7 +39,7 @@ def get_current_weather(location, unit="fahrenheit"):
|
|||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
"gpt-3.5-turbo-1106",
|
||||
"gpt-6-luna",
|
||||
"mistral/mistral-large-latest",
|
||||
"claude-haiku-4-5-20251001",
|
||||
"gemini/gemini-2.5-flash-lite",
|
||||
|
|
@ -386,7 +386,7 @@ def test_parallel_function_call_stream():
|
|||
}
|
||||
]
|
||||
response = litellm.completion(
|
||||
model="gpt-3.5-turbo-1106",
|
||||
model="gpt-6-luna",
|
||||
messages=messages,
|
||||
tools=tools,
|
||||
stream=True,
|
||||
|
|
@ -435,7 +435,7 @@ def test_parallel_function_call_stream():
|
|||
) # extend conversation with function response
|
||||
print(f"messages: {messages}")
|
||||
second_response = litellm.completion(
|
||||
model="gpt-3.5-turbo-1106", messages=messages, temperature=0.2, seed=22
|
||||
model="gpt-6-luna", messages=messages, temperature=0.2, seed=22, reasoning_effort="none"
|
||||
) # get a new response from the model where it can see the function response
|
||||
print("second response\n", second_response)
|
||||
return second_response
|
||||
|
|
|
|||
|
|
@ -1,14 +1,18 @@
|
|||
# What is this?
|
||||
## Unit testing for the 'get_model_info()' function
|
||||
import os
|
||||
import re
|
||||
from collections.abc import Collection, Mapping
|
||||
|
||||
|
||||
from typing import List, Dict, Any
|
||||
from typing import List, Dict, Any, Final, Literal
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm import get_model_info
|
||||
from litellm.llms.bedrock.common_utils import BedrockModelInfo
|
||||
from litellm.types.utils import ModelInfoBase
|
||||
from litellm.utils import _invalidate_model_cost_lowercase_map
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
|
|
@ -116,26 +120,31 @@ def test_get_model_info_ft_model_with_provider_prefix():
|
|||
|
||||
|
||||
def _enforce_bedrock_converse_models(
|
||||
model_cost: List[Dict[str, Any]], whitelist_models: List[str]
|
||||
):
|
||||
model_cost: Mapping[str, ModelInfoBase], whitelist_models: Collection[str]
|
||||
) -> None:
|
||||
"""
|
||||
Assert all new bedrock chat models are added as `bedrock_converse` unless explicitly whitelisted.
|
||||
Assert unlisted Bedrock chat models declare or inherit Converse routing.
|
||||
"""
|
||||
# Check for unwhitelisted models
|
||||
for model, info in litellm.model_cost.items():
|
||||
for model, info in model_cost.items():
|
||||
if (
|
||||
info["litellm_provider"] == "bedrock"
|
||||
and info["mode"] == "chat"
|
||||
and model not in whitelist_models
|
||||
and not (
|
||||
(base_model := BedrockModelInfo.get_base_model(model)) != model
|
||||
and model_cost.get(base_model, {}).get("litellm_provider") == "bedrock_converse"
|
||||
and BedrockModelInfo.get_bedrock_route(model) == "converse"
|
||||
)
|
||||
):
|
||||
raise AssertionError(
|
||||
f"New bedrock chat model detected: {model}. Please set `litellm_provider='bedrock_converse'` for this model."
|
||||
f"Unlisted Bedrock chat model does not route to Converse: {model}"
|
||||
)
|
||||
|
||||
|
||||
def test_model_info_bedrock_converse(monkeypatch):
|
||||
"""
|
||||
Assert all new bedrock chat models are added as `bedrock_converse` unless explicitly whitelisted.
|
||||
Assert unlisted Bedrock chat models declare or inherit Converse routing.
|
||||
|
||||
This ensures they are automatically routed to the converse endpoint.
|
||||
"""
|
||||
|
|
@ -173,7 +182,7 @@ def test_model_info_bedrock_converse_enforcement(monkeypatch):
|
|||
whitelist_models = [line.strip() for line in file.readlines()]
|
||||
|
||||
# Check for unwhitelisted models
|
||||
with pytest.raises(AssertionError):
|
||||
with pytest.raises(AssertionError, match=r"fake\.bedrock-chat-model"):
|
||||
_enforce_bedrock_converse_models(
|
||||
model_cost=litellm.model_cost, whitelist_models=whitelist_models
|
||||
)
|
||||
|
|
@ -181,6 +190,27 @@ def test_model_info_bedrock_converse_enforcement(monkeypatch):
|
|||
pytest.skip("whitelisted_bedrock_models.txt not found")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("region", ("us-gov-east-1", "us-gov-west-1"))
|
||||
@pytest.mark.parametrize("base_provider", ("bedrock_converse", "bedrock"))
|
||||
def test_regional_bedrock_alias_requires_canonical_converse_metadata(
|
||||
region: str, base_provider: Literal["bedrock_converse", "bedrock"]
|
||||
) -> None:
|
||||
base_model: Final = next(
|
||||
model for model in sorted(litellm.bedrock_converse_models) if BedrockModelInfo.get_base_model(model) == model
|
||||
)
|
||||
model: Final = f"bedrock/{region}/{base_model}"
|
||||
model_cost: Final[Mapping[str, ModelInfoBase]] = {
|
||||
model: {"litellm_provider": "bedrock", "mode": "chat"},
|
||||
base_model: {"litellm_provider": base_provider, "mode": "chat"},
|
||||
}
|
||||
assert BedrockModelInfo.get_bedrock_route(model) == "converse"
|
||||
if base_provider == "bedrock":
|
||||
with pytest.raises(AssertionError, match=re.escape(model)):
|
||||
_enforce_bedrock_converse_models(model_cost, ())
|
||||
return
|
||||
_enforce_bedrock_converse_models(model_cost, ())
|
||||
|
||||
|
||||
def test_get_model_info_custom_provider():
|
||||
# Custom provider example copied from https://docs.litellm.ai/docs/providers/custom_llm_server:
|
||||
import litellm
|
||||
|
|
|
|||
|
|
@ -83,13 +83,15 @@ def test_lunary_with_tools():
|
|||
]
|
||||
|
||||
response = litellm.completion(
|
||||
model="gpt-3.5-turbo-1106",
|
||||
model="gpt-6-luna",
|
||||
messages=messages,
|
||||
tools=tools,
|
||||
tool_choice="auto", # auto is default, but we'll be explicit
|
||||
)
|
||||
|
||||
response_message = response.choices[0].message
|
||||
assert response.choices[0].message.tool_calls
|
||||
assert all(call.function.name == "get_current_weather" for call in response.choices[0].message.tool_calls)
|
||||
print("\nLLM Response:\n", response.choices[0].message)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -157,6 +157,7 @@ async def test_create_mcp_server_direct():
|
|||
# Mock server manager
|
||||
mock_manager.add_server = mock.AsyncMock()
|
||||
mock_manager.reload_servers_from_database = mock.AsyncMock()
|
||||
mock_manager.get_mcp_server_by_id.return_value = None
|
||||
|
||||
# Set up test data
|
||||
server_id = str(uuid.uuid4())
|
||||
|
|
|
|||
|
|
@ -0,0 +1,559 @@
|
|||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._experimental.mcp_server import mcp_server_manager
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_UserTable, UserAPIKeyAuth
|
||||
from litellm.proxy.auth import auth_checks
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import ManagedAgentContext
|
||||
|
||||
|
||||
def actor(tools: tuple[str, ...] | None, *, delegated: bool = False) -> UserAPIKeyAuth:
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="agent-permissions",
|
||||
mcp_servers=["slack", "linear"],
|
||||
mcp_tool_permissions={"slack": list(tools)} if tools is not None else None,
|
||||
)
|
||||
agent: Final = AgentResponse(
|
||||
agent_id="publisher",
|
||||
agent_name="Publisher",
|
||||
agent_card_params={},
|
||||
object_permission=permission.model_dump(),
|
||||
identity_managed=True,
|
||||
)
|
||||
auth: Final = UserAPIKeyAuth(agent_id=agent.agent_id)
|
||||
auth.managed_agent_policy = agent
|
||||
auth.managed_agent_context = ManagedAgentContext(
|
||||
agent_id=agent.agent_id,
|
||||
mode="delegated" if delegated else "autonomous",
|
||||
user_id="human" if delegated else None,
|
||||
)
|
||||
return auth
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def isolated_manager(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", mcp_server_manager.MCPServerManager())
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", MagicMock())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("tools", (None, (), ("read",), ("read", "write")))
|
||||
async def test_autonomous_agent_uses_only_its_own_tool_grants(tools: tuple[str, ...] | None) -> None:
|
||||
auth: Final = actor(tools)
|
||||
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"slack", "linear"}
|
||||
actual: Final = await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)
|
||||
assert (frozenset(actual) if actual is not None else None) == (frozenset(tools) if tools is not None else None)
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("ungranted-server", auth) == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"agent_tools,user_tools,expected",
|
||||
(
|
||||
(None, ("read",), ("read",)),
|
||||
(("read",), None, ("read",)),
|
||||
(("read", "write"), ("read",), ("read",)),
|
||||
(("read",), ("write",), ()),
|
||||
((), None, ()),
|
||||
),
|
||||
)
|
||||
async def test_delegated_server_and_tool_intersections(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
agent_tools: tuple[str, ...] | None,
|
||||
user_tools: tuple[str, ...] | None,
|
||||
expected: tuple[str, ...],
|
||||
) -> None:
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="user-permissions",
|
||||
mcp_servers=["slack", "user-only"],
|
||||
mcp_tool_permissions={"slack": list(user_tools)} if user_tools is not None else None,
|
||||
)
|
||||
user: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=permission)
|
||||
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=user))
|
||||
auth: Final = actor(agent_tools, delegated=True)
|
||||
assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == ["slack"]
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == list(expected)
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == []
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("user-only", auth) == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unavailable_delegated_user_never_leaves_agent_permissions_unrestricted(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(side_effect=RuntimeError("DB unavailable")))
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await MCPRequestHandler.get_allowed_tools_for_server("slack", actor(None, delegated=True))
|
||||
assert failure.value.status_code == 503
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("servers,expected", (((), ()), (("slack",), ("slack",)), (("user-only",), ())))
|
||||
async def test_access_groups_cap_agent_servers_without_granting_new_ones(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
servers: tuple[str, ...],
|
||||
expected: tuple[str, ...],
|
||||
) -> None:
|
||||
from litellm.proxy._types import LiteLLM_AccessGroupTable
|
||||
|
||||
group: Final = LiteLLM_AccessGroupTable(
|
||||
access_group_id="group", access_group_name="Restricted", access_mcp_server_ids=list(servers)
|
||||
)
|
||||
monkeypatch.setattr(auth_checks, "get_access_object", AsyncMock(return_value=group))
|
||||
auth: Final = actor(None)
|
||||
assert auth.managed_agent_policy is not None
|
||||
auth.managed_agent_policy = auth.managed_agent_policy.model_copy(update={"access_group_ids": ["group"]})
|
||||
assert tuple(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == expected
|
||||
if "slack" not in expected:
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("change", ["tools", "servers", "disabled", "outage"])
|
||||
async def test_delegated_mcp_revokes_warm_human_policy_before_tool_execution(
|
||||
monkeypatch: pytest.MonkeyPatch, change: str
|
||||
) -> None:
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key
|
||||
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="user-grant", mcp_servers=["slack"], mcp_tool_permissions={"slack": ["read", "write"]}
|
||||
)
|
||||
user: Final = LiteLLM_UserTable(
|
||||
user_id="human", teams=[], organization_memberships=[], object_permission_id="user-grant"
|
||||
)
|
||||
cache: Final = UserApiKeyCache()
|
||||
cache.set_cache("human", user)
|
||||
cache.set_cache(object_permission_cache_key("user-grant"), permission)
|
||||
client: Final = MagicMock()
|
||||
client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=user)
|
||||
client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", client)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
|
||||
auth: Final = actor(("read", "write"), delegated=True)
|
||||
assert set(await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)) == {"read", "write"}
|
||||
if change == "disabled":
|
||||
client.writer_db.litellm_usertable.find_unique.return_value = user.model_copy(
|
||||
update={"metadata": {"scim_active": False}}
|
||||
)
|
||||
elif change == "outage":
|
||||
client.writer_db.litellm_usertable.find_unique.side_effect = RuntimeError("writer unavailable")
|
||||
elif change == "servers":
|
||||
client.writer_db.litellm_objectpermissiontable.find_unique.return_value = permission.model_copy(
|
||||
update={"mcp_servers": [], "mcp_tool_permissions": {}}
|
||||
)
|
||||
else:
|
||||
client.writer_db.litellm_objectpermissiontable.find_unique.return_value = permission.model_copy(
|
||||
update={"mcp_tool_permissions": {"slack": ["read"]}}
|
||||
)
|
||||
if change in ("disabled", "outage"):
|
||||
with pytest.raises(HTTPException):
|
||||
await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)
|
||||
else:
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == (
|
||||
["read"] if change == "tools" else []
|
||||
)
|
||||
client.db.litellm_usertable.find_unique.assert_not_called()
|
||||
client.db.litellm_objectpermissiontable.find_unique.assert_not_called()
|
||||
|
||||
|
||||
def _server_row(server_id: str, access_groups: tuple[str, ...]) -> MagicMock:
|
||||
row: Final = MagicMock()
|
||||
row.server_id = server_id
|
||||
row.mcp_access_groups = list(access_groups)
|
||||
return row
|
||||
|
||||
|
||||
def _toolset_row(server_id: str, tool_name: str) -> MagicMock:
|
||||
row: Final = MagicMock()
|
||||
row.tools = [{"server_id": server_id, "tool_name": tool_name}]
|
||||
return row
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("change", ["tool", "server", "outage"])
|
||||
async def test_autonomous_agent_toolset_and_access_group_revocations_bind_on_the_next_request(
|
||||
monkeypatch: pytest.MonkeyPatch, change: str
|
||||
) -> None:
|
||||
"""The agent's entitlements are read through the shared toolset and access-group resolvers. Once the
|
||||
writer revokes a tool or drops the server from the group, the next managed request must be denied
|
||||
even though the legacy cache still holds the warm grant and the replica still shows the old rows"""
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._experimental.mcp_server import toolset_db
|
||||
|
||||
warm_toolset: Final = _toolset_row("slack", "read")
|
||||
list_toolsets: Final = AsyncMock(return_value=[warm_toolset])
|
||||
monkeypatch.setattr(toolset_db, "list_mcp_toolsets", list_toolsets)
|
||||
client: Final = MagicMock()
|
||||
client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[_server_row("linear", ("grp",))])
|
||||
client.writer_db.litellm_mcpservertable.find_many = AsyncMock(return_value=[_server_row("linear", ("grp",))])
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", client)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", DualCache())
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="agent-permissions", mcp_toolsets=["ts"], mcp_access_groups=["grp"]
|
||||
)
|
||||
auth: Final = actor(None)
|
||||
assert auth.managed_agent_policy is not None
|
||||
auth.managed_agent_policy = auth.managed_agent_policy.model_copy(
|
||||
update={"object_permission": permission.model_dump()}
|
||||
)
|
||||
auth.requires_fresh_policy = True
|
||||
|
||||
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"slack", "linear"}
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == ["read"]
|
||||
|
||||
if change == "tool":
|
||||
list_toolsets.return_value = [_toolset_row("slack", "other")]
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == ["other"]
|
||||
elif change == "server":
|
||||
client.writer_db.litellm_mcpservertable.find_many.return_value = []
|
||||
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"slack"}
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == []
|
||||
else:
|
||||
list_toolsets.side_effect = RuntimeError("writer unavailable")
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)
|
||||
assert failure.value.status_code == 503
|
||||
for call in list_toolsets.await_args_list:
|
||||
assert call.kwargs["use_writer"] is True, "managed agent toolsets must be read from the writer"
|
||||
client.db.litellm_mcpservertable.find_many.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("role", ["proxy_admin", "proxy_admin_viewer", "internal_user"])
|
||||
@pytest.mark.parametrize("open_channel", ["none", "operator", "submitted"])
|
||||
@pytest.mark.parametrize("has_grant", [True, False])
|
||||
@pytest.mark.parametrize("agent_tools", [("read", "write"), None])
|
||||
async def test_delegated_mcp_uses_explicit_team_grants_even_for_dashboard_admins(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
role: str,
|
||||
open_channel: str,
|
||||
has_grant: bool,
|
||||
agent_tools: tuple[str, ...] | None,
|
||||
) -> None:
|
||||
from litellm.proxy._types import LiteLLM_TeamTable
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
manager: Final = mcp_server_manager.global_mcp_server_manager
|
||||
manager.registry = {
|
||||
name: MCPServer(
|
||||
server_id=name,
|
||||
name=name,
|
||||
transport="http",
|
||||
url="https://example.com/mcp",
|
||||
allow_all_keys=open_channel == "operator",
|
||||
)
|
||||
for name in ("slack", "linear")
|
||||
}
|
||||
from litellm.proxy._experimental.mcp_server import db
|
||||
|
||||
monkeypatch.setattr(
|
||||
db,
|
||||
"get_active_submitted_mcp_server_ids_for_user",
|
||||
AsyncMock(return_value=["slack", "linear"] if open_channel == "submitted" else []),
|
||||
)
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="team-grant", mcp_servers=["slack"], mcp_tool_permissions={"slack": ["read"]}
|
||||
)
|
||||
user: Final = LiteLLM_UserTable(
|
||||
user_id="human", user_role=role, teams=["team"] if has_grant else [], organization_memberships=[]
|
||||
)
|
||||
team: Final = LiteLLM_TeamTable(
|
||||
team_id="team",
|
||||
models=[],
|
||||
members_with_roles=[{"user_id": "human", "role": "user"}],
|
||||
object_permission_id="team-grant",
|
||||
)
|
||||
client: Final = MagicMock()
|
||||
client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=user)
|
||||
client.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=team)
|
||||
client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", client)
|
||||
auth: Final = actor(agent_tools, delegated=True)
|
||||
auth.team_id = "team"
|
||||
assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == (["slack"] if has_grant else [])
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == (["read"] if has_grant else [])
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == []
|
||||
admitted: Final = await MCPRequestHandler.reload_admitted_user("human", requires_fresh_policy=True)
|
||||
assert admitted.user_role == role
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_explicit_grants_never_fall_back_to_open_servers_on_resolution_failure(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from litellm.proxy._experimental.mcp_server import db
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
manager: Final = mcp_server_manager.global_mcp_server_manager
|
||||
manager.registry = {"slack": MCPServer(server_id="slack", name="slack", transport="http", allow_all_keys=True)}
|
||||
monkeypatch.setattr(db, "get_active_submitted_mcp_server_ids_for_user", AsyncMock(return_value=["slack"]))
|
||||
auth: Final = UserAPIKeyAuth(user_id="human")
|
||||
auth.mcp_explicit_grants_only = True
|
||||
with pytest.MonkeyPatch.context() as patcher:
|
||||
patcher.setattr(MCPRequestHandler, "get_mcp_server_access", AsyncMock(side_effect=RuntimeError("unavailable")))
|
||||
assert await manager.get_allowed_mcp_servers(auth) == []
|
||||
auth.mcp_explicit_grants_only = False
|
||||
assert await manager.get_allowed_mcp_servers(auth) == ["slack"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_absent_agent_policy_and_missing_delegated_subject_grant_no_servers() -> None:
|
||||
from litellm.proxy._experimental.mcp_server.auth.managed_agent_access import managed_agent_servers
|
||||
|
||||
assert await managed_agent_servers(UserAPIKeyAuth()) == ()
|
||||
auth: Final = actor(None, delegated=True)
|
||||
assert auth.managed_agent_context is not None
|
||||
auth.managed_agent_context = auth.managed_agent_context.model_copy(update={"user_id": None})
|
||||
assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == []
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_policy_outage_after_server_admission_fails_closed(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="human-grant", mcp_servers=["slack"])
|
||||
user: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=permission)
|
||||
monkeypatch.setattr(
|
||||
auth_checks, "get_user_object", AsyncMock(side_effect=[user, RuntimeError("tool lookup unavailable")])
|
||||
)
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await MCPRequestHandler.get_allowed_tools_for_server("slack", actor(None, delegated=True))
|
||||
assert failure.value.status_code == 503
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("role", (None, "proxy_admin", "internal_user"))
|
||||
@pytest.mark.parametrize("scoped", (False, True))
|
||||
async def test_manager_preserves_managed_server_grants_across_open_channels(
|
||||
monkeypatch: pytest.MonkeyPatch, role: str | None, scoped: bool
|
||||
) -> None:
|
||||
from litellm.proxy._experimental.mcp_server import db
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPServerAccess
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
manager: Final = mcp_server_manager.global_mcp_server_manager
|
||||
manager.registry = {
|
||||
"open": MCPServer(server_id="open", name="open", transport="http", allow_all_keys=True),
|
||||
"submitted": MCPServer(server_id="submitted", name="submitted", transport="http"),
|
||||
"passthrough": MCPServer(
|
||||
server_id="passthrough", name="passthrough", transport="http", auth_type="true_passthrough"
|
||||
),
|
||||
}
|
||||
monkeypatch.setattr(db, "get_active_submitted_mcp_server_ids_for_user", AsyncMock(return_value=["submitted"]))
|
||||
auth: Final = actor(None)
|
||||
auth.user_role = role
|
||||
assert not auth.mcp_explicit_grants_only
|
||||
access: Final = MCPServerAccess(server_ids=("slack", "open")) if scoped else None
|
||||
assert set(await manager.get_allowed_mcp_servers(auth, access=access)) == (
|
||||
{"slack"} if scoped else {"slack", "linear"}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_manager_does_not_replace_managed_policy_failure_with_open_servers(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
manager: Final = mcp_server_manager.global_mcp_server_manager
|
||||
manager.registry = {"open": MCPServer(server_id="open", name="open", transport="http", allow_all_keys=True)}
|
||||
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(side_effect=RuntimeError("writer unavailable")))
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await manager.get_allowed_mcp_servers(actor(None, delegated=True))
|
||||
assert failure.value.status_code == 503
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_inline_tool_grant_admits_its_server_without_widening_tools() -> None:
|
||||
auth: Final = actor(("read",))
|
||||
assert auth.managed_agent_policy is not None
|
||||
auth.managed_agent_policy = auth.managed_agent_policy.model_copy(
|
||||
update={"object_permission": {"object_permission_id": "tools", "mcp_tool_permissions": {"slack": ["read"]}}}
|
||||
)
|
||||
assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == ["slack"]
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == ["read"]
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("selected_team", (None, "selected"))
|
||||
@pytest.mark.parametrize("selected_grant", (False, True))
|
||||
async def test_delegation_never_borrows_another_teams_server_or_tools(
|
||||
monkeypatch: pytest.MonkeyPatch, selected_team: str | None, selected_grant: bool
|
||||
) -> None:
|
||||
from litellm.proxy._types import LiteLLM_TeamTable
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
|
||||
user: Final = LiteLLM_UserTable(user_id="human", teams=["selected", "other"], organization_memberships=[])
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="selected-grant",
|
||||
mcp_servers=["slack"] if selected_grant else [],
|
||||
mcp_tool_permissions={"slack": ["read"]} if selected_grant else {},
|
||||
)
|
||||
teams: Final = {
|
||||
name: LiteLLM_TeamTable(
|
||||
team_id=name,
|
||||
models=[],
|
||||
members_with_roles=[{"user_id": "human", "role": "user"}],
|
||||
object_permission=permission if name == "selected" else LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="other-grant", mcp_servers=["slack", "linear"]
|
||||
),
|
||||
)
|
||||
for name in ("selected", "other")
|
||||
}
|
||||
|
||||
async def get_team(team_id: str, **kwargs: object) -> LiteLLM_TeamTable:
|
||||
return teams[team_id]
|
||||
|
||||
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=user))
|
||||
monkeypatch.setattr(auth_checks, "get_team_object", get_team)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", UserApiKeyCache())
|
||||
auth: Final = actor(None, delegated=True)
|
||||
auth.team_id = selected_team
|
||||
expected: Final = ["slack"] if selected_team and selected_grant else []
|
||||
assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == expected
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == (["read"] if expected else [])
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == []
|
||||
ordinary: Final = await MCPRequestHandler.reload_admitted_user("human", requires_fresh_policy=True)
|
||||
assert set(await MCPRequestHandler.resolve_admitted_subject_servers(ordinary)) == {"slack", "linear"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("entitlement", ("group", "toolset"))
|
||||
async def test_managed_mcp_rejects_unavailable_authoritative_entitlements(
|
||||
monkeypatch: pytest.MonkeyPatch, entitlement: str
|
||||
) -> None:
|
||||
client: Final = MagicMock()
|
||||
client.writer_db.litellm_mcpservertable.find_many = AsyncMock(side_effect=RuntimeError("writer unavailable"))
|
||||
client.writer_db.litellm_mcptoolsettable.find_many = AsyncMock(side_effect=RuntimeError("writer unavailable"))
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", client)
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="entitlements",
|
||||
mcp_access_groups=["group"] if entitlement == "group" else [],
|
||||
mcp_toolsets=["toolset"] if entitlement == "toolset" else [],
|
||||
)
|
||||
auth: Final = actor(None)
|
||||
assert auth.managed_agent_policy is not None
|
||||
auth.managed_agent_policy = auth.managed_agent_policy.model_copy(update={"object_permission": permission.model_dump()})
|
||||
auth.requires_fresh_policy = True
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)
|
||||
assert failure.value.status_code == 503
|
||||
client.db.litellm_mcpservertable.find_many.assert_not_called()
|
||||
client.db.litellm_mcptoolsettable.find_many.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_managed_agent_mcp_access_is_capped_at_the_invoking_callers_grants(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""The managed MCP path must honour the agent_caller ceiling the same way the unmanaged path does:
|
||||
the agent's own policy grants slack and linear, but the team echoed back on the request reaches
|
||||
only slack, so the agent may use slack alone."""
|
||||
from litellm.proxy._types import AgentCaller
|
||||
|
||||
monkeypatch.setattr(
|
||||
MCPRequestHandler,
|
||||
"_get_allowed_mcp_servers_for_team",
|
||||
AsyncMock(return_value=["slack"]),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
MCPRequestHandler,
|
||||
"_apply_user_server_ceiling",
|
||||
AsyncMock(side_effect=lambda servers, _auth: (tuple(servers), False)),
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
MCPRequestHandler,
|
||||
"_get_team_object_permission",
|
||||
AsyncMock(
|
||||
return_value=LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="caller-team-permissions",
|
||||
mcp_servers=["slack"],
|
||||
mcp_tool_permissions={"slack": ["read"]},
|
||||
)
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
MCPRequestHandler,
|
||||
"_apply_user_tool_ceiling",
|
||||
AsyncMock(side_effect=lambda tools, _server_id, _auth: tools),
|
||||
)
|
||||
|
||||
auth: Final = actor(("read", "write"))
|
||||
auth.agent_caller = AgentCaller(user_id="alice", team_id="callers")
|
||||
|
||||
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"slack"}
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == ["read"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("fresh", [False, True])
|
||||
@pytest.mark.parametrize("caller_kind", ["team", "user"])
|
||||
async def test_caller_mcp_revocation_uses_fresh_policy(
|
||||
monkeypatch: pytest.MonkeyPatch, fresh: bool, caller_kind: str,
|
||||
) -> None:
|
||||
from litellm.proxy._types import LiteLLM_TeamTable
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key
|
||||
from litellm.types.agents import AgentCaller
|
||||
|
||||
cached_permission: Final = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="caller-permission", mcp_servers=["slack", "linear"],
|
||||
mcp_tool_permissions={"slack": ["read", "write"]},
|
||||
)
|
||||
current_permission: Final = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="caller-permission", mcp_servers=["slack"],
|
||||
mcp_tool_permissions={"slack": ["read"]},
|
||||
)
|
||||
team: Final = LiteLLM_TeamTable(
|
||||
team_id="caller", object_permission_id="caller-permission", object_permission=current_permission,
|
||||
)
|
||||
user: Final = LiteLLM_UserTable(
|
||||
user_id="caller", teams=[], object_permission_id="caller-permission", object_permission=current_permission,
|
||||
)
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=team)
|
||||
database.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=user)
|
||||
database.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=current_permission)
|
||||
cache: Final = UserApiKeyCache()
|
||||
cache.set_cache("team_id:caller", team.model_copy(update={"object_permission": cached_permission}))
|
||||
cache.set_cache("caller", user.model_copy(update={"object_permission": cached_permission}))
|
||||
cache.set_cache(object_permission_cache_key("caller-permission"), cached_permission)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", database)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
|
||||
auth: Final = actor(("read", "write"))
|
||||
auth.requires_fresh_policy = fresh
|
||||
auth.agent_caller = AgentCaller(team_id="caller") if caller_kind == "team" else AgentCaller(user_id="caller")
|
||||
|
||||
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == ({"slack"} if fresh else {"slack", "linear"})
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == (["read"] if fresh else ["read", "write"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("fresh", [False, True])
|
||||
async def test_caller_team_outage_cannot_remove_authoritative_server_ceiling(
|
||||
monkeypatch: pytest.MonkeyPatch, fresh: bool,
|
||||
) -> None:
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.types.agents import AgentCaller
|
||||
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_teamtable.find_unique = AsyncMock(side_effect=RuntimeError("writer unavailable"))
|
||||
database.db.litellm_teamtable.find_unique = AsyncMock(side_effect=RuntimeError("reader unavailable"))
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", database)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", UserApiKeyCache())
|
||||
auth: Final = actor(("read",))
|
||||
auth.agent_caller = AgentCaller(team_id="caller")
|
||||
auth.requires_fresh_policy = fresh
|
||||
|
||||
if fresh:
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await MCPRequestHandler.get_allowed_mcp_servers(auth)
|
||||
assert failure.value.status_code == 503
|
||||
else:
|
||||
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"slack", "linear"}
|
||||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue