diff --git a/.github/workflows/test-e2e-changed.yml b/.github/workflows/test-e2e-changed.yml index 8e03a902383..228e23f60d7 100644 --- a/.github/workflows/test-e2e-changed.yml +++ b/.github/workflows/test-e2e-changed.yml @@ -176,6 +176,7 @@ jobs: TESTS: ${{ needs.detect.outputs.tests }} E2E_FIXTURE_MODE: live E2E_PROVIDER_EDGE_HOST_REACHABLE: '1' + E2E_OWNED_GATEWAY: '1' COLUMNS: '400' run: | umask 077 diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20261001200000_add_autorouter_daily_spend/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261001200000_add_autorouter_daily_spend/migration.sql new file mode 100644 index 00000000000..ce166b4df45 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261001200000_add_autorouter_daily_spend/migration.sql @@ -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") +); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 6f285e9dc39..aba89526cf6 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -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 diff --git a/litellm-proxy-extras/litellm_proxy_extras/utils.py b/litellm-proxy-extras/litellm_proxy_extras/utils.py index 3acc19d397d..2062ca93fb3 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/utils.py +++ b/litellm-proxy-extras/litellm_proxy_extras/utils.py @@ -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: diff --git a/litellm-rust/crates/traces/query/list_traces.sql b/litellm-rust/crates/traces/query/list_traces.sql index c0c1b28aa7f..163d9b03d6b 100644 --- a/litellm-rust/crates/traces/query/list_traces.sql +++ b/litellm-rust/crates/traces/query/list_traces.sql @@ -1,11 +1,13 @@ +WITH page AS ( SELECT TraceId AS trace_id, hex(SHA256(concat(TeamId, char(0), ApiKeyHash, char(0), TraceId))) AS trace_ref, TeamId AS team_id, ApiKeyHash AS api_key_hash, ifNull(any(RootName), '') AS name, any(ServiceName) AS service, ifNull(any(RootInput), '') AS input_preview, ifNull(any(RootStatus), '') AS status, toUnixTimestamp64Milli(min(StartTs)) AS start_ms, + min(StartTs) AS trace_start, max(EndTs) AS trace_end, dateDiff('millisecond', min(StartTs), max(EndTs)) AS duration_ms, - sum(SpanCount) AS span_count, length(groupUniqArrayArray(AgentNames)) AS agent_count, + sum(SpanCount) AS span_count, sum(AgentCount) AS agent_invocations, sum(LlmCount) AS llm_calls, sum(ToolCount) AS tool_calls, sum(InputTokens) AS input_tokens, sum(OutputTokens) AS output_tokens, @@ -21,3 +23,21 @@ HAVING min(StartTs) >= fromUnixTimestamp64Milli({start_ms:Int64}) < ({cursor_ms:Int64}, {cursor_trace_id:String})) ORDER BY start_ms DESC, trace_ref DESC LIMIT {limit:UInt32} +) +SELECT page.* EXCEPT (trace_start, trace_end), + identities.agent_names AS agent_names, identities.agent_count AS agent_count +FROM page +LEFT JOIN ( + SELECT TeamId, ApiKeyHash, TraceId, + arraySort(groupUniqArrayIf(AgentName, AgentName != '')) AS agent_names, + uniqExactIf(if(AgentName = '', SpanName, AgentName), ObservationType = 'agent') AS agent_count + FROM otel_traces + WHERE Timestamp >= (SELECT min(trace_start) FROM page) + AND Timestamp <= (SELECT max(trace_end) FROM page) + AND TraceId IN (SELECT trace_id FROM page) + AND (TeamId, ApiKeyHash, TraceId) IN (SELECT team_id, api_key_hash, trace_id FROM page) + GROUP BY TeamId, ApiKeyHash, TraceId +) AS identities +ON page.team_id = identities.TeamId AND page.api_key_hash = identities.ApiKeyHash + AND page.trace_id = identities.TraceId +ORDER BY page.start_ms DESC, page.trace_ref DESC diff --git a/litellm-rust/crates/traces/src/normalize/mod.rs b/litellm-rust/crates/traces/src/normalize/mod.rs index b4d5a4a9d06..60dc8f816f2 100644 --- a/litellm-rust/crates/traces/src/normalize/mod.rs +++ b/litellm-rust/crates/traces/src/normalize/mod.rs @@ -1,7 +1,7 @@ use std::collections::BTreeMap; use crate::DecodeError; -use serde::Serialize; +use serde::{Deserialize, Serialize}; #[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize)] #[serde(rename_all = "lowercase")] @@ -148,6 +148,60 @@ fn usage_tokens(attributes: &BTreeMap) -> Result<(u32, u32), Dec )) } +#[derive(Default, Deserialize)] +struct AgentMetadata { + #[serde(default)] + lc_agent_name: String, + #[serde(default)] + ls_integration: String, +} + +fn recorded_agent_name( + name: &str, + attributes: &BTreeMap, + span: &NormalizedSpan, +) -> String { + let explicit = [ + span.agent_name.as_str(), + attr(attributes, "gen_ai.agent.name"), + attr(attributes, "agent.name"), + attr(attributes, "openclaw.agent"), + ] + .into_iter() + .find(|value| !value.is_empty()); + if let Some(value) = explicit { + return value.to_owned(); + } + let metadata = + serde_json::from_str::(attr(attributes, "metadata")).unwrap_or_default(); + if !metadata.lc_agent_name.is_empty() { + return metadata.lc_agent_name; + } + if span.observation_type == ObservationType::Agent { + let node = attr(attributes, "graph.node.id"); + if !node.is_empty() { + return node.to_owned(); + } + if metadata.ls_integration == "langgraph" && name != "LangGraph" && !is_middleware(name) { + return name.to_owned(); + } + } + String::new() +} + +fn is_middleware(name: &str) -> bool { + [ + ".wrap_model_call", + ".wrap_tool_call", + ".before_agent", + ".after_agent", + ".before_model", + ".after_model", + ] + .iter() + .any(|suffix| name.ends_with(suffix)) +} + pub fn normalize( scope_name: &str, name: &str, @@ -163,8 +217,22 @@ pub fn normalize( .into_iter() .find(|normalizer| normalizer.matches(scope_name, attributes)) .expect("GenAI fallback always matches"); + let span = normalizer.normalize(name, parent_span_id, attributes)?; + let agent_name = recorded_agent_name(name, attributes, &span); + let observation_type = if !parent_span_id.is_empty() + && scope_name == "openinference.instrumentation.langchain" + && is_middleware(name) + { + ObservationType::Framework + } else { + span.observation_type + }; Ok(Normalization { - span: normalizer.normalize(name, parent_span_id, attributes)?, + span: NormalizedSpan { + agent_name, + observation_type, + ..span + }, consumed_attributes: normalizer.consumed_attributes(attributes), }) } diff --git a/litellm-rust/crates/traces/src/otlp/span.rs b/litellm-rust/crates/traces/src/otlp/span.rs index 1c1e53e756c..0ae5725ba47 100644 --- a/litellm-rust/crates/traces/src/otlp/span.rs +++ b/litellm-rust/crates/traces/src/otlp/span.rs @@ -133,7 +133,18 @@ fn decoded_span( &parent_span_id, &span_attributes, )?; - let normalized = normalization.span; + let resource_agent_name = resource_attributes + .get("gen_ai.agent.name") + .filter(|name| !name.is_empty()); + let agent_name = match (resource_agent_name, normalization.span.agent_name.as_str()) { + (Some(name), "") => name.clone(), + (Some(name), "hermes-agent") if scope_name.as_ref() == "hermes-otel-plugin" => name.clone(), + (_, name) => name.to_owned(), + }; + let normalized = crate::normalize::NormalizedSpan { + agent_name, + ..normalization.span + }; budget.consume( normalized.input.len() + normalized.output.len() diff --git a/litellm-rust/crates/traces/tests/migrations.rs b/litellm-rust/crates/traces/tests/migrations.rs index ccefa8b0b9f..c62e3538ddb 100644 --- a/litellm-rust/crates/traces/tests/migrations.rs +++ b/litellm-rust/crates/traces/tests/migrations.rs @@ -362,6 +362,143 @@ async fn keyed_rollup_keeps_same_trace_ids_separate_by_api_key( Ok(()) } +#[rstest] +#[tokio::test] +async fn listed_agent_names_preserve_scope_and_cursor( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let writer = Connection::writer(&database.url)?; + ensure_schema(&database.client, &writer, "trace_test", 7).await?; + let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64; + for (team, key, trace, agent, span, parent) in [ + ("alpha", "one", "shared", "research_agent", "root", ""), + ("alpha", "one", "shared", "reviewer", "child", "root"), + ("alpha", "one", "shared", "reviewer", "repeated", "root"), + ("alpha", "one", "shared", "", "unnamed", "root"), + ("alpha", "one", "second", "support_agent", "root", ""), + ("alpha", "two", "shared", "private_agent", "root", ""), + ("beta", "one", "shared", "other_agent", "root", ""), + ] { + insert_rows( + &database, + "otel_traces", + vec![serde_json::from_value(serde_json::json!({ + "Timestamp": timestamp, "TraceId": trace, "SpanId": span, "ParentSpanId": parent, + "ServiceName": "shared-app", "SpanName": span, "AgentName": agent, + "ObservationType": "agent", + "ResourceAttributes": {"litellm.team_id": team, "litellm.api_key_hash": key} + }))?], + ) + .await?; + } + let historical_rows = (0..5000) + .map(|index| { + serde_json::from_value(serde_json::json!({ + "Timestamp": timestamp - 86_400_000_000_000_i64, + "TraceId": "shared", "SpanId": format!("historical-{index}"), + "ParentSpanId": "", "SpanName": "historical", "AgentName": "private_agent", + "ObservationType": "agent", "ServiceName": "shared-app", + "ResourceAttributes": {"litellm.team_id": "alpha", "litellm.api_key_hash": "history"} + })) + }) + .collect::, _>>()?; + insert_rows(&database, "otel_traces", historical_rows).await?; + let connection = Connection::configured(&database.url, "trace_test", "default", "")?; + let parameters = BTreeMap::from([ + ("team_ids".into(), Parameter::Strings(vec!["alpha".into()])), + ("api_key_hash".into(), Parameter::Text("one".into())), + ( + "start_ms".into(), + Parameter::Integer(timestamp / 1_000_000 - 1000), + ), + ( + "end_ms".into(), + Parameter::Integer(timestamp / 1_000_000 + 1000), + ), + ("cursor_ms".into(), Parameter::Integer(0)), + ("cursor_trace_id".into(), Parameter::Text(String::new())), + ("limit".into(), Parameter::Integer(1)), + ]); + let first: serde_json::Value = serde_json::from_str( + &execute_named_read( + &database.client, + &connection, + ReadQuery::ListTraces, + ¶meters, + ) + .await?, + )?; + let cursor = first["data"][0]["trace_ref"] + .as_str() + .ok_or("missing cursor")?; + let next_parameters = parameters + .into_iter() + .chain([ + ( + "cursor_ms".into(), + Parameter::Integer(timestamp / 1_000_000), + ), + ("cursor_trace_id".into(), Parameter::Text(cursor.into())), + ]) + .collect(); + let second: serde_json::Value = serde_json::from_str( + &execute_named_read( + &database.client, + &connection, + ReadQuery::ListTraces, + &next_parameters, + ) + .await?, + )?; + assert_eq!( + first["data"].as_array().ok_or("missing first page")?.len(), + 1 + ); + assert_eq!( + second["data"] + .as_array() + .ok_or("missing second page")? + .len(), + 1 + ); + assert_ne!(first["data"][0]["trace_id"], second["data"][0]["trace_id"]); + let names = [&first["data"][0], &second["data"][0]] + .into_iter() + .map(|row| { + ( + row["trace_id"].as_str().unwrap(), + row["agent_names"].clone(), + ) + }) + .collect::>(); + assert_eq!( + names["shared"], + serde_json::json!(["research_agent", "reviewer"]) + ); + assert_eq!(names["second"], serde_json::json!(["support_agent"])); + let counts = [&first["data"][0], &second["data"][0]] + .into_iter() + .map(|row| { + ( + row["trace_id"].as_str().unwrap(), + row["agent_count"].as_u64(), + ) + }) + .collect::>(); + assert_eq!(counts["shared"], Some(3)); + assert_eq!(counts["second"], Some(1)); + for page in [&first, &second] { + assert!( + page["statistics"]["rows_read"] + .as_u64() + .ok_or("missing read statistics")? + < 5000 + ); + } + Ok(()) +} + #[rstest] #[tokio::test] async fn rollup_merges_spans_across_days_without_losing_root_fields( @@ -376,6 +513,7 @@ async fn rollup_merges_spans_across_days_without_losing_root_fields( let root = serde_json::from_value(serde_json::json!({ "Timestamp": day_start - 1_000_000_000, "TraceId": "cross-day", "SpanId": "span-root", "ParentSpanId": "", "ServiceName": "proxy", "SpanName": "root", "Input": "root input", + "AgentName": "lead", "ObservationType": "agent", "StatusCode": "STATUS_CODE_ERROR", "ResourceAttributes": {"litellm.team_id": "team-1"} }))?; @@ -383,6 +521,7 @@ async fn rollup_merges_spans_across_days_without_losing_root_fields( let child = serde_json::from_value(serde_json::json!({ "Timestamp": day_start + 1_000_000_000, "TraceId": "cross-day", "SpanId": "span-child", "ParentSpanId": "span-root", "ServiceName": "proxy", "SpanName": "child", + "AgentName": "researcher", "ObservationType": "agent", "StatusCode": "STATUS_CODE_UNSET", "ResourceAttributes": {"litellm.team_id": "team-1"} }))?; @@ -406,6 +545,33 @@ async fn rollup_merges_spans_across_days_without_losing_root_fields( "RootStatus": "STATUS_CODE_ERROR", "SpanCount": 2 }]) ); + let connection = Connection::configured(&database.url, "trace_test", "default", "")?; + let parameters = BTreeMap::from([ + ("team_ids".into(), Parameter::Strings(vec!["team-1".into()])), + ("api_key_hash".into(), Parameter::Text(String::new())), + ( + "start_ms".into(), + Parameter::Integer(day_start / 1_000_000 - 2000), + ), + ("end_ms".into(), Parameter::Integer(day_start / 1_000_000)), + ("cursor_ms".into(), Parameter::Integer(0)), + ("cursor_trace_id".into(), Parameter::Text(String::new())), + ("limit".into(), Parameter::Integer(10)), + ]); + let listed: serde_json::Value = serde_json::from_str( + &execute_named_read( + &database.client, + &connection, + ReadQuery::ListTraces, + ¶meters, + ) + .await?, + )?; + assert_eq!( + listed["data"][0]["agent_names"], + serde_json::json!(["lead", "researcher"]) + ); + assert_eq!(listed["data"][0]["agent_count"], 2); Ok(()) } diff --git a/litellm/integrations/otel/logger.py b/litellm/integrations/otel/logger.py index 55eb8e8fb71..e1983a44451 100644 --- a/litellm/integrations/otel/logger.py +++ b/litellm/integrations/otel/logger.py @@ -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) diff --git a/litellm/integrations/otel/model/config.py b/litellm/integrations/otel/model/config.py index 9eb29157d6f..ddb8e127408 100644 --- a/litellm/integrations/otel/model/config.py +++ b/litellm/integrations/otel/model/config.py @@ -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 " diff --git a/litellm/llms/laya/__init__.py b/litellm/llms/laya/__init__.py new file mode 100644 index 00000000000..8b137891791 --- /dev/null +++ b/litellm/llms/laya/__init__.py @@ -0,0 +1 @@ + diff --git a/litellm/llms/laya/common_utils.py b/litellm/llms/laya/common_utils.py new file mode 100644 index 00000000000..f400eef22d3 --- /dev/null +++ b/litellm/llms/laya/common_utils.py @@ -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" diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 6a548e7a82d..3985a232f62 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -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", diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index 0119ce4430e..837552e522f 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -229,6 +229,7 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = ( "/tinyfish/", "/transcribe", "/typesafe/", + "/laya/", "/openrouter/", "/vertex-ai/", "/vertex_ai/", diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 5de6e91a0f2..3346b0c9ff8 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -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.", diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index fe505ad6b06..0dfbf2c5081 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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: diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 813d72ed7ae..e5a1430b3d8 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -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) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 3400dccf2a7..e82f3eed7cc 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -140,6 +140,7 @@ from litellm.proxy.utils import ( normalize_route_for_root_path, ) from litellm.repositories.table_repositories import TeamMembershipRepository +from litellm.repositories.verification_token_repository import VerificationTokenRepository from litellm.router_utils.common_utils import resolve_model_group_alias from litellm.secret_managers.main import get_secret_bool from litellm.types.services import ServiceTypes @@ -939,6 +940,24 @@ class _PendingAutoRegister(NamedTuple): jwt_issuer: str | None = None +def _claim_identifies_user(jwt_handler: JWTHandler, claim_field: str, jwt_issuer: str | None) -> bool: + if not jwt_handler.litellm_jwtauth.auto_register_map_existing_key: + return False + if jwt_handler.litellm_jwtauth.is_user_identity_claim(claim_field, jwt_issuer): + return True + verbose_proxy_logger.warning( + "JWT Key Mapping (auto_register_map_existing_key): claim '%s' is not the user_id or user_email JWT field " + "and may be shared by several users, so a new key is minted instead of reusing one the user owns.", + claim_field, + ) + return False + + +async def _reusable_key_hash_for_user(prisma_client: PrismaClient, user_id: str, team_id: str | None) -> str | None: + key: Final = await VerificationTokenRepository(prisma_client).find_newest_reusable_llm_api_key(user_id, team_id) + return None if key is None else key.token + + async def _auto_register_jwt_mapping( virtual_key_claim_field: str, claim_value: str, @@ -957,8 +976,10 @@ async def _auto_register_jwt_mapping( ) -> UserAPIKeyAuth | None: """ Auto-register: create a new virtual key + mapping for an unrecognised JWT - claim value. ``team_id`` and ``user_id`` must come from a successful - ``JWTAuthManager.auth_builder`` run — they encode the JWT identity AFTER + claim value, or point the mapping at a key the resolved user already owns + when ``auto_register_map_existing_key`` is set. ``team_id`` and ``user_id`` + must come from a successful ``JWTAuthManager.auth_builder`` run — they + encode the JWT identity AFTER RBAC/scope/custom_validate/email-domain policy has been enforced. The key is stamped with those values so the cached future-request path inherits the same team/user/org limits the auth_builder path would have applied. @@ -974,29 +995,38 @@ async def _auto_register_jwt_mapping( generate_key_helper_fn, ) - # ``table_name="key"`` is required: without it, generate_key_helper_fn - # falls into the user-upsert branch (`table_name is None or "user"`) and - # attempts to insert into LiteLLM_UserTable with user_id=None, which fails - # the NOT NULL @id constraint. Every successful key-creation caller (e.g. - # /key/generate) passes table_name="key" explicitly. - key_data: Final = await generate_key_helper_fn( - llm_router=None, - request_type="key", - table_name="key", - team_id=team_id, - user_id=user_id, - organization_id=org_id, - agent_id=agent_id, - metadata={ - "auto_registered": True, - "jwt_claim_field": virtual_key_claim_field, - "jwt_claim_value": claim_value, - }, + existing_token_hash: Final = ( + await _reusable_key_hash_for_user(prisma_client, user_id, team_id) + if user_id is not None and _claim_identifies_user(jwt_handler, virtual_key_claim_field, jwt_issuer) + else None ) - # generate_key_helper_fn returns the plaintext key in "token"; the persisted - # row in LiteLLM_VerificationToken uses its hash, so hash here to get the FK - # value referenced by LiteLLM_JWTKeyMapping.token. - token_hash = hash_token(key_data["token"]) + minted: Final = existing_token_hash is None + if existing_token_hash is not None: + token_hash = existing_token_hash + else: + # ``table_name="key"`` is required: without it, generate_key_helper_fn + # falls into the user-upsert branch (`table_name is None or "user"`) and + # attempts to insert into LiteLLM_UserTable with user_id=None, which fails + # the NOT NULL @id constraint. Every successful key-creation caller (e.g. + # /key/generate) passes table_name="key" explicitly. + key_data: Final = await generate_key_helper_fn( + llm_router=None, + request_type="key", + table_name="key", + team_id=team_id, + user_id=user_id, + organization_id=org_id, + agent_id=agent_id, + metadata={ + "auto_registered": True, + "jwt_claim_field": virtual_key_claim_field, + "jwt_claim_value": claim_value, + }, + ) + # generate_key_helper_fn returns the plaintext key in "token"; the persisted + # row in LiteLLM_VerificationToken uses its hash, so hash here to get the FK + # value referenced by LiteLLM_JWTKeyMapping.token. + token_hash = hash_token(key_data["token"]) try: await prisma_client.db.litellm_jwtkeymapping.create( @@ -1023,15 +1053,16 @@ async def _auto_register_jwt_mapping( virtual_key_claim_field, claim_value, ) - try: - await prisma_client.db.litellm_verificationtoken.delete(where={"token": token_hash}) - except Exception as delete_err: - # Don't fail the request if cleanup fails — the orphan is - # unmapped and inert. Log so an operator can prune it later. - verbose_proxy_logger.warning( - "JWT Key Mapping (auto_register): failed to delete orphaned key after race: %s", - delete_err, - ) + if minted: + try: + await prisma_client.db.litellm_verificationtoken.delete(where={"token": token_hash}) + except Exception as delete_err: + # Don't fail the request if cleanup fails — the orphan is + # unmapped and inert. Log so an operator can prune it later. + verbose_proxy_logger.warning( + "JWT Key Mapping (auto_register): failed to delete orphaned key after race: %s", + delete_err, + ) token_hash = await get_jwt_key_mapping_object( jwt_claim_name=virtual_key_claim_field, jwt_claim_value=claim_value, @@ -1061,7 +1092,8 @@ async def _auto_register_jwt_mapping( ) verbose_proxy_logger.info( - "JWT Key Mapping (auto_register): created new virtual key for %s='%s'.", + "JWT Key Mapping (auto_register): %s virtual key for %s='%s'.", + "created new" if minted else "mapped existing", virtual_key_claim_field, claim_value, ) @@ -1075,7 +1107,8 @@ async def _auto_register_jwt_mapping( ).resolve(hashed_token=token_hash) ) if auto_registered_key is not None: - auto_registered_key.org_id = org_id + if minted: + auto_registered_key.org_id = org_id auto_registered_key.end_user_id = end_user_id auto_registered_key.api_key = auto_registered_key.token return auto_registered_key @@ -1771,8 +1804,8 @@ async def _user_api_key_auth_builder( # mapping + virtual key from the *validated* identity, then # replace valid_token with the new key so downstream checks # use the key-scoped path. - if pending_auto_register is not None and prisma_client is not None: - auto_registered: Final = await _auto_register_jwt_mapping( + auto_registered: Final = ( + await _auto_register_jwt_mapping( virtual_key_claim_field=pending_auto_register.claim_field, claim_value=pending_auto_register.claim_value, jwt_handler=jwt_handler, @@ -1788,72 +1821,81 @@ async def _user_api_key_auth_builder( end_user_id=end_user_id, agent_id=agent_id, ) - if auto_registered is not None: - auto_registered.jwt_claims = jwt_claims - auto_registered.user_email = user_email - # The auto-registered token is built from the new key's - # columns, which carry no user budget. Carry over the - # already-loaded user row rather than re-reading it, or - # the budget check below has nothing to enforce. - auto_registered.user_model_max_budget = ( - user_object.model_max_budget if user_object is not None else None - ) - valid_token = auto_registered - api_key = valid_token.token or "" - - # Check if model has zero cost - if so, skip all budget checks - model = _get_model_from_request_context( - request_data=request_data, - route=route, - request=request, - llm_router=llm_router, - team_id=valid_token.team_id, + if pending_auto_register is not None and prisma_client is not None + else None ) - skip_budget_checks = False - if model is not None and llm_router is not None: - from litellm.proxy.auth.auth_checks import _is_model_cost_zero - - skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router) - if skip_budget_checks: - verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model) - - # Fetch project object for JWT path if project_id is set - _jwt_project_obj = None - if valid_token.project_id is not None: - _jwt_project_obj = await get_project_object( - project_id=valid_token.project_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - proxy_logging_obj=proxy_logging_obj, + if auto_registered is not None: + auto_registered.jwt_claims = jwt_claims + auto_registered.user_email = user_email + # The auto-registered token is built from the new key's + # columns, which carry no user budget. Carry over the + # already-loaded user row rather than re-reading it, or + # the budget check below has nothing to enforce. + auto_registered.user_model_max_budget = ( + user_object.model_max_budget if user_object is not None else None ) - if _jwt_project_obj is not None: - valid_token.project_metadata = _jwt_project_obj.metadata - valid_token.project_alias = _jwt_project_obj.project_alias + valid_token = auto_registered + api_key = valid_token.token or "" - # JWT auth returns here rather than falling through to the - # virtual-key checks below, so the user's per-model budget - # has to be enforced on this path too. Without it the - # post-call increment still charges the counter and nothing - # ever reads it, which is worse than not tracking at all. - # Guarded by the same flag the virtual-key path uses, or a - # zero-cost model would be refused here and allowed there, - # while the log above claims all budget checks were skipped. - if not skip_budget_checks: - await _check_user_model_budget( - valid_token=cast(UserAPIKeyAuth, valid_token), - model_max_budget_limiter=model_max_budget_limiter, - models=_get_model_names_for_budget_checks( - model=_get_model_from_request_context( - request_data=request_data, - route=route, - request=request, - llm_router=llm_router, - team_id=valid_token.team_id, - ) - ), + falls_through_to_key_checks: Final = ( + auto_registered is not None + and jwt_handler.litellm_jwtauth.auto_register_map_existing_key + and master_key is not None + ) + if not falls_through_to_key_checks: + # Check if model has zero cost - if so, skip all budget checks + model = _get_model_from_request_context( + request_data=request_data, + route=route, + request=request, + llm_router=llm_router, + team_id=valid_token.team_id, ) + skip_budget_checks = False + if model is not None and llm_router is not None: + from litellm.proxy.auth.auth_checks import _is_model_cost_zero - return cast(UserAPIKeyAuth, valid_token) + skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router) + if skip_budget_checks: + verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model) + + # Fetch project object for JWT path if project_id is set + _jwt_project_obj = None + if valid_token.project_id is not None: + _jwt_project_obj = await get_project_object( + project_id=valid_token.project_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + if _jwt_project_obj is not None: + valid_token.project_metadata = _jwt_project_obj.metadata + valid_token.project_alias = _jwt_project_obj.project_alias + + # JWT auth returns here rather than falling through to the + # virtual-key checks below, so the user's per-model budget + # has to be enforced on this path too. Without it the + # post-call increment still charges the counter and nothing + # ever reads it, which is worse than not tracking at all. + # Guarded by the same flag the virtual-key path uses, or a + # zero-cost model would be refused here and allowed there, + # while the log above claims all budget checks were skipped. + if not skip_budget_checks: + await _check_user_model_budget( + valid_token=cast(UserAPIKeyAuth, valid_token), + model_max_budget_limiter=model_max_budget_limiter, + models=_get_model_names_for_budget_checks( + model=_get_model_from_request_context( + request_data=request_data, + route=route, + request=request, + llm_router=llm_router, + team_id=valid_token.team_id, + ) + ), + ) + + return cast(UserAPIKeyAuth, valid_token) #### ELSE #### ## CHECK PASS-THROUGH ENDPOINTS ## diff --git a/litellm/proxy/db/autorouter_session_rollup.py b/litellm/proxy/db/autorouter_session_rollup.py index b762a40f344..cdc949a6dca 100644 --- a/litellm/proxy/db/autorouter_session_rollup.py +++ b/litellm/proxy/db/autorouter_session_rollup.py @@ -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: diff --git a/litellm/proxy/db/baseline_accounting.py b/litellm/proxy/db/baseline_accounting.py index 05a4a989152..9536f8d740a 100644 --- a/litellm/proxy/db/baseline_accounting.py +++ b/litellm/proxy/db/baseline_accounting.py @@ -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) diff --git a/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py b/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py index 06e4d06fca4..679e286cedd 100644 --- a/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py +++ b/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py @@ -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 diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index d6daf6ebe59..a809f53aa85 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -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 diff --git a/litellm/proxy/management_endpoints/auto_router_endpoints.py b/litellm/proxy/management_endpoints/auto_router_endpoints.py index 8b73b8177f4..35ff9186f5e 100644 --- a/litellm/proxy/management_endpoints/auto_router_endpoints.py +++ b/litellm/proxy/management_endpoints/auto_router_endpoints.py @@ -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, ) diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 0f0d156d649..50bb831d169 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -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): diff --git a/litellm/proxy/management_helpers/auto_router_permissions.py b/litellm/proxy/management_helpers/auto_router_permissions.py index c8bb95eb3bb..e0d8fda5b1c 100644 --- a/litellm/proxy/management_helpers/auto_router_permissions.py +++ b/litellm/proxy/management_helpers/auto_router_permissions.py @@ -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 diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index c7b597b584f..ed2ea475c7a 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -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"], diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/batch_attribution.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/batch_attribution.py index e7b608e162e..880fdad92bf 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/batch_attribution.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/batch_attribution.py @@ -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: diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/typesafe_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/typesafe_passthrough_logging_handler.py index 3ad92acb48a..03ec559b83d 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/typesafe_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/typesafe_passthrough_logging_handler.py @@ -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() diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 22ecdc06ed9..865374a0430 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -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 {} diff --git a/litellm/proxy/pass_through_endpoints/success_handler.py b/litellm/proxy/pass_through_endpoints/success_handler.py index 6bba879b6c1..3c4733d0bf0 100644 --- a/litellm/proxy/pass_through_endpoints/success_handler.py +++ b/litellm/proxy/pass_through_endpoints/success_handler.py @@ -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, diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 6f285e9dc39..aba89526cf6 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -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 diff --git a/litellm/repositories/verification_token_repository.py b/litellm/repositories/verification_token_repository.py index d02c2114136..b20fb47306e 100644 --- a/litellm/repositories/verification_token_repository.py +++ b/litellm/repositories/verification_token_repository.py @@ -8,6 +8,7 @@ from datetime import datetime from types import TracebackType from typing import TYPE_CHECKING, Final, Protocol +from litellm.constants import UI_SESSION_TOKEN_TEAM_ID from litellm.models.verification_token import ( LiteLLM_VerificationToken, ) @@ -123,6 +124,37 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]): records: Final[Sequence[PrismaVerificationToken]] = await self.table.find_many(where={"user_id": user_id}) return self._to_model_list(records) + async def find_newest_reusable_llm_api_key( + self, user_id: str, team_id: str | None + ) -> LiteLLM_VerificationToken | None: + records: Final[Sequence[PrismaVerificationToken]] = await self.table.find_many( + where={ + "user_id": user_id, + "team_id": team_id, + "expires": None, + "AND": [ + {"OR": [{"blocked": False}, {"blocked": None}]}, + { + "OR": [ + {"team_id": None}, + {"team_id": {"not": UI_SESSION_TOKEN_TEAM_ID}}, + ] + }, + { + "OR": [ + {"allowed_routes": {"is_empty": True}}, + {"allowed_routes": {"has": "llm_api_routes"}}, + ] + }, + ], + }, + order={"created_at": "desc"}, + ) + return next( + (key for key in self._to_model_list(records) if key.metadata.get("auto_registered") is not True), + None, + ) + async def find_by_team_id(self, team_id: str) -> list[LiteLLM_VerificationToken]: """Find all tokens belonging to a team.""" records: Final[Sequence[PrismaVerificationToken]] = await self.table.find_many(where={"team_id": team_id}) diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 3f56141b21c..39fb237917c 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -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") diff --git a/litellm/router_strategy/complexity_router/config.py b/litellm/router_strategy/complexity_router/config.py index 00cff661d2f..41f389db7d8 100644 --- a/litellm/router_strategy/complexity_router/config.py +++ b/litellm/router_strategy/complexity_router/config.py @@ -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 diff --git a/litellm/router_strategy/complexity_router/jev_classifier.py b/litellm/router_strategy/complexity_router/jev_classifier.py index a2f03b07e3a..073d87c25a6 100644 --- a/litellm/router_strategy/complexity_router/jev_classifier.py +++ b/litellm/router_strategy/complexity_router/jev_classifier.py @@ -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: diff --git a/litellm/router_utils/auto_router_model_naming.py b/litellm/router_utils/auto_router_model_naming.py index 6b589c3bfc0..3423836e4fd 100644 --- a/litellm/router_utils/auto_router_model_naming.py +++ b/litellm/router_utils/auto_router_model_naming.py @@ -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}')" ), ) diff --git a/litellm/tracing/store.py b/litellm/tracing/store.py index 91420ffd025..edfe1285fc7 100644 --- a/litellm/tracing/store.py +++ b/litellm/tracing/store.py @@ -119,6 +119,7 @@ def trace_summary_from_row(row: dict[str, Any], spend_rows: Sequence[_SpendRow] trace_ref=row.get("trace_ref", ""), name=row["name"], service=row["service"], + agent_names=tuple(row.get("agent_names") or ()), input_preview=row["input_preview"], start_time=_iso(int(row["start_ms"])), duration_ms=float(row["duration_ms"]), @@ -169,8 +170,8 @@ def _parent_agent_of(span: Span, by_id: Mapping[str, Span]) -> str | None: if parent_id is None or parent_id not in by_id or parent_id == span["span_id"]: return None parent = by_id[parent_id] - if parent["type"] == "agent" and parent["name"] != span["name"]: - return parent["name"] + if parent["type"] == "agent" and (parent["agent"] or parent["name"]) != (span["agent"] or span["name"]): + return parent["agent"] or parent["name"] parent_id = parent["parent_span_id"] return None @@ -183,9 +184,9 @@ def agent_nodes(spans: Sequence[Span]) -> tuple[AgentNode, ...]: if span["type"] != "agent": continue node = agents.setdefault( - span["name"], + span["agent"] or span["name"], AgentNode( - name=span["name"], + name=span["agent"] or span["name"], parent_agent=_parent_agent_of(span, by_id), invocations=0, llm_calls=0, @@ -250,6 +251,7 @@ def trace_from_rows( trace_ref=trace_ref, name=root["name"], service=rows[0]["service"], + agent_names=tuple(sorted(frozenset(s["agent"] for s in spans if s["agent"]))), input_preview=root["input_preview"], start_time=_iso(trace_start_ns // NANOS_PER_MS), duration_ms=(trace_end_ns - trace_start_ns) / NANOS_PER_MS, diff --git a/litellm/tracing/types.py b/litellm/tracing/types.py index ff965483013..a0e824982fa 100644 --- a/litellm/tracing/types.py +++ b/litellm/tracing/types.py @@ -56,6 +56,7 @@ class TraceSummary(TypedDict): trace_ref: ReadOnly[NotRequired[str]] name: ReadOnly[str] service: ReadOnly[str] + agent_names: ReadOnly[NotRequired[tuple[str, ...]]] input_preview: ReadOnly[str] start_time: ReadOnly[str] # ISO 8601 duration_ms: ReadOnly[float] diff --git a/litellm/types/management_endpoints/auto_router_endpoints.py b/litellm/types/management_endpoints/auto_router_endpoints.py index 00083e01f54..ded971f6705 100644 --- a/litellm/types/management_endpoints/auto_router_endpoints.py +++ b/litellm/types/management_endpoints/auto_router_endpoints.py @@ -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") diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 6a548e7a82d..3985a232f62 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -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", diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index 9cbd326277e..d18f8d2e6d1 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -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", diff --git a/schema.prisma b/schema.prisma index 6f285e9dc39..aba89526cf6 100644 --- a/schema.prisma +++ b/schema.prisma @@ -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 diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py index 62153e38a83..37f0bf00da6 100644 --- a/tests/e2e/conftest.py +++ b/tests/e2e/conftest.py @@ -32,6 +32,7 @@ from e2e_config import ( MCP_OAUTH_LIVE_OPT_IN_ENV, OTEL_TLS_OPT_IN_ENV, OTEL_V2_OPT_IN_ENV, + OWNED_GATEWAY_OPT_IN_ENV, PROMPT_CACHING_OPT_IN_ENV, PROVIDER_EDGE_HOST_OPT_IN_ENV, PROXY_BASE_URL, @@ -70,6 +71,7 @@ OPT_IN_MARKERS: Final = MappingProxyType( "cli_determinism": CLI_DETERMINISM_OPT_IN_ENV, "mcp_oauth_live": MCP_OAUTH_LIVE_OPT_IN_ENV, "provider_edge_host": PROVIDER_EDGE_HOST_OPT_IN_ENV, + "owned_gateway": OWNED_GATEWAY_OPT_IN_ENV, "otel_v2": OTEL_V2_OPT_IN_ENV, "otel_tls": OTEL_TLS_OPT_IN_ENV, "secret_manager": SECRET_MANAGER_OPT_IN_ENV, @@ -172,6 +174,11 @@ def pytest_configure(config: pytest.Config) -> None: "provider_edge_host: routes provider traffic through the pytest host's edge in every fixture mode, so the " "gateway must reach the pytest host; deselected unless E2E_PROVIDER_EDGE_HOST_REACHABLE is set", ) + config.addinivalue_line( + "markers", + "owned_gateway: boots its own proxy from source against the stack's Postgres, so it needs DATABASE_URL " + "on the pytest host; deselected unless E2E_OWNED_GATEWAY is set", + ) config.addinivalue_line( "markers", "otel_v2: needs a proxy running with LITELLM_OTEL_V2=true; deselected unless E2E_OTEL_V2 is set", diff --git a/tests/e2e/coverage_registry/other.yaml b/tests/e2e/coverage_registry/other.yaml index 0b9249d7420..39747607531 100644 --- a/tests/e2e/coverage_registry/other.yaml +++ b/tests/e2e/coverage_registry/other.yaml @@ -60,6 +60,9 @@ - {id: other.auth.jwt.wrong_issuer_denied, module: other, tier: P0, area: auth, assertions: [wrong_issuer_denied], source: "auth/handle_jwt.py", rationale: "A signed token with the correct audience and an unexpected issuer is rejected"} - {id: other.auth.jwt.wrong_audience_denied, module: other, tier: P0, area: auth, assertions: [wrong_audience_denied], source: "auth/handle_jwt.py", rationale: "A signed token from the trusted issuer intended for another app is rejected"} +- {id: other.auth.jwt.auto_register_maps_existing_key, module: other, tier: P0, area: auth, assertions: [maps_existing_key], source: "user_api_key_auth.py _auto_register_jwt_mapping", rationale: "With auto_register_map_existing_key: true, the first JWT call of a user who already owns a key writes the sub-claim mapping to that existing key hash and mints nothing; the spend row lands on the pre-existing key (LIT-5378)", fail_before_fix: proven} +- {id: other.auth.jwt.auto_register_mints_when_keyless, module: other, tier: P0, area: auth, assertions: [mints_when_keyless], source: "user_api_key_auth.py _auto_register_jwt_mapping", rationale: "With auto_register_map_existing_key: true, a user with no keys still gets exactly one minted key and a sub-claim mapping on their first JWT call (LIT-5378)", fail_before_fix: proven} +- {id: other.auth.jwt.auto_register_default_mints, module: other, tier: P0, area: auth, assertions: [default_mints], source: "user_api_key_auth.py _auto_register_jwt_mapping", rationale: "Without auto_register_map_existing_key, auto_register keeps the current behavior: it mints a second key for a user who already has one and bills the minted key (LIT-5378)", fail_before_fix: proven} - {id: other.auth.session_token.valid_allows, module: other, tier: P0, area: auth, assertions: [valid_allows], source: "auth/auth_checks.py ExperimentalUIJWTToken", rationale: "An unexpired LiteLLM-minted session token authenticates with the role it carries"} - {id: other.auth.session_token.expired_denied, module: other, tier: P0, area: auth, assertions: [expired_denied], source: "auth/user_api_key_auth.py expiry check", rationale: "An expired session token is rejected with the expired-key error"} - {id: other.auth.session_token.encrypted_value_denied, module: other, tier: P0, area: auth, assertions: [encrypted_value_denied], source: "auth/auth_checks.py ExperimentalUIJWTToken", rationale: "An encrypted value read back from a management route is not accepted as a bearer token"} diff --git a/tests/e2e/e2e_config.py b/tests/e2e/e2e_config.py index 3fa9f534ffd..e88bfad8388 100644 --- a/tests/e2e/e2e_config.py +++ b/tests/e2e/e2e_config.py @@ -7,6 +7,7 @@ environment so the same tests run against localhost or a deployed proxy. from __future__ import annotations import os +import socket from dataclasses import dataclass import time import uuid @@ -16,6 +17,7 @@ from typing import Final from dotenv import load_dotenv from fixture_mode import deterministic_marker, parse_fixture_mode, registration_owner from provider_edge import provider_edge_api_base +from pydantic import TypeAdapter # Local runs keep provider / DataDog keys in tests/e2e/.env (see CONTRIBUTING.md). # Compose injects them into the proxy container, but pytest on the host does not @@ -206,6 +208,7 @@ REDIS_CHAOS_OPT_IN_ENV = "E2E_REDIS_CHAOS" CLI_DETERMINISM_OPT_IN_ENV = "E2E_CLI_DETERMINISM" MCP_OAUTH_LIVE_OPT_IN_ENV: Final = "E2E_MCP_OAUTH_LIVE" PROVIDER_EDGE_HOST_OPT_IN_ENV: Final = "E2E_PROVIDER_EDGE_HOST_REACHABLE" +OWNED_GATEWAY_OPT_IN_ENV: Final = "E2E_OWNED_GATEWAY" OTEL_V2_OPT_IN_ENV: Final = "E2E_OTEL_V2" OTEL_TLS_OPT_IN_ENV: Final = "E2E_OTEL_EXPORTER_ENDPOINT" SECRET_MANAGER_OPT_IN_ENV: Final = "E2E_SECRET_MANAGER" @@ -296,6 +299,15 @@ def unique_marker() -> str: return uuid.uuid4().hex[:12] +INHERITED_ENV_PREFIXES: Final = ("REDIS_", "MICROSOFT_", "GOOGLE_", "GENERIC_", "PROXY_") + + +def available_port() -> int: + with socket.socket() as listener: + listener.bind(("127.0.0.1", 0)) + return TypeAdapter(tuple[str, int]).validate_python(listener.getsockname())[1] + + def settle_propagation(written_at: float) -> None: """Block until PROPAGATION_TIMEOUT has elapsed since `written_at`, a `time.monotonic()` stamp taken the moment a control-plane write returned. diff --git a/tests/e2e/mcp/oauth_gateway.py b/tests/e2e/mcp/oauth_gateway.py index b328c81687b..029b0135900 100644 --- a/tests/e2e/mcp/oauth_gateway.py +++ b/tests/e2e/mcp/oauth_gateway.py @@ -8,7 +8,6 @@ The optional live edge measures headers without recording credentials or bodies. from __future__ import annotations import os -import socket import subprocess import sys import threading @@ -20,13 +19,12 @@ from pathlib import Path from typing import Final import psycopg +from e2e_config import INHERITED_ENV_PREFIXES, available_port from e2e_http import NoBody from idp import Keycloak, stop_process_group from proxy_client import ProxyClient, build_proxy_client from psycopg.rows import class_row -from pydantic import BaseModel, SecretStr, TypeAdapter, ValidationError - -INHERITED_ENV_PREFIXES: Final = ("REDIS_", "MICROSOFT_", "GOOGLE_", "GENERIC_", "PROXY_") +from pydantic import BaseModel, SecretStr, ValidationError class StoredOAuth(BaseModel): @@ -101,12 +99,6 @@ class OAuthObservation: assert all(not item[2] for item in snapshot), "gateway bearer leaked to the upstream" -def available_port() -> int: - with socket.socket() as listener: - listener.bind(("127.0.0.1", 0)) - return TypeAdapter(tuple[str, int]).validate_python(listener.getsockname())[1] - - @dataclass(slots=True) class OAuthGateway: base_url: str diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 83d8a9884a3..e027c410e44 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -1537,6 +1537,7 @@ class UserNewBody(BaseModel): class UserNewResponse(BaseModel): user_id: str + key: str | None = None class UserUpdateBody(BaseModel): @@ -1580,6 +1581,40 @@ class UserListResponse(BaseModel): total: int +class UserKeyRow(BaseModel): + token: str + key_alias: str | None = None + + +class UserInfoWithKeysResponse(BaseModel): + user_id: str | None = None + keys: list[UserKeyRow] = [] + + +class JwtKeyMappingRow(BaseModel): + id: str + jwt_claim_name: str + jwt_claim_value: str + created_by: str | None = None + + +class JwtKeyMappingListParams(BaseModel): + size: int = 100 + + +class JwtKeyMappingListResponse(BaseModel): + mappings: list[JwtKeyMappingRow] + total_count: int + + +class JwtKeyMappingDeleteBody(BaseModel): + id: str + + +class JwtKeyMappingDeleteResponse(BaseModel): + status: str + + class OrgNewBody(BaseModel): organization_alias: str models: list[str] = [] diff --git a/tests/e2e/other/other_client.py b/tests/e2e/other/other_client.py index 93c198586f6..d7bad4f1ed1 100644 --- a/tests/e2e/other/other_client.py +++ b/tests/e2e/other/other_client.py @@ -21,12 +21,20 @@ from idp import Keycloak, keycloak_from_env from models import ( ChatBody, ChatResponse, + JwtKeyMappingDeleteBody, + JwtKeyMappingDeleteResponse, + JwtKeyMappingListParams, + JwtKeyMappingListResponse, ModelsListParams, ModelsListResponse, ReadinessDetailsResponse, ReadinessResponse, + UserInfoParams, + UserInfoWithKeysResponse, UserListParams, UserListResponse, + UserNewBody, + UserNewResponse, ) from proxy_client import ProxyClient from pydantic import Field @@ -79,6 +87,44 @@ class OtherClient: response_type=ReadinessDetailsResponse, ) + def user_new(self, body: UserNewBody) -> Result[UserNewResponse]: + """POST /user/new under the master key: seed the litellm user a JWT + `sub` claim resolves to, before that token ever reaches the proxy.""" + return self.proxy.transport.post( + "/user/new", + headers=self.proxy.transport.master, + json=body, + response_type=UserNewResponse, + ) + + def user_info(self, user_id: str) -> Result[UserInfoWithKeysResponse]: + """GET /user/info under the master key. Only the user's key rows are + modelled: `token` is the stored key hash, never the plaintext key.""" + return self.proxy.transport.get( + "/user/info", + headers=self.proxy.transport.master, + params=UserInfoParams(user_id=user_id), + response_type=UserInfoWithKeysResponse, + ) + + def jwt_mapping_list(self) -> Result[JwtKeyMappingListResponse]: + """GET /jwt/key/mapping/list under the master key.""" + return self.proxy.transport.get( + "/jwt/key/mapping/list", + headers=self.proxy.transport.master, + params=JwtKeyMappingListParams(size=100), + response_type=JwtKeyMappingListResponse, + ) + + def jwt_mapping_delete(self, mapping_id: str) -> Result[JwtKeyMappingDeleteResponse]: + """POST /jwt/key/mapping/delete under the master key.""" + return self.proxy.transport.post( + "/jwt/key/mapping/delete", + headers=self.proxy.transport.master, + json=JwtKeyMappingDeleteBody(id=mapping_id), + response_type=JwtKeyMappingDeleteResponse, + ) + def chat_as_team(self, token: str, team: str, body: ChatBody) -> Result[ChatResponse]: """POST /chat/completions under `token` with `x-litellm-team-id: team`.""" return self.proxy.transport.post( diff --git a/tests/e2e/other/owned_jwt_gateway.py b/tests/e2e/other/owned_jwt_gateway.py new file mode 100644 index 00000000000..1af348cac60 --- /dev/null +++ b/tests/e2e/other/owned_jwt_gateway.py @@ -0,0 +1,108 @@ +"""An owned, source-built proxy whose `litellm_jwtauth` block a test controls. + +The shared proxy on :4000 runs the CONTRIBUTING.md JWT block, so a test that +needs a different `litellm_jwtauth` config boots its own gateway on a free port +against the same database and the same Keycloak realm. The caller supplies the +`litellm_jwtauth` mapping verbatim, which is exactly what makes a config an +unfixed proxy rejects observable as a boot failure in this gateway's own log. +""" + +from __future__ import annotations + +import os +import subprocess +import sys +import time +from collections.abc import Mapping +from contextlib import ExitStack +from dataclasses import dataclass, field +from pathlib import Path +from typing import Final + +from e2e_config import INHERITED_ENV_PREFIXES, available_port +from e2e_http import NoBody +from idp import Keycloak, stop_process_group +from proxy_client import ProxyClient, build_proxy_client + +MODEL_NAME: Final = "gemini-3.8-flash" + + +@dataclass(slots=True) +class OwnedJwtGateway: + base_url: str + proxy: ProxyClient + _environment: Mapping[str, str] = field(repr=False) + _command: tuple[str, ...] = field(repr=False) + _log_path: Path + _child: subprocess.Popen[bytes] | None = field(default=None, init=False, repr=False) + + def start(self) -> None: + with self._log_path.open("ab") as log: + self._child = subprocess.Popen( + self._command, + env=self._environment, + stdout=log, + stderr=log, + start_new_session=True, + ) + deadline: Final = time.monotonic() + 120 + while time.monotonic() < deadline: + assert self._child.poll() is None, "owned JWT gateway exited; inspect its private log" + result = self.proxy.transport.probe("/health/liveliness", params=NoBody()) + if result.status_code == 200: + return + time.sleep(0.5) + raise AssertionError("owned JWT gateway did not become ready") + + def stop(self) -> None: + if self._child is not None: + stop_process_group(self._child) + assert self._child.poll() is not None, "old gateway process is still alive" + + +def owned_jwt_gateway( + idp: Keycloak, directory: Path, cleanup: ExitStack, *, litellm_jwtauth: str, name: str +) -> OwnedJwtGateway: + for env_name in ("DATABASE_URL", "LITELLM_LICENSE", "LITELLM_MASTER_KEY"): + assert os.environ.get(env_name), f"{env_name} is required for the owned JWT gateway" + port: Final = available_port() + base_url: Final = f"http://127.0.0.1:{port}" + config: Final = directory / f"{name}.yaml" + config.write_text( + "model_list:\n" + f" - model_name: {MODEL_NAME}\n" + " litellm_params:\n" + f" model: gemini/{MODEL_NAME}\n" + " api_key: os.environ/GEMINI_API_KEY\n" + "general_settings:\n" + " master_key: os.environ/LITELLM_MASTER_KEY\n" + " database_url: os.environ/DATABASE_URL\n" + " proxy_batch_write_at: 5\n" + " enable_jwt_auth: true\n" + " litellm_jwtauth:\n" + "".join(f" {line}\n" for line in litellm_jwtauth.strip().splitlines()) + ) + environment: Final = { + **{key: value for key, value in os.environ.items() if not key.startswith(INHERITED_ENV_PREFIXES)}, + "JWT_PUBLIC_KEY_URL": idp.jwks_url, + "JWT_ISSUER": idp.issuer, + "JWT_AUDIENCE": "litellm-e2e", + "LITELLM_DANGEROUSLY_PERMIT_WEAK_OR_UNSET_MASTER_KEY": "true", + "DISABLE_SCHEMA_UPDATE": "true", + "STORE_MODEL_IN_DB": "True", + "PYTHONPATH": str(Path(__file__).resolve().parents[3]), + } + gateway: Final = OwnedJwtGateway( + base_url=base_url, + proxy=build_proxy_client( + base_url=base_url, + control_plane_base_url=base_url, + replica_urls=(base_url,), + master_key=os.environ["LITELLM_MASTER_KEY"], + ), + _environment=environment, + _command=(sys.executable, "-m", "litellm.proxy.proxy_cli", "--config", str(config), "--port", str(port)), + _log_path=directory / f"{name}.log", + ) + cleanup.callback(gateway.stop) + gateway.start() + return gateway diff --git a/tests/e2e/other/test_jwt_auto_register_e2e.py b/tests/e2e/other/test_jwt_auto_register_e2e.py new file mode 100644 index 00000000000..8f7c2c6a693 --- /dev/null +++ b/tests/e2e/other/test_jwt_auto_register_e2e.py @@ -0,0 +1,182 @@ +"""auto_register with auto_register_map_existing_key binds the JWT claim to the user's existing key. + +`unregistered_jwt_client_behavior: auto_register` on `virtual_key_claim_field: sub` mints a fresh +virtual key on the user's first JWT call. With `auto_register_map_existing_key: true` the proxy must +instead point the new JWT mapping at a key the resolved user already owns, and mint only when the +user has none. Each behavior gets its own gateway because the flag lives in `litellm_jwtauth`, so +this file boots two owned proxies against the shared database and Keycloak realm. +""" + +from __future__ import annotations + +import hashlib +from collections.abc import Iterator +from contextlib import ExitStack +from typing import Final + +import pytest +from e2e_config import unique_marker +from e2e_http import unwrap +from idp import Identity, Keycloak +from lifecycle import ResourceManager +from models import ChatBody, ChatMessage, JwtKeyMappingRow, KeyGenerateBody, TeamNewBody, UserNewBody +from other_client import OtherClient +from owned_jwt_gateway import MODEL_NAME, OwnedJwtGateway, owned_jwt_gateway + +pytestmark = pytest.mark.e2e + +_JWT_COMMON: Final = ( + "user_id_jwt_field: sub\n" + "user_email_jwt_field: email\n" + "team_ids_jwt_field: groups\n" + "user_id_upsert: true\n" + "virtual_key_claim_field: sub\n" + "unregistered_jwt_client_behavior: auto_register" +) + + +def _key_hash(key: str) -> str: + return hashlib.sha256(key.encode()).hexdigest() + + +def _ping() -> ChatBody: + return ChatBody( + model=MODEL_NAME, + messages=[ChatMessage(role="user", content=f"Reply with the single word ok. {unique_marker()}")], + max_tokens=5, + ) + + +def _identity_with_user(idp: Keycloak, client: OtherClient, resources: ResourceManager) -> Identity: + """An IdP identity plus the litellm user and team its claims resolve to, with + teardown that also sweeps the user's keys and JWT mapping rows the proxy + wrote, since those outlive the user row itself.""" + marker: Final = unique_marker() + identity: Final = idp.provision(marker=marker, group=f"e2e-jwt-team-{marker}", defer=resources.defer) + resources.defer(lambda: client.proxy.delete_user(identity.user_id)) + team_id: Final = client.proxy.create_team(TeamNewBody(team_alias=f"e2e-jwt-{marker}", team_id=identity.group)) + resources.defer(lambda: client.proxy.delete_team(team_id)) + unwrap( + client.user_new( + UserNewBody( + user_id=identity.user_id, + user_email=f"{identity.username}@example.com", + user_role="internal_user", + auto_create_key=False, + ) + ) + ) + + def delete_user_keys() -> None: + for row in unwrap(client.user_info(identity.user_id)).keys: + client.proxy.delete_key(row.token) + + def delete_user_mappings() -> None: + for mapping in unwrap(client.jwt_mapping_list()).mappings: + if mapping.jwt_claim_value == identity.user_id: + _ = client.jwt_mapping_delete(mapping.id) + + resources.defer(delete_user_keys) + resources.defer(delete_user_mappings) + return identity + + +def _mapping_for(client: OtherClient, claim_value: str) -> JwtKeyMappingRow | None: + return next( + (row for row in unwrap(client.jwt_mapping_list()).mappings if row.jwt_claim_value == claim_value), + None, + ) + + +@pytest.fixture(scope="module") +def mapping_gateway(idp: Keycloak, tmp_path_factory: pytest.TempPathFactory) -> Iterator[OwnedJwtGateway]: + with ExitStack() as cleanup: + yield owned_jwt_gateway( + idp, + tmp_path_factory.mktemp("jwt-mapping"), + cleanup, + litellm_jwtauth=f"{_JWT_COMMON}\nauto_register_map_existing_key: true", + name="jwt-mapping-gateway", + ) + + +@pytest.fixture(scope="module") +def minting_gateway(idp: Keycloak, tmp_path_factory: pytest.TempPathFactory) -> Iterator[OwnedJwtGateway]: + with ExitStack() as cleanup: + yield owned_jwt_gateway( + idp, + tmp_path_factory.mktemp("jwt-minting"), + cleanup, + litellm_jwtauth=_JWT_COMMON, + name="jwt-minting-gateway", + ) + + +@pytest.mark.owned_gateway +class TestJwtAutoRegisterMapExistingKey: + @pytest.mark.covers("other.auth.jwt.auto_register_maps_existing_key") + def test_first_jwt_call_maps_to_the_users_existing_key_and_mints_none( + self, client: OtherClient, idp: Keycloak, resources: ResourceManager, mapping_gateway: OwnedJwtGateway + ) -> None: + identity: Final = _identity_with_user(idp, client, resources) + existing_key: Final = client.proxy.generate_key( + KeyGenerateBody( + user_id=identity.user_id, team_id=identity.group, key_alias=f"e2e-jwt-existing-{unique_marker()}" + ) + ) + resources.defer(lambda: client.proxy.delete_key(existing_key)) + + response: Final = unwrap(mapping_gateway.proxy.chat(idp.access_token(identity), _ping())) + + keys: Final = unwrap(client.user_info(identity.user_id)).keys + assert [row.token for row in keys] == [_key_hash(existing_key)], ( + f"map_existing_key must leave the user with only their pre-existing key, got {keys}" + ) + mapping: Final = _mapping_for(client, identity.user_id) + assert mapping is not None, ( + f"no JWT mapping row for sub={identity.user_id}: {unwrap(client.jwt_mapping_list())}" + ) + assert mapping.jwt_claim_name == "sub", f"mapping must bind the sub claim, got {mapping}" + assert mapping.created_by == "auto_register", f"mapping must be written by auto_register, got {mapping}" + rows: Final = client.proxy.poll_logs_for_key(existing_key) + assert any(row.request_id == response.id for row in rows), ( + f"the JWT chat must be billed to the user's existing key, spend rows for it: {rows}" + ) + + @pytest.mark.covers("other.auth.jwt.auto_register_mints_when_keyless") + def test_first_jwt_call_mints_a_key_when_the_user_has_none( + self, client: OtherClient, idp: Keycloak, resources: ResourceManager, mapping_gateway: OwnedJwtGateway + ) -> None: + identity: Final = _identity_with_user(idp, client, resources) + + response: Final = unwrap(mapping_gateway.proxy.chat(idp.access_token(identity), _ping())) + assert response.choices, f"JWT chat returned no completion: {response}" + + keys: Final = unwrap(client.user_info(identity.user_id)).keys + assert len(keys) == 1, f"a keyless user must get exactly one minted key, got {keys}" + mapping: Final = _mapping_for(client, identity.user_id) + assert mapping is not None and mapping.jwt_claim_name == "sub", ( + f"the minted key must be recorded as a sub-claim mapping, mappings: {unwrap(client.jwt_mapping_list())}" + ) + + @pytest.mark.covers("other.auth.jwt.auto_register_default_mints") + def test_default_behavior_still_mints_when_the_user_already_has_a_key( + self, client: OtherClient, idp: Keycloak, resources: ResourceManager, minting_gateway: OwnedJwtGateway + ) -> None: + identity: Final = _identity_with_user(idp, client, resources) + existing_key: Final = client.proxy.generate_key( + KeyGenerateBody(user_id=identity.user_id, key_alias=f"e2e-jwt-existing-{unique_marker()}") + ) + resources.defer(lambda: client.proxy.delete_key(existing_key)) + + response: Final = unwrap(minting_gateway.proxy.chat(idp.access_token(identity), _ping())) + assert response.id is not None, f"JWT chat returned no response id: {response}" + + keys: Final = unwrap(client.user_info(identity.user_id)).keys + assert len(keys) == 2, ( + f"default auto_register must mint a second key for a user who already has one, got {keys}" + ) + rows: Final = client.proxy.poll_logs_for_request_id(response.id) + assert rows and all(row.api_key != _key_hash(existing_key) for row in rows), ( + f"the default path must bill the minted key, not the user's existing one: {rows}" + ) diff --git a/tests/e2e/pytest.ini b/tests/e2e/pytest.ini index e795ebe5721..dbd3ff47daa 100644 --- a/tests/e2e/pytest.ini +++ b/tests/e2e/pytest.ini @@ -16,6 +16,7 @@ markers = quiet_stack: measures the proxy itself, so it runs while no other test on this host is hitting the stack; every other test waits for it to finish mcp_oauth_live: real Linear OAuth consent via a captured browser session; deselected unless E2E_MCP_OAUTH_LIVE is set provider_edge_host: routes provider traffic through the pytest host's edge in every fixture mode, so the gateway must reach the pytest host; deselected unless E2E_PROVIDER_EDGE_HOST_REACHABLE is set + owned_gateway: boots its own proxy from source against the stack's Postgres, so it needs DATABASE_URL on the pytest host; deselected unless E2E_OWNED_GATEWAY is set otel_v2: needs a proxy running with LITELLM_OTEL_V2=true; deselected unless E2E_OTEL_V2 is set otel_tls: needs a stack whose gateway exports OTLP over TLS signed by the CA in SSL_CERT_FILE; deselected unless E2E_OTEL_EXPORTER_ENDPOINT is set secret_manager: needs a proxy booted from gateway/secret_manager__ci_config.yml against that live secret manager; deselected unless E2E_SECRET_MANAGER names the backend (see secret_manager/secret_backends.py) diff --git a/tests/integration/_support/database_relay.py b/tests/integration/_support/database_relay.py index 1b3bc0183a5..cb8b9098a64 100644 --- a/tests/integration/_support/database_relay.py +++ b/tests/integration/_support/database_relay.py @@ -88,15 +88,89 @@ class DatabaseRelay: ) +class HeldStatementRelay: + def __init__(self, upstream_host: str, upstream_port: int, trigger: bytes) -> None: + self.port: Final = _free_port() + self._upstream_host: Final = upstream_host + self._upstream_port: Final = upstream_port + self._trigger: Final = trigger + self._loop: Final = asyncio.new_event_loop() + self._released: Final = asyncio.Event() + self.held: Final = threading.Event() + self._ready: Final = threading.Event() + self._thread: Final = threading.Thread(target=self._run, daemon=True) + + def release(self) -> None: + self._loop.call_soon_threadsafe(self._released.set) + + def start(self) -> None: + self._thread.start() + assert self._ready.wait(10), "Database relay did not start" + + def stop(self) -> None: + self.release() + self._loop.call_soon_threadsafe(self._loop.stop) + self._thread.join(10) + + def _run(self) -> None: + asyncio.set_event_loop(self._loop) + self._loop.run_until_complete(asyncio.start_server(self._serve, "127.0.0.1", self.port)) + self._ready.set() + self._loop.run_forever() + + def _holds(self, window: bytes) -> bool: + return not self.held.is_set() and self._trigger in window + + async def _serve(self, client_reader: asyncio.StreamReader, client_writer: asyncio.StreamWriter) -> None: + server_reader, server_writer = await asyncio.open_connection(self._upstream_host, self._upstream_port) + + async def forward(reader: asyncio.StreamReader, writer: asyncio.StreamWriter, inspect: bool) -> None: + tail = b"" # rebind-ok: carries the previous read's end so a trigger split across reads still matches + try: + while chunk := await reader.read(65536): + window: Final = tail + chunk + if inspect and self._holds(window): + self.held.set() + await self._released.wait() + tail = window[-(len(self._trigger) - 1) :] + writer.write(chunk) + await writer.drain() + except (ConnectionError, asyncio.IncompleteReadError): + return + finally: + writer.close() + + await asyncio.gather( + forward(client_reader, server_writer, True), + forward(server_reader, client_writer, False), + ) + + +def _relayed_url(database_url: str, port: int) -> str: + parts: Final = urlsplit(database_url) + credentials: Final = f"{parts.username}:{parts.password}@" if parts.username else "" + return urlunsplit(parts._replace(netloc=f"{credentials}127.0.0.1:{port}")) + + @contextmanager def database_relay(database_url: str, trigger: bytes) -> Generator[tuple[DatabaseRelay, str]]: parts: Final = urlsplit(database_url) assert parts.hostname is not None and parts.port is not None, database_url relay: Final = DatabaseRelay(parts.hostname, parts.port, trigger) relay.start() - credentials: Final = f"{parts.username}:{parts.password}@" if parts.username else "" - relayed: Final = urlunsplit(parts._replace(netloc=f"{credentials}127.0.0.1:{relay.port}")) try: - yield relay, relayed + yield relay, _relayed_url(database_url, relay.port) + finally: + relay.stop() + + +@contextmanager +def held_statement_relay(database_url: str, trigger: bytes) -> Generator[tuple[HeldStatementRelay, str]]: + parts: Final = urlsplit(database_url) + assert parts.hostname is not None and parts.port is not None, database_url + relay: Final = HeldStatementRelay(parts.hostname, parts.port, trigger) + relay.start() + try: + yield relay, _relayed_url(database_url, relay.port) finally: relay.stop() diff --git a/tests/integration/authorization/test_jwt_auto_register_map_existing_key.py b/tests/integration/authorization/test_jwt_auto_register_map_existing_key.py new file mode 100644 index 00000000000..5215054a364 --- /dev/null +++ b/tests/integration/authorization/test_jwt_auto_register_map_existing_key.py @@ -0,0 +1,219 @@ +import json +import os +import time +import uuid +from collections.abc import Iterator +from concurrent.futures import ThreadPoolExecutor +from contextlib import contextmanager +from hashlib import sha256 +from pathlib import Path +from typing import Final + +import httpx +import jwt +import pytest +import yaml +from cryptography.hazmat.primitives.asymmetric import rsa + +from tests.integration._support.client import Gateway, eventually, string_value +from tests.integration._support.database import read_rows +from tests.integration._support.database_relay import held_statement_relay +from tests.integration._support.process import owned_proxy +from tests.integration._support.wire import Reply, Request, wire_server + +KEY_ID: Final = "integration-jwt-map-existing-key" +MAPPING_INSERT: Final = b'INSERT INTO "public"."LiteLLM_JWTKeyMapping"' + +pytestmark = pytest.mark.timeout(240) + + +def _hash(key: str) -> str: + return sha256(key.encode()).hexdigest() + + +def _config(directory: Path, claim_field: str) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["general_settings"] = { + **config["general_settings"], + "enable_jwt_auth": True, + "litellm_jwtauth": { + "user_id_jwt_field": "sub", + "user_email_jwt_field": "email", + "virtual_key_claim_field": claim_field, + "unregistered_jwt_client_behavior": "auto_register", + "auto_register_map_existing_key": True, + }, + } + path: Final = directory / f"jwt_map_existing_key_{claim_field}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +@contextmanager +def _issuer() -> Iterator[tuple[rsa.RSAPrivateKey, str]]: + private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + public_jwk: Final = json.loads(jwt.algorithms.RSAAlgorithm.to_jwk(private_key.public_key())) + jwks: Final = json.dumps({"keys": [{**public_jwk, "kid": KEY_ID, "use": "sig", "alg": "RS256"}]}).encode() + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return Reply(body=jwks) + + with wire_server(respond) as server: + yield private_key, server.url + + +def _token(private_key: rsa.RSAPrivateKey, subject: str, **claims: str) -> str: + now: Final = int(time.time()) + return jwt.encode( + {"sub": subject, **claims, "iat": now, "exp": now + 300}, + private_key, + algorithm="RS256", + headers={"kid": KEY_ID}, + ) + + +def _chat(candidate: Gateway, model: str, token: str) -> httpx.Response: + return candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "map existing key control"}]}, + key=token, + ) + + +def _mapped_token(claim_name: str, claim_value: str) -> str: + rows: Final = read_rows( + 'SELECT token FROM "LiteLLM_JWTKeyMapping" WHERE jwt_claim_name = %s AND jwt_claim_value = %s', + (claim_name, claim_value), + ) + assert len(rows) == 1, rows + return string_value(rows[0]["token"]) + + +def _user_key_hashes(user: str) -> frozenset[str]: + rows: Final = read_rows('SELECT token FROM "LiteLLM_VerificationToken" WHERE user_id = %s', (user,)) + return frozenset(string_value(row["token"]) for row in rows) + + +def _billed_key(response: httpx.Response) -> str: + rows: Final = eventually( + lambda: read_rows( + 'SELECT api_key FROM "LiteLLM_SpendLogs" WHERE request_id = %s', (str(response.json()["id"]),) + ), + lambda values: len(values) == 1, + seconds=70, + ) + return string_value(rows[0]["api_key"]) + + +def test_first_jwt_call_reuses_the_newest_durable_llm_key_and_skips_every_ineligible_newer_key( + gateway: Gateway, tmp_path: Path +) -> None: + with _issuer() as (private_key, jwks_url), gateway.scenario() as scenario: + model: Final = scenario.model() + user: Final = scenario.user(user_role="internal_user") + older_durable: Final = scenario.key(user_id=user) + durable: Final = scenario.key(user_id=user) + skipped: Final = { + "older_durable": older_durable, + "expiring": scenario.key(user_id=user, duration="1h"), + "management_only": scenario.key(user_id=user, allowed_routes=["management_routes"]), + "auto_registered_look_alike": scenario.key(user_id=user, metadata={"auto_registered": True}), + "other_team": scenario.key(user_id=user, team_id=scenario.team()), + "blocked": scenario.key(user_id=user), + } + gateway.post("/key/block", {"key": skipped["blocked"]}) + keys_before: Final = _user_key_hashes(user) + + with owned_proxy( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks_url}, config=_config(tmp_path, "sub") + ) as candidate: + response: Final = _chat(candidate, model, _token(private_key, user)) + + assert response.status_code == 200, response.text + mapped: Final = _mapped_token("sub", user) + assert mapped == _hash(durable), { + "mapped_to": next((name for name, key in skipped.items() if _hash(key) == mapped), mapped) + } + assert _user_key_hashes(user) == keys_before, "a key was minted although a reusable one existed" + assert _billed_key(response) == _hash(durable) + + +def test_user_matched_by_email_instead_of_sub_still_reuses_their_existing_key(gateway: Gateway, tmp_path: Path) -> None: + with _issuer() as (private_key, jwks_url), gateway.scenario() as scenario: + model: Final = scenario.model() + email: Final = f"integration-{uuid.uuid4().hex}@example.com" + user: Final = scenario.user(user_role="internal_user", user_email=email) + existing: Final = scenario.key(user_id=user) + subject: Final = f"integration-idp-subject-{uuid.uuid4().hex}" + + with owned_proxy( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks_url}, config=_config(tmp_path, "sub") + ) as candidate: + response: Final = _chat(candidate, model, _token(private_key, subject, email=email.upper())) + + assert response.status_code == 200, response.text + assert _mapped_token("sub", subject) == _hash(existing) + assert _user_key_hashes(user) == frozenset({_hash(existing)}), "a key was minted for an email-matched user" + assert _billed_key(response) == _hash(existing) + + +def test_shared_client_claim_never_maps_a_second_user_onto_the_first_users_personal_key( + gateway: Gateway, tmp_path: Path +) -> None: + with _issuer() as (private_key, jwks_url), gateway.scenario() as scenario: + model: Final = scenario.model() + first_user: Final = scenario.user(user_role="internal_user") + second_user: Final = scenario.user(user_role="internal_user") + personal: Final = scenario.key(user_id=first_user) + client_id: Final = f"integration-shared-client-{uuid.uuid4().hex}" + + with owned_proxy( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks_url}, config=_config(tmp_path, "client_id") + ) as candidate: + first: Final = _chat(candidate, model, _token(private_key, first_user, client_id=client_id)) + second: Final = _chat(candidate, model, _token(private_key, second_user, client_id=client_id)) + + assert first.status_code == 200, first.text + assert second.status_code == 200, second.text + mapped: Final = _mapped_token("client_id", client_id) + assert mapped != _hash(personal), "the shared client claim was mapped to the first user's personal key" + assert (_billed_key(first), _billed_key(second)) == (mapped, mapped) + assert read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE api_key = %s', (_hash(personal),)) == [] + + +def test_concurrent_first_jwt_calls_of_a_keyless_user_both_succeed_on_one_surviving_mapped_key( + gateway: Gateway, tmp_path: Path +) -> None: + writer_url: Final = os.environ.get("INTEGRATION_PROXY_DATABASE_URL") or os.environ["DATABASE_URL"] + with ( + _issuer() as (private_key, jwks_url), + gateway.scenario() as scenario, + held_statement_relay(writer_url, MAPPING_INSERT) as (relay, relayed_url), + ): + model: Final = scenario.model() + user: Final = scenario.user(user_role="internal_user") + token: Final = _token(private_key, user) + overrides: Final = { + "JWT_PUBLIC_KEY_URL": jwks_url, + "DATABASE_URL": relayed_url, + "PRISMA_HEALTH_WATCHDOG_ENABLED": "false", + } + + with ( + owned_proxy(gateway, tmp_path, overrides, config=_config(tmp_path, "sub")) as candidate, + ThreadPoolExecutor(max_workers=1) as pool, + ): + held_call: Final = pool.submit(_chat, candidate, model, token) + assert relay.held.wait(60), "the first call never reached its mapping insert" + racing: Final = _chat(candidate, model, token) + relay.release() + held: Final = held_call.result(timeout=60) + + assert racing.status_code == 200, racing.text + assert held.status_code == 200, held.text + keys: Final = _user_key_hashes(user) + assert len(keys) == 1, keys + assert _mapped_token("sub", user) in keys + assert (_billed_key(held), _billed_key(racing)) == (_mapped_token("sub", user),) * 2 diff --git a/tests/integration/observability/test_otel_excluded_services.py b/tests/integration/observability/test_otel_excluded_services.py index 48b20651b97..8768052b488 100644 --- a/tests/integration/observability/test_otel_excluded_services.py +++ b/tests/integration/observability/test_otel_excluded_services.py @@ -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, diff --git a/tests/integration/observability/test_otel_excluded_services_matrix.py b/tests/integration/observability/test_otel_excluded_services_matrix.py index 0d4b5d087c6..7bca7ffc0dc 100644 --- a/tests/integration/observability/test_otel_excluded_services_matrix.py +++ b/tests/integration/observability/test_otel_excluded_services_matrix.py @@ -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] diff --git a/tests/integration/spend/test_passthrough_request_tags.py b/tests/integration/spend/test_passthrough_request_tags.py new file mode 100644 index 00000000000..e6f5e15e161 --- /dev/null +++ b/tests/integration/spend/test_passthrough_request_tags.py @@ -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 diff --git a/tests/pass_through_unit_tests/test_pass_through_unit_tests.py b/tests/pass_through_unit_tests/test_pass_through_unit_tests.py index 6c57e59f7e3..7fb23223845 100644 --- a/tests/pass_through_unit_tests/test_pass_through_unit_tests.py +++ b/tests/pass_through_unit_tests/test_pass_through_unit_tests.py @@ -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"}, } diff --git a/tests/proxy_behavior/spend/test_autorouter_session_rollup.py b/tests/proxy_behavior/spend/test_autorouter_session_rollup.py index 2c648f309f6..a5c6f5962a2 100644 --- a/tests/proxy_behavior/spend/test_autorouter_session_rollup.py +++ b/tests/proxy_behavior/spend/test_autorouter_session_rollup.py @@ -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 diff --git a/tests/proxy_behavior/spend/test_baseline_accounting.py b/tests/proxy_behavior/spend/test_baseline_accounting.py index 3504751d132..8fb82d0c80e 100644 --- a/tests/proxy_behavior/spend/test_baseline_accounting.py +++ b/tests/proxy_behavior/spend/test_baseline_accounting.py @@ -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 diff --git a/tests/test_litellm/tracing/test_decode.py b/tests/test_litellm/tracing/test_decode.py index 21a79dd6b87..dda075863ac 100644 --- a/tests/test_litellm/tracing/test_decode.py +++ b/tests/test_litellm/tracing/test_decode.py @@ -55,13 +55,94 @@ def _kv(key: str, value: str | int) -> KeyValue: return KeyValue(key=key, value=AnyValue(string_value=value)) -def _export(*spans: Span, service: str = "svc", scope: str = "test") -> bytes: +def _export(*spans: Span, service: str = "svc", scope: str = "test", agent_name: str = "") -> bytes: resource_spans = ResourceSpans(scope_spans=[ScopeSpans(spans=list(spans))]) resource_spans.resource.attributes.append(_kv("service.name", service)) + if agent_name: + resource_spans.resource.attributes.append(_kv("gen_ai.agent.name", agent_name)) resource_spans.scope_spans[0].scope.name = scope return ExportTraceServiceRequest(resource_spans=[resource_spans]).SerializeToString() +@pytest.mark.parametrize( + ("name", "attributes"), + [ + ("research_agent", {"openinference.span.kind": "AGENT", "metadata": '{"lc_agent_name":"research_agent"}'}), + ("research_agent", {"openinference.span.kind": "AGENT", "metadata": '{"ls_integration":"langgraph"}'}), + ("research_agent._execute_core", {"openinference.span.kind": "AGENT", "graph.node.id": "research_agent"}), + ("agent", {"openinference.span.kind": "AGENT", "gen_ai.agent.name": "research_agent"}), + ("openclaw.harness.run", {"openclaw.agent": "research_agent"}), + ( + "invoke_agent research_agent", + {"gen_ai.operation.name": "invoke_agent", "gen_ai.agent.name": "research_agent"}, + ), + ], + ids=["deepagents", "langgraph", "crewai", "hermes", "openclaw", "genai"], +) +def test_framework_agent_identity_is_independent_of_service(name: str, attributes: dict[str, str]): + span = _span(name, b"\x02" * 8, **attributes) + row = decode_otlp(_export(span, service="shared-deployment"), "application/x-protobuf")[0] + assert row["AgentName"] == "research_agent" + assert row["ServiceName"] == "shared-deployment" + assert row["SpanName"] == name + + +@pytest.mark.parametrize("name", ["ClaudeAgentSDK.query", "FunctionAgent.run"]) +def test_resource_agent_name_labels_instrumentors_without_an_agent_attribute(name: str): + span = _span(name, b"\x02" * 8, openinference__span__kind="AGENT") + row = decode_otlp(_export(span, agent_name="research_agent"), "application/x-protobuf")[0] + assert row["AgentName"] == "research_agent" + + +def test_span_agent_name_takes_precedence_over_resource_default(): + span = _span("invoke_agent child", b"\x02" * 8, gen_ai__agent__name="child") + row = decode_otlp(_export(span, agent_name="research_agent"), "application/x-protobuf")[0] + assert row["AgentName"] == "child" + + +@pytest.mark.parametrize( + ("scope", "span_name", "configured_name", "expected"), + [ + ("hermes-otel-plugin", "hermes-agent", "research_agent", "research_agent"), + ("hermes-otel-plugin", "child", "research_agent", "child"), + ("hermes-otel-plugin", "hermes-agent", "", "hermes-agent"), + ("other-plugin", "hermes-agent", "research_agent", "hermes-agent"), + ], +) +def test_hermes_resource_name_replaces_only_its_plugin_default( + scope: str, span_name: str, configured_name: str, expected: str +): + span = _span("agent", b"\x02" * 8, gen_ai__agent__name=span_name) + row = decode_otlp(_export(span, scope=scope, agent_name=configured_name), "application/x-protobuf")[0] + assert row["AgentName"] == expected + + +@pytest.mark.parametrize("agent_name", ["research_agent", ""]) +def test_openinference_middleware_is_not_a_separate_agent(agent_name: str): + span = _span( + "PatchToolCallsMiddleware.before_agent", b"\x02" * 8, b"\x01" * 8, + openinference__span__kind="AGENT", metadata=json.dumps({"lc_agent_name": agent_name}), + ) + row = decode_otlp(_export(span, scope="openinference.instrumentation.langchain"), "application/x-protobuf")[0] + assert (row["ObservationType"], row["AgentName"]) == ("framework", agent_name) + + +@pytest.mark.parametrize("scope", ["test", "openinference.instrumentation.langchain"]) +@pytest.mark.parametrize("kind", ["CHAIN", "AGENT"]) +@pytest.mark.parametrize("metadata", ["not json", "[]", '{"lc_agent_name":null}', "{}"]) +def test_unnamed_framework_does_not_invent_an_agent_from_service(metadata: str, scope: str, kind: str): + span = _span("workflow", b"\x02" * 8, openinference__span__kind=kind, metadata=metadata) + row = decode_otlp(_export(span, scope=scope), "application/x-protobuf")[0] + assert row["AgentName"] == "" + + +@pytest.mark.parametrize("name,expected", [("support", "support"), ("LangGraph", "")]) +def test_langgraph_distinguishes_configured_graph_name_from_default(name: str, expected: str): + span = _span(name, b"\x02" * 8, openinference__span__kind="CHAIN", metadata='{"ls_integration":"langgraph"}') + row = decode_otlp(_export(span, scope="openinference.instrumentation.langchain"), "application/x-protobuf")[0] + assert row["AgentName"] == expected + + def _span(name: str, span_id: bytes, parent: bytes = b"", **attributes: str | int) -> Span: return Span( trace_id=bytes.fromhex(TRACE_ID), diff --git a/tests/test_litellm/tracing/test_store.py b/tests/test_litellm/tracing/test_store.py index 3f43e42842c..0dd615a96ac 100644 --- a/tests/test_litellm/tracing/test_store.py +++ b/tests/test_litellm/tracing/test_store.py @@ -221,6 +221,23 @@ def test_agent_nodes_ignores_spans_of_unknown_agents(): assert agent_nodes(spans) == () +def test_trace_groups_normalized_names_and_preserves_span_labels(): + rows = [ + _row("root", "", "invoke_agent research_agent", "agent", "research_agent"), + _row("r1", "root", "researcher._execute_core", "agent", "researcher"), + _row("r2", "r1", "invoke_agent researcher", "agent", "researcher"), + _row("llm", "r2", "chat", "llm", "researcher"), + ] + result = trace_from_rows("t1", rows) + assert result is not None + assert result["summary"]["agent_names"] == ("research_agent", "researcher") + assert result["summary"]["name"] == "invoke_agent research_agent" + agents = {agent["name"]: agent for agent in result["agents"]} + assert agents["researcher"]["parent_agent"] == "research_agent" + assert agents["researcher"]["invocations"] == 2 + assert agents["researcher"]["llm_calls"] == 1 + + # ---------------------------------------------------------------- list helpers diff --git a/tests/unit/integrations/otel/test_otel_v2_config_baggage_parenting_guardrails.py b/tests/unit/integrations/otel/test_otel_v2_config_baggage_parenting_guardrails.py index 86837f7f46c..20f8c89b5ff 100644 --- a/tests/unit/integrations/otel/test_otel_v2_config_baggage_parenting_guardrails.py +++ b/tests/unit/integrations/otel/test_otel_v2_config_baggage_parenting_guardrails.py @@ -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"}) diff --git a/tests/unit/integrations/otel/test_otel_v2_destinations.py b/tests/unit/integrations/otel/test_otel_v2_destinations.py index 5a7057203e4..cb986e61229 100644 --- a/tests/unit/integrations/otel/test_otel_v2_destinations.py +++ b/tests/unit/integrations/otel/test_otel_v2_destinations.py @@ -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 diff --git a/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py b/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py index 5540cf54193..29c9ec56d91 100644 --- a/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py +++ b/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py @@ -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.""" diff --git a/tests/unit/llms/laya/__init__.py b/tests/unit/llms/laya/__init__.py new file mode 100644 index 00000000000..8b137891791 --- /dev/null +++ b/tests/unit/llms/laya/__init__.py @@ -0,0 +1 @@ + diff --git a/tests/unit/llms/laya/test_common_utils.py b/tests/unit/llms/laya/test_common_utils.py new file mode 100644 index 00000000000..c9ee0062cd2 --- /dev/null +++ b/tests/unit/llms/laya/test_common_utils.py @@ -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 diff --git a/tests/unit/proxy/auth/test_auth_utils.py b/tests/unit/proxy/auth/test_auth_utils.py index 83ac56c4c85..6cf4456a0ff 100644 --- a/tests/unit/proxy/auth/test_auth_utils.py +++ b/tests/unit/proxy/auth/test_auth_utils.py @@ -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 diff --git a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py index 90078fa836f..eb17ff0823c 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py @@ -2020,6 +2020,454 @@ async def test_auto_register_binds_api_key_to_token_hash(): assert result.end_user_id == "validated-end-user" +def _auto_register_patches(*, plaintext_key: str | None = "sk-minted-plaintext"): + from litellm.proxy.auth.auth_method import AuthMethod + from litellm.proxy.auth.resolvers.models import CredentialRef + from litellm.proxy.auth.resolvers.store import IdentityStore + from litellm.proxy.proxy_server import hash_token + + resolved_key = UserAPIKeyAuth( + token="existing-hash" if plaintext_key is None else hash_token(plaintext_key), + user_id="validated-user", + team_id="validated-team", + org_id="key-own-org", + ) + principal = IdentityStore._principal_from_key( + resolved_key, + auth_method=AuthMethod.API_KEY, + credential_ref=CredentialRef(token_id=resolved_key.token), + ) + return ( + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn", + new_callable=AsyncMock, + return_value={"token": plaintext_key}, + ), + patch( + "litellm.proxy.auth.resolvers.store.IdentityStore.resolve", + new_callable=AsyncMock, + return_value=principal, + ), + ) + + +def _auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler, **over): + kwargs = { + "virtual_key_claim_field": "sub", + "claim_value": "validated-user", + "jwt_handler": jwt_handler, + "prisma_client": prisma_client, + "user_api_key_cache": user_api_key_cache, + "parent_otel_span": None, + "proxy_logging_obj": MagicMock(), + "cache_key": "jwt_key_mapping:sub:validated-user", + "team_id": "validated-team", + "user_id": "validated-user", + "org_id": "jwt-org", + "end_user_id": "validated-end-user", + } + kwargs.update(over) + return kwargs + + +@pytest.mark.asyncio +async def test_auto_register_map_existing_key_reuses_users_key_but_never_an_auto_registered_one(): + from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping + + prisma_client = MagicMock() + prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[ + {"token": "auto-registered-hash", "metadata": {"auto_registered": True}}, + {"token": "existing-hash", "metadata": {}}, + ] + ) + prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() + + user_api_key_cache = MagicMock() + user_api_key_cache.async_set_cache = AsyncMock() + + jwt_handler = MagicMock() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + user_id_jwt_field="sub", + auto_register_map_existing_key=True, + virtual_key_mapping_cache_ttl=300, + ) + + generate_patch, resolve_patch = _auto_register_patches(plaintext_key=None) + with generate_patch as generate_key, resolve_patch: + result = await _auto_register_jwt_mapping( + **_auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler) + ) + + generate_key.assert_not_awaited() + + create_data = prisma_client.db.litellm_jwtkeymapping.create.await_args.kwargs["data"] + assert create_data["token"] == "existing-hash" + assert create_data["created_by"] == "auto_register" + assert user_api_key_cache.async_set_cache.await_args.kwargs["value"] == "existing-hash" + assert result is not None + assert result.token == "existing-hash" + assert result.api_key == "existing-hash" + assert result.org_id == "key-own-org" + + +@pytest.mark.asyncio +async def test_auto_register_map_existing_key_mints_when_user_has_no_key(): + from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping + from litellm.proxy.proxy_server import hash_token + + prisma_client = MagicMock() + prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() + + user_api_key_cache = MagicMock() + user_api_key_cache.async_set_cache = AsyncMock() + + jwt_handler = MagicMock() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + user_id_jwt_field="sub", + auto_register_map_existing_key=True, + virtual_key_mapping_cache_ttl=300, + ) + + generate_patch, resolve_patch = _auto_register_patches() + with generate_patch as generate_key, resolve_patch: + result = await _auto_register_jwt_mapping( + **_auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler) + ) + + generate_key.assert_awaited_once() + create_data = prisma_client.db.litellm_jwtkeymapping.create.await_args.kwargs["data"] + assert create_data["token"] == hash_token("sk-minted-plaintext") + assert result is not None + assert result.token == hash_token("sk-minted-plaintext") + + +@pytest.mark.asyncio +async def test_auto_register_default_never_looks_up_existing_keys(): + from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping + + prisma_client = MagicMock() + prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[{"token": "existing-hash", "metadata": {}}] + ) + prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() + + user_api_key_cache = MagicMock() + user_api_key_cache.async_set_cache = AsyncMock() + + jwt_handler = MagicMock() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_id_jwt_field="sub", virtual_key_mapping_cache_ttl=300) + + generate_patch, resolve_patch = _auto_register_patches() + with generate_patch as generate_key, resolve_patch: + await _auto_register_jwt_mapping(**_auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler)) + + prisma_client.db.litellm_verificationtoken.find_many.assert_not_awaited() + generate_key.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_auto_register_map_existing_key_race_loser_keeps_reused_key(): + from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping + + prisma_client = MagicMock() + prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[{"token": "existing-hash", "metadata": {}}] + ) + prisma_client.db.litellm_verificationtoken.delete = AsyncMock() + prisma_client.db.litellm_jwtkeymapping.create = AsyncMock(side_effect=Exception("Unique constraint failed (P2002)")) + + user_api_key_cache = MagicMock() + user_api_key_cache.async_set_cache = AsyncMock() + + jwt_handler = MagicMock() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + user_id_jwt_field="sub", + auto_register_map_existing_key=True, + virtual_key_mapping_cache_ttl=300, + ) + + generate_patch, resolve_patch = _auto_register_patches(plaintext_key=None) + with ( + generate_patch, + resolve_patch, + patch( + "litellm.proxy.auth.user_api_key_auth.get_jwt_key_mapping_object", + new_callable=AsyncMock, + return_value="winner-hash", + ), + ): + result = await _auto_register_jwt_mapping( + **_auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler) + ) + + assert result is not None + assert result.org_id == "key-own-org" + prisma_client.db.litellm_verificationtoken.delete.assert_not_awaited() + assert user_api_key_cache.async_set_cache.await_args.kwargs["value"] == "winner-hash" + + +@pytest.mark.asyncio +async def test_auto_register_map_existing_key_user_id_none_mints(): + from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping + + prisma_client = MagicMock() + prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[{"token": "existing-hash", "metadata": {}}] + ) + prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() + + user_api_key_cache = MagicMock() + user_api_key_cache.async_set_cache = AsyncMock() + + jwt_handler = MagicMock() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + user_id_jwt_field="sub", + auto_register_map_existing_key=True, + virtual_key_mapping_cache_ttl=300, + ) + + generate_patch, resolve_patch = _auto_register_patches() + with generate_patch as generate_key, resolve_patch: + await _auto_register_jwt_mapping( + **_auto_register_kwargs(prisma_client, user_api_key_cache, jwt_handler, user_id=None) + ) + + prisma_client.db.litellm_verificationtoken.find_many.assert_not_awaited() + generate_key.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_auto_register_map_existing_key_reuses_when_the_user_was_matched_by_a_fallback_lookup(): + from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping + + prisma_client = MagicMock() + prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[{"token": "existing-hash", "metadata": {}}] + ) + prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() + + user_api_key_cache = MagicMock() + user_api_key_cache.async_set_cache = AsyncMock() + + jwt_handler = MagicMock() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + user_id_jwt_field="sub", + user_email_jwt_field="email", + auto_register_map_existing_key=True, + virtual_key_mapping_cache_ttl=300, + ) + + generate_patch, resolve_patch = _auto_register_patches(plaintext_key=None) + with generate_patch as generate_key, resolve_patch: + await _auto_register_jwt_mapping( + **_auto_register_kwargs( + prisma_client, + user_api_key_cache, + jwt_handler, + claim_value="idp-subject-not-the-db-user-id", + cache_key="jwt_key_mapping:sub:idp-subject-not-the-db-user-id", + ) + ) + + generate_key.assert_not_awaited() + assert prisma_client.db.litellm_jwtkeymapping.create.await_args.kwargs["data"]["token"] == "existing-hash" + + +@pytest.mark.asyncio +async def test_auto_register_map_existing_key_mints_when_the_claim_is_not_a_user_identity_field(): + from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping + from litellm.proxy.proxy_server import hash_token + + prisma_client = MagicMock() + prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[{"token": "existing-hash", "metadata": {}}] + ) + prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() + + user_api_key_cache = MagicMock() + user_api_key_cache.async_set_cache = AsyncMock() + + jwt_handler = MagicMock() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + user_id_jwt_field="sub", + auto_register_map_existing_key=True, + virtual_key_mapping_cache_ttl=300, + ) + + generate_patch, resolve_patch = _auto_register_patches() + with generate_patch as generate_key, resolve_patch: + await _auto_register_jwt_mapping( + **_auto_register_kwargs( + prisma_client, + user_api_key_cache, + jwt_handler, + virtual_key_claim_field="azp", + claim_value="shared-client-app", + cache_key="jwt_key_mapping:azp:shared-client-app", + ) + ) + + prisma_client.db.litellm_verificationtoken.find_many.assert_not_awaited() + generate_key.assert_awaited_once() + assert prisma_client.db.litellm_jwtkeymapping.create.await_args.kwargs["data"]["token"] == hash_token( + "sk-minted-plaintext" + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("issuer_user_id_field", "expect_reuse"), + [("uid", False), (None, True)], +) +async def test_auto_register_map_existing_key_uses_the_issuers_own_user_field_over_the_global_one( + issuer_user_id_field, expect_reuse +): + from litellm.proxy._types import JWTIssuerConfig + from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping + from litellm.proxy.proxy_server import hash_token + + prisma_client = MagicMock() + prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[{"token": "existing-hash", "metadata": {}}] + ) + prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() + + user_api_key_cache = MagicMock() + user_api_key_cache.async_set_cache = AsyncMock() + + jwt_handler = MagicMock() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + user_id_jwt_field="sub", + auto_register_map_existing_key=True, + virtual_key_mapping_cache_ttl=300, + issuers=[ + JWTIssuerConfig( + issuer="https://idp.example.com", audience="litellm", user_id_jwt_field=issuer_user_id_field + ) + ], + ) + + generate_patch, resolve_patch = _auto_register_patches() + with generate_patch as generate_key, resolve_patch: + await _auto_register_jwt_mapping( + **_auto_register_kwargs( + prisma_client, user_api_key_cache, jwt_handler, jwt_issuer="https://idp.example.com" + ) + ) + + mapped_token = prisma_client.db.litellm_jwtkeymapping.create.await_args.kwargs["data"]["token"] + assert mapped_token == ("existing-hash" if expect_reuse else hash_token("sk-minted-plaintext")) + assert generate_key.await_count == (0 if expect_reuse else 1) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("map_existing_key", "master_key", "reused_key_models", "expect_denied"), + [ + (True, "sk-master", ["some-other-model"], True), + (True, "sk-master", [], False), + (False, "sk-master", ["some-other-model"], False), + (True, None, ["some-other-model"], False), + ], +) +async def test_auto_register_map_existing_key_first_request_runs_key_checks( + map_existing_key: bool, master_key: str | None, reused_key_models: list[str], expect_denied: bool +) -> None: + jwt_token = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.signature" + user_api_key_cache = DualCache() + prisma_client = MagicMock() + jwt_handler = MagicMock() + jwt_handler.is_jwt.return_value = True + jwt_handler.auth_jwt = AsyncMock(return_value={"sub": "user1"}) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + virtual_key_claim_field="sub", + virtual_key_mapping_cache_ttl=300, + auto_register_map_existing_key=map_existing_key, + ) + reused_key = UserAPIKeyAuth( + token="hashed-existing-key", + api_key="hashed-existing-key", + user_id="validated-user", + team_id="validated-team", + models=reused_key_models, + ) + mock_jwt_result = { + "is_proxy_admin": False, + "team_object": None, + "user_object": LiteLLM_UserTable(user_id="validated-user", user_role="internal_user"), + "end_user_object": None, + "org_object": None, + "token": jwt_token, + "team_id": "validated-team", + "user_id": "validated-user", + "user_email": None, + "end_user_id": None, + "org_id": None, + "team_membership": None, + "jwt_claims": {"sub": "user1"}, + } + + mock_request = MagicMock() + mock_request.url.path = "/v1/chat/completions" + mock_request.method = "POST" + mock_request.headers = {"authorization": f"Bearer {jwt_token}"} + mock_request.query_params = {} + mock_request.state = SimpleNamespace() + + with ( + patch("litellm.proxy.proxy_server.general_settings", {"enable_jwt_auth": True}), + patch("litellm.proxy.proxy_server.premium_user", True), + patch("litellm.proxy.proxy_server.master_key", master_key), + patch("litellm.proxy.proxy_server.prisma_client", prisma_client), + patch("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache), + patch( + "litellm.proxy.proxy_server.proxy_logging_obj", + MagicMock(post_call_failure_hook=AsyncMock(return_value=None)), + ), + patch("litellm.proxy.proxy_server.jwt_handler", jwt_handler), + patch( + "litellm.proxy.auth.user_api_key_auth._resolve_jwt_to_virtual_key", + new_callable=AsyncMock, + return_value=_PendingAutoRegister( + claim_field="sub", + claim_value="user1", + cache_key="jwt_key_mapping:sub:user1", + ), + ), + patch( + "litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder", + new_callable=AsyncMock, + return_value=mock_jwt_result, + ), + patch( + "litellm.proxy.auth.user_api_key_auth._auto_register_jwt_mapping", + new_callable=AsyncMock, + return_value=reused_key, + ), + ): + call = _user_api_key_auth_builder( + request=mock_request, + api_key=jwt_token, + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={"model": "gpt-4o-mini"}, + ) + if expect_denied: + with pytest.raises(ProxyException, match="not available for this API key"): + await call + return + result = await call + + assert result.api_key == "hashed-existing-key" + assert result.user_id == "validated-user" + assert result.team_id == "validated-team" + assert result.models == reused_key_models + + @pytest.mark.asyncio @pytest.mark.parametrize("active", [True, False]) async def test_auto_register_first_request_propagates_user_email(active: bool) -> None: diff --git a/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py b/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py index 2cbba9da8b3..01e41e8b03f 100644 --- a/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py @@ -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 diff --git a/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py index 5f7807650e1..7eda03c560b 100644 --- a/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py @@ -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"]) diff --git a/tests/unit/proxy/management_helpers/test_auto_router_permissions.py b/tests/unit/proxy/management_helpers/test_auto_router_permissions.py index e16271a5189..b60fd4ac7ad 100644 --- a/tests/unit/proxy/management_helpers/test_auto_router_permissions.py +++ b/tests/unit/proxy/management_helpers/test_auto_router_permissions.py @@ -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}") diff --git a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_typesafe_passthrough_logging_handler.py b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_typesafe_passthrough_logging_handler.py index e0a5ef063e8..acf05dcdfde 100644 --- a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_typesafe_passthrough_logging_handler.py +++ b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_typesafe_passthrough_logging_handler.py @@ -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"] diff --git a/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index a0772c7d4f3..22171f4ffb0 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -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]: diff --git a/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 52acdf93f35..c6c81c14b16 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -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( diff --git a/tests/unit/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py b/tests/unit/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py index 73927e92c15..09987b2781c 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py +++ b/tests/unit/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py @@ -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: diff --git a/tests/unit/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py b/tests/unit/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py index 29a635e9b27..d6b69c7c010 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py +++ b/tests/unit/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py @@ -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" diff --git a/tests/unit/proxy/test_spend_log_cleanup.py b/tests/unit/proxy/test_spend_log_cleanup.py index 46ac1234615..399c76d97c1 100644 --- a/tests/unit/proxy/test_spend_log_cleanup.py +++ b/tests/unit/proxy/test_spend_log_cleanup.py @@ -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 diff --git a/tests/unit/router_strategy/complexity_router/test_jev_classifier.py b/tests/unit/router_strategy/complexity_router/test_jev_classifier.py index 45070dfd3a7..affcfdc789c 100644 --- a/tests/unit/router_strategy/complexity_router/test_jev_classifier.py +++ b/tests/unit/router_strategy/complexity_router/test_jev_classifier.py @@ -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"), ], diff --git a/tests/unit/router_utils/test_auto_router_model_naming.py b/tests/unit/router_utils/test_auto_router_model_naming.py index 645f9e5e62a..4881b850f2a 100644 --- a/tests/unit/router_utils/test_auto_router_model_naming.py +++ b/tests/unit/router_utils/test_auto_router_model_naming.py @@ -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 diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index 4c36f99d2c8..ab7ecfcda05 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -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", ], }, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.test.tsx index b4b34dbfaf3..081eb7f6e09 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.test.tsx @@ -75,7 +75,6 @@ const totals = (overrides: Partial = {}): 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('[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(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.tsx index 24a97587e32..27e8df6db87 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.tsx @@ -105,6 +105,12 @@ const HeroCard: React.FC<{ view: BenchmarkView }> = ({ view }) => { adaptive and quality routers are excluded

)} + {stats.unattributed_saved_spend != null && ( +

+ 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 +

+ )}
@@ -297,22 +303,34 @@ const BenchmarksBody: React.FC = ({ isPending, error, data, -
- - - - -
+

+ Savings and spend count requests on the selected UTC days. Actual spend covers every request on complexity + routers, including LLM classification cost. Baseline is actual spend plus recorded savings, so savings can be + zero or negative. +

- Actual spend covers every request on complexity routers, including LLM classification cost. Baseline is actual - spend plus recorded savings, so savings can be zero or negative. The range counts whole sessions that overlap - it, so totals can differ from savings views that group usage by UTC day. + Session metrics cover every session that overlaps the range, including its turns outside the range.

+
+ + + +
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/TierTurnsChart.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/TierTurnsChart.test.tsx index e4417d77463..42444fd8f06 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/TierTurnsChart.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/TierTurnsChart.test.tsx @@ -12,21 +12,21 @@ vi.mock("@/components/shared/charts", () => ({ })); import TierTurnsChart, { tierDisplayLabel } from "./TierTurnsChart"; -import type { AutoRouterBenchmarkGroup, BenchmarkView } from "./autoRouterBenchmarks"; +import type { AutoRouterBenchmarkGroup, AutoRouterBenchmarkTotals, BenchmarkView } from "./autoRouterBenchmarks"; -const totalsOnly = { +const totalsOnly: AutoRouterBenchmarkTotals = { sessions: 3, turns: 9, avg_turns_per_session: 3, avg_session_seconds: 60, avg_tokens_per_session: 100, spend: 1, + classifier_cost: 0, savings_estimated_turns: 9, savings_estimated_actual_spend: 1, saved_spend: 1, baseline_spend: 2, saved_pct: 50, - saved_per_session: 0.33, cache: { coverage_pct: 0, hit_rate_pct: 0, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/autoRouterBenchmarks.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/autoRouterBenchmarks.test.ts index 0586163e77e..62d0c0e4e5e 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/autoRouterBenchmarks.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/autoRouterBenchmarks.test.ts @@ -43,7 +43,6 @@ const totals = (overrides: Partial = {}) => ({ saved_spend: 2174.59, baseline_spend: 2534.45, saved_pct: 85.8, - saved_per_session: 23.13, cache: cache(), ...overrides, }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.test.ts index 79c4243271e..153b77666a7 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.test.ts @@ -85,7 +85,8 @@ describe("autoRouterRows", () => { it.each([ ["llm", "LLM Classifier"], - ["jev", "JEV Classifier"], + ["jev", "OSS Classifier"], + ["oss_classifier", "OSS Classifier"], ])("labels a router using the %s classifier", (classifierType, label) => { const row = toAutoRouterRow( { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.ts b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.ts index 1faf3408c23..3c3366a767d 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AutoRouters/autoRouterRows.ts @@ -57,7 +57,8 @@ const dedupe = (models: string[]): string[] => Array.from(new Set(models)); const COMPLEXITY_TYPE_LABELS: Record = { llm: "LLM Classifier", - jev: "JEV Classifier", + jev: "OSS Classifier", + oss_classifier: "OSS Classifier", capability: "Capability", llm_v2: "Fuse v2", heuristic_first: "Heuristic first", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/users/_components/view_users/user_info_view.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/users/_components/view_users/user_info_view.integration.test.tsx index 676bc7d29eb..4195a3c9913 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/users/_components/view_users/user_info_view.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/users/_components/view_users/user_info_view.integration.test.tsx @@ -347,7 +347,6 @@ const routerUsageResponse = (saved: number): AutoRouterBenchmarksResponse => ({ saved_spend: saved, baseline_spend: 10 + saved, saved_pct: (100 * saved) / (10 + saved), - saved_per_session: saved / 2, cache: { coverage_pct: 100, hit_rate_pct: 0, diff --git a/ui/litellm-dashboard/src/components/add_model/AutoRouterAvailability.tsx b/ui/litellm-dashboard/src/components/add_model/AutoRouterAvailability.tsx index b1233ef03b8..135ebb958db 100644 --- a/ui/litellm-dashboard/src/components/add_model/AutoRouterAvailability.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AutoRouterAvailability.tsx @@ -136,7 +136,7 @@ export const AutoRouterLimits = () => { Routing and customization limits

- Rule-based, Complexity, and Jev are unlimited with built-in settings. Choose or change tier models freely. + Rule-based, Complexity, and OSS are unlimited with built-in settings. Choose or change tier models freely. Customization allowances are shared across this proxy.

diff --git a/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.integration.test.tsx b/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.integration.test.tsx index f843a472d15..08f329e4ad7 100644 --- a/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.integration.test.tsx @@ -68,7 +68,7 @@ describe("Auto-router classifier selection", () => { llm: "LLM", heuristic_first: "LLM", hybrid: "LLM", - jev: "Jev", + jev: "OSS Classifier", }[classifier_type]; expect(screen.getByRole("radio", { name: new RegExp(`^${family}$`) })).toBeChecked(); fireEvent.click(screen.getByRole("radio", { name: new RegExp(`^${family}$`) })); diff --git a/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.tsx b/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.tsx index 92fc8d2a335..3a1e0065530 100644 --- a/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.tsx @@ -16,6 +16,7 @@ import { type ClassifierType, type ComplexityRouterConfigValue, } from "./ComplexityRouterConfig"; +import { defaultJevClassifierConfig, normalizeJevClassifierConfig } from "./jev_classifier_config"; import { transitionClassifierType } from "./classifier_type_transition"; import { isForecastClassifier } from "./forecast_classifier_config"; import { @@ -148,6 +149,14 @@ const AutoRouterClassifierTabs: React.FC = ({ val if (next === "llm") changeType("llm"); if (next === "jev") changeType("jev"); }; + const changeProvider = (provider: unknown) => { + if (provider !== "jev" && provider !== "laya") return; + const defaults = defaultJevClassifierConfig(provider); + onChange({ + ...value, + jev_classifier_config: { ...defaults, ...value.jev_classifier_config, provider, model: defaults.model }, + }); + }; const approachLabels: Partial> = { capability: "Capability", llm_v2: "Fuse v2" }; const approachDescription: Partial> = { capability: "Use the efficient model when it is likely to succeed", @@ -164,7 +173,7 @@ const AutoRouterClassifierTabs: React.FC = ({ val {[ { value: "heuristics", label: "Heuristics", description: "Classify locally, with no API call" }, { value: "llm", label: "LLM", description: "Use a judge model to choose a solver" }, - { value: "jev", label: "Jev", description: "Use TypeSafe System One Choice to choose a tier" }, + { value: "jev", label: "OSS Classifier", description: "Use Jev or Laya to choose a tier" }, ].map((option) => (
diff --git a/ui/litellm-dashboard/src/components/add_model/ClassifierTypeRadios.tsx b/ui/litellm-dashboard/src/components/add_model/ClassifierTypeRadios.tsx index 1fd6dfa6a20..2eda2c63f64 100644 --- a/ui/litellm-dashboard/src/components/add_model/ClassifierTypeRadios.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ClassifierTypeRadios.tsx @@ -54,8 +54,8 @@ const ClassifierTypeRadios: React.FC = ({ value, clas diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx index 32f9ebf97ad..e324886549b 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx @@ -237,7 +237,7 @@ const TierSetToolbar: React.FC<{ {editing && ( 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 )} {editing && keywordRulesError && ( diff --git a/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.integration.test.tsx b/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.integration.test.tsx index 7da8b12c2d7..e175fc934b9 100644 --- a/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.integration.test.tsx @@ -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(
); - 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(); + 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 ; }; renderWithProviders(); - 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(""); }); }); diff --git a/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.tsx b/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.tsx index 97609bcd8c3..22e8708acc7 100644 --- a/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.tsx @@ -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) => onChange({ ...value, jev_classifier_config: { ...config, ...patch } }); return (

- 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"}

- - update({ model: event.target.value })} /> + + {isLaya ? ( + + ) : ( + update({ model: event.target.value })} /> + )}
- +
- + {config.instructions && ( )}

- 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

diff --git a/ui/litellm-dashboard/src/components/add_model/JevConnectionTest.integration.test.tsx b/ui/litellm-dashboard/src/components/add_model/JevConnectionTest.integration.test.tsx index ff147e20d1a..02b065083a0 100644 --- a/ui/litellm-dashboard/src/components/add_model/JevConnectionTest.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/JevConnectionTest.integration.test.tsx @@ -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(); diff --git a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.integration.test.tsx b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.integration.test.tsx index 0337d35d698..4a871d5b103 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.integration.test.tsx @@ -274,13 +274,13 @@ describe("AddAutoRouterTab", () => { mockFetchAvailableModels.mockResolvedValue(ALL_FAMILY_MODELS); renderWithProviders(); 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(); 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(); - 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(); + 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, }); }); diff --git a/ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.tsx b/ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.tsx index aa2ce91766a..eafa5122b59 100644 --- a/ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.tsx @@ -50,7 +50,7 @@ const AutoRouterConnectionTest: React.FC = ({ ? { 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 = ({ classifier probe includes its reasoning effort override.

{jevRequest && ( -
- Jev Classifier +
+ OSS Classifier

- {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}

diff --git a/ui/litellm-dashboard/src/components/add_model/build_auto_router_routing_test_request.test.ts b/ui/litellm-dashboard/src/components/add_model/build_auto_router_routing_test_request.test.ts index de0fb6fe6e1..3686103f234 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_auto_router_routing_test_request.test.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_auto_router_routing_test_request.test.ts @@ -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", diff --git a/ui/litellm-dashboard/src/components/add_model/build_auto_router_routing_test_request.ts b/ui/litellm-dashboard/src/components/add_model/build_auto_router_routing_test_request.ts index 6a9d1ce7d92..9a5be0837a0 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_auto_router_routing_test_request.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_auto_router_routing_test_request.ts @@ -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 }), }; diff --git a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts index 6d458d61797..c054f6bbe9d 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts @@ -64,6 +64,7 @@ describe("buildComplexityRouterConfig", () => { it.each([ { model: "" }, { model: " " }, + { provider: "laya" as const, model: "unsupported" }, { timeout_ms: 0 }, { timeout_ms: 1.5 }, { timeout_ms: Number.NaN }, @@ -75,15 +76,16 @@ describe("buildComplexityRouterConfig", () => { classifier_type: "jev", jev_classifier_config: { model: "jev-latest", timeout_ms: 3000, ...patch }, }), - ).toBe("Enter a JEV model, a positive whole-number timeout and a positive cooldown"); + ).toBe("Enter a valid classifier model, a positive whole-number timeout and a positive cooldown"); }); - it.each([false, true])("serializes JEV with shared context and no LLM config, custom tiers: %s", (custom) => { + it.each([false, true])("serializes Laya with shared context and no LLM config, custom tiers: %s", (custom) => { const params: BuildComplexityRouterConfigParams = { ...baseParams, classifierType: "jev", jevClassifierConfig: { - model: "jev-test", + provider: "laya", + model: "english", timeout_ms: 4500, instructions: " Choose the configured tier ", circuit_breaker_enabled: false, @@ -108,15 +110,17 @@ describe("buildComplexityRouterConfig", () => { }), }; const config = buildComplexityRouterConfig(params); - expect(config.classifier_type).toBe("jev"); + expect(config.classifier_type).toBe("oss_classifier"); const expectedJevConfig = { - model: "jev-test", + provider: "laya", + model: "english", timeout_ms: 4500, instructions: "Choose the configured tier", circuit_breaker_enabled: false, circuit_breaker_cooldown_seconds: 12.5, }; - expect(config.jev_classifier_config).toEqual(expectedJevConfig); + expect(config.opensource_classifier_config).toEqual(expectedJevConfig); + expect(config).not.toHaveProperty("jev_classifier_config"); expect(config.classifier_context_window_size).toBe(4); expect(config.classifier_context_budget_chars).toBe(2000); expect(config.classifier_context_per_turn_chars).toBe(450); @@ -139,15 +143,15 @@ describe("buildComplexityRouterConfig", () => { classifierType: "jev", jevClassifierConfig: { model: "jev-latest", timeout_ms: 3000, instructions: " " }, }); - expect(jev.jev_classifier_config).toEqual({ model: "jev-latest", timeout_ms: 3000 }); + expect(jev.opensource_classifier_config).toEqual({ provider: "jev", model: "jev-latest", timeout_ms: 3000 }); const llmParams: BuildComplexityRouterConfigParams = { ...baseParams, classifierType: "llm", classifierLlmConfig: { model: "judge", timeout_ms: 1000 }, - jevClassifierConfig: jev.jev_classifier_config, + jevClassifierConfig: jev.opensource_classifier_config, }; const llm = buildComplexityRouterConfig(llmParams); - expect(llm).not.toHaveProperty("jev_classifier_config"); + expect(llm).not.toHaveProperty("opensource_classifier_config"); }); it("forwards preset references and explicit overrides without materializing absent text on create", () => { diff --git a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts index 13df40464f0..7756d0ee8fa 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts @@ -153,12 +153,13 @@ export interface StoredComplexityRouterConfig { heuristic_first_max_tier?: unknown; hybrid_boundary_margin?: unknown; tier_labels?: unknown; - classifier_type?: ClassifierType; + classifier_type?: ClassifierType | "oss_classifier"; heuristic_v2_success_threshold?: unknown; capability_classifier_config?: unknown; llm_v2_config?: unknown; classifier_llm_config?: ClassifierLLMConfig; jev_classifier_config?: unknown; + opensource_classifier_config?: unknown; classifier_context_window_size?: unknown; classifier_context_budget_chars?: unknown; classifier_context_per_turn_chars?: unknown; @@ -283,12 +284,13 @@ export interface ComplexityRouterConfigPayload { default_model?: string; plan_mode_min_tier?: string; tier_labels?: ComplexityTierLabels; - classifier_type: ClassifierType; + classifier_type: ClassifierType | "oss_classifier"; heuristic_v2_success_threshold?: number; capability_classifier_config?: CapabilitySettings; llm_v2_config?: FuseSettings; classifier_llm_config?: ClassifierLLMConfig; jev_classifier_config?: JevClassifierConfig; + opensource_classifier_config?: JevClassifierConfig; classifier_context_window_size?: number; classifier_context_budget_chars?: number; classifier_context_per_turn_chars?: number; @@ -438,7 +440,9 @@ export const getClassifierModelError = ( ): string | null => { if (effectiveClassifierType(config) === "jev") { const parsed = jevClassifierConfigSchema.safeParse(config.jev_classifier_config ?? {}); - return parsed.success ? null : "Enter a JEV model, a positive whole-number timeout and a positive cooldown"; + return parsed.success + ? null + : "Enter a valid classifier model, a positive whole-number timeout and a positive cooldown"; } if (!usesLlmClassifier(effectiveClassifierType(config)) || config.classifier_llm_config?.model) return null; return config.custom_tier_set @@ -498,7 +502,7 @@ export const customTierWireFields = ( tiers: Object.fromEntries(rows.map((row) => [activeTierName(row), row.models])), tier_definitions: tierDefinitionsFromRows(rows), ...(fallback && { fallback_tier: activeTierName(fallback) }), - classifier_type: classifierType === "jev" ? "jev" : "llm", + classifier_type: classifierType === "jev" ? "oss_classifier" : "llm", // Rebuilt from the fields an edited tier set allows. The backend rejects system_prompt and // classification_rubric beside tier_definitions, and both live inside this object rather than at // the top level the omit list covers. The opening instructions ride classification_prompt below. @@ -769,8 +773,8 @@ export const buildComplexityRouterConfig = ({ ...(defaultModel?.trim() && { default_model: defaultModel }), ...(planModeMinTier?.trim() && { plan_mode_min_tier: planModeMinTier }), ...(cleanedTierLabels && { tier_labels: cleanedTierLabels }), - classifier_type: classifierType, - ...(effectiveType === "jev" && { jev_classifier_config: normalizeJevClassifierConfig(jevClassifierConfig) }), + classifier_type: classifierType === "jev" ? "oss_classifier" : classifierType, + ...(effectiveType === "jev" && { opensource_classifier_config: normalizeJevClassifierConfig(jevClassifierConfig) }), ...(heuristicV2SuccessThreshold !== undefined && { heuristic_v2_success_threshold: heuristicV2SuccessThreshold, }), diff --git a/ui/litellm-dashboard/src/components/add_model/jev_classifier_config.ts b/ui/litellm-dashboard/src/components/add_model/jev_classifier_config.ts index 478c763351c..d1e1c96863c 100644 --- a/ui/litellm-dashboard/src/components/add_model/jev_classifier_config.ts +++ b/ui/litellm-dashboard/src/components/add_model/jev_classifier_config.ts @@ -1,7 +1,11 @@ import { z } from "zod"; +import type { ClassifierType } from "./classifier_types"; + +export const LAYA_MODELS = ["english", "multilingual", "typed-decisions"] as const; const jevClassifierConfigFields = { - model: z.string().trim().min(1).default("jev-latest"), + provider: z.preprocess((value) => (value === "typesafe" ? "jev" : value), z.enum(["jev", "laya"]).optional()), + model: z.string().trim().min(1).optional(), timeout_ms: z.number().int().positive().default(3000), instructions: z .string() @@ -11,15 +15,39 @@ const jevClassifierConfigFields = { circuit_breaker_cooldown_seconds: z.number().finite().positive().optional(), }; -export const jevClassifierConfigSchema = z.object(jevClassifierConfigFields); +export const jevClassifierConfigSchema = z + .object(jevClassifierConfigFields) + .transform((config) => ({ + ...config, + model: config.model ?? (config.provider === "laya" ? "english" : "jev-latest"), + })) + .refine((config) => config.provider !== "laya" || LAYA_MODELS.some((model) => model === config.model), { + message: "Select a supported Laya model", + path: ["model"], + }); export type JevClassifierConfig = z.infer; -export const defaultJevClassifierConfig = (): JevClassifierConfig => jevClassifierConfigSchema.parse({}); +export const defaultJevClassifierConfig = (provider: "jev" | "laya" = "jev"): JevClassifierConfig => + jevClassifierConfigSchema.parse({ provider }); + +export const hydrateOssClassifier = (config: { + classifier_type?: ClassifierType | "oss_classifier"; + opensource_classifier_config?: unknown; + jev_classifier_config?: unknown; +}): { classifier_type: ClassifierType; jev_classifier_config?: JevClassifierConfig } => ({ + classifier_type: config.classifier_type === "oss_classifier" ? "jev" : config.classifier_type ?? "heuristic", + jev_classifier_config: + config.classifier_type === "oss_classifier" || config.classifier_type === "jev" + ? jevClassifierConfigSchema.safeParse(config.opensource_classifier_config ?? config.jev_classifier_config ?? {}) + .data ?? defaultJevClassifierConfig() + : undefined, +}); export const normalizeJevClassifierConfig = ( config: JevClassifierConfig = defaultJevClassifierConfig(), ): JevClassifierConfig => ({ + provider: config.provider ?? "jev", model: config.model.trim(), timeout_ms: config.timeout_ms, ...(config.instructions?.trim() && { instructions: config.instructions.trim() }), diff --git a/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts b/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts index 4b8e298191d..91514c3a7d6 100644 --- a/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts +++ b/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts @@ -48,7 +48,7 @@ const hydratedState: KeywordMatchingState = { }; describe("buildUpdatedComplexityRouterConfig keyword matching", () => { - it.each([false, true])("omits masked JEV credentials from dashboard saves, edited: %s", (edited) => { + it.each([false, true])("omits masked Jev credentials from legacy/canonical saves, edited: %s", (edited) => { const stored = { classifier_type: "jev" as const, tiers: FORM_VALUE.tiers, @@ -60,7 +60,15 @@ describe("buildUpdatedComplexityRouterConfig keyword matching", () => { api_base: "https://jev.example.com", }, }; - const hydrated = hydrateComplexityRouterConfig(stored, undefined); + const source = edited + ? { + ...stored, + classifier_type: "oss_classifier" as const, + jev_classifier_config: undefined, + opensource_classifier_config: { ...stored.jev_classifier_config, provider: "typesafe" }, + } + : stored; + const hydrated = hydrateComplexityRouterConfig(source, undefined); expect(hydrated.jev_classifier_config).not.toHaveProperty("api_key"); expect(hydrated.jev_classifier_config).not.toHaveProperty("api_base"); const value = edited @@ -69,8 +77,9 @@ describe("buildUpdatedComplexityRouterConfig keyword matching", () => { jev_classifier_config: { model: "jev-updated", timeout_ms: 8100, instructions: "" }, } : hydrated; - const saved = buildUpdatedComplexityRouterConfig(stored, value); - expect(saved.jev_classifier_config).toEqual({ + const saved = buildUpdatedComplexityRouterConfig(source, value); + expect(saved.opensource_classifier_config).toEqual({ + provider: "jev", ...(edited ? { model: "jev-updated", timeout_ms: 8100 } : { model: "jev-configured", timeout_ms: 6100, instructions: "Existing instructions" }), @@ -78,7 +87,7 @@ describe("buildUpdatedComplexityRouterConfig keyword matching", () => { for (const classifierType of ["llm", "heuristic"] as const) { expect( buildUpdatedComplexityRouterConfig(saved, transitionClassifierType(value, classifierType)), - ).not.toHaveProperty("jev_classifier_config"); + ).not.toHaveProperty("opensource_classifier_config"); } }); @@ -94,19 +103,21 @@ describe("buildUpdatedComplexityRouterConfig keyword matching", () => { tiers: FORM_VALUE.tiers, }; const saved = buildUpdatedComplexityRouterConfig(stored, hydrateComplexityRouterConfig(stored, undefined)); - expect(saved.jev_classifier_config).toEqual({ + expect(saved.opensource_classifier_config).toEqual({ + provider: "jev", model: "jev-configured", timeout_ms: 6100, circuit_breaker_enabled: false, }); }); - it.each([false, true])("round trips JEV settings and preserves unmanaged fields, custom: %s", (custom) => { + it.each([false, true])("round trips Laya settings and preserves unmanaged fields, custom: %s", (custom) => { const stored = { ...(custom ? storedCustomConfig() : STORED), classifier_llm_config: { model: "stale-judge", timeout_ms: 3000 }, - classifier_type: "jev" as const, - jev_classifier_config: { - model: "jev-test", + classifier_type: "oss_classifier" as const, + opensource_classifier_config: { + provider: "laya" as const, + model: "english", timeout_ms: 4100, instructions: "Judge the request", circuit_breaker_enabled: false, @@ -121,12 +132,12 @@ describe("buildUpdatedComplexityRouterConfig keyword matching", () => { const hydrated = hydrateComplexityRouterConfig(stored, undefined); expect(effectiveClassifierType(hydrated)).toBe("jev"); expect(hydrated.classifier_llm_config).toBeUndefined(); - expect(hydrated.jev_classifier_config).toEqual(stored.jev_classifier_config); + expect(hydrated.jev_classifier_config).toEqual(stored.opensource_classifier_config); expect(hydrated.classifier_context_per_turn_chars).toBe(450); const saved = buildUpdatedComplexityRouterConfig(stored, hydrated); const expectedSavedConfig = { - classifier_type: "jev", - jev_classifier_config: stored.jev_classifier_config, + classifier_type: "oss_classifier", + opensource_classifier_config: stored.opensource_classifier_config, classifier_context_window_size: 7, classifier_context_budget_chars: 9000, classifier_context_per_turn_chars: 450, @@ -135,12 +146,13 @@ describe("buildUpdatedComplexityRouterConfig keyword matching", () => { }; expect(saved).toMatchObject(expectedSavedConfig); expect(saved).not.toHaveProperty("classifier_llm_config"); + expect(saved).not.toHaveProperty("jev_classifier_config"); const reloaded = hydrateComplexityRouterConfig(saved, undefined); expect(reloaded.jev_classifier_config).toEqual(hydrated.jev_classifier_config); expect(reloaded.classifier_context_per_turn_chars).toBe(450); expect(effectiveClassifierType(reloaded)).toBe("jev"); const llm = buildUpdatedComplexityRouterConfig(saved, transitionClassifierType(reloaded, "llm")); - expect(llm).not.toHaveProperty("jev_classifier_config"); + expect(llm).not.toHaveProperty("opensource_classifier_config"); }); it.each([0, 0.92, 1])("hydrates and saves a success threshold of %s without changing the artifact", (threshold) => { @@ -883,6 +895,7 @@ describe("managed keys survive an untouched open-and-save", () => { "fallback_tier", "hybrid_boundary_margin", "jev_classifier_config", + "opensource_classifier_config", "classifier_plugin_timeout_ms", ]); diff --git a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx index 2fd4fba2b31..e91936ebd36 100644 --- a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx +++ b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx @@ -101,6 +101,7 @@ export const MANAGED_COMPLEXITY_ROUTER_KEYS = new Set([ "llm_v2_config", "classifier_llm_config", "jev_classifier_config", + "opensource_classifier_config", "classifier_context_window_size", "classifier_context_budget_chars", "classifier_context_include_assistant_turns", diff --git a/ui/litellm-dashboard/src/components/edit_auto_router/hydrate_complexity_router_config.ts b/ui/litellm-dashboard/src/components/edit_auto_router/hydrate_complexity_router_config.ts index 6dbd2b19b52..f357fa77cea 100644 --- a/ui/litellm-dashboard/src/components/edit_auto_router/hydrate_complexity_router_config.ts +++ b/ui/litellm-dashboard/src/components/edit_auto_router/hydrate_complexity_router_config.ts @@ -1,4 +1,4 @@ -import { defaultJevClassifierConfig, jevClassifierConfigSchema } from "../add_model/jev_classifier_config"; +import { hydrateOssClassifier } from "../add_model/jev_classifier_config"; import { capabilitySettingsSchema, fuseSettingsSchema } from "../add_model/forecast_classifier_config"; import type { StoredComplexityRouterConfig } from "../add_model/build_complexity_router_config"; import { @@ -54,6 +54,7 @@ export const hydrateComplexityRouterConfig = ( parsedConfig: StoredComplexityRouterConfig, complexityRouterDefaultModel: string | null | undefined, ): ComplexityRouterConfigValue => { + const classifier = hydrateOssClassifier(parsedConfig); const builtIn = hydrateBuiltInTiers(parsedConfig.tiers, parsedConfig.enable_non_reasoning_tier); const { tiers: hydratedTiers, enable_non_reasoning_tier } = builtIn; const custom_tier_set = hydrateCustomTierSet(parsedConfig); @@ -70,19 +71,14 @@ export const hydrateComplexityRouterConfig = ( default_model: hydratePinnedDefaultModel(parsedConfig.default_model, complexityRouterDefaultModel, activeTiers), plan_mode_min_tier: hydratePlanModeMinTier(parsedConfig.plan_mode_min_tier, custom_tier_set), tier_labels: hydrateTierLabels(parsedConfig.tier_labels), - classifier_type: parsedConfig.classifier_type || "heuristic", + ...classifier, heuristic_v2_success_threshold: typeof parsedConfig.heuristic_v2_success_threshold === "number" ? parsedConfig.heuristic_v2_success_threshold : undefined, capability_classifier_config: capabilitySettingsSchema.safeParse(parsedConfig.capability_classifier_config).data, llm_v2_config: fuseSettingsSchema.safeParse(parsedConfig.llm_v2_config).data, - classifier_llm_config: parsedConfig.classifier_type === "jev" ? undefined : parsedConfig.classifier_llm_config, - jev_classifier_config: - parsedConfig.classifier_type === "jev" - ? jevClassifierConfigSchema.safeParse(parsedConfig.jev_classifier_config ?? {}).data ?? - defaultJevClassifierConfig() - : undefined, + classifier_llm_config: classifier.classifier_type === "jev" ? undefined : parsedConfig.classifier_llm_config, classifier_context_window_size: typeof parsedConfig.classifier_context_window_size === "number" ? parsedConfig.classifier_context_window_size diff --git a/ui/litellm-dashboard/src/components/templates/KeyAutoRouterUsageTab.integration.test.tsx b/ui/litellm-dashboard/src/components/templates/KeyAutoRouterUsageTab.integration.test.tsx index 5c7b9b28789..2ce585e895e 100644 --- a/ui/litellm-dashboard/src/components/templates/KeyAutoRouterUsageTab.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/templates/KeyAutoRouterUsageTab.integration.test.tsx @@ -39,7 +39,6 @@ const stats = { saved_spend: 8.75, baseline_spend: 10, saved_pct: 87.5, - saved_per_session: 4.375, cache, }; diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/RoutingDecisionCard.tsx b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/RoutingDecisionCard.tsx index 6a8cedd3746..b10aaff7cd7 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/RoutingDecisionCard.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/RoutingDecisionCard.tsx @@ -143,7 +143,7 @@ function describeCause(decision: RoutingDecision): string { case "llm_classifier": return classifierModel ? `LLM classifier (${classifierModel})` : "LLM classifier"; case "jev_classifier": - return "JEV classifier"; + return "OSS classifier"; case "literal_keyword_match": case "keyword": return matchedKeyword ? `Keyword match: "${matchedKeyword}"` : "Keyword match"; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.test.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.test.tsx index 5db34ef9788..d9cb4ee0399 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.test.tsx @@ -271,10 +271,13 @@ describe("AgentTracesSection", () => { expect(rows[0]).toHaveTextContent("Should we store OTEL agent spans"); }); - it("labels the OTEL service as the agent and filters runs by it", async () => { + it("uses recorded agent names for the column and filter even when services are shared", async () => { vi.mocked(agentTraceListCall).mockResolvedValue({ ...(traceList as TracePage), - data: [...runs.slice(1), { ...runs[0], service: "billing-agent" }], + data: [ + ...runs.slice(1).map((run) => ({ ...run, service: "shared-app", agent_names: ["research-agent"] })), + { ...runs[0], service: "shared-app", agent_names: ["billing-agent", "review-agent"] }, + ], }); const user = userEvent.setup(); renderSection(); @@ -288,6 +291,10 @@ describe("AgentTracesSection", () => { const rows = screen.getAllByTestId("agent-trace-row"); expect(rows).toHaveLength(1); expect(rows[0]).toHaveTextContent("billing-agent"); + expect(rows[0]).not.toHaveTextContent("shared-app"); + + await chooseSelectOption(user, agentFilter, "review-agent"); + expect(screen.getAllByTestId("agent-trace-row")).toHaveLength(1); await chooseSelectOption(user, agentFilter, "All agents"); expect(screen.getAllByTestId("agent-trace-row")).toHaveLength(runs.length); diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx index 06094864230..4ec8fe7e0ec 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx @@ -9,7 +9,7 @@ import { AgentTracesTable } from "./AgentTracesTable"; import { RunDrawer } from "./RunDrawer"; import { ALL_AGENTS, RunsToolbar, type RunStatusFilter } from "./RunsToolbar"; import type { TraceSummary } from "./traceTypes"; -import { previewText } from "./traceUtils"; +import { previewText, traceAgentNames } from "./traceUtils"; import { TimeRangeControls } from "./TimeRangeControls"; import { TracesTimeline, type TimeWindow } from "./TracesTimeline"; import { ActiveDot } from "./ActiveDot"; @@ -27,7 +27,7 @@ export function filterRuns( return runs.filter((run) => { const haystack = [run.trace_id, previewText(run.input_preview), run.name].map((s) => s.toLowerCase()); const matchesQuery = !q || haystack.some((text) => text.includes(q)); - const matchesAgent = agent === ALL_AGENTS || run.service === agent; + const matchesAgent = agent === ALL_AGENTS || traceAgentNames(run).includes(agent); const failed = run.error_count > 0; const matchesStatus = status === "all" || (status === "error" ? failed : !failed); return matchesQuery && matchesAgent && matchesStatus; @@ -121,7 +121,7 @@ export function AgentTracesSection({ if (setup.disabledDetail == null) void history.refetch(); }; - const agents = useMemo(() => Array.from(new Set(traces.traces.map((t) => t.service))).sort(), [traces.traces]); + const agents = useMemo(() => Array.from(new Set(traces.traces.flatMap(traceAgentNames))).sort(), [traces.traces]); // Relative ranges end "now" (the list query uses Date.now() too); round to the minute so the histogram is stable. const endMs = isCustomDate ? moment(endTime).valueOf() : moment().endOf("minute").valueOf(); const range = useMemo( diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx index 60f6816bd7a..6f0c8eac5f0 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx @@ -8,7 +8,7 @@ import { cn } from "@/lib/cva.config"; import { StatusMark } from "./StatusMark"; import type { TraceSummary } from "./traceTypes"; -import { fmtMs, previewText, traceDisplayName } from "./traceUtils"; +import { fmtMs, previewText, traceDisplayName, traceAgentNames } from "./traceUtils"; interface AgentTracesTableProps { traces: TraceSummary[]; @@ -83,8 +83,8 @@ export function AgentTracesTable({ > {formatActivityTimestamp(run.start_time)} - - {run.service} + + {traceAgentNames(run).join(", ") || "—"}
diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.test.ts b/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.test.ts index 9353f2024c8..f1394cec980 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.test.ts +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.test.ts @@ -287,6 +287,19 @@ describe("payload helpers", () => { expect(messageText(image)).toBe(image); expect(messageText("[not json")).toBe("[not json"); }); + + it("reads GenAI message parts and native content arrays without crashing previews", () => { + const question = "What is an agent trace?"; + const parts = [{ type: "text", content: question }]; + const input = JSON.stringify([{ role: "user", parts }]); + expect(parseMessages(input)).toEqual([{ role: "user", parts, content: question }]); + expect(previewText(input)).toBe(question); + expect( + parseMessages(JSON.stringify({ role: "assistant", content: [{ type: "text", text: "An execution record" }] })), + ).toEqual([{ role: "assistant", content: "An execution record" }]); + expect(parseMessages('[{"role":"assistant","tool_calls":[]}]')).toBeNull(); + expect(parseMessages('[{"role":"user","content":42}]')).toBeNull(); + }); }); describe("treeGuides", () => { diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.ts b/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.ts index 00d24d4eaef..4379255c430 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.ts +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/traceUtils.ts @@ -9,6 +9,9 @@ import type { Span, TraceMessage, TraceSummary } from "./traceTypes"; /* Formatting */ /* ------------------------------------------------------------------ */ +export const traceAgentNames = (trace: TraceSummary): readonly string[] => + trace.agent_names ?? (trace.service ? [trace.service] : []); + export const fmtMs = (ms: number): string => { if (ms >= 60_000) return `${(ms / 60_000).toFixed(1)}m`; if (ms >= 1000) return `${(ms / 1000).toFixed(2)}s`; @@ -267,14 +270,10 @@ export const parseJson = (value: string): unknown => { } }; -const isMessage = (value: unknown): value is TraceMessage => { - const isObject = typeof value === "object" && value !== null; - return isObject && "role" in value && typeof (value as TraceMessage).role === "string"; -}; - const blockText = (block: unknown): string | null => { if (typeof block !== "object" || block === null) return null; - const text: unknown = Reflect.get(block, "text"); + const text: unknown = + Reflect.get(block, "text") ?? (Reflect.get(block, "type") === "text" ? Reflect.get(block, "content") : undefined); return typeof text === "string" ? text : null; }; @@ -296,13 +295,19 @@ export function messageText(content: string): string { .join("\n\n"); } -const withText = (message: TraceMessage): TraceMessage => ({ ...message, content: messageText(message.content) }); +const parseMessage = (value: unknown): TraceMessage | null => { + if (typeof value !== "object" || value === null) return null; + const role: unknown = Reflect.get(value, "role"); + const content: unknown = Reflect.get(value, "content") ?? Reflect.get(value, "parts"); + if (typeof role !== "string" || (typeof content !== "string" && !Array.isArray(content))) return null; + return { ...value, role, content: messageText(typeof content === "string" ? content : JSON.stringify(content)) }; +}; /** An llm span's input (array of messages) or output (one message); null when it isn't one. */ export function parseMessages(value: string): TraceMessage[] | null { const parsed = parseJson(value); - if (Array.isArray(parsed)) return parsed.every(isMessage) ? parsed.map(withText) : null; - return isMessage(parsed) ? [withText(parsed)] : null; + const messages = (Array.isArray(parsed) ? parsed : [parsed]).map(parseMessage); + return messages.every((message) => message !== null) ? messages : null; } /** Pretty JSON when the payload is JSON, else the raw string. */ diff --git a/ui/litellm-dashboard/src/lib/autorouter_presets.test.ts b/ui/litellm-dashboard/src/lib/autorouter_presets.test.ts index c67befe8fa5..85dcd7df716 100644 --- a/ui/litellm-dashboard/src/lib/autorouter_presets.test.ts +++ b/ui/litellm-dashboard/src/lib/autorouter_presets.test.ts @@ -859,23 +859,28 @@ describe("autorouter_presets", () => { }); describe("buildPresetPrefill", () => { - it("preserves JEV settings and drops inactive classifier settings when prefilling", () => { + it("preserves Laya settings and drops inactive classifier settings when prefilling", () => { const config = { tiers: { SIMPLE: ["fast"], MEDIUM: [], COMPLEX: [], REASONING: [] }, - classifier_type: "jev" as const, + classifier_type: "oss_classifier" as const, classification_mode: "every_request" as const, session_affinity: false, deployment_affinity: true, modality_routing: false, modality_pin_override: false, - jev_classifier_config: { model: "jev-test", timeout_ms: 4000, circuit_breaker_enabled: false }, + opensource_classifier_config: { + provider: "laya" as const, + model: "english", + timeout_ms: 4000, + circuit_breaker_enabled: false, + }, classifier_llm_config: { model: "stale-judge", timeout_ms: 6000 }, classifier_context_window_size: 6, }; const prefill = buildPresetPrefill(config, groupsOnly(["fast"])); const expectedJevConfig = { classifier_type: "jev", - jev_classifier_config: config.jev_classifier_config, + jev_classifier_config: config.opensource_classifier_config, classifier_context_window_size: 6, classifier_llm_config: undefined, }; diff --git a/ui/litellm-dashboard/src/lib/autorouter_presets.ts b/ui/litellm-dashboard/src/lib/autorouter_presets.ts index bbbf151e49b..ac3894f9d8d 100644 --- a/ui/litellm-dashboard/src/lib/autorouter_presets.ts +++ b/ui/litellm-dashboard/src/lib/autorouter_presets.ts @@ -1,3 +1,4 @@ +import { hydrateOssClassifier } from "@/components/add_model/jev_classifier_config"; import { ComplexityRouterConfigPayload, hydrateTierLabels, @@ -302,6 +303,7 @@ export const buildPresetPrefill = ( config: ComplexityRouterConfigPayload, availability: ModelAvailability, ): PresetPrefill => { + const classifier = hydrateOssClassifier(config); const resolve = (model: string): string => resolveAvailableModel(model, availability) ?? model; const resolveTier = (models: string[]): string[] => models.map(resolve); // Params key on the model name the preset spells while every tier entry is rewritten to the @@ -334,11 +336,10 @@ export const buildPresetPrefill = ( }, tier_model_params: resolveParamKeys(hydrateTierModelParams(config.tiers, config.tier_model_configs)), tier_labels: hydrateTierLabels(config.tier_labels), - classifier_type: config.classifier_type, + ...classifier, heuristic_v2_success_threshold: config.heuristic_v2_success_threshold, - jev_classifier_config: config.classifier_type === "jev" ? config.jev_classifier_config : undefined, classifier_llm_config: - config.classifier_type !== "jev" && config.classifier_llm_config + classifier.classifier_type !== "jev" && config.classifier_llm_config ? { ...config.classifier_llm_config, model: resolve(config.classifier_llm_config.model) } : undefined, classifier_context_window_size: config.classifier_context_window_size, diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 43098982a5e..a376f8d7910 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -1350,9 +1350,10 @@ export interface paths { * * Reads session rollups folded once per request at spend-write time, so this endpoint * never scans LiteLLM_SpendLogs. A user filter selects only turns attributed to that - * internal user when written; older key-only history remains outside user views. A session - * is in the window when it overlaps it: its last turn is on or after start_date and its first turn is on or before - * end_date. Overall hit rate is over telemetry-bearing turns; each bucket's hit rate is + * internal user when written; older key-only history remains outside user views. Money counts + * only requests on the selected UTC days, and the all-router savings headline is the same daily + * total the Overall view reads. Session shape and caching cover every session that overlaps the + * window, whole. Overall hit rate is over telemetry-bearing turns; each bucket's hit rate is * over that bucket's turns. * * The rollup supplies the measures, never the list. Which routers appear comes from the @@ -8830,6 +8831,23 @@ export interface paths { patch: operations["langfuse_proxy_route_langfuse__endpoint__patch"]; trace?: never; }; + "/laya/v1/systemone": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Laya Proxy Route */ + post: operations["laya_proxy_route_laya_v1_systemone_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/lazy/warm/{name}": { parameters: { query?: never; @@ -26152,12 +26170,21 @@ export interface components { * @description One auto-router's slice of the benchmarks. */ AutoRouterBenchmarkGroup: { - /** Avg Session Seconds */ - avg_session_seconds: number; - /** Avg Tokens Per Session */ - avg_tokens_per_session: number; - /** Avg Turns Per Session */ - avg_turns_per_session: number; + /** + * Avg Session Seconds + * @description Lifetime seconds per overlapping session; null as above + */ + avg_session_seconds: number | null; + /** + * Avg Tokens Per Session + * @description Lifetime tokens per overlapping session; null as above + */ + avg_tokens_per_session: number | null; + /** + * Avg Turns Per Session + * @description Lifetime turns per overlapping session; null when the window has routed requests but no session rows for this router type, such as an alias whose router type changed mid-session + */ + avg_turns_per_session: number | null; /** * Baseline Spend * @description Estimated single-model cost: compared actual spend plus recorded savings; null when traffic has no recorded savings @@ -26184,14 +26211,9 @@ export interface components { * @description Recorded savings over baseline_spend, as a percentage */ saved_pct: number | null; - /** - * Saved Per Session - * @description Recorded savings per session, including historical estimates - */ - saved_per_session: number | null; /** * Saved Spend - * @description Recorded historical savings plus newer estimates; null when traffic has no recorded savings estimates + * @description Recorded savings on the selected UTC days; null when traffic has no recorded savings estimates. On totals this is the same daily figure the Overall savings view reports */ saved_spend: number | null; /** @@ -26209,11 +26231,14 @@ export interface components { * @description Requests compared against the baseline: every request on complexity routers that recorded savings */ savings_estimated_turns: number; - /** Sessions */ + /** + * Sessions + * @description Sessions overlapping the window, counted whole + */ sessions: number; /** * Spend - * @description What the routed traffic actually cost + * @description What the selected days' routed traffic actually cost */ spend: number; /** @@ -26223,20 +26248,38 @@ export interface components { tier_turns?: { [key: string]: number; }; - /** Turns */ + /** + * Turns + * @description Auto-routed requests on the selected UTC days + */ turns: number; + /** + * Unattributed Saved Spend + * @description Part of saved_spend no router's daily rows account for, such as history recorded before per-router daily tracking; when set, baseline_spend and saved_pct are null + */ + unattributed_saved_spend?: number | null; }; /** * AutoRouterBenchmarkTotals - * @description Session-shape and savings aggregates over auto-routed traffic in the window. + * @description Auto-routed traffic in the window. Turns, spend and savings count requests on the selected UTC days; + * the session averages and cache stats describe every session overlapping the window, whole. */ AutoRouterBenchmarkTotals: { - /** Avg Session Seconds */ - avg_session_seconds: number; - /** Avg Tokens Per Session */ - avg_tokens_per_session: number; - /** Avg Turns Per Session */ - avg_turns_per_session: number; + /** + * Avg Session Seconds + * @description Lifetime seconds per overlapping session; null as above + */ + avg_session_seconds: number | null; + /** + * Avg Tokens Per Session + * @description Lifetime tokens per overlapping session; null as above + */ + avg_tokens_per_session: number | null; + /** + * Avg Turns Per Session + * @description Lifetime turns per overlapping session; null when the window has routed requests but no session rows for this router type, such as an alias whose router type changed mid-session + */ + avg_turns_per_session: number | null; /** * Baseline Spend * @description Estimated single-model cost: compared actual spend plus recorded savings; null when traffic has no recorded savings @@ -26253,14 +26296,9 @@ export interface components { * @description Recorded savings over baseline_spend, as a percentage */ saved_pct: number | null; - /** - * Saved Per Session - * @description Recorded savings per session, including historical estimates - */ - saved_per_session: number | null; /** * Saved Spend - * @description Recorded historical savings plus newer estimates; null when traffic has no recorded savings estimates + * @description Recorded savings on the selected UTC days; null when traffic has no recorded savings estimates. On totals this is the same daily figure the Overall savings view reports */ saved_spend: number | null; /** @@ -26278,19 +26316,30 @@ export interface components { * @description Requests compared against the baseline: every request on complexity routers that recorded savings */ savings_estimated_turns: number; - /** Sessions */ + /** + * Sessions + * @description Sessions overlapping the window, counted whole + */ sessions: number; /** * Spend - * @description What the routed traffic actually cost + * @description What the selected days' routed traffic actually cost */ spend: number; - /** Turns */ + /** + * Turns + * @description Auto-routed requests on the selected UTC days + */ turns: number; + /** + * Unattributed Saved Spend + * @description Part of saved_spend no router's daily rows account for, such as history recorded before per-router daily tracking; when set, baseline_spend and saved_pct are null + */ + unattributed_saved_spend?: number | null; }; /** * AutoRouterBenchmarksResponse - * @description Benchmarks for the auto-router dashboard, aggregated from the per-session rollup. + * @description Benchmarks for the auto-router dashboard, aggregated from the per-session and per-day rollups. */ AutoRouterBenchmarksResponse: { /** @@ -33152,44 +33201,6 @@ export interface components { /** Updated By */ updated_by?: string | null; }; - /** JevClassifierConfig */ - JevClassifierConfig: { - /** - * Api Base - * @description TypeSafe API base, falling back to TYPESAFE_API_BASE and then https://api.typesafe.ai - */ - api_base?: string | null; - /** - * Api Key - * @description TypeSafe API key, falling back to TYPESAFE_API_KEY - */ - api_key?: string | null; - /** - * Circuit Breaker Cooldown Seconds - * @default 30 - */ - circuit_breaker_cooldown_seconds: number; - /** - * Circuit Breaker Enabled - * @default true - */ - circuit_breaker_enabled: boolean; - /** - * Instructions - * @description Replaces the built-in Jev question instructions - */ - instructions?: string | null; - /** - * Model - * @default jev-latest - */ - model: string; - /** - * Timeout Ms - * @default 3000 - */ - timeout_ms: number; - }; /** Job */ Job: { /** @@ -39134,6 +39145,50 @@ export interface components { */ type: "openIdConnect"; }; + /** OpenSourceClassifierConfig */ + OpenSourceClassifierConfig: { + /** + * Api Base + * @description Provider API base; defaults to TYPESAFE_API_BASE or LAYA_API_BASE for the selected provider + */ + api_base?: string | null; + /** + * Api Key + * @description Provider API key; optional for self-hosted Laya + */ + api_key?: string | null; + /** + * Circuit Breaker Cooldown Seconds + * @default 30 + */ + circuit_breaker_cooldown_seconds: number; + /** + * Circuit Breaker Enabled + * @default true + */ + circuit_breaker_enabled: boolean; + /** + * Instructions + * @description Replaces the built-in Jev question instructions + */ + instructions?: string | null; + /** + * Model + * @default jev-latest + */ + model: string; + /** + * Provider + * @default jev + * @enum {string} + */ + provider: "jev" | "laya"; + /** + * Timeout Ms + * @default 3000 + */ + timeout_ms: number; + }; /** * OperationCreateFile * @description Instruction describing how to create a file via the apply_patch tool. @@ -41934,11 +41989,11 @@ export interface components { classifier_plugin_timeout_ms: number; /** * Classifier Type - * @description Classification strategy: local regex/keyword scoring, the bundled trained four-tier heuristic, 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 + * @description Classification strategy: local regex/keyword scoring, the bundled trained four-tier heuristic, 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 'oss_classifier', a structured choice call using Jev or Laya * @default heuristic * @enum {string} */ - classifier_type: "heuristic" | "heuristic_v2" | "llm" | "capability" | "llm_v2" | "custom" | "heuristic_first" | "hybrid" | "jev"; + classifier_type: "heuristic" | "heuristic_v2" | "llm" | "capability" | "llm_v2" | "custom" | "heuristic_first" | "hybrid" | "oss_classifier"; /** * Code Keywords * @description Keywords indicating code-related content @@ -42037,7 +42092,6 @@ export interface components { * @description How close to a tier boundary a heuristic score has to land before the LLM classifier breaks the tie; required when classifier_type is 'hybrid' and rejected otherwise. Everything further than this from every active boundary routes on the scorer's own tier with no classifier call, at any tier, which is what separates 'hybrid' from 'heuristic_first' and its cheap-tier ceiling. A prompt where no dimension fired still goes to the classifier, since the scorer has no opinion to be near a boundary with. 0 escalates only scores sitting exactly on a boundary. */ hybrid_boundary_margin?: number | null; - jev_classifier_config?: components["schemas"]["JevClassifierConfig"] | null; /** * Keyword Tier Rules * @description Rules that force a specific tier when their keywords match the prompt @@ -42069,6 +42123,7 @@ export interface components { * @default false */ modality_routing: boolean; + opensource_classifier_config?: components["schemas"]["OpenSourceClassifierConfig"] | null; /** * Plan Mode Min Tier * @description When set, requests carrying a coding-agent plan-mode sentinel (Claude Code plan mode, VS Code Copilot Plan mode, Copilot CLI's exit_plan_mode tool) are routed to at least this tier: the classified tier still wins when it is higher, and the floor also overrides a session-affinity pin to a lower tier for exactly the turns carrying the sentinel, without rewriting the pin -- the first turn after plan mode exits routes as if plan mode had never happened. Names a built-in tier, or with tier_definitions set, one of the defined tier names (list order is ascending severity, same as keyword_tier_rules). Unset disables detection entirely. The sentinels ride in client-injected prompt text, so a caller who pastes one can spend up to this tier's models -- never down, and never outside the configured pools. @@ -42166,7 +42221,7 @@ export interface components { }; /** * Tier Definitions - * @description Operator-defined tier set replacing the built-in SIMPLE/MEDIUM/COMPLEX/REASONING. Each entry's name becomes a value the LLM classifier can return and its description 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, 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. + * @description Operator-defined tier set replacing the built-in SIMPLE/MEDIUM/COMPLEX/REASONING. Each entry's name becomes a value the LLM classifier can return and its description 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', '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. */ tier_definitions?: components["schemas"]["TierDefinition"][] | null; /** @@ -46849,6 +46904,8 @@ export interface components { agent_count: number; /** Agent Invocations */ agent_invocations: number; + /** Agent Names */ + agent_names?: string[]; /** Duration Ms */ duration_ms: number; /** Error Count */ @@ -61639,6 +61696,26 @@ export interface operations { }; }; }; + laya_proxy_route_laya_v1_systemone_post: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + }; + }; warm_lazy_warm__name__post: { parameters: { query?: never;