mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
commit
fe169152a9
204 changed files with 15725 additions and 1585 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-enterprise"
|
||||
version = "0.1.71"
|
||||
version = "0.1.72"
|
||||
description = "Package for LiteLLM Enterprise features"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
|
|
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
|
|||
module-root = ""
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.1.71"
|
||||
version = "0.1.72"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-enterprise==",
|
||||
|
|
|
|||
|
|
@ -0,0 +1,97 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN IF NOT EXISTS "enabled" BOOLEAN NOT NULL DEFAULT true,
|
||||
ADD COLUMN IF NOT EXISTS "execution_mode" TEXT NOT NULL DEFAULT 'autonomous',
|
||||
ADD COLUMN IF NOT EXISTS "identity_managed" BOOLEAN NOT NULL DEFAULT false;
|
||||
|
||||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_SpendLogs" ADD COLUMN IF NOT EXISTS "billing_agent_id" TEXT;
|
||||
|
||||
-- CreateTable
|
||||
CREATE TABLE IF NOT EXISTS "LiteLLM_AgentIdentity" (
|
||||
"agent_id" TEXT NOT NULL,
|
||||
"active" BOOLEAN NOT NULL DEFAULT true,
|
||||
"provider" TEXT NOT NULL,
|
||||
"issuer" TEXT NOT NULL,
|
||||
"tenant_id" TEXT NOT NULL,
|
||||
"client_id" TEXT NOT NULL,
|
||||
"service_principal_id" TEXT,
|
||||
"required_roles" TEXT[] DEFAULT ARRAY[]::TEXT[],
|
||||
"required_scopes" TEXT[] DEFAULT ARRAY['user_impersonation']::TEXT[],
|
||||
"revision" TEXT NOT NULL,
|
||||
"last_authenticated_at" TIMESTAMP(3),
|
||||
|
||||
CONSTRAINT "LiteLLM_AgentIdentity_pkey" PRIMARY KEY ("agent_id")
|
||||
);
|
||||
|
||||
-- CreateTable
|
||||
CREATE TABLE IF NOT EXISTS "LiteLLM_RetiredAgentIdentity" (
|
||||
"binding_id" TEXT NOT NULL,
|
||||
"agent_id" TEXT,
|
||||
"provider" TEXT NOT NULL,
|
||||
"issuer" TEXT NOT NULL,
|
||||
"tenant_id" TEXT NOT NULL,
|
||||
"client_id" TEXT NOT NULL,
|
||||
|
||||
CONSTRAINT "LiteLLM_RetiredAgentIdentity_pkey" PRIMARY KEY ("binding_id")
|
||||
);
|
||||
|
||||
-- CreateTable
|
||||
CREATE TABLE IF NOT EXISTS "LiteLLM_RetiredAgent" (
|
||||
"original_agent_id" TEXT NOT NULL,
|
||||
"retired_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
|
||||
CONSTRAINT "LiteLLM_RetiredAgent_pkey" PRIMARY KEY ("original_agent_id")
|
||||
);
|
||||
|
||||
-- CreateTable
|
||||
CREATE TABLE IF NOT EXISTS "LiteLLM_VerifiedSubject" (
|
||||
"subject_id" TEXT NOT NULL,
|
||||
"issuer" TEXT NOT NULL,
|
||||
"tenant_id" TEXT NOT NULL,
|
||||
"oid" TEXT NOT NULL,
|
||||
"kind" TEXT NOT NULL DEFAULT 'human',
|
||||
"user_id" TEXT,
|
||||
"verified_via" TEXT NOT NULL DEFAULT 'sso_interactive',
|
||||
"verified_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
|
||||
CONSTRAINT "LiteLLM_VerifiedSubject_pkey" PRIMARY KEY ("subject_id")
|
||||
);
|
||||
|
||||
-- CreateIndex
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_AgentIdentity_provider_tenant_id_client_id_key" ON "LiteLLM_AgentIdentity"("provider", "tenant_id", "client_id");
|
||||
|
||||
-- CreateIndex
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_AgentIdentity_issuer_service_principal_id_key" ON "LiteLLM_AgentIdentity"("issuer", "service_principal_id");
|
||||
|
||||
-- CreateIndex
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_RetiredAgentIdentity_provider_tenant_id_client_id_key" ON "LiteLLM_RetiredAgentIdentity"("provider", "tenant_id", "client_id");
|
||||
|
||||
-- CreateIndex
|
||||
CREATE INDEX IF NOT EXISTS "LiteLLM_VerifiedSubject_user_id_idx" ON "LiteLLM_VerifiedSubject"("user_id");
|
||||
|
||||
-- CreateIndex
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_VerifiedSubject_issuer_tenant_id_oid_key" ON "LiteLLM_VerifiedSubject"("issuer", "tenant_id", "oid");
|
||||
|
||||
-- AddForeignKey
|
||||
DO $$
|
||||
BEGIN
|
||||
IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_AgentIdentity_agent_id_fkey') THEN
|
||||
ALTER TABLE "LiteLLM_AgentIdentity" ADD CONSTRAINT "LiteLLM_AgentIdentity_agent_id_fkey" FOREIGN KEY ("agent_id") REFERENCES "LiteLLM_AgentsTable"("agent_id") ON DELETE CASCADE ON UPDATE CASCADE;
|
||||
END IF;
|
||||
END $$;
|
||||
|
||||
-- AddForeignKey
|
||||
DO $$
|
||||
BEGIN
|
||||
IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_RetiredAgentIdentity_agent_id_fkey') THEN
|
||||
ALTER TABLE "LiteLLM_RetiredAgentIdentity" ADD CONSTRAINT "LiteLLM_RetiredAgentIdentity_agent_id_fkey" FOREIGN KEY ("agent_id") REFERENCES "LiteLLM_AgentsTable"("agent_id") ON DELETE SET NULL ON UPDATE CASCADE;
|
||||
END IF;
|
||||
END $$;
|
||||
|
||||
-- AddForeignKey
|
||||
DO $$
|
||||
BEGIN
|
||||
IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_VerifiedSubject_user_id_fkey') THEN
|
||||
ALTER TABLE "LiteLLM_VerifiedSubject" ADD CONSTRAINT "LiteLLM_VerifiedSubject_user_id_fkey" FOREIGN KEY ("user_id") REFERENCES "LiteLLM_UserTable"("user_id") ON DELETE CASCADE ON UPDATE CASCADE;
|
||||
END IF;
|
||||
END $$;
|
||||
|
|
@ -78,6 +78,11 @@ model LiteLLM_AgentsTable {
|
|||
object_permission_id String?
|
||||
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
|
||||
spend Float @default(0.0)
|
||||
identity_managed Boolean @default(false)
|
||||
enabled Boolean @default(true)
|
||||
execution_mode String @default("autonomous")
|
||||
identity LiteLLM_AgentIdentity?
|
||||
retired_identities LiteLLM_RetiredAgentIdentity[]
|
||||
tpm_limit Int?
|
||||
rpm_limit Int?
|
||||
session_tpm_limit Int?
|
||||
|
|
@ -88,6 +93,56 @@ model LiteLLM_AgentsTable {
|
|||
updated_by String
|
||||
}
|
||||
|
||||
model LiteLLM_AgentIdentity {
|
||||
agent_id String @id
|
||||
active Boolean @default(true)
|
||||
agent LiteLLM_AgentsTable @relation(fields: [agent_id], references: [agent_id], onDelete: Cascade)
|
||||
provider String
|
||||
issuer String
|
||||
tenant_id String
|
||||
client_id String
|
||||
service_principal_id String?
|
||||
required_roles String[] @default([])
|
||||
required_scopes String[] @default(["user_impersonation"])
|
||||
revision String @default(uuid())
|
||||
last_authenticated_at DateTime?
|
||||
@@unique([provider, tenant_id, client_id])
|
||||
@@unique([issuer, service_principal_id])
|
||||
}
|
||||
|
||||
model LiteLLM_RetiredAgentIdentity {
|
||||
binding_id String @id @default(uuid())
|
||||
agent_id String?
|
||||
agent LiteLLM_AgentsTable? @relation(fields: [agent_id], references: [agent_id], onDelete: SetNull)
|
||||
provider String
|
||||
issuer String
|
||||
tenant_id String
|
||||
client_id String
|
||||
@@unique([provider, tenant_id, client_id])
|
||||
}
|
||||
|
||||
model LiteLLM_RetiredAgent {
|
||||
original_agent_id String @id
|
||||
retired_at DateTime @default(now())
|
||||
}
|
||||
|
||||
model LiteLLM_VerifiedSubject {
|
||||
subject_id String @id @default(uuid())
|
||||
issuer String
|
||||
tenant_id String
|
||||
oid String
|
||||
kind String @default("human")
|
||||
user_id String?
|
||||
user LiteLLM_UserTable? @relation(fields: [user_id], references: [user_id], onDelete: Cascade)
|
||||
verified_via String @default("sso_interactive")
|
||||
verified_at DateTime @default(now())
|
||||
@@unique([issuer, tenant_id, oid])
|
||||
@@index([user_id])
|
||||
}
|
||||
|
||||
|
||||
|
||||
|
||||
model LiteLLM_OrganizationTable {
|
||||
organization_id String @id @default(uuid())
|
||||
organization_alias String
|
||||
|
|
@ -241,6 +296,7 @@ model LiteLLM_DeletedTeamTable {
|
|||
|
||||
// Track spend, rate limit, budget Users
|
||||
model LiteLLM_UserTable {
|
||||
verified_subjects LiteLLM_VerifiedSubject[]
|
||||
user_id String @id
|
||||
user_alias String?
|
||||
team_id String?
|
||||
|
|
@ -675,6 +731,7 @@ model LiteLLM_SpendLogs {
|
|||
session_id String?
|
||||
status String?
|
||||
mcp_namespaced_tool_name String?
|
||||
billing_agent_id String?
|
||||
agent_id String?
|
||||
proxy_server_request Json? @default("{}")
|
||||
litellm_call_id String?
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.102"
|
||||
version = "0.4.103"
|
||||
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
|
|
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
|
|||
module-root = ""
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.4.102"
|
||||
version = "0.4.103"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-proxy-extras==",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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())),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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> {
|
||||
|
|
|
|||
|
|
@ -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> {
|
||||
|
|
|
|||
|
|
@ -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> {
|
||||
|
|
|
|||
48
litellm-rust/crates/core/src/context.rs
Normal file
48
litellm-rust/crates/core/src/context.rs
Normal 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)
|
||||
}
|
||||
}
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
},
|
||||
)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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> {
|
||||
|
|
|
|||
|
|
@ -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> {
|
||||
|
|
|
|||
|
|
@ -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>,
|
||||
|
|
|
|||
|
|
@ -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) {
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -69,7 +69,7 @@ async fn handle(
|
|||
extra_headers: None,
|
||||
timeout: deployment.timeout,
|
||||
},
|
||||
cache_options,
|
||||
cache_options.policy,
|
||||
),
|
||||
(),
|
||||
headers.clone(),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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!(
|
||||
|
|
|
|||
|
|
@ -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")]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
},
|
||||
))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
555
litellm/caching/redis_batch.py
Normal file
555
litellm/caching/redis_batch.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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": [
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
||||
|
||||
|
|
|
|||
17
litellm/proxy/agent_endpoints/identity.py
Normal file
17
litellm/proxy/agent_endpoints/identity.py
Normal file
|
|
@ -0,0 +1,17 @@
|
|||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
LEGACY_IDENTITY_MESSAGE: Final = (
|
||||
"litellm_params.identity is not supported: bind an Entra application through the top-level identity field"
|
||||
)
|
||||
|
||||
|
||||
def has_legacy_identity(params: Mapping[str, object] | None) -> bool:
|
||||
return params is not None and "identity" in params
|
||||
|
||||
|
||||
def reject_legacy_identity(params: Mapping[str, object] | None) -> None:
|
||||
if has_legacy_identity(params):
|
||||
raise HTTPException(400, LEGACY_IDENTITY_MESSAGE)
|
||||
250
litellm/proxy/agent_endpoints/identity_store.py
Normal file
250
litellm/proxy/agent_endpoints/identity_store.py
Normal 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
|
||||
220
litellm/proxy/agent_endpoints/managed_identity.py
Normal file
220
litellm/proxy/agent_endpoints/managed_identity.py
Normal file
|
|
@ -0,0 +1,220 @@
|
|||
from collections.abc import Mapping
|
||||
from datetime import datetime
|
||||
from typing import Final, NoReturn, TypedDict
|
||||
from uuid import uuid4
|
||||
|
||||
from fastapi import HTTPException
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import (
|
||||
AgentExecutionMode,
|
||||
AgentIdentityBinding,
|
||||
AgentIdentityFailure,
|
||||
AgentSubject,
|
||||
EntraIdentityConfig,
|
||||
)
|
||||
|
||||
_MODE: Final = TypeAdapter(AgentExecutionMode)
|
||||
|
||||
|
||||
class IdentityFields(TypedDict, total=False):
|
||||
provider: ReadOnly[str]
|
||||
tenant_id: ReadOnly[str]
|
||||
client_id: ReadOnly[str]
|
||||
issuer: ReadOnly[str]
|
||||
service_principal_id: ReadOnly[str | None]
|
||||
required_roles: ReadOnly[tuple[str, ...]]
|
||||
required_scopes: ReadOnly[tuple[str, ...]]
|
||||
active: ReadOnly[bool]
|
||||
revision: ReadOnly[str]
|
||||
last_authenticated_at: ReadOnly[datetime | None]
|
||||
|
||||
|
||||
class IdentityUpsert(TypedDict):
|
||||
create: ReadOnly[IdentityFields]
|
||||
update: ReadOnly[IdentityFields]
|
||||
|
||||
|
||||
class IdentityRelationWrite(TypedDict, total=False):
|
||||
create: ReadOnly[IdentityFields]
|
||||
update: ReadOnly[IdentityFields]
|
||||
upsert: ReadOnly[IdentityUpsert]
|
||||
|
||||
|
||||
class IdentityHistoryKey(TypedDict):
|
||||
provider: ReadOnly[str]
|
||||
tenant_id: ReadOnly[str]
|
||||
client_id: ReadOnly[str]
|
||||
|
||||
|
||||
class IdentityHistoryWhere(TypedDict):
|
||||
provider_tenant_id_client_id: ReadOnly[IdentityHistoryKey]
|
||||
|
||||
|
||||
class IdentityHistoryEntry(IdentityHistoryKey):
|
||||
issuer: ReadOnly[str]
|
||||
|
||||
|
||||
class IdentityHistoryConnect(TypedDict):
|
||||
where: ReadOnly[IdentityHistoryWhere]
|
||||
create: ReadOnly[IdentityHistoryEntry]
|
||||
|
||||
|
||||
class IdentityHistoryWrite(TypedDict):
|
||||
connectOrCreate: ReadOnly[IdentityHistoryConnect]
|
||||
|
||||
|
||||
class ManagedWriteFields(TypedDict, total=False):
|
||||
enabled: ReadOnly[bool]
|
||||
execution_mode: ReadOnly[AgentExecutionMode]
|
||||
identity_managed: ReadOnly[bool]
|
||||
identity: ReadOnly[IdentityRelationWrite]
|
||||
retired_identities: ReadOnly[IdentityHistoryWrite]
|
||||
|
||||
|
||||
def raise_identity_failure(failure: AgentIdentityFailure, status_code: int = 403) -> NoReturn:
|
||||
raise HTTPException(503 if failure.code == "policy_unavailable" else status_code, failure.message)
|
||||
|
||||
|
||||
def _configuration_failure(
|
||||
identity: EntraIdentityConfig | AgentIdentityBinding | None,
|
||||
mode: AgentExecutionMode,
|
||||
enabling_without_binding: bool,
|
||||
) -> AgentIdentityFailure | None:
|
||||
if identity is not None and mode != "delegated" and not identity.service_principal_id:
|
||||
return AgentIdentityFailure(
|
||||
message="Autonomous mode requires the Enterprise application service-principal object ID"
|
||||
)
|
||||
if enabling_without_binding and (
|
||||
identity is None or isinstance(identity, AgentIdentityBinding) and not identity.active
|
||||
):
|
||||
return AgentIdentityFailure(message="Bind an identity before enabling this managed agent")
|
||||
return None
|
||||
|
||||
|
||||
def managed_write_fields(
|
||||
incoming: Mapping[str, object],
|
||||
existing: AgentResponse | None,
|
||||
updated_by: str,
|
||||
) -> ManagedWriteFields | AgentIdentityFailure:
|
||||
try:
|
||||
identity: Final = (
|
||||
EntraIdentityConfig.model_validate(incoming["identity"]) if incoming.get("identity") is not None else None
|
||||
)
|
||||
mode: Final = _MODE.validate_python(
|
||||
incoming.get("execution_mode", existing.execution_mode if existing else "autonomous")
|
||||
)
|
||||
current_identity: Final = identity if "identity" in incoming else existing.identity if existing else None
|
||||
failure: Final = _configuration_failure(
|
||||
current_identity,
|
||||
mode,
|
||||
incoming.get("enabled") is True
|
||||
and "identity" not in incoming
|
||||
and bool(existing and existing.identity_managed),
|
||||
)
|
||||
if failure is not None:
|
||||
return failure
|
||||
empty: Final[ManagedWriteFields] = {}
|
||||
identity_fields: Final = _identity_write(identity, existing) if "identity" in incoming else empty
|
||||
result: Final[ManagedWriteFields] = {
|
||||
**({"enabled": incoming["enabled"] is True} if "enabled" in incoming else {}),
|
||||
**({"execution_mode": mode} if "execution_mode" in incoming else {}),
|
||||
**identity_fields,
|
||||
}
|
||||
return result
|
||||
except (ValidationError, ValueError) as exc:
|
||||
return AgentIdentityFailure(message=f"Invalid agent identity configuration: {exc}")
|
||||
|
||||
|
||||
def _identity_write(identity: EntraIdentityConfig | None, existing: AgentResponse | None) -> ManagedWriteFields:
|
||||
if identity is None:
|
||||
unbind: Final[ManagedWriteFields] = {
|
||||
**(
|
||||
{"identity": {"update": {"active": False, "revision": str(uuid4()), "last_authenticated_at": None}}}
|
||||
if existing and existing.identity
|
||||
else {}
|
||||
),
|
||||
**({"identity_managed": True, "enabled": False} if existing and existing.identity_managed else {}),
|
||||
}
|
||||
return unbind
|
||||
if (
|
||||
existing
|
||||
and existing.identity
|
||||
and existing.identity.active
|
||||
and all(getattr(existing.identity, name) == value for name, value in identity.model_dump().items())
|
||||
):
|
||||
unchanged: Final[ManagedWriteFields] = {}
|
||||
return unchanged
|
||||
binding: Final[IdentityFields] = {
|
||||
"provider": identity.provider,
|
||||
"tenant_id": identity.tenant_id,
|
||||
"client_id": identity.client_id,
|
||||
"service_principal_id": identity.service_principal_id,
|
||||
"required_roles": identity.required_roles,
|
||||
"required_scopes": identity.required_scopes,
|
||||
"issuer": identity.issuer,
|
||||
"active": True,
|
||||
"revision": str(uuid4()),
|
||||
"last_authenticated_at": None,
|
||||
}
|
||||
result: Final[ManagedWriteFields] = {
|
||||
"retired_identities": {
|
||||
"connectOrCreate": {
|
||||
"where": {
|
||||
"provider_tenant_id_client_id": {
|
||||
"provider": identity.provider,
|
||||
"tenant_id": identity.tenant_id,
|
||||
"client_id": identity.client_id,
|
||||
}
|
||||
},
|
||||
"create": {
|
||||
"provider": identity.provider,
|
||||
"issuer": identity.issuer,
|
||||
"tenant_id": identity.tenant_id,
|
||||
"client_id": identity.client_id,
|
||||
},
|
||||
}
|
||||
},
|
||||
"identity_managed": True,
|
||||
"identity": {"upsert": {"create": binding, "update": binding}} if existing else {"create": binding},
|
||||
}
|
||||
return result
|
||||
|
||||
|
||||
def classify_agent_subject(
|
||||
binding: AgentIdentityBinding,
|
||||
claims: Mapping[str, object],
|
||||
allowed_mode: AgentExecutionMode,
|
||||
) -> AgentSubject | AgentIdentityFailure:
|
||||
if (claims.get("iss"), claims.get("tid"), claims.get("azp")) != (
|
||||
binding.issuer,
|
||||
binding.tenant_id,
|
||||
binding.client_id,
|
||||
):
|
||||
return AgentIdentityFailure(message="Token does not match the registered Entra application")
|
||||
oid: Final = claims.get("oid")
|
||||
if not isinstance(oid, str) or not oid:
|
||||
return AgentIdentityFailure(message="Entra token must identify its object subject")
|
||||
scope: Final = claims.get("scp")
|
||||
facets: Final = claims.get("xms_sub_fct")
|
||||
if facets is not None and (not isinstance(facets, str) or "13" in facets.split()):
|
||||
return AgentIdentityFailure(message="Native agent-user authentication is not supported by this binding")
|
||||
if scope is not None and not isinstance(scope, str):
|
||||
return AgentIdentityFailure(message="Invalid delegated scope claim")
|
||||
if isinstance(scope, str) and scope:
|
||||
if allowed_mode == "autonomous" or oid == binding.service_principal_id or claims.get("idtyp") == "app":
|
||||
return AgentIdentityFailure(message="Delegated token contradicts the configured agent identity or mode")
|
||||
granted_scopes: Final = frozenset(scope.split())
|
||||
if not granted_scopes or not frozenset(binding.required_scopes).issubset(granted_scopes):
|
||||
return AgentIdentityFailure(message="Token lacks the required delegated scopes")
|
||||
return AgentSubject(kind="delegated_subject", oid=oid, mode="delegated")
|
||||
if allowed_mode == "delegated" or oid != binding.service_principal_id or claims.get("idtyp") == "user":
|
||||
return AgentIdentityFailure(message="Application token contradicts the configured agent identity or mode")
|
||||
roles: Final = claims.get("roles", ())
|
||||
if not isinstance(roles, (list, tuple)) or any(not isinstance(role, str) for role in roles):
|
||||
return AgentIdentityFailure(message="Invalid application roles claim")
|
||||
if not frozenset(binding.required_roles).issubset(roles):
|
||||
return AgentIdentityFailure(message="Token lacks the required application roles")
|
||||
return AgentSubject(kind="application", oid=oid, mode="autonomous")
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -72,26 +72,55 @@ def _derive_key(signing_key: str) -> bytes:
|
|||
return hashlib.sha256(signing_key.encode()).digest()
|
||||
|
||||
|
||||
def _encrypt_aes_gcm(value: str, signing_key: str) -> str:
|
||||
"""Encrypt under AES-256-GCM and return the versioned ``v2:gcm:`` string."""
|
||||
def _seal_aes_gcm(value: str, signing_key: str, aad: bytes | None) -> bytes:
|
||||
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
|
||||
|
||||
nonce: Final = os.urandom(12)
|
||||
# AESGCM.encrypt returns ciphertext || tag(16); wire format is nonce || that.
|
||||
blob: Final = AESGCM(_derive_key(signing_key)).encrypt(nonce, value.encode("utf-8"), None)
|
||||
return _V2_GCM_PREFIX + base64.urlsafe_b64encode(nonce + blob).decode("utf-8")
|
||||
return nonce + AESGCM(_derive_key(signing_key)).encrypt(nonce, value.encode("utf-8"), aad)
|
||||
|
||||
|
||||
def _open_aes_gcm(sealed: bytes, signing_key: str, aad: bytes | None) -> str:
|
||||
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
|
||||
|
||||
# An empty plaintext still serializes to nonce(12) || tag(16) = 28 bytes, so a
|
||||
# short/empty buffer here is a corrupt value: let AESGCM.decrypt raise and be
|
||||
# swallowed by the caller (returns None/original), same as legacy.
|
||||
return AESGCM(_derive_key(signing_key)).decrypt(sealed[:12], sealed[12:], aad).decode("utf-8")
|
||||
|
||||
|
||||
def _encrypt_aes_gcm(value: str, signing_key: str) -> str:
|
||||
"""Encrypt under AES-256-GCM and return the versioned ``v2:gcm:`` string."""
|
||||
sealed: Final = _seal_aes_gcm(value=value, signing_key=signing_key, aad=None)
|
||||
return _V2_GCM_PREFIX + base64.urlsafe_b64encode(sealed).decode("utf-8")
|
||||
|
||||
|
||||
def _decrypt_aes_gcm(value: str, signing_key: str) -> str:
|
||||
"""Decrypt a versioned ``v2:gcm:`` string produced by :func:`_encrypt_aes_gcm`."""
|
||||
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
|
||||
sealed: Final = base64.urlsafe_b64decode(value[len(_V2_GCM_PREFIX) :])
|
||||
return _open_aes_gcm(sealed=sealed, signing_key=signing_key, aad=None)
|
||||
|
||||
raw: Final = base64.urlsafe_b64decode(value[len(_V2_GCM_PREFIX) :])
|
||||
# An empty plaintext still serializes to nonce(12) || tag(16) = 28 bytes, so a
|
||||
# short/empty buffer here is a corrupt value: let AESGCM.decrypt raise and be
|
||||
# swallowed by decrypt_value_helper (returns None/original), same as legacy.
|
||||
nonce, blob = raw[:12], raw[12:]
|
||||
return AESGCM(_derive_key(signing_key)).decrypt(nonce, blob, None).decode("utf-8")
|
||||
|
||||
def encrypt_bearer_token(value: str, prefix: str) -> str:
|
||||
"""AES-256-GCM as unpadded base64url behind ``prefix``, which is also the AAD so a token can't change kind."""
|
||||
salt_key: Final = _get_salt_key()
|
||||
if not isinstance(salt_key, str):
|
||||
raise ValueError("Set LITELLM_SALT_KEY or a master key to mint bearer tokens")
|
||||
sealed: Final = _seal_aes_gcm(value=value, signing_key=salt_key, aad=prefix.encode("utf-8"))
|
||||
return prefix + base64.urlsafe_b64encode(sealed).decode("ascii").rstrip("=")
|
||||
|
||||
|
||||
def decrypt_bearer_token(token: str, prefix: str) -> str | None:
|
||||
"""None unless ``token`` came from :func:`encrypt_bearer_token` with the same ``prefix``."""
|
||||
salt_key: Final = _get_salt_key()
|
||||
if not isinstance(salt_key, str) or not token.startswith(prefix):
|
||||
return None
|
||||
encoded: Final = token.removeprefix(prefix)
|
||||
try:
|
||||
sealed: Final = base64.b64decode(encoded + "=" * (-len(encoded) % 4), altchars=b"-_", validate=True)
|
||||
return _open_aes_gcm(sealed=sealed, signing_key=salt_key, aad=prefix.encode("utf-8"))
|
||||
except Exception: # noqa: BLE001 # base64 and AES-GCM each raise their own "not a token" type
|
||||
return None
|
||||
|
||||
|
||||
def encrypt_value_helper(value: str, new_encryption_key: str | None = None):
|
||||
|
|
|
|||
|
|
@ -17,6 +17,8 @@ from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec
|
|||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span
|
||||
|
||||
from litellm.caching.redis_batch import BatchResult
|
||||
|
||||
T = TypeVar("T", bound=BaseModel)
|
||||
|
||||
_HASHED_TOKEN_CACHE_KEY: Final = re.compile(r"[0-9a-f]{64}")
|
||||
|
|
@ -27,6 +29,9 @@ def is_user_key_cache_key(key: str) -> bool:
|
|||
return _HASHED_TOKEN_CACHE_KEY.fullmatch(key) is not None
|
||||
|
||||
|
||||
_PIPELINED_SET_OPTIONS: Final = frozenset(("ttl",))
|
||||
|
||||
|
||||
class UserApiKeyCache(DualCache):
|
||||
"""
|
||||
DualCache wrapper for UserAPIKeyAuth-like payloads.
|
||||
|
|
@ -208,10 +213,23 @@ class UserApiKeyCache(DualCache):
|
|||
return super().set_cache(key=key, value=payload, local_only=local_only, **kwargs)
|
||||
|
||||
async def async_set_cache(self, key: str | None, value: object, local_only: bool = False, **kwargs: object):
|
||||
"""Inside a request the Redis SET rides the request's pipeline (memory is written at once); anywhere
|
||||
else, or with options the pipeline does not carry, it goes to Redis directly as before."""
|
||||
model_type: Final = cast(type[BaseModel] | None, kwargs.pop("model_type", None))
|
||||
payload: Final[object] = CacheCodec.serialize(value, model_type=model_type)
|
||||
ttl: Final = kwargs.get("ttl")
|
||||
pipelined: Final = (
|
||||
key is not None
|
||||
and not local_only
|
||||
and kwargs.keys() <= _PIPELINED_SET_OPTIONS
|
||||
and (ttl is None or isinstance(ttl, (int, float)))
|
||||
)
|
||||
if key is not None and is_user_key_cache_key(key):
|
||||
if pipelined and await self.key_object_cache.async_set_cache_pre_call(key, payload, ttl) is not None:
|
||||
return None
|
||||
return await self.key_object_cache.async_set_cache(key=key, value=payload, local_only=local_only, **kwargs)
|
||||
if pipelined and await super().async_set_cache_pre_call(key, payload, ttl) is not None:
|
||||
return None
|
||||
return await super().async_set_cache(key=key, value=payload, local_only=local_only, **kwargs)
|
||||
|
||||
def delete_cache(self, key: str) -> None:
|
||||
|
|
@ -226,6 +244,11 @@ class UserApiKeyCache(DualCache):
|
|||
return
|
||||
await super().async_delete_cache(key)
|
||||
|
||||
async def async_delete_cache_pre_call(self, key: str) -> BatchResult[None] | None:
|
||||
if is_user_key_cache_key(key):
|
||||
return await self.key_object_cache.async_delete_cache_pre_call(key)
|
||||
return await super().async_delete_cache_pre_call(key)
|
||||
|
||||
async def async_delete_cache_keys(self, keys: Sequence[str]) -> None:
|
||||
"""Batch twin of ``async_delete_cache``, partitioned like
|
||||
``async_set_cache_pipeline``.
|
||||
|
|
|
|||
147
litellm/proxy/db/autorouter_savings_comparison.py
Normal file
147
litellm/proxy/db/autorouter_savings_comparison.py
Normal 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
|
||||
"""
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -2012,4 +2012,5 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
GuardrailEventHooks.logging_only,
|
||||
GuardrailEventHooks.pre_mcp_call,
|
||||
GuardrailEventHooks.during_mcp_call,
|
||||
GuardrailEventHooks.post_mcp_call,
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
25
litellm/proxy/middleware/redis_request_batch_middleware.py
Normal file
25
litellm/proxy/middleware/redis_request_batch_middleware.py
Normal 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()
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -78,6 +78,11 @@ model LiteLLM_AgentsTable {
|
|||
object_permission_id String?
|
||||
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
|
||||
spend Float @default(0.0)
|
||||
identity_managed Boolean @default(false)
|
||||
enabled Boolean @default(true)
|
||||
execution_mode String @default("autonomous")
|
||||
identity LiteLLM_AgentIdentity?
|
||||
retired_identities LiteLLM_RetiredAgentIdentity[]
|
||||
tpm_limit Int?
|
||||
rpm_limit Int?
|
||||
session_tpm_limit Int?
|
||||
|
|
@ -88,6 +93,56 @@ model LiteLLM_AgentsTable {
|
|||
updated_by String
|
||||
}
|
||||
|
||||
model LiteLLM_AgentIdentity {
|
||||
agent_id String @id
|
||||
active Boolean @default(true)
|
||||
agent LiteLLM_AgentsTable @relation(fields: [agent_id], references: [agent_id], onDelete: Cascade)
|
||||
provider String
|
||||
issuer String
|
||||
tenant_id String
|
||||
client_id String
|
||||
service_principal_id String?
|
||||
required_roles String[] @default([])
|
||||
required_scopes String[] @default(["user_impersonation"])
|
||||
revision String @default(uuid())
|
||||
last_authenticated_at DateTime?
|
||||
@@unique([provider, tenant_id, client_id])
|
||||
@@unique([issuer, service_principal_id])
|
||||
}
|
||||
|
||||
model LiteLLM_RetiredAgentIdentity {
|
||||
binding_id String @id @default(uuid())
|
||||
agent_id String?
|
||||
agent LiteLLM_AgentsTable? @relation(fields: [agent_id], references: [agent_id], onDelete: SetNull)
|
||||
provider String
|
||||
issuer String
|
||||
tenant_id String
|
||||
client_id String
|
||||
@@unique([provider, tenant_id, client_id])
|
||||
}
|
||||
|
||||
model LiteLLM_RetiredAgent {
|
||||
original_agent_id String @id
|
||||
retired_at DateTime @default(now())
|
||||
}
|
||||
|
||||
model LiteLLM_VerifiedSubject {
|
||||
subject_id String @id @default(uuid())
|
||||
issuer String
|
||||
tenant_id String
|
||||
oid String
|
||||
kind String @default("human")
|
||||
user_id String?
|
||||
user LiteLLM_UserTable? @relation(fields: [user_id], references: [user_id], onDelete: Cascade)
|
||||
verified_via String @default("sso_interactive")
|
||||
verified_at DateTime @default(now())
|
||||
@@unique([issuer, tenant_id, oid])
|
||||
@@index([user_id])
|
||||
}
|
||||
|
||||
|
||||
|
||||
|
||||
model LiteLLM_OrganizationTable {
|
||||
organization_id String @id @default(uuid())
|
||||
organization_alias String
|
||||
|
|
@ -241,6 +296,7 @@ model LiteLLM_DeletedTeamTable {
|
|||
|
||||
// Track spend, rate limit, budget Users
|
||||
model LiteLLM_UserTable {
|
||||
verified_subjects LiteLLM_VerifiedSubject[]
|
||||
user_id String @id
|
||||
user_alias String?
|
||||
team_id String?
|
||||
|
|
@ -675,6 +731,7 @@ model LiteLLM_SpendLogs {
|
|||
session_id String?
|
||||
status String?
|
||||
mcp_namespaced_tool_name String?
|
||||
billing_agent_id String?
|
||||
agent_id String?
|
||||
proxy_server_request Json? @default("{}")
|
||||
litellm_call_id String?
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ 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:
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
),
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
156
litellm/router_utils/routing_read_batch.py
Normal file
156
litellm/router_utils/routing_read_batch.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
96
litellm/types/proxy/agent_identity.py
Normal file
96
litellm/types/proxy/agent_identity.py
Normal file
|
|
@ -0,0 +1,96 @@
|
|||
from datetime import datetime
|
||||
from typing import Literal, TypeAlias
|
||||
from uuid import UUID
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator
|
||||
|
||||
AgentExecutionMode: TypeAlias = Literal["autonomous", "delegated", "both"]
|
||||
|
||||
|
||||
class EntraIdentityConfig(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="forbid")
|
||||
|
||||
provider: Literal["microsoft_entra"]
|
||||
tenant_id: str
|
||||
client_id: str
|
||||
service_principal_id: str | None = None
|
||||
required_roles: tuple[str, ...] = ()
|
||||
required_scopes: tuple[str, ...] = Field(
|
||||
default=("user_impersonation",),
|
||||
description="Required delegated scopes. An empty list accepts any nonempty scope granted for this gateway.",
|
||||
)
|
||||
|
||||
@field_validator("tenant_id", "client_id", "service_principal_id")
|
||||
@classmethod
|
||||
def normalize_identifier(cls, value: str | None) -> str | None:
|
||||
return str(UUID(value)) if value is not None else None
|
||||
|
||||
@property
|
||||
def issuer(self) -> str:
|
||||
return f"https://login.microsoftonline.com/{self.tenant_id}/v2.0"
|
||||
|
||||
|
||||
class AgentIdentityBinding(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
agent_id: str
|
||||
active: bool = True
|
||||
provider: Literal["microsoft_entra"]
|
||||
tenant_id: str
|
||||
client_id: str
|
||||
service_principal_id: str | None = None
|
||||
issuer: str
|
||||
required_roles: tuple[str, ...] = ()
|
||||
required_scopes: tuple[str, ...] = ("user_impersonation",)
|
||||
revision: str
|
||||
last_authenticated_at: datetime | None = None
|
||||
|
||||
|
||||
class AgentSubject(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
kind: Literal["application", "delegated_subject"]
|
||||
oid: str
|
||||
mode: Literal["autonomous", "delegated"]
|
||||
|
||||
|
||||
class AgentIdentityFailure(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
code: Literal["identity_denied", "policy_unavailable"] = "identity_denied"
|
||||
message: str
|
||||
|
||||
|
||||
class ManagedAgentContext(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
agent_id: str
|
||||
binding_revision: str | None = None
|
||||
mode: Literal["autonomous", "delegated"]
|
||||
user_id: str | None = None
|
||||
subject_oid: str | None = None
|
||||
|
||||
|
||||
class VerifiedHumanSubject(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
issuer: str
|
||||
tenant_id: str
|
||||
oid: str
|
||||
user_id: str
|
||||
|
||||
|
||||
class MicrosoftInteractiveSubject(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
issuer: str
|
||||
tenant_id: str
|
||||
oid: str
|
||||
|
||||
|
||||
class ManagedAgentIdentityStatus(BaseModel):
|
||||
identity: AgentIdentityBinding | None = None
|
||||
identity_managed: bool = False
|
||||
enabled: bool = True
|
||||
execution_mode: AgentExecutionMode = "autonomous"
|
||||
last_authenticated_at: datetime | None = None
|
||||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue