mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge remote-tracking branch 'origin/main' into litellm_vector_store_deny_by_default
This commit is contained in:
commit
b94047f2ac
116 changed files with 5724 additions and 907 deletions
1
.github/workflows/test-e2e-changed.yml
vendored
1
.github/workflows/test-e2e-changed.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -0,0 +1,17 @@
|
|||
CREATE TABLE IF NOT EXISTS "LiteLLM_AutoRouterDailySpend" (
|
||||
"date" TEXT NOT NULL,
|
||||
"api_key" TEXT NOT NULL,
|
||||
"user_id" TEXT NOT NULL,
|
||||
"router_name" TEXT NOT NULL,
|
||||
"router_type" TEXT NOT NULL,
|
||||
"turns" INTEGER NOT NULL DEFAULT 0,
|
||||
"spend" DOUBLE PRECISION NOT NULL DEFAULT 0,
|
||||
"saved_spend" DOUBLE PRECISION NOT NULL DEFAULT 0,
|
||||
"savings_estimated_turns" INTEGER NOT NULL DEFAULT 0,
|
||||
"savings_estimated_actual_spend" DOUBLE PRECISION NOT NULL DEFAULT 0,
|
||||
"savings_estimated_saved_spend" DOUBLE PRECISION NOT NULL DEFAULT 0,
|
||||
"classifier_cost" DOUBLE PRECISION NOT NULL DEFAULT 0,
|
||||
"classifier_cost_recorded_turns" INTEGER NOT NULL DEFAULT 0,
|
||||
|
||||
CONSTRAINT "LiteLLM_AutoRouterDailySpend_pkey" PRIMARY KEY ("date", "api_key", "user_id", "router_name", "router_type")
|
||||
);
|
||||
|
|
@ -1744,6 +1744,27 @@ model LiteLLM_AutoRouterUserSession {
|
|||
@@index([user_id, last_turn_at], map: "idx_autorouter_user_session_user_last_turn")
|
||||
}
|
||||
|
||||
// Auto-routed requests per UTC request day and router: the selected-day money behind the
|
||||
// auto-router usage view. Written in the same statement as the session rollup, so a day row
|
||||
// and its session row never disagree; corrected in the same transaction as late baselines.
|
||||
model LiteLLM_AutoRouterDailySpend {
|
||||
date String
|
||||
api_key String
|
||||
user_id String
|
||||
router_name String
|
||||
router_type String
|
||||
turns Int @default(0)
|
||||
spend Float @default(0)
|
||||
saved_spend Float @default(0)
|
||||
savings_estimated_turns Int @default(0)
|
||||
savings_estimated_actual_spend Float @default(0)
|
||||
savings_estimated_saved_spend Float @default(0)
|
||||
classifier_cost Float @default(0)
|
||||
classifier_cost_recorded_turns Int @default(0)
|
||||
|
||||
@@id([date, api_key, user_id, router_name, router_type])
|
||||
}
|
||||
|
||||
// Shadow eval: evaluation of an auto-router against one or more keys' live traffic, in
|
||||
// either direction. forward duplicates the requests the keys did not route through the
|
||||
// router through it, answering whether they should adopt it; reverse duplicates the
|
||||
|
|
|
|||
|
|
@ -78,6 +78,23 @@ class _InvalidIndex:
|
|||
table_size: str
|
||||
|
||||
MAX_MIGRATE_DEPLOY_ATTEMPTS = 4
|
||||
LIBPQ_URL_PARAMS: Final = frozenset(
|
||||
{
|
||||
"sslmode",
|
||||
"sslcert",
|
||||
"sslkey",
|
||||
"sslrootcert",
|
||||
"sslpassword",
|
||||
"application_name",
|
||||
"connect_timeout",
|
||||
"client_encoding",
|
||||
"options",
|
||||
"service",
|
||||
"gssencmode",
|
||||
"krbsrvname",
|
||||
"target_session_attrs",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
|
|
@ -689,30 +706,43 @@ class ProxyExtrasDBManager:
|
|||
|
||||
@staticmethod
|
||||
def _strip_prisma_query_params(url: str) -> str:
|
||||
"""Remove Prisma-specific query params (connection_limit, pool_timeout,
|
||||
schema, etc.) from DATABASE_URL so psycopg can parse it."""
|
||||
"""Rewrite a Prisma-dialect URL for libpq: drop the Prisma-only params
|
||||
(connection_limit, pool_timeout, schema, pgbouncer, sslaccept, ...) and
|
||||
translate Prisma's TLS params back, since libpq reads ``sslcert`` as a
|
||||
client certificate where Prisma reads it as the CA."""
|
||||
from urllib.parse import parse_qsl, quote, urlencode, urlparse, urlunparse
|
||||
|
||||
parsed = urlparse(url)
|
||||
parsed: Final = urlparse(url)
|
||||
if not parsed.query:
|
||||
return url
|
||||
libpq_params = {
|
||||
"sslmode",
|
||||
"sslcert",
|
||||
"sslkey",
|
||||
"sslrootcert",
|
||||
"sslpassword",
|
||||
"application_name",
|
||||
"connect_timeout",
|
||||
"client_encoding",
|
||||
"options",
|
||||
"service",
|
||||
"gssencmode",
|
||||
"krbsrvname",
|
||||
"target_session_attrs",
|
||||
}
|
||||
kept = [(k, v) for k, v in parse_qsl(parsed.query) if k in libpq_params]
|
||||
return urlunparse(parsed._replace(query=urlencode(kept, quote_via=quote)))
|
||||
pairs: Final = tuple(parse_qsl(parsed.query))
|
||||
kept: Final = tuple((k, v) for k, v in pairs if k in LIBPQ_URL_PARAMS)
|
||||
sslaccept: Final = next((v for k, v in pairs if k == "sslaccept"), None)
|
||||
libpq_pairs: Final = ProxyExtrasDBManager._libpq_tls_params(kept, sslaccept)
|
||||
return urlunparse(parsed._replace(query=urlencode(libpq_pairs, quote_via=quote)))
|
||||
|
||||
@staticmethod
|
||||
def _libpq_tls_params(
|
||||
pairs: "tuple[tuple[str, str], ...]", sslaccept: "str | None"
|
||||
) -> "tuple[tuple[str, str], ...]":
|
||||
"""Undo ``translate_libpq_ssl_params``. Prisma's ``sslcert`` is the CA and
|
||||
``sslaccept=strict`` checks chain and hostname, which libpq only does in
|
||||
``sslmode=verify-full``, so strict becomes ``sslrootcert`` plus
|
||||
``verify-full`` whatever ``sslmode`` said (``disable`` stays off). Prisma
|
||||
defaults an absent ``sslaccept`` to ``accept_invalid_certs`` and anything
|
||||
else to strict. Without strict it checks nothing, so the CA is dropped and
|
||||
``sslmode`` is kept as is: libpq only verifies when a root cert is present.
|
||||
A URL that also carries ``sslkey`` is libpq's own client-certificate form
|
||||
and is kept."""
|
||||
keys: Final = frozenset(k for k, _ in pairs)
|
||||
if "sslcert" not in keys or "sslkey" in keys:
|
||||
return pairs
|
||||
sslmode: Final = next((v for k, v in pairs if k == "sslmode"), None)
|
||||
rest: Final = tuple((k, v) for k, v in pairs if k not in ("sslcert", "sslmode"))
|
||||
if sslaccept in (None, "accept_invalid_certs") or sslmode == "disable":
|
||||
return rest if sslmode is None else rest + (("sslmode", sslmode),)
|
||||
root_cert: Final = tuple(("sslrootcert", v) for k, v in pairs if k == "sslcert" and "sslrootcert" not in keys)
|
||||
return rest + root_cert + (("sslmode", "verify-full"),)
|
||||
|
||||
@staticmethod
|
||||
def _warn_if_db_ahead_of_head(migrations_dir: str) -> None:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
use std::collections::BTreeMap;
|
||||
|
||||
use crate::DecodeError;
|
||||
use serde::Serialize;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
|
|
@ -148,6 +148,60 @@ fn usage_tokens(attributes: &BTreeMap<String, String>) -> Result<(u32, u32), Dec
|
|||
))
|
||||
}
|
||||
|
||||
#[derive(Default, Deserialize)]
|
||||
struct AgentMetadata {
|
||||
#[serde(default)]
|
||||
lc_agent_name: String,
|
||||
#[serde(default)]
|
||||
ls_integration: String,
|
||||
}
|
||||
|
||||
fn recorded_agent_name(
|
||||
name: &str,
|
||||
attributes: &BTreeMap<String, String>,
|
||||
span: &NormalizedSpan,
|
||||
) -> String {
|
||||
let explicit = [
|
||||
span.agent_name.as_str(),
|
||||
attr(attributes, "gen_ai.agent.name"),
|
||||
attr(attributes, "agent.name"),
|
||||
attr(attributes, "openclaw.agent"),
|
||||
]
|
||||
.into_iter()
|
||||
.find(|value| !value.is_empty());
|
||||
if let Some(value) = explicit {
|
||||
return value.to_owned();
|
||||
}
|
||||
let metadata =
|
||||
serde_json::from_str::<AgentMetadata>(attr(attributes, "metadata")).unwrap_or_default();
|
||||
if !metadata.lc_agent_name.is_empty() {
|
||||
return metadata.lc_agent_name;
|
||||
}
|
||||
if span.observation_type == ObservationType::Agent {
|
||||
let node = attr(attributes, "graph.node.id");
|
||||
if !node.is_empty() {
|
||||
return node.to_owned();
|
||||
}
|
||||
if metadata.ls_integration == "langgraph" && name != "LangGraph" && !is_middleware(name) {
|
||||
return name.to_owned();
|
||||
}
|
||||
}
|
||||
String::new()
|
||||
}
|
||||
|
||||
fn is_middleware(name: &str) -> bool {
|
||||
[
|
||||
".wrap_model_call",
|
||||
".wrap_tool_call",
|
||||
".before_agent",
|
||||
".after_agent",
|
||||
".before_model",
|
||||
".after_model",
|
||||
]
|
||||
.iter()
|
||||
.any(|suffix| name.ends_with(suffix))
|
||||
}
|
||||
|
||||
pub fn normalize(
|
||||
scope_name: &str,
|
||||
name: &str,
|
||||
|
|
@ -163,8 +217,22 @@ pub fn normalize(
|
|||
.into_iter()
|
||||
.find(|normalizer| normalizer.matches(scope_name, attributes))
|
||||
.expect("GenAI fallback always matches");
|
||||
let span = normalizer.normalize(name, parent_span_id, attributes)?;
|
||||
let agent_name = recorded_agent_name(name, attributes, &span);
|
||||
let observation_type = if !parent_span_id.is_empty()
|
||||
&& scope_name == "openinference.instrumentation.langchain"
|
||||
&& is_middleware(name)
|
||||
{
|
||||
ObservationType::Framework
|
||||
} else {
|
||||
span.observation_type
|
||||
};
|
||||
Ok(Normalization {
|
||||
span: normalizer.normalize(name, parent_span_id, attributes)?,
|
||||
span: NormalizedSpan {
|
||||
agent_name,
|
||||
observation_type,
|
||||
..span
|
||||
},
|
||||
consumed_attributes: normalizer.consumed_attributes(attributes),
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -362,6 +362,143 @@ async fn keyed_rollup_keeps_same_trace_ids_separate_by_api_key(
|
|||
Ok(())
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn listed_agent_names_preserve_scope_and_cursor(
|
||||
#[future(awt)] database: TestResult<ClickHouseDatabase>,
|
||||
) -> TestResult {
|
||||
let database = database?;
|
||||
let writer = Connection::writer(&database.url)?;
|
||||
ensure_schema(&database.client, &writer, "trace_test", 7).await?;
|
||||
let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64;
|
||||
for (team, key, trace, agent, span, parent) in [
|
||||
("alpha", "one", "shared", "research_agent", "root", ""),
|
||||
("alpha", "one", "shared", "reviewer", "child", "root"),
|
||||
("alpha", "one", "shared", "reviewer", "repeated", "root"),
|
||||
("alpha", "one", "shared", "", "unnamed", "root"),
|
||||
("alpha", "one", "second", "support_agent", "root", ""),
|
||||
("alpha", "two", "shared", "private_agent", "root", ""),
|
||||
("beta", "one", "shared", "other_agent", "root", ""),
|
||||
] {
|
||||
insert_rows(
|
||||
&database,
|
||||
"otel_traces",
|
||||
vec![serde_json::from_value(serde_json::json!({
|
||||
"Timestamp": timestamp, "TraceId": trace, "SpanId": span, "ParentSpanId": parent,
|
||||
"ServiceName": "shared-app", "SpanName": span, "AgentName": agent,
|
||||
"ObservationType": "agent",
|
||||
"ResourceAttributes": {"litellm.team_id": team, "litellm.api_key_hash": key}
|
||||
}))?],
|
||||
)
|
||||
.await?;
|
||||
}
|
||||
let historical_rows = (0..5000)
|
||||
.map(|index| {
|
||||
serde_json::from_value(serde_json::json!({
|
||||
"Timestamp": timestamp - 86_400_000_000_000_i64,
|
||||
"TraceId": "shared", "SpanId": format!("historical-{index}"),
|
||||
"ParentSpanId": "", "SpanName": "historical", "AgentName": "private_agent",
|
||||
"ObservationType": "agent", "ServiceName": "shared-app",
|
||||
"ResourceAttributes": {"litellm.team_id": "alpha", "litellm.api_key_hash": "history"}
|
||||
}))
|
||||
})
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
insert_rows(&database, "otel_traces", historical_rows).await?;
|
||||
let connection = Connection::configured(&database.url, "trace_test", "default", "")?;
|
||||
let parameters = BTreeMap::from([
|
||||
("team_ids".into(), Parameter::Strings(vec!["alpha".into()])),
|
||||
("api_key_hash".into(), Parameter::Text("one".into())),
|
||||
(
|
||||
"start_ms".into(),
|
||||
Parameter::Integer(timestamp / 1_000_000 - 1000),
|
||||
),
|
||||
(
|
||||
"end_ms".into(),
|
||||
Parameter::Integer(timestamp / 1_000_000 + 1000),
|
||||
),
|
||||
("cursor_ms".into(), Parameter::Integer(0)),
|
||||
("cursor_trace_id".into(), Parameter::Text(String::new())),
|
||||
("limit".into(), Parameter::Integer(1)),
|
||||
]);
|
||||
let first: serde_json::Value = serde_json::from_str(
|
||||
&execute_named_read(
|
||||
&database.client,
|
||||
&connection,
|
||||
ReadQuery::ListTraces,
|
||||
¶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::<BTreeMap<_, _>>();
|
||||
assert_eq!(
|
||||
names["shared"],
|
||||
serde_json::json!(["research_agent", "reviewer"])
|
||||
);
|
||||
assert_eq!(names["second"], serde_json::json!(["support_agent"]));
|
||||
let counts = [&first["data"][0], &second["data"][0]]
|
||||
.into_iter()
|
||||
.map(|row| {
|
||||
(
|
||||
row["trace_id"].as_str().unwrap(),
|
||||
row["agent_count"].as_u64(),
|
||||
)
|
||||
})
|
||||
.collect::<BTreeMap<_, _>>();
|
||||
assert_eq!(counts["shared"], Some(3));
|
||||
assert_eq!(counts["second"], Some(1));
|
||||
for page in [&first, &second] {
|
||||
assert!(
|
||||
page["statistics"]["rows_read"]
|
||||
.as_u64()
|
||||
.ok_or("missing read statistics")?
|
||||
< 5000
|
||||
);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn rollup_merges_spans_across_days_without_losing_root_fields(
|
||||
|
|
@ -376,6 +513,7 @@ async fn rollup_merges_spans_across_days_without_losing_root_fields(
|
|||
let root = serde_json::from_value(serde_json::json!({
|
||||
"Timestamp": day_start - 1_000_000_000, "TraceId": "cross-day", "SpanId": "span-root",
|
||||
"ParentSpanId": "", "ServiceName": "proxy", "SpanName": "root", "Input": "root input",
|
||||
"AgentName": "lead", "ObservationType": "agent",
|
||||
"StatusCode": "STATUS_CODE_ERROR",
|
||||
"ResourceAttributes": {"litellm.team_id": "team-1"}
|
||||
}))?;
|
||||
|
|
@ -383,6 +521,7 @@ async fn rollup_merges_spans_across_days_without_losing_root_fields(
|
|||
let child = serde_json::from_value(serde_json::json!({
|
||||
"Timestamp": day_start + 1_000_000_000, "TraceId": "cross-day", "SpanId": "span-child",
|
||||
"ParentSpanId": "span-root", "ServiceName": "proxy", "SpanName": "child",
|
||||
"AgentName": "researcher", "ObservationType": "agent",
|
||||
"StatusCode": "STATUS_CODE_UNSET",
|
||||
"ResourceAttributes": {"litellm.team_id": "team-1"}
|
||||
}))?;
|
||||
|
|
@ -406,6 +545,33 @@ async fn rollup_merges_spans_across_days_without_losing_root_fields(
|
|||
"RootStatus": "STATUS_CODE_ERROR", "SpanCount": 2
|
||||
}])
|
||||
);
|
||||
let connection = Connection::configured(&database.url, "trace_test", "default", "")?;
|
||||
let parameters = BTreeMap::from([
|
||||
("team_ids".into(), Parameter::Strings(vec!["team-1".into()])),
|
||||
("api_key_hash".into(), Parameter::Text(String::new())),
|
||||
(
|
||||
"start_ms".into(),
|
||||
Parameter::Integer(day_start / 1_000_000 - 2000),
|
||||
),
|
||||
("end_ms".into(), Parameter::Integer(day_start / 1_000_000)),
|
||||
("cursor_ms".into(), Parameter::Integer(0)),
|
||||
("cursor_trace_id".into(), Parameter::Text(String::new())),
|
||||
("limit".into(), Parameter::Integer(10)),
|
||||
]);
|
||||
let listed: serde_json::Value = serde_json::from_str(
|
||||
&execute_named_read(
|
||||
&database.client,
|
||||
&connection,
|
||||
ReadQuery::ListTraces,
|
||||
¶meters,
|
||||
)
|
||||
.await?,
|
||||
)?;
|
||||
assert_eq!(
|
||||
listed["data"][0]["agent_names"],
|
||||
serde_json::json!(["lead", "researcher"])
|
||||
);
|
||||
assert_eq!(listed["data"][0]["agent_count"], 2);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -915,7 +915,8 @@ def _excluded_db_systems(logger: "OpenTelemetryV2") -> frozenset[str]:
|
|||
logger got published: with ``callbacks: [langfuse_otel, otel]`` the ``otel``
|
||||
callback folds into the preset, whose config is env-only.
|
||||
"""
|
||||
configured: Final = litellm.callback_settings.get("otel", {}).get("excluded_services")
|
||||
otel_settings: Final = (litellm.callback_settings or {}).get("otel")
|
||||
configured: Final = otel_settings.get("excluded_services") if isinstance(otel_settings, dict) else None
|
||||
if configured is None:
|
||||
return logger.config.excluded_services
|
||||
return excluded_db_systems_from(configured)
|
||||
|
|
|
|||
|
|
@ -5,7 +5,8 @@ from functools import lru_cache
|
|||
from typing import Annotated, Any, Final
|
||||
|
||||
from pydantic import AliasChoices, BaseModel, Field, TypeAdapter, ValidationError, field_validator, model_validator
|
||||
from pydantic_settings import BaseSettings, NoDecode, SettingsConfigDict
|
||||
from pydantic.fields import FieldInfo
|
||||
from pydantic_settings import BaseSettings, NoDecode, PydanticBaseSettingsSource, SettingsConfigDict
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.otel.model.baggage import (
|
||||
|
|
@ -121,9 +122,37 @@ class ExporterSpec(BaseModel):
|
|||
)
|
||||
|
||||
|
||||
class _EnvWithoutBareExcludedServices(PydanticBaseSettingsSource):
|
||||
def __init__(self, settings_cls: type[BaseSettings], env_settings: PydanticBaseSettingsSource) -> None:
|
||||
super().__init__(settings_cls)
|
||||
self._env_settings: Final = env_settings
|
||||
|
||||
def get_field_value(self, field: FieldInfo, field_name: str) -> tuple[object, str, bool]:
|
||||
return self._env_settings.get_field_value(field, field_name)
|
||||
|
||||
def __call__(self) -> dict[str, object]:
|
||||
return {key: value for key, value in self._env_settings().items() if key != "excluded_services"}
|
||||
|
||||
|
||||
class OpenTelemetryV2Config(BaseSettings):
|
||||
model_config = SettingsConfigDict(populate_by_name=True, extra="ignore")
|
||||
|
||||
@classmethod
|
||||
def settings_customise_sources(
|
||||
cls,
|
||||
settings_cls: type[BaseSettings],
|
||||
init_settings: PydanticBaseSettingsSource,
|
||||
env_settings: PydanticBaseSettingsSource,
|
||||
dotenv_settings: PydanticBaseSettingsSource,
|
||||
file_secret_settings: PydanticBaseSettingsSource,
|
||||
) -> tuple[PydanticBaseSettingsSource, ...]:
|
||||
return (
|
||||
init_settings,
|
||||
_EnvWithoutBareExcludedServices(settings_cls, env_settings),
|
||||
dotenv_settings,
|
||||
file_secret_settings,
|
||||
)
|
||||
|
||||
# ----- single-destination shorthand, read from standard OTEL_* envs ----- #
|
||||
exporter: str = Field(
|
||||
default="console",
|
||||
|
|
@ -178,7 +207,7 @@ class OpenTelemetryV2Config(BaseSettings):
|
|||
)
|
||||
excluded_services: Annotated[frozenset[str], NoDecode] = Field(
|
||||
default_factory=frozenset,
|
||||
validation_alias=AliasChoices("excluded_services", "LITELLM_OTEL_EXCLUDED_SERVICES"),
|
||||
validation_alias=AliasChoices("LITELLM_OTEL_EXCLUDED_SERVICES"),
|
||||
description=(
|
||||
"Datastore services whose spans are withheld from key/team ``callback_vars`` "
|
||||
"OTel destinations (the operator's own exporters still receive them). Accepted "
|
||||
|
|
|
|||
1
litellm/llms/laya/__init__.py
Normal file
1
litellm/llms/laya/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
|
||||
60
litellm/llms/laya/common_utils.py
Normal file
60
litellm/llms/laya/common_utils.py
Normal file
|
|
@ -0,0 +1,60 @@
|
|||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Final, Literal, TypeAlias
|
||||
|
||||
from pydantic import AnyHttpUrl, BaseModel, TypeAdapter, ValidationError
|
||||
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
LayaCheckpoint: TypeAlias = Literal["english", "multilingual", "typed-decisions"]
|
||||
|
||||
|
||||
def validate_laya_model(value: object) -> LayaCheckpoint:
|
||||
try:
|
||||
return TypeAdapter(LayaCheckpoint).validate_python(value)
|
||||
except ValidationError as exc:
|
||||
raise ValueError("Laya model must be 'english', 'multilingual', or 'typed-decisions'") from exc
|
||||
|
||||
|
||||
def validate_laya_request(body: Mapping[str, object]) -> LayaCheckpoint:
|
||||
if "custom_body" in body:
|
||||
raise ValueError("custom_body is not supported for Laya requests")
|
||||
if body.get("stream"):
|
||||
raise ValueError("Streaming is not supported for Laya requests")
|
||||
return validate_laya_model(body.get("model"))
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class LayaConnection:
|
||||
api_base: str
|
||||
api_key: str | None = field(repr=False)
|
||||
|
||||
|
||||
def validate_laya_api_base(value: str) -> str:
|
||||
try:
|
||||
url: Final = TypeAdapter(AnyHttpUrl).validate_python(value)
|
||||
except ValidationError as exc:
|
||||
raise ValueError("Laya api_base must be an HTTP or HTTPS server URL") from exc
|
||||
if url.username or url.password or url.query or url.fragment:
|
||||
raise ValueError("Laya api_base must not contain credentials, a query, or a fragment")
|
||||
return str(url).rstrip("/")
|
||||
|
||||
|
||||
def laya_connection(api_base: str | None = None, api_key: str | None = None) -> LayaConnection:
|
||||
base: Final = api_base if api_base is not None else get_secret_str("LAYA_API_BASE")
|
||||
if not base:
|
||||
raise ValueError("Laya requires api_base or LAYA_API_BASE pointing to a self-hosted server")
|
||||
key: Final = api_key if api_base is not None else api_key or get_secret_str("LAYA_API_KEY")
|
||||
return LayaConnection(api_base=validate_laya_api_base(base), api_key=key)
|
||||
|
||||
|
||||
class _LayaRouting(BaseModel):
|
||||
model: str | None = None
|
||||
|
||||
|
||||
def laya_response_model(response: Mapping[str, object], requested_model: str | None) -> str:
|
||||
try:
|
||||
routing: Final = TypeAdapter(_LayaRouting).validate_python(response.get("routing") or _LayaRouting())
|
||||
except ValidationError:
|
||||
return requested_model or "unknown"
|
||||
return routing.model or requested_model or "unknown"
|
||||
|
|
@ -72622,6 +72622,45 @@
|
|||
"supports_audio_input": true,
|
||||
"supports_video_input": true
|
||||
},
|
||||
"laya/english": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "laya",
|
||||
"mode": "evaluation",
|
||||
"output_cost_per_token": 0.0,
|
||||
"source": "https://github.com/NandhaKishorM/laya",
|
||||
"supported_endpoints": [
|
||||
"/v1/systemone"
|
||||
],
|
||||
"metadata": {
|
||||
"notes": "Self-hosted decision model; infrastructure costs are paid separately"
|
||||
}
|
||||
},
|
||||
"laya/multilingual": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "laya",
|
||||
"mode": "evaluation",
|
||||
"output_cost_per_token": 0.0,
|
||||
"source": "https://github.com/NandhaKishorM/laya",
|
||||
"supported_endpoints": [
|
||||
"/v1/systemone"
|
||||
],
|
||||
"metadata": {
|
||||
"notes": "Self-hosted decision model; infrastructure costs are paid separately"
|
||||
}
|
||||
},
|
||||
"laya/typed-decisions": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "laya",
|
||||
"mode": "evaluation",
|
||||
"output_cost_per_token": 0.0,
|
||||
"source": "https://github.com/NandhaKishorM/laya",
|
||||
"supported_endpoints": [
|
||||
"/v1/systemone"
|
||||
],
|
||||
"metadata": {
|
||||
"notes": "Self-hosted decision model; infrastructure costs are paid separately"
|
||||
}
|
||||
},
|
||||
"typesafe/jev-1.13.0": {
|
||||
"input_cost_per_token": 4.2e-08,
|
||||
"litellm_provider": "typesafe",
|
||||
|
|
|
|||
|
|
@ -229,6 +229,7 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = (
|
|||
"/tinyfish/",
|
||||
"/transcribe",
|
||||
"/typesafe/",
|
||||
"/laya/",
|
||||
"/openrouter/",
|
||||
"/vertex-ai/",
|
||||
"/vertex_ai/",
|
||||
|
|
|
|||
|
|
@ -27761,6 +27761,30 @@
|
|||
]
|
||||
}
|
||||
},
|
||||
"/laya/v1/systemone": {
|
||||
"post": {
|
||||
"operationId": "laya_proxy_route_laya_v1_systemone_post",
|
||||
"responses": {
|
||||
"200": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {}
|
||||
}
|
||||
},
|
||||
"description": "Successful Response"
|
||||
}
|
||||
},
|
||||
"security": [
|
||||
{
|
||||
"APIKeyHeader": []
|
||||
}
|
||||
],
|
||||
"summary": "Laya Proxy Route",
|
||||
"tags": [
|
||||
"llm_passthrough"
|
||||
]
|
||||
}
|
||||
},
|
||||
"/milvus/{endpoint}": {
|
||||
"delete": {
|
||||
"description": "Enable using Milvus `/vectors` endpoint as a pass-through endpoint.",
|
||||
|
|
|
|||
|
|
@ -507,6 +507,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/vllm",
|
||||
"/mistral",
|
||||
"/typesafe",
|
||||
"/laya",
|
||||
"/openrouter",
|
||||
"/milvus",
|
||||
"/gigachat",
|
||||
|
|
@ -5474,6 +5475,17 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
|
|||
"'auto_register': auto-create a virtual key and mapping on first encounter."
|
||||
),
|
||||
)
|
||||
auto_register_map_existing_key: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"Only used with unregistered_jwt_client_behavior='auto_register'. When True and the virtual key claim "
|
||||
"field is the user_id_jwt_field or user_email_jwt_field, the JWT claim is mapped to a virtual key the "
|
||||
"JWT-resolved user already owns instead of minting a new one. If the user owns several, the most recently created key in the "
|
||||
"JWT-resolved team (or with no team when the JWT resolves none) is chosen among keys that never "
|
||||
"expire, are not blocked, are not Admin UI session keys, were not minted by auto_register, and "
|
||||
"have no allowed_routes or include llm_api_routes. Otherwise a new key is minted as usual."
|
||||
),
|
||||
)
|
||||
routing_overrides: list[JWTRoutingOverride] | None = Field(
|
||||
default=None,
|
||||
description="Optional claim-based routing overrides for JWT-shaped tokens. Matching rules route requests to oauth2 before default JWT flow.",
|
||||
|
|
@ -5574,6 +5586,15 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
|
|||
return issuer_config.virtual_key_claim_field
|
||||
return self.virtual_key_claim_field
|
||||
|
||||
def is_user_identity_claim(self, claim_field: str, issuer: str | None) -> bool:
|
||||
issuer_config: Final = self.get_issuer_config(issuer)
|
||||
if issuer_config is None:
|
||||
return claim_field in (self.user_id_jwt_field, self.user_email_jwt_field)
|
||||
return claim_field in (
|
||||
issuer_config.user_id_jwt_field or self.user_id_jwt_field,
|
||||
issuer_config.user_email_jwt_field or self.user_email_jwt_field,
|
||||
)
|
||||
|
||||
def get_unregistered_jwt_client_behavior(self, issuer: str | None) -> UnregisteredJWTClientBehavior:
|
||||
issuer_config: Final = self.get_issuer_config(issuer)
|
||||
if issuer_config is not None and issuer_config.unregistered_jwt_client_behavior is not None:
|
||||
|
|
|
|||
|
|
@ -1883,6 +1883,15 @@ def _extract_model_candidates_from_request(
|
|||
llm_router: Router | None = None,
|
||||
team_id: str | None = None,
|
||||
) -> list[str]:
|
||||
if route.rstrip("/") == "/laya/v1/systemone":
|
||||
from litellm.llms.laya.common_utils import validate_laya_model
|
||||
|
||||
try:
|
||||
laya_request: Final = TypeAdapter(Mapping[str, object]).validate_python(request_data)
|
||||
laya_model: Final = validate_laya_model(laya_request.get("model"))
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
return _dedupe_model_candidates((f"laya/{laya_model}",))
|
||||
if route == "/cost/predict-cache":
|
||||
prediction_models: Final = _cache_prediction_model_candidates(request_data, llm_router, team_id) # pyright: ignore[reportUnknownArgumentType] # the typed reader validates each deployment ID from this legacy payload
|
||||
return _dedupe_model_candidates(prediction_models)
|
||||
|
|
|
|||
|
|
@ -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 ##
|
||||
|
|
|
|||
|
|
@ -4,11 +4,12 @@ Per-session auto-router benchmarks rollup.
|
|||
At request time the spend writer builds one AutoRouterTurnTransaction per successful
|
||||
auto-routed request (a request whose metadata carries a routing_decision) and queues it
|
||||
on the prisma client. The spend-log flush job drains the queue into
|
||||
key and user session rollups with one atomic statement per turn: each upsert classifies
|
||||
key and user session rollups, plus the per-day router rollup, with one atomic statement
|
||||
per turn: each upsert classifies
|
||||
the turn (same model, first visit, return to a model the session already used, out of
|
||||
order) against the row's own columns, so nothing is read before the write and concurrent
|
||||
pods compose. The benchmarks endpoint aggregates these rows and can recover matching historical
|
||||
costs from retained spend logs when estimate coverage predates these columns.
|
||||
pods compose. The benchmarks endpoint reads session shape from the session rows and money from the
|
||||
day rows, so spend and savings count only requests on the selected UTC days.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
|
@ -71,45 +72,82 @@ tier_maps AS (
|
|||
GROUP BY router_name, router_type, kv.key
|
||||
) per_tier
|
||||
GROUP BY router_name, router_type
|
||||
),
|
||||
sessions AS (
|
||||
SELECT
|
||||
router_name,
|
||||
router_type,
|
||||
COUNT(*)::int AS sessions,
|
||||
SUM(turns)::int AS session_turns,
|
||||
SUM(unordered_turns)::int AS unordered_turns,
|
||||
SUM(covered_turns)::int AS covered_turns,
|
||||
SUM(cache_hits)::int AS cache_hits,
|
||||
SUM(same_model_turns)::int AS same_model_turns,
|
||||
SUM(same_model_hits)::int AS same_model_hits,
|
||||
SUM(first_visit_turns)::int AS first_visit_turns,
|
||||
SUM(first_visit_hits)::int AS first_visit_hits,
|
||||
SUM(return_turns)::int AS return_turns,
|
||||
SUM(return_hits)::int AS return_hits,
|
||||
SUM(return_expired_misses)::int AS return_expired_misses,
|
||||
SUM(return_within_ttl_misses)::int AS return_within_ttl_misses,
|
||||
SUM(ttl_5m_turns)::int AS ttl_5m_turns,
|
||||
SUM(ttl_1h_turns)::int AS ttl_1h_turns,
|
||||
SUM(total_tokens)::bigint AS total_tokens,
|
||||
SUM(EXTRACT(EPOCH FROM (last_turn_at - first_turn_at)))::float8 AS session_seconds
|
||||
FROM windowed
|
||||
GROUP BY router_name, router_type
|
||||
),
|
||||
days AS (
|
||||
SELECT
|
||||
router_name,
|
||||
router_type,
|
||||
SUM(turns)::int AS turns,
|
||||
SUM(spend)::float8 AS spend,
|
||||
SUM(saved_spend)::float8 AS saved_spend,
|
||||
SUM(savings_estimated_turns)::int AS savings_estimated_turns,
|
||||
SUM(savings_estimated_actual_spend)::float8 AS savings_estimated_actual_spend,
|
||||
SUM(savings_estimated_saved_spend)::float8 AS savings_estimated_saved_spend,
|
||||
SUM(classifier_cost)::float8 AS classifier_cost,
|
||||
SUM(classifier_cost_recorded_turns)::int AS classifier_cost_recorded_turns
|
||||
FROM "LiteLLM_AutoRouterDailySpend"
|
||||
WHERE date >= $5 AND date <= $6
|
||||
AND ($3::text IS NULL OR api_key = $3::text)
|
||||
AND ($4::text IS NULL OR user_id = $4::text)
|
||||
GROUP BY router_name, router_type
|
||||
)
|
||||
SELECT
|
||||
agg.*,
|
||||
COALESCE(tier_maps.tier_turns, '{{}}'::jsonb) AS tier_turns
|
||||
FROM (
|
||||
SELECT
|
||||
router_name,
|
||||
router_type,
|
||||
COUNT(*)::int AS sessions,
|
||||
COALESCE(SUM(turns), 0)::int AS turns,
|
||||
COALESCE(SUM(unordered_turns), 0)::int AS unordered_turns,
|
||||
COALESCE(SUM(covered_turns), 0)::int AS covered_turns,
|
||||
COALESCE(SUM(cache_hits), 0)::int AS cache_hits,
|
||||
COALESCE(SUM(same_model_turns), 0)::int AS same_model_turns,
|
||||
COALESCE(SUM(same_model_hits), 0)::int AS same_model_hits,
|
||||
COALESCE(SUM(first_visit_turns), 0)::int AS first_visit_turns,
|
||||
COALESCE(SUM(first_visit_hits), 0)::int AS first_visit_hits,
|
||||
COALESCE(SUM(return_turns), 0)::int AS return_turns,
|
||||
COALESCE(SUM(return_hits), 0)::int AS return_hits,
|
||||
COALESCE(SUM(return_expired_misses), 0)::int AS return_expired_misses,
|
||||
COALESCE(SUM(return_within_ttl_misses), 0)::int AS return_within_ttl_misses,
|
||||
COALESCE(SUM(ttl_5m_turns), 0)::int AS ttl_5m_turns,
|
||||
COALESCE(SUM(ttl_1h_turns), 0)::int AS ttl_1h_turns,
|
||||
COALESCE(SUM(total_tokens), 0)::bigint AS total_tokens,
|
||||
COALESCE(SUM(spend), 0)::float8 AS spend,
|
||||
COALESCE(SUM(saved_spend), 0)::float8 AS saved_spend,
|
||||
COALESCE(SUM(savings_estimated_turns), 0)::int AS savings_estimated_turns,
|
||||
COALESCE(SUM(savings_estimated_actual_spend), 0)::float8 AS savings_estimated_actual_spend,
|
||||
CASE WHEN BOOL_AND(savings_estimated_turns = turns AND classifier_cost_recorded_turns = turns)
|
||||
THEN SUM(classifier_cost)::float8 END AS savings_estimated_classifier_cost,
|
||||
COALESCE(SUM(savings_estimated_saved_spend), 0)::float8 AS savings_estimated_saved_spend,
|
||||
COALESCE(SUM(classifier_cost), 0)::float8 AS classifier_cost,
|
||||
COALESCE(SUM(classifier_cost_recorded_turns), 0)::int AS classifier_cost_recorded_turns,
|
||||
COALESCE(SUM(EXTRACT(EPOCH FROM (last_turn_at - first_turn_at))), 0)::float8 AS session_seconds
|
||||
FROM windowed
|
||||
GROUP BY router_name, router_type
|
||||
) agg
|
||||
COALESCE(tier_maps.tier_turns, '{{}}'::jsonb) AS tier_turns,
|
||||
COALESCE(sessions.sessions, 0) AS sessions,
|
||||
COALESCE(sessions.session_turns, 0) AS session_turns,
|
||||
COALESCE(sessions.unordered_turns, 0) AS unordered_turns,
|
||||
COALESCE(sessions.covered_turns, 0) AS covered_turns,
|
||||
COALESCE(sessions.cache_hits, 0) AS cache_hits,
|
||||
COALESCE(sessions.same_model_turns, 0) AS same_model_turns,
|
||||
COALESCE(sessions.same_model_hits, 0) AS same_model_hits,
|
||||
COALESCE(sessions.first_visit_turns, 0) AS first_visit_turns,
|
||||
COALESCE(sessions.first_visit_hits, 0) AS first_visit_hits,
|
||||
COALESCE(sessions.return_turns, 0) AS return_turns,
|
||||
COALESCE(sessions.return_hits, 0) AS return_hits,
|
||||
COALESCE(sessions.return_expired_misses, 0) AS return_expired_misses,
|
||||
COALESCE(sessions.return_within_ttl_misses, 0) AS return_within_ttl_misses,
|
||||
COALESCE(sessions.ttl_5m_turns, 0) AS ttl_5m_turns,
|
||||
COALESCE(sessions.ttl_1h_turns, 0) AS ttl_1h_turns,
|
||||
COALESCE(sessions.total_tokens, 0) AS total_tokens,
|
||||
COALESCE(sessions.session_seconds, 0) AS session_seconds,
|
||||
COALESCE(days.turns, 0) AS turns,
|
||||
COALESCE(days.spend, 0) AS spend,
|
||||
COALESCE(days.saved_spend, 0) AS saved_spend,
|
||||
COALESCE(days.savings_estimated_turns, 0) AS savings_estimated_turns,
|
||||
COALESCE(days.savings_estimated_actual_spend, 0) AS savings_estimated_actual_spend,
|
||||
COALESCE(days.savings_estimated_saved_spend, 0) AS savings_estimated_saved_spend,
|
||||
COALESCE(days.classifier_cost, 0) AS classifier_cost,
|
||||
COALESCE(days.classifier_cost_recorded_turns, 0) AS classifier_cost_recorded_turns
|
||||
FROM sessions
|
||||
FULL OUTER JOIN days USING (router_name, router_type)
|
||||
LEFT JOIN tier_maps USING (router_name, router_type)
|
||||
ORDER BY agg.spend DESC
|
||||
ORDER BY spend DESC, router_name, router_type
|
||||
"""
|
||||
|
||||
|
||||
|
|
@ -391,15 +429,43 @@ ON CONFLICT ({user_column}api_key, session_id, router_name) DO UPDATE SET
|
|||
"""
|
||||
|
||||
|
||||
_DAY_UPSERT_SQL: Final = f"""
|
||||
day_rollup AS (
|
||||
INSERT INTO "LiteLLM_AutoRouterDailySpend" AS d (
|
||||
date, api_key, user_id, router_name, router_type, turns, spend, saved_spend, savings_estimated_turns,
|
||||
savings_estimated_actual_spend, savings_estimated_saved_spend, classifier_cost, classifier_cost_recorded_turns
|
||||
)
|
||||
VALUES (
|
||||
({_TURN_AT}::timestamp)::date::text, {_p("api_key")}::text, {_p("user_id")}::text, {_p("router_name")},
|
||||
{_p("router_type")}, 1, {_p("spend")}::float8, {_p("saved_spend")}::float8, {_p("savings_estimated_turns")}::int,
|
||||
{_p("savings_estimated_actual_spend")}::float8, {_p("savings_estimated_saved_spend")}::float8,
|
||||
{_p("classifier_cost")}::float8, 1
|
||||
)
|
||||
ON CONFLICT (date, api_key, user_id, router_name, router_type) DO UPDATE SET
|
||||
turns = d.turns + 1,
|
||||
spend = d.spend + EXCLUDED.spend,
|
||||
saved_spend = d.saved_spend + EXCLUDED.saved_spend,
|
||||
savings_estimated_turns = d.savings_estimated_turns + EXCLUDED.savings_estimated_turns,
|
||||
savings_estimated_actual_spend = d.savings_estimated_actual_spend + EXCLUDED.savings_estimated_actual_spend,
|
||||
savings_estimated_saved_spend = d.savings_estimated_saved_spend + EXCLUDED.savings_estimated_saved_spend,
|
||||
classifier_cost = d.classifier_cost + EXCLUDED.classifier_cost,
|
||||
classifier_cost_recorded_turns = d.classifier_cost_recorded_turns + 1
|
||||
RETURNING 1
|
||||
)
|
||||
"""
|
||||
|
||||
UPSERT_AUTOROUTER_SESSION_SQL: Final = f"""
|
||||
WITH key_rollup AS (
|
||||
{_session_upsert_sql(user_scoped=False)}
|
||||
RETURNING 1
|
||||
)
|
||||
), {_DAY_UPSERT_SQL}
|
||||
{_session_upsert_sql(user_scoped=True)}
|
||||
"""
|
||||
|
||||
UPSERT_AUTOROUTER_USER_SESSION_SQL: Final = _session_upsert_sql(user_scoped=True)
|
||||
UPSERT_AUTOROUTER_USER_SESSION_SQL: Final = f"""
|
||||
WITH {_DAY_UPSERT_SQL}
|
||||
{_session_upsert_sql(user_scoped=True)}
|
||||
"""
|
||||
|
||||
|
||||
def _as_sql_param(value: str | float | bool | datetime | None) -> str | float | None:
|
||||
|
|
|
|||
|
|
@ -179,6 +179,8 @@ class _Change(BaseModel):
|
|||
actual_delta: float
|
||||
savings_delta: float
|
||||
daily: DailyBaselineAttribution | None
|
||||
date: str | None = None
|
||||
router_type: str | None = None
|
||||
|
||||
|
||||
class _TransactionManager(Protocol):
|
||||
|
|
@ -303,6 +305,26 @@ WHERE {user_match}session.api_key = totals.api_key AND session.session_id = tota
|
|||
|
||||
_UPDATE_SESSIONS: Final = _session_correction_sql(user_scoped=False)
|
||||
_UPDATE_USER_SESSIONS: Final = _session_correction_sql(user_scoped=True)
|
||||
_UPDATE_DAYS: Final = """
|
||||
WITH totals AS (
|
||||
SELECT date, api_key, user_id, router_name, router_type, SUM(covered_delta)::int AS covered_delta,
|
||||
SUM(actual_delta) AS actual_delta, SUM(savings_delta) AS savings_delta
|
||||
FROM jsonb_to_recordset($1::jsonb) AS x(
|
||||
date text, api_key text, user_id text, router_name text, router_type text,
|
||||
covered_delta int, actual_delta float8, savings_delta float8
|
||||
)
|
||||
WHERE date IS NOT NULL
|
||||
GROUP BY date, api_key, user_id, router_name, router_type
|
||||
)
|
||||
UPDATE "LiteLLM_AutoRouterDailySpend" AS day
|
||||
SET saved_spend = day.saved_spend + totals.savings_delta,
|
||||
savings_estimated_turns = day.savings_estimated_turns + totals.covered_delta,
|
||||
savings_estimated_actual_spend = day.savings_estimated_actual_spend + totals.actual_delta,
|
||||
savings_estimated_saved_spend = day.savings_estimated_saved_spend + totals.savings_delta
|
||||
FROM totals
|
||||
WHERE day.date = totals.date AND day.api_key = totals.api_key AND day.user_id = totals.user_id
|
||||
AND day.router_name = totals.router_name AND day.router_type = totals.router_type
|
||||
"""
|
||||
|
||||
|
||||
def _primary_transaction(client: PrismaClient) -> _TransactionManager:
|
||||
|
|
@ -331,6 +353,8 @@ def _change(record: BaselineAccountingRecord, old: BaselinePublication | None, n
|
|||
savings_delta=(current.savings if current is not None else 0.0)
|
||||
- (previous.savings if previous is not None else 0.0),
|
||||
daily=record.daily,
|
||||
date=record.turn.turn_at.date().isoformat() if record.turn is not None else None,
|
||||
router_type=record.turn.router_type if record.turn is not None else None,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -373,6 +397,7 @@ async def _publish(db: SupportsRawQueries, changes: Sequence[_Change]) -> None:
|
|||
await db.execute_raw(_UPDATE_SESSIONS, serialized)
|
||||
if any(change.user_id for change in changes):
|
||||
await db.execute_raw(_UPDATE_USER_SESSIONS, serialized)
|
||||
await db.execute_raw(_UPDATE_DAYS, serialized)
|
||||
for entity, table in DAILY_SPEND_TABLES.items():
|
||||
if adjustments := tuple(
|
||||
change.daily.adjustment(target, change.savings_delta, change.request_id)
|
||||
|
|
|
|||
|
|
@ -540,6 +540,18 @@ class SpendLogCleanup:
|
|||
deadline=deadline,
|
||||
)
|
||||
|
||||
async def _delete_old_autorouter_daily_rows(
|
||||
self, prisma_client: PrismaClient, cutoff_day: str, deadline: float
|
||||
) -> TableCleanupResult:
|
||||
return await self._delete_old_rows_batched(
|
||||
prisma_client,
|
||||
cutoff_day,
|
||||
table_name="LiteLLM_AutoRouterDailySpend",
|
||||
key_columns=("date", "api_key", "user_id", "router_name", "router_type"),
|
||||
time_column="date",
|
||||
deadline=deadline,
|
||||
)
|
||||
|
||||
async def _delete_old_health_check_rows(
|
||||
self, prisma_client: PrismaClient, cutoff_date: datetime, deadline: float
|
||||
) -> TableCleanupResult:
|
||||
|
|
@ -623,16 +635,20 @@ class SpendLogCleanup:
|
|||
except Exception: # noqa: BLE001 # retained observations are retried by the next cleanup job
|
||||
verbose_proxy_logger.warning("Auto-router baseline retention remains pending")
|
||||
sessions_result: Final = await self._delete_old_autorouter_session_rows(
|
||||
prisma_client, session_cutoff, self._group_deadline(deadline, 2)
|
||||
prisma_client, session_cutoff, self._group_deadline(deadline, 3)
|
||||
)
|
||||
verbose_proxy_logger.info("Deleted %s expired auto-router session rollup rows", sessions_result.rows_deleted)
|
||||
user_sessions_result: Final = await self._delete_old_autorouter_user_session_rows(
|
||||
prisma_client, session_cutoff, deadline
|
||||
prisma_client, session_cutoff, self._group_deadline(deadline, 2)
|
||||
)
|
||||
verbose_proxy_logger.info(
|
||||
"Deleted %s expired auto-router user session rollup rows", user_sessions_result.rows_deleted
|
||||
)
|
||||
return (sessions_result, user_sessions_result)
|
||||
days_result: Final = await self._delete_old_autorouter_daily_rows(
|
||||
prisma_client, session_cutoff.date().isoformat(), deadline
|
||||
)
|
||||
verbose_proxy_logger.info("Deleted %s expired auto-router daily rollup rows", days_result.rows_deleted)
|
||||
return (sessions_result, user_sessions_result, days_result)
|
||||
|
||||
async def _clean_health_checks(
|
||||
self, prisma_client: PrismaClient, retention_seconds: int, deadline: float
|
||||
|
|
|
|||
|
|
@ -1760,7 +1760,7 @@ class LiteLLMProxyRequestSetup:
|
|||
): # don't override k-v pair sent by request (user request)
|
||||
data[_metadata_variable_name]["spend_logs_metadata"][key] = value
|
||||
else:
|
||||
data[_metadata_variable_name]["spend_logs_metadata"] = key_metadata["spend_logs_metadata"]
|
||||
data[_metadata_variable_name]["spend_logs_metadata"] = dict(key_metadata["spend_logs_metadata"])
|
||||
|
||||
## KEY-LEVEL DISABLE FALLBACKS
|
||||
if "disable_fallbacks" in key_metadata and isinstance(key_metadata["disable_fallbacks"], bool):
|
||||
|
|
@ -1777,6 +1777,53 @@ class LiteLLMProxyRequestSetup:
|
|||
)
|
||||
return data
|
||||
|
||||
@staticmethod
|
||||
def add_team_and_project_level_controls(
|
||||
user_api_key_dict: UserAPIKeyAuth, metadata: dict[str, object]
|
||||
) -> dict[str, object]:
|
||||
team_metadata: Final = user_api_key_dict.team_metadata or MappingProxyType({})
|
||||
project_metadata: Final = user_api_key_dict.project_metadata or MappingProxyType({})
|
||||
request_tags: Final = metadata.get("tags")
|
||||
team_tags: Final = team_metadata.get("tags")
|
||||
project_tags: Final = project_metadata.get("tags")
|
||||
disable_global_guardrails: Final = team_metadata.get("disable_global_guardrails")
|
||||
opted_out_global_guardrails: Final = team_metadata.get("opted_out_global_guardrails")
|
||||
spend_logs_metadata: Final = LiteLLMProxyRequestSetup._merge_spend_logs_metadata(
|
||||
team_spend_logs_metadata=team_metadata.get("spend_logs_metadata"),
|
||||
request_spend_logs_metadata=metadata.get("spend_logs_metadata"),
|
||||
)
|
||||
tags: Final = LiteLLMProxyRequestSetup._merge_tags(
|
||||
request_tags=LiteLLMProxyRequestSetup._merge_tags(
|
||||
request_tags=request_tags if isinstance(request_tags, list) else None,
|
||||
tags_to_add=team_tags if isinstance(team_tags, list) else None,
|
||||
),
|
||||
tags_to_add=project_tags if isinstance(project_tags, list) else None,
|
||||
)
|
||||
controls: Final = (
|
||||
("tags", tags or None),
|
||||
("spend_logs_metadata", spend_logs_metadata),
|
||||
(
|
||||
"disable_global_guardrails",
|
||||
disable_global_guardrails if isinstance(disable_global_guardrails, bool) else None,
|
||||
),
|
||||
(
|
||||
"opted_out_global_guardrails",
|
||||
opted_out_global_guardrails if isinstance(opted_out_global_guardrails, list) else None,
|
||||
),
|
||||
)
|
||||
return {**metadata, **{key: value for key, value in controls if value is not None}}
|
||||
|
||||
@staticmethod
|
||||
def _merge_spend_logs_metadata(
|
||||
team_spend_logs_metadata: object, request_spend_logs_metadata: object
|
||||
) -> dict[str, object] | None:
|
||||
"""Team values as defaults, the request's own values win on the same key. None when neither is a dict"""
|
||||
team_values: Final = team_spend_logs_metadata if isinstance(team_spend_logs_metadata, dict) else None
|
||||
request_values: Final = request_spend_logs_metadata if isinstance(request_spend_logs_metadata, dict) else None
|
||||
if team_values is None and request_values is None:
|
||||
return None
|
||||
return {**(team_values or {}), **(request_values or {})}
|
||||
|
||||
@staticmethod
|
||||
def _merge_tags(request_tags: list | None, tags_to_add: list | None) -> list:
|
||||
"""
|
||||
|
|
@ -2312,38 +2359,12 @@ async def add_litellm_data_to_request(
|
|||
data=data,
|
||||
_metadata_variable_name=_metadata_variable_name,
|
||||
)
|
||||
## TEAM-LEVEL SPEND LOGS/TAGS
|
||||
data[_metadata_variable_name] = LiteLLMProxyRequestSetup.add_team_and_project_level_controls(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
metadata=data[_metadata_variable_name],
|
||||
)
|
||||
team_metadata: Final = user_api_key_dict.team_metadata or {}
|
||||
if "tags" in team_metadata and team_metadata["tags"] is not None:
|
||||
data[_metadata_variable_name]["tags"] = LiteLLMProxyRequestSetup._merge_tags(
|
||||
request_tags=data[_metadata_variable_name].get("tags"),
|
||||
tags_to_add=team_metadata["tags"],
|
||||
)
|
||||
if "disable_global_guardrails" in team_metadata and isinstance(team_metadata["disable_global_guardrails"], bool):
|
||||
data[_metadata_variable_name]["disable_global_guardrails"] = team_metadata["disable_global_guardrails"]
|
||||
if "opted_out_global_guardrails" in team_metadata and isinstance(
|
||||
team_metadata["opted_out_global_guardrails"], list
|
||||
):
|
||||
data[_metadata_variable_name]["opted_out_global_guardrails"] = team_metadata["opted_out_global_guardrails"]
|
||||
if "spend_logs_metadata" in team_metadata and isinstance(team_metadata["spend_logs_metadata"], dict):
|
||||
if "spend_logs_metadata" in data[_metadata_variable_name] and isinstance(
|
||||
data[_metadata_variable_name]["spend_logs_metadata"], dict
|
||||
):
|
||||
for key, value in team_metadata["spend_logs_metadata"].items():
|
||||
if (
|
||||
key not in data[_metadata_variable_name]["spend_logs_metadata"]
|
||||
): # don't override k-v pair sent by request (user request)
|
||||
data[_metadata_variable_name]["spend_logs_metadata"][key] = value
|
||||
else:
|
||||
data[_metadata_variable_name]["spend_logs_metadata"] = team_metadata["spend_logs_metadata"]
|
||||
|
||||
## PROJECT-LEVEL TAGS
|
||||
project_metadata: Final = user_api_key_dict.project_metadata or {}
|
||||
if "tags" in project_metadata and project_metadata["tags"] is not None:
|
||||
data[_metadata_variable_name]["tags"] = LiteLLMProxyRequestSetup._merge_tags(
|
||||
request_tags=data[_metadata_variable_name].get("tags"),
|
||||
tags_to_add=project_metadata["tags"],
|
||||
)
|
||||
|
||||
# inherited_tags: every tag key/team/project policy contributed, read
|
||||
# directly from those three sources rather than snapshotted off the shared
|
||||
|
|
|
|||
|
|
@ -5,6 +5,8 @@ POST /auto_router/test_routing - Route one request through an unsaved complexity
|
|||
POST /auto_router/validate_complexity_router_config - Dry-run the complexity-router write gate without saving
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import math
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from itertools import chain, groupby
|
||||
|
|
@ -40,6 +42,7 @@ from litellm.proxy.litellm_pre_call_utils import (
|
|||
refresh_proxy_server_request_body_snapshot,
|
||||
)
|
||||
from litellm.proxy.management.teams.access import is_team_admin
|
||||
from litellm.proxy.management_endpoints.common_daily_activity import daily_activity_scope
|
||||
from litellm.proxy.management_helpers.auto_router_permissions import (
|
||||
authorize_member_auto_router_dependencies,
|
||||
authorize_member_auto_router_team,
|
||||
|
|
@ -47,6 +50,7 @@ from litellm.proxy.management_helpers.auto_router_permissions import (
|
|||
)
|
||||
from litellm.repositories.autorouter_session_repository import AutoRouterSessionRepository
|
||||
from litellm.repositories.base_repository import SupportsModelDump
|
||||
from litellm.repositories.daily_activity_sql import build_where_clause
|
||||
from litellm.repositories.team_repository import TeamRepository
|
||||
from litellm.router_strategy.complexity_router import ComplexityRouter
|
||||
from litellm.router_utils.auto_router_model_naming import (
|
||||
|
|
@ -316,7 +320,7 @@ async def _authorize_models_this_test_can_call(
|
|||
its calls through the proxy. Team and member budgets are already enforced on every route.
|
||||
"""
|
||||
models: Final = _models_this_test_can_call(config)
|
||||
if not models and config.classifier_type != "jev":
|
||||
if not models and config.classifier_type != "oss_classifier":
|
||||
return
|
||||
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
|
@ -342,9 +346,9 @@ async def _authorize_models_this_test_can_call(
|
|||
code=status.HTTP_400_BAD_REQUEST,
|
||||
) from e
|
||||
|
||||
if config.classifier_type == "jev" and user_api_key_dict.budget_throttle_pct is not None:
|
||||
if config.classifier_type == "oss_classifier" and user_api_key_dict.budget_throttle_pct is not None:
|
||||
raise ProxyException(
|
||||
message="Budget has been exceeded! JEV Test Routing requires available budget.",
|
||||
message="Budget has been exceeded! OSS Classifier Test Routing requires available budget.",
|
||||
type=ProxyErrorTypes.budget_exceeded,
|
||||
param=None,
|
||||
code=status.HTTP_400_BAD_REQUEST,
|
||||
|
|
@ -616,34 +620,37 @@ async def preview_auto_router_routing(
|
|||
|
||||
|
||||
class _SessionAggRow(BaseModel):
|
||||
"""One router's window: session shape from overlapping sessions, money from the selected days."""
|
||||
|
||||
router_name: str
|
||||
router_type: str
|
||||
tier_turns: Mapping[str, int]
|
||||
sessions: int
|
||||
turns: int
|
||||
unordered_turns: int
|
||||
covered_turns: int
|
||||
cache_hits: int
|
||||
same_model_turns: int
|
||||
same_model_hits: int
|
||||
first_visit_turns: int
|
||||
first_visit_hits: int
|
||||
return_turns: int
|
||||
return_hits: int
|
||||
return_expired_misses: int
|
||||
return_within_ttl_misses: int
|
||||
ttl_5m_turns: int
|
||||
ttl_1h_turns: int
|
||||
total_tokens: int
|
||||
spend: float
|
||||
saved_spend: float
|
||||
tier_turns: Mapping[str, int] = MappingProxyType({})
|
||||
sessions: int = 0
|
||||
session_turns: int = 0
|
||||
unordered_turns: int = 0
|
||||
covered_turns: int = 0
|
||||
cache_hits: int = 0
|
||||
same_model_turns: int = 0
|
||||
same_model_hits: int = 0
|
||||
first_visit_turns: int = 0
|
||||
first_visit_hits: int = 0
|
||||
return_turns: int = 0
|
||||
return_hits: int = 0
|
||||
return_expired_misses: int = 0
|
||||
return_within_ttl_misses: int = 0
|
||||
ttl_5m_turns: int = 0
|
||||
ttl_1h_turns: int = 0
|
||||
total_tokens: int = 0
|
||||
session_seconds: float = 0.0
|
||||
turns: int = 0
|
||||
spend: float = 0.0
|
||||
saved_spend: float = 0.0
|
||||
savings_estimated_turns: int = 0
|
||||
savings_estimated_actual_spend: float = 0.0
|
||||
savings_estimated_classifier_cost: float | None = None
|
||||
savings_estimated_saved_spend: float = 0.0
|
||||
classifier_cost: float
|
||||
classifier_cost_recorded_turns: int
|
||||
session_seconds: float
|
||||
classifier_cost: float = 0.0
|
||||
classifier_cost_recorded_turns: int = 0
|
||||
|
||||
|
||||
_SESSION_AGG_ROWS: Final = TypeAdapter(list[_SessionAggRow])
|
||||
|
|
@ -692,6 +699,13 @@ def _compared_row(row: _SessionAggRow) -> _SessionAggRow:
|
|||
)
|
||||
|
||||
|
||||
def _per_session(row: _SessionAggRow, total: float) -> float | None:
|
||||
"""Unknown, not zero, when routed requests have no session rows of their own to average over."""
|
||||
if row.sessions:
|
||||
return total / row.sessions
|
||||
return None if row.turns else 0.0
|
||||
|
||||
|
||||
def _benchmark_totals(row: _SessionAggRow) -> AutoRouterBenchmarkTotals:
|
||||
return_misses: Final = row.return_turns - row.return_hits
|
||||
saved_spend, baseline_spend = _savings_cohort(
|
||||
|
|
@ -701,9 +715,9 @@ def _benchmark_totals(row: _SessionAggRow) -> AutoRouterBenchmarkTotals:
|
|||
return AutoRouterBenchmarkTotals(
|
||||
sessions=sessions,
|
||||
turns=row.turns,
|
||||
avg_turns_per_session=row.turns / sessions if sessions else 0.0,
|
||||
avg_session_seconds=row.session_seconds / sessions if sessions else 0.0,
|
||||
avg_tokens_per_session=row.total_tokens / sessions if sessions else 0.0,
|
||||
avg_turns_per_session=_per_session(row, row.session_turns),
|
||||
avg_session_seconds=_per_session(row, row.session_seconds),
|
||||
avg_tokens_per_session=_per_session(row, row.total_tokens),
|
||||
spend=row.spend,
|
||||
savings_estimated_turns=row.savings_estimated_turns,
|
||||
savings_estimated_actual_spend=row.savings_estimated_actual_spend,
|
||||
|
|
@ -712,9 +726,8 @@ def _benchmark_totals(row: _SessionAggRow) -> AutoRouterBenchmarkTotals:
|
|||
classifier_cost=row.classifier_cost if row.classifier_cost_recorded_turns == row.turns else None,
|
||||
baseline_spend=baseline_spend,
|
||||
saved_pct=_pct(saved_spend, baseline_spend) if saved_spend is not None and baseline_spend is not None else None,
|
||||
saved_per_session=(saved_spend / sessions if sessions else 0.0) if saved_spend is not None else None,
|
||||
cache=AutoRouterCacheStats(
|
||||
coverage_pct=_pct(row.covered_turns, row.turns),
|
||||
coverage_pct=_pct(row.covered_turns, row.session_turns),
|
||||
hit_rate_pct=_pct(row.cache_hits, row.covered_turns),
|
||||
same_model=_cache_bucket(row.same_model_turns, row.same_model_hits),
|
||||
first_visit=_cache_bucket(row.first_visit_turns, row.first_visit_hits),
|
||||
|
|
@ -748,7 +761,6 @@ def _benchmark_group(row: _SessionAggRow) -> AutoRouterBenchmarkGroup:
|
|||
classifier_cost=totals.classifier_cost,
|
||||
baseline_spend=totals.baseline_spend,
|
||||
saved_pct=totals.saved_pct,
|
||||
saved_per_session=totals.saved_per_session,
|
||||
cache=totals.cache,
|
||||
)
|
||||
|
||||
|
|
@ -759,6 +771,7 @@ def _summed_agg_row(rows: Sequence[_SessionAggRow]) -> _SessionAggRow:
|
|||
router_type="",
|
||||
tier_turns=MappingProxyType({}),
|
||||
sessions=sum(row.sessions for row in rows),
|
||||
session_turns=sum(row.session_turns for row in rows),
|
||||
turns=sum(row.turns for row in rows),
|
||||
unordered_turns=sum(row.unordered_turns for row in rows),
|
||||
covered_turns=sum(row.covered_turns for row in rows),
|
||||
|
|
@ -790,6 +803,49 @@ def _summed_agg_row(rows: Sequence[_SessionAggRow]) -> _SessionAggRow:
|
|||
)
|
||||
|
||||
|
||||
async def _recorded_autorouter_savings(
|
||||
prisma_client: "PrismaClient", start_day: str, end_day: str, api_key: str | None, user_id: str | None
|
||||
) -> float:
|
||||
"""The selected days' auto-router savings exactly as the Overall view sums them: same table, same filters."""
|
||||
where, params = build_where_clause(
|
||||
daily_activity_scope(
|
||||
table="litellm_dailyuserspend",
|
||||
entity_id_field="user_id",
|
||||
entity_id=user_id,
|
||||
exclude_entity_ids=None,
|
||||
api_key=api_key,
|
||||
start_date=start_day,
|
||||
end_date=end_day,
|
||||
model=None,
|
||||
timezone_offset_minutes=None,
|
||||
)
|
||||
)
|
||||
rows: Final = await _query_raw(
|
||||
prisma_client,
|
||||
f'SELECT COALESCE(SUM(autorouter_savings_spend), 0)::float8 AS saved FROM "LiteLLM_DailyUserSpend" WHERE {where}',
|
||||
*params,
|
||||
)
|
||||
return float(rows[0]["saved"]) if rows else 0.0
|
||||
|
||||
|
||||
def _with_recorded_savings(
|
||||
totals: AutoRouterBenchmarkTotals, rows: Sequence[_SessionAggRow], recorded: float
|
||||
) -> AutoRouterBenchmarkTotals:
|
||||
"""The headline is the recorded total. Savings outside the compared routers void the cost comparison,
|
||||
and the part no router's day rows account for is reported as unattributed."""
|
||||
if math.isclose(recorded, totals.saved_spend or 0.0, abs_tol=1e-9):
|
||||
return totals
|
||||
unattributed: Final = recorded - sum(row.saved_spend for row in rows)
|
||||
return totals.model_copy(
|
||||
update={
|
||||
"saved_spend": recorded,
|
||||
"unattributed_saved_spend": None if math.isclose(unattributed, 0.0, abs_tol=1e-9) else unattributed,
|
||||
"baseline_spend": None,
|
||||
"saved_pct": None,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _strategy_router_key(deployment: object) -> tuple[str, str] | None:
|
||||
"""``(model_name, kind)`` for a deployment whose routing the session rollup records.
|
||||
|
||||
|
|
@ -849,7 +905,7 @@ async def get_auto_router_benchmarks(
|
|||
str | None, Query(description="YYYY-MM-DD UTC, inclusive (defaults to 30 days before end_date)")
|
||||
] = None,
|
||||
end_date: Annotated[str | None, Query(description="YYYY-MM-DD UTC, inclusive (defaults to today)")] = None,
|
||||
api_key: Annotated[str | None, Query(description="Filter to one virtual key token hash")] = None,
|
||||
api_key: Annotated[str | None, Query(min_length=1, description="Filter to one virtual key token hash")] = None,
|
||||
user_id: Annotated[
|
||||
str | None, Query(min_length=1, description="Filter to one canonical internal user recorded on each turn")
|
||||
] = None,
|
||||
|
|
@ -860,9 +916,10 @@ async def get_auto_router_benchmarks(
|
|||
|
||||
Reads session rollups folded once per request at spend-write time, so this endpoint
|
||||
never scans LiteLLM_SpendLogs. A user filter selects only turns attributed to that
|
||||
internal user when written; older key-only history remains outside user views. A session
|
||||
is in the window when it overlaps it: its last turn is on or after start_date and its first turn is on or before
|
||||
end_date. Overall hit rate is over telemetry-bearing turns; each bucket's hit rate is
|
||||
internal user when written; older key-only history remains outside user views. Money counts
|
||||
only requests on the selected UTC days, and the all-router savings headline is the same daily
|
||||
total the Overall view reads. Session shape and caching cover every session that overlaps the
|
||||
window, whole. Overall hit rate is over telemetry-bearing turns; each bucket's hit rate is
|
||||
over that bucket's turns.
|
||||
|
||||
The rollup supplies the measures, never the list. Which routers appear comes from the
|
||||
|
|
@ -885,24 +942,35 @@ async def get_auto_router_benchmarks(
|
|||
if end_day < start_day:
|
||||
raise HTTPException(status_code=400, detail="end_date must not be earlier than start_date")
|
||||
|
||||
raw_rows: Final = await _query_raw(
|
||||
prisma_client,
|
||||
AUTOROUTER_BENCHMARKS_SQL,
|
||||
start_day.isoformat(),
|
||||
(end_day + timedelta(days=1)).isoformat(),
|
||||
api_key,
|
||||
user_id,
|
||||
first_day: Final = start_day.strftime("%Y-%m-%d")
|
||||
last_day: Final = end_day.strftime("%Y-%m-%d")
|
||||
raw_rows, recorded = await asyncio.gather(
|
||||
_query_raw(
|
||||
prisma_client,
|
||||
AUTOROUTER_BENCHMARKS_SQL,
|
||||
start_day.isoformat(),
|
||||
(end_day + timedelta(days=1)).isoformat(),
|
||||
api_key,
|
||||
user_id,
|
||||
first_day,
|
||||
last_day,
|
||||
),
|
||||
_recorded_autorouter_savings(prisma_client, first_day, last_day, api_key, user_id),
|
||||
)
|
||||
rows: Final = tuple(_compared_row(row) for row in _SESSION_AGG_ROWS.validate_python(raw_rows or ()))
|
||||
totals: Final = _with_recorded_savings(_benchmark_totals(_summed_agg_row(rows)), rows, recorded)
|
||||
unattributed: Final = MappingProxyType(
|
||||
{"baseline_spend": None, "saved_pct": None} if totals.unattributed_saved_spend is not None else {}
|
||||
)
|
||||
groups: Final = (
|
||||
*(_benchmark_group(row) for row in rows),
|
||||
*(_benchmark_group(row).model_copy(update=unattributed) for row in rows),
|
||||
*_idle_router_groups(llm_router, frozenset((row.router_name, row.router_type) for row in rows)),
|
||||
)
|
||||
return AutoRouterBenchmarksResponse(
|
||||
start_date=start_day.strftime("%Y-%m-%d"),
|
||||
end_date=end_day.strftime("%Y-%m-%d"),
|
||||
start_date=first_day,
|
||||
end_date=last_day,
|
||||
routers_in_scope=len(groups),
|
||||
totals=_benchmark_totals(_summed_agg_row(rows)),
|
||||
totals=totals,
|
||||
groups=groups,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ from litellm.litellm_core_utils.ptu_pricing import (
|
|||
SEARCH_CONTEXT_SIZES,
|
||||
ptu_config_error,
|
||||
)
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import redact_credentials_in_payload
|
||||
from litellm.proxy._types import (
|
||||
BlockModelRequest,
|
||||
CommonProxyErrors,
|
||||
|
|
@ -115,6 +116,7 @@ from litellm.router_strategy.complexity_router import (
|
|||
normalize_classification_examples,
|
||||
normalize_classification_prompt,
|
||||
)
|
||||
from litellm.router_strategy.complexity_router.config import resolve_complexity_router_config_write
|
||||
from litellm.router_utils.auto_router_model_naming import (
|
||||
GATED_AUTO_ROUTER_CAPABILITIES,
|
||||
STRATEGY_ROUTER_PARAM_FIELDS,
|
||||
|
|
@ -187,6 +189,25 @@ class _ProxyModelRow(Protocol):
|
|||
def model_dump_json(self, *, exclude_none: bool = False) -> str: ...
|
||||
|
||||
|
||||
def _model_write_response(
|
||||
row: _ProxyModelRow, member_write: MemberAutoRouterWrite | None
|
||||
) -> _ProxyModelRow | Mapping[str, object]:
|
||||
if member_write is None:
|
||||
return row
|
||||
payload: Final = TypeAdapter(dict[str, object]).validate_json(row.model_dump_json())
|
||||
stored_params: Final = payload.get("litellm_params")
|
||||
params: Final = (
|
||||
TypeAdapter(dict[str, object]).validate_json(stored_params)
|
||||
if isinstance(stored_params, str)
|
||||
else TypeAdapter(dict[str, object]).validate_python(stored_params)
|
||||
)
|
||||
redacted: Final = redact_credentials_in_payload(params)
|
||||
return {
|
||||
**payload,
|
||||
"litellm_params": json.dumps(redacted) if isinstance(stored_params, str) else redacted,
|
||||
}
|
||||
|
||||
|
||||
class _ProxyModelTable(Protocol):
|
||||
def find_unique(self, *, where: Mapping[str, object]) -> Awaitable[BaseModel | None]: ...
|
||||
|
||||
|
|
@ -407,34 +428,13 @@ WHERE model_id <> $1
|
|||
|
||||
def _effective_complexity_router_config(
|
||||
incoming_params: GenericLiteLLMParams | None, existing_params: GenericLiteLLMParams | None
|
||||
) -> object:
|
||||
) -> Mapping[str, object] | None:
|
||||
incoming: Final = None if incoming_params is None else incoming_params.complexity_router_config
|
||||
existing: Final = None if existing_params is None else existing_params.complexity_router_config
|
||||
if incoming is None:
|
||||
return existing
|
||||
if existing is None or incoming.get("classifier_type") != "jev" or existing.get("classifier_type") != "jev":
|
||||
return incoming
|
||||
incoming_jev: Final[object] = incoming.get("jev_classifier_config")
|
||||
existing_jev: Final[object] = existing.get("jev_classifier_config")
|
||||
if not isinstance(incoming_jev, Mapping) or not isinstance(existing_jev, Mapping):
|
||||
return incoming
|
||||
supplied: Final = TypeAdapter(dict[str, object]).validate_python(incoming_jev)
|
||||
stored: Final = TypeAdapter(dict[str, object]).validate_python(existing_jev)
|
||||
same_base: Final = "api_base" not in supplied or supplied["api_base"] == stored.get("api_base")
|
||||
transport: Final = MappingProxyType(
|
||||
{
|
||||
key: value
|
||||
for key, value in stored.items()
|
||||
if key in ("api_key", "api_base") and (key != "api_key" or same_base)
|
||||
}
|
||||
)
|
||||
return {
|
||||
**incoming,
|
||||
"jev_classifier_config": {
|
||||
**transport,
|
||||
**supplied,
|
||||
},
|
||||
}
|
||||
config_adapter: Final[TypeAdapter[Mapping[str, object] | None]] = TypeAdapter(Mapping[str, object] | None)
|
||||
return resolve_complexity_router_config_write(
|
||||
config_adapter.validate_python(incoming), config_adapter.validate_python(existing)
|
||||
).effective
|
||||
|
||||
|
||||
def _effective_model(
|
||||
|
|
@ -1304,7 +1304,7 @@ async def patch_model(
|
|||
live_after=reload_outcome.live_after,
|
||||
)
|
||||
|
||||
return updated_model
|
||||
return _model_write_response(updated_model, member_write)
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error in patch_model: %s", e)
|
||||
|
|
@ -1501,10 +1501,18 @@ async def _add_model_to_db(
|
|||
slot: AbstractAsyncContextManager[_ProxyModelTable] | None = None,
|
||||
) -> "_ProxyModelRow | LiteLLM_ProxyModelTable":
|
||||
# encrypt litellm params #
|
||||
_litellm_params_dict: Final = model_params.litellm_params.dict(exclude_none=True)
|
||||
_litellm_params_dict: Final = TypeAdapter(dict[str, object]).validate_python(
|
||||
model_params.litellm_params.model_dump(exclude_none=True)
|
||||
)
|
||||
if "complexity_router_config" in _litellm_params_dict:
|
||||
_litellm_params_dict["complexity_router_config"] = _effective_complexity_router_config(
|
||||
model_params.litellm_params, None
|
||||
)
|
||||
_original_litellm_model_name: Final = model_params.litellm_params.model
|
||||
for k, v in _litellm_params_dict.items():
|
||||
encrypted_value = encrypt_value_helper(value=v, new_encryption_key=new_encryption_key)
|
||||
encrypted_value = (
|
||||
encrypt_value_helper(value=v, new_encryption_key=new_encryption_key) if isinstance(v, str) else v
|
||||
)
|
||||
model_params.litellm_params[k] = encrypted_value
|
||||
_data: Final[dict] = {
|
||||
"model_id": model_params.model_info.id,
|
||||
|
|
@ -2536,7 +2544,7 @@ async def add_new_model(
|
|||
live_after=reload_outcome.live_after,
|
||||
)
|
||||
|
||||
return model_response
|
||||
return _model_write_response(model_response, member_write)
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("litellm.proxy.proxy_server.add_new_model(): Exception occured - %s", e)
|
||||
|
|
@ -2760,7 +2768,7 @@ async def update_model(
|
|||
live_after=reload_outcome.live_after,
|
||||
)
|
||||
|
||||
return model_response
|
||||
return None if model_response is None else _model_write_response(model_response, member_write)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("litellm.proxy.proxy_server.update_model(): Exception occured - %s", e)
|
||||
if isinstance(e, HTTPException):
|
||||
|
|
|
|||
|
|
@ -33,6 +33,10 @@ from litellm.repositories.prisma_protocols import DatabaseClient
|
|||
from litellm.repositories.project_repository import ProjectRepository
|
||||
from litellm.repositories.table_repositories import TeamMembershipRepository
|
||||
from litellm.router import Router
|
||||
from litellm.router_strategy.complexity_router.config import (
|
||||
ComplexityRouterConfigWrite,
|
||||
resolve_complexity_router_config_write,
|
||||
)
|
||||
from litellm.router_utils.auto_router_model_naming import classify_strategy_router_model, strategy_router_dependencies
|
||||
from litellm.types.management_endpoints.auto_router_endpoints import RequestComplexityRouterConfig
|
||||
from litellm.types.router import Deployment, updateDeployment
|
||||
|
|
@ -65,12 +69,12 @@ class _MemberRouterGenerationParams(BaseModel):
|
|||
stop: str | tuple[str, ...] | None = None
|
||||
|
||||
|
||||
class _MemberJevClassifierConfig(BaseModel):
|
||||
"""The Jev classifier settings a team member may set. Credentials stay the proxy's own: a member-chosen
|
||||
api_base would receive the proxy's TYPESAFE_API_KEY, and a member-chosen api_key would be sent from the proxy."""
|
||||
class _MemberOpenSourceClassifierConfig(BaseModel):
|
||||
"""Classifier settings a team member may set while the gateway owns the connection."""
|
||||
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
provider: Literal["jev", "laya"] = "jev"
|
||||
model: str
|
||||
api_key: None = None
|
||||
api_base: None = None
|
||||
|
|
@ -123,14 +127,21 @@ def authorize_member_auto_router_team(
|
|||
|
||||
|
||||
def validate_member_auto_router_config(config: Mapping[str, object]) -> RequestComplexityRouterConfig:
|
||||
return _validate_member_auto_router_config_write(resolve_complexity_router_config_write(config, None))
|
||||
|
||||
|
||||
def _validate_member_auto_router_config_write(write: ComplexityRouterConfigWrite) -> RequestComplexityRouterConfig:
|
||||
if write.effective is None:
|
||||
raise HTTPException(status_code=400, detail="A complexity_router_config is required.")
|
||||
try:
|
||||
validated: Final = _MemberComplexityRouterConfig.model_validate(config)
|
||||
for entries in validated.tier_model_configs.values():
|
||||
for entry in entries:
|
||||
_MemberRouterGenerationParams.model_validate(entry.litellm_params)
|
||||
if validated.jev_classifier_config is not None:
|
||||
_MemberJevClassifierConfig.model_validate(validated.jev_classifier_config.model_dump())
|
||||
return validated
|
||||
if write.submitted is not None:
|
||||
validated: Final = _MemberComplexityRouterConfig.model_validate(write.submitted)
|
||||
for entries in validated.tier_model_configs.values():
|
||||
for entry in entries:
|
||||
_MemberRouterGenerationParams.model_validate(entry.litellm_params)
|
||||
if validated.opensource_classifier_config is not None:
|
||||
_MemberOpenSourceClassifierConfig.model_validate(validated.opensource_classifier_config.model_dump())
|
||||
return RequestComplexityRouterConfig.model_validate(write.effective)
|
||||
except ValidationError as exc:
|
||||
location: Final = ".".join(str(part) for part in exc.errors()[0]["loc"])
|
||||
raise HTTPException(status_code=400, detail=f"Invalid member auto-router configuration at {location}.") from exc
|
||||
|
|
@ -332,16 +343,15 @@ async def authorize_member_auto_router_write(
|
|||
if existing is not None and incoming.model_name not in (None, public_name, existing.model_name):
|
||||
raise HTTPException(status_code=403, detail="Team members cannot rename an auto router.")
|
||||
supplied_config: Final = _RouterConfigSource.model_validate(params.model_dump()).complexity_router_config
|
||||
raw_config: Final = (
|
||||
supplied_config
|
||||
if supplied_config is not None
|
||||
else _RouterConfigSource.model_validate(existing.litellm_params.model_dump()).complexity_router_config
|
||||
stored_config: Final = (
|
||||
_RouterConfigSource.model_validate(existing.litellm_params.model_dump()).complexity_router_config
|
||||
if existing is not None
|
||||
else None
|
||||
)
|
||||
if raw_config is None:
|
||||
raise HTTPException(status_code=400, detail="A complexity_router_config is required.")
|
||||
config: Final = validate_member_auto_router_config(raw_config)
|
||||
resolved_config: Final = resolve_complexity_router_config_write(supplied_config, stored_config)
|
||||
if resolved_config.supplied_connection_fields:
|
||||
raise HTTPException(status_code=403, detail="Team members cannot change classifier connections.")
|
||||
config: Final = _validate_member_auto_router_config_write(resolved_config)
|
||||
stored_default: Final = existing.litellm_params.complexity_router_default_model if existing is not None else None
|
||||
default_model: Final = (
|
||||
params.complexity_router_default_model
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, cast
|
|||
import httpx
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response, WebSocket
|
||||
from fastapi.responses import StreamingResponse
|
||||
from pydantic import ConfigDict, TypeAdapter
|
||||
from starlette.websockets import WebSocketState
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
|
|
@ -57,6 +58,7 @@ from litellm.llms.deepgram.common_utils import (
|
|||
deepgram_listen_websocket_target,
|
||||
)
|
||||
from litellm.llms.fal_ai.cost_calculator import fal_ai_passthrough_cost, fal_ai_queue_base
|
||||
from litellm.llms.laya.common_utils import laya_connection, validate_laya_request
|
||||
from litellm.llms.nvidia_nim.passthrough.transformation import nvidia_nim_model_group_in_path
|
||||
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
|
||||
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
|
||||
|
|
@ -636,6 +638,47 @@ async def typesafe_proxy_route(
|
|||
return await endpoint_func(request, fastapi_response, user_api_key_dict)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/laya/v1/systemone",
|
||||
tags=["Laya Pass-through", "pass-through"],
|
||||
)
|
||||
async def laya_proxy_route(
|
||||
request: Request,
|
||||
fastapi_response: Response,
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
) -> Response:
|
||||
body: Final = TypeAdapter(dict[str, object]).validate_python(await _read_request_body(request))
|
||||
try:
|
||||
_ = validate_laya_request(body)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
try:
|
||||
connection: Final = laya_connection()
|
||||
except ValueError as exc:
|
||||
raise HTTPException(
|
||||
status_code=503, detail="Laya server is not configured correctly; check LAYA_API_BASE"
|
||||
) from exc
|
||||
base_url: Final = httpx.URL(connection.api_base)
|
||||
updated_url: Final = base_url.copy_with(
|
||||
path=HttpPassThroughEndpointHelpers.join_base_and_endpoint_path(base_url, "/v1/systemone"),
|
||||
)
|
||||
authorization: Final[Mapping[str, str]] = (
|
||||
MappingProxyType({"Authorization": f"Bearer {connection.api_key}"})
|
||||
if connection.api_key
|
||||
else MappingProxyType({})
|
||||
)
|
||||
endpoint_func: Final = create_pass_through_route(
|
||||
endpoint="v1/systemone",
|
||||
target=str(updated_url),
|
||||
custom_headers=MappingProxyType({**authorization, "Content-Type": "application/json"}),
|
||||
custom_llm_provider="laya",
|
||||
is_streaming_request=False,
|
||||
)
|
||||
return TypeAdapter(Response, config=ConfigDict(arbitrary_types_allowed=True)).validate_python(
|
||||
await endpoint_func(request, fastapi_response, user_api_key_dict)
|
||||
)
|
||||
|
||||
|
||||
@router.api_route(
|
||||
"/openrouter/{endpoint:path}",
|
||||
methods=["GET", "POST", "PUT", "DELETE", "PATCH"],
|
||||
|
|
|
|||
|
|
@ -34,9 +34,7 @@ def is_collection_route(url_route: str, collection_suffix: str) -> bool:
|
|||
|
||||
def request_tags_from_metadata(request_metadata: Mapping[str, object]) -> tuple[str, ...] | None:
|
||||
"""Tags for the batch-cost spend row: the request's own tags when it sent any,
|
||||
otherwise the key's tags, which auth exposes as user_api_key_auth_metadata (a
|
||||
tagged key does not put its tags in the top-level metadata "tags" on the
|
||||
passthrough path)
|
||||
otherwise the key's tags, which auth exposes as user_api_key_auth_metadata
|
||||
"""
|
||||
tags: Final = _sanitized_str_tuple(request_metadata.get("tags"))
|
||||
if tags:
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
|||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
get_standard_logging_object_payload, # pyright: ignore[reportUnknownVariableType] # legacy helper has an untyped signature
|
||||
)
|
||||
from litellm.llms.laya.common_utils import laya_response_model
|
||||
from litellm.proxy._types import PassThroughEndpointLoggingTypedDict
|
||||
from litellm.types.utils import ModelResponse, StandardPassThroughResponseObject, Usage
|
||||
|
||||
|
|
@ -69,9 +70,11 @@ class TypeSafePassthroughLoggingHandler:
|
|||
**kwargs: object,
|
||||
) -> PassThroughEndpointLoggingTypedDict:
|
||||
response: Final = _parse_typesafe_response(response_body)
|
||||
response_model: Final = response.model
|
||||
request_model_value: Final = request_body.get("model")
|
||||
request_model: Final = request_model_value if isinstance(request_model_value, str) else None
|
||||
response_model: Final = (
|
||||
laya_response_model(response_body, request_model) if custom_llm_provider == "laya" else response.model
|
||||
)
|
||||
logged_model: Final = response_model or request_model or "unknown"
|
||||
model_name: Final = f"{custom_llm_provider}/{logged_model}"
|
||||
usage: Final = response.usage or _TypeSafeUsage()
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ from fastapi import (
|
|||
status,
|
||||
)
|
||||
from fastapi.responses import StreamingResponse
|
||||
from pydantic import TypeAdapter
|
||||
from starlette.datastructures import UploadFile as StarletteUploadFile
|
||||
from starlette.websockets import WebSocketState
|
||||
from websockets.asyncio.client import connect
|
||||
|
|
@ -64,6 +65,7 @@ from litellm.llms.base_llm.managed_resources.utils import (
|
|||
resolve_passthrough_managed_id_provider,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.llms.laya.common_utils import validate_laya_request
|
||||
from litellm.passthrough import BasePassthroughUtils
|
||||
from litellm.proxy._types import (
|
||||
ConfigFieldInfo,
|
||||
|
|
@ -74,7 +76,11 @@ from litellm.proxy._types import (
|
|||
ProxyException,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.auth_utils import request_dispatched_to_pass_through_endpoint
|
||||
from litellm.proxy.auth.auth_utils import (
|
||||
get_model_from_request,
|
||||
get_request_route,
|
||||
request_dispatched_to_pass_through_endpoint,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_request_processing import (
|
||||
ProxyBaseLLMRequestProcessing,
|
||||
|
|
@ -100,6 +106,8 @@ from litellm.proxy.common_utils.sse_keepalive import (
|
|||
from litellm.proxy.litellm_pre_call_utils import (
|
||||
LiteLLMProxyRequestSetup,
|
||||
_get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage] # shared proxy helper, same import style as _read_request_body above
|
||||
_key_or_team_allows_client_pricing_override, # pyright: ignore[reportPrivateUsage] # reuse the proxy's pricing trust policy
|
||||
_strip_client_pricing_overrides, # pyright: ignore[reportPrivateUsage] # sanitize before trusted hooks add guardrail costs
|
||||
)
|
||||
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
|
||||
from litellm.proxy.utils import normalize_route_for_root_path
|
||||
|
|
@ -585,7 +593,18 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
|
|||
"""
|
||||
Filter out litellm params from the request body
|
||||
"""
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
|
||||
_parsed_body = _parsed_body or {}
|
||||
managed_model: Final = get_model_from_request(
|
||||
request_data=_parsed_body,
|
||||
route=get_request_route(request),
|
||||
request_headers=request.headers,
|
||||
request_query_params=request.query_params,
|
||||
llm_router=llm_router,
|
||||
request=request,
|
||||
team_id=user_api_key_dict.team_id,
|
||||
)
|
||||
|
||||
litellm_keys_in_body: Final = MappingProxyType(
|
||||
{k: _parsed_body.pop(k) for k in types_utils.all_litellm_params if k in _parsed_body}
|
||||
|
|
@ -600,11 +619,11 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
|
|||
|
||||
litellm_metadata: Final = litellm_keys_in_body.get("litellm_metadata")
|
||||
metadata: Final = litellm_keys_in_body.get("metadata")
|
||||
if litellm_metadata:
|
||||
_metadata.update(litellm_metadata)
|
||||
if metadata:
|
||||
_metadata.update(metadata)
|
||||
for client_metadata in (litellm_metadata, metadata):
|
||||
if isinstance(client_metadata, dict):
|
||||
_metadata.update({k: v for k, v in client_metadata.items() if not k.startswith("user_api_key_")})
|
||||
|
||||
_metadata = _apply_key_team_project_controls(user_api_key_dict=user_api_key_dict, metadata=_metadata)
|
||||
_metadata = _update_metadata_with_tags_in_header(
|
||||
request=request,
|
||||
metadata=_metadata,
|
||||
|
|
@ -631,10 +650,19 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
|
|||
# would attribute it to a budget the operator scoped to a LiteLLM model that
|
||||
# merely shares the name.
|
||||
if not request_dispatched_to_pass_through_endpoint(request):
|
||||
_metadata["model_group"] = managed_model if isinstance(managed_model, str) else None
|
||||
_metadata["user_api_key_model_max_budget"] = user_api_key_dict.model_max_budget
|
||||
_metadata["user_api_key_team_model_max_budget"] = user_api_key_dict.team_model_max_budget
|
||||
_metadata["user_api_key_user_model_max_budget"] = user_api_key_dict.user_model_max_budget
|
||||
_metadata["user_api_key_end_user_model_max_budget"] = user_api_key_dict.end_user_model_max_budget
|
||||
else:
|
||||
for field in (
|
||||
"user_api_key_model_max_budget",
|
||||
"user_api_key_team_model_max_budget",
|
||||
"user_api_key_user_model_max_budget",
|
||||
"user_api_key_end_user_model_max_budget",
|
||||
):
|
||||
_metadata.pop(field, None)
|
||||
_metadata.update(
|
||||
LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_dict)
|
||||
)
|
||||
|
|
@ -1131,6 +1159,15 @@ async def pass_through_request(
|
|||
_parsed_body,
|
||||
)
|
||||
|
||||
if not _key_or_team_allows_client_pricing_override(user_api_key_dict):
|
||||
pricing_body: Final = TypeAdapter(dict[str, object]).validate_python(_parsed_body)
|
||||
_strip_client_pricing_overrides(pricing_body)
|
||||
_parsed_body = pricing_body
|
||||
if custom_llm_provider == "laya":
|
||||
laya_request: Final = TypeAdapter(Mapping[str, object]).validate_python(_parsed_body)
|
||||
checkpoint: Final = validate_laya_request(laya_request)
|
||||
_parsed_body["model"] = f"laya/{checkpoint}"
|
||||
|
||||
### COLLECT GUARDRAILS FOR PASSTHROUGH ENDPOINT ###
|
||||
# Passthrough endpoints are opt-in only for guardrails
|
||||
# When enabled, collect guardrails from org/team/key levels + passthrough-specific
|
||||
|
|
@ -1186,6 +1223,17 @@ async def pass_through_request(
|
|||
call_type="pass_through_endpoint",
|
||||
endpoint_type=endpoint_type,
|
||||
)
|
||||
if custom_llm_provider == "laya":
|
||||
hook_body: Final = TypeAdapter(dict[str, object]).validate_python(_parsed_body)
|
||||
hook_model: Final = hook_body.get("model")
|
||||
laya_body: Final = MappingProxyType(
|
||||
{
|
||||
**hook_body,
|
||||
"model": hook_model.removeprefix("laya/") if isinstance(hook_model, str) else hook_model,
|
||||
}
|
||||
)
|
||||
_ = validate_laya_request(laya_body)
|
||||
_parsed_body = TypeAdapter(dict[str, object]).validate_python(laya_body)
|
||||
resolved_timeout: Final = resolve_pass_through_request_timeout(timeout)
|
||||
async_client_obj: Final = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.PassThroughEndpoint,
|
||||
|
|
@ -1886,6 +1934,20 @@ async def pass_through_request(
|
|||
)
|
||||
|
||||
|
||||
def _apply_key_team_project_controls(
|
||||
user_api_key_dict: UserAPIKeyAuth, metadata: dict[str, object]
|
||||
) -> dict[str, object]:
|
||||
data: Final = LiteLLMProxyRequestSetup.add_key_level_controls(
|
||||
key_metadata=user_api_key_dict.metadata,
|
||||
data={"metadata": metadata},
|
||||
_metadata_variable_name="metadata",
|
||||
)
|
||||
return LiteLLMProxyRequestSetup.add_team_and_project_level_controls(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
metadata=data["metadata"],
|
||||
)
|
||||
|
||||
|
||||
def _update_metadata_with_tags_in_header(request: Request, metadata: dict) -> dict:
|
||||
"""
|
||||
If tags are in the request headers, add them to the metadata
|
||||
|
|
@ -1906,9 +1968,10 @@ def _update_metadata_with_tags_in_header(request: Request, metadata: dict) -> di
|
|||
|
||||
# Only add tags key if there are tags to add
|
||||
if tags_to_add:
|
||||
if "tags" not in metadata:
|
||||
metadata["tags"] = []
|
||||
metadata["tags"].extend(tags_to_add)
|
||||
metadata["tags"] = LiteLLMProxyRequestSetup._merge_tags(
|
||||
request_tags=metadata.get("tags"),
|
||||
tags_to_add=tags_to_add,
|
||||
)
|
||||
|
||||
return metadata
|
||||
|
||||
|
|
@ -2389,7 +2452,9 @@ async def websocket_passthrough_request(
|
|||
# with the existing _init_kwargs_for_pass_through_endpoint function
|
||||
class DummyRequest:
|
||||
def __init__(self, url: str, method: str = "WEBSOCKET", headers: dict | None = None):
|
||||
self.url = url
|
||||
self.url = httpx.URL(url)
|
||||
self.scope = websocket.scope
|
||||
self.query_params = websocket.query_params
|
||||
self.method = method
|
||||
self.headers = headers or {}
|
||||
|
||||
|
|
|
|||
|
|
@ -334,8 +334,10 @@ class PassThroughEndpointLogging:
|
|||
)
|
||||
standard_logging_response_object = transcribe_handler_result["result"] # rebind-ok: elif-chain
|
||||
kwargs = transcribe_handler_result["kwargs"] # rebind-ok: elif-chain contract
|
||||
elif self.is_typesafe_route(custom_llm_provider) or self.is_openrouter_decisions_route(
|
||||
url_route, custom_llm_provider
|
||||
elif (
|
||||
self.is_typesafe_route(custom_llm_provider)
|
||||
or custom_llm_provider == "laya"
|
||||
or self.is_openrouter_decisions_route(url_route, custom_llm_provider)
|
||||
):
|
||||
from .llm_provider_handlers.typesafe_passthrough_logging_handler import (
|
||||
TypeSafePassthroughLoggingHandler,
|
||||
|
|
|
|||
|
|
@ -1744,6 +1744,27 @@ model LiteLLM_AutoRouterUserSession {
|
|||
@@index([user_id, last_turn_at], map: "idx_autorouter_user_session_user_last_turn")
|
||||
}
|
||||
|
||||
// Auto-routed requests per UTC request day and router: the selected-day money behind the
|
||||
// auto-router usage view. Written in the same statement as the session rollup, so a day row
|
||||
// and its session row never disagree; corrected in the same transaction as late baselines.
|
||||
model LiteLLM_AutoRouterDailySpend {
|
||||
date String
|
||||
api_key String
|
||||
user_id String
|
||||
router_name String
|
||||
router_type String
|
||||
turns Int @default(0)
|
||||
spend Float @default(0)
|
||||
saved_spend Float @default(0)
|
||||
savings_estimated_turns Int @default(0)
|
||||
savings_estimated_actual_spend Float @default(0)
|
||||
savings_estimated_saved_spend Float @default(0)
|
||||
classifier_cost Float @default(0)
|
||||
classifier_cost_recorded_turns Int @default(0)
|
||||
|
||||
@@id([date, api_key, user_id, router_name, router_type])
|
||||
}
|
||||
|
||||
// Shadow eval: evaluation of an auto-router against one or more keys' live traffic, in
|
||||
// either direction. forward duplicates the requests the keys did not route through the
|
||||
// router through it, answering whether they should adopt it; reverse duplicates the
|
||||
|
|
|
|||
|
|
@ -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})
|
||||
|
|
|
|||
|
|
@ -108,7 +108,7 @@ from .config import (
|
|||
ComplexityRouterConfig,
|
||||
ComplexityTier,
|
||||
CustomDimension,
|
||||
JevClassifierConfig,
|
||||
OpenSourceClassifierConfig,
|
||||
TierDefinition,
|
||||
)
|
||||
from .jev_classifier import (
|
||||
|
|
@ -1308,10 +1308,22 @@ class ComplexityRouter(CustomLogger):
|
|||
"""
|
||||
|
||||
@staticmethod
|
||||
def _build_jev_client(config: JevClassifierConfig) -> JevClassifierClient:
|
||||
def _build_jev_client(config: OpenSourceClassifierConfig) -> JevClassifierClient:
|
||||
if config.provider == "laya":
|
||||
from litellm.llms.laya.common_utils import laya_connection
|
||||
|
||||
connection: Final = laya_connection(config.api_base, config.api_key)
|
||||
return HttpJevClassifierClient(
|
||||
api_key=connection.api_key,
|
||||
api_base=connection.api_base,
|
||||
http_client=get_async_httpx_client(httpxSpecialProvider.PassThroughEndpoint),
|
||||
provider="laya",
|
||||
)
|
||||
api_key: Final = config.api_key or get_secret_str("TYPESAFE_API_KEY")
|
||||
if not api_key:
|
||||
raise ValueError("jev_classifier_config.api_key or TYPESAFE_API_KEY is required for classifier_type 'jev'")
|
||||
raise ValueError(
|
||||
"opensource_classifier_config.api_key or TYPESAFE_API_KEY is required for classifier_type 'oss_classifier'"
|
||||
)
|
||||
api_base: Final = config.api_base or get_secret_str("TYPESAFE_API_BASE") or "https://api.typesafe.ai"
|
||||
return HttpJevClassifierClient(
|
||||
api_key=api_key,
|
||||
|
|
@ -1354,12 +1366,12 @@ class ComplexityRouter(CustomLogger):
|
|||
if default_model:
|
||||
self.config.default_model = default_model
|
||||
|
||||
jev_config: Final = self.config.jev_classifier_config
|
||||
jev_config: Final = self.config.opensource_classifier_config
|
||||
self._jev_client: JevClassifierClient | None = (
|
||||
jev_client
|
||||
if jev_client is not None
|
||||
else self._build_jev_client(jev_config)
|
||||
if self.config.classifier_type == "jev" and jev_config is not None
|
||||
if self.config.classifier_type == "oss_classifier" and jev_config is not None
|
||||
else None
|
||||
)
|
||||
|
||||
|
|
@ -1459,7 +1471,11 @@ class ComplexityRouter(CustomLogger):
|
|||
and self.config.classifier_llm_config.circuit_breaker_enabled
|
||||
)
|
||||
else jev_config.circuit_breaker_cooldown_seconds
|
||||
if (self.config.classifier_type == "jev" and jev_config is not None and jev_config.circuit_breaker_enabled)
|
||||
if (
|
||||
self.config.classifier_type == "oss_classifier"
|
||||
and jev_config is not None
|
||||
and jev_config.circuit_breaker_enabled
|
||||
)
|
||||
else None
|
||||
)
|
||||
self._classifier_circuit_breaker: _ClassifierCircuitBreaker | None = (
|
||||
|
|
@ -1909,7 +1925,7 @@ class ComplexityRouter(CustomLogger):
|
|||
return self._classify_with_heuristic_v2(prompt)
|
||||
if self.config.classifier_type == "custom":
|
||||
return await self._classify_with_plugin(prompt, system_prompt, request_kwargs, raw_messages)
|
||||
if self.config.classifier_type == "jev":
|
||||
if self.config.classifier_type == "oss_classifier":
|
||||
return await self._jev_classifier_outcome(prompt, system_prompt, request_kwargs, messages)
|
||||
if self.config.classifier_type in ("heuristic_first", "hybrid") and _encrypted_classifier_task(
|
||||
request_kwargs, self._reminder_markers_for_request(request_kwargs or EMPTY_MAPPING)
|
||||
|
|
@ -2161,7 +2177,7 @@ class ComplexityRouter(CustomLogger):
|
|||
request_kwargs: Mapping[str, object] | None,
|
||||
messages: Sequence[Mapping[str, object]] | None,
|
||||
) -> ClassificationOutcome:
|
||||
config: Final = self.config.jev_classifier_config
|
||||
config: Final = self.config.opensource_classifier_config
|
||||
client: Final = self._jev_client
|
||||
if config is None or client is None:
|
||||
return self._classifier_failure_outcome("jev classifier is not configured", prompt, system_prompt)
|
||||
|
|
@ -2212,12 +2228,14 @@ class ComplexityRouter(CustomLogger):
|
|||
if not self._tier_pools().get(tier_name):
|
||||
raise ValueError(f"Jev classifier returned tier {tier_name!r}, which has no models configured")
|
||||
model: Final = response.model or config.model
|
||||
accounting_provider: Final = "laya" if config.provider == "laya" else "typesafe"
|
||||
verdict: Final = JevVerdict(
|
||||
label=answer.choice,
|
||||
probabilities=answer.probabilities,
|
||||
confidence=answer.confidence,
|
||||
model=model,
|
||||
cost=jev_classifier_cost(response, config.model),
|
||||
cost=jev_classifier_cost(response, config.model, accounting_provider),
|
||||
provider=accounting_provider,
|
||||
)
|
||||
if breaker is not None and permit is not None:
|
||||
breaker.record_success(permit)
|
||||
|
|
@ -2225,8 +2243,8 @@ class ComplexityRouter(CustomLogger):
|
|||
tier=tier,
|
||||
score=None,
|
||||
signals=(
|
||||
f"jev-classifier:{tier_name}",
|
||||
f"jev-confidence={answer.confidence:.6f}",
|
||||
f"{'laya' if config.provider == 'laya' else 'jev'}-classifier:{tier_name}",
|
||||
f"{'laya' if config.provider == 'laya' else 'jev'}-confidence={answer.confidence:.6f}",
|
||||
*(
|
||||
f"tier-probability:{label}={probability:.6f}"
|
||||
for label, probability in answer.probabilities.items()
|
||||
|
|
@ -4757,7 +4775,7 @@ class ComplexityRouter(CustomLogger):
|
|||
|
||||
tier_litellm_params: Final = self._litellm_params_for_model(tier, routed_model)
|
||||
classifier_model: Final = (
|
||||
f"typesafe/{outcome.jev_verdict.model}"
|
||||
f"{outcome.jev_verdict.provider}/{outcome.jev_verdict.model}"
|
||||
if outcome.cause == "jev_classifier" and outcome.jev_verdict is not None
|
||||
else self.config.classifier_llm_config.model
|
||||
if outcome.cause in ("llm_classifier", "capability_classifier", "llm_v2_classifier", "llm_v2_fallback")
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ import math
|
|||
import re
|
||||
import warnings
|
||||
from collections.abc import Iterable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from types import MappingProxyType
|
||||
from typing import Annotated, Final, Literal, NamedTuple
|
||||
|
|
@ -19,6 +20,7 @@ from pydantic import (
|
|||
Field,
|
||||
SkipValidation,
|
||||
StrictFloat,
|
||||
TypeAdapter,
|
||||
field_serializer,
|
||||
field_validator,
|
||||
model_validator,
|
||||
|
|
@ -674,14 +676,34 @@ class CapabilityClassifierConfig(BaseModel):
|
|||
return self
|
||||
|
||||
|
||||
class JevClassifierConfig(BaseModel):
|
||||
def normalize_classifier_config_aliases(config: Mapping[str, object]) -> Mapping[str, object]:
|
||||
if "jev_classifier_config" in config and "opensource_classifier_config" in config:
|
||||
return config
|
||||
normalized: Final = dict(config)
|
||||
if "jev_classifier_config" in normalized:
|
||||
normalized["opensource_classifier_config"] = normalized.pop("jev_classifier_config")
|
||||
if normalized.get("classifier_type") == "jev":
|
||||
normalized["classifier_type"] = "oss_classifier"
|
||||
classifier: Final = normalized.get("opensource_classifier_config")
|
||||
if isinstance(classifier, Mapping):
|
||||
classifier_fields: Final = TypeAdapter(Mapping[str, object]).validate_python(classifier)
|
||||
if classifier_fields.get("provider") == "typesafe":
|
||||
normalized["opensource_classifier_config"] = {
|
||||
**classifier_fields,
|
||||
"provider": "jev",
|
||||
}
|
||||
return normalized
|
||||
|
||||
|
||||
class OpenSourceClassifierConfig(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid", frozen=True)
|
||||
|
||||
provider: Literal["jev", "laya"] = "jev"
|
||||
model: str = "jev-latest"
|
||||
api_key: str | None = Field(default=None, description="TypeSafe API key, falling back to TYPESAFE_API_KEY")
|
||||
api_key: str | None = Field(default=None, description="Provider API key; optional for self-hosted Laya")
|
||||
api_base: str | None = Field(
|
||||
default=None,
|
||||
description="TypeSafe API base, falling back to TYPESAFE_API_BASE and then https://api.typesafe.ai",
|
||||
description="Provider API base; defaults to TYPESAFE_API_BASE or LAYA_API_BASE for the selected provider",
|
||||
)
|
||||
timeout_ms: int = Field(default=3000, ge=1)
|
||||
instructions: str | None = Field(
|
||||
|
|
@ -691,30 +713,112 @@ class JevClassifierConfig(BaseModel):
|
|||
circuit_breaker_enabled: bool = True
|
||||
circuit_breaker_cooldown_seconds: float = Field(default=30.0, gt=0.0)
|
||||
|
||||
@field_validator("provider", mode="before")
|
||||
@classmethod
|
||||
def _normalize_provider_alias(cls, value: object) -> object:
|
||||
return "jev" if value == "typesafe" else value
|
||||
|
||||
@field_validator("instructions")
|
||||
@classmethod
|
||||
def _reject_blank_instructions(cls, value: str | None) -> str | None:
|
||||
if value is not None and not value.strip():
|
||||
raise ValueError("jev_classifier_config.instructions must be non-empty; omit it to use the default")
|
||||
raise ValueError("opensource_classifier_config.instructions must be non-empty; omit it to use the default")
|
||||
return value
|
||||
|
||||
@field_validator("api_key")
|
||||
@classmethod
|
||||
def _reject_blank_api_key(cls, value: str | None) -> str | None:
|
||||
if value is not None and not value.strip():
|
||||
raise ValueError("jev_classifier_config.api_key must be non-empty; omit it to use TYPESAFE_API_KEY")
|
||||
raise ValueError("opensource_classifier_config.api_key must be non-empty; omit it to use TYPESAFE_API_KEY")
|
||||
return value
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _keep_the_environment_key_on_the_environment_base(self) -> "JevClassifierConfig":
|
||||
def _keep_the_environment_key_on_the_environment_base(self) -> "OpenSourceClassifierConfig":
|
||||
if self.provider == "laya":
|
||||
from litellm.llms.laya.common_utils import validate_laya_api_base, validate_laya_model
|
||||
|
||||
_ = validate_laya_model(self.model)
|
||||
if self.api_base is not None:
|
||||
_ = validate_laya_api_base(self.api_base)
|
||||
return self
|
||||
if self.api_base is not None and self.api_key is None:
|
||||
raise ValueError(
|
||||
"jev_classifier_config.api_base requires jev_classifier_config.api_key: TYPESAFE_API_KEY is only sent "
|
||||
"opensource_classifier_config.api_base requires opensource_classifier_config.api_key: TYPESAFE_API_KEY is only sent "
|
||||
"to TYPESAFE_API_BASE or https://api.typesafe.ai"
|
||||
)
|
||||
return self
|
||||
|
||||
|
||||
JevClassifierConfig = OpenSourceClassifierConfig
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ComplexityRouterConfigWrite:
|
||||
submitted: Mapping[str, object] | None
|
||||
effective: Mapping[str, object] | None
|
||||
|
||||
@property
|
||||
def supplied_connection_fields(self) -> frozenset[str]:
|
||||
classifier: Final = self.submitted.get("opensource_classifier_config") if self.submitted is not None else None
|
||||
return frozenset(
|
||||
field for field in ("api_base", "api_key") if isinstance(classifier, Mapping) and field in classifier
|
||||
)
|
||||
|
||||
|
||||
def resolve_complexity_router_config_write(
|
||||
incoming: Mapping[str, object] | None, stored: Mapping[str, object] | None
|
||||
) -> ComplexityRouterConfigWrite:
|
||||
if incoming is None:
|
||||
return ComplexityRouterConfigWrite(submitted=None, effective=stored)
|
||||
return _resolve_normalized_complexity_router_config_write(
|
||||
normalize_classifier_config_aliases(incoming),
|
||||
normalize_classifier_config_aliases(stored) if stored is not None else None,
|
||||
)
|
||||
|
||||
|
||||
def _resolve_normalized_complexity_router_config_write(
|
||||
incoming: Mapping[str, object], stored: Mapping[str, object] | None
|
||||
) -> ComplexityRouterConfigWrite:
|
||||
if (
|
||||
stored is None
|
||||
or incoming.get("classifier_type") != "oss_classifier"
|
||||
or stored.get("classifier_type") != "oss_classifier"
|
||||
):
|
||||
return ComplexityRouterConfigWrite(submitted=incoming, effective=incoming)
|
||||
incoming_classifier: Final = incoming.get("opensource_classifier_config")
|
||||
stored_classifier: Final = stored.get("opensource_classifier_config")
|
||||
if not isinstance(incoming_classifier, Mapping) or not isinstance(stored_classifier, Mapping):
|
||||
return ComplexityRouterConfigWrite(submitted=incoming, effective=incoming)
|
||||
existing: Final = TypeAdapter(dict[str, object]).validate_python(stored_classifier)
|
||||
supplied: Final = TypeAdapter(dict[str, object]).validate_python(incoming_classifier)
|
||||
classifier: Final = (
|
||||
MappingProxyType({**supplied, "provider": existing["provider"]})
|
||||
if "provider" not in supplied and "provider" in existing
|
||||
else supplied
|
||||
)
|
||||
same_provider: Final = classifier.get("provider", "jev") == existing.get("provider", "jev")
|
||||
same_base: Final = "api_base" not in classifier or (
|
||||
classifier["api_base"] is not None and classifier["api_base"] == existing.get("api_base")
|
||||
)
|
||||
transport: Final = MappingProxyType(
|
||||
{
|
||||
key: value
|
||||
for key, value in existing.items()
|
||||
if same_provider and key in ("api_key", "api_base") and (key != "api_key" or same_base)
|
||||
}
|
||||
)
|
||||
return ComplexityRouterConfigWrite(
|
||||
submitted=MappingProxyType({**incoming, "opensource_classifier_config": classifier}),
|
||||
effective={
|
||||
**incoming,
|
||||
"opensource_classifier_config": {
|
||||
**transport,
|
||||
**classifier,
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
MAX_CUSTOM_PATTERN_REPEAT: Final[int] = 64
|
||||
MAX_CUSTOM_PATTERN_WORK: Final[int] = 2048
|
||||
MAX_CUSTOM_DIMENSIONS_WORK: Final[int] = 8192
|
||||
|
|
@ -846,6 +950,20 @@ class ContextCompactionConfig(BaseModel):
|
|||
class ComplexityRouterConfig(BaseModel):
|
||||
"""Configuration for the ComplexityRouter."""
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def _normalize_classifier_aliases(cls, value: object) -> object:
|
||||
if not isinstance(value, Mapping):
|
||||
return value
|
||||
config: Final = TypeAdapter(dict[str, object]).validate_python(value)
|
||||
if "jev_classifier_config" in config and "opensource_classifier_config" in config:
|
||||
raise ValueError("Use only opensource_classifier_config; do not also supply jev_classifier_config")
|
||||
return normalize_classifier_config_aliases(config)
|
||||
|
||||
@property
|
||||
def jev_classifier_config(self) -> OpenSourceClassifierConfig | None:
|
||||
return self.opensource_classifier_config
|
||||
|
||||
# string = pin; list = random pick when adaptive=False, soft-floor home pool when adaptive=True
|
||||
tiers: dict[str, str | list[str]] = Field(
|
||||
default_factory=lambda: DEFAULT_TIER_MODELS.copy(),
|
||||
|
|
@ -880,7 +998,7 @@ class ComplexityRouterConfig(BaseModel):
|
|||
"becomes that tier's rubric bullet; entries named after a built-in tier may omit the "
|
||||
"description and inherit the built-in criteria. List order is ascending severity and "
|
||||
"decides which tier wins when several keyword_tier_rules match. Requires classifier_type "
|
||||
"'llm', 'jev' or 'custom', a fallback_tier, and `tiers` keys matching the defined names exactly. Escalation, "
|
||||
"'llm', 'oss_classifier' or 'custom', a fallback_tier, and `tiers` keys matching the defined names exactly. Escalation, "
|
||||
"adaptive selection, session affinity, plugins, tier_labels, and the calibration-example "
|
||||
"rubric presets are unavailable with a custom tier set: the first four are built on the "
|
||||
"built-in tier ladder, and the last two rename or exemplify tiers the set replaces."
|
||||
|
|
@ -1024,7 +1142,7 @@ class ComplexityRouterConfig(BaseModel):
|
|||
"custom",
|
||||
"heuristic_first",
|
||||
"hybrid",
|
||||
"jev",
|
||||
"oss_classifier",
|
||||
] = Field(
|
||||
default="heuristic",
|
||||
description=(
|
||||
|
|
@ -1032,7 +1150,7 @@ class ComplexityRouterConfig(BaseModel):
|
|||
"an LLM tier-selection call, a Switchyard-compatible capability forecast, a joint Fuse V2 forecast, "
|
||||
"a custom classifier plugin, 'heuristic_first', which scores locally and only pays for the LLM classifier when the "
|
||||
"local scorer does not confidently land a cheap tier, or 'hybrid', which trusts the local scorer "
|
||||
"everywhere except when its score lands near a tier boundary, or 'jev', a TypeSafe AI Jev structured choice call"
|
||||
"everywhere except when its score lands near a tier boundary, or 'oss_classifier', a structured choice call using Jev or Laya"
|
||||
),
|
||||
)
|
||||
llm_v2_config: LLMV2Config | None = Field(
|
||||
|
|
@ -1073,7 +1191,7 @@ class ComplexityRouterConfig(BaseModel):
|
|||
"and otherwise routes to capable_tier"
|
||||
),
|
||||
)
|
||||
jev_classifier_config: JevClassifierConfig | None = None
|
||||
opensource_classifier_config: OpenSourceClassifierConfig | None = None
|
||||
heuristic_first_max_tier: str | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
|
|
@ -1639,14 +1757,16 @@ class ComplexityRouterConfig(BaseModel):
|
|||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _validate_jev_classifier_config(self) -> "ComplexityRouterConfig":
|
||||
jev: Final = self.jev_classifier_config
|
||||
if self.classifier_type != "jev":
|
||||
def _validate_opensource_classifier_config(self) -> "ComplexityRouterConfig":
|
||||
jev: Final = self.opensource_classifier_config
|
||||
if self.classifier_type != "oss_classifier":
|
||||
if jev is not None:
|
||||
raise ValueError("jev_classifier_config requires classifier_type 'jev'; otherwise it has no effect")
|
||||
raise ValueError(
|
||||
"opensource_classifier_config requires classifier_type 'oss_classifier'; otherwise it has no effect"
|
||||
)
|
||||
return self
|
||||
if jev is None:
|
||||
raise ValueError("jev_classifier_config is required when classifier_type is 'jev'")
|
||||
raise ValueError("opensource_classifier_config is required when classifier_type is 'oss_classifier'")
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
|
|
@ -1962,9 +2082,9 @@ class ComplexityRouterConfig(BaseModel):
|
|||
"enable_non_reasoning_tier cannot be combined with tier_definitions: a custom tier set "
|
||||
f"replaces the built-in ladder, so name a tier {non_reasoning_key} in tier_definitions instead"
|
||||
)
|
||||
if self.classifier_type not in ("llm", "custom", "jev"):
|
||||
if self.classifier_type not in ("llm", "custom", "oss_classifier"):
|
||||
raise ValueError(
|
||||
f"enable_non_reasoning_tier requires classifier_type 'llm', 'jev' or 'custom', got "
|
||||
f"enable_non_reasoning_tier requires classifier_type 'llm', 'oss_classifier' or 'custom', got "
|
||||
f"{self.classifier_type!r}: the heuristic scorers only produce the four tiers from SIMPLE up, "
|
||||
f"so nothing would ever classify as {non_reasoning_key}"
|
||||
)
|
||||
|
|
@ -1997,7 +2117,7 @@ class ComplexityRouterConfig(BaseModel):
|
|||
raise ValueError(f"tier_definitions names must be unique (case-insensitive): {', '.join(duplicated)}")
|
||||
if self.classifier_type in ("heuristic", "heuristic_v2", "capability", "heuristic_first", "hybrid"):
|
||||
raise ValueError(
|
||||
"tier_definitions requires classifier_type 'llm', 'jev' or 'custom': the heuristic scorer only "
|
||||
"tier_definitions requires classifier_type 'llm', 'oss_classifier' or 'custom': the heuristic scorer only "
|
||||
"produces the built-in tiers from SIMPLE up, as does heuristic_v2"
|
||||
)
|
||||
conflicts: Final = self._tier_definition_conflicts()
|
||||
|
|
@ -2164,7 +2284,9 @@ class ComplexityRouterConfig(BaseModel):
|
|||
)
|
||||
|
||||
|
||||
COMPLEXITY_ROUTER_CONFIG_KEYS: Final[frozenset[str]] = frozenset(ComplexityRouterConfig.model_fields)
|
||||
COMPLEXITY_ROUTER_CONFIG_KEYS: Final[frozenset[str]] = frozenset(ComplexityRouterConfig.model_fields) | frozenset(
|
||||
("jev_classifier_config",)
|
||||
)
|
||||
"""Every setting name this config owns, derived from the model so a field added later is covered.
|
||||
|
||||
These names are disjoint from the OpenAI request params, from ``all_litellm_params``, and from the
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ from litellm.litellm_core_utils.internal_call_metadata import (
|
|||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.llms.laya.common_utils import laya_response_model
|
||||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.typesafe_passthrough_logging_handler import (
|
||||
TypeSafePassthroughLoggingHandler,
|
||||
)
|
||||
|
|
@ -78,10 +79,17 @@ class JevClassifierClient(Protocol):
|
|||
|
||||
|
||||
class HttpJevClassifierClient:
|
||||
def __init__(self, api_key: str, api_base: str, http_client: AsyncHTTPHandler) -> None:
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str | None,
|
||||
api_base: str,
|
||||
http_client: AsyncHTTPHandler,
|
||||
provider: Literal["typesafe", "laya"] = "typesafe",
|
||||
) -> None:
|
||||
self._api_key = api_key
|
||||
self._api_base = api_base.rstrip("/")
|
||||
self._http_client = http_client
|
||||
self._provider = provider
|
||||
|
||||
async def evaluate(
|
||||
self,
|
||||
|
|
@ -90,26 +98,30 @@ class HttpJevClassifierClient:
|
|||
request_kwargs: Mapping[str, object] | None = None,
|
||||
) -> JevSystemOneResponse:
|
||||
start_time: Final = datetime.now(timezone.utc)
|
||||
authorization: Final[Mapping[str, str]] = (
|
||||
MappingProxyType({"Authorization": f"Bearer {self._api_key}"}) if self._api_key else MappingProxyType({})
|
||||
)
|
||||
response: Final = await self._http_client.post( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler has a dynamic post signature
|
||||
f"{self._api_base}/v1/systemone",
|
||||
json=request.model_dump(mode="json"),
|
||||
headers=MappingProxyType(
|
||||
{
|
||||
"Authorization": f"Bearer {self._api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
), # pyright: ignore[reportArgumentType] # HTTP headers are not mutated by AsyncHTTPHandler
|
||||
headers=MappingProxyType({**authorization, "Content-Type": "application/json"}), # pyright: ignore[reportArgumentType] # HTTP headers are not mutated by AsyncHTTPHandler
|
||||
timeout=timeout_s,
|
||||
)
|
||||
response.raise_for_status()
|
||||
body: Final = TypeAdapter(dict[str, object]).validate_json(response.content)
|
||||
normalized_body: Final = (
|
||||
MappingProxyType({**body, "model": laya_response_model(body, request.model)})
|
||||
if self._provider == "laya"
|
||||
else body
|
||||
)
|
||||
try:
|
||||
self._log_response(request, response, request_kwargs, start_time)
|
||||
except Exception as exc: # noqa: BLE001 # logging integrations must not discard a provider verdict
|
||||
verbose_router_logger.warning("JEV response logging failed (%s)", type(exc).__name__)
|
||||
return TypeAdapter(JevSystemOneResponse).validate_python(response.json())
|
||||
return TypeAdapter(JevSystemOneResponse).validate_python(normalized_body)
|
||||
|
||||
@staticmethod
|
||||
def _log_response(
|
||||
self,
|
||||
request: JevSystemOneRequest,
|
||||
response: httpx.Response,
|
||||
request_kwargs: Mapping[str, object] | None,
|
||||
|
|
@ -139,7 +151,7 @@ class HttpJevClassifierClient:
|
|||
"turn_off_message_logging": effective_turn_off_message_logging(request_kwargs),
|
||||
}
|
||||
logging_obj: Final = Logging(
|
||||
model=f"typesafe/{request.model}",
|
||||
model=f"{self._provider}/{request.model}",
|
||||
messages=[{"role": "user", "content": request.state}],
|
||||
stream=False,
|
||||
call_type="pass_through_endpoint",
|
||||
|
|
@ -150,7 +162,7 @@ class HttpJevClassifierClient:
|
|||
kwargs=params,
|
||||
)
|
||||
logging_obj.update_environment_variables(
|
||||
model=f"typesafe/{request.model}",
|
||||
model=f"{self._provider}/{request.model}",
|
||||
user=parent_user if isinstance(parent_user := parent.get("user"), str) else None,
|
||||
optional_params={},
|
||||
litellm_params=params,
|
||||
|
|
@ -165,7 +177,7 @@ class HttpJevClassifierClient:
|
|||
end_time=end_time,
|
||||
cache_hit=False,
|
||||
request_body=MappingProxyType({"model": request.model}),
|
||||
custom_llm_provider="typesafe",
|
||||
custom_llm_provider=self._provider,
|
||||
litellm_params=params,
|
||||
)
|
||||
success_handlers: Final = logging_obj.dispatch_success_handlers(
|
||||
|
|
@ -189,6 +201,7 @@ class JevVerdict(NamedTuple):
|
|||
confidence: float
|
||||
model: str
|
||||
cost: float | None
|
||||
provider: Literal["typesafe", "laya"] = "typesafe"
|
||||
|
||||
|
||||
class _RegistryPricing(BaseModel):
|
||||
|
|
@ -211,12 +224,14 @@ def build_jev_request(
|
|||
return JevSystemOneRequest(state=state, model=model, questions=MappingProxyType({"tier": question}))
|
||||
|
||||
|
||||
def jev_classifier_cost(response: JevSystemOneResponse, configured_model: str) -> float | None:
|
||||
def jev_classifier_cost(
|
||||
response: JevSystemOneResponse, configured_model: str, provider: Literal["typesafe", "laya"] = "typesafe"
|
||||
) -> float | None:
|
||||
usage: Final = response.usage
|
||||
if usage is None:
|
||||
return None
|
||||
model: Final = response.model or configured_model
|
||||
model_key: Final = f"typesafe/{model}"
|
||||
model_key: Final = f"{provider}/{model}"
|
||||
if model_key not in litellm.model_cost: # pyright: ignore[reportUnknownMemberType] # registry is dynamically typed
|
||||
return None
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ from litellm.router_strategy.complexity_router.config import (
|
|||
COMPLEXITY_ROUTER_CONFIG_KEYS,
|
||||
DEFAULT_JEV_INSTRUCTIONS,
|
||||
LLM_CLASSIFIER_TYPES,
|
||||
normalize_classifier_config_aliases,
|
||||
)
|
||||
|
||||
AUTO_ROUTER_MODEL_PREFIX: Final = "auto_router/"
|
||||
|
|
@ -151,8 +152,11 @@ def strategy_router_dependencies(
|
|||
)
|
||||
)
|
||||
)
|
||||
complexity: Final = _mapping(litellm_params.get("complexity_router_config"))
|
||||
complexity: Final = normalize_classifier_config_aliases(_mapping(litellm_params.get("complexity_router_config")))
|
||||
classifier: Final = _mapping(complexity.get("classifier_llm_config"))
|
||||
decision_classifier: Final = _mapping(complexity.get("opensource_classifier_config"))
|
||||
decision_provider: Final = decision_classifier.get("provider", "jev")
|
||||
accounting_provider: Final = "typesafe" if decision_provider == "jev" else decision_provider
|
||||
return tuple(
|
||||
dict.fromkeys(
|
||||
tuple(dep for tier in _mapping(complexity.get("tiers")).values() for dep in _pool(tier, "tier"))
|
||||
|
|
@ -165,10 +169,10 @@ def strategy_router_dependencies(
|
|||
)
|
||||
+ (
|
||||
_named(
|
||||
f"typesafe/{_mapping(complexity.get('jev_classifier_config')).get('model', 'jev-latest')}",
|
||||
f"{accounting_provider}/{decision_classifier.get('model', 'jev-latest')}",
|
||||
"evaluation",
|
||||
)
|
||||
if complexity.get("classifier_type") == "jev"
|
||||
if complexity.get("classifier_type") == "oss_classifier"
|
||||
else ()
|
||||
)
|
||||
+ (
|
||||
|
|
@ -206,9 +210,9 @@ def defines_custom_classifier_prompt(complexity_router_config: object) -> bool:
|
|||
Scoped to the classifier types that actually call an LLM, which is also where the config validator
|
||||
accepts these fields: the heuristic scorers never read them.
|
||||
"""
|
||||
config: Final = _mapping(complexity_router_config)
|
||||
if config.get("classifier_type") == "jev":
|
||||
instructions: Final = _mapping(config.get("jev_classifier_config")).get("instructions")
|
||||
config: Final = normalize_classifier_config_aliases(_mapping(complexity_router_config))
|
||||
if config.get("classifier_type") == "oss_classifier":
|
||||
instructions: Final = _mapping(config.get("opensource_classifier_config")).get("instructions")
|
||||
return isinstance(instructions, str) and instructions != DEFAULT_JEV_INSTRUCTIONS
|
||||
if config.get("classifier_type") not in LLM_CLASSIFIER_TYPES:
|
||||
return False
|
||||
|
|
@ -272,6 +276,9 @@ _OPERATOR_PROMPT_FIELDS_SQL: Final = " OR ".join(
|
|||
f"{{config}} ->> '{field}' IS NOT NULL" for field in OPERATOR_CLASSIFIER_PROMPT_FIELDS
|
||||
)
|
||||
_DEFAULT_JEV_INSTRUCTIONS_SQL: Final = DEFAULT_JEV_INSTRUCTIONS.replace("'", "''")
|
||||
_OPENSOURCE_CLASSIFIER_CONFIG_SQL: Final = (
|
||||
"COALESCE({config} -> 'opensource_classifier_config', {config} -> 'jev_classifier_config')"
|
||||
)
|
||||
|
||||
CUSTOMIZATION_CAPABILITY: Final = GatedAutoRouterCapability(
|
||||
key="tier_or_classifier_prompt",
|
||||
|
|
@ -286,9 +293,9 @@ CUSTOMIZATION_CAPABILITY: Final = GatedAutoRouterCapability(
|
|||
f"({{config}} ->> 'classifier_type' IN ({_LLM_CLASSIFIER_TYPES_SQL}) AND ("
|
||||
"{config} -> 'classifier_llm_config' ->> 'system_prompt' IS NOT NULL OR "
|
||||
f"{_OPERATOR_PROMPT_FIELDS_SQL})) OR "
|
||||
"({config} ->> 'classifier_type' = 'jev' AND "
|
||||
"jsonb_typeof({config} -> 'jev_classifier_config' -> 'instructions') = 'string' AND "
|
||||
f"{{config}} -> 'jev_classifier_config' ->> 'instructions' <> '{_DEFAULT_JEV_INSTRUCTIONS_SQL}')"
|
||||
"({config} ->> 'classifier_type' IN ('oss_classifier', 'jev') AND "
|
||||
f"jsonb_typeof({_OPENSOURCE_CLASSIFIER_CONFIG_SQL} -> 'instructions') = 'string' AND "
|
||||
f"{_OPENSOURCE_CLASSIFIER_CONFIG_SQL} ->> 'instructions' <> '{_DEFAULT_JEV_INSTRUCTIONS_SQL}')"
|
||||
),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -205,14 +205,18 @@ class AutoRouterCacheStats(BaseModel):
|
|||
|
||||
|
||||
class AutoRouterBenchmarkTotals(BaseModel):
|
||||
"""Session-shape and savings aggregates over auto-routed traffic in the window."""
|
||||
"""Auto-routed traffic in the window. Turns, spend and savings count requests on the selected UTC days;
|
||||
the session averages and cache stats describe every session overlapping the window, whole."""
|
||||
|
||||
sessions: int
|
||||
turns: int
|
||||
avg_turns_per_session: float
|
||||
avg_session_seconds: float
|
||||
avg_tokens_per_session: float
|
||||
spend: float = Field(description="What the routed traffic actually cost")
|
||||
sessions: int = Field(description="Sessions overlapping the window, counted whole")
|
||||
turns: int = Field(description="Auto-routed requests on the selected UTC days")
|
||||
avg_turns_per_session: float | None = Field(
|
||||
description="Lifetime turns per overlapping session; null when the window has routed requests but no session "
|
||||
"rows for this router type, such as an alias whose router type changed mid-session"
|
||||
)
|
||||
avg_session_seconds: float | None = Field(description="Lifetime seconds per overlapping session; null as above")
|
||||
avg_tokens_per_session: float | None = Field(description="Lifetime tokens per overlapping session; null as above")
|
||||
spend: float = Field(description="What the selected days' routed traffic actually cost")
|
||||
classifier_cost: float | None = Field(
|
||||
description="Recorded LLM classifier cost already included in spend; null when any session turns predate "
|
||||
"subtotal recording, and zero for an empty window"
|
||||
|
|
@ -229,14 +233,19 @@ class AutoRouterBenchmarkTotals(BaseModel):
|
|||
"null when classification costs for those requests are unavailable",
|
||||
)
|
||||
saved_spend: float | None = Field(
|
||||
description="Recorded historical savings plus newer estimates; null when traffic has no recorded savings estimates"
|
||||
description="Recorded savings on the selected UTC days; null when traffic has no recorded savings estimates. "
|
||||
"On totals this is the same daily figure the Overall savings view reports"
|
||||
)
|
||||
unattributed_saved_spend: float | None = Field(
|
||||
default=None,
|
||||
description="Part of saved_spend no router's daily rows account for, such as history recorded before "
|
||||
"per-router daily tracking; when set, baseline_spend and saved_pct are null",
|
||||
)
|
||||
baseline_spend: float | None = Field(
|
||||
description="Estimated single-model cost: compared actual spend plus recorded savings; "
|
||||
"null when traffic has no recorded savings"
|
||||
)
|
||||
saved_pct: float | None = Field(description="Recorded savings over baseline_spend, as a percentage")
|
||||
saved_per_session: float | None = Field(description="Recorded savings per session, including historical estimates")
|
||||
cache: AutoRouterCacheStats
|
||||
|
||||
|
||||
|
|
@ -291,7 +300,7 @@ class AutoRouterSessionResponse(BaseModel):
|
|||
|
||||
|
||||
class AutoRouterBenchmarksResponse(BaseModel):
|
||||
"""Benchmarks for the auto-router dashboard, aggregated from the per-session rollup."""
|
||||
"""Benchmarks for the auto-router dashboard, aggregated from the per-session and per-day rollups."""
|
||||
|
||||
start_date: str = Field(description="Window start day, YYYY-MM-DD UTC, inclusive")
|
||||
end_date: str = Field(description="Window end day, YYYY-MM-DD UTC, inclusive")
|
||||
|
|
|
|||
|
|
@ -72622,6 +72622,45 @@
|
|||
"supports_audio_input": true,
|
||||
"supports_video_input": true
|
||||
},
|
||||
"laya/english": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "laya",
|
||||
"mode": "evaluation",
|
||||
"output_cost_per_token": 0.0,
|
||||
"source": "https://github.com/NandhaKishorM/laya",
|
||||
"supported_endpoints": [
|
||||
"/v1/systemone"
|
||||
],
|
||||
"metadata": {
|
||||
"notes": "Self-hosted decision model; infrastructure costs are paid separately"
|
||||
}
|
||||
},
|
||||
"laya/multilingual": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "laya",
|
||||
"mode": "evaluation",
|
||||
"output_cost_per_token": 0.0,
|
||||
"source": "https://github.com/NandhaKishorM/laya",
|
||||
"supported_endpoints": [
|
||||
"/v1/systemone"
|
||||
],
|
||||
"metadata": {
|
||||
"notes": "Self-hosted decision model; infrastructure costs are paid separately"
|
||||
}
|
||||
},
|
||||
"laya/typed-decisions": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "laya",
|
||||
"mode": "evaluation",
|
||||
"output_cost_per_token": 0.0,
|
||||
"source": "https://github.com/NandhaKishorM/laya",
|
||||
"supported_endpoints": [
|
||||
"/v1/systemone"
|
||||
],
|
||||
"metadata": {
|
||||
"notes": "Self-hosted decision model; infrastructure costs are paid separately"
|
||||
}
|
||||
},
|
||||
"typesafe/jev-1.13.0": {
|
||||
"input_cost_per_token": 4.2e-08,
|
||||
"litellm_provider": "typesafe",
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@
|
|||
"url": "Link to provider documentation",
|
||||
"endpoints": {
|
||||
"chat_completions": "Supports /chat/completions endpoint",
|
||||
"systemone": "Supports native System One typed decisions",
|
||||
"messages": "Supports /messages endpoint (Anthropic format)",
|
||||
"responses": "Supports /responses endpoint (OpenAI/Anthropic unified)",
|
||||
"embeddings": "Supports /embeddings endpoint",
|
||||
|
|
@ -1476,6 +1477,13 @@
|
|||
"rerank": false
|
||||
}
|
||||
},
|
||||
"laya": {
|
||||
"display_name": "Laya (`laya`)",
|
||||
"url": "https://docs.litellm.ai/docs/auto_router/decision_classifiers",
|
||||
"endpoints": {
|
||||
"systemone": true
|
||||
}
|
||||
},
|
||||
"lambda_ai": {
|
||||
"display_name": "Lambda AI (`lambda_ai`)",
|
||||
"url": "https://docs.litellm.ai/docs/providers/lambda_ai",
|
||||
|
|
@ -3354,6 +3362,13 @@
|
|||
"provider_json_field": "skills",
|
||||
"url": "https://docs.litellm.ai/docs/skills"
|
||||
},
|
||||
"systemone": {
|
||||
"docs_label": "systemone",
|
||||
"display_name": "System One Decision API",
|
||||
"leftnav_label": "/laya/v1/systemone",
|
||||
"provider_json_field": "systemone",
|
||||
"url": "https://docs.litellm.ai/docs/auto_router/decision_classifiers"
|
||||
},
|
||||
"text_completion": {
|
||||
"docs_label": "text_completion",
|
||||
"display_name": "OpenAI Completions API",
|
||||
|
|
|
|||
|
|
@ -1744,6 +1744,27 @@ model LiteLLM_AutoRouterUserSession {
|
|||
@@index([user_id, last_turn_at], map: "idx_autorouter_user_session_user_last_turn")
|
||||
}
|
||||
|
||||
// Auto-routed requests per UTC request day and router: the selected-day money behind the
|
||||
// auto-router usage view. Written in the same statement as the session rollup, so a day row
|
||||
// and its session row never disagree; corrected in the same transaction as late baselines.
|
||||
model LiteLLM_AutoRouterDailySpend {
|
||||
date String
|
||||
api_key String
|
||||
user_id String
|
||||
router_name String
|
||||
router_type String
|
||||
turns Int @default(0)
|
||||
spend Float @default(0)
|
||||
saved_spend Float @default(0)
|
||||
savings_estimated_turns Int @default(0)
|
||||
savings_estimated_actual_spend Float @default(0)
|
||||
savings_estimated_saved_spend Float @default(0)
|
||||
classifier_cost Float @default(0)
|
||||
classifier_cost_recorded_turns Int @default(0)
|
||||
|
||||
@@id([date, api_key, user_id, router_name, router_type])
|
||||
}
|
||||
|
||||
// Shadow eval: evaluation of an auto-router against one or more keys' live traffic, in
|
||||
// either direction. forward duplicates the requests the keys did not route through the
|
||||
// router through it, answering whether they should adopt it; reverse duplicates the
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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] = []
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
108
tests/e2e/other/owned_jwt_gateway.py
Normal file
108
tests/e2e/other/owned_jwt_gateway.py
Normal file
|
|
@ -0,0 +1,108 @@
|
|||
"""An owned, source-built proxy whose `litellm_jwtauth` block a test controls.
|
||||
|
||||
The shared proxy on :4000 runs the CONTRIBUTING.md JWT block, so a test that
|
||||
needs a different `litellm_jwtauth` config boots its own gateway on a free port
|
||||
against the same database and the same Keycloak realm. The caller supplies the
|
||||
`litellm_jwtauth` mapping verbatim, which is exactly what makes a config an
|
||||
unfixed proxy rejects observable as a boot failure in this gateway's own log.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
from collections.abc import Mapping
|
||||
from contextlib import ExitStack
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
from e2e_config import INHERITED_ENV_PREFIXES, available_port
|
||||
from e2e_http import NoBody
|
||||
from idp import Keycloak, stop_process_group
|
||||
from proxy_client import ProxyClient, build_proxy_client
|
||||
|
||||
MODEL_NAME: Final = "gemini-3.8-flash"
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class OwnedJwtGateway:
|
||||
base_url: str
|
||||
proxy: ProxyClient
|
||||
_environment: Mapping[str, str] = field(repr=False)
|
||||
_command: tuple[str, ...] = field(repr=False)
|
||||
_log_path: Path
|
||||
_child: subprocess.Popen[bytes] | None = field(default=None, init=False, repr=False)
|
||||
|
||||
def start(self) -> None:
|
||||
with self._log_path.open("ab") as log:
|
||||
self._child = subprocess.Popen(
|
||||
self._command,
|
||||
env=self._environment,
|
||||
stdout=log,
|
||||
stderr=log,
|
||||
start_new_session=True,
|
||||
)
|
||||
deadline: Final = time.monotonic() + 120
|
||||
while time.monotonic() < deadline:
|
||||
assert self._child.poll() is None, "owned JWT gateway exited; inspect its private log"
|
||||
result = self.proxy.transport.probe("/health/liveliness", params=NoBody())
|
||||
if result.status_code == 200:
|
||||
return
|
||||
time.sleep(0.5)
|
||||
raise AssertionError("owned JWT gateway did not become ready")
|
||||
|
||||
def stop(self) -> None:
|
||||
if self._child is not None:
|
||||
stop_process_group(self._child)
|
||||
assert self._child.poll() is not None, "old gateway process is still alive"
|
||||
|
||||
|
||||
def owned_jwt_gateway(
|
||||
idp: Keycloak, directory: Path, cleanup: ExitStack, *, litellm_jwtauth: str, name: str
|
||||
) -> OwnedJwtGateway:
|
||||
for env_name in ("DATABASE_URL", "LITELLM_LICENSE", "LITELLM_MASTER_KEY"):
|
||||
assert os.environ.get(env_name), f"{env_name} is required for the owned JWT gateway"
|
||||
port: Final = available_port()
|
||||
base_url: Final = f"http://127.0.0.1:{port}"
|
||||
config: Final = directory / f"{name}.yaml"
|
||||
config.write_text(
|
||||
"model_list:\n"
|
||||
f" - model_name: {MODEL_NAME}\n"
|
||||
" litellm_params:\n"
|
||||
f" model: gemini/{MODEL_NAME}\n"
|
||||
" api_key: os.environ/GEMINI_API_KEY\n"
|
||||
"general_settings:\n"
|
||||
" master_key: os.environ/LITELLM_MASTER_KEY\n"
|
||||
" database_url: os.environ/DATABASE_URL\n"
|
||||
" proxy_batch_write_at: 5\n"
|
||||
" enable_jwt_auth: true\n"
|
||||
" litellm_jwtauth:\n" + "".join(f" {line}\n" for line in litellm_jwtauth.strip().splitlines())
|
||||
)
|
||||
environment: Final = {
|
||||
**{key: value for key, value in os.environ.items() if not key.startswith(INHERITED_ENV_PREFIXES)},
|
||||
"JWT_PUBLIC_KEY_URL": idp.jwks_url,
|
||||
"JWT_ISSUER": idp.issuer,
|
||||
"JWT_AUDIENCE": "litellm-e2e",
|
||||
"LITELLM_DANGEROUSLY_PERMIT_WEAK_OR_UNSET_MASTER_KEY": "true",
|
||||
"DISABLE_SCHEMA_UPDATE": "true",
|
||||
"STORE_MODEL_IN_DB": "True",
|
||||
"PYTHONPATH": str(Path(__file__).resolve().parents[3]),
|
||||
}
|
||||
gateway: Final = OwnedJwtGateway(
|
||||
base_url=base_url,
|
||||
proxy=build_proxy_client(
|
||||
base_url=base_url,
|
||||
control_plane_base_url=base_url,
|
||||
replica_urls=(base_url,),
|
||||
master_key=os.environ["LITELLM_MASTER_KEY"],
|
||||
),
|
||||
_environment=environment,
|
||||
_command=(sys.executable, "-m", "litellm.proxy.proxy_cli", "--config", str(config), "--port", str(port)),
|
||||
_log_path=directory / f"{name}.log",
|
||||
)
|
||||
cleanup.callback(gateway.stop)
|
||||
gateway.start()
|
||||
return gateway
|
||||
182
tests/e2e/other/test_jwt_auto_register_e2e.py
Normal file
182
tests/e2e/other/test_jwt_auto_register_e2e.py
Normal file
|
|
@ -0,0 +1,182 @@
|
|||
"""auto_register with auto_register_map_existing_key binds the JWT claim to the user's existing key.
|
||||
|
||||
`unregistered_jwt_client_behavior: auto_register` on `virtual_key_claim_field: sub` mints a fresh
|
||||
virtual key on the user's first JWT call. With `auto_register_map_existing_key: true` the proxy must
|
||||
instead point the new JWT mapping at a key the resolved user already owns, and mint only when the
|
||||
user has none. Each behavior gets its own gateway because the flag lives in `litellm_jwtauth`, so
|
||||
this file boots two owned proxies against the shared database and Keycloak realm.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
from collections.abc import Iterator
|
||||
from contextlib import ExitStack
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import unwrap
|
||||
from idp import Identity, Keycloak
|
||||
from lifecycle import ResourceManager
|
||||
from models import ChatBody, ChatMessage, JwtKeyMappingRow, KeyGenerateBody, TeamNewBody, UserNewBody
|
||||
from other_client import OtherClient
|
||||
from owned_jwt_gateway import MODEL_NAME, OwnedJwtGateway, owned_jwt_gateway
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
_JWT_COMMON: Final = (
|
||||
"user_id_jwt_field: sub\n"
|
||||
"user_email_jwt_field: email\n"
|
||||
"team_ids_jwt_field: groups\n"
|
||||
"user_id_upsert: true\n"
|
||||
"virtual_key_claim_field: sub\n"
|
||||
"unregistered_jwt_client_behavior: auto_register"
|
||||
)
|
||||
|
||||
|
||||
def _key_hash(key: str) -> str:
|
||||
return hashlib.sha256(key.encode()).hexdigest()
|
||||
|
||||
|
||||
def _ping() -> ChatBody:
|
||||
return ChatBody(
|
||||
model=MODEL_NAME,
|
||||
messages=[ChatMessage(role="user", content=f"Reply with the single word ok. {unique_marker()}")],
|
||||
max_tokens=5,
|
||||
)
|
||||
|
||||
|
||||
def _identity_with_user(idp: Keycloak, client: OtherClient, resources: ResourceManager) -> Identity:
|
||||
"""An IdP identity plus the litellm user and team its claims resolve to, with
|
||||
teardown that also sweeps the user's keys and JWT mapping rows the proxy
|
||||
wrote, since those outlive the user row itself."""
|
||||
marker: Final = unique_marker()
|
||||
identity: Final = idp.provision(marker=marker, group=f"e2e-jwt-team-{marker}", defer=resources.defer)
|
||||
resources.defer(lambda: client.proxy.delete_user(identity.user_id))
|
||||
team_id: Final = client.proxy.create_team(TeamNewBody(team_alias=f"e2e-jwt-{marker}", team_id=identity.group))
|
||||
resources.defer(lambda: client.proxy.delete_team(team_id))
|
||||
unwrap(
|
||||
client.user_new(
|
||||
UserNewBody(
|
||||
user_id=identity.user_id,
|
||||
user_email=f"{identity.username}@example.com",
|
||||
user_role="internal_user",
|
||||
auto_create_key=False,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
def delete_user_keys() -> None:
|
||||
for row in unwrap(client.user_info(identity.user_id)).keys:
|
||||
client.proxy.delete_key(row.token)
|
||||
|
||||
def delete_user_mappings() -> None:
|
||||
for mapping in unwrap(client.jwt_mapping_list()).mappings:
|
||||
if mapping.jwt_claim_value == identity.user_id:
|
||||
_ = client.jwt_mapping_delete(mapping.id)
|
||||
|
||||
resources.defer(delete_user_keys)
|
||||
resources.defer(delete_user_mappings)
|
||||
return identity
|
||||
|
||||
|
||||
def _mapping_for(client: OtherClient, claim_value: str) -> JwtKeyMappingRow | None:
|
||||
return next(
|
||||
(row for row in unwrap(client.jwt_mapping_list()).mappings if row.jwt_claim_value == claim_value),
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def mapping_gateway(idp: Keycloak, tmp_path_factory: pytest.TempPathFactory) -> Iterator[OwnedJwtGateway]:
|
||||
with ExitStack() as cleanup:
|
||||
yield owned_jwt_gateway(
|
||||
idp,
|
||||
tmp_path_factory.mktemp("jwt-mapping"),
|
||||
cleanup,
|
||||
litellm_jwtauth=f"{_JWT_COMMON}\nauto_register_map_existing_key: true",
|
||||
name="jwt-mapping-gateway",
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def minting_gateway(idp: Keycloak, tmp_path_factory: pytest.TempPathFactory) -> Iterator[OwnedJwtGateway]:
|
||||
with ExitStack() as cleanup:
|
||||
yield owned_jwt_gateway(
|
||||
idp,
|
||||
tmp_path_factory.mktemp("jwt-minting"),
|
||||
cleanup,
|
||||
litellm_jwtauth=_JWT_COMMON,
|
||||
name="jwt-minting-gateway",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.owned_gateway
|
||||
class TestJwtAutoRegisterMapExistingKey:
|
||||
@pytest.mark.covers("other.auth.jwt.auto_register_maps_existing_key")
|
||||
def test_first_jwt_call_maps_to_the_users_existing_key_and_mints_none(
|
||||
self, client: OtherClient, idp: Keycloak, resources: ResourceManager, mapping_gateway: OwnedJwtGateway
|
||||
) -> None:
|
||||
identity: Final = _identity_with_user(idp, client, resources)
|
||||
existing_key: Final = client.proxy.generate_key(
|
||||
KeyGenerateBody(
|
||||
user_id=identity.user_id, team_id=identity.group, key_alias=f"e2e-jwt-existing-{unique_marker()}"
|
||||
)
|
||||
)
|
||||
resources.defer(lambda: client.proxy.delete_key(existing_key))
|
||||
|
||||
response: Final = unwrap(mapping_gateway.proxy.chat(idp.access_token(identity), _ping()))
|
||||
|
||||
keys: Final = unwrap(client.user_info(identity.user_id)).keys
|
||||
assert [row.token for row in keys] == [_key_hash(existing_key)], (
|
||||
f"map_existing_key must leave the user with only their pre-existing key, got {keys}"
|
||||
)
|
||||
mapping: Final = _mapping_for(client, identity.user_id)
|
||||
assert mapping is not None, (
|
||||
f"no JWT mapping row for sub={identity.user_id}: {unwrap(client.jwt_mapping_list())}"
|
||||
)
|
||||
assert mapping.jwt_claim_name == "sub", f"mapping must bind the sub claim, got {mapping}"
|
||||
assert mapping.created_by == "auto_register", f"mapping must be written by auto_register, got {mapping}"
|
||||
rows: Final = client.proxy.poll_logs_for_key(existing_key)
|
||||
assert any(row.request_id == response.id for row in rows), (
|
||||
f"the JWT chat must be billed to the user's existing key, spend rows for it: {rows}"
|
||||
)
|
||||
|
||||
@pytest.mark.covers("other.auth.jwt.auto_register_mints_when_keyless")
|
||||
def test_first_jwt_call_mints_a_key_when_the_user_has_none(
|
||||
self, client: OtherClient, idp: Keycloak, resources: ResourceManager, mapping_gateway: OwnedJwtGateway
|
||||
) -> None:
|
||||
identity: Final = _identity_with_user(idp, client, resources)
|
||||
|
||||
response: Final = unwrap(mapping_gateway.proxy.chat(idp.access_token(identity), _ping()))
|
||||
assert response.choices, f"JWT chat returned no completion: {response}"
|
||||
|
||||
keys: Final = unwrap(client.user_info(identity.user_id)).keys
|
||||
assert len(keys) == 1, f"a keyless user must get exactly one minted key, got {keys}"
|
||||
mapping: Final = _mapping_for(client, identity.user_id)
|
||||
assert mapping is not None and mapping.jwt_claim_name == "sub", (
|
||||
f"the minted key must be recorded as a sub-claim mapping, mappings: {unwrap(client.jwt_mapping_list())}"
|
||||
)
|
||||
|
||||
@pytest.mark.covers("other.auth.jwt.auto_register_default_mints")
|
||||
def test_default_behavior_still_mints_when_the_user_already_has_a_key(
|
||||
self, client: OtherClient, idp: Keycloak, resources: ResourceManager, minting_gateway: OwnedJwtGateway
|
||||
) -> None:
|
||||
identity: Final = _identity_with_user(idp, client, resources)
|
||||
existing_key: Final = client.proxy.generate_key(
|
||||
KeyGenerateBody(user_id=identity.user_id, key_alias=f"e2e-jwt-existing-{unique_marker()}")
|
||||
)
|
||||
resources.defer(lambda: client.proxy.delete_key(existing_key))
|
||||
|
||||
response: Final = unwrap(minting_gateway.proxy.chat(idp.access_token(identity), _ping()))
|
||||
assert response.id is not None, f"JWT chat returned no response id: {response}"
|
||||
|
||||
keys: Final = unwrap(client.user_info(identity.user_id)).keys
|
||||
assert len(keys) == 2, (
|
||||
f"default auto_register must mint a second key for a user who already has one, got {keys}"
|
||||
)
|
||||
rows: Final = client.proxy.poll_logs_for_request_id(response.id)
|
||||
assert rows and all(row.api_key != _key_hash(existing_key) for row in rows), (
|
||||
f"the default path must bill the minted key, not the user's existing one: {rows}"
|
||||
)
|
||||
|
|
@ -16,6 +16,7 @@ markers =
|
|||
quiet_stack: measures the proxy itself, so it runs while no other test on this host is hitting the stack; every other test waits for it to finish
|
||||
mcp_oauth_live: real Linear OAuth consent via a captured browser session; deselected unless E2E_MCP_OAUTH_LIVE is set
|
||||
provider_edge_host: routes provider traffic through the pytest host's edge in every fixture mode, so the gateway must reach the pytest host; deselected unless E2E_PROVIDER_EDGE_HOST_REACHABLE is set
|
||||
owned_gateway: boots its own proxy from source against the stack's Postgres, so it needs DATABASE_URL on the pytest host; deselected unless E2E_OWNED_GATEWAY is set
|
||||
otel_v2: needs a proxy running with LITELLM_OTEL_V2=true; deselected unless E2E_OTEL_V2 is set
|
||||
otel_tls: needs a stack whose gateway exports OTLP over TLS signed by the CA in SSL_CERT_FILE; deselected unless E2E_OTEL_EXPORTER_ENDPOINT is set
|
||||
secret_manager: needs a proxy booted from gateway/secret_manager_<system>_ci_config.yml against that live secret manager; deselected unless E2E_SECRET_MANAGER names the backend (see secret_manager/secret_backends.py)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -124,6 +124,20 @@ def _trace_spans(sink_url: str, trace_id: str, seconds: float = 30) -> tuple[Spa
|
|||
return group
|
||||
|
||||
|
||||
def _trace_spans_when(
|
||||
sink_url: str,
|
||||
trace_id: str,
|
||||
ready: Callable[[tuple[Span, ...]], bool],
|
||||
seconds: float = 30,
|
||||
) -> tuple[Span, ...]:
|
||||
spans: Final = eventually(
|
||||
lambda: spans_for_trace(recorded_spans(sink_url)[1], trace_id),
|
||||
ready,
|
||||
seconds=seconds,
|
||||
)
|
||||
return spans
|
||||
|
||||
|
||||
def _await_db_span(sink_url: str, trace_id: str | None, needle: str, seconds: float = 40, since: int = 0) -> None:
|
||||
def seen() -> bool:
|
||||
_, spans = recorded_spans(sink_url, since)
|
||||
|
|
@ -242,6 +256,196 @@ def test_without_excluded_services_the_tenant_still_gets_redis_and_postgres_span
|
|||
assert {"redis", "postgresql"} <= systems, f"datastore spans missing at tenant: {systems}"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("otel", [None, True, "on", "", []], ids=["null", "true", "on", "empty_string", "empty_list"])
|
||||
@pytest.mark.timeout(180)
|
||||
def test_a_non_mapping_otel_block_still_publishes_the_tenant_fan_out(
|
||||
gateway: Gateway,
|
||||
audit_sinks: SpanSinks,
|
||||
otel_audit_config: AuditConfigWriter,
|
||||
langfuse_vars: dict[str, JsonValue],
|
||||
tmp_path: Path,
|
||||
otel: JsonValue,
|
||||
) -> None:
|
||||
def with_callback_settings(config: dict) -> None:
|
||||
config["litellm_settings"]["callbacks"] = ["langfuse_otel"]
|
||||
config["callback_settings"]["otel"] = otel
|
||||
|
||||
config: Final = _config_with(tmp_path, otel_audit_config, extra=with_callback_settings)
|
||||
overrides: Final = {"LITELLM_OTEL_V2": "1", **_operator_langfuse(audit_sinks)}
|
||||
with owned_proxy(gateway, tmp_path, overrides, config=config, workers=2) as candidate:
|
||||
traffic: Final = _drive(candidate, langfuse_vars)
|
||||
tenant_trace: Final = _trace_id(audit_sinks.tenant, traffic)
|
||||
_await_db_span(audit_sinks.tenant, tenant_trace, "redis")
|
||||
tenant_spans: Final = _trace_spans_when(
|
||||
audit_sinks.tenant,
|
||||
tenant_trace,
|
||||
lambda spans: any(span["kind"] == 2 for span in spans) and "redis" in _db_systems(spans),
|
||||
seconds=15,
|
||||
)
|
||||
assert any(span["kind"] == 2 for span in tenant_spans), "tenant SERVER root span missing"
|
||||
assert "redis" in _db_systems(tenant_spans), f"tenant redis span missing: {_db_systems(tenant_spans)}"
|
||||
operator_trace: Final = _trace_id(audit_sinks.operator, traffic)
|
||||
operator_spans: Final = _trace_spans_when(
|
||||
audit_sinks.operator,
|
||||
operator_trace,
|
||||
lambda spans: any(span["kind"] == 2 for span in spans),
|
||||
seconds=15,
|
||||
)
|
||||
assert any(span["kind"] == 2 for span in operator_spans), "operator SERVER root span missing"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("name", ["EXCLUDED_SERVICES", "excluded_services"])
|
||||
@pytest.mark.timeout(180)
|
||||
def test_a_bare_excluded_services_env_var_is_ignored(
|
||||
gateway: Gateway,
|
||||
audit_sinks: SpanSinks,
|
||||
otel_audit_config: AuditConfigWriter,
|
||||
langfuse_vars: dict[str, JsonValue],
|
||||
tmp_path: Path,
|
||||
name: str,
|
||||
) -> None:
|
||||
config: Final = _config_with(tmp_path, otel_audit_config)
|
||||
overrides: Final = {"LITELLM_OTEL_V2": "1", name: "redis,postgres"}
|
||||
with owned_proxy(gateway, tmp_path, overrides, config=config, workers=2) as candidate:
|
||||
tenant_start, _ = recorded_spans(audit_sinks.tenant)
|
||||
traffic: Final = _drive(candidate, langfuse_vars)
|
||||
tenant_trace: Final = _trace_id(audit_sinks.tenant, traffic)
|
||||
_await_db_span(audit_sinks.tenant, tenant_trace, "redis")
|
||||
tenant_spans: Final = _trace_spans_when(
|
||||
audit_sinks.tenant,
|
||||
tenant_trace,
|
||||
lambda spans: "redis" in _db_systems(spans),
|
||||
seconds=15,
|
||||
)
|
||||
assert "redis" in _db_systems(tenant_spans), f"redis span missing at tenant: {_db_systems(tenant_spans)}"
|
||||
_await_db_span(audit_sinks.tenant, None, "postgresql", since=tenant_start)
|
||||
_, all_tenant = recorded_spans(audit_sinks.tenant, tenant_start)
|
||||
systems: Final = _db_systems(all_tenant)
|
||||
assert {"redis", "postgresql"} <= systems, f"datastore spans missing at tenant: {systems}"
|
||||
|
||||
|
||||
@pytest.mark.timeout(180)
|
||||
def test_the_documented_env_var_wins_over_a_bare_excluded_services(
|
||||
gateway: Gateway,
|
||||
audit_sinks: SpanSinks,
|
||||
otel_audit_config: AuditConfigWriter,
|
||||
langfuse_vars: dict[str, JsonValue],
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
config: Final = _config_with(tmp_path, otel_audit_config)
|
||||
overrides: Final = {
|
||||
"LITELLM_OTEL_V2": "1",
|
||||
"LITELLM_OTEL_EXCLUDED_SERVICES": "redis",
|
||||
"EXCLUDED_SERVICES": "postgres",
|
||||
}
|
||||
with owned_proxy(gateway, tmp_path, overrides, config=config, workers=2) as candidate:
|
||||
tenant_start, _ = recorded_spans(audit_sinks.tenant)
|
||||
traffic: Final = _drive(candidate, langfuse_vars)
|
||||
tenant_trace: Final = _trace_id(audit_sinks.tenant, traffic)
|
||||
_trace_spans(audit_sinks.tenant, tenant_trace, seconds=15)
|
||||
_await_db_span(audit_sinks.tenant, None, "postgresql", since=tenant_start)
|
||||
_, tenant_spans = recorded_spans(audit_sinks.tenant, tenant_start)
|
||||
systems: Final = _db_systems(tenant_spans)
|
||||
assert "postgresql" in systems, f"postgresql spans missing at tenant: {systems}"
|
||||
assert "redis" not in systems, f"redis spans reached tenant: {systems}"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("env_name", "redis_reaches_tenant"),
|
||||
[
|
||||
pytest.param("LITELLM_OTEL_EXCLUDED_SERVICES", False, id="exact-case"),
|
||||
pytest.param("litellm_otel_excluded_services", True, id="wrong-case"),
|
||||
],
|
||||
)
|
||||
@pytest.mark.timeout(180)
|
||||
def test_case_sensitive_otel_settings_read_only_the_exact_env_name(
|
||||
gateway: Gateway,
|
||||
audit_sinks: SpanSinks,
|
||||
otel_audit_config: AuditConfigWriter,
|
||||
langfuse_vars: Mapping[str, JsonValue],
|
||||
tmp_path: Path,
|
||||
env_name: str,
|
||||
redis_reaches_tenant: bool,
|
||||
) -> None:
|
||||
config: Final = _config_with(tmp_path, otel_audit_config, otel={"_case_sensitive": True})
|
||||
overrides: Final = {"LITELLM_OTEL_V2": "1", env_name: "redis"}
|
||||
with owned_proxy(
|
||||
gateway,
|
||||
tmp_path,
|
||||
overrides,
|
||||
config=config,
|
||||
remove_environment=("LITELLM_OTEL_EXCLUDED_SERVICES", "litellm_otel_excluded_services"),
|
||||
workers=2,
|
||||
) as candidate:
|
||||
tenant_start, _ = recorded_spans(audit_sinks.tenant)
|
||||
traffic: Final = _drive(candidate, langfuse_vars)
|
||||
tenant_trace: Final = _trace_id(audit_sinks.tenant, traffic)
|
||||
if redis_reaches_tenant:
|
||||
_await_db_span(audit_sinks.tenant, tenant_trace, "redis")
|
||||
_trace_spans_when(
|
||||
audit_sinks.tenant,
|
||||
tenant_trace,
|
||||
lambda spans: "redis" in _db_systems(spans),
|
||||
seconds=15,
|
||||
)
|
||||
else:
|
||||
_trace_spans(audit_sinks.tenant, tenant_trace, seconds=15)
|
||||
_await_db_span(audit_sinks.tenant, None, "postgresql", seconds=60, since=tenant_start)
|
||||
_, tenant_spans = recorded_spans(audit_sinks.tenant, tenant_start)
|
||||
systems: Final = _db_systems(tenant_spans)
|
||||
assert "postgresql" in systems, f"postgresql spans missing at tenant: {systems}"
|
||||
assert ("redis" in systems) is redis_reaches_tenant, (
|
||||
f"tenant redis presence={('redis' in systems)}; expected={redis_reaches_tenant}; systems={systems}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.timeout(180)
|
||||
def test_env_ignore_empty_keeps_the_default_service_name(
|
||||
gateway: Gateway,
|
||||
audit_sinks: SpanSinks,
|
||||
otel_audit_config: AuditConfigWriter,
|
||||
langfuse_vars: Mapping[str, JsonValue],
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
config: Final = _config_with(tmp_path, otel_audit_config, otel={"_env_ignore_empty": True})
|
||||
overrides: Final = {"LITELLM_OTEL_V2": "1", "OTEL_SERVICE_NAME": ""}
|
||||
with owned_proxy(gateway, tmp_path, overrides, config=config, workers=2) as candidate:
|
||||
traffic: Final = _drive(candidate, langfuse_vars)
|
||||
operator_trace: Final = _trace_id(audit_sinks.operator, traffic)
|
||||
operator_spans: Final = _trace_spans_when(
|
||||
audit_sinks.operator,
|
||||
operator_trace,
|
||||
lambda spans: any(span["kind"] == 2 for span in spans),
|
||||
seconds=15,
|
||||
)
|
||||
service_names: Final = tuple(span["resource"].get("service.name") for span in operator_spans)
|
||||
assert service_names and all(name == "litellm" for name in service_names), (
|
||||
f"operator service.name values={service_names}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.timeout(180)
|
||||
def test_env_parse_none_str_reads_a_null_traces_endpoint_as_unset(
|
||||
gateway: Gateway,
|
||||
audit_sinks: SpanSinks,
|
||||
otel_audit_config: AuditConfigWriter,
|
||||
langfuse_vars: Mapping[str, JsonValue],
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
config: Final = _config_with(tmp_path, otel_audit_config, otel={"_env_parse_none_str": "null"})
|
||||
overrides: Final = {"LITELLM_OTEL_V2": "1", "OTEL_TRACES_ENDPOINT": "null"}
|
||||
with owned_proxy(gateway, tmp_path, overrides, config=config, workers=2) as candidate:
|
||||
traffic: Final = _drive(candidate, langfuse_vars)
|
||||
operator_trace: Final = _trace_id(audit_sinks.operator, traffic)
|
||||
operator_spans: Final = _trace_spans_when(
|
||||
audit_sinks.operator,
|
||||
operator_trace,
|
||||
lambda spans: any(span["kind"] == 2 for span in spans),
|
||||
seconds=15,
|
||||
)
|
||||
assert any(span["kind"] == 2 for span in operator_spans), "operator SERVER root span missing"
|
||||
|
||||
|
||||
def test_env_excluded_services_drops_only_redis(
|
||||
gateway: Gateway,
|
||||
audit_sinks: SpanSinks,
|
||||
|
|
|
|||
|
|
@ -9,7 +9,8 @@ from concurrent.futures import ThreadPoolExecutor
|
|||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Final, Literal
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal, Protocol, cast
|
||||
|
||||
import anthropic
|
||||
import httpx
|
||||
|
|
@ -26,6 +27,9 @@ from pydantic import JsonValue, TypeAdapter
|
|||
MARKER: Final = re.compile(rb"excl-[0-9a-f]{32}")
|
||||
FAILING: Final = re.compile(rb"excl-fail-[0-9a-f]{32}")
|
||||
JSON: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
|
||||
UNCONFIGURED_VARIANT: Final[TypeAdapter[Literal["null_block", "bare_env"]]] = TypeAdapter(
|
||||
Literal["null_block", "bare_env"]
|
||||
)
|
||||
REPLY_TEXT: Final = "excluded ok"
|
||||
SERVER: Final = 2
|
||||
INVALID_NAME_LOG: Final = "is not a datastore service"
|
||||
|
|
@ -37,6 +41,11 @@ CLIENTS: Final[tuple[Client, ...]] = ("raw", "sdk", "async_sdk")
|
|||
AuditConfigWriter = Callable[[Path, Mapping[str, JsonValue]], Path]
|
||||
|
||||
|
||||
class _FixtureRequestParam(Protocol):
|
||||
@property
|
||||
def param(self) -> object: ...
|
||||
|
||||
|
||||
def _marker() -> str:
|
||||
return "excl-" + uuid.uuid4().hex
|
||||
|
||||
|
|
@ -339,6 +348,11 @@ def _db_systems(spans: tuple[Span, ...]) -> set[str]:
|
|||
}
|
||||
|
||||
|
||||
def _post_auth_datastore_spans(spans: tuple[Span, ...]) -> tuple[Span, ...]:
|
||||
auth_ids: Final = frozenset(span["span_id"] for span in spans if span["name"].startswith("auth "))
|
||||
return tuple(span for span in spans if _db_systems((span,)) and span["parent_span_id"] not in auth_ids)
|
||||
|
||||
|
||||
def _names(spans: tuple[Span, ...]) -> list[str]:
|
||||
return sorted(span["name"] for span in spans)
|
||||
|
||||
|
|
@ -408,6 +422,29 @@ def _config(directory: Path, otel_audit_config: AuditConfigWriter, otel: Mapping
|
|||
return path
|
||||
|
||||
|
||||
def _null_otel_config(directory: Path, otel_audit_config: AuditConfigWriter, name: str) -> Path:
|
||||
written: Final = otel_audit_config(directory, {})
|
||||
loaded: Final = object_value(JSON.validate_python(yaml.safe_load(written.read_text())))
|
||||
config: Final = {
|
||||
**loaded,
|
||||
"litellm_settings": {**object_value(loaded["litellm_settings"]), "callbacks": ["langfuse_otel"]},
|
||||
"callback_settings": {**object_value(loaded["callback_settings"]), "otel": None},
|
||||
}
|
||||
path: Final = directory / f"{name}.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
return path
|
||||
|
||||
|
||||
def _operator_langfuse(sinks: SpanSinks) -> dict[str, str]:
|
||||
return {
|
||||
"LANGFUSE_HOST": sinks.operator,
|
||||
"LANGFUSE_PUBLIC_KEY": "pk-lf-operator",
|
||||
"LANGFUSE_SECRET_KEY": "sk-lf-operator",
|
||||
"OTEL_EXPORTER": "http/json",
|
||||
"OTEL_ENDPOINT": sinks.operator,
|
||||
}
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _started(
|
||||
provider: Wire,
|
||||
|
|
@ -416,13 +453,14 @@ def _started(
|
|||
directory: Path,
|
||||
langfuse_vars: Mapping[str, JsonValue],
|
||||
workers: int,
|
||||
environment: Mapping[str, str] = MappingProxyType({}),
|
||||
) -> Generator[Rig]:
|
||||
with (
|
||||
gateway_from_environment() as gateway,
|
||||
owned_proxy_process(
|
||||
gateway,
|
||||
directory,
|
||||
{"LITELLM_OTEL_V2": "1", "OTEL_BSP_SCHEDULE_DELAY": "300"},
|
||||
{"LITELLM_OTEL_V2": "1", "OTEL_BSP_SCHEDULE_DELAY": "300", **environment},
|
||||
config=config,
|
||||
remove_environment=("LITELLM_OTEL_EXCLUDED_SERVICES",),
|
||||
workers=workers,
|
||||
|
|
@ -458,6 +496,32 @@ def rig(
|
|||
yield started
|
||||
|
||||
|
||||
@pytest.fixture(scope="module", params=["null_block", "bare_env"], ids=["null_block", "bare_env"])
|
||||
def unconfigured_rig(
|
||||
request: pytest.FixtureRequest,
|
||||
provider: Wire,
|
||||
audit_sinks: SpanSinks,
|
||||
otel_audit_config: AuditConfigWriter,
|
||||
langfuse_vars: dict[str, JsonValue],
|
||||
tmp_path_factory: pytest.TempPathFactory,
|
||||
) -> Iterator[Rig]:
|
||||
parameter: Final = cast(_FixtureRequestParam, request).param
|
||||
variant: Final = UNCONFIGURED_VARIANT.validate_python(parameter)
|
||||
directory: Final = tmp_path_factory.mktemp(f"excluded-{variant}")
|
||||
config: Final = (
|
||||
_null_otel_config(directory, otel_audit_config, variant)
|
||||
if variant == "null_block"
|
||||
else _config(directory, otel_audit_config, {}, variant)
|
||||
)
|
||||
environment: Final = (
|
||||
_operator_langfuse(audit_sinks) if variant == "null_block" else {"EXCLUDED_SERVICES": "redis,postgres"}
|
||||
)
|
||||
with _started(
|
||||
provider, audit_sinks, config, directory, langfuse_vars, workers=2, environment=environment
|
||||
) as started:
|
||||
yield started
|
||||
|
||||
|
||||
@pytest.mark.timeout(120)
|
||||
@pytest.mark.parametrize("stream", [False, True], ids=["unary", "stream"])
|
||||
@pytest.mark.parametrize("client", CLIENTS)
|
||||
|
|
@ -654,6 +718,237 @@ def test_killing_one_of_two_workers_mid_burst_keeps_the_filter_on_the_survivor(r
|
|||
_assert_withheld(rig, rig.raw("chat", _marker(), stream=False), after)
|
||||
|
||||
|
||||
def _assert_tenant_kept(
|
||||
rig: Rig, trace_id: str, cursors: Cursors, *, needs_model_span: bool = False
|
||||
) -> tuple[Span, ...]:
|
||||
def ready(spans: tuple[Span, ...]) -> bool:
|
||||
return (
|
||||
sum(1 for span in spans if span["kind"] == SERVER) == 1
|
||||
and "redis" in _db_systems(spans)
|
||||
and (not needs_model_span or any("gen_ai.operation.name" in span["attributes"] for span in spans))
|
||||
)
|
||||
|
||||
tenant: Final = eventually(
|
||||
lambda: spans_for_trace(recorded_spans(rig.sinks.tenant, cursors.tenant)[1], trace_id),
|
||||
ready,
|
||||
seconds=40,
|
||||
return_last_on_timeout=True,
|
||||
)
|
||||
assert sum(1 for span in tenant if span["kind"] == SERVER) == 1, _names(tenant)
|
||||
assert "redis" in _db_systems(tenant), f"redis spans missing at the tenant: {_names(tenant)}"
|
||||
assert not needs_model_span or any("gen_ai.operation.name" in span["attributes"] for span in tenant), _names(tenant)
|
||||
return tenant
|
||||
|
||||
|
||||
def _assert_kept(rig: Rig, sent: Sent, cursors: Cursors) -> tuple[Span, ...]:
|
||||
operator: Final = _operator_trace(rig, sent, cursors)
|
||||
return _assert_tenant_kept(rig, operator[0]["trace_id"], cursors, needs_model_span=True)
|
||||
|
||||
|
||||
@pytest.mark.timeout(120)
|
||||
@pytest.mark.parametrize("stream", [False, True], ids=["unary", "stream"])
|
||||
@pytest.mark.parametrize("client", CLIENTS)
|
||||
@pytest.mark.parametrize("endpoint", ENDPOINTS)
|
||||
def test_unconfigured_tenant_trace_keeps_datastore_spans(
|
||||
unconfigured_rig: Rig, endpoint: Endpoint, client: Client, stream: bool
|
||||
) -> None:
|
||||
cursors: Final = unconfigured_rig.cursors()
|
||||
marker: Final = _marker()
|
||||
sent: Final = unconfigured_rig.send(endpoint, client, marker, stream)
|
||||
assert sent.text == REPLY_TEXT, sent
|
||||
assert unconfigured_rig.upstream_hits(marker) == 1
|
||||
_assert_kept(unconfigured_rig, sent, cursors)
|
||||
|
||||
|
||||
@pytest.mark.timeout(120)
|
||||
@pytest.mark.parametrize("endpoint", ["chat", "messages"])
|
||||
def test_unconfigured_cache_hit_twin_keeps_datastore_spans(unconfigured_rig: Rig, endpoint: Endpoint) -> None:
|
||||
cursors: Final = unconfigured_rig.cursors()
|
||||
marker: Final = _marker()
|
||||
first_result: Final = _traced_raw(unconfigured_rig, endpoint, marker)
|
||||
first: Final = first_result[1]
|
||||
assert first.text == REPLY_TEXT, first
|
||||
assert unconfigured_rig.upstream_hits(marker) == 1
|
||||
_assert_kept(unconfigured_rig, first, cursors)
|
||||
hit_cursors: Final = unconfigured_rig.cursors()
|
||||
|
||||
def read_hit() -> tuple[str, Sent, tuple[Span, ...]]:
|
||||
trace_id, sent = _traced_raw(unconfigured_rig, endpoint, marker)
|
||||
return trace_id, sent, _operator_trace_by_id(unconfigured_rig, trace_id, hit_cursors)
|
||||
|
||||
trace_id, hit, operator = eventually(
|
||||
read_hit,
|
||||
lambda result: (
|
||||
unconfigured_rig.upstream_hits(marker) == 0
|
||||
and "redis" in _db_systems(_post_auth_datastore_spans(result[2]))
|
||||
),
|
||||
seconds=60,
|
||||
)
|
||||
assert hit.text == REPLY_TEXT, hit
|
||||
post_auth_datastore: Final = _post_auth_datastore_spans(operator)
|
||||
post_auth_span_ids: Final = frozenset(span["span_id"] for span in post_auth_datastore)
|
||||
non_datastore_names: Final = frozenset(span["name"] for span in operator if not _db_systems((span,)))
|
||||
tenant: Final = eventually(
|
||||
lambda: spans_for_trace(recorded_spans(unconfigured_rig.sinks.tenant, hit_cursors.tenant)[1], trace_id),
|
||||
lambda spans: (
|
||||
non_datastore_names <= frozenset(span["name"] for span in spans)
|
||||
and post_auth_span_ids <= frozenset(span["span_id"] for span in spans)
|
||||
),
|
||||
seconds=40,
|
||||
return_last_on_timeout=True,
|
||||
)
|
||||
tenant_span_ids: Final = frozenset(span["span_id"] for span in tenant)
|
||||
missing_post_auth_names: Final = tuple(
|
||||
span["name"] for span in post_auth_datastore if span["span_id"] not in tenant_span_ids
|
||||
)
|
||||
assert sum(1 for span in tenant if span["kind"] == SERVER) == 1, _names(tenant)
|
||||
assert "redis" in _db_systems(tenant), (
|
||||
f"operator datastore systems={sorted(_db_systems(post_auth_datastore))}; "
|
||||
f"tenant datastore systems={sorted(_db_systems(tenant))}; tenant spans={_names(tenant)}"
|
||||
)
|
||||
assert not missing_post_auth_names, (
|
||||
f"missing post-auth datastore span names={missing_post_auth_names}; "
|
||||
f"operator={_names(post_auth_datastore)}; tenant={_names(tenant)}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.timeout(120)
|
||||
@pytest.mark.parametrize("endpoint", ENDPOINTS)
|
||||
def test_unconfigured_failed_upstream_keeps_datastore_spans(unconfigured_rig: Rig, endpoint: Endpoint) -> None:
|
||||
cursors: Final = unconfigured_rig.cursors()
|
||||
marker: Final = "excl-fail-" + uuid.uuid4().hex
|
||||
trace_id: Final = uuid.uuid4().hex
|
||||
path, body = _body(unconfigured_rig.model, endpoint, marker, stream=False)
|
||||
failed: Final = unconfigured_rig.proxy.client.post(
|
||||
path,
|
||||
json=body,
|
||||
headers={
|
||||
"Authorization": f"Bearer {unconfigured_rig.key}",
|
||||
"traceparent": f"00-{trace_id}-{uuid.uuid4().hex[:16]}-01",
|
||||
},
|
||||
)
|
||||
assert failed.status_code == 500, failed.text
|
||||
assert unconfigured_rig.upstream_hits(marker) >= 1
|
||||
operator: Final = eventually(
|
||||
lambda: spans_for_trace(recorded_spans(unconfigured_rig.sinks.operator, cursors.operator)[1], trace_id),
|
||||
lambda spans: _has_root(spans) and "redis" in _db_systems(spans),
|
||||
seconds=40,
|
||||
)
|
||||
assert "redis" in _db_systems(operator), _names(operator)
|
||||
_assert_tenant_kept(unconfigured_rig, trace_id, cursors)
|
||||
|
||||
|
||||
@pytest.mark.timeout(120)
|
||||
@pytest.mark.parametrize("status", [403, 404])
|
||||
def test_unconfigured_rejecting_tenant_destination_recovers(unconfigured_rig: Rig, status: int) -> None:
|
||||
configure_sink(unconfigured_rig.sinks.tenant, status=status)
|
||||
try:
|
||||
cursors: Final = unconfigured_rig.cursors()
|
||||
marker: Final = _marker()
|
||||
sent: Final = unconfigured_rig.raw("chat", marker, stream=True)
|
||||
assert sent.text == REPLY_TEXT, sent
|
||||
assert unconfigured_rig.upstream_hits(marker) == 1
|
||||
_operator_trace(unconfigured_rig, sent, cursors)
|
||||
finally:
|
||||
configure_sink(unconfigured_rig.sinks.tenant, status=200)
|
||||
after: Final = unconfigured_rig.cursors()
|
||||
recovered: Final = unconfigured_rig.raw("responses", _marker(), stream=False)
|
||||
assert recovered.text == REPLY_TEXT, recovered
|
||||
_assert_kept(unconfigured_rig, recovered, after)
|
||||
|
||||
|
||||
@pytest.mark.timeout(120)
|
||||
def test_unconfigured_key_level_destination_keeps_datastore_spans(
|
||||
unconfigured_rig: Rig, langfuse_vars: dict[str, JsonValue]
|
||||
) -> None:
|
||||
key: Final = unconfigured_rig.scenario.key(
|
||||
metadata={
|
||||
"logging": [
|
||||
{"callback_name": "langfuse_otel", "callback_type": "success", "callback_vars": dict(langfuse_vars)}
|
||||
]
|
||||
}
|
||||
)
|
||||
cursors: Final = unconfigured_rig.cursors()
|
||||
marker: Final = _marker()
|
||||
sent: Final = unconfigured_rig.raw("chat", marker, stream=False, key=key)
|
||||
assert sent.text == REPLY_TEXT, sent
|
||||
assert unconfigured_rig.upstream_hits(marker) == 1
|
||||
_assert_kept(unconfigured_rig, sent, cursors)
|
||||
|
||||
|
||||
def _assert_tenant_kept_the_burst(rig: Rig, cursors: Cursors, traces: set[str]) -> None:
|
||||
def ready(spans: tuple[Span, ...]) -> bool:
|
||||
def trace_kept(trace: str) -> bool:
|
||||
trace_spans: Final = spans_for_trace(spans, trace)
|
||||
return any(span["kind"] == SERVER for span in trace_spans) and "redis" in _db_systems(trace_spans)
|
||||
|
||||
return all(trace_kept(trace) for trace in traces)
|
||||
|
||||
tenant: Final = eventually(
|
||||
lambda: recorded_spans(rig.sinks.tenant, cursors.tenant)[1],
|
||||
ready,
|
||||
seconds=90,
|
||||
return_last_on_timeout=True,
|
||||
)
|
||||
burst: Final = tuple(span for span in tenant if span["trace_id"] in traces)
|
||||
missing_roots: Final = tuple(
|
||||
trace for trace in traces if not any(span["kind"] == SERVER for span in spans_for_trace(burst, trace))
|
||||
)
|
||||
missing_redis: Final = tuple(trace for trace in traces if "redis" not in _db_systems(spans_for_trace(burst, trace)))
|
||||
assert not missing_roots, f"SERVER root missing from tenant burst traces: {missing_roots}, {_names(burst)}"
|
||||
assert not missing_redis, f"redis spans missing from tenant burst traces: {missing_redis}, {_names(burst)}"
|
||||
|
||||
|
||||
@pytest.mark.timeout(300)
|
||||
def test_unconfigured_tenant_outage_during_a_mixed_burst(unconfigured_rig: Rig) -> None:
|
||||
cursors: Final = unconfigured_rig.cursors()
|
||||
configure_sink(unconfigured_rig.sinks.tenant, status=503)
|
||||
try:
|
||||
results: Final = _burst(unconfigured_rig, 30)
|
||||
finally:
|
||||
configure_sink(unconfigured_rig.sinks.tenant, status=200)
|
||||
served: Final = _served(results)
|
||||
assert len(served) == 30, [result for result in results if isinstance(result, str)]
|
||||
assert all(sent.text == REPLY_TEXT for sent in served), served
|
||||
traces: Final = _assert_operator_exactly_once(unconfigured_rig, served, cursors)
|
||||
_assert_tenant_kept_the_burst(unconfigured_rig, cursors, traces)
|
||||
after: Final = unconfigured_rig.cursors()
|
||||
_assert_kept(unconfigured_rig, unconfigured_rig.raw("messages", _marker(), stream=True), after)
|
||||
|
||||
|
||||
@pytest.mark.timeout(300)
|
||||
def test_unconfigured_killing_one_of_two_workers_keeps_the_fan_out(unconfigured_rig: Rig) -> None:
|
||||
root: Final = psutil.Process(unconfigured_rig.owned.process.pid)
|
||||
workers: Final = eventually(
|
||||
lambda: tuple(child for child in root.children() if "resource_tracker" not in " ".join(child.cmdline())),
|
||||
lambda found: len(found) == 2,
|
||||
seconds=30,
|
||||
)
|
||||
cursors: Final = unconfigured_rig.cursors()
|
||||
|
||||
def one(index: int) -> Sent | str:
|
||||
if index == 6:
|
||||
os.kill(workers[0].pid, signal.SIGKILL)
|
||||
try:
|
||||
return unconfigured_rig.raw("chat", _marker(), stream=index % 2 == 0)
|
||||
except (httpx.HTTPError, AssertionError) as error:
|
||||
return repr(error)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=6) as pool:
|
||||
results: Final = tuple(pool.map(one, range(18)))
|
||||
assert unconfigured_rig.owned.process.poll() is None, "Proxy root exited after a worker was killed"
|
||||
failures: Final = tuple(result for result in results if isinstance(result, str))
|
||||
assert all(failure.startswith(("ReadError(", "RemoteProtocolError(", "ConnectError(")) for failure in failures), (
|
||||
failures
|
||||
)
|
||||
assert len(failures) <= 6, failures
|
||||
settled: Final = tuple(result for index, result in enumerate(results) if index > 12 and isinstance(result, Sent))
|
||||
traces: Final = _assert_operator_exactly_once(unconfigured_rig, settled, cursors)
|
||||
_assert_tenant_kept_the_burst(unconfigured_rig, cursors, traces)
|
||||
after: Final = unconfigured_rig.cursors()
|
||||
_assert_kept(unconfigured_rig, unconfigured_rig.raw("chat", _marker(), stream=False), after)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Setting:
|
||||
otel: Mapping[str, JsonValue]
|
||||
|
|
|
|||
421
tests/integration/spend/test_passthrough_request_tags.py
Normal file
421
tests/integration/spend/test_passthrough_request_tags.py
Normal file
|
|
@ -0,0 +1,421 @@
|
|||
import json
|
||||
import uuid
|
||||
from collections.abc import Callable, Mapping
|
||||
from hashlib import sha256
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from integration._support.client import Gateway, JsonValue, Scenario, eventually, object_value
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.process import owned_proxy
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
|
||||
def _chat_reply(marker: str) -> dict[str, JsonValue]:
|
||||
return {
|
||||
"id": f"chatcmpl-{marker}",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": "gpt-4o-mini",
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": marker}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8},
|
||||
}
|
||||
|
||||
|
||||
def _anthropic_reply(marker: str) -> dict[str, JsonValue]:
|
||||
return {
|
||||
"id": f"msg_{marker}",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-sonnet-4-5",
|
||||
"content": [{"type": "text", "text": marker}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 5, "output_tokens": 3},
|
||||
}
|
||||
|
||||
|
||||
def _chat_stream_frames(marker: str) -> tuple[bytes, ...]:
|
||||
chunk: Final = {"id": f"chatcmpl-{marker}", "object": "chat.completion.chunk", "created": 1, "model": "gpt-4o-mini"}
|
||||
return (
|
||||
f"data: {json.dumps({**chunk, 'choices': [{'index': 0, 'delta': {'role': 'assistant', 'content': marker}}]})}\n\n".encode(),
|
||||
f"data: {json.dumps({**chunk, 'choices': [{'index': 0, 'delta': {}, 'finish_reason': 'stop'}], 'usage': {'prompt_tokens': 5, 'completion_tokens': 3, 'total_tokens': 8}})}\n\n".encode(),
|
||||
b"data: [DONE]\n\n",
|
||||
)
|
||||
|
||||
|
||||
def _spend_row(digest: str, call_type: str) -> dict[str, JsonValue]:
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT request_tags, metadata, team_id FROM "LiteLLM_SpendLogs" WHERE api_key=%s AND call_type=%s',
|
||||
(digest, call_type),
|
||||
),
|
||||
lambda values: len(values) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
return rows[0]
|
||||
|
||||
|
||||
def _spend_row_tagged(tag: str) -> dict[str, JsonValue]:
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT request_tags, metadata, team_id, api_key FROM "LiteLLM_SpendLogs" WHERE request_tags::text LIKE %s',
|
||||
(f'%"{tag}"%',),
|
||||
),
|
||||
lambda values: len(values) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
return rows[0]
|
||||
|
||||
|
||||
def _policy_tags(row: Mapping[str, JsonValue]) -> list[JsonValue]:
|
||||
raw: Final = row["request_tags"]
|
||||
tags: Final = json.loads(raw) if isinstance(raw, str) else raw
|
||||
assert isinstance(tags, list), row
|
||||
return [tag for tag in tags if not (isinstance(tag, str) and tag.startswith("User-Agent: "))]
|
||||
|
||||
|
||||
def _spend_logs_metadata(row: Mapping[str, JsonValue]) -> JsonValue:
|
||||
metadata: Final = row["metadata"]
|
||||
return object_value(json.loads(metadata) if isinstance(metadata, str) else metadata).get("spend_logs_metadata")
|
||||
|
||||
|
||||
def _tagged_key(scenario: Scenario, marker: str, **fields: JsonValue) -> tuple[str, str]:
|
||||
team: Final = scenario.team(metadata={"tags": [f"team-{marker}"], "spend_logs_metadata": {"team_field": marker}})
|
||||
project: Final = scenario.project(team, metadata={"tags": [f"project-{marker}"]})
|
||||
key: Final = scenario.key(
|
||||
team_id=team,
|
||||
project_id=project,
|
||||
metadata={"tags": [f"key-{marker}"], "spend_logs_metadata": {"cost_center": marker}},
|
||||
**fields,
|
||||
)
|
||||
return key, sha256(key.encode()).hexdigest()
|
||||
|
||||
|
||||
def _digest(key: str) -> str:
|
||||
return sha256(key.encode()).hexdigest()
|
||||
|
||||
|
||||
def _configured_passthrough(gateway: Gateway, scenario: Scenario, marker: str, target: str, *, auth: bool) -> str:
|
||||
path: Final = f"/integration-passthrough-{marker}"
|
||||
created: Final = gateway.post("/config/pass_through_endpoint", {"path": path, "target": target, "auth": auth})
|
||||
endpoints: Final = TypeAdapter(list[JsonValue]).validate_python(created["endpoints"])
|
||||
endpoint_id: Final = object_value(endpoints[0])["id"]
|
||||
scenario.cleanups.callback(
|
||||
lambda: gateway.request("DELETE", "/config/pass_through_endpoint", params={"endpoint_id": str(endpoint_id)})
|
||||
)
|
||||
return path
|
||||
|
||||
|
||||
def _responses_reply(marker: str, stream: bool) -> Reply:
|
||||
response: Final[dict[str, JsonValue]] = {
|
||||
"id": f"resp_{marker}",
|
||||
"object": "response",
|
||||
"created_at": 1,
|
||||
"status": "completed",
|
||||
"model": "gpt-4o-mini",
|
||||
"output": [
|
||||
{
|
||||
"id": f"msg_{marker}",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"status": "completed",
|
||||
"content": [{"type": "output_text", "text": marker, "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 5, "output_tokens": 3, "total_tokens": 8},
|
||||
}
|
||||
if not stream:
|
||||
return Reply(body=json.dumps(response).encode())
|
||||
events: Final[tuple[dict[str, JsonValue], ...]] = (
|
||||
{
|
||||
"type": "response.created",
|
||||
"sequence_number": 0,
|
||||
"response": {**response, "status": "in_progress", "output": []},
|
||||
},
|
||||
{
|
||||
"type": "response.output_text.delta",
|
||||
"sequence_number": 1,
|
||||
"item_id": f"msg_{marker}",
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"delta": marker,
|
||||
},
|
||||
{"type": "response.completed", "sequence_number": 2, "response": response},
|
||||
)
|
||||
return Reply(
|
||||
content_type="text/event-stream",
|
||||
chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events),
|
||||
)
|
||||
|
||||
|
||||
def _echo_upstream(marker: str) -> Callable[[Request], Reply]:
|
||||
def respond(request: Request) -> Reply:
|
||||
if request.method == "GET" and request.target == "/v1/models":
|
||||
return Reply(body=json.dumps({"object": "list", "data": []}).encode())
|
||||
assert request.method == "POST", request
|
||||
body: Final = object_value(json.loads(request.body))
|
||||
assert marker in json.dumps(body), request
|
||||
if request.target == "/v1/responses":
|
||||
return _responses_reply(marker, body.get("stream") is True)
|
||||
assert body["messages"] == [{"role": "user", "content": marker}], request
|
||||
if body.get("stream") is True:
|
||||
return Reply(chunks=_chat_stream_frames(marker), content_type="text/event-stream")
|
||||
return Reply(body=json.dumps(_chat_reply(marker)).encode())
|
||||
|
||||
return respond
|
||||
|
||||
|
||||
def test_configured_passthrough_spend_row_matches_native_route_tags_and_spend_logs_metadata(gateway: Gateway) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(api_base=wire.url + "/v1")
|
||||
path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=True)
|
||||
key, digest = _tagged_key(scenario, marker, models=[model], allowed_passthrough_routes=[path])
|
||||
headers: Final = {"x-litellm-tags": f"caller-{marker},key-{marker}", "User-Agent": "integration-tags/1"}
|
||||
body: Final[dict[str, JsonValue]] = {"model": model, "messages": [{"role": "user", "content": marker}]}
|
||||
|
||||
native: Final = gateway.request("POST", "/v1/chat/completions", body, key=key, headers=headers)
|
||||
assert native.status_code == 200, native.text
|
||||
passthrough: Final = gateway.request("POST", path, body, key=key, headers=headers)
|
||||
assert passthrough.status_code == 200, passthrough.text
|
||||
assert json.loads(passthrough.content) == _chat_reply(marker)
|
||||
|
||||
native_row: Final = _spend_row(digest, "acompletion")
|
||||
passthrough_row: Final = _spend_row(digest, "pass_through_endpoint")
|
||||
expected: Final = [f"key-{marker}", f"team-{marker}", f"project-{marker}", f"caller-{marker}"]
|
||||
assert _policy_tags(native_row) == expected, native_row
|
||||
assert _policy_tags(passthrough_row) == expected, passthrough_row
|
||||
assert _spend_logs_metadata(native_row) == {"cost_center": marker, "team_field": marker}, native_row
|
||||
assert _spend_logs_metadata(passthrough_row) == {"cost_center": marker, "team_field": marker}, passthrough_row
|
||||
|
||||
|
||||
@pytest.mark.parametrize("bucket", ["metadata", "litellm_metadata"])
|
||||
def test_configured_passthrough_body_tags_lead_and_body_spend_logs_metadata_wins_over_key_and_team(
|
||||
gateway: Gateway, bucket: str
|
||||
) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario:
|
||||
path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=True)
|
||||
key, digest = _tagged_key(scenario, marker, allowed_passthrough_routes=[path])
|
||||
body: Final[dict[str, JsonValue]] = {
|
||||
"messages": [{"role": "user", "content": marker}],
|
||||
bucket: {
|
||||
"tags": [f"body-{marker}", f"team-{marker}"],
|
||||
"spend_logs_metadata": {"cost_center": f"body-{marker}"},
|
||||
},
|
||||
}
|
||||
response: Final = gateway.request("POST", path, body, key=key)
|
||||
assert response.status_code == 200, response.text
|
||||
row: Final = _spend_row(digest, "pass_through_endpoint")
|
||||
assert _policy_tags(row) == [f"body-{marker}", f"team-{marker}", f"key-{marker}", f"project-{marker}"], row
|
||||
assert _spend_logs_metadata(row) == {"cost_center": f"body-{marker}", "team_field": marker}, row
|
||||
|
||||
|
||||
def test_configured_passthrough_streaming_upstream_row_carries_key_team_project_and_caller_tags(
|
||||
gateway: Gateway,
|
||||
) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario:
|
||||
path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=True)
|
||||
key, digest = _tagged_key(scenario, marker, allowed_passthrough_routes=[path])
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
path,
|
||||
{"stream": True, "messages": [{"role": "user", "content": marker}]},
|
||||
key=key,
|
||||
headers={"x-litellm-tags": f"caller-{marker}"},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.content == b"".join(_chat_stream_frames(marker)), response.text
|
||||
row: Final = _spend_row(digest, "pass_through_endpoint")
|
||||
assert _policy_tags(row) == [f"key-{marker}", f"team-{marker}", f"project-{marker}", f"caller-{marker}"], row
|
||||
assert _spend_logs_metadata(row) == {"cost_center": marker, "team_field": marker}, row
|
||||
|
||||
|
||||
def test_configured_passthrough_key_outside_any_team_carries_its_own_tags_and_spend_logs_metadata(
|
||||
gateway: Gateway,
|
||||
) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario:
|
||||
path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=True)
|
||||
key: Final = scenario.key(
|
||||
allowed_passthrough_routes=[path],
|
||||
metadata={"tags": [f"key-{marker}"], "spend_logs_metadata": {"cost_center": marker}},
|
||||
)
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
path,
|
||||
{"messages": [{"role": "user", "content": marker}]},
|
||||
key=key,
|
||||
headers={"x-litellm-tags": f"caller-{marker}"},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
row: Final = _spend_row(_digest(key), "pass_through_endpoint")
|
||||
assert _policy_tags(row) == [f"key-{marker}", f"caller-{marker}"], row
|
||||
assert _spend_logs_metadata(row) == {"cost_center": marker}, row
|
||||
|
||||
|
||||
def test_configured_passthrough_untagged_key_row_keeps_only_caller_tag_and_no_spend_logs_metadata(
|
||||
gateway: Gateway,
|
||||
) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario:
|
||||
path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=True)
|
||||
team: Final = scenario.team()
|
||||
key: Final = scenario.key(team_id=team, allowed_passthrough_routes=[path])
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
path,
|
||||
{"messages": [{"role": "user", "content": marker}]},
|
||||
key=key,
|
||||
headers={"x-litellm-tags": f"caller-{marker}"},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
row: Final = _spend_row(_digest(key), "pass_through_endpoint")
|
||||
assert _policy_tags(row) == [f"caller-{marker}"], row
|
||||
assert _spend_logs_metadata(row) is None, row
|
||||
assert row["team_id"] == team, row
|
||||
|
||||
|
||||
def test_open_passthrough_without_auth_row_carries_only_caller_tag(gateway: Gateway) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario:
|
||||
path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=False)
|
||||
response: Final = gateway.client.post(
|
||||
path,
|
||||
json={"messages": [{"role": "user", "content": marker}]},
|
||||
headers={"x-litellm-tags": f"caller-{marker}"},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
row: Final = _spend_row_tagged(f"caller-{marker}")
|
||||
assert _policy_tags(row) == [f"caller-{marker}"], row
|
||||
assert _spend_logs_metadata(row) is None, row
|
||||
assert row["api_key"] == "", row
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("metadata", "leading_tags"),
|
||||
[
|
||||
({"tags": "string-not-list"}, []),
|
||||
({"tags": [1, None, "z"]}, [1, None, "z"]),
|
||||
({"spend_logs_metadata": "string-not-object"}, []),
|
||||
],
|
||||
)
|
||||
def test_configured_passthrough_hostile_body_metadata_shapes_still_carry_key_team_project_tags(
|
||||
gateway: Gateway, metadata: JsonValue, leading_tags: list[JsonValue]
|
||||
) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario:
|
||||
path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=True)
|
||||
key, digest = _tagged_key(scenario, marker, allowed_passthrough_routes=[path])
|
||||
response: Final = gateway.request(
|
||||
"POST", path, {"messages": [{"role": "user", "content": marker}], "metadata": metadata}, key=key
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
row: Final = _spend_row(digest, "pass_through_endpoint")
|
||||
assert _policy_tags(row) == [*leading_tags, f"key-{marker}", f"team-{marker}", f"project-{marker}"], row
|
||||
assert _spend_logs_metadata(row) == {"cost_center": marker, "team_field": marker}, row
|
||||
|
||||
|
||||
def test_configured_passthrough_body_cannot_forge_user_api_key_attribution_fields(gateway: Gateway) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario:
|
||||
path: Final = _configured_passthrough(gateway, scenario, marker, wire.url + "/echo", auth=True)
|
||||
forged_team: Final = scenario.team()
|
||||
key, digest = _tagged_key(scenario, marker, allowed_passthrough_routes=[path])
|
||||
forged: Final[dict[str, JsonValue]] = {
|
||||
"user_api_key": "forged-" + marker,
|
||||
"user_api_key_team_id": forged_team,
|
||||
"user_api_key_user_id": "forged-" + marker,
|
||||
"user_api_key_alias": "forged-" + marker,
|
||||
}
|
||||
body: Final[dict[str, JsonValue]] = {"messages": [{"role": "user", "content": marker}], "metadata": forged}
|
||||
response: Final = gateway.request("POST", path, body, key=key)
|
||||
assert response.status_code == 200, response.text
|
||||
row: Final = _spend_row(digest, "pass_through_endpoint")
|
||||
assert row["team_id"] != forged_team, row
|
||||
assert _policy_tags(row) == [f"key-{marker}", f"team-{marker}", f"project-{marker}"], row
|
||||
assert read_rows('SELECT api_key FROM "LiteLLM_SpendLogs" WHERE team_id=%s', (forged_team,)) == [], forged_team
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stream", [False, True])
|
||||
@pytest.mark.parametrize("route", ["/v1/chat/completions", "/v1/messages"])
|
||||
def test_native_routes_carry_key_team_project_and_caller_tags_and_key_over_team_spend_logs_metadata(
|
||||
gateway: Gateway, route: str, stream: bool
|
||||
) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
with wire_server(_echo_upstream(marker)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(api_base=wire.url + "/v1")
|
||||
key, digest = _tagged_key(scenario, marker, models=[model])
|
||||
body: Final[dict[str, JsonValue]] = {
|
||||
"model": model,
|
||||
"max_tokens": 16,
|
||||
"stream": stream,
|
||||
"messages": [{"role": "user", "content": marker}],
|
||||
}
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
route,
|
||||
body,
|
||||
key=key,
|
||||
headers={"x-litellm-tags": f"caller-{marker}", "User-Agent": "integration-tags/1"},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows('SELECT request_tags, metadata FROM "LiteLLM_SpendLogs" WHERE api_key=%s', (digest,)),
|
||||
lambda values: len(values) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
assert _policy_tags(rows[0]) == [f"key-{marker}", f"team-{marker}", f"project-{marker}", f"caller-{marker}"], (
|
||||
rows
|
||||
)
|
||||
assert _spend_logs_metadata(rows[0]) == {"cost_center": marker, "team_field": marker}, rows
|
||||
|
||||
|
||||
def test_anthropic_passthrough_spend_row_carries_key_team_project_tags_and_spend_logs_metadata(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
marker: Final = uuid.uuid4().hex
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.method == "POST" and request.target == "/v1/messages", request
|
||||
assert request.headers["x-api-key"] == "synthetic-anthropic-key"
|
||||
return Reply(body=json.dumps(_anthropic_reply(marker)).encode())
|
||||
|
||||
config: Final = tmp_path / "proxy_config.yaml"
|
||||
config.write_text(
|
||||
"model_list: []\n"
|
||||
"general_settings:\n"
|
||||
" master_key: os.environ/LITELLM_MASTER_KEY\n"
|
||||
" database_url: os.environ/DATABASE_URL\n"
|
||||
" store_model_in_db: true\n"
|
||||
" disable_spend_logs: false\n"
|
||||
" proxy_batch_write_at: 1\n"
|
||||
"router_settings:\n"
|
||||
" disable_cooldowns: true\n"
|
||||
)
|
||||
with wire_server(respond) as wire:
|
||||
overrides: Final = {"ANTHROPIC_API_BASE": wire.url, "ANTHROPIC_API_KEY": "synthetic-anthropic-key"}
|
||||
with owned_proxy(gateway, tmp_path, overrides, config=config) as candidate, candidate.scenario() as scenario:
|
||||
key, digest = _tagged_key(scenario, marker)
|
||||
response: Final = candidate.request(
|
||||
"POST",
|
||||
"/anthropic/v1/messages",
|
||||
{
|
||||
"model": "claude-sonnet-4-5",
|
||||
"max_tokens": 16,
|
||||
"messages": [{"role": "user", "content": marker}],
|
||||
},
|
||||
key=key,
|
||||
headers={"x-litellm-tags": f"caller-{marker}", "User-Agent": "integration-tags/1"},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert json.loads(response.content) == _anthropic_reply(marker)
|
||||
row: Final = _spend_row(digest, "pass_through_endpoint")
|
||||
assert _policy_tags(row) == [f"key-{marker}", f"team-{marker}", f"project-{marker}", f"caller-{marker}"], (
|
||||
row
|
||||
)
|
||||
assert _spend_logs_metadata(row) == {"cost_center": marker, "team_field": marker}, row
|
||||
|
|
@ -63,7 +63,8 @@ def mock_request():
|
|||
self.method = method
|
||||
self.request_body = request_body or {}
|
||||
# Add url attribute that the actual code expects
|
||||
self.url = "http://localhost:8000/test"
|
||||
self.url = httpx.URL("http://localhost:8000/test")
|
||||
self.scope = {"type": "http", "method": method, "path": "/test"}
|
||||
# Add state attribute that FastAPI requests have
|
||||
self.state = type("State", (), {})()
|
||||
|
||||
|
|
@ -414,6 +415,7 @@ PROTOCOL_CONSTRAINED_PASS_THROUGH_ROUTES = {
|
|||
"/transcribe": {"POST"},
|
||||
"/transcribe/{operation}": {"POST"},
|
||||
"/tinyfish/{endpoint:path}": {"GET", "POST"},
|
||||
"/laya/v1/systemone": {"POST"},
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -80,6 +80,25 @@ async def _turn(
|
|||
)
|
||||
|
||||
|
||||
async def _benchmark_rows(
|
||||
db, start: datetime, end: datetime, key: str | None = None, user_id: str | None = None
|
||||
) -> list[dict]:
|
||||
return await db.query_raw(
|
||||
AUTOROUTER_BENCHMARKS_SQL,
|
||||
start.isoformat(),
|
||||
end.isoformat(),
|
||||
key,
|
||||
user_id,
|
||||
start.date().isoformat(),
|
||||
(end - timedelta(days=1)).date().isoformat(),
|
||||
)
|
||||
|
||||
|
||||
async def _days(db, key: str | None = None, user_id: str | None = None, router: str | None = None) -> list[dict]:
|
||||
rows = await _benchmark_rows(db, T0 - timedelta(days=1), T0 + timedelta(days=2), key, user_id)
|
||||
return [row for row in rows if row["turns"] and (router is None or row["router_name"] == router)]
|
||||
|
||||
|
||||
async def _row(db, key: str, session_id: str = "s1", router: str = "auto-1") -> dict:
|
||||
rows = await db.query_raw(
|
||||
'SELECT * FROM "LiteLLM_AutoRouterSession" WHERE api_key = $1 AND session_id = $2 AND router_name = $3',
|
||||
|
|
@ -225,18 +244,15 @@ async def test_subtotal_coverage_survives_legacy_and_rolling_writers(db, writers
|
|||
assert row["savings_estimated_turns"] == sum(writers)
|
||||
assert row["savings_estimated_actual_spend"] == pytest.approx(0.01 * sum(writers))
|
||||
assert row["savings_estimated_saved_spend"] == pytest.approx(0.02 * sum(writers))
|
||||
groups: Final = await db.query_raw(
|
||||
AUTOROUTER_BENCHMARKS_SQL, T0.isoformat(), (T0 + timedelta(days=1)).isoformat(), key, None
|
||||
)
|
||||
assert len(groups) == 1
|
||||
assert groups[0]["classifier_cost"] == row["classifier_cost"]
|
||||
assert groups[0]["classifier_cost_recorded_turns"] == sum(writers)
|
||||
assert groups[0]["turns"] == len(writers)
|
||||
assert groups[0]["spend"] == row["spend"]
|
||||
assert groups[0]["saved_spend"] == row["saved_spend"]
|
||||
assert groups[0]["savings_estimated_turns"] == sum(writers)
|
||||
assert groups[0]["savings_estimated_actual_spend"] == row["savings_estimated_actual_spend"]
|
||||
assert groups[0]["savings_estimated_saved_spend"] == row["savings_estimated_saved_spend"]
|
||||
days: Final = await _days(db, key)
|
||||
assert len(days) == int(any(writers))
|
||||
for day in days:
|
||||
assert day["classifier_cost"] == row["classifier_cost"]
|
||||
assert day["classifier_cost_recorded_turns"] == day["turns"] == sum(writers)
|
||||
assert day["spend"] == pytest.approx(0.01 * sum(writers))
|
||||
assert day["saved_spend"] == pytest.approx(0.02 * sum(writers))
|
||||
assert day["savings_estimated_actual_spend"] == row["savings_estimated_actual_spend"]
|
||||
assert day["savings_estimated_saved_spend"] == row["savings_estimated_saved_spend"]
|
||||
|
||||
|
||||
async def test_unknown_and_legacy_turns_preserve_actual_spend_without_entering_the_estimated_cohort(db: Prisma) -> None:
|
||||
|
|
@ -250,13 +266,10 @@ async def test_unknown_and_legacy_turns_preserve_actual_spend_without_entering_t
|
|||
row: Final = await _row(db, key)
|
||||
assert row["saved_spend"] == pytest.approx(-0.03)
|
||||
assert row["savings_estimated_baseline_models"] == {"opus": 1}
|
||||
groups: Final = await db.query_raw(
|
||||
AUTOROUTER_BENCHMARKS_SQL, T0.isoformat(), (T0 + timedelta(days=1)).isoformat(), key, None
|
||||
)
|
||||
assert len(groups) == 1
|
||||
for actual in (row, groups[0]):
|
||||
assert actual["turns"] == 3
|
||||
assert actual["spend"] == pytest.approx(0.96)
|
||||
(day,) = await _days(db, key)
|
||||
assert (row["turns"], day["turns"]) == (3, 2)
|
||||
assert (row["spend"], day["spend"]) == (pytest.approx(0.96), pytest.approx(0.95))
|
||||
for actual in (row, day):
|
||||
assert actual["savings_estimated_turns"] == 1
|
||||
assert actual["savings_estimated_actual_spend"] == pytest.approx(0.25)
|
||||
assert actual["savings_estimated_saved_spend"] == pytest.approx(-0.05)
|
||||
|
|
@ -281,25 +294,20 @@ async def test_the_benchmarks_aggregate_reads_only_overlapping_sessions(db):
|
|||
await _turn(db, key, "A", T0, session_id=in_window, router=router, saved=0.5, spend=0.25, classifier_cost=0.02)
|
||||
await _turn(db, key, "A", T0 - timedelta(days=40), session_id=out_of_window, router=router, classifier_cost=9.0)
|
||||
|
||||
rows = await db.query_raw(
|
||||
AUTOROUTER_BENCHMARKS_SQL,
|
||||
(T0 - timedelta(days=1)).isoformat(),
|
||||
(T0 + timedelta(days=1)).isoformat(),
|
||||
None,
|
||||
None,
|
||||
)
|
||||
rows = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), None, None)
|
||||
matching = [row for row in rows if row["router_name"] == router]
|
||||
assert len(matching) == 1
|
||||
grouped = matching[0]
|
||||
assert grouped["router_type"] == "complexity"
|
||||
assert grouped["sessions"] == 1
|
||||
assert grouped["turns"] == 2
|
||||
assert grouped["spend"] == pytest.approx(0.5)
|
||||
assert grouped["saved_spend"] == pytest.approx(1.0)
|
||||
assert grouped["classifier_cost"] == pytest.approx(0.03)
|
||||
assert grouped["classifier_cost_recorded_turns"] == 2
|
||||
assert grouped["session_turns"] == 2
|
||||
assert grouped["unordered_turns"] == 1
|
||||
assert grouped["session_seconds"] == pytest.approx(60.0)
|
||||
(day,) = await _days(db, router=router)
|
||||
assert (day["turns"], day["classifier_cost_recorded_turns"]) == (2, 2)
|
||||
assert day["spend"] == pytest.approx(0.5)
|
||||
assert day["saved_spend"] == pytest.approx(1.0)
|
||||
assert day["classifier_cost"] == pytest.approx(0.03)
|
||||
|
||||
|
||||
async def test_the_benchmarks_aggregate_can_filter_to_one_key(db):
|
||||
|
|
@ -309,32 +317,22 @@ async def test_the_benchmarks_aggregate_can_filter_to_one_key(db):
|
|||
await _turn(db, first_key, "A", T0, router=router, saved=0.5, classifier_cost=0.01)
|
||||
await _turn(db, second_key, "A", T0, router=router, saved=9.0, classifier_cost=0.09)
|
||||
|
||||
rows = await db.query_raw(
|
||||
AUTOROUTER_BENCHMARKS_SQL,
|
||||
(T0 - timedelta(days=1)).isoformat(),
|
||||
(T0 + timedelta(days=1)).isoformat(),
|
||||
first_key,
|
||||
None,
|
||||
)
|
||||
rows = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), first_key, None)
|
||||
matching = [row for row in rows if row["router_name"] == router]
|
||||
assert len(matching) == 1
|
||||
assert matching[0]["sessions"] == 1
|
||||
assert matching[0]["saved_spend"] == pytest.approx(0.5)
|
||||
assert matching[0]["classifier_cost"] == pytest.approx(0.01)
|
||||
assert matching[0]["classifier_cost_recorded_turns"] == 1
|
||||
(day,) = await _days(db, first_key, router=router)
|
||||
assert day["saved_spend"] == pytest.approx(0.5)
|
||||
assert day["classifier_cost"] == pytest.approx(0.01)
|
||||
assert day["classifier_cost_recorded_turns"] == 1
|
||||
|
||||
unknown_key_rows = await db.query_raw(
|
||||
AUTOROUTER_BENCHMARKS_SQL,
|
||||
(T0 - timedelta(days=1)).isoformat(),
|
||||
(T0 + timedelta(days=1)).isoformat(),
|
||||
f"k-{uuid.uuid4()}",
|
||||
None,
|
||||
)
|
||||
unknown_key_rows = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), f"k-{uuid.uuid4()}", None)
|
||||
assert [row for row in unknown_key_rows if row["router_name"] == router] == []
|
||||
|
||||
|
||||
class _BenchmarkRow(TypedDict):
|
||||
sessions: ReadOnly[int]
|
||||
session_turns: ReadOnly[int]
|
||||
turns: ReadOnly[int]
|
||||
same_model_turns: ReadOnly[int]
|
||||
first_visit_turns: ReadOnly[int]
|
||||
|
|
@ -350,14 +348,11 @@ class _BenchmarkRow(TypedDict):
|
|||
async def _scoped_benchmarks(
|
||||
db: Prisma, router: str, user_id: str | None = None, key: str | None = None
|
||||
) -> tuple[_BenchmarkRow, ...]:
|
||||
rows: Final = await db.query_raw(
|
||||
AUTOROUTER_BENCHMARKS_SQL,
|
||||
(T0 - timedelta(days=1)).isoformat(),
|
||||
(T0 + timedelta(days=1)).isoformat(),
|
||||
key,
|
||||
user_id,
|
||||
rows: Final = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), key, user_id)
|
||||
days: Final = await _days(db, key, user_id, router)
|
||||
return tuple(
|
||||
cast(_BenchmarkRow, {**row, **next(iter(days), {})}) for row in rows if row["router_name"] == router
|
||||
)
|
||||
return tuple(cast(_BenchmarkRow, row) for row in rows if row["router_name"] == router)
|
||||
|
||||
|
||||
async def test_users_keep_written_identity_across_shared_keys_and_keyless_sessions(db: Prisma) -> None:
|
||||
|
|
@ -384,28 +379,33 @@ async def test_users_keep_written_identity_across_shared_keys_and_keyless_sessio
|
|||
intersection: Final = await _scoped_benchmarks(db, router, user_id=alice, key=first_key)
|
||||
assert len(alice_rows) == len(bob_rows) == len(global_rows) == len(key_rows) == len(intersection) == 1
|
||||
assert (alice_rows[0]["sessions"], alice_rows[0]["turns"], alice_rows[0]["same_model_turns"]) == (3, 4, 1)
|
||||
assert (alice_rows[0]["session_turns"], bob_rows[0]["session_turns"]) == (4, 2)
|
||||
assert (bob_rows[0]["sessions"], bob_rows[0]["turns"], bob_rows[0]["first_visit_turns"]) == (2, 2, 2)
|
||||
assert alice_rows[0]["spend"] == pytest.approx(0.05)
|
||||
assert bob_rows[0]["spend"] == pytest.approx(0.07)
|
||||
assert alice_rows[0]["tier_turns"] == {"simple": 1}
|
||||
assert bob_rows[0]["tier_turns"] == {"complex": 1}
|
||||
assert (alice_rows[0]["cache_hits"], bob_rows[0]["cache_hits"]) == (1, 0)
|
||||
assert (global_rows[0]["sessions"], global_rows[0]["turns"]) == (4, 7)
|
||||
assert (global_rows[0]["sessions"], global_rows[0]["session_turns"], global_rows[0]["turns"]) == (4, 7, 6)
|
||||
assert (alice_rows[0]["savings_estimated_turns"], bob_rows[0]["savings_estimated_turns"]) == (4, 2)
|
||||
assert global_rows[0]["savings_estimated_turns"] == 6
|
||||
for scoped in (alice_rows[0], bob_rows[0]):
|
||||
assert scoped["savings_estimated_actual_spend"] == pytest.approx(scoped["spend"])
|
||||
assert scoped["savings_estimated_saved_spend"] == pytest.approx(scoped["saved_spend"])
|
||||
assert global_rows[0]["spend"] == pytest.approx(alice_rows[0]["spend"] + bob_rows[0]["spend"] + 0.01)
|
||||
assert global_rows[0]["saved_spend"] == pytest.approx(alice_rows[0]["saved_spend"] + bob_rows[0]["saved_spend"] + 0.02)
|
||||
assert global_rows[0]["spend"] == pytest.approx(alice_rows[0]["spend"] + bob_rows[0]["spend"])
|
||||
assert global_rows[0]["saved_spend"] == pytest.approx(alice_rows[0]["saved_spend"] + bob_rows[0]["saved_spend"])
|
||||
assert global_rows[0]["tier_turns"] == {"simple": 1, "complex": 1}
|
||||
assert (key_rows[0]["sessions"], key_rows[0]["turns"]) == (1, 3)
|
||||
assert key_rows[0]["spend"] == pytest.approx(0.05)
|
||||
assert (key_rows[0]["sessions"], key_rows[0]["session_turns"], key_rows[0]["turns"]) == (1, 3, 2)
|
||||
assert key_rows[0]["spend"] == pytest.approx(0.04)
|
||||
assert (intersection[0]["sessions"], intersection[0]["turns"]) == (1, 1)
|
||||
assert intersection[0]["spend"] == pytest.approx(0.01)
|
||||
assert await _scoped_benchmarks(db, router, user_id=bob, key=second_key) == ()
|
||||
assert await _scoped_benchmarks(db, router, user_id=f"u-{uuid.uuid4()}") == ()
|
||||
assert await _scoped_benchmarks(db, router, user_id="") == ()
|
||||
assert [
|
||||
row
|
||||
for row in await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), None, "")
|
||||
if row["router_name"] == router
|
||||
] == []
|
||||
|
||||
|
||||
async def test_a_failed_user_projection_rolls_back_the_keys_increment(db: Prisma) -> None:
|
||||
|
|
@ -419,6 +419,7 @@ async def test_a_failed_user_projection_rolls_back_the_keys_increment(db: Prisma
|
|||
|
||||
assert await _row(db, key) == before
|
||||
assert await db.query_raw('SELECT user_id FROM "LiteLLM_AutoRouterUserSession" WHERE user_id = $1', user_id) == []
|
||||
assert [day["turns"] for day in await _days(db, key)] == [1]
|
||||
|
||||
first_user: Final = f"u-{uuid.uuid4()}"
|
||||
second_user: Final = f"u-{uuid.uuid4()}"
|
||||
|
|
@ -463,6 +464,14 @@ async def test_a_failed_user_projection_rolls_back_the_keys_increment(db: Prisma
|
|||
assert (row["turns"], row["same_model_turns"], row["unordered_turns"], row["last_model"]) == (count, 1, 0, model)
|
||||
assert row["spend"] == pytest.approx(count * 0.01)
|
||||
assert row["saved_spend"] == pytest.approx(count * 0.02)
|
||||
days: Final = await db.query_raw(
|
||||
'SELECT user_id, turns, saved_spend FROM "LiteLLM_AutoRouterDailySpend" WHERE api_key = $1', key
|
||||
)
|
||||
assert {day["user_id"]: (day["turns"], day["saved_spend"]) for day in days} == {
|
||||
"": (1, pytest.approx(0.02)),
|
||||
first_user: (3, pytest.approx(0.06)),
|
||||
second_user: (2, pytest.approx(0.04)),
|
||||
}
|
||||
|
||||
|
||||
async def test_user_session_cleanup_keeps_another_users_recent_keyless_session(db: Prisma) -> None:
|
||||
|
|
@ -490,13 +499,7 @@ async def test_a_reconfigured_alias_reports_each_router_type_as_its_own_group(db
|
|||
db, key, "A", T0 + timedelta(seconds=10), session_id=f"s-{uuid.uuid4()}", router=router, router_type="quality"
|
||||
)
|
||||
|
||||
rows = await db.query_raw(
|
||||
AUTOROUTER_BENCHMARKS_SQL,
|
||||
(T0 - timedelta(days=1)).isoformat(),
|
||||
(T0 + timedelta(days=1)).isoformat(),
|
||||
None,
|
||||
None,
|
||||
)
|
||||
rows = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), None, None)
|
||||
matching = sorted(
|
||||
(row for row in rows if row["router_name"] == router),
|
||||
key=lambda row: row["router_type"],
|
||||
|
|
@ -575,16 +578,10 @@ async def test_the_benchmarks_aggregate_sums_tier_turns_across_sessions(db):
|
|||
await _turn(db, key, "B", T0 + timedelta(seconds=20), session_id=f"s-{uuid.uuid4()}", router=router, tier="complex")
|
||||
await _turn(db, key, "C", T0 + timedelta(seconds=30), session_id=f"s-{uuid.uuid4()}", router=router, tier=None)
|
||||
|
||||
rows = await db.query_raw(
|
||||
AUTOROUTER_BENCHMARKS_SQL,
|
||||
(T0 - timedelta(days=1)).isoformat(),
|
||||
(T0 + timedelta(days=1)).isoformat(),
|
||||
None,
|
||||
None,
|
||||
)
|
||||
rows = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), None, None)
|
||||
grouped = next(row for row in rows if row["router_name"] == router)
|
||||
assert grouped["tier_turns"] == {"simple": 2, "complex": 1}
|
||||
assert grouped["turns"] == 4
|
||||
assert grouped["session_turns"] == 4
|
||||
|
||||
|
||||
async def test_tier_maps_stay_separate_per_router_type_on_a_reconfigured_alias(db):
|
||||
|
|
@ -604,13 +601,7 @@ async def test_tier_maps_stay_separate_per_router_type_on_a_reconfigured_alias(d
|
|||
tier="2",
|
||||
)
|
||||
|
||||
rows = await db.query_raw(
|
||||
AUTOROUTER_BENCHMARKS_SQL,
|
||||
(T0 - timedelta(days=1)).isoformat(),
|
||||
(T0 + timedelta(days=1)).isoformat(),
|
||||
None,
|
||||
None,
|
||||
)
|
||||
rows = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), None, None)
|
||||
by_type = {row["router_type"]: row["tier_turns"] for row in rows if row["router_name"] == router}
|
||||
assert by_type == {"complexity": {"medium": 1}, "quality": {"2": 1}}
|
||||
|
||||
|
|
@ -620,13 +611,7 @@ async def test_a_window_with_no_tiered_turns_aggregates_to_an_empty_map(db):
|
|||
router = f"r-{uuid.uuid4()}"
|
||||
await _turn(db, key, "A", T0, session_id=f"s-{uuid.uuid4()}", router=router, tier=None)
|
||||
|
||||
rows = await db.query_raw(
|
||||
AUTOROUTER_BENCHMARKS_SQL,
|
||||
(T0 - timedelta(days=1)).isoformat(),
|
||||
(T0 + timedelta(days=1)).isoformat(),
|
||||
None,
|
||||
None,
|
||||
)
|
||||
rows = await _benchmark_rows(db, (T0 - timedelta(days=1)), (T0 + timedelta(days=1)), None, None)
|
||||
grouped = next(row for row in rows if row["router_name"] == router)
|
||||
assert grouped["tier_turns"] == {}
|
||||
|
||||
|
|
@ -653,3 +638,49 @@ async def test_an_out_of_order_hit_still_counts_toward_the_overall_hit_rate(db):
|
|||
assert row["unordered_turns"] == 1
|
||||
assert row["cache_hits"] == 1
|
||||
assert row["same_model_hits"] + row["first_visit_hits"] + row["return_hits"] == 0
|
||||
|
||||
|
||||
async def test_a_cross_midnight_session_splits_its_money_by_request_day(db):
|
||||
key = f"k-{uuid.uuid4()}"
|
||||
router = f"auto-{uuid.uuid4()}"
|
||||
midnight = datetime(2026, 9, 2)
|
||||
await _turn(db, key, "A", midnight - timedelta(minutes=10), router=router, spend=1.0, saved=7.0, user_id="u1")
|
||||
await _turn(db, key, "A", midnight + timedelta(minutes=10), router=router, spend=1.0, saved=3.0, user_id="u1")
|
||||
await _turn(db, key, "B", midnight + timedelta(days=1), router=router, spend=1.0, saved=11.0, user_id="u1")
|
||||
|
||||
assert (await _row(db, key, router=router))["saved_spend"] == 21.0
|
||||
days = await db.query_raw(
|
||||
'SELECT date, turns, saved_spend FROM "LiteLLM_AutoRouterDailySpend" WHERE api_key = $1 ORDER BY date', key
|
||||
)
|
||||
assert [(d["date"], d["turns"], d["saved_spend"]) for d in days] == [
|
||||
("2026-09-01", 1, 7.0),
|
||||
("2026-09-02", 1, 3.0),
|
||||
("2026-09-03", 1, 11.0),
|
||||
]
|
||||
for user_id in (None, "u1"):
|
||||
(selected,) = await _benchmark_rows(db, midnight, midnight + timedelta(days=1), key, user_id)
|
||||
assert (selected["sessions"], selected["session_turns"]) == (1, 3)
|
||||
assert (selected["turns"], selected["spend"], selected["saved_spend"]) == (1, 1.0, 3.0)
|
||||
|
||||
|
||||
async def test_a_router_type_change_within_a_day_keeps_each_types_money_apart(db):
|
||||
key = f"k-{uuid.uuid4()}"
|
||||
router = f"auto-{uuid.uuid4()}"
|
||||
await _turn(db, key, "A", T0, router=router, router_type="complexity", spend=1.0, saved=4.0)
|
||||
await _turn(db, key, "A", T0 + timedelta(hours=1), router=router, router_type="quality", spend=2.0, saved=0.0)
|
||||
|
||||
days = {day["router_type"]: (day["turns"], day["spend"], day["saved_spend"]) for day in await _days(db, key)}
|
||||
assert days == {"complexity": (1, 1.0, 4.0), "quality": (1, 2.0, 0.0)}
|
||||
|
||||
|
||||
async def test_a_router_type_change_mid_session_keeps_session_shape_with_the_sessions_type(db):
|
||||
key = f"k-{uuid.uuid4()}"
|
||||
router = f"auto-{uuid.uuid4()}"
|
||||
await _turn(db, key, "A", T0, router=router, router_type="complexity", spend=1.0, saved=4.0)
|
||||
await _turn(db, key, "B", T0 + timedelta(hours=1), router=router, router_type="quality", spend=2.0, saved=0.0)
|
||||
|
||||
rows = {row["router_type"]: row for row in await _benchmark_rows(db, T0, T0 + timedelta(days=1), key)}
|
||||
assert set(rows) == {"complexity", "quality"}
|
||||
assert (rows["complexity"]["sessions"], rows["complexity"]["session_turns"], rows["complexity"]["turns"]) == (1, 2, 1)
|
||||
assert (rows["quality"]["sessions"], rows["quality"]["session_turns"], rows["quality"]["turns"]) == (0, 0, 1)
|
||||
assert rows["quality"]["spend"] == 2.0
|
||||
|
|
|
|||
|
|
@ -151,6 +151,13 @@ async def test_late_replay_updates_all_projections_without_rebilling(db: Prisma,
|
|||
):
|
||||
assert after_users["late-user"][field] == after[field]
|
||||
assert after_users["late-user"]["turns"] == 1 and after_users["late-user"]["spend"] == 0.17
|
||||
days: Final = await db.query_raw(
|
||||
'SELECT * FROM "LiteLLM_AutoRouterDailySpend" WHERE api_key=$1 ORDER BY user_id', late.api_key
|
||||
)
|
||||
assert [(day["date"], day["user_id"]) for day in days] == [("1970-01-01", "early-user"), ("1970-01-01", "late-user")]
|
||||
assert days[0]["saved_spend"] == days[0]["savings_estimated_turns"] == 0
|
||||
for field in ("saved_spend", "savings_estimated_turns", "savings_estimated_actual_spend", "savings_estimated_saved_spend"):
|
||||
assert days[1][field] == after[field]
|
||||
for table in ("DailyUserSpend", "DailyTeamSpend", "DailyOrganizationSpend", "DailyEndUserSpend", "DailyAgentSpend", "DailyTagSpend"):
|
||||
rows: Final = await db.query_raw(f'SELECT spend,api_requests,autorouter_savings_spend FROM "LiteLLM_{table}" WHERE api_key=$1', late.api_key)
|
||||
assert rows[0]["spend"] == rows[0]["api_requests"] == 0
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@
|
|||
|
||||
import asyncio
|
||||
import logging
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -110,6 +111,61 @@ def test_excluded_services_from_env_csv(monkeypatch):
|
|||
assert OpenTelemetryV2Config().excluded_services == frozenset({"redis", "postgresql"})
|
||||
|
||||
|
||||
@pytest.mark.parametrize("name", ["EXCLUDED_SERVICES", "excluded_services", "Excluded_Services"])
|
||||
def test_a_bare_excluded_services_env_var_is_ignored(monkeypatch, name):
|
||||
for env_name in ("LITELLM_OTEL_EXCLUDED_SERVICES", "EXCLUDED_SERVICES", "excluded_services", "Excluded_Services"):
|
||||
monkeypatch.delenv(env_name, raising=False)
|
||||
monkeypatch.setenv(name, "redis,postgres")
|
||||
assert OpenTelemetryV2Config().excluded_services == frozenset()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("set_env_name", "env_value", "case_sensitive", "env_ignore_empty", "env_parse_none_str"),
|
||||
[
|
||||
pytest.param("otel_service_name", "lower", True, False, None, id="case-sensitive"),
|
||||
pytest.param("OTEL_SERVICE_NAME", "", False, True, None, id="ignore-empty"),
|
||||
pytest.param("OTEL_ENDPOINT", "null", False, False, "null", id="parse-none"),
|
||||
pytest.param("excluded_services", "redis", True, False, None, id="bare-exclusion"),
|
||||
],
|
||||
)
|
||||
def test_env_source_preserves_runtime_options(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
set_env_name: str,
|
||||
env_value: str,
|
||||
case_sensitive: bool,
|
||||
env_ignore_empty: bool,
|
||||
env_parse_none_str: str | None,
|
||||
) -> None:
|
||||
for env_name in (
|
||||
"OTEL_SERVICE_NAME",
|
||||
"otel_service_name",
|
||||
"OTEL_ENDPOINT",
|
||||
"OTEL_EXPORTER_OTLP_ENDPOINT",
|
||||
"LITELLM_OTEL_EXCLUDED_SERVICES",
|
||||
"EXCLUDED_SERVICES",
|
||||
"excluded_services",
|
||||
"Excluded_Services",
|
||||
):
|
||||
monkeypatch.delenv(env_name, raising=False)
|
||||
monkeypatch.setenv(set_env_name, env_value)
|
||||
config: Final = OpenTelemetryV2Config(
|
||||
_case_sensitive=case_sensitive,
|
||||
_env_ignore_empty=env_ignore_empty,
|
||||
_env_parse_none_str=env_parse_none_str,
|
||||
)
|
||||
assert config.service_name == "litellm"
|
||||
assert config.endpoint is None
|
||||
assert config.excluded_services == frozenset()
|
||||
|
||||
|
||||
def test_the_documented_env_var_wins_over_a_bare_excluded_services(monkeypatch):
|
||||
for env_name in ("LITELLM_OTEL_EXCLUDED_SERVICES", "EXCLUDED_SERVICES", "excluded_services", "Excluded_Services"):
|
||||
monkeypatch.delenv(env_name, raising=False)
|
||||
monkeypatch.setenv("EXCLUDED_SERVICES", "postgres")
|
||||
monkeypatch.setenv("LITELLM_OTEL_EXCLUDED_SERVICES", "redis")
|
||||
assert OpenTelemetryV2Config().excluded_services == frozenset({"redis"})
|
||||
|
||||
|
||||
def test_excluded_services_config_wins_over_env(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_OTEL_EXCLUDED_SERVICES", "redis")
|
||||
assert OpenTelemetryV2Config(excluded_services=["postgres"]).excluded_services == frozenset({"postgresql"})
|
||||
|
|
|
|||
|
|
@ -1116,6 +1116,18 @@ class TestProviderWiring:
|
|||
|
||||
assert self._fan_out_of(preset)._excluded_db_systems == frozenset({"redis"})
|
||||
|
||||
@pytest.mark.parametrize("otel", [None, True, "on", "", []], ids=["null", "true", "on", "empty_string", "empty_list"])
|
||||
def test_a_non_mapping_otel_block_falls_back_to_the_published_logger_config(self, monkeypatch, otel):
|
||||
monkeypatch.setattr(litellm, "callback_settings", {"otel": otel}, raising=False)
|
||||
preset = OpenTelemetryV2(
|
||||
config=OpenTelemetryV2Config(exporters=[ExporterSpec(kind="in_memory")], excluded_services=["redis"]),
|
||||
callback_name="langfuse_otel",
|
||||
)
|
||||
|
||||
publish_global_otel_v2_provider([], lambda _p: None, registered=preset)
|
||||
|
||||
assert self._fan_out_of(preset)._excluded_db_systems == frozenset({"redis"})
|
||||
|
||||
def test_otel_after_a_preset_reuses_it_and_still_takes_callback_settings_exclusions(self, monkeypatch):
|
||||
"""``callbacks: [langfuse_otel, otel]`` keeps one v2 logger, exactly as
|
||||
before ``excluded_services`` existed, and the exclusion still comes from
|
||||
|
|
|
|||
|
|
@ -1029,6 +1029,84 @@ class TestJWTKeyMappingCascade:
|
|||
|
||||
|
||||
|
||||
class TestStripPrismaQueryParams:
|
||||
"""The psycopg URL the job connects with is derived from the Prisma-dialect
|
||||
DATABASE_URL, whose TLS params mean something else to libpq."""
|
||||
|
||||
@staticmethod
|
||||
def _query(url: str) -> dict[str, str]:
|
||||
from urllib.parse import parse_qsl, urlparse
|
||||
|
||||
return dict(parse_qsl(urlparse(url).query))
|
||||
|
||||
def test_prisma_ca_sslcert_becomes_sslrootcert_with_verify_full(self):
|
||||
url = "postgresql://u:p@writer:5432/db?schema=public&sslmode=require&sslcert=/tmp/pinned.pem&sslaccept=strict"
|
||||
|
||||
cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url)
|
||||
|
||||
assert self._query(cleaned) == {"sslmode": "verify-full", "sslrootcert": "/tmp/pinned.pem"}
|
||||
assert cleaned.startswith("postgresql://u:p@writer:5432/db?")
|
||||
|
||||
@pytest.mark.parametrize("sslmode", ["prefer", "require"])
|
||||
@pytest.mark.parametrize("sslaccept", ["strict", "unknown-mode-prisma-treats-as-strict"])
|
||||
def test_strict_verifies_chain_and_hostname_whatever_sslmode_prisma_was_given(self, sslmode, sslaccept):
|
||||
url = f"postgresql://writer/db?sslmode={sslmode}&sslcert=/certs/ca.pem&sslaccept={sslaccept}"
|
||||
|
||||
cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url)
|
||||
|
||||
assert self._query(cleaned) == {"sslmode": "verify-full", "sslrootcert": "/certs/ca.pem"}
|
||||
|
||||
def test_strict_with_tls_disabled_stays_off(self):
|
||||
url = "postgresql://writer/db?sslmode=disable&sslcert=/certs/ca.pem&sslaccept=strict"
|
||||
|
||||
cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url)
|
||||
|
||||
assert self._query(cleaned) == {"sslmode": "disable"}
|
||||
|
||||
@pytest.mark.parametrize("sslaccept", ["&sslaccept=accept_invalid_certs", ""])
|
||||
def test_without_strict_the_ca_is_dropped_so_libpq_checks_nothing_like_prisma(self, sslaccept):
|
||||
url = f"postgresql://writer/db?sslmode=require&sslcert=/certs/ca.pem{sslaccept}"
|
||||
|
||||
cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url)
|
||||
|
||||
assert self._query(cleaned) == {"sslmode": "require"}
|
||||
|
||||
def test_a_ca_alone_without_strict_or_sslmode_leaves_libpq_its_defaults(self):
|
||||
cleaned = ProxyExtrasDBManager._strip_prisma_query_params("postgresql://writer/db?sslcert=/certs/ca.pem")
|
||||
|
||||
assert cleaned == "postgresql://writer/db"
|
||||
|
||||
def test_a_libpq_client_certificate_pair_is_left_alone(self):
|
||||
url = "postgresql://writer/db?sslmode=verify-full&sslrootcert=/ca.pem&sslcert=/client.crt&sslkey=/client.key"
|
||||
|
||||
cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url)
|
||||
|
||||
assert self._query(cleaned) == {
|
||||
"sslmode": "verify-full",
|
||||
"sslrootcert": "/ca.pem",
|
||||
"sslcert": "/client.crt",
|
||||
"sslkey": "/client.key",
|
||||
}
|
||||
|
||||
def test_an_explicit_sslrootcert_wins_over_the_prisma_sslcert(self):
|
||||
url = "postgresql://writer/db?sslmode=require&sslrootcert=/ca.pem&sslcert=/pinned.pem&sslaccept=strict"
|
||||
|
||||
cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url)
|
||||
|
||||
assert self._query(cleaned) == {"sslmode": "verify-full", "sslrootcert": "/ca.pem"}
|
||||
|
||||
def test_prisma_only_params_are_dropped_and_plain_urls_pass_through(self):
|
||||
url = "postgresql://u:p@pooler:6543/db?schema=tenant&pgbouncer=true&connection_limit=5&connect_timeout=3"
|
||||
|
||||
cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url)
|
||||
|
||||
assert cleaned == "postgresql://u:p@pooler:6543/db?connect_timeout=3"
|
||||
assert (
|
||||
ProxyExtrasDBManager._strip_prisma_query_params("postgresql://u:p@writer/db")
|
||||
== "postgresql://u:p@writer/db"
|
||||
)
|
||||
|
||||
|
||||
class TestBuildRequestLogIndexes:
|
||||
"""The migration job hands the index build the direct database URL and the schema
|
||||
the migrations target, waits for it, and reports its result."""
|
||||
|
|
|
|||
1
tests/unit/llms/laya/__init__.py
Normal file
1
tests/unit/llms/laya/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
|
||||
60
tests/unit/llms/laya/test_common_utils.py
Normal file
60
tests/unit/llms/laya/test_common_utils.py
Normal file
|
|
@ -0,0 +1,60 @@
|
|||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.llms.laya.common_utils import laya_connection, laya_response_model
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("base", "key", "expected_base", "expected_key"),
|
||||
[
|
||||
(None, None, "http://laya.test/root", "laya-env-key"),
|
||||
("http://custom.test/", None, "http://custom.test", None),
|
||||
("http://custom.test/", "explicit-key", "http://custom.test", "explicit-key"),
|
||||
],
|
||||
)
|
||||
def test_laya_credentials_stay_with_their_configured_destination(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
base: str | None,
|
||||
key: str | None,
|
||||
expected_base: str,
|
||||
expected_key: str | None,
|
||||
) -> None:
|
||||
monkeypatch.setenv("LAYA_API_BASE", "http://laya.test/root/")
|
||||
monkeypatch.setenv("LAYA_API_KEY", "laya-env-key")
|
||||
monkeypatch.setenv("TYPESAFE_API_KEY", "never-send-this")
|
||||
connection: Final = laya_connection(base, key)
|
||||
assert (connection.api_base, connection.api_key) == (expected_base, expected_key)
|
||||
assert "key" not in repr(connection)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"base",
|
||||
["", "ftp://laya.test", "http://user:password@laya.test", "https://laya.test?key=x", "http://laya.test/#x"],
|
||||
)
|
||||
def test_laya_rejects_ambiguous_server_urls(base: str) -> None:
|
||||
with pytest.raises(ValueError, match="Laya"):
|
||||
laya_connection(base)
|
||||
|
||||
|
||||
def test_laya_missing_server_does_not_fall_back_to_typesafe(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.delenv("LAYA_API_BASE", raising=False)
|
||||
monkeypatch.setenv("TYPESAFE_API_BASE", "https://typesafe.test")
|
||||
with pytest.raises(ValueError, match="LAYA_API_BASE"):
|
||||
laya_connection()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("routing", "requested", "expected"),
|
||||
[
|
||||
({"model": "multilingual"}, "english", "multilingual"),
|
||||
(None, "english", "english"),
|
||||
({"model": 42}, "english", "english"),
|
||||
(None, None, "unknown"),
|
||||
],
|
||||
)
|
||||
def test_laya_identity_tracks_the_checkpoint_not_the_shared_agent_name(
|
||||
routing: Mapping[str, object] | None, requested: str | None, expected: str
|
||||
) -> None:
|
||||
assert laya_response_model({"model": "laya-rl-agent", "routing": routing}, requested) == expected
|
||||
|
|
@ -8,7 +8,7 @@ from typing import Optional
|
|||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import Request
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_utils import (
|
||||
|
|
@ -463,6 +463,24 @@ def test_get_model_from_request_no_request_extracts_model():
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", ["english", "multilingual", "typed-decisions"])
|
||||
@pytest.mark.parametrize("route", ["/laya/v1/systemone", "/laya/v1/systemone/"])
|
||||
def test_laya_native_model_uses_the_classifier_permission_identity(model: str, route: str) -> None:
|
||||
assert get_model_from_request(request_data={"model": model}, route=route) == f"laya/{model}"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", [None, "", "auto", "laya/english", "unknown", ["english"], 7])
|
||||
def test_laya_native_model_cannot_implicitly_select_an_unauthorized_checkpoint(model: object) -> None:
|
||||
with pytest.raises(HTTPException) as denied:
|
||||
get_model_from_request(request_data={"model": model}, route="/laya/v1/systemone")
|
||||
assert denied.value.status_code == 400
|
||||
|
||||
|
||||
def test_laya_model_normalization_does_not_change_other_provider_routes() -> None:
|
||||
assert get_model_from_request(request_data={"model": "jev-latest"}, route="/typesafe/v1/systemone") == "jev-latest"
|
||||
assert get_model_from_request(request_data={}, route="/laya/health") is None
|
||||
|
||||
|
||||
def _cache_prediction_router():
|
||||
from litellm.router import Router
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -619,6 +619,18 @@ def test_classifier_plugin_is_not_settable_over_http():
|
|||
_request("what is 2+2", classifier_type="custom", classifier_plugin="my_module.instance")
|
||||
|
||||
|
||||
def _benchmark_db(rows: Sequence[Mapping[str, object]], recorded: float | None = None) -> SimpleNamespace:
|
||||
"""The joined benchmark statement returns the rows as given; any other statement is the Overall total."""
|
||||
from litellm.proxy.db.autorouter_session_rollup import AUTOROUTER_BENCHMARKS_SQL
|
||||
|
||||
total: Final = recorded if recorded is not None else sum(float(row.get("saved_spend") or 0.0) for row in rows)
|
||||
|
||||
async def query_raw(sql: str, *params: object) -> Sequence[Mapping[str, object]]:
|
||||
return rows if sql == AUTOROUTER_BENCHMARKS_SQL else ({"saved": total},)
|
||||
|
||||
return SimpleNamespace(db=SimpleNamespace(query_raw=AsyncMock(side_effect=query_raw)))
|
||||
|
||||
|
||||
class TestAutoRouterBenchmarks:
|
||||
from litellm.proxy.management_endpoints.auto_router_endpoints import _SessionAggRow
|
||||
|
||||
|
|
@ -635,15 +647,12 @@ class TestAutoRouterBenchmarks:
|
|||
rows: Sequence[Mapping[str, object]],
|
||||
model_list: Sequence[object],
|
||||
api_key: str | None = None,
|
||||
recorded: float | None = None,
|
||||
) -> AutoRouterBenchmarksResponse:
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_benchmarks
|
||||
|
||||
class _DB:
|
||||
async def query_raw(self, sql: str, *params: object):
|
||||
return rows
|
||||
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", type("P", (), {"db": _DB()})())
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", _benchmark_db(rows, recorded))
|
||||
monkeypatch.setattr(proxy_server, "llm_router", type("R", (), {"model_list": model_list})())
|
||||
return await get_auto_router_benchmarks(
|
||||
user_api_key_dict=ADMIN,
|
||||
|
|
@ -657,6 +666,7 @@ class TestAutoRouterBenchmarks:
|
|||
router_type="complexity",
|
||||
tier_turns={},
|
||||
sessions=4,
|
||||
session_turns=40,
|
||||
turns=40,
|
||||
unordered_turns=1,
|
||||
covered_turns=38,
|
||||
|
|
@ -703,7 +713,6 @@ class TestAutoRouterBenchmarks:
|
|||
assert totals.baseline_spend == 40.0
|
||||
assert totals.saved_pct == 75.0
|
||||
assert totals.savings_estimated_classifier_cost == 0.4
|
||||
assert totals.saved_per_session == 7.5
|
||||
assert totals.cache.coverage_pct == 95.0
|
||||
assert totals.cache.hit_rate_pct == pytest.approx(73.7)
|
||||
assert totals.cache.same_model.hit_rate_pct == 95.0
|
||||
|
|
@ -742,7 +751,6 @@ class TestAutoRouterBenchmarks:
|
|||
assert (totals.spend, totals.saved_spend, totals.baseline_spend, totals.saved_pct) == (10.0, 30.0, 40.0, 75.0)
|
||||
assert (totals.savings_estimated_turns, totals.savings_estimated_actual_spend) == (40, 10.0)
|
||||
assert totals.savings_estimated_classifier_cost == 0.4
|
||||
assert totals.saved_per_session == 7.5
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("router_type, saved", [("adaptive", 0.0), ("quality", 0.0), ("quality", 2.0)])
|
||||
|
|
@ -772,9 +780,50 @@ class TestAutoRouterBenchmarks:
|
|||
totals: Final = response.totals
|
||||
assert (totals.turns, totals.spend) == (50, 13.0)
|
||||
assert (totals.savings_estimated_turns, totals.savings_estimated_actual_spend) == (40, 10.0)
|
||||
assert (totals.saved_spend, totals.baseline_spend, totals.saved_pct) == (30.0, 40.0, 75.0)
|
||||
assert totals.unattributed_saved_spend is None
|
||||
assert (totals.saved_spend, totals.baseline_spend, totals.saved_pct) == (
|
||||
(30.0, 40.0, 75.0) if saved == 0.0 else (32.0, None, None)
|
||||
)
|
||||
assert totals.savings_estimated_classifier_cost == 0.4
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("recorded, unattributed", [(30.0, None), (33.0, 3.0), (27.0, -3.0)])
|
||||
async def test_the_headline_is_the_overall_daily_total_and_untracked_savings_void_the_baseline(
|
||||
self, recorded: float, unattributed: float | None, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
response: Final = await self._benchmarks(
|
||||
monkeypatch, rows=[self.ROW.model_dump()], model_list=[], recorded=recorded
|
||||
)
|
||||
totals: Final = response.totals
|
||||
assert (totals.saved_spend, totals.unattributed_saved_spend) == (recorded, unattributed)
|
||||
assert (totals.baseline_spend, totals.saved_pct) == ((40.0, 75.0) if unattributed is None else (None, None))
|
||||
group: Final = response.groups[0]
|
||||
assert group.saved_spend == 30.0
|
||||
assert (group.baseline_spend, group.saved_pct) == ((40.0, 75.0) if unattributed is None else (None, None))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_window_holding_only_untracked_history_shows_no_router_baseline(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
history_only: Final = self.ROW.model_dump(
|
||||
exclude={
|
||||
"turns",
|
||||
"spend",
|
||||
"saved_spend",
|
||||
"savings_estimated_turns",
|
||||
"savings_estimated_actual_spend",
|
||||
"savings_estimated_classifier_cost",
|
||||
"savings_estimated_saved_spend",
|
||||
"classifier_cost",
|
||||
"classifier_cost_recorded_turns",
|
||||
}
|
||||
)
|
||||
response: Final = await self._benchmarks(monkeypatch, rows=[history_only], model_list=[], recorded=3.0)
|
||||
assert (response.totals.saved_spend, response.totals.unattributed_saved_spend) == (3.0, 3.0)
|
||||
group: Final = response.groups[0]
|
||||
assert (group.sessions, group.turns, group.saved_spend) == (4, 0, 0.0)
|
||||
assert (group.baseline_spend, group.saved_pct) == (None, None)
|
||||
|
||||
def test_an_empty_window_folds_to_zeros(self):
|
||||
from litellm.proxy.management_endpoints.auto_router_endpoints import (
|
||||
_benchmark_totals,
|
||||
|
|
@ -805,7 +854,7 @@ class TestAutoRouterBenchmarks:
|
|||
"savings_estimated_classifier_cost": 0.0,
|
||||
}
|
||||
)
|
||||
summed = _summed_agg_row([self.ROW, other])
|
||||
summed = _summed_agg_row([self.ROW, other.model_copy(update={"session_turns": 10})])
|
||||
totals = _benchmark_totals(summed)
|
||||
assert summed.sessions == 5
|
||||
assert summed.turns == 50
|
||||
|
|
@ -868,6 +917,28 @@ class TestAutoRouterBenchmarks:
|
|||
assert response.status_code == 422
|
||||
query.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_an_empty_key_filter_is_rejected_before_querying_deployment_data(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
import httpx
|
||||
from fastapi import FastAPI
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_benchmarks
|
||||
|
||||
query: Final = AsyncMock(return_value=[])
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=SimpleNamespace(query_raw=query)))
|
||||
app: Final = FastAPI()
|
||||
app.get("/auto_router/benchmarks")(get_auto_router_benchmarks)
|
||||
app.dependency_overrides[user_api_key_auth] = lambda: ADMIN
|
||||
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://test") as client:
|
||||
response: Final = await client.get("/auto_router/benchmarks", params={"api_key": ""})
|
||||
|
||||
assert response.status_code == 422
|
||||
query.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_reversed_window_is_rejected(self, monkeypatch: pytest.MonkeyPatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
|
@ -891,15 +962,8 @@ class TestAutoRouterBenchmarks:
|
|||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_benchmarks
|
||||
|
||||
captured: dict = {}
|
||||
|
||||
class _DB:
|
||||
async def query_raw(self, sql: str, *params: object):
|
||||
captured["sql"] = sql
|
||||
captured["params"] = params
|
||||
return [TestAutoRouterBenchmarks.ROW.model_dump()]
|
||||
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", type("P", (), {"db": _DB()})())
|
||||
prisma_client: Final = _benchmark_db([TestAutoRouterBenchmarks.ROW.model_dump()])
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma_client)
|
||||
|
||||
response = await get_auto_router_benchmarks(
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=role, api_key="sk-admin", user_id="viewer"),
|
||||
|
|
@ -908,7 +972,11 @@ class TestAutoRouterBenchmarks:
|
|||
api_key="key-hash",
|
||||
user_id=user_id,
|
||||
)
|
||||
assert captured["params"] == ("2026-07-01T00:00:00", "2026-08-02T00:00:00", "key-hash", user_id)
|
||||
params: Final = tuple(call.args[1:] for call in prisma_client.db.query_raw.await_args_list)
|
||||
assert params == (
|
||||
("2026-07-01T00:00:00", "2026-08-02T00:00:00", "key-hash", user_id, "2026-07-01", "2026-08-01"),
|
||||
("2026-07-01", "2026-08-01", *(([user_id],) if user_id else ()), ["key-hash"]),
|
||||
)
|
||||
assert response.routers_in_scope == 1
|
||||
assert response.groups[0].router_name == "live-auto"
|
||||
assert response.groups[0].saved_pct == response.totals.saved_pct == 75.0
|
||||
|
|
@ -946,7 +1014,6 @@ class TestAutoRouterBenchmarks:
|
|||
assert response.totals.saved_spend == 29.5
|
||||
assert response.totals.baseline_spend == 41.5
|
||||
assert response.totals.saved_pct == 71.1
|
||||
assert response.totals.saved_per_session == 5.9
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
|
|
@ -958,11 +1025,9 @@ class TestAutoRouterBenchmarks:
|
|||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_benchmarks
|
||||
|
||||
class _DB:
|
||||
async def query_raw(self, sql: str, *params: object):
|
||||
return [{**TestAutoRouterBenchmarks.ROW.model_dump(), "tier_turns": wire_value}]
|
||||
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", type("P", (), {"db": _DB()})())
|
||||
monkeypatch.setattr(
|
||||
proxy_server, "prisma_client", _benchmark_db([{**TestAutoRouterBenchmarks.ROW.model_dump(), "tier_turns": wire_value}])
|
||||
)
|
||||
|
||||
response = await get_auto_router_benchmarks(
|
||||
user_api_key_dict=ADMIN,
|
||||
|
|
@ -1006,7 +1071,7 @@ class TestAutoRouterBenchmarks:
|
|||
0.0,
|
||||
0.0,
|
||||
)
|
||||
assert (idle.saved_pct, idle.saved_per_session, idle.avg_turns_per_session) == (0.0, 0.0, 0.0)
|
||||
assert (idle.saved_pct, idle.avg_turns_per_session) == (0.0, 0.0)
|
||||
assert (idle.cache.hit_rate_pct, idle.cache.coverage_pct) == (0.0, 0.0)
|
||||
assert idle.cache.same_model.turns == idle.cache.return_to_tier.hits == 0
|
||||
assert idle.tier_turns == {}
|
||||
|
|
@ -3724,3 +3789,18 @@ async def test_availability_waits_for_the_first_complete_catalog(monkeypatch):
|
|||
with pytest.raises(HTTPException) as error:
|
||||
await auto_router_endpoints.get_auto_router_availability(AutoRouterAvailabilityRequest(), ADMIN)
|
||||
assert error.value.status_code == 503
|
||||
|
||||
|
||||
class TestPerSessionAverages:
|
||||
@pytest.mark.parametrize(
|
||||
"sessions, turns, expected",
|
||||
[(4, 40, (10.0, 100.0, 1000.0)), (0, 0, (0.0, 0.0, 0.0)), (0, 3, (None, None, None))],
|
||||
)
|
||||
def test_requests_without_session_rows_have_unknown_averages_not_zero(
|
||||
self, sessions: int, turns: int, expected: tuple[float | None, ...]
|
||||
) -> None:
|
||||
from litellm.proxy.management_endpoints.auto_router_endpoints import _benchmark_totals
|
||||
|
||||
row: Final = TestAutoRouterBenchmarks.ROW.model_copy(update={"sessions": sessions, "turns": turns})
|
||||
totals: Final = _benchmark_totals(row)
|
||||
assert (totals.avg_turns_per_session, totals.avg_session_seconds, totals.avg_tokens_per_session) == expected
|
||||
|
|
|
|||
|
|
@ -8,6 +8,8 @@ from typing import Dict, Final, Optional
|
|||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from fastapi.encoders import jsonable_encoder
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from litellm._uuid import uuid
|
||||
|
|
@ -7515,6 +7517,96 @@ class TestTeamMemberAutoRouterWrites:
|
|||
"model_info": {"id": "allowed-id"},
|
||||
}])
|
||||
|
||||
@staticmethod
|
||||
def _classifier_config(classifier: Mapping[str, object], legacy: bool) -> Mapping[str, object]:
|
||||
return {
|
||||
"classifier_type": "jev" if legacy else "oss_classifier",
|
||||
"tiers": {"SIMPLE": "allowed"},
|
||||
"jev_classifier_config" if legacy else "opensource_classifier_config": classifier,
|
||||
}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("team_id", [None, "member-team"])
|
||||
@pytest.mark.parametrize(
|
||||
"legacy,provider,model",
|
||||
[(True, "typesafe", "jev-latest"), (False, "jev", "jev-latest"), (True, "laya", "english"), (False, "laya", "english")],
|
||||
)
|
||||
async def test_classifier_create_stores_only_canonical_configuration(
|
||||
self, team_id: str | None, legacy: bool, provider: str, model: str
|
||||
) -> None:
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import add_new_model
|
||||
|
||||
row: Final = self._row()
|
||||
database: Final = self._database(self._team(), row)
|
||||
classifier: Final = {
|
||||
"provider": provider, "model": model,
|
||||
"api_base": "https://decision.test", "api_key": "stored-secret",
|
||||
}
|
||||
deployment: Final = Deployment(
|
||||
model_name="new-classifier-router",
|
||||
litellm_params=LiteLLM_Params(
|
||||
model="auto_router/complexity_router",
|
||||
complexity_router_config=self._classifier_config(classifier, legacy),
|
||||
),
|
||||
model_info=ModelInfo(id=row.model_id, team_id=team_id),
|
||||
)
|
||||
actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
with (
|
||||
self._environment(database, row),
|
||||
patch("litellm.proxy.proxy_server.proxy_config.add_deployment", new=AsyncMock(return_value=ReconcileOutcome( # test-quality-ok: [TQ008] model reload I/O boundary
|
||||
still_desired=frozenset((row.model_id,)), live_after=frozenset((row.model_id,))
|
||||
))),
|
||||
patch("litellm.proxy.management_endpoints.model_management_endpoints.append_team_models", new=AsyncMock()), # test-quality-ok: [TQ008] team allowlist persistence boundary
|
||||
):
|
||||
await add_new_model(deployment, actor)
|
||||
written: Final = database.db.litellm_proxymodeltable.create.await_args.kwargs["data"]
|
||||
saved: Final = json.loads(written["litellm_params"])["complexity_router_config"]
|
||||
assert saved == {
|
||||
"classifier_type": "oss_classifier",
|
||||
"tiers": {"SIMPLE": "allowed"},
|
||||
"opensource_classifier_config": {**classifier, "provider": "laya" if provider == "laya" else "jev"},
|
||||
}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("endpoint", ["create", "patch", "legacy"])
|
||||
@pytest.mark.parametrize("legacy_config", [None, {"provider": "laya", "model": "english"}])
|
||||
async def test_ambiguous_classifier_blocks_are_rejected_before_persistence(
|
||||
self, endpoint: str, legacy_config: Mapping[str, object] | None
|
||||
) -> None:
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import add_new_model
|
||||
|
||||
row: Final = self._row()
|
||||
database: Final = self._database(self._team(), row)
|
||||
config: Final = {
|
||||
**self._classifier_config({"provider": "laya", "model": "english"}, False),
|
||||
"jev_classifier_config": legacy_config,
|
||||
}
|
||||
actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
request: Final = updateDeployment(
|
||||
litellm_params=updateLiteLLMParams(complexity_router_config=config), model_info=ModelInfo(id=row.model_id),
|
||||
)
|
||||
operation: Final = (
|
||||
add_new_model(
|
||||
Deployment(
|
||||
model_name="ambiguous-classifier-router",
|
||||
litellm_params=LiteLLM_Params(model="auto_router/complexity_router", complexity_router_config=config),
|
||||
model_info=ModelInfo(id=row.model_id),
|
||||
),
|
||||
actor,
|
||||
)
|
||||
if endpoint == "create"
|
||||
else patch_model(row.model_id, request, actor)
|
||||
if endpoint == "patch"
|
||||
else update_model(request, actor)
|
||||
)
|
||||
with self._environment(database, row), pytest.raises(ProxyException) as denied:
|
||||
await operation
|
||||
assert denied.value.code == "400"
|
||||
assert "opensource_classifier_config" in denied.value.message
|
||||
assert "jev_classifier_config" in denied.value.message
|
||||
database.db.litellm_proxymodeltable.create.assert_not_awaited()
|
||||
database.db.litellm_proxymodeltable.update.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("endpoint,change", [("patch", "config"), ("legacy", "strategy"), ("patch", "unrelated")])
|
||||
async def test_admin_router_changes_release_member_scope(self, endpoint: str, change: str) -> None:
|
||||
|
|
@ -7546,15 +7638,16 @@ class TestTeamMemberAutoRouterWrites:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("endpoint", ["patch", "legacy"])
|
||||
@pytest.mark.parametrize("stored_legacy,supplied_legacy", [(True, True), (True, False), (False, True), (False, False)])
|
||||
@pytest.mark.parametrize("change", ["save", "rotate", "move", "move-without-key", "reset", "heuristic"])
|
||||
async def test_jev_dashboard_save_preserves_server_transport(self, endpoint: str, change: str) -> None:
|
||||
async def test_jev_dashboard_save_preserves_server_transport(
|
||||
self, endpoint: str, change: str, stored_legacy: bool, supplied_legacy: bool
|
||||
) -> None:
|
||||
original: Final = self._row()
|
||||
transport: Final = {"api_key": "synthetic-original-jev-key", "api_base": "https://jev.example.com"}
|
||||
stored_config: Final = {
|
||||
"classifier_type": "jev",
|
||||
"tiers": {"SIMPLE": "allowed"},
|
||||
"jev_classifier_config": {**transport, "instructions": "Old instructions", "timeout_ms": 6100},
|
||||
}
|
||||
stored_config: Final = self._classifier_config(
|
||||
{**transport, "instructions": "Old instructions", "timeout_ms": 6100}, stored_legacy
|
||||
)
|
||||
row: Final = original.model_copy(
|
||||
update={
|
||||
"litellm_params": {
|
||||
|
|
@ -7572,11 +7665,11 @@ class TestTeamMemberAutoRouterWrites:
|
|||
"reset": {"api_key": None, "api_base": None},
|
||||
"heuristic": {},
|
||||
}[change]
|
||||
config: Final = {
|
||||
"tiers": {"SIMPLE": "allowed"},
|
||||
"classifier_type": "heuristic" if change == "heuristic" else "jev",
|
||||
**({} if change == "heuristic" else {"jev_classifier_config": {"timeout_ms": 8100, **overrides}}),
|
||||
}
|
||||
config: Final = (
|
||||
{"tiers": {"SIMPLE": "allowed"}, "classifier_type": "heuristic"}
|
||||
if change == "heuristic"
|
||||
else self._classifier_config({"timeout_ms": 8100, **overrides}, supplied_legacy)
|
||||
)
|
||||
request: Final = updateDeployment(
|
||||
litellm_params=updateLiteLLMParams(complexity_router_config=config),
|
||||
model_info=ModelInfo(id=row.model_id),
|
||||
|
|
@ -7597,12 +7690,179 @@ class TestTeamMemberAutoRouterWrites:
|
|||
expected: Final = (
|
||||
config
|
||||
if change == "heuristic"
|
||||
else {**config, "jev_classifier_config": {**transport, "timeout_ms": 8100, **overrides}}
|
||||
else {
|
||||
"classifier_type": "oss_classifier",
|
||||
"tiers": {"SIMPLE": "allowed"},
|
||||
"opensource_classifier_config": {**transport, "timeout_ms": 8100, **overrides},
|
||||
}
|
||||
)
|
||||
assert saved == expected
|
||||
assert row.litellm_params["complexity_router_config"] == stored_config
|
||||
assert request.litellm_params.complexity_router_config == config
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("endpoint", ["patch", "legacy"])
|
||||
@pytest.mark.parametrize("stored_legacy,supplied_legacy", [(True, True), (True, False), (False, True), (False, False)])
|
||||
@pytest.mark.parametrize(
|
||||
"stored_provider,stored_base,supplied,expected_transport",
|
||||
[
|
||||
("laya", "https://decision.test", {"provider": "laya", "model": "english"}, {"api_base": "https://decision.test", "api_key": "stored-secret"}),
|
||||
("laya", "https://decision.test", {"provider": "laya", "model": "english", "api_base": "https://decision.test"}, {"api_base": "https://decision.test", "api_key": "stored-secret"}),
|
||||
(
|
||||
"laya",
|
||||
"https://decision.test",
|
||||
{"provider": "laya", "model": "english", "api_key": None},
|
||||
{"api_base": "https://decision.test"},
|
||||
),
|
||||
("laya", "https://decision.test", {"provider": "laya", "model": "english", "api_base": "https://new.test"}, {}),
|
||||
("laya", "https://decision.test", {"provider": "laya", "model": "english", "api_base": None}, {}),
|
||||
("laya", None, {"provider": "laya", "model": "english", "api_base": None}, {}),
|
||||
("laya", "https://decision.test", {"provider": "typesafe", "model": "jev-latest"}, {}),
|
||||
(
|
||||
"laya", "https://decision.test", {"model": "english", "timeout_ms": 8100},
|
||||
{"provider": "laya", "api_base": "https://decision.test", "api_key": "stored-secret"},
|
||||
),
|
||||
("typesafe", "https://decision.test", {"provider": "laya", "model": "english"}, {}),
|
||||
(
|
||||
"typesafe", "https://decision.test", {"provider": "jev", "model": "jev-latest"},
|
||||
{"api_base": "https://decision.test", "api_key": "stored-secret"},
|
||||
),
|
||||
(
|
||||
"jev", "https://decision.test", {"provider": "typesafe", "model": "jev-latest"},
|
||||
{"api_base": "https://decision.test", "api_key": "stored-secret"},
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_decision_provider_changes_cannot_reuse_a_stored_key(
|
||||
self, endpoint: str, stored_provider: str, stored_base: str | None,
|
||||
supplied: Mapping[str, object], expected_transport: Mapping[str, object],
|
||||
stored_legacy: bool, supplied_legacy: bool,
|
||||
) -> None:
|
||||
original: Final = self._row()
|
||||
row: Final = original.model_copy(update={"litellm_params": {
|
||||
"model": "auto_router/complexity_router",
|
||||
"complexity_router_config": self._classifier_config(
|
||||
{
|
||||
"provider": stored_provider, "model": "english" if stored_provider == "laya" else "jev-latest",
|
||||
"api_base": stored_base, "api_key": "stored-secret",
|
||||
},
|
||||
stored_legacy,
|
||||
),
|
||||
}})
|
||||
database: Final = self._database(self._team(), row)
|
||||
config: Final = self._classifier_config(supplied, supplied_legacy)
|
||||
request: Final = updateDeployment(
|
||||
litellm_params=updateLiteLLMParams(complexity_router_config=config), model_info=ModelInfo(id=row.model_id),
|
||||
)
|
||||
actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
with self._environment(database, row):
|
||||
await (patch_model(row.model_id, request, actor) if endpoint == "patch" else update_model(request, actor))
|
||||
written: Final = database.db.litellm_proxymodeltable.update.await_args.kwargs["data"]
|
||||
saved: Final = json.loads(written["litellm_params"])["complexity_router_config"]
|
||||
expected_provider: Final = supplied.get("provider", stored_provider)
|
||||
assert saved == {
|
||||
"classifier_type": "oss_classifier",
|
||||
"tiers": {"SIMPLE": "allowed"},
|
||||
"opensource_classifier_config": {
|
||||
**expected_transport, **supplied,
|
||||
"provider": "jev" if expected_provider == "typesafe" else expected_provider,
|
||||
},
|
||||
}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("endpoint", ["patch", "legacy"])
|
||||
@pytest.mark.parametrize(
|
||||
"string_params,reset_field,config_shape",
|
||||
[
|
||||
(False, None, "full"), (True, None, "full"), (False, "api_key", "full"),
|
||||
(False, "api_base", "full"), (False, None, "omit-provider"),
|
||||
(False, None, "omit-config"), (False, None, "null-config"),
|
||||
],
|
||||
)
|
||||
async def test_member_save_protects_stored_classifier_connection(
|
||||
self, endpoint: str, string_params: bool, reset_field: str | None, config_shape: str
|
||||
) -> None:
|
||||
original: Final = self._row()
|
||||
config: Final = {
|
||||
"classifier_type": "jev", "tiers": {"SIMPLE": "allowed"},
|
||||
"jev_classifier_config": {"provider": "laya", "model": "english"},
|
||||
}
|
||||
secret_params: Final = {
|
||||
"model": "auto_router/complexity_router",
|
||||
"complexity_router_config": {
|
||||
**config, "jev_classifier_config": {
|
||||
**config["jev_classifier_config"], "api_key": "retained-laya-secret", "api_base": "https://laya.test",
|
||||
},
|
||||
},
|
||||
}
|
||||
row: Final = original.model_copy(update={"litellm_params": secret_params})
|
||||
team: Final = self._team().model_copy(update={"models": ["allowed", "laya/english"]})
|
||||
database: Final = self._database(team, row)
|
||||
database.transaction.litellm_proxymodeltable.update.return_value = row.model_copy(
|
||||
update={"litellm_params": json.dumps(secret_params) if string_params else secret_params}
|
||||
)
|
||||
supplied_config: Final = {
|
||||
**config, "jev_classifier_config": {
|
||||
**{
|
||||
key: value for key, value in config["jev_classifier_config"].items()
|
||||
if key != "provider" or config_shape != "omit-provider"
|
||||
},
|
||||
**({reset_field: None} if reset_field is not None else {}),
|
||||
},
|
||||
}
|
||||
patch_params: Final = (
|
||||
{"complexity_router_default_model": "allowed"}
|
||||
if config_shape == "omit-config"
|
||||
else {"complexity_router_config": None, "complexity_router_default_model": "allowed"}
|
||||
if config_shape == "null-config"
|
||||
else {"complexity_router_config": supplied_config}
|
||||
)
|
||||
request: Final = updateDeployment(
|
||||
litellm_params=updateLiteLLMParams.model_validate(patch_params),
|
||||
model_info=ModelInfo(id=row.model_id, team_id="member-team"),
|
||||
)
|
||||
actor: Final = UserAPIKeyAuth(
|
||||
user_id="owner", user_role=LitellmUserRoles.INTERNAL_USER, models=["allowed", "laya/english"], config={"timeout": 60},
|
||||
)
|
||||
with self._environment(database, row):
|
||||
if reset_field is not None:
|
||||
expected_error: Final = HTTPException if endpoint == "patch" else ProxyException
|
||||
with pytest.raises(expected_error, match="Team members cannot change classifier connections") as denied:
|
||||
await (
|
||||
patch_model(row.model_id, request, actor) if endpoint == "patch" else update_model(request, actor)
|
||||
)
|
||||
assert (
|
||||
denied.value.status_code if isinstance(denied.value, HTTPException) else int(denied.value.code)
|
||||
) == 403
|
||||
database.transaction.litellm_proxymodeltable.update.assert_not_awaited()
|
||||
assert row.litellm_params == secret_params
|
||||
return
|
||||
response: Final = await (patch_model(row.model_id, request, actor) if endpoint == "patch" else update_model(request, actor))
|
||||
written: Final = database.transaction.litellm_proxymodeltable.update.await_args.kwargs["data"]
|
||||
saved_config: Final = json.loads(written["litellm_params"])["complexity_router_config"]
|
||||
untouched: Final = config_shape in ("omit-config", "null-config")
|
||||
saved: Final = saved_config["jev_classifier_config" if untouched else "opensource_classifier_config"]
|
||||
assert saved == secret_params["complexity_router_config"]["jev_classifier_config"]
|
||||
assert saved_config["classifier_type"] == ("jev" if untouched else "oss_classifier")
|
||||
if untouched:
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
|
||||
|
||||
assert decrypt_value_helper(
|
||||
json.loads(written["litellm_params"])["complexity_router_default_model"],
|
||||
key="complexity_router_default_model", return_original_value=True,
|
||||
) == "allowed"
|
||||
response_payload: Final = jsonable_encoder(response)
|
||||
assert "retained-laya-secret" not in json.dumps(response_payload)
|
||||
response_params: Final = json.loads(response_payload["litellm_params"]) if string_params else response_payload["litellm_params"]
|
||||
assert response_params == {
|
||||
**secret_params, "complexity_router_config": {
|
||||
**config, "jev_classifier_config": {
|
||||
**config["jev_classifier_config"], "api_key": "REDACTED", "api_base": "https://laya.test",
|
||||
},
|
||||
},
|
||||
}
|
||||
assert "retained-laya-secret" in row.model_dump_json()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("endpoint", ["patch", "legacy"])
|
||||
@pytest.mark.parametrize("access", ["owner", "peer", "limited-key"])
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import json
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
|
|
@ -137,33 +138,48 @@ def test_tier_config_is_normalized_and_unknown_router_extras_are_rejected() -> N
|
|||
@pytest.mark.parametrize(
|
||||
("jev_override", "rejected_at"),
|
||||
[
|
||||
({"api_base": "https://collector.invalid"}, "jev_classifier_config"),
|
||||
({"api_base": "https://collector.invalid"}, "opensource_classifier_config"),
|
||||
({"api_key": "sk-member"}, "api_key"),
|
||||
({"api_base": "https://collector.invalid", "api_key": "sk-member"}, "api_key"),
|
||||
({"api_base": "https://collector.invalid", "api_key": ""}, "jev_classifier_config.api_key"),
|
||||
({"api_base": "https://collector.invalid", "api_key": ""}, "opensource_classifier_config.api_key"),
|
||||
({"provider": "laya", "model": "english", "api_base": "https://collector.invalid"}, "api_base"),
|
||||
({"provider": "laya", "model": "english", "api_key": "sk-member"}, "api_key"),
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize("legacy", [False, True])
|
||||
def test_members_cannot_move_the_jev_classifier_off_the_proxys_typesafe_account(
|
||||
jev_override: Mapping[str, str], rejected_at: str
|
||||
jev_override: Mapping[str, str], rejected_at: str, legacy: bool
|
||||
) -> None:
|
||||
with pytest.raises(HTTPException) as denied:
|
||||
validate_member_auto_router_config(
|
||||
{"tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev", "jev_classifier_config": jev_override}
|
||||
{
|
||||
"tiers": {"SIMPLE": "allowed"},
|
||||
"classifier_type": "jev" if legacy else "oss_classifier",
|
||||
"jev_classifier_config" if legacy else "opensource_classifier_config": jev_override,
|
||||
}
|
||||
)
|
||||
assert denied.value.status_code == 400
|
||||
assert denied.value.detail == f"Invalid member auto-router configuration at {rejected_at}."
|
||||
|
||||
|
||||
def test_members_can_still_tune_the_jev_classifier() -> None:
|
||||
@pytest.mark.parametrize(("provider", "model"), [("typesafe", "jev-preview"), ("laya", "english")])
|
||||
@pytest.mark.parametrize("legacy", [False, True])
|
||||
def test_members_can_still_tune_the_jev_classifier(provider: str, model: str, legacy: bool) -> None:
|
||||
validated: Final = validate_member_auto_router_config(
|
||||
{
|
||||
"tiers": {"SIMPLE": "allowed"},
|
||||
"classifier_type": "jev",
|
||||
"jev_classifier_config": {"model": "jev-preview", "timeout_ms": 500},
|
||||
"classifier_type": "jev" if legacy else "oss_classifier",
|
||||
"jev_classifier_config" if legacy else "opensource_classifier_config": {
|
||||
"provider": provider, "model": model, "timeout_ms": 500,
|
||||
},
|
||||
}
|
||||
)
|
||||
assert validated.jev_classifier_config is not None
|
||||
assert (validated.jev_classifier_config.model, validated.jev_classifier_config.timeout_ms) == ("jev-preview", 500)
|
||||
assert (
|
||||
validated.jev_classifier_config.provider,
|
||||
validated.jev_classifier_config.model,
|
||||
validated.jev_classifier_config.timeout_ms,
|
||||
) == ("jev" if provider == "typesafe" else provider, model, 500)
|
||||
assert validate_member_auto_router_config(validated.model_dump()).jev_classifier_config is not None
|
||||
|
||||
|
||||
|
|
@ -217,6 +233,92 @@ async def test_member_updates_restrict_fields_and_preserve_an_inherited_default(
|
|||
assert granted.default_model == "allowed"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"nested,expected_identity,restricted",
|
||||
[
|
||||
("omit-config", "laya/english", False),
|
||||
("omit-config", "laya/english", True),
|
||||
("omit-block", None, False),
|
||||
(None, None, False),
|
||||
({}, None, False),
|
||||
({"timeout_ms": 500}, None, False),
|
||||
({"model": "english", "timeout_ms": 500}, "laya/english", False),
|
||||
({"model": "english", "timeout_ms": 500}, "laya/english", True),
|
||||
({"model": "multilingual"}, "laya/multilingual", False),
|
||||
({"provider": "typesafe", "model": "jev-latest"}, "typesafe/jev-latest", False),
|
||||
({"provider": "typesafe", "model": "jev-latest"}, "typesafe/jev-latest", True),
|
||||
],
|
||||
)
|
||||
async def test_member_authorization_and_persistence_resolve_the_same_classifier(
|
||||
catalog: Router, monkeypatch: pytest.MonkeyPatch, nested: object, expected_identity: str | None, restricted: bool
|
||||
) -> None:
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_strategy_router_write_violation,
|
||||
update_db_model,
|
||||
)
|
||||
from litellm.types.management_endpoints.auto_router_endpoints import RequestComplexityRouterConfig
|
||||
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "member-router-test-salt")
|
||||
stored_config: Final = {
|
||||
"classifier_type": "jev", "tiers": {"SIMPLE": "allowed"},
|
||||
"jev_classifier_config": {
|
||||
"provider": "laya", "model": "english", "timeout_ms": 12000,
|
||||
"api_base": "https://laya.test", "api_key": "stored-classifier-key",
|
||||
},
|
||||
}
|
||||
existing: Final = Deployment(
|
||||
model_name="member-router",
|
||||
litellm_params=LiteLLM_Params(model="auto_router/complexity_router", complexity_router_config=stored_config),
|
||||
model_info=ModelInfo(id="router-a", team_id="team-a"), created_by="owner",
|
||||
)
|
||||
incoming_config: Final = (
|
||||
None if nested == "omit-config" else {
|
||||
"classifier_type": "jev", "tiers": {"SIMPLE": "allowed"},
|
||||
**({} if nested == "omit-block" else {"jev_classifier_config": nested}),
|
||||
}
|
||||
)
|
||||
patch: Final = updateDeployment.model_validate({"litellm_params": {
|
||||
"complexity_router_config": incoming_config, "complexity_router_default_model": "allowed",
|
||||
}})
|
||||
operation: Final = authorize_member_auto_router_write(
|
||||
incoming=patch, existing=existing, user_api_key_dict=_actor(
|
||||
models=["allowed"] if restricted or expected_identity is None else ["allowed", expected_identity],
|
||||
),
|
||||
team=_team(models=["allowed", "laya/english", "laya/multilingual", "typesafe/jev-latest"]),
|
||||
premium_user=True, prisma_client=_Client(), llm_router=catalog,
|
||||
)
|
||||
violation: Final = _strategy_router_write_violation(patch.litellm_params, existing.litellm_params)
|
||||
if expected_identity is None:
|
||||
assert violation is not None
|
||||
with pytest.raises(HTTPException) as rejected:
|
||||
await operation
|
||||
assert rejected.value.status_code == 400
|
||||
return
|
||||
assert violation is None
|
||||
if restricted:
|
||||
with pytest.raises(ProxyException, match=expected_identity):
|
||||
await operation
|
||||
return
|
||||
grant: Final = await operation
|
||||
persisted: Final = update_db_model(existing, patch)
|
||||
saved: Final = RequestComplexityRouterConfig.model_validate(
|
||||
json.loads(persisted["litellm_params"])["complexity_router_config"]
|
||||
)
|
||||
assert grant.config == saved
|
||||
assert saved.jev_classifier_config is not None
|
||||
assert (
|
||||
"typesafe" if saved.jev_classifier_config.provider == "jev" else saved.jev_classifier_config.provider
|
||||
) + f"/{saved.jev_classifier_config.model}" == expected_identity
|
||||
assert saved.jev_classifier_config.api_key == (
|
||||
"stored-classifier-key" if expected_identity.startswith("laya/") else None
|
||||
)
|
||||
assert saved.jev_classifier_config.timeout_ms == (
|
||||
12000 if nested == "omit-config" else 500 if nested == {"model": "english", "timeout_ms": 500} else 3000
|
||||
)
|
||||
assert existing.litellm_params.complexity_router_config == stored_config
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("target", ["missing", "nested"])
|
||||
async def test_member_dependencies_require_plain_configured_models(target: str) -> None:
|
||||
|
|
@ -246,13 +348,17 @@ async def test_member_dependencies_require_plain_configured_models(target: str)
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("restricted", ["key", "team", None])
|
||||
@pytest.mark.parametrize(("provider", "model"), [("typesafe", "jev-latest"), ("laya", "english")])
|
||||
async def test_jev_evaluation_requires_model_access_but_no_completion_deployment(
|
||||
catalog: Router, restricted: str | None
|
||||
catalog: Router, restricted: str | None, provider: str, model: str
|
||||
) -> None:
|
||||
permitted: Final = ["allowed", "typesafe/jev-latest"]
|
||||
permitted: Final = ["allowed", f"{provider}/{model}"]
|
||||
operation: Final = authorize_member_auto_router_dependencies(
|
||||
config=validate_member_auto_router_config(
|
||||
{"tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev", "jev_classifier_config": {}}
|
||||
{
|
||||
"tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev",
|
||||
"jev_classifier_config": {"provider": provider, "model": model},
|
||||
}
|
||||
),
|
||||
default_model=None,
|
||||
user_api_key_dict=_actor(models=["allowed"] if restricted == "key" else permitted),
|
||||
|
|
@ -261,17 +367,20 @@ async def test_jev_evaluation_requires_model_access_but_no_completion_deployment
|
|||
llm_router=catalog,
|
||||
)
|
||||
if restricted is not None:
|
||||
with pytest.raises(ProxyException, match="jev-latest"):
|
||||
with pytest.raises(ProxyException, match=model):
|
||||
await operation
|
||||
return
|
||||
await operation
|
||||
assert not catalog.get_model_list("typesafe/jev-latest")
|
||||
assert not catalog.get_model_list(f"{provider}/{model}")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("restricted", ["member", "project", "organization", None])
|
||||
async def test_jev_evaluation_obeys_each_containing_scope(catalog: Router, restricted: str | None) -> None:
|
||||
allowed: Final = ["allowed", "typesafe/jev-latest"]
|
||||
@pytest.mark.parametrize(("provider", "model"), [("typesafe", "jev-latest"), ("laya", "english")])
|
||||
async def test_jev_evaluation_obeys_each_containing_scope(
|
||||
catalog: Router, restricted: str | None, provider: str, model: str
|
||||
) -> None:
|
||||
allowed: Final = ["allowed", f"{provider}/{model}"]
|
||||
membership: Final = LiteLLM_TeamMembership.model_validate(
|
||||
{
|
||||
"user_id": "owner",
|
||||
|
|
@ -293,7 +402,10 @@ async def test_jev_evaluation_obeys_each_containing_scope(catalog: Router, restr
|
|||
)
|
||||
operation: Final = authorize_member_auto_router_dependencies(
|
||||
config=validate_member_auto_router_config(
|
||||
{"tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev", "jev_classifier_config": {}}
|
||||
{
|
||||
"tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev",
|
||||
"jev_classifier_config": {"provider": provider, "model": model},
|
||||
}
|
||||
),
|
||||
default_model=None,
|
||||
user_api_key_dict=_actor(models=allowed, project_id="project-a"),
|
||||
|
|
@ -303,8 +415,8 @@ async def test_jev_evaluation_obeys_each_containing_scope(catalog: Router, restr
|
|||
dependency_objects=MemberAutoRouterDependencyObjects(membership, organization, project),
|
||||
)
|
||||
if restricted is not None:
|
||||
with pytest.raises(ProxyException, match="jev-latest"):
|
||||
with pytest.raises(ProxyException, match=model):
|
||||
await operation
|
||||
return
|
||||
await operation
|
||||
assert not catalog.get_model_list("typesafe/jev-latest")
|
||||
assert not catalog.get_model_list(f"{provider}/{model}")
|
||||
|
|
|
|||
|
|
@ -1,10 +1,12 @@
|
|||
from datetime import datetime
|
||||
from typing import Final
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.typesafe_passthrough_logging_handler import (
|
||||
TypeSafePassthroughLoggingHandler,
|
||||
)
|
||||
|
|
@ -137,6 +139,84 @@ def test_success_handler_dispatches_to_typesafe_handler():
|
|||
assert normalized["kwargs"]["model"] == "typesafe/jev-1.13.0"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("guardrail_cost", [0.0, 0.25])
|
||||
@pytest.mark.parametrize("metadata_slot", ["metadata", "litellm_metadata"])
|
||||
@pytest.mark.parametrize("routing_model", ["multilingual", None])
|
||||
async def test_laya_gateway_accounts_for_checkpoint_usage_and_registered_cost(
|
||||
monkeypatch: pytest.MonkeyPatch, routing_model: str | None, metadata_slot: str, guardrail_cost: float
|
||||
) -> None:
|
||||
checkpoint: Final = routing_model or "english"
|
||||
model: Final = f"laya/{checkpoint}"
|
||||
input_rate: Final = 0.002
|
||||
output_rate: Final = 0.005
|
||||
monkeypatch.setitem(litellm.model_cost, model, {
|
||||
"input_cost_per_token": input_rate, "output_cost_per_token": output_rate,
|
||||
"litellm_provider": "laya", "mode": "evaluation",
|
||||
})
|
||||
start: Final = datetime.now()
|
||||
logging_obj: Final = Logging(
|
||||
model="english", messages=[], stream=False, call_type="pass_through_endpoint",
|
||||
start_time=start, litellm_call_id="laya-accounting", function_id="laya-accounting", kwargs={},
|
||||
)
|
||||
from fastapi import Request
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import HttpPassThroughEndpointHelpers
|
||||
|
||||
request: Final = Request({
|
||||
"type": "http", "method": "POST", "path": "/laya/v1/systemone",
|
||||
"headers": [], "query_string": b"",
|
||||
})
|
||||
auth: Final = UserAPIKeyAuth(
|
||||
api_key="laya-budget-key", token="laya-budget-key",
|
||||
model_max_budget={"laya/english": {"budget_limit": 0.01, "time_period": "1d"}},
|
||||
)
|
||||
request_body: Final = {"model": "english", metadata_slot: {"model_group": "unbounded-client-choice"}}
|
||||
logging_kwargs: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
|
||||
request=request, user_api_key_dict=auth, logging_obj=logging_obj,
|
||||
passthrough_logging_payload={"url": "https://laya.test/v1/systemone"}, _parsed_body=request_body,
|
||||
)
|
||||
logging_kwargs["litellm_params"]["metadata"]["standard_logging_guardrail_information"] = [
|
||||
{"guardrail_name": "trusted-hook", "guardrail_cost": guardrail_cost},
|
||||
]
|
||||
logging_obj.update_environment_variables(
|
||||
model="english", user="unknown", optional_params={},
|
||||
litellm_params=logging_kwargs["litellm_params"], call_type="pass_through_endpoint",
|
||||
)
|
||||
body: Final = {
|
||||
"model": "laya-rl-agent", "usage": {"input_tokens": 10, "output_tokens": 3},
|
||||
**({"routing": {"model": routing_model}} if routing_model else {}),
|
||||
}
|
||||
normalized: Final = PassThroughEndpointLogging().normalize_llm_passthrough_logging_payload(
|
||||
httpx_response=httpx.Response(200, request=httpx.Request("POST", "https://laya.test/v1/systemone"), json=body),
|
||||
response_body=body, request_body={"model": "english"}, logging_obj=logging_obj,
|
||||
url_route="https://laya.test/v1/systemone", result="{}", start_time=start,
|
||||
end_time=datetime.now(), cache_hit=False, custom_llm_provider="laya", **logging_kwargs,
|
||||
)
|
||||
logged: Final = normalized["kwargs"]
|
||||
expected_cost: Final = 10 * input_rate + 3 * output_rate
|
||||
assert (logged["model"], logged["custom_llm_provider"]) == (model, "laya")
|
||||
assert logged["response_cost"] == pytest.approx(expected_cost)
|
||||
assert logged["combined_usage_object"].model_dump(exclude_none=True) == {
|
||||
"prompt_tokens": 10, "completion_tokens": 3, "total_tokens": 13,
|
||||
}
|
||||
assert logging_obj.model_call_details["model"] == model
|
||||
assert logging_obj.model_call_details["response_cost"] == pytest.approx(expected_cost)
|
||||
assert logged["standard_logging_object"]["model"] == model
|
||||
assert logged["standard_logging_object"]["model_group"] == "laya/english"
|
||||
assert logged["standard_logging_object"]["response_cost"] == pytest.approx(expected_cost + guardrail_cost)
|
||||
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.exceptions import BudgetExceededError
|
||||
from litellm.proxy.hooks.model_max_budget_limiter import _PROXY_VirtualKeyModelMaxBudgetLimiter
|
||||
|
||||
budget_limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(DualCache())
|
||||
assert await budget_limiter.is_key_within_model_budget(auth, "laya/english")
|
||||
await budget_limiter.async_log_success_event(logged, None, start, datetime.now())
|
||||
with pytest.raises(BudgetExceededError):
|
||||
await budget_limiter.is_key_within_model_budget(auth, "laya/english")
|
||||
|
||||
|
||||
def test_openrouter_decisions_response_is_priced_from_request_model_registry_row():
|
||||
logging_obj = _logging_obj()
|
||||
model_cost = litellm.model_cost["openrouter/typesafe/jev-1.13"]
|
||||
|
|
|
|||
|
|
@ -23,6 +23,8 @@ from starlette.datastructures import FormData
|
|||
|
||||
|
||||
import litellm
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.types.utils import CallTypesLiteral
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from tests.unit.llms.bedrock.event_loop_probe import EventLoopProbe
|
||||
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
|
||||
|
|
@ -7407,6 +7409,152 @@ class TestTypeSafePassthroughRoute:
|
|||
)
|
||||
|
||||
|
||||
class TestLayaPassthroughRoute:
|
||||
@pytest.fixture
|
||||
def client(self, monkeypatch: pytest.MonkeyPatch) -> Iterator[TestClient]:
|
||||
from litellm.proxy.proxy_server import app
|
||||
|
||||
monkeypatch.setenv("LAYA_API_BASE", "http://laya.test/base")
|
||||
monkeypatch.setenv("TYPESAFE_API_KEY", "never-send-typesafe-key")
|
||||
monkeypatch.delenv("LAYA_API_KEY", raising=False)
|
||||
monkeypatch.delenv("SERVER_ROOT_PATH", raising=False)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
monkeypatch.setitem(app.dependency_overrides, user_api_key_auth, lambda: UserAPIKeyAuth(api_key="sk-virtual"))
|
||||
yield TestClient(app)
|
||||
|
||||
@pytest.mark.parametrize("api_key", [None, "laya-provider-key"])
|
||||
def test_laya_forwards_native_decisions_without_gateway_or_typesafe_credentials(
|
||||
self, client: TestClient, monkeypatch: pytest.MonkeyPatch, api_key: str | None
|
||||
) -> None:
|
||||
if api_key is not None:
|
||||
monkeypatch.setenv("LAYA_API_KEY", api_key)
|
||||
body: Final = {
|
||||
"model": "english",
|
||||
"state": "refund",
|
||||
"questions": {"department": {"type": "choice", "criteria": {"billing": "refunds"}}},
|
||||
}
|
||||
answer: Final = {"model": "laya-rl-agent", "routing": {"model": "english"}, "answers": {}}
|
||||
with respx.mock(assert_all_called=True) as upstream:
|
||||
route: Final = upstream.post("http://laya.test/base/v1/systemone?trace=yes").respond(200, json=answer)
|
||||
response: Final = client.post(
|
||||
"/laya/v1/systemone?trace=yes",
|
||||
json=body,
|
||||
headers={"Authorization": "Bearer sk-virtual", "x-pass-authorization": "Bearer attacker"},
|
||||
)
|
||||
|
||||
assert (response.status_code, response.json()) == (200, answer)
|
||||
sent: Final = route.calls.last.request
|
||||
assert sent.headers.get("authorization") == (f"Bearer {api_key}" if api_key else None)
|
||||
assert json.loads(sent.content) == body
|
||||
|
||||
def test_laya_missing_server_fails_without_contacting_another_provider(
|
||||
self, client: TestClient, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.delenv("LAYA_API_BASE")
|
||||
with respx.mock(assert_all_called=False) as upstream:
|
||||
response: Final = client.post("/laya/v1/systemone", json={"model": "english"})
|
||||
assert response.status_code == 503
|
||||
assert "LAYA_API_BASE" in response.text
|
||||
assert len(upstream.calls) == 0
|
||||
|
||||
def test_laya_does_not_forward_unsupported_endpoints(self, client: TestClient) -> None:
|
||||
with respx.mock(assert_all_called=False) as upstream:
|
||||
response: Final = client.post("/laya/v1/evaluate", json={"model": "english"})
|
||||
assert response.status_code == 404
|
||||
assert len(upstream.calls) == 0
|
||||
|
||||
@pytest.mark.parametrize("model", [None, "auto", "jev-latest"])
|
||||
def test_laya_rejects_implicit_checkpoint_selection(self, client: TestClient, model: str | None) -> None:
|
||||
with respx.mock(assert_all_called=False) as upstream:
|
||||
response: Final = client.post("/laya/v1/systemone", json={"model": model})
|
||||
assert response.status_code == 400
|
||||
assert len(upstream.calls) == 0
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"controls",
|
||||
[{"custom_body": {"model": "multilingual", "state": "refund"}}, {"stream": True}, {"stream": "true"}],
|
||||
)
|
||||
def test_laya_rejects_controls_that_change_authorized_body_or_usage_accounting(
|
||||
self, client: TestClient, controls: Mapping[str, object]
|
||||
) -> None:
|
||||
with respx.mock(assert_all_called=False) as upstream:
|
||||
route: Final = upstream.post("http://laya.test/base/v1/systemone").respond(200, json={"answers": {}})
|
||||
response: Final = client.post("/laya/v1/systemone", json={"model": "english", **controls})
|
||||
assert response.status_code == 400
|
||||
assert not route.called
|
||||
|
||||
|
||||
@pytest.mark.parametrize("metadata_slot", ["metadata", "litellm_metadata"])
|
||||
def test_laya_hooks_enforce_canonical_model_limits_and_keep_native_wire_body(
|
||||
self, client: TestClient, monkeypatch: pytest.MonkeyPatch, metadata_slot: str
|
||||
) -> None:
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3
|
||||
from litellm.proxy.utils import InternalUsageCache
|
||||
from litellm.proxy.proxy_server import app
|
||||
|
||||
cache: Final = DualCache()
|
||||
limiter: Final = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(cache))
|
||||
auth: Final = UserAPIKeyAuth(
|
||||
api_key="laya-native-rpm", metadata={"model_rpm_limit": {"laya/english": 1}},
|
||||
)
|
||||
def authenticated_key() -> UserAPIKeyAuth:
|
||||
return auth
|
||||
|
||||
monkeypatch.setitem(app.dependency_overrides, user_api_key_auth, authenticated_key)
|
||||
|
||||
class LimitHook(CustomLogger):
|
||||
async def async_pre_call_hook(
|
||||
self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache,
|
||||
data: dict[str, object], call_type: CallTypesLiteral,
|
||||
) -> dict[str, object]:
|
||||
assert data["model"] == "laya/english"
|
||||
metadata: Final = data.get(metadata_slot)
|
||||
assert isinstance(metadata, dict)
|
||||
assert "standard_logging_guardrail_information" not in metadata
|
||||
assert metadata["customer_label"] == "retained"
|
||||
await limiter.async_pre_call_hook(user_api_key_dict, cache, data, call_type)
|
||||
return data
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [LimitHook()])
|
||||
body: Final = {
|
||||
"model": "english", "state": "refund",
|
||||
metadata_slot: {
|
||||
"customer_label": "retained", "model_group": "unbounded-client-choice",
|
||||
"standard_logging_guardrail_information": [{"guardrail_cost": 25.0}],
|
||||
},
|
||||
}
|
||||
with respx.mock(assert_all_called=True) as upstream:
|
||||
route: Final = upstream.post("http://laya.test/base/v1/systemone").respond(200, json={"answers": {}})
|
||||
first: Final = client.post("/laya/v1/systemone", json=body)
|
||||
second: Final = client.post("/laya/v1/systemone", json=body)
|
||||
assert first.status_code == 200, first.text
|
||||
assert second.status_code == 429, second.text
|
||||
assert route.call_count == 1
|
||||
assert json.loads(route.calls.last.request.content) == {"model": "english", "state": "refund"}
|
||||
|
||||
def test_laya_preserves_trusted_hook_checkpoint_changes(
|
||||
self, client: TestClient, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
class CheckpointHook(CustomLogger):
|
||||
async def async_pre_call_hook(
|
||||
self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache,
|
||||
data: dict[str, object], call_type: CallTypesLiteral,
|
||||
) -> dict[str, object]:
|
||||
assert data["model"] == "laya/english"
|
||||
return {**data, "model": "laya/multilingual"}
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [CheckpointHook()])
|
||||
with respx.mock(assert_all_called=True) as upstream:
|
||||
route: Final = upstream.post("http://laya.test/base/v1/systemone").respond(200, json={"answers": {}})
|
||||
response: Final = client.post("/laya/v1/systemone", json={"model": "english", "state": "refund"})
|
||||
assert response.status_code == 200, response.text
|
||||
assert json.loads(route.calls.last.request.content) == {"model": "multilingual", "state": "refund"}
|
||||
|
||||
|
||||
class TestFalAIPassthroughRoute:
|
||||
@pytest.fixture
|
||||
def client(self, monkeypatch: pytest.MonkeyPatch) -> Iterator[TestClient]:
|
||||
|
|
|
|||
|
|
@ -1470,7 +1470,7 @@ async def test_pass_through_request_contains_proxy_server_request_in_kwargs():
|
|||
# Create mock request
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = "http://test-proxy.com/api/endpoint"
|
||||
mock_request.url = httpx.URL("http://test-proxy.com/api/endpoint")
|
||||
mock_request.body = AsyncMock(return_value=b'{"message": "test request"}')
|
||||
mock_request.headers = Headers({})
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
|
@ -1575,7 +1575,7 @@ async def test_pass_through_request_streaming_marks_logging_obj_as_stream():
|
|||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = "http://test-proxy.com/v1/messages"
|
||||
mock_request.url = httpx.URL("http://test-proxy.com/v1/messages")
|
||||
mock_request.body = AsyncMock(return_value=b'{"model": "claude-3", "stream": true}')
|
||||
mock_request.headers = Headers({})
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
|
@ -1637,7 +1637,7 @@ async def test_pass_through_request_sse_response_marks_logging_obj_as_stream():
|
|||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = "http://test-proxy.com/v1/messages"
|
||||
mock_request.url = httpx.URL("http://test-proxy.com/v1/messages")
|
||||
mock_request.body = AsyncMock(return_value=b'{"model": "claude-3"}')
|
||||
mock_request.headers = Headers({})
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
|
@ -2507,7 +2507,7 @@ async def test_pass_through_request_query_params_forwarding():
|
|||
# Create mock request with query parameters (Azure API version)
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = "http://localhost:4000/azure-assistant/openai/assistants"
|
||||
mock_request.url = httpx.URL("http://localhost:4000/azure-assistant/openai/assistants")
|
||||
mock_request.body = AsyncMock(return_value=json.dumps(test_body).encode())
|
||||
mock_request.headers = Headers({"Content-Type": "application/json"})
|
||||
|
||||
|
|
@ -3016,7 +3016,7 @@ async def test_bedrock_router_passthrough_metadata_initialization():
|
|||
# Create mock request with headers
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = "http://localhost:4000/bedrock/model/my-model/invoke"
|
||||
mock_request.url = httpx.URL("http://localhost:4000/bedrock/model/my-model/invoke")
|
||||
mock_request.headers = Headers(
|
||||
{
|
||||
"content-type": "application/json",
|
||||
|
|
@ -3850,7 +3850,7 @@ def _lit3538_request():
|
|||
r = MagicMock()
|
||||
r.method = "POST"
|
||||
r.query_params = {}
|
||||
r.url = "http://testserver/mock/echo"
|
||||
r.url = httpx.URL("http://testserver/mock/echo")
|
||||
r.state = SimpleNamespace()
|
||||
headers = MagicMock()
|
||||
headers.copy.return_value = {}
|
||||
|
|
@ -3983,7 +3983,7 @@ async def test_pass_through_request_non_streaming_upstream_error_returned_unchan
|
|||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = "http://test-proxy.com/mock-upstream/api/denied"
|
||||
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/denied")
|
||||
mock_request.body = AsyncMock(return_value=b'{"action": "read"}')
|
||||
mock_request.headers = Headers({"content-type": "application/json"})
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
|
@ -4069,7 +4069,7 @@ async def test_pass_through_request_upstream_error_failure_hook_exception_is_swa
|
|||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = "http://test-proxy.com/mock-upstream/api/denied"
|
||||
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/denied")
|
||||
mock_request.body = AsyncMock(return_value=b'{"action": "read"}')
|
||||
mock_request.headers = Headers({"content-type": "application/json"})
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
|
@ -4118,7 +4118,7 @@ async def test_pass_through_request_streaming_upstream_error_returned_unchanged(
|
|||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "GET"
|
||||
mock_request.url = "http://test-proxy.com/mock-upstream/api/stream-denied"
|
||||
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/stream-denied")
|
||||
mock_request.body = AsyncMock(return_value=b"")
|
||||
mock_request.headers = Headers({})
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
|
@ -4169,7 +4169,7 @@ class _UpstreamErrorBodyStream(httpx.AsyncByteStream):
|
|||
def _upstream_error_request() -> MagicMock:
|
||||
mock_request: Final = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = "http://test-proxy.com/mock-upstream/v1beta/models/claude-nope-9:generateContent"
|
||||
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/v1beta/models/claude-nope-9:generateContent")
|
||||
mock_request.body = AsyncMock(return_value=b'{"contents": []}')
|
||||
mock_request.headers = Headers({"content-type": "application/json"})
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
|
@ -4966,7 +4966,7 @@ async def test_pass_through_request_non_streaming_success_unchanged():
|
|||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "GET"
|
||||
mock_request.url = "http://test-proxy.com/mock-upstream/api/success"
|
||||
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/success")
|
||||
mock_request.body = AsyncMock(return_value=b"")
|
||||
mock_request.headers = Headers({})
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
|
@ -5029,7 +5029,7 @@ async def test_pass_through_request_claims_the_budget_reservation_only_when_its_
|
|||
mock_get_client.return_value = MagicMock(client=async_client)
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = "http://test-proxy.com/mock-upstream/api/generate"
|
||||
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/generate")
|
||||
mock_request.body = AsyncMock(return_value=b'{"prompt": "hi"}')
|
||||
mock_request.headers = Headers({"content-type": "application/json"})
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
|
@ -5081,7 +5081,7 @@ async def test_pass_through_request_leaves_the_budget_reservation_for_the_reques
|
|||
mock_get_client.return_value = MagicMock(client=async_client)
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = "http://test-proxy.com/mock-upstream/api/generate"
|
||||
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/generate")
|
||||
mock_request.body = AsyncMock(return_value=b'{"prompt": "hi"}')
|
||||
mock_request.headers = Headers({"content-type": "application/json"})
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
|
@ -5112,7 +5112,7 @@ async def test_pass_through_request_internal_failure_still_raises_proxy_exceptio
|
|||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "GET"
|
||||
mock_request.url = "http://test-proxy.com/mock-upstream/api/success"
|
||||
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/success")
|
||||
mock_request.body = AsyncMock(return_value=b"")
|
||||
mock_request.headers = Headers({})
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
|
@ -5213,7 +5213,7 @@ def _enter_relay_logging_mocks(stack, parsed_body):
|
|||
def _relay_client_request(method="GET"):
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = method
|
||||
mock_request.url = "http://localhost:4000/passthrough-relay/results"
|
||||
mock_request.url = httpx.URL("http://localhost:4000/passthrough-relay/results")
|
||||
mock_request.body = AsyncMock(return_value=b"")
|
||||
mock_request.headers = Headers({})
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
|
@ -6650,7 +6650,7 @@ def _passthrough_kwargs_for_reservation(
|
|||
) -> dict:
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = "http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent"
|
||||
mock_request.url = httpx.URL("http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent")
|
||||
mock_request.headers = Headers({})
|
||||
mock_request.scope = {"endpoint": _marked_pass_through_endpoint()} if user_defined_route else {}
|
||||
|
||||
|
|
@ -6797,7 +6797,7 @@ async def _drive_streaming_pass_through(upstream_content_type, chunk_delay_secon
|
|||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = "http://test-proxy.com/v1/messages"
|
||||
mock_request.url = httpx.URL("http://test-proxy.com/v1/messages")
|
||||
mock_request.body = AsyncMock(
|
||||
return_value=b'{"model": "claude-3", "stream": true}'
|
||||
if client_asked_for_stream
|
||||
|
|
@ -6985,36 +6985,82 @@ def _marked_pass_through_endpoint():
|
|||
return _endpoint
|
||||
|
||||
|
||||
def test_user_defined_passthrough_is_neither_tracked_nor_enforced():
|
||||
"""
|
||||
`get_model_from_request` returns None for a user-defined pass-through on
|
||||
purpose: the body is forwarded verbatim, so its `model` names an UPSTREAM
|
||||
model rather than a LiteLLM-managed one, and enforcing key/team allowlists
|
||||
against it would reject valid requests. Enforcement is therefore skipped
|
||||
on those routes.
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("metadata_slot", ["metadata", "litellm_metadata"])
|
||||
async def test_user_defined_passthrough_is_neither_tracked_nor_enforced(metadata_slot: str) -> None:
|
||||
from datetime import datetime
|
||||
|
||||
Attaching the budget metadata anyway would charge a counter that nothing on
|
||||
that route can refuse, and would attribute the spend to a budget the operator
|
||||
scoped to a LiteLLM model that merely shares the name. Tracking and
|
||||
enforcement have to agree: both on for the built-in provider routes, both off
|
||||
here.
|
||||
"""
|
||||
kwargs = _passthrough_kwargs_for_reservation(
|
||||
UserAPIKeyAuth(
|
||||
token="hash",
|
||||
user_id="u-1",
|
||||
model_max_budget={"claude-opus-4-8": {"budget_limit": 1.0, "time_period": "18h"}},
|
||||
),
|
||||
user_defined_route=True,
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy.auth.auth_utils import get_model_from_request
|
||||
from litellm.proxy.hooks.model_max_budget_limiter import _PROXY_VirtualKeyModelMaxBudgetLimiter
|
||||
|
||||
budget: Final = {"managed-model": {"budget_limit": 0.1, "time_period": "1d"}}
|
||||
limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(DualCache())
|
||||
auth: Final = UserAPIKeyAuth(
|
||||
api_key="custom-key", token="custom-key", team_id="shared-team", team_model_max_budget=budget,
|
||||
)
|
||||
endpoint: Final = create_pass_through_route(
|
||||
endpoint="/custom-budget-test", target="https://upstream.test/echo", custom_headers={}, cost_per_request=0.25,
|
||||
)
|
||||
request: Final = Request({
|
||||
"type": "http", "method": "POST", "path": "/custom-budget-test", "headers": [],
|
||||
"query_string": b"", "endpoint": endpoint,
|
||||
})
|
||||
body: Final = {
|
||||
"model": "upstream-only-model", metadata_slot: {
|
||||
"model_group": "managed-model", "customer_label": "retained",
|
||||
"user_api_key_team_model_max_budget": budget,
|
||||
},
|
||||
}
|
||||
assert get_model_from_request(body, "/custom-budget-test", request=request) is None
|
||||
assert await limiter.is_team_within_model_budget("shared-team", budget, None, "managed-model")
|
||||
start: Final = datetime.now()
|
||||
logging_obj: Final = LiteLLMLoggingObj(
|
||||
model="upstream-only-model", messages=[], stream=False, call_type="pass_through_endpoint",
|
||||
start_time=start, litellm_call_id="custom-budget", function_id="custom-budget", kwargs={},
|
||||
dynamic_async_success_callbacks=[limiter],
|
||||
)
|
||||
payload: Final = {
|
||||
"url": "https://upstream.test/echo", "request_body": body, "request_method": "POST", "cost_per_request": 0.25,
|
||||
}
|
||||
kwargs: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
|
||||
request=request, user_api_key_dict=auth, passthrough_logging_payload=payload, logging_obj=logging_obj,
|
||||
_parsed_body=body, litellm_call_id="custom-budget",
|
||||
)
|
||||
logging_obj.update_environment_variables(
|
||||
model="upstream-only-model", user="unknown", optional_params={},
|
||||
litellm_params=kwargs["litellm_params"], call_type="pass_through_endpoint",
|
||||
)
|
||||
response: Final = httpx.Response(
|
||||
200, request=httpx.Request("POST", "https://upstream.test/echo"), json={"ok": True},
|
||||
)
|
||||
await PassThroughEndpointLogging().pass_through_async_success_handler(
|
||||
httpx_response=response, response_body={"ok": True}, request_body=body, logging_obj=logging_obj,
|
||||
url_route="https://upstream.test/echo", result=response.text, start_time=start, end_time=datetime.now(),
|
||||
cache_hit=False, **kwargs,
|
||||
)
|
||||
assert logging_obj.model_call_details["response_cost"] == 0.25
|
||||
assert await limiter.is_team_within_model_budget("shared-team", budget, None, "managed-model")
|
||||
metadata: Final = kwargs["litellm_params"]["metadata"]
|
||||
assert (metadata["model_group"], metadata["customer_label"]) == ("managed-model", "retained")
|
||||
assert metadata.keys().isdisjoint({
|
||||
"user_api_key_model_max_budget", "user_api_key_team_model_max_budget",
|
||||
"user_api_key_user_model_max_budget", "user_api_key_end_user_model_max_budget",
|
||||
})
|
||||
|
||||
metadata = kwargs["litellm_params"]["metadata"]
|
||||
for field in (
|
||||
"user_api_key_model_max_budget",
|
||||
"user_api_key_user_model_max_budget",
|
||||
"user_api_key_end_user_model_max_budget",
|
||||
):
|
||||
assert field not in metadata, f"{field} was attached on a route that never enforces it"
|
||||
|
||||
@pytest.mark.parametrize("metadata_slot", ["metadata", "litellm_metadata"])
|
||||
def test_builtin_passthrough_pins_model_group_to_the_resolved_model(metadata_slot: str) -> None:
|
||||
request: Final = Request({
|
||||
"type": "http", "method": "POST", "path": "/gemini/v1beta/models/gemini-2.5-flash:generateContent",
|
||||
"headers": [], "query_string": b"",
|
||||
})
|
||||
kwargs: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
|
||||
request=request, user_api_key_dict=UserAPIKeyAuth(token="hash", user_id="u-1"),
|
||||
passthrough_logging_payload=MagicMock(), logging_obj=MagicMock(),
|
||||
_parsed_body={"contents": [], metadata_slot: {"model_group": "unbounded-client-choice"}},
|
||||
)
|
||||
assert kwargs["litellm_params"]["metadata"]["model_group"] == "gemini-2.5-flash"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
@ -7344,7 +7390,7 @@ def test_passthrough_client_cannot_forge_session_id_omission(client_metadata_key
|
|||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = "http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent"
|
||||
mock_request.url = httpx.URL("http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent")
|
||||
mock_request.headers = Headers({})
|
||||
mock_request.scope = {}
|
||||
|
||||
|
|
@ -7377,7 +7423,7 @@ def test_passthrough_logs_the_resolved_deployment_model_info_over_the_request_bo
|
|||
the call to (LIT-1761: passthrough successes carried model_id="")."""
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = "http://0.0.0.0:4000/vertex_ai/v1/projects/p/locations/global/publishers/google/models/gemini-3.8-flash:generateContent"
|
||||
mock_request.url = httpx.URL("http://0.0.0.0:4000/vertex_ai/v1/projects/p/locations/global/publishers/google/models/gemini-3.8-flash:generateContent")
|
||||
mock_request.headers = Headers({})
|
||||
mock_request.scope = {}
|
||||
mock_request.state = SimpleNamespace(
|
||||
|
|
@ -7409,7 +7455,7 @@ _PROXY_SERVER_REQUEST: Final = TypeAdapter(dict[str, object])
|
|||
def _split_pass_through_body(body: str) -> _PassThroughSplit:
|
||||
mock_request: Final = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = "http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent"
|
||||
mock_request.url = httpx.URL("http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent")
|
||||
mock_request.headers = Headers()
|
||||
mock_request.scope = MappingProxyType({})
|
||||
|
||||
|
|
@ -7546,6 +7592,61 @@ def test_passthrough_sees_the_public_list_rebound_after_import(monkeypatch: pyte
|
|||
assert split.forwarded_body == {"contents": [{"parts": [{"text": "hi"}]}]}
|
||||
|
||||
|
||||
def test_passthrough_metadata_carries_key_team_project_tags_and_key_spend_logs_metadata():
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = httpx.URL("http://0.0.0.0:4000/anthropic/v1/messages")
|
||||
mock_request.headers = Headers({"x-litellm-tags": "caller-tag,key-tag"})
|
||||
mock_request.scope = {}
|
||||
|
||||
cached_key = UserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
metadata={"tags": ["key-tag", "shared-tag"], "spend_logs_metadata": {"cost_center": "key"}},
|
||||
team_metadata={
|
||||
"tags": ["team-tag", "shared-tag"],
|
||||
"spend_logs_metadata": {"cost_center": "team", "team_field": "team"},
|
||||
},
|
||||
project_metadata={"tags": ["project-tag"]},
|
||||
)
|
||||
|
||||
kwargs = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
|
||||
request=mock_request,
|
||||
user_api_key_dict=cached_key,
|
||||
passthrough_logging_payload=MagicMock(),
|
||||
logging_obj=MagicMock(),
|
||||
_parsed_body={
|
||||
"metadata": {
|
||||
"tags": ["body-tag"],
|
||||
"spend_logs_metadata": {"request_id": "body"},
|
||||
"user_api_key_auth_metadata": "forged",
|
||||
}
|
||||
},
|
||||
litellm_call_id="lit-5359-call-id",
|
||||
)
|
||||
second = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
|
||||
request=mock_request,
|
||||
user_api_key_dict=cached_key,
|
||||
passthrough_logging_payload=MagicMock(),
|
||||
logging_obj=MagicMock(),
|
||||
_parsed_body={},
|
||||
litellm_call_id="lit-5359-second-call-id",
|
||||
)
|
||||
|
||||
metadata = kwargs["litellm_params"]["metadata"]
|
||||
assert metadata["tags"] == ["body-tag", "key-tag", "shared-tag", "team-tag", "project-tag", "caller-tag"]
|
||||
assert metadata["spend_logs_metadata"] == {"request_id": "body", "cost_center": "key", "team_field": "team"}
|
||||
assert metadata["user_api_key_auth_metadata"] == {
|
||||
"tags": ["key-tag", "shared-tag"],
|
||||
"spend_logs_metadata": {"cost_center": "key"},
|
||||
}
|
||||
assert second["litellm_params"]["metadata"]["spend_logs_metadata"] == {"cost_center": "key", "team_field": "team"}
|
||||
assert cached_key.metadata == {"tags": ["key-tag", "shared-tag"], "spend_logs_metadata": {"cost_center": "key"}}
|
||||
assert cached_key.team_metadata == {
|
||||
"tags": ["team-tag", "shared-tag"],
|
||||
"spend_logs_metadata": {"cost_center": "team", "team_field": "team"},
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_completion_pass_through_endpoint_answers_an_openai_typed_error_for_an_unknown_model(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
|
|
@ -7665,7 +7766,7 @@ def test_passthrough_attributes_a_cli_session_to_its_alias_not_the_login_token()
|
|||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = "http://0.0.0.0:4000/anthropic/v1/messages"
|
||||
mock_request.url = httpx.URL("http://0.0.0.0:4000/anthropic/v1/messages")
|
||||
mock_request.headers = Headers({})
|
||||
mock_request.scope = {}
|
||||
session = UserAPIKeyAuth(
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
pytest.importorskip("opentelemetry")
|
||||
|
||||
|
|
@ -81,15 +81,16 @@ def _user_api_key_dict():
|
|||
return d
|
||||
|
||||
|
||||
def _mock_request():
|
||||
r = MagicMock()
|
||||
r.method = "POST"
|
||||
r.query_params = {}
|
||||
r.url = "http://testserver/mock/echo"
|
||||
headers = MagicMock()
|
||||
headers.copy.return_value = {}
|
||||
r.headers = headers
|
||||
return r
|
||||
def _mock_request() -> Request:
|
||||
return Request({
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"scheme": "http",
|
||||
"server": ("testserver", 80),
|
||||
"path": "/mock/echo",
|
||||
"headers": [],
|
||||
"query_string": b"",
|
||||
})
|
||||
|
||||
|
||||
def _httpx_response(text: str) -> httpx.Response:
|
||||
|
|
|
|||
|
|
@ -1,8 +1,8 @@
|
|||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import Request
|
||||
from starlette.datastructures import Headers, State
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
|
||||
|
|
@ -771,12 +771,15 @@ async def test_vertex_passthrough_attributes_the_call_to_the_resolved_deployment
|
|||
"""The router deployment that rewrote the upstream URL is the one the logging kwargs must name, so
|
||||
the Prometheus model_id label (and SpendLogs.model_id) on a Vertex passthrough success reads the
|
||||
deployment's id instead of "" (LIT-1761)."""
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = "http://0.0.0.0:4000/vertex_ai/v1/projects/p/locations/global/publishers/google/models/gemini-3.8-flash:generateContent"
|
||||
mock_request.headers = Headers({})
|
||||
mock_request.scope = {}
|
||||
mock_request.state = State()
|
||||
mock_request: Final = Request({
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"scheme": "http",
|
||||
"server": ("0.0.0.0", 4000),
|
||||
"path": "/vertex_ai/v1/projects/p/locations/global/publishers/google/models/gemini-3.8-flash:generateContent",
|
||||
"headers": [],
|
||||
"query_string": b"",
|
||||
})
|
||||
mock_handler = MagicMock()
|
||||
mock_handler.get_default_base_target_url.return_value = "https://aiplatform.googleapis.com"
|
||||
|
||||
|
|
|
|||
|
|
@ -796,19 +796,23 @@ async def test_spend_logs_retention_alone_does_not_touch_the_session_rollup():
|
|||
assert any('"LiteLLM_SpendLogs"' in sql for sql in tables)
|
||||
assert not any('"LiteLLM_AutoRouterSession"' in sql for sql in tables)
|
||||
assert not any('"LiteLLM_AutoRouterUserSession"' in sql for sql in tables)
|
||||
assert not any('"LiteLLM_AutoRouterDailySpend"' in sql for sql in tables)
|
||||
assert not any('"LiteLLM_HealthCheckTable"' in sql for sql in tables)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_retention_alone_cleans_both_session_rollups():
|
||||
client = _mock_prisma_for_retention([0, 0])
|
||||
async def test_session_retention_alone_cleans_both_session_rollups_and_the_daily_rollup():
|
||||
client = _mock_prisma_for_retention([0, 0, 0])
|
||||
cleaner = SpendLogCleanup(general_settings={"maximum_autorouter_session_retention_period": "365d"})
|
||||
cleaner.pod_lock_manager = None
|
||||
await cleaner.cleanup_old_spend_logs(client)
|
||||
tables = [call[0][0] for call in client.db.execute_raw.call_args_list]
|
||||
assert len(tables) == 2
|
||||
calls = client.db.execute_raw.call_args_list
|
||||
tables = [call[0][0] for call in calls]
|
||||
assert len(tables) == 3
|
||||
assert '"LiteLLM_AutoRouterSession"' in tables[0]
|
||||
assert '"LiteLLM_AutoRouterUserSession"' in tables[1]
|
||||
assert '"LiteLLM_AutoRouterDailySpend"' in tables[2]
|
||||
assert calls[2][0][1] == calls[0][0][1].date().isoformat()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -852,7 +856,7 @@ async def test_spend_logs_retention_alone_keeps_daily_tag_spend_forever():
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_each_retention_key_cuts_off_at_its_own_horizon():
|
||||
client = _mock_prisma_for_retention([0, 0, 0, 0, 0])
|
||||
client = _mock_prisma_for_retention([0, 0, 0, 0, 0, 0])
|
||||
cleaner = SpendLogCleanup(
|
||||
general_settings={
|
||||
"maximum_spend_logs_retention_period": "7d",
|
||||
|
|
@ -868,6 +872,8 @@ async def test_each_retention_key_cuts_off_at_its_own_horizon():
|
|||
if '"LiteLLM_AutoRouterSession"' in call[0][0]
|
||||
else "LiteLLM_AutoRouterUserSession"
|
||||
if '"LiteLLM_AutoRouterUserSession"' in call[0][0]
|
||||
else "LiteLLM_AutoRouterDailySpend"
|
||||
if '"LiteLLM_AutoRouterDailySpend"' in call[0][0]
|
||||
else "LiteLLM_HealthCheckTable"
|
||||
if '"LiteLLM_HealthCheckTable"' in call[0][0]
|
||||
else "logs"
|
||||
|
|
@ -878,6 +884,7 @@ async def test_each_retention_key_cuts_off_at_its_own_horizon():
|
|||
assert (now - cutoffs["logs"]).days == 7
|
||||
assert (now - cutoffs["LiteLLM_AutoRouterSession"]).days == 365
|
||||
assert cutoffs["LiteLLM_AutoRouterUserSession"] == cutoffs["LiteLLM_AutoRouterSession"]
|
||||
assert cutoffs["LiteLLM_AutoRouterDailySpend"] == cutoffs["LiteLLM_AutoRouterSession"].date().isoformat()
|
||||
assert (now - cutoffs["LiteLLM_HealthCheckTable"]).days == 30
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ from unittest.mock import create_autospec
|
|||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_router_logger
|
||||
|
|
@ -30,14 +31,15 @@ from litellm.types.utils import AUTOROUTER_CLASSIFIER_CALL_ORIGIN
|
|||
|
||||
|
||||
class _UsageRecorder(CustomLogger):
|
||||
def __init__(self) -> None:
|
||||
def __init__(self, model_key: str = "typesafe/jev-accounting") -> None:
|
||||
super().__init__()
|
||||
self.model_key = model_key
|
||||
self.calls: tuple[Mapping[str, object], ...] = ()
|
||||
|
||||
async def async_log_success_event(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
|
||||
) -> None:
|
||||
if str(kwargs.get("model", "")).removeprefix("typesafe/") != "jev-accounting":
|
||||
if str(kwargs.get("model", "")) != self.model_key:
|
||||
return
|
||||
self.calls = (*self.calls, kwargs)
|
||||
|
||||
|
|
@ -167,8 +169,9 @@ async def test_jev_invalid_usage_never_reaches_spend_callbacks(
|
|||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("answer", ["SIMPLE", "UNAVAILABLE", "malformed"])
|
||||
@pytest.mark.parametrize("private", [False, True])
|
||||
@pytest.mark.parametrize("legacy", [False, True])
|
||||
async def test_jev_accounts_once_with_parent_identity_even_when_the_verdict_fails(
|
||||
monkeypatch: pytest.MonkeyPatch, answer: str, private: bool
|
||||
monkeypatch: pytest.MonkeyPatch, answer: str, private: bool, legacy: bool
|
||||
) -> None:
|
||||
recorder: Final = _UsageRecorder()
|
||||
monkeypatch.setattr(litellm, "_async_success_callback", [recorder])
|
||||
|
|
@ -196,7 +199,15 @@ async def test_jev_accounts_once_with_parent_identity_even_when_the_verdict_fail
|
|||
router: Final = ComplexityRouter(
|
||||
"jev-router",
|
||||
litellm.Router(model_list=[]),
|
||||
{"classifier_type": "jev", "jev_classifier_config": {}, "tiers": {"SIMPLE": "cheap"}},
|
||||
{
|
||||
"classifier_type": "jev" if legacy else "oss_classifier",
|
||||
"jev_classifier_config" if legacy else "opensource_classifier_config": {
|
||||
"provider": "typesafe" if legacy else "jev",
|
||||
},
|
||||
"tiers": {"SIMPLE": "cheap"},
|
||||
"session_affinity": False,
|
||||
"deployment_affinity": False,
|
||||
},
|
||||
jev_client=provider,
|
||||
derive_savings_baseline=False,
|
||||
)
|
||||
|
|
@ -209,8 +220,9 @@ async def test_jev_accounts_once_with_parent_identity_even_when_the_verdict_fail
|
|||
"user_api_key_budget_reservation": {"reservation_id": "parent-reservation"},
|
||||
"user_api_key_auth": {"budget_reservation": {"reservation_id": "parent-reservation"}},
|
||||
}
|
||||
outcome: Final = await router.aclassify(
|
||||
"private current ask",
|
||||
result: Final = await router.async_pre_routing_hook(
|
||||
model="jev-router",
|
||||
messages=[{"role": "user", "content": "private current ask"}],
|
||||
request_kwargs={
|
||||
"metadata": metadata,
|
||||
"litellm_session_id": "session-a",
|
||||
|
|
@ -221,7 +233,15 @@ async def test_jev_accounts_once_with_parent_identity_even_when_the_verdict_fail
|
|||
await GLOBAL_LOGGING_WORKER.flush()
|
||||
await handler.client.aclose()
|
||||
|
||||
assert (outcome.cause == "jev_classifier") is (answer == "SIMPLE")
|
||||
assert result is not None and result.model == "cheap"
|
||||
assert result.routing_decision is not None
|
||||
decision: Final = result.routing_decision
|
||||
assert (decision["cause"] == "jev_classifier") is (answer == "SIMPLE")
|
||||
if answer == "SIMPLE":
|
||||
assert decision["classifier_model"] == "typesafe/jev-accounting"
|
||||
assert decision["classifier_cost"] == pytest.approx(0.007)
|
||||
assert "jev-classifier:SIMPLE" in decision["signals"]
|
||||
assert "jev-confidence=1.000000" in decision["signals"]
|
||||
assert len(recorder.calls) == 1
|
||||
event: Final = recorder.calls[0]
|
||||
assert event["response_cost"] == pytest.approx(0.007)
|
||||
|
|
@ -416,10 +436,101 @@ def _answer(choice: str = "SIMPLE") -> JevChoiceAnswer:
|
|||
|
||||
|
||||
def test_jev_config_requires_classifier_config() -> None:
|
||||
with pytest.raises(ValueError, match="jev_classifier_config is required"):
|
||||
with pytest.raises(ValueError, match="opensource_classifier_config is required"):
|
||||
ComplexityRouterConfig.model_validate({"classifier_type": "jev"})
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("classifier_type", "config_key"),
|
||||
[
|
||||
("oss_classifier", "opensource_classifier_config"),
|
||||
("jev", "jev_classifier_config"),
|
||||
("oss_classifier", "jev_classifier_config"),
|
||||
("jev", "opensource_classifier_config"),
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
("provider", "model", "canonical_provider"),
|
||||
[(None, "jev-latest", "jev"), ("typesafe", "jev-latest", "jev"), ("jev", "jev-latest", "jev"), ("laya", "english", "laya")],
|
||||
)
|
||||
def test_classifier_aliases_load_and_serialize_one_canonical_config(
|
||||
classifier_type: str, config_key: str, provider: str | None, model: str, canonical_provider: str
|
||||
) -> None:
|
||||
incoming: Final = {
|
||||
"classifier_type": classifier_type,
|
||||
config_key: {"model": model, "api_key": None, **({"provider": provider} if provider is not None else {})},
|
||||
}
|
||||
original: Final = deepcopy(incoming)
|
||||
config: Final = ComplexityRouterConfig.model_validate(incoming)
|
||||
assert config.classifier_type == "oss_classifier"
|
||||
assert config.opensource_classifier_config is not None
|
||||
assert config.opensource_classifier_config.provider == canonical_provider
|
||||
assert config.opensource_classifier_config.model == model
|
||||
assert config.opensource_classifier_config.api_key is None
|
||||
assert "api_key" in config.opensource_classifier_config.model_fields_set
|
||||
assert "api_base" not in config.opensource_classifier_config.model_fields_set
|
||||
assert "jev_classifier_config" not in config.model_dump()
|
||||
assert config.jev_classifier_config is config.opensource_classifier_config
|
||||
assert incoming == original
|
||||
|
||||
|
||||
@pytest.mark.parametrize("config", [{"provider": "laya"}, {"provider": "laya", "model": " "}])
|
||||
def test_laya_requires_its_own_checkpoint(config: Mapping[str, object]) -> None:
|
||||
with pytest.raises(ValueError, match="Laya model must be"):
|
||||
JevClassifierConfig.model_validate(config)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("custom_base", [False, True])
|
||||
@pytest.mark.parametrize("legacy", [False, True])
|
||||
async def test_laya_routes_with_its_own_credentials_and_accounts_the_checkpoint(
|
||||
monkeypatch: pytest.MonkeyPatch, custom_base: bool, legacy: bool
|
||||
) -> None:
|
||||
monkeypatch.setenv("TYPESAFE_API_KEY", "never-send-typesafe-key")
|
||||
monkeypatch.setenv("LAYA_API_BASE", "https://laya.test")
|
||||
monkeypatch.setenv("LAYA_API_KEY", "laya-env-key")
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.setitem(litellm.model_cost, "laya/english", {"input_cost_per_token": 0.01})
|
||||
recorder: Final = _UsageRecorder("laya/english")
|
||||
monkeypatch.setattr(litellm, "_async_success_callback", [recorder])
|
||||
router: Final = ComplexityRouter(
|
||||
"laya-route",
|
||||
litellm.Router(model_list=[]),
|
||||
{
|
||||
"classifier_type": "jev" if legacy else "oss_classifier",
|
||||
"jev_classifier_config" if legacy else "opensource_classifier_config": {
|
||||
"provider": "laya",
|
||||
"model": "english",
|
||||
**({"api_base": "https://laya.test"} if custom_base else {}),
|
||||
},
|
||||
"tiers": {"SIMPLE": "cheap"},
|
||||
},
|
||||
derive_savings_baseline=False,
|
||||
)
|
||||
with respx.mock(assert_all_called=True) as upstream:
|
||||
route: Final = upstream.post("https://laya.test/v1/systemone").respond(
|
||||
200,
|
||||
json={
|
||||
"model": "laya-rl-agent",
|
||||
"routing": {"model": "english"},
|
||||
"answers": {"tier": _answer().model_dump()},
|
||||
"usage": {"input_tokens": 31, "output_tokens": 0},
|
||||
},
|
||||
)
|
||||
outcome: Final = await router.aclassify("choose a tier")
|
||||
await GLOBAL_LOGGING_WORKER.flush()
|
||||
|
||||
assert outcome.cause == "jev_classifier"
|
||||
assert outcome.jev_verdict is not None
|
||||
assert (outcome.jev_verdict.provider, outcome.jev_verdict.model) == ("laya", "english")
|
||||
assert outcome.classifier_cost == pytest.approx(0.31)
|
||||
sent: Final = route.calls.last.request
|
||||
assert sent.headers.get("authorization") == (None if custom_base else "Bearer laya-env-key")
|
||||
assert json.loads(sent.content)["model"] == "english"
|
||||
assert len(recorder.calls) == 1
|
||||
assert recorder.calls[0]["response_cost"] == pytest.approx(0.31)
|
||||
|
||||
|
||||
def test_jev_config_is_rejected_for_other_classifier_types() -> None:
|
||||
with pytest.raises(ValueError, match="has no effect"):
|
||||
ComplexityRouterConfig.model_validate(
|
||||
|
|
@ -437,7 +548,7 @@ def test_jev_instructions_reject_blank_values() -> None:
|
|||
@pytest.mark.parametrize(
|
||||
("missing_key", "rejection"),
|
||||
[
|
||||
({}, r"api_base requires jev_classifier_config\.api_key"),
|
||||
({}, r"api_base requires opensource_classifier_config\.api_key"),
|
||||
({"api_key": ""}, r"api_key must be non-empty"),
|
||||
({"api_key": " "}, r"api_key must be non-empty"),
|
||||
],
|
||||
|
|
|
|||
|
|
@ -6,11 +6,11 @@ import pytest
|
|||
from litellm.router_strategy.complexity_router.fuse_presets import get_fuse_presets
|
||||
from litellm.router_strategy.complexity_router.jev_classifier import DEFAULT_JEV_INSTRUCTIONS
|
||||
from litellm.router_utils.auto_router_model_naming import (
|
||||
carries_complexity_router_settings,
|
||||
classify_strategy_router_model,
|
||||
GATED_AUTO_ROUTER_CAPABILITIES,
|
||||
capability_limit_violation,
|
||||
carries_complexity_router_settings,
|
||||
claimed_capability,
|
||||
classify_strategy_router_model,
|
||||
count_capability_routers,
|
||||
gated_capability_of,
|
||||
strategy_router_dependencies,
|
||||
|
|
@ -23,27 +23,59 @@ COMPLEXITY_FIELDS = frozenset({"complexity_router_config"})
|
|||
SEMANTIC_FIELDS = frozenset({"auto_router_config", "auto_router_default_model", "auto_router_embedding_model"})
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", ["jev-latest", "jev-preview"])
|
||||
def test_jev_enumerates_a_paid_evaluation_without_a_completion_classifier(model: str) -> None:
|
||||
found = strategy_router_dependencies(
|
||||
@pytest.mark.parametrize(
|
||||
("classifier_type", "config_key"),
|
||||
[
|
||||
("jev", "jev_classifier_config"),
|
||||
("oss_classifier", "opensource_classifier_config"),
|
||||
("jev", "opensource_classifier_config"),
|
||||
("oss_classifier", "jev_classifier_config"),
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
("provider", "model", "accounting_provider"),
|
||||
[
|
||||
(None, "jev-latest", "typesafe"),
|
||||
("typesafe", "jev-preview", "typesafe"),
|
||||
("jev", "jev-preview", "typesafe"),
|
||||
("laya", "english", "laya"),
|
||||
],
|
||||
)
|
||||
def test_open_source_classifier_enumerates_its_accounting_model(
|
||||
classifier_type: str, config_key: str, provider: str | None, model: str, accounting_provider: str
|
||||
) -> None:
|
||||
found: Final = strategy_router_dependencies(
|
||||
{
|
||||
"model": "auto_router/complexity_router",
|
||||
"complexity_router_config": {
|
||||
"classifier_type": "jev",
|
||||
"jev_classifier_config": {"model": model},
|
||||
"classifier_type": classifier_type,
|
||||
config_key: {"model": model, **({"provider": provider} if provider else {})},
|
||||
"tiers": {"SIMPLE": "cheap"},
|
||||
},
|
||||
}
|
||||
)
|
||||
assert tuple((dep.model_name, dep.role) for dep in found) == (
|
||||
("cheap", "tier"),
|
||||
(f"typesafe/{model}", "evaluation"),
|
||||
(f"{accounting_provider}/{model}", "evaluation"),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("instructions", [None, DEFAULT_JEV_INSTRUCTIONS, "Route conservatively"])
|
||||
def test_only_non_default_jev_instructions_claim_the_shared_customization_slot(instructions: str | None) -> None:
|
||||
capability = claimed_capability({"classifier_type": "jev", "jev_classifier_config": {"instructions": instructions}})
|
||||
@pytest.mark.parametrize(
|
||||
("classifier_type", "config_key"),
|
||||
[
|
||||
("jev", "jev_classifier_config"),
|
||||
("oss_classifier", "opensource_classifier_config"),
|
||||
("jev", "opensource_classifier_config"),
|
||||
("oss_classifier", "jev_classifier_config"),
|
||||
],
|
||||
)
|
||||
def test_only_non_default_open_source_instructions_claim_the_shared_customization_slot(
|
||||
instructions: str | None, classifier_type: str, config_key: str
|
||||
) -> None:
|
||||
capability: Final = claimed_capability(
|
||||
{"classifier_type": classifier_type, config_key: {"instructions": instructions}}
|
||||
)
|
||||
assert (capability.key if capability else None) == (
|
||||
"tier_or_classifier_prompt" if instructions == "Route conservatively" else None
|
||||
)
|
||||
|
|
@ -123,6 +155,21 @@ VALID_TIERS = {
|
|||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("legacy_config", [None, {}, {"provider": "laya", "model": "english"}])
|
||||
def test_dual_classifier_blocks_return_a_write_validation_error(legacy_config: Mapping[str, object] | None) -> None:
|
||||
violation: Final = validate_complexity_router_config_write(
|
||||
{
|
||||
"tiers": VALID_TIERS,
|
||||
"classifier_type": "oss_classifier",
|
||||
"opensource_classifier_config": {"provider": "laya", "model": "english"},
|
||||
"jev_classifier_config": legacy_config,
|
||||
}
|
||||
)
|
||||
assert violation is not None
|
||||
assert "opensource_classifier_config" in violation
|
||||
assert "jev_classifier_config" in violation
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"keyword_tier_rules,expected_fragment",
|
||||
[
|
||||
|
|
@ -408,6 +455,8 @@ def test_complexity_embedding_model_is_a_dependency_only_when_semantic_matching_
|
|||
("token_thresholds", "dimension_weights"),
|
||||
("reasoning_override_min_score",),
|
||||
("tiers",),
|
||||
("jev_classifier_config",),
|
||||
("opensource_classifier_config",),
|
||||
],
|
||||
)
|
||||
def test_placement_rejects_settings_written_beside_the_config(misplaced):
|
||||
|
|
@ -447,7 +496,7 @@ def test_placement_guards_every_setting_the_config_owns():
|
|||
ComplexityRouterConfig,
|
||||
)
|
||||
|
||||
assert COMPLEXITY_ROUTER_CONFIG_KEYS == frozenset(ComplexityRouterConfig.model_fields)
|
||||
assert COMPLEXITY_ROUTER_CONFIG_KEYS == frozenset(ComplexityRouterConfig.model_fields) | {"jev_classifier_config"}
|
||||
assert {"tier_boundaries", "token_thresholds", "dimension_weights"} <= COMPLEXITY_ROUTER_CONFIG_KEYS
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1020,6 +1020,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
|
|||
"/v1/videos",
|
||||
"/vertex_ai/live",
|
||||
"/v1/listen",
|
||||
"/v1/systemone",
|
||||
"/v1beta/interactions",
|
||||
],
|
||||
},
|
||||
|
|
|
|||
|
|
@ -75,7 +75,6 @@ const totals = (overrides: Partial<Totals> = {}): Totals => ({
|
|||
saved_spend: 2174.59,
|
||||
baseline_spend: 2534.45,
|
||||
saved_pct: 85.8,
|
||||
saved_per_session: 23.13,
|
||||
cache: cache(),
|
||||
...overrides,
|
||||
});
|
||||
|
|
@ -110,7 +109,6 @@ const zeroTotals: Totals = {
|
|||
saved_spend: 0,
|
||||
baseline_spend: 0,
|
||||
saved_pct: 0,
|
||||
saved_per_session: 0,
|
||||
cache: zeroCache,
|
||||
};
|
||||
|
||||
|
|
@ -173,7 +171,6 @@ describe("AutoRouterBenchmarksTab", () => {
|
|||
saved_spend: saved,
|
||||
baseline_spend: estimatedTurns ? actual + (saved ?? 0) : null,
|
||||
saved_pct: pct,
|
||||
saved_per_session: null,
|
||||
};
|
||||
mockHook({
|
||||
data: response([], totals(comparison)),
|
||||
|
|
@ -204,18 +201,15 @@ describe("AutoRouterBenchmarksTab", () => {
|
|||
}
|
||||
});
|
||||
|
||||
it("leads with total estimated savings, before the four session-shape metrics", () => {
|
||||
it("leads with total estimated savings, before the three session-shape metrics", () => {
|
||||
mockHook({ data: response([group(), group({ router_name: "gpt-auto" })]) });
|
||||
renderTab();
|
||||
|
||||
const labels = screen
|
||||
.getAllByText(
|
||||
/Total estimated savings|Avg saved per session|Avg turns per session|Avg session length|Avg tokens per session/,
|
||||
)
|
||||
.getAllByText(/Total estimated savings|Avg turns per session|Avg session length|Avg tokens per session/)
|
||||
.map((node) => node.textContent);
|
||||
expect(labels).toEqual([
|
||||
"Total estimated savings",
|
||||
"Avg saved per session",
|
||||
"Avg turns per session",
|
||||
"Avg session length",
|
||||
"Avg tokens per session",
|
||||
|
|
@ -271,15 +265,40 @@ describe("AutoRouterBenchmarksTab", () => {
|
|||
},
|
||||
);
|
||||
|
||||
it("pairs the savings with the session count it was earned over, in its own tile", () => {
|
||||
it("labels selected-day money apart from whole-session metrics, with no savings-per-session tile", () => {
|
||||
mockHook({ data: response([group(), group({ router_name: "gpt-auto" })]) });
|
||||
renderTab();
|
||||
|
||||
const tile = screen.getByText("Avg saved per session").closest('[data-slot="card"]');
|
||||
if (!tile) throw new Error("expected avg saved per session to render as a metric tile");
|
||||
|
||||
expect(within(tile).getByText("$23.13")).toBeInTheDocument();
|
||||
const tile = screen.getByText("Avg turns per session").closest<HTMLElement>('[data-slot="card"]');
|
||||
if (!tile) throw new Error("expected avg turns per session to render as a metric tile");
|
||||
expect(within(tile).getByText("· 94 sessions")).toBeInTheDocument();
|
||||
expect(screen.queryByText("Avg saved per session")).not.toBeInTheDocument();
|
||||
expect(screen.getByText(/Savings and spend count requests on the selected UTC days/)).toBeInTheDocument();
|
||||
expect(screen.getByText(/Session metrics cover every session that overlaps the range/)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("shows session averages as unavailable, not zero, when routed requests have no session rows", () => {
|
||||
const noSessions = {
|
||||
sessions: 0,
|
||||
avg_turns_per_session: null,
|
||||
avg_session_seconds: null,
|
||||
avg_tokens_per_session: null,
|
||||
};
|
||||
mockHook({ data: response([], totals(noSessions)) });
|
||||
renderTab();
|
||||
|
||||
expect(screen.getAllByText("Unavailable")).toHaveLength(3);
|
||||
expect(screen.queryByText("0.0")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it.each([3, -3])("explains a %s gap between router records and recorded savings instead of comparing", (gap) => {
|
||||
const residual = { saved_spend: 5, unattributed_saved_spend: gap, baseline_spend: null, saved_pct: null };
|
||||
mockHook({ data: response([], totals(residual)) });
|
||||
renderTab();
|
||||
|
||||
expect(screen.getByText("$5.00")).toBeInTheDocument();
|
||||
expect(screen.getByText(/Per-router records differ from recorded savings by \$3\.00/)).toBeInTheDocument();
|
||||
expect(screen.getByText("Estimated baseline spend").nextSibling?.textContent).toBe("Unavailable");
|
||||
});
|
||||
|
||||
it("exposes each spend row as a term and its value, not as loose text", () => {
|
||||
|
|
@ -431,7 +450,7 @@ describe("AutoRouterBenchmarksTab", () => {
|
|||
renderTab();
|
||||
|
||||
expect(screen.getByText("Total estimated savings")).toBeInTheDocument();
|
||||
expect(screen.getAllByText("$0.00")).toHaveLength(6);
|
||||
expect(screen.getAllByText("$0.00")).toHaveLength(5);
|
||||
expect(screen.getByText("· 0 sessions")).toBeInTheDocument();
|
||||
expect(screen.getByText("0s")).toBeInTheDocument();
|
||||
expect(screen.getByText(/turns measured/)).toBeInTheDocument();
|
||||
|
|
|
|||
|
|
@ -105,6 +105,12 @@ const HeroCard: React.FC<{ view: BenchmarkView }> = ({ view }) => {
|
|||
adaptive and quality routers are excluded
|
||||
</p>
|
||||
)}
|
||||
{stats.unattributed_saved_spend != null && (
|
||||
<p className="text-center text-xs text-muted-foreground">
|
||||
Per-router records differ from recorded savings by {usd(Math.abs(stats.unattributed_saved_spend))}, for
|
||||
example history from before per-router tracking, so the baseline comparison is unavailable
|
||||
</p>
|
||||
)}
|
||||
</div>
|
||||
|
||||
<div className="flex flex-col justify-center border-t p-6 md:border-t-0 md:border-l">
|
||||
|
|
@ -297,22 +303,34 @@ const BenchmarksBody: React.FC<BenchmarksBodyProps> = ({ isPending, error, data,
|
|||
|
||||
<TierTurnsChart view={view} autoRouters={autoRouters} />
|
||||
|
||||
<div className="grid grid-cols-1 gap-4 sm:grid-cols-2 lg:grid-cols-4">
|
||||
<Metric
|
||||
label="Avg saved per session"
|
||||
value={stats.saved_per_session == null ? "Unavailable" : usd(stats.saved_per_session)}
|
||||
hint={`· ${stats.sessions.toLocaleString()} sessions`}
|
||||
/>
|
||||
<Metric label="Avg turns per session" value={stats.avg_turns_per_session.toFixed(1)} />
|
||||
<Metric label="Avg session length" value={durationLabel(stats.avg_session_seconds)} />
|
||||
<Metric label="Avg tokens per session" value={formatNumberWithCommas(stats.avg_tokens_per_session, 1, true)} />
|
||||
</div>
|
||||
<p className="text-xs text-muted-foreground">
|
||||
Savings and spend count requests on the selected UTC days. Actual spend covers every request on complexity
|
||||
routers, including LLM classification cost. Baseline is actual spend plus recorded savings, so savings can be
|
||||
zero or negative.
|
||||
</p>
|
||||
|
||||
<p className="text-xs text-muted-foreground">
|
||||
Actual spend covers every request on complexity routers, including LLM classification cost. Baseline is actual
|
||||
spend plus recorded savings, so savings can be zero or negative. The range counts whole sessions that overlap
|
||||
it, so totals can differ from savings views that group usage by UTC day.
|
||||
Session metrics cover every session that overlaps the range, including its turns outside the range.
|
||||
</p>
|
||||
<div className="grid grid-cols-1 gap-4 sm:grid-cols-3">
|
||||
<Metric
|
||||
label="Avg turns per session"
|
||||
value={stats.avg_turns_per_session == null ? "Unavailable" : stats.avg_turns_per_session.toFixed(1)}
|
||||
hint={`· ${stats.sessions.toLocaleString()} sessions`}
|
||||
/>
|
||||
<Metric
|
||||
label="Avg session length"
|
||||
value={stats.avg_session_seconds == null ? "Unavailable" : durationLabel(stats.avg_session_seconds)}
|
||||
/>
|
||||
<Metric
|
||||
label="Avg tokens per session"
|
||||
value={
|
||||
stats.avg_tokens_per_session == null
|
||||
? "Unavailable"
|
||||
: formatNumberWithCommas(stats.avg_tokens_per_session, 1, true)
|
||||
}
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div className="space-y-4">
|
||||
<div className="flex flex-wrap items-baseline gap-2">
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -43,7 +43,6 @@ const totals = (overrides: Partial<AutoRouterBenchmarkGroup> = {}) => ({
|
|||
saved_spend: 2174.59,
|
||||
baseline_spend: 2534.45,
|
||||
saved_pct: 85.8,
|
||||
saved_per_session: 23.13,
|
||||
cache: cache(),
|
||||
...overrides,
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
{
|
||||
|
|
|
|||
|
|
@ -57,7 +57,8 @@ const dedupe = (models: string[]): string[] => Array.from(new Set(models));
|
|||
|
||||
const COMPLEXITY_TYPE_LABELS: Record<string, string> = {
|
||||
llm: "LLM Classifier",
|
||||
jev: "JEV Classifier",
|
||||
jev: "OSS Classifier",
|
||||
oss_classifier: "OSS Classifier",
|
||||
capability: "Capability",
|
||||
llm_v2: "Fuse v2",
|
||||
heuristic_first: "Heuristic first",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -136,7 +136,7 @@ export const AutoRouterLimits = () => {
|
|||
<PopoverContent align="end" className="w-96 max-w-[calc(100vw-2rem)] gap-3">
|
||||
<PopoverTitle>Routing and customization limits</PopoverTitle>
|
||||
<p className="text-xs leading-5 text-muted-foreground">
|
||||
Rule-based, Complexity, and Jev are unlimited with built-in settings. Choose or change tier models freely.
|
||||
Rule-based, Complexity, and OSS are unlimited with built-in settings. Choose or change tier models freely.
|
||||
Customization allowances are shared across this proxy.
|
||||
</p>
|
||||
<dl className="space-y-2 text-xs">
|
||||
|
|
|
|||
|
|
@ -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}$`) }));
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ import {
|
|||
type ClassifierType,
|
||||
type ComplexityRouterConfigValue,
|
||||
} from "./ComplexityRouterConfig";
|
||||
import { defaultJevClassifierConfig, normalizeJevClassifierConfig } from "./jev_classifier_config";
|
||||
import { transitionClassifierType } from "./classifier_type_transition";
|
||||
import { isForecastClassifier } from "./forecast_classifier_config";
|
||||
import {
|
||||
|
|
@ -148,6 +149,14 @@ const AutoRouterClassifierTabs: React.FC<AutoRouterClassifierTabsProps> = ({ val
|
|||
if (next === "llm") changeType("llm");
|
||||
if (next === "jev") changeType("jev");
|
||||
};
|
||||
const changeProvider = (provider: unknown) => {
|
||||
if (provider !== "jev" && provider !== "laya") return;
|
||||
const defaults = defaultJevClassifierConfig(provider);
|
||||
onChange({
|
||||
...value,
|
||||
jev_classifier_config: { ...defaults, ...value.jev_classifier_config, provider, model: defaults.model },
|
||||
});
|
||||
};
|
||||
const approachLabels: Partial<Record<ClassifierType, string>> = { capability: "Capability", llm_v2: "Fuse v2" };
|
||||
const approachDescription: Partial<Record<ClassifierType, string>> = {
|
||||
capability: "Use the efficient model when it is likely to succeed",
|
||||
|
|
@ -164,7 +173,7 @@ const AutoRouterClassifierTabs: React.FC<AutoRouterClassifierTabsProps> = ({ val
|
|||
{[
|
||||
{ value: "heuristics", label: "Heuristics", description: "Classify locally, with no API call" },
|
||||
{ value: "llm", label: "LLM", description: "Use a judge model to choose a solver" },
|
||||
{ value: "jev", label: "Jev", description: "Use TypeSafe System One Choice to choose a tier" },
|
||||
{ value: "jev", label: "OSS Classifier", description: "Use Jev or Laya to choose a tier" },
|
||||
].map((option) => (
|
||||
<Label
|
||||
key={option.value}
|
||||
|
|
@ -189,6 +198,25 @@ const AutoRouterClassifierTabs: React.FC<AutoRouterClassifierTabsProps> = ({ val
|
|||
))}
|
||||
</RadioGroup>
|
||||
</fieldset>
|
||||
{family === "jev" && (
|
||||
<fieldset className="space-y-2">
|
||||
<legend className="text-sm font-medium">OSS provider</legend>
|
||||
<RadioGroup
|
||||
value={normalizeJevClassifierConfig(value.jev_classifier_config).provider}
|
||||
onValueChange={changeProvider}
|
||||
className="flex gap-6"
|
||||
>
|
||||
<Label>
|
||||
<RadioGroupItem value="jev" />
|
||||
Jev
|
||||
</Label>
|
||||
<Label>
|
||||
<RadioGroupItem value="laya" />
|
||||
Laya
|
||||
</Label>
|
||||
</RadioGroup>
|
||||
</fieldset>
|
||||
)}
|
||||
{family === "custom" && (
|
||||
<p className="text-sm text-muted-foreground">This router uses a custom classifier plugin</p>
|
||||
)}
|
||||
|
|
|
|||
|
|
@ -642,8 +642,8 @@ const ClassificationMethodConfig: React.FC<ClassificationMethodConfigProps> = ({
|
|||
/>
|
||||
<span className="text-xs text-muted-foreground">
|
||||
Number of prior user turns sent to the classifier provider, excluding tool output and harness reminders.
|
||||
LLM and Jev default to 3 turns; Jev sends them to the configured TypeSafe endpoint. Set to 0 to omit
|
||||
conversation history. The current message and selected system text are still sent.
|
||||
LLM and OSS classifiers default to 3 turns. Set to 0 to omit conversation history. The current message and
|
||||
selected system text are still sent.
|
||||
</span>
|
||||
</div>
|
||||
<div>
|
||||
|
|
|
|||
|
|
@ -54,8 +54,8 @@ const ClassifierTypeRadios: React.FC<ClassifierTypeRadiosProps> = ({ value, clas
|
|||
<Label className="items-start font-normal leading-normal">
|
||||
<RadioGroupItem value="jev" className="mt-0.5" />
|
||||
<span>
|
||||
<strong className="font-semibold">Jev Classifier</strong>{" "}
|
||||
<span className="text-muted-foreground">uses TypeSafe System One Choice to decide the tier</span>
|
||||
<strong className="font-semibold">OSS Classifier</strong>{" "}
|
||||
<span className="text-muted-foreground">uses Jev or Laya to decide the tier</span>
|
||||
</span>
|
||||
</Label>
|
||||
<SimpleTooltip content={scorerLockedReason}>
|
||||
|
|
|
|||
|
|
@ -237,7 +237,7 @@ const TierSetToolbar: React.FC<{
|
|||
{editing && (
|
||||
<span className="block mt-1 text-xs text-muted-foreground">
|
||||
Add or remove tiers to define your own set. Every custom tier needs a definition the classifier routes on, and
|
||||
an edited set requires the LLM or Jev classification method
|
||||
an edited set requires the LLM or OSS classification method
|
||||
</span>
|
||||
)}
|
||||
{editing && keywordRulesError && (
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import React, { useState } from "react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { afterEach, describe, expect, it, vi } from "vitest";
|
||||
import { fireEvent, renderWithProviders, screen } from "../../../tests/test-utils";
|
||||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
|
|
@ -96,49 +97,65 @@ function Form() {
|
|||
|
||||
describe("JEV classifier editor", () => {
|
||||
afterEach(() => vi.mocked(useAuthorized).mockReset());
|
||||
it("uses built-in JEV without a license and preserves custom tiers and context through reload", () => {
|
||||
renderWithProviders(<Form />);
|
||||
expect(screen.getByLabelText("Judge model")).toBeInTheDocument();
|
||||
expect(screen.getByText("Reasoning Effort")).toBeInTheDocument();
|
||||
expect(screen.getByText("Classifier Prompt")).toBeInTheDocument();
|
||||
expect(screen.getByRole("switch", { name: "Use images for classification" })).toBeInTheDocument();
|
||||
fireEvent.click(screen.getByRole("radio", { name: /Jev Classifier/ }));
|
||||
expect(screen.getByRole("radio", { name: /^Jev Classifier/ })).toBeChecked();
|
||||
expect(screen.getByLabelText("Jev Model")).toHaveValue("jev-latest");
|
||||
expect(screen.getByLabelText("Jev Instructions")).toBeEnabled();
|
||||
expect(screen.queryByLabelText("Judge model")).not.toBeInTheDocument();
|
||||
expect(screen.queryByText("Reasoning Effort")).not.toBeInTheDocument();
|
||||
expect(screen.queryByText("Classifier Prompt")).not.toBeInTheDocument();
|
||||
expect(screen.queryByRole("switch", { name: "Use images for classification" })).not.toBeInTheDocument();
|
||||
fireEvent.change(screen.getByLabelText("Jev Model"), { target: { value: "jev-test" } });
|
||||
fireEvent.change(screen.getByLabelText("Jev Timeout (ms)"), { target: { value: "4200" } });
|
||||
fireEvent.change(screen.getByLabelText("Context Window Size"), { target: { value: "6" } });
|
||||
fireEvent.change(screen.getByLabelText("Circuit breaker cooldown (seconds)"), { target: { value: "50" } });
|
||||
fireEvent.click(screen.getByRole("switch", { name: "Classifier circuit breaker" }));
|
||||
fireEvent.click(screen.getByRole("button", { name: "Customize tiers" }));
|
||||
fireEvent.click(screen.getByRole("button", { name: "Save and reload" }));
|
||||
expect(screen.getByRole("radio", { name: /Jev Classifier/ })).toBeChecked();
|
||||
expect(screen.getByLabelText("Jev Model")).toHaveValue("jev-test");
|
||||
expect(screen.getByLabelText("Jev Timeout (ms)")).toHaveValue(4200);
|
||||
expect(screen.getByLabelText("Context Window Size")).toHaveValue("6");
|
||||
expect(screen.getByRole("switch", { name: "Classifier circuit breaker" })).not.toBeChecked();
|
||||
fireEvent.click(screen.getByRole("button", { name: "Probe current config" }));
|
||||
expect(testAutoRouterRouting).toHaveBeenCalledWith(
|
||||
"token",
|
||||
expect.objectContaining({
|
||||
complexity_router_config: expect.objectContaining({
|
||||
classifier_type: "jev",
|
||||
jev_classifier_config: {
|
||||
model: "jev-test",
|
||||
timeout_ms: 4200,
|
||||
circuit_breaker_enabled: false,
|
||||
circuit_breaker_cooldown_seconds: 50,
|
||||
},
|
||||
tiers: expect.objectContaining({ QUICK: ["fast"] }),
|
||||
it.each(["jev", "laya"] as const)(
|
||||
"preserves %s, custom tiers and context through save, reload and probe",
|
||||
async (provider) => {
|
||||
renderWithProviders(<Form />);
|
||||
expect(screen.getByLabelText("Judge model")).toBeInTheDocument();
|
||||
expect(screen.getByText("Reasoning Effort")).toBeInTheDocument();
|
||||
expect(screen.getByText("Classifier Prompt")).toBeInTheDocument();
|
||||
expect(screen.getByRole("switch", { name: "Use images for classification" })).toBeInTheDocument();
|
||||
fireEvent.click(screen.getByRole("radio", { name: /^OSS Classifier$/ }));
|
||||
expect(screen.getByRole("radio", { name: /^OSS Classifier$/ })).toBeChecked();
|
||||
expect(screen.getByLabelText("Classifier Model")).toHaveValue("jev-latest");
|
||||
expect(screen.getByLabelText("Classifier Instructions")).toBeEnabled();
|
||||
expect(screen.queryByLabelText("Judge model")).not.toBeInTheDocument();
|
||||
expect(screen.queryByText("Reasoning Effort")).not.toBeInTheDocument();
|
||||
expect(screen.queryByText("Classifier Prompt")).not.toBeInTheDocument();
|
||||
expect(screen.queryByRole("switch", { name: "Use images for classification" })).not.toBeInTheDocument();
|
||||
fireEvent.click(screen.getByRole("radio", { name: "Laya" }));
|
||||
expect(screen.getByLabelText("Classifier Model")).toHaveTextContent("english");
|
||||
fireEvent.click(screen.getByRole("radio", { name: "Jev" }));
|
||||
expect(screen.getByLabelText("Classifier Model")).toHaveValue("jev-latest");
|
||||
if (provider === "laya") {
|
||||
fireEvent.click(screen.getByRole("radio", { name: "Laya" }));
|
||||
await userEvent.click(screen.getByLabelText("Classifier Model"));
|
||||
await userEvent.click(screen.getByRole("option", { name: "multilingual" }));
|
||||
} else {
|
||||
fireEvent.change(screen.getByLabelText("Classifier Model"), { target: { value: "jev-test" } });
|
||||
}
|
||||
fireEvent.change(screen.getByLabelText("Classifier Timeout (ms)"), { target: { value: "4200" } });
|
||||
fireEvent.change(screen.getByLabelText("Context Window Size"), { target: { value: "6" } });
|
||||
fireEvent.change(screen.getByLabelText("Circuit breaker cooldown (seconds)"), { target: { value: "50" } });
|
||||
fireEvent.click(screen.getByRole("switch", { name: "Classifier circuit breaker" }));
|
||||
fireEvent.click(screen.getByRole("button", { name: "Customize tiers" }));
|
||||
fireEvent.click(screen.getByRole("button", { name: "Save and reload" }));
|
||||
expect(screen.getByRole("radio", { name: /^OSS Classifier$/ })).toBeChecked();
|
||||
expect(screen.getByRole("radio", { name: provider === "laya" ? "Laya" : "Jev" })).toBeChecked();
|
||||
if (provider === "laya") expect(screen.getByLabelText("Classifier Model")).toHaveTextContent("multilingual");
|
||||
else expect(screen.getByLabelText("Classifier Model")).toHaveValue("jev-test");
|
||||
expect(screen.getByLabelText("Classifier Timeout (ms)")).toHaveValue(4200);
|
||||
expect(screen.getByLabelText("Context Window Size")).toHaveValue("6");
|
||||
expect(screen.getByRole("switch", { name: "Classifier circuit breaker" })).not.toBeChecked();
|
||||
fireEvent.click(screen.getByRole("button", { name: "Probe current config" }));
|
||||
expect(testAutoRouterRouting).toHaveBeenCalledWith(
|
||||
"token",
|
||||
expect.objectContaining({
|
||||
complexity_router_config: expect.objectContaining({
|
||||
classifier_type: "oss_classifier",
|
||||
opensource_classifier_config: {
|
||||
provider,
|
||||
model: provider === "laya" ? "multilingual" : "jev-test",
|
||||
timeout_ms: 4200,
|
||||
circuit_breaker_enabled: false,
|
||||
circuit_breaker_cooldown_seconds: 50,
|
||||
},
|
||||
tiers: expect.objectContaining({ QUICK: ["fast"] }),
|
||||
}),
|
||||
}),
|
||||
}),
|
||||
);
|
||||
});
|
||||
);
|
||||
},
|
||||
);
|
||||
|
||||
it("allows licensed instructions and can restore built-in instructions", () => {
|
||||
const authorized = useAuthorized();
|
||||
|
|
@ -152,10 +169,10 @@ describe("JEV classifier editor", () => {
|
|||
return <JevEditor value={value} onChange={setValue} />;
|
||||
};
|
||||
renderWithProviders(<LicensedForm />);
|
||||
expect(screen.getByLabelText("Jev Instructions")).toBeEnabled();
|
||||
fireEvent.change(screen.getByLabelText("Jev Instructions"), { target: { value: "New instructions" } });
|
||||
expect(screen.getByLabelText("Jev Instructions")).toHaveValue("New instructions");
|
||||
fireEvent.click(screen.getByRole("button", { name: "Restore built-in Jev instructions" }));
|
||||
expect(screen.getByLabelText("Jev Instructions")).toHaveValue("");
|
||||
expect(screen.getByLabelText("Classifier Instructions")).toBeEnabled();
|
||||
fireEvent.change(screen.getByLabelText("Classifier Instructions"), { target: { value: "New instructions" } });
|
||||
expect(screen.getByLabelText("Classifier Instructions")).toHaveValue("New instructions");
|
||||
fireEvent.click(screen.getByRole("button", { name: "Restore built-in instructions" }));
|
||||
expect(screen.getByLabelText("Classifier Instructions")).toHaveValue("");
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -6,7 +6,8 @@ import { Label } from "@/components/ui/label";
|
|||
import { Textarea } from "@/components/ui/textarea";
|
||||
import ClassifierCircuitBreakerConfig from "./ClassifierCircuitBreakerConfig";
|
||||
import type { ComplexityRouterConfigValue } from "./ComplexityRouterConfig";
|
||||
import { defaultJevClassifierConfig } from "./jev_classifier_config";
|
||||
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
|
||||
import { defaultJevClassifierConfig, LAYA_MODELS } from "./jev_classifier_config";
|
||||
|
||||
export default function JevClassifierConfig({
|
||||
value,
|
||||
|
|
@ -17,20 +18,38 @@ export default function JevClassifierConfig({
|
|||
}) {
|
||||
const id = useId();
|
||||
const config = value.jev_classifier_config ?? defaultJevClassifierConfig();
|
||||
const isLaya = config.provider === "laya";
|
||||
const update = (patch: Partial<typeof config>) =>
|
||||
onChange({ ...value, jev_classifier_config: { ...config, ...patch } });
|
||||
|
||||
return (
|
||||
<div className="mt-4 space-y-3">
|
||||
<p className="text-sm text-muted-foreground">
|
||||
Uses TypeSafe System One Choice evaluation with your configured tiers
|
||||
{isLaya
|
||||
? "Uses Laya with your configured tiers. Set LAYA_API_BASE on the gateway to connect your Laya server."
|
||||
: "Uses TypeSafe System One Choice evaluation with your configured tiers"}
|
||||
</p>
|
||||
<div>
|
||||
<Label htmlFor={`${id}-model`}>Jev Model</Label>
|
||||
<Input id={`${id}-model`} value={config.model} onChange={(event) => update({ model: event.target.value })} />
|
||||
<Label htmlFor={`${id}-model`}>Classifier Model</Label>
|
||||
{isLaya ? (
|
||||
<Select value={config.model} onValueChange={(model) => model && update({ model })}>
|
||||
<SelectTrigger id={`${id}-model`} className="w-full">
|
||||
<SelectValue />
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
{LAYA_MODELS.map((model) => (
|
||||
<SelectItem key={model} value={model}>
|
||||
{model}
|
||||
</SelectItem>
|
||||
))}
|
||||
</SelectContent>
|
||||
</Select>
|
||||
) : (
|
||||
<Input id={`${id}-model`} value={config.model} onChange={(event) => update({ model: event.target.value })} />
|
||||
)}
|
||||
</div>
|
||||
<div>
|
||||
<Label htmlFor={`${id}-timeout`}>Jev Timeout (ms)</Label>
|
||||
<Label htmlFor={`${id}-timeout`}>Classifier Timeout (ms)</Label>
|
||||
<Input
|
||||
id={`${id}-timeout`}
|
||||
type="number"
|
||||
|
|
@ -50,7 +69,7 @@ export default function JevClassifierConfig({
|
|||
}
|
||||
/>
|
||||
<div>
|
||||
<Label htmlFor={`${id}-instructions`}>Jev Instructions</Label>
|
||||
<Label htmlFor={`${id}-instructions`}>Classifier Instructions</Label>
|
||||
<AutoRouterAllowanceNote
|
||||
feature="tier_or_classifier_prompt"
|
||||
label="Custom instructions share the custom-tier allowance"
|
||||
|
|
@ -63,11 +82,11 @@ export default function JevClassifierConfig({
|
|||
/>
|
||||
{config.instructions && (
|
||||
<Button variant="outline" type="button" onClick={() => update({ instructions: undefined })}>
|
||||
Restore built-in Jev instructions
|
||||
Restore built-in instructions
|
||||
</Button>
|
||||
)}
|
||||
<p className="text-xs text-muted-foreground">
|
||||
Built-in Jev is available without a license and uses the shipped tier criteria
|
||||
Built-in OSS classification is available without a license and uses the shipped tier criteria
|
||||
</p>
|
||||
</div>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -107,10 +107,10 @@ describe("JEV network probes", () => {
|
|||
expect(JSON.parse(String(routingCall?.[1]?.body))).toEqual(expectedRequest);
|
||||
expect(fetchMock).toHaveBeenCalledTimes(5);
|
||||
expect(screen.getAllByTestId("test-status-success")).toHaveLength(4);
|
||||
expect(screen.getByRole("status", { name: "Jev connection" })).toHaveTextContent(
|
||||
expect(screen.getByRole("status", { name: "OSS classifier connection" })).toHaveTextContent(
|
||||
cause === "jev_classifier"
|
||||
? "Jev classification succeeded"
|
||||
: `Jev was not reached successfully (routing cause: ${cause})`,
|
||||
? "OSS classification succeeded"
|
||||
: `OSS classifier was not reached successfully (routing cause: ${cause})`,
|
||||
);
|
||||
},
|
||||
);
|
||||
|
|
@ -131,7 +131,7 @@ describe("JEV network probes", () => {
|
|||
);
|
||||
fireEvent.change(screen.getByTestId("auto-router-routing-test-prompt"), { target: { value: "Hello" } });
|
||||
fireEvent.click(screen.getByTestId("auto-router-routing-test-send"));
|
||||
expect(await screen.findByText("JEV classifier")).toBeInTheDocument();
|
||||
expect(await screen.findByText("OSS classifier")).toBeInTheDocument();
|
||||
expect(screen.getByText("jev-latest")).toBeInTheDocument();
|
||||
expect(screen.getByText("80.0%")).toBeInTheDocument();
|
||||
expect(screen.getByText("SIMPLE: 80.0%")).toBeInTheDocument();
|
||||
|
|
|
|||
|
|
@ -274,13 +274,13 @@ describe("AddAutoRouterTab", () => {
|
|||
mockFetchAvailableModels.mockResolvedValue(ALL_FAMILY_MODELS);
|
||||
renderWithProviders(<Harness />);
|
||||
await user.click(await screen.findByRole("button", { name: "Choose models for me" }));
|
||||
await user.click(screen.getByRole("radio", { name: "Jev" }));
|
||||
await user.click(screen.getByRole("radio", { name: "OSS Classifier" }));
|
||||
await waitFor(() =>
|
||||
expect(apiClient.post).toHaveBeenLastCalledWith(
|
||||
"/auto_router/availability",
|
||||
expect.objectContaining({
|
||||
body: expect.objectContaining({
|
||||
complexity_router_config: expect.objectContaining({ classifier_type: "jev" }),
|
||||
complexity_router_config: expect.objectContaining({ classifier_type: "oss_classifier" }),
|
||||
}),
|
||||
}),
|
||||
),
|
||||
|
|
@ -301,13 +301,13 @@ describe("AddAutoRouterTab", () => {
|
|||
expect(within(screen.getByRole("alert")).getByRole("link", { name: "Talk to our team" })).toBeVisible();
|
||||
await user.click(screen.getByRole("button", { name: "Restore defaults" }));
|
||||
await waitFor(() => expect(screen.queryByRole("alert")).not.toBeInTheDocument());
|
||||
expect(screen.getByRole("radio", { name: "Jev" })).toBeChecked();
|
||||
expect(screen.getByRole("radio", { name: "OSS Classifier" })).toBeChecked();
|
||||
await waitFor(() => expect(screen.getByRole("button", { name: "Add Auto Router" })).toBeEnabled());
|
||||
await user.click(screen.getByRole("button", { name: "Add Auto Router" }));
|
||||
await waitFor(() => expect(handleAddAutoRouterSubmit).toHaveBeenCalled());
|
||||
const saved = vi.mocked(handleAddAutoRouterSubmit).mock.calls.at(-1)?.[0].complexity_router_config;
|
||||
expect(saved).not.toHaveProperty("tier_definitions");
|
||||
expect(saved?.classifier_type).toBe("jev");
|
||||
expect(saved?.classifier_type).toBe("oss_classifier");
|
||||
expect(Object.keys(saved?.tiers ?? {})).toEqual(["SIMPLE", "MEDIUM", "COMPLEX", "REASONING"]);
|
||||
expect(saved?.tiers).toEqual(initialRequest.complexity_router_config.tiers);
|
||||
});
|
||||
|
|
@ -317,7 +317,7 @@ describe("AddAutoRouterTab", () => {
|
|||
mockFetchAvailableModels.mockResolvedValue(ALL_FAMILY_MODELS);
|
||||
renderWithProviders(<Harness />);
|
||||
await user.click(await screen.findByRole("button", { name: "Choose models for me" }));
|
||||
await user.click(screen.getByRole("radio", { name: "Jev" }));
|
||||
await user.click(screen.getByRole("radio", { name: "OSS Classifier" }));
|
||||
fireEvent.change(screen.getByLabelText("Auto Router Name"), { target: { value: "checked-router" } });
|
||||
await waitFor(() => expect(screen.getByRole("button", { name: "Add Auto Router" })).toBeEnabled());
|
||||
let complete: ((result: unknown) => void) | undefined;
|
||||
|
|
@ -357,17 +357,22 @@ describe("AddAutoRouterTab", () => {
|
|||
expect(screen.getByRole("button", { name: "Routing approach" })).toHaveTextContent("Complexity");
|
||||
});
|
||||
|
||||
it.each(["LLM", "Jev"])("keeps %s and the frequency when choosing models automatically", async (family) => {
|
||||
mockFetchAvailableModels.mockResolvedValue(ALL_FAMILY_MODELS);
|
||||
renderWithProviders(<Harness />);
|
||||
const automatic = await screen.findByRole("button", { name: "Choose models for me" });
|
||||
await userEvent.click(screen.getByRole("radio", { name: family }));
|
||||
await selectAutoRouterOption("How often to classify", "Every new user message");
|
||||
await userEvent.click(automatic);
|
||||
expect(screen.getByRole("radio", { name: family })).toBeChecked();
|
||||
expect(screen.getByRole("combobox", { name: "How often to classify" })).toHaveTextContent("Every new user message");
|
||||
expect(screen.getByRole("button", { name: "Advanced settings" })).toHaveAttribute("aria-expanded", "false");
|
||||
});
|
||||
it.each(["LLM", "OSS Classifier"])(
|
||||
"keeps %s and the frequency when choosing models automatically",
|
||||
async (family) => {
|
||||
mockFetchAvailableModels.mockResolvedValue(ALL_FAMILY_MODELS);
|
||||
renderWithProviders(<Harness />);
|
||||
const automatic = await screen.findByRole("button", { name: "Choose models for me" });
|
||||
await userEvent.click(screen.getByRole("radio", { name: family }));
|
||||
await selectAutoRouterOption("How often to classify", "Every new user message");
|
||||
await userEvent.click(automatic);
|
||||
expect(screen.getByRole("radio", { name: family })).toBeChecked();
|
||||
expect(screen.getByRole("combobox", { name: "How often to classify" })).toHaveTextContent(
|
||||
"Every new user message",
|
||||
);
|
||||
expect(screen.getByRole("button", { name: "Advanced settings" })).toHaveAttribute("aria-expanded", "false");
|
||||
},
|
||||
);
|
||||
|
||||
it.each(["Capability", "Fuse v2"])(
|
||||
"creates %s from its dedicated tab without complexity templates",
|
||||
|
|
@ -1902,7 +1907,7 @@ describe("preset catalog fetch states", () => {
|
|||
|
||||
await waitFor(() => expect(handleAddAutoRouterSubmit).toHaveBeenCalledOnce());
|
||||
expect(vi.mocked(handleAddAutoRouterSubmit).mock.calls[0][0].complexity_router_config).toMatchObject({
|
||||
classifier_type: "jev",
|
||||
classifier_type: "oss_classifier",
|
||||
classifier_context_per_turn_chars: 450,
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -50,7 +50,7 @@ const AutoRouterConnectionTest: React.FC<AutoRouterConnectionTestProps> = ({
|
|||
? { status: "success" }
|
||||
: {
|
||||
status: "error",
|
||||
error: `Jev was not reached successfully (routing cause: ${decision.cause ?? "unknown"})`,
|
||||
error: `OSS classifier was not reached successfully (routing cause: ${decision.cause ?? "unknown"})`,
|
||||
},
|
||||
);
|
||||
};
|
||||
|
|
@ -91,11 +91,11 @@ const AutoRouterConnectionTest: React.FC<AutoRouterConnectionTestProps> = ({
|
|||
classifier probe includes its reasoning effort override.
|
||||
</p>
|
||||
{jevRequest && (
|
||||
<div role="status" aria-label="Jev connection" className="rounded-lg border p-3 text-sm">
|
||||
<strong>Jev Classifier</strong>
|
||||
<div role="status" aria-label="OSS classifier connection" className="rounded-lg border p-3 text-sm">
|
||||
<strong>OSS Classifier</strong>
|
||||
<p>
|
||||
{jevResult.status === "pending" && "Testing Jev classification"}
|
||||
{jevResult.status === "success" && "Jev classification succeeded"}
|
||||
{jevResult.status === "pending" && "Testing OSS classification"}
|
||||
{jevResult.status === "success" && "OSS classification succeeded"}
|
||||
{jevResult.status === "error" && jevResult.error}
|
||||
</p>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -33,20 +33,20 @@ describe("buildAutoRouterRoutingTestRequest", () => {
|
|||
const expectedRequest = {
|
||||
prompt: JEV_CONNECTION_TEST_PROMPT,
|
||||
complexity_router_config: {
|
||||
classifier_type: "jev",
|
||||
classifier_type: "oss_classifier",
|
||||
tiers: CONFIG.tiers,
|
||||
jev_classifier_config: defaultJevClassifierConfig(),
|
||||
opensource_classifier_config: defaultJevClassifierConfig(),
|
||||
},
|
||||
saved_model_id: "saved-id",
|
||||
};
|
||||
expect(request).toEqual(expectedRequest);
|
||||
expect(request?.complexity_router_config.jev_classifier_config).not.toHaveProperty("api_key");
|
||||
expect(request?.complexity_router_config.jev_classifier_config).not.toHaveProperty("api_base");
|
||||
expect(request?.complexity_router_config.opensource_classifier_config).not.toHaveProperty("api_key");
|
||||
expect(request?.complexity_router_config.opensource_classifier_config).not.toHaveProperty("api_base");
|
||||
});
|
||||
it.each(["object", "json"])("probes saved JEV %s configuration with custom tiers and team context", (format) => {
|
||||
it.each(["object", "json"])("probes saved Laya %s configuration with custom tiers and team context", (format) => {
|
||||
const config = {
|
||||
classifier_type: "jev",
|
||||
jev_classifier_config: { model: "jev-test", timeout_ms: 900 },
|
||||
classifier_type: "oss_classifier",
|
||||
opensource_classifier_config: { provider: "laya", model: "english", timeout_ms: 900 },
|
||||
tiers: { QUICK: ["fast"], DEEP: ["strong"] },
|
||||
tier_definitions: { QUICK: "Simple questions", DEEP: "Complex questions" },
|
||||
fallback_tier: "DEEP",
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import { AutoRouterRoutingTestRequest } from "../networking";
|
||||
import { ComplexityRouterConfigPayload } from "./build_complexity_router_config";
|
||||
import { z } from "zod";
|
||||
import { jevClassifierConfigSchema } from "./jev_classifier_config";
|
||||
import { hydrateOssClassifier, jevClassifierConfigSchema, normalizeJevClassifierConfig } from "./jev_classifier_config";
|
||||
|
||||
export const JEV_CONNECTION_TEST_PROMPT = "What is 2 plus 2?";
|
||||
|
||||
|
|
@ -23,16 +23,23 @@ export const buildSavedJevConnectionTestRequest = (
|
|||
: rawConfig;
|
||||
const result = z
|
||||
.object({
|
||||
classifier_type: z.literal("jev"),
|
||||
classifier_type: z.enum(["jev", "oss_classifier"]),
|
||||
tiers: z.record(z.unknown()),
|
||||
jev_classifier_config: jevClassifierConfigSchema.default({}),
|
||||
jev_classifier_config: jevClassifierConfigSchema.optional(),
|
||||
opensource_classifier_config: jevClassifierConfigSchema.optional(),
|
||||
})
|
||||
.passthrough()
|
||||
.safeParse(parsed);
|
||||
if (!result.success) return undefined;
|
||||
const { jev_classifier_config, opensource_classifier_config, ...config } = result.data;
|
||||
const classifier = hydrateOssClassifier({ ...config, jev_classifier_config, opensource_classifier_config });
|
||||
return {
|
||||
prompt: JEV_CONNECTION_TEST_PROMPT,
|
||||
complexity_router_config: result.data,
|
||||
complexity_router_config: {
|
||||
...config,
|
||||
classifier_type: "oss_classifier",
|
||||
opensource_classifier_config: normalizeJevClassifierConfig(classifier.jev_classifier_config),
|
||||
},
|
||||
saved_model_id: savedModelId,
|
||||
...(teamId && { team_id: teamId }),
|
||||
};
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue