Merge remote-tracking branch 'origin/main' into litellm_vector_store_deny_by_default

This commit is contained in:
mrinal 2026-10-02 20:08:38 +00:00
commit b94047f2ac
116 changed files with 5724 additions and 907 deletions

View file

@ -176,6 +176,7 @@ jobs:
TESTS: ${{ needs.detect.outputs.tests }}
E2E_FIXTURE_MODE: live
E2E_PROVIDER_EDGE_HOST_REACHABLE: '1'
E2E_OWNED_GATEWAY: '1'
COLUMNS: '400'
run: |
umask 077

View file

@ -0,0 +1,17 @@
CREATE TABLE IF NOT EXISTS "LiteLLM_AutoRouterDailySpend" (
"date" TEXT NOT NULL,
"api_key" TEXT NOT NULL,
"user_id" TEXT NOT NULL,
"router_name" TEXT NOT NULL,
"router_type" TEXT NOT NULL,
"turns" INTEGER NOT NULL DEFAULT 0,
"spend" DOUBLE PRECISION NOT NULL DEFAULT 0,
"saved_spend" DOUBLE PRECISION NOT NULL DEFAULT 0,
"savings_estimated_turns" INTEGER NOT NULL DEFAULT 0,
"savings_estimated_actual_spend" DOUBLE PRECISION NOT NULL DEFAULT 0,
"savings_estimated_saved_spend" DOUBLE PRECISION NOT NULL DEFAULT 0,
"classifier_cost" DOUBLE PRECISION NOT NULL DEFAULT 0,
"classifier_cost_recorded_turns" INTEGER NOT NULL DEFAULT 0,
CONSTRAINT "LiteLLM_AutoRouterDailySpend_pkey" PRIMARY KEY ("date", "api_key", "user_id", "router_name", "router_type")
);

View file

@ -1744,6 +1744,27 @@ model LiteLLM_AutoRouterUserSession {
@@index([user_id, last_turn_at], map: "idx_autorouter_user_session_user_last_turn")
}
// Auto-routed requests per UTC request day and router: the selected-day money behind the
// auto-router usage view. Written in the same statement as the session rollup, so a day row
// and its session row never disagree; corrected in the same transaction as late baselines.
model LiteLLM_AutoRouterDailySpend {
date String
api_key String
user_id String
router_name String
router_type String
turns Int @default(0)
spend Float @default(0)
saved_spend Float @default(0)
savings_estimated_turns Int @default(0)
savings_estimated_actual_spend Float @default(0)
savings_estimated_saved_spend Float @default(0)
classifier_cost Float @default(0)
classifier_cost_recorded_turns Int @default(0)
@@id([date, api_key, user_id, router_name, router_type])
}
// Shadow eval: evaluation of an auto-router against one or more keys' live traffic, in
// either direction. forward duplicates the requests the keys did not route through the
// router through it, answering whether they should adopt it; reverse duplicates the

View file

@ -78,6 +78,23 @@ class _InvalidIndex:
table_size: str
MAX_MIGRATE_DEPLOY_ATTEMPTS = 4
LIBPQ_URL_PARAMS: Final = frozenset(
{
"sslmode",
"sslcert",
"sslkey",
"sslrootcert",
"sslpassword",
"application_name",
"connect_timeout",
"client_encoding",
"options",
"service",
"gssencmode",
"krbsrvname",
"target_session_attrs",
}
)
@dataclass(frozen=True)
@ -689,30 +706,43 @@ class ProxyExtrasDBManager:
@staticmethod
def _strip_prisma_query_params(url: str) -> str:
"""Remove Prisma-specific query params (connection_limit, pool_timeout,
schema, etc.) from DATABASE_URL so psycopg can parse it."""
"""Rewrite a Prisma-dialect URL for libpq: drop the Prisma-only params
(connection_limit, pool_timeout, schema, pgbouncer, sslaccept, ...) and
translate Prisma's TLS params back, since libpq reads ``sslcert`` as a
client certificate where Prisma reads it as the CA."""
from urllib.parse import parse_qsl, quote, urlencode, urlparse, urlunparse
parsed = urlparse(url)
parsed: Final = urlparse(url)
if not parsed.query:
return url
libpq_params = {
"sslmode",
"sslcert",
"sslkey",
"sslrootcert",
"sslpassword",
"application_name",
"connect_timeout",
"client_encoding",
"options",
"service",
"gssencmode",
"krbsrvname",
"target_session_attrs",
}
kept = [(k, v) for k, v in parse_qsl(parsed.query) if k in libpq_params]
return urlunparse(parsed._replace(query=urlencode(kept, quote_via=quote)))
pairs: Final = tuple(parse_qsl(parsed.query))
kept: Final = tuple((k, v) for k, v in pairs if k in LIBPQ_URL_PARAMS)
sslaccept: Final = next((v for k, v in pairs if k == "sslaccept"), None)
libpq_pairs: Final = ProxyExtrasDBManager._libpq_tls_params(kept, sslaccept)
return urlunparse(parsed._replace(query=urlencode(libpq_pairs, quote_via=quote)))
@staticmethod
def _libpq_tls_params(
pairs: "tuple[tuple[str, str], ...]", sslaccept: "str | None"
) -> "tuple[tuple[str, str], ...]":
"""Undo ``translate_libpq_ssl_params``. Prisma's ``sslcert`` is the CA and
``sslaccept=strict`` checks chain and hostname, which libpq only does in
``sslmode=verify-full``, so strict becomes ``sslrootcert`` plus
``verify-full`` whatever ``sslmode`` said (``disable`` stays off). Prisma
defaults an absent ``sslaccept`` to ``accept_invalid_certs`` and anything
else to strict. Without strict it checks nothing, so the CA is dropped and
``sslmode`` is kept as is: libpq only verifies when a root cert is present.
A URL that also carries ``sslkey`` is libpq's own client-certificate form
and is kept."""
keys: Final = frozenset(k for k, _ in pairs)
if "sslcert" not in keys or "sslkey" in keys:
return pairs
sslmode: Final = next((v for k, v in pairs if k == "sslmode"), None)
rest: Final = tuple((k, v) for k, v in pairs if k not in ("sslcert", "sslmode"))
if sslaccept in (None, "accept_invalid_certs") or sslmode == "disable":
return rest if sslmode is None else rest + (("sslmode", sslmode),)
root_cert: Final = tuple(("sslrootcert", v) for k, v in pairs if k == "sslcert" and "sslrootcert" not in keys)
return rest + root_cert + (("sslmode", "verify-full"),)
@staticmethod
def _warn_if_db_ahead_of_head(migrations_dir: str) -> None:

View file

@ -1,11 +1,13 @@
WITH page AS (
SELECT TraceId AS trace_id,
hex(SHA256(concat(TeamId, char(0), ApiKeyHash, char(0), TraceId))) AS trace_ref,
TeamId AS team_id, ApiKeyHash AS api_key_hash,
ifNull(any(RootName), '') AS name, any(ServiceName) AS service,
ifNull(any(RootInput), '') AS input_preview, ifNull(any(RootStatus), '') AS status,
toUnixTimestamp64Milli(min(StartTs)) AS start_ms,
min(StartTs) AS trace_start, max(EndTs) AS trace_end,
dateDiff('millisecond', min(StartTs), max(EndTs)) AS duration_ms,
sum(SpanCount) AS span_count, length(groupUniqArrayArray(AgentNames)) AS agent_count,
sum(SpanCount) AS span_count,
sum(AgentCount) AS agent_invocations,
sum(LlmCount) AS llm_calls, sum(ToolCount) AS tool_calls,
sum(InputTokens) AS input_tokens, sum(OutputTokens) AS output_tokens,
@ -21,3 +23,21 @@ HAVING min(StartTs) >= fromUnixTimestamp64Milli({start_ms:Int64})
< ({cursor_ms:Int64}, {cursor_trace_id:String}))
ORDER BY start_ms DESC, trace_ref DESC
LIMIT {limit:UInt32}
)
SELECT page.* EXCEPT (trace_start, trace_end),
identities.agent_names AS agent_names, identities.agent_count AS agent_count
FROM page
LEFT JOIN (
SELECT TeamId, ApiKeyHash, TraceId,
arraySort(groupUniqArrayIf(AgentName, AgentName != '')) AS agent_names,
uniqExactIf(if(AgentName = '', SpanName, AgentName), ObservationType = 'agent') AS agent_count
FROM otel_traces
WHERE Timestamp >= (SELECT min(trace_start) FROM page)
AND Timestamp <= (SELECT max(trace_end) FROM page)
AND TraceId IN (SELECT trace_id FROM page)
AND (TeamId, ApiKeyHash, TraceId) IN (SELECT team_id, api_key_hash, trace_id FROM page)
GROUP BY TeamId, ApiKeyHash, TraceId
) AS identities
ON page.team_id = identities.TeamId AND page.api_key_hash = identities.ApiKeyHash
AND page.trace_id = identities.TraceId
ORDER BY page.start_ms DESC, page.trace_ref DESC

View file

@ -1,7 +1,7 @@
use std::collections::BTreeMap;
use crate::DecodeError;
use serde::Serialize;
use serde::{Deserialize, Serialize};
#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize)]
#[serde(rename_all = "lowercase")]
@ -148,6 +148,60 @@ fn usage_tokens(attributes: &BTreeMap<String, String>) -> Result<(u32, u32), Dec
))
}
#[derive(Default, Deserialize)]
struct AgentMetadata {
#[serde(default)]
lc_agent_name: String,
#[serde(default)]
ls_integration: String,
}
fn recorded_agent_name(
name: &str,
attributes: &BTreeMap<String, String>,
span: &NormalizedSpan,
) -> String {
let explicit = [
span.agent_name.as_str(),
attr(attributes, "gen_ai.agent.name"),
attr(attributes, "agent.name"),
attr(attributes, "openclaw.agent"),
]
.into_iter()
.find(|value| !value.is_empty());
if let Some(value) = explicit {
return value.to_owned();
}
let metadata =
serde_json::from_str::<AgentMetadata>(attr(attributes, "metadata")).unwrap_or_default();
if !metadata.lc_agent_name.is_empty() {
return metadata.lc_agent_name;
}
if span.observation_type == ObservationType::Agent {
let node = attr(attributes, "graph.node.id");
if !node.is_empty() {
return node.to_owned();
}
if metadata.ls_integration == "langgraph" && name != "LangGraph" && !is_middleware(name) {
return name.to_owned();
}
}
String::new()
}
fn is_middleware(name: &str) -> bool {
[
".wrap_model_call",
".wrap_tool_call",
".before_agent",
".after_agent",
".before_model",
".after_model",
]
.iter()
.any(|suffix| name.ends_with(suffix))
}
pub fn normalize(
scope_name: &str,
name: &str,
@ -163,8 +217,22 @@ pub fn normalize(
.into_iter()
.find(|normalizer| normalizer.matches(scope_name, attributes))
.expect("GenAI fallback always matches");
let span = normalizer.normalize(name, parent_span_id, attributes)?;
let agent_name = recorded_agent_name(name, attributes, &span);
let observation_type = if !parent_span_id.is_empty()
&& scope_name == "openinference.instrumentation.langchain"
&& is_middleware(name)
{
ObservationType::Framework
} else {
span.observation_type
};
Ok(Normalization {
span: normalizer.normalize(name, parent_span_id, attributes)?,
span: NormalizedSpan {
agent_name,
observation_type,
..span
},
consumed_attributes: normalizer.consumed_attributes(attributes),
})
}

View file

@ -133,7 +133,18 @@ fn decoded_span(
&parent_span_id,
&span_attributes,
)?;
let normalized = normalization.span;
let resource_agent_name = resource_attributes
.get("gen_ai.agent.name")
.filter(|name| !name.is_empty());
let agent_name = match (resource_agent_name, normalization.span.agent_name.as_str()) {
(Some(name), "") => name.clone(),
(Some(name), "hermes-agent") if scope_name.as_ref() == "hermes-otel-plugin" => name.clone(),
(_, name) => name.to_owned(),
};
let normalized = crate::normalize::NormalizedSpan {
agent_name,
..normalization.span
};
budget.consume(
normalized.input.len()
+ normalized.output.len()

View file

@ -362,6 +362,143 @@ async fn keyed_rollup_keeps_same_trace_ids_separate_by_api_key(
Ok(())
}
#[rstest]
#[tokio::test]
async fn listed_agent_names_preserve_scope_and_cursor(
#[future(awt)] database: TestResult<ClickHouseDatabase>,
) -> TestResult {
let database = database?;
let writer = Connection::writer(&database.url)?;
ensure_schema(&database.client, &writer, "trace_test", 7).await?;
let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64;
for (team, key, trace, agent, span, parent) in [
("alpha", "one", "shared", "research_agent", "root", ""),
("alpha", "one", "shared", "reviewer", "child", "root"),
("alpha", "one", "shared", "reviewer", "repeated", "root"),
("alpha", "one", "shared", "", "unnamed", "root"),
("alpha", "one", "second", "support_agent", "root", ""),
("alpha", "two", "shared", "private_agent", "root", ""),
("beta", "one", "shared", "other_agent", "root", ""),
] {
insert_rows(
&database,
"otel_traces",
vec![serde_json::from_value(serde_json::json!({
"Timestamp": timestamp, "TraceId": trace, "SpanId": span, "ParentSpanId": parent,
"ServiceName": "shared-app", "SpanName": span, "AgentName": agent,
"ObservationType": "agent",
"ResourceAttributes": {"litellm.team_id": team, "litellm.api_key_hash": key}
}))?],
)
.await?;
}
let historical_rows = (0..5000)
.map(|index| {
serde_json::from_value(serde_json::json!({
"Timestamp": timestamp - 86_400_000_000_000_i64,
"TraceId": "shared", "SpanId": format!("historical-{index}"),
"ParentSpanId": "", "SpanName": "historical", "AgentName": "private_agent",
"ObservationType": "agent", "ServiceName": "shared-app",
"ResourceAttributes": {"litellm.team_id": "alpha", "litellm.api_key_hash": "history"}
}))
})
.collect::<Result<Vec<_>, _>>()?;
insert_rows(&database, "otel_traces", historical_rows).await?;
let connection = Connection::configured(&database.url, "trace_test", "default", "")?;
let parameters = BTreeMap::from([
("team_ids".into(), Parameter::Strings(vec!["alpha".into()])),
("api_key_hash".into(), Parameter::Text("one".into())),
(
"start_ms".into(),
Parameter::Integer(timestamp / 1_000_000 - 1000),
),
(
"end_ms".into(),
Parameter::Integer(timestamp / 1_000_000 + 1000),
),
("cursor_ms".into(), Parameter::Integer(0)),
("cursor_trace_id".into(), Parameter::Text(String::new())),
("limit".into(), Parameter::Integer(1)),
]);
let first: serde_json::Value = serde_json::from_str(
&execute_named_read(
&database.client,
&connection,
ReadQuery::ListTraces,
&parameters,
)
.await?,
)?;
let cursor = first["data"][0]["trace_ref"]
.as_str()
.ok_or("missing cursor")?;
let next_parameters = parameters
.into_iter()
.chain([
(
"cursor_ms".into(),
Parameter::Integer(timestamp / 1_000_000),
),
("cursor_trace_id".into(), Parameter::Text(cursor.into())),
])
.collect();
let second: serde_json::Value = serde_json::from_str(
&execute_named_read(
&database.client,
&connection,
ReadQuery::ListTraces,
&next_parameters,
)
.await?,
)?;
assert_eq!(
first["data"].as_array().ok_or("missing first page")?.len(),
1
);
assert_eq!(
second["data"]
.as_array()
.ok_or("missing second page")?
.len(),
1
);
assert_ne!(first["data"][0]["trace_id"], second["data"][0]["trace_id"]);
let names = [&first["data"][0], &second["data"][0]]
.into_iter()
.map(|row| {
(
row["trace_id"].as_str().unwrap(),
row["agent_names"].clone(),
)
})
.collect::<BTreeMap<_, _>>();
assert_eq!(
names["shared"],
serde_json::json!(["research_agent", "reviewer"])
);
assert_eq!(names["second"], serde_json::json!(["support_agent"]));
let counts = [&first["data"][0], &second["data"][0]]
.into_iter()
.map(|row| {
(
row["trace_id"].as_str().unwrap(),
row["agent_count"].as_u64(),
)
})
.collect::<BTreeMap<_, _>>();
assert_eq!(counts["shared"], Some(3));
assert_eq!(counts["second"], Some(1));
for page in [&first, &second] {
assert!(
page["statistics"]["rows_read"]
.as_u64()
.ok_or("missing read statistics")?
< 5000
);
}
Ok(())
}
#[rstest]
#[tokio::test]
async fn rollup_merges_spans_across_days_without_losing_root_fields(
@ -376,6 +513,7 @@ async fn rollup_merges_spans_across_days_without_losing_root_fields(
let root = serde_json::from_value(serde_json::json!({
"Timestamp": day_start - 1_000_000_000, "TraceId": "cross-day", "SpanId": "span-root",
"ParentSpanId": "", "ServiceName": "proxy", "SpanName": "root", "Input": "root input",
"AgentName": "lead", "ObservationType": "agent",
"StatusCode": "STATUS_CODE_ERROR",
"ResourceAttributes": {"litellm.team_id": "team-1"}
}))?;
@ -383,6 +521,7 @@ async fn rollup_merges_spans_across_days_without_losing_root_fields(
let child = serde_json::from_value(serde_json::json!({
"Timestamp": day_start + 1_000_000_000, "TraceId": "cross-day", "SpanId": "span-child",
"ParentSpanId": "span-root", "ServiceName": "proxy", "SpanName": "child",
"AgentName": "researcher", "ObservationType": "agent",
"StatusCode": "STATUS_CODE_UNSET",
"ResourceAttributes": {"litellm.team_id": "team-1"}
}))?;
@ -406,6 +545,33 @@ async fn rollup_merges_spans_across_days_without_losing_root_fields(
"RootStatus": "STATUS_CODE_ERROR", "SpanCount": 2
}])
);
let connection = Connection::configured(&database.url, "trace_test", "default", "")?;
let parameters = BTreeMap::from([
("team_ids".into(), Parameter::Strings(vec!["team-1".into()])),
("api_key_hash".into(), Parameter::Text(String::new())),
(
"start_ms".into(),
Parameter::Integer(day_start / 1_000_000 - 2000),
),
("end_ms".into(), Parameter::Integer(day_start / 1_000_000)),
("cursor_ms".into(), Parameter::Integer(0)),
("cursor_trace_id".into(), Parameter::Text(String::new())),
("limit".into(), Parameter::Integer(10)),
]);
let listed: serde_json::Value = serde_json::from_str(
&execute_named_read(
&database.client,
&connection,
ReadQuery::ListTraces,
&parameters,
)
.await?,
)?;
assert_eq!(
listed["data"][0]["agent_names"],
serde_json::json!(["lead", "researcher"])
);
assert_eq!(listed["data"][0]["agent_count"], 2);
Ok(())
}

View file

@ -915,7 +915,8 @@ def _excluded_db_systems(logger: "OpenTelemetryV2") -> frozenset[str]:
logger got published: with ``callbacks: [langfuse_otel, otel]`` the ``otel``
callback folds into the preset, whose config is env-only.
"""
configured: Final = litellm.callback_settings.get("otel", {}).get("excluded_services")
otel_settings: Final = (litellm.callback_settings or {}).get("otel")
configured: Final = otel_settings.get("excluded_services") if isinstance(otel_settings, dict) else None
if configured is None:
return logger.config.excluded_services
return excluded_db_systems_from(configured)

View file

@ -5,7 +5,8 @@ from functools import lru_cache
from typing import Annotated, Any, Final
from pydantic import AliasChoices, BaseModel, Field, TypeAdapter, ValidationError, field_validator, model_validator
from pydantic_settings import BaseSettings, NoDecode, SettingsConfigDict
from pydantic.fields import FieldInfo
from pydantic_settings import BaseSettings, NoDecode, PydanticBaseSettingsSource, SettingsConfigDict
from litellm._logging import verbose_logger
from litellm.integrations.otel.model.baggage import (
@ -121,9 +122,37 @@ class ExporterSpec(BaseModel):
)
class _EnvWithoutBareExcludedServices(PydanticBaseSettingsSource):
def __init__(self, settings_cls: type[BaseSettings], env_settings: PydanticBaseSettingsSource) -> None:
super().__init__(settings_cls)
self._env_settings: Final = env_settings
def get_field_value(self, field: FieldInfo, field_name: str) -> tuple[object, str, bool]:
return self._env_settings.get_field_value(field, field_name)
def __call__(self) -> dict[str, object]:
return {key: value for key, value in self._env_settings().items() if key != "excluded_services"}
class OpenTelemetryV2Config(BaseSettings):
model_config = SettingsConfigDict(populate_by_name=True, extra="ignore")
@classmethod
def settings_customise_sources(
cls,
settings_cls: type[BaseSettings],
init_settings: PydanticBaseSettingsSource,
env_settings: PydanticBaseSettingsSource,
dotenv_settings: PydanticBaseSettingsSource,
file_secret_settings: PydanticBaseSettingsSource,
) -> tuple[PydanticBaseSettingsSource, ...]:
return (
init_settings,
_EnvWithoutBareExcludedServices(settings_cls, env_settings),
dotenv_settings,
file_secret_settings,
)
# ----- single-destination shorthand, read from standard OTEL_* envs ----- #
exporter: str = Field(
default="console",
@ -178,7 +207,7 @@ class OpenTelemetryV2Config(BaseSettings):
)
excluded_services: Annotated[frozenset[str], NoDecode] = Field(
default_factory=frozenset,
validation_alias=AliasChoices("excluded_services", "LITELLM_OTEL_EXCLUDED_SERVICES"),
validation_alias=AliasChoices("LITELLM_OTEL_EXCLUDED_SERVICES"),
description=(
"Datastore services whose spans are withheld from key/team ``callback_vars`` "
"OTel destinations (the operator's own exporters still receive them). Accepted "

View file

@ -0,0 +1 @@

View file

@ -0,0 +1,60 @@
from collections.abc import Mapping
from dataclasses import dataclass, field
from typing import Final, Literal, TypeAlias
from pydantic import AnyHttpUrl, BaseModel, TypeAdapter, ValidationError
from litellm.secret_managers.main import get_secret_str
LayaCheckpoint: TypeAlias = Literal["english", "multilingual", "typed-decisions"]
def validate_laya_model(value: object) -> LayaCheckpoint:
try:
return TypeAdapter(LayaCheckpoint).validate_python(value)
except ValidationError as exc:
raise ValueError("Laya model must be 'english', 'multilingual', or 'typed-decisions'") from exc
def validate_laya_request(body: Mapping[str, object]) -> LayaCheckpoint:
if "custom_body" in body:
raise ValueError("custom_body is not supported for Laya requests")
if body.get("stream"):
raise ValueError("Streaming is not supported for Laya requests")
return validate_laya_model(body.get("model"))
@dataclass(frozen=True, slots=True)
class LayaConnection:
api_base: str
api_key: str | None = field(repr=False)
def validate_laya_api_base(value: str) -> str:
try:
url: Final = TypeAdapter(AnyHttpUrl).validate_python(value)
except ValidationError as exc:
raise ValueError("Laya api_base must be an HTTP or HTTPS server URL") from exc
if url.username or url.password or url.query or url.fragment:
raise ValueError("Laya api_base must not contain credentials, a query, or a fragment")
return str(url).rstrip("/")
def laya_connection(api_base: str | None = None, api_key: str | None = None) -> LayaConnection:
base: Final = api_base if api_base is not None else get_secret_str("LAYA_API_BASE")
if not base:
raise ValueError("Laya requires api_base or LAYA_API_BASE pointing to a self-hosted server")
key: Final = api_key if api_base is not None else api_key or get_secret_str("LAYA_API_KEY")
return LayaConnection(api_base=validate_laya_api_base(base), api_key=key)
class _LayaRouting(BaseModel):
model: str | None = None
def laya_response_model(response: Mapping[str, object], requested_model: str | None) -> str:
try:
routing: Final = TypeAdapter(_LayaRouting).validate_python(response.get("routing") or _LayaRouting())
except ValidationError:
return requested_model or "unknown"
return routing.model or requested_model or "unknown"

View file

@ -72622,6 +72622,45 @@
"supports_audio_input": true,
"supports_video_input": true
},
"laya/english": {
"input_cost_per_token": 0.0,
"litellm_provider": "laya",
"mode": "evaluation",
"output_cost_per_token": 0.0,
"source": "https://github.com/NandhaKishorM/laya",
"supported_endpoints": [
"/v1/systemone"
],
"metadata": {
"notes": "Self-hosted decision model; infrastructure costs are paid separately"
}
},
"laya/multilingual": {
"input_cost_per_token": 0.0,
"litellm_provider": "laya",
"mode": "evaluation",
"output_cost_per_token": 0.0,
"source": "https://github.com/NandhaKishorM/laya",
"supported_endpoints": [
"/v1/systemone"
],
"metadata": {
"notes": "Self-hosted decision model; infrastructure costs are paid separately"
}
},
"laya/typed-decisions": {
"input_cost_per_token": 0.0,
"litellm_provider": "laya",
"mode": "evaluation",
"output_cost_per_token": 0.0,
"source": "https://github.com/NandhaKishorM/laya",
"supported_endpoints": [
"/v1/systemone"
],
"metadata": {
"notes": "Self-hosted decision model; infrastructure costs are paid separately"
}
},
"typesafe/jev-1.13.0": {
"input_cost_per_token": 4.2e-08,
"litellm_provider": "typesafe",

View file

@ -229,6 +229,7 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = (
"/tinyfish/",
"/transcribe",
"/typesafe/",
"/laya/",
"/openrouter/",
"/vertex-ai/",
"/vertex_ai/",

View file

@ -27761,6 +27761,30 @@
]
}
},
"/laya/v1/systemone": {
"post": {
"operationId": "laya_proxy_route_laya_v1_systemone_post",
"responses": {
"200": {
"content": {
"application/json": {
"schema": {}
}
},
"description": "Successful Response"
}
},
"security": [
{
"APIKeyHeader": []
}
],
"summary": "Laya Proxy Route",
"tags": [
"llm_passthrough"
]
}
},
"/milvus/{endpoint}": {
"delete": {
"description": "Enable using Milvus `/vectors` endpoint as a pass-through endpoint.",

View file

@ -507,6 +507,7 @@ class LiteLLMRoutes(enum.Enum):
"/vllm",
"/mistral",
"/typesafe",
"/laya",
"/openrouter",
"/milvus",
"/gigachat",
@ -5474,6 +5475,17 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
"'auto_register': auto-create a virtual key and mapping on first encounter."
),
)
auto_register_map_existing_key: bool = Field(
default=False,
description=(
"Only used with unregistered_jwt_client_behavior='auto_register'. When True and the virtual key claim "
"field is the user_id_jwt_field or user_email_jwt_field, the JWT claim is mapped to a virtual key the "
"JWT-resolved user already owns instead of minting a new one. If the user owns several, the most recently created key in the "
"JWT-resolved team (or with no team when the JWT resolves none) is chosen among keys that never "
"expire, are not blocked, are not Admin UI session keys, were not minted by auto_register, and "
"have no allowed_routes or include llm_api_routes. Otherwise a new key is minted as usual."
),
)
routing_overrides: list[JWTRoutingOverride] | None = Field(
default=None,
description="Optional claim-based routing overrides for JWT-shaped tokens. Matching rules route requests to oauth2 before default JWT flow.",
@ -5574,6 +5586,15 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
return issuer_config.virtual_key_claim_field
return self.virtual_key_claim_field
def is_user_identity_claim(self, claim_field: str, issuer: str | None) -> bool:
issuer_config: Final = self.get_issuer_config(issuer)
if issuer_config is None:
return claim_field in (self.user_id_jwt_field, self.user_email_jwt_field)
return claim_field in (
issuer_config.user_id_jwt_field or self.user_id_jwt_field,
issuer_config.user_email_jwt_field or self.user_email_jwt_field,
)
def get_unregistered_jwt_client_behavior(self, issuer: str | None) -> UnregisteredJWTClientBehavior:
issuer_config: Final = self.get_issuer_config(issuer)
if issuer_config is not None and issuer_config.unregistered_jwt_client_behavior is not None:

View file

@ -1883,6 +1883,15 @@ def _extract_model_candidates_from_request(
llm_router: Router | None = None,
team_id: str | None = None,
) -> list[str]:
if route.rstrip("/") == "/laya/v1/systemone":
from litellm.llms.laya.common_utils import validate_laya_model
try:
laya_request: Final = TypeAdapter(Mapping[str, object]).validate_python(request_data)
laya_model: Final = validate_laya_model(laya_request.get("model"))
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
return _dedupe_model_candidates((f"laya/{laya_model}",))
if route == "/cost/predict-cache":
prediction_models: Final = _cache_prediction_model_candidates(request_data, llm_router, team_id) # pyright: ignore[reportUnknownArgumentType] # the typed reader validates each deployment ID from this legacy payload
return _dedupe_model_candidates(prediction_models)

View file

@ -140,6 +140,7 @@ from litellm.proxy.utils import (
normalize_route_for_root_path,
)
from litellm.repositories.table_repositories import TeamMembershipRepository
from litellm.repositories.verification_token_repository import VerificationTokenRepository
from litellm.router_utils.common_utils import resolve_model_group_alias
from litellm.secret_managers.main import get_secret_bool
from litellm.types.services import ServiceTypes
@ -939,6 +940,24 @@ class _PendingAutoRegister(NamedTuple):
jwt_issuer: str | None = None
def _claim_identifies_user(jwt_handler: JWTHandler, claim_field: str, jwt_issuer: str | None) -> bool:
if not jwt_handler.litellm_jwtauth.auto_register_map_existing_key:
return False
if jwt_handler.litellm_jwtauth.is_user_identity_claim(claim_field, jwt_issuer):
return True
verbose_proxy_logger.warning(
"JWT Key Mapping (auto_register_map_existing_key): claim '%s' is not the user_id or user_email JWT field "
"and may be shared by several users, so a new key is minted instead of reusing one the user owns.",
claim_field,
)
return False
async def _reusable_key_hash_for_user(prisma_client: PrismaClient, user_id: str, team_id: str | None) -> str | None:
key: Final = await VerificationTokenRepository(prisma_client).find_newest_reusable_llm_api_key(user_id, team_id)
return None if key is None else key.token
async def _auto_register_jwt_mapping(
virtual_key_claim_field: str,
claim_value: str,
@ -957,8 +976,10 @@ async def _auto_register_jwt_mapping(
) -> UserAPIKeyAuth | None:
"""
Auto-register: create a new virtual key + mapping for an unrecognised JWT
claim value. ``team_id`` and ``user_id`` must come from a successful
``JWTAuthManager.auth_builder`` run — they encode the JWT identity AFTER
claim value, or point the mapping at a key the resolved user already owns
when ``auto_register_map_existing_key`` is set. ``team_id`` and ``user_id``
must come from a successful ``JWTAuthManager.auth_builder`` run — they
encode the JWT identity AFTER
RBAC/scope/custom_validate/email-domain policy has been enforced. The key
is stamped with those values so the cached future-request path inherits
the same team/user/org limits the auth_builder path would have applied.
@ -974,29 +995,38 @@ async def _auto_register_jwt_mapping(
generate_key_helper_fn,
)
# ``table_name="key"`` is required: without it, generate_key_helper_fn
# falls into the user-upsert branch (`table_name is None or "user"`) and
# attempts to insert into LiteLLM_UserTable with user_id=None, which fails
# the NOT NULL @id constraint. Every successful key-creation caller (e.g.
# /key/generate) passes table_name="key" explicitly.
key_data: Final = await generate_key_helper_fn(
llm_router=None,
request_type="key",
table_name="key",
team_id=team_id,
user_id=user_id,
organization_id=org_id,
agent_id=agent_id,
metadata={
"auto_registered": True,
"jwt_claim_field": virtual_key_claim_field,
"jwt_claim_value": claim_value,
},
existing_token_hash: Final = (
await _reusable_key_hash_for_user(prisma_client, user_id, team_id)
if user_id is not None and _claim_identifies_user(jwt_handler, virtual_key_claim_field, jwt_issuer)
else None
)
# generate_key_helper_fn returns the plaintext key in "token"; the persisted
# row in LiteLLM_VerificationToken uses its hash, so hash here to get the FK
# value referenced by LiteLLM_JWTKeyMapping.token.
token_hash = hash_token(key_data["token"])
minted: Final = existing_token_hash is None
if existing_token_hash is not None:
token_hash = existing_token_hash
else:
# ``table_name="key"`` is required: without it, generate_key_helper_fn
# falls into the user-upsert branch (`table_name is None or "user"`) and
# attempts to insert into LiteLLM_UserTable with user_id=None, which fails
# the NOT NULL @id constraint. Every successful key-creation caller (e.g.
# /key/generate) passes table_name="key" explicitly.
key_data: Final = await generate_key_helper_fn(
llm_router=None,
request_type="key",
table_name="key",
team_id=team_id,
user_id=user_id,
organization_id=org_id,
agent_id=agent_id,
metadata={
"auto_registered": True,
"jwt_claim_field": virtual_key_claim_field,
"jwt_claim_value": claim_value,
},
)
# generate_key_helper_fn returns the plaintext key in "token"; the persisted
# row in LiteLLM_VerificationToken uses its hash, so hash here to get the FK
# value referenced by LiteLLM_JWTKeyMapping.token.
token_hash = hash_token(key_data["token"])
try:
await prisma_client.db.litellm_jwtkeymapping.create(
@ -1023,15 +1053,16 @@ async def _auto_register_jwt_mapping(
virtual_key_claim_field,
claim_value,
)
try:
await prisma_client.db.litellm_verificationtoken.delete(where={"token": token_hash})
except Exception as delete_err:
# Don't fail the request if cleanup fails — the orphan is
# unmapped and inert. Log so an operator can prune it later.
verbose_proxy_logger.warning(
"JWT Key Mapping (auto_register): failed to delete orphaned key after race: %s",
delete_err,
)
if minted:
try:
await prisma_client.db.litellm_verificationtoken.delete(where={"token": token_hash})
except Exception as delete_err:
# Don't fail the request if cleanup fails — the orphan is
# unmapped and inert. Log so an operator can prune it later.
verbose_proxy_logger.warning(
"JWT Key Mapping (auto_register): failed to delete orphaned key after race: %s",
delete_err,
)
token_hash = await get_jwt_key_mapping_object(
jwt_claim_name=virtual_key_claim_field,
jwt_claim_value=claim_value,
@ -1061,7 +1092,8 @@ async def _auto_register_jwt_mapping(
)
verbose_proxy_logger.info(
"JWT Key Mapping (auto_register): created new virtual key for %s='%s'.",
"JWT Key Mapping (auto_register): %s virtual key for %s='%s'.",
"created new" if minted else "mapped existing",
virtual_key_claim_field,
claim_value,
)
@ -1075,7 +1107,8 @@ async def _auto_register_jwt_mapping(
).resolve(hashed_token=token_hash)
)
if auto_registered_key is not None:
auto_registered_key.org_id = org_id
if minted:
auto_registered_key.org_id = org_id
auto_registered_key.end_user_id = end_user_id
auto_registered_key.api_key = auto_registered_key.token
return auto_registered_key
@ -1771,8 +1804,8 @@ async def _user_api_key_auth_builder(
# mapping + virtual key from the *validated* identity, then
# replace valid_token with the new key so downstream checks
# use the key-scoped path.
if pending_auto_register is not None and prisma_client is not None:
auto_registered: Final = await _auto_register_jwt_mapping(
auto_registered: Final = (
await _auto_register_jwt_mapping(
virtual_key_claim_field=pending_auto_register.claim_field,
claim_value=pending_auto_register.claim_value,
jwt_handler=jwt_handler,
@ -1788,72 +1821,81 @@ async def _user_api_key_auth_builder(
end_user_id=end_user_id,
agent_id=agent_id,
)
if auto_registered is not None:
auto_registered.jwt_claims = jwt_claims
auto_registered.user_email = user_email
# The auto-registered token is built from the new key's
# columns, which carry no user budget. Carry over the
# already-loaded user row rather than re-reading it, or
# the budget check below has nothing to enforce.
auto_registered.user_model_max_budget = (
user_object.model_max_budget if user_object is not None else None
)
valid_token = auto_registered
api_key = valid_token.token or ""
# Check if model has zero cost - if so, skip all budget checks
model = _get_model_from_request_context(
request_data=request_data,
route=route,
request=request,
llm_router=llm_router,
team_id=valid_token.team_id,
if pending_auto_register is not None and prisma_client is not None
else None
)
skip_budget_checks = False
if model is not None and llm_router is not None:
from litellm.proxy.auth.auth_checks import _is_model_cost_zero
skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router)
if skip_budget_checks:
verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model)
# Fetch project object for JWT path if project_id is set
_jwt_project_obj = None
if valid_token.project_id is not None:
_jwt_project_obj = await get_project_object(
project_id=valid_token.project_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
if auto_registered is not None:
auto_registered.jwt_claims = jwt_claims
auto_registered.user_email = user_email
# The auto-registered token is built from the new key's
# columns, which carry no user budget. Carry over the
# already-loaded user row rather than re-reading it, or
# the budget check below has nothing to enforce.
auto_registered.user_model_max_budget = (
user_object.model_max_budget if user_object is not None else None
)
if _jwt_project_obj is not None:
valid_token.project_metadata = _jwt_project_obj.metadata
valid_token.project_alias = _jwt_project_obj.project_alias
valid_token = auto_registered
api_key = valid_token.token or ""
# JWT auth returns here rather than falling through to the
# virtual-key checks below, so the user's per-model budget
# has to be enforced on this path too. Without it the
# post-call increment still charges the counter and nothing
# ever reads it, which is worse than not tracking at all.
# Guarded by the same flag the virtual-key path uses, or a
# zero-cost model would be refused here and allowed there,
# while the log above claims all budget checks were skipped.
if not skip_budget_checks:
await _check_user_model_budget(
valid_token=cast(UserAPIKeyAuth, valid_token),
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=valid_token.team_id,
)
),
falls_through_to_key_checks: Final = (
auto_registered is not None
and jwt_handler.litellm_jwtauth.auto_register_map_existing_key
and master_key is not None
)
if not falls_through_to_key_checks:
# Check if model has zero cost - if so, skip all budget checks
model = _get_model_from_request_context(
request_data=request_data,
route=route,
request=request,
llm_router=llm_router,
team_id=valid_token.team_id,
)
skip_budget_checks = False
if model is not None and llm_router is not None:
from litellm.proxy.auth.auth_checks import _is_model_cost_zero
return cast(UserAPIKeyAuth, valid_token)
skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router)
if skip_budget_checks:
verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model)
# Fetch project object for JWT path if project_id is set
_jwt_project_obj = None
if valid_token.project_id is not None:
_jwt_project_obj = await get_project_object(
project_id=valid_token.project_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
if _jwt_project_obj is not None:
valid_token.project_metadata = _jwt_project_obj.metadata
valid_token.project_alias = _jwt_project_obj.project_alias
# JWT auth returns here rather than falling through to the
# virtual-key checks below, so the user's per-model budget
# has to be enforced on this path too. Without it the
# post-call increment still charges the counter and nothing
# ever reads it, which is worse than not tracking at all.
# Guarded by the same flag the virtual-key path uses, or a
# zero-cost model would be refused here and allowed there,
# while the log above claims all budget checks were skipped.
if not skip_budget_checks:
await _check_user_model_budget(
valid_token=cast(UserAPIKeyAuth, valid_token),
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=valid_token.team_id,
)
),
)
return cast(UserAPIKeyAuth, valid_token)
#### ELSE ####
## CHECK PASS-THROUGH ENDPOINTS ##

View file

@ -4,11 +4,12 @@ Per-session auto-router benchmarks rollup.
At request time the spend writer builds one AutoRouterTurnTransaction per successful
auto-routed request (a request whose metadata carries a routing_decision) and queues it
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
key and user session rollups, plus the per-day router rollup, 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 can recover matching historical
costs from retained spend logs when estimate coverage predates these columns.
pods compose. The benchmarks endpoint reads session shape from the session rows and money from the
day rows, so spend and savings count only requests on the selected UTC days.
"""
from __future__ import annotations
@ -71,45 +72,82 @@ tier_maps AS (
GROUP BY router_name, router_type, kv.key
) per_tier
GROUP BY router_name, router_type
),
sessions AS (
SELECT
router_name,
router_type,
COUNT(*)::int AS sessions,
SUM(turns)::int AS session_turns,
SUM(unordered_turns)::int AS unordered_turns,
SUM(covered_turns)::int AS covered_turns,
SUM(cache_hits)::int AS cache_hits,
SUM(same_model_turns)::int AS same_model_turns,
SUM(same_model_hits)::int AS same_model_hits,
SUM(first_visit_turns)::int AS first_visit_turns,
SUM(first_visit_hits)::int AS first_visit_hits,
SUM(return_turns)::int AS return_turns,
SUM(return_hits)::int AS return_hits,
SUM(return_expired_misses)::int AS return_expired_misses,
SUM(return_within_ttl_misses)::int AS return_within_ttl_misses,
SUM(ttl_5m_turns)::int AS ttl_5m_turns,
SUM(ttl_1h_turns)::int AS ttl_1h_turns,
SUM(total_tokens)::bigint AS total_tokens,
SUM(EXTRACT(EPOCH FROM (last_turn_at - first_turn_at)))::float8 AS session_seconds
FROM windowed
GROUP BY router_name, router_type
),
days AS (
SELECT
router_name,
router_type,
SUM(turns)::int AS turns,
SUM(spend)::float8 AS spend,
SUM(saved_spend)::float8 AS saved_spend,
SUM(savings_estimated_turns)::int AS savings_estimated_turns,
SUM(savings_estimated_actual_spend)::float8 AS savings_estimated_actual_spend,
SUM(savings_estimated_saved_spend)::float8 AS savings_estimated_saved_spend,
SUM(classifier_cost)::float8 AS classifier_cost,
SUM(classifier_cost_recorded_turns)::int AS classifier_cost_recorded_turns
FROM "LiteLLM_AutoRouterDailySpend"
WHERE date >= $5 AND date <= $6
AND ($3::text IS NULL OR api_key = $3::text)
AND ($4::text IS NULL OR user_id = $4::text)
GROUP BY router_name, router_type
)
SELECT
agg.*,
COALESCE(tier_maps.tier_turns, '{{}}'::jsonb) AS tier_turns
FROM (
SELECT
router_name,
router_type,
COUNT(*)::int AS sessions,
COALESCE(SUM(turns), 0)::int AS turns,
COALESCE(SUM(unordered_turns), 0)::int AS unordered_turns,
COALESCE(SUM(covered_turns), 0)::int AS covered_turns,
COALESCE(SUM(cache_hits), 0)::int AS cache_hits,
COALESCE(SUM(same_model_turns), 0)::int AS same_model_turns,
COALESCE(SUM(same_model_hits), 0)::int AS same_model_hits,
COALESCE(SUM(first_visit_turns), 0)::int AS first_visit_turns,
COALESCE(SUM(first_visit_hits), 0)::int AS first_visit_hits,
COALESCE(SUM(return_turns), 0)::int AS return_turns,
COALESCE(SUM(return_hits), 0)::int AS return_hits,
COALESCE(SUM(return_expired_misses), 0)::int AS return_expired_misses,
COALESCE(SUM(return_within_ttl_misses), 0)::int AS return_within_ttl_misses,
COALESCE(SUM(ttl_5m_turns), 0)::int AS ttl_5m_turns,
COALESCE(SUM(ttl_1h_turns), 0)::int AS ttl_1h_turns,
COALESCE(SUM(total_tokens), 0)::bigint AS total_tokens,
COALESCE(SUM(spend), 0)::float8 AS spend,
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,
COALESCE(SUM(EXTRACT(EPOCH FROM (last_turn_at - first_turn_at))), 0)::float8 AS session_seconds
FROM windowed
GROUP BY router_name, router_type
) agg
COALESCE(tier_maps.tier_turns, '{{}}'::jsonb) AS tier_turns,
COALESCE(sessions.sessions, 0) AS sessions,
COALESCE(sessions.session_turns, 0) AS session_turns,
COALESCE(sessions.unordered_turns, 0) AS unordered_turns,
COALESCE(sessions.covered_turns, 0) AS covered_turns,
COALESCE(sessions.cache_hits, 0) AS cache_hits,
COALESCE(sessions.same_model_turns, 0) AS same_model_turns,
COALESCE(sessions.same_model_hits, 0) AS same_model_hits,
COALESCE(sessions.first_visit_turns, 0) AS first_visit_turns,
COALESCE(sessions.first_visit_hits, 0) AS first_visit_hits,
COALESCE(sessions.return_turns, 0) AS return_turns,
COALESCE(sessions.return_hits, 0) AS return_hits,
COALESCE(sessions.return_expired_misses, 0) AS return_expired_misses,
COALESCE(sessions.return_within_ttl_misses, 0) AS return_within_ttl_misses,
COALESCE(sessions.ttl_5m_turns, 0) AS ttl_5m_turns,
COALESCE(sessions.ttl_1h_turns, 0) AS ttl_1h_turns,
COALESCE(sessions.total_tokens, 0) AS total_tokens,
COALESCE(sessions.session_seconds, 0) AS session_seconds,
COALESCE(days.turns, 0) AS turns,
COALESCE(days.spend, 0) AS spend,
COALESCE(days.saved_spend, 0) AS saved_spend,
COALESCE(days.savings_estimated_turns, 0) AS savings_estimated_turns,
COALESCE(days.savings_estimated_actual_spend, 0) AS savings_estimated_actual_spend,
COALESCE(days.savings_estimated_saved_spend, 0) AS savings_estimated_saved_spend,
COALESCE(days.classifier_cost, 0) AS classifier_cost,
COALESCE(days.classifier_cost_recorded_turns, 0) AS classifier_cost_recorded_turns
FROM sessions
FULL OUTER JOIN days USING (router_name, router_type)
LEFT JOIN tier_maps USING (router_name, router_type)
ORDER BY agg.spend DESC
ORDER BY spend DESC, router_name, router_type
"""
@ -391,15 +429,43 @@ ON CONFLICT ({user_column}api_key, session_id, router_name) DO UPDATE SET
"""
_DAY_UPSERT_SQL: Final = f"""
day_rollup AS (
INSERT INTO "LiteLLM_AutoRouterDailySpend" AS d (
date, api_key, user_id, router_name, router_type, turns, spend, saved_spend, savings_estimated_turns,
savings_estimated_actual_spend, savings_estimated_saved_spend, classifier_cost, classifier_cost_recorded_turns
)
VALUES (
({_TURN_AT}::timestamp)::date::text, {_p("api_key")}::text, {_p("user_id")}::text, {_p("router_name")},
{_p("router_type")}, 1, {_p("spend")}::float8, {_p("saved_spend")}::float8, {_p("savings_estimated_turns")}::int,
{_p("savings_estimated_actual_spend")}::float8, {_p("savings_estimated_saved_spend")}::float8,
{_p("classifier_cost")}::float8, 1
)
ON CONFLICT (date, api_key, user_id, router_name, router_type) DO UPDATE SET
turns = d.turns + 1,
spend = d.spend + EXCLUDED.spend,
saved_spend = d.saved_spend + EXCLUDED.saved_spend,
savings_estimated_turns = d.savings_estimated_turns + EXCLUDED.savings_estimated_turns,
savings_estimated_actual_spend = d.savings_estimated_actual_spend + EXCLUDED.savings_estimated_actual_spend,
savings_estimated_saved_spend = d.savings_estimated_saved_spend + EXCLUDED.savings_estimated_saved_spend,
classifier_cost = d.classifier_cost + EXCLUDED.classifier_cost,
classifier_cost_recorded_turns = d.classifier_cost_recorded_turns + 1
RETURNING 1
)
"""
UPSERT_AUTOROUTER_SESSION_SQL: Final = f"""
WITH key_rollup AS (
{_session_upsert_sql(user_scoped=False)}
RETURNING 1
)
), {_DAY_UPSERT_SQL}
{_session_upsert_sql(user_scoped=True)}
"""
UPSERT_AUTOROUTER_USER_SESSION_SQL: Final = _session_upsert_sql(user_scoped=True)
UPSERT_AUTOROUTER_USER_SESSION_SQL: Final = f"""
WITH {_DAY_UPSERT_SQL}
{_session_upsert_sql(user_scoped=True)}
"""
def _as_sql_param(value: str | float | bool | datetime | None) -> str | float | None:

View file

@ -179,6 +179,8 @@ class _Change(BaseModel):
actual_delta: float
savings_delta: float
daily: DailyBaselineAttribution | None
date: str | None = None
router_type: str | None = None
class _TransactionManager(Protocol):
@ -303,6 +305,26 @@ WHERE {user_match}session.api_key = totals.api_key AND session.session_id = tota
_UPDATE_SESSIONS: Final = _session_correction_sql(user_scoped=False)
_UPDATE_USER_SESSIONS: Final = _session_correction_sql(user_scoped=True)
_UPDATE_DAYS: Final = """
WITH totals AS (
SELECT date, api_key, user_id, router_name, router_type, SUM(covered_delta)::int AS covered_delta,
SUM(actual_delta) AS actual_delta, SUM(savings_delta) AS savings_delta
FROM jsonb_to_recordset($1::jsonb) AS x(
date text, api_key text, user_id text, router_name text, router_type text,
covered_delta int, actual_delta float8, savings_delta float8
)
WHERE date IS NOT NULL
GROUP BY date, api_key, user_id, router_name, router_type
)
UPDATE "LiteLLM_AutoRouterDailySpend" AS day
SET saved_spend = day.saved_spend + totals.savings_delta,
savings_estimated_turns = day.savings_estimated_turns + totals.covered_delta,
savings_estimated_actual_spend = day.savings_estimated_actual_spend + totals.actual_delta,
savings_estimated_saved_spend = day.savings_estimated_saved_spend + totals.savings_delta
FROM totals
WHERE day.date = totals.date AND day.api_key = totals.api_key AND day.user_id = totals.user_id
AND day.router_name = totals.router_name AND day.router_type = totals.router_type
"""
def _primary_transaction(client: PrismaClient) -> _TransactionManager:
@ -331,6 +353,8 @@ def _change(record: BaselineAccountingRecord, old: BaselinePublication | None, n
savings_delta=(current.savings if current is not None else 0.0)
- (previous.savings if previous is not None else 0.0),
daily=record.daily,
date=record.turn.turn_at.date().isoformat() if record.turn is not None else None,
router_type=record.turn.router_type if record.turn is not None else None,
)
@ -373,6 +397,7 @@ async def _publish(db: SupportsRawQueries, changes: Sequence[_Change]) -> None:
await db.execute_raw(_UPDATE_SESSIONS, serialized)
if any(change.user_id for change in changes):
await db.execute_raw(_UPDATE_USER_SESSIONS, serialized)
await db.execute_raw(_UPDATE_DAYS, serialized)
for entity, table in DAILY_SPEND_TABLES.items():
if adjustments := tuple(
change.daily.adjustment(target, change.savings_delta, change.request_id)

View file

@ -540,6 +540,18 @@ class SpendLogCleanup:
deadline=deadline,
)
async def _delete_old_autorouter_daily_rows(
self, prisma_client: PrismaClient, cutoff_day: str, deadline: float
) -> TableCleanupResult:
return await self._delete_old_rows_batched(
prisma_client,
cutoff_day,
table_name="LiteLLM_AutoRouterDailySpend",
key_columns=("date", "api_key", "user_id", "router_name", "router_type"),
time_column="date",
deadline=deadline,
)
async def _delete_old_health_check_rows(
self, prisma_client: PrismaClient, cutoff_date: datetime, deadline: float
) -> TableCleanupResult:
@ -623,16 +635,20 @@ class SpendLogCleanup:
except Exception: # noqa: BLE001 # retained observations are retried by the next cleanup job
verbose_proxy_logger.warning("Auto-router baseline retention remains pending")
sessions_result: Final = await self._delete_old_autorouter_session_rows(
prisma_client, session_cutoff, self._group_deadline(deadline, 2)
prisma_client, session_cutoff, self._group_deadline(deadline, 3)
)
verbose_proxy_logger.info("Deleted %s expired auto-router session rollup rows", sessions_result.rows_deleted)
user_sessions_result: Final = await self._delete_old_autorouter_user_session_rows(
prisma_client, session_cutoff, deadline
prisma_client, session_cutoff, self._group_deadline(deadline, 2)
)
verbose_proxy_logger.info(
"Deleted %s expired auto-router user session rollup rows", user_sessions_result.rows_deleted
)
return (sessions_result, user_sessions_result)
days_result: Final = await self._delete_old_autorouter_daily_rows(
prisma_client, session_cutoff.date().isoformat(), deadline
)
verbose_proxy_logger.info("Deleted %s expired auto-router daily rollup rows", days_result.rows_deleted)
return (sessions_result, user_sessions_result, days_result)
async def _clean_health_checks(
self, prisma_client: PrismaClient, retention_seconds: int, deadline: float

View file

@ -1760,7 +1760,7 @@ class LiteLLMProxyRequestSetup:
): # don't override k-v pair sent by request (user request)
data[_metadata_variable_name]["spend_logs_metadata"][key] = value
else:
data[_metadata_variable_name]["spend_logs_metadata"] = key_metadata["spend_logs_metadata"]
data[_metadata_variable_name]["spend_logs_metadata"] = dict(key_metadata["spend_logs_metadata"])
## KEY-LEVEL DISABLE FALLBACKS
if "disable_fallbacks" in key_metadata and isinstance(key_metadata["disable_fallbacks"], bool):
@ -1777,6 +1777,53 @@ class LiteLLMProxyRequestSetup:
)
return data
@staticmethod
def add_team_and_project_level_controls(
user_api_key_dict: UserAPIKeyAuth, metadata: dict[str, object]
) -> dict[str, object]:
team_metadata: Final = user_api_key_dict.team_metadata or MappingProxyType({})
project_metadata: Final = user_api_key_dict.project_metadata or MappingProxyType({})
request_tags: Final = metadata.get("tags")
team_tags: Final = team_metadata.get("tags")
project_tags: Final = project_metadata.get("tags")
disable_global_guardrails: Final = team_metadata.get("disable_global_guardrails")
opted_out_global_guardrails: Final = team_metadata.get("opted_out_global_guardrails")
spend_logs_metadata: Final = LiteLLMProxyRequestSetup._merge_spend_logs_metadata(
team_spend_logs_metadata=team_metadata.get("spend_logs_metadata"),
request_spend_logs_metadata=metadata.get("spend_logs_metadata"),
)
tags: Final = LiteLLMProxyRequestSetup._merge_tags(
request_tags=LiteLLMProxyRequestSetup._merge_tags(
request_tags=request_tags if isinstance(request_tags, list) else None,
tags_to_add=team_tags if isinstance(team_tags, list) else None,
),
tags_to_add=project_tags if isinstance(project_tags, list) else None,
)
controls: Final = (
("tags", tags or None),
("spend_logs_metadata", spend_logs_metadata),
(
"disable_global_guardrails",
disable_global_guardrails if isinstance(disable_global_guardrails, bool) else None,
),
(
"opted_out_global_guardrails",
opted_out_global_guardrails if isinstance(opted_out_global_guardrails, list) else None,
),
)
return {**metadata, **{key: value for key, value in controls if value is not None}}
@staticmethod
def _merge_spend_logs_metadata(
team_spend_logs_metadata: object, request_spend_logs_metadata: object
) -> dict[str, object] | None:
"""Team values as defaults, the request's own values win on the same key. None when neither is a dict"""
team_values: Final = team_spend_logs_metadata if isinstance(team_spend_logs_metadata, dict) else None
request_values: Final = request_spend_logs_metadata if isinstance(request_spend_logs_metadata, dict) else None
if team_values is None and request_values is None:
return None
return {**(team_values or {}), **(request_values or {})}
@staticmethod
def _merge_tags(request_tags: list | None, tags_to_add: list | None) -> list:
"""
@ -2312,38 +2359,12 @@ async def add_litellm_data_to_request(
data=data,
_metadata_variable_name=_metadata_variable_name,
)
## TEAM-LEVEL SPEND LOGS/TAGS
data[_metadata_variable_name] = LiteLLMProxyRequestSetup.add_team_and_project_level_controls(
user_api_key_dict=user_api_key_dict,
metadata=data[_metadata_variable_name],
)
team_metadata: Final = user_api_key_dict.team_metadata or {}
if "tags" in team_metadata and team_metadata["tags"] is not None:
data[_metadata_variable_name]["tags"] = LiteLLMProxyRequestSetup._merge_tags(
request_tags=data[_metadata_variable_name].get("tags"),
tags_to_add=team_metadata["tags"],
)
if "disable_global_guardrails" in team_metadata and isinstance(team_metadata["disable_global_guardrails"], bool):
data[_metadata_variable_name]["disable_global_guardrails"] = team_metadata["disable_global_guardrails"]
if "opted_out_global_guardrails" in team_metadata and isinstance(
team_metadata["opted_out_global_guardrails"], list
):
data[_metadata_variable_name]["opted_out_global_guardrails"] = team_metadata["opted_out_global_guardrails"]
if "spend_logs_metadata" in team_metadata and isinstance(team_metadata["spend_logs_metadata"], dict):
if "spend_logs_metadata" in data[_metadata_variable_name] and isinstance(
data[_metadata_variable_name]["spend_logs_metadata"], dict
):
for key, value in team_metadata["spend_logs_metadata"].items():
if (
key not in data[_metadata_variable_name]["spend_logs_metadata"]
): # don't override k-v pair sent by request (user request)
data[_metadata_variable_name]["spend_logs_metadata"][key] = value
else:
data[_metadata_variable_name]["spend_logs_metadata"] = team_metadata["spend_logs_metadata"]
## PROJECT-LEVEL TAGS
project_metadata: Final = user_api_key_dict.project_metadata or {}
if "tags" in project_metadata and project_metadata["tags"] is not None:
data[_metadata_variable_name]["tags"] = LiteLLMProxyRequestSetup._merge_tags(
request_tags=data[_metadata_variable_name].get("tags"),
tags_to_add=project_metadata["tags"],
)
# inherited_tags: every tag key/team/project policy contributed, read
# directly from those three sources rather than snapshotted off the shared

View file

@ -5,6 +5,8 @@ POST /auto_router/test_routing - Route one request through an unsaved complexity
POST /auto_router/validate_complexity_router_config - Dry-run the complexity-router write gate without saving
"""
import asyncio
import math
from collections.abc import Mapping, Sequence
from datetime import datetime, timedelta, timezone
from itertools import chain, groupby
@ -40,6 +42,7 @@ from litellm.proxy.litellm_pre_call_utils import (
refresh_proxy_server_request_body_snapshot,
)
from litellm.proxy.management.teams.access import is_team_admin
from litellm.proxy.management_endpoints.common_daily_activity import daily_activity_scope
from litellm.proxy.management_helpers.auto_router_permissions import (
authorize_member_auto_router_dependencies,
authorize_member_auto_router_team,
@ -47,6 +50,7 @@ from litellm.proxy.management_helpers.auto_router_permissions import (
)
from litellm.repositories.autorouter_session_repository import AutoRouterSessionRepository
from litellm.repositories.base_repository import SupportsModelDump
from litellm.repositories.daily_activity_sql import build_where_clause
from litellm.repositories.team_repository import TeamRepository
from litellm.router_strategy.complexity_router import ComplexityRouter
from litellm.router_utils.auto_router_model_naming import (
@ -316,7 +320,7 @@ async def _authorize_models_this_test_can_call(
its calls through the proxy. Team and member budgets are already enforced on every route.
"""
models: Final = _models_this_test_can_call(config)
if not models and config.classifier_type != "jev":
if not models and config.classifier_type != "oss_classifier":
return
from litellm.proxy.proxy_server import proxy_logging_obj
@ -342,9 +346,9 @@ async def _authorize_models_this_test_can_call(
code=status.HTTP_400_BAD_REQUEST,
) from e
if config.classifier_type == "jev" and user_api_key_dict.budget_throttle_pct is not None:
if config.classifier_type == "oss_classifier" and user_api_key_dict.budget_throttle_pct is not None:
raise ProxyException(
message="Budget has been exceeded! JEV Test Routing requires available budget.",
message="Budget has been exceeded! OSS Classifier Test Routing requires available budget.",
type=ProxyErrorTypes.budget_exceeded,
param=None,
code=status.HTTP_400_BAD_REQUEST,
@ -616,34 +620,37 @@ async def preview_auto_router_routing(
class _SessionAggRow(BaseModel):
"""One router's window: session shape from overlapping sessions, money from the selected days."""
router_name: str
router_type: str
tier_turns: Mapping[str, int]
sessions: int
turns: int
unordered_turns: int
covered_turns: int
cache_hits: int
same_model_turns: int
same_model_hits: int
first_visit_turns: int
first_visit_hits: int
return_turns: int
return_hits: int
return_expired_misses: int
return_within_ttl_misses: int
ttl_5m_turns: int
ttl_1h_turns: int
total_tokens: int
spend: float
saved_spend: float
tier_turns: Mapping[str, int] = MappingProxyType({})
sessions: int = 0
session_turns: int = 0
unordered_turns: int = 0
covered_turns: int = 0
cache_hits: int = 0
same_model_turns: int = 0
same_model_hits: int = 0
first_visit_turns: int = 0
first_visit_hits: int = 0
return_turns: int = 0
return_hits: int = 0
return_expired_misses: int = 0
return_within_ttl_misses: int = 0
ttl_5m_turns: int = 0
ttl_1h_turns: int = 0
total_tokens: int = 0
session_seconds: float = 0.0
turns: int = 0
spend: float = 0.0
saved_spend: float = 0.0
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
classifier_cost: float
classifier_cost_recorded_turns: int
session_seconds: float
classifier_cost: float = 0.0
classifier_cost_recorded_turns: int = 0
_SESSION_AGG_ROWS: Final = TypeAdapter(list[_SessionAggRow])
@ -692,6 +699,13 @@ def _compared_row(row: _SessionAggRow) -> _SessionAggRow:
)
def _per_session(row: _SessionAggRow, total: float) -> float | None:
"""Unknown, not zero, when routed requests have no session rows of their own to average over."""
if row.sessions:
return total / row.sessions
return None if row.turns else 0.0
def _benchmark_totals(row: _SessionAggRow) -> AutoRouterBenchmarkTotals:
return_misses: Final = row.return_turns - row.return_hits
saved_spend, baseline_spend = _savings_cohort(
@ -701,9 +715,9 @@ def _benchmark_totals(row: _SessionAggRow) -> AutoRouterBenchmarkTotals:
return AutoRouterBenchmarkTotals(
sessions=sessions,
turns=row.turns,
avg_turns_per_session=row.turns / sessions if sessions else 0.0,
avg_session_seconds=row.session_seconds / sessions if sessions else 0.0,
avg_tokens_per_session=row.total_tokens / sessions if sessions else 0.0,
avg_turns_per_session=_per_session(row, row.session_turns),
avg_session_seconds=_per_session(row, row.session_seconds),
avg_tokens_per_session=_per_session(row, row.total_tokens),
spend=row.spend,
savings_estimated_turns=row.savings_estimated_turns,
savings_estimated_actual_spend=row.savings_estimated_actual_spend,
@ -712,9 +726,8 @@ def _benchmark_totals(row: _SessionAggRow) -> AutoRouterBenchmarkTotals:
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=(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),
coverage_pct=_pct(row.covered_turns, row.session_turns),
hit_rate_pct=_pct(row.cache_hits, row.covered_turns),
same_model=_cache_bucket(row.same_model_turns, row.same_model_hits),
first_visit=_cache_bucket(row.first_visit_turns, row.first_visit_hits),
@ -748,7 +761,6 @@ def _benchmark_group(row: _SessionAggRow) -> AutoRouterBenchmarkGroup:
classifier_cost=totals.classifier_cost,
baseline_spend=totals.baseline_spend,
saved_pct=totals.saved_pct,
saved_per_session=totals.saved_per_session,
cache=totals.cache,
)
@ -759,6 +771,7 @@ def _summed_agg_row(rows: Sequence[_SessionAggRow]) -> _SessionAggRow:
router_type="",
tier_turns=MappingProxyType({}),
sessions=sum(row.sessions for row in rows),
session_turns=sum(row.session_turns for row in rows),
turns=sum(row.turns for row in rows),
unordered_turns=sum(row.unordered_turns for row in rows),
covered_turns=sum(row.covered_turns for row in rows),
@ -790,6 +803,49 @@ def _summed_agg_row(rows: Sequence[_SessionAggRow]) -> _SessionAggRow:
)
async def _recorded_autorouter_savings(
prisma_client: "PrismaClient", start_day: str, end_day: str, api_key: str | None, user_id: str | None
) -> float:
"""The selected days' auto-router savings exactly as the Overall view sums them: same table, same filters."""
where, params = build_where_clause(
daily_activity_scope(
table="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id=user_id,
exclude_entity_ids=None,
api_key=api_key,
start_date=start_day,
end_date=end_day,
model=None,
timezone_offset_minutes=None,
)
)
rows: Final = await _query_raw(
prisma_client,
f'SELECT COALESCE(SUM(autorouter_savings_spend), 0)::float8 AS saved FROM "LiteLLM_DailyUserSpend" WHERE {where}',
*params,
)
return float(rows[0]["saved"]) if rows else 0.0
def _with_recorded_savings(
totals: AutoRouterBenchmarkTotals, rows: Sequence[_SessionAggRow], recorded: float
) -> AutoRouterBenchmarkTotals:
"""The headline is the recorded total. Savings outside the compared routers void the cost comparison,
and the part no router's day rows account for is reported as unattributed."""
if math.isclose(recorded, totals.saved_spend or 0.0, abs_tol=1e-9):
return totals
unattributed: Final = recorded - sum(row.saved_spend for row in rows)
return totals.model_copy(
update={
"saved_spend": recorded,
"unattributed_saved_spend": None if math.isclose(unattributed, 0.0, abs_tol=1e-9) else unattributed,
"baseline_spend": None,
"saved_pct": None,
}
)
def _strategy_router_key(deployment: object) -> tuple[str, str] | None:
"""``(model_name, kind)`` for a deployment whose routing the session rollup records.
@ -849,7 +905,7 @@ async def get_auto_router_benchmarks(
str | None, Query(description="YYYY-MM-DD UTC, inclusive (defaults to 30 days before end_date)")
] = None,
end_date: Annotated[str | None, Query(description="YYYY-MM-DD UTC, inclusive (defaults to today)")] = None,
api_key: Annotated[str | None, Query(description="Filter to one virtual key token hash")] = None,
api_key: Annotated[str | None, Query(min_length=1, description="Filter to one virtual key token hash")] = None,
user_id: Annotated[
str | None, Query(min_length=1, description="Filter to one canonical internal user recorded on each turn")
] = None,
@ -860,9 +916,10 @@ async def get_auto_router_benchmarks(
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
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
internal user when written; older key-only history remains outside user views. Money counts
only requests on the selected UTC days, and the all-router savings headline is the same daily
total the Overall view reads. Session shape and caching cover every session that overlaps the
window, whole. Overall hit rate is over telemetry-bearing turns; each bucket's hit rate is
over that bucket's turns.
The rollup supplies the measures, never the list. Which routers appear comes from the
@ -885,24 +942,35 @@ async def get_auto_router_benchmarks(
if end_day < start_day:
raise HTTPException(status_code=400, detail="end_date must not be earlier than start_date")
raw_rows: Final = await _query_raw(
prisma_client,
AUTOROUTER_BENCHMARKS_SQL,
start_day.isoformat(),
(end_day + timedelta(days=1)).isoformat(),
api_key,
user_id,
first_day: Final = start_day.strftime("%Y-%m-%d")
last_day: Final = end_day.strftime("%Y-%m-%d")
raw_rows, recorded = await asyncio.gather(
_query_raw(
prisma_client,
AUTOROUTER_BENCHMARKS_SQL,
start_day.isoformat(),
(end_day + timedelta(days=1)).isoformat(),
api_key,
user_id,
first_day,
last_day,
),
_recorded_autorouter_savings(prisma_client, first_day, last_day, api_key, user_id),
)
rows: Final = tuple(_compared_row(row) for row in _SESSION_AGG_ROWS.validate_python(raw_rows or ()))
totals: Final = _with_recorded_savings(_benchmark_totals(_summed_agg_row(rows)), rows, recorded)
unattributed: Final = MappingProxyType(
{"baseline_spend": None, "saved_pct": None} if totals.unattributed_saved_spend is not None else {}
)
groups: Final = (
*(_benchmark_group(row) for row in rows),
*(_benchmark_group(row).model_copy(update=unattributed) for row in rows),
*_idle_router_groups(llm_router, frozenset((row.router_name, row.router_type) for row in rows)),
)
return AutoRouterBenchmarksResponse(
start_date=start_day.strftime("%Y-%m-%d"),
end_date=end_day.strftime("%Y-%m-%d"),
start_date=first_day,
end_date=last_day,
routers_in_scope=len(groups),
totals=_benchmark_totals(_summed_agg_row(rows)),
totals=totals,
groups=groups,
)

View file

@ -39,6 +39,7 @@ from litellm.litellm_core_utils.ptu_pricing import (
SEARCH_CONTEXT_SIZES,
ptu_config_error,
)
from litellm.litellm_core_utils.sensitive_data_masker import redact_credentials_in_payload
from litellm.proxy._types import (
BlockModelRequest,
CommonProxyErrors,
@ -115,6 +116,7 @@ from litellm.router_strategy.complexity_router import (
normalize_classification_examples,
normalize_classification_prompt,
)
from litellm.router_strategy.complexity_router.config import resolve_complexity_router_config_write
from litellm.router_utils.auto_router_model_naming import (
GATED_AUTO_ROUTER_CAPABILITIES,
STRATEGY_ROUTER_PARAM_FIELDS,
@ -187,6 +189,25 @@ class _ProxyModelRow(Protocol):
def model_dump_json(self, *, exclude_none: bool = False) -> str: ...
def _model_write_response(
row: _ProxyModelRow, member_write: MemberAutoRouterWrite | None
) -> _ProxyModelRow | Mapping[str, object]:
if member_write is None:
return row
payload: Final = TypeAdapter(dict[str, object]).validate_json(row.model_dump_json())
stored_params: Final = payload.get("litellm_params")
params: Final = (
TypeAdapter(dict[str, object]).validate_json(stored_params)
if isinstance(stored_params, str)
else TypeAdapter(dict[str, object]).validate_python(stored_params)
)
redacted: Final = redact_credentials_in_payload(params)
return {
**payload,
"litellm_params": json.dumps(redacted) if isinstance(stored_params, str) else redacted,
}
class _ProxyModelTable(Protocol):
def find_unique(self, *, where: Mapping[str, object]) -> Awaitable[BaseModel | None]: ...
@ -407,34 +428,13 @@ WHERE model_id <> $1
def _effective_complexity_router_config(
incoming_params: GenericLiteLLMParams | None, existing_params: GenericLiteLLMParams | None
) -> object:
) -> Mapping[str, object] | None:
incoming: Final = None if incoming_params is None else incoming_params.complexity_router_config
existing: Final = None if existing_params is None else existing_params.complexity_router_config
if incoming is None:
return existing
if existing is None or incoming.get("classifier_type") != "jev" or existing.get("classifier_type") != "jev":
return incoming
incoming_jev: Final[object] = incoming.get("jev_classifier_config")
existing_jev: Final[object] = existing.get("jev_classifier_config")
if not isinstance(incoming_jev, Mapping) or not isinstance(existing_jev, Mapping):
return incoming
supplied: Final = TypeAdapter(dict[str, object]).validate_python(incoming_jev)
stored: Final = TypeAdapter(dict[str, object]).validate_python(existing_jev)
same_base: Final = "api_base" not in supplied or supplied["api_base"] == stored.get("api_base")
transport: Final = MappingProxyType(
{
key: value
for key, value in stored.items()
if key in ("api_key", "api_base") and (key != "api_key" or same_base)
}
)
return {
**incoming,
"jev_classifier_config": {
**transport,
**supplied,
},
}
config_adapter: Final[TypeAdapter[Mapping[str, object] | None]] = TypeAdapter(Mapping[str, object] | None)
return resolve_complexity_router_config_write(
config_adapter.validate_python(incoming), config_adapter.validate_python(existing)
).effective
def _effective_model(
@ -1304,7 +1304,7 @@ async def patch_model(
live_after=reload_outcome.live_after,
)
return updated_model
return _model_write_response(updated_model, member_write)
except Exception as e:
verbose_proxy_logger.exception("Error in patch_model: %s", e)
@ -1501,10 +1501,18 @@ async def _add_model_to_db(
slot: AbstractAsyncContextManager[_ProxyModelTable] | None = None,
) -> "_ProxyModelRow | LiteLLM_ProxyModelTable":
# encrypt litellm params #
_litellm_params_dict: Final = model_params.litellm_params.dict(exclude_none=True)
_litellm_params_dict: Final = TypeAdapter(dict[str, object]).validate_python(
model_params.litellm_params.model_dump(exclude_none=True)
)
if "complexity_router_config" in _litellm_params_dict:
_litellm_params_dict["complexity_router_config"] = _effective_complexity_router_config(
model_params.litellm_params, None
)
_original_litellm_model_name: Final = model_params.litellm_params.model
for k, v in _litellm_params_dict.items():
encrypted_value = encrypt_value_helper(value=v, new_encryption_key=new_encryption_key)
encrypted_value = (
encrypt_value_helper(value=v, new_encryption_key=new_encryption_key) if isinstance(v, str) else v
)
model_params.litellm_params[k] = encrypted_value
_data: Final[dict] = {
"model_id": model_params.model_info.id,
@ -2536,7 +2544,7 @@ async def add_new_model(
live_after=reload_outcome.live_after,
)
return model_response
return _model_write_response(model_response, member_write)
except Exception as e:
verbose_proxy_logger.exception("litellm.proxy.proxy_server.add_new_model(): Exception occured - %s", e)
@ -2760,7 +2768,7 @@ async def update_model(
live_after=reload_outcome.live_after,
)
return model_response
return None if model_response is None else _model_write_response(model_response, member_write)
except Exception as e:
verbose_proxy_logger.exception("litellm.proxy.proxy_server.update_model(): Exception occured - %s", e)
if isinstance(e, HTTPException):

View file

@ -33,6 +33,10 @@ from litellm.repositories.prisma_protocols import DatabaseClient
from litellm.repositories.project_repository import ProjectRepository
from litellm.repositories.table_repositories import TeamMembershipRepository
from litellm.router import Router
from litellm.router_strategy.complexity_router.config import (
ComplexityRouterConfigWrite,
resolve_complexity_router_config_write,
)
from litellm.router_utils.auto_router_model_naming import classify_strategy_router_model, strategy_router_dependencies
from litellm.types.management_endpoints.auto_router_endpoints import RequestComplexityRouterConfig
from litellm.types.router import Deployment, updateDeployment
@ -65,12 +69,12 @@ class _MemberRouterGenerationParams(BaseModel):
stop: str | tuple[str, ...] | None = None
class _MemberJevClassifierConfig(BaseModel):
"""The Jev classifier settings a team member may set. Credentials stay the proxy's own: a member-chosen
api_base would receive the proxy's TYPESAFE_API_KEY, and a member-chosen api_key would be sent from the proxy."""
class _MemberOpenSourceClassifierConfig(BaseModel):
"""Classifier settings a team member may set while the gateway owns the connection."""
model_config = ConfigDict(extra="forbid")
provider: Literal["jev", "laya"] = "jev"
model: str
api_key: None = None
api_base: None = None
@ -123,14 +127,21 @@ def authorize_member_auto_router_team(
def validate_member_auto_router_config(config: Mapping[str, object]) -> RequestComplexityRouterConfig:
return _validate_member_auto_router_config_write(resolve_complexity_router_config_write(config, None))
def _validate_member_auto_router_config_write(write: ComplexityRouterConfigWrite) -> RequestComplexityRouterConfig:
if write.effective is None:
raise HTTPException(status_code=400, detail="A complexity_router_config is required.")
try:
validated: Final = _MemberComplexityRouterConfig.model_validate(config)
for entries in validated.tier_model_configs.values():
for entry in entries:
_MemberRouterGenerationParams.model_validate(entry.litellm_params)
if validated.jev_classifier_config is not None:
_MemberJevClassifierConfig.model_validate(validated.jev_classifier_config.model_dump())
return validated
if write.submitted is not None:
validated: Final = _MemberComplexityRouterConfig.model_validate(write.submitted)
for entries in validated.tier_model_configs.values():
for entry in entries:
_MemberRouterGenerationParams.model_validate(entry.litellm_params)
if validated.opensource_classifier_config is not None:
_MemberOpenSourceClassifierConfig.model_validate(validated.opensource_classifier_config.model_dump())
return RequestComplexityRouterConfig.model_validate(write.effective)
except ValidationError as exc:
location: Final = ".".join(str(part) for part in exc.errors()[0]["loc"])
raise HTTPException(status_code=400, detail=f"Invalid member auto-router configuration at {location}.") from exc
@ -332,16 +343,15 @@ async def authorize_member_auto_router_write(
if existing is not None and incoming.model_name not in (None, public_name, existing.model_name):
raise HTTPException(status_code=403, detail="Team members cannot rename an auto router.")
supplied_config: Final = _RouterConfigSource.model_validate(params.model_dump()).complexity_router_config
raw_config: Final = (
supplied_config
if supplied_config is not None
else _RouterConfigSource.model_validate(existing.litellm_params.model_dump()).complexity_router_config
stored_config: Final = (
_RouterConfigSource.model_validate(existing.litellm_params.model_dump()).complexity_router_config
if existing is not None
else None
)
if raw_config is None:
raise HTTPException(status_code=400, detail="A complexity_router_config is required.")
config: Final = validate_member_auto_router_config(raw_config)
resolved_config: Final = resolve_complexity_router_config_write(supplied_config, stored_config)
if resolved_config.supplied_connection_fields:
raise HTTPException(status_code=403, detail="Team members cannot change classifier connections.")
config: Final = _validate_member_auto_router_config_write(resolved_config)
stored_default: Final = existing.litellm_params.complexity_router_default_model if existing is not None else None
default_model: Final = (
params.complexity_router_default_model

View file

@ -24,6 +24,7 @@ from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, cast
import httpx
from fastapi import APIRouter, Depends, HTTPException, Request, Response, WebSocket
from fastapi.responses import StreamingResponse
from pydantic import ConfigDict, TypeAdapter
from starlette.websockets import WebSocketState
from typing_extensions import ReadOnly, TypedDict
@ -57,6 +58,7 @@ from litellm.llms.deepgram.common_utils import (
deepgram_listen_websocket_target,
)
from litellm.llms.fal_ai.cost_calculator import fal_ai_passthrough_cost, fal_ai_queue_base
from litellm.llms.laya.common_utils import laya_connection, validate_laya_request
from litellm.llms.nvidia_nim.passthrough.transformation import nvidia_nim_model_group_in_path
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
@ -636,6 +638,47 @@ async def typesafe_proxy_route(
return await endpoint_func(request, fastapi_response, user_api_key_dict)
@router.post(
"/laya/v1/systemone",
tags=["Laya Pass-through", "pass-through"],
)
async def laya_proxy_route(
request: Request,
fastapi_response: Response,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
) -> Response:
body: Final = TypeAdapter(dict[str, object]).validate_python(await _read_request_body(request))
try:
_ = validate_laya_request(body)
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
try:
connection: Final = laya_connection()
except ValueError as exc:
raise HTTPException(
status_code=503, detail="Laya server is not configured correctly; check LAYA_API_BASE"
) from exc
base_url: Final = httpx.URL(connection.api_base)
updated_url: Final = base_url.copy_with(
path=HttpPassThroughEndpointHelpers.join_base_and_endpoint_path(base_url, "/v1/systemone"),
)
authorization: Final[Mapping[str, str]] = (
MappingProxyType({"Authorization": f"Bearer {connection.api_key}"})
if connection.api_key
else MappingProxyType({})
)
endpoint_func: Final = create_pass_through_route(
endpoint="v1/systemone",
target=str(updated_url),
custom_headers=MappingProxyType({**authorization, "Content-Type": "application/json"}),
custom_llm_provider="laya",
is_streaming_request=False,
)
return TypeAdapter(Response, config=ConfigDict(arbitrary_types_allowed=True)).validate_python(
await endpoint_func(request, fastapi_response, user_api_key_dict)
)
@router.api_route(
"/openrouter/{endpoint:path}",
methods=["GET", "POST", "PUT", "DELETE", "PATCH"],

View file

@ -34,9 +34,7 @@ def is_collection_route(url_route: str, collection_suffix: str) -> bool:
def request_tags_from_metadata(request_metadata: Mapping[str, object]) -> tuple[str, ...] | None:
"""Tags for the batch-cost spend row: the request's own tags when it sent any,
otherwise the key's tags, which auth exposes as user_api_key_auth_metadata (a
tagged key does not put its tags in the top-level metadata "tags" on the
passthrough path)
otherwise the key's tags, which auth exposes as user_api_key_auth_metadata
"""
tags: Final = _sanitized_str_tuple(request_metadata.get("tags"))
if tags:

View file

@ -10,6 +10,7 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.litellm_core_utils.litellm_logging import (
get_standard_logging_object_payload, # pyright: ignore[reportUnknownVariableType] # legacy helper has an untyped signature
)
from litellm.llms.laya.common_utils import laya_response_model
from litellm.proxy._types import PassThroughEndpointLoggingTypedDict
from litellm.types.utils import ModelResponse, StandardPassThroughResponseObject, Usage
@ -69,9 +70,11 @@ class TypeSafePassthroughLoggingHandler:
**kwargs: object,
) -> PassThroughEndpointLoggingTypedDict:
response: Final = _parse_typesafe_response(response_body)
response_model: Final = response.model
request_model_value: Final = request_body.get("model")
request_model: Final = request_model_value if isinstance(request_model_value, str) else None
response_model: Final = (
laya_response_model(response_body, request_model) if custom_llm_provider == "laya" else response.model
)
logged_model: Final = response_model or request_model or "unknown"
model_name: Final = f"{custom_llm_provider}/{logged_model}"
usage: Final = response.usage or _TypeSafeUsage()

View file

@ -26,6 +26,7 @@ from fastapi import (
status,
)
from fastapi.responses import StreamingResponse
from pydantic import TypeAdapter
from starlette.datastructures import UploadFile as StarletteUploadFile
from starlette.websockets import WebSocketState
from websockets.asyncio.client import connect
@ -64,6 +65,7 @@ from litellm.llms.base_llm.managed_resources.utils import (
resolve_passthrough_managed_id_provider,
)
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.llms.laya.common_utils import validate_laya_request
from litellm.passthrough import BasePassthroughUtils
from litellm.proxy._types import (
ConfigFieldInfo,
@ -74,7 +76,11 @@ from litellm.proxy._types import (
ProxyException,
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_utils import request_dispatched_to_pass_through_endpoint
from litellm.proxy.auth.auth_utils import (
get_model_from_request,
get_request_route,
request_dispatched_to_pass_through_endpoint,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
@ -100,6 +106,8 @@ from litellm.proxy.common_utils.sse_keepalive import (
from litellm.proxy.litellm_pre_call_utils import (
LiteLLMProxyRequestSetup,
_get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage] # shared proxy helper, same import style as _read_request_body above
_key_or_team_allows_client_pricing_override, # pyright: ignore[reportPrivateUsage] # reuse the proxy's pricing trust policy
_strip_client_pricing_overrides, # pyright: ignore[reportPrivateUsage] # sanitize before trusted hooks add guardrail costs
)
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
from litellm.proxy.utils import normalize_route_for_root_path
@ -585,7 +593,18 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
"""
Filter out litellm params from the request body
"""
from litellm.proxy.proxy_server import llm_router
_parsed_body = _parsed_body or {}
managed_model: Final = get_model_from_request(
request_data=_parsed_body,
route=get_request_route(request),
request_headers=request.headers,
request_query_params=request.query_params,
llm_router=llm_router,
request=request,
team_id=user_api_key_dict.team_id,
)
litellm_keys_in_body: Final = MappingProxyType(
{k: _parsed_body.pop(k) for k in types_utils.all_litellm_params if k in _parsed_body}
@ -600,11 +619,11 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
litellm_metadata: Final = litellm_keys_in_body.get("litellm_metadata")
metadata: Final = litellm_keys_in_body.get("metadata")
if litellm_metadata:
_metadata.update(litellm_metadata)
if metadata:
_metadata.update(metadata)
for client_metadata in (litellm_metadata, metadata):
if isinstance(client_metadata, dict):
_metadata.update({k: v for k, v in client_metadata.items() if not k.startswith("user_api_key_")})
_metadata = _apply_key_team_project_controls(user_api_key_dict=user_api_key_dict, metadata=_metadata)
_metadata = _update_metadata_with_tags_in_header(
request=request,
metadata=_metadata,
@ -631,10 +650,19 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
# would attribute it to a budget the operator scoped to a LiteLLM model that
# merely shares the name.
if not request_dispatched_to_pass_through_endpoint(request):
_metadata["model_group"] = managed_model if isinstance(managed_model, str) else None
_metadata["user_api_key_model_max_budget"] = user_api_key_dict.model_max_budget
_metadata["user_api_key_team_model_max_budget"] = user_api_key_dict.team_model_max_budget
_metadata["user_api_key_user_model_max_budget"] = user_api_key_dict.user_model_max_budget
_metadata["user_api_key_end_user_model_max_budget"] = user_api_key_dict.end_user_model_max_budget
else:
for field in (
"user_api_key_model_max_budget",
"user_api_key_team_model_max_budget",
"user_api_key_user_model_max_budget",
"user_api_key_end_user_model_max_budget",
):
_metadata.pop(field, None)
_metadata.update(
LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_dict)
)
@ -1131,6 +1159,15 @@ async def pass_through_request(
_parsed_body,
)
if not _key_or_team_allows_client_pricing_override(user_api_key_dict):
pricing_body: Final = TypeAdapter(dict[str, object]).validate_python(_parsed_body)
_strip_client_pricing_overrides(pricing_body)
_parsed_body = pricing_body
if custom_llm_provider == "laya":
laya_request: Final = TypeAdapter(Mapping[str, object]).validate_python(_parsed_body)
checkpoint: Final = validate_laya_request(laya_request)
_parsed_body["model"] = f"laya/{checkpoint}"
### COLLECT GUARDRAILS FOR PASSTHROUGH ENDPOINT ###
# Passthrough endpoints are opt-in only for guardrails
# When enabled, collect guardrails from org/team/key levels + passthrough-specific
@ -1186,6 +1223,17 @@ async def pass_through_request(
call_type="pass_through_endpoint",
endpoint_type=endpoint_type,
)
if custom_llm_provider == "laya":
hook_body: Final = TypeAdapter(dict[str, object]).validate_python(_parsed_body)
hook_model: Final = hook_body.get("model")
laya_body: Final = MappingProxyType(
{
**hook_body,
"model": hook_model.removeprefix("laya/") if isinstance(hook_model, str) else hook_model,
}
)
_ = validate_laya_request(laya_body)
_parsed_body = TypeAdapter(dict[str, object]).validate_python(laya_body)
resolved_timeout: Final = resolve_pass_through_request_timeout(timeout)
async_client_obj: Final = get_async_httpx_client(
llm_provider=httpxSpecialProvider.PassThroughEndpoint,
@ -1886,6 +1934,20 @@ async def pass_through_request(
)
def _apply_key_team_project_controls(
user_api_key_dict: UserAPIKeyAuth, metadata: dict[str, object]
) -> dict[str, object]:
data: Final = LiteLLMProxyRequestSetup.add_key_level_controls(
key_metadata=user_api_key_dict.metadata,
data={"metadata": metadata},
_metadata_variable_name="metadata",
)
return LiteLLMProxyRequestSetup.add_team_and_project_level_controls(
user_api_key_dict=user_api_key_dict,
metadata=data["metadata"],
)
def _update_metadata_with_tags_in_header(request: Request, metadata: dict) -> dict:
"""
If tags are in the request headers, add them to the metadata
@ -1906,9 +1968,10 @@ def _update_metadata_with_tags_in_header(request: Request, metadata: dict) -> di
# Only add tags key if there are tags to add
if tags_to_add:
if "tags" not in metadata:
metadata["tags"] = []
metadata["tags"].extend(tags_to_add)
metadata["tags"] = LiteLLMProxyRequestSetup._merge_tags(
request_tags=metadata.get("tags"),
tags_to_add=tags_to_add,
)
return metadata
@ -2389,7 +2452,9 @@ async def websocket_passthrough_request(
# with the existing _init_kwargs_for_pass_through_endpoint function
class DummyRequest:
def __init__(self, url: str, method: str = "WEBSOCKET", headers: dict | None = None):
self.url = url
self.url = httpx.URL(url)
self.scope = websocket.scope
self.query_params = websocket.query_params
self.method = method
self.headers = headers or {}

View file

@ -334,8 +334,10 @@ class PassThroughEndpointLogging:
)
standard_logging_response_object = transcribe_handler_result["result"] # rebind-ok: elif-chain
kwargs = transcribe_handler_result["kwargs"] # rebind-ok: elif-chain contract
elif self.is_typesafe_route(custom_llm_provider) or self.is_openrouter_decisions_route(
url_route, custom_llm_provider
elif (
self.is_typesafe_route(custom_llm_provider)
or custom_llm_provider == "laya"
or self.is_openrouter_decisions_route(url_route, custom_llm_provider)
):
from .llm_provider_handlers.typesafe_passthrough_logging_handler import (
TypeSafePassthroughLoggingHandler,

View file

@ -1744,6 +1744,27 @@ model LiteLLM_AutoRouterUserSession {
@@index([user_id, last_turn_at], map: "idx_autorouter_user_session_user_last_turn")
}
// Auto-routed requests per UTC request day and router: the selected-day money behind the
// auto-router usage view. Written in the same statement as the session rollup, so a day row
// and its session row never disagree; corrected in the same transaction as late baselines.
model LiteLLM_AutoRouterDailySpend {
date String
api_key String
user_id String
router_name String
router_type String
turns Int @default(0)
spend Float @default(0)
saved_spend Float @default(0)
savings_estimated_turns Int @default(0)
savings_estimated_actual_spend Float @default(0)
savings_estimated_saved_spend Float @default(0)
classifier_cost Float @default(0)
classifier_cost_recorded_turns Int @default(0)
@@id([date, api_key, user_id, router_name, router_type])
}
// Shadow eval: evaluation of an auto-router against one or more keys' live traffic, in
// either direction. forward duplicates the requests the keys did not route through the
// router through it, answering whether they should adopt it; reverse duplicates the

View file

@ -8,6 +8,7 @@ from datetime import datetime
from types import TracebackType
from typing import TYPE_CHECKING, Final, Protocol
from litellm.constants import UI_SESSION_TOKEN_TEAM_ID
from litellm.models.verification_token import (
LiteLLM_VerificationToken,
)
@ -123,6 +124,37 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]):
records: Final[Sequence[PrismaVerificationToken]] = await self.table.find_many(where={"user_id": user_id})
return self._to_model_list(records)
async def find_newest_reusable_llm_api_key(
self, user_id: str, team_id: str | None
) -> LiteLLM_VerificationToken | None:
records: Final[Sequence[PrismaVerificationToken]] = await self.table.find_many(
where={
"user_id": user_id,
"team_id": team_id,
"expires": None,
"AND": [
{"OR": [{"blocked": False}, {"blocked": None}]},
{
"OR": [
{"team_id": None},
{"team_id": {"not": UI_SESSION_TOKEN_TEAM_ID}},
]
},
{
"OR": [
{"allowed_routes": {"is_empty": True}},
{"allowed_routes": {"has": "llm_api_routes"}},
]
},
],
},
order={"created_at": "desc"},
)
return next(
(key for key in self._to_model_list(records) if key.metadata.get("auto_registered") is not True),
None,
)
async def find_by_team_id(self, team_id: str) -> list[LiteLLM_VerificationToken]:
"""Find all tokens belonging to a team."""
records: Final[Sequence[PrismaVerificationToken]] = await self.table.find_many(where={"team_id": team_id})

View file

@ -108,7 +108,7 @@ from .config import (
ComplexityRouterConfig,
ComplexityTier,
CustomDimension,
JevClassifierConfig,
OpenSourceClassifierConfig,
TierDefinition,
)
from .jev_classifier import (
@ -1308,10 +1308,22 @@ class ComplexityRouter(CustomLogger):
"""
@staticmethod
def _build_jev_client(config: JevClassifierConfig) -> JevClassifierClient:
def _build_jev_client(config: OpenSourceClassifierConfig) -> JevClassifierClient:
if config.provider == "laya":
from litellm.llms.laya.common_utils import laya_connection
connection: Final = laya_connection(config.api_base, config.api_key)
return HttpJevClassifierClient(
api_key=connection.api_key,
api_base=connection.api_base,
http_client=get_async_httpx_client(httpxSpecialProvider.PassThroughEndpoint),
provider="laya",
)
api_key: Final = config.api_key or get_secret_str("TYPESAFE_API_KEY")
if not api_key:
raise ValueError("jev_classifier_config.api_key or TYPESAFE_API_KEY is required for classifier_type 'jev'")
raise ValueError(
"opensource_classifier_config.api_key or TYPESAFE_API_KEY is required for classifier_type 'oss_classifier'"
)
api_base: Final = config.api_base or get_secret_str("TYPESAFE_API_BASE") or "https://api.typesafe.ai"
return HttpJevClassifierClient(
api_key=api_key,
@ -1354,12 +1366,12 @@ class ComplexityRouter(CustomLogger):
if default_model:
self.config.default_model = default_model
jev_config: Final = self.config.jev_classifier_config
jev_config: Final = self.config.opensource_classifier_config
self._jev_client: JevClassifierClient | None = (
jev_client
if jev_client is not None
else self._build_jev_client(jev_config)
if self.config.classifier_type == "jev" and jev_config is not None
if self.config.classifier_type == "oss_classifier" and jev_config is not None
else None
)
@ -1459,7 +1471,11 @@ class ComplexityRouter(CustomLogger):
and self.config.classifier_llm_config.circuit_breaker_enabled
)
else jev_config.circuit_breaker_cooldown_seconds
if (self.config.classifier_type == "jev" and jev_config is not None and jev_config.circuit_breaker_enabled)
if (
self.config.classifier_type == "oss_classifier"
and jev_config is not None
and jev_config.circuit_breaker_enabled
)
else None
)
self._classifier_circuit_breaker: _ClassifierCircuitBreaker | None = (
@ -1909,7 +1925,7 @@ class ComplexityRouter(CustomLogger):
return self._classify_with_heuristic_v2(prompt)
if self.config.classifier_type == "custom":
return await self._classify_with_plugin(prompt, system_prompt, request_kwargs, raw_messages)
if self.config.classifier_type == "jev":
if self.config.classifier_type == "oss_classifier":
return await self._jev_classifier_outcome(prompt, system_prompt, request_kwargs, messages)
if self.config.classifier_type in ("heuristic_first", "hybrid") and _encrypted_classifier_task(
request_kwargs, self._reminder_markers_for_request(request_kwargs or EMPTY_MAPPING)
@ -2161,7 +2177,7 @@ class ComplexityRouter(CustomLogger):
request_kwargs: Mapping[str, object] | None,
messages: Sequence[Mapping[str, object]] | None,
) -> ClassificationOutcome:
config: Final = self.config.jev_classifier_config
config: Final = self.config.opensource_classifier_config
client: Final = self._jev_client
if config is None or client is None:
return self._classifier_failure_outcome("jev classifier is not configured", prompt, system_prompt)
@ -2212,12 +2228,14 @@ class ComplexityRouter(CustomLogger):
if not self._tier_pools().get(tier_name):
raise ValueError(f"Jev classifier returned tier {tier_name!r}, which has no models configured")
model: Final = response.model or config.model
accounting_provider: Final = "laya" if config.provider == "laya" else "typesafe"
verdict: Final = JevVerdict(
label=answer.choice,
probabilities=answer.probabilities,
confidence=answer.confidence,
model=model,
cost=jev_classifier_cost(response, config.model),
cost=jev_classifier_cost(response, config.model, accounting_provider),
provider=accounting_provider,
)
if breaker is not None and permit is not None:
breaker.record_success(permit)
@ -2225,8 +2243,8 @@ class ComplexityRouter(CustomLogger):
tier=tier,
score=None,
signals=(
f"jev-classifier:{tier_name}",
f"jev-confidence={answer.confidence:.6f}",
f"{'laya' if config.provider == 'laya' else 'jev'}-classifier:{tier_name}",
f"{'laya' if config.provider == 'laya' else 'jev'}-confidence={answer.confidence:.6f}",
*(
f"tier-probability:{label}={probability:.6f}"
for label, probability in answer.probabilities.items()
@ -4757,7 +4775,7 @@ class ComplexityRouter(CustomLogger):
tier_litellm_params: Final = self._litellm_params_for_model(tier, routed_model)
classifier_model: Final = (
f"typesafe/{outcome.jev_verdict.model}"
f"{outcome.jev_verdict.provider}/{outcome.jev_verdict.model}"
if outcome.cause == "jev_classifier" and outcome.jev_verdict is not None
else self.config.classifier_llm_config.model
if outcome.cause in ("llm_classifier", "capability_classifier", "llm_v2_classifier", "llm_v2_fallback")

View file

@ -9,6 +9,7 @@ import math
import re
import warnings
from collections.abc import Iterable, Mapping
from dataclasses import dataclass
from enum import Enum
from types import MappingProxyType
from typing import Annotated, Final, Literal, NamedTuple
@ -19,6 +20,7 @@ from pydantic import (
Field,
SkipValidation,
StrictFloat,
TypeAdapter,
field_serializer,
field_validator,
model_validator,
@ -674,14 +676,34 @@ class CapabilityClassifierConfig(BaseModel):
return self
class JevClassifierConfig(BaseModel):
def normalize_classifier_config_aliases(config: Mapping[str, object]) -> Mapping[str, object]:
if "jev_classifier_config" in config and "opensource_classifier_config" in config:
return config
normalized: Final = dict(config)
if "jev_classifier_config" in normalized:
normalized["opensource_classifier_config"] = normalized.pop("jev_classifier_config")
if normalized.get("classifier_type") == "jev":
normalized["classifier_type"] = "oss_classifier"
classifier: Final = normalized.get("opensource_classifier_config")
if isinstance(classifier, Mapping):
classifier_fields: Final = TypeAdapter(Mapping[str, object]).validate_python(classifier)
if classifier_fields.get("provider") == "typesafe":
normalized["opensource_classifier_config"] = {
**classifier_fields,
"provider": "jev",
}
return normalized
class OpenSourceClassifierConfig(BaseModel):
model_config = ConfigDict(extra="forbid", frozen=True)
provider: Literal["jev", "laya"] = "jev"
model: str = "jev-latest"
api_key: str | None = Field(default=None, description="TypeSafe API key, falling back to TYPESAFE_API_KEY")
api_key: str | None = Field(default=None, description="Provider API key; optional for self-hosted Laya")
api_base: str | None = Field(
default=None,
description="TypeSafe API base, falling back to TYPESAFE_API_BASE and then https://api.typesafe.ai",
description="Provider API base; defaults to TYPESAFE_API_BASE or LAYA_API_BASE for the selected provider",
)
timeout_ms: int = Field(default=3000, ge=1)
instructions: str | None = Field(
@ -691,30 +713,112 @@ class JevClassifierConfig(BaseModel):
circuit_breaker_enabled: bool = True
circuit_breaker_cooldown_seconds: float = Field(default=30.0, gt=0.0)
@field_validator("provider", mode="before")
@classmethod
def _normalize_provider_alias(cls, value: object) -> object:
return "jev" if value == "typesafe" else value
@field_validator("instructions")
@classmethod
def _reject_blank_instructions(cls, value: str | None) -> str | None:
if value is not None and not value.strip():
raise ValueError("jev_classifier_config.instructions must be non-empty; omit it to use the default")
raise ValueError("opensource_classifier_config.instructions must be non-empty; omit it to use the default")
return value
@field_validator("api_key")
@classmethod
def _reject_blank_api_key(cls, value: str | None) -> str | None:
if value is not None and not value.strip():
raise ValueError("jev_classifier_config.api_key must be non-empty; omit it to use TYPESAFE_API_KEY")
raise ValueError("opensource_classifier_config.api_key must be non-empty; omit it to use TYPESAFE_API_KEY")
return value
@model_validator(mode="after")
def _keep_the_environment_key_on_the_environment_base(self) -> "JevClassifierConfig":
def _keep_the_environment_key_on_the_environment_base(self) -> "OpenSourceClassifierConfig":
if self.provider == "laya":
from litellm.llms.laya.common_utils import validate_laya_api_base, validate_laya_model
_ = validate_laya_model(self.model)
if self.api_base is not None:
_ = validate_laya_api_base(self.api_base)
return self
if self.api_base is not None and self.api_key is None:
raise ValueError(
"jev_classifier_config.api_base requires jev_classifier_config.api_key: TYPESAFE_API_KEY is only sent "
"opensource_classifier_config.api_base requires opensource_classifier_config.api_key: TYPESAFE_API_KEY is only sent "
"to TYPESAFE_API_BASE or https://api.typesafe.ai"
)
return self
JevClassifierConfig = OpenSourceClassifierConfig
@dataclass(frozen=True, slots=True)
class ComplexityRouterConfigWrite:
submitted: Mapping[str, object] | None
effective: Mapping[str, object] | None
@property
def supplied_connection_fields(self) -> frozenset[str]:
classifier: Final = self.submitted.get("opensource_classifier_config") if self.submitted is not None else None
return frozenset(
field for field in ("api_base", "api_key") if isinstance(classifier, Mapping) and field in classifier
)
def resolve_complexity_router_config_write(
incoming: Mapping[str, object] | None, stored: Mapping[str, object] | None
) -> ComplexityRouterConfigWrite:
if incoming is None:
return ComplexityRouterConfigWrite(submitted=None, effective=stored)
return _resolve_normalized_complexity_router_config_write(
normalize_classifier_config_aliases(incoming),
normalize_classifier_config_aliases(stored) if stored is not None else None,
)
def _resolve_normalized_complexity_router_config_write(
incoming: Mapping[str, object], stored: Mapping[str, object] | None
) -> ComplexityRouterConfigWrite:
if (
stored is None
or incoming.get("classifier_type") != "oss_classifier"
or stored.get("classifier_type") != "oss_classifier"
):
return ComplexityRouterConfigWrite(submitted=incoming, effective=incoming)
incoming_classifier: Final = incoming.get("opensource_classifier_config")
stored_classifier: Final = stored.get("opensource_classifier_config")
if not isinstance(incoming_classifier, Mapping) or not isinstance(stored_classifier, Mapping):
return ComplexityRouterConfigWrite(submitted=incoming, effective=incoming)
existing: Final = TypeAdapter(dict[str, object]).validate_python(stored_classifier)
supplied: Final = TypeAdapter(dict[str, object]).validate_python(incoming_classifier)
classifier: Final = (
MappingProxyType({**supplied, "provider": existing["provider"]})
if "provider" not in supplied and "provider" in existing
else supplied
)
same_provider: Final = classifier.get("provider", "jev") == existing.get("provider", "jev")
same_base: Final = "api_base" not in classifier or (
classifier["api_base"] is not None and classifier["api_base"] == existing.get("api_base")
)
transport: Final = MappingProxyType(
{
key: value
for key, value in existing.items()
if same_provider and key in ("api_key", "api_base") and (key != "api_key" or same_base)
}
)
return ComplexityRouterConfigWrite(
submitted=MappingProxyType({**incoming, "opensource_classifier_config": classifier}),
effective={
**incoming,
"opensource_classifier_config": {
**transport,
**classifier,
},
},
)
MAX_CUSTOM_PATTERN_REPEAT: Final[int] = 64
MAX_CUSTOM_PATTERN_WORK: Final[int] = 2048
MAX_CUSTOM_DIMENSIONS_WORK: Final[int] = 8192
@ -846,6 +950,20 @@ class ContextCompactionConfig(BaseModel):
class ComplexityRouterConfig(BaseModel):
"""Configuration for the ComplexityRouter."""
@model_validator(mode="before")
@classmethod
def _normalize_classifier_aliases(cls, value: object) -> object:
if not isinstance(value, Mapping):
return value
config: Final = TypeAdapter(dict[str, object]).validate_python(value)
if "jev_classifier_config" in config and "opensource_classifier_config" in config:
raise ValueError("Use only opensource_classifier_config; do not also supply jev_classifier_config")
return normalize_classifier_config_aliases(config)
@property
def jev_classifier_config(self) -> OpenSourceClassifierConfig | None:
return self.opensource_classifier_config
# string = pin; list = random pick when adaptive=False, soft-floor home pool when adaptive=True
tiers: dict[str, str | list[str]] = Field(
default_factory=lambda: DEFAULT_TIER_MODELS.copy(),
@ -880,7 +998,7 @@ class ComplexityRouterConfig(BaseModel):
"becomes that tier's rubric bullet; entries named after a built-in tier may omit the "
"description and inherit the built-in criteria. List order is ascending severity and "
"decides which tier wins when several keyword_tier_rules match. Requires classifier_type "
"'llm', 'jev' or 'custom', a fallback_tier, and `tiers` keys matching the defined names exactly. Escalation, "
"'llm', 'oss_classifier' or 'custom', a fallback_tier, and `tiers` keys matching the defined names exactly. Escalation, "
"adaptive selection, session affinity, plugins, tier_labels, and the calibration-example "
"rubric presets are unavailable with a custom tier set: the first four are built on the "
"built-in tier ladder, and the last two rename or exemplify tiers the set replaces."
@ -1024,7 +1142,7 @@ class ComplexityRouterConfig(BaseModel):
"custom",
"heuristic_first",
"hybrid",
"jev",
"oss_classifier",
] = Field(
default="heuristic",
description=(
@ -1032,7 +1150,7 @@ class ComplexityRouterConfig(BaseModel):
"an LLM tier-selection call, a Switchyard-compatible capability forecast, a joint Fuse V2 forecast, "
"a custom classifier plugin, 'heuristic_first', which scores locally and only pays for the LLM classifier when the "
"local scorer does not confidently land a cheap tier, or 'hybrid', which trusts the local scorer "
"everywhere except when its score lands near a tier boundary, or 'jev', a TypeSafe AI Jev structured choice call"
"everywhere except when its score lands near a tier boundary, or 'oss_classifier', a structured choice call using Jev or Laya"
),
)
llm_v2_config: LLMV2Config | None = Field(
@ -1073,7 +1191,7 @@ class ComplexityRouterConfig(BaseModel):
"and otherwise routes to capable_tier"
),
)
jev_classifier_config: JevClassifierConfig | None = None
opensource_classifier_config: OpenSourceClassifierConfig | None = None
heuristic_first_max_tier: str | None = Field(
default=None,
description=(
@ -1639,14 +1757,16 @@ class ComplexityRouterConfig(BaseModel):
return self
@model_validator(mode="after")
def _validate_jev_classifier_config(self) -> "ComplexityRouterConfig":
jev: Final = self.jev_classifier_config
if self.classifier_type != "jev":
def _validate_opensource_classifier_config(self) -> "ComplexityRouterConfig":
jev: Final = self.opensource_classifier_config
if self.classifier_type != "oss_classifier":
if jev is not None:
raise ValueError("jev_classifier_config requires classifier_type 'jev'; otherwise it has no effect")
raise ValueError(
"opensource_classifier_config requires classifier_type 'oss_classifier'; otherwise it has no effect"
)
return self
if jev is None:
raise ValueError("jev_classifier_config is required when classifier_type is 'jev'")
raise ValueError("opensource_classifier_config is required when classifier_type is 'oss_classifier'")
return self
@model_validator(mode="after")
@ -1962,9 +2082,9 @@ class ComplexityRouterConfig(BaseModel):
"enable_non_reasoning_tier cannot be combined with tier_definitions: a custom tier set "
f"replaces the built-in ladder, so name a tier {non_reasoning_key} in tier_definitions instead"
)
if self.classifier_type not in ("llm", "custom", "jev"):
if self.classifier_type not in ("llm", "custom", "oss_classifier"):
raise ValueError(
f"enable_non_reasoning_tier requires classifier_type 'llm', 'jev' or 'custom', got "
f"enable_non_reasoning_tier requires classifier_type 'llm', 'oss_classifier' or 'custom', got "
f"{self.classifier_type!r}: the heuristic scorers only produce the four tiers from SIMPLE up, "
f"so nothing would ever classify as {non_reasoning_key}"
)
@ -1997,7 +2117,7 @@ class ComplexityRouterConfig(BaseModel):
raise ValueError(f"tier_definitions names must be unique (case-insensitive): {', '.join(duplicated)}")
if self.classifier_type in ("heuristic", "heuristic_v2", "capability", "heuristic_first", "hybrid"):
raise ValueError(
"tier_definitions requires classifier_type 'llm', 'jev' or 'custom': the heuristic scorer only "
"tier_definitions requires classifier_type 'llm', 'oss_classifier' or 'custom': the heuristic scorer only "
"produces the built-in tiers from SIMPLE up, as does heuristic_v2"
)
conflicts: Final = self._tier_definition_conflicts()
@ -2164,7 +2284,9 @@ class ComplexityRouterConfig(BaseModel):
)
COMPLEXITY_ROUTER_CONFIG_KEYS: Final[frozenset[str]] = frozenset(ComplexityRouterConfig.model_fields)
COMPLEXITY_ROUTER_CONFIG_KEYS: Final[frozenset[str]] = frozenset(ComplexityRouterConfig.model_fields) | frozenset(
("jev_classifier_config",)
)
"""Every setting name this config owns, derived from the model so a field added later is covered.
These names are disjoint from the OpenAI request params, from ``all_litellm_params``, and from the

View file

@ -18,6 +18,7 @@ from litellm.litellm_core_utils.internal_call_metadata import (
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.llms.laya.common_utils import laya_response_model
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.typesafe_passthrough_logging_handler import (
TypeSafePassthroughLoggingHandler,
)
@ -78,10 +79,17 @@ class JevClassifierClient(Protocol):
class HttpJevClassifierClient:
def __init__(self, api_key: str, api_base: str, http_client: AsyncHTTPHandler) -> None:
def __init__(
self,
api_key: str | None,
api_base: str,
http_client: AsyncHTTPHandler,
provider: Literal["typesafe", "laya"] = "typesafe",
) -> None:
self._api_key = api_key
self._api_base = api_base.rstrip("/")
self._http_client = http_client
self._provider = provider
async def evaluate(
self,
@ -90,26 +98,30 @@ class HttpJevClassifierClient:
request_kwargs: Mapping[str, object] | None = None,
) -> JevSystemOneResponse:
start_time: Final = datetime.now(timezone.utc)
authorization: Final[Mapping[str, str]] = (
MappingProxyType({"Authorization": f"Bearer {self._api_key}"}) if self._api_key else MappingProxyType({})
)
response: Final = await self._http_client.post( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler has a dynamic post signature
f"{self._api_base}/v1/systemone",
json=request.model_dump(mode="json"),
headers=MappingProxyType(
{
"Authorization": f"Bearer {self._api_key}",
"Content-Type": "application/json",
}
), # pyright: ignore[reportArgumentType] # HTTP headers are not mutated by AsyncHTTPHandler
headers=MappingProxyType({**authorization, "Content-Type": "application/json"}), # pyright: ignore[reportArgumentType] # HTTP headers are not mutated by AsyncHTTPHandler
timeout=timeout_s,
)
response.raise_for_status()
body: Final = TypeAdapter(dict[str, object]).validate_json(response.content)
normalized_body: Final = (
MappingProxyType({**body, "model": laya_response_model(body, request.model)})
if self._provider == "laya"
else body
)
try:
self._log_response(request, response, request_kwargs, start_time)
except Exception as exc: # noqa: BLE001 # logging integrations must not discard a provider verdict
verbose_router_logger.warning("JEV response logging failed (%s)", type(exc).__name__)
return TypeAdapter(JevSystemOneResponse).validate_python(response.json())
return TypeAdapter(JevSystemOneResponse).validate_python(normalized_body)
@staticmethod
def _log_response(
self,
request: JevSystemOneRequest,
response: httpx.Response,
request_kwargs: Mapping[str, object] | None,
@ -139,7 +151,7 @@ class HttpJevClassifierClient:
"turn_off_message_logging": effective_turn_off_message_logging(request_kwargs),
}
logging_obj: Final = Logging(
model=f"typesafe/{request.model}",
model=f"{self._provider}/{request.model}",
messages=[{"role": "user", "content": request.state}],
stream=False,
call_type="pass_through_endpoint",
@ -150,7 +162,7 @@ class HttpJevClassifierClient:
kwargs=params,
)
logging_obj.update_environment_variables(
model=f"typesafe/{request.model}",
model=f"{self._provider}/{request.model}",
user=parent_user if isinstance(parent_user := parent.get("user"), str) else None,
optional_params={},
litellm_params=params,
@ -165,7 +177,7 @@ class HttpJevClassifierClient:
end_time=end_time,
cache_hit=False,
request_body=MappingProxyType({"model": request.model}),
custom_llm_provider="typesafe",
custom_llm_provider=self._provider,
litellm_params=params,
)
success_handlers: Final = logging_obj.dispatch_success_handlers(
@ -189,6 +201,7 @@ class JevVerdict(NamedTuple):
confidence: float
model: str
cost: float | None
provider: Literal["typesafe", "laya"] = "typesafe"
class _RegistryPricing(BaseModel):
@ -211,12 +224,14 @@ def build_jev_request(
return JevSystemOneRequest(state=state, model=model, questions=MappingProxyType({"tier": question}))
def jev_classifier_cost(response: JevSystemOneResponse, configured_model: str) -> float | None:
def jev_classifier_cost(
response: JevSystemOneResponse, configured_model: str, provider: Literal["typesafe", "laya"] = "typesafe"
) -> float | None:
usage: Final = response.usage
if usage is None:
return None
model: Final = response.model or configured_model
model_key: Final = f"typesafe/{model}"
model_key: Final = f"{provider}/{model}"
if model_key not in litellm.model_cost: # pyright: ignore[reportUnknownMemberType] # registry is dynamically typed
return None
try:

View file

@ -19,6 +19,7 @@ from litellm.router_strategy.complexity_router.config import (
COMPLEXITY_ROUTER_CONFIG_KEYS,
DEFAULT_JEV_INSTRUCTIONS,
LLM_CLASSIFIER_TYPES,
normalize_classifier_config_aliases,
)
AUTO_ROUTER_MODEL_PREFIX: Final = "auto_router/"
@ -151,8 +152,11 @@ def strategy_router_dependencies(
)
)
)
complexity: Final = _mapping(litellm_params.get("complexity_router_config"))
complexity: Final = normalize_classifier_config_aliases(_mapping(litellm_params.get("complexity_router_config")))
classifier: Final = _mapping(complexity.get("classifier_llm_config"))
decision_classifier: Final = _mapping(complexity.get("opensource_classifier_config"))
decision_provider: Final = decision_classifier.get("provider", "jev")
accounting_provider: Final = "typesafe" if decision_provider == "jev" else decision_provider
return tuple(
dict.fromkeys(
tuple(dep for tier in _mapping(complexity.get("tiers")).values() for dep in _pool(tier, "tier"))
@ -165,10 +169,10 @@ def strategy_router_dependencies(
)
+ (
_named(
f"typesafe/{_mapping(complexity.get('jev_classifier_config')).get('model', 'jev-latest')}",
f"{accounting_provider}/{decision_classifier.get('model', 'jev-latest')}",
"evaluation",
)
if complexity.get("classifier_type") == "jev"
if complexity.get("classifier_type") == "oss_classifier"
else ()
)
+ (
@ -206,9 +210,9 @@ def defines_custom_classifier_prompt(complexity_router_config: object) -> bool:
Scoped to the classifier types that actually call an LLM, which is also where the config validator
accepts these fields: the heuristic scorers never read them.
"""
config: Final = _mapping(complexity_router_config)
if config.get("classifier_type") == "jev":
instructions: Final = _mapping(config.get("jev_classifier_config")).get("instructions")
config: Final = normalize_classifier_config_aliases(_mapping(complexity_router_config))
if config.get("classifier_type") == "oss_classifier":
instructions: Final = _mapping(config.get("opensource_classifier_config")).get("instructions")
return isinstance(instructions, str) and instructions != DEFAULT_JEV_INSTRUCTIONS
if config.get("classifier_type") not in LLM_CLASSIFIER_TYPES:
return False
@ -272,6 +276,9 @@ _OPERATOR_PROMPT_FIELDS_SQL: Final = " OR ".join(
f"{{config}} ->> '{field}' IS NOT NULL" for field in OPERATOR_CLASSIFIER_PROMPT_FIELDS
)
_DEFAULT_JEV_INSTRUCTIONS_SQL: Final = DEFAULT_JEV_INSTRUCTIONS.replace("'", "''")
_OPENSOURCE_CLASSIFIER_CONFIG_SQL: Final = (
"COALESCE({config} -> 'opensource_classifier_config', {config} -> 'jev_classifier_config')"
)
CUSTOMIZATION_CAPABILITY: Final = GatedAutoRouterCapability(
key="tier_or_classifier_prompt",
@ -286,9 +293,9 @@ CUSTOMIZATION_CAPABILITY: Final = GatedAutoRouterCapability(
f"({{config}} ->> 'classifier_type' IN ({_LLM_CLASSIFIER_TYPES_SQL}) AND ("
"{config} -> 'classifier_llm_config' ->> 'system_prompt' IS NOT NULL OR "
f"{_OPERATOR_PROMPT_FIELDS_SQL})) OR "
"({config} ->> 'classifier_type' = 'jev' AND "
"jsonb_typeof({config} -> 'jev_classifier_config' -> 'instructions') = 'string' AND "
f"{{config}} -> 'jev_classifier_config' ->> 'instructions' <> '{_DEFAULT_JEV_INSTRUCTIONS_SQL}')"
"({config} ->> 'classifier_type' IN ('oss_classifier', 'jev') AND "
f"jsonb_typeof({_OPENSOURCE_CLASSIFIER_CONFIG_SQL} -> 'instructions') = 'string' AND "
f"{_OPENSOURCE_CLASSIFIER_CONFIG_SQL} ->> 'instructions' <> '{_DEFAULT_JEV_INSTRUCTIONS_SQL}')"
),
)

View file

@ -119,6 +119,7 @@ def trace_summary_from_row(row: dict[str, Any], spend_rows: Sequence[_SpendRow]
trace_ref=row.get("trace_ref", ""),
name=row["name"],
service=row["service"],
agent_names=tuple(row.get("agent_names") or ()),
input_preview=row["input_preview"],
start_time=_iso(int(row["start_ms"])),
duration_ms=float(row["duration_ms"]),
@ -169,8 +170,8 @@ def _parent_agent_of(span: Span, by_id: Mapping[str, Span]) -> str | None:
if parent_id is None or parent_id not in by_id or parent_id == span["span_id"]:
return None
parent = by_id[parent_id]
if parent["type"] == "agent" and parent["name"] != span["name"]:
return parent["name"]
if parent["type"] == "agent" and (parent["agent"] or parent["name"]) != (span["agent"] or span["name"]):
return parent["agent"] or parent["name"]
parent_id = parent["parent_span_id"]
return None
@ -183,9 +184,9 @@ def agent_nodes(spans: Sequence[Span]) -> tuple[AgentNode, ...]:
if span["type"] != "agent":
continue
node = agents.setdefault(
span["name"],
span["agent"] or span["name"],
AgentNode(
name=span["name"],
name=span["agent"] or span["name"],
parent_agent=_parent_agent_of(span, by_id),
invocations=0,
llm_calls=0,
@ -250,6 +251,7 @@ def trace_from_rows(
trace_ref=trace_ref,
name=root["name"],
service=rows[0]["service"],
agent_names=tuple(sorted(frozenset(s["agent"] for s in spans if s["agent"]))),
input_preview=root["input_preview"],
start_time=_iso(trace_start_ns // NANOS_PER_MS),
duration_ms=(trace_end_ns - trace_start_ns) / NANOS_PER_MS,

View file

@ -56,6 +56,7 @@ class TraceSummary(TypedDict):
trace_ref: ReadOnly[NotRequired[str]]
name: ReadOnly[str]
service: ReadOnly[str]
agent_names: ReadOnly[NotRequired[tuple[str, ...]]]
input_preview: ReadOnly[str]
start_time: ReadOnly[str] # ISO 8601
duration_ms: ReadOnly[float]

View file

@ -205,14 +205,18 @@ class AutoRouterCacheStats(BaseModel):
class AutoRouterBenchmarkTotals(BaseModel):
"""Session-shape and savings aggregates over auto-routed traffic in the window."""
"""Auto-routed traffic in the window. Turns, spend and savings count requests on the selected UTC days;
the session averages and cache stats describe every session overlapping the window, whole."""
sessions: int
turns: int
avg_turns_per_session: float
avg_session_seconds: float
avg_tokens_per_session: float
spend: float = Field(description="What the routed traffic actually cost")
sessions: int = Field(description="Sessions overlapping the window, counted whole")
turns: int = Field(description="Auto-routed requests on the selected UTC days")
avg_turns_per_session: float | None = Field(
description="Lifetime turns per overlapping session; null when the window has routed requests but no session "
"rows for this router type, such as an alias whose router type changed mid-session"
)
avg_session_seconds: float | None = Field(description="Lifetime seconds per overlapping session; null as above")
avg_tokens_per_session: float | None = Field(description="Lifetime tokens per overlapping session; null as above")
spend: float = Field(description="What the selected days' routed traffic actually cost")
classifier_cost: float | None = Field(
description="Recorded LLM classifier cost already included in spend; null when any session turns predate "
"subtotal recording, and zero for an empty window"
@ -229,14 +233,19 @@ class AutoRouterBenchmarkTotals(BaseModel):
"null when classification costs for those requests are unavailable",
)
saved_spend: float | None = Field(
description="Recorded historical savings plus newer estimates; null when traffic has no recorded savings estimates"
description="Recorded savings on the selected UTC days; null when traffic has no recorded savings estimates. "
"On totals this is the same daily figure the Overall savings view reports"
)
unattributed_saved_spend: float | None = Field(
default=None,
description="Part of saved_spend no router's daily rows account for, such as history recorded before "
"per-router daily tracking; when set, baseline_spend and saved_pct are null",
)
baseline_spend: float | None = Field(
description="Estimated single-model cost: compared actual spend plus recorded savings; "
"null when traffic has no recorded savings"
)
saved_pct: float | None = Field(description="Recorded savings over baseline_spend, as a percentage")
saved_per_session: float | None = Field(description="Recorded savings per session, including historical estimates")
cache: AutoRouterCacheStats
@ -291,7 +300,7 @@ class AutoRouterSessionResponse(BaseModel):
class AutoRouterBenchmarksResponse(BaseModel):
"""Benchmarks for the auto-router dashboard, aggregated from the per-session rollup."""
"""Benchmarks for the auto-router dashboard, aggregated from the per-session and per-day rollups."""
start_date: str = Field(description="Window start day, YYYY-MM-DD UTC, inclusive")
end_date: str = Field(description="Window end day, YYYY-MM-DD UTC, inclusive")

View file

@ -72622,6 +72622,45 @@
"supports_audio_input": true,
"supports_video_input": true
},
"laya/english": {
"input_cost_per_token": 0.0,
"litellm_provider": "laya",
"mode": "evaluation",
"output_cost_per_token": 0.0,
"source": "https://github.com/NandhaKishorM/laya",
"supported_endpoints": [
"/v1/systemone"
],
"metadata": {
"notes": "Self-hosted decision model; infrastructure costs are paid separately"
}
},
"laya/multilingual": {
"input_cost_per_token": 0.0,
"litellm_provider": "laya",
"mode": "evaluation",
"output_cost_per_token": 0.0,
"source": "https://github.com/NandhaKishorM/laya",
"supported_endpoints": [
"/v1/systemone"
],
"metadata": {
"notes": "Self-hosted decision model; infrastructure costs are paid separately"
}
},
"laya/typed-decisions": {
"input_cost_per_token": 0.0,
"litellm_provider": "laya",
"mode": "evaluation",
"output_cost_per_token": 0.0,
"source": "https://github.com/NandhaKishorM/laya",
"supported_endpoints": [
"/v1/systemone"
],
"metadata": {
"notes": "Self-hosted decision model; infrastructure costs are paid separately"
}
},
"typesafe/jev-1.13.0": {
"input_cost_per_token": 4.2e-08,
"litellm_provider": "typesafe",

View file

@ -6,6 +6,7 @@
"url": "Link to provider documentation",
"endpoints": {
"chat_completions": "Supports /chat/completions endpoint",
"systemone": "Supports native System One typed decisions",
"messages": "Supports /messages endpoint (Anthropic format)",
"responses": "Supports /responses endpoint (OpenAI/Anthropic unified)",
"embeddings": "Supports /embeddings endpoint",
@ -1476,6 +1477,13 @@
"rerank": false
}
},
"laya": {
"display_name": "Laya (`laya`)",
"url": "https://docs.litellm.ai/docs/auto_router/decision_classifiers",
"endpoints": {
"systemone": true
}
},
"lambda_ai": {
"display_name": "Lambda AI (`lambda_ai`)",
"url": "https://docs.litellm.ai/docs/providers/lambda_ai",
@ -3354,6 +3362,13 @@
"provider_json_field": "skills",
"url": "https://docs.litellm.ai/docs/skills"
},
"systemone": {
"docs_label": "systemone",
"display_name": "System One Decision API",
"leftnav_label": "/laya/v1/systemone",
"provider_json_field": "systemone",
"url": "https://docs.litellm.ai/docs/auto_router/decision_classifiers"
},
"text_completion": {
"docs_label": "text_completion",
"display_name": "OpenAI Completions API",

View file

@ -1744,6 +1744,27 @@ model LiteLLM_AutoRouterUserSession {
@@index([user_id, last_turn_at], map: "idx_autorouter_user_session_user_last_turn")
}
// Auto-routed requests per UTC request day and router: the selected-day money behind the
// auto-router usage view. Written in the same statement as the session rollup, so a day row
// and its session row never disagree; corrected in the same transaction as late baselines.
model LiteLLM_AutoRouterDailySpend {
date String
api_key String
user_id String
router_name String
router_type String
turns Int @default(0)
spend Float @default(0)
saved_spend Float @default(0)
savings_estimated_turns Int @default(0)
savings_estimated_actual_spend Float @default(0)
savings_estimated_saved_spend Float @default(0)
classifier_cost Float @default(0)
classifier_cost_recorded_turns Int @default(0)
@@id([date, api_key, user_id, router_name, router_type])
}
// Shadow eval: evaluation of an auto-router against one or more keys' live traffic, in
// either direction. forward duplicates the requests the keys did not route through the
// router through it, answering whether they should adopt it; reverse duplicates the

View file

@ -32,6 +32,7 @@ from e2e_config import (
MCP_OAUTH_LIVE_OPT_IN_ENV,
OTEL_TLS_OPT_IN_ENV,
OTEL_V2_OPT_IN_ENV,
OWNED_GATEWAY_OPT_IN_ENV,
PROMPT_CACHING_OPT_IN_ENV,
PROVIDER_EDGE_HOST_OPT_IN_ENV,
PROXY_BASE_URL,
@ -70,6 +71,7 @@ OPT_IN_MARKERS: Final = MappingProxyType(
"cli_determinism": CLI_DETERMINISM_OPT_IN_ENV,
"mcp_oauth_live": MCP_OAUTH_LIVE_OPT_IN_ENV,
"provider_edge_host": PROVIDER_EDGE_HOST_OPT_IN_ENV,
"owned_gateway": OWNED_GATEWAY_OPT_IN_ENV,
"otel_v2": OTEL_V2_OPT_IN_ENV,
"otel_tls": OTEL_TLS_OPT_IN_ENV,
"secret_manager": SECRET_MANAGER_OPT_IN_ENV,
@ -172,6 +174,11 @@ def pytest_configure(config: pytest.Config) -> None:
"provider_edge_host: routes provider traffic through the pytest host's edge in every fixture mode, so the "
"gateway must reach the pytest host; deselected unless E2E_PROVIDER_EDGE_HOST_REACHABLE is set",
)
config.addinivalue_line(
"markers",
"owned_gateway: boots its own proxy from source against the stack's Postgres, so it needs DATABASE_URL "
"on the pytest host; deselected unless E2E_OWNED_GATEWAY is set",
)
config.addinivalue_line(
"markers",
"otel_v2: needs a proxy running with LITELLM_OTEL_V2=true; deselected unless E2E_OTEL_V2 is set",

View file

@ -60,6 +60,9 @@
- {id: other.auth.jwt.wrong_issuer_denied, module: other, tier: P0, area: auth, assertions: [wrong_issuer_denied], source: "auth/handle_jwt.py", rationale: "A signed token with the correct audience and an unexpected issuer is rejected"}
- {id: other.auth.jwt.wrong_audience_denied, module: other, tier: P0, area: auth, assertions: [wrong_audience_denied], source: "auth/handle_jwt.py", rationale: "A signed token from the trusted issuer intended for another app is rejected"}
- {id: other.auth.jwt.auto_register_maps_existing_key, module: other, tier: P0, area: auth, assertions: [maps_existing_key], source: "user_api_key_auth.py _auto_register_jwt_mapping", rationale: "With auto_register_map_existing_key: true, the first JWT call of a user who already owns a key writes the sub-claim mapping to that existing key hash and mints nothing; the spend row lands on the pre-existing key (LIT-5378)", fail_before_fix: proven}
- {id: other.auth.jwt.auto_register_mints_when_keyless, module: other, tier: P0, area: auth, assertions: [mints_when_keyless], source: "user_api_key_auth.py _auto_register_jwt_mapping", rationale: "With auto_register_map_existing_key: true, a user with no keys still gets exactly one minted key and a sub-claim mapping on their first JWT call (LIT-5378)", fail_before_fix: proven}
- {id: other.auth.jwt.auto_register_default_mints, module: other, tier: P0, area: auth, assertions: [default_mints], source: "user_api_key_auth.py _auto_register_jwt_mapping", rationale: "Without auto_register_map_existing_key, auto_register keeps the current behavior: it mints a second key for a user who already has one and bills the minted key (LIT-5378)", fail_before_fix: proven}
- {id: other.auth.session_token.valid_allows, module: other, tier: P0, area: auth, assertions: [valid_allows], source: "auth/auth_checks.py ExperimentalUIJWTToken", rationale: "An unexpired LiteLLM-minted session token authenticates with the role it carries"}
- {id: other.auth.session_token.expired_denied, module: other, tier: P0, area: auth, assertions: [expired_denied], source: "auth/user_api_key_auth.py expiry check", rationale: "An expired session token is rejected with the expired-key error"}
- {id: other.auth.session_token.encrypted_value_denied, module: other, tier: P0, area: auth, assertions: [encrypted_value_denied], source: "auth/auth_checks.py ExperimentalUIJWTToken", rationale: "An encrypted value read back from a management route is not accepted as a bearer token"}

View file

@ -7,6 +7,7 @@ environment so the same tests run against localhost or a deployed proxy.
from __future__ import annotations
import os
import socket
from dataclasses import dataclass
import time
import uuid
@ -16,6 +17,7 @@ from typing import Final
from dotenv import load_dotenv
from fixture_mode import deterministic_marker, parse_fixture_mode, registration_owner
from provider_edge import provider_edge_api_base
from pydantic import TypeAdapter
# Local runs keep provider / DataDog keys in tests/e2e/.env (see CONTRIBUTING.md).
# Compose injects them into the proxy container, but pytest on the host does not
@ -206,6 +208,7 @@ REDIS_CHAOS_OPT_IN_ENV = "E2E_REDIS_CHAOS"
CLI_DETERMINISM_OPT_IN_ENV = "E2E_CLI_DETERMINISM"
MCP_OAUTH_LIVE_OPT_IN_ENV: Final = "E2E_MCP_OAUTH_LIVE"
PROVIDER_EDGE_HOST_OPT_IN_ENV: Final = "E2E_PROVIDER_EDGE_HOST_REACHABLE"
OWNED_GATEWAY_OPT_IN_ENV: Final = "E2E_OWNED_GATEWAY"
OTEL_V2_OPT_IN_ENV: Final = "E2E_OTEL_V2"
OTEL_TLS_OPT_IN_ENV: Final = "E2E_OTEL_EXPORTER_ENDPOINT"
SECRET_MANAGER_OPT_IN_ENV: Final = "E2E_SECRET_MANAGER"
@ -296,6 +299,15 @@ def unique_marker() -> str:
return uuid.uuid4().hex[:12]
INHERITED_ENV_PREFIXES: Final = ("REDIS_", "MICROSOFT_", "GOOGLE_", "GENERIC_", "PROXY_")
def available_port() -> int:
with socket.socket() as listener:
listener.bind(("127.0.0.1", 0))
return TypeAdapter(tuple[str, int]).validate_python(listener.getsockname())[1]
def settle_propagation(written_at: float) -> None:
"""Block until PROPAGATION_TIMEOUT has elapsed since `written_at`, a
`time.monotonic()` stamp taken the moment a control-plane write returned.

View file

@ -8,7 +8,6 @@ The optional live edge measures headers without recording credentials or bodies.
from __future__ import annotations
import os
import socket
import subprocess
import sys
import threading
@ -20,13 +19,12 @@ from pathlib import Path
from typing import Final
import psycopg
from e2e_config import INHERITED_ENV_PREFIXES, available_port
from e2e_http import NoBody
from idp import Keycloak, stop_process_group
from proxy_client import ProxyClient, build_proxy_client
from psycopg.rows import class_row
from pydantic import BaseModel, SecretStr, TypeAdapter, ValidationError
INHERITED_ENV_PREFIXES: Final = ("REDIS_", "MICROSOFT_", "GOOGLE_", "GENERIC_", "PROXY_")
from pydantic import BaseModel, SecretStr, ValidationError
class StoredOAuth(BaseModel):
@ -101,12 +99,6 @@ class OAuthObservation:
assert all(not item[2] for item in snapshot), "gateway bearer leaked to the upstream"
def available_port() -> int:
with socket.socket() as listener:
listener.bind(("127.0.0.1", 0))
return TypeAdapter(tuple[str, int]).validate_python(listener.getsockname())[1]
@dataclass(slots=True)
class OAuthGateway:
base_url: str

View file

@ -1537,6 +1537,7 @@ class UserNewBody(BaseModel):
class UserNewResponse(BaseModel):
user_id: str
key: str | None = None
class UserUpdateBody(BaseModel):
@ -1580,6 +1581,40 @@ class UserListResponse(BaseModel):
total: int
class UserKeyRow(BaseModel):
token: str
key_alias: str | None = None
class UserInfoWithKeysResponse(BaseModel):
user_id: str | None = None
keys: list[UserKeyRow] = []
class JwtKeyMappingRow(BaseModel):
id: str
jwt_claim_name: str
jwt_claim_value: str
created_by: str | None = None
class JwtKeyMappingListParams(BaseModel):
size: int = 100
class JwtKeyMappingListResponse(BaseModel):
mappings: list[JwtKeyMappingRow]
total_count: int
class JwtKeyMappingDeleteBody(BaseModel):
id: str
class JwtKeyMappingDeleteResponse(BaseModel):
status: str
class OrgNewBody(BaseModel):
organization_alias: str
models: list[str] = []

View file

@ -21,12 +21,20 @@ from idp import Keycloak, keycloak_from_env
from models import (
ChatBody,
ChatResponse,
JwtKeyMappingDeleteBody,
JwtKeyMappingDeleteResponse,
JwtKeyMappingListParams,
JwtKeyMappingListResponse,
ModelsListParams,
ModelsListResponse,
ReadinessDetailsResponse,
ReadinessResponse,
UserInfoParams,
UserInfoWithKeysResponse,
UserListParams,
UserListResponse,
UserNewBody,
UserNewResponse,
)
from proxy_client import ProxyClient
from pydantic import Field
@ -79,6 +87,44 @@ class OtherClient:
response_type=ReadinessDetailsResponse,
)
def user_new(self, body: UserNewBody) -> Result[UserNewResponse]:
"""POST /user/new under the master key: seed the litellm user a JWT
`sub` claim resolves to, before that token ever reaches the proxy."""
return self.proxy.transport.post(
"/user/new",
headers=self.proxy.transport.master,
json=body,
response_type=UserNewResponse,
)
def user_info(self, user_id: str) -> Result[UserInfoWithKeysResponse]:
"""GET /user/info under the master key. Only the user's key rows are
modelled: `token` is the stored key hash, never the plaintext key."""
return self.proxy.transport.get(
"/user/info",
headers=self.proxy.transport.master,
params=UserInfoParams(user_id=user_id),
response_type=UserInfoWithKeysResponse,
)
def jwt_mapping_list(self) -> Result[JwtKeyMappingListResponse]:
"""GET /jwt/key/mapping/list under the master key."""
return self.proxy.transport.get(
"/jwt/key/mapping/list",
headers=self.proxy.transport.master,
params=JwtKeyMappingListParams(size=100),
response_type=JwtKeyMappingListResponse,
)
def jwt_mapping_delete(self, mapping_id: str) -> Result[JwtKeyMappingDeleteResponse]:
"""POST /jwt/key/mapping/delete under the master key."""
return self.proxy.transport.post(
"/jwt/key/mapping/delete",
headers=self.proxy.transport.master,
json=JwtKeyMappingDeleteBody(id=mapping_id),
response_type=JwtKeyMappingDeleteResponse,
)
def chat_as_team(self, token: str, team: str, body: ChatBody) -> Result[ChatResponse]:
"""POST /chat/completions under `token` with `x-litellm-team-id: team`."""
return self.proxy.transport.post(

View file

@ -0,0 +1,108 @@
"""An owned, source-built proxy whose `litellm_jwtauth` block a test controls.
The shared proxy on :4000 runs the CONTRIBUTING.md JWT block, so a test that
needs a different `litellm_jwtauth` config boots its own gateway on a free port
against the same database and the same Keycloak realm. The caller supplies the
`litellm_jwtauth` mapping verbatim, which is exactly what makes a config an
unfixed proxy rejects observable as a boot failure in this gateway's own log.
"""
from __future__ import annotations
import os
import subprocess
import sys
import time
from collections.abc import Mapping
from contextlib import ExitStack
from dataclasses import dataclass, field
from pathlib import Path
from typing import Final
from e2e_config import INHERITED_ENV_PREFIXES, available_port
from e2e_http import NoBody
from idp import Keycloak, stop_process_group
from proxy_client import ProxyClient, build_proxy_client
MODEL_NAME: Final = "gemini-3.8-flash"
@dataclass(slots=True)
class OwnedJwtGateway:
base_url: str
proxy: ProxyClient
_environment: Mapping[str, str] = field(repr=False)
_command: tuple[str, ...] = field(repr=False)
_log_path: Path
_child: subprocess.Popen[bytes] | None = field(default=None, init=False, repr=False)
def start(self) -> None:
with self._log_path.open("ab") as log:
self._child = subprocess.Popen(
self._command,
env=self._environment,
stdout=log,
stderr=log,
start_new_session=True,
)
deadline: Final = time.monotonic() + 120
while time.monotonic() < deadline:
assert self._child.poll() is None, "owned JWT gateway exited; inspect its private log"
result = self.proxy.transport.probe("/health/liveliness", params=NoBody())
if result.status_code == 200:
return
time.sleep(0.5)
raise AssertionError("owned JWT gateway did not become ready")
def stop(self) -> None:
if self._child is not None:
stop_process_group(self._child)
assert self._child.poll() is not None, "old gateway process is still alive"
def owned_jwt_gateway(
idp: Keycloak, directory: Path, cleanup: ExitStack, *, litellm_jwtauth: str, name: str
) -> OwnedJwtGateway:
for env_name in ("DATABASE_URL", "LITELLM_LICENSE", "LITELLM_MASTER_KEY"):
assert os.environ.get(env_name), f"{env_name} is required for the owned JWT gateway"
port: Final = available_port()
base_url: Final = f"http://127.0.0.1:{port}"
config: Final = directory / f"{name}.yaml"
config.write_text(
"model_list:\n"
f" - model_name: {MODEL_NAME}\n"
" litellm_params:\n"
f" model: gemini/{MODEL_NAME}\n"
" api_key: os.environ/GEMINI_API_KEY\n"
"general_settings:\n"
" master_key: os.environ/LITELLM_MASTER_KEY\n"
" database_url: os.environ/DATABASE_URL\n"
" proxy_batch_write_at: 5\n"
" enable_jwt_auth: true\n"
" litellm_jwtauth:\n" + "".join(f" {line}\n" for line in litellm_jwtauth.strip().splitlines())
)
environment: Final = {
**{key: value for key, value in os.environ.items() if not key.startswith(INHERITED_ENV_PREFIXES)},
"JWT_PUBLIC_KEY_URL": idp.jwks_url,
"JWT_ISSUER": idp.issuer,
"JWT_AUDIENCE": "litellm-e2e",
"LITELLM_DANGEROUSLY_PERMIT_WEAK_OR_UNSET_MASTER_KEY": "true",
"DISABLE_SCHEMA_UPDATE": "true",
"STORE_MODEL_IN_DB": "True",
"PYTHONPATH": str(Path(__file__).resolve().parents[3]),
}
gateway: Final = OwnedJwtGateway(
base_url=base_url,
proxy=build_proxy_client(
base_url=base_url,
control_plane_base_url=base_url,
replica_urls=(base_url,),
master_key=os.environ["LITELLM_MASTER_KEY"],
),
_environment=environment,
_command=(sys.executable, "-m", "litellm.proxy.proxy_cli", "--config", str(config), "--port", str(port)),
_log_path=directory / f"{name}.log",
)
cleanup.callback(gateway.stop)
gateway.start()
return gateway

View file

@ -0,0 +1,182 @@
"""auto_register with auto_register_map_existing_key binds the JWT claim to the user's existing key.
`unregistered_jwt_client_behavior: auto_register` on `virtual_key_claim_field: sub` mints a fresh
virtual key on the user's first JWT call. With `auto_register_map_existing_key: true` the proxy must
instead point the new JWT mapping at a key the resolved user already owns, and mint only when the
user has none. Each behavior gets its own gateway because the flag lives in `litellm_jwtauth`, so
this file boots two owned proxies against the shared database and Keycloak realm.
"""
from __future__ import annotations
import hashlib
from collections.abc import Iterator
from contextlib import ExitStack
from typing import Final
import pytest
from e2e_config import unique_marker
from e2e_http import unwrap
from idp import Identity, Keycloak
from lifecycle import ResourceManager
from models import ChatBody, ChatMessage, JwtKeyMappingRow, KeyGenerateBody, TeamNewBody, UserNewBody
from other_client import OtherClient
from owned_jwt_gateway import MODEL_NAME, OwnedJwtGateway, owned_jwt_gateway
pytestmark = pytest.mark.e2e
_JWT_COMMON: Final = (
"user_id_jwt_field: sub\n"
"user_email_jwt_field: email\n"
"team_ids_jwt_field: groups\n"
"user_id_upsert: true\n"
"virtual_key_claim_field: sub\n"
"unregistered_jwt_client_behavior: auto_register"
)
def _key_hash(key: str) -> str:
return hashlib.sha256(key.encode()).hexdigest()
def _ping() -> ChatBody:
return ChatBody(
model=MODEL_NAME,
messages=[ChatMessage(role="user", content=f"Reply with the single word ok. {unique_marker()}")],
max_tokens=5,
)
def _identity_with_user(idp: Keycloak, client: OtherClient, resources: ResourceManager) -> Identity:
"""An IdP identity plus the litellm user and team its claims resolve to, with
teardown that also sweeps the user's keys and JWT mapping rows the proxy
wrote, since those outlive the user row itself."""
marker: Final = unique_marker()
identity: Final = idp.provision(marker=marker, group=f"e2e-jwt-team-{marker}", defer=resources.defer)
resources.defer(lambda: client.proxy.delete_user(identity.user_id))
team_id: Final = client.proxy.create_team(TeamNewBody(team_alias=f"e2e-jwt-{marker}", team_id=identity.group))
resources.defer(lambda: client.proxy.delete_team(team_id))
unwrap(
client.user_new(
UserNewBody(
user_id=identity.user_id,
user_email=f"{identity.username}@example.com",
user_role="internal_user",
auto_create_key=False,
)
)
)
def delete_user_keys() -> None:
for row in unwrap(client.user_info(identity.user_id)).keys:
client.proxy.delete_key(row.token)
def delete_user_mappings() -> None:
for mapping in unwrap(client.jwt_mapping_list()).mappings:
if mapping.jwt_claim_value == identity.user_id:
_ = client.jwt_mapping_delete(mapping.id)
resources.defer(delete_user_keys)
resources.defer(delete_user_mappings)
return identity
def _mapping_for(client: OtherClient, claim_value: str) -> JwtKeyMappingRow | None:
return next(
(row for row in unwrap(client.jwt_mapping_list()).mappings if row.jwt_claim_value == claim_value),
None,
)
@pytest.fixture(scope="module")
def mapping_gateway(idp: Keycloak, tmp_path_factory: pytest.TempPathFactory) -> Iterator[OwnedJwtGateway]:
with ExitStack() as cleanup:
yield owned_jwt_gateway(
idp,
tmp_path_factory.mktemp("jwt-mapping"),
cleanup,
litellm_jwtauth=f"{_JWT_COMMON}\nauto_register_map_existing_key: true",
name="jwt-mapping-gateway",
)
@pytest.fixture(scope="module")
def minting_gateway(idp: Keycloak, tmp_path_factory: pytest.TempPathFactory) -> Iterator[OwnedJwtGateway]:
with ExitStack() as cleanup:
yield owned_jwt_gateway(
idp,
tmp_path_factory.mktemp("jwt-minting"),
cleanup,
litellm_jwtauth=_JWT_COMMON,
name="jwt-minting-gateway",
)
@pytest.mark.owned_gateway
class TestJwtAutoRegisterMapExistingKey:
@pytest.mark.covers("other.auth.jwt.auto_register_maps_existing_key")
def test_first_jwt_call_maps_to_the_users_existing_key_and_mints_none(
self, client: OtherClient, idp: Keycloak, resources: ResourceManager, mapping_gateway: OwnedJwtGateway
) -> None:
identity: Final = _identity_with_user(idp, client, resources)
existing_key: Final = client.proxy.generate_key(
KeyGenerateBody(
user_id=identity.user_id, team_id=identity.group, key_alias=f"e2e-jwt-existing-{unique_marker()}"
)
)
resources.defer(lambda: client.proxy.delete_key(existing_key))
response: Final = unwrap(mapping_gateway.proxy.chat(idp.access_token(identity), _ping()))
keys: Final = unwrap(client.user_info(identity.user_id)).keys
assert [row.token for row in keys] == [_key_hash(existing_key)], (
f"map_existing_key must leave the user with only their pre-existing key, got {keys}"
)
mapping: Final = _mapping_for(client, identity.user_id)
assert mapping is not None, (
f"no JWT mapping row for sub={identity.user_id}: {unwrap(client.jwt_mapping_list())}"
)
assert mapping.jwt_claim_name == "sub", f"mapping must bind the sub claim, got {mapping}"
assert mapping.created_by == "auto_register", f"mapping must be written by auto_register, got {mapping}"
rows: Final = client.proxy.poll_logs_for_key(existing_key)
assert any(row.request_id == response.id for row in rows), (
f"the JWT chat must be billed to the user's existing key, spend rows for it: {rows}"
)
@pytest.mark.covers("other.auth.jwt.auto_register_mints_when_keyless")
def test_first_jwt_call_mints_a_key_when_the_user_has_none(
self, client: OtherClient, idp: Keycloak, resources: ResourceManager, mapping_gateway: OwnedJwtGateway
) -> None:
identity: Final = _identity_with_user(idp, client, resources)
response: Final = unwrap(mapping_gateway.proxy.chat(idp.access_token(identity), _ping()))
assert response.choices, f"JWT chat returned no completion: {response}"
keys: Final = unwrap(client.user_info(identity.user_id)).keys
assert len(keys) == 1, f"a keyless user must get exactly one minted key, got {keys}"
mapping: Final = _mapping_for(client, identity.user_id)
assert mapping is not None and mapping.jwt_claim_name == "sub", (
f"the minted key must be recorded as a sub-claim mapping, mappings: {unwrap(client.jwt_mapping_list())}"
)
@pytest.mark.covers("other.auth.jwt.auto_register_default_mints")
def test_default_behavior_still_mints_when_the_user_already_has_a_key(
self, client: OtherClient, idp: Keycloak, resources: ResourceManager, minting_gateway: OwnedJwtGateway
) -> None:
identity: Final = _identity_with_user(idp, client, resources)
existing_key: Final = client.proxy.generate_key(
KeyGenerateBody(user_id=identity.user_id, key_alias=f"e2e-jwt-existing-{unique_marker()}")
)
resources.defer(lambda: client.proxy.delete_key(existing_key))
response: Final = unwrap(minting_gateway.proxy.chat(idp.access_token(identity), _ping()))
assert response.id is not None, f"JWT chat returned no response id: {response}"
keys: Final = unwrap(client.user_info(identity.user_id)).keys
assert len(keys) == 2, (
f"default auto_register must mint a second key for a user who already has one, got {keys}"
)
rows: Final = client.proxy.poll_logs_for_request_id(response.id)
assert rows and all(row.api_key != _key_hash(existing_key) for row in rows), (
f"the default path must bill the minted key, not the user's existing one: {rows}"
)

View file

@ -16,6 +16,7 @@ markers =
quiet_stack: measures the proxy itself, so it runs while no other test on this host is hitting the stack; every other test waits for it to finish
mcp_oauth_live: real Linear OAuth consent via a captured browser session; deselected unless E2E_MCP_OAUTH_LIVE is set
provider_edge_host: routes provider traffic through the pytest host's edge in every fixture mode, so the gateway must reach the pytest host; deselected unless E2E_PROVIDER_EDGE_HOST_REACHABLE is set
owned_gateway: boots its own proxy from source against the stack's Postgres, so it needs DATABASE_URL on the pytest host; deselected unless E2E_OWNED_GATEWAY is set
otel_v2: needs a proxy running with LITELLM_OTEL_V2=true; deselected unless E2E_OTEL_V2 is set
otel_tls: needs a stack whose gateway exports OTLP over TLS signed by the CA in SSL_CERT_FILE; deselected unless E2E_OTEL_EXPORTER_ENDPOINT is set
secret_manager: needs a proxy booted from gateway/secret_manager_<system>_ci_config.yml against that live secret manager; deselected unless E2E_SECRET_MANAGER names the backend (see secret_manager/secret_backends.py)

View file

@ -88,15 +88,89 @@ class DatabaseRelay:
)
class HeldStatementRelay:
def __init__(self, upstream_host: str, upstream_port: int, trigger: bytes) -> None:
self.port: Final = _free_port()
self._upstream_host: Final = upstream_host
self._upstream_port: Final = upstream_port
self._trigger: Final = trigger
self._loop: Final = asyncio.new_event_loop()
self._released: Final = asyncio.Event()
self.held: Final = threading.Event()
self._ready: Final = threading.Event()
self._thread: Final = threading.Thread(target=self._run, daemon=True)
def release(self) -> None:
self._loop.call_soon_threadsafe(self._released.set)
def start(self) -> None:
self._thread.start()
assert self._ready.wait(10), "Database relay did not start"
def stop(self) -> None:
self.release()
self._loop.call_soon_threadsafe(self._loop.stop)
self._thread.join(10)
def _run(self) -> None:
asyncio.set_event_loop(self._loop)
self._loop.run_until_complete(asyncio.start_server(self._serve, "127.0.0.1", self.port))
self._ready.set()
self._loop.run_forever()
def _holds(self, window: bytes) -> bool:
return not self.held.is_set() and self._trigger in window
async def _serve(self, client_reader: asyncio.StreamReader, client_writer: asyncio.StreamWriter) -> None:
server_reader, server_writer = await asyncio.open_connection(self._upstream_host, self._upstream_port)
async def forward(reader: asyncio.StreamReader, writer: asyncio.StreamWriter, inspect: bool) -> None:
tail = b"" # rebind-ok: carries the previous read's end so a trigger split across reads still matches
try:
while chunk := await reader.read(65536):
window: Final = tail + chunk
if inspect and self._holds(window):
self.held.set()
await self._released.wait()
tail = window[-(len(self._trigger) - 1) :]
writer.write(chunk)
await writer.drain()
except (ConnectionError, asyncio.IncompleteReadError):
return
finally:
writer.close()
await asyncio.gather(
forward(client_reader, server_writer, True),
forward(server_reader, client_writer, False),
)
def _relayed_url(database_url: str, port: int) -> str:
parts: Final = urlsplit(database_url)
credentials: Final = f"{parts.username}:{parts.password}@" if parts.username else ""
return urlunsplit(parts._replace(netloc=f"{credentials}127.0.0.1:{port}"))
@contextmanager
def database_relay(database_url: str, trigger: bytes) -> Generator[tuple[DatabaseRelay, str]]:
parts: Final = urlsplit(database_url)
assert parts.hostname is not None and parts.port is not None, database_url
relay: Final = DatabaseRelay(parts.hostname, parts.port, trigger)
relay.start()
credentials: Final = f"{parts.username}:{parts.password}@" if parts.username else ""
relayed: Final = urlunsplit(parts._replace(netloc=f"{credentials}127.0.0.1:{relay.port}"))
try:
yield relay, relayed
yield relay, _relayed_url(database_url, relay.port)
finally:
relay.stop()
@contextmanager
def held_statement_relay(database_url: str, trigger: bytes) -> Generator[tuple[HeldStatementRelay, str]]:
parts: Final = urlsplit(database_url)
assert parts.hostname is not None and parts.port is not None, database_url
relay: Final = HeldStatementRelay(parts.hostname, parts.port, trigger)
relay.start()
try:
yield relay, _relayed_url(database_url, relay.port)
finally:
relay.stop()

View file

@ -0,0 +1,219 @@
import json
import os
import time
import uuid
from collections.abc import Iterator
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
from hashlib import sha256
from pathlib import Path
from typing import Final
import httpx
import jwt
import pytest
import yaml
from cryptography.hazmat.primitives.asymmetric import rsa
from tests.integration._support.client import Gateway, eventually, string_value
from tests.integration._support.database import read_rows
from tests.integration._support.database_relay import held_statement_relay
from tests.integration._support.process import owned_proxy
from tests.integration._support.wire import Reply, Request, wire_server
KEY_ID: Final = "integration-jwt-map-existing-key"
MAPPING_INSERT: Final = b'INSERT INTO "public"."LiteLLM_JWTKeyMapping"'
pytestmark = pytest.mark.timeout(240)
def _hash(key: str) -> str:
return sha256(key.encode()).hexdigest()
def _config(directory: Path, claim_field: str) -> Path:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["general_settings"] = {
**config["general_settings"],
"enable_jwt_auth": True,
"litellm_jwtauth": {
"user_id_jwt_field": "sub",
"user_email_jwt_field": "email",
"virtual_key_claim_field": claim_field,
"unregistered_jwt_client_behavior": "auto_register",
"auto_register_map_existing_key": True,
},
}
path: Final = directory / f"jwt_map_existing_key_{claim_field}.yaml"
path.write_text(yaml.safe_dump(config))
return path
@contextmanager
def _issuer() -> Iterator[tuple[rsa.RSAPrivateKey, str]]:
private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048)
public_jwk: Final = json.loads(jwt.algorithms.RSAAlgorithm.to_jwk(private_key.public_key()))
jwks: Final = json.dumps({"keys": [{**public_jwk, "kid": KEY_ID, "use": "sig", "alg": "RS256"}]}).encode()
def respond(request: Request) -> Reply:
assert request.method == "GET", request
return Reply(body=jwks)
with wire_server(respond) as server:
yield private_key, server.url
def _token(private_key: rsa.RSAPrivateKey, subject: str, **claims: str) -> str:
now: Final = int(time.time())
return jwt.encode(
{"sub": subject, **claims, "iat": now, "exp": now + 300},
private_key,
algorithm="RS256",
headers={"kid": KEY_ID},
)
def _chat(candidate: Gateway, model: str, token: str) -> httpx.Response:
return candidate.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": "map existing key control"}]},
key=token,
)
def _mapped_token(claim_name: str, claim_value: str) -> str:
rows: Final = read_rows(
'SELECT token FROM "LiteLLM_JWTKeyMapping" WHERE jwt_claim_name = %s AND jwt_claim_value = %s',
(claim_name, claim_value),
)
assert len(rows) == 1, rows
return string_value(rows[0]["token"])
def _user_key_hashes(user: str) -> frozenset[str]:
rows: Final = read_rows('SELECT token FROM "LiteLLM_VerificationToken" WHERE user_id = %s', (user,))
return frozenset(string_value(row["token"]) for row in rows)
def _billed_key(response: httpx.Response) -> str:
rows: Final = eventually(
lambda: read_rows(
'SELECT api_key FROM "LiteLLM_SpendLogs" WHERE request_id = %s', (str(response.json()["id"]),)
),
lambda values: len(values) == 1,
seconds=70,
)
return string_value(rows[0]["api_key"])
def test_first_jwt_call_reuses_the_newest_durable_llm_key_and_skips_every_ineligible_newer_key(
gateway: Gateway, tmp_path: Path
) -> None:
with _issuer() as (private_key, jwks_url), gateway.scenario() as scenario:
model: Final = scenario.model()
user: Final = scenario.user(user_role="internal_user")
older_durable: Final = scenario.key(user_id=user)
durable: Final = scenario.key(user_id=user)
skipped: Final = {
"older_durable": older_durable,
"expiring": scenario.key(user_id=user, duration="1h"),
"management_only": scenario.key(user_id=user, allowed_routes=["management_routes"]),
"auto_registered_look_alike": scenario.key(user_id=user, metadata={"auto_registered": True}),
"other_team": scenario.key(user_id=user, team_id=scenario.team()),
"blocked": scenario.key(user_id=user),
}
gateway.post("/key/block", {"key": skipped["blocked"]})
keys_before: Final = _user_key_hashes(user)
with owned_proxy(
gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks_url}, config=_config(tmp_path, "sub")
) as candidate:
response: Final = _chat(candidate, model, _token(private_key, user))
assert response.status_code == 200, response.text
mapped: Final = _mapped_token("sub", user)
assert mapped == _hash(durable), {
"mapped_to": next((name for name, key in skipped.items() if _hash(key) == mapped), mapped)
}
assert _user_key_hashes(user) == keys_before, "a key was minted although a reusable one existed"
assert _billed_key(response) == _hash(durable)
def test_user_matched_by_email_instead_of_sub_still_reuses_their_existing_key(gateway: Gateway, tmp_path: Path) -> None:
with _issuer() as (private_key, jwks_url), gateway.scenario() as scenario:
model: Final = scenario.model()
email: Final = f"integration-{uuid.uuid4().hex}@example.com"
user: Final = scenario.user(user_role="internal_user", user_email=email)
existing: Final = scenario.key(user_id=user)
subject: Final = f"integration-idp-subject-{uuid.uuid4().hex}"
with owned_proxy(
gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks_url}, config=_config(tmp_path, "sub")
) as candidate:
response: Final = _chat(candidate, model, _token(private_key, subject, email=email.upper()))
assert response.status_code == 200, response.text
assert _mapped_token("sub", subject) == _hash(existing)
assert _user_key_hashes(user) == frozenset({_hash(existing)}), "a key was minted for an email-matched user"
assert _billed_key(response) == _hash(existing)
def test_shared_client_claim_never_maps_a_second_user_onto_the_first_users_personal_key(
gateway: Gateway, tmp_path: Path
) -> None:
with _issuer() as (private_key, jwks_url), gateway.scenario() as scenario:
model: Final = scenario.model()
first_user: Final = scenario.user(user_role="internal_user")
second_user: Final = scenario.user(user_role="internal_user")
personal: Final = scenario.key(user_id=first_user)
client_id: Final = f"integration-shared-client-{uuid.uuid4().hex}"
with owned_proxy(
gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks_url}, config=_config(tmp_path, "client_id")
) as candidate:
first: Final = _chat(candidate, model, _token(private_key, first_user, client_id=client_id))
second: Final = _chat(candidate, model, _token(private_key, second_user, client_id=client_id))
assert first.status_code == 200, first.text
assert second.status_code == 200, second.text
mapped: Final = _mapped_token("client_id", client_id)
assert mapped != _hash(personal), "the shared client claim was mapped to the first user's personal key"
assert (_billed_key(first), _billed_key(second)) == (mapped, mapped)
assert read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE api_key = %s', (_hash(personal),)) == []
def test_concurrent_first_jwt_calls_of_a_keyless_user_both_succeed_on_one_surviving_mapped_key(
gateway: Gateway, tmp_path: Path
) -> None:
writer_url: Final = os.environ.get("INTEGRATION_PROXY_DATABASE_URL") or os.environ["DATABASE_URL"]
with (
_issuer() as (private_key, jwks_url),
gateway.scenario() as scenario,
held_statement_relay(writer_url, MAPPING_INSERT) as (relay, relayed_url),
):
model: Final = scenario.model()
user: Final = scenario.user(user_role="internal_user")
token: Final = _token(private_key, user)
overrides: Final = {
"JWT_PUBLIC_KEY_URL": jwks_url,
"DATABASE_URL": relayed_url,
"PRISMA_HEALTH_WATCHDOG_ENABLED": "false",
}
with (
owned_proxy(gateway, tmp_path, overrides, config=_config(tmp_path, "sub")) as candidate,
ThreadPoolExecutor(max_workers=1) as pool,
):
held_call: Final = pool.submit(_chat, candidate, model, token)
assert relay.held.wait(60), "the first call never reached its mapping insert"
racing: Final = _chat(candidate, model, token)
relay.release()
held: Final = held_call.result(timeout=60)
assert racing.status_code == 200, racing.text
assert held.status_code == 200, held.text
keys: Final = _user_key_hashes(user)
assert len(keys) == 1, keys
assert _mapped_token("sub", user) in keys
assert (_billed_key(held), _billed_key(racing)) == (_mapped_token("sub", user),) * 2

View file

@ -124,6 +124,20 @@ def _trace_spans(sink_url: str, trace_id: str, seconds: float = 30) -> tuple[Spa
return group
def _trace_spans_when(
sink_url: str,
trace_id: str,
ready: Callable[[tuple[Span, ...]], bool],
seconds: float = 30,
) -> tuple[Span, ...]:
spans: Final = eventually(
lambda: spans_for_trace(recorded_spans(sink_url)[1], trace_id),
ready,
seconds=seconds,
)
return spans
def _await_db_span(sink_url: str, trace_id: str | None, needle: str, seconds: float = 40, since: int = 0) -> None:
def seen() -> bool:
_, spans = recorded_spans(sink_url, since)
@ -242,6 +256,196 @@ def test_without_excluded_services_the_tenant_still_gets_redis_and_postgres_span
assert {"redis", "postgresql"} <= systems, f"datastore spans missing at tenant: {systems}"
@pytest.mark.parametrize("otel", [None, True, "on", "", []], ids=["null", "true", "on", "empty_string", "empty_list"])
@pytest.mark.timeout(180)
def test_a_non_mapping_otel_block_still_publishes_the_tenant_fan_out(
gateway: Gateway,
audit_sinks: SpanSinks,
otel_audit_config: AuditConfigWriter,
langfuse_vars: dict[str, JsonValue],
tmp_path: Path,
otel: JsonValue,
) -> None:
def with_callback_settings(config: dict) -> None:
config["litellm_settings"]["callbacks"] = ["langfuse_otel"]
config["callback_settings"]["otel"] = otel
config: Final = _config_with(tmp_path, otel_audit_config, extra=with_callback_settings)
overrides: Final = {"LITELLM_OTEL_V2": "1", **_operator_langfuse(audit_sinks)}
with owned_proxy(gateway, tmp_path, overrides, config=config, workers=2) as candidate:
traffic: Final = _drive(candidate, langfuse_vars)
tenant_trace: Final = _trace_id(audit_sinks.tenant, traffic)
_await_db_span(audit_sinks.tenant, tenant_trace, "redis")
tenant_spans: Final = _trace_spans_when(
audit_sinks.tenant,
tenant_trace,
lambda spans: any(span["kind"] == 2 for span in spans) and "redis" in _db_systems(spans),
seconds=15,
)
assert any(span["kind"] == 2 for span in tenant_spans), "tenant SERVER root span missing"
assert "redis" in _db_systems(tenant_spans), f"tenant redis span missing: {_db_systems(tenant_spans)}"
operator_trace: Final = _trace_id(audit_sinks.operator, traffic)
operator_spans: Final = _trace_spans_when(
audit_sinks.operator,
operator_trace,
lambda spans: any(span["kind"] == 2 for span in spans),
seconds=15,
)
assert any(span["kind"] == 2 for span in operator_spans), "operator SERVER root span missing"
@pytest.mark.parametrize("name", ["EXCLUDED_SERVICES", "excluded_services"])
@pytest.mark.timeout(180)
def test_a_bare_excluded_services_env_var_is_ignored(
gateway: Gateway,
audit_sinks: SpanSinks,
otel_audit_config: AuditConfigWriter,
langfuse_vars: dict[str, JsonValue],
tmp_path: Path,
name: str,
) -> None:
config: Final = _config_with(tmp_path, otel_audit_config)
overrides: Final = {"LITELLM_OTEL_V2": "1", name: "redis,postgres"}
with owned_proxy(gateway, tmp_path, overrides, config=config, workers=2) as candidate:
tenant_start, _ = recorded_spans(audit_sinks.tenant)
traffic: Final = _drive(candidate, langfuse_vars)
tenant_trace: Final = _trace_id(audit_sinks.tenant, traffic)
_await_db_span(audit_sinks.tenant, tenant_trace, "redis")
tenant_spans: Final = _trace_spans_when(
audit_sinks.tenant,
tenant_trace,
lambda spans: "redis" in _db_systems(spans),
seconds=15,
)
assert "redis" in _db_systems(tenant_spans), f"redis span missing at tenant: {_db_systems(tenant_spans)}"
_await_db_span(audit_sinks.tenant, None, "postgresql", since=tenant_start)
_, all_tenant = recorded_spans(audit_sinks.tenant, tenant_start)
systems: Final = _db_systems(all_tenant)
assert {"redis", "postgresql"} <= systems, f"datastore spans missing at tenant: {systems}"
@pytest.mark.timeout(180)
def test_the_documented_env_var_wins_over_a_bare_excluded_services(
gateway: Gateway,
audit_sinks: SpanSinks,
otel_audit_config: AuditConfigWriter,
langfuse_vars: dict[str, JsonValue],
tmp_path: Path,
) -> None:
config: Final = _config_with(tmp_path, otel_audit_config)
overrides: Final = {
"LITELLM_OTEL_V2": "1",
"LITELLM_OTEL_EXCLUDED_SERVICES": "redis",
"EXCLUDED_SERVICES": "postgres",
}
with owned_proxy(gateway, tmp_path, overrides, config=config, workers=2) as candidate:
tenant_start, _ = recorded_spans(audit_sinks.tenant)
traffic: Final = _drive(candidate, langfuse_vars)
tenant_trace: Final = _trace_id(audit_sinks.tenant, traffic)
_trace_spans(audit_sinks.tenant, tenant_trace, seconds=15)
_await_db_span(audit_sinks.tenant, None, "postgresql", since=tenant_start)
_, tenant_spans = recorded_spans(audit_sinks.tenant, tenant_start)
systems: Final = _db_systems(tenant_spans)
assert "postgresql" in systems, f"postgresql spans missing at tenant: {systems}"
assert "redis" not in systems, f"redis spans reached tenant: {systems}"
@pytest.mark.parametrize(
("env_name", "redis_reaches_tenant"),
[
pytest.param("LITELLM_OTEL_EXCLUDED_SERVICES", False, id="exact-case"),
pytest.param("litellm_otel_excluded_services", True, id="wrong-case"),
],
)
@pytest.mark.timeout(180)
def test_case_sensitive_otel_settings_read_only_the_exact_env_name(
gateway: Gateway,
audit_sinks: SpanSinks,
otel_audit_config: AuditConfigWriter,
langfuse_vars: Mapping[str, JsonValue],
tmp_path: Path,
env_name: str,
redis_reaches_tenant: bool,
) -> None:
config: Final = _config_with(tmp_path, otel_audit_config, otel={"_case_sensitive": True})
overrides: Final = {"LITELLM_OTEL_V2": "1", env_name: "redis"}
with owned_proxy(
gateway,
tmp_path,
overrides,
config=config,
remove_environment=("LITELLM_OTEL_EXCLUDED_SERVICES", "litellm_otel_excluded_services"),
workers=2,
) as candidate:
tenant_start, _ = recorded_spans(audit_sinks.tenant)
traffic: Final = _drive(candidate, langfuse_vars)
tenant_trace: Final = _trace_id(audit_sinks.tenant, traffic)
if redis_reaches_tenant:
_await_db_span(audit_sinks.tenant, tenant_trace, "redis")
_trace_spans_when(
audit_sinks.tenant,
tenant_trace,
lambda spans: "redis" in _db_systems(spans),
seconds=15,
)
else:
_trace_spans(audit_sinks.tenant, tenant_trace, seconds=15)
_await_db_span(audit_sinks.tenant, None, "postgresql", seconds=60, since=tenant_start)
_, tenant_spans = recorded_spans(audit_sinks.tenant, tenant_start)
systems: Final = _db_systems(tenant_spans)
assert "postgresql" in systems, f"postgresql spans missing at tenant: {systems}"
assert ("redis" in systems) is redis_reaches_tenant, (
f"tenant redis presence={('redis' in systems)}; expected={redis_reaches_tenant}; systems={systems}"
)
@pytest.mark.timeout(180)
def test_env_ignore_empty_keeps_the_default_service_name(
gateway: Gateway,
audit_sinks: SpanSinks,
otel_audit_config: AuditConfigWriter,
langfuse_vars: Mapping[str, JsonValue],
tmp_path: Path,
) -> None:
config: Final = _config_with(tmp_path, otel_audit_config, otel={"_env_ignore_empty": True})
overrides: Final = {"LITELLM_OTEL_V2": "1", "OTEL_SERVICE_NAME": ""}
with owned_proxy(gateway, tmp_path, overrides, config=config, workers=2) as candidate:
traffic: Final = _drive(candidate, langfuse_vars)
operator_trace: Final = _trace_id(audit_sinks.operator, traffic)
operator_spans: Final = _trace_spans_when(
audit_sinks.operator,
operator_trace,
lambda spans: any(span["kind"] == 2 for span in spans),
seconds=15,
)
service_names: Final = tuple(span["resource"].get("service.name") for span in operator_spans)
assert service_names and all(name == "litellm" for name in service_names), (
f"operator service.name values={service_names}"
)
@pytest.mark.timeout(180)
def test_env_parse_none_str_reads_a_null_traces_endpoint_as_unset(
gateway: Gateway,
audit_sinks: SpanSinks,
otel_audit_config: AuditConfigWriter,
langfuse_vars: Mapping[str, JsonValue],
tmp_path: Path,
) -> None:
config: Final = _config_with(tmp_path, otel_audit_config, otel={"_env_parse_none_str": "null"})
overrides: Final = {"LITELLM_OTEL_V2": "1", "OTEL_TRACES_ENDPOINT": "null"}
with owned_proxy(gateway, tmp_path, overrides, config=config, workers=2) as candidate:
traffic: Final = _drive(candidate, langfuse_vars)
operator_trace: Final = _trace_id(audit_sinks.operator, traffic)
operator_spans: Final = _trace_spans_when(
audit_sinks.operator,
operator_trace,
lambda spans: any(span["kind"] == 2 for span in spans),
seconds=15,
)
assert any(span["kind"] == 2 for span in operator_spans), "operator SERVER root span missing"
def test_env_excluded_services_drops_only_redis(
gateway: Gateway,
audit_sinks: SpanSinks,

View file

@ -9,7 +9,8 @@ from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
from typing import Final, Literal
from types import MappingProxyType
from typing import Final, Literal, Protocol, cast
import anthropic
import httpx
@ -26,6 +27,9 @@ from pydantic import JsonValue, TypeAdapter
MARKER: Final = re.compile(rb"excl-[0-9a-f]{32}")
FAILING: Final = re.compile(rb"excl-fail-[0-9a-f]{32}")
JSON: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
UNCONFIGURED_VARIANT: Final[TypeAdapter[Literal["null_block", "bare_env"]]] = TypeAdapter(
Literal["null_block", "bare_env"]
)
REPLY_TEXT: Final = "excluded ok"
SERVER: Final = 2
INVALID_NAME_LOG: Final = "is not a datastore service"
@ -37,6 +41,11 @@ CLIENTS: Final[tuple[Client, ...]] = ("raw", "sdk", "async_sdk")
AuditConfigWriter = Callable[[Path, Mapping[str, JsonValue]], Path]
class _FixtureRequestParam(Protocol):
@property
def param(self) -> object: ...
def _marker() -> str:
return "excl-" + uuid.uuid4().hex
@ -339,6 +348,11 @@ def _db_systems(spans: tuple[Span, ...]) -> set[str]:
}
def _post_auth_datastore_spans(spans: tuple[Span, ...]) -> tuple[Span, ...]:
auth_ids: Final = frozenset(span["span_id"] for span in spans if span["name"].startswith("auth "))
return tuple(span for span in spans if _db_systems((span,)) and span["parent_span_id"] not in auth_ids)
def _names(spans: tuple[Span, ...]) -> list[str]:
return sorted(span["name"] for span in spans)
@ -408,6 +422,29 @@ def _config(directory: Path, otel_audit_config: AuditConfigWriter, otel: Mapping
return path
def _null_otel_config(directory: Path, otel_audit_config: AuditConfigWriter, name: str) -> Path:
written: Final = otel_audit_config(directory, {})
loaded: Final = object_value(JSON.validate_python(yaml.safe_load(written.read_text())))
config: Final = {
**loaded,
"litellm_settings": {**object_value(loaded["litellm_settings"]), "callbacks": ["langfuse_otel"]},
"callback_settings": {**object_value(loaded["callback_settings"]), "otel": None},
}
path: Final = directory / f"{name}.yaml"
path.write_text(yaml.safe_dump(config))
return path
def _operator_langfuse(sinks: SpanSinks) -> dict[str, str]:
return {
"LANGFUSE_HOST": sinks.operator,
"LANGFUSE_PUBLIC_KEY": "pk-lf-operator",
"LANGFUSE_SECRET_KEY": "sk-lf-operator",
"OTEL_EXPORTER": "http/json",
"OTEL_ENDPOINT": sinks.operator,
}
@contextmanager
def _started(
provider: Wire,
@ -416,13 +453,14 @@ def _started(
directory: Path,
langfuse_vars: Mapping[str, JsonValue],
workers: int,
environment: Mapping[str, str] = MappingProxyType({}),
) -> Generator[Rig]:
with (
gateway_from_environment() as gateway,
owned_proxy_process(
gateway,
directory,
{"LITELLM_OTEL_V2": "1", "OTEL_BSP_SCHEDULE_DELAY": "300"},
{"LITELLM_OTEL_V2": "1", "OTEL_BSP_SCHEDULE_DELAY": "300", **environment},
config=config,
remove_environment=("LITELLM_OTEL_EXCLUDED_SERVICES",),
workers=workers,
@ -458,6 +496,32 @@ def rig(
yield started
@pytest.fixture(scope="module", params=["null_block", "bare_env"], ids=["null_block", "bare_env"])
def unconfigured_rig(
request: pytest.FixtureRequest,
provider: Wire,
audit_sinks: SpanSinks,
otel_audit_config: AuditConfigWriter,
langfuse_vars: dict[str, JsonValue],
tmp_path_factory: pytest.TempPathFactory,
) -> Iterator[Rig]:
parameter: Final = cast(_FixtureRequestParam, request).param
variant: Final = UNCONFIGURED_VARIANT.validate_python(parameter)
directory: Final = tmp_path_factory.mktemp(f"excluded-{variant}")
config: Final = (
_null_otel_config(directory, otel_audit_config, variant)
if variant == "null_block"
else _config(directory, otel_audit_config, {}, variant)
)
environment: Final = (
_operator_langfuse(audit_sinks) if variant == "null_block" else {"EXCLUDED_SERVICES": "redis,postgres"}
)
with _started(
provider, audit_sinks, config, directory, langfuse_vars, workers=2, environment=environment
) as started:
yield started
@pytest.mark.timeout(120)
@pytest.mark.parametrize("stream", [False, True], ids=["unary", "stream"])
@pytest.mark.parametrize("client", CLIENTS)
@ -654,6 +718,237 @@ def test_killing_one_of_two_workers_mid_burst_keeps_the_filter_on_the_survivor(r
_assert_withheld(rig, rig.raw("chat", _marker(), stream=False), after)
def _assert_tenant_kept(
rig: Rig, trace_id: str, cursors: Cursors, *, needs_model_span: bool = False
) -> tuple[Span, ...]:
def ready(spans: tuple[Span, ...]) -> bool:
return (
sum(1 for span in spans if span["kind"] == SERVER) == 1
and "redis" in _db_systems(spans)
and (not needs_model_span or any("gen_ai.operation.name" in span["attributes"] for span in spans))
)
tenant: Final = eventually(
lambda: spans_for_trace(recorded_spans(rig.sinks.tenant, cursors.tenant)[1], trace_id),
ready,
seconds=40,
return_last_on_timeout=True,
)
assert sum(1 for span in tenant if span["kind"] == SERVER) == 1, _names(tenant)
assert "redis" in _db_systems(tenant), f"redis spans missing at the tenant: {_names(tenant)}"
assert not needs_model_span or any("gen_ai.operation.name" in span["attributes"] for span in tenant), _names(tenant)
return tenant
def _assert_kept(rig: Rig, sent: Sent, cursors: Cursors) -> tuple[Span, ...]:
operator: Final = _operator_trace(rig, sent, cursors)
return _assert_tenant_kept(rig, operator[0]["trace_id"], cursors, needs_model_span=True)
@pytest.mark.timeout(120)
@pytest.mark.parametrize("stream", [False, True], ids=["unary", "stream"])
@pytest.mark.parametrize("client", CLIENTS)
@pytest.mark.parametrize("endpoint", ENDPOINTS)
def test_unconfigured_tenant_trace_keeps_datastore_spans(
unconfigured_rig: Rig, endpoint: Endpoint, client: Client, stream: bool
) -> None:
cursors: Final = unconfigured_rig.cursors()
marker: Final = _marker()
sent: Final = unconfigured_rig.send(endpoint, client, marker, stream)
assert sent.text == REPLY_TEXT, sent
assert unconfigured_rig.upstream_hits(marker) == 1
_assert_kept(unconfigured_rig, sent, cursors)
@pytest.mark.timeout(120)
@pytest.mark.parametrize("endpoint", ["chat", "messages"])
def test_unconfigured_cache_hit_twin_keeps_datastore_spans(unconfigured_rig: Rig, endpoint: Endpoint) -> None:
cursors: Final = unconfigured_rig.cursors()
marker: Final = _marker()
first_result: Final = _traced_raw(unconfigured_rig, endpoint, marker)
first: Final = first_result[1]
assert first.text == REPLY_TEXT, first
assert unconfigured_rig.upstream_hits(marker) == 1
_assert_kept(unconfigured_rig, first, cursors)
hit_cursors: Final = unconfigured_rig.cursors()
def read_hit() -> tuple[str, Sent, tuple[Span, ...]]:
trace_id, sent = _traced_raw(unconfigured_rig, endpoint, marker)
return trace_id, sent, _operator_trace_by_id(unconfigured_rig, trace_id, hit_cursors)
trace_id, hit, operator = eventually(
read_hit,
lambda result: (
unconfigured_rig.upstream_hits(marker) == 0
and "redis" in _db_systems(_post_auth_datastore_spans(result[2]))
),
seconds=60,
)
assert hit.text == REPLY_TEXT, hit
post_auth_datastore: Final = _post_auth_datastore_spans(operator)
post_auth_span_ids: Final = frozenset(span["span_id"] for span in post_auth_datastore)
non_datastore_names: Final = frozenset(span["name"] for span in operator if not _db_systems((span,)))
tenant: Final = eventually(
lambda: spans_for_trace(recorded_spans(unconfigured_rig.sinks.tenant, hit_cursors.tenant)[1], trace_id),
lambda spans: (
non_datastore_names <= frozenset(span["name"] for span in spans)
and post_auth_span_ids <= frozenset(span["span_id"] for span in spans)
),
seconds=40,
return_last_on_timeout=True,
)
tenant_span_ids: Final = frozenset(span["span_id"] for span in tenant)
missing_post_auth_names: Final = tuple(
span["name"] for span in post_auth_datastore if span["span_id"] not in tenant_span_ids
)
assert sum(1 for span in tenant if span["kind"] == SERVER) == 1, _names(tenant)
assert "redis" in _db_systems(tenant), (
f"operator datastore systems={sorted(_db_systems(post_auth_datastore))}; "
f"tenant datastore systems={sorted(_db_systems(tenant))}; tenant spans={_names(tenant)}"
)
assert not missing_post_auth_names, (
f"missing post-auth datastore span names={missing_post_auth_names}; "
f"operator={_names(post_auth_datastore)}; tenant={_names(tenant)}"
)
@pytest.mark.timeout(120)
@pytest.mark.parametrize("endpoint", ENDPOINTS)
def test_unconfigured_failed_upstream_keeps_datastore_spans(unconfigured_rig: Rig, endpoint: Endpoint) -> None:
cursors: Final = unconfigured_rig.cursors()
marker: Final = "excl-fail-" + uuid.uuid4().hex
trace_id: Final = uuid.uuid4().hex
path, body = _body(unconfigured_rig.model, endpoint, marker, stream=False)
failed: Final = unconfigured_rig.proxy.client.post(
path,
json=body,
headers={
"Authorization": f"Bearer {unconfigured_rig.key}",
"traceparent": f"00-{trace_id}-{uuid.uuid4().hex[:16]}-01",
},
)
assert failed.status_code == 500, failed.text
assert unconfigured_rig.upstream_hits(marker) >= 1
operator: Final = eventually(
lambda: spans_for_trace(recorded_spans(unconfigured_rig.sinks.operator, cursors.operator)[1], trace_id),
lambda spans: _has_root(spans) and "redis" in _db_systems(spans),
seconds=40,
)
assert "redis" in _db_systems(operator), _names(operator)
_assert_tenant_kept(unconfigured_rig, trace_id, cursors)
@pytest.mark.timeout(120)
@pytest.mark.parametrize("status", [403, 404])
def test_unconfigured_rejecting_tenant_destination_recovers(unconfigured_rig: Rig, status: int) -> None:
configure_sink(unconfigured_rig.sinks.tenant, status=status)
try:
cursors: Final = unconfigured_rig.cursors()
marker: Final = _marker()
sent: Final = unconfigured_rig.raw("chat", marker, stream=True)
assert sent.text == REPLY_TEXT, sent
assert unconfigured_rig.upstream_hits(marker) == 1
_operator_trace(unconfigured_rig, sent, cursors)
finally:
configure_sink(unconfigured_rig.sinks.tenant, status=200)
after: Final = unconfigured_rig.cursors()
recovered: Final = unconfigured_rig.raw("responses", _marker(), stream=False)
assert recovered.text == REPLY_TEXT, recovered
_assert_kept(unconfigured_rig, recovered, after)
@pytest.mark.timeout(120)
def test_unconfigured_key_level_destination_keeps_datastore_spans(
unconfigured_rig: Rig, langfuse_vars: dict[str, JsonValue]
) -> None:
key: Final = unconfigured_rig.scenario.key(
metadata={
"logging": [
{"callback_name": "langfuse_otel", "callback_type": "success", "callback_vars": dict(langfuse_vars)}
]
}
)
cursors: Final = unconfigured_rig.cursors()
marker: Final = _marker()
sent: Final = unconfigured_rig.raw("chat", marker, stream=False, key=key)
assert sent.text == REPLY_TEXT, sent
assert unconfigured_rig.upstream_hits(marker) == 1
_assert_kept(unconfigured_rig, sent, cursors)
def _assert_tenant_kept_the_burst(rig: Rig, cursors: Cursors, traces: set[str]) -> None:
def ready(spans: tuple[Span, ...]) -> bool:
def trace_kept(trace: str) -> bool:
trace_spans: Final = spans_for_trace(spans, trace)
return any(span["kind"] == SERVER for span in trace_spans) and "redis" in _db_systems(trace_spans)
return all(trace_kept(trace) for trace in traces)
tenant: Final = eventually(
lambda: recorded_spans(rig.sinks.tenant, cursors.tenant)[1],
ready,
seconds=90,
return_last_on_timeout=True,
)
burst: Final = tuple(span for span in tenant if span["trace_id"] in traces)
missing_roots: Final = tuple(
trace for trace in traces if not any(span["kind"] == SERVER for span in spans_for_trace(burst, trace))
)
missing_redis: Final = tuple(trace for trace in traces if "redis" not in _db_systems(spans_for_trace(burst, trace)))
assert not missing_roots, f"SERVER root missing from tenant burst traces: {missing_roots}, {_names(burst)}"
assert not missing_redis, f"redis spans missing from tenant burst traces: {missing_redis}, {_names(burst)}"
@pytest.mark.timeout(300)
def test_unconfigured_tenant_outage_during_a_mixed_burst(unconfigured_rig: Rig) -> None:
cursors: Final = unconfigured_rig.cursors()
configure_sink(unconfigured_rig.sinks.tenant, status=503)
try:
results: Final = _burst(unconfigured_rig, 30)
finally:
configure_sink(unconfigured_rig.sinks.tenant, status=200)
served: Final = _served(results)
assert len(served) == 30, [result for result in results if isinstance(result, str)]
assert all(sent.text == REPLY_TEXT for sent in served), served
traces: Final = _assert_operator_exactly_once(unconfigured_rig, served, cursors)
_assert_tenant_kept_the_burst(unconfigured_rig, cursors, traces)
after: Final = unconfigured_rig.cursors()
_assert_kept(unconfigured_rig, unconfigured_rig.raw("messages", _marker(), stream=True), after)
@pytest.mark.timeout(300)
def test_unconfigured_killing_one_of_two_workers_keeps_the_fan_out(unconfigured_rig: Rig) -> None:
root: Final = psutil.Process(unconfigured_rig.owned.process.pid)
workers: Final = eventually(
lambda: tuple(child for child in root.children() if "resource_tracker" not in " ".join(child.cmdline())),
lambda found: len(found) == 2,
seconds=30,
)
cursors: Final = unconfigured_rig.cursors()
def one(index: int) -> Sent | str:
if index == 6:
os.kill(workers[0].pid, signal.SIGKILL)
try:
return unconfigured_rig.raw("chat", _marker(), stream=index % 2 == 0)
except (httpx.HTTPError, AssertionError) as error:
return repr(error)
with ThreadPoolExecutor(max_workers=6) as pool:
results: Final = tuple(pool.map(one, range(18)))
assert unconfigured_rig.owned.process.poll() is None, "Proxy root exited after a worker was killed"
failures: Final = tuple(result for result in results if isinstance(result, str))
assert all(failure.startswith(("ReadError(", "RemoteProtocolError(", "ConnectError(")) for failure in failures), (
failures
)
assert len(failures) <= 6, failures
settled: Final = tuple(result for index, result in enumerate(results) if index > 12 and isinstance(result, Sent))
traces: Final = _assert_operator_exactly_once(unconfigured_rig, settled, cursors)
_assert_tenant_kept_the_burst(unconfigured_rig, cursors, traces)
after: Final = unconfigured_rig.cursors()
_assert_kept(unconfigured_rig, unconfigured_rig.raw("chat", _marker(), stream=False), after)
@dataclass(frozen=True, slots=True)
class Setting:
otel: Mapping[str, JsonValue]

View file

@ -0,0 +1,421 @@
import json
import uuid
from collections.abc import Callable, Mapping
from hashlib import sha256
from pathlib import Path
from typing import Final
import pytest
from integration._support.client import Gateway, JsonValue, Scenario, eventually, object_value
from integration._support.database import read_rows
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, wire_server
from pydantic import TypeAdapter
def _chat_reply(marker: str) -> dict[str, JsonValue]:
return {
"id": f"chatcmpl-{marker}",
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [{"index": 0, "message": {"role": "assistant", "content": marker}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8},
}
def _anthropic_reply(marker: str) -> dict[str, JsonValue]:
return {
"id": f"msg_{marker}",
"type": "message",
"role": "assistant",
"model": "claude-sonnet-4-5",
"content": [{"type": "text", "text": marker}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 5, "output_tokens": 3},
}
def _chat_stream_frames(marker: str) -> tuple[bytes, ...]:
chunk: Final = {"id": f"chatcmpl-{marker}", "object": "chat.completion.chunk", "created": 1, "model": "gpt-4o-mini"}
return (
f"data: {json.dumps({**chunk, 'choices': [{'index': 0, 'delta': {'role': 'assistant', 'content': marker}}]})}\n\n".encode(),
f"data: {json.dumps({**chunk, 'choices': [{'index': 0, 'delta': {}, 'finish_reason': 'stop'}], 'usage': {'prompt_tokens': 5, 'completion_tokens': 3, 'total_tokens': 8}})}\n\n".encode(),
b"data: [DONE]\n\n",
)
def _spend_row(digest: str, call_type: str) -> dict[str, JsonValue]:
rows: Final = eventually(
lambda: read_rows(
'SELECT request_tags, metadata, team_id FROM "LiteLLM_SpendLogs" WHERE api_key=%s AND call_type=%s',
(digest, call_type),
),
lambda values: len(values) == 1,
seconds=70,
)
return rows[0]
def _spend_row_tagged(tag: str) -> dict[str, JsonValue]:
rows: Final = eventually(
lambda: read_rows(
'SELECT request_tags, metadata, team_id, api_key FROM "LiteLLM_SpendLogs" WHERE request_tags::text LIKE %s',
(f'%"{tag}"%',),
),
lambda values: len(values) == 1,
seconds=70,
)
return rows[0]
def _policy_tags(row: Mapping[str, JsonValue]) -> list[JsonValue]:
raw: Final = row["request_tags"]
tags: Final = json.loads(raw) if isinstance(raw, str) else raw
assert isinstance(tags, list), row
return [tag for tag in tags if not (isinstance(tag, str) and tag.startswith("User-Agent: "))]
def _spend_logs_metadata(row: Mapping[str, JsonValue]) -> JsonValue:
metadata: Final = row["metadata"]
return object_value(json.loads(metadata) if isinstance(metadata, str) else metadata).get("spend_logs_metadata")
def _tagged_key(scenario: Scenario, marker: str, **fields: JsonValue) -> tuple[str, str]:
team: Final = scenario.team(metadata={"tags": [f"team-{marker}"], "spend_logs_metadata": {"team_field": marker}})
project: Final = scenario.project(team, metadata={"tags": [f"project-{marker}"]})
key: Final = scenario.key(
team_id=team,
project_id=project,
metadata={"tags": [f"key-{marker}"], "spend_logs_metadata": {"cost_center": marker}},
**fields,
)
return key, sha256(key.encode()).hexdigest()
def _digest(key: str) -> str:
return sha256(key.encode()).hexdigest()
def _configured_passthrough(gateway: Gateway, scenario: Scenario, marker: str, target: str, *, auth: bool) -> str:
path: Final = f"/integration-passthrough-{marker}"
created: Final = gateway.post("/config/pass_through_endpoint", {"path": path, "target": target, "auth": auth})
endpoints: Final = TypeAdapter(list[JsonValue]).validate_python(created["endpoints"])
endpoint_id: Final = object_value(endpoints[0])["id"]
scenario.cleanups.callback(
lambda: gateway.request("DELETE", "/config/pass_through_endpoint", params={"endpoint_id": str(endpoint_id)})
)
return path
def _responses_reply(marker: str, stream: bool) -> Reply:
response: Final[dict[str, JsonValue]] = {
"id": f"resp_{marker}",
"object": "response",
"created_at": 1,
"status": "completed",
"model": "gpt-4o-mini",
"output": [
{
"id": f"msg_{marker}",
"type": "message",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": marker, "annotations": []}],
}
],
"usage": {"input_tokens": 5, "output_tokens": 3, "total_tokens": 8},
}
if not stream:
return Reply(body=json.dumps(response).encode())
events: Final[tuple[dict[str, JsonValue], ...]] = (
{
"type": "response.created",
"sequence_number": 0,
"response": {**response, "status": "in_progress", "output": []},
},
{
"type": "response.output_text.delta",
"sequence_number": 1,
"item_id": f"msg_{marker}",
"output_index": 0,
"content_index": 0,
"delta": marker,
},
{"type": "response.completed", "sequence_number": 2, "response": response},
)
return Reply(
content_type="text/event-stream",
chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events),
)
def _echo_upstream(marker: str) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
if request.method == "GET" and request.target == "/v1/models":
return Reply(body=json.dumps({"object": "list", "data": []}).encode())
assert request.method == "POST", request
body: Final = object_value(json.loads(request.body))
assert marker in json.dumps(body), request
if request.target == "/v1/responses":
return _responses_reply(marker, body.get("stream") is True)
assert body["messages"] == [{"role": "user", "content": marker}], request
if body.get("stream") is True:
return Reply(chunks=_chat_stream_frames(marker), content_type="text/event-stream")
return Reply(body=json.dumps(_chat_reply(marker)).encode())
return respond
def test_configured_passthrough_spend_row_matches_native_route_tags_and_spend_logs_metadata(gateway: Gateway) -> None:
marker: Final = uuid.uuid4().hex
with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(api_base=wire.url + "/v1")
path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=True)
key, digest = _tagged_key(scenario, marker, models=[model], allowed_passthrough_routes=[path])
headers: Final = {"x-litellm-tags": f"caller-{marker},key-{marker}", "User-Agent": "integration-tags/1"}
body: Final[dict[str, JsonValue]] = {"model": model, "messages": [{"role": "user", "content": marker}]}
native: Final = gateway.request("POST", "/v1/chat/completions", body, key=key, headers=headers)
assert native.status_code == 200, native.text
passthrough: Final = gateway.request("POST", path, body, key=key, headers=headers)
assert passthrough.status_code == 200, passthrough.text
assert json.loads(passthrough.content) == _chat_reply(marker)
native_row: Final = _spend_row(digest, "acompletion")
passthrough_row: Final = _spend_row(digest, "pass_through_endpoint")
expected: Final = [f"key-{marker}", f"team-{marker}", f"project-{marker}", f"caller-{marker}"]
assert _policy_tags(native_row) == expected, native_row
assert _policy_tags(passthrough_row) == expected, passthrough_row
assert _spend_logs_metadata(native_row) == {"cost_center": marker, "team_field": marker}, native_row
assert _spend_logs_metadata(passthrough_row) == {"cost_center": marker, "team_field": marker}, passthrough_row
@pytest.mark.parametrize("bucket", ["metadata", "litellm_metadata"])
def test_configured_passthrough_body_tags_lead_and_body_spend_logs_metadata_wins_over_key_and_team(
gateway: Gateway, bucket: str
) -> None:
marker: Final = uuid.uuid4().hex
with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario:
path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=True)
key, digest = _tagged_key(scenario, marker, allowed_passthrough_routes=[path])
body: Final[dict[str, JsonValue]] = {
"messages": [{"role": "user", "content": marker}],
bucket: {
"tags": [f"body-{marker}", f"team-{marker}"],
"spend_logs_metadata": {"cost_center": f"body-{marker}"},
},
}
response: Final = gateway.request("POST", path, body, key=key)
assert response.status_code == 200, response.text
row: Final = _spend_row(digest, "pass_through_endpoint")
assert _policy_tags(row) == [f"body-{marker}", f"team-{marker}", f"key-{marker}", f"project-{marker}"], row
assert _spend_logs_metadata(row) == {"cost_center": f"body-{marker}", "team_field": marker}, row
def test_configured_passthrough_streaming_upstream_row_carries_key_team_project_and_caller_tags(
gateway: Gateway,
) -> None:
marker: Final = uuid.uuid4().hex
with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario:
path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=True)
key, digest = _tagged_key(scenario, marker, allowed_passthrough_routes=[path])
response: Final = gateway.request(
"POST",
path,
{"stream": True, "messages": [{"role": "user", "content": marker}]},
key=key,
headers={"x-litellm-tags": f"caller-{marker}"},
)
assert response.status_code == 200, response.text
assert response.content == b"".join(_chat_stream_frames(marker)), response.text
row: Final = _spend_row(digest, "pass_through_endpoint")
assert _policy_tags(row) == [f"key-{marker}", f"team-{marker}", f"project-{marker}", f"caller-{marker}"], row
assert _spend_logs_metadata(row) == {"cost_center": marker, "team_field": marker}, row
def test_configured_passthrough_key_outside_any_team_carries_its_own_tags_and_spend_logs_metadata(
gateway: Gateway,
) -> None:
marker: Final = uuid.uuid4().hex
with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario:
path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=True)
key: Final = scenario.key(
allowed_passthrough_routes=[path],
metadata={"tags": [f"key-{marker}"], "spend_logs_metadata": {"cost_center": marker}},
)
response: Final = gateway.request(
"POST",
path,
{"messages": [{"role": "user", "content": marker}]},
key=key,
headers={"x-litellm-tags": f"caller-{marker}"},
)
assert response.status_code == 200, response.text
row: Final = _spend_row(_digest(key), "pass_through_endpoint")
assert _policy_tags(row) == [f"key-{marker}", f"caller-{marker}"], row
assert _spend_logs_metadata(row) == {"cost_center": marker}, row
def test_configured_passthrough_untagged_key_row_keeps_only_caller_tag_and_no_spend_logs_metadata(
gateway: Gateway,
) -> None:
marker: Final = uuid.uuid4().hex
with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario:
path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=True)
team: Final = scenario.team()
key: Final = scenario.key(team_id=team, allowed_passthrough_routes=[path])
response: Final = gateway.request(
"POST",
path,
{"messages": [{"role": "user", "content": marker}]},
key=key,
headers={"x-litellm-tags": f"caller-{marker}"},
)
assert response.status_code == 200, response.text
row: Final = _spend_row(_digest(key), "pass_through_endpoint")
assert _policy_tags(row) == [f"caller-{marker}"], row
assert _spend_logs_metadata(row) is None, row
assert row["team_id"] == team, row
def test_open_passthrough_without_auth_row_carries_only_caller_tag(gateway: Gateway) -> None:
marker: Final = uuid.uuid4().hex
with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario:
path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=False)
response: Final = gateway.client.post(
path,
json={"messages": [{"role": "user", "content": marker}]},
headers={"x-litellm-tags": f"caller-{marker}"},
)
assert response.status_code == 200, response.text
row: Final = _spend_row_tagged(f"caller-{marker}")
assert _policy_tags(row) == [f"caller-{marker}"], row
assert _spend_logs_metadata(row) is None, row
assert row["api_key"] == "", row
@pytest.mark.parametrize(
("metadata", "leading_tags"),
[
({"tags": "string-not-list"}, []),
({"tags": [1, None, "z"]}, [1, None, "z"]),
({"spend_logs_metadata": "string-not-object"}, []),
],
)
def test_configured_passthrough_hostile_body_metadata_shapes_still_carry_key_team_project_tags(
gateway: Gateway, metadata: JsonValue, leading_tags: list[JsonValue]
) -> None:
marker: Final = uuid.uuid4().hex
with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario:
path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=True)
key, digest = _tagged_key(scenario, marker, allowed_passthrough_routes=[path])
response: Final = gateway.request(
"POST", path, {"messages": [{"role": "user", "content": marker}], "metadata": metadata}, key=key
)
assert response.status_code == 200, response.text
row: Final = _spend_row(digest, "pass_through_endpoint")
assert _policy_tags(row) == [*leading_tags, f"key-{marker}", f"team-{marker}", f"project-{marker}"], row
assert _spend_logs_metadata(row) == {"cost_center": marker, "team_field": marker}, row
def test_configured_passthrough_body_cannot_forge_user_api_key_attribution_fields(gateway: Gateway) -> None:
marker: Final = uuid.uuid4().hex
with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario:
path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=True)
forged_team: Final = scenario.team()
key, digest = _tagged_key(scenario, marker, allowed_passthrough_routes=[path])
forged: Final[dict[str, JsonValue]] = {
"user_api_key": "forged-" + marker,
"user_api_key_team_id": forged_team,
"user_api_key_user_id": "forged-" + marker,
"user_api_key_alias": "forged-" + marker,
}
body: Final[dict[str, JsonValue]] = {"messages": [{"role": "user", "content": marker}], "metadata": forged}
response: Final = gateway.request("POST", path, body, key=key)
assert response.status_code == 200, response.text
row: Final = _spend_row(digest, "pass_through_endpoint")
assert row["team_id"] != forged_team, row
assert _policy_tags(row) == [f"key-{marker}", f"team-{marker}", f"project-{marker}"], row
assert read_rows('SELECT api_key FROM "LiteLLM_SpendLogs" WHERE team_id=%s', (forged_team,)) == [], forged_team
@pytest.mark.parametrize("stream", [False, True])
@pytest.mark.parametrize("route", ["/v1/chat/completions", "/v1/messages"])
def test_native_routes_carry_key_team_project_and_caller_tags_and_key_over_team_spend_logs_metadata(
gateway: Gateway, route: str, stream: bool
) -> None:
marker: Final = uuid.uuid4().hex
with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(api_base=wire.url + "/v1")
key, digest = _tagged_key(scenario, marker, models=[model])
body: Final[dict[str, JsonValue]] = {
"model": model,
"max_tokens": 16,
"stream": stream,
"messages": [{"role": "user", "content": marker}],
}
response: Final = gateway.request(
"POST",
route,
body,
key=key,
headers={"x-litellm-tags": f"caller-{marker}", "User-Agent": "integration-tags/1"},
)
assert response.status_code == 200, response.text
rows: Final = eventually(
lambda: read_rows('SELECT request_tags, metadata FROM "LiteLLM_SpendLogs" WHERE api_key=%s', (digest,)),
lambda values: len(values) == 1,
seconds=70,
)
assert _policy_tags(rows[0]) == [f"key-{marker}", f"team-{marker}", f"project-{marker}", f"caller-{marker}"], (
rows
)
assert _spend_logs_metadata(rows[0]) == {"cost_center": marker, "team_field": marker}, rows
def test_anthropic_passthrough_spend_row_carries_key_team_project_tags_and_spend_logs_metadata(
gateway: Gateway, tmp_path: Path
) -> None:
marker: Final = uuid.uuid4().hex
def respond(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/v1/messages", request
assert request.headers["x-api-key"] == "synthetic-anthropic-key"
return Reply(body=json.dumps(_anthropic_reply(marker)).encode())
config: Final = tmp_path / "proxy_config.yaml"
config.write_text(
"model_list: []\n"
"general_settings:\n"
" master_key: os.environ/LITELLM_MASTER_KEY\n"
" database_url: os.environ/DATABASE_URL\n"
" store_model_in_db: true\n"
" disable_spend_logs: false\n"
" proxy_batch_write_at: 1\n"
"router_settings:\n"
" disable_cooldowns: true\n"
)
with wire_server(respond) as wire:
overrides: Final = {"ANTHROPIC_API_BASE": wire.url, "ANTHROPIC_API_KEY": "synthetic-anthropic-key"}
with owned_proxy(gateway, tmp_path, overrides, config=config) as candidate, candidate.scenario() as scenario:
key, digest = _tagged_key(scenario, marker)
response: Final = candidate.request(
"POST",
"/anthropic/v1/messages",
{
"model": "claude-sonnet-4-5",
"max_tokens": 16,
"messages": [{"role": "user", "content": marker}],
},
key=key,
headers={"x-litellm-tags": f"caller-{marker}", "User-Agent": "integration-tags/1"},
)
assert response.status_code == 200, response.text
assert json.loads(response.content) == _anthropic_reply(marker)
row: Final = _spend_row(digest, "pass_through_endpoint")
assert _policy_tags(row) == [f"key-{marker}", f"team-{marker}", f"project-{marker}", f"caller-{marker}"], (
row
)
assert _spend_logs_metadata(row) == {"cost_center": marker, "team_field": marker}, row

View file

@ -63,7 +63,8 @@ def mock_request():
self.method = method
self.request_body = request_body or {}
# Add url attribute that the actual code expects
self.url = "http://localhost:8000/test"
self.url = httpx.URL("http://localhost:8000/test")
self.scope = {"type": "http", "method": method, "path": "/test"}
# Add state attribute that FastAPI requests have
self.state = type("State", (), {})()
@ -414,6 +415,7 @@ PROTOCOL_CONSTRAINED_PASS_THROUGH_ROUTES = {
"/transcribe": {"POST"},
"/transcribe/{operation}": {"POST"},
"/tinyfish/{endpoint:path}": {"GET", "POST"},
"/laya/v1/systemone": {"POST"},
}

View file

@ -80,6 +80,25 @@ async def _turn(
)
async def _benchmark_rows(
db, start: datetime, end: datetime, key: str | None = None, user_id: str | None = None
) -> list[dict]:
return await db.query_raw(
AUTOROUTER_BENCHMARKS_SQL,
start.isoformat(),
end.isoformat(),
key,
user_id,
start.date().isoformat(),
(end - timedelta(days=1)).date().isoformat(),
)
async def _days(db, key: str | None = None, user_id: str | None = None, router: str | None = None) -> list[dict]:
rows = await _benchmark_rows(db, T0 - timedelta(days=1), T0 + timedelta(days=2), key, user_id)
return [row for row in rows if row["turns"] and (router is None or row["router_name"] == router)]
async def _row(db, key: str, session_id: str = "s1", router: str = "auto-1") -> dict:
rows = await db.query_raw(
'SELECT * FROM "LiteLLM_AutoRouterSession" WHERE api_key = $1 AND session_id = $2 AND router_name = $3',
@ -225,18 +244,15 @@ async def test_subtotal_coverage_survives_legacy_and_rolling_writers(db, writers
assert row["savings_estimated_turns"] == sum(writers)
assert row["savings_estimated_actual_spend"] == pytest.approx(0.01 * sum(writers))
assert row["savings_estimated_saved_spend"] == pytest.approx(0.02 * sum(writers))
groups: Final = await db.query_raw(
AUTOROUTER_BENCHMARKS_SQL, T0.isoformat(), (T0 + timedelta(days=1)).isoformat(), key, None
)
assert len(groups) == 1
assert groups[0]["classifier_cost"] == row["classifier_cost"]
assert groups[0]["classifier_cost_recorded_turns"] == sum(writers)
assert groups[0]["turns"] == len(writers)
assert groups[0]["spend"] == row["spend"]
assert groups[0]["saved_spend"] == row["saved_spend"]
assert groups[0]["savings_estimated_turns"] == sum(writers)
assert groups[0]["savings_estimated_actual_spend"] == row["savings_estimated_actual_spend"]
assert groups[0]["savings_estimated_saved_spend"] == row["savings_estimated_saved_spend"]
days: Final = await _days(db, key)
assert len(days) == int(any(writers))
for day in days:
assert day["classifier_cost"] == row["classifier_cost"]
assert day["classifier_cost_recorded_turns"] == day["turns"] == sum(writers)
assert day["spend"] == pytest.approx(0.01 * sum(writers))
assert day["saved_spend"] == pytest.approx(0.02 * sum(writers))
assert day["savings_estimated_actual_spend"] == row["savings_estimated_actual_spend"]
assert day["savings_estimated_saved_spend"] == row["savings_estimated_saved_spend"]
async def test_unknown_and_legacy_turns_preserve_actual_spend_without_entering_the_estimated_cohort(db: Prisma) -> None:
@ -250,13 +266,10 @@ async def test_unknown_and_legacy_turns_preserve_actual_spend_without_entering_t
row: Final = await _row(db, key)
assert row["saved_spend"] == pytest.approx(-0.03)
assert row["savings_estimated_baseline_models"] == {"opus": 1}
groups: Final = await db.query_raw(
AUTOROUTER_BENCHMARKS_SQL, T0.isoformat(), (T0 + timedelta(days=1)).isoformat(), key, None
)
assert len(groups) == 1
for actual in (row, groups[0]):
assert actual["turns"] == 3
assert actual["spend"] == pytest.approx(0.96)
(day,) = await _days(db, key)
assert (row["turns"], day["turns"]) == (3, 2)
assert (row["spend"], day["spend"]) == (pytest.approx(0.96), pytest.approx(0.95))
for actual in (row, day):
assert actual["savings_estimated_turns"] == 1
assert actual["savings_estimated_actual_spend"] == pytest.approx(0.25)
assert actual["savings_estimated_saved_spend"] == pytest.approx(-0.05)
@ -281,25 +294,20 @@ async def test_the_benchmarks_aggregate_reads_only_overlapping_sessions(db):
await _turn(db, key, "A", T0, session_id=in_window, router=router, saved=0.5, spend=0.25, classifier_cost=0.02)
await _turn(db, key, "A", T0 - timedelta(days=40), session_id=out_of_window, router=router, classifier_cost=9.0)
rows = await db.query_raw(
AUTOROUTER_BENCHMARKS_SQL,
(T0 - timedelta(days=1)).isoformat(),
(T0 + timedelta(days=1)).isoformat(),
None,
None,
)
rows = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), None, None)
matching = [row for row in rows if row["router_name"] == router]
assert len(matching) == 1
grouped = matching[0]
assert grouped["router_type"] == "complexity"
assert grouped["sessions"] == 1
assert grouped["turns"] == 2
assert grouped["spend"] == pytest.approx(0.5)
assert grouped["saved_spend"] == pytest.approx(1.0)
assert grouped["classifier_cost"] == pytest.approx(0.03)
assert grouped["classifier_cost_recorded_turns"] == 2
assert grouped["session_turns"] == 2
assert grouped["unordered_turns"] == 1
assert grouped["session_seconds"] == pytest.approx(60.0)
(day,) = await _days(db, router=router)
assert (day["turns"], day["classifier_cost_recorded_turns"]) == (2, 2)
assert day["spend"] == pytest.approx(0.5)
assert day["saved_spend"] == pytest.approx(1.0)
assert day["classifier_cost"] == pytest.approx(0.03)
async def test_the_benchmarks_aggregate_can_filter_to_one_key(db):
@ -309,32 +317,22 @@ async def test_the_benchmarks_aggregate_can_filter_to_one_key(db):
await _turn(db, first_key, "A", T0, router=router, saved=0.5, classifier_cost=0.01)
await _turn(db, second_key, "A", T0, router=router, saved=9.0, classifier_cost=0.09)
rows = await db.query_raw(
AUTOROUTER_BENCHMARKS_SQL,
(T0 - timedelta(days=1)).isoformat(),
(T0 + timedelta(days=1)).isoformat(),
first_key,
None,
)
rows = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), first_key, None)
matching = [row for row in rows if row["router_name"] == router]
assert len(matching) == 1
assert matching[0]["sessions"] == 1
assert matching[0]["saved_spend"] == pytest.approx(0.5)
assert matching[0]["classifier_cost"] == pytest.approx(0.01)
assert matching[0]["classifier_cost_recorded_turns"] == 1
(day,) = await _days(db, first_key, router=router)
assert day["saved_spend"] == pytest.approx(0.5)
assert day["classifier_cost"] == pytest.approx(0.01)
assert day["classifier_cost_recorded_turns"] == 1
unknown_key_rows = await db.query_raw(
AUTOROUTER_BENCHMARKS_SQL,
(T0 - timedelta(days=1)).isoformat(),
(T0 + timedelta(days=1)).isoformat(),
f"k-{uuid.uuid4()}",
None,
)
unknown_key_rows = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), f"k-{uuid.uuid4()}", None)
assert [row for row in unknown_key_rows if row["router_name"] == router] == []
class _BenchmarkRow(TypedDict):
sessions: ReadOnly[int]
session_turns: ReadOnly[int]
turns: ReadOnly[int]
same_model_turns: ReadOnly[int]
first_visit_turns: ReadOnly[int]
@ -350,14 +348,11 @@ class _BenchmarkRow(TypedDict):
async def _scoped_benchmarks(
db: Prisma, router: str, user_id: str | None = None, key: str | None = None
) -> tuple[_BenchmarkRow, ...]:
rows: Final = await db.query_raw(
AUTOROUTER_BENCHMARKS_SQL,
(T0 - timedelta(days=1)).isoformat(),
(T0 + timedelta(days=1)).isoformat(),
key,
user_id,
rows: Final = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), key, user_id)
days: Final = await _days(db, key, user_id, router)
return tuple(
cast(_BenchmarkRow, {**row, **next(iter(days), {})}) for row in rows if row["router_name"] == router
)
return tuple(cast(_BenchmarkRow, row) for row in rows if row["router_name"] == router)
async def test_users_keep_written_identity_across_shared_keys_and_keyless_sessions(db: Prisma) -> None:
@ -384,28 +379,33 @@ async def test_users_keep_written_identity_across_shared_keys_and_keyless_sessio
intersection: Final = await _scoped_benchmarks(db, router, user_id=alice, key=first_key)
assert len(alice_rows) == len(bob_rows) == len(global_rows) == len(key_rows) == len(intersection) == 1
assert (alice_rows[0]["sessions"], alice_rows[0]["turns"], alice_rows[0]["same_model_turns"]) == (3, 4, 1)
assert (alice_rows[0]["session_turns"], bob_rows[0]["session_turns"]) == (4, 2)
assert (bob_rows[0]["sessions"], bob_rows[0]["turns"], bob_rows[0]["first_visit_turns"]) == (2, 2, 2)
assert alice_rows[0]["spend"] == pytest.approx(0.05)
assert bob_rows[0]["spend"] == pytest.approx(0.07)
assert alice_rows[0]["tier_turns"] == {"simple": 1}
assert bob_rows[0]["tier_turns"] == {"complex": 1}
assert (alice_rows[0]["cache_hits"], bob_rows[0]["cache_hits"]) == (1, 0)
assert (global_rows[0]["sessions"], global_rows[0]["turns"]) == (4, 7)
assert (global_rows[0]["sessions"], global_rows[0]["session_turns"], global_rows[0]["turns"]) == (4, 7, 6)
assert (alice_rows[0]["savings_estimated_turns"], bob_rows[0]["savings_estimated_turns"]) == (4, 2)
assert global_rows[0]["savings_estimated_turns"] == 6
for scoped in (alice_rows[0], bob_rows[0]):
assert scoped["savings_estimated_actual_spend"] == pytest.approx(scoped["spend"])
assert scoped["savings_estimated_saved_spend"] == pytest.approx(scoped["saved_spend"])
assert global_rows[0]["spend"] == pytest.approx(alice_rows[0]["spend"] + bob_rows[0]["spend"] + 0.01)
assert global_rows[0]["saved_spend"] == pytest.approx(alice_rows[0]["saved_spend"] + bob_rows[0]["saved_spend"] + 0.02)
assert global_rows[0]["spend"] == pytest.approx(alice_rows[0]["spend"] + bob_rows[0]["spend"])
assert global_rows[0]["saved_spend"] == pytest.approx(alice_rows[0]["saved_spend"] + bob_rows[0]["saved_spend"])
assert global_rows[0]["tier_turns"] == {"simple": 1, "complex": 1}
assert (key_rows[0]["sessions"], key_rows[0]["turns"]) == (1, 3)
assert key_rows[0]["spend"] == pytest.approx(0.05)
assert (key_rows[0]["sessions"], key_rows[0]["session_turns"], key_rows[0]["turns"]) == (1, 3, 2)
assert key_rows[0]["spend"] == pytest.approx(0.04)
assert (intersection[0]["sessions"], intersection[0]["turns"]) == (1, 1)
assert intersection[0]["spend"] == pytest.approx(0.01)
assert await _scoped_benchmarks(db, router, user_id=bob, key=second_key) == ()
assert await _scoped_benchmarks(db, router, user_id=f"u-{uuid.uuid4()}") == ()
assert await _scoped_benchmarks(db, router, user_id="") == ()
assert [
row
for row in await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), None, "")
if row["router_name"] == router
] == []
async def test_a_failed_user_projection_rolls_back_the_keys_increment(db: Prisma) -> None:
@ -419,6 +419,7 @@ async def test_a_failed_user_projection_rolls_back_the_keys_increment(db: Prisma
assert await _row(db, key) == before
assert await db.query_raw('SELECT user_id FROM "LiteLLM_AutoRouterUserSession" WHERE user_id = $1', user_id) == []
assert [day["turns"] for day in await _days(db, key)] == [1]
first_user: Final = f"u-{uuid.uuid4()}"
second_user: Final = f"u-{uuid.uuid4()}"
@ -463,6 +464,14 @@ async def test_a_failed_user_projection_rolls_back_the_keys_increment(db: Prisma
assert (row["turns"], row["same_model_turns"], row["unordered_turns"], row["last_model"]) == (count, 1, 0, model)
assert row["spend"] == pytest.approx(count * 0.01)
assert row["saved_spend"] == pytest.approx(count * 0.02)
days: Final = await db.query_raw(
'SELECT user_id, turns, saved_spend FROM "LiteLLM_AutoRouterDailySpend" WHERE api_key = $1', key
)
assert {day["user_id"]: (day["turns"], day["saved_spend"]) for day in days} == {
"": (1, pytest.approx(0.02)),
first_user: (3, pytest.approx(0.06)),
second_user: (2, pytest.approx(0.04)),
}
async def test_user_session_cleanup_keeps_another_users_recent_keyless_session(db: Prisma) -> None:
@ -490,13 +499,7 @@ async def test_a_reconfigured_alias_reports_each_router_type_as_its_own_group(db
db, key, "A", T0 + timedelta(seconds=10), session_id=f"s-{uuid.uuid4()}", router=router, router_type="quality"
)
rows = await db.query_raw(
AUTOROUTER_BENCHMARKS_SQL,
(T0 - timedelta(days=1)).isoformat(),
(T0 + timedelta(days=1)).isoformat(),
None,
None,
)
rows = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), None, None)
matching = sorted(
(row for row in rows if row["router_name"] == router),
key=lambda row: row["router_type"],
@ -575,16 +578,10 @@ async def test_the_benchmarks_aggregate_sums_tier_turns_across_sessions(db):
await _turn(db, key, "B", T0 + timedelta(seconds=20), session_id=f"s-{uuid.uuid4()}", router=router, tier="complex")
await _turn(db, key, "C", T0 + timedelta(seconds=30), session_id=f"s-{uuid.uuid4()}", router=router, tier=None)
rows = await db.query_raw(
AUTOROUTER_BENCHMARKS_SQL,
(T0 - timedelta(days=1)).isoformat(),
(T0 + timedelta(days=1)).isoformat(),
None,
None,
)
rows = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), None, None)
grouped = next(row for row in rows if row["router_name"] == router)
assert grouped["tier_turns"] == {"simple": 2, "complex": 1}
assert grouped["turns"] == 4
assert grouped["session_turns"] == 4
async def test_tier_maps_stay_separate_per_router_type_on_a_reconfigured_alias(db):
@ -604,13 +601,7 @@ async def test_tier_maps_stay_separate_per_router_type_on_a_reconfigured_alias(d
tier="2",
)
rows = await db.query_raw(
AUTOROUTER_BENCHMARKS_SQL,
(T0 - timedelta(days=1)).isoformat(),
(T0 + timedelta(days=1)).isoformat(),
None,
None,
)
rows = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), None, None)
by_type = {row["router_type"]: row["tier_turns"] for row in rows if row["router_name"] == router}
assert by_type == {"complexity": {"medium": 1}, "quality": {"2": 1}}
@ -620,13 +611,7 @@ async def test_a_window_with_no_tiered_turns_aggregates_to_an_empty_map(db):
router = f"r-{uuid.uuid4()}"
await _turn(db, key, "A", T0, session_id=f"s-{uuid.uuid4()}", router=router, tier=None)
rows = await db.query_raw(
AUTOROUTER_BENCHMARKS_SQL,
(T0 - timedelta(days=1)).isoformat(),
(T0 + timedelta(days=1)).isoformat(),
None,
None,
)
rows = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), None, None)
grouped = next(row for row in rows if row["router_name"] == router)
assert grouped["tier_turns"] == {}
@ -653,3 +638,49 @@ async def test_an_out_of_order_hit_still_counts_toward_the_overall_hit_rate(db):
assert row["unordered_turns"] == 1
assert row["cache_hits"] == 1
assert row["same_model_hits"] + row["first_visit_hits"] + row["return_hits"] == 0
async def test_a_cross_midnight_session_splits_its_money_by_request_day(db):
key = f"k-{uuid.uuid4()}"
router = f"auto-{uuid.uuid4()}"
midnight = datetime(2026, 9, 2)
await _turn(db, key, "A", midnight - timedelta(minutes=10), router=router, spend=1.0, saved=7.0, user_id="u1")
await _turn(db, key, "A", midnight + timedelta(minutes=10), router=router, spend=1.0, saved=3.0, user_id="u1")
await _turn(db, key, "B", midnight + timedelta(days=1), router=router, spend=1.0, saved=11.0, user_id="u1")
assert (await _row(db, key, router=router))["saved_spend"] == 21.0
days = await db.query_raw(
'SELECT date, turns, saved_spend FROM "LiteLLM_AutoRouterDailySpend" WHERE api_key = $1 ORDER BY date', key
)
assert [(d["date"], d["turns"], d["saved_spend"]) for d in days] == [
("2026-09-01", 1, 7.0),
("2026-09-02", 1, 3.0),
("2026-09-03", 1, 11.0),
]
for user_id in (None, "u1"):
(selected,) = await _benchmark_rows(db, midnight, midnight + timedelta(days=1), key, user_id)
assert (selected["sessions"], selected["session_turns"]) == (1, 3)
assert (selected["turns"], selected["spend"], selected["saved_spend"]) == (1, 1.0, 3.0)
async def test_a_router_type_change_within_a_day_keeps_each_types_money_apart(db):
key = f"k-{uuid.uuid4()}"
router = f"auto-{uuid.uuid4()}"
await _turn(db, key, "A", T0, router=router, router_type="complexity", spend=1.0, saved=4.0)
await _turn(db, key, "A", T0 + timedelta(hours=1), router=router, router_type="quality", spend=2.0, saved=0.0)
days = {day["router_type"]: (day["turns"], day["spend"], day["saved_spend"]) for day in await _days(db, key)}
assert days == {"complexity": (1, 1.0, 4.0), "quality": (1, 2.0, 0.0)}
async def test_a_router_type_change_mid_session_keeps_session_shape_with_the_sessions_type(db):
key = f"k-{uuid.uuid4()}"
router = f"auto-{uuid.uuid4()}"
await _turn(db, key, "A", T0, router=router, router_type="complexity", spend=1.0, saved=4.0)
await _turn(db, key, "B", T0 + timedelta(hours=1), router=router, router_type="quality", spend=2.0, saved=0.0)
rows = {row["router_type"]: row for row in await _benchmark_rows(db, T0, T0 + timedelta(days=1), key)}
assert set(rows) == {"complexity", "quality"}
assert (rows["complexity"]["sessions"], rows["complexity"]["session_turns"], rows["complexity"]["turns"]) == (1, 2, 1)
assert (rows["quality"]["sessions"], rows["quality"]["session_turns"], rows["quality"]["turns"]) == (0, 0, 1)
assert rows["quality"]["spend"] == 2.0

View file

@ -151,6 +151,13 @@ async def test_late_replay_updates_all_projections_without_rebilling(db: Prisma,
):
assert after_users["late-user"][field] == after[field]
assert after_users["late-user"]["turns"] == 1 and after_users["late-user"]["spend"] == 0.17
days: Final = await db.query_raw(
'SELECT * FROM "LiteLLM_AutoRouterDailySpend" WHERE api_key=$1 ORDER BY user_id', late.api_key
)
assert [(day["date"], day["user_id"]) for day in days] == [("1970-01-01", "early-user"), ("1970-01-01", "late-user")]
assert days[0]["saved_spend"] == days[0]["savings_estimated_turns"] == 0
for field in ("saved_spend", "savings_estimated_turns", "savings_estimated_actual_spend", "savings_estimated_saved_spend"):
assert days[1][field] == after[field]
for table in ("DailyUserSpend", "DailyTeamSpend", "DailyOrganizationSpend", "DailyEndUserSpend", "DailyAgentSpend", "DailyTagSpend"):
rows: Final = await db.query_raw(f'SELECT spend,api_requests,autorouter_savings_spend FROM "LiteLLM_{table}" WHERE api_key=$1', late.api_key)
assert rows[0]["spend"] == rows[0]["api_requests"] == 0

View file

@ -55,13 +55,94 @@ def _kv(key: str, value: str | int) -> KeyValue:
return KeyValue(key=key, value=AnyValue(string_value=value))
def _export(*spans: Span, service: str = "svc", scope: str = "test") -> bytes:
def _export(*spans: Span, service: str = "svc", scope: str = "test", agent_name: str = "") -> bytes:
resource_spans = ResourceSpans(scope_spans=[ScopeSpans(spans=list(spans))])
resource_spans.resource.attributes.append(_kv("service.name", service))
if agent_name:
resource_spans.resource.attributes.append(_kv("gen_ai.agent.name", agent_name))
resource_spans.scope_spans[0].scope.name = scope
return ExportTraceServiceRequest(resource_spans=[resource_spans]).SerializeToString()
@pytest.mark.parametrize(
("name", "attributes"),
[
("research_agent", {"openinference.span.kind": "AGENT", "metadata": '{"lc_agent_name":"research_agent"}'}),
("research_agent", {"openinference.span.kind": "AGENT", "metadata": '{"ls_integration":"langgraph"}'}),
("research_agent._execute_core", {"openinference.span.kind": "AGENT", "graph.node.id": "research_agent"}),
("agent", {"openinference.span.kind": "AGENT", "gen_ai.agent.name": "research_agent"}),
("openclaw.harness.run", {"openclaw.agent": "research_agent"}),
(
"invoke_agent research_agent",
{"gen_ai.operation.name": "invoke_agent", "gen_ai.agent.name": "research_agent"},
),
],
ids=["deepagents", "langgraph", "crewai", "hermes", "openclaw", "genai"],
)
def test_framework_agent_identity_is_independent_of_service(name: str, attributes: dict[str, str]):
span = _span(name, b"\x02" * 8, **attributes)
row = decode_otlp(_export(span, service="shared-deployment"), "application/x-protobuf")[0]
assert row["AgentName"] == "research_agent"
assert row["ServiceName"] == "shared-deployment"
assert row["SpanName"] == name
@pytest.mark.parametrize("name", ["ClaudeAgentSDK.query", "FunctionAgent.run"])
def test_resource_agent_name_labels_instrumentors_without_an_agent_attribute(name: str):
span = _span(name, b"\x02" * 8, openinference__span__kind="AGENT")
row = decode_otlp(_export(span, agent_name="research_agent"), "application/x-protobuf")[0]
assert row["AgentName"] == "research_agent"
def test_span_agent_name_takes_precedence_over_resource_default():
span = _span("invoke_agent child", b"\x02" * 8, gen_ai__agent__name="child")
row = decode_otlp(_export(span, agent_name="research_agent"), "application/x-protobuf")[0]
assert row["AgentName"] == "child"
@pytest.mark.parametrize(
("scope", "span_name", "configured_name", "expected"),
[
("hermes-otel-plugin", "hermes-agent", "research_agent", "research_agent"),
("hermes-otel-plugin", "child", "research_agent", "child"),
("hermes-otel-plugin", "hermes-agent", "", "hermes-agent"),
("other-plugin", "hermes-agent", "research_agent", "hermes-agent"),
],
)
def test_hermes_resource_name_replaces_only_its_plugin_default(
scope: str, span_name: str, configured_name: str, expected: str
):
span = _span("agent", b"\x02" * 8, gen_ai__agent__name=span_name)
row = decode_otlp(_export(span, scope=scope, agent_name=configured_name), "application/x-protobuf")[0]
assert row["AgentName"] == expected
@pytest.mark.parametrize("agent_name", ["research_agent", ""])
def test_openinference_middleware_is_not_a_separate_agent(agent_name: str):
span = _span(
"PatchToolCallsMiddleware.before_agent", b"\x02" * 8, b"\x01" * 8,
openinference__span__kind="AGENT", metadata=json.dumps({"lc_agent_name": agent_name}),
)
row = decode_otlp(_export(span, scope="openinference.instrumentation.langchain"), "application/x-protobuf")[0]
assert (row["ObservationType"], row["AgentName"]) == ("framework", agent_name)
@pytest.mark.parametrize("scope", ["test", "openinference.instrumentation.langchain"])
@pytest.mark.parametrize("kind", ["CHAIN", "AGENT"])
@pytest.mark.parametrize("metadata", ["not json", "[]", '{"lc_agent_name":null}', "{}"])
def test_unnamed_framework_does_not_invent_an_agent_from_service(metadata: str, scope: str, kind: str):
span = _span("workflow", b"\x02" * 8, openinference__span__kind=kind, metadata=metadata)
row = decode_otlp(_export(span, scope=scope), "application/x-protobuf")[0]
assert row["AgentName"] == ""
@pytest.mark.parametrize("name,expected", [("support", "support"), ("LangGraph", "")])
def test_langgraph_distinguishes_configured_graph_name_from_default(name: str, expected: str):
span = _span(name, b"\x02" * 8, openinference__span__kind="CHAIN", metadata='{"ls_integration":"langgraph"}')
row = decode_otlp(_export(span, scope="openinference.instrumentation.langchain"), "application/x-protobuf")[0]
assert row["AgentName"] == expected
def _span(name: str, span_id: bytes, parent: bytes = b"", **attributes: str | int) -> Span:
return Span(
trace_id=bytes.fromhex(TRACE_ID),

View file

@ -221,6 +221,23 @@ def test_agent_nodes_ignores_spans_of_unknown_agents():
assert agent_nodes(spans) == ()
def test_trace_groups_normalized_names_and_preserves_span_labels():
rows = [
_row("root", "", "invoke_agent research_agent", "agent", "research_agent"),
_row("r1", "root", "researcher._execute_core", "agent", "researcher"),
_row("r2", "r1", "invoke_agent researcher", "agent", "researcher"),
_row("llm", "r2", "chat", "llm", "researcher"),
]
result = trace_from_rows("t1", rows)
assert result is not None
assert result["summary"]["agent_names"] == ("research_agent", "researcher")
assert result["summary"]["name"] == "invoke_agent research_agent"
agents = {agent["name"]: agent for agent in result["agents"]}
assert agents["researcher"]["parent_agent"] == "research_agent"
assert agents["researcher"]["invocations"] == 2
assert agents["researcher"]["llm_calls"] == 1
# ---------------------------------------------------------------- list helpers

View file

@ -12,6 +12,7 @@
import asyncio
import logging
from typing import Final
import pytest
@ -110,6 +111,61 @@ def test_excluded_services_from_env_csv(monkeypatch):
assert OpenTelemetryV2Config().excluded_services == frozenset({"redis", "postgresql"})
@pytest.mark.parametrize("name", ["EXCLUDED_SERVICES", "excluded_services", "Excluded_Services"])
def test_a_bare_excluded_services_env_var_is_ignored(monkeypatch, name):
for env_name in ("LITELLM_OTEL_EXCLUDED_SERVICES", "EXCLUDED_SERVICES", "excluded_services", "Excluded_Services"):
monkeypatch.delenv(env_name, raising=False)
monkeypatch.setenv(name, "redis,postgres")
assert OpenTelemetryV2Config().excluded_services == frozenset()
@pytest.mark.parametrize(
("set_env_name", "env_value", "case_sensitive", "env_ignore_empty", "env_parse_none_str"),
[
pytest.param("otel_service_name", "lower", True, False, None, id="case-sensitive"),
pytest.param("OTEL_SERVICE_NAME", "", False, True, None, id="ignore-empty"),
pytest.param("OTEL_ENDPOINT", "null", False, False, "null", id="parse-none"),
pytest.param("excluded_services", "redis", True, False, None, id="bare-exclusion"),
],
)
def test_env_source_preserves_runtime_options(
monkeypatch: pytest.MonkeyPatch,
set_env_name: str,
env_value: str,
case_sensitive: bool,
env_ignore_empty: bool,
env_parse_none_str: str | None,
) -> None:
for env_name in (
"OTEL_SERVICE_NAME",
"otel_service_name",
"OTEL_ENDPOINT",
"OTEL_EXPORTER_OTLP_ENDPOINT",
"LITELLM_OTEL_EXCLUDED_SERVICES",
"EXCLUDED_SERVICES",
"excluded_services",
"Excluded_Services",
):
monkeypatch.delenv(env_name, raising=False)
monkeypatch.setenv(set_env_name, env_value)
config: Final = OpenTelemetryV2Config(
_case_sensitive=case_sensitive,
_env_ignore_empty=env_ignore_empty,
_env_parse_none_str=env_parse_none_str,
)
assert config.service_name == "litellm"
assert config.endpoint is None
assert config.excluded_services == frozenset()
def test_the_documented_env_var_wins_over_a_bare_excluded_services(monkeypatch):
for env_name in ("LITELLM_OTEL_EXCLUDED_SERVICES", "EXCLUDED_SERVICES", "excluded_services", "Excluded_Services"):
monkeypatch.delenv(env_name, raising=False)
monkeypatch.setenv("EXCLUDED_SERVICES", "postgres")
monkeypatch.setenv("LITELLM_OTEL_EXCLUDED_SERVICES", "redis")
assert OpenTelemetryV2Config().excluded_services == frozenset({"redis"})
def test_excluded_services_config_wins_over_env(monkeypatch):
monkeypatch.setenv("LITELLM_OTEL_EXCLUDED_SERVICES", "redis")
assert OpenTelemetryV2Config(excluded_services=["postgres"]).excluded_services == frozenset({"postgresql"})

View file

@ -1116,6 +1116,18 @@ class TestProviderWiring:
assert self._fan_out_of(preset)._excluded_db_systems == frozenset({"redis"})
@pytest.mark.parametrize("otel", [None, True, "on", "", []], ids=["null", "true", "on", "empty_string", "empty_list"])
def test_a_non_mapping_otel_block_falls_back_to_the_published_logger_config(self, monkeypatch, otel):
monkeypatch.setattr(litellm, "callback_settings", {"otel": otel}, raising=False)
preset = OpenTelemetryV2(
config=OpenTelemetryV2Config(exporters=[ExporterSpec(kind="in_memory")], excluded_services=["redis"]),
callback_name="langfuse_otel",
)
publish_global_otel_v2_provider([], lambda _p: None, registered=preset)
assert self._fan_out_of(preset)._excluded_db_systems == frozenset({"redis"})
def test_otel_after_a_preset_reuses_it_and_still_takes_callback_settings_exclusions(self, monkeypatch):
"""``callbacks: [langfuse_otel, otel]`` keeps one v2 logger, exactly as
before ``excluded_services`` existed, and the exclusion still comes from

View file

@ -1029,6 +1029,84 @@ class TestJWTKeyMappingCascade:
class TestStripPrismaQueryParams:
"""The psycopg URL the job connects with is derived from the Prisma-dialect
DATABASE_URL, whose TLS params mean something else to libpq."""
@staticmethod
def _query(url: str) -> dict[str, str]:
from urllib.parse import parse_qsl, urlparse
return dict(parse_qsl(urlparse(url).query))
def test_prisma_ca_sslcert_becomes_sslrootcert_with_verify_full(self):
url = "postgresql://u:p@writer:5432/db?schema=public&sslmode=require&sslcert=/tmp/pinned.pem&sslaccept=strict"
cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url)
assert self._query(cleaned) == {"sslmode": "verify-full", "sslrootcert": "/tmp/pinned.pem"}
assert cleaned.startswith("postgresql://u:p@writer:5432/db?")
@pytest.mark.parametrize("sslmode", ["prefer", "require"])
@pytest.mark.parametrize("sslaccept", ["strict", "unknown-mode-prisma-treats-as-strict"])
def test_strict_verifies_chain_and_hostname_whatever_sslmode_prisma_was_given(self, sslmode, sslaccept):
url = f"postgresql://writer/db?sslmode={sslmode}&sslcert=/certs/ca.pem&sslaccept={sslaccept}"
cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url)
assert self._query(cleaned) == {"sslmode": "verify-full", "sslrootcert": "/certs/ca.pem"}
def test_strict_with_tls_disabled_stays_off(self):
url = "postgresql://writer/db?sslmode=disable&sslcert=/certs/ca.pem&sslaccept=strict"
cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url)
assert self._query(cleaned) == {"sslmode": "disable"}
@pytest.mark.parametrize("sslaccept", ["&sslaccept=accept_invalid_certs", ""])
def test_without_strict_the_ca_is_dropped_so_libpq_checks_nothing_like_prisma(self, sslaccept):
url = f"postgresql://writer/db?sslmode=require&sslcert=/certs/ca.pem{sslaccept}"
cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url)
assert self._query(cleaned) == {"sslmode": "require"}
def test_a_ca_alone_without_strict_or_sslmode_leaves_libpq_its_defaults(self):
cleaned = ProxyExtrasDBManager._strip_prisma_query_params("postgresql://writer/db?sslcert=/certs/ca.pem")
assert cleaned == "postgresql://writer/db"
def test_a_libpq_client_certificate_pair_is_left_alone(self):
url = "postgresql://writer/db?sslmode=verify-full&sslrootcert=/ca.pem&sslcert=/client.crt&sslkey=/client.key"
cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url)
assert self._query(cleaned) == {
"sslmode": "verify-full",
"sslrootcert": "/ca.pem",
"sslcert": "/client.crt",
"sslkey": "/client.key",
}
def test_an_explicit_sslrootcert_wins_over_the_prisma_sslcert(self):
url = "postgresql://writer/db?sslmode=require&sslrootcert=/ca.pem&sslcert=/pinned.pem&sslaccept=strict"
cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url)
assert self._query(cleaned) == {"sslmode": "verify-full", "sslrootcert": "/ca.pem"}
def test_prisma_only_params_are_dropped_and_plain_urls_pass_through(self):
url = "postgresql://u:p@pooler:6543/db?schema=tenant&pgbouncer=true&connection_limit=5&connect_timeout=3"
cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url)
assert cleaned == "postgresql://u:p@pooler:6543/db?connect_timeout=3"
assert (
ProxyExtrasDBManager._strip_prisma_query_params("postgresql://u:p@writer/db")
== "postgresql://u:p@writer/db"
)
class TestBuildRequestLogIndexes:
"""The migration job hands the index build the direct database URL and the schema
the migrations target, waits for it, and reports its result."""

View file

@ -0,0 +1 @@

View file

@ -0,0 +1,60 @@
from collections.abc import Mapping
from typing import Final
import pytest
from litellm.llms.laya.common_utils import laya_connection, laya_response_model
@pytest.mark.parametrize(
("base", "key", "expected_base", "expected_key"),
[
(None, None, "http://laya.test/root", "laya-env-key"),
("http://custom.test/", None, "http://custom.test", None),
("http://custom.test/", "explicit-key", "http://custom.test", "explicit-key"),
],
)
def test_laya_credentials_stay_with_their_configured_destination(
monkeypatch: pytest.MonkeyPatch,
base: str | None,
key: str | None,
expected_base: str,
expected_key: str | None,
) -> None:
monkeypatch.setenv("LAYA_API_BASE", "http://laya.test/root/")
monkeypatch.setenv("LAYA_API_KEY", "laya-env-key")
monkeypatch.setenv("TYPESAFE_API_KEY", "never-send-this")
connection: Final = laya_connection(base, key)
assert (connection.api_base, connection.api_key) == (expected_base, expected_key)
assert "key" not in repr(connection)
@pytest.mark.parametrize(
"base",
["", "ftp://laya.test", "http://user:password@laya.test", "https://laya.test?key=x", "http://laya.test/#x"],
)
def test_laya_rejects_ambiguous_server_urls(base: str) -> None:
with pytest.raises(ValueError, match="Laya"):
laya_connection(base)
def test_laya_missing_server_does_not_fall_back_to_typesafe(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.delenv("LAYA_API_BASE", raising=False)
monkeypatch.setenv("TYPESAFE_API_BASE", "https://typesafe.test")
with pytest.raises(ValueError, match="LAYA_API_BASE"):
laya_connection()
@pytest.mark.parametrize(
("routing", "requested", "expected"),
[
({"model": "multilingual"}, "english", "multilingual"),
(None, "english", "english"),
({"model": 42}, "english", "english"),
(None, None, "unknown"),
],
)
def test_laya_identity_tracks_the_checkpoint_not_the_shared_agent_name(
routing: Mapping[str, object] | None, requested: str | None, expected: str
) -> None:
assert laya_response_model({"model": "laya-rl-agent", "routing": routing}, requested) == expected

View file

@ -8,7 +8,7 @@ from typing import Optional
from unittest.mock import MagicMock, patch
import pytest
from fastapi import Request
from fastapi import HTTPException, Request
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_utils import (
@ -463,6 +463,24 @@ def test_get_model_from_request_no_request_extracts_model():
)
@pytest.mark.parametrize("model", ["english", "multilingual", "typed-decisions"])
@pytest.mark.parametrize("route", ["/laya/v1/systemone", "/laya/v1/systemone/"])
def test_laya_native_model_uses_the_classifier_permission_identity(model: str, route: str) -> None:
assert get_model_from_request(request_data={"model": model}, route=route) == f"laya/{model}"
@pytest.mark.parametrize("model", [None, "", "auto", "laya/english", "unknown", ["english"], 7])
def test_laya_native_model_cannot_implicitly_select_an_unauthorized_checkpoint(model: object) -> None:
with pytest.raises(HTTPException) as denied:
get_model_from_request(request_data={"model": model}, route="/laya/v1/systemone")
assert denied.value.status_code == 400
def test_laya_model_normalization_does_not_change_other_provider_routes() -> None:
assert get_model_from_request(request_data={"model": "jev-latest"}, route="/typesafe/v1/systemone") == "jev-latest"
assert get_model_from_request(request_data={}, route="/laya/health") is None
def _cache_prediction_router():
from litellm.router import Router

View file

@ -2020,6 +2020,454 @@ async def test_auto_register_binds_api_key_to_token_hash():
assert result.end_user_id == "validated-end-user"
def _auto_register_patches(*, plaintext_key: str | None = "sk-minted-plaintext"):
from litellm.proxy.auth.auth_method import AuthMethod
from litellm.proxy.auth.resolvers.models import CredentialRef
from litellm.proxy.auth.resolvers.store import IdentityStore
from litellm.proxy.proxy_server import hash_token
resolved_key = UserAPIKeyAuth(
token="existing-hash" if plaintext_key is None else hash_token(plaintext_key),
user_id="validated-user",
team_id="validated-team",
org_id="key-own-org",
)
principal = IdentityStore._principal_from_key(
resolved_key,
auth_method=AuthMethod.API_KEY,
credential_ref=CredentialRef(token_id=resolved_key.token),
)
return (
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
new_callable=AsyncMock,
return_value={"token": plaintext_key},
),
patch(
"litellm.proxy.auth.resolvers.store.IdentityStore.resolve",
new_callable=AsyncMock,
return_value=principal,
),
)
def _auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler, **over):
kwargs = {
"virtual_key_claim_field": "sub",
"claim_value": "validated-user",
"jwt_handler": jwt_handler,
"prisma_client": prisma_client,
"user_api_key_cache": user_api_key_cache,
"parent_otel_span": None,
"proxy_logging_obj": MagicMock(),
"cache_key": "jwt_key_mapping:sub:validated-user",
"team_id": "validated-team",
"user_id": "validated-user",
"org_id": "jwt-org",
"end_user_id": "validated-end-user",
}
kwargs.update(over)
return kwargs
@pytest.mark.asyncio
async def test_auto_register_map_existing_key_reuses_users_key_but_never_an_auto_registered_one():
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[
{"token": "auto-registered-hash", "metadata": {"auto_registered": True}},
{"token": "existing-hash", "metadata": {}},
]
)
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
user_api_key_cache = MagicMock()
user_api_key_cache.async_set_cache = AsyncMock()
jwt_handler = MagicMock()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
user_id_jwt_field="sub",
auto_register_map_existing_key=True,
virtual_key_mapping_cache_ttl=300,
)
generate_patch, resolve_patch = _auto_register_patches(plaintext_key=None)
with generate_patch as generate_key, resolve_patch:
result = await _auto_register_jwt_mapping(
**_auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler)
)
generate_key.assert_not_awaited()
create_data = prisma_client.db.litellm_jwtkeymapping.create.await_args.kwargs["data"]
assert create_data["token"] == "existing-hash"
assert create_data["created_by"] == "auto_register"
assert user_api_key_cache.async_set_cache.await_args.kwargs["value"] == "existing-hash"
assert result is not None
assert result.token == "existing-hash"
assert result.api_key == "existing-hash"
assert result.org_id == "key-own-org"
@pytest.mark.asyncio
async def test_auto_register_map_existing_key_mints_when_user_has_no_key():
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
from litellm.proxy.proxy_server import hash_token
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
user_api_key_cache = MagicMock()
user_api_key_cache.async_set_cache = AsyncMock()
jwt_handler = MagicMock()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
user_id_jwt_field="sub",
auto_register_map_existing_key=True,
virtual_key_mapping_cache_ttl=300,
)
generate_patch, resolve_patch = _auto_register_patches()
with generate_patch as generate_key, resolve_patch:
result = await _auto_register_jwt_mapping(
**_auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler)
)
generate_key.assert_awaited_once()
create_data = prisma_client.db.litellm_jwtkeymapping.create.await_args.kwargs["data"]
assert create_data["token"] == hash_token("sk-minted-plaintext")
assert result is not None
assert result.token == hash_token("sk-minted-plaintext")
@pytest.mark.asyncio
async def test_auto_register_default_never_looks_up_existing_keys():
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[{"token": "existing-hash", "metadata": {}}]
)
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
user_api_key_cache = MagicMock()
user_api_key_cache.async_set_cache = AsyncMock()
jwt_handler = MagicMock()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_id_jwt_field="sub", virtual_key_mapping_cache_ttl=300)
generate_patch, resolve_patch = _auto_register_patches()
with generate_patch as generate_key, resolve_patch:
await _auto_register_jwt_mapping(**_auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler))
prisma_client.db.litellm_verificationtoken.find_many.assert_not_awaited()
generate_key.assert_awaited_once()
@pytest.mark.asyncio
async def test_auto_register_map_existing_key_race_loser_keeps_reused_key():
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[{"token": "existing-hash", "metadata": {}}]
)
prisma_client.db.litellm_verificationtoken.delete = AsyncMock()
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock(side_effect=Exception("Unique constraint failed (P2002)"))
user_api_key_cache = MagicMock()
user_api_key_cache.async_set_cache = AsyncMock()
jwt_handler = MagicMock()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
user_id_jwt_field="sub",
auto_register_map_existing_key=True,
virtual_key_mapping_cache_ttl=300,
)
generate_patch, resolve_patch = _auto_register_patches(plaintext_key=None)
with (
generate_patch,
resolve_patch,
patch(
"litellm.proxy.auth.user_api_key_auth.get_jwt_key_mapping_object",
new_callable=AsyncMock,
return_value="winner-hash",
),
):
result = await _auto_register_jwt_mapping(
**_auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler)
)
assert result is not None
assert result.org_id == "key-own-org"
prisma_client.db.litellm_verificationtoken.delete.assert_not_awaited()
assert user_api_key_cache.async_set_cache.await_args.kwargs["value"] == "winner-hash"
@pytest.mark.asyncio
async def test_auto_register_map_existing_key_user_id_none_mints():
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[{"token": "existing-hash", "metadata": {}}]
)
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
user_api_key_cache = MagicMock()
user_api_key_cache.async_set_cache = AsyncMock()
jwt_handler = MagicMock()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
user_id_jwt_field="sub",
auto_register_map_existing_key=True,
virtual_key_mapping_cache_ttl=300,
)
generate_patch, resolve_patch = _auto_register_patches()
with generate_patch as generate_key, resolve_patch:
await _auto_register_jwt_mapping(
**_auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler, user_id=None)
)
prisma_client.db.litellm_verificationtoken.find_many.assert_not_awaited()
generate_key.assert_awaited_once()
@pytest.mark.asyncio
async def test_auto_register_map_existing_key_reuses_when_the_user_was_matched_by_a_fallback_lookup():
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[{"token": "existing-hash", "metadata": {}}]
)
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
user_api_key_cache = MagicMock()
user_api_key_cache.async_set_cache = AsyncMock()
jwt_handler = MagicMock()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
user_id_jwt_field="sub",
user_email_jwt_field="email",
auto_register_map_existing_key=True,
virtual_key_mapping_cache_ttl=300,
)
generate_patch, resolve_patch = _auto_register_patches(plaintext_key=None)
with generate_patch as generate_key, resolve_patch:
await _auto_register_jwt_mapping(
**_auto_register_kwargs(
prisma_client,
user_api_key_cache,
jwt_handler,
claim_value="idp-subject-not-the-db-user-id",
cache_key="jwt_key_mapping:sub:idp-subject-not-the-db-user-id",
)
)
generate_key.assert_not_awaited()
assert prisma_client.db.litellm_jwtkeymapping.create.await_args.kwargs["data"]["token"] == "existing-hash"
@pytest.mark.asyncio
async def test_auto_register_map_existing_key_mints_when_the_claim_is_not_a_user_identity_field():
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
from litellm.proxy.proxy_server import hash_token
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[{"token": "existing-hash", "metadata": {}}]
)
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
user_api_key_cache = MagicMock()
user_api_key_cache.async_set_cache = AsyncMock()
jwt_handler = MagicMock()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
user_id_jwt_field="sub",
auto_register_map_existing_key=True,
virtual_key_mapping_cache_ttl=300,
)
generate_patch, resolve_patch = _auto_register_patches()
with generate_patch as generate_key, resolve_patch:
await _auto_register_jwt_mapping(
**_auto_register_kwargs(
prisma_client,
user_api_key_cache,
jwt_handler,
virtual_key_claim_field="azp",
claim_value="shared-client-app",
cache_key="jwt_key_mapping:azp:shared-client-app",
)
)
prisma_client.db.litellm_verificationtoken.find_many.assert_not_awaited()
generate_key.assert_awaited_once()
assert prisma_client.db.litellm_jwtkeymapping.create.await_args.kwargs["data"]["token"] == hash_token(
"sk-minted-plaintext"
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("issuer_user_id_field", "expect_reuse"),
[("uid", False), (None, True)],
)
async def test_auto_register_map_existing_key_uses_the_issuers_own_user_field_over_the_global_one(
issuer_user_id_field, expect_reuse
):
from litellm.proxy._types import JWTIssuerConfig
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
from litellm.proxy.proxy_server import hash_token
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[{"token": "existing-hash", "metadata": {}}]
)
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
user_api_key_cache = MagicMock()
user_api_key_cache.async_set_cache = AsyncMock()
jwt_handler = MagicMock()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
user_id_jwt_field="sub",
auto_register_map_existing_key=True,
virtual_key_mapping_cache_ttl=300,
issuers=[
JWTIssuerConfig(
issuer="https://idp.example.com", audience="litellm", user_id_jwt_field=issuer_user_id_field
)
],
)
generate_patch, resolve_patch = _auto_register_patches()
with generate_patch as generate_key, resolve_patch:
await _auto_register_jwt_mapping(
**_auto_register_kwargs(
prisma_client, user_api_key_cache, jwt_handler, jwt_issuer="https://idp.example.com"
)
)
mapped_token = prisma_client.db.litellm_jwtkeymapping.create.await_args.kwargs["data"]["token"]
assert mapped_token == ("existing-hash" if expect_reuse else hash_token("sk-minted-plaintext"))
assert generate_key.await_count == (0 if expect_reuse else 1)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("map_existing_key", "master_key", "reused_key_models", "expect_denied"),
[
(True, "sk-master", ["some-other-model"], True),
(True, "sk-master", [], False),
(False, "sk-master", ["some-other-model"], False),
(True, None, ["some-other-model"], False),
],
)
async def test_auto_register_map_existing_key_first_request_runs_key_checks(
map_existing_key: bool, master_key: str | None, reused_key_models: list[str], expect_denied: bool
) -> None:
jwt_token = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.signature"
user_api_key_cache = DualCache()
prisma_client = MagicMock()
jwt_handler = MagicMock()
jwt_handler.is_jwt.return_value = True
jwt_handler.auth_jwt = AsyncMock(return_value={"sub": "user1"})
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
virtual_key_claim_field="sub",
virtual_key_mapping_cache_ttl=300,
auto_register_map_existing_key=map_existing_key,
)
reused_key = UserAPIKeyAuth(
token="hashed-existing-key",
api_key="hashed-existing-key",
user_id="validated-user",
team_id="validated-team",
models=reused_key_models,
)
mock_jwt_result = {
"is_proxy_admin": False,
"team_object": None,
"user_object": LiteLLM_UserTable(user_id="validated-user", user_role="internal_user"),
"end_user_object": None,
"org_object": None,
"token": jwt_token,
"team_id": "validated-team",
"user_id": "validated-user",
"user_email": None,
"end_user_id": None,
"org_id": None,
"team_membership": None,
"jwt_claims": {"sub": "user1"},
}
mock_request = MagicMock()
mock_request.url.path = "/v1/chat/completions"
mock_request.method = "POST"
mock_request.headers = {"authorization": f"Bearer {jwt_token}"}
mock_request.query_params = {}
mock_request.state = SimpleNamespace()
with (
patch("litellm.proxy.proxy_server.general_settings", {"enable_jwt_auth": True}),
patch("litellm.proxy.proxy_server.premium_user", True),
patch("litellm.proxy.proxy_server.master_key", master_key),
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
patch("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache),
patch(
"litellm.proxy.proxy_server.proxy_logging_obj",
MagicMock(post_call_failure_hook=AsyncMock(return_value=None)),
),
patch("litellm.proxy.proxy_server.jwt_handler", jwt_handler),
patch(
"litellm.proxy.auth.user_api_key_auth._resolve_jwt_to_virtual_key",
new_callable=AsyncMock,
return_value=_PendingAutoRegister(
claim_field="sub",
claim_value="user1",
cache_key="jwt_key_mapping:sub:user1",
),
),
patch(
"litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder",
new_callable=AsyncMock,
return_value=mock_jwt_result,
),
patch(
"litellm.proxy.auth.user_api_key_auth._auto_register_jwt_mapping",
new_callable=AsyncMock,
return_value=reused_key,
),
):
call = _user_api_key_auth_builder(
request=mock_request,
api_key=jwt_token,
azure_api_key_header="",
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
request_data={"model": "gpt-4o-mini"},
)
if expect_denied:
with pytest.raises(ProxyException, match="not available for this API key"):
await call
return
result = await call
assert result.api_key == "hashed-existing-key"
assert result.user_id == "validated-user"
assert result.team_id == "validated-team"
assert result.models == reused_key_models
@pytest.mark.asyncio
@pytest.mark.parametrize("active", [True, False])
async def test_auto_register_first_request_propagates_user_email(active: bool) -> None:

View file

@ -619,6 +619,18 @@ def test_classifier_plugin_is_not_settable_over_http():
_request("what is 2+2", classifier_type="custom", classifier_plugin="my_module.instance")
def _benchmark_db(rows: Sequence[Mapping[str, object]], recorded: float | None = None) -> SimpleNamespace:
"""The joined benchmark statement returns the rows as given; any other statement is the Overall total."""
from litellm.proxy.db.autorouter_session_rollup import AUTOROUTER_BENCHMARKS_SQL
total: Final = recorded if recorded is not None else sum(float(row.get("saved_spend") or 0.0) for row in rows)
async def query_raw(sql: str, *params: object) -> Sequence[Mapping[str, object]]:
return rows if sql == AUTOROUTER_BENCHMARKS_SQL else ({"saved": total},)
return SimpleNamespace(db=SimpleNamespace(query_raw=AsyncMock(side_effect=query_raw)))
class TestAutoRouterBenchmarks:
from litellm.proxy.management_endpoints.auto_router_endpoints import _SessionAggRow
@ -635,15 +647,12 @@ class TestAutoRouterBenchmarks:
rows: Sequence[Mapping[str, object]],
model_list: Sequence[object],
api_key: str | None = None,
recorded: float | None = None,
) -> AutoRouterBenchmarksResponse:
from litellm.proxy import proxy_server
from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_benchmarks
class _DB:
async def query_raw(self, sql: str, *params: object):
return rows
monkeypatch.setattr(proxy_server, "prisma_client", type("P", (), {"db": _DB()})())
monkeypatch.setattr(proxy_server, "prisma_client", _benchmark_db(rows, recorded))
monkeypatch.setattr(proxy_server, "llm_router", type("R", (), {"model_list": model_list})())
return await get_auto_router_benchmarks(
user_api_key_dict=ADMIN,
@ -657,6 +666,7 @@ class TestAutoRouterBenchmarks:
router_type="complexity",
tier_turns={},
sessions=4,
session_turns=40,
turns=40,
unordered_turns=1,
covered_turns=38,
@ -703,7 +713,6 @@ class TestAutoRouterBenchmarks:
assert totals.baseline_spend == 40.0
assert totals.saved_pct == 75.0
assert totals.savings_estimated_classifier_cost == 0.4
assert totals.saved_per_session == 7.5
assert totals.cache.coverage_pct == 95.0
assert totals.cache.hit_rate_pct == pytest.approx(73.7)
assert totals.cache.same_model.hit_rate_pct == 95.0
@ -742,7 +751,6 @@ class TestAutoRouterBenchmarks:
assert (totals.spend, totals.saved_spend, totals.baseline_spend, totals.saved_pct) == (10.0, 30.0, 40.0, 75.0)
assert (totals.savings_estimated_turns, totals.savings_estimated_actual_spend) == (40, 10.0)
assert totals.savings_estimated_classifier_cost == 0.4
assert totals.saved_per_session == 7.5
@pytest.mark.asyncio
@pytest.mark.parametrize("router_type, saved", [("adaptive", 0.0), ("quality", 0.0), ("quality", 2.0)])
@ -772,9 +780,50 @@ class TestAutoRouterBenchmarks:
totals: Final = response.totals
assert (totals.turns, totals.spend) == (50, 13.0)
assert (totals.savings_estimated_turns, totals.savings_estimated_actual_spend) == (40, 10.0)
assert (totals.saved_spend, totals.baseline_spend, totals.saved_pct) == (30.0, 40.0, 75.0)
assert totals.unattributed_saved_spend is None
assert (totals.saved_spend, totals.baseline_spend, totals.saved_pct) == (
(30.0, 40.0, 75.0) if saved == 0.0 else (32.0, None, None)
)
assert totals.savings_estimated_classifier_cost == 0.4
@pytest.mark.asyncio
@pytest.mark.parametrize("recorded, unattributed", [(30.0, None), (33.0, 3.0), (27.0, -3.0)])
async def test_the_headline_is_the_overall_daily_total_and_untracked_savings_void_the_baseline(
self, recorded: float, unattributed: float | None, monkeypatch: pytest.MonkeyPatch
) -> None:
response: Final = await self._benchmarks(
monkeypatch, rows=[self.ROW.model_dump()], model_list=[], recorded=recorded
)
totals: Final = response.totals
assert (totals.saved_spend, totals.unattributed_saved_spend) == (recorded, unattributed)
assert (totals.baseline_spend, totals.saved_pct) == ((40.0, 75.0) if unattributed is None else (None, None))
group: Final = response.groups[0]
assert group.saved_spend == 30.0
assert (group.baseline_spend, group.saved_pct) == ((40.0, 75.0) if unattributed is None else (None, None))
@pytest.mark.asyncio
async def test_a_window_holding_only_untracked_history_shows_no_router_baseline(
self, monkeypatch: pytest.MonkeyPatch
) -> None:
history_only: Final = self.ROW.model_dump(
exclude={
"turns",
"spend",
"saved_spend",
"savings_estimated_turns",
"savings_estimated_actual_spend",
"savings_estimated_classifier_cost",
"savings_estimated_saved_spend",
"classifier_cost",
"classifier_cost_recorded_turns",
}
)
response: Final = await self._benchmarks(monkeypatch, rows=[history_only], model_list=[], recorded=3.0)
assert (response.totals.saved_spend, response.totals.unattributed_saved_spend) == (3.0, 3.0)
group: Final = response.groups[0]
assert (group.sessions, group.turns, group.saved_spend) == (4, 0, 0.0)
assert (group.baseline_spend, group.saved_pct) == (None, None)
def test_an_empty_window_folds_to_zeros(self):
from litellm.proxy.management_endpoints.auto_router_endpoints import (
_benchmark_totals,
@ -805,7 +854,7 @@ class TestAutoRouterBenchmarks:
"savings_estimated_classifier_cost": 0.0,
}
)
summed = _summed_agg_row([self.ROW, other])
summed = _summed_agg_row([self.ROW, other.model_copy(update={"session_turns": 10})])
totals = _benchmark_totals(summed)
assert summed.sessions == 5
assert summed.turns == 50
@ -868,6 +917,28 @@ class TestAutoRouterBenchmarks:
assert response.status_code == 422
query.assert_not_awaited()
@pytest.mark.asyncio
async def test_an_empty_key_filter_is_rejected_before_querying_deployment_data(
self, monkeypatch: pytest.MonkeyPatch
) -> None:
import httpx
from fastapi import FastAPI
from litellm.proxy import proxy_server
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_benchmarks
query: Final = AsyncMock(return_value=[])
monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=SimpleNamespace(query_raw=query)))
app: Final = FastAPI()
app.get("/auto_router/benchmarks")(get_auto_router_benchmarks)
app.dependency_overrides[user_api_key_auth] = lambda: ADMIN
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://test") as client:
response: Final = await client.get("/auto_router/benchmarks", params={"api_key": ""})
assert response.status_code == 422
query.assert_not_awaited()
@pytest.mark.asyncio
async def test_a_reversed_window_is_rejected(self, monkeypatch: pytest.MonkeyPatch):
from litellm.proxy import proxy_server
@ -891,15 +962,8 @@ class TestAutoRouterBenchmarks:
from litellm.proxy import proxy_server
from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_benchmarks
captured: dict = {}
class _DB:
async def query_raw(self, sql: str, *params: object):
captured["sql"] = sql
captured["params"] = params
return [TestAutoRouterBenchmarks.ROW.model_dump()]
monkeypatch.setattr(proxy_server, "prisma_client", type("P", (), {"db": _DB()})())
prisma_client: Final = _benchmark_db([TestAutoRouterBenchmarks.ROW.model_dump()])
monkeypatch.setattr(proxy_server, "prisma_client", prisma_client)
response = await get_auto_router_benchmarks(
user_api_key_dict=UserAPIKeyAuth(user_role=role, api_key="sk-admin", user_id="viewer"),
@ -908,7 +972,11 @@ class TestAutoRouterBenchmarks:
api_key="key-hash",
user_id=user_id,
)
assert captured["params"] == ("2026-07-01T00:00:00", "2026-08-02T00:00:00", "key-hash", user_id)
params: Final = tuple(call.args[1:] for call in prisma_client.db.query_raw.await_args_list)
assert params == (
("2026-07-01T00:00:00", "2026-08-02T00:00:00", "key-hash", user_id, "2026-07-01", "2026-08-01"),
("2026-07-01", "2026-08-01", *(([user_id],) if user_id else ()), ["key-hash"]),
)
assert response.routers_in_scope == 1
assert response.groups[0].router_name == "live-auto"
assert response.groups[0].saved_pct == response.totals.saved_pct == 75.0
@ -946,7 +1014,6 @@ class TestAutoRouterBenchmarks:
assert response.totals.saved_spend == 29.5
assert response.totals.baseline_spend == 41.5
assert response.totals.saved_pct == 71.1
assert response.totals.saved_per_session == 5.9
@pytest.mark.asyncio
@pytest.mark.parametrize(
@ -958,11 +1025,9 @@ class TestAutoRouterBenchmarks:
from litellm.proxy import proxy_server
from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_benchmarks
class _DB:
async def query_raw(self, sql: str, *params: object):
return [{**TestAutoRouterBenchmarks.ROW.model_dump(), "tier_turns": wire_value}]
monkeypatch.setattr(proxy_server, "prisma_client", type("P", (), {"db": _DB()})())
monkeypatch.setattr(
proxy_server, "prisma_client", _benchmark_db([{**TestAutoRouterBenchmarks.ROW.model_dump(), "tier_turns": wire_value}])
)
response = await get_auto_router_benchmarks(
user_api_key_dict=ADMIN,
@ -1006,7 +1071,7 @@ class TestAutoRouterBenchmarks:
0.0,
0.0,
)
assert (idle.saved_pct, idle.saved_per_session, idle.avg_turns_per_session) == (0.0, 0.0, 0.0)
assert (idle.saved_pct, idle.avg_turns_per_session) == (0.0, 0.0)
assert (idle.cache.hit_rate_pct, idle.cache.coverage_pct) == (0.0, 0.0)
assert idle.cache.same_model.turns == idle.cache.return_to_tier.hits == 0
assert idle.tier_turns == {}
@ -3724,3 +3789,18 @@ async def test_availability_waits_for_the_first_complete_catalog(monkeypatch):
with pytest.raises(HTTPException) as error:
await auto_router_endpoints.get_auto_router_availability(AutoRouterAvailabilityRequest(), ADMIN)
assert error.value.status_code == 503
class TestPerSessionAverages:
@pytest.mark.parametrize(
"sessions, turns, expected",
[(4, 40, (10.0, 100.0, 1000.0)), (0, 0, (0.0, 0.0, 0.0)), (0, 3, (None, None, None))],
)
def test_requests_without_session_rows_have_unknown_averages_not_zero(
self, sessions: int, turns: int, expected: tuple[float | None, ...]
) -> None:
from litellm.proxy.management_endpoints.auto_router_endpoints import _benchmark_totals
row: Final = TestAutoRouterBenchmarks.ROW.model_copy(update={"sessions": sessions, "turns": turns})
totals: Final = _benchmark_totals(row)
assert (totals.avg_turns_per_session, totals.avg_session_seconds, totals.avg_tokens_per_session) == expected

View file

@ -8,6 +8,8 @@ from typing import Dict, Final, Optional
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
from fastapi.encoders import jsonable_encoder
from fastapi.testclient import TestClient
from litellm._uuid import uuid
@ -7515,6 +7517,96 @@ class TestTeamMemberAutoRouterWrites:
"model_info": {"id": "allowed-id"},
}])
@staticmethod
def _classifier_config(classifier: Mapping[str, object], legacy: bool) -> Mapping[str, object]:
return {
"classifier_type": "jev" if legacy else "oss_classifier",
"tiers": {"SIMPLE": "allowed"},
"jev_classifier_config" if legacy else "opensource_classifier_config": classifier,
}
@pytest.mark.asyncio
@pytest.mark.parametrize("team_id", [None, "member-team"])
@pytest.mark.parametrize(
"legacy,provider,model",
[(True, "typesafe", "jev-latest"), (False, "jev", "jev-latest"), (True, "laya", "english"), (False, "laya", "english")],
)
async def test_classifier_create_stores_only_canonical_configuration(
self, team_id: str | None, legacy: bool, provider: str, model: str
) -> None:
from litellm.proxy.management_endpoints.model_management_endpoints import add_new_model
row: Final = self._row()
database: Final = self._database(self._team(), row)
classifier: Final = {
"provider": provider, "model": model,
"api_base": "https://decision.test", "api_key": "stored-secret",
}
deployment: Final = Deployment(
model_name="new-classifier-router",
litellm_params=LiteLLM_Params(
model="auto_router/complexity_router",
complexity_router_config=self._classifier_config(classifier, legacy),
),
model_info=ModelInfo(id=row.model_id, team_id=team_id),
)
actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
with (
self._environment(database, row),
patch("litellm.proxy.proxy_server.proxy_config.add_deployment", new=AsyncMock(return_value=ReconcileOutcome( # test-quality-ok: [TQ008] model reload I/O boundary
still_desired=frozenset((row.model_id,)), live_after=frozenset((row.model_id,))
))),
patch("litellm.proxy.management_endpoints.model_management_endpoints.append_team_models", new=AsyncMock()), # test-quality-ok: [TQ008] team allowlist persistence boundary
):
await add_new_model(deployment, actor)
written: Final = database.db.litellm_proxymodeltable.create.await_args.kwargs["data"]
saved: Final = json.loads(written["litellm_params"])["complexity_router_config"]
assert saved == {
"classifier_type": "oss_classifier",
"tiers": {"SIMPLE": "allowed"},
"opensource_classifier_config": {**classifier, "provider": "laya" if provider == "laya" else "jev"},
}
@pytest.mark.asyncio
@pytest.mark.parametrize("endpoint", ["create", "patch", "legacy"])
@pytest.mark.parametrize("legacy_config", [None, {"provider": "laya", "model": "english"}])
async def test_ambiguous_classifier_blocks_are_rejected_before_persistence(
self, endpoint: str, legacy_config: Mapping[str, object] | None
) -> None:
from litellm.proxy.management_endpoints.model_management_endpoints import add_new_model
row: Final = self._row()
database: Final = self._database(self._team(), row)
config: Final = {
**self._classifier_config({"provider": "laya", "model": "english"}, False),
"jev_classifier_config": legacy_config,
}
actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
request: Final = updateDeployment(
litellm_params=updateLiteLLMParams(complexity_router_config=config), model_info=ModelInfo(id=row.model_id),
)
operation: Final = (
add_new_model(
Deployment(
model_name="ambiguous-classifier-router",
litellm_params=LiteLLM_Params(model="auto_router/complexity_router", complexity_router_config=config),
model_info=ModelInfo(id=row.model_id),
),
actor,
)
if endpoint == "create"
else patch_model(row.model_id, request, actor)
if endpoint == "patch"
else update_model(request, actor)
)
with self._environment(database, row), pytest.raises(ProxyException) as denied:
await operation
assert denied.value.code == "400"
assert "opensource_classifier_config" in denied.value.message
assert "jev_classifier_config" in denied.value.message
database.db.litellm_proxymodeltable.create.assert_not_awaited()
database.db.litellm_proxymodeltable.update.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("endpoint,change", [("patch", "config"), ("legacy", "strategy"), ("patch", "unrelated")])
async def test_admin_router_changes_release_member_scope(self, endpoint: str, change: str) -> None:
@ -7546,15 +7638,16 @@ class TestTeamMemberAutoRouterWrites:
@pytest.mark.asyncio
@pytest.mark.parametrize("endpoint", ["patch", "legacy"])
@pytest.mark.parametrize("stored_legacy,supplied_legacy", [(True, True), (True, False), (False, True), (False, False)])
@pytest.mark.parametrize("change", ["save", "rotate", "move", "move-without-key", "reset", "heuristic"])
async def test_jev_dashboard_save_preserves_server_transport(self, endpoint: str, change: str) -> None:
async def test_jev_dashboard_save_preserves_server_transport(
self, endpoint: str, change: str, stored_legacy: bool, supplied_legacy: bool
) -> None:
original: Final = self._row()
transport: Final = {"api_key": "synthetic-original-jev-key", "api_base": "https://jev.example.com"}
stored_config: Final = {
"classifier_type": "jev",
"tiers": {"SIMPLE": "allowed"},
"jev_classifier_config": {**transport, "instructions": "Old instructions", "timeout_ms": 6100},
}
stored_config: Final = self._classifier_config(
{**transport, "instructions": "Old instructions", "timeout_ms": 6100}, stored_legacy
)
row: Final = original.model_copy(
update={
"litellm_params": {
@ -7572,11 +7665,11 @@ class TestTeamMemberAutoRouterWrites:
"reset": {"api_key": None, "api_base": None},
"heuristic": {},
}[change]
config: Final = {
"tiers": {"SIMPLE": "allowed"},
"classifier_type": "heuristic" if change == "heuristic" else "jev",
**({} if change == "heuristic" else {"jev_classifier_config": {"timeout_ms": 8100, **overrides}}),
}
config: Final = (
{"tiers": {"SIMPLE": "allowed"}, "classifier_type": "heuristic"}
if change == "heuristic"
else self._classifier_config({"timeout_ms": 8100, **overrides}, supplied_legacy)
)
request: Final = updateDeployment(
litellm_params=updateLiteLLMParams(complexity_router_config=config),
model_info=ModelInfo(id=row.model_id),
@ -7597,12 +7690,179 @@ class TestTeamMemberAutoRouterWrites:
expected: Final = (
config
if change == "heuristic"
else {**config, "jev_classifier_config": {**transport, "timeout_ms": 8100, **overrides}}
else {
"classifier_type": "oss_classifier",
"tiers": {"SIMPLE": "allowed"},
"opensource_classifier_config": {**transport, "timeout_ms": 8100, **overrides},
}
)
assert saved == expected
assert row.litellm_params["complexity_router_config"] == stored_config
assert request.litellm_params.complexity_router_config == config
@pytest.mark.asyncio
@pytest.mark.parametrize("endpoint", ["patch", "legacy"])
@pytest.mark.parametrize("stored_legacy,supplied_legacy", [(True, True), (True, False), (False, True), (False, False)])
@pytest.mark.parametrize(
"stored_provider,stored_base,supplied,expected_transport",
[
("laya", "https://decision.test", {"provider": "laya", "model": "english"}, {"api_base": "https://decision.test", "api_key": "stored-secret"}),
("laya", "https://decision.test", {"provider": "laya", "model": "english", "api_base": "https://decision.test"}, {"api_base": "https://decision.test", "api_key": "stored-secret"}),
(
"laya",
"https://decision.test",
{"provider": "laya", "model": "english", "api_key": None},
{"api_base": "https://decision.test"},
),
("laya", "https://decision.test", {"provider": "laya", "model": "english", "api_base": "https://new.test"}, {}),
("laya", "https://decision.test", {"provider": "laya", "model": "english", "api_base": None}, {}),
("laya", None, {"provider": "laya", "model": "english", "api_base": None}, {}),
("laya", "https://decision.test", {"provider": "typesafe", "model": "jev-latest"}, {}),
(
"laya", "https://decision.test", {"model": "english", "timeout_ms": 8100},
{"provider": "laya", "api_base": "https://decision.test", "api_key": "stored-secret"},
),
("typesafe", "https://decision.test", {"provider": "laya", "model": "english"}, {}),
(
"typesafe", "https://decision.test", {"provider": "jev", "model": "jev-latest"},
{"api_base": "https://decision.test", "api_key": "stored-secret"},
),
(
"jev", "https://decision.test", {"provider": "typesafe", "model": "jev-latest"},
{"api_base": "https://decision.test", "api_key": "stored-secret"},
),
],
)
async def test_decision_provider_changes_cannot_reuse_a_stored_key(
self, endpoint: str, stored_provider: str, stored_base: str | None,
supplied: Mapping[str, object], expected_transport: Mapping[str, object],
stored_legacy: bool, supplied_legacy: bool,
) -> None:
original: Final = self._row()
row: Final = original.model_copy(update={"litellm_params": {
"model": "auto_router/complexity_router",
"complexity_router_config": self._classifier_config(
{
"provider": stored_provider, "model": "english" if stored_provider == "laya" else "jev-latest",
"api_base": stored_base, "api_key": "stored-secret",
},
stored_legacy,
),
}})
database: Final = self._database(self._team(), row)
config: Final = self._classifier_config(supplied, supplied_legacy)
request: Final = updateDeployment(
litellm_params=updateLiteLLMParams(complexity_router_config=config), model_info=ModelInfo(id=row.model_id),
)
actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
with self._environment(database, row):
await (patch_model(row.model_id, request, actor) if endpoint == "patch" else update_model(request, actor))
written: Final = database.db.litellm_proxymodeltable.update.await_args.kwargs["data"]
saved: Final = json.loads(written["litellm_params"])["complexity_router_config"]
expected_provider: Final = supplied.get("provider", stored_provider)
assert saved == {
"classifier_type": "oss_classifier",
"tiers": {"SIMPLE": "allowed"},
"opensource_classifier_config": {
**expected_transport, **supplied,
"provider": "jev" if expected_provider == "typesafe" else expected_provider,
},
}
@pytest.mark.asyncio
@pytest.mark.parametrize("endpoint", ["patch", "legacy"])
@pytest.mark.parametrize(
"string_params,reset_field,config_shape",
[
(False, None, "full"), (True, None, "full"), (False, "api_key", "full"),
(False, "api_base", "full"), (False, None, "omit-provider"),
(False, None, "omit-config"), (False, None, "null-config"),
],
)
async def test_member_save_protects_stored_classifier_connection(
self, endpoint: str, string_params: bool, reset_field: str | None, config_shape: str
) -> None:
original: Final = self._row()
config: Final = {
"classifier_type": "jev", "tiers": {"SIMPLE": "allowed"},
"jev_classifier_config": {"provider": "laya", "model": "english"},
}
secret_params: Final = {
"model": "auto_router/complexity_router",
"complexity_router_config": {
**config, "jev_classifier_config": {
**config["jev_classifier_config"], "api_key": "retained-laya-secret", "api_base": "https://laya.test",
},
},
}
row: Final = original.model_copy(update={"litellm_params": secret_params})
team: Final = self._team().model_copy(update={"models": ["allowed", "laya/english"]})
database: Final = self._database(team, row)
database.transaction.litellm_proxymodeltable.update.return_value = row.model_copy(
update={"litellm_params": json.dumps(secret_params) if string_params else secret_params}
)
supplied_config: Final = {
**config, "jev_classifier_config": {
**{
key: value for key, value in config["jev_classifier_config"].items()
if key != "provider" or config_shape != "omit-provider"
},
**({reset_field: None} if reset_field is not None else {}),
},
}
patch_params: Final = (
{"complexity_router_default_model": "allowed"}
if config_shape == "omit-config"
else {"complexity_router_config": None, "complexity_router_default_model": "allowed"}
if config_shape == "null-config"
else {"complexity_router_config": supplied_config}
)
request: Final = updateDeployment(
litellm_params=updateLiteLLMParams.model_validate(patch_params),
model_info=ModelInfo(id=row.model_id, team_id="member-team"),
)
actor: Final = UserAPIKeyAuth(
user_id="owner", user_role=LitellmUserRoles.INTERNAL_USER, models=["allowed", "laya/english"], config={"timeout": 60},
)
with self._environment(database, row):
if reset_field is not None:
expected_error: Final = HTTPException if endpoint == "patch" else ProxyException
with pytest.raises(expected_error, match="Team members cannot change classifier connections") as denied:
await (
patch_model(row.model_id, request, actor) if endpoint == "patch" else update_model(request, actor)
)
assert (
denied.value.status_code if isinstance(denied.value, HTTPException) else int(denied.value.code)
) == 403
database.transaction.litellm_proxymodeltable.update.assert_not_awaited()
assert row.litellm_params == secret_params
return
response: Final = await (patch_model(row.model_id, request, actor) if endpoint == "patch" else update_model(request, actor))
written: Final = database.transaction.litellm_proxymodeltable.update.await_args.kwargs["data"]
saved_config: Final = json.loads(written["litellm_params"])["complexity_router_config"]
untouched: Final = config_shape in ("omit-config", "null-config")
saved: Final = saved_config["jev_classifier_config" if untouched else "opensource_classifier_config"]
assert saved == secret_params["complexity_router_config"]["jev_classifier_config"]
assert saved_config["classifier_type"] == ("jev" if untouched else "oss_classifier")
if untouched:
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
assert decrypt_value_helper(
json.loads(written["litellm_params"])["complexity_router_default_model"],
key="complexity_router_default_model", return_original_value=True,
) == "allowed"
response_payload: Final = jsonable_encoder(response)
assert "retained-laya-secret" not in json.dumps(response_payload)
response_params: Final = json.loads(response_payload["litellm_params"]) if string_params else response_payload["litellm_params"]
assert response_params == {
**secret_params, "complexity_router_config": {
**config, "jev_classifier_config": {
**config["jev_classifier_config"], "api_key": "REDACTED", "api_base": "https://laya.test",
},
},
}
assert "retained-laya-secret" in row.model_dump_json()
@pytest.mark.asyncio
@pytest.mark.parametrize("endpoint", ["patch", "legacy"])
@pytest.mark.parametrize("access", ["owner", "peer", "limited-key"])

View file

@ -1,3 +1,4 @@
import json
from collections.abc import Mapping
from dataclasses import dataclass
from typing import Final
@ -137,33 +138,48 @@ def test_tier_config_is_normalized_and_unknown_router_extras_are_rejected() -> N
@pytest.mark.parametrize(
("jev_override", "rejected_at"),
[
({"api_base": "https://collector.invalid"}, "jev_classifier_config"),
({"api_base": "https://collector.invalid"}, "opensource_classifier_config"),
({"api_key": "sk-member"}, "api_key"),
({"api_base": "https://collector.invalid", "api_key": "sk-member"}, "api_key"),
({"api_base": "https://collector.invalid", "api_key": ""}, "jev_classifier_config.api_key"),
({"api_base": "https://collector.invalid", "api_key": ""}, "opensource_classifier_config.api_key"),
({"provider": "laya", "model": "english", "api_base": "https://collector.invalid"}, "api_base"),
({"provider": "laya", "model": "english", "api_key": "sk-member"}, "api_key"),
],
)
@pytest.mark.parametrize("legacy", [False, True])
def test_members_cannot_move_the_jev_classifier_off_the_proxys_typesafe_account(
jev_override: Mapping[str, str], rejected_at: str
jev_override: Mapping[str, str], rejected_at: str, legacy: bool
) -> None:
with pytest.raises(HTTPException) as denied:
validate_member_auto_router_config(
{"tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev", "jev_classifier_config": jev_override}
{
"tiers": {"SIMPLE": "allowed"},
"classifier_type": "jev" if legacy else "oss_classifier",
"jev_classifier_config" if legacy else "opensource_classifier_config": jev_override,
}
)
assert denied.value.status_code == 400
assert denied.value.detail == f"Invalid member auto-router configuration at {rejected_at}."
def test_members_can_still_tune_the_jev_classifier() -> None:
@pytest.mark.parametrize(("provider", "model"), [("typesafe", "jev-preview"), ("laya", "english")])
@pytest.mark.parametrize("legacy", [False, True])
def test_members_can_still_tune_the_jev_classifier(provider: str, model: str, legacy: bool) -> None:
validated: Final = validate_member_auto_router_config(
{
"tiers": {"SIMPLE": "allowed"},
"classifier_type": "jev",
"jev_classifier_config": {"model": "jev-preview", "timeout_ms": 500},
"classifier_type": "jev" if legacy else "oss_classifier",
"jev_classifier_config" if legacy else "opensource_classifier_config": {
"provider": provider, "model": model, "timeout_ms": 500,
},
}
)
assert validated.jev_classifier_config is not None
assert (validated.jev_classifier_config.model, validated.jev_classifier_config.timeout_ms) == ("jev-preview", 500)
assert (
validated.jev_classifier_config.provider,
validated.jev_classifier_config.model,
validated.jev_classifier_config.timeout_ms,
) == ("jev" if provider == "typesafe" else provider, model, 500)
assert validate_member_auto_router_config(validated.model_dump()).jev_classifier_config is not None
@ -217,6 +233,92 @@ async def test_member_updates_restrict_fields_and_preserve_an_inherited_default(
assert granted.default_model == "allowed"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"nested,expected_identity,restricted",
[
("omit-config", "laya/english", False),
("omit-config", "laya/english", True),
("omit-block", None, False),
(None, None, False),
({}, None, False),
({"timeout_ms": 500}, None, False),
({"model": "english", "timeout_ms": 500}, "laya/english", False),
({"model": "english", "timeout_ms": 500}, "laya/english", True),
({"model": "multilingual"}, "laya/multilingual", False),
({"provider": "typesafe", "model": "jev-latest"}, "typesafe/jev-latest", False),
({"provider": "typesafe", "model": "jev-latest"}, "typesafe/jev-latest", True),
],
)
async def test_member_authorization_and_persistence_resolve_the_same_classifier(
catalog: Router, monkeypatch: pytest.MonkeyPatch, nested: object, expected_identity: str | None, restricted: bool
) -> None:
from litellm.proxy.management_endpoints.model_management_endpoints import (
_strategy_router_write_violation,
update_db_model,
)
from litellm.types.management_endpoints.auto_router_endpoints import RequestComplexityRouterConfig
monkeypatch.setenv("LITELLM_SALT_KEY", "member-router-test-salt")
stored_config: Final = {
"classifier_type": "jev", "tiers": {"SIMPLE": "allowed"},
"jev_classifier_config": {
"provider": "laya", "model": "english", "timeout_ms": 12000,
"api_base": "https://laya.test", "api_key": "stored-classifier-key",
},
}
existing: Final = Deployment(
model_name="member-router",
litellm_params=LiteLLM_Params(model="auto_router/complexity_router", complexity_router_config=stored_config),
model_info=ModelInfo(id="router-a", team_id="team-a"), created_by="owner",
)
incoming_config: Final = (
None if nested == "omit-config" else {
"classifier_type": "jev", "tiers": {"SIMPLE": "allowed"},
**({} if nested == "omit-block" else {"jev_classifier_config": nested}),
}
)
patch: Final = updateDeployment.model_validate({"litellm_params": {
"complexity_router_config": incoming_config, "complexity_router_default_model": "allowed",
}})
operation: Final = authorize_member_auto_router_write(
incoming=patch, existing=existing, user_api_key_dict=_actor(
models=["allowed"] if restricted or expected_identity is None else ["allowed", expected_identity],
),
team=_team(models=["allowed", "laya/english", "laya/multilingual", "typesafe/jev-latest"]),
premium_user=True, prisma_client=_Client(), llm_router=catalog,
)
violation: Final = _strategy_router_write_violation(patch.litellm_params, existing.litellm_params)
if expected_identity is None:
assert violation is not None
with pytest.raises(HTTPException) as rejected:
await operation
assert rejected.value.status_code == 400
return
assert violation is None
if restricted:
with pytest.raises(ProxyException, match=expected_identity):
await operation
return
grant: Final = await operation
persisted: Final = update_db_model(existing, patch)
saved: Final = RequestComplexityRouterConfig.model_validate(
json.loads(persisted["litellm_params"])["complexity_router_config"]
)
assert grant.config == saved
assert saved.jev_classifier_config is not None
assert (
"typesafe" if saved.jev_classifier_config.provider == "jev" else saved.jev_classifier_config.provider
) + f"/{saved.jev_classifier_config.model}" == expected_identity
assert saved.jev_classifier_config.api_key == (
"stored-classifier-key" if expected_identity.startswith("laya/") else None
)
assert saved.jev_classifier_config.timeout_ms == (
12000 if nested == "omit-config" else 500 if nested == {"model": "english", "timeout_ms": 500} else 3000
)
assert existing.litellm_params.complexity_router_config == stored_config
@pytest.mark.asyncio
@pytest.mark.parametrize("target", ["missing", "nested"])
async def test_member_dependencies_require_plain_configured_models(target: str) -> None:
@ -246,13 +348,17 @@ async def test_member_dependencies_require_plain_configured_models(target: str)
@pytest.mark.asyncio
@pytest.mark.parametrize("restricted", ["key", "team", None])
@pytest.mark.parametrize(("provider", "model"), [("typesafe", "jev-latest"), ("laya", "english")])
async def test_jev_evaluation_requires_model_access_but_no_completion_deployment(
catalog: Router, restricted: str | None
catalog: Router, restricted: str | None, provider: str, model: str
) -> None:
permitted: Final = ["allowed", "typesafe/jev-latest"]
permitted: Final = ["allowed", f"{provider}/{model}"]
operation: Final = authorize_member_auto_router_dependencies(
config=validate_member_auto_router_config(
{"tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev", "jev_classifier_config": {}}
{
"tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev",
"jev_classifier_config": {"provider": provider, "model": model},
}
),
default_model=None,
user_api_key_dict=_actor(models=["allowed"] if restricted == "key" else permitted),
@ -261,17 +367,20 @@ async def test_jev_evaluation_requires_model_access_but_no_completion_deployment
llm_router=catalog,
)
if restricted is not None:
with pytest.raises(ProxyException, match="jev-latest"):
with pytest.raises(ProxyException, match=model):
await operation
return
await operation
assert not catalog.get_model_list("typesafe/jev-latest")
assert not catalog.get_model_list(f"{provider}/{model}")
@pytest.mark.asyncio
@pytest.mark.parametrize("restricted", ["member", "project", "organization", None])
async def test_jev_evaluation_obeys_each_containing_scope(catalog: Router, restricted: str | None) -> None:
allowed: Final = ["allowed", "typesafe/jev-latest"]
@pytest.mark.parametrize(("provider", "model"), [("typesafe", "jev-latest"), ("laya", "english")])
async def test_jev_evaluation_obeys_each_containing_scope(
catalog: Router, restricted: str | None, provider: str, model: str
) -> None:
allowed: Final = ["allowed", f"{provider}/{model}"]
membership: Final = LiteLLM_TeamMembership.model_validate(
{
"user_id": "owner",
@ -293,7 +402,10 @@ async def test_jev_evaluation_obeys_each_containing_scope(catalog: Router, restr
)
operation: Final = authorize_member_auto_router_dependencies(
config=validate_member_auto_router_config(
{"tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev", "jev_classifier_config": {}}
{
"tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev",
"jev_classifier_config": {"provider": provider, "model": model},
}
),
default_model=None,
user_api_key_dict=_actor(models=allowed, project_id="project-a"),
@ -303,8 +415,8 @@ async def test_jev_evaluation_obeys_each_containing_scope(catalog: Router, restr
dependency_objects=MemberAutoRouterDependencyObjects(membership, organization, project),
)
if restricted is not None:
with pytest.raises(ProxyException, match="jev-latest"):
with pytest.raises(ProxyException, match=model):
await operation
return
await operation
assert not catalog.get_model_list("typesafe/jev-latest")
assert not catalog.get_model_list(f"{provider}/{model}")

View file

@ -1,10 +1,12 @@
from datetime import datetime
from typing import Final
from unittest.mock import MagicMock
import httpx
import pytest
import litellm
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.typesafe_passthrough_logging_handler import (
TypeSafePassthroughLoggingHandler,
)
@ -137,6 +139,84 @@ def test_success_handler_dispatches_to_typesafe_handler():
assert normalized["kwargs"]["model"] == "typesafe/jev-1.13.0"
@pytest.mark.asyncio
@pytest.mark.parametrize("guardrail_cost", [0.0, 0.25])
@pytest.mark.parametrize("metadata_slot", ["metadata", "litellm_metadata"])
@pytest.mark.parametrize("routing_model", ["multilingual", None])
async def test_laya_gateway_accounts_for_checkpoint_usage_and_registered_cost(
monkeypatch: pytest.MonkeyPatch, routing_model: str | None, metadata_slot: str, guardrail_cost: float
) -> None:
checkpoint: Final = routing_model or "english"
model: Final = f"laya/{checkpoint}"
input_rate: Final = 0.002
output_rate: Final = 0.005
monkeypatch.setitem(litellm.model_cost, model, {
"input_cost_per_token": input_rate, "output_cost_per_token": output_rate,
"litellm_provider": "laya", "mode": "evaluation",
})
start: Final = datetime.now()
logging_obj: Final = Logging(
model="english", messages=[], stream=False, call_type="pass_through_endpoint",
start_time=start, litellm_call_id="laya-accounting", function_id="laya-accounting", kwargs={},
)
from fastapi import Request
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import HttpPassThroughEndpointHelpers
request: Final = Request({
"type": "http", "method": "POST", "path": "/laya/v1/systemone",
"headers": [], "query_string": b"",
})
auth: Final = UserAPIKeyAuth(
api_key="laya-budget-key", token="laya-budget-key",
model_max_budget={"laya/english": {"budget_limit": 0.01, "time_period": "1d"}},
)
request_body: Final = {"model": "english", metadata_slot: {"model_group": "unbounded-client-choice"}}
logging_kwargs: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
request=request, user_api_key_dict=auth, logging_obj=logging_obj,
passthrough_logging_payload={"url": "https://laya.test/v1/systemone"}, _parsed_body=request_body,
)
logging_kwargs["litellm_params"]["metadata"]["standard_logging_guardrail_information"] = [
{"guardrail_name": "trusted-hook", "guardrail_cost": guardrail_cost},
]
logging_obj.update_environment_variables(
model="english", user="unknown", optional_params={},
litellm_params=logging_kwargs["litellm_params"], call_type="pass_through_endpoint",
)
body: Final = {
"model": "laya-rl-agent", "usage": {"input_tokens": 10, "output_tokens": 3},
**({"routing": {"model": routing_model}} if routing_model else {}),
}
normalized: Final = PassThroughEndpointLogging().normalize_llm_passthrough_logging_payload(
httpx_response=httpx.Response(200, request=httpx.Request("POST", "https://laya.test/v1/systemone"), json=body),
response_body=body, request_body={"model": "english"}, logging_obj=logging_obj,
url_route="https://laya.test/v1/systemone", result="{}", start_time=start,
end_time=datetime.now(), cache_hit=False, custom_llm_provider="laya", **logging_kwargs,
)
logged: Final = normalized["kwargs"]
expected_cost: Final = 10 * input_rate + 3 * output_rate
assert (logged["model"], logged["custom_llm_provider"]) == (model, "laya")
assert logged["response_cost"] == pytest.approx(expected_cost)
assert logged["combined_usage_object"].model_dump(exclude_none=True) == {
"prompt_tokens": 10, "completion_tokens": 3, "total_tokens": 13,
}
assert logging_obj.model_call_details["model"] == model
assert logging_obj.model_call_details["response_cost"] == pytest.approx(expected_cost)
assert logged["standard_logging_object"]["model"] == model
assert logged["standard_logging_object"]["model_group"] == "laya/english"
assert logged["standard_logging_object"]["response_cost"] == pytest.approx(expected_cost + guardrail_cost)
from litellm.caching.caching import DualCache
from litellm.exceptions import BudgetExceededError
from litellm.proxy.hooks.model_max_budget_limiter import _PROXY_VirtualKeyModelMaxBudgetLimiter
budget_limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(DualCache())
assert await budget_limiter.is_key_within_model_budget(auth, "laya/english")
await budget_limiter.async_log_success_event(logged, None, start, datetime.now())
with pytest.raises(BudgetExceededError):
await budget_limiter.is_key_within_model_budget(auth, "laya/english")
def test_openrouter_decisions_response_is_priced_from_request_model_registry_row():
logging_obj = _logging_obj()
model_cost = litellm.model_cost["openrouter/typesafe/jev-1.13"]

View file

@ -23,6 +23,8 @@ from starlette.datastructures import FormData
import litellm
from litellm.caching.caching import DualCache
from litellm.types.utils import CallTypesLiteral
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from tests.unit.llms.bedrock.event_loop_probe import EventLoopProbe
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
@ -7407,6 +7409,152 @@ class TestTypeSafePassthroughRoute:
)
class TestLayaPassthroughRoute:
@pytest.fixture
def client(self, monkeypatch: pytest.MonkeyPatch) -> Iterator[TestClient]:
from litellm.proxy.proxy_server import app
monkeypatch.setenv("LAYA_API_BASE", "http://laya.test/base")
monkeypatch.setenv("TYPESAFE_API_KEY", "never-send-typesafe-key")
monkeypatch.delenv("LAYA_API_KEY", raising=False)
monkeypatch.delenv("SERVER_ROOT_PATH", raising=False)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
monkeypatch.setitem(app.dependency_overrides, user_api_key_auth, lambda: UserAPIKeyAuth(api_key="sk-virtual"))
yield TestClient(app)
@pytest.mark.parametrize("api_key", [None, "laya-provider-key"])
def test_laya_forwards_native_decisions_without_gateway_or_typesafe_credentials(
self, client: TestClient, monkeypatch: pytest.MonkeyPatch, api_key: str | None
) -> None:
if api_key is not None:
monkeypatch.setenv("LAYA_API_KEY", api_key)
body: Final = {
"model": "english",
"state": "refund",
"questions": {"department": {"type": "choice", "criteria": {"billing": "refunds"}}},
}
answer: Final = {"model": "laya-rl-agent", "routing": {"model": "english"}, "answers": {}}
with respx.mock(assert_all_called=True) as upstream:
route: Final = upstream.post("http://laya.test/base/v1/systemone?trace=yes").respond(200, json=answer)
response: Final = client.post(
"/laya/v1/systemone?trace=yes",
json=body,
headers={"Authorization": "Bearer sk-virtual", "x-pass-authorization": "Bearer attacker"},
)
assert (response.status_code, response.json()) == (200, answer)
sent: Final = route.calls.last.request
assert sent.headers.get("authorization") == (f"Bearer {api_key}" if api_key else None)
assert json.loads(sent.content) == body
def test_laya_missing_server_fails_without_contacting_another_provider(
self, client: TestClient, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.delenv("LAYA_API_BASE")
with respx.mock(assert_all_called=False) as upstream:
response: Final = client.post("/laya/v1/systemone", json={"model": "english"})
assert response.status_code == 503
assert "LAYA_API_BASE" in response.text
assert len(upstream.calls) == 0
def test_laya_does_not_forward_unsupported_endpoints(self, client: TestClient) -> None:
with respx.mock(assert_all_called=False) as upstream:
response: Final = client.post("/laya/v1/evaluate", json={"model": "english"})
assert response.status_code == 404
assert len(upstream.calls) == 0
@pytest.mark.parametrize("model", [None, "auto", "jev-latest"])
def test_laya_rejects_implicit_checkpoint_selection(self, client: TestClient, model: str | None) -> None:
with respx.mock(assert_all_called=False) as upstream:
response: Final = client.post("/laya/v1/systemone", json={"model": model})
assert response.status_code == 400
assert len(upstream.calls) == 0
@pytest.mark.parametrize(
"controls",
[{"custom_body": {"model": "multilingual", "state": "refund"}}, {"stream": True}, {"stream": "true"}],
)
def test_laya_rejects_controls_that_change_authorized_body_or_usage_accounting(
self, client: TestClient, controls: Mapping[str, object]
) -> None:
with respx.mock(assert_all_called=False) as upstream:
route: Final = upstream.post("http://laya.test/base/v1/systemone").respond(200, json={"answers": {}})
response: Final = client.post("/laya/v1/systemone", json={"model": "english", **controls})
assert response.status_code == 400
assert not route.called
@pytest.mark.parametrize("metadata_slot", ["metadata", "litellm_metadata"])
def test_laya_hooks_enforce_canonical_model_limits_and_keep_native_wire_body(
self, client: TestClient, monkeypatch: pytest.MonkeyPatch, metadata_slot: str
) -> None:
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3
from litellm.proxy.utils import InternalUsageCache
from litellm.proxy.proxy_server import app
cache: Final = DualCache()
limiter: Final = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(cache))
auth: Final = UserAPIKeyAuth(
api_key="laya-native-rpm", metadata={"model_rpm_limit": {"laya/english": 1}},
)
def authenticated_key() -> UserAPIKeyAuth:
return auth
monkeypatch.setitem(app.dependency_overrides, user_api_key_auth, authenticated_key)
class LimitHook(CustomLogger):
async def async_pre_call_hook(
self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache,
data: dict[str, object], call_type: CallTypesLiteral,
) -> dict[str, object]:
assert data["model"] == "laya/english"
metadata: Final = data.get(metadata_slot)
assert isinstance(metadata, dict)
assert "standard_logging_guardrail_information" not in metadata
assert metadata["customer_label"] == "retained"
await limiter.async_pre_call_hook(user_api_key_dict, cache, data, call_type)
return data
monkeypatch.setattr(litellm, "callbacks", [LimitHook()])
body: Final = {
"model": "english", "state": "refund",
metadata_slot: {
"customer_label": "retained", "model_group": "unbounded-client-choice",
"standard_logging_guardrail_information": [{"guardrail_cost": 25.0}],
},
}
with respx.mock(assert_all_called=True) as upstream:
route: Final = upstream.post("http://laya.test/base/v1/systemone").respond(200, json={"answers": {}})
first: Final = client.post("/laya/v1/systemone", json=body)
second: Final = client.post("/laya/v1/systemone", json=body)
assert first.status_code == 200, first.text
assert second.status_code == 429, second.text
assert route.call_count == 1
assert json.loads(route.calls.last.request.content) == {"model": "english", "state": "refund"}
def test_laya_preserves_trusted_hook_checkpoint_changes(
self, client: TestClient, monkeypatch: pytest.MonkeyPatch
) -> None:
from litellm.integrations.custom_logger import CustomLogger
class CheckpointHook(CustomLogger):
async def async_pre_call_hook(
self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache,
data: dict[str, object], call_type: CallTypesLiteral,
) -> dict[str, object]:
assert data["model"] == "laya/english"
return {**data, "model": "laya/multilingual"}
monkeypatch.setattr(litellm, "callbacks", [CheckpointHook()])
with respx.mock(assert_all_called=True) as upstream:
route: Final = upstream.post("http://laya.test/base/v1/systemone").respond(200, json={"answers": {}})
response: Final = client.post("/laya/v1/systemone", json={"model": "english", "state": "refund"})
assert response.status_code == 200, response.text
assert json.loads(route.calls.last.request.content) == {"model": "multilingual", "state": "refund"}
class TestFalAIPassthroughRoute:
@pytest.fixture
def client(self, monkeypatch: pytest.MonkeyPatch) -> Iterator[TestClient]:

View file

@ -1470,7 +1470,7 @@ async def test_pass_through_request_contains_proxy_server_request_in_kwargs():
# Create mock request
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://test-proxy.com/api/endpoint"
mock_request.url = httpx.URL("http://test-proxy.com/api/endpoint")
mock_request.body = AsyncMock(return_value=b'{"message": "test request"}')
mock_request.headers = Headers({})
mock_request.query_params = QueryParams({})
@ -1575,7 +1575,7 @@ async def test_pass_through_request_streaming_marks_logging_obj_as_stream():
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://test-proxy.com/v1/messages"
mock_request.url = httpx.URL("http://test-proxy.com/v1/messages")
mock_request.body = AsyncMock(return_value=b'{"model": "claude-3", "stream": true}')
mock_request.headers = Headers({})
mock_request.query_params = QueryParams({})
@ -1637,7 +1637,7 @@ async def test_pass_through_request_sse_response_marks_logging_obj_as_stream():
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://test-proxy.com/v1/messages"
mock_request.url = httpx.URL("http://test-proxy.com/v1/messages")
mock_request.body = AsyncMock(return_value=b'{"model": "claude-3"}')
mock_request.headers = Headers({})
mock_request.query_params = QueryParams({})
@ -2507,7 +2507,7 @@ async def test_pass_through_request_query_params_forwarding():
# Create mock request with query parameters (Azure API version)
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://localhost:4000/azure-assistant/openai/assistants"
mock_request.url = httpx.URL("http://localhost:4000/azure-assistant/openai/assistants")
mock_request.body = AsyncMock(return_value=json.dumps(test_body).encode())
mock_request.headers = Headers({"Content-Type": "application/json"})
@ -3016,7 +3016,7 @@ async def test_bedrock_router_passthrough_metadata_initialization():
# Create mock request with headers
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://localhost:4000/bedrock/model/my-model/invoke"
mock_request.url = httpx.URL("http://localhost:4000/bedrock/model/my-model/invoke")
mock_request.headers = Headers(
{
"content-type": "application/json",
@ -3850,7 +3850,7 @@ def _lit3538_request():
r = MagicMock()
r.method = "POST"
r.query_params = {}
r.url = "http://testserver/mock/echo"
r.url = httpx.URL("http://testserver/mock/echo")
r.state = SimpleNamespace()
headers = MagicMock()
headers.copy.return_value = {}
@ -3983,7 +3983,7 @@ async def test_pass_through_request_non_streaming_upstream_error_returned_unchan
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://test-proxy.com/mock-upstream/api/denied"
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/denied")
mock_request.body = AsyncMock(return_value=b'{"action": "read"}')
mock_request.headers = Headers({"content-type": "application/json"})
mock_request.query_params = QueryParams({})
@ -4069,7 +4069,7 @@ async def test_pass_through_request_upstream_error_failure_hook_exception_is_swa
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://test-proxy.com/mock-upstream/api/denied"
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/denied")
mock_request.body = AsyncMock(return_value=b'{"action": "read"}')
mock_request.headers = Headers({"content-type": "application/json"})
mock_request.query_params = QueryParams({})
@ -4118,7 +4118,7 @@ async def test_pass_through_request_streaming_upstream_error_returned_unchanged(
mock_request = MagicMock(spec=Request)
mock_request.method = "GET"
mock_request.url = "http://test-proxy.com/mock-upstream/api/stream-denied"
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/stream-denied")
mock_request.body = AsyncMock(return_value=b"")
mock_request.headers = Headers({})
mock_request.query_params = QueryParams({})
@ -4169,7 +4169,7 @@ class _UpstreamErrorBodyStream(httpx.AsyncByteStream):
def _upstream_error_request() -> MagicMock:
mock_request: Final = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://test-proxy.com/mock-upstream/v1beta/models/claude-nope-9:generateContent"
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/v1beta/models/claude-nope-9:generateContent")
mock_request.body = AsyncMock(return_value=b'{"contents": []}')
mock_request.headers = Headers({"content-type": "application/json"})
mock_request.query_params = QueryParams({})
@ -4966,7 +4966,7 @@ async def test_pass_through_request_non_streaming_success_unchanged():
mock_request = MagicMock(spec=Request)
mock_request.method = "GET"
mock_request.url = "http://test-proxy.com/mock-upstream/api/success"
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/success")
mock_request.body = AsyncMock(return_value=b"")
mock_request.headers = Headers({})
mock_request.query_params = QueryParams({})
@ -5029,7 +5029,7 @@ async def test_pass_through_request_claims_the_budget_reservation_only_when_its_
mock_get_client.return_value = MagicMock(client=async_client)
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://test-proxy.com/mock-upstream/api/generate"
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/generate")
mock_request.body = AsyncMock(return_value=b'{"prompt": "hi"}')
mock_request.headers = Headers({"content-type": "application/json"})
mock_request.query_params = QueryParams({})
@ -5081,7 +5081,7 @@ async def test_pass_through_request_leaves_the_budget_reservation_for_the_reques
mock_get_client.return_value = MagicMock(client=async_client)
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://test-proxy.com/mock-upstream/api/generate"
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/generate")
mock_request.body = AsyncMock(return_value=b'{"prompt": "hi"}')
mock_request.headers = Headers({"content-type": "application/json"})
mock_request.query_params = QueryParams({})
@ -5112,7 +5112,7 @@ async def test_pass_through_request_internal_failure_still_raises_proxy_exceptio
mock_request = MagicMock(spec=Request)
mock_request.method = "GET"
mock_request.url = "http://test-proxy.com/mock-upstream/api/success"
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/success")
mock_request.body = AsyncMock(return_value=b"")
mock_request.headers = Headers({})
mock_request.query_params = QueryParams({})
@ -5213,7 +5213,7 @@ def _enter_relay_logging_mocks(stack, parsed_body):
def _relay_client_request(method="GET"):
mock_request = MagicMock(spec=Request)
mock_request.method = method
mock_request.url = "http://localhost:4000/passthrough-relay/results"
mock_request.url = httpx.URL("http://localhost:4000/passthrough-relay/results")
mock_request.body = AsyncMock(return_value=b"")
mock_request.headers = Headers({})
mock_request.query_params = QueryParams({})
@ -6650,7 +6650,7 @@ def _passthrough_kwargs_for_reservation(
) -> dict:
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent"
mock_request.url = httpx.URL("http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent")
mock_request.headers = Headers({})
mock_request.scope = {"endpoint": _marked_pass_through_endpoint()} if user_defined_route else {}
@ -6797,7 +6797,7 @@ async def _drive_streaming_pass_through(upstream_content_type, chunk_delay_secon
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://test-proxy.com/v1/messages"
mock_request.url = httpx.URL("http://test-proxy.com/v1/messages")
mock_request.body = AsyncMock(
return_value=b'{"model": "claude-3", "stream": true}'
if client_asked_for_stream
@ -6985,36 +6985,82 @@ def _marked_pass_through_endpoint():
return _endpoint
def test_user_defined_passthrough_is_neither_tracked_nor_enforced():
"""
`get_model_from_request` returns None for a user-defined pass-through on
purpose: the body is forwarded verbatim, so its `model` names an UPSTREAM
model rather than a LiteLLM-managed one, and enforcing key/team allowlists
against it would reject valid requests. Enforcement is therefore skipped
on those routes.
@pytest.mark.asyncio
@pytest.mark.parametrize("metadata_slot", ["metadata", "litellm_metadata"])
async def test_user_defined_passthrough_is_neither_tracked_nor_enforced(metadata_slot: str) -> None:
from datetime import datetime
Attaching the budget metadata anyway would charge a counter that nothing on
that route can refuse, and would attribute the spend to a budget the operator
scoped to a LiteLLM model that merely shares the name. Tracking and
enforcement have to agree: both on for the built-in provider routes, both off
here.
"""
kwargs = _passthrough_kwargs_for_reservation(
UserAPIKeyAuth(
token="hash",
user_id="u-1",
model_max_budget={"claude-opus-4-8": {"budget_limit": 1.0, "time_period": "18h"}},
),
user_defined_route=True,
from litellm.caching.caching import DualCache
from litellm.proxy.auth.auth_utils import get_model_from_request
from litellm.proxy.hooks.model_max_budget_limiter import _PROXY_VirtualKeyModelMaxBudgetLimiter
budget: Final = {"managed-model": {"budget_limit": 0.1, "time_period": "1d"}}
limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(DualCache())
auth: Final = UserAPIKeyAuth(
api_key="custom-key", token="custom-key", team_id="shared-team", team_model_max_budget=budget,
)
endpoint: Final = create_pass_through_route(
endpoint="/custom-budget-test", target="https://upstream.test/echo", custom_headers={}, cost_per_request=0.25,
)
request: Final = Request({
"type": "http", "method": "POST", "path": "/custom-budget-test", "headers": [],
"query_string": b"", "endpoint": endpoint,
})
body: Final = {
"model": "upstream-only-model", metadata_slot: {
"model_group": "managed-model", "customer_label": "retained",
"user_api_key_team_model_max_budget": budget,
},
}
assert get_model_from_request(body, "/custom-budget-test", request=request) is None
assert await limiter.is_team_within_model_budget("shared-team", budget, None, "managed-model")
start: Final = datetime.now()
logging_obj: Final = LiteLLMLoggingObj(
model="upstream-only-model", messages=[], stream=False, call_type="pass_through_endpoint",
start_time=start, litellm_call_id="custom-budget", function_id="custom-budget", kwargs={},
dynamic_async_success_callbacks=[limiter],
)
payload: Final = {
"url": "https://upstream.test/echo", "request_body": body, "request_method": "POST", "cost_per_request": 0.25,
}
kwargs: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
request=request, user_api_key_dict=auth, passthrough_logging_payload=payload, logging_obj=logging_obj,
_parsed_body=body, litellm_call_id="custom-budget",
)
logging_obj.update_environment_variables(
model="upstream-only-model", user="unknown", optional_params={},
litellm_params=kwargs["litellm_params"], call_type="pass_through_endpoint",
)
response: Final = httpx.Response(
200, request=httpx.Request("POST", "https://upstream.test/echo"), json={"ok": True},
)
await PassThroughEndpointLogging().pass_through_async_success_handler(
httpx_response=response, response_body={"ok": True}, request_body=body, logging_obj=logging_obj,
url_route="https://upstream.test/echo", result=response.text, start_time=start, end_time=datetime.now(),
cache_hit=False, **kwargs,
)
assert logging_obj.model_call_details["response_cost"] == 0.25
assert await limiter.is_team_within_model_budget("shared-team", budget, None, "managed-model")
metadata: Final = kwargs["litellm_params"]["metadata"]
assert (metadata["model_group"], metadata["customer_label"]) == ("managed-model", "retained")
assert metadata.keys().isdisjoint({
"user_api_key_model_max_budget", "user_api_key_team_model_max_budget",
"user_api_key_user_model_max_budget", "user_api_key_end_user_model_max_budget",
})
metadata = kwargs["litellm_params"]["metadata"]
for field in (
"user_api_key_model_max_budget",
"user_api_key_user_model_max_budget",
"user_api_key_end_user_model_max_budget",
):
assert field not in metadata, f"{field} was attached on a route that never enforces it"
@pytest.mark.parametrize("metadata_slot", ["metadata", "litellm_metadata"])
def test_builtin_passthrough_pins_model_group_to_the_resolved_model(metadata_slot: str) -> None:
request: Final = Request({
"type": "http", "method": "POST", "path": "/gemini/v1beta/models/gemini-2.5-flash:generateContent",
"headers": [], "query_string": b"",
})
kwargs: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
request=request, user_api_key_dict=UserAPIKeyAuth(token="hash", user_id="u-1"),
passthrough_logging_payload=MagicMock(), logging_obj=MagicMock(),
_parsed_body={"contents": [], metadata_slot: {"model_group": "unbounded-client-choice"}},
)
assert kwargs["litellm_params"]["metadata"]["model_group"] == "gemini-2.5-flash"
@pytest.mark.parametrize(
@ -7344,7 +7390,7 @@ def test_passthrough_client_cannot_forge_session_id_omission(client_metadata_key
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent"
mock_request.url = httpx.URL("http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent")
mock_request.headers = Headers({})
mock_request.scope = {}
@ -7377,7 +7423,7 @@ def test_passthrough_logs_the_resolved_deployment_model_info_over_the_request_bo
the call to (LIT-1761: passthrough successes carried model_id="")."""
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://0.0.0.0:4000/vertex_ai/v1/projects/p/locations/global/publishers/google/models/gemini-3.8-flash:generateContent"
mock_request.url = httpx.URL("http://0.0.0.0:4000/vertex_ai/v1/projects/p/locations/global/publishers/google/models/gemini-3.8-flash:generateContent")
mock_request.headers = Headers({})
mock_request.scope = {}
mock_request.state = SimpleNamespace(
@ -7409,7 +7455,7 @@ _PROXY_SERVER_REQUEST: Final = TypeAdapter(dict[str, object])
def _split_pass_through_body(body: str) -> _PassThroughSplit:
mock_request: Final = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent"
mock_request.url = httpx.URL("http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent")
mock_request.headers = Headers()
mock_request.scope = MappingProxyType({})
@ -7546,6 +7592,61 @@ def test_passthrough_sees_the_public_list_rebound_after_import(monkeypatch: pyte
assert split.forwarded_body == {"contents": [{"parts": [{"text": "hi"}]}]}
def test_passthrough_metadata_carries_key_team_project_tags_and_key_spend_logs_metadata():
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = httpx.URL("http://0.0.0.0:4000/anthropic/v1/messages")
mock_request.headers = Headers({"x-litellm-tags": "caller-tag,key-tag"})
mock_request.scope = {}
cached_key = UserAPIKeyAuth(
api_key="hashed-key",
metadata={"tags": ["key-tag", "shared-tag"], "spend_logs_metadata": {"cost_center": "key"}},
team_metadata={
"tags": ["team-tag", "shared-tag"],
"spend_logs_metadata": {"cost_center": "team", "team_field": "team"},
},
project_metadata={"tags": ["project-tag"]},
)
kwargs = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
request=mock_request,
user_api_key_dict=cached_key,
passthrough_logging_payload=MagicMock(),
logging_obj=MagicMock(),
_parsed_body={
"metadata": {
"tags": ["body-tag"],
"spend_logs_metadata": {"request_id": "body"},
"user_api_key_auth_metadata": "forged",
}
},
litellm_call_id="lit-5359-call-id",
)
second = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
request=mock_request,
user_api_key_dict=cached_key,
passthrough_logging_payload=MagicMock(),
logging_obj=MagicMock(),
_parsed_body={},
litellm_call_id="lit-5359-second-call-id",
)
metadata = kwargs["litellm_params"]["metadata"]
assert metadata["tags"] == ["body-tag", "key-tag", "shared-tag", "team-tag", "project-tag", "caller-tag"]
assert metadata["spend_logs_metadata"] == {"request_id": "body", "cost_center": "key", "team_field": "team"}
assert metadata["user_api_key_auth_metadata"] == {
"tags": ["key-tag", "shared-tag"],
"spend_logs_metadata": {"cost_center": "key"},
}
assert second["litellm_params"]["metadata"]["spend_logs_metadata"] == {"cost_center": "key", "team_field": "team"}
assert cached_key.metadata == {"tags": ["key-tag", "shared-tag"], "spend_logs_metadata": {"cost_center": "key"}}
assert cached_key.team_metadata == {
"tags": ["team-tag", "shared-tag"],
"spend_logs_metadata": {"cost_center": "team", "team_field": "team"},
}
@pytest.mark.asyncio
async def test_chat_completion_pass_through_endpoint_answers_an_openai_typed_error_for_an_unknown_model(
monkeypatch: pytest.MonkeyPatch,
@ -7665,7 +7766,7 @@ def test_passthrough_attributes_a_cli_session_to_its_alias_not_the_login_token()
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://0.0.0.0:4000/anthropic/v1/messages"
mock_request.url = httpx.URL("http://0.0.0.0:4000/anthropic/v1/messages")
mock_request.headers = Headers({})
mock_request.scope = {}
session = UserAPIKeyAuth(

View file

@ -18,7 +18,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from fastapi import HTTPException
from fastapi import HTTPException, Request
pytest.importorskip("opentelemetry")
@ -81,15 +81,16 @@ def _user_api_key_dict():
return d
def _mock_request():
r = MagicMock()
r.method = "POST"
r.query_params = {}
r.url = "http://testserver/mock/echo"
headers = MagicMock()
headers.copy.return_value = {}
r.headers = headers
return r
def _mock_request() -> Request:
return Request({
"type": "http",
"method": "POST",
"scheme": "http",
"server": ("testserver", 80),
"path": "/mock/echo",
"headers": [],
"query_string": b"",
})
def _httpx_response(text: str) -> httpx.Response:

View file

@ -1,8 +1,8 @@
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import Request
from starlette.datastructures import Headers, State
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
@ -771,12 +771,15 @@ async def test_vertex_passthrough_attributes_the_call_to_the_resolved_deployment
"""The router deployment that rewrote the upstream URL is the one the logging kwargs must name, so
the Prometheus model_id label (and SpendLogs.model_id) on a Vertex passthrough success reads the
deployment's id instead of "" (LIT-1761)."""
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://0.0.0.0:4000/vertex_ai/v1/projects/p/locations/global/publishers/google/models/gemini-3.8-flash:generateContent"
mock_request.headers = Headers({})
mock_request.scope = {}
mock_request.state = State()
mock_request: Final = Request({
"type": "http",
"method": "POST",
"scheme": "http",
"server": ("0.0.0.0", 4000),
"path": "/vertex_ai/v1/projects/p/locations/global/publishers/google/models/gemini-3.8-flash:generateContent",
"headers": [],
"query_string": b"",
})
mock_handler = MagicMock()
mock_handler.get_default_base_target_url.return_value = "https://aiplatform.googleapis.com"

View file

@ -796,19 +796,23 @@ async def test_spend_logs_retention_alone_does_not_touch_the_session_rollup():
assert any('"LiteLLM_SpendLogs"' in sql for sql in tables)
assert not any('"LiteLLM_AutoRouterSession"' in sql for sql in tables)
assert not any('"LiteLLM_AutoRouterUserSession"' in sql for sql in tables)
assert not any('"LiteLLM_AutoRouterDailySpend"' in sql for sql in tables)
assert not any('"LiteLLM_HealthCheckTable"' in sql for sql in tables)
@pytest.mark.asyncio
async def test_session_retention_alone_cleans_both_session_rollups():
client = _mock_prisma_for_retention([0, 0])
async def test_session_retention_alone_cleans_both_session_rollups_and_the_daily_rollup():
client = _mock_prisma_for_retention([0, 0, 0])
cleaner = SpendLogCleanup(general_settings={"maximum_autorouter_session_retention_period": "365d"})
cleaner.pod_lock_manager = None
await cleaner.cleanup_old_spend_logs(client)
tables = [call[0][0] for call in client.db.execute_raw.call_args_list]
assert len(tables) == 2
calls = client.db.execute_raw.call_args_list
tables = [call[0][0] for call in calls]
assert len(tables) == 3
assert '"LiteLLM_AutoRouterSession"' in tables[0]
assert '"LiteLLM_AutoRouterUserSession"' in tables[1]
assert '"LiteLLM_AutoRouterDailySpend"' in tables[2]
assert calls[2][0][1] == calls[0][0][1].date().isoformat()
@pytest.mark.asyncio
@ -852,7 +856,7 @@ async def test_spend_logs_retention_alone_keeps_daily_tag_spend_forever():
@pytest.mark.asyncio
async def test_each_retention_key_cuts_off_at_its_own_horizon():
client = _mock_prisma_for_retention([0, 0, 0, 0, 0])
client = _mock_prisma_for_retention([0, 0, 0, 0, 0, 0])
cleaner = SpendLogCleanup(
general_settings={
"maximum_spend_logs_retention_period": "7d",
@ -868,6 +872,8 @@ async def test_each_retention_key_cuts_off_at_its_own_horizon():
if '"LiteLLM_AutoRouterSession"' in call[0][0]
else "LiteLLM_AutoRouterUserSession"
if '"LiteLLM_AutoRouterUserSession"' in call[0][0]
else "LiteLLM_AutoRouterDailySpend"
if '"LiteLLM_AutoRouterDailySpend"' in call[0][0]
else "LiteLLM_HealthCheckTable"
if '"LiteLLM_HealthCheckTable"' in call[0][0]
else "logs"
@ -878,6 +884,7 @@ async def test_each_retention_key_cuts_off_at_its_own_horizon():
assert (now - cutoffs["logs"]).days == 7
assert (now - cutoffs["LiteLLM_AutoRouterSession"]).days == 365
assert cutoffs["LiteLLM_AutoRouterUserSession"] == cutoffs["LiteLLM_AutoRouterSession"]
assert cutoffs["LiteLLM_AutoRouterDailySpend"] == cutoffs["LiteLLM_AutoRouterSession"].date().isoformat()
assert (now - cutoffs["LiteLLM_HealthCheckTable"]).days == 30

View file

@ -8,6 +8,7 @@ from unittest.mock import create_autospec
import httpx
import pytest
import respx
import litellm
from litellm._logging import verbose_router_logger
@ -30,14 +31,15 @@ from litellm.types.utils import AUTOROUTER_CLASSIFIER_CALL_ORIGIN
class _UsageRecorder(CustomLogger):
def __init__(self) -> None:
def __init__(self, model_key: str = "typesafe/jev-accounting") -> None:
super().__init__()
self.model_key = model_key
self.calls: tuple[Mapping[str, object], ...] = ()
async def async_log_success_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
) -> None:
if str(kwargs.get("model", "")).removeprefix("typesafe/") != "jev-accounting":
if str(kwargs.get("model", "")) != self.model_key:
return
self.calls = (*self.calls, kwargs)
@ -167,8 +169,9 @@ async def test_jev_invalid_usage_never_reaches_spend_callbacks(
@pytest.mark.asyncio
@pytest.mark.parametrize("answer", ["SIMPLE", "UNAVAILABLE", "malformed"])
@pytest.mark.parametrize("private", [False, True])
@pytest.mark.parametrize("legacy", [False, True])
async def test_jev_accounts_once_with_parent_identity_even_when_the_verdict_fails(
monkeypatch: pytest.MonkeyPatch, answer: str, private: bool
monkeypatch: pytest.MonkeyPatch, answer: str, private: bool, legacy: bool
) -> None:
recorder: Final = _UsageRecorder()
monkeypatch.setattr(litellm, "_async_success_callback", [recorder])
@ -196,7 +199,15 @@ async def test_jev_accounts_once_with_parent_identity_even_when_the_verdict_fail
router: Final = ComplexityRouter(
"jev-router",
litellm.Router(model_list=[]),
{"classifier_type": "jev", "jev_classifier_config": {}, "tiers": {"SIMPLE": "cheap"}},
{
"classifier_type": "jev" if legacy else "oss_classifier",
"jev_classifier_config" if legacy else "opensource_classifier_config": {
"provider": "typesafe" if legacy else "jev",
},
"tiers": {"SIMPLE": "cheap"},
"session_affinity": False,
"deployment_affinity": False,
},
jev_client=provider,
derive_savings_baseline=False,
)
@ -209,8 +220,9 @@ async def test_jev_accounts_once_with_parent_identity_even_when_the_verdict_fail
"user_api_key_budget_reservation": {"reservation_id": "parent-reservation"},
"user_api_key_auth": {"budget_reservation": {"reservation_id": "parent-reservation"}},
}
outcome: Final = await router.aclassify(
"private current ask",
result: Final = await router.async_pre_routing_hook(
model="jev-router",
messages=[{"role": "user", "content": "private current ask"}],
request_kwargs={
"metadata": metadata,
"litellm_session_id": "session-a",
@ -221,7 +233,15 @@ async def test_jev_accounts_once_with_parent_identity_even_when_the_verdict_fail
await GLOBAL_LOGGING_WORKER.flush()
await handler.client.aclose()
assert (outcome.cause == "jev_classifier") is (answer == "SIMPLE")
assert result is not None and result.model == "cheap"
assert result.routing_decision is not None
decision: Final = result.routing_decision
assert (decision["cause"] == "jev_classifier") is (answer == "SIMPLE")
if answer == "SIMPLE":
assert decision["classifier_model"] == "typesafe/jev-accounting"
assert decision["classifier_cost"] == pytest.approx(0.007)
assert "jev-classifier:SIMPLE" in decision["signals"]
assert "jev-confidence=1.000000" in decision["signals"]
assert len(recorder.calls) == 1
event: Final = recorder.calls[0]
assert event["response_cost"] == pytest.approx(0.007)
@ -416,10 +436,101 @@ def _answer(choice: str = "SIMPLE") -> JevChoiceAnswer:
def test_jev_config_requires_classifier_config() -> None:
with pytest.raises(ValueError, match="jev_classifier_config is required"):
with pytest.raises(ValueError, match="opensource_classifier_config is required"):
ComplexityRouterConfig.model_validate({"classifier_type": "jev"})
@pytest.mark.parametrize(
("classifier_type", "config_key"),
[
("oss_classifier", "opensource_classifier_config"),
("jev", "jev_classifier_config"),
("oss_classifier", "jev_classifier_config"),
("jev", "opensource_classifier_config"),
],
)
@pytest.mark.parametrize(
("provider", "model", "canonical_provider"),
[(None, "jev-latest", "jev"), ("typesafe", "jev-latest", "jev"), ("jev", "jev-latest", "jev"), ("laya", "english", "laya")],
)
def test_classifier_aliases_load_and_serialize_one_canonical_config(
classifier_type: str, config_key: str, provider: str | None, model: str, canonical_provider: str
) -> None:
incoming: Final = {
"classifier_type": classifier_type,
config_key: {"model": model, "api_key": None, **({"provider": provider} if provider is not None else {})},
}
original: Final = deepcopy(incoming)
config: Final = ComplexityRouterConfig.model_validate(incoming)
assert config.classifier_type == "oss_classifier"
assert config.opensource_classifier_config is not None
assert config.opensource_classifier_config.provider == canonical_provider
assert config.opensource_classifier_config.model == model
assert config.opensource_classifier_config.api_key is None
assert "api_key" in config.opensource_classifier_config.model_fields_set
assert "api_base" not in config.opensource_classifier_config.model_fields_set
assert "jev_classifier_config" not in config.model_dump()
assert config.jev_classifier_config is config.opensource_classifier_config
assert incoming == original
@pytest.mark.parametrize("config", [{"provider": "laya"}, {"provider": "laya", "model": " "}])
def test_laya_requires_its_own_checkpoint(config: Mapping[str, object]) -> None:
with pytest.raises(ValueError, match="Laya model must be"):
JevClassifierConfig.model_validate(config)
@pytest.mark.asyncio
@pytest.mark.parametrize("custom_base", [False, True])
@pytest.mark.parametrize("legacy", [False, True])
async def test_laya_routes_with_its_own_credentials_and_accounts_the_checkpoint(
monkeypatch: pytest.MonkeyPatch, custom_base: bool, legacy: bool
) -> None:
monkeypatch.setenv("TYPESAFE_API_KEY", "never-send-typesafe-key")
monkeypatch.setenv("LAYA_API_BASE", "https://laya.test")
monkeypatch.setenv("LAYA_API_KEY", "laya-env-key")
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
monkeypatch.setitem(litellm.model_cost, "laya/english", {"input_cost_per_token": 0.01})
recorder: Final = _UsageRecorder("laya/english")
monkeypatch.setattr(litellm, "_async_success_callback", [recorder])
router: Final = ComplexityRouter(
"laya-route",
litellm.Router(model_list=[]),
{
"classifier_type": "jev" if legacy else "oss_classifier",
"jev_classifier_config" if legacy else "opensource_classifier_config": {
"provider": "laya",
"model": "english",
**({"api_base": "https://laya.test"} if custom_base else {}),
},
"tiers": {"SIMPLE": "cheap"},
},
derive_savings_baseline=False,
)
with respx.mock(assert_all_called=True) as upstream:
route: Final = upstream.post("https://laya.test/v1/systemone").respond(
200,
json={
"model": "laya-rl-agent",
"routing": {"model": "english"},
"answers": {"tier": _answer().model_dump()},
"usage": {"input_tokens": 31, "output_tokens": 0},
},
)
outcome: Final = await router.aclassify("choose a tier")
await GLOBAL_LOGGING_WORKER.flush()
assert outcome.cause == "jev_classifier"
assert outcome.jev_verdict is not None
assert (outcome.jev_verdict.provider, outcome.jev_verdict.model) == ("laya", "english")
assert outcome.classifier_cost == pytest.approx(0.31)
sent: Final = route.calls.last.request
assert sent.headers.get("authorization") == (None if custom_base else "Bearer laya-env-key")
assert json.loads(sent.content)["model"] == "english"
assert len(recorder.calls) == 1
assert recorder.calls[0]["response_cost"] == pytest.approx(0.31)
def test_jev_config_is_rejected_for_other_classifier_types() -> None:
with pytest.raises(ValueError, match="has no effect"):
ComplexityRouterConfig.model_validate(
@ -437,7 +548,7 @@ def test_jev_instructions_reject_blank_values() -> None:
@pytest.mark.parametrize(
("missing_key", "rejection"),
[
({}, r"api_base requires jev_classifier_config\.api_key"),
({}, r"api_base requires opensource_classifier_config\.api_key"),
({"api_key": ""}, r"api_key must be non-empty"),
({"api_key": " "}, r"api_key must be non-empty"),
],

View file

@ -6,11 +6,11 @@ import pytest
from litellm.router_strategy.complexity_router.fuse_presets import get_fuse_presets
from litellm.router_strategy.complexity_router.jev_classifier import DEFAULT_JEV_INSTRUCTIONS
from litellm.router_utils.auto_router_model_naming import (
carries_complexity_router_settings,
classify_strategy_router_model,
GATED_AUTO_ROUTER_CAPABILITIES,
capability_limit_violation,
carries_complexity_router_settings,
claimed_capability,
classify_strategy_router_model,
count_capability_routers,
gated_capability_of,
strategy_router_dependencies,
@ -23,27 +23,59 @@ COMPLEXITY_FIELDS = frozenset({"complexity_router_config"})
SEMANTIC_FIELDS = frozenset({"auto_router_config", "auto_router_default_model", "auto_router_embedding_model"})
@pytest.mark.parametrize("model", ["jev-latest", "jev-preview"])
def test_jev_enumerates_a_paid_evaluation_without_a_completion_classifier(model: str) -> None:
found = strategy_router_dependencies(
@pytest.mark.parametrize(
("classifier_type", "config_key"),
[
("jev", "jev_classifier_config"),
("oss_classifier", "opensource_classifier_config"),
("jev", "opensource_classifier_config"),
("oss_classifier", "jev_classifier_config"),
],
)
@pytest.mark.parametrize(
("provider", "model", "accounting_provider"),
[
(None, "jev-latest", "typesafe"),
("typesafe", "jev-preview", "typesafe"),
("jev", "jev-preview", "typesafe"),
("laya", "english", "laya"),
],
)
def test_open_source_classifier_enumerates_its_accounting_model(
classifier_type: str, config_key: str, provider: str | None, model: str, accounting_provider: str
) -> None:
found: Final = strategy_router_dependencies(
{
"model": "auto_router/complexity_router",
"complexity_router_config": {
"classifier_type": "jev",
"jev_classifier_config": {"model": model},
"classifier_type": classifier_type,
config_key: {"model": model, **({"provider": provider} if provider else {})},
"tiers": {"SIMPLE": "cheap"},
},
}
)
assert tuple((dep.model_name, dep.role) for dep in found) == (
("cheap", "tier"),
(f"typesafe/{model}", "evaluation"),
(f"{accounting_provider}/{model}", "evaluation"),
)
@pytest.mark.parametrize("instructions", [None, DEFAULT_JEV_INSTRUCTIONS, "Route conservatively"])
def test_only_non_default_jev_instructions_claim_the_shared_customization_slot(instructions: str | None) -> None:
capability = claimed_capability({"classifier_type": "jev", "jev_classifier_config": {"instructions": instructions}})
@pytest.mark.parametrize(
("classifier_type", "config_key"),
[
("jev", "jev_classifier_config"),
("oss_classifier", "opensource_classifier_config"),
("jev", "opensource_classifier_config"),
("oss_classifier", "jev_classifier_config"),
],
)
def test_only_non_default_open_source_instructions_claim_the_shared_customization_slot(
instructions: str | None, classifier_type: str, config_key: str
) -> None:
capability: Final = claimed_capability(
{"classifier_type": classifier_type, config_key: {"instructions": instructions}}
)
assert (capability.key if capability else None) == (
"tier_or_classifier_prompt" if instructions == "Route conservatively" else None
)
@ -123,6 +155,21 @@ VALID_TIERS = {
}
@pytest.mark.parametrize("legacy_config", [None, {}, {"provider": "laya", "model": "english"}])
def test_dual_classifier_blocks_return_a_write_validation_error(legacy_config: Mapping[str, object] | None) -> None:
violation: Final = validate_complexity_router_config_write(
{
"tiers": VALID_TIERS,
"classifier_type": "oss_classifier",
"opensource_classifier_config": {"provider": "laya", "model": "english"},
"jev_classifier_config": legacy_config,
}
)
assert violation is not None
assert "opensource_classifier_config" in violation
assert "jev_classifier_config" in violation
@pytest.mark.parametrize(
"keyword_tier_rules,expected_fragment",
[
@ -408,6 +455,8 @@ def test_complexity_embedding_model_is_a_dependency_only_when_semantic_matching_
("token_thresholds", "dimension_weights"),
("reasoning_override_min_score",),
("tiers",),
("jev_classifier_config",),
("opensource_classifier_config",),
],
)
def test_placement_rejects_settings_written_beside_the_config(misplaced):
@ -447,7 +496,7 @@ def test_placement_guards_every_setting_the_config_owns():
ComplexityRouterConfig,
)
assert COMPLEXITY_ROUTER_CONFIG_KEYS == frozenset(ComplexityRouterConfig.model_fields)
assert COMPLEXITY_ROUTER_CONFIG_KEYS == frozenset(ComplexityRouterConfig.model_fields) | {"jev_classifier_config"}
assert {"tier_boundaries", "token_thresholds", "dimension_weights"} <= COMPLEXITY_ROUTER_CONFIG_KEYS

View file

@ -1020,6 +1020,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"/v1/videos",
"/vertex_ai/live",
"/v1/listen",
"/v1/systemone",
"/v1beta/interactions",
],
},

View file

@ -75,7 +75,6 @@ const totals = (overrides: Partial<Totals> = {}): Totals => ({
saved_spend: 2174.59,
baseline_spend: 2534.45,
saved_pct: 85.8,
saved_per_session: 23.13,
cache: cache(),
...overrides,
});
@ -110,7 +109,6 @@ const zeroTotals: Totals = {
saved_spend: 0,
baseline_spend: 0,
saved_pct: 0,
saved_per_session: 0,
cache: zeroCache,
};
@ -173,7 +171,6 @@ describe("AutoRouterBenchmarksTab", () => {
saved_spend: saved,
baseline_spend: estimatedTurns ? actual + (saved ?? 0) : null,
saved_pct: pct,
saved_per_session: null,
};
mockHook({
data: response([], totals(comparison)),
@ -204,18 +201,15 @@ describe("AutoRouterBenchmarksTab", () => {
}
});
it("leads with total estimated savings, before the four session-shape metrics", () => {
it("leads with total estimated savings, before the three session-shape metrics", () => {
mockHook({ data: response([group(), group({ router_name: "gpt-auto" })]) });
renderTab();
const labels = screen
.getAllByText(
/Total estimated savings|Avg saved per session|Avg turns per session|Avg session length|Avg tokens per session/,
)
.getAllByText(/Total estimated savings|Avg turns per session|Avg session length|Avg tokens per session/)
.map((node) => node.textContent);
expect(labels).toEqual([
"Total estimated savings",
"Avg saved per session",
"Avg turns per session",
"Avg session length",
"Avg tokens per session",
@ -271,15 +265,40 @@ describe("AutoRouterBenchmarksTab", () => {
},
);
it("pairs the savings with the session count it was earned over, in its own tile", () => {
it("labels selected-day money apart from whole-session metrics, with no savings-per-session tile", () => {
mockHook({ data: response([group(), group({ router_name: "gpt-auto" })]) });
renderTab();
const tile = screen.getByText("Avg saved per session").closest('[data-slot="card"]');
if (!tile) throw new Error("expected avg saved per session to render as a metric tile");
expect(within(tile).getByText("$23.13")).toBeInTheDocument();
const tile = screen.getByText("Avg turns per session").closest<HTMLElement>('[data-slot="card"]');
if (!tile) throw new Error("expected avg turns per session to render as a metric tile");
expect(within(tile).getByText("· 94 sessions")).toBeInTheDocument();
expect(screen.queryByText("Avg saved per session")).not.toBeInTheDocument();
expect(screen.getByText(/Savings and spend count requests on the selected UTC days/)).toBeInTheDocument();
expect(screen.getByText(/Session metrics cover every session that overlaps the range/)).toBeInTheDocument();
});
it("shows session averages as unavailable, not zero, when routed requests have no session rows", () => {
const noSessions = {
sessions: 0,
avg_turns_per_session: null,
avg_session_seconds: null,
avg_tokens_per_session: null,
};
mockHook({ data: response([], totals(noSessions)) });
renderTab();
expect(screen.getAllByText("Unavailable")).toHaveLength(3);
expect(screen.queryByText("0.0")).not.toBeInTheDocument();
});
it.each([3, -3])("explains a %s gap between router records and recorded savings instead of comparing", (gap) => {
const residual = { saved_spend: 5, unattributed_saved_spend: gap, baseline_spend: null, saved_pct: null };
mockHook({ data: response([], totals(residual)) });
renderTab();
expect(screen.getByText("$5.00")).toBeInTheDocument();
expect(screen.getByText(/Per-router records differ from recorded savings by \$3\.00/)).toBeInTheDocument();
expect(screen.getByText("Estimated baseline spend").nextSibling?.textContent).toBe("Unavailable");
});
it("exposes each spend row as a term and its value, not as loose text", () => {
@ -431,7 +450,7 @@ describe("AutoRouterBenchmarksTab", () => {
renderTab();
expect(screen.getByText("Total estimated savings")).toBeInTheDocument();
expect(screen.getAllByText("$0.00")).toHaveLength(6);
expect(screen.getAllByText("$0.00")).toHaveLength(5);
expect(screen.getByText("· 0 sessions")).toBeInTheDocument();
expect(screen.getByText("0s")).toBeInTheDocument();
expect(screen.getByText(/turns measured/)).toBeInTheDocument();

View file

@ -105,6 +105,12 @@ const HeroCard: React.FC<{ view: BenchmarkView }> = ({ view }) => {
adaptive and quality routers are excluded
</p>
)}
{stats.unattributed_saved_spend != null && (
<p className="text-center text-xs text-muted-foreground">
Per-router records differ from recorded savings by {usd(Math.abs(stats.unattributed_saved_spend))}, for
example history from before per-router tracking, so the baseline comparison is unavailable
</p>
)}
</div>
<div className="flex flex-col justify-center border-t p-6 md:border-t-0 md:border-l">
@ -297,22 +303,34 @@ const BenchmarksBody: React.FC<BenchmarksBodyProps> = ({ isPending, error, data,
<TierTurnsChart view={view} autoRouters={autoRouters} />
<div className="grid grid-cols-1 gap-4 sm:grid-cols-2 lg:grid-cols-4">
<Metric
label="Avg saved per session"
value={stats.saved_per_session == null ? "Unavailable" : usd(stats.saved_per_session)}
hint={`· ${stats.sessions.toLocaleString()} sessions`}
/>
<Metric label="Avg turns per session" value={stats.avg_turns_per_session.toFixed(1)} />
<Metric label="Avg session length" value={durationLabel(stats.avg_session_seconds)} />
<Metric label="Avg tokens per session" value={formatNumberWithCommas(stats.avg_tokens_per_session, 1, true)} />
</div>
<p className="text-xs text-muted-foreground">
Savings and spend count requests on the selected UTC days. Actual spend covers every request on complexity
routers, including LLM classification cost. Baseline is actual spend plus recorded savings, so savings can be
zero or negative.
</p>
<p className="text-xs text-muted-foreground">
Actual spend covers every request on complexity routers, including LLM classification cost. Baseline is actual
spend plus recorded savings, so savings can be zero or negative. The range counts whole sessions that overlap
it, so totals can differ from savings views that group usage by UTC day.
Session metrics cover every session that overlaps the range, including its turns outside the range.
</p>
<div className="grid grid-cols-1 gap-4 sm:grid-cols-3">
<Metric
label="Avg turns per session"
value={stats.avg_turns_per_session == null ? "Unavailable" : stats.avg_turns_per_session.toFixed(1)}
hint={`· ${stats.sessions.toLocaleString()} sessions`}
/>
<Metric
label="Avg session length"
value={stats.avg_session_seconds == null ? "Unavailable" : durationLabel(stats.avg_session_seconds)}
/>
<Metric
label="Avg tokens per session"
value={
stats.avg_tokens_per_session == null
? "Unavailable"
: formatNumberWithCommas(stats.avg_tokens_per_session, 1, true)
}
/>
</div>
<div className="space-y-4">
<div className="flex flex-wrap items-baseline gap-2">

View file

@ -12,21 +12,21 @@ vi.mock("@/components/shared/charts", () => ({
}));
import TierTurnsChart, { tierDisplayLabel } from "./TierTurnsChart";
import type { AutoRouterBenchmarkGroup, BenchmarkView } from "./autoRouterBenchmarks";
import type { AutoRouterBenchmarkGroup, AutoRouterBenchmarkTotals, BenchmarkView } from "./autoRouterBenchmarks";
const totalsOnly = {
const totalsOnly: AutoRouterBenchmarkTotals = {
sessions: 3,
turns: 9,
avg_turns_per_session: 3,
avg_session_seconds: 60,
avg_tokens_per_session: 100,
spend: 1,
classifier_cost: 0,
savings_estimated_turns: 9,
savings_estimated_actual_spend: 1,
saved_spend: 1,
baseline_spend: 2,
saved_pct: 50,
saved_per_session: 0.33,
cache: {
coverage_pct: 0,
hit_rate_pct: 0,

View file

@ -43,7 +43,6 @@ const totals = (overrides: Partial<AutoRouterBenchmarkGroup> = {}) => ({
saved_spend: 2174.59,
baseline_spend: 2534.45,
saved_pct: 85.8,
saved_per_session: 23.13,
cache: cache(),
...overrides,
});

View file

@ -85,7 +85,8 @@ describe("autoRouterRows", () => {
it.each([
["llm", "LLM Classifier"],
["jev", "JEV Classifier"],
["jev", "OSS Classifier"],
["oss_classifier", "OSS Classifier"],
])("labels a router using the %s classifier", (classifierType, label) => {
const row = toAutoRouterRow(
{

View file

@ -57,7 +57,8 @@ const dedupe = (models: string[]): string[] => Array.from(new Set(models));
const COMPLEXITY_TYPE_LABELS: Record<string, string> = {
llm: "LLM Classifier",
jev: "JEV Classifier",
jev: "OSS Classifier",
oss_classifier: "OSS Classifier",
capability: "Capability",
llm_v2: "Fuse v2",
heuristic_first: "Heuristic first",

View file

@ -347,7 +347,6 @@ const routerUsageResponse = (saved: number): AutoRouterBenchmarksResponse => ({
saved_spend: saved,
baseline_spend: 10 + saved,
saved_pct: (100 * saved) / (10 + saved),
saved_per_session: saved / 2,
cache: {
coverage_pct: 100,
hit_rate_pct: 0,

View file

@ -136,7 +136,7 @@ export const AutoRouterLimits = () => {
<PopoverContent align="end" className="w-96 max-w-[calc(100vw-2rem)] gap-3">
<PopoverTitle>Routing and customization limits</PopoverTitle>
<p className="text-xs leading-5 text-muted-foreground">
Rule-based, Complexity, and Jev are unlimited with built-in settings. Choose or change tier models freely.
Rule-based, Complexity, and OSS are unlimited with built-in settings. Choose or change tier models freely.
Customization allowances are shared across this proxy.
</p>
<dl className="space-y-2 text-xs">

View file

@ -68,7 +68,7 @@ describe("Auto-router classifier selection", () => {
llm: "LLM",
heuristic_first: "LLM",
hybrid: "LLM",
jev: "Jev",
jev: "OSS Classifier",
}[classifier_type];
expect(screen.getByRole("radio", { name: new RegExp(`^${family}$`) })).toBeChecked();
fireEvent.click(screen.getByRole("radio", { name: new RegExp(`^${family}$`) }));

View file

@ -16,6 +16,7 @@ import {
type ClassifierType,
type ComplexityRouterConfigValue,
} from "./ComplexityRouterConfig";
import { defaultJevClassifierConfig, normalizeJevClassifierConfig } from "./jev_classifier_config";
import { transitionClassifierType } from "./classifier_type_transition";
import { isForecastClassifier } from "./forecast_classifier_config";
import {
@ -148,6 +149,14 @@ const AutoRouterClassifierTabs: React.FC<AutoRouterClassifierTabsProps> = ({ val
if (next === "llm") changeType("llm");
if (next === "jev") changeType("jev");
};
const changeProvider = (provider: unknown) => {
if (provider !== "jev" && provider !== "laya") return;
const defaults = defaultJevClassifierConfig(provider);
onChange({
...value,
jev_classifier_config: { ...defaults, ...value.jev_classifier_config, provider, model: defaults.model },
});
};
const approachLabels: Partial<Record<ClassifierType, string>> = { capability: "Capability", llm_v2: "Fuse v2" };
const approachDescription: Partial<Record<ClassifierType, string>> = {
capability: "Use the efficient model when it is likely to succeed",
@ -164,7 +173,7 @@ const AutoRouterClassifierTabs: React.FC<AutoRouterClassifierTabsProps> = ({ val
{[
{ value: "heuristics", label: "Heuristics", description: "Classify locally, with no API call" },
{ value: "llm", label: "LLM", description: "Use a judge model to choose a solver" },
{ value: "jev", label: "Jev", description: "Use TypeSafe System One Choice to choose a tier" },
{ value: "jev", label: "OSS Classifier", description: "Use Jev or Laya to choose a tier" },
].map((option) => (
<Label
key={option.value}
@ -189,6 +198,25 @@ const AutoRouterClassifierTabs: React.FC<AutoRouterClassifierTabsProps> = ({ val
))}
</RadioGroup>
</fieldset>
{family === "jev" && (
<fieldset className="space-y-2">
<legend className="text-sm font-medium">OSS provider</legend>
<RadioGroup
value={normalizeJevClassifierConfig(value.jev_classifier_config).provider}
onValueChange={changeProvider}
className="flex gap-6"
>
<Label>
<RadioGroupItem value="jev" />
Jev
</Label>
<Label>
<RadioGroupItem value="laya" />
Laya
</Label>
</RadioGroup>
</fieldset>
)}
{family === "custom" && (
<p className="text-sm text-muted-foreground">This router uses a custom classifier plugin</p>
)}

View file

@ -642,8 +642,8 @@ const ClassificationMethodConfig: React.FC<ClassificationMethodConfigProps> = ({
/>
<span className="text-xs text-muted-foreground">
Number of prior user turns sent to the classifier provider, excluding tool output and harness reminders.
LLM and Jev default to 3 turns; Jev sends them to the configured TypeSafe endpoint. Set to 0 to omit
conversation history. The current message and selected system text are still sent.
LLM and OSS classifiers default to 3 turns. Set to 0 to omit conversation history. The current message and
selected system text are still sent.
</span>
</div>
<div>

View file

@ -54,8 +54,8 @@ const ClassifierTypeRadios: React.FC<ClassifierTypeRadiosProps> = ({ value, clas
<Label className="items-start font-normal leading-normal">
<RadioGroupItem value="jev" className="mt-0.5" />
<span>
<strong className="font-semibold">Jev Classifier</strong>{" "}
<span className="text-muted-foreground">uses TypeSafe System One Choice to decide the tier</span>
<strong className="font-semibold">OSS Classifier</strong>{" "}
<span className="text-muted-foreground">uses Jev or Laya to decide the tier</span>
</span>
</Label>
<SimpleTooltip content={scorerLockedReason}>

View file

@ -237,7 +237,7 @@ const TierSetToolbar: React.FC<{
{editing && (
<span className="block mt-1 text-xs text-muted-foreground">
Add or remove tiers to define your own set. Every custom tier needs a definition the classifier routes on, and
an edited set requires the LLM or Jev classification method
an edited set requires the LLM or OSS classification method
</span>
)}
{editing && keywordRulesError && (

View file

@ -1,4 +1,5 @@
import React, { useState } from "react";
import userEvent from "@testing-library/user-event";
import { afterEach, describe, expect, it, vi } from "vitest";
import { fireEvent, renderWithProviders, screen } from "../../../tests/test-utils";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
@ -96,49 +97,65 @@ function Form() {
describe("JEV classifier editor", () => {
afterEach(() => vi.mocked(useAuthorized).mockReset());
it("uses built-in JEV without a license and preserves custom tiers and context through reload", () => {
renderWithProviders(<Form />);
expect(screen.getByLabelText("Judge model")).toBeInTheDocument();
expect(screen.getByText("Reasoning Effort")).toBeInTheDocument();
expect(screen.getByText("Classifier Prompt")).toBeInTheDocument();
expect(screen.getByRole("switch", { name: "Use images for classification" })).toBeInTheDocument();
fireEvent.click(screen.getByRole("radio", { name: /Jev Classifier/ }));
expect(screen.getByRole("radio", { name: /^Jev Classifier/ })).toBeChecked();
expect(screen.getByLabelText("Jev Model")).toHaveValue("jev-latest");
expect(screen.getByLabelText("Jev Instructions")).toBeEnabled();
expect(screen.queryByLabelText("Judge model")).not.toBeInTheDocument();
expect(screen.queryByText("Reasoning Effort")).not.toBeInTheDocument();
expect(screen.queryByText("Classifier Prompt")).not.toBeInTheDocument();
expect(screen.queryByRole("switch", { name: "Use images for classification" })).not.toBeInTheDocument();
fireEvent.change(screen.getByLabelText("Jev Model"), { target: { value: "jev-test" } });
fireEvent.change(screen.getByLabelText("Jev Timeout (ms)"), { target: { value: "4200" } });
fireEvent.change(screen.getByLabelText("Context Window Size"), { target: { value: "6" } });
fireEvent.change(screen.getByLabelText("Circuit breaker cooldown (seconds)"), { target: { value: "50" } });
fireEvent.click(screen.getByRole("switch", { name: "Classifier circuit breaker" }));
fireEvent.click(screen.getByRole("button", { name: "Customize tiers" }));
fireEvent.click(screen.getByRole("button", { name: "Save and reload" }));
expect(screen.getByRole("radio", { name: /Jev Classifier/ })).toBeChecked();
expect(screen.getByLabelText("Jev Model")).toHaveValue("jev-test");
expect(screen.getByLabelText("Jev Timeout (ms)")).toHaveValue(4200);
expect(screen.getByLabelText("Context Window Size")).toHaveValue("6");
expect(screen.getByRole("switch", { name: "Classifier circuit breaker" })).not.toBeChecked();
fireEvent.click(screen.getByRole("button", { name: "Probe current config" }));
expect(testAutoRouterRouting).toHaveBeenCalledWith(
"token",
expect.objectContaining({
complexity_router_config: expect.objectContaining({
classifier_type: "jev",
jev_classifier_config: {
model: "jev-test",
timeout_ms: 4200,
circuit_breaker_enabled: false,
circuit_breaker_cooldown_seconds: 50,
},
tiers: expect.objectContaining({ QUICK: ["fast"] }),
it.each(["jev", "laya"] as const)(
"preserves %s, custom tiers and context through save, reload and probe",
async (provider) => {
renderWithProviders(<Form />);
expect(screen.getByLabelText("Judge model")).toBeInTheDocument();
expect(screen.getByText("Reasoning Effort")).toBeInTheDocument();
expect(screen.getByText("Classifier Prompt")).toBeInTheDocument();
expect(screen.getByRole("switch", { name: "Use images for classification" })).toBeInTheDocument();
fireEvent.click(screen.getByRole("radio", { name: /^OSS Classifier$/ }));
expect(screen.getByRole("radio", { name: /^OSS Classifier$/ })).toBeChecked();
expect(screen.getByLabelText("Classifier Model")).toHaveValue("jev-latest");
expect(screen.getByLabelText("Classifier Instructions")).toBeEnabled();
expect(screen.queryByLabelText("Judge model")).not.toBeInTheDocument();
expect(screen.queryByText("Reasoning Effort")).not.toBeInTheDocument();
expect(screen.queryByText("Classifier Prompt")).not.toBeInTheDocument();
expect(screen.queryByRole("switch", { name: "Use images for classification" })).not.toBeInTheDocument();
fireEvent.click(screen.getByRole("radio", { name: "Laya" }));
expect(screen.getByLabelText("Classifier Model")).toHaveTextContent("english");
fireEvent.click(screen.getByRole("radio", { name: "Jev" }));
expect(screen.getByLabelText("Classifier Model")).toHaveValue("jev-latest");
if (provider === "laya") {
fireEvent.click(screen.getByRole("radio", { name: "Laya" }));
await userEvent.click(screen.getByLabelText("Classifier Model"));
await userEvent.click(screen.getByRole("option", { name: "multilingual" }));
} else {
fireEvent.change(screen.getByLabelText("Classifier Model"), { target: { value: "jev-test" } });
}
fireEvent.change(screen.getByLabelText("Classifier Timeout (ms)"), { target: { value: "4200" } });
fireEvent.change(screen.getByLabelText("Context Window Size"), { target: { value: "6" } });
fireEvent.change(screen.getByLabelText("Circuit breaker cooldown (seconds)"), { target: { value: "50" } });
fireEvent.click(screen.getByRole("switch", { name: "Classifier circuit breaker" }));
fireEvent.click(screen.getByRole("button", { name: "Customize tiers" }));
fireEvent.click(screen.getByRole("button", { name: "Save and reload" }));
expect(screen.getByRole("radio", { name: /^OSS Classifier$/ })).toBeChecked();
expect(screen.getByRole("radio", { name: provider === "laya" ? "Laya" : "Jev" })).toBeChecked();
if (provider === "laya") expect(screen.getByLabelText("Classifier Model")).toHaveTextContent("multilingual");
else expect(screen.getByLabelText("Classifier Model")).toHaveValue("jev-test");
expect(screen.getByLabelText("Classifier Timeout (ms)")).toHaveValue(4200);
expect(screen.getByLabelText("Context Window Size")).toHaveValue("6");
expect(screen.getByRole("switch", { name: "Classifier circuit breaker" })).not.toBeChecked();
fireEvent.click(screen.getByRole("button", { name: "Probe current config" }));
expect(testAutoRouterRouting).toHaveBeenCalledWith(
"token",
expect.objectContaining({
complexity_router_config: expect.objectContaining({
classifier_type: "oss_classifier",
opensource_classifier_config: {
provider,
model: provider === "laya" ? "multilingual" : "jev-test",
timeout_ms: 4200,
circuit_breaker_enabled: false,
circuit_breaker_cooldown_seconds: 50,
},
tiers: expect.objectContaining({ QUICK: ["fast"] }),
}),
}),
}),
);
});
);
},
);
it("allows licensed instructions and can restore built-in instructions", () => {
const authorized = useAuthorized();
@ -152,10 +169,10 @@ describe("JEV classifier editor", () => {
return <JevEditor value={value} onChange={setValue} />;
};
renderWithProviders(<LicensedForm />);
expect(screen.getByLabelText("Jev Instructions")).toBeEnabled();
fireEvent.change(screen.getByLabelText("Jev Instructions"), { target: { value: "New instructions" } });
expect(screen.getByLabelText("Jev Instructions")).toHaveValue("New instructions");
fireEvent.click(screen.getByRole("button", { name: "Restore built-in Jev instructions" }));
expect(screen.getByLabelText("Jev Instructions")).toHaveValue("");
expect(screen.getByLabelText("Classifier Instructions")).toBeEnabled();
fireEvent.change(screen.getByLabelText("Classifier Instructions"), { target: { value: "New instructions" } });
expect(screen.getByLabelText("Classifier Instructions")).toHaveValue("New instructions");
fireEvent.click(screen.getByRole("button", { name: "Restore built-in instructions" }));
expect(screen.getByLabelText("Classifier Instructions")).toHaveValue("");
});
});

View file

@ -6,7 +6,8 @@ import { Label } from "@/components/ui/label";
import { Textarea } from "@/components/ui/textarea";
import ClassifierCircuitBreakerConfig from "./ClassifierCircuitBreakerConfig";
import type { ComplexityRouterConfigValue } from "./ComplexityRouterConfig";
import { defaultJevClassifierConfig } from "./jev_classifier_config";
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
import { defaultJevClassifierConfig, LAYA_MODELS } from "./jev_classifier_config";
export default function JevClassifierConfig({
value,
@ -17,20 +18,38 @@ export default function JevClassifierConfig({
}) {
const id = useId();
const config = value.jev_classifier_config ?? defaultJevClassifierConfig();
const isLaya = config.provider === "laya";
const update = (patch: Partial<typeof config>) =>
onChange({ ...value, jev_classifier_config: { ...config, ...patch } });
return (
<div className="mt-4 space-y-3">
<p className="text-sm text-muted-foreground">
Uses TypeSafe System One Choice evaluation with your configured tiers
{isLaya
? "Uses Laya with your configured tiers. Set LAYA_API_BASE on the gateway to connect your Laya server."
: "Uses TypeSafe System One Choice evaluation with your configured tiers"}
</p>
<div>
<Label htmlFor={`${id}-model`}>Jev Model</Label>
<Input id={`${id}-model`} value={config.model} onChange={(event) => update({ model: event.target.value })} />
<Label htmlFor={`${id}-model`}>Classifier Model</Label>
{isLaya ? (
<Select value={config.model} onValueChange={(model) => model && update({ model })}>
<SelectTrigger id={`${id}-model`} className="w-full">
<SelectValue />
</SelectTrigger>
<SelectContent>
{LAYA_MODELS.map((model) => (
<SelectItem key={model} value={model}>
{model}
</SelectItem>
))}
</SelectContent>
</Select>
) : (
<Input id={`${id}-model`} value={config.model} onChange={(event) => update({ model: event.target.value })} />
)}
</div>
<div>
<Label htmlFor={`${id}-timeout`}>Jev Timeout (ms)</Label>
<Label htmlFor={`${id}-timeout`}>Classifier Timeout (ms)</Label>
<Input
id={`${id}-timeout`}
type="number"
@ -50,7 +69,7 @@ export default function JevClassifierConfig({
}
/>
<div>
<Label htmlFor={`${id}-instructions`}>Jev Instructions</Label>
<Label htmlFor={`${id}-instructions`}>Classifier Instructions</Label>
<AutoRouterAllowanceNote
feature="tier_or_classifier_prompt"
label="Custom instructions share the custom-tier allowance"
@ -63,11 +82,11 @@ export default function JevClassifierConfig({
/>
{config.instructions && (
<Button variant="outline" type="button" onClick={() => update({ instructions: undefined })}>
Restore built-in Jev instructions
Restore built-in instructions
</Button>
)}
<p className="text-xs text-muted-foreground">
Built-in Jev is available without a license and uses the shipped tier criteria
Built-in OSS classification is available without a license and uses the shipped tier criteria
</p>
</div>
</div>

View file

@ -107,10 +107,10 @@ describe("JEV network probes", () => {
expect(JSON.parse(String(routingCall?.[1]?.body))).toEqual(expectedRequest);
expect(fetchMock).toHaveBeenCalledTimes(5);
expect(screen.getAllByTestId("test-status-success")).toHaveLength(4);
expect(screen.getByRole("status", { name: "Jev connection" })).toHaveTextContent(
expect(screen.getByRole("status", { name: "OSS classifier connection" })).toHaveTextContent(
cause === "jev_classifier"
? "Jev classification succeeded"
: `Jev was not reached successfully (routing cause: ${cause})`,
? "OSS classification succeeded"
: `OSS classifier was not reached successfully (routing cause: ${cause})`,
);
},
);
@ -131,7 +131,7 @@ describe("JEV network probes", () => {
);
fireEvent.change(screen.getByTestId("auto-router-routing-test-prompt"), { target: { value: "Hello" } });
fireEvent.click(screen.getByTestId("auto-router-routing-test-send"));
expect(await screen.findByText("JEV classifier")).toBeInTheDocument();
expect(await screen.findByText("OSS classifier")).toBeInTheDocument();
expect(screen.getByText("jev-latest")).toBeInTheDocument();
expect(screen.getByText("80.0%")).toBeInTheDocument();
expect(screen.getByText("SIMPLE: 80.0%")).toBeInTheDocument();

View file

@ -274,13 +274,13 @@ describe("AddAutoRouterTab", () => {
mockFetchAvailableModels.mockResolvedValue(ALL_FAMILY_MODELS);
renderWithProviders(<Harness />);
await user.click(await screen.findByRole("button", { name: "Choose models for me" }));
await user.click(screen.getByRole("radio", { name: "Jev" }));
await user.click(screen.getByRole("radio", { name: "OSS Classifier" }));
await waitFor(() =>
expect(apiClient.post).toHaveBeenLastCalledWith(
"/auto_router/availability",
expect.objectContaining({
body: expect.objectContaining({
complexity_router_config: expect.objectContaining({ classifier_type: "jev" }),
complexity_router_config: expect.objectContaining({ classifier_type: "oss_classifier" }),
}),
}),
),
@ -301,13 +301,13 @@ describe("AddAutoRouterTab", () => {
expect(within(screen.getByRole("alert")).getByRole("link", { name: "Talk to our team" })).toBeVisible();
await user.click(screen.getByRole("button", { name: "Restore defaults" }));
await waitFor(() => expect(screen.queryByRole("alert")).not.toBeInTheDocument());
expect(screen.getByRole("radio", { name: "Jev" })).toBeChecked();
expect(screen.getByRole("radio", { name: "OSS Classifier" })).toBeChecked();
await waitFor(() => expect(screen.getByRole("button", { name: "Add Auto Router" })).toBeEnabled());
await user.click(screen.getByRole("button", { name: "Add Auto Router" }));
await waitFor(() => expect(handleAddAutoRouterSubmit).toHaveBeenCalled());
const saved = vi.mocked(handleAddAutoRouterSubmit).mock.calls.at(-1)?.[0].complexity_router_config;
expect(saved).not.toHaveProperty("tier_definitions");
expect(saved?.classifier_type).toBe("jev");
expect(saved?.classifier_type).toBe("oss_classifier");
expect(Object.keys(saved?.tiers ?? {})).toEqual(["SIMPLE", "MEDIUM", "COMPLEX", "REASONING"]);
expect(saved?.tiers).toEqual(initialRequest.complexity_router_config.tiers);
});
@ -317,7 +317,7 @@ describe("AddAutoRouterTab", () => {
mockFetchAvailableModels.mockResolvedValue(ALL_FAMILY_MODELS);
renderWithProviders(<Harness />);
await user.click(await screen.findByRole("button", { name: "Choose models for me" }));
await user.click(screen.getByRole("radio", { name: "Jev" }));
await user.click(screen.getByRole("radio", { name: "OSS Classifier" }));
fireEvent.change(screen.getByLabelText("Auto Router Name"), { target: { value: "checked-router" } });
await waitFor(() => expect(screen.getByRole("button", { name: "Add Auto Router" })).toBeEnabled());
let complete: ((result: unknown) => void) | undefined;
@ -357,17 +357,22 @@ describe("AddAutoRouterTab", () => {
expect(screen.getByRole("button", { name: "Routing approach" })).toHaveTextContent("Complexity");
});
it.each(["LLM", "Jev"])("keeps %s and the frequency when choosing models automatically", async (family) => {
mockFetchAvailableModels.mockResolvedValue(ALL_FAMILY_MODELS);
renderWithProviders(<Harness />);
const automatic = await screen.findByRole("button", { name: "Choose models for me" });
await userEvent.click(screen.getByRole("radio", { name: family }));
await selectAutoRouterOption("How often to classify", "Every new user message");
await userEvent.click(automatic);
expect(screen.getByRole("radio", { name: family })).toBeChecked();
expect(screen.getByRole("combobox", { name: "How often to classify" })).toHaveTextContent("Every new user message");
expect(screen.getByRole("button", { name: "Advanced settings" })).toHaveAttribute("aria-expanded", "false");
});
it.each(["LLM", "OSS Classifier"])(
"keeps %s and the frequency when choosing models automatically",
async (family) => {
mockFetchAvailableModels.mockResolvedValue(ALL_FAMILY_MODELS);
renderWithProviders(<Harness />);
const automatic = await screen.findByRole("button", { name: "Choose models for me" });
await userEvent.click(screen.getByRole("radio", { name: family }));
await selectAutoRouterOption("How often to classify", "Every new user message");
await userEvent.click(automatic);
expect(screen.getByRole("radio", { name: family })).toBeChecked();
expect(screen.getByRole("combobox", { name: "How often to classify" })).toHaveTextContent(
"Every new user message",
);
expect(screen.getByRole("button", { name: "Advanced settings" })).toHaveAttribute("aria-expanded", "false");
},
);
it.each(["Capability", "Fuse v2"])(
"creates %s from its dedicated tab without complexity templates",
@ -1902,7 +1907,7 @@ describe("preset catalog fetch states", () => {
await waitFor(() => expect(handleAddAutoRouterSubmit).toHaveBeenCalledOnce());
expect(vi.mocked(handleAddAutoRouterSubmit).mock.calls[0][0].complexity_router_config).toMatchObject({
classifier_type: "jev",
classifier_type: "oss_classifier",
classifier_context_per_turn_chars: 450,
});
});

View file

@ -50,7 +50,7 @@ const AutoRouterConnectionTest: React.FC<AutoRouterConnectionTestProps> = ({
? { status: "success" }
: {
status: "error",
error: `Jev was not reached successfully (routing cause: ${decision.cause ?? "unknown"})`,
error: `OSS classifier was not reached successfully (routing cause: ${decision.cause ?? "unknown"})`,
},
);
};
@ -91,11 +91,11 @@ const AutoRouterConnectionTest: React.FC<AutoRouterConnectionTestProps> = ({
classifier probe includes its reasoning effort override.
</p>
{jevRequest && (
<div role="status" aria-label="Jev connection" className="rounded-lg border p-3 text-sm">
<strong>Jev Classifier</strong>
<div role="status" aria-label="OSS classifier connection" className="rounded-lg border p-3 text-sm">
<strong>OSS Classifier</strong>
<p>
{jevResult.status === "pending" && "Testing Jev classification"}
{jevResult.status === "success" && "Jev classification succeeded"}
{jevResult.status === "pending" && "Testing OSS classification"}
{jevResult.status === "success" && "OSS classification succeeded"}
{jevResult.status === "error" && jevResult.error}
</p>
</div>

View file

@ -33,20 +33,20 @@ describe("buildAutoRouterRoutingTestRequest", () => {
const expectedRequest = {
prompt: JEV_CONNECTION_TEST_PROMPT,
complexity_router_config: {
classifier_type: "jev",
classifier_type: "oss_classifier",
tiers: CONFIG.tiers,
jev_classifier_config: defaultJevClassifierConfig(),
opensource_classifier_config: defaultJevClassifierConfig(),
},
saved_model_id: "saved-id",
};
expect(request).toEqual(expectedRequest);
expect(request?.complexity_router_config.jev_classifier_config).not.toHaveProperty("api_key");
expect(request?.complexity_router_config.jev_classifier_config).not.toHaveProperty("api_base");
expect(request?.complexity_router_config.opensource_classifier_config).not.toHaveProperty("api_key");
expect(request?.complexity_router_config.opensource_classifier_config).not.toHaveProperty("api_base");
});
it.each(["object", "json"])("probes saved JEV %s configuration with custom tiers and team context", (format) => {
it.each(["object", "json"])("probes saved Laya %s configuration with custom tiers and team context", (format) => {
const config = {
classifier_type: "jev",
jev_classifier_config: { model: "jev-test", timeout_ms: 900 },
classifier_type: "oss_classifier",
opensource_classifier_config: { provider: "laya", model: "english", timeout_ms: 900 },
tiers: { QUICK: ["fast"], DEEP: ["strong"] },
tier_definitions: { QUICK: "Simple questions", DEEP: "Complex questions" },
fallback_tier: "DEEP",

View file

@ -1,7 +1,7 @@
import { AutoRouterRoutingTestRequest } from "../networking";
import { ComplexityRouterConfigPayload } from "./build_complexity_router_config";
import { z } from "zod";
import { jevClassifierConfigSchema } from "./jev_classifier_config";
import { hydrateOssClassifier, jevClassifierConfigSchema, normalizeJevClassifierConfig } from "./jev_classifier_config";
export const JEV_CONNECTION_TEST_PROMPT = "What is 2 plus 2?";
@ -23,16 +23,23 @@ export const buildSavedJevConnectionTestRequest = (
: rawConfig;
const result = z
.object({
classifier_type: z.literal("jev"),
classifier_type: z.enum(["jev", "oss_classifier"]),
tiers: z.record(z.unknown()),
jev_classifier_config: jevClassifierConfigSchema.default({}),
jev_classifier_config: jevClassifierConfigSchema.optional(),
opensource_classifier_config: jevClassifierConfigSchema.optional(),
})
.passthrough()
.safeParse(parsed);
if (!result.success) return undefined;
const { jev_classifier_config, opensource_classifier_config, ...config } = result.data;
const classifier = hydrateOssClassifier({ ...config, jev_classifier_config, opensource_classifier_config });
return {
prompt: JEV_CONNECTION_TEST_PROMPT,
complexity_router_config: result.data,
complexity_router_config: {
...config,
classifier_type: "oss_classifier",
opensource_classifier_config: normalizeJevClassifierConfig(classifier.jev_classifier_config),
},
saved_model_id: savedModelId,
...(teamId && { team_id: teamId }),
};

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