+ 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
+
+ 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.
+
- 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.
+
+
+
+
+
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/TierTurnsChart.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/TierTurnsChart.test.tsx
index e4417d77463..42444fd8f06 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/TierTurnsChart.test.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/TierTurnsChart.test.tsx
@@ -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,
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/autoRouterBenchmarks.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/autoRouterBenchmarks.test.ts
index 0586163e77e..62d0c0e4e5e 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/autoRouterBenchmarks.test.ts
+++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/autoRouterBenchmarks.test.ts
@@ -43,7 +43,6 @@ const totals = (overrides: Partial = {}) => ({
saved_spend: 2174.59,
baseline_spend: 2534.45,
saved_pct: 85.8,
- saved_per_session: 23.13,
cache: cache(),
...overrides,
});
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/users/_components/view_users/user_info_view.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/users/_components/view_users/user_info_view.integration.test.tsx
index 676bc7d29eb..4195a3c9913 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/users/_components/view_users/user_info_view.integration.test.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/users/_components/view_users/user_info_view.integration.test.tsx
@@ -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,
diff --git a/ui/litellm-dashboard/src/components/templates/KeyAutoRouterUsageTab.integration.test.tsx b/ui/litellm-dashboard/src/components/templates/KeyAutoRouterUsageTab.integration.test.tsx
index 5c7b9b28789..2ce585e895e 100644
--- a/ui/litellm-dashboard/src/components/templates/KeyAutoRouterUsageTab.integration.test.tsx
+++ b/ui/litellm-dashboard/src/components/templates/KeyAutoRouterUsageTab.integration.test.tsx
@@ -39,7 +39,6 @@ const stats = {
saved_spend: 8.75,
baseline_spend: 10,
saved_pct: 87.5,
- saved_per_session: 4.375,
cache,
};
diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts
index 1d50f37e573..fc1a9c67946 100644
--- a/ui/litellm-dashboard/src/lib/http/schema.d.ts
+++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts
@@ -1350,9 +1350,10 @@ export interface paths {
*
* 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
@@ -26169,12 +26170,21 @@ export interface components {
* @description One auto-router's slice of the benchmarks.
*/
AutoRouterBenchmarkGroup: {
- /** Avg Session Seconds */
- avg_session_seconds: number;
- /** Avg Tokens Per Session */
- avg_tokens_per_session: number;
- /** Avg Turns Per Session */
- avg_turns_per_session: number;
+ /**
+ * Avg Session Seconds
+ * @description Lifetime seconds per overlapping session; null as above
+ */
+ avg_session_seconds: number | null;
+ /**
+ * Avg Tokens Per Session
+ * @description Lifetime tokens per overlapping session; null as above
+ */
+ avg_tokens_per_session: number | null;
+ /**
+ * Avg Turns Per Session
+ * @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_turns_per_session: number | null;
/**
* Baseline Spend
* @description Estimated single-model cost: compared actual spend plus recorded savings; null when traffic has no recorded savings
@@ -26201,14 +26211,9 @@ export interface components {
* @description Recorded savings over baseline_spend, as a percentage
*/
saved_pct: number | null;
- /**
- * Saved Per Session
- * @description Recorded savings per session, including historical estimates
- */
- saved_per_session: number | null;
/**
* Saved Spend
- * @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
*/
saved_spend: number | null;
/**
@@ -26226,11 +26231,14 @@ export interface components {
* @description Requests compared against the baseline: every request on complexity routers that recorded savings
*/
savings_estimated_turns: number;
- /** Sessions */
+ /**
+ * Sessions
+ * @description Sessions overlapping the window, counted whole
+ */
sessions: number;
/**
* Spend
- * @description What the routed traffic actually cost
+ * @description What the selected days' routed traffic actually cost
*/
spend: number;
/**
@@ -26240,20 +26248,38 @@ export interface components {
tier_turns?: {
[key: string]: number;
};
- /** Turns */
+ /**
+ * Turns
+ * @description Auto-routed requests on the selected UTC days
+ */
turns: number;
+ /**
+ * Unattributed Saved Spend
+ * @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
+ */
+ unattributed_saved_spend?: number | null;
};
/**
* AutoRouterBenchmarkTotals
- * @description Session-shape and savings aggregates over auto-routed traffic in the window.
+ * @description 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.
*/
AutoRouterBenchmarkTotals: {
- /** Avg Session Seconds */
- avg_session_seconds: number;
- /** Avg Tokens Per Session */
- avg_tokens_per_session: number;
- /** Avg Turns Per Session */
- avg_turns_per_session: number;
+ /**
+ * Avg Session Seconds
+ * @description Lifetime seconds per overlapping session; null as above
+ */
+ avg_session_seconds: number | null;
+ /**
+ * Avg Tokens Per Session
+ * @description Lifetime tokens per overlapping session; null as above
+ */
+ avg_tokens_per_session: number | null;
+ /**
+ * Avg Turns Per Session
+ * @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_turns_per_session: number | null;
/**
* Baseline Spend
* @description Estimated single-model cost: compared actual spend plus recorded savings; null when traffic has no recorded savings
@@ -26270,14 +26296,9 @@ export interface components {
* @description Recorded savings over baseline_spend, as a percentage
*/
saved_pct: number | null;
- /**
- * Saved Per Session
- * @description Recorded savings per session, including historical estimates
- */
- saved_per_session: number | null;
/**
* Saved Spend
- * @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
*/
saved_spend: number | null;
/**
@@ -26295,19 +26316,30 @@ export interface components {
* @description Requests compared against the baseline: every request on complexity routers that recorded savings
*/
savings_estimated_turns: number;
- /** Sessions */
+ /**
+ * Sessions
+ * @description Sessions overlapping the window, counted whole
+ */
sessions: number;
/**
* Spend
- * @description What the routed traffic actually cost
+ * @description What the selected days' routed traffic actually cost
*/
spend: number;
- /** Turns */
+ /**
+ * Turns
+ * @description Auto-routed requests on the selected UTC days
+ */
turns: number;
+ /**
+ * Unattributed Saved Spend
+ * @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
+ */
+ unattributed_saved_spend?: number | null;
};
/**
* AutoRouterBenchmarksResponse
- * @description Benchmarks for the auto-router dashboard, aggregated from the per-session rollup.
+ * @description Benchmarks for the auto-router dashboard, aggregated from the per-session and per-day rollups.
*/
AutoRouterBenchmarksResponse: {
/**
From 2584721ca3e88dbbb4eaf4a79c00d3ead1b52832 Mon Sep 17 00:00:00 2001
From: moe-berri
Date: Fri, 2 Oct 2026 11:58:32 -0700
Subject: [PATCH 4/8] fix(lens): preserve framework agent names and GenAI
message content (#44218)
* fix(lens): use recorded agent identities across framework traces
* fix(lens): tighten agent identity and bound trace lookups
* style(tracing): wrap framework agent identity test case
---
.../crates/traces/query/list_traces.sql | 22 ++-
.../crates/traces/src/normalize/mod.rs | 72 +++++++-
litellm-rust/crates/traces/src/otlp/span.rs | 13 +-
.../crates/traces/tests/migrations.rs | 166 ++++++++++++++++++
litellm/tracing/store.py | 10 +-
litellm/tracing/types.py | 1 +
tests/test_litellm/tracing/test_decode.py | 83 ++++++++-
tests/test_litellm/tracing/test_store.py | 17 ++
.../TraceView/AgentTracesSection.test.tsx | 11 +-
.../TraceView/AgentTracesSection.tsx | 6 +-
.../view_logs/TraceView/AgentTracesTable.tsx | 6 +-
.../view_logs/TraceView/traceUtils.test.ts | 13 ++
.../view_logs/TraceView/traceUtils.ts | 23 ++-
ui/litellm-dashboard/src/lib/http/schema.d.ts | 2 +
14 files changed, 419 insertions(+), 26 deletions(-)
diff --git a/litellm-rust/crates/traces/query/list_traces.sql b/litellm-rust/crates/traces/query/list_traces.sql
index c0c1b28aa7f..163d9b03d6b 100644
--- a/litellm-rust/crates/traces/query/list_traces.sql
+++ b/litellm-rust/crates/traces/query/list_traces.sql
@@ -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
diff --git a/litellm-rust/crates/traces/src/normalize/mod.rs b/litellm-rust/crates/traces/src/normalize/mod.rs
index b4d5a4a9d06..60dc8f816f2 100644
--- a/litellm-rust/crates/traces/src/normalize/mod.rs
+++ b/litellm-rust/crates/traces/src/normalize/mod.rs
@@ -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) -> 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,
+ 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::(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),
})
}
diff --git a/litellm-rust/crates/traces/src/otlp/span.rs b/litellm-rust/crates/traces/src/otlp/span.rs
index 1c1e53e756c..0ae5725ba47 100644
--- a/litellm-rust/crates/traces/src/otlp/span.rs
+++ b/litellm-rust/crates/traces/src/otlp/span.rs
@@ -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()
diff --git a/litellm-rust/crates/traces/tests/migrations.rs b/litellm-rust/crates/traces/tests/migrations.rs
index ccefa8b0b9f..c62e3538ddb 100644
--- a/litellm-rust/crates/traces/tests/migrations.rs
+++ b/litellm-rust/crates/traces/tests/migrations.rs
@@ -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,
+) -> 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::, _>>()?;
+ 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,
+ ¶meters,
+ )
+ .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::>();
+ 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::>();
+ 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,
+ ¶meters,
+ )
+ .await?,
+ )?;
+ assert_eq!(
+ listed["data"][0]["agent_names"],
+ serde_json::json!(["lead", "researcher"])
+ );
+ assert_eq!(listed["data"][0]["agent_count"], 2);
Ok(())
}
diff --git a/litellm/tracing/store.py b/litellm/tracing/store.py
index 91420ffd025..edfe1285fc7 100644
--- a/litellm/tracing/store.py
+++ b/litellm/tracing/store.py
@@ -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,
diff --git a/litellm/tracing/types.py b/litellm/tracing/types.py
index ff965483013..a0e824982fa 100644
--- a/litellm/tracing/types.py
+++ b/litellm/tracing/types.py
@@ -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]
diff --git a/tests/test_litellm/tracing/test_decode.py b/tests/test_litellm/tracing/test_decode.py
index 21a79dd6b87..dda075863ac 100644
--- a/tests/test_litellm/tracing/test_decode.py
+++ b/tests/test_litellm/tracing/test_decode.py
@@ -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),
diff --git a/tests/test_litellm/tracing/test_store.py b/tests/test_litellm/tracing/test_store.py
index 3f43e42842c..0dd615a96ac 100644
--- a/tests/test_litellm/tracing/test_store.py
+++ b/tests/test_litellm/tracing/test_store.py
@@ -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
diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.test.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.test.tsx
index 5db34ef9788..d9cb4ee0399 100644
--- a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.test.tsx
+++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.test.tsx
@@ -271,10 +271,13 @@ describe("AgentTracesSection", () => {
expect(rows[0]).toHaveTextContent("Should we store OTEL agent spans");
});
- it("labels the OTEL service as the agent and filters runs by it", async () => {
+ it("uses recorded agent names for the column and filter even when services are shared", async () => {
vi.mocked(agentTraceListCall).mockResolvedValue({
...(traceList as TracePage),
- data: [...runs.slice(1), { ...runs[0], service: "billing-agent" }],
+ data: [
+ ...runs.slice(1).map((run) => ({ ...run, service: "shared-app", agent_names: ["research-agent"] })),
+ { ...runs[0], service: "shared-app", agent_names: ["billing-agent", "review-agent"] },
+ ],
});
const user = userEvent.setup();
renderSection();
@@ -288,6 +291,10 @@ describe("AgentTracesSection", () => {
const rows = screen.getAllByTestId("agent-trace-row");
expect(rows).toHaveLength(1);
expect(rows[0]).toHaveTextContent("billing-agent");
+ expect(rows[0]).not.toHaveTextContent("shared-app");
+
+ await chooseSelectOption(user, agentFilter, "review-agent");
+ expect(screen.getAllByTestId("agent-trace-row")).toHaveLength(1);
await chooseSelectOption(user, agentFilter, "All agents");
expect(screen.getAllByTestId("agent-trace-row")).toHaveLength(runs.length);
diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx
index 06094864230..4ec8fe7e0ec 100644
--- a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx
+++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx
@@ -9,7 +9,7 @@ import { AgentTracesTable } from "./AgentTracesTable";
import { RunDrawer } from "./RunDrawer";
import { ALL_AGENTS, RunsToolbar, type RunStatusFilter } from "./RunsToolbar";
import type { TraceSummary } from "./traceTypes";
-import { previewText } from "./traceUtils";
+import { previewText, traceAgentNames } from "./traceUtils";
import { TimeRangeControls } from "./TimeRangeControls";
import { TracesTimeline, type TimeWindow } from "./TracesTimeline";
import { ActiveDot } from "./ActiveDot";
@@ -27,7 +27,7 @@ export function filterRuns(
return runs.filter((run) => {
const haystack = [run.trace_id, previewText(run.input_preview), run.name].map((s) => s.toLowerCase());
const matchesQuery = !q || haystack.some((text) => text.includes(q));
- const matchesAgent = agent === ALL_AGENTS || run.service === agent;
+ const matchesAgent = agent === ALL_AGENTS || traceAgentNames(run).includes(agent);
const failed = run.error_count > 0;
const matchesStatus = status === "all" || (status === "error" ? failed : !failed);
return matchesQuery && matchesAgent && matchesStatus;
@@ -121,7 +121,7 @@ export function AgentTracesSection({
if (setup.disabledDetail == null) void history.refetch();
};
- const agents = useMemo(() => Array.from(new Set(traces.traces.map((t) => t.service))).sort(), [traces.traces]);
+ const agents = useMemo(() => Array.from(new Set(traces.traces.flatMap(traceAgentNames))).sort(), [traces.traces]);
// Relative ranges end "now" (the list query uses Date.now() too); round to the minute so the histogram is stable.
const endMs = isCustomDate ? moment(endTime).valueOf() : moment().endOf("minute").valueOf();
const range = useMemo(
diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx
index 60f6816bd7a..6f0c8eac5f0 100644
--- a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx
+++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx
@@ -8,7 +8,7 @@ import { cn } from "@/lib/cva.config";
import { StatusMark } from "./StatusMark";
import type { TraceSummary } from "./traceTypes";
-import { fmtMs, previewText, traceDisplayName } from "./traceUtils";
+import { fmtMs, previewText, traceDisplayName, traceAgentNames } from "./traceUtils";
interface AgentTracesTableProps {
traces: TraceSummary[];
@@ -83,8 +83,8 @@ export function AgentTracesTable({
>
{formatActivityTimestamp(run.start_time)}
-
- {run.service}
+
+ {traceAgentNames(run).join(", ") || "—"}
diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.test.ts b/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.test.ts
index 9353f2024c8..f1394cec980 100644
--- a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.test.ts
+++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.test.ts
@@ -287,6 +287,19 @@ describe("payload helpers", () => {
expect(messageText(image)).toBe(image);
expect(messageText("[not json")).toBe("[not json");
});
+
+ it("reads GenAI message parts and native content arrays without crashing previews", () => {
+ const question = "What is an agent trace?";
+ const parts = [{ type: "text", content: question }];
+ const input = JSON.stringify([{ role: "user", parts }]);
+ expect(parseMessages(input)).toEqual([{ role: "user", parts, content: question }]);
+ expect(previewText(input)).toBe(question);
+ expect(
+ parseMessages(JSON.stringify({ role: "assistant", content: [{ type: "text", text: "An execution record" }] })),
+ ).toEqual([{ role: "assistant", content: "An execution record" }]);
+ expect(parseMessages('[{"role":"assistant","tool_calls":[]}]')).toBeNull();
+ expect(parseMessages('[{"role":"user","content":42}]')).toBeNull();
+ });
});
describe("treeGuides", () => {
diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.ts b/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.ts
index 00d24d4eaef..4379255c430 100644
--- a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.ts
+++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.ts
@@ -9,6 +9,9 @@ import type { Span, TraceMessage, TraceSummary } from "./traceTypes";
/* Formatting */
/* ------------------------------------------------------------------ */
+export const traceAgentNames = (trace: TraceSummary): readonly string[] =>
+ trace.agent_names ?? (trace.service ? [trace.service] : []);
+
export const fmtMs = (ms: number): string => {
if (ms >= 60_000) return `${(ms / 60_000).toFixed(1)}m`;
if (ms >= 1000) return `${(ms / 1000).toFixed(2)}s`;
@@ -267,14 +270,10 @@ export const parseJson = (value: string): unknown => {
}
};
-const isMessage = (value: unknown): value is TraceMessage => {
- const isObject = typeof value === "object" && value !== null;
- return isObject && "role" in value && typeof (value as TraceMessage).role === "string";
-};
-
const blockText = (block: unknown): string | null => {
if (typeof block !== "object" || block === null) return null;
- const text: unknown = Reflect.get(block, "text");
+ const text: unknown =
+ Reflect.get(block, "text") ?? (Reflect.get(block, "type") === "text" ? Reflect.get(block, "content") : undefined);
return typeof text === "string" ? text : null;
};
@@ -296,13 +295,19 @@ export function messageText(content: string): string {
.join("\n\n");
}
-const withText = (message: TraceMessage): TraceMessage => ({ ...message, content: messageText(message.content) });
+const parseMessage = (value: unknown): TraceMessage | null => {
+ if (typeof value !== "object" || value === null) return null;
+ const role: unknown = Reflect.get(value, "role");
+ const content: unknown = Reflect.get(value, "content") ?? Reflect.get(value, "parts");
+ if (typeof role !== "string" || (typeof content !== "string" && !Array.isArray(content))) return null;
+ return { ...value, role, content: messageText(typeof content === "string" ? content : JSON.stringify(content)) };
+};
/** An llm span's input (array of messages) or output (one message); null when it isn't one. */
export function parseMessages(value: string): TraceMessage[] | null {
const parsed = parseJson(value);
- if (Array.isArray(parsed)) return parsed.every(isMessage) ? parsed.map(withText) : null;
- return isMessage(parsed) ? [withText(parsed)] : null;
+ const messages = (Array.isArray(parsed) ? parsed : [parsed]).map(parseMessage);
+ return messages.every((message) => message !== null) ? messages : null;
}
/** Pretty JSON when the payload is JSON, else the raw string. */
diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts
index fc1a9c67946..a376f8d7910 100644
--- a/ui/litellm-dashboard/src/lib/http/schema.d.ts
+++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts
@@ -46904,6 +46904,8 @@ export interface components {
agent_count: number;
/** Agent Invocations */
agent_invocations: number;
+ /** Agent Names */
+ agent_names?: string[];
/** Duration Ms */
duration_ms: number;
/** Error Count */
From b21e44cbf97b1b1b5ce71f7edc31059c7836cf17 Mon Sep 17 00:00:00 2001
From: "devin-ai-integration[bot]"
<158243242+devin-ai-integration[bot]@users.noreply.github.com>
Date: Fri, 2 Oct 2026 12:11:33 -0700
Subject: [PATCH 5/8] feat(jwt): auto_register_map_existing_key maps JWT to the
user's existing virtual key (#42375)
* test(e2e): jwt auto_register map-existing-key repro
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* feat(jwt): auto_register_map_existing_key maps JWT to the user's existing virtual key
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(jwt): exclude blocked keys from auto_register_map_existing_key reuse
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* refactor(jwt): route existing-key lookup through VerificationTokenRepository
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(e2e): stop requiring LITELLM_SALT_KEY for the owned JWT gateway
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(e2e): gate the owned JWT gateway tests behind E2E_OWNED_GATEWAY
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(jwt): only reuse keys that can call LLM routes in auto_register_map_existing_key
Skip Admin UI session keys and keys whose allowed_routes restrict them to
anything other than llm_api_routes (management, read_only, password-reset
sessions). Mapping a JWT to one of those left the user with 401s or 403s on
every LLM call, since the mapping persists.
* fix(jwt): scope auto_register_map_existing_key reuse to the JWT-resolved team
Only reuse a key whose team_id matches the team auth_builder resolved for
the JWT (no team matches no team), so a personal key can no longer bypass
the resolved team's model and budget limits.
With the flag on, the first JWT request now falls through to the same
virtual-key checks later mapped requests get, instead of returning early,
so a reused key's own limits apply from request one rather than 200 then
403. Flag off keeps the early return unchanged.
* fix(jwt): keep the early return when no master key is set
Without a master key the generic virtual-key path returns a bare
INTERNAL_USER object, so falling through on the first auto-registered
request dropped the key's team, models and budgets. Only fall through when
a master key is configured.
Tests now assert the reused key per team rather than the query shape, and
cover the flag-off early return and the no-master-key case.
* test(jwt): assert on race-loser's returned key, not only mocks (TQ002)
Co-Authored-By: Claude Opus 5.5
* fix(jwt): close the auto_register_map_existing_key race, shared-claim and expiry holes
A key auto_register just minted is never adopted by a concurrent request, so the race loser's cleanup can no longer delete a key another request mapped and cascade its mapping away (503, user left with no key)
Reuse only happens when the claim value is the JWT-resolved user_id. A shared claim such as azp or client_id falls back to minting, so one user can no longer land on another user's personal key and budget
Only keys that never expire are reused, so an expiring key can no longer pin the claim to a permanent 401
Integration tests on a real proxy and Postgres cover all three. The race test holds the first mapping insert in a Postgres relay, so the interleaving is forced rather than timed. The where-clause shape unit tests are replaced by these, since only a real database proves the filter
* test(e2e): create the reused key in the team the JWT resolves to
The flag only reuses a key in the JWT-resolved team, and this identity's groups claim resolves to its team, so a teamless key was never eligible and the test could not pass
* test(integration): match the held statement across TCP reads
The relay looked for the trigger inside one read, so an insert split across two reads was never held and the race test would fail waiting for it. It now matches one exact trigger over a window that keeps the end of the previous read
* fix(jwt): gate key reuse on the claim field, not on the claim value
Requiring the claim value to equal the resolved user_id skipped reuse for users matched through the sso_user_id or case-insensitive email fallback, whose stored user_id differs from the JWT sub. That is the lookup LIT-5378 asks for. Reuse is now allowed when the virtual key claim is the user_id or user_email JWT field, globally or for the token's issuer, which still keeps shared claims such as azp or client_id on the mint path
* fix(jwt): let an issuer's own user field replace the global one when gating key reuse
An issuer that identifies users by uid no longer treats the global sub field as a user identity claim, so a shared sub under that issuer mints instead of reusing a personal key
* test(jwt): make the flag-off test fail when the flag no longer gates key reuse
The flag-off test used a config where sub was not a user identity claim, so deleting the flag check still passed. Configure user_id_jwt_field=sub so only the flag keeps the lookup off, and drop test docstrings
* chore(lint): drop mutable-ok suppressions that LIT013 flags as no-ops
---------
Co-authored-by: yuneng
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: mrinal
Co-authored-by: Mrinal Chanshetty
Co-authored-by: Claude Opus 5.5
---
.github/workflows/test-e2e-changed.yml | 1 +
litellm/proxy/_types.py | 20 +
litellm/proxy/auth/user_api_key_auth.py | 238 ++++++----
.../verification_token_repository.py | 32 ++
tests/e2e/conftest.py | 7 +
tests/e2e/coverage_registry/other.yaml | 3 +
tests/e2e/e2e_config.py | 12 +
tests/e2e/mcp/oauth_gateway.py | 12 +-
tests/e2e/models.py | 35 ++
tests/e2e/other/other_client.py | 46 ++
tests/e2e/other/owned_jwt_gateway.py | 108 +++++
tests/e2e/other/test_jwt_auto_register_e2e.py | 182 +++++++
tests/e2e/pytest.ini | 1 +
tests/integration/_support/database_relay.py | 80 +++-
...test_jwt_auto_register_map_existing_key.py | 219 +++++++++
.../test_user_api_key_auth_request_flow.py | 448 ++++++++++++++++++
16 files changed, 1333 insertions(+), 111 deletions(-)
create mode 100644 tests/e2e/other/owned_jwt_gateway.py
create mode 100644 tests/e2e/other/test_jwt_auto_register_e2e.py
create mode 100644 tests/integration/authorization/test_jwt_auto_register_map_existing_key.py
diff --git a/.github/workflows/test-e2e-changed.yml b/.github/workflows/test-e2e-changed.yml
index 8e03a902383..228e23f60d7 100644
--- a/.github/workflows/test-e2e-changed.yml
+++ b/.github/workflows/test-e2e-changed.yml
@@ -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
diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py
index ef4c545507b..cabb04cc52b 100644
--- a/litellm/proxy/_types.py
+++ b/litellm/proxy/_types.py
@@ -5464,6 +5464,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.",
@@ -5564,6 +5575,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:
diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py
index 3400dccf2a7..e82f3eed7cc 100644
--- a/litellm/proxy/auth/user_api_key_auth.py
+++ b/litellm/proxy/auth/user_api_key_auth.py
@@ -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 ##
diff --git a/litellm/repositories/verification_token_repository.py b/litellm/repositories/verification_token_repository.py
index d02c2114136..b20fb47306e 100644
--- a/litellm/repositories/verification_token_repository.py
+++ b/litellm/repositories/verification_token_repository.py
@@ -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})
diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py
index 62153e38a83..37f0bf00da6 100644
--- a/tests/e2e/conftest.py
+++ b/tests/e2e/conftest.py
@@ -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",
diff --git a/tests/e2e/coverage_registry/other.yaml b/tests/e2e/coverage_registry/other.yaml
index 0b9249d7420..39747607531 100644
--- a/tests/e2e/coverage_registry/other.yaml
+++ b/tests/e2e/coverage_registry/other.yaml
@@ -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"}
diff --git a/tests/e2e/e2e_config.py b/tests/e2e/e2e_config.py
index 3fa9f534ffd..e88bfad8388 100644
--- a/tests/e2e/e2e_config.py
+++ b/tests/e2e/e2e_config.py
@@ -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.
diff --git a/tests/e2e/mcp/oauth_gateway.py b/tests/e2e/mcp/oauth_gateway.py
index b328c81687b..029b0135900 100644
--- a/tests/e2e/mcp/oauth_gateway.py
+++ b/tests/e2e/mcp/oauth_gateway.py
@@ -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
diff --git a/tests/e2e/models.py b/tests/e2e/models.py
index 83d8a9884a3..e027c410e44 100644
--- a/tests/e2e/models.py
+++ b/tests/e2e/models.py
@@ -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] = []
diff --git a/tests/e2e/other/other_client.py b/tests/e2e/other/other_client.py
index 93c198586f6..d7bad4f1ed1 100644
--- a/tests/e2e/other/other_client.py
+++ b/tests/e2e/other/other_client.py
@@ -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(
diff --git a/tests/e2e/other/owned_jwt_gateway.py b/tests/e2e/other/owned_jwt_gateway.py
new file mode 100644
index 00000000000..1af348cac60
--- /dev/null
+++ b/tests/e2e/other/owned_jwt_gateway.py
@@ -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
diff --git a/tests/e2e/other/test_jwt_auto_register_e2e.py b/tests/e2e/other/test_jwt_auto_register_e2e.py
new file mode 100644
index 00000000000..8f7c2c6a693
--- /dev/null
+++ b/tests/e2e/other/test_jwt_auto_register_e2e.py
@@ -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}"
+ )
diff --git a/tests/e2e/pytest.ini b/tests/e2e/pytest.ini
index e795ebe5721..dbd3ff47daa 100644
--- a/tests/e2e/pytest.ini
+++ b/tests/e2e/pytest.ini
@@ -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__ci_config.yml against that live secret manager; deselected unless E2E_SECRET_MANAGER names the backend (see secret_manager/secret_backends.py)
diff --git a/tests/integration/_support/database_relay.py b/tests/integration/_support/database_relay.py
index 1b3bc0183a5..cb8b9098a64 100644
--- a/tests/integration/_support/database_relay.py
+++ b/tests/integration/_support/database_relay.py
@@ -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()
diff --git a/tests/integration/authorization/test_jwt_auto_register_map_existing_key.py b/tests/integration/authorization/test_jwt_auto_register_map_existing_key.py
new file mode 100644
index 00000000000..5215054a364
--- /dev/null
+++ b/tests/integration/authorization/test_jwt_auto_register_map_existing_key.py
@@ -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
diff --git a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py
index 781d0a13bfd..a7219ac059b 100644
--- a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py
+++ b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py
@@ -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:
From b1e0e9e84b6710f45d746111b47fd21e04bada06 Mon Sep 17 00:00:00 2001
From: tin-berri
Date: Fri, 2 Oct 2026 12:17:55 -0700
Subject: [PATCH 6/8] feat(ui): select Laya for OSS classification (#43768)
---
.../AutoRouters/autoRouterRows.test.ts | 3 +-
.../components/AutoRouters/autoRouterRows.ts | 3 +-
.../add_model/AutoRouterAvailability.tsx | 2 +-
...oRouterClassifierTabs.integration.test.tsx | 2 +-
.../add_model/AutoRouterClassifierTabs.tsx | 30 ++++-
.../add_model/ClassificationMethodConfig.tsx | 4 +-
.../add_model/ClassifierTypeRadios.tsx | 4 +-
.../add_model/ComplexityRouterConfig.tsx | 2 +-
.../JevClassifierConfig.integration.test.tsx | 111 ++++++++++--------
.../add_model/JevClassifierConfig.tsx | 35 ++++--
.../JevConnectionTest.integration.test.tsx | 8 +-
.../add_auto_router_tab.integration.test.tsx | 39 +++---
.../add_model/auto_router_connection_test.tsx | 10 +-
...d_auto_router_routing_test_request.test.ts | 14 +--
.../build_auto_router_routing_test_request.ts | 15 ++-
.../build_complexity_router_config.test.ts | 22 ++--
.../build_complexity_router_config.ts | 16 ++-
.../add_model/jev_classifier_config.ts | 34 +++++-
...d_updated_complexity_router_config.test.ts | 41 ++++---
.../edit_auto_router_modal.tsx | 1 +
.../hydrate_complexity_router_config.ts | 12 +-
.../LogDetailsDrawer/RoutingDecisionCard.tsx | 2 +-
.../src/lib/autorouter_presets.test.ts | 13 +-
.../src/lib/autorouter_presets.ts | 7 +-
24 files changed, 280 insertions(+), 150 deletions(-)
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.test.ts
index 79c4243271e..153b77666a7 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.test.ts
+++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.test.ts
@@ -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(
{
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.ts b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.ts
index 1faf3408c23..3c3366a767d 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.ts
+++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.ts
@@ -57,7 +57,8 @@ const dedupe = (models: string[]): string[] => Array.from(new Set(models));
const COMPLEXITY_TYPE_LABELS: Record = {
llm: "LLM Classifier",
- jev: "JEV Classifier",
+ jev: "OSS Classifier",
+ oss_classifier: "OSS Classifier",
capability: "Capability",
llm_v2: "Fuse v2",
heuristic_first: "Heuristic first",
diff --git a/ui/litellm-dashboard/src/components/add_model/AutoRouterAvailability.tsx b/ui/litellm-dashboard/src/components/add_model/AutoRouterAvailability.tsx
index b1233ef03b8..135ebb958db 100644
--- a/ui/litellm-dashboard/src/components/add_model/AutoRouterAvailability.tsx
+++ b/ui/litellm-dashboard/src/components/add_model/AutoRouterAvailability.tsx
@@ -136,7 +136,7 @@ export const AutoRouterLimits = () => {
Routing and customization limits
- 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.
diff --git a/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.integration.test.tsx b/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.integration.test.tsx
index f843a472d15..08f329e4ad7 100644
--- a/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.integration.test.tsx
+++ b/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.integration.test.tsx
@@ -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}$`) }));
diff --git a/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.tsx b/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.tsx
index 92fc8d2a335..3a1e0065530 100644
--- a/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.tsx
+++ b/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.tsx
@@ -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 = ({ 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> = { capability: "Capability", llm_v2: "Fuse v2" };
const approachDescription: Partial> = {
capability: "Use the efficient model when it is likely to succeed",
@@ -164,7 +173,7 @@ const AutoRouterClassifierTabs: React.FC = ({ 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) => (
- 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"}
- 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