diff --git a/helm/litellm-helm/templates/migrations-job.yaml b/helm/litellm-helm/templates/migrations-job.yaml index 5a873cbb965..b199ad2d8f3 100644 --- a/helm/litellm-helm/templates/migrations-job.yaml +++ b/helm/litellm-helm/templates/migrations-job.yaml @@ -6,6 +6,9 @@ metadata: name: {{ include "litellm.fullname" . }}-migrations labels: {{- include "litellm.labels" . | nindent 4 }} + {{- with .Values.migrationJob.jobLabels }} + {{- toYaml . | nindent 4 }} + {{- end }} annotations: {{- if .Values.migrationJob.hooks.argocd.enabled }} argocd.argoproj.io/hook: PreSync @@ -17,6 +20,9 @@ metadata: helm.sh/hook-weight: {{ .Values.migrationJob.hooks.helm.weight | default "1" | quote }} {{- end }} checksum/config: {{ toYaml .Values | sha256sum }} + {{- with .Values.migrationJob.jobAnnotations }} + {{- toYaml . | nindent 4 }} + {{- end }} spec: template: metadata: @@ -25,6 +31,9 @@ spec: {{- with .Values.podLabels }} {{- toYaml . | nindent 8 }} {{- end }} + {{- with .Values.migrationJob.podLabels }} + {{- toYaml . | nindent 8 }} + {{- end }} annotations: {{- with .Values.migrationJob.annotations }} {{- toYaml . | nindent 8 }} @@ -47,7 +56,16 @@ spec: imagePullPolicy: {{ .Values.image.pullPolicy }} securityContext: {{- toYaml .Values.securityContext | nindent 12 }} + {{- if .Values.migrationJob.command }} + command: {{ toYaml .Values.migrationJob.command | nindent 12 }} + {{- else }} command: ["python", "litellm/proxy/prisma_migration.py"] + {{- end }} + {{- if .Values.migrationJob.args }} + args: {{ toYaml .Values.migrationJob.args | nindent 12 }} + {{- else if .Values.migrationJob.command }} + args: [] + {{- end }} workingDir: "/app" env: {{- if .Values.db.useExisting }} diff --git a/helm/litellm-helm/tests/migrations-job_tests.yaml b/helm/litellm-helm/tests/migrations-job_tests.yaml index dd4276ac60f..bd5dd457f95 100644 --- a/helm/litellm-helm/tests/migrations-job_tests.yaml +++ b/helm/litellm-helm/tests/migrations-job_tests.yaml @@ -360,3 +360,73 @@ tests: asserts: - notExists: path: spec.activeDeadlineSeconds + + - it: should set custom jobLabels and jobAnnotations on Job metadata + template: migrations-job.yaml + set: + migrationJob: + enabled: true + jobLabels: + environment: production + team: platform + jobAnnotations: + example.com/cost-center: "1234" + asserts: + - equal: + path: metadata.labels.environment + value: production + - equal: + path: metadata.labels.team + value: platform + - equal: + path: metadata.annotations['example.com/cost-center'] + value: "1234" + + - it: should set custom podLabels on Pod template + template: migrations-job.yaml + set: + migrationJob: + enabled: true + podLabels: + custom.io/pod-role: migration + asserts: + - equal: + path: spec.template.metadata.labels['custom.io/pod-role'] + value: migration + + - it: should override container command and args + template: migrations-job.yaml + set: + migrationJob: + enabled: true + command: + - sh + args: + - -c + - echo migrating + asserts: + - equal: + path: spec.template.spec.containers[0].command + value: + - sh + - equal: + path: spec.template.spec.containers[0].args + value: + - -c + - echo migrating + + - it: should clear container args when only command is specified + template: migrations-job.yaml + set: + migrationJob: + enabled: true + command: + - sh + asserts: + - equal: + path: spec.template.spec.containers[0].command + value: + - sh + - equal: + path: spec.template.spec.containers[0].args + value: [] diff --git a/helm/litellm-helm/values.yaml b/helm/litellm-helm/values.yaml index c53fd5b0315..2e2b73ddc9a 100644 --- a/helm/litellm-helm/values.yaml +++ b/helm/litellm-helm/values.yaml @@ -572,6 +572,11 @@ migrationJob: # In that case, pre-install/pre-upgrade hooks run before normal resources, so this defaults to "default". serviceAccountName: "" annotations: {} + jobLabels: {} # Custom labels for the Job metadata + jobAnnotations: {} # Custom annotations for the Job metadata + podLabels: {} # Custom labels for the Job pod template + command: [] # Override container command (defaults to ["python", "litellm/proxy/prisma_migration.py"]) + args: [] # Override container args ttlSecondsAfterFinished: 120 resources: {} # Unset by default. This job runs the database migration and exits, so it does not diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008000000_lens_due_at/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008000000_lens_due_at/migration.sql new file mode 100644 index 00000000000..676d9b124c9 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008000000_lens_due_at/migration.sql @@ -0,0 +1 @@ +ALTER TABLE "LiteLLM_Lens" ADD COLUMN IF NOT EXISTS "due_at" TIMESTAMP(3) DEFAULT '1970-01-01 00:00:00'; diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008000100_lens_due_at_index/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008000100_lens_due_at_index/migration.sql new file mode 100644 index 00000000000..36311c4e749 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008000100_lens_due_at_index/migration.sql @@ -0,0 +1,2 @@ +CREATE INDEX CONCURRENTLY IF NOT EXISTS "LiteLLM_Lens_due_at_idx" +ON "LiteLLM_Lens" ("due_at"); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 28061792187..e4df567111a 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -1939,6 +1939,9 @@ model LiteLLM_Lens { id String @id version Int @default(0) data Json + due_at DateTime? @default(dbgenerated("'1970-01-01 00:00:00'::timestamp without time zone")) + + @@index([due_at], map: "LiteLLM_Lens_due_at_idx") } model LiteLLM_LensRun { diff --git a/litellm-rust/crates/lens/src/ingest.rs b/litellm-rust/crates/lens/src/ingest.rs index db9323d7855..e09c61a2a43 100644 --- a/litellm-rust/crates/lens/src/ingest.rs +++ b/litellm-rust/crates/lens/src/ingest.rs @@ -78,7 +78,7 @@ pub fn response(content_type: Option<&str>, outcome: Result<(), Error>) -> Respo ) }; let mut response = (status, [(http::header::CONTENT_TYPE, media_type)], body).into_response(); - if status == StatusCode::SERVICE_UNAVAILABLE { + if matches!(status, StatusCode::SERVICE_UNAVAILABLE | StatusCode::TOO_MANY_REQUESTS) { response .headers_mut() .insert("retry-after", http::HeaderValue::from_static("5")); diff --git a/litellm-rust/crates/storage-clickhouse/src/read.rs b/litellm-rust/crates/storage-clickhouse/src/read.rs index c4bfdef393a..03a5522e595 100644 --- a/litellm-rust/crates/storage-clickhouse/src/read.rs +++ b/litellm-rust/crates/storage-clickhouse/src/read.rs @@ -61,6 +61,16 @@ pub async fn execute_read( connection: &Connection, sql: &str, parameters: &BTreeMap, +) -> Result { + execute_read_with_limits(client, connection, sql, parameters, READ_LIMITS).await +} + +async fn execute_read_with_limits( + client: &Client, + connection: &Connection, + sql: &str, + parameters: &BTreeMap, + limits: ReadLimits, ) -> Result { if sql.trim().is_empty() { return Err(Error::EmptySql); @@ -89,12 +99,9 @@ pub async fn execute_read( .clear() .extend_pairs(existing_pairs) .append_pair("readonly", "1") - .append_pair("max_result_rows", &READ_LIMITS.result_rows.to_string()) + .append_pair("max_result_rows", &limits.result_rows.to_string()) .append_pair("result_overflow_mode", "throw") - .append_pair( - "max_execution_time", - &READ_LIMITS.execution_seconds.to_string(), - ) + .append_pair("max_execution_time", &limits.execution_seconds.to_string()) .append_pair("wait_end_of_query", "1") .append_pair("default_format", "JSON"); @@ -122,7 +129,7 @@ pub async fn execute_read( let mut body = Vec::new(); while let Some(chunk) = response.chunk().await.map_err(|_| Error::Transport)? { - if body.len() + chunk.len() > READ_LIMITS.response_bytes { + if body.len() + chunk.len() > limits.response_bytes { return Err(Error::ResponseTooLarge); } body.extend_from_slice(&chunk); @@ -141,6 +148,7 @@ pub trait Query { type Params: Serialize; type Row: DeserializeOwned; + const READ_LIMITS: ReadLimits = crate::read::READ_LIMITS; const SQL: &'static str; } @@ -159,7 +167,14 @@ pub async fn fetch( connection: &Connection, params: &Q::Params, ) -> Result, Error> { - let body = execute_read(client, connection, Q::SQL, ¶meters(params)?).await?; + let body = execute_read_with_limits( + client, + connection, + Q::SQL, + ¶meters(params)?, + Q::READ_LIMITS, + ) + .await?; decode_rows::(&body) } @@ -168,7 +183,14 @@ pub async fn fetch_json( connection: &Connection, params: &Q::Params, ) -> Result { - let body = execute_read(client, connection, Q::SQL, ¶meters(params)?).await?; + let body = execute_read_with_limits( + client, + connection, + Q::SQL, + ¶meters(params)?, + Q::READ_LIMITS, + ) + .await?; decode_rows::(&body)?; Ok(body) } diff --git a/litellm-rust/crates/traces-clickhouse/query/lens_sample.sql b/litellm-rust/crates/traces-clickhouse/query/lens_sample.sql index 92086c33c13..5df0c8a1145 100644 --- a/litellm-rust/crates/traces-clickhouse/query/lens_sample.sql +++ b/litellm-rust/crates/traces-clickhouse/query/lens_sample.sql @@ -18,10 +18,14 @@ SELECT *, selection_key FROM ( WHERE {source:String} IN ('traces','both') AND ({all_teams:UInt8}=1 OR TeamId={team:String}) AND ({key_hash:String}='' OR ApiKeyHash={key_hash:String}) + -- The 7 day slack covers spans that started before the window and late ingestion + AND Timestamp >= fromUnixTimestamp64Milli(toInt64({start:UInt64})) - INTERVAL 7 DAY AND (TeamId,ApiKeyHash,TraceId) IN ( SELECT TeamId,ApiKeyHash,TraceId FROM otel_traces WHERE ({all_teams:UInt8}=1 OR TeamId={team:String}) AND ({key_hash:String}='' OR ApiKeyHash={key_hash:String}) + AND Timestamp >= fromUnixTimestamp64Milli(toInt64({start:UInt64})) - INTERVAL 7 DAY + AND Timestamp < fromUnixTimestamp64Milli(toInt64({end:UInt64})) AND if(EngineReceivedMs>0,toInt64(EngineReceivedMs), toUnixTimestamp64Milli(Timestamp)+toInt64(intDiv(Duration,1000000))) >= {start:UInt64} ) @@ -42,6 +46,8 @@ SELECT *, selection_key FROM ( WHERE {source:String} IN ('requests','both') AND ({all_teams:UInt8}=1 OR team_id={team:String}) AND ({key_hash:String}='' OR api_key={key_hash:String}) + AND spend_logs.start_time >= fromUnixTimestamp64Milli(toInt64({start:UInt64})) - INTERVAL 7 DAY + AND spend_logs.start_time < fromUnixTimestamp64Milli(toInt64({end:UInt64})) AND if(EngineReceivedMs>0,toInt64(EngineReceivedMs),toUnixTimestamp64Milli(end_time)) >= {start:UInt64} AND EngineReceivedMs < {end:UInt64} AND toUnixTimestamp64Milli(end_time) < {end:UInt64} @@ -55,6 +61,8 @@ SELECT *, selection_key FROM ( SELECT TeamId,ApiKeyHash,LiteLLMRequestId FROM otel_traces WHERE ({all_teams:UInt8}=1 OR TeamId={team:String}) AND ({key_hash:String}='' OR ApiKeyHash={key_hash:String}) AND LiteLLMRequestId!='' + AND Timestamp >= fromUnixTimestamp64Milli(toInt64({start:UInt64})) - INTERVAL 7 DAY + AND Timestamp < fromUnixTimestamp64Milli(toInt64({end:UInt64})) )) ) WHERE ({selected_team:String}='' OR team_id={selected_team:String}) diff --git a/litellm-rust/crates/traces-clickhouse/src/query/lens.rs b/litellm-rust/crates/traces-clickhouse/src/query/lens.rs index 622a014599e..6439e696dad 100644 --- a/litellm-rust/crates/traces-clickhouse/src/query/lens.rs +++ b/litellm-rust/crates/traces-clickhouse/src/query/lens.rs @@ -1,4 +1,10 @@ -use litellm_storage_clickhouse::Query; +use litellm_storage_clickhouse::{Query, ReadLimits}; + +const SAMPLE_READ_LIMITS: ReadLimits = ReadLimits { + result_rows: 10_000, + response_bytes: 16 * 1024 * 1024, + ..litellm_storage_clickhouse::READ_LIMITS +}; pub const LENS_QUERIES: [litellm_traces::ReadQuery; 5] = [ litellm_traces::ReadQuery::Availability, @@ -188,6 +194,7 @@ impl Query for LensSample { type Params = LensSampleParams; type Row = LensSampleRow; + const READ_LIMITS: ReadLimits = SAMPLE_READ_LIMITS; const SQL: &'static str = include_str!("../../query/lens_sample.sql"); } diff --git a/litellm-rust/crates/traces-clickhouse/tests/load.rs b/litellm-rust/crates/traces-clickhouse/tests/load.rs new file mode 100644 index 00000000000..07c1095dfc3 --- /dev/null +++ b/litellm-rust/crates/traces-clickhouse/tests/load.rs @@ -0,0 +1,146 @@ +use std::collections::BTreeMap; + +use litellm_storage_clickhouse::READ_LIMITS; +use litellm_traces_clickhouse::{Connection, Parameter, ReadQuery, execute_named_read}; +use rstest::rstest; +use serde_json::Value; + +#[path = "queries/support.rs"] +#[expect( + dead_code, + reason = "load tests share the query fixture but do not read through QueryReaders" +)] +mod fixtures; +mod support; + +use fixtures::{DATABASE, SeededDatabase, migrated_database}; +use support::TestResult; + +const SPANS_PER_DAY: u64 = 2_000; + +async fn seed_days(fixture: &SeededDatabase, first_day: u64, days: u64) -> TestResult { + let count = SPANS_PER_DAY * days; + let first_row = SPANS_PER_DAY * first_day; + let query = format!( + "INSERT INTO {DATABASE}.otel_traces \ + (Timestamp, TraceId, SpanId, ParentSpanId, SpanName, ServiceName, ObservationType, TeamId, ApiKeyHash, Duration, SpanAttributes) \ + SELECT now64(9) - toIntervalHour(intDiv(number, {SPANS_PER_DAY}) * 24 + 12 + {first_day} * 24), \ + concat('load-', toString(number + {first_row})), concat('span-', toString(number + {first_row})), \ + '', 'span', 'service', 'agent', 'load-team', '', 0, \ + if({first_day}=0 AND number < {SPANS_PER_DAY}, map('payload', repeat('x', 3000)), map()) \ + FROM numbers({count})" + ); + fixture + .database + .client + .post(&fixture.database.url) + .body(query) + .send() + .await? + .error_for_status()?; + Ok(()) +} + +fn sample_parameters(start: u64, end: u64) -> BTreeMap { + BTreeMap::from([ + ("source".into(), Parameter::Text("traces".into())), + ("all_teams".into(), Parameter::Integer(0)), + ("team".into(), Parameter::Text("load-team".into())), + ("key_hash".into(), Parameter::Text(String::new())), + ("start".into(), Parameter::Unsigned(start)), + ("end".into(), Parameter::Unsigned(end)), + ("agent_name".into(), Parameter::Text(String::new())), + ("service".into(), Parameter::Text(String::new())), + ("filter_keys".into(), Parameter::Strings(Vec::new())), + ("filter_values".into(), Parameter::Strings(Vec::new())), + ("selected_team".into(), Parameter::Text(String::new())), + ("execution_ids".into(), Parameter::Strings(Vec::new())), + ("sample_cap".into(), Parameter::Unsigned(0)), + ("sample_percent".into(), Parameter::Integer(100)), + ("preview".into(), Parameter::Integer(0)), + ("after".into(), Parameter::Text(String::new())), + ("limit".into(), Parameter::Unsigned(10_000)), + ("offset".into(), Parameter::Unsigned(0)), + ]) +} + +async fn sample( + fixture: &SeededDatabase, + start: u64, + end: u64, + query_id: &str, +) -> TestResult<(usize, usize)> { + let mut url = Connection::configured(&fixture.database.url, DATABASE, "default", "")? + .url() + .clone(); + url.query_pairs_mut().append_pair("query_id", query_id); + let connection = Connection::parse(url.as_str())?; + let response = execute_named_read( + &fixture.database.client, + &connection, + ReadQuery::Sample, + &sample_parameters(start, end), + ) + .await?; + let result: Value = serde_json::from_str(&response)?; + Ok(( + result["data"].as_array().ok_or("sample rows")?.len(), + response.len(), + )) +} + +async fn query_read_rows(fixture: &SeededDatabase, query_id: &str) -> TestResult { + fixture + .database + .client + .post(&fixture.database.url) + .body("SYSTEM FLUSH LOGS") + .send() + .await? + .error_for_status()?; + let response = fixture + .database + .client + .post(&fixture.database.url) + .body(format!( + "SELECT read_rows FROM system.query_log WHERE type = 'QueryFinish' \ + AND query_id = '{query_id}' ORDER BY event_time DESC LIMIT 1 FORMAT JSON" + )) + .send() + .await? + .error_for_status()? + .text() + .await?; + let result: Value = serde_json::from_str(&response)?; + result["data"][0]["read_rows"] + .as_u64() + .ok_or_else(|| "query log read_rows missing".into()) +} + +#[rstest] +#[tokio::test] +async fn lens_sample_reads_scale_with_window_not_retention( + #[future(awt)] migrated_database: TestResult, +) -> TestResult { + let fixture = migrated_database?; + seed_days(&fixture, 0, 8).await?; + let now_ms = time::OffsetDateTime::now_utc().unix_timestamp() as u64 * 1000; + let start = now_ms - 86_400_000; + let end = now_ms + 60_000; + let before_id = format!("lens_sample_before_{}", std::process::id()); + let (before_rows, response_bytes) = sample(&fixture, start, end, &before_id).await?; + assert_eq!(before_rows, SPANS_PER_DAY as usize); + assert!(response_bytes > READ_LIMITS.response_bytes); + let before = query_read_rows(&fixture, &before_id).await?; + + seed_days(&fixture, 8, 24).await?; + let after_id = format!("lens_sample_after_{}", std::process::id()); + let (after_rows, _) = sample(&fixture, start, end, &after_id).await?; + assert_eq!(after_rows, SPANS_PER_DAY as usize); + let after = query_read_rows(&fixture, &after_id).await?; + assert!( + after * 100 <= before * 105, + "read_rows grew from {before} to {after}" + ); + Ok(()) +} diff --git a/litellm-rust/crates/traces-clickhouse/tests/queries.rs b/litellm-rust/crates/traces-clickhouse/tests/queries.rs index 5cb0ee4bfcd..6e63adfc347 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/queries.rs +++ b/litellm-rust/crates/traces-clickhouse/tests/queries.rs @@ -3,7 +3,8 @@ use std::collections::BTreeMap; use litellm_storage_clickhouse::fetch; use litellm_traces::query::named as contracts; use litellm_traces_clickhouse::{ - QueryScope, + Connection, InsertTable, Parameter, QueryScope, ReadQuery, execute_named_read, execute_read, + insert_rows, query::named::{ListTraces, ListTracesParams, TraceSpans, TraceSpansParams}, query_help, query_sql, }; @@ -18,6 +19,112 @@ mod support; use fixtures::{SeededDatabase, insert_export, migrated_database, seeded_database}; use support::TestResult; +#[rstest] +#[tokio::test] +async fn lens_sample_keeps_spans_before_window_start_and_excludes_old_only_traces( + #[future(awt)] migrated_database: TestResult, +) -> TestResult { + let fixture = migrated_database?; + let start_ms = time::OffsetDateTime::now_utc().unix_timestamp() * 1000 - 86_400_000; + let end_ms = start_ms + 86_460_000; + let rows = [ + ( + "late-root", + "trace-with-slack", + start_ms - 2 * 86_400_000, + "", + ), + ( + "in-window", + "trace-with-slack", + start_ms + 1_000, + "late-root", + ), + ("old-span", "trace-too-old", start_ms - 8 * 86_400_000, ""), + ] + .into_iter() + .map(|(span_id, trace_id, timestamp_ms, parent_span_id)| { + BTreeMap::from([ + ( + "Timestamp".into(), + serde_json::json!(timestamp_ms * 1_000_000), + ), + ("Duration".into(), serde_json::json!(1_000_000)), + ("TraceId".into(), serde_json::json!(trace_id)), + ("SpanId".into(), serde_json::json!(span_id)), + ("ParentSpanId".into(), serde_json::json!(parent_span_id)), + ("SpanName".into(), serde_json::json!(span_id)), + ("ObservationType".into(), serde_json::json!("agent")), + ("TeamId".into(), serde_json::json!("team-lens")), + ("ApiKeyHash".into(), serde_json::json!("")), + ]) + }) + .collect(); + let writer = Connection::writer(&fixture.database.url)?; + insert_rows( + &fixture.database.client, + &writer, + fixtures::DATABASE, + InsertTable::OtelTraces, + rows, + ) + .await?; + let connection = + Connection::configured(&fixture.database.url, fixtures::DATABASE, "default", "")?; + let parameters = BTreeMap::from([ + ("source".into(), Parameter::Text("traces".into())), + ("all_teams".into(), Parameter::Integer(0)), + ("team".into(), Parameter::Text("team-lens".into())), + ("key_hash".into(), Parameter::Text(String::new())), + ("start".into(), Parameter::Unsigned(start_ms as u64)), + ("end".into(), Parameter::Unsigned(end_ms as u64)), + ("agent_name".into(), Parameter::Text(String::new())), + ("service".into(), Parameter::Text(String::new())), + ("filter_keys".into(), Parameter::Strings(Vec::new())), + ("filter_values".into(), Parameter::Strings(Vec::new())), + ("selected_team".into(), Parameter::Text(String::new())), + ("execution_ids".into(), Parameter::Strings(Vec::new())), + ("sample_cap".into(), Parameter::Unsigned(0)), + ("sample_percent".into(), Parameter::Integer(100)), + ("preview".into(), Parameter::Integer(0)), + ("after".into(), Parameter::Text(String::new())), + ("limit".into(), Parameter::Unsigned(10_000)), + ("offset".into(), Parameter::Unsigned(0)), + ]); + let body = execute_named_read( + &fixture.database.client, + &connection, + ReadQuery::Sample, + ¶meters, + ) + .await?; + let result: serde_json::Value = serde_json::from_str(&body)?; + let executions = result["data"].as_array().ok_or("sample rows")?; + let trace = executions + .iter() + .find(|row| row["trace_id"] == "trace-with-slack") + .ok_or("sampled trace missing")?; + let original_start = execute_read( + &fixture.database.client, + &connection, + "SELECT toString(fromUnixTimestamp64Nano({timestamp:Int64})) AS start_time FORMAT JSON", + &BTreeMap::from([( + "timestamp".into(), + Parameter::Integer((start_ms - 2 * 86_400_000) * 1_000_000), + )]), + ) + .await?; + let original_start: serde_json::Value = serde_json::from_str(&original_start)?; + assert_eq!(trace["span_count"].as_u64(), Some(2)); + assert_eq!(trace["start_time"], original_start["data"][0]["start_time"]); + assert!( + !executions + .iter() + .any(|row| row["trace_id"] == "trace-too-old") + ); + Ok(()) +} + #[derive(Clone, Copy, strum::AsRefStr)] #[strum(serialize_all = "snake_case")] enum ScopeCase { diff --git a/litellm/integrations/anthropic_cache_control_hook.py b/litellm/integrations/anthropic_cache_control_hook.py index 893cb0d84d3..22d76242f2f 100644 --- a/litellm/integrations/anthropic_cache_control_hook.py +++ b/litellm/integrations/anthropic_cache_control_hook.py @@ -1029,7 +1029,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): ) AnthropicCacheControlHook.record_gateway_injection(kwargs, breakpoints_added) if openai_dialect and breakpoints_added > 0: - kwargs.setdefault("prompt_cache_options", PromptCacheOptions(mode="explicit")) + kwargs.setdefault("prompt_cache_options", PromptCacheOptions(mode="implicit")) if remaining: kwargs["cache_control_injection_points"] = remaining return messages, system diff --git a/litellm/litellm_core_utils/logging_utils.py b/litellm/litellm_core_utils/logging_utils.py index 5e247324cea..5ac6b5435dc 100644 --- a/litellm/litellm_core_utils/logging_utils.py +++ b/litellm/litellm_core_utils/logging_utils.py @@ -43,7 +43,7 @@ Helper utils used for logging callbacks # Regex matching data-URI base64 content: "data:;base64," # Captures: group(1)=mime_type, group(2)=base64_payload -_DATA_URI_RE: Final = re.compile(r"data:([^;]+);base64,([A-Za-z0-9+/=]+)") +_DATA_URI_RE: Final = re.compile(r"data:([^;,\s]{1,255});base64,([A-Za-z0-9+/=]+)") # Maximum nesting depth for _truncate_base64_in_value to guard against # pathological payloads. OpenAI message format is typically 3-4 levels deep. diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 62b8ce95b22..a781be610a6 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -896,8 +896,8 @@ class RealTimeStreaming: # clientContent / cancel messages are sent. if pre_block_backend_message is not None: await self._send_to_backend(pre_block_backend_message) - # Cancel any in-progress LLM response (e.g. VAD auto-response). - await self._send_to_backend(json.dumps({"type": "response.cancel"})) + if not self._is_transcription_session: + await self._send_to_backend(json.dumps({"type": "response.cancel"})) # Send the policy violation hint (shows as small gray status text in UI). await self.websocket.send_text( json.dumps( @@ -911,25 +911,26 @@ class RealTimeStreaming: } ) ) - # Ask the LLM to voice the exact guardrail message so the - # user hears it as audio in voice sessions (not just text). - guardrail_prompt = ( - f"Say exactly the following message to the user, word for word, " - f"do not add anything else: {error_msg}" - ) - await self._send_to_backend( - json.dumps( - { - "type": "conversation.item.create", - "item": { - "type": "message", - "role": "user", - "content": [{"type": "input_text", "text": guardrail_prompt}], - }, - } + if not self._is_transcription_session: + # Ask the LLM to voice the exact guardrail message so the + # user hears it as audio in voice sessions (not just text). + guardrail_prompt = ( + f"Say exactly the following message to the user, word for word, " + f"do not add anything else: {error_msg}" ) - ) - await self._send_to_backend(json.dumps({"type": "response.create"})) + await self._send_to_backend( + json.dumps( + { + "type": "conversation.item.create", + "item": { + "type": "message", + "role": "user", + "content": [{"type": "input_text", "text": guardrail_prompt}], + }, + } + ) + ) + await self._send_to_backend(json.dumps({"type": "response.create"})) self._violation_count += 1 end_session_after: int | None = getattr(callback, "end_session_after_n_fails", None) @@ -1070,18 +1071,14 @@ class RealTimeStreaming: self.store_message(event_obj) await self.websocket.send_text(self._event_to_client_json(event_obj)) - # Transcription-only sessions (e.g. gpt-realtime-whisper) have no - # assistant turn: capture audio-duration usage for cost and never - # trigger response.create. if self._is_transcription_session: self._capture_transcription_usage(event_obj) - return True blocked: Final = await self.run_realtime_guardrails( transcript, item_id=event_obj.get("item_id"), ) - if not blocked: + if not blocked and not self._is_transcription_session: await self._send_to_backend(json.dumps({"type": "response.create"})) return True return False diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index dbc1bb0de46..cecc7d0856f 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -15229,8 +15229,8 @@ "cache_creation_input_token_cost": 2.5e-06, "cache_creation_input_token_cost_above_1hr": 4e-06, "cache_creation_input_token_cost_batches": 1.25e-06, - "cache_read_input_token_cost": 2e-07, - "cache_read_input_token_cost_batches": 1e-07, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_batches": 5e-08, "input_cost_per_token": 2e-06, "input_cost_per_token_batches": 1e-06, "litellm_provider": "anthropic", @@ -15267,7 +15267,7 @@ "supports_forced_tool_use": false, "thinking_always_on": true, "prompt_cache_min_tokens": 512, - "source": "https://platform.claude.com/docs/en/models/sonnet-5-5/overview", + "source": "https://platform.claude.com/docs/en/about-claude/pricing", "supports_web_search": true }, "claude-sonnet-4-6": { @@ -80768,5 +80768,648 @@ "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true + }, +"claude-haiku-5-5": { + "supports_anthropic_compaction": true, + "cache_creation_input_token_cost": 1.25e-07, + "cache_creation_input_token_cost_above_1hr": 2e-07, + "cache_creation_input_token_cost_batches": 6.25e-08, + "cache_read_input_token_cost": 1e-08, + "cache_read_input_token_cost_batches": 5e-09, + "input_cost_per_token": 1e-07, + "input_cost_per_token_batches": 5e-08, + "litellm_provider": "anthropic", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "output_cost_per_token_batches": 2.5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_native_structured_output": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true, + "provider_specific_entry": { + "us": 1.1 + }, + "supports_output_config": true, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "source": "https://platform.claude.com/docs/en/models/haiku-5-5/overview", + "supports_web_search": true, + "input_cost_per_token_above_100k_tokens": 5e-07, + "output_cost_per_token_above_100k_tokens": 2.5e-06, + "cache_creation_input_token_cost_above_100k_tokens": 6.25e-07, + "cache_creation_input_token_cost_above_1hr_above_100k_tokens": 1e-06, + "cache_read_input_token_cost_above_100k_tokens": 5e-08 + }, + "bedrock_mantle/anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "cache_creation_input_token_cost": 1.375e-07, + "cache_creation_input_token_cost_above_1hr": 2.2e-07, + "cache_read_input_token_cost": 1.1e-08, + "input_cost_per_token": 1.1e-07, + "litellm_provider": "bedrock_mantle", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5.5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": false, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh", + "supports_parallel_tool_use_config": true, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-anthropic-claude-haiku-5-5.html" + }, + "bedrock_mantle/us-gov-west-1/anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "bedrock_output_config_effort_ceiling": "xhigh", + "cache_creation_input_token_cost": 1.5e-07, + "cache_creation_input_token_cost_above_1hr": 2.4e-07, + "cache_read_input_token_cost": 1.2e-08, + "input_cost_per_token": 1.2e-07, + "litellm_provider": "bedrock_mantle", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 6e-07, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_mid_conversation_system": true, + "supports_native_structured_output": false, + "supports_output_config": true, + "supports_parallel_tool_use_config": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-anthropic-claude-haiku-5-5.html" + }, + "anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "cache_creation_input_token_cost": 1.25e-07, + "cache_creation_input_token_cost_above_1hr": 2e-07, + "cache_read_input_token_cost": 1e-08, + "input_cost_per_token": 1e-07, + "litellm_provider": "bedrock_converse", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": false, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh", + "supports_parallel_tool_use_config": true, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + }, + "apac.anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "litellm_provider": "bedrock_converse", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": false, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh", + "supports_parallel_tool_use_config": true, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "cache_creation_input_token_cost_above_1hr": 2.2e-07, + "cache_creation_input_token_cost": 1.375e-07, + "input_cost_per_token": 1.1e-07, + "output_cost_per_token": 5.5e-07, + "cache_read_input_token_cost": 1.1e-08, + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + }, + "au.anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "cache_creation_input_token_cost": 1.375e-07, + "cache_creation_input_token_cost_above_1hr": 2.2e-07, + "cache_read_input_token_cost": 1.1e-08, + "input_cost_per_token": 1.1e-07, + "litellm_provider": "bedrock_converse", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5.5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": false, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh", + "supports_parallel_tool_use_config": true, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "source": "https://aws.amazon.com/bedrock/pricing/" + }, + "azure_ai/claude-haiku-5-5": { + "supports_mid_conversation_system": true, + "cache_creation_input_token_cost": 1.25e-07, + "cache_creation_input_token_cost_above_1hr": 2e-07, + "cache_read_input_token_cost": 1e-08, + "deprecation_date": "2027-09-29", + "input_cost_per_token": 1e-07, + "litellm_provider": "azure_ai", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true, + "supports_forced_tool_use": true, + "thinking_always_on": false, + "prompt_cache_min_tokens": 512, + "source": "https://platform.claude.com/docs/en/models/haiku-5-5/migration-guide" + }, + "bedrock/us-gov-east-1/anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "bedrock_output_config_effort_ceiling": "xhigh", + "cache_creation_input_token_cost": 1.5e-07, + "cache_creation_input_token_cost_above_1hr": 2.4e-07, + "cache_read_input_token_cost": 1.2e-08, + "input_cost_per_token": 1.2e-07, + "litellm_provider": "bedrock", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 6e-07, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_mid_conversation_system": true, + "supports_native_structured_output": false, + "supports_output_config": true, + "supports_parallel_tool_use_config": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true + }, + "bedrock/us-gov-west-1/anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "bedrock_output_config_effort_ceiling": "xhigh", + "cache_creation_input_token_cost": 1.5e-07, + "cache_creation_input_token_cost_above_1hr": 2.4e-07, + "cache_read_input_token_cost": 1.2e-08, + "input_cost_per_token": 1.2e-07, + "litellm_provider": "bedrock", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 6e-07, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_mid_conversation_system": true, + "supports_native_structured_output": false, + "supports_output_config": true, + "supports_parallel_tool_use_config": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true + }, + "eu.anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "cache_creation_input_token_cost": 1.375e-07, + "cache_creation_input_token_cost_above_1hr": 2.2e-07, + "cache_read_input_token_cost": 1.1e-08, + "input_cost_per_token": 1.1e-07, + "litellm_provider": "bedrock_converse", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5.5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": false, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh", + "supports_parallel_tool_use_config": true, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "source": "https://aws.amazon.com/bedrock/pricing/" + }, + "global.anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "cache_creation_input_token_cost": 1.25e-07, + "cache_creation_input_token_cost_above_1hr": 2e-07, + "cache_read_input_token_cost": 1e-08, + "input_cost_per_token": 1e-07, + "litellm_provider": "bedrock_converse", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": false, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh", + "supports_parallel_tool_use_config": true, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "source": "https://aws.amazon.com/bedrock/pricing/" + }, + "jp.anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "cache_creation_input_token_cost": 1.375e-07, + "cache_creation_input_token_cost_above_1hr": 2.2e-07, + "cache_read_input_token_cost": 1.1e-08, + "input_cost_per_token": 1.1e-07, + "litellm_provider": "bedrock_converse", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5.5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": false, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh", + "supports_parallel_tool_use_config": true, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "source": "https://aws.amazon.com/bedrock/pricing/" + }, + "perplexity/anthropic/claude-haiku-5-5": { + "litellm_provider": "perplexity", + "mode": "responses", + "supports_adaptive_thinking": true, + "supports_web_search": true, + "supports_function_calling": true, + "input_cost_per_token": 1e-07, + "output_cost_per_token": 5e-07, + "cache_read_input_token_cost": 1e-08, + "source": "https://docs.perplexity.ai/docs/agent-api/models" + }, + "us-gov.anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "bedrock_output_config_effort_ceiling": "xhigh", + "cache_creation_input_token_cost": 1.5e-07, + "cache_creation_input_token_cost_above_1hr": 2.4e-07, + "cache_read_input_token_cost": 1.2e-08, + "input_cost_per_token": 1.2e-07, + "litellm_provider": "bedrock_converse", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 6e-07, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_mid_conversation_system": true, + "supports_native_structured_output": false, + "supports_output_config": true, + "supports_parallel_tool_use_config": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true + }, + "us.anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "cache_creation_input_token_cost": 1.375e-07, + "cache_creation_input_token_cost_above_1hr": 2.2e-07, + "cache_read_input_token_cost": 1.1e-08, + "input_cost_per_token": 1.1e-07, + "litellm_provider": "bedrock_converse", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5.5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": false, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh", + "supports_parallel_tool_use_config": true, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "source": "https://aws.amazon.com/bedrock/pricing/" + }, + "vertex_ai/claude-haiku-5-5": { + "regional_endpoint_uplift_multiplier": 1.1, + "supports_mid_conversation_system": true, + "cache_creation_input_token_cost": 1.25e-07, + "cache_creation_input_token_cost_above_1hr": 2e-07, + "cache_creation_input_token_cost_batches": 6.25e-08, + "cache_read_input_token_cost": 1e-08, + "cache_read_input_token_cost_batches": 5e-09, + "input_cost_per_token": 1e-07, + "input_cost_per_token_batches": 5e-08, + "litellm_provider": "vertex_ai-anthropic_models", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "output_cost_per_token_batches": 2.5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true, + "supports_forced_tool_use": true, + "thinking_always_on": false, + "prompt_cache_min_tokens": 512, + "source": "https://platform.claude.com/docs/en/about-claude/pricing" + }, + "vertex_ai/claude-haiku-5-5@default": { + "regional_endpoint_uplift_multiplier": 1.1, + "supports_mid_conversation_system": true, + "cache_creation_input_token_cost": 1.25e-07, + "cache_creation_input_token_cost_above_1hr": 2e-07, + "cache_creation_input_token_cost_batches": 6.25e-08, + "cache_read_input_token_cost": 1e-08, + "cache_read_input_token_cost_batches": 5e-09, + "input_cost_per_token": 1e-07, + "input_cost_per_token_batches": 5e-08, + "litellm_provider": "vertex_ai-anthropic_models", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "output_cost_per_token_batches": 2.5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true, + "supports_forced_tool_use": true, + "thinking_always_on": false, + "prompt_cache_min_tokens": 512, + "source": "https://platform.claude.com/docs/en/about-claude/pricing" } } diff --git a/litellm/proxy/lens/endpoints.py b/litellm/proxy/lens/endpoints.py index c830ca008c3..e070324f004 100644 --- a/litellm/proxy/lens/endpoints.py +++ b/litellm/proxy/lens/endpoints.py @@ -1,10 +1,11 @@ import hashlib import secrets +from collections.abc import Awaitable, Callable from datetime import datetime, timedelta, timezone from functools import reduce from itertools import chain from types import MappingProxyType -from typing import Annotated, Final, TypeAlias +from typing import Annotated, Final, Protocol, TypeAlias from uuid import uuid4 from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, Response @@ -59,7 +60,7 @@ from litellm.proxy.lens.models import ( WorkerCreated, ) from litellm.proxy.lens.release import PROTOCOL_VERSION, release_tag, worker_image -from litellm.proxy.lens.repository import LensRepository, WriterDatabase +from litellm.proxy.lens.repository import DueLens, LensRepository, WriterDatabase from litellm.proxy.lens.reviews import criteria_key from litellm.proxy.lens.sources import ActivityAvailability, SourceReader, Storage, parse_execution from litellm.proxy.lens.state import ( @@ -82,9 +83,25 @@ from litellm.tracing.remote import LensConnection, bounded_response from litellm.types.llms.base import LiteLLMBaseModel router: Final = APIRouter(prefix="/lens", tags=["Lens"]) +CLAIM_CANDIDATES: Final = 20 _bearer: Final = HTTPBearer() Auth: TypeAlias = Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)] StorageDep: TypeAlias = Annotated[Storage | None, Depends(provide_storage)] +SAMPLE_PAGE_SIZE: Final = 10_000 +SAMPLE_PAGE_SIZES: Final = (SAMPLE_PAGE_SIZE, 5_000, 2_500, 1_250, 625, 312, 156, 100) +SAMPLE_RESPONSE_TOO_LARGE: Final = "ClickHouse query exceeded the response size limit" + + +class _ClaimRepository(Protocol): + async def due( + self, scope: Scope, now: datetime, limit: int, after: DueLens | None = None + ) -> tuple[DueLens, ...]: ... + + async def sync_due(self, lens: Lens) -> None: ... + + async def update( + self, lens_id: str, transform: Callable[[Lens], Lens], attempts: int, *, changed_only: bool + ) -> Lens | None: ... def repository() -> LensRepository: @@ -615,13 +632,29 @@ async def claim(worker: WorkerAuth, protocol_version: int = 1, worker_release: s if worker.analysis_key_id is None: raise HTTPException(409, "Assign an analysis key to this worker in Lens setup") now: Final = datetime.now(timezone.utc) - await repository().heartbeat(worker.id, now.isoformat()) - async for candidate in repository().claim_candidates(worker.scope, now): - if not can_access(worker.scope, candidate.scope): - continue - if claimed := await claim_candidate(candidate, worker, now): - return claimed - return None + lens_repository: Final = repository() + await lens_repository.heartbeat(worker.id, now.isoformat()) + return await claim_due(worker, now, lens_repository) + + +async def claim_due( + worker: Worker, + now: datetime, + lens_repository: _ClaimRepository, + supports_model: Callable[[Worker, LensSettings], Awaitable[bool]] = worker_supports_model, +) -> Claim | None: + after: DueLens | None = None # rebind-ok: keyset cursor advances one page at a time + while True: + page = await lens_repository.due(worker.scope, now, CLAIM_CANDIDATES, after) + for candidate in page: + if not can_access(worker.scope, candidate.lens.scope): + continue + if claimed := await claim_candidate(candidate.lens, worker, now, lens_repository, supports_model): + return claimed + await lens_repository.sync_due(candidate.lens) + if len(page) < CLAIM_CANDIDATES: + return None + after = page[-1] @router.post("/worker/{lens_id}/{job_id}/progress", response_model=bool) @@ -654,24 +687,37 @@ async def sample(lens_id: str, job_id: str, worker: WorkerAuth, storage: Storage lens, job = await assigned(lens_id, job_id, worker, attempt) if job.sample is not None: return job.sample - pages: list[Sample] = [] # mutable-ok: freeze selection after stable cursor traversal + + async def read_page(cursor: str, sizes: tuple[int, ...]) -> tuple[Sample, tuple[int, ...]]: + page_size: Final = sizes[0] + try: + page: Final = await source_reader(storage).sample( + lens.scope, + job.settings, + int(job.start.timestamp() * 1000), + int(job.end.timestamp() * 1000), + page_size=page_size, + cursor=cursor, + ) + except RuntimeError as error: + if type(error) is not RuntimeError or str(error) != SAMPLE_RESPONSE_TOO_LARGE or len(sizes) == 1: + raise + return await read_page(cursor, sizes[1:]) + return page, sizes + + pages: list[tuple[Sample, tuple[int, ...]]] = [] # mutable-ok: freeze selection after stable cursor traversal cursor = "" # rebind-ok: advance by immutable identity, never by shifting row positions while True: - page = await source_reader(storage).sample( - lens.scope, - job.settings, - int(job.start.timestamp() * 1000), - int(job.end.timestamp() * 1000), - cursor=cursor, - ) - pages.append(page) - if not page.next_cursor or sum(len(p.executions) for p in pages) >= pages[0].selected: + sizes: Final = pages[-1][1] if pages else SAMPLE_PAGE_SIZES + page, usable_sizes = await read_page(cursor, sizes) + pages.append((page, usable_sizes)) + if not page.next_cursor or sum(len(p.executions) for p, _ in pages) >= pages[0][0].selected: break cursor = page.next_cursor executions: Final = tuple( - execution for p in pages for execution in p.executions + execution for p, _ in pages for execution in p.executions ) # comprehension-ok: flatten query pages - selected: Final = Sample(executions=executions, eligible=pages[0].eligible, selected=len(executions)) + selected: Final = Sample(executions=executions, eligible=pages[0][0].eligible, selected=len(executions)) def freeze(e: Lens) -> Lens: active: Final = current_job(e) @@ -894,9 +940,15 @@ async def heartbeat(lens_id: str, job_id: str, worker: WorkerAuth, attempt: Atte return await progress(lens_id, job_id, Progress(), worker, attempt) -async def claim_candidate(candidate: Lens, worker: Worker, now: datetime) -> Claim | None: +async def claim_candidate( + candidate: Lens, + worker: Worker, + now: datetime, + lens_repository: _ClaimRepository, + supports_model: Callable[[Worker, LensSettings], Awaitable[bool]] = worker_supports_model, +) -> Claim | None: active: Final = current_job(candidate) - if not await worker_supports_model(worker, active.settings if active else candidate.settings): + if not await supports_model(worker, active.settings if active else candidate.settings): return None job_id: Final = str(uuid4()) @@ -907,7 +959,7 @@ async def claim_candidate(candidate: Lens, worker: Worker, now: datetime) -> Cla return e return claim_job(scheduled, worker, now) - updated: Final = await repository().update(candidate.id, schedule, changed_only=True) + updated: Final = await lens_repository.update(candidate.id, schedule, attempts=1, changed_only=True) if updated is None: return None job: Final = current_job(updated) diff --git a/litellm/proxy/lens/repository.py b/litellm/proxy/lens/repository.py index 05099014091..219f88fd11d 100644 --- a/litellm/proxy/lens/repository.py +++ b/litellm/proxy/lens/repository.py @@ -1,8 +1,9 @@ import asyncio import json import random -from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable +from collections.abc import AsyncGenerator, Awaitable, Callable from contextlib import AbstractAsyncContextManager, asynccontextmanager +from dataclasses import dataclass from datetime import datetime, timedelta, timezone from types import MappingProxyType from typing import TYPE_CHECKING, Final, Protocol @@ -25,7 +26,7 @@ from litellm.proxy.lens.models import ( Worker, ) from litellm.proxy.lens.reviews import criteria_key -from litellm.proxy.lens.state import apply_progress, current_job, replace_job +from litellm.proxy.lens.state import apply_progress, current_job, due_at, replace_job from litellm.types.llms.base import LiteLLMBaseModel if TYPE_CHECKING: @@ -40,6 +41,18 @@ class Database(Protocol): class Row(LiteLLMBaseModel): data: JsonValue + due_at: datetime | None = None + + +class DueRow(LiteLLMBaseModel): + data: JsonValue + due_at: datetime + + +@dataclass(frozen=True, slots=True) +class DueLens: + lens: Lens + due_at: datetime class FindingRun(LiteLLMBaseModel): @@ -48,6 +61,26 @@ class FindingRun(LiteLLMBaseModel): _ROWS: Final = TypeAdapter(tuple[Row, ...]) +_DUE_ROWS: Final = TypeAdapter(tuple[DueRow, ...]) +_DUE_QUERY: Final[LiteralString] = """SELECT data, due_at FROM "LiteLLM_Lens" +WHERE due_at IS NOT NULL AND due_at <= ($4::timestamptz AT TIME ZONE 'UTC') +AND ($1::boolean OR ( + COALESCE((data->'scope'->>'all_teams')::boolean, false) IS NOT TRUE + AND COALESCE(data->'scope'->>'team_id', '')=$2 + AND ($2 <> '' OR COALESCE(data->'scope'->>'api_key_hash', '')=$3) +)) +ORDER BY due_at, id +LIMIT $5""" +_DUE_AFTER_QUERY: Final[LiteralString] = """SELECT data, due_at FROM "LiteLLM_Lens" +WHERE due_at IS NOT NULL AND due_at <= ($4::timestamptz AT TIME ZONE 'UTC') +AND (due_at, id) > ($6::timestamp, $7) +AND ($1::boolean OR ( + COALESCE((data->'scope'->>'all_teams')::boolean, false) IS NOT TRUE + AND COALESCE(data->'scope'->>'team_id', '')=$2 + AND ($2 <> '' OR COALESCE(data->'scope'->>'api_key_hash', '')=$3) +)) +ORDER BY due_at, id +LIMIT $5""" UPDATE_ATTEMPTS: Final = 40 UPDATE_BACKOFF_SECONDS: Final = 0.02 @@ -170,38 +203,29 @@ class LensRepository: rows: Final = _ROWS.validate_python(await self.db.query_raw('SELECT data FROM "LiteLLM_Lens" ORDER BY id')) return tuple(Lens.model_validate(row.data) for row in rows) - async def claim_candidates(self, scope: Scope, now: datetime) -> AsyncIterator[Lens]: - cursor = "" # rebind-ok: advance a bounded keyset page - while True: - rows = _ROWS.validate_python( # rebind-ok: fetch the next bounded keyset page - await self.db.query_raw( - """SELECT data FROM "LiteLLM_Lens" - WHERE id > $1 AND ($2::boolean OR ( - COALESCE((data->'scope'->>'all_teams')::boolean, false)=false - AND data->'scope'->>'team_id'=$3 - AND ($3<>'' OR data->'scope'->>'api_key_hash'=$4))) - AND ( - EXISTS (SELECT 1 FROM jsonb_array_elements(data->'jobs') AS job - WHERE job->>'status'='queued' OR (job->>'status'='running' - AND (job->>'lease_until' IS NULL OR (job->>'lease_until')::timestamptz<=$5::timestamptz))) - OR ((data->'settings'->>'enabled')::boolean - AND (data->>'next_run_at')::timestamptz<=$5::timestamptz - AND NOT EXISTS (SELECT 1 FROM jsonb_array_elements(data->'jobs') AS job - WHERE job->>'status' IN ('queued', 'running')))) - ORDER BY id LIMIT 50""", - cursor, - scope.all_teams, - scope.team_id, - scope.api_key_hash, - now.isoformat(), - ) + async def due(self, scope: Scope, now: datetime, limit: int, after: DueLens | None = None) -> tuple[DueLens, ...]: + query: Final[LiteralString] = _DUE_QUERY if after is None else _DUE_AFTER_QUERY + parameters: Final[tuple[object, ...]] = ( + ( + scope.all_teams, + scope.team_id, + scope.api_key_hash, + now.isoformat(), + limit, ) - candidates = tuple(Lens.model_validate(row.data) for row in rows) # rebind-ok: decode this page - for candidate in candidates: - yield candidate - if len(candidates) < 50: - return - cursor = candidates[-1].id + if after is None + else ( + scope.all_teams, + scope.team_id, + scope.api_key_hash, + now.isoformat(), + limit, + after.due_at, + after.lens.id, + ) + ) + rows: Final = _DUE_ROWS.validate_python(await self.db.query_raw(query, *parameters), from_attributes=True) + return tuple(DueLens(lens=Lens.model_validate(row.data), due_at=row.due_at) for row in rows) async def get(self, lens_id: str) -> Lens | None: rows: Final = _ROWS.validate_python( @@ -214,12 +238,25 @@ class LensRepository: async def create(self, lens: Lens) -> Lens: await self.db.execute_raw( - 'INSERT INTO "LiteLLM_Lens" (id, version, data) VALUES ($1,0,$2::jsonb)', + """INSERT INTO "LiteLLM_Lens" (id, version, data, due_at) + VALUES ($1,0,$2::jsonb,($3::timestamptz AT TIME ZONE 'UTC'))""", lens.id, lens.model_dump_json(), + scheduled_at.isoformat() if (scheduled_at := due_at(lens)) else None, ) return lens + async def sync_due(self, lens: Lens) -> None: + await self.db.execute_raw( + """UPDATE "LiteLLM_Lens" + SET due_at=($3::timestamptz AT TIME ZONE 'UTC') + WHERE id=$1 AND version=$2 + AND due_at IS DISTINCT FROM ($3::timestamptz AT TIME ZONE 'UTC')""", + lens.id, + lens.version, + scheduled_at.isoformat() if (scheduled_at := due_at(lens)) else None, + ) + async def update( self, lens_id: str, @@ -250,7 +287,8 @@ class LensRepository: """WITH previous AS MATERIALIZED ( SELECT data FROM "LiteLLM_Lens" WHERE id=$2 AND version=$3 FOR UPDATE ), updated AS ( - UPDATE "LiteLLM_Lens" SET data=$1::jsonb, version=version+1 + UPDATE "LiteLLM_Lens" SET data=$1::jsonb, version=version+1, + due_at=($4::timestamptz AT TIME ZONE 'UTC') WHERE id=$2 AND version=$3 AND EXISTS (SELECT 1 FROM previous) RETURNING id ) , archived AS (INSERT INTO "LiteLLM_LensRun" (id, lens_id, created_at, data) @@ -264,6 +302,7 @@ class LensRepository: updated.model_dump_json(), lens_id, previous.version, + scheduled_at.isoformat() if (scheduled_at := due_at(updated)) else None, ) ) return bool(rows and rows[0].data == 1), updated diff --git a/litellm/proxy/lens/state.py b/litellm/proxy/lens/state.py index 38efc7f23bb..599dca8f1c6 100644 --- a/litellm/proxy/lens/state.py +++ b/litellm/proxy/lens/state.py @@ -37,6 +37,15 @@ def current_job(lens: Lens) -> Job | None: return next((job for job in lens.jobs if job.status in ("queued", "running")), None) +def due_at(lens: Lens) -> datetime | None: + job: Final = current_job(lens) + if job is None: + return lens.next_run_at if lens.settings.enabled else None + if job.status == "queued": + return job.created_at + return job.lease_until or job.created_at + + def replace_job(lens: Lens, job: Job) -> Lens: return lens.model_copy( update=MappingProxyType({"jobs": tuple(job if old.id == job.id else old for old in lens.jobs)}) diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 514df905866..ceb127c31bd 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -1939,6 +1939,9 @@ model LiteLLM_Lens { id String @id version Int @default(0) data Json + due_at DateTime? @default(dbgenerated("'1970-01-01 00:00:00'::timestamp without time zone")) + + @@index([due_at], map: "LiteLLM_Lens_due_at_idx") } model LiteLLM_LensRun { diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index dbc1bb0de46..cecc7d0856f 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -15229,8 +15229,8 @@ "cache_creation_input_token_cost": 2.5e-06, "cache_creation_input_token_cost_above_1hr": 4e-06, "cache_creation_input_token_cost_batches": 1.25e-06, - "cache_read_input_token_cost": 2e-07, - "cache_read_input_token_cost_batches": 1e-07, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_batches": 5e-08, "input_cost_per_token": 2e-06, "input_cost_per_token_batches": 1e-06, "litellm_provider": "anthropic", @@ -15267,7 +15267,7 @@ "supports_forced_tool_use": false, "thinking_always_on": true, "prompt_cache_min_tokens": 512, - "source": "https://platform.claude.com/docs/en/models/sonnet-5-5/overview", + "source": "https://platform.claude.com/docs/en/about-claude/pricing", "supports_web_search": true }, "claude-sonnet-4-6": { @@ -80768,5 +80768,648 @@ "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true + }, +"claude-haiku-5-5": { + "supports_anthropic_compaction": true, + "cache_creation_input_token_cost": 1.25e-07, + "cache_creation_input_token_cost_above_1hr": 2e-07, + "cache_creation_input_token_cost_batches": 6.25e-08, + "cache_read_input_token_cost": 1e-08, + "cache_read_input_token_cost_batches": 5e-09, + "input_cost_per_token": 1e-07, + "input_cost_per_token_batches": 5e-08, + "litellm_provider": "anthropic", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "output_cost_per_token_batches": 2.5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_native_structured_output": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true, + "provider_specific_entry": { + "us": 1.1 + }, + "supports_output_config": true, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "source": "https://platform.claude.com/docs/en/models/haiku-5-5/overview", + "supports_web_search": true, + "input_cost_per_token_above_100k_tokens": 5e-07, + "output_cost_per_token_above_100k_tokens": 2.5e-06, + "cache_creation_input_token_cost_above_100k_tokens": 6.25e-07, + "cache_creation_input_token_cost_above_1hr_above_100k_tokens": 1e-06, + "cache_read_input_token_cost_above_100k_tokens": 5e-08 + }, + "bedrock_mantle/anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "cache_creation_input_token_cost": 1.375e-07, + "cache_creation_input_token_cost_above_1hr": 2.2e-07, + "cache_read_input_token_cost": 1.1e-08, + "input_cost_per_token": 1.1e-07, + "litellm_provider": "bedrock_mantle", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5.5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": false, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh", + "supports_parallel_tool_use_config": true, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-anthropic-claude-haiku-5-5.html" + }, + "bedrock_mantle/us-gov-west-1/anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "bedrock_output_config_effort_ceiling": "xhigh", + "cache_creation_input_token_cost": 1.5e-07, + "cache_creation_input_token_cost_above_1hr": 2.4e-07, + "cache_read_input_token_cost": 1.2e-08, + "input_cost_per_token": 1.2e-07, + "litellm_provider": "bedrock_mantle", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 6e-07, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_mid_conversation_system": true, + "supports_native_structured_output": false, + "supports_output_config": true, + "supports_parallel_tool_use_config": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-anthropic-claude-haiku-5-5.html" + }, + "anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "cache_creation_input_token_cost": 1.25e-07, + "cache_creation_input_token_cost_above_1hr": 2e-07, + "cache_read_input_token_cost": 1e-08, + "input_cost_per_token": 1e-07, + "litellm_provider": "bedrock_converse", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": false, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh", + "supports_parallel_tool_use_config": true, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + }, + "apac.anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "litellm_provider": "bedrock_converse", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": false, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh", + "supports_parallel_tool_use_config": true, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "cache_creation_input_token_cost_above_1hr": 2.2e-07, + "cache_creation_input_token_cost": 1.375e-07, + "input_cost_per_token": 1.1e-07, + "output_cost_per_token": 5.5e-07, + "cache_read_input_token_cost": 1.1e-08, + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + }, + "au.anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "cache_creation_input_token_cost": 1.375e-07, + "cache_creation_input_token_cost_above_1hr": 2.2e-07, + "cache_read_input_token_cost": 1.1e-08, + "input_cost_per_token": 1.1e-07, + "litellm_provider": "bedrock_converse", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5.5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": false, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh", + "supports_parallel_tool_use_config": true, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "source": "https://aws.amazon.com/bedrock/pricing/" + }, + "azure_ai/claude-haiku-5-5": { + "supports_mid_conversation_system": true, + "cache_creation_input_token_cost": 1.25e-07, + "cache_creation_input_token_cost_above_1hr": 2e-07, + "cache_read_input_token_cost": 1e-08, + "deprecation_date": "2027-09-29", + "input_cost_per_token": 1e-07, + "litellm_provider": "azure_ai", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true, + "supports_forced_tool_use": true, + "thinking_always_on": false, + "prompt_cache_min_tokens": 512, + "source": "https://platform.claude.com/docs/en/models/haiku-5-5/migration-guide" + }, + "bedrock/us-gov-east-1/anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "bedrock_output_config_effort_ceiling": "xhigh", + "cache_creation_input_token_cost": 1.5e-07, + "cache_creation_input_token_cost_above_1hr": 2.4e-07, + "cache_read_input_token_cost": 1.2e-08, + "input_cost_per_token": 1.2e-07, + "litellm_provider": "bedrock", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 6e-07, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_mid_conversation_system": true, + "supports_native_structured_output": false, + "supports_output_config": true, + "supports_parallel_tool_use_config": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true + }, + "bedrock/us-gov-west-1/anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "bedrock_output_config_effort_ceiling": "xhigh", + "cache_creation_input_token_cost": 1.5e-07, + "cache_creation_input_token_cost_above_1hr": 2.4e-07, + "cache_read_input_token_cost": 1.2e-08, + "input_cost_per_token": 1.2e-07, + "litellm_provider": "bedrock", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 6e-07, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_mid_conversation_system": true, + "supports_native_structured_output": false, + "supports_output_config": true, + "supports_parallel_tool_use_config": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true + }, + "eu.anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "cache_creation_input_token_cost": 1.375e-07, + "cache_creation_input_token_cost_above_1hr": 2.2e-07, + "cache_read_input_token_cost": 1.1e-08, + "input_cost_per_token": 1.1e-07, + "litellm_provider": "bedrock_converse", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5.5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": false, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh", + "supports_parallel_tool_use_config": true, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "source": "https://aws.amazon.com/bedrock/pricing/" + }, + "global.anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "cache_creation_input_token_cost": 1.25e-07, + "cache_creation_input_token_cost_above_1hr": 2e-07, + "cache_read_input_token_cost": 1e-08, + "input_cost_per_token": 1e-07, + "litellm_provider": "bedrock_converse", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": false, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh", + "supports_parallel_tool_use_config": true, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "source": "https://aws.amazon.com/bedrock/pricing/" + }, + "jp.anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "cache_creation_input_token_cost": 1.375e-07, + "cache_creation_input_token_cost_above_1hr": 2.2e-07, + "cache_read_input_token_cost": 1.1e-08, + "input_cost_per_token": 1.1e-07, + "litellm_provider": "bedrock_converse", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5.5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": false, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh", + "supports_parallel_tool_use_config": true, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "source": "https://aws.amazon.com/bedrock/pricing/" + }, + "perplexity/anthropic/claude-haiku-5-5": { + "litellm_provider": "perplexity", + "mode": "responses", + "supports_adaptive_thinking": true, + "supports_web_search": true, + "supports_function_calling": true, + "input_cost_per_token": 1e-07, + "output_cost_per_token": 5e-07, + "cache_read_input_token_cost": 1e-08, + "source": "https://docs.perplexity.ai/docs/agent-api/models" + }, + "us-gov.anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "bedrock_output_config_effort_ceiling": "xhigh", + "cache_creation_input_token_cost": 1.5e-07, + "cache_creation_input_token_cost_above_1hr": 2.4e-07, + "cache_read_input_token_cost": 1.2e-08, + "input_cost_per_token": 1.2e-07, + "litellm_provider": "bedrock_converse", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 6e-07, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_mid_conversation_system": true, + "supports_native_structured_output": false, + "supports_output_config": true, + "supports_parallel_tool_use_config": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true + }, + "us.anthropic.claude-haiku-5-5": { + "bedrock_converse_supports_strict_tools": false, + "cache_creation_input_token_cost": 1.375e-07, + "cache_creation_input_token_cost_above_1hr": 2.2e-07, + "cache_read_input_token_cost": 1.1e-08, + "input_cost_per_token": 1.1e-07, + "litellm_provider": "bedrock_converse", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5.5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": false, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh", + "supports_parallel_tool_use_config": true, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "source": "https://aws.amazon.com/bedrock/pricing/" + }, + "vertex_ai/claude-haiku-5-5": { + "regional_endpoint_uplift_multiplier": 1.1, + "supports_mid_conversation_system": true, + "cache_creation_input_token_cost": 1.25e-07, + "cache_creation_input_token_cost_above_1hr": 2e-07, + "cache_creation_input_token_cost_batches": 6.25e-08, + "cache_read_input_token_cost": 1e-08, + "cache_read_input_token_cost_batches": 5e-09, + "input_cost_per_token": 1e-07, + "input_cost_per_token_batches": 5e-08, + "litellm_provider": "vertex_ai-anthropic_models", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "output_cost_per_token_batches": 2.5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true, + "supports_forced_tool_use": true, + "thinking_always_on": false, + "prompt_cache_min_tokens": 512, + "source": "https://platform.claude.com/docs/en/about-claude/pricing" + }, + "vertex_ai/claude-haiku-5-5@default": { + "regional_endpoint_uplift_multiplier": 1.1, + "supports_mid_conversation_system": true, + "cache_creation_input_token_cost": 1.25e-07, + "cache_creation_input_token_cost_above_1hr": 2e-07, + "cache_creation_input_token_cost_batches": 6.25e-08, + "cache_read_input_token_cost": 1e-08, + "cache_read_input_token_cost_batches": 5e-09, + "input_cost_per_token": 1e-07, + "input_cost_per_token_batches": 5e-08, + "litellm_provider": "vertex_ai-anthropic_models", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "output_cost_per_token_batches": 2.5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true, + "supports_forced_tool_use": true, + "thinking_always_on": false, + "prompt_cache_min_tokens": 512, + "source": "https://platform.claude.com/docs/en/about-claude/pricing" } } diff --git a/model_prices_and_context_window.schema.json b/model_prices_and_context_window.schema.json index de3d760d57a..41b616a6b1f 100644 --- a/model_prices_and_context_window.schema.json +++ b/model_prices_and_context_window.schema.json @@ -83,6 +83,11 @@ "minimum": 0, "description": "USD per token written to the provider's prompt cache." }, + "cache_creation_input_token_cost_above_100k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, "cache_creation_input_token_cost_above_128k_tokens": { "type": "number", "minimum": 0, @@ -93,6 +98,11 @@ "minimum": 0, "description": "Rate applied once the prompt exceeds the token threshold in the field name." }, + "cache_creation_input_token_cost_above_1hr_above_100k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, "cache_creation_input_token_cost_above_1hr_above_200k_tokens": { "type": "number", "minimum": 0, @@ -174,6 +184,11 @@ "minimum": 0, "description": "USD per prompt token served from the provider's prompt cache." }, + "cache_read_input_token_cost_above_100k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, "cache_read_input_token_cost_above_128k_tokens": { "type": "number", "minimum": 0, @@ -377,6 +392,11 @@ "minimum": 0, "description": "USD per prompt token." }, + "input_cost_per_token_above_100k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, "input_cost_per_token_above_128k_tokens": { "type": "number", "minimum": 0, @@ -756,6 +776,11 @@ "minimum": 0, "description": "USD per generated token." }, + "output_cost_per_token_above_100k_tokens": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, "output_cost_per_token_above_128k_tokens": { "type": "number", "minimum": 0, diff --git a/schema.prisma b/schema.prisma index 9797322eed2..2f1d4aab78e 100644 --- a/schema.prisma +++ b/schema.prisma @@ -1939,6 +1939,9 @@ model LiteLLM_Lens { id String @id version Int @default(0) data Json + due_at DateTime? @default(dbgenerated("'1970-01-01 00:00:00'::timestamp without time zone")) + + @@index([due_at], map: "LiteLLM_Lens_due_at_idx") } model LiteLLM_LensRun { diff --git a/tests/e2e/a2a/a2a_client.py b/tests/e2e/a2a/a2a_client.py index 605dd8fb7e5..4017602c959 100644 --- a/tests/e2e/a2a/a2a_client.py +++ b/tests/e2e/a2a/a2a_client.py @@ -19,6 +19,7 @@ from pydantic import BaseModel, ConfigDict, Field from e2e_config import settle_propagation from e2e_http import NoBody, Result, Success, get_external, is_ok +from e2e_metadata import STEP_FRAMES, step from proxy_client import ProxyClient @@ -87,7 +88,7 @@ class A2ABridgeParams(BaseModel): custom_llm_provider: str model: str - api_key: str | None = None + api_key: str | None = Field(default=None, repr=False) class AgentRegisterBody(BaseModel): @@ -291,6 +292,7 @@ class A2AResponse(BaseModel): class A2AClient: proxy: ProxyClient + @step("Register the A2A agent {body.agent_name} through /v1/agents") def register_agent(self, body: AgentRegisterBody) -> Result[AgentResponse]: """Register an agent and, on success, wait until the data plane serves it. @@ -337,6 +339,7 @@ class A2AClient: ) time.sleep(self.proxy.poll_interval) + @step("Read the A2A agent back from /v1/agents/{{agent_id}}") def get_agent(self, agent_id: str) -> Result[AgentResponse]: return self.proxy.transport.get( f"/v1/agents/{agent_id}", @@ -345,6 +348,7 @@ class A2AClient: response_type=AgentResponse, ) + @step("Delete the A2A agent") def delete_agent(self, agent_id: str) -> None: result = self.proxy.transport.delete( f"/v1/agents/{agent_id}", @@ -353,8 +357,9 @@ class A2AClient: response_type=NoBody, ) if not is_ok(result): - warnings.warn(f"delete_agent({agent_id!r}) failed: {result}", stacklevel=2) + warnings.warn(f"delete_agent({agent_id!r}) failed: {result}", stacklevel=2 + STEP_FRAMES) + @step("Read the A2A agent's card from /a2a/{{agent_id}}/.well-known/agent-card.json with the given key") def agent_card(self, agent_id: str, key: str) -> Result[ServedAgentCard]: return self.proxy.transport.get( f"/a2a/{agent_id}/.well-known/agent-card.json", @@ -363,6 +368,7 @@ class A2AClient: response_type=ServedAgentCard, ) + @step("Send an A2A message to /a2a/{{agent_id}} with {body.params.message.parts}") def send_message(self, agent_id: str, key: str, body: A2AJsonRpcRequest) -> Result[A2AResponse]: return self.proxy.transport.post( f"/a2a/{agent_id}", @@ -376,6 +382,7 @@ def build_a2a_client(proxy: ProxyClient) -> A2AClient: return A2AClient(proxy=proxy) +@step("Fetch a published A2A agent card from its /.well-known endpoint") def fetch_agent_card(url: str, *, timeout: float = 20.0) -> Result[UpstreamAgentCard]: """Fetch a live A2A agent card from its /.well-known endpoint and parse it into the registration model, so a test can register a real published card verbatim rather diff --git a/tests/e2e/a2a/test_a2a_agent_e2e.py b/tests/e2e/a2a/test_a2a_agent_e2e.py index 8b89ce91806..c0427368a67 100644 --- a/tests/e2e/a2a/test_a2a_agent_e2e.py +++ b/tests/e2e/a2a/test_a2a_agent_e2e.py @@ -10,6 +10,8 @@ protocol version, and an unsupported version is refused at registration). from __future__ import annotations +from typing import Final + import pytest from a2a_client import ( @@ -31,6 +33,9 @@ from a2a_client import ( from e2e_config import unique_marker from e2e_http import Result, UnknownApiError, unwrap from lifecycle import ResourceManager +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta + +BRIDGE_MODEL: Final = "claude-haiku-4-5" # No api_key: litellm resolves ANTHROPIC_API_KEY from the proxy's own environment # for this provider, which is what the agent-owner flow relies on. Pinning @@ -42,7 +47,7 @@ from lifecycle import ResourceManager # omitted -> 200, "os.environ/..." -> 500 invalid x-api-key, literal key -> 200. BRIDGE = A2ABridgeParams( custom_llm_provider="anthropic", - model="claude-haiku-4-5", + model=BRIDGE_MODEL, ) MOVEHOME_AGENT_CARD_URL = "https://movehome.org/.well-known/agent.json" @@ -96,6 +101,12 @@ def _ask(text: str) -> A2AJsonRpcRequest: class TestA2AAgentLifecycle: @pytest.mark.covers("other.a2a.register.persists") + @meta( + Subject( + domain=Domain.AGENTS_API, + route=Route.A2A, + ) + ) def test_register_persists(self, client: A2AClient, resources: ResourceManager) -> None: agent = _register(client, resources, "0.3") fetched = unwrap(client.get_agent(agent.agent_id)) @@ -104,6 +115,15 @@ class TestA2AAgentLifecycle: assert fetched.agent_card_params.protocol_version == "0.3" @pytest.mark.covers("other.a2a.register.semver_version_accepted") + @meta( + Subject( + domain=Domain.AGENTS_API, + route=Route.A2A, + providers=(Provider.ANTHROPIC,), + models=(BRIDGE_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_semver_protocol_version_registers_and_serves(self, client: A2AClient, resources: ResourceManager, scoped_key: str) -> None: agent = _register(client, resources, "0.3.0") assert agent.agent_card_params.protocol_version == "0.3" @@ -116,6 +136,12 @@ class TestA2AAgentLifecycle: assert result.text != "" @pytest.mark.covers("other.a2a.message_send.real_world_agent_replies") + @meta( + Subject( + domain=Domain.AGENTS_API, + route=Route.A2A, + ) + ) def test_real_world_agent_replies_to_property_query(self, client: A2AClient, resources: ResourceManager, scoped_key: str) -> None: upstream = unwrap(fetch_agent_card(MOVEHOME_AGENT_CARD_URL)).model_copy(update={"url": MOVEHOME_ORIGIN}) assert upstream.protocol_version == "0.3.0" @@ -152,6 +178,12 @@ class TestA2AAgentLifecycle: assert all(listing.location.un_locode == location for listing in results.listings) @pytest.mark.covers("other.a2a.discovery.proxy_fronted_card") + @meta( + Subject( + domain=Domain.AGENTS_API, + route=Route.A2A, + ) + ) def test_discovery_card_is_proxy_fronted(self, client: A2AClient, resources: ResourceManager, scoped_key: str) -> None: agent = _register(client, resources, "0.3") card = unwrap(client.agent_card(agent.agent_id, scoped_key)) @@ -163,6 +195,15 @@ class TestA2AAgentLifecycle: assert card.supported_interfaces[0].url == card.url @pytest.mark.covers("other.a2a.message_send.bridge_invokes") + @meta( + Subject( + domain=Domain.AGENTS_API, + route=Route.A2A, + providers=(Provider.ANTHROPIC,), + models=(BRIDGE_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_message_send_runs_completion_bridge(self, client: A2AClient, resources: ResourceManager, scoped_key: str) -> None: agent = _register(client, resources, "0.3") request = _ask("Reply with exactly the word PONG and nothing else") @@ -177,6 +218,15 @@ class TestA2AAgentLifecycle: assert rows[0].model == f"a2a_agent/{agent.agent_card_params.name}" @pytest.mark.covers("other.a2a.version.serves_pinned_0_3") + @meta( + Subject( + domain=Domain.AGENTS_API, + route=Route.A2A, + providers=(Provider.ANTHROPIC,), + models=(BRIDGE_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_pinned_v0_3_serves_flat_message_shape(self, client: A2AClient, resources: ResourceManager, scoped_key: str) -> None: agent = _register(client, resources, "0.3") request = _ask("Say hi in one word") @@ -188,6 +238,15 @@ class TestA2AAgentLifecycle: assert result.text != "" @pytest.mark.covers("other.a2a.version.serves_pinned_1_0") + @meta( + Subject( + domain=Domain.AGENTS_API, + route=Route.A2A, + providers=(Provider.ANTHROPIC,), + models=(BRIDGE_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_pinned_v1_0_serves_nested_message_shape(self, client: A2AClient, resources: ResourceManager, scoped_key: str) -> None: agent = _register(client, resources, "1.0") request = _ask("Say hi in one word") @@ -199,6 +258,12 @@ class TestA2AAgentLifecycle: assert result.text != "" @pytest.mark.covers("other.a2a.register.unsupported_version_rejected") + @meta( + Subject( + domain=Domain.AGENTS_API, + route=Route.A2A, + ) + ) def test_unsupported_protocol_version_rejected(self, client: A2AClient) -> None: result = _register_rejection(client, "9.9") match result: @@ -209,6 +274,12 @@ class TestA2AAgentLifecycle: pytest.fail(f"expected 400 for unsupported protocolVersion, got {result}") @pytest.mark.covers("other.a2a.register.malformed_version_rejected") + @meta( + Subject( + domain=Domain.AGENTS_API, + route=Route.A2A, + ) + ) def test_malformed_protocol_version_rejected(self, client: A2AClient) -> None: result = _register_rejection(client, "0.3.garbage") match result: diff --git a/tests/e2e/access_control/access_control_client.py b/tests/e2e/access_control/access_control_client.py index 5f459c09767..90ddf061bca 100644 --- a/tests/e2e/access_control/access_control_client.py +++ b/tests/e2e/access_control/access_control_client.py @@ -8,6 +8,7 @@ from dataclasses import dataclass from pydantic import BaseModel, ValidationError from proxy_client import ProxyClient +from e2e_metadata import step from e2e_http import NoBody, StreamingResponse, is_ok, unwrap from models import ( ChatBody, @@ -59,14 +60,17 @@ def error_envelope(body: str) -> ApiErrorEnvelope | None: class AccessControlClient: proxy: ProxyClient + @step("Generate a virtual key that can only call LLM API routes") def llm_only_key(self) -> str: return self.proxy.generate_key( KeyGenerateBody(models=[], allowed_routes=["llm_api_routes"]) ) + @step("Delete the virtual key") def delete_key(self, key: str) -> None: self.proxy.delete_key(key) + @step('Send a /chat/completions request to {model} with the prompt "{content}"') def chat_status( self, key: str, model: str, content: str, max_completion_tokens: int | None = None ) -> StreamingResponse: @@ -80,6 +84,7 @@ class AccessControlClient: ), ) + @step("Create the team {team_alias} with models: {models}") def create_team(self, team_alias: str, models: list[str]) -> str: team_id = unwrap( self.proxy.transport.post( @@ -92,6 +97,7 @@ class AccessControlClient: self._await_team(team_id) return team_id + @step("Set the team {team_alias}'s models to {models} through /team/update") def set_team_models(self, team_id: str, team_alias: str, models: list[str]) -> None: """Replace the team's allow-list. /model/new appends a team-scoped deployment's public name to it, so a test that means to grant only an access group has to @@ -105,6 +111,7 @@ class AccessControlClient: ) ) + @step("Delete the team") def delete_team(self, team_id: str) -> None: _ = self.proxy.transport.post( "/team/delete", @@ -113,6 +120,7 @@ class AccessControlClient: response_type=NoBody, ) + @step("List the deployments in the model access group {access_group}") def access_group_info(self, access_group: str) -> AccessGroupInfoResponse | None: result = self.proxy.transport.get( f"/access_group/{access_group}/info", @@ -122,6 +130,7 @@ class AccessControlClient: ) return unwrap(result) if is_ok(result) else None + @step("Read the team's models from /team/info") def team_models(self, team_id: str) -> list[str] | None: result = self.proxy.transport.get( "/team/info", @@ -139,6 +148,7 @@ class AccessControlClient: time.sleep(self.proxy.poll_interval) raise AssertionError(f"/team/info never resolved team {team_id!r} created by /team/new") + @step("Add a deployment named {model_name} that calls openai/gpt-4o-mini with the given key") def create_model_status(self, key: str, model_name: str) -> StreamingResponse: return self.proxy.transport.send( "/model/new", diff --git a/tests/e2e/access_control/test_access_control_e2e.py b/tests/e2e/access_control/test_access_control_e2e.py index 9d01f2915e7..419cae78df2 100644 --- a/tests/e2e/access_control/test_access_control_e2e.py +++ b/tests/e2e/access_control/test_access_control_e2e.py @@ -26,6 +26,7 @@ from e2e_http import Success, UnauthorizedError, UnknownApiError, unwrap from lifecycle import ResourceManager from models import ChatBody, ChatMessage, ChatResponse, EmbedBody, LiteLLMParamsBody from proxy_client import ProxyClient +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta pytestmark = pytest.mark.e2e @@ -36,6 +37,14 @@ EMBEDDING_MODEL = "openai-text-embedding-3-small" class TestAccessControl: + @meta( + Subject( + domain=Domain.PROXY_AUTH, + providers=(Provider.GEMINI,), + models=(ALLOWED_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_allowed_model_is_permitted( self, client: AccessControlClient, resources: ResourceManager ) -> None: @@ -57,6 +66,13 @@ class TestAccessControl: f"200 must carry a real completion, not an error envelope: {result.body[:300]}" ) + @meta( + Subject( + domain=Domain.PROXY_AUTH, + models=(DISALLOWED_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_disallowed_model_is_denied_403( self, client: AccessControlClient, resources: ResourceManager ) -> None: @@ -73,6 +89,13 @@ class TestAccessControl: ) @pytest.mark.covers("other.auth.virtual_key.route_group_allowed") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + providers=(Provider.GEMINI, Provider.OPENAI,), + models=(ALLOWED_MODEL, EMBEDDING_MODEL,), + ) + ) def test_llm_api_routes_group_grants_every_llm_endpoint( self, client: AccessControlClient, resources: ResourceManager ) -> None: @@ -97,6 +120,12 @@ class TestAccessControl: f"the same key must still be shut out of /model/new, got {denied.status_code}: {denied.body[:300]}" ) + @meta( + Subject( + domain=Domain.PROXY_AUTH, + route=Route.MODEL_MANAGEMENT, + ) + ) def test_llm_only_key_forbidden_from_management_route_403( self, client: AccessControlClient, resources: ResourceManager ) -> None: @@ -111,6 +140,12 @@ class TestAccessControl: f"403 body must be a route-permission denial, got: {result.body[:300]}" ) + @meta( + Subject( + domain=Domain.PROXY_AUTH, + mode=Mode.NONSTREAM, + ) + ) def test_unknown_model_returns_400( self, client: AccessControlClient, resources: ResourceManager ) -> None: @@ -140,6 +175,14 @@ class TestVirtualKeyAuth: "mgmt.virtual_key.invalid_denied", exercised_on=[], ) + @meta( + Subject( + domain=Domain.PROXY_AUTH, + providers=(Provider.ANTHROPIC,), + models=(VIRTUAL_KEY_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_valid_key_allows_and_invalid_key_denied( self, proxy: ProxyClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/access_control/test_chat_auth_headers_e2e.py b/tests/e2e/access_control/test_chat_auth_headers_e2e.py index 197a54cc3a3..c0de14a0961 100644 --- a/tests/e2e/access_control/test_chat_auth_headers_e2e.py +++ b/tests/e2e/access_control/test_chat_auth_headers_e2e.py @@ -11,6 +11,7 @@ import pytest from e2e_http import AuthHeaders, NoBody, StreamingResponse, assert_auth_denied from models import ChatBody, ChatMessage from proxy_client import ProxyClient +from e2e_metadata import Domain, Route, Subject, meta pytestmark = pytest.mark.e2e @@ -32,26 +33,56 @@ def _chat_with_headers(proxy: ProxyClient, headers: AuthHeaders | NoBody) -> Str class TestChatAuthHeaders: @pytest.mark.covers("other.auth.llm_chat.missing_header_denied") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + route=Route.CHAT_COMPLETIONS, + ) + ) def test_missing_authorization_header_is_denied(self, proxy: ProxyClient) -> None: result = _chat_with_headers(proxy, NoBody()) assert_auth_denied(result, "missing Authorization") @pytest.mark.covers("other.auth.llm_chat.invalid_bearer_denied") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + route=Route.CHAT_COMPLETIONS, + ) + ) def test_bearer_invalid_token_is_denied(self, proxy: ProxyClient) -> None: result = _chat_with_headers(proxy, AuthHeaders(authorization="Bearer invalid_token")) assert_auth_denied(result, "Bearer invalid_token") @pytest.mark.covers("other.auth.llm_chat.no_bearer_prefix_denied") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + route=Route.CHAT_COMPLETIONS, + ) + ) def test_token_without_bearer_prefix_is_denied(self, proxy: ProxyClient) -> None: result = _chat_with_headers(proxy, AuthHeaders(authorization="invalid_token")) assert_auth_denied(result, "token without Bearer prefix") @pytest.mark.covers("other.auth.llm_chat.empty_bearer_denied") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + route=Route.CHAT_COMPLETIONS, + ) + ) def test_empty_bearer_token_is_denied(self, proxy: ProxyClient) -> None: result = _chat_with_headers(proxy, AuthHeaders(authorization="Bearer ")) assert_auth_denied(result, "empty Bearer token") @pytest.mark.covers("other.auth.llm_chat.not_bearer_scheme_denied") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + route=Route.CHAT_COMPLETIONS, + ) + ) def test_not_bearer_scheme_is_denied(self, proxy: ProxyClient) -> None: result = _chat_with_headers(proxy, AuthHeaders(authorization="NotBearer validtoken123")) assert_auth_denied(result, "NotBearer scheme") diff --git a/tests/e2e/access_control/test_model_access_group_e2e.py b/tests/e2e/access_control/test_model_access_group_e2e.py index 6dab51805fb..cfe2f061693 100644 --- a/tests/e2e/access_control/test_model_access_group_e2e.py +++ b/tests/e2e/access_control/test_model_access_group_e2e.py @@ -33,6 +33,7 @@ from models import ( ModelNewBody, TeamInfoResponse, ) +from e2e_metadata import Domain, Mode, Provider, Subject, meta pytestmark = pytest.mark.e2e @@ -177,6 +178,14 @@ class TestKeyScopedToAccessGroup: "other.auth.model_access_group.member_allowed", ) @pytest.mark.parametrize(("case", "select_model"), ALLOWED) + @meta( + Subject( + domain=Domain.PROXY_AUTH, + providers=(Provider.OPENAI,), + models=(GROUP_BACKEND, WILDCARD_BARE_MODEL, WILDCARD_PREFIXED_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_group_grants_every_deployment_in_it( self, case: str, @@ -202,6 +211,13 @@ class TestKeyScopedToAccessGroup: @pytest.mark.covers("other.auth.model_access_group.non_member_denied") @pytest.mark.parametrize(("case", "select_model"), DENIED) + @meta( + Subject( + domain=Domain.PROXY_AUTH, + models=(GROUP_BACKEND, UNCOVERED_OPENAI_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_group_grants_nothing_outside_it( self, case: str, @@ -228,6 +244,14 @@ class TestKeyScopedToAccessGroup: class TestTeamScopedToAccessGroup: @pytest.mark.covers("other.auth.model_access_group.team_wildcard_bare_name_allowed") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + providers=(Provider.OPENAI,), + models=(TEAM_WILDCARD_BARE_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_group_grants_the_teams_own_wildcard( self, client: AccessControlClient, team_grant: TeamGrant ) -> None: @@ -248,6 +272,12 @@ class TestTeamScopedToAccessGroup: ) @pytest.mark.covers("other.auth.model_access_group.team_non_member_denied") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + mode=Mode.NONSTREAM, + ) + ) def test_group_grants_the_team_nothing_outside_it( self, client: AccessControlClient, team_grant: TeamGrant ) -> None: diff --git a/tests/e2e/llm_translation/realtime/realtime_client.py b/tests/e2e/llm_translation/realtime/realtime_client.py index 6eae4fdf54d..c9f06e787f3 100644 --- a/tests/e2e/llm_translation/realtime/realtime_client.py +++ b/tests/e2e/llm_translation/realtime/realtime_client.py @@ -308,7 +308,7 @@ class RealtimeSession: connection: Connection @step("Send the realtime event {event.type} over the websocket") - def send(self, event: SessionUpdate | ConversationItemCreate | ResponseCreate) -> None: + def send(self, event: SessionUpdate | ConversationItemCreate | ResponseCreate | InputAudioBufferAppend) -> None: self.connection.send(event.model_dump_json(by_alias=True, exclude_none=True)) @step("Wait for a {stop_type} event on the realtime websocket") diff --git a/tests/e2e/migrations/checks.py b/tests/e2e/migrations/checks.py index 619ad3b0e6c..4d67dbaa20d 100644 --- a/tests/e2e/migrations/checks.py +++ b/tests/e2e/migrations/checks.py @@ -3,6 +3,7 @@ from contextlib import ExitStack from typing import Final from uuid import uuid4 +from e2e_metadata import step from psycopg import sql from .containers import Containers, Replica, failed, until @@ -21,12 +22,14 @@ GATED: Final = Migration( ) +@step("Start {count} proxy containers on the test database") def start_replicas( stack: ExitStack, containers: Containers, database: Database, migrations: tuple[Migration, ...] = (), count: int = 3 ) -> tuple[Replica, ...]: return tuple(stack.enter_context(containers.start(database, migrations)) for _ in range(count)) +@step("Check that the migration {migration.name} ran exactly once") def assert_completed(database: Database, migration: Migration = COMPLETE) -> None: assert database.query( 'SELECT finished_at IS NOT NULL, rolled_back_at IS NULL, applied_steps_count FROM ' @@ -36,6 +39,7 @@ def assert_completed(database: Database, migration: Migration = COMPLETE) -> Non assert database.query("SELECT id FROM migration_effect") == ((1,),) +@step("Apply the test migration by hand and record it in _prisma_migrations") def confirmed_history(database: Database) -> str: database.execute(COMPLETE_SQL) row_id: Final = str(uuid4()) @@ -46,6 +50,7 @@ def confirmed_history(database: Database) -> str: return row_id +@step("Check that the original _prisma_migrations row and its effect survived, with the row marked finished: {finished}") def assert_original_proof(database: Database, row_id: str, finished: bool) -> None: assert database.query( 'SELECT id, applied_steps_count, finished_at IS NOT NULL, rolled_back_at IS NULL FROM ' @@ -55,6 +60,7 @@ def assert_original_proof(database: Database, row_id: str, finished: bool) -> No assert database.query("SELECT id FROM migration_effect") == ((1,),) +@step("Install a trigger that pauses the migration before it is marked finished") def pause_completion(database: Database) -> None: database.execute( sql.SQL( @@ -67,6 +73,10 @@ def pause_completion(database: Database) -> None: ) +@step( + "Start a proxy container on the migration and kill it at its crash point, " + "with the migration SQL committed: {after_commit}" +) def interrupt_owner( containers: Containers, database: Database, after_commit: bool, *, stop_database_session: bool = True ) -> None: @@ -104,6 +114,7 @@ def interrupt_owner( ) +@step("Wait for every proxy container to refuse an unconfirmed migration and log recovery guidance") def unconfirmed(replicas: tuple[Replica, ...], database: Database) -> None: failed(replicas, "Migration completion could not be verified") started: Final = str( diff --git a/tests/e2e/migrations/containers.py b/tests/e2e/migrations/containers.py index 30374dedcf3..a05941373be 100644 --- a/tests/e2e/migrations/containers.py +++ b/tests/e2e/migrations/containers.py @@ -11,6 +11,7 @@ from typing import Final from uuid import uuid4 from e2e_http import NoBody, Success, unwrap +from e2e_metadata import step from models import KeyGenerateBody, KeyGenerateResponse, KeyInfoParams, KeyInfoResponse from transport import HttpTransport @@ -20,12 +21,14 @@ from .startup_models import ContainerState, Migration, Observation, Readiness MASTER_KEY: Final = "sk-migration-ci-fixture" +@step("Run a docker command") def docker(*args: str) -> str: result: Final = subprocess.run(("docker", *args), capture_output=True, text=True, timeout=90) assert result.returncode == 0, f"Docker operation failed: {result.stderr}" return result.stdout.strip() +@step("Wait for {description}") def until(description: str, condition: Callable[[], bool], seconds: float = 150) -> None: deadline: Final = time.monotonic() + seconds while time.monotonic() < deadline: @@ -41,9 +44,11 @@ class Replica: transport: HttpTransport output: Path + @step("Read the proxy container's state from docker inspect") def state(self) -> ContainerState: return ContainerState.model_validate_json(docker("inspect", "--format", "{{json .State}}", self.name)) + @step("Check whether the proxy container is running and ready on /health/readiness") def observe(self) -> Observation: state: Final = self.state() result: Final = self.transport.get( @@ -52,15 +57,18 @@ class Replica: ready: Final = isinstance(result, Success) and result.data.status == "healthy" and result.data.db == "connected" return Observation(None if state.Running else state.ExitCode, ready) + @step("Read the proxy container's logs") def logs(self) -> str: result: Final = subprocess.run(("docker", "logs", self.name), capture_output=True, text=True, timeout=30) assert result.returncode == 0, result.stderr return result.stdout + result.stderr + @step("Kill the proxy container") def kill(self) -> None: if self.state().Running: docker("kill", self.name) + @step("Generate a virtual key on the proxy container and read it back from /key/info and the database") def usable(self, database: Database) -> None: alias: Final = f"migration-{uuid4().hex}" key: Final = unwrap( @@ -86,6 +94,7 @@ class Replica: ) == ((alias,),) +@step("Wait for every proxy container to be ready, then generate and read back a virtual key on each") def ready(replicas: tuple[Replica, ...], database: Database) -> None: def all_ready() -> bool: observations: Final = tuple(replica.observe() for replica in replicas) @@ -97,6 +106,7 @@ def ready(replicas: tuple[Replica, ...], database: Database) -> None: replica.usable(database) +@step("Wait for the seed proxy container to be ready and finish building its request-log indexes") def seeded(seed: Replica, database: Database) -> None: ready((seed,), database) until("the seed replica to finish its request-log indexes", lambda: request_log_indexes_built(database)) @@ -110,6 +120,7 @@ def request_log_indexes_built(database: Database) -> bool: ) == ((2,),) +@step('Wait for every proxy container to refuse to start, logging "{marker}"') def failed(replicas: tuple[Replica, ...], marker: str) -> None: def all_stopped() -> bool: observations: Final = tuple(replica.observe() for replica in replicas) @@ -122,6 +133,7 @@ def failed(replicas: tuple[Replica, ...], marker: str) -> None: assert marker in replica.logs(), f"Startup failed outside the expected migration: {marker}" +@step("Check that every proxy container keeps waiting without serving for {seconds}s") def waiting(replicas: tuple[Replica, ...], seconds: float) -> None: deadline: Final = time.monotonic() + seconds while time.monotonic() < deadline: @@ -139,6 +151,7 @@ class Containers: def using(self, image: str) -> "Containers": return replace(self, image=image) + @step("Start a proxy container on the test database") @contextmanager def start( self, @@ -213,6 +226,7 @@ class Containers: subprocess.run(("docker", "rm", "-f", name), capture_output=True, text=True, timeout=30, check=True) +@step("Write the migration {migration.name} into the proxy container's migration directory") def write_migration(directory: Path, migration: Migration) -> None: path: Final = directory / "prisma" / "migrations" / migration.name path.mkdir(parents=True) diff --git a/tests/e2e/migrations/database.py b/tests/e2e/migrations/database.py index a370c21ba0b..929a6aee817 100644 --- a/tests/e2e/migrations/database.py +++ b/tests/e2e/migrations/database.py @@ -11,6 +11,8 @@ import psycopg from psycopg import sql from pydantic import TypeAdapter +from e2e_metadata import step + Scalar = str | int | bool | None ROWS: Final = TypeAdapter(tuple[tuple[Scalar, ...], ...]) GATE_KEY: Final = 39178002 @@ -35,6 +37,7 @@ class Database: container_url: str schema: str = "public" + @step("Open a connection to the test database") @contextmanager def connection(self) -> Generator[psycopg.Connection[tuple[object, ...]]]: with psycopg.connect(self.url, autocommit=True, connect_timeout=5) as connection: @@ -42,19 +45,23 @@ class Database: connection.execute("SET statement_timeout = '15s'") yield connection + @step("Run a SQL statement on the test database") def execute(self, statement: LiteralString | sql.Composed, params: tuple[Scalar, ...] = ()) -> None: with self.connection() as connection: connection.execute(statement, params or None) + @step("Query the test database") def query( self, statement: LiteralString | sql.Composed, params: tuple[Scalar, ...] = () ) -> tuple[tuple[Scalar, ...], ...]: with self.connection() as connection: return ROWS.validate_python(connection.execute(statement, params or None).fetchall()) + @step("Check whether {name} exists in the test database") def exists(self, name: str) -> bool: return self.query("SELECT to_regclass(%s) IS NOT NULL", (name,)) == ((True,),) + @step("Read the migration history from _prisma_migrations") def history(self) -> tuple[tuple[Scalar, ...], ...]: if not self.exists("_prisma_migrations"): return () @@ -63,6 +70,7 @@ class Database: "applied_steps_count, logs FROM _prisma_migrations ORDER BY id" ) + @step("List the database sessions waiting on an advisory lock") def blocked(self, key: int = GATE_KEY) -> tuple[tuple[Scalar, ...], ...]: return self.query( "SELECT pid FROM pg_locks WHERE locktype = 'advisory' AND NOT granted " @@ -71,6 +79,7 @@ class Database: (key >> 32, key & 0xFFFFFFFF), ) + @step("Hold an advisory lock on the test database") @contextmanager def lock(self, key: int = GATE_KEY) -> Generator[None]: with self.connection() as connection: @@ -86,6 +95,7 @@ class Databases: admin_url: str container_admin_url: str + @step("Create a test database") @contextmanager def create(self, template: Database | None = None, schema: str = "public") -> Generator[Database]: name: Final = f"litellm_migration_test_{uuid4().hex[:20]}" @@ -105,6 +115,7 @@ class Databases: connection.execute(sql.SQL("DROP DATABASE {} WITH (FORCE)").format(sql.Identifier(name))) +@step("Create a read-only database role on the test database") @contextmanager def restricted_user(database: Database) -> Generator[Database]: role: Final = f"migration_reader_{uuid4().hex[:16]}" diff --git a/tests/e2e/migrations/test_legacy.py b/tests/e2e/migrations/test_legacy.py index 4cad808d3ff..8d43fcd2deb 100644 --- a/tests/e2e/migrations/test_legacy.py +++ b/tests/e2e/migrations/test_legacy.py @@ -7,6 +7,7 @@ import pytest from .checks import COMPLETE, assert_completed, confirmed_history, assert_original_proof, start_replicas from .containers import Containers, failed, ready, seeded from .database import Database, Databases +from e2e_metadata import Domain, Subject, meta pytestmark: Final = [pytest.mark.e2e, pytest.mark.migration_startup] @@ -42,10 +43,20 @@ def adopt_legacy(containers: Containers, database: Database) -> None: class TestLegacyMigrations: + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_matching_schema_warns_and_starts(self, containers: Containers, database: Database) -> None: adopt_legacy(containers, database) @pytest.mark.parametrize("fault", ("schema_drift", "custom_migrations", "empty_ledger")) + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_unrecognized_legacy_state_is_not_baselined( self, containers: Containers, database: Database, fault: str ) -> None: @@ -64,6 +75,11 @@ class TestLegacyMigrations: ) == ((0,),) @pytest.mark.parametrize("scenario", ("upgrade", "recovery", "legacy")) + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_non_default_schema( self, containers: Containers, databases: Databases, scenario: Literal["upgrade", "recovery", "legacy"] ) -> None: diff --git a/tests/e2e/migrations/test_pooling.py b/tests/e2e/migrations/test_pooling.py index c4015549de3..6f6b8911f94 100644 --- a/tests/e2e/migrations/test_pooling.py +++ b/tests/e2e/migrations/test_pooling.py @@ -13,6 +13,7 @@ from psycopg import sql from .checks import COMPLETE, assert_completed from .containers import Containers, docker, ready, until from .database import Database, Databases, prisma_url, restricted_user +from e2e_metadata import Domain, Subject, meta POOL_IMAGE: Final = ( "ghcr.io/cloudnative-pg/pgbouncer@sha256:e6ddfe22d845e603825e235dd8334b21ecd125abea2a2172478f556b8dee2bb8" @@ -94,6 +95,11 @@ def pool(database: Database, output: Path) -> Generator[str]: class TestMigrationPooling: @pytest.mark.parametrize("scenario,replica_count", (("fresh", 3), ("upgrade", 3), ("legacy", 3), ("upgrade", 6))) + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_direct_migrations_with_one_application_backend( self, containers: Containers, diff --git a/tests/e2e/migrations/test_recovery.py b/tests/e2e/migrations/test_recovery.py index 80e5747eaac..b1f019acbde 100644 --- a/tests/e2e/migrations/test_recovery.py +++ b/tests/e2e/migrations/test_recovery.py @@ -20,12 +20,18 @@ from .checks import ( from .containers import Containers, failed, ready, until, waiting from .database import COORDINATOR_LOCK, GATE_KEY, Database from .startup_models import Migration +from e2e_metadata import Domain, Subject, meta pytestmark: Final = [pytest.mark.e2e, pytest.mark.migration_startup] class TestMigrationRecovery: @pytest.mark.parametrize("after_commit", (False, True)) + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_container_owner_crash(self, containers: Containers, database: Database, after_commit: bool) -> None: interrupt_owner(containers, database, after_commit, stop_database_session=False) history: Final = database.history() @@ -42,6 +48,11 @@ class TestMigrationRecovery: assert database.history() == history @pytest.mark.parametrize("after_commit", (False, True)) + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_owner_and_database_session_crash( self, containers: Containers, database: Database, after_commit: bool ) -> None: @@ -60,6 +71,11 @@ class TestMigrationRecovery: assert database.history() == history @pytest.mark.parametrize("later_failure", (False, True)) + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_remaining_migrations_after_recovery( self, containers: Containers, database: Database, later_failure: bool ) -> None: @@ -98,6 +114,11 @@ class TestMigrationRecovery: assert database.query("SELECT id FROM migration_next") == ((2,),) assert_original_proof(database, original, True) + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_second_crash_during_recovery_is_atomic(self, containers: Containers, database: Database) -> None: original: Final = confirmed_history(database) pause_completion(database) @@ -114,6 +135,11 @@ class TestMigrationRecovery: ready(start_replicas(stack, containers, database, (COMPLETE,)), database) assert_original_proof(database, original, True) + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_competing_recovery_rechecks_stale_failures(self, containers: Containers, database: Database) -> None: original: Final = confirmed_history(database) with ExitStack() as stack: @@ -132,6 +158,11 @@ class TestMigrationRecovery: @pytest.mark.parametrize( "fault", ("no_steps", "extra_steps", "failure_logs", "checksum", "duplicate_history", "missing_script") ) + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_unproven_history_is_never_repaired( self, containers: Containers, @@ -172,6 +203,11 @@ class TestMigrationRecovery: assert database.history() == history assert database.query("SELECT id FROM migration_effect") == ((1,),) + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_coordinator_timeout_preserves_proof(self, containers: Containers, database: Database) -> None: original: Final = confirmed_history(database) with database.lock(COORDINATOR_LOCK): diff --git a/tests/e2e/migrations/test_rolling_upgrade.py b/tests/e2e/migrations/test_rolling_upgrade.py index 5ad74e0ba8c..03f12cada3a 100644 --- a/tests/e2e/migrations/test_rolling_upgrade.py +++ b/tests/e2e/migrations/test_rolling_upgrade.py @@ -14,11 +14,17 @@ from .upgrade import ( migration_names, provision, ) +from e2e_metadata import Domain, Subject, meta pytestmark: Final = [pytest.mark.e2e, pytest.mark.migration_startup] class TestRollingUpgrade: + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_baseline_replica_keeps_serving_while_the_candidate_migrates( self, containers: Containers, baseline_image: str, baseline_database: Database ) -> None: @@ -38,6 +44,11 @@ class TestRollingUpgrade: assert CACHED_PLAN not in old.logs(), "The baseline replica hit a stale prepared statement" assert old.state().Running, "The baseline replica died during the upgrade" + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_both_releases_serve_and_share_keys_during_the_overlap( self, containers: Containers, baseline_image: str, baseline_database: Database ) -> None: diff --git a/tests/e2e/migrations/test_shaped_database.py b/tests/e2e/migrations/test_shaped_database.py index 20c4368ae33..07cc4b69597 100644 --- a/tests/e2e/migrations/test_shaped_database.py +++ b/tests/e2e/migrations/test_shaped_database.py @@ -5,6 +5,7 @@ import pytest from .containers import Containers, ready from .database import Database from .upgrade import assert_history_clean, assert_upgraded, confirm, migration_names, provision +from e2e_metadata import Domain, Subject, meta SPEND_ROWS: Final = 20_000 @@ -22,6 +23,11 @@ def seed_spend_logs(database: Database, rows: int) -> None: class TestPopulatedDatabaseUpgrade: + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_upgrade_completes_and_preserves_a_populated_spend_log( self, containers: Containers, baseline_image: str, baseline_database: Database ) -> None: diff --git a/tests/e2e/migrations/test_startup.py b/tests/e2e/migrations/test_startup.py index dc628a8ed7a..ff7112c3881 100644 --- a/tests/e2e/migrations/test_startup.py +++ b/tests/e2e/migrations/test_startup.py @@ -7,12 +7,18 @@ from .checks import COMPLETE, FATAL, GATED, assert_completed, start_replicas from .containers import Containers, failed, ready, until, waiting from .database import PRISMA_LOCK, Database, Databases, restricted_user from .startup_models import Migration +from e2e_metadata import Domain, Subject, meta pytestmark: Final = [pytest.mark.e2e, pytest.mark.migration_startup] class TestMigrationStartup: @pytest.mark.parametrize("replicas,v2", ((1, True), (3, True), (1, False))) + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_fresh_database(self, containers: Containers, databases: Databases, replicas: int, v2: bool) -> None: with databases.create() as database, ExitStack() as stack: ready(tuple(stack.enter_context(containers.start(database, v2=v2)) for _ in range(replicas)), database) @@ -21,11 +27,21 @@ class TestMigrationStartup: ) == ((0,),) assert database.query("SELECT count(*) > 0 FROM _prisma_migrations") == ((True,),) + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_concurrent_upgrade(self, containers: Containers, database: Database) -> None: with ExitStack() as stack: ready(start_replicas(stack, containers, database, (COMPLETE,)), database) assert_completed(database) + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_waiters_survive_prolonged_contention(self, containers: Containers, database: Database) -> None: with ExitStack() as stack: with database.lock(): @@ -37,6 +53,11 @@ class TestMigrationStartup: ready((owner, *followers), database) assert_completed(database, GATED) + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_lock_deadline_then_restart(self, containers: Containers, database: Database) -> None: history: Final = database.history() with database.lock(PRISMA_LOCK): @@ -51,6 +72,11 @@ class TestMigrationStartup: ready((restarted,), database) assert_completed(database) + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_fatal_sql(self, containers: Containers, database: Database) -> None: with ExitStack() as stack: replicas: Final = start_replicas(stack, containers, database, (FATAL,)) @@ -61,6 +87,11 @@ class TestMigrationStartup: (COMPLETE.name, "%MIGRATION_TEST_FATAL%"), ) == ((1,),) + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_duplicate_object_does_not_hide_incomplete_sql(self, containers: Containers, database: Database) -> None: database.execute( "CREATE TABLE migration_existing (id int PRIMARY KEY); INSERT INTO migration_existing VALUES (42)" @@ -77,6 +108,11 @@ class TestMigrationStartup: ) == ((True,),) @pytest.mark.parametrize("v2", (True, False)) + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_restart_preserves_history_and_data(self, containers: Containers, database: Database, v2: bool) -> None: history: Final = database.history() before: Final = database.query('SELECT token FROM "LiteLLM_VerificationToken" ORDER BY token') @@ -86,12 +122,22 @@ class TestMigrationStartup: assert database.history() == history assert set(before).issubset(database.query('SELECT token FROM "LiteLLM_VerificationToken" ORDER BY token')) + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_disabled_migrations(self, containers: Containers, database: Database) -> None: history: Final = database.history() with containers.start(database, (FATAL,), disabled=True) as replica: ready((replica,), database) assert database.history() == history + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_insufficient_privileges(self, containers: Containers, database: Database) -> None: history: Final = database.history() with restricted_user(database) as limited: diff --git a/tests/e2e/migrations/test_upgrade.py b/tests/e2e/migrations/test_upgrade.py index 23f0bbe9124..e9a29c970e9 100644 --- a/tests/e2e/migrations/test_upgrade.py +++ b/tests/e2e/migrations/test_upgrade.py @@ -7,11 +7,17 @@ from .checks import start_replicas from .containers import Containers, ready from .database import Database from .upgrade import assert_history_clean, assert_upgraded, confirm, migration_names, provision +from e2e_metadata import Domain, Subject, meta pytestmark: Final = [pytest.mark.e2e, pytest.mark.migration_startup] class TestReleaseUpgrade: + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_candidate_applies_the_pending_release_migrations( self, containers: Containers, baseline_database: Database ) -> None: @@ -21,6 +27,11 @@ class TestReleaseUpgrade: assert_upgraded(before, migration_names(baseline_database)) assert_history_clean(baseline_database) + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_upgrade_preserves_keys_minted_by_the_baseline_release( self, containers: Containers, baseline_image: str, baseline_database: Database ) -> None: @@ -34,6 +45,11 @@ class TestReleaseUpgrade: assert_upgraded(before, migration_names(baseline_database)) confirm(new, key, alias) + @meta( + Subject( + domain=Domain.DB, + ) + ) def test_concurrent_replicas_upgrade_a_baseline_database_once( self, containers: Containers, baseline_database: Database ) -> None: diff --git a/tests/e2e/migrations/upgrade.py b/tests/e2e/migrations/upgrade.py index 2123f86450a..3dc6c348587 100644 --- a/tests/e2e/migrations/upgrade.py +++ b/tests/e2e/migrations/upgrade.py @@ -8,6 +8,7 @@ from typing import Final from uuid import uuid4 from e2e_http import Result, Success, unwrap +from e2e_metadata import step from models import ( KeyGenerateBody, KeyGenerateResponse, @@ -24,6 +25,7 @@ from .database import Database CACHED_PLAN: Final = "cached plan must not change result type" +@step("Generate a virtual key on the proxy container") def provision(replica: Replica) -> tuple[str, str]: alias: Final = f"upgrade-{uuid4().hex}" key: Final = unwrap( @@ -37,6 +39,7 @@ def provision(replica: Replica) -> tuple[str, str]: return key, alias +@step("Check that the key {alias} resolves on the proxy container through /key/info") def confirm(replica: Replica, key: str, alias: str) -> None: info: Final = unwrap( replica.transport.get( @@ -62,6 +65,7 @@ class Outcomes: self.failures.append(result.model_dump_json()) +@step("Send /v1/models requests with the virtual key to proxy container {replica.name} in the background") @contextmanager def auth_traffic(replica: Replica, key: str, interval: float = 0.05) -> Generator[Outcomes]: outcomes: Final = Outcomes() @@ -93,6 +97,7 @@ def auth_traffic(replica: Replica, key: str, interval: float = 0.05) -> Generato ) +@step("Wait for {calls} more successful /v1/models calls from {description}") def keep_serving(outcomes: Outcomes, description: str, calls: int = 20) -> int: target: Final = outcomes.served + calls until(description, lambda: outcomes.served >= target or bool(outcomes.failures)) @@ -100,10 +105,12 @@ def keep_serving(outcomes: Outcomes, description: str, calls: int = 20) -> int: return outcomes.served +@step("Read the applied migration names from _prisma_migrations") def migration_names(database: Database) -> frozenset[str]: return frozenset(str(row[0]) for row in database.query("SELECT migration_name FROM _prisma_migrations")) +@step("Check that _prisma_migrations holds no unfinished, rolled-back or duplicated migration") def assert_history_clean(database: Database) -> None: assert database.query( "SELECT count(*) FROM _prisma_migrations WHERE finished_at IS NULL OR rolled_back_at IS NOT NULL" diff --git a/tests/e2e/other/other_client.py b/tests/e2e/other/other_client.py index d7bad4f1ed1..17c5bcfd069 100644 --- a/tests/e2e/other/other_client.py +++ b/tests/e2e/other/other_client.py @@ -17,6 +17,7 @@ from dataclasses import dataclass from typing import Final from e2e_http import AnthropicHeaders, AuthHeaders, NoBody, ProbeResult, Result +from e2e_metadata import step from idp import Keycloak, keycloak_from_env from models import ( ChatBody, @@ -56,11 +57,13 @@ class OtherClient: """Resolved per use, so the suite's non-JWT tests never need the IdP env.""" return keycloak_from_env() + @step("Call /health/liveliness without credentials") def liveness(self) -> ProbeResult: """GET /health/liveliness. Unauthenticated; the probe returns status + raw body so the test can assert the worker reports itself alive.""" return self.proxy.transport.probe("/health/liveliness", params=NoBody()) + @step("Call /health/readiness without credentials") def readiness_public(self) -> Result[ReadinessResponse]: """GET /health/readiness with no credential at all, proving the probe is safe to expose to an unauthenticated load balancer.""" @@ -71,6 +74,7 @@ class OtherClient: response_type=ReadinessResponse, ) + @step("Call /health/readiness/details with the given key") def readiness_details(self, key: str) -> Result[ReadinessDetailsResponse]: return self.proxy.transport.get( "/health/readiness/details", @@ -79,6 +83,7 @@ class OtherClient: response_type=ReadinessDetailsResponse, ) + @step("Call /health/readiness/details without credentials") def readiness_details_unauthenticated(self) -> Result[ReadinessDetailsResponse]: return self.proxy.transport.get( "/health/readiness/details", @@ -87,6 +92,7 @@ class OtherClient: response_type=ReadinessDetailsResponse, ) + @step("Create the {body.user_role} user {body.user_email} through /user/new") 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.""" @@ -97,6 +103,7 @@ class OtherClient: response_type=UserNewResponse, ) + @step("Read the user's keys from /user/info") 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.""" @@ -107,6 +114,7 @@ class OtherClient: response_type=UserInfoWithKeysResponse, ) + @step("List the JWT-to-key mappings from /jwt/key/mapping/list") def jwt_mapping_list(self) -> Result[JwtKeyMappingListResponse]: """GET /jwt/key/mapping/list under the master key.""" return self.proxy.transport.get( @@ -116,6 +124,7 @@ class OtherClient: response_type=JwtKeyMappingListResponse, ) + @step("Delete the JWT-to-key mapping") def jwt_mapping_delete(self, mapping_id: str) -> Result[JwtKeyMappingDeleteResponse]: """POST /jwt/key/mapping/delete under the master key.""" return self.proxy.transport.post( @@ -125,6 +134,7 @@ class OtherClient: response_type=JwtKeyMappingDeleteResponse, ) + @step("Send a /chat/completions request to {body.model} as team {team} with the given token") 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( @@ -137,6 +147,7 @@ class OtherClient: response_type=ChatResponse, ) + @step("List the models from /v1/models with the given token, in the Anthropic shape: {anthropic}") def list_models_as(self, token: str, *, anthropic: bool = False) -> Result[ModelsListResponse]: """GET /v1/models under `token`, in the OpenAI shape or, with `anthropic`, the Anthropic Models API shape Claude Code reads. Both carry `data[].id`.""" @@ -148,6 +159,7 @@ class OtherClient: response_type=ModelsListResponse, ) + @step("List users from /user/list with the given key") def list_users_as(self, key: str) -> Result[UserListResponse]: """GET /user/list under `key`. Admin-only, so it doubles as the master key's authorization proof: the master key (proxy admin) reads it, a diff --git a/tests/e2e/other/owned_jwt_gateway.py b/tests/e2e/other/owned_jwt_gateway.py index b6c31479bd2..17d49f591d7 100644 --- a/tests/e2e/other/owned_jwt_gateway.py +++ b/tests/e2e/other/owned_jwt_gateway.py @@ -21,6 +21,7 @@ from typing import Final from e2e_config import INHERITED_ENV_PREFIXES, available_port from e2e_http import NoBody +from e2e_metadata import step from idp import Keycloak, stop_process_group from proxy_client import ProxyClient, build_proxy_client @@ -36,6 +37,7 @@ class OwnedJwtGateway: _log_path: Path _child: subprocess.Popen[bytes] | None = field(default=None, init=False, repr=False) + @step("Start the dedicated JWT proxy and wait for /health/liveliness") def start(self) -> None: with self._log_path.open("ab") as log: self._child = subprocess.Popen( @@ -54,12 +56,14 @@ class OwnedJwtGateway: time.sleep(0.5) raise AssertionError("owned JWT gateway did not become ready") + @step("Stop the dedicated JWT proxy") 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" +@step("Boot a dedicated proxy {name} with its own litellm_jwtauth config") def owned_jwt_gateway( idp: Keycloak, directory: Path, cleanup: ExitStack, *, litellm_jwtauth: str, name: str ) -> OwnedJwtGateway: diff --git a/tests/e2e/other/test_health_lifecycle_e2e.py b/tests/e2e/other/test_health_lifecycle_e2e.py index 2551352e8fa..3a63851a504 100644 --- a/tests/e2e/other/test_health_lifecycle_e2e.py +++ b/tests/e2e/other/test_health_lifecycle_e2e.py @@ -17,12 +17,19 @@ import pytest from e2e_config import MASTER_KEY from e2e_http import UnauthorizedError, unwrap from other_client import OtherClient +from e2e_metadata import Domain, Route, Subject, meta pytestmark = pytest.mark.e2e class TestHealthLifecycle: @pytest.mark.covers("other.lifecycle.liveness.ping") + @meta( + Subject( + domain=Domain.DEPLOY_OPS, + route=Route.HEALTH, + ) + ) def test_liveness_reports_alive_without_auth(self, client: OtherClient) -> None: probe = client.liveness() assert probe.status_code == 200, ( @@ -34,6 +41,12 @@ class TestHealthLifecycle: ) @pytest.mark.covers("other.lifecycle.readiness.public_probe") + @meta( + Subject( + domain=Domain.DEPLOY_OPS, + route=Route.HEALTH, + ) + ) def test_readiness_is_reachable_without_credentials(self, client: OtherClient) -> None: readiness = unwrap(client.readiness_public()) assert readiness.status == "healthy", ( @@ -41,6 +54,12 @@ class TestHealthLifecycle: ) @pytest.mark.covers("other.lifecycle.readiness.reports_db_status") + @meta( + Subject( + domain=Domain.DEPLOY_OPS, + route=Route.HEALTH, + ) + ) def test_readiness_reports_connected_db(self, client: OtherClient) -> None: readiness = unwrap(client.readiness_public()) assert readiness.db == "connected", ( @@ -49,6 +68,12 @@ class TestHealthLifecycle: ) @pytest.mark.covers("other.lifecycle.readiness_details.authenticated_diagnostics") + @meta( + Subject( + domain=Domain.DEPLOY_OPS, + route=Route.HEALTH, + ) + ) def test_readiness_details_require_auth_and_expose_diagnostics(self, client: OtherClient) -> None: anonymous = client.readiness_details_unauthenticated() assert isinstance(anonymous, UnauthorizedError), ( diff --git a/tests/e2e/other/test_jwt_auth_e2e.py b/tests/e2e/other/test_jwt_auth_e2e.py index a8aa474f04f..6da3924c284 100644 --- a/tests/e2e/other/test_jwt_auth_e2e.py +++ b/tests/e2e/other/test_jwt_auth_e2e.py @@ -15,6 +15,7 @@ from lifecycle import ResourceManager from models import ChatBody, ChatMessage, TeamNewBody from other_client import OtherClient from pydantic import BaseModel +from e2e_metadata import Domain, Mode, Provider, Subject, meta pytestmark = pytest.mark.e2e @@ -120,6 +121,14 @@ def _corrupt_signature(token: str) -> str: class TestJwtAuth: @pytest.mark.covers("other.auth.jwt.valid_token_allows", "other.auth.jwt.spend_attributed_to_claims") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + providers=(Provider.OPENAI,), + models=(CHEAP_OPENAI_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_valid_token_for_an_existing_team_is_accepted_and_attributed( self, client: OtherClient, identity: Identity ) -> None: @@ -142,6 +151,13 @@ class TestJwtAuth: ) @pytest.mark.covers("other.auth.jwt.invalid_signature_denied") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + models=(CHEAP_OPENAI_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_tampered_signature_is_rejected(self, client: OtherClient, identity: Identity) -> None: tampered: Final = _corrupt_signature(client.idp.access_token(identity)) @@ -154,6 +170,13 @@ class TestJwtAuth: ) @pytest.mark.covers("other.auth.jwt.expired_denied") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + models=(CHEAP_OPENAI_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_expired_token_is_rejected(self, client: OtherClient, identity: Identity) -> None: expiring: Final = client.idp.access_token(identity, client_id=SHORT_LIVED_CLIENT_ID) delay: Final = _claims(expiring).exp - time.time() + 1 @@ -167,6 +190,13 @@ class TestJwtAuth: assert "expired" in result.body.lower(), f"the 401 must say the token expired, got {result.body[:300]}" @pytest.mark.covers("other.auth.jwt.wrong_issuer_denied") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + models=(CHEAP_OPENAI_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_signed_token_from_the_wrong_issuer_is_rejected(self, client: OtherClient, identity: Identity) -> None: token: Final = client.idp.access_token(identity, issuer_host="unexpected-issuer.invalid") claims: Final = _claims(token) @@ -177,6 +207,13 @@ class TestJwtAuth: assert "issuer" in result.body.lower(), f"expected issuer validation to reject the token: {result}" @pytest.mark.covers("other.auth.jwt.wrong_audience_denied") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + models=(CHEAP_OPENAI_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_signed_token_for_another_application_is_rejected(self, client: OtherClient, identity: Identity) -> None: token: Final = client.idp.access_token(identity, client_id=WRONG_AUDIENCE_CLIENT_ID) claims: Final = _claims(token) @@ -189,6 +226,13 @@ class TestJwtAuth: assert "audience" in result.body.lower(), f"expected audience validation to reject the token: {result}" @pytest.mark.covers("other.auth.jwt.unknown_team_denied") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + models=(CHEAP_OPENAI_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_token_naming_a_team_that_does_not_exist_is_rejected( self, client: OtherClient, resources: ResourceManager ) -> None: @@ -204,6 +248,14 @@ class TestJwtAuth: ) @pytest.mark.covers("other.auth.jwt.virtual_key_unaffected") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + providers=(Provider.OPENAI,), + models=(CHEAP_OPENAI_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_plain_virtual_key_still_works_with_jwt_auth_enabled(self, client: OtherClient, scoped_key: str) -> None: response: Final = unwrap(client.proxy.chat(scoped_key, _ping())) assert response.choices, f"an sk- key must keep working on a proxy with enable_jwt_auth, got {response}" @@ -230,6 +282,14 @@ def _denial(client: OtherClient, token: str, team: str) -> str: class TestJwtTeamHeader: @pytest.mark.covers("other.auth.jwt.team_header_alias_binds_team") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + providers=(Provider.OPENAI,), + models=(CHEAP_OPENAI_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_team_header_with_the_team_alias_binds_the_same_team_as_the_team_id( self, client: OtherClient, bound_team: BoundTeam ) -> None: @@ -249,6 +309,14 @@ class TestJwtTeamHeader: @pytest.mark.covers("other.auth.jwt.team_model_alias_listed_and_routes") @pytest.mark.parametrize("anthropic", [False, True], ids=["openai_shape", "anthropic_shape"]) + @meta( + Subject( + domain=Domain.PROXY_AUTH, + providers=(Provider.OPENAI,), + models=(CHEAP_OPENAI_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_team_model_alias_is_listed_by_v1_models_under_the_same_token_that_routes_it( self, client: OtherClient, aliased_team: AliasedTeam, anthropic: bool ) -> None: @@ -265,6 +333,13 @@ class TestJwtTeamHeader: assert aliased_team.target in listed, f"the alias target {aliased_team.target!r} must stay listed, got {listed}" @pytest.mark.covers("other.auth.jwt.team_header_non_member_alias_denied") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + models=(CHEAP_OPENAI_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_team_header_with_the_alias_of_a_team_the_caller_is_not_in_is_rejected_like_an_unknown_value( self, client: OtherClient, resources: ResourceManager, bound_team: BoundTeam ) -> None: diff --git a/tests/e2e/other/test_jwt_auto_register_e2e.py b/tests/e2e/other/test_jwt_auto_register_e2e.py index 8f7c2c6a693..45a7ce2c182 100644 --- a/tests/e2e/other/test_jwt_auto_register_e2e.py +++ b/tests/e2e/other/test_jwt_auto_register_e2e.py @@ -22,6 +22,7 @@ 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 +from e2e_metadata import Domain, Mode, Provider, Subject, meta pytestmark = pytest.mark.e2e @@ -115,6 +116,14 @@ def minting_gateway(idp: Keycloak, tmp_path_factory: pytest.TempPathFactory) -> @pytest.mark.owned_gateway class TestJwtAutoRegisterMapExistingKey: @pytest.mark.covers("other.auth.jwt.auto_register_maps_existing_key") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + providers=(Provider.GEMINI,), + models=(MODEL_NAME,), + mode=Mode.NONSTREAM, + ) + ) 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: @@ -144,6 +153,14 @@ class TestJwtAutoRegisterMapExistingKey: ) @pytest.mark.covers("other.auth.jwt.auto_register_mints_when_keyless") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + providers=(Provider.GEMINI,), + models=(MODEL_NAME,), + mode=Mode.NONSTREAM, + ) + ) def test_first_jwt_call_mints_a_key_when_the_user_has_none( self, client: OtherClient, idp: Keycloak, resources: ResourceManager, mapping_gateway: OwnedJwtGateway ) -> None: @@ -160,6 +177,14 @@ class TestJwtAutoRegisterMapExistingKey: ) @pytest.mark.covers("other.auth.jwt.auto_register_default_mints") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + providers=(Provider.GEMINI,), + models=(MODEL_NAME,), + mode=Mode.NONSTREAM, + ) + ) def test_default_behavior_still_mints_when_the_user_already_has_a_key( self, client: OtherClient, idp: Keycloak, resources: ResourceManager, minting_gateway: OwnedJwtGateway ) -> None: diff --git a/tests/e2e/other/test_master_key_auth_e2e.py b/tests/e2e/other/test_master_key_auth_e2e.py index 6ab33c9b62a..cd506f746ba 100644 --- a/tests/e2e/other/test_master_key_auth_e2e.py +++ b/tests/e2e/other/test_master_key_auth_e2e.py @@ -15,12 +15,18 @@ import pytest from e2e_config import MASTER_KEY, unique_marker from e2e_http import UnauthorizedError, unwrap from other_client import OtherClient +from e2e_metadata import Domain, Subject, meta pytestmark = pytest.mark.e2e class TestMasterKeyAuth: @pytest.mark.covers("other.auth.master_key.valid_allows") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + ) + ) def test_master_key_authenticates_and_grants_admin_route(self, client: OtherClient) -> None: listing = unwrap(client.list_users_as(MASTER_KEY)) assert listing.total >= 0, ( @@ -29,6 +35,11 @@ class TestMasterKeyAuth: ) @pytest.mark.covers("other.auth.master_key.invalid_denied") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + ) + ) def test_non_matching_master_key_is_denied(self, client: OtherClient) -> None: bogus = f"sk-{unique_marker()}" result = client.list_users_as(bogus) diff --git a/tests/e2e/other/test_session_token_e2e.py b/tests/e2e/other/test_session_token_e2e.py index 51791278026..16740390fd6 100644 --- a/tests/e2e/other/test_session_token_e2e.py +++ b/tests/e2e/other/test_session_token_e2e.py @@ -20,6 +20,7 @@ from e2e_http import UnauthorizedError, unwrap from lifecycle import ResourceManager from models import KeyGenerateBody, KeyLoggingCallback, KeyLoggingCallbackVars, KeyMetadata from other_client import OtherClient +from e2e_metadata import Domain, Subject, meta pytestmark = pytest.mark.e2e @@ -47,12 +48,22 @@ def _admin_session_token(expires_at: datetime) -> str: class TestSessionToken: @pytest.mark.covers("other.auth.session_token.valid_allows") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + ) + ) def test_unexpired_session_token_reaches_admin_route(self, client: OtherClient) -> None: token: Final = _admin_session_token(datetime.now(timezone.utc) + timedelta(minutes=10)) listing: Final = unwrap(client.list_users_as(token)) assert listing.total >= 0, f"an unexpired admin session token did not reach /user/list: {listing}" @pytest.mark.covers("other.auth.session_token.expired_denied") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + ) + ) def test_expired_session_token_is_denied(self, client: OtherClient) -> None: token: Final = _admin_session_token(datetime.now(timezone.utc) - timedelta(minutes=1)) result: Final = client.list_users_as(token) @@ -60,6 +71,11 @@ class TestSessionToken: assert "expired" in result.body.lower(), f"expected the expired-key error, got {result.body[:300]}" @pytest.mark.covers("other.auth.session_token.encrypted_value_denied") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + ) + ) def test_encrypted_stored_value_is_not_a_bearer_token( self, client: OtherClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/secret_manager/secret_store_cyberark.py b/tests/e2e/secret_manager/secret_store_cyberark.py index 375b997e02c..22c7c458d86 100644 --- a/tests/e2e/secret_manager/secret_store_cyberark.py +++ b/tests/e2e/secret_manager/secret_store_cyberark.py @@ -10,6 +10,7 @@ from urllib.parse import quote import pytest import yaml from e2e_http import ExternalWrite, Headers, send_text_external +from e2e_metadata import step from pydantic import Field from secret_store import SecretBackend @@ -94,6 +95,7 @@ class Conjur: if not result.ok: pytest.fail(f"Conjur refused to {action}: HTTP {result.status_code} {result.body[:300]}") + @step("Write the secret {name} to CyberArk Conjur") def write(self, name: str, value: str) -> None: self._update_root_policy("POST", f"- !variable {_policy_scalar(name)}\n", f"declare {name}") result: Final = send_text_external("POST", self._secret_url(name), headers=self._headers(), content=value) @@ -101,6 +103,7 @@ class Conjur: if not result.ok: pytest.fail(f"Conjur refused to write {name}: HTTP {result.status_code} {result.body[:300]}") + @step("Read the secret {name} from CyberArk Conjur") def read(self, name: str) -> str | None: result: Final = send_text_external("GET", self._secret_url(name), headers=self._headers()) self._fail_unless_reached(result, f"read {name}") @@ -110,6 +113,7 @@ class Conjur: pytest.fail(f"Conjur refused to read {name}: HTTP {result.status_code} {result.body[:300]}") return result.body + @step("Delete the secret {name} from CyberArk Conjur") def destroy(self, name: str) -> None: self._update_root_policy("PATCH", f"- !delete\n record: !variable {_policy_scalar(name)}\n", f"destroy {name}") diff --git a/tests/e2e/secret_manager/secret_store_hashicorp_vault.py b/tests/e2e/secret_manager/secret_store_hashicorp_vault.py index ccf8cefe716..5954719ccbb 100644 --- a/tests/e2e/secret_manager/secret_store_hashicorp_vault.py +++ b/tests/e2e/secret_manager/secret_store_hashicorp_vault.py @@ -14,6 +14,7 @@ from e2e_http import ( get_external, post_json_external, ) +from e2e_metadata import step from pydantic import BaseModel, Field from secret_store import SecretBackend @@ -68,6 +69,7 @@ class Vault: def _metadata_url(self, name: str) -> str: return f"{self.base_url}/v1/{self.mount}/metadata/{name}" + @step("Write the secret {name} to HashiCorp Vault") def write(self, name: str, value: str) -> None: write: Final = post_json_external( self._data_url(name), headers=self._headers(), json=KvWriteBody(data=KvData(key=value)) @@ -77,6 +79,7 @@ class Vault: if not write.ok: pytest.fail(f"Vault refused to write {name}: HTTP {write.status_code} {write.body[:300]}") + @step("Read the secret {name} from HashiCorp Vault") def read(self, name: str) -> str | None: result: Final = get_external(self._data_url(name), headers=self._headers(), response_type=KvReadResponse) match result: @@ -89,6 +92,7 @@ class Vault: case _: return pytest.fail(f"Vault refused to read {name}: {result}") + @step("Delete the secret {name} from HashiCorp Vault") def destroy(self, name: str) -> None: write: Final = delete_external(self._metadata_url(name), headers=self._headers()) if not write.ok and write.status_code != 404: diff --git a/tests/e2e/secret_manager/test_secret_manager_e2e.py b/tests/e2e/secret_manager/test_secret_manager_e2e.py index a9c9024718d..8eaabcaae6f 100644 --- a/tests/e2e/secret_manager/test_secret_manager_e2e.py +++ b/tests/e2e/secret_manager/test_secret_manager_e2e.py @@ -13,6 +13,7 @@ from lifecycle import ResourceManager from models import ChatBody, ChatMessage, ChatResponse, KeyGenerateBody, LiteLLMParamsBody from proxy_client import ProxyClient from secret_store import SecretStore +from e2e_metadata import Domain, Mode, Provider, Subject, meta pytestmark = [pytest.mark.e2e, pytest.mark.secret_manager] @@ -73,6 +74,14 @@ def _eventually(proxy: ProxyClient, read: Callable[[], str | None], expected: st class TestSecretManager: @pytest.mark.covers("other.config.secret_resolution.kms_integration") + @meta( + Subject( + domain=Domain.DEPLOY_OPS, + providers=(Provider.OPENAI,), + models=(BACKEND_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_deployment_key_resolves_from_the_manager( self, proxy: ProxyClient, resources: ResourceManager, store: SecretStore, scoped_key: str ) -> None: @@ -83,6 +92,14 @@ class TestSecretManager: assert response.choices, f"the manager-backed deployment answered with no choices: {response}" @pytest.mark.covers("other.config.secret_resolution.manager_value_used") + @meta( + Subject( + domain=Domain.DEPLOY_OPS, + providers=(Provider.OPENAI,), + models=(BACKEND_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_deployment_uses_the_value_the_manager_holds( self, proxy: ProxyClient, resources: ResourceManager, store: SecretStore, scoped_key: str ) -> None: @@ -100,6 +117,11 @@ class TestSecretManager: pytest.fail(f"expected the provider to reject the manager-held key with 401, got {result}") @pytest.mark.covers("other.config.secret_manager.virtual_key_stored") + @meta( + Subject( + domain=Domain.DEPLOY_OPS, + ) + ) def test_generated_key_is_written_to_the_manager( self, proxy: ProxyClient, resources: ResourceManager, store: SecretStore ) -> None: @@ -113,6 +135,11 @@ class TestSecretManager: @pytest.mark.requires_capability("deletes_stored_keys") @pytest.mark.covers("other.config.secret_manager.virtual_key_deleted") + @meta( + Subject( + domain=Domain.DEPLOY_OPS, + ) + ) def test_deleted_key_is_removed_from_the_manager( self, proxy: ProxyClient, resources: ResourceManager, store: SecretStore ) -> None: diff --git a/tests/integration/_support/wire.py b/tests/integration/_support/wire.py index 10bb7787cc9..d118da03daf 100644 --- a/tests/integration/_support/wire.py +++ b/tests/integration/_support/wire.py @@ -47,6 +47,13 @@ class Wire: return self.connected.qsize() +def _await_release(gate: threading.Event, closing: threading.Event) -> bool: + while not closing.is_set(): + if gate.wait(timeout=0.05): + return True + return gate.is_set() + + @contextmanager def wire_server( respond: Callable[[Request], Reply], @@ -60,6 +67,7 @@ def wire_server( errors: Final[SimpleQueue[Exception]] = SimpleQueue() disconnected: Final[SimpleQueue[str]] = SimpleQueue() connected: Final[SimpleQueue[str]] = SimpleQueue() + closing: Final = threading.Event() class Handler(BaseHTTPRequestHandler): protocol_version = "HTTP/1.1" @@ -108,8 +116,12 @@ def wire_server( break self.wfile.write(b"%x\r\n%s\r\n" % (len(chunk), chunk)) self.wfile.flush() - if index == 0 and reply.gate_after_first is not None: - assert reply.gate_after_first.wait(timeout=5), "Stream barrier was never released" + if ( + index == 0 + and reply.gate_after_first is not None + and not _await_release(reply.gate_after_first, closing) + ): + break if reply.pause_between_chunks and index + 1 < len(reply.chunks): time.sleep(reply.pause_between_chunks) else: @@ -151,6 +163,7 @@ def wire_server( connected, ) finally: + closing.set() server.shutdown() thread.join(timeout=6) assert not thread.is_alive(), "Owned HTTP server survived cleanup" diff --git a/tests/integration/coordination_redis_proxy_config.yaml b/tests/integration/coordination_redis_proxy_config.yaml index 30294c291bf..e28e0f5f2ad 100644 --- a/tests/integration/coordination_redis_proxy_config.yaml +++ b/tests/integration/coordination_redis_proxy_config.yaml @@ -3,6 +3,7 @@ general_settings: master_key: os.environ/LITELLM_MASTER_KEY database_url: os.environ/DATABASE_URL store_model_in_db: true + disable_model_info_refresh: true disable_spend_logs: false proxy_batch_write_at: 1 coordination_redis: diff --git a/tests/integration/database/test_lens_repository.py b/tests/integration/database/test_lens_repository.py index dd9f467fd76..a46e37668ab 100644 --- a/tests/integration/database/test_lens_repository.py +++ b/tests/integration/database/test_lens_repository.py @@ -14,6 +14,7 @@ import pytest_asyncio from fastapi import HTTPException from prisma import Prisma from psycopg import sql +from pydantic import TypeAdapter from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.db.prisma_client import PrismaWrapper @@ -35,7 +36,7 @@ from litellm.proxy.lens.models import ( Worker, ) from litellm.proxy.lens.repository import LensRepository, WriterDatabase -from litellm.proxy.lens.state import claim_job, queue_job +from litellm.proxy.lens.state import cancel_job, claim_job, current_job, due_at, end_job, queue_job, replace_job @pytest_asyncio.fixture(loop_scope="function") @@ -44,6 +45,274 @@ async def lens_db() -> AsyncIterator[Prisma]: yield db +def _scheduled_lens( + lens_id: str, + scope: Scope, + now: datetime, + next_run_at: datetime, + *, + enabled: bool = True, + jobs: tuple[Job, ...] = (), +) -> Lens: + return Lens( + id=lens_id, + scope=scope, + settings=LensSettings( + name="Scheduling test", + model="analysis", + context="Find unexpected behavior", + enabled=enabled, + ), + created_at=now, + next_run_at=next_run_at, + jobs=jobs, + budget_month=now.strftime("%Y-%m"), + ) + + +def _stored_due_at(lens_id: str) -> datetime | None: + with psycopg.connect(os.environ["DATABASE_URL"]) as connection: + row: Final = connection.execute('SELECT due_at FROM "LiteLLM_Lens" WHERE id=%s', (lens_id,)).fetchone() + return TypeAdapter(datetime | None).validate_python(row[0]) if row else None + + +async def _assert_due_column(repo: LensRepository, lens_id: str) -> None: + stored: Final = await repo.get(lens_id) + assert stored is not None + expected: Final = due_at(stored) + actual: Final = _stored_due_at(lens_id) + if expected is None: + assert actual is None + return + assert actual is not None + difference: Final = actual.replace(tzinfo=timezone.utc) - expected.astimezone(timezone.utc) + assert abs(difference.total_seconds()) <= 0.001 + + +@pytest.mark.asyncio +async def test_due_filters_by_schedule_and_scope(lens_db: Prisma) -> None: + utc_now: Final = datetime.now(timezone.utc).replace(microsecond=0) + worker_now: Final = utc_now.astimezone(timezone(timedelta(hours=3))) + team_id: Final = uuid4().hex + worker_scope: Final = Scope(team_id=team_id) + worker: Final = Worker(id=uuid4().hex, name="worker", scope=worker_scope, last_seen=worker_now) + repo: Final = LensRepository(WriterDatabase(PrismaWrapper(lens_db))) + due_lens: Final = _scheduled_lens(uuid4().hex, worker_scope, utc_now, utc_now - timedelta(minutes=20)) + future_lens: Final = _scheduled_lens(uuid4().hex, worker_scope, utc_now, utc_now + timedelta(minutes=20)) + disabled_lens: Final = _scheduled_lens( + uuid4().hex, worker_scope, utc_now, utc_now - timedelta(minutes=10), enabled=False + ) + live_queued: Final = queue_job( + _scheduled_lens(uuid4().hex, worker_scope, utc_now - timedelta(minutes=5), utc_now - timedelta(minutes=5)), + utc_now - timedelta(minutes=5), + uuid4().hex, + ) + live_lens: Final = claim_job(live_queued, worker, worker_now) + expired_queued: Final = queue_job( + _scheduled_lens(uuid4().hex, worker_scope, utc_now - timedelta(minutes=10), utc_now - timedelta(minutes=10)), + utc_now - timedelta(minutes=10), + uuid4().hex, + ) + expired_claimed: Final = claim_job(expired_queued, worker, utc_now - timedelta(minutes=10)) + expired_job: Final = expired_claimed.jobs[0].model_copy(update={"lease_until": utc_now - timedelta(minutes=5)}) + expired_lens: Final = expired_claimed.model_copy(update={"jobs": (expired_job,)}) + other_lens: Final = _scheduled_lens( + uuid4().hex, Scope(team_id=uuid4().hex), utc_now, utc_now - timedelta(minutes=3) + ) + worker_key: Final = uuid4().hex + key_lens: Final = _scheduled_lens( + uuid4().hex, Scope(api_key_hash=worker_key), utc_now, utc_now - timedelta(minutes=2) + ) + candidates: Final = (due_lens, future_lens, disabled_lens, live_lens, expired_lens, other_lens, key_lens) + await asyncio.gather(*(repo.create(candidate) for candidate in candidates)) + try: + await lens_db.execute_raw( + """UPDATE "LiteLLM_Lens" + SET data=jsonb_set(data, '{scope}', jsonb_build_object('team_id', $2)) + WHERE id=$1""", + due_lens.id, + team_id, + ) + await lens_db.execute_raw( + """UPDATE "LiteLLM_Lens" + SET data=jsonb_set(data, '{scope}', jsonb_build_object('api_key_hash', $2)) + WHERE id=$1""", + key_lens.id, + worker_key, + ) + team_due: Final = await repo.due(worker_scope, worker_now, 20) + assert tuple(candidate.lens.id for candidate in team_due) == tuple( + lens.id for lens in sorted((due_lens, expired_lens), key=lambda lens: (due_at(lens), lens.id)) + ) + assert team_due[0].lens.scope == worker_scope + key_due: Final = await repo.due(Scope(api_key_hash=worker_key), worker_now, 20) + assert tuple(candidate.lens.id for candidate in key_due) == (key_lens.id,) + all_due: Final = await repo.due(Scope(all_teams=True), worker_now, 20) + assert {candidate.lens.id for candidate in all_due} == { + due_lens.id, + expired_lens.id, + other_lens.id, + key_lens.id, + } + finally: + await lens_db.execute_raw( + 'DELETE FROM "LiteLLM_Lens" WHERE id=ANY($1::text[])', + tuple(lens.id for lens in candidates), + ) + + +@pytest.mark.asyncio +async def test_due_pages_lenses_with_equal_due_at_without_skipping_or_repeating(lens_db: Prisma) -> None: + now: Final = datetime.now(timezone.utc).replace(microsecond=0) + scope: Final = Scope(team_id=uuid4().hex) + repo: Final = LensRepository(WriterDatabase(PrismaWrapper(lens_db))) + lenses: Final = tuple(_scheduled_lens(uuid4().hex, scope, now, now - timedelta(minutes=1)) for _ in range(45)) + await asyncio.gather(*(repo.create(lens) for lens in lenses)) + try: + await lens_db.execute_raw( + """UPDATE "LiteLLM_Lens" + SET due_at=$2::timestamp + WHERE id=ANY($1::text[])""", + tuple(lens.id for lens in lenses), + "1970-01-01 00:00:00", + ) + first: Final = await repo.due(scope, now, 20) + second: Final = await repo.due(scope, now, 20, first[-1]) + third: Final = await repo.due(scope, now, 20, second[-1]) + assert tuple(len(page) for page in (first, second, third)) == (20, 20, 5) + ids: Final = tuple(candidate.lens.id for candidate in (*first, *second, *third)) + assert ids == tuple(sorted(lens.id for lens in lenses)) + finally: + await lens_db.execute_raw( + 'DELETE FROM "LiteLLM_Lens" WHERE id=ANY($1::text[])', + tuple(lens.id for lens in lenses), + ) + + +@pytest.mark.asyncio +async def test_due_at_stays_consistent_through_job_lifecycle(lens_db: Prisma) -> None: + now: Final = datetime.now(timezone.utc).replace(microsecond=0) + scope: Final = Scope(team_id=uuid4().hex) + lens: Final = _scheduled_lens(uuid4().hex, scope, now, now) + repo: Final = LensRepository(WriterDatabase(PrismaWrapper(lens_db))) + worker: Final = Worker(id=uuid4().hex, name="worker", scope=scope, last_seen=now) + await repo.create(lens) + try: + await _assert_due_column(repo, lens.id) + job_id: Final = uuid4().hex + claimed: Final = await repo.update( + lens.id, + lambda candidate: claim_job(queue_job(candidate, now, job_id), worker, now), + attempts=1, + ) + assert claimed is not None + await _assert_due_column(repo, lens.id) + active: Final = current_job(claimed) + assert active is not None + progressed: Final = await repo.progress(lens.id, active, Progress()) + assert progressed is not None + await _assert_due_column(repo, lens.id) + result_at: Final = datetime.now(timezone.utc) + + def finish(candidate: Lens) -> Lens: + active_job: Final = current_job(candidate) + if active_job is None: + return candidate + return replace_job(candidate, end_job(active_job, "completed", result_at)).model_copy( + update={"next_run_at": result_at + timedelta(minutes=candidate.settings.interval_minutes)} + ) + + completed: Final = await repo.update(lens.id, finish, attempts=1) + assert completed is not None + await _assert_due_column(repo, lens.id) + cancelled_at: Final = datetime.now(timezone.utc) + cancelled: Final = await repo.update( + lens.id, + lambda candidate: cancel_job( + queue_job(candidate, cancelled_at, uuid4().hex, trigger="manual"), + cancelled_at, + ), + attempts=1, + ) + assert cancelled is not None + await _assert_due_column(repo, lens.id) + finally: + await lens_db.execute_raw('DELETE FROM "LiteLLM_Lens" WHERE id=$1', lens.id) + + +@pytest.mark.asyncio +async def test_sync_due_repairs_legacy_rows_and_ignores_stale_versions(lens_db: Prisma) -> None: + now: Final = datetime.now(timezone.utc).replace(microsecond=0) + team_id: Final = uuid4().hex + scope: Final = Scope(team_id=team_id) + worker: Final = Worker(id=uuid4().hex, name="worker", scope=scope, last_seen=now) + repo: Final = LensRepository(WriterDatabase(PrismaWrapper(lens_db))) + due_idle: Final = _scheduled_lens(uuid4().hex, scope, now, now - timedelta(minutes=20)) + future_idle: Final = _scheduled_lens(uuid4().hex, scope, now, now + timedelta(minutes=20)) + disabled_idle: Final = _scheduled_lens(uuid4().hex, scope, now, now - timedelta(minutes=10), enabled=False) + queued_lens: Final = queue_job( + _scheduled_lens(uuid4().hex, scope, now, now + timedelta(minutes=20), enabled=False), + now - timedelta(minutes=3), + uuid4().hex, + trigger="manual", + ) + live_lens: Final = claim_job( + queue_job( + _scheduled_lens(uuid4().hex, scope, now, now + timedelta(minutes=20)), + now - timedelta(minutes=10), + uuid4().hex, + ), + worker, + now, + ) + expired_claimed: Final = claim_job( + queue_job( + _scheduled_lens(uuid4().hex, scope, now, now + timedelta(minutes=20)), + now - timedelta(minutes=10), + uuid4().hex, + ), + worker, + now - timedelta(minutes=10), + ) + expired_lens: Final = expired_claimed.model_copy( + update={"jobs": (expired_claimed.jobs[0].model_copy(update={"lease_until": now - timedelta(minutes=5)}),)} + ) + candidates: Final = (due_idle, future_idle, disabled_idle, queued_lens, live_lens, expired_lens) + await asyncio.gather(*(repo.create(candidate) for candidate in candidates)) + try: + past: Final = now - timedelta(hours=1) + await lens_db.execute_raw( + """UPDATE "LiteLLM_Lens" + SET due_at=($2::timestamptz AT TIME ZONE 'UTC') + WHERE id=ANY($1::text[])""", + tuple(lens.id for lens in candidates), + past.isoformat(), + ) + legacy_due: Final = await repo.due(scope, now, 20) + assert {candidate.lens.id for candidate in legacy_due} == {lens.id for lens in candidates} + for candidate in legacy_due: + await repo.sync_due(candidate.lens) + repaired_due: Final = await repo.due(scope, now, 20) + assert {candidate.lens.id for candidate in repaired_due} == {due_idle.id, queued_lens.id, expired_lens.id} + await asyncio.gather(*(_assert_due_column(repo, lens.id) for lens in candidates)) + stale: Final = await repo.get(future_idle.id) + assert stale is not None + await lens_db.execute_raw( + """UPDATE "LiteLLM_Lens" + SET version=version+1, due_at=($2::timestamptz AT TIME ZONE 'UTC') + WHERE id=$1""", + stale.id, + past.isoformat(), + ) + await repo.sync_due(stale) + assert _stored_due_at(stale.id) == past.replace(tzinfo=None) + finally: + await lens_db.execute_raw( + 'DELETE FROM "LiteLLM_Lens" WHERE id=ANY($1::text[])', + tuple(lens.id for lens in candidates), + ) + + @pytest.mark.asyncio async def test_concurrent_workers_cannot_both_acquire_the_same_job(lens_db: Prisma) -> None: now: Final = datetime.now(timezone.utc) @@ -146,7 +415,10 @@ async def test_claim_pages_only_yield_work_the_worker_can_claim(lens_db: Prisma) try: for row in rows: await repo.create(row) - found: Final = tuple([candidate async for candidate in repo.claim_candidates(scope, now)]) + first: Final = await repo.due(scope, now, 50) + second: Final = await repo.due(scope, now, 50, first[-1]) + assert len(first) == 50 + found: Final = tuple(candidate.lens for candidate in (*first, *second)) assert frozenset(candidate.id for candidate in found) == frozenset( candidate.id for candidate in (*queued[1:], due, expired) ) @@ -358,6 +630,19 @@ def test_db_push_creates_fresh_lens_tables_and_preserves_them_on_restart(monkeyp assert connection.execute( sql.SQL("SELECT id, data FROM {}").format(sql.Identifier(schema, "LiteLLM_Lens")) ).fetchall() == [("saved", {"keep": True})] + assert ( + connection.execute( + sql.SQL("SELECT due_at FROM {} WHERE id='saved'").format(sql.Identifier(schema, "LiteLLM_Lens")) + ).fetchone()[0] + is not None + ) + due_index: Final = connection.execute( + """SELECT indexdef FROM pg_indexes + WHERE schemaname=%s AND tablename='LiteLLM_Lens' AND indexname='LiteLLM_Lens_due_at_idx'""", + (schema,), + ).fetchone() + assert due_index is not None + assert "WHERE" not in due_index[0] finally: connection.execute(sql.SQL("DROP SCHEMA {} CASCADE").format(sql.Identifier(schema))) diff --git a/tests/integration/database/test_lens_scheduler_load.py b/tests/integration/database/test_lens_scheduler_load.py new file mode 100644 index 00000000000..0afa41bea74 --- /dev/null +++ b/tests/integration/database/test_lens_scheduler_load.py @@ -0,0 +1,199 @@ +import asyncio +import json +import os +import sys +from collections.abc import AsyncIterator +from contextlib import AbstractAsyncContextManager +from datetime import datetime, timedelta, timezone +from time import perf_counter +from typing import Final +from uuid import uuid4 + +import pytest +import pytest_asyncio +from prisma import Prisma +from pydantic import TypeAdapter +from typing_extensions import LiteralString + +from litellm.proxy.db.prisma_client import PrismaWrapper +from litellm.proxy.lens.endpoints import claim_due +from litellm.proxy.lens.models import Evidence, Finding, Lens, LensSettings, Scope, Worker +from litellm.proxy.lens.repository import Database, LensRepository, Row, WriterDatabase +from litellm.proxy.lens.state import current_job + + +@pytest_asyncio.fixture(loop_scope="function") +async def lens_db() -> AsyncIterator[Prisma]: + async with Prisma(datasource={"url": os.environ["DATABASE_URL"]}) as db: + yield db + + +class ReadMeter: + def __init__(self) -> None: + self.batches: tuple[tuple[int, ...], ...] = () + + def record(self, document_sizes: tuple[int, ...]) -> None: + self.batches = (*self.batches, document_sizes) + + @property + def document_count(self) -> int: + return sum(len(batch) for batch in self.batches) + + @property + def total_bytes(self) -> int: + return sum(sum(batch) for batch in self.batches) + + +class MeasuredDatabase: + def __init__(self, database: WriterDatabase, meter: ReadMeter) -> None: + self.database: Final = database + self.meter: Final = meter + + async def query_raw(self, query: LiteralString, *args: object) -> object: + rows: Final = await self.database.query_raw(query, *args) + if 'FROM "LiteLLM_Lens"' in query and "WHERE id" not in query: + documents: Final = TypeAdapter(tuple[Row, ...]).validate_python(rows) + self.meter.record( + tuple(len(json.dumps(row.data, separators=(",", ":")).encode("utf-8")) for row in documents) + ) + return rows + + async def execute_raw(self, query: LiteralString, *args: object) -> int: + return await self.database.execute_raw(query, *args) + + def transaction(self) -> AbstractAsyncContextManager[Database]: + return self.database.transaction() + + +def _large_lens(lens_id: str, scope: Scope, now: datetime, next_run_at: datetime) -> Lens: + findings: Final = tuple( + Finding( + id=f"f{index}", + title=f"Issue {index}", + description="Repeated operation returns an unexpected result.", + check_id="behavior", + evidence=( + Evidence( + execution_id=f"t{index}", + span_id=f"s{index}", + quote="Unexpected result", + ), + ), + first_seen=now, + last_seen=now, + revision=1, + ) + for index in range(100) + ) + return Lens( + id=lens_id, + scope=scope, + settings=LensSettings( + name="Claim scheduler load", + model="analysis", + context="Find unexpected behavior", + enabled=True, + ), + created_at=now, + next_run_at=next_run_at, + findings=findings, + budget_month=now.strftime("%Y-%m"), + ) + + +def _due_lens(lens_id: str, scope: Scope, now: datetime, model: str, next_run_at: datetime) -> Lens: + return Lens( + id=lens_id, + scope=scope, + settings=LensSettings( + name="Claim paging test", + model=model, + context="Find unexpected behavior", + enabled=True, + ), + created_at=now, + next_run_at=next_run_at, + budget_month=now.strftime("%Y-%m"), + ) + + +async def _supports_model(_worker: Worker, _settings: LensSettings) -> bool: + return True + + +async def _supports_supported_model(_worker: Worker, settings: LensSettings) -> bool: + return settings.model == "supported" + + +@pytest.mark.asyncio +async def test_claim_due_reaches_a_supported_lens_behind_a_full_page_of_unsupported_ones( + lens_db: Prisma, +) -> None: + now: Final = datetime.now(timezone.utc).replace(microsecond=0) + scope: Final = Scope(team_id=uuid4().hex) + worker: Final = Worker(id=uuid4().hex, name="paging-test-worker", scope=scope, last_seen=now) + unsupported_at: Final = now - timedelta(minutes=5) + supported_at: Final = now - timedelta(minutes=1) + unsupported: Final = tuple(_due_lens(uuid4().hex, scope, now, "unsupported", unsupported_at) for _ in range(25)) + supported: Final = _due_lens(uuid4().hex, scope, now, "supported", supported_at) + candidates: Final = (*unsupported, supported) + repository: Final = LensRepository(WriterDatabase(PrismaWrapper(lens_db))) + await asyncio.gather(*(repository.create(candidate) for candidate in candidates)) + try: + claim: Final = await claim_due(worker, now, repository, _supports_supported_model) + assert claim is not None + assert claim.lens_id == supported.id + assert claim.job.status == "running" + finally: + await lens_db.execute_raw( + 'DELETE FROM "LiteLLM_Lens" WHERE id=ANY($1::text[])', + tuple(candidate.id for candidate in candidates), + ) + + +@pytest.mark.asyncio +async def test_lens_claim_reads_scale_with_due_lenses_not_total_lenses(lens_db: Prisma) -> None: + now: Final = datetime.now(timezone.utc).replace(microsecond=0) + scope: Final = Scope(team_id=uuid4().hex) + worker: Final = Worker(id=uuid4().hex, name="load-test-worker", scope=scope, last_seen=now) + due_lens: Final = _large_lens(uuid4().hex, scope, now, now - timedelta(seconds=1)) + initial_future: Final = tuple(_large_lens(uuid4().hex, scope, now, now + timedelta(days=1)) for _ in range(20)) + additional_future: Final = tuple(_large_lens(uuid4().hex, scope, now, now + timedelta(days=1)) for _ in range(200)) + ids: Final = tuple(lens.id for lens in (due_lens, *initial_future, *additional_future)) + writer: Final = WriterDatabase(PrismaWrapper(lens_db)) + seed_repository: Final = LensRepository(writer) + await asyncio.gather(*(seed_repository.create(lens) for lens in (due_lens, *initial_future))) + try: + before_meter: Final = ReadMeter() + before_repository: Final = LensRepository(MeasuredDatabase(writer, before_meter)) + before_started: Final = perf_counter() + before_claim: Final = await claim_due(worker, now, before_repository, _supports_model) + before_seconds: Final = perf_counter() - before_started + assert before_claim is not None + assert before_claim.lens_id == due_lens.id + assert before_claim.job.status == "running" + claimed_lens: Final = await seed_repository.get(due_lens.id) + assert claimed_lens is not None + assert current_job(claimed_lens) == before_claim.job + await seed_repository.update( + due_lens.id, + lambda lens: lens.model_copy(update={"jobs": (), "next_run_at": now - timedelta(seconds=1)}), + attempts=1, + ) + await asyncio.gather(*(seed_repository.create(lens) for lens in additional_future)) + after_meter: Final = ReadMeter() + after_repository: Final = LensRepository(MeasuredDatabase(writer, after_meter)) + after_started: Final = perf_counter() + after_claim: Final = await claim_due(worker, now, after_repository, _supports_model) + after_seconds: Final = perf_counter() - after_started + assert after_claim is not None + assert after_claim.lens_id == due_lens.id + assert after_claim.job.status == "running" + sys.stdout.write( + f"claim read: before={before_meter.total_bytes} bytes, {before_seconds:.4f}s; " + f"after={after_meter.total_bytes} bytes, {after_seconds:.4f}s\n" + ) + assert before_meter.document_count == after_meter.document_count == 1 + assert before_meter.total_bytes == after_meter.total_bytes + finally: + await lens_db.execute_raw('DELETE FROM "LiteLLM_Lens" WHERE id=ANY($1::text[])', ids) diff --git a/tests/integration/observability/test_langtrace_delivery.py b/tests/integration/observability/test_langtrace_delivery.py index 84c9ed5bec0..03e42087dc7 100644 --- a/tests/integration/observability/test_langtrace_delivery.py +++ b/tests/integration/observability/test_langtrace_delivery.py @@ -132,9 +132,8 @@ def _upstream(request: Request) -> Reply: def _config(tmp_path: Path, **litellm_settings: object) -> Path: config: Final = _PROXY_CONFIG.validate_python(yaml.safe_load(STOCK_CONFIG.read_text())) settings: Final = {**_SETTINGS.validate_python(config["litellm_settings"]), **litellm_settings} - general: Final = {**_SETTINGS.validate_python(config["general_settings"]), "disable_model_info_refresh": True} path: Final = tmp_path / "langtrace.yaml" - path.write_text(yaml.safe_dump({**config, "litellm_settings": settings, "general_settings": general})) + path.write_text(yaml.safe_dump({**config, "litellm_settings": settings})) return path diff --git a/tests/integration/observability/test_otel_conversation_id.py b/tests/integration/observability/test_otel_conversation_id.py index 78ff40927e5..e77e5603148 100644 --- a/tests/integration/observability/test_otel_conversation_id.py +++ b/tests/integration/observability/test_otel_conversation_id.py @@ -289,7 +289,7 @@ class RigFactory: def start(self) -> Iterator[Rig]: config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) config["litellm_settings"].update({"callbacks": ["otel"]}) - config["general_settings"].update({"disable_model_info_refresh": True, **self.settings}) + config["general_settings"].update(self.settings) config["callback_settings"] = { "otel": {"exporter": "http/json", "endpoint": self.sink.wire.url, "mapper_names": ["genai"]}, } diff --git a/tests/integration/observability/test_signoz_delivery.py b/tests/integration/observability/test_signoz_delivery.py index f3d715fe5cf..56a98b3f17c 100644 --- a/tests/integration/observability/test_signoz_delivery.py +++ b/tests/integration/observability/test_signoz_delivery.py @@ -394,7 +394,6 @@ class RigFactory: "callbacks": ["signoz"], "provider_url_destination_allowed_hosts": [self.tenant_sink.wire.url], }, - "general_settings": {**object_value(loaded["general_settings"]), "disable_model_info_refresh": True}, } path: Final = self.directory / f"signoz-{uuid.uuid4().hex}.yaml" path.write_text(yaml.safe_dump(config)) diff --git a/tests/integration/oci_proxy_test_config.yaml b/tests/integration/oci_proxy_test_config.yaml index 95e74963cd2..f09216c4f94 100644 --- a/tests/integration/oci_proxy_test_config.yaml +++ b/tests/integration/oci_proxy_test_config.yaml @@ -19,6 +19,7 @@ model_list: general_settings: master_key: os.environ/LITELLM_MASTER_KEY + disable_model_info_refresh: true litellm_settings: drop_params: True diff --git a/tests/integration/providers/test_realtime_transcription_guardrail_session_update.py b/tests/integration/providers/test_realtime_transcription_guardrail_session_update.py index b60fbcdbdeb..fd8879e02d4 100644 --- a/tests/integration/providers/test_realtime_transcription_guardrail_session_update.py +++ b/tests/integration/providers/test_realtime_transcription_guardrail_session_update.py @@ -6,6 +6,12 @@ every assistant turn. A transcription session has no assistant turn and the vend injection is skipped when the route intent or a backend session event says the session is transcription-only, while a client frame alone never flips a voice session into one. Every row runs against the scripted upstream, which answers the injected and rewritten updates the way the vendors do. + +A blocked transcript on a transcription session reaches the client as a ``guardrail_violation`` error and +nothing else: the ``response.cancel``, the voiced block prompt and the ``response.create`` a voice session gets +never go to a backend that cannot speak, while ``on_violation: end_session`` and ``end_session_after_n_fails`` +still close the session. A guardrail that raises anything but a block closes the session the way it does on a +voice session. """ from __future__ import annotations @@ -130,6 +136,69 @@ APPEND: Final[dict[str, JsonValue]] = {"type": "input_audio_buffer.append", "aud PUSH_TO_TALK: Final = "PUSH_TO_TALK" ENDPOINTING: Final = "ENDPOINTING" END_STREAM: Final[dict[str, JsonValue]] = {"type": "endStream"} +RESPONSE_CREATE: Final[dict[str, JsonValue]] = {"type": "response.create"} +TRANSCRIPTION_UPDATE_FRAME: Final[dict[str, JsonValue]] = {"type": "session.update", "session": GA_TRANSCRIPTION_UPDATE} +VOICE_UPDATE_FRAME: Final[dict[str, JsonValue]] = {"type": "session.update", "session": VOICE_UPDATE} +REALTIME_HOOK: Final = "realtime_input_transcription" +ENDER_GUARDRAIL: Final = "transcript-ender" +ENDER_WORD: Final = "anchovy" +ENDER_MESSAGE: Final = "The session was ended by the transcript policy." +TWO_STRIKES_GUARDRAIL: Final = "transcript-two-strikes" +TWO_STRIKES_WORD: Final = "olives" +ONE_STRIKE_GUARDRAIL: Final = "transcript-one-strike" +ONE_STRIKE_WORD: Final = "radish" +WARN_GUARDRAIL: Final = "transcript-warn" +WARN_WORD: Final = "capers" +VALUE_ERROR_GUARDRAIL: Final = "transcript-value-error" +VALUE_ERROR_WORD: Final = "durian" +RUNTIME_ERROR_GUARDRAIL: Final = "transcript-runtime-error" +RUNTIME_ERROR_WORD: Final = "lychee" +RAISING_MODULE: Final = "raising_guardrails" +CONFIGURED_GUARDRAILS: Final = frozenset( + { + TRANSCRIPT_GUARDRAIL, + OPTIN_GUARDRAIL, + PROMPT_GUARDRAIL, + ENDER_GUARDRAIL, + TWO_STRIKES_GUARDRAIL, + ONE_STRIKE_GUARDRAIL, + WARN_GUARDRAIL, + VALUE_ERROR_GUARDRAIL, + RUNTIME_ERROR_GUARDRAIL, + } +) +RELAYED_CLOSE: Final = ("server_error", "") +GUARDRAIL_END_CLOSE: Final = 1000 +PROXY_FAILURE_CLOSE: Final = 1011 +PROXY_FAILURE_REASON: Final = "proxy failed while relaying the upstream websocket" +TOOL_OUTPUT_BLOCKED: Final = json.dumps({"error": "Tool output blocked by content policy"}) +VOICE_BLOCK_FRAMES: Final = ("response.cancel", "conversation.item.create", "response.create") +CREATED_SECONDS: Final = 60.0 +STEP_SECONDS: Final = 15.0 +SDK_SECONDS: Final = 45.0 +RAISING_GUARDRAILS_SOURCE: Final = f"""\ +from litellm.integrations.custom_guardrail import CustomGuardrail + + +class WordRaiser(CustomGuardrail): + word = "" + error = Exception + + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): + if any(self.word in str(text) for text in inputs["texts"]): + raise self.error(f"{{self.word}} is not allowed") + return inputs + + +class ValueErrorGuardrail(WordRaiser): + word = "{VALUE_ERROR_WORD}" + error = ValueError + + +class RuntimeErrorGuardrail(WordRaiser): + word = "{RUNTIME_ERROR_WORD}" + error = RuntimeError +""" @dataclass(frozen=True, slots=True) @@ -183,7 +252,15 @@ def _owned_url(owned: OwnedProxy) -> str: return str(owned.gateway.client.base_url).rstrip("/") -def _transcript_event(transcript: str) -> dict[str, JsonValue]: +def _blocked(word: str) -> str: + return f"please add {word} to the order" + + +def _optin(guardrail: str) -> str: + return f"guardrails={guardrail}" + + +def _transcript_event(transcript: JsonValue) -> dict[str, JsonValue]: return { "type": TRANSCRIPT_COMPLETED, "event_id": "evt_$UNIQUE_ID", @@ -208,15 +285,20 @@ def _done_event() -> dict[str, JsonValue]: } -def _transcription_scenario(transcript: str, *, repeats: int = 1) -> RealtimeResponse: +def _transcript_event_without_the_field() -> dict[str, JsonValue]: + return {key: value for key, value in _transcript_event("").items() if key != "transcript"} + + +def _transcription_events(events: tuple[dict[str, JsonValue], ...], *, repeats: int = 1) -> RealtimeResponse: return RealtimeResponse( - content_type="application/x-realtime", - events=(_transcript_event(transcript),), - session_type=TRANSCRIPTION, - created_repeats=repeats, + content_type="application/x-realtime", events=events, session_type=TRANSCRIPTION, created_repeats=repeats ) +def _transcription_scenario(*transcripts: JsonValue, repeats: int = 1) -> RealtimeResponse: + return _transcription_events(tuple(_transcript_event(transcript) for transcript in transcripts), repeats=repeats) + + def _older_transcription_scenario(transcript: str) -> RealtimeResponse: return RealtimeResponse( content_type="application/x-realtime", @@ -226,10 +308,14 @@ def _older_transcription_scenario(transcript: str) -> RealtimeResponse: ) -def _voice_scenario(transcript: str) -> RealtimeResponse: - return RealtimeResponse( - content_type="application/x-realtime", events=(_transcript_event(transcript), _done_event()) - ) +def _voice_turns(transcripts: tuple[JsonValue, ...]) -> Iterator[dict[str, JsonValue]]: + for transcript in transcripts: + yield _transcript_event(transcript) + yield _done_event() + + +def _voice_scenario(*transcripts: JsonValue) -> RealtimeResponse: + return RealtimeResponse(content_type="application/x-realtime", events=tuple(_voice_turns(transcripts))) def _muse_scenario(transcript: str) -> RealtimeResponse: @@ -345,6 +431,97 @@ async def _session( return Session((), refusal.response.status_code) +@dataclass(frozen=True, slots=True) +class Step: + frames: tuple[dict[str, JsonValue], ...] + until: str + + +def _step(*frames: dict[str, JsonValue], until: str) -> Step: + return Step(frames, until) + + +def _update_frame(update: JsonValue) -> dict[str, JsonValue]: + return {"type": "session.update", "session": update} + + +def _probe(update: JsonValue = GA_TRANSCRIPTION_UPDATE) -> Step: + return Step((_update_frame(update),), "session.updated") + + +async def _until(socket: ClientConnection, until: str, seconds: float) -> AsyncIterator[dict[str, JsonValue]]: + deadline: Final = asyncio.get_running_loop().time() + seconds + while True: + event: Final = await _next_event(socket, deadline) + if event is None: + yield {"type": "timeout"} + return + yield event + if event.get("type") == until: + return + + +async def _stepped( + socket: ClientConnection, steps: tuple[Step, ...], seconds: float +) -> AsyncIterator[dict[str, JsonValue]]: + try: + first: Final = JSON_OBJECT.validate_json(await asyncio.wait_for(socket.recv(), CREATED_SECONDS)) + yield first + if first.get("type") not in CREATED_TYPES: + async for message in socket: + yield JSON_OBJECT.validate_json(message) + return + for step in steps: + for frame in step.frames: + await socket.send(json.dumps(frame)) + async for event in _until(socket, step.until, seconds): + yield event + if event["type"] == "timeout": + return + except ConnectionClosed as closed: + yield {"type": "closed", "code": _close_code(closed)} + + +async def _driven( + ws_base: str, + path: str, + query: str, + key: str, + steps: tuple[Step, ...], + headers: Mapping[str, str] | None, + seconds: float, +) -> Session: + request_headers: Final = {"Authorization": f"Bearer {key}", **(headers or {})} + async with websockets.connect(f"{ws_base}{path}?{query}", additional_headers=request_headers) as socket: + return Session(tuple([event async for event in _stepped(socket, steps, seconds)]), None) + + +def _drive( + ws_base: str, + query: str, + key: str, + steps: tuple[Step, ...], + *, + path: str = "/v1/realtime", + headers: Mapping[str, str] | None = None, + seconds: float = STEP_SECONDS, +) -> Session: + return asyncio.run(_driven(ws_base, path, query, key, steps, headers, seconds)) + + +def _transcribe_blocked( + ws_base: str, + query: str, + key: str, + *, + path: str = "/v1/realtime", + update: JsonValue = GA_TRANSCRIPTION_UPDATE, + headers: Mapping[str, str] | None = None, +) -> Session: + steps: Final = (_step(_update_frame(update), COMMIT, until="error"), _probe(update)) + return _drive(ws_base, query, key, steps, path=path, headers=headers) + + def _transcribe( ws_base: str, query: str, @@ -439,27 +616,70 @@ def _spend_rows(key: str, count: int) -> list[dict[str, JsonValue]]: ) -def _content_filter(name: str, mode: str, *, default_on: bool) -> dict[str, JsonValue]: +def _content_filter( + name: str, mode: str, *, default_on: bool, word: str = BLOCKED_WORD, **settings: JsonValue +) -> dict[str, JsonValue]: return { "guardrail_name": name, "litellm_params": { "guardrail": "litellm_content_filter", "mode": mode, "default_on": default_on, - "blocked_words": [{"keyword": BLOCKED_WORD, "action": "BLOCK"}], + "blocked_words": [{"keyword": word, "action": "BLOCK"}], + **settings, }, } +def _raising_guardrail(name: str, class_name: str) -> dict[str, JsonValue]: + return { + "guardrail_name": name, + "litellm_params": {"guardrail": f"{RAISING_MODULE}.{class_name}", "mode": REALTIME_HOOK, "default_on": False}, + } + + def _write_config(directory: Path, ca_bundle: Path) -> Path: config: Final = directory / f"realtime_guardrails_{uuid.uuid4().hex[:8]}.yaml" + (directory / f"{RAISING_MODULE}.py").write_text(RAISING_GUARDRAILS_SOURCE) config.write_text( json.dumps( { "guardrails": [ - _content_filter(TRANSCRIPT_GUARDRAIL, "realtime_input_transcription", default_on=True), - _content_filter(OPTIN_GUARDRAIL, "realtime_input_transcription", default_on=False), + _content_filter(TRANSCRIPT_GUARDRAIL, REALTIME_HOOK, default_on=True), + _content_filter(OPTIN_GUARDRAIL, REALTIME_HOOK, default_on=False), _content_filter(PROMPT_GUARDRAIL, "pre_call", default_on=True), + _content_filter( + ENDER_GUARDRAIL, + REALTIME_HOOK, + default_on=False, + word=ENDER_WORD, + on_violation="end_session", + realtime_violation_message=ENDER_MESSAGE, + ), + _content_filter( + TWO_STRIKES_GUARDRAIL, + REALTIME_HOOK, + default_on=False, + word=TWO_STRIKES_WORD, + end_session_after_n_fails=2, + ), + _content_filter( + ONE_STRIKE_GUARDRAIL, + REALTIME_HOOK, + default_on=False, + word=ONE_STRIKE_WORD, + end_session_after_n_fails=1, + ), + _content_filter( + WARN_GUARDRAIL, + REALTIME_HOOK, + default_on=False, + word=WARN_WORD, + on_violation="warn", + end_session_after_n_fails=None, + ), + _raising_guardrail(VALUE_ERROR_GUARDRAIL, "ValueErrorGuardrail"), + _raising_guardrail(RUNTIME_ERROR_GUARDRAIL, "RuntimeErrorGuardrail"), ], "general_settings": { "master_key": "os.environ/LITELLM_MASTER_KEY", @@ -524,6 +744,35 @@ def _assert_transcription_left_alone( assert _session_updates(observed) == (update,), _session_updates(observed) +def _assert_transcription_blocked( + session: Session, observed: tuple[dict[str, JsonValue], ...], transcript: str, *, update: JsonValue +) -> None: + assert session.types == ("session.created", "session.updated", TRANSCRIPT_COMPLETED, "error", "session.updated"), ( + session + ) + assert session.session_type == TRANSCRIPTION, session + assert session.transcripts == (transcript,), session + assert session.errors == (GUARDRAIL_VIOLATION,), session + assert _sent_types(observed) == ("session.update", "input_audio_buffer.commit", "session.update"), _sent(observed) + assert _session_updates(observed) == (update, update), _session_updates(observed) + + +def _relayed_close(session: Session) -> str: + return f"upstream websocket closed with code {session.close_code}" + + +def _assert_ended_by_the_guardrail(session: Session, *, violations: int) -> None: + assert session.errors == (GUARDRAIL_VIOLATION,) * violations, session + assert session.types[-1] == "closed", session + assert session.close_code == GUARDRAIL_END_CLOSE, session + + +def _assert_closed_by_a_proxy_failure(session: Session) -> None: + assert session.errors[-1] == RELAYED_CLOSE, session + assert session.close_code == PROXY_FAILURE_CLOSE, session + assert session.error_messages[-1] == f"{_relayed_close(session)}: {PROXY_FAILURE_REASON}", session + + @pytest.mark.parametrize("path", ["/v1/realtime", "/realtime", "/openai/v1/realtime"]) def test_transcription_session_update_reaches_the_upstream_verbatim(guardrail_proxy: OwnedProxy, path: str) -> None: with guardrail_proxy.gateway.scenario() as scenario: @@ -542,17 +791,21 @@ def test_transcription_session_update_reaches_the_upstream_verbatim(guardrail_pr assert rows[0]["call_type"] == "_arealtime", rows -async def _sdk_async_events(connection: AsyncRealtimeConnection) -> AsyncIterator[dict[str, JsonValue]]: +async def _sdk_async_events( + connection: AsyncRealtimeConnection, until: str = TRANSCRIPT_COMPLETED +) -> AsyncIterator[dict[str, JsonValue]]: async for event in connection: yield JSON_OBJECT.validate_python(event.model_dump()) - if event.type == TRANSCRIPT_COMPLETED: + if event.type == until: return -def _sdk_sync_events(connection: RealtimeConnection) -> Iterator[dict[str, JsonValue]]: +def _sdk_sync_events( + connection: RealtimeConnection, until: str = TRANSCRIPT_COMPLETED +) -> Iterator[dict[str, JsonValue]]: for event in connection: yield JSON_OBJECT.validate_python(event.model_dump()) - if event.type == TRANSCRIPT_COMPLETED: + if event.type == until: return @@ -750,6 +1003,14 @@ def test_backend_session_created_typed_transcription_skips_the_injection_on_the_ assert [_query(upgrade) for upgrade in _upgrades(observed)] == [[["model", TRANSCRIBE_MODEL]]] +MUSE_PUSH_TO_TALK_FRAMES: Final[tuple[dict[str, JsonValue], ...]] = ( + {"binary_bytes": MUSE_PACKET_BYTES}, + {"binary_bytes": MUSE_PACKET_BYTES}, + {"binary_bytes": MUSE_REMAINDER_BYTES}, + END_STREAM, +) + + def _muse_handshake(observed: tuple[dict[str, JsonValue], ...]) -> dict[str, JsonValue]: upgrades: Final = _upgrades(observed) assert len(upgrades) == 1, upgrades @@ -767,13 +1028,19 @@ def _muse_session( transcript: str, query: str, turn_detection: JsonValue, + *, + until: str | None = None, ) -> tuple[Session, tuple[dict[str, JsonValue], ...]]: handle: Final = _scripted(scenario, _muse_scenario(transcript), control_url=tls_upstream) key: Final = scenario.key() model: Final = _muse_deployment(scenario, handle.scenario_id, tls_upstream) - until: Final = TRANSCRIPT_COMPLETED if BLOCKED_WORD not in transcript else "error" + verdict: Final = TRANSCRIPT_COMPLETED if BLOCKED_WORD not in transcript else "error" session: Final = _talk( - _ws_base(_owned_url(guardrail_proxy)), f"model={model}{query}", key, _muse_frames(turn_detection), until=until + _ws_base(_owned_url(guardrail_proxy)), + f"model={model}{query}", + key, + _muse_frames(turn_detection), + until=verdict if until is None else until, ) return session, _observed(tls_upstream, handle.scenario_id) @@ -789,12 +1056,7 @@ def test_muse_push_to_talk_transcription_session_keeps_push_to_talk( assert session.transcripts == (CLEAN_TRANSCRIPT,), session assert session.errors == (), session assert _muse_handshake(observed)["mode"] == PUSH_TO_TALK, _muse_handshake(observed) - assert _muse_audio_frames(observed) == ( - {"binary_bytes": MUSE_PACKET_BYTES}, - {"binary_bytes": MUSE_PACKET_BYTES}, - {"binary_bytes": MUSE_REMAINDER_BYTES}, - END_STREAM, - ), _muse_audio_frames(observed) + assert _muse_audio_frames(observed) == MUSE_PUSH_TO_TALK_FRAMES, _muse_audio_frames(observed) def test_muse_transcription_session_blocked_transcript_reaches_the_client_as_a_violation( @@ -807,7 +1069,27 @@ def test_muse_transcription_session_blocked_transcript_reaches_the_client_as_a_v assert session.transcripts == (BLOCKED_TRANSCRIPT,), session assert session.errors == (GUARDRAIL_VIOLATION,), session assert _muse_handshake(observed)["mode"] == PUSH_TO_TALK, _muse_handshake(observed) - assert _muse_audio_frames(observed)[-1] == END_STREAM, _muse_audio_frames(observed) + assert _muse_audio_frames(observed) == MUSE_PUSH_TO_TALK_FRAMES, _muse_audio_frames(observed) + + +def test_muse_transcription_session_on_violation_end_session_closes_the_session( + guardrail_proxy: OwnedProxy, tls_upstream: str +) -> None: + blocked: Final = _blocked(ENDER_WORD) + with guardrail_proxy.gateway.scenario() as scenario: + session, observed = _muse_session( + guardrail_proxy, + scenario, + tls_upstream, + blocked, + f"&{TRANSCRIPTION_QUERY}&{_optin(ENDER_GUARDRAIL)}", + None, + until="closed", + ) + assert session.transcripts == (blocked,), session + _assert_ended_by_the_guardrail(session, violations=1) + assert session.error_messages[0] == ENDER_MESSAGE, session + assert _muse_audio_frames(observed) == MUSE_PUSH_TO_TALK_FRAMES, _muse_audio_frames(observed) def test_muse_server_vad_transcription_session_keeps_endpointing( @@ -949,10 +1231,10 @@ def test_opt_in_guardrail_leaves_a_transcription_session_alone_on_an_opted_out_k _assert_transcription_left_alone(session, observed, CLEAN_TRANSCRIPT, update=GA_TRANSCRIPTION_UPDATE) -def test_guardrails_list_names_the_three_configured_guardrails(guardrail_proxy: OwnedProxy) -> None: +def test_guardrails_list_names_every_configured_guardrail(guardrail_proxy: OwnedProxy) -> None: listed: Final = guardrail_proxy.gateway.get("/guardrails/list") names: Final = {string_value(object_value(entry)["guardrail_name"]) for entry in _list(listed["guardrails"])} - assert names == {TRANSCRIPT_GUARDRAIL, OPTIN_GUARDRAIL, PROMPT_GUARDRAIL}, listed + assert names == CONFIGURED_GUARDRAILS, listed def _list(value: JsonValue) -> list[JsonValue]: @@ -1089,6 +1371,583 @@ def test_repeated_transcription_sessions_write_one_spend_row_each(guardrail_prox assert {str(row["call_type"]) for row in rows} == {"_arealtime"}, rows +@pytest.mark.parametrize("path", ["/v1/realtime", "/realtime", "/openai/v1/realtime"]) +def test_blocked_transcript_on_a_transcription_session_reports_a_violation_and_sends_nothing_upstream( + guardrail_proxy: OwnedProxy, path: str +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(BLOCKED_TRANSCRIPT)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _transcribe_blocked( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}", key, path=path + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + _assert_transcription_blocked(session, observed, BLOCKED_TRANSCRIPT, update=GA_TRANSCRIPTION_UPDATE) + assert [_query(upgrade) for upgrade in _upgrades(observed)] == [ + [["model", TRANSCRIBE_MODEL], ["intent", TRANSCRIPTION]] + ] + rows: Final = _spend_rows(key, 1) + assert rows[0]["call_type"] == "_arealtime", rows + + +async def _sdk_async_blocked_transcription(proxy_url: str, key: str, model: str) -> Session: + client: Final = AsyncOpenAI(api_key=key, base_url=f"{proxy_url}/v1", websocket_base_url=f"{_ws_base(proxy_url)}/v1") + async with client.realtime.connect(model=model, extra_query={"intent": TRANSCRIPTION}) as connection: + await connection.session.update(session=SDK_TRANSCRIPTION_SESSION) + await connection.input_audio_buffer.commit() + verdict: Final = tuple([event async for event in _sdk_async_events(connection, until="error")]) + await connection.session.update(session=SDK_TRANSCRIPTION_SESSION) + probe: Final = tuple([event async for event in _sdk_async_events(connection, until="session.updated")]) + return Session((*verdict, *probe), None) + + +def _sdk_sync_blocked_transcription(proxy_url: str, key: str, model: str) -> Session: + client: Final = OpenAI(api_key=key, base_url=f"{proxy_url}/v1", websocket_base_url=f"{_ws_base(proxy_url)}/v1") + with client.realtime.connect(model=model, extra_query={"intent": TRANSCRIPTION}) as connection: + connection.session.update(session=SDK_TRANSCRIPTION_SESSION) + connection.input_audio_buffer.commit() + verdict: Final = tuple(_sdk_sync_events(connection, until="error")) + connection.session.update(session=SDK_TRANSCRIPTION_SESSION) + probe: Final = tuple(_sdk_sync_events(connection, until="session.updated")) + return Session((*verdict, *probe), None) + + +@pytest.mark.parametrize("client", ["async", pytest.param("sync", marks=pytest.mark.timeout(90))]) +def test_openai_sdk_transcription_session_gets_the_violation_and_stays_open( + guardrail_proxy: OwnedProxy, client: str +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(BLOCKED_TRANSCRIPT)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + proxy_url: Final = _owned_url(guardrail_proxy) + session: Final = ( + asyncio.run(asyncio.wait_for(_sdk_async_blocked_transcription(proxy_url, key, model), SDK_SECONDS)) + if client == "async" + else _sdk_sync_blocked_transcription(proxy_url, key, model) + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + _assert_transcription_blocked(session, observed, BLOCKED_TRANSCRIPT, update=GA_TRANSCRIPTION_UPDATE) + + +def test_beta_protocol_transcription_session_blocked_transcript_reports_a_violation( + guardrail_proxy: OwnedProxy, +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(BLOCKED_TRANSCRIPT)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _transcribe_blocked( + _ws_base(_owned_url(guardrail_proxy)), + f"model={model}&{TRANSCRIPTION_QUERY}", + key, + update=BETA_TRANSCRIPTION_UPDATE, + headers=BETA_HEADERS, + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + _assert_transcription_blocked(session, observed, BLOCKED_TRANSCRIPT, update=BETA_TRANSCRIPTION_UPDATE) + + +def test_intent_without_model_blocked_transcript_reports_a_violation_on_the_whisper_default( + guardrail_proxy: OwnedProxy, +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(BLOCKED_TRANSCRIPT)) + key: Final = scenario.key() + _named_deployment( + guardrail_proxy.gateway, scenario, WHISPER_DEFAULT, f"openai/{WHISPER_DEFAULT}", handle.scenario_id + ) + session: Final = _transcribe_blocked(_ws_base(_owned_url(guardrail_proxy)), TRANSCRIPTION_QUERY, key) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + forwarded: Final = _with_transcription_model(GA_TRANSCRIPTION_UPDATE, WHISPER_DEFAULT) + _assert_transcription_blocked(session, observed, BLOCKED_TRANSCRIPT, update=forwarded) + assert [_query(upgrade) for upgrade in _upgrades(observed)] == [[["intent", TRANSCRIPTION]]] + + +def test_azure_transcription_session_blocked_transcript_reports_a_violation_and_sends_nothing_upstream( + guardrail_proxy: OwnedProxy, tls_upstream: str +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(BLOCKED_TRANSCRIPT), control_url=tls_upstream) + key: Final = scenario.key() + model: Final = _azure_deployment(scenario, handle.scenario_id, tls_upstream) + session: Final = _transcribe_blocked( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}", key + ) + observed: Final = _observed(tls_upstream, handle.scenario_id) + _assert_transcription_blocked(session, observed, BLOCKED_TRANSCRIPT, update=GA_TRANSCRIPTION_UPDATE) + upgrades: Final = tuple(request for request in observed if request["method"] == "WEBSOCKET") + assert [upgrade["path"] for upgrade in upgrades] == ["/openai/v1/realtime"], upgrades + assert [_query(object_value(upgrade["body"])) for upgrade in upgrades] == [[["intent", TRANSCRIPTION]]], ( + upgrades + ) + + +def test_second_violation_under_end_session_after_n_fails_closes_the_transcription_session( + guardrail_proxy: OwnedProxy, +) -> None: + blocked: Final = _blocked(TWO_STRIKES_WORD) + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(blocked, blocked)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + steps: Final = (_step(TRANSCRIPTION_UPDATE_FRAME, COMMIT, until="error"), _step(COMMIT, until="closed")) + session: Final = _drive( + _ws_base(_owned_url(guardrail_proxy)), + f"model={model}&{TRANSCRIPTION_QUERY}&{_optin(TWO_STRIKES_GUARDRAIL)}", + key, + steps, + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ( + "session.created", + "session.updated", + TRANSCRIPT_COMPLETED, + "error", + TRANSCRIPT_COMPLETED, + "error", + "closed", + ), session + assert session.transcripts == (blocked, blocked), session + _assert_ended_by_the_guardrail(session, violations=2) + assert _sent_types(observed) == ( + "session.update", + "input_audio_buffer.commit", + "input_audio_buffer.commit", + ), _sent(observed) + rows: Final = _spend_rows(key, 1) + assert rows[0]["call_type"] == "_arealtime", rows + + +def test_on_violation_end_session_closes_the_transcription_session_with_the_configured_message( + guardrail_proxy: OwnedProxy, +) -> None: + blocked: Final = _blocked(ENDER_WORD) + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(blocked)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _drive( + _ws_base(_owned_url(guardrail_proxy)), + f"model={model}&{TRANSCRIPTION_QUERY}&{_optin(ENDER_GUARDRAIL)}", + key, + (_step(TRANSCRIPTION_UPDATE_FRAME, COMMIT, until="closed"),), + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("session.created", "session.updated", TRANSCRIPT_COMPLETED, "error", "closed"), session + assert session.transcripts == (blocked,), session + _assert_ended_by_the_guardrail(session, violations=1) + assert session.error_messages[0] == ENDER_MESSAGE, session + assert _sent_types(observed) == ("session.update", "input_audio_buffer.commit"), _sent(observed) + + +def test_end_session_after_one_fail_closes_the_transcription_session_on_the_first_violation( + guardrail_proxy: OwnedProxy, +) -> None: + blocked: Final = _blocked(ONE_STRIKE_WORD) + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(blocked)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _drive( + _ws_base(_owned_url(guardrail_proxy)), + f"model={model}&{TRANSCRIPTION_QUERY}&{_optin(ONE_STRIKE_GUARDRAIL)}", + key, + (_step(TRANSCRIPTION_UPDATE_FRAME, COMMIT, until="closed"),), + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("session.created", "session.updated", TRANSCRIPT_COMPLETED, "error", "closed"), session + _assert_ended_by_the_guardrail(session, violations=1) + assert _sent_types(observed) == ("session.update", "input_audio_buffer.commit"), _sent(observed) + + +def _twice_blocked_and_open(guardrail_proxy: OwnedProxy, scenario: Scenario, blocked: str, query: str) -> None: + handle: Final = _scripted(scenario, _transcription_scenario(blocked, blocked)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + steps: Final = ( + _step(TRANSCRIPTION_UPDATE_FRAME, COMMIT, until="error"), + _step(COMMIT, until="error"), + _probe(), + ) + session: Final = _drive( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}{query}", key, steps + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ( + "session.created", + "session.updated", + TRANSCRIPT_COMPLETED, + "error", + TRANSCRIPT_COMPLETED, + "error", + "session.updated", + ), session + assert session.transcripts == (blocked, blocked), session + assert session.errors == (GUARDRAIL_VIOLATION, GUARDRAIL_VIOLATION), session + assert _sent_types(observed) == ( + "session.update", + "input_audio_buffer.commit", + "input_audio_buffer.commit", + "session.update", + ), _sent(observed) + rows: Final = _spend_rows(key, 1) + assert rows[0]["call_type"] == "_arealtime", rows + + +def test_same_blocked_transcript_twice_reports_two_violations_and_keeps_the_session_open( + guardrail_proxy: OwnedProxy, +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + _twice_blocked_and_open(guardrail_proxy, scenario, BLOCKED_TRANSCRIPT, "") + + +def test_on_violation_warn_with_a_null_end_rule_reports_each_violation_and_keeps_the_session_open( + guardrail_proxy: OwnedProxy, +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + _twice_blocked_and_open(guardrail_proxy, scenario, _blocked(WARN_WORD), f"&{_optin(WARN_GUARDRAIL)}") + + +def test_first_configured_guardrail_wins_when_two_match_one_transcript(guardrail_proxy: OwnedProxy) -> None: + transcript: Final = f"please add {BLOCKED_WORD} and {ENDER_WORD} to the order" + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(transcript)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _transcribe_blocked( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}&{_optin(ENDER_GUARDRAIL)}", key + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + _assert_transcription_blocked(session, observed, transcript, update=GA_TRANSCRIPTION_UPDATE) + assert ENDER_MESSAGE not in session.error_messages, session + + +def test_guardrail_raising_value_error_reports_the_exception_text_and_keeps_the_transcription_session_open( + guardrail_proxy: OwnedProxy, +) -> None: + blocked: Final = _blocked(VALUE_ERROR_WORD) + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(blocked)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _transcribe_blocked( + _ws_base(_owned_url(guardrail_proxy)), + f"model={model}&{TRANSCRIPTION_QUERY}&{_optin(VALUE_ERROR_GUARDRAIL)}", + key, + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + _assert_transcription_blocked(session, observed, blocked, update=GA_TRANSCRIPTION_UPDATE) + assert session.error_messages == (f"{VALUE_ERROR_WORD} is not allowed",), session + + +def test_guardrail_raising_runtime_error_closes_the_transcription_session_with_a_proxy_failure( + guardrail_proxy: OwnedProxy, +) -> None: + blocked: Final = _blocked(RUNTIME_ERROR_WORD) + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(blocked)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _drive( + _ws_base(_owned_url(guardrail_proxy)), + f"model={model}&{TRANSCRIPTION_QUERY}&{_optin(RUNTIME_ERROR_GUARDRAIL)}", + key, + (_step(TRANSCRIPTION_UPDATE_FRAME, COMMIT, until="closed"),), + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("session.created", "session.updated", TRANSCRIPT_COMPLETED, "error", "closed"), session + assert session.transcripts == (blocked,), session + _assert_closed_by_a_proxy_failure(session) + assert _sent_types(observed) == ("session.update", "input_audio_buffer.commit"), _sent(observed) + rows: Final = _spend_rows(key, 1) + assert rows[0]["call_type"] == "_arealtime", rows + + +@pytest.mark.parametrize("guardrail", [VALUE_ERROR_GUARDRAIL, RUNTIME_ERROR_GUARDRAIL]) +def test_raising_guardrails_leave_a_clean_transcription_session_alone( + guardrail_proxy: OwnedProxy, guardrail: str +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(CLEAN_TRANSCRIPT)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _transcribe( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}&{_optin(guardrail)}", key + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + _assert_transcription_left_alone(session, observed, CLEAN_TRANSCRIPT, update=GA_TRANSCRIPTION_UPDATE) + + +def _voice_session( + guardrail_proxy: OwnedProxy, scenario: Scenario, response: RealtimeResponse, query: str, steps: tuple[Step, ...] +) -> tuple[Session, tuple[dict[str, JsonValue], ...]]: + handle: Final = _scripted(scenario, response) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id, model=VOICE_MODEL) + session: Final = _drive(_ws_base(_owned_url(guardrail_proxy)), f"model={model}{query}", key, steps) + return session, _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + + +def test_voice_session_guardrail_raising_runtime_error_closes_with_the_same_proxy_failure( + guardrail_proxy: OwnedProxy, +) -> None: + blocked: Final = _blocked(RUNTIME_ERROR_WORD) + with guardrail_proxy.gateway.scenario() as scenario: + session, observed = _voice_session( + guardrail_proxy, + scenario, + _voice_scenario(blocked), + f"&{_optin(RUNTIME_ERROR_GUARDRAIL)}", + (_step(VOICE_UPDATE_FRAME, COMMIT, until="closed"),), + ) + assert session.types == ( + "session.created", + "session.updated", + "error", + TRANSCRIPT_COMPLETED, + "error", + "closed", + ), session + assert session.errors[0] == MISSING_TURN_DETECTION_TYPE, session + _assert_closed_by_a_proxy_failure(session) + assert _sent_types(observed) == ("session.update", "session.update", "input_audio_buffer.commit"), _sent( + observed + ) + + +def test_voice_session_guardrail_raising_value_error_is_voiced_through_the_backend(guardrail_proxy: OwnedProxy) -> None: + blocked: Final = _blocked(VALUE_ERROR_WORD) + with guardrail_proxy.gateway.scenario() as scenario: + session, observed = _voice_session( + guardrail_proxy, + scenario, + _voice_scenario(blocked), + f"&{_optin(VALUE_ERROR_GUARDRAIL)}", + (_step(VOICE_UPDATE_FRAME, COMMIT, until=RESPONSE_DONE),), + ) + assert session.transcripts == (blocked,), session + assert session.errors == (MISSING_TURN_DETECTION_TYPE, GUARDRAIL_VIOLATION), session + assert session.error_messages[1] == f"{VALUE_ERROR_WORD} is not allowed", session + assert session.types[-1] == RESPONSE_DONE, session + assert _sent_types(observed) == ( + "session.update", + "session.update", + "input_audio_buffer.commit", + *VOICE_BLOCK_FRAMES, + ), _sent(observed) + + +def test_voice_session_second_violation_under_end_session_after_n_fails_closes_after_voicing_both( + guardrail_proxy: OwnedProxy, +) -> None: + blocked: Final = _blocked(TWO_STRIKES_WORD) + with guardrail_proxy.gateway.scenario() as scenario: + session, observed = _voice_session( + guardrail_proxy, + scenario, + _voice_scenario(blocked, blocked), + f"&{_optin(TWO_STRIKES_GUARDRAIL)}", + (_step(VOICE_UPDATE_FRAME, COMMIT, until=RESPONSE_DONE), _step(COMMIT, until="closed")), + ) + assert session.transcripts == (blocked, blocked), session + assert session.errors == (MISSING_TURN_DETECTION_TYPE, GUARDRAIL_VIOLATION, GUARDRAIL_VIOLATION), session + assert session.types[-1] == "closed", session + assert session.close_code == GUARDRAIL_END_CLOSE, session + assert _sent_types(observed) == ( + "session.update", + "session.update", + "input_audio_buffer.commit", + *VOICE_BLOCK_FRAMES, + "input_audio_buffer.commit", + *VOICE_BLOCK_FRAMES, + ), _sent(observed) + + +def _user_text_item(text: str) -> dict[str, JsonValue]: + return { + "type": "conversation.item.create", + "item": {"type": "message", "role": "user", "content": [{"type": "input_text", "text": text}]}, + } + + +def _tool_output_item(output: str) -> dict[str, JsonValue]: + return { + "type": "conversation.item.create", + "item": {"type": "function_call_output", "call_id": "call_realtime_guard", "output": output}, + } + + +def test_blocked_user_text_on_a_transcription_session_is_dropped_without_a_voiced_block( + guardrail_proxy: OwnedProxy, +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(CLEAN_TRANSCRIPT)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + steps: Final = ( + _probe(), + _step(_user_text_item(BLOCKED_TRANSCRIPT), RESPONSE_CREATE, until="error"), + _probe(), + ) + session: Final = _drive( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}", key, steps + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("session.created", "session.updated", "error", "session.updated"), session + assert session.errors == (GUARDRAIL_VIOLATION,), session + assert _sent_types(observed) == ("session.update", "session.update"), _sent(observed) + + +def test_blocked_tool_output_on_a_transcription_session_is_sanitized_without_a_voiced_block( + guardrail_proxy: OwnedProxy, +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(CLEAN_TRANSCRIPT)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + steps: Final = ( + _probe(), + _step(_tool_output_item(BLOCKED_TRANSCRIPT), until="error"), + _probe(), + ) + session: Final = _drive( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}", key, steps + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("session.created", "session.updated", "error", "session.updated"), session + assert session.errors == (GUARDRAIL_VIOLATION,), session + assert _sent_types(observed) == ("session.update", "conversation.item.create", "session.update"), _sent( + observed + ) + assert object_value(_sent(observed)[1]["item"])["output"] == TOOL_OUTPUT_BLOCKED, _sent(observed) + + +def test_clean_user_text_on_a_transcription_session_is_forwarded_verbatim(guardrail_proxy: OwnedProxy) -> None: + item: Final = _user_text_item(CLEAN_TRANSCRIPT) + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(CLEAN_TRANSCRIPT)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + steps: Final = (_step(TRANSCRIPTION_UPDATE_FRAME, item, RESPONSE_CREATE, until=TRANSCRIPT_COMPLETED),) + session: Final = _drive( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}", key, steps + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("session.created", "session.updated", TRANSCRIPT_COMPLETED), session + assert session.errors == (), session + assert _sent(observed) == (TRANSCRIPTION_UPDATE_FRAME, item, RESPONSE_CREATE), _sent(observed) + + +CLEAN_TRANSCRIPT_SHAPES: Final = (pytest.param("", id="empty"), pytest.param(FIVE_KB, id="five_kilobytes")) +NON_STRING_TRANSCRIPTS: Final = ( + pytest.param(None, id="null"), + pytest.param(123, id="integer"), + pytest.param(["a"], id="list"), +) + + +@pytest.mark.parametrize("transcript", CLEAN_TRANSCRIPT_SHAPES) +def test_clean_transcript_field_shapes_are_relayed_and_the_session_stays_open( + guardrail_proxy: OwnedProxy, transcript: str +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(transcript)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + steps: Final = (_step(TRANSCRIPTION_UPDATE_FRAME, COMMIT, until=TRANSCRIPT_COMPLETED), _probe()) + session: Final = _drive( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}", key, steps + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("session.created", "session.updated", TRANSCRIPT_COMPLETED, "session.updated"), session + assert session.transcripts == (transcript,), session + assert session.errors == (), session + assert _sent_types(observed) == ("session.update", "input_audio_buffer.commit", "session.update"), _sent( + observed + ) + + +def test_transcript_event_without_the_field_is_relayed_and_the_session_stays_open(guardrail_proxy: OwnedProxy) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_events((_transcript_event_without_the_field(),))) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + steps: Final = (_step(TRANSCRIPTION_UPDATE_FRAME, COMMIT, until=TRANSCRIPT_COMPLETED), _probe()) + session: Final = _drive( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}", key, steps + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("session.created", "session.updated", TRANSCRIPT_COMPLETED, "session.updated"), session + assert "transcript" not in session.events[2], session + assert session.errors == (), session + assert _sent_types(observed) == ("session.update", "input_audio_buffer.commit", "session.update"), _sent( + observed + ) + + +def test_five_kilobyte_transcript_ending_in_the_blocked_word_reports_a_violation(guardrail_proxy: OwnedProxy) -> None: + transcript: Final = f"{FIVE_KB} {BLOCKED_TRANSCRIPT}" + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(transcript)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _transcribe_blocked( + _ws_base(_owned_url(guardrail_proxy)), f"model={model}&{TRANSCRIPTION_QUERY}", key + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + _assert_transcription_blocked(session, observed, transcript, update=GA_TRANSCRIPTION_UPDATE) + + +@pytest.mark.parametrize("transcript", NON_STRING_TRANSCRIPTS) +def test_non_string_transcript_field_closes_the_transcription_session_with_a_proxy_failure( + guardrail_proxy: OwnedProxy, transcript: JsonValue +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + handle: Final = _scripted(scenario, _transcription_scenario(transcript)) + key: Final = scenario.key() + model: Final = _openai_deployment(scenario, handle.scenario_id) + session: Final = _drive( + _ws_base(_owned_url(guardrail_proxy)), + f"model={model}&{TRANSCRIPTION_QUERY}", + key, + (_step(TRANSCRIPTION_UPDATE_FRAME, COMMIT, until="closed"),), + ) + observed: Final = _observed(guardrail_proxy.gateway.upstream_url, handle.scenario_id) + assert session.types == ("session.created", "session.updated", TRANSCRIPT_COMPLETED, "error", "closed"), session + assert session.events[2]["transcript"] == transcript, session + _assert_closed_by_a_proxy_failure(session) + assert _sent_types(observed) == ("session.update", "input_audio_buffer.commit"), _sent(observed) + + +@pytest.mark.parametrize("transcript", NON_STRING_TRANSCRIPTS) +def test_non_string_transcript_field_closes_a_voice_session_with_the_same_proxy_failure( + guardrail_proxy: OwnedProxy, transcript: JsonValue +) -> None: + with guardrail_proxy.gateway.scenario() as scenario: + session, observed = _voice_session( + guardrail_proxy, + scenario, + _voice_scenario(transcript), + "", + (_step(VOICE_UPDATE_FRAME, COMMIT, until="closed"),), + ) + assert session.types == ( + "session.created", + "session.updated", + "error", + TRANSCRIPT_COMPLETED, + "error", + "closed", + ), session + assert session.events[3]["transcript"] == transcript, session + _assert_closed_by_a_proxy_failure(session) + assert _sent_types(observed) == ("session.update", "session.update", "input_audio_buffer.commit"), _sent( + observed + ) + + async def _hold_until_closed(ws_base: str, query: str, key: str, opened: asyncio.Queue[str]) -> Session: async with websockets.connect( f"{ws_base}/v1/realtime?{query}", additional_headers={"Authorization": f"Bearer {key}"} @@ -1108,7 +1967,7 @@ async def _frames_until_closed(socket: ClientConnection) -> AsyncIterator[dict[s def _relays_the_upstream_close(session: Session) -> bool: - return f"upstream websocket closed with code {session.close_code}" in session.error_messages[0] + return _relayed_close(session) in session.error_messages[-1] async def _drain(opened: asyncio.Queue[str], count: int) -> tuple[str, ...]: @@ -1249,6 +2108,8 @@ def test_worker_kill_then_proxy_restart_keep_transcription_sessions_uninjected( config: Final = _write_config(tmp_path, ca_bundle) with gateway.scenario() as scenario: handle: Final = _scripted(scenario, _transcription_scenario(CLEAN_TRANSCRIPT)) + blocked_after_kill: Final = _scripted(scenario, _transcription_scenario(BLOCKED_TRANSCRIPT)) + blocked_after_restart: Final = _scripted(scenario, _transcription_scenario(BLOCKED_TRANSCRIPT)) key: Final = scenario.key() with owned_proxy_process(gateway, tmp_path, _overrides(ca_bundle), config=config, workers=WORKERS) as owned: owned_url: Final = _owned_url(owned) @@ -1259,6 +2120,20 @@ def test_worker_kill_then_proxy_restart_keep_transcription_sessions_uninjected( f"openai/{TRANSCRIBE_MODEL}", handle.scenario_id, ) + blocked_model_after_kill: Final = _named_deployment( + owned.gateway, + scenario, + f"realtime-guard-owned-blocked-{uuid.uuid4().hex[:8]}", + f"openai/{TRANSCRIBE_MODEL}", + blocked_after_kill.scenario_id, + ) + blocked_model_after_restart: Final = _named_deployment( + owned.gateway, + scenario, + f"realtime-guard-owned-blocked-{uuid.uuid4().hex[:8]}", + f"openai/{TRANSCRIBE_MODEL}", + blocked_after_restart.scenario_id, + ) root: Final = psutil.Process(owned.process.pid) outcome: Final = asyncio.run(_sessions_through_worker_kill(_ws_base(owned_url), model, key, root)) record_property( @@ -1284,6 +2159,15 @@ def test_worker_kill_then_proxy_restart_keep_transcription_sessions_uninjected( CLEAN_TRANSCRIPT, update=GA_TRANSCRIPTION_UPDATE, ) + blocked_on_the_survivor: Final = _transcribe_blocked( + _ws_base(owned_url), f"model={blocked_model_after_kill}&{TRANSCRIPTION_QUERY}", key + ) + _assert_transcription_blocked( + blocked_on_the_survivor, + _observed(gateway.upstream_url, blocked_after_kill.scenario_id), + BLOCKED_TRANSCRIPT, + update=GA_TRANSCRIPTION_UPDATE, + ) codes: Final = asyncio.run( _sessions_through_proxy_shutdown( _ws_base(owned_url), model, key, lambda: stop_root_process(owned.process) @@ -1299,3 +2183,106 @@ def test_worker_kill_then_proxy_restart_keep_transcription_sessions_uninjected( CLEAN_TRANSCRIPT, update=GA_TRANSCRIPTION_UPDATE, ) + blocked_on_the_restart: Final = _transcribe_blocked( + _ws_base(_owned_url(restarted)), f"model={blocked_model_after_restart}&{TRANSCRIPTION_QUERY}", key + ) + _assert_transcription_blocked( + blocked_on_the_restart, + _observed(gateway.upstream_url, blocked_after_restart.scenario_id), + BLOCKED_TRANSCRIPT, + update=GA_TRANSCRIPTION_UPDATE, + ) + + +async def _frames_reporting_verdicts( + socket: ClientConnection, session_id: str, verdicts: asyncio.Queue[str] +) -> AsyncIterator[dict[str, JsonValue]]: + async for frame in _frames_until_closed(socket): + if frame.get("type") == "error" and object_value(frame["error"]).get("type") == GUARDRAIL_VIOLATION[0]: + await verdicts.put(session_id) + yield frame + + +async def _hold_blocked_until_closed( + ws_base: str, query: str, key: str, opened: asyncio.Queue[str], verdicts: asyncio.Queue[str] +) -> Session: + async with websockets.connect( + f"{ws_base}/v1/realtime?{query}", additional_headers={"Authorization": f"Bearer {key}"} + ) as socket: + created: Final = JSON_OBJECT.validate_json(await socket.recv()) + assert created.get("type") == "session.created", created + session_id: Final = string_value(object_value(created["session"])["id"]) + await opened.put(session_id) + await socket.send(json.dumps(TRANSCRIPTION_UPDATE_FRAME)) + await socket.send(json.dumps(COMMIT)) + return Session(tuple([frame async for frame in _frames_reporting_verdicts(socket, session_id, verdicts)]), None) + + +@dataclass(frozen=True, slots=True) +class BlockedBurst: + sessions: tuple[Session, ...] + before_the_outage: tuple[dict[str, JsonValue], ...] + + +async def _blocked_burst_through_outage( + ws_base: str, + proxy_url: str, + upstream_url: str, + scenario_id: str, + model: str, + key: str, + stop_upstream: Callable[[], None], +) -> BlockedBurst: + opened: Final[asyncio.Queue[str]] = asyncio.Queue() + verdicts: Final[asyncio.Queue[str]] = asyncio.Queue() + query: Final = f"model={model}&{TRANSCRIPTION_QUERY}" + holders: Final = tuple( + asyncio.ensure_future(_hold_blocked_until_closed(ws_base, query, key, opened, verdicts)) for _ in range(BURST) + ) + opened_sessions: Final = await asyncio.wait_for(_drain(opened, BURST), 60) + assert len(opened_sessions) == BURST, opened_sessions + judged_sessions: Final = await asyncio.wait_for(_drain(verdicts, BURST), 60) + assert sorted(judged_sessions) == sorted(opened_sessions), judged_sessions + before_the_outage: Final = await asyncio.to_thread(_observed, upstream_url, scenario_id) + await asyncio.to_thread(stop_upstream) + async with httpx.AsyncClient(base_url=proxy_url, timeout=15, trust_env=False) as client: + liveliness: Final = await client.get("/health/liveliness") + assert liveliness.status_code == 200, liveliness.text + return BlockedBurst(tuple(await asyncio.wait_for(asyncio.gather(*holders), 90)), before_the_outage) + + +@pytest.mark.timeout(240) +def test_upstream_outage_closes_every_blocked_transcription_session_and_the_verdict_survives_the_restart( + guardrail_proxy: OwnedProxy, tmp_path: Path, record_property: RecordProperty +) -> None: + with guardrail_proxy.gateway.scenario() as scenario, owned_upstream(tmp_path) as slot: + scenario_id: Final = f"realtime-guard-blocked-outage-{uuid.uuid4().hex[:12]}" + register_scenario(scenario_id, _transcription_scenario(BLOCKED_TRANSCRIPT), control_url=slot.url) + key: Final = scenario.key() + model: Final = scenario.model(model=f"openai/{TRANSCRIBE_MODEL}", api_key=scenario_id, api_base=slot.url) + proxy_url: Final = _owned_url(guardrail_proxy) + burst: Final = asyncio.run( + _blocked_burst_through_outage(_ws_base(proxy_url), proxy_url, slot.url, scenario_id, model, key, slot.stop) + ) + record_property( + "close_codes_during_upstream_outage", sorted(session.close_code or 0 for session in burst.sessions) + ) + assert [session.types for session in burst.sessions] == [ + ("session.updated", TRANSCRIPT_COMPLETED, "error", "error", "closed") + ] * BURST, burst.sessions + assert [session.errors for session in burst.sessions] == [(GUARDRAIL_VIOLATION, RELAYED_CLOSE)] * BURST, ( + burst.sessions + ) + assert all(_relays_the_upstream_close(session) for session in burst.sessions), burst.sessions + assert len({session.close_code for session in burst.sessions}) == 1, burst.sessions + assert sorted(map(str, _sent_types(burst.before_the_outage))) == sorted( + ("session.update", "input_audio_buffer.commit") * BURST + ), _sent(burst.before_the_outage) + slot.start() + register_scenario(scenario_id, _transcription_scenario(BLOCKED_TRANSCRIPT), control_url=slot.url) + recovered: Final = _transcribe_blocked(_ws_base(proxy_url), f"model={model}&{TRANSCRIPTION_QUERY}", key) + _assert_transcription_blocked( + recovered, _observed(slot.url, scenario_id), BLOCKED_TRANSCRIPT, update=GA_TRANSCRIPTION_UPDATE + ) + rows: Final = _spend_rows(key, BURST + 1) + assert {str(row["call_type"]) for row in rows} == {"_arealtime"}, rows diff --git a/tests/integration/proxy_config.yaml b/tests/integration/proxy_config.yaml index 9e4a90fc608..6924b641f13 100644 --- a/tests/integration/proxy_config.yaml +++ b/tests/integration/proxy_config.yaml @@ -326,6 +326,7 @@ general_settings: master_key: os.environ/LITELLM_MASTER_KEY database_url: os.environ/DATABASE_URL store_model_in_db: true + disable_model_info_refresh: true disable_spend_logs: false proxy_batch_write_at: 1 proxy_batch_polling_interval: 1 diff --git a/tests/proxy_behavior/lens/test_lifecycle.py b/tests/proxy_behavior/lens/test_lifecycle.py index e3027474baa..5f30a11cf70 100644 --- a/tests/proxy_behavior/lens/test_lifecycle.py +++ b/tests/proxy_behavior/lens/test_lifecycle.py @@ -197,7 +197,10 @@ async def test_team_route_requires_a_worker_with_matching_model_access(lens_data ) try: wrong_team: Final = await endpoints.register_worker(endpoints.WorkerName(analysis_key_id=key_b), admin) - assert await endpoints.claim_candidate(lens, wrong_team.worker, datetime.now(timezone.utc)) is None + assert ( + await endpoints.claim_candidate(lens, wrong_team.worker, datetime.now(timezone.utc), endpoints.repository()) + is None + ) for operation in ( endpoints.create_lens(settings, admin), endpoints.run_lens(lens.id, RunRequest(), admin), @@ -212,7 +215,9 @@ async def test_team_route_requires_a_worker_with_matching_model_access(lens_data assert edited.settings.context == "Use sources" right_team: Final = await endpoints.register_worker(endpoints.WorkerName(analysis_key_id=key_a), admin) await endpoints.validate_workers(settings, lens.scope) - claim: Final = await endpoints.claim_candidate(lens, right_team.worker, datetime.now(timezone.utc)) + claim: Final = await endpoints.claim_candidate( + lens, right_team.worker, datetime.now(timezone.utc), endpoints.repository() + ) assert claim is not None and claim.job.worker_id == right_team.worker.id finally: await lens_database.db.execute_raw( @@ -257,7 +262,10 @@ async def test_scan_lifecycle_persists_results_and_revokes_worker(lens_database: assert lens.id in tuple(e.id for e in listing.lenses) assert worker.id in tuple(w.id for w in listing.workers) claims: Final = await asyncio.gather( - *(endpoints.claim_candidate(lens, worker, datetime.now(timezone.utc)) for _ in range(8)) + *( + endpoints.claim_candidate(lens, worker, datetime.now(timezone.utc), endpoints.repository()) + for _ in range(8) + ) ) winners: Final = tuple(claim for claim in claims if claim is not None) assert len(winners) == 1 @@ -265,7 +273,10 @@ async def test_scan_lifecycle_persists_results_and_revokes_worker(lens_database: assert claimed.job.worker_id == worker.id assert ( await endpoints.claim_candidate( - await endpoints.get_lens(lens.id, worker.scope), worker, datetime.now(timezone.utc) + await endpoints.get_lens(lens.id, worker.scope), + worker, + datetime.now(timezone.utc), + endpoints.repository(), ) is None ) @@ -420,7 +431,9 @@ async def test_failed_model_requests_release_lens_budget_reservations(lens_datab registration: Final = await endpoints.register_worker(endpoints.WorkerName(analysis_key_id=key_id), admin) worker: Final = registration.worker try: - claimed: Final = await endpoints.claim_candidate(lens, worker, datetime.now(timezone.utc)) + claimed: Final = await endpoints.claim_candidate( + lens, worker, datetime.now(timezone.utc), endpoints.repository() + ) assert claimed is not None for _ in range(3): with pytest.raises(HTTPException) as failed: diff --git a/tests/unit/integrations/test_anthropic_cache_control_hook.py b/tests/unit/integrations/test_anthropic_cache_control_hook.py index 23fa2e36bce..8392cef8a7b 100644 --- a/tests/unit/integrations/test_anthropic_cache_control_hook.py +++ b/tests/unit/integrations/test_anthropic_cache_control_hook.py @@ -3088,7 +3088,7 @@ class TestOpenAIPromptCacheBreakpoint: messages, system = self._inject([{"role": "user", "content": "hi"}], "sys", kwargs) assert system == [{"type": "text", "text": "sys", "prompt_cache_breakpoint": self.EXPLICIT}] assert messages == [{"role": "user", "content": [{"type": "text", "text": "hi"}]}] - assert kwargs == {"prompt_cache_options": self.EXPLICIT} + assert kwargs == {"prompt_cache_options": {"mode": "implicit"}} assert not _contains_key(system, "cache_control") def test_v1_messages_list_system_marks_last_block_only(self): @@ -3099,7 +3099,7 @@ class TestOpenAIPromptCacheBreakpoint: {"type": "text", "text": "a"}, {"type": "text", "text": "b", "prompt_cache_breakpoint": self.EXPLICIT}, ] - assert kwargs["prompt_cache_options"] == self.EXPLICIT + assert kwargs["prompt_cache_options"] == {"mode": "implicit"} def test_v1_messages_targets_by_role(self): messages = [ @@ -3115,7 +3115,7 @@ class TestOpenAIPromptCacheBreakpoint: ] assert result[1] == messages[1] assert result[2]["content"] == [{"type": "text", "text": "last", "prompt_cache_breakpoint": self.EXPLICIT}] - assert kwargs["prompt_cache_options"] == self.EXPLICIT + assert kwargs["prompt_cache_options"] == {"mode": "implicit"} def test_v1_messages_targets_by_index(self): messages = [ @@ -3142,8 +3142,9 @@ class TestOpenAIPromptCacheBreakpoint: assert not _contains_key(system, "cache_control") assert not _contains_key(messages, "cache_control") - def test_v1_messages_keeps_caller_prompt_cache_options(self): - caller_options = {"mode": "explicit", "ttl": "30m"} + @pytest.mark.parametrize("mode", ["explicit", "implicit"]) + def test_v1_messages_keeps_caller_prompt_cache_options(self, mode): + caller_options = {"mode": mode, "ttl": "30m"} kwargs = { "cache_control_injection_points": copy.deepcopy(self.SYSTEM_POINT), "prompt_cache_options": dict(caller_options), @@ -3152,6 +3153,15 @@ class TestOpenAIPromptCacheBreakpoint: assert system[0]["prompt_cache_breakpoint"] == self.EXPLICIT assert kwargs["prompt_cache_options"] == caller_options + def test_v1_messages_and_chat_paths_default_to_the_same_implicit_mode(self): + messages_kwargs = {"cache_control_injection_points": copy.deepcopy(self.SYSTEM_POINT)} + self._inject([{"role": "user", "content": "hi"}], "sys", messages_kwargs) + _, _, chat_params = self._chat( + [{"role": "system", "content": "sys"}, {"role": "user", "content": "hi"}], + {"cache_control_injection_points": copy.deepcopy(self.SYSTEM_POINT)}, + ) + assert messages_kwargs["prompt_cache_options"] == chat_params["prompt_cache_options"] == {"mode": "implicit"} + def test_v1_messages_no_prompt_cache_options_when_nothing_injected(self): kwargs = {"cache_control_injection_points": copy.deepcopy(self.SYSTEM_POINT)} messages, system = self._inject([{"role": "user", "content": "hi"}], None, kwargs) @@ -3185,7 +3195,7 @@ class TestOpenAIPromptCacheBreakpoint: result, system = self._inject(messages, "sys", kwargs) assert result == messages assert system == [{"type": "text", "text": "sys", "prompt_cache_breakpoint": self.EXPLICIT}] - assert kwargs == {"prompt_cache_options": self.EXPLICIT} + assert kwargs == {"prompt_cache_options": {"mode": "implicit"}} def test_v1_messages_tail_point_applies_beside_client_system_breakpoint(self): system = [{"type": "text", "text": "sys", "prompt_cache_breakpoint": self.EXPLICIT}] @@ -3196,7 +3206,7 @@ class TestOpenAIPromptCacheBreakpoint: {"role": "user", "content": [{"type": "text", "text": "hi", "prompt_cache_breakpoint": self.EXPLICIT}]} ] assert result_system == system - assert kwargs == {"prompt_cache_options": self.EXPLICIT} + assert kwargs == {"prompt_cache_options": {"mode": "implicit"}} def test_chat_system_string_wrapped_with_block_breakpoint(self): params = {"cache_control_injection_points": copy.deepcopy(self.SYSTEM_POINT)} @@ -3375,7 +3385,7 @@ class TestOpenAIPromptCacheBreakpointPlacementRules: {"type": "tool_result", "tool_use_id": "t1", "content": "sunny"}, {"type": "text", "text": "thanks", "prompt_cache_breakpoint": self.EXPLICIT}, ] - assert kwargs["prompt_cache_options"] == self.EXPLICIT + assert kwargs["prompt_cache_options"] == {"mode": "implicit"} def test_marker_walks_back_to_last_eligible_block(self): messages = [ @@ -3656,12 +3666,12 @@ class TestMessagesPathApiBaseGate: def test_regional_openai_api_base_uses_openai_dialect(self): block, kwargs = self._inject("gpt-5.6", api_base="https://eu.api.openai.com/v1") assert block == self.BREAKPOINT_BLOCK - assert kwargs["prompt_cache_options"] == self.EXPLICIT + assert kwargs["prompt_cache_options"] == {"mode": "implicit"} def test_default_api_base_uses_openai_dialect(self): block, kwargs = self._inject("openai/gpt-5.6") assert block == self.BREAKPOINT_BLOCK - assert kwargs["prompt_cache_options"] == self.EXPLICIT + assert kwargs["prompt_cache_options"] == {"mode": "implicit"} class TestToolConfigSlotInOpenAIDialect: @@ -3752,7 +3762,7 @@ class TestPromptCacheBreakpointCapability: [{"role": "user", "content": "hi"}], "sys", kwargs, model="gpt-5.6", custom_llm_provider="openai" ) assert system == [{"type": "text", "text": "sys", "prompt_cache_breakpoint": {"mode": "explicit"}}] - assert kwargs == {"prompt_cache_options": {"mode": "explicit"}} + assert kwargs == {"prompt_cache_options": {"mode": "implicit"}} @pytest.mark.parametrize("model,expected", [("gpt-5.6-2026-01-01", True), ("gpt-5.5-preview-unlisted", False)]) def test_unlisted_model_falls_back_to_the_version_rule(self, model, expected): diff --git a/tests/unit/litellm_core_utils/test_logging_utils.py b/tests/unit/litellm_core_utils/test_logging_utils.py index 7b8db097d47..edf0dc7960b 100644 --- a/tests/unit/litellm_core_utils/test_logging_utils.py +++ b/tests/unit/litellm_core_utils/test_logging_utils.py @@ -95,6 +95,18 @@ class TestTruncateBase64InString: result = _truncate_base64_in_string(text) assert result.count("base64_data truncated") == 2 + @pytest.mark.timeout(10) + @pytest.mark.parametrize( + "text", + [ + 'data: {"choices": [{"delta": {"content": "hi"}}]}\n\n' * 50_000, + "data:" * 200_000, + ], + ids=["sse_lines", "whitespace_free_prefixes"], + ) + def test_repeated_data_prefixes_without_data_uris_are_scanned_in_linear_time(self, text: str): + assert _truncate_base64_in_string(text) == text + def test_no_data_uri(self): text = "hello world, no base64 here" assert _truncate_base64_in_string(text) == text diff --git a/tests/unit/litellm_core_utils/test_realtime_streaming.py b/tests/unit/litellm_core_utils/test_realtime_streaming.py index 190a7adb6e5..9f9f0d340d4 100644 --- a/tests/unit/litellm_core_utils/test_realtime_streaming.py +++ b/tests/unit/litellm_core_utils/test_realtime_streaming.py @@ -1,11 +1,12 @@ import asyncio import json -from collections.abc import Coroutine +from collections.abc import Coroutine, Mapping from dataclasses import dataclass -from typing import Final +from typing import Final, Literal from unittest.mock import AsyncMock, MagicMock, patch import pytest +from typing_extensions import ReadOnly, TypedDict from websockets.exceptions import ConnectionClosed from websockets.frames import Close @@ -18,6 +19,7 @@ from litellm.litellm_core_utils.realtime_streaming import ( ) from litellm.llms.xai.realtime.transformation import XAIRealtimeNormalizer from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.utils import GenericGuardrailAPIInputs def _make_transcript_event(text: str, item_id: str = "item_x") -> bytes: @@ -3708,3 +3710,112 @@ async def test_transcription_guardrail_still_disables_auto_response_on_realtime_ forwarded: Final = json.loads(backend_ws.send.await_args.args[0]) assert forwarded["session"]["audio"]["input"]["turn_detection"]["create_response"] is False, forwarded + + +class _ViolationSettings(TypedDict, total=False): + on_violation: ReadOnly[str] + end_session_after_n_fails: ReadOnly[int] + + +def _passthrough_transcription_config() -> MagicMock: + def transform_response( + message: str | bytes, + model: str, + logging_obj: object, + realtime_response_transform_input: object, + ) -> dict[str, object]: + return { + "response": json.loads(message), + "current_output_item_id": None, + "current_response_id": None, + "current_delta_chunks": None, + "current_conversation_id": None, + "current_item_chunks": None, + "current_delta_type": None, + "session_configuration_request": None, + } + + def transform_request(message: str, model: str, session_configuration_request: str | None = None) -> list[str]: + return [message] + + provider_config: Final = MagicMock() + provider_config.requires_session_configuration.return_value = False + provider_config.transform_realtime_response.side_effect = transform_response + provider_config.transform_realtime_request.side_effect = transform_request + return provider_config + + +@pytest.mark.asyncio +@pytest.mark.parametrize("uses_provider_config", [False, True]) +@pytest.mark.parametrize( + ("violation_settings", "expect_session_closed"), + [ + ({}, False), + ({"on_violation": "end_session"}, True), + ({"end_session_after_n_fails": 1}, True), + ], +) +async def test_transcription_session_guardrail_block_only_reports_violation( + monkeypatch: pytest.MonkeyPatch, + uses_provider_config: bool, + violation_settings: _ViolationSettings, + expect_session_closed: bool, +) -> None: + class BlockingGuardrail(CustomGuardrail): + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: Mapping[str, object], + input_type: Literal["request", "response"], + logging_obj: object | None = None, + ) -> GenericGuardrailAPIInputs: + if any("blocked" in text for text in inputs.get("texts", [])): + raise ValueError("blocked transcript") + return inputs + + monkeypatch.setattr( + litellm, + "callbacks", + [ + BlockingGuardrail( + guardrail_name="transcription-blocker", + event_hook=GuardrailEventHooks.realtime_input_transcription, + default_on=True, + **violation_settings, + ) + ], + ) + completed_type: Final = "conversation.item.input_audio_transcription.completed" + client_ws: Final = MagicMock() + client_ws.send_text = AsyncMock() + backend_ws: Final = MagicMock() + blocked_event: Final = _make_transcript_event("a blocked transcript", item_id="item_1") + follow_up_events: Final = ( + () if expect_session_closed else (_make_transcript_event("a clean follow-up", item_id="item_2"),) + ) + backend_ws.recv = AsyncMock(side_effect=[blocked_event, *follow_up_events, ConnectionClosed(None, None)]) + backend_ws.send = AsyncMock() + backend_ws.close = AsyncMock() + streaming: Final = RealTimeStreaming( + client_ws, + backend_ws, + MagicMock(), + provider_config=_passthrough_transcription_config() if uses_provider_config else None, + model="gpt-4o-transcribe", + force_transcription_model="gpt-4o-transcribe", + ) + + await streaming.backend_to_client_send_messages() + + sent_to_client: Final = [json.loads(call.args[0]) for call in client_ws.send_text.await_args_list] + expected_follow_up: Final = () if expect_session_closed else ((completed_type, "a clean follow-up"),) + assert [(event["type"], event.get("transcript")) for event in sent_to_client] == [ + (completed_type, "a blocked transcript"), + ("error", None), + *expected_follow_up, + ], sent_to_client + assert sent_to_client[1]["error"]["type"] == "guardrail_violation", sent_to_client + assert streaming._violation_count == 1 + sent_to_backend: Final = [call.args[0] for call in backend_ws.send.await_args_list] + assert sent_to_backend == [], sent_to_backend + assert backend_ws.close.await_count == (1 if expect_session_closed else 0) diff --git a/tests/unit/proxy/lens/test_endpoints.py b/tests/unit/proxy/lens/test_endpoints.py index 2238da7e00f..74eb649aad0 100644 --- a/tests/unit/proxy/lens/test_endpoints.py +++ b/tests/unit/proxy/lens/test_endpoints.py @@ -1,3 +1,4 @@ +from collections.abc import Callable from datetime import datetime, timedelta, timezone from types import SimpleNamespace from typing import Final @@ -11,6 +12,7 @@ import litellm from litellm import Router from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.lens.endpoints import ( + claim_due, list_agents, read_reviews, result, @@ -23,6 +25,9 @@ from litellm.proxy.lens.endpoints import ( watching, worker_supports_model, ) +from litellm.proxy.lens.endpoints import ( + sample as worker_sample, +) from litellm.proxy.lens.models import ( ActivitySelection, Coverage, @@ -36,11 +41,13 @@ from litellm.proxy.lens.models import ( Scope, TraceFindingsRequest, TraceIdentity, + Worker, ) -from litellm.proxy.lens.repository import Row +from litellm.proxy.lens.repository import DueLens, Row from litellm.proxy.lens.state import claim_job, queue_job, replace_job from litellm.rust_bridge.trace.storage import ClickHouseStorage from litellm.tracing.remote import RemoteTraceStore +from litellm.rust_bridge.trace.generated.models import ExecutionRow, LensSampleParams from tests.unit.proxy.lens.test_agent_workspace import execution from tests.unit.proxy.lens.test_state import NOW, lens, worker @@ -115,6 +122,74 @@ async def test_result_cannot_commit_after_losing_ownership_during_evidence_valid assert db.completed == () +@pytest.mark.asyncio +async def test_worker_sample_retries_oversized_pages_and_keeps_all_executions( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.proxy import proxy_server + + claimed: Final = claim_job(queue_job(lens(), NOW, "job"), worker(), NOW) + active: Final = claimed.jobs[0].model_copy(update={"lease_until": datetime.max.replace(tzinfo=timezone.utc)}) + db: Final = ResultDatabase(replace_job(claimed, active)) + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db)) + rows: Final = tuple( + ExecutionRow( + source="traces", + trace_id=trace_id, + team_id="team", + name=trace_id, + start_time="", + span_count=1, + root_seen=1, + eligible=3, + selected=3, + selection_key=trace_id, + ) + for trace_id in ("trace-1", "trace-2", "trace-3") + ) + + class SampleStorage: + def __init__(self) -> None: + self.limits: tuple[int, ...] = () + + async def lens_sample(self, parameters: LensSampleParams) -> tuple[ExecutionRow, ...]: + self.limits = (*self.limits, parameters.limit) + if parameters.limit > 2_500: + raise RuntimeError("ClickHouse query exceeded the response size limit") + return rows + + storage: Final = SampleStorage() + selected: Final = await worker_sample("lens", "job", worker(), storage) + assert storage.limits == (10_000, 5_000, 2_500) + assert tuple(execution.trace_id for execution in selected.executions) == ("trace-1", "trace-2", "trace-3") + assert selected.selected == 3 + + +@pytest.mark.asyncio +async def test_worker_sample_propagates_response_too_large_at_minimum_page_size( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.proxy import proxy_server + + claimed: Final = claim_job(queue_job(lens(), NOW, "job"), worker(), NOW) + active: Final = claimed.jobs[0].model_copy(update={"lease_until": datetime.max.replace(tzinfo=timezone.utc)}) + db: Final = ResultDatabase(replace_job(claimed, active)) + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db)) + + class SampleStorage: + def __init__(self) -> None: + self.limits: tuple[int, ...] = () + + async def lens_sample(self, parameters: LensSampleParams) -> tuple[ExecutionRow, ...]: + self.limits = (*self.limits, parameters.limit) + raise RuntimeError("ClickHouse query exceeded the response size limit") + + storage: Final = SampleStorage() + with pytest.raises(RuntimeError, match="response size limit"): + await worker_sample("lens", "job", worker(), storage) + assert storage.limits == (10_000, 5_000, 2_500, 1_250, 625, 312, 156, 100) + + @pytest.mark.asyncio @pytest.mark.parametrize( "selected,check_id,quoted", @@ -675,3 +750,53 @@ async def test_unknown_gateway_release_refuses_registration_and_claims(monkeypat await claim(worker(), protocol_version=PROTOCOL_VERSION, worker_release="") assert claim_error.value.status_code == 503 assert claim_error.value.detail == registration_error.value.detail + + +@pytest.mark.asyncio +async def test_claim_due_pages_through_more_than_a_thousand_full_pages() -> None: + candidate_lens: Final = lens() + + def candidate_page(page_number: int, size: int) -> tuple[DueLens, ...]: + return tuple( + DueLens( + lens=candidate_lens.model_copy(update={"id": f"lens-{page_number * 20 + offset:05}"}), + due_at=NOW, + ) + for offset in range(size) + ) + + full_pages: Final = tuple(candidate_page(page_number, 20) for page_number in range(1_200)) + pages: Final = (*full_pages, candidate_page(1_200, 1)) + assigned_worker: Final = worker() + + class PagingRepository: + def __init__(self) -> None: + self.after_calls: tuple[DueLens | None, ...] = () + + async def due( + self, scope: Scope, now: datetime, limit: int, after: DueLens | None = None + ) -> tuple[DueLens, ...]: + assert scope == assigned_worker.scope + assert now == NOW + assert limit == 20 + self.after_calls = (*self.after_calls, after) + return pages[len(self.after_calls) - 1] + + async def sync_due(self, lens: Lens) -> None: + return None + + async def update( + self, lens_id: str, transform: Callable[[Lens], Lens], attempts: int, *, changed_only: bool + ) -> Lens | None: + raise AssertionError("Unsupported models must not update candidates") + + async def reject_model(_worker: Worker, _settings: LensSettings) -> bool: + return False + + repository: Final = PagingRepository() + claim: Final = await claim_due(assigned_worker, NOW, repository, reject_model) + expected_after: Final = (None, *(page[-1] for page in pages[:-1])) + + assert claim is None + assert len(repository.after_calls) == 1_201 + assert repository.after_calls == expected_after diff --git a/tests/unit/proxy/lens/test_state.py b/tests/unit/proxy/lens/test_state.py index 9db69557e28..80eb638a610 100644 --- a/tests/unit/proxy/lens/test_state.py +++ b/tests/unit/proxy/lens/test_state.py @@ -37,6 +37,7 @@ from litellm.proxy.lens.state import ( cancel_job, claim_job, current_job, + due_at, end_job, merge_finding, next_scan_start, @@ -89,6 +90,22 @@ def worker(team: str = "alpha", identity: str = "worker") -> Worker: return Worker(id=identity, name=identity, scope=Scope(team_id=team), last_seen=NOW) +def lens_with_job( + status: Literal["queued", "running", "completed"], + lease_until: datetime | None = None, + *, + enabled: bool = True, + trigger: Literal["schedule", "manual"] = "schedule", +) -> Lens: + original: Final = lens() + configured: Final = original.model_copy( + update={"settings": original.settings.model_copy(update={"enabled": enabled})} + ) + queued: Final = queue_job(configured, NOW, "job", trigger=trigger) + job: Final = queued.jobs[0].model_copy(update={"status": status, "lease_until": lease_until}) + return queued.model_copy(update={"jobs": (job,)}) + + def finding(execution: str) -> FindingDraft: return FindingDraft( title="Repeated failed searches", @@ -98,6 +115,29 @@ def finding(execution: str) -> FindingDraft: ) +@pytest.mark.parametrize( + ("candidate", "expected"), + ( + pytest.param(lens(), NOW, id="idle-enabled"), + pytest.param( + lens().model_copy(update={"settings": lens().settings.model_copy(update={"enabled": False})}), + None, + id="idle-disabled", + ), + pytest.param(lens_with_job("queued", enabled=False, trigger="manual"), NOW, id="queued-manual-while-disabled"), + pytest.param( + lens_with_job("running", NOW + timedelta(minutes=5)), + NOW + timedelta(minutes=5), + id="running-with-lease", + ), + pytest.param(lens_with_job("running"), NOW, id="running-without-lease"), + pytest.param(lens_with_job("completed"), NOW, id="completed-only"), + ), +) +def test_due_at_matches_the_current_scheduling_state(candidate: Lens, expected: datetime | None) -> None: + assert due_at(candidate) == expected + + @pytest.mark.parametrize( ("viewer", "target", "allowed"), ( diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index c899083ca80..1f420d0e8dc 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -777,6 +777,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "cache_creation_input_token_cost": {"type": "number"}, "cache_creation_input_token_cost_above_1hr": {"type": "number"}, "cache_creation_input_token_cost_above_32k_tokens": {"type": "number"}, + "cache_creation_input_token_cost_above_100k_tokens": {"type": "number"}, "cache_creation_input_token_cost_above_128k_tokens": {"type": "number"}, "cache_creation_input_token_cost_above_200k_tokens": {"type": "number"}, "cache_creation_input_token_cost_above_256k_tokens": {"type": "number"}, @@ -792,6 +793,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "cache_creation_input_token_cost_ultrafast": {"type": "number"}, "cache_read_input_token_cost": {"type": "number"}, "cache_read_input_token_cost_above_32k_tokens": {"type": "number"}, + "cache_read_input_token_cost_above_100k_tokens": {"type": "number"}, "cache_read_input_token_cost_above_128k_tokens": {"type": "number"}, "cache_read_input_token_cost_above_200k_tokens": {"type": "number"}, "cache_read_input_token_cost_above_200k_tokens_batches": {"type": "number"}, @@ -803,6 +805,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "input_cost_per_token_above_272k_tokens_ultrafast": {"type": "number"}, "cache_read_input_token_cost_batches": {"type": "number"}, "cache_read_input_token_cost_above_272k_tokens_batches": {"type": "number"}, + "cache_creation_input_token_cost_above_1hr_above_100k_tokens": {"type": "number"}, "cache_creation_input_token_cost_above_1hr_above_200k_tokens": {"type": "number"}, "cache_read_input_audio_token_cost": {"type": "number"}, "cache_read_input_image_token_cost": {"type": "number"}, @@ -819,6 +822,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "input_cost_per_image_above_128k_tokens": {"type": "number"}, "input_cost_per_video_token": {"type": "number"}, "input_cost_per_token_above_32k_tokens": {"type": "number"}, + "input_cost_per_token_above_100k_tokens": {"type": "number"}, "input_cost_per_token_above_200k_tokens": {"type": "number"}, "input_cost_per_token_above_200k_tokens_batches": {"type": "number"}, "input_cost_per_token_above_256k_tokens": {"type": "number"}, @@ -928,6 +932,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "output_cost_per_second_4k": {"type": "number"}, "output_cost_per_token": {"type": "number"}, "output_cost_per_token_above_32k_tokens": {"type": "number"}, + "output_cost_per_token_above_100k_tokens": {"type": "number"}, "output_cost_per_token_above_128k_tokens": {"type": "number"}, "output_cost_per_token_above_200k_tokens": {"type": "number"}, "output_cost_per_token_above_200k_tokens_batches": {"type": "number"},