Merge remote-tracking branch 'upstream/main' into litellm_scheduler_remove_admitted_requests

# Conflicts:
#	tests/code_coverage_tests/router_code_coverage.py
This commit is contained in:
RachelHuangZW 2026-09-29 22:51:04 -04:00
commit fe169152a9
204 changed files with 15725 additions and 1585 deletions

View file

@ -89,6 +89,7 @@ legacy_paths() {
proxy-db-auth-checks)
echo tests/unit/proxy/auth/test_auth_checks.py
echo tests/unit/proxy/auth/test_user_api_key_auth.py
echo tests/unit/proxy/test_credential_slot_registry.py
echo tests/unit/proxy/test_deprecated_key_grace_period.py ;;
proxy-db-budgets)
echo tests/unit/proxy/auth/test_default_end_user_budget_simple.py

View file

@ -24,7 +24,7 @@ model_list:
- model_name: sagemaker-completion-model
litellm_params:
model: sagemaker/berri-benchmarking-Llama-2-70b-chat-hf-4
input_cost_per_second: 0.000420
cost_per_second: 0.000420
- model_name: text-embedding-ada-002
litellm_params:
model: azure/azure-embedding-model

View file

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

View file

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

View file

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

View file

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

View file

@ -24,6 +24,6 @@ Keep unary caching independent of stream-only methods. Store streams only after
Test each contract in its owner: storage capabilities in backend tests, envelopes and freshness here, reuse and replay in core, Python callback and fallback behavior at the bridge, and HTTP behavior at the gateway. Run backend contract checks and Python response-codec fixtures before exposing a new backend
`ScopedCache` requires an explicit shared or isolated scope at construction. `CacheOptions` has no default sharing policy. Callers may override policy per invocation without replacing the attached service. Versioned native envelopes reject incompatible API surfaces and versions as misses; this envelope is distinct from the legacy Python response codec
`ScopedCache` requires an explicit shared or isolated scope at construction. Per-call `CachePolicy` controls reads, writes, expiry, and freshness without replacing the attached scope or service. `CacheOptions` binds that policy to an explicit scope for storage requests and has no default sharing policy. Versioned native envelopes reject incompatible API surfaces and versions as misses; this envelope is distinct from the legacy Python response codec
Response storage is not the source of budget or rate-limit coordination dependencies. Keep counters, reservations, and atomic admission operations out of `ResponseCacheService`, including when both services happen to use Redis

View file

@ -17,6 +17,6 @@ pub use exact::{ConnectionProbe, ExactResponseCache};
pub use response::{ResponseCache, ResponseCacheRequest};
pub use service::{
CacheOptions, CacheScope, ResponseCacheConfig, ResponseCacheService, ResponseEnvelope,
ScopedCache,
CacheOptions, CachePolicy, CacheScope, ResponseCacheConfig, ResponseCacheService,
ResponseEnvelope, ScopedCache,
};

View file

@ -73,32 +73,35 @@ pub enum CacheScope {
Isolated(String),
}
#[derive(Clone)]
pub struct CacheOptions {
#[derive(Clone, Copy, Default)]
pub struct CachePolicy {
pub caching: Option<bool>,
pub no_cache: bool,
pub no_store: bool,
pub ttl: Option<Duration>,
pub max_age: Option<Duration>,
}
impl CachePolicy {
pub fn enabled(&self) -> bool {
self.caching != Some(false) && !(self.no_cache && self.no_store)
}
}
#[derive(Clone)]
pub struct CacheOptions {
pub policy: CachePolicy,
pub scope: CacheScope,
}
impl CacheOptions {
pub fn new(scope: CacheScope) -> Self {
Self {
caching: None,
no_cache: false,
no_store: false,
ttl: None,
max_age: None,
policy: CachePolicy::default(),
scope,
}
}
pub fn enabled(&self) -> bool {
self.caching != Some(false) && !(self.no_cache && self.no_store)
}
pub fn request(self, namespace: &str, surface: &str, mut input: Value) -> ResponseCacheRequest {
input.sort_all_objects();
let scope = match self.scope {
@ -128,13 +131,15 @@ impl CacheOptions {
supported_call_type: true,
native_backend: true,
default_on: true,
caching: self.caching,
no_cache: self.no_cache,
no_store: self.no_store,
caching: self.policy.caching,
no_cache: self.policy.no_cache,
no_store: self.policy.no_store,
..Default::default()
},
context: ExactCacheContext { ttl: self.ttl },
max_age: self.max_age,
context: ExactCacheContext {
ttl: self.policy.ttl,
},
max_age: self.policy.max_age,
}
}
}
@ -171,7 +176,10 @@ impl ScopedCache {
Self { service, scope }
}
pub fn options(&self, overrides: Option<CacheOptions>) -> CacheOptions {
overrides.unwrap_or_else(|| CacheOptions::new(self.scope.clone()))
pub fn options(&self, policy: Option<CachePolicy>) -> CacheOptions {
CacheOptions {
policy: policy.unwrap_or_default(),
scope: self.scope.clone(),
}
}
}

View file

@ -129,11 +129,20 @@ async fn isolated_policy_controls_actual_entry_reuse(
#[case] first: &str,
#[case] second: &str,
#[case] hit: bool,
#[values(false, true)] override_policy: bool,
) {
use litellm_cache_response::{CacheOptions, CacheScope};
let service = ResponseCache::new(Arc::new(InMemoryCache::<CacheEntry>::default()));
let request =
|scope| CacheOptions::new(scope).request("test", "messages", json!({"prompt":"hello"}));
use litellm_cache_response::{CachePolicy, CacheScope, ScopedCache};
let service = Arc::new(ResponseCache::new(Arc::new(
InMemoryCache::<CacheEntry>::default(),
)));
let request = |scope| {
ScopedCache::new(service.clone(), scope)
.options(override_policy.then_some(CachePolicy {
ttl: Some(Duration::from_secs(30)),
..CachePolicy::default()
}))
.request("test", "messages", json!({"prompt":"hello"}))
};
service
.async_store(
&request(CacheScope::Isolated(first.into())),

View file

@ -36,7 +36,9 @@ Not here: serving HTTP (axum routes, extractors), config file reading, rollout s
## Response caching and accounting boundary
Attach a `litellm_cache_response::ScopedCache` with `route.with_cache(cache)`. Cached and uncached routes use the same `execute` and `machine` methods. `CallOptions` carries per-call cache overrides and observation; attaching a service does not change the execution contract
Attach a `litellm_cache_response::ScopedCache` with `route.with_cache(cache)`. Cached and uncached routes use the same `execute` and `machine` methods. `CallOptions` carries a scope-free `CachePolicy` and observation; per-call policy never replaces the attached scope or service
Messages groups per-call dependencies in `CallContext` and explicitly sequences cache lookup, provider execution, result acceptance, and cache storage. Provider transport does not own cache orchestration. Stream capture remains in the shared cache implementation
Core owns request identity, typed response reconstruction and stream capture/replay. `cache-response` owns cache policy, namespacing, scope encoding, versioned envelopes and freshness. The SDK explicitly chooses shared scope. The gateway derives isolated scope from authenticated identity before attaching its service

View file

@ -1,5 +1,6 @@
use std::{
future::Future,
marker::PhantomData,
sync::Arc,
time::{Duration, SystemTime, UNIX_EPOCH},
};
@ -7,7 +8,8 @@ use std::{
use bytes::{Bytes, BytesMut};
use futures_util::{StreamExt, TryStreamExt, stream};
use litellm_cache_response::{
CacheOptions, ResponseCacheRequest, ResponseCacheService, ResponseEnvelope, cache_key,
CacheOptions, CachePolicy, ResponseCacheRequest, ResponseCacheService, ResponseEnvelope,
ScopedCache, cache_key,
};
use litellm_host::{
call::{CallOutput, OutputOf},
@ -77,7 +79,7 @@ impl CacheSession {
options: Option<CacheOptions>,
request: &CacheRequest,
) -> Option<Self> {
let options = options.filter(CacheOptions::enabled)?;
let options = options.filter(|options| options.policy.enabled())?;
let service = service?;
let input = request.input.clone();
let request = options.request(&service.config().namespace, P::SURFACE, input);
@ -198,81 +200,137 @@ where
let identity = request.identity.clone();
crate::diagnostic::provider(&identity.model, &identity.provider);
let session = CacheSession::prepare::<P>(cache, options, &request);
let hit = match &session {
Some(session) => session.lookup::<P>().await.and_then(|entry| {
let output = match entry {
CachedOutput::Response(response) => Some(CallOutput::Complete(response)),
CachedOutput::Stream(data) => P::replay(Bytes::from(data)),
};
output.map(|output| (output, cache_key(&session.request.key)))
}),
None => None,
let cache = CallCache::<P> {
session,
protocol: PhantomData,
};
let hit = cache.lookup().await;
let (output, source) = match hit {
Some((output, key)) => (output, ResultSource::Cache { key }),
Some(hit) => hit,
None => (provider().await?, ResultSource::Provider),
};
let from_provider = source == ResultSource::Provider;
publish(
ExecutionFacts {
provider: identity,
source,
source: source.clone(),
},
interceptors,
observers,
)
.await?;
let Some(session) =
session.filter(|session| from_provider && session.request.controls.writes())
else {
return Ok(output);
};
match output {
CallOutput::Complete(response) => {
session.store_response::<P>(&response).await;
Ok(CallOutput::Complete(response))
}
CallOutput::Stream { head, chunks } => {
let captured = stream::try_unfold(
(chunks, Some(Vec::<u8>::new()), session),
|(mut chunks, captured, session)| async move {
match chunks.try_next().await? {
Some(chunk) => {
let captured = captured.and_then(|mut data| {
let bytes = P::bytes(&chunk);
if data.len().saturating_add(bytes.len())
> session.service.config().max_entry_bytes
{
return None;
}
data.extend_from_slice(bytes);
Some(data)
});
Ok(Some((chunk, (chunks, captured, session))))
}
None => {
if let Some(data) = captured
&& let Ok(text) = String::from_utf8(data)
&& successful_stream(&text, P::TERMINAL_EVENT)
&& let Ok(entry) = serde_json::to_value(ResponseEnvelope::new(
P::SURFACE,
CachedOutput::<Value>::Stream(text),
))
{
session.store(entry).await;
}
Ok::<_, RouteError>(None)
}
}
},
)
.boxed();
Ok(CallOutput::Stream {
head,
chunks: captured,
Ok(cache.finish(output, &source).await)
}
pub(crate) struct CallCache<P> {
session: Option<CacheSession>,
protocol: PhantomData<P>,
}
impl<P: StreamCachable> CallCache<P> {
pub(crate) fn from_wire(
cache: Option<&ScopedCache>,
policy: CachePolicy,
identity: &ProviderIdentity,
wire: &WireRequest,
) -> Self {
let session = cache.and_then(|cache| {
if !policy.enabled() {
return None;
}
let options = cache.options(Some(policy));
let request = CacheRequest::from_wire(identity.clone(), Some(wire));
Some(CacheSession {
request: options.request(
&cache.service.config().namespace,
P::SURFACE,
request.input,
),
service: cache.service.clone(),
})
});
Self {
session,
protocol: PhantomData,
}
}
pub(crate) async fn lookup(&self) -> Option<(OutputOf<P>, ResultSource)>
where
P::Response: DeserializeOwned,
{
let session = self.session.as_ref()?;
let output = match session.lookup::<P>().await? {
CachedOutput::Response(response) => CallOutput::Complete(response),
CachedOutput::Stream(data) => P::replay(Bytes::from(data))?,
};
Some((
output,
ResultSource::Cache {
key: cache_key(&session.request.key),
},
))
}
pub(crate) async fn finish(self, output: OutputOf<P>, source: &ResultSource) -> OutputOf<P>
where
P::Response: Serialize,
{
let Some(session) = self.session.filter(|session| {
*source == ResultSource::Provider && session.request.controls.writes()
}) else {
return output;
};
match output {
CallOutput::Complete(response) => {
session.store_response::<P>(&response).await;
CallOutput::Complete(response)
}
CallOutput::Stream { head, chunks } => CallOutput::Stream {
head,
chunks: capture_stream::<P>(chunks, session),
},
}
}
}
fn capture_stream<P: StreamCachable>(
chunks: futures_util::stream::BoxStream<'static, Result<P::Chunk, RouteError>>,
session: CacheSession,
) -> futures_util::stream::BoxStream<'static, Result<P::Chunk, RouteError>> {
stream::try_unfold(
(chunks, Some(Vec::<u8>::new()), session),
|(mut chunks, captured, session)| async move {
match chunks.try_next().await? {
Some(chunk) => {
let captured = captured.and_then(|mut data| {
let bytes = P::bytes(&chunk);
if data.len().saturating_add(bytes.len())
> session.service.config().max_entry_bytes
{
return None;
}
data.extend_from_slice(bytes);
Some(data)
});
Ok(Some((chunk, (chunks, captured, session))))
}
None => {
if let Some(data) = captured
&& let Ok(text) = String::from_utf8(data)
&& successful_stream(&text, P::TERMINAL_EVENT)
&& let Ok(entry) = serde_json::to_value(ResponseEnvelope::new(
P::SURFACE,
CachedOutput::<Value>::Stream(text),
))
{
session.store(entry).await;
}
Ok::<_, RouteError>(None)
}
}
},
)
.boxed()
}
fn now() -> Duration {

View file

@ -22,7 +22,7 @@ pub(super) async fn execute(
auth: &AuthServices,
request: ProviderChatCompletionsRequest,
cache: Option<litellm_cache_response::ScopedCache>,
cache_options: Option<litellm_cache_response::CacheOptions>,
cache_options: Option<litellm_cache_response::CachePolicy>,
interceptors: &impl Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<ChatCompletionsResponse, Error> {

View file

@ -67,7 +67,7 @@ impl ChatCompletionsRoute {
async fn run(
&self,
request: ChatCompletionsRequest<'_>,
cache_options: Option<litellm_cache_response::CacheOptions>,
cache_options: Option<litellm_cache_response::CachePolicy>,
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<ChatCompletionsResponse, Error> {

View file

@ -55,7 +55,7 @@ impl ChatCompletionsRoute {
pub(super) async fn run_call(
&self,
call: ChatCompletionsCall,
cache_options: Option<litellm_cache_response::CacheOptions>,
cache_options: Option<litellm_cache_response::CachePolicy>,
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<ChatCompletionsResponse, Error> {

View file

@ -0,0 +1,48 @@
use litellm_cache_response::CachePolicy;
use litellm_host::{
interceptors::{ExecutionFacts, Interceptors, RawResponse},
lifecycle::{CallEvent, ExecutionEvent},
observation::ObservationSender,
};
use crate::{CallOptions, RouteError};
pub(crate) struct CallContext<'a, I> {
pub interceptors: &'a I,
pub observers: Option<ObservationSender>,
pub cache: CachePolicy,
}
impl<'a, I: Interceptors<RouteError>> CallContext<'a, I> {
pub fn new(interceptors: &'a I, options: CallOptions) -> Self {
Self {
interceptors,
observers: options.observers,
cache: options.cache.unwrap_or_default(),
}
}
pub async fn result_ready(&self, facts: ExecutionFacts) -> Result<(), RouteError> {
if let Some(observers) = &self.observers {
observers.emit(CallEvent::Execution(ExecutionEvent::ResultReady {
facts: facts.clone(),
}));
}
self.interceptors.result_ready(facts).await
}
pub async fn response_received(&self, body: &str) -> Result<(), RouteError> {
let raw = RawResponse {
body: body.to_owned(),
};
if let Some(observers) = &self.observers {
observers.emit(CallEvent::Execution(
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
));
}
self.interceptors
.after_provider_response(raw)
.await
.map_err(RouteError::post_call)
}
}

View file

@ -1,3 +1,4 @@
mod context;
mod diagnostic;
pub mod audio_transcription;
@ -16,7 +17,7 @@ pub use error::RouteError;
#[derive(Clone, Default)]
pub struct CallOptions {
pub cache: Option<litellm_cache_response::CacheOptions>,
pub cache: Option<litellm_cache_response::CachePolicy>,
pub observers: Option<litellm_host::observation::ObservationSender>,
}
@ -29,8 +30,8 @@ impl From<Option<litellm_host::observation::ObservationSender>> for CallOptions
}
}
impl From<litellm_cache_response::CacheOptions> for CallOptions {
fn from(cache: litellm_cache_response::CacheOptions) -> Self {
impl From<litellm_cache_response::CachePolicy> for CallOptions {
fn from(cache: litellm_cache_response::CachePolicy) -> Self {
Self {
cache: Some(cache),
observers: None,

View file

@ -1,10 +1,8 @@
use litellm_host::{lifecycle::ExecutionEvent, observation::ObservationSender};
use std::time::Duration;
use bytes::Bytes;
use futures_util::{StreamExt, TryStreamExt, stream::BoxStream};
use litellm_auth::AuthServices;
use litellm_host::interceptors::{Interceptors, RawResponse, RequestContext, WireRequest};
use litellm_host::interceptors::{Interceptors, ProviderIdentity, RequestContext, WireRequest};
use litellm_http::transport::Error as TransportError;
use litellm_llms::base_llm::{
auth::{Authenticated, resolve_auth},
@ -18,102 +16,130 @@ use litellm_tracing::ByteChunk;
use serde_json::Value;
use super::{
Error, MessagesCallResponse, common_utils::truncate_error_body,
Error, MessagesCallResponse, MessagesRoute, common_utils::truncate_error_body,
prepare::ProviderMessagesRequest,
};
use crate::{constants::MESSAGES_TIMEOUT_SECS, outbound::outbound_request};
use crate::{constants::MESSAGES_TIMEOUT_SECS, context::CallContext, outbound::outbound_request};
pub(super) async fn execute(
http: &litellm_http::Client,
auth: &AuthServices,
request: ProviderMessagesRequest,
cache: Option<litellm_cache_response::ScopedCache>,
cache_options: Option<litellm_cache_response::CacheOptions>,
interceptors: &impl Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<MessagesCallResponse, Error> {
let ProviderMessagesRequest {
provider,
url,
body,
environment,
timeout,
api_key,
} = request;
let stream = body.params.stream == Some(true);
let context = RequestContext {
model: body.model.clone(),
custom_llm_provider: provider.as_str().to_string(),
optional_params: serde_json::to_value(&body.params).map_err(serialize_failure)?,
secret_fields: Vec::new(),
api_key,
};
let authenticated = resolve_auth(auth, environment, &|key| std::env::var(key).ok()).await?;
let identity = litellm_host::interceptors::ProviderIdentity {
model: context.model.clone(),
provider: context.custom_llm_provider.clone(),
};
let wire = interceptors
.before_provider_request(
WireRequest {
url,
headers: authenticated.headers,
body: serde_json::to_value(&body).map_err(serialize_failure)?,
},
context,
)
.await?;
let cache = cache.filter(|_| authenticated.signer.is_none());
let cache_request =
crate::caching::CacheRequest::from_wire(identity, cache.as_ref().map(|_| &wire));
crate::caching::execute_streaming::<super::route::Messages, _, _>(
cache_request,
cache.as_ref().map(|cache| cache.service.clone()),
cache.as_ref().map(|cache| cache.options(cache_options)),
interceptors,
observers,
|| async move {
let provider_name = provider.as_str();
log_request_body(provider_name, stream, &wire.body);
let response = send(
http,
Authenticated {
headers: wire.headers,
signer: authenticated.signer,
pub(super) struct ProviderCall {
pub identity: ProviderIdentity,
pub wire: WireRequest,
provider: super::common_utils::MessagesProvider,
signer: Option<litellm_auth_aws::SigV4Signer>,
timeout: Option<Duration>,
stream: bool,
}
impl ProviderCall {
pub fn cacheable(&self) -> bool {
self.signer.is_none()
}
}
impl MessagesRoute {
pub(super) async fn prepare_outbound(
&self,
request: ProviderMessagesRequest,
context: &CallContext<'_, impl Interceptors<Error>>,
) -> Result<ProviderCall, Error> {
let ProviderMessagesRequest {
provider,
url,
body,
environment,
timeout,
api_key,
} = request;
let request_context = RequestContext {
model: body.model.clone(),
custom_llm_provider: provider.as_str().to_string(),
optional_params: serde_json::to_value(&body.params).map_err(serialize_failure)?,
secret_fields: Vec::new(),
api_key,
};
let authenticated =
resolve_auth(&self.auth, environment, &|key| std::env::var(key).ok()).await?;
let identity = ProviderIdentity {
model: request_context.model.clone(),
provider: request_context.custom_llm_provider.clone(),
};
let wire = context
.interceptors
.before_provider_request(
WireRequest {
url,
headers: authenticated.headers,
body: serde_json::to_value(&body).map_err(serialize_failure)?,
},
&wire.url,
&wire.body,
timeout,
request_context,
)
.await?;
if !response.status().is_success() {
return Err(provider_error(response).await);
}
let config = provider.config();
if stream {
return Ok(streaming_response(
response,
config.stream_decoder(),
provider_name,
let stream = match wire.body.get("stream") {
None | Some(Value::Null) => false,
Some(Value::Bool(stream)) => *stream,
Some(value) => {
return Err(Error::InvalidRequest(
litellm_llms::ErrorDetail::InvalidValue {
field: "stream",
expected: "a boolean",
actual: value.clone(),
},
));
}
let text = response.text().await.map_err(network)?;
log_response_body(&text);
let raw = RawResponse { body: text.clone() };
if let Some(observers) = observers {
observers.emit(litellm_host::lifecycle::CallEvent::Execution(
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
));
}
interceptors
.after_provider_response(raw)
.await
.map_err(Error::post_call)?;
decode_response(config, &body.model, &text)
.map(|message| MessagesCallResponse::Complete(Box::new(message)))
},
)
.await
};
Ok(ProviderCall {
identity,
wire,
provider,
signer: authenticated.signer,
timeout,
stream,
})
}
pub(super) async fn call_provider(
&self,
request: ProviderCall,
context: &CallContext<'_, impl Interceptors<Error>>,
) -> Result<MessagesCallResponse, Error> {
let ProviderCall {
identity,
wire,
provider,
signer,
timeout,
stream,
} = request;
let provider_name = provider.as_str();
log_request_body(provider_name, stream, &wire.body);
let response = send(
&self.http,
Authenticated {
headers: wire.headers,
signer,
},
&wire.url,
&wire.body,
timeout,
)
.await?;
if !response.status().is_success() {
return Err(provider_error(response).await);
}
let config = provider.config();
if stream {
return Ok(streaming_response(
response,
config.stream_decoder(),
provider_name,
));
}
let text = response.text().await.map_err(network)?;
log_response_body(&text);
context.response_received(&text).await?;
decode_response(config, &identity.model, &text)
.map(|message| MessagesCallResponse::Complete(Box::new(message)))
}
}
fn serialize_failure(err: serde_json::Error) -> Error {

View file

@ -1,11 +1,14 @@
use litellm_host::observation::ObservationSender;
mod common_utils;
mod handler;
mod prepare;
pub mod route;
mod types;
use futures_util::FutureExt;
use litellm_auth::AuthServices;
use litellm_host::interceptors::{ExecutionFacts, Interceptors, ResultSource};
use crate::{caching::CallCache, context::CallContext};
use litellm_secrets::source::SecretSource;
use std::sync::Arc;
@ -20,74 +23,18 @@ pub struct MessagesRoute {
cache: Option<litellm_cache_response::ScopedCache>,
}
#[must_use]
#[derive(Clone, Default)]
pub struct MessagesRouteBuilder<Http = (), Auth = (), Secrets = ()> {
http: Http,
auth: Auth,
secrets: Secrets,
cache: Option<litellm_cache_response::ScopedCache>,
}
impl<Http, Auth, Secrets> MessagesRouteBuilder<Http, Auth, Secrets> {
pub fn with_http(
self,
http: litellm_http::Client,
) -> MessagesRouteBuilder<litellm_http::Client, Auth, Secrets> {
MessagesRouteBuilder {
http,
auth: self.auth,
secrets: self.secrets,
cache: self.cache,
}
}
pub fn with_auth(
self,
auth: Arc<AuthServices>,
) -> MessagesRouteBuilder<Http, Arc<AuthServices>, Secrets> {
MessagesRouteBuilder {
http: self.http,
auth,
secrets: self.secrets,
cache: self.cache,
}
}
pub fn with_secrets(
self,
secrets: Arc<dyn SecretSource>,
) -> MessagesRouteBuilder<Http, Auth, Arc<dyn SecretSource>> {
MessagesRouteBuilder {
http: self.http,
auth: self.auth,
secrets,
cache: self.cache,
}
}
pub fn with_cache(self, cache: litellm_cache_response::ScopedCache) -> Self {
Self {
cache: Some(cache),
..self
}
}
}
impl MessagesRouteBuilder<litellm_http::Client, Arc<AuthServices>, Arc<dyn SecretSource>> {
pub fn build(self) -> MessagesRoute {
MessagesRoute {
http: self.http,
auth: self.auth,
secrets: self.secrets,
cache: self.cache,
}
}
}
impl MessagesRoute {
pub fn builder() -> MessagesRouteBuilder {
MessagesRouteBuilder::default()
pub fn new(
http: litellm_http::Client,
auth: Arc<AuthServices>,
secrets: Arc<dyn SecretSource>,
) -> Self {
Self {
http,
auth,
secrets,
cache: None,
}
}
#[must_use]
@ -104,15 +51,9 @@ impl MessagesRoute {
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
options: impl Into<crate::CallOptions>,
) -> Result<MessagesCallResponse, Error> {
let crate::CallOptions {
cache: cache_options,
observers,
} = options.into();
litellm_host::lifecycle::observe_call(
observers.clone(),
self.run(call, cache_options, interceptors, observers.as_ref()),
)
.await
let context = CallContext::new(interceptors, options.into());
litellm_host::lifecycle::observe_call(context.observers.clone(), self.run(call, context))
.await
}
#[tracing::instrument(name = "litellm.route", skip_all, fields(
@ -126,36 +67,34 @@ impl MessagesRoute {
async fn run(
&self,
call: MessagesCall,
cache_options: Option<litellm_cache_response::CacheOptions>,
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
observers: Option<&ObservationSender>,
context: CallContext<'_, impl Interceptors<Error>>,
) -> Result<MessagesCallResponse, Error> {
crate::diagnostic::call(async {
self.run_provider(call, cache_options, interceptors, observers)
.await
let prepared = prepare::prepare(call, self.secrets.as_ref()).await?;
crate::diagnostic::provider(&prepared.body.model, prepared.provider.as_str());
let request = self.prepare_outbound(prepared, &context).boxed().await?;
let cache = CallCache::<route::Messages>::from_wire(
self.cache.as_ref().filter(|_| request.cacheable()),
context.cache,
&request.identity,
&request.wire,
);
let identity = request.identity.clone();
let (output, source) = match cache.lookup().await {
Some(hit) => hit,
None => (
self.call_provider(request, &context).await?,
ResultSource::Provider,
),
};
context
.result_ready(ExecutionFacts {
provider: identity,
source: source.clone(),
})
.await?;
Ok(cache.finish(output, &source).await)
})
.await
}
async fn run_provider(
&self,
call: MessagesCall,
cache_options: Option<litellm_cache_response::CacheOptions>,
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<MessagesCallResponse, Error> {
let request = prepare::prepare(call, self.secrets.as_ref()).await?;
crate::diagnostic::provider(&request.body.model, request.provider.as_str());
let execute: futures_util::future::BoxFuture<'_, Result<MessagesCallResponse, Error>> =
Box::pin(handler::execute(
&self.http,
&self.auth,
request,
self.cache.clone(),
cache_options,
interceptors,
observers,
));
execute.await
}
}

View file

@ -43,8 +43,14 @@ impl super::MessagesRoute {
request,
observers,
move |call, _, interceptors, observers| async move {
self.run(call, cache_options, &interceptors, observers.as_ref())
.await
let context = crate::context::CallContext::new(
&interceptors,
crate::CallOptions {
cache: cache_options,
observers,
},
);
self.run(call, context).await
},
)
}

View file

@ -16,7 +16,7 @@ pub(super) async fn execute(
auth: &litellm_auth::AuthServices,
request: ProviderResponsesRequest,
cache: Option<litellm_cache_response::ScopedCache>,
cache_options: Option<litellm_cache_response::CacheOptions>,
cache_options: Option<litellm_cache_response::CachePolicy>,
interceptors: &impl Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<ResponsesOutput, Error> {

View file

@ -71,7 +71,7 @@ impl ResponsesRoute {
async fn run(
&self,
call: ResponsesCall,
cache_options: Option<litellm_cache_response::CacheOptions>,
cache_options: Option<litellm_cache_response::CachePolicy>,
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<ResponsesOutput, Error> {
@ -85,7 +85,7 @@ impl ResponsesRoute {
async fn run_provider(
&self,
call: ResponsesCall,
cache_options: Option<litellm_cache_response::CacheOptions>,
cache_options: Option<litellm_cache_response::CachePolicy>,
interceptors: &impl Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<ResponsesOutput, Error> {

View file

@ -12,8 +12,8 @@ use bytes::Bytes;
use futures_util::{StreamExt, TryStreamExt, stream};
use litellm_cache_memory::InMemoryCache;
use litellm_cache_response::{
CacheOptions, CacheScope, ResponseCache, ResponseCacheConfig, ResponseCacheService,
ResponseEnvelope,
CacheOptions, CachePolicy, CacheScope, ResponseCache, ResponseCacheConfig,
ResponseCacheService, ResponseEnvelope,
};
use litellm_core::{
RouteError,
@ -117,9 +117,9 @@ async fn call(
#[rstest]
#[case::normal(CacheOptions::new(CacheScope::Shared), true, true)]
#[case::no_cache(CacheOptions { no_cache: true, ..CacheOptions::new(CacheScope::Shared) }, false, true)]
#[case::no_store(CacheOptions { no_store: true, ..CacheOptions::new(CacheScope::Shared) }, true, false)]
#[case::disabled(CacheOptions { caching: Some(false), ..CacheOptions::new(CacheScope::Shared) }, false, false)]
#[case::no_cache(CacheOptions { policy: CachePolicy { no_cache: true, ..CachePolicy::default() }, ..CacheOptions::new(CacheScope::Shared) }, false, true)]
#[case::no_store(CacheOptions { policy: CachePolicy { no_store: true, ..CachePolicy::default() }, ..CacheOptions::new(CacheScope::Shared) }, true, false)]
#[case::disabled(CacheOptions { policy: CachePolicy { caching: Some(false), ..CachePolicy::default() }, ..CacheOptions::new(CacheScope::Shared) }, false, false)]
#[tokio::test]
async fn cache_controls_apply_to_both_reads_and_writes(
cache: Arc<dyn ResponseCacheService>,
@ -723,9 +723,9 @@ async fn unary_call(
#[rstest]
#[case::normal(CacheOptions::new(CacheScope::Shared), true, true)]
#[case::no_cache(CacheOptions { no_cache: true, ..CacheOptions::new(CacheScope::Shared) }, false, true)]
#[case::no_store(CacheOptions { no_store: true, ..CacheOptions::new(CacheScope::Shared) }, true, false)]
#[case::disabled(CacheOptions { caching: Some(false), ..CacheOptions::new(CacheScope::Shared) }, false, false)]
#[case::no_cache(CacheOptions { policy: CachePolicy { no_cache: true, ..CachePolicy::default() }, ..CacheOptions::new(CacheScope::Shared) }, false, true)]
#[case::no_store(CacheOptions { policy: CachePolicy { no_store: true, ..CachePolicy::default() }, ..CacheOptions::new(CacheScope::Shared) }, true, false)]
#[case::disabled(CacheOptions { policy: CachePolicy { caching: Some(false), ..CachePolicy::default() }, ..CacheOptions::new(CacheScope::Shared) }, false, false)]
#[tokio::test]
async fn unary_cache_controls_do_not_change_the_shared_service(
cache: Arc<dyn ResponseCacheService>,

View file

@ -1,9 +1,9 @@
use litellm_host::lifecycle::ExecutionEvent;
use std::sync::Mutex;
use litellm_core::messages::route::Messages;
use litellm_core::messages::{MessagesCallResponse, route::Messages};
use litellm_host::{
interceptors::{RequestContext, WireRequest},
interceptors::{ExecutionFacts, RequestContext, ResultSource, WireRequest},
lifecycle::CallEvent,
};
use litellm_llms::base_llm::messages::context::MessagesModelCapabilities as AnthropicModelCapabilities;
@ -20,6 +20,8 @@ struct RecordingHost {
rewrite: Rewrite,
events: super::support::Observations,
optional_params: Mutex<Vec<Value>>,
facts: Mutex<Vec<ExecutionFacts>>,
reject_result: bool,
}
impl RecordingHost {
@ -29,6 +31,8 @@ impl RecordingHost {
rewrite,
events: super::support::Observations::default(),
optional_params: Mutex::new(Vec::new()),
facts: Mutex::new(Vec::new()),
reject_result: false,
}
}
@ -73,6 +77,14 @@ impl litellm_host::lifecycle::CallObserver for RecordingHost {
impl litellm_host::interceptors::Interceptors<<Messages as litellm_host::protocol::Protocol>::Error>
for RecordingHost
{
async fn result_ready(&self, facts: ExecutionFacts) -> Result<(), Error> {
self.facts.lock().unwrap().push(facts);
if self.reject_result {
return Err(Error::Unsupported("result rejected"));
}
Ok(())
}
async fn before_provider_request(
&self,
wire: WireRequest,
@ -98,6 +110,91 @@ impl litellm_host::interceptors::Interceptors<<Messages as litellm_host::protoco
}
}
#[rstest]
#[case::native_unary(false, false)]
#[case::native_stream(true, false)]
#[case::hosted_unary(false, true)]
#[case::hosted_stream(true, true)]
#[tokio::test]
async fn rejected_results_are_not_delivered_or_cached(
call: MessagesCall,
#[case] streaming: bool,
#[case] hosted: bool,
) {
use futures_util::TryStreamExt;
use litellm_cache_memory::InMemoryCache;
use litellm_cache_response::{CacheScope, ResponseCache, ScopedCache};
let response = if streaming {
ResponseTemplate::new(200).set_body_raw(
"event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n",
"text/event-stream",
)
} else {
message_response()
};
let upstream = upstream([response.clone(), response]).await;
let route = messages_route(no_secrets()).with_cache(ScopedCache::new(
Arc::new(ResponseCache::new(Arc::new(InMemoryCache::new(
Some(100),
Some(Duration::from_secs(60)),
)))),
CacheScope::Shared,
));
for (reject, expected_requests, cached) in [
(true, 1, false),
(false, 2, false),
(true, 2, true),
(false, 2, true),
] {
let request = authenticated(
with_fields(
MessagesCall {
body: call.body.clone(),
..super::call()
},
json!({"stream": streaming}),
),
upstream.uri(),
);
let host = RecordingHost {
reject_result: reject,
..RecordingHost::passthrough(request)
};
let result = if hosted {
litellm_host_native::in_process::run_hosted(
route.clone().machine(host.request().unwrap(), None),
host.runtime(),
)
.await
.map(|_| ())
} else {
match route.execute(host.request().unwrap(), &host, None).await {
Ok(MessagesCallResponse::Complete(_)) => Ok(()),
Ok(MessagesCallResponse::Stream { chunks, .. }) => {
chunks.try_collect::<Vec<_>>().await.map(|_| ())
}
Err(error) => Err(error),
}
};
assert_eq!(
result,
if reject {
Err(Error::Unsupported("result rejected"))
} else {
Ok(())
}
);
assert_eq!(received(&upstream).await.len(), expected_requests);
let facts = host.facts.lock().unwrap();
assert_eq!(facts.len(), 1);
assert_eq!(
matches!(facts[0].source, ResultSource::Cache { .. }),
cached
);
}
}
async fn run_through(host: &RecordingHost) -> Result<MessagesOutput, Error> {
litellm_host_native::in_process::run_hosted(
machine(Arc::new(RecordingSecrets::empty()))(host.request()?),
@ -143,6 +240,81 @@ async fn what_before_send_returns_is_what_the_provider_receives(call: MessagesCa
assert_eq!(request.header("x-api-key"), Some("sk-ant"));
}
#[rstest]
#[case::enable(false, json!(true), Some(true))]
#[case::disable(true, json!(false), Some(false))]
#[case::null(true, Value::Null, Some(false))]
#[case::invalid(false, json!("true"), None)]
#[tokio::test]
async fn response_mode_follows_the_intercepted_request(
call: MessagesCall,
traces: TraceCapture,
#[case] original_stream: bool,
#[case] rewritten_stream: Value,
#[case] expected_stream: Option<bool>,
) {
use futures_util::TryStreamExt;
let sse = "event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n";
let response = if expected_stream == Some(true) {
ResponseTemplate::new(200).set_body_raw(sse, "text/event-stream")
} else {
message_response()
};
let upstream = upstream([response]).await;
let rewrite = rewritten_stream.clone();
let host = RecordingHost::new(
authenticated(
with_fields(call, json!({"stream": original_stream})),
upstream.uri(),
),
Box::new(move |wire| {
let mut body = wire.body;
body["stream"] = rewrite.clone();
Ok(WireRequest { body, ..wire })
}),
);
let result = traces
.logger()
.instrument(async {
let output = messages_route(no_secrets())
.execute(host.request()?, &host, None)
.await?;
match output {
MessagesCallResponse::Stream { chunks, .. } => {
assert_eq!(expected_stream, Some(true));
assert_eq!(
chunks.try_collect::<Vec<_>>().await?.concat(),
sse.as_bytes()
);
}
MessagesCallResponse::Complete(message) => {
assert_eq!(expected_stream, Some(false));
assert_eq!(*message, serde_json::from_value(message_body()).unwrap());
}
}
Ok::<_, Error>(())
})
.await;
let summaries = traces.summaries("litellm.route");
assert_eq!(summaries.len(), 1);
let Some(expected_stream) = expected_stream else {
assert!(matches!(result, Err(Error::InvalidRequest(_))));
assert!(received(&upstream).await.is_empty());
assert_eq!(summaries[0]["outcome"], "failure");
return;
};
result.unwrap();
assert_eq!(
only_request(&upstream).await.json()["stream"],
rewritten_stream
);
assert_eq!(host.raw_responses().len(), usize::from(!expected_stream));
assert_eq!(summaries[0]["stream"], expected_stream);
assert_eq!(summaries[0]["outcome"], "success");
}
#[rstest]
#[tokio::test]
async fn a_before_send_failure_never_sends(call: MessagesCall) {

View file

@ -269,25 +269,22 @@ async fn the_facade_sends_through_the_injected_http_pool_configuration(call: Mes
};
let resources = support::resources();
let response = litellm_core::messages::MessagesRoute::builder()
.with_http(provider_http(
&resources,
&Resolution::from(&settings).config,
))
.with_auth(resources.auth)
.with_secrets(no_secrets())
.build()
.execute(
MessagesCall {
api_key: Some("sk-ant".into()),
api_base: Some(base),
..call
},
&(),
None,
)
.await
.expect("messages request succeeds");
let response = litellm_core::messages::MessagesRoute::new(
provider_http(&resources, &Resolution::from(&settings).config),
resources.auth,
no_secrets(),
)
.execute(
MessagesCall {
api_key: Some("sk-ant".into()),
api_base: Some(base),
..call
},
&(),
None,
)
.await
.expect("messages request succeeds");
let MessagesCallResponse::Complete(message) = response else {
panic!("a non-streaming request returns a message");
@ -345,7 +342,7 @@ async fn message_route_summary_excludes_payload_diagnostics(
#[case::uncached(false, 2)]
#[case::cached(true, 1)]
#[tokio::test]
async fn builder_preserves_dependencies_and_optional_cache(
async fn route_uses_injected_dependencies_and_optional_cache(
#[case] caching: bool,
#[case] expected_requests: usize,
) {
@ -355,9 +352,13 @@ async fn builder_preserves_dependencies_and_optional_cache(
let upstream = upstream([message_response(), message_response()]).await;
let resources = resources();
let builder = MessagesRoute::builder();
let builder = if caching {
builder.with_cache(ScopedCache::new(
let route = MessagesRoute::new(
provider_http(&resources, &http_config()),
resources.auth.clone(),
Arc::new(RecordingSecrets::new([("ANTHROPIC_API_KEY", "route-key")])),
);
let route = if caching {
route.with_cache(ScopedCache::new(
Arc::new(ResponseCache::new(Arc::new(InMemoryCache::new(
Some(100),
Some(Duration::from_secs(60)),
@ -365,16 +366,8 @@ async fn builder_preserves_dependencies_and_optional_cache(
CacheScope::Shared,
))
} else {
builder
route
};
let route = builder
.with_secrets(Arc::new(RecordingSecrets::new([(
"ANTHROPIC_API_KEY",
"builder-key",
)])))
.with_auth(resources.auth.clone())
.with_http(provider_http(&resources, &http_config()))
.build();
for _ in 0..2 {
let request = MessagesCall {
api_base: Some(upstream.uri()),
@ -392,5 +385,72 @@ async fn builder_preserves_dependencies_and_optional_cache(
}
let requests = received(&upstream).await;
assert_eq!(requests.len(), expected_requests);
assert_eq!(requests[0].header("x-api-key"), Some("builder-key"));
assert_eq!(requests[0].header("x-api-key"), Some("route-key"));
}
#[rstest]
#[tokio::test]
async fn cache_overrides_preserve_the_routes_isolated_scope(call: MessagesCall) {
use litellm_cache_memory::InMemoryCache;
use litellm_cache_response::{CachePolicy, CacheScope, ResponseCache, ScopedCache};
let first_body = message_body();
let second_body = Value::Object(
first_body
.as_object()
.unwrap()
.iter()
.map(|(key, value)| {
(
key.clone(),
if key == "id" {
json!("msg_second")
} else {
value.clone()
},
)
})
.collect(),
);
let upstream = upstream([
json_response(first_body.clone()),
json_response(second_body.clone()),
])
.await;
let service = Arc::new(ResponseCache::new(Arc::new(InMemoryCache::new(
Some(100),
Some(Duration::from_secs(60)),
))));
let first = messages_route(no_secrets()).with_cache(ScopedCache::new(
service.clone(),
CacheScope::Isolated("first".into()),
));
let second = messages_route(no_secrets()).with_cache(ScopedCache::new(
service,
CacheScope::Isolated("second".into()),
));
for (route, expected) in [
(&first, &first_body),
(&second, &second_body),
(&first, &first_body),
(&second, &second_body),
] {
let request = MessagesCall {
body: call.body.clone(),
api_key: Some("same-key".into()),
api_base: Some(upstream.uri()),
..super::call()
};
let override_options = CachePolicy {
ttl: Some(Duration::from_secs(30)),
..CachePolicy::default()
};
let MessagesCallResponse::Complete(response) =
route.execute(request, &(), override_options).await.unwrap()
else {
panic!("expected a completed message");
};
assert_eq!(response.id, expected["id"].as_str().unwrap());
}
assert_eq!(received(&upstream).await.len(), 2);
}

View file

@ -44,11 +44,11 @@ pub fn provider_http(
pub fn messages_route(secrets: Arc<dyn SecretSource>) -> litellm_core::messages::MessagesRoute {
let resources = resources();
litellm_core::messages::MessagesRoute::builder()
.with_http(provider_http(&resources, &http_config()))
.with_auth(resources.auth)
.with_secrets(secrets)
.build()
litellm_core::messages::MessagesRoute::new(
provider_http(&resources, &http_config()),
resources.auth,
secrets,
)
}
pub fn chat_completions_route() -> litellm_core::chat_completions::ChatCompletionsRoute {

View file

@ -1,6 +1,6 @@
use std::time::Duration;
use litellm_cache_response::{CacheOptions, CacheScope};
use litellm_cache_response::{CacheOptions, CachePolicy, CacheScope};
use litellm_gateway_auth::AuthenticatedRequest;
use serde::Deserialize;
use serde_json::{Map, Value};
@ -38,11 +38,13 @@ pub(crate) fn prepare(
.map_err(|error| Error::InvalidBody(error.to_string()))?;
let caller = identity.caller();
let options = CacheOptions {
caching,
no_cache: controls.no_cache,
no_store: controls.no_store,
ttl: controls.ttl.map(duration).transpose()?,
max_age: controls.max_age.map(duration).transpose()?,
policy: CachePolicy {
caching,
no_cache: controls.no_cache,
no_store: controls.no_store,
ttl: controls.ttl.map(duration).transpose()?,
max_age: controls.max_age.map(duration).transpose()?,
},
scope: CacheScope::Isolated(
serde_json::json!([
caller.principal().authority(),

View file

@ -69,7 +69,7 @@ async fn handle(
extra_headers: None,
timeout: deployment.timeout,
},
cache_options,
cache_options.policy,
),
(),
headers.clone(),

View file

@ -68,11 +68,7 @@ impl Gateway {
auth.clone(),
secrets.clone(),
),
messages: MessagesRoute::builder()
.with_http(provider.clone())
.with_auth(auth.clone())
.with_secrets(secrets.clone())
.build(),
messages: MessagesRoute::new(provider.clone(), auth.clone(), secrets.clone()),
responses: ResponsesRoute::new(provider, auth.clone(), secrets.clone()),
ocr: OcrRoute::new(OcrClient::new(
&resources.pool,

View file

@ -54,7 +54,7 @@ async fn handle(
};
let call = project(deployment, body, headers)?;
let machine = route.machine(call, cache_options);
let machine = route.machine(call, cache_options.policy);
let stream =
Sse::<Messages, _, _>::new(Json, |error| Bytes::from(Error::from(error).sse_frame()));
let headers = crate::caching::CacheHeaders::default();

View file

@ -38,7 +38,7 @@ pub(crate) async fn create(
extra_headers: None,
timeout: deployment.timeout,
};
let machine = route.machine(call, cache_options);
let machine = route.machine(call, cache_options.policy);
let stream = Sse::<Responses, _, _>::new(Json, |error| {
let error = Error::from(error);
Bytes::from(format!(

View file

@ -57,6 +57,9 @@ pub struct ModelInfo {
/// Priority service-tier rate for the same-named base field.
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_creation_input_token_cost_above_272k_tokens_priority: Option<f64>,
/// Ultrafast service-tier rate for the same-named base field.
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_creation_input_token_cost_above_272k_tokens_ultrafast: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_creation_input_token_cost_batches: Option<f64>,
/// Flex service-tier rate for the same-named base field.
@ -65,6 +68,9 @@ pub struct ModelInfo {
/// Priority service-tier rate for the same-named base field.
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_creation_input_token_cost_priority: Option<f64>,
/// Ultrafast service-tier rate for the same-named base field.
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_creation_input_token_cost_ultrafast: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_read_input_audio_token_cost: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
@ -101,6 +107,9 @@ pub struct ModelInfo {
/// Priority service-tier rate for the same-named base field.
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_read_input_token_cost_above_272k_tokens_priority: Option<f64>,
/// Ultrafast service-tier rate for the same-named base field.
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_read_input_token_cost_above_272k_tokens_ultrafast: Option<f64>,
/// Rate applied once the prompt exceeds the token threshold in the field name.
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_read_input_token_cost_above_512k_tokens: Option<f64>,
@ -115,6 +124,9 @@ pub struct ModelInfo {
/// Priority service-tier rate for the same-named base field.
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_read_input_token_cost_priority: Option<f64>,
/// Ultrafast service-tier rate for the same-named base field.
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_read_input_token_cost_ultrafast: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub citation_cost_per_token: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
@ -125,6 +137,8 @@ pub struct ModelInfo {
pub computer_use_input_cost_per_1k_tokens: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub computer_use_output_cost_per_1k_tokens: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub cost_per_second: Option<f64>,
/// Reasoning effort the provider applies when the request omits reasoning_effort. Gates whether a non-default temperature or the top_p/logprobs sampling params are accepted, which hold only when the effort resolves to 'none'.
#[serde(skip_serializing_if = "Option::is_none")]
pub default_reasoning_effort: Option<ReasoningEffort>,
@ -211,6 +225,9 @@ pub struct ModelInfo {
/// Priority service-tier rate for the same-named base field.
#[serde(skip_serializing_if = "Option::is_none")]
pub input_cost_per_token_above_272k_tokens_priority: Option<f64>,
/// Ultrafast service-tier rate for the same-named base field.
#[serde(skip_serializing_if = "Option::is_none")]
pub input_cost_per_token_above_272k_tokens_ultrafast: Option<f64>,
/// Rate applied once the prompt exceeds the token threshold in the field name.
#[serde(skip_serializing_if = "Option::is_none")]
pub input_cost_per_token_above_512k_tokens: Option<f64>,
@ -228,6 +245,9 @@ pub struct ModelInfo {
/// Priority service-tier rate for the same-named base field.
#[serde(skip_serializing_if = "Option::is_none")]
pub input_cost_per_token_priority: Option<f64>,
/// Ultrafast service-tier rate for the same-named base field.
#[serde(skip_serializing_if = "Option::is_none")]
pub input_cost_per_token_ultrafast: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub input_cost_per_video_per_second: Option<f64>,
/// Rate applied once the prompt exceeds the token threshold in the field name.
@ -360,6 +380,9 @@ pub struct ModelInfo {
/// Priority service-tier rate for the same-named base field.
#[serde(skip_serializing_if = "Option::is_none")]
pub output_cost_per_token_above_272k_tokens_priority: Option<f64>,
/// Ultrafast service-tier rate for the same-named base field.
#[serde(skip_serializing_if = "Option::is_none")]
pub output_cost_per_token_above_272k_tokens_ultrafast: Option<f64>,
/// Rate applied once the prompt exceeds the token threshold in the field name.
#[serde(skip_serializing_if = "Option::is_none")]
pub output_cost_per_token_above_512k_tokens: Option<f64>,
@ -375,6 +398,9 @@ pub struct ModelInfo {
/// Priority service-tier rate for the same-named base field.
#[serde(skip_serializing_if = "Option::is_none")]
pub output_cost_per_token_priority: Option<f64>,
/// Ultrafast service-tier rate for the same-named base field.
#[serde(skip_serializing_if = "Option::is_none")]
pub output_cost_per_token_ultrafast: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub output_cost_per_video_per_second: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]

View file

@ -330,15 +330,17 @@ pub(in crate::cache) fn configured(
Ok((
Some(cache.service.clone()),
litellm_cache_response::CacheOptions {
caching: kwargs
.get_item("caching")?
.filter(|value| !value.is_none())
.map(|value| value.extract())
.transpose()?,
no_cache: boolean("no-cache")?,
no_store: boolean("no-store")?,
ttl: seconds("ttl")?,
max_age: seconds("s-max-age")?.or(seconds("s-maxage")?),
policy: litellm_cache_response::CachePolicy {
caching: kwargs
.get_item("caching")?
.filter(|value| !value.is_none())
.map(|value| value.extract())
.transpose()?,
no_cache: boolean("no-cache")?,
no_store: boolean("no-store")?,
ttl: seconds("ttl")?,
max_age: seconds("s-max-age")?.or(seconds("s-maxage")?),
},
scope: litellm_cache_response::CacheScope::Shared,
},
))

View file

@ -1,5 +1,7 @@
use super::{native, python};
use litellm_cache_response::{CacheOptions, CacheScope, ResponseCacheService, ScopedCache};
use litellm_cache_response::{
CacheOptions, CachePolicy, CacheScope, ResponseCacheService, ScopedCache,
};
use litellm_host::{
machine::{HostServices, MachineFault},
protocol::Protocol,
@ -154,8 +156,11 @@ pub(crate) fn configure(
.map(|value| value.unwrap_or(false))
};
let options = CacheOptions {
no_cache: boolean("no-cache")?,
no_store: boolean("no-store")?,
policy: CachePolicy {
no_cache: boolean("no-cache")?,
no_store: boolean("no-store")?,
..CachePolicy::default()
},
..CacheOptions::new(CacheScope::Shared)
};
let namespace = cache

View file

@ -175,7 +175,7 @@ fn run_public(
)),
None => route,
};
Ok(route.machine(request, cache_options))
Ok(route.machine(request, cache_options.policy))
},
host::ChatCompletionsPythonHost(host),
hooks,

View file

@ -26,14 +26,12 @@ fn run_messages(
py,
arguments,
move |py, arguments, request| {
let builder = litellm_core::messages::MessagesRoute::builder()
.with_http(
crate::http::provider_client(py, arguments, asynchronous)?
.map_err(crate::http::client_error)?,
)
.with_auth(crate::http::resources().auth.clone())
.with_secrets(crate::secrets::source(py)?);
let route = builder.build();
let route = litellm_core::messages::MessagesRoute::new(
crate::http::provider_client(py, arguments, asynchronous)?
.map_err(crate::http::client_error)?,
crate::http::resources().auth.clone(),
crate::secrets::source(py)?,
);
Ok(litellm_host::call::hosted_call(
request,
None,
@ -51,7 +49,7 @@ fn run_messages(
call,
&interceptors,
litellm_core::CallOptions {
cache: Some(options),
cache: Some(options.policy),
observers,
},
)

View file

@ -96,7 +96,7 @@ fn run_public(
)),
None => route,
};
Ok(route.machine(request, cache_options))
Ok(route.machine(request, cache_options.policy))
},
host::ResponsesPythonHost(host),
hooks,

View file

@ -8,6 +8,7 @@
# Thank you users! We ❤️ you! - Krrish & Ishaan
import ast
import asyncio
import hashlib
import json
import logging
@ -30,10 +31,11 @@ from litellm.types.utils import EmbeddingResponse, is_litellm_owned_kwarg
from .azure_blob_cache import AzureBlobCache
from .base_cache import BaseCache
from .disk_cache import DiskCache
from .dual_cache import DualCache # noqa: F401
from .dual_cache import DualCache
from .gcs_cache import GCSCache
from .in_memory_cache import InMemoryCache
from .qdrant_semantic_cache import QdrantSemanticCache
from .redis_batch import active_post_call_redis_batch
from .redis_cache import RedisCache, log_redis_failure
from .redis_cluster_cache import RedisClusterCache
from .redis_semantic_cache import RedisSemanticCache
@ -68,6 +70,15 @@ def print_verbose(print_statement):
pass
def _ttl_seconds(raw: object) -> int | None:
if not isinstance(raw, (int, float, str)):
return None
try:
return int(raw)
except ValueError:
return None
class CacheMode(str, Enum):
default_on = "default_on"
default_off = "default_off"
@ -759,6 +770,8 @@ class Cache:
await self.batch_cache_write(result, **kwargs)
else:
cache_key, cached_data, kwargs = self._add_cache_logic(result=result, **kwargs)
if await self._defer_set_to_post_call_batch(cache_key, cached_data, kwargs, dynamic_cache_object):
return
if dynamic_cache_object is not None:
await dynamic_cache_object.async_set_cache(cache_key, cached_data, **kwargs)
else:
@ -766,6 +779,39 @@ class Cache:
except Exception as e:
self._log_add_cache_failure(e)
async def _defer_set_to_post_call_batch(
self,
cache_key: str,
cached_data: object,
kwargs: Mapping[str, object],
dynamic_cache_object: BaseCache | None,
) -> bool:
"""A plain SET on the Redis response cache rides the request's post-call pipeline with the counters,
instead of its own round trip. Anything with SET options keeps the direct path."""
if kwargs.get("nx"):
return False
ttl: Final = _ttl_seconds(kwargs.get("ttl"))
if isinstance(dynamic_cache_object, DualCache):
deferred: Final = await dynamic_cache_object.async_set_cache_post_call(cache_key, cached_data, ttl)
if deferred is None:
return False
deferred.on_settled(self._log_deferred_add_cache_failure)
return True
if dynamic_cache_object is not None or not isinstance(self.cache, RedisCache):
return False
batch: Final = active_post_call_redis_batch(self.cache)
if batch is None:
return False
batch.set(cache_key, cached_data, ttl).on_settled(self._log_deferred_add_cache_failure)
return True
def _log_deferred_add_cache_failure(self, future: asyncio.Future[None]) -> None:
if future.cancelled():
return
failure: Final = future.exception()
if isinstance(failure, Exception):
self._log_add_cache_failure(failure)
def _convert_to_cached_embedding(
self,
embedding_response: Any,

View file

@ -1129,11 +1129,8 @@ class LLMCachingHandler:
Returns:
bool: True if the result should be stored in the cache, False otherwise.
"""
return (
(litellm.cache is not None)
and litellm.cache.supported_call_types is not None
and (str(original_function.__name__) in litellm.cache.supported_call_types)
and (kwargs.get("cache", {}).get("no-store", False) is not True)
return self._is_call_type_supported_by_cache(original_function=original_function) and (
kwargs.get("cache", {}).get("no-store", False) is not True
)
def wrap_streaming_result_for_cache(
@ -1170,13 +1167,11 @@ class LLMCachingHandler:
Returns:
bool: True if the call type is supported by the cache, False otherwise.
"""
if (
litellm.cache is not None
and litellm.cache.supported_call_types is not None
and str(original_function.__name__) in litellm.cache.supported_call_types
):
return True
return False
if litellm.cache is None or litellm.cache.supported_call_types is None:
return False
call_type: Final = str(original_function.__name__)
covering_call_types: Final = ("aresponses", "responses") if call_type == "aresponses" else (call_type,)
return any(name in litellm.cache.supported_call_types for name in covering_call_types)
async def _add_streaming_response_to_cache(self, processed_chunk: ModelResponse):
"""

View file

@ -8,21 +8,23 @@ Has 4 primary methods:
- async_get_cache
"""
import asyncio
import itertools
import logging
import time
from collections.abc import Sequence
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from threading import Lock
from typing import TYPE_CHECKING, Any, Final
if TYPE_CHECKING:
from litellm.types.caching import RedisPipelineIncrementOperation
import litellm
from litellm._logging import print_verbose, verbose_logger
from litellm.constants import DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE
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, active_request_redis_batch
from .redis_cache import RedisCache, RedisCircuitBreakerOpenError, log_redis_failure
if TYPE_CHECKING:
@ -47,6 +49,34 @@ class LimitedSizeOrderedDict(OrderedDict):
super().__setitem__(key, value)
@dataclass(frozen=True)
class PendingBatchRead:
"""A batch read that has consulted the in-memory tier and reserved its Redis keys, but not hit Redis yet."""
keys: list[str]
result: list[object | None]
redis_keys: list[str]
previous_access_times: dict[str, float | None]
@dataclass(frozen=True, slots=True)
class DeclaredBatchRead:
"""A ``async_batch_get_cache`` split in two: the memory half done, the Redis half declared on a ``RedisBatch``
so it rides that batch's next round trip, resolved later with ``async_resolve_batch_get``."""
keys: tuple[str, ...]
pending: PendingBatchRead
result: BatchResult[Mapping[str, object]] | None
def _log_deferred_increment_failure(future: asyncio.Future[float]) -> None:
if future.cancelled():
return
failure: Final = future.exception()
if failure is not None:
log_redis_failure(verbose_logger, logging.WARNING, "post-call Redis increment failed", failure)
class DualCache(BaseCache):
"""
DualCache is a cache implementation that updates both Redis and an in-memory cache simultaneously.
@ -249,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)
@ -301,59 +334,85 @@ class DualCache(BaseCache):
else:
self.last_redis_batch_access_time[key] = previous_time
async def _prepare_batch_get(
self, keys: list[str], local_only: bool, throttle_redis: bool = True, **kwargs: object
) -> PendingBatchRead:
result: list[object | None] = [None] * len(keys)
if self.in_memory_cache is not None:
in_memory_result: Final = await self.in_memory_cache.async_batch_get_cache(keys, **kwargs)
if in_memory_result is not None:
result = in_memory_result
redis_keys: list[str] = []
previous_access_times: dict[str, float | None] = {}
if None in result and self.redis_cache is not None and local_only is False:
if throttle_redis:
redis_keys, previous_access_times = self._reserve_redis_batch_keys(time.time(), keys, result)
else:
redis_keys = [key for key, value in zip(keys, result) if value is None]
return PendingBatchRead(
keys=keys, result=result, redis_keys=redis_keys, previous_access_times=previous_access_times
)
async def _apply_batch_get(
self, pending: PendingBatchRead, redis_result: Mapping[str, object] | None, **kwargs: object
) -> list[object | None]:
if redis_result is None or all(v is None for v in redis_result.values()):
return pending.result
merged: Final[list[object | None]] = [
redis_result.get(key, value) for key, value in zip(pending.keys, pending.result)
]
if self.in_memory_cache is not None:
for key, value in redis_result.items():
if value is not None:
await self.in_memory_cache.async_set_cache(key, value, **self._backfill_kwargs(kwargs))
return merged
async def declare_batch_get(self, keys: Sequence[str], batch: RedisBatch) -> DeclaredBatchRead:
pending: Final = await self._prepare_batch_get(
list(keys), # mutable-ok: the shared batch read takes a list
local_only=False,
throttle_redis=False,
)
return DeclaredBatchRead(
keys=tuple(keys),
pending=pending,
result=batch.mget(pending.redis_keys) if pending.redis_keys else None,
)
async def async_resolve_batch_get(self, declared: DeclaredBatchRead) -> list[object | None]:
redis_result: Final = None if declared.result is None else await declared.result
return await self._apply_batch_get(declared.pending, redis_result)
async def async_batch_get_cache(
self,
keys: list,
parent_otel_span: Span | None = None,
local_only: bool = False,
throttle_redis: bool = True,
**kwargs,
):
"""With ``throttle_redis`` False every key memory cannot serve is read from Redis, exactly as a per-key
``async_get_cache`` would read it, instead of skipping keys that missed within ``redis_batch_cache_expiry``."""
try:
result = [None] * len(keys)
if self.in_memory_cache is not None:
in_memory_result: Final = await self.in_memory_cache.async_batch_get_cache(keys, **kwargs)
if in_memory_result is not None:
result = in_memory_result
if None in result and self.redis_cache is not None and local_only is False:
"""
- for the none values in the result
- check the redis cache
"""
current_time: Final = time.time()
sublist_keys, previous_access_times = self._reserve_redis_batch_keys(current_time, keys, result)
# Only hit Redis if enough time has passed since last access.
if len(sublist_keys) > 0:
try:
# If not found in in-memory cache, try fetching from Redis
redis_result: Final = await self.redis_cache.async_batch_get_cache(
sublist_keys, parent_otel_span=parent_otel_span
)
except Exception as e:
# Do not throttle subsequent callers if the Redis read fails.
self._rollback_redis_batch_key_reservations(previous_access_times)
if isinstance(e, RedisCircuitBreakerOpenError):
verbose_logger.debug("LiteLLM Cache: async_batch_get_cache served from memory only: %s", e)
return result
raise
# Short-circuit if redis_result is None or contains only None values
if redis_result is None or all(v is None for v in redis_result.values()):
return result
# Pre-compute key-to-index mapping for O(1) lookup
key_to_index: Final = {key: i for i, key in enumerate(keys)}
# Update both result and in-memory cache in a single loop
for key, value in redis_result.items():
result[key_to_index[key]] = value
if value is not None and self.in_memory_cache is not None:
await self.in_memory_cache.async_set_cache(key, value, **self._backfill_kwargs(kwargs))
return result
pending: Final = await self._prepare_batch_get(keys, local_only, throttle_redis, **kwargs)
# Only hit Redis for keys memory could not serve and enough time has passed since last access.
if not pending.redis_keys or self.redis_cache is None:
return pending.result
try:
redis_result: Final = await self.redis_cache.async_batch_get_cache(
pending.redis_keys, parent_otel_span=parent_otel_span
)
except Exception as e:
# Do not throttle subsequent callers if the Redis read fails.
self._rollback_redis_batch_key_reservations(pending.previous_access_times)
if isinstance(e, RedisCircuitBreakerOpenError):
verbose_logger.debug("LiteLLM Cache: async_batch_get_cache served from memory only: %s", e)
return pending.result
raise
return await self._apply_batch_get(pending, redis_result, **kwargs)
except Exception as e:
log_redis_failure(
verbose_logger,
@ -363,6 +422,74 @@ class DualCache(BaseCache):
with_traceback=True,
)
@staticmethod
async def async_batch_get_cache_shared(
reads: Sequence[tuple["DualCache", list[str]]],
parent_otel_span: Span | None = None,
) -> list[list[object | None] | None]:
"""
`async_batch_get_cache` for several caches in one Redis round trip.
Each cache still serves what it can from its own in-memory tier, applies its own Redis read
throttle and backfills its own memory; only the Redis MGET is shared. A failed MGET is reported
to every cache that took part in it exactly as its own failed `async_batch_get_cache` would be:
None when the read raised, the in-memory result when the circuit breaker is open. A cache whose
Redis client is not the one the first cache uses falls back to its own read.
"""
results: Final[list[list[object | None] | None]] = [None] * len(reads)
shared_redis: Final = reads[0][0].redis_cache if reads else None
pendings: Final[list[tuple[int, DualCache, PendingBatchRead]]] = []
for index, (cache, keys) in enumerate(reads):
if shared_redis is None or cache.redis_cache is not shared_redis:
results[index] = await cache.async_batch_get_cache(keys=keys, parent_otel_span=parent_otel_span)
continue
try:
pending = await cache._prepare_batch_get(keys, local_only=False)
except Exception as e:
DualCache._log_shared_batch_get_failure(e)
continue
pendings.append((index, cache, pending))
results[index] = pending.result
redis_keys: Final = list(
dict.fromkeys(itertools.chain.from_iterable(pending.redis_keys for _, _, pending in pendings))
)
if shared_redis is None or not redis_keys:
return results
try:
redis_result: Final = await shared_redis.async_batch_get_cache(
redis_keys, parent_otel_span=parent_otel_span
)
except Exception as e:
for index, cache, pending in pendings:
cache._rollback_redis_batch_key_reservations(pending.previous_access_times)
if pending.redis_keys and not isinstance(e, RedisCircuitBreakerOpenError):
results[index] = None
if isinstance(e, RedisCircuitBreakerOpenError):
verbose_logger.debug("LiteLLM Cache: async_batch_get_cache_shared served from memory only: %s", e)
else:
DualCache._log_shared_batch_get_failure(e)
return results
for index, cache, pending in pendings:
own_result = {key: redis_result[key] for key in pending.redis_keys if key in redis_result}
try:
results[index] = await cache._apply_batch_get(pending, own_result)
except Exception as e:
results[index] = None
DualCache._log_shared_batch_get_failure(e)
return results
@staticmethod
def _log_shared_batch_get_failure(e: Exception) -> None:
log_redis_failure(
verbose_logger,
logging.ERROR,
"LiteLLM Cache: exception in async_batch_get_cache_shared",
e,
with_traceback=True,
)
async def async_set_cache(self, key, value, local_only: bool = False, **kwargs):
print_verbose(f"async set cache: cache key: {key}; local_only: {local_only}; value: {value}")
try:
@ -378,6 +505,34 @@ 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)
return batch.set(key, value, effective_ttl)
# async_batch_set_cache
async def async_set_cache_pipeline(
self, cache_list: Sequence[tuple[str, object]], local_only: bool = False, **kwargs
@ -445,6 +600,41 @@ class DualCache(BaseCache):
)
return result
async def async_increment_cache_post_call(
self,
key: str,
value: float,
ttl: int | None,
parent_otel_span: Span | None = None,
) -> None:
"""Memory is incremented now; the Redis increment rides the request's post-call pipeline when one is
open, and runs on its own as ``async_increment_cache`` otherwise."""
await self.async_increment_cache_pipeline_post_call(
(RedisPipelineIncrementOperation(key=key, increment_value=value, ttl=ttl),), parent_otel_span
)
async def async_increment_cache_pipeline_post_call(
self,
increment_list: Sequence["RedisPipelineIncrementOperation"],
parent_otel_span: Span | None = None,
) -> None:
batch: Final = None if self.redis_cache is None else active_post_call_redis_batch(self.redis_cache)
operations: Final = list(increment_list) # mutable-ok: both increment pipelines take a list
if batch is None:
await self.async_increment_cache_pipeline(operations, parent_otel_span=parent_otel_span)
return
try:
if self.in_memory_cache is not None:
await self.in_memory_cache.async_increment_pipeline(
increment_list=operations, parent_otel_span=parent_otel_span
)
except Exception as e: # noqa: BLE001 # same tolerance as async_increment_cache_pipeline
log_redis_failure(verbose_logger, logging.WARNING, "in-memory increment failed", e)
for operation in increment_list:
batch.increment(operation["key"], operation["increment_value"], operation["ttl"]).on_settled(
_log_deferred_increment_failure
)
async def async_increment_cache_pipeline(
self,
increment_list: list["RedisPipelineIncrementOperation"],

View file

@ -0,0 +1,555 @@
"""One Redis pipeline for several independent operations, each with its own result and its own failure.
A ``RedisBatch`` collects MGETs, Lua scripts and increments declared by unrelated callers and sends them
in one ``pipeline(transaction=False)`` round trip. Every declaration returns an awaitable; awaiting one
flushes whatever has been declared so far, so callers keep their existing ``await`` shape and their own
error handling while sharing the wire. Redis Cluster clients run each operation on its own, as before:
a cluster pipeline is per node anyway and the existing per-operation paths already group by slot.
"""
from __future__ import annotations
import asyncio
import hashlib
import json
import logging
import time
import weakref
from collections.abc import Awaitable, Callable, Generator, Mapping, Sequence
from contextvars import ContextVar, Token
from dataclasses import dataclass, field
from datetime import timedelta
from types import MappingProxyType, TracebackType
from typing import Final, Generic, Protocol, TypeVar
from litellm._logging import verbose_logger
from litellm.caching.redis_cache import (
RedisCache,
_run_under_circuit_breaker, # pyright: ignore[reportPrivateUsage] # same health signal as every RedisCache method
log_redis_failure,
)
from litellm.caching.redis_cluster_cache import RedisClusterCache
from litellm.types.services import ServiceTypes
_T = TypeVar("_T")
_ScriptArg = str | bytes | int | float
SettledHook = Callable[[asyncio.Future[_T]], Awaitable[None] | None] # mutable-ok: Callable params
POST_CALL_FLUSH_DEADLINE_SECONDS: Final = 1.0
class RegisteredScript(Protocol):
def __call__(self, keys: Sequence[str], args: Sequence[_ScriptArg]) -> Awaitable[object]: ...
class _RedisPipeline(Protocol):
def mget(self, keys: Sequence[str]) -> object: ...
def evalsha(self, sha: str, numkeys: int, *keys_and_args: _ScriptArg) -> object: ...
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]: ...
class _Op(Generic[_T]):
"""One declared operation: how many pipeline replies it consumes, how to turn them into a result, and
how to run on its own when the batch cannot pipeline (cluster client, or a reply the pipeline cannot
settle, like NOSCRIPT)."""
__slots__ = ("future", "settled_hooks")
def __init__(self) -> None:
self.future: Final[asyncio.Future[_T]] = asyncio.get_running_loop().create_future()
self.future.add_done_callback(_mark_retrieved)
self.settled_hooks: Final[list[SettledHook[_T]]] = [] # mutable-ok: append-only registry
async def run_settled_hooks(self) -> None:
for hook in self.settled_hooks:
await self._run_settled_hook(hook)
async def _run_settled_hook(self, hook: SettledHook[_T]) -> None:
try:
follow_up: Final = hook(self.future)
if follow_up is not None:
await follow_up
except Exception as e: # noqa: BLE001 # one owner's follow-up must not stop the others
verbose_logger.warning("redis batch settled hook failed: %s", e)
def enqueue(self, pipe: _RedisPipeline) -> int:
raise NotImplementedError
def resolve(self, replies: Sequence[object]) -> _T:
raise NotImplementedError
async def run_alone(self) -> _T:
raise NotImplementedError
def settle(self, replies: Sequence[object]) -> Awaitable[None] | None:
"""Resolve from pipeline replies; return a coroutine when the op has to be retried on its own."""
failure: Final = next((reply for reply in replies if isinstance(reply, Exception)), None)
if failure is None:
try:
self.future.set_result(self.resolve(replies))
except Exception as e: # noqa: BLE001 # a reply this op cannot decode fails this op alone
self.future.set_exception(e)
return None
if _is_missing_script(failure):
return self._settle_alone()
self.future.set_exception(failure)
return None
async def _settle_alone(self) -> None:
try:
self.future.set_result(await self.run_alone())
except Exception as e: # noqa: BLE001 # the declaring caller owns the failure of its own operation
self.future.set_exception(e)
def _is_missing_script(failure: Exception) -> bool:
"""Imported lazily: this module is reachable from a base ``import litellm`` while redis is not a base dependency."""
from redis.exceptions import NoScriptError
return isinstance(failure, NoScriptError)
def _mark_retrieved(future: asyncio.Future[object]) -> None:
"""A caller that stops awaiting (cancelled request) must not leave an 'exception never retrieved' log."""
if not future.cancelled():
future.exception()
class _MGet(_Op[Mapping[str, object]]):
__slots__ = ("_keys", "_redis_cache")
def __init__(self, redis_cache: RedisCache, keys: Sequence[str]) -> None:
super().__init__()
self._redis_cache: Final = redis_cache
self._keys: Final[tuple[str, ...]] = tuple(dict.fromkeys(keys))
def enqueue(self, pipe: _RedisPipeline) -> int:
pipe.mget(tuple(self._redis_cache.check_and_fix_namespace(key=key) for key in self._keys))
return 1
def resolve(self, replies: Sequence[object]) -> Mapping[str, object]:
values: Final = replies[0]
if not isinstance(values, (list, tuple)):
raise TypeError(f"MGET reply is not a list: {type(values).__name__}")
return MappingProxyType(
{key: self._redis_cache._get_cache_logic(value) for key, value in zip(self._keys, values)} # pyright: ignore[reportPrivateUsage, reportUnknownMemberType, reportUnknownArgumentType] # shared decode with async_batch_get_cache
)
async def run_alone(self) -> Mapping[str, object]:
found: Mapping[str, object] = await self._redis_cache.async_batch_get_cache(key_list=list(self._keys)) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # untyped cache API # mutable-ok: the cache API takes a list
if any(key not in found for key in self._keys):
raise ConnectionError("batch get did not return every key")
return found
class _Script(_Op[object]):
__slots__ = ("_args", "_keys", "_redis_cache", "_run", "_sha")
def __init__(
self,
redis_cache: RedisCache,
source: str,
run: RegisteredScript,
keys: Sequence[str],
args: Sequence[_ScriptArg],
) -> None:
super().__init__()
self._redis_cache: Final = redis_cache
self._sha: Final = hashlib.sha1(source.encode()).hexdigest() # noqa: S324 # EVALSHA identifies scripts by SHA-1
self._run: Final = run
self._keys: Final[tuple[str, ...]] = tuple(keys)
self._args: Final[tuple[_ScriptArg, ...]] = tuple(args)
def enqueue(self, pipe: _RedisPipeline) -> int:
namespaced: Final = tuple(self._redis_cache.check_and_fix_namespace(key=key) for key in self._keys)
pipe.evalsha(self._sha, len(namespaced), *namespaced, *self._args)
return 1
def resolve(self, replies: Sequence[object]) -> object:
return replies[0]
async def run_alone(self) -> object:
return await self._run(keys=self._keys, args=self._args)
class _Increment(_Op[float]):
__slots__ = ("_key", "_redis_cache", "_ttl", "_value")
def __init__(self, redis_cache: RedisCache, key: str, value: float, ttl: int | None) -> None:
super().__init__()
self._redis_cache: Final = redis_cache
self._key: Final = key
self._value: Final = value
self._ttl: Final = ttl
def enqueue(self, pipe: _RedisPipeline) -> int:
name: Final = self._redis_cache.check_and_fix_namespace(key=self._key)
pipe.incrbyfloat(name, self._value)
if self._ttl is None:
return 1
pipe.expire(name, timedelta(seconds=self._ttl))
return 2
def resolve(self, replies: Sequence[object]) -> float:
reply: Final = replies[0]
if not isinstance(reply, (int, float, str, bytes)):
raise TypeError(f"INCRBYFLOAT reply is not numeric: {type(reply).__name__}")
return float(reply)
async def run_alone(self) -> float:
value: object = await self._redis_cache.async_increment(key=self._key, value=self._value, ttl=self._ttl) # pyright: ignore[reportUnknownMemberType] # untyped cache API
if not isinstance(value, (int, float)):
raise TypeError(f"increment did not return a number: {type(value).__name__}")
return float(value)
class _Set(_Op[None]):
"""SET with the cache's TTL rules, same encoding as ``async_set_cache_pipeline_with_ttls``."""
__slots__ = ("_key", "_redis_cache", "_ttl", "_value")
def __init__(self, redis_cache: RedisCache, key: str, value: object, ttl: float | None) -> None:
super().__init__()
self._redis_cache: Final = redis_cache
self._key: Final = key
self._value: Final = value
self._ttl: Final = ttl
def enqueue(self, pipe: _RedisPipeline) -> int:
ttl: Final = self._redis_cache.get_ttl(ttl=self._ttl) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # untyped cache API
pipe.set(
self._redis_cache.check_and_fix_namespace(key=self._key),
json.dumps(self._value),
ex=None if ttl is None else timedelta(seconds=ttl),
)
return 1
def resolve(self, replies: Sequence[object]) -> None:
return None
async def run_alone(self) -> 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."""
__slots__ = ("_batch", "_op")
def __init__(self, batch: RedisBatch, op: _Op[_T]) -> None:
self._batch: Final = batch
self._op: Final = op
def __await__(self) -> Generator[object, None, _T]:
return self._wait().__await__()
async def _wait(self) -> _T:
if not self._op.future.done():
await self._batch.flush()
return self._op.future.result()
@property
def done(self) -> bool:
return self._op.future.done()
def on_settled(self, hook: SettledHook[_T]) -> None:
"""For an owner that does not await: runs inside the flush once this operation has its result or
failure (or was cancelled with the pipeline), so the flush completes with the follow-up done."""
self._op.settled_hooks.append(hook)
@dataclass(slots=True)
class RedisBatch:
"""Operations declared here go out in one pipeline the next time any of them is awaited or ``flush`` runs."""
redis_cache: RedisCache
name: str = "redis_batch"
_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]]:
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]
) -> BatchResult[object]:
return self._declare(_Script(self.redis_cache, source, run, keys, args))
def increment(self, key: str, value: float, ttl: int | None = None) -> BatchResult[float]:
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)
@property
def pending(self) -> int:
return len(self._pending)
def _declare(self, op: _Op[_T]) -> BatchResult[_T]:
self._pending.append(op) # pyright: ignore[reportArgumentType] # heterogeneous ops share the flush loop
return BatchResult(self, op)
async def flush(self) -> None:
async with self._lock:
for hook in self._flush_hooks:
hook()
ops: Final = tuple(self._pending)
self._pending.clear()
if not ops:
return
self.flushes += 1
try:
if isinstance(self.redis_cache, RedisClusterCache):
await asyncio.gather(*(op._settle_alone() for op in ops)) # pyright: ignore[reportPrivateUsage] # batch owns its ops
else:
await self._flush_pipeline(ops)
finally:
for op in ops:
if not op.future.done():
op.future.cancel()
await asyncio.gather(*(op.run_settled_hooks() for op in ops))
async def _flush_pipeline(self, ops: Sequence[_Op[object]]) -> None:
start_time: Final = time.time()
widths: list[int] = [] # mutable-ok: filled while enqueuing
async def run() -> list[object]:
client: Final = self.redis_cache.init_async_client()
async with client.pipeline(transaction=False) as pipe:
widths.extend(op.enqueue(pipe) for op in ops)
return await pipe.execute(raise_on_error=False)
try:
replies: Final = await _run_under_circuit_breaker(self.redis_cache._circuit_breaker, self.name, run) # pyright: ignore[reportPrivateUsage] # same breaker as the cache's own methods
except Exception as e: # noqa: BLE001 # each declaring caller applies its own Redis fallback
log_redis_failure(verbose_logger, logging.WARNING, f"{self.name}: pipeline of {len(ops)} ops failed", e)
asyncio.create_task(
self.redis_cache.service_logger_obj.async_service_failure_hook(
service=ServiceTypes.REDIS,
duration=time.time() - start_time,
error=e,
call_type=f"{self.name}[{len(ops)}]",
start_time=start_time,
end_time=time.time(),
)
)
for op in ops:
op.future.set_exception(e)
return
asyncio.create_task(
self.redis_cache.service_logger_obj.async_service_success_hook(
service=ServiceTypes.REDIS,
duration=time.time() - start_time,
call_type=f"{self.name}[{len(ops)}]",
start_time=start_time,
end_time=time.time(),
)
)
retries: list[Awaitable[None]] = [] # mutable-ok: collected while slicing replies
offset = 0
for op, width in zip(ops, widths):
retry = op.settle(replies[offset : offset + width])
offset += width
if retry is not None:
retries.append(retry)
if retries:
await asyncio.gather(*retries)
def _backend_key(redis_cache: RedisCache) -> object:
"""Two ``RedisCache`` instances built from the same connection settings and namespace talk to the same server
under the same key prefix, so the proxy's cache and the router's cache share one pipeline (the router gets its
port as a string, hence the ``str`` comparison); a cache whose settings cannot be compared (a test double) gets
its own."""
try:
settings: Final = tuple(sorted((str(k), str(v)) for k, v in redis_cache.redis_kwargs.items() if v is not None)) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType, reportUnknownArgumentType] # untyped cache API
except AttributeError:
return ("instance", id(redis_cache))
return (type(redis_cache), redis_cache.namespace, settings)
_open_post_call: Final[weakref.WeakSet[RequestRedisBatches]] = weakref.WeakSet()
"""Requests whose post-call batch still holds declared ops, so a shutdown can send them before Redis goes away."""
class RequestRedisBatches:
"""One ``RedisBatch`` per Redis backend for the current request, so readers of different caches that
share a server (the proxy's and the router's) share the pipeline.
The post-call batches hold the writes nothing waits on (counters, token scripts, the response cache).
They flush once, when the success or failure callbacks have all run, or at ``post_call_deadline``
seconds after the first declaration when no callback phase closes them."""
__slots__ = (
"__weakref__",
"_batches",
"_deadline",
"_deadline_flush",
"_post_call",
"post_call_deadline",
"prefetched",
)
def __init__(self, post_call_deadline: float = POST_CALL_FLUSH_DEADLINE_SECONDS) -> None:
self._batches: Final[dict[object, RedisBatch]] = {} # mutable-ok: lazily filled per backend
self._post_call: Final[dict[object, RedisBatch]] = {} # mutable-ok: lazily filled per backend
self.post_call_deadline: Final = post_call_deadline
self._deadline: asyncio.TimerHandle | None = None
self._deadline_flush: asyncio.Task[None] | None = None
# Reads declared early for a consumer that runs later in the request, keyed by consumer name.
self.prefetched: Final[dict[str, object]] = {} # mutable-ok: armed pre-admission, taken at use
def batch(self, redis_cache: RedisCache) -> RedisBatch:
key: Final = _backend_key(redis_cache)
batch = self._batches.get(key)
if batch is None:
batch = RedisBatch(redis_cache, name="request_redis_batch")
self._batches[key] = batch
return batch
def post_call(self, redis_cache: RedisCache) -> RedisBatch:
key: Final = _backend_key(redis_cache)
existing: Final = self._post_call.get(key)
batch: Final = (
existing
if existing is not None
else self._post_call.setdefault(key, RedisBatch(redis_cache, name="post_call_redis_batch"))
)
if self._deadline is None:
self._deadline = asyncio.get_running_loop().call_later(self.post_call_deadline, self._flush_on_deadline)
_open_post_call.add(self)
return batch
def _flush_on_deadline(self) -> None:
self._deadline = None
self._deadline_flush = asyncio.ensure_future(self.flush_post_call())
async def flush_all(self) -> None:
"""Send whatever is still declared (write-backs nobody awaits) before the request scope closes."""
await asyncio.gather(*(batch.flush() for batch in self._batches.values() if batch.pending))
async def flush_post_call(self) -> None:
"""One pipeline per backend for the post-call writes; the deadline is disarmed since this is that flush."""
if self._deadline is not None:
self._deadline.cancel()
self._deadline = None
await asyncio.gather(*(batch.flush() for batch in self._post_call.values() if batch.pending))
if not any(batch.pending for batch in self._post_call.values()):
_open_post_call.discard(self)
@property
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
)
def active_request_redis_batch(redis_cache: RedisCache) -> RedisBatch | None:
"""The request's batch for this backend, or None outside a ``request_redis_batch_scope``."""
batches: Final = _active_request_batches.get()
if batches is None:
return None
return batches.batch(redis_cache)
def active_request_redis_batches() -> RequestRedisBatches | None:
return _active_request_batches.get()
def active_post_call_redis_batch(redis_cache: RedisCache) -> RedisBatch | None:
"""The request's post-call batch for this backend, or None outside a ``request_redis_batch_scope``."""
batches: Final = _active_request_batches.get()
if batches is None:
return None
return batches.post_call(redis_cache)
async def flush_post_call_redis_batches() -> None:
"""Called where the success and failure callbacks of a request have all run."""
batches: Final = _active_request_batches.get()
if batches is not None:
await batches.flush_post_call()
async def drain_post_call_redis_batches() -> None:
"""Sends every post-call batch still waiting on its callbacks or deadline; for the shutdown path."""
await asyncio.gather(*(batches.flush_post_call() for batches in tuple(_open_post_call)))
class request_redis_batch_scope:
"""Redis reads declared inside share one pipeline per backend; nested scopes join the outer one."""
__slots__ = ("_post_call_deadline", "_token")
def __init__(self, post_call_deadline: float = POST_CALL_FLUSH_DEADLINE_SECONDS) -> None:
self._token: Token[RequestRedisBatches | None] | None = None
self._post_call_deadline: Final = post_call_deadline
def __enter__(self) -> RequestRedisBatches:
outer: Final = _active_request_batches.get()
if outer is not None:
return outer
batches: Final = RequestRedisBatches(post_call_deadline=self._post_call_deadline)
self._token = _active_request_batches.set(batches)
return batches
def __exit__(
self, exc_type: type[BaseException] | None, exc: BaseException | None, tb: TracebackType | None
) -> None:
if self._token is not None:
_active_request_batches.reset(self._token)

View file

@ -1334,6 +1334,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
):
super().__init__(streaming_response, sync_stream, json_mode)
self._chat_completion_id: str | None = None
self._served_service_tier: str | None = None
self._tool_call_index_map: dict[int, int] = {} # mutable-ok: per-stream accumulator state
def _handle_string_chunk(
@ -1598,6 +1599,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(response_data.get("usage"))
provider_metadata: Final = _provider_metadata(response_data)
served_service_tier: Final = response_data.get("service_tier")
return ModelResponseStream(
choices=[
StreamingChoices(
@ -1611,6 +1613,11 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
],
usage=usage,
provider_specific_fields=dict(provider_metadata) or None, # mutable-ok: field is typed dict
**(
MappingProxyType({"service_tier": served_service_tier})
if isinstance(served_service_tier, str)
else MappingProxyType({})
),
)
else:
pass
@ -1639,12 +1646,28 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
ModelResponseStream: OpenAI-formatted streaming chunk
"""
verbose_logger.debug("Chat provider: transform_streaming_response called with chunk: %s", chunk)
return self._with_stream_scoped_id(
OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(
chunk, tool_call_index_map=self._tool_call_index_map
self._remember_served_service_tier(chunk)
return self._with_served_service_tier(
self._with_stream_scoped_id(
OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(
chunk, tool_call_index_map=self._tool_call_index_map
)
)
)
def _remember_served_service_tier(self, chunk: dict[str, object]) -> None:
response_payload: Final = chunk.get("response")
if not isinstance(response_payload, dict):
return
served_tier: Final = response_payload.get("service_tier")
if isinstance(served_tier, str) and served_tier:
self._served_service_tier = served_tier
def _with_served_service_tier(self, chunk: "ModelResponseStream") -> "ModelResponseStream":
if self._served_service_tier is not None and chunk.model_dump().get("service_tier") is None:
setattr(chunk, "service_tier", self._served_service_tier) # noqa: B010 # pydantic extra, not a declared field
return chunk
def _with_stream_scoped_id(self, chunk: "ModelResponseStream") -> "ModelResponseStream":
if self._chat_completion_id is None:
self._chat_completion_id = chunk.id

View file

@ -26,6 +26,7 @@ from litellm.litellm_core_utils.llm_cost_calc.usage_object_transformation import
TranscriptionUsageObjectTransformation,
)
from litellm.litellm_core_utils.llm_cost_calc.utils import (
_SERVICE_TIER_TO_COST_KEY_SUFFIX,
BilledTokenRates,
CostCalculatorUtils,
_generic_cost_per_character,
@ -351,19 +352,27 @@ def _per_second_pricing_cost(
return None
if _has_token_or_tiered_pricing(model_info) or not _bills_wall_clock_seconds(model_info):
return None
cost_per_second: Final = model_info.get("cost_per_second")
input_cost_per_second: Final = model_info.get("input_cost_per_second")
output_cost_per_second: Final = model_info.get("output_cost_per_second")
if input_cost_per_second is None and output_cost_per_second is None:
resolved_cost_per_second: Final = (
cost_per_second
if cost_per_second is not None
else input_cost_per_second
if input_cost_per_second is not None
else output_cost_per_second
)
if resolved_cost_per_second is None:
return None
seconds: Final = (response_time_ms or 0.0) / 1000
verbose_logger.debug(
"For model=%s - input_cost_per_second: %s; output_cost_per_second: %s; response time: %s",
"For model=%s - cost_per_second: %s; response time: %s",
model,
input_cost_per_second,
output_cost_per_second,
resolved_cost_per_second,
response_time_ms,
)
return (input_cost_per_second or 0.0) * seconds, (output_cost_per_second or 0.0) * seconds
return resolved_cost_per_second * seconds, 0.0
def cost_per_token(
@ -696,7 +705,7 @@ def cost_per_token(
data_residency=data_residency,
)
elif custom_llm_provider == "databricks":
return databricks_cost_per_token(model=model, usage=usage_block)
return databricks_cost_per_token(model=model, usage=usage_block, service_tier=service_tier)
elif custom_llm_provider == "fireworks_ai":
return fireworks_ai_cost_per_token(model=model, usage=usage_block)
elif custom_llm_provider == "azure":
@ -790,7 +799,9 @@ def _get_hidden_str_for_cost_calc(hidden_params: object, key: str) -> str | None
return value if isinstance(value, str) and value else None
_NON_TOKEN_RATE_FIELDS: Final = frozenset({"input_cost_per_second", "input_cost_per_query", "tiered_pricing"})
_NON_TOKEN_RATE_FIELDS: Final = frozenset(
{"cost_per_second", "input_cost_per_second", "output_cost_per_second", "input_cost_per_query", "tiered_pricing"}
)
def _cost_map_entry_prices_anything(entry: Mapping[str, object]) -> bool:
@ -959,6 +970,37 @@ def _normalize_service_tier(service_tier: object) -> str | None:
return service_tier
_BASE_PRICING_SERVICE_TIERS: Final[frozenset[str]] = frozenset({"default", "standard"})
def _resolve_billable_service_tier(requested: object, served: object) -> str | None:
"""Served tier wins when it names a priced tier or explicitly says base pricing; otherwise the request decides."""
served_lower: Final = served.lower() if isinstance(served, str) else None
if served_lower is not None and served_lower in _SERVICE_TIER_TO_COST_KEY_SUFFIX:
return served_lower
if served_lower in _BASE_PRICING_SERVICE_TIERS:
return None
return _normalize_service_tier(requested)
def _served_service_tier(completion_response: object, usage_object: Usage | None) -> str | None:
"""Find the tier the provider actually served: response, then usage, then Gemini trafficType."""
response_tier: Final = _extract_service_tier(completion_response)
if isinstance(response_tier, str):
return response_tier
usage_tier: Final = _extract_service_tier(usage_object)
if isinstance(usage_tier, str):
return usage_tier
hidden_params: Final = getattr(completion_response, "_hidden_params", None)
if hidden_params is None:
return None
provider_specific: Final = hidden_params.get("provider_specific_fields") or {}
raw_traffic_type: Final = provider_specific.get("traffic_type")
if not raw_traffic_type:
return None
return _map_traffic_type_to_service_tier(raw_traffic_type) or "default"
def _extract_service_tier(source: object) -> str | None:
"""Read a raw ``service_tier`` off a response body or usage object, dict or pydantic model alike."""
if isinstance(source, BaseModel):
@ -1378,23 +1420,14 @@ def completion_cost(
)
rerank_billed_units: RerankBilledUnits | None = None
# Extract service_tier from optional_params if not provided directly
if service_tier is None and optional_params is not None:
service_tier = optional_params.get("service_tier")
service_tier = _normalize_service_tier(service_tier)
# Extract service_tier from completion_response if not provided
if service_tier is None and completion_response is not None:
service_tier = _extract_service_tier(completion_response)
service_tier = _normalize_service_tier(service_tier)
# Extract service_tier from usage object if not provided
if service_tier is None and cost_per_token_usage_object is not None:
service_tier = _extract_service_tier(cost_per_token_usage_object)
service_tier = _normalize_service_tier(service_tier)
explicit_tier: Final = _normalize_service_tier(service_tier)
if explicit_tier is not None:
service_tier = explicit_tier
else:
service_tier = _resolve_billable_service_tier( # rebind-ok: resolved from request then response
requested=optional_params.get("service_tier") if optional_params is not None else None,
served=_served_service_tier(completion_response, cost_per_token_usage_object),
)
explicit_pricing: Final = custom_pricing is True or base_model is not None
selected_model: Final = _select_model_name_for_cost_calc(
@ -1484,15 +1517,6 @@ def completion_cost(
custom_llm_provider = hidden_params.get("custom_llm_provider", custom_llm_provider or None)
region_name = hidden_params.get("region_name", region_name)
# For Gemini/Vertex AI responses, trafficType is stored in
# provider_specific_fields. Map it to the service_tier used
# by the cost key lookup (_priority / _flex suffixes) so that
# ON_DEMAND_PRIORITY requests are billed at priority prices.
if service_tier is None:
provider_specific = hidden_params.get("provider_specific_fields") or {}
raw_traffic_type = provider_specific.get("traffic_type")
if raw_traffic_type:
service_tier = _map_traffic_type_to_service_tier(raw_traffic_type)
else:
if model is None:
raise ValueError(

View file

@ -376,8 +376,12 @@ class SlackAlerting(CustomBatchLogger):
if combined_metrics_values is None:
return False
metric_values: Final[list[float | None]] = [
val if isinstance(val, (int, float)) else None for val in combined_metrics_values
]
all_none = True
for val in combined_metrics_values:
for val in metric_values:
if val is not None and val > 0:
all_none = False
break
@ -385,8 +389,8 @@ class SlackAlerting(CustomBatchLogger):
if all_none:
return False
failed_request_values: Final = combined_metrics_values[: len(failed_request_keys)] # # [1, 2, None, ..]
latency_values: Final = combined_metrics_values[len(failed_request_keys) :]
failed_request_values: Final = metric_values[: len(failed_request_keys)] # # [1, 2, None, ..]
latency_values: Final = metric_values[len(failed_request_keys) :]
# find top 5 failed
## Replace None values with a placeholder value (-1 in this case)

View file

@ -155,6 +155,7 @@ def get_litellm_params(
allm_passthrough_route=None,
preset_cache_key=None,
no_log=None,
cost_per_second: float | None = None,
input_cost_per_second=None,
input_cost_per_token=None,
output_cost_per_token=None,
@ -216,6 +217,7 @@ def get_litellm_params(
"preset_cache_key": preset_cache_key,
"no-log": no_log or kwargs.get("no-log"),
"stream_response": {}, # litellm_call_id: ModelResponse Dict
"cost_per_second": cost_per_second,
"input_cost_per_token": input_cost_per_token,
"input_cost_per_second": input_cost_per_second,
"output_cost_per_token": output_cost_per_token,

View file

@ -35,6 +35,7 @@ from litellm._uuid import uuid
from litellm.batches.batch_utils import _handle_completed_batch, batch_cost_is_final
from litellm.caching.caching import DualCache
from litellm.caching.caching_handler import LLMCachingHandler
from litellm.caching.redis_batch import flush_post_call_redis_batches
from litellm.constants import (
DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT,
DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT,
@ -1884,7 +1885,6 @@ class Logging(LiteLLMLoggingBaseClass):
"standard_built_in_tools_params": self.standard_built_in_tools_params,
"router_model_id": router_model_id,
"litellm_logging_obj": self,
"service_tier": (self.optional_params.get("service_tier") if self.optional_params else None),
"data_residency": (
self.litellm_params.get("data_residency")
if hasattr(self, "litellm_params") and self.litellm_params
@ -3553,6 +3553,7 @@ class Logging(LiteLLMLoggingBaseClass):
traceback.format_exc(),
)
self._handle_callback_failure(callback=callback)
await flush_post_call_redis_batches()
def _handle_callback_failure(self, callback: object):
"""
@ -3938,6 +3939,7 @@ class Logging(LiteLLMLoggingBaseClass):
)
# Track callback logging failures in Prometheus
self._handle_callback_failure(callback=callback)
await flush_post_call_redis_batches()
def _get_trace_id(self, service_name: Literal["langfuse"]) -> str | None:
"""

View file

@ -109,6 +109,7 @@ class _BaseChunk(TypedDict, total=False):
created: ReadOnly[int]
model: ReadOnly[str]
system_fingerprint: ReadOnly[str | None]
service_tier: ReadOnly[str | None]
choices: ReadOnly[Required[Sequence[StreamingChoices]]]
_hidden_params: ReadOnly[_ChunkHiddenParams]
@ -369,6 +370,13 @@ class ChunkProcessor:
# Fall back to first chunk's model if no different model found
return first_chunk_model
@staticmethod
def _get_service_tier_from_chunks(chunks: Sequence["_BaseChunk"]) -> str | None:
return next(
(tier for chunk in reversed(chunks) if isinstance(tier := chunk.get("service_tier"), str) and tier),
None,
)
def build_base_response(self, chunks: Sequence["_BaseChunk"]) -> ModelResponse:
chunk = self.first_chunk
id: Final = ChunkProcessor._get_chunk_id(chunks)
@ -378,6 +386,7 @@ class ChunkProcessor:
# Get the actual model - for Azure Model Router, this finds the real model from later chunks
model: Final = ChunkProcessor._get_model_from_chunks(chunks, first_chunk_model)
system_fingerprint: Final = chunk.get("system_fingerprint", None)
service_tier: Final = ChunkProcessor._get_service_tier_from_chunks(chunks)
role: Final = ChunkProcessor._get_role_from_chunks(chunks)
finish_reason = "stop"
@ -399,6 +408,11 @@ class ChunkProcessor:
"created": created,
"model": model,
"system_fingerprint": system_fingerprint,
**(
MappingProxyType({"service_tier": service_tier})
if service_tier is not None
else MappingProxyType({})
),
"choices": [
{
"index": 0,

View file

@ -72,6 +72,12 @@ def _next_sync_or_exhausted(it: Any) -> object:
return _SYNC_ITER_EXHAUSTED
def _stamp_served_service_tier(response: ModelResponseStream, complete_streaming_response: ModelResponse) -> None:
served_tier: Final = complete_streaming_response.model_dump().get("service_tier")
if isinstance(served_tier, str) and served_tier:
setattr(response, "service_tier", served_tier) # noqa: B010 # pydantic extra, not a declared field
def is_async_iterable(obj: object) -> bool:
"""
Check if an object is an async iterable (can be used with 'async for').
@ -1876,6 +1882,7 @@ class CustomStreamWrapper:
"usage",
getattr(complete_streaming_response, "usage"),
)
_stamp_served_service_tier(response, complete_streaming_response)
try:
_cache_copy = complete_streaming_response.model_copy(deep=True)
_log_copy = complete_streaming_response.model_copy(deep=True)
@ -2127,6 +2134,7 @@ class CustomStreamWrapper:
"usage",
getattr(complete_streaming_response, "usage"),
)
_stamp_served_service_tier(response, complete_streaming_response)
try:
_copy = complete_streaming_response.model_copy(deep=True)
except RuntimeError:

View file

@ -11,6 +11,7 @@ from typing import (
Final,
Literal,
Protocol,
cast,
get_args,
)
@ -35,6 +36,7 @@ from litellm.types.utils import AdapterCompletionStreamWrapper, Delta
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObject
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ModelResponseStream
@ -115,6 +117,18 @@ class _CombinedChunkSplitter:
self._async_iter: AsyncIterator[ModelResponseStream] | None = None
self._buffer: deque[ModelResponseStream] = deque()
@property
def chunks(self) -> "list[ModelResponseStream] | None":
return cast( # cast-ok: chunks is a list of ModelResponseStream on the inner stream
"list[ModelResponseStream] | None", getattr(self._stream, "chunks", None)
)
@property
def messages(self) -> "list[AllMessageValues] | None":
return cast( # cast-ok: messages is a list of AllMessageValues on the inner stream
"list[AllMessageValues] | None", getattr(self._stream, "messages", None)
)
@staticmethod
def _is_combined(chunk: "ModelResponseStream") -> bool:
"""True if ``chunk`` carries response content AND a finish_reason."""
@ -351,6 +365,18 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
text="",
)
@property
def chunks(self) -> "list[ModelResponseStream] | None":
return cast( # cast-ok: chunks is a list of ModelResponseStream on the inner stream
"list[ModelResponseStream] | None", getattr(self.completion_stream, "chunks", None)
)
@property
def messages(self) -> "list[AllMessageValues] | None":
return cast( # cast-ok: messages is a list of AllMessageValues on the inner stream
"list[AllMessageValues] | None", getattr(self.completion_stream, "messages", None)
)
def _merge_usage_into_held_stop_reason_chunk(self, chunk: Any) -> MessageBlockDelta:
"""Merge usage data from ``chunk`` into the held ``message_delta`` chunk.
@ -1173,3 +1199,37 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
return True
return False
class AnthropicSSEStream(AsyncIterator[bytes]):
"""
AsyncIterator[bytes] view of AnthropicStreamWrapper returned to callers of
translate_completion_output_params_streaming. Keeps the wrapper reachable so
the proxy's disconnect-time partial billing can read the inner chat stream's
collected chunks, messages, and model; a bare async generator would hide them.
"""
def __init__(self, anthropic_wrapper: AnthropicStreamWrapper) -> None:
self._anthropic_wrapper = anthropic_wrapper
self._byte_stream: Final[AsyncIterator[bytes]] = anthropic_wrapper.async_anthropic_sse_wrapper()
self._hidden_params: dict[
str, object
] = {} # mutable-ok: the proxy merges provider headers onto _hidden_params in place
@property
def chunks(self) -> "list[ModelResponseStream] | None":
return self._anthropic_wrapper.chunks
@property
def messages(self) -> "list[AllMessageValues] | None":
return self._anthropic_wrapper.messages
@property
def model(self) -> str:
return self._anthropic_wrapper.model
async def __anext__(self) -> bytes:
return await self._byte_stream.__anext__()
async def aclose(self) -> None:
await self._byte_stream.aclose()

View file

@ -201,7 +201,7 @@ from litellm.types.llms.openai import (
from litellm.types.utils import Choices, ModelResponse, StreamingChoices, Usage
from litellm.utils import supports_mid_conversation_system
from .streaming_iterator import AnthropicStreamWrapper
from .streaming_iterator import AnthropicSSEStream, AnthropicStreamWrapper
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObject
@ -341,7 +341,7 @@ class AnthropicAdapter:
)
# Return the SSE-wrapped version for proper event formatting.
if is_async:
return anthropic_wrapper.async_anthropic_sse_wrapper()
return AnthropicSSEStream(anthropic_wrapper)
return anthropic_wrapper.anthropic_sse_wrapper()

View file

@ -1,7 +1,7 @@
import re
from collections.abc import AsyncIterator, Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
from typing import TYPE_CHECKING, Final, cast
import litellm
from litellm._logging import verbose_logger
@ -17,6 +17,8 @@ from litellm.llms.anthropic.pass_through.messages.streaming_iterator import (
if TYPE_CHECKING:
from litellm.caching.caching_handler import LLMCachingHandler
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ModelResponseStream
CACHED_STREAM_EVENTS_KEY: Final = "litellm_cached_anthropic_sse_events"
@ -51,6 +53,24 @@ class AnthropicMessagesStreamCacheWriter:
def has_buffered_provider_output(self) -> bool:
return getattr(self.stream, "has_buffered_provider_output", False) is True
@property
def chunks(self) -> "list[ModelResponseStream] | None":
return cast( # cast-ok: chunks is a list of ModelResponseStream on the inner stream
"list[ModelResponseStream] | None", getattr(self.stream, "chunks", None)
)
@property
def messages(self) -> "list[AllMessageValues] | None":
return cast( # cast-ok: messages is a list of AllMessageValues on the inner stream
"list[AllMessageValues] | None", getattr(self.stream, "messages", None)
)
@property
def model(self) -> str | None:
return cast( # cast-ok: model is a str on the inner stream
"str | None", getattr(self.stream, "model", None)
)
def __aiter__(self) -> "AnthropicMessagesStreamCacheWriter":
return self

View file

@ -3,12 +3,8 @@ Helper util for handling azure openai-specific cost calculation
- e.g.: prompt caching, audio tokens
"""
from typing import Final
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token
from litellm.types.utils import Usage
from litellm.utils import get_model_info
def cost_per_token(
@ -27,26 +23,6 @@ def cost_per_token(
Returns:
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
"""
## GET MODEL INFO
model_info: Final = get_model_info(model=model, custom_llm_provider="azure")
## Speech / Audio cost calculation (cost per second for TTS models)
if (
"output_cost_per_second" in model_info
and model_info["output_cost_per_second"] is not None
and response_time_ms is not None
):
verbose_logger.debug(
"For model=%s - output_cost_per_second: %s; response time: %s",
model,
model_info.get("output_cost_per_second"),
response_time_ms,
)
## COST PER SECOND ##
prompt_cost: Final = 0.0
completion_cost: Final = model_info["output_cost_per_second"] * response_time_ms / 1000
return prompt_cost, completion_cost
## Use generic cost calculator for all other cases
## This properly handles: text tokens, audio tokens, cached tokens, reasoning tokens, etc.
return generic_cost_per_token(

View file

@ -778,6 +778,17 @@ class DatabricksChatResponseIterator(BaseModelResponseIterator):
)
choice["delta"]["thinking_blocks"] = thinking_blocks
translated_choices.append(choice)
service_tier: Final = chunk.get("service_tier")
if isinstance(service_tier, str) and service_tier:
return ModelResponseStream(
id=chunk["id"],
object="chat.completion.chunk",
created=chunk["created"],
model=chunk["model"],
choices=translated_choices,
usage=chunk.get("usage"),
service_tier=service_tier,
)
return ModelResponseStream(
id=chunk["id"],
object="chat.completion.chunk",

View file

@ -30,7 +30,7 @@ def _registry_key(model: str) -> str:
)
def cost_per_token(model: str, usage: Usage) -> tuple[float, float]:
def cost_per_token(model: str, usage: Usage, service_tier: str | None = None) -> tuple[float, float]:
"""
Calculates the cost per token for a given model, prompt tokens, and completion tokens.
@ -45,4 +45,5 @@ def cost_per_token(model: str, usage: Usage) -> tuple[float, float]:
model=_registry_key(model),
usage=usage,
custom_llm_provider="databricks",
service_tier=service_tier,
)

View file

@ -890,6 +890,9 @@ class OpenAIChatCompletionStreamingHandler(BaseModelResponseIterator):
}
if "usage" in chunk and chunk["usage"] is not None:
kwargs["usage"] = chunk["usage"]
service_tier: Final = chunk.get("service_tier")
if isinstance(service_tier, str) and service_tier:
kwargs["service_tier"] = service_tier
return ModelResponseStream(**kwargs)
except Exception as e:
raise e

View file

@ -5353,6 +5353,7 @@ def completion(
### CUSTOM MODEL COST ###
input_cost_per_token: Final = kwargs.get("input_cost_per_token", None)
output_cost_per_token: Final = kwargs.get("output_cost_per_token", None)
cost_per_second: Final = kwargs.get("cost_per_second", None)
input_cost_per_second: Final = kwargs.get("input_cost_per_second", None)
output_cost_per_second: Final = kwargs.get("output_cost_per_second", None)
### CUSTOM PROMPT TEMPLATE ###
@ -5514,8 +5515,11 @@ def completion(
### REGISTER CUSTOM MODEL PRICING -- IF GIVEN ###
if (
input_cost_per_token is not None and output_cost_per_token is not None
) or input_cost_per_second is not None:
(input_cost_per_token is not None and output_cost_per_token is not None)
or input_cost_per_second is not None
or output_cost_per_second is not None
or cost_per_second is not None
):
_register_custom_pricing_for_request(
model=model,
custom_llm_provider=custom_llm_provider,
@ -5657,6 +5661,7 @@ def completion(
proxy_server_request=proxy_server_request,
preset_cache_key=preset_cache_key,
no_log=no_log,
cost_per_second=cost_per_second,
input_cost_per_second=input_cost_per_second,
input_cost_per_token=input_cost_per_token,
output_cost_per_second=output_cost_per_second,
@ -6354,7 +6359,9 @@ def embedding(
### CUSTOM MODEL COST ###
input_cost_per_token: Final = kwargs.get("input_cost_per_token", None)
output_cost_per_token: Final = kwargs.get("output_cost_per_token", None)
cost_per_second: Final = kwargs.get("cost_per_second", None)
input_cost_per_second: Final = kwargs.get("input_cost_per_second", None)
output_cost_per_second: Final = kwargs.get("output_cost_per_second", None)
openai_params: Final = [
"user",
"dimensions",
@ -6395,7 +6402,12 @@ def embedding(
)
### REGISTER CUSTOM MODEL PRICING -- IF GIVEN ###
if (input_cost_per_token is not None and output_cost_per_token is not None) or input_cost_per_second is not None:
if (
(input_cost_per_token is not None and output_cost_per_token is not None)
or input_cost_per_second is not None
or output_cost_per_second is not None
or cost_per_second is not None
):
_register_custom_pricing_for_request(
model=model,
custom_llm_provider=custom_llm_provider,

File diff suppressed because it is too large Load diff

View file

@ -34,15 +34,12 @@ class LiteLLM_AutoRouterSession(LiteLLMPydanticObjectBase):
@property
def baseline_model(self) -> str | None:
"""The baseline most covered turns were priced against, or None when none were estimated.
A router reconfigured mid-session leaves turns priced against two baselines; the row keeps both
counts, and the label is the one that priced the most money-carrying turns rather than whatever the
router is configured with now.
"""
if not self.savings_estimated_baseline_models:
"""A recorded baseline label when excluded turns cannot change the selected model."""
if not self.baseline_models:
return None
if self.savings_estimated_turns < self.turns and len(self.baseline_models) > 1:
return None
return max(
self.savings_estimated_baseline_models,
key=lambda model: (self.savings_estimated_baseline_models[model], model),
self.baseline_models,
key=lambda model: (self.baseline_models[model], model),
)

View file

@ -3,7 +3,8 @@ import html as _html
import json
import secrets
import time
from collections.abc import Callable, Mapping
from collections.abc import AsyncIterator, Callable, Mapping
from contextlib import asynccontextmanager
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
@ -107,6 +108,14 @@ _OAUTH_METADATA_CACHE_MAX_SIZE: Final = 128
# Per-(server_id, resource_url) async locks so concurrent discovery requests
# coalesce onto a single upstream fetch instead of issuing N parallel calls.
_OAUTH_METADATA_FETCH_LOCKS: Final[dict[tuple[str, str], asyncio.Lock]] = {}
# Callers inside ``_oauth_metadata_fetch_slot`` per cache key, lock waiters included. ``Lock.locked()``
# reads False between one holder's release and the next waiter's wake-up, so it cannot tell an
# idle lock from one being handed off.
_OAUTH_METADATA_FETCHERS: Final[dict[tuple[str, str], int]] = {}
# Per-server_id generation, bumped on invalidation so a fetch that started before the server
# definition changed cannot repopulate the cache with the stale reply. Only servers with a fetch
# in flight carry an entry; the rest are pruned with the cache.
_OAUTH_METADATA_GENERATIONS: Final[dict[str, int]] = {}
router: Final = APIRouter(
tags=["mcp"],
@ -130,13 +139,52 @@ def _prune_oauth_metadata_cache(now: float | None = None) -> None:
for cache_key in cache_keys_by_expiry[:overflow]:
_OAUTH_METADATA_CACHE.pop(cache_key, None)
# Drop locks whose cache entry has been evicted and that aren't currently
# held; held locks stay so in-flight callers continue to coalesce.
# Drop locks whose cache entry has been evicted and that nobody holds or
# waits on; the rest stay so in-flight callers continue to coalesce.
for cache_key in list(_OAUTH_METADATA_FETCH_LOCKS):
if cache_key in _OAUTH_METADATA_CACHE:
if cache_key in _OAUTH_METADATA_CACHE or not _oauth_metadata_lock_idle(cache_key):
continue
lock = _OAUTH_METADATA_FETCH_LOCKS.get(cache_key)
if lock is None or lock.locked():
_OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None)
for server_id in [sid for sid in _OAUTH_METADATA_GENERATIONS if not _oauth_metadata_fetch_in_flight(sid)]:
_OAUTH_METADATA_GENERATIONS.pop(server_id, None)
def _oauth_metadata_fetch_in_flight(server_id: str) -> bool:
return any(cache_key[0] == server_id for cache_key in _OAUTH_METADATA_FETCHERS)
def _oauth_metadata_lock_idle(cache_key: tuple[str, str]) -> bool:
if cache_key in _OAUTH_METADATA_FETCHERS:
return False
lock: Final = _OAUTH_METADATA_FETCH_LOCKS.get(cache_key)
return lock is None or not lock.locked()
@asynccontextmanager
async def _oauth_metadata_fetch_slot(cache_key: tuple[str, str]) -> AsyncIterator[None]:
_OAUTH_METADATA_FETCHERS[cache_key] = _OAUTH_METADATA_FETCHERS.get(cache_key, 0) + 1
try:
async with _OAUTH_METADATA_FETCH_LOCKS.setdefault(cache_key, asyncio.Lock()):
yield
finally:
remaining: Final = _OAUTH_METADATA_FETCHERS.get(cache_key, 0) - 1
if remaining > 0:
_OAUTH_METADATA_FETCHERS[cache_key] = remaining
else:
_OAUTH_METADATA_FETCHERS.pop(cache_key, None)
def invalidate_oauth_metadata_cache(server_id: str) -> None:
"""Drop cached upstream IdP metadata for a server whose definition changed."""
if _oauth_metadata_fetch_in_flight(server_id):
_OAUTH_METADATA_GENERATIONS[server_id] = _OAUTH_METADATA_GENERATIONS.get(server_id, 0) + 1
else:
_OAUTH_METADATA_GENERATIONS.pop(server_id, None)
for cache_key in [key for key in _OAUTH_METADATA_CACHE if key[0] == server_id]:
del _OAUTH_METADATA_CACHE[cache_key]
for cache_key in [key for key in _OAUTH_METADATA_FETCH_LOCKS if key[0] == server_id]:
if not _oauth_metadata_lock_idle(cache_key):
continue
_OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None)
@ -2360,12 +2408,19 @@ async def fetch_upstream_oauth_protected_resource(
if cached is not None and cached[0] > now:
return cached[1]
lock: Final = _OAUTH_METADATA_FETCH_LOCKS.setdefault(cache_key, asyncio.Lock())
async with lock:
async with _oauth_metadata_fetch_slot(cache_key):
now = time.time()
cached = _OAUTH_METADATA_CACHE.get(cache_key)
if cached is not None and cached[0] > now:
return cached[1]
generation: Final = _OAUTH_METADATA_GENERATIONS.get(mcp_server.server_id, 0)
def store(payload: dict | None, ttl_seconds: int) -> None:
if _OAUTH_METADATA_GENERATIONS.get(mcp_server.server_id, 0) != generation:
return
stored_at: Final = time.time()
_OAUTH_METADATA_CACHE[cache_key] = (stored_at + ttl_seconds, payload)
_prune_oauth_metadata_cache(stored_at)
host_base: Final = f"{upstream.scheme}://{upstream.netloc}"
candidates: Final = [f"{host_base}/.well-known/oauth-protected-resource"]
@ -2407,12 +2462,7 @@ async def fetch_upstream_oauth_protected_resource(
)
continue
if isinstance(payload, dict):
now = time.time()
_OAUTH_METADATA_CACHE[cache_key] = (
now + _OAUTH_METADATA_CACHE_TTL_SECONDS,
payload,
)
_prune_oauth_metadata_cache(now)
store(payload, _OAUTH_METADATA_CACHE_TTL_SECONDS)
return payload
if len(network_errors) == len(candidates):
@ -2421,12 +2471,7 @@ async def fetch_upstream_oauth_protected_resource(
# Negative-result caching: when no candidate yielded a usable payload,
# remember that for a shorter TTL so we don't re-fetch on every
# subsequent discovery request (and so the per-key lock can be pruned).
now = time.time()
_OAUTH_METADATA_CACHE[cache_key] = (
now + _OAUTH_METADATA_NEGATIVE_CACHE_TTL_SECONDS,
None,
)
_prune_oauth_metadata_cache(now)
store(None, _OAUTH_METADATA_NEGATIVE_CACHE_TTL_SECONDS)
return None

View file

@ -2673,7 +2673,7 @@ class MCPServerManager:
self._assign_unique_short_prefix(new_server)
_warn_legacy_delegate_auth_if_applicable(new_server, source="config")
_warn_config_id_jag_server_outruns_sso(new_server)
self._invalidate_discovery_lists(server_id)
self._invalidate_server_definition_caches(server_id)
self.config_mcp_servers[server_id] = new_server
self._set_oauth_discovery_deferred(
server_id,
@ -2877,7 +2877,7 @@ class MCPServerManager:
global_mcp_tool_registry,
)
self._invalidate_discovery_lists(server.server_id)
self._invalidate_server_definition_caches(server.server_id)
prefix_root: Final = normalize_server_name(get_server_prefix(server))
if server.spec_path and prefix_root:
openapi_key_prefix: Final = prefix_root + MCP_TOOL_PREFIX_SEPARATOR
@ -3285,7 +3285,7 @@ class MCPServerManager:
# env_vars_are_encrypted=False.
new_server: Final = await self.build_mcp_server_from_table(mcp_server, env_vars_are_encrypted=False)
self._assign_unique_short_prefix(new_server)
self._invalidate_discovery_lists(mcp_server.server_id)
self._invalidate_server_definition_caches(mcp_server.server_id)
self.registry[mcp_server.server_id] = new_server
await self._maybe_register_openapi_tools(new_server)
self.prime_oauth_metadata_discovery(new_server)
@ -3322,7 +3322,7 @@ class MCPServerManager:
previous_server=self.registry[mcp_server.server_id],
)
self._assign_unique_short_prefix(new_server)
self._invalidate_discovery_lists(mcp_server.server_id)
self._invalidate_server_definition_caches(mcp_server.server_id)
self.registry[mcp_server.server_id] = new_server
await self._maybe_register_openapi_tools(new_server)
self.prime_oauth_metadata_discovery(new_server)
@ -4504,16 +4504,16 @@ class MCPServerManager:
if server.spec_path:
# OpenAPI tools were stored in the registry under the prefix
# active at registration time — fetch by that same prefix.
registered_prefix: Final = f"{get_server_prefix(server)}{MCP_TOOL_PREFIX_SEPARATOR}"
registry_prefix: Final = normalize_server_name(get_server_prefix(server)) + MCP_TOOL_PREFIX_SEPARATOR
registered: Final = global_mcp_tool_registry.convert_tools_to_mcp_sdk_tool_type(
global_mcp_tool_registry.list_tools(tool_prefix=get_server_prefix(server))
global_mcp_tool_registry.list_tools(tool_prefix=registry_prefix)
)
registered_names: Final = MappingProxyType(
{t.name.removeprefix(registered_prefix): t.name for t in registered}
{t.name.removeprefix(registry_prefix): t.name for t in registered}
)
guarded_openapi: Final = await self._guard_tool_catalog(
server=server,
tools=[t.model_copy(update={"name": t.name.removeprefix(registered_prefix)}) for t in registered],
tools=[t.model_copy(update={"name": t.name.removeprefix(registry_prefix)}) for t in registered],
proxy_logging_obj=proxy_logging_obj,
user_api_key_auth=user_api_key_auth,
raw_headers=raw_headers,
@ -4582,6 +4582,14 @@ class MCPServerManager:
self._resource_discovery_cache.invalidate(server_id)
self._template_discovery_cache.invalidate(server_id)
def _invalidate_server_definition_caches(self, server_id: str) -> None:
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( # noqa: PLC0415 # lazy: discoverable_endpoints lazily imports this module's manager singleton
invalidate_oauth_metadata_cache,
)
self._invalidate_discovery_lists(server_id)
invalidate_oauth_metadata_cache(server_id)
def _discovery_key(
self,
server: MCPServer,
@ -6792,7 +6800,7 @@ class MCPServerManager:
for server_id in previous_registry.keys() | registered_registry.keys():
if previous_registry.get(server_id) != registered_registry.get(server_id):
self._invalidate_discovery_lists(server_id)
self._invalidate_server_definition_caches(server_id)
self.registry = registered_registry
_warn_on_shared_identifier_prefixes(registered_registry.values())
# A discovery task may have published into ``previous_registry`` while

View file

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

View file

@ -27,7 +27,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 +46,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 +568,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 +3304,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 +3330,12 @@ 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_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 +3373,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 +4081,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 +4149,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 +5072,7 @@ class JWTAuthBuilderResult(TypedDict):
org_id: str | None
team_membership: LiteLLM_TeamMembership | None
jwt_claims: dict # Decoded JWT token claims (avoids re-decoding)
managed_agent_context: ReadOnly[NotRequired[ManagedAgentContext | None]]
agent_id: ReadOnly[str | None]

View file

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

View file

@ -0,0 +1,250 @@
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_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:
return await self.retired_agents.table.find_unique(where={"original_agent_id": agent_id}) is not None
except Exception:
return AgentIdentityFailure(code="policy_unavailable", message="Agent history is unavailable")
async def record_authentication(self, context: ManagedAgentContext) -> AgentIdentityFailure | None:
try:
if context.binding_revision is None:
return AgentIdentityFailure(message="Agent authentication requires a binding revision")
where: Final[LiteLLM_AgentIdentityWhereInput] = {
"agent_id": context.agent_id,
"revision": context.binding_revision,
"active": True,
"agent": {"is": {"enabled": True, "identity_managed": True}},
}
data: Final[LiteLLM_AgentIdentityUpdateManyMutationInput] = {
"last_authenticated_at": datetime.now(timezone.utc)
}
count: Final = await self.identities.table.update_many(where=where, data=data)
if count != 1:
return AgentIdentityFailure(message="Agent identity changed during authentication; retry")
return None
except Exception:
return AgentIdentityFailure(code="policy_unavailable", message="Agent authentication could not be recorded")
async def enroll_interactive_human(
self,
subject: MicrosoftInteractiveSubject,
user_id: str,
) -> AgentIdentityFailure | None:
try:
where: Final[LiteLLM_VerifiedSubjectWhereUniqueInput] = {
"issuer_tenant_id_oid": {"issuer": subject.issuer, "tenant_id": subject.tenant_id, "oid": subject.oid}
}
create_data: Final[LiteLLM_VerifiedSubjectCreateInput] = {
"issuer": subject.issuer,
"tenant_id": subject.tenant_id,
"oid": subject.oid,
"user_id": user_id,
"verified_via": "sso_interactive",
}
data: Final[LiteLLM_VerifiedSubjectUpsertInput] = {"create": create_data, "update": {}}
row: Final = await self.humans.table.upsert(where=where, data=data)
if row.kind != "human" or row.user_id != user_id or row.verified_via != "sso_interactive":
return AgentIdentityFailure(message="Microsoft subject is already bound to another local identity")
return None
except Exception:
return AgentIdentityFailure(
code="policy_unavailable", message="Microsoft subject enrollment is unavailable"
)
async def resolve_managed_agent(
claims: Mapping[str, object],
client: object,
*,
cache: UserApiKeyCache | None = None,
) -> ManagedAgentContext | None:
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
if client is None:
return None
result: Final = await AgentIdentityStore.from_client(client, cache=cache).resolve_verified_claims(claims)
if isinstance(result, AgentIdentityFailure):
raise_identity_failure(result)
return result

View file

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

View file

@ -23,7 +23,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,
@ -1784,11 +1784,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,
@ -2782,17 +2783,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 +2815,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 +2827,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(
@ -3565,13 +3587,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 +3622,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 +3653,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 +3691,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 +3701,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:

View file

@ -13,8 +13,9 @@ from typing import Final, Literal, Protocol, TypeAlias
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
@ -218,11 +219,23 @@ def _set_in_memory(memory: _InMemoryCache, cache_key: str, value: object, ttl: f
memory.set_cache(key=cache_key, value=value, ttl=ttl)
async def _read_redis_rows(keys: list[str], redis_cache: RedisCache) -> Mapping[str, object]:
"""On the request pipeline when one is open; a failed pipeline reads as a miss, like ``async_batch_get_cache``."""
batch: Final = active_request_redis_batch(redis_cache)
if batch is None:
return await redis_cache.async_batch_get_cache(key_list=keys) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # untyped cache API
try:
return await batch.mget(keys)
except Exception as e: # noqa: BLE001 # the DB fill below takes over, as it does after a failed MGET today
verbose_proxy_logger.debug("auth prefetch Redis read failed, filling from the database: %s", e)
return MappingProxyType({})
async def _fill_from_redis(entries: Sequence[_CacheEntry], redis_cache: RedisCache, memory: _InMemoryCache) -> None:
if not entries:
return
found: Final = _RowValues.validate_python(
await redis_cache.async_batch_get_cache(key_list=sorted(entry.cache_key for entry in entries)) # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # untyped cache API
await _read_redis_rows(sorted(entry.cache_key for entry in entries), redis_cache)
)
for entry, value in ((entry, found.get(entry.cache_key)) for entry in entries):
if value is not None:
@ -267,8 +280,14 @@ async def _write_back(entries: Sequence[tuple[_CacheEntry, BaseModel]], cache: U
memory: Final[_InMemoryCache] = cache.in_memory_cache
for cache_key, payload, ttl in payloads:
_set_in_memory(memory, cache_key, payload, cache.default_in_memory_ttl if ttl is None else ttl)
if cache.redis_cache is not None:
if cache.redis_cache is None:
return
batch: Final = active_request_redis_batch(cache.redis_cache)
if batch is None:
await cache.redis_cache.async_set_cache_pipeline_with_ttls(payloads)
return
for cache_key, payload, ttl in payloads: # rides the request's next round trip; the scope drains leftovers
batch.set(cache_key, payload, ttl)
async def _fill_from_db(
@ -305,3 +324,35 @@ async def prefetch_auth_objects(
await _fill_from_db(refs, _missing_in_memory(missing, memory), user_api_key_cache, prisma_client)
except Exception as e: # noqa: BLE001 # warm-up only; the getters enforce and fail closed on their own
verbose_proxy_logger.warning("auth prefetch skipped, falling back to per-object lookups: %s", e)
def _identity_memory_ttl(value: object, management_ttl: float) -> float:
"""A registry stored as a string is a sentinel, written with the shorter of the two registry TTLs."""
return min(REGISTRY_ERROR_NEGATIVE_CACHE_TTL, management_ttl) if isinstance(value, str) else management_ttl
async def prefetch_identity_keys(cache_keys: Sequence[str], user_api_key_cache: UserApiKeyCache) -> None:
"""Warm the entries auth reads before it knows the key's owners (the key object, the end user and the two
registries) in one MGET on the request pipeline. Keys the MGET finds absent stay noted on the pipeline, so the
per-key getters that follow go to the database without a GET of their own. Best effort, like the
owner prefetch: the getters read and enforce on their own."""
try:
redis_cache: Final = user_api_key_cache.redis_cache
if redis_cache is None:
return
missing: Final = tuple(
key
for key in dict.fromkeys(cache_keys)
if user_api_key_cache.in_memory_cache_for(key).get_cache(key=key) is None
)
if not missing:
return
found: Final = _RowValues.validate_python(await _read_redis_rows(sorted(missing), redis_cache))
management_ttl: Final = get_management_object_ttl(user_api_key_cache)
except Exception as e: # noqa: BLE001 # warm-up only; the getters read Redis and the database on their own
verbose_proxy_logger.warning("auth identity prefetch skipped, falling back to per-key lookups: %s", e)
return
for key, value in ((key, found.get(key)) for key in missing):
if value is not None:
memory: _InMemoryCache = user_api_key_cache.in_memory_cache_for(key)
_set_in_memory(memory, key, value, _identity_memory_ttl(value, management_ttl))

View file

@ -73,7 +73,7 @@ from litellm.proxy.auth.auth_checks import (
)
from litellm.proxy.auth.auth_exception_handler import UserAPIKeyAuthExceptionHandler
from litellm.proxy.auth.auth_method import AuthMethod
from litellm.proxy.auth.auth_object_prefetch import AuthObjectRefs, prefetch_auth_objects
from litellm.proxy.auth.auth_object_prefetch import AuthObjectRefs, prefetch_auth_objects, prefetch_identity_keys
from litellm.proxy.auth.auth_utils import (
abbreviate_api_key,
get_end_user_id_from_request_body,
@ -120,6 +120,9 @@ from litellm.proxy.common_utils.model_listing_utils import claude_code_requested
from litellm.proxy.common_utils.realtime_utils import _realtime_request_body
from litellm.proxy.common_utils.user_api_key_cache import (
UserApiKeyCache,
end_user_cache_key,
end_user_restricted_registry_cache_key,
model_access_group_registry_cache_key,
team_membership_auth_cache_key,
)
from litellm.proxy.db.db_lookup_gate import bounded_db_lookup
@ -1892,6 +1895,11 @@ async def _user_api_key_auth_builder(
proxy_logging_obj=proxy_logging_obj,
route=route,
)
if prisma_client is not None:
await prefetch_identity_keys(
_identity_cache_keys(api_key, end_user_id=end_user_id, key_is_resolved=valid_token is not None),
user_api_key_cache=user_api_key_cache,
)
if end_user_id:
try:
end_user_params["end_user_id"] = end_user_id
@ -3006,41 +3014,40 @@ async def _run_centralized_common_checks(
skip_budget_checks=skip_budget_checks,
project_object=project_object,
)
if not skip_budget_checks:
await _check_team_model_budget(
valid_token=user_api_key_auth_obj,
model_max_budget_limiter=model_max_budget_limiter,
models=_get_model_names_for_budget_checks(
model=_get_model_from_request_context(
request_data=request_data,
route=route,
request=request,
llm_router=llm_router,
team_id=user_api_key_auth_obj.team_id,
)
),
)
await _reserve_budget_after_common_checks(
user_api_key_auth_obj=user_api_key_auth_obj,
request=request,
request_data=request_data,
route=route,
llm_router=llm_router,
team_object=team_object,
user_object=user_object,
end_user_id=end_user_id,
end_user_object=end_user_object,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
skip_budget_checks=skip_budget_checks,
general_settings=general_settings,
)
finally:
release_spend_counter_batch()
if not skip_budget_checks:
await _check_team_model_budget(
valid_token=user_api_key_auth_obj,
model_max_budget_limiter=model_max_budget_limiter,
models=_get_model_names_for_budget_checks(
model=_get_model_from_request_context(
request_data=request_data,
route=route,
request=request,
llm_router=llm_router,
team_id=user_api_key_auth_obj.team_id,
)
),
)
await _reserve_budget_after_common_checks(
user_api_key_auth_obj=user_api_key_auth_obj,
request=request,
request_data=request_data,
route=route,
llm_router=llm_router,
team_object=team_object,
user_object=user_object,
end_user_id=end_user_id,
end_user_object=end_user_object,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
skip_budget_checks=skip_budget_checks,
general_settings=general_settings,
)
async def _noop_none() -> None:
"""Sentinel coroutine for asyncio.gather when a fetch is unnecessary
@ -3249,6 +3256,21 @@ def _spend_counter_redis_cache() -> RedisCache | None:
return spend_counter_cache.redis_cache
def _identity_cache_keys(api_key: str, *, end_user_id: str | None, key_is_resolved: bool) -> tuple[str, ...]:
"""Cache keys auth reads before it knows the key's owners, all known from the request alone. A key object is
cached under the hash of the bearer, so the bearer itself never reaches Redis."""
return tuple(
key
for key in (
None if key_is_resolved else hash_token(api_key),
None if not end_user_id else end_user_cache_key(end_user_id),
None if not end_user_id else end_user_restricted_registry_cache_key(),
model_access_group_registry_cache_key(),
)
if key is not None
)
async def _prefetch_referenced_auth_objects(
valid_token: UserAPIKeyAuth,
end_user_id: str | None,

View file

@ -2199,6 +2199,9 @@ class ProxyBaseLLMRequestProcessing:
if self._tags_before_guardrails is None:
self._tags_before_guardrails = frozenset(get_tags_from_request_body(request_body=self.data))
prefetch_model = self.data.get("model")
if llm_router is not None and isinstance(prefetch_model, str):
llm_router.arm_routing_read_prefetch(prefetch_model, self.data)
self.data = await proxy_logging_obj.pre_call_hook(
user_api_key_dict=user_api_key_dict,
data=self.data,

View file

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

View file

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

View file

@ -0,0 +1,147 @@
from collections.abc import Mapping
from contextlib import AbstractAsyncContextManager
from datetime import timedelta
from math import isclose
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Protocol, cast
from pydantic import BaseModel, ConfigDict, TypeAdapter
from litellm._logging import verbose_proxy_logger
from litellm.constants import MAX_SPENDLOG_ROWS_TO_QUERY
from litellm.proxy.db.autorouter_session_rollup import AUTOROUTER_SESSION_WINDOW_SQL
from litellm.proxy.db.create_views import SupportsRawQueries
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient
class SessionSavingsComparison(BaseModel):
model_config = ConfigDict(frozen=True, allow_inf_nan=False)
router_name: str
router_type: str
turns: int
estimated_turns: int
actual_spend: float
classifier_cost: float | None
saved_spend: float
complete: bool
def coverage_fields(self, recorded_savings: float, recorded_turns: int) -> Mapping[str, float | int]:
if self.turns != recorded_turns or not self.complete:
return MappingProxyType({})
if not isclose(self.saved_spend, recorded_savings, rel_tol=1e-9, abs_tol=1e-9):
return MappingProxyType({})
return MappingProxyType(
{
"savings_estimated_turns": self.estimated_turns,
"savings_estimated_actual_spend": self.actual_spend,
"savings_estimated_saved_spend": self.saved_spend,
}
)
class _ReadTransactions(Protocol):
def tx(self, *, timeout: timedelta, max_wait: timedelta) -> AbstractAsyncContextManager[SupportsRawQueries]: ...
_COMPARISONS: Final = TypeAdapter(tuple[SessionSavingsComparison, ...])
async def historical_session_comparisons(
prisma_client: "PrismaClient",
start_date: str,
end_date: str,
api_key: str | None,
user_id: str | None,
session_id: str | None = None,
) -> Mapping[tuple[str, str], SessionSavingsComparison]:
try:
reader: Final = cast(_ReadTransactions, prisma_client.read_db) # cast-ok: untyped Prisma transaction delegate
async with reader.tx(timeout=timedelta(seconds=3), max_wait=timedelta(seconds=1)) as transaction:
await transaction.execute_raw("SET TRANSACTION READ ONLY")
await transaction.execute_raw("SET LOCAL statement_timeout = 2000")
rows: Final = await transaction.query_raw(
HISTORICAL_SESSION_COMPARISONS_SQL,
start_date,
end_date,
api_key,
user_id,
session_id,
)
comparisons: Final = _COMPARISONS.validate_python(rows or ())
return MappingProxyType({(row.router_name, row.router_type): row for row in comparisons})
except Exception: # noqa: BLE001 # missing retained logs must not discard recorded dollar savings
verbose_proxy_logger.warning("Historical auto-router cost comparison unavailable; preserving recorded savings")
return MappingProxyType({})
HISTORICAL_SESSION_COMPARISONS_SQL: Final = f"""
WITH {AUTOROUTER_SESSION_WINDOW_SQL}, scoped AS MATERIALIZED (
SELECT * FROM windowed WHERE $5::text IS NULL OR session_id = $5::text
), limited_logs AS MATERIALIZED (
SELECT session.api_key, session.session_id, session.router_name, session.router_type, session.comparison_user_id,
session.classifier_cost_recorded_turns = session.turns AS classifier_cost_tracked,
logs.spend, logs.prompt_tokens + logs.completion_tokens AS tokens,
logs.metadata::jsonb -> 'routing_decision' AS decision,
logs.metadata::jsonb -> 'autorouter_savings' AS savings,
logs.metadata::jsonb -> 'autorouter_savings_estimate' AS estimate
FROM scoped AS session JOIN "LiteLLM_SpendLogs" AS logs
ON logs.api_key = session.api_key
AND CASE WHEN char_length(logs.session_id) > 256
THEN 'sha256:' || encode(sha256(convert_to(logs.session_id, 'UTF8')), 'hex')
ELSE logs.session_id END = session.session_id
AND (session.comparison_user_id IS NULL OR logs."user" = session.comparison_user_id)
AND logs."startTime" BETWEEN session.first_turn_at AND session.last_turn_at
AND COALESCE(logs.metadata::jsonb #>> '{{routing_decision,router_model_name}}', logs.model_group)
= session.router_name
WHERE session.savings_estimated_turns < session.turns
AND logs.status = 'success' AND COALESCE(logs.metadata::jsonb ->> 'internal_call_origin', '') = ''
LIMIT {MAX_SPENDLOG_ROWS_TO_QUERY + 1}
), facts AS (
SELECT *,
CASE WHEN jsonb_typeof(decision -> 'classifier_cost') = 'number'
THEN (decision ->> 'classifier_cost')::float8
WHEN classifier_cost_tracked THEN 0 END AS classifier,
CASE WHEN jsonb_typeof(savings) = 'number' AND (
estimate IS NULL OR estimate = 'null'::jsonb OR (
jsonb_typeof(estimate -> 'version') = 'number' AND estimate ->> 'version' IN ('1', '2', '3')
AND estimate ->> 'status' = 'estimated'
)
) THEN savings::text::float8 END AS saved
FROM limited_logs
), compared AS (
SELECT api_key, session_id, router_name, router_type, comparison_user_id,
COUNT(*) AS turns, SUM(spend + COALESCE(classifier, 0)) AS spend, SUM(tokens) AS total_tokens,
COUNT(saved) AS estimated_turns,
COALESCE(SUM(spend + COALESCE(classifier, 0)) FILTER (WHERE saved IS NOT NULL), 0)::float8 AS actual_spend,
CASE WHEN COUNT(saved) = COUNT(classifier) FILTER (WHERE saved IS NOT NULL)
THEN COALESCE(SUM(classifier) FILTER (WHERE saved IS NOT NULL), 0)::float8
END AS estimated_classifier_cost,
COALESCE(SUM(saved), 0)::float8 AS saved_spend
FROM facts GROUP BY 1, 2, 3, 4, 5
), reconciled AS (
SELECT session.*, logs.estimated_turns, logs.actual_spend, logs.estimated_classifier_cost,
COALESCE((SELECT COUNT(*) FROM limited_logs) <= {MAX_SPENDLOG_ROWS_TO_QUERY}
AND logs.turns = session.turns AND logs.total_tokens = session.total_tokens
AND ABS(logs.spend - session.spend) <= GREATEST(1e-9, ABS(session.spend) * 1e-9)
AND ABS(logs.saved_spend - session.saved_spend) <= GREATEST(1e-9, ABS(session.saved_spend) * 1e-9), FALSE
) AS recovered
FROM scoped AS session LEFT JOIN compared AS logs
ON logs.api_key = session.api_key AND logs.session_id = session.session_id
AND logs.router_name = session.router_name AND logs.router_type = session.router_type
AND logs.comparison_user_id IS NOT DISTINCT FROM session.comparison_user_id
)
SELECT router_name, router_type,
SUM(turns)::bigint AS turns,
SUM(CASE WHEN recovered THEN estimated_turns ELSE savings_estimated_turns END)::bigint AS estimated_turns,
SUM(CASE WHEN recovered THEN actual_spend ELSE savings_estimated_actual_spend END)::float8 AS actual_spend,
CASE WHEN BOOL_AND(CASE WHEN recovered THEN estimated_classifier_cost IS NOT NULL
ELSE savings_estimated_turns = turns AND classifier_cost_recorded_turns = turns END)
THEN SUM(CASE WHEN recovered THEN estimated_classifier_cost ELSE classifier_cost END)::float8
END AS classifier_cost,
SUM(saved_spend)::float8 AS saved_spend,
BOOL_AND(recovered OR savings_estimated_turns = turns) AS complete
FROM reconciled GROUP BY router_name, router_type
"""

View file

@ -7,8 +7,8 @@ on the prisma client. The spend-log flush job drains the queue into
key and user session rollups with one atomic statement per turn: each upsert classifies
the turn (same model, first visit, return to a model the session already used, out of
order) against the row's own columns, so nothing is read before the write and concurrent
pods compose. The benchmarks endpoint aggregates these rows and never touches
LiteLLM_SpendLogs.
pods compose. The benchmarks endpoint aggregates these rows and can recover matching historical
costs from retained spend logs when estimate coverage predates these columns.
"""
from __future__ import annotations
@ -45,20 +45,24 @@ _SESSION_COLUMNS: Final = """
savings_estimated_baseline_models
"""
AUTOROUTER_BENCHMARKS_SQL: Final = f"""
WITH windowed AS (
SELECT {_SESSION_COLUMNS} FROM "LiteLLM_AutoRouterSession"
AUTOROUTER_SESSION_WINDOW_SQL: Final = f"""
windowed AS (
SELECT {_SESSION_COLUMNS}, NULL::text AS comparison_user_id FROM "LiteLLM_AutoRouterSession"
WHERE $4::text IS NULL
AND last_turn_at >= $1::timestamp
AND first_turn_at < $2::timestamp
AND ($3::text IS NULL OR api_key = $3::text)
UNION ALL
SELECT {_SESSION_COLUMNS} FROM "LiteLLM_AutoRouterUserSession"
SELECT {_SESSION_COLUMNS}, user_id AS comparison_user_id FROM "LiteLLM_AutoRouterUserSession"
WHERE (($4::text IS NOT NULL AND user_id = $4::text) OR ($4::text IS NULL AND api_key = ''))
AND last_turn_at >= $1::timestamp
AND first_turn_at < $2::timestamp
AND ($3::text IS NULL OR api_key = $3::text)
),
)
"""
AUTOROUTER_BENCHMARKS_SQL: Final = f"""
WITH {AUTOROUTER_SESSION_WINDOW_SQL},
tier_maps AS (
SELECT router_name, router_type, jsonb_object_agg(tier, tier_turns) AS tier_turns
FROM (
@ -95,6 +99,8 @@ SELECT
COALESCE(SUM(saved_spend), 0)::float8 AS saved_spend,
COALESCE(SUM(savings_estimated_turns), 0)::int AS savings_estimated_turns,
COALESCE(SUM(savings_estimated_actual_spend), 0)::float8 AS savings_estimated_actual_spend,
CASE WHEN BOOL_AND(savings_estimated_turns = turns AND classifier_cost_recorded_turns = turns)
THEN SUM(classifier_cost)::float8 END AS savings_estimated_classifier_cost,
COALESCE(SUM(savings_estimated_saved_spend), 0)::float8 AS savings_estimated_saved_spend,
COALESCE(SUM(classifier_cost), 0)::float8 AS classifier_cost,
COALESCE(SUM(classifier_cost_recorded_turns), 0)::int AS classifier_cost_recorded_turns,

View file

@ -2012,4 +2012,5 @@ class PanwPrismaAirsHandler(CustomGuardrail):
GuardrailEventHooks.logging_only,
GuardrailEventHooks.pre_mcp_call,
GuardrailEventHooks.during_mcp_call,
GuardrailEventHooks.post_mcp_call,
]

View file

@ -33,6 +33,12 @@ from typing_extensions import NotRequired, ReadOnly
from litellm import DualCache
from litellm._logging import verbose_proxy_logger
from litellm.caching.redis_batch import (
BatchResult,
RegisteredScript,
active_post_call_redis_batch,
active_request_redis_batch,
)
from litellm.caching.redis_cache import log_redis_failure
from litellm.constants import DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE, INTERNAL_CALL_ORIGIN_METADATA_KEY
from litellm.integrations.custom_logger import CustomLogger
@ -474,6 +480,19 @@ CacheCounterValue: TypeAlias = int | float | str | bytes
CacheCounterValues: TypeAlias = Sequence[CacheCounterValue | None]
def _as_counter_values(reply: object) -> list[CacheCounterValue]:
"""A Lua reply read back off the pipeline is the same array the script returns when called directly."""
if not isinstance(reply, (list, tuple)):
raise TypeError(f"rate limiter script reply is not a list: {type(reply).__name__}")
values: Final[list[CacheCounterValue]] = [] # mutable-ok: each element is narrowed before it is kept
for value in reply: # pyright: ignore[reportUnknownVariableType] # raw Redis reply
if not isinstance(value, (int, float, str, bytes)):
raise TypeError(f"rate limiter script reply holds {type(value).__name__}") # pyright: ignore[reportUnknownArgumentType] # raw Redis reply
values.append(value)
return values
ReservationWindowIdentity: TypeAlias = tuple[str, str, Literal["redis", "local"]]
ParallelGaugeCacheValue: TypeAlias = dict[str, object] | int | float | str | bytes
@ -1323,6 +1342,21 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
crc: Final = binascii.crc_hqx(key.encode("utf-8"), 0)
return crc % REDIS_CLUSTER_SLOTS
def _pipeline_scripts(
self,
source: str,
run: RegisteredScript,
calls: Sequence[tuple[Sequence[str], Sequence[int]]],
) -> tuple[BatchResult[object] | None, ...]:
"""Declare one Lua call per group on the request's Redis batch, so all groups share one round trip
with whatever else the request declared (the routing read). Returns ``None`` per call when no batch
is open, and the caller runs the script directly as before."""
redis_cache: Final = self.internal_usage_cache.dual_cache.redis_cache
batch: Final = None if redis_cache is None else active_request_redis_batch(redis_cache)
if batch is None:
return (None,) * len(calls)
return tuple(batch.script(source, run, keys, args) for keys, args in calls)
def _group_keys_by_hash_tag(self, keys: list[str]) -> dict[str, list[str]]:
"""
Group keys by their Redis hash tag to ensure cluster compatibility.
@ -1404,7 +1438,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
)
return await self._batch_get_counter_values(keys=keys, parent_otel_span=parent_otel_span, local_only=True)
def _reject_if_rate_limit_unverifiable(self, failed_operation: str, error: Exception) -> None:
def _reject_if_rate_limit_unverifiable(self, failed_operation: str, error: BaseException) -> None:
if not self._fail_closed_resolver():
return
log_redis_failure(
@ -1436,12 +1470,19 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
key_groups: Final = list(self._group_keys_by_hash_tag(keys_to_fetch).items())
all_cache_values: Final[list[CacheCounterValue | None]] = []
args: Final = (now_int, self.window_size)
pipelined: Final = self._pipeline_scripts(
BATCH_RATE_LIMITER_SCRIPT,
self.batch_rate_limiter_script,
tuple((group_keys, args) for _tag, group_keys in key_groups),
)
for index, (hash_tag, group_keys) in enumerate(key_groups):
for index, ((hash_tag, group_keys), group_result) in enumerate(zip(key_groups, pipelined)):
try:
group_cache_values: CacheCounterValues = await self.batch_rate_limiter_script(
keys=group_keys,
args=[now_int, self.window_size], # Use integer timestamp
group_cache_values: CacheCounterValues = (
await self.batch_rate_limiter_script(keys=group_keys, args=args)
if group_result is None
else _as_counter_values(await group_result)
)
all_cache_values.extend(group_cache_values)
except Exception as e:
@ -1450,6 +1491,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
await self._refund_counter_increments(
self._counter_refunds_from_batch_values(applied_keys, all_cache_values)
)
await self._refund_later_pipelined_groups(key_groups[index + 1 :], pipelined[index + 1 :])
self._reject_if_rate_limit_unverifiable("batch_rate_limiter_script", e)
log_redis_failure(
verbose_proxy_logger, logging.WARNING, f"Redis Lua script failed for hash tag {hash_tag}", e
@ -1464,6 +1506,22 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
return all_cache_values
async def _refund_later_pipelined_groups(
self,
key_groups: Sequence[tuple[str, list[str]]],
pipelined: Sequence[BatchResult[object] | None],
) -> None:
"""Groups declared on the request batch ran in the same round trip as the one that failed, so their
increments landed even though the loop never read them."""
for (_tag, group_keys), group_result in zip(key_groups, pipelined):
if group_result is None:
continue
try:
group_values = _as_counter_values(await group_result)
except Exception: # noqa: BLE001 # a group that failed in Redis incremented nothing to refund
continue
await self._refund_counter_increments(self._counter_refunds_from_batch_values(group_keys, group_values))
async def should_rate_limit(
self,
descriptors: Sequence[RateLimitDescriptor],
@ -1840,6 +1898,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
self,
stash: RequestRateLimiterStash | None,
parent_otel_span: Span | None,
*,
in_logging_callback: bool = False,
) -> None:
if stash is None:
return
@ -1847,7 +1907,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
acquisition: Final = stash.parallel_slot
if acquisition is None:
return
await self._release_parallel_request_slots(acquisition, parent_otel_span)
deferred: Final = in_logging_callback and await self._defer_parallel_slot_release(
acquisition, parent_otel_span
)
if not deferred:
await self._release_parallel_request_slots(acquisition, parent_otel_span)
stash.parallel_slot = None # rebind-ok: marks this request's slot as released
async def _release_parallel_request_slots(
@ -1873,14 +1937,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
keys=counter_keys,
args=[slot_id for _ in counter_keys],
)
for counter_key, remaining in zip(counter_keys, raw):
await self.internal_usage_cache.async_set_cache(
key=counter_key,
value=max(0, int(remaining)),
ttl=PARALLEL_REQUEST_SLOT_TTL_SECONDS,
litellm_parent_otel_span=parent_otel_span,
local_only=True,
)
await self._mirror_released_parallel_slots(counter_keys, raw, parent_otel_span)
return
except Exception as e: # noqa: BLE001 - any Redis/Lua failure degrades to the in-memory release, never a 500
log_redis_failure(
@ -1889,7 +1946,55 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
"parallel_release_script failed, falling back to in-memory release",
e,
)
await self._release_parallel_request_slots_in_memory(counter_keys, slot_id, parent_otel_span)
async def _defer_parallel_slot_release(
self, acquisition: ParallelSlotAcquisition, parent_otel_span: Span | None
) -> bool:
"""Only for a release from the logging callbacks: the response has left and the callbacks' end flushes
the pipeline. A release before the response goes to Redis at once, so another worker's next acquire
never counts a finished request. The local gauge frees the slot at once, so admission on this worker
sees the capacity before the pipeline goes out. The count Redis returns from the pipeline is not
mirrored: by then a newer acquire on this worker may have written a fresher count, and the next
acquire refreshes the gauge anyway."""
counter_keys: Final = acquisition["counter_keys"]
slot_id: Final = acquisition["slot_id"]
redis_cache: Final = self.internal_usage_cache.dual_cache.redis_cache
script: Final = self.parallel_release_script
batch: Final = None if redis_cache is None else active_post_call_redis_batch(redis_cache)
if batch is None or script is None or not counter_keys or not slot_id:
return False
await self._release_parallel_request_slots_in_memory(counter_keys, slot_id, parent_otel_span)
async def settle(future: asyncio.Future[object]) -> None:
if future.cancelled() or future.exception() is not None:
log_redis_failure(
verbose_proxy_logger,
logging.WARNING,
"parallel_release_script failed, the slot stays released in memory only",
future.exception() if not future.cancelled() else asyncio.CancelledError(),
)
batch.script(PARALLEL_RELEASE_SCRIPT, script, counter_keys, (slot_id,) * len(counter_keys)).on_settled(settle)
return True
async def _mirror_released_parallel_slots(
self, counter_keys: list[str], remaining_by_key: Sequence[object], parent_otel_span: Span | None
) -> None:
for counter_key, remaining in zip(counter_keys, remaining_by_key):
if not isinstance(remaining, (int, float, str, bytes)):
continue
await self.internal_usage_cache.async_set_cache(
key=counter_key,
value=max(0, int(remaining)),
ttl=PARALLEL_REQUEST_SLOT_TTL_SECONDS,
litellm_parent_otel_span=parent_otel_span,
local_only=True,
)
async def _release_parallel_request_slots_in_memory(
self, counter_keys: list[str], slot_id: str, parent_otel_span: Span | None
) -> None:
async with self._check_and_increment_lock:
for counter_key in counter_keys:
raw_value: ParallelGaugeCacheValue | None = await self.internal_usage_cache.async_get_cache(
@ -2061,7 +2166,16 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
reservation_windows: Final[set[ReservationWindowIdentity]] = set() # mutable-ok: filled by the group loop
raw: list[CacheCounterValue]
for _idx, (keys, args, meta) in enumerate(descriptor_groups):
pipelined: Final = self._pipeline_scripts(
CHECK_AND_INCREMENT_BY_N_SCRIPT,
self.check_and_increment_by_n_script, # pyright: ignore[reportArgumentType] # sole caller guards it is not None
tuple((keys, args) for keys, args, _meta in descriptor_groups),
)
batched: Final = tuple(result for result in pipelined if result is not None)
if len(batched) == len(descriptor_groups):
return await self._settle_pipelined_descriptor_groups(descriptor_groups, batched, parent_otel_span)
for keys, args, meta in descriptor_groups:
try:
raw = await self.check_and_increment_by_n_script( # pyright: ignore[reportOptionalCall] # sole caller guards it is not None
keys=keys,
@ -2105,6 +2219,76 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
reservation_windows=frozenset(reservation_windows),
)
async def _settle_pipelined_descriptor_groups(
self,
descriptor_groups: list[DescriptorAtomicGroup],
results: Sequence[BatchResult[object]],
parent_otel_span: Span | None,
) -> RateLimitResponse:
"""Every group's Lua call left in one pipeline, so each group has already checked and incremented on
its own before any result is read. A failed or over-limit group therefore refunds every group that
incremented, after it as well as before it, where the one-at-a-time loop only unwinds the groups it ran.
A Redis denial stands even when another group failed: the in-memory fallback only replaces a verdict
Redis never gave."""
replies: Final = await asyncio.gather(*results, return_exceptions=True)
responses: Final = tuple(
self._pipelined_group_response(reply, meta)
for reply, (_keys, _args, meta) in zip(replies, descriptor_groups)
)
applied: Final[list[tuple[CounterRefund, ...]]] = [] # mutable-ok: filled by the group loop
statuses: Final[list[RateLimitStatus]] = [] # mutable-ok: filled by the group loop
reservation_windows: Final[set[ReservationWindowIdentity]] = set() # mutable-ok: filled by the group loop
for reply, response, (_keys, _args, meta) in zip(replies, responses, descriptor_groups):
if isinstance(response, BaseException) or response["overall_code"] != "OK":
continue
applied.append(self._counter_refunds_from_atomic_response(_as_counter_values(reply), meta))
statuses.extend(response["statuses"])
reservation_windows.update(response.get("reservation_windows", frozenset()))
over_limit: Final = next(
(r for r in responses if not isinstance(r, BaseException) and r["overall_code"] == "OVER_LIMIT"), None
)
if over_limit is not None:
await self._refund_applied_descriptor_groups(applied)
return over_limit
failure: Final = next((r for r in responses if isinstance(r, BaseException)), None)
if failure is not None:
await self._refund_applied_descriptor_groups(applied)
self._reject_if_rate_limit_unverifiable("check_and_increment_by_n_script", failure)
log_redis_failure(
verbose_proxy_logger,
logging.ERROR,
f"atomic_check_and_increment_by_n: Redis Lua execution failed ({type(failure).__name__}). Refunding "
f"{len(applied)} pipelined descriptors and falling back to in-memory enforcement, counters will "
f"diverge from Redis until window expires (window_size={self.window_size}s)",
failure,
)
flat_meta: Final = tuple(
itertools.chain.from_iterable(group_meta for _k, _a, group_meta in descriptor_groups)
)
async with self._check_and_increment_lock:
return await self._atomic_check_and_increment_in_memory(
per_counter_meta=flat_meta,
parent_otel_span=parent_otel_span,
)
if len(responses) == 1 and not isinstance(responses[0], BaseException):
return responses[0]
return RateLimitResponse(
overall_code="OK",
statuses=statuses,
reservation_windows=frozenset(reservation_windows),
)
def _pipelined_group_response(
self, reply: object, per_counter_meta: list[AtomicCounterMeta]
) -> RateLimitResponse | BaseException:
if isinstance(reply, BaseException):
return reply
try:
return self._build_atomic_response(_as_counter_values(reply), per_counter_meta)
except Exception as e: # noqa: BLE001 # a reply this group cannot read is that group's Lua failure
return e
async def _refund_applied_descriptor_groups(
self,
applied: Sequence[Sequence[CounterRefund]],
@ -2233,7 +2417,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
async def _atomic_check_and_increment_in_memory(
self,
per_counter_meta: list[AtomicCounterMeta],
per_counter_meta: Sequence[AtomicCounterMeta],
parent_otel_span: Span | None = None,
) -> RateLimitResponse:
"""In-memory all-or-nothing check-and-increment. Caller holds lock.
@ -4169,11 +4353,43 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
keys.append(op["key"])
args.extend([op["increment_value"], ttl_value])
if self._defer_token_increment_script(keys, args, group_operations):
continue
await self.token_increment_script(
keys=keys,
args=args,
)
def _defer_token_increment_script(
self,
keys: list[str],
args: list[int],
group_operations: list["RedisPipelineIncrementOperation"],
) -> bool:
"""Declared into the request's post-call pipeline instead of its own EVALSHA round trip; a failed
script falls back to the plain increment pipeline for its own group, as the direct path does."""
redis_cache: Final = self.internal_usage_cache.dual_cache.redis_cache
script: Final = self.token_increment_script
batch: Final = None if redis_cache is None else active_post_call_redis_batch(redis_cache)
if batch is None or script is None:
return False
async def fall_back(future: asyncio.Future[object]) -> None:
if future.cancelled() or future.exception() is None:
return
log_redis_failure(
verbose_proxy_logger,
logging.WARNING,
"TTL preservation failed, falling back to regular pipeline",
future.exception(),
)
await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline(
increment_list=group_operations,
)
batch.script(TOKEN_INCREMENT_SCRIPT, script, keys, args).on_settled(fall_back)
return True
async def async_increment_tokens_with_ttl_preservation(
self,
pipeline_operations: list["RedisPipelineIncrementOperation"],
@ -4787,7 +5003,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
verbose_proxy_logger.debug("INSIDE parallel request limiter ASYNC SUCCESS LOGGING")
stash: Final = get_request_stash_for_call(_call_id_from_callback_kwargs(kwargs))
await self._release_stashed_parallel_slot(stash, litellm_parent_otel_span)
await self._release_stashed_parallel_slot(stash, litellm_parent_otel_span, in_logging_callback=True)
pipeline_operations: Final = self._build_success_event_pipeline_operations(
kwargs=kwargs,
@ -4907,7 +5123,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
pipeline_operations: Final[list[RedisPipelineIncrementOperation]] = []
stash: Final = get_request_stash_for_call(_call_id_from_callback_kwargs(kwargs))
await self._release_stashed_parallel_slot(stash, litellm_parent_otel_span)
await self._release_stashed_parallel_slot(stash, litellm_parent_otel_span, in_logging_callback=True)
# Skip the reservation refund if async_post_call_failure_hook
# already released it (proxy-level rejection that also bubbles up
@ -4977,15 +5193,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
)
if pipeline_operations:
await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline(
increment_list=pipeline_operations,
litellm_parent_otel_span=litellm_parent_otel_span,
await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline_post_call(
pipeline_operations, parent_otel_span=litellm_parent_otel_span
)
for project_operations in (itpm_operations, otpm_operations):
if isinstance(project_operations, list):
await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline(
increment_list=project_operations,
litellm_parent_otel_span=litellm_parent_otel_span,
await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline_post_call(
project_operations, parent_otel_span=litellm_parent_otel_span
)
elif project_operations:
await self.async_increment_reservation_aware_tokens(

View file

@ -30,6 +30,7 @@ from litellm.proxy.db.db_spend_update_writer import (
get_llm_router,
)
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.spend_tracking.spend_counter_batch import post_call_counter_keys, spend_counter_batch_scope
from litellm.proxy.spend_tracking.spend_event import (
ObjectMapping,
SpendEventBuildError,
@ -284,6 +285,7 @@ class _ProxyDBLogger(CustomLogger):
increment_spend_counters,
proxy_logging_obj,
update_cache,
update_cache_read_keys,
)
verbose_proxy_logger.debug("INSIDE _PROXY_track_cost_callback")
@ -377,6 +379,13 @@ class _ProxyDBLogger(CustomLogger):
request_tags=tags,
model_access_groups=model_access_groups,
project_id=project_id,
update_cache_read_keys=update_cache_read_keys(
user_id=user_id,
end_user_id=end_user_id,
team_id=team_id,
tags=tags,
response_cost=response_cost,
),
)
if not charged:
return
@ -694,11 +703,73 @@ async def _update_database_and_spend_counters(
request_tags: list[str] | None = None,
model_access_groups: Sequence[str] | None = None,
project_id: str | None = None,
update_cache_read_keys: Sequence[str] = (),
) -> bool:
"""The reservation is reconciled before the spend is persisted, from its own read. One spend counter batch then
spans the database write and the counter update, so the post-call counters are read with a single MGET after the
write and their increments leave in a single pipeline."""
from litellm.proxy.proxy_server import spend_counter_cache
from litellm.proxy.spend_tracking.budget_reservation import get_reserved_counter_keys
if budget_reservation is not None:
await _reconcile_budget_reservation_before_db_update(
budget_reservation=budget_reservation, response_cost=response_cost
)
counter_keys: Final = frozenset(
get_reserved_counter_keys(budget_reservation=budget_reservation)
) | post_call_counter_keys(
token=user_api_key,
team_id=team_id,
user_id=user_id,
org_id=org_id,
end_user_id=end_user_id,
tags=request_tags,
model_access_groups=model_access_groups,
project_id=project_id,
)
with spend_counter_batch_scope(spend_counter_cache.redis_cache, counter_keys=counter_keys):
return await _update_database_and_spend_counters_in_batch(
proxy_logging_obj=proxy_logging_obj,
increment_spend_counters=increment_spend_counters,
user_api_key=user_api_key,
user_id=user_id,
end_user_id=end_user_id,
team_id=team_id,
org_id=org_id,
kwargs=kwargs,
completion_response=completion_response,
start_time=start_time,
end_time=end_time,
response_cost=response_cost,
budget_reservation=budget_reservation,
request_tags=request_tags,
model_access_groups=model_access_groups,
project_id=project_id,
update_cache_read_keys=update_cache_read_keys,
)
async def _update_database_and_spend_counters_in_batch(
proxy_logging_obj: "ProxyLogging",
increment_spend_counters: _IncrementSpendCounters,
user_api_key: str | None,
user_id: str | None,
end_user_id: str | None,
team_id: str | None,
org_id: str | None,
kwargs: dict,
completion_response: object,
start_time: datetime | None,
end_time: datetime | None,
response_cost: float,
budget_reservation: dict | None,
request_tags: list[str] | None,
model_access_groups: Sequence[str] | None,
project_id: str | None,
update_cache_read_keys: Sequence[str],
) -> bool:
from litellm.proxy.proxy_server import arm_update_cache_read
try:
charged: Final = await proxy_logging_obj.db_spend_update_writer.update_database(
token=user_api_key,
@ -730,6 +801,7 @@ async def _update_database_and_spend_counters(
await _release_budget_reservation(budget_reservation=budget_reservation)
return False
await arm_update_cache_read(update_cache_read_keys)
try:
await increment_spend_counters(
token=user_api_key,
@ -762,11 +834,13 @@ async def _reconcile_budget_reservation_before_db_update(
budget_reservation: dict, # mutable-ok: reconcile_budget_reservation stamps applied_adjustment on the caller's shared reservation dict
response_cost: float,
) -> None:
"""Reseeds the reserved counters that were flushed since reservation; the adjustments themselves are written by ``increment_spend_counters`` in the same pipeline as its increments, or by
the release / invalidation that runs when the spend write fails."""
from litellm.proxy.spend_tracking.budget_reservation import reconcile_budget_reservation
try:
await reconcile_budget_reservation(
budget_reservation=budget_reservation, actual_cost=response_cost, finalize=False
_ = await reconcile_budget_reservation(
budget_reservation=budget_reservation, actual_cost=response_cost, finalize=False, apply_consistent=False
)
except Exception: # noqa: BLE001 # a failed reconcile must not block the spend write; the counters are dropped instead
verbose_proxy_logger.warning(

View file

@ -8,6 +8,7 @@ POST /auto_router/validate_complexity_router_config - Dry-run the complexity-rou
from collections.abc import Mapping, Sequence
from datetime import datetime, timedelta, timezone
from itertools import chain, groupby
from math import isclose
from types import MappingProxyType
from typing import TYPE_CHECKING, Annotated, Final, Protocol
from uuid import uuid4
@ -31,6 +32,7 @@ from litellm.proxy.auth.auth_checks import (
can_key_call_resolved_model,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.db.autorouter_savings_comparison import historical_session_comparisons
from litellm.proxy.db.autorouter_session_rollup import (
AUTOROUTER_BENCHMARKS_SQL,
bounded_session_id,
@ -651,7 +653,9 @@ class _SessionAggRow(BaseModel):
saved_spend: float
savings_estimated_turns: int = 0
savings_estimated_actual_spend: float = 0.0
savings_estimated_classifier_cost: float | None = None
savings_estimated_saved_spend: float = 0.0
savings_comparison_complete: bool = True
classifier_cost: float
classifier_cost_recorded_turns: int
session_seconds: float
@ -679,18 +683,25 @@ def _cache_bucket(turns: int, hits: int) -> AutoRouterCacheBucket:
def _savings_cohort(
turns: int, estimated_turns: int, actual_spend: float, saved_spend: float
turns: int, estimated_turns: int, actual_spend: float, saved_spend: float, recorded_savings: float
) -> tuple[float | None, float | None]:
if turns > 0 and estimated_turns == 0:
if turns > 0 and estimated_turns == 0 and recorded_savings == 0:
return None, None
return saved_spend, actual_spend + saved_spend
if not isclose(saved_spend, recorded_savings, rel_tol=1e-9, abs_tol=1e-9):
return recorded_savings, None
return recorded_savings, actual_spend + recorded_savings
def _benchmark_totals(row: _SessionAggRow) -> AutoRouterBenchmarkTotals:
return_misses: Final = row.return_turns - row.return_hits
saved_spend, baseline_spend = _savings_cohort(
row.turns, row.savings_estimated_turns, row.savings_estimated_actual_spend, row.savings_estimated_saved_spend
saved_spend, compared_baseline = _savings_cohort(
row.turns,
row.savings_estimated_turns,
row.savings_estimated_actual_spend,
row.savings_estimated_saved_spend,
row.saved_spend,
)
baseline_spend: Final = compared_baseline if row.savings_comparison_complete else None
sessions: Final = row.sessions
return AutoRouterBenchmarkTotals(
sessions=sessions,
@ -701,13 +712,12 @@ def _benchmark_totals(row: _SessionAggRow) -> AutoRouterBenchmarkTotals:
spend=row.spend,
savings_estimated_turns=row.savings_estimated_turns,
savings_estimated_actual_spend=row.savings_estimated_actual_spend,
savings_estimated_classifier_cost=row.savings_estimated_classifier_cost if baseline_spend is not None else None,
saved_spend=saved_spend,
classifier_cost=row.classifier_cost if row.classifier_cost_recorded_turns == row.turns else None,
baseline_spend=baseline_spend,
saved_pct=_pct(saved_spend, baseline_spend) if saved_spend is not None and baseline_spend is not None else None,
saved_per_session=(row.savings_estimated_saved_spend / sessions if sessions else 0.0)
if row.savings_estimated_turns == row.turns
else None,
saved_per_session=(saved_spend / sessions if sessions else 0.0) if saved_spend is not None else None,
cache=AutoRouterCacheStats(
coverage_pct=_pct(row.covered_turns, row.turns),
hit_rate_pct=_pct(row.cache_hits, row.covered_turns),
@ -739,6 +749,7 @@ def _benchmark_group(row: _SessionAggRow) -> AutoRouterBenchmarkGroup:
saved_spend=totals.saved_spend,
savings_estimated_turns=totals.savings_estimated_turns,
savings_estimated_actual_spend=totals.savings_estimated_actual_spend,
savings_estimated_classifier_cost=totals.savings_estimated_classifier_cost,
classifier_cost=totals.classifier_cost,
baseline_spend=totals.baseline_spend,
saved_pct=totals.saved_pct,
@ -772,7 +783,13 @@ def _summed_agg_row(rows: Sequence[_SessionAggRow]) -> _SessionAggRow:
saved_spend=sum(row.saved_spend for row in rows),
savings_estimated_turns=sum(row.savings_estimated_turns for row in rows),
savings_estimated_actual_spend=sum(row.savings_estimated_actual_spend for row in rows),
savings_estimated_classifier_cost=(
sum(row.savings_estimated_classifier_cost or 0.0 for row in rows)
if all(row.savings_estimated_classifier_cost is not None for row in rows)
else None
),
savings_estimated_saved_spend=sum(row.savings_estimated_saved_spend for row in rows),
savings_comparison_complete=all(row.savings_comparison_complete for row in rows),
classifier_cost=sum(row.classifier_cost for row in rows),
classifier_cost_recorded_turns=sum(row.classifier_cost_recorded_turns for row in rows),
session_seconds=sum(row.session_seconds for row in rows),
@ -847,8 +864,8 @@ async def get_auto_router_benchmarks(
Benchmarks for the auto-router dashboard: session shape, savings against the configured
baseline, and prompt-caching behaviour bucketed by what the router did.
Reads session rollups folded once per request at spend-write time, so this endpoint
never scans LiteLLM_SpendLogs. A user filter selects only turns attributed to that
Reads session rollups folded once per request at spend-write time, with bounded
retained-log recovery for historical comparisons. A user filter selects only turns attributed to that
internal user when written; older key-only history remains outside user views. A session
is in the window when it overlaps it: its last turn is on or after start_date and its first turn is on or before
end_date. Overall hit rate is over telemetry-bearing turns; each bucket's hit rate is
@ -882,7 +899,44 @@ async def get_auto_router_benchmarks(
api_key,
user_id,
)
rows: Final = _SESSION_AGG_ROWS.validate_python(raw_rows or ())
recorded_rows: Final = _SESSION_AGG_ROWS.validate_python(raw_rows or ())
comparisons: Final = (
await historical_session_comparisons(
prisma_client,
start_day.isoformat(),
(end_day + timedelta(days=1)).isoformat(),
api_key,
user_id,
)
if any(row.savings_estimated_turns < row.turns for row in recorded_rows)
else MappingProxyType({})
)
covered_rows: Final = tuple(
row.model_copy(
update={
**comparison.coverage_fields(row.saved_spend, row.turns),
"savings_estimated_classifier_cost": comparison.classifier_cost,
"savings_comparison_complete": comparison.complete and comparison.turns == row.turns,
}
)
if (comparison := comparisons.get((row.router_name, row.router_type)))
else row.model_copy(update={"savings_comparison_complete": row.savings_estimated_turns == row.turns})
for row in recorded_rows
)
rows: Final = tuple(
row.model_copy(
update={
"savings_comparison_complete": row.savings_comparison_complete
and isclose(
row.saved_spend,
row.savings_estimated_saved_spend,
rel_tol=1e-9,
abs_tol=1e-9,
),
}
)
for row in covered_rows
)
groups: Final = (
*(_benchmark_group(row) for row in rows),
*_idle_router_groups(llm_router, frozenset((row.router_name, row.router_type) for row in rows)),
@ -920,15 +974,43 @@ async def get_auto_router_session(
if prisma_client is None:
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
row: Final = await AutoRouterSessionRepository(prisma_client).find_latest_for_key(
recorded: Final = await AutoRouterSessionRepository(prisma_client).find_latest_for_key(
user_api_key_dict.api_key, bounded_session_id(session_id)
)
if row is None:
if recorded is None:
raise HTTPException(
status_code=404, detail=f"No auto-routed turns recorded for session {session_id!r} under this key"
)
saved_spend, baseline_spend = _savings_cohort(
row.turns, row.savings_estimated_turns, row.savings_estimated_actual_spend, row.savings_estimated_saved_spend
comparisons: Final = (
await historical_session_comparisons(
prisma_client,
recorded.first_turn_at.isoformat(),
(recorded.last_turn_at + timedelta(microseconds=1)).isoformat(),
user_api_key_dict.api_key,
None,
bounded_session_id(session_id),
)
if recorded.savings_estimated_turns < recorded.turns
else MappingProxyType({})
)
comparison: Final = comparisons.get((recorded.router_name, recorded.router_type))
row: Final = (
recorded.model_copy(update=comparison.coverage_fields(recorded.saved_spend, recorded.turns))
if comparison
else recorded
)
saved_spend, compared_baseline = _savings_cohort(
row.turns,
row.savings_estimated_turns,
row.savings_estimated_actual_spend,
row.savings_estimated_saved_spend,
row.saved_spend,
)
baseline_spend: Final = (
compared_baseline
if row.savings_estimated_turns == row.turns
or (comparison and comparison.complete and comparison.turns == row.turns)
else None
)
return AutoRouterSessionResponse(
session_id=session_id,
@ -943,7 +1025,7 @@ async def get_auto_router_session(
baseline_spend=baseline_spend if row.savings_estimated_turns == row.turns else None,
savings_estimated_baseline_spend=baseline_spend,
baseline_model=row.baseline_model,
baseline_models=row.savings_estimated_baseline_models,
baseline_models=row.baseline_models,
)

View file

@ -16,6 +16,7 @@ from litellm.proxy.spend_tracking.key_metadata_recovery import (
recover_cli_session_key_metadata,
recover_double_hashed_key_metadata,
recover_key_metadata_from_spend_logs,
recover_key_owner_from_daily_spend,
)
from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled
from litellm.proxy.utils import PrismaClient
@ -468,6 +469,17 @@ def _parse_spend_date(raw: str | None) -> datetime | None:
_EMPTY_KEY_METADATA: Final[Mapping[str, _KeyMetadataDict]] = MappingProxyType({})
def _metadata_with_recovered_owner(
metadata: Mapping[str, _KeyMetadataDict],
key: str,
owner: str,
) -> _KeyMetadataDict:
current: Final = metadata.get(key)
if current is None:
return {"user_id": owner}
return {**current, "user_id": owner}
async def get_api_key_metadata(
prisma_client: PrismaClient,
api_keys: AbstractSet[str],
@ -530,7 +542,19 @@ async def get_api_key_metadata(
else _EMPTY_KEY_METADATA
)
combined: Final = MappingProxyType({**after_token_recovery, **from_spend_logs})
return await attach_user_details(prisma_client, combined)
ownerless: Final = frozenset(
key
for key in api_keys
if not combined.get(key, {}).get("user_id") and not combined.get(key, {}).get("key_exists")
)
owners: Final = await recover_key_owner_from_daily_spend(prisma_client, ownerless)
metadata_with_owners: Final[Mapping[str, _KeyMetadataDict]] = MappingProxyType(
{
**combined,
**{key: _metadata_with_recovered_owner(combined, key, owner) for key, owner in owners.items()},
}
)
return await attach_user_details(prisma_client, metadata_with_owners)
def _adjust_dates_for_timezone(
@ -944,7 +968,7 @@ async def _aggregate_spend_records(
record.api_key for record in records if record.api_key and record.api_key != PTU_SENTINEL_API_KEY
}
api_key_metadata: dict[str, _KeyMetadataDict] = {}
api_key_metadata: Mapping[str, _KeyMetadataDict] = MappingProxyType({})
if api_keys:
api_key_metadata = await get_api_key_metadata(
prisma_client, api_keys, _spend_logs_window(frozenset(record.date for record in records))
@ -1144,7 +1168,7 @@ async def _aggregate_grouping_sets_records(
"""Async wrapper: fetch api_key_metadata, then dispatch on a worker thread."""
api_keys: Final[set[str]] = {r.api_key for r in records if r.api_key and r.api_key != PTU_SENTINEL_API_KEY}
api_key_metadata: dict[str, _KeyMetadataDict] = {}
api_key_metadata: Mapping[str, _KeyMetadataDict] = MappingProxyType({})
if api_keys:
api_key_metadata = await get_api_key_metadata(
prisma_client, api_keys, _spend_logs_window(frozenset(r.date for r in records))

View file

@ -0,0 +1,25 @@
from typing import Final
from starlette.types import ASGIApp, Receive, Scope, Send
from litellm.caching.redis_batch import request_redis_batch_scope
_REQUEST_SCOPES: Final = frozenset({"http", "websocket"})
class RedisRequestBatchMiddleware:
"""Opens the request's Redis batch scope so auth, admission and routing reads issued anywhere in the
request (dependencies, the endpoint, tasks it spawns) share one pipeline per Redis backend."""
def __init__(self, app: ASGIApp) -> None:
self.app = app
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
if scope["type"] not in _REQUEST_SCOPES:
await self.app(scope, receive, send)
return
with request_redis_batch_scope() as batches:
try:
await self.app(scope, receive, send)
finally:
await batches.flush_all()

View file

@ -269,6 +269,12 @@ import litellm._redis
from litellm import Router
from litellm._logging import _redact_string, verbose_proxy_logger, verbose_router_logger
from litellm.caching.caching import DualCache, RedisCache
from litellm.caching.dual_cache import DeclaredBatchRead
from litellm.caching.redis_batch import (
active_post_call_redis_batch,
active_request_redis_batches,
drain_post_call_redis_batches,
)
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError, is_redis_timeout_failure
from litellm.caching.redis_cluster_cache import RedisClusterCache
from litellm.constants import (
@ -681,6 +687,7 @@ from litellm.proxy.middleware.billable_request_metrics_middleware import (
from litellm.proxy.middleware.budget_reservation_release_middleware import (
BudgetReservationReleaseMiddleware,
)
from litellm.proxy.middleware.redis_request_batch_middleware import RedisRequestBatchMiddleware
from litellm.proxy.plugin_routes import (
register_plugins_from_config,
)
@ -759,6 +766,7 @@ from litellm.proxy.shutdown.scheduled_jobs import (
from litellm.proxy.spend_tracking.budget_reservation import (
get_budget_window_start,
release_unbound_budget_reservation,
stamp_budget_reservation_actual_cost,
)
from litellm.proxy.spend_tracking.spend_capture_rate import (
run_scheduled_spend_capture_rate_check,
@ -1110,6 +1118,7 @@ async def proxy_shutdown_event(worker_heartbeat: ProxyWorkerHeartbeat | None = N
verbose_proxy_logger.debug("Disconnecting from Prisma")
await prisma_client.disconnect()
await drain_post_call_redis_batches()
if litellm.cache is not None:
await litellm.cache.disconnect()
@ -2416,6 +2425,7 @@ app.add_middleware(
sink_factory=lambda: gateway_request_accumulator if prisma_client is not None else None,
)
app.add_middleware(BudgetReservationReleaseMiddleware, release=release_unbound_budget_reservation)
app.add_middleware(RedisRequestBatchMiddleware)
app.add_middleware(InFlightRequestsMiddleware)
app.add_middleware(SecurityHeadersMiddleware)
@ -2846,13 +2856,16 @@ async def _repair_stale_spend_counter(counter_key: str, db_spend: float) -> None
if spend_counter_cache.redis_cache is not None:
forget_spend_counter(counter_key)
try:
await spend_counter_cache.redis_cache.async_set_max(key=counter_key, value=db_spend)
repaired: Final = await spend_counter_cache.redis_cache.async_set_max(key=counter_key, value=db_spend)
except Exception:
verbose_proxy_logger.debug(
"Unable to repair stale spend counter %s in Redis",
counter_key,
exc_info=True,
)
return
if repaired is not None:
record_spend_counter_value(counter_key, repaired)
async def reseed_spend_counter_from_db(counter_key: str) -> bool:
@ -3049,13 +3062,17 @@ async def _increment_spend_counters_batched(
model_access_groups: Sequence[str] | None,
project_id: str | None = None,
):
"""Runs inside one spend counter batch: the reservation reconcile and the warm checks share a single MGET."""
reserved_counter_keys: Final = await _reconcile_budget_reservation_for_counter_update(
"""Runs inside one spend counter batch: the reservation reconcile and the warm checks share a single MGET, and
the reconcile adjustments go out in the same INCRBYFLOAT pipeline as the counter increments."""
reservation_update: Final = await _reconcile_budget_reservation_for_counter_update(
budget_reservation=budget_reservation,
response_cost=response_cost,
)
reserved_counter_keys: Final = reservation_update.reserved_counter_keys
if response_cost is None or response_cost == 0:
await _apply_spend_counter_increments(pending=reservation_update.pending)
stamp_budget_reservation_actual_cost(budget_reservation=budget_reservation, actual_cost=response_cost)
if budget_reservation is not None:
budget_reservation["finalized"] = True
return
@ -3276,7 +3293,8 @@ async def _increment_spend_counters_batched(
for item in scope
if not isinstance(item, BaseException)
)
await _apply_spend_counter_increments(pending=pending)
await _apply_spend_counter_increments(pending=reservation_update.pending + pending)
stamp_budget_reservation_actual_cost(budget_reservation=budget_reservation, actual_cost=response_cost)
if scope_errors:
raise scope_errors[0]
@ -3284,12 +3302,21 @@ async def _increment_spend_counters_batched(
budget_reservation["finalized"] = True
@dataclass(frozen=True, slots=True)
class _ReservationCounterUpdate:
"""The reserved counters the direct increment must skip, and the adjustments that settle them on the actual
cost, still to be written; both empty when the reservation could not be reconciled and was dropped."""
reserved_counter_keys: frozenset[str] = frozenset()
pending: tuple[PendingSpendIncrement, ...] = ()
async def _reconcile_budget_reservation_for_counter_update(
budget_reservation: dict | None,
response_cost: float | None,
) -> set[str]:
) -> _ReservationCounterUpdate:
if budget_reservation is None or budget_reservation.get("finalized") is True:
return set()
return _ReservationCounterUpdate()
from litellm.proxy.spend_tracking.budget_reservation import (
get_reserved_counter_keys,
@ -3299,10 +3326,11 @@ async def _reconcile_budget_reservation_for_counter_update(
reserved_counter_keys: Final = get_reserved_counter_keys(budget_reservation=budget_reservation)
try:
await reconcile_budget_reservation(
pending: Final = await reconcile_budget_reservation(
budget_reservation=budget_reservation,
actual_cost=response_cost or 0.0,
finalize=False,
apply_consistent=False,
)
except Exception:
verbose_proxy_logger.warning(
@ -3315,8 +3343,8 @@ async def _reconcile_budget_reservation_for_counter_update(
verbose_proxy_logger.exception(
"Failed to invalidate reserved counters after reservation reconciliation failed"
)
return set()
return reserved_counter_keys
return _ReservationCounterUpdate()
return _ReservationCounterUpdate(reserved_counter_keys=frozenset(reserved_counter_keys), pending=pending)
async def _prepare_end_user_and_tag_spend_increments(
@ -3694,6 +3722,8 @@ async def _invalidate_spend_counter(counter_key: str):
async def _apply_spend_counter_increments(pending: Sequence[PendingSpendIncrement]) -> None:
if _defer_spend_counter_increments(pending):
return
try:
await increment_spend_counters_pipeline(pending=pending)
except Exception as e:
@ -3702,31 +3732,148 @@ async def _apply_spend_counter_increments(pending: Sequence[PendingSpendIncremen
raise
async def increment_spend_counters_pipeline(pending: Sequence[PendingSpendIncrement]) -> None:
"""One INCRBYFLOAT+EXPIRE pipeline for every pending counter; on failure every counter is invalidated
before the error propagates, so no caller can read a half-applied batch."""
def _defer_spend_counter_increments(pending: Sequence[PendingSpendIncrement]) -> bool:
"""Post-call increments ride the request's post-call pipeline with the other counters. Each counter's
new value lands in memory when the pipeline settles; a failed one is invalidated so no reader trusts a
counter whose increment may not have applied, as ``increment_spend_counters_pipeline`` does."""
redis_cache: Final = spend_counter_cache.redis_cache
if redis_cache is None or not pending:
return False
batch: Final = active_post_call_redis_batch(redis_cache)
if batch is None:
return False
ttl: Final = redis_cache.get_ttl()
for item in pending:
batch.increment(item.counter_key, item.increment, ttl).on_settled(_settle_spend_counter_increment(item))
return True
def _settle_spend_counter_increment(item: PendingSpendIncrement) -> Callable[[asyncio.Future[float]], Awaitable[None]]:
async def settle(future: asyncio.Future[float]) -> None:
if not future.cancelled() and future.exception() is None:
current_value: Final = float(future.result())
spend_counter_cache.in_memory_cache.set_cache(key=item.counter_key, value=current_value)
record_spend_counter_value(item.counter_key, current_value)
return
if future.cancelled():
if spend_counter_cache.in_memory_cache.get_cache(key=item.counter_key) is not None:
spend_counter_cache.in_memory_cache.increment_cache(key=item.counter_key, value=item.increment)
return
verbose_proxy_logger.warning(
"Spend counter %s increment did not land in the post-call pipeline; invalidating it", item.counter_key
)
await _invalidate_spend_counter(counter_key=item.counter_key)
return settle
async def increment_spend_counters_pipeline(pending: Sequence[PendingSpendIncrement]) -> tuple[float | None, ...]:
"""One INCRBYFLOAT+EXPIRE pipeline for every pending counter, returning each counter's new value in order; on
failure every counter is invalidated before the error propagates, so no caller can read a half-applied batch."""
if spend_counter_cache.redis_cache is None:
return await run_spend_counter_pipeline(pending=pending)
try:
return await run_spend_counter_pipeline(pending=pending)
except Exception:
await asyncio.gather(*(_invalidate_spend_counter(counter_key=item.counter_key) for item in pending))
raise
async def run_spend_counter_pipeline(pending: Sequence[PendingSpendIncrement]) -> tuple[float | None, ...]:
"""The pipeline behind ``increment_spend_counters_pipeline`` without its invalidation: the caller decides what
happens to counters whose increment may or may not have landed when the pipeline fails."""
if not pending:
return
return ()
redis_cache: Final = spend_counter_cache.redis_cache
if redis_cache is None:
for item in pending:
await SpendCounterReseed.increment_in_memory(
spend_counter_cache=spend_counter_cache, counter_key=item.counter_key, increment=item.increment
)
return
return tuple(
[
await SpendCounterReseed.increment_in_memory(
spend_counter_cache=spend_counter_cache, counter_key=item.counter_key, increment=item.increment
)
for item in pending
]
)
ttl: Final = redis_cache.get_ttl()
increment_list: Final = [ # mutable-ok: async_increment_pipeline signature requires list[RedisPipelineIncrementOperation]
RedisPipelineIncrementOperation(key=item.counter_key, increment_value=item.increment, ttl=ttl)
for item in pending
]
try:
results: Final = await redis_cache.async_increment_pipeline(increment_list=increment_list)
except Exception:
await asyncio.gather(*(_invalidate_spend_counter(counter_key=item.counter_key) for item in pending))
raise
results: Final = await redis_cache.async_increment_pipeline(increment_list=increment_list)
for item, current_value in zip(pending, results or ()):
spend_counter_cache.in_memory_cache.set_cache(key=item.counter_key, value=current_value)
record_spend_counter_value(item.counter_key, float(current_value))
return tuple(float(current_value) for current_value in results or ())
def update_cache_read_keys(
user_id: str | None,
end_user_id: str | None,
team_id: str | None,
tags: Sequence[object] | None,
response_cost: float | None,
) -> tuple[str, ...]:
if response_cost is None:
return ()
user_keys: tuple[str, ...] = (user_id, GLOBAL_PROXY_SPEND_CACHE_KEY) if user_id is not None else ()
end_user_keys: tuple[str, ...] = (end_user_cache_key(end_user_id),) if end_user_id is not None else ()
team_keys: tuple[str, ...] = (f"team_id:{team_id}",) if team_id is not None else ()
tag_keys: tuple[str, ...] = tuple(tag_cache_key(tag) for tag in tags or () if isinstance(tag, str) and tag)
return user_keys + end_user_keys + team_keys + tag_keys
_UPDATE_CACHE_PREFETCH_SLOT: Final = "update_cache_read"
async def arm_update_cache_read(keys: Sequence[str], cache: DualCache | None = None) -> None:
"""Declares the ``update_cache`` read on the request pipeline once the spend is persisted, so it rides the same
round trip as the post-call spend counter read instead of its own."""
request: Final = active_request_redis_batches()
target: Final = user_api_key_cache if cache is None else cache
if request is None or target.redis_cache is None or not keys:
return
request.prefetched[_UPDATE_CACHE_PREFETCH_SLOT] = await target.declare_batch_get(
keys, request.batch(target.redis_cache)
)
async def _take_armed_update_cache_read(keys: Sequence[str], cache: DualCache) -> Mapping[str, object] | None:
request: Final = active_request_redis_batches()
if request is None:
return None
armed: Final = request.prefetched.pop(_UPDATE_CACHE_PREFETCH_SLOT, None)
if not isinstance(armed, DeclaredBatchRead) or armed.keys != tuple(keys):
return None
values: Final = await cache.async_resolve_batch_get(armed)
return MappingProxyType({key: value for key, value in zip(keys, values) if value is not None})
async def _read_update_cache_values(
keys: Sequence[str], parent_otel_span: Span | None, cache: DualCache | None = None
) -> Mapping[str, object]:
"""One batched read for every object ``update_cache`` refreshes; a failed read leaves them all untouched,
exactly as a failed per-object GET left that object untouched."""
if not keys:
return MappingProxyType({})
target: Final = user_api_key_cache if cache is None else cache
try:
armed: Final = await _take_armed_update_cache_read(keys, target)
if armed is not None:
return armed
values: Final = await target.async_batch_get_cache(
keys=list(keys), parent_otel_span=parent_otel_span, throttle_redis=False
)
except Exception as e:
verbose_proxy_logger.warning(
"Spend tracking - failed to read cached spend objects. Budget enforcement may use stale spend values. "
"keys=%s - %s",
keys,
str(e),
)
return MappingProxyType({})
if values is None:
return MappingProxyType({})
return MappingProxyType({key: value for key, value in zip(keys, values) if value is not None})
async def update_cache(
@ -3745,6 +3892,12 @@ async def update_cache(
"""
values_to_update_in_cache: Final[list[tuple[str, object]]] = []
cached_values: Final = await _read_update_cache_values(
keys=update_cache_read_keys(
user_id=user_id, end_user_id=end_user_id, team_id=team_id, tags=tags, response_cost=response_cost
),
parent_otel_span=parent_otel_span,
)
### UPDATE KEY SPEND ###
async def _update_key_cache(token: str, response_cost: float):
@ -3810,7 +3963,7 @@ async def update_cache(
# Fetch the existing cost for the given user
if _id is None:
continue
cached_user = await user_api_key_cache.async_get_cache(key=_id)
cached_user = cached_values.get(_id)
if cached_user is None:
# do nothing if there is no cache value
return
@ -3833,11 +3986,11 @@ async def update_cache(
)
)
## UPDATE GLOBAL PROXY ##
global_proxy_spend: Final = await user_api_key_cache.async_get_cache(key=GLOBAL_PROXY_SPEND_CACHE_KEY)
if global_proxy_spend is None:
global_proxy_spend: Final = cached_values.get(GLOBAL_PROXY_SPEND_CACHE_KEY)
if not isinstance(global_proxy_spend, (int, float)):
# do nothing if not in cache
return
elif response_cost is not None and global_proxy_spend is not None:
elif response_cost is not None:
increment: Final = global_proxy_spend + response_cost
values_to_update_in_cache.append((GLOBAL_PROXY_SPEND_CACHE_KEY, increment))
except Exception as e:
@ -3859,7 +4012,7 @@ async def update_cache(
_id: Final = end_user_cache_key(end_user_id)
try:
# Fetch the existing cost for the given user
cached_end_user: Final = await user_api_key_cache.async_get_cache(key=_id)
cached_end_user: Final = cached_values.get(_id)
if cached_end_user is None:
# if user does not exist in LiteLLM_UserTable, create a new user
# do nothing if end-user not in api key cache
@ -3900,7 +4053,7 @@ async def update_cache(
_id: Final = f"team_id:{team_id}"
try:
cached_team: Final = await user_api_key_cache.async_get_cache(key=_id)
cached_team: Final = cached_values.get(_id)
if cached_team is None:
# do nothing if team not in api key cache
return
@ -3950,7 +4103,7 @@ async def update_cache(
cache_key = tag_cache_key(tag_name)
# Fetch the existing tag object from cache
cached_tag = await user_api_key_cache.async_get_cache(key=cache_key)
cached_tag = cached_values.get(cache_key)
if cached_tag is None:
# do nothing if tag not in api key cache
continue
@ -9296,9 +9449,10 @@ def _fast_serialize_simple_model_response_stream(
"object": getattr(chunk, "object", None),
"created": getattr(chunk, "created", None),
"model": model,
"service_tier": getattr(chunk, "service_tier", None),
"choices": [choice_dict],
}
for top_level_key in ("id", "object", "created"):
for top_level_key in ("id", "object", "created", "service_tier"):
if payload[top_level_key] is None:
payload.pop(top_level_key)
return orjson.dumps(payload)

View file

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

View file

@ -4,7 +4,7 @@ import asyncio
import json
import math
import time
from collections.abc import Mapping, Sequence
from collections.abc import AsyncIterator, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from types import MappingProxyType
@ -290,7 +290,6 @@ async def reserve_budget_for_request(
raw_body=raw_body,
)
current_spend_by_counter_key: Final[dict[str, float]] = {}
reservation_cost = estimate_request_max_cost(
request_body=request_body,
route=route,
@ -306,46 +305,17 @@ async def reserve_budget_for_request(
applied_entries: Final[list[dict[str, float | str]]] = []
try:
with _counters_batch_scope(frozenset(counter.counter_key for counter in counters)):
for counter in counters:
entry = _counter_to_reservation_entry(
counter=counter,
reserved_cost=reservation_cost,
)
applied_entries.append(entry)
try:
reserved_value = await _reserve_counter(
counter=counter,
reservation_cost=reservation_cost,
)
except _CounterReservationUnavailable as exc:
if exc.touched_counter and not exc.counter_invalidated:
await _release_applied_entries_best_effort(
entries=[entry],
default_reserved_cost=reservation_cost,
)
applied_entries.remove(entry)
if fail_closed_budget_enforcement:
_raise_reservation_unavailable(counter_key=counter.counter_key)
continue
if reserved_value is not None:
current_spend = reserved_value
else:
cached_spend = current_spend_by_counter_key.get(counter.counter_key)
if cached_spend is None:
cached_spend = await _get_current_counter_value(counter=counter)
current_spend = cached_spend + reservation_cost
if current_spend > counter.max_budget:
reservation_cost = await _apply_over_budget_reservation_policy(
counter=counter,
valid_token=valid_token,
entry=entry,
applied_entries=applied_entries,
reservation_cost=reservation_cost,
current_spend=current_spend,
fail_closed_budget_enforcement=fail_closed_budget_enforcement,
)
continue
reservable: Final = await _initialize_reservation_counters(
counters=counters,
fail_closed_budget_enforcement=fail_closed_budget_enforcement,
)
reservation_cost = await _reserve_reservable_counters(
reservable=reservable,
valid_token=valid_token,
applied_entries=applied_entries,
reservation_cost=reservation_cost,
fail_closed_budget_enforcement=fail_closed_budget_enforcement,
)
except Exception:
await _release_applied_entries_best_effort(
entries=applied_entries,
@ -381,19 +351,39 @@ async def reconcile_budget_reservation(
budget_reservation: dict | None,
actual_cost: float | None,
finalize: bool = True,
) -> None:
apply_consistent: bool = True,
) -> tuple[PendingSpendIncrement, ...]:
"""Settle every reserved counter on ``actual_cost``. With ``apply_consistent`` False the adjustments for
counters that still hold the reservation are returned instead of written, so the caller can pipeline them with
its own increments and then call ``stamp_budget_reservation_actual_cost``."""
if not budget_reservation or budget_reservation.get("finalized") is True:
return
return ()
reserved_cost: Final = float(budget_reservation.get("reserved_cost") or 0.0)
actual: Final = float(actual_cost or 0.0)
await _set_reserved_entries_actual_cost(
pending: Final = await _set_reserved_entries_actual_cost(
entries=budget_reservation.get("entries") or [],
actual_cost=actual,
default_reserved_cost=reserved_cost,
apply_consistent=apply_consistent,
)
if finalize:
budget_reservation["finalized"] = True
return pending
def stamp_budget_reservation_actual_cost(budget_reservation: dict | None, actual_cost: float | None) -> None:
"""Record that every reserved counter now holds ``actual_cost``, once the adjustments handed back by
``reconcile_budget_reservation(apply_consistent=False)`` have been written."""
if not budget_reservation:
return
reserved_cost: Final = float(budget_reservation.get("reserved_cost") or 0.0)
actual: Final = float(actual_cost or 0.0)
for entry in budget_reservation.get("entries") or []:
if "counter_key" in entry:
entry["applied_adjustment"] = actual - _get_entry_reserved_cost(
entry=entry, default_reserved_cost=reserved_cost
)
async def release_budget_reservation(budget_reservation: dict | None) -> None:
@ -917,18 +907,40 @@ def _coerce_window(window: object) -> Mapping[str, object]:
return dumped if isinstance(dumped, Mapping) else {}
async def _reserve_counter(
counter: _BudgetCounter,
reservation_cost: float,
) -> float | None:
async def _initialize_reservation_counters(
counters: Sequence[_BudgetCounter],
fail_closed_budget_enforcement: bool,
) -> tuple[_BudgetCounter, ...]:
"""The counters whose current value is loaded, in order; one that cannot be loaded is skipped (or rejects the
request under fail-closed enforcement) exactly as it was when each counter was reserved on its own."""
return tuple([counter async for counter in _loaded_reservation_counters(counters, fail_closed_budget_enforcement)])
async def _loaded_reservation_counters(
counters: Sequence[_BudgetCounter], fail_closed_budget_enforcement: bool
) -> AsyncIterator[_BudgetCounter]:
for counter in counters:
if await _reservation_counter_loaded(counter, fail_closed_budget_enforcement):
yield counter
async def _reservation_counter_loaded(counter: _BudgetCounter, fail_closed_budget_enforcement: bool) -> bool:
try:
await _initialize_reservation_counter(counter=counter)
except _CounterReservationUnavailable:
if fail_closed_budget_enforcement:
_raise_reservation_unavailable(counter_key=counter.counter_key)
return False
return True
async def _initialize_reservation_counter(counter: _BudgetCounter) -> None:
from litellm.proxy.proxy_server import (
_ensure_spend_counter_initialized,
_ensure_window_spend_counter_initialized,
_increment_spend_counter_cache,
_invalidate_spend_counter,
)
attempted_increment = False
try:
if counter.source_cache_key is not None:
await _ensure_spend_counter_initialized(
@ -949,13 +961,6 @@ async def _reserve_counter(
counter.counter_key,
)
raise _CounterReservationUnavailable
attempted_increment = True
reserved_value: Final = await _increment_spend_counter_cache(
counter_key=counter.counter_key,
increment=reservation_cost,
)
return float(reserved_value) if reserved_value is not None else None
except _CounterReservationUnavailable:
raise
except Exception:
@ -964,20 +969,121 @@ async def _reserve_counter(
counter.counter_key,
exc_info=True,
)
counter_invalidated = False
try:
await _invalidate_spend_counter(counter_key=counter.counter_key)
counter_invalidated = True
except Exception:
verbose_proxy_logger.warning(
"Failed to invalidate spend counter after budget reservation failure for %s",
counter.counter_key,
exc_info=True,
)
raise _CounterReservationUnavailable(
touched_counter=attempted_increment,
counter_invalidated=counter_invalidated,
raise _CounterReservationUnavailable
async def _reserve_reservable_counters(
reservable: Sequence[_BudgetCounter],
valid_token: UserAPIKeyAuth | None,
applied_entries: list[dict[str, float | str]],
reservation_cost: float,
fail_closed_budget_enforcement: bool,
) -> float:
"""Charge the counters group by group (see ``_reservation_groups``), settling the over-budget policy on each
group before the next is charged, and hand back the reservation cost the policy left standing."""
current_spend_by_counter_key: Final = {
counter.counter_key: await _get_current_counter_value(counter=counter) for counter in reservable
}
for group in _reservation_groups(
counters=reservable,
current_spend_by_counter_key=current_spend_by_counter_key,
reservation_cost=reservation_cost,
):
charged_cost = reservation_cost
entries = tuple(_counter_to_reservation_entry(counter=counter, reserved_cost=charged_cost) for counter in group)
applied_entries.extend(entries)
reserved_values = await _reserve_counters(counters=group, entries=entries, reservation_cost=charged_cost)
if reserved_values is None:
for entry in entries:
applied_entries.remove(entry)
if fail_closed_budget_enforcement:
_raise_reservation_unavailable(counter_key=group[0].counter_key)
continue
for counter, entry, reserved_value in zip(group, entries, reserved_values):
if entry not in applied_entries:
continue
if reserved_value is not None:
current_spend = reserved_value - (charged_cost - reservation_cost)
else:
current_spend = current_spend_by_counter_key[counter.counter_key] + reservation_cost
if current_spend > counter.max_budget:
reservation_cost = await _apply_over_budget_reservation_policy(
counter=counter,
valid_token=valid_token,
entry=entry,
applied_entries=applied_entries,
reservation_cost=reservation_cost,
current_spend=current_spend,
fail_closed_budget_enforcement=fail_closed_budget_enforcement,
)
return reservation_cost
def _reservation_groups(
counters: Sequence[_BudgetCounter],
current_spend_by_counter_key: Mapping[str, float],
reservation_cost: float,
) -> tuple[tuple[_BudgetCounter, ...], ...]:
"""Every counter the batch read says still has room for the estimate is charged in one pipeline. As soon as one
does not, the counters are charged one at a time so the over-budget policy settles each before the next is
touched, and a rejection charges nothing after it."""
if not counters:
return ()
if all(
current_spend_by_counter_key[counter.counter_key] + reservation_cost <= counter.max_budget
for counter in counters
):
return (tuple(counters),)
return tuple((counter,) for counter in counters)
async def _reserve_counters(
counters: Sequence[_BudgetCounter],
entries: Sequence[dict[str, float | str]],
reservation_cost: float,
) -> tuple[float | None, ...] | None:
"""One INCRBYFLOAT pipeline reserves every counter. When it fails each counter is dropped, and one that cannot
be dropped is released instead in case its increment landed, so nothing is left to release by the caller."""
from litellm.proxy.proxy_server import _invalidate_spend_counter, run_spend_counter_pipeline
if not counters:
return ()
try:
reserved: Final = await run_spend_counter_pipeline(
pending=tuple(
PendingSpendIncrement(counter_key=counter.counter_key, increment=reservation_cost)
for counter in counters
)
)
except Exception:
verbose_proxy_logger.warning(
"Skipping budget reservation for %s because spend counter reservation failed",
tuple(counter.counter_key for counter in counters),
exc_info=True,
)
for counter, entry in zip(counters, entries):
try:
await _invalidate_spend_counter(counter_key=counter.counter_key)
except Exception:
verbose_proxy_logger.warning(
"Failed to invalidate spend counter after budget reservation failure for %s",
counter.counter_key,
exc_info=True,
)
await _release_applied_entries_best_effort(
entries=[entry], # mutable-ok: the release takes the reservation's list of entries
default_reserved_cost=reservation_cost,
)
return None
return tuple(reserved) + (None,) * (len(counters) - len(reserved))
async def _get_current_counter_value(counter: _BudgetCounter) -> float:
@ -1026,9 +1132,11 @@ async def _set_reserved_entries_actual_cost(
actual_cost: float,
default_reserved_cost: float,
reseed_on_inconsistent: bool = True,
) -> None:
"""Every reserved counter is read from one MGET and the consistent adjustments go out in one pipeline.
A counter that was flushed or reseeded since reservation is settled on its own after the pipeline."""
apply_consistent: bool = True,
) -> tuple[PendingSpendIncrement, ...]:
"""Every reserved counter is read from one MGET and the consistent adjustments go out in one pipeline, or are
returned unwritten when ``apply_consistent`` is False. A counter that was flushed or reseeded since reservation
is settled on its own after the pipeline."""
from litellm.proxy.proxy_server import increment_spend_counters_pipeline
with _counters_batch_scope(frozenset(str(entry["counter_key"]) for entry in entries if "counter_key" in entry)):
@ -1055,15 +1163,16 @@ async def _set_reserved_entries_actual_cost(
f"Cannot resize budget reservation against inconsistent counter {inconsistent[0].counter_key}"
)
applicable: Final = tuple(item for item, ok in zip(adjustments, consistent) if ok)
await increment_spend_counters_pipeline(
pending=tuple(
PendingSpendIncrement(counter_key=item.counter_key, increment=item.adjustment) for item in applicable
)
applicable_pending: Final = tuple(
PendingSpendIncrement(counter_key=item.counter_key, increment=item.adjustment) for item in applicable
)
if apply_consistent:
await increment_spend_counters_pipeline(pending=applicable_pending)
for item in inconsistent:
await _reseed_reserved_entry(item=item, actual_cost=actual_cost)
for item in adjustments:
for item in adjustments if apply_consistent else inconsistent:
item.entry["applied_adjustment"] = item.target_adjustment
return () if apply_consistent else applicable_pending
async def _reseed_reserved_entry(item: _EntryAdjustment, actual_cost: float) -> None:

View file

@ -61,6 +61,13 @@ WHERE COALESCE(key_alias, user_id, team_id) IS NOT NULL
GROUP BY api_key
"""
_DAILY_USER_SPEND_OWNER_SQL: Final = """
SELECT api_key, MIN(user_id) AS first_owner, MAX(user_id) AS last_owner
FROM "LiteLLM_DailyUserSpend"
WHERE api_key = ANY($1::text[]) AND user_id IS NOT NULL AND user_id <> ''
GROUP BY api_key
"""
_SPEND_LOG_STATEMENT_TIMEOUT_SQL: Final = f"SET LOCAL statement_timeout = {SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS}"
_SPEND_LOG_TRANSACTION_TIMEOUT: Final = timedelta(milliseconds=2 * SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS)
@ -104,8 +111,15 @@ class _SpendLogDigestRow(BaseModel):
)
class _DailyUserSpendOwnerRow(BaseModel):
api_key: str
first_owner: str | None = None
last_owner: str | None = None
_TOKEN_DIGEST_ROWS: Final = TypeAdapter(tuple[_TokenDigestRow, ...])
_SPEND_LOG_DIGEST_ROWS: Final = TypeAdapter(tuple[_SpendLogDigestRow, ...])
_DAILY_USER_SPEND_OWNER_ROWS: Final = TypeAdapter(tuple[_DailyUserSpendOwnerRow, ...])
_CACHED_KEY_METADATA: Final = TypeAdapter(KeyMetadataDict)
_SPEND_LOG_METADATA_CACHE: Final = InMemoryCache(
max_size_in_memory=SPEND_LOG_KEY_METADATA_CACHE_MAX_ITEMS,
@ -113,6 +127,7 @@ _SPEND_LOG_METADATA_CACHE: Final = InMemoryCache(
)
_SPEND_LOG_QUERY_LOCK: Final = asyncio.Lock()
_EMPTY_KEY_METADATA: Final[Mapping[str, KeyMetadataDict]] = MappingProxyType({})
_EMPTY_KEY_OWNERS: Final[Mapping[str, str]] = MappingProxyType({})
async def _db_or_empty(
@ -129,6 +144,16 @@ async def _db_or_empty(
return None
async def _rows_within_the_statement_timeout(
prisma_client: PrismaClient,
sql: str,
*params: object,
) -> 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)
return await transaction.query_raw(sql, *params)
async def _reverse_hash_key_metadata(
prisma_client: PrismaClient,
sql: str,
@ -152,6 +177,29 @@ async def _reverse_hash_key_metadata(
)
async def recover_key_owner_from_daily_spend(
prisma_client: PrismaClient,
keys: AbstractSet[str],
) -> Mapping[str, str]:
if not keys:
return _EMPTY_KEY_OWNERS
rows: Final = await _db_or_empty(
lambda: _rows_within_the_statement_timeout(prisma_client, _DAILY_USER_SPEND_OWNER_SQL, sorted(keys)),
"Failed daily-spend key owner recovery for %d keys: %s",
len(keys),
)
if rows is None:
return _EMPTY_KEY_OWNERS
return MappingProxyType(
{
row.api_key: owner
for row in _DAILY_USER_SPEND_OWNER_ROWS.validate_python(rows)
for owner in (_unanimous(row.first_owner, row.last_owner),)
if row.api_key in keys and owner is not None
}
)
@dataclass(frozen=True, slots=True)
class _UserDetails:
email: str | None
@ -309,24 +357,14 @@ def _cached_spend_log_metadata(
)
async def _spend_log_rows_within_the_statement_timeout(
prisma_client: PrismaClient,
digests: AbstractSet[str],
window: tuple[datetime, datetime],
) -> Sequence[Mapping[str, object]]:
start, end = window
async with prisma_client.db.tx(timeout=_SPEND_LOG_TRANSACTION_TIMEOUT) as transaction:
await transaction.execute_raw(_SPEND_LOG_STATEMENT_TIMEOUT_SQL)
return await transaction.query_raw(_SPEND_LOG_ALIAS_SQL, sorted(digests), start, end)
async def _query_spend_log_metadata(
prisma_client: PrismaClient,
digests: AbstractSet[str],
window: tuple[datetime, datetime],
) -> Mapping[str, KeyMetadataDict] | None:
start, end = window
rows: Final = await _db_or_empty(
lambda: _spend_log_rows_within_the_statement_timeout(prisma_client, digests, window),
lambda: _rows_within_the_statement_timeout(prisma_client, _SPEND_LOG_ALIAS_SQL, sorted(digests), start, end),
"Failed spend-log alias recovery for %d missing keys: %s",
len(digests),
)

View file

@ -10,6 +10,7 @@ from typing import Final
from pydantic import TypeAdapter
from litellm._logging import verbose_proxy_logger
from litellm.caching.redis_batch import BatchResult, RedisBatch, active_request_redis_batch
from litellm.caching.redis_cache import RedisCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.user_api_key_cache import (
@ -30,9 +31,13 @@ class PendingSpendIncrement:
class SpendCounterBatch:
"""Bound counters are read with one MGET on first use; counters bound later join the next MGET.
``async_batch_get_cache`` maps a clean miss to ``None`` and drops keys only when Redis failed, so an absent
key means "read it yourself" and a present ``None`` is an authoritative miss."""
key means "read it yourself" and a present ``None`` is an authoritative miss.
__slots__ = ("_fetched", "_keys", "_loaded", "_lock", "_open", "_redis_cache")
Inside a ``request_redis_batch_scope`` the MGET rides the request's pipeline instead: the batch's flush
hook declares whatever is bound but unread, so whoever flushes first (the auth object prefetch, usually)
carries the spend counters in the same round trip."""
__slots__ = ("_fetched", "_inflight", "_keys", "_loaded", "_lock", "_open", "_redis_cache", "_request_batch")
def __init__(self, redis_cache: RedisCache) -> None:
self._redis_cache: Final = redis_cache
@ -41,6 +46,10 @@ class SpendCounterBatch:
self._keys: frozenset[str] = frozenset()
self._fetched: frozenset[str] = frozenset()
self._loaded: Mapping[str, float | None] = _NO_VALUES
self._inflight: Final[list[BatchResult[Mapping[str, object]]]] = [] # mutable-ok: drained by _load
self._request_batch: Final[RedisBatch | None] = active_request_redis_batch(redis_cache)
if self._request_batch is not None:
self._request_batch.add_flush_hook(self._declare_pending)
@property
def counter_keys(self) -> frozenset[str]:
@ -85,6 +94,10 @@ class SpendCounterBatch:
async def _load(self) -> Mapping[str, float | None]:
async with self._lock:
if self._request_batch is not None:
self._declare_pending()
await self._collect_inflight()
return self._loaded
pending: Final = self._keys - self._fetched
if pending:
self._fetched = self._fetched | pending
@ -92,6 +105,26 @@ class SpendCounterBatch:
self._loaded = MappingProxyType({**fetched, **self._loaded})
return self._loaded
def _declare_pending(self) -> None:
"""Flush hook: put every bound-but-unread counter on the request pipeline that is about to go out."""
if self._request_batch is None or not self._open:
return
pending: Final = self._keys - self._fetched
if pending:
self._fetched = self._fetched | pending
self._inflight.append(self._request_batch.mget(sorted(pending)))
async def _collect_inflight(self) -> None:
results: Final = tuple(self._inflight)
self._inflight.clear()
for result in results:
try:
fetched: Mapping[str, float | None] = _CounterValues.validate_python(await result)
except Exception as e: # noqa: BLE001 # per-key reads take over and apply their own Redis fallback
verbose_proxy_logger.debug("spend counter batch read failed, falling back to per-key reads: %s", e)
continue
self._loaded = MappingProxyType({**fetched, **self._loaded})
async def _fetch(self, keys: frozenset[str]) -> Mapping[str, float | None]:
try:
return _CounterValues.validate_python(
@ -144,25 +177,43 @@ def release_spend_counter_batch() -> None:
batch.close()
def _iter_admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) -> Iterator[str]:
if token.token is not None:
yield f"spend:key:{token.token}"
if token.team_id is not None:
yield f"spend:team:{token.team_id}"
if token.user_id is not None:
yield f"spend:team_member:{token.user_id}:{token.team_id}"
if token.user_id is not None:
yield f"spend:user:{token.user_id}"
if end_user_id is not None:
def _iter_entity_counter_keys(
token: object,
team_id: object,
user_id: object,
org_id: object,
project_id: object,
end_user_id: object,
) -> Iterator[str]:
"""Only string ids name a counter; anything else (None, or an unresolved placeholder in synthetic
logging payloads) simply has no counter to bind."""
if isinstance(token, str):
yield f"spend:key:{token}"
if isinstance(team_id, str):
yield f"spend:team:{team_id}"
if isinstance(user_id, str):
yield f"spend:team_member:{user_id}:{team_id}"
if isinstance(user_id, str):
yield f"spend:user:{user_id}"
if isinstance(end_user_id, str):
yield f"spend:end_user:{end_user_id}"
if token.org_id is not None:
yield f"spend:org:{token.org_id}"
if token.project_id is not None:
yield project_spend_counter_key(token.project_id)
if isinstance(org_id, str):
yield f"spend:org:{org_id}"
if isinstance(project_id, str):
yield project_spend_counter_key(project_id)
def admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) -> frozenset[str]:
return frozenset(_iter_admission_counter_keys(token, end_user_id))
return frozenset(
_iter_entity_counter_keys(
token=token.token,
team_id=token.team_id,
user_id=token.user_id,
org_id=token.org_id,
project_id=token.project_id,
end_user_id=end_user_id,
)
)
def post_call_counter_keys(
@ -176,9 +227,15 @@ def post_call_counter_keys(
project_id: str | None = None,
) -> frozenset[str]:
"""Every counter ``increment_spend_counters`` warm-checks, except budget windows which bind on read."""
entity_keys: Final = admission_counter_keys(
UserAPIKeyAuth(token=token, team_id=team_id, user_id=user_id, org_id=org_id, project_id=project_id),
end_user_id,
entity_keys: Final = frozenset(
_iter_entity_counter_keys(
token=token,
team_id=team_id,
user_id=user_id,
org_id=org_id,
project_id=project_id,
end_user_id=end_user_id,
)
)
tag_keys: Final = frozenset(f"spend:tag:{tag}" for tag in tags or () if tag and isinstance(tag, str))
group_keys: Final = frozenset(

View file

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

View file

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

View file

@ -134,7 +134,7 @@ from litellm.router_strategy.least_busy import LeastBusyLoggingHandler
from litellm.router_strategy.lowest_cost import LowestCostLoggingHandler
from litellm.router_strategy.lowest_latency import LowestLatencyLoggingHandler
from litellm.router_strategy.lowest_tpm_rpm import LowestTPMLoggingHandler
from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2
from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2, PrefetchedUsage
from litellm.router_strategy.simple_shuffle import simple_shuffle
from litellm.router_strategy.tag_based_routing import (
_get_tags_from_request_kwargs,
@ -259,6 +259,7 @@ from litellm.router_utils.routing_groups import (
parse_routing_groups,
validate_routing_strategy,
)
from litellm.router_utils.routing_read_batch import RoutingPrefetch, RoutingReadBatch
from litellm.scheduler import FlowItem, Scheduler
from litellm.types.litellm_params import RoutingStrategyName
from litellm.types.llms.openai import (
@ -538,6 +539,10 @@ def _is_retriable_anthropic_status(status_code: int) -> bool:
return status_code == 429 or status_code >= 500
def _without_line_breaks(value: object) -> str:
return str(value).replace("\r", "").replace("\n", "")
def _anthropic_stream_error_is_gateway_verdict(chunk: object) -> bool:
"""AgenticAnthropicStreamingIterator's own retrieval-failure frame is the gateway's verdict, not a provider
failure: another deployment would rerun the same failed hook, so it reaches the client instead of falling back."""
@ -638,6 +643,24 @@ class FallbackAwareAnthropicMessagesStream:
def has_buffered_provider_output(self) -> bool:
return getattr(self._source_iterator, "has_buffered_provider_output", False) is True
@property
def chunks(self) -> list[ModelResponseStream] | None:
return cast( # cast-ok: chunks is a list of ModelResponseStream on the inner stream
"list[ModelResponseStream] | None", getattr(self._source_iterator, "chunks", None)
)
@property
def messages(self) -> list[AllMessageValues] | None:
return cast( # cast-ok: messages is a list of AllMessageValues on the inner stream
"list[AllMessageValues] | None", getattr(self._source_iterator, "messages", None)
)
@property
def model(self) -> str | None:
return cast( # cast-ok: model is a str on the inner stream
"str | None", getattr(self._source_iterator, "model", None)
)
def adopt_fallback_source(self, fallback_response: object) -> None:
self._source_iterator = fallback_response
self.fallback_headers_adopted = True
@ -1711,6 +1734,25 @@ class Router:
normalized for normalized in map(self._normalize_strategy, configured) if normalized is not None
)
def arm_routing_read_prefetch(self, model: str, request_kwargs: dict[str, object] | None = None) -> None:
"""Declare the cooldown read (and, for usage-based routing, the usage read) that
`async_get_available_deployment` will make for `model` on the request's Redis batch, so admission's
flush carries it. A miss (alias, no batch) costs nothing: routing then reads as it always has."""
try:
strategy, selector = self._get_routing_context(model, request_kwargs)
usage_selector: Final = (
selector
if strategy == "usage-based-routing-v2" and isinstance(selector, LowestTPMLoggingHandler_v2)
else None
)
deployments: Final = self.get_model_list(model_name=model)
if deployments:
RoutingPrefetch.arm(self, usage_selector, deployments)
except Exception as e: # noqa: BLE001 # a prefetch is an optimisation, never a reason to fail the request
verbose_router_logger.debug(
"routing read prefetch not armed for %s: %s", _without_line_breaks(model), _without_line_breaks(e)
)
def _get_routing_context(
self, model: str, request_kwargs: dict | None = None
) -> tuple[str | None, RouterStrategySelector | None]:
@ -1771,6 +1813,7 @@ 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
@ -1796,6 +1839,14 @@ 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,
@ -8730,7 +8781,7 @@ class Router:
return
if any(
model_info.get(field) is not None
for field in ("input_cost_per_token", "input_cost_per_second", "tiered_pricing")
for field in ("input_cost_per_token", "input_cost_per_second", "cost_per_second", "tiered_pricing")
):
return
try:
@ -12836,6 +12887,7 @@ 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.
@ -12888,8 +12940,14 @@ class Router:
health_check_probe=health_check_probe,
)
cooldown_deployments: Final = await _async_get_cooldown_deployments(
litellm_router_instance=self, parent_otel_span=parent_otel_span
cooldown_deployments: Final = (
await _async_get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=parent_otel_span)
if routing_read_batch is None
else await routing_read_batch.async_get_cooldown_deployments(
litellm_router_instance=self,
healthy_deployments=healthy_deployments,
parent_otel_span=parent_otel_span,
)
)
if verbose_router_logger.isEnabledFor(logging.DEBUG):
verbose_router_logger.debug("cooldown deployments: %s", cooldown_deployments)
@ -13167,6 +13225,7 @@ class Router:
# the hook can replace `model` and routing-group lookup must key
# off the final model name.
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,
@ -13175,6 +13234,7 @@ class Router:
input=input,
specific_deployment=specific_deployment,
parent_otel_span=parent_otel_span,
routing_read_batch=routing_read_batch,
)
if isinstance(healthy_deployments, dict):
await self._async_override_selector_pre_call_check(
@ -13205,6 +13265,7 @@ class Router:
messages=messages,
input=input,
request_kwargs=request_kwargs,
prefetched_usage=routing_read_batch.prefetched_usage if routing_read_batch is not None else None,
)
if deployment is None:
exception: Final = await async_raise_no_deployment_exception(
@ -13607,8 +13668,6 @@ class Router:
self._stamp_or_clear_metadata_key(request_kwargs, "model_group", bound_model)
return bound_registered_model
if self._request_header(request_kwargs, "x-app") != "cli":
return registered_model_name
if self._select_pre_routing_strategy(registered_model_name, request_kwargs) is None:
return registered_model_name
await self._claude_code_session_router_cache.async_set_cache(

View file

@ -1,7 +1,8 @@
#### What this does ####
# identifies lowest tpm deployment
import random
from collections.abc import Sequence
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Final
import httpx
@ -31,6 +32,26 @@ class RoutingArgs(LiteLLMPydanticObjectBase):
ttl: int = 1 * 60 # 1min (RPM/TPM expire key)
@dataclass(frozen=True)
class PrefetchedUsage:
"""
tpm/rpm counter values another read of this request already fetched from the router cache.
`values` is None when that read failed, which is what `async_batch_get_cache` returns on failure.
"""
keys: frozenset[str]
values: Mapping[str, object] | None
def covers(self, keys: Sequence[str]) -> bool:
return self.keys.issuperset(keys)
def values_for(self, keys: Sequence[str]) -> list[object | None] | None:
if self.values is None:
return None
return [self.values.get(key) for key in keys]
class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
"""
Updated version of TPM/RPM Logging.
@ -283,7 +304,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
# update cache
parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs)
## TPM
await self.router_cache.async_increment_cache(
await self.router_cache.async_increment_cache_post_call(
key=tpm_key,
value=total_tokens,
ttl=self.routing_args.ttl,
@ -412,17 +433,35 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
else:
return None
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
async def async_get_available_deployments(
self,
model_group: str,
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
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`).
"""
# get list of potential deployments
verbose_router_logger.debug(
@ -431,28 +470,15 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
healthy_deployments,
)
dt: Final = get_utc_datetime()
current_minute: Final = dt.strftime("%H-%M")
tpm_keys: Final = []
rpm_keys: Final = []
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_key = f"{id}:{deployment_name}:tpm:{current_minute}"
rpm_key = f"{id}:{deployment_name}:rpm:{current_minute}"
tpm_keys.append(tpm_key)
rpm_keys.append(rpm_key)
tpm_keys, rpm_keys = self.usage_counter_keys(healthy_deployments)
combined_tpm_rpm_keys: Final = tpm_keys + rpm_keys
combined_tpm_rpm_values: Final = await self.router_cache.async_batch_get_cache(
keys=combined_tpm_rpm_keys
) # [1, 2, None, ..]
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:
combined_tpm_rpm_values = await self.router_cache.async_batch_get_cache(
keys=combined_tpm_rpm_keys
) # [1, 2, None, ..]
if combined_tpm_rpm_values is not None:
tpm_values = combined_tpm_rpm_values[: len(tpm_keys)]

View file

@ -4,7 +4,7 @@ Wrapper around router cache. Meant to handle model cooldown logic
import functools
import time
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final
from typing_extensions import TypedDict
@ -163,6 +163,12 @@ class CooldownCache:
keys: Final = [CooldownCache.get_cooldown_cache_key(model_id) for model_id in model_ids]
results: Final = await self.cooldown_store.async_batch_get_cache(keys=keys, parent_otel_span=parent_otel_span)
return self.active_cooldowns_from_results(model_ids, results)
def active_cooldowns_from_results(
self, model_ids: list[str], results: Sequence[object] | None
) -> list[tuple[str, CooldownCacheValue]]:
"""The cooldowns still active in a `cooldown_store` batch read of `get_cooldown_cache_key(model_id)` per id."""
active_cooldowns: Final[list[tuple[str, CooldownCacheValue]]] = []
if results is None or all(v is None for v in results):

View file

@ -0,0 +1,156 @@
"""
One Redis round trip for the reads a request needs before a deployment can be picked.
The cooldown filter (`CooldownCache`, its own `DualCache`) and usage-based selection
(`LowestTPMLoggingHandler_v2`, the router cache) each issue their own MGET because they live in
different objects. `RoutingReadBatch` fetches both key sets in one
`DualCache.async_batch_get_cache_shared` while the healthy deployments are being resolved and hands
the usage slice to the strategy, so selection does not read again.
"""
import itertools
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final
from litellm._logging import verbose_router_logger
from litellm.caching.dual_cache import DualCache
from litellm.caching.redis_batch import BatchResult, active_request_redis_batches
from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2, PrefetchedUsage
from litellm.router_utils.cooldown_cache import CooldownCache
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
from litellm.router import Router as _Router
LitellmRouter = _Router
Span = _Span
else:
LitellmRouter = Any
Span = Any
_PREFETCH_SLOT: Final = "routing_read"
@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]
result: BatchResult[Mapping[str, object]]
@staticmethod
def arm(
litellm_router_instance: LitellmRouter,
usage_selector: LowestTPMLoggingHandler_v2 | None,
deployments: list,
) -> None:
request: Final = active_request_redis_batches()
redis_cache: Final = litellm_router_instance.cache.redis_cache
if request is None or redis_cache is None or _PREFETCH_SLOT in request.prefetched:
return
cooldown_keys: Final = tuple(
CooldownCache.get_cooldown_cache_key(model_id) for model_id in litellm_router_instance.get_model_ids()
)
usage_keys: Final = (
() 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)
)
@staticmethod
def armed() -> bool:
request: Final = active_request_redis_batches()
return request is not None and _PREFETCH_SLOT in request.prefetched
@staticmethod
def take(needed: Sequence[str]) -> "RoutingPrefetch | None":
"""The armed prefetch when it covers every key this read needs; taken once, so a retry reads fresh."""
request: Final = active_request_redis_batches()
if request is None:
return None
armed: Final = request.prefetched.pop(_PREFETCH_SLOT, None)
if isinstance(armed, RoutingPrefetch) and armed.keys.issuperset(needed):
return armed
return 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
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
cooldown state, and only through this batch when the request armed a prefetch for it. Otherwise the
router's plain cooldown read stays in charge."""
if strategy == "usage-based-routing-v2" and isinstance(selector, LowestTPMLoggingHandler_v2):
return RoutingReadBatch(usage_selector=selector)
return RoutingReadBatch(usage_selector=None) if RoutingPrefetch.armed() else None
async def async_get_cooldown_deployments(
self,
litellm_router_instance: LitellmRouter,
healthy_deployments: list,
parent_otel_span: Span | None,
) -> list[str]:
"""
`_async_get_cooldown_deployments`, with the strategy's tpm/rpm counters for
`healthy_deployments` fetched in the same MGET and kept as `prefetched_usage`.
"""
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))
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:
usage_values: Final = results[1]
self.prefetched_usage = PrefetchedUsage(
keys=frozenset(usage_keys),
values=None if usage_values is None else MappingProxyType(dict(zip(usage_keys, usage_values))),
)
cooldown_models: Final = litellm_router_instance.cooldown_cache.active_cooldowns_from_results(
model_ids, cooldown_results
)
verbose_router_logger.debug("retrieve cooldown models: %s", cooldown_models)
return [model_id for model_id, _ in cooldown_models]
@staticmethod
async def _read_prefetched(
reads: list[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."""
prefetch: Final = RoutingPrefetch.take(tuple(itertools.chain.from_iterable(keys for _, keys in reads)))
if prefetch is None:
return None
try:
values: Final = await prefetch.result
except Exception as e: # noqa: BLE001 # the shared read below applies the caches' own Redis fallback
verbose_router_logger.debug("routing prefetch failed, reading again: %s", e)
return None
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
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
}
results.append(await cache._apply_batch_get(pending, missed)) # pyright: ignore[reportPrivateUsage] # same two-step read as async_batch_get_cache_shared
return results

View file

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

View file

@ -224,19 +224,24 @@ class AutoRouterBenchmarkTotals(BaseModel):
"subtotal recording, and zero for an empty window"
)
savings_estimated_turns: int = Field(
description="Turns covered by the current savings estimator; legacy estimates are excluded"
description="Requests with a matching savings comparison, including historical recorded estimates"
)
savings_estimated_actual_spend: float = Field(
description="Actual spend, including classifier cost, for covered turns only"
)
savings_estimated_classifier_cost: float | None = Field(
default=None,
description="Classifier cost included in the matching historical and newer savings comparison; "
"null when classification costs for those requests are unavailable",
)
saved_spend: float | None = Field(
description="Signed savings for covered turns only; null when traffic has no current estimates"
description="Recorded historical savings plus newer estimates; null when traffic has no recorded savings estimates"
)
baseline_spend: float | None = Field(description="Estimated single-model cost for covered turns only")
saved_pct: float | None = Field(description="Covered savings over covered baseline spend, as a percentage")
saved_per_session: float | None = Field(
description="Average session savings; unavailable unless every turn is covered"
saved_pct: float | None = Field(
description="Total recorded savings over the matching historical and current baseline; null when costs are unavailable"
)
saved_per_session: float | None = Field(description="Recorded savings per session, including historical estimates")
cache: AutoRouterCacheStats
@ -268,12 +273,14 @@ class AutoRouterSessionResponse(BaseModel):
last_model: str = Field(description="The deployment model the most recent turn was routed to")
spend: float = Field(description="What the session's routed traffic actually cost, classifier calls included")
savings_estimated_turns: int = Field(
description="Turns covered by the current savings estimator; legacy estimates are excluded"
description="Requests with a matching savings comparison, including historical recorded estimates"
)
savings_estimated_actual_spend: float = Field(
description="Actual spend, including classifier cost, for covered turns only"
)
saved_spend: float | None = Field(description="Estimated savings for covered turns only, net of classifier cost")
saved_spend: float | None = Field(
description="Recorded historical savings plus newer estimates, net of classifier cost"
)
baseline_spend: float | None = Field(
description="Estimated single-model cost; unavailable unless every turn is covered"
)
@ -281,14 +288,14 @@ class AutoRouterSessionResponse(BaseModel):
description="Estimated single-model cost for covered turns only"
)
baseline_model: str | None = Field(
description="The savings baseline most covered turns were priced against, recorded turn by "
description="The savings baseline recorded by most session turns, including historical turns, recorded turn by "
"turn, so it still names the counterfactual after the router is reconfigured or removed. None when no "
"turn recorded one: rows from before the baseline was recorded, and adaptive and quality routers, "
"which derive no baseline and so report no savings"
)
baseline_models: Mapping[str, int] = Field(
description="Covered turns priced against each baseline model; more than one entry means the router's "
"baseline changed mid-session and baseline_spend mixes both"
description="Session turns recording each baseline model; more than one entry means the router's "
"baseline changed mid-session; these counts do not imply savings coverage"
)

View file

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

View file

@ -606,6 +606,7 @@ class LiteLLMParamsTypedDict(TypedDict, total=False):
## CUSTOM PRICING ##
input_cost_per_token: float | None
output_cost_per_token: float | None
cost_per_second: ReadOnly[float | None]
input_cost_per_second: float | None
output_cost_per_second: float | None
output_cost_per_second_480p: ReadOnly[float | None]

View file

@ -329,6 +329,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
input_cost_per_video_per_second: float | None # only for vertex ai models
input_cost_per_audio_token_batches: ReadOnly[float | None]
input_cost_per_image_token_batches: ReadOnly[float | None]
cost_per_second: ReadOnly[float | None]
input_cost_per_second: float | None # for OpenAI Speech models
input_cost_per_token_batches: float | None
input_cost_per_video_token_batches: ReadOnly[float | None]
@ -2784,6 +2785,7 @@ class LoggedLiteLLMParams(TypedDict, total=False):
acompletion: bool | None
preset_cache_key: str | None
no_log: bool | None
cost_per_second: ReadOnly[float | None]
input_cost_per_second: float | None
input_cost_per_token: float | None
output_cost_per_token: float | None
@ -3709,6 +3711,7 @@ class MirroredPricingParams(BaseModel):
class CustomPricingLiteLLMParams(MirroredPricingParams):
## CUSTOM PRICING ##
cost_per_second: float | None = None
input_cost_per_second: float | None = None
output_cost_per_second: float | None = None
output_cost_per_second_1080p: float | None = None

View file

@ -2309,6 +2309,7 @@ def _is_async_request(
or kwargs.get("_arealtime", False) is True
or kwargs.get("acreate_batch", False) is True
or kwargs.get("acreate_fine_tuning_job", False) is True
or kwargs.get("aresponses", False) is True
or is_pass_through is True
):
return True
@ -6168,6 +6169,7 @@ def _get_model_info_helper(
),
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),
input_cost_per_second=_model_info.get("input_cost_per_second", None),
input_cost_per_audio_token=_model_info.get("input_cost_per_audio_token", None),
input_cost_per_image_token=_model_info.get("input_cost_per_image_token", None),

File diff suppressed because it is too large Load diff

View file

@ -133,6 +133,11 @@
"minimum": 0,
"description": "Priority service-tier rate for the same-named base field."
},
"cache_creation_input_token_cost_above_272k_tokens_ultrafast": {
"type": "number",
"minimum": 0,
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
},
"cache_creation_input_token_cost_above_32k_tokens": {
"type": "number",
"minimum": 0,
@ -152,6 +157,10 @@
"minimum": 0,
"description": "Priority service-tier rate for the same-named base field."
},
"cache_creation_input_token_cost_ultrafast": {
"type": "number",
"minimum": 0
},
"cache_read_input_audio_token_cost": {
"type": "number",
"minimum": 0
@ -210,6 +219,11 @@
"minimum": 0,
"description": "Priority service-tier rate for the same-named base field."
},
"cache_read_input_token_cost_above_272k_tokens_ultrafast": {
"type": "number",
"minimum": 0,
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
},
"cache_read_input_token_cost_above_32k_tokens": {
"type": "number",
"minimum": 0,
@ -238,6 +252,10 @@
"minimum": 0,
"description": "Priority service-tier rate for the same-named base field."
},
"cache_read_input_token_cost_ultrafast": {
"type": "number",
"minimum": 0
},
"citation_cost_per_token": {
"type": "number",
"minimum": 0
@ -249,6 +267,10 @@
"comment": {
"type": "string"
},
"cost_per_second": {
"type": "number",
"minimum": 0
},
"default_reasoning_effort": {
"type": "string",
"description": "Reasoning effort the provider applies when the request omits reasoning_effort. Gates whether a non-default temperature or the top_p/logprobs sampling params are accepted, which hold only when the effort resolves to 'none'.",
@ -400,6 +422,11 @@
"minimum": 0,
"description": "Priority service-tier rate for the same-named base field."
},
"input_cost_per_token_above_272k_tokens_ultrafast": {
"type": "number",
"minimum": 0,
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
},
"input_cost_per_token_above_32k_tokens": {
"type": "number",
"minimum": 0,
@ -433,6 +460,10 @@
"minimum": 0,
"description": "Priority service-tier rate for the same-named base field."
},
"input_cost_per_token_ultrafast": {
"type": "number",
"minimum": 0
},
"input_cost_per_video_per_second": {
"type": "number",
"minimum": 0
@ -766,6 +797,11 @@
"minimum": 0,
"description": "Priority service-tier rate for the same-named base field."
},
"output_cost_per_token_above_272k_tokens_ultrafast": {
"type": "number",
"minimum": 0,
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
},
"output_cost_per_token_above_32k_tokens": {
"type": "number",
"minimum": 0,
@ -795,6 +831,10 @@
"minimum": 0,
"description": "Priority service-tier rate for the same-named base field."
},
"output_cost_per_token_ultrafast": {
"type": "number",
"minimum": 0
},
"output_cost_per_video_per_second": {
"type": "number",
"minimum": 0

View file

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

View file

@ -31,7 +31,7 @@ model_list:
- model_name: sagemaker-completion-model
litellm_params:
model: sagemaker/berri-benchmarking-Llama-2-70b-chat-hf-4
input_cost_per_second: 0.000420
cost_per_second: 0.000420
- model_name: text-embedding-ada-002
litellm_params:
model: openai/text-embedding-3-small

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