diff --git a/.github/codeql/codeql-config.yml b/.github/codeql/codeql-config.yml index 6e15c1069a3..61982be90a9 100644 --- a/.github/codeql/codeql-config.yml +++ b/.github/codeql/codeql-config.yml @@ -1,8 +1,5 @@ name: "LiteLLM CodeQL config" -queries: - - uses: security-and-quality - # Known OOM queries on large Python codebases: # CodeQL builds a full data flow graph in memory. These two queries trace # sensitive data through every log call / regex pattern, causing combinatorial @@ -14,17 +11,6 @@ query-filters: id: py/clear-text-logging-sensitive-data # CWE-312 - exclude: id: py/polynomial-redos # CWE-730 - # Import resolution confuses stdlib types with management_endpoints/types.py. - # The generic cycle query also reports intentional deferred imports. - - exclude: - id: py/cyclic-import - - exclude: - id: py/unsafe-cyclic-import - # Known false positives on live settings and Protocol placeholders. - - exclude: - id: py/unused-global-variable - - exclude: - id: py/ineffectual-statement paths-ignore: - tests diff --git a/.github/workflows/codeql.yml b/.github/workflows/codeql.yml index 9a85ced57f6..7fd767abf59 100644 --- a/.github/workflows/codeql.yml +++ b/.github/workflows/codeql.yml @@ -1,8 +1,6 @@ name: "CodeQL" on: - push: - branches: [main] pull_request: branches: [main] schedule: @@ -43,14 +41,15 @@ jobs: persist-credentials: false - name: Initialize CodeQL - uses: github/codeql-action/init@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3.34.1 + uses: github/codeql-action/init@2892aa5e19bbd11bc0cff5427e3b750a04d9e3c2 # v4.38.2 with: languages: ${{ matrix.language }} build-mode: ${{ matrix.build-mode }} config-file: ./.github/codeql/codeql-config.yml + queries: ${{ github.event_name == 'pull_request' && '+security-extended' || '' }} - name: Perform CodeQL Analysis - uses: github/codeql-action/analyze@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3.34.1 + uses: github/codeql-action/analyze@2892aa5e19bbd11bc0cff5427e3b750a04d9e3c2 # v4.38.2 with: category: "/language:${{ matrix.language }}" output: sarif-results @@ -83,7 +82,7 @@ jobs: output: sarif-results/python.sarif - name: Upload SARIF - uses: github/codeql-action/upload-sarif@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3.34.1 + uses: github/codeql-action/upload-sarif@2892aa5e19bbd11bc0cff5427e3b750a04d9e3c2 # v4.38.2 with: sarif_file: sarif-results category: "/language:${{ matrix.language }}" 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 83dbb3c5aa0..42ed777e6bc 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/migrations/20261008000200_lens_trace_signals/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008000200_lens_trace_signals/migration.sql new file mode 100644 index 00000000000..67043eccb5f --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008000200_lens_trace_signals/migration.sql @@ -0,0 +1,16 @@ +CREATE TABLE IF NOT EXISTS "LiteLLM_LensSignalConfig" ( + "id" TEXT NOT NULL, + "data" JSONB NOT NULL, + CONSTRAINT "LiteLLM_LensSignalConfig_pkey" PRIMARY KEY ("id") +); + +CREATE TABLE IF NOT EXISTS "LiteLLM_LensTraceSignal" ( + "trace_id" TEXT NOT NULL, + "trace_ref" TEXT NOT NULL DEFAULT '', + "config_key" TEXT NOT NULL, + "span_count" INTEGER NOT NULL, + "claimed_until" TIMESTAMP(3), + "classified_at" TIMESTAMP(3), + "data" JSONB NOT NULL, + CONSTRAINT "LiteLLM_LensTraceSignal_pkey" PRIMARY KEY ("trace_id", "trace_ref") +); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 514df905866..3b83c5b09cc 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 { @@ -1974,3 +1977,20 @@ model LiteLLM_LensDataset { @@id([id, revision]) } + +model LiteLLM_LensSignalConfig { + id String @id + data Json +} + +model LiteLLM_LensTraceSignal { + trace_id String + trace_ref String @default("") + config_key String + span_count Int + claimed_until DateTime? + classified_at DateTime? + data Json + + @@id([trace_id, trace_ref]) +} diff --git a/litellm-rust/crates/cache/src/semantic.rs b/litellm-rust/crates/cache/src/semantic.rs index c88706a213b..6200645e555 100644 --- a/litellm-rust/crates/cache/src/semantic.rs +++ b/litellm-rust/crates/cache/src/semantic.rs @@ -83,17 +83,16 @@ impl Embedder for PreparedEmbedding { } } -/// `get_str_from_messages`: every message's text content followed by its search results. +/// `get_semantic_cache_prompt_from_messages`: every message's text content, including the text of +/// Messages API `tool_result` blocks, followed by its search results. pub fn str_from_messages(messages: &[Value]) -> String { let mut text = String::new(); for message in messages.iter().filter_map(Value::as_object) { match message.get("content") { Some(Value::String(content)) => text.push_str(content), - Some(Value::Array(parts)) => { - for part in parts { - if let Some(part_text) = part.get("text").and_then(Value::as_str) { - text.push_str(part_text); - } + Some(Value::Array(blocks)) => { + for block in blocks { + push_block_text(&mut text, block); } } _ => {} @@ -103,6 +102,28 @@ pub fn str_from_messages(messages: &[Value]) -> String { text } +fn push_block_text(text: &mut String, block: &Value) { + if block.get("type").and_then(Value::as_str) != Some("tool_result") { + push_text_field(text, block); + return; + } + match block.get("content") { + Some(Value::String(result)) => text.push_str(result), + Some(Value::Array(blocks)) => { + for inner in blocks { + push_text_field(text, inner); + } + } + _ => {} + } +} + +fn push_text_field(text: &mut String, block: &Value) { + if let Some(block_text) = block.get("text").and_then(Value::as_str) { + text.push_str(block_text); + } +} + /// The messages prompt Qdrant embeds: `None` when the request carries no messages. pub fn prompt_from_messages(context: &SemanticCacheContext) -> Option { let messages = context.messages.as_ref()?.as_array()?; @@ -163,6 +184,10 @@ fn collect_input_text(value: &Value, parts: &mut Vec) { collect_input_text(content, parts); return; } + if let Some(output) = map.get("output").filter(|output| output.is_array()) { + collect_input_text(output, parts); + return; + } for key in ["text", "output", "input_text", "output_text"] { if let Some(Value::String(text)) = map.get(key) && push_trimmed(text, parts) diff --git a/litellm-rust/crates/cache/tests/semantic.rs b/litellm-rust/crates/cache/tests/semantic.rs index 97a552a8010..76af1863e1f 100644 --- a/litellm-rust/crates/cache/tests/semantic.rs +++ b/litellm-rust/crates/cache/tests/semantic.rs @@ -30,6 +30,31 @@ fn context(messages: Option, input: Option) -> SemanticCacheContex ]}]), "What is this?", )] +#[case::tool_result_string( + json!([ + {"role": "user", "content": "list the files"}, + {"role": "assistant", "content": [ + {"type": "tool_use", "id": "toolu_1", "name": "Bash", "input": {"command": "ls"}}, + ]}, + {"role": "user", "content": [ + {"type": "tool_result", "tool_use_id": "toolu_1", "content": "calc.py test_calc.py"}, + ]}, + ]), + "list the filescalc.py test_calc.py", +)] +#[case::tool_result_blocks( + json!([{"role": "user", "content": [ + {"type": "tool_result", "tool_use_id": "toolu_1", "content": [ + {"type": "text", "text": "x = 1"}, + {"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": ""}}, + ]}, + ]}]), + "x = 1", +)] +#[case::tool_result_without_content( + json!([{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_1"}]}]), + "", +)] #[case::missing_null_and_empty_content( json!([{"role": "assistant"}, {"role": "assistant", "content": null}, {"role": "user", "content": ""}]), "", @@ -166,6 +191,15 @@ fn prompt_from_messages_reads_messages_only( ])), Some("model dump prompt\ndict prompt\ninline prompt"), )] +#[case::function_call_output_blocks( + None, + Some(json!([ + {"role": "user", "content": "update the config"}, + {"type": "function_call", "call_id": "c1", "name": "write_file", "arguments": "{\"path\": \"a\"}"}, + {"type": "function_call_output", "call_id": "c1", "output": [{"type": "input_text", "text": "wrote a"}]}, + ])), + Some("update the config\nwrote a"), +)] #[case::object_content( None, Some(json!({"content": [{"text": "object content prompt"}]})), 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_content.sql b/litellm-rust/crates/traces-clickhouse/query/lens_content.sql index 99fb56a5f48..52355a11061 100644 --- a/litellm-rust/crates/traces-clickhouse/query/lens_content.sql +++ b/litellm-rust/crates/traces-clickhouse/query/lens_content.sql @@ -17,6 +17,7 @@ SELECT * FROM ( FROM otel_traces WHERE {source:String}='traces' AND ({all_teams:UInt8}=1 OR TeamId={team:String}) AND ({key_hash:String}='' OR ApiKeyHash={key_hash:String}) + AND Timestamp >= parseDateTime64BestEffortOrZero({start_time:String}, 9) - INTERVAL 7 DAY AND ({trace_ref:String}='' OR hex(SHA256(concat(TeamId, char(0), ApiKeyHash, char(0), TraceId)))={trace_ref:String}) AND TraceId={id:String} AND TeamId={record_team:String} AND SpanId > {cursor:String} ORDER BY SpanId LIMIT 1 BY SpanId LIMIT 40 @@ -35,5 +36,6 @@ SELECT * FROM ( FROM spend_logs FINAL WHERE {source:String}='requests' 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 >= parseDateTime64BestEffortOrZero({start_time:String}, 3) - INTERVAL 7 DAY AND request_id={id:String} AND team_id={record_team:String} LIMIT 1 ) diff --git a/litellm-rust/crates/traces-clickhouse/query/lens_evidence.sql b/litellm-rust/crates/traces-clickhouse/query/lens_evidence.sql index a0d600cdfde..b53617364cf 100644 --- a/litellm-rust/crates/traces-clickhouse/query/lens_evidence.sql +++ b/litellm-rust/crates/traces-clickhouse/query/lens_evidence.sql @@ -2,6 +2,7 @@ SELECT sum(matches) AS count FROM ( SELECT count() AS matches FROM otel_traces WHERE {source:String}='traces' AND ({all_teams:UInt8}=1 OR TeamId={team:String}) AND ({key_hash:String}='' OR ApiKeyHash={key_hash:String}) + AND Timestamp >= parseDateTime64BestEffortOrZero({start_time:String}, 9) - INTERVAL 7 DAY AND ({trace_ref:String}='' OR hex(SHA256(concat(TeamId, char(0), ApiKeyHash, char(0), TraceId)))={trace_ref:String}) AND TraceId={id:String} AND TeamId={record_team:String} AND SpanId={span:String} AND position(concat('Input: ',Input,'\nOutput: ',Output,'\nStatus: ',StatusCode,' ',StatusMessage),{quote:String})>0 @@ -9,6 +10,7 @@ SELECT sum(matches) AS count FROM ( SELECT count() AS matches FROM spend_logs FINAL WHERE {source:String}='requests' 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 >= parseDateTime64BestEffortOrZero({start_time:String}, 3) - INTERVAL 7 DAY AND request_id={id:String} AND team_id={record_team:String} AND request_id={span:String} AND position(concat('Input: ',messages,'\nOutput: ',response,'\nError: ',error_str),{quote:String})>0 ) 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..77474c7143d 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"); } @@ -202,6 +209,7 @@ pub struct LensContentParams { pub source: ContentSource, pub id: String, pub record_team: String, + pub start_time: String, pub trace_ref: String, pub cursor: String, #[serde(deserialize_with = "super::number::deserialize")] @@ -245,6 +253,7 @@ pub struct LensEvidenceParams { pub source: ContentSource, pub id: String, pub record_team: String, + pub start_time: String, pub trace_ref: String, pub span: String, pub quote: String, diff --git a/litellm-rust/crates/traces-clickhouse/src/query/number.rs b/litellm-rust/crates/traces-clickhouse/src/query/number.rs index 9283903fee1..7a07845d0cd 100644 --- a/litellm-rust/crates/traces-clickhouse/src/query/number.rs +++ b/litellm-rust/crates/traces-clickhouse/src/query/number.rs @@ -99,7 +99,7 @@ mod tests { fn content_rejects_unsupported_sources(#[case] source: &str, #[case] valid: bool) { let parameters = serde_json::json!({ "all_teams": 0, "team": "team", "key_hash": "", "source": source, "id": "id", - "record_team": "team", "trace_ref": "", "cursor": "", "offset": 0 + "record_team": "team", "start_time": "", "trace_ref": "", "cursor": "", "offset": 0 }); assert_eq!( serde_json::from_value::(parameters).is_ok(), 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..aaec2c17e54 --- /dev/null +++ b/litellm-rust/crates/traces-clickhouse/tests/load.rs @@ -0,0 +1,227 @@ +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), \ + if({first_day} = 0, concat('load-', toString(number + {first_row})), 'load-0'), \ + 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(()) +} + +async fn trace_start_time(fixture: &SeededDatabase) -> TestResult { + let query = format!( + "SELECT toString(Timestamp, 'UTC') AS start_time FROM {DATABASE}.otel_traces \ + WHERE TraceId = 'load-0' LIMIT 1 FORMAT JSON" + ); + let response = fixture + .database + .client + .post(&fixture.database.url) + .body(query) + .send() + .await? + .error_for_status()? + .text() + .await?; + let result: Value = serde_json::from_str(&response)?; + result["data"][0]["start_time"] + .as_str() + .map(str::to_owned) + .ok_or_else(|| "trace start time missing".into()) +} + +fn content_parameters(start_time: &str) -> 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())), + ("id".into(), Parameter::Text("load-0".into())), + ("record_team".into(), Parameter::Text("load-team".into())), + ("start_time".into(), Parameter::Text(start_time.into())), + ("trace_ref".into(), Parameter::Text(String::new())), + ("cursor".into(), Parameter::Text(String::new())), + ("offset".into(), Parameter::Integer(1)), + ]) +} + +async fn content(fixture: &SeededDatabase, start_time: &str, query_id: &str) -> TestResult { + let connection = Connection::configured( + &format!("{}?query_id={query_id}", fixture.database.url), + DATABASE, + "default", + "", + )?; + let response = execute_named_read( + &fixture.database.client, + &connection, + ReadQuery::Content, + &content_parameters(start_time), + ) + .await?; + let result: Value = serde_json::from_str(&response)?; + assert!(!result["data"].as_array().ok_or("content rows")?.is_empty()); + 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(()) +} + +#[rstest] +#[tokio::test] +async fn lens_content_reads_scale_with_trace_not_retention( + #[future(awt)] migrated_database: TestResult, +) -> TestResult { + let fixture = migrated_database?; + seed_days(&fixture, 0, 8).await?; + let start_time = trace_start_time(&fixture).await?; + let before_id = format!("lens_content_before_{}", std::process::id()); + content(&fixture, &start_time, &before_id).await?; + let before = query_read_rows(&fixture, &before_id).await?; + + seed_days(&fixture, 8, 24).await?; + let after_id = format!("lens_content_after_{}", std::process::id()); + content(&fixture, &start_time, &after_id).await?; + let after = query_read_rows(&fixture, &after_id).await?; + println!("lens_content read_rows: before={before}, after={after}"); + assert!( + after * 100 <= before * 105, + "read_rows grew from {before} to {after}" + ); + Ok(()) +} diff --git a/litellm-rust/crates/traces-clickhouse/tests/migrations.rs b/litellm-rust/crates/traces-clickhouse/tests/migrations.rs index e51d9083c59..7630b3033a7 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/migrations.rs +++ b/litellm-rust/crates/traces-clickhouse/tests/migrations.rs @@ -1150,6 +1150,7 @@ async fn lens_filters_reads_and_evidence_keep_reused_trace_ids_separate( ("source".into(), Parameter::Text("traces".into())), ("id".into(), Parameter::Text("shared".into())), ("record_team".into(), Parameter::Text("team".into())), + ("start_time".into(), Parameter::Text(String::new())), ("trace_ref".into(), Parameter::Text(first_ref.into())), ("cursor".into(), Parameter::Text(String::new())), ("offset".into(), Parameter::Integer(1)), @@ -1177,6 +1178,7 @@ async fn lens_filters_reads_and_evidence_keep_reused_trace_ids_separate( ("source".into(), Parameter::Text("traces".into())), ("id".into(), Parameter::Text("shared".into())), ("record_team".into(), Parameter::Text("team".into())), + ("start_time".into(), Parameter::Text(String::new())), ("trace_ref".into(), Parameter::Text(first_ref.into())), ("span".into(), Parameter::Text("root".into())), ("quote".into(), Parameter::Text(opposite.into())), @@ -1339,7 +1341,7 @@ async fn lens_selection_pages_without_losing_or_repeating_runs( #[case::traces("traces", 9)] #[case::requests("requests", 3)] #[tokio::test] -async fn lens_content_keeps_original_span_and_request_timestamps( +async fn lens_content_keeps_original_timestamps_with_start_time_slack( #[future(awt)] database: TestResult, #[case] source: &str, #[case] precision: usize, @@ -1379,6 +1381,30 @@ async fn lens_content_keeps_original_span_and_request_timestamps( ) .await?; let connection = Connection::configured(&database.url, "trace_test", "default", "")?; + let start_time_body = execute_read( + &database.client, + &connection, + "SELECT toString(fromUnixTimestamp64Nano({timestamp:Int64})) AS start_time FORMAT JSON", + &BTreeMap::from([( + "timestamp".into(), + Parameter::Integer(root_start + 86_400_000_000_000), + )]), + ) + .await?; + let start_time: serde_json::Value = serde_json::from_str(&start_time_body)?; + let start_time = start_time["data"][0]["start_time"] + .as_str() + .ok_or("start time missing")? + .to_owned(); + let parsed_time_body = execute_read( + &database.client, + &connection, + "SELECT toString(parseDateTime64BestEffortOrZero({start_time:String}, 9)) AS start_time FORMAT JSON", + &BTreeMap::from([("start_time".into(), Parameter::Text(start_time.clone()))]), + ) + .await?; + let parsed_time: serde_json::Value = serde_json::from_str(&parsed_time_body)?; + assert_eq!(parsed_time["data"][0]["start_time"], start_time); let parameters = BTreeMap::from([ ("source".into(), Parameter::Text(source.into())), ("all_teams".into(), Parameter::Integer(0)), @@ -1386,6 +1412,7 @@ async fn lens_content_keeps_original_span_and_request_timestamps( ("record_team".into(), Parameter::Text("team".into())), ("key_hash".into(), Parameter::Text(String::new())), ("trace_ref".into(), Parameter::Text(String::new())), + ("start_time".into(), Parameter::Text(start_time)), ("id".into(), Parameter::Text("run".into())), ("cursor".into(), Parameter::Text(String::new())), ("offset".into(), Parameter::Integer(1)), @@ -1453,6 +1480,7 @@ async fn lens_content_keeps_output_visible_after_long_input( ("record_team".into(), Parameter::Text("team".into())), ("key_hash".into(), Parameter::Text(String::new())), ("trace_ref".into(), Parameter::Text(String::new())), + ("start_time".into(), Parameter::Text(String::new())), ("id".into(), Parameter::Text("request".into())), ("cursor".into(), Parameter::Text(String::new())), ("offset".into(), Parameter::Integer(1)), 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/caching/caching.py b/litellm/caching/caching.py index 1dc4de04dc1..85ef5a93937 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -71,6 +71,17 @@ class CacheMode(str, Enum): #### LiteLLM.Completion / Embedding Cache #### +def _request_message_count(kwargs: Mapping[str, object]) -> int: + """Chat and Messages API `messages`, else Responses API `input` items; embedding `input` strings count as none""" + messages: Final = kwargs.get("messages") + if isinstance(messages, list): + return len(messages) + input_items: Final = kwargs.get("input") + if not isinstance(input_items, list): + return 0 + return sum(1 for item in input_items if isinstance(item, (Mapping, BaseModel))) + + class Cache: def __init__( self, @@ -119,6 +130,7 @@ class Cache: semantic_cache_embedding_max_input_tokens: int | None = None, semantic_cache_embedding_timeout: float | None = None, semantic_cache_scope: str = SemanticCacheScope.KEY.value, + max_messages: int | None = 4, # GCP IAM authentication parameters gcp_service_account: str | None = None, gcp_ssl_ca_certs: str | None = None, @@ -148,6 +160,7 @@ class Cache: semantic_cache_embedding_max_input_tokens (int, optional): Truncate prompts to this many tokens before embedding them for semantic caching. Defaults to the embedding deployment's configured max_input_tokens. semantic_cache_embedding_timeout (float, optional): Seconds a semantic-cache lookup may spend embedding the prompt before it gives up and lets the request continue to the LLM. Defaults to SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS. semantic_cache_scope (str, optional): "key" isolates semantic-cache buckets per key/team/org. "end_user" additionally isolates per end user (falls back to the key scope when the request carries no end-user id). Defaults to "key". + max_messages (int, optional): Requests with more `messages` (or Responses API `input` items) than this are neither looked up nor stored, so long agent conversations never serve or create a cache entry. None disables the limit. Defaults to 4. # Disk Cache Args disk_cache_dir (str, optional): The directory for the disk cache. Defaults to None. @@ -298,6 +311,7 @@ class Cache: self.ttl = ttl self.mode: CacheMode = mode or CacheMode.default_on self.semantic_cache_scope: str = SemanticCacheScope(semantic_cache_scope).value + self.max_messages: int | None = max_messages if self.type == LiteLLMCacheType.LOCAL and default_in_memory_ttl is not None: self.ttl = default_in_memory_ttl @@ -933,7 +947,10 @@ class Cache: If cache is default_on then this is True If cache is default_off then this is only true when user has opted in to use cache + Always False once the request carries more than `max_messages` messages """ + if self.max_messages is not None and _request_message_count(kwargs) > self.max_messages: + return False if self.mode == CacheMode.default_on: return True diff --git a/litellm/caching/qdrant_semantic_cache.py b/litellm/caching/qdrant_semantic_cache.py index 8dac3f2eef9..eb88df066f6 100644 --- a/litellm/caching/qdrant_semantic_cache.py +++ b/litellm/caching/qdrant_semantic_cache.py @@ -24,7 +24,7 @@ from litellm.constants import ( ) from litellm.litellm_core_utils.asyncify import asyncify from litellm.litellm_core_utils.prompt_templates.common_utils import ( - get_str_from_messages, + get_semantic_cache_prompt_from_messages, ) from litellm.types.utils import EmbeddingResponse @@ -286,7 +286,7 @@ class QdrantSemanticCache(BaseCache): # get the prompt messages: Final = kwargs["messages"] - prompt: Final = get_str_from_messages(messages) + prompt: Final = get_semantic_cache_prompt_from_messages(messages) # create an embedding for prompt embedding_response: Final = cast( @@ -325,7 +325,7 @@ class QdrantSemanticCache(BaseCache): # get the messages messages: Final = kwargs["messages"] - prompt: Final = get_str_from_messages(messages) + prompt: Final = get_semantic_cache_prompt_from_messages(messages) # convert to embedding embedding_response: Final = cast( @@ -400,7 +400,7 @@ class QdrantSemanticCache(BaseCache): # get the prompt messages: Final = kwargs["messages"] - prompt: Final = get_str_from_messages(messages) + prompt: Final = get_semantic_cache_prompt_from_messages(messages) embedding_response: Final = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata")) # get the embedding @@ -435,7 +435,7 @@ class QdrantSemanticCache(BaseCache): # get the messages messages: Final = kwargs["messages"] - prompt: Final = get_str_from_messages(messages) + prompt: Final = get_semantic_cache_prompt_from_messages(messages) embedding_response: Final = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata")) diff --git a/litellm/caching/redis_semantic_cache.py b/litellm/caching/redis_semantic_cache.py index d4c815e15b7..8f99d76ba46 100644 --- a/litellm/caching/redis_semantic_cache.py +++ b/litellm/caching/redis_semantic_cache.py @@ -21,7 +21,7 @@ from litellm._logging import print_verbose, verbose_logger from litellm.constants import SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS from litellm.litellm_core_utils.asyncify import asyncify from litellm.litellm_core_utils.prompt_templates.common_utils import ( - get_str_from_messages, + get_semantic_cache_prompt_from_messages, ) from litellm.types.utils import EmbeddingResponse @@ -263,7 +263,7 @@ class RedisSemanticCache(BaseCache): """ messages: Final = kwargs.get("messages") if messages: - return get_str_from_messages(messages) + return get_semantic_cache_prompt_from_messages(messages) if "input" not in kwargs: return None @@ -274,7 +274,7 @@ class RedisSemanticCache(BaseCache): return prompt or None @classmethod - def _collect_responses_input_text(cls, value: object, prompt_parts: list[str]) -> None: + def _collect_responses_input_text(cls, value: object, prompt_parts: list[str]) -> None: # noqa: C901 # one branch per Responses input shape value = cls._coerce_response_input_value(value) if value is None: return @@ -296,6 +296,11 @@ class RedisSemanticCache(BaseCache): cls._collect_responses_input_text(content, prompt_parts) return + output = value.get("output") + if isinstance(output, list): + cls._collect_responses_input_text(output, prompt_parts) + return + for text_key in ("text", "output", "input_text", "output_text"): text_value = value.get(text_key) if isinstance(text_value, str): 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/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index a315d88d552..6af82d9546b 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -24,6 +24,7 @@ from litellm.types.guardrails import ( DynamicGuardrailParams, GuardrailEventHooks, LitellmParams, + LoggingOnlyScope, Mode, ) from litellm.types.llms.openai import AllMessageValues @@ -180,6 +181,7 @@ class CustomGuardrail(CustomLogger): use_native_lifecycle_hooks: ClassVar[bool] = False records_own_guardrail_information: ClassVar[bool] = False + logging_only_scope: LoggingOnlyScope | None timeout: float | httpx.Timeout | None = None @@ -256,6 +258,7 @@ class CustomGuardrail(CustomLogger): self.run_in_parallel: bool = run_in_parallel self.scan_raw_request: bool = scan_raw_request self.only_scan_new_messages: bool = only_scan_new_messages + self.logging_only_scope = None if timeout is not None: self.timeout = timeout @@ -817,6 +820,13 @@ class CustomGuardrail(CustomLogger): def uses_apply_guardrail_interface(self) -> bool: return type(self).apply_guardrail is not CustomGuardrail.apply_guardrail + @classmethod + def supports_logging_only_scope(cls) -> bool: + return ( + cls.apply_guardrail is not CustomGuardrail.apply_guardrail + and cls.async_logging_hook is CustomGuardrail.async_logging_hook + ) + def _deployment_hook_target(self) -> "CustomLogger": if not self.uses_apply_guardrail_interface() or self.use_native_lifecycle_hooks: return self @@ -948,7 +958,7 @@ class CustomGuardrail(CustomLogger): result: object, call_type: str, ) -> tuple[dict, object]: - """logging_only: run apply_guardrail on copies of the logged request/response and record the verdict.""" + """logging_only: scan copies of the logged request and/or response according to logging_only_scope.""" from litellm.llms import get_guardrail_translation_mapping if not self.uses_apply_guardrail_interface(): @@ -995,6 +1005,28 @@ class CustomGuardrail(CustomLogger): "standard_logging_object": {**standard_logging_object, "guardrail_information": [*existing, *entries]}, }, result + def _copy_scratch_request_fields( + self, + kwargs: Mapping[str, object], + ) -> tuple[object, object] | None: + optional_params: Final = kwargs.get("optional_params") + try: + return ( + copy.deepcopy(kwargs.get("messages") or kwargs.get("input")), + copy.deepcopy(optional_params.get("tools") if isinstance(optional_params, Mapping) else None), + ) + except Exception as e: + if self.logging_only_scope == "output": + return None + if self.logging_only_scope == "both": + verbose_logger.warning( + "Guardrail %s: logging_only request copy failed, skipping request scan: %s", + self.guardrail_name, + e, + ) + return None + raise + async def _scan_logged_call( self, kwargs: dict, @@ -1003,18 +1035,25 @@ class CustomGuardrail(CustomLogger): output_translation: "BaseTranslation", scratch_metadata: dict, ) -> None: - optional_params: Final = kwargs.get("optional_params") or {} - scratch_input: Final = copy.deepcopy(kwargs.get("messages") or kwargs.get("input")) + scratch_fields: Final = self._copy_scratch_request_fields(kwargs) + scratch_input, scratch_tools = scratch_fields or (None, None) scratch_request: Final = { "model": kwargs.get("model"), "messages": scratch_input, "input": scratch_input, - "tools": copy.deepcopy(optional_params.get("tools")), + "tools": scratch_tools, "litellm_call_id": kwargs.get("litellm_call_id"), "metadata": scratch_metadata, } - await translation.process_input_messages(data=scratch_request, guardrail_to_apply=self) - if response is None: + if self.logging_only_scope != "output" and scratch_fields is not None: + if self.logging_only_scope == "both": + try: + await translation.process_input_messages(data=scratch_request, guardrail_to_apply=self) + except Exception as e: # noqa: BLE001 # one direction's scan failure must not drop the other direction's verdict + verbose_logger.warning("Guardrail %s: logging_only scan raised: %s", self.guardrail_name, e) + else: + await translation.process_input_messages(data=scratch_request, guardrail_to_apply=self) + if response is None or self.logging_only_scope == "input": return await output_translation.process_output_response( response=copy.deepcopy(response), guardrail_to_apply=self, request_data=scratch_request diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index 6a4987298c4..45fce665bf0 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -909,7 +909,8 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac """ import litellm from litellm import Choices, Message, ModelResponse - from litellm.litellm_core_utils.classifier_logging import CLASSIFIER_AUDIT_FIELDS, without_classifier_audit + from litellm.litellm_core_utils.classifier_logging import CLASSIFIER_AUDIT_FIELDS + from litellm.litellm_core_utils.redact_messages import redacted_litellm_params turn_off_message_logging: Final[bool] = getattr(self, "turn_off_message_logging", False) excluded_fields: Final[list[str] | None] = getattr(litellm, "standard_logging_payload_excluded_fields", None) @@ -918,9 +919,15 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac if turn_off_message_logging is False and not excluded_fields: return model_call_details + params: Final = model_call_details.get("litellm_params") + redacted_params: Final = ( + MappingProxyType({"litellm_params": redacted_litellm_params(params)}) + if turn_off_message_logging and isinstance(params, Mapping) + else EMPTY_MAPPING + ) standard_logging_object: Final = model_call_details.get("standard_logging_object") if standard_logging_object is None: - return model_call_details.copy() + return {**model_call_details, **redacted_params} # Make a copy of just the standard_logging_object to avoid modifying the original standard_logging_object_copy: Final = { @@ -960,13 +967,6 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac model_response_dict: Final = model_response.model_dump() standard_logging_object_copy["response"] = model_response_dict - params: Final = model_call_details.get("litellm_params") - request: Final = params.get("proxy_server_request") if isinstance(params, dict) else None - redacted_params: Final = ( - MappingProxyType({"litellm_params": {**params, "proxy_server_request": without_classifier_audit(request)}}) - if turn_off_message_logging and isinstance(params, dict) and isinstance(request, dict) - else EMPTY_MAPPING - ) return { **model_call_details, **redacted_params, diff --git a/litellm/integrations/otel/model/spans.py b/litellm/integrations/otel/model/spans.py index 3b17aea3cd8..b22a8b9415b 100644 --- a/litellm/integrations/otel/model/spans.py +++ b/litellm/integrations/otel/model/spans.py @@ -353,6 +353,8 @@ _PRISMA_MODELS: Final[frozenset[str]] = frozenset( "LiteLLM_LensRun", "LiteLLM_LensReview", "LiteLLM_LensWorker", + "LiteLLM_LensSignalConfig", + "LiteLLM_LensTraceSignal", ) ) PRISMA_RELATIONS: Final[frozenset[str]] = _PRISMA_MODELS | _PRISMA_VIEWS 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/logging_worker.py b/litellm/litellm_core_utils/logging_worker.py index 16cc352cda0..1602fb81e26 100644 --- a/litellm/litellm_core_utils/logging_worker.py +++ b/litellm/litellm_core_utils/logging_worker.py @@ -23,6 +23,19 @@ from litellm.constants import ( MAX_TIME_TO_CLEAR_QUEUE, ) +_CALLBACK_DEADLINE: Final[contextvars.ContextVar[float | None]] = contextvars.ContextVar( + "logging_callback_deadline", default=None +) + + +def optional_callback_budget(maximum: float, *, fraction: float = 0.25) -> float: + deadline: Final = _CALLBACK_DEADLINE.get() + return ( + maximum + if deadline is None + else max(0.0, min(maximum, (deadline - asyncio.get_running_loop().time()) * fraction)) + ) + def _coroutine_name(coroutine: Coroutine) -> str: return getattr(coroutine, "__qualname__", None) or getattr(coroutine, "__name__", None) or type(coroutine).__name__ @@ -100,12 +113,20 @@ class LoggingWorker: return len(revived) def _run_coroutine_silently(self, loop: asyncio.AbstractEventLoop, coroutine: Coroutine) -> bool: + token: Final = _CALLBACK_DEADLINE.set(loop.time() + self.timeout) try: loop.run_until_complete(asyncio.wait_for(coroutine, timeout=self.timeout)) except (Exception, asyncio.CancelledError): # noqa: BLE001 # atexit flush must never break the user's program return False + finally: + _CALLBACK_DEADLINE.reset(token) return True + def _create_callback_task(self, task: LoggingTask) -> asyncio.Task[object]: + context: Final = task["context"].copy() + context.run(_CALLBACK_DEADLINE.set, asyncio.get_running_loop().time() + self.timeout) + return context.run(asyncio.create_task, task["coroutine"]) + @staticmethod def _drain_pending(queue: "asyncio.Queue[LoggingTask]") -> tuple[LoggingTask, ...]: """Pop every task still queued, without awaiting them, so they can be moved to another queue.""" @@ -172,7 +193,7 @@ class LoggingWorker: try: if self._queue is not None: # Run the coroutine in its original context - callback_task: Final = task["context"].run(asyncio.create_task, task["coroutine"]) + callback_task: Final = self._create_callback_task(task) try: await asyncio.wait_for(callback_task, timeout=self.timeout) except asyncio.TimeoutError as e: @@ -424,7 +445,7 @@ class LoggingWorker: try: await asyncio.wait_for( - task["context"].run(asyncio.create_task, task["coroutine"]), + self._create_callback_task(task), timeout=self.timeout, ) except Exception: @@ -517,7 +538,7 @@ class LoggingWorker: # Await the coroutine to properly execute and avoid "never awaited" warnings try: await asyncio.wait_for( - task["context"].run(asyncio.create_task, task["coroutine"]), + self._create_callback_task(task), timeout=self.timeout, ) except Exception: diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index becbc5fb1c5..0e0f51e5f8b 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -192,6 +192,33 @@ def get_str_from_messages(messages: list[AllMessageValues]) -> str: return text +def get_semantic_cache_prompt_from_messages(messages: Sequence[Mapping[str, object]]) -> str: + """ + The text a semantic cache embeds for a request: `get_str_from_messages` plus the text inside + Messages API `tool_result` blocks, so a tool turn does not embed identically to the turn before it + """ + return "".join( + _semantic_cache_content_text(message.get("content")) + + extract_search_results_text(message.get("search_results")) + for message in messages + ) + + +def _semantic_cache_content_text(content: object) -> str: + if isinstance(content, str): + return content + if not isinstance(content, list): + return "" + return "".join(_semantic_cache_block_text(block) for block in content if isinstance(block, Mapping)) + + +def _semantic_cache_block_text(block: Mapping[str, object]) -> str: + if block.get("type") == "tool_result": + return _semantic_cache_content_text(block.get("content")) + text: Final = block.get("text") + return text if isinstance(text, str) else "" + + def is_non_content_values_set(message: AllMessageValues) -> bool: ignore_keys: Final = ["content", "role", "name"] return any(message.get(key, None) is not None for key in message if key not in ignore_keys) diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 973a2787b6a..4580f9bd01b 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -5188,7 +5188,7 @@ def function_call_prompt( messages: list[dict[str, object]], functions: list[object], ) -> list[dict[str, object]]: - function_prompt = """Produce JSON OUTPUT ONLY! Adhere to this format {"name": "function_name", "arguments":{"argument_name": "argument_value"}} The following functions are available to you:""" + function_prompt = """To call a function, reply with JSON ONLY in this format {"name": "function_name", "arguments":{"argument_name": "argument_value"}}. Once a function result answers the request, reply to the user in plain text instead of calling a function again. The following functions are available to you:""" for function in functions: function_prompt += f"""\n{function}\n""" 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/litellm_core_utils/redact_messages.py b/litellm/litellm_core_utils/redact_messages.py index 85ed0a40687..7c6c39abb76 100644 --- a/litellm/litellm_core_utils/redact_messages.py +++ b/litellm/litellm_core_utils/redact_messages.py @@ -11,6 +11,7 @@ import asyncio import copy import inspect from collections.abc import Mapping +from dataclasses import replace from typing import TYPE_CHECKING, Any, Final import litellm @@ -26,6 +27,7 @@ from litellm.llms.vertex_ai.common_utils import ( redact_vertex_ai_metadata_from_logged_object, ) from litellm.secret_managers.main import str_to_bool +from litellm.types.router import BaselineRouteStamp from litellm.types.utils import StandardCallbackDynamicParams if TYPE_CHECKING: @@ -252,6 +254,26 @@ def _redact_model_response_dict_choices(choices, redacted_str: str): _redact_choice_content(choice) +def _redacted_baseline_metadata(metadata: Mapping[str, object]) -> Mapping[str, object]: + route: Final = metadata.get("_autorouter_baseline_route") + if not isinstance(route, BaselineRouteStamp): + return metadata + return {**metadata, "_autorouter_baseline_route": replace(route, request_parameters=None)} + + +def redacted_litellm_params(params: Mapping[str, object]) -> dict[str, object]: + request: Final = params.get("proxy_server_request") + return { + **params, + **{ + key: _redacted_baseline_metadata(value) + for key, value in params.items() + if key in ("metadata", "litellm_metadata") and isinstance(value, Mapping) + }, + **({"proxy_server_request": without_classifier_audit(request)} if isinstance(request, Mapping) else {}), + } + + def perform_redaction(model_call_details: dict, result, redact_streaming_responses: bool = True): """ Performs the actual redaction on the logging object and result. @@ -262,9 +284,8 @@ def perform_redaction(model_call_details: dict, result, redact_streaming_respons """ # Redact model_call_details params: Final = model_call_details.get("litellm_params") - request: Final = params.get("proxy_server_request") if isinstance(params, dict) else None - if isinstance(params, dict) and isinstance(request, Mapping): - model_call_details["litellm_params"] = {**params, "proxy_server_request": without_classifier_audit(request)} + if isinstance(params, Mapping): + model_call_details["litellm_params"] = redacted_litellm_params(params) model_call_details["messages"] = [{"role": "user", "content": REDACTED_BY_LITELLM}] model_call_details["prompt"] = "" model_call_details["input"] = "" diff --git a/litellm/llms/anthropic/pass_through/messages/handler.py b/litellm/llms/anthropic/pass_through/messages/handler.py index 6d13e38aa45..4e7a154be67 100644 --- a/litellm/llms/anthropic/pass_through/messages/handler.py +++ b/litellm/llms/anthropic/pass_through/messages/handler.py @@ -15,9 +15,6 @@ import litellm from litellm.litellm_core_utils.exception_mapping_utils import exception_type from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.anthropic.common_utils import ( - flatten_unencrypted_web_search_results_in_anthropic_messages, - sanitize_tool_use_ids_in_anthropic_messages, - strip_empty_content_blocks_from_anthropic_messages, strip_provider_specific_fields_from_anthropic_messages, ) from litellm.llms.base_llm.anthropic_messages.transformation import ( @@ -36,9 +33,8 @@ from litellm.utils import ProviderConfigManager, client from ..adapters.handler import LiteLLMMessagesToCompletionTransformationHandler from ..responses_adapters.handler import LiteLLMMessagesToResponsesAPIHandler -from ..utils import is_reasoning_auto_summary_enabled from .interceptors import get_messages_interceptors -from .utils import AnthropicMessagesRequestUtils, mock_response +from .utils import AnthropicMessagesRequestUtils, mock_response, prepare_native_messages __all__ = ("anthropic_messages", "anthropic_messages_handler") @@ -251,28 +247,7 @@ async def anthropic_messages( Runs the empty-content-block sanitizer before any backend dispatch. """ - # Anthropic's API rejects requests containing empty / whitespace-only - # text content blocks ("messages: text content blocks must be - # non-empty") and empty thinking blocks ("each thinking block must - # contain thinking"). Multi-turn tool-use clients (e.g. Claude Code) - # routinely loop assistant responses that contain such blocks — an empty - # text block alongside tool_use, or an empty thinking block from a turn - # a non-Anthropic reasoning model served through the bridge — back as - # conversation history, which then causes the next /v1/messages call to - # 400. /v1/chat/completions already handles this in - # anthropic_messages_pt; sanitize the native Anthropic Messages path - # here for the same guarantee. See #22930. - messages = strip_empty_content_blocks_from_anthropic_messages(messages) - # Replay of cross-provider tool history (e.g. kimi -> Anthropic) may carry - # ids like ``functions.Bash:0`` that violate Anthropic's id pattern. - messages = sanitize_tool_use_ids_in_anthropic_messages(messages) - messages = flatten_unencrypted_web_search_results_in_anthropic_messages(messages) - - from litellm.integrations.anthropic_cache_control_hook import ( - AnthropicCacheControlHook, - ) - - messages, system = AnthropicCacheControlHook.maybe_inject_cache_control( + messages, system = prepare_native_messages( messages, system, kwargs, model=model, custom_llm_provider=custom_llm_provider, tools=tools, api_base=api_base ) @@ -454,23 +429,15 @@ def anthropic_messages_handler( """ from litellm.types.utils import LlmProviders - # Sanitize empty text blocks so the sync entry point - # (litellm.messages.create -> anthropic_messages_handler) gets the same - # protection as the async wrapper. The async wrapper already sanitized and - # does not reassign messages before dispatch, so it sets - # ``_litellm_messages_presanitized`` to skip this redundant second - # full-messages scan. Pop it so it never leaks into provider params. - if not kwargs.pop("_litellm_messages_presanitized", False): - messages = strip_empty_content_blocks_from_anthropic_messages(messages) - messages = sanitize_tool_use_ids_in_anthropic_messages(messages) - messages = flatten_unencrypted_web_search_results_in_anthropic_messages(messages) - - from litellm.integrations.anthropic_cache_control_hook import ( - AnthropicCacheControlHook, - ) - - messages, system = AnthropicCacheControlHook.maybe_inject_cache_control( - messages, system, kwargs, model=model, custom_llm_provider=custom_llm_provider, tools=tools, api_base=api_base + messages, system = prepare_native_messages( + messages, + system, + kwargs, + model=model, + custom_llm_provider=custom_llm_provider, + tools=tools, + api_base=api_base, + presanitized=bool(kwargs.pop("_litellm_messages_presanitized", False)), ) metadata = validate_anthropic_api_metadata(metadata) @@ -645,14 +612,6 @@ def anthropic_messages_handler( custom_llm_provider=custom_llm_provider, ) ) - if is_reasoning_auto_summary_enabled(): - thinking_param: Final = anthropic_messages_optional_request_params.get("thinking") - if isinstance(thinking_param, dict) and thinking_param.get("type") != "disabled": - anthropic_messages_optional_request_params["thinking"] = { - **thinking_param, - "display": "summarized", - } - resolved_api_base: Final = ( dynamic_api_base if dynamic_api_base is not None and anthropic_messages_provider_config.uses_get_llm_provider_api_base() diff --git a/litellm/llms/anthropic/pass_through/messages/utils.py b/litellm/llms/anthropic/pass_through/messages/utils.py index dfab0af8eaa..8371615baea 100644 --- a/litellm/llms/anthropic/pass_through/messages/utils.py +++ b/litellm/llms/anthropic/pass_through/messages/utils.py @@ -2,6 +2,15 @@ from collections.abc import Iterable, Mapping, Sequence from functools import lru_cache from typing import TYPE_CHECKING, Any, Final, cast, get_type_hints +from pydantic import JsonValue + +from litellm.integrations.anthropic_cache_control_hook import AnthropicCacheControlHook +from litellm.llms.anthropic.common_utils import ( + flatten_unencrypted_web_search_results_in_anthropic_messages, + sanitize_tool_use_ids_in_anthropic_messages, + strip_empty_content_blocks_from_anthropic_messages, +) +from litellm.llms.anthropic.pass_through.utils import is_reasoning_auto_summary_enabled from litellm.types.llms.anthropic import ( AnthropicMessagesRequestOptionalParams, AnthropicStopDetails, @@ -119,8 +128,40 @@ def anthropic_system_to_openai_message(system: object) -> ChatCompletionSystemMe return ChatCompletionSystemMessage(role="system", content=system) +def prepare_native_messages( + messages: list[dict[str, JsonValue]], + system: str | list[dict[str, JsonValue]] | None, + kwargs: dict[str, object], + *, + model: str, + custom_llm_provider: str | None = None, + tools: list[dict[str, JsonValue]] | None = None, + api_base: str | None = None, + presanitized: bool = False, +) -> tuple[list[dict[str, JsonValue]], str | list[dict[str, JsonValue]] | None]: + normalized: Final = ( + messages + if presanitized + else flatten_unencrypted_web_search_results_in_anthropic_messages( + sanitize_tool_use_ids_in_anthropic_messages(strip_empty_content_blocks_from_anthropic_messages(messages)) + ) + ) + return cast( # cast-ok: legacy normalizers and injection preserve the JSON message and system shapes + tuple[list[dict[str, JsonValue]], str | list[dict[str, JsonValue]] | None], + AnthropicCacheControlHook.maybe_inject_cache_control( + normalized, + system, + kwargs, + model=model, + custom_llm_provider=custom_llm_provider, + tools=tools, + api_base=api_base, + ), + ) + + @lru_cache(maxsize=1) -def _anthropic_messages_optional_param_keys() -> frozenset[str]: +def anthropic_messages_optional_param_keys() -> frozenset[str]: """ Valid AnthropicMessagesRequestOptionalParams keys. @@ -152,7 +193,7 @@ class AnthropicMessagesRequestUtils: Returns: AnthropicMessagesRequestOptionalParams instance with only the valid parameters """ - valid_keys: Final = _anthropic_messages_optional_param_keys() + valid_keys: Final = anthropic_messages_optional_param_keys() filtered_params: Final = {k: v for k, v in params.items() if k in valid_keys and v is not None} if model is not None: from litellm.llms.anthropic.chat.transformation import AnthropicConfig @@ -174,6 +215,13 @@ class AnthropicMessagesRequestUtils: drop_params=drop_params, output_key=param, ) + if is_reasoning_auto_summary_enabled(): + thinking_param: Final = filtered_params.get("thinking") + if isinstance(thinking_param, dict) and thinking_param.get("type") != "disabled": + return cast( + AnthropicMessagesRequestOptionalParams, + {**filtered_params, "thinking": {**thinking_param, "display": "summarized"}}, + ) return cast(AnthropicMessagesRequestOptionalParams, filtered_params) diff --git a/litellm/llms/anthropic/prompt_cache_prediction.py b/litellm/llms/anthropic/prompt_cache_prediction.py index c346f83839f..7ef6ca58a5a 100644 --- a/litellm/llms/anthropic/prompt_cache_prediction.py +++ b/litellm/llms/anthropic/prompt_cache_prediction.py @@ -5,6 +5,7 @@ import hashlib import json from collections.abc import Mapping, Sequence from dataclasses import dataclass, field +from functools import reduce from itertools import accumulate, groupby from types import MappingProxyType from typing import Annotated, Final, Literal, Protocol, TypeAlias @@ -13,20 +14,32 @@ import httpx from pydantic import ConfigDict, Field, JsonValue, StrictInt, TypeAdapter, ValidationError import litellm -from litellm.llms.anthropic.common_utils import AnthropicModelInfo, is_anthropic_oauth_key +from litellm.litellm_core_utils.dot_notation_indexing import delete_nested_value +from litellm.llms.anthropic.common_utils import ( + AnthropicModelInfo, + is_anthropic_oauth_key, + strip_provider_specific_fields_from_anthropic_messages, +) from litellm.llms.anthropic.count_tokens.handler import AnthropicCountTokensHandler from litellm.llms.anthropic.count_tokens.transformation import COUNT_TOKEN_OPTION_NAMES from litellm.llms.anthropic.pass_through.messages.transformation import ( DEFAULT_ANTHROPIC_API_VERSION, AnthropicMessagesConfig, ) +from litellm.llms.anthropic.pass_through.messages.utils import AnthropicMessagesRequestUtils, prepare_native_messages +from litellm.router_utils.baseline_request import ( + BASELINE_PARAMETERS, + capture_baseline_parameters, +) from litellm.types.llms.base import LiteLLMBaseModel -from litellm.types.router import LiteLLM_Params +from litellm.types.router import GenericLiteLLMParams, LiteLLM_Params from litellm.types.utils import ModelResponse from litellm.utils import supports_thinking_cache_preservation _JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) _HEADERS: Final = TypeAdapter(dict[str, str]) +_MESSAGES: Final = TypeAdapter(list[dict[str, JsonValue]]) +_SYSTEM: Final = TypeAdapter(str | list[dict[str, JsonValue]] | None) _counter: Final = AnthropicCountTokensHandler() @@ -325,7 +338,7 @@ def _entry_fingerprint(fingerprint: str, ttl_seconds: int) -> str: def parse_cache_plan(body: Mapping[str, JsonValue]) -> PromptCachePlan | UnsupportedCachePlan: try: - request: Final = _PlanRequest.model_validate(body) + request: Final = _PlanRequest.model_validate(dict(body)) positions: Final = _positions(body) except ValidationError: return UnsupportedCachePlan("unsupported_prompt_shape") @@ -618,13 +631,56 @@ def resolve_baseline_prediction_target(params: LiteLLM_Params) -> NativePredicti return _resolve_prediction_target(params, allow_configured_endpoint=True) +def prepare_native_baseline_body(request: Mapping[str, object], model: str) -> Mapping[str, JsonValue] | None: + parameters: Final = capture_baseline_parameters(request) + if parameters is None: + return None + source: Final = {**parameters, "messages": request.get("messages"), "stream": request.get("stream", False)} + try: + owned: Final = _JSON_OBJECT.validate_python(source) + context: Final = {**{k: v for k, v in request.items() if k not in ("metadata", "litellm_metadata")}, **owned} + resolved_model: Final = litellm.get_llm_provider(model=model, custom_llm_provider="anthropic")[0] + messages, system = prepare_native_messages( + _MESSAGES.validate_python(owned.get("messages")), + _SYSTEM.validate_python(owned.get("system")), + context, + model=resolved_model, + custom_llm_provider="anthropic", + tools=_MESSAGES.validate_python(owned.get("tools") or []), + ) + options: Final = AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param( + {**owned, "system": system}, + model=resolved_model, + custom_llm_provider="anthropic", + drop_params=owned.get("drop_params") is True, + ) + filtered: Final = reduce( + delete_nested_value, + TypeAdapter(tuple[str, ...]).validate_python(owned.get("additional_drop_params") or ()), + dict(options), + ) + body: Final = AnthropicMessagesConfig().transform_anthropic_messages_request( + model=resolved_model, + messages=strip_provider_specific_fields_from_anthropic_messages(messages), + anthropic_messages_optional_request_params=filtered, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + return MappingProxyType(_JSON_OBJECT.validate_python(body)) + except Exception: # noqa: BLE001 # an unsupported hypothetical request is unavailable, never an inference failure + return None + + def _resolve_prediction_target( params: LiteLLM_Params, *, allow_configured_endpoint: bool, ) -> NativePredictionTarget | UnsupportedPredictionTarget: configured_options: Final = frozenset(params.model_dump(exclude_defaults=True, exclude_none=True)) - if configured_options - _DEPLOYMENT_OPTIONS: + allowed: Final = ( + _DEPLOYMENT_OPTIONS | frozenset(BASELINE_PARAMETERS) if allow_configured_endpoint else _DEPLOYMENT_OPTIONS + ) + if configured_options - allowed: return UnsupportedPredictionTarget("unsupported_deployment_configuration") api_base: Final = AnthropicModelInfo.get_api_base(params.api_base) if not allow_configured_endpoint and api_base not in ( diff --git a/litellm/llms/base_llm/chat/transformation.py b/litellm/llms/base_llm/chat/transformation.py index 5bae66c2681..ce7b0796b23 100644 --- a/litellm/llms/base_llm/chat/transformation.py +++ b/litellm/llms/base_llm/chat/transformation.py @@ -65,6 +65,26 @@ class BaseLLMException(Exception): super().__init__(self.message) # Call the base class constructor with the parameters it needs +_NO_ATTRIBUTION_HEADERS: Final[Mapping[str, str]] = types.MappingProxyType({}) + + +def with_attribution_headers( + attribution_headers: Mapping[str, str], + headers: dict[str, str] | None, # mutable-ok: returned as-is when there is nothing to add +) -> dict[str, str] | None: # mutable-ok: becomes the request's outbound headers + """ + `headers` plus any attribution header the caller didn't already set (names + compared case-insensitively). Builds a new dict; `headers` is never mutated. + """ + if not attribution_headers: + return headers + caller_names: Final = {name.lower() for name in headers or {}} + return { + **{name: value for name, value in attribution_headers.items() if name.lower() not in caller_names}, + **(headers or {}), + } + + class BaseConfig(ABC): def __init__(self): pass @@ -89,6 +109,15 @@ class BaseConfig(ABC): and not callable(v) # Filter out any callable objects including mocks } + def get_attribution_headers(self) -> Mapping[str, str]: + """ + Headers that tell the provider a request came through LiteLLM. + + Sent by default on every request; a caller header with the same name + (any casing) wins. Override in a provider config to opt in. + """ + return _NO_ATTRIBUTION_HEADERS + def get_json_schema_from_pydantic_object(self, response_format: type[BaseModel] | dict | None) -> dict | None: return type_to_response_format_param(response_format=response_format) diff --git a/litellm/llms/novita/chat/transformation.py b/litellm/llms/novita/chat/transformation.py index acdfa7e8790..f1e1e3c91b1 100644 --- a/litellm/llms/novita/chat/transformation.py +++ b/litellm/llms/novita/chat/transformation.py @@ -6,11 +6,20 @@ Calls done in OpenAI/openai.py as Novita AI is openai-compatible. Docs: https://novita.ai/docs/guides/llm-api """ +from collections.abc import Mapping +from types import MappingProxyType +from typing import Final + from ....types.llms.openai import AllMessageValues from ...openai.chat.gpt_transformation import OpenAIGPTConfig +_NOVITA_ATTRIBUTION_HEADERS: Final[Mapping[str, str]] = MappingProxyType({"X-Novita-Source": "litellm"}) + class NovitaConfig(OpenAIGPTConfig): + def get_attribution_headers(self) -> Mapping[str, str]: + return _NOVITA_ATTRIBUTION_HEADERS + def validate_environment( self, headers: dict, @@ -27,5 +36,6 @@ class NovitaConfig(OpenAIGPTConfig): ) headers["Authorization"] = f"Bearer {api_key}" headers["Content-Type"] = "application/json" - headers["X-Novita-Source"] = "litellm" + if not any(name.lower() == "x-novita-source" for name in headers): + headers["X-Novita-Source"] = "litellm" return headers diff --git a/litellm/llms/perplexity/chat/transformation.py b/litellm/llms/perplexity/chat/transformation.py index 2894218da67..5b85e1e8ae6 100644 --- a/litellm/llms/perplexity/chat/transformation.py +++ b/litellm/llms/perplexity/chat/transformation.py @@ -2,6 +2,8 @@ Translate from OpenAI's `/v1/chat/completions` to Perplexity's `/v1/chat/completions` """ +from collections.abc import Mapping +from types import MappingProxyType from typing import TYPE_CHECKING, Final import httpx @@ -17,12 +19,17 @@ from litellm.types.utils import ModelResponse, PromptTokensDetailsWrapper, Usage if TYPE_CHECKING: from litellm.litellm_core_utils.tokenizer import Encoding as Tokenizer +_PERPLEXITY_ATTRIBUTION_HEADERS: Final[Mapping[str, str]] = MappingProxyType({"X-Pplx-Integration": "litellm"}) + class PerplexityChatConfig(OpenAIGPTConfig): @property def custom_llm_provider(self) -> str | None: return "perplexity" + def get_attribution_headers(self) -> Mapping[str, str]: + return _PERPLEXITY_ATTRIBUTION_HEADERS + def _get_openai_compatible_provider_info( self, api_base: str | None, api_key: str | None ) -> tuple[str | None, str | None]: diff --git a/litellm/main.py b/litellm/main.py index 519b8aad0ef..37b290f51ba 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -118,6 +118,7 @@ from litellm.llms.base_llm import BaseConfig, BaseImageGenerationConfig from litellm.llms.base_llm.base_model_iterator import ( convert_model_response_to_streaming, ) +from litellm.llms.base_llm.chat.transformation import with_attribution_headers from litellm.llms.bedrock.common_utils import ( BedrockModelInfo, bedrock_route_for_request, @@ -2634,6 +2635,11 @@ def _complete_custom_openai( ) headers = headers or litellm.headers + outbound_headers: Final = ( + headers + if provider_config is None + else with_attribution_headers(provider_config.get_attribution_headers(), headers) + ) # Add GitHub Copilot headers (same as /responses endpoint does) if custom_llm_provider == "github_copilot": @@ -2685,7 +2691,7 @@ def _complete_custom_openai( acompletion=acompletion, stream=stream, api_key=api_key, - headers=headers, + headers=outbound_headers, client=client, provider_config=provider_config, ) @@ -2693,7 +2699,7 @@ def _complete_custom_openai( response = openai_chat_completions.completion( model=model, messages=messages, - headers=headers, + headers=outbound_headers, model_response=model_response, print_verbose=print_verbose, api_key=api_key, @@ -2716,7 +2722,7 @@ def _complete_custom_openai( input=messages, api_key=api_key, original_response=str(e), - additional_args={"headers": headers}, + additional_args={"headers": outbound_headers}, ) raise e @@ -2726,7 +2732,7 @@ def _complete_custom_openai( input=messages, api_key=api_key, original_response=response, - additional_args={"headers": headers}, + additional_args={"headers": outbound_headers}, ) return response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract 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/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index d864ec04a8d..40f644995bd 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -11572,6 +11572,23 @@ "description": "Google Cloud location/region (e.g., us-central1)", "title": "Location" }, + "logging_only_scope": { + "anyOf": [ + { + "enum": [ + "input", + "output", + "both" + ], + "type": "string" + }, + { + "type": "null" + } + ], + "description": "which direction a logging_only scan observes: 'input' (request), 'output' (response), or 'both' (default). Only applies to mode logging_only; pre_call/post_call on the same guardrail keep blocking.", + "title": "Logging Only Scope" + }, "mask_request_content": { "anyOf": [ { @@ -12727,6 +12744,77 @@ "title": "GuardrailSubmissionSummary", "type": "object" }, + "GuardrailUIAddGuardrailSettings": { + "properties": { + "content_filter_settings": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "title": "Content Filter Settings" + }, + "pii_entity_categories": { + "items": { + "$ref": "#/components/schemas/PiiEntityCategoryMap" + }, + "title": "Pii Entity Categories", + "type": "array" + }, + "providers_without_directional_logging_only_scope": { + "items": { + "type": "string" + }, + "title": "Providers Without Directional Logging Only Scope", + "type": "array" + }, + "supported_actions": { + "items": { + "type": "string" + }, + "title": "Supported Actions", + "type": "array" + }, + "supported_entities": { + "items": { + "type": "string" + }, + "title": "Supported Entities", + "type": "array" + }, + "supported_modes": { + "items": { + "type": "string" + }, + "title": "Supported Modes", + "type": "array" + }, + "supported_modes_by_provider": { + "additionalProperties": { + "items": { + "type": "string" + }, + "type": "array" + }, + "title": "Supported Modes By Provider", + "type": "object" + } + }, + "required": [ + "supported_entities", + "supported_actions", + "supported_modes", + "supported_modes_by_provider", + "providers_without_directional_logging_only_scope", + "pii_entity_categories" + ], + "title": "GuardrailUIAddGuardrailSettings", + "type": "object" + }, "HTTPValidationError": { "properties": { "detail": { @@ -13822,6 +13910,23 @@ "description": "Google Cloud location/region (e.g., us-central1)", "title": "Location" }, + "logging_only_scope": { + "anyOf": [ + { + "enum": [ + "input", + "output", + "both" + ], + "type": "string" + }, + { + "type": "null" + } + ], + "description": "which direction a logging_only scan observes: 'input' (request), 'output' (response), or 'both' (default). Only applies to mode logging_only; pre_call/post_call on the same guardrail keep blocking.", + "title": "Logging Only Scope" + }, "mask": { "anyOf": [ { @@ -14882,6 +14987,27 @@ "title": "PiiAction", "type": "string" }, + "PiiEntityCategoryMap": { + "properties": { + "category": { + "title": "Category", + "type": "string" + }, + "entities": { + "items": { + "type": "string" + }, + "title": "Entities", + "type": "array" + } + }, + "required": [ + "category", + "entities" + ], + "title": "PiiEntityCategoryMap", + "type": "object" + }, "PiiEntityType": { "enum": [ "CREDIT_CARD", @@ -16323,7 +16449,9 @@ "200": { "content": { "application/json": { - "schema": {} + "schema": { + "$ref": "#/components/schemas/GuardrailUIAddGuardrailSettings" + } } }, "description": "Successful Response" diff --git a/litellm/proxy/db/baseline_accounting.py b/litellm/proxy/db/baseline_accounting.py index b21b7c8a2b9..486717f9b93 100644 --- a/litellm/proxy/db/baseline_accounting.py +++ b/litellm/proxy/db/baseline_accounting.py @@ -223,7 +223,8 @@ ON CONFLICT (request_id) DO NOTHING _MARK_CONFLICT: Final = """ UPDATE "LiteLLM_AutoRouterBaselineObservation" SET conflicted = TRUE, revision = $4::bigint -WHERE request_id = $1 AND scope = $2 AND data <> $3 AND NOT conflicted +WHERE request_id = $1 AND scope = $2 AND NOT conflicted + AND (data::jsonb #- '{turn,turn_at}') <> ($3::jsonb #- '{turn,turn_at}') """ _READ_PAGE: Final = """ WITH times AS ( diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index 1dcc7cf5485..1c9639fa9dc 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -36,9 +36,11 @@ from litellm.proxy.guardrails.guardrail_hooks.custom_code.sandbox import ( ) from litellm.proxy.guardrails.guardrail_registry import ( GuardrailRegistry, + configured_event_hooks, contains_encrypted_marker, decrypt_guardrail_litellm_params, encrypt_guardrail_litellm_params, + parse_tolerant_litellm_params, ) from litellm.proxy.guardrails.usage_endpoints import router as guardrails_usage_router from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view @@ -283,7 +285,13 @@ async def list_guardrails_v2( number_of_asterisks=4, ) masked_litellm_params = ( - BaseLitellmParams.model_validate(masked_litellm_params_dict) if masked_litellm_params_dict else None + parse_tolerant_litellm_params( + masked_litellm_params_dict, + guardrail.get("guardrail_name") or "Unknown", + params_model=BaseLitellmParams, + ) + if masked_litellm_params_dict + else None ) guardrail_configs.append( GuardrailInfoResponse( @@ -324,7 +332,11 @@ async def list_guardrails_v2( number_of_asterisks=4, ) masked_in_memory_litellm_params_typed = ( - BaseLitellmParams.model_validate(masked_in_memory_litellm_params) + parse_tolerant_litellm_params( + masked_in_memory_litellm_params, + guardrail.get("guardrail_name") or "Unknown", + params_model=BaseLitellmParams, + ) if masked_in_memory_litellm_params else None ) @@ -425,7 +437,11 @@ async def create_guardrail( guardrail_id: Final = result.get("guardrail_id", "Unknown") try: - IN_MEMORY_GUARDRAIL_HANDLER.initialize_guardrail(guardrail=cast(Guardrail, result), source="db") + IN_MEMORY_GUARDRAIL_HANDLER.initialize_guardrail( + guardrail=cast(Guardrail, result), + source="db", + reject_invalid_logging_only_scope=True, + ) verbose_proxy_logger.info( "Immediate sync: Successfully initialized guardrail '%s' (ID: %s)", guardrail_name, guardrail_id ) @@ -550,7 +566,10 @@ async def update_guardrail( guardrail_name: Final = result.get("guardrail_name", "Unknown") try: - IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(guardrail=cast(Guardrail, result)) + IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db( + guardrail=cast(Guardrail, result), + reject_invalid_logging_only_scope=True, + ) verbose_proxy_logger.info( "Immediate sync: Successfully updated guardrail '%s' (ID: %s)", guardrail_name, guardrail_id ) @@ -1240,19 +1259,35 @@ async def patch_guardrail( # Update litellm_params if default_on is provided or pii_entities_config is provided existing_litellm_params: Final = _as_str_object_mapping(dict(existing_guardrail.get("litellm_params", {}))) - litellm_params = LitellmParams(**existing_litellm_params) - if request.litellm_params is not None: - requested_litellm_params: Final = request.litellm_params.model_dump(exclude_unset=True) - litellm_params_dict: Final = litellm_params.model_dump(exclude_unset=True) - litellm_params_dict.update(requested_litellm_params) - merged_litellm_params: Final = _as_str_object_mapping(litellm_params_dict) - try: - litellm_params = LitellmParams(**merged_litellm_params) - except ValidationError as validation_error: - raise HTTPException( - status_code=422, - detail=f"Invalid guardrail configuration, update rejected: {validation_error}", - ) from validation_error + current_litellm_params: Final = parse_tolerant_litellm_params( + existing_litellm_params, + existing_guardrail.get("guardrail_name") or "Unknown", + ) + requested_litellm_params: Final[Mapping[str, object]] = ( + MappingProxyType(request.litellm_params.model_dump(exclude_unset=True)) + if request.litellm_params is not None + else MappingProxyType({}) + ) + merged_litellm_params: Final = _as_str_object_mapping( + MappingProxyType({**current_litellm_params.model_dump(exclude_unset=True), **requested_litellm_params}) + ) + try: + parsed_litellm_params: Final = LitellmParams(**merged_litellm_params) + except ValidationError as validation_error: + raise HTTPException( + status_code=422, + detail=f"Invalid guardrail configuration, update rejected: {validation_error}", + ) from validation_error + clear_stored_scope: Final = ( + "logging_only_scope" not in requested_litellm_params + and parsed_litellm_params.logging_only_scope is not None + and GuardrailEventHooks.logging_only.value not in configured_event_hooks(parsed_litellm_params.mode) + ) + litellm_params: Final = ( + LitellmParams(**MappingProxyType({**merged_litellm_params, "logging_only_scope": None})) + if clear_stored_scope + else parsed_litellm_params + ) # Update guardrail_info if provided guardrail_info: Final = ( @@ -1281,6 +1316,7 @@ async def patch_guardrail( try: IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db( guardrail=guardrail, + reject_invalid_logging_only_scope="logging_only_scope" in requested_litellm_params, ) verbose_proxy_logger.info( "Immediate sync: Successfully updated guardrail '%s' (ID: %s)", guardrail_name, guardrail_id @@ -1294,15 +1330,7 @@ async def patch_guardrail( # the caller instead of a misleading 200. await GUARDRAIL_REGISTRY.update_guardrail_in_db( guardrail_id=guardrail_id, - guardrail=Guardrail( - guardrail_id=guardrail_id, - guardrail_name=existing_guardrail.get("guardrail_name") or "", - litellm_params=LitellmParams(**existing_litellm_params), - guardrail_info=existing_guardrail.get( - "guardrail_info", - {}, - ), - ), + guardrail=existing_guardrail, prisma_client=prisma_client, ) raise HTTPException( @@ -1404,7 +1432,13 @@ async def get_guardrail_info(guardrail_id: str): number_of_asterisks=4, ) masked_litellm_params = ( - BaseLitellmParams.model_validate(masked_litellm_params_dict) if masked_litellm_params_dict else None + parse_tolerant_litellm_params( + masked_litellm_params_dict, + result.get("guardrail_name") or "Unknown", + params_model=BaseLitellmParams, + ) + if masked_litellm_params_dict + else None ) return GuardrailInfoResponse( @@ -1427,7 +1461,7 @@ async def get_guardrail_info(guardrail_id: str): tags=["Guardrails"], dependencies=[Depends(user_api_key_auth)], ) -async def get_guardrail_ui_settings(): +async def get_guardrail_ui_settings() -> GuardrailUIAddGuardrailSettings: """ Get the UI settings for the guardrails @@ -1461,12 +1495,18 @@ async def get_guardrail_ui_settings(): # above; it only runs on pre_call. {SupportedGuardrailIntegrations.HIDE_SECRETS.value: [GuardrailEventHooks.pre_call.value]} ) + providers_without_directional_logging_only_scope: Final = tuple( + provider + for provider, guardrail_class in guardrail_class_registry.items() + if not guardrail_class.supports_logging_only_scope() + ) return GuardrailUIAddGuardrailSettings( supported_entities=[entity.value for entity in PiiEntityType], supported_actions=[action.value for action in PiiAction], supported_modes=[mode.value for mode in GuardrailEventHooks], supported_modes_by_provider=supported_modes_by_provider, + providers_without_directional_logging_only_scope=providers_without_directional_logging_only_scope, pii_entity_categories=category_maps, content_filter_settings={ "prebuilt_patterns": get_pattern_metadata(), diff --git a/litellm/proxy/guardrails/guardrail_initializers.py b/litellm/proxy/guardrails/guardrail_initializers.py index 0c1e22463ff..bacbd728d89 100644 --- a/litellm/proxy/guardrails/guardrail_initializers.py +++ b/litellm/proxy/guardrails/guardrail_initializers.py @@ -100,7 +100,7 @@ _MCP_EVENT_HOOKS: Final = frozenset( ) -def _configured_event_hooks(mode: str | list[str] | Mode) -> tuple[str, ...]: +def configured_event_hooks(mode: str | list[str] | Mode) -> tuple[str, ...]: if isinstance(mode, str): return (mode,) if isinstance(mode, list): @@ -114,7 +114,7 @@ def _configured_event_hooks(mode: str | list[str] | Mode) -> tuple[str, ...]: def _is_mcp_only_mode(mode: str | list[str] | Mode) -> bool: - hooks: Final = _configured_event_hooks(mode) + hooks: Final = configured_event_hooks(mode) return bool(hooks) and all(hook in _MCP_EVENT_HOOKS for hook in hooks) diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index 58b6b390576..8a430c71326 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -6,7 +6,8 @@ import os from collections.abc import Callable, Iterable, Iterator, Mapping, Sequence from datetime import datetime, timezone from itertools import chain, count -from typing import TYPE_CHECKING, Final, Literal, Optional, Protocol, TypeAlias, cast +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, Literal, Optional, Protocol, TypeAlias, TypeVar, cast from pydantic import BaseModel, TypeAdapter, ValidationError @@ -58,6 +59,7 @@ from .guardrail_hooks.llm_as_a_judge import ( initialize_guardrail as initialize_llm_as_a_judge, ) from .guardrail_initializers import ( + configured_event_hooks, initialize_bedrock, initialize_hide_secrets, initialize_lakera, @@ -572,9 +574,46 @@ def _as_callback_tuple( return (initialized,) -def _configure_callback_scoping( +def _logging_only_scope_error( custom_guardrail_callback: CustomGuardrail, guardrail_name: str, litellm_params: LitellmParams +) -> str | None: + logging_only_scope: Final = litellm_params.logging_only_scope + if logging_only_scope is not None and GuardrailEventHooks.logging_only.value not in configured_event_hooks( + litellm_params.mode + ): + return ( + f"Guardrail {guardrail_name}: logging_only_scope is set, but mode does not include logging_only, " + "so it would never apply. Add logging_only to mode or remove logging_only_scope." + ) + if logging_only_scope in ("input", "output") and not custom_guardrail_callback.supports_logging_only_scope(): + return ( + f"Guardrail {guardrail_name}: logging_only_scope={logging_only_scope!r} is not supported by this " + "guardrail, whose logging_only hook scans on its own. Remove logging_only_scope." + ) + return None + + +def _configure_callback_scoping( + custom_guardrail_callback: CustomGuardrail, + guardrail_name: str, + litellm_params: LitellmParams, + *, + reject_invalid_logging_only_scope: bool = False, ) -> None: + logging_only_scope: Final = litellm_params.logging_only_scope + logging_only_scope_error: Final = _logging_only_scope_error( + custom_guardrail_callback, guardrail_name, litellm_params + ) + if logging_only_scope_error is not None: + if reject_invalid_logging_only_scope: + raise ValueError(logging_only_scope_error) + verbose_proxy_logger.error( + "%s Ignoring logging_only_scope; the guardrail keeps its configured mode.", + logging_only_scope_error.replace("\r", "").replace("\n", ""), + ) + custom_guardrail_callback.logging_only_scope = None + else: + custom_guardrail_callback.logging_only_scope = logging_only_scope for scoping_param in ( "skip_system_message_in_guardrail", "skip_tool_message_in_guardrail", @@ -597,6 +636,28 @@ def _configure_callback_scoping( _apply_configured_bool_overrides(custom_guardrail_callback, litellm_params) +_ParamsT = TypeVar("_ParamsT", bound=BaseModel) + + +def parse_tolerant_litellm_params( + litellm_params_data: Mapping[str, object], + guardrail_name: str, + params_model: type[_ParamsT] = LitellmParams, +) -> _ParamsT: + try: + return params_model(**litellm_params_data) + except ValidationError as validation_error: + if any(tuple(error["loc"]) != ("logging_only_scope",) for error in validation_error.errors()): + raise + verbose_proxy_logger.error( + "Guardrail %s: logging_only_scope=%r is not one of 'input', 'output' or 'both'. " + "Ignoring logging_only_scope; the guardrail keeps its configured mode.", + guardrail_name.replace("\r", "").replace("\n", ""), + str(litellm_params_data.get("logging_only_scope")).replace("\r", "").replace("\n", "")[:100], + ) + return params_model(**MappingProxyType({**litellm_params_data, "logging_only_scope": None})) + + class InMemoryGuardrailHandler: """ Class that handles initializing guardrails and adding them to the CallbackManager @@ -633,6 +694,8 @@ class InMemoryGuardrailHandler: config_file_path: str | None = None, llm_router: Optional["Router"] = None, source: Literal["db", "config"] = "config", + *, + reject_invalid_logging_only_scope: bool = False, ) -> Guardrail | None: """ Initialize a guardrail from a dictionary and add it to the litellm callback manager @@ -653,7 +716,10 @@ class InMemoryGuardrailHandler: verbose_proxy_logger.debug("litellm_params= %s", litellm_params_data) if isinstance(litellm_params_data, dict): - litellm_params = LitellmParams(**litellm_params_data) + if reject_invalid_logging_only_scope: + litellm_params = LitellmParams(**litellm_params_data) + else: + litellm_params = parse_tolerant_litellm_params(litellm_params_data, guardrail["guardrail_name"]) else: litellm_params = litellm_params_data @@ -679,8 +745,18 @@ class InMemoryGuardrailHandler: config_file_path=config_file_path, llm_router=llm_router, ) - for custom_guardrail_callback in created_callbacks: - _configure_callback_scoping(custom_guardrail_callback, guardrail["guardrail_name"], litellm_params) + try: + for custom_guardrail_callback in created_callbacks: + _configure_callback_scoping( + custom_guardrail_callback, + guardrail["guardrail_name"], + litellm_params, + reject_invalid_logging_only_scope=reject_invalid_logging_only_scope, + ) + except Exception: + for custom_guardrail_callback in created_callbacks: + litellm.logging_callback_manager.remove_callback_from_all_lists(custom_guardrail_callback) + raise parsed_guardrail: Final = Guardrail( guardrail_id=guardrail.get("guardrail_id"), @@ -729,6 +805,28 @@ class InMemoryGuardrailHandler: siblings: Final = self.guardrail_id_to_sibling_callbacks.get(guardrail_id, ()) return (() if primary is None else (primary,)) + siblings + def _reject_invalid_logging_only_scope(self, guardrail_id: str, guardrail: Guardrail) -> None: + """ + Strictly validate logging_only_scope on a row whose params are otherwise + unchanged, without rebuilding the live callback. + + API write paths send the whole object, so an invalid scope must still be + rejected even when the write changed nothing else. But an unchanged row + must not force a teardown + re-append: initialize_guardrail appends the + rebuilt callback at the END of litellm.callbacks, so a no-op PUT would + reorder guardrails and change which one wins between a BLOCK and a MASK + guardrail over the same content. + """ + params: Final = guardrail.get("litellm_params") + if not isinstance(params, (dict, LitellmParams)): + return + litellm_params: Final = LitellmParams(**params) if isinstance(params, dict) else params + guardrail_name: Final = guardrail.get("guardrail_name", "Unknown") + for custom_guardrail_callback in self._tracked_callbacks(guardrail_id): + scope_error = _logging_only_scope_error(custom_guardrail_callback, guardrail_name, litellm_params) + if scope_error is not None: + raise ValueError(scope_error) + def initialize_custom_guardrail( self, guardrail: Guardrail, @@ -788,6 +886,8 @@ class InMemoryGuardrailHandler: guardrail_id: str, guardrail: Guardrail, source: Literal["db", "config"] = "db", + *, + reject_invalid_logging_only_scope: bool = False, ) -> None: """ Update a guardrail in memory: a changed name or litellm_params rebuilds the @@ -796,8 +896,14 @@ class InMemoryGuardrailHandler: """ updated_guardrail: Final = cast(Guardrail, {**guardrail, "guardrail_id": guardrail_id}) if self._has_guardrail_params_changed(guardrail_id, updated_guardrail): - self.reinitialize_guardrail(guardrail=updated_guardrail, source=source) + self.reinitialize_guardrail( + guardrail=updated_guardrail, + source=source, + reject_invalid_logging_only_scope=reject_invalid_logging_only_scope, + ) return + if reject_invalid_logging_only_scope: + self._reject_invalid_logging_only_scope(guardrail_id, updated_guardrail) self.IN_MEMORY_GUARDRAILS[guardrail_id] = updated_guardrail self._sources[guardrail_id] = source @@ -883,6 +989,7 @@ class InMemoryGuardrailHandler: @staticmethod def _normalize_litellm_params_for_comparison( params: LitellmParams | Mapping[str, object] | None, + guardrail_name: str, ) -> Mapping[str, object] | None: """ Render litellm_params to a canonical dict so an in-memory LitellmParams and @@ -899,7 +1006,7 @@ class InMemoryGuardrailHandler: return params.model_dump() if isinstance(params, dict): try: - return LitellmParams(**params).model_dump() + return parse_tolerant_litellm_params(params, guardrail_name).model_dump() except ValidationError as e: verbose_proxy_logger.warning( "Could not normalize guardrail litellm_params for comparison; treating the guardrail as changed. Error: %s", @@ -922,8 +1029,12 @@ class InMemoryGuardrailHandler: return True # Compare litellm_params - existing_dict: Final = self._normalize_litellm_params_for_comparison(existing.get("litellm_params")) - new_dict: Final = self._normalize_litellm_params_for_comparison(new_guardrail.get("litellm_params")) + existing_dict: Final = self._normalize_litellm_params_for_comparison( + existing.get("litellm_params"), existing.get("guardrail_name", "Unknown") + ) + new_dict: Final = self._normalize_litellm_params_for_comparison( + new_guardrail.get("litellm_params"), new_guardrail.get("guardrail_name", "Unknown") + ) # Compare and identify specific differences changed_fields = {} @@ -949,6 +1060,8 @@ class InMemoryGuardrailHandler: guardrail: Guardrail, config_file_path: str | None = None, source: Literal["db", "config"] = "config", + *, + reject_invalid_logging_only_scope: bool = False, ) -> Guardrail | None: """ Force re-initialization of a guardrail even if it exists in memory. @@ -978,7 +1091,12 @@ class InMemoryGuardrailHandler: # instance instead of leaving the guardrail silently removed: a guardrail # that was enforcing must never fail open because an update was bad. try: - return self.initialize_guardrail(guardrail=guardrail, config_file_path=config_file_path, source=source) + return self.initialize_guardrail( + guardrail=guardrail, + config_file_path=config_file_path, + source=source, + reject_invalid_logging_only_scope=reject_invalid_logging_only_scope, + ) except Exception as init_error: if previous_guardrail is not None: verbose_proxy_logger.exception( @@ -1003,7 +1121,9 @@ class InMemoryGuardrailHandler: ) if existing is None or db_params is None or not contains_encrypted_marker(db_params): return guardrail - loaded_params: Final = self._normalize_litellm_params_for_comparison(existing.get("litellm_params")) + loaded_params: Final = self._normalize_litellm_params_for_comparison( + existing.get("litellm_params"), guardrail.get("guardrail_name", "Unknown") + ) verbose_proxy_logger.warning( "Guardrail %s has litellm_params that do not decrypt with the current key; keeping the loaded values for " "them. Restart the proxy if the master key was rotated.", @@ -1023,7 +1143,13 @@ class InMemoryGuardrailHandler: } ) - def sync_guardrail_from_db(self, guardrail: Guardrail, config_file_path: str | None = None) -> Guardrail | None: + def sync_guardrail_from_db( + self, + guardrail: Guardrail, + config_file_path: str | None = None, + *, + reject_invalid_logging_only_scope: bool = False, + ) -> Guardrail | None: """ Sync a guardrail from DB - initializes if new, re-initializes if changed. DB values that do not decrypt with the current key keep the loaded guardrail's values. @@ -1044,8 +1170,12 @@ class InMemoryGuardrailHandler: guardrail=synced, config_file_path=config_file_path, source="db", + reject_invalid_logging_only_scope=reject_invalid_logging_only_scope, ) + if reject_invalid_logging_only_scope: + self._reject_invalid_logging_only_scope(guardrail_id, synced) + # Params unchanged but the entry is still DB-backed; make sure the # source marker reflects that even if it was previously set differently # (e.g. a config entry whose UUID later collided with a DB row). diff --git a/litellm/proxy/hooks/autorouter_baseline_cache.py b/litellm/proxy/hooks/autorouter_baseline_cache.py index e3cd6c67aa0..c4232609124 100644 --- a/litellm/proxy/hooks/autorouter_baseline_cache.py +++ b/litellm/proxy/hooks/autorouter_baseline_cache.py @@ -5,7 +5,7 @@ import hashlib import json import time from collections.abc import Callable, Mapping -from dataclasses import dataclass, replace +from dataclasses import dataclass, field, replace from datetime import datetime from types import MappingProxyType from typing import TYPE_CHECKING, Final @@ -16,9 +16,8 @@ from pydantic import ConfigDict, Field, JsonValue, TypeAdapter from litellm._logging import verbose_proxy_logger from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY from litellm.integrations.custom_logger import CustomLogger -from litellm.litellm_core_utils.core_helpers import ( - get_litellm_metadata_from_kwargs, # pyright: ignore[reportUnknownVariableType] # legacy metadata boundary validated below -) +from litellm.litellm_core_utils.core_helpers import get_metadata_variable_name_from_kwargs +from litellm.litellm_core_utils.logging_worker import optional_callback_budget from litellm.llms.anthropic.prompt_cache_prediction import ( CountedPromptCachePlan, NativePredictionTarget, @@ -28,6 +27,7 @@ from litellm.llms.anthropic.prompt_cache_prediction import ( count_cache_plan, count_prompt_tokens, parse_cache_plan, + prepare_native_baseline_body, resolve_baseline_prediction_target, supported_baseline_recipient, supported_prediction_headers, @@ -37,6 +37,8 @@ from litellm.proxy.spend_tracking.savings import ( _effective_model_info, # pyright: ignore[reportPrivateUsage] # existing deployment-price owner _proxy_llm_router, # pyright: ignore[reportPrivateUsage] # existing optional proxy-router owner ) +from litellm.router_strategy.complexity_router.context_compaction import compaction_applied +from litellm.router_utils.baseline_request import baseline_request from litellm.types.llms.base import LiteLLMBaseModel from litellm.types.router import BaselineRouteStamp from litellm.types.utils import CallTypes, ModelInfo, Usage @@ -66,6 +68,9 @@ class CapturedBaselineObservation(LiteLLMBaseModel): prices: ModelInfo | None observation: BaselineObservation + def with_observation(self, observation: BaselineObservation) -> CapturedBaselineObservation: + return self.model_copy(update={"observation": observation}) + @dataclass(frozen=True, slots=True) class BaselineCacheContext: @@ -73,7 +78,10 @@ class BaselineCacheContext: capture: CapturedBaselineObservation target: NativePredictionTarget | UnsupportedPredictionTarget baseline_deployment_id: str + baseline_body: Mapping[str, JsonValue] | None = field(default=None, repr=False) + selected_body_digest: str | None = field(default=None, repr=False) invalidated: str | None = None + finalization: asyncio.Task[CapturedBaselineObservation] | None = field(default=None, repr=False, compare=False) class _Metadata(LiteLLMBaseModel): @@ -102,6 +110,10 @@ def _digest(value: object) -> str: return hashlib.sha256(json.dumps(value, sort_keys=True, separators=(",", ":")).encode()).hexdigest() +def _native_body_digest(body: Mapping[str, JsonValue]) -> str: + return _digest({key: value for key, value in body.items() if key not in ("metadata", "stream")}) + + class AutoRouterBaselineCache(CustomLogger): def __init__( self, @@ -124,12 +136,15 @@ class AutoRouterBaselineCache(CustomLogger): if not isinstance(logging_obj, Logging) or call_type != CallTypes.anthropic_messages: return try: - metadata: Final = _METADATA.validate_python(get_litellm_metadata_from_kwargs({"litellm_params": kwargs})) + raw_metadata: Final = kwargs.get(get_metadata_variable_name_from_kwargs(kwargs)) + metadata: Final = _METADATA.validate_python(raw_metadata) if isinstance(raw_metadata, Mapping) else {} if metadata.get(INTERNAL_CALL_ORIGIN_METADATA_KEY): return if logging_obj.baseline_cache_context is not None: await invalidate_baseline_cache(logging_obj, "retried_request") return + if not isinstance(metadata.get("_autorouter_baseline_route"), BaselineRouteStamp): + return request: Final = _Metadata.model_validate(metadata) session: Final = kwargs.get("litellm_session_id") or request.session_id or logging_obj.litellm_session_id if not isinstance(session, str) or not session or len(session) > 256: @@ -142,13 +157,27 @@ class AutoRouterBaselineCache(CustomLogger): prices: Final = _PRICES.validate_python( _effective_model_info(router, request.route.baseline_deployment_id, request.route.baseline_model) ) + params: Final = ( + _METADATA.validate_python(deployment.litellm_params.model_dump(mode="json")) if deployment else {} + ) + projected: Final = ( + baseline_request( + kwargs, + request.route.request_parameters, + params, + include_extra_body=False, + ) + if request.route.request_parameters is not None + else None + ) scope: Final = "autorouter-baseline:v3:" + _digest( ( + "baseline_request_v4", request.user_api_key_hash, session, request.route.router_name, request.route.baseline_deployment_id, - deployment.litellm_params.model_dump(mode="json"), + params, prices, ) ) @@ -170,8 +199,21 @@ class AutoRouterBaselineCache(CustomLogger): reason="incomplete_response", ), ) + selected_model: Final = kwargs.get("model") + selected_body: Final = prepare_native_baseline_body( + kwargs, selected_model if isinstance(selected_model, str) else logging_obj.model + ) logging_obj.baseline_cache_context = BaselineCacheContext( - self, capture, target, request.route.baseline_deployment_id + self, + capture, + target, + request.route.baseline_deployment_id, + prepare_native_baseline_body(projected, target.model) + if projected is not None and isinstance(target, NativePredictionTarget) + else None, + _native_body_digest(selected_body) + if selected_body is not None and not compaction_applied(kwargs) + else None, ) except Exception: # noqa: BLE001 # optional observation cannot fail inference verbose_proxy_logger.warning("Auto-router baseline observation could not be initialized") @@ -197,14 +239,16 @@ class AutoRouterBaselineCache(CustomLogger): async def plan( self, target: NativePredictionTarget, wire: httpx.Request, body: Mapping[str, JsonValue], usage: Usage | None ) -> tuple[CountedPromptCachePlan | None, str | None]: + deadline: Final = asyncio.get_running_loop().time() + optional_callback_budget(_COUNT_TIMEOUT, fraction=0.75) if not supported_prediction_headers(wire.headers): return None, "unsupported_request_headers" plan: Final = parse_cache_plan(body) if isinstance(plan, UnsupportedCachePlan): return None, plan.reason details: Final = usage.prompt_tokens_details if usage is not None else None + selected: Final = parse_cache_plan(_JSON_BODY.validate_json(wire.content)) if ( - not plan.breakpoints + (isinstance(selected, UnsupportedCachePlan) or not selected.breakpoints) and details is not None and ((details.cached_tokens or 0) + (details.cache_creation_tokens or 0)) ): @@ -215,7 +259,8 @@ class AutoRouterBaselineCache(CustomLogger): try: counted: Final = await asyncio.wait_for( - count_cache_plan(target.model, target.api_key, plan, token_counter=count), timeout=_COUNT_TIMEOUT + count_cache_plan(target.model, target.api_key, plan, token_counter=count), + timeout=max(0.0, deadline - asyncio.get_running_loop().time()), ) return (None, counted.reason) if isinstance(counted, UnsupportedCachePlan) else (counted, None) except TimeoutError: @@ -229,111 +274,121 @@ async def invalidate_baseline_cache(logging_obj: Logging, reason: str, *, comple if context is not None: logging_obj.baseline_cache_context = replace(context, invalidated=reason) logging_obj.baseline_observation = context.capture.model_copy( - update=MappingProxyType( - { - "observation": context.capture.observation.model_copy( - update=MappingProxyType( - { - "available_at": max(context.capture.observation.started_at, context.collector.clock()), - "reason": reason, - } - ) - ), - } - ) + update={ + "observation": context.capture.observation.model_copy( + update={ + "available_at": max(context.capture.observation.started_at, context.collector.clock()), + "reason": reason, + } + ), + } ) async def finalize_baseline_cache(logging_obj: Logging, response_obj: object) -> None: context: Final = logging_obj.baseline_cache_context - if context is None: + if context is None or logging_obj.baseline_observation is not None: return + task: Final = context.finalization or asyncio.create_task(_capture(context, logging_obj, response_obj)) + active: Final = context if context.finalization is not None else replace(context, finalization=task) + if context.finalization is None: + task.add_done_callback(_consume_finalization) + logging_obj.baseline_cache_context = active try: - capture: Final = await _capture(context, logging_obj, response_obj) - if logging_obj.baseline_cache_context is context: - logging_obj.baseline_observation = capture # rebind-ok: attach only to the captured request owner - except Exception: # noqa: BLE001 # observation failures must preserve inference and billing + capture: Final = await asyncio.shield(task) + if logging_obj.baseline_cache_context is active: + logging_obj.baseline_observation = capture # rebind-ok: publish only for the current attempt + except Exception: # noqa: BLE001 # estimation must preserve inference and billing await invalidate_baseline_cache(logging_obj, "observation_unavailable") -async def _capture( +def _consume_finalization(task: asyncio.Task[CapturedBaselineObservation]) -> None: + if not task.cancelled(): + task.exception() + + +async def _capture_native( context: BaselineCacheContext, logging_obj: Logging, response_obj: object ) -> CapturedBaselineObservation: - original: Final = context.capture.observation - details: Final = _METADATA.validate_python(logging_obj.model_call_details) - if details.get("cache_hit") is True: - return context.capture.model_copy( - update=MappingProxyType( - { - "observation": original.model_copy( - update=MappingProxyType({"outcome": "response_cache", "reason": "response_cache_hit"}) - ) - } - ) - ) - event: Final = _WireEvent.model_validate(details) + capture: Final = context.capture + original: Final = capture.observation + event: Final = _WireEvent.model_validate(logging_obj.model_call_details) wire: Final = event.httpx_response.request usage: Final = _ResponseUsage.model_validate(response_obj).usage + available: Final = event.completion_start_time.timestamp() complete: Final = ( event.custom_llm_provider == "anthropic" and event.httpx_response.status_code == 200 and (not event.stream or event.prompt_cache_response_complete) ) - started: Final = original.started_at - available: Final = event.completion_start_time.timestamp() - if context.invalidated or not complete or not started <= available <= context.collector.clock(): - return context.capture.model_copy( - update=MappingProxyType( - { - "observation": original.model_copy( - update=MappingProxyType( - { - "available_at": max(started, context.collector.clock()), - "reason": context.invalidated or "incomplete_response", - } - ) - ) + if context.invalidated or not complete or not original.started_at <= available <= context.collector.clock(): + return capture.with_observation( + original.model_copy( + update={ + "available_at": max(original.started_at, context.collector.clock()), + "reason": context.invalidated or "incomplete_response", } ) ) target: Final = context.target if isinstance(target, UnsupportedPredictionTarget) or not supported_baseline_recipient(target, wire): - return context.capture.model_copy( - update=MappingProxyType( - { - "observation": original.model_copy( - update=MappingProxyType( - { - "available_at": available, - "reason": target.reason - if isinstance(target, UnsupportedPredictionTarget) - else "unsupported_baseline_recipient", - } - ) - ) + return capture.with_observation( + original.model_copy( + update={ + "available_at": available, + "reason": target.reason + if isinstance(target, UnsupportedPredictionTarget) + else "unsupported_baseline_recipient", } ) ) body: Final = _JSON_BODY.validate_json(wire.content) - same: Final = ( - logging_obj.get_router_model_id() == context.baseline_deployment_id and body.get("model") == target.model - ) - plan, reason = await context.collector.plan(target, wire, body, usage) - minimum: Final = get_prompt_cache_min_tokens(target.model) - return context.capture.model_copy( - update=MappingProxyType( - { - "observation": BaselineObservation( - request_id=original.request_id, - started_at=started, - available_at=available, - outcome="complete", - baseline_equivalent=same, - usage=usage, - plan=plan, - minimum_cache_tokens=minimum, - reason=reason, - ) - } + projected: Final = context.baseline_body + if projected is None or context.selected_body_digest != _native_body_digest(body): + return capture.with_observation( + original.model_copy( + update={ + "available_at": available, + "usage": usage, + "reason": "unsupported_baseline_settings" + if projected is None + else "unsupported_request_transformation", + } + ) + ) + same: Final = logging_obj.get_router_model_id() == context.baseline_deployment_id and _native_body_digest( + projected + ) == _native_body_digest(body) + plan, reason = await context.collector.plan(target, wire, projected, usage) + return capture.with_observation( + BaselineObservation( + request_id=original.request_id, + started_at=original.started_at, + available_at=available, + outcome="complete", + baseline_equivalent=same, + usage=usage.model_copy(update={key: projected.get(key) for key in ("speed", "inference_geo")}) + if usage is not None and not same + else usage, + plan=plan, + reason=reason, + minimum_cache_tokens=get_prompt_cache_min_tokens(target.model), ) ) + + +async def _capture( + context: BaselineCacheContext, logging_obj: Logging, response_obj: object +) -> CapturedBaselineObservation: + if _METADATA.validate_python(logging_obj.model_call_details).get("cache_hit") is True: + return context.capture.model_copy( + update={ + "observation": context.capture.observation.model_copy( + update={ + "outcome": "response_cache", + "reason": "response_cache_hit", + } + ), + } + ) + return await _capture_native(context, logging_obj, response_obj) diff --git a/litellm/proxy/lens/endpoints.py b/litellm/proxy/lens/endpoints.py index 21fde94edfe..a6381e10e39 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, HTTPException, Query, Request, Response @@ -49,8 +50,10 @@ 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.signal_repository import SignalRepository +from litellm.proxy.lens.signals import SignalConfig, TraceSignals, trace_signals from litellm.proxy.lens.sources import ActivityAvailability, SourceReader, Storage, parse_execution from litellm.proxy.lens.state import ( can_access, @@ -68,12 +71,29 @@ from litellm.proxy.lens.state import ( summarized, ) from litellm.proxy.tracing_runtime import provide_storage +from litellm.router import Router 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: @@ -84,6 +104,14 @@ def repository() -> LensRepository: return LensRepository(WriterDatabase(writer_wrapper(prisma_client.db))) +def signals_repository() -> SignalRepository: + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(503, "Lens needs a connected Postgres database") + return SignalRepository(WriterDatabase(writer_wrapper(prisma_client.db))) + + def source_reader(storage: Storage | None) -> SourceReader: if storage is None: raise HTTPException( @@ -101,6 +129,20 @@ def user_scope(auth: UserAPIKeyAuth, write: bool = False) -> Scope: raise HTTPException(403, "Lens requires proxy administrator access") +def validate_signal_model(config: SignalConfig, llm_router: Router | None) -> None: + if not config.model: + return + message: Final = "Choose a System 1 model (evaluation mode) configured on this proxy" + if llm_router is None: + raise HTTPException(400, message) + try: + model_group: Final = llm_router.get_model_group_info(model_group=config.model) + except Exception as error: + raise HTTPException(400, message) from error + if model_group is None or model_group.mode != "evaluation": + raise HTTPException(400, message) + + async def get_lens(lens_id: str, scope: Scope) -> Lens: lens: Final = await repository().get(lens_id) if lens is None or not can_access(scope, lens.scope): @@ -244,6 +286,39 @@ async def list_agents(auth: Auth, storage: StorageDep) -> tuple[str, ...]: return await source_reader(storage).agents(scope) if storage is not None else () +@router.get("/signals", response_model=SignalConfig) +async def get_signals(auth: Auth) -> SignalConfig: + user_scope(auth) + return await signals_repository().get_config() + + +@router.put("/signals", response_model=SignalConfig) +async def put_signals(body: SignalConfig, auth: Auth) -> SignalConfig: + user_scope(auth, write=True) + from litellm.proxy.proxy_server import llm_router + + validate_signal_model(body, llm_router) + await signals_repository().save_config(body) + return body + + +@router.post("/traces/signals", response_model=tuple[TraceSignals, ...]) +async def trace_signal_statuses(body: TraceFindingsRequest, auth: Auth) -> tuple[TraceSignals, ...]: + user_scope(auth) + repo: Final = signals_repository() + config: Final = await repo.get_config() + existing: Final = await repo.traces(body.traces) + rows: Final = MappingProxyType({(row.trace_id, row.trace_ref): row for row in existing}) + return tuple( + trace_signals( + trace, + rows.get((trace.trace_id, trace.trace_ref)), + config, + ) + for trace in body.traces + ) + + @router.post("/traces/findings", response_model=tuple[TraceFindingCount, ...]) async def trace_findings(body: TraceFindingsRequest, auth: Auth) -> tuple[TraceFindingCount, ...]: user_scope(auth) @@ -501,13 +576,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()) - for candidate in await repository().lenses(): - 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) @@ -540,24 +631,37 @@ async def sample(lens_id: str, job_id: str, worker: WorkerAuth, storage: Storage lens, job = await assigned(lens_id, job_id, worker) 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) @@ -752,9 +856,15 @@ async def heartbeat(lens_id: str, job_id: str, worker: WorkerAuth) -> bool: return await progress(lens_id, job_id, Progress(), worker) -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()) @@ -765,7 +875,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 ddbc1aad44a..cbb66f338fa 100644 --- a/litellm/proxy/lens/repository.py +++ b/litellm/proxy/lens/repository.py @@ -3,6 +3,7 @@ import json import random from collections.abc import AsyncGenerator, AsyncIterator, 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 @@ -24,7 +25,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: @@ -39,6 +40,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): @@ -47,6 +60,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 @@ -146,6 +179,30 @@ 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 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, + ) + 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( await self.db.query_raw( @@ -157,12 +214,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, @@ -193,7 +263,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) @@ -207,6 +278,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/signal_repository.py b/litellm/proxy/lens/signal_repository.py new file mode 100644 index 00000000000..462cbdf13f4 --- /dev/null +++ b/litellm/proxy/lens/signal_repository.py @@ -0,0 +1,142 @@ +import json +from datetime import datetime +from typing import Final + +from pydantic import TypeAdapter + +from litellm.proxy.lens.models import Execution, TraceIdentity +from litellm.proxy.lens.repository import Database, Row +from litellm.proxy.lens.signals import ( + SIGNAL_RECLASSIFY_AFTER, + SIGNAL_RETRY_FAILED_AFTER, + SignalAttempt, + SignalConfig, + StoredTraceSignal, +) + +_ROWS: Final[TypeAdapter[tuple[Row, ...]]] = TypeAdapter(tuple[Row, ...]) + + +class SignalRepository: + def __init__(self, db: Database) -> None: + self.db: Final = db + + async def get_config(self) -> SignalConfig: + rows: Final = _ROWS.validate_python( + await self.db.query_raw('SELECT data FROM "LiteLLM_LensSignalConfig" WHERE id=$1', "global") + ) + return SignalConfig() if not rows else SignalConfig.model_validate(rows[0].data) + + async def save_config(self, config: SignalConfig) -> None: + await self.db.execute_raw( + """INSERT INTO "LiteLLM_LensSignalConfig" (id, data) + VALUES ($1, $2::jsonb) + ON CONFLICT (id) DO UPDATE SET data=EXCLUDED.data""", + "global", + json.dumps(config.model_dump(mode="json")), + ) + + async def traces(self, identities: tuple[TraceIdentity, ...]) -> tuple[StoredTraceSignal, ...]: + if not identities: + return () + payload: Final = json.dumps( + tuple({"trace_id": trace.trace_id, "trace_ref": trace.trace_ref} for trace in identities) + ) + rows: Final = _ROWS.validate_python( + await self.db.query_raw( + """SELECT jsonb_build_object( + 'trace_id', trace_id, + 'trace_ref', trace_ref, + 'config_key', config_key, + 'span_count', span_count, + 'claimed_until', claimed_until, + 'classified_at', classified_at, + 'data', data + ) AS data + FROM "LiteLLM_LensTraceSignal" + WHERE (trace_id, trace_ref) IN ( + SELECT trace_id, trace_ref FROM jsonb_to_recordset($1::jsonb) AS requested( + trace_id text, trace_ref text + ) + )""", + payload, + ) + ) + return tuple(StoredTraceSignal.model_validate(row.data) for row in rows) + + async def claim( + self, + execution: Execution, + config: SignalConfig, + claimed_until: datetime, + now: datetime, + ) -> bool: + data: Final = json.dumps({"status": "pending", "scores": {}, "model": config.model, "error": ""}) + rows: Final = _ROWS.validate_python( + await self.db.query_raw( + """INSERT INTO "LiteLLM_LensTraceSignal" AS stored + (trace_id, trace_ref, config_key, span_count, claimed_until, classified_at, data) + VALUES ($1, $2, $3, $4, $5::timestamp, NULL, $6::jsonb) + ON CONFLICT (trace_id, trace_ref) DO UPDATE SET + config_key=EXCLUDED.config_key, + span_count=EXCLUDED.span_count, + claimed_until=EXCLUDED.claimed_until, + classified_at=NULL, + data=EXCLUDED.data + WHERE (stored.claimed_until IS NULL OR stored.claimed_until < $7::timestamp) + AND ( + stored.config_key IS DISTINCT FROM EXCLUDED.config_key + OR ( + stored.data->>'status'='pending' + AND stored.claimed_until < $7::timestamp + ) + OR ( + EXCLUDED.span_count > stored.span_count + AND stored.classified_at < $8::timestamp + ) + OR ( + stored.data->>'status'='failed' + AND stored.classified_at < $9::timestamp + ) + ) + RETURNING jsonb_build_object('trace_id', trace_id) AS data""", + execution.trace_id, + execution.trace_ref, + config.key(), + execution.span_count, + claimed_until, + data, + now, + now - SIGNAL_RECLASSIFY_AFTER, + now - SIGNAL_RETRY_FAILED_AFTER, + ) + ) + return bool(rows) + + async def store( + self, + execution: Execution, + config: SignalConfig, + claimed_until: datetime, + classified_at: datetime, + attempt: SignalAttempt, + ) -> None: + payload: Final = json.dumps( + { + "status": attempt.status, + "scores": dict(attempt.scores), + "model": attempt.model, + "error": attempt.error, + } + ) + await self.db.execute_raw( + """UPDATE "LiteLLM_LensTraceSignal" + SET classified_at=$1::timestamp, claimed_until=NULL, data=$2::jsonb + WHERE trace_id=$3 AND trace_ref=$4 AND config_key=$5 AND claimed_until=$6::timestamp""", + classified_at, + payload, + execution.trace_id, + execution.trace_ref, + config.key(), + claimed_until, + ) diff --git a/litellm/proxy/lens/signals.py b/litellm/proxy/lens/signals.py new file mode 100644 index 00000000000..cfe273eb6af --- /dev/null +++ b/litellm/proxy/lens/signals.py @@ -0,0 +1,586 @@ +import asyncio +import hashlib +import json +from collections.abc import Callable, Mapping +from datetime import datetime, timedelta, timezone +from itertools import accumulate +from types import MappingProxyType +from typing import Annotated, Final, Literal, Protocol, TypeAlias + +from pydantic import ConfigDict, Field, JsonValue, ValidationError, field_validator, model_validator + +from litellm.integrations.clickhouse.context import lens_analysis +from litellm.litellm_core_utils.initialize_dynamic_callback_params import inherit_message_logging_privacy +from litellm.litellm_core_utils.secret_redaction import redact_internal_details +from litellm.proxy.lens.models import ActivitySelection, Execution, Record, Scope, TraceIdentity +from litellm.proxy.lens.sources import SourceReader, Storage + +SIGNAL_INTERVAL_SECONDS: Final = 60 +SIGNAL_PAGE_SIZE: Final = 100 +SIGNAL_MAX_PER_TICK: Final = 50 +SIGNAL_CONCURRENCY: Final = 8 +SIGNAL_CLAIM_LEASE: Final = timedelta(minutes=5) +SIGNAL_RECLASSIFY_AFTER: Final = timedelta(minutes=5) +SIGNAL_RETRY_FAILED_AFTER: Final = timedelta(minutes=30) +SIGNAL_MAX_CONTENT_PAGES: Final = 3 +SIGNAL_PART_MAX_CHARS: Final = 2000 +SIGNAL_PART_HEAD_CHARS: Final = 800 +SIGNAL_PART_TAIL_CHARS: Final = 1200 +SIGNAL_TRANSCRIPT_MAX_CHARS: Final = 40000 +SIGNAL_TRANSCRIPT_HEAD_CHARS: Final = 15000 +SIGNAL_TRANSCRIPT_TAIL_CHARS: Final = 25000 +SIGNAL_MAX_SCAN_PAGES: Final = 10 +SIGNAL_TASK: Final = ( + "An AI agent run recorded as a trace. Judge only what the user and the agent said and did in these steps." +) + + +class Signal(Record): + id: str = Field(pattern=r"^[a-z][a-z0-9_]{0,63}$") + name: str = Field(min_length=1, max_length=60) + question: str = Field(min_length=3, max_length=500) + + +DEFAULT_SIGNALS: Final[tuple[Signal, ...]] = ( + Signal( + id="user_frustration", + name="User frustration", + question=( + "Does the user show frustration, annoyance or dissatisfaction with the agent in this run, for example " + "complaints, irritated corrections, all caps, profanity, or giving up on the task?" + ), + ), + Signal( + id="missing_capability", + name="Missing capability", + question=( + "Does the user ask for something the agent cannot do in this run, so that the agent refuses, says it " + "lacks a tool, permission, integration or data source, or fails because the capability does not exist?" + ), + ), + Signal( + id="repeated_request", + name="Repeated request", + question=( + "Does the user ask for the same thing more than once in this run, usually because the agent did not " + "deliver it the first time?" + ), + ), +) + + +class SignalConfig(Record): + model: str = "" + threshold: float = Field(default=0.5, ge=0.05, le=0.95, allow_inf_nan=False) + signals: tuple[Signal, ...] = DEFAULT_SIGNALS + + @model_validator(mode="after") + def validate_signals(self) -> "SignalConfig": + if len(self.signals) > 20: + raise ValueError("A maximum of 20 signals is allowed") + if len(frozenset(signal.id for signal in self.signals)) != len(self.signals): + raise ValueError("Signal IDs must be unique") + return self + + @property + def enabled(self) -> bool: + return bool(self.model) and bool(self.signals) + + def key(self) -> str: + payload: Final = json.dumps( + { + "model": self.model, + "signals": tuple({"id": signal.id, "question": signal.question} for signal in self.signals), + }, + ensure_ascii=False, + sort_keys=True, + separators=(",", ":"), + ) + return hashlib.sha256(payload.encode()).hexdigest() + + +Score: TypeAlias = Annotated[float, Field(ge=0, le=1, allow_inf_nan=False)] + + +class SignalFlag(Record): + signal_id: str + name: str + score: Score + + +class TraceSignals(TraceIdentity): + status: Literal["unclassified", "pending", "classified", "failed"] + flags: tuple[SignalFlag, ...] = () + model: str = "" + classified_at: datetime | None = None + + +class SignalStep(Record): + kind: str + name: str + content: str + + +class SignalData(Record): + model_config = ConfigDict(extra="ignore") + + status: Literal["pending", "classified", "failed"] = "pending" + scores: Mapping[str, Score] = Field(default_factory=lambda: MappingProxyType({})) + model: str = "" + error: str = "" + + +class SignalAttempt(Record): + status: Literal["classified", "failed"] + scores: Mapping[str, Score] = Field(default_factory=lambda: MappingProxyType({})) + model: str + error: str = "" + + +class StoredTraceSignal(Record): + trace_id: str + trace_ref: str = "" + config_key: str + span_count: int + claimed_until: datetime | None = None + classified_at: datetime | None = None + data: JsonValue + + @field_validator("claimed_until", "classified_at") + @classmethod + def normalize_database_timestamp(cls, value: datetime | None) -> datetime | None: + if value is not None and value.tzinfo is None: + return value.replace(tzinfo=timezone.utc) + return value + + +class NoulAnswer(Record): + model_config = ConfigDict(extra="ignore", allow_inf_nan=False, from_attributes=True) + + type: Literal["noul"] + noul: float = Field(ge=0, le=1, allow_inf_nan=False) + + +class DecisionsOutput(Record): + model_config = ConfigDict(extra="ignore", from_attributes=True) + + answers: Mapping[str, object] + + +DecisionState: TypeAlias = Mapping[str, object] +DecisionQuestions: TypeAlias = Mapping[str, Mapping[str, str]] +Clock: TypeAlias = Callable[[], datetime] +RouterReady: TypeAlias = Callable[[], bool] + + +class DecisionsCall(Protocol): + async def __call__( + self, + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: ... + + +class SignalRepositoryProtocol(Protocol): + async def get_config(self) -> SignalConfig: ... + async def traces(self, identities: tuple[TraceIdentity, ...]) -> tuple[StoredTraceSignal, ...]: ... + async def claim( + self, + execution: Execution, + config: SignalConfig, + claimed_until: datetime, + now: datetime, + ) -> bool: ... + async def store( + self, + execution: Execution, + config: SignalConfig, + claimed_until: datetime, + classified_at: datetime, + attempt: SignalAttempt, + ) -> None: ... + + +def signal_identity(trace: TraceIdentity | StoredTraceSignal | Execution) -> tuple[str, str]: + return trace.trace_id, trace.trace_ref + + +def candidate( + trace: Execution, + existing: StoredTraceSignal | None, + config_key: str, + now: datetime, +) -> bool: + if existing is None: + return True + if existing.claimed_until is not None and existing.claimed_until > now: + return False + if existing.config_key != config_key: + return True + status: Final = existing.data.get("status") if isinstance(existing.data, dict) else "" + if status == "pending": + return existing.claimed_until is not None and existing.claimed_until <= now + if existing.span_count > trace.span_count: + return False + if existing.span_count < trace.span_count: + return existing.classified_at is not None and existing.classified_at < now - SIGNAL_RECLASSIFY_AFTER + return ( + status == "failed" + and existing.classified_at is not None + and existing.classified_at < now - SIGNAL_RETRY_FAILED_AFTER + ) + + +def _take_head(steps: tuple[SignalStep, ...], remaining: int) -> tuple[SignalStep, ...]: + if remaining <= 0: + return () + cumulative_lengths: Final = tuple(accumulate(len(step.content) for step in steps)) + boundary: Final = next((index for index, total in enumerate(cumulative_lengths) if total >= remaining), None) + if boundary is None: + return steps + preceding: Final = steps[:boundary] + last: Final = steps[boundary] + used: Final = cumulative_lengths[boundary - 1] if boundary > 0 else 0 + last_length: Final = remaining - used + return ( + *preceding, + last + if last_length == len(last.content) + else last.model_copy(update=MappingProxyType({"content": last.content[:last_length]})), + ) + + +def _take_tail(steps: tuple[SignalStep, ...], remaining: int) -> tuple[SignalStep, ...]: + if remaining <= 0: + return () + reversed_steps: Final = tuple(reversed(steps)) + cumulative_lengths: Final = tuple(accumulate(len(step.content) for step in reversed_steps)) + boundary: Final = next((index for index, total in enumerate(cumulative_lengths) if total >= remaining), None) + if boundary is None: + return steps + preceding: Final = reversed_steps[:boundary] + last: Final = reversed_steps[boundary] + used: Final = cumulative_lengths[boundary - 1] if boundary > 0 else 0 + last_length: Final = remaining - used + selected: Final = ( + *preceding, + last + if last_length == len(last.content) + else last.model_copy(update=MappingProxyType({"content": last.content[-last_length:]})), + ) + return tuple(reversed(selected)) + + +def _bounded_steps(steps: tuple[SignalStep, ...]) -> tuple[SignalStep, ...]: + if sum(len(step.content) for step in steps) <= SIGNAL_TRANSCRIPT_MAX_CHARS: + return steps + head: Final = _take_head(steps, SIGNAL_TRANSCRIPT_HEAD_CHARS) + tail: Final = _take_tail(steps, SIGNAL_TRANSCRIPT_TAIL_CHARS) + omitted_count: Final = len(steps) - len(head) - len(tail) + marker: Final = SignalStep(kind="omitted", name="", content=f"{omitted_count} steps omitted") + return (*head, marker, *tail) + + +def _part_excerpt(content: str) -> str: + if len(content) <= SIGNAL_PART_MAX_CHARS: + return content + omitted: Final = len(content) - SIGNAL_PART_MAX_CHARS + marker: Final = f"\n[... {omitted} characters omitted ...]\n" + return f"{content[:SIGNAL_PART_HEAD_CHARS]}{marker}{content[-SIGNAL_PART_TAIL_CHARS:]}" + + +async def _content_pages( + reader: SourceReader, + scope: Scope, + execution: Execution, + cursor: str, + pages_left: int, +) -> tuple[SignalStep, ...]: + if pages_left == 0: + return () + content: Final = await reader.content(scope, execution, cursor) + current: Final = tuple( + SignalStep(kind=part.kind, name=part.name, content=_part_excerpt(part.content)) for part in content.parts + ) + rest: Final = ( + await _content_pages(reader, scope, execution, content.next_cursor, pages_left - 1) + if content.next_cursor is not None + else () + ) + return (*current, *rest) + + +async def signal_state(reader: SourceReader, scope: Scope, execution: Execution) -> DecisionState: + steps: Final = _bounded_steps(await _content_pages(reader, scope, execution, "", SIGNAL_MAX_CONTENT_PAGES)) + return { + "task": SIGNAL_TASK, + "steps": tuple(step.model_dump(mode="json") for step in steps), + } + + +def _noul_score(value: object) -> float | None: + try: + return NoulAnswer.model_validate(value).noul + except ValidationError: + return None + + +class SignalClassifier: + def __init__(self, reader: SourceReader, completion: DecisionsCall, clock: Clock) -> None: + self.reader: Final = reader + self.completion: Final = completion + self.clock: Final = clock + + async def classify(self, scope: Scope, execution: Execution, config: SignalConfig) -> SignalAttempt: + try: + state: Final = await signal_state(self.reader, scope, execution) + questions: Final = { + signal.id: {"type": "noul", "instructions": signal.question} for signal in config.signals + } + with lens_analysis(), inherit_message_logging_privacy(True): + response: Final = await self.completion( + model=config.model, + state=state, + questions=questions, + timeout=60, + metadata={"tags": ["litellm-lens-signals"]}, + ) + output: Final = DecisionsOutput.model_validate(response) + scores: Final = MappingProxyType( + { + signal.id: score + for signal in config.signals + if (score := _noul_score(output.answers.get(signal.id))) is not None + } + ) + if len(scores) != len(config.signals): + return SignalAttempt( + status="failed", + scores=scores, + model=config.model, + error="Decisions response omitted a configured noul answer", + ) + return SignalAttempt(status="classified", scores=scores, model=config.model) + except Exception as error: + detail: Final = redact_internal_details(str(error))[:300] + return SignalAttempt(status="failed", model=config.model, error=detail) + + +def trace_signals( + trace: TraceIdentity, + existing: StoredTraceSignal | None, + config: SignalConfig, +) -> TraceSignals: + if existing is None or existing.config_key != config.key(): + return TraceSignals(trace_id=trace.trace_id, trace_ref=trace.trace_ref, status="unclassified") + data: Final = SignalData.model_validate(existing.data) + if data.status == "pending": + return TraceSignals( + trace_id=trace.trace_id, + trace_ref=trace.trace_ref, + status="pending", + model=data.model, + ) + if data.status == "failed" or data.error: + return TraceSignals( + trace_id=trace.trace_id, + trace_ref=trace.trace_ref, + status="failed", + model=data.model, + classified_at=existing.classified_at, + ) + flags: Final = tuple( + sorted( + ( + SignalFlag(signal_id=signal.id, name=signal.name, score=data.scores[signal.id]) + for signal in config.signals + if signal.id in data.scores and data.scores[signal.id] >= config.threshold + ), + key=lambda flag: flag.score, + reverse=True, + ) + ) + return TraceSignals( + trace_id=trace.trace_id, + trace_ref=trace.trace_ref, + status="classified", + flags=flags, + model=data.model, + classified_at=existing.classified_at, + ) + + +async def _process_claimed( + classifier: SignalClassifier, + repository: SignalRepositoryProtocol, + scope: Scope, + execution: Execution, + config: SignalConfig, + claimed_until: datetime, +) -> None: + from litellm._logging import verbose_proxy_logger + + attempt: Final = await classifier.classify(scope, execution, config) + try: + await repository.store(execution, config, claimed_until, classifier.clock(), attempt) + except Exception as error: + verbose_proxy_logger.error("Lens signal result could not be stored: %s", redact_internal_details(str(error))) + + +class _SignalScan: + def __init__( + self, + reader: SourceReader, + repository: SignalRepositoryProtocol, + scope: Scope, + config: SignalConfig, + now: datetime, + cursor: str, + limit: int, + ) -> None: + self.reader: Final = reader + self.repository: Final = repository + self.scope: Final = scope + self.config: Final = config + self.now: Final = now + self.cursor: str = cursor + self.limit: Final = limit + self.executions: tuple[Execution, ...] = () + self.finished: bool = False + + async def _read_page(self, start: int, end: int) -> tuple[tuple[Execution, ...], str | None]: + page_cursor: Final = self.cursor + sample: Final = await self.reader.sample( + self.scope, + ActivitySelection(source="traces"), + start, + end, + page_size=SIGNAL_PAGE_SIZE, + cursor=page_cursor, + ) + identities: Final = tuple( + TraceIdentity(trace_id=trace.trace_id, trace_ref=trace.trace_ref) for trace in sample.executions + ) + existing_rows: Final = await self.repository.traces(identities) + existing: Final = MappingProxyType({signal_identity(row): row for row in existing_rows}) + remaining: Final = self.limit - len(self.executions) + all_eligible: Final = tuple( + execution + for execution in sample.executions + if candidate(execution, existing.get(signal_identity(execution)), self.config.key(), self.now) + ) + eligible: Final = all_eligible[:remaining] + next_cursor: Final = page_cursor if len(all_eligible) > remaining else sample.next_cursor + return eligible, next_cursor + + async def run(self) -> tuple[tuple[Execution, ...], str]: + start: Final = int((self.now - timedelta(hours=24)).timestamp() * 1000) + end: Final = int((self.now - timedelta(minutes=2)).timestamp() * 1000) + for _ in range(SIGNAL_MAX_SCAN_PAGES): + if self.finished or len(self.executions) >= self.limit: + break + eligible, next_cursor = await self._read_page(start, end) + self.executions = (*self.executions, *eligible) + if next_cursor is None: + self.cursor = "" + self.finished = True + else: + self.cursor = next_cursor + return self.executions, self.cursor + + +async def _scan_pages( + reader: SourceReader, + repository: SignalRepositoryProtocol, + scope: Scope, + config: SignalConfig, + now: datetime, + cursor: str, + remaining: int, +) -> tuple[tuple[Execution, ...], str]: + if remaining <= 0: + return (), cursor + scan: Final = _SignalScan(reader, repository, scope, config, now, cursor, remaining) + return await scan.run() + + +async def run_signal_tick( + storage: Storage, + repository: SignalRepositoryProtocol | None, + completion: DecisionsCall | None, + clock: Clock, + router_ready: RouterReady = lambda: True, + cursor: str = "", +) -> str: + if repository is None or completion is None or not router_ready(): + return cursor + now: Final = clock() + config: Final = await repository.get_config() + if not config.enabled: + return cursor + reader: Final = SourceReader(storage) + scope: Final = Scope(all_teams=True) + candidates: Final = await _scan_pages( + reader, + repository, + scope, + config, + now, + cursor, + SIGNAL_MAX_PER_TICK, + ) + executions, next_cursor = candidates + classifier: Final = SignalClassifier(reader, completion, clock) + semaphore: Final = asyncio.Semaphore(SIGNAL_CONCURRENCY) + + async def process(execution: Execution) -> None: + from litellm._logging import verbose_proxy_logger + + async with semaphore: + claimed_at: Final = classifier.clock() + claimed_until: Final = claimed_at + SIGNAL_CLAIM_LEASE + try: + claimed: Final = await repository.claim(execution, config, claimed_until, claimed_at) + except Exception as error: + verbose_proxy_logger.error("Lens signal claim failed: %s", redact_internal_details(str(error))) + return + if not claimed: + return + await _process_claimed(classifier, repository, scope, execution, config, claimed_until) + + await asyncio.gather(*(process(execution) for execution in executions)) + return next_cursor + + +class _SignalLoopState: + def __init__(self) -> None: + self.cursor: str = "" + + +async def run_signal_loop( + storage: Storage, + repository: SignalRepositoryProtocol | None, + completion: DecisionsCall | None, + clock: Clock = lambda: datetime.now(timezone.utc), + router_ready: RouterReady = lambda: True, +) -> None: + from litellm._logging import verbose_proxy_logger + + state: Final = _SignalLoopState() + while True: + try: + state.cursor = await run_signal_tick( + storage, + repository, + completion, + clock, + router_ready, + cursor=state.cursor, + ) + except Exception as error: + verbose_proxy_logger.error("Lens signal tick failed: %s", redact_internal_details(str(error))) + await asyncio.sleep(SIGNAL_INTERVAL_SECONDS) diff --git a/litellm/proxy/lens/sources.py b/litellm/proxy/lens/sources.py index 015be69ecca..b1d142a6748 100644 --- a/litellm/proxy/lens/sources.py +++ b/litellm/proxy/lens/sources.py @@ -142,6 +142,7 @@ class SourceReader: id=execution.trace_id, trace_ref=execution.trace_ref, record_team=execution.team_id, + start_time=execution.start_time, cursor=cursor, offset=offset + 1, ) @@ -175,6 +176,7 @@ class SourceReader: id=execution.trace_id, trace_ref=execution.trace_ref, record_team=execution.team_id, + start_time=execution.start_time, span=evidence.span_id, quote=evidence.quote, ) 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/proxy_server.py b/litellm/proxy/proxy_server.py index b2c5977ccc4..3d11d237665 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -581,6 +581,14 @@ from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger, run_sp from litellm.proxy.image_endpoints.endpoints import router as image_router from litellm.proxy.lens.dataset_endpoints import router as lens_dataset_router from litellm.proxy.lens.endpoints import router as lens_router +from litellm.proxy.lens.repository import WriterDatabase +from litellm.proxy.lens.signal_repository import SignalRepository +from litellm.proxy.lens.signals import ( + DecisionQuestions, + DecisionsCall, + DecisionState, + run_signal_loop, +) from litellm.proxy.list_api.common import ( ManagementProblem, problem_response, @@ -1275,6 +1283,27 @@ async def _connect_to_count_stored_values() -> SupportsRawQueries: return client.writer_db +async def _call_current_lens_signal_router( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], +) -> object: + current_router: Final = llm_router + if current_router is None: + raise RuntimeError("The proxy router is not initialized") + decisions: Final[DecisionsCall] = cast(DecisionsCall, current_router.adecisions) + return await decisions( + model=model, + state=state, + questions=questions, + timeout=timeout, + metadata=metadata, + ) + + @asynccontextmanager async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState, None]: global \ @@ -1645,12 +1674,30 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState state: Final[ProxyLifespanState] = {"tracing_receiver": receiver} from litellm.proxy.admin_mcp import admin_mcp_lifespan + signal_completion: Final[DecisionsCall] = _call_current_lens_signal_router + signal_task: Final = ( + asyncio.create_task( + run_signal_loop( + receiver.storage, + SignalRepository(WriterDatabase(writer_wrapper(prisma_client.db))), + signal_completion, + router_ready=lambda: llm_router is not None, + ) + ) + if receiver is not None and prisma_client is not None + else None + ) + try: async with AsyncExitStack() as admin_mcp_stack: try: await admin_mcp_stack.enter_async_context(admin_mcp_lifespan(app)) yield state finally: + if signal_task is not None: + signal_task.cancel() + await asyncio.gather(signal_task, return_exceptions=True) + if model_info_scheduler is not None and model_info_scheduler.running: model_info_scheduler.remove_job("refresh_model_info") if model_info_scheduler is not scheduler: @@ -13951,7 +13998,7 @@ async def run_thread( # ) # async def get_available_routes(user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth)): from litellm.llms.base_llm.base_utils import BaseTokenCounter -from litellm.proxy.db.routing_prisma_wrapper import WriterPinnedClient +from litellm.proxy.db.routing_prisma_wrapper import WriterPinnedClient, writer_wrapper from litellm.repositories.config_repository import ConfigRepository from litellm.repositories.model_repository import ModelRepository from litellm.repositories.table_repositories import ( diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 514df905866..3b83c5b09cc 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 { @@ -1974,3 +1977,20 @@ model LiteLLM_LensDataset { @@id([id, revision]) } + +model LiteLLM_LensSignalConfig { + id String @id + data Json +} + +model LiteLLM_LensTraceSignal { + trace_id String + trace_ref String @default("") + config_key String + span_count Int + claimed_until DateTime? + classified_at DateTime? + data Json + + @@id([trace_id, trace_ref]) +} diff --git a/litellm/proxy/spend_tracking/baseline_accounting.py b/litellm/proxy/spend_tracking/baseline_accounting.py index 46a38f71260..e344c0d7169 100644 --- a/litellm/proxy/spend_tracking/baseline_accounting.py +++ b/litellm/proxy/spend_tracking/baseline_accounting.py @@ -71,12 +71,12 @@ def _complete_usage(usage: Usage | None) -> bool: if usage is None or usage.prompt_tokens < 0 or usage.completion_tokens < 0: return False details: Final = usage.prompt_tokens_details - if details is None: + if details is None or not hasattr(details, "cache_creation_tokens"): return False values: Final = (details.text_tokens, details.cached_tokens, details.cache_creation_tokens) if any(value is None or value < 0 for value in values): return False - split: Final = details.cache_creation_token_details + split: Final = details.cache_creation_token_details if hasattr(details, "cache_creation_token_details") else None writes: Final = details.cache_creation_tokens or 0 return ( usage.total_tokens == usage.prompt_tokens + usage.completion_tokens @@ -133,10 +133,13 @@ def _matches(entry: CacheEntry, markers: tuple[CountedBreakpoint, ...], started: def _ambiguous(entry: CacheEntry, markers: tuple[CountedBreakpoint, ...], started: float) -> bool: - return entry.available_at <= started < entry.expires_at and any( - entry.content_fingerprint in marker.lookback_content_fingerprints - and (entry.uncertain or entry.ttl_seconds != marker.ttl_seconds) - for marker in markers + matching: Final = tuple( + marker for marker in markers if entry.content_fingerprint in marker.lookback_content_fingerprints + ) + return ( + entry.available_at <= started < entry.expires_at + and bool(matching) + and (entry.uncertain or all(entry.ttl_seconds != marker.ttl_seconds for marker in matching)) ) @@ -261,7 +264,7 @@ def _writes(history: BaselineHistory, observation: BaselineObservation) -> tuple observation.started_at + hit.ttl_seconds, ), ) - if hit is not None and all(marker.fingerprint != hit.fingerprint for marker in markers) + if hit is not None else () ) return ( @@ -277,6 +280,7 @@ def _writes(history: BaselineHistory, observation: BaselineObservation) -> tuple uncertain=bool(ambiguous), ) for marker in markers + if hit is None or marker.prefix_tokens > hit.tokens ), ) diff --git a/litellm/router.py b/litellm/router.py index e5358324868..5e52d81314b 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -14582,17 +14582,21 @@ class Router: to the deployment that actually served the request. Every attempt therefore writes or clears, never just writes. """ + from litellm.router_utils.baseline_request import capture_baseline_parameters from litellm.types.router import BaselineRouteStamp phase_attributes(routing_decision_attributes(routing_decision)) baseline_model: Final = routing_decision.get("savings_baseline_model") if routing_decision else None baseline_id: Final = routing_decision.get("savings_baseline_deployment_id") if routing_decision else None router_name: Final = routing_decision.get("router_model_name") if routing_decision else None + caller_parameters: Final = ( + capture_baseline_parameters(request_kwargs) if router_name and baseline_model else None + ) Router._stamp_or_clear_metadata_key( request_kwargs=request_kwargs, key="_autorouter_baseline_route", value=( - BaselineRouteStamp(router_name, baseline_model, baseline_id) + BaselineRouteStamp(router_name, baseline_model, baseline_id, caller_parameters) if router_name and baseline_model and baseline_id else None ), diff --git a/litellm/router_strategy/complexity_router/context_compaction.py b/litellm/router_strategy/complexity_router/context_compaction.py index d82090f3b99..6a1a63ea4b6 100644 --- a/litellm/router_strategy/complexity_router/context_compaction.py +++ b/litellm/router_strategy/complexity_router/context_compaction.py @@ -164,6 +164,11 @@ def compaction_pending(kwargs: Mapping[str, object] | None) -> bool: return isinstance(state, CompactionState) and state.config is not None and not _client_managed(kwargs or _EMPTY) +def compaction_applied(kwargs: Mapping[str, object]) -> bool: + state: Final = kwargs.get(_STATE_KEY) + return isinstance(state, CompactionState) and state.summary is not None + + def _reject(model: str, reason: str) -> NoReturn: from litellm.exceptions import BadRequestError diff --git a/litellm/router_utils/baseline_request.py b/litellm/router_utils/baseline_request.py new file mode 100644 index 00000000000..2f6262197a1 --- /dev/null +++ b/litellm/router_utils/baseline_request.py @@ -0,0 +1,153 @@ +from __future__ import annotations + +from collections.abc import Iterator, Mapping +from itertools import accumulate +from types import MappingProxyType +from typing import Final, cast + +from pydantic import JsonValue, TypeAdapter, ValidationError + +from litellm.llms.anthropic.pass_through.messages.utils import anthropic_messages_optional_param_keys + +CACHE_SETTINGS: Final = ( + "system", + "instructions", + "tools", + "tool_choice", + "parallel_tool_calls", + "response_format", + "text", + "reasoning", + "reasoning_effort", + "thinking", + "verbosity", + "output_config", + "output_format", + "speed", + "prompt_cache_key", + "cache_key", + "cached_content", + "previous_response_id", + "conversation", + "context_management", + "compaction", +) +_GENERIC_PARAMETERS: Final = ( + *CACHE_SETTINGS, + "prompt_cache_options", + "prompt_cache_retention", + "cache_control", + "max_tokens", + "max_completion_tokens", + "max_output_tokens", + "temperature", + "top_p", + "top_k", + "stop_sequences", + "enable_prompt_caching", + "cache_control_injection_points", + "drop_params", + "additional_drop_params", +) +NATIVE_ONLY_PARAMETERS: Final = tuple( + key + for key in sorted(anthropic_messages_optional_param_keys()) + if key not in (*_GENERIC_PARAMETERS, "metadata", "stream") +) +BASELINE_PARAMETERS: Final = (*_GENERIC_PARAMETERS, *NATIVE_ONLY_PARAMETERS) +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_MAX_BYTES: Final = 4 * 1024 * 1024 +_MAX_NODES: Final = 32768 +_MAX_DEPTH: Final = 32 + + +def _json_cost(value: object, depth: int = 0) -> Iterator[int]: + if depth > _MAX_DEPTH: + yield _MAX_BYTES + 1 + elif isinstance(value, str): + yield (6 if value.isascii() else 12) * len(value) + 2 + elif isinstance(value, dict): + yield 2 + for key, item in cast(dict[object, object], value).items(): + yield from _json_cost(key, depth + 1) + yield from _json_cost(item, depth + 1) + yield 2 + elif isinstance(value, (list, tuple)): + yield 2 + for item in cast(list[object] | tuple[object, ...], value): + yield from _json_cost(item, depth + 1) + yield 1 + elif isinstance(value, int) and value.bit_length() > 64: + yield _MAX_BYTES + 1 + elif value is None or isinstance(value, (bool, int, float)): + yield 32 + else: + yield _MAX_BYTES + 1 + + +def within_baseline_budget(value: object) -> bool: + return all( + size <= _MAX_BYTES and nodes <= _MAX_NODES for nodes, size in enumerate(accumulate(_json_cost(value)), 1) + ) + + +def _parameters(value: object, *, envelope: bool = False) -> dict[str, object]: + if not isinstance(value, Mapping): + return {} + mapping: Final = cast(Mapping[str, object], value) + keys: Final = (*BASELINE_PARAMETERS, "messages") if envelope else BASELINE_PARAMETERS + return {key: mapping[key] for key in keys if key in mapping} + + +def capture_baseline_parameters( + kwargs: Mapping[str, object], *, include_extra_body: bool = True +) -> Mapping[str, JsonValue] | None: + extra: Final = ( + {"extra_body": _parameters(kwargs.get("extra_body"), envelope=True)} + if include_extra_body and "extra_body" in kwargs + else {} + ) + parameters: Final = {**_parameters(kwargs), **extra} + if not within_baseline_budget(parameters): + return None + try: + return MappingProxyType(_JSON_OBJECT.validate_python(parameters)) + except ValidationError: + return None + + +def baseline_request( + kwargs: Mapping[str, object], + caller: Mapping[str, JsonValue], + deployment: Mapping[str, object], + *, + include_extra_body: bool = True, +) -> Mapping[str, object] | None: + snapshot: Final = capture_baseline_parameters(deployment) + if snapshot is None: + return None + configured: Final = { + **_parameters(snapshot), + **(_parameters(snapshot.get("extra_body")) if include_extra_body else {}), + } + requested: Final = {**_parameters(caller), **(_parameters(caller.get("extra_body")) if include_extra_body else {})} + configured_tools: Final = configured.get("tools") or [] + caller_tools: Final = requested.get("tools") or [] + merged_tools: Final = ( + {"tools": [*configured_tools, *caller_tools]} + if (configured_tools or caller_tools) and isinstance(configured_tools, list) and isinstance(caller_tools, list) + else {} + ) + return MappingProxyType( + { + **{key: value for key, value in kwargs.items() if key not in (*BASELINE_PARAMETERS, "extra_body")}, + **configured, + **requested, + **merged_tools, + **( + {"extra_body": caller.get("extra_body", snapshot.get("extra_body"))} + if not include_extra_body and ("extra_body" in caller or "extra_body" in snapshot) + else {} + ), + } + ) diff --git a/litellm/rust_bridge/trace/generated/models.py b/litellm/rust_bridge/trace/generated/models.py index ea2c8bda648..5d84003aba2 100644 --- a/litellm/rust_bridge/trace/generated/models.py +++ b/litellm/rust_bridge/trace/generated/models.py @@ -215,6 +215,7 @@ class LensContentParams(LiteLLMBaseModel): source: ContentSource id: str record_team: str + start_time: str trace_ref: str cursor: str offset: int = Field(..., ge=0, le=4294967295) @@ -232,6 +233,7 @@ class LensEvidenceParams(LiteLLMBaseModel): source: ContentSource id: str record_team: str + start_time: str trace_ref: str span: str quote: str diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 0c7695d0852..201b92f8107 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -901,6 +901,8 @@ class ContentFilterConfigModel(LiteLLMBaseModel): MCP_SECURITY_ON_VIOLATION: Final = frozenset({"block", "alert"}) +LoggingOnlyScope = Literal["input", "output", "both"] + class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch update guardrails api_key: str | None = Field(default=None, description="API key for the guardrail service") @@ -1142,6 +1144,14 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up ), ) + logging_only_scope: LoggingOnlyScope | None = Field( + default=None, + description=( + "which direction a logging_only scan observes: 'input' (request), 'output' (response), or 'both' " + "(default). Only applies to mode logging_only; pre_call/post_call on the same guardrail keep blocking." + ), + ) + @field_validator( "mode", "default_action", @@ -1310,6 +1320,7 @@ class GuardrailUIAddGuardrailSettings(LiteLLMBaseModel): supported_actions: list[str] supported_modes: list[str] supported_modes_by_provider: dict[str, list[str]] + providers_without_directional_logging_only_scope: tuple[str, ...] pii_entity_categories: list[PiiEntityCategoryMap] content_filter_settings: dict[str, object] | None = None diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index dc552a730ea..a951381a2c9 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -416,6 +416,7 @@ def _resolve_deployment_and_latency_caller_identity_labels( class PrometheusMetricLabels: litellm_llm_api_latency_metric = [ + UserAPIKeyLabelNames.MODEL_GROUP.value, UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value, UserAPIKeyLabelNames.API_KEY_HASH.value, UserAPIKeyLabelNames.API_KEY_ALIAS.value, @@ -430,6 +431,7 @@ class PrometheusMetricLabels: ] litellm_llm_api_time_to_first_token_metric = [ + UserAPIKeyLabelNames.MODEL_GROUP.value, UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value, UserAPIKeyLabelNames.API_KEY_HASH.value, UserAPIKeyLabelNames.API_KEY_ALIAS.value, @@ -444,6 +446,7 @@ class PrometheusMetricLabels: ] litellm_request_total_latency_metric = [ + UserAPIKeyLabelNames.MODEL_GROUP.value, UserAPIKeyLabelNames.END_USER.value, UserAPIKeyLabelNames.API_KEY_HASH.value, UserAPIKeyLabelNames.API_KEY_ALIAS.value, @@ -516,6 +519,7 @@ class PrometheusMetricLabels: ] litellm_deployment_latency_per_output_token = [ + UserAPIKeyLabelNames.MODEL_GROUP.value, UserAPIKeyLabelNames.v2_LITELLM_MODEL_NAME.value, UserAPIKeyLabelNames.MODEL_ID.value, UserAPIKeyLabelNames.API_BASE.value, diff --git a/litellm/types/router.py b/litellm/types/router.py index 66a5b3540f9..a66c4571b39 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -5,7 +5,7 @@ litellm.Router Types - includes RouterConfig, UpdateRouterConfig, ModelInfo etc import datetime import enum from collections.abc import Container, Mapping, Sequence -from dataclasses import dataclass +from dataclasses import dataclass, field from typing import ( TYPE_CHECKING, Annotated, @@ -21,7 +21,7 @@ from typing import ( from zoneinfo import ZoneInfo, ZoneInfoNotFoundError import httpx -from pydantic import ConfigDict, Field, field_validator, model_validator +from pydantic import ConfigDict, Field, JsonValue, field_validator, model_validator from typing_extensions import Protocol, ReadOnly, Required, TypedDict, runtime_checkable from litellm._logging import verbose_logger @@ -1223,6 +1223,7 @@ class BaselineRouteStamp: router_name: str baseline_model: str baseline_deployment_id: str + request_parameters: Mapping[str, JsonValue] | None = field(default=None, repr=False) @dataclass(frozen=True, slots=True) 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 514df905866..3b83c5b09cc 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 { @@ -1974,3 +1977,20 @@ model LiteLLM_LensDataset { @@id([id, revision]) } + +model LiteLLM_LensSignalConfig { + id String @id + data Json +} + +model LiteLLM_LensTraceSignal { + trace_id String + trace_ref String @default("") + config_key String + span_count Int + claimed_until DateTime? + classified_at DateTime? + data Json + + @@id([trace_id, trace_ref]) +} diff --git a/scripts/trace_codegen/schemas/traces-clickhouse/LensContentParams.json b/scripts/trace_codegen/schemas/traces-clickhouse/LensContentParams.json index 5ee5ab558ce..6026ccd26e1 100644 --- a/scripts/trace_codegen/schemas/traces-clickhouse/LensContentParams.json +++ b/scripts/trace_codegen/schemas/traces-clickhouse/LensContentParams.json @@ -39,6 +39,9 @@ "source": { "$ref": "#/$defs/ContentSource" }, + "start_time": { + "type": "string" + }, "team": { "type": "string" }, @@ -53,6 +56,7 @@ "source", "id", "record_team", + "start_time", "trace_ref", "cursor", "offset" diff --git a/scripts/trace_codegen/schemas/traces-clickhouse/LensEvidenceParams.json b/scripts/trace_codegen/schemas/traces-clickhouse/LensEvidenceParams.json index 07b9c216083..dbe9b32fdd6 100644 --- a/scripts/trace_codegen/schemas/traces-clickhouse/LensEvidenceParams.json +++ b/scripts/trace_codegen/schemas/traces-clickhouse/LensEvidenceParams.json @@ -36,6 +36,9 @@ "span": { "type": "string" }, + "start_time": { + "type": "string" + }, "team": { "type": "string" }, @@ -50,6 +53,7 @@ "source", "id", "record_team", + "start_time", "trace_ref", "span", "quote" diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py index 0c53137cfab..7166891bb1a 100644 --- a/tests/code_coverage_tests/recursive_detector.py +++ b/tests/code_coverage_tests/recursive_detector.py @@ -2,6 +2,7 @@ import ast import os IGNORE_FUNCTIONS = [ + "_json_cost", # bounded at depth 32 and consumed under byte/node limits. "_format_type", "remove_additional_properties", "remove_strict_from_schema", diff --git a/tests/e2e/AGENTS.md b/tests/e2e/AGENTS.md index 2e2169f8604..98a56bf1860 100644 --- a/tests/e2e/AGENTS.md +++ b/tests/e2e/AGENTS.md @@ -215,7 +215,7 @@ llm..... | rerank | images_generations | audio_speech | audio_transcriptions | moderations | realtime route : openai | azure_openai | anthropic | bedrock_converse | bedrock_invoke | vertex - | azure_foundry | cohere | together_ai + | azure_foundry | cohere | together_ai | ollama | ollama_chat (vocab varies per endpoint; messages is anthropic-format only) capability : basic | tool_use | prompt_cache_5m | vision | thinking | structured_output | service_tier | mid_conversation_system 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/batches/batch_cleanup.py b/tests/e2e/batches/batch_cleanup.py index f1142a60782..f8fc1bbf01f 100644 --- a/tests/e2e/batches/batch_cleanup.py +++ b/tests/e2e/batches/batch_cleanup.py @@ -8,6 +8,7 @@ from typing import Final, Protocol from batch_client import BatchObject, FileDeleteResponse from capabilities import is_cloud_storage_id, is_managed_id from e2e_http import NetworkError, RateLimitedError, Result, Success, UnknownApiError +from e2e_metadata import STEP_FRAMES, step from pydantic import BaseModel CLEANUP_DELAYS: Final = (1.0, 2.0, 4.0) @@ -52,6 +53,7 @@ def _require_cleanup_success[R: BaseModel](result: Result[R], operation: str) -> raise AssertionError(f"{operation} failed: {result.kind}") +@step("Clean up the uploaded file") def cleanup_file(client: BatchCleanupClient, file_id: str, *, key: str, provider: str | None = None) -> None: delete: Final[Callable[[], Result[FileDeleteResponse]]] = ( (lambda: client.delete_file_as_admin(file_id, provider=provider)) @@ -65,7 +67,7 @@ def cleanup_file(client: BatchCleanupClient, file_id: str, *, key: str, provider warnings.warn( f"Left file {file_id} in place: LiteLLM refused to delete it while a batch still references it", UserWarning, - stacklevel=2, + stacklevel=2 + STEP_FRAMES, ) return deleted: Final = _require_cleanup_success(result, f"Delete file {file_id}") @@ -74,6 +76,7 @@ def cleanup_file(client: BatchCleanupClient, file_id: str, *, key: str, provider ), f"Delete file {file_id} did not confirm deletion" +@step("Cancel the batch if it is still running") def cleanup_batch( client: BatchCleanupClient, batch_id: str, @@ -137,7 +140,7 @@ def cleanup_batch( warnings.warn( f"Left batch {batch_id} cancelling after {BATCH_CANCEL_TIMEOUT_SECONDS}s for the provider to finish", UserWarning, - stacklevel=2, + stacklevel=2 + STEP_FRAMES, ) return wait(BATCH_CANCEL_POLL_SECONDS) diff --git a/tests/e2e/batches/batch_client.py b/tests/e2e/batches/batch_client.py index 8745140a818..b02b09e557b 100644 --- a/tests/e2e/batches/batch_client.py +++ b/tests/e2e/batches/batch_client.py @@ -17,6 +17,7 @@ from typing import Final, Literal from pydantic import BaseModel, Field +from e2e_metadata import step from proxy_client import ProxyClient from e2e_http import ( FileUploadForm, @@ -136,12 +137,15 @@ def is_result_access_denied[R: BaseModel](result: Result[R]) -> bool: class BatchClient: proxy: ProxyClient + @step("Add a batch deployment named {model_name} that calls {litellm_params.model}") def create_model(self, model_name: str, litellm_params: LiteLLMParamsBody) -> str: return self.proxy.create_model(model_name, litellm_params, mode="batch") + @step("Delete the batch deployment") def delete_model(self, model_id: str) -> None: self.proxy.delete_model(model_id) + @step("Upload a batch input file to /v1/files") def upload_file( self, *, @@ -161,6 +165,7 @@ class BatchClient: response_type=FileObject, ) + @step("Retrieve the uploaded file") def retrieve_file( self, file_id: str, *, key: str, provider: str | None = None ) -> Result[FileObject]: @@ -171,6 +176,7 @@ class BatchClient: response_type=FileObject, ) + @step("List the files the key can see from /v1/files") def list_files(self, *, key: str, provider: str | None = None) -> Result[FileList]: return self.proxy.transport.get( _files_path(provider), @@ -179,6 +185,7 @@ class BatchClient: response_type=FileList, ) + @step("Create a batch of {body.endpoint} requests from the uploaded file") def create_batch( self, *, body: BatchCreateBody, key: str, provider: str | None = None ) -> StreamingResponse: @@ -188,6 +195,7 @@ class BatchClient: json=body, ) + @step("Retrieve the batch") def retrieve_batch( self, batch_id: str, *, key: str, provider: str | None = None ) -> Result[BatchObject]: @@ -198,6 +206,7 @@ class BatchClient: response_type=BatchObject, ) + @step("Cancel the batch") def cancel_batch( self, batch_id: str, *, key: str, provider: str | None = None ) -> Result[BatchObject]: @@ -208,6 +217,7 @@ class BatchClient: response_type=BatchObject, ) + @step("List the batches the key can see from /v1/batches") def list_batches( self, *, @@ -223,6 +233,7 @@ class BatchClient: response_type=BatchList, ) + @step("Delete the uploaded file") def delete_file( self, file_id: str, *, key: str, provider: str | None = None ) -> Result[FileDeleteResponse]: @@ -233,6 +244,7 @@ class BatchClient: response_type=FileDeleteResponse, ) + @step("Delete the uploaded file as the proxy admin") def delete_file_as_admin(self, file_id: str, *, provider: str | None = None) -> Result[FileDeleteResponse]: return self.proxy.transport.delete( f"{_files_path(provider)}/{file_id}", diff --git a/tests/e2e/batches/capabilities.py b/tests/e2e/batches/capabilities.py index d510426dee2..c03b481f060 100644 --- a/tests/e2e/batches/capabilities.py +++ b/tests/e2e/batches/capabilities.py @@ -7,7 +7,11 @@ import os from dataclasses import dataclass from typing import Final, Literal +import pytest + from e2e_config import provider_edge_base, unique_marker +from e2e_metadata import Domain, Mode, Route, Subject, meta +from e2e_metadata import Provider as MetaProvider from models import LiteLLMParamsBody _BATCH_RUN = unique_marker() @@ -18,6 +22,9 @@ def batch_model_name(base: str) -> str: OPENAI_BATCH_BACKEND: Final = "gpt-4o-mini" +AZURE_BATCH_BACKEND: Final = "gpt-5.4-mini-batch" +VERTEX_BATCH_BACKEND: Final = "gemini-2.5-flash" +BEDROCK_BATCH_BACKEND: Final = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0" def openai_batch_params() -> LiteLLMParamsBody: @@ -65,14 +72,14 @@ class Provider: return openai_batch_params() case "azure": return LiteLLMParamsBody( - model="azure/gpt-5.4-mini-batch", + model=f"azure/{AZURE_BATCH_BACKEND}", api_base="os.environ/AZURE_API_BASE", api_key="os.environ/AZURE_API_KEY", api_version="2025-04-01-preview", ) case "vertex_ai": return LiteLLMParamsBody( - model="vertex_ai/gemini-2.5-flash", + model=f"vertex_ai/{VERTEX_BATCH_BACKEND}", vertex_project="os.environ/VERTEXAI_PROJECT", vertex_location="us-central1", vertex_credentials="os.environ/VERTEXAI_CREDENTIALS", @@ -81,7 +88,7 @@ class Provider: ) case "bedrock": return LiteLLMParamsBody( - model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", + model=BEDROCK_BATCH_BACKEND, aws_access_key_id="os.environ/AWS_ACCESS_KEY_ID", aws_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY", aws_region_name="os.environ/AWS_REGION", @@ -132,21 +139,21 @@ PROVIDERS: tuple[Provider, ...] = ( Provider( "azure", batch_model_name("azure-batch"), - "gpt-5.4-mini-batch", + AZURE_BATCH_BACKEND, can_cancel=True, can_list=True, ), Provider( "vertex_ai", batch_model_name("vertex-batch"), - "gemini-2.5-flash", + VERTEX_BATCH_BACKEND, can_cancel=True, can_list=True, ), Provider( "bedrock", batch_model_name("bedrock-batch"), - "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", + BEDROCK_BATCH_BACKEND, can_cancel=True, can_list=True, ), @@ -181,6 +188,29 @@ CAPABILITIES: tuple[Capability, ...] = tuple( ) +def lifecycle_meta(cap: Capability) -> pytest.MarkDecorator: + return meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.BATCHES, + providers=(MetaProvider(cap.provider),), + models=(cap.raw_model,), + mode=Mode.BATCH, + ) + ) + + +def file_content_meta(provider: Provider) -> pytest.MarkDecorator: + return meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.FILES, + providers=(MetaProvider(provider.name),), + models=(provider.raw_model,), + ) + ) + + def raw_id_matches_provider(provider: str, batch_id: str) -> bool: if provider in ("openai", "azure"): return batch_id.startswith("batch") diff --git a/tests/e2e/batches/test_batches_e2e.py b/tests/e2e/batches/test_batches_e2e.py index 8da2deb4010..54bcb80ed24 100644 --- a/tests/e2e/batches/test_batches_e2e.py +++ b/tests/e2e/batches/test_batches_e2e.py @@ -44,17 +44,22 @@ from capabilities import ( OPENAI_BATCH_BACKEND, OPENAI_BATCH_MODEL, PROVIDERS, + VERTEX_BATCH_BACKEND, Capability, Provider, batch_model_name, coverage_cells_for_lifecycle, decoded_model_from_id, + file_content_meta, is_managed_id, + lifecycle_meta, matches_id_shape, openai_batch_params, raw_id_matches_provider, ) from e2e_config import MASTER_KEY, PROXY_BASE_URL, unique_marker +from e2e_metadata import Domain, Mode, Route, Subject, meta +from e2e_metadata import Provider as MetaProvider from e2e_http import ( FileUploadForm, Result, @@ -249,7 +254,7 @@ def assert_batch_object(batch: BatchObject) -> None: pytest.param( cap, id=cap.id, - marks=pytest.mark.covers(*coverage_cells_for_lifecycle(cap)), + marks=(pytest.mark.covers(*coverage_cells_for_lifecycle(cap)), lifecycle_meta(cap)), ) for cap in CAPABILITIES ], @@ -350,6 +355,15 @@ def test_batch_lifecycle( @pytest.mark.covers("llm.batches.openai.key_model_access_denied.nonstream.works") +@meta( + Subject( + domain=Domain.PROXY_AUTH, + route=Route.BATCHES, + providers=(MetaProvider.OPENAI,), + models=(OPENAI_BATCH_BACKEND,), + mode=Mode.BATCH, + ) +) def test_batch_key_model_access_denied( client: BatchClient, resources: ResourceManager, batch_deployments: None ) -> None: @@ -389,6 +403,14 @@ def test_batch_key_model_access_denied( "llm.files.openai.upload.nonstream.works", "llm.files.openai.delete.nonstream.works", ) +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.FILES, + providers=(MetaProvider.OPENAI,), + models=(OPENAI_BATCH_BACKEND,), + ) +) def test_file_upload_and_delete_outputs( client: BatchClient, resources: ResourceManager, batch_deployments: None ) -> None: @@ -433,6 +455,15 @@ def unattributed_rows(rows: list[SpendLogRow]) -> list[SpendLogRow]: "once the fetch is bounded." ) ) +@meta( + Subject( + domain=Domain.SPEND_BUDGETS, + route=Route.BATCHES, + providers=(MetaProvider.OPENAI,), + models=(OPENAI_BATCH_BACKEND,), + mode=Mode.BATCH, + ) +) def test_rate_limited_batch_create_leaves_no_unattributed_spend_row( client: BatchClient, resources: ResourceManager, batch_deployments: None ) -> None: @@ -471,7 +502,7 @@ def test_rate_limited_batch_create_leaves_no_unattributed_spend_row( file = unwrap( client.upload_file( - content=render_jsonl("gpt-4o-mini"), + content=render_jsonl(OPENAI_BATCH_BACKEND), form=FileUploadForm(purpose="batch"), model=OPENAI_BATCH_MODEL, key=key, @@ -520,6 +551,14 @@ class TestBatchFileContent: "llm.files.openai.content.nonstream.works", exercised_on=["files"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.FILES, + providers=(MetaProvider.OPENAI,), + models=(OPENAI_BATCH_BACKEND,), + ) + ) def test_file_content_matches_upload( self, client: BatchClient, resources: ResourceManager ) -> None: @@ -558,8 +597,9 @@ class TestBatchFileContent: pytest.param( p, id=p.name, - marks=pytest.mark.covers( - FILE_CONTENT_CELLS[p.name], exercised_on=["files"] + marks=( + pytest.mark.covers(FILE_CONTENT_CELLS[p.name], exercised_on=["files"]), + file_content_meta(p), ), ) for p in PROVIDERS @@ -632,6 +672,14 @@ class TestOpenAIFiles: "marker when LIT-4820 is fixed; do not relax the assertion to make it pass." ) ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.FILES, + providers=(MetaProvider.OPENAI,), + models=(OPENAI_BATCH_BACKEND,), + ) + ) def test_uploaded_file_appears_in_list( self, client: BatchClient, resources: ResourceManager, batch_deployments: None ) -> None: @@ -662,6 +710,7 @@ class TestOpenAIFiles: "llm.files.openai.list_isolation.nonstream.works", exercised_on=["files"], ) + @meta(Subject(domain=Domain.LLM_TRANSLATION, route=Route.FILES, providers=(MetaProvider.OPENAI,))) def test_list_page_cursors_address_only_the_callers_own_files( self, client: BatchClient, resources: ResourceManager ) -> None: @@ -697,6 +746,14 @@ class TestOpenAIFiles: "llm.files.openai.retrieve.nonstream.works", exercised_on=["files"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.FILES, + providers=(MetaProvider.OPENAI,), + models=(OPENAI_BATCH_BACKEND,), + ) + ) def test_retrieve_round_trips_metadata( self, client: BatchClient, resources: ResourceManager, batch_deployments: None ) -> None: @@ -760,6 +817,15 @@ class TestBatchRateLimitErrorMapping: "quota_management.ratelimit.batch_rpm.blocks_over_limit", exercised_on=["batches"], ) + @meta( + Subject( + domain=Domain.SPEND_BUDGETS, + route=Route.BATCHES, + providers=(MetaProvider.OPENAI,), + models=(OPENAI_BATCH_BACKEND,), + mode=Mode.BATCH, + ) + ) def test_batch_create_over_rpm_returns_mapped_429( self, client: BatchClient, resources: ResourceManager, batch_deployments: None ) -> None: @@ -773,7 +839,7 @@ class TestBatchRateLimitErrorMapping: file = unwrap( client.upload_file( - content=_multi_request_jsonl("gpt-4o-mini", BATCH_RL_REQUEST_LINES), + content=_multi_request_jsonl(OPENAI_BATCH_BACKEND, BATCH_RL_REQUEST_LINES), form=FileUploadForm(purpose="batch"), model=OPENAI_BATCH_MODEL, key=key, @@ -826,7 +892,7 @@ class TestBatchEnqueuedTokenLimit: ) -> FileObject: file = unwrap( client.upload_file( - content=_multi_request_jsonl("gpt-4o-mini", BATCH_RL_REQUEST_LINES), + content=_multi_request_jsonl(OPENAI_BATCH_BACKEND, BATCH_RL_REQUEST_LINES), form=FileUploadForm(purpose="batch"), model=OPENAI_BATCH_MODEL, key=key, @@ -859,6 +925,15 @@ class TestBatchEnqueuedTokenLimit: "quota_management.ratelimit.batch_enqueued_tokens.accepts_over_rpm", exercised_on=["batches"], ) + @meta( + Subject( + domain=Domain.SPEND_BUDGETS, + route=Route.BATCHES, + providers=(MetaProvider.OPENAI,), + models=(OPENAI_BATCH_BACKEND,), + mode=Mode.BATCH, + ) + ) def test_enqueued_allowance_accepts_batch_over_key_rpm( self, client: BatchClient, resources: ResourceManager, batch_deployments: None ) -> None: @@ -890,6 +965,15 @@ class TestBatchEnqueuedTokenLimit: "quota_management.ratelimit.batch_enqueued_tokens.refunds_on_cancel", exercised_on=["batches"], ) + @meta( + Subject( + domain=Domain.SPEND_BUDGETS, + route=Route.BATCHES, + providers=(MetaProvider.OPENAI,), + models=(OPENAI_BATCH_BACKEND,), + mode=Mode.BATCH, + ) + ) def test_exhausted_allowance_blocks_until_cancel_refunds( self, client: BatchClient, resources: ResourceManager, batch_deployments: None ) -> None: @@ -985,6 +1069,15 @@ class TestBedrockBatchAssumeRole: "llm.files.bedrock.upload.nonstream.works", exercised_on=["batches", "files"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.BATCHES, + providers=(MetaProvider.BEDROCK,), + models=(ASSUME_ROLE_RAW_MODEL,), + mode=Mode.BATCH, + ) + ) def test_unified_batch_create_with_assume_role( self, client: BatchClient, resources: ResourceManager ) -> None: @@ -1052,6 +1145,14 @@ class TestBedrockBatchSplitS3Credentials: "llm.files.bedrock.split_s3_credentials.nonstream.works", exercised_on=["files"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.FILES, + providers=(MetaProvider.BEDROCK,), + models=(ASSUME_ROLE_RAW_MODEL,), + ) + ) def test_file_lifecycle_signs_s3_with_s3_credentials( self, client: BatchClient, resources: ResourceManager ) -> None: @@ -1123,6 +1224,15 @@ class TestBedrockBatchGovCloud: "llm.files.bedrock.govcloud_partition.nonstream.works", exercised_on=["batches", "files"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.BATCHES, + providers=(MetaProvider.BEDROCK,), + models=(GOVCLOUD_RAW_MODEL,), + mode=Mode.BATCH, + ) + ) def test_unified_file_upload_and_batch_create_in_govcloud( self, client: BatchClient, resources: ResourceManager ) -> None: @@ -1193,6 +1303,14 @@ class TestGeminiFiles: "llm.files.gemini.upload.nonstream.works", exercised_on=["files"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.FILES, + providers=(MetaProvider.GEMINI,), + models=(GEMINI_FILES_RAW_MODEL,), + ) + ) def test_gemini_file_upload( self, client: BatchClient, resources: ResourceManager ) -> None: @@ -1227,7 +1345,7 @@ def _vllm_params(api_base: str, api_key: str | None, model_id: str) -> LiteLLMPa ) -HOSTED_VLLM_DEFAULT_MODEL = "Qwen/Qwen2.5-0.5B-Instruct" +HOSTED_VLLM_MODEL: Final = (os.environ.get("HOSTED_VLLM_MODEL") or "Qwen/Qwen2.5-0.5B-Instruct").strip() HOSTED_VLLM_BAD_LINE_CUSTOM_ID = "req-bad" @@ -1236,9 +1354,8 @@ def _hosted_vllm_deployment(client: BatchClient, resources: ResourceManager) -> if api_base is None: pytest.skip("set HOSTED_VLLM_API_BASE (the live vLLM server this deployment targets)") api_key = (os.environ.get("HOSTED_VLLM_API_KEY") or "").strip() or None - model_id = (os.environ.get("HOSTED_VLLM_MODEL") or HOSTED_VLLM_DEFAULT_MODEL).strip() proxy_name = batch_model_name("hosted-vllm-batch") - model_row_id = client.create_model(proxy_name, _vllm_params(api_base, api_key, model_id)) + model_row_id = client.create_model(proxy_name, _vllm_params(api_base, api_key, HOSTED_VLLM_MODEL)) resources.defer(lambda: client.delete_model(model_row_id)) return proxy_name @@ -1290,6 +1407,15 @@ class TestHostedVllmBatch: "llm.files.hosted_vllm.upload.nonstream.works", exercised_on=["batches", "files"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.BATCHES, + providers=(MetaProvider.HOSTED_VLLM,), + models=(HOSTED_VLLM_MODEL,), + mode=Mode.BATCH, + ) + ) def test_batch_runs_to_completion_with_a_downloadable_output( self, client: BatchClient, resources: ResourceManager, upload_route: str ) -> None: @@ -1337,6 +1463,15 @@ class TestHostedVllmBatch: ) @pytest.mark.covers("llm.batches.hosted_vllm.basic.nonstream.works", exercised_on=["batches", "files"]) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.BATCHES, + providers=(MetaProvider.HOSTED_VLLM,), + models=(HOSTED_VLLM_MODEL,), + mode=Mode.BATCH, + ) + ) def test_failing_line_lands_in_the_error_file_not_the_batch_status( self, client: BatchClient, resources: ResourceManager ) -> None: @@ -1417,6 +1552,7 @@ class TestBatchFailurePaths: "llm.batches.openai.malformed_jsonl.nonstream.works", exercised_on=["files"], ) + @meta(Subject(domain=Domain.LLM_TRANSLATION, route=Route.FILES)) def test_malformed_jsonl_upload_rejected( self, client: BatchClient, resources: ResourceManager, batch_deployments: None ) -> None: @@ -1439,13 +1575,22 @@ class TestBatchFailurePaths: "llm.batches.openai.cancel_terminal.nonstream.works", exercised_on=["batches", "files"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.BATCHES, + providers=(MetaProvider.OPENAI,), + models=(OPENAI_BATCH_BACKEND,), + mode=Mode.BATCH, + ) + ) def test_endpoint_mismatch_fails_batch_and_cancel_conflicts( self, client: BatchClient, resources: ResourceManager, batch_deployments: None ) -> None: key = resources.key() file = unwrap( client.upload_file( - content=_mismatched_endpoint_jsonl("gpt-4o-mini"), + content=_mismatched_endpoint_jsonl(OPENAI_BATCH_BACKEND), form=FileUploadForm(purpose="batch"), model=OPENAI_BATCH_MODEL, key=key, @@ -1495,6 +1640,15 @@ class TestBatchFailurePaths: "llm.batches.openai.foreign_file_id.nonstream.works", exercised_on=["batches", "files"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.BATCHES, + providers=(MetaProvider.AZURE,), + models=(AZURE_BATCH_RAW_MODEL,), + mode=Mode.BATCH, + ) + ) def test_foreign_encoded_file_id_routes_by_file_model( self, client: BatchClient, resources: ResourceManager, batch_deployments: None ) -> None: @@ -1544,6 +1698,15 @@ class TestBatchSecondHop: "llm.batches.openai.second_hop.nonstream.works", exercised_on=["batches", "files"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.BATCHES, + providers=(MetaProvider.LITELLM_PROXY, MetaProvider.OPENAI), + models=(OPENAI_BATCH_BACKEND,), + mode=Mode.BATCH, + ) + ) def test_unified_create_and_retrieve_via_chained_gateway( self, client: BatchClient, resources: ResourceManager, batch_deployments: None ) -> None: @@ -1561,7 +1724,7 @@ class TestBatchSecondHop: file = unwrap( client.upload_file( - content=render_jsonl("gpt-4o-mini"), + content=render_jsonl(OPENAI_BATCH_BACKEND), form=FileUploadForm(purpose="batch", target_model_names=hop_name), key=key, ) @@ -1680,13 +1843,22 @@ class TestBatchTerminalState: "llm.batches.openai.terminal_state.nonstream.cost_logged", exercised_on=["batches", "files"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.BATCHES, + providers=(MetaProvider.OPENAI,), + models=(OPENAI_BATCH_BACKEND,), + mode=Mode.BATCH, + ) + ) def test_completed_batch_downloads_output_and_books_cost( self, client: BatchClient, resources: ResourceManager, batch_deployments: None ) -> None: key = resources.key() file = unwrap( client.upload_file( - content=render_jsonl("gpt-4o-mini"), + content=render_jsonl(OPENAI_BATCH_BACKEND), form=FileUploadForm(purpose="batch"), model=OPENAI_BATCH_MODEL, key=key, @@ -1786,6 +1958,15 @@ class TestVertexNativePassthrough: "llm.batches.vertex.native_passthrough.nonstream.works", exercised_on=["files", "batches"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.BATCHES, + providers=(MetaProvider.VERTEX_AI,), + models=(VERTEX_BATCH_BACKEND,), + mode=Mode.BATCH, + ) + ) def test_native_jsonl_round_trips_untouched_and_starts_a_batch( self, client: BatchClient, resources: ResourceManager, batch_deployments: None ) -> None: @@ -1848,6 +2029,7 @@ class TestVertexNativePassthrough: ), ], ) + @meta(Subject(domain=Domain.LLM_TRANSLATION, route=Route.FILES)) def test_passthrough_upload_is_rejected_outside_a_native_vertex_batch( self, content: bytes, diff --git a/tests/e2e/batches/test_managed_files_enforcement_e2e.py b/tests/e2e/batches/test_managed_files_enforcement_e2e.py index 2f5d0588aca..43a2488289c 100644 --- a/tests/e2e/batches/test_managed_files_enforcement_e2e.py +++ b/tests/e2e/batches/test_managed_files_enforcement_e2e.py @@ -22,9 +22,10 @@ import pytest from batch_client import BatchClient, FileObject from batch_cleanup import cleanup_file -from capabilities import batch_model_name, is_managed_id, openai_batch_params +from capabilities import OPENAI_BATCH_BACKEND, batch_model_name, is_managed_id, openai_batch_params from e2e_config import unique_marker from e2e_http import FileUploadForm, Result, UnknownApiError, unwrap +from e2e_metadata import Domain, Provider, Route, Subject, meta from lifecycle import ResourceManager pytestmark = [pytest.mark.e2e, pytest.mark.managed_files] @@ -64,6 +65,12 @@ def managed_model(client: BatchClient) -> Iterator[str]: @pytest.mark.covers(UPLOAD_ROW) +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.FILES, + ) +) def test_upload_without_target_model_names_rejected( client: BatchClient, scoped_key: str, managed_model: str ) -> None: @@ -76,6 +83,12 @@ def test_upload_without_target_model_names_rejected( @pytest.mark.covers(UPLOAD_ROW) +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.FILES, + ) +) def test_upload_with_model_param_rejected( client: BatchClient, scoped_key: str, managed_model: str ) -> None: @@ -89,12 +102,26 @@ def test_upload_with_model_param_rejected( @pytest.mark.covers(ISOLATION_ROW) +@meta( + Subject( + domain=Domain.PROXY_AUTH, + route=Route.FILES, + ) +) def test_raw_provider_file_id_rejected(client: BatchClient, scoped_key: str) -> None: result = client.retrieve_file("file-e2e-raw-provider-id", key=scoped_key) expect_api_error(result, 400, "Raw provider file ids cannot be used") @pytest.mark.covers(ISOLATION_ROW) +@meta( + Subject( + domain=Domain.PROXY_AUTH, + route=Route.FILES, + providers=(Provider.OPENAI,), + models=(OPENAI_BATCH_BACKEND,), + ) +) def test_cross_user_managed_id_denied_owner_allowed( client: BatchClient, resources: ResourceManager, managed_model: str ) -> None: diff --git a/tests/e2e/claude_code/_basic_messaging.py b/tests/e2e/claude_code/_basic_messaging.py index 7c581cc5e38..b207bb6808f 100644 --- a/tests/e2e/claude_code/_basic_messaging.py +++ b/tests/e2e/claude_code/_basic_messaging.py @@ -31,6 +31,8 @@ from typing import Any, Callable, Mapping, Sequence import pytest +from e2e_metadata import step + from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -74,6 +76,7 @@ def _count_stream_event_deltas(events: Sequence[Mapping[str, Any]]) -> int: return count +@step("Run Claude Code headless against {models} through the proxy and check every model replies") def run_basic_messaging_cell( *, compat_result, diff --git a/tests/e2e/claude_code/_passthrough.py b/tests/e2e/claude_code/_passthrough.py index be7a475dff7..3693ce25a9c 100644 --- a/tests/e2e/claude_code/_passthrough.py +++ b/tests/e2e/claude_code/_passthrough.py @@ -58,6 +58,8 @@ from typing import Any, Callable, Dict, Mapping, Optional, Sequence import pytest +from e2e_metadata import step + from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -118,6 +120,10 @@ def foundry_extra_env(proxy_base_url: str) -> Dict[str, str]: } +@step( + "Run Claude Code headless against {models} through the proxy's native provider passthrough route" + " and check every model replies" +) def run_passthrough_cell( *, compat_result, diff --git a/tests/e2e/claude_code/basic_messaging_non_streaming/test_anthropic.py b/tests/e2e/claude_code/basic_messaging_non_streaming/test_anthropic.py index 21383b85da5..1a7e26d37c7 100644 --- a/tests/e2e/claude_code/basic_messaging_non_streaming/test_anthropic.py +++ b/tests/e2e/claude_code/basic_messaging_non_streaming/test_anthropic.py @@ -21,6 +21,7 @@ the matrix builder still sees three rows for this (feature, provider). from __future__ import annotations import pytest +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._basic_messaging import run_basic_messaging_cell # Per the PRD: each cell is exercised against three Claude tiers via the @@ -34,6 +35,15 @@ ANTHROPIC_MODELS = [ @pytest.mark.covers("llm.messages.anthropic.basic.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=tuple(ANTHROPIC_MODELS), + mode=Mode.NONSTREAM, + ) +) def test_basic_messaging_non_streaming_anthropic(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a reply. diff --git a/tests/e2e/claude_code/basic_messaging_non_streaming/test_azure.py b/tests/e2e/claude_code/basic_messaging_non_streaming/test_azure.py index 19e88dbe3cb..1fb1cc6fac7 100644 --- a/tests/e2e/claude_code/basic_messaging_non_streaming/test_azure.py +++ b/tests/e2e/claude_code/basic_messaging_non_streaming/test_azure.py @@ -26,6 +26,7 @@ the matrix builder still sees three rows for this (feature, provider). from __future__ import annotations import pytest +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._basic_messaging import run_basic_messaging_cell # Per-model aliases registered in the LiteLLM proxy's routing config to @@ -40,6 +41,15 @@ AZURE_MODELS = [ @pytest.mark.covers("llm.messages.azure_foundry.basic.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE_AI,), + models=tuple(AZURE_MODELS), + mode=Mode.NONSTREAM, + ) +) def test_basic_messaging_non_streaming_azure(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a reply. diff --git a/tests/e2e/claude_code/basic_messaging_non_streaming/test_azure_openai.py b/tests/e2e/claude_code/basic_messaging_non_streaming/test_azure_openai.py index 77876c8f7ee..2b90b4e9f02 100644 --- a/tests/e2e/claude_code/basic_messaging_non_streaming/test_azure_openai.py +++ b/tests/e2e/claude_code/basic_messaging_non_streaming/test_azure_openai.py @@ -25,6 +25,7 @@ green if all three pass. from __future__ import annotations +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._basic_messaging import run_basic_messaging_cell AZURE_OPENAI_MODELS = [ @@ -34,6 +35,15 @@ AZURE_OPENAI_MODELS = [ ] +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE,), + models=tuple(AZURE_OPENAI_MODELS), + mode=Mode.NONSTREAM, + ) +) def test_basic_messaging_non_streaming_azure_openai(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a non-empty reply from each GPT-5.6 tier.""" diff --git a/tests/e2e/claude_code/basic_messaging_non_streaming/test_bedrock_converse.py b/tests/e2e/claude_code/basic_messaging_non_streaming/test_bedrock_converse.py index 2b0f49bc205..4b02f10e6a9 100644 --- a/tests/e2e/claude_code/basic_messaging_non_streaming/test_bedrock_converse.py +++ b/tests/e2e/claude_code/basic_messaging_non_streaming/test_bedrock_converse.py @@ -21,6 +21,7 @@ the matrix builder still sees three rows for this (feature, provider). from __future__ import annotations import pytest +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._basic_messaging import run_basic_messaging_cell # Per-model aliases registered in the LiteLLM proxy's routing config to @@ -35,6 +36,15 @@ BEDROCK_CONVERSE_MODELS = [ @pytest.mark.covers("llm.messages.bedrock_converse.basic.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_CONVERSE_MODELS), + mode=Mode.NONSTREAM, + ) +) def test_basic_messaging_non_streaming_bedrock_converse(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a reply.""" run_basic_messaging_cell( diff --git a/tests/e2e/claude_code/basic_messaging_non_streaming/test_bedrock_invoke.py b/tests/e2e/claude_code/basic_messaging_non_streaming/test_bedrock_invoke.py index 937ea5ee27e..874dcebdc35 100644 --- a/tests/e2e/claude_code/basic_messaging_non_streaming/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/basic_messaging_non_streaming/test_bedrock_invoke.py @@ -21,6 +21,7 @@ the matrix builder still sees three rows for this (feature, provider). from __future__ import annotations import pytest +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._basic_messaging import run_basic_messaging_cell # Per-model aliases registered in the LiteLLM proxy's routing config to @@ -35,6 +36,15 @@ BEDROCK_INVOKE_MODELS = [ @pytest.mark.covers("llm.messages.bedrock_invoke.basic.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_INVOKE_MODELS), + mode=Mode.NONSTREAM, + ) +) def test_basic_messaging_non_streaming_bedrock_invoke(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a reply.""" run_basic_messaging_cell( diff --git a/tests/e2e/claude_code/basic_messaging_non_streaming/test_bedrock_mantle.py b/tests/e2e/claude_code/basic_messaging_non_streaming/test_bedrock_mantle.py index 8a64547a732..28e723bb450 100644 --- a/tests/e2e/claude_code/basic_messaging_non_streaming/test_bedrock_mantle.py +++ b/tests/e2e/claude_code/basic_messaging_non_streaming/test_bedrock_mantle.py @@ -25,6 +25,7 @@ COMPAT_MANTLE_CELLS=1 (see `claude_code._gpt_cells`). from __future__ import annotations +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._basic_messaging import run_basic_messaging_cell from claude_code._gpt_cells import skip_unless_mantle_cells_enabled @@ -35,6 +36,15 @@ BEDROCK_MANTLE_MODELS = [ ] +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK_MANTLE,), + models=tuple(BEDROCK_MANTLE_MODELS), + mode=Mode.NONSTREAM, + ) +) def test_basic_messaging_non_streaming_bedrock_mantle(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a non-empty reply from each GPT-5.6 tier.""" diff --git a/tests/e2e/claude_code/basic_messaging_non_streaming/test_openai.py b/tests/e2e/claude_code/basic_messaging_non_streaming/test_openai.py index 323c2f11173..27ab244c52a 100644 --- a/tests/e2e/claude_code/basic_messaging_non_streaming/test_openai.py +++ b/tests/e2e/claude_code/basic_messaging_non_streaming/test_openai.py @@ -22,6 +22,7 @@ green if all three pass. from __future__ import annotations +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._basic_messaging import run_basic_messaging_cell from claude_code._gpt_cells import skip_unless_openai_gpt_cells_enabled @@ -32,6 +33,15 @@ OPENAI_MODELS = [ ] +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.OPENAI,), + models=tuple(OPENAI_MODELS), + mode=Mode.NONSTREAM, + ) +) def test_basic_messaging_non_streaming_openai(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a non-empty reply from each GPT-5.6 tier.""" diff --git a/tests/e2e/claude_code/basic_messaging_non_streaming/test_vertex_ai.py b/tests/e2e/claude_code/basic_messaging_non_streaming/test_vertex_ai.py index c46e5a8f762..b0b41dc3f87 100644 --- a/tests/e2e/claude_code/basic_messaging_non_streaming/test_vertex_ai.py +++ b/tests/e2e/claude_code/basic_messaging_non_streaming/test_vertex_ai.py @@ -21,6 +21,7 @@ the matrix builder still sees three rows for this (feature, provider). from __future__ import annotations import pytest +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._basic_messaging import run_basic_messaging_cell # Per-model aliases registered in the LiteLLM proxy's routing config to @@ -35,6 +36,15 @@ VERTEX_AI_MODELS = [ @pytest.mark.covers("llm.messages.vertex.basic.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.VERTEX_AI,), + models=tuple(VERTEX_AI_MODELS), + mode=Mode.NONSTREAM, + ) +) def test_basic_messaging_non_streaming_vertex_ai(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a reply.""" run_basic_messaging_cell( diff --git a/tests/e2e/claude_code/basic_messaging_non_streaming/test_vertex_ai_gpt.py b/tests/e2e/claude_code/basic_messaging_non_streaming/test_vertex_ai_gpt.py index 3b155b6ac9d..7e8358a7975 100644 --- a/tests/e2e/claude_code/basic_messaging_non_streaming/test_vertex_ai_gpt.py +++ b/tests/e2e/claude_code/basic_messaging_non_streaming/test_vertex_ai_gpt.py @@ -16,9 +16,11 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations +from e2e_metadata import Domain, Subject, meta from claude_code._gpt_cells import VERTEX_AI_GPT_NOT_APPLICABLE_REASON +@meta(Subject(domain=Domain.LLM_TRANSLATION)) def test_basic_messaging_non_streaming_vertex_ai_gpt(compat_result): """Record the static not_applicable outcome for this cell.""" compat_result.set( diff --git a/tests/e2e/claude_code/basic_messaging_streaming/test_anthropic.py b/tests/e2e/claude_code/basic_messaging_streaming/test_anthropic.py index ce453f3e523..69ca0ee475a 100644 --- a/tests/e2e/claude_code/basic_messaging_streaming/test_anthropic.py +++ b/tests/e2e/claude_code/basic_messaging_streaming/test_anthropic.py @@ -26,6 +26,7 @@ sees three rows for this (feature, provider). from __future__ import annotations import pytest +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._basic_messaging import run_basic_messaging_cell ANTHROPIC_MODELS = [ @@ -36,6 +37,15 @@ ANTHROPIC_MODELS = [ @pytest.mark.covers("llm.messages.anthropic.basic.stream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=tuple(ANTHROPIC_MODELS), + mode=Mode.STREAM, + ) +) def test_basic_messaging_streaming_anthropic(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a non-empty streamed reply (one row per Claude tier). diff --git a/tests/e2e/claude_code/basic_messaging_streaming/test_azure.py b/tests/e2e/claude_code/basic_messaging_streaming/test_azure.py index 3307194e862..f9fc39c8f65 100644 --- a/tests/e2e/claude_code/basic_messaging_streaming/test_azure.py +++ b/tests/e2e/claude_code/basic_messaging_streaming/test_azure.py @@ -20,6 +20,7 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations import pytest +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._basic_messaging import run_basic_messaging_cell AZURE_MODELS = [ @@ -30,6 +31,15 @@ AZURE_MODELS = [ @pytest.mark.covers("llm.messages.azure_foundry.basic.stream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE_AI,), + models=tuple(AZURE_MODELS), + mode=Mode.STREAM, + ) +) def test_basic_messaging_streaming_azure(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a non-empty streamed reply (one row per Claude tier). diff --git a/tests/e2e/claude_code/basic_messaging_streaming/test_azure_openai.py b/tests/e2e/claude_code/basic_messaging_streaming/test_azure_openai.py index 357596590c7..d1ff8578f09 100644 --- a/tests/e2e/claude_code/basic_messaging_streaming/test_azure_openai.py +++ b/tests/e2e/claude_code/basic_messaging_streaming/test_azure_openai.py @@ -24,6 +24,7 @@ green if all three pass. from __future__ import annotations +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._basic_messaging import run_basic_messaging_cell AZURE_OPENAI_MODELS = [ @@ -33,6 +34,15 @@ AZURE_OPENAI_MODELS = [ ] +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE,), + models=tuple(AZURE_OPENAI_MODELS), + mode=Mode.STREAM, + ) +) def test_basic_messaging_streaming_azure_openai(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a non-empty streamed reply from each GPT-5.6 tier.""" diff --git a/tests/e2e/claude_code/basic_messaging_streaming/test_bedrock_converse.py b/tests/e2e/claude_code/basic_messaging_streaming/test_bedrock_converse.py index a8bc0b77a5d..3ad0b930df4 100644 --- a/tests/e2e/claude_code/basic_messaging_streaming/test_bedrock_converse.py +++ b/tests/e2e/claude_code/basic_messaging_streaming/test_bedrock_converse.py @@ -16,6 +16,7 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations import pytest +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._basic_messaging import run_basic_messaging_cell BEDROCK_CONVERSE_MODELS = [ @@ -26,6 +27,15 @@ BEDROCK_CONVERSE_MODELS = [ @pytest.mark.covers("llm.messages.bedrock_converse.basic.stream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_CONVERSE_MODELS), + mode=Mode.STREAM, + ) +) def test_basic_messaging_streaming_bedrock_converse(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a non-empty streamed reply (one row per Claude tier). diff --git a/tests/e2e/claude_code/basic_messaging_streaming/test_bedrock_invoke.py b/tests/e2e/claude_code/basic_messaging_streaming/test_bedrock_invoke.py index c0ece0e0721..1de3236d2aa 100644 --- a/tests/e2e/claude_code/basic_messaging_streaming/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/basic_messaging_streaming/test_bedrock_invoke.py @@ -16,6 +16,7 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations import pytest +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._basic_messaging import run_basic_messaging_cell BEDROCK_INVOKE_MODELS = [ @@ -26,6 +27,15 @@ BEDROCK_INVOKE_MODELS = [ @pytest.mark.covers("llm.messages.bedrock_invoke.basic.stream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_INVOKE_MODELS), + mode=Mode.STREAM, + ) +) def test_basic_messaging_streaming_bedrock_invoke(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a non-empty streamed reply (one row per Claude tier). diff --git a/tests/e2e/claude_code/basic_messaging_streaming/test_bedrock_mantle.py b/tests/e2e/claude_code/basic_messaging_streaming/test_bedrock_mantle.py index 38297e6a3e5..679731003bb 100644 --- a/tests/e2e/claude_code/basic_messaging_streaming/test_bedrock_mantle.py +++ b/tests/e2e/claude_code/basic_messaging_streaming/test_bedrock_mantle.py @@ -25,6 +25,7 @@ COMPAT_MANTLE_CELLS=1 (see `claude_code._gpt_cells`). from __future__ import annotations +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._basic_messaging import run_basic_messaging_cell from claude_code._gpt_cells import skip_unless_mantle_cells_enabled @@ -35,6 +36,15 @@ BEDROCK_MANTLE_MODELS = [ ] +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK_MANTLE,), + models=tuple(BEDROCK_MANTLE_MODELS), + mode=Mode.STREAM, + ) +) def test_basic_messaging_streaming_bedrock_mantle(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a non-empty streamed reply from each GPT-5.6 tier.""" diff --git a/tests/e2e/claude_code/basic_messaging_streaming/test_openai.py b/tests/e2e/claude_code/basic_messaging_streaming/test_openai.py index a7945fb92c0..c1617253f39 100644 --- a/tests/e2e/claude_code/basic_messaging_streaming/test_openai.py +++ b/tests/e2e/claude_code/basic_messaging_streaming/test_openai.py @@ -24,6 +24,7 @@ green if all three pass. from __future__ import annotations +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._basic_messaging import run_basic_messaging_cell from claude_code._gpt_cells import skip_unless_openai_gpt_cells_enabled @@ -34,6 +35,15 @@ OPENAI_MODELS = [ ] +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.OPENAI,), + models=tuple(OPENAI_MODELS), + mode=Mode.STREAM, + ) +) def test_basic_messaging_streaming_openai(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a non-empty streamed reply from each GPT-5.6 tier.""" diff --git a/tests/e2e/claude_code/basic_messaging_streaming/test_vertex_ai.py b/tests/e2e/claude_code/basic_messaging_streaming/test_vertex_ai.py index 13f1a0abf40..1201c71dce6 100644 --- a/tests/e2e/claude_code/basic_messaging_streaming/test_vertex_ai.py +++ b/tests/e2e/claude_code/basic_messaging_streaming/test_vertex_ai.py @@ -16,6 +16,7 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations import pytest +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._basic_messaging import run_basic_messaging_cell VERTEX_AI_MODELS = [ @@ -26,6 +27,15 @@ VERTEX_AI_MODELS = [ @pytest.mark.covers("llm.messages.vertex.basic.stream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.VERTEX_AI,), + models=tuple(VERTEX_AI_MODELS), + mode=Mode.STREAM, + ) +) def test_basic_messaging_streaming_vertex_ai(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a non-empty streamed reply (one row per Claude tier). diff --git a/tests/e2e/claude_code/basic_messaging_streaming/test_vertex_ai_gpt.py b/tests/e2e/claude_code/basic_messaging_streaming/test_vertex_ai_gpt.py index f6aa01de521..6a6cec8411a 100644 --- a/tests/e2e/claude_code/basic_messaging_streaming/test_vertex_ai_gpt.py +++ b/tests/e2e/claude_code/basic_messaging_streaming/test_vertex_ai_gpt.py @@ -16,9 +16,11 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations +from e2e_metadata import Domain, Subject, meta from claude_code._gpt_cells import VERTEX_AI_GPT_NOT_APPLICABLE_REASON +@meta(Subject(domain=Domain.LLM_TRANSLATION)) def test_basic_messaging_streaming_vertex_ai_gpt(compat_result): """Record the static not_applicable outcome for this cell.""" compat_result.set( diff --git a/tests/e2e/claude_code/cli_driver.py b/tests/e2e/claude_code/cli_driver.py index 6996849de80..6f7a7b39a39 100644 --- a/tests/e2e/claude_code/cli_driver.py +++ b/tests/e2e/claude_code/cli_driver.py @@ -25,6 +25,8 @@ from concurrent.futures import ThreadPoolExecutor, as_completed from dataclasses import dataclass, field from typing import Any, Callable, Dict, List, Mapping, Optional, Sequence, Tuple, Union +from e2e_metadata import step + from claude_code.rate_limiter import ( RateLimiter, get_default_limiter, @@ -211,6 +213,7 @@ class DriverResult: duration_ms: Optional[int] = None +@step("Run Claude Code headless against {model} through the proxy") def run_claude( *, prompt: Optional[str], @@ -395,6 +398,7 @@ def _matches_failure_shape(outcome: ModelResult, pattern: "re.Pattern[str]") -> return bool(pattern.search(failure_diagnostic(outcome))) +@step("Run Claude Code headless against {models} in parallel through the proxy") def run_claude_models_parallel( *, models: Sequence[str], diff --git a/tests/e2e/claude_code/count_tokens/test_anthropic.py b/tests/e2e/claude_code/count_tokens/test_anthropic.py index 05110d24e86..19e7bcfbc81 100644 --- a/tests/e2e/claude_code/count_tokens/test_anthropic.py +++ b/tests/e2e/claude_code/count_tokens/test_anthropic.py @@ -39,6 +39,7 @@ from __future__ import annotations import pytest +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy_client from claude_code.http_probe import ( assert_count_tokens_shape, @@ -54,6 +55,15 @@ ANTHROPIC_MODELS = [ @pytest.mark.covers("llm.messages.anthropic.count_tokens.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.COUNT_TOKENS, + providers=(Provider.ANTHROPIC,), + models=tuple(ANTHROPIC_MODELS), + mode=Mode.NONSTREAM, + ) +) def test_count_tokens_anthropic(compat_result): """Probe `/v1/messages/count_tokens` for each Anthropic tier and assert the response shape.""" diff --git a/tests/e2e/claude_code/count_tokens/test_azure.py b/tests/e2e/claude_code/count_tokens/test_azure.py index c60c623ae89..256babf3cd7 100644 --- a/tests/e2e/claude_code/count_tokens/test_azure.py +++ b/tests/e2e/claude_code/count_tokens/test_azure.py @@ -39,6 +39,7 @@ from __future__ import annotations import pytest +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy_client from claude_code.http_probe import ( assert_count_tokens_shape, @@ -54,6 +55,15 @@ AZURE_MODELS = [ @pytest.mark.covers("llm.messages.azure_foundry.count_tokens.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.COUNT_TOKENS, + providers=(Provider.AZURE_AI,), + models=tuple(AZURE_MODELS), + mode=Mode.NONSTREAM, + ) +) def test_count_tokens_azure(compat_result): """Probe `/v1/messages/count_tokens` for each Azure (Microsoft Foundry) tier and assert the response shape.""" diff --git a/tests/e2e/claude_code/count_tokens/test_bedrock_converse.py b/tests/e2e/claude_code/count_tokens/test_bedrock_converse.py index 0cb4766bb31..de5f4eae3ee 100644 --- a/tests/e2e/claude_code/count_tokens/test_bedrock_converse.py +++ b/tests/e2e/claude_code/count_tokens/test_bedrock_converse.py @@ -39,6 +39,7 @@ from __future__ import annotations import pytest +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy_client from claude_code.http_probe import ( assert_count_tokens_shape, @@ -54,6 +55,15 @@ BEDROCK_CONVERSE_MODELS = [ @pytest.mark.covers("llm.messages.bedrock_converse.count_tokens.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.COUNT_TOKENS, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_CONVERSE_MODELS), + mode=Mode.NONSTREAM, + ) +) def test_count_tokens_bedrock_converse(compat_result): """Probe `/v1/messages/count_tokens` for each Bedrock (Converse) tier and assert the response shape.""" diff --git a/tests/e2e/claude_code/count_tokens/test_bedrock_invoke.py b/tests/e2e/claude_code/count_tokens/test_bedrock_invoke.py index f1389574527..56dea8a5ad6 100644 --- a/tests/e2e/claude_code/count_tokens/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/count_tokens/test_bedrock_invoke.py @@ -39,6 +39,7 @@ from __future__ import annotations import pytest +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy_client from claude_code.http_probe import ( assert_count_tokens_shape, @@ -54,6 +55,15 @@ BEDROCK_INVOKE_MODELS = [ @pytest.mark.covers("llm.messages.bedrock_invoke.count_tokens.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.COUNT_TOKENS, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_INVOKE_MODELS), + mode=Mode.NONSTREAM, + ) +) def test_count_tokens_bedrock_invoke(compat_result): """Probe `/v1/messages/count_tokens` for each Bedrock (Invoke) tier and assert the response shape.""" diff --git a/tests/e2e/claude_code/count_tokens/test_vertex_ai.py b/tests/e2e/claude_code/count_tokens/test_vertex_ai.py index 0894214d4f0..63d12fc5dae 100644 --- a/tests/e2e/claude_code/count_tokens/test_vertex_ai.py +++ b/tests/e2e/claude_code/count_tokens/test_vertex_ai.py @@ -39,6 +39,7 @@ from __future__ import annotations import pytest +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy_client from claude_code.http_probe import ( assert_count_tokens_shape, @@ -55,6 +56,15 @@ VERTEX_AI_MODELS = [ @pytest.mark.skip(reason="stage red: Vertex returns not supported for token counting for Claude aliases") @pytest.mark.covers("llm.messages.vertex.count_tokens.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.COUNT_TOKENS, + providers=(Provider.VERTEX_AI,), + models=tuple(VERTEX_AI_MODELS), + mode=Mode.NONSTREAM, + ) +) def test_count_tokens_vertex_ai(compat_result): """Probe `/v1/messages/count_tokens` for each Vertex AI tier and assert the response shape.""" diff --git a/tests/e2e/claude_code/http_probe.py b/tests/e2e/claude_code/http_probe.py index 8aba54576c4..a6c587f98cf 100644 --- a/tests/e2e/claude_code/http_probe.py +++ b/tests/e2e/claude_code/http_probe.py @@ -42,6 +42,7 @@ from e2e_http import ( UnknownApiError, ValidationError, ) +from e2e_metadata import step from models import ( AnthropicAssistantTurn, AnthropicCustomTool, @@ -109,6 +110,7 @@ def _acquire(model: str, rate_limiter: RateLimiter | None) -> None: limiter.acquire(infer_provider(model)) +@step('Count tokens with /v1/messages/count_tokens for {model} on the message "{message}"') def probe_count_tokens( *, client: ProxyClient, @@ -132,6 +134,7 @@ def probe_count_tokens( ) +@step("Send a /v1/messages request to {model} with the tool_search tool declared") def probe_tool_search( *, client: ProxyClient, @@ -228,6 +231,10 @@ def _replay_history(answer: AnthropicMessagesResponse) -> tuple[AnthropicMessage ) +@step( + "Send a /v1/messages request to {model} with the tool_search tool declared," + " then send its answer back as history in a second request" +) def probe_tool_search_multiturn( *, client: ProxyClient, diff --git a/tests/e2e/claude_code/long_context_1m/test_anthropic.py b/tests/e2e/claude_code/long_context_1m/test_anthropic.py index 0f53e512ace..1de515c11f0 100644 --- a/tests/e2e/claude_code/long_context_1m/test_anthropic.py +++ b/tests/e2e/claude_code/long_context_1m/test_anthropic.py @@ -58,6 +58,7 @@ from typing import Sequence import pytest +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -155,6 +156,15 @@ def _build_long_prompt(target_tokens: int = TARGET_INPUT_TOKENS) -> str: @pytest.mark.skip(reason="stage red: 1M long_context not green on stage Anthropic path yet (200k sonnet / model alias)") @pytest.mark.covers("llm.messages.anthropic.long_context_1m.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=tuple(ANTHROPIC_MODELS), + mode=Mode.NONSTREAM, + ) +) def test_long_context_1m_anthropic(compat_result): """Drive the `claude` CLI with a ~210k-token prompt and the `context-1m-2025-08-07` beta header; assert no 400 / 413 and a diff --git a/tests/e2e/claude_code/long_context_1m/test_azure.py b/tests/e2e/claude_code/long_context_1m/test_azure.py index cdaa7f08178..eaad5e3de9f 100644 --- a/tests/e2e/claude_code/long_context_1m/test_azure.py +++ b/tests/e2e/claude_code/long_context_1m/test_azure.py @@ -58,6 +58,7 @@ from typing import Sequence import pytest +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -155,6 +156,15 @@ def _build_long_prompt(target_tokens: int = TARGET_INPUT_TOKENS) -> str: @pytest.mark.skip(reason="stage red: 1M long_context not green on stage Azure Foundry deployments yet") @pytest.mark.covers("llm.messages.azure_foundry.long_context_1m.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE_AI,), + models=tuple(AZURE_MODELS), + mode=Mode.NONSTREAM, + ) +) def test_long_context_1m_azure(compat_result): """Drive the `claude` CLI (Azure (Microsoft Foundry)) with a ~210k-token prompt and the `context-1m-2025-08-07` beta header; assert no 400 / 413 and a diff --git a/tests/e2e/claude_code/long_context_1m/test_bedrock_converse.py b/tests/e2e/claude_code/long_context_1m/test_bedrock_converse.py index 38aeef2ae63..eb72a36bf14 100644 --- a/tests/e2e/claude_code/long_context_1m/test_bedrock_converse.py +++ b/tests/e2e/claude_code/long_context_1m/test_bedrock_converse.py @@ -58,6 +58,7 @@ from typing import Sequence import pytest +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -155,6 +156,15 @@ def _build_long_prompt(target_tokens: int = TARGET_INPUT_TOKENS) -> str: @pytest.mark.skip(reason="stage red: 1M long_context not green on stage Bedrock Converse deployments yet") @pytest.mark.covers("llm.messages.bedrock_converse.long_context_1m.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_CONVERSE_MODELS), + mode=Mode.NONSTREAM, + ) +) def test_long_context_1m_bedrock_converse(compat_result): """Drive the `claude` CLI (Bedrock (Converse)) with a ~210k-token prompt and the `context-1m-2025-08-07` beta header; assert no 400 / 413 and a diff --git a/tests/e2e/claude_code/long_context_1m/test_bedrock_invoke.py b/tests/e2e/claude_code/long_context_1m/test_bedrock_invoke.py index f652af4aa22..dd960c15428 100644 --- a/tests/e2e/claude_code/long_context_1m/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/long_context_1m/test_bedrock_invoke.py @@ -58,6 +58,7 @@ from typing import Sequence import pytest +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -155,6 +156,15 @@ def _build_long_prompt(target_tokens: int = TARGET_INPUT_TOKENS) -> str: @pytest.mark.skip(reason="stage red: 1M long_context not green on stage Bedrock Invoke deployments yet") @pytest.mark.covers("llm.messages.bedrock_invoke.long_context_1m.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_INVOKE_MODELS), + mode=Mode.NONSTREAM, + ) +) def test_long_context_1m_bedrock_invoke(compat_result): """Drive the `claude` CLI (Bedrock (Invoke)) with a ~210k-token prompt and the `context-1m-2025-08-07` beta header; assert no 400 / 413 and a diff --git a/tests/e2e/claude_code/long_context_1m/test_vertex_ai.py b/tests/e2e/claude_code/long_context_1m/test_vertex_ai.py index 0ad68aac138..1f344d80ba8 100644 --- a/tests/e2e/claude_code/long_context_1m/test_vertex_ai.py +++ b/tests/e2e/claude_code/long_context_1m/test_vertex_ai.py @@ -58,6 +58,7 @@ from typing import Sequence import pytest +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -155,6 +156,15 @@ def _build_long_prompt(target_tokens: int = TARGET_INPUT_TOKENS) -> str: @pytest.mark.skip(reason="stage red: 1M long_context not green on stage Vertex deployments yet") @pytest.mark.covers("llm.messages.vertex.long_context_1m.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.VERTEX_AI,), + models=tuple(VERTEX_AI_MODELS), + mode=Mode.NONSTREAM, + ) +) def test_long_context_1m_vertex_ai(compat_result): """Drive the `claude` CLI (Vertex AI) with a ~210k-token prompt and the `context-1m-2025-08-07` beta header; assert no 400 / 413 and a diff --git a/tests/e2e/claude_code/passthrough/test_anthropic.py b/tests/e2e/claude_code/passthrough/test_anthropic.py index 8382342ae12..bf6ab0524d0 100644 --- a/tests/e2e/claude_code/passthrough/test_anthropic.py +++ b/tests/e2e/claude_code/passthrough/test_anthropic.py @@ -22,6 +22,7 @@ no per-provider transformation is involved. from __future__ import annotations +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._passthrough import ( ANTHROPIC_PASSTHROUGH_BASE_PATH, run_passthrough_cell, @@ -34,6 +35,15 @@ ANTHROPIC_MODELS = [ ] +@meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.ANTHROPIC,), + models=tuple(ANTHROPIC_MODELS), + mode=Mode.STREAM, + ) +) def test_passthrough_anthropic(compat_result): """Drive the `claude` CLI through `{proxy}/anthropic` and assert a reply.""" run_passthrough_cell( diff --git a/tests/e2e/claude_code/passthrough/test_azure.py b/tests/e2e/claude_code/passthrough/test_azure.py index 7365b4f50da..d2690025f08 100644 --- a/tests/e2e/claude_code/passthrough/test_azure.py +++ b/tests/e2e/claude_code/passthrough/test_azure.py @@ -42,6 +42,7 @@ from __future__ import annotations import pytest +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._passthrough import foundry_extra_env, run_passthrough_cell AZURE_MODELS = [ @@ -52,6 +53,15 @@ AZURE_MODELS = [ @pytest.mark.skip(reason="stage red: /azure passthrough drops client headers (e.g. anthropic-version); product gap") +@meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.AZURE,), + models=tuple(AZURE_MODELS), + mode=Mode.STREAM, + ) +) def test_passthrough_azure(compat_result): """Drive the `claude` CLI through `{proxy}/azure` and assert a reply.""" run_passthrough_cell( diff --git a/tests/e2e/claude_code/passthrough/test_bedrock_converse.py b/tests/e2e/claude_code/passthrough/test_bedrock_converse.py index d1093a7a958..e1605be9e49 100644 --- a/tests/e2e/claude_code/passthrough/test_bedrock_converse.py +++ b/tests/e2e/claude_code/passthrough/test_bedrock_converse.py @@ -18,7 +18,10 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations +from e2e_metadata import Domain, Subject, meta + +@meta(Subject(domain=Domain.PASSTHROUGH)) def test_passthrough_bedrock_converse(compat_result): """Report not_applicable: Claude Code has no Converse-wire mode.""" compat_result.set( diff --git a/tests/e2e/claude_code/passthrough/test_bedrock_invoke.py b/tests/e2e/claude_code/passthrough/test_bedrock_invoke.py index f1f28ab5b4c..d95118991fa 100644 --- a/tests/e2e/claude_code/passthrough/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/passthrough/test_bedrock_invoke.py @@ -23,6 +23,7 @@ cell. from __future__ import annotations +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._passthrough import bedrock_extra_env, run_passthrough_cell BEDROCK_INVOKE_MODELS = [ @@ -32,6 +33,15 @@ BEDROCK_INVOKE_MODELS = [ ] +@meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_INVOKE_MODELS), + mode=Mode.STREAM, + ) +) def test_passthrough_bedrock_invoke(compat_result): """Drive the `claude` CLI through `{proxy}/bedrock` and assert a reply.""" run_passthrough_cell( diff --git a/tests/e2e/claude_code/passthrough/test_vertex_ai.py b/tests/e2e/claude_code/passthrough/test_vertex_ai.py index 790f8b60c8f..3f84cdf02fa 100644 --- a/tests/e2e/claude_code/passthrough/test_vertex_ai.py +++ b/tests/e2e/claude_code/passthrough/test_vertex_ai.py @@ -26,6 +26,7 @@ Google and every tier fails with a 401. from __future__ import annotations +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from claude_code._passthrough import run_passthrough_cell, vertex_extra_env VERTEX_MODELS = [ @@ -35,6 +36,15 @@ VERTEX_MODELS = [ ] +@meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.VERTEX_AI,), + models=tuple(VERTEX_MODELS), + mode=Mode.STREAM, + ) +) def test_passthrough_vertex_ai(compat_result): """Drive the `claude` CLI through `{proxy}/vertex_ai` and assert a reply.""" run_passthrough_cell( diff --git a/tests/e2e/claude_code/pdf_input/test_anthropic.py b/tests/e2e/claude_code/pdf_input/test_anthropic.py index 21c8028ef1c..59a91036f26 100644 --- a/tests/e2e/claude_code/pdf_input/test_anthropic.py +++ b/tests/e2e/claude_code/pdf_input/test_anthropic.py @@ -24,6 +24,7 @@ from __future__ import annotations import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -106,6 +107,16 @@ def _build_minimal_pdf(marker: str) -> bytes: @pytest.mark.covers("llm.messages.anthropic.pdf_input.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=tuple(ANTHROPIC_MODELS), + capabilities=(Capability.PDF_INPUT,), + mode=Mode.NONSTREAM, + ) +) def test_pdf_input_anthropic(compat_result, tmp_path): """Drive the `claude` CLI against the LiteLLM proxy with a PDF attached via the Read tool and assert the reply references it.""" diff --git a/tests/e2e/claude_code/pdf_input/test_azure.py b/tests/e2e/claude_code/pdf_input/test_azure.py index 34ae3732b99..3f7642321cf 100644 --- a/tests/e2e/claude_code/pdf_input/test_azure.py +++ b/tests/e2e/claude_code/pdf_input/test_azure.py @@ -17,6 +17,7 @@ from __future__ import annotations import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -83,6 +84,16 @@ def _build_minimal_pdf(marker: str) -> bytes: @pytest.mark.covers("llm.messages.azure_foundry.pdf_input.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE_AI,), + models=tuple(AZURE_MODELS), + capabilities=(Capability.PDF_INPUT,), + mode=Mode.NONSTREAM, + ) +) def test_pdf_input_azure(compat_result, tmp_path): base_url, api_key = require_proxy(compat_result) diff --git a/tests/e2e/claude_code/pdf_input/test_bedrock_converse.py b/tests/e2e/claude_code/pdf_input/test_bedrock_converse.py index 76aa84f0f47..14020799c21 100644 --- a/tests/e2e/claude_code/pdf_input/test_bedrock_converse.py +++ b/tests/e2e/claude_code/pdf_input/test_bedrock_converse.py @@ -23,6 +23,7 @@ from __future__ import annotations import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -89,6 +90,16 @@ def _build_minimal_pdf(marker: str) -> bytes: @pytest.mark.covers("llm.messages.bedrock_converse.pdf_input.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_CONVERSE_MODELS), + capabilities=(Capability.PDF_INPUT,), + mode=Mode.NONSTREAM, + ) +) def test_pdf_input_bedrock_converse(compat_result, tmp_path): base_url, api_key = require_proxy(compat_result) diff --git a/tests/e2e/claude_code/pdf_input/test_bedrock_invoke.py b/tests/e2e/claude_code/pdf_input/test_bedrock_invoke.py index 4450266bb6b..5b3ba208421 100644 --- a/tests/e2e/claude_code/pdf_input/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/pdf_input/test_bedrock_invoke.py @@ -22,6 +22,7 @@ from __future__ import annotations import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -88,6 +89,16 @@ def _build_minimal_pdf(marker: str) -> bytes: @pytest.mark.covers("llm.messages.bedrock_invoke.pdf_input.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_INVOKE_MODELS), + capabilities=(Capability.PDF_INPUT,), + mode=Mode.NONSTREAM, + ) +) def test_pdf_input_bedrock_invoke(compat_result, tmp_path): base_url, api_key = require_proxy(compat_result) diff --git a/tests/e2e/claude_code/pdf_input/test_vertex_ai.py b/tests/e2e/claude_code/pdf_input/test_vertex_ai.py index b78f58cfda1..5f42c00efcb 100644 --- a/tests/e2e/claude_code/pdf_input/test_vertex_ai.py +++ b/tests/e2e/claude_code/pdf_input/test_vertex_ai.py @@ -17,6 +17,7 @@ from __future__ import annotations import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -83,6 +84,16 @@ def _build_minimal_pdf(marker: str) -> bytes: @pytest.mark.covers("llm.messages.vertex.pdf_input.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.VERTEX_AI,), + models=tuple(VERTEX_AI_MODELS), + capabilities=(Capability.PDF_INPUT,), + mode=Mode.NONSTREAM, + ) +) def test_pdf_input_vertex_ai(compat_result, tmp_path): base_url, api_key = require_proxy(compat_result) diff --git a/tests/e2e/claude_code/prompt_caching_1h/test_anthropic.py b/tests/e2e/claude_code/prompt_caching_1h/test_anthropic.py index 637be1c551d..60045eff8b9 100644 --- a/tests/e2e/claude_code/prompt_caching_1h/test_anthropic.py +++ b/tests/e2e/claude_code/prompt_caching_1h/test_anthropic.py @@ -28,6 +28,7 @@ from typing import Any, Mapping, Optional import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -64,6 +65,16 @@ def _cache_tokens(usage: Optional[Mapping[str, Any]]) -> int: @pytest.mark.covers("llm.messages.anthropic.prompt_cache_1h.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=tuple(ANTHROPIC_MODELS), + capabilities=(Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) +) def test_prompt_caching_1h_anthropic(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with the 1h TTL opt-in env var set, and assert the upstream usage block diff --git a/tests/e2e/claude_code/prompt_caching_1h/test_azure.py b/tests/e2e/claude_code/prompt_caching_1h/test_azure.py index f34557b3c5f..86dc4c8e50a 100644 --- a/tests/e2e/claude_code/prompt_caching_1h/test_azure.py +++ b/tests/e2e/claude_code/prompt_caching_1h/test_azure.py @@ -19,6 +19,7 @@ from typing import Any, Mapping, Optional import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -48,6 +49,16 @@ def _cache_tokens(usage: Optional[Mapping[str, Any]]) -> int: @pytest.mark.covers("llm.messages.azure_foundry.prompt_cache_1h.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE_AI,), + models=tuple(AZURE_MODELS), + capabilities=(Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) +) def test_prompt_caching_1h_azure(compat_result): base_url, api_key = require_proxy(compat_result) diff --git a/tests/e2e/claude_code/prompt_caching_1h/test_bedrock_converse.py b/tests/e2e/claude_code/prompt_caching_1h/test_bedrock_converse.py index bf62a49444c..d1106ca4350 100644 --- a/tests/e2e/claude_code/prompt_caching_1h/test_bedrock_converse.py +++ b/tests/e2e/claude_code/prompt_caching_1h/test_bedrock_converse.py @@ -24,6 +24,7 @@ from typing import Any, Mapping, Optional import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -56,6 +57,16 @@ def _cache_tokens(usage: Optional[Mapping[str, Any]]) -> int: @pytest.mark.covers("llm.messages.bedrock_converse.prompt_cache_1h.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_CONVERSE_MODELS), + capabilities=(Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) +) def test_prompt_caching_1h_bedrock_converse(compat_result): base_url, api_key = require_proxy(compat_result) diff --git a/tests/e2e/claude_code/prompt_caching_1h/test_bedrock_invoke.py b/tests/e2e/claude_code/prompt_caching_1h/test_bedrock_invoke.py index dc3468702d4..9c1166eb3b6 100644 --- a/tests/e2e/claude_code/prompt_caching_1h/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/prompt_caching_1h/test_bedrock_invoke.py @@ -26,6 +26,7 @@ from typing import Any, Mapping, Optional import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -60,6 +61,16 @@ def _cache_tokens(usage: Optional[Mapping[str, Any]]) -> int: @pytest.mark.covers("llm.messages.bedrock_invoke.prompt_cache_1h.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_INVOKE_MODELS), + capabilities=(Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) +) def test_prompt_caching_1h_bedrock_invoke(compat_result): base_url, api_key = require_proxy(compat_result) diff --git a/tests/e2e/claude_code/prompt_caching_1h/test_vertex_ai.py b/tests/e2e/claude_code/prompt_caching_1h/test_vertex_ai.py index 66cf961fcfc..89b6d7e1d4e 100644 --- a/tests/e2e/claude_code/prompt_caching_1h/test_vertex_ai.py +++ b/tests/e2e/claude_code/prompt_caching_1h/test_vertex_ai.py @@ -19,6 +19,7 @@ from typing import Any, Mapping, Optional import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -48,6 +49,16 @@ def _cache_tokens(usage: Optional[Mapping[str, Any]]) -> int: @pytest.mark.covers("llm.messages.vertex.prompt_cache_1h.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.VERTEX_AI,), + models=tuple(VERTEX_AI_MODELS), + capabilities=(Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) +) def test_prompt_caching_1h_vertex_ai(compat_result): base_url, api_key = require_proxy(compat_result) diff --git a/tests/e2e/claude_code/prompt_caching_5m/test_anthropic.py b/tests/e2e/claude_code/prompt_caching_5m/test_anthropic.py index ef551beb45c..c97df9afe2c 100644 --- a/tests/e2e/claude_code/prompt_caching_5m/test_anthropic.py +++ b/tests/e2e/claude_code/prompt_caching_5m/test_anthropic.py @@ -26,6 +26,7 @@ from typing import Any, Mapping, Optional import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -55,6 +56,16 @@ def _cache_tokens(usage: Optional[Mapping[str, Any]]) -> int: @pytest.mark.covers("llm.messages.anthropic.prompt_cache_5m.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=tuple(ANTHROPIC_MODELS), + capabilities=(Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) +) def test_prompt_caching_5m_anthropic(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert the upstream usage block surfaces a non-zero cache token count.""" diff --git a/tests/e2e/claude_code/prompt_caching_5m/test_azure.py b/tests/e2e/claude_code/prompt_caching_5m/test_azure.py index 9d4137e0726..7705af62ab4 100644 --- a/tests/e2e/claude_code/prompt_caching_5m/test_azure.py +++ b/tests/e2e/claude_code/prompt_caching_5m/test_azure.py @@ -26,6 +26,7 @@ from typing import Any, Mapping, Optional import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -53,6 +54,16 @@ def _cache_tokens(usage: Optional[Mapping[str, Any]]) -> int: @pytest.mark.covers("llm.messages.azure_foundry.prompt_cache_5m.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE_AI,), + models=tuple(AZURE_MODELS), + capabilities=(Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) +) def test_prompt_caching_5m_azure(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert the upstream usage block surfaces a non-zero cache token count.""" diff --git a/tests/e2e/claude_code/prompt_caching_5m/test_bedrock_converse.py b/tests/e2e/claude_code/prompt_caching_5m/test_bedrock_converse.py index c9b34c010b0..4a5a7dfd660 100644 --- a/tests/e2e/claude_code/prompt_caching_5m/test_bedrock_converse.py +++ b/tests/e2e/claude_code/prompt_caching_5m/test_bedrock_converse.py @@ -19,6 +19,7 @@ from typing import Any, Mapping, Optional import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -46,6 +47,16 @@ def _cache_tokens(usage: Optional[Mapping[str, Any]]) -> int: @pytest.mark.covers("llm.messages.bedrock_converse.prompt_cache_5m.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_CONVERSE_MODELS), + capabilities=(Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) +) def test_prompt_caching_5m_bedrock_converse(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert the upstream usage block surfaces a non-zero cache token count.""" diff --git a/tests/e2e/claude_code/prompt_caching_5m/test_bedrock_invoke.py b/tests/e2e/claude_code/prompt_caching_5m/test_bedrock_invoke.py index b95c509ba3c..bbe483242bd 100644 --- a/tests/e2e/claude_code/prompt_caching_5m/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/prompt_caching_5m/test_bedrock_invoke.py @@ -19,6 +19,7 @@ from typing import Any, Mapping, Optional import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -46,6 +47,16 @@ def _cache_tokens(usage: Optional[Mapping[str, Any]]) -> int: @pytest.mark.covers("llm.messages.bedrock_invoke.prompt_cache_5m.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_INVOKE_MODELS), + capabilities=(Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) +) def test_prompt_caching_5m_bedrock_invoke(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert the upstream usage block surfaces a non-zero cache token count.""" diff --git a/tests/e2e/claude_code/prompt_caching_5m/test_vertex_ai.py b/tests/e2e/claude_code/prompt_caching_5m/test_vertex_ai.py index f79377b7372..3f325efcc93 100644 --- a/tests/e2e/claude_code/prompt_caching_5m/test_vertex_ai.py +++ b/tests/e2e/claude_code/prompt_caching_5m/test_vertex_ai.py @@ -19,6 +19,7 @@ from typing import Any, Mapping, Optional import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -46,6 +47,16 @@ def _cache_tokens(usage: Optional[Mapping[str, Any]]) -> int: @pytest.mark.covers("llm.messages.vertex.prompt_cache_5m.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.VERTEX_AI,), + models=tuple(VERTEX_AI_MODELS), + capabilities=(Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) +) def test_prompt_caching_5m_vertex_ai(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert the upstream usage block surfaces a non-zero cache token count.""" diff --git a/tests/e2e/claude_code/structured_outputs/test_anthropic.py b/tests/e2e/claude_code/structured_outputs/test_anthropic.py index 3dc4c7ab8f2..c8cb7e7386c 100644 --- a/tests/e2e/claude_code/structured_outputs/test_anthropic.py +++ b/tests/e2e/claude_code/structured_outputs/test_anthropic.py @@ -53,6 +53,7 @@ from typing import Any, Mapping, Optional, Sequence, Tuple import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -151,6 +152,16 @@ def _validate_against_schema( @pytest.mark.covers("llm.messages.anthropic.structured_output.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=tuple(ANTHROPIC_MODELS), + capabilities=(Capability.RESPONSE_SCHEMA,), + mode=Mode.NONSTREAM, + ) +) def test_structured_outputs_anthropic(compat_result): """Drive `claude --json-schema ...` against the LiteLLM proxy and assert the trailing `result` event contains a schema-conforming diff --git a/tests/e2e/claude_code/structured_outputs/test_azure.py b/tests/e2e/claude_code/structured_outputs/test_azure.py index 7a776ed55ad..0708920f5e7 100644 --- a/tests/e2e/claude_code/structured_outputs/test_azure.py +++ b/tests/e2e/claude_code/structured_outputs/test_azure.py @@ -53,6 +53,7 @@ from typing import Any, Mapping, Optional, Sequence, Tuple import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -151,6 +152,16 @@ def _validate_against_schema( @pytest.mark.covers("llm.messages.azure_foundry.structured_output.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE_AI,), + models=tuple(AZURE_MODELS), + capabilities=(Capability.RESPONSE_SCHEMA,), + mode=Mode.NONSTREAM, + ) +) def test_structured_outputs_azure(compat_result): """Drive `claude --json-schema ...` against the LiteLLM proxy and assert the trailing `result` event contains a schema-conforming diff --git a/tests/e2e/claude_code/structured_outputs/test_bedrock_converse.py b/tests/e2e/claude_code/structured_outputs/test_bedrock_converse.py index 345d7c327cf..7228c194f85 100644 --- a/tests/e2e/claude_code/structured_outputs/test_bedrock_converse.py +++ b/tests/e2e/claude_code/structured_outputs/test_bedrock_converse.py @@ -53,6 +53,8 @@ from typing import Any, Mapping, Optional, Sequence, Tuple import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta + from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -151,6 +153,16 @@ def _validate_against_schema( @pytest.mark.covers("llm.messages.bedrock_converse.structured_output.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_CONVERSE_MODELS), + capabilities=(Capability.RESPONSE_SCHEMA,), + mode=Mode.NONSTREAM, + ) +) def test_structured_outputs_bedrock_converse(compat_result): """Drive `claude --json-schema ...` against the LiteLLM proxy and assert the trailing `result` event contains a schema-conforming diff --git a/tests/e2e/claude_code/structured_outputs/test_bedrock_invoke.py b/tests/e2e/claude_code/structured_outputs/test_bedrock_invoke.py index 0cf48c72d4f..c415de9e396 100644 --- a/tests/e2e/claude_code/structured_outputs/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/structured_outputs/test_bedrock_invoke.py @@ -53,6 +53,8 @@ from typing import Any, Mapping, Optional, Sequence, Tuple import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta + from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -151,6 +153,16 @@ def _validate_against_schema( @pytest.mark.covers("llm.messages.bedrock_invoke.structured_output.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_INVOKE_MODELS), + capabilities=(Capability.RESPONSE_SCHEMA,), + mode=Mode.NONSTREAM, + ) +) def test_structured_outputs_bedrock_invoke(compat_result): """Drive `claude --json-schema ...` against the LiteLLM proxy and assert the trailing `result` event contains a schema-conforming diff --git a/tests/e2e/claude_code/structured_outputs/test_vertex_ai.py b/tests/e2e/claude_code/structured_outputs/test_vertex_ai.py index 24f5a0c35d4..a4cf669c19f 100644 --- a/tests/e2e/claude_code/structured_outputs/test_vertex_ai.py +++ b/tests/e2e/claude_code/structured_outputs/test_vertex_ai.py @@ -53,6 +53,8 @@ from typing import Any, Mapping, Optional, Sequence, Tuple import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta + from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -151,6 +153,16 @@ def _validate_against_schema( @pytest.mark.covers("llm.messages.vertex.structured_output.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.VERTEX_AI,), + models=tuple(VERTEX_AI_MODELS), + capabilities=(Capability.RESPONSE_SCHEMA,), + mode=Mode.NONSTREAM, + ) +) def test_structured_outputs_vertex_ai(compat_result): """Drive `claude --json-schema ...` against the LiteLLM proxy and assert the trailing `result` event contains a schema-conforming diff --git a/tests/e2e/claude_code/thinking/test_anthropic.py b/tests/e2e/claude_code/thinking/test_anthropic.py index ebb2445fb6d..2c3a8f0a8a9 100644 --- a/tests/e2e/claude_code/thinking/test_anthropic.py +++ b/tests/e2e/claude_code/thinking/test_anthropic.py @@ -24,6 +24,8 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta + from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -75,6 +77,16 @@ def _has_thinking_block(events: Sequence[Mapping[str, Any]]) -> bool: @pytest.mark.covers("llm.messages.anthropic.thinking.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=tuple(ANTHROPIC_MODELS), + capabilities=(Capability.REASONING,), + mode=Mode.NONSTREAM, + ) +) def test_thinking_anthropic(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with thinking enabled and assert a `thinking` content block was emitted.""" diff --git a/tests/e2e/claude_code/thinking/test_azure.py b/tests/e2e/claude_code/thinking/test_azure.py index ffd5ca92df0..0828dd033f4 100644 --- a/tests/e2e/claude_code/thinking/test_azure.py +++ b/tests/e2e/claude_code/thinking/test_azure.py @@ -27,6 +27,8 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta + from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -63,6 +65,16 @@ def _has_thinking_block(events: Sequence[Mapping[str, Any]]) -> bool: @pytest.mark.covers("llm.messages.azure_foundry.thinking.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE_AI,), + models=tuple(AZURE_MODELS), + capabilities=(Capability.REASONING,), + mode=Mode.NONSTREAM, + ) +) def test_thinking_azure(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with thinking enabled and assert a `thinking` content block was emitted.""" diff --git a/tests/e2e/claude_code/thinking/test_bedrock_converse.py b/tests/e2e/claude_code/thinking/test_bedrock_converse.py index 0b409f18ea7..fd075298234 100644 --- a/tests/e2e/claude_code/thinking/test_bedrock_converse.py +++ b/tests/e2e/claude_code/thinking/test_bedrock_converse.py @@ -19,6 +19,8 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta + from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -55,6 +57,16 @@ def _has_thinking_block(events: Sequence[Mapping[str, Any]]) -> bool: @pytest.mark.covers("llm.messages.bedrock_converse.thinking.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_CONVERSE_MODELS), + capabilities=(Capability.REASONING,), + mode=Mode.NONSTREAM, + ) +) def test_thinking_bedrock_converse(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with thinking enabled and assert a `thinking` content block was emitted.""" diff --git a/tests/e2e/claude_code/thinking/test_bedrock_invoke.py b/tests/e2e/claude_code/thinking/test_bedrock_invoke.py index a2c97eae321..c115ba9e408 100644 --- a/tests/e2e/claude_code/thinking/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/thinking/test_bedrock_invoke.py @@ -19,6 +19,8 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta + from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -55,6 +57,16 @@ def _has_thinking_block(events: Sequence[Mapping[str, Any]]) -> bool: @pytest.mark.covers("llm.messages.bedrock_invoke.thinking.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_INVOKE_MODELS), + capabilities=(Capability.REASONING,), + mode=Mode.NONSTREAM, + ) +) def test_thinking_bedrock_invoke(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with thinking enabled and assert a `thinking` content block was emitted.""" diff --git a/tests/e2e/claude_code/thinking/test_vertex_ai.py b/tests/e2e/claude_code/thinking/test_vertex_ai.py index f1a1c5b6cee..b947f97e7ac 100644 --- a/tests/e2e/claude_code/thinking/test_vertex_ai.py +++ b/tests/e2e/claude_code/thinking/test_vertex_ai.py @@ -19,6 +19,8 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta + from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -55,6 +57,16 @@ def _has_thinking_block(events: Sequence[Mapping[str, Any]]) -> bool: @pytest.mark.covers("llm.messages.vertex.thinking.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.VERTEX_AI,), + models=tuple(VERTEX_AI_MODELS), + capabilities=(Capability.REASONING,), + mode=Mode.NONSTREAM, + ) +) def test_thinking_vertex_ai(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with thinking enabled and assert a `thinking` content block was emitted.""" diff --git a/tests/e2e/claude_code/thinking_with_tool_use/test_anthropic.py b/tests/e2e/claude_code/thinking_with_tool_use/test_anthropic.py index 7e39ea26d42..8ddbdba04f9 100644 --- a/tests/e2e/claude_code/thinking_with_tool_use/test_anthropic.py +++ b/tests/e2e/claude_code/thinking_with_tool_use/test_anthropic.py @@ -28,6 +28,8 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta + from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -89,6 +91,16 @@ def _has_block_type( @pytest.mark.covers("llm.messages.anthropic.thinking_with_tool_use.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=tuple(ANTHROPIC_MODELS), + capabilities=(Capability.REASONING, Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) +) def test_thinking_with_tool_use_anthropic(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with thinking enabled and tool use, and assert both `thinking` and `tool_use` diff --git a/tests/e2e/claude_code/thinking_with_tool_use/test_azure.py b/tests/e2e/claude_code/thinking_with_tool_use/test_azure.py index 0371a10f8a6..6dbd6b27fde 100644 --- a/tests/e2e/claude_code/thinking_with_tool_use/test_azure.py +++ b/tests/e2e/claude_code/thinking_with_tool_use/test_azure.py @@ -22,6 +22,8 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta + from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -70,6 +72,16 @@ def _has_block_type( @pytest.mark.covers("llm.messages.azure_foundry.thinking_with_tool_use.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE_AI,), + models=tuple(AZURE_MODELS), + capabilities=(Capability.REASONING, Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) +) def test_thinking_with_tool_use_azure(compat_result): base_url, api_key = require_proxy(compat_result) diff --git a/tests/e2e/claude_code/thinking_with_tool_use/test_bedrock_converse.py b/tests/e2e/claude_code/thinking_with_tool_use/test_bedrock_converse.py index 026d2a3707f..b3869f157b5 100644 --- a/tests/e2e/claude_code/thinking_with_tool_use/test_bedrock_converse.py +++ b/tests/e2e/claude_code/thinking_with_tool_use/test_bedrock_converse.py @@ -27,6 +27,8 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta + from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -75,6 +77,16 @@ def _has_block_type( @pytest.mark.covers("llm.messages.bedrock_converse.thinking_with_tool_use.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_CONVERSE_MODELS), + capabilities=(Capability.REASONING, Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) +) def test_thinking_with_tool_use_bedrock_converse(compat_result): base_url, api_key = require_proxy(compat_result) diff --git a/tests/e2e/claude_code/thinking_with_tool_use/test_bedrock_invoke.py b/tests/e2e/claude_code/thinking_with_tool_use/test_bedrock_invoke.py index 1dd4cf0a73c..384ce0bb3c6 100644 --- a/tests/e2e/claude_code/thinking_with_tool_use/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/thinking_with_tool_use/test_bedrock_invoke.py @@ -29,6 +29,8 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta + from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -77,6 +79,16 @@ def _has_block_type( @pytest.mark.covers("llm.messages.bedrock_invoke.thinking_with_tool_use.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_INVOKE_MODELS), + capabilities=(Capability.REASONING, Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) +) def test_thinking_with_tool_use_bedrock_invoke(compat_result): base_url, api_key = require_proxy(compat_result) diff --git a/tests/e2e/claude_code/thinking_with_tool_use/test_vertex_ai.py b/tests/e2e/claude_code/thinking_with_tool_use/test_vertex_ai.py index b25228edb55..c3a3dd58533 100644 --- a/tests/e2e/claude_code/thinking_with_tool_use/test_vertex_ai.py +++ b/tests/e2e/claude_code/thinking_with_tool_use/test_vertex_ai.py @@ -27,6 +27,8 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta + from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -75,6 +77,16 @@ def _has_block_type( @pytest.mark.covers("llm.messages.vertex.thinking_with_tool_use.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.VERTEX_AI,), + models=tuple(VERTEX_AI_MODELS), + capabilities=(Capability.REASONING, Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) +) def test_thinking_with_tool_use_vertex_ai(compat_result): base_url, api_key = require_proxy(compat_result) diff --git a/tests/e2e/claude_code/tool_search/test_anthropic.py b/tests/e2e/claude_code/tool_search/test_anthropic.py index a23d1b3d2bb..f08c30b507e 100644 --- a/tests/e2e/claude_code/tool_search/test_anthropic.py +++ b/tests/e2e/claude_code/tool_search/test_anthropic.py @@ -45,6 +45,8 @@ from __future__ import annotations import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta + from claude_code._env import require_proxy_client from claude_code.http_probe import ( assert_tool_search_shape, @@ -60,6 +62,16 @@ ANTHROPIC_MODELS = [ @pytest.mark.covers("llm.messages.anthropic.tool_search.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=tuple(ANTHROPIC_MODELS), + capabilities=(Capability.TOOL_SEARCH,), + mode=Mode.NONSTREAM, + ) +) def test_tool_search_anthropic(compat_result): """Probe `/v1/messages` with a `tool_search_tool_regex_20251119` tool and assert the proxy + upstream accept it for every Anthropic diff --git a/tests/e2e/claude_code/tool_search/test_azure.py b/tests/e2e/claude_code/tool_search/test_azure.py index b094a35ea63..87628bc02d5 100644 --- a/tests/e2e/claude_code/tool_search/test_azure.py +++ b/tests/e2e/claude_code/tool_search/test_azure.py @@ -45,6 +45,8 @@ from __future__ import annotations import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta + from claude_code._env import require_proxy_client from claude_code.http_probe import ( assert_tool_search_shape, @@ -61,6 +63,16 @@ AZURE_MODELS = [ @pytest.mark.skip(reason="stage red: Azure Foundry tool_search_server not supported in workspace for probed models") @pytest.mark.covers("llm.messages.azure_foundry.tool_search.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE_AI,), + models=tuple(AZURE_MODELS), + capabilities=(Capability.TOOL_SEARCH,), + mode=Mode.NONSTREAM, + ) +) def test_tool_search_azure(compat_result): """Probe `/v1/messages` with a `tool_search_tool_regex_20251119` tool and assert the proxy + upstream accept it for every Azure (Microsoft Foundry) diff --git a/tests/e2e/claude_code/tool_search/test_bedrock_converse.py b/tests/e2e/claude_code/tool_search/test_bedrock_converse.py index f395122a5ab..84fd47bcb39 100644 --- a/tests/e2e/claude_code/tool_search/test_bedrock_converse.py +++ b/tests/e2e/claude_code/tool_search/test_bedrock_converse.py @@ -45,6 +45,8 @@ from __future__ import annotations import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta + from claude_code._env import require_proxy_client from claude_code.http_probe import ( assert_tool_search_shape, @@ -60,6 +62,16 @@ BEDROCK_CONVERSE_MODELS = [ @pytest.mark.covers("llm.messages.bedrock_converse.tool_search.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_CONVERSE_MODELS), + capabilities=(Capability.TOOL_SEARCH,), + mode=Mode.NONSTREAM, + ) +) def test_tool_search_bedrock_converse(compat_result): """Probe `/v1/messages` with a `tool_search_tool_regex_20251119` tool and assert the proxy + upstream accept it for every Bedrock (Converse) diff --git a/tests/e2e/claude_code/tool_search/test_bedrock_invoke.py b/tests/e2e/claude_code/tool_search/test_bedrock_invoke.py index 5b4c50e9dc5..b4a7f721a92 100644 --- a/tests/e2e/claude_code/tool_search/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/tool_search/test_bedrock_invoke.py @@ -50,6 +50,8 @@ from __future__ import annotations import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta + from claude_code._env import require_proxy_client from claude_code.http_probe import ( assert_tool_search_replay_shape, @@ -67,6 +69,16 @@ BEDROCK_INVOKE_MODELS = [ @pytest.mark.covers("llm.messages.bedrock_invoke.tool_search.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_INVOKE_MODELS), + capabilities=(Capability.TOOL_SEARCH,), + mode=Mode.NONSTREAM, + ) +) def test_tool_search_bedrock_invoke(compat_result): """Probe `/v1/messages` with a `tool_search_tool_regex_20251119` tool and assert the proxy + upstream accept it for every Bedrock (Invoke) @@ -90,6 +102,16 @@ def test_tool_search_bedrock_invoke(compat_result): @pytest.mark.covers("llm.messages.bedrock_invoke.tool_search_history.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_INVOKE_MODELS), + capabilities=(Capability.TOOL_SEARCH,), + mode=Mode.NONSTREAM, + ) +) def test_tool_search_history_bedrock_invoke(compat_result): """Send the tool-search request, take the real assistant turn back, and replay it as history with the tools still declared. diff --git a/tests/e2e/claude_code/tool_search/test_vertex_ai.py b/tests/e2e/claude_code/tool_search/test_vertex_ai.py index 7d0d35b1c1d..af629a7692e 100644 --- a/tests/e2e/claude_code/tool_search/test_vertex_ai.py +++ b/tests/e2e/claude_code/tool_search/test_vertex_ai.py @@ -45,6 +45,8 @@ from __future__ import annotations import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta + from claude_code._env import require_proxy_client from claude_code.http_probe import ( assert_tool_search_shape, @@ -61,6 +63,16 @@ VERTEX_AI_MODELS = [ @pytest.mark.skip(reason="stage red: Vertex rejects tool_search when deployment extra_headers inject context-1m beta; product/config") @pytest.mark.covers("llm.messages.vertex.tool_search.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.VERTEX_AI,), + models=tuple(VERTEX_AI_MODELS), + capabilities=(Capability.TOOL_SEARCH,), + mode=Mode.NONSTREAM, + ) +) def test_tool_search_vertex_ai(compat_result): """Probe `/v1/messages` with a `tool_search_tool_regex_20251119` tool and assert the proxy + upstream accept it for every Vertex AI diff --git a/tests/e2e/claude_code/tool_use/test_anthropic.py b/tests/e2e/claude_code/tool_use/test_anthropic.py index 9ff4c58907f..429557322f6 100644 --- a/tests/e2e/claude_code/tool_use/test_anthropic.py +++ b/tests/e2e/claude_code/tool_use/test_anthropic.py @@ -19,6 +19,7 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -72,6 +73,16 @@ def _has_tool_use_event(events: Sequence[Mapping[str, Any]]) -> bool: @pytest.mark.covers("llm.messages.anthropic.tool_use.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=tuple(ANTHROPIC_MODELS), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) +) def test_tool_use_anthropic(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a tool call was emitted on the wire.""" diff --git a/tests/e2e/claude_code/tool_use/test_azure.py b/tests/e2e/claude_code/tool_use/test_azure.py index 9e7398267c4..96946093d9d 100644 --- a/tests/e2e/claude_code/tool_use/test_azure.py +++ b/tests/e2e/claude_code/tool_use/test_azure.py @@ -23,6 +23,7 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -66,6 +67,16 @@ def _has_tool_use_event(events: Sequence[Mapping[str, Any]]) -> bool: @pytest.mark.covers("llm.messages.azure_foundry.tool_use.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE_AI,), + models=tuple(AZURE_MODELS), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) +) def test_tool_use_azure(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a tool call was emitted on the wire.""" diff --git a/tests/e2e/claude_code/tool_use/test_azure_openai.py b/tests/e2e/claude_code/tool_use/test_azure_openai.py index 7e1eecdbc03..3cb12b4023e 100644 --- a/tests/e2e/claude_code/tool_use/test_azure_openai.py +++ b/tests/e2e/claude_code/tool_use/test_azure_openai.py @@ -29,6 +29,7 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -67,6 +68,16 @@ def _has_tool_use_event(events: Sequence[Mapping[str, Any]]) -> bool: return False +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE,), + models=tuple(AZURE_OPENAI_MODELS), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) +) def test_tool_use_azure_openai(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a tool call was emitted on the wire by each GPT-5.6 tier.""" diff --git a/tests/e2e/claude_code/tool_use/test_bedrock_converse.py b/tests/e2e/claude_code/tool_use/test_bedrock_converse.py index 33d4d3820d2..f361757042b 100644 --- a/tests/e2e/claude_code/tool_use/test_bedrock_converse.py +++ b/tests/e2e/claude_code/tool_use/test_bedrock_converse.py @@ -19,6 +19,7 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -62,6 +63,16 @@ def _has_tool_use_event(events: Sequence[Mapping[str, Any]]) -> bool: @pytest.mark.covers("llm.messages.bedrock_converse.tool_use.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_CONVERSE_MODELS), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) +) def test_tool_use_bedrock_converse(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a tool call was emitted on the wire.""" diff --git a/tests/e2e/claude_code/tool_use/test_bedrock_invoke.py b/tests/e2e/claude_code/tool_use/test_bedrock_invoke.py index 47ae3aef1da..16150c5eb9d 100644 --- a/tests/e2e/claude_code/tool_use/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/tool_use/test_bedrock_invoke.py @@ -19,6 +19,7 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -62,6 +63,16 @@ def _has_tool_use_event(events: Sequence[Mapping[str, Any]]) -> bool: @pytest.mark.covers("llm.messages.bedrock_invoke.tool_use.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_INVOKE_MODELS), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) +) def test_tool_use_bedrock_invoke(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a tool call was emitted on the wire.""" diff --git a/tests/e2e/claude_code/tool_use/test_bedrock_mantle.py b/tests/e2e/claude_code/tool_use/test_bedrock_mantle.py index e9cb70e74e9..b04b19e5542 100644 --- a/tests/e2e/claude_code/tool_use/test_bedrock_mantle.py +++ b/tests/e2e/claude_code/tool_use/test_bedrock_mantle.py @@ -33,6 +33,7 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code._gpt_cells import skip_unless_mantle_cells_enabled from claude_code.cli_driver import ( @@ -72,6 +73,16 @@ def _has_tool_use_event(events: Sequence[Mapping[str, Any]]) -> bool: return False +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK_MANTLE,), + models=tuple(BEDROCK_MANTLE_MODELS), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) +) def test_tool_use_bedrock_mantle(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a tool call was emitted on the wire by each GPT-5.6 tier.""" diff --git a/tests/e2e/claude_code/tool_use/test_openai.py b/tests/e2e/claude_code/tool_use/test_openai.py index ffb7e795c2b..d8a9ba83480 100644 --- a/tests/e2e/claude_code/tool_use/test_openai.py +++ b/tests/e2e/claude_code/tool_use/test_openai.py @@ -28,6 +28,7 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code._gpt_cells import skip_unless_openai_gpt_cells_enabled from claude_code.cli_driver import ( @@ -67,6 +68,16 @@ def _has_tool_use_event(events: Sequence[Mapping[str, Any]]) -> bool: return False +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.OPENAI,), + models=tuple(OPENAI_MODELS), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) +) def test_tool_use_openai(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a tool call was emitted on the wire by each GPT-5.6 tier.""" diff --git a/tests/e2e/claude_code/tool_use/test_vertex_ai.py b/tests/e2e/claude_code/tool_use/test_vertex_ai.py index 79a3016345c..a082a61920e 100644 --- a/tests/e2e/claude_code/tool_use/test_vertex_ai.py +++ b/tests/e2e/claude_code/tool_use/test_vertex_ai.py @@ -19,6 +19,7 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -62,6 +63,16 @@ def _has_tool_use_event(events: Sequence[Mapping[str, Any]]) -> bool: @pytest.mark.covers("llm.messages.vertex.tool_use.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.VERTEX_AI,), + models=tuple(VERTEX_AI_MODELS), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) +) def test_tool_use_vertex_ai(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a tool call was emitted on the wire.""" diff --git a/tests/e2e/claude_code/tool_use/test_vertex_ai_gpt.py b/tests/e2e/claude_code/tool_use/test_vertex_ai_gpt.py index d1ebbced9dc..c15ba00d325 100644 --- a/tests/e2e/claude_code/tool_use/test_vertex_ai_gpt.py +++ b/tests/e2e/claude_code/tool_use/test_vertex_ai_gpt.py @@ -20,9 +20,11 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations +from e2e_metadata import Domain, Subject, meta from claude_code._gpt_cells import VERTEX_AI_GPT_NOT_APPLICABLE_REASON +@meta(Subject(domain=Domain.LLM_TRANSLATION)) def test_tool_use_vertex_ai_gpt(compat_result): """Record the static not_applicable outcome for this cell.""" compat_result.set( diff --git a/tests/e2e/claude_code/tool_use_streaming/test_anthropic.py b/tests/e2e/claude_code/tool_use_streaming/test_anthropic.py index 152652dcf3c..91002e7ebf9 100644 --- a/tests/e2e/claude_code/tool_use_streaming/test_anthropic.py +++ b/tests/e2e/claude_code/tool_use_streaming/test_anthropic.py @@ -29,6 +29,7 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -98,6 +99,16 @@ def _count_input_json_deltas(events: Sequence[Mapping[str, Any]]) -> int: @pytest.mark.covers("llm.messages.anthropic.tool_use.stream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=tuple(ANTHROPIC_MODELS), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.STREAM, + ) +) def test_tool_use_streaming_anthropic(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert the proxy preserves fine-grained tool streaming end-to-end.""" diff --git a/tests/e2e/claude_code/tool_use_streaming/test_azure.py b/tests/e2e/claude_code/tool_use_streaming/test_azure.py index 8a1cc1852dd..4ec93f5a667 100644 --- a/tests/e2e/claude_code/tool_use_streaming/test_azure.py +++ b/tests/e2e/claude_code/tool_use_streaming/test_azure.py @@ -21,6 +21,7 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -83,6 +84,16 @@ def _count_input_json_deltas(events: Sequence[Mapping[str, Any]]) -> int: @pytest.mark.covers("llm.messages.azure_foundry.tool_use.stream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE_AI,), + models=tuple(AZURE_MODELS), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.STREAM, + ) +) def test_tool_use_streaming_azure(compat_result): base_url, api_key = require_proxy(compat_result) diff --git a/tests/e2e/claude_code/tool_use_streaming/test_azure_openai.py b/tests/e2e/claude_code/tool_use_streaming/test_azure_openai.py index ad5d4e0f613..020f4e79eb4 100644 --- a/tests/e2e/claude_code/tool_use_streaming/test_azure_openai.py +++ b/tests/e2e/claude_code/tool_use_streaming/test_azure_openai.py @@ -31,6 +31,7 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -88,6 +89,16 @@ def _count_input_json_deltas(events: Sequence[Mapping[str, Any]]) -> int: ) +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE,), + models=tuple(AZURE_OPENAI_MODELS), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.STREAM, + ) +) def test_tool_use_streaming_azure_openai(compat_result): proxy = require_proxy(compat_result) diff --git a/tests/e2e/claude_code/tool_use_streaming/test_bedrock_converse.py b/tests/e2e/claude_code/tool_use_streaming/test_bedrock_converse.py index 3b04ed5962f..2d4b5764eb9 100644 --- a/tests/e2e/claude_code/tool_use_streaming/test_bedrock_converse.py +++ b/tests/e2e/claude_code/tool_use_streaming/test_bedrock_converse.py @@ -27,6 +27,7 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -89,6 +90,16 @@ def _count_input_json_deltas(events: Sequence[Mapping[str, Any]]) -> int: @pytest.mark.covers("llm.messages.bedrock_converse.tool_use.stream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_CONVERSE_MODELS), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.STREAM, + ) +) def test_tool_use_streaming_bedrock_converse(compat_result): base_url, api_key = require_proxy(compat_result) diff --git a/tests/e2e/claude_code/tool_use_streaming/test_bedrock_invoke.py b/tests/e2e/claude_code/tool_use_streaming/test_bedrock_invoke.py index c7b61129782..2cd7a68d57a 100644 --- a/tests/e2e/claude_code/tool_use_streaming/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/tool_use_streaming/test_bedrock_invoke.py @@ -25,6 +25,7 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -87,6 +88,16 @@ def _count_input_json_deltas(events: Sequence[Mapping[str, Any]]) -> int: @pytest.mark.covers("llm.messages.bedrock_invoke.tool_use.stream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_INVOKE_MODELS), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.STREAM, + ) +) def test_tool_use_streaming_bedrock_invoke(compat_result): base_url, api_key = require_proxy(compat_result) diff --git a/tests/e2e/claude_code/tool_use_streaming/test_bedrock_mantle.py b/tests/e2e/claude_code/tool_use_streaming/test_bedrock_mantle.py index 20fae5d48db..190eb51adc0 100644 --- a/tests/e2e/claude_code/tool_use_streaming/test_bedrock_mantle.py +++ b/tests/e2e/claude_code/tool_use_streaming/test_bedrock_mantle.py @@ -34,6 +34,7 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code._gpt_cells import skip_unless_mantle_cells_enabled from claude_code.cli_driver import ( @@ -92,6 +93,16 @@ def _count_input_json_deltas(events: Sequence[Mapping[str, Any]]) -> int: ) +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK_MANTLE,), + models=tuple(BEDROCK_MANTLE_MODELS), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.STREAM, + ) +) def test_tool_use_streaming_bedrock_mantle(compat_result): skip_unless_mantle_cells_enabled() proxy = require_proxy(compat_result) diff --git a/tests/e2e/claude_code/tool_use_streaming/test_openai.py b/tests/e2e/claude_code/tool_use_streaming/test_openai.py index a5ce31b1fd6..3555741b197 100644 --- a/tests/e2e/claude_code/tool_use_streaming/test_openai.py +++ b/tests/e2e/claude_code/tool_use_streaming/test_openai.py @@ -29,6 +29,7 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code._gpt_cells import skip_unless_openai_gpt_cells_enabled from claude_code.cli_driver import ( @@ -87,6 +88,16 @@ def _count_input_json_deltas(events: Sequence[Mapping[str, Any]]) -> int: ) +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.OPENAI,), + models=tuple(OPENAI_MODELS), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.STREAM, + ) +) def test_tool_use_streaming_openai(compat_result): skip_unless_openai_gpt_cells_enabled() proxy = require_proxy(compat_result) diff --git a/tests/e2e/claude_code/tool_use_streaming/test_vertex_ai.py b/tests/e2e/claude_code/tool_use_streaming/test_vertex_ai.py index 2912e3aae3d..7bb29742cfd 100644 --- a/tests/e2e/claude_code/tool_use_streaming/test_vertex_ai.py +++ b/tests/e2e/claude_code/tool_use_streaming/test_vertex_ai.py @@ -24,6 +24,7 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -86,6 +87,16 @@ def _count_input_json_deltas(events: Sequence[Mapping[str, Any]]) -> int: @pytest.mark.covers("llm.messages.vertex.tool_use.stream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.VERTEX_AI,), + models=tuple(VERTEX_AI_MODELS), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.STREAM, + ) +) def test_tool_use_streaming_vertex_ai(compat_result): base_url, api_key = require_proxy(compat_result) diff --git a/tests/e2e/claude_code/tool_use_streaming/test_vertex_ai_gpt.py b/tests/e2e/claude_code/tool_use_streaming/test_vertex_ai_gpt.py index 7037e91fee0..5e6c203cb57 100644 --- a/tests/e2e/claude_code/tool_use_streaming/test_vertex_ai_gpt.py +++ b/tests/e2e/claude_code/tool_use_streaming/test_vertex_ai_gpt.py @@ -20,9 +20,11 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations +from e2e_metadata import Domain, Subject, meta from claude_code._gpt_cells import VERTEX_AI_GPT_NOT_APPLICABLE_REASON +@meta(Subject(domain=Domain.LLM_TRANSLATION)) def test_tool_use_streaming_vertex_ai_gpt(compat_result): """Record the static not_applicable outcome for this cell.""" compat_result.set( diff --git a/tests/e2e/claude_code/vision/test_anthropic.py b/tests/e2e/claude_code/vision/test_anthropic.py index f681b2be5ae..b37bbb4ebf5 100644 --- a/tests/e2e/claude_code/vision/test_anthropic.py +++ b/tests/e2e/claude_code/vision/test_anthropic.py @@ -27,6 +27,7 @@ from __future__ import annotations import json import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -83,6 +84,16 @@ def _build_stdin_input() -> str: @pytest.mark.covers("llm.messages.anthropic.vision.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=tuple(ANTHROPIC_MODELS), + capabilities=(Capability.VISION,), + mode=Mode.NONSTREAM, + ) +) def test_vision_anthropic(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with an image attached via stream-json input and assert a non-empty reply.""" diff --git a/tests/e2e/claude_code/vision/test_azure.py b/tests/e2e/claude_code/vision/test_azure.py index f0eaaad84a2..739c8aba74d 100644 --- a/tests/e2e/claude_code/vision/test_azure.py +++ b/tests/e2e/claude_code/vision/test_azure.py @@ -27,6 +27,7 @@ from __future__ import annotations import json import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -83,6 +84,16 @@ def _build_stdin_input() -> str: @pytest.mark.covers("llm.messages.azure_foundry.vision.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE_AI,), + models=tuple(AZURE_MODELS), + capabilities=(Capability.VISION,), + mode=Mode.NONSTREAM, + ) +) def test_vision_azure(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with an image attached via stream-json input and assert a non-empty reply.""" diff --git a/tests/e2e/claude_code/vision/test_bedrock_converse.py b/tests/e2e/claude_code/vision/test_bedrock_converse.py index 2a5aba5a393..07f21b7d657 100644 --- a/tests/e2e/claude_code/vision/test_bedrock_converse.py +++ b/tests/e2e/claude_code/vision/test_bedrock_converse.py @@ -27,6 +27,7 @@ from __future__ import annotations import json import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -83,6 +84,16 @@ def _build_stdin_input() -> str: @pytest.mark.covers("llm.messages.bedrock_converse.vision.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_CONVERSE_MODELS), + capabilities=(Capability.VISION,), + mode=Mode.NONSTREAM, + ) +) def test_vision_bedrock_converse(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with an image attached via stream-json input and assert a non-empty reply.""" diff --git a/tests/e2e/claude_code/vision/test_bedrock_invoke.py b/tests/e2e/claude_code/vision/test_bedrock_invoke.py index 5c995cd479e..847f2c80484 100644 --- a/tests/e2e/claude_code/vision/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/vision/test_bedrock_invoke.py @@ -27,6 +27,7 @@ from __future__ import annotations import json import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -83,6 +84,16 @@ def _build_stdin_input() -> str: @pytest.mark.covers("llm.messages.bedrock_invoke.vision.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_INVOKE_MODELS), + capabilities=(Capability.VISION,), + mode=Mode.NONSTREAM, + ) +) def test_vision_bedrock_invoke(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with an image attached via stream-json input and assert a non-empty reply.""" diff --git a/tests/e2e/claude_code/vision/test_vertex_ai.py b/tests/e2e/claude_code/vision/test_vertex_ai.py index 8d385e295d0..256b762f06a 100644 --- a/tests/e2e/claude_code/vision/test_vertex_ai.py +++ b/tests/e2e/claude_code/vision/test_vertex_ai.py @@ -27,6 +27,7 @@ from __future__ import annotations import json import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -83,6 +84,16 @@ def _build_stdin_input() -> str: @pytest.mark.covers("llm.messages.vertex.vision.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.VERTEX_AI,), + models=tuple(VERTEX_AI_MODELS), + capabilities=(Capability.VISION,), + mode=Mode.NONSTREAM, + ) +) def test_vision_vertex_ai(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with an image attached via stream-json input and assert a non-empty reply.""" diff --git a/tests/e2e/claude_code/web_search/test_anthropic.py b/tests/e2e/claude_code/web_search/test_anthropic.py index a20a2133dc9..8b7f4259a68 100644 --- a/tests/e2e/claude_code/web_search/test_anthropic.py +++ b/tests/e2e/claude_code/web_search/test_anthropic.py @@ -31,6 +31,7 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -85,6 +86,16 @@ def _has_web_search_tool_use(events: Sequence[Mapping[str, Any]]) -> bool: @pytest.mark.covers("llm.messages.anthropic.web_search.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=tuple(ANTHROPIC_MODELS), + capabilities=(Capability.WEB_SEARCH,), + mode=Mode.NONSTREAM, + ) +) def test_web_search_anthropic(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert the upstream emitted a `tool_use` block calling `WebSearch`, proving diff --git a/tests/e2e/claude_code/web_search/test_azure.py b/tests/e2e/claude_code/web_search/test_azure.py index 8f9f638fbee..da3d1ebeb0b 100644 --- a/tests/e2e/claude_code/web_search/test_azure.py +++ b/tests/e2e/claude_code/web_search/test_azure.py @@ -31,6 +31,7 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -85,6 +86,16 @@ def _has_web_search_tool_use(events: Sequence[Mapping[str, Any]]) -> bool: @pytest.mark.covers("llm.messages.azure_foundry.web_search.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE_AI,), + models=tuple(AZURE_MODELS), + capabilities=(Capability.WEB_SEARCH,), + mode=Mode.NONSTREAM, + ) +) def test_web_search_azure(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert the upstream emitted a `tool_use` block calling `WebSearch`, proving diff --git a/tests/e2e/claude_code/web_search/test_bedrock_converse.py b/tests/e2e/claude_code/web_search/test_bedrock_converse.py index 32f37b2be79..36c163dd2d6 100644 --- a/tests/e2e/claude_code/web_search/test_bedrock_converse.py +++ b/tests/e2e/claude_code/web_search/test_bedrock_converse.py @@ -31,6 +31,7 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -85,6 +86,16 @@ def _has_web_search_tool_use(events: Sequence[Mapping[str, Any]]) -> bool: @pytest.mark.covers("llm.messages.bedrock_converse.web_search.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_CONVERSE_MODELS), + capabilities=(Capability.WEB_SEARCH,), + mode=Mode.NONSTREAM, + ) +) def test_web_search_bedrock_converse(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert the upstream emitted a `tool_use` block calling `WebSearch`, proving diff --git a/tests/e2e/claude_code/web_search/test_bedrock_invoke.py b/tests/e2e/claude_code/web_search/test_bedrock_invoke.py index 68d1b30e83f..1328d361404 100644 --- a/tests/e2e/claude_code/web_search/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/web_search/test_bedrock_invoke.py @@ -31,6 +31,7 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -85,6 +86,16 @@ def _has_web_search_tool_use(events: Sequence[Mapping[str, Any]]) -> bool: @pytest.mark.covers("llm.messages.bedrock_invoke.web_search.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=tuple(BEDROCK_INVOKE_MODELS), + capabilities=(Capability.WEB_SEARCH,), + mode=Mode.NONSTREAM, + ) +) def test_web_search_bedrock_invoke(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert the upstream emitted a `tool_use` block calling `WebSearch`, proving diff --git a/tests/e2e/claude_code/web_search/test_vertex_ai.py b/tests/e2e/claude_code/web_search/test_vertex_ai.py index 540a8396c98..7ab352b1953 100644 --- a/tests/e2e/claude_code/web_search/test_vertex_ai.py +++ b/tests/e2e/claude_code/web_search/test_vertex_ai.py @@ -31,6 +31,7 @@ from typing import Any, Mapping, Sequence import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from claude_code._env import require_proxy from claude_code.cli_driver import ( ClaudeCLIError, @@ -85,6 +86,16 @@ def _has_web_search_tool_use(events: Sequence[Mapping[str, Any]]) -> bool: @pytest.mark.covers("llm.messages.vertex.web_search.nonstream.works") +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.VERTEX_AI,), + models=tuple(VERTEX_AI_MODELS), + capabilities=(Capability.WEB_SEARCH,), + mode=Mode.NONSTREAM, + ) +) def test_web_search_vertex_ai(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert the upstream emitted a `tool_use` block calling `WebSearch`, proving diff --git a/tests/e2e/coverage_registry/llm_conversational.yaml b/tests/e2e/coverage_registry/llm_conversational.yaml index 919884b66f0..cd41a4548ee 100644 --- a/tests/e2e/coverage_registry/llm_conversational.yaml +++ b/tests/e2e/coverage_registry/llm_conversational.yaml @@ -101,6 +101,32 @@ - {id: llm.messages.together_ai.basic.stream.works, module: llm, tier: P1, subject_endpoint: messages, route: together_ai, capability: basic, streaming: stream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together over /v1/messages streaming"} - {id: llm.messages.together_ai.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: together_ai, capability: tool_use, streaming: nonstream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together tool calls over /v1/messages"} - {id: llm.messages.together_ai.multi_turn.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: together_ai, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together tool result round trip over /v1/messages"} +- {id: llm.chat_completions.ollama_chat.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: ollama_chat, capability: basic, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /chat/completions returns text and usage"} +- {id: llm.chat_completions.ollama_chat.basic.stream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: ollama_chat, capability: basic, streaming: stream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /chat/completions streams text deltas, usage and a terminal event"} +- {id: llm.chat_completions.ollama_chat.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: ollama_chat, capability: tool_use, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /chat/completions returns one addressable get_weather call"} +- {id: llm.chat_completions.ollama_chat.multi_turn.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: ollama_chat, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /chat/completions tool result round trip reaches the model"} +- {id: llm.chat_completions.ollama_chat.tool_use.stream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: ollama_chat, capability: tool_use, streaming: stream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) streams the call as tool_call deltas ending in finish_reason tool_calls, the shape coding agents like OpenCode execute"} +- {id: llm.messages.ollama_chat.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: ollama_chat, capability: basic, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /v1/messages returns text and usage"} +- {id: llm.messages.ollama_chat.basic.stream.works, module: llm, tier: P1, subject_endpoint: messages, route: ollama_chat, capability: basic, streaming: stream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /v1/messages streams text deltas, usage and a terminal event"} +- {id: llm.messages.ollama_chat.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: ollama_chat, capability: tool_use, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /v1/messages returns one addressable get_weather call"} +- {id: llm.messages.ollama_chat.multi_turn.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: ollama_chat, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /v1/messages tool result round trip reaches the model"} +- {id: llm.responses.ollama_chat.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: ollama_chat, capability: basic, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /v1/responses returns text and usage"} +- {id: llm.responses.ollama_chat.basic.stream.works, module: llm, tier: P1, subject_endpoint: responses, route: ollama_chat, capability: basic, streaming: stream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /v1/responses streams text deltas, usage and a terminal event"} +- {id: llm.responses.ollama_chat.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: ollama_chat, capability: tool_use, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /v1/responses returns one addressable get_weather call"} +- {id: llm.responses.ollama_chat.multi_turn.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: ollama_chat, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /v1/responses tool result round trip reaches the model"} +- {id: llm.chat_completions.ollama.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: ollama, capability: basic, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /chat/completions returns text and usage"} +- {id: llm.chat_completions.ollama.basic.stream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: ollama, capability: basic, streaming: stream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /chat/completions streams text deltas, usage and a terminal event"} +- {id: llm.chat_completions.ollama.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: ollama, capability: tool_use, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /chat/completions returns one addressable get_weather call"} +- {id: llm.chat_completions.ollama.multi_turn.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: ollama, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /chat/completions tool result round trip reaches the model"} +- {id: llm.chat_completions.ollama.tool_use.stream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: ollama, capability: tool_use, streaming: stream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) streams the call as tool_call deltas ending in finish_reason tool_calls, the shape coding agents like OpenCode execute; prompt-based JSON used to arrive as plain text with finish_reason stop (GitHub issue #35711)"} +- {id: llm.messages.ollama.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: ollama, capability: basic, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /v1/messages returns text and usage"} +- {id: llm.messages.ollama.basic.stream.works, module: llm, tier: P1, subject_endpoint: messages, route: ollama, capability: basic, streaming: stream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /v1/messages streams text deltas, usage and a terminal event"} +- {id: llm.messages.ollama.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: ollama, capability: tool_use, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /v1/messages returns one addressable get_weather call"} +- {id: llm.messages.ollama.multi_turn.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: ollama, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /v1/messages tool result round trip reaches the model"} +- {id: llm.responses.ollama.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: ollama, capability: basic, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /v1/responses returns text and usage"} +- {id: llm.responses.ollama.basic.stream.works, module: llm, tier: P1, subject_endpoint: responses, route: ollama, capability: basic, streaming: stream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /v1/responses streams text deltas, usage and a terminal event"} +- {id: llm.responses.ollama.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: ollama, capability: tool_use, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /v1/responses returns one addressable get_weather call"} +- {id: llm.responses.ollama.multi_turn.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: ollama, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /v1/responses tool result round trip reaches the model"} - {id: llm.chat_completions.sail.service_tier.nonstream.cost_logged, module: llm, tier: P1, subject_endpoint: chat_completions, route: sail, capability: service_tier, streaming: nonstream, assertions: [cost_logged], source: "llm_translation/test_sail_e2e.py", rationale: "service_tier balanced and auto map to Sail completion windows and bill the matching price columns; Sail serves flex only to background responses and Batch"} - {id: llm.chat_completions.sail.service_tier.nonstream.rejects_unknown_tier, module: llm, tier: P1, subject_endpoint: chat_completions, route: sail, capability: service_tier, streaming: nonstream, assertions: [rejects_unknown_tier], source: "llm_translation/test_sail_e2e.py", rationale: "A service_tier Sail has no completion window for is a 400 without drop_params"} - {id: llm.chat_completions.sail.service_tier.nonstream.drops_unknown_tier_and_bills_asap, module: llm, tier: P1, subject_endpoint: chat_completions, route: sail, capability: service_tier, streaming: nonstream, assertions: [drops_unknown_tier_and_bills_asap], source: "llm_translation/test_sail_e2e.py", rationale: "An unknown service_tier under drop_params is dropped and billed at asap in both the cost header and spend log"} diff --git a/tests/e2e/coverage_registry/schema.py b/tests/e2e/coverage_registry/schema.py index d089b9c1ed8..931085cbbc2 100644 --- a/tests/e2e/coverage_registry/schema.py +++ b/tests/e2e/coverage_registry/schema.py @@ -55,6 +55,8 @@ LlmRoute = Literal[ "cohere", "gemini", "hosted_vllm", + "ollama", + "ollama_chat", "openai", "sail", "together_ai", diff --git a/tests/e2e/guardrails/guardrails_client.py b/tests/e2e/guardrails/guardrails_client.py index 7772a6a1e85..b193896bf6c 100644 --- a/tests/e2e/guardrails/guardrails_client.py +++ b/tests/e2e/guardrails/guardrails_client.py @@ -11,6 +11,7 @@ from typing import Final, Literal from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, SLOW_PROVIDER_TIMEOUT_SECONDS, settle_propagation, unique_marker from e2e_http import NoBody, Result, StreamingResponse, Success, unwrap +from e2e_metadata import step from lifecycle import ResourceManager from models import ( AnthropicMessagesBody, @@ -35,7 +36,9 @@ from models import ( VideoCreateResponse, ) from proxy_client import ProxyClient -from pydantic import BaseModel +from pydantic import BaseModel, Field + +GUARDRAIL_BACKEND: Final = "gemini/gemini-2.5-flash" GuardrailMode = Literal["pre_call", "post_call", "during_call", "logging_only"] PiiEntity = Literal["EMAIL_ADDRESS", "PHONE_NUMBER", "PERSON", "CREDIT_CARD", "US_SSN"] @@ -62,14 +65,14 @@ class BedrockGuardrailParamsBody(GuardrailParamsBase): guardrail: Literal["bedrock"] = "bedrock" guardrailIdentifier: str guardrailVersion: str - aws_access_key_id: str | None = None - aws_secret_access_key: str | None = None + aws_access_key_id: str | None = Field(default=None, repr=False) + aws_secret_access_key: str | None = Field(default=None, repr=False) aws_region_name: str | None = None class OpenAIModerationParamsBody(GuardrailParamsBase): guardrail: Literal["openai_moderation"] = "openai_moderation" - api_key: str | None = None + api_key: str | None = Field(default=None, repr=False) model: str | None = None @@ -193,6 +196,7 @@ class _ResponsesGuardrailBody(BaseModel): class GuardrailsClient: proxy: ProxyClient + @step("Register the content filter guardrail {name} that blocks prompts containing {blocked_keyword}") def create_content_filter_guardrail(self, name: str, blocked_keyword: str, *, default_on: bool = True) -> str: return self.register( name, @@ -203,6 +207,7 @@ class GuardrailsClient: ), ) + @step("Register the Bedrock guardrail {name}") def create_bedrock_guardrail( self, name: str, @@ -235,7 +240,7 @@ class GuardrailsClient: resources: ResourceManager, prefix: str = "e2e-guard-backend", *, - backend: str = "gemini/gemini-2.5-flash", + backend: str = GUARDRAIL_BACKEND, api_key: str = "os.environ/GEMINI_API_KEY", ) -> str: """Register a chat deployment for a guardrail test to run against @@ -250,6 +255,7 @@ class GuardrailsClient: resources.defer(lambda: self.proxy.delete_model(model_id)) return model_name + @step("Register the {params.guardrail} guardrail {name} with mode {params.mode}") def register(self, name: str, params: GuardrailParamsBody) -> str: """Register any guardrail via POST /guardrails and return its id, once every replica can be expected to serve it. New built-ins register with @@ -274,6 +280,7 @@ class GuardrailsClient: settle_propagation(time.monotonic()) return guardrail_id + @step("Delete the guardrail") def delete_guardrail(self, guardrail_id: str) -> None: _ = self.proxy.transport.delete( f"/guardrails/{guardrail_id}", @@ -282,6 +289,7 @@ class GuardrailsClient: response_type=NoBody, ) + @step("Create the guardrail policy {body.policy_name} that adds the guardrails {body.guardrails_add}") def create_policy(self, body: PolicyCreateBody) -> str: """Create a policy via POST /policies and return its name once every replica can be expected to serve it (policies reach the data plane on the periodic @@ -297,6 +305,7 @@ class GuardrailsClient: settle_propagation(time.monotonic()) return created.policy_name + @step("Delete every version of the guardrail policy {policy_name}") def delete_policy(self, policy_name: str) -> None: _ = self.proxy.transport.delete( f"/policies/name/{policy_name}/all-versions", @@ -305,6 +314,7 @@ class GuardrailsClient: response_type=NoBody, ) + @step("Attach the guardrail policy {policy_name} to requests tagged {tags}") def attach_policy_to_tags(self, policy_name: str, tags: list[str]) -> str: attachment_id = unwrap( self.proxy.transport.post( @@ -317,6 +327,7 @@ class GuardrailsClient: settle_propagation(time.monotonic()) return attachment_id + @step("Delete the guardrail policy attachment") def delete_policy_attachment(self, attachment_id: str) -> None: _ = self.proxy.transport.delete( f"/policies/attachments/{attachment_id}", @@ -325,6 +336,7 @@ class GuardrailsClient: response_type=NoBody, ) + @step("Create the team {alias} opted out of global guardrails and wait until /team/info returns it") def create_team_opted_out_of_global_guardrails(self, alias: str) -> str: team_id = unwrap( self.proxy.transport.post( @@ -340,6 +352,7 @@ class GuardrailsClient: self._await_team(team_id) return team_id + @step("Delete the team") def delete_team(self, team_id: str) -> None: _ = self.proxy.transport.post( "/team/delete", @@ -348,9 +361,11 @@ class GuardrailsClient: response_type=NoBody, ) + @step("Generate a virtual key in the team") def create_key_in_team(self, team_id: str) -> str: return self.proxy.generate_key(KeyGenerateBody(team_id=team_id, user_id="e2e-guardrails-user")) + @step("Generate a virtual key with the guardrails {guardrails}") def create_key_with_guardrails(self, resources: ResourceManager, guardrails: list[str]) -> str: key = self.proxy.generate_key( KeyGenerateBody(user_id="e2e-guardrails-user", metadata=KeyMetadata(guardrails=guardrails)) @@ -358,6 +373,7 @@ class GuardrailsClient: resources.defer(lambda: self.proxy.delete_key(key)) return key + @step("Send a /v1/videos request to {model}") def create_video(self, key: str, model: str, prompt: str) -> Result[VideoCreateResponse]: return self.proxy.transport.post( "/v1/videos", @@ -366,6 +382,7 @@ class GuardrailsClient: response_type=VideoCreateResponse, ) + @step("Send a /v1/images/edits request to {model}") def edit_image(self, key: str, model: str, prompt: str, image: bytes) -> Result[ImageGenerationResponse]: return self.proxy.transport.upload( "/v1/images/edits", @@ -379,6 +396,7 @@ class GuardrailsClient: timeout=SLOW_PROVIDER_TIMEOUT_SECONDS, ) + @step("Send a /chat/completions request to {model}") def chat( self, key: str, @@ -407,6 +425,7 @@ class GuardrailsClient: ), ) + @step("Send a /chat/completions request to {model}") def chat_raw( self, key: str, @@ -438,6 +457,7 @@ class GuardrailsClient: ), ) + @step("Send a streaming /chat/completions request to {model}") def chat_stream_raw( self, key: str, @@ -462,6 +482,7 @@ class GuardrailsClient: ), ) + @step("Send a /v1/messages request to {model}") def messages( self, key: str, @@ -481,6 +502,7 @@ class GuardrailsClient: ), ) + @step("Send a /v1/messages request to {model}") def messages_raw( self, key: str, @@ -501,6 +523,7 @@ class GuardrailsClient: ), ) + @step("Send a streaming /v1/messages request to {model}") def messages_stream_raw( self, key: str, @@ -521,6 +544,7 @@ class GuardrailsClient: ), ) + @step("Send a /v1/responses request to {model}") def responses( self, key: str, @@ -535,6 +559,7 @@ class GuardrailsClient: json=_ResponsesGuardrailBody(model=model, input=text, guardrails=guardrails), ) + @step("Send a streaming /v1/responses request to {model}") def responses_stream_raw( self, key: str, @@ -553,6 +578,7 @@ class GuardrailsClient: stream=True, ) + @step("Apply the guardrail {name} to a piece of text with /guardrails/apply_guardrail") def apply_guardrail(self, key: str, *, name: str, text: str) -> Result[ApplyGuardrailResponse]: return self.proxy.transport.post( "/guardrails/apply_guardrail", diff --git a/tests/e2e/guardrails/test_apply_guardrail_e2e.py b/tests/e2e/guardrails/test_apply_guardrail_e2e.py index ee691db22da..7a4db80e379 100644 --- a/tests/e2e/guardrails/test_apply_guardrail_e2e.py +++ b/tests/e2e/guardrails/test_apply_guardrail_e2e.py @@ -11,6 +11,7 @@ import pytest from e2e_config import MASTER_KEY, unique_marker from e2e_http import Success, UnauthorizedError, UnknownApiError +from e2e_metadata import Domain, Route, Subject, meta from guardrails_client import GuardrailsClient from lifecycle import ResourceManager @@ -23,6 +24,12 @@ class TestApplyGuardrailEndpoint: "guardrail.litellm_content_filter.apply_endpoint.allows", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.GUARDRAILS, + ) + ) def test_apply_guardrail_blocks_banned_and_allows_clean( self, client: GuardrailsClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/guardrails/test_bedrock_guardrail_e2e.py b/tests/e2e/guardrails/test_bedrock_guardrail_e2e.py index aeffec24c61..fac15645f83 100644 --- a/tests/e2e/guardrails/test_bedrock_guardrail_e2e.py +++ b/tests/e2e/guardrails/test_bedrock_guardrail_e2e.py @@ -22,6 +22,7 @@ from typing import Final import pytest from e2e_config import unique_marker from e2e_http import StreamingResponse, UnknownApiError +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from guardrails_client import ( BedrockGuardrailParamsBody, GuardrailsClient, @@ -59,6 +60,15 @@ class TestBedrockGuardrail: "guardrail.bedrock.pre_call.blocks", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.GEMINI,), + models=(MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_bedrock_pre_call_blocks_harmful_prompt( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -96,6 +106,14 @@ class TestBedrockGuardrail: "guardrail.bedrock.post_call.blocks", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + providers=(Provider.GEMINI,), + models=(MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_bedrock_post_call_blocks_denied_model_output( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -138,6 +156,15 @@ class TestBedrockGuardrail: pytest.fail(f"bedrock post_call guardrail did not block denied model output; got {result}") @pytest.mark.covers("guardrail.bedrock.pre_call.blocks", exercised_on=["messages"]) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.MESSAGES, + providers=(Provider.GEMINI,), + models=(MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_bedrock_pre_call_blocks_on_messages( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -149,6 +176,15 @@ class TestBedrockGuardrail: _assert_policy_block(result, "/v1/messages") @pytest.mark.covers("guardrail.bedrock.pre_call.blocks", exercised_on=["responses"]) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.RESPONSES, + providers=(Provider.GEMINI,), + models=(MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_bedrock_pre_call_blocks_on_responses( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -160,6 +196,14 @@ class TestBedrockGuardrail: _assert_policy_block(result, "/v1/responses") @pytest.mark.covers("guardrail.bedrock.post_call.blocks", exercised_on=["chat_completions"]) + @meta( + Subject( + domain=Domain.GUARDRAILS, + providers=(Provider.GEMINI,), + models=(MODEL,), + mode=Mode.STREAM, + ) + ) def test_bedrock_post_call_blocks_denied_streamed_output_and_passes_clean_streams( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: diff --git a/tests/e2e/guardrails/test_block_code_execution_guardrail_e2e.py b/tests/e2e/guardrails/test_block_code_execution_guardrail_e2e.py index 7cf4c195424..7cda38d9aa4 100644 --- a/tests/e2e/guardrails/test_block_code_execution_guardrail_e2e.py +++ b/tests/e2e/guardrails/test_block_code_execution_guardrail_e2e.py @@ -20,7 +20,8 @@ import pytest from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, unique_marker from e2e_http import unwrap -from guardrails_client import BlockCodeExecutionParamsBody, GuardrailsClient +from e2e_metadata import Domain, Mode, Provider, Subject, meta +from guardrails_client import GUARDRAIL_BACKEND, BlockCodeExecutionParamsBody, GuardrailsClient from lifecycle import ResourceManager from models import ChatResponse @@ -45,6 +46,14 @@ class TestBlockCodeExecutionGuardrail: "guardrail.block_code_execution.pre_call.blocks", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + providers=(Provider.GEMINI,), + models=(GUARDRAIL_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_blocks_execution_request_but_allows_explanation( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: diff --git a/tests/e2e/guardrails/test_guardrail_dispatch_e2e.py b/tests/e2e/guardrails/test_guardrail_dispatch_e2e.py index 793974ccdb1..c8d295ee459 100644 --- a/tests/e2e/guardrails/test_guardrail_dispatch_e2e.py +++ b/tests/e2e/guardrails/test_guardrail_dispatch_e2e.py @@ -11,6 +11,7 @@ from __future__ import annotations import pytest from e2e_config import unique_marker from e2e_http import UnknownApiError, ValidationError +from e2e_metadata import Domain, Mode, Provider, Subject, meta from guardrails_client import GuardrailsClient pytestmark = pytest.mark.e2e @@ -28,6 +29,14 @@ MODEL = "gemini-2.5-flash" "guardrail.dispatch.pre_call.rejects_unknown_name", exercised_on=["chat_completions"], ) +@meta( + Subject( + domain=Domain.GUARDRAILS, + providers=(Provider.GEMINI,), + models=(MODEL,), + mode=Mode.NONSTREAM, + ) +) def test_request_naming_an_unknown_guardrail_fails_closed(client: GuardrailsClient, scoped_key: str) -> None: result = client.chat(scoped_key, MODEL, "say hi", guardrails=[f"e2e-no-such-guardrail-{unique_marker()}"]) diff --git a/tests/e2e/guardrails/test_guardrail_information_response_e2e.py b/tests/e2e/guardrails/test_guardrail_information_response_e2e.py index 9701ea35819..e0475610e8f 100644 --- a/tests/e2e/guardrails/test_guardrail_information_response_e2e.py +++ b/tests/e2e/guardrails/test_guardrail_information_response_e2e.py @@ -9,6 +9,7 @@ import pytest from e2e_config import unique_marker from e2e_http import unwrap +from e2e_metadata import Domain, Mode, Provider, Subject, meta from guardrails_client import ( BlockedWordBody, ContentFilterParamsBody, @@ -19,6 +20,8 @@ from models import ChatResponse, GuardrailInformationEntry pytestmark = pytest.mark.e2e +BACKEND_MODEL: Final = "openai/gpt-4.1-mini" + GUARDRAIL_PROPAGATION_DEADLINE_SECONDS: Final = 40.0 GUARDRAIL_PROPAGATION_POLL_INTERVAL_SECONDS: Final = 5.0 @@ -60,6 +63,14 @@ class TestGuardrailInformationResponse: "guardrail.litellm_content_filter.pre_call.returns_guardrail_information", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + providers=(Provider.OPENAI,), + models=(BACKEND_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_flag_returns_guardrail_information_for_the_guardrail_that_ran( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -68,7 +79,7 @@ class TestGuardrailInformationResponse: model = client.create_backend_model( resources, prefix="e2e-guardrail-info-backend", - backend="openai/gpt-4.1-mini", + backend=BACKEND_MODEL, api_key="os.environ/OPENAI_API_KEY", ) deadline = time.monotonic() + GUARDRAIL_PROPAGATION_DEADLINE_SECONDS @@ -91,6 +102,14 @@ class TestGuardrailInformationResponse: ) time.sleep(GUARDRAIL_PROPAGATION_POLL_INTERVAL_SECONDS) + @meta( + Subject( + domain=Domain.GUARDRAILS, + providers=(Provider.OPENAI,), + models=(BACKEND_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_without_flag_response_has_no_guardrail_information( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -99,7 +118,7 @@ class TestGuardrailInformationResponse: model = client.create_backend_model( resources, prefix="e2e-guardrail-info-backend", - backend="openai/gpt-4.1-mini", + backend=BACKEND_MODEL, api_key="os.environ/OPENAI_API_KEY", ) diff --git a/tests/e2e/guardrails/test_key_guardrail_image_edit_e2e.py b/tests/e2e/guardrails/test_key_guardrail_image_edit_e2e.py index 388a554cde7..153d5fd0aa7 100644 --- a/tests/e2e/guardrails/test_key_guardrail_image_edit_e2e.py +++ b/tests/e2e/guardrails/test_key_guardrail_image_edit_e2e.py @@ -6,6 +6,7 @@ from typing import Final import pytest from e2e_config import unique_marker from e2e_http import Success, UnknownApiError +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from guardrails_client import GuardrailsClient, poll_until_blocked from lifecycle import ResourceManager from models import LiteLLMParamsBody @@ -41,6 +42,15 @@ class TestKeyAttachedGuardrailOnImageEdits: "guardrail.litellm_content_filter.pre_call.blocks_image_edit", exercised_on=["images_edits"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.IMAGES, + providers=(Provider.GEMINI, Provider.OPENAI,), + models=(CHAT_MODEL, IMAGE_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_key_attached_content_filter_blocks_banned_image_edit_prompt( self, client: GuardrailsClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/guardrails/test_key_guardrail_video_e2e.py b/tests/e2e/guardrails/test_key_guardrail_video_e2e.py index 5f318e141a5..233cbfb8881 100644 --- a/tests/e2e/guardrails/test_key_guardrail_video_e2e.py +++ b/tests/e2e/guardrails/test_key_guardrail_video_e2e.py @@ -3,6 +3,7 @@ from __future__ import annotations import pytest from e2e_config import unique_marker from e2e_http import Success, UnknownApiError +from e2e_metadata import Domain, Mode, Provider, Subject, meta from guardrails_client import GuardrailsClient, poll_until_blocked from lifecycle import ResourceManager from models import LiteLLMParamsBody @@ -38,6 +39,14 @@ class TestKeyAttachedGuardrailOnVideos: "guardrail.litellm_content_filter.pre_call.blocks_video", exercised_on=["videos"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + providers=(Provider.GEMINI, Provider.VERTEX_AI,), + models=(CHAT_MODEL, VIDEO_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_key_attached_content_filter_blocks_banned_video_prompt( self, client: GuardrailsClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/guardrails/test_openai_moderation_category_matrix_e2e.py b/tests/e2e/guardrails/test_openai_moderation_category_matrix_e2e.py index 1f1af818290..10155cf96a5 100644 --- a/tests/e2e/guardrails/test_openai_moderation_category_matrix_e2e.py +++ b/tests/e2e/guardrails/test_openai_moderation_category_matrix_e2e.py @@ -7,15 +7,22 @@ body that names moderation; a refine-wrapper bypass must also be blocked. from __future__ import annotations +from typing import Final + import pytest from e2e_config import unique_marker from e2e_http import Result, UnknownApiError +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from guardrails_client import GuardrailsClient, OpenAIModerationParamsBody from lifecycle import ResourceManager from models import AnthropicMessagesResponse, ChatResponse pytestmark = pytest.mark.e2e +GEMINI_BACKEND: Final = "gemini/gemini-2.5-flash" +ANTHROPIC_BACKEND: Final = "anthropic/claude-haiku-4-5" +OPENAI_BACKEND: Final = "openai/gpt-4o-mini" + CATEGORY_PROMPTS: tuple[tuple[str, str], ...] = ( ( "violence", @@ -79,6 +86,15 @@ class TestOpenAIModerationCategoryMatrix: "guardrail.openai_moderations.pre_call.blocks", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.GEMINI,), + models=(GEMINI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_chat_blocks_category( self, client: GuardrailsClient, @@ -89,7 +105,7 @@ class TestOpenAIModerationCategoryMatrix: client, resources, prefix="e2e-mod-cat-chat", - backend="gemini/gemini-2.5-flash", + backend=GEMINI_BACKEND, api_key="os.environ/GEMINI_API_KEY", ) for category, prompt in CATEGORY_PROMPTS: @@ -99,6 +115,15 @@ class TestOpenAIModerationCategoryMatrix: "guardrail.openai_moderations.pre_call.blocks", exercised_on=["messages"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_messages_blocks_category( self, client: GuardrailsClient, @@ -109,7 +134,7 @@ class TestOpenAIModerationCategoryMatrix: client, resources, prefix="e2e-mod-cat-msg", - backend="anthropic/claude-haiku-4-5", + backend=ANTHROPIC_BACKEND, api_key="os.environ/ANTHROPIC_API_KEY", ) for category, prompt in CATEGORY_PROMPTS: @@ -119,6 +144,15 @@ class TestOpenAIModerationCategoryMatrix: "guardrail.openai_moderations.pre_call.blocks", exercised_on=["responses"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(OPENAI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_responses_blocks_category( self, client: GuardrailsClient, @@ -129,7 +163,7 @@ class TestOpenAIModerationCategoryMatrix: client, resources, prefix="e2e-mod-cat-resp", - backend="openai/gpt-4o-mini", + backend=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY", ) for category, prompt in CATEGORY_PROMPTS: diff --git a/tests/e2e/guardrails/test_openai_moderation_guardrail_e2e.py b/tests/e2e/guardrails/test_openai_moderation_guardrail_e2e.py index 43deb279bc8..b9593feb2ed 100644 --- a/tests/e2e/guardrails/test_openai_moderation_guardrail_e2e.py +++ b/tests/e2e/guardrails/test_openai_moderation_guardrail_e2e.py @@ -18,7 +18,9 @@ import pytest from e2e_config import unique_marker from e2e_http import UnknownApiError, unwrap +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from guardrails_client import ( + GUARDRAIL_BACKEND, GuardrailsClient, OpenAIModerationParamsBody, poll_until_blocked, @@ -37,6 +39,15 @@ class TestOpenAIModerationGuardrail: "guardrail.openai_moderations.pre_call.blocks", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.GEMINI,), + models=(GUARDRAIL_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_moderation_blocks_flagged_input( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -76,6 +87,15 @@ class TestOpenAIModerationGuardrail: "guardrail.openai_moderations.pre_call.blocks", exercised_on=["messages"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.MESSAGES, + providers=(Provider.GEMINI,), + models=(GUARDRAIL_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_moderation_blocks_flagged_input_on_messages( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: diff --git a/tests/e2e/guardrails/test_policy_inherited_guardrail_e2e.py b/tests/e2e/guardrails/test_policy_inherited_guardrail_e2e.py index 6298a1de038..8c8e86ebe7a 100644 --- a/tests/e2e/guardrails/test_policy_inherited_guardrail_e2e.py +++ b/tests/e2e/guardrails/test_policy_inherited_guardrail_e2e.py @@ -16,6 +16,7 @@ from __future__ import annotations import pytest from e2e_config import CHEAP_OPENAI_MODEL, unique_marker from e2e_http import StreamingResponse +from e2e_metadata import Domain, Mode, Provider, Subject, meta from guardrails_client import ( GuardrailsClient, PolicyConditionBody, @@ -73,6 +74,14 @@ def _setup_child_policy_attached_to_tag( class TestPolicyInheritedGuardrail: + @meta( + Subject( + domain=Domain.GUARDRAILS, + providers=(Provider.OPENAI,), + models=(MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_child_condition_miss_still_applies_inherited_parent_guardrail( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -101,6 +110,14 @@ class TestPolicyInheritedGuardrail: f"the child's own guardrail must not run when its condition fails; got {outcome.headers}" ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + providers=(Provider.OPENAI,), + models=(MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_child_condition_match_applies_child_and_inherited_parent_guardrails( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: diff --git a/tests/e2e/guardrails/test_presidio_masking_e2e.py b/tests/e2e/guardrails/test_presidio_masking_e2e.py index 35eb4884ddf..9fc2c7d3ea2 100644 --- a/tests/e2e/guardrails/test_presidio_masking_e2e.py +++ b/tests/e2e/guardrails/test_presidio_masking_e2e.py @@ -40,6 +40,7 @@ from pydantic import BaseModel, JsonValue, TypeAdapter from e2e_config import unique_marker from e2e_http import Result, StreamingResponse, Success +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from guardrails_client import GuardrailMode, GuardrailsClient, PiiAction, PiiEntity, PresidioParamsBody from lifecycle import ResourceManager from models import ( @@ -248,6 +249,15 @@ class TestPresidioPreCallMasking: "guardrail.presidio.pre_call.masks", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.GEMINI,), + models=(MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_pre_call_masks_pii_on_chat_completions( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -267,6 +277,15 @@ class TestPresidioPreCallMasking: "guardrail.presidio.pre_call.masks", exercised_on=["messages"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.MESSAGES, + providers=(Provider.GEMINI,), + models=(MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_pre_call_masks_pii_on_messages( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -312,6 +331,14 @@ class TestPresidioPostCallMasking: "guardrail.presidio.post_call.masks", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + providers=(Provider.GEMINI,), + models=(MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_post_call_masks_pii_in_model_output( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -393,6 +420,15 @@ class TestPresidioCreditCardOutputMasking: "guardrail.presidio.post_call.masks_generated_output", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.GEMINI,), + models=(MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_ui_default_scope_masks_a_card_number_the_model_generates_on_chat_completions( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -421,6 +457,15 @@ class TestPresidioCreditCardOutputMasking: "guardrail.presidio.post_call.masks_generated_output", exercised_on=["chat_completions_stream"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.GEMINI,), + models=(MODEL,), + mode=Mode.STREAM, + ) + ) def test_ui_default_scope_masks_a_card_number_the_model_generates_on_streaming_chat_completions( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -453,6 +498,15 @@ class TestPresidioCreditCardOutputMasking: "guardrail.presidio.post_call.masks_generated_output", exercised_on=["anthropic_messages_stream"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.MESSAGES, + providers=(Provider.GEMINI,), + models=(MODEL,), + mode=Mode.STREAM, + ) + ) def test_ui_default_scope_masks_a_card_number_the_model_generates_on_streaming_anthropic_messages( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -571,6 +625,15 @@ class TestPresidioSpendLogStoresMaskedOutput: ) @pytest.mark.covers(_CELL, exercised_on=["chat_completions"]) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.GEMINI,), + models=(MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_spend_log_stores_masked_output_on_chat_completions( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -585,6 +648,15 @@ class TestPresidioSpendLogStoresMaskedOutput: ) @pytest.mark.covers(_CELL, exercised_on=["chat_completions_stream"]) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.GEMINI,), + models=(MODEL,), + mode=Mode.STREAM, + ) + ) def test_spend_log_stores_masked_output_on_streaming_chat_completions( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -599,6 +671,15 @@ class TestPresidioSpendLogStoresMaskedOutput: ) @pytest.mark.covers(_CELL, exercised_on=["messages"]) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.MESSAGES, + providers=(Provider.GEMINI,), + models=(MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_spend_log_stores_masked_output_on_anthropic_messages( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -613,6 +694,15 @@ class TestPresidioSpendLogStoresMaskedOutput: ) @pytest.mark.covers(_CELL, exercised_on=["anthropic_messages_stream"]) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.MESSAGES, + providers=(Provider.GEMINI,), + models=(MODEL,), + mode=Mode.STREAM, + ) + ) def test_spend_log_stores_masked_output_on_streaming_anthropic_messages( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -627,6 +717,15 @@ class TestPresidioSpendLogStoresMaskedOutput: ) @pytest.mark.covers(_CELL, exercised_on=["responses"]) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.RESPONSES, + providers=(Provider.GEMINI,), + models=(MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_spend_log_stores_masked_output_on_responses( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -673,6 +772,14 @@ class TestPresidioSpendLogRecord: "guardrail.presidio.pre_call.logs_masked_entities", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + providers=(Provider.GEMINI,), + models=(MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_masking_run_is_recorded_on_the_spend_log( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: diff --git a/tests/e2e/guardrails/test_responses_pre_call_block_stream_e2e.py b/tests/e2e/guardrails/test_responses_pre_call_block_stream_e2e.py index 93512e2a64c..cdf3098e85d 100644 --- a/tests/e2e/guardrails/test_responses_pre_call_block_stream_e2e.py +++ b/tests/e2e/guardrails/test_responses_pre_call_block_stream_e2e.py @@ -7,12 +7,15 @@ from typing import Final import pytest from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, unique_marker from e2e_http import StreamingResponse +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from guardrails_client import CustomCodeParamsBody, GuardrailsClient from lifecycle import ResourceManager from pydantic import BaseModel, TypeAdapter pytestmark = pytest.mark.e2e +BACKEND_MODEL: Final = "openai/gpt-4.1-mini" + DENIAL: Final = "This model is not currently available. Please contact support if you think this is a mistake." CUSTOM_CODE: Final = f''' @@ -109,12 +112,21 @@ class TestResponsesPreCallBlock: return name @pytest.mark.covers("guardrail.custom_code.pre_call.blocks", exercised_on=["responses"]) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(BACKEND_MODEL,), + mode=Mode.STREAM, + ) + ) def test_stream_block_is_sse_with_completed_assistant_message( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: name: Final = self._register_block(client, resources) model: Final = client.create_backend_model( - resources, prefix="e2e-responses-block", backend="openai/gpt-4.1-mini", api_key="os.environ/OPENAI_API_KEY" + resources, prefix="e2e-responses-block", backend=BACKEND_MODEL, api_key="os.environ/OPENAI_API_KEY" ) result: Final = _poll_for_block( @@ -137,12 +149,21 @@ class TestResponsesPreCallBlock: _assert_blocked_response(completed[0].response) @pytest.mark.covers("guardrail.custom_code.pre_call.blocks", exercised_on=["responses"]) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(BACKEND_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_non_stream_block_is_schema_valid_json( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: name: Final = self._register_block(client, resources) model: Final = client.create_backend_model( - resources, prefix="e2e-responses-block", backend="openai/gpt-4.1-mini", api_key="os.environ/OPENAI_API_KEY" + resources, prefix="e2e-responses-block", backend=BACKEND_MODEL, api_key="os.environ/OPENAI_API_KEY" ) result: Final = _poll_for_block(lambda: client.responses(scoped_key, model, "say hi", guardrails=[name])) diff --git a/tests/e2e/guardrails/test_streaming_guardrail_e2e.py b/tests/e2e/guardrails/test_streaming_guardrail_e2e.py index 911ddf9304b..324dd4107a4 100644 --- a/tests/e2e/guardrails/test_streaming_guardrail_e2e.py +++ b/tests/e2e/guardrails/test_streaming_guardrail_e2e.py @@ -22,6 +22,7 @@ import os import pytest from e2e_config import unique_marker +from e2e_metadata import Domain, Mode, Provider, Subject, meta from guardrails_client import ( BedrockGuardrailParamsBody, GuardrailsClient, @@ -39,6 +40,14 @@ class TestBedrockDuringCallStreaming: "guardrail.bedrock.during.blocks", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + providers=(Provider.GEMINI,), + models=(MODEL,), + mode=Mode.STREAM, + ) + ) def test_during_call_blocks_stream_before_first_chunk( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: diff --git a/tests/e2e/guardrails/test_team_disable_global_guardrail_e2e.py b/tests/e2e/guardrails/test_team_disable_global_guardrail_e2e.py index db917d6ede9..0cc6106abfd 100644 --- a/tests/e2e/guardrails/test_team_disable_global_guardrail_e2e.py +++ b/tests/e2e/guardrails/test_team_disable_global_guardrail_e2e.py @@ -14,6 +14,7 @@ import pytest from e2e_config import unique_marker from e2e_http import UnknownApiError, unwrap +from e2e_metadata import Domain, Mode, Provider, Subject, meta from guardrails_client import GuardrailsClient from lifecycle import ResourceManager @@ -59,6 +60,14 @@ class TestTeamDisableGlobalGuardrail: "guardrail.litellm_content_filter.pre_call.blocks", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + providers=(Provider.GEMINI,), + models=(MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_global_guardrail_blocks_key_without_team_opt_out( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -72,6 +81,14 @@ class TestTeamDisableGlobalGuardrail: "guardrail.litellm_content_filter.pre_call.allows", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + providers=(Provider.GEMINI,), + models=(MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_team_with_disable_flag_bypasses_global_guardrail( self, client: GuardrailsClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/guardrails/test_tool_permission_guardrail_e2e.py b/tests/e2e/guardrails/test_tool_permission_guardrail_e2e.py index 8d1047e53c7..d5a6a9d9866 100644 --- a/tests/e2e/guardrails/test_tool_permission_guardrail_e2e.py +++ b/tests/e2e/guardrails/test_tool_permission_guardrail_e2e.py @@ -25,6 +25,7 @@ import pytest from e2e_config import unique_marker from e2e_http import StreamingResponse, UnknownApiError +from e2e_metadata import Capability, Domain, Mode, Provider, Subject, meta from guardrails_client import ( GuardrailsClient, ToolPermissionParamsBody, @@ -101,6 +102,15 @@ def _tool_call_names(response: ChatResponse) -> tuple[str, ...]: class TestToolPermissionPreCall: @pytest.mark.covers("guardrail.tool_permission.pre_call.blocks", exercised_on=["chat_completions"]) + @meta( + Subject( + domain=Domain.GUARDRAILS, + providers=(Provider.GEMINI,), + models=(MODEL,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_pre_call_blocks_tool_outside_the_allow_list( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -135,6 +145,15 @@ class TestToolPermissionPreCall: pytest.fail(f"tool_permission let a tool outside the allow-list through; got {result}") @pytest.mark.covers("guardrail.tool_permission.pre_call.allows", exercised_on=["chat_completions"]) + @meta( + Subject( + domain=Domain.GUARDRAILS, + providers=(Provider.GEMINI,), + models=(MODEL,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_pre_call_allows_permitted_tool( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: diff --git a/tests/e2e/llm_translation/conversational_matrix.py b/tests/e2e/llm_translation/conversational_matrix.py index 0d6f6ed3d4e..20dbe236fae 100644 --- a/tests/e2e/llm_translation/conversational_matrix.py +++ b/tests/e2e/llm_translation/conversational_matrix.py @@ -33,6 +33,8 @@ from anthropic.types import ( ToolUseBlockParam, ) from e2e_config import provider_edge_base, unique_marker +from e2e_metadata import Capability as MetaCapability +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta, step from lifecycle import ResourceManager from llm_translation.sdk_clients import NO_PROXY_CACHE, SdkClients, response_header from models import CredentialCreateBody, LiteLLMParamsBody @@ -65,6 +67,10 @@ Streaming = Literal["stream", "nonstream"] Assertion = Literal["works", "cost_logged"] ToolMode = Literal["none", "forced", "offered"] +GPT_4O_MINI_BACKEND: Final = "openai/gpt-4o-mini" +GPT_5_4_MINI_BACKEND: Final = "openai/gpt-5.4-mini" +CLAUDE_HAIKU_BACKEND: Final = "anthropic/claude-haiku-4-5" + SURFACES: Final[tuple[SurfaceName, ...]] = ("chat_completions", "messages", "responses") AUTH_METHODS: Final[tuple[AuthMethod, ...]] = ("env_ref", "stored_credential") @@ -104,12 +110,19 @@ class Deployment: assert key, f"{self.api_key_env} is not set in the test process environment" return key + def provider(self) -> Provider: + match self.route: + case "openai": + return Provider.OPENAI + case "anthropic": + return Provider.ANTHROPIC + DEPLOYMENTS: Final[tuple[Deployment, ...]] = ( Deployment( route="openai", label="gpt-4o-mini", - backend="openai/gpt-4o-mini", + backend=GPT_4O_MINI_BACKEND, api_key_env="OPENAI_API_KEY", edge_mount="openai", edge_suffix="/v1", @@ -117,7 +130,7 @@ DEPLOYMENTS: Final[tuple[Deployment, ...]] = ( Deployment( route="openai", label="gpt-5.4-mini", - backend="openai/gpt-5.4-mini", + backend=GPT_5_4_MINI_BACKEND, api_key_env="OPENAI_API_KEY", edge_mount="openai", edge_suffix="/v1", @@ -125,7 +138,7 @@ DEPLOYMENTS: Final[tuple[Deployment, ...]] = ( Deployment( route="anthropic", label="claude-haiku-4-5", - backend="anthropic/claude-haiku-4-5", + backend=CLAUDE_HAIKU_BACKEND, api_key_env="ANTHROPIC_API_KEY", edge_mount="anthropic", edge_suffix="", @@ -146,6 +159,26 @@ class Cell: def registry_id(self, capability: Capability, streaming: Streaming, assertion: Assertion) -> str: return f"llm.{self.surface}.{self.deployment.route}.{capability}.{streaming}.{assertion}" + def subject(self, capability: Capability, streaming: Streaming, assertion: Assertion) -> Subject: + return Subject( + domain=Domain.SPEND_BUDGETS if assertion == "cost_logged" else Domain.LLM_TRANSLATION, + route=_surface_route(self.surface), + providers=(self.deployment.provider(),), + models=(self.deployment.backend,), + capabilities=() if capability == "basic" else (MetaCapability.FUNCTION_CALLING,), + mode=Mode.STREAM if streaming == "stream" else Mode.NONSTREAM, + ) + + +def _surface_route(surface: SurfaceName) -> Route: + match surface: + case "chat_completions": + return Route.CHAT_COMPLETIONS + case "messages": + return Route.MESSAGES + case "responses": + return Route.RESPONSES + CELLS: Final[tuple[Cell, ...]] = tuple( Cell(surface=surface, deployment=deployment, auth=auth) @@ -158,7 +191,14 @@ CELLS: Final[tuple[Cell, ...]] = tuple( def cells_covering(capability: Capability, streaming: Streaming, assertion: Assertion) -> tuple[ParameterSet, ...]: """Every cell as a pytest param carrying the registry id its test proves.""" return tuple( - pytest.param(cell, id=cell.id, marks=pytest.mark.covers(cell.registry_id(capability, streaming, assertion))) + pytest.param( + cell, + id=cell.id, + marks=( + pytest.mark.covers(cell.registry_id(capability, streaming, assertion)), + meta(cell.subject(capability, streaming, assertion)), + ), + ) for cell in CELLS ) @@ -352,9 +392,14 @@ class ChatCompletionsSurface: cost_header=response_header(raw.headers, "x-litellm-response-cost"), ) + @step( + 'Send a /chat/completions request to {model} with the prompt "{prompt}"' + " and forced weather tool use set to {with_tool}" + ) def reply(self, key: str, model: str, prompt: str, *, with_tool: bool = False) -> Reply: return self._turn(key, model, _chat_history(prompt), "forced" if with_tool else "none") + @step('Send a streaming /chat/completions request to {model} with the prompt "{prompt}"') def stream(self, key: str, model: str, prompt: str) -> StreamedReply: chunks: Final[tuple[ChatCompletionChunk, ...]] = tuple( self.sdk.openai(key).chat.completions.create( @@ -373,6 +418,7 @@ class ChatCompletionsSurface: event_count=len(chunks), ) + @step("Send the {call.name} tool result back to {model} over /chat/completions") def reply_to_tool_result(self, key: str, model: str, prompt: str, call: ToolCall, result: str) -> Reply: tool_call: Final[ChatCompletionMessageFunctionToolCallParam] = { "id": call.call_id, @@ -421,9 +467,14 @@ class MessagesSurface: cost_header=response_header(raw.headers, "x-litellm-response-cost"), ) + @step( + 'Send a /v1/messages request to {model} with the prompt "{prompt}"' + " and forced weather tool use set to {with_tool}" + ) def reply(self, key: str, model: str, prompt: str, *, with_tool: bool = False) -> Reply: return self._turn(key, model, ({"role": "user", "content": prompt},), "forced" if with_tool else "none") + @step('Send a streaming /v1/messages request to {model} with the prompt "{prompt}"') def stream(self, key: str, model: str, prompt: str) -> StreamedReply: events: Final[tuple[RawMessageStreamEvent, ...]] = tuple( self.sdk.anthropic(key).messages.create( @@ -446,6 +497,7 @@ class MessagesSurface: event_count=len(events), ) + @step("Send the {call.name} tool result back to {model} over /v1/messages") def reply_to_tool_result(self, key: str, model: str, prompt: str, call: ToolCall, result: str) -> Reply: tool_use: Final[ToolUseBlockParam] = { "type": "tool_use", @@ -496,9 +548,14 @@ class ResponsesSurface: cost_header=response_header(raw.headers, "x-litellm-response-cost"), ) + @step( + 'Send a /v1/responses request to {model} with the prompt "{prompt}"' + " and forced weather tool use set to {with_tool}" + ) def reply(self, key: str, model: str, prompt: str, *, with_tool: bool = False) -> Reply: return self._turn(key, model, [{"role": "user", "content": prompt}], "forced" if with_tool else "none") + @step('Send a streaming /v1/responses request to {model} with the prompt "{prompt}"') def stream(self, key: str, model: str, prompt: str) -> StreamedReply: events: Final[tuple[ResponseStreamEvent, ...]] = tuple( self.sdk.openai(key).responses.create( @@ -519,6 +576,7 @@ class ResponsesSurface: event_count=len(events), ) + @step("Send the {call.name} tool result back to {model} over /v1/responses") def reply_to_tool_result(self, key: str, model: str, prompt: str, call: ToolCall, result: str) -> Reply: function_call: Final[ResponseFunctionToolCallParam] = { "type": "function_call", diff --git a/tests/e2e/llm_translation/passthrough_client.py b/tests/e2e/llm_translation/passthrough_client.py index a56d3dc077e..478148e45aa 100644 --- a/tests/e2e/llm_translation/passthrough_client.py +++ b/tests/e2e/llm_translation/passthrough_client.py @@ -18,6 +18,7 @@ from websockets.exceptions import InvalidStatus from websockets.sync.client import connect from e2e_config import ws_base_url +from e2e_metadata import step from proxy_client import ProxyClient from e2e_http import FileUploadForm, Headers, NoBody, Result, StreamingResponse from models import ChatMessage @@ -34,7 +35,7 @@ class JsonSchema(BaseModel): class GeminiHeaders(Headers): - x_goog_api_key: str = Field(serialization_alias="x-goog-api-key") + x_goog_api_key: str = Field(serialization_alias="x-goog-api-key", repr=False) content_type: str = Field( default="application/json", serialization_alias="Content-Type" ) @@ -42,7 +43,7 @@ class GeminiHeaders(Headers): class AnthropicHeaders(Headers): - x_api_key: str = Field(serialization_alias="x-api-key") + x_api_key: str = Field(serialization_alias="x-api-key", repr=False) anthropic_version: str = Field( default="2023-06-01", serialization_alias="anthropic-version" ) @@ -56,7 +57,7 @@ class VertexHeaders(Headers): # Only the litellm virtual key; the /vertex_ai passthrough mints the Vertex token # from the proxy's own service account (the deployment marked use_in_pass_through), # so no upstream Authorization bearer is sent from the client. - x_litellm_api_key: str = Field(serialization_alias="x-litellm-api-key") + x_litellm_api_key: str = Field(serialization_alias="x-litellm-api-key", repr=False) content_type: str = Field( default="application/json", serialization_alias="Content-Type" ) @@ -246,6 +247,7 @@ class PassthroughClient: # ---- Gemini native passthrough (/gemini/v1beta/...) ----------------- + @step("Send a Gemini generateContent request to {model} through /gemini") def gemini_generate( self, key: str, @@ -263,6 +265,7 @@ class PassthroughClient: ), ) + @step("Send a Gemini streamGenerateContent request to {model} through /gemini") def gemini_stream( self, key: str, model: str, text: str, *, tags: list[str] | None = None ) -> StreamingResponse: @@ -278,6 +281,7 @@ class PassthroughClient: # ---- Vertex AI native passthrough (/vertex_ai/v1/projects/...) ------- + @step("Send a Vertex AI generateContent request to {model} in {location} through /vertex_ai") def vertex_generate( self, key: str, project: str, location: str, model: str, text: str ) -> StreamingResponse: @@ -295,6 +299,7 @@ class PassthroughClient: # ---- Anthropic native passthrough (/anthropic/v1/messages) ---------- + @step("Send a /v1/messages request to {model} through /anthropic with streaming set to {stream}") def anthropic_message( self, key: str, @@ -324,6 +329,7 @@ class PassthroughClient: # Relayed to OpenAI untouched, which is the whole point of the prefix: the # customer opts out of the gateway's managed-file handling here. + @step("Upload {filename} to /openai_passthrough/v1/files") def openai_passthrough_upload_file( self, key: str, *, content: bytes, filename: str ) -> Result[PassthroughFileObject]: @@ -336,6 +342,7 @@ class PassthroughClient: response_type=PassthroughFileObject, ) + @step("Delete the uploaded file through /openai_passthrough/v1/files") def openai_passthrough_delete_file( self, key: str, file_id: str ) -> Result[PassthroughFileDeleted]: @@ -346,6 +353,7 @@ class PassthroughClient: response_type=PassthroughFileDeleted, ) + @step("List batches from /openai_passthrough/v1/batches") def openai_passthrough_list_batches(self, key: str) -> Result[PassthroughBatchList]: return self.proxy.transport.get( "/openai_passthrough/v1/batches", @@ -360,6 +368,7 @@ class PassthroughClient: # budgets against this traffic, so a 200 that logs no spend is money the # gateway never sees. + @step("Send a /v1/responses request to {model} through /openai_passthrough with streaming set to {stream}") def openai_passthrough_responses( self, key: str, model: str, text: str, *, stream: bool = False ) -> StreamingResponse: @@ -370,6 +379,7 @@ class PassthroughClient: stream=stream, ) + @step('Send a /v1/embeddings request to {model} through /openai_passthrough for "{text}"') def openai_passthrough_embed( self, key: str, model: str, text: str ) -> StreamingResponse: @@ -379,6 +389,7 @@ class PassthroughClient: json=OpenAIEmbeddingBody(model=model, input=text), ) + @step("Send a /v1/chat/completions request to {model} through /openai") def openai_chat( self, key: str, model: str, text: str, *, max_completion_tokens: int = 64 ) -> StreamingResponse: @@ -397,6 +408,7 @@ class PassthroughClient: # The same prefixes over an upgrade instead of a POST, for the provider APIs # that only speak websocket (realtime, responses.connect). + @step("Open a websocket to {path} and wait for its first event") def openai_passthrough_websocket( self, key: str, diff --git a/tests/e2e/llm_translation/realtime/realtime_client.py b/tests/e2e/llm_translation/realtime/realtime_client.py index a280f4bc26b..c9f06e787f3 100644 --- a/tests/e2e/llm_translation/realtime/realtime_client.py +++ b/tests/e2e/llm_translation/realtime/realtime_client.py @@ -307,9 +307,11 @@ def as_text(message: str | bytes) -> str: class RealtimeSession: connection: Connection - def send(self, event: BaseModel) -> None: + @step("Send the realtime event {event.type} over the websocket") + 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") def collect_until( self, stop_type: str, *, timeout: float ) -> tuple[ReceivedEvent, ...]: @@ -381,6 +383,7 @@ class RealtimeSession: class RealtimeClient: proxy: ProxyClient + @step("Add a realtime deployment that calls {provider.litellm_params.model}") def provision(self, provider: RealtimeProvider) -> tuple[str, str]: """Register this provider's realtime deployment through /model/new and return (model_name, model_id). The name is marker-unique so it never collides with a @@ -393,6 +396,7 @@ class RealtimeClient: ) return model_name, model_id + @step("Open a /v1/realtime websocket session to {model}") @contextmanager def connect( self, *, key: str, model: str, timeout: float = 15.0 diff --git a/tests/e2e/llm_translation/realtime/test_realtime_bedrock_e2e.py b/tests/e2e/llm_translation/realtime/test_realtime_bedrock_e2e.py index 656882a4d92..a86f1daff63 100644 --- a/tests/e2e/llm_translation/realtime/test_realtime_bedrock_e2e.py +++ b/tests/e2e/llm_translation/realtime/test_realtime_bedrock_e2e.py @@ -19,6 +19,7 @@ from __future__ import annotations import pytest from e2e_config import unique_marker +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import LiteLLMParamsBody from realtime_client import ( @@ -42,6 +43,15 @@ class TestNovaSonicRealtime: "llm.realtime.bedrock_converse.basic.stream.works", exercised_on=["realtime"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.REALTIME, + providers=(Provider.BEDROCK,), + models=(NOVA_SONIC,), + mode=Mode.WEBSOCKET, + ) + ) def test_nova_sonic_response_create_completes( self, client: RealtimeClient, resources: ResourceManager, scoped_key: str ) -> None: diff --git a/tests/e2e/llm_translation/realtime/test_realtime_e2e.py b/tests/e2e/llm_translation/realtime/test_realtime_e2e.py index d7870b26497..622ed9d507f 100644 --- a/tests/e2e/llm_translation/realtime/test_realtime_e2e.py +++ b/tests/e2e/llm_translation/realtime/test_realtime_e2e.py @@ -12,7 +12,10 @@ hard failure, not a skip; once configured, a protocol failure is likewise a hard failure. See REALTIME_COVERAGE_MATRIX.md. """ +from typing import Final + import pytest +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import LiteLLMParamsBody from pydantic import BaseModel @@ -42,7 +45,42 @@ from websockets.exceptions import ConnectionClosedError pytestmark = pytest.mark.e2e -PROVIDER_PARAMS = [pytest.param(p, id=p.id) for p in PROVIDERS] +AZURE_REALTIME_MODEL: Final = "azure/gpt-realtime" + +TEXT_PARAMS: Final = tuple( + pytest.param( + p, + id=p.id, + marks=meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.REALTIME, + providers=(Provider(p.id),), + models=(p.litellm_params.model,), + mode=Mode.WEBSOCKET, + ) + ), + ) + for p in PROVIDERS +) + +TOOL_PARAMS: Final = tuple( + pytest.param( + p, + id=p.id, + marks=meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.REALTIME, + providers=(Provider(p.id),), + models=(p.litellm_params.model,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.WEBSOCKET, + ) + ), + ) + for p in PROVIDERS +) WEATHER_TOOL = FunctionTool( name="get_weather", @@ -62,7 +100,7 @@ class WeatherResult(BaseModel): temperature_f: int -@pytest.mark.parametrize("provider", PROVIDER_PARAMS) +@pytest.mark.parametrize("provider", TEXT_PARAMS) def test_text_conversation( client: RealtimeClient, scoped_key: str, @@ -99,7 +137,7 @@ def test_text_conversation( assert done.response.usage is not None, "response.done missing normalized usage" -@pytest.mark.parametrize("provider", PROVIDER_PARAMS) +@pytest.mark.parametrize("provider", TOOL_PARAMS) def test_tool_call_round_trip( client: RealtimeClient, scoped_key: str, @@ -158,7 +196,7 @@ _REFUSED_UPSTREAMS = ( "azure-bad-key", "azure-realtime-refused", LiteLLMParamsBody( - model="azure/gpt-realtime", + model=AZURE_REALTIME_MODEL, api_key="invalid-e2e-key", api_version="2025-08-28", realtime_protocol="GA", @@ -168,6 +206,15 @@ _REFUSED_UPSTREAMS = ( @pytest.mark.parametrize("provider", _REFUSED_UPSTREAMS, ids=[p.id for p in _REFUSED_UPSTREAMS]) +@meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.REALTIME, + providers=(Provider.AZURE,), + models=(AZURE_REALTIME_MODEL,), + mode=Mode.WEBSOCKET, + ) +) def test_upstream_handshake_refusal_is_an_error_event_and_policy_close( client: RealtimeClient, resources: ResourceManager, diff --git a/tests/e2e/llm_translation/realtime/test_realtime_pipecat_audio_e2e.py b/tests/e2e/llm_translation/realtime/test_realtime_pipecat_audio_e2e.py index 2e9cfcfe648..27d5abd1828 100644 --- a/tests/e2e/llm_translation/realtime/test_realtime_pipecat_audio_e2e.py +++ b/tests/e2e/llm_translation/realtime/test_realtime_pipecat_audio_e2e.py @@ -24,10 +24,12 @@ Three test scenarios per provider: import asyncio import wave from pathlib import Path +from typing import Final import pytest from e2e_config import ws_base_url +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from realtime_client import ( PROVIDERS, RealtimeProvider, @@ -73,7 +75,59 @@ from pipecat.services.openai.realtime.llm import OpenAIRealtimeLLMService # noq from pipecat_service import LiteLLMRealtimeLLMService # noqa: E402 -PROVIDER_PARAMS = [pytest.param(p, id=p.id) for p in PROVIDERS] +TOOL_PARAMS: Final = tuple( + pytest.param( + p, + id=p.id, + marks=meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.REALTIME, + providers=(Provider(p.id),), + models=(p.litellm_params.model,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.WEBSOCKET, + ) + ), + ) + for p in PROVIDERS +) + +AUDIO_OUTPUT_PARAMS: Final = tuple( + pytest.param( + p, + id=p.id, + marks=meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.REALTIME, + providers=(Provider(p.id),), + models=(p.litellm_params.model,), + capabilities=(Capability.AUDIO_OUTPUT,), + mode=Mode.WEBSOCKET, + ) + ), + ) + for p in PROVIDERS +) + +AUDIO_INPUT_PARAMS: Final = tuple( + pytest.param( + p, + id=p.id, + marks=meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.REALTIME, + providers=(Provider(p.id),), + models=(p.litellm_params.model,), + capabilities=(Capability.AUDIO_INPUT, Capability.AUDIO_OUTPUT), + mode=Mode.WEBSOCKET, + ) + ), + ) + for p in PROVIDERS +) # PCM16 24 kHz mono WAV of "What is the weather in Paris?" (generated via macOS # `say` and resampled with audioop). Used by the server-VAD audio-input test. @@ -193,7 +247,7 @@ async def _run_pipeline( # --------------------------------------------------------------------------- -@pytest.mark.parametrize("provider", PROVIDER_PARAMS) +@pytest.mark.parametrize("provider", TOOL_PARAMS) def test_pipecat_server_vad( scoped_key: str, realtime_models: dict[str, str], @@ -208,7 +262,7 @@ def test_pipecat_server_vad( assert got_text, "no assistant text frames produced" -@pytest.mark.parametrize("provider", PROVIDER_PARAMS) +@pytest.mark.parametrize("provider", AUDIO_OUTPUT_PARAMS) def test_pipecat_audio_output( scoped_key: str, realtime_models: dict[str, str], @@ -332,7 +386,7 @@ async def _run_audio_input_pipeline( return bool(capture.texts), capture.audio_bytes -@pytest.mark.parametrize("provider", PROVIDER_PARAMS) +@pytest.mark.parametrize("provider", AUDIO_INPUT_PARAMS) def test_pipecat_server_vad_audio_input( scoped_key: str, realtime_models: dict[str, str], diff --git a/tests/e2e/llm_translation/realtime/test_realtime_pipecat_e2e.py b/tests/e2e/llm_translation/realtime/test_realtime_pipecat_e2e.py index f84ce197f88..3c22ae1a928 100644 --- a/tests/e2e/llm_translation/realtime/test_realtime_pipecat_e2e.py +++ b/tests/e2e/llm_translation/realtime/test_realtime_pipecat_e2e.py @@ -26,6 +26,7 @@ import asyncio import pytest from e2e_config import ws_base_url +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from realtime_client import ( PROVIDERS, RealtimeProvider, @@ -64,7 +65,21 @@ from pipecat_service import LiteLLMRealtimeLLMService # noqa: E402 # pipecat-ai/pipecat#2544); raw-ws tool_call_round_trip[vertex_ai] is the # source of truth for that provider. Keep openai/azure/gemini here. PROVIDER_PARAMS = [ - pytest.param(p, id=p.id) for p in PROVIDERS if p.id != "vertex_ai" + pytest.param( + p, + id=p.id, + marks=meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.REALTIME, + providers=(Provider(p.id),), + models=(p.litellm_params.model,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.WEBSOCKET, + ) + ), + ) + for p in PROVIDERS if p.id != "vertex_ai" ] WEATHER_TOOL = ToolsSchema( diff --git a/tests/e2e/llm_translation/test_audio_speech_e2e.py b/tests/e2e/llm_translation/test_audio_speech_e2e.py index 75a3de86d13..c4c3f751828 100644 --- a/tests/e2e/llm_translation/test_audio_speech_e2e.py +++ b/tests/e2e/llm_translation/test_audio_speech_e2e.py @@ -10,8 +10,11 @@ the SDK refuses to send a request missing its required fields. from __future__ import annotations +from typing import Final + import pytest from e2e_config import unique_marker +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from e2e_http import assert_client_error from lifecycle import ResourceManager from models import LiteLLMParamsBody @@ -21,6 +24,9 @@ from sdk_clients import SdkClients, response_header pytestmark = pytest.mark.e2e +OPENAI_TTS_MODEL: Final = "openai/gpt-4o-mini-tts" +AWS_POLLY_MODEL: Final = "aws_polly/generative" + class _OptionalSpeechBody(BaseModel): model: str | None = None @@ -32,7 +38,7 @@ def _register_tts(proxy: ProxyClient, resources: ResourceManager) -> tuple[str, model = f"e2e-speech-{unique_marker()}" model_id = proxy.create_model( model, - LiteLLMParamsBody(model="openai/gpt-4o-mini-tts", api_key="os.environ/OPENAI_API_KEY"), + LiteLLMParamsBody(model=OPENAI_TTS_MODEL, api_key="os.environ/OPENAI_API_KEY"), ) resources.defer(lambda: proxy.delete_model(model_id)) return model, resources.key() @@ -40,6 +46,15 @@ def _register_tts(proxy: ProxyClient, resources: ResourceManager) -> tuple[str, class TestAudioSpeech: @pytest.mark.covers("llm.audio_speech.openai.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.AUDIO, + providers=(Provider.OPENAI,), + models=(OPENAI_TTS_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_audio_speech_returns_audio( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -56,6 +71,15 @@ class TestAudioSpeech: assert response.content, "/audio/speech returned an empty body" @pytest.mark.covers("llm.audio_speech.openai.basic.stream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.AUDIO, + providers=(Provider.OPENAI,), + models=(OPENAI_TTS_MODEL,), + mode=Mode.STREAM, + ) + ) def test_audio_speech_streams_audio_chunks( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -90,6 +114,15 @@ class TestAudioSpeech: @pytest.mark.skip(reason="stage red: product gap, /v1/audio/speech 500s on missing input instead of 400") @pytest.mark.covers("llm.audio_speech.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.AUDIO, + providers=(Provider.OPENAI,), + models=(OPENAI_TTS_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_missing_input_returns_error( self, proxy: ProxyClient, resources: ResourceManager ) -> None: @@ -103,6 +136,12 @@ class TestAudioSpeech: @pytest.mark.skip(reason="stage red: product gap, /v1/audio/speech 500s on missing model instead of 400") @pytest.mark.covers("llm.audio_speech.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.AUDIO, + ) + ) def test_missing_model_returns_error( self, proxy: ProxyClient, resources: ResourceManager ) -> None: @@ -116,6 +155,15 @@ class TestAudioSpeech: @pytest.mark.skip(reason="stage red: product gap, /v1/audio/speech 500s on invalid voice instead of surfacing the provider 4xx") @pytest.mark.covers("llm.audio_speech.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.AUDIO, + providers=(Provider.OPENAI,), + models=(OPENAI_TTS_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_invalid_voice_returns_error( self, proxy: ProxyClient, resources: ResourceManager ) -> None: @@ -129,6 +177,15 @@ class TestAudioSpeech: @pytest.mark.skip(reason="stage red: product gap, /v1/audio/speech 500s on empty input instead of surfacing the provider 4xx") @pytest.mark.covers("llm.audio_speech.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.AUDIO, + providers=(Provider.OPENAI,), + models=(OPENAI_TTS_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_empty_input_returns_error( self, proxy: ProxyClient, resources: ResourceManager ) -> None: @@ -145,6 +202,15 @@ MP3_PREFIXES = (b"ID3", b"\xff\xfb", b"\xff\xf3", b"\xff\xf2") class TestAwsPollySpeech: + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.AUDIO, + providers=(Provider.AWS_POLLY,), + models=(AWS_POLLY_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_polly_generative_voice_returns_mp3( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -152,7 +218,7 @@ class TestAwsPollySpeech: model_id = proxy.create_model( model, LiteLLMParamsBody( - model="aws_polly/generative", + model=AWS_POLLY_MODEL, aws_access_key_id="os.environ/AWS_ACCESS_KEY_ID", aws_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY", aws_region_name="os.environ/AWS_REGION", diff --git a/tests/e2e/llm_translation/test_audio_transcriptions_e2e.py b/tests/e2e/llm_translation/test_audio_transcriptions_e2e.py index 725e15a0209..cd1ee8fb03c 100644 --- a/tests/e2e/llm_translation/test_audio_transcriptions_e2e.py +++ b/tests/e2e/llm_translation/test_audio_transcriptions_e2e.py @@ -17,6 +17,7 @@ from typing import Final import pytest from e2e_config import unique_marker +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from e2e_http import UnknownApiError, unwrap from lifecycle import ResourceManager from models import LiteLLMParamsBody @@ -30,6 +31,8 @@ WEATHER_WAV = ( Path(__file__).resolve().parent / "realtime" / "fixtures" / "weather_question_24k.wav" ) +OPENAI_TRANSCRIBE_MODEL: Final = "openai/gpt-4o-mini-transcribe" +OPENAI_WHISPER_MODEL: Final = "openai/whisper-1" MISSING_MODEL_PHRASES: Final = ("model=none", "invalid model", "model is required") @@ -47,7 +50,7 @@ def _register(proxy: ProxyClient, resources: ResourceManager) -> tuple[str, str] model_id = proxy.create_model( model, LiteLLMParamsBody( - model="openai/gpt-4o-mini-transcribe", api_key="os.environ/OPENAI_API_KEY" + model=OPENAI_TRANSCRIBE_MODEL, api_key="os.environ/OPENAI_API_KEY" ), ) resources.defer(lambda: proxy.delete_model(model_id)) @@ -56,6 +59,15 @@ def _register(proxy: ProxyClient, resources: ResourceManager) -> tuple[str, str] class TestAudioTranscriptions: @pytest.mark.covers("llm.audio_transcriptions.openai.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.AUDIO, + providers=(Provider.OPENAI,), + models=(OPENAI_TRANSCRIBE_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_audio_transcriptions_returns_text( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -72,6 +84,15 @@ class TestAudioTranscriptions: ) @pytest.mark.covers("llm.audio_transcriptions.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.AUDIO, + providers=(Provider.OPENAI,), + models=(OPENAI_TRANSCRIBE_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_missing_file_returns_error( self, proxy: ProxyClient, resources: ResourceManager ) -> None: @@ -98,6 +119,12 @@ class TestAudioTranscriptions: pytest.fail(f"empty audio expected a file-specific 400, got {other!r}") @pytest.mark.covers("llm.audio_transcriptions.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.AUDIO, + ) + ) def test_missing_model_returns_error( self, proxy: ProxyClient, resources: ResourceManager ) -> None: @@ -143,7 +170,7 @@ class TestWhisperTranscriptionFormats: self, proxy: ProxyClient, resources: ResourceManager, form: _WhisperForm, response_type: type[R] ) -> R: model_id = proxy.create_model( - form.model, LiteLLMParamsBody(model="openai/whisper-1", api_key="os.environ/OPENAI_API_KEY") + form.model, LiteLLMParamsBody(model=OPENAI_WHISPER_MODEL, api_key="os.environ/OPENAI_API_KEY") ) resources.defer(lambda: proxy.delete_model(model_id)) return unwrap( @@ -158,12 +185,30 @@ class TestWhisperTranscriptionFormats: ) ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.AUDIO, + providers=(Provider.OPENAI,), + models=(OPENAI_WHISPER_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_vtt_format_returns_webvtt_transcript(self, proxy: ProxyClient, resources: ResourceManager) -> None: form = _WhisperForm(model=f"e2e-whisper-vtt-{unique_marker()}", response_format="vtt") transcript = self._upload(proxy, resources, form, _TranscriptionResult) assert transcript.text.lstrip().startswith("WEBVTT"), f"vtt transcript is not WebVTT: {transcript.text[:200]!r}" assert "weather" in transcript.text.lower(), f"vtt transcript lost the spoken words: {transcript.text!r}" + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.AUDIO, + providers=(Provider.OPENAI,), + models=(OPENAI_WHISPER_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_verbose_json_returns_word_timestamps(self, proxy: ProxyClient, resources: ResourceManager) -> None: form = _WhisperForm( model=f"e2e-whisper-verbose-{unique_marker()}", diff --git a/tests/e2e/llm_translation/test_bedrock_native_e2e.py b/tests/e2e/llm_translation/test_bedrock_native_e2e.py index 19c1be7b6db..a9b9e45858d 100644 --- a/tests/e2e/llm_translation/test_bedrock_native_e2e.py +++ b/tests/e2e/llm_translation/test_bedrock_native_e2e.py @@ -8,6 +8,7 @@ from __future__ import annotations import pytest from e2e_config import unique_marker +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from e2e_http import ( assert_client_error, require_successful_call, @@ -100,6 +101,15 @@ def _default_invoke() -> InvokeBody: class TestBedrockNative: @pytest.mark.covers("llm.bedrock_native.bedrock_converse.basic.nonstream.works") + @meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.BEDROCK,), + models=(BEDROCK_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_converse_returns_assistant(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, key = _register(proxy, resources) result = proxy.transport.send( @@ -113,6 +123,15 @@ class TestBedrockNative: assert any(part.text.strip() for part in response.output.message.content) @pytest.mark.covers("llm.bedrock_native.bedrock_converse.basic.stream.works") + @meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.BEDROCK,), + models=(BEDROCK_BACKEND,), + mode=Mode.STREAM, + ) + ) def test_converse_stream_returns_chunks(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, key = _register(proxy, resources) result = proxy.transport.send( @@ -126,6 +145,15 @@ class TestBedrockNative: assert result.chunks > 0, "converse-stream returned no events" @pytest.mark.covers("llm.bedrock_native.bedrock_invoke.basic.nonstream.works") + @meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.BEDROCK,), + models=(BEDROCK_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_invoke_returns_message(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, key = _register(proxy, resources) result = proxy.transport.send( @@ -138,6 +166,15 @@ class TestBedrockNative: assert any(part.text.strip() for part in response.content) @pytest.mark.covers("llm.bedrock_native.bedrock_invoke.basic.stream.works") + @meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.BEDROCK,), + models=(BEDROCK_BACKEND,), + mode=Mode.STREAM, + ) + ) def test_invoke_stream_returns_chunks(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, key = _register(proxy, resources) result = proxy.transport.send( @@ -151,6 +188,15 @@ class TestBedrockNative: assert result.chunks > 0, "invoke stream returned no events" @pytest.mark.covers("llm.bedrock_native.bedrock_converse.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.BEDROCK,), + models=(BEDROCK_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_converse_missing_messages_returns_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, key = _register(proxy, resources) result = proxy.transport.send( @@ -161,6 +207,15 @@ class TestBedrockNative: assert_client_error(result, "converse missing messages") @pytest.mark.covers("llm.bedrock_native.bedrock_converse.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.BEDROCK,), + models=(BEDROCK_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_converse_empty_messages_returns_client_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, key = _register(proxy, resources) result = proxy.transport.send( @@ -171,6 +226,14 @@ class TestBedrockNative: assert_client_error(result, "converse empty messages") @pytest.mark.covers("llm.bedrock_native.bedrock_converse.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.BEDROCK,), + mode=Mode.NONSTREAM, + ) + ) def test_converse_invalid_model_returns_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: _, key = _register(proxy, resources) result = proxy.transport.send( @@ -183,6 +246,15 @@ class TestBedrockNative: ) @pytest.mark.covers("llm.bedrock_native.bedrock_invoke.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.BEDROCK,), + models=(BEDROCK_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_invoke_missing_messages_returns_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, key = _register(proxy, resources) result = proxy.transport.send( @@ -193,6 +265,15 @@ class TestBedrockNative: assert_client_error(result, "invoke missing messages") @pytest.mark.covers("llm.bedrock_native.bedrock_invoke.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.BEDROCK,), + models=(BEDROCK_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_invoke_missing_max_tokens_returns_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, key = _register(proxy, resources) result = proxy.transport.send( @@ -206,6 +287,15 @@ class TestBedrockNative: assert_client_error(result, "invoke missing max_tokens") @pytest.mark.covers("llm.bedrock_native.bedrock_invoke.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.BEDROCK,), + models=(BEDROCK_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_invoke_invalid_temperature_returns_client_error( self, proxy: ProxyClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/llm_translation/test_bedrock_provider_matrix_e2e.py b/tests/e2e/llm_translation/test_bedrock_provider_matrix_e2e.py index 21333d39849..c80609edcdd 100644 --- a/tests/e2e/llm_translation/test_bedrock_provider_matrix_e2e.py +++ b/tests/e2e/llm_translation/test_bedrock_provider_matrix_e2e.py @@ -18,6 +18,7 @@ import pytest from pydantic import BaseModel from e2e_config import unique_marker +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from e2e_http import StreamingResponse, unwrap from lifecycle import ResourceManager from models import ChatBody, ChatMessage, ChatResponse, LiteLLMParamsBody @@ -98,6 +99,15 @@ class TestBedrockResponseHeaders: "llm.chat_completions.bedrock_converse.response_headers.nonstream.works", exercised_on=[], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.BEDROCK,), + models=(CONVERSE_REGIONAL_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_bedrock_request_id_header_surfaces( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -117,6 +127,15 @@ class TestBedrockResponseHeaders: "llm.chat_completions.bedrock_converse.response_headers.stream.works", exercised_on=[], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.BEDROCK,), + models=(CONVERSE_REGIONAL_BACKEND,), + mode=Mode.STREAM, + ) + ) def test_bedrock_request_id_header_surfaces_on_stream( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -158,6 +177,15 @@ class TestBedrockBatchDeploymentServesChat: "llm.chat_completions.bedrock_converse.batch_deployment.nonstream.works", exercised_on=[], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.BEDROCK,), + models=(CONVERSE_REGIONAL_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_batch_s3_keys_do_not_break_chat( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -179,6 +207,15 @@ class TestBedrockBatchDeploymentServesChat: class TestBedrockInvokeRegionalModelIds: @pytest.mark.covers("llm.chat_completions.bedrock_invoke.basic.nonstream.works", exercised_on=[]) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.BEDROCK,), + models=(INVOKE_REGIONAL_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_invoke_regional_id_completes( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -190,6 +227,15 @@ class TestBedrockInvokeRegionalModelIds: _assert_completion(response) @pytest.mark.covers("llm.chat_completions.bedrock_invoke.basic.stream.works", exercised_on=[]) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.BEDROCK,), + models=(INVOKE_REGIONAL_BACKEND,), + mode=Mode.STREAM, + ) + ) def test_invoke_regional_id_streams( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -205,6 +251,15 @@ class TestBedrockInvokeRegionalModelIds: class TestBedrockOpenAIFamilyDefaultRoute: @pytest.mark.covers("llm.chat_completions.bedrock_converse.basic.nonstream.works", exercised_on=[]) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.BEDROCK,), + models=(OPENAI_FAMILY_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_openai_family_model_id_completes_with_max_tokens( self, client: PassthroughClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/llm_translation/test_bedrock_web_search_server_tool_e2e.py b/tests/e2e/llm_translation/test_bedrock_web_search_server_tool_e2e.py index b4253a82dd8..6133fa960ff 100644 --- a/tests/e2e/llm_translation/test_bedrock_web_search_server_tool_e2e.py +++ b/tests/e2e/llm_translation/test_bedrock_web_search_server_tool_e2e.py @@ -36,6 +36,7 @@ from __future__ import annotations import pytest from anthropic.types import WebSearchTool20250305Param from e2e_config import unique_marker +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import LiteLLMParamsBody from proxy_client import ProxyClient @@ -62,6 +63,16 @@ class TestBedrockWebSearchServerTool: "ephemeral stack ships the config in this module's docstring." ) @pytest.mark.covers("llm.messages.bedrock_invoke.web_search_server_tool.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=(BEDROCK_INVOKE_BACKEND,), + capabilities=(Capability.WEB_SEARCH,), + mode=Mode.NONSTREAM, + ) + ) def test_web_search_server_tool_is_served( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: diff --git a/tests/e2e/llm_translation/test_cache_control.py b/tests/e2e/llm_translation/test_cache_control.py index bec65144c9f..33c52efe37e 100644 --- a/tests/e2e/llm_translation/test_cache_control.py +++ b/tests/e2e/llm_translation/test_cache_control.py @@ -43,6 +43,7 @@ import pytest from pydantic import BaseModel from e2e_config import unique_marker +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from e2e_http import Result, UnknownApiError, unwrap from lifecycle import ResourceManager from models import CacheControl, ChatBody, ChatMessage, ChatResponse, LiteLLMParamsBody, RichMessage, TextBlock, Usage @@ -226,6 +227,16 @@ class TestCacheControl: "llm.chat_completions.bedrock_converse.prompt_cache_5m.nonstream.works", exercised_on=[], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.BEDROCK,), + models=(BEDROCK_MODEL,), + capabilities=(Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) + ) def test_bedrock_prompt_caching_reads_cache( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -242,6 +253,16 @@ class TestCacheControl: "llm.chat_completions.vertex.prompt_cache_5m.nonstream.works", exercised_on=[], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.VERTEX_AI,), + models=(VERTEX_MODEL,), + capabilities=(Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) + ) def test_vertex_prompt_caching_reads_cache( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -266,6 +287,16 @@ class TestCacheControl: "llm.chat_completions.anthropic.prompt_cache_5m.nonstream.works", exercised_on=[], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_MODEL,), + capabilities=(Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) + ) def test_anthropic_prompt_caching_reads_cache( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -282,6 +313,16 @@ class TestCacheControl: "llm.chat_completions.openai.prompt_cache_5m.nonstream.works", exercised_on=[], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(OPENAI_MODEL,), + capabilities=(Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) + ) def test_openai_prompt_caching_reads_cache( self, client: PassthroughClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/llm_translation/test_cache_control_injection_tool_calls_e2e.py b/tests/e2e/llm_translation/test_cache_control_injection_tool_calls_e2e.py index 1481185a601..0d796a32ef3 100644 --- a/tests/e2e/llm_translation/test_cache_control_injection_tool_calls_e2e.py +++ b/tests/e2e/llm_translation/test_cache_control_injection_tool_calls_e2e.py @@ -5,6 +5,7 @@ from typing import Final, Literal, TypeAlias import pytest from e2e_config import unique_marker from e2e_http import Result, unwrap +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import ( CacheControl, @@ -162,7 +163,36 @@ def _assert_normal_completion(response: ChatResponse, model_name: str) -> None: @pytest.mark.parametrize( "backend", - (pytest.param("azure_foundry", id="azure-foundry"), pytest.param("vertex", id="vertex")), + ( + pytest.param( + "azure_foundry", + id="azure-foundry", + marks=meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.AZURE_AI,), + models=(AZURE_MODEL,), + capabilities=(Capability.FUNCTION_CALLING, Capability.PROMPT_CACHING), + mode=Mode.NONSTREAM, + ) + ), + ), + pytest.param( + "vertex", + id="vertex", + marks=meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.VERTEX_AI,), + models=(VERTEX_MODEL,), + capabilities=(Capability.FUNCTION_CALLING, Capability.PROMPT_CACHING), + mode=Mode.NONSTREAM, + ) + ), + ), + ), ) @pytest.mark.provider_live @pytest.mark.covers("llm.chat_completions.azure_foundry.basic.nonstream.works") diff --git a/tests/e2e/llm_translation/test_chat_completions_contract_e2e.py b/tests/e2e/llm_translation/test_chat_completions_contract_e2e.py index 09b484eb120..45b9fa44e52 100644 --- a/tests/e2e/llm_translation/test_chat_completions_contract_e2e.py +++ b/tests/e2e/llm_translation/test_chat_completions_contract_e2e.py @@ -8,6 +8,7 @@ from __future__ import annotations import pytest from e2e_config import provider_edge_base, unique_marker from e2e_http import StreamingResponse, assert_client_error, require_successful_call, unwrap +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import ChatBody, ChatMessage, ChatResponse, LiteLLMParamsBody from proxy_client import ProxyClient @@ -62,6 +63,15 @@ def _chat_status(proxy: ProxyClient, key: str, body: BaseModel) -> StreamingResp class TestChatCompletionsContract: @pytest.mark.covers("llm.chat_completions.openai.multi_turn.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(OPENAI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_multi_turn_history_is_honored(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, key = _register_chat_model(proxy, resources) turn1 = unwrap( @@ -106,6 +116,15 @@ class TestChatCompletionsContract: assert "84" in second, f"turn2 must answer 84 from history, got: {second!r}" @pytest.mark.covers("llm.chat_completions.openai.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(OPENAI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_success_response_matches_chat_completion_contract( self, proxy: ProxyClient, resources: ResourceManager ) -> None: @@ -131,6 +150,12 @@ class TestChatCompletionsContract: assert (message.content or "").strip(), f"content must be non-empty: {result.body[:300]}" @pytest.mark.covers("llm.chat_completions.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + ) + ) def test_missing_model_returns_client_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: _, key = _register_chat_model(proxy, resources) result = _chat_status( @@ -145,12 +170,30 @@ class TestChatCompletionsContract: ) @pytest.mark.covers("llm.chat_completions.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(OPENAI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_missing_messages_returns_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, key = _register_chat_model(proxy, resources) result = _chat_status(proxy, key, ChatMissingMessagesBody(model=model)) assert_client_error(result, "missing messages") @pytest.mark.covers("llm.chat_completions.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(OPENAI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_empty_messages_returns_client_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, key = _register_chat_model(proxy, resources) result = _chat_status( @@ -161,6 +204,15 @@ class TestChatCompletionsContract: assert_client_error(result, "empty messages") @pytest.mark.covers("llm.chat_completions.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(OPENAI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_invalid_role_returns_client_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, key = _register_chat_model(proxy, resources) result = _chat_status( @@ -175,6 +227,15 @@ class TestChatCompletionsContract: assert_client_error(result, "invalid role") @pytest.mark.covers("llm.chat_completions.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(OPENAI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_invalid_temperatures_return_client_errors(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, key = _register_chat_model(proxy, resources) for temperature in (-0.1, 2.1, 3.0, 100.0): @@ -191,6 +252,15 @@ class TestChatCompletionsContract: assert_client_error(result, f"temperature={temperature}") @pytest.mark.covers("llm.chat_completions.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(OPENAI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_invalid_max_completion_tokens_return_client_errors( self, proxy: ProxyClient, resources: ResourceManager ) -> None: @@ -208,6 +278,15 @@ class TestChatCompletionsContract: assert_client_error(result, f"max_completion_tokens={max_completion_tokens}") @pytest.mark.covers("llm.chat_completions.openai.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(OPENAI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_temperature_boundaries_succeed(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, key = _register_chat_model(proxy, resources) for temperature in (0.0, 2.0): diff --git a/tests/e2e/llm_translation/test_chat_completions_regression_e2e.py b/tests/e2e/llm_translation/test_chat_completions_regression_e2e.py index 90ac18c16ac..1e0a8cbb412 100644 --- a/tests/e2e/llm_translation/test_chat_completions_regression_e2e.py +++ b/tests/e2e/llm_translation/test_chat_completions_regression_e2e.py @@ -26,6 +26,7 @@ from pydantic import BaseModel from e2e_config import unique_marker from e2e_http import StreamingResponse, unwrap +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import ( ChatBody, @@ -53,16 +54,40 @@ OPENAI_BACKEND = "openai/gpt-5.6" ANTHROPIC_BACKEND = "anthropic/claude-haiku-4-5-20251001" BEDROCK_CONVERSE_BACKEND = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0" BEDROCK_NOVA_BACKEND: Final = "bedrock/us.amazon.nova-2-lite-v1:0" +VERTEX_MISTRAL_BACKEND: Final = "vertex_ai/mistral-small-2503" +VERTEX_GPT_OSS_BACKEND: Final = "vertex_ai/openai/gpt-oss-120b-maas" VERTEX_PARTNER_BACKENDS: Final = ( pytest.param( - "vertex_ai/mistral-small-2503", - marks=pytest.mark.skip( - reason="the e2e Vertex project has no access to mistral-small-2503 (404 publisher model not found)" + VERTEX_MISTRAL_BACKEND, + marks=( + pytest.mark.skip( + reason="the e2e Vertex project has no access to mistral-small-2503 (404 publisher model not found)" + ), + meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.VERTEX_AI,), + models=(VERTEX_MISTRAL_BACKEND,), + mode=Mode.STREAM, + ) + ), ), ), pytest.param( - "vertex_ai/openai/gpt-oss-120b-maas", - marks=pytest.mark.skip(reason="never served by the e2e Vertex project (60s read timeout, no headers)"), + VERTEX_GPT_OSS_BACKEND, + marks=( + pytest.mark.skip(reason="never served by the e2e Vertex project (60s read timeout, no headers)"), + meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.VERTEX_AI,), + models=(VERTEX_GPT_OSS_BACKEND,), + mode=Mode.STREAM, + ) + ), + ), ), ) PDF_DOCUMENT_URL: Final = ( @@ -216,19 +241,58 @@ _PERSON_SCHEMA: dict[str, object] = { }, } -CHAT_MODELS: tuple[tuple[str, str], ...] = ( - ("gpt-5.5", "openai"), - ("claude-haiku-4-5", "anthropic"), - ("gemini-2.5-flash", "gemini"), +OPENAI_CHAT_MODEL: Final = "gpt-5.5" +ANTHROPIC_CHAT_MODEL: Final = "claude-haiku-4-5" +GEMINI_FLASH_MODEL: Final = "gemini-2.5-flash" + +CHAT_MODELS: Final = ( + pytest.param( + OPENAI_CHAT_MODEL, + "openai", + id=f"{OPENAI_CHAT_MODEL}-openai", + marks=meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(OPENAI_CHAT_MODEL,), + mode=Mode.NONSTREAM, + ) + ), + ), + pytest.param( + ANTHROPIC_CHAT_MODEL, + "anthropic", + id=f"{ANTHROPIC_CHAT_MODEL}-anthropic", + marks=meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_CHAT_MODEL,), + mode=Mode.NONSTREAM, + ) + ), + ), + pytest.param( + GEMINI_FLASH_MODEL, + "gemini", + id=f"{GEMINI_FLASH_MODEL}-gemini", + marks=meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.GEMINI,), + models=(GEMINI_FLASH_MODEL,), + mode=Mode.NONSTREAM, + ) + ), + ), ) class TestChatCompletionsRegression: - @pytest.mark.parametrize( - ("model", "route"), - CHAT_MODELS, - ids=[f"{model}-{route}" for model, route in CHAT_MODELS], - ) + @pytest.mark.parametrize(("model", "route"), CHAT_MODELS) @pytest.mark.covers( "llm.chat_completions.openai.basic.nonstream.works", "llm.chat_completions.anthropic.basic.nonstream.works", @@ -272,6 +336,15 @@ class TestCohereChat: "llm.chat_completions.cohere.basic.nonstream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.COHERE,), + models=(COHERE_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_cohere_chat_returns_content( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -316,6 +389,15 @@ class TestGeminiChatCompletions: "llm.chat_completions.gemini.basic.nonstream.cost_logged", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.GEMINI,), + models=(GEMINI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_gemini_chat_returns_content_and_logs_cost( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -377,6 +459,15 @@ class TestVertexChatCompletions: "llm.chat_completions.vertex.basic.nonstream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.VERTEX_AI,), + models=(VERTEX_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_vertex_chat_returns_content( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -406,6 +497,16 @@ class TestVertexChatCompletions: "llm.chat_completions.vertex.tool_use.nonstream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.VERTEX_AI,), + models=(VERTEX_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_vertex_chat_returns_tool_call( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -435,6 +536,16 @@ class TestVertexChatCompletions: "llm.chat_completions.vertex.vision.nonstream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.VERTEX_AI,), + models=(VERTEX_BACKEND,), + capabilities=(Capability.VISION,), + mode=Mode.NONSTREAM, + ) + ) def test_vertex_chat_vision_describes_image( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -452,6 +563,15 @@ class TestVertexChatCompletions: "llm.chat_completions.vertex.basic.stream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.VERTEX_AI,), + models=(VERTEX_BACKEND,), + mode=Mode.STREAM, + ) + ) def test_vertex_chat_streams_real_content( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -494,6 +614,15 @@ class TestAzureOpenAIChatCompletions: "llm.chat_completions.azure_openai.basic.nonstream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.AZURE,), + models=(AZURE_OPENAI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_azure_openai_chat_returns_content( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -523,6 +652,16 @@ class TestAzureOpenAIChatCompletions: "llm.chat_completions.azure_openai.tool_use.nonstream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.AZURE,), + models=(AZURE_OPENAI_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_azure_openai_chat_returns_tool_call( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -554,6 +693,15 @@ class TestAzureFoundryChatCompletions: "llm.chat_completions.azure_foundry.basic.nonstream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.AZURE_AI,), + models=(AZURE_FOUNDRY_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_azure_foundry_chat_returns_content( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -596,6 +744,14 @@ class TestHostedVllmChat: "llm.chat_completions.hosted_vllm.basic.nonstream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.HOSTED_VLLM,), + mode=Mode.NONSTREAM, + ) + ) def test_hosted_vllm_chat_returns_content( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -651,6 +807,15 @@ class TestOpenAIChatCompletions: "llm.chat_completions.openai.basic.stream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(OPENAI_BACKEND,), + mode=Mode.STREAM, + ) + ) def test_openai_chat_streams_real_content( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -678,6 +843,15 @@ class TestOpenAIChatCompletions: "llm.chat_completions.openai.basic.nonstream.cost_logged", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.SPEND_BUDGETS, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(OPENAI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_openai_chat_logs_cost( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -711,6 +885,16 @@ class TestOpenAIChatCompletions: "llm.chat_completions.openai.tool_use.nonstream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(OPENAI_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_openai_chat_returns_tool_call( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -742,6 +926,16 @@ class TestOpenAIChatCompletions: "llm.chat_completions.openai.structured_output.nonstream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(OPENAI_BACKEND,), + capabilities=(Capability.RESPONSE_SCHEMA,), + mode=Mode.NONSTREAM, + ) + ) def test_openai_chat_structured_output_conforms_to_schema( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -775,6 +969,16 @@ class TestOpenAIChatCompletions: "llm.chat_completions.openai.thinking.nonstream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(OPENAI_BACKEND,), + capabilities=(Capability.REASONING,), + mode=Mode.NONSTREAM, + ) + ) def test_openai_chat_reasoning_reports_reasoning_tokens( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -825,6 +1029,16 @@ class TestOpenAIChatCompletions: "llm.chat_completions.openai.vision.nonstream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(OPENAI_VISION_BACKEND,), + capabilities=(Capability.VISION,), + mode=Mode.NONSTREAM, + ) + ) def test_openai_chat_vision_describes_image( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -842,6 +1056,16 @@ class TestOpenAIChatCompletions: "llm.chat_completions.openai.tool_use.stream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(OPENAI_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.STREAM, + ) + ) def test_openai_chat_streams_tool_call( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -889,6 +1113,15 @@ class TestBedrockConverseChatCompletions: "llm.chat_completions.bedrock_converse.basic.nonstream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.BEDROCK,), + models=(BEDROCK_CONVERSE_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_bedrock_converse_chat_returns_content( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -913,6 +1146,15 @@ class TestBedrockConverseChatCompletions: "llm.chat_completions.bedrock_converse.basic.stream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.BEDROCK,), + models=(BEDROCK_CONVERSE_BACKEND,), + mode=Mode.STREAM, + ) + ) def test_bedrock_converse_chat_streams_real_content( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -936,6 +1178,16 @@ class TestBedrockConverseChatCompletions: "llm.chat_completions.bedrock_converse.tool_use.nonstream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.BEDROCK,), + models=(BEDROCK_CONVERSE_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_bedrock_converse_chat_returns_tool_call( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -962,6 +1214,16 @@ class TestBedrockConverseChatCompletions: "llm.chat_completions.bedrock_converse.thinking.nonstream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.BEDROCK,), + models=(BEDROCK_CONVERSE_BACKEND,), + capabilities=(Capability.REASONING,), + mode=Mode.NONSTREAM, + ) + ) def test_bedrock_converse_chat_returns_reasoning( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -992,6 +1254,16 @@ class TestBedrockConverseChatCompletions: "llm.chat_completions.bedrock_converse.vision.nonstream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.BEDROCK,), + models=(BEDROCK_CONVERSE_BACKEND,), + capabilities=(Capability.VISION,), + mode=Mode.NONSTREAM, + ) + ) def test_bedrock_converse_chat_vision_describes_image( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -1001,6 +1273,16 @@ class TestBedrockConverseChatCompletions: response = unwrap(client.proxy.chat(key, ChatBody(model=model, messages=_vision_messages(), max_tokens=32))) _assert_describes_cat(response) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.BEDROCK,), + models=(BEDROCK_NOVA_BACKEND,), + capabilities=(Capability.PDF_INPUT,), + mode=Mode.NONSTREAM, + ) + ) def test_bedrock_converse_reads_a_pdf_sent_by_url( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -1107,6 +1389,16 @@ class TestAnthropicChatCompletions: "llm.chat_completions.anthropic.structured_output.nonstream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_BACKEND,), + capabilities=(Capability.RESPONSE_SCHEMA,), + mode=Mode.NONSTREAM, + ) + ) def test_anthropic_chat_structured_output_conforms_to_schema( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -1140,6 +1432,16 @@ class TestAnthropicChatCompletions: "llm.chat_completions.anthropic.thinking.nonstream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_BACKEND,), + capabilities=(Capability.REASONING,), + mode=Mode.NONSTREAM, + ) + ) def test_anthropic_chat_returns_thinking_content( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -1178,6 +1480,16 @@ class TestAnthropicChatCompletions: "llm.chat_completions.anthropic.vision.nonstream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_BACKEND,), + capabilities=(Capability.VISION,), + mode=Mode.NONSTREAM, + ) + ) def test_anthropic_chat_vision_describes_image( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -1191,6 +1503,15 @@ class TestAnthropicChatCompletions: "llm.chat_completions.anthropic.basic.stream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_BACKEND,), + mode=Mode.STREAM, + ) + ) def test_anthropic_chat_streams_real_content( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -1214,6 +1535,16 @@ class TestAnthropicChatCompletions: "llm.chat_completions.anthropic.tool_use.nonstream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_anthropic_chat_returns_tool_call( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -1240,6 +1571,16 @@ class TestAnthropicChatCompletions: "llm.chat_completions.anthropic.tool_use.stream.works", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.STREAM, + ) + ) def test_anthropic_chat_streams_tool_call( self, client: PassthroughClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/llm_translation/test_chat_mid_conversation_system_e2e.py b/tests/e2e/llm_translation/test_chat_mid_conversation_system_e2e.py index 480225b502e..71e716b2255 100644 --- a/tests/e2e/llm_translation/test_chat_mid_conversation_system_e2e.py +++ b/tests/e2e/llm_translation/test_chat_mid_conversation_system_e2e.py @@ -37,6 +37,7 @@ from pydantic import BaseModel from e2e_config import unique_marker from e2e_http import Result, unwrap +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import CacheControl, ChatResponse, LiteLLMParamsBody, RichMessage, TextBlock, Usage from passthrough_client import PassthroughClient @@ -281,6 +282,16 @@ class TestAnthropicChatMidConversationSystem: "llm.chat_completions.anthropic.mid_conversation_system.nonstream.cache_hit", exercised_on=[], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(FLAGGED_MODEL,), + capabilities=(Capability.MID_CONVERSATION_SYSTEM, Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) + ) def test_flagged_model_keeps_prompt_cache_across_system_reminder( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -290,6 +301,16 @@ class TestAnthropicChatMidConversationSystem: "llm.chat_completions.anthropic.mid_conversation_system.nonstream.works", exercised_on=[], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(UNFLAGGED_MODEL,), + capabilities=(Capability.MID_CONVERSATION_SYSTEM, Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) + ) def test_unflagged_model_converts_system_reminder_and_succeeds( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -305,6 +326,16 @@ class TestBedrockInvokeChatMidConversationSystem: "llm.chat_completions.bedrock_invoke.mid_conversation_system.nonstream.cache_hit", exercised_on=[], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.BEDROCK,), + models=(FLAGGED_MODEL,), + capabilities=(Capability.MID_CONVERSATION_SYSTEM, Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) + ) def test_flagged_model_keeps_prompt_cache_across_system_reminder( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -314,6 +345,16 @@ class TestBedrockInvokeChatMidConversationSystem: "llm.chat_completions.bedrock_invoke.mid_conversation_system.nonstream.works", exercised_on=[], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.BEDROCK,), + models=(UNFLAGGED_MODEL,), + capabilities=(Capability.MID_CONVERSATION_SYSTEM, Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) + ) def test_unflagged_model_converts_system_reminder_and_succeeds( self, client: PassthroughClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/llm_translation/test_chat_stream_contract_e2e.py b/tests/e2e/llm_translation/test_chat_stream_contract_e2e.py index fdb76df703d..1e353d11283 100644 --- a/tests/e2e/llm_translation/test_chat_stream_contract_e2e.py +++ b/tests/e2e/llm_translation/test_chat_stream_contract_e2e.py @@ -5,6 +5,7 @@ from typing import Final import pytest from e2e_config import provider_edge_base, unique_marker from e2e_http import require_successful_call +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import ChatBody, ChatMessage, ChatStreamOptions, LiteLLMParamsBody, Usage from proxy_client import ProxyClient @@ -12,6 +13,8 @@ from pydantic import BaseModel pytestmark = [pytest.mark.e2e, pytest.mark.replayable] +OPENAI_BACKEND: Final = "openai/gpt-5.6" + class _Delta(BaseModel): content: str | None = None @@ -30,13 +33,22 @@ class _Chunk(BaseModel): class TestChatStreamContract: @pytest.mark.covers("llm.chat_completions.openai.basic.stream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(OPENAI_BACKEND,), + mode=Mode.STREAM, + ) + ) def test_chat_stream_is_sse_and_ends_with_done(self, proxy: ProxyClient, resources: ResourceManager) -> None: model: Final = f"e2e-chat-stream-{unique_marker()}" base: Final = provider_edge_base("openai") model_id: Final = proxy.create_model( model, LiteLLMParamsBody( - model="openai/gpt-5.6", + model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY", api_base=f"{base}/v1" if base else None, ), diff --git a/tests/e2e/llm_translation/test_chat_tool_round_trip_e2e.py b/tests/e2e/llm_translation/test_chat_tool_round_trip_e2e.py index 75ba5c23ff8..01d15aed443 100644 --- a/tests/e2e/llm_translation/test_chat_tool_round_trip_e2e.py +++ b/tests/e2e/llm_translation/test_chat_tool_round_trip_e2e.py @@ -6,6 +6,7 @@ from typing import Final import pytest from e2e_config import unique_marker from e2e_http import unwrap +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import ( ChatAssistantTurn, @@ -138,22 +139,72 @@ def _assert_tool_results_reach_the_model( class TestChatToolResultRoundTrip: + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.GEMINI,), + models=(GEMINI_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_gemini(self, client: PassthroughClient, resources: ResourceManager) -> None: model, key = _register(client, resources, _api_key_params(GEMINI_BACKEND, "GEMINI_API_KEY")) _assert_tool_results_reach_the_model(client, key, model, thinking=None, tool_choice="required") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.MISTRAL,), + models=(MISTRAL_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_mistral(self, client: PassthroughClient, resources: ResourceManager) -> None: model, key = _register(client, resources, _api_key_params(MISTRAL_BACKEND, "MISTRAL_API_KEY")) _assert_tool_results_reach_the_model(client, key, model, thinking=None, tool_choice="required") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.BEDROCK,), + models=(BEDROCK_CONVERSE_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_bedrock_converse(self, client: PassthroughClient, resources: ResourceManager) -> None: model, key = _register(client, resources, _bedrock_params(BEDROCK_CONVERSE_BACKEND)) _assert_tool_results_reach_the_model(client, key, model, thinking=None, tool_choice="required") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING, Capability.REASONING,), + mode=Mode.NONSTREAM, + ) + ) def test_anthropic_with_extended_thinking(self, client: PassthroughClient, resources: ResourceManager) -> None: model, key = _register(client, resources, _api_key_params(ANTHROPIC_BACKEND, "ANTHROPIC_API_KEY")) _assert_tool_results_reach_the_model(client, key, model, thinking=THINKING, tool_choice=None) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.BEDROCK,), + models=(BEDROCK_LEGACY_THINKING_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING, Capability.REASONING,), + mode=Mode.NONSTREAM, + ) + ) def test_bedrock_converse_with_extended_thinking( self, client: PassthroughClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/llm_translation/test_completions_endpoint_e2e.py b/tests/e2e/llm_translation/test_completions_endpoint_e2e.py index 6202dada599..7f89d8a7ed8 100644 --- a/tests/e2e/llm_translation/test_completions_endpoint_e2e.py +++ b/tests/e2e/llm_translation/test_completions_endpoint_e2e.py @@ -10,8 +10,11 @@ the completion fails here. from __future__ import annotations +from typing import Final + import pytest from e2e_config import unique_marker +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import LiteLLMParamsBody from proxy_client import ProxyClient @@ -19,9 +22,20 @@ from sdk_clients import NO_PROXY_CACHE, SdkClients pytestmark = pytest.mark.e2e +OPENAI_COMPLETIONS_BACKEND: Final = "openai/gpt-5.4-nano" + class TestCompletionsEndpoint: @pytest.mark.covers("llm.completions.openai.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.COMPLETIONS, + providers=(Provider.OPENAI,), + models=(OPENAI_COMPLETIONS_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_text_completion_returns_text( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -29,7 +43,7 @@ class TestCompletionsEndpoint: model_id = proxy.create_model( model, LiteLLMParamsBody( - model="openai/gpt-5.4-nano", + model=OPENAI_COMPLETIONS_BACKEND, api_key="os.environ/OPENAI_API_KEY", ), ) diff --git a/tests/e2e/llm_translation/test_containers_e2e.py b/tests/e2e/llm_translation/test_containers_e2e.py index 1c3e37ec8bb..0a494d11108 100644 --- a/tests/e2e/llm_translation/test_containers_e2e.py +++ b/tests/e2e/llm_translation/test_containers_e2e.py @@ -50,6 +50,7 @@ import openai import pytest from e2e_config import REQUEST_TIMEOUT, unique_marker from e2e_http import unwrap +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from management.management_client import ManagementClient, build_client from models import KeyGenerateBody, KeyGenerateResponse, LiteLLMParamsBody, TeamNewBody, UserNewBody @@ -172,6 +173,15 @@ def _assert_file_round_trip(client: OpenAI, native_id: str, marker: str) -> None class TestAzureContainerFiles: @pytest.mark.covers("llm.responses.azure_openai.code_interpreter.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CONTAINERS, + providers=(Provider.AZURE,), + models=(AZURE_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_service_account_key_reads_container_file_by_native_id( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -187,6 +197,15 @@ class TestAzureContainerFiles: _assert_file_round_trip(client, native_id, marker) @pytest.mark.covers("llm.responses.azure_openai.code_interpreter.stream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CONTAINERS, + providers=(Provider.AZURE,), + models=(AZURE_BACKEND,), + mode=Mode.STREAM, + ) + ) def test_service_account_key_reads_container_file_created_by_a_streamed_response( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -202,6 +221,13 @@ class TestAzureContainerFiles: class TestOpenAIContainerFiles: + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CONTAINERS, + providers=(Provider.OPENAI,), + ) + ) def test_container_file_lifecycle_through_the_gateway(self, resources: ResourceManager, sdk: SdkClients) -> None: client: Final = sdk.openai(resources.key()) marker: Final = unique_marker() diff --git a/tests/e2e/llm_translation/test_credential_messages_e2e.py b/tests/e2e/llm_translation/test_credential_messages_e2e.py index 58e17f20bb8..55fd1cf7f2f 100644 --- a/tests/e2e/llm_translation/test_credential_messages_e2e.py +++ b/tests/e2e/llm_translation/test_credential_messages_e2e.py @@ -3,10 +3,12 @@ from __future__ import annotations import os +from typing import Final import pytest from e2e_config import unique_marker +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import CredentialCreateBody, LiteLLMParamsBody from proxy_client import ProxyClient @@ -14,9 +16,20 @@ from sdk_clients import NO_PROXY_CACHE, SdkClients pytestmark = pytest.mark.e2e +CLAUDE_BACKEND: Final = "anthropic/claude-haiku-4-5" + class TestCredentialBackedMessages: @pytest.mark.covers("mgmt.credential.new.serves_request") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=(CLAUDE_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_credential_backed_messages(self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients) -> None: marker = unique_marker() credential_name = f"e2e-cred-{marker}" @@ -35,7 +48,7 @@ class TestCredentialBackedMessages: model_id = proxy.create_model( model, LiteLLMParamsBody( - model="anthropic/claude-haiku-4-5", + model=CLAUDE_BACKEND, litellm_credential_name=credential_name, ), ) diff --git a/tests/e2e/llm_translation/test_custom_pricing_e2e.py b/tests/e2e/llm_translation/test_custom_pricing_e2e.py index 1cebf90fa21..6d8cd24648e 100644 --- a/tests/e2e/llm_translation/test_custom_pricing_e2e.py +++ b/tests/e2e/llm_translation/test_custom_pricing_e2e.py @@ -23,6 +23,7 @@ import pytest from pydantic import BaseModel, RootModel from e2e_config import unique_marker +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from proxy_client import ProxyClient from e2e_http import Success, unwrap from lifecycle import ResourceManager @@ -148,6 +149,15 @@ def _poll_breakdown_row(proxy: ProxyClient, key: str, response_id: str | None) - class TestCustomPricing: + @meta( + Subject( + domain=Domain.COST_MAP, + route=Route.SPEND_REPORTING, + providers=(Provider.GEMINI,), + models=(BACKEND_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_custom_pricing_is_billed_at_configured_rate( self, proxy: ProxyClient, @@ -193,6 +203,12 @@ class TestCustomPricing: f"= {completion * CUSTOM_OUTPUT_RATE}" ) + @meta( + Subject( + domain=Domain.COST_MAP, + route=Route.MODEL_MANAGEMENT, + ) + ) def test_model_info_reports_custom_pricing( self, proxy: ProxyClient, resources: ResourceManager ) -> None: @@ -208,6 +224,12 @@ class TestCustomPricing: f"{entry.litellm_params.output_cost_per_token} != configured {CUSTOM_OUTPUT_RATE}" ) + @meta( + Subject( + domain=Domain.COST_MAP, + route=Route.MODEL_MANAGEMENT, + ) + ) def test_custom_pricing_is_isolated_from_sibling_deployment( self, proxy: ProxyClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/llm_translation/test_deepseek_reasoning_e2e.py b/tests/e2e/llm_translation/test_deepseek_reasoning_e2e.py index 8dfccf0d74b..8008345a3b6 100644 --- a/tests/e2e/llm_translation/test_deepseek_reasoning_e2e.py +++ b/tests/e2e/llm_translation/test_deepseek_reasoning_e2e.py @@ -20,6 +20,7 @@ from __future__ import annotations import pytest from e2e_config import unique_marker +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from e2e_http import unwrap from lifecycle import ResourceManager from models import ChatBody, ChatMessage, ChatResponse, LiteLLMParamsBody, ThinkingParam @@ -49,6 +50,16 @@ def _reasoning_content(response: ChatResponse) -> str | None: class TestDeepSeekReasoningDisable: + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.DEEPSEEK,), + models=(REASONER,), + capabilities=(Capability.REASONING,), + mode=Mode.NONSTREAM, + ) + ) def test_reasoner_returns_reasoning_by_default( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -71,6 +82,16 @@ class TestDeepSeekReasoningDisable: f"disable param, so the disable assertions below can't be trusted: {response}" ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.DEEPSEEK,), + models=(REASONER,), + capabilities=(Capability.REASONING,), + mode=Mode.NONSTREAM, + ) + ) def test_reasoning_effort_none_disables_reasoning( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -93,6 +114,16 @@ class TestDeepSeekReasoningDisable: f"is still present: {response}" ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.DEEPSEEK,), + models=(REASONER,), + capabilities=(Capability.REASONING,), + mode=Mode.NONSTREAM, + ) + ) def test_thinking_disabled_disables_reasoning( self, client: PassthroughClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/llm_translation/test_embeddings_endpoint_e2e.py b/tests/e2e/llm_translation/test_embeddings_endpoint_e2e.py index 0e5bac556cc..7b59ffbf15f 100644 --- a/tests/e2e/llm_translation/test_embeddings_endpoint_e2e.py +++ b/tests/e2e/llm_translation/test_embeddings_endpoint_e2e.py @@ -16,6 +16,7 @@ from typing import Final import pytest from e2e_config import provider_edge_base, unique_marker from e2e_http import assert_client_error +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import LiteLLMParamsBody from proxy_client import ProxyClient @@ -24,6 +25,10 @@ from sdk_clients import NO_PROXY_CACHE, SdkClients, response_header pytestmark = pytest.mark.e2e +OPENAI_EMBEDDING: Final = "openai/text-embedding-3-small" +BEDROCK_TITAN_EMBEDDING: Final = "bedrock/amazon.titan-embed-text-v2:0" +COHERE_EMBEDDING: Final = "cohere/embed-v4.0" +MISTRAL_EMBEDDING: Final = "mistral/mistral-embed" VERTEX_TEXT_EMBEDDING: Final = "vertex_ai/text-embedding-005" VERTEX_MULTIMODAL_EMBEDDING: Final = "vertex_ai/multimodalembedding@001" TOKENS_TEXT: Final = "The quick brown fox jumps over the lazy dog" @@ -45,7 +50,7 @@ def _cosine(left: list[float], right: list[float]) -> float: def _titan_params() -> LiteLLMParamsBody: return LiteLLMParamsBody( - model="bedrock/amazon.titan-embed-text-v2:0", + model=BEDROCK_TITAN_EMBEDDING, aws_access_key_id="os.environ/AWS_ACCESS_KEY_ID", aws_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY", aws_region_name="os.environ/AWS_REGION", @@ -62,7 +67,7 @@ def _openai_embeddings_params() -> LiteLLMParamsBody: Vertex stay live: SigV4 signs the Host header, and neither has an edge mount.""" base = provider_edge_base("openai") return LiteLLMParamsBody( - model="openai/text-embedding-3-small", + model=OPENAI_EMBEDDING, api_key="os.environ/OPENAI_API_KEY", api_base=None if base is None else f"{base}/v1", ) @@ -97,10 +102,28 @@ def _assert_embedding_vector( class TestEmbeddingsEndpoint: @pytest.mark.replayable @pytest.mark.covers("llm.embeddings.openai.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.EMBEDDINGS, + providers=(Provider.OPENAI,), + models=(OPENAI_EMBEDDING,), + mode=Mode.NONSTREAM, + ) + ) def test_embeddings_returns_vector(self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients) -> None: _assert_embedding_vector(proxy, resources, sdk, "e2e-embeddings", _openai_embeddings_params()) @pytest.mark.covers("llm.embeddings.bedrock.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.EMBEDDINGS, + providers=(Provider.BEDROCK,), + models=(BEDROCK_TITAN_EMBEDDING,), + mode=Mode.NONSTREAM, + ) + ) def test_bedrock_embeddings_returns_vector( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -113,6 +136,15 @@ class TestEmbeddingsEndpoint: ) @pytest.mark.covers("llm.embeddings.cohere.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.EMBEDDINGS, + providers=(Provider.COHERE,), + models=(COHERE_EMBEDDING,), + mode=Mode.NONSTREAM, + ) + ) def test_cohere_embeddings_returns_vector( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -121,10 +153,19 @@ class TestEmbeddingsEndpoint: resources, sdk, "e2e-embeddings-cohere", - LiteLLMParamsBody(model="cohere/embed-v4.0", api_key="os.environ/COHERE_API_KEY"), + LiteLLMParamsBody(model=COHERE_EMBEDDING, api_key="os.environ/COHERE_API_KEY"), ) @pytest.mark.covers("llm.embeddings.vertex.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.EMBEDDINGS, + providers=(Provider.VERTEX_AI,), + models=(VERTEX_TEXT_EMBEDDING,), + mode=Mode.NONSTREAM, + ) + ) def test_vertex_embeddings_returns_vector( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -134,12 +175,21 @@ class TestEmbeddingsEndpoint: sdk, "e2e-embeddings-vertex", LiteLLMParamsBody( - model="vertex_ai/text-embedding-005", + model=VERTEX_TEXT_EMBEDDING, vertex_project="os.environ/VERTEXAI_PROJECT", vertex_location="us-central1", ), ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.EMBEDDINGS, + providers=(Provider.MISTRAL,), + models=(MISTRAL_EMBEDDING,), + mode=Mode.NONSTREAM, + ) + ) def test_mistral_embeddings_returns_vector( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -148,10 +198,19 @@ class TestEmbeddingsEndpoint: resources, sdk, "e2e-embeddings-mistral", - LiteLLMParamsBody(model="mistral/mistral-embed", api_key="os.environ/MISTRAL_API_KEY"), + LiteLLMParamsBody(model=MISTRAL_EMBEDDING, api_key="os.environ/MISTRAL_API_KEY"), ) @pytest.mark.covers("llm.embeddings.vertex.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.EMBEDDINGS, + providers=(Provider.VERTEX_AI,), + models=(VERTEX_TEXT_EMBEDDING,), + mode=Mode.NONSTREAM, + ) + ) def test_vertex_embeddings_honor_requested_dimensions( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -166,6 +225,15 @@ class TestEmbeddingsEndpoint: assert len(embeddings.data[0].embedding) == 8, f"dimensions=8 was not honored: {embeddings!r}" assert embeddings.usage.prompt_tokens > 0, f"vertex embeddings reported no prompt usage: {embeddings.usage!r}" + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.EMBEDDINGS, + providers=(Provider.VERTEX_AI,), + models=(VERTEX_MULTIMODAL_EMBEDDING,), + mode=Mode.NONSTREAM, + ) + ) def test_vertex_multimodal_embeddings_honor_dimensions_and_are_costed( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -181,6 +249,15 @@ class TestEmbeddingsEndpoint: cost = response_header(raw.headers, "x-litellm-response-cost") assert cost is not None and float(cost) > 0, f"multimodal embedding was not costed: {cost!r}" + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.EMBEDDINGS, + providers=(Provider.BEDROCK,), + models=(BEDROCK_TITAN_EMBEDDING,), + mode=Mode.NONSTREAM, + ) + ) def test_bedrock_titan_embeds_token_array_input_as_its_decoded_text( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -199,6 +276,15 @@ class TestEmbeddingsEndpoint: @pytest.mark.replayable @pytest.mark.covers("llm.embeddings.openai.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.EMBEDDINGS, + providers=(Provider.OPENAI,), + models=(OPENAI_EMBEDDING,), + mode=Mode.NONSTREAM, + ) + ) def test_array_input_returns_vectors(self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients) -> None: model, key = _register(proxy, resources, "e2e-embeddings-array", _openai_embeddings_params()) embeddings = sdk.openai(key).embeddings.create( @@ -208,6 +294,12 @@ class TestEmbeddingsEndpoint: @pytest.mark.replayable @pytest.mark.covers("llm.embeddings.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.EMBEDDINGS, + ) + ) def test_missing_model_returns_client_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: key = resources.key() result = proxy.transport.send( @@ -219,6 +311,12 @@ class TestEmbeddingsEndpoint: @pytest.mark.replayable @pytest.mark.covers("llm.embeddings.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.EMBEDDINGS, + ) + ) def test_missing_input_returns_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, key = _register(proxy, resources, "e2e-embeddings-missin", _openai_embeddings_params()) result = proxy.transport.send( diff --git a/tests/e2e/llm_translation/test_files_batches_contract_e2e.py b/tests/e2e/llm_translation/test_files_batches_contract_e2e.py index 5627fa1c0bf..5e925f222b7 100644 --- a/tests/e2e/llm_translation/test_files_batches_contract_e2e.py +++ b/tests/e2e/llm_translation/test_files_batches_contract_e2e.py @@ -8,6 +8,7 @@ from __future__ import annotations import pytest from e2e_http import NoBody, Success, UnknownApiError, assert_client_error +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from proxy_client import ProxyClient from pydantic import BaseModel @@ -28,6 +29,13 @@ class BatchObject(BaseModel): class TestFilesBatchesContract: @pytest.mark.covers("llm.files.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.FILES, + mode=Mode.BATCH, + ) + ) def test_upload_without_purpose_returns_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: key = resources.key() result = proxy.transport.upload( @@ -47,6 +55,13 @@ class TestFilesBatchesContract: pytest.fail(f"upload without purpose expected 4xx, got {other!r}") @pytest.mark.covers("llm.batches.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.BATCHES, + mode=Mode.BATCH, + ) + ) def test_create_batch_missing_input_file_id_returns_error( self, proxy: ProxyClient, resources: ResourceManager ) -> None: @@ -59,6 +74,14 @@ class TestFilesBatchesContract: assert_client_error(result, "batch missing input_file_id") @pytest.mark.covers("llm.batches.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.BATCHES, + providers=(Provider.OPENAI,), + mode=Mode.BATCH, + ) + ) def test_retrieve_invalid_batch_id_returns_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: key = resources.key() result = proxy.transport.get( diff --git a/tests/e2e/llm_translation/test_google_native_e2e.py b/tests/e2e/llm_translation/test_google_native_e2e.py index 6910519c6df..b18683bd4ad 100644 --- a/tests/e2e/llm_translation/test_google_native_e2e.py +++ b/tests/e2e/llm_translation/test_google_native_e2e.py @@ -13,6 +13,7 @@ from typing import Literal import pytest from e2e_config import unique_marker from e2e_http import StreamingResponse, require_successful_call +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import LiteLLMParamsBody from proxy_client import ProxyClient @@ -85,6 +86,15 @@ def _streamed_text(result: StreamingResponse) -> str: class TestGoogleNativeGenerateContent: @pytest.mark.covers("llm.google_native.gemini.basic.nonstream.cost_logged") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.GOOGLE_GENAI, + providers=(Provider.GEMINI,), + models=(UPSTREAM_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_generate_content_returns_response_cost_header( self, proxy: ProxyClient, @@ -104,6 +114,15 @@ class TestGoogleNativeGenerateContent: assert result.response_cost > 0, f"x-litellm-response-cost must be a real cost, got {result.response_cost}" @pytest.mark.covers("llm.google_native.gemini.basic.stream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.GOOGLE_GENAI, + providers=(Provider.GEMINI,), + models=(UPSTREAM_MODEL,), + mode=Mode.STREAM, + ) + ) def test_stream_generate_content_frames_sse_the_way_google_sdks_expect( self, proxy: ProxyClient, diff --git a/tests/e2e/llm_translation/test_image_edits_e2e.py b/tests/e2e/llm_translation/test_image_edits_e2e.py index e95b054862e..66b0ebfc538 100644 --- a/tests/e2e/llm_translation/test_image_edits_e2e.py +++ b/tests/e2e/llm_translation/test_image_edits_e2e.py @@ -11,10 +11,12 @@ as the `image` part, not a JSON body. The fixture image is a small generated from __future__ import annotations import base64 +from typing import Final import openai import pytest from e2e_config import SLOW_PROVIDER_TIMEOUT_SECONDS, unique_marker +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import LiteLLMParamsBody from proxy_client import ProxyClient @@ -22,6 +24,8 @@ from sdk_clients import SdkClients pytestmark = pytest.mark.e2e +IMAGE_EDIT_BACKEND: Final = "openai/gpt-image-1" + _TEST_PNG = base64.b64decode( "iVBORw0KGgoAAAANSUhEUgAAAEAAAABACAIAAAAlC+aJAAAAS0lEQVR42u3PMQ0AAAwDoPo3" "3UrYvQQckD4XAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEB" @@ -33,7 +37,7 @@ def _register_image_model(proxy: ProxyClient, resources: ResourceManager) -> tup model = f"e2e-image-edit-{unique_marker()}" model_id = proxy.create_model( model, - LiteLLMParamsBody(model="openai/gpt-image-1", api_key="os.environ/OPENAI_API_KEY"), + LiteLLMParamsBody(model=IMAGE_EDIT_BACKEND, api_key="os.environ/OPENAI_API_KEY"), ) resources.defer(lambda: proxy.delete_model(model_id)) return model, resources.key() @@ -49,6 +53,15 @@ def _assert_client_error(error: openai.APIStatusError, context: str) -> None: class TestImageEdit: @pytest.mark.covers("llm.images_edits.openai.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.IMAGES, + providers=(Provider.OPENAI,), + models=(IMAGE_EDIT_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_image_edit_returns_image(self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients) -> None: model, key = _register_image_model(proxy, resources) client = sdk.openai(key) @@ -64,6 +77,15 @@ class TestImageEdit: assert first.b64_json or first.url, f"edited image has neither b64_json nor url: {first!r}" @pytest.mark.covers("llm.images_edits.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.IMAGES, + providers=(Provider.OPENAI,), + models=(IMAGE_EDIT_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_empty_prompt_returns_error(self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients) -> None: model, key = _register_image_model(proxy, resources) client = sdk.openai(key) @@ -73,6 +95,15 @@ class TestImageEdit: _assert_client_error(raised.value, "empty image-edit prompt") @pytest.mark.covers("llm.images_edits.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.IMAGES, + providers=(Provider.OPENAI,), + models=(IMAGE_EDIT_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_empty_image_returns_error(self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients) -> None: model, key = _register_image_model(proxy, resources) client = sdk.openai(key) diff --git a/tests/e2e/llm_translation/test_image_generation_e2e.py b/tests/e2e/llm_translation/test_image_generation_e2e.py index 1db40e7e15a..15cf20aa65a 100644 --- a/tests/e2e/llm_translation/test_image_generation_e2e.py +++ b/tests/e2e/llm_translation/test_image_generation_e2e.py @@ -8,9 +8,12 @@ from litellm-regression-tests/tests/test_inference_endpoints.py. from __future__ import annotations +from typing import Final + import pytest from e2e_config import unique_marker from e2e_http import assert_client_error +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import LiteLLMParamsBody from openai.types import ImagesResponse @@ -20,6 +23,9 @@ from sdk_clients import SdkClients pytestmark = pytest.mark.e2e +OPENAI_IMAGE_BACKEND: Final = "openai/gpt-image-1-mini" +BEDROCK_IMAGE_BACKEND: Final = "bedrock/amazon.nova-canvas-v1:0" + class _OptionalImageBody(BaseModel): model: str | None = None @@ -47,12 +53,21 @@ def _register_openai_image(proxy: ProxyClient, resources: ResourceManager) -> tu proxy, resources, "e2e-image", - LiteLLMParamsBody(model="openai/gpt-image-1-mini", api_key="os.environ/OPENAI_API_KEY"), + LiteLLMParamsBody(model=OPENAI_IMAGE_BACKEND, api_key="os.environ/OPENAI_API_KEY"), ) class TestImageGeneration: @pytest.mark.covers("llm.images_generations.openai.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.IMAGES, + providers=(Provider.OPENAI,), + models=(OPENAI_IMAGE_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_image_generation_returns_image( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -61,6 +76,15 @@ class TestImageGeneration: _assert_image_returned(images) @pytest.mark.covers("llm.images_generations.bedrock.basic.nonstream.works", exercised_on=["images_generations"]) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.IMAGES, + providers=(Provider.BEDROCK,), + models=(BEDROCK_IMAGE_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_bedrock_image_generation_returns_image( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -69,7 +93,7 @@ class TestImageGeneration: resources, "e2e-bedrock-image", LiteLLMParamsBody( - model="bedrock/amazon.nova-canvas-v1:0", + model=BEDROCK_IMAGE_BACKEND, aws_access_key_id="os.environ/AWS_ACCESS_KEY_ID", aws_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY", aws_region_name="os.environ/AWS_REGION", @@ -80,6 +104,12 @@ class TestImageGeneration: @pytest.mark.skip(reason="stage red: product gap, /v1/images/generations 500s (aimage_generation TypeError) on missing prompt instead of 400") @pytest.mark.covers("llm.images_generations.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.IMAGES, + ) + ) def test_missing_prompt_returns_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, key = _register_openai_image(proxy, resources) result = proxy.transport.send( @@ -90,6 +120,15 @@ class TestImageGeneration: assert_client_error(result, "images missing prompt") @pytest.mark.covers("llm.images_generations.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.IMAGES, + providers=(Provider.OPENAI,), + models=(OPENAI_IMAGE_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_empty_prompt_returns_client_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, key = _register_openai_image(proxy, resources) result = proxy.transport.send( @@ -100,6 +139,15 @@ class TestImageGeneration: assert_client_error(result, "images empty prompt") @pytest.mark.covers("llm.images_generations.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.IMAGES, + providers=(Provider.OPENAI,), + models=(OPENAI_IMAGE_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_invalid_size_returns_client_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, key = _register_openai_image(proxy, resources) result = proxy.transport.send( @@ -110,6 +158,15 @@ class TestImageGeneration: assert_client_error(result, "images invalid size") @pytest.mark.covers("llm.images_generations.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.IMAGES, + providers=(Provider.OPENAI,), + models=(OPENAI_IMAGE_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_invalid_n_returns_client_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, key = _register_openai_image(proxy, resources) result = proxy.transport.send( diff --git a/tests/e2e/llm_translation/test_messages_azure_foundry_e2e.py b/tests/e2e/llm_translation/test_messages_azure_foundry_e2e.py index 7a99f9c45e1..b10b8fe8bea 100644 --- a/tests/e2e/llm_translation/test_messages_azure_foundry_e2e.py +++ b/tests/e2e/llm_translation/test_messages_azure_foundry_e2e.py @@ -14,6 +14,7 @@ import pytest from anthropic.types import RawMessageStreamEvent, ToolParam from e2e_config import unique_marker +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import LiteLLMParamsBody from proxy_client import ProxyClient @@ -56,6 +57,15 @@ class TestAzureFoundryMessages: return model @pytest.mark.covers("llm.messages.azure_foundry.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE_AI,), + models=(AZURE_FOUNDRY_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_basic_nonstream(self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients) -> None: model = self._register(proxy, resources) client = sdk.anthropic(resources.key(models=[model])) @@ -71,6 +81,15 @@ class TestAzureFoundryMessages: assert text.strip(), f"/v1/messages returned no text: {message.content!r}" @pytest.mark.covers("llm.messages.azure_foundry.basic.stream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE_AI,), + models=(AZURE_FOUNDRY_MODEL,), + mode=Mode.STREAM, + ) + ) def test_basic_stream(self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients) -> None: model = self._register(proxy, resources) client = sdk.anthropic(resources.key(models=[model])) @@ -85,6 +104,16 @@ class TestAzureFoundryMessages: _assert_streamed_ok([event.type for event in stream]) @pytest.mark.covers("llm.messages.azure_foundry.tool_use.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE_AI,), + models=(AZURE_FOUNDRY_MODEL,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_tool_use_nonstream(self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients) -> None: model = self._register(proxy, resources) client = sdk.anthropic(resources.key(models=[model])) @@ -102,6 +131,16 @@ class TestAzureFoundryMessages: ) @pytest.mark.covers("llm.messages.azure_foundry.tool_use.stream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE_AI,), + models=(AZURE_FOUNDRY_MODEL,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.STREAM, + ) + ) def test_tool_use_stream(self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients) -> None: model = self._register(proxy, resources) client = sdk.anthropic(resources.key(models=[model])) @@ -122,6 +161,16 @@ class TestAzureFoundryMessages: ), "stream carried no tool_use block" assert "message_stop" in event_types, "stream never reached message_stop" + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE_AI,), + models=(AZURE_FOUNDRY_MODEL,), + capabilities=(Capability.RESPONSE_SCHEMA,), + mode=Mode.NONSTREAM, + ) + ) def test_output_format_returns_schema_json( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: diff --git a/tests/e2e/llm_translation/test_messages_bedrock_e2e.py b/tests/e2e/llm_translation/test_messages_bedrock_e2e.py index 61a37af7fd0..4e37f01886a 100644 --- a/tests/e2e/llm_translation/test_messages_bedrock_e2e.py +++ b/tests/e2e/llm_translation/test_messages_bedrock_e2e.py @@ -5,6 +5,7 @@ from typing import Final import pytest from anthropic.types import RawContentBlockDeltaEvent, RawMessageDeltaEvent, TextBlock, TextDelta from e2e_config import unique_marker +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import LiteLLMParamsBody from proxy_client import ProxyClient @@ -33,6 +34,16 @@ def _register(proxy: ProxyClient, resources: ResourceManager, backend: str) -> s class TestBedrockMessages: + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=(CONVERSE_CLAUDE_BACKEND,), + capabilities=(Capability.RESPONSE_SCHEMA,), + mode=Mode.NONSTREAM, + ) + ) def test_converse_output_format_returns_schema_json_text( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -50,6 +61,15 @@ class TestBedrockMessages: assert_sentiment_json("".join(texts)) @pytest.mark.covers("llm.messages.bedrock_converse.basic.stream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=(NOVA_BACKEND,), + mode=Mode.STREAM, + ) + ) def test_nova_stream_relays_text_usage_and_stop( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: diff --git a/tests/e2e/llm_translation/test_messages_e2e.py b/tests/e2e/llm_translation/test_messages_e2e.py index 90e74474d3b..6d2f690b947 100644 --- a/tests/e2e/llm_translation/test_messages_e2e.py +++ b/tests/e2e/llm_translation/test_messages_e2e.py @@ -17,6 +17,7 @@ from typing import Final import anthropic import pytest +from _pytest.mark.structures import ParameterSet from anthropic import Anthropic from anthropic.types import ( InputJSONDelta, @@ -42,6 +43,7 @@ from e2e_config import ( unique_marker, ) from e2e_http import assert_client_error +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import AnthropicErrorEvent, AnthropicMessagesBody, ChatMessage, LiteLLMParamsBody, SpendLogRow from provider_edge import EDGE_MOUNTS, LiveEdge, RunningEdge, StreamCut, start_provider_edge @@ -61,6 +63,7 @@ class _OptionalMessagesBody(BaseModel): ANTHROPIC_BACKEND = "anthropic/claude-haiku-4-5" +OPENAI_BRIDGE_BACKEND: Final = "openai/gpt-5.6" WEATHER_TOOL: ToolParam = { "name": "get_weather", @@ -110,6 +113,15 @@ def _user_turn(text: str) -> MessageParam: class TestAnthropicMessages: @pytest.mark.covers("llm.messages.anthropic.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_messages_returns_completion(self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients) -> None: model, key = _register(proxy, resources) client = sdk.anthropic(key) @@ -121,6 +133,15 @@ class TestAnthropicMessages: assert _text(message).strip(), f"/v1/messages returned no text: {message.content!r}" @pytest.mark.covers("llm.messages.anthropic.basic.nonstream.cost_logged") + @meta( + Subject( + domain=Domain.SPEND_BUDGETS, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_messages_logs_cost_matching_the_response_header( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -169,6 +190,15 @@ class TestAnthropicMessages: @pytest.mark.covers("llm.messages.anthropic.basic.stream.works") @pytest.mark.provider_live + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_BACKEND,), + mode=Mode.STREAM, + ) + ) def test_messages_streams_completion(self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients) -> None: """Edge-wired like its non-streaming siblings, so record and replay both carry the streamed response. @@ -226,6 +256,16 @@ class TestAnthropicMessages: ) @pytest.mark.covers("llm.messages.anthropic.tool_use.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_messages_tool_use(self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients) -> None: model, key = _register(proxy, resources) client = sdk.anthropic(key) @@ -243,6 +283,16 @@ class TestAnthropicMessages: ) @pytest.mark.covers("llm.messages.anthropic.structured_output.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_BACKEND,), + capabilities=(Capability.RESPONSE_SCHEMA,), + mode=Mode.NONSTREAM, + ) + ) def test_messages_output_format_returns_schema_json( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -261,6 +311,14 @@ class TestAnthropicMessages: reason="stage red: product gap, /v1/messages 500s (anthropic_messages TypeError) on missing messages instead of 400" ) @pytest.mark.covers("llm.messages.anthropic.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(), + models=(), + ) + ) def test_missing_messages_returns_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, key = _register(proxy, resources) result = proxy.transport.send( @@ -274,6 +332,14 @@ class TestAnthropicMessages: reason="stage red: product gap, /v1/messages 500s (anthropic_messages TypeError) on missing max_tokens instead of 400" ) @pytest.mark.covers("llm.messages.anthropic.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(), + models=(), + ) + ) def test_missing_max_tokens_returns_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, key = _register(proxy, resources) result = proxy.transport.send( @@ -284,6 +350,14 @@ class TestAnthropicMessages: assert_client_error(result, "messages missing max_tokens") @pytest.mark.covers("llm.messages.anthropic.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(), + models=(), + ) + ) def test_missing_model_returns_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: _, key = _register(proxy, resources) result = proxy.transport.send( @@ -364,9 +438,26 @@ def _request_tool(client: Anthropic, model: str, question: MessageParam, tool: T return blocks[0] +def _openai_bridge_subject(mode: Mode) -> Subject: + return Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.OPENAI,), + models=(OPENAI_BRIDGE_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=mode, + ) + + class TestOpenAIMessagesToolContinuation: @pytest.mark.provider_live - @pytest.mark.parametrize("stream", [True, False], ids=["stream", "nonstream"]) + @pytest.mark.parametrize( + "stream", + [ + pytest.param(stream, marks=meta(_openai_bridge_subject(mode)), id=name) + for stream, name, mode in ((True, "stream", Mode.STREAM), (False, "nonstream", Mode.NONSTREAM)) + ], + ) def test_required_tool_arguments_and_correlated_result( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients, stream: bool ) -> None: @@ -375,7 +466,7 @@ class TestOpenAIMessagesToolContinuation: model_id: Final = proxy.create_model( model, LiteLLMParamsBody( - model="openai/gpt-5.6", api_key="os.environ/OPENAI_API_KEY", api_base=f"{base}/v1" if base else None + model=OPENAI_BRIDGE_BACKEND, api_key="os.environ/OPENAI_API_KEY", api_base=f"{base}/v1" if base else None ), ) resources.defer(lambda: proxy.delete_model(model_id)) @@ -478,6 +569,30 @@ _DROPPED_BEFORE_FIRST_BYTE: Final[tuple[tuple[str, _CutRegistration, StreamCut], ) +_CUT_SUBJECTS: Final[MappingProxyType[_CutRegistration, Subject]] = MappingProxyType( + { + _register_cut_bedrock: Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=(BEDROCK_BACKEND,), + mode=Mode.STREAM, + ), + _register_cut_anthropic: Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_BACKEND,), + mode=Mode.STREAM, + ), + } +) + + +def _cut_params(cases: tuple[tuple[str, _CutRegistration, StreamCut], ...]) -> list[ParameterSet]: + return [pytest.param(register, cut, id=name, marks=meta(_CUT_SUBJECTS[register])) for name, register, cut in cases] + + def _payload(frame: str) -> JsonValue | None: try: return _FRAME_PAYLOAD.validate_json(frame) @@ -495,7 +610,7 @@ def _bare_error_frame(frame: str) -> bool: class TestMessagesUpstreamStreamFailure: @pytest.mark.covers("llm.messages.anthropic.upstream_stream_failure.stream.error_event") @pytest.mark.parametrize( - ("register", "cut"), [case[1:] for case in _DROPPED_UPSTREAMS], ids=[case[0] for case in _DROPPED_UPSTREAMS] + ("register", "cut"), _cut_params(_DROPPED_UPSTREAMS) ) def test_interrupted_upstream_stream_raises_in_the_anthropic_sdk( self, @@ -533,7 +648,7 @@ class TestMessagesUpstreamStreamFailure: @pytest.mark.covers("llm.messages.anthropic.upstream_stream_failure.stream.error_event") @pytest.mark.parametrize( - ("register", "cut"), [case[1:] for case in _DROPPED_UPSTREAMS], ids=[case[0] for case in _DROPPED_UPSTREAMS] + ("register", "cut"), _cut_params(_DROPPED_UPSTREAMS) ) def test_interrupted_upstream_stream_is_an_anthropic_error_event( self, proxy: ProxyClient, resources: ResourceManager, register: _CutRegistration, cut: StreamCut @@ -585,11 +700,7 @@ class TestMessagesUpstreamStreamFailure: ) @pytest.mark.covers("llm.messages.anthropic.upstream_stream_failure.stream.error_status") - @pytest.mark.parametrize( - ("register", "cut"), - [case[1:] for case in _DROPPED_BEFORE_FIRST_BYTE], - ids=[case[0] for case in _DROPPED_BEFORE_FIRST_BYTE], - ) + @pytest.mark.parametrize(("register", "cut"), _cut_params(_DROPPED_BEFORE_FIRST_BYTE)) def test_upstream_that_hangs_up_before_the_first_byte_raises_with_its_status_in_the_anthropic_sdk( self, proxy: ProxyClient, @@ -622,11 +733,7 @@ class TestMessagesUpstreamStreamFailure: ) @pytest.mark.covers("llm.messages.anthropic.upstream_stream_failure.stream.error_status") - @pytest.mark.parametrize( - ("register", "cut"), - [case[1:] for case in _DROPPED_BEFORE_FIRST_BYTE], - ids=[case[0] for case in _DROPPED_BEFORE_FIRST_BYTE], - ) + @pytest.mark.parametrize(("register", "cut"), _cut_params(_DROPPED_BEFORE_FIRST_BYTE)) def test_upstream_that_hangs_up_before_the_first_byte_is_a_json_error_with_its_status( self, proxy: ProxyClient, resources: ResourceManager, register: _CutRegistration, cut: StreamCut ) -> None: diff --git a/tests/e2e/llm_translation/test_messages_mid_conversation_system_e2e.py b/tests/e2e/llm_translation/test_messages_mid_conversation_system_e2e.py index e9b4b394996..84e0674d959 100644 --- a/tests/e2e/llm_translation/test_messages_mid_conversation_system_e2e.py +++ b/tests/e2e/llm_translation/test_messages_mid_conversation_system_e2e.py @@ -34,6 +34,7 @@ import pytest from anthropic import Anthropic from anthropic.types import Message, MessageParam, TextBlockParam from e2e_config import unique_marker +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import LiteLLMParamsBody from proxy_client import ProxyClient @@ -208,6 +209,16 @@ class TestBedrockInvokeMidConversationSystem: "llm.messages.bedrock_invoke.mid_conversation_system.nonstream.cache_hit", exercised_on=[], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=(FLAGGED_INVOKE_MODEL,), + capabilities=(Capability.MID_CONVERSATION_SYSTEM, Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) + ) def test_flagged_model_keeps_prompt_cache_across_system_reminder( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -234,6 +245,16 @@ class TestBedrockInvokeMidConversationSystem: "llm.messages.bedrock_invoke.mid_conversation_system.nonstream.works", exercised_on=[], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.BEDROCK,), + models=(UNFLAGGED_INVOKE_MODEL,), + capabilities=(Capability.MID_CONVERSATION_SYSTEM, Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) + ) def test_unflagged_model_converts_system_reminder_and_succeeds( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: diff --git a/tests/e2e/llm_translation/test_messages_mid_conversation_system_native_providers_e2e.py b/tests/e2e/llm_translation/test_messages_mid_conversation_system_native_providers_e2e.py index 9f5ed8b05da..44701b827f0 100644 --- a/tests/e2e/llm_translation/test_messages_mid_conversation_system_native_providers_e2e.py +++ b/tests/e2e/llm_translation/test_messages_mid_conversation_system_native_providers_e2e.py @@ -41,6 +41,7 @@ import pytest from anthropic import Anthropic from anthropic.types import Message, MessageParam, TextBlockParam from e2e_config import unique_marker +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import LiteLLMParamsBody from proxy_client import ProxyClient @@ -290,6 +291,16 @@ class TestAzureFoundryMidConversationSystem: "llm.messages.azure_foundry.mid_conversation_system.nonstream.cache_hit", exercised_on=[], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE_AI,), + models=(FLAGGED_MODEL,), + capabilities=(Capability.MID_CONVERSATION_SYSTEM, Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) + ) def test_flagged_model_keeps_prompt_cache_across_system_reminder( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -299,6 +310,16 @@ class TestAzureFoundryMidConversationSystem: "llm.messages.azure_foundry.mid_conversation_system.nonstream.works", exercised_on=[], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.AZURE_AI,), + models=(UNFLAGGED_MODEL,), + capabilities=(Capability.MID_CONVERSATION_SYSTEM, Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) + ) def test_unflagged_model_converts_system_reminder_and_succeeds( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -322,6 +343,16 @@ class TestVertexMidConversationSystem: "llm.messages.vertex.mid_conversation_system.nonstream.cache_hit", exercised_on=[], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.VERTEX_AI,), + models=(FLAGGED_MODEL,), + capabilities=(Capability.MID_CONVERSATION_SYSTEM, Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) + ) def test_flagged_model_keeps_prompt_cache_across_system_reminder( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -333,6 +364,16 @@ class TestVertexMidConversationSystem: "llm.messages.vertex.mid_conversation_system.nonstream.works", exercised_on=[], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.VERTEX_AI,), + models=(UNFLAGGED_MODEL,), + capabilities=(Capability.MID_CONVERSATION_SYSTEM, Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) + ) def test_unflagged_model_converts_system_reminder_and_succeeds( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: diff --git a/tests/e2e/llm_translation/test_moderations_e2e.py b/tests/e2e/llm_translation/test_moderations_e2e.py index e936f7b335a..4d2d5ea8b79 100644 --- a/tests/e2e/llm_translation/test_moderations_e2e.py +++ b/tests/e2e/llm_translation/test_moderations_e2e.py @@ -9,9 +9,12 @@ negative stays on the shared transport, since the SDK refuses to send it. from __future__ import annotations +from typing import Final + import pytest from e2e_config import unique_marker from e2e_http import assert_client_error +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import LiteLLMParamsBody from openai.types import Moderation @@ -21,6 +24,7 @@ from sdk_clients import SdkClients pytestmark = pytest.mark.e2e +OPENAI_MODERATION_BACKEND: Final = "openai/omni-moderation-latest" VIOLENT_TEXT = "I am going to find you and kill you, and I will hurt everyone you love." BENIGN_TEXT = "I enjoyed the sunny afternoon and a relaxing walk in the park today." @@ -35,7 +39,7 @@ def _register_moderation_model(proxy: ProxyClient, resources: ResourceManager) - model_id = proxy.create_model( model, LiteLLMParamsBody( - model="openai/omni-moderation-latest", api_key="os.environ/OPENAI_API_KEY" + model=OPENAI_MODERATION_BACKEND, api_key="os.environ/OPENAI_API_KEY" ), ) resources.defer(lambda: proxy.delete_model(model_id)) @@ -52,6 +56,15 @@ def _flagged_categories(item: Moderation) -> tuple[str, ...]: class TestModerations: @pytest.mark.covers("llm.moderations.openai.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MODERATIONS, + providers=(Provider.OPENAI,), + models=(OPENAI_MODERATION_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_moderations_flags_violent_content( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -64,6 +77,15 @@ class TestModerations: assert item.flagged, f"violent text was not flagged: {item!r}" assert _flagged_categories(item), f"flagged result reported no true category: {item!r}" + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MODERATIONS, + providers=(Provider.OPENAI,), + models=(OPENAI_MODERATION_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_moderations_passes_benign_content( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -79,6 +101,14 @@ class TestModerations: @pytest.mark.skip(reason="stage red: product gap, /v1/moderations 500s (KeyError 'input') on missing input instead of 400") @pytest.mark.covers("llm.moderations.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MODERATIONS, + providers=(), + models=(), + ) + ) def test_missing_input_returns_error( self, proxy: ProxyClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/llm_translation/test_ocr_rust_e2e.py b/tests/e2e/llm_translation/test_ocr_rust_e2e.py index b54a6ea010b..cdc3fe87431 100644 --- a/tests/e2e/llm_translation/test_ocr_rust_e2e.py +++ b/tests/e2e/llm_translation/test_ocr_rust_e2e.py @@ -25,6 +25,7 @@ from typing import Final, Protocol import pytest from e2e_config import unique_marker +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from e2e_http import ( PROVIDER_RATE_LIMIT_ATTEMPTS, RateLimitedError, @@ -60,6 +61,13 @@ TEST_IMAGE_URL = ( ) +MISTRAL_OCR_MODEL: Final = "mistral/mistral-ocr-latest" +AZURE_AI_OCR_MODEL: Final = "azure_ai/mistral-document-ai-2512" +AZURE_DOC_INTELLIGENCE_MODEL: Final = "azure_ai/doc-intelligence/prebuilt-layout" +VERTEX_OCR_MODEL: Final = "vertex_ai/mistral-ocr-2505" +COHERE_OCR_MODEL: Final = "cohere/parse-v5.0" + + class OcrProvider(Protocol): """One OCR provider's deployment config: its model id plus the os.environ/* credential references the proxy resolves at call time. Each provider owns which @@ -70,7 +78,7 @@ class OcrProvider(Protocol): @dataclass(frozen=True, slots=True) class MistralOcr: - model: str = "mistral/mistral-ocr-latest" + model: str = MISTRAL_OCR_MODEL def litellm_params(self) -> LiteLLMParamsBody: return LiteLLMParamsBody(model=self.model, api_key="os.environ/MISTRAL_API_KEY") @@ -96,7 +104,7 @@ class AzureDocIntelligenceOcr: AZURE_DOCUMENT_INTELLIGENCE_API_KEY, which the OCR config resolves from the doc-intelligence model name when api_base/api_key are left unset.""" - model: str = "azure_ai/doc-intelligence/prebuilt-layout" + model: str = AZURE_DOC_INTELLIGENCE_MODEL def litellm_params(self) -> LiteLLMParamsBody: return LiteLLMParamsBody(model=self.model) @@ -120,7 +128,7 @@ class VertexOcr: @dataclass(frozen=True, slots=True) class CohereOcr: - model: str = "cohere/parse-v5.0" + model: str = COHERE_OCR_MODEL def litellm_params(self) -> LiteLLMParamsBody: return LiteLLMParamsBody(model=self.model, api_key="os.environ/COHERE_API_KEY") @@ -141,7 +149,7 @@ RUST_OCR_CASES: tuple[_OcrCase, ...] = ( ), _OcrCase( "azure-ai", - AzureAiOcr("azure_ai/mistral-document-ai-2512"), + AzureAiOcr(AZURE_AI_OCR_MODEL), OcrDocument(type="document_url", document_url=TEST_PDF_URL), ), _OcrCase( @@ -151,13 +159,11 @@ RUST_OCR_CASES: tuple[_OcrCase, ...] = ( ), _OcrCase( "vertex-mistral", - VertexOcr("vertex_ai/mistral-ocr-2505", "us-central1"), + VertexOcr(VERTEX_OCR_MODEL, "us-central1"), OcrDocument(type="document_url", document_url=TEST_PDF_URL), ), ) -_CASE_IDS = tuple(case.suffix for case in RUST_OCR_CASES) - PDF_TEXT: Final = "test pdf file" IMAGE_TEXT: Final = "litellm" PDF_DOCUMENT: Final = OcrDocument(type="document_url", document_url=TEST_PDF_URL) @@ -175,14 +181,37 @@ class _OcrContentCase: OCR_CONTENT_CASES: Final = ( _OcrContentCase("mistral-pdf", MistralOcr(), PDF_DOCUMENT, PDF_TEXT), _OcrContentCase("mistral-image", MistralOcr(), IMAGE_DOCUMENT, IMAGE_TEXT), - _OcrContentCase("azure-ai-image", AzureAiOcr("azure_ai/mistral-document-ai-2512"), IMAGE_DOCUMENT, IMAGE_TEXT), + _OcrContentCase("azure-ai-image", AzureAiOcr(AZURE_AI_OCR_MODEL), IMAGE_DOCUMENT, IMAGE_TEXT), _OcrContentCase( - "vertex-mistral-image", VertexOcr("vertex_ai/mistral-ocr-2505", "us-central1"), IMAGE_DOCUMENT, IMAGE_TEXT + "vertex-mistral-image", VertexOcr(VERTEX_OCR_MODEL, "us-central1"), IMAGE_DOCUMENT, IMAGE_TEXT ), _OcrContentCase("cohere-image", CohereOcr(), IMAGE_DOCUMENT, IMAGE_TEXT), ) +def _ocr_subject(provider: OcrProvider) -> Subject: + match provider: + case MistralOcr(): + vendor, model = Provider.MISTRAL, MISTRAL_OCR_MODEL + case AzureAiOcr(): + vendor, model = Provider.AZURE_AI, AZURE_AI_OCR_MODEL + case AzureDocIntelligenceOcr(): + vendor, model = Provider.AZURE_AI, AZURE_DOC_INTELLIGENCE_MODEL + case VertexOcr(): + vendor, model = Provider.VERTEX_AI, VERTEX_OCR_MODEL + case CohereOcr(): + vendor, model = Provider.COHERE, COHERE_OCR_MODEL + case _: + raise TypeError(f"no OCR subject for {provider!r}") + return Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.OCR, + providers=(vendor,), + models=(model,), + mode=Mode.NONSTREAM, + ) + + def _assert_ocr_document(response: OcrResponse) -> None: assert response.object == "ocr", f"expected object='ocr', got {response.object!r}" assert response.model, "response missing the resolved model name" @@ -202,7 +231,10 @@ def _assert_provider_rate_limit_relayed(model: str, outcome: RateLimitedError) - class TestRustOcrGateway: - @pytest.mark.parametrize("case", RUST_OCR_CASES, ids=_CASE_IDS) + @pytest.mark.parametrize( + "case", + [pytest.param(case, marks=meta(_ocr_subject(case.provider)), id=case.suffix) for case in RUST_OCR_CASES], + ) def test_rust_ocr_response(self, proxy: ProxyClient, resources: ResourceManager, case: _OcrCase) -> None: model = f"rust-ocr-{case.suffix}-{unique_marker()}" model_id = proxy.create_model(model, case.provider.litellm_params()) @@ -219,6 +251,14 @@ class TestRustOcrGateway: @pytest.mark.skip(reason="stage red: product gap, /v1/ocr 500s (aocr TypeError) on missing document instead of 400") @pytest.mark.covers("llm.ocr.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.OCR, + providers=(), + models=(), + ) + ) def test_missing_document_returns_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: model = f"rust-ocr-val-{unique_marker()}" model_id = proxy.create_model(model, MistralOcr().litellm_params()) @@ -233,7 +273,13 @@ class TestRustOcrGateway: class TestOcrDocumentContent: - @pytest.mark.parametrize("case", OCR_CONTENT_CASES, ids=tuple(case.suffix for case in OCR_CONTENT_CASES)) + @pytest.mark.parametrize( + "case", + [ + pytest.param(case, marks=meta(_ocr_subject(case.provider)), id=case.suffix) + for case in OCR_CONTENT_CASES + ], + ) def test_ocr_reads_the_document_and_bills_its_pages( self, proxy: ProxyClient, resources: ResourceManager, case: _OcrContentCase ) -> None: diff --git a/tests/e2e/llm_translation/test_ollama_e2e.py b/tests/e2e/llm_translation/test_ollama_e2e.py new file mode 100644 index 00000000000..6dea5a51562 --- /dev/null +++ b/tests/e2e/llm_translation/test_ollama_e2e.py @@ -0,0 +1,258 @@ +"""Ollama behind the proxy on /chat/completions, /v1/messages and /v1/responses. + +Ollama has two litellm routes with different tool plumbing: `ollama_chat/` calls +/api/chat and forwards native tools, while `ollama/` calls /api/generate, which +has no tools field, so litellm prompts the model for a JSON function call and +turns that JSON back into a tool call. Each route runs the same conversation +contract as the conversational matrix on every surface, through the matrix's +SDK-backed surfaces, plus a streamed tool call on chat completions, the shape +coding agents such as OpenCode consume. + +The deployments set drop_params because Ollama has no parallel_tool_calls, +which the matrix surfaces send alongside a forced tool_choice. Live only: Ollama +Cloud has no provider edge mount, and its requests are not priced in the cost +map, so there is no cost cell here. +""" + +from __future__ import annotations + +import json +from collections.abc import Iterator, Mapping +from dataclasses import dataclass +from itertools import chain, product +from types import MappingProxyType +from typing import Final, Literal, cast + +import pytest +from _pytest.mark.structures import ParameterSet +from e2e_config import unique_marker +from e2e_metadata import Capability as SubjectCapability +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta +from lifecycle import ResourceManager +from llm_translation.conversational_matrix import ( + GREETING_PROMPT, + INSTRUCTIONS, + MAX_OUTPUT_TOKENS, + SURFACES, + WEATHER_PROMPT, + WEATHER_REPORT, + WEATHER_TOOL_DESCRIPTION, + WEATHER_TOOL_NAME, + WEATHER_TOOL_SCHEMA, + Surface, + SurfaceName, + ToolCall, + WeatherArgs, + build_surfaces, +) +from llm_translation.sdk_clients import NO_PROXY_CACHE, SdkClients +from models import LiteLLMParamsBody +from openai.types.chat import ChatCompletionChunk, ChatCompletionToolParam +from proxy_client import ProxyClient + +pytestmark = pytest.mark.e2e + +OllamaRoute = Literal["ollama_chat", "ollama"] +Capability = Literal["basic", "tool_use", "multi_turn"] +Streaming = Literal["stream", "nonstream"] + +OLLAMA_API_BASE: Final = "https://ollama.com" +OLLAMA_MODEL: Final = "gemma4:31b" +ROUTES: Final[tuple[OllamaRoute, ...]] = ("ollama_chat", "ollama") + + +@dataclass(frozen=True, slots=True) +class Cell: + surface: SurfaceName + route: OllamaRoute + + @property + def id(self) -> str: + return f"{self.surface}-{self.route}" + + +def _cells(capability: Capability, streaming: Streaming) -> tuple[ParameterSet, ...]: + return tuple( + pytest.param( + Cell(surface=surface, route=route), + id=f"{surface}-{route}", + marks=pytest.mark.covers(f"llm.{surface}.{route}.{capability}.{streaming}.works"), + ) + for surface, route in product(SURFACES, ROUTES) + ) + + +def _streamed_tool_cells() -> tuple[ParameterSet, ...]: + return tuple( + pytest.param(route, id=route, marks=pytest.mark.covers(f"llm.chat_completions.{route}.tool_use.stream.works")) + for route in ROUTES + ) + + +def _register(proxy: ProxyClient, resources: ResourceManager, route: OllamaRoute) -> str: + alias: Final = f"e2e-ollama-{route}-{unique_marker()}" + model_id: Final = proxy.create_model( + alias, + LiteLLMParamsBody( + model=f"{route}/{OLLAMA_MODEL}", + api_base=OLLAMA_API_BASE, + api_key="os.environ/OLLAMA_API_KEY", + drop_params=True, + ), + ) + resources.defer(lambda: proxy.delete_model(model_id)) + return alias + + +@pytest.fixture(scope="module") +def aliases(proxy: ProxyClient) -> Iterator[Mapping[OllamaRoute, str]]: + resources: Final = ResourceManager(client=proxy) + try: + yield MappingProxyType({route: _register(proxy, resources, route) for route in ROUTES}) + finally: + resources.teardown() + + +@pytest.fixture(scope="module") +def surfaces(sdk: SdkClients) -> Mapping[SurfaceName, Surface]: + return build_surfaces(sdk) + + +def _weather_call(cell: Cell, surface: Surface, key: str, model: str) -> ToolCall: + reply: Final = surface.reply(key, model, WEATHER_PROMPT, with_tool=True) + assert len(reply.tool_calls) == 1, ( + f"{cell.id}: expected one {WEATHER_TOOL_NAME} call, got {reply.tool_calls} text={reply.text!r}" + ) + call: Final = reply.tool_calls[0] + assert call.name == WEATHER_TOOL_NAME, f"{cell.id}: called {call.name!r}, not {WEATHER_TOOL_NAME!r}" + assert call.call_id, f"{cell.id}: tool call has no id, so the caller cannot answer it: {call}" + assert "paris" in call.parsed().location.lower(), f"{cell.id}: tool arguments lost the location: {call}" + return call + + +def _weather_tool() -> ChatCompletionToolParam: + return { + "type": "function", + "function": { + "name": WEATHER_TOOL_NAME, + "description": WEATHER_TOOL_DESCRIPTION, + "parameters": dict(WEATHER_TOOL_SCHEMA), + }, + } + + +def _subject(mode: Mode, *, tools: bool, route: Route | None = None) -> Subject: + return Subject( + domain=Domain.LLM_TRANSLATION, + route=route, + providers=(Provider.OLLAMA,), + models=(OLLAMA_MODEL,), + capabilities=(SubjectCapability.FUNCTION_CALLING,) if tools else (), + mode=mode, + ) + + +class TestOllamaConversation: + @pytest.mark.parametrize("cell", _cells("basic", "nonstream")) + @meta(_subject(Mode.NONSTREAM, tools=False)) + def test_reply_carries_assistant_text_and_usage( + self, + cell: Cell, + aliases: Mapping[OllamaRoute, str], + surfaces: Mapping[SurfaceName, Surface], + resources: ResourceManager, + ) -> None: + reply: Final = surfaces[cell.surface].reply(resources.key(), aliases[cell.route], GREETING_PROMPT) + + assert reply.response_id, f"{cell.id}: response has no id" + assert reply.text.strip(), f"{cell.id}: response carried no assistant text" + assert reply.usage is not None and reply.usage.input_tokens > 0 and reply.usage.output_tokens > 0, ( + f"{cell.id}: usage missing or zero: {reply.usage}" + ) + assert reply.call_id_header, f"{cell.id}: x-litellm-call-id header missing" + + @pytest.mark.parametrize("cell", _cells("basic", "stream")) + @meta(_subject(Mode.STREAM, tools=False)) + def test_stream_delivers_text_usage_and_a_terminal_event( + self, + cell: Cell, + aliases: Mapping[OllamaRoute, str], + surfaces: Mapping[SurfaceName, Surface], + resources: ResourceManager, + ) -> None: + streamed: Final = surfaces[cell.surface].stream(resources.key(), aliases[cell.route], GREETING_PROMPT) + + assert streamed.event_count > 1, f"{cell.id}: stream arrived as {streamed.event_count} event(s)" + assert streamed.text.strip(), f"{cell.id}: stream carried no text deltas" + assert streamed.finished, f"{cell.id}: stream never sent its terminal event" + assert streamed.usage_reported, f"{cell.id}: stream never reported usage" + + @pytest.mark.parametrize("cell", _cells("tool_use", "nonstream")) + @meta(_subject(Mode.NONSTREAM, tools=True)) + def test_tool_call_is_returned_named_and_addressable( + self, + cell: Cell, + aliases: Mapping[OllamaRoute, str], + surfaces: Mapping[SurfaceName, Surface], + resources: ResourceManager, + ) -> None: + _ = _weather_call(cell, surfaces[cell.surface], resources.key(), aliases[cell.route]) + + @pytest.mark.parametrize("cell", _cells("multi_turn", "nonstream")) + @meta(_subject(Mode.NONSTREAM, tools=True)) + def test_tool_result_round_trip_reaches_the_model( + self, + cell: Cell, + aliases: Mapping[OllamaRoute, str], + surfaces: Mapping[SurfaceName, Surface], + resources: ResourceManager, + ) -> None: + key: Final = resources.key() + model: Final = aliases[cell.route] + surface: Final = surfaces[cell.surface] + call: Final = _weather_call(cell, surface, key, model) + + answer: Final = surface.reply_to_tool_result(key, model, WEATHER_PROMPT, call, WEATHER_REPORT) + assert "22" in answer.text, f"{cell.id}: the model never saw the tool result: {answer.text!r}" + + +class TestOllamaStreamedToolCall: + @pytest.mark.parametrize("route", _streamed_tool_cells()) + @meta(_subject(Mode.STREAM, tools=True, route=Route.CHAT_COMPLETIONS)) + def test_tool_call_streams_as_tool_call_deltas( + self, + route: OllamaRoute, + aliases: Mapping[OllamaRoute, str], + sdk: SdkClients, + resources: ResourceManager, + ) -> None: + chunks: Final[tuple[ChatCompletionChunk, ...]] = tuple( + sdk.openai(resources.key()).chat.completions.create( + model=aliases[route], + messages=[ + {"role": "system", "content": INSTRUCTIONS}, + {"role": "user", "content": WEATHER_PROMPT}, + ], + tools=[_weather_tool()], + max_completion_tokens=MAX_OUTPUT_TOKENS, + stream=True, + extra_body=NO_PROXY_CACHE, + ) + ) + choices: Final = tuple(chunk.choices[0] for chunk in chunks if chunk.choices) + text: Final = "".join(choice.delta.content or "" for choice in choices) + deltas: Final = tuple(chain.from_iterable(choice.delta.tool_calls or () for choice in choices)) + call_ids: Final = tuple(delta.id for delta in deltas if delta.id) + indexes: Final = frozenset(delta.index for delta in deltas) + functions: Final = tuple(delta.function for delta in deltas if delta.function is not None) + names: Final = tuple(function.name for function in functions if function.name) + arguments: Final = "".join(function.arguments or "" for function in functions) + finish_reasons: Final = tuple(choice.finish_reason for choice in choices if choice.finish_reason is not None) + + assert names == (WEATHER_TOOL_NAME,), f"{route}: streamed tool names {names}, text={text!r}" + assert len(call_ids) == 1, f"{route}: expected one streamed tool call id, got {call_ids}" + assert indexes == {0}, f"{route}: streamed tool call deltas used indexes {sorted(indexes)}" + assert WEATHER_TOOL_NAME not in text, f"{route}: the tool call leaked into assistant text: {text!r}" + location: Final = WeatherArgs.model_validate(cast(object, json.loads(arguments))).location + assert "paris" in location.lower(), f"{route}: streamed tool arguments lost the location: {arguments!r}" + assert finish_reasons[-1:] == ("tool_calls",), f"{route}: stream finished with {finish_reasons}" diff --git a/tests/e2e/llm_translation/test_passthrough_e2e.py b/tests/e2e/llm_translation/test_passthrough_e2e.py index 447fe7d30d9..52d1c10d829 100644 --- a/tests/e2e/llm_translation/test_passthrough_e2e.py +++ b/tests/e2e/llm_translation/test_passthrough_e2e.py @@ -12,10 +12,13 @@ A passthrough call returning non-2xx fails hard (never a skip); once it returns 2xx, a missing or zero-cost SpendLogs row fails too. """ +from typing import Final + import pytest from e2e_config import CHEAP_OPENAI_MODEL, unique_marker from e2e_http import require_successful_call, unwrap +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import ChatResponse, KeyGenerateBody, SpendLogRow from passthrough_client import ( @@ -29,6 +32,8 @@ from passthrough_client import ( completed_responses_object, ) +GEMINI_MODEL: Final = "gemini-2.5-flash" +ANTHROPIC_PASSTHROUGH_MODEL: Final = "claude-haiku-4-5" EMBEDDING_MODEL = "text-embedding-3-small" REALTIME_MODEL = "gpt-realtime-2" @@ -57,12 +62,21 @@ def _fetch_cost_breakdown(client: PassthroughClient, request_id: str | None) -> # ---- Gemini passthrough ------------------------------------------------ +@meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.GEMINI,), + models=(GEMINI_MODEL,), + mode=Mode.NONSTREAM, + ) +) def test_gemini_passthrough_nonstreaming_logs_cost( client: PassthroughClient, scoped_key: str ) -> None: tag = f"e2e-passthrough-{unique_marker()}" result = client.gemini_generate( - scoped_key, "gemini-2.5-flash", "Say hello in one word", tags=[tag, "gemini"] + scoped_key, GEMINI_MODEL, "Say hello in one word", tags=[tag, "gemini"] ) require_successful_call(result) @@ -73,6 +87,15 @@ def test_gemini_passthrough_nonstreaming_logs_cost( @pytest.mark.skip(reason="stage red: product gap, native passthrough returns no x-litellm-response-cost or x-ratelimit-* headers") +@meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.GEMINI,), + models=(GEMINI_MODEL,), + mode=Mode.NONSTREAM, + ) +) def test_gemini_passthrough_returns_the_same_header_contract_as_the_managed_route( client: PassthroughClient, scoped_key: str ) -> None: @@ -82,7 +105,7 @@ def test_gemini_passthrough_returns_the_same_header_contract_as_the_managed_rout today, which makes native traffic invisible to the same tooling. """ result = client.gemini_generate( - scoped_key, "gemini-2.5-flash", f"Say hello in one word. {unique_marker()}" + scoped_key, GEMINI_MODEL, f"Say hello in one word. {unique_marker()}" ) require_successful_call(result) @@ -102,10 +125,19 @@ def test_gemini_passthrough_returns_the_same_header_contract_as_the_managed_rout ) +@meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.GEMINI,), + models=(GEMINI_MODEL,), + mode=Mode.STREAM, + ) +) def test_gemini_passthrough_streaming_logs_cost( client: PassthroughClient, scoped_key: str ) -> None: - result = client.gemini_stream(scoped_key, "gemini-2.5-flash", "Count to five") + result = client.gemini_stream(scoped_key, GEMINI_MODEL, "Count to five") require_successful_call(result) assert result.chunks > 0, "streaming passthrough produced no events" @@ -113,12 +145,22 @@ def test_gemini_passthrough_streaming_logs_cost( assert row.custom_llm_provider == "gemini" +@meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.GEMINI,), + models=(GEMINI_MODEL,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) +) def test_gemini_passthrough_tool_call_logs_cost( client: PassthroughClient, scoped_key: str ) -> None: result = client.gemini_generate( scoped_key, - "gemini-2.5-flash", + GEMINI_MODEL, "What is the weather in Paris? Use the get_weather tool.", tools=[ GeminiTool( @@ -146,10 +188,19 @@ def test_gemini_passthrough_tool_call_logs_cost( # ---- Anthropic passthrough --------------------------------------------- +@meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_PASSTHROUGH_MODEL,), + mode=Mode.NONSTREAM, + ) +) def test_anthropic_passthrough_nonstreaming_logs_cost( client: PassthroughClient, scoped_key: str ) -> None: - result = client.anthropic_message(scoped_key, "claude-haiku-4-5", "Say hello") + result = client.anthropic_message(scoped_key, ANTHROPIC_PASSTHROUGH_MODEL, "Say hello") require_successful_call(result) row = _fetch_cost_breakdown(client, anthropic_message_id(result)) @@ -157,11 +208,20 @@ def test_anthropic_passthrough_nonstreaming_logs_cost( assert "claude" in (row.model or "") +@meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_PASSTHROUGH_MODEL,), + mode=Mode.STREAM, + ) +) def test_anthropic_passthrough_streaming_logs_cost( client: PassthroughClient, scoped_key: str ) -> None: result = client.anthropic_message( - scoped_key, "claude-haiku-4-5", "Count to five", stream=True + scoped_key, ANTHROPIC_PASSTHROUGH_MODEL, "Count to five", stream=True ) require_successful_call(result) assert result.chunks > 0, "streaming passthrough produced no events" @@ -170,12 +230,22 @@ def test_anthropic_passthrough_streaming_logs_cost( assert row.custom_llm_provider == "anthropic" +@meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_PASSTHROUGH_MODEL,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) +) def test_anthropic_passthrough_tool_call_logs_cost( client: PassthroughClient, scoped_key: str ) -> None: result = client.anthropic_message( scoped_key, - "claude-haiku-4-5", + ANTHROPIC_PASSTHROUGH_MODEL, "What is the weather in Paris? Use the get_weather tool.", tools=[ AnthropicTool( @@ -205,13 +275,21 @@ class TestPassthroughModelAllowlist: """ @pytest.mark.covers("other.auth.passthrough.model_allowlist_enforced") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + route=Route.PASSTHROUGH, + providers=(), + models=(), + ) + ) def test_passthrough_denies_model_outside_key_allowlist( self, client: PassthroughClient, resources: ResourceManager ) -> None: - key = client.proxy.generate_key(KeyGenerateBody(models=["gemini-2.5-flash"])) + key = client.proxy.generate_key(KeyGenerateBody(models=[GEMINI_MODEL])) resources.defer(lambda: client.proxy.delete_key(key)) - result = client.anthropic_message(key, "claude-haiku-4-5", f"say hi {unique_marker()}") + result = client.anthropic_message(key, ANTHROPIC_PASSTHROUGH_MODEL, f"say hi {unique_marker()}") assert result.status_code == 403, ( "a key restricted to gemini-2.5-flash must be denied a claude passthrough call, " f"got {result.status_code}: {result.body[:300]}" @@ -230,6 +308,14 @@ class TestOpenAIPassthroughPrefix: """ @pytest.mark.covers("llm.files.openai.passthrough.nonstream.works") + @meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.OPENAI,), + models=(), + ) + ) def test_passthrough_prefix_uploads_a_file_to_openai( self, client: PassthroughClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -252,6 +338,14 @@ class TestOpenAIPassthroughPrefix: assert uploaded.bytes == len(content) @pytest.mark.covers("llm.batches.openai.passthrough.nonstream.works") + @meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.OPENAI,), + models=(), + ) + ) def test_passthrough_prefix_lists_batches_from_openai( self, client: PassthroughClient, scoped_key: str ) -> None: @@ -274,6 +368,15 @@ class TestOpenAIPassthroughSpend: """ @pytest.mark.covers("llm.responses.openai.passthrough.stream.cost_logged") + @meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.OPENAI,), + models=(CHEAP_OPENAI_MODEL,), + mode=Mode.STREAM, + ) + ) def test_streamed_responses_call_logs_its_cost( self, client: PassthroughClient, scoped_key: str ) -> None: @@ -318,6 +421,15 @@ class TestOpenAIPassthroughSpend: ) @pytest.mark.covers("llm.embeddings.openai.passthrough.nonstream.cost_logged") + @meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.OPENAI,), + models=(EMBEDDING_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_embeddings_call_logs_its_cost( self, client: PassthroughClient, scoped_key: str ) -> None: @@ -353,6 +465,15 @@ class TestOpenAIProviderPrefixChat: """ @pytest.mark.covers("llm.chat_completions.openai.passthrough.nonstream.cost_logged") + @meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.OPENAI,), + models=(CHEAP_OPENAI_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_openai_prefix_chat_returns_completion_and_logs_its_cost( self, client: PassthroughClient, scoped_key: str ) -> None: @@ -393,6 +514,15 @@ class TestOpenAIPassthroughWebsocket: """ @pytest.mark.covers("llm.realtime.openai.passthrough.stream.works") + @meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.OPENAI,), + models=(REALTIME_MODEL,), + mode=Mode.WEBSOCKET, + ) + ) def test_realtime_upgrade_reaches_openai_through_the_passthrough_prefix( self, client: PassthroughClient, scoped_key: str ) -> None: @@ -414,6 +544,15 @@ class TestOpenAIPassthroughWebsocket: ) @pytest.mark.covers("llm.responses.openai.passthrough_websocket.stream.works") + @meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.OPENAI,), + models=(), + mode=Mode.WEBSOCKET, + ) + ) def test_responses_upgrade_is_accepted_on_the_openai_prefix( self, client: PassthroughClient, scoped_key: str ) -> None: diff --git a/tests/e2e/llm_translation/test_passthrough_headers_e2e.py b/tests/e2e/llm_translation/test_passthrough_headers_e2e.py index 26f98774c63..66d51d122da 100644 --- a/tests/e2e/llm_translation/test_passthrough_headers_e2e.py +++ b/tests/e2e/llm_translation/test_passthrough_headers_e2e.py @@ -20,6 +20,7 @@ from pydantic import BaseModel, Field from e2e_config import unique_marker from e2e_http import AuthHeaders, NoBody, require_successful_call, unwrap +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import AnthropicMessagesResponse, ChatMessage, KeyGenerateBody from passthrough_client import PassthroughClient @@ -136,6 +137,15 @@ class TestPassthroughHeaders: "other.config.passthrough.headers_forwarded", exercised_on=[], ) + @meta( + Subject( + domain=Domain.PASSTHROUGH, + route=Route.PASSTHROUGH, + providers=(Provider.ANTHROPIC,), + models=(MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_static_and_x_pass_headers_reach_upstream( self, client: PassthroughClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/llm_translation/test_provider_features_e2e.py b/tests/e2e/llm_translation/test_provider_features_e2e.py index 2ea1d28748d..9a3de0aba7d 100644 --- a/tests/e2e/llm_translation/test_provider_features_e2e.py +++ b/tests/e2e/llm_translation/test_provider_features_e2e.py @@ -16,10 +16,13 @@ Prompt caching lives in test_cache_control.py. from __future__ import annotations +from typing import Final + import pytest from e2e_config import unique_marker from e2e_http import unwrap +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import ChatBody, ChatMessage, LiteLLMParamsBody from passthrough_client import PassthroughClient @@ -27,12 +30,22 @@ from passthrough_client import PassthroughClient pytestmark = pytest.mark.e2e SERVICE_TIER = "priority" +OPENAI_BACKEND: Final = "openai/gpt-5.5" class TestServiceTier: @pytest.mark.covers( "llm.chat_completions.openai.service_tier.nonstream.works", exercised_on=[] ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(OPENAI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_openai_service_tier_is_echoed( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -40,7 +53,7 @@ class TestServiceTier: model_id = client.proxy.create_model( model, LiteLLMParamsBody( - model="openai/gpt-5.5", api_key="os.environ/OPENAI_API_KEY" + model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY" ), ) resources.defer(lambda: client.proxy.delete_model(model_id)) diff --git a/tests/e2e/llm_translation/test_realtime_http_e2e.py b/tests/e2e/llm_translation/test_realtime_http_e2e.py index 9579ae13bbc..ee3ce3893c2 100644 --- a/tests/e2e/llm_translation/test_realtime_http_e2e.py +++ b/tests/e2e/llm_translation/test_realtime_http_e2e.py @@ -9,6 +9,7 @@ from __future__ import annotations import pytest from e2e_config import unique_marker from e2e_http import NoBody, assert_auth_denied, unwrap +from e2e_metadata import Domain, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import LiteLLMParamsBody from proxy_client import ProxyClient @@ -59,6 +60,14 @@ def _register(proxy: ProxyClient, resources: ResourceManager) -> tuple[str, str] class TestRealtimeHttp: @pytest.mark.covers("llm.realtime.openai.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.REALTIME, + providers=(Provider.OPENAI,), + models=(REALTIME_BACKEND,), + ) + ) def test_create_client_secret(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, key = _register(proxy, resources) secret = unwrap( @@ -82,6 +91,7 @@ class TestRealtimeHttp: assert secret.session.type in (None, "realtime"), f"unexpected session type: {secret.session.type}" @pytest.mark.covers("other.auth.realtime.missing_header_denied") + @meta(Subject(domain=Domain.PROXY_AUTH, route=Route.REALTIME)) def test_client_secret_missing_auth_is_denied(self, proxy: ProxyClient, resources: ResourceManager) -> None: model, _ = _register(proxy, resources) result = proxy.transport.send( @@ -92,6 +102,7 @@ class TestRealtimeHttp: assert_auth_denied(result, "realtime client_secrets missing auth") @pytest.mark.covers("other.auth.realtime.missing_header_denied") + @meta(Subject(domain=Domain.PROXY_AUTH, route=Route.REALTIME)) def test_calls_without_auth_is_denied(self, proxy: ProxyClient) -> None: result = proxy.transport.send( "/v1/realtime/calls", diff --git a/tests/e2e/llm_translation/test_rerank_e2e.py b/tests/e2e/llm_translation/test_rerank_e2e.py index 87b8618e6fb..9fafff03004 100644 --- a/tests/e2e/llm_translation/test_rerank_e2e.py +++ b/tests/e2e/llm_translation/test_rerank_e2e.py @@ -9,9 +9,12 @@ litellm-regression-tests/tests/test_inference_endpoints.py. from __future__ import annotations +from typing import Final + import pytest from e2e_config import unique_marker from e2e_http import unwrap +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import LiteLLMParamsBody, RerankBody, RerankResponse from proxy_client import ProxyClient @@ -24,6 +27,8 @@ DOCUMENTS = [ "Washington, D.C. is the capital of the United States.", "Capital punishment has existed in the United States since before it was a country.", ] +COHERE_RERANK_BACKEND: Final = "cohere/rerank-v3.5" +BEDROCK_RERANK_BACKEND: Final = "bedrock/arn:aws:bedrock:us-east-1::foundation-model/cohere.rerank-v3-5:0" QUERY = "What is the capital of the United States?" @@ -43,11 +48,20 @@ def _rerank_top_3(proxy: ProxyClient, key: str, model: str) -> RerankResponse: class TestRerank: @pytest.mark.covers("llm.rerank.cohere.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RERANK, + providers=(Provider.COHERE,), + models=(COHERE_RERANK_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_rerank_scores_top_n(self, proxy: ProxyClient, resources: ResourceManager) -> None: model = f"e2e-rerank-{unique_marker()}" model_id = proxy.create_model( model, - LiteLLMParamsBody(model="cohere/rerank-v3.5", api_key="os.environ/COHERE_API_KEY"), + LiteLLMParamsBody(model=COHERE_RERANK_BACKEND, api_key="os.environ/COHERE_API_KEY"), ) resources.defer(lambda: proxy.delete_model(model_id)) key = resources.key() @@ -55,6 +69,15 @@ class TestRerank: _assert_top_n_scored(_rerank_top_3(proxy, key, model)) @pytest.mark.covers("llm.rerank.bedrock.basic.nonstream.works", exercised_on=["rerank"]) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RERANK, + providers=(Provider.BEDROCK,), + models=(BEDROCK_RERANK_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_bedrock_rerank_scores_top_n( self, proxy: ProxyClient, resources: ResourceManager ) -> None: @@ -62,7 +85,7 @@ class TestRerank: model_id = proxy.create_model( model, LiteLLMParamsBody( - model="bedrock/arn:aws:bedrock:us-east-1::foundation-model/cohere.rerank-v3-5:0", + model=BEDROCK_RERANK_BACKEND, aws_access_key_id="os.environ/AWS_ACCESS_KEY_ID", aws_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY", aws_region_name="os.environ/AWS_REGION", diff --git a/tests/e2e/llm_translation/test_responses_bridge_streaming_e2e.py b/tests/e2e/llm_translation/test_responses_bridge_streaming_e2e.py index 75817340876..7826dc388ea 100644 --- a/tests/e2e/llm_translation/test_responses_bridge_streaming_e2e.py +++ b/tests/e2e/llm_translation/test_responses_bridge_streaming_e2e.py @@ -23,13 +23,14 @@ from pydantic import BaseModel, Field from e2e_config import unique_marker from e2e_http import StreamingResponse +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import ChatBody, ChatMessage, ChatTool, ChatToolFunction, LiteLLMParamsBody from passthrough_client import PassthroughClient pytestmark = pytest.mark.e2e -RESPONSES_ONLY_BACKEND = "openai/gpt-5.3-codex" +RESPONSES_ONLY_BACKEND: Final = "openai/gpt-5.3-codex" class _BridgeToolCallFunction(BaseModel): @@ -99,6 +100,15 @@ class TestResponsesBridgeChatCompletionsStreaming: "llm.chat_completions.openai.basic.stream.bridge_shares_chunk_id", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(RESPONSES_ONLY_BACKEND,), + mode=Mode.STREAM, + ) + ) def test_bridged_stream_shares_one_chunk_id( self, client: PassthroughClient, resources: ResourceManager, bridged_model: str ) -> None: @@ -124,6 +134,15 @@ class TestResponsesBridgeChatCompletionsStreaming: "llm.chat_completions.openai.basic.stream.bridge_streams_sse", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(RESPONSES_ONLY_BACKEND,), + mode=Mode.STREAM, + ) + ) def test_bridged_stream_delivers_content_finish_reason_and_done( self, client: PassthroughClient, resources: ResourceManager, bridged_model: str ) -> None: @@ -149,6 +168,16 @@ class TestResponsesBridgeChatCompletionsStreaming: "llm.chat_completions.openai.tool_use.stream.bridge_streams_tool_call", exercised_on=["chat_completions"], ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(RESPONSES_ONLY_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.STREAM, + ) + ) def test_bridged_stream_reassembles_tool_call( self, client: PassthroughClient, resources: ResourceManager, bridged_model: str ) -> None: diff --git a/tests/e2e/llm_translation/test_responses_e2e.py b/tests/e2e/llm_translation/test_responses_e2e.py index cc6f98dad50..00a1d3611f9 100644 --- a/tests/e2e/llm_translation/test_responses_e2e.py +++ b/tests/e2e/llm_translation/test_responses_e2e.py @@ -21,6 +21,7 @@ import openai import pytest from e2e_config import PROVIDER_EDGE_ADVERTISE_HOST, PROVIDER_EDGE_BIND_HOST, unique_marker from e2e_http import assert_client_error +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import ChatBody, ChatMessage, LiteLLMParamsBody from openai.types.responses import ( @@ -52,7 +53,10 @@ class _OptionalResponsesBody(BaseModel): max_output_tokens: int | None = None -BEDROCK_CONVERSE_BACKEND = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0" +OPENAI_MINI_BACKEND: Final = "openai/gpt-4o-mini" +OPENAI_VISION_BACKEND: Final = "openai/gpt-4o" +ANTHROPIC_BACKEND: Final = "anthropic/claude-haiku-4-5" +BEDROCK_CONVERSE_BACKEND: Final = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0" VERTEX_BACKEND: Final = "vertex_ai/gemini-2.5-flash" AZURE_OPENAI_BACKEND: Final = "azure/gpt-5.4-nano" AZURE_OPENAI_API_VERSION: Final = "v1" @@ -100,11 +104,11 @@ WEATHER_TOOL: FunctionToolParam = { def _openai_params() -> LiteLLMParamsBody: - return LiteLLMParamsBody(model="openai/gpt-4o-mini", api_key="os.environ/OPENAI_API_KEY") + return LiteLLMParamsBody(model=OPENAI_MINI_BACKEND, api_key="os.environ/OPENAI_API_KEY") def _anthropic_params() -> LiteLLMParamsBody: - return LiteLLMParamsBody(model="anthropic/claude-haiku-4-5", api_key="os.environ/ANTHROPIC_API_KEY") + return LiteLLMParamsBody(model=ANTHROPIC_BACKEND, api_key="os.environ/ANTHROPIC_API_KEY") def _bedrock_params() -> LiteLLMParamsBody: @@ -160,6 +164,15 @@ class WeatherArguments(BaseModel): class TestResponses: @pytest.mark.covers("llm.responses.openai.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(OPENAI_MINI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_responses_returns_completion( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -172,6 +185,15 @@ class TestResponses: assert response.output_text.strip(), f"/responses returned no output text: {response.output!r}" @pytest.mark.covers("llm.responses.openai.basic.stream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(OPENAI_MINI_BACKEND,), + mode=Mode.STREAM, + ) + ) def test_responses_streaming_returns_completion( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -194,6 +216,15 @@ class TestResponses: ) @pytest.mark.covers("llm.responses.openai.basic.nonstream.cost_logged") + @meta( + Subject( + domain=Domain.SPEND_BUDGETS, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(OPENAI_MINI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_responses_logs_cost(self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients) -> None: model = _register(proxy, resources, _openai_params()) client = sdk.openai(resources.key()) @@ -219,6 +250,16 @@ class TestResponses: assert "gpt-4o-mini" in (row.model or ""), f"unexpected spend row model: {row.model}" @pytest.mark.covers("llm.responses.openai.tool_use.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(OPENAI_MINI_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_responses_returns_function_call( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -235,13 +276,23 @@ class TestResponses: _assert_weather_call(response) @pytest.mark.covers("llm.responses.openai.vision.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(OPENAI_VISION_BACKEND,), + capabilities=(Capability.VISION,), + mode=Mode.NONSTREAM, + ) + ) def test_responses_vision_describes_image( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: model = _register( proxy, resources, - LiteLLMParamsBody(model="openai/gpt-4o", api_key="os.environ/OPENAI_API_KEY"), + LiteLLMParamsBody(model=OPENAI_VISION_BACKEND, api_key="os.environ/OPENAI_API_KEY"), ) client = sdk.openai(resources.key()) @@ -264,6 +315,15 @@ class TestResponses: ) @pytest.mark.covers("llm.responses.anthropic.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_responses_anthropic_returns_completion( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -276,6 +336,16 @@ class TestResponses: assert response.output_text.strip(), f"/responses returned no output text: {response.output!r}" @pytest.mark.covers("llm.responses.anthropic.tool_use.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_responses_anthropic_returns_function_call( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -292,6 +362,15 @@ class TestResponses: _assert_weather_call(response) @pytest.mark.covers("llm.responses.bedrock_converse.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.BEDROCK,), + models=(BEDROCK_CONVERSE_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_responses_bedrock_returns_completion( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -304,6 +383,16 @@ class TestResponses: assert response.output_text.strip(), f"/responses over bedrock returned no output text: {response.output!r}" @pytest.mark.covers("llm.responses.bedrock_converse.tool_use.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.BEDROCK,), + models=(BEDROCK_CONVERSE_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_responses_bedrock_returns_function_call( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -320,6 +409,15 @@ class TestResponses: _assert_weather_call(response) @pytest.mark.covers("llm.responses.vertex.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.VERTEX_AI,), + models=(VERTEX_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_responses_vertex_returns_completion( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -332,6 +430,16 @@ class TestResponses: assert response.output_text.strip(), f"/responses over vertex returned no output text: {response.output!r}" @pytest.mark.covers("llm.responses.vertex.tool_use.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.VERTEX_AI,), + models=(VERTEX_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_responses_vertex_returns_function_call( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -349,6 +457,15 @@ class TestResponses: _assert_weather_call(response) @pytest.mark.covers("llm.responses.azure_openai.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.AZURE,), + models=(AZURE_OPENAI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_responses_azure_openai_returns_completion( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -363,6 +480,16 @@ class TestResponses: ) @pytest.mark.covers("llm.responses.azure_openai.tool_use.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.AZURE,), + models=(AZURE_OPENAI_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_responses_azure_openai_returns_function_call( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -380,7 +507,37 @@ class TestResponses: _assert_weather_call(response) @pytest.mark.provider_edge_host - @pytest.mark.parametrize("endpoint", ["/v1/responses", "/v1/chat/completions"]) + @pytest.mark.parametrize( + "endpoint", + [ + pytest.param( + "/v1/responses", + marks=meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.BEDROCK,), + models=(BEDROCK_CONVERSE_BACKEND,), + mode=Mode.NONSTREAM, + ) + ), + id="/v1/responses", + ), + pytest.param( + "/v1/chat/completions", + marks=meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.BEDROCK,), + models=(BEDROCK_CONVERSE_BACKEND,), + mode=Mode.NONSTREAM, + ) + ), + id="/v1/chat/completions", + ), + ], + ) def test_bedrock_forwards_allowed_safety_identifier_as_additional_model_request_field( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients, endpoint: str ) -> None: @@ -442,6 +599,14 @@ class TestResponses: reason="stage red: product gap, /v1/responses 500s (aresponses TypeError) on missing input instead of 400" ) @pytest.mark.covers("llm.responses.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(OPENAI_MINI_BACKEND,), + ) + ) def test_missing_input_returns_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: model = _register(proxy, resources, _openai_params(), prefix="e2e-responses-val") key = resources.key() @@ -453,6 +618,12 @@ class TestResponses: assert_client_error(result, "responses missing input") @pytest.mark.covers("llm.responses.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + ) + ) def test_missing_model_returns_client_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: key = resources.key() result = proxy.transport.send( @@ -463,6 +634,14 @@ class TestResponses: assert_client_error(result, "responses missing model") @pytest.mark.covers("llm.responses.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(OPENAI_MINI_BACKEND,), + ) + ) def test_empty_input_returns_client_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: model = _register(proxy, resources, _openai_params(), prefix="e2e-responses-val") key = resources.key() @@ -509,6 +688,16 @@ class TodayReport(BaseModel): class TestResponsesOpenAIHostedFeatures: + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(REASONING_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING, Capability.REASONING, Capability.RESPONSE_SCHEMA), + mode=Mode.NONSTREAM, + ) + ) def test_reasoning_items_replay_into_structured_output_after_tool_call( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -561,6 +750,15 @@ class TestResponsesOpenAIHostedFeatures: assert TOOL_DATE in report.today, f"structured output ignored the tool result: {report!r}" @pytest.mark.provider_live + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(SHELL_BACKEND,), + mode=Mode.STREAM, + ) + ) def test_shell_tool_stream_surfaces_shell_call_and_its_output( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: diff --git a/tests/e2e/llm_translation/test_responses_retrieve_e2e.py b/tests/e2e/llm_translation/test_responses_retrieve_e2e.py index 063d014d1f5..b4a9634fbef 100644 --- a/tests/e2e/llm_translation/test_responses_retrieve_e2e.py +++ b/tests/e2e/llm_translation/test_responses_retrieve_e2e.py @@ -12,6 +12,7 @@ import openai import pytest from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, unique_marker from e2e_http import NoBody, Success, UnknownApiError, unwrap +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import LiteLLMParamsBody from openai.types.responses import ( @@ -27,6 +28,7 @@ from sdk_clients import NO_PROXY_CACHE, SdkClients pytestmark = pytest.mark.e2e OPENAI_BACKEND: Final = "openai/gpt-5.5" +OPENAI_MINI_BACKEND: Final = "openai/gpt-4o-mini" LONG_TASK: Final = "Write a numbered list counting from 1 to 400, one number per line, with a short word after each." CANCELLABLE_STATUSES: Final = frozenset({"queued", "in_progress"}) @@ -69,11 +71,20 @@ class TestResponsesRetrieve: reason="stage red: product gap (LIT-5446), retrieve returns a different id than the stored response (non-idempotent response-id re-encryption)" ) @pytest.mark.covers("llm.responses.openai.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(OPENAI_MINI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_store_and_retrieve_by_id(self, proxy: ProxyClient, resources: ResourceManager) -> None: model = f"e2e-resp-store-{unique_marker()}" model_id = proxy.create_model( model, - LiteLLMParamsBody(model="openai/gpt-4o-mini", api_key="os.environ/OPENAI_API_KEY"), + LiteLLMParamsBody(model=OPENAI_MINI_BACKEND, api_key="os.environ/OPENAI_API_KEY"), ) resources.defer(lambda: proxy.delete_model(model_id)) key = resources.key() @@ -103,6 +114,12 @@ class TestResponsesRetrieve: reason="stage red: product gap (LIT-5447), retrieving an unknown response id returns 400 (model=None) instead of 404" ) @pytest.mark.covers("llm.responses.openai.input_validation.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + ) + ) def test_invalid_response_id_returns_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: key = resources.key() get_result = proxy.transport.get( @@ -135,6 +152,15 @@ def _input_texts(item: object) -> tuple[str, ...]: @pytest.mark.provider_live class TestStoredResponseLifecycle: + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(OPENAI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_input_items_list_the_stored_prompt( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -150,6 +176,15 @@ class TestStoredResponseLifecycle: texts = tuple(text for item in items for text in _input_texts(item)) assert any(marker in text for text in texts), f"input_items did not list the stored prompt: {items!r}" + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(OPENAI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_deleted_response_is_no_longer_retrievable( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -171,6 +206,15 @@ class TestStoredResponseLifecycle: @pytest.mark.provider_live class TestBackgroundResponseCancel: + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(OPENAI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_cancel_background_response( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -185,6 +229,15 @@ class TestBackgroundResponseCancel: cancelled = client.responses.cancel(created.id) assert cancelled.status == "cancelled", f"cancel did not stop the response: {cancelled.status}" + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(OPENAI_BACKEND,), + mode=Mode.STREAM, + ) + ) def test_cancel_background_streaming_response_by_streamed_id( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: diff --git a/tests/e2e/llm_translation/test_sail_e2e.py b/tests/e2e/llm_translation/test_sail_e2e.py index 7267052e12c..77003ec4198 100644 --- a/tests/e2e/llm_translation/test_sail_e2e.py +++ b/tests/e2e/llm_translation/test_sail_e2e.py @@ -15,6 +15,7 @@ from typing import Final, Literal import pytest from e2e_config import SLOW_PROVIDER_TIMEOUT_SECONDS, unique_marker +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import LiteLLMParamsBody, SpendLogRow from openai import OpenAI @@ -116,6 +117,15 @@ def _assert_spend_row_matches(proxy: ProxyClient, key: str, header_cost: float) class TestSailChatCompletions: @pytest.mark.covers("llm.chat_completions.sail.service_tier.nonstream.cost_logged") + @meta( + Subject( + domain=Domain.SPEND_BUDGETS, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.SAIL,), + models=(BACKEND,), + mode=Mode.NONSTREAM, + ) + ) @pytest.mark.parametrize( ("service_tier", "billed_tier"), [("balanced", "balanced"), ("auto", "base")] ) @@ -150,6 +160,15 @@ class TestSailChatCompletions: _assert_spend_row_matches(proxy, key, header_cost) @pytest.mark.covers("llm.chat_completions.sail.service_tier.nonstream.drops_unknown_tier_and_bills_asap") + @meta( + Subject( + domain=Domain.SPEND_BUDGETS, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.SAIL,), + models=(BACKEND,), + mode=Mode.NONSTREAM, + ) + ) @pytest.mark.parametrize("service_tier", ["bogus", 5]) def test_unknown_service_tier_is_dropped_and_billed_asap( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients, service_tier: str | int @@ -176,6 +195,15 @@ class TestSailChatCompletions: class TestSailResponses: @pytest.mark.covers("llm.responses.sail.service_tier.nonstream.cost_logged") + @meta( + Subject( + domain=Domain.SPEND_BUDGETS, + route=Route.RESPONSES, + providers=(Provider.SAIL,), + models=(BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_caller_completion_window_bills_its_rates( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: @@ -204,6 +232,15 @@ class TestSailResponses: class TestSailMessages: @pytest.mark.covers("llm.messages.sail.basic.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.SAIL,), + models=(BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_plain_call_returns_a_message( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: diff --git a/tests/e2e/llm_translation/test_together_ai_e2e.py b/tests/e2e/llm_translation/test_together_ai_e2e.py index 874c6d77d19..d1ec2764a70 100644 --- a/tests/e2e/llm_translation/test_together_ai_e2e.py +++ b/tests/e2e/llm_translation/test_together_ai_e2e.py @@ -26,6 +26,7 @@ from typing import Final import pytest from e2e_config import STREAM_MIN_LEAD_SECONDS, provider_paces_stream, unique_marker from e2e_http import StreamingResponse, require_successful_call, unwrap +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import ( AnthropicAssistantTurn, @@ -324,6 +325,15 @@ def _weather_call(client: PassthroughClient, key: str, model: str) -> OutMessage class TestTogetherChatCompletions: @pytest.mark.covers("llm.chat_completions.together_ai.thinking.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.TOGETHER_AI,), + capabilities=(Capability.REASONING,), + mode=Mode.NONSTREAM, + ) + ) def test_reasoning_surfaces_as_reasoning_content( self, client: PassthroughClient, resources: ResourceManager, reasoning_tool_backend: str ) -> None: @@ -347,6 +357,15 @@ class TestTogetherChatCompletions: assert message.content and "43" in message.content, f"answer lost: {message}" @pytest.mark.covers("llm.chat_completions.together_ai.thinking.stream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.TOGETHER_AI,), + capabilities=(Capability.REASONING,), + mode=Mode.STREAM, + ) + ) def test_reasoning_streams_as_reasoning_content_deltas( self, client: PassthroughClient, resources: ResourceManager, reasoning_tool_backend: str ) -> None: @@ -369,6 +388,15 @@ class TestTogetherChatCompletions: assert "43" in content, f"streamed answer lost: {content!r}" @pytest.mark.covers("llm.chat_completions.together_ai.tool_use.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.TOGETHER_AI,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_tool_call_is_returned( self, client: PassthroughClient, resources: ResourceManager, reasoning_tool_backend: str ) -> None: @@ -376,6 +404,15 @@ class TestTogetherChatCompletions: _ = _weather_call_ids(_weather_call(client, key, model)) @pytest.mark.covers("llm.chat_completions.together_ai.tool_use.stream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.TOGETHER_AI,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.STREAM, + ) + ) def test_tool_call_is_streamed( self, client: PassthroughClient, resources: ResourceManager, reasoning_tool_backend: str ) -> None: @@ -407,6 +444,15 @@ class TestTogetherChatCompletions: assert "paris" in args.location.lower(), f"streamed tool arguments lost the location: {args}" @pytest.mark.covers("llm.chat_completions.together_ai.multi_turn.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.TOGETHER_AI,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_tool_result_round_trip( self, client: PassthroughClient, resources: ResourceManager, reasoning_tool_backend: str ) -> None: @@ -443,6 +489,16 @@ class TestTogetherChatCompletions: ) @pytest.mark.covers("llm.chat_completions.together_ai.thinking.nonstream.template_kwargs_forwarded") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.TOGETHER_AI,), + models=(HYBRID_REASONING_BACKEND,), + capabilities=(Capability.REASONING,), + mode=Mode.NONSTREAM, + ) + ) def test_chat_template_kwargs_reach_together( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -476,6 +532,16 @@ class TestTogetherChatCompletions: assert treatment.content and "43" in treatment.content, f"answer lost: {treatment}" @pytest.mark.covers("llm.chat_completions.together_ai.thinking.nonstream.replayed_reasoning_forwarded") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.TOGETHER_AI,), + models=(REASONING_REPLAY_BACKEND,), + capabilities=(Capability.REASONING,), + mode=Mode.NONSTREAM, + ) + ) def test_replayed_reasoning_content_reaches_together( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -496,6 +562,14 @@ class TestTogetherChatCompletions: ) @pytest.mark.covers("llm.chat_completions.together_ai.basic.nonstream.cost_logged") + @meta( + Subject( + domain=Domain.SPEND_BUDGETS, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.TOGETHER_AI,), + mode=Mode.NONSTREAM, + ) + ) def test_cost_header_and_spend_row_match_the_registry_price( self, client: PassthroughClient, @@ -551,6 +625,16 @@ class TestTogetherChatCompletions: ) @pytest.mark.covers("llm.chat_completions.together_ai.thinking.nonstream.effort_none_disables") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.TOGETHER_AI,), + models=(HYBRID_REASONING_BACKEND,), + capabilities=(Capability.REASONING,), + mode=Mode.NONSTREAM, + ) + ) def test_reasoning_effort_none_reaches_together( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -584,6 +668,16 @@ class TestTogetherChatCompletions: assert treatment.content and "43" in treatment.content, f"answer lost: {treatment}" @pytest.mark.covers("llm.chat_completions.together_ai.structured_output.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.TOGETHER_AI,), + models=(HYBRID_REASONING_BACKEND,), + capabilities=(Capability.RESPONSE_SCHEMA,), + mode=Mode.NONSTREAM, + ) + ) def test_response_format_json_schema_shapes_the_reply( self, client: PassthroughClient, resources: ResourceManager ) -> None: @@ -608,6 +702,15 @@ class TestTogetherChatCompletions: assert person.name, f"schema-shaped reply carries an empty name: {message.content!r}" @pytest.mark.covers("llm.chat_completions.together_ai.prompt_cache_5m.nonstream.cost_logged") + @meta( + Subject( + domain=Domain.SPEND_BUDGETS, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.TOGETHER_AI,), + capabilities=(Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) + ) def test_cache_read_tokens_bill_at_the_cache_read_rate( self, client: PassthroughClient, @@ -705,6 +808,15 @@ def _messages_weather_call( class TestTogetherMessages: @pytest.mark.covers("llm.messages.together_ai.tool_use.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.TOGETHER_AI,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_tool_use_block_is_returned( self, client: PassthroughClient, resources: ResourceManager, reasoning_tool_backend: str ) -> None: @@ -712,6 +824,15 @@ class TestTogetherMessages: _messages_weather_call(client, key, model) @pytest.mark.covers("llm.messages.together_ai.multi_turn.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.TOGETHER_AI,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_tool_result_round_trip( self, client: PassthroughClient, resources: ResourceManager, reasoning_tool_backend: str ) -> None: @@ -744,6 +865,14 @@ class TestTogetherMessages: @pytest.mark.covers("llm.messages.together_ai.basic.stream.works") @pytest.mark.provider_live + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.MESSAGES, + providers=(Provider.TOGETHER_AI,), + mode=Mode.STREAM, + ) + ) def test_streams_text_deltas( self, client: PassthroughClient, resources: ResourceManager, reasoning_tool_backend: str ) -> None: diff --git a/tests/e2e/llm_translation/test_token_counter_gemini_contents_e2e.py b/tests/e2e/llm_translation/test_token_counter_gemini_contents_e2e.py index b2b8f46ab0a..d50a7bae9d1 100644 --- a/tests/e2e/llm_translation/test_token_counter_gemini_contents_e2e.py +++ b/tests/e2e/llm_translation/test_token_counter_gemini_contents_e2e.py @@ -9,15 +9,43 @@ route the claude_code rows never reach from __future__ import annotations +from typing import Final + import pytest from e2e_config import unique_marker from e2e_http import require_successful_call +from e2e_metadata import Domain, Provider, Route, Subject, meta from proxy_client import ProxyClient from pydantic import BaseModel pytestmark = pytest.mark.e2e -GEMINI_DEPLOYMENTS = ("gemini-2.5-flash", "gemini-2.5-flash-vertex") +GEMINI_STUDIO_DEPLOYMENT: Final = "gemini-2.5-flash" +GEMINI_VERTEX_DEPLOYMENT: Final = "gemini-2.5-flash-vertex" +GEMINI_DEPLOYMENTS = ( + pytest.param( + GEMINI_STUDIO_DEPLOYMENT, + marks=meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.COUNT_TOKENS, + providers=(Provider.GEMINI,), + models=(GEMINI_STUDIO_DEPLOYMENT,), + ) + ), + ), + pytest.param( + GEMINI_VERTEX_DEPLOYMENT, + marks=meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.COUNT_TOKENS, + providers=(Provider.VERTEX_AI,), + models=(GEMINI_VERTEX_DEPLOYMENT,), + ) + ), + ), +) class _Part(BaseModel): diff --git a/tests/e2e/llm_translation/test_vector_stores_e2e.py b/tests/e2e/llm_translation/test_vector_stores_e2e.py index 71015d28d9f..2ae56594335 100644 --- a/tests/e2e/llm_translation/test_vector_stores_e2e.py +++ b/tests/e2e/llm_translation/test_vector_stores_e2e.py @@ -12,6 +12,7 @@ from typing import Literal import pytest from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, unique_marker +from e2e_metadata import Domain, Provider, Route, Subject, meta from e2e_http import ( FileUploadForm, NoBody, @@ -162,6 +163,7 @@ def _await_store_in_list(proxy: ProxyClient, key: str, store_id: str) -> None: class TestVectorStores: @pytest.mark.covers("llm.vector_stores.openai.basic.nonstream.works") + @meta(Subject(domain=Domain.LLM_TRANSLATION, route=Route.VECTOR_STORES, providers=(Provider.OPENAI,))) def test_create_list_retrieve_delete_lifecycle(self, proxy: ProxyClient, resources: ResourceManager) -> None: key = resources.key() name = f"e2e-vector-store-{unique_marker()}" @@ -204,6 +206,7 @@ class TestVectorStores: reason="stage red: product gap, vector store search 500s (asearch TypeError) on missing query instead of 400" ) @pytest.mark.covers("llm.vector_stores.openai.input_validation.nonstream.works") + @meta(Subject(domain=Domain.LLM_TRANSLATION, route=Route.VECTOR_STORES, providers=(Provider.OPENAI,))) def test_search_missing_query_returns_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: key = resources.key() created = unwrap( @@ -223,6 +226,7 @@ class TestVectorStores: assert_client_error(result, "vector store search missing query") @pytest.mark.covers("llm.vector_stores.openai.basic.nonstream.works") + @meta(Subject(domain=Domain.LLM_TRANSLATION, route=Route.VECTOR_STORES, providers=(Provider.OPENAI,))) def test_file_attach_poll_and_search(self, proxy: ProxyClient, resources: ResourceManager) -> None: key = resources.key() marker = f"azure-falcon-{unique_marker()}" @@ -309,6 +313,7 @@ class TestVectorStores: reason="stage red: product gap, retrieving a nonexistent vector store returns 2xx with an error envelope in the body instead of 404" ) @pytest.mark.covers("llm.vector_stores.openai.input_validation.nonstream.works") + @meta(Subject(domain=Domain.LLM_TRANSLATION, route=Route.VECTOR_STORES, providers=(Provider.OPENAI,))) def test_retrieve_invalid_id_returns_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: key = resources.key() result = proxy.transport.get( @@ -328,6 +333,7 @@ class TestVectorStores: pytest.fail(f"invalid vector store id must be a client error, got {other!r}") @pytest.mark.covers("llm.vector_stores.openai.input_validation.nonstream.works") + @meta(Subject(domain=Domain.LLM_TRANSLATION, route=Route.VECTOR_STORES, providers=(Provider.OPENAI,))) def test_invalid_chunking_returns_error(self, proxy: ProxyClient, resources: ResourceManager) -> None: key = resources.key() result = proxy.transport.send( diff --git a/tests/e2e/llm_translation/test_vertex_passthrough_e2e.py b/tests/e2e/llm_translation/test_vertex_passthrough_e2e.py index 5e9c9f614e5..3daf78a31b2 100644 --- a/tests/e2e/llm_translation/test_vertex_passthrough_e2e.py +++ b/tests/e2e/llm_translation/test_vertex_passthrough_e2e.py @@ -30,6 +30,7 @@ from pydantic import BaseModel from e2e_config import settle_propagation, unique_marker from e2e_http import NoBody, require_successful_call, unwrap +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import SpendLogRow from passthrough_client import PassthroughClient @@ -149,6 +150,15 @@ def _costed_row(client: PassthroughClient, call_id: str | None) -> SpendLogRow: class TestVertexPassthroughSpendTracking: + @meta( + Subject( + domain=Domain.SPEND_BUDGETS, + route=Route.PASSTHROUGH, + providers=(Provider.VERTEX_AI,), + models=(VERTEX_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_vertex_passthrough_via_managed_model_logs_cost( self, client: PassthroughClient, diff --git a/tests/e2e/logging/datadog_reader.py b/tests/e2e/logging/datadog_reader.py index 368c20cb6aa..61b1b5a1536 100644 --- a/tests/e2e/logging/datadog_reader.py +++ b/tests/e2e/logging/datadog_reader.py @@ -32,6 +32,7 @@ from e2e_config import ( POLL_TIMEOUT, ) from e2e_http import URL, Headers, StreamingResponse, send +from e2e_metadata import step type SearchCall = Callable[[str, float], StreamingResponse] @@ -115,6 +116,7 @@ class DdLogsReader: sleep: Callable[[float], None] = field(default=time.sleep, repr=False) jitter: Callable[[], float] = field(default=random.random, repr=False) + @step("Search DataDog for logs carrying the marker {marker}") def events_for_marker(self, marker: str) -> list[DdLogEvent]: """Every ingested event whose attributes carry the marker. DataDog consumes the shipped JSON message into ``attributes`` and leaves the @@ -124,6 +126,7 @@ class DdLogsReader: it).""" return self.events_for_query(f"*:*{marker}*") + @step("Search DataDog for logs matching {query}") def events_for_query(self, query: str) -> list[DdLogEvent]: """Every ingested event the search query matches (failure payloads carry no prompt to mark, so failure scenarios query indexed attributes @@ -156,10 +159,12 @@ class DdLogsReader: timeout=timeout, ) + @step("Wait for DataDog to ingest logs carrying the marker {marker}, then watch for duplicates") def poll_events_for_marker(self, marker: str) -> list[DdLogEvent]: """``poll_events_for_query`` over the every-attribute marker scan.""" return self.poll_events_for_query(f"*:*{marker}*") + @step("Wait for DataDog to ingest logs matching {query}, then watch for duplicates") def poll_events_for_query(self, query: str) -> list[DdLogEvent]: """Poll until at least one matching event is searchable (the callback flushes in periodic batches and DataDog ingestion adds seconds of lag), diff --git a/tests/e2e/logging/gcs_reader.py b/tests/e2e/logging/gcs_reader.py index 60622c121ac..8abde23e221 100644 --- a/tests/e2e/logging/gcs_reader.py +++ b/tests/e2e/logging/gcs_reader.py @@ -30,6 +30,7 @@ from pydantic import BaseModel, ConfigDict, Field from e2e_config import POLL_INTERVAL, POLL_TIMEOUT from e2e_http import URL, Headers, probe +from e2e_metadata import step _GCS_API = "https://storage.googleapis.com" #: Tolerance for clock skew between this host and GCS object timestamps. @@ -44,7 +45,7 @@ class _ServiceAccount(BaseModel): model_config = ConfigDict(extra="ignore") client_email: str - private_key: str + private_key: str = Field(repr=False) class _GcsAuthHeaders(Headers): @@ -146,6 +147,7 @@ class GcsLogReader: ) return result.body + @step("Read the GCS bucket's log objects for the response {response_id}") def records_for_response_id(self, response_id: str, *, since: datetime) -> list[GcsLogRecord]: """Every payload written for ``response_id``: the direct ``{date}/{response_id}`` object plus any hit inside batch NDJSON @@ -168,6 +170,7 @@ class GcsLogReader: ) return records + @step("Wait for the log of the response {response_id} to land in the GCS bucket, then watch for duplicates") def poll_records_for_response_id(self, response_id: str, *, since: datetime) -> list[GcsLogRecord]: """Poll until the payload is readable (the gcs_bucket callback flushes on a ~20s timer), then keep re-reading for GCS_SETTLE_SECONDS - past a diff --git a/tests/e2e/logging/logging_client.py b/tests/e2e/logging/logging_client.py index c2f987ea33d..d522c01c054 100644 --- a/tests/e2e/logging/logging_client.py +++ b/tests/e2e/logging/logging_client.py @@ -24,6 +24,7 @@ import pytest from pydantic import BaseModel, ConfigDict, Field, JsonValue, TypeAdapter, ValidationError from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, settle_propagation +from e2e_metadata import step from proxy_client import ProxyClient from e2e_http import ( URL, @@ -326,6 +327,7 @@ def observation_has_guardrail(obs: LangfuseObservation, *, guardrail_name: str) class LoggingClient: proxy: ProxyClient + @step("Generate a virtual key named {alias} with models: {models}") def key_with_alias( self, alias: str, @@ -347,9 +349,11 @@ class LoggingClient: ) ) + @step("Delete the virtual key") def delete_key(self, key: str) -> None: self.proxy.delete_key(key) + @step("Create the team {alias} with models: {models}") def create_team( self, alias: str, @@ -370,6 +374,7 @@ class LoggingClient: ) ).team_id + @step("Delete the team") def delete_team(self, team_id: str) -> None: _ = self.proxy.transport.post( "/team/delete", @@ -378,6 +383,7 @@ class LoggingClient: response_type=NoBody, ) + @step("Create the internal user {user_email}") def create_user(self, *, user_email: str, user_id: str | None = None) -> str: return unwrap( self.proxy.transport.post( @@ -392,6 +398,7 @@ class LoggingClient: ) ).user_id + @step("Delete the internal user") def delete_user(self, user_id: str) -> None: _ = self.proxy.transport.post( "/user/delete", @@ -400,6 +407,7 @@ class LoggingClient: response_type=NoBody, ) + @step("Create the organization {alias} with models: {models}") def create_org(self, alias: str, *, models: list[str]) -> str: return unwrap( self.proxy.transport.post( @@ -410,6 +418,7 @@ class LoggingClient: ) ).organization_id + @step("Delete the organization") def delete_org(self, organization_id: str) -> None: _ = self.proxy.transport.delete( "/organization/delete", @@ -418,6 +427,7 @@ class LoggingClient: response_type=NoBody, ) + @step("Add a Langfuse OTel logging callback for {callback_type} events to the team") def add_team_langfuse_callback( self, team_id: str, @@ -441,6 +451,7 @@ class LoggingClient: f"POST /team/{team_id}/callback must return status=success; got {response.status!r}" ) + @step("Create the tool_permission guardrail {name} that allows only the tool {allowed_tool}") def create_tool_permission_guardrail(self, name: str, *, allowed_tool: str) -> str: """Register a tool_permission guardrail that allows one tool and denies the rest.""" response = unwrap( @@ -474,6 +485,7 @@ class LoggingClient: settle_propagation(time.monotonic()) return guardrail_id + @step("Delete the guardrail") def delete_guardrail(self, guardrail_id: str) -> None: _ = self.proxy.transport.delete( f"/guardrails/{guardrail_id}", @@ -482,12 +494,15 @@ class LoggingClient: response_type=NoBody, ) + @step("Add a deployment named {model_name} that calls {litellm_params.model}") def create_model(self, model_name: str, litellm_params: LiteLLMParamsBody) -> str: return self.proxy.create_model(model_name, litellm_params) + @step("Delete the deployment") def delete_model(self, model_id: str) -> None: self.proxy.delete_model(model_id) + @step('Send a /chat/completions request to {model} with the prompt "{text}"') def chat(self, key: str, model: str, text: str) -> ChatResponse: return unwrap( self.proxy.chat( @@ -500,6 +515,7 @@ class LoggingClient: ) ) + @step('Send a /chat/completions request to {model} with stream={stream} and the prompt "{text}"') def chat_raw( self, key: str, @@ -529,6 +545,7 @@ class LoggingClient: json=body, ) + @step('Send a /v1/messages request to {model} with stream={stream} and the prompt "{text}"') def messages_raw( self, key: str, model: str, text: str, *, max_tokens: int = 16, stream: bool = False ) -> StreamingResponse: @@ -545,6 +562,7 @@ class LoggingClient: return self.proxy.transport.stream("/v1/messages", headers=self.proxy.transport.bearer(key), json=body) return self.proxy.transport.send("/v1/messages", headers=self.proxy.transport.bearer(key), json=body) + @step('Send a /v1/responses request to {model} with stream={stream} and the prompt "{text}"') def responses_raw( self, key: str, model: str, text: str, *, max_output_tokens: int = 64, stream: bool = False ) -> StreamingResponse: @@ -560,9 +578,11 @@ class LoggingClient: return self.proxy.transport.stream("/v1/responses", headers=self.proxy.transport.bearer(key), json=body) return self.proxy.transport.send("/v1/responses", headers=self.proxy.transport.bearer(key), json=body) + @step("Scrape the Prometheus metrics from /metrics") def scrape_metrics(self) -> str: return self.proxy.probe("/metrics", params=NoBody()).body + @step("Wait for the key's spend log in /spend/logs") def poll_proxy_spend_for_key( self, key: str, @@ -590,6 +610,7 @@ class LoggingClient: return row return None + @step("List observations from Langfuse") def list_langfuse_observations( self, creds: LangfuseCreds, @@ -616,6 +637,7 @@ class LoggingClient: case _: return [] + @step("Look up the Langfuse generation for the key {key_alias}") def find_langfuse_observation( self, creds: LangfuseCreds, @@ -634,6 +656,7 @@ class LoggingClient: return obs return None + @step("Wait for the Langfuse generation for the key {key_alias}") def poll_langfuse_observation( self, creds: LangfuseCreds, @@ -653,6 +676,7 @@ class LoggingClient: time.sleep(POLL_INTERVAL) return last + @step("Wait for the OTel v2 Langfuse generation for the key {key_alias}") def poll_langfuse_generation( self, creds: LangfuseCreds, *, key_alias: str, from_start_time: str ) -> LangfuseObservation | None: @@ -665,6 +689,7 @@ class LoggingClient: time.sleep(POLL_INTERVAL) return None + @step("Wait for the Langfuse trace of the key {key_alias} and every observation in it") def poll_langfuse_trace_observations( self, creds: LangfuseCreds, @@ -698,6 +723,7 @@ def build_logging_client(proxy: ProxyClient) -> LoggingClient: return LoggingClient(proxy=proxy) +@step("Read the proxy's callback list from /health/readiness/details") def readiness_details_body(client: LoggingClient) -> str: """/health/readiness/details, tolerating the 503 it serves while the ephemeral stack's DB leg blips: the recorded state the logging suites check diff --git a/tests/e2e/logging/s3_reader.py b/tests/e2e/logging/s3_reader.py index d605dec6096..572086c2695 100644 --- a/tests/e2e/logging/s3_reader.py +++ b/tests/e2e/logging/s3_reader.py @@ -24,6 +24,7 @@ import pytest from pydantic import BaseModel, ConfigDict from e2e_config import POLL_INTERVAL, POLL_TIMEOUT +from e2e_metadata import step if TYPE_CHECKING: from types_boto3_s3.client import S3Client @@ -53,17 +54,21 @@ class S3LogReader: bucket: str client: S3Client + @step("List the log objects in the S3 bucket under {prefix}") def list_keys(self, prefix: str) -> list[str]: response = self.client.list_objects_v2(Bucket=self.bucket, Prefix=prefix) return [obj["Key"] for obj in response.get("Contents", []) if "Key" in obj] + @step("Download a log object from the S3 bucket") def read_record(self, key: str) -> S3LogRecord: body = self.client.get_object(Bucket=self.bucket, Key=key)["Body"].read() return S3LogRecord.model_validate_json(body) + @step("Read the log objects in the S3 bucket under {prefix}") def records_matching(self, *, prefix: str, predicate: Callable[[S3LogRecord], bool]) -> list[S3LogRecord]: return [record for record in map(self.read_record, self.list_keys(prefix)) if predicate(record)] + @step("Wait for the request's log object to land in the S3 bucket under {prefix}, then watch for duplicates") def poll_records(self, *, prefix: str, predicate: Callable[[S3LogRecord], bool]) -> list[S3LogRecord]: """Poll until at least one matching object is listed (the s3_v2 callback flushes on a ~10s timer), then keep re-reading for diff --git a/tests/e2e/logging/test_datadog_log_e2e.py b/tests/e2e/logging/test_datadog_log_e2e.py index 2d57181c4bf..108985a3953 100644 --- a/tests/e2e/logging/test_datadog_log_e2e.py +++ b/tests/e2e/logging/test_datadog_log_e2e.py @@ -20,10 +20,12 @@ from __future__ import annotations import math import time +from typing import Final import pytest from datadog_reader import DdLogEvent, DdLogsReader from e2e_config import CHEAP_ANTHROPIC_MODEL, CHEAP_OPENAI_MODEL, unique_marker +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from logging_client import INVALID_UPSTREAM_API_KEY, LoggingClient, first_ok, readiness_details_body from models import ChatMessage, LiteLLMParamsBody, ReliabilityChatBody, RouterSettingsOverride @@ -33,6 +35,7 @@ pytestmark = pytest.mark.e2e #: The active DataDog callback's name in /health/readiness/details success_callbacks. DD_LOGGER_NAME = "DataDogLogger" +FAILING_BACKEND_MODEL: Final = "anthropic/claude-haiku-4-5" class _DdMessagePayload(BaseModel): @@ -109,6 +112,15 @@ def _assert_exactly_one_event( class TestDataDogLogDelivery: @pytest.mark.covers("logging.datadog.success.exports_metric", exercised_on=["chat_completions"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(CHEAP_ANTHROPIC_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_chat_completions_emits_one_log_event( self, client: LoggingClient, dd_logs: DdLogsReader, resources: ResourceManager ) -> None: @@ -134,6 +146,15 @@ class TestDataDogLogDelivery: ) @pytest.mark.covers("logging.datadog.success.exports_metric", exercised_on=["messages"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=(CHEAP_ANTHROPIC_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_messages_emits_one_log_event( self, client: LoggingClient, dd_logs: DdLogsReader, resources: ResourceManager ) -> None: @@ -161,6 +182,15 @@ class TestDataDogLogDelivery: ) @pytest.mark.covers("logging.datadog.success.exports_metric", exercised_on=["responses"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(CHEAP_OPENAI_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_responses_emits_one_log_event( self, client: LoggingClient, dd_logs: DdLogsReader, resources: ResourceManager ) -> None: @@ -186,6 +216,15 @@ class TestDataDogLogDelivery: ) @pytest.mark.covers("logging.datadog.stream.exports_metric", exercised_on=["chat_completions"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(CHEAP_ANTHROPIC_MODEL,), + mode=Mode.STREAM, + ) + ) def test_chat_completions_stream_emits_one_log_event( self, client: LoggingClient, dd_logs: DdLogsReader, resources: ResourceManager ) -> None: @@ -232,6 +271,15 @@ class TestDataDogLogDelivery: ) @pytest.mark.covers("logging.datadog.stream.exports_metric", exercised_on=["messages"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=(CHEAP_ANTHROPIC_MODEL,), + mode=Mode.STREAM, + ) + ) def test_messages_stream_emits_one_log_event( self, client: LoggingClient, dd_logs: DdLogsReader, resources: ResourceManager ) -> None: @@ -276,6 +324,15 @@ class TestDataDogLogDelivery: ) @pytest.mark.covers("logging.datadog.stream.exports_metric", exercised_on=["responses"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(CHEAP_OPENAI_MODEL,), + mode=Mode.STREAM, + ) + ) def test_responses_stream_emits_one_log_event( self, client: LoggingClient, dd_logs: DdLogsReader, resources: ResourceManager ) -> None: @@ -347,6 +404,15 @@ def _assert_exactly_one_failure_event(events: list[DdLogEvent], *, model_group: class TestDataDogFailureDelivery: @pytest.mark.covers("logging.datadog.failure.exports_metric", exercised_on=["chat_completions"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(FAILING_BACKEND_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_failed_chat_completions_emits_one_error_event( self, client: LoggingClient, dd_logs: DdLogsReader, resources: ResourceManager ) -> None: @@ -366,7 +432,7 @@ class TestDataDogFailureDelivery: model_name = f"dd-err-{unique_marker()}" model_id = client.create_model( model_name, - LiteLLMParamsBody(model="anthropic/claude-haiku-4-5", api_key=INVALID_UPSTREAM_API_KEY), + LiteLLMParamsBody(model=FAILING_BACKEND_MODEL, api_key=INVALID_UPSTREAM_API_KEY), ) resources.defer(lambda: client.delete_model(model_id)) key = client.key_with_alias(f"dd-err-key-{unique_marker()}", models=[model_name]) @@ -399,6 +465,15 @@ class TestDataDogFailureDelivery: ) @pytest.mark.covers("logging.datadog.stream_failure.exports_metric", exercised_on=["chat_completions"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(FAILING_BACKEND_MODEL,), + mode=Mode.STREAM, + ) + ) def test_failed_chat_completions_stream_emits_one_error_event( self, client: LoggingClient, dd_logs: DdLogsReader, resources: ResourceManager ) -> None: @@ -420,7 +495,7 @@ class TestDataDogFailureDelivery: model_id = client.create_model( model_name, LiteLLMParamsBody( - model="anthropic/claude-haiku-4-5", + model=FAILING_BACKEND_MODEL, api_key=INVALID_UPSTREAM_API_KEY, api_base="http://localhost:1", ), diff --git a/tests/e2e/logging/test_gcs_log_e2e.py b/tests/e2e/logging/test_gcs_log_e2e.py index 17ad1507049..259735e6cad 100644 --- a/tests/e2e/logging/test_gcs_log_e2e.py +++ b/tests/e2e/logging/test_gcs_log_e2e.py @@ -22,6 +22,7 @@ import math import pytest from e2e_config import CHEAP_ANTHROPIC_MODEL, unique_marker +from e2e_metadata import Domain, Mode, Provider, Subject, meta from gcs_reader import GcsLogReader, build_gcs_reader, utc_now from lifecycle import ResourceManager from logging_client import LoggingClient, completion_response_id, first_ok, readiness_details_body @@ -52,6 +53,14 @@ def _assert_gcs_configured(client: LoggingClient) -> None: class TestGcsLogDelivery: @pytest.mark.covers("logging.gcs_bucket.success.writes_object", exercised_on=["chat_completions"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + providers=(Provider.ANTHROPIC,), + models=(CHEAP_ANTHROPIC_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_chat_completions_writes_one_success_record( self, client: LoggingClient, gcs_logs: GcsLogReader, resources: ResourceManager ) -> None: diff --git a/tests/e2e/logging/test_langsmith_batch_serialization_e2e.py b/tests/e2e/logging/test_langsmith_batch_serialization_e2e.py index 874b2b6a045..e2dadf9e696 100644 --- a/tests/e2e/logging/test_langsmith_batch_serialization_e2e.py +++ b/tests/e2e/logging/test_langsmith_batch_serialization_e2e.py @@ -21,6 +21,7 @@ from typing import Final import pytest from e2e_config import CHEAP_OPENAI_MODEL, POLL_INTERVAL, POLL_TIMEOUT, unique_marker from e2e_http import Headers, Success, get_external +from e2e_metadata import Domain, Mode, Provider, Subject, meta from pydantic import BaseModel, ConfigDict, Field, JsonValue import litellm @@ -90,6 +91,14 @@ def _poll_run(creds: LangsmithCreds, run_id: uuid.UUID) -> LangsmithRun: class TestLangsmithBatchSerialization: @pytest.mark.asyncio @pytest.mark.covers("logging.langsmith.success.serializes_non_native_metadata") + @meta( + Subject( + domain=Domain.OBSERVABILITY, + providers=(Provider.OPENAI,), + models=(CHEAP_OPENAI_MODEL,), + mode=Mode.NONSTREAM, + ) + ) async def test_non_json_native_metadata_reaches_langsmith(self) -> None: creds: Final = load_langsmith_creds() logger: Final = LangsmithLogger( diff --git a/tests/e2e/logging/test_otel_trace_e2e.py b/tests/e2e/logging/test_otel_trace_e2e.py index 4902b0703c3..01eaf5ce80c 100644 --- a/tests/e2e/logging/test_otel_trace_e2e.py +++ b/tests/e2e/logging/test_otel_trace_e2e.py @@ -25,6 +25,7 @@ from typing import Final import pytest from e2e_config import CHEAP_ANTHROPIC_MODEL, CHEAP_OPENAI_MODEL, OTEL_EXPORTER_ENDPOINT, unique_marker +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from logging_client import INVALID_UPSTREAM_API_KEY, LoggingClient, first_ok, readiness_details_body from models import LiteLLMParamsBody @@ -34,6 +35,7 @@ from pydantic import BaseModel, ConfigDict, ValidationError pytestmark = pytest.mark.e2e MODEL = CHEAP_ANTHROPIC_MODEL +FAILING_BACKEND_MODEL: Final = "anthropic/claude-haiku-4-5" DB_SPAN_PREFIX = "postgres." #: The active OTEL v2 logger's name in /health/readiness/details success_callbacks. OTEL_V2_LOGGER_NAME = "OpenTelemetryV2" @@ -279,6 +281,15 @@ def _assert_error_span_contract(span: JaegerSpan) -> None: class TestOtelTraceCompleteness: @pytest.mark.covers("logging.otel.success.exports_metric", exercised_on=["chat_completions"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_chat_completions_exports_complete_trace( self, client: LoggingClient, otel_reader: OtelReader, resources: ResourceManager ) -> None: @@ -312,6 +323,14 @@ class TestOtelTraceCompleteness: _assert_complete_trace(traces, route=route, genai_span=f"chat {MODEL}") @pytest.mark.covers("logging.otel.success.exports_metric", exercised_on=["chat_completions"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + providers=(Provider.ANTHROPIC,), + models=(MODEL,), + mode=Mode.NONSTREAM, + ) + ) @pytest.mark.otel_tls def test_otel_export_over_tls_with_internal_ca_reaches_destination( self, client: LoggingClient, otel_reader: OtelReader, resources: ResourceManager @@ -337,6 +356,15 @@ class TestOtelTraceCompleteness: _assert_complete_trace(hits, route=route, genai_span=f"chat {MODEL}") @pytest.mark.covers("logging.otel.success.exports_metric", exercised_on=["messages"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=(MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_messages_exports_complete_trace( self, client: LoggingClient, otel_reader: OtelReader, resources: ResourceManager ) -> None: @@ -366,6 +394,15 @@ class TestOtelTraceCompleteness: _assert_complete_trace(traces, route=route, genai_span=f"chat {MODEL}") @pytest.mark.covers("logging.otel.success.exports_metric", exercised_on=["responses"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(CHEAP_OPENAI_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_responses_exports_complete_trace( self, client: LoggingClient, otel_reader: OtelReader, resources: ResourceManager ) -> None: @@ -396,6 +433,15 @@ class TestOtelTraceCompleteness: _assert_complete_trace(traces, route=route, genai_span=genai_span) @pytest.mark.covers("logging.otel.stream.exports_metric", exercised_on=["chat_completions"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(MODEL,), + mode=Mode.STREAM, + ) + ) def test_chat_completions_stream_exports_complete_trace( self, client: LoggingClient, otel_reader: OtelReader, resources: ResourceManager ) -> None: @@ -444,6 +490,15 @@ class TestOtelTraceCompleteness: ) @pytest.mark.covers("logging.otel.stream.exports_metric", exercised_on=["messages"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=(MODEL,), + mode=Mode.STREAM, + ) + ) def test_messages_stream_exports_complete_trace( self, client: LoggingClient, otel_reader: OtelReader, resources: ResourceManager ) -> None: @@ -492,6 +547,15 @@ class TestOtelTraceCompleteness: ) @pytest.mark.covers("logging.otel.stream.exports_metric", exercised_on=["responses"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(CHEAP_OPENAI_MODEL,), + mode=Mode.STREAM, + ) + ) def test_responses_stream_exports_complete_trace( self, client: LoggingClient, otel_reader: OtelReader, resources: ResourceManager ) -> None: @@ -544,6 +608,15 @@ class TestOtelTraceCompleteness: ) @pytest.mark.covers("logging.otel.stream.records_ttft", exercised_on=["chat_completions"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(MODEL,), + mode=Mode.STREAM, + ) + ) def test_chat_completions_stream_records_real_ttft( self, client: LoggingClient, otel_reader: OtelReader, resources: ResourceManager ) -> None: @@ -582,6 +655,15 @@ class TestOtelTraceCompleteness: _assert_real_ttft(traces.hits, genai_span=genai_span) @pytest.mark.covers("logging.otel.stream.records_ttft", exercised_on=["messages"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=(MODEL,), + mode=Mode.STREAM, + ) + ) def test_messages_stream_records_real_ttft( self, client: LoggingClient, otel_reader: OtelReader, resources: ResourceManager ) -> None: @@ -620,6 +702,15 @@ class TestOtelTraceCompleteness: _assert_real_ttft(traces.hits, genai_span=genai_span) @pytest.mark.covers("logging.otel.stream.records_ttft", exercised_on=["responses"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(CHEAP_OPENAI_MODEL,), + mode=Mode.STREAM, + ) + ) def test_responses_stream_records_real_ttft( self, client: LoggingClient, otel_reader: OtelReader, resources: ResourceManager ) -> None: @@ -658,6 +749,15 @@ class TestOtelTraceCompleteness: _assert_real_ttft(traces.hits, genai_span=genai_span) @pytest.mark.covers("logging.otel.failure.exports_metric", exercised_on=["chat_completions"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(FAILING_BACKEND_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_failed_chat_completions_error_span_attributes( self, client: LoggingClient, otel_reader: OtelReader, resources: ResourceManager ) -> None: @@ -677,7 +777,7 @@ class TestOtelTraceCompleteness: model_name = f"otel-err-{unique_marker()}" model_id = client.create_model( model_name, - LiteLLMParamsBody(model="anthropic/claude-haiku-4-5", api_key=INVALID_UPSTREAM_API_KEY), + LiteLLMParamsBody(model=FAILING_BACKEND_MODEL, api_key=INVALID_UPSTREAM_API_KEY), ) resources.defer(lambda: client.delete_model(model_id)) key = client.key_with_alias(f"otel-err-{unique_marker()}", models=[model_name]) @@ -711,6 +811,15 @@ class TestOtelTraceCompleteness: _assert_error_span_contract(genai) @pytest.mark.covers("logging.otel.failure.exports_metric", exercised_on=["messages"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=(FAILING_BACKEND_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_failed_messages_error_span_attributes( self, client: LoggingClient, otel_reader: OtelReader, resources: ResourceManager ) -> None: @@ -729,7 +838,7 @@ class TestOtelTraceCompleteness: model_name = f"otel-err-{unique_marker()}" model_id = client.create_model( model_name, - LiteLLMParamsBody(model="anthropic/claude-haiku-4-5", api_key=INVALID_UPSTREAM_API_KEY), + LiteLLMParamsBody(model=FAILING_BACKEND_MODEL, api_key=INVALID_UPSTREAM_API_KEY), ) resources.defer(lambda: client.delete_model(model_id)) key = client.key_with_alias(f"otel-err-{unique_marker()}", models=[model_name]) diff --git a/tests/e2e/logging/test_otel_v2_langfuse_generation_output_e2e.py b/tests/e2e/logging/test_otel_v2_langfuse_generation_output_e2e.py index 9e08cf21da0..852e10636d0 100644 --- a/tests/e2e/logging/test_otel_v2_langfuse_generation_output_e2e.py +++ b/tests/e2e/logging/test_otel_v2_langfuse_generation_output_e2e.py @@ -24,6 +24,7 @@ from typing import Final import pytest from e2e_config import unique_marker from e2e_http import unwrap +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from logging_client import LangfuseCreds, LangfuseObservation, LoggingClient, load_langfuse_creds from models import ( @@ -69,6 +70,13 @@ RED_SQUARE_PNG: Final = base64.b64decode( ) BOUNDED_OUTPUT_CHARS: Final = 1024 PLACEHOLDER_INPUT: Final = "default-message-value" +COMPLETION_BACKEND: Final = "openai/gpt-3.5-turbo-instruct" +IMAGE_BACKEND: Final = "openai/gpt-image-1-mini" +SPEECH_BACKEND: Final = "openai/gpt-4o-mini-tts" +TRANSCRIPTION_BACKEND: Final = "openai/gpt-4o-mini-transcribe" +MODERATION_BACKEND: Final = "openai/omni-moderation-latest" +MISTRAL_OCR_BACKEND: Final = "mistral/mistral-ocr-latest" +RERANK_BACKEND: Final = "cohere/rerank-v4.0-fast" class _OutputMessage(BaseModel): @@ -142,7 +150,7 @@ def _openai(model: str) -> LiteLLMParamsBody: def _mistral_ocr() -> LiteLLMParamsBody: - return LiteLLMParamsBody(model="mistral/mistral-ocr-latest", api_key="os.environ/MISTRAL_API_KEY") + return LiteLLMParamsBody(model=MISTRAL_OCR_BACKEND, api_key="os.environ/MISTRAL_API_KEY") def _langfuse_search_tool( @@ -171,10 +179,19 @@ def _langfuse_search_tool( class TestOtelV2LangfuseGenerationOutput: @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["completions"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.COMPLETIONS, + providers=(Provider.OPENAI,), + models=(COMPLETION_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_completions_output_is_the_completion_text( self, client: LoggingClient, langfuse_creds: LangfuseCreds, resources: ResourceManager ) -> None: - model, key, alias = _langfuse_key(client, langfuse_creds, resources, _openai("openai/gpt-3.5-turbo-instruct")) + model, key, alias = _langfuse_key(client, langfuse_creds, resources, _openai(COMPLETION_BACKEND)) started: Final = datetime.now(timezone.utc) response: Final = unwrap( client.proxy.transport.post( @@ -193,10 +210,19 @@ class TestOtelV2LangfuseGenerationOutput: ) @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["images_generations"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.IMAGES, + providers=(Provider.OPENAI,), + models=(IMAGE_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_images_output_is_a_bounded_summary_without_base64( self, client: LoggingClient, langfuse_creds: LangfuseCreds, resources: ResourceManager ) -> None: - model, key, alias = _langfuse_key(client, langfuse_creds, resources, _openai("openai/gpt-image-1-mini")) + model, key, alias = _langfuse_key(client, langfuse_creds, resources, _openai(IMAGE_BACKEND)) started: Final = datetime.now(timezone.utc) response: Final = unwrap( client.proxy.transport.post( @@ -220,10 +246,19 @@ class TestOtelV2LangfuseGenerationOutput: ) @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["audio_speech"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.AUDIO, + providers=(Provider.OPENAI,), + models=(SPEECH_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_speech_output_is_a_bounded_summary_without_audio_bytes( self, client: LoggingClient, langfuse_creds: LangfuseCreds, resources: ResourceManager ) -> None: - model, key, alias = _langfuse_key(client, langfuse_creds, resources, _openai("openai/gpt-4o-mini-tts")) + model, key, alias = _langfuse_key(client, langfuse_creds, resources, _openai(SPEECH_BACKEND)) started: Final = datetime.now(timezone.utc) audio: Final = client.proxy.transport.stream_binary( "/v1/audio/speech", @@ -239,10 +274,19 @@ class TestOtelV2LangfuseGenerationOutput: ) @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["audio_transcriptions"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.AUDIO, + providers=(Provider.OPENAI,), + models=(TRANSCRIPTION_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_transcription_output_is_the_transcript( self, client: LoggingClient, langfuse_creds: LangfuseCreds, resources: ResourceManager ) -> None: - model, key, alias = _langfuse_key(client, langfuse_creds, resources, _openai("openai/gpt-4o-mini-transcribe")) + model, key, alias = _langfuse_key(client, langfuse_creds, resources, _openai(TRANSCRIPTION_BACKEND)) started: Final = datetime.now(timezone.utc) response: Final = unwrap( client.proxy.transport.upload( @@ -262,10 +306,19 @@ class TestOtelV2LangfuseGenerationOutput: assert transcript in output, f"generation output lacks the transcript {transcript!r}: {output!r}" @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["moderations"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.MODERATIONS, + providers=(Provider.OPENAI,), + models=(MODERATION_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_moderations_output_is_the_verdict( self, client: LoggingClient, langfuse_creds: LangfuseCreds, resources: ResourceManager ) -> None: - model, key, alias = _langfuse_key(client, langfuse_creds, resources, _openai("openai/omni-moderation-latest")) + model, key, alias = _langfuse_key(client, langfuse_creds, resources, _openai(MODERATION_BACKEND)) started: Final = datetime.now(timezone.utc) response: Final = unwrap( client.proxy.transport.post( @@ -284,6 +337,15 @@ class TestOtelV2LangfuseGenerationOutput: ) @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["rerank"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.RERANK, + providers=(Provider.COHERE,), + models=(RERANK_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_rerank_output_is_the_ranked_indices_and_scores( self, client: LoggingClient, langfuse_creds: LangfuseCreds, resources: ResourceManager ) -> None: @@ -291,7 +353,7 @@ class TestOtelV2LangfuseGenerationOutput: client, langfuse_creds, resources, - LiteLLMParamsBody(model="cohere/rerank-v4.0-fast", api_key="os.environ/COHERE_API_KEY"), + LiteLLMParamsBody(model=RERANK_BACKEND, api_key="os.environ/COHERE_API_KEY"), ) query: Final = f"What is the capital of France? {unique_marker()}" started: Final = datetime.now(timezone.utc) @@ -317,6 +379,15 @@ class TestOtelV2LangfuseGenerationOutput: assert output == "\n\n".join(ranked), f"generation output is not the ranked indices and scores: {output!r}" @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["ocr"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.OCR, + providers=(Provider.MISTRAL,), + models=(MISTRAL_OCR_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_ocr_input_is_the_document_url( self, client: LoggingClient, langfuse_creds: LangfuseCreds, resources: ResourceManager ) -> None: @@ -334,6 +405,15 @@ class TestOtelV2LangfuseGenerationOutput: assert response.pages[0].markdown in _output_text(generation) @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["ocr"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.OCR, + providers=(Provider.MISTRAL,), + models=(MISTRAL_OCR_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_ocr_upload_input_is_a_bounded_document_summary_without_base64( self, client: LoggingClient, langfuse_creds: LangfuseCreds, resources: ResourceManager ) -> None: @@ -360,10 +440,19 @@ class TestOtelV2LangfuseGenerationOutput: ) @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["images_edits"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.IMAGES, + providers=(Provider.OPENAI,), + models=(IMAGE_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_image_edit_input_is_the_edit_prompt( self, client: LoggingClient, langfuse_creds: LangfuseCreds, resources: ResourceManager ) -> None: - model, key, alias = _langfuse_key(client, langfuse_creds, resources, _openai("openai/gpt-image-1-mini")) + model, key, alias = _langfuse_key(client, langfuse_creds, resources, _openai(IMAGE_BACKEND)) prompt: Final = f"make the square blue {unique_marker()}" started: Final = datetime.now(timezone.utc) response: Final = unwrap( @@ -388,6 +477,11 @@ class TestOtelV2LangfuseGenerationOutput: assert _output_text(generation).startswith("b64_json image (") @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["search"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + ) + ) def test_search_input_is_the_query_and_output_the_results( self, client: LoggingClient, langfuse_creds: LangfuseCreds, resources: ResourceManager ) -> None: diff --git a/tests/e2e/logging/test_prometheus_cardinality_e2e.py b/tests/e2e/logging/test_prometheus_cardinality_e2e.py index e2d164d8b4c..597e1e814e3 100644 --- a/tests/e2e/logging/test_prometheus_cardinality_e2e.py +++ b/tests/e2e/logging/test_prometheus_cardinality_e2e.py @@ -24,6 +24,7 @@ import pytest from prometheus_client.parser import text_string_to_metric_families from e2e_config import unique_marker +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from logging_client import LoggingClient @@ -47,6 +48,15 @@ def _aliases_in_metric(exposition: str, metric: str, label: str) -> frozenset[st class TestPrometheusPerKeyCardinality: @pytest.mark.covers("logging.prometheus.success.exports_metric", exercised_on=[]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.METRICS, + providers=(Provider.GEMINI,), + models=(DRIVER_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_distinct_key_aliases_produce_distinct_series( self, client: LoggingClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/logging/test_prometheus_queue_time_e2e.py b/tests/e2e/logging/test_prometheus_queue_time_e2e.py index 1f3c111bb65..9668322b362 100644 --- a/tests/e2e/logging/test_prometheus_queue_time_e2e.py +++ b/tests/e2e/logging/test_prometheus_queue_time_e2e.py @@ -6,6 +6,7 @@ import pytest from prometheus_client.parser import text_string_to_metric_families from e2e_config import unique_marker +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from logging_client import LoggingClient @@ -30,6 +31,15 @@ def _observation_count(exposition: str, alias: str) -> float | None: class TestPrometheusRequestQueueTime: @pytest.mark.covers("logging.prometheus.success.records_queue_time") + @meta( + Subject( + domain=Domain.OBSERVABILITY, + route=Route.METRICS, + providers=(Provider.GEMINI,), + models=(DRIVER_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_queue_time_histogram_records_an_observation( self, client: LoggingClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/logging/test_s3_log_e2e.py b/tests/e2e/logging/test_s3_log_e2e.py index 1612fa315a6..e917276edb0 100644 --- a/tests/e2e/logging/test_s3_log_e2e.py +++ b/tests/e2e/logging/test_s3_log_e2e.py @@ -24,10 +24,12 @@ from __future__ import annotations import math import re import time +from typing import Final import pytest from e2e_config import CHEAP_ANTHROPIC_MODEL, S3_PARTITION_GRANULARITY, unique_marker +from e2e_metadata import Domain, Mode, Provider, Subject, meta from lifecycle import ResourceManager from logging_client import ( INVALID_UPSTREAM_API_KEY, @@ -43,6 +45,7 @@ pytestmark = pytest.mark.e2e #: The active s3_v2 callback's name in /health/readiness/details success_callbacks. S3_LOGGER_NAME = "S3Logger" +UNREACHABLE_ANTHROPIC_BACKEND: Final = "anthropic/claude-haiku-4-5" @pytest.fixture(scope="session") @@ -64,6 +67,14 @@ def _assert_s3_configured(client: LoggingClient) -> None: class TestS3LogDelivery: @pytest.mark.covers("logging.s3.success.writes_object", exercised_on=["chat_completions"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + providers=(Provider.ANTHROPIC,), + models=(CHEAP_ANTHROPIC_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_chat_completions_writes_one_success_object( self, client: LoggingClient, s3_logs: S3LogReader, resources: ResourceManager ) -> None: @@ -108,6 +119,14 @@ class TestS3LogDelivery: ), f"payload response_cost {record.response_cost!r} must equal the header cost {outcome.response_cost}" @pytest.mark.covers("logging.s3.success.partition_layout", exercised_on=["chat_completions"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + providers=(Provider.ANTHROPIC,), + models=(CHEAP_ANTHROPIC_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_chat_completions_object_key_follows_the_partition_granularity( self, client: LoggingClient, s3_logs: S3LogReader, resources: ResourceManager ) -> None: @@ -148,6 +167,14 @@ class TestS3LogDelivery: ) @pytest.mark.covers("logging.s3.failure.writes_object", exercised_on=["chat_completions"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + providers=(Provider.ANTHROPIC,), + models=(UNREACHABLE_ANTHROPIC_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_chat_completions_failure_writes_one_object( self, client: LoggingClient, s3_logs: S3LogReader, resources: ResourceManager ) -> None: @@ -167,7 +194,7 @@ class TestS3LogDelivery: model_name = f"s3-err-{unique_marker()}" model_id = client.create_model( model_name, - LiteLLMParamsBody(model="anthropic/claude-haiku-4-5", api_key=INVALID_UPSTREAM_API_KEY), + LiteLLMParamsBody(model=UNREACHABLE_ANTHROPIC_BACKEND, api_key=INVALID_UPSTREAM_API_KEY), ) resources.defer(lambda: client.delete_model(model_id)) alias = f"s3-err-key-{unique_marker()}" diff --git a/tests/e2e/logging/test_team_langfuse_callback_e2e.py b/tests/e2e/logging/test_team_langfuse_callback_e2e.py index 89cd45c9f16..f2ece7f5d90 100644 --- a/tests/e2e/logging/test_team_langfuse_callback_e2e.py +++ b/tests/e2e/logging/test_team_langfuse_callback_e2e.py @@ -19,6 +19,7 @@ import time import pytest from e2e_config import CHEAP_ANTHROPIC_MODEL, unique_marker +from e2e_metadata import Domain, Mode, Provider, Subject, meta from lifecycle import ResourceManager from logging_client import ( LangfuseCreds, @@ -46,6 +47,14 @@ def langfuse_creds() -> LangfuseCreds: class TestTeamLangfuseCallback: @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["chat_completions"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + providers=(Provider.ANTHROPIC,), + models=(CHEAP_ANTHROPIC_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_team_callback_delivers_and_isolates( self, client: LoggingClient, langfuse_creds: LangfuseCreds, resources: ResourceManager ) -> None: diff --git a/tests/e2e/logging/test_weave_log_e2e.py b/tests/e2e/logging/test_weave_log_e2e.py index dab5993c87a..edd8ec7c3c3 100644 --- a/tests/e2e/logging/test_weave_log_e2e.py +++ b/tests/e2e/logging/test_weave_log_e2e.py @@ -21,11 +21,13 @@ query API; nothing is mocked. from __future__ import annotations import time +from typing import Final import pytest from e2e_config import CHEAP_ANTHROPIC_MODEL, unique_marker from e2e_http import StreamingResponse +from e2e_metadata import Domain, Mode, Provider, Subject, meta from lifecycle import ResourceManager from logging_client import ( INVALID_UPSTREAM_API_KEY, @@ -40,6 +42,8 @@ from weave_reader import WeaveCall, WeaveReader, build_weave_reader pytestmark = pytest.mark.e2e +UNREACHABLE_ANTHROPIC_BACKEND: Final = "anthropic/claude-haiku-4-5" + @pytest.fixture(scope="session") def weave_creds() -> WeaveCreds: @@ -79,6 +83,14 @@ WEAVE_STAGE_RED_REASON = ( class TestWeaveLogDelivery: @pytest.mark.skip(reason=WEAVE_STAGE_RED_REASON) @pytest.mark.covers("logging.niche_integrations.success.logs_spend", exercised_on=["chat_completions"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + providers=(Provider.ANTHROPIC,), + models=(CHEAP_ANTHROPIC_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_chat_completions_delivers_one_call_with_spend( self, client: LoggingClient, @@ -121,6 +133,14 @@ class TestWeaveLogDelivery: @pytest.mark.skip(reason=WEAVE_STAGE_RED_REASON) @pytest.mark.covers("logging.niche_integrations.failure.logs_spend", exercised_on=["chat_completions"]) + @meta( + Subject( + domain=Domain.OBSERVABILITY, + providers=(Provider.ANTHROPIC,), + models=(UNREACHABLE_ANTHROPIC_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_failed_chat_completions_delivers_one_error_call( self, client: LoggingClient, @@ -135,7 +155,7 @@ class TestWeaveLogDelivery: model_name = f"weave-err-{unique_marker()}" model_id = client.create_model( model_name, - LiteLLMParamsBody(model="anthropic/claude-haiku-4-5", api_key=INVALID_UPSTREAM_API_KEY), + LiteLLMParamsBody(model=UNREACHABLE_ANTHROPIC_BACKEND, api_key=INVALID_UPSTREAM_API_KEY), ) resources.defer(lambda: client.delete_model(model_id)) key = client.key_with_alias( diff --git a/tests/e2e/logging/weave_reader.py b/tests/e2e/logging/weave_reader.py index 2f8f759d299..c43e3280b39 100644 --- a/tests/e2e/logging/weave_reader.py +++ b/tests/e2e/logging/weave_reader.py @@ -28,7 +28,7 @@ import base64 import json import os import time -from dataclasses import dataclass +from dataclasses import dataclass, field from itertools import count, takewhile from typing import Final @@ -37,6 +37,7 @@ from pydantic import BaseModel, ConfigDict, Field from e2e_config import POLL_INTERVAL, POLL_TIMEOUT from e2e_http import URL, AuthHeaders, send +from e2e_metadata import step _WEAVE_TRACE_API: Final = "https://trace.wandb.ai" @@ -201,7 +202,7 @@ class WeaveCall(BaseModel): @dataclass(frozen=True, slots=True) class WeaveReader: project_id: str - api_key: str + api_key: str = field(repr=False) @property def _headers(self) -> AuthHeaders: @@ -232,6 +233,7 @@ class WeaveReader: ) return tuple(WeaveCall.model_validate_json(line) for line in outcome.body.splitlines() if line.strip()) + @step("Read the Weave {op} calls carrying the marker {marker}") def calls_matching(self, marker: str, *, since: float, op: str = LITELLM_REQUEST_OP) -> tuple[WeaveCall, ...]: """Every call under ``op`` started after ``since`` whose inputs carry ``marker``, paging until the window is exhausted. @@ -247,6 +249,7 @@ class WeaveReader: ) return tuple(call for page in pages for call in page if call.mentions(marker)) + @step("Wait for Weave to ingest a {op} call carrying the marker {marker}, then watch for duplicates") def poll_calls_matching(self, marker: str, *, since: float, op: str = LITELLM_REQUEST_OP) -> tuple[WeaveCall, ...]: """Poll until the call is readable, then keep re-reading for WEAVE_SETTLE_SECONDS so a duplicate exported by a later batch flush diff --git a/tests/e2e/management/jwt_actors.py b/tests/e2e/management/jwt_actors.py index 2d23549fe71..ae33f583195 100644 --- a/tests/e2e/management/jwt_actors.py +++ b/tests/e2e/management/jwt_actors.py @@ -5,6 +5,7 @@ from typing import Final, Literal from e2e_config import unique_marker from e2e_http import NoBody, unwrap +from e2e_metadata import step from idp import ADMIN_CLIENT_ID, TESTS_CLIENT_ID, Identity, Keycloak from lifecycle import ResourceManager from management.management_client import ManagementClient @@ -53,6 +54,7 @@ class Actor: profile: ActorProfile tenants: tuple[Tenant, ...] + @step("Get a JWT from the identity provider for the actor in the {self.role} role") def mint_caller(self, idp: Keycloak) -> Caller: return Caller( credential=idp.access_token( @@ -74,6 +76,7 @@ class ActorFactory: if self.bootstrap.proxy.caller is not None: raise ValueError("Actor bootstrap requires a separately held master client") + @step("Generate a virtual key as the proxy admin") def key(self, tenant: Tenant | None = None, *, user_id: str | None = None) -> KeyGenerateResponse: created: Final = unwrap( self.bootstrap.generate_key( @@ -87,6 +90,7 @@ class ActorFactory: self.resources.defer(lambda: self.bootstrap.delete_key_strict(created.key, missing_ok=True)) return created + @step("Create an organization, a team in it and a matching identity provider group") def tenant(self) -> Tenant: marker: Final = unique_marker() organization_id: Final = self.bootstrap.create_org(OrgNewBody(organization_alias=f"e2e-organization-{marker}")) @@ -118,6 +122,7 @@ class ActorFactory: self.resources.defer(lambda: self.idp.with_strict_cleanup().delete_group(group_id)) return Tenant(organization_id=organization_id, team_id=team_id, group_id=group_id) + @step("Create an actor in the {role} role, with its identity provider user, internal user and any tenant memberships") def create( self, role: ActorRole, *, tenants: tuple[Tenant, ...] = (), profile: ActorProfile = "database_role" ) -> Actor: diff --git a/tests/e2e/management/management_client.py b/tests/e2e/management/management_client.py index 3da9bea12a3..ffa310a9f4a 100644 --- a/tests/e2e/management/management_client.py +++ b/tests/e2e/management/management_client.py @@ -24,6 +24,7 @@ from e2e_http import ( retry_attempts, unwrap, ) +from e2e_metadata import STEP_FRAMES, step from models import ( AuditLogPage, AuditLogParams, @@ -116,9 +117,11 @@ class ManagementClient: def with_caller(self, caller: Caller) -> ManagementClient: return replace(self, proxy=self.proxy.with_caller(caller)) + @step("Generate a virtual key limited to the LLM API routes") def llm_only_key(self) -> str: return self.proxy.generate_key(KeyGenerateBody(models=[], allowed_routes=["llm_api_routes"])) + @step("Generate a virtual key with {body}") def generate_key(self, body: KeyGenerateBody, *, caller_key: str | None = None) -> Result[KeyGenerateResponse]: """POST /key/generate. `caller_key` is who is creating the key: the master key by default, or a virtual key (an admin filling in Create New Key on the @@ -133,6 +136,7 @@ class ManagementClient: response_type=KeyGenerateResponse, ) + @step("Update the virtual key's settings with /key/update") def update_key(self, body: KeyUpdateBody, *, caller_key: str | None = None) -> Result[NoBody]: """POST /key/update. `caller_key` is who is editing: the master key by default, or a virtual key (the dashboard edits under the session key its @@ -152,16 +156,22 @@ class ManagementClient: case UnknownApiError(body=error_body) if any( marker in error_body.lower() for marker in _TRANSIENT_BACKEND_MARKERS ): - warnings.warn(f"Transient backend response on attempt {attempt + 1}", RuntimeWarning, stacklevel=2) + warnings.warn( + f"Transient backend response on attempt {attempt + 1}", + RuntimeWarning, + stacklevel=2 + STEP_FRAMES, + ) time.sleep(0.5 * (attempt + 1)) continue case _: break return last + @step("Set the virtual key's models to [{models}]") def update_key_models(self, key: str, models: list[str]) -> None: _ = unwrap(self.update_key(KeyUpdateBody(key=key, models=models))) + @step("Delete the virtual key with the alias {key_alias}") def delete_key_by_alias(self, key_alias: str) -> None: _ = unwrap( self.proxy.transport.post( @@ -172,6 +182,7 @@ class ManagementClient: ) ) + @step("Read the key's deletion entries from the /audit log") def key_deleted_audit_logs(self, token_hash: str) -> AuditLogPage: return unwrap( self.proxy.transport.get( @@ -187,6 +198,7 @@ class ManagementClient: ) ) + @step("Read the key's settings back from /key/info") def key_info_as(self, key: str, *, caller_key: str | None = None) -> Result[KeyInfoResponse]: return self.proxy.transport.get( "/key/info", @@ -195,6 +207,7 @@ class ManagementClient: response_type=KeyInfoResponse, ) + @step("Delete the virtual key") def delete_key_strict(self, key: str, *, caller_key: str | None = None, missing_ok: bool = False) -> None: """Strict delete for the act phase of a test: a failed delete is a hard failure, unlike the warn-only ProxyClient.delete_key used at teardown.""" @@ -208,6 +221,7 @@ class ManagementClient: return _ = unwrap(result) + @step("Delete the deployment") def delete_model_strict(self, model_id: str) -> None: """Strict delete for the act phase of a test: a failed delete is a hard failure, unlike the warn-only ProxyClient.delete_model used at teardown.""" @@ -220,6 +234,7 @@ class ManagementClient: ) ) + @step("Run Test Connection on {body.litellm_params.model} in {body.mode} mode with /health/test_connection") def connection_test(self, body: ConnectionTestBody) -> Result[ConnectionTestResponse]: """POST /health/test_connection, the call behind the Admin UI's Test Connection button, probing the live provider with the supplied params.""" @@ -231,6 +246,7 @@ class ManagementClient: timeout=120.0, ) + @step("Block the virtual key") def block_key(self, key: str) -> None: _ = unwrap( self.proxy.transport.post( @@ -240,6 +256,7 @@ class ManagementClient: response_type=NoBody, ) ) + @step("Regenerate the virtual key with /key/regenerate") def regenerate_key(self, key: str, *, grace_period: str | None = None) -> str: return unwrap( self.proxy.transport.post( @@ -250,6 +267,7 @@ class ManagementClient: ) ).key + @step("Reset the virtual key's spend to {reset_to}") def reset_key_spend(self, key: str, reset_to: float) -> KeyResetSpendResponse: return unwrap( self.proxy.transport.post( @@ -260,6 +278,7 @@ class ManagementClient: ) ) + @step("List the keys with the alias {key_alias} from /key/list") def key_list(self, key_alias: str, *, caller_key: str | None = None) -> Result[KeyListResponse]: """GET /key/list, the Virtual Keys page's own inventory call. `caller_key` is who is asking: the master key by default, or a virtual key.""" @@ -271,9 +290,11 @@ class ManagementClient: response_type=KeyListResponse, ) + @step("Count the keys with the alias {key_alias} in /key/list") def key_alias_count(self, key_alias: str) -> int: return unwrap(self.key_list(key_alias)).total_count + @step("Sign in to the Admin UI with /v2/login") def dashboard_login(self, username: str, password: str) -> DashboardSession: """POST /v2/login, the call the Admin UI's sign-in form makes. @@ -297,6 +318,7 @@ class ManagementClient: redirect_url=response.redirect_url, ) + @step("Create a team with {body}") def create_team(self, body: TeamNewBody) -> str: team_id = unwrap( self.proxy.transport.post( @@ -309,6 +331,7 @@ class ManagementClient: self._wait_for_team(team_id) return team_id + @step("Update a team with {body}") def update_team(self, body: TeamUpdateBody) -> None: last: Result[NoBody] | None = None for attempt in range(retry_attempts(5)): @@ -324,7 +347,11 @@ class ManagementClient: case UnknownApiError(body=body_text) if ( "connecting to redis" in body_text.lower() or "name resolution" in body_text.lower() ): - warnings.warn(f"Transient backend response on attempt {attempt + 1}", RuntimeWarning, stacklevel=2) + warnings.warn( + f"Transient backend response on attempt {attempt + 1}", + RuntimeWarning, + stacklevel=2 + STEP_FRAMES, + ) time.sleep(0.5 * (attempt + 1)) continue case _: @@ -332,6 +359,7 @@ class ManagementClient: assert last is not None raise AssertionError(last) + @step("Delete the team") def delete_team(self, team_id: str) -> None: _ = self.proxy.transport.post( "/team/delete", @@ -340,6 +368,7 @@ class ManagementClient: response_type=NoBody, ) + @step("Read the team back from /team/info") def team_info(self, team_id: str) -> TeamData: return unwrap( self.proxy.transport.get( @@ -350,6 +379,7 @@ class ManagementClient: ) ).team_info + @step("List the teams from /team/list") def team_list_ids(self) -> tuple[str, ...]: return tuple( entry.team_id @@ -363,6 +393,7 @@ class ManagementClient: ).root ) + @step("Check whether /team/info finds the team") def team_info_status(self, team_id: str) -> ProbeResult: return self.proxy.transport.probe( "/team/info", params=TeamInfoParams(team_id=team_id), headers=self.proxy.management_headers() @@ -386,6 +417,7 @@ class ManagementClient: assert last is not None raise AssertionError(last) + @step("Add a user to the team with /team/member_add") def add_team_member(self, team_id: str, user_id: str) -> None: last: Result[NoBody] | None = None for attempt in range(retry_attempts(_TEAM_READY_ATTEMPTS)): @@ -402,7 +434,9 @@ class ManagementClient: _TEAM_READY_ATTEMPTS ): warnings.warn( - "Retrying team membership while the team becomes available", RuntimeWarning, stacklevel=2 + "Retrying team membership while the team becomes available", + RuntimeWarning, + stacklevel=2 + STEP_FRAMES, ) time.sleep(_TEAM_READY_SLEEP_SECONDS) continue @@ -411,6 +445,7 @@ class ManagementClient: assert last is not None raise AssertionError(last) + @step("Add a roster of members to the team with /team/member_add") def add_team_members(self, team_id: str, members: list[TeamMemberEntry]) -> None: """Bulk form of /team/member_add: `member` accepts a list, so one call seeds a whole roster the way an admin import does.""" @@ -423,6 +458,7 @@ class ManagementClient: ) ) + @step("Try to delete the team with /team/delete") def delete_team_status(self, team_id: str) -> StreamingResponse: """POST /team/delete judged by HTTP outcome: the raw status and body, so a test can assert on what a caller actually sees when the delete fails.""" @@ -432,6 +468,7 @@ class ManagementClient: json=TeamDeleteBody(team_ids=[team_id]), ) + @step("Remove a user from the team with /team/member_delete") def delete_team_member(self, team_id: str, user_id: str) -> None: _ = unwrap( self.proxy.transport.post( @@ -442,6 +479,7 @@ class ManagementClient: ) ) + @step("Create an internal user with {body}") def create_user(self, body: UserNewBody) -> str: return unwrap( self.proxy.transport.post( @@ -452,6 +490,7 @@ class ManagementClient: ) ).user_id + @step("Create the end user {user_id}") def create_customer(self, user_id: str) -> str: _ = unwrap( self.proxy.transport.post( @@ -463,6 +502,7 @@ class ManagementClient: ) return user_id + @step("Read the end user {end_user_id} back from /customer/info") def customer_info(self, end_user_id: str) -> CustomerResponse: return unwrap( self.proxy.transport.get( @@ -473,6 +513,7 @@ class ManagementClient: ) ) + @step("Delete the end user {user_id}") def delete_customer(self, user_id: str) -> None: _ = self.proxy.transport.post( "/customer/delete", @@ -481,6 +522,7 @@ class ManagementClient: response_type=NoBody, ) + @step("Update an internal user with {body}") def update_user(self, body: UserUpdateBody) -> None: _ = unwrap( self.proxy.transport.post( @@ -491,6 +533,7 @@ class ManagementClient: ) ) + @step("Delete the internal user") def delete_user(self, user_id: str) -> None: _ = self.proxy.transport.post( "/user/delete", @@ -499,6 +542,7 @@ class ManagementClient: response_type=NoBody, ) + @step("Delete the internal user") def delete_user_strict(self, user_id: str) -> None: """Strict delete for the act phase of a test: a failed delete is a hard failure, unlike the warn-only delete_user used at teardown.""" @@ -511,6 +555,7 @@ class ManagementClient: ) ) + @step("Read the user back from /user/info") def user_info(self, user_id: str | None = None) -> UserInfoResponse: return unwrap( self.proxy.transport.get( @@ -521,6 +566,7 @@ class ManagementClient: ) ) + @step("Count the matching users in /user/list") def user_count(self, user_id: str) -> int: return unwrap( self.proxy.transport.get( @@ -531,6 +577,7 @@ class ManagementClient: ) ).total + @step("List the matching users from /user/list") def user_list_ids(self, user_id: str) -> tuple[str, ...]: listing = unwrap( self.proxy.transport.get( @@ -542,6 +589,7 @@ class ManagementClient: ) return tuple(row.user_id for row in listing.users) + @step("Create an organization with {body}") def create_org(self, body: OrgNewBody) -> str: return unwrap( self.proxy.transport.post( @@ -552,6 +600,7 @@ class ManagementClient: ) ).organization_id + @step("Update an organization with {body}") def update_org(self, body: OrgUpdateBody) -> None: _ = unwrap( self.proxy.transport.patch( @@ -562,6 +611,7 @@ class ManagementClient: ) ) + @step("Delete the organization") def delete_org(self, organization_id: str) -> None: _ = self.proxy.transport.delete( "/organization/delete", @@ -570,6 +620,7 @@ class ManagementClient: response_type=NoBody, ) + @step("Read the organization back from /organization/info") def org_info(self, organization_id: str) -> OrgInfoResponse: return unwrap( self.proxy.transport.get( @@ -580,6 +631,7 @@ class ManagementClient: ) ) + @step("Check whether /organization/info finds the organization") def org_info_status(self, organization_id: str) -> ProbeResult: return self.proxy.transport.probe( "/organization/info", @@ -587,6 +639,7 @@ class ManagementClient: headers=self.proxy.management_headers(), ) + @step("Create a tag with {body}") def create_tag(self, body: TagNewBody) -> None: _ = unwrap( self.proxy.transport.post( @@ -597,6 +650,7 @@ class ManagementClient: ) ) + @step("Delete the tag {name}") def delete_tag(self, name: str) -> None: _ = self.proxy.transport.post( "/tag/delete", @@ -605,6 +659,7 @@ class ManagementClient: response_type=NoBody, ) + @step("List the tags from /tag/list") def tag_list(self) -> tuple[TagListEntry, ...]: return tuple( unwrap( @@ -617,6 +672,7 @@ class ManagementClient: ).root ) + @step("Create an MCP server named {body.alias}") def create_mcp_server(self, body: McpServerCreateBody) -> McpServerRow: return unwrap( self.proxy.transport.post( @@ -627,6 +683,7 @@ class ManagementClient: ) ) + @step("Update the MCP server's settings with PUT /v1/mcp/server") def update_mcp_server(self, body: McpServerUpdateBody) -> McpServerRow: """PUT /v1/mcp/server, the call behind the dashboard's Save Changes: a partial update where a field left unset keeps its stored value and None clears it.""" @@ -639,6 +696,7 @@ class ManagementClient: ) ) + @step("Delete the MCP server") def delete_mcp_server(self, server_id: str) -> Result[NoBody]: """DELETE /v1/mcp/server/{server_id}. Returns the outcome so the act phase can unwrap it while a deferred teardown can ignore an already-deleted server.""" @@ -649,6 +707,7 @@ class ManagementClient: response_type=NoBody, ) + @step("Send a /chat/completions request to {model}") def chat_status(self, key: str, model: str, content: str) -> StreamingResponse: return self.proxy.transport.send( "/chat/completions", @@ -656,12 +715,15 @@ class ManagementClient: json=ChatBody(model=model, messages=[ChatMessage(role="user", content=content)], max_tokens=16), ) + @step("Try to generate a virtual key with {body}, calling with a virtual key") def key_generate_status(self, key: str, body: KeyGenerateBody) -> StreamingResponse: return self.proxy.transport.send("/key/generate", headers=self.proxy.transport.bearer(key), json=body) + @step("Try to create a team with {body}, calling with a virtual key") def team_new_status(self, key: str, body: TeamNewBody) -> StreamingResponse: return self.proxy.transport.send("/team/new", headers=self.proxy.transport.bearer(key), json=body) + @step("Try to create an internal user with {body}, calling with a virtual key") def user_new_status(self, key: str, body: UserNewBody) -> StreamingResponse: return self.proxy.transport.send("/user/new", headers=self.proxy.transport.bearer(key), json=body) diff --git a/tests/e2e/management/test_budget_customer_user_org_e2e.py b/tests/e2e/management/test_budget_customer_user_org_e2e.py index 6e14a2d5745..ae4e6db4e15 100644 --- a/tests/e2e/management/test_budget_customer_user_org_e2e.py +++ b/tests/e2e/management/test_budget_customer_user_org_e2e.py @@ -23,6 +23,7 @@ from pydantic import BaseModel, Field, RootModel from e2e_config import unique_marker from e2e_http import NoBody, Success, UnauthorizedError, UnknownApiError, is_ok, unwrap +from e2e_metadata import Domain, Route, Subject, meta from lifecycle import ResourceManager from management_client import ManagementClient from models import KeyGenerateBody, ModelBudgetEntry, OrgInfoParams, OrgNewBody, UserNewBody @@ -152,6 +153,7 @@ _UPDATED_MAX_BUDGET = 91.25 class TestBudgetManagement: @pytest.mark.covers("mgmt.budget.list.happy_path") + @meta(Subject(domain=Domain.SPEND_BUDGETS, route=Route.BUDGET_MANAGEMENT)) def test_created_budget_appears_in_budget_list( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -164,6 +166,7 @@ class TestBudgetManagement: ) @pytest.mark.covers("mgmt.budget.update.accepts_model_max_budget") + @meta(Subject(domain=Domain.SPEND_BUDGETS, route=Route.BUDGET_MANAGEMENT)) def test_update_accepts_per_model_budgets_including_punctuated_names( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -213,6 +216,7 @@ class TestBudgetManagement: ) @pytest.mark.covers("mgmt.budget.update.persists") + @meta(Subject(domain=Domain.SPEND_BUDGETS, route=Route.BUDGET_MANAGEMENT)) def test_update_max_budget_persists_to_budget_info( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -247,6 +251,7 @@ class TestBudgetManagement: ) @pytest.mark.covers("mgmt.budget.new.admin_only") + @meta(Subject(domain=Domain.PROXY_AUTH, route=Route.BUDGET_MANAGEMENT)) def test_new_is_refused_for_a_non_admin_key( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -352,6 +357,7 @@ class TestBudgetListV1: """ @pytest.mark.covers("mgmt.budget.list_v1.happy_path") + @meta(Subject(domain=Domain.SPEND_BUDGETS, route=Route.BUDGET_MANAGEMENT)) def test_sorts_pages_and_filters_the_budgets_it_created( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -392,6 +398,7 @@ class TestBudgetListV1: assert [row.tpm_limit for row in limits] == [60000, 60000, 60000] @pytest.mark.covers("mgmt.budget.list_v1.happy_path") + @meta(Subject(domain=Domain.SPEND_BUDGETS, route=Route.BUDGET_MANAGEMENT)) def test_is_null_finds_the_budget_left_uncapped( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -409,11 +416,13 @@ class TestBudgetListV1: assert [row.max_budget for row in _list_budgets(client, found).data] == [None] @pytest.mark.covers("mgmt.budget.list_v1.happy_path") + @meta(Subject(domain=Domain.SPEND_BUDGETS, route=Route.BUDGET_MANAGEMENT)) def test_refuses_a_sort_field_and_a_parameter_it_does_not_support(self, client: ManagementClient) -> None: assert _list_status(client, BudgetPageParams(sort="budget_duration")) == 400 assert _list_status(client, BudgetPageParams(not_a_parameter="b-1")) == 400 @pytest.mark.covers("mgmt.budget.list_v1.admin_only") + @meta(Subject(domain=Domain.PROXY_AUTH, route=Route.BUDGET_MANAGEMENT)) def test_is_refused_for_a_non_admin_key(self, client: ManagementClient, resources: ResourceManager) -> None: key = client.proxy.generate_key(KeyGenerateBody()) resources.defer(lambda: client.proxy.delete_key(key)) @@ -482,6 +491,7 @@ def _customer_info(client: ManagementClient, route: str, user_id: str) -> Custom class TestCustomerManagement: @pytest.mark.covers("mgmt.customer.new.happy_path") + @meta(Subject(domain=Domain.MANAGEMENT, route=Route.CUSTOMER_MANAGEMENT)) def test_new_persists_to_customer_info(self, client: ManagementClient, resources: ResourceManager) -> None: customer_id = f"e2e-mgmt-cust-{unique_marker()}" created = _create_customer( @@ -495,6 +505,7 @@ class TestCustomerManagement: ) @pytest.mark.covers("mgmt.customer.delete.persists") + @meta(Subject(domain=Domain.MANAGEMENT, route=Route.CUSTOMER_MANAGEMENT)) def test_delete_removes_the_customer(self, client: ManagementClient, resources: ResourceManager) -> None: """The teardown's deferred delete fires again on the already-deleted customer by design: it is the safety net if this test fails before the in-body delete, @@ -524,6 +535,7 @@ class TestCustomerManagement: _ = _poll(client, gone, f"customer {customer_id} still resolved on /customer/info after /customer/delete") @pytest.mark.covers("mgmt.end_user.new.happy_path") + @meta(Subject(domain=Domain.MANAGEMENT, route=Route.CUSTOMER_MANAGEMENT)) def test_end_user_new_persists_to_end_user_info( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -542,6 +554,7 @@ class TestCustomerManagement: class TestUserManagement: @pytest.mark.covers("mgmt.user.info.happy_path") + @meta(Subject(domain=Domain.MANAGEMENT, route=Route.USER_MANAGEMENT)) def test_new_user_is_readable_via_user_info( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -586,6 +599,7 @@ class OrgInfoMembersResponse(BaseModel): class TestOrganizationMembership: @pytest.mark.covers("mgmt.organization.member_add.happy_path") + @meta(Subject(domain=Domain.MANAGEMENT, route=Route.ORGANIZATION_MANAGEMENT)) def test_member_add_records_membership( self, client: ManagementClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/management/test_config_misc_endpoints_e2e.py b/tests/e2e/management/test_config_misc_endpoints_e2e.py index a3be0a64e7f..b7ab311cc18 100644 --- a/tests/e2e/management/test_config_misc_endpoints_e2e.py +++ b/tests/e2e/management/test_config_misc_endpoints_e2e.py @@ -32,6 +32,7 @@ from pydantic import BaseModel from e2e_config import unique_marker from e2e_http import NoBody, Success, unwrap, unwrap_status +from e2e_metadata import Domain, Route, Subject, meta from lifecycle import ResourceManager from management_client import ManagementClient from models import KeyGenerateBody, LiteLLMParamsBody, TeamNewBody @@ -231,6 +232,7 @@ class McpServerResponse(BaseModel): class TestInventoryRoutes: @pytest.mark.covers("mgmt.callback.list.happy_path") + @meta(Subject(domain=Domain.OBSERVABILITY)) def test_callbacks_list_reports_active_logging_callbacks(self, client: ManagementClient) -> None: listing = unwrap( client.proxy.transport.get( @@ -247,6 +249,7 @@ class TestInventoryRoutes: ) @pytest.mark.covers("mgmt.tool_management.list.happy_path") + @meta(Subject(domain=Domain.MANAGEMENT)) def test_tool_list_returns_catalog_with_consistent_total(self, client: ManagementClient) -> None: listing = unwrap( client.proxy.transport.get( @@ -261,6 +264,7 @@ class TestInventoryRoutes: ) @pytest.mark.covers("mgmt.workflow.list.happy_path") + @meta(Subject(domain=Domain.MANAGEMENT)) def test_workflow_runs_list_returns_consistent_count(self, client: ManagementClient) -> None: listing = unwrap( client.proxy.transport.get( @@ -275,6 +279,7 @@ class TestInventoryRoutes: ) @pytest.mark.covers("mgmt.credential_migration.check.happy_path") + @meta(Subject(domain=Domain.DEPLOY_OPS)) def test_credential_migration_check_reports_residual_scan(self, client: ManagementClient) -> None: report = unwrap( client.proxy.transport.get( @@ -295,6 +300,7 @@ class TestInventoryRoutes: class TestCostEstimate: @pytest.mark.covers("mgmt.cost_tracking.estimate.happy_path") + @meta(Subject(domain=Domain.COST_MAP)) def test_estimate_computes_cost_from_token_counts(self, client: ManagementClient) -> None: estimate = unwrap( client.proxy.transport.post( @@ -325,6 +331,7 @@ class TestCostEstimate: class TestComplianceRoutes: @pytest.mark.covers("mgmt.compliance.gdpr.happy_path") + @meta(Subject(domain=Domain.MANAGEMENT)) def test_gdpr_check_derives_verdict_from_the_request(self, client: ManagementClient) -> None: result = unwrap( client.proxy.transport.post( @@ -356,6 +363,7 @@ class TestComplianceRoutes: class TestFallbackManagement: @pytest.mark.covers("mgmt.fallback_management.update.happy_path") + @meta(Subject(domain=Domain.ROUTING)) def test_create_persists_and_is_read_back(self, client: ManagementClient, resources: ResourceManager) -> None: primary = f"e2e-fallback-primary-{unique_marker()}" secondary = f"e2e-fallback-secondary-{unique_marker()}" @@ -410,6 +418,7 @@ class TestFallbackManagement: class TestJwtKeyMapping: @pytest.mark.covers("mgmt.jwt_key_mapping.new.happy_path") + @meta(Subject(domain=Domain.PROXY_AUTH)) def test_new_persists_mapping_and_is_read_back( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -461,6 +470,7 @@ class TestJwtKeyMapping: class TestRouterSettings: @pytest.mark.covers("mgmt.router_settings.update.happy_path") + @meta(Subject(domain=Domain.ROUTING, route=Route.PROXY_CONFIG)) def test_config_update_persists_router_setting_to_get( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -531,6 +541,7 @@ class TestRouterSettings: class TestMcpServerSubmission: @pytest.mark.covers("mgmt.mcp_server.register.happy_path") + @meta(Subject(domain=Domain.MCP, route=Route.MCP)) def test_register_submits_pending_server(self, client: ManagementClient, resources: ResourceManager) -> None: """A non-admin, team-scoped key submits an MCP server for review; the proxy stores it as pending_review without loading it into the runtime registry.""" @@ -564,6 +575,7 @@ class TestMcpServerSubmission: ) @pytest.mark.covers("mgmt.mcp_server.approve.persists") + @meta(Subject(domain=Domain.MCP, route=Route.MCP)) def test_approve_activates_submission_and_persists( self, client: ManagementClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/management/test_jwt_management_e2e.py b/tests/e2e/management/test_jwt_management_e2e.py index 5898073a4e6..3853ae2c4e6 100644 --- a/tests/e2e/management/test_jwt_management_e2e.py +++ b/tests/e2e/management/test_jwt_management_e2e.py @@ -7,6 +7,7 @@ from typing import Final, Literal import pytest from e2e_config import CHEAP_OPENAI_MODEL, PROXY_BASE_URL, unique_marker from e2e_http import UnauthorizedError, UnknownApiError, unwrap +from e2e_metadata import Domain, Route, Subject, meta from idp import ADMIN_CLIENT_ID, Identity, Keycloak, token_claims from lifecycle import ResourceManager from management.jwt_actors import ActorFactory, ActorRole @@ -32,6 +33,7 @@ class TestJwtManagement: ), ) @pytest.mark.covers("mgmt.user.jwt.database_roles") + @meta(Subject(domain=Domain.PROXY_AUTH, route=Route.USER_MANAGEMENT)) def test_actor_subject_and_database_role(self, actor_factory: ActorFactory, role: ActorRole) -> None: tenants: Final = ( (actor_factory.tenant(),) if role in ("organization_admin", "team_admin", "team_member") else () @@ -70,6 +72,7 @@ class TestJwtManagement: } == {(actor.identity.user_id, "org_admin" if role == "organization_admin" else "internal_user")} @pytest.mark.covers("mgmt.key.jwt.viewer_denied") + @meta(Subject(domain=Domain.PROXY_AUTH, route=Route.KEY_MANAGEMENT)) def test_admin_viewer_reads_but_cannot_update(self, actor_factory: ActorFactory) -> None: actor: Final = actor_factory.create("proxy_admin_viewer") viewer: Final = actor_factory.bootstrap.with_caller(actor.mint_caller(actor_factory.idp)) @@ -83,6 +86,7 @@ class TestJwtManagement: assert actor_factory.bootstrap.proxy.key_info(key).key_alias == alias @pytest.mark.covers("mgmt.user.oidc.identity_mapping") + @meta(Subject(domain=Domain.PROXY_AUTH)) def test_oidc_browser_profile_identity_mapping(self, actor_factory: ActorFactory) -> None: actor: Final = actor_factory.create("internal_user") idp: Final = actor_factory.idp.with_strict_cleanup() @@ -100,6 +104,7 @@ class TestJwtManagement: @pytest.mark.covers("mgmt.key.jwt.lifecycle") @pytest.mark.parametrize("credential_kind", ("direct_jwt", "virtual_key")) + @meta(Subject(domain=Domain.PROXY_AUTH, route=Route.KEY_MANAGEMENT)) def test_admin_creates_reads_updates_clears_and_deletes_a_key( self, actor_factory: ActorFactory, @@ -142,6 +147,7 @@ class TestJwtManagement: assert unwrap(bound.key_list(updated_alias)).total_count == 0 @pytest.mark.covers("mgmt.team.jwt.tenant_isolation") + @meta(Subject(domain=Domain.PROXY_AUTH, route=Route.KEY_MANAGEMENT)) def test_two_actor_sets_keep_tenants_and_keys_isolated(self, actor_factory: ActorFactory) -> None: first: Final = actor_factory.tenant() second: Final = actor_factory.tenant() @@ -163,6 +169,7 @@ class TestJwtManagement: assert tuple(actor.identity.groups for actor in actors) == ((first.team_id,), (second.team_id,)) @pytest.mark.covers("mgmt.team.jwt.multiple_memberships") + @meta(Subject(domain=Domain.PROXY_AUTH, route=Route.TEAM_MANAGEMENT)) def test_multi_group_actor_keeps_exact_memberships(self, actor_factory: ActorFactory) -> None: tenants: Final = (actor_factory.tenant(), actor_factory.tenant()) actor: Final = actor_factory.create("team_member", tenants=tenants, profile="group_scoped") @@ -177,6 +184,7 @@ class TestJwtManagement: } == {(actor.identity.user_id, "user")} @pytest.mark.covers("mgmt.user.jwt.cleanup") + @meta(Subject(domain=Domain.PROXY_AUTH)) def test_successful_actor_cleanup_removes_owned_state(self, actor_factory: ActorFactory) -> None: resources: Final = ResourceManager(client=actor_factory.bootstrap.proxy, strict_cleanup=True) factory: Final = ActorFactory(bootstrap=actor_factory.bootstrap, idp=actor_factory.idp, resources=resources) @@ -197,6 +205,7 @@ class TestJwtManagement: @pytest.mark.parametrize("stage", ("group", "user")) @pytest.mark.covers("mgmt.user.jwt.partial_cleanup") + @meta(Subject(domain=Domain.PROXY_AUTH)) def test_partial_setup_removes_previously_created_identities( self, actor_factory: ActorFactory, @@ -236,6 +245,7 @@ class TestJwtManagement: idp.assert_absent("users", identity.user_id) @pytest.mark.covers("mgmt.key.jwt.member_denied", "mgmt.key.jwt.other_team_denied") + @meta(Subject(domain=Domain.PROXY_AUTH, route=Route.KEY_MANAGEMENT)) def test_member_cannot_write_and_another_team_cannot_read_the_key( self, client: ManagementClient, idp: Keycloak, jwt_identity: Identity, resources: ResourceManager ) -> None: diff --git a/tests/e2e/management/test_key_lifecycle_e2e.py b/tests/e2e/management/test_key_lifecycle_e2e.py index fb153f2a7f3..956719e0ba3 100644 --- a/tests/e2e/management/test_key_lifecycle_e2e.py +++ b/tests/e2e/management/test_key_lifecycle_e2e.py @@ -23,6 +23,7 @@ import pytest from e2e_config import unique_marker from e2e_http import Result, StreamingResponse, Success, unwrap +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from management_client import MODEL_ACCESS_DENIED_MARKER, ManagementClient from models import ( @@ -188,6 +189,7 @@ def _assert_chat_rejected_everywhere(client: ManagementClient, key: str, model: class TestKeyLifecycle: + @meta(Subject(domain=Domain.MANAGEMENT, route=Route.KEY_MANAGEMENT)) def test_create_echoes_every_field_written( self, client: ManagementClient, resources: ResourceManager, mock_deployment: str ) -> None: @@ -206,6 +208,7 @@ class TestKeyLifecycle: ): assert observed == wanted, f"/key/generate echoed {field}={observed!r}, sent {wanted!r}" + @meta(Subject(domain=Domain.MANAGEMENT, route=Route.KEY_MANAGEMENT)) def test_read_reflects_the_create_on_every_replica( self, client: ManagementClient, resources: ResourceManager, mock_deployment: str ) -> None: @@ -221,6 +224,7 @@ class TestKeyLifecycle: ) @pytest.mark.covers("mgmt.key.update.preserves_unrelated_fields") + @meta(Subject(domain=Domain.MANAGEMENT, route=Route.KEY_MANAGEMENT)) def test_partial_update_changes_only_the_named_field( self, client: ManagementClient, resources: ResourceManager, mock_deployment: str ) -> None: @@ -238,6 +242,7 @@ class TestKeyLifecycle: ) @pytest.mark.covers("mgmt.key.update.clear_persists") + @meta(Subject(domain=Domain.MANAGEMENT, route=Route.KEY_MANAGEMENT)) def test_explicit_null_clears_the_budget_and_its_reset_time( self, client: ManagementClient, resources: ResourceManager, mock_deployment: str ) -> None: @@ -258,6 +263,14 @@ class TestKeyLifecycle: info, created.written.model_copy(update={"max_budget": None, "budget_duration": None}), replica ) + @meta( + Subject( + domain=Domain.PROXY_AUTH, + providers=(Provider.OPENAI,), + models=(BACKING_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_key_serves_its_model_and_is_denied_others( self, client: ManagementClient, resources: ResourceManager, mock_deployment: str ) -> None: @@ -274,6 +287,15 @@ class TestKeyLifecycle: f"403 body must be a model-access denial, got: {denied.body[:300]}" ) + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.KEY_MANAGEMENT, + providers=(Provider.OPENAI,), + models=(BACKING_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_delete_revokes_info_and_chat_on_every_replica( self, client: ManagementClient, resources: ResourceManager, mock_deployment: str ) -> None: diff --git a/tests/e2e/management/test_key_management_e2e.py b/tests/e2e/management/test_key_management_e2e.py index 39a9e657b8c..34f912fe604 100644 --- a/tests/e2e/management/test_key_management_e2e.py +++ b/tests/e2e/management/test_key_management_e2e.py @@ -19,6 +19,7 @@ import pytest from e2e_config import unique_marker from e2e_http import NoBody, StreamingResponse, unwrap +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from management_client import ManagementClient from models import ( @@ -31,6 +32,7 @@ pytestmark = pytest.mark.e2e TINY_BUDGET = 3e-6 SPEND_MODEL = "claude-haiku-4-5" +SYNTHETIC_BACKEND: Final = "openai/synthetic-detachment" class KeyToggleBlockBody(BaseModel): @@ -161,13 +163,22 @@ def project_resources(client: ManagementClient) -> Iterator[ResourceManager]: class TestKeyManagementRoutes: @pytest.mark.covers("mgmt.key.update.persists") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.KEY_MANAGEMENT, + providers=(Provider.OPENAI,), + models=(SYNTHETIC_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_project_detachment_preserves_key_scope_and_refreshes_auth( self, client: ManagementClient, project_resources: ResourceManager ) -> None: resources: Final = project_resources name: Final = f"e2e-detach-{unique_marker()}" model_id: Final = client.proxy.create_model( - name, LiteLLMParamsBody(model="openai/synthetic-detachment", api_key="synthetic", mock_response="orbit") + name, LiteLLMParamsBody(model=SYNTHETIC_BACKEND, api_key="synthetic", mock_response="orbit") ) resources.defer(lambda: client.proxy.delete_model(model_id)) org_id: Final = client.create_org(OrgNewBody(organization_alias=name, models=[name])) @@ -222,6 +233,7 @@ class TestKeyManagementRoutes: assert denied.status_code in (401, 403), denied.body @pytest.mark.covers("mgmt.key.info.persists") + @meta(Subject(domain=Domain.MANAGEMENT, route=Route.KEY_MANAGEMENT)) def test_info_reflects_the_fields_the_key_was_created_with( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -246,6 +258,7 @@ class TestKeyManagementRoutes: assert info.rpm_limit == 141414, f"/key/info reports rpm_limit {info.rpm_limit}, configured 141414" @pytest.mark.covers("mgmt.key.unblock.persists") + @meta(Subject(domain=Domain.MANAGEMENT, route=Route.KEY_MANAGEMENT)) def test_unblock_flips_key_info_blocked_back( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -266,6 +279,7 @@ class TestKeyManagementRoutes: ) @pytest.mark.covers("mgmt.key.health.happy_path") + @meta(Subject(domain=Domain.MANAGEMENT, route=Route.KEY_MANAGEMENT)) def test_health_reports_the_calling_key_healthy( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -285,6 +299,7 @@ class TestKeyManagementRoutes: ) @pytest.mark.covers("mgmt.key.bulk_update.happy_path") + @meta(Subject(domain=Domain.MANAGEMENT, route=Route.KEY_MANAGEMENT)) def test_bulk_update_applies_max_budget_to_target_key( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -314,6 +329,15 @@ class TestKeyManagementRoutes: ) @pytest.mark.covers("other.key_mgmt.spend_reset.resets_to_value") + @meta( + Subject( + domain=Domain.SPEND_BUDGETS, + route=Route.KEY_MANAGEMENT, + providers=(Provider.ANTHROPIC,), + models=(SPEND_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_reset_spend_zeroes_recorded_spend_and_lifts_the_budget_block( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -340,6 +364,7 @@ class TestKeyManagementRoutes: _ = _poll(client, call_allowed_again, "the key stayed budget-blocked after its spend was reset to 0") @pytest.mark.covers("mgmt.key.generate.admin_only") + @meta(Subject(domain=Domain.PROXY_AUTH, route=Route.KEY_MANAGEMENT)) def test_generate_forbidden_for_non_admin_key( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -355,6 +380,7 @@ class TestKeyManagementRoutes: ) @pytest.mark.covers("mgmt.key.delete.admin_only") + @meta(Subject(domain=Domain.PROXY_AUTH, route=Route.KEY_MANAGEMENT)) def test_delete_forbidden_for_non_admin_key( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -374,6 +400,7 @@ class TestKeyManagementRoutes: ) @pytest.mark.covers("mgmt.key.update.admin_only") + @meta(Subject(domain=Domain.PROXY_AUTH, route=Route.KEY_MANAGEMENT)) def test_update_forbidden_for_non_admin_key( self, client: ManagementClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/management/test_management_e2e.py b/tests/e2e/management/test_management_e2e.py index 908eb752611..c167a1323cb 100644 --- a/tests/e2e/management/test_management_e2e.py +++ b/tests/e2e/management/test_management_e2e.py @@ -18,6 +18,7 @@ import pytest from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, UI_PASSWORD, UI_USERNAME, unique_marker from e2e_http import StreamingResponse, Success, unwrap +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from management_client import ( DASHBOARD_SESSION_TEAM_ID, @@ -47,6 +48,8 @@ from proxy_client import Converged, await_converged pytestmark = pytest.mark.e2e +GEMINI_MODEL: Final = "gemini-2.5-flash" +OPENAI_MODEL: Final = "gpt-5.5" REGENERATE_GRACE_PERIOD = "15s" REGENERATE_GRACE_SECONDS = 15.0 TEAM_DELETE_POOL_OVERFLOW_MEMBERS = 250 @@ -130,6 +133,15 @@ def _poll_model_access_granted(client: ManagementClient, key: str, model: str) - class TestKeyRoutes: @pytest.mark.covers("mgmt.key.generate.persists") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.KEY_MANAGEMENT, + providers=(Provider.GEMINI,), + models=(GEMINI_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_generate_persists_to_key_info_and_scopes_chat( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -137,12 +149,12 @@ class TestKeyRoutes: key = _generate_key( client, resources, - KeyGenerateBody(models=["gemini-2.5-flash"], key_alias=alias, tpm_limit=424242, rpm_limit=424243), + KeyGenerateBody(models=[GEMINI_MODEL], key_alias=alias, tpm_limit=424242, rpm_limit=424243), ) info = client.proxy.key_info(key) assert info.key_alias == alias, f"/key/info reports key_alias {info.key_alias!r}, configured {alias!r}" - assert info.models == ["gemini-2.5-flash"], ( + assert info.models == [GEMINI_MODEL], ( f"/key/info reports models {info.models}, configured ['gemini-2.5-flash']" ) assert info.tpm_limit == 424242, ( @@ -152,49 +164,73 @@ class TestKeyRoutes: f"/key/info reports rpm_limit {info.rpm_limit}, configured 424243" ) - _poll_chat_ok(client, key, "gemini-2.5-flash") + _poll_chat_ok(client, key, GEMINI_MODEL) _assert_model_denied( - client.chat_status(key, "gpt-5.5", f"say hi {unique_marker()}"), "gpt-5.5" + client.chat_status(key, OPENAI_MODEL, f"say hi {unique_marker()}"), OPENAI_MODEL ) @pytest.mark.covers("mgmt.key.update.persists") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.KEY_MANAGEMENT, + providers=(Provider.GEMINI, Provider.OPENAI), + models=(GEMINI_MODEL, OPENAI_MODEL), + mode=Mode.NONSTREAM, + ) + ) def test_update_models_persists_and_flips_enforcement( self, client: ManagementClient, resources: ResourceManager ) -> None: - key = _generate_key(client, resources, KeyGenerateBody(models=["gemini-2.5-flash"])) - _poll_chat_ok(client, key, "gemini-2.5-flash") + key = _generate_key(client, resources, KeyGenerateBody(models=[GEMINI_MODEL])) + _poll_chat_ok(client, key, GEMINI_MODEL) _assert_model_denied( - client.chat_status(key, "gpt-5.5", f"say hi {unique_marker()}"), "gpt-5.5" + client.chat_status(key, OPENAI_MODEL, f"say hi {unique_marker()}"), OPENAI_MODEL ) - client.update_key_models(key, ["gpt-5.5"]) + client.update_key_models(key, [OPENAI_MODEL]) info = client.proxy.key_info(key) - assert info.models == ["gpt-5.5"], ( + assert info.models == [OPENAI_MODEL], ( f"/key/info reports models {info.models} after /key/update to ['gpt-5.5']" ) - _poll_model_access_granted(client, key, "gpt-5.5") - _poll_chat_denied(client, key, "gemini-2.5-flash") + _poll_model_access_granted(client, key, OPENAI_MODEL) + _poll_chat_denied(client, key, GEMINI_MODEL) @pytest.mark.covers("mgmt.key.delete.persists") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.KEY_MANAGEMENT, + providers=(Provider.GEMINI,), + models=(GEMINI_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_delete_revokes_the_key_on_chat(self, client: ManagementClient, resources: ResourceManager) -> None: """The teardown's deferred delete fires again on the already-deleted key by design: the deferred cleanup must survive this test failing before the in-body delete, and a repeat /key/delete is a cheap no-op the warn-only teardown absorbs.""" - key = _generate_key(client, resources, KeyGenerateBody(models=["gemini-2.5-flash"])) - _poll_chat_ok(client, key, "gemini-2.5-flash") + key = _generate_key(client, resources, KeyGenerateBody(models=[GEMINI_MODEL])) + _poll_chat_ok(client, key, GEMINI_MODEL) client.delete_key_strict(key) def rejected() -> bool | None: - outcome = client.chat_status(key, "gemini-2.5-flash", f"say hi {unique_marker()}") + outcome = client.chat_status(key, GEMINI_MODEL, f"say hi {unique_marker()}") return True if outcome.status_code == 401 else None _ = _poll(client, rejected, "deleted key was still accepted on chat (never rejected 401) at the deadline") @pytest.mark.covers("mgmt.key.list.happy_path") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.KEY_MANAGEMENT, + ) + ) def test_created_key_appears_in_key_list_inventory( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -214,8 +250,14 @@ class TestKeyRoutes: @pytest.mark.covers("mgmt.key.block.persists") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.KEY_MANAGEMENT, + ) + ) def test_block_persists_to_key_info(self, client: ManagementClient, resources: ResourceManager) -> None: - key = _generate_key(client, resources, KeyGenerateBody(models=["gemini-2.5-flash"])) + key = _generate_key(client, resources, KeyGenerateBody(models=[GEMINI_MODEL])) assert not client.proxy.key_info(key).blocked, "/key/info reports the key blocked before /key/block ran" client.block_key(key) @@ -233,6 +275,15 @@ class TestDashboardKeyRoutes: are the same routes the API-surface tests cover with a different caller.""" @pytest.mark.covers("mgmt.key.generate.happy_path") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.KEY_MANAGEMENT, + providers=(Provider.GEMINI,), + models=(GEMINI_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_creating_a_key_from_the_dashboard_persists_and_works( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -260,7 +311,7 @@ class TestDashboardKeyRoutes: def dashboard_creates_the_key() -> str | None: match client.generate_key( - KeyGenerateBody(models=["gemini-2.5-flash"], key_alias=alias, tpm_limit=100), + KeyGenerateBody(models=[GEMINI_MODEL], key_alias=alias, tpm_limit=100), caller_key=session.session_key, ): case Success(data=created): @@ -280,7 +331,7 @@ class TestDashboardKeyRoutes: f"/key/info reports key_alias {created_info.key_alias!r} for the key the dashboard created, " f"expected {alias!r}" ) - assert created_info.models == ["gemini-2.5-flash"], ( + assert created_info.models == [GEMINI_MODEL], ( f"/key/info reports models {created_info.models} for the key the dashboard created" ) assert created_info.tpm_limit == 100, ( @@ -301,10 +352,19 @@ class TestDashboardKeyRoutes: "would render no keys", ) - _poll_chat_ok(client, created, "gemini-2.5-flash") - _assert_model_denied(client.chat_status(created, "gpt-5.5", f"say hi {unique_marker()}"), "gpt-5.5") + _poll_chat_ok(client, created, GEMINI_MODEL) + _assert_model_denied(client.chat_status(created, OPENAI_MODEL, f"say hi {unique_marker()}"), OPENAI_MODEL) @pytest.mark.covers("mgmt.key.update.happy_path") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.KEY_MANAGEMENT, + providers=(Provider.GEMINI, Provider.OPENAI), + models=(GEMINI_MODEL, OPENAI_MODEL), + mode=Mode.NONSTREAM, + ) + ) def test_editing_a_key_from_the_dashboard_persists_and_is_enforced( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -312,17 +372,17 @@ class TestDashboardKeyRoutes: target = _generate_key( client, resources, - KeyGenerateBody(models=["gemini-2.5-flash"], key_alias=alias, tpm_limit=100, rpm_limit=200), + KeyGenerateBody(models=[GEMINI_MODEL], key_alias=alias, tpm_limit=100, rpm_limit=200), ) - _poll_chat_ok(client, target, "gemini-2.5-flash") - _assert_model_denied(client.chat_status(target, "gpt-5.5", f"say hi {unique_marker()}"), "gpt-5.5") + _poll_chat_ok(client, target, GEMINI_MODEL) + _assert_model_denied(client.chat_status(target, OPENAI_MODEL, f"say hi {unique_marker()}"), OPENAI_MODEL) session = client.dashboard_login(UI_USERNAME, UI_PASSWORD) resources.defer(lambda: client.proxy.delete_key(session.session_key)) def dashboard_saves_the_edit() -> bool | None: match client.update_key( - KeyUpdateBody(key=target, models=["gpt-5.5"], tpm_limit=300, rpm_limit=400), + KeyUpdateBody(key=target, models=[OPENAI_MODEL], tpm_limit=300, rpm_limit=400), caller_key=session.session_key, ): case Success(): @@ -337,7 +397,7 @@ class TestDashboardKeyRoutes: ) info = client.proxy.key_info(target) - assert info.models == ["gpt-5.5"], ( + assert info.models == [OPENAI_MODEL], ( f"/key/info reports models {info.models} after the dashboard edit to ['gpt-5.5']" ) assert info.tpm_limit == 300, f"/key/info reports tpm_limit {info.tpm_limit} after the dashboard edit to 300" @@ -346,29 +406,38 @@ class TestDashboardKeyRoutes: f"the dashboard edit renamed the key to {info.key_alias!r}, it should still be {alias!r}" ) - _poll_model_access_granted(client, target, "gpt-5.5") - _poll_chat_denied(client, target, "gemini-2.5-flash") + _poll_model_access_granted(client, target, OPENAI_MODEL) + _poll_chat_denied(client, target, GEMINI_MODEL) class TestKeyRegeneration: @pytest.mark.covers("mgmt.key.regenerate.happy_path") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.KEY_MANAGEMENT, + providers=(Provider.OPENAI,), + models=(OPENAI_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_regenerate_rotates_to_a_working_new_key( self, client: ManagementClient, resources: ResourceManager ) -> None: - old_key = _generate_key(client, resources, KeyGenerateBody(models=["gpt-5.5"])) + old_key = _generate_key(client, resources, KeyGenerateBody(models=[OPENAI_MODEL])) new_key = client.regenerate_key(old_key) resources.defer(lambda: client.proxy.delete_key(new_key)) assert new_key != old_key, "regenerate returned the same key string, so no rotation happened" def new_accepted() -> bool | None: - outcome = client.chat_status(new_key, "gpt-5.5", f"say hi {unique_marker()}") + outcome = client.chat_status(new_key, OPENAI_MODEL, f"say hi {unique_marker()}") return True if outcome.status_code != 401 else None _ = _poll(client, new_accepted, "regenerated key was never accepted at auth (still 401) at the deadline") def old_rejected() -> bool | None: - outcome = client.chat_status(old_key, "gpt-5.5", f"say hi {unique_marker()}") + outcome = client.chat_status(old_key, OPENAI_MODEL, f"say hi {unique_marker()}") return True if outcome.status_code == 401 else None _ = _poll( @@ -376,10 +445,19 @@ class TestKeyRegeneration: ) @pytest.mark.covers("other.key_mgmt.regenerate.grace_period_honored") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.KEY_MANAGEMENT, + providers=(Provider.OPENAI,), + models=(OPENAI_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_regenerate_with_grace_period_keeps_old_key_until_revoked( self, client: ManagementClient, resources: ResourceManager ) -> None: - old_key = _generate_key(client, resources, KeyGenerateBody(models=["gpt-5.5"])) + old_key = _generate_key(client, resources, KeyGenerateBody(models=[OPENAI_MODEL])) new_key = client.regenerate_key(old_key, grace_period=REGENERATE_GRACE_PERIOD) resources.defer(lambda: client.proxy.delete_key(new_key)) @@ -387,7 +465,7 @@ class TestKeyRegeneration: assert new_key != old_key, "regenerate returned the same key string, so no rotation happened" def old_accepted() -> bool | None: - outcome = client.chat_status(old_key, "gpt-5.5", f"say hi {unique_marker()}") + outcome = client.chat_status(old_key, OPENAI_MODEL, f"say hi {unique_marker()}") return True if outcome.ok else None _ = _poll(client, old_accepted, "old key was rejected 401 inside its grace period at the deadline") @@ -396,7 +474,7 @@ class TestKeyRegeneration: ) def old_rejected() -> bool | None: - outcome = client.chat_status(old_key, "gpt-5.5", f"say hi {unique_marker()}") + outcome = client.chat_status(old_key, OPENAI_MODEL, f"say hi {unique_marker()}") return True if outcome.status_code == 401 else None _ = _poll( @@ -408,15 +486,21 @@ class TestKeyRegeneration: class TestTeamRoutes: @pytest.mark.covers("mgmt.team.new.persists") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.TEAM_MANAGEMENT, + ) + ) def test_new_persists_to_team_info_and_binds_keys( self, client: ManagementClient, resources: ResourceManager ) -> None: alias = f"e2e-mgmt-team-{unique_marker()}" - team_id = _create_team(client, resources, alias, ["gemini-2.5-flash"]) + team_id = _create_team(client, resources, alias, [GEMINI_MODEL]) info = client.team_info(team_id) assert info.team_alias == alias, f"/team/info reports team_alias {info.team_alias!r}, configured {alias!r}" - assert info.models == ["gemini-2.5-flash"], ( + assert info.models == [GEMINI_MODEL], ( f"/team/info reports models {info.models}, configured ['gemini-2.5-flash']" ) @@ -427,8 +511,14 @@ class TestTeamRoutes: ) @pytest.mark.covers("mgmt.team.update.persists") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.TEAM_MANAGEMENT, + ) + ) def test_update_persists_to_team_info(self, client: ManagementClient, resources: ResourceManager) -> None: - team_id = _create_team(client, resources, f"e2e-mgmt-team-{unique_marker()}", ["gemini-2.5-flash"]) + team_id = _create_team(client, resources, f"e2e-mgmt-team-{unique_marker()}", [GEMINI_MODEL]) updated_alias = f"e2e-mgmt-team-updated-{unique_marker()}" client.update_team(TeamUpdateBody(team_id=team_id, team_alias=updated_alias)) @@ -438,11 +528,17 @@ class TestTeamRoutes: _ = _poll(client, reflected, f"/team/info never reflected team_alias {updated_alias!r} after /team/update") @pytest.mark.covers("mgmt.team.list.happy_path") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.TEAM_MANAGEMENT, + ) + ) def test_created_team_appears_in_team_list( self, client: ManagementClient, resources: ResourceManager ) -> None: alias = f"e2e-mgmt-team-{unique_marker()}" - team_id = _create_team(client, resources, alias, ["gemini-2.5-flash"]) + team_id = _create_team(client, resources, alias, [GEMINI_MODEL]) _ = _poll( client, @@ -451,17 +547,26 @@ class TestTeamRoutes: ) @pytest.mark.covers("mgmt.team.delete.persists") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.TEAM_MANAGEMENT, + providers=(Provider.OPENAI,), + models=(OPENAI_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_delete_persists_and_revokes_team_bound_key( self, client: ManagementClient, resources: ResourceManager ) -> None: """The teardown's deferred delete_team/delete_key fire again on the already- deleted team and key by design: both are warn-only no-ops, and the deferred cleanup must survive this test failing before the in-body delete.""" - team_id = _create_team(client, resources, f"e2e-mgmt-team-{unique_marker()}", ["gpt-5.5"]) + team_id = _create_team(client, resources, f"e2e-mgmt-team-{unique_marker()}", [OPENAI_MODEL]) key = _generate_key(client, resources, KeyGenerateBody(team_id=team_id)) def accepted() -> bool | None: - outcome = client.chat_status(key, "gpt-5.5", f"say hi {unique_marker()}") + outcome = client.chat_status(key, OPENAI_MODEL, f"say hi {unique_marker()}") return True if outcome.status_code != 401 else None _ = _poll(client, accepted, "team-bound key was never accepted at auth before team deletion") @@ -474,7 +579,7 @@ class TestTeamRoutes: ) def rejected() -> bool | None: - outcome = client.chat_status(key, "gpt-5.5", f"say hi {unique_marker()}") + outcome = client.chat_status(key, OPENAI_MODEL, f"say hi {unique_marker()}") return True if outcome.status_code == 401 else None _ = _poll( @@ -482,6 +587,12 @@ class TestTeamRoutes: ) @pytest.mark.covers("mgmt.team.delete.membership_larger_than_db_pool") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.TEAM_MANAGEMENT, + ) + ) def test_team_delete_succeeds_for_team_larger_than_db_pool( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -519,6 +630,12 @@ class TestTeamRoutes: ) @pytest.mark.covers("mgmt.team.member_add.persists") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.TEAM_MANAGEMENT, + ) + ) def test_member_add_and_delete_persist_to_team_info( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -527,7 +644,7 @@ class TestTeamRoutes: resources, UserNewBody(user_email=f"e2e-mgmt-{unique_marker()}@example.com", user_role="internal_user"), ) - team_id = _create_team(client, resources, f"e2e-mgmt-team-{unique_marker()}", ["gemini-2.5-flash"]) + team_id = _create_team(client, resources, f"e2e-mgmt-team-{unique_marker()}", [GEMINI_MODEL]) client.add_team_member(team_id, user_id) member = next( @@ -545,6 +662,12 @@ class TestTeamRoutes: class TestUserRoutes: @pytest.mark.covers("mgmt.user.new.happy_path") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.USER_MANAGEMENT, + ) + ) def test_new_persists_to_user_info(self, client: ManagementClient, resources: ResourceManager) -> None: email = f"e2e-mgmt-{unique_marker()}@example.com" user_id = _create_user(client, resources, UserNewBody(user_email=email, user_role="internal_user")) @@ -556,6 +679,12 @@ class TestUserRoutes: ) @pytest.mark.covers("mgmt.user.update.persists") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.USER_MANAGEMENT, + ) + ) def test_update_persists_to_user_info(self, client: ManagementClient, resources: ResourceManager) -> None: email = f"e2e-mgmt-{unique_marker()}@example.com" user_id = _create_user(client, resources, UserNewBody(user_email=email, user_role="internal_user")) @@ -572,6 +701,12 @@ class TestUserRoutes: f"/user/info reports user_role {info.user_role!r} after /user/update to 'internal_user_viewer'" ) @pytest.mark.covers("mgmt.user.delete.persists") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.USER_MANAGEMENT, + ) + ) def test_delete_removes_the_user_from_inventory( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -594,6 +729,12 @@ class TestUserRoutes: _ = _poll(client, removed, f"user {user_id} still present in /user/list after /user/delete at the deadline") @pytest.mark.covers("mgmt.user.list.happy_path") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.USER_MANAGEMENT, + ) + ) def test_created_users_appear_in_user_list( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -616,22 +757,34 @@ class TestUserRoutes: class TestOrganizationRoutes: @pytest.mark.covers("mgmt.organization.new.happy_path") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.ORGANIZATION_MANAGEMENT, + ) + ) def test_new_persists_to_organization_info( self, client: ManagementClient, resources: ResourceManager ) -> None: alias = f"e2e-mgmt-org-{unique_marker()}" - org_id = client.create_org(OrgNewBody(organization_alias=alias, models=["gemini-2.5-flash"])) + org_id = client.create_org(OrgNewBody(organization_alias=alias, models=[GEMINI_MODEL])) resources.defer(lambda: client.delete_org(org_id)) info = client.org_info(org_id) assert info.organization_alias == alias, ( f"/organization/info reports alias {info.organization_alias!r}, configured {alias!r}" ) - assert info.models == ["gemini-2.5-flash"], ( + assert info.models == [GEMINI_MODEL], ( f"/organization/info reports models {info.models}, configured ['gemini-2.5-flash']" ) @pytest.mark.covers("mgmt.organization.update.persists") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.ORGANIZATION_MANAGEMENT, + ) + ) def test_update_alias_persists_to_organization_info( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -650,6 +803,12 @@ class TestOrganizationRoutes: ) @pytest.mark.covers("mgmt.organization.delete.persists") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.ORGANIZATION_MANAGEMENT, + ) + ) def test_delete_removes_from_organization_info( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -674,6 +833,12 @@ class TestOrganizationRoutes: class TestTagRoutes: @pytest.mark.covers("mgmt.tag.new.happy_path") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.TAG_MANAGEMENT, + ) + ) def test_new_persists_to_tag_list(self, client: ManagementClient, resources: ResourceManager) -> None: name = f"e2e-mgmt-tag-{unique_marker()}" description = "Tag for spend categorization" @@ -704,6 +869,12 @@ def _model_entry(client: ManagementClient, model_name: str) -> ModelInfoEntry | class TestModelRoutes: @pytest.mark.covers("mgmt.model.update.persists") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.MODEL_MANAGEMENT, + ) + ) def test_update_persists_input_cost_to_model_info( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -747,6 +918,12 @@ class TestModelRoutes: ) @pytest.mark.covers("mgmt.model.delete.persists") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.MODEL_MANAGEMENT, + ) + ) def test_delete_removes_from_model_info_catalog( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -769,6 +946,12 @@ class TestModelRoutes: _ = _poll(client, absent, f"{model_name} still present in /model/info after /model/delete at the deadline") @pytest.mark.covers("mgmt.model.add.persists") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.MODEL_MANAGEMENT, + ) + ) def test_new_persists_to_model_info_catalog( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -797,6 +980,11 @@ def _assert_route_forbidden(route: str, outcome: StreamingResponse) -> None: class TestManagementRoutePermissions: @pytest.mark.covers("other.auth.virtual_key.route_permission_enforced") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + ) + ) def test_llm_only_key_forbidden_from_management_writes( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -832,6 +1020,12 @@ class TestManagementRoutePermissions: class TestCustomer: @pytest.mark.covers("mgmt.end_user.new.happy_path") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.CUSTOMER_MANAGEMENT, + ) + ) def test_customer_create_persists_to_info( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -896,6 +1090,12 @@ def _generate_response( class TestKeyDeletionAuditLog: @pytest.mark.covers("mgmt.key.delete.audit_logged") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.KEY_MANAGEMENT, + ) + ) def test_key_delete_by_key_writes_audit_row( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -908,6 +1108,12 @@ class TestKeyDeletionAuditLog: _assert_single_deleted_row(_await_deleted_audit_rows(client, token), token) @pytest.mark.covers("mgmt.key.delete.audit_logged") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.KEY_MANAGEMENT, + ) + ) def test_key_delete_by_alias_writes_audit_row( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -921,6 +1127,12 @@ class TestKeyDeletionAuditLog: _assert_single_deleted_row(_await_deleted_audit_rows(client, token), token) @pytest.mark.covers("mgmt.team.member_delete.audit_logs_keys") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.TEAM_MANAGEMENT, + ) + ) def test_team_member_delete_writes_audit_row_for_member_keys( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -940,6 +1152,12 @@ class TestKeyDeletionAuditLog: _assert_single_deleted_row(_await_deleted_audit_rows(client, token), token) @pytest.mark.covers("mgmt.team.delete.audit_logs_keys") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.TEAM_MANAGEMENT, + ) + ) def test_team_delete_writes_audit_row_for_team_keys( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -953,6 +1171,12 @@ class TestKeyDeletionAuditLog: _assert_single_deleted_row(_await_deleted_audit_rows(client, token), token) @pytest.mark.covers("mgmt.user.delete.audit_logs_keys") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.USER_MANAGEMENT, + ) + ) def test_user_delete_writes_audit_row_for_user_keys( self, client: ManagementClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/management/test_mcp_lifecycle_e2e.py b/tests/e2e/management/test_mcp_lifecycle_e2e.py index 9257d697647..ca2e99cad8b 100644 --- a/tests/e2e/management/test_mcp_lifecycle_e2e.py +++ b/tests/e2e/management/test_mcp_lifecycle_e2e.py @@ -19,6 +19,7 @@ from typing import Final import pytest from e2e_config import unique_marker from e2e_http import unwrap +from e2e_metadata import Domain, Route, Subject, meta from lifecycle import ResourceManager from management_client import ManagementClient from models import ( @@ -88,6 +89,12 @@ def _listed_server_everywhere(client: ManagementClient, server_id: str) -> Mappi class TestMcpServerLifecycle: @pytest.mark.covers("mgmt.mcp_server.new.persists") + @meta( + Subject( + domain=Domain.MCP, + route=Route.MCP, + ) + ) def test_create_persists_every_field_on_every_replica( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -107,6 +114,12 @@ class TestMcpServerLifecycle: ) ) @pytest.mark.covers("mgmt.mcp_server.list.persists") + @meta( + Subject( + domain=Domain.MCP, + route=Route.MCP, + ) + ) def test_created_server_is_listed_with_every_field( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -116,6 +129,12 @@ class TestMcpServerLifecycle: _assert_server_matches(row, body, where=f"GET /v1/mcp/server on {replica}") @pytest.mark.covers("mgmt.mcp_server.update.preserves_unrelated_fields") + @meta( + Subject( + domain=Domain.MCP, + route=Route.MCP, + ) + ) def test_updating_only_the_alias_keeps_every_other_field_on_every_replica( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -133,6 +152,12 @@ class TestMcpServerLifecycle: ) @pytest.mark.covers("mgmt.mcp_server.update.clear_persists") + @meta( + Subject( + domain=Domain.MCP, + route=Route.MCP, + ) + ) def test_clearing_the_description_with_null_reads_back_null( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -149,6 +174,12 @@ class TestMcpServerLifecycle: ) @pytest.mark.covers("mgmt.mcp_server.delete.persists") + @meta( + Subject( + domain=Domain.MCP, + route=Route.MCP, + ) + ) def test_delete_removes_the_server_from_every_replica( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -199,6 +230,12 @@ def _toolset_everywhere( class TestMcpToolsetLifecycle: @pytest.mark.covers("mgmt.mcp_toolset.new.persists") + @meta( + Subject( + domain=Domain.MCP, + route=Route.MCP, + ) + ) def test_create_persists_both_tools_under_the_exact_names_written( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -221,6 +258,12 @@ class TestMcpToolsetLifecycle: ) @pytest.mark.covers("mgmt.mcp_toolset.update.preserves_unrelated_fields") + @meta( + Subject( + domain=Domain.MCP, + route=Route.MCP, + ) + ) def test_updating_only_the_description_keeps_the_tools_and_name( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -238,6 +281,12 @@ class TestMcpToolsetLifecycle: ) @pytest.mark.covers("mgmt.mcp_toolset.update.persists") + @meta( + Subject( + domain=Domain.MCP, + route=Route.MCP, + ) + ) def test_updating_the_tools_to_one_entry_reads_back_exactly_that_entry( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -256,6 +305,12 @@ class TestMcpToolsetLifecycle: ) @pytest.mark.covers("mgmt.mcp_toolset.update.clear_persists") + @meta( + Subject( + domain=Domain.MCP, + route=Route.MCP, + ) + ) def test_clearing_the_description_with_null_reads_back_null( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -273,6 +328,12 @@ class TestMcpToolsetLifecycle: ) @pytest.mark.covers("mgmt.mcp_toolset.delete.persists") + @meta( + Subject( + domain=Domain.MCP, + route=Route.MCP, + ) + ) def test_delete_removes_the_toolset_from_every_replica( self, client: ManagementClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/management/test_model_tag_accessgroup_e2e.py b/tests/e2e/management/test_model_tag_accessgroup_e2e.py index eb3a6093c69..5c104db18c1 100644 --- a/tests/e2e/management/test_model_tag_accessgroup_e2e.py +++ b/tests/e2e/management/test_model_tag_accessgroup_e2e.py @@ -22,6 +22,7 @@ from pydantic import BaseModel, ConfigDict, RootModel from e2e_config import unique_marker from e2e_http import NoBody, unwrap +from e2e_metadata import Domain, Route, Subject, meta from lifecycle import ResourceManager from management_client import ManagementClient from models import KeyGenerateBody, LiteLLMParamsBody, ModelInfoBody, ModelNewBody @@ -216,6 +217,12 @@ def _model_blocked_flag(client: ManagementClient, model_id: str) -> bool | None: class TestModelRoutes: @pytest.mark.covers("mgmt.model.add.admin_only") + @meta( + Subject( + domain=Domain.PROXY_AUTH, + route=Route.MODEL_MANAGEMENT, + ) + ) def test_non_admin_key_cannot_add_global_model( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -248,6 +255,12 @@ class TestModelRoutes: ) @pytest.mark.covers("mgmt.model.block.persists") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.MODEL_MANAGEMENT, + ) + ) def test_block_then_unblock_persists_to_model_info( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -281,6 +294,12 @@ class TestModelRoutes: class TestTagRoutes: @pytest.mark.covers("mgmt.tag.list.happy_path") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.TAG_MANAGEMENT, + ) + ) def test_tag_list_reports_created_tag(self, client: ManagementClient, resources: ResourceManager) -> None: name = f"e2e-mgmt-tag-{unique_marker()}" description = "coverage: tag inventory" @@ -301,6 +320,12 @@ class TestTagRoutes: ) @pytest.mark.covers("mgmt.tag.delete.persists") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.TAG_MANAGEMENT, + ) + ) def test_tag_delete_removes_from_list(self, client: ManagementClient, resources: ResourceManager) -> None: """The teardown's deferred delete fires again on the already-deleted tag by design: it is the safety net if this test fails before the in-body delete, @@ -326,6 +351,12 @@ class TestTagRoutes: class TestModelAccessGroupRoutes: @pytest.mark.covers("mgmt.access_group.new.happy_path") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.MODEL_MANAGEMENT, + ) + ) def test_new_access_group_tags_the_deployment( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -356,6 +387,12 @@ class TestModelAccessGroupRoutes: ) @pytest.mark.covers("mgmt.access_group.info.happy_path") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.MODEL_MANAGEMENT, + ) + ) def test_access_group_info_reports_membership( self, client: ManagementClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/management/test_model_test_connection_e2e.py b/tests/e2e/management/test_model_test_connection_e2e.py index 25b0b4f24e6..b0d788d582c 100644 --- a/tests/e2e/management/test_model_test_connection_e2e.py +++ b/tests/e2e/management/test_model_test_connection_e2e.py @@ -23,6 +23,7 @@ import time import pytest from e2e_http import unwrap +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from management_client import ManagementClient from models import ConnectionTestBody, ConnectionTestResponse, LiteLLMParamsBody @@ -50,6 +51,15 @@ def _probe_mantle(client: ManagementClient) -> ConnectionTestResponse: class TestModelTestConnection: @pytest.mark.covers("mgmt.model.test_connection.happy_path") + @meta( + Subject( + domain=Domain.MANAGEMENT, + route=Route.HEALTH, + providers=(Provider.BEDROCK_MANTLE,), + models=(MANTLE_RESPONSES_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_bedrock_mantle_responses_connection_succeeds(self, client: ManagementClient) -> None: for attempt in range(1, PROBE_ATTEMPTS + 1): response = _probe_mantle(client) diff --git a/tests/e2e/management/test_team_management_e2e.py b/tests/e2e/management/test_team_management_e2e.py index f30dc6990a9..68a26838ebd 100644 --- a/tests/e2e/management/test_team_management_e2e.py +++ b/tests/e2e/management/test_team_management_e2e.py @@ -27,6 +27,7 @@ import pytest from pydantic import BaseModel from e2e_config import settle_propagation, unique_marker +from e2e_metadata import Domain, Route, Subject, meta from e2e_http import NoBody, PartialBody, StreamingResponse, unwrap from lifecycle import ResourceManager from management_client import ManagementClient @@ -241,6 +242,7 @@ def _member_delete_status(client: ManagementClient, key: str, team_id: str, user class TestTeamManagementRoutes: @pytest.mark.covers("mgmt.team.info.happy_path") + @meta(Subject(domain=Domain.MANAGEMENT, route=Route.TEAM_MANAGEMENT)) def test_info_returns_created_team_fields( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -257,6 +259,7 @@ class TestTeamManagementRoutes: ) @pytest.mark.covers("mgmt.team.block.persists") + @meta(Subject(domain=Domain.MANAGEMENT, route=Route.TEAM_MANAGEMENT)) def test_block_then_unblock_persists_to_team_info( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -278,6 +281,7 @@ class TestTeamManagementRoutes: ) @pytest.mark.covers("mgmt.team.member_update.persists") + @meta(Subject(domain=Domain.MANAGEMENT, route=Route.TEAM_MANAGEMENT)) def test_member_update_persists_role_and_budget( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -302,6 +306,7 @@ class TestTeamManagementRoutes: ) @pytest.mark.covers("mgmt.team.member_delete.persists") + @meta(Subject(domain=Domain.MANAGEMENT, route=Route.TEAM_MANAGEMENT)) def test_member_delete_persists_to_team_info( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -320,6 +325,7 @@ class TestTeamManagementRoutes: ) @pytest.mark.covers("mgmt.team.new.admin_only") + @meta(Subject(domain=Domain.PROXY_AUTH, route=Route.TEAM_MANAGEMENT)) def test_new_is_denied_to_non_admin_keys( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -334,6 +340,7 @@ class TestTeamManagementRoutes: ) @pytest.mark.covers("mgmt.team.member_add.member_forbidden") + @meta(Subject(domain=Domain.PROXY_AUTH, route=Route.TEAM_MANAGEMENT)) def test_member_add_forbidden_to_plain_member( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -348,6 +355,7 @@ class TestTeamManagementRoutes: ) @pytest.mark.covers("mgmt.team.member_delete.member_forbidden") + @meta(Subject(domain=Domain.PROXY_AUTH, route=Route.TEAM_MANAGEMENT)) def test_member_delete_forbidden_to_plain_member( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -458,6 +466,7 @@ class TestTeamAdminWithNoEditableFields: """No proxy admin has enabled a team field for team admins, which is how every proxy starts.""" @pytest.mark.covers("mgmt.team.update.team_admin_forbidden_until_enabled") + @meta(Subject(domain=Domain.PROXY_AUTH, route=Route.TEAM_MANAGEMENT)) def test_team_admin_cannot_change_any_team_setting( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -483,6 +492,7 @@ class TestTeamAdminWithTpmLimitEnabled: """A proxy admin has enabled tpm_limit, so a team admin may change that setting and no other.""" @pytest.mark.covers("mgmt.team.update.team_admin_limited_to_enabled_fields") + @meta(Subject(domain=Domain.PROXY_AUTH, route=Route.TEAM_MANAGEMENT)) def test_team_admin_saves_the_settings_form_with_a_new_tpm_limit( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -522,6 +532,7 @@ class TestTeamAdminWithTpmLimitEnabled: pytest.param(TeamSettingsChange(metadata=TeamCustomMetadata(cost_center="team-admin")), id="metadata"), ], ) + @meta(Subject(domain=Domain.PROXY_AUTH, route=Route.TEAM_MANAGEMENT)) def test_team_admin_cannot_change_a_setting_that_is_not_enabled( self, client: ManagementClient, resources: ResourceManager, change: TeamSettingsChange ) -> None: @@ -548,6 +559,7 @@ class TestTeamAdminWithTpmLimitEnabled: ) @pytest.mark.covers("mgmt.team.update.team_admin_resend_keeps_budget_reset") + @meta(Subject(domain=Domain.SPEND_BUDGETS, route=Route.TEAM_MANAGEMENT)) def test_team_admin_resending_the_budget_settings_keeps_the_next_budget_reset( self, client: ManagementClient, resources: ResourceManager ) -> None: @@ -613,6 +625,7 @@ class TestTeamAdminWithRpmLimitAndMaxBudgetEnabled: "current_budget", [pytest.param(_TEAM_MAX_BUDGET, id="lower"), pytest.param(None, id="first-budget")], ) + @meta(Subject(domain=Domain.PROXY_AUTH, route=Route.TEAM_MANAGEMENT)) def test_team_admin_saves_a_new_rpm_limit_and_a_tighter_budget( self, client: ManagementClient, resources: ResourceManager, current_budget: float | None ) -> None: @@ -649,6 +662,7 @@ class TestTeamAdminWithRpmLimitAndMaxBudgetEnabled: pytest.param(None, "Only a proxy admin can remove", id="remove"), ], ) + @meta(Subject(domain=Domain.PROXY_AUTH, route=Route.TEAM_MANAGEMENT)) def test_team_admin_cannot_raise_or_remove_the_budget( self, client: ManagementClient, resources: ResourceManager, max_budget: float | None, refusal: str ) -> None: @@ -670,6 +684,7 @@ class TestTeamAdminWithRpmLimitAndMaxBudgetEnabled: ) @pytest.mark.covers("mgmt.team.update.team_admin_cannot_grow_budget") + @meta(Subject(domain=Domain.PROXY_AUTH, route=Route.TEAM_MANAGEMENT)) def test_team_admin_cannot_raise_an_org_team_budget_under_the_org_cap( self, client: ManagementClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/mcp/conftest.py b/tests/e2e/mcp/conftest.py index e6094ab95ea..91f55a528b5 100644 --- a/tests/e2e/mcp/conftest.py +++ b/tests/e2e/mcp/conftest.py @@ -15,14 +15,11 @@ from typing import Protocol, cast import pytest +from datadog_mcp import DdLogsReader from mcp_client import McpClient, build_client from proxy_client import ProxyClient -class DdLogsReader(Protocol): - def poll_events_for_marker(self, marker: str) -> list[object]: ... - - class _DdLogsReaderBuilder(Protocol): def __call__(self) -> DdLogsReader: ... diff --git a/tests/e2e/mcp/datadog_mcp.py b/tests/e2e/mcp/datadog_mcp.py index 352b4446cfd..63194af221e 100644 --- a/tests/e2e/mcp/datadog_mcp.py +++ b/tests/e2e/mcp/datadog_mcp.py @@ -4,14 +4,20 @@ from __future__ import annotations import os from collections.abc import Sequence +from typing import Protocol from e2e_config import datadog_mcp_url, unique_marker +from e2e_metadata import step from lifecycle import ResourceManager from mcp_client import McpClient SEARCH_LOGS_TOOL = "search_datadog_logs" +class DdLogsReader(Protocol): + def poll_events_for_marker(self, marker: str) -> list[object]: ... + + def _dd_api_key() -> str: return os.environ.get("DD_API_KEY", "").strip() @@ -31,6 +37,7 @@ def assert_dd_mcp_creds() -> None: ) +@step("Register the Datadog remote MCP server with its credentials from the environment") def register_datadog_mcp( client: McpClient, resources: ResourceManager, diff --git a/tests/e2e/mcp/mcp_client.py b/tests/e2e/mcp/mcp_client.py index 56f7fffba29..e373a8e31ae 100644 --- a/tests/e2e/mcp/mcp_client.py +++ b/tests/e2e/mcp/mcp_client.py @@ -21,6 +21,7 @@ from pydantic import BaseModel, ConfigDict, Field, RootModel from e2e_config import settle_propagation from e2e_http import Headers, NoBody, Result, Success, UnknownApiError, unwrap +from e2e_metadata import step from models import KeyGenerateBody, McpServerListResponse, McpServerRow, ObjectPermission from proxy_client import ProxyClient @@ -153,6 +154,7 @@ class McpCallToolResponse(BaseModel): class McpClient: proxy: ProxyClient + @step("Register the MCP server {server_name} with the alias {alias}") def register_server( self, *, @@ -183,6 +185,7 @@ class McpClient: ) ).server_id + @step("Delete the MCP server") def delete_server(self, server_id: str) -> None: _ = self.proxy.transport.delete( f"/v1/mcp/server/{server_id}", @@ -191,6 +194,7 @@ class McpClient: response_type=NoBody, ) + @step("List the MCP servers from /v1/mcp/server") def registered_servers(self) -> list[McpServerRow]: return unwrap( self.proxy.transport.get( @@ -201,6 +205,7 @@ class McpClient: ) ).root + @step("List the MCP servers the key can see from /v1/mcp/server") def list_servers(self, key: str) -> Result[McpServerListResponse]: return self.proxy.transport.get( "/v1/mcp/server", @@ -209,6 +214,7 @@ class McpClient: response_type=McpServerListResponse, ) + @step("Check the health of the MCP servers the key can see from /v1/mcp/server/health") def server_health(self, key: str, server_ids: list[str] | None = None) -> Result[McpHealthResponse]: return self.proxy.transport.get( "/v1/mcp/server/health", @@ -217,6 +223,7 @@ class McpClient: response_type=McpHealthResponse, ) + @step("Wait for every proxy replica to list the MCP server in /v1/mcp/server") def await_registered(self, server_id: str) -> McpServerRow: """Wait for every configured replica to list the server and return its row.""" registered = self.proxy.read_body_back_everywhere( @@ -228,6 +235,7 @@ class McpClient: row for response in registered.values() for row in response.root if row.server_id == server_id ) + @step("Generate a virtual key for the user {user_id}") def generate_key( self, *, @@ -254,6 +262,7 @@ class McpClient: ) ) + @step("List the MCP tools the key can see from /mcp-rest/tools/list") def list_tools(self, key: str) -> Result[McpToolsListResponse]: return self.proxy.transport.get( "/mcp-rest/tools/list", @@ -262,6 +271,7 @@ class McpClient: response_type=McpToolsListResponse, ) + @step('Wait for /mcp-rest/tools/list to show the MCP server\'s tool matching "{needle}"') def await_tool(self, key: str, server_id: str, needle: str) -> str: """Poll tools/list until `server_id` serves a tool matching `needle`, and return its fully-qualified name. Fails at poll_timeout. @@ -287,6 +297,7 @@ class McpClient: ) time.sleep(self.proxy.poll_interval) + @step("Wait for /mcp-rest/tools/list to show the key exactly the expected tools on the MCP server") def await_tools(self, key: str, server_id: str, *, expected: frozenset[str]) -> frozenset[str]: """Poll tools/list until `server_id`'s tools as `key` sees them are exactly `expected`, and return the last listing either way, so the caller's equality @@ -301,6 +312,7 @@ class McpClient: return unwrap(result).tool_names_for_server(server_id) time.sleep(self.proxy.poll_interval) + @step("Call the MCP tool {name} through /mcp-rest/tools/call") def await_call_tool( self, key: str, @@ -329,6 +341,7 @@ class McpClient: ) time.sleep(self.proxy.poll_interval) + @step("Call the MCP tool {name} through /mcp-rest/tools/call and wait for a 403") def await_call_tool_denied( self, key: str, @@ -355,6 +368,7 @@ class McpClient: ) time.sleep(self.proxy.poll_interval) + @step('Create the guardrail {name} that blocks MCP tool calls containing "{blocked_keyword}"') def register_mcp_content_filter(self, *, name: str, blocked_keyword: str) -> str: """Register a default-on content-filter guardrail that runs on the MCP tool-call hook (pre_mcp_call) and blocks a single keyword. The keyword is @@ -378,6 +392,7 @@ class McpClient: settle_propagation(time.monotonic()) return guardrail_id + @step("Delete the guardrail") def delete_guardrail(self, guardrail_id: str) -> None: _ = self.proxy.transport.delete( f"/guardrails/{guardrail_id}", @@ -386,6 +401,7 @@ class McpClient: response_type=NoBody, ) + @step("Call the MCP tool {name} through /mcp-rest/tools/call with {arguments}") def call_tool( self, key: str, diff --git a/tests/e2e/mcp/oauth_chat_client.py b/tests/e2e/mcp/oauth_chat_client.py index e8c4c1d7620..c8496cc7592 100644 --- a/tests/e2e/mcp/oauth_chat_client.py +++ b/tests/e2e/mcp/oauth_chat_client.py @@ -26,6 +26,7 @@ import httpx2 import pytest from e2e_config import PROXY_BASE_URL, REQUEST_TIMEOUT from e2e_http import AuthHeaders, NoBody, unwrap +from e2e_metadata import step from idp import Identity from mcp import ClientSession from mcp.client.auth import OAuthClientProvider @@ -318,6 +319,7 @@ async def _list_and_call( class ChatMcpClient: proxy: ProxyClient + @step("Register the MCP server with the alias {body.alias}") def create_server(self, body: McpServerCreateBody) -> McpServerInfo: return unwrap( self.proxy.transport.post( @@ -328,6 +330,7 @@ class ChatMcpClient: ) ) + @step("Read the MCP server back from /v1/mcp/server") def server_info(self, server_id: str) -> McpServerInfo: return unwrap( self.proxy.transport.get( @@ -338,6 +341,7 @@ class ChatMcpClient: ) ) + @step("Delete the MCP server") def delete_server(self, server_id: str) -> None: _ = self.proxy.transport.delete( f"/v1/mcp/server/{server_id}", @@ -346,6 +350,7 @@ class ChatMcpClient: response_type=NoBody, ) + @step("Sign the key's user in to the MCP server {alias} through the OAuth consent flow") def seed_user_token(self, alias: str, key: str, storage_state_path: str) -> tuple[str, ...]: """Drive the interactive authorize dance for `key`'s user so the gateway stores their upstream token, retried to the shared deadline since the @@ -367,6 +372,7 @@ class ChatMcpClient: f"last error: {last_error!r}" ) + @step("List the tools on the MCP server {alias} and call {tool} over the MCP protocol") def list_and_call( self, alias: str, @@ -394,6 +400,7 @@ class ChatMcpClient: ) ) + @step("List the users with a stored OAuth token for the MCP server") def server_user_credentials(self, server_id: str) -> tuple[McpServerUserCredentialRow, ...]: return unwrap( self.proxy.transport.get( @@ -404,6 +411,7 @@ class ChatMcpClient: ) ).root + @step("Revoke the user's stored OAuth token for the MCP server") def revoke_user_token(self, server_id: str, headers: AuthHeaders) -> None: _ = unwrap( self.proxy.transport.delete( @@ -414,6 +422,7 @@ class ChatMcpClient: ) ) + @step("Send a /chat/completions request to {body.model} with an MCP server attached as a tool") def chat_with_mcp(self, headers: AuthHeaders, body: ChatBody) -> ChatResponse: """POST /chat/completions carrying the LiteLLM key in `headers` (either ingress form) with an MCP server attached in `body.tools`. The gateway diff --git a/tests/e2e/mcp/oauth_gateway.py b/tests/e2e/mcp/oauth_gateway.py index 029b0135900..2abe8262b6e 100644 --- a/tests/e2e/mcp/oauth_gateway.py +++ b/tests/e2e/mcp/oauth_gateway.py @@ -21,6 +21,7 @@ from typing import Final import psycopg 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 from psycopg.rows import class_row @@ -37,6 +38,7 @@ class CredentialRow: credential_b64: str = field(repr=False) +@step("Read the user's stored OAuth credential for the MCP server from the database and decrypt it") def stored_oauth(user_id: str, server_id: str) -> StoredOAuth: """Read the encrypted credential because management APIs omit the plaintext token.""" from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper @@ -108,6 +110,7 @@ class OAuthGateway: _log_path: Path _child: subprocess.Popen[bytes] | None = field(default=None, init=False, repr=False) + @step("Start the separate LiteLLM proxy and wait for /health/liveliness") def start(self) -> None: with self._log_path.open("ab") as log: self._child = subprocess.Popen( @@ -126,11 +129,13 @@ class OAuthGateway: time.sleep(0.5) raise AssertionError("owned OAuth gateway did not become ready") + @step("Stop the separate LiteLLM 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("Restart the separate LiteLLM proxy process so its in-memory caches start empty") def restart(self) -> None: assert self._child is not None previous: Final = self._child.pid @@ -139,6 +144,7 @@ class OAuthGateway: assert self._child.pid != previous, "gateway restart did not create a new process" +@step("Start a separate LiteLLM proxy from source with JWT auth against Keycloak") def owned_gateway(idp: Keycloak, directory: Path, cleanup: ExitStack) -> OAuthGateway: for name in ("DATABASE_URL", "LITELLM_LICENSE", "LITELLM_SALT_KEY", "LITELLM_MASTER_KEY"): assert os.environ.get(name), f"{name} is required for the owned OAuth gateway" diff --git a/tests/e2e/mcp/test_mcp_access_group_e2e.py b/tests/e2e/mcp/test_mcp_access_group_e2e.py index f72b75fd43d..744acb78aea 100644 --- a/tests/e2e/mcp/test_mcp_access_group_e2e.py +++ b/tests/e2e/mcp/test_mcp_access_group_e2e.py @@ -16,6 +16,7 @@ import pytest from datadog_mcp import SEARCH_LOGS_TOOL, register_datadog_mcp from e2e_config import unique_marker from e2e_http import unwrap +from e2e_metadata import Domain, Route, Subject, meta from lifecycle import ResourceManager from mcp_client import McpClient @@ -24,6 +25,12 @@ pytestmark = pytest.mark.e2e class TestMcpAccessGroupToolSelection: @pytest.mark.covers("mcp.list_tools.api_key.access_group_scoped") + @meta( + Subject( + domain=Domain.MCP, + route=Route.MCP, + ) + ) def test_access_group_scopes_tool_selection( self, client: McpClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/mcp/test_mcp_chat_completion_oauth_e2e.py b/tests/e2e/mcp/test_mcp_chat_completion_oauth_e2e.py index 01e94f7b86f..9d7e4713963 100644 --- a/tests/e2e/mcp/test_mcp_chat_completion_oauth_e2e.py +++ b/tests/e2e/mcp/test_mcp_chat_completion_oauth_e2e.py @@ -30,6 +30,7 @@ import pytest from e2e_config import CHEAP_ANTHROPIC_MODEL, LINEAR_MCP_URL, LINEAR_STORAGE_STATE, unique_marker from e2e_http import AuthHeaders +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import ChatBody, ChatMessage, KeyGenerateBody, McpChatTool, McpServerCreateBody, ObjectPermission from proxy_client import ProxyClient @@ -70,6 +71,16 @@ class TestMcpChatCompletionOauth: @pytest.mark.covers("mcp.list_tools.oauth.succeeds") @pytest.mark.covers("mcp.call_tool.oauth.succeeds") + @meta( + Subject( + domain=Domain.MCP, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(CHEAP_ANTHROPIC_MODEL,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_chat_completion_uses_linear_with_x_litellm_api_key_header( self, chat_client: ChatMcpClient, resources: ResourceManager ) -> None: @@ -134,6 +145,16 @@ class TestMcpChatCompletionOauth: @pytest.mark.covers("mcp.list_tools.oauth.succeeds") @pytest.mark.covers("mcp.call_tool.oauth.succeeds") + @meta( + Subject( + domain=Domain.MCP, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(CHEAP_ANTHROPIC_MODEL,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) def test_chat_completion_uses_linear_with_authorization_bearer_header( self, chat_client: ChatMcpClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/mcp/test_mcp_datadog_e2e.py b/tests/e2e/mcp/test_mcp_datadog_e2e.py index 031fbf6d936..e383e9ab766 100644 --- a/tests/e2e/mcp/test_mcp_datadog_e2e.py +++ b/tests/e2e/mcp/test_mcp_datadog_e2e.py @@ -13,10 +13,10 @@ from __future__ import annotations import pytest -from conftest import DdLogsReader -from datadog_mcp import SEARCH_LOGS_TOOL, assert_dd_mcp_creds, register_datadog_mcp +from datadog_mcp import SEARCH_LOGS_TOOL, DdLogsReader, assert_dd_mcp_creds, register_datadog_mcp from e2e_config import CHEAP_ANTHROPIC_MODEL, DD_SEARCH_FROM, unique_marker from e2e_http import NoBody, unwrap +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from mcp_client import McpClient from models import ChatBody, ChatMessage @@ -50,6 +50,15 @@ def _seed_completion(proxy: ProxyClient, *, key: str, marker: str) -> None: class TestDatadogMcpRoundTrip: @pytest.mark.covers("mcp.list_tools.api_key.succeeds", "mcp.call_tool.api_key.succeeds") + @meta( + Subject( + domain=Domain.MCP, + route=Route.MCP, + providers=(Provider.ANTHROPIC,), + models=(CHEAP_ANTHROPIC_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_search_logs_finds_seeded_completion( self, client: McpClient, diff --git a/tests/e2e/mcp/test_mcp_guardrail_e2e.py b/tests/e2e/mcp/test_mcp_guardrail_e2e.py index 92c632cb316..d34bd80e564 100644 --- a/tests/e2e/mcp/test_mcp_guardrail_e2e.py +++ b/tests/e2e/mcp/test_mcp_guardrail_e2e.py @@ -23,6 +23,7 @@ import pytest from datadog_mcp import SEARCH_LOGS_TOOL, assert_dd_mcp_creds, register_datadog_mcp from e2e_config import DD_SEARCH_FROM, unique_marker from e2e_http import Result, Success, UnknownApiError +from e2e_metadata import Domain, Route, Subject, meta from lifecycle import ResourceManager from mcp_client import McpCallToolResponse, McpClient, McpToolArguments @@ -82,6 +83,12 @@ class TestMcpToolCallGuardrail: "guardrail.litellm_content_filter.pre_mcp_call.blocks", exercised_on=["mcp_operations"], ) + @meta( + Subject( + domain=Domain.GUARDRAILS, + route=Route.MCP, + ) + ) def test_content_filter_blocks_banned_keyword_in_tool_args( self, client: McpClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/mcp/test_mcp_key_access_e2e.py b/tests/e2e/mcp/test_mcp_key_access_e2e.py index c00d67bc9cf..9f5e04ccb9b 100644 --- a/tests/e2e/mcp/test_mcp_key_access_e2e.py +++ b/tests/e2e/mcp/test_mcp_key_access_e2e.py @@ -18,6 +18,7 @@ from typing import Final from datadog_mcp import SEARCH_LOGS_TOOL, register_datadog_mcp from e2e_config import DD_SEARCH_FROM, unique_marker from e2e_http import unwrap +from e2e_metadata import Domain, Route, Subject, meta from lifecycle import ResourceManager from mcp_client import McpClient from models import KeyGenerateBody, ObjectPermission @@ -33,6 +34,12 @@ def _key(client: McpClient, resources: ResourceManager, *, mcp_servers: list[str class TestMcpKeyGrantByAlias: + @meta( + Subject( + domain=Domain.MCP, + route=Route.MCP, + ) + ) def test_alias_grant_persists_verbatim_and_lists_tools( self, client: McpClient, @@ -62,6 +69,12 @@ class TestMcpKeyGrantByAlias: class TestMcpKeyWithoutAccessIsDenied: @pytest.mark.covers("mcp.list_tools.api_key.denied_without_permission") + @meta( + Subject( + domain=Domain.MCP, + route=Route.MCP, + ) + ) def test_list_tools_denied_without_permission( self, client: McpClient, @@ -82,6 +95,12 @@ class TestMcpKeyWithoutAccessIsDenied: ) @pytest.mark.covers("mcp.call_tool.api_key.denied_without_permission") + @meta( + Subject( + domain=Domain.MCP, + route=Route.MCP, + ) + ) def test_call_tool_denied_without_permission( self, client: McpClient, @@ -113,6 +132,12 @@ class TestMcpKeyWithoutAccessIsDenied: class TestMcpHealthVisibility: + @meta( + Subject( + domain=Domain.MCP, + route=Route.MCP, + ) + ) def test_route_restricted_health_matches_server_grants( self, client: McpClient, diff --git a/tests/e2e/mcp/test_mcp_oauth_happy_path_e2e.py b/tests/e2e/mcp/test_mcp_oauth_happy_path_e2e.py index 81470c21d51..b57c941cb18 100644 --- a/tests/e2e/mcp/test_mcp_oauth_happy_path_e2e.py +++ b/tests/e2e/mcp/test_mcp_oauth_happy_path_e2e.py @@ -18,6 +18,7 @@ from typing import Final, Literal import pytest from e2e_config import LINEAR_MCP_URL, LINEAR_READONLY_TOOL, LINEAR_STORAGE_STATE, unique_marker from e2e_http import AuthHeaders, NoBody, get_external, unwrap +from e2e_metadata import Domain, Route, Subject, meta from idp import Identity, Keycloak from lifecycle import ResourceManager from models import ( @@ -88,6 +89,12 @@ class TestMcpOauthHappyPath: @pytest.mark.covers("mcp.list_tools.oauth.succeeds") @pytest.mark.covers("mcp.call_tool.oauth.succeeds") @pytest.mark.covers("mcp.call_tool.oauth.persists_across_processes") + @meta( + Subject( + domain=Domain.MCP, + route=Route.MCP, + ) + ) @pytest.mark.parametrize("route", ("aggregate_sso", "explicit_header_jwt")) @pytest.mark.parametrize("observed", (False, True), ids=("direct", "observed")) def test_consent_list_call_and_cold_restart( diff --git a/tests/e2e/mcp/test_mcp_toolset_enforcement_e2e.py b/tests/e2e/mcp/test_mcp_toolset_enforcement_e2e.py index 6b901145eb1..cff72f3c95d 100644 --- a/tests/e2e/mcp/test_mcp_toolset_enforcement_e2e.py +++ b/tests/e2e/mcp/test_mcp_toolset_enforcement_e2e.py @@ -18,6 +18,7 @@ import pytest from datadog_mcp import SEARCH_LOGS_TOOL, register_datadog_mcp from e2e_config import unique_marker from e2e_http import unwrap +from e2e_metadata import Domain, Route, Subject, meta from lifecycle import ResourceManager from mcp_client import McpClient from models import ToolsetCreateBody, ToolsetTool @@ -60,6 +61,12 @@ def _wire_prefix(wire_name: str, tool_name: str, catalog: frozenset[str]) -> str class TestMcpToolsetEnforcement: @pytest.mark.covers("mcp.list_tools.api_key.toolset_scoped") + @meta( + Subject( + domain=Domain.MCP, + route=Route.MCP, + ) + ) def test_key_granted_a_toolset_lists_exactly_its_tools(self, client: McpClient, resources: ResourceManager) -> None: server_id: Final = register_datadog_mcp(client, resources, allowed_tools=None) client.await_registered(server_id) 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/models.py b/tests/e2e/models.py index 8796393f269..3ae200ff5e7 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -1257,6 +1257,7 @@ class LiteLLMParamsBody(BaseModel): api_version: str | None = None realtime_protocol: str | None = None allowed_openai_params: list[str] | None = None + drop_params: bool | None = None aws_access_key_id: str | None = Field(default=None, repr=False) aws_secret_access_key: str | None = Field(default=None, repr=False) aws_region_name: str | None = None 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/quota_management/budgets/budget_client.py b/tests/e2e/quota_management/budgets/budget_client.py index 087dc8ca522..926dae0d954 100644 --- a/tests/e2e/quota_management/budgets/budget_client.py +++ b/tests/e2e/quota_management/budgets/budget_client.py @@ -17,6 +17,7 @@ from datetime import datetime from pydantic import AliasPath, BaseModel, Field, RootModel from e2e_http import NoBody, Result, StreamingResponse, Success, unwrap +from e2e_metadata import step from proxy_client import ProxyClient from models import ( AnthropicMessagesBody, @@ -253,15 +254,18 @@ class BudgetClient: ) ) + @step("Delete the virtual key") def delete_key(self, key: str) -> None: self.proxy.delete_key(key) + @step("Read the key's budget windows from /key/info") def key_budget_windows(self, key: str) -> list[BudgetWindowState]: """A key's budget_limits windows as /key/info stores them. Each window's reset_at is advanced by the reset job in the same pass that zeroes the window's spend counter, so a strictly-later value proves the wipe ran.""" return self.proxy.key_info(key).budget_limits or [] + @step("Read the team's budget windows from /team/info") def team_budget_windows(self, team_id: str) -> list[BudgetWindowState]: """Team analog of key_budget_windows, read from /team/info.""" match self._team_info(team_id): @@ -270,11 +274,13 @@ class BudgetClient: case _: return [] + @step("Delete the end users {user_ids}") def delete_customers(self, user_ids: list[str]) -> None: self.proxy.delete_customers(user_ids) # ---- chat (raw HTTP outcome: a budget block surfaces as a non-2xx) -- + @step('Send a /chat/completions request to {model} with the prompt "{content}"') def chat( self, key: str, @@ -297,6 +303,7 @@ class BudgetClient: ), ) + @step('Send a /v1/messages request to {model} with the prompt "{content}"') def messages( self, key: str, @@ -317,6 +324,7 @@ class BudgetClient: # ---- internal user -------------------------------------------------- + @step("Create an internal user with max budget: {max_budget}") def create_user(self, *, max_budget: float, budget_duration: str | None = None) -> str: return unwrap( self.proxy.transport.post( @@ -327,6 +335,7 @@ class BudgetClient: ) ).user_id + @step("Delete the internal user") def delete_user(self, user_id: str) -> None: _ = self.proxy.transport.post( "/user/delete", @@ -335,6 +344,7 @@ class BudgetClient: response_type=NoBody, ) + @step("Read the internal user's spend and budget from /user/info") def user_info(self, user_id: str) -> UserInfoRow | None: result = self.proxy.transport.get( "/user/info", @@ -350,6 +360,7 @@ class BudgetClient: # ---- customer / end-user ------------------------------------------- + @step("Create the end user {customer_id}") def create_customer( self, customer_id: str, @@ -369,6 +380,7 @@ class BudgetClient: # ---- organization --------------------------------------------------- + @step("Create the organization {alias} with max budget: {max_budget}") def create_org(self, *, max_budget: float, alias: str, budget_duration: str | None = None) -> str: return unwrap( self.proxy.transport.post( @@ -383,6 +395,7 @@ class BudgetClient: ) ).organization_id + @step("Read the organization's budget id from /organization/info") def org_budget_id(self, org_id: str) -> str | None: """The id of the budget row backing an org; its budget_reset_at is read via budget_info (LIT-4570: /organization/new stores budget_duration without @@ -399,6 +412,7 @@ class BudgetClient: case _: return None + @step("Delete the organization") def delete_org(self, org_id: str) -> None: _ = self.proxy.transport.delete( "/organization/delete", @@ -409,6 +423,7 @@ class BudgetClient: # ---- team ----------------------------------------------------------- + @step("Create the team {alias} and wait until /team/info returns it") def create_team( self, *, @@ -435,6 +450,7 @@ class BudgetClient: self._wait_for_team(team_id) return team_id + @step("Delete the team") def delete_team(self, team_id: str) -> None: _ = self.proxy.transport.post( "/team/delete", @@ -463,6 +479,7 @@ class BudgetClient: assert last is not None raise AssertionError(last) + @step("Add the internal user to the team") def add_team_member(self, team_id: str, user_id: str, *, max_budget_in_team: float | None = None) -> None: last_body = "" for attempt in range(_TEAM_READY_ATTEMPTS): @@ -484,6 +501,7 @@ class BudgetClient: break raise AssertionError(last_body) + @step("Update the team member's budget with /team/member_update") def update_team_member( self, team_id: str, @@ -504,6 +522,7 @@ class BudgetClient: ) assert resp.ok, resp.body + @step("Read the team member's budget reset time from /team/info") def member_budget_reset_at(self, team_id: str, user_id: str) -> str | None: """The member's per-team budget_reset_at as /team/info reports it, or None if no reset is scheduled. The reset job advances this each time the window @@ -519,6 +538,7 @@ class BudgetClient: # ---- tag ------------------------------------------------------------ + @step("Create the tag {name} with max budget: {max_budget}") def create_tag(self, name: str, *, max_budget: float) -> str: resp = self.proxy.transport.send( "/tag/new", @@ -528,6 +548,7 @@ class BudgetClient: assert resp.ok, resp.body return name + @step("Delete the tag {name}") def delete_tag(self, name: str) -> None: _ = self.proxy.transport.post( "/tag/delete", @@ -538,6 +559,7 @@ class BudgetClient: # ---- model access group --------------------------------------------- + @step("Set a shared budget on the model access group {access_group}") def set_access_group_budget( self, access_group: str, @@ -561,6 +583,7 @@ class BudgetClient: ) ) + @step("Read the budget and spend of the model access group {access_group}") def access_group_budget(self, access_group: str) -> AccessGroupBudgetResponse: return unwrap( self.proxy.transport.get( @@ -571,6 +594,7 @@ class BudgetClient: ) ) + @step("Delete the budget on the model access group {access_group}") def delete_access_group_budget(self, access_group: str) -> None: _ = self.proxy.transport.delete( f"/access_group/{access_group}/budget", @@ -581,6 +605,7 @@ class BudgetClient: # ---- budget table --------------------------------------------------- + @step("Create a budget with /budget/new") def create_budget( self, *, @@ -603,6 +628,7 @@ class BudgetClient: ) ).budget_id + @step("Delete the budget") def delete_budget(self, budget_id: str) -> None: _ = self.proxy.transport.post( "/budget/delete", @@ -611,6 +637,7 @@ class BudgetClient: response_type=NoBody, ) + @step("Read the budget from /budget/info") def budget_info(self, budget_id: str) -> tuple[BudgetRow, ...]: result = self.proxy.transport.post( "/budget/info", diff --git a/tests/e2e/quota_management/ratelimit/test_model_group_alias_rate_limit_e2e.py b/tests/e2e/quota_management/ratelimit/test_model_group_alias_rate_limit_e2e.py index 1c3fff47b78..523436f8d1e 100644 --- a/tests/e2e/quota_management/ratelimit/test_model_group_alias_rate_limit_e2e.py +++ b/tests/e2e/quota_management/ratelimit/test_model_group_alias_rate_limit_e2e.py @@ -21,6 +21,7 @@ import time import pytest from e2e_config import unique_marker from e2e_http import StreamingResponse, require_successful_call +from e2e_metadata import Domain, Mode, Provider, Subject, meta from quota_client import QuotaClient pytestmark = pytest.mark.e2e @@ -75,11 +76,27 @@ def _assert_blocked_inside_window( class TestModelGroupAliasRateLimit: @pytest.mark.covers("quota_management.ratelimit.model_group_alias.shares_bucket") + @meta( + Subject( + domain=Domain.SPEND_BUDGETS, + providers=(Provider.ANTHROPIC,), + models=(MODEL_GROUP, MODEL_ALIAS), + mode=Mode.NONSTREAM, + ) + ) def test_alias_shares_rpm_bucket_with_model_group(self, client: QuotaClient, scoped_key: str) -> None: opened_at = _exhaust_rpm(client, scoped_key, MODEL_GROUP) _assert_blocked_inside_window(client, scoped_key, MODEL_ALIAS, opened_at) @pytest.mark.covers("quota_management.ratelimit.model_group_alias.shares_bucket") + @meta( + Subject( + domain=Domain.SPEND_BUDGETS, + providers=(Provider.ANTHROPIC,), + models=(MODEL_GROUP, MODEL_ALIAS), + mode=Mode.NONSTREAM, + ) + ) def test_model_group_shares_rpm_bucket_with_alias(self, client: QuotaClient, scoped_key: str) -> None: opened_at = _exhaust_rpm(client, scoped_key, MODEL_ALIAS) _assert_blocked_inside_window(client, scoped_key, MODEL_GROUP, opened_at) diff --git a/tests/e2e/quota_management/spend_tracking/cost_rows.py b/tests/e2e/quota_management/spend_tracking/cost_rows.py index 87af54fe83f..8b15c555c66 100644 --- a/tests/e2e/quota_management/spend_tracking/cost_rows.py +++ b/tests/e2e/quota_management/spend_tracking/cost_rows.py @@ -36,6 +36,7 @@ from pydantic import BaseModel, RootModel from e2e_config import unique_marker from e2e_http import Success +from e2e_metadata import step from lifecycle import ResourceManager from models import LiteLLMParamsBody, SpendLogsParams from proxy_client import ProxyClient @@ -133,6 +134,7 @@ def assert_fresh_tokens_billed_at(row: CostRow, input_rate: float) -> None: ) +@step("Wait for the request's cost breakdown in /spend/logs") def poll_cost_row(proxy: ProxyClient, request_id: str) -> CostRow | None: """Poll /spend/logs for the call's row until it lands with a cost breakdown (rows flush ~60s behind the call via proxy_batch_write_at); None on timeout.""" @@ -156,6 +158,7 @@ def poll_cost_row(proxy: ProxyClient, request_id: str) -> CostRow | None: return None +@step("Wait for a matching cost breakdown in the key's /spend/logs") def poll_cost_row_where( proxy: ProxyClient, api_key: str, predicate: Callable[[CostRow], bool] ) -> CostRow | None: @@ -182,6 +185,7 @@ def poll_cost_row_where( return None +@step("Add a deployment with custom rates that calls {litellm_params.model}") def register_priced_model( proxy: ProxyClient, resources: ResourceManager, diff --git a/tests/e2e/quota_management/spend_tracking/spend_e2e_client.py b/tests/e2e/quota_management/spend_tracking/spend_e2e_client.py index e607c12b731..60372a35afe 100644 --- a/tests/e2e/quota_management/spend_tracking/spend_e2e_client.py +++ b/tests/e2e/quota_management/spend_tracking/spend_e2e_client.py @@ -31,6 +31,7 @@ from e2e_http import ( is_ok, unwrap, ) +from e2e_metadata import step from models import ( AnthropicMessagesBody, ChatBody, @@ -264,6 +265,7 @@ def _chat_body( class SpendClient: proxy: ProxyClient + @step('Send a /chat/completions request to {model} with the prompt "{content}"') def chat( self, key: str, @@ -280,6 +282,7 @@ class SpendClient: _chat_body(model, content, max_tokens=max_tokens, tags=tags, user=user, cache=cache), ) + @step('Send a streaming /chat/completions request to {model} with the prompt "{content}"') def chat_stream( self, key: str, model: str, content: str, *, max_tokens: int | None = None ) -> StreamingResponse: @@ -287,6 +290,7 @@ class SpendClient: key, _chat_body(model, content, max_tokens=max_tokens, stream=True) ) + @step('Send a streaming /v1/messages request to {model} with the prompt "{content}"') def messages_stream( self, key: str, model: str, content: str, *, max_tokens: int ) -> StreamingResponse: @@ -300,9 +304,11 @@ class SpendClient: ), ) + @step('Send an /embeddings request to {model} for "{content}"') def embed(self, key: str, model: str, content: str) -> Result[EmbedResponse]: return self.proxy.embed(key, EmbedBody(model=model, input=content)) + @step("Wait for at least {min_rows} of the key's spend logs in /spend/logs") def poll_logs_for_key( self, key: str, @@ -314,6 +320,7 @@ class SpendClient: key, min_rows=min_rows, predicate=predicate ) + @step('Estimate the cost of sending "{content}" to {model} with /spend/calculate') def calculate_spend(self, model: str, content: str) -> float: return unwrap( self.proxy.transport.post( @@ -326,6 +333,7 @@ class SpendClient: ) ).cost + @step("Read the spend per tag from /spend/tags") def spend_by_tags(self) -> list[TagSpend]: result = self.proxy.transport.get( "/spend/tags", @@ -339,6 +347,7 @@ class SpendClient: case _: return [] + @step("Wait for the tag {tag} to reach the expected spend in /spend/tags") def poll_tag_spend(self, tag: str, *, minimum: float = 0.0) -> TagSpend | None: """Poll /spend/tags until the tag's aggregate reaches `minimum`; last seen.""" deadline = time.monotonic() + self.proxy.poll_timeout @@ -354,6 +363,7 @@ class SpendClient: time.sleep(self.proxy.poll_interval) return entry + @step("Wait for the key's spend in /key/info to reach the expected minimum") def poll_key_spend(self, key: str, *, minimum: float = 0.0) -> float: deadline = time.monotonic() + self.proxy.poll_timeout spend = 0.0 @@ -364,6 +374,7 @@ class SpendClient: time.sleep(self.proxy.poll_interval) return spend + @step("Read the team's spend from /team/info") def team_spend(self, team_id: str) -> float: return ( unwrap( @@ -377,6 +388,7 @@ class SpendClient: or 0.0 ) + @step("Wait for the team's spend in /team/info to reach the expected minimum") def poll_team_spend(self, team_id: str, *, minimum: float = 0.0) -> float: outcome: Final = await_converged( lambda: self.team_spend(team_id), @@ -388,6 +400,7 @@ class SpendClient: ) return outcome.result if isinstance(outcome, Converged) else outcome.last_result + @step("Read the end user's spend from /customer/info") def customer_spend(self, customer_id: str) -> float: """0.0 until the spend writer has upserted the end-user row, which /customer/info 404s before.""" looked_up: Final = self.proxy.transport.get( @@ -402,6 +415,7 @@ class SpendClient: case _: return 0.0 + @step("Wait for the end user's spend in /customer/info to go above {minimum}") def poll_customer_spend(self, customer_id: str, *, minimum: float = 0.0) -> float: outcome: Final = await_converged( lambda: self.customer_spend(customer_id), @@ -413,6 +427,7 @@ class SpendClient: ) return outcome.result if isinstance(outcome, Converged) else outcome.last_result + @step("Scrape /metrics/ on every proxy replica") def scrape_metrics(self) -> Mapping[str, ProbeResult]: """GET /metrics/ on every replica in PROXY_REPLICA_URLS, keyed by replica. The counter is per pod, so the union of the replicas is the fleet's exposition; the @@ -425,6 +440,7 @@ class SpendClient: } ) + @step("Read page {page} of /spend/logs/v2 at a page size of {page_size}") def spend_logs_page( self, *, api_key: str | None, page: int, page_size: int ) -> SpendLogsPage: @@ -447,9 +463,11 @@ class SpendClient: ) ) + @step("Call the management route {path}") def probe(self, path: str, *, params: DateRangeParams) -> ProbeResult: return self.proxy.transport.probe(path, params=params) + @step("Call the management route {path} until it answers successfully") def probe_until_healthy(self, path: str, *, params: DateRangeParams) -> ProbeResult: outcome: Final = await_converged( lambda: self.probe(path, params=params), @@ -461,6 +479,7 @@ class SpendClient: ) return outcome.result if isinstance(outcome, Converged) else outcome.last_result + @step("Create an internal user with the role {role}") def create_user(self, *, email: str, role: UserRole, user_id: str) -> str: return unwrap( self.proxy.transport.post( @@ -471,6 +490,7 @@ class SpendClient: ) ).user_id + @step("Delete the internal user") def delete_user(self, user_id: str) -> None: _ = unwrap( self.proxy.transport.post( @@ -481,6 +501,7 @@ class SpendClient: ) ) + @step("Generate a virtual key with {body}") def generate_key_record(self, body: KeyGenerateBody) -> KeyGenerateResponse: return unwrap( self.proxy.transport.post( @@ -491,6 +512,7 @@ class SpendClient: ) ) + @step('Send a /chat/completions request to {model} with the prompt "{content}"') def send_chat(self, key: str, model: str, content: str, *, max_tokens: int) -> StreamingResponse: return self.proxy.transport.send( "/chat/completions", @@ -498,6 +520,7 @@ class SpendClient: json=_chat_body(model, content, max_tokens=max_tokens), ) + @step('Send a /queue/chat/completions request to {model} with the prompt "{content}"') def send_queued_chat(self, key: str, model: str, content: str, *, max_tokens: int) -> StreamingResponse: return self.proxy.transport.send( "/queue/chat/completions", @@ -509,6 +532,7 @@ class SpendClient: ), ) + @step('Send a /v1/messages request to {model} with the prompt "{content}"') def send_messages(self, key: str, model: str, content: str, *, max_tokens: int) -> StreamingResponse: return self.proxy.transport.send( "/v1/messages", @@ -520,9 +544,11 @@ class SpendClient: ), ) + @step('Send a /v1/responses request to {model} with the prompt "{content}"') def send_responses(self, key: str, model: str, content: str) -> StreamingResponse: return self.send_responses_with_headers(self.proxy.transport.bearer(key), model, content) + @step('Send a /v1/responses request to {model} with custom headers and the prompt "{content}"') def send_responses_with_headers(self, headers: AuthHeaders, model: str, content: str) -> StreamingResponse: return self.proxy.transport.send( "/v1/responses", @@ -530,6 +556,7 @@ class SpendClient: json=ResponsesBody(model=model, input=content), ) + @step('Send an /embeddings request to {model} for "{content}"') def send_embed(self, key: str, model: str, content: str) -> StreamingResponse: return self.proxy.transport.send( "/embeddings", @@ -537,6 +564,7 @@ class SpendClient: json=EmbedBody(model=model, input=content), ) + @step('Send a Gemini generateContent request to {model} through /gemini with the prompt "{content}"') def send_gemini_generate(self, key: str, model: str, content: str, *, max_tokens: int) -> StreamingResponse: return self.proxy.transport.send( f"/gemini/v1beta/models/{model}:generateContent", @@ -547,6 +575,7 @@ class SpendClient: ), ) + @step("Upload a batch input file for {model} to /v1/files") def upload_batch_file(self, key: str, model: str, content: bytes) -> FileObject: return unwrap( self.proxy.transport.upload( @@ -560,6 +589,7 @@ class SpendClient: ) ) + @step("Create a batch for {body.model} on /v1/batches") def create_batch(self, key: str, body: BatchCreateBody) -> BatchObject: return unwrap( self.proxy.transport.post( @@ -570,6 +600,7 @@ class SpendClient: ) ) + @step("Retrieve the {provider} batch from /v1/batches") def retrieve_batch(self, key: str, batch_id: str, *, provider: str) -> BatchObject: return unwrap( self.proxy.transport.get( @@ -580,6 +611,7 @@ class SpendClient: ) ) + @step("Post a callback log for {payload.model} to /v1/rust_control_plane/logs") def replay_callback_log(self, key: str, payload: CallbackLogPayload) -> CallbackLogsResponse: return unwrap( self.proxy.transport.post( @@ -590,12 +622,15 @@ class SpendClient: ) ) + @step("Run a health check on {model} with /health") def health(self, model: str) -> ProbeResult: return self.proxy.transport.probe("/health", params=HealthParams(model=model)) + @step("Read the key's daily activity from /user/daily/activity") def daily_activity_for_key(self, token: str, *, start: datetime, end: datetime) -> DailyActivityKeyBreakdown | None: return self._key_breakdown("/user/daily/activity", token, start=start, end=end) + @step("Read the key's usage export row from /user/daily/activity/aggregated") def usage_export_row_for_key( self, token: str, *, start: datetime, end: datetime ) -> DailyActivityKeyBreakdown | None: @@ -623,11 +658,13 @@ class SpendClient: None, ) + @step("Wait for at least {min_requests} of the key's requests in /user/daily/activity") def poll_daily_activity_for_key( self, token: str, *, start: datetime, end: datetime, min_requests: int ) -> DailyActivityKeyBreakdown | None: return self._poll_key_breakdown(lambda: self.daily_activity_for_key(token, start=start, end=end), min_requests) + @step("Wait for at least {min_requests} of the key's requests in /user/daily/activity/aggregated") def poll_usage_export_row_for_key( self, token: str, *, start: datetime, end: datetime, min_requests: int ) -> DailyActivityKeyBreakdown | None: @@ -648,6 +685,7 @@ class SpendClient: ) return outcome.result if isinstance(outcome, Converged) else outcome.last_result + @step("Read the OpenAPI schema from /openapi.json") def openapi(self) -> OpenAPISchema: return unwrap( self.proxy.transport.get( diff --git a/tests/e2e/quota_management/spend_tracking/spend_reconciliation.py b/tests/e2e/quota_management/spend_tracking/spend_reconciliation.py index f313325dbda..96f1fbd4a7a 100644 --- a/tests/e2e/quota_management/spend_tracking/spend_reconciliation.py +++ b/tests/e2e/quota_management/spend_tracking/spend_reconciliation.py @@ -5,6 +5,7 @@ from typing import Final from e2e_config import provider_edge_base, unique_marker from e2e_http import unwrap +from e2e_metadata import step from lifecycle import ResourceManager from models import ChatBody, ChatMessage, ChatResponse, KeyGenerateBody, LiteLLMParamsBody, TeamNewBody from spend_e2e_client import SpendClient @@ -33,6 +34,10 @@ class TeamTraffic: return self.prompt_tokens * INPUT_RATE + self.completion_tokens * OUTPUT_RATE +@step( + "Add a priced deployment, then create two teams with one key each" + " and send 7 /chat/completions requests per key, 6 of them at once" +) def create_traffic(client: SpendClient, resources: ResourceManager) -> tuple[TeamTraffic, ...]: base: Final = provider_edge_base("openai") model: Final = f"e2e-reconciliation-{unique_marker()}" @@ -85,6 +90,7 @@ def create_traffic(client: SpendClient, resources: ResourceManager) -> tuple[Tea return tuple(team_traffic() for _ in range(2)) +@step("Check that the /spend/logs rows of team {traffic.team_id} match each response's tokens and cost") def assert_logs_match(client: SpendClient, traffic: TeamTraffic) -> None: expected_ids: Final = frozenset(response.id for response in traffic.responses) assert len(expected_ids) == len(traffic.responses), "responses must have distinct IDs" diff --git a/tests/e2e/quota_management/spend_tracking/test_service_tier_pricing_e2e.py b/tests/e2e/quota_management/spend_tracking/test_service_tier_pricing_e2e.py index 76d80b1aab8..31c9b90d209 100644 --- a/tests/e2e/quota_management/spend_tracking/test_service_tier_pricing_e2e.py +++ b/tests/e2e/quota_management/spend_tracking/test_service_tier_pricing_e2e.py @@ -36,7 +36,7 @@ from cost_rows import ( ) from e2e_config import CHEAP_OPENAI_MODEL, unique_marker from e2e_http import unwrap -from e2e_metadata import Capability, Domain, Mode, Provider, Subject, meta +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import ( AnthropicMessagesBody, @@ -184,6 +184,15 @@ class TestServiceTierPricing: assert_total_is_sum_of_components(row) @pytest.mark.covers("quota_management.spend_tracking.service_tier_stream.records_served_tier") + @meta( + Subject( + domain=Domain.SPEND_BUDGETS, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(BACKEND,), + mode=Mode.STREAM, + ) + ) def test_streamed_call_records_and_bills_the_served_tier( self, client: SpendClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -232,6 +241,15 @@ class TestServiceTierPricing: assert_total_is_sum_of_components(row) @pytest.mark.covers("llm.chat_completions.openai.service_tier.stream.echoes_served_tier") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.OPENAI,), + models=(STREAM_BACKEND,), + mode=Mode.STREAM, + ) + ) def test_every_streamed_chunk_carries_the_served_tier( self, client: SpendClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -262,6 +280,15 @@ class TestServiceTierPricing: ) @pytest.mark.covers("quota_management.spend_tracking.service_tier_stream.responses_records_served_tier") + @meta( + Subject( + domain=Domain.SPEND_BUDGETS, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(STREAM_BACKEND,), + mode=Mode.STREAM, + ) + ) def test_responses_stream_records_the_served_tier( self, client: SpendClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -308,6 +335,15 @@ class TestServiceTierPricing: assert_fresh_tokens_billed_at(row, INPUT_RATE_FOR_PRICING_BASIS[pricing_basis]) @pytest.mark.covers("quota_management.spend_tracking.service_tier_stream.messages_records_served_tier") + @meta( + Subject( + domain=Domain.SPEND_BUDGETS, + route=Route.MESSAGES, + providers=(Provider.OPENAI,), + models=(STREAM_BACKEND,), + mode=Mode.STREAM, + ) + ) def test_messages_stream_records_the_served_tier( self, client: SpendClient, resources: ResourceManager, scoped_key: str ) -> None: diff --git a/tests/e2e/quota_management/spend_tracking/test_spend_routes.py b/tests/e2e/quota_management/spend_tracking/test_spend_routes.py index 7b5db9ccd27..1b2069abdc6 100644 --- a/tests/e2e/quota_management/spend_tracking/test_spend_routes.py +++ b/tests/e2e/quota_management/spend_tracking/test_spend_routes.py @@ -163,6 +163,12 @@ def test_schema_listed_spend_routes_are_responsive(client: SpendClient) -> None: assert not offenders, "non-responsive schema spend routes:\n" + "\n".join(offenders) +@meta( + Subject( + domain=Domain.SPEND_BUDGETS, + route=Route.SPEND_REPORTING, + ) +) def test_capture_rate_reports_or_names_the_missing_billing_key(client: SpendClient) -> None: result: Final = client.probe(_CAPTURE_RATE_ROUTE, params=_date_range()) print(f"{_CAPTURE_RATE_ROUTE} -> {result.status_code}\n{result.body[:600]}") diff --git a/tests/e2e/quota_management/spend_tracking/test_spend_surface_consistency_e2e.py b/tests/e2e/quota_management/spend_tracking/test_spend_surface_consistency_e2e.py index 9033b9d75c4..63ba785a2ee 100644 --- a/tests/e2e/quota_management/spend_tracking/test_spend_surface_consistency_e2e.py +++ b/tests/e2e/quota_management/spend_tracking/test_spend_surface_consistency_e2e.py @@ -32,6 +32,7 @@ from typing import Final import pytest from e2e_config import provider_edge_base, unique_marker from e2e_http import ProbeResult +from e2e_metadata import Domain, Mode, Provider, Subject, meta from lifecycle import ResourceManager from models import ChatBody, ChatMessage, KeyGenerateBody, LiteLLMParamsBody, TeamNewBody from prometheus_client.parser import text_string_to_metric_families @@ -41,6 +42,7 @@ from spend_reconciliation import INPUT_RATE, OUTPUT_RATE pytestmark = pytest.mark.e2e +BACKEND: Final = "openai/gpt-5.6-luna" SPEND_METRIC: Final = "litellm_spend_metric_total" KEY_HASH_LABEL: Final = "hashed_api_key" TEAM_LABEL: Final = "team" @@ -86,6 +88,14 @@ def _same_spend(actual: float | None, expected: float) -> bool: class TestSpendSurfaceConsistency: @pytest.mark.replayable @pytest.mark.covers("quota_management.spend_tracking.surface_consistency.matches_every_surface") + @meta( + Subject( + domain=Domain.SPEND_BUDGETS, + providers=(Provider.OPENAI,), + models=(BACKEND,), + mode=Mode.NONSTREAM, + ) + ) def test_one_request_lands_the_same_spend_on_every_surface( self, client: SpendClient, resources: ResourceManager ) -> None: @@ -96,7 +106,7 @@ class TestSpendSurfaceConsistency: model_id: Final = client.proxy.create_model( model, LiteLLMParamsBody( - model="openai/gpt-5.6-luna", + model=BACKEND, api_key="os.environ/OPENAI_API_KEY", api_base=None if base is None else f"{base}/v1", input_cost_per_token=INPUT_RATE, diff --git a/tests/e2e/quota_management/spend_tracking/test_spend_tracking_e2e.py b/tests/e2e/quota_management/spend_tracking/test_spend_tracking_e2e.py index a4c37c2df94..c7ef81f826a 100644 --- a/tests/e2e/quota_management/spend_tracking/test_spend_tracking_e2e.py +++ b/tests/e2e/quota_management/spend_tracking/test_spend_tracking_e2e.py @@ -42,6 +42,7 @@ CLAUDE_MODEL = "claude-haiku-4-5" CODEX_MODEL = "openai-responses-codex" EMBEDDING_MODEL = "openai-text-embedding-3-small" OPENAI_BACKEND = "openai/gpt-5.5" +ANTHROPIC_BACKEND: Final = "anthropic/claude-haiku-4-5" def _approx_equal(actual: float, expected: float) -> bool: @@ -534,6 +535,15 @@ def test_end_user_spend_attributed_on_row( @pytest.mark.covers("quota_management.spend_tracking.end_user.attributes_responses_header") +@meta( + Subject( + domain=Domain.SPEND_BUDGETS, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(CODEX_MODEL,), + mode=Mode.NONSTREAM, + ) +) @pytest.mark.parametrize("header", ["x-litellm-customer-id", "x-litellm-end-user-id"]) def test_end_user_header_attributes_responses_row( client: SpendClient, scoped_key: str, resources: ResourceManager, header: str @@ -659,6 +669,14 @@ def test_failure_call_writes_failure_status_row( @pytest.mark.covers("quota_management.spend_tracking.failure.writes_normalized_error") +@meta( + Subject( + domain=Domain.SPEND_BUDGETS, + providers=(Provider.ANTHROPIC, Provider.OPENAI), + models=(OPENAI_BACKEND, ANTHROPIC_BACKEND), + mode=Mode.NONSTREAM, + ) +) def test_failure_rows_share_normalized_error_across_provider_wording( client: SpendClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -668,7 +686,7 @@ def test_failure_rows_share_normalized_error_across_provider_wording( marker = unique_marker() deployments: Final = ( (f"e2e-norm-openai-{marker}", OPENAI_BACKEND), - (f"e2e-norm-anthropic-{marker}", "anthropic/claude-haiku-4-5"), + (f"e2e-norm-anthropic-{marker}", ANTHROPIC_BACKEND), ) for name, provider_model in deployments: model_id = client.proxy.create_model( @@ -700,6 +718,14 @@ def test_failure_rows_share_normalized_error_across_provider_wording( @pytest.mark.covers("quota_management.spend_tracking.failure.attributes_provider") +@meta( + Subject( + domain=Domain.SPEND_BUDGETS, + providers=(Provider.OPENAI,), + models=(OPENAI_BACKEND,), + mode=Mode.NONSTREAM, + ) +) def test_pre_call_rejection_row_attributes_provider_and_model_id( client: SpendClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/quota_management/spend_tracking/test_websearch_interception_session_e2e.py b/tests/e2e/quota_management/spend_tracking/test_websearch_interception_session_e2e.py index 6352ab67c3c..ab7c2370365 100644 --- a/tests/e2e/quota_management/spend_tracking/test_websearch_interception_session_e2e.py +++ b/tests/e2e/quota_management/spend_tracking/test_websearch_interception_session_e2e.py @@ -17,6 +17,7 @@ from typing import Final, Literal import pytest from e2e_config import unique_marker from e2e_http import unwrap +from e2e_metadata import Capability, Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import ( AnthropicContentBlock, @@ -56,6 +57,16 @@ class TestWebSearchInterceptionSession: "quota_management.spend_tracking.websearch_interception.bills_under_request_session", exercised_on=("messages",), ) + @meta( + Subject( + domain=Domain.SPEND_BUDGETS, + route=Route.MESSAGES, + providers=(Provider.BEDROCK, Provider.PERPLEXITY), + models=(BEDROCK_INVOKE_BACKEND,), + capabilities=(Capability.WEB_SEARCH,), + mode=Mode.NONSTREAM, + ) + ) def test_intercepted_search_is_billed_under_the_request_session( self, proxy: ProxyClient, resources: ResourceManager ) -> None: diff --git a/tests/e2e/router/reliability_support.py b/tests/e2e/router/reliability_support.py index 04fa30a0d14..03e585273c0 100644 --- a/tests/e2e/router/reliability_support.py +++ b/tests/e2e/router/reliability_support.py @@ -23,6 +23,7 @@ from pydantic import BaseModel, ValidationError from proxy_client import ProxyClient from e2e_config import CHEAP_OPENAI_MODEL, PROXY_BASE_URL, unique_marker from e2e_http import NetworkError, StreamHead, StreamingResponse +from e2e_metadata import step from models import ( CacheControl, ChatMessage, @@ -79,6 +80,7 @@ def cached_system_turn(marker: str) -> ChatMessage: return ChatMessage(role="system", content=[TextContentPart(text=filler, cache_control=CacheControl())]) +@step(f"Add a deployment named {{name}} that calls {REAL_MODEL} at an unreachable address") def create_bad_base_deployment(proxy: ProxyClient, name: str) -> str: """Register a deployment pointing at an unreachable base, so every call to it fails with a real connection error the fallback can reroute around.""" @@ -87,6 +89,7 @@ def create_bad_base_deployment(proxy: ProxyClient, name: str) -> str: ) +@step(f"Add a deployment named {{name}} that calls {REAL_MODEL} at an unreachable address and is never benched") def create_never_benched_refusing_deployment(proxy: ProxyClient, name: str) -> str: return proxy.create_model( name, @@ -94,6 +97,7 @@ def create_never_benched_refusing_deployment(proxy: ProxyClient, name: str) -> s ) +@step(f"Add a deployment named {{name}} that calls {REAL_MODEL} with a 1ms timeout") def create_timeout_deployment(proxy: ProxyClient, name: str) -> str: """Register a deployment with a 1ms deadline the real backend always exceeds.""" return proxy.create_model( @@ -101,12 +105,14 @@ def create_timeout_deployment(proxy: ProxyClient, name: str) -> str: ) +@step(f"Add a deployment named {{name}} that calls the small-context model {SMALL_CONTEXT_MODEL}") def create_small_context_deployment(proxy: ProxyClient, name: str) -> str: """Register a deployment on the smallest-context model OpenAI still serves, so an oversized prompt earns a real context-window refusal from the provider.""" return proxy.create_model(name, LiteLLMParamsBody(model=SMALL_CONTEXT_MODEL, api_key=REAL_KEY)) +@step(f"Add a deployment named {{name}} that calls {AZURE_MODEL} behind Azure's content filter") def create_content_filtered_deployment(proxy: ProxyClient, name: str) -> str: """Register the Azure OpenAI deployment whose content filter refuses CONTENT_POLICY_PROMPT with a real policy-violation 400 (the one live trigger @@ -124,6 +130,10 @@ def create_content_filtered_deployment(proxy: ProxyClient, name: str) -> str: ) +@step( + f"Add a deployment named {{name}} that calls {AZURE_MODEL} and is benched for {{cooldown_time}}s" + " on its first failure" +) def create_azure_benched_on_first_failure_deployment(proxy: ProxyClient, name: str, cooldown_time: float) -> str: """The live Azure OpenAI deployment holding all of the group's shuffle weight, benched on its first failure of any class, with the client's own retries off.""" @@ -144,6 +154,7 @@ def create_azure_benched_on_first_failure_deployment(proxy: ProxyClient, name: s ) +@step(f"Add a deployment named {{name}} that calls {CACHING_MODEL} with prompt caching") def create_caching_deployment(proxy: ProxyClient, name: str) -> str: """Register the Anthropic deployment whose prompt cache the affinity check pins to.""" return proxy.create_model(name, LiteLLMParamsBody(model=CACHING_MODEL, api_key=CACHING_KEY, weight=1)) @@ -165,6 +176,7 @@ def _register_benched_on_first_failure( ) +@step(f"Add a deployment named {{name}} that calls {REAL_MODEL} with a 1ms timeout and is benched on its first timeout") def create_always_timing_out_deployment(proxy: ProxyClient, name: str, cooldown_time: float | None = None) -> str: """A 1ms deadline the real backend always exceeds, benched on its first Timeout.""" return _register_benched_on_first_failure( @@ -176,6 +188,7 @@ def create_always_timing_out_deployment(proxy: ProxyClient, name: str, cooldown_ ) +@step(f"Add a deployment named {{name}} that calls {REAL_MODEL} with an invalid key and is benched on its first 401") def create_always_unauthorized_deployment(proxy: ProxyClient, name: str, cooldown_time: float | None = None) -> str: """A key the real backend rejects with a 401, benched on its first AuthenticationError.""" return _register_benched_on_first_failure( @@ -204,6 +217,10 @@ def _nested_proxy_params(upstream_group: str, upstream_key: str, cooldown_time: ) +@step( + "Add a deployment named {name} that fronts {upstream_group} on this proxy, so it always gets a 500" + " and is benched on the first one" +) def create_always_5xx_deployment( proxy: ProxyClient, name: str, upstream_group: str, upstream_key: str, cooldown_time: float | None = None ) -> str: @@ -217,6 +234,10 @@ def create_always_5xx_deployment( ) +@step( + "Add a deployment named {name} that fronts {upstream_group} on this proxy with a key out of rpm," + " so it always gets a 429 and is benched on the first one" +) def create_always_rate_limited_deployment( proxy: ProxyClient, name: str, upstream_group: str, upstream_key: str, cooldown_time: float | None = None ) -> str: @@ -227,6 +248,7 @@ def create_always_rate_limited_deployment( ) +@step(f"Use up the rpm-limited key's one allowed request with a /chat/completions call to {CHEAP_OPENAI_MODEL}") def spend_only_request_of(proxy: ProxyClient, spent_key: str) -> None: """Uses up the one request an rpm_limit=1 key allows. The proxy's rate limiter opens the key's 60s window on this call, so it goes right before the calls that @@ -239,6 +261,7 @@ def spend_only_request_of(proxy: ProxyClient, spent_key: str) -> None: ) +@step(f"Add a deployment named {{name}} that calls {SMALL_CONTEXT_MODEL} and takes all of its group's traffic") def create_always_picked_small_context_deployment(proxy: ProxyClient, name: str) -> str: """The always-picked half of a retry pair on the smallest-context model OpenAI still serves: it holds all of the model group's shuffle weight, so an oversized @@ -253,12 +276,14 @@ def create_always_picked_small_context_deployment(proxy: ProxyClient, name: str) ) +@step(f"Add a deployment named {{name}} for {REAL_MODEL} that answers with a canned reply") def create_canned_deployment(proxy: ProxyClient, name: str) -> str: """A deployment that answers from a canned reply, so a call to it goes through the router's deployment pick like any other but never reaches a provider.""" return proxy.create_model(name, LiteLLMParamsBody(model=REAL_MODEL, mock_response="ok")) +@step(f"Add a zero-weight backup deployment named {{name}} that calls {REAL_MODEL}") def create_zero_weight_backup_deployment(proxy: ProxyClient, name: str) -> str: """The other half of a retry pair: healthy, but weight 0, so the weighted shuffle never opens on it. It is reachable only once its sibling is out of the running, @@ -273,6 +298,7 @@ def create_zero_weight_backup_deployment(proxy: ProxyClient, name: str) -> str: ) +@step("Send a /chat/completions request to {model} with a full message history and stream set to {stream}") def chat_turns_override( proxy: ProxyClient, key: str, @@ -300,6 +326,7 @@ def chat_turns_override( ) +@step("Send a /chat/completions request to {model} with stream set to {stream}") def chat_override( proxy: ProxyClient, key: str, @@ -322,6 +349,7 @@ def chat_override( ) +@step('Open a streaming /chat/completions request to {model} with the prompt "{content}" and leave it in flight') def open_chat_stream( proxy: ProxyClient, key: str, diff --git a/tests/e2e/router/test_auto_router_regressions_e2e.py b/tests/e2e/router/test_auto_router_regressions_e2e.py index 374badcf5fc..898b164e713 100644 --- a/tests/e2e/router/test_auto_router_regressions_e2e.py +++ b/tests/e2e/router/test_auto_router_regressions_e2e.py @@ -49,6 +49,7 @@ import pytest from pydantic import BaseModel, ConfigDict, Field from e2e_config import unique_marker +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from e2e_http import AnthropicHeaders, AuthHeaders, UnauthorizedError, unwrap from lifecycle import ResourceManager from models import ( @@ -341,6 +342,15 @@ def credentialed_alias(proxy: ProxyClient, router_stack: ExitStack) -> Credentia class TestTagSplitRouting: @pytest.mark.covers("reliability.routing.tagged_marker.request_tag_selects_marker") + @meta( + Subject( + domain=Domain.ROUTING, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(CHEAP_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_body_tagged_chat_routes_through_the_marker_to_its_tier( self, proxy: ProxyClient, resources: ResourceManager, plain_first_split: TagSplitDeployment ) -> None: @@ -355,6 +365,15 @@ class TestTagSplitRouting: _assert_served_only_by(rows, CHEAP_SERVED | {plain_first_split.tier}, "body-tagged chat on the shared name") @pytest.mark.covers("reliability.routing.tagged_marker.untagged_request_served_by_plain_deployment") + @meta( + Subject( + domain=Domain.ROUTING, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(PLAIN_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_untagged_chat_is_always_served_by_the_plain_deployment( self, proxy: ProxyClient, resources: ResourceManager, plain_first_split: TagSplitDeployment ) -> None: @@ -370,6 +389,15 @@ class TestTagSplitRouting: _assert_served_only_by(rows, PLAIN_SERVED | {plain_first_split.shared}, "untagged chat on the shared name") @pytest.mark.covers("reliability.routing.tagged_marker.untagged_request_served_by_plain_deployment") + @meta( + Subject( + domain=Domain.ROUTING, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=(PLAIN_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_untagged_messages_is_served_by_the_plain_deployment( self, proxy: ProxyClient, resources: ResourceManager, plain_first_split: TagSplitDeployment ) -> None: @@ -387,6 +415,15 @@ class TestTagSplitRouting: class TestUntaggedTierDeployments: @pytest.mark.covers("reliability.routing.tagged_marker.header_tag_selects_marker") + @meta( + Subject( + domain=Domain.ROUTING, + route=Route.MESSAGES, + providers=(Provider.ANTHROPIC,), + models=(CHEAP_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_header_tagged_messages_routes_through_the_marker_to_an_untagged_tier( self, proxy: ProxyClient, resources: ResourceManager, marker_first_split: TagSplitDeployment ) -> None: @@ -413,6 +450,15 @@ class TestUntaggedTierDeployments: ) @pytest.mark.covers("reliability.routing.tagged_marker.untagged_tier_deployments_still_served") + @meta( + Subject( + domain=Domain.ROUTING, + route=Route.CHAT_COMPLETIONS, + providers=(Provider.ANTHROPIC,), + models=(CHEAP_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_body_tagged_chat_reaches_the_untagged_tier_after_marker_rewrite( self, proxy: ProxyClient, resources: ResourceManager, marker_first_split: TagSplitDeployment ) -> None: @@ -431,6 +477,12 @@ class TestUntaggedTierDeployments: _assert_served_only_by(rows, CHEAP_SERVED | {marker_first_split.tier}, "body-tagged chat with untagged tier") @pytest.mark.covers("reliability.routing.tagged_marker.tag_semantics_stay_strict") + @meta( + Subject( + domain=Domain.ROUTING, + route=Route.CHAT_COMPLETIONS, + ) + ) def test_tagged_call_straight_at_an_untagged_deployment_stays_denied( self, proxy: ProxyClient, resources: ResourceManager, marker_first_split: TagSplitDeployment ) -> None: @@ -449,6 +501,15 @@ class TestUntaggedTierDeployments: class TestResponsesApiTagRouting: @pytest.mark.covers("reliability.routing.tagged_marker.responses_input_routes_through_marker") + @meta( + Subject( + domain=Domain.ROUTING, + route=Route.RESPONSES, + providers=(Provider.ANTHROPIC,), + models=(CHEAP_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_header_tagged_responses_with_string_input_routes_to_the_tier( self, proxy: ProxyClient, resources: ResourceManager, plain_first_split: TagSplitDeployment ) -> None: @@ -471,6 +532,15 @@ class TestResponsesApiTagRouting: ) @pytest.mark.covers("reliability.routing.tagged_marker.responses_input_routes_through_marker") + @meta( + Subject( + domain=Domain.ROUTING, + route=Route.RESPONSES, + providers=(Provider.ANTHROPIC,), + models=(CHEAP_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_body_tagged_responses_with_list_input_routes_to_the_tier( self, proxy: ProxyClient, resources: ResourceManager, plain_first_split: TagSplitDeployment ) -> None: @@ -497,6 +567,15 @@ class TestResponsesApiTagRouting: _assert_served_only_by(rows, CHEAP_SERVED | {plain_first_split.tier}, "body-tagged /v1/responses list input") @pytest.mark.covers("reliability.routing.tagged_marker.untagged_request_served_by_plain_deployment") + @meta( + Subject( + domain=Domain.ROUTING, + route=Route.RESPONSES, + providers=(Provider.ANTHROPIC,), + models=(PLAIN_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_untagged_responses_is_served_by_the_plain_deployment( self, proxy: ProxyClient, resources: ResourceManager, plain_first_split: TagSplitDeployment ) -> None: @@ -524,6 +603,14 @@ class TestResponsesApiTagRouting: class TestStrategyAliasPricing: @pytest.mark.covers("reliability.routing.strategy_alias.custom_pricing_ignored") + @meta( + Subject( + domain=Domain.SPEND_BUDGETS, + providers=(Provider.ANTHROPIC,), + models=(CHEAP_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_zero_priced_alias_still_logs_spend_at_the_tier_rate( self, proxy: ProxyClient, resources: ResourceManager, zero_priced_alias: ZeroPricedAlias ) -> None: @@ -546,6 +633,14 @@ class TestStrategyAliasPricing: class TestComplexityHeuristicScope: @pytest.mark.covers("reliability.routing.complexity_heuristic.scores_current_ask_only") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.ANTHROPIC,), + models=(CHEAP_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_trivial_ask_behind_keyword_heavy_system_prompt_stays_on_the_cheap_tier( self, proxy: ProxyClient, resources: ResourceManager, heuristic_split: HeuristicSplit ) -> None: @@ -573,6 +668,15 @@ class TestComplexityHeuristicScope: class TestSemanticAutoRouterResponses: @pytest.mark.covers("reliability.routing.semantic_auto_router.responses_input_routed") + @meta( + Subject( + domain=Domain.ROUTING, + route=Route.RESPONSES, + providers=(Provider.ANTHROPIC, Provider.OPENAI,), + models=(CHEAP_MODEL, EMBEDDING_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_responses_input_reaches_the_semantic_auto_router( self, proxy: ProxyClient, resources: ResourceManager, semantic_auto_router: SemanticAutoRouter ) -> None: @@ -617,6 +721,14 @@ class TestSemanticAutoRouterResponses: class TestAliasParamForwarding: @pytest.mark.covers("reliability.routing.tagged_marker.alias_connection_params_stay_with_tier") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.ANTHROPIC,), + models=(CHEAP_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_alias_api_key_never_overrides_the_tier_credential( self, proxy: ProxyClient, resources: ResourceManager, credentialed_alias: CredentialedAlias ) -> None: diff --git a/tests/e2e/router/test_complexity_router_e2e.py b/tests/e2e/router/test_complexity_router_e2e.py index e8508c963b8..c8f20102223 100644 --- a/tests/e2e/router/test_complexity_router_e2e.py +++ b/tests/e2e/router/test_complexity_router_e2e.py @@ -19,9 +19,12 @@ anthropic proves the classifier ran and openai proves it silently fell back - th exact failure before the fix. """ +from typing import Final + import pytest from complexity_router_client import ComplexityRouterClient +from e2e_metadata import Domain, Mode, Provider, Subject, meta from e2e_http import unwrap from models import ChatBody, ChatMessage @@ -33,9 +36,11 @@ LEXICALLY_SIMPLE_HARD_PROMPT = "Should I pay off my mortgage early or invest the # SIMPLE tier backend; served only when the classifier silently falls back to heuristic. # Spend logs may store the alias (gpt-5.5) or the provider-prefixed form depending on # how the deployment is registered (compose vs /model/new). -HEURISTIC_TIER_MODELS = frozenset({"openai/gpt-5.5", "gpt-5.5"}) +HEURISTIC_TIER_BACKEND: Final = "openai/gpt-5.5" +HEURISTIC_TIER_MODELS = frozenset({HEURISTIC_TIER_BACKEND, "gpt-5.5"}) # MEDIUM/COMPLEX/REASONING tier backend; served only when the LLM classifier runs. -LLM_TIER_MODELS = frozenset({"anthropic/claude-haiku-4-5", "claude-haiku-4-5"}) +LLM_TIER_BACKEND: Final = "anthropic/claude-haiku-4-5" +LLM_TIER_MODELS = frozenset({LLM_TIER_BACKEND, "claude-haiku-4-5"}) @pytest.mark.usefixtures("_ensure_complexity_smart_router") @@ -45,6 +50,14 @@ class TestComplexityRouterLlmClassifier: "(e.g. Is P equal to NP?); re-enable when classifier tier quality is fixed" ) @pytest.mark.covers("reliability.routing.complexity_llm_classifier.routes_by_llm_tier") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.OPENAI, Provider.ANTHROPIC), + models=(HEURISTIC_TIER_BACKEND, LLM_TIER_BACKEND), + mode=Mode.NONSTREAM, + ) + ) def test_llm_classifier_runs_and_routes_by_semantic_tier( self, client: ComplexityRouterClient, complexity_key: str ) -> None: diff --git a/tests/e2e/router/test_reliability_cache_e2e.py b/tests/e2e/router/test_reliability_cache_e2e.py index f7a2f2ffeb7..5452853a57e 100644 --- a/tests/e2e/router/test_reliability_cache_e2e.py +++ b/tests/e2e/router/test_reliability_cache_e2e.py @@ -18,6 +18,7 @@ from e2e_config import ( REQUEST_TIMEOUT, unique_marker, ) +from e2e_metadata import Domain, Mode, Provider, Subject, meta from lifecycle import ResourceManager from models import ChatBody, ChatMessage, ChatResponse, LiteLLMParamsBody from provider_edge import ProviderRequestObservation, observed_provider_edge @@ -25,6 +26,8 @@ from pydantic import BaseModel, JsonValue pytestmark = [pytest.mark.e2e, pytest.mark.replayable] +CACHE_MODEL: Final = "openai/gpt-5.6" + class _CacheChatBody(ChatBody): ttl: int = 600 @@ -38,6 +41,14 @@ class _CachedAnswer(BaseModel): class TestReliabilityCache: @pytest.mark.covers("reliability.cache.exact.returns_cached") + @meta( + Subject( + domain=Domain.CACHING, + providers=(Provider.OPENAI,), + models=(CACHE_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_exact_cache_returns_cached( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -57,7 +68,7 @@ class TestReliabilityCache: model_id: Final = client.proxy.create_model( model, LiteLLMParamsBody( - model="openai/gpt-5.6", + model=CACHE_MODEL, api_key="os.environ/OPENAI_API_KEY", api_base=f"{edge.api_base('openai')}/v1", ), diff --git a/tests/e2e/router/test_reliability_cancel_on_disconnect_e2e.py b/tests/e2e/router/test_reliability_cancel_on_disconnect_e2e.py index 06174e97d20..e39f1ddcb90 100644 --- a/tests/e2e/router/test_reliability_cancel_on_disconnect_e2e.py +++ b/tests/e2e/router/test_reliability_cancel_on_disconnect_e2e.py @@ -29,9 +29,11 @@ import pytest from complexity_router_client import ComplexityRouterClient from e2e_config import unique_marker from e2e_http import AbandonedRequest, StreamingResponse +from e2e_metadata import Domain, Mode, Provider, Subject, meta from lifecycle import ResourceManager from models import ChatMessage, ReliabilityChatBody, RouterSettingsOverride from reliability_support import ( + AZURE_MODEL, REPLICA_PROPAGATION_SECONDS, chat_override, create_azure_benched_on_first_failure_deployment, @@ -102,6 +104,14 @@ def _hang_up_mid_answer(client: ComplexityRouterClient, key: str, group: str) -> class TestReliabilityCancelOnDisconnect: @pytest.mark.covers("reliability.cooldown.client_disconnect.stays_healthy") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.AZURE,), + models=(AZURE_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_client_hanging_up_never_benches_the_deployment( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: diff --git a/tests/e2e/router/test_reliability_cooldowns_e2e.py b/tests/e2e/router/test_reliability_cooldowns_e2e.py index 2456bfb5f85..bdb02256976 100644 --- a/tests/e2e/router/test_reliability_cooldowns_e2e.py +++ b/tests/e2e/router/test_reliability_cooldowns_e2e.py @@ -71,10 +71,12 @@ import pytest from complexity_router_client import ComplexityRouterClient from e2e_config import CHEAP_OPENAI_MODEL, unique_marker from e2e_http import StreamingResponse +from e2e_metadata import Domain, Mode, Provider, Subject, meta from lifecycle import ResourceManager from models import KeyGenerateBody, RouterSettingsOverride from reliability_support import ( COOLDOWN_SECONDS, + REAL_MODEL, REPLICA_PROPAGATION_SECONDS, chat_override, create_always_5xx_deployment, @@ -212,6 +214,14 @@ def _assert_trips_then_recovers( class TestReliabilityCooldowns: @pytest.mark.covers("reliability.cooldown.5xx.trips_then_recovers") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.OPENAI,), + models=(REAL_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_5xx_trips_cooldown_then_recovers( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -230,6 +240,14 @@ class TestReliabilityCooldowns: _assert_trips_then_recovers(client, scoped_key, group, failing, backup, failure_status=500) @pytest.mark.covers("reliability.cooldown.sibling_replica.serves_backup_within_read_interval") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.OPENAI,), + models=(REAL_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_sibling_replica_serves_backup_within_redis_read_interval( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -272,6 +290,14 @@ class TestReliabilityCooldowns: ) @pytest.mark.covers("reliability.cooldown.429.trips_then_recovers") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.OPENAI,), + models=(CHEAP_OPENAI_MODEL, REAL_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_429_trips_cooldown_then_recovers( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -292,6 +318,14 @@ class TestReliabilityCooldowns: _assert_trips_then_recovers(client, scoped_key, group, failing, backup, failure_status=429) @pytest.mark.covers("reliability.cooldown.auth.trips_then_recovers") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.OPENAI,), + models=(REAL_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_auth_failure_trips_cooldown_then_recovers( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -304,6 +338,14 @@ class TestReliabilityCooldowns: _assert_trips_then_recovers(client, scoped_key, group, failing, backup, failure_status=401) @pytest.mark.covers("reliability.cooldown.timeout.trips_then_recovers") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.OPENAI,), + models=(REAL_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_timeout_trips_cooldown_then_recovers( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: diff --git a/tests/e2e/router/test_reliability_fallbacks_e2e.py b/tests/e2e/router/test_reliability_fallbacks_e2e.py index 54c11b163d0..61115e00d07 100644 --- a/tests/e2e/router/test_reliability_fallbacks_e2e.py +++ b/tests/e2e/router/test_reliability_fallbacks_e2e.py @@ -31,10 +31,14 @@ import pytest from complexity_router_client import ComplexityRouterClient from e2e_config import unique_marker from e2e_http import StreamingResponse +from e2e_metadata import Domain, Mode, Provider, Subject, meta from lifecycle import ResourceManager from models import RouterSettingsOverride from reliability_support import ( + AZURE_MODEL, CONTENT_POLICY_PROMPT, + REAL_MODEL, + SMALL_CONTEXT_MODEL, azure_prompt_filter_skipped, chat_override, completion_tokens_of, @@ -50,6 +54,8 @@ from reliability_support import ( pytestmark = pytest.mark.e2e +FALLBACK_MODEL: Final = "gpt-5.5" + def _assert_served_by_fallback(resp: StreamingResponse) -> None: assert resp.status_code == 200, f"expected 200 after fallback, got {resp.status_code}: {resp.body[:300]}" @@ -95,6 +101,14 @@ def _filter_verdict(resp: StreamingResponse) -> str: class TestReliabilityFallbacks: @pytest.mark.covers("reliability.fallback.5xx.routes_to_fallback") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.OPENAI,), + models=(FALLBACK_MODEL, REAL_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_5xx_routes_to_fallback( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -107,11 +121,19 @@ class TestReliabilityFallbacks: scoped_key, primary, f"say hi {unique_marker()}", - override=RouterSettingsOverride(fallbacks=[{primary: ["gpt-5.5"]}]), + override=RouterSettingsOverride(fallbacks=[{primary: [FALLBACK_MODEL]}]), ) _assert_served_by_fallback(resp) @pytest.mark.covers("reliability.fallback.timeout.routes_to_fallback") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.OPENAI,), + models=(FALLBACK_MODEL, REAL_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_timeout_routes_to_fallback( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -124,11 +146,19 @@ class TestReliabilityFallbacks: scoped_key, primary, f"say hi {unique_marker()}", - override=RouterSettingsOverride(fallbacks=[{primary: ["gpt-5.5"]}]), + override=RouterSettingsOverride(fallbacks=[{primary: [FALLBACK_MODEL]}]), ) _assert_served_by_fallback(resp) @pytest.mark.covers("reliability.fallback.context_window.routes_to_fallback") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.OPENAI,), + models=(FALLBACK_MODEL, SMALL_CONTEXT_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_context_window_routes_to_fallback( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -141,11 +171,19 @@ class TestReliabilityFallbacks: scoped_key, primary, oversized_prompt(unique_marker()), - override=RouterSettingsOverride(context_window_fallbacks=[{primary: ["gpt-5.5"]}]), + override=RouterSettingsOverride(context_window_fallbacks=[{primary: [FALLBACK_MODEL]}]), ) _assert_served_by_fallback(resp) @pytest.mark.covers("reliability.fallback.content_policy.routes_to_fallback") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.AZURE, Provider.OPENAI,), + models=(AZURE_MODEL, FALLBACK_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_content_policy_routes_to_fallback( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -167,7 +205,7 @@ class TestReliabilityFallbacks: scoped_key, primary, f"{CONTENT_POLICY_PROMPT} {unique_marker()}", - override=RouterSettingsOverride(content_policy_fallbacks=[{primary: ["gpt-5.5"]}]), + override=RouterSettingsOverride(content_policy_fallbacks=[{primary: [FALLBACK_MODEL]}]), ) ) _assert_served_by_fallback(resp) diff --git a/tests/e2e/router/test_reliability_memory_e2e.py b/tests/e2e/router/test_reliability_memory_e2e.py index 77d3a68cae5..9ae6edb3eb2 100644 --- a/tests/e2e/router/test_reliability_memory_e2e.py +++ b/tests/e2e/router/test_reliability_memory_e2e.py @@ -80,11 +80,12 @@ from e2e_config import ( PROXY_REPLICA_URLS, unique_marker, ) +from e2e_metadata import Domain, Mode, Provider, Subject, meta from lifecycle import ResourceManager from memory_readings import RssCapture, RssReading, WorkerKey, read_rss_everywhere from models import ChatMessage, RouterSettingsOverride, SpendLogRow from proxy_client import ProxyClient -from reliability_support import chat_override, create_never_benched_refusing_deployment +from reliability_support import REAL_MODEL, chat_override, create_never_benched_refusing_deployment pytestmark = [pytest.mark.e2e, pytest.mark.quiet_stack] @@ -222,6 +223,11 @@ def _stored_request_kb(proxy: ProxyClient, call: FailedCall) -> float: class TestReliabilityMemory: @pytest.mark.covers("reliability.perf.idle_memory.under_slo") + @meta( + Subject( + domain=Domain.DEPLOY_OPS, + ) + ) def test_workers_idle_under_rss_budget_before_traffic(self, idle_rss: RssCapture) -> None: assert not idle_rss.failures, ( f"{len(idle_rss.failures)} replica(s) gave no RSS reading when the session started, so their idle " @@ -242,6 +248,14 @@ class TestReliabilityMemory: ) @pytest.mark.covers("reliability.perf.memory.under_slo") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.OPENAI,), + models=(REAL_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_failing_requests_do_not_grow_rss_or_stored_request( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: diff --git a/tests/e2e/router/test_reliability_prompt_caching_e2e.py b/tests/e2e/router/test_reliability_prompt_caching_e2e.py index 667b398cd16..7bf23bf967f 100644 --- a/tests/e2e/router/test_reliability_prompt_caching_e2e.py +++ b/tests/e2e/router/test_reliability_prompt_caching_e2e.py @@ -23,9 +23,11 @@ import pytest from complexity_router_client import ComplexityRouterClient from e2e_config import unique_marker +from e2e_metadata import Capability, Domain, Mode, Provider, Subject, meta from lifecycle import ResourceManager from models import ChatMessage, LiteLLMParamsBody, ModelInfoBody, ModelNewBody from reliability_support import ( + CACHING_MODEL, REAL_KEY, REAL_MODEL, cached_system_turn, @@ -42,6 +44,15 @@ FOLLOW_UPS = 3 class TestReliabilityPromptCachingAffinity: @pytest.mark.covers("reliability.cache.prompt_caching_model_select.returns_cached") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.ANTHROPIC, Provider.OPENAI), + models=(CACHING_MODEL, REAL_MODEL), + capabilities=(Capability.PROMPT_CACHING,), + mode=Mode.NONSTREAM, + ) + ) def test_cached_conversation_stays_on_deployment_holding_its_cache( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: diff --git a/tests/e2e/router/test_reliability_retries_e2e.py b/tests/e2e/router/test_reliability_retries_e2e.py index a90efa52b5f..c6ae8c457e4 100644 --- a/tests/e2e/router/test_reliability_retries_e2e.py +++ b/tests/e2e/router/test_reliability_retries_e2e.py @@ -30,9 +30,12 @@ import pytest from complexity_router_client import ComplexityRouterClient from e2e_config import CHEAP_OPENAI_MODEL, unique_marker from e2e_http import StreamingResponse +from e2e_metadata import Domain, Mode, Provider, Subject, meta from lifecycle import ResourceManager from models import KeyGenerateBody, RouterSettingsOverride from reliability_support import ( + REAL_MODEL, + SMALL_CONTEXT_MODEL, chat_override, completion_tokens_of, content_of, @@ -84,6 +87,14 @@ def _retry_once(client: ComplexityRouterClient, key: str, group: str) -> Streami class TestReliabilityRetries: @pytest.mark.covers("reliability.retry.timeout.succeeds_within_retries") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.OPENAI,), + models=(REAL_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_timeout_on_first_deployment_succeeds_on_retry( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -96,6 +107,14 @@ class TestReliabilityRetries: _assert_served_after_retry(_retry_once(client, scoped_key, group)) @pytest.mark.covers("reliability.retry.5xx.succeeds_within_retries") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.OPENAI,), + models=(REAL_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_5xx_on_first_deployment_succeeds_on_retry( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -112,6 +131,14 @@ class TestReliabilityRetries: _assert_served_after_retry(_retry_once(client, scoped_key, group)) @pytest.mark.covers("reliability.retry.429.succeeds_within_retries") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.OPENAI,), + models=(CHEAP_OPENAI_MODEL, REAL_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_429_on_first_deployment_succeeds_on_retry( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -130,6 +157,14 @@ class TestReliabilityRetries: _assert_served_after_retry(_retry_once(client, scoped_key, group)) @pytest.mark.covers("reliability.retry.auth.succeeds_within_retries") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.OPENAI,), + models=(REAL_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_auth_failure_on_first_deployment_succeeds_on_retry( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -142,6 +177,14 @@ class TestReliabilityRetries: _assert_served_after_retry(_retry_once(client, scoped_key, group)) @pytest.mark.covers("reliability.retry.context_window.succeeds_within_retries") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.OPENAI,), + models=(REAL_MODEL, SMALL_CONTEXT_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_context_window_refusal_on_first_deployment_succeeds_on_retry( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: diff --git a/tests/e2e/router/test_reliability_routing_strategies_e2e.py b/tests/e2e/router/test_reliability_routing_strategies_e2e.py index 2abc2ee5f54..ee0334f989d 100644 --- a/tests/e2e/router/test_reliability_routing_strategies_e2e.py +++ b/tests/e2e/router/test_reliability_routing_strategies_e2e.py @@ -63,6 +63,7 @@ import pytest from complexity_router_client import ComplexityRouterClient from e2e_config import unique_marker from e2e_http import StreamChunk, StreamHead, StreamStep, StreamTruncation +from e2e_metadata import Domain, Mode, Provider, Subject, meta from lifecycle import ResourceManager from models import LiteLLMParamsBody, ModelInfoBody, ModelNewBody, RouterSettingsOverride, RoutingStrategy from reliability_support import REAL_KEY, REAL_MODEL, chat_override, model_id_of, open_chat_stream @@ -173,6 +174,14 @@ def _assert_shuffle_control_lands_on(client: ComplexityRouterClient, key: str, g class TestReliabilityRoutingStrategies: @pytest.mark.covers("reliability.routing.simple_shuffle.picks_healthy_deployment") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.OPENAI,), + models=(REAL_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_simple_shuffle_honors_weights( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -191,6 +200,14 @@ class TestReliabilityRoutingStrategies: ) @pytest.mark.covers("reliability.routing.cost_based.picks_lowest_cost") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.OPENAI,), + models=(REAL_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_cost_based_picks_cheapest_deployment( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -206,6 +223,14 @@ class TestReliabilityRoutingStrategies: _assert_shuffle_control_lands_on(client, scoped_key, group, pricey) @pytest.mark.covers("reliability.routing.usage_based.picks_under_tpm") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.OPENAI,), + models=(REAL_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_usage_based_picks_deployment_with_tpm_headroom( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -223,6 +248,14 @@ class TestReliabilityRoutingStrategies: "so latency-based has no signal to route on" ) @pytest.mark.covers("reliability.routing.latency_based.picks_lowest_latency") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.OPENAI,), + models=(REAL_MODEL,), + mode=Mode.NONSTREAM, + ) + ) def test_latency_based_routes_around_deployment_that_times_out( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -253,6 +286,13 @@ class TestReliabilityRoutingStrategies: "so least-busy has no signal to route on" ) @pytest.mark.covers("reliability.routing.least_busy.picks_lowest_traffic") + @meta( + Subject( + domain=Domain.ROUTING, + providers=(Provider.OPENAI,), + models=(REAL_MODEL,), + ) + ) def test_least_busy_avoids_deployment_with_request_in_flight( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: diff --git a/tests/e2e/router/test_reliability_timeouts_e2e.py b/tests/e2e/router/test_reliability_timeouts_e2e.py index f24d5139e66..926d8539d6d 100644 --- a/tests/e2e/router/test_reliability_timeouts_e2e.py +++ b/tests/e2e/router/test_reliability_timeouts_e2e.py @@ -13,14 +13,16 @@ import pytest from complexity_router_client import ComplexityRouterClient from e2e_config import unique_marker +from e2e_metadata import Domain, Mode, Provider, Subject, meta from lifecycle import ResourceManager -from reliability_support import chat_override, create_timeout_deployment +from reliability_support import REAL_MODEL, chat_override, create_timeout_deployment pytestmark = pytest.mark.e2e class TestReliabilityTimeouts: @pytest.mark.covers("reliability.timeout.request_timeout.exceeds_deadline") + @meta(Subject(domain=Domain.ROUTING, providers=(Provider.OPENAI,), models=(REAL_MODEL,), mode=Mode.NONSTREAM)) def test_request_timeout_exceeds_deadline( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> None: @@ -35,6 +37,7 @@ class TestReliabilityTimeouts: assert "timeout" in resp.body.lower(), f"the 408 body should name the timeout, got: {resp.body[:300]}" @pytest.mark.covers("reliability.timeout.stream_timeout.exceeds_deadline") + @meta(Subject(domain=Domain.ROUTING, providers=(Provider.OPENAI,), models=(REAL_MODEL,), mode=Mode.STREAM)) def test_stream_timeout_exceeds_deadline( self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str ) -> 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/e2e/ui/helpers/roundTrip.ts b/tests/e2e/ui/helpers/roundTrip.ts index 91e48b7087c..4dc6d17a369 100644 --- a/tests/e2e/ui/helpers/roundTrip.ts +++ b/tests/e2e/ui/helpers/roundTrip.ts @@ -7,18 +7,18 @@ import { masterKey } from "./traffic"; * `action` is a callback so the listener is armed before the click; awaiting the * click first lets the request go by, and the test then hangs until timeout. */ -export async function captureRequestBody( +export async function captureRequestBody>( page: Page, match: { method: string; urlIncludes: string }, action: () => Promise, -): Promise> { +): Promise { const pending = page.waitForRequest( (req) => req.method() === match.method && req.url().includes(match.urlIncludes), ); await action(); const request = await pending; - return JSON.parse(request.postData() ?? "{}") as Record; + return JSON.parse(request.postData() ?? "{}") as T; } /** Reads an endpoint as the master key, so a failure is bad data and not an expired UI token. */ diff --git a/tests/e2e/ui/tests/integrationCritical/anthropicFederationCredential.spec.ts b/tests/e2e/ui/tests/integrationCritical/anthropicFederationCredential.spec.ts new file mode 100644 index 00000000000..c00c61b9e6c --- /dev/null +++ b/tests/e2e/ui/tests/integrationCritical/anthropicFederationCredential.spec.ts @@ -0,0 +1,799 @@ +import { + test, + expect, + type APIRequestContext, + type Locator, + type Page as PlaywrightPage, +} from "@playwright/test"; +import { Page } from "../../fixtures/pages"; +import { navigateToPage } from "../../helpers/navigation"; +import { captureRequestBody, readBack } from "../../helpers/roundTrip"; +import { uniqueSuffix } from "../../helpers/traffic"; +import { + logInThroughLoginPage, + setInvitedUserPassword, +} from "../../helpers/userOnboarding"; + +const master = process.env.LITELLM_MASTER_KEY ?? "sk-integration-master"; +const headers = { Authorization: `Bearer ${master}` }; + +const ANTHROPIC_LABEL = "Anthropic"; +const OPENAI_LABEL = "OpenAI"; +const FEDERATION_BADGE = "Workload identity federation"; +const FEDERATION_BUTTON = "Use workload identity federation"; +const FEDERATION_HELP = + "Workload identity federation is saved as a credential, then attached to this model."; +const ADDED_TOAST = "Credential added successfully"; +const UPDATED_TOAST = "Credential updated successfully"; +const ALLOWLIST_VARIABLE = "LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS"; +const NO_IDS_MESSAGE = + "Enter at least one of these ids, or pick an identity source that stores a token. The proxy rejects a credential with no values"; +const RAW_TOKEN_MESSAGE = + "Enter an oidc/ secret reference such as oidc/env/VAR_NAME. Raw tokens and oidc/env_path/ references are not accepted"; +const TTL_MESSAGE = "Enter a whole number of seconds from 1 to 3600"; + +type CredentialValues = Record; + +interface StoredCredential { + credential_name: string; + credential_values: CredentialValues; + credential_info: { custom_llm_provider: string }; +} + +interface Deployment { + model_name: string; + litellm_params: Record; + model_info: { id: string }; +} + +interface CredentialWrite { + credential_name: string; + credential_values: CredentialValues; + credential_info: { custom_llm_provider: string }; + credential_values_to_delete?: string[]; +} + +interface ModelWrite { + model_name: string; + litellm_params: Record; +} + +interface ModelCreated { + model_id: string; +} + +const masked = (value: string): string => + value.length <= 4 ? "*****" : `${value.slice(0, 2)}****${value.slice(-2)}`; + +const exactText = (text: string): RegExp => + new RegExp(`^${text.replace(/[.*+?^${}()|[\]\\]/g, "\\$&")}$`); + +async function logInAsAdmin(page: PlaywrightPage): Promise { + await page.goto("/ui/login"); + await page.getByPlaceholder("Enter your username").fill("admin"); + await page.getByPlaceholder("Enter your password").fill(master); + await page.getByRole("button", { name: "Login", exact: true }).click(); + await expect(page).toHaveURL( + (url) => url.pathname.startsWith("/ui") && !url.pathname.includes("login"), + ); +} + +async function openModelsTab(page: PlaywrightPage, tab: string): Promise { + await navigateToPage(page, Page.Models); + await page.getByRole("tab", { name: tab, exact: true }).click(); +} + +async function pickOption( + page: PlaywrightPage, + trigger: Locator, + label: string, +): Promise { + await trigger.click(); + await page + .getByRole("option") + .filter({ hasText: exactText(label) }) + .first() + .click(); +} + +async function pickProvider( + page: PlaywrightPage, + scope: Locator, + provider: string, +): Promise { + await pickOption(page, scope.getByPlaceholder("Select a provider"), provider); +} + +async function pickAuthMethod( + page: PlaywrightPage, + dialog: Locator, + label: string, +): Promise { + await pickOption(page, dialog.locator("#anthropic_auth_method"), label); +} + +async function pickIdentitySource( + page: PlaywrightPage, + dialog: Locator, + label: string, +): Promise { + await pickOption( + page, + dialog.locator("#anthropic_federation_identity_source"), + label, + ); +} + +async function fillFields( + dialog: Locator, + values: Record, +): Promise { + for (const [key, value] of Object.entries(values)) { + await dialog.locator(`#${key}`).fill(value); + } +} + +async function openAddCredentialDialog( + page: PlaywrightPage, + name: string, +): Promise { + await page + .getByRole("button", { name: "Add Credential", exact: true }) + .click(); + const dialog = page.getByRole("dialog", { name: "Add New Credential" }); + await expect(dialog).toBeVisible(); + await dialog + .getByPlaceholder("Enter a friendly name for these credentials") + .fill(name); + await pickProvider(page, dialog, ANTHROPIC_LABEL); + await pickAuthMethod(page, dialog, FEDERATION_BADGE); + return dialog; +} + +async function openEditDialog( + page: PlaywrightPage, + name: string, +): Promise { + await expect(page.getByText(UPDATED_TOAST)).toHaveCount(0); + const row = page.locator("tr", { hasText: name }); + await expect(row).toBeVisible({ timeout: 15_000 }); + await row.getByTestId(`credential-actions-${name}`).click(); + await page.getByTestId("credential-action-edit").click(); + const dialog = page.getByRole("dialog", { name: "Edit Credential" }); + await expect(dialog).toBeVisible(); + return dialog; +} + +const submitCredential = ( + page: PlaywrightPage, + dialog: Locator, + method: "POST" | "PATCH", + urlIncludes: string, + button: string, +): Promise => + captureRequestBody(page, { method, urlIncludes }, () => + dialog.getByRole("button", { name: button, exact: true }).click(), + ); + +function countRequests( + page: PlaywrightPage, + method: string, + pathSuffix: string, +): () => number { + let seen = 0; + page.on("request", (request) => { + if ( + request.method() === method && + new URL(request.url()).pathname.endsWith(pathSuffix) + ) { + seen += 1; + } + }); + return () => seen; +} + +async function expectBlocked( + page: PlaywrightPage, + dialog: Locator, + button: string, + message: string, + postsBefore: number, + posts: () => number, +): Promise { + await dialog.getByRole("button", { name: button, exact: true }).click(); + await expect(dialog.getByText(message)).toBeVisible(); + await expect(dialog).toBeVisible(); + expect(posts(), "a refused form must send nothing").toBe(postsBefore); +} + +async function createCredential( + request: APIRequestContext, + name: string, + values: CredentialValues, + provider = "anthropic", +): Promise { + const response = await request.post("/credentials", { + headers, + data: { + credential_name: name, + credential_values: values, + credential_info: { custom_llm_provider: provider }, + }, + }); + expect(response.status(), await response.text()).toBe(200); +} + +async function deleteCredentials( + request: APIRequestContext, + names: readonly string[], +): Promise { + for (const name of names) { + await request.delete(`/credentials/${name}`, { headers }); + } +} + +const sameValues = (left: CredentialValues, right: CredentialValues): boolean => + JSON.stringify(Object.entries(left).sort()) === + JSON.stringify(Object.entries(right).sort()); + +async function readStoredCredential( + request: APIRequestContext, + name: string, +): Promise { + const response = await request.get(`/credentials/by_name/${name}`, { + headers, + }); + return response.status() === 200 + ? ((await response.json()) as StoredCredential) + : null; +} + +async function expectStored( + request: APIRequestContext, + name: string, + expected: CredentialValues, +): Promise { + let last: StoredCredential | null = null; + await expect + .poll( + async () => { + last = await readStoredCredential(request, name); + return last !== null && sameValues(last.credential_values, expected); + }, + { + timeout: 70_000, + message: `the proxy must serve ${name} as ${JSON.stringify(expected)}; last read ${JSON.stringify(last)}`, + }, + ) + .toBe(true); + if (last === null) { + throw new Error(`${name} was never read back`); + } + return last; +} + +async function postAsMaster>( + request: APIRequestContext, + route: string, + data: Record, +): Promise { + const response = await request.post(route, { headers, data }); + expect(response.status(), `POST ${route}: ${await response.text()}`).toBe( + 200, + ); + return (await response.json()) as T; +} + +test("the Add Credential dialog stores each federation identity source with only the values it needs and refuses the invalid ones", async ({ + page, + request, +}) => { + const suffix = uniqueSuffix(); + const names = { + tokenFile: `int-wif-file-${suffix}`, + secretReference: `int-wif-ref-${suffix}`, + internalIssuer: `int-wif-issuer-${suffix}`, + environment: `int-wif-env-${suffix}`, + }; + const ids = { + anthropic_federation_rule_id: `fdrl_${suffix}`, + anthropic_organization_id: `org-${suffix}`, + anthropic_service_account_id: `svac_${suffix}`, + anthropic_federation_workspace_id: `wrkspc_${suffix}`, + }; + const tokenFile = `/var/run/secrets/anthropic/${suffix}/token`; + const posts = countRequests(page, "POST", "/credentials"); + try { + await logInAsAdmin(page); + await openModelsTab(page, "LLM Credentials"); + + const fileDialog = await openAddCredentialDialog(page, names.tokenFile); + await pickIdentitySource(page, fileDialog, "Identity token file"); + await fillFields(fileDialog, { + ...ids, + anthropic_identity_token_file: tokenFile, + }); + const created = await submitCredential( + page, + fileDialog, + "POST", + "/credentials", + "Add Credential", + ); + expect(created).toEqual({ + credential_name: names.tokenFile, + credential_values: { ...ids, anthropic_identity_token_file: tokenFile }, + credential_info: { custom_llm_provider: ANTHROPIC_LABEL }, + }); + await expect(page.getByText(ADDED_TOAST)).toBeVisible(); + const row = page.locator("tr", { hasText: names.tokenFile }); + await expect(row).toBeVisible({ timeout: 15_000 }); + await expect(row).toContainText(FEDERATION_BADGE); + await expectStored(request, names.tokenFile, { + ...ids, + anthropic_identity_token_file: masked(tokenFile), + }); + + const referenceDialog = await openAddCredentialDialog( + page, + names.secretReference, + ); + await pickIdentitySource( + page, + referenceDialog, + "Identity token secret reference", + ); + await fillFields(referenceDialog, { + anthropic_federation_rule_id: ids.anthropic_federation_rule_id, + anthropic_identity_token: `eyJhbGciOiJSUzI1NiJ9.${suffix}.signature`, + }); + await expectBlocked( + page, + referenceDialog, + "Add Credential", + RAW_TOKEN_MESSAGE, + 1, + posts, + ); + const reference = `oidc/env/ANTHROPIC_IDENTITY_${suffix.replace(/\W/g, "_")}`; + await fillFields(referenceDialog, { anthropic_identity_token: reference }); + const referenced = await submitCredential( + page, + referenceDialog, + "POST", + "/credentials", + "Add Credential", + ); + expect(referenced.credential_values).toEqual({ + anthropic_federation_rule_id: ids.anthropic_federation_rule_id, + anthropic_identity_token: reference, + }); + await expect( + page.locator("tr", { hasText: names.secretReference }), + ).toBeVisible({ timeout: 15_000 }); + + const issuerDialog = await openAddCredentialDialog( + page, + names.internalIssuer, + ); + await pickIdentitySource( + page, + issuerDialog, + "Token signed by LiteLLM (internal issuer)", + ); + const issuer = { + anthropic_issuer_url: `https://issuer-${suffix}.example`, + anthropic_issuer_subject: `proxy-${suffix}`, + anthropic_issuer_audience: `anthropic-${suffix}`, + anthropic_issuer_signing_key_ref: + "os.environ/ANTHROPIC_ISSUER_SIGNING_KEY", + }; + await fillFields(issuerDialog, { + anthropic_federation_rule_id: ids.anthropic_federation_rule_id, + ...issuer, + anthropic_issuer_ttl_seconds: "7200", + }); + await expectBlocked( + page, + issuerDialog, + "Add Credential", + TTL_MESSAGE, + 2, + posts, + ); + await fillFields(issuerDialog, { anthropic_issuer_ttl_seconds: "900" }); + const issued = await submitCredential( + page, + issuerDialog, + "POST", + "/credentials", + "Add Credential", + ); + expect(issued.credential_values).toEqual({ + anthropic_federation_rule_id: ids.anthropic_federation_rule_id, + ...issuer, + anthropic_issuer_ttl_seconds: 900, + anthropic_identity_source: "internal_issuer", + }); + await expect( + page.locator("tr", { hasText: names.internalIssuer }), + ).toBeVisible({ timeout: 15_000 }); + + const environmentDialog = await openAddCredentialDialog( + page, + names.environment, + ); + await pickIdentitySource( + page, + environmentDialog, + "Proxy environment variables", + ); + await expectBlocked( + page, + environmentDialog, + "Add Credential", + NO_IDS_MESSAGE, + 3, + posts, + ); + await fillFields(environmentDialog, { + anthropic_organization_id: ids.anthropic_organization_id, + }); + const fromEnvironment = await submitCredential( + page, + environmentDialog, + "POST", + "/credentials", + "Add Credential", + ); + expect(fromEnvironment.credential_values).toEqual({ + anthropic_organization_id: ids.anthropic_organization_id, + }); + await expect( + page.locator("tr", { hasText: names.environment }), + ).toBeVisible({ timeout: 15_000 }); + expect(posts()).toBe(4); + } finally { + await deleteCredentials(request, Object.values(names)); + } +}); + +test("the Edit Credential dialog sends only the changed values and deletes the keys the new choice leaves behind", async ({ + page, + request, +}) => { + const suffix = uniqueSuffix(); + const names = { + fromApiKey: `int-wif-edit-key-${suffix}`, + clearWorkspace: `int-wif-edit-clear-${suffix}`, + switchSource: `int-wif-edit-source-${suffix}`, + switchProvider: `int-wif-edit-provider-${suffix}`, + }; + const ruleId = `fdrl_${suffix}`; + const tokenFile = `/var/run/secrets/anthropic/${suffix}/token`; + try { + await createCredential(request, names.fromApiKey, { + api_key: `sk-ant-${suffix}`, + }); + await createCredential(request, names.clearWorkspace, { + anthropic_federation_rule_id: ruleId, + anthropic_federation_workspace_id: `wrkspc_${suffix}`, + anthropic_identity_token_file: tokenFile, + }); + await createCredential(request, names.switchSource, { + anthropic_federation_rule_id: ruleId, + anthropic_identity_token_file: tokenFile, + }); + await createCredential(request, names.switchProvider, { + anthropic_federation_rule_id: ruleId, + anthropic_organization_id: `org-${suffix}`, + anthropic_identity_token_file: tokenFile, + }); + + await logInAsAdmin(page); + await openModelsTab(page, "LLM Credentials"); + + const keyDialog = await openEditDialog(page, names.fromApiKey); + await pickAuthMethod(page, keyDialog, FEDERATION_BADGE); + await fillFields(keyDialog, { + anthropic_federation_rule_id: ruleId, + anthropic_identity_token_file: tokenFile, + }); + const federated = await submitCredential( + page, + keyDialog, + "PATCH", + `/credentials/${names.fromApiKey}`, + "Update Credential", + ); + expect(federated).toEqual({ + credential_name: names.fromApiKey, + credential_values: { + anthropic_federation_rule_id: ruleId, + anthropic_identity_token_file: tokenFile, + }, + credential_info: { custom_llm_provider: "anthropic" }, + credential_values_to_delete: ["api_key"], + }); + await expect(page.getByText(UPDATED_TOAST)).toBeVisible(); + await expect( + page.locator("tr", { hasText: names.fromApiKey }), + ).toContainText(FEDERATION_BADGE); + await expectStored(request, names.fromApiKey, { + anthropic_federation_rule_id: ruleId, + anthropic_identity_token_file: masked(tokenFile), + }); + + const clearDialog = await openEditDialog(page, names.clearWorkspace); + await fillFields(clearDialog, { anthropic_federation_workspace_id: "" }); + const cleared = await submitCredential( + page, + clearDialog, + "PATCH", + `/credentials/${names.clearWorkspace}`, + "Update Credential", + ); + expect(cleared).toEqual({ + credential_name: names.clearWorkspace, + credential_values: {}, + credential_info: { custom_llm_provider: "anthropic" }, + credential_values_to_delete: ["anthropic_federation_workspace_id"], + }); + await expect(page.getByText(UPDATED_TOAST)).toBeVisible(); + await expectStored(request, names.clearWorkspace, { + anthropic_federation_rule_id: ruleId, + anthropic_identity_token_file: masked(tokenFile), + }); + + const sourceDialog = await openEditDialog(page, names.switchSource); + await pickIdentitySource( + page, + sourceDialog, + "Identity token secret reference", + ); + const reference = `oidc/env/ANTHROPIC_IDENTITY_${suffix.replace(/\W/g, "_")}`; + await fillFields(sourceDialog, { anthropic_identity_token: reference }); + const switched = await submitCredential( + page, + sourceDialog, + "PATCH", + `/credentials/${names.switchSource}`, + "Update Credential", + ); + expect(switched).toEqual({ + credential_name: names.switchSource, + credential_values: { anthropic_identity_token: reference }, + credential_info: { custom_llm_provider: "anthropic" }, + credential_values_to_delete: ["anthropic_identity_token_file"], + }); + await expect(page.getByText(UPDATED_TOAST)).toBeVisible(); + await expectStored(request, names.switchSource, { + anthropic_federation_rule_id: ruleId, + anthropic_identity_token: masked(reference), + }); + + const providerDialog = await openEditDialog(page, names.switchProvider); + await pickProvider(page, providerDialog, OPENAI_LABEL); + const openai = { + api_key: `sk-openai-${suffix}`, + api_base: `https://openai-${suffix}.example/v1`, + }; + await fillFields(providerDialog, openai); + const moved = await submitCredential( + page, + providerDialog, + "PATCH", + `/credentials/${names.switchProvider}`, + "Update Credential", + ); + expect(moved.credential_values).toEqual(openai); + expect(moved.credential_info).toEqual({ + custom_llm_provider: OPENAI_LABEL, + }); + expect([...(moved.credential_values_to_delete ?? [])].sort()).toEqual([ + "anthropic_federation_rule_id", + "anthropic_identity_token_file", + "anthropic_organization_id", + ]); + await expect(page.getByText(UPDATED_TOAST)).toBeVisible(); + const onOpenai = await expectStored(request, names.switchProvider, { + api_key: masked(openai.api_key), + api_base: openai.api_base, + }); + expect(onOpenai.credential_info.custom_llm_provider).toBe(OPENAI_LABEL); + await expect( + page.locator("tr", { hasText: names.switchProvider }), + ).not.toContainText(FEDERATION_BADGE); + } finally { + await deleteCredentials(request, Object.values(names)); + } +}); + +test("a proxy admin saves a federation credential from the Add Model tab and the model is created against it", async ({ + page, + request, +}) => { + test.slow(); + const suffix = uniqueSuffix(); + const credentialName = `int-wif-model-${suffix}`; + const modelName = `int-wif-model-${suffix}`; + const ids = { + anthropic_federation_rule_id: `fdrl_${suffix}`, + anthropic_organization_id: `org-${suffix}`, + }; + const tokenFile = `/opt/int-wif-${suffix}/token`; + let deploymentId = ""; + try { + await logInAsAdmin(page); + await openModelsTab(page, "Add Model"); + await pickProvider(page, page.locator("body"), ANTHROPIC_LABEL); + await expect(page.getByText(FEDERATION_HELP)).toBeVisible(); + await page + .getByRole("button", { name: FEDERATION_BUTTON, exact: true }) + .click(); + + const dialog = page.getByRole("dialog", { name: "Add New Credential" }); + await expect(dialog).toBeVisible(); + const provider = dialog.getByPlaceholder("Select a provider"); + await expect(provider).toHaveValue(ANTHROPIC_LABEL); + await expect(provider).toBeDisabled(); + await expect(dialog.locator("#anthropic_auth_method")).toContainText( + FEDERATION_BADGE, + ); + await dialog + .getByPlaceholder("Enter a friendly name for these credentials") + .fill(credentialName); + await fillFields(dialog, { + ...ids, + anthropic_identity_token_file: tokenFile, + }); + const created = await submitCredential( + page, + dialog, + "POST", + "/credentials", + "Add Credential", + ); + expect(created).toEqual({ + credential_name: credentialName, + credential_values: { ...ids, anthropic_identity_token_file: tokenFile }, + credential_info: { custom_llm_provider: ANTHROPIC_LABEL }, + }); + await expect(page.getByText(ADDED_TOAST)).toBeVisible(); + await expect(dialog).toBeHidden(); + await expect(page.locator("#litellm_credential_name")).toHaveValue( + credentialName, + ); + await expect(page.locator("#api_key")).toHaveCount(0); + + await pickOption( + page, + page.getByRole("combobox", { name: "Select models" }), + "Custom Model Name (Enter below)", + ); + await page.keyboard.press("Escape"); + await page.getByPlaceholder("Enter custom model name").fill(modelName); + + await expectStored(request, credentialName, { + ...ids, + anthropic_identity_token_file: masked(tokenFile), + }); + + const probe = await captureRequestBody( + page, + { method: "POST", urlIncludes: "/health/test_connection" }, + () => page.getByTestId("test-connect-btn").click(), + ); + expect(probe.litellm_params.litellm_credential_name).toBe(credentialName); + expect(probe.litellm_params.custom_llm_provider).toBe("anthropic"); + expect("api_key" in probe.litellm_params).toBe(false); + const results = page.getByRole("dialog", { + name: "Connection Test Results", + }); + await expect(results.getByTestId("connection-failure-msg")).toBeVisible({ + timeout: 60_000, + }); + await expect(results.getByText(ALLOWLIST_VARIABLE)).toBeVisible(); + await expect(results.getByText(tokenFile)).toBeVisible(); + await page.keyboard.press("Escape"); + await expect(results).toBeHidden(); + + const creation = page.waitForResponse( + (response) => + response.request().method() === "POST" && + response.url().includes("/model/new"), + ); + const added = await captureRequestBody( + page, + { method: "POST", urlIncludes: "/model/new" }, + () => page.getByTestId("add-model-btn").click(), + ); + const creationResponse = await creation; + expect( + creationResponse.status(), + `POST /model/new: ${await creationResponse.text()}`, + ).toBe(200); + deploymentId = ((await creationResponse.json()) as ModelCreated).model_id; + expect(deploymentId).not.toBe(""); + expect(added.model_name).toBe(modelName); + expect(added.litellm_params.litellm_credential_name).toBe(credentialName); + expect(added.litellm_params.custom_llm_provider).toBe("anthropic"); + expect("api_key" in added.litellm_params).toBe(false); + + const findDeployment = async (): Promise => { + const info = await readBack<{ data: Deployment[] }>(page, "/model/info"); + return info.data.find( + (deployment) => deployment.model_name === modelName, + ); + }; + await expect + .poll(async () => (await findDeployment())?.model_info.id ?? "", { + timeout: 70_000, + message: "the deployment never appeared in /model/info", + }) + .not.toBe(""); + const deployment = await findDeployment(); + expect(deployment?.model_info.id).toBe(deploymentId); + expect(deployment?.litellm_params.litellm_credential_name).toBe( + credentialName, + ); + expect(deployment?.litellm_params.custom_llm_provider).toBe("anthropic"); + expect("api_key" in (deployment?.litellm_params ?? {})).toBe(false); + } finally { + if (deploymentId) { + await postAsMaster(request, "/model/delete", { id: deploymentId }); + } + await deleteCredentials(request, [credentialName]); + } +}); + +test("a team admin who is not a proxy admin gets no federation shortcut on the Add Model tab", async ({ + page, + request, +}) => { + const suffix = uniqueSuffix(); + const userId = `int-wif-team-admin-${suffix}`; + const email = `${userId}@integration.example`; + const password = `Int-Wif-${suffix}!`; + const teamAlias = `int-wif-team-${suffix}`; + let teamId = ""; + try { + await postAsMaster(request, "/user/new", { + user_id: userId, + user_email: email, + user_role: "internal_user", + auto_create_key: false, + }); + await setInvitedUserPassword(request, userId, password); + teamId = ( + await postAsMaster<{ team_id: string }>(request, "/team/new", { + team_alias: teamAlias, + members_with_roles: [{ role: "admin", user_id: userId }], + }) + ).team_id; + + await logInThroughLoginPage(page, email, password); + await openModelsTab(page, "Add Model"); + await expect(page.getByText("Team Selection Required")).toBeVisible(); + await page.getByPlaceholder("Search or select a team").click(); + await page + .getByRole("option") + .filter({ hasText: teamAlias }) + .first() + .click(); + await pickProvider(page, page.locator("body"), ANTHROPIC_LABEL); + await expect(page.getByPlaceholder("Select a provider")).toHaveValue( + ANTHROPIC_LABEL, + ); + await expect(page.locator("#api_key")).toBeVisible(); + await expect( + page.getByRole("button", { name: FEDERATION_BUTTON, exact: true }), + ).toHaveCount(0); + await expect(page.getByText(FEDERATION_HELP)).toHaveCount(0); + } finally { + if (teamId) { + await postAsMaster(request, "/team/delete", { team_ids: [teamId] }); + } + await postAsMaster(request, "/user/delete", { user_ids: [userId] }); + } +}); diff --git a/tests/e2e/ui/tests/integrationCritical/expected.json b/tests/e2e/ui/tests/integrationCritical/expected.json index 81d3be247fd..117df8d4892 100644 --- a/tests/e2e/ui/tests/integrationCritical/expected.json +++ b/tests/e2e/ui/tests/integrationCritical/expected.json @@ -10,5 +10,9 @@ "tests/e2e/ui/tests/integrationCritical/costOptimizationModelGroups.spec.ts::cache leakage by model merges a deployment's resolved and requested model names into its model group", "tests/e2e/ui/tests/integrationCritical/logsDrawerCredentialCanary.spec.ts::the Logs drawer renders the stored request without the deployment api_key", "tests/e2e/ui/tests/integrationCritical/toolPoliciesUserColumn.spec.ts::the Tool Policies page names the user behind the key that discovered a tool", - "tests/e2e/ui/tests/integrationCritical/teamMetadataEmptyKey.spec.ts::the team settings form skips a metadata row with an empty key and saving drops the key" + "tests/e2e/ui/tests/integrationCritical/teamMetadataEmptyKey.spec.ts::the team settings form skips a metadata row with an empty key and saving drops the key", + "tests/e2e/ui/tests/integrationCritical/anthropicFederationCredential.spec.ts::the Add Credential dialog stores each federation identity source with only the values it needs and refuses the invalid ones", + "tests/e2e/ui/tests/integrationCritical/anthropicFederationCredential.spec.ts::the Edit Credential dialog sends only the changed values and deletes the keys the new choice leaves behind", + "tests/e2e/ui/tests/integrationCritical/anthropicFederationCredential.spec.ts::a proxy admin saves a federation credential from the Add Model tab and the model is created against it", + "tests/e2e/ui/tests/integrationCritical/anthropicFederationCredential.spec.ts::a team admin who is not a proxy admin gets no federation shortcut on the Add Model tab" ] diff --git a/tests/integration/_support/process.py b/tests/integration/_support/process.py index 8e9f689ec2d..12780540380 100644 --- a/tests/integration/_support/process.py +++ b/tests/integration/_support/process.py @@ -119,7 +119,7 @@ def _stop(process: subprocess.Popen[bytes]) -> None: _PORT_ATTEMPTS: Final = 3 -_BIND_COLLISION: Final = os.strerror(errno.EADDRINUSE) +_BIND_COLLISION: Final = os.strerror(errno.EADDRINUSE).lower() def _free_port() -> int: @@ -151,7 +151,7 @@ def _launch(command: tuple[str, ...], root: Path, environment: Mapping[str, str] def _lost_port_race(exit_code: int | None, log: Path) -> bool: - return exit_code is not None and _BIND_COLLISION in log.read_text() + return exit_code is not None and _BIND_COLLISION in log.read_text().lower() def _wait_until_ready(launch: _Launch) -> 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/authorization/test_team_scoped_models.py b/tests/integration/authorization/test_team_scoped_models.py index aea6b9e4ba4..1323c1ee004 100644 --- a/tests/integration/authorization/test_team_scoped_models.py +++ b/tests/integration/authorization/test_team_scoped_models.py @@ -15,10 +15,16 @@ def upstream(gateway: Gateway) -> Iterator[httpx.Client]: yield client -def _observed_models(upstream: httpx.Client) -> list[JsonValue]: +def _observed_requests(upstream: httpx.Client) -> list[JsonValue]: observed: Final = upstream.get("/__observations") observed.raise_for_status() - return [request["body"]["model"] for request in observed.json()["requests"]] + requests: Final = object_value(observed.json())["requests"] + assert isinstance(requests, list) + return requests + + +def _calls_to(observed: list[JsonValue], provider_model: str) -> int: + return sum(object_value(object_value(request)["body"]).get("model") == provider_model for request in observed) def _chat(gateway: Gateway, model: str, key: str) -> httpx.Response: @@ -62,7 +68,8 @@ def test_a_team_model_is_listed_and_served_only_for_keys_of_its_team(gateway: Ga assert served.status_code == 200, served.text refused: Final = _chat(gateway, model, other_key) assert refused.status_code == 400, refused.text - assert _observed_models(upstream) == [provider_model] + observed: Final = _observed_requests(upstream) + assert _calls_to(observed, provider_model) == 1, observed def _v2_team_public_names(gateway: Gateway, key: str, model: str) -> list[JsonValue]: @@ -98,8 +105,10 @@ def test_team_model_alias_routes_a_team_key_to_its_target( response: Final = _chat(gateway, alias, key) assert response.status_code == 200, response.text assert string_value(response.json()["model"]) == alias - assert _observed_models(upstream) == [provider_model] + observed: Final = _observed_requests(upstream) + assert _calls_to(observed, provider_model) == 1, observed unaliased: Final = _chat(gateway, f"alias-{uuid.uuid4().hex}", key) assert unaliased.status_code == 403, unaliased.text assert unaliased.json()["error"]["type"] == "key_model_access_denied" - assert _observed_models(upstream) == [] + after_refusal: Final = _observed_requests(upstream) + assert _calls_to(after_refusal, provider_model) == 0, after_refusal diff --git a/tests/integration/caching/test_cache_max_messages.py b/tests/integration/caching/test_cache_max_messages.py new file mode 100644 index 00000000000..89401b7c4ec --- /dev/null +++ b/tests/integration/caching/test_cache_max_messages.py @@ -0,0 +1,144 @@ +import json +import os +import uuid +from collections.abc import Callable +from typing import Final + +import pytest +from integration._support.anthropic_sse import message_json +from integration._support.client import Gateway, eventually, object_value, string_value +from integration._support.openai_wire import chat_reply, responses_reply +from integration._support.provider import SharedProvider +from integration._support.wire import Reply +from pydantic import JsonValue +from redis import Redis + +_Turns = Callable[[str], tuple[list[JsonValue], list[JsonValue]]] + + +def _claude_code_turns(task: str) -> tuple[list[JsonValue], list[JsonValue]]: + """Turns 1 and 3 of a Claude Code session on /v1/messages: 1 and 5 messages""" + first: Final[list[JsonValue]] = [{"role": "user", "content": task}] + third: Final[list[JsonValue]] = [ + *first, + { + "role": "assistant", + "content": [{"type": "tool_use", "id": "toolu_1", "name": "Bash", "input": {"command": "ls"}}], + }, + { + "role": "user", + "content": [{"type": "tool_result", "tool_use_id": "toolu_1", "content": "calc.py test_calc.py"}], + }, + { + "role": "assistant", + "content": [{"type": "tool_use", "id": "toolu_2", "name": "Read", "input": {"file_path": "calc.py"}}], + }, + { + "role": "user", + "content": [{"type": "tool_result", "tool_use_id": "toolu_2", "content": "def add(a, b): return a - b"}], + }, + ] + return first, third + + +def _agent_turns(task: str) -> tuple[list[JsonValue], list[JsonValue]]: + """Turns 1 and 3 of an OpenAI tool loop on /v1/chat/completions: 2 and 6 messages""" + + def call(call_id: str, path: str) -> list[JsonValue]: + function: Final[JsonValue] = {"name": "write_file", "arguments": json.dumps({"path": path})} + return [ + { + "role": "assistant", + "content": None, + "tool_calls": [{"id": call_id, "type": "function", "function": function}], + }, + {"role": "tool", "tool_call_id": call_id, "content": f"wrote {path}"}, + ] + + first: Final[list[JsonValue]] = [ + {"role": "system", "content": "You are a coding agent"}, + {"role": "user", "content": task}, + ] + return first, [*first, *call("call_1", "a.yaml"), *call("call_2", "b.yaml")] + + +def _responses_turns(task: str) -> tuple[list[JsonValue], list[JsonValue]]: + """Turns 1 and 3 of an agent on /v1/responses: 1 and 5 input items""" + + def call(call_id: str, path: str) -> list[JsonValue]: + return [ + { + "type": "function_call", + "call_id": call_id, + "name": "write_file", + "arguments": json.dumps({"path": path}), + }, + {"type": "function_call_output", "call_id": call_id, "output": f"wrote {path}"}, + ] + + first: Final[list[JsonValue]] = [{"role": "user", "content": task}] + return first, [*first, *call("call_1", "a.yaml"), *call("call_2", "b.yaml")] + + +def _body(path: str, model: str, conversation: list[JsonValue]) -> dict[str, JsonValue]: + if path == "/v1/responses": + return {"model": model, "input": conversation} + return {"model": model, "max_tokens": 16, "messages": conversation} + + +def _reply(path: str, text: str) -> Reply: + identity: Final = f"id_{uuid.uuid4().hex}" + if path == "/v1/messages": + return Reply(body=message_json(identity, "claude-sonnet-5-5", text)) + if path == "/v1/responses": + return responses_reply(identity, "gpt-5.6-sol", text, stream=False) + return chat_reply(identity, "gpt-5.4", text, stream=False) + + +def _answer(path: str, payload: dict[str, JsonValue]) -> str: + if path == "/v1/messages": + return string_value(_first(payload["content"])["text"]) + if path == "/v1/responses": + return string_value(_first(_first(payload["output"])["content"])["text"]) + return string_value(object_value(_first(payload["choices"])["message"])["content"]) + + +def _first(value: JsonValue) -> dict[str, JsonValue]: + assert isinstance(value, list), value + return object_value(value[0]) + + +def _cached_responses(redis: Redis) -> frozenset[bytes]: + digests: Final = tuple(key for key in redis.scan_iter() if len(key) == 64) + return frozenset(key for key in digests if b'"response"' in (redis.get(key) or b"")) + + +@pytest.mark.parametrize( + ("path", "model", "turns"), + [ + pytest.param("/v1/messages", "anthropic/claude-sonnet-5-5", _claude_code_turns, id="messages"), + pytest.param("/v1/chat/completions", "openai/gpt-5.4", _agent_turns, id="chat-completions"), + pytest.param("/v1/responses", "openai/responses/gpt-5.6-sol", _responses_turns, id="responses"), + ], +) +def test_cache_serves_a_turn_under_max_messages_and_skips_one_past_it( + gateway: Gateway, provider: SharedProvider, path: str, model: str, turns: _Turns +) -> None: + under_cap, past_cap = turns(f"update the config {uuid.uuid4().hex}") + provider.expect(_reply(path, "first answer"), _reply(path, "second answer"), _reply(path, "third answer")) + + with Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"])) as redis: + cached_before: Final = _cached_responses(redis) + gateway.post(path, _body(path, model, under_cap)) + eventually(lambda: _cached_responses(redis) - cached_before, lambda written: len(written) == 1) + repeated: Final = gateway.post(path, _body(path, model, under_cap)) + past_cap_twice: Final = ( + gateway.post(path, _body(path, model, past_cap)), + gateway.post(path, _body(path, model, past_cap)), + ) + + assert _answer(path, repeated) == "first answer", "a repeated turn under max_messages was not served from the cache" + assert tuple(_answer(path, answer) for answer in past_cap_twice) == ("second answer", "third answer"), ( + "a turn past max_messages was served from the cache" + ) + assert len(provider.received()) == 3 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 e29e92a2505..cd2f4996657 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) @@ -291,6 +560,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/management/test_credential_federation_values.py b/tests/integration/management/test_credential_federation_values.py new file mode 100644 index 00000000000..3f812c87c59 --- /dev/null +++ b/tests/integration/management/test_credential_federation_values.py @@ -0,0 +1,819 @@ +from __future__ import annotations + +import json +import uuid +from collections.abc import Callable, Iterator, Mapping, Sequence +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from functools import partial +from pathlib import Path +from types import MappingProxyType +from typing import Final, TypeVar +from urllib.parse import parse_qs + +import anthropic +import httpx +import jwt +import openai +import psutil +import pytest +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import ec +from pydantic import JsonValue + +from litellm.types.llms.anthropic import ANTHROPIC_OAUTH_TOKEN_PREFIX +from tests.integration._support.anthropic_thinking import ( + answer, + identity, + marker_of, + message_body, + message_events, + prompt, + stream_reply, + streams, + text_events, +) +from tests.integration._support.client import ( + JSON_OBJECT, + Gateway, + Scenario, + eventually, + gateway_from_environment, + object_value, + string_value, +) +from tests.integration._support.database import read_rows +from tests.integration._support.process import OwnedProxy, owned_proxy_process +from tests.integration._support.wire import Reply, Request, Wire, wire_server + +T = TypeVar("T") + +SOURCES: Final = ("token_file", "secret_reference", "internal_issuer", "keycloak", "environment") +CLIENTS: Final = ("chat", "messages", "messages_stream") + +_CLOSE: Final = MappingProxyType({"Connection": "close"}) +_PROVIDER: Final = MappingProxyType({"custom_llm_provider": "anthropic"}) +_SENSITIVE: Final = ( + "authorization", + "token", + "key", + "secret", + "vertex_credentials", + "credentials", + "password", + "passwd", +) +_DEFAULT_TOKEN_FILE: Final = "/var/run/secrets/integration/anthropic-identity-token" +_DEFAULT_KEYCLOAK_URL: Final = "https://keycloak.integration.invalid/realms/integration/protocol/openid-connect/token" +_KEYCLOAK_TARGET: Final = "/realms/integration/protocol/openid-connect/token" +_KEYCLOAK_ASSERTION: Final = "keycloak-scripted-assertion" +_ISSUER: Final = "https://issuer.integration.invalid" +_SUBJECT: Final = "workload-integration" +_AUDIENCE: Final = "https://api.anthropic.com" +_TTL_SECONDS: Final = 900 +_IDENTITY_TOKEN_VARIABLE: Final = "INTEGRATION_WIF_IDENTITY_TOKEN" +_SIGNING_KEY_VARIABLE: Final = "INTEGRATION_WIF_SIGNING_KEY" +_KEYCLOAK_SECRET_VARIABLE: Final = "INTEGRATION_WIF_KEYCLOAK_SECRET" +_JWT_BEARER: Final = "urn:ietf:params:oauth:grant-type:jwt-bearer" +_MODEL: Final = "anthropic/claude-haiku-4-5" +_CREDENTIAL_QUERY: Final = ( + 'SELECT credential_values, credential_info FROM "LiteLLM_CredentialsTable" WHERE credential_name = %s' +) + + +def _ids(rule_id: JsonValue) -> dict[str, JsonValue]: + return { + "anthropic_federation_rule_id": rule_id, + "anthropic_organization_id": "org-integration", + "anthropic_service_account_id": "svac-integration", + "anthropic_federation_workspace_id": "wrkspc-integration", + } + + +def _shape( + source: str, + rule_id: JsonValue, + *, + token_file: str = _DEFAULT_TOKEN_FILE, + keycloak_token_url: str = _DEFAULT_KEYCLOAK_URL, +) -> dict[str, JsonValue]: + match source: + case "token_file": + return {**_ids(rule_id), "anthropic_identity_token_file": token_file} + case "secret_reference": + return {**_ids(rule_id), "anthropic_identity_token": f"oidc/env/{_IDENTITY_TOKEN_VARIABLE}"} + case "internal_issuer": + return { + **_ids(rule_id), + "anthropic_identity_source": "internal_issuer", + "anthropic_issuer_url": _ISSUER, + "anthropic_issuer_subject": _SUBJECT, + "anthropic_issuer_audience": _AUDIENCE, + "anthropic_issuer_ttl_seconds": _TTL_SECONDS, + "anthropic_issuer_signing_key_ref": f"os.environ/{_SIGNING_KEY_VARIABLE}", + } + case "keycloak": + return { + **_ids(rule_id), + "anthropic_identity_source": "keycloak", + "anthropic_keycloak_token_url": keycloak_token_url, + "anthropic_keycloak_client_id": "litellm-integration", + "anthropic_keycloak_client_secret_ref": f"os.environ/{_KEYCLOAK_SECRET_VARIABLE}", + "anthropic_keycloak_auth_method": "client_secret_post", + "anthropic_keycloak_scope": "openid", + } + case "environment": + return _ids(rule_id) + case _: + pytest.fail(f"unknown identity source {source!r}") + + +def _sensitive(key: str) -> bool: + return any(word in key.lower() for word in _SENSITIVE) + + +def _masked(value: JsonValue) -> JsonValue: + if not isinstance(value, str): + return value + return "*****" if len(value) <= 4 else f"{value[:2]}****{value[-2:]}" + + +def _masked_view(values: Mapping[str, JsonValue]) -> dict[str, JsonValue]: + return {key: _masked(value) if _sensitive(key) else value for key, value in values.items()} + + +def _credential_name() -> str: + return f"federation-{uuid.uuid4().hex}" + + +def _create_body(name: str, values: Mapping[str, JsonValue]) -> dict[str, JsonValue]: + return {"credential_name": name, "credential_values": dict(values), "credential_info": dict(_PROVIDER)} + + +def _create(gateway: Gateway, scenario: Scenario, values: Mapping[str, JsonValue]) -> str: + name: Final = _credential_name() + scenario.cleanups.callback(_delete_if_present, gateway, name) + gateway.post("/credentials", _create_body(name, values)) + return name + + +def _delete_if_present(gateway: Gateway, name: str) -> None: + response: Final = gateway.request("DELETE", f"/credentials/{name}") + assert response.status_code in (200, 404), response.text + + +def _patch( + gateway: Gateway, + name: str, + values: Mapping[str, JsonValue], + delete: Sequence[str] = (), + *, + key: str | None = None, +) -> httpx.Response: + return gateway.request( + "PATCH", + f"/credentials/{name}", + { + "credential_name": name, + "credential_values": dict(values), + "credential_info": dict(_PROVIDER), + "credential_values_to_delete": list(delete), + }, + key=key, + ) + + +def _stored_values(gateway: Gateway, name: str, *, key: str | None = None) -> tuple[int, JsonValue]: + response: Final = gateway.request("GET", f"/credentials/by_name/{name}", key=key, headers=_CLOSE) + if response.status_code != 200: + return response.status_code, response.text + return 200, JSON_OBJECT.validate_json(response.content)["credential_values"] + + +def _listed_values(gateway: Gateway, name: str) -> JsonValue: + response: Final = gateway.request("GET", "/credentials", headers=_CLOSE) + assert response.status_code == 200, response.text + listed: Final = JSON_OBJECT.validate_json(response.content)["credentials"] + assert isinstance(listed, list), listed + matching: Final = tuple(object_value(entry) for entry in listed if object_value(entry)["credential_name"] == name) + return matching[0]["credential_values"] if matching else None + + +def _stable(read: Callable[[], T], satisfied: Callable[[T], bool], *, seconds: float = 70) -> T: + def spread_over_workers() -> tuple[T, ...]: + with ThreadPoolExecutor(max_workers=10) as pool: + return tuple(pool.map(lambda _: read(), range(10))) + + return eventually( + spread_over_workers, + lambda samples: all(satisfied(sample) for sample in samples), + seconds=seconds, + )[-1] + + +def _converged(gateway: Gateway, name: str, expected: Mapping[str, JsonValue], *, key: str | None = None) -> None: + masked: Final = _masked_view(expected) + _stable(partial(_stored_values, gateway, name, key=key), lambda sample: sample == (200, masked)) + + +def _db_row(name: str) -> tuple[dict[str, JsonValue], dict[str, JsonValue]]: + rows: Final = read_rows(_CREDENTIAL_QUERY, (name,)) + assert len(rows) == 1, rows + return object_value(rows[0]["credential_values"]), object_value(rows[0]["credential_info"]) + + +def _assert_encrypted_at_rest(name: str, submitted: Mapping[str, JsonValue]) -> None: + stored, info = _db_row(name) + assert info == dict(_PROVIDER), info + assert stored.keys() == submitted.keys(), stored + dumped: Final = json.dumps(stored) + for key, value in submitted.items(): + if not isinstance(value, str): + assert stored[key] == value, (key, stored[key]) + continue + assert stored[key] != value, key + if _sensitive(key): + assert value not in dumped, key + + +def _without(values: Mapping[str, JsonValue], *keys: str) -> dict[str, JsonValue]: + return {key: value for key, value in values.items() if key not in keys} + + +def _listed_models(gateway: Gateway) -> frozenset[str]: + response: Final = gateway.request("GET", "/model/info", headers=_CLOSE) + assert response.status_code == 200, response.text + entries: Final = JSON_OBJECT.validate_json(response.content)["data"] + assert isinstance(entries, list), entries + return frozenset(string_value(object_value(entry)["model_name"]) for entry in entries) + + +def _deployments_visible(gateway: Gateway, names: Sequence[str]) -> None: + wanted: Final = frozenset(names) + _stable(partial(_listed_models, gateway), lambda listed: wanted <= listed) + + +def _federated_deployment(gateway: Gateway, scenario: Scenario, credential: str, api_base: str) -> str: + name: Final = f"integration-{uuid.uuid4().hex}" + created: Final = gateway.post( + "/model/new", + { + "model_name": name, + "litellm_params": {"model": _MODEL, "api_base": api_base, "litellm_credential_name": credential}, + "model_info": {}, + }, + ) + scenario.cleanups.callback(scenario.delete_model, string_value(object_value(created["model_info"])["id"])) + return name + + +def _chat(gateway: Gateway, model: str, marker: str) -> httpx.Response: + return gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt(marker)}]}, + headers=_CLOSE, + ) + + +def _chat_outcome(gateway: Gateway, model: str) -> tuple[int, bool]: + marker: Final = uuid.uuid4().hex + response: Final = _chat(gateway, model, marker) + if response.status_code != 200: + return response.status_code, False + choices: Final = JSON_OBJECT.validate_json(response.content)["choices"] + assert isinstance(choices, list), choices + return 200, object_value(object_value(choices[0])["message"])["content"] == answer(marker) + + +@pytest.mark.parametrize("source", SOURCES) +def test_federation_shape_round_trips(gateway: Gateway, source: str) -> None: + with gateway.scenario() as scenario: + values: Final = _shape(source, f"fdrl-{uuid.uuid4().hex}") + name: Final = _create(gateway, scenario, values) + _converged(gateway, name, values) + _stable(partial(_listed_values, gateway, name), lambda listed: listed == _masked_view(values)) + _assert_encrypted_at_rest(name, values) + + +def test_patch_sets_and_deletes_values(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + values: Final = _shape("token_file", f"fdrl-{uuid.uuid4().hex}") + name: Final = _create(gateway, scenario, values) + _converged(gateway, name, values) + response: Final = _patch( + gateway, + name, + {"anthropic_federation_workspace_id": "wrkspc-updated"}, + ("anthropic_service_account_id",), + ) + assert response.status_code == 200, response.text + expected: Final = _without( + {**values, "anthropic_federation_workspace_id": "wrkspc-updated"}, "anthropic_service_account_id" + ) + _converged(gateway, name, expected) + _assert_encrypted_at_rest(name, expected) + + +def test_patch_with_empty_values_only_deletes(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + values: Final = _shape("token_file", f"fdrl-{uuid.uuid4().hex}") + name: Final = _create(gateway, scenario, values) + _converged(gateway, name, values) + response: Final = _patch(gateway, name, {}, ("anthropic_federation_workspace_id",)) + assert response.status_code == 200, response.text + expected: Final = _without(values, "anthropic_federation_workspace_id") + _converged(gateway, name, expected) + _assert_encrypted_at_rest(name, expected) + + +def test_patch_overlapping_set_and_delete_is_refused(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + values: Final = _shape("token_file", f"fdrl-{uuid.uuid4().hex}") + name: Final = _create(gateway, scenario, values) + _converged(gateway, name, values) + before: Final = _db_row(name) + response: Final = _patch( + gateway, + name, + {"anthropic_federation_workspace_id": "wrkspc-overlap"}, + ("anthropic_federation_workspace_id",), + ) + assert response.status_code == 400, response.text + assert "credential_values_to_delete overlaps credential_values for key(s)" in response.text, response.text + assert "anthropic_federation_workspace_id" in response.text, response.text + assert _db_row(name) == before + _converged(gateway, name, values) + + +_MALFORMED: Final[Mapping[str, JsonValue]] = MappingProxyType( + {"string": "fdrl-malformed", "list": ["fdrl-malformed"], "empty": {}, "null": None} +) + + +@pytest.mark.parametrize("shape", tuple(_MALFORMED)) +def test_malformed_credential_values_are_refused(gateway: Gateway, shape: str) -> None: + name: Final = _credential_name() + with gateway.scenario() as scenario: + scenario.cleanups.callback(_delete_if_present, gateway, name) + response: Final = gateway.request( + "POST", + "/credentials", + {"credential_name": name, "credential_values": _MALFORMED[shape], "credential_info": dict(_PROVIDER)}, + ) + assert response.status_code == 422, response.text + assert read_rows(_CREDENTIAL_QUERY, (name,)) == [] + assert _stored_values(gateway, name)[0] == 404 + + +def test_oversized_and_non_string_ids_round_trip(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + oversized: Final = _shape("environment", "f" * 5120) + numeric: Final = _shape("environment", 424242) + oversized_name: Final = _create(gateway, scenario, oversized) + numeric_name: Final = _create(gateway, scenario, numeric) + _converged(gateway, oversized_name, oversized) + _converged(gateway, numeric_name, numeric) + _assert_encrypted_at_rest(oversized_name, oversized) + _assert_encrypted_at_rest(numeric_name, numeric) + + +def test_duplicate_create_conflicts_and_repeated_patch_is_idempotent(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + values: Final = _shape("token_file", f"fdrl-{uuid.uuid4().hex}") + name: Final = _create(gateway, scenario, values) + duplicate: Final = gateway.request( + "POST", "/credentials", _create_body(name, _shape("token_file", f"fdrl-{uuid.uuid4().hex}")) + ) + assert duplicate.status_code == 409, duplicate.text + assert ( + f"Credential '{name}' already exists. Update it with PATCH /credentials/{name}, or delete it first." + in duplicate.text + ), duplicate.text + _converged(gateway, name, values) + _assert_encrypted_at_rest(name, values) + statuses: Final = tuple( + _patch(gateway, name, {"anthropic_federation_workspace_id": "wrkspc-twice"}).status_code for _ in range(2) + ) + assert statuses == (200, 200), statuses + expected: Final = {**values, "anthropic_federation_workspace_id": "wrkspc-twice"} + _converged(gateway, name, expected) + _assert_encrypted_at_rest(name, expected) + + +@pytest.mark.parametrize("grant", ("routed", "plain")) +def test_non_admin_cannot_write_federation_fields(gateway: Gateway, grant: str) -> None: + with gateway.scenario() as scenario: + values: Final = _shape("token_file", f"fdrl-{uuid.uuid4().hex}") + name: Final = _create(gateway, scenario, values) + _converged(gateway, name, values) + before: Final = _db_row(name) + model: Final = scenario.model() + user_id: Final = scenario.user(user_role="internal_user") + key: Final = ( + scenario.key(user_id=user_id, models=[model], allowed_routes=["/credentials*", "/v1/chat/completions"]) + if grant == "routed" + else scenario.key(user_id=user_id, models=[model]) + ) + intruder: Final = _credential_name() + scenario.cleanups.callback(_delete_if_present, gateway, intruder) + attempts: Final = ( + ( + "create", + gateway.request( + "POST", + "/credentials", + _create_body(intruder, _shape("token_file", f"fdrl-{uuid.uuid4().hex}")), + key=key, + ), + ), + ("set", _patch(gateway, name, {"anthropic_federation_rule_id": "fdrl-hijacked"}, key=key)), + ("unset", _patch(gateway, name, {}, ("anthropic_identity_token_file",), key=key)), + ("delete", gateway.request("DELETE", f"/credentials/{name}", key=key)), + ("jwks", gateway.request("GET", f"/credentials/{name}/jwks", key=key)), + ) + refused: Final = 403 if grant == "routed" else 401 + assert tuple((label, response.status_code) for label, response in attempts) == tuple( + (label, refused) for label, _ in attempts + ), tuple((label, response.text) for label, response in attempts) + if grant == "routed": + assert all("Only proxy admins" in response.text for _, response in attempts), tuple( + response.text for _, response in attempts + ) + listing: Final = gateway.request("GET", "/credentials", key=key, headers=_CLOSE) + assert listing.status_code == (200 if grant == "routed" else 401), listing.text + assert _DEFAULT_TOKEN_FILE not in listing.text, listing.text + if grant == "routed": + _converged(gateway, name, values, key=key) + else: + assert _stored_values(gateway, name, key=key)[0] == 401 + assert _db_row(name) == before + assert read_rows(_CREDENTIAL_QUERY, (intruder,)) == [] + _converged(gateway, name, values) + gateway.chat(model, key=key) + + +def test_untrusted_exchange_host_is_refused_before_any_exchange(gateway: Gateway) -> None: + with wire_server(lambda request: Reply(status=500, body=b'{"error": "never reached"}')) as wire: + with gateway.scenario() as scenario: + values: Final = _shape("token_file", f"fdrl-{uuid.uuid4().hex}") + name: Final = _create(gateway, scenario, values) + model: Final = _federated_deployment(gateway, scenario, name, wire.url) + _deployments_visible(gateway, (model,)) + response: Final = _chat(gateway, model, uuid.uuid4().hex) + assert response.status_code == 401, response.text + port: Final = httpx.URL(wire.url).port + assert "refused to use host '" in response.text, response.text + assert f":{port}'" in response.text, response.text + assert "LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS" in response.text, response.text + assert wire.drain() == () + assert wire.connections() == 0 + + +def test_concurrent_credential_writes_converge_across_workers(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + names: Final = tuple(_credential_name() for _ in range(10)) + for name in names: + scenario.cleanups.callback(_delete_if_present, gateway, name) + + def lane(name: str) -> tuple[int, int, int, int]: + created: Final = gateway.request( + "POST", "/credentials", _create_body(name, _shape("token_file", f"fdrl-{name}")) + ) + updated: Final = _patch(gateway, name, {"anthropic_federation_workspace_id": f"wrkspc-{name}"}) + trimmed: Final = _patch(gateway, name, {}, ("anthropic_service_account_id",)) + read: Final = gateway.request("GET", f"/credentials/by_name/{name}", headers=_CLOSE) + return created.status_code, updated.status_code, trimmed.status_code, read.status_code + + with ThreadPoolExecutor(max_workers=10) as pool: + outcomes: Final = tuple(pool.map(lane, names)) + assert all(outcome[:3] == (200, 200, 200) for outcome in outcomes), outcomes + assert all(outcome[3] in (200, 404) for outcome in outcomes), outcomes + for name in names: + expected: Final = _without( + {**_shape("token_file", f"fdrl-{name}"), "anthropic_federation_workspace_id": f"wrkspc-{name}"}, + "anthropic_service_account_id", + ) + _converged(gateway, name, expected) + _assert_encrypted_at_rest(name, expected) + + +@dataclass(frozen=True, slots=True) +class _Peer: + wire: Wire + seen: list[Request] + + def requests(self) -> tuple[Request, ...]: + self.seen.extend(self.wire.drain()) + return tuple(self.seen) + + +def _exchange_reply(request: Request) -> Reply: + grant: Final = JSON_OBJECT.validate_json(request.body) + minted: Final = { + "access_token": f"{ANTHROPIC_OAUTH_TOKEN_PREFIX}01-exchanged-{string_value(grant['federation_rule_id'])}", + "token_type": "Bearer", + "expires_in": 3600, + } + return Reply(body=json.dumps(minted).encode()) + + +def _federation_peer(request: Request) -> Reply: + if request.target == "/v1/oauth/token": + return _exchange_reply(request) + if request.target == _KEYCLOAK_TARGET: + return Reply( + body=json.dumps({"access_token": _KEYCLOAK_ASSERTION, "token_type": "Bearer", "expires_in": 60}).encode() + ) + marker: Final = marker_of(request) + if streams(request): + return stream_reply(request, message_events(marker, (text_events(0, answer(marker)),))) + return Reply(body=message_body(marker)) + + +def _es256_private_key_pem() -> str: + return ( + ec.generate_private_key(ec.SECP256R1()) + .private_bytes(serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption()) + .decode() + ) + + +@dataclass(frozen=True, slots=True) +class _Secrets: + allowed_dir: Path + token_file: Path + identity_token: str + environment_token: str + keycloak_secret: str + signing_key_pem: str + + def overrides(self) -> Mapping[str, str]: + return MappingProxyType( + { + "LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS": "127.0.0.1", + "LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS": str(self.allowed_dir), + "PROXY_CONFIG_RELOAD_INTERVAL_SECONDS": "2", + _IDENTITY_TOKEN_VARIABLE: self.identity_token, + _SIGNING_KEY_VARIABLE: self.signing_key_pem, + _KEYCLOAK_SECRET_VARIABLE: self.keycloak_secret, + "ANTHROPIC_IDENTITY_TOKEN": self.environment_token, + } + ) + + +_REMOVED_ENVIRONMENT: Final = ("ANTHROPIC_IDENTITY_TOKEN_FILE", "ANTHROPIC_API_KEY", "ANTHROPIC_API_BASE") + + +def _secrets(directory: Path) -> _Secrets: + allowed: Final = directory / "secrets" + allowed.mkdir() + token_file: Final = allowed / "anthropic-identity-token" + token_file.write_text(f"file-token-{uuid.uuid4().hex}") + return _Secrets( + allowed_dir=allowed, + token_file=token_file, + identity_token=f"env-token-{uuid.uuid4().hex}", + environment_token=f"ambient-token-{uuid.uuid4().hex}", + keycloak_secret=f"keycloak-secret-{uuid.uuid4().hex}", + signing_key_pem=_es256_private_key_pem(), + ) + + +@dataclass(frozen=True, slots=True) +class FederationRig: + owned: OwnedProxy + peer: _Peer + secrets: _Secrets + rule_ids: Mapping[str, str] + credentials: Mapping[str, str] + deployments: Mapping[str, str] + + +@pytest.fixture(scope="module") +def federation(tmp_path_factory: pytest.TempPathFactory) -> Iterator[FederationRig]: + directory: Final = tmp_path_factory.mktemp("federation").resolve() + secrets: Final = _secrets(directory) + with gateway_from_environment() as gateway, wire_server(_federation_peer) as wire: + with owned_proxy_process( + gateway, directory, secrets.overrides(), remove_environment=_REMOVED_ENVIRONMENT, workers=2 + ) as owned: + with owned.gateway.scenario() as scenario: + rule_ids: Final = {source: f"fdrl-{source}-{uuid.uuid4().hex}" for source in SOURCES} + credentials: Final = { + source: _create( + owned.gateway, + scenario, + _shape( + source, + rule_ids[source], + token_file=str(secrets.token_file), + keycloak_token_url=f"{wire.url}{_KEYCLOAK_TARGET}", + ), + ) + for source in SOURCES + } + deployments: Final = { + source: _federated_deployment(owned.gateway, scenario, credentials[source], wire.url) + for source in SOURCES + } + _deployments_visible(owned.gateway, tuple(deployments.values())) + yield FederationRig( + owned=owned, + peer=_Peer(wire, []), + secrets=secrets, + rule_ids=MappingProxyType(rule_ids), + credentials=MappingProxyType(credentials), + deployments=MappingProxyType(deployments), + ) + + +def _call(rig: FederationRig, client: str, model: str, marker: str) -> None: + base_url: Final = str(rig.owned.gateway.client.base_url) + key: Final = rig.owned.gateway.key + match client: + case "chat": + completion: Final = openai.OpenAI( + base_url=base_url + "/v1", api_key=key, max_retries=0 + ).chat.completions.create(model=model, messages=[{"role": "user", "content": prompt(marker)}]) + assert completion.choices[0].message.content == answer(marker), completion + case "messages": + message: Final = anthropic.Anthropic(base_url=base_url, api_key=key, max_retries=0).messages.create( + model=model, max_tokens=64, messages=[{"role": "user", "content": prompt(marker)}] + ) + assert message.id == identity(marker), message + assert [block.text for block in message.content if block.type == "text"] == [answer(marker)], message + case "messages_stream": + with anthropic.Anthropic(base_url=base_url, api_key=key, max_retries=0).messages.stream( + model=model, max_tokens=64, messages=[{"role": "user", "content": prompt(marker)}] + ) as stream: + final: Final = stream.get_final_message() + assert final.id == identity(marker), final + assert [block.text for block in final.content if block.type == "text"] == [answer(marker)], final + case _: + pytest.fail(f"unknown client {client!r}") + + +def _assert_assertion(rig: FederationRig, source: str, assertion: str, upstream: Sequence[Request]) -> None: + match source: + case "token_file": + assert assertion == rig.secrets.token_file.read_text() + case "secret_reference": + assert assertion == rig.secrets.identity_token + case "environment": + assert assertion == rig.secrets.environment_token + case "keycloak": + assert assertion == _KEYCLOAK_ASSERTION + grants: Final = tuple(request for request in upstream if request.target == _KEYCLOAK_TARGET) + assert grants, upstream + for grant in grants: + assert grant.headers.get("content-type") == "application/x-www-form-urlencoded", grant.headers + assert "authorization" not in grant.headers, grant.headers + assert parse_qs(grant.body.decode()) == { + "grant_type": ["client_credentials"], + "client_id": ["litellm-integration"], + "client_secret": [rig.secrets.keycloak_secret], + "scope": ["openid"], + }, grant.body + case "internal_issuer": + exported: Final = rig.owned.gateway.get(f"/credentials/{rig.credentials[source]}/jwks")["keys"] + assert isinstance(exported, list) and len(exported) == 1, exported + jwk: Final = object_value(exported[0]) + header: Final = jwt.get_unverified_header(assertion) + assert (header["alg"], header["kid"]) == ("ES256", jwk["kid"]), (header, jwk) + claims: Final = jwt.decode( + assertion, + jwt.PyJWK(dict(jwk)).key, + algorithms=["ES256"], + audience=_AUDIENCE, + issuer=_ISSUER, + options={"verify_exp": False}, + ) + assert claims["sub"] == _SUBJECT, claims + assert claims["exp"] - claims["iat"] == _TTL_SECONDS, claims + assert claims["jti"], claims + case _: + pytest.fail(f"unknown identity source {source!r}") + + +@pytest.mark.timeout(240) +@pytest.mark.parametrize("client", CLIENTS) +@pytest.mark.parametrize("source", SOURCES) +def test_federated_exchange(federation: FederationRig, source: str, client: str) -> None: + marker: Final = uuid.uuid4().hex + rule_id: Final = federation.rule_ids[source] + _call(federation, client, federation.deployments[source], marker) + upstream: Final = federation.peer.requests() + sent: Final = tuple( + request for request in upstream if request.target == "/v1/messages" and marker in request.body.decode() + ) + assert len(sent) == 1, upstream + assert sent[0].headers.get("authorization") == f"Bearer {ANTHROPIC_OAUTH_TOKEN_PREFIX}01-exchanged-{rule_id}", sent[ + 0 + ].headers + assert "oauth-2025-04-20" in sent[0].headers.get("anthropic-beta", ""), sent[0].headers + assert "x-api-key" not in sent[0].headers, sent[0].headers + exchanges: Final = tuple( + JSON_OBJECT.validate_json(request.body) for request in upstream if request.target == "/v1/oauth/token" + ) + mine: Final = tuple(grant for grant in exchanges if grant["federation_rule_id"] == rule_id) + assert mine, exchanges + for grant in mine: + assert grant["grant_type"] == _JWT_BEARER, grant + assert (grant["organization_id"], grant["service_account_id"], grant["workspace_id"]) == ( + "org-integration", + "svac-integration", + "wrkspc-integration", + ), grant + _assert_assertion(federation, source, string_value(grant["assertion"]), upstream) + + +@pytest.mark.timeout(240) +def test_token_file_outside_allowed_dirs_is_refused_before_any_exchange( + federation: FederationRig, tmp_path: Path +) -> None: + stray: Final = tmp_path.resolve() / "anthropic-identity-token" + stray.write_text("stray-token") + rule_id: Final = f"fdrl-stray-{uuid.uuid4().hex}" + marker: Final = uuid.uuid4().hex + owned: Final = federation.owned + with owned.gateway.scenario() as scenario: + name: Final = _create(owned.gateway, scenario, _shape("token_file", rule_id, token_file=str(stray))) + model: Final = _federated_deployment(owned.gateway, scenario, name, federation.peer.wire.url) + _deployments_visible(owned.gateway, (model,)) + response: Final = _chat(owned.gateway, model, marker) + assert response.status_code == 401, response.text + assert "LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS" in response.text, response.text + upstream: Final = federation.peer.requests() + assert not any(rule_id in request.body.decode() for request in upstream), upstream + assert not any(marker in request.body.decode() for request in upstream), upstream + assert owned.gateway.request("GET", "/health/liveliness").status_code == 200 + control: Final = scenario.model() + assert _chat_outcome(owned.gateway, control)[0] == 200 + + +def _is_worker(child: psutil.Process) -> bool: + try: + return "spawn_main" in " ".join(child.cmdline()) and child.status() != psutil.STATUS_ZOMBIE + except psutil.Error: + return False + + +def _workers(owned: OwnedProxy) -> tuple[psutil.Process, ...]: + return tuple(child for child in psutil.Process(owned.process.pid).children() if _is_worker(child)) + + +@pytest.mark.timeout(240) +def test_worker_kill_mid_credential_burst(gateway: Gateway, tmp_path: Path) -> None: + directory: Final = tmp_path.resolve() + secrets: Final = _secrets(directory) + with wire_server(_federation_peer) as wire: + with owned_proxy_process( + gateway, directory, secrets.overrides(), remove_environment=_REMOVED_ENVIRONMENT, workers=2 + ) as owned: + with owned.gateway.scenario() as scenario: + rule_id: Final = f"fdrl-burst-{uuid.uuid4().hex}" + credential: Final = _create( + owned.gateway, scenario, _shape("token_file", rule_id, token_file=str(secrets.token_file)) + ) + model: Final = _federated_deployment(owned.gateway, scenario, credential, wire.url) + _deployments_visible(owned.gateway, (model,)) + assert _chat_outcome(owned.gateway, model) == (200, True) + names: Final = tuple(_credential_name() for _ in range(24)) + for name in names: + scenario.cleanups.callback(_delete_if_present, owned.gateway, name) + + def create(name: str) -> tuple[str, int | str]: + try: + response: Final = owned.gateway.request( + "POST", "/credentials", _create_body(name, _shape("token_file", f"fdrl-{name}")) + ) + except httpx.TransportError as error: + return name, type(error).__name__ + return name, response.status_code + + victim: Final = eventually(lambda: _workers(owned), lambda workers: len(workers) == 2)[0] + with ThreadPoolExecutor(max_workers=8) as pool: + futures: Final = tuple(pool.submit(create, name) for name in names) + victim.kill() + outcomes: Final = tuple(future.result() for future in futures) + assert all(status == 200 or isinstance(status, str) for _, status in outcomes), outcomes + landed: Final = tuple(name for name, status in outcomes if status == 200) + for name in landed: + assert len(read_rows(_CREDENTIAL_QUERY, (name,))) == 1, name + respawned: Final = eventually( + lambda: frozenset(worker.pid for worker in _workers(owned)), + lambda pids: len(pids) == 2 and victim.pid not in pids, + seconds=30, + ) + assert f"Child process [{victim.pid}] died" in owned.log.read_text(), respawned + with httpx.Client( + base_url=owned.gateway.client.base_url, + timeout=15, + trust_env=False, + limits=httpx.Limits(max_keepalive_connections=0), + ) as fresh: + survivor: Final = Gateway(fresh, owned.gateway.key, owned.gateway.upstream_url) + for name in landed: + _converged(survivor, name, _shape("token_file", f"fdrl-{name}")) + _stable(partial(_chat_outcome, survivor, model), lambda outcome: outcome == (200, True)) diff --git a/tests/integration/mcp/test_mcp_accounting_guardrails.py b/tests/integration/mcp/test_mcp_accounting_guardrails.py index a276718c280..0d1c707d1ff 100644 --- a/tests/integration/mcp/test_mcp_accounting_guardrails.py +++ b/tests/integration/mcp/test_mcp_accounting_guardrails.py @@ -468,7 +468,7 @@ def test_generic_sink_and_native_hooks_receive_listed_metadata_on_typed_keys_wit assert row["status"] == "success", row assert _tool_metadata(row)["name"] == "lookup", row metadata: Final = row["metadata"] - assert isinstance(metadata, dict) and metadata["applied_guardrails"] == [hooks_rig.guardrail], metadata + assert isinstance(metadata, dict) and metadata["applied_guardrails"].count(hooks_rig.guardrail) == 1, metadata def test_pre_call_mask_reaches_the_peer_and_post_call_mask_reaches_the_caller_on_one_call_id( diff --git a/tests/integration/mcp/test_mcp_llm_endpoints.py b/tests/integration/mcp/test_mcp_llm_endpoints.py index af3bff1fd8b..9e946dfc253 100644 --- a/tests/integration/mcp/test_mcp_llm_endpoints.py +++ b/tests/integration/mcp/test_mcp_llm_endpoints.py @@ -934,7 +934,7 @@ def test_identical_nonstream_repeat_is_a_cache_hit_without_new_model_peer_or_hoo assert _spend_row(key, repeat.call_id)["cache_hit"] == "True" -def test_messages_bridge_hook_keeps_the_base_shape_without_request_local_metadata(hooked: Hooked) -> None: +def test_messages_bridge_hook_sees_the_definition_the_request_served(hooked: Hooked) -> None: with _bridge_rig(hooked, "messages") as rig: key: Final = _bridge_key(rig) marker: Final = "m" + uuid.uuid4().hex @@ -943,6 +943,6 @@ def test_messages_bridge_hook_keeps_the_base_shape_without_request_local_metadat assert found.status_code == 200, found.text blocked: Final = rig.post(key, probe, [rig.mcp("lookup")]) assert blocked.status_code == 200, blocked.text - assert _echoed(JSON_VALUE.validate_json(blocked.content)) == COLD, blocked.text + assert _echoed(JSON_VALUE.validate_json(blocked.content)) == _served(LOOKUP), blocked.text assert rig.peer_calls() == (("lookup", {"query": marker}),) assert rig.hook_messages(marker) == (f"Tool: lookup\nArguments: {dict(query=marker)}",) diff --git a/tests/integration/mcp/test_pagination.py b/tests/integration/mcp/test_pagination.py index 84689232038..b8b28bd39ce 100644 --- a/tests/integration/mcp/test_pagination.py +++ b/tests/integration/mcp/test_pagination.py @@ -1,7 +1,7 @@ import asyncio from contextlib import asynccontextmanager from pathlib import Path -from typing import Literal +from typing import Final, Literal import httpx import pytest @@ -138,7 +138,7 @@ def test_continuations_reauthorize_and_reject_registry_changes(tmp_path: Path, m assert os.environ.get("DATABASE_URL"), "This integration case requires disposable-database access" - async def exercise(a, b, peer, identity, owner, stranger, policy): + async def exercise(a, b, peer, identity, owner, stranger, policy, spare): owner_a = Gateway(a.client, owner, peer.url) owner_b = Gateway(b.client, owner, peer.url) stranger_b = Gateway(b.client, stranger, peer.url) @@ -165,9 +165,9 @@ def test_continuations_reauthorize_and_reject_registry_changes(tmp_path: Path, m { "key": owner, **( - {"access_group_ids": []} + {"access_group_ids": [], "object_permission": {"mcp_servers": [spare]}} if grant == "access_group" - else {"object_permission": {"mcp_servers": ["no-mcp-servers"]}} + else {"object_permission": {"mcp_servers": [spare]}} ), }, ) @@ -177,7 +177,7 @@ def test_continuations_reauthorize_and_reject_registry_changes(tmp_path: Path, m with pytest.raises(MCPError, match="fresh listing"): await getattr(session, method)(params=PaginatedRequestParams(cursor=first.next_cursor)) assert not any(call["body"].get("method", "").endswith("/list") for call in peer.drain()) - a.post("/key/update", {"key": owner, **policy}) + a.post("/key/update", {"key": owner, "object_permission": {"mcp_servers": []}, **policy}) changed = a.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "new catalog generation"}) assert changed.status_code == 202, changed.text for method, first in first_pages.items(): @@ -207,7 +207,7 @@ def test_continuations_reauthorize_and_reject_registry_changes(tmp_path: Path, m capture_output=True, text=True, ) - with paginated_mcp_peer() as peer, httpx.Client() as client: + with paginated_mcp_peer() as peer, paginated_mcp_peer() as spare_peer, httpx.Client() as client: seed = Gateway(client, "sk-pagination-test", peer.url) config = tmp_path / "database-proxy.yaml" config.write_text( @@ -231,6 +231,7 @@ def test_continuations_reauthorize_and_reject_registry_changes(tmp_path: Path, m a.scenario() as scenario, ): identity = register_mcp(scenario, peer, "pages") + spare: Final = register_mcp(scenario, spare_peer, "spare") group = a.request( "POST", "/v1/access_group", @@ -248,7 +249,7 @@ def test_continuations_reauthorize_and_reject_registry_changes(tmp_path: Path, m owner = scenario.key(**policy) stranger = scenario.key(object_permission={"mcp_servers": [identity]}) assert owner != stranger - asyncio.run(exercise(a, b, peer, identity, owner, stranger, policy)) + asyncio.run(exercise(a, b, peer, identity, owner, stranger, policy, spare)) @pytest.mark.parametrize("changed", ["key", "snapshot"]) diff --git a/tests/integration/observability/_logging_only_scope_support.py b/tests/integration/observability/_logging_only_scope_support.py new file mode 100644 index 00000000000..4c7d76d6564 --- /dev/null +++ b/tests/integration/observability/_logging_only_scope_support.py @@ -0,0 +1,1098 @@ +from __future__ import annotations + +import asyncio +import base64 +import binascii +import json +import os +import re +import uuid +from collections.abc import Callable, Iterator, Mapping +from contextlib import contextmanager +from contextvars import ContextVar +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final, Literal + +import httpx +import pytest +import yaml +from anthropic import Anthropic, AsyncAnthropic +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import rsa +from integration._support.client import Gateway, Scenario, eventually, object_value +from integration._support.database import read_rows, write_rows +from integration._support.upstream import delete_scenario, register_scenario +from integration._support.wire import Reply, Request, Wire +from integration._support.wire import wire_server as _wire_server +from openai import AsyncOpenAI, OpenAI +from openai.types.chat import ChatCompletionChunk +from pydantic import JsonValue, TypeAdapter + +from litellm.proxy.common_utils.callback_utils import CALLBACK_VAR_ENCRYPTED_PREFIX +from litellm.proxy.guardrails.guardrail_registry import decrypt_guardrail_litellm_params +from tests.integration.cost_calculation.cost_tracking_case import JsonResponse, SseResponse + +JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) + +Endpoint = Literal["chat", "messages", "responses"] + +ClientKind = Literal["openai_sync", "openai_async", "anthropic_sync", "anthropic_async", "httpx"] + +Direction = Literal["request", "response"] + +BASE_DEFAULT_NORMAL_DIRECTIONS: Final[Mapping[tuple[Endpoint, bool], tuple[Direction, ...]]] = MappingProxyType( + { + ("chat", False): ("request", "response"), + ("chat", True): ("request", "response"), + ("messages", False): ("request", "response"), + ("messages", True): ("request", "response"), + ("responses", False): ("request", "response"), + ("responses", True): ("request", "response"), + } +) + +BASE_DEFAULT_CACHE_HIT_DIRECTIONS: Final[Mapping[Endpoint, tuple[Direction, ...]]] = MappingProxyType( + { + "chat": ("request", "response"), + "messages": ("request", "response"), + "responses": ("request", "response"), + } +) + +BASE_DEFAULT_UPSTREAM_FAILURE_DIRECTIONS: Final[Mapping[Endpoint, tuple[Direction, ...]]] = MappingProxyType( + { + "chat": (), + "messages": (), + "responses": (), + } +) + +_AUDIT_RESPONSE_IDS: Final[ContextVar[tuple[str, ...]]] = ContextVar("audit_response_ids", default=()) + +_AUDIT_POLICY_REQUEST_COUNT: Final[ContextVar[int]] = ContextVar("audit_policy_request_count", default=0) + +_AUDIT_UPSTREAM_REQUEST_COUNT: Final[ContextVar[int]] = ContextVar("audit_upstream_request_count", default=0) + + +@dataclass(frozen=True, slots=True) +class CallerResult: + status: int + body: dict[str, JsonValue] + response_id: str + text: str + + +@dataclass(frozen=True, slots=True) +class ChaosCall: + index: int + endpoint: Endpoint + client_kind: ClientKind + model: str + stream: bool + prompt: str + call_id: str + + +@dataclass(frozen=True, slots=True) +class ChaosDeployment: + model_name: str + model: str + api_base: str + + +def _chaos_models(scenario: Scenario, marker: str) -> tuple[ChaosDeployment, ...]: + endpoints: Final[tuple[Endpoint, Endpoint, Endpoint]] = ("chat", "messages", "responses") + specs: Final = tuple((endpoint, False) for endpoint in endpoints) + tuple( + (endpoint, True) for endpoint in endpoints + ) + handles: Final = tuple( + register_scenario( + f"{marker}-{endpoint}-{'stream' if stream else 'complete'}", + _provider_response(endpoint, marker, f"synthetic K response {marker}", stream), + ) + for endpoint, stream in specs + ) + for handle in handles: + scenario.cleanups.callback(delete_scenario, handle) + + return tuple( + ChaosDeployment( + model_name=f"integration-{marker}-{endpoint}-{'stream' if stream else 'complete'}", + model={ + "chat": "openai/gpt-4o-mini", + "messages": "anthropic/claude-3-7-sonnet-20250219", + "responses": "openai/gpt-4.1-mini", + }[endpoint], + api_base=handle.api_base() if endpoint == "messages" else f"{handle.api_base()}/v1", + ) + for (endpoint, stream), handle in zip(specs, handles) + ) + + +def _chaos_model_list(deployments: tuple[ChaosDeployment, ...]) -> tuple[dict[str, JsonValue], ...]: + return tuple( + { + "model_name": deployment.model_name, + "litellm_params": { + "model": deployment.model, + "api_base": deployment.api_base, + "api_key": "synthetic-provider-key", + }, + } + for deployment in deployments + ) + + +def _chaos_control_configuration( + tmp_path: Path, + identity: str, + model_list: tuple[dict[str, JsonValue], ...], +) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"]["cache"] = False + config["model_list"] = list(model_list) + config["guardrails"] = [] + path: Final = tmp_path / f"{identity}-models.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _chaos_calls(deployments: tuple[ChaosDeployment, ...], marker: str) -> tuple[ChaosCall, ...]: + endpoints: Final[tuple[Endpoint, Endpoint, Endpoint]] = ("chat", "messages", "responses") + + def one(index: int) -> ChaosCall: + endpoint_index: Final = index % 3 + endpoint: Final = endpoints[endpoint_index] + stream: Final = (index // 3) % 2 == 1 + client_kind: Final[ClientKind] = ( + "openai_async" + if endpoint == "chat" and stream + else "openai_sync" + if endpoint in ("chat", "responses") + else "anthropic_async" + if stream + else "anthropic_sync" + ) + return ChaosCall( + index=index, + endpoint=endpoint, + client_kind=client_kind, + model=deployments[endpoint_index + 3 * int(stream)].model_name, + stream=stream, + prompt=f"synthetic K burst {marker}-{index}", + call_id=f"{marker}-k-{index}", + ) + + return tuple(one(index) for index in range(30)) + + +def _chaos_spend_minimums( + models: tuple[str, ...], calls: tuple[ChaosCall, ...], baseline_count: int +) -> tuple[int, ...]: + return tuple(baseline_count + sum(call.model == model for call in calls) for model in models) + + +def _proxy_url(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + + +def _is_base_audit_leg() -> bool: + leg: Final = os.environ.get("LITELLM_LOGGING_ONLY_SCOPE_AUDIT_LEG", "head") + assert leg in ("base", "head"), leg + return leg == "base" + + +def _record_response_id(response_id: str) -> None: + response_ids: Final = _AUDIT_RESPONSE_IDS.get() + _record_response_ids((response_id,) if response_id not in response_ids else ()) + + +def _record_response_ids(response_ids: tuple[str, ...]) -> None: + current: Final = _AUDIT_RESPONSE_IDS.get() + _AUDIT_RESPONSE_IDS.set(tuple(dict.fromkeys((*current, *response_ids)))) + + +def _record_policy_request_count(count: int) -> None: + _AUDIT_POLICY_REQUEST_COUNT.set(_AUDIT_POLICY_REQUEST_COUNT.get() + count) + + +def _record_upstream_request_count(count: int) -> None: + _AUDIT_UPSTREAM_REQUEST_COUNT.set(_AUDIT_UPSTREAM_REQUEST_COUNT.get() + count) + + +@contextmanager +def wire_server(respond: Callable[[Request], Reply], port: int = 0, *, policy_edge: bool = True) -> Iterator[Wire]: + received: Final[SimpleQueue[Request]] = SimpleQueue() + + def record(request: Request) -> Reply: + received.put(request) + return respond(request) + + try: + with _wire_server(record, port=port) as server: + yield server + finally: + if policy_edge: + _record_policy_request_count(received.qsize()) + + +def _directions_for_scope(base_default: tuple[Direction, ...], scope: str | None) -> tuple[Direction, ...]: + if scope is None or scope == "both": + return base_default + selected_direction: Final = "request" if scope == "input" else "response" + return tuple(direction for direction in base_default if direction == selected_direction) + + +def _directions_for_audit_leg(base_default: tuple[Direction, ...], scope: str | None) -> tuple[Direction, ...]: + if _is_base_audit_leg(): + return base_default + return _directions_for_scope(base_default, scope) + + +def _response_text(endpoint: Endpoint, body: Mapping[str, JsonValue]) -> str: + if endpoint == "chat": + choices: Final = body.get("choices") + assert isinstance(choices, list) and choices, body + message: Final = object_value(object_value(choices[0])["message"]) + return str(message["content"]) + if endpoint == "messages": + content: Final = body.get("content") + assert isinstance(content, list), body + return "".join( + str(object_value(block)["text"]) + for block in content + if isinstance(block, dict) and isinstance(block.get("text"), str) + ) + output: Final = body.get("output") + assert isinstance(output, list), body + return "".join( + _response_text_from_blocks(object_value(item).get("content")) + for item in output + if isinstance(item, dict) and object_value(item).get("type") == "message" + ) + + +def _response_text_from_blocks(value: JsonValue | None) -> str: + if not isinstance(value, list): + return "" + return "".join( + str(object_value(block)["text"]) + for block in value + if isinstance(block, dict) and isinstance(block.get("text"), str) + ) + + +def _chat_chunk_text(chunk: ChatCompletionChunk) -> str: + return "".join(choice.delta.content for choice in chunk.choices if isinstance(choice.delta.content, str)) + + +def _caller_result(endpoint: Endpoint, status: int, body: Mapping[str, JsonValue]) -> CallerResult: + response_id: Final = body.get("id") + assert isinstance(response_id, str), body + _record_response_id(response_id) + normalized: Final = JSON_OBJECT.validate_python(dict(body)) + return CallerResult(status, normalized, response_id, _response_text(endpoint, normalized)) + + +def _stream_result(endpoint: Endpoint, response_id: str, text: str) -> CallerResult: + _record_response_id(response_id) + body: Final = JSON_OBJECT.validate_python({"id": response_id, "text": text}) + return CallerResult(200, body, response_id, text) + + +def _response_body_without_ids(value: JsonValue) -> JsonValue: + if isinstance(value, dict): + return {key: _response_body_without_ids(item) for key, item in value.items() if key != "id"} + if isinstance(value, list): + return [_response_body_without_ids(item) for item in value] + return value + + +def _call_sync( + client_kind: ClientKind, + endpoint: Endpoint, + proxy_url: str, + key: str, + model: str, + prompt: str, + stream: bool, + call_id: str, +) -> CallerResult: + if client_kind == "httpx": + body: Final[dict[str, JsonValue]] = {"model": model, "messages": [{"role": "user", "content": prompt}]} + with httpx.Client(base_url=proxy_url, timeout=30, trust_env=False) as client: + response: Final = client.post( + "/v1/chat/completions", + json=body, + headers={"Authorization": f"Bearer {key}", "x-litellm-call-id": call_id}, + ) + return _caller_result(endpoint, response.status_code, JSON_OBJECT.validate_json(response.content)) + if client_kind == "anthropic_sync": + with Anthropic( + base_url=proxy_url, + api_key=key, + max_retries=0, + http_client=httpx.Client(timeout=30, trust_env=False), + ) as client: + if stream: + with client.messages.stream( + model=model, + max_tokens=32, + messages=[{"role": "user", "content": prompt}], + extra_headers={"x-litellm-call-id": call_id}, + ) as stream_response: + message: Final = stream_response.get_final_message() + return _stream_result( + endpoint, + message.id, + "".join(block.text for block in message.content if block.type == "text"), + ) + message: Final = client.messages.create( + model=model, + max_tokens=32, + messages=[{"role": "user", "content": prompt}], + extra_headers={"x-litellm-call-id": call_id}, + ) + return _caller_result(endpoint, 200, JSON_OBJECT.validate_python(message.model_dump(mode="json"))) + assert client_kind == "openai_sync", client_kind + with OpenAI( + base_url=f"{proxy_url}/v1", + api_key=key, + max_retries=0, + http_client=httpx.Client(timeout=30, trust_env=False), + ) as client: + headers: Final = {"x-litellm-call-id": call_id} + if endpoint == "chat": + if stream: + chunks: Final = tuple( + client.chat.completions.create( + model=model, + messages=[{"role": "user", "content": prompt}], + stream=True, + extra_headers=headers, + ) + ) + response_id: Final = chunks[0].id + text: Final = "".join(_chat_chunk_text(chunk) for chunk in chunks) + return _stream_result(endpoint, response_id, text) + response: Final = client.chat.completions.create( + model=model, + messages=[{"role": "user", "content": prompt}], + extra_headers=headers, + ) + return _caller_result(endpoint, 200, JSON_OBJECT.validate_python(response.model_dump(mode="json"))) + assert endpoint == "responses", endpoint + if stream: + events: Final = tuple( + client.responses.create(model=model, input=prompt, stream=True, extra_headers=headers) + ) + completed: Final = next(event.response for event in events if event.type == "response.completed") + body: Final = JSON_OBJECT.validate_python(completed.model_dump(mode="json")) + return _stream_result(endpoint, str(body["id"]), _response_text(endpoint, body)) + completion: Final = client.responses.create(model=model, input=prompt, extra_headers=headers) + return _caller_result(endpoint, 200, JSON_OBJECT.validate_python(completion.model_dump(mode="json"))) + + +async def _call_async( + client_kind: ClientKind, + endpoint: Endpoint, + proxy_url: str, + key: str, + model: str, + prompt: str, + stream: bool, + call_id: str, +) -> CallerResult: + if client_kind == "anthropic_async": + async with AsyncAnthropic( + base_url=proxy_url, + api_key=key, + max_retries=0, + http_client=httpx.AsyncClient(timeout=30, trust_env=False), + ) as client: + if stream: + async with client.messages.stream( + model=model, + max_tokens=32, + messages=[{"role": "user", "content": prompt}], + extra_headers={"x-litellm-call-id": call_id}, + ) as stream_response: + message: Final = await stream_response.get_final_message() + return _stream_result( + endpoint, + message.id, + "".join(block.text for block in message.content if block.type == "text"), + ) + message: Final = await client.messages.create( + model=model, + max_tokens=32, + messages=[{"role": "user", "content": prompt}], + extra_headers={"x-litellm-call-id": call_id}, + ) + return _caller_result(endpoint, 200, JSON_OBJECT.validate_python(message.model_dump(mode="json"))) + assert client_kind == "openai_async", client_kind + async with AsyncOpenAI( + base_url=f"{proxy_url}/v1", + api_key=key, + max_retries=0, + http_client=httpx.AsyncClient(timeout=30, trust_env=False), + ) as client: + headers: Final = {"x-litellm-call-id": call_id} + if endpoint == "chat": + assert stream, "The audit only uses the async OpenAI chat client for streaming rows" + stream_response: Final = await client.chat.completions.create( + model=model, + messages=[{"role": "user", "content": prompt}], + stream=True, + extra_headers=headers, + ) + chunks: Final = tuple([chunk async for chunk in stream_response]) + response_id: Final = chunks[0].id + text: Final = "".join( + choice.delta.content + for chunk in chunks + for choice in chunk.choices + if isinstance(choice.delta.content, str) + ) + return _stream_result(endpoint, response_id, text) + assert endpoint == "responses", endpoint + if stream: + responses_stream: Final = await client.responses.create( + model=model, input=prompt, stream=True, extra_headers=headers + ) + events: Final = tuple([event async for event in responses_stream]) + completed: Final = next(event.response for event in events if event.type == "response.completed") + body: Final = JSON_OBJECT.validate_python(completed.model_dump(mode="json")) + return _stream_result(endpoint, str(body["id"]), _response_text(endpoint, body)) + completion: Final = await client.responses.create(model=model, input=prompt, extra_headers=headers) + return _caller_result(endpoint, 200, JSON_OBJECT.validate_python(completion.model_dump(mode="json"))) + + +def _call_client( + client_kind: ClientKind, + endpoint: Endpoint, + gateway: Gateway, + model: str, + prompt: str, + stream: bool, + call_id: str, +) -> CallerResult: + if client_kind in ("openai_async", "anthropic_async"): + return _record_caller_result( + asyncio.run( + _call_async(client_kind, endpoint, _proxy_url(gateway), gateway.key, model, prompt, stream, call_id) + ) + ) + return _record_caller_result( + _call_sync(client_kind, endpoint, _proxy_url(gateway), gateway.key, model, prompt, stream, call_id) + ) + + +def _record_caller_result(result: CallerResult) -> CallerResult: + _record_response_id(result.response_id) + return result + + +def _call_cache_client(endpoint: Endpoint, gateway: Gateway, model: str, prompt: str, call_id: str) -> CallerResult: + path: Final = { + "chat": "/v1/chat/completions", + "messages": "/v1/messages", + "responses": "/v1/responses", + }[endpoint] + body: Final = { + "chat": {"model": model, "messages": [{"role": "user", "content": prompt}]}, + "messages": {"model": model, "max_tokens": 32, "messages": [{"role": "user", "content": prompt}]}, + "responses": {"model": model, "input": prompt}, + }[endpoint] + with httpx.Client(timeout=30, trust_env=False) as client: + response: Final = client.post( + f"{_proxy_url(gateway)}{path}", + json=body, + headers={"Authorization": f"Bearer {gateway.key}", "x-litellm-call-id": call_id}, + ) + assert response.status_code == 200, response.text + return _caller_result( + endpoint, + response.status_code, + JSON_OBJECT.validate_json(response.content), + ) + + +def _provider_response(endpoint: Endpoint, _scenario_id: str, reply: str, stream: bool) -> JsonResponse | SseResponse: + response_id: Final = { + "chat": "chatcmpl-$UNIQUE_ID", + "messages": "msg_$UNIQUE_ID", + "responses": "resp_$UNIQUE_ID", + }[endpoint] + if endpoint == "chat": + if stream: + return SseResponse( + content_type="text/event-stream", + frames=( + f"data: {json.dumps({'id': response_id, 'object': 'chat.completion.chunk', 'created': 1, 'model': 'gpt-4o-mini', 'choices': [{'index': 0, 'delta': {'role': 'assistant', 'content': reply}, 'finish_reason': None}]})}", + f"data: {json.dumps({'id': response_id, 'object': 'chat.completion.chunk', 'created': 1, 'model': 'gpt-4o-mini', 'choices': [{'index': 0, 'delta': {}, 'finish_reason': 'stop'}]})}", + "data: [DONE]", + ), + ) + return JsonResponse( + content_type="application/json", + body={ + "id": response_id, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": reply}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 9, "completion_tokens": 5, "total_tokens": 14}, + }, + ) + if endpoint == "messages": + if stream: + return SseResponse( + content_type="text/event-stream", + frames=( + f"event: message_start\ndata: {json.dumps({'type': 'message_start', 'message': {'id': response_id, 'type': 'message', 'role': 'assistant', 'content': [], 'model': 'claude-3-7-sonnet-20250219', 'stop_reason': None, 'stop_sequence': None, 'usage': {'input_tokens': 9, 'output_tokens': 0}}})}", + f"event: content_block_start\ndata: {json.dumps({'type': 'content_block_start', 'index': 0, 'content_block': {'type': 'text', 'text': ''}})}", + f"event: content_block_delta\ndata: {json.dumps({'type': 'content_block_delta', 'index': 0, 'delta': {'type': 'text_delta', 'text': reply}})}", + f"event: content_block_stop\ndata: {json.dumps({'type': 'content_block_stop', 'index': 0})}", + f"event: message_delta\ndata: {json.dumps({'type': 'message_delta', 'delta': {'stop_reason': 'end_turn', 'stop_sequence': None}, 'usage': {'output_tokens': 5}})}", + f"event: message_stop\ndata: {json.dumps({'type': 'message_stop'})}", + ), + ) + return JsonResponse( + content_type="application/json", + body={ + "id": response_id, + "type": "message", + "role": "assistant", + "model": "claude-3-7-sonnet-20250219", + "content": [{"type": "text", "text": reply}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 9, "output_tokens": 5}, + }, + ) + if stream: + completed: Final = { + "id": response_id, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4.1-mini", + "output": [ + { + "id": "msg_$UNIQUE_ID", + "type": "message", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": reply, "annotations": []}], + } + ], + "usage": {"input_tokens": 9, "output_tokens": 5, "total_tokens": 14}, + } + return SseResponse( + content_type="text/event-stream", + frames=( + f"data: {json.dumps({'type': 'response.created', 'response': {'id': response_id, 'object': 'response', 'created_at': 1, 'status': 'in_progress', 'model': 'gpt-4.1-mini', 'output': []}})}", + f"data: {json.dumps({'type': 'response.output_text.delta', 'item_id': 'msg_$UNIQUE_ID', 'output_index': 0, 'content_index': 0, 'delta': reply})}", + f"data: {json.dumps({'type': 'response.completed', 'response': completed})}", + ), + ) + return JsonResponse( + content_type="application/json", + body={ + "id": response_id, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4.1-mini", + "output": [ + { + "id": "msg_$UNIQUE_ID", + "type": "message", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": reply, "annotations": []}], + } + ], + "usage": {"input_tokens": 9, "output_tokens": 5, "total_tokens": 14}, + }, + ) + + +def _configuration( + tmp_path: Path, + identity: str, + policy_url: str, + scope: str | None, + *, + include_scope: bool = True, + default_on: bool = True, + mode: str | list[str] = "logging_only", + cache: bool = False, + num_retries: int | None = None, + model_list: tuple[dict[str, JsonValue], ...] = (), +) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"]["cache"] = cache + if num_retries is not None: + config["litellm_settings"]["num_retries"] = num_retries + config["model_list"] = list(model_list) + params: Final = { + "guardrail": "generic_guardrail_api", + "mode": mode, + "default_on": default_on, + "api_base": policy_url, + "api_key": "synthetic-guardrail-key", + "extra_headers": ["x-litellm-call-id"], + **({"logging_only_scope": scope} if include_scope else {}), + } + config["guardrails"] = [{"guardrail_name": identity, "litellm_params": params}] + scope_name: Final = scope if scope is not None else "unset" + path: Final = tmp_path / f"{identity}-{scope_name}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _content_filter_configuration(tmp_path: Path, identity: str, scope: str, blocked_word: str) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"]["cache"] = False + config["guardrails"] = [ + { + "guardrail_name": identity, + "litellm_params": { + "guardrail": "litellm_content_filter", + "mode": "logging_only", + "logging_only_scope": scope, + "default_on": True, + "blocked_words": [{"keyword": blocked_word, "action": "BLOCK"}], + }, + } + ] + path: Final = tmp_path / f"{identity}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _presidio_configuration( + tmp_path: Path, + identity: str, + analyzer_api_base: str, + anonymizer_api_base: str, + scope: str | None, +) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"]["cache"] = False + params: Final = { + "guardrail": "presidio", + "mode": "logging_only", + "default_on": True, + "presidio_analyzer_api_base": analyzer_api_base, + "presidio_anonymizer_api_base": anonymizer_api_base, + "pii_entities_config": {"PERSON": "MASK"}, + **({"logging_only_scope": scope} if scope is not None else {}), + } + config["guardrails"] = [{"guardrail_name": identity, "litellm_params": params}] + scope_name: Final = scope if scope is not None else "unset" + path: Final = tmp_path / f"{identity}-{scope_name}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _empty_proxy_configuration(tmp_path: Path, identity: str, reload_seconds: int = 30) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"]["cache"] = False + config["guardrails"] = [] + config["general_settings"]["proxy_config_reload_interval_seconds"] = reload_seconds + path: Final = tmp_path / f"{identity}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _model_armor_configuration(tmp_path: Path, identity: str, api_endpoint: str, token_uri: str, scope: str) -> Path: + private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + private_key_pem: Final = private_key.private_bytes( + serialization.Encoding.PEM, + serialization.PrivateFormat.PKCS8, + serialization.NoEncryption(), + ).decode() + credentials: Final = { + "type": "service_account", + "project_id": "synthetic-model-armor-project", + "private_key_id": "synthetic-key-id", + "private_key": private_key_pem, + "client_email": "integration-model-armor@synthetic-project.iam.gserviceaccount.com", + "client_id": "123456789012345678901", + "token_uri": token_uri, + "auth_uri": "https://accounts.google.com/o/oauth2/auth", + "auth_provider_x509_cert_url": "https://www.googleapis.com/oauth2/v1/certs", + "client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/integration", + } + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"]["cache"] = False + config["guardrails"] = [ + { + "guardrail_name": identity, + "litellm_params": { + "guardrail": "model_armor", + "mode": "logging_only", + "logging_only_scope": scope, + "default_on": True, + "template_id": "synthetic-template", + "project_id": "synthetic-model-armor-project", + "location": "us-central1", + "credentials": json.dumps(credentials), + "api_endpoint": api_endpoint, + }, + } + ] + path: Final = tmp_path / f"{identity}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _insert_database_guardrail( + identity: str, + policy_url: str, + scope: str | None, + *, + mode: str = "pre_call", + default_on: bool = True, +) -> None: + params: Final = { + "guardrail": "generic_guardrail_api", + "mode": mode, + "default_on": default_on, + "api_base": policy_url, + "api_key": "synthetic-guardrail-key", + "extra_headers": ["x-litellm-call-id"], + **({"logging_only_scope": scope} if scope is not None else {}), + } + write_rows( + 'INSERT INTO "LiteLLM_GuardrailsTable" ' + "(guardrail_id, guardrail_name, litellm_params, guardrail_info, updated_at) " + "VALUES (%s, %s, %s::jsonb, %s::jsonb, NOW())", + (str(uuid.uuid5(uuid.NAMESPACE_URL, identity)), identity, json.dumps(params), "{}"), + ) + + +@contextmanager +def _database_guardrail( + identity: str, + policy_url: str, + scope: str | None, + *, + mode: str = "pre_call", + default_on: bool = True, +) -> Iterator[None]: + _insert_database_guardrail(identity, policy_url, scope, mode=mode, default_on=default_on) + try: + yield + finally: + _delete_database_guardrail(identity) + + +def _delete_database_guardrail(identity: str) -> None: + write_rows('DELETE FROM "LiteLLM_GuardrailsTable" WHERE guardrail_name=%s', (identity,)) + + +def _post_guardrail_body( + identity: str, + provider: str, + mode: str | list[str], + api_base: str, + scope: JsonValue, + include_scope: bool = True, +) -> dict[str, JsonValue]: + params: Final = { + "guardrail": provider, + "mode": mode, + "default_on": True, + **( + { + "presidio_analyzer_api_base": api_base, + "presidio_anonymizer_api_base": api_base, + "pii_entities_config": {"PERSON": "MASK"}, + } + if provider == "presidio" + else {"api_base": api_base, "api_key": "synthetic-guardrail-key"} + ), + **({"extra_headers": ["x-litellm-call-id"]} if provider == "generic_guardrail_api" else {}), + **({"logging_only_scope": scope} if include_scope else {}), + } + return { + "guardrail": { + "guardrail_name": identity, + "litellm_params": params, + "guardrail_info": {"description": "phase-12 logging scope audit"}, + } + } + + +def _create_guardrail(candidate: Gateway, identity: str, params: Mapping[str, JsonValue]) -> dict[str, JsonValue]: + guardrail_params: Final = { + **params, + **({"extra_headers": ["x-litellm-call-id"]} if params.get("guardrail") == "generic_guardrail_api" else {}), + } + response: Final = candidate.request( + "POST", + "/guardrails", + { + "guardrail": { + "guardrail_name": identity, + "litellm_params": guardrail_params, + "guardrail_info": {"description": "phase-12 logging scope audit"}, + } + }, + ) + assert response.status_code == 200, response.text + return JSON_OBJECT.validate_json(response.content) + + +def _management_guardrail_rows(identity: str) -> tuple[dict[str, JsonValue], ...]: + rows: Final = tuple( + object_value(row) + for row in read_rows( + "SELECT guardrail_id, guardrail_name, litellm_params, guardrail_info " + 'FROM "LiteLLM_GuardrailsTable" WHERE guardrail_name=%s', + (identity,), + ) + ) + return tuple({**row, "litellm_params": _decrypted_management_litellm_params(row["litellm_params"])} for row in rows) + + +def _decrypted_management_litellm_params(stored_value: JsonValue) -> dict[str, JsonValue]: + stored: Final = object_value(stored_value) + salt_key: Final = os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt") + with pytest.MonkeyPatch.context() as monkeypatch: + monkeypatch.setenv("LITELLM_SALT_KEY", salt_key) + decrypted: Final = JSON_OBJECT.validate_python(decrypt_guardrail_litellm_params(stored)) + assert _decryption_only_changes_encrypted_values(stored, decrypted) + return decrypted + + +def _decryption_only_changes_encrypted_values(stored: JsonValue, decrypted: JsonValue) -> bool: + if isinstance(stored, dict) and isinstance(decrypted, dict): + return stored.keys() == decrypted.keys() and all( + _decryption_only_changes_encrypted_values(value, decrypted[key]) for key, value in stored.items() + ) + if isinstance(stored, list) and isinstance(decrypted, list): + return len(stored) == len(decrypted) and all( + _decryption_only_changes_encrypted_values(stored_value, decrypted_value) + for stored_value, decrypted_value in zip(stored, decrypted) + ) + if isinstance(stored, str) and stored.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX): + return isinstance(decrypted, str) and not decrypted.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX) + return stored == decrypted + + +def _drain_upstream(upstream_url: str) -> tuple[dict[str, JsonValue], ...]: + response: Final = httpx.get(f"{upstream_url.rstrip('/')}/__observations", trust_env=False, timeout=15) + response.raise_for_status() + requests: Final = object_value(JSON_OBJECT.validate_python(response.json())).get("requests") + assert isinstance(requests, list), response.text + observations: Final = tuple(object_value(request) for request in requests) + forwarded_requests: Final = tuple(request for request in observations if request.get("method", "POST") != "GET") + _record_upstream_request_count(len(forwarded_requests)) + return forwarded_requests + + +def _json_contains_exact_string(value: JsonValue, expected: str) -> bool: + if isinstance(value, str): + return value == expected + if isinstance(value, list): + return any(_json_contains_exact_string(item, expected) for item in value) + if isinstance(value, dict): + return any(_json_contains_exact_string(item, expected) for item in value.values()) + return False + + +@pytest.fixture(autouse=True) +def _record_audit_properties( + request: pytest.FixtureRequest, + record_property: Callable[[str, object], None], + gateway: Gateway, +) -> Iterator[None]: + response_ids_token: Final = _AUDIT_RESPONSE_IDS.set(()) + policy_count_token: Final = _AUDIT_POLICY_REQUEST_COUNT.set(0) + upstream_count_token: Final = _AUDIT_UPSTREAM_REQUEST_COUNT.set(0) + try: + response: Final = httpx.get(f"{gateway.upstream_url.rstrip('/')}/__observations", trust_env=False, timeout=15) + response.raise_for_status() + yield + finally: + node_id: Final = request.node.nodeid + inventory_ids: Final = re.findall(r"[A-Z]{1,2}\d+", node_id) + record_property("node_id", node_id) + record_property("inventory_id", inventory_ids[0] if inventory_ids else "support") + record_property("response_ids", ",".join(_AUDIT_RESPONSE_IDS.get())) + record_property("policy_edge_request_count", str(_AUDIT_POLICY_REQUEST_COUNT.get())) + record_property("upstream_request_count", str(_AUDIT_UPSTREAM_REQUEST_COUNT.get())) + _AUDIT_RESPONSE_IDS.reset(response_ids_token) + _AUDIT_POLICY_REQUEST_COUNT.reset(policy_count_token) + _AUDIT_UPSTREAM_REQUEST_COUNT.reset(upstream_count_token) + + +def _policy_call_id_matches(payload: Mapping[str, JsonValue], call_id: str) -> bool: + actual: Final = payload.get("litellm_call_id") + if actual == call_id: + return True + headers: Final = payload.get("request_headers") + return isinstance(headers, dict) and any( + key.lower() == "x-litellm-call-id" and value == call_id for key, value in headers.items() + ) + + +def _policy_call_id(payload: Mapping[str, JsonValue]) -> str | None: + actual: Final = payload.get("litellm_call_id") + if isinstance(actual, str): + return actual + headers: Final = payload.get("request_headers") + if not isinstance(headers, dict): + return None + return next( + (value for key, value in headers.items() if key.lower() == "x-litellm-call-id" and isinstance(value, str)), + None, + ) + + +def _chat_request(gateway: Gateway, model: str, prompt: str, call_id: str) -> httpx.Response: + return gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": call_id}, + ) + + +def _direction(payload: Mapping[str, JsonValue]) -> str: + value: Final = payload.get("input_type") + assert value in ("request", "response"), payload + return str(value) + + +def _cache_hit(value: JsonValue) -> bool: + return value is True or value == "True" + + +def _guardrail_mode_values(value: JsonValue) -> tuple[str, ...]: + if isinstance(value, list): + return tuple(str(mode) for mode in value) + return (str(value),) + + +def _guardrail_mode_status_pairs( + entries: tuple[dict[str, JsonValue], ...], +) -> tuple[tuple[tuple[str, ...], str], ...]: + return tuple((_guardrail_mode_values(entry["guardrail_mode"]), str(entry["guardrail_status"])) for entry in entries) + + +def _spend_row_for_response_id(response_id: str) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id, metadata FROM "LiteLLM_SpendLogs" WHERE request_id=%s', + (response_id,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + row: Final = object_value(rows[0]) + _record_response_ids((str(row["request_id"]),)) + return row + + +def _spend_rows(model: str, minimum: int) -> tuple[dict[str, JsonValue], ...]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id, metadata, cache_hit FROM "LiteLLM_SpendLogs" WHERE model_group=%s', + (model,), + ), + lambda rows: len(rows) >= minimum, + seconds=70, + ) + response_ids: Final = tuple(str(row["request_id"]) for row in rows) + _record_response_ids(response_ids) + return rows + + +def _spend_rows_for_calls( + models: tuple[str, ...], + expected: tuple[tuple[str, str, str], ...], + *, + tolerate_missing: bool = False, +) -> tuple[dict[str, JsonValue], ...]: + model_placeholders: Final = ", ".join("%s" for _ in models) + query: Final = ( + "SELECT model_group, request_id, metadata, cache_hit " + f'FROM "LiteLLM_SpendLogs" WHERE model_group IN ({model_placeholders})' + ) + rows: Final = eventually( + lambda: tuple(read_rows(query, models)), + lambda values: all( + len(_spend_rows_matching_call(values, model, call_id)) == 1 for model, _, call_id in expected + ), + seconds=45 if tolerate_missing else 70, + return_last_on_timeout=tolerate_missing, + ) + _record_response_ids(tuple(response_id for _, response_id, _ in expected)) + return rows + + +def _spend_rows_matching_call( + rows: tuple[dict[str, JsonValue], ...], + model: str, + call_id: str, +) -> tuple[dict[str, JsonValue], ...]: + return tuple( + row + for row in rows + if row["model_group"] == model and object_value(row["metadata"]).get("litellm_call_id") == call_id + ) + + +def _spend_row_for_call_id(call_id: str) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id, metadata FROM "LiteLLM_SpendLogs" WHERE request_id=%s', + (call_id,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + assert str(rows[0]["request_id"]) == call_id, rows + row: Final = object_value(rows[0]) + _record_response_ids((str(row["request_id"]),)) + return row + + +def _guardrail_entries(row: Mapping[str, JsonValue]) -> tuple[dict[str, JsonValue], ...]: + metadata: Final = object_value(row["metadata"]) + entries: Final = metadata.get("guardrail_information") + if not isinstance(entries, list): + return () + return tuple(object_value(entry) for entry in entries) + + +def _response_id_matches(endpoint: Endpoint, request_id: str, response_id: str, scenario_id: str) -> bool: + if request_id == response_id: + return True + if endpoint != "responses" or not request_id.startswith("resp_"): + return False + encoded: Final = request_id.removeprefix("resp_") + padding: Final = "=" * (-len(encoded) % 4) + try: + decoded: Final = base64.urlsafe_b64decode(encoded + padding).decode("utf-8") + except (binascii.Error, UnicodeDecodeError): + return False + return response_id in decoded or scenario_id in decoded + + +def _assert_response_id(endpoint: Endpoint, request_id: str, response_id: str, scenario_id: str) -> None: + assert _response_id_matches(endpoint, request_id, response_id, scenario_id), ( + endpoint, + request_id, + response_id, + scenario_id, + ) diff --git a/tests/integration/observability/conftest.py b/tests/integration/observability/conftest.py index c5150877857..fcac8eaba3c 100644 --- a/tests/integration/observability/conftest.py +++ b/tests/integration/observability/conftest.py @@ -8,8 +8,9 @@ from urllib.parse import urlparse import pytest import yaml +from integration._support.client import eventually from integration._support.otlp_sink import SpanSinks, owned_sinks -from integration._support.prometheus_series import CapRig, series_cap_rig +from integration._support.prometheus_series import CapRig, series_cap_rig, spend_rows from pydantic import JsonValue AuditConfigWriter = Callable[[Path, Mapping[str, JsonValue]], Path] @@ -63,4 +64,9 @@ def capped(tmp_path_factory: pytest.TempPathFactory) -> Iterator[CapRig]: workers=2, warm_keys=3, ) as rig: + eventually( + lambda: tuple(len(spend_rows(key.alias)) for key in rig.warm), + lambda counts: all(count == 1 for count in counts), + seconds=70, + ) yield rig diff --git a/tests/integration/observability/test_cache_hit_guardrail_metrics_chaos.py b/tests/integration/observability/test_cache_hit_guardrail_metrics_chaos.py index c77b8eebe33..7e3b479e5f0 100644 --- a/tests/integration/observability/test_cache_hit_guardrail_metrics_chaos.py +++ b/tests/integration/observability/test_cache_hit_guardrail_metrics_chaos.py @@ -263,7 +263,7 @@ def test_worker_kill_mid_burst_keeps_counting(gateway: Gateway, tmp_path: Path) def test_proxy_restart_mid_burst_keeps_counting(gateway: Gateway, tmp_path: Path) -> None: - """X4: restart the owned proxy between the two halves; pre-restart count asserted, then recounted.""" + """X4: the boot wipes the kept directory, so the second proxy counts only the second half.""" marker: Final = uuid.uuid4().hex prom_dir: Final = tmp_path / "prom" prom_dir.mkdir() @@ -321,8 +321,8 @@ def test_proxy_restart_mid_burst_keeps_counting(gateway: Gateway, tmp_path: Path _populated(_samples(owned_two.gateway, (model,)), deployment), _blank(_samples(owned_two.gateway, (model,))), ), - lambda observed: observed[0] == len(named) and observed[1] == 0, + lambda observed: observed[0] == len(second_half) and observed[1] == 0, seconds=70, ) - assert post[0] == len(named), (pre, post, outcomes_two) + assert post[0] == len(second_half), (pre, post, outcomes_two) owned_two.gateway.post("/model/delete", {"id": deployment}) diff --git a/tests/integration/observability/test_guardrail_effects.py b/tests/integration/observability/test_guardrail_effects.py index 1df4427bb4e..79feb4f0d92 100644 --- a/tests/integration/observability/test_guardrail_effects.py +++ b/tests/integration/observability/test_guardrail_effects.py @@ -8,6 +8,7 @@ import threading import uuid from collections.abc import Callable, Iterator, Mapping from concurrent.futures import ThreadPoolExecutor +from datetime import datetime, timezone from contextlib import contextmanager from dataclasses import dataclass from pathlib import Path @@ -1805,8 +1806,161 @@ def test_responses_pre_call_denial_stream_survives_worker_kill(gateway: Gateway, for response in responses: assert response.status_code == 200, response.text assert response.headers["content-type"].startswith("text/event-stream"), response.text +@pytest.mark.parametrize( + ("logging_only_scope", "scanned_directions"), + (("input", ("request",)), ("output", ("response",)), ("both", ("request", "response"))), +) +def test_logging_only_scope_observes_only_the_configured_direction_without_blocking( + gateway: Gateway, tmp_path: Path, logging_only_scope: str, scanned_directions: tuple[str, ...] +) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + prompt: Final = "synthetic observed prompt " + identity + reply: Final = "synthetic observed reply " + identity + texts_by_direction: Final = {"request": [prompt], "response": [reply]} + + def guardrail(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api" + return Reply(body=json.dumps({"action": "BLOCKED", "blocked_reason": "synthetic observed denial"}).encode()) + + def provider(request: Request) -> Reply: + assert request.target == "/v1/chat/completions" + assert json.loads(request.body)["messages"] == [{"role": "user", "content": prompt}] + return Reply( + body=json.dumps( + { + "id": identity, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": reply}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 9, "completion_tokens": 5, "total_tokens": 14}, + } + ).encode() + ) + + with wire_server(guardrail) as policy, wire_server(provider) as upstream: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [ + { + "guardrail_name": identity, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "logging_only", + "logging_only_scope": logging_only_scope, + "default_on": True, + "api_base": policy.url, + "api_key": "synthetic-guardrail-key", + }, + } + ] + path: Final = tmp_path / "logging-only-scope.yaml" + path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(api_base=upstream.url + "/v1") + response: Final = candidate.request( + "POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": prompt}]} + ) + assert response.status_code == 200, response.text + assert response.json()["choices"][0]["message"]["content"] == reply, response.text + assert len(upstream.drain()) == 1 + rows: Final = eventually( + lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)), + lambda values: len(values) == 1, + seconds=70, + ) + scans: Final = tuple(json.loads(scan.body) for scan in policy.drain()) + assert [(scan["input_type"], scan["texts"]) for scan in scans] == [ + (direction, texts_by_direction[direction]) for direction in scanned_directions + ], scans + entries: Final = object_value(rows[0]["metadata"])["guardrail_information"] + assert isinstance(entries, list), rows[0] + assert [ + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in map(object_value, entries) + ] == [(identity, "logging_only", "guardrail_intervened")] * len(scanned_directions), entries + today: Final = datetime.now(timezone.utc).date().isoformat() + guardrail_id: Final = next( + object_value(row)["guardrail_id"] + for row in candidate.get("/v2/guardrails/list")["guardrails"] + if object_value(row)["guardrail_name"] == identity + ) + detail: Final = eventually( + lambda: candidate.request( + "GET", + f"/guardrails/usage/detail/{guardrail_id}", + params={"start_date": today, "end_date": today}, + ).json(), + lambda body: body["requestsEvaluated"] >= len(scanned_directions), + seconds=30, + return_last_on_timeout=True, + ) + assert detail["requestsEvaluated"] == len(scanned_directions), detail +@pytest.mark.parametrize("logging_only_scope", ("input", "Input")) +def test_logging_only_scope_literal_or_mode_mismatch_is_ignored_at_load_and_keeps_blocking( + gateway: Gateway, tmp_path: Path, logging_only_scope: str +) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + prompt: Final = "synthetic invalid-scope prompt pineapple " + identity + + def guardrail(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api" + return Reply(body=json.dumps({"action": "BLOCKED", "blocked_reason": "synthetic policy denial"}).encode()) + + def provider(request: Request) -> Reply: + assert request.target == "/v1/chat/completions" + assert json.loads(request.body)["messages"] == [{"role": "user", "content": prompt}] + return Reply( + body=json.dumps( + { + "id": identity, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "unchanged provider reply"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 9, "completion_tokens": 5, "total_tokens": 14}, + } + ).encode() + ) + + with wire_server(guardrail) as policy, wire_server(provider) as upstream: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [ + { + "guardrail_name": identity, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "logging_only_scope": logging_only_scope, + "default_on": True, + "blocked_words": [{"keyword": "pineapple", "action": "BLOCK"}], + "api_base": policy.url, + "api_key": "synthetic-guardrail-key", + }, + } + ] + path: Final = tmp_path / "invalid-scope-pre-call.yaml" + path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(api_base=upstream.url + "/v1") + response: Final = candidate.request( + "POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": prompt}]} + ) + assert response.status_code == 400, response.text + assert "synthetic policy denial" in response.text, response.text + assert len(policy.drain()) == 1 + assert len(upstream.drain()) == 0 + guardrails: Final = candidate.get("/v2/guardrails/list")["guardrails"] + assert any(object_value(row)["guardrail_name"] == identity for row in guardrails), guardrails _TOKEN: Final = re.compile(rb"token-[0-9a-f]{32}-\d+") diff --git a/tests/integration/observability/test_langtrace_delivery.py b/tests/integration/observability/test_langtrace_delivery.py index 1e27538120b..8463460436f 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_logging_only_scope_chaos.py b/tests/integration/observability/test_logging_only_scope_chaos.py new file mode 100644 index 00000000000..d70073a5071 --- /dev/null +++ b/tests/integration/observability/test_logging_only_scope_chaos.py @@ -0,0 +1,762 @@ +from __future__ import annotations + +import signal +import socket +import threading +import time +import uuid +from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor +from itertools import accumulate, repeat +from pathlib import Path +from queue import SimpleQueue +from typing import Final + +import httpx +import psutil +import pytest +from _logging_only_scope_support import ( + JSON_OBJECT, + CallerResult, + ChaosCall, + _assert_response_id, + _call_client, + _chaos_calls, + _chaos_control_configuration, + _chaos_model_list, + _chaos_models, + _chaos_spend_minimums, + _configuration, + _direction, + _directions_for_audit_leg, + _drain_upstream, + _guardrail_entries, + _is_base_audit_leg, + _json_contains_exact_string, + _policy_call_id_matches, + _spend_rows, + _spend_rows_for_calls, + _spend_rows_matching_call, + wire_server, +) +from _logging_only_scope_support import ( + _record_audit_properties as _record_audit_properties, +) +from anthropic import APIConnectionError as AnthropicAPIConnectionError +from integration._support.client import Gateway, eventually +from integration._support.process import owned_proxy, owned_proxy_process +from integration._support.wire import Reply, Request +from openai import APIConnectionError as OpenAIAPIConnectionError +from pydantic import JsonValue + + +def test_K1_policy_edge_restart_mid_burst_keeps_output_observation_fail_open(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-k1-{uuid.uuid4().hex}" + marker: Final = uuid.uuid4().hex + + def policy(_request: Request) -> Reply: + return Reply(body=b'{"action":"NONE"}') + + with socket.socket() as reservation: + reservation.bind(("127.0.0.1", 0)) + policy_port: Final = reservation.getsockname()[1] + with gateway.scenario() as scenario: + deployments: Final = _chaos_models(scenario, marker) + models: Final = tuple(deployment.model_name for deployment in deployments) + model_list: Final = _chaos_model_list(deployments) + chaos_calls: Final = _chaos_calls(deployments, marker) + expected_directions: Final = _directions_for_audit_leg(("request", "response"), "output") + control_config: Final = _chaos_control_configuration(tmp_path, identity, model_list) + with owned_proxy(gateway, tmp_path, {}, config=control_config, workers=2) as control_proxy: + controls: Final = tuple( + _call_client( + call.client_kind, + call.endpoint, + control_proxy, + call.model, + call.prompt, + call.stream, + f"{call.call_id}-control", + ) + for call in chaos_calls + ) + control_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(control_upstream) == 30, control_upstream + assert ( + tuple( + sum(_json_contains_exact_string(observation["body"], call.prompt) for observation in control_upstream) + for call in chaos_calls + ) + == (1,) * 30 + ), control_upstream + config: Final = _configuration( + tmp_path, + identity, + f"http://127.0.0.1:{policy_port}", + "output", + model_list=model_list, + ) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=2) as candidate: + starts: Final = tuple(threading.Event() for _ in range(3)) + + def run(index: int) -> tuple[int, CallerResult]: + call: Final = chaos_calls[index] + phase: Final = index // 10 + assert starts[phase].wait(timeout=90), (index, phase) + return index, _call_client( + call.client_kind, + call.endpoint, + candidate, + call.model, + call.prompt, + call.stream, + call.call_id, + ) + + with ThreadPoolExecutor(max_workers=30) as pool: + futures: Final = tuple(pool.submit(run, index) for index in range(30)) + try: + with wire_server(policy, port=policy_port) as initial_edge: + starts[0].set() + first: Final = tuple(futures[index].result(timeout=90) for index in range(10)) + eventually( + lambda: initial_edge.received.qsize(), + lambda count: count == 10 * len(expected_directions), + seconds=30, + ) + tuple( + _spend_rows(model, minimum) + for model, minimum in zip(models, _chaos_spend_minimums(models, chaos_calls[:10], 5)) + ) + starts[1].set() + middle: Final = tuple(futures[index].result(timeout=90) for index in range(10, 20)) + tuple( + _spend_rows(model, minimum) + for model, minimum in zip(models, _chaos_spend_minimums(models, chaos_calls[:20], 5)) + ) + with wire_server(policy, port=policy_port) as recovered_edge: + starts[2].set() + recovered: Final = tuple(futures[index].result(timeout=90) for index in range(20, 30)) + eventually( + lambda: recovered_edge.received.qsize(), + lambda count: count == 10 * len(expected_directions), + seconds=30, + ) + tuple(_spend_rows(model, 10) for model in models) + finally: + for start in starts: + start.set() + results: Final = first + middle + recovered + assert tuple(index for index, _ in results) == tuple(range(30)), results + assert all(result.status == controls[index].status for index, result in results), results + assert all(result.text == controls[index].text for index, result in results), results + candidate_ids: Final = tuple(result.response_id for _, result in results) + assert len(set(candidate_ids)) == 30, candidate_ids + observed_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(observed_upstream) == 30, observed_upstream + assert ( + tuple( + sum( + _json_contains_exact_string(observation["body"], call.prompt) + for observation in observed_upstream + ) + for call in chaos_calls + ) + == (1,) * 30 + ), observed_upstream + expected_success_ids: Final = frozenset( + call.call_id for call in chaos_calls if call.index < 10 or call.index >= 20 + ) + edge_payloads: Final = tuple( + JSON_OBJECT.validate_json(call.body) for call in initial_edge.drain() + recovered_edge.drain() + ) + assert len(edge_payloads) == 20 * len(expected_directions), edge_payloads + successful_calls: Final = tuple(call for call in chaos_calls if call.call_id in expected_success_ids) + for call in successful_calls: + payloads_for_call: Final = tuple( + payload for payload in edge_payloads if _policy_call_id_matches(payload, call.call_id) + ) + assert tuple(sorted(_direction(payload) for payload in payloads_for_call)) == tuple( + sorted(expected_directions) + ), ( + call.call_id, + payloads_for_call, + ) + assert all( + payload["texts"] + == ([call.prompt] if _direction(payload) == "request" else [results[call.index][1].text]) + for payload in payloads_for_call + ), payloads_for_call + rows: Final = _spend_rows_for_calls( + models, + tuple( + ( + chaos_calls[index].model, + response_id, + chaos_calls[index].call_id, + ) + for index, response_id in enumerate(candidate_ids) + ), + ) + for index, call in enumerate(chaos_calls): + matching_rows: Final = _spend_rows_matching_call(rows, call.model, call.call_id) + assert len(matching_rows) == 1, (call.call_id, matching_rows) + row: Final = matching_rows[0] + expected_status: Final = "guardrail_failed_to_respond" if 10 <= index < 20 else "success" + entries: Final = _guardrail_entries(row) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) for entry in entries + ) == tuple((identity, "logging_only", expected_status) for _ in expected_directions), (index, entries) + + +def test_K2_policy_edge_delay_does_not_delay_concurrent_callers(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-k2-{uuid.uuid4().hex}" + marker: Final = uuid.uuid4().hex + + def policy(_request: Request) -> Reply: + time.sleep(2) + return Reply(body=b'{"action":"NONE"}') + + with gateway.scenario() as scenario: + deployments: Final = _chaos_models(scenario, marker) + models: Final = tuple(deployment.model_name for deployment in deployments) + model_list: Final = _chaos_model_list(deployments) + chaos_calls: Final = _chaos_calls(deployments, marker) + expected_directions: Final = _directions_for_audit_leg(("request", "response"), "output") + control_config: Final = _chaos_control_configuration(tmp_path, identity, model_list) + with owned_proxy(gateway, tmp_path, {}, config=control_config, workers=2) as control_proxy: + controls: Final = tuple( + _call_client( + call.client_kind, + call.endpoint, + control_proxy, + call.model, + call.prompt, + call.stream, + f"{call.call_id}-control", + ) + for call in chaos_calls + ) + control_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(control_upstream) == 30, control_upstream + assert ( + tuple( + sum(_json_contains_exact_string(observation["body"], call.prompt) for observation in control_upstream) + for call in chaos_calls + ) + == (1,) * 30 + ), control_upstream + with wire_server(policy) as edge: + config: Final = _configuration(tmp_path, identity, edge.url, "output", model_list=model_list) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=2) as candidate: + + def run(call: ChaosCall) -> tuple[int, CallerResult, float]: + started: Final = time.monotonic() + result: Final = _call_client( + call.client_kind, + call.endpoint, + candidate, + call.model, + call.prompt, + call.stream, + call.call_id, + ) + return call.index, result, time.monotonic() - started + + with ThreadPoolExecutor(max_workers=30) as pool: + results: Final = tuple(pool.map(run, chaos_calls)) + assert tuple(index for index, _, _ in results) == tuple(range(30)), results + assert all(result.status == controls[index].status for index, result, _ in results), results + assert all(result.text == controls[index].text for index, result, _ in results), results + assert all(duration < 2 for _, _, duration in results), results + eventually( + lambda: edge.received.qsize(), + lambda count: count == 30 * len(expected_directions), + seconds=30, + ) + edge_calls: Final = edge.drain() + payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in edge_calls) + for call in chaos_calls: + payloads_for_call: Final = tuple( + payload for payload in payloads if _policy_call_id_matches(payload, call.call_id) + ) + assert tuple(sorted(_direction(payload) for payload in payloads_for_call)) == tuple( + sorted(expected_directions) + ), ( + call, + payloads_for_call, + ) + assert all( + payload["texts"] + == ([call.prompt] if _direction(payload) == "request" else [controls[call.index].text]) + for payload in payloads_for_call + ), payloads_for_call + response_ids: Final = tuple(result.response_id for _, result, _ in results) + assert len(set(response_ids)) == 30, response_ids + upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(upstream) == 30, upstream + assert ( + tuple( + sum(_json_contains_exact_string(observation["body"], call.prompt) for observation in upstream) + for call in chaos_calls + ) + == (1,) * 30 + ), upstream + rows: Final = _spend_rows_for_calls( + models, + tuple( + (call.model, result.response_id, call.call_id) + for call, (_, result, _) in zip(chaos_calls, results) + ), + ) + for call, (_, result, _) in zip(chaos_calls, results): + matching_rows: Final = _spend_rows_matching_call(rows, call.model, call.call_id) + assert len(matching_rows) == 1, (call, result.response_id, matching_rows) + row: Final = matching_rows[0] + entries: Final = _guardrail_entries(row) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == tuple((identity, "logging_only", "success") for _ in expected_directions), entries + + +@pytest.mark.timeout(180) +def test_K3_two_worker_sigkill_checks_post_kill_spend_rows( + gateway: Gateway, + tmp_path: Path, + record_property: Callable[[str, object], None], +) -> None: + identity: Final = f"logging-scope-k3-{uuid.uuid4().hex}" + marker: Final = uuid.uuid4().hex + scan_started: Final = threading.Event() + release_scans: Final = threading.Event() + + def policy(_request: Request) -> Reply: + scan_started.set() + assert release_scans.wait(timeout=60), identity + return Reply(body=b'{"action":"NONE"}') + + with gateway.scenario() as scenario: + deployments: Final = _chaos_models(scenario, marker) + models: Final = tuple(deployment.model_name for deployment in deployments) + model_list: Final = _chaos_model_list(deployments) + calls: Final = _chaos_calls(deployments, marker) + expected_directions: Final = _directions_for_audit_leg(("request", "response"), "output") + expected_entries: Final = tuple((identity, "logging_only", "success") for _ in expected_directions) + control_config: Final = _chaos_control_configuration(tmp_path, identity, model_list) + with owned_proxy(gateway, tmp_path, {}, config=control_config, workers=2) as control_proxy: + controls: Final = tuple( + _call_client( + call.client_kind, + call.endpoint, + control_proxy, + call.model, + call.prompt, + call.stream, + f"{call.call_id}-control", + ) + for call in calls + ) + assert len(_drain_upstream(gateway.upstream_url)) == 30 + with wire_server(policy) as edge: + config: Final = _configuration(tmp_path, identity, edge.url, "output", model_list=model_list) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + candidate: Final = owned.gateway + root: Final = psutil.Process(owned.process.pid) + workers: Final = eventually( + lambda: tuple( + child + for child in root.children(recursive=True) + if any("spawn_main" in part for part in child.cmdline()) + ), + lambda children: len(children) == 2, + seconds=30, + ) + + def run(call: ChaosCall) -> tuple[int, CallerResult | None, str | None]: + try: + result: Final = _call_client( + call.client_kind, + call.endpoint, + candidate, + call.model, + call.prompt, + call.stream, + call.call_id, + ) + return call.index, result, None + except (OpenAIAPIConnectionError, AnthropicAPIConnectionError, httpx.RemoteProtocolError) as error: + return call.index, None, str(error) + + with ThreadPoolExecutor(max_workers=30) as pool: + futures: Final = tuple(pool.submit(run, call) for call in calls) + try: + assert eventually(lambda: scan_started.is_set(), bool, seconds=30) + eventually(lambda: edge.received.qsize(), lambda count: count >= 5, seconds=30) + workers[0].send_signal(signal.SIGKILL) + killed_workers, surviving_workers = psutil.wait_procs((workers[0],), timeout=10) + assert len(killed_workers) == 1 and not surviving_workers, ( + killed_workers, + surviving_workers, + ) + finally: + release_scans.set() + outcomes: Final = tuple(future.result(timeout=90) for future in futures) + assert owned.process.poll() is None, "Proxy supervisor exited after a worker was killed" + successful: Final = tuple( + (calls[index], result) for index, result, error in outcomes if result is not None and error is None + ) + assert successful, outcomes + assert all( + result.status == controls[call.index].status and result.text == controls[call.index].text + for call, result in successful + ), successful + pre_kill_expected: Final = tuple( + (call.model, result.response_id, call.call_id) for call, result in successful + ) + pre_kill_rows: Final = _spend_rows_for_calls( + models, + pre_kill_expected, + tolerate_missing=True, + ) + pre_kill_rows_by_call: Final = tuple( + (call, _spend_rows_matching_call(pre_kill_rows, call.model, call.call_id)) for call in calls + ) + assert all(len(rows) <= 1 for _, rows in pre_kill_rows_by_call), pre_kill_rows_by_call + pre_kill_missing_rows: Final = sum(not rows for _, rows in pre_kill_rows_by_call) + record_property("k3_pre_kill_missing_spend_rows", pre_kill_missing_rows) + for call, matching_rows in pre_kill_rows_by_call: + if not matching_rows: + continue + entries: Final = _guardrail_entries(matching_rows[0]) + assert tuple( + sorted( + ( + entry["guardrail_name"], + entry["guardrail_mode"], + entry["guardrail_status"], + ) + for entry in entries + ) + ) == tuple(sorted(expected_entries)), (call, entries) + + post_kill_templates: Final = calls[:6] + post_kill_calls: Final = tuple( + ChaosCall( + index=call.index, + endpoint=call.endpoint, + client_kind=call.client_kind, + model=call.model, + stream=call.stream, + prompt=f"synthetic K post-kill burst {marker}-{call.index}", + call_id=f"{marker}-k-post-kill-{call.index}", + ) + for call in post_kill_templates + ) + + def run_post_kill(call: ChaosCall) -> CallerResult: + return _call_client( + call.client_kind, + call.endpoint, + candidate, + call.model, + call.prompt, + call.stream, + call.call_id, + ) + + with ThreadPoolExecutor(max_workers=len(post_kill_calls)) as pool: + post_kill_futures: Final = tuple(pool.submit(run_post_kill, call) for call in post_kill_calls) + post_kill_results: Final = tuple(future.result(timeout=90) for future in post_kill_futures) + assert all( + result.status == controls[call.index].status and result.text == controls[call.index].text + for call, result in zip(post_kill_calls, post_kill_results) + ), post_kill_results + served: Final = successful + tuple(zip(post_kill_calls, post_kill_results)) + response_ids: Final = tuple(result.response_id for _, result in served) + assert len(set(response_ids)) == len(response_ids), response_ids + served_calls: Final = tuple(call for call, _ in served) + requested_call_ids: Final = frozenset(call.call_id for call in calls + post_kill_calls) + upstream: Final = _drain_upstream(gateway.upstream_url) + assert all( + sum(_json_contains_exact_string(observation["body"], call.prompt) for observation in upstream) == 1 + for call in served_calls + ), upstream + + def accumulate_policy_payloads( + collected: tuple[dict[str, JsonValue], ...], _: None + ) -> tuple[dict[str, JsonValue], ...]: + return collected + tuple(JSON_OBJECT.validate_json(call.body) for call in edge.drain()) + + payload_batches: Final = accumulate(repeat(None), accumulate_policy_payloads, initial=()) + + def has_expected_post_kill_scans(collected: tuple[dict[str, JsonValue], ...], call: ChaosCall) -> bool: + payloads_for_call: Final = tuple( + payload for payload in collected if _policy_call_id_matches(payload, call.call_id) + ) + return all( + sum(_direction(payload) == direction for payload in payloads_for_call) + >= expected_directions.count(direction) + for direction in expected_directions + ) + + payloads: Final = eventually( + lambda: next(payload_batches), + lambda collected: all(has_expected_post_kill_scans(collected, call) for call in post_kill_calls), + seconds=30, + ) + assert all( + any(_policy_call_id_matches(payload, call_id) for call_id in requested_call_ids) + and _direction(payload) in expected_directions + for payload in payloads + ), payloads + for call, result in zip(post_kill_calls, post_kill_results): + payloads_for_call: Final = tuple( + payload for payload in payloads if _policy_call_id_matches(payload, call.call_id) + ) + assert tuple(sorted(_direction(payload) for payload in payloads_for_call)) == tuple( + sorted(expected_directions) + ), (call, payloads_for_call) + assert all( + payload["texts"] + == ([call.prompt] if _direction(payload) == "request" else [controls[call.index].text]) + for payload in payloads_for_call + ), payloads_for_call + + post_kill_expected: Final = tuple( + (call.model, result.response_id, call.call_id) + for call, result in zip(post_kill_calls, post_kill_results) + ) + rows: Final = _spend_rows_for_calls( + models, + post_kill_expected, + tolerate_missing=True, + ) + all_candidate_calls: Final = calls + post_kill_calls + rows_by_call: Final = tuple( + (call, _spend_rows_matching_call(rows, call.model, call.call_id)) for call in all_candidate_calls + ) + assert all(len(matching_rows) <= 1 for _, matching_rows in rows_by_call), rows_by_call + for call, matching_rows in rows_by_call: + if not matching_rows: + continue + entries: Final = _guardrail_entries(matching_rows[0]) + assert tuple( + sorted( + ( + entry["guardrail_name"], + entry["guardrail_mode"], + entry["guardrail_status"], + ) + for entry in entries + ) + ) == tuple(sorted(expected_entries)), (call, entries) + for call, result in successful: + matching_rows: Final = _spend_rows_matching_call(rows, call.model, call.call_id) + if matching_rows: + _assert_response_id( + call.endpoint, + str(matching_rows[0]["request_id"]), + result.response_id, + marker if call.endpoint == "responses" else call.call_id, + ) + for call, result in zip(post_kill_calls, post_kill_results): + matching_rows: Final = _spend_rows_matching_call(rows, call.model, call.call_id) + assert len(matching_rows) == 1, (result.response_id, matching_rows) + _assert_response_id( + call.endpoint, + str(matching_rows[0]["request_id"]), + result.response_id, + marker if call.endpoint == "responses" else call.call_id, + ) + + +@pytest.mark.timeout(180) +def test_K4_proxy_restart_after_fifteen_responses_records_lost_ids( + gateway: Gateway, + tmp_path: Path, + record_property: Callable[[str, object], None], +) -> None: + identity: Final = f"logging-scope-k4-{uuid.uuid4().hex}" + marker: Final = uuid.uuid4().hex + + def policy(_request: Request) -> Reply: + return Reply(body=b'{"action":"NONE"}') + + with gateway.scenario() as scenario: + deployments: Final = _chaos_models(scenario, marker) + models: Final = tuple(deployment.model_name for deployment in deployments) + model_list: Final = _chaos_model_list(deployments) + calls: Final = _chaos_calls(deployments, marker) + expected_directions: Final = _directions_for_audit_leg(("request", "response"), "output") + control_config: Final = _chaos_control_configuration(tmp_path, identity, model_list) + with owned_proxy(gateway, tmp_path, {}, config=control_config, workers=2) as control_proxy: + controls: Final = tuple( + _call_client( + call.client_kind, + call.endpoint, + control_proxy, + call.model, + call.prompt, + call.stream, + f"{call.call_id}-control", + ) + for call in calls + ) + control_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(control_upstream) == 30, control_upstream + assert ( + tuple( + sum(_json_contains_exact_string(observation["body"], call.prompt) for observation in control_upstream) + for call in calls + ) + == (1,) * 30 + ), control_upstream + with wire_server(policy) as edge: + config: Final = _configuration(tmp_path, identity, edge.url, "output", model_list=model_list) + restart_gate: Final = threading.Event() + second_wave_ready: Final = threading.Event() + second_wave_barrier: Final = threading.Barrier(15, action=second_wave_ready.set) + restarted_gateways: Final[SimpleQueue[Gateway]] = SimpleQueue() + + def gateway_for_call(call: ChaosCall, first_gateway: Gateway) -> Gateway: + if call.index < 15: + return first_gateway + second_wave_barrier.wait(timeout=90) + assert restart_gate.wait(timeout=90) + return restarted_gateways.get() + + def run(call: ChaosCall, first_gateway: Gateway) -> tuple[int, CallerResult]: + candidate: Final = gateway_for_call(call, first_gateway) + result: Final = _call_client( + call.client_kind, + call.endpoint, + candidate, + call.model, + call.prompt, + call.stream, + call.call_id, + ) + return call.index, result + + with ThreadPoolExecutor(max_workers=30) as pool: + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as first_proxy: + first_proxy_port: Final = first_proxy.gateway.client.base_url.port + assert first_proxy_port is not None + futures: Final = tuple(pool.submit(run, call, first_proxy.gateway) for call in calls) + first_results: Final = tuple(futures[index].result(timeout=90)[1] for index in range(15)) + assert all( + result.status == controls[call.index].status and result.text == controls[call.index].text + for call, result in zip(calls[:15], first_results) + ), first_results + assert eventually(lambda: second_wave_ready.is_set(), bool, seconds=30) + eventually( + lambda: edge.received.qsize(), + lambda count: count == 15 * len(expected_directions), + seconds=30, + ) + first_wave_expected: Final = tuple( + (call.model, result.response_id, call.call_id) + for call, result in zip(calls[:15], first_results) + ) + first_wave_rows: Final = _spend_rows_for_calls( + models, + first_wave_expected, + tolerate_missing=True, + ) + assert all( + len(_spend_rows_matching_call(first_wave_rows, model, call_id)) <= 1 + for model, _, call_id in first_wave_expected + ), first_wave_rows + first_wave_present_response_ids: Final = frozenset( + response_id + for model, response_id, call_id in first_wave_expected + if len(_spend_rows_matching_call(first_wave_rows, model, call_id)) == 1 + ) + first_wave_lost_response_ids: Final = ( + frozenset(response_id for _, response_id, _ in first_wave_expected) + - first_wave_present_response_ids + ) + record_property( + "K4_PRE_RESTART_LOST_RESPONSE_IDS", + tuple(sorted(first_wave_lost_response_ids)), + ) + record_property( + f"K4_PRE_RESTART_LOST_ROW_COUNT_{'base' if _is_base_audit_leg() else 'head'}", + len(first_wave_lost_response_ids), + ) + assert first_proxy.process.poll() is not None, first_proxy.process.pid + eventually( + lambda: tuple( + connection + for connection in psutil.net_connections(kind="tcp") + if connection.status == psutil.CONN_LISTEN and connection.laddr.port == first_proxy_port + ), + lambda listeners: not listeners, + seconds=30, + ) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as restarted_proxy: + for _ in range(15): + restarted_gateways.put(restarted_proxy.gateway) + restart_gate.set() + second_results: Final = tuple(futures[index].result(timeout=90)[1] for index in range(15, 30)) + assert all( + result.status == controls[call.index].status and result.text == controls[call.index].text + for call, result in zip(calls[15:], second_results) + ), second_results + eventually( + lambda: edge.received.qsize(), + lambda count: count == 30 * len(expected_directions), + seconds=30, + ) + results: Final = first_results + second_results + expected_response_ids: Final = frozenset(result.response_id for result in results) + assert len(expected_response_ids) == 30, results + upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(upstream) == 30, upstream + assert ( + tuple( + sum(_json_contains_exact_string(observation["body"], call.prompt) for observation in upstream) + for call in calls + ) + == (1,) * 30 + ), upstream + edge_payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in edge.drain()) + assert len(edge_payloads) == 30 * len(expected_directions), edge_payloads + for call, result in zip(calls, results): + payloads_for_call: Final = tuple( + payload for payload in edge_payloads if _policy_call_id_matches(payload, call.call_id) + ) + assert tuple(sorted(_direction(payload) for payload in payloads_for_call)) == tuple( + sorted(expected_directions) + ), ( + call, + payloads_for_call, + ) + assert all( + payload["texts"] == ([call.prompt] if _direction(payload) == "request" else [result.text]) + for payload in payloads_for_call + ), payloads_for_call + post_restart_expected: Final = tuple( + (call.model, result.response_id, call.call_id) for call, result in zip(calls[15:], second_results) + ) + rows: Final = _spend_rows_for_calls( + models, + post_restart_expected, + tolerate_missing=True, + ) + record_property("K4_RESPONSE_IDS", tuple(sorted(expected_response_ids))) + for call, result in zip(calls, results): + matching_rows: Final = _spend_rows_matching_call(rows, call.model, call.call_id) + assert len(matching_rows) <= 1, (result.response_id, matching_rows) + if call.index >= 15: + assert len(matching_rows) == 1, (call.call_id, result.response_id, matching_rows) + if matching_rows: + entries: Final = _guardrail_entries(matching_rows[0]) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == tuple((identity, "logging_only", "success") for _ in expected_directions), ( + call.call_id, + entries, + ) diff --git a/tests/integration/observability/test_logging_only_scope_config.py b/tests/integration/observability/test_logging_only_scope_config.py new file mode 100644 index 00000000000..f2445c3a3a5 --- /dev/null +++ b/tests/integration/observability/test_logging_only_scope_config.py @@ -0,0 +1,1591 @@ +from __future__ import annotations + +import json +import threading +import uuid +from collections.abc import Callable +from pathlib import Path +from typing import Final +from urllib.parse import parse_qs + +import pytest +from _logging_only_scope_support import ( + JSON_OBJECT, + Direction, + _assert_response_id, + _call_client, + _configuration, + _content_filter_configuration, + _create_guardrail, + _delete_database_guardrail, + _direction, + _directions_for_audit_leg, + _drain_upstream, + _empty_proxy_configuration, + _guardrail_entries, + _guardrail_mode_values, + _insert_database_guardrail, + _is_base_audit_leg, + _management_guardrail_rows, + _model_armor_configuration, + _policy_call_id, + _policy_call_id_matches, + _post_guardrail_body, + _presidio_configuration, + _provider_response, + _response_body_without_ids, + _spend_row_for_call_id, + _spend_rows, + wire_server, +) +from _logging_only_scope_support import ( + _record_audit_properties as _record_audit_properties, +) +from integration._support.client import Gateway, eventually, object_value, string_value +from integration._support.database import read_rows, write_rows +from integration._support.process import owned_proxy, owned_proxy_process +from integration._support.upstream import delete_scenario, register_scenario +from integration._support.wire import Reply, Request +from pydantic import JsonValue + + +def test_G7_database_invalid_scope_reads_preserve_pre_call_blocking(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-g7-{uuid.uuid4().hex}" + blocked_word: Final = f"pineapple{uuid.uuid4().hex[:8]}" + prompt: Final = f"synthetic request containing {blocked_word}" + reply: Final = f"synthetic upstream response {identity}" + scenario_id: Final = f"phase12-g7-{uuid.uuid4().hex}" + guardrail_id: Final = str(uuid.uuid5(uuid.NAMESPACE_URL, identity)) + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + write_rows( + 'INSERT INTO "LiteLLM_GuardrailsTable" ' + "(guardrail_id, guardrail_name, litellm_params, guardrail_info, updated_at) " + "VALUES (%s, %s, %s::jsonb, %s::jsonb, NOW())", + ( + guardrail_id, + identity, + json.dumps( + { + "guardrail": "litellm_content_filter", + "mode": "pre_call", + "logging_only_scope": "Input", + "default_on": True, + "blocked_words": [{"keyword": blocked_word, "action": "BLOCK"}], + } + ), + "{}", + ), + ) + try: + config: Final = _empty_proxy_configuration(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + listing_response: Final = candidate.request("GET", "/v2/guardrails/list") + assert listing_response.status_code == 200, listing_response.text + listing: Final = JSON_OBJECT.validate_json(listing_response.content) + rows: Final = tuple(object_value(row) for row in listing["guardrails"]) + stored_guardrail: Final = next(row for row in rows if row.get("guardrail_id") == guardrail_id) + stored_params: Final = object_value(stored_guardrail["litellm_params"]) + expected_scope: Final = "Input" if _is_base_audit_leg() else None + assert stored_params.get("logging_only_scope") == expected_scope, stored_guardrail + assert stored_params["mode"] == "pre_call", stored_guardrail + assert stored_params["guardrail"] == "litellm_content_filter", stored_guardrail + + info_response: Final = candidate.request("GET", f"/guardrails/{guardrail_id}/info") + assert info_response.status_code == 200, info_response.text + info: Final = JSON_OBJECT.validate_json(info_response.content) + info_params: Final = object_value(info["litellm_params"]) + assert info_params.get("logging_only_scope") == expected_scope, info + + blocked_response: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + ) + assert blocked_response.status_code == 400, blocked_response.text + assert _drain_upstream(gateway.upstream_url) == () + finally: + _delete_database_guardrail(identity) + finally: + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize( + ("row_id", "scope", "blocked_side"), + ( + pytest.param("F1", "output", "response", id="F1-native-filter-response"), + pytest.param("F2", "input", "request", id="F2-native-filter-request"), + ), +) +def test_native_content_filter_scope_logs_without_blocking( + gateway: Gateway, tmp_path: Path, row_id: str, scope: str, blocked_side: str +) -> None: + identity: Final = f"logging-scope-{row_id.lower()}-{uuid.uuid4().hex}" + blocked_word: Final = f"pineapple{uuid.uuid4().hex[:6]}" + prompt: Final = ( + f"synthetic {blocked_word} request {identity}" + if blocked_side == "request" + else f"synthetic clean request {identity}" + ) + reply: Final = ( + f"synthetic {blocked_word} response {identity}" + if blocked_side == "response" + else f"synthetic clean response {identity}" + ) + scenario_id: Final = f"phase12-{row_id.lower()}-{uuid.uuid4().hex}" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + config: Final = _content_filter_configuration(tmp_path, identity, scope, blocked_word) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + baseline: Final = _call_client( + "openai_sync", "chat", gateway, model, prompt, False, f"{scenario_id}-baseline" + ) + assert baseline.status == 200 and baseline.text == reply, baseline + assert len(_drain_upstream(gateway.upstream_url)) == 1 + guarded: Final = _call_client( + "openai_sync", "chat", candidate, model, prompt, False, f"{scenario_id}-guarded" + ) + assert (guarded.status, _response_body_without_ids(guarded.body)) == ( + baseline.status, + _response_body_without_ids(baseline.body), + ), (guarded, baseline) + assert guarded.text == reply, guarded + observed_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(observed_upstream) == 1, observed_upstream + assert prompt in json.dumps(observed_upstream[0]["body"]), observed_upstream + rows: Final = _spend_rows(model, 2) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 1, rows + _assert_response_id("chat", str(guarded_rows[0]["request_id"]), guarded.response_id, scenario_id) + entries: Final = _guardrail_entries(guarded_rows[0]) + expected_directions: Final = _directions_for_audit_leg(("request", "response"), scope) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) for entry in entries + ) == tuple( + ( + identity, + "logging_only", + "guardrail_intervened" if direction == blocked_side else "success", + ) + for direction in expected_directions + ), entries + intervened_entries: Final = tuple( + entry for entry in entries if entry["guardrail_status"] == "guardrail_intervened" + ) + assert len(intervened_entries) == 1, entries + assert intervened_entries[0]["guardrail_response"] is not None, entries + assert intervened_entries[0]["guardrail_response"] == "REDACTED_BY_LITELM", entries + finally: + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize( + ("row_id", "scope"), + ( + pytest.param("F3", "output", id="F3-model-armor-response-only"), + pytest.param("F4", "input", id="F4-model-armor-request-only"), + ), +) +def test_model_armor_directional_scope_uses_real_service_account_oauth( + gateway: Gateway, + tmp_path: Path, + row_id: str, + scope: str, +) -> None: + identity: Final = f"logging-scope-{row_id.lower()}-{uuid.uuid4().hex}" + prompt: Final = f"synthetic Model Armor prompt {identity}" + reply: Final = f"synthetic Model Armor response {identity}" + scenario_id: Final = f"phase12-{row_id.lower()}-{uuid.uuid4().hex}" + expected_directions: Final = _directions_for_audit_leg(("request", "response"), scope) + directions_to_collect: Final = _directions_for_audit_leg(("request", "response"), scope) + expected_fields: Final = tuple( + "userPromptData" if direction == "request" else "modelResponseData" for direction in expected_directions + ) + fields_to_collect: Final = tuple( + "userPromptData" if direction == "request" else "modelResponseData" for direction in directions_to_collect + ) + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def oauth(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/token", request + form: Final = parse_qs(request.body.decode()) + assert form.get("grant_type") == ["urn:ietf:params:oauth:grant-type:jwt-bearer"], form + assertion: Final = form.get("assertion") + assert assertion is not None and len(assertion[0].split(".")) == 3, form + return Reply(body=b'{"access_token":"synthetic-model-armor-token","expires_in":3600,"token_type":"Bearer"}') + + def model_armor(request: Request) -> Reply: + assert request.method == "POST", request + assert request.headers.get("authorization") == "Bearer synthetic-model-armor-token", request.headers + payload: Final = JSON_OBJECT.validate_json(request.body) + assert len(payload) == 1, payload + field: Final = next(iter(payload)) + assert field in ("userPromptData", "modelResponseData"), payload + expected_text: Final = prompt if field == "userPromptData" else reply + assert payload[field] == {"text": expected_text}, payload + return Reply(body=b'{"sanitizationResult":{"filterMatchState":"NO_MATCH_FOUND"}}') + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + baseline: Final = _call_client( + "openai_sync", "chat", gateway, model, prompt, False, f"{scenario_id}-baseline" + ) + assert baseline.status == 200 and baseline.text == reply, baseline + assert len(_drain_upstream(gateway.upstream_url)) == 1 + with wire_server(oauth, policy_edge=False) as token_edge, wire_server(model_armor) as armor_edge: + config: Final = _model_armor_configuration( + tmp_path, identity, armor_edge.url, token_edge.url + "/token", scope + ) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + result: Final = _call_client( + "openai_sync", "chat", candidate, model, prompt, False, f"{scenario_id}-candidate" + ) + assert (result.status, _response_body_without_ids(result.body)) == ( + baseline.status, + _response_body_without_ids(baseline.body), + ), (result, baseline) + assert result.text == reply, result + upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(upstream) == 1 and prompt in json.dumps(upstream[0]["body"]), upstream + eventually( + lambda: armor_edge.received.qsize(), + lambda count: count >= len(fields_to_collect), + seconds=30, + ) + eventually( + lambda: token_edge.received.qsize(), + lambda count: count >= 1, + seconds=30, + ) + armor_calls: Final = armor_edge.drain() + assert len(armor_calls) == len(fields_to_collect), armor_calls + armor_payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in armor_calls) + observed_fields: Final = tuple(next(iter(payload)) for payload in armor_payloads) + assert tuple(sorted(observed_fields)) == tuple(sorted(fields_to_collect)), armor_calls + assert all( + call.target.endswith( + ":sanitizeUserPrompt" if field == "userPromptData" else ":sanitizeModelResponse" + ) + for call, field in zip(armor_calls, observed_fields) + ), armor_calls + token_calls: Final = token_edge.drain() + assert token_calls and all(call.target == "/token" for call in token_calls), token_calls + rows: Final = _spend_rows(model, 2) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 1, rows + _assert_response_id("chat", str(guarded_rows[0]["request_id"]), result.response_id, scenario_id) + entries: Final = _guardrail_entries(guarded_rows[0]) + observed_entries: Final = tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) + expected_entries: Final = tuple((identity, "logging_only", "success") for _ in expected_directions) + assert ( + tuple(sorted(observed_fields)), + tuple(sorted(observed_entries)), + ) == (tuple(sorted(expected_fields)), tuple(sorted(expected_entries))), ( + armor_calls, + entries, + rows, + ) + finally: + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize( + ("row_id", "scope"), + ( + pytest.param("F5", "input", id="F5-presidio-input-scope-ignored"), + pytest.param("F6", "both", id="F6-presidio-both-scope-ignored"), + ), +) +def test_presidio_scope_matches_no_scope_behavior(gateway: Gateway, tmp_path: Path, row_id: str, scope: str) -> None: + identity: Final = f"logging-scope-{row_id.lower()}-{uuid.uuid4().hex}" + scenario_id: Final = f"phase12-{row_id.lower()}-{uuid.uuid4().hex}" + person: Final = f"synthetic Person {identity}" + prompt: Final = f"synthetic Presidio prompt {person}" + reply: Final = f"synthetic Presidio response {person}" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def analyzer(response_seen: threading.Event) -> Callable[[Request], Reply]: + def handle(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/analyze", request + payload: Final = JSON_OBJECT.validate_json(request.body) + text: Final = payload["text"] + assert isinstance(text, str), payload + if reply in text: + response_seen.set() + start: Final = text.index(person) + return Reply( + body=json.dumps( + [{"entity_type": "PERSON", "start": start, "end": start + len(person), "score": 0.99}] + ).encode() + ) + + return handle + + def anonymizer(response_seen: threading.Event) -> Callable[[Request], Reply]: + def handle(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/anonymize", request + payload: Final = JSON_OBJECT.validate_json(request.body) + text: Final = payload["text"] + results: Final = payload["analyzer_results"] + assert isinstance(text, str) and isinstance(results, list) and len(results) == 1, payload + if reply in text: + response_seen.set() + return Reply(body=json.dumps({"text": text, "items": [{"entity_type": "PERSON"}]}).encode()) + + return handle + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + no_scope_analyzer_response_seen: Final = threading.Event() + no_scope_anonymizer_response_seen: Final = threading.Event() + scoped_analyzer_response_seen: Final = threading.Event() + scoped_anonymizer_response_seen: Final = threading.Event() + with ( + wire_server(analyzer(no_scope_analyzer_response_seen)) as no_scope_analyzer, + wire_server(anonymizer(no_scope_anonymizer_response_seen)) as no_scope_anonymizer, + wire_server(analyzer(scoped_analyzer_response_seen)) as scoped_analyzer, + wire_server(anonymizer(scoped_anonymizer_response_seen)) as scoped_anonymizer, + ): + no_scope_config: Final = _presidio_configuration( + tmp_path, identity, no_scope_analyzer.url, no_scope_anonymizer.url, None + ) + scoped_config: Final = _presidio_configuration( + tmp_path, identity, scoped_analyzer.url, scoped_anonymizer.url, scope + ) + with owned_proxy_process(gateway, tmp_path, {}, config=no_scope_config, workers=1) as no_scope_owned: + no_scope_proxy: Final = no_scope_owned.gateway + no_scope_guardrails: Final = no_scope_proxy.get("/v2/guardrails/list")["guardrails"] + assert identity in {object_value(item)["guardrail_name"] for item in no_scope_guardrails}, ( + no_scope_guardrails + ) + no_scope_result: Final = _call_client( + "openai_sync", "chat", no_scope_proxy, model, prompt, False, f"{scenario_id}-no-scope" + ) + assert no_scope_result.status == 200 and no_scope_result.text == reply, no_scope_result + no_scope_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(no_scope_upstream) == 1 and prompt in json.dumps(no_scope_upstream[0]["body"]), ( + no_scope_upstream + ) + eventually( + lambda: no_scope_analyzer_response_seen.is_set() and no_scope_anonymizer_response_seen.is_set(), + bool, + seconds=30, + ) + no_scope_analyzer_calls: Final = no_scope_analyzer.drain() + no_scope_anonymizer_calls: Final = no_scope_anonymizer.drain() + no_scope_call_counts: Final = ( + len(no_scope_analyzer_calls), + len(no_scope_anonymizer_calls), + ) + assert no_scope_call_counts[0] == no_scope_call_counts[1] > 0, no_scope_call_counts + with owned_proxy_process(gateway, tmp_path, {}, config=scoped_config, workers=1) as scoped_owned: + scoped_proxy: Final = scoped_owned.gateway + guardrails: Final = scoped_proxy.get("/v2/guardrails/list")["guardrails"] + assert identity in {object_value(item)["guardrail_name"] for item in guardrails}, guardrails + scoped_result: Final = _call_client( + "openai_sync", "chat", scoped_proxy, model, prompt, False, f"{scenario_id}-scoped" + ) + assert (scoped_result.status, _response_body_without_ids(scoped_result.body)) == ( + no_scope_result.status, + _response_body_without_ids(no_scope_result.body), + ), (scoped_result, no_scope_result) + scoped_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(scoped_upstream) == 1 and prompt in json.dumps(scoped_upstream[0]["body"]), ( + scoped_upstream + ) + eventually( + lambda: scoped_analyzer_response_seen.is_set() and scoped_anonymizer_response_seen.is_set(), + bool, + seconds=30, + ) + scoped_analyzer_calls: Final = scoped_analyzer.drain() + scoped_anonymizer_calls: Final = scoped_anonymizer.drain() + scoped_call_counts: Final = ( + len(scoped_analyzer_calls), + len(scoped_anonymizer_calls), + ) + assert scoped_call_counts == no_scope_call_counts, ( + scoped_call_counts, + no_scope_call_counts, + ) + assert tuple(sorted((call.method, call.target, call.body) for call in scoped_analyzer_calls)) == ( + tuple(sorted((call.method, call.target, call.body) for call in no_scope_analyzer_calls)) + ), (scoped_analyzer_calls, no_scope_analyzer_calls) + assert tuple( + sorted((call.method, call.target, call.body) for call in scoped_anonymizer_calls) + ) == tuple(sorted((call.method, call.target, call.body) for call in no_scope_anonymizer_calls)), ( + scoped_anonymizer_calls, + no_scope_anonymizer_calls, + ) + log_text: Final = scoped_owned.log.read_text() + if row_id == "F5" and not _is_base_audit_leg(): + assert "whose logging_only hook scans on its own" in log_text, log_text + if row_id == "F6": + assert "whose logging_only hook scans on its own" not in log_text, log_text + finally: + delete_scenario(upstream_handle) + + +def test_F7_guardrail_ui_settings_classify_directional_scope_support(gateway: Gateway, tmp_path: Path) -> None: + config: Final = _empty_proxy_configuration(tmp_path, f"logging-scope-f7-{uuid.uuid4().hex}") + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + response: Final = candidate.get("/guardrails/ui/add_guardrail_settings") + if _is_base_audit_leg(): + assert "providers_without_directional_logging_only_scope" not in response, response + return + unsupported: Final = response.get("providers_without_directional_logging_only_scope") + assert unsupported is not None, response + assert isinstance(unsupported, list), response + assert set(unsupported) == { + "lakera", + "lakera_v2", + "presidio", + "tool_permission", + "cisco_ai_defense", + "xecguard", + "repelloai", + "noma", + "mcp_jwt_signer", + "microsoft_purview", + "agent_365", + "guardrails_ai", + "mcp_security", + "conduct", + "javelin", + "pillar", + "lasso", + "dynamoai", + "pangea", + "aporia", + "aim", + "ibm_guardrails", + "semantic_guard", + "cato_networks", + }, response + assert not {"generic_guardrail_api", "litellm_content_filter", "model_armor"}.intersection(unsupported), response + + +@pytest.mark.parametrize( + ("row_id", "mode", "scope", "expected_status", "expected_directions", "expected_guardrail_mode"), + ( + pytest.param("G1", "pre_call", "input", 400, ("request",), "pre_call", id="G1-yaml-blocking-valid-scope"), + pytest.param("G2", "pre_call", "Input", 400, ("request",), "pre_call", id="G2-yaml-blocking-invalid-literal"), + pytest.param( + "G3", + "logging_only", + "sideways", + 200, + ("request",), + "logging_only", + id="G3-yaml-logging-invalid-literal", + ), + ), +) +def test_yaml_scope_loading_keeps_guardrail_enforcement( + gateway: Gateway, + tmp_path: Path, + row_id: str, + mode: str, + scope: str, + expected_status: int, + expected_directions: tuple[Direction, ...], + expected_guardrail_mode: str, +) -> None: + identity: Final = f"logging-scope-{row_id.lower()}-{uuid.uuid4().hex}" + prompt: Final = f"synthetic yaml blocking marker {identity}" + reply: Final = f"synthetic yaml response {identity}" + scenario_id: Final = f"phase12-{row_id.lower()}-{uuid.uuid4().hex}" + call_id: Final = f"{scenario_id}-guarded" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"BLOCKED","blocked_reason":"synthetic yaml denial"}') + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + with wire_server(policy) as guardrail: + config: Final = _configuration(tmp_path, identity, guardrail.url, scope, mode=mode) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + baseline: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": f"{scenario_id}-baseline"}, + ) + assert baseline.status_code == 200, baseline.text + assert baseline.json()["choices"][0]["message"]["content"] == reply, baseline.text + assert len(_drain_upstream(gateway.upstream_url)) == 1 + guarded: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": call_id}, + ) + assert guarded.status_code == expected_status, guarded.text + caller_body: Final = JSON_OBJECT.validate_python(guarded.json()) + candidate_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(candidate_upstream) == (0 if expected_status == 400 else 1), candidate_upstream + if expected_status == 200: + assert _response_body_without_ids(caller_body) == _response_body_without_ids( + JSON_OBJECT.validate_python(baseline.json()) + ), (guarded.text, baseline.text) + assert prompt in json.dumps(candidate_upstream[0]["body"]), candidate_upstream + else: + assert "synthetic yaml denial" in guarded.text, guarded.text + if expected_status == 200: + eventually( + lambda: guardrail.received.qsize(), + lambda count: count >= len(expected_directions), + seconds=30, + ) + payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert tuple(sorted(_direction(payload) for payload in payloads)) == tuple( + sorted(expected_directions) + ), payloads + assert all(_policy_call_id_matches(payload, call_id) for payload in payloads), payloads + assert all( + payload["texts"] == ([prompt] if _direction(payload) == "request" else [reply]) + for payload in payloads + ), payloads + rows: Final = _spend_rows(model, 2) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 1, rows + entries: Final = _guardrail_entries(guarded_rows[0]) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == tuple( + (identity, expected_guardrail_mode, "guardrail_intervened") for _ in expected_directions + ), entries + if expected_status == 400: + blocked_row: Final = _spend_row_for_call_id(call_id) + assert _guardrail_entries(blocked_row) == entries, blocked_row + else: + _assert_response_id( + "chat", + str(guarded_rows[0]["request_id"]), + str(caller_body["id"]), + scenario_id, + ) + finally: + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize( + ("row_id", "scope"), + ( + pytest.param("G4", "output", id="G4-database-invalid-combination-before-boot"), + pytest.param("G5", "sideways", id="G5-database-invalid-literal-before-boot"), + ), +) +def test_database_guardrail_load_keeps_pre_call_blocking( + gateway: Gateway, tmp_path: Path, row_id: str, scope: str +) -> None: + identity: Final = f"logging-scope-{row_id.lower()}-{uuid.uuid4().hex}" + prompt: Final = f"synthetic DB blocked marker {identity}" + reply: Final = f"synthetic DB response {identity}" + scenario_id: Final = f"phase12-{row_id.lower()}-{uuid.uuid4().hex}" + call_id: Final = f"{scenario_id}-guarded" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"BLOCKED","blocked_reason":"synthetic DB denial"}') + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + baseline: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": f"{scenario_id}-baseline"}, + ) + assert baseline.status_code == 200, baseline.text + assert len(_drain_upstream(gateway.upstream_url)) == 1 + with wire_server(policy) as guardrail: + _insert_database_guardrail(identity, guardrail.url, scope) + try: + config: Final = _empty_proxy_configuration(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + loaded: Final = read_rows( + 'SELECT guardrail_name, litellm_params FROM "LiteLLM_GuardrailsTable" ' + "WHERE guardrail_name=%s", + (identity,), + ) + assert len(loaded) == 1, loaded + guarded: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": call_id}, + ) + assert guarded.status_code == 400 and "synthetic DB denial" in guarded.text, guarded.text + assert _drain_upstream(gateway.upstream_url) == () + payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert tuple(_direction(payload) for payload in payloads) == ("request",), payloads + assert _policy_call_id_matches(payloads[0], call_id), payloads + assert payloads[0]["texts"] == [prompt], payloads + spend_row: Final = _spend_row_for_call_id(call_id) + entries: Final = _guardrail_entries(spend_row) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == ((identity, "pre_call", "guardrail_intervened"),), entries + finally: + _delete_database_guardrail(identity) + finally: + delete_scenario(upstream_handle) + + +def test_G6_database_guardrail_polling_normalizes_invalid_scope_without_reinitializing( + gateway: Gateway, tmp_path: Path +) -> None: + identity: Final = f"logging-scope-g6-{uuid.uuid4().hex}" + prompt: Final = f"synthetic DB polling marker {identity}" + reply: Final = f"synthetic DB polling response {identity}" + scenario_id: Final = f"phase12-g6-{uuid.uuid4().hex}" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def denial(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"BLOCKED","blocked_reason":"synthetic polling denial"}') + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + baseline: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": f"{scenario_id}-baseline"}, + ) + assert baseline.status_code == 200, baseline.text + assert len(_drain_upstream(gateway.upstream_url)) == 1 + with wire_server(denial) as original_policy, wire_server(denial) as updated_policy: + _insert_database_guardrail(identity, original_policy.url, None) + try: + config: Final = _empty_proxy_configuration(tmp_path, identity, reload_seconds=1) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + first_call_id: Final = f"{scenario_id}-before-update" + first: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": first_call_id}, + ) + assert first.status_code == 400 and "synthetic polling denial" in first.text, first.text + first_policy_calls: Final = original_policy.drain() + assert len(first_policy_calls) == 1, first_policy_calls + updated_params: Final = { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": True, + "api_base": updated_policy.url, + "api_key": "synthetic-guardrail-key", + "extra_headers": ["x-litellm-call-id"], + "logging_only_scope": "sideways", + } + write_rows( + 'UPDATE "LiteLLM_GuardrailsTable" SET litellm_params=%s::jsonb, updated_at=NOW() ' + "WHERE guardrail_name=%s", + (json.dumps(updated_params), identity), + ) + + def probe_updated_policy() -> tuple[str, int]: + call_id: Final = f"{scenario_id}-poll-{uuid.uuid4().hex}" + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": call_id}, + ) + assert response.status_code == 400 and "synthetic polling denial" in response.text, ( + response.text + ) + return call_id, updated_policy.received.qsize() + + observed_call_id: Final = eventually( + probe_updated_policy, + lambda result: result[1] >= 1, + seconds=30, + )[0] + final_call_id: Final = f"{scenario_id}-after-sync" + final: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": final_call_id}, + ) + assert final.status_code == 400 and "synthetic polling denial" in final.text, final.text + old_payloads: Final = tuple( + JSON_OBJECT.validate_json(call.body) + for call in (*first_policy_calls, *original_policy.drain()) + ) + new_payloads: Final = tuple( + JSON_OBJECT.validate_json(call.body) for call in updated_policy.drain() + ) + assert all(payload["texts"] == [prompt] for payload in (*old_payloads, *new_payloads)), ( + old_payloads, + new_payloads, + ) + all_call_ids: Final = tuple( + str(_policy_call_id(payload)) for payload in (*old_payloads, *new_payloads) + ) + assert len(all_call_ids) == len(set(all_call_ids)), all_call_ids + assert first_call_id in all_call_ids and observed_call_id in all_call_ids, all_call_ids + assert final_call_id in all_call_ids, all_call_ids + assert len(new_payloads) >= 2, new_payloads + for call_id in (first_call_id, observed_call_id, final_call_id): + spend_row: Final = _spend_row_for_call_id(call_id) + entries: Final = _guardrail_entries(spend_row) + assert tuple( + ( + entry["guardrail_name"], + entry["guardrail_mode"], + entry["guardrail_status"], + ) + for entry in entries + ) == ((identity, "pre_call", "guardrail_intervened"),), (call_id, entries) + assert _drain_upstream(gateway.upstream_url) == () + finally: + _delete_database_guardrail(identity) + finally: + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize( + ("row_id", "provider", "mode", "scope", "expected_status", "expected_scope", "expected_message"), + ( + pytest.param( + "H1", + "generic_guardrail_api", + "pre_call", + "input", + 400, + None, + "mode does not include logging_only", + id="H1-post-rejects-scope-outside-logging-only", + ), + pytest.param( + "H2", + "generic_guardrail_api", + "logging_only", + "sideways", + 422, + None, + "logging_only_scope", + id="H2-post-rejects-invalid-string", + ), + pytest.param( + "H2", + "generic_guardrail_api", + "logging_only", + 5, + 422, + None, + "logging_only_scope", + id="H2-post-rejects-invalid-number", + ), + pytest.param( + "H2", + "generic_guardrail_api", + "logging_only", + "", + 422, + None, + "logging_only_scope", + id="H2-post-rejects-invalid-empty-string", + ), + pytest.param( + "H2", + "generic_guardrail_api", + "logging_only", + ["input"], + 422, + None, + "logging_only_scope", + id="H2-post-rejects-invalid-list", + ), + pytest.param( + "H2", + "generic_guardrail_api", + "logging_only", + "x" * 5000, + 422, + None, + "logging_only_scope", + id="H2-post-rejects-oversized-string", + ), + pytest.param( + "H3", + "generic_guardrail_api", + "logging_only", + None, + 200, + None, + "", + id="H3-post-accepts-null-scope", + ), + pytest.param( + "H4", + "presidio", + "logging_only", + "input", + 400, + None, + "whose logging_only hook scans on its own", + id="H4-post-rejects-presidio-input-scope", + ), + pytest.param( + "H4", + "presidio", + "logging_only", + "both", + 200, + "both", + "", + id="H4-post-accepts-presidio-both-scope", + ), + ), +) +def test_management_post_validates_logging_only_scope( + gateway: Gateway, + tmp_path: Path, + row_id: str, + provider: str, + mode: str | list[str], + scope: JsonValue, + expected_status: int, + expected_scope: JsonValue, + expected_message: str, +) -> None: + identity: Final = f"logging-scope-{row_id.lower()}-{uuid.uuid4().hex}" + api_base: Final = "http://127.0.0.1:9" + config: Final = _empty_proxy_configuration(tmp_path, identity) + try: + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + before_rows: Final = _management_guardrail_rows(identity) + before_list: Final = candidate.get("/v2/guardrails/list")["guardrails"] + expected_leg_status: Final = 200 if _is_base_audit_leg() else expected_status + expected_leg_scope: Final = scope if _is_base_audit_leg() else expected_scope + response: Final = candidate.request( + "POST", + "/guardrails", + _post_guardrail_body(identity, provider, mode, api_base, scope), + ) + assert response.status_code == expected_leg_status, response.text + if expected_message and not _is_base_audit_leg(): + assert expected_message in response.text, response.text + if expected_leg_status == 200: + body: Final = JSON_OBJECT.validate_json(response.content) + params: Final = object_value(body["litellm_params"]) + assert params.get("logging_only_scope") == expected_leg_scope, body + rows: Final = _management_guardrail_rows(identity) + assert ( + len(rows) == 1 + and object_value(rows[0]["litellm_params"]).get("logging_only_scope") == expected_leg_scope + ), rows + else: + if row_id == "H2" and not _is_base_audit_leg(): + details: Final = object_value(JSON_OBJECT.validate_python(response.json())["detail"][0]) + assert details["type"] == "literal_error", details + assert details["loc"] == [ + "body", + "guardrail", + "litellm_params", + "logging_only_scope", + ], details + assert _management_guardrail_rows(identity) == before_rows == () + after_list: Final = candidate.get("/v2/guardrails/list")["guardrails"] + assert after_list == before_list, (before_list, after_list) + assert isinstance(after_list, list), after_list + assert identity not in {object_value(item)["guardrail_name"] for item in after_list}, after_list + finally: + _delete_database_guardrail(identity) + + +def test_H5_management_put_rejection_preserves_database_and_runtime(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-h5-{uuid.uuid4().hex}" + prompt: Final = f"synthetic management marker {identity}" + reply: Final = f"synthetic management response {identity}" + scenario_id: Final = f"phase12-h5-{uuid.uuid4().hex}" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"BLOCKED","blocked_reason":"synthetic management denial"}') + + try: + with wire_server(policy) as guardrail: + config: Final = _empty_proxy_configuration(tmp_path, identity) + with ( + owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + created: Final = _create_guardrail( + candidate, + identity, + { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": True, + "api_base": guardrail.url, + "api_key": "synthetic-guardrail-key", + }, + ) + guardrail_id: Final = string_value(created["guardrail_id"]) + before_rows: Final = _management_guardrail_rows(identity) + assert len(before_rows) == 1, before_rows + before_info: Final = candidate.get(f"/guardrails/{guardrail_id}/info") + before_list: Final = tuple( + object_value(item) + for item in candidate.get("/v2/guardrails/list")["guardrails"] + if object_value(item)["guardrail_id"] == guardrail_id + ) + before_call_id: Final = f"{scenario_id}-before-put" + before_response: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": before_call_id}, + ) + assert before_response.status_code == 400, before_response.text + before_policy: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert len(before_policy) == 1 and _policy_call_id_matches(before_policy[0], before_call_id), ( + before_policy + ) + assert before_policy[0]["texts"] == [prompt], before_policy + update: Final = candidate.request( + "PUT", + f"/guardrails/{guardrail_id}", + _post_guardrail_body( + identity, + "generic_guardrail_api", + "pre_call", + guardrail.url, + "output", + ), + ) + if _is_base_audit_leg(): + assert update.status_code == 200, update.text + updated_rows: Final = _management_guardrail_rows(identity) + assert len(updated_rows) == 1, updated_rows + assert object_value(updated_rows[0]["litellm_params"]).get("logging_only_scope") == "output", ( + updated_rows + ) + assert object_value(updated_rows[0]["litellm_params"])["mode"] == "pre_call", updated_rows + else: + assert update.status_code == 422, update.text + assert "logging_only_scope" in update.text and "logging_only" in update.text, update.text + assert _management_guardrail_rows(identity) == before_rows + after_info: Final = candidate.get(f"/guardrails/{guardrail_id}/info") + assert {key: value for key, value in before_info.items() if key != "updated_at"} == { + key: value for key, value in after_info.items() if key != "updated_at" + } + after_list: Final = tuple( + object_value(item) + for item in candidate.get("/v2/guardrails/list")["guardrails"] + if object_value(item)["guardrail_id"] == guardrail_id + ) + if _is_base_audit_leg(): + assert ( + len(after_list) == 1 + and object_value(after_list[0]["litellm_params"]).get("logging_only_scope") == "output" + ), after_list + else: + assert tuple( + {key: value for key, value in item.items() if key != "updated_at"} for item in after_list + ) == tuple( + {key: value for key, value in item.items() if key != "updated_at"} for item in before_list + ) + after_call_id: Final = f"{scenario_id}-after-put" + after_response: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": after_call_id}, + ) + assert after_response.status_code == before_response.status_code, after_response.text + assert after_response.text == before_response.text, (after_response.text, before_response.text) + after_policy: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert len(after_policy) == 1 and _policy_call_id_matches(after_policy[0], after_call_id), after_policy + assert after_policy[0]["texts"] == [prompt], after_policy + assert _drain_upstream(gateway.upstream_url) == () + for call_id in (before_call_id, after_call_id): + entries: Final = _guardrail_entries(_spend_row_for_call_id(call_id)) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == ((identity, "pre_call", "guardrail_intervened"),), (call_id, entries) + finally: + _delete_database_guardrail(identity) + delete_scenario(upstream_handle) + + +def test_H6_management_patch_clears_scope_when_switching_to_blocking_only_mode( + gateway: Gateway, tmp_path: Path +) -> None: + identity: Final = f"logging-scope-h6-{uuid.uuid4().hex}" + prompt: Final = f"synthetic mode patch marker {identity}" + scenario_id: Final = f"phase12-h6-{uuid.uuid4().hex}" + upstream_handle: Final = register_scenario( + scenario_id, _provider_response("chat", scenario_id, f"response {identity}", False) + ) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"BLOCKED","blocked_reason":"synthetic mode patch denial"}') + + try: + with wire_server(policy) as guardrail: + config: Final = _empty_proxy_configuration(tmp_path, identity) + with ( + owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + created: Final = _create_guardrail( + candidate, + identity, + { + "guardrail": "generic_guardrail_api", + "mode": ["pre_call", "logging_only"], + "default_on": True, + "api_base": guardrail.url, + "api_key": "synthetic-guardrail-key", + "logging_only_scope": "output", + }, + ) + guardrail_id: Final = string_value(created["guardrail_id"]) + before: Final = _management_guardrail_rows(identity) + assert len(before) == 1, before + patched: Final = candidate.request( + "PATCH", + f"/guardrails/{guardrail_id}", + {"litellm_params": {"mode": ["pre_call"]}}, + ) + assert patched.status_code == 200, patched.text + persisted: Final = _management_guardrail_rows(identity) + assert len(persisted) == 1, persisted + params: Final = object_value(persisted[0]["litellm_params"]) + assert params.get("logging_only_scope") == ("output" if _is_base_audit_leg() else None), persisted + assert params["mode"] == ["pre_call"], persisted + call_id: Final = f"{scenario_id}-after-patch" + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": call_id}, + ) + assert response.status_code == 400 and "synthetic mode patch denial" in response.text, response.text + calls: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert tuple(_direction(call) for call in calls) == ("request",), calls + assert _policy_call_id_matches(calls[0], call_id), calls + assert calls[0]["texts"] == [prompt], calls + assert _drain_upstream(gateway.upstream_url) == () + entries: Final = _guardrail_entries(_spend_row_for_call_id(call_id)) + assert tuple( + ( + entry["guardrail_name"], + _guardrail_mode_values(entry["guardrail_mode"]), + entry["guardrail_status"], + ) + for entry in entries + ) == ((identity, ("pre_call",), "guardrail_intervened"),), entries + finally: + _delete_database_guardrail(identity) + delete_scenario(upstream_handle) + + +def test_H7_management_patch_rejection_preserves_pre_call_guardrail(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-h7-{uuid.uuid4().hex}" + prompt: Final = f"synthetic invalid patch marker {identity}" + scenario_id: Final = f"phase12-h7-{uuid.uuid4().hex}" + upstream_handle: Final = register_scenario( + scenario_id, _provider_response("chat", scenario_id, f"response {identity}", False) + ) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"BLOCKED","blocked_reason":"synthetic invalid patch denial"}') + + try: + with wire_server(policy) as guardrail: + config: Final = _empty_proxy_configuration(tmp_path, identity) + with ( + owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + created: Final = _create_guardrail( + candidate, + identity, + { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": True, + "api_base": guardrail.url, + "api_key": "synthetic-guardrail-key", + }, + ) + guardrail_id: Final = string_value(created["guardrail_id"]) + before: Final = _management_guardrail_rows(identity) + before_call_id: Final = f"{scenario_id}-before-patch" + response_before: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": before_call_id}, + ) + assert response_before.status_code == 400, response_before.text + policy_before: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + rejected: Final = candidate.request( + "PATCH", + f"/guardrails/{guardrail_id}", + {"litellm_params": {"logging_only_scope": "input"}}, + ) + if _is_base_audit_leg(): + assert rejected.status_code == 200, rejected.text + updated: Final = _management_guardrail_rows(identity) + assert len(updated) == 1, updated + assert object_value(updated[0]["litellm_params"]).get("logging_only_scope") == "input", updated + else: + assert rejected.status_code == 422, rejected.text + assert "logging_only_scope" in rejected.text and "logging_only" in rejected.text, rejected.text + assert _management_guardrail_rows(identity) == before + after_call_id: Final = f"{scenario_id}-after-patch" + response_after: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": after_call_id}, + ) + assert response_after.status_code == response_before.status_code, response_after.text + policy_after: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert tuple(_direction(call) for call in policy_before) == ("request",), policy_before + assert tuple(_direction(call) for call in policy_after) == ("request",), policy_after + assert _policy_call_id_matches(policy_before[0], before_call_id), policy_before + assert _policy_call_id_matches(policy_after[0], after_call_id), policy_after + assert policy_before[0]["texts"] == [prompt], policy_before + assert policy_after[0]["texts"] == [prompt], policy_after + assert _drain_upstream(gateway.upstream_url) == () + for call_id in (before_call_id, after_call_id): + entries: Final = _guardrail_entries(_spend_row_for_call_id(call_id)) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == ((identity, "pre_call", "guardrail_intervened"),), (call_id, entries) + finally: + _delete_database_guardrail(identity) + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize( + ("row_id", "stored_scope"), + ( + pytest.param("H8", "output", id="H8-patch-heals-invalid-mode-combination"), + pytest.param("H9", "sideways", id="H9-patch-heals-invalid-stored-literal"), + ), +) +def test_management_patch_default_on_heals_stored_scope( + gateway: Gateway, tmp_path: Path, row_id: str, stored_scope: str +) -> None: + identity: Final = f"logging-scope-{row_id.lower()}-{uuid.uuid4().hex}" + prompt: Final = f"synthetic healed scope marker {identity}" + scenario_id: Final = f"phase12-{row_id.lower()}-{uuid.uuid4().hex}" + upstream_handle: Final = register_scenario( + scenario_id, _provider_response("chat", scenario_id, f"response {identity}", False) + ) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"BLOCKED","blocked_reason":"synthetic healed scope denial"}') + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + baseline: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": f"{scenario_id}-baseline"}, + ) + assert baseline.status_code == 200, baseline.text + assert len(_drain_upstream(gateway.upstream_url)) == 1 + with wire_server(policy) as guardrail: + _insert_database_guardrail( + identity, + guardrail.url, + stored_scope, + default_on=False, + ) + try: + config: Final = _empty_proxy_configuration(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + before: Final = _management_guardrail_rows(identity) + assert len(before) == 1, before + guardrail_id: Final = string_value(before[0]["guardrail_id"]) + patched: Final = candidate.request( + "PATCH", + f"/guardrails/{guardrail_id}", + {"litellm_params": {"default_on": True}}, + ) + assert patched.status_code == 200, patched.text + body: Final = JSON_OBJECT.validate_json(patched.content) + params: Final = object_value(body["litellm_params"]) + assert params["default_on"] is True, body + persisted: Final = _management_guardrail_rows(identity) + assert len(persisted) == 1, persisted + persisted_params: Final = object_value(persisted[0]["litellm_params"]) + assert persisted_params.get("logging_only_scope") == ( + stored_scope if _is_base_audit_leg() else None + ), persisted + assert persisted_params["default_on"] is True, persisted + call_id: Final = f"{scenario_id}-after-patch" + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": call_id}, + ) + assert response.status_code == 400 and "synthetic healed scope denial" in response.text, ( + response.text + ) + calls: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert tuple(_direction(call) for call in calls) == ("request",), calls + assert _policy_call_id_matches(calls[0], call_id), calls + assert calls[0]["texts"] == [prompt], calls + assert _drain_upstream(gateway.upstream_url) == () + entries: Final = _guardrail_entries(_spend_row_for_call_id(call_id)) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == ((identity, "pre_call", "guardrail_intervened"),), entries + finally: + _delete_database_guardrail(identity) + finally: + delete_scenario(upstream_handle) + + +def test_H10_management_patch_null_scope_restores_both_directions(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-h10-{uuid.uuid4().hex}" + prompt: Final = f"synthetic reset-scope prompt {identity}" + reply: Final = f"synthetic reset-scope response {identity}" + scenario_id: Final = f"phase12-h10-{uuid.uuid4().hex}" + control_call_id: Final = f"{scenario_id}-control" + guarded_call_id: Final = f"{scenario_id}-guarded" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + payload: Final = JSON_OBJECT.validate_json(request.body) + direction: Final = _direction(payload) + verdict: Final = ( + {"action": "NONE"} + if direction == "request" + else {"action": "BLOCKED", "blocked_reason": "synthetic reset-scope monitor"} + ) + return Reply(body=json.dumps(verdict).encode()) + + try: + with wire_server(policy) as guardrail: + config: Final = _empty_proxy_configuration(tmp_path, identity) + with ( + owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + control: Final = _call_client("openai_sync", "chat", candidate, model, prompt, False, control_call_id) + assert control.status == 200 and control.text == reply, control + assert len(_drain_upstream(gateway.upstream_url)) == 1 + created: Final = _create_guardrail( + candidate, + identity, + { + "guardrail": "generic_guardrail_api", + "mode": "logging_only", + "default_on": True, + "api_base": guardrail.url, + "api_key": "synthetic-guardrail-key", + "logging_only_scope": "output", + }, + ) + guardrail_id: Final = string_value(created["guardrail_id"]) + patched: Final = candidate.request( + "PATCH", + f"/guardrails/{guardrail_id}", + {"litellm_params": {"logging_only_scope": None}}, + ) + assert patched.status_code == 200, patched.text + persisted: Final = _management_guardrail_rows(identity) + assert len(persisted) == 1, persisted + params: Final = object_value(persisted[0]["litellm_params"]) + assert params.get("logging_only_scope") is None, persisted + result: Final = _call_client("openai_sync", "chat", candidate, model, prompt, False, guarded_call_id) + assert (result.status, _response_body_without_ids(result.body)) == ( + control.status, + _response_body_without_ids(control.body), + ), (result, control) + upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(upstream) == 1, upstream + assert prompt in json.dumps(upstream[0]["body"]), upstream + expected_directions: Final = ("request", "response") + eventually( + lambda: guardrail.received.qsize(), + lambda count: count == len(expected_directions), + seconds=20, + ) + policy_calls: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert tuple(sorted(_direction(call) for call in policy_calls)) == tuple(sorted(expected_directions)), ( + policy_calls + ) + assert all(_policy_call_id_matches(call, guarded_call_id) for call in policy_calls), policy_calls + assert all( + call["texts"] == ([prompt] if _direction(call) == "request" else [reply]) for call in policy_calls + ), policy_calls + rows: Final = _spend_rows(model, 2) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 1, rows + _assert_response_id("chat", str(guarded_rows[0]["request_id"]), result.response_id, scenario_id) + entries: Final = _guardrail_entries(guarded_rows[0]) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) for entry in entries + ) == ( + (identity, "logging_only", "success"), + (identity, "logging_only", "guardrail_intervened"), + ), entries + finally: + _delete_database_guardrail(identity) + delete_scenario(upstream_handle) + + +def test_H11_management_patch_same_scope_is_idempotent_and_output_only(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-h11-{uuid.uuid4().hex}" + prompt: Final = f"synthetic idempotent prompt {identity}" + reply: Final = f"synthetic idempotent response {identity}" + scenario_id: Final = f"phase12-h11-{uuid.uuid4().hex}" + control_call_id: Final = f"{scenario_id}-control" + guarded_call_id: Final = f"{scenario_id}-guarded" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"BLOCKED","blocked_reason":"synthetic idempotent monitor"}') + + try: + with wire_server(policy) as guardrail: + config: Final = _empty_proxy_configuration(tmp_path, identity) + with ( + owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + control: Final = _call_client("openai_sync", "chat", candidate, model, prompt, False, control_call_id) + assert control.status == 200 and control.text == reply, control + assert len(_drain_upstream(gateway.upstream_url)) == 1 + created: Final = _create_guardrail( + candidate, + identity, + { + "guardrail": "generic_guardrail_api", + "mode": "logging_only", + "default_on": True, + "api_base": guardrail.url, + "api_key": "synthetic-guardrail-key", + "logging_only_scope": "output", + }, + ) + guardrail_id: Final = string_value(created["guardrail_id"]) + expected_directions: Final = _directions_for_audit_leg(("request", "response"), "output") + before: Final = _management_guardrail_rows(identity) + first: Final = candidate.request( + "PATCH", + f"/guardrails/{guardrail_id}", + {"litellm_params": {"logging_only_scope": "output"}}, + ) + second: Final = candidate.request( + "PATCH", + f"/guardrails/{guardrail_id}", + {"litellm_params": {"logging_only_scope": "output"}}, + ) + assert first.status_code == 200, first.text + assert second.status_code == 200, second.text + assert _management_guardrail_rows(identity) == before + result: Final = _call_client("openai_sync", "chat", candidate, model, prompt, False, guarded_call_id) + assert (result.status, _response_body_without_ids(result.body)) == ( + control.status, + _response_body_without_ids(control.body), + ), (result, control) + upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(upstream) == 1 and prompt in json.dumps(upstream[0]["body"]), upstream + eventually( + lambda: guardrail.received.qsize(), + lambda count: count == len(expected_directions), + seconds=20, + ) + calls: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert tuple(sorted(_direction(call) for call in calls)) == tuple(sorted(expected_directions)), calls + assert all(_policy_call_id_matches(call, guarded_call_id) for call in calls), calls + assert all( + call["texts"] == ([prompt] if _direction(call) == "request" else [reply]) for call in calls + ), calls + rows: Final = _spend_rows(model, 2) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 1, rows + _assert_response_id("chat", str(guarded_rows[0]["request_id"]), result.response_id, scenario_id) + entries: Final = _guardrail_entries(guarded_rows[0]) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) for entry in entries + ) == tuple((identity, "logging_only", "guardrail_intervened") for _ in expected_directions), entries + finally: + _delete_database_guardrail(identity) + delete_scenario(upstream_handle) + + +def test_H12_management_reads_expose_typed_logging_only_scope(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-h12-{uuid.uuid4().hex}" + try: + with owned_proxy( + gateway, tmp_path, {}, config=_empty_proxy_configuration(tmp_path, identity), workers=1 + ) as candidate: + created: Final = _create_guardrail( + candidate, + identity, + { + "guardrail": "generic_guardrail_api", + "mode": "logging_only", + "default_on": True, + "api_base": "http://127.0.0.1:9", + "api_key": "synthetic-guardrail-key", + "logging_only_scope": "input", + }, + ) + guardrail_id: Final = string_value(created["guardrail_id"]) + info: Final = candidate.get(f"/guardrails/{guardrail_id}/info") + listed: Final = candidate.get("/v2/guardrails/list")["guardrails"] + assert isinstance(listed, list), listed + matches: Final = tuple( + object_value(item) for item in listed if object_value(item)["guardrail_id"] == guardrail_id + ) + assert len(matches) == 1, listed + assert object_value(info["litellm_params"])["logging_only_scope"] == "input", info + assert object_value(matches[0]["litellm_params"])["logging_only_scope"] == "input", matches + finally: + _delete_database_guardrail(identity) + + +def test_H13_unauthenticated_management_and_chat_requests_do_not_scan(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-h13-{uuid.uuid4().hex}" + prompt: Final = f"synthetic unauthorized prompt {identity}" + scenario_id: Final = f"phase12-h13-{uuid.uuid4().hex}" + upstream_handle: Final = register_scenario( + scenario_id, _provider_response("chat", scenario_id, f"response {identity}", False) + ) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"NONE"}') + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + with wire_server(policy) as guardrail: + config: Final = _configuration(tmp_path, identity, guardrail.url, "input") + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + unauthorized_post: Final = candidate.request( + "POST", + "/guardrails", + _post_guardrail_body( + f"{identity}-unauthorized", + "generic_guardrail_api", + "logging_only", + guardrail.url, + "input", + ), + key="synthetic-invalid-key", + ) + assert unauthorized_post.status_code == 401, unauthorized_post.text + unauthorized_chat: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + key="synthetic-invalid-key", + headers={"x-litellm-call-id": f"{scenario_id}-unauthorized"}, + ) + assert unauthorized_chat.status_code == 401, unauthorized_chat.text + assert _management_guardrail_rows(f"{identity}-unauthorized") == () + assert guardrail.drain() == () + assert _drain_upstream(gateway.upstream_url) == () + finally: + delete_scenario(upstream_handle) diff --git a/tests/integration/observability/test_logging_only_scope_runtime.py b/tests/integration/observability/test_logging_only_scope_runtime.py new file mode 100644 index 00000000000..dd00378401e --- /dev/null +++ b/tests/integration/observability/test_logging_only_scope_runtime.py @@ -0,0 +1,1579 @@ +from __future__ import annotations + +import json +import threading +import uuid +from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor +from datetime import datetime, timezone +from pathlib import Path +from typing import Final + +import pytest +import yaml +from _logging_only_scope_support import ( + BASE_DEFAULT_CACHE_HIT_DIRECTIONS, + BASE_DEFAULT_NORMAL_DIRECTIONS, + BASE_DEFAULT_UPSTREAM_FAILURE_DIRECTIONS, + JSON_OBJECT, + ChaosCall, + ClientKind, + Direction, + Endpoint, + _assert_response_id, + _cache_hit, + _call_cache_client, + _call_client, + _configuration, + _database_guardrail, + _direction, + _directions_for_audit_leg, + _directions_for_scope, + _drain_upstream, + _empty_proxy_configuration, + _guardrail_entries, + _guardrail_mode_status_pairs, + _is_base_audit_leg, + _json_contains_exact_string, + _policy_call_id, + _policy_call_id_matches, + _provider_response, + _response_body_without_ids, + _response_text, + _spend_row_for_call_id, + _spend_row_for_response_id, + _spend_rows, + _spend_rows_for_calls, + _spend_rows_matching_call, + wire_server, +) +from _logging_only_scope_support import ( + _record_audit_properties as _record_audit_properties, +) +from integration._support.client import Gateway, eventually, object_value, string_value +from integration._support.process import owned_proxy +from integration._support.upstream import delete_scenario, register_scenario +from integration._support.wire import Reply, Request +from pydantic import JsonValue + +from tests.integration.cost_calculation.cost_tracking_case import JsonResponse + + +@pytest.mark.parametrize( + ("row_id", "endpoint", "stream", "client_kind", "scope", "include_scope", "block_directions"), + ( + pytest.param("A1", "chat", False, "openai_sync", "input", True, (), id="A1-chat-input"), + pytest.param("A2", "chat", False, "openai_sync", "output", True, (), id="A2-chat-output"), + pytest.param("A3", "chat", False, "openai_sync", "both", True, (), id="A3-chat-both"), + pytest.param("A4", "chat", False, "httpx", None, False, (), id="A4-chat-missing-scope"), + pytest.param("A5", "chat", False, "httpx", None, True, (), id="A5-chat-null-scope"), + pytest.param("A6", "chat", True, "openai_async", "input", True, (), id="A6-chat-stream-async-input"), + pytest.param("A7", "chat", True, "openai_async", "output", True, (), id="A7-chat-stream-async-output"), + pytest.param("A8", "messages", False, "anthropic_sync", "input", True, (), id="A8-messages-input"), + pytest.param("A9", "messages", False, "anthropic_sync", "output", True, (), id="A9-messages-output"), + pytest.param( + "A10", + "messages", + True, + "anthropic_async", + "output", + True, + (), + id="A10-messages-stream-async-output", + ), + pytest.param( + "A11", + "messages", + True, + "anthropic_async", + "input", + True, + (), + id="A11-messages-stream-async-input", + ), + pytest.param("A12", "responses", False, "openai_async", "input", True, (), id="A12-responses-async-input"), + pytest.param("A13", "responses", False, "openai_async", "output", True, (), id="A13-responses-async-output"), + pytest.param("A14", "responses", True, "openai_sync", "output", True, (), id="A14-responses-stream-output"), + pytest.param("A15", "responses", True, "openai_sync", "input", True, (), id="A15-responses-stream-input"), + pytest.param("B1", "chat", False, "openai_sync", "input", True, ("request",), id="B1-logging-block-input"), + pytest.param("B2", "chat", False, "openai_sync", "output", True, ("response",), id="B2-logging-block-output"), + pytest.param( + "B3", + "chat", + False, + "openai_sync", + "both", + True, + ("request",), + id="B3-logging-block-both", + ), + ), +) +def test_runtime_directional_scope_matches_client_call_and_spend_log( + gateway: Gateway, + tmp_path: Path, + row_id: str, + endpoint: Endpoint, + stream: bool, + client_kind: ClientKind, + scope: str | None, + include_scope: bool, + block_directions: tuple[Direction, ...], +) -> None: + identity: Final = f"logging-scope-{row_id.lower()}-{uuid.uuid4().hex}" + prompt: Final = f"synthetic request {identity}" + reply: Final = f"synthetic response {identity}" + scenario_id: Final = f"phase12-{row_id.lower()}-{uuid.uuid4().hex}" + baseline_call_id: Final = f"{scenario_id}-baseline" + guarded_call_id: Final = f"{scenario_id}-guarded" + provider_response: Final = _provider_response(endpoint, scenario_id, reply, stream) + upstream_handle: Final = register_scenario(scenario_id, provider_response) + + def policy(request: Request) -> Reply: + payload: Final = JSON_OBJECT.validate_json(request.body) + direction: Final = _direction(payload) + verdict: Final = ( + {"action": "BLOCKED", "blocked_reason": "synthetic logging-only denial"} + if direction in block_directions + else {"action": "NONE"} + ) + return Reply(body=json.dumps(verdict).encode()) + + try: + with gateway.scenario() as scenario: + api_base: Final = ( + upstream_handle.api_base() if endpoint == "messages" else f"{upstream_handle.api_base()}/v1" + ) + model: Final = scenario.model( + model={ + "chat": "openai/gpt-4o-mini", + "messages": "anthropic/claude-3-7-sonnet-20250219", + "responses": "openai/gpt-4.1-mini", + }[endpoint], + api_base=api_base, + api_key="synthetic-provider-key", + ) + with wire_server(policy) as guardrail: + config: Final = _configuration(tmp_path, identity, guardrail.url, scope, include_scope=include_scope) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + baseline: Final = _call_client( + client_kind, endpoint, gateway, model, prompt, stream, baseline_call_id + ) + baseline_upstream: Final = _drain_upstream(gateway.upstream_url) + assert baseline.status == 200, baseline.body + assert baseline.text == reply, baseline + assert len(baseline_upstream) == 1, baseline_upstream + result: Final = _call_client( + client_kind, endpoint, candidate, model, prompt, stream, guarded_call_id + ) + assert (result.status, _response_body_without_ids(result.body)) == ( + baseline.status, + _response_body_without_ids(baseline.body), + ), (result, baseline) + assert result.text == reply, result + observed_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(observed_upstream) == 1, observed_upstream + assert prompt in json.dumps(observed_upstream[0]["body"]), observed_upstream + base_default_directions: Final = BASE_DEFAULT_NORMAL_DIRECTIONS[(endpoint, stream)] + expected_directions: Final = _directions_for_audit_leg(base_default_directions, scope) + directions_to_collect: Final = _directions_for_audit_leg(base_default_directions, scope) + eventually( + lambda: guardrail.received.qsize(), + lambda count: count >= len(directions_to_collect), + seconds=20, + ) + calls: Final = guardrail.drain() + payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in calls) + observed_directions: Final = tuple(_direction(payload) for payload in payloads) + assert all(_policy_call_id_matches(payload, guarded_call_id) for payload in payloads), payloads + assert tuple(payload["texts"] for payload in payloads) == tuple( + [prompt] if direction == "request" else [reply] for direction in observed_directions + ), payloads + rows: Final = _spend_rows(model, 2) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 1, rows + _assert_response_id(endpoint, str(guarded_rows[0]["request_id"]), result.response_id, scenario_id) + entries: Final = _guardrail_entries(guarded_rows[0]) + observed_entries: Final = tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) + expected_entries: Final = tuple( + ( + identity, + "logging_only", + "guardrail_intervened" if direction in block_directions else "success", + ) + for direction in expected_directions + ) + assert ( + tuple(sorted(observed_directions)), + tuple(sorted(observed_entries)), + ) == (tuple(sorted(expected_directions)), tuple(sorted(expected_entries))), ( + payloads, + entries, + rows, + ) + finally: + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize( + ("endpoint", "stream", "client_kind"), + ( + pytest.param("chat", False, "openai_sync", id="A-default-chat-nonstream"), + pytest.param("chat", True, "openai_async", id="A-default-chat-stream"), + pytest.param("messages", False, "anthropic_sync", id="A-default-messages-nonstream"), + pytest.param("messages", True, "anthropic_async", id="A-default-messages-stream"), + pytest.param("responses", False, "openai_async", id="A-default-responses-nonstream"), + pytest.param("responses", True, "openai_sync", id="A-default-responses-stream"), + ), +) +def test_A_unset_scope_matches_measured_endpoint_stream_default( + gateway: Gateway, + tmp_path: Path, + endpoint: Endpoint, + stream: bool, + client_kind: ClientKind, +) -> None: + identity: Final = f"logging-scope-a-default-{endpoint}-{stream}-{uuid.uuid4().hex}" + scenario_id: Final = f"phase12-a-default-{endpoint}-{stream}-{uuid.uuid4().hex}" + prompt: Final = f"synthetic unset-scope request {identity}" + reply: Final = f"synthetic unset-scope response {identity}" + upstream_handle: Final = register_scenario(scenario_id, _provider_response(endpoint, scenario_id, reply, stream)) + + def policy(_request: Request) -> Reply: + return Reply(body=b'{"action":"NONE"}') + + try: + with gateway.scenario() as scenario: + api_base: Final = ( + upstream_handle.api_base() if endpoint == "messages" else f"{upstream_handle.api_base()}/v1" + ) + model: Final = scenario.model( + model={ + "chat": "openai/gpt-4o-mini", + "messages": "anthropic/claude-3-7-sonnet-20250219", + "responses": "openai/gpt-4.1-mini", + }[endpoint], + api_base=api_base, + api_key="synthetic-provider-key", + ) + with wire_server(policy) as guardrail: + config: Final = _configuration(tmp_path, identity, guardrail.url, None, include_scope=False) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + baseline: Final = _call_client( + client_kind, + endpoint, + gateway, + model, + prompt, + stream, + f"{scenario_id}-baseline", + ) + assert baseline.status == 200 and baseline.text == reply, baseline + assert len(_drain_upstream(gateway.upstream_url)) == 1 + result: Final = _call_client( + client_kind, + endpoint, + candidate, + model, + prompt, + stream, + f"{scenario_id}-candidate", + ) + assert (result.status, _response_body_without_ids(result.body)) == ( + baseline.status, + _response_body_without_ids(baseline.body), + ), (result, baseline) + assert result.text == reply, result + upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(upstream) == 1 and prompt in json.dumps(upstream[0]["body"]), upstream + expected_directions: Final = BASE_DEFAULT_NORMAL_DIRECTIONS[(endpoint, stream)] + eventually( + lambda: guardrail.received.qsize(), + lambda count: count == len(expected_directions), + seconds=20, + ) + payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert tuple(sorted(_direction(payload) for payload in payloads)) == tuple( + sorted(expected_directions) + ), payloads + assert all(_policy_call_id_matches(payload, f"{scenario_id}-candidate") for payload in payloads), ( + payloads + ) + assert tuple(payload["texts"] for payload in payloads) == tuple( + [prompt] if _direction(payload) == "request" else [reply] for payload in payloads + ), payloads + rows: Final = _spend_rows(model, 2) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 1, rows + _assert_response_id( + endpoint, + str(guarded_rows[0]["request_id"]), + result.response_id, + scenario_id, + ) + entries: Final = _guardrail_entries(guarded_rows[0]) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == tuple((identity, "logging_only", "success") for _ in expected_directions), entries + finally: + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize("inventory_id", (pytest.param("B4", id="B4-monitor-usage-detail"),)) +def test_logging_only_monitor_counts_only_the_observed_direction( + gateway: Gateway, tmp_path: Path, inventory_id: str +) -> None: + input_identity: Final = f"logging-scope-b4-input-{uuid.uuid4().hex}" + output_identity: Final = f"logging-scope-b4-output-{uuid.uuid4().hex}" + scenario_id: Final = f"phase12-{inventory_id.lower()}-{uuid.uuid4().hex}" + prompt_by_identity: Final = { + input_identity: f"synthetic input {input_identity}", + output_identity: f"synthetic input {output_identity}", + } + provider_reply: Final = f"synthetic response {scenario_id}" + upstream_handle: Final = register_scenario( + scenario_id, _provider_response("chat", scenario_id, provider_reply, False) + ) + + def policy(request: Request) -> Reply: + return Reply(body=json.dumps({"action": "BLOCKED", "blocked_reason": "synthetic monitor denial"}).encode()) + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + with wire_server(policy) as input_guardrail, wire_server(policy) as output_guardrail: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"]["cache"] = False + config["guardrails"] = [ + { + "guardrail_name": input_identity, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "logging_only", + "logging_only_scope": "input", + "default_on": False, + "api_base": input_guardrail.url, + "api_key": "synthetic-guardrail-key", + "extra_headers": ["x-litellm-call-id"], + }, + }, + { + "guardrail_name": output_identity, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "logging_only", + "logging_only_scope": "output", + "default_on": False, + "api_base": output_guardrail.url, + "api_key": "synthetic-guardrail-key", + "extra_headers": ["x-litellm-call-id"], + }, + }, + ] + config_path: Final = tmp_path / "b4.yaml" + config_path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {}, config=config_path, workers=1) as candidate: + cases: Final = tuple( + ( + identity, + gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": prompt_by_identity[identity]}], + }, + headers={"x-litellm-call-id": f"{scenario_id}-{identity}-baseline"}, + ), + candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": prompt_by_identity[identity]}], + "guardrails": [identity], + }, + headers={"x-litellm-call-id": f"{scenario_id}-{identity}"}, + ), + ) + for identity in (input_identity, output_identity) + ) + observed_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(observed_upstream) == 4, observed_upstream + for identity, baseline_response, guarded_response in cases: + assert baseline_response.status_code == guarded_response.status_code == 200, ( + baseline_response.text, + guarded_response.text, + ) + assert _response_body_without_ids(JSON_OBJECT.validate_python(baseline_response.json())) == ( + _response_body_without_ids(JSON_OBJECT.validate_python(guarded_response.json())) + ), ( + baseline_response.text, + guarded_response.text, + ) + assert baseline_response.json()["choices"][0]["message"]["content"] == provider_reply + assert ( + sum( + prompt_by_identity[identity] in json.dumps(observation["body"]) + for observation in observed_upstream + ) + == 2 + ), observed_upstream + expected_call_ids: Final = tuple( + f"{scenario_id}-{identity}" for identity in (input_identity, output_identity) + ) + eventually( + lambda: input_guardrail.received.qsize(), + lambda count: count >= len(cases), + seconds=20, + ) + eventually( + lambda: output_guardrail.received.qsize(), + lambda count: count >= len(cases), + seconds=20, + ) + policy_payloads: Final = { + input_identity: tuple(JSON_OBJECT.validate_json(call.body) for call in input_guardrail.drain()), + output_identity: tuple( + JSON_OBJECT.validate_json(call.body) for call in output_guardrail.drain() + ), + } + for identity in (input_identity, output_identity): + expected_direction: Final = ( + "request" if _is_base_audit_leg() or identity == input_identity else "response" + ) + assert len(policy_payloads[identity]) == len(expected_call_ids), policy_payloads[identity] + for case_identity in (input_identity, output_identity): + call_id: Final = f"{scenario_id}-{case_identity}" + calls_for_id: Final = tuple( + payload + for payload in policy_payloads[identity] + if _policy_call_id_matches(payload, call_id) + ) + call_summary: Final = tuple( + ( + _direction(payload), + payload.get("litellm_call_id"), + _policy_call_id(payload), + tuple(payload["texts"]), + ) + for payload in calls_for_id + ) + assert tuple(_direction(payload) for payload in calls_for_id) == (expected_direction,), ( + call_id, + call_summary, + ) + expected_text: Final = ( + prompt_by_identity[case_identity] if expected_direction == "request" else provider_reply + ) + assert tuple(tuple(payload["texts"]) for payload in calls_for_id) == ((expected_text,),), ( + call_id, + call_summary, + ) + rows: Final = _spend_rows_for_calls( + (model,), + tuple( + (model, str(response.json()["id"]), f"{scenario_id}-{identity}") + for identity, _baseline_response, response in cases + ), + ) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 2, rows + expected_entries: Final = tuple( + sorted( + ( + (input_identity, "logging_only", "guardrail_intervened"), + (output_identity, "logging_only", "guardrail_intervened"), + ) + ) + ) + for identity, _baseline_response, response in cases: + call_id: Final = f"{scenario_id}-{identity}" + matching_rows: Final = _spend_rows_matching_call(rows, model, call_id) + assert len(matching_rows) == 1, (call_id, rows) + _assert_response_id( + "chat", + str(matching_rows[0]["request_id"]), + str(response.json()["id"]), + scenario_id, + ) + entries: Final = _guardrail_entries(matching_rows[0]) + assert ( + tuple( + sorted( + ( + entry["guardrail_name"], + entry["guardrail_mode"], + entry["guardrail_status"], + ) + for entry in entries + ) + ) + == expected_entries + ), entries + listed: Final = candidate.get("/v2/guardrails/list")["guardrails"] + assert isinstance(listed, list), listed + guardrail_ids: Final = { + object_value(row)["guardrail_name"]: str(object_value(row)["guardrail_id"]) + for row in listed + if object_value(row)["guardrail_name"] in (input_identity, output_identity) + } + assert set(guardrail_ids) == {input_identity, output_identity}, listed + today: Final = datetime.now(timezone.utc).date().isoformat() + for identity in (input_identity, output_identity): + detail: Final = eventually( + lambda identity=identity: candidate.request( + "GET", + f"/guardrails/usage/detail/{guardrail_ids[identity]}", + params={"start_date": today, "end_date": today}, + ).json(), + lambda body: body["requestsEvaluated"] >= 1, + seconds=30, + return_last_on_timeout=True, + ) + assert detail["requestsEvaluated"] == len(cases), detail + finally: + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize( + ("row_id", "endpoint", "scope", "expected_direction"), + ( + pytest.param("C1", "chat", "output", "response", id="C1-chat-cache-output"), + pytest.param("C2", "chat", "input", "request", id="C2-chat-cache-input"), + pytest.param("C3", "messages", "output", "response", id="C3-messages-cache-output"), + pytest.param("C4", "responses", "output", "response", id="C4-responses-cache-output"), + ), +) +def test_cache_hit_directional_scope_uses_measured_base_default( + gateway: Gateway, + tmp_path: Path, + row_id: str, + endpoint: Endpoint, + scope: str, + expected_direction: Direction, +) -> None: + identity: Final = f"logging-scope-{row_id.lower()}-{uuid.uuid4().hex}" + scenario_id: Final = f"phase12-{row_id.lower()}-{uuid.uuid4().hex}" + baseline_prompt: Final = f"uncached control {identity}" + cached_prompt: Final = f"repeated cache prompt {identity}" + reply: Final = f"cache response {identity}" + base_default_on_hit: Final = BASE_DEFAULT_CACHE_HIT_DIRECTIONS[endpoint] + expected_miss_directions: Final = _directions_for_audit_leg( + BASE_DEFAULT_NORMAL_DIRECTIONS[(endpoint, False)], scope + ) + expected_hit_directions: Final = _directions_for_audit_leg(base_default_on_hit, scope) + if _is_base_audit_leg(): + assert expected_hit_directions == base_default_on_hit, (base_default_on_hit, expected_hit_directions) + upstream_handle: Final = register_scenario(scenario_id, _provider_response(endpoint, scenario_id, reply, False)) + + def policy(_request: Request) -> Reply: + return Reply(body=b'{"action":"NONE"}') + + try: + with gateway.scenario() as scenario: + api_base: Final = ( + upstream_handle.api_base() if endpoint == "messages" else f"{upstream_handle.api_base()}/v1" + ) + model: Final = scenario.model( + model={ + "chat": "openai/gpt-4o-mini", + "messages": "anthropic/claude-3-7-sonnet-20250219", + "responses": "openai/gpt-4.1-mini", + }[endpoint], + api_base=api_base, + api_key="synthetic-provider-key", + ) + with wire_server(policy) as guardrail: + config: Final = _configuration(tmp_path, identity, guardrail.url, scope, cache=True) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + baseline: Final = _call_cache_client( + endpoint, gateway, model, baseline_prompt, f"{scenario_id}-baseline" + ) + assert baseline.status == 200, baseline.body + assert baseline.text == reply, baseline + assert len(_drain_upstream(gateway.upstream_url)) == 1 + first: Final = _call_cache_client(endpoint, candidate, model, cached_prompt, f"{scenario_id}-first") + assert (first.status, _response_body_without_ids(first.body)) == ( + baseline.status, + _response_body_without_ids(baseline.body), + ), (first, baseline) + assert first.text == reply, first + first_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(first_upstream) == 1, first_upstream + assert cached_prompt in json.dumps(first_upstream[0]["body"]), first_upstream + eventually( + lambda: guardrail.received.qsize(), + lambda count: count >= len(expected_miss_directions), + seconds=20, + ) + miss_payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert all(_policy_call_id_matches(payload, f"{scenario_id}-first") for payload in miss_payloads), ( + miss_payloads + ) + second: Final = _call_cache_client( + endpoint, candidate, model, cached_prompt, f"{scenario_id}-second" + ) + assert (second.status, _response_body_without_ids(second.body)) == ( + first.status, + _response_body_without_ids(first.body), + ), (second, first) + assert _drain_upstream(gateway.upstream_url) == (), "The identical second request must hit Redis" + rows: Final = _spend_rows(model, 3) + eventually( + lambda: guardrail.received.qsize(), + lambda count: count >= len(expected_hit_directions), + seconds=20, + ) + hit_payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + for response_id, policy_payloads, expected_directions in ( + (first.response_id, miss_payloads, expected_miss_directions), + (second.response_id, hit_payloads, expected_hit_directions), + ): + assert tuple(sorted(_direction(payload) for payload in policy_payloads)) == tuple( + sorted(expected_directions) + ), (response_id, policy_payloads) + assert all( + payload["texts"] == ([cached_prompt] if _direction(payload) == "request" else [reply]) + for payload in policy_payloads + ), (response_id, policy_payloads) + assert sum(_cache_hit(row["cache_hit"]) for row in rows) == 1, rows + guarded_rows: Final = tuple( + row for row in rows if identity in object_value(row["metadata"]).get("applied_guardrails", []) + ) + assert len(guarded_rows) == 2, rows + miss_rows: Final = tuple(row for row in guarded_rows if not _cache_hit(row["cache_hit"])) + hit_rows: Final = tuple(row for row in guarded_rows if _cache_hit(row["cache_hit"])) + assert len(miss_rows) == len(hit_rows) == 1, rows + _assert_response_id(endpoint, str(miss_rows[0]["request_id"]), first.response_id, scenario_id) + assert str(hit_rows[0]["request_id"]).startswith(f"{second.response_id}_cache_hit"), hit_rows + for row, expected_directions in ( + (miss_rows[0], expected_miss_directions), + (hit_rows[0], expected_hit_directions), + ): + entries: Final = _guardrail_entries(row) + assert len(entries) == len(expected_directions), entries + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == tuple((identity, "logging_only", "success") for _ in expected_directions), entries + finally: + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize( + ("row_id", "scope", "modes", "expected_directions", "expected_statuses", "expected_status"), + ( + pytest.param( + "D1", + "output", + ["pre_call", "logging_only"], + ("request",), + ("guardrail_intervened",), + 400, + id="D1-pre-call-blocks-before-upstream", + ), + pytest.param( + "D2", + "output", + ["pre_call", "logging_only"], + ("request", "response"), + ("success", "guardrail_intervened"), + 200, + id="D2-pre-call-and-output-observation", + ), + pytest.param( + "D3", + "input", + ["logging_only", "post_call"], + ("response",), + ("guardrail_intervened",), + 400, + id="D3-post-call-block-remains-enforced", + ), + ), +) +def test_combined_modes_preserve_blocking_and_directional_observation( + gateway: Gateway, + tmp_path: Path, + row_id: str, + scope: str, + modes: list[str], + expected_directions: tuple[Direction, ...], + expected_statuses: tuple[str, ...], + expected_status: int, +) -> None: + identity: Final = f"logging-scope-{row_id.lower()}-{uuid.uuid4().hex}" + prompt: Final = f"synthetic combined-mode prompt {identity}" + reply: Final = f"synthetic combined-mode response {identity}" + scenario_id: Final = f"phase12-{row_id.lower()}-{uuid.uuid4().hex}" + call_id: Final = f"{scenario_id}-guarded" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def policy(request: Request) -> Reply: + payload: Final = JSON_OBJECT.validate_json(request.body) + direction: Final = _direction(payload) + blocked: Final = row_id == "D1" or direction == "response" + verdict: Final = ( + {"action": "BLOCKED", "blocked_reason": f"synthetic denial from {identity}"} + if blocked + else {"action": "NONE"} + ) + return Reply(body=json.dumps(verdict).encode()) + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + with wire_server(policy) as guardrail: + config: Final = _configuration(tmp_path, identity, guardrail.url, scope, mode=modes) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + baseline: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": f"{scenario_id}-baseline"}, + ) + assert baseline.status_code == 200, baseline.text + assert baseline.json()["choices"][0]["message"]["content"] == reply, baseline.text + assert len(_drain_upstream(gateway.upstream_url)) == 1 + guarded: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": call_id}, + ) + leg_expected_directions: Final = ( + ("request", "request", "response") + if _is_base_audit_leg() and row_id == "D2" + else expected_directions + ) + leg_expected_statuses: Final = ( + ("success", "success", "guardrail_intervened") + if _is_base_audit_leg() and row_id == "D2" + else expected_statuses + ) + assert guarded.status_code == expected_status, guarded.text + candidate_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(candidate_upstream) == (0 if row_id == "D1" else 1), candidate_upstream + if candidate_upstream: + assert prompt in json.dumps(candidate_upstream[0]["body"]), candidate_upstream + if row_id == "D2": + assert _response_body_without_ids(JSON_OBJECT.validate_python(guarded.json())) == ( + _response_body_without_ids(JSON_OBJECT.validate_python(baseline.json())) + ), (guarded.text, baseline.text) + else: + assert identity in guarded.text, guarded.text + assert f"synthetic denial from {identity}" in guarded.text, guarded.text + expected_policy_call_count: Final = len(leg_expected_directions) + eventually( + lambda: guardrail.received.qsize(), + lambda count: count >= expected_policy_call_count, + seconds=20, + ) + calls: Final = guardrail.drain() + payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in calls) + assert all(_policy_call_id_matches(payload, call_id) for payload in payloads), payloads + observed_directions: Final = tuple(_direction(payload) for payload in payloads) + assert tuple(payload["texts"] for payload in payloads) == tuple( + [prompt] if direction == "request" else [reply] for direction in observed_directions + ), payloads + rows: Final = _spend_rows(model, 2) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 1, rows + entries: Final = _guardrail_entries(guarded_rows[0]) + observed_modes: Final = _guardrail_mode_status_pairs(entries) + assert all(mode_values == tuple(modes) for mode_values, _ in observed_modes), entries + observed_statuses: Final = tuple(status for _, status in observed_modes) + assert ( + tuple(sorted(observed_directions)), + tuple(sorted(observed_statuses)), + ) == (tuple(sorted(leg_expected_directions)), tuple(sorted(leg_expected_statuses))), ( + payloads, + entries, + rows, + ) + if row_id == "D2": + _assert_response_id( + "chat", + str(guarded_rows[0]["request_id"]), + str(guarded.json()["id"]), + scenario_id, + ) + else: + assert _guardrail_entries(_spend_row_for_call_id(call_id)) == entries, entries + finally: + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize( + ("row_id", "scope", "selection"), + ( + pytest.param("E1", "output", "request", id="E1-selected-by-request-body"), + pytest.param("E2", "input", "virtual-key", id="E2-selected-by-key-metadata"), + pytest.param("E3", "output", "unselected", id="E3-no-request-or-key-selection"), + ), +) +def test_logging_only_scope_respects_guardrail_selection_level( + gateway: Gateway, tmp_path: Path, row_id: str, scope: str, selection: str +) -> None: + identity: Final = f"logging-scope-{row_id.lower()}-{uuid.uuid4().hex}" + prompt: Final = f"synthetic selected request {identity}" + reply: Final = f"synthetic selected response {identity}" + scenario_id: Final = f"phase12-{row_id.lower()}-{uuid.uuid4().hex}" + call_id: Final = f"{scenario_id}-guarded" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def policy(_request: Request) -> Reply: + return Reply(body=b'{"action":"NONE"}') + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + key: Final = ( + scenario.key(metadata={"guardrails": [identity]}) if selection == "virtual-key" else gateway.key + ) + with wire_server(policy) as guardrail: + config: Final = _configuration(tmp_path, identity, guardrail.url, scope, default_on=False) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + request_body: Final[dict[str, JsonValue]] = { + "model": model, + "messages": [{"role": "user", "content": prompt}], + **({"guardrails": [identity]} if selection == "request" else {}), + } + baseline: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + key=key, + headers={"x-litellm-call-id": f"{scenario_id}-baseline"}, + ) + assert baseline.status_code == 200, baseline.text + assert baseline.json()["choices"][0]["message"]["content"] == reply, baseline.text + assert len(_drain_upstream(gateway.upstream_url)) == 1 + guarded: Final = candidate.request( + "POST", + "/v1/chat/completions", + request_body, + key=key, + headers={"x-litellm-call-id": call_id}, + ) + assert guarded.status_code == 200, guarded.text + assert _response_body_without_ids(JSON_OBJECT.validate_python(guarded.json())) == ( + _response_body_without_ids(JSON_OBJECT.validate_python(baseline.json())) + ), (guarded.text, baseline.text) + observed_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(observed_upstream) == 1, observed_upstream + assert prompt in json.dumps(observed_upstream[0]["body"]), observed_upstream + expected_directions: Final = _directions_for_audit_leg(("request", "response"), scope) + directions_to_collect: Final = expected_directions + eventually( + lambda: guardrail.received.qsize(), + lambda count: count >= len(directions_to_collect), + seconds=20, + ) + payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + if expected_directions: + assert all(_policy_call_id_matches(payload, call_id) for payload in payloads), payloads + assert tuple(sorted(_direction(payload) for payload in payloads)) == tuple( + sorted(expected_directions) + ), payloads + if expected_directions: + rows: Final = _spend_rows(model, 2) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 1, rows + row: Final = guarded_rows[0] + _assert_response_id("chat", str(row["request_id"]), str(guarded.json()["id"]), scenario_id) + entries: Final = _guardrail_entries(row) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == tuple((identity, "logging_only", "success") for _ in expected_directions), entries + else: + row: Final = _spend_row_for_response_id(str(guarded.json()["id"])) + assert all(entry["guardrail_name"] != identity for entry in _guardrail_entries(row)), row + finally: + delete_scenario(upstream_handle) + + +def test_X1_missing_null_and_both_scope_have_identical_scans(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-x1-{uuid.uuid4().hex}" + variants: Final = ( + ("both", "both", True), + ("missing", None, False), + ("null", None, True), + ) + + def policy(_request: Request) -> Reply: + return Reply(body=b'{"action":"NONE"}') + + with gateway.scenario() as scenario: + model: Final = scenario.model() + baseline: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": identity}]}, + headers={"x-litellm-call-id": f"{identity}-baseline"}, + ) + assert baseline.status_code == 200, baseline.text + baseline_body: Final = JSON_OBJECT.validate_json(baseline.content) + baseline_reply: Final = _response_text("chat", baseline_body) + assert len(_drain_upstream(gateway.upstream_url)) == 1 + for suffix, scope, include_scope in variants: + name: Final = f"{identity}-{suffix}" + call_id: Final = f"{identity}-{suffix}" + with wire_server(policy) as edge: + config: Final = _configuration( + tmp_path, + name, + edge.url, + scope, + include_scope=include_scope, + default_on=False, + ) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": identity}], + "guardrails": [name], + }, + headers={"x-litellm-call-id": call_id}, + ) + assert response.status_code == baseline.status_code, response.text + body: Final = JSON_OBJECT.validate_json(response.content) + assert body["choices"] == baseline_body["choices"], response.text + assert _response_text("chat", body) == baseline_reply, response.text + upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(upstream) == 1 and identity in json.dumps(upstream[0]["body"]), upstream + expected_directions: Final = _directions_for_audit_leg(("request", "response"), scope) + eventually( + lambda: edge.received.qsize(), + lambda count, expected_directions=expected_directions: count == len(expected_directions), + seconds=20, + ) + payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in edge.drain()) + assert tuple(sorted(_direction(payload) for payload in payloads)) == tuple( + sorted(expected_directions) + ), payloads + assert all(_policy_call_id_matches(payload, call_id) for payload in payloads), payloads + assert tuple(sorted((_direction(payload), tuple(payload["texts"])) for payload in payloads)) == ( + ("request", (identity,)), + ("response", (baseline_reply,)), + ), payloads + row: Final = _spend_row_for_response_id(str(body["id"])) + entries: Final = _guardrail_entries(row) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == tuple((name, "logging_only", "success") for _ in expected_directions), entries + + +def test_X3_five_identical_requests_each_receive_one_response_scan(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-x3-{uuid.uuid4().hex}" + prompt: Final = f"synthetic repeated request {identity}" + + def policy(_request: Request) -> Reply: + return Reply(body=b'{"action":"NONE"}') + + with gateway.scenario() as scenario: + model: Final = scenario.model() + baseline: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": f"{identity}-baseline"}, + ) + assert baseline.status_code == 200, baseline.text + assert len(_drain_upstream(gateway.upstream_url)) == 1 + with wire_server(policy) as edge: + config: Final = _configuration(tmp_path, identity, edge.url, "output") + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + expected_directions: Final = _directions_for_audit_leg(("request", "response"), "output") + results: Final = tuple( + candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": f"{identity}-{index}"}, + ) + for index in range(5) + ) + assert all(result.status_code == baseline.status_code for result in results), results + assert all(result.json()["choices"] == baseline.json()["choices"] for result in results), results + response_ids: Final = tuple(str(result.json()["id"]) for result in results) + assert len(set(response_ids)) == 5, response_ids + upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(upstream) == 5, upstream + eventually( + lambda: edge.received.qsize(), + lambda count: count == 5 * len(expected_directions), + seconds=20, + ) + calls: Final = edge.drain() + payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in calls) + for index, result in enumerate(results): + call_id: Final = f"{identity}-{index}" + matching_payloads: Final = tuple( + payload for payload in payloads if _policy_call_id_matches(payload, call_id) + ) + assert tuple(sorted(_direction(payload) for payload in matching_payloads)) == tuple( + sorted(expected_directions) + ), ( + call_id, + matching_payloads, + ) + expected_reply: Final = _response_text("chat", JSON_OBJECT.validate_json(result.content)) + assert all( + payload["texts"] == ([prompt] if _direction(payload) == "request" else [expected_reply]) + for payload in matching_payloads + ), matching_payloads + rows: Final = _spend_rows(model, 6) + for index, result in enumerate(results): + matching_rows: Final = tuple(row for row in rows if row["request_id"] == result.json()["id"]) + assert len(matching_rows) == 1, (result.json()["id"], matching_rows) + _assert_response_id( + "chat", + str(matching_rows[0]["request_id"]), + str(result.json()["id"]), + f"{identity}-{index}", + ) + entries: Final = _guardrail_entries(matching_rows[0]) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == tuple((identity, "logging_only", "success") for _ in expected_directions), entries + rows: Final = _spend_rows(model, 6) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 5, rows + assert {str(row["request_id"]) for row in guarded_rows} == set(response_ids), guarded_rows + assert all( + tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in _guardrail_entries(row) + ) + == tuple((identity, "logging_only", "success") for _ in expected_directions) + for row in guarded_rows + ), guarded_rows + + +def test_X2_scope_patch_toggles_during_concurrent_requests_keep_one_scan_per_response( + gateway: Gateway, tmp_path: Path +) -> None: + identity: Final = f"logging-scope-x2-{uuid.uuid4().hex}" + marker: Final = uuid.uuid4().hex + scan_started: Final = threading.Event() + release_scans: Final = threading.Event() + + def policy(_request: Request) -> Reply: + scan_started.set() + assert release_scans.wait(timeout=60), identity + return Reply(body=b'{"action":"NONE"}') + + with gateway.scenario() as scenario: + model: Final = scenario.model() + baseline: Final = _call_client( + "openai_sync", "chat", gateway, model, f"synthetic X2 control {marker}", False, f"{marker}-control" + ) + assert baseline.status == 200, baseline + assert len(_drain_upstream(gateway.upstream_url)) == 1 + with wire_server(policy) as edge: + config: Final = _empty_proxy_configuration(tmp_path, identity) + with ( + _database_guardrail(identity, edge.url, "output", mode="logging_only"), + owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate, + ): + guardrails: Final = candidate.get("/v2/guardrails/list")["guardrails"] + guardrail_id: Final = next( + string_value(object_value(guardrail)["guardrail_id"]) + for guardrail in guardrails + if object_value(guardrail)["guardrail_name"] == identity + ) + calls: Final = tuple( + ChaosCall( + index=index, + endpoint="chat", + client_kind="openai_sync", + model=model, + stream=False, + prompt=f"synthetic X2 request {marker}-{index}", + call_id=f"{marker}-x2-{index}", + ) + for index in range(20) + ) + expected_scan_count: Final = len(calls) * (2 if _is_base_audit_leg() else 1) + with ThreadPoolExecutor(max_workers=20) as pool: + futures: Final = tuple( + pool.submit( + _call_client, + call.client_kind, + call.endpoint, + candidate, + call.model, + call.prompt, + call.stream, + call.call_id, + ) + for call in calls + ) + try: + assert eventually(lambda: scan_started.is_set(), bool, seconds=30) + patch_responses: Final = tuple( + candidate.request( + "PATCH", + f"/guardrails/{guardrail_id}", + {"litellm_params": {"logging_only_scope": "input" if index % 2 == 0 else "output"}}, + ) + for index in range(10) + ) + assert all(response.status_code == 200 for response in patch_responses), patch_responses + finally: + release_scans.set() + results: Final = tuple(future.result(timeout=90) for future in futures) + assert all(result.status == baseline.status for result in results), results + assert all(result.text == baseline.text for result in results), results + response_ids: Final = tuple(result.response_id for result in results) + assert len(set(response_ids)) == 20, response_ids + eventually( + lambda: edge.received.qsize(), + lambda count: count == expected_scan_count, + seconds=30, + ) + edge_calls: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in edge.drain()) + for call in calls: + payloads_for_call: Final = tuple( + payload for payload in edge_calls if _policy_call_id_matches(payload, call.call_id) + ) + directions: Final = tuple(_direction(payload) for payload in payloads_for_call) + if _is_base_audit_leg(): + assert tuple(sorted(directions)) == ("request", "response"), (call, payloads_for_call) + else: + assert len(directions) == 1 and directions[0] in ("request", "response"), ( + call, + payloads_for_call, + ) + assert all( + payload["texts"] == ([call.prompt] if _direction(payload) == "request" else [baseline.text]) + for payload in payloads_for_call + ), payloads_for_call + upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(upstream) == 20, upstream + assert ( + tuple( + sum(_json_contains_exact_string(observation["body"], call.prompt) for observation in upstream) + for call in calls + ) + == (1,) * 20 + ), upstream + rows: Final = _spend_rows(model, 21) + for call, result in zip(calls, results): + row: Final = next(row for row in rows if row["request_id"] == result.response_id) + entries: Final = _guardrail_entries(row) + payloads_for_call: Final = tuple( + payload for payload in edge_calls if _policy_call_id_matches(payload, call.call_id) + ) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == tuple( + (identity, "logging_only", "success") + for _ in tuple(_direction(payload) for payload in payloads_for_call) + ), (call.call_id, entries) + + +def test_S1_logging_only_output_scope_fails_open_on_policy_500(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-s1-{uuid.uuid4().hex}" + prompt: Final = f"synthetic policy outage prompt {identity}" + reply: Final = f"synthetic policy outage response {identity}" + scenario_id: Final = f"phase12-s1-{uuid.uuid4().hex}" + control_call_id: Final = f"{scenario_id}-control" + guarded_call_id: Final = f"{scenario_id}-guarded" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(status=500, body=b'{"error":"synthetic policy outage"}') + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + control: Final = _call_client("openai_sync", "chat", gateway, model, prompt, False, control_call_id) + assert control.status == 200 and control.text == reply, control + assert len(_drain_upstream(gateway.upstream_url)) == 1 + with wire_server(policy) as guardrail: + config: Final = _configuration(tmp_path, identity, guardrail.url, "output") + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + result: Final = _call_client( + "openai_sync", "chat", candidate, model, prompt, False, guarded_call_id + ) + assert (result.status, _response_body_without_ids(result.body)) == ( + control.status, + _response_body_without_ids(control.body), + ), (result, control) + expected_directions: Final = _directions_for_scope(("request", "response"), "output") + directions_to_collect: Final = _directions_for_audit_leg(("request", "response"), "output") + upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(upstream) == 1 and prompt in json.dumps(upstream[0]["body"]), upstream + eventually( + lambda: guardrail.received.qsize(), + lambda count: count >= len(directions_to_collect), + seconds=20, + ) + calls: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert all(_policy_call_id_matches(call, guarded_call_id) for call in calls), calls + assert tuple(call["texts"] for call in calls) == tuple( + [prompt] if _direction(call) == "request" else [reply] for call in calls + ), calls + guarded_row: Final = _spend_row_for_response_id(result.response_id) + _assert_response_id("chat", str(guarded_row["request_id"]), result.response_id, scenario_id) + entries: Final = _guardrail_entries(guarded_row) + observed_directions: Final = tuple(_direction(call) for call in calls) + observed_entries: Final = tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) + expected_entries: Final = tuple( + (identity, "logging_only", "guardrail_failed_to_respond") for _ in expected_directions + ) + assert ( + tuple(sorted(observed_directions)), + tuple(sorted(observed_entries)), + ) == (tuple(sorted(expected_directions)), tuple(sorted(expected_entries))), ( + calls, + entries, + guarded_row, + ) + finally: + delete_scenario(upstream_handle) + + +def test_S2_logging_only_output_scope_scans_both_chat_choices(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-s2-{uuid.uuid4().hex}" + prompt: Final = f"synthetic multiple choice prompt {identity}" + first_reply: Final = f"synthetic first choice {identity}" + second_reply: Final = f"synthetic second choice {identity}" + scenario_id: Final = f"phase12-s2-{uuid.uuid4().hex}" + control_call_id: Final = f"{scenario_id}-control" + guarded_call_id: Final = f"{scenario_id}-guarded" + provider_response: Final = JsonResponse( + content_type="application/json", + body={ + "id": "chatcmpl-$UNIQUE_ID", + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": first_reply}, "finish_reason": "stop"}, + {"index": 1, "message": {"role": "assistant", "content": second_reply}, "finish_reason": "stop"}, + ], + "usage": {"prompt_tokens": 9, "completion_tokens": 10, "total_tokens": 19}, + }, + ) + upstream_handle: Final = register_scenario(scenario_id, provider_response) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"BLOCKED","blocked_reason":"synthetic multiple choice monitor"}') + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + control: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}], "n": 2}, + headers={"x-litellm-call-id": control_call_id}, + ) + assert control.status_code == 200, control.text + assert len(control.json()["choices"]) == 2, control.text + assert len(_drain_upstream(gateway.upstream_url)) == 1 + with wire_server(policy) as guardrail: + config: Final = _configuration(tmp_path, identity, guardrail.url, "output") + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + result: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}], "n": 2}, + headers={"x-litellm-call-id": guarded_call_id}, + ) + assert result.status_code == control.status_code, result.text + assert _response_body_without_ids(JSON_OBJECT.validate_python(result.json())) == ( + _response_body_without_ids(JSON_OBJECT.validate_python(control.json())) + ), (result.text, control.text) + upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(upstream) == 1 and prompt in json.dumps(upstream[0]["body"]), upstream + expected_directions: Final = ("response",) + eventually( + lambda: guardrail.received.qsize(), + lambda count: count == len(expected_directions), + seconds=20, + ) + calls: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert tuple(sorted(_direction(call) for call in calls)) == tuple(sorted(expected_directions)), ( + calls + ) + assert all(_policy_call_id_matches(call, guarded_call_id) for call in calls), calls + assert tuple(sorted((_direction(call), tuple(call["texts"])) for call in calls)) == tuple( + sorted( + (direction, tuple([prompt] if direction == "request" else [first_reply, second_reply])) + for direction in expected_directions + ) + ), calls + rows: Final = _spend_rows(model, 2) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 1, rows + _assert_response_id( + "chat", + str(guarded_rows[0]["request_id"]), + str(result.json()["id"]), + scenario_id, + ) + entries: Final = _guardrail_entries(guarded_rows[0]) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == tuple((identity, "logging_only", "guardrail_intervened") for _ in expected_directions), entries + finally: + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize( + ("endpoint", "scope"), + ( + pytest.param("chat", "input", id="S3-chat-upstream-401-input"), + pytest.param("chat", "output", id="S3-chat-upstream-401-output"), + pytest.param("messages", "input", id="S3-messages-upstream-401-input"), + pytest.param("messages", "output", id="S3-messages-upstream-401-output"), + pytest.param("responses", "input", id="S3-responses-upstream-401-input"), + pytest.param("responses", "output", id="S3-responses-upstream-401-output"), + ), +) +def test_logging_only_scope_on_upstream_401_uses_measured_base_default( + gateway: Gateway, + tmp_path: Path, + record_property: Callable[[str, object], None], + endpoint: Endpoint, + scope: str, +) -> None: + identity: Final = f"logging-scope-s3-{endpoint}-{scope}-{uuid.uuid4().hex}" + prompt: Final = f"synthetic upstream unauthorized marker {identity}" + scenario_id: Final = f"phase12-s3-{endpoint}-{scope}-{uuid.uuid4().hex}" + base_default_on_failure: Final = BASE_DEFAULT_UPSTREAM_FAILURE_DIRECTIONS[endpoint] + expected_directions: Final = _directions_for_audit_leg(base_default_on_failure, scope) + provider_response: Final = JsonResponse( + content_type="application/json", + body={ + "error": { + "message": f"synthetic upstream unauthorized {identity}", + "type": "invalid_request_error", + "code": "401", + } + }, + status=401, + ) + upstream_handle: Final = register_scenario(scenario_id, provider_response) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"NONE"}') + + try: + with gateway.scenario() as scenario: + api_base: Final = ( + upstream_handle.api_base() if endpoint == "messages" else f"{upstream_handle.api_base()}/v1" + ) + model: Final = scenario.model( + model={ + "chat": "openai/gpt-4o-mini", + "messages": "anthropic/claude-3-7-sonnet-20250219", + "responses": "openai/gpt-4.1-mini", + }[endpoint], + api_base=api_base, + api_key="synthetic-provider-key", + ) + path: Final = { + "chat": "/v1/chat/completions", + "messages": "/v1/messages", + "responses": "/v1/responses", + }[endpoint] + request_body: Final = { + "chat": {"model": model, "messages": [{"role": "user", "content": prompt}]}, + "messages": { + "model": model, + "max_tokens": 1000, + "messages": [{"role": "user", "content": prompt}], + }, + "responses": {"model": model, "input": prompt}, + }[endpoint] + control_call_id: Final = f"{scenario_id}-control" + candidate_call_id: Final = f"{scenario_id}-candidate" + control: Final = gateway.client.request( + "POST", + path, + json=request_body, + headers={ + "Authorization": f"Bearer {gateway.key}", + "x-litellm-call-id": control_call_id, + }, + timeout=60, + ) + control_upstream: Final = _drain_upstream(gateway.upstream_url) + control_upstream_count: Final = len(control_upstream) + record_property("s3_control_upstream_request_count", control_upstream_count) + assert control_upstream_count >= 1, control_upstream + assert control.status_code >= 400, control.text + assert all(prompt in json.dumps(observation["body"]) for observation in control_upstream), control_upstream + with wire_server(policy) as guardrail: + config: Final = _configuration(tmp_path, identity, guardrail.url, scope) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + result: Final = candidate.client.request( + "POST", + path, + json=request_body, + headers={ + "Authorization": f"Bearer {candidate.key}", + "x-litellm-call-id": candidate_call_id, + }, + timeout=60, + ) + upstream: Final = _drain_upstream(gateway.upstream_url) + upstream_count: Final = len(upstream) + record_property("s3_candidate_upstream_request_count", upstream_count) + assert upstream_count >= 1, upstream + assert result.status_code >= 400, result.text + assert result.status_code == control.status_code, result.text + assert _response_body_without_ids(JSON_OBJECT.validate_python(result.json())) == ( + _response_body_without_ids(JSON_OBJECT.validate_python(control.json())) + ), (result.text, control.text) + assert all(prompt in json.dumps(observation["body"]) for observation in upstream), upstream + calls: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert all(_policy_call_id_matches(call, candidate_call_id) for call in calls), calls + assert all(call["texts"] == [prompt] for call in calls if _direction(call) == "request"), calls + spend_rows: Final = _spend_rows(model, 2) + assert all(not _guardrail_entries(row) for row in spend_rows), spend_rows + assert {str(row["request_id"]) for row in spend_rows} >= { + control_call_id, + candidate_call_id, + }, spend_rows + assert tuple(sorted(_direction(call) for call in calls)) == tuple(sorted(expected_directions)), ( + calls, + spend_rows, + ) + finally: + delete_scenario(upstream_handle) + + +def test_S4_logging_only_input_scope_scans_every_multipart_text_part(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-s4-{uuid.uuid4().hex}" + first_part: Final = f"synthetic first text part {identity}" + second_part: Final = f"synthetic second text part {identity}" + reply: Final = f"synthetic multipart response {identity}" + scenario_id: Final = f"phase12-s4-{uuid.uuid4().hex}" + control_call_id: Final = f"{scenario_id}-control" + guarded_call_id: Final = f"{scenario_id}-guarded" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"BLOCKED","blocked_reason":"synthetic multipart monitor"}') + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + request_body: Final = { + "model": model, + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": first_part}, + {"type": "text", "text": second_part}, + ], + } + ], + } + control: Final = gateway.request( + "POST", + "/v1/chat/completions", + request_body, + headers={"x-litellm-call-id": control_call_id}, + ) + assert control.status_code == 200, control.text + assert control.json()["choices"][0]["message"]["content"] == reply, control.text + assert len(_drain_upstream(gateway.upstream_url)) == 1 + with wire_server(policy) as guardrail: + config: Final = _configuration(tmp_path, identity, guardrail.url, "input") + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + result: Final = candidate.request( + "POST", + "/v1/chat/completions", + request_body, + headers={"x-litellm-call-id": guarded_call_id}, + ) + assert result.status_code == control.status_code, result.text + assert _response_body_without_ids(JSON_OBJECT.validate_python(result.json())) == ( + _response_body_without_ids(JSON_OBJECT.validate_python(control.json())) + ), (result.text, control.text) + observed_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(observed_upstream) == 1, observed_upstream + assert first_part in json.dumps(observed_upstream[0]["body"]), observed_upstream + assert second_part in json.dumps(observed_upstream[0]["body"]), observed_upstream + expected_directions: Final = ("request",) + eventually( + lambda: guardrail.received.qsize(), + lambda count: count == len(expected_directions), + seconds=20, + ) + calls: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert tuple(sorted(_direction(call) for call in calls)) == tuple(sorted(expected_directions)), ( + calls + ) + assert all(_policy_call_id_matches(call, guarded_call_id) for call in calls), calls + assert tuple(sorted((_direction(call), tuple(call["texts"])) for call in calls)) == tuple( + sorted( + (direction, tuple([first_part, second_part] if direction == "request" else [reply])) + for direction in expected_directions + ) + ), calls + rows: Final = _spend_rows(model, 2) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 1, rows + _assert_response_id( + "chat", str(guarded_rows[0]["request_id"]), str(result.json()["id"]), scenario_id + ) + entries: Final = _guardrail_entries(guarded_rows[0]) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == tuple((identity, "logging_only", "guardrail_intervened") for _ in expected_directions), entries + finally: + delete_scenario(upstream_handle) 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_ollama_prompt_tools_chaos.py b/tests/integration/providers/test_ollama_prompt_tools_chaos.py new file mode 100644 index 00000000000..243c85b2a64 --- /dev/null +++ b/tests/integration/providers/test_ollama_prompt_tools_chaos.py @@ -0,0 +1,404 @@ +import asyncio +import json +import re +import signal +import threading +import uuid +from collections.abc import Callable +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final, Literal +from urllib.parse import urlsplit + +import httpx +import psutil +import pytest +import yaml +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.process import graceful_stop_seconds, owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue, TypeAdapter + +_BACKEND: Final = "llama3-prompt-tools-chaos" +_API_KEY: Final = "synthetic-ollama-key" +_CONFIG_MODEL: Final = "ollama-prompt-tools-chaos" +_INSTRUCTION: Final = ( + 'To call a function, reply with JSON ONLY in this format {"name": "function_name", ' + '"arguments":{"argument_name": "argument_value"}}. Once a function result answers the request, ' + "reply to the user in plain text instead of calling a function again. " + "The following functions are available to you:" +) +_CALL_ID: Final = "call_prompt_tools_chaos_1" +_ARGUMENTS: Final[dict[str, JsonValue]] = {"city": "Paris"} +_PARAMETERS: Final[dict[str, JsonValue]] = { + "type": "object", + "properties": {"city": {"type": "string"}}, + "required": ["city"], +} +_CHAT_TOOL: Final[dict[str, JsonValue]] = { + "type": "function", + "function": {"name": "get_weather", "description": "Weather for a city", "parameters": _PARAMETERS}, +} +_ANTHROPIC_TOOL: Final[dict[str, JsonValue]] = { + "name": "get_weather", + "description": "Weather for a city", + "input_schema": _PARAMETERS, +} +_RESPONSES_TOOL: Final[dict[str, JsonValue]] = { + "type": "function", + "name": "get_weather", + "description": "Weather for a city", + "parameters": _PARAMETERS, +} +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_MARKER: Final = re.compile(r"marker-([0-9a-f]{32})") +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +_OWNED_PROXY_CELL_SECONDS: Final = 2 * graceful_stop_seconds() + 120 + +Endpoint = Literal["chat", "messages", "responses"] + + +@dataclass(frozen=True, slots=True) +class _Call: + endpoint: Endpoint + stream: bool + marker: str + + +@dataclass(frozen=True, slots=True) +class _Served: + call: _Call + status: int + text: str + + +def _result(marker: str) -> str: + return f"Paris: 22 degrees Celsius marker-{marker}" + + +def _answer(marker: str) -> str: + return f"answer marker-{marker}" + + +def _path(endpoint: Endpoint) -> str: + match endpoint: + case "chat": + return "/v1/chat/completions" + case "messages": + return "/v1/messages" + case "responses": + return "/v1/responses" + + +def _body(model: str, call: _Call) -> dict[str, JsonValue]: + question: Final = "What is the weather in Paris?" + common: Final[dict[str, JsonValue]] = { + "model": model, + "stream": call.stream, + "num_retries": 0, + "cache": {"no-cache": True}, + } + match call.endpoint: + case "chat": + return { + **common, + "tools": [_CHAT_TOOL], + "messages": [ + {"role": "user", "content": question}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": _CALL_ID, + "type": "function", + "function": {"name": "get_weather", "arguments": json.dumps(_ARGUMENTS)}, + } + ], + }, + {"role": "tool", "tool_call_id": _CALL_ID, "content": _result(call.marker)}, + ], + } + case "messages": + return { + **common, + "max_tokens": 64, + "tools": [_ANTHROPIC_TOOL], + "messages": [ + {"role": "user", "content": question}, + { + "role": "assistant", + "content": [{"type": "tool_use", "id": _CALL_ID, "name": "get_weather", "input": _ARGUMENTS}], + }, + { + "role": "user", + "content": [{"type": "tool_result", "tool_use_id": _CALL_ID, "content": _result(call.marker)}], + }, + ], + } + case "responses": + return { + **common, + "store": False, + "tools": [_RESPONSES_TOOL], + "input": [ + {"role": "user", "content": question}, + { + "type": "function_call", + "call_id": _CALL_ID, + "name": "get_weather", + "arguments": json.dumps(_ARGUMENTS), + }, + {"type": "function_call_output", "call_id": _CALL_ID, "output": _result(call.marker)}, + ], + } + + +def _generate_reply(marker: str, stream: bool, drop_connection: bool = False) -> Reply: + done: Final = { + "model": _BACKEND, + "created_at": "2026-10-07T00:00:00Z", + "response": "", + "done": True, + "done_reason": "stop", + "prompt_eval_count": 30, + "eval_count": 12, + } + if not stream: + return Reply(body=json.dumps({**done, "response": _answer(marker)}).encode(), drop_connection=drop_connection) + pieces: Final = ("answer ", f"marker-{marker}") + frames: Final = ( + *( + {"model": _BACKEND, "created_at": "2026-10-07T00:00:00Z", "response": piece, "done": False} + for piece in pieces + ), + done, + ) + return Reply( + content_type="application/x-ndjson", + chunks=tuple(json.dumps(frame).encode() + b"\n" for frame in frames), + drop_connection=drop_connection, + ) + + +def _is_generate(request: Request) -> bool: + return (request.method, request.target) == ("POST", "/api/generate") + + +def _marker_of(request: Request) -> str: + found: Final = _MARKER.search(request.body.decode()) + assert found is not None, request.body + return found.group(1) + + +def _echo(request: Request) -> Reply: + if not _is_generate(request): + return Reply(body=b"{}") + stream: Final = _JSON_OBJECT.validate_json(request.body).get("stream") is True + return _generate_reply(_marker_of(request), stream) + + +def _assert_each_prompt_is_instructed_once(received: tuple[Request, ...], markers: frozenset[str]) -> None: + generates: Final = tuple(request for request in received if _is_generate(request)) + assert sorted(_marker_of(request) for request in generates) == sorted(markers) + for request in generates: + body: Final = _JSON_OBJECT.validate_json(request.body) + assert body["format"] == "json", sorted(body) + prompt: Final = body["prompt"] + assert isinstance(prompt, str) + assert prompt.count(_INSTRUCTION) == 1, prompt + assert set(_MARKER.findall(prompt)) == {_marker_of(request)}, prompt + + +def _response_id(served: _Served) -> str | None: + if served.call.endpoint == "responses": + return None + if not served.call.stream: + identity: Final = _JSON_OBJECT.validate_json(served.text)["id"] + assert isinstance(identity, str) + return identity + for line in served.text.splitlines(): + if not line.startswith("data: ") or line == "data: [DONE]": + continue + payload: Final = _JSON_OBJECT.validate_json(line.removeprefix("data: ")) + if served.call.endpoint == "chat": + first: Final = payload["id"] + assert isinstance(first, str) + return first + if payload.get("type") == "message_start": + message: Final = payload["message"] + assert isinstance(message, dict) and isinstance(message["id"], str) + return message["id"] + raise AssertionError(served.text) + + +def _spend_statuses(model: str, expected: int) -> MappingProxyType[str, str]: + rows: Final = eventually( + lambda: read_rows('SELECT request_id, status FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)), + lambda found: len(found) >= expected, + seconds=70, + ) + statuses: Final = MappingProxyType({str(row["request_id"]): str(row["status"]) for row in rows}) + assert len(statuses) == len(rows) == expected, rows + return statuses + + +def _successes(statuses: MappingProxyType[str, str]) -> frozenset[str]: + return frozenset(identity for identity, status in statuses.items() if status == "success") + + +async def _send(client: httpx.AsyncClient, key: str, model: str, call: _Call) -> _Served: + async with client.stream( + "POST", + _path(call.endpoint), + json=_body(model, call), + headers={"Authorization": f"Bearer {key}", "anthropic-version": "2023-06-01"}, + ) as response: + raw: Final = await response.aread() + return _Served(call=call, status=response.status_code, text=raw.decode()) + + +async def _burst( + base_url: str, key: str, model: str, calls: tuple[_Call, ...], *, tolerate_transport_errors: bool = False +) -> tuple[_Served, ...]: + async with httpx.AsyncClient(base_url=base_url, timeout=60, trust_env=False) as client: + results: Final = await asyncio.gather( + *(_send(client, key, model, call) for call in calls), return_exceptions=tolerate_transport_errors + ) + for result in results: + assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result) + return tuple(result for result in results if isinstance(result, _Served)) + + +def _calls(count: int, endpoints: tuple[Endpoint, ...], stream: Callable[[int], bool]) -> tuple[_Call, ...]: + return tuple( + _Call(endpoint=endpoints[index % len(endpoints)], stream=stream(index), marker=uuid.uuid4().hex) + for index in range(count) + ) + + +def _assert_answered_with_its_own_marker(served: _Served) -> None: + assert served.status == 200, served.text + assert set(_MARKER.findall(served.text)) == {served.call.marker}, served.text + + +async def test_concurrent_tool_result_turns_across_endpoints_each_get_their_own_instructed_prompt( + gateway: Gateway, +) -> None: + calls: Final = _calls(24, ("chat", "messages", "responses"), lambda index: index % 2 == 0) + with wire_server(_echo) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + served: Final = await _burst(str(gateway.client.base_url), gateway.key, model, calls) + assert len(served) == 24 + for item in served: + _assert_answered_with_its_own_marker(item) + _assert_each_prompt_is_instructed_once(wire.drain(), frozenset(call.marker for call in calls)) + known: Final = frozenset(identity for identity in map(_response_id, served) if identity is not None) + assert len(known) == 16, known + statuses: Final = _spend_statuses(model, 24) + assert _successes(statuses) == frozenset(statuses), statuses + assert known <= _successes(statuses) + + +async def test_dropped_ollama_connections_fail_their_callers_and_the_rest_keep_their_prompts(gateway: Gateway) -> None: + calls: Final = _calls(12, ("chat",), lambda index: index % 2 == 1) + dropped: Final = frozenset(call.marker for index, call in enumerate(calls) if index % 3 == 0) + + def respond(request: Request) -> Reply: + if not _is_generate(request): + return Reply(body=b"{}") + marker: Final = _marker_of(request) + stream: Final = _JSON_OBJECT.validate_json(request.body).get("stream") is True + return _generate_reply(marker, stream, drop_connection=marker in dropped) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + served: Final = await _burst(str(gateway.client.base_url), gateway.key, model, calls) + assert len(served) == 12 + for item in served: + if item.call.marker in dropped: + assert item.status == 500, item.text + assert "APIConnectionError" in item.text and "marker-" not in item.text, item.text + else: + _assert_answered_with_its_own_marker(item) + recovery: Final = _Call(endpoint="chat", stream=False, marker=uuid.uuid4().hex) + (recovered,) = await _burst(str(gateway.client.base_url), gateway.key, model, (recovery,)) + _assert_answered_with_its_own_marker(recovered) + _assert_each_prompt_is_instructed_once(wire.drain(), frozenset(call.marker for call in (*calls, recovery))) + answered: Final = tuple(item for item in served if item.call.marker not in dropped) + survivors: Final = frozenset( + identity for identity in map(_response_id, (*answered, recovered)) if identity is not None + ) + assert len(survivors) == 9, survivors + statuses: Final = _spend_statuses(model, 13) + assert _successes(statuses) == survivors, statuses + assert sum(status == "failure" for status in statuses.values()) == len(dropped), statuses + + +def _chaos_config(wire: Wire, tmp_path: Path) -> Path: + base: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + deployment: Final = { + "model_name": _CONFIG_MODEL, + "litellm_params": {"model": f"ollama/{_BACKEND}", "api_base": wire.url, "api_key": _API_KEY}, + } + path: Final = tmp_path / "ollama-prompt-tools-chaos.yaml" + path.write_text(yaml.safe_dump({**base, "model_list": [deployment]})) + return path + + +def _open_upstream_connections(pid: int, upstream: str) -> int: + port: Final = urlsplit(upstream).port + return sum( + 1 + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port + ) + + +@pytest.mark.timeout(_OWNED_PROXY_CELL_SECONDS) +async def test_worker_sigkill_mid_burst_leaves_the_sibling_instructing_ollama(gateway: Gateway, tmp_path: Path) -> None: + calls: Final = _calls(20, ("chat",), lambda _: False) + release: Final = threading.Event() + held_markers: Final[SimpleQueue[str]] = SimpleQueue() + + def held(request: Request) -> Reply: + if not _is_generate(request): + return Reply(body=b"{}") + held_markers.put(_marker_of(request)) + assert release.wait(timeout=60), "The burst was never released" + return _echo(request) + + with wire_server(held) as wire: + path: Final = _chaos_config(wire, tmp_path) + with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned: + candidate: Final = owned.gateway + workers: Final = eventually( + lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())), + lambda pids: len(pids) == 2, + seconds=30, + ) + burst: Final = asyncio.create_task( + _burst( + str(candidate.client.base_url), candidate.key, _CONFIG_MODEL, calls, tolerate_transport_errors=True + ) + ) + await asyncio.to_thread(eventually, held_markers.qsize, lambda size: size == 20, 60) + held_by: Final = MappingProxyType({pid: _open_upstream_connections(pid, wire.url) for pid in workers}) + assert sum(held_by.values()) == 20, held_by + victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__) + victim: Final = psutil.Process(victim_pid) + victim.suspend() + victim.send_signal(signal.SIGKILL) + release.set() + served: Final = await burst + assert held_by[survivor_pid] >= 10, held_by + assert len(served) == held_by[survivor_pid], (held_by, len(served)) + for item in served: + _assert_answered_with_its_own_marker(item) + follow_up: Final = _Call(endpoint="chat", stream=False, marker=uuid.uuid4().hex) + (answered,) = await _burst(str(candidate.client.base_url), candidate.key, _CONFIG_MODEL, (follow_up,)) + _assert_answered_with_its_own_marker(answered) + _assert_each_prompt_is_instructed_once(wire.drain(), frozenset(call.marker for call in (*calls, follow_up))) diff --git a/tests/integration/providers/test_ollama_prompt_tools_wire.py b/tests/integration/providers/test_ollama_prompt_tools_wire.py new file mode 100644 index 00000000000..115fdc65e33 --- /dev/null +++ b/tests/integration/providers/test_ollama_prompt_tools_wire.py @@ -0,0 +1,796 @@ +import itertools +import json +import uuid +from collections.abc import Callable, Iterator, Sequence +from contextlib import contextmanager +from typing import Final + +import anthropic +import openai +import pytest +from openai.types.chat import ChatCompletionChunk +from openai.types.chat.chat_completion_chunk import Choice as ChunkChoice +from openai.types.chat.chat_completion_chunk import ChoiceDeltaToolCall +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue, TypeAdapter + +_BACKEND: Final = "llama3-prompt-tools" +_API_KEY: Final = "synthetic-ollama-key" +_INSTRUCTION: Final = ( + 'To call a function, reply with JSON ONLY in this format {"name": "function_name", ' + '"arguments":{"argument_name": "argument_value"}}. Once a function result answers the request, ' + "reply to the user in plain text instead of calling a function again. " + "The following functions are available to you:" +) +_QUESTION: Final = "What is the weather in Paris?" +_RESULT: Final = "Paris: 22 degrees Celsius, clear skies" +_ANSWER: Final = "Paris is 22 degrees Celsius with clear skies." +_CALL_ID: Final = "call_prompt_tools_1" +_ARGUMENTS: Final[dict[str, JsonValue]] = {"city": "Paris"} +_CALL_JSON: Final = json.dumps({"name": "get_weather", "arguments": _ARGUMENTS}) +_CALL_JSON_FIELDS: Final = frozenset({"get_weather"}) +_PARAMETERS: Final[dict[str, JsonValue]] = { + "type": "object", + "properties": {"city": {"type": "string"}}, + "required": ["city"], +} +_WEATHER_TOOL: Final[dict[str, JsonValue]] = { + "type": "function", + "function": {"name": "get_weather", "description": "Weather for a city", "parameters": _PARAMETERS}, +} +_TIME_TOOL: Final[dict[str, JsonValue]] = { + "type": "function", + "function": {"name": "get_time", "description": "Local time for a city", "parameters": _PARAMETERS}, +} +_ANTHROPIC_TOOL: Final[dict[str, JsonValue]] = { + "name": "get_weather", + "description": "Weather for a city", + "input_schema": _PARAMETERS, +} +_RESPONSES_TOOL: Final[dict[str, JsonValue]] = { + "type": "function", + "name": "get_weather", + "description": "Weather for a city", + "parameters": _PARAMETERS, +} +_NO_CACHE: Final[dict[str, JsonValue]] = {"no-cache": True} +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_RAW_ANTHROPIC_EVENTS: Final = frozenset( + { + "message_start", + "content_block_start", + "content_block_delta", + "content_block_stop", + "message_delta", + "message_stop", + } +) + + +def _generate_reply(text: str) -> Reply: + return Reply( + body=json.dumps( + { + "model": _BACKEND, + "created_at": "2026-10-07T00:00:00Z", + "response": text, + "done": True, + "done_reason": "stop", + "prompt_eval_count": 30, + "eval_count": 12, + } + ).encode() + ) + + +def _streamed_reply(text: str) -> Reply: + pieces: Final = tuple(text[index : index + 7] for index in range(0, len(text), 7)) + lines: Final = tuple( + json.dumps({"model": _BACKEND, "created_at": "2026-10-07T00:00:00Z", "response": piece, "done": False}).encode() + + b"\n" + for piece in pieces + ) + final: Final = ( + json.dumps( + { + "model": _BACKEND, + "created_at": "2026-10-07T00:00:00Z", + "response": "", + "done": True, + "done_reason": "stop", + "prompt_eval_count": 30, + "eval_count": 12, + } + ).encode() + + b"\n" + ) + return Reply(content_type="application/x-ndjson", chunks=(*lines, final)) + + +def _is_generate(request: Request) -> bool: + return (request.method, request.target) == ("POST", "/api/generate") + + +@contextmanager +def _ollama_server(respond: Callable[[Request], Reply]) -> Iterator[Wire]: + with wire_server(lambda request: respond(request) if _is_generate(request) else Reply(body=b"{}")) as wire: + yield wire + + +def _generate_calls(wire: Wire) -> tuple[Request, ...]: + return tuple(request for request in wire.drain() if _is_generate(request)) + + +def _only_generate(wire: Wire) -> dict[str, JsonValue]: + received: Final = _generate_calls(wire) + assert len(received) == 1, [(request.method, request.target) for request in received] + assert received[0].headers["authorization"] == f"Bearer {_API_KEY}" + return _JSON_OBJECT.validate_json(received[0].body) + + +def _prompt_of(body: dict[str, JsonValue]) -> str: + assert body["model"] == _BACKEND + assert body["format"] == "json" + assert "tools" not in body and "messages" not in body, sorted(body) + prompt: Final = body["prompt"] + assert isinstance(prompt, str) + return prompt + + +def _assert_instructed_once(prompt: str, *tool_names: str) -> None: + assert prompt.count(_INSTRUCTION) == 1, prompt + assert prompt.count("### System:") == 1, prompt + for name in tool_names: + assert f"'name': '{name}'" in prompt, prompt + + +def _assert_tool_turn(prompt: str, result: str = _RESULT) -> None: + _assert_instructed_once(prompt, "get_weather") + assert f"### User:\n{_QUESTION}\n\n### Assistant:\n{_CALL_JSON}\n\n### User:\n{result}\n\n" in prompt, prompt + + +def _spend_row(identity: str) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT model_group, status, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" WHERE request_id=%s', + (identity,), + ), + lambda found: len(found) == 1, + seconds=70, + ) + return rows[0] + + +def _model_spend_rows(model: str, expected: int) -> list[dict[str, JsonValue]]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id, status, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" WHERE model_group=%s', + (model,), + ), + lambda found: len(found) >= expected, + seconds=70, + ) + assert len({row["request_id"] for row in rows}) == len(rows) == expected, rows + return rows + + +def _billed(model: str) -> dict[str, JsonValue]: + return {"model_group": model, "status": "success", "prompt_tokens": 30, "completion_tokens": 12} + + +def _openai_client(gateway: Gateway) -> openai.OpenAI: + return openai.OpenAI(base_url=str(gateway.client.base_url) + "/v1", api_key=gateway.key, max_retries=0) + + +def _async_openai_client(gateway: Gateway) -> openai.AsyncOpenAI: + return openai.AsyncOpenAI(base_url=str(gateway.client.base_url) + "/v1", api_key=gateway.key, max_retries=0) + + +def _anthropic_client(gateway: Gateway) -> anthropic.Anthropic: + return anthropic.Anthropic(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0) + + +def _async_anthropic_client(gateway: Gateway) -> anthropic.AsyncAnthropic: + return anthropic.AsyncAnthropic(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0) + + +def _stream_choices(chunks: Sequence[ChatCompletionChunk]) -> Iterator[ChunkChoice]: + for chunk in chunks: + yield from chunk.choices + + +def _delta_tool_calls(choices: Sequence[ChunkChoice]) -> Iterator[ChoiceDeltaToolCall]: + for choice in choices: + yield from choice.delta.tool_calls or () + + +def _first_turn() -> list[dict[str, JsonValue]]: + return [{"role": "user", "content": _QUESTION}] + + +def _second_turn(result: JsonValue = _RESULT) -> list[dict[str, JsonValue]]: + return [ + {"role": "user", "content": _QUESTION}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": _CALL_ID, + "type": "function", + "function": {"name": "get_weather", "arguments": json.dumps(_ARGUMENTS)}, + } + ], + }, + {"role": "tool", "tool_call_id": _CALL_ID, "content": result}, + ] + + +def _anthropic_second_turn() -> list[dict[str, JsonValue]]: + return [ + {"role": "user", "content": _QUESTION}, + { + "role": "assistant", + "content": [{"type": "tool_use", "id": _CALL_ID, "name": "get_weather", "input": _ARGUMENTS}], + }, + {"role": "user", "content": [{"type": "tool_result", "tool_use_id": _CALL_ID, "content": _RESULT}]}, + ] + + +def _responses_second_turn() -> list[dict[str, JsonValue]]: + return [ + {"role": "user", "content": _QUESTION}, + {"type": "function_call", "call_id": _CALL_ID, "name": "get_weather", "arguments": json.dumps(_ARGUMENTS)}, + {"type": "function_call_output", "call_id": _CALL_ID, "output": _RESULT}, + ] + + +def _post(gateway: Gateway, path: str, body: dict[str, JsonValue], key: str | None = None) -> tuple[int, str]: + response: Final = gateway.request("POST", path, {**body, "cache": _NO_CACHE}, key=key) + return response.status_code, response.text + + +def _post_chat( + gateway: Gateway, model: str, messages: Sequence[dict[str, JsonValue]], **extra: JsonValue +) -> dict[str, JsonValue]: + code, text = _post(gateway, "/v1/chat/completions", {"model": model, "messages": list(messages), **extra}) + assert code == 200, text + return _JSON_OBJECT.validate_json(text) + + +def test_openai_sdk_tool_request_reaches_ollama_as_an_instructed_prompt(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + completion: Final = _openai_client(gateway).chat.completions.create( + model=model, + messages=_first_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_WEATHER_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + extra_body={"cache": _NO_CACHE}, + ) + choice: Final = completion.choices[0] + assert choice.finish_reason == "tool_calls" + assert choice.message.tool_calls is not None and len(choice.message.tool_calls) == 1 + call: Final = choice.message.tool_calls[0] + assert call.type == "function" + assert call.function.name == "get_weather" + assert json.loads(call.function.arguments) == _ARGUMENTS + assert completion.usage is not None + assert (completion.usage.prompt_tokens, completion.usage.completion_tokens) == (30, 12) + body: Final = _only_generate(wire) + assert body["stream"] is False + prompt: Final = _prompt_of(body) + _assert_instructed_once(prompt, "get_weather") + assert f"### User:\n{_QUESTION}\n\n" in prompt, prompt + assert "Weather for a city" in prompt, prompt + assert _spend_row(completion.id) == _billed(model) + + +def test_openai_sdk_tool_result_turn_gets_a_plain_text_answer(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + completion: Final = _openai_client(gateway).chat.completions.create( + model=model, + messages=_second_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_WEATHER_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + extra_body={"cache": _NO_CACHE}, + ) + choice: Final = completion.choices[0] + assert choice.finish_reason == "stop" + assert choice.message.content == _ANSWER + assert choice.message.tool_calls is None + _assert_tool_turn(_prompt_of(_only_generate(wire))) + assert _spend_row(completion.id) == _billed(model) + + +async def test_async_openai_sdk_stream_flushes_the_held_tool_call_once(gateway: Gateway) -> None: + with _ollama_server(lambda _: _streamed_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + stream: Final = await _async_openai_client(gateway).chat.completions.create( + model=model, + messages=_first_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_WEATHER_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + stream=True, + stream_options={"include_usage": True}, + extra_body={"cache": _NO_CACHE}, + ) + chunks: Final = [chunk async for chunk in stream] + assert {chunk.id for chunk in chunks} == {chunks[0].id} + choices: Final = tuple(_stream_choices(chunks)) + deltas: Final = tuple(_delta_tool_calls(choices)) + assert len(deltas) == 1, deltas + assert deltas[0].function is not None and deltas[0].function.name == "get_weather" + assert json.loads(deltas[0].function.arguments or "") == _ARGUMENTS + assert [choice.finish_reason for choice in choices if choice.finish_reason] == ["tool_calls"] + usages: Final = [chunk.usage for chunk in chunks if chunk.usage is not None] + assert [(usage.prompt_tokens, usage.completion_tokens) for usage in usages] == [(30, 12)] + body: Final = _only_generate(wire) + assert body["stream"] is True + _assert_instructed_once(_prompt_of(body), "get_weather") + assert _spend_row(chunks[0].id) == _billed(model) + + +async def test_async_openai_sdk_stream_answers_the_tool_result_in_plain_text(gateway: Gateway) -> None: + with _ollama_server(lambda _: _streamed_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + stream: Final = await _async_openai_client(gateway).chat.completions.create( + model=model, + messages=_second_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_WEATHER_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + stream=True, + stream_options={"include_usage": True}, + extra_body={"cache": _NO_CACHE}, + ) + chunks: Final = [chunk async for chunk in stream] + choices: Final = tuple(_stream_choices(chunks)) + assert "".join(choice.delta.content or "" for choice in choices) == _ANSWER + assert tuple(_delta_tool_calls(choices)) == () + assert [choice.finish_reason for choice in choices if choice.finish_reason] == ["stop"] + body: Final = _only_generate(wire) + assert body["stream"] is True + _assert_tool_turn(_prompt_of(body)) + assert _spend_row(chunks[0].id) == _billed(model) + + +def test_anthropic_sdk_tool_request_comes_back_as_tool_use(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + message: Final = _anthropic_client(gateway).messages.create( + model=model, + max_tokens=64, + messages=_first_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_ANTHROPIC_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + extra_body={"cache": _NO_CACHE}, + ) + assert message.stop_reason == "tool_use" + assert [block.type for block in message.content] == ["tool_use"] + block: Final = message.content[0] + assert block.type == "tool_use" + assert block.name == "get_weather" + assert block.input == _ARGUMENTS + assert (message.usage.input_tokens, message.usage.output_tokens) == (30, 12) + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather") + assert f"### User:\n{_QUESTION}\n\n" in prompt, prompt + assert _spend_row(message.id) == _billed(model) + + +def test_anthropic_sdk_tool_result_turn_ends_with_text(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + message: Final = _anthropic_client(gateway).messages.create( + model=model, + max_tokens=64, + messages=_anthropic_second_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_ANTHROPIC_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + extra_body={"cache": _NO_CACHE}, + ) + assert message.stop_reason == "end_turn" + assert [(block.type, getattr(block, "text", None)) for block in message.content] == [("text", _ANSWER)] + _assert_tool_turn(_prompt_of(_only_generate(wire))) + assert _spend_row(message.id) == _billed(model) + + +async def test_async_anthropic_sdk_stream_emits_the_tool_use_block(gateway: Gateway) -> None: + with _ollama_server(lambda _: _streamed_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + async with _async_anthropic_client(gateway).messages.stream( + model=model, + max_tokens=64, + messages=_first_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_ANTHROPIC_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + extra_body={"cache": _NO_CACHE}, + ) as stream: + events: Final = [event async for event in stream if event.type in _RAW_ANTHROPIC_EVENTS] + final: Final = await stream.get_final_message() + starts: Final = [event for event in events if event.type == "content_block_start"] + assert [event.content_block.type for event in starts] == ["tool_use"], [event.type for event in events] + assert any( + event.type == "content_block_delta" and event.delta.type == "input_json_delta" for event in events + ), [event.type for event in events] + assert final.stop_reason == "tool_use" + block: Final = final.content[0] + assert block.type == "tool_use" and block.name == "get_weather" and block.input == _ARGUMENTS + body: Final = _only_generate(wire) + assert body["stream"] is True + _assert_instructed_once(_prompt_of(body), "get_weather") + assert _spend_row(final.id) == _billed(model) + + +async def test_async_anthropic_sdk_stream_answers_the_tool_result_in_text(gateway: Gateway) -> None: + with _ollama_server(lambda _: _streamed_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + async with _async_anthropic_client(gateway).messages.stream( + model=model, + max_tokens=64, + messages=_anthropic_second_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_ANTHROPIC_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + extra_body={"cache": _NO_CACHE}, + ) as stream: + texts: Final = [event.text async for event in stream if event.type == "text"] + final: Final = await stream.get_final_message() + assert "".join(texts) == _ANSWER + assert final.stop_reason == "end_turn" + assert [(block.type, getattr(block, "text", None)) for block in final.content] == [("text", _ANSWER)] + body: Final = _only_generate(wire) + assert body["stream"] is True + _assert_tool_turn(_prompt_of(body)) + assert _spend_row(final.id) == _billed(model) + + +def test_openai_sdk_responses_tool_request_comes_back_as_a_function_call(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + response: Final = _openai_client(gateway).responses.create( + model=model, + input=_QUESTION, + tools=[_RESPONSES_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + store=False, + extra_body={"cache": _NO_CACHE}, + ) + assert [item.type for item in response.output] == ["function_call"] + item: Final = response.output[0] + assert item.type == "function_call" + assert item.name == "get_weather" + assert json.loads(item.arguments) == _ARGUMENTS + assert response.usage is not None and (response.usage.input_tokens, response.usage.output_tokens) == (30, 12) + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather") + assert f"### User:\n{_QUESTION}\n\n" in prompt, prompt + assert _model_spend_rows(model, 1)[0]["status"] == "success" + + +def test_openai_sdk_responses_function_output_turn_gets_a_message(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + response: Final = _openai_client(gateway).responses.create( + model=model, + input=_responses_second_turn(), # pyright: ignore[reportArgumentType] # plain JSON input items + tools=[_RESPONSES_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + store=False, + extra_body={"cache": _NO_CACHE}, + ) + assert [item.type for item in response.output] == ["message"] + assert response.output_text == _ANSWER + _assert_tool_turn(_prompt_of(_only_generate(wire))) + rows: Final = _model_spend_rows(model, 1) + assert (rows[0]["status"], rows[0]["prompt_tokens"], rows[0]["completion_tokens"]) == ("success", 30, 12) + + +async def test_async_openai_sdk_responses_stream_emits_the_function_call_item(gateway: Gateway) -> None: + with _ollama_server(lambda _: _streamed_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + stream: Final = await _async_openai_client(gateway).responses.create( + model=model, + input=_QUESTION, + tools=[_RESPONSES_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + store=False, + stream=True, + extra_body={"cache": _NO_CACHE}, + ) + events: Final = [event async for event in stream] + done_items: Final = [event.item for event in events if event.type == "response.output_item.done"] + assert [item.type for item in done_items] == ["function_call"], [event.type for event in events] + item: Final = done_items[0] + assert item.type == "function_call" and item.name == "get_weather" + assert json.loads(item.arguments) == _ARGUMENTS + completed: Final = [event for event in events if event.type == "response.completed"] + assert len(completed) == 1 + final_calls: Final = [item for item in completed[0].response.output if item.type == "function_call"] + assert [(item.name, json.loads(item.arguments)) for item in final_calls] == [("get_weather", _ARGUMENTS)] + assert completed[0].response.output_text == "" + body: Final = _only_generate(wire) + assert body["stream"] is True + _assert_instructed_once(_prompt_of(body), "get_weather") + assert _model_spend_rows(model, 1)[0]["status"] == "success" + + +async def test_async_openai_sdk_responses_stream_answers_the_function_output_in_text(gateway: Gateway) -> None: + with _ollama_server(lambda _: _streamed_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + stream: Final = await _async_openai_client(gateway).responses.create( + model=model, + input=_responses_second_turn(), # pyright: ignore[reportArgumentType] # plain JSON input items + tools=[_RESPONSES_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + store=False, + stream=True, + extra_body={"cache": _NO_CACHE}, + ) + events: Final = [event async for event in stream] + assert "".join(event.delta for event in events if event.type == "response.output_text.delta") == _ANSWER + completed: Final = [event for event in events if event.type == "response.completed"] + assert len(completed) == 1 and completed[0].response.output_text == _ANSWER + done_types: Final = [event.type for event in events if event.type == "response.output_item.done"] + assert done_types == ["response.output_item.done"] + body: Final = _only_generate(wire) + assert body["stream"] is True + _assert_tool_turn(_prompt_of(body)) + assert _model_spend_rows(model, 1)[0]["status"] == "success" + + +def test_legacy_functions_param_is_instructed_the_same_way(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + function: Final = _WEATHER_TOOL["function"] + payload: Final = _post_chat(gateway, model, _first_turn(), functions=[function]) + choices: Final = payload["choices"] + assert isinstance(choices, list) and len(choices) == 1 + assert _CALL_JSON_FIELDS <= set(json.dumps(choices[0]).split('"')), choices + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather") + assert "Weather for a city" in prompt, prompt + identity: Final = payload["id"] + assert isinstance(identity, str) + assert _spend_row(identity) == _billed(model) + + +def test_string_system_message_keeps_one_system_section_with_the_instruction(gateway: Gateway) -> None: + system: Final = f"You are a terse weather bot {uuid.uuid4().hex}." + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + _post_chat(gateway, model, [{"role": "system", "content": system}, *_first_turn()], tools=[_WEATHER_TOOL]) + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather") + assert f"### System:\n{system} {_INSTRUCTION}\n" in prompt, prompt + + +def test_list_system_message_keeps_one_system_section_with_the_instruction(gateway: Gateway) -> None: + system: Final = f"You are a terse weather bot {uuid.uuid4().hex}." + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + _post_chat( + gateway, + model, + [{"role": "system", "content": [{"type": "text", "text": system}]}, *_first_turn()], + tools=[_WEATHER_TOOL], + ) + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather") + section: Final = prompt.split("### System:\n", 1)[1] + assert section.startswith(system), section + assert section.count(_INSTRUCTION) == 1, section + + +def test_two_tools_are_both_listed_under_one_instruction(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + _post_chat(gateway, model, _first_turn(), tools=[_WEATHER_TOOL, _TIME_TOOL]) + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather", "get_time") + assert prompt.index("'name': 'get_weather'") < prompt.index("'name': 'get_time'"), prompt + + +def test_unauthenticated_tool_request_never_reaches_ollama(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + code, text = _post( + gateway, + "/v1/chat/completions", + {"model": model, "messages": _first_turn(), "tools": [_WEATHER_TOOL]}, + key=f"sk-not-a-key-{uuid.uuid4().hex}", + ) + assert code == 401, text + assert _generate_calls(wire) == () + + +def test_ollama_model_not_found_reaches_the_caller_after_one_attempt(gateway: Gateway) -> None: + message: Final = f"model '{_BACKEND}' not found {uuid.uuid4().hex}" + reply: Final = Reply(status=404, body=json.dumps({"error": message}).encode()) + with _ollama_server(lambda _: reply) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + code, text = _post( + gateway, "/v1/chat/completions", {"model": model, "messages": _first_turn(), "tools": [_WEATHER_TOOL]} + ) + assert code == 404, text + assert message in text, text + _assert_instructed_once(_prompt_of(_only_generate(wire)), "get_weather") + + +def test_ollama_server_error_on_the_tool_result_turn_does_not_take_the_deployment_down(gateway: Gateway) -> None: + attempts: Final = itertools.count() + failure: Final = f"internal failure {uuid.uuid4().hex}" + + def respond(_: Request) -> Reply: + if next(attempts) == 0: + return Reply(status=500, body=json.dumps({"error": failure}).encode()) + return _generate_reply(_ANSWER) + + with _ollama_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + code, text = _post( + gateway, "/v1/chat/completions", {"model": model, "messages": _second_turn(), "tools": [_WEATHER_TOOL]} + ) + assert code == 500, text + assert failure in text, text + payload: Final = _post_chat(gateway, model, _second_turn(), tools=[_WEATHER_TOOL]) + choices: Final = payload["choices"] + assert isinstance(choices, list) and len(choices) == 1 + assert json.dumps(choices[0]).count(_ANSWER) == 1, choices + prompts: Final = [_prompt_of(_JSON_OBJECT.validate_json(request.body)) for request in _generate_calls(wire)] + assert len(prompts) == 2, prompts + for prompt in prompts: + _assert_tool_turn(prompt) + + +@pytest.mark.parametrize( + ("result", "forwarded"), + [ + pytest.param("", None, id="empty-string-drops-the-section"), + pytest.param("r" * 5120, "r" * 5120, id="5kb-string-forwarded-intact"), + pytest.param( + [{"type": "text", "text": "Paris: 22 degrees"}, {"type": "text", "text": "clear skies"}], + "Paris: 22 degreesclear skies", + id="text-parts-joined", + ), + ], +) +def test_tool_result_content_shapes_reach_the_prompt( + gateway: Gateway, result: JsonValue, forwarded: str | None +) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + payload: Final = _post_chat(gateway, model, _second_turn(result), tools=[_WEATHER_TOOL]) + assert json.dumps(payload["choices"]).count(_ANSWER) == 1, payload + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather") + if forwarded is None: + assert f"### User:\n{_QUESTION}\n\n### Assistant:\n{_CALL_JSON}\n\n### System:\n" in prompt, prompt + else: + _assert_tool_turn(prompt, forwarded) + + +def test_the_same_tool_result_twice_is_forwarded_twice_under_one_instruction(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + messages: Final = [*_second_turn(), {"role": "tool", "tool_call_id": _CALL_ID, "content": _RESULT}] + _post_chat(gateway, model, messages, tools=[_WEATHER_TOOL]) + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather") + assert prompt.count(_RESULT) == 2, prompt + assert prompt.count("### User:") == 2 and prompt.count("### Assistant:") == 1, prompt + + +def test_a_second_function_call_after_the_result_is_surfaced_as_tool_calls(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + completion: Final = _openai_client(gateway).chat.completions.create( + model=model, + messages=_second_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_WEATHER_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + extra_body={"cache": _NO_CACHE}, + ) + choice: Final = completion.choices[0] + assert choice.finish_reason == "tool_calls" + assert choice.message.tool_calls is not None and [call.function.name for call in choice.message.tool_calls] == [ + "get_weather" + ] + _assert_tool_turn(_prompt_of(_only_generate(wire))) + assert _spend_row(completion.id) == _billed(model) + + +def test_a_non_function_json_answer_is_returned_as_text(gateway: Gateway) -> None: + answer: Final = json.dumps({"city": "Paris", "temperature_c": 22, "sky": "clear"}) + with _ollama_server(lambda _: _generate_reply(answer)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + completion: Final = _openai_client(gateway).chat.completions.create( + model=model, + messages=_second_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_WEATHER_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + extra_body={"cache": _NO_CACHE}, + ) + choice: Final = completion.choices[0] + assert choice.finish_reason == "stop" + assert choice.message.tool_calls is None + assert choice.message.content is not None and json.loads(choice.message.content) == json.loads(answer) + _assert_tool_turn(_prompt_of(_only_generate(wire))) + assert _spend_row(completion.id) == _billed(model) + + +def test_int_tool_result_content_fails_in_the_response_body_and_leaves_the_deployment_serving(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + code, text = _post( + gateway, "/v1/chat/completions", {"model": model, "messages": _second_turn(22), "tools": [_WEATHER_TOOL]} + ) + assert code >= 400, text + error: Final = _JSON_OBJECT.validate_json(text)["error"] + assert isinstance(error, dict) and isinstance(error["message"], str) and error["message"], text + assert _generate_calls(wire) == () + payload: Final = _post_chat(gateway, model, _second_turn(), tools=[_WEATHER_TOOL]) + assert json.dumps(payload["choices"]).count(_ANSWER) == 1, payload + _assert_tool_turn(_prompt_of(_only_generate(wire))) + + +def test_ollama_chat_keeps_native_tools_and_gets_no_instruction(gateway: Gateway) -> None: + reply: Final = Reply( + body=json.dumps( + { + "model": _BACKEND, + "created_at": "2026-10-07T00:00:00Z", + "message": {"role": "assistant", "content": _ANSWER}, + "done": True, + "prompt_eval_count": 30, + "eval_count": 12, + } + ).encode() + ) + with wire_server(lambda _: reply) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama_chat/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + completion: Final = _openai_client(gateway).chat.completions.create( + model=model, + messages=_second_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_WEATHER_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + extra_body={"cache": _NO_CACHE}, + ) + assert completion.choices[0].message.content == _ANSWER + received: Final = tuple(request for request in wire.drain() if request.method == "POST") + assert [request.target for request in received] == ["/api/chat"] + body: Final = _JSON_OBJECT.validate_json(received[0].body) + assert "prompt" not in body and "format" not in body, sorted(body) + tools: Final = body["tools"] + assert isinstance(tools, list) and len(tools) == 1 + messages: Final = body["messages"] + assert isinstance(messages, list) and [item["role"] for item in messages if isinstance(item, dict)] == [ + "user", + "assistant", + "tool", + ] + assert "function_name" not in received[0].body.decode(), received[0].body + assert _spend_row(completion.id) == _billed(model) + + +def test_identical_uncached_tool_result_turns_are_each_forwarded_and_billed_once(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + first: Final = _post_chat(gateway, model, _second_turn(), tools=[_WEATHER_TOOL]) + second: Final = _post_chat(gateway, model, _second_turn(), tools=[_WEATHER_TOOL]) + assert first["id"] != second["id"], (first["id"], second["id"]) + prompts: Final = [_prompt_of(_JSON_OBJECT.validate_json(request.body)) for request in _generate_calls(wire)] + assert len(prompts) == 2, prompts + for prompt in prompts: + _assert_tool_turn(prompt) + for payload in (first, second): + identity: Final = payload["id"] + assert isinstance(identity, str) + assert _spend_row(identity) == _billed(model) + + +def test_the_cell_deployment_is_gone_after_its_scenario(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire: + with gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + _post_chat(gateway, model, _first_turn(), tools=[_WEATHER_TOOL]) + _only_generate(wire) + listed: Final = eventually( + lambda: [entry["model_name"] for entry in _deployments(gateway) if isinstance(entry, dict)], + lambda names: model not in names, + seconds=70, + ) + assert model not in listed + + +def _deployments(gateway: Gateway) -> list[JsonValue]: + data: Final = gateway.get("/model/info")["data"] + assert isinstance(data, list) + return data 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/proxy_behavior/spend/test_baseline_accounting.py b/tests/proxy_behavior/spend/test_baseline_accounting.py index dbaf32d579f..4298b85cab4 100644 --- a/tests/proxy_behavior/spend/test_baseline_accounting.py +++ b/tests/proxy_behavior/spend/test_baseline_accounting.py @@ -3,7 +3,8 @@ import json import uuid from collections.abc import AsyncIterator, Callable from contextlib import asynccontextmanager -from datetime import datetime, timezone +from dataclasses import replace +from datetime import datetime, timedelta, timezone from typing import Final import pytest @@ -174,6 +175,10 @@ async def test_commit_ack_loss_and_concurrent_duplicate_delivery_are_idempotent( assert await _store(db, after_commit=True).append(event) == "unavailable" store: Final = _store(db) assert set(await asyncio.gather(*(store.append(event) for _ in range(4)))) == {"recorded"} + duplicate: Final = event.model_copy(update={ + "turn": replace(event.turn, turn_at=event.turn.turn_at + timedelta(seconds=1)) + }) + assert await store.append(duplicate) == "recorded" await _log(db, other) assert await store.append(other) == "recorded" if not attributed: diff --git a/tests/unit/caching/test_caching.py b/tests/unit/caching/test_caching.py index 0a7ac3ecad1..0adb6b9a6f4 100644 --- a/tests/unit/caching/test_caching.py +++ b/tests/unit/caching/test_caching.py @@ -1,6 +1,7 @@ import asyncio import logging import re +import uuid from typing import Final from unittest.mock import MagicMock @@ -9,7 +10,7 @@ import pytest import litellm import litellm.caching.redis_cache as redis_cache_module from litellm._internal_context import current_service_target -from litellm.caching.caching import Cache, response_cache_phase +from litellm.caching.caching import Cache, CacheMode, response_cache_phase from litellm.caching.caching_handler import _PENDING_CACHE_WRITES from litellm.caching.in_memory_cache import InMemoryCache from litellm.caching.redis_cache import RedisCache, _RedisTimeoutLogThrottle @@ -473,3 +474,65 @@ async def test_a_lookup_already_inside_the_phase_does_not_open_a_second_one(v2_s await cache.async_get_cache(dynamic_cache_object=backend, **_REQUEST) assert backend.seen == [("llm_response", "cache.get llm_response")] assert [s.name for s in v2_span_exporter.get_finished_spans()] == ["cache.get llm_response"] + + +_TOOL_TURN_ITEM: Final = {"role": "user", "content": "hi"} + + +@pytest.mark.parametrize( + ("kwargs", "expected"), + [ + pytest.param({"messages": [_TOOL_TURN_ITEM] * 4}, True, id="four-messages-are-cached"), + pytest.param({"messages": [_TOOL_TURN_ITEM] * 5}, False, id="five-messages-skip-the-cache"), + pytest.param({"input": [_TOOL_TURN_ITEM] * 4}, True, id="four-responses-items-are-cached"), + pytest.param({"input": [_TOOL_TURN_ITEM] * 5}, False, id="five-responses-items-skip-the-cache"), + pytest.param({"input": "one prompt"}, True, id="string-input-is-one-message"), + pytest.param({"input": ["a", "b", "c", "d", "e"]}, True, id="embedding-strings-are-not-messages"), + ], +) +def test_should_use_cache_stops_past_the_default_max_messages(kwargs: dict[str, object], expected: bool) -> None: + assert Cache(type=LiteLLMCacheType.LOCAL).should_use_cache(**kwargs) is expected + + +def test_responses_sdk_items_count_toward_max_messages() -> None: + from openai.types.responses import ResponseFunctionToolCall + + call: Final = ResponseFunctionToolCall(type="function_call", call_id="c1", name="ls", arguments="{}") + + assert Cache(type=LiteLLMCacheType.LOCAL).should_use_cache(input=[_TOOL_TURN_ITEM, call, call, call, call]) is False + + +def test_max_messages_is_configurable_and_none_disables_it() -> None: + three: Final = [_TOOL_TURN_ITEM] * 3 + + assert Cache(type=LiteLLMCacheType.LOCAL, max_messages=2).should_use_cache(messages=three) is False + assert Cache(type=LiteLLMCacheType.LOCAL, max_messages=3).should_use_cache(messages=three) is True + assert Cache(type=LiteLLMCacheType.LOCAL, max_messages=None).should_use_cache(messages=three * 50) is True + + +def test_max_messages_beats_an_explicit_use_cache_opt_in() -> None: + cache: Final = Cache(type=LiteLLMCacheType.LOCAL, mode=CacheMode.default_off) + + assert cache.should_use_cache(messages=[_TOOL_TURN_ITEM] * 4, cache={"use-cache": True}) is True + assert cache.should_use_cache(messages=[_TOOL_TURN_ITEM] * 5, cache={"use-cache": True}) is False + + +def test_completion_past_max_messages_is_neither_served_from_nor_written_to_the_cache( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(litellm, "cache", Cache(type=LiteLLMCacheType.LOCAL)) + tag: Final = uuid.uuid4().hex + four: Final = [{"role": "user", "content": f"{tag} turn {index}"} for index in range(4)] + five: Final = [*four, {"role": "user", "content": f"{tag} turn 4"}] + + def answer(messages: list[dict[str, str]], mock_response: str) -> str: + response: Final = litellm.completion(model="gpt-4o-mini", messages=messages, mock_response=mock_response) + assert isinstance(response, litellm.ModelResponse), response + choice: Final = response.choices[0] + assert isinstance(choice, litellm.Choices), choice + return str(choice.message.content) + + assert answer(four, "four first") == "four first" + assert answer(four, "four second") == "four first", "a 4-message repeat missed the cache" + assert answer(five, "five first") == "five first" + assert answer(five, "five second") == "five second", "a 5-message repeat was served from the cache" diff --git a/tests/unit/caching/test_redis_semantic_cache.py b/tests/unit/caching/test_redis_semantic_cache.py index 99844c695cb..93db26fc9b2 100644 --- a/tests/unit/caching/test_redis_semantic_cache.py +++ b/tests/unit/caching/test_redis_semantic_cache.py @@ -568,6 +568,20 @@ def test_redis_semantic_cache_set_cache_flattens_structured_responses_input(): ) +def test_redis_semantic_cache_prompt_extraction_reads_function_call_output_blocks(): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + prompt = RedisSemanticCache._get_prompt_from_kwargs( + input=[ + {"role": "user", "content": "update the config"}, + {"type": "function_call", "call_id": "c1", "name": "write_file", "arguments": '{"path": "a"}'}, + {"type": "function_call_output", "call_id": "c1", "output": [{"type": "input_text", "text": "wrote a"}]}, + ] + ) + + assert prompt == "update the config\nwrote a" + + def test_redis_semantic_cache_prompt_extraction_prefers_messages(): from litellm.caching.redis_semantic_cache import RedisSemanticCache @@ -1416,3 +1430,23 @@ async def test_redis_async_embedding_truncates_off_the_event_loop(monkeypatch): assert embedding == [0.1, 0.2] assert _token_count("sem-embed", router.aembedding.call_args.kwargs["input"]) == 5 assert_loop_stayed_free(took, lags) + + +def test_redis_semantic_cache_prompt_extraction_keeps_tool_result_text(): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + prompt = RedisSemanticCache._get_prompt_from_kwargs( + messages=[ + {"role": "user", "content": "list the files"}, + { + "role": "assistant", + "content": [{"type": "tool_use", "id": "toolu_1", "name": "Bash", "input": {"command": "ls"}}], + }, + { + "role": "user", + "content": [{"type": "tool_result", "tool_use_id": "toolu_1", "content": "calc.py test_calc.py"}], + }, + ] + ) + + assert prompt == "list the filescalc.py test_calc.py" diff --git a/tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py b/tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py index f1c80bb11ea..78545d3fb62 100644 --- a/tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py +++ b/tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py @@ -433,6 +433,7 @@ def test_set_latency_metrics(prometheus_logger): team_alias="test_team_alias", org_id=None, org_alias=None, + model_group="openai-gpt", requested_model="openai-gpt", model="gpt-5-mini", model_id="model-123", @@ -453,6 +454,7 @@ def test_set_latency_metrics(prometheus_logger): team_alias="test_team_alias", org_id=None, org_alias=None, + model_group="openai-gpt", requested_model="openai-gpt", model="gpt-5-mini", model_id="model-123", @@ -473,6 +475,7 @@ def test_set_latency_metrics(prometheus_logger): team_alias="test_team_alias", org_id=None, org_alias=None, + model_group="openai-gpt", requested_model="openai-gpt", model="gpt-5-mini", model_id="model-123", @@ -1054,6 +1057,7 @@ def test_set_llm_deployment_success_metrics(prometheus_logger): # Verify latency per output token metric prometheus_logger.litellm_deployment_latency_per_output_token.labels.assert_called_once_with( + model_group="my_custom_model_group", litellm_model_name="gpt-5-mini", model_id="model-123", api_base="https://api.openai.com", diff --git a/tests/unit/integration_support/test_process.py b/tests/unit/integration_support/test_process.py index 042a2447dff..7896786fce8 100644 --- a/tests/unit/integration_support/test_process.py +++ b/tests/unit/integration_support/test_process.py @@ -1,8 +1,10 @@ from __future__ import annotations +import asyncio import errno import importlib -import os +import socket +from collections.abc import Iterator from pathlib import Path from types import ModuleType from typing import Final @@ -10,7 +12,6 @@ from typing import Final import pytest TESTS_DIR: Final = Path(__file__).resolve().parents[2] -BIND_ERROR_LINE: Final = f"ERROR: {OSError(errno.EADDRINUSE, os.strerror(errno.EADDRINUSE))}\n" UNRELATED_CRASH: Final = "Traceback (most recent call last):\nModuleNotFoundError: No module named 'litellm'\n" @@ -20,19 +21,41 @@ def process_module(monkeypatch: pytest.MonkeyPatch) -> ModuleType: return importlib.import_module("integration._support.process") +async def _refused_bind(port: int) -> OSError | None: + try: + await asyncio.get_running_loop().create_server(asyncio.Protocol, "127.0.0.1", port) + except OSError as refused: + return refused + return None + + +@pytest.fixture +def bind_error_line() -> Iterator[str]: + with socket.socket() as held: + held.bind(("127.0.0.1", 0)) + held.listen() + refused: Final = asyncio.run(_refused_bind(held.getsockname()[1])) + assert refused is not None and refused.errno == errno.EADDRINUSE + yield f"ERROR: {refused}\n" + + def _written_log(directory: Path, text: str) -> Path: log: Final = directory / "owned-proxy.log" log.write_text(text) return log -def test_lost_port_race_matches_the_bind_error_the_server_logs(process_module: ModuleType, tmp_path: Path) -> None: - assert process_module._lost_port_race(1, _written_log(tmp_path, BIND_ERROR_LINE)) +def test_lost_port_race_matches_the_bind_error_the_server_logs( + process_module: ModuleType, bind_error_line: str, tmp_path: Path +) -> None: + assert process_module._lost_port_race(1, _written_log(tmp_path, bind_error_line)) def test_lost_port_race_ignores_an_exit_for_another_reason(process_module: ModuleType, tmp_path: Path) -> None: assert not process_module._lost_port_race(1, _written_log(tmp_path, UNRELATED_CRASH)) -def test_lost_port_race_needs_the_process_to_have_exited(process_module: ModuleType, tmp_path: Path) -> None: - assert not process_module._lost_port_race(None, _written_log(tmp_path, BIND_ERROR_LINE)) +def test_lost_port_race_needs_the_process_to_have_exited( + process_module: ModuleType, bind_error_line: str, tmp_path: Path +) -> None: + assert not process_module._lost_port_race(None, _written_log(tmp_path, bind_error_line)) 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/integrations/test_custom_guardrail.py b/tests/unit/integrations/test_custom_guardrail.py index 7bfdfb00faf..72c362425bb 100644 --- a/tests/unit/integrations/test_custom_guardrail.py +++ b/tests/unit/integrations/test_custom_guardrail.py @@ -12,7 +12,7 @@ from litellm.integrations.custom_guardrail import ( ) from litellm.litellm_core_utils.litellm_logging import Logging from litellm.proxy._types import CallTypes, UserAPIKeyAuth -from litellm.types.guardrails import GuardrailEventHooks, Mode +from litellm.types.guardrails import GuardrailEventHooks, LoggingOnlyScope, Mode from litellm.types.utils import ( Choices, GenericGuardrailAPIInputs, @@ -2577,6 +2577,32 @@ class TestLoggingOnlyApplyGuardrail: assert "standard_logging_guardrail_information" not in kwargs["litellm_params"]["metadata"] assert kwargs["standard_logging_object"] == {"guardrail_information": None} + @pytest.mark.parametrize( + "scope,expected_calls", + ( + (None, [("request", ["hello there"]), ("response", ["general kenobi"])]), + ("both", [("request", ["hello there"]), ("response", ["general kenobi"])]), + ("input", [("request", ["hello there"])]), + ("output", [("response", ["general kenobi"])]), + ), + ) + @pytest.mark.asyncio + async def test_logging_only_scope_scans_configured_directions( + self, + scope: LoggingOnlyScope | None, + expected_calls: list[tuple[str, list[str]]], + ) -> None: + guardrail: Final = _ApplyOnlyObserver() + guardrail.logging_only_scope = scope + kwargs, response = _logged_call([{"role": "user", "content": "hello there"}]) + + out_kwargs, _ = await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value) + + assert guardrail.calls == expected_calls + entries: Final = out_kwargs["standard_logging_object"]["guardrail_information"] + assert len(entries) == len(expected_calls) + assert [entry["guardrail_mode"] for entry in entries] == ["logging_only"] * len(expected_calls) + @pytest.mark.asyncio async def test_appends_to_pre_call_verdicts_without_duplicating_them(self): guardrail = _ApplyOnlyObserver() @@ -2604,6 +2630,39 @@ class TestLoggingOnlyApplyGuardrail: assert out_kwargs is kwargs assert out_response is response + @pytest.mark.asyncio + async def test_output_scope_scans_response_when_request_copy_fails(self): + import threading + + guardrail: Final = _ApplyOnlyObserver() + guardrail.logging_only_scope = "output" + call: Final = _logged_call([{"role": "user", "content": "hello there", "lock": threading.Lock()}]) + kwargs: Final = call[0] + response: Final = call[1] + + out_kwargs, _ = await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value) + + assert guardrail.calls == [("response", ["general kenobi"])] + entries: Final = out_kwargs["standard_logging_object"]["guardrail_information"] + assert [entry["guardrail_name"] for entry in entries] == ["apply-only-observer"] + assert [entry["guardrail_status"] for entry in entries] == ["success"] + + @pytest.mark.asyncio + async def test_request_copy_failure_does_not_drop_the_response_scan_for_explicit_both_scope(self): + import threading + + guardrail: Final = _ApplyOnlyObserver() + guardrail.logging_only_scope = "both" + call: Final = _logged_call([{"role": "user", "content": "hello there", "lock": threading.Lock()}]) + kwargs: Final = call[0] + response: Final = call[1] + + out_kwargs, _ = await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value) + + assert guardrail.calls == [("response", ["general kenobi"])] + entries: Final = out_kwargs["standard_logging_object"]["guardrail_information"] + assert [entry["guardrail_status"] for entry in entries] == ["success"] + @pytest.mark.asyncio async def test_block_verdict_is_recorded_without_raising(self): guardrail = _ApplyOnlyObserver(block=True) @@ -2615,6 +2674,45 @@ class TestLoggingOnlyApplyGuardrail: entries = out_kwargs["standard_logging_object"]["guardrail_information"] assert [e["guardrail_status"] for e in entries] == ["guardrail_intervened"] + @pytest.mark.asyncio + async def test_input_scan_error_aborts_the_response_scan(self): + class _FailingObserver(_ApplyOnlyObserver): + @log_guardrail_information + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): + self.calls.append((input_type, list(inputs.get("texts") or []))) + if input_type == "request": + raise RuntimeError("guardrail service unavailable") + return GenericGuardrailAPIInputs(texts=[]) + + guardrail = _FailingObserver() + kwargs, response = _logged_call([{"role": "user", "content": "hello there"}]) + + out_kwargs, _ = await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value) + + assert guardrail.calls == [("request", ["hello there"])] + entries = out_kwargs["standard_logging_object"]["guardrail_information"] + assert [e["guardrail_status"] for e in entries] == ["guardrail_failed_to_respond"] + + @pytest.mark.asyncio + async def test_input_scan_error_does_not_drop_the_response_scan_for_explicit_both_scope(self): + class _FailingBothObserver(_ApplyOnlyObserver): + @log_guardrail_information + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): + self.calls.append((input_type, list(inputs.get("texts") or []))) + if input_type == "request": + raise RuntimeError("guardrail service unavailable") + return GenericGuardrailAPIInputs(texts=[]) + + guardrail = _FailingBothObserver() + guardrail.logging_only_scope = "both" + kwargs, response = _logged_call([{"role": "user", "content": "hello there"}]) + + out_kwargs, _ = await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value) + + assert guardrail.calls == [("request", ["hello there"]), ("response", ["general kenobi"])] + entries = out_kwargs["standard_logging_object"]["guardrail_information"] + assert [e["guardrail_status"] for e in entries] == ["guardrail_failed_to_respond", "success"] + @pytest.mark.asyncio async def test_call_type_without_translation_is_skipped(self): guardrail = _ApplyOnlyObserver() @@ -2932,6 +3030,17 @@ class _NativeLifecycleLoggingGuardrail(CustomGuardrail): return inputs +@pytest.mark.asyncio +async def test_native_lifecycle_guardrail_logging_only_scope_scans_only_input(): + guardrail: Final = _NativeLifecycleLoggingGuardrail() + guardrail.logging_only_scope = "input" + kwargs, response = _logged_call([{"role": "user", "content": "native lifecycle input"}]) + + await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value) + + assert guardrail.calls == [("request", ["native lifecycle input"])] + + @pytest.mark.asyncio async def test_native_lifecycle_guardrail_logging_only_scans_assembled_response(): """A use_native_lifecycle_hooks guardrail accepts mode logging_only and its diff --git a/tests/unit/integrations/test_prometheus_labels.py b/tests/unit/integrations/test_prometheus_labels.py index 8a1f5f5a0a8..ead0ce8ecff 100644 --- a/tests/unit/integrations/test_prometheus_labels.py +++ b/tests/unit/integrations/test_prometheus_labels.py @@ -980,6 +980,193 @@ def test_deployment_tpm_rpm_limit_metrics_emit_model_group_from_enum_values(): _clear_prometheus_registry() +def test_model_group_in_latency_metrics(): + """ + Test that model_group label is present on the end-to-end / per-call + latency metrics needed to build model-group latency dashboards. These + metrics previously only carried requested_model, litellm_model_name and + model_id, none of which identify the model_group a pooled deployment + belongs to -- only the proxy-overhead-only latency metrics + (litellm_overhead_latency_metric and friends) carried model_group. + """ + model_group_label = UserAPIKeyLabelNames.MODEL_GROUP.value + + metrics_with_model_group = [ + "litellm_llm_api_latency_metric", + "litellm_llm_api_time_to_first_token_metric", + "litellm_request_total_latency_metric", + "litellm_deployment_latency_per_output_token", + ] + + for metric_name in metrics_with_model_group: + labels = PrometheusMetricLabels.get_labels(metric_name) + assert ( + model_group_label in labels + ), f"Metric {metric_name} should contain model_group label" + print(f"✅ {metric_name} contains model_group label") + + +def test_model_group_value_flows_through_latency_metrics_label_factory(): + """ + The label being in the allow-list is necessary but not sufficient: the + factory must also carry the value from the enum through to the emitted + label. This would fail if the label were dropped from a metric's list or + if the value plumbing regressed, which the allow-list assertion above + cannot catch on its own. + """ + from unittest.mock import MagicMock + + from litellm.integrations.prometheus import ( + PrometheusLogger, + UserAPIKeyLabelValues, + prometheus_label_factory, + ) + + prometheus_logger = MagicMock() + prometheus_logger._cached_metric_labels = {} + prometheus_logger.label_filters = {} + prometheus_logger.get_labels_for_metric = ( + PrometheusLogger.get_labels_for_metric.__get__(prometheus_logger) + ) + + enum_values = UserAPIKeyLabelValues( + model_group="example-model-group", + litellm_model_name="gpt-4o-mini", + requested_model="example-model-group", + status_code="200", + ) + + for metric_name in [ + "litellm_llm_api_latency_metric", + "litellm_llm_api_time_to_first_token_metric", + "litellm_request_total_latency_metric", + "litellm_deployment_latency_per_output_token", + ]: + labels = prometheus_label_factory( + supported_enum_labels=prometheus_logger.get_labels_for_metric( + metric_name=metric_name + ), + enum_values=enum_values, + ) + assert ( + labels.get("model_group") == "example-model-group" + ), f"{metric_name} should emit model_group=example-model-group, got {labels.get('model_group')!r}" + + +def test_latency_metrics_emit_model_group_from_set_latency_metrics(): + """ + End-to-end emit wiring for _set_latency_metrics. + + The label-list and factory tests above prove the label exists and that + the factory carries a value handed to it, but neither drives the real + _set_latency_metrics code path, so deleting the production model_group + plumbing there would still pass them. This calls it directly with a + streaming request (so the time-to-first-token branch also fires) and + asserts the real litellm_llm_api_latency_metric, + litellm_llm_api_time_to_first_token_metric and + litellm_request_total_latency_metric Histogram series actually carry it. + """ + import datetime + + from litellm.integrations.prometheus import PrometheusLogger, UserAPIKeyLabelValues + + _clear_prometheus_registry() + try: + logger = PrometheusLogger() + start_time = datetime.datetime(2024, 1, 1, 0, 0, 0) + api_call_start_time = datetime.datetime(2024, 1, 1, 0, 0, 1) + completion_start_time = datetime.datetime(2024, 1, 1, 0, 0, 2) + end_time = datetime.datetime(2024, 1, 1, 0, 0, 3) + + enum_values = UserAPIKeyLabelValues( + model_group="example-model-group", + litellm_model_name="gpt-4o-mini", + requested_model="example-model-group", + status_code="200", + ) + + logger._set_latency_metrics( + kwargs={ + "start_time": start_time, + "end_time": end_time, + "api_call_start_time": api_call_start_time, + "completion_start_time": completion_start_time, + "stream": True, + "litellm_params": {"metadata": {}}, + }, + model="gpt-4o-mini", + user_api_key=None, + user_api_key_alias=None, + user_api_team=None, + user_api_team_alias=None, + enum_values=enum_values, + ) + + for metric in ( + logger.litellm_llm_api_latency_metric, + logger.litellm_llm_api_time_to_first_token_metric, + logger.litellm_request_total_latency_metric, + ): + index = metric._labelnames.index("model_group") + values = {sample_key[index] for sample_key in metric._metrics} + assert values == {"example-model-group"}, ( + f"expected model_group=example-model-group on {metric._name}, got {values}" + ) + finally: + _clear_prometheus_registry() + + +def test_deployment_latency_per_output_token_emits_model_group_from_enum_values(): + """ + End-to-end emit wiring for litellm_deployment_latency_per_output_token. + + Drives set_llm_deployment_success_metrics directly (its only caller) with + output_tokens > 0 so the latency-per-token branch fires, and asserts the + real Histogram series carries model_group; fails if that label-list + addition or the enum_values plumbing is removed. + """ + import datetime + + from litellm.integrations.prometheus import PrometheusLogger, UserAPIKeyLabelValues + + _clear_prometheus_registry() + try: + logger = PrometheusLogger() + start_time = datetime.datetime(2024, 1, 1, 0, 0, 0) + end_time = datetime.datetime(2024, 1, 1, 0, 0, 2) + enum_values = UserAPIKeyLabelValues( + model_group="example-model-group", + litellm_model_name="gpt-4o-mini", + requested_model="example-model-group", + status_code="200", + ) + logger.set_llm_deployment_success_metrics( + request_kwargs={ + "model": "gpt-4o-mini", + "litellm_params": {"metadata": {"model_info": {"id": "model-123"}}}, + "standard_logging_object": { + "model_group": "example-model-group", + "model_id": "model-123", + "api_base": "https://api.openai.com", + "hidden_params": {"additional_headers": None, "litellm_overhead_time_ms": None}, + }, + }, + start_time=start_time, + end_time=end_time, + enum_values=enum_values, + output_tokens=10.0, + ) + + metric = logger.litellm_deployment_latency_per_output_token + index = metric._labelnames.index("model_group") + values = {sample_key[index] for sample_key in metric._metrics} + assert values == {"example-model-group"}, ( + f"expected model_group=example-model-group on {metric._name}, got {values}" + ) + finally: + _clear_prometheus_registry() + + if __name__ == "__main__": test_user_email_in_required_metrics() test_user_email_label_exists() diff --git a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py index 7415c74226d..1d19264fb9e 100644 --- a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py +++ b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py @@ -16,6 +16,8 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( encrypted_reasoning_signature, get_file_ids_from_messages, get_format_from_file_id, + get_semantic_cache_prompt_from_messages, + get_str_from_messages, handle_any_messages_to_chat_completion_str_messages_conversion, hoist_images_from_tool_messages, is_encrypted_reasoning_block, @@ -2171,3 +2173,93 @@ class TestMergeConsecutiveSystemMessages: ) assert merged == [{"role": "system"}, {"role": "user", "content": "Hi"}] + + +_CLAUDE_CODE_TOOL_TURN: Final = [ + {"role": "user", "content": "list the files"}, + { + "role": "assistant", + "content": [{"type": "tool_use", "id": "toolu_1", "name": "Bash", "input": {"command": "ls"}}], + }, + { + "role": "user", + "content": [{"type": "tool_result", "tool_use_id": "toolu_1", "content": "calc.py test_calc.py"}], + }, +] + + +@pytest.mark.parametrize( + ("messages", "expected"), + [ + pytest.param(_CLAUDE_CODE_TOOL_TURN, "list the filescalc.py test_calc.py", id="tool-result-string"), + pytest.param( + [ + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "toolu_1", + "content": [ + {"type": "text", "text": "x = 1"}, + {"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": ""}}, + ], + } + ], + } + ], + "x = 1", + id="tool-result-blocks", + ), + pytest.param( + [{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_1"}]}], + "", + id="tool-result-without-content", + ), + ], +) +def test_get_semantic_cache_prompt_from_messages_keeps_tool_result_text( + messages: list[dict[str, object]], expected: str +) -> None: + assert get_semantic_cache_prompt_from_messages(messages) == expected + + +def test_get_semantic_cache_prompt_from_messages_differs_from_the_turn_before_it() -> None: + assert get_str_from_messages(_CLAUDE_CODE_TOOL_TURN) == get_str_from_messages(_CLAUDE_CODE_TOOL_TURN[:1]) + assert get_semantic_cache_prompt_from_messages(_CLAUDE_CODE_TOOL_TURN) != get_semantic_cache_prompt_from_messages( + _CLAUDE_CODE_TOOL_TURN[:1] + ) + + +@pytest.mark.parametrize( + "messages", + [ + pytest.param([{"role": "system", "content": "be brief. "}, {"role": "user", "content": "hello"}], id="strings"), + pytest.param( + [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is "}, + {"type": "image_url", "image_url": {"url": "https://example.com/a.png"}}, + {"type": "text", "text": "this?"}, + ], + } + ], + id="text-parts", + ), + pytest.param( + [ + {"role": "assistant"}, + {"role": "assistant", "content": None}, + {"role": "user", "content": ""}, + {"role": "tool", "content": "small", "search_results": [{"source": "s", "title": "t", "content": []}]}, + ], + id="empty-content-and-search-results", + ), + ], +) +def test_get_semantic_cache_prompt_from_messages_matches_get_str_from_messages_without_tool_results( + messages: list[dict[str, object]], +) -> None: + assert get_semantic_cache_prompt_from_messages(messages) == get_str_from_messages(messages) diff --git a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index a2987730851..61b59f7dbe6 100644 --- a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -51,6 +51,17 @@ def test_function_call_prompt_preserves_append_failure_for_non_string_content() function_call_prompt(messages, []) +def test_function_call_prompt_lets_the_model_answer_after_a_function_result() -> None: + messages: Final[list[dict[str, object]]] = [{"role": "system", "content": "Be terse."}] + + prompted: Final = function_call_prompt(messages, [{"name": "get_weather"}]) + + system: Final = str(prompted[0]["content"]) + assert "JSON OUTPUT ONLY" not in system + assert "reply to the user in plain text instead of calling a function again" in system + assert "{'name': 'get_weather'}" in system + + @pytest.mark.parametrize( ("thought_signature", "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_logging_worker.py b/tests/unit/litellm_core_utils/test_logging_worker.py index 5d4c9e65d9b..f48495d1062 100644 --- a/tests/unit/litellm_core_utils/test_logging_worker.py +++ b/tests/unit/litellm_core_utils/test_logging_worker.py @@ -6,6 +6,7 @@ import asyncio import contextvars import io import logging +from typing import Final from unittest.mock import AsyncMock, patch import pytest @@ -14,6 +15,61 @@ from litellm.constants import LOGGING_WORKER_AGGRESSIVE_CLEAR_COOLDOWN_SECONDS from litellm.litellm_core_utils.logging_worker import LoggingWorker +@pytest.mark.asyncio +@pytest.mark.parametrize("dispatch", ("worker", "flush", "extracted")) +async def test_optional_work_budget_preserves_callback_context_and_reserves_logging_time(dispatch: str) -> None: + from litellm.litellm_core_utils.logging_worker import optional_callback_budget + + worker: Final = LoggingWorker(timeout=1.0) + identity: Final = contextvars.ContextVar("test_callback_identity", default="outside") + results: Final[asyncio.Queue[tuple[str, float]]] = asyncio.Queue() + + async def callback() -> None: + results.put_nowait((identity.get(), optional_callback_budget(3.0))) + + token: Final = identity.set("request") + worker._ensure_queue() + worker.enqueue(callback()) + identity.reset(token) + try: + if dispatch == "worker": + worker.start() + elif dispatch == "flush": + await worker.flush() + else: + assert worker._queue is not None + await worker._process_single_task(worker._queue.get_nowait()) + restored_identity, budget = await asyncio.wait_for(results.get(), timeout=2) + assert restored_identity == "request" + assert 0 < budget <= worker.timeout / 4 + assert identity.get() == "outside" + assert optional_callback_budget(3.0) == 3.0 + finally: + await worker.stop() + + +def test_exit_flush_bounds_optional_work_and_restores_callers_budget() -> None: + from queue import SimpleQueue + + from litellm.litellm_core_utils.logging_worker import optional_callback_budget + + worker: Final = LoggingWorker(timeout=1.0) + observed: Final[SimpleQueue[float]] = SimpleQueue() + + async def callback() -> None: + observed.put(optional_callback_budget(3.0)) + + async def enqueue() -> None: + worker._ensure_queue() + worker.enqueue(callback()) + + asyncio.run(enqueue()) + worker._flush_on_exit() + assert observed.qsize() == 1 + assert 0 < observed.get_nowait() <= worker.timeout / 4 + assert optional_callback_budget(3.0) == 3.0 + + class _RecordCollector(logging.Handler): """Captures emitted log records so a test can assert on real logging output (level, message args, traceback) instead of patching the logger object.""" @@ -205,7 +261,9 @@ class TestLoggingWorker: asyncio.run(log_on_second_loop()) first_loop.run_until_complete(asyncio.sleep(0.1)) failures = [ - task.exception() for task in first_loop_tasks if task.done() and not task.cancelled() and task.exception() + task.exception() + for task in first_loop_tasks + if task.done() and not task.cancelled() and task.exception() ] finally: first_loop.close() 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/litellm_core_utils/test_redact_messages.py b/tests/unit/litellm_core_utils/test_redact_messages.py index 3bb6b379873..6fbcba170f1 100644 --- a/tests/unit/litellm_core_utils/test_redact_messages.py +++ b/tests/unit/litellm_core_utils/test_redact_messages.py @@ -5,14 +5,26 @@ Covers the proxy flow where headers arrive in litellm_params["metadata"]["header but litellm_params["litellm_metadata"] is None. """ -import asyncio, httpx, importlib, json, os, pytest_asyncio, threading +import asyncio +import importlib +import json +import os +import threading +from collections.abc import AsyncIterator, Mapping +from datetime import datetime +from types import MappingProxyType, SimpleNamespace from typing import Final, Optional, Union -from types import SimpleNamespace +from unittest.mock import patch +import httpx import pytest +import pytest_asyncio +from pydantic import JsonValue import litellm +from litellm.constants import LOGGING_WORKER_MAX_TIME_PER_COROUTINE from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.litellm_core_utils.redact_messages import ( _redact_responses_api_output, perform_redaction, @@ -21,18 +33,14 @@ from litellm.litellm_core_utils.redact_messages import ( should_redact_message_logging, ) from litellm.responses.main import mock_responses_api_response -from collections.abc import AsyncIterator -from datetime import datetime -from litellm.constants import LOGGING_WORKER_MAX_TIME_PER_COROUTINE -from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER -from litellm.types.utils import( +from litellm.types.router import BaselineRouteStamp +from litellm.types.utils import ( ModelResponse, ResponsesAPIResponse, StandardLoggingPayload, TextCompletionResponse, ) from tests._vcr_conftest_common import install_live_call_probe, record_vcr_outcome -from unittest.mock import patch @pytest.fixture(autouse=True) @@ -113,9 +121,7 @@ class TestShouldRedactMessageLogging: def test_enable_redaction_via_header_in_litellm_metadata(self): """Headers inside litellm_metadata (SDK direct call) should work.""" details = _make_model_call_details( - litellm_metadata={ - "headers": {"x-litellm-enable-message-redaction": "true"} - }, + litellm_metadata={"headers": {"x-litellm-enable-message-redaction": "true"}}, ) assert should_redact_message_logging(details) is True @@ -217,21 +223,15 @@ class TestPerformRedaction: redacted = perform_redaction(details, result) - assert details["messages"] == [ - {"role": "user", "content": "redacted-by-litellm"} - ] + assert details["messages"] == [{"role": "user", "content": "redacted-by-litellm"}] assert details["prompt"] == "" assert details["input"] == "" logged_response = details["standard_logging_object"]["response"] assert logged_response["usage"] == {"total_tokens": 1} assert logged_response["output"][0]["text"] == "redacted-by-litellm" - assert logged_response["output"][1]["content"][0]["text"] == ( - "redacted-by-litellm" - ) - assert logged_response["output"][2]["summary"][0]["text"] == ( - "redacted-by-litellm" - ) + assert logged_response["output"][1]["content"][0]["text"] == ("redacted-by-litellm") + assert logged_response["output"][2]["summary"][0]["text"] == ("redacted-by-litellm") assert redacted["usage"] == {"total_tokens": 1} assert redacted["output"][0]["text"] == "redacted-by-litellm" @@ -444,9 +444,7 @@ class TestPerformRedaction: tool_call = redacted.choices[0].message.tool_calls[0] assert tool_call.function.arguments == "redacted-by-litellm" assert tool_call.function.name == "get_weather" - assert result.choices[0].message.tool_calls[0].function.arguments == ( - '{"city": "sensitive-city"}' - ) + assert result.choices[0].message.tool_calls[0].function.arguments == ('{"city": "sensitive-city"}') def test_redacts_tool_call_arguments_on_streaming_response_object(self): """Reproduces the Stream=True path where tool calls arrive as deltas.""" @@ -714,12 +712,8 @@ class TestPerformRedaction: } } ], - "vertex_ai_grounding_metadata": [ - {"webSearchQueries": ["sensitive search term"]} - ], - "vertex_ai_url_context_metadata": [ - {"urlMetadata": [{"retrievedUrl": "https://example.com"}]} - ], + "vertex_ai_grounding_metadata": [{"webSearchQueries": ["sensitive search term"]}], + "vertex_ai_url_context_metadata": [{"urlMetadata": [{"retrievedUrl": "https://example.com"}]}], }, } } @@ -749,9 +743,7 @@ class TestPerformRedaction: "vertex_ai_grounding_metadata", [{"webSearchQueries": ["sensitive search term"]}], ) - response._hidden_params["vertex_ai_grounding_metadata"] = [ - {"webSearchQueries": ["sensitive search term"]} - ] + response._hidden_params["vertex_ai_grounding_metadata"] = [{"webSearchQueries": ["sensitive search term"]}] details = { "stream": True, @@ -772,12 +764,8 @@ class TestPerformRedaction: "metadata": { "hidden_params": { "response_cost": 0.01, - "vertex_ai_grounding_metadata": [ - {"webSearchQueries": ["sensitive search term"]} - ], - "vertex_ai_url_context_metadata": [ - {"urlMetadata": [{"retrievedUrl": "https://example.com"}]} - ], + "vertex_ai_grounding_metadata": [{"webSearchQueries": ["sensitive search term"]}], + "vertex_ai_url_context_metadata": [{"urlMetadata": [{"retrievedUrl": "https://example.com"}]}], "vertex_ai_safety_ratings": [{"category": "HARM"}], "vertex_ai_citation_metadata": [{"citations": ["source"]}], } @@ -797,11 +785,7 @@ class TestPerformRedaction: def test_redact_async_complete_streaming_response(self): """Test that async_complete_streaming_response is properly redacted.""" response_obj = litellm.ModelResponse( - choices=[ - litellm.Choices( - message=litellm.Message(content="secret content", role="assistant") - ) - ] + choices=[litellm.Choices(message=litellm.Message(content="secret content", role="assistant"))] ) model_call_details = { @@ -820,11 +804,7 @@ class TestPerformRedaction: def test_redact_complete_streaming_response(self): """Test that complete_streaming_response is properly redacted.""" response_obj = litellm.ModelResponse( - choices=[ - litellm.Choices( - message=litellm.Message(content="secret content", role="assistant") - ) - ] + choices=[litellm.Choices(message=litellm.Message(content="secret content", role="assistant"))] ) model_call_details = { @@ -842,11 +822,7 @@ class TestPerformRedaction: def test_streaming_responses_untouched_when_disabled(self): response_obj = litellm.ModelResponse( - choices=[ - litellm.Choices( - message=litellm.Message(content="secret content", role="assistant") - ) - ] + choices=[litellm.Choices(message=litellm.Message(content="secret content", role="assistant"))] ) model_call_details = { @@ -909,11 +885,7 @@ class TestPerformRedaction: class TestRedactStreamingResponsesForCustomLogger: def _model_call_details(self): response_obj = litellm.ModelResponse( - choices=[ - litellm.Choices( - message=litellm.Message(content="secret content", role="assistant") - ) - ] + choices=[litellm.Choices(message=litellm.Message(content="secret content", role="assistant"))] ) return { "stream": True, @@ -947,7 +919,10 @@ class TestRedactStreamingResponsesForCustomLogger: @pytest.mark.parametrize("callback_only", [False, True]) def test_classifier_audit_redaction_removes_both_fields_and_source_carrier(callback_only: bool) -> None: - audit: Final = {"classifier_input": {"system": "private rubric"}, "originating_request_masked": {"input": "private source"}} + audit: Final = { + "classifier_input": {"system": "private rubric"}, + "originating_request_masked": {"input": "private source"}, + } standard_payload: Final = { **audit, "messages": [{"role": "user", "content": "private prompt"}], @@ -956,7 +931,9 @@ def test_classifier_audit_redaction_removes_both_fields_and_source_carrier(callb } details: Final = { "standard_logging_object": standard_payload, - "litellm_params": {"proxy_server_request": {"body": {}, "originating_request_masked": audit["originating_request_masked"]}}, + "litellm_params": { + "proxy_server_request": {"body": {}, "originating_request_masked": audit["originating_request_masked"]} + }, } logger: Final = CustomLogger() logger.turn_off_message_logging = True @@ -966,7 +943,10 @@ def test_classifier_audit_redaction_removes_both_fields_and_source_carrier(callb assert "originating_request_masked" not in redacted["standard_logging_object"] assert "originating_request_masked" not in redacted["litellm_params"]["proxy_server_request"] assert details["standard_logging_object"]["classifier_input"] == audit["classifier_input"] - assert details["litellm_params"]["proxy_server_request"]["originating_request_masked"] == audit["originating_request_masked"] + assert ( + details["litellm_params"]["proxy_server_request"]["originating_request_masked"] + == audit["originating_request_masked"] + ) else: perform_redaction(details, result=None) assert "classifier_input" not in details["standard_logging_object"] @@ -981,7 +961,9 @@ def test_classifier_audit_redaction_removes_both_fields_and_source_carrier(callb @pytest.mark.parametrize("excluded", [False, True]) def test_classifier_callback_redaction_preserves_exclusions(monkeypatch: pytest.MonkeyPatch, excluded: bool) -> None: - monkeypatch.setattr(litellm, "standard_logging_payload_excluded_fields", ["messages", "response"] if excluded else []) + monkeypatch.setattr( + litellm, "standard_logging_payload_excluded_fields", ["messages", "response"] if excluded else [] + ) payload: Final = { "classifier_input": {"system": "private rubric"}, "originating_request_masked": {"input": "private source"}, @@ -991,7 +973,9 @@ def test_classifier_callback_redaction_preserves_exclusions(monkeypatch: pytest. } logger: Final = CustomLogger() logger.turn_off_message_logging = True - redacted: Final = logger.redact_standard_logging_payload_from_model_call_details({"standard_logging_object": payload}) + redacted: Final = logger.redact_standard_logging_payload_from_model_call_details( + {"standard_logging_object": payload} + ) stored: Final = redacted["standard_logging_object"] assert "classifier_input" not in stored assert "originating_request_masked" not in stored @@ -1017,16 +1001,18 @@ class _SelfRedactingLogger(CustomLogger): @pytest.mark.parametrize("logger", [CustomLogger(), _SelfRedactingLogger()], ids=["default", "redacts_itself"]) -def test_field_exclusion_alone_leaves_messages_and_responses_intact(monkeypatch: pytest.MonkeyPatch, logger: CustomLogger) -> None: +def test_field_exclusion_alone_leaves_messages_and_responses_intact( + monkeypatch: pytest.MonkeyPatch, logger: CustomLogger +) -> None: monkeypatch.setattr(litellm, "standard_logging_payload_excluded_fields", ["model"]) payload: Final = { "messages": [{"role": "user", "content": "private prompt"}], "response": {"choices": [{"message": {"content": "private answer"}}]}, "model": "classifier", } - stored: Final = logger.redact_standard_logging_payload_from_model_call_details({"standard_logging_object": payload})[ - "standard_logging_object" - ] + stored: Final = logger.redact_standard_logging_payload_from_model_call_details( + {"standard_logging_object": payload} + )["standard_logging_object"] assert stored == {"messages": payload["messages"], "response": payload["response"]} @@ -1038,9 +1024,9 @@ def test_a_callback_that_redacts_itself_keeps_its_messages_but_not_the_classifie } logger: Final = _SelfRedactingLogger() logger.turn_off_message_logging = True - stored: Final = logger.redact_standard_logging_payload_from_model_call_details({"standard_logging_object": payload})[ - "standard_logging_object" - ] + stored: Final = logger.redact_standard_logging_payload_from_model_call_details( + {"standard_logging_object": payload} + )["standard_logging_object"] assert "classifier_input" not in stored assert stored["messages"] == payload["messages"] assert stored["response"] == payload["response"] @@ -1054,19 +1040,54 @@ def test_perform_redaction_drops_the_served_output_texts_from_the_callback_kwarg assert SERVED_OUTPUT_TEXTS_KEY not in details +@pytest.mark.parametrize("callback_only", (False, True)) +@pytest.mark.parametrize("with_standard_payload", (False, True)) +def test_baseline_snapshots_are_redacted_without_mutating_request_state( + callback_only: bool, with_standard_payload: bool +) -> None: + snapshot: Final[Mapping[str, JsonValue]] = MappingProxyType({"system": "private system"}) + route: Final = BaselineRouteStamp("router", "baseline", "deployment", snapshot) + metadata: Final = {"_autorouter_baseline_route": route, "session_id": "session"} + params: Final = {"metadata": metadata, "litellm_metadata": metadata} + details: Final = { + "litellm_params": params, + **({"standard_logging_object": {"model": "model"}} if with_standard_payload else {}), + } + logger: Final = CustomLogger() + logger.turn_off_message_logging = True + if not callback_only: + perform_redaction(details, None) + redacted: Final = ( + logger.redact_standard_logging_payload_from_model_call_details(details) if callback_only else details + ) + expected: Final = BaselineRouteStamp(route.router_name, route.baseline_model, route.baseline_deployment_id) + assert redacted["litellm_params"] == { + key: {"_autorouter_baseline_route": expected, "session_id": "session"} + for key in ("metadata", "litellm_metadata") + } + assert route.request_parameters is snapshot + assert params["metadata"]["_autorouter_baseline_route"] is route + assert params["litellm_metadata"]["_autorouter_baseline_route"] is route + if callback_only: + assert details["litellm_params"] is params + + @pytest.fixture() def _vcr_outcome_gate(request, vcr): install_live_call_probe(request, vcr) yield record_vcr_outcome(request, vcr) + @pytest_asyncio.fixture(loop_scope="function") async def drain_logging_worker(isolate_litellm_state: None) -> AsyncIterator[None]: yield await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), timeout=LOGGING_WORKER_DRAIN_TIMEOUT_SECONDS) + LOGGING_WORKER_DRAIN_TIMEOUT_SECONDS: Final = LOGGING_WORKER_MAX_TIME_PER_COROUTINE + 5.0 + @pytest.fixture(scope="function") def isolate_litellm_state(): """ @@ -1104,6 +1125,7 @@ def isolate_litellm_state(): if attr in _DEFAULTS: setattr(litellm, attr, _DEFAULTS[attr]) + _LIST_ATTRS = ( "callbacks", "success_callback", @@ -1131,6 +1153,7 @@ _SCALAR_ATTRS = ( _DEFAULTS: dict = {} + @pytest.fixture(scope="module") def setup_and_teardown(): """ @@ -1153,6 +1176,7 @@ def setup_and_teardown(): litellm.in_memory_llm_clients_cache.flush_cache() yield + class TestCustomLogger(CustomLogger): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) @@ -1164,6 +1188,7 @@ class TestCustomLogger(CustomLogger): self.logged_standard_logging_payload = standard_logging_payload self.response_obj = response_obj + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_global_redaction_on(): @@ -1187,6 +1212,7 @@ async def test_global_redaction_on(): json.dumps(standard_logging_payload, indent=2), ) + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.parametrize( "dynamic_turn_off, expect_redacted", @@ -1213,6 +1239,7 @@ async def test_dynamic_turn_off_message_logging_overrides_global_on(dynamic_turn assert standard_logging_payload["response"]["choices"][0]["message"]["content"] == expected_response_content assert standard_logging_payload["messages"][0]["content"] == expected_message_content + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.parametrize( "dynamic_turn_off, expect_redacted", @@ -1239,6 +1266,7 @@ async def test_dynamic_turn_off_message_logging_overrides_global_off(dynamic_tur assert standard_logging_payload["response"]["choices"][0]["message"]["content"] == expected_response_content assert standard_logging_payload["messages"][0]["content"] == expected_message_content + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_redaction_with_custom_logger_streaming(): @@ -1284,6 +1312,7 @@ async def test_redaction_with_custom_logger_streaming(): finally: litellm.turn_off_message_logging = False + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_streaming_redaction_scoped_to_opted_out_logger(): @@ -1311,6 +1340,7 @@ async def test_streaming_redaction_scoped_to_opted_out_logger(): finally: litellm.callbacks = [] + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_redaction_responses_api(): @@ -1355,6 +1385,7 @@ async def test_redaction_responses_api(): json.dumps(standard_logging_payload, indent=2), ) + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_redaction_responses_api_stream(): @@ -1430,6 +1461,7 @@ async def test_redaction_responses_api_stream(): json.dumps(standard_logging_payload, indent=2), ) + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_redaction_responses_api_with_reasoning_summary(): @@ -1490,6 +1522,7 @@ async def test_redaction_responses_api_with_reasoning_summary(): assert model_call_details["messages"][0]["content"] == "redacted-by-litellm", "Input messages should be redacted" + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_redaction_with_coroutine_objects(): @@ -1535,6 +1568,7 @@ async def test_redaction_with_coroutine_objects(): result = perform_redaction({}, mock_iter) assert result == {"text": "redacted-by-litellm"} + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_redaction_with_streaming_response(): @@ -1570,6 +1604,7 @@ async def test_redaction_with_streaming_response(): json.dumps(standard_logging_payload, indent=2), ) + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_disable_redaction_header_responses_api(): @@ -1605,6 +1640,7 @@ async def test_disable_redaction_header_responses_api(): assert response["output"][0]["content"][0]["text"] == "This is a test response" assert standard_logging_payload["messages"][0]["content"] == "hi" + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_redaction_with_metadata_completion_api(): diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py b/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py index c3d4dba7376..ac55e8e9350 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py @@ -684,12 +684,12 @@ def _empty_block_msgs(): def test_handler_strips_when_no_presanitized_flag(): """Sync entry point (no async wrapper): handler must still sanitize.""" - from litellm.llms.anthropic.pass_through.messages import handler + from litellm.llms.anthropic.pass_through.messages import handler, utils with patch.object( - handler, + utils, "strip_empty_content_blocks_from_anthropic_messages", - wraps=handler.strip_empty_content_blocks_from_anthropic_messages, + wraps=utils.strip_empty_content_blocks_from_anthropic_messages, ) as spy: result = handler.anthropic_messages_handler( max_tokens=10, @@ -704,12 +704,12 @@ def test_handler_strips_when_no_presanitized_flag(): def test_handler_skips_strip_when_presanitized(): """Async wrapper already sanitized -> handler must NOT rescan.""" - from litellm.llms.anthropic.pass_through.messages import handler + from litellm.llms.anthropic.pass_through.messages import handler, utils with patch.object( - handler, + utils, "strip_empty_content_blocks_from_anthropic_messages", - wraps=handler.strip_empty_content_blocks_from_anthropic_messages, + wraps=utils.strip_empty_content_blocks_from_anthropic_messages, ) as spy: result = handler.anthropic_messages_handler( max_tokens=10, @@ -809,7 +809,7 @@ def test_presanitized_flag_not_leaked_to_provider_params(): @pytest.mark.asyncio async def test_async_wrapper_sets_presanitized_and_sanitizes_once(): """End-to-end: wrapper sanitizes (once) AND signals the handler to skip.""" - from litellm.llms.anthropic.pass_through.messages import handler + from litellm.llms.anthropic.pass_through.messages import handler, utils captured = {} @@ -825,9 +825,9 @@ async def test_async_wrapper_sets_presanitized_and_sanitizes_once(): patch.object(handler, "anthropic_messages_handler", side_effect=fake_handler), patch("asyncio.get_event_loop", return_value=fake_loop), patch.object( - handler, + utils, "strip_empty_content_blocks_from_anthropic_messages", - wraps=handler.strip_empty_content_blocks_from_anthropic_messages, + wraps=utils.strip_empty_content_blocks_from_anthropic_messages, ) as spy, ): await handler.anthropic_messages( diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_request_optional_param_utils.py b/tests/unit/llms/anthropic/pass_through/messages/test_request_optional_param_utils.py index dd744ca66a1..7a4f3aa5f90 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_request_optional_param_utils.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_request_optional_param_utils.py @@ -11,7 +11,7 @@ import pytest import litellm from litellm.llms.anthropic.pass_through.messages.utils import ( AnthropicMessagesRequestUtils, - _anthropic_messages_optional_param_keys, + anthropic_messages_optional_param_keys, ) @@ -30,16 +30,16 @@ def test_optional_param_filtering_unchanged(): def test_valid_keys_are_memoized(): - _anthropic_messages_optional_param_keys.cache_clear() - first = _anthropic_messages_optional_param_keys() + anthropic_messages_optional_param_keys.cache_clear() + first = anthropic_messages_optional_param_keys() for _ in range(50): AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param({"temperature": 0.1}) - info = _anthropic_messages_optional_param_keys.cache_info() + info = anthropic_messages_optional_param_keys.cache_info() # Resolved exactly once despite many calls. assert info.misses == 1 assert info.hits >= 50 # Stable identity (frozenset) returned each call. - assert _anthropic_messages_optional_param_keys() is first + assert anthropic_messages_optional_param_keys() is first assert isinstance(first, frozenset) assert "temperature" in first and "tools" in first diff --git a/tests/unit/llms/base_llm/chat/test_attribution_headers.py b/tests/unit/llms/base_llm/chat/test_attribution_headers.py new file mode 100644 index 00000000000..f31b4e846bf --- /dev/null +++ b/tests/unit/llms/base_llm/chat/test_attribution_headers.py @@ -0,0 +1,210 @@ +""" +Provider attribution headers (`BaseConfig.get_attribution_headers`) must reach +the outbound request on every OpenAI-compatible chat path, and a caller header +with the same name must win. + +Requests go through a real `litellm.completion` into an in-process httpx +transport that records what would have been sent, because the default path +(OpenAI SDK) never calls `validate_environment`. +""" + +import json +from collections.abc import AsyncIterable, Iterable +from typing import Final, cast + +import httpx +import openai +import pytest + +import litellm +from litellm.llms.base_llm.chat.transformation import with_attribution_headers +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler + +_API_BASE: Final = "https://provider.invalid/v1" +_COMPLETION_BODY: Final = json.dumps( + { + "id": "chatcmpl-1", + "object": "chat.completion", + "created": 1, + "model": "m", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + } +).encode() +_STREAM_BODY: Final = ( + "data: " + + json.dumps( + { + "id": "chatcmpl-1", + "object": "chat.completion.chunk", + "created": 1, + "model": "m", + "choices": [{"index": 0, "delta": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}], + } + ) + + "\n\ndata: [DONE]\n\n" +).encode() + + +class _HeaderCapturingTransport(httpx.BaseTransport, httpx.AsyncBaseTransport): + """Records each outbound request's headers and answers like a chat completions server.""" + + def __init__(self) -> None: + self.sent: tuple[httpx.Headers, ...] = () + + def handle_request(self, request: httpx.Request) -> httpx.Response: + return self._respond(request, request.read()) + + async def handle_async_request(self, request: httpx.Request) -> httpx.Response: + return self._respond(request, await request.aread()) + + def _respond(self, request: httpx.Request, body: bytes) -> httpx.Response: + self.sent = (*self.sent, request.headers) + if json.loads(body).get("stream"): + return httpx.Response(200, content=_STREAM_BODY, headers={"content-type": "text/event-stream"}) + return httpx.Response(200, content=_COMPLETION_BODY, headers={"content-type": "application/json"}) + + def last(self, header: str) -> list[str]: + return self.sent[-1].get_list(header) + + +def _client(transport: _HeaderCapturingTransport, path: str, is_async: bool) -> object: + if path == "sdk": + if is_async: + return openai.AsyncOpenAI( + api_key="k", base_url=_API_BASE, http_client=httpx.AsyncClient(transport=transport) + ) + return openai.OpenAI(api_key="k", base_url=_API_BASE, http_client=httpx.Client(transport=transport)) + if is_async: + return AsyncHTTPHandler(transport=transport) + return HTTPHandler(client=httpx.Client(transport=transport)) + + +@pytest.fixture +def transport() -> _HeaderCapturingTransport: + return _HeaderCapturingTransport() + + +@pytest.fixture(params=["sdk", "http_handler"]) +def handler_path(request: pytest.FixtureRequest, monkeypatch: pytest.MonkeyPatch) -> str: + if request.param == "http_handler": + monkeypatch.setenv("EXPERIMENTAL_OPENAI_BASE_LLM_HTTP_HANDLER", "true") + else: + monkeypatch.delenv("EXPERIMENTAL_OPENAI_BASE_LLM_HTTP_HANDLER", raising=False) + return request.param + + +_NOVITA_MODEL: Final = "novita/meta-llama/llama-3.3-70b-instruct" + +_ATTRIBUTED: Final = [ + pytest.param(_NOVITA_MODEL, "x-novita-source", id="novita"), + pytest.param("perplexity/sonar", "x-pplx-integration", id="perplexity"), +] + + +def _drain(response: object) -> None: + for _ in cast(Iterable[object], response): + pass + + +async def _adrain(response: object) -> None: + async for _ in cast(AsyncIterable[object], response): + pass + + +@pytest.mark.parametrize(("model", "header"), _ATTRIBUTED) +@pytest.mark.parametrize("stream", [False, True]) +def test_attribution_header_sent_sync( + transport: _HeaderCapturingTransport, handler_path: str, model: str, header: str, stream: bool +) -> None: + response: Final = litellm.completion( + model=model, + messages=[{"role": "user", "content": "hi"}], + api_base=_API_BASE, + api_key="k", + stream=stream, + client=_client(transport, handler_path, is_async=False), + ) + if stream: + _drain(response) + + assert transport.last(header) == ["litellm"] + + +@pytest.mark.asyncio +@pytest.mark.parametrize(("model", "header"), _ATTRIBUTED) +@pytest.mark.parametrize("stream", [False, True]) +async def test_attribution_header_sent_async( + transport: _HeaderCapturingTransport, handler_path: str, model: str, header: str, stream: bool +) -> None: + response: Final = await litellm.acompletion( + model=model, + messages=[{"role": "user", "content": "hi"}], + api_base=_API_BASE, + api_key="k", + stream=stream, + client=_client(transport, handler_path, is_async=True), + ) + if stream: + await _adrain(response) + + assert transport.last(header) == ["litellm"] + + +@pytest.mark.parametrize(("model", "header"), _ATTRIBUTED) +@pytest.mark.parametrize("header_kwarg", ["headers", "extra_headers"]) +def test_caller_header_overrides_attribution_any_casing( + transport: _HeaderCapturingTransport, handler_path: str, model: str, header: str, header_kwarg: str +) -> None: + caller_headers: Final = {header.upper(): "my-app"} + + litellm.completion( + model=model, + messages=[{"role": "user", "content": "hi"}], + api_base=_API_BASE, + api_key="k", + client=_client(transport, handler_path, is_async=False), + **{header_kwarg: caller_headers}, + ) + + assert transport.last(header) == ["my-app"] + assert caller_headers == {header.upper(): "my-app"} + + +def test_provider_without_attribution_sends_none(transport: _HeaderCapturingTransport, handler_path: str) -> None: + litellm.completion( + model="deepinfra/meta-llama/Meta-Llama-3-8B-Instruct", + messages=[{"role": "user", "content": "hi"}], + api_base=_API_BASE, + api_key="k", + client=_client(transport, handler_path, is_async=False), + ) + + assert transport.last("x-novita-source") == [] + assert transport.last("x-pplx-integration") == [] + + +def test_global_litellm_headers_still_apply_and_are_not_mutated( + transport: _HeaderCapturingTransport, handler_path: str, monkeypatch: pytest.MonkeyPatch +) -> None: + global_headers: Final = {"X-Global": "1"} + monkeypatch.setattr(litellm, "headers", global_headers) + + litellm.completion( + model=_NOVITA_MODEL, + messages=[{"role": "user", "content": "hi"}], + api_base=_API_BASE, + api_key="k", + client=_client(transport, handler_path, is_async=False), + ) + + assert transport.last("x-global") == ["1"] + assert transport.last("x-novita-source") == ["litellm"] + assert global_headers == {"X-Global": "1"} + + +def test_with_attribution_headers_returns_headers_unchanged_when_nothing_to_add() -> None: + headers: Final = {"A": "1"} + + assert with_attribution_headers({}, headers) is headers + assert with_attribution_headers({}, None) is None diff --git a/tests/unit/proxy/guardrails/test_guardrail_endpoints.py b/tests/unit/proxy/guardrails/test_guardrail_endpoints.py index 7397b0ec6e2..fa33c2d462b 100644 --- a/tests/unit/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/unit/proxy/guardrails/test_guardrail_endpoints.py @@ -1,7 +1,9 @@ import json import time +from collections.abc import Callable from datetime import datetime -from typing import Dict, List, Optional +from types import MappingProxyType +from typing import Dict, Final, List, Optional from unittest.mock import AsyncMock import pytest @@ -10,6 +12,7 @@ import yaml from fastapi import HTTPException +from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.guardrails.guardrail_endpoints import ( CreateGuardrailRequest, @@ -46,10 +49,12 @@ from litellm.proxy.guardrails.guardrail_registry import ( from litellm.types.guardrails import ( ApplyGuardrailRequest, BaseLitellmParams, + GuardrailEventHooks, Guardrail, GuardrailInfoResponse, LitellmParams, ) +from litellm.types.utils import GenericGuardrailAPIInputs # Mock data for testing MOCK_DB_GUARDRAIL = { @@ -64,6 +69,46 @@ MOCK_DB_GUARDRAIL = { "updated_at": datetime.now(), } +_INVALID_SCOPE_LITELLM_PARAMS: Final = MappingProxyType( + { + "guardrail": "litellm_content_filter", + "mode": "pre_call", + "logging_only_scope": "Input", + "blocked_words": [{"keyword": "synthetic blocked phrase", "action": "BLOCK"}], + } +) +_INVALID_SCOPE_DB_GUARDRAIL: Final = MappingProxyType( + { + "guardrail_id": "invalid-scope-db-guardrail", + "guardrail_name": "Invalid scope DB guardrail", + "litellm_params": _INVALID_SCOPE_LITELLM_PARAMS, + "guardrail_info": MappingProxyType({}), + } +) +_INVALID_SCOPE_IN_MEMORY_GUARDRAIL: Final = MappingProxyType( + { + "guardrail_id": "invalid-scope-in-memory-guardrail", + "guardrail_name": "Invalid scope in-memory guardrail", + "litellm_params": _INVALID_SCOPE_LITELLM_PARAMS, + "guardrail_info": MappingProxyType({}), + } +) +_VALID_SCOPE_DB_GUARDRAIL: Final = MappingProxyType( + { + "guardrail_id": "valid-scope-db-guardrail", + "guardrail_name": "Valid scope DB guardrail", + "litellm_params": MappingProxyType( + { + "guardrail": "litellm_content_filter", + "mode": "pre_call", + "logging_only_scope": "output", + "blocked_words": [{"keyword": "synthetic blocked phrase", "action": "BLOCK"}], + } + ), + "guardrail_info": MappingProxyType({}), + } +) + MOCK_CONFIG_GUARDRAIL = { "guardrail_id": "test-config-guardrail", "guardrail_name": "Test Config Guardrail", @@ -89,6 +134,44 @@ MOCK_PATCH_REQUEST = PatchGuardrailRequest( ) +class _PatchScopeSupportedGuardrail(CustomGuardrail): + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: str, + logging_obj: object | None = None, + ) -> GenericGuardrailAPIInputs: + return inputs + + +class _PatchScopeUnsupportedGuardrail(_PatchScopeSupportedGuardrail): + async def async_logging_hook( + self, + kwargs: dict[str, object], + result: object, + call_type: str, + ) -> tuple[dict[str, object], object]: + return kwargs, result + + +def _patch_scope_initializer( + callback_type: type[CustomGuardrail], +) -> Callable[[LitellmParams, Guardrail], CustomGuardrail]: + def _initializer(litellm_params: LitellmParams, guardrail: Guardrail) -> CustomGuardrail: + import litellm + + callback = callback_type( + guardrail_name=guardrail["guardrail_name"], + event_hook=litellm_params.mode, + default_on=litellm_params.default_on, + ) + litellm.logging_callback_manager.add_litellm_callback(callback) + return callback + + return _initializer + + @pytest.fixture def mock_prisma_client(mocker): """Mock Prisma client for testing""" @@ -128,6 +211,37 @@ def mock_guardrail_registry(mocker): return mock_registry +def _setup_patch_scope_guardrail( + mocker, + monkeypatch, + mock_guardrail_registry, + callback_type: type[CustomGuardrail], + guardrail_type: str, + litellm_params: dict[str, object], +) -> tuple[InMemoryGuardrailHandler, Guardrail]: + from litellm.proxy.guardrails import guardrail_registry as registry_module + + guardrail: Guardrail = { + "guardrail_id": "patch-scope-test", + "guardrail_name": "Patch scope test", + "litellm_params": {"guardrail": guardrail_type, **litellm_params}, + "guardrail_info": {}, + } + mock_guardrail_registry.get_guardrail_by_id_from_db.return_value = guardrail + mock_guardrail_registry.update_guardrail_in_db.return_value = guardrail + monkeypatch.setitem( + registry_module.guardrail_initializer_registry, + guardrail_type, + _patch_scope_initializer(callback_type), + ) + handler = InMemoryGuardrailHandler() + handler.initialize_guardrail(guardrail=guardrail, source="db") + mocker.patch("litellm.proxy.proxy_server.prisma_client", mocker.Mock()) + mocker.patch("litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY", mock_guardrail_registry) + mocker.patch("litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER", handler) + return handler, guardrail + + @pytest.mark.asyncio async def test_list_guardrails_v2_with_db_and_config(mocker, mock_prisma_client, mock_in_memory_handler): """Test listing guardrails from both DB and config""" @@ -157,6 +271,54 @@ async def test_list_guardrails_v2_with_db_and_config(mocker, mock_prisma_client, assert isinstance(config_guardrail.litellm_params, BaseLitellmParams) +@pytest.mark.asyncio +async def test_list_guardrails_v2_normalizes_invalid_scope_and_keeps_other_db_rows( + mocker, mock_prisma_client, mock_in_memory_handler +): + mock_prisma_client.db.litellm_guardrailstable.find_many.return_value = ( + _INVALID_SCOPE_DB_GUARDRAIL, + _VALID_SCOPE_DB_GUARDRAIL, + ) + mock_in_memory_handler.list_in_memory_guardrails.return_value = () + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mocker.patch( + "litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER", + mock_in_memory_handler, + ) + + response: Final = await list_guardrails_v2(user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)) + + invalid_scope_row: Final = next( + guardrail for guardrail in response.guardrails if guardrail.guardrail_id == "invalid-scope-db-guardrail" + ) + assert invalid_scope_row.litellm_params is not None + assert invalid_scope_row.litellm_params.logging_only_scope is None + assert invalid_scope_row.litellm_params.mode == "pre_call" + assert invalid_scope_row.litellm_params.guardrail == "litellm_content_filter" + assert any(guardrail.guardrail_id == "valid-scope-db-guardrail" for guardrail in response.guardrails) + + +@pytest.mark.asyncio +async def test_list_guardrails_v2_normalizes_invalid_scope_in_memory( + mocker, mock_prisma_client, mock_in_memory_handler +): + mock_prisma_client.db.litellm_guardrailstable.find_many.return_value = () + mock_in_memory_handler.list_in_memory_guardrails.return_value = (_INVALID_SCOPE_IN_MEMORY_GUARDRAIL,) + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mocker.patch( + "litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER", + mock_in_memory_handler, + ) + + response: Final = await list_guardrails_v2(user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)) + + assert len(response.guardrails) == 1 + assert response.guardrails[0].litellm_params is not None + assert response.guardrails[0].litellm_params.logging_only_scope is None + assert response.guardrails[0].litellm_params.mode == "pre_call" + assert response.guardrails[0].litellm_params.guardrail == "litellm_content_filter" + + @pytest.mark.asyncio async def test_list_guardrails_v2_skips_stale_db_backed_in_memory_entries(mocker): """ @@ -421,6 +583,28 @@ async def test_get_guardrail_info_from_db(mocker, mock_prisma_client): assert response.guardrail_info == {"description": "Test guardrail from DB"} +@pytest.mark.asyncio +async def test_get_guardrail_info_normalizes_invalid_scope_from_db( + mocker, mock_guardrail_registry, mock_in_memory_handler +): + mock_guardrail_registry.get_guardrail_by_id_from_db.return_value = _INVALID_SCOPE_DB_GUARDRAIL + mock_in_memory_handler.get_guardrail_by_id.return_value = None + mocker.patch("litellm.proxy.proxy_server.prisma_client", mocker.Mock()) + mocker.patch("litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY", mock_guardrail_registry) + mocker.patch( + "litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER", + mock_in_memory_handler, + ) + + response: Final = await get_guardrail_info("invalid-scope-db-guardrail") + + assert response.guardrail_id == "invalid-scope-db-guardrail" + assert response.litellm_params is not None + assert response.litellm_params.logging_only_scope is None + assert response.litellm_params.mode == "pre_call" + assert response.litellm_params.guardrail == "litellm_content_filter" + + @pytest.mark.asyncio async def test_get_guardrail_info_from_config(mocker, mock_prisma_client, mock_in_memory_handler): """Test getting guardrail info from config when not found in DB""" @@ -1127,7 +1311,10 @@ async def test_update_guardrail_endpoint( prisma_client=mocker.ANY, ) - mock_in_memory_handler.sync_guardrail_from_db.assert_called_once_with(guardrail=mocker.ANY) + mock_in_memory_handler.sync_guardrail_from_db.assert_called_once_with( + guardrail=mocker.ANY, + reject_invalid_logging_only_scope=True, + ) if scenario == "success_sync_fails_unexpected_error": assert mock_logger is not None @@ -1256,7 +1443,10 @@ async def test_patch_guardrail_endpoint( mock_guardrail_registry.update_guardrail_in_db.assert_called_once() - mock_in_memory_handler.sync_guardrail_from_db.assert_called_once_with(guardrail=mocker.ANY) + mock_in_memory_handler.sync_guardrail_from_db.assert_called_once_with( + guardrail=mocker.ANY, + reject_invalid_logging_only_scope=False, + ) if scenario == "success_sync_fails_unexpected_error": assert mock_logger is not None @@ -1280,6 +1470,192 @@ async def test_patch_guardrail_rejects_mcp_only_on_violation_with_422(mocker, mo mock_guardrail_registry.update_guardrail_in_db.assert_not_called() +@pytest.mark.asyncio +async def test_patch_guardrail_rejects_invalid_logging_only_scope_with_422(mocker, mock_guardrail_registry): + mocker.patch("litellm.proxy.proxy_server.prisma_client", mocker.Mock()) + mocker.patch( + "litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY", + mock_guardrail_registry, + ) + mock_in_memory_handler = mocker.Mock(spec=InMemoryGuardrailHandler) + mock_in_memory_handler.sync_guardrail_from_db.side_effect = ValueError( + "Guardrail test-db-guardrail: logging_only_scope is set, but mode does not include logging_only" + ) + mocker.patch( + "litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER", + mock_in_memory_handler, + ) + request = PatchGuardrailRequest(litellm_params=BaseLitellmParams(mode="pre_call", logging_only_scope="input")) + + with pytest.raises(HTTPException) as exc_info: + await patch_guardrail("test-guardrail-id", request, user_api_key_dict=MOCK_ADMIN_USER) + + assert exc_info.value.status_code == 422 + assert "update rejected" in str(exc_info.value.detail) + mock_in_memory_handler.sync_guardrail_from_db.assert_called_once_with( + guardrail=mocker.ANY, + reject_invalid_logging_only_scope=True, + ) + + +@pytest.mark.asyncio +async def test_patch_guardrail_clears_scope_when_logging_only_mode_is_removed( + mocker, monkeypatch, mock_guardrail_registry +): + handler, stored_guardrail = _setup_patch_scope_guardrail( + mocker, + monkeypatch, + mock_guardrail_registry, + _PatchScopeSupportedGuardrail, + "patch_scope_supported_test", + { + "mode": ["pre_call", "logging_only"], + "logging_only_scope": "output", + "default_on": True, + }, + ) + request = PatchGuardrailRequest(litellm_params=BaseLitellmParams(mode=["pre_call"])) + + try: + result = await patch_guardrail( + stored_guardrail["guardrail_id"], + request, + user_api_key_dict=MOCK_ADMIN_USER, + ) + + assert result["guardrail_id"] == stored_guardrail["guardrail_id"] + persisted_guardrail = mock_guardrail_registry.update_guardrail_in_db.call_args.kwargs["guardrail"] + assert persisted_guardrail["litellm_params"].logging_only_scope is None + callback = handler.guardrail_id_to_custom_guardrail[stored_guardrail["guardrail_id"]] + assert callback.logging_only_scope is None + assert callback.should_run_guardrail(data={}, event_type=GuardrailEventHooks.pre_call) is True + finally: + handler.delete_in_memory_guardrail(stored_guardrail["guardrail_id"]) + + +@pytest.mark.asyncio +async def test_patch_guardrail_tolerates_stored_unsupported_scope_on_unrelated_update( + mocker, monkeypatch, mock_guardrail_registry, caplog +): + handler, stored_guardrail = _setup_patch_scope_guardrail( + mocker, + monkeypatch, + mock_guardrail_registry, + _PatchScopeUnsupportedGuardrail, + "patch_scope_unsupported_test", + {"mode": "logging_only", "logging_only_scope": "output", "default_on": True}, + ) + caplog.clear() + request = PatchGuardrailRequest(litellm_params=BaseLitellmParams(default_on=False)) + + try: + result = await patch_guardrail( + stored_guardrail["guardrail_id"], + request, + user_api_key_dict=MOCK_ADMIN_USER, + ) + + assert result["guardrail_id"] == stored_guardrail["guardrail_id"] + callback = handler.guardrail_id_to_custom_guardrail[stored_guardrail["guardrail_id"]] + assert callback.logging_only_scope is None + assert any("Ignoring logging_only_scope" in record.getMessage() for record in caplog.records) + finally: + handler.delete_in_memory_guardrail(stored_guardrail["guardrail_id"]) + + +@pytest.mark.asyncio +async def test_patch_guardrail_tolerates_invalid_stored_scope_on_unrelated_update( + mocker, monkeypatch, mock_guardrail_registry +): + handler, stored_guardrail = _setup_patch_scope_guardrail( + mocker, + monkeypatch, + mock_guardrail_registry, + _PatchScopeUnsupportedGuardrail, + "patch_scope_invalid_literal_test", + {"mode": "logging_only", "logging_only_scope": "sideways", "default_on": True}, + ) + request = PatchGuardrailRequest(litellm_params=BaseLitellmParams(default_on=False)) + + try: + result = await patch_guardrail( + stored_guardrail["guardrail_id"], + request, + user_api_key_dict=MOCK_ADMIN_USER, + ) + + assert result["guardrail_id"] == stored_guardrail["guardrail_id"] + persisted_guardrail = mock_guardrail_registry.update_guardrail_in_db.call_args.kwargs["guardrail"] + assert persisted_guardrail["litellm_params"].logging_only_scope is None + assert persisted_guardrail["litellm_params"].default_on is False + finally: + handler.delete_in_memory_guardrail(stored_guardrail["guardrail_id"]) + + +@pytest.mark.asyncio +async def test_patch_guardrail_rejects_explicit_unsupported_scope_and_rolls_back( + mocker, monkeypatch, mock_guardrail_registry +): + handler, stored_guardrail = _setup_patch_scope_guardrail( + mocker, + monkeypatch, + mock_guardrail_registry, + _PatchScopeUnsupportedGuardrail, + "patch_scope_unsupported_test", + {"mode": "logging_only", "logging_only_scope": "output", "default_on": True}, + ) + request = PatchGuardrailRequest(litellm_params=BaseLitellmParams(logging_only_scope="output")) + + try: + with pytest.raises(HTTPException) as exc_info: + await patch_guardrail( + stored_guardrail["guardrail_id"], + request, + user_api_key_dict=MOCK_ADMIN_USER, + ) + + assert exc_info.value.status_code == 422 + assert mock_guardrail_registry.update_guardrail_in_db.call_count == 2 + restored_guardrail = mock_guardrail_registry.update_guardrail_in_db.call_args_list[-1].kwargs["guardrail"] + assert restored_guardrail["litellm_params"] == stored_guardrail["litellm_params"] + callback = handler.guardrail_id_to_custom_guardrail[stored_guardrail["guardrail_id"]] + assert callback.logging_only_scope is None + finally: + handler.delete_in_memory_guardrail(stored_guardrail["guardrail_id"]) + + +@pytest.mark.asyncio +async def test_patch_guardrail_rejected_update_restores_invalid_stored_scope_verbatim( + mocker, monkeypatch, mock_guardrail_registry +): + handler, stored_guardrail = _setup_patch_scope_guardrail( + mocker, + monkeypatch, + mock_guardrail_registry, + _PatchScopeUnsupportedGuardrail, + "patch_scope_invalid_rollback_test", + {"mode": "logging_only", "logging_only_scope": "sideways", "default_on": True}, + ) + request = PatchGuardrailRequest(litellm_params=BaseLitellmParams(logging_only_scope="output")) + + try: + with pytest.raises(HTTPException) as exc_info: + await patch_guardrail( + stored_guardrail["guardrail_id"], + request, + user_api_key_dict=MOCK_ADMIN_USER, + ) + + assert exc_info.value.status_code == 422 + assert mock_guardrail_registry.update_guardrail_in_db.call_count == 2 + restored_guardrail = mock_guardrail_registry.update_guardrail_in_db.call_args_list[-1].kwargs["guardrail"] + assert restored_guardrail["litellm_params"] == stored_guardrail["litellm_params"] + callback = handler.guardrail_id_to_custom_guardrail[stored_guardrail["guardrail_id"]] + assert callback.logging_only_scope is None + finally: + handler.delete_in_memory_guardrail(stored_guardrail["guardrail_id"]) + + @pytest.mark.parametrize( "scenario,expected_result,expected_exception", [ @@ -2465,6 +2841,13 @@ async def test_ui_settings_map_matches_runtime_supported_event_hooks(): from litellm.proxy.guardrails.guardrail_registry import guardrail_class_registry result = await get_guardrail_ui_settings() + expected_without_directional_scope = { + provider + for provider, guardrail_class in guardrail_class_registry.items() + if not guardrail_class.supports_logging_only_scope() + } + assert set(result.providers_without_directional_logging_only_scope) == expected_without_directional_scope + assert "xecguard" in result.providers_without_directional_logging_only_scope for provider, guardrail_class in guardrail_class_registry.items(): declared = guardrail_class.get_supported_event_hooks() diff --git a/tests/unit/proxy/guardrails/test_guardrail_registry.py b/tests/unit/proxy/guardrails/test_guardrail_registry.py index 2f171677d12..c7b13f0bd00 100644 --- a/tests/unit/proxy/guardrails/test_guardrail_registry.py +++ b/tests/unit/proxy/guardrails/test_guardrail_registry.py @@ -1,15 +1,20 @@ +import json from collections.abc import Iterable, Iterator -from unittest.mock import AsyncMock, MagicMock, patch +from typing import ClassVar, Final +from unittest.mock import AsyncMock, MagicMock import pytest +from pydantic import ValidationError from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.proxy.guardrails.guardrail_registry import ( - get_guardrail_initializer_from_hooks, GuardrailRegistry, InMemoryGuardrailHandler, + get_guardrail_initializer_from_hooks, + parse_tolerant_litellm_params, ) -from litellm.types.guardrails import GuardrailEventHooks, Guardrail, LitellmParams +from litellm.types.guardrails import Guardrail, GuardrailEventHooks, LitellmParams, LoggingOnlyScope, Mode +from litellm.types.utils import GenericGuardrailAPIInputs def test_get_guardrail_initializer_from_hooks(): @@ -400,6 +405,85 @@ def test_sync_guardrail_from_db_marks_source_db_when_unchanged(): assert handler.get_source("collide") == "db" +def test_sync_guardrail_from_db_reject_flag_keeps_callback_order_on_noop_update(): + """ + The PUT endpoint syncs the whole object with reject_invalid_logging_only_scope=True. + That strictness must not force a teardown + re-append of an unchanged guardrail: + initialize_guardrail appends the rebuilt callback at the END of litellm.callbacks, + so a description-only PUT would reorder guardrails and change which one wins + between a BLOCK and a MASK guardrail over the same content. + """ + import litellm + + registry_module = _register_mode_following_initializer("mode_following_test") + lists = _all_callback_lists() + snapshots = [list(cb_list) for cb_list in lists] + sentinel: Final = CustomGuardrail( + guardrail_name="order-sentinel", + supported_event_hooks=[GuardrailEventHooks.pre_call], + event_hook=GuardrailEventHooks.pre_call, + ) + try: + handler = InMemoryGuardrailHandler() + handler.initialize_guardrail(guardrail=_mode_following_db_row("123", "pre_call"), source="db") + original = handler.guardrail_id_to_custom_guardrail["123"] + assert original is not None + litellm.callbacks.append(sentinel) + index_before = litellm.callbacks.index(original) + + handler.sync_guardrail_from_db( + guardrail=_mode_following_db_row("123", "pre_call", "description-only edit"), + reject_invalid_logging_only_scope=True, + ) + + assert handler.guardrail_id_to_custom_guardrail["123"] is original + assert litellm.callbacks.index(original) == index_before + assert _live_instances_named("mode-following") == 1 + finally: + registry_module.guardrail_initializer_registry.pop("mode_following_test", None) + for cb_list, snapshot in zip(lists, snapshots): + cb_list[:] = snapshot + + +def test_sync_guardrail_from_db_reject_flag_still_rejects_invalid_unchanged_scope(): + """ + A PUT sends the whole object, so an unchanged row that already carries an + invalid logging_only_scope (tolerated at load) must still be rejected on the + strict sync path, without rebuilding the live callback. + """ + registry_module = _register_mode_following_initializer("mode_following_test") + lists = _all_callback_lists() + snapshots = [list(cb_list) for cb_list in lists] + row: Final = Guardrail( + guardrail_id="123", + guardrail_name="mode-following", + litellm_params={ + "guardrail": "mode_following_test", + "mode": "pre_call", + "default_on": True, + "logging_only_scope": "input", + }, + guardrail_info={}, + ) + try: + handler = InMemoryGuardrailHandler() + handler.initialize_guardrail(guardrail=row, source="db") + original = handler.guardrail_id_to_custom_guardrail["123"] + assert original is not None + assert original.logging_only_scope is None # tolerated at load + + with pytest.raises(ValueError, match="logging_only_scope is set"): + handler.sync_guardrail_from_db(guardrail=row, reject_invalid_logging_only_scope=True) + + # Rejected without touching the live instance. + assert handler.guardrail_id_to_custom_guardrail["123"] is original + assert _live_instances_named("mode-following") == 1 + finally: + registry_module.guardrail_initializer_registry.pop("mode_following_test", None) + for cb_list, snapshot in zip(lists, snapshots): + cb_list[:] = snapshot + + @pytest.fixture def rotation_handler() -> Iterator[InMemoryGuardrailHandler]: registry_module = _register_mode_following_initializer("rotation_test") @@ -528,6 +612,49 @@ def test_unchanged_db_params_do_not_register_as_changed(): assert handler._has_guardrail_params_changed(gid, new) is False +def test_db_poll_does_not_reinitialize_config_guardrail_without_default_on(): + handler = InMemoryGuardrailHandler() + guardrail_id: Final = "config-default-on-guardrail" + guardrail_name: Final = "config-default-on-guardrail" + params: Final = { + "guardrail": "litellm_content_filter", + "mode": "pre_call", + "logging_only_scope": "Input", + "blocked_words": [{"keyword": "synthetic blocked phrase", "action": "BLOCK"}], + } + callback_lists: Final = _all_callback_lists() + callback_snapshots: Final = [list(callback_list) for callback_list in callback_lists] + + try: + existing: Final = handler.initialize_guardrail( + guardrail=Guardrail( + guardrail_id=guardrail_id, + guardrail_name=guardrail_name, + litellm_params=params, + ), + source="config", + ) + assert existing is not None + assert existing["litellm_params"].default_on is False + assert existing["litellm_params"].logging_only_scope is None + + synced: Final = handler.sync_guardrail_from_db( + Guardrail( + guardrail_id=guardrail_id, + guardrail_name=guardrail_name, + litellm_params=params, + ) + ) + + assert synced is existing + assert handler.IN_MEMORY_GUARDRAILS[guardrail_id] is existing + assert handler._sources[guardrail_id] == "db" + finally: + handler.delete_in_memory_guardrail(guardrail_id) + for callback_list, snapshot in zip(callback_lists, callback_snapshots): + callback_list[:] = snapshot + + def test_changed_db_params_register_as_changed(): """Normalizing both sides must still surface a genuine config change.""" handler = InMemoryGuardrailHandler() @@ -565,6 +692,24 @@ def test_unnormalizable_db_params_register_as_changed_without_raising(): assert handler._has_guardrail_params_changed(gid, new) is True +def test_invalid_scope_literal_db_params_compare_equal_after_normalization(): + handler = InMemoryGuardrailHandler() + raw = _db_litellm_params() + gid = "77777777-7777-7777-7777-777777777777" + handler.IN_MEMORY_GUARDRAILS[gid] = Guardrail( + guardrail_id=gid, + guardrail_name="cf", + litellm_params=LitellmParams(**{**raw, "logging_only_scope": None}), + ) + new = Guardrail( + guardrail_id=gid, + guardrail_name="cf", + litellm_params={**raw, "logging_only_scope": "Input"}, + ) + + assert handler._has_guardrail_params_changed(gid, new) is False + + def _all_callback_lists(): import litellm @@ -1025,6 +1170,250 @@ class TestScanOnlyToolResultsInitRefusal: ) +class _LoggingOnlyScopeSupportedGuardrail(CustomGuardrail): + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: str, + logging_obj: object | None = None, + ) -> GenericGuardrailAPIInputs: + return inputs + + +class _LoggingOnlyScopeUnsupportedGuardrail(_LoggingOnlyScopeSupportedGuardrail): + async def async_logging_hook( + self, + kwargs: dict[str, object], + result: object, + call_type: str, + ) -> tuple[dict[str, object], object]: + return kwargs, result + + +class _LoggingOnlyScopeNativeGuardrail(_LoggingOnlyScopeSupportedGuardrail): + use_native_lifecycle_hooks: ClassVar[bool] = True + + +def _invalid_scope_content_filter_guardrail() -> Guardrail: + return Guardrail( + guardrail_id="invalid-scope-content-filter-test", + guardrail_name="invalid-scope-content-filter", + litellm_params={ + "guardrail": "litellm_content_filter", + "mode": "pre_call", + "logging_only_scope": "Input", + "blocked_words": [{"keyword": "pineapple", "action": "BLOCK"}], + }, + ) + + +class TestLoggingOnlyScopeValidation: + @pytest.mark.parametrize( + ("scope", "expected_scope"), + (("input", "input"), ("Input", None)), + ) + def test_tolerant_parser_preserves_default_on_constructor_coercion( + self, scope: str, expected_scope: str | None + ) -> None: + params: Final = { + "guardrail": "litellm_content_filter", + "mode": "pre_call", + "logging_only_scope": scope, + "blocked_words": [{"keyword": "synthetic blocked phrase", "action": "BLOCK"}], + } + + parsed: Final = parse_tolerant_litellm_params(params, "test-content-filter") + expected: Final = LitellmParams(**{**params, "logging_only_scope": expected_scope}).model_dump() + + assert parsed.default_on is False + assert parsed.model_dump() == expected + + def _initialize( + self, + mode: str | list[str] | Mode, + scope: LoggingOnlyScope | None, + callback_type: type[CustomGuardrail] = _LoggingOnlyScopeSupportedGuardrail, + reject_invalid_logging_only_scope: bool = False, + assert_registered: bool = False, + ) -> CustomGuardrail: + import litellm + from litellm.proxy.guardrails import guardrail_registry as registry_module + + guardrail_type: Final = "logging_only_scope_test" + created_callbacks: Final[list[CustomGuardrail]] = [] + + def _initializer(litellm_params: LitellmParams, guardrail: Guardrail) -> CustomGuardrail: + supported_event_hooks: Final = ( + [GuardrailEventHooks.logging_only] if callback_type.use_native_lifecycle_hooks else None + ) + callback: Final = callback_type( + guardrail_name=guardrail["guardrail_name"], + event_hook=litellm_params.mode, + default_on=True, + supported_event_hooks=supported_event_hooks, + ) + litellm.logging_callback_manager.add_litellm_callback(callback) + created_callbacks.append(callback) + return callback + + registry_module.guardrail_initializer_registry[guardrail_type] = _initializer + lists: Final = _all_callback_lists() + snapshots: Final = [list(callback_list) for callback_list in lists] + try: + handler: Final = InMemoryGuardrailHandler() + result: Final = handler.initialize_guardrail( + guardrail={ + "guardrail_name": "logging-only-scope-guardrail", + "litellm_params": { + "guardrail": guardrail_type, + "mode": mode, + "logging_only_scope": scope, + }, + }, + reject_invalid_logging_only_scope=reject_invalid_logging_only_scope, + ) + assert result is not None + callback: Final = handler.guardrail_id_to_custom_guardrail[result["guardrail_id"]] + assert callback is not None + if assert_registered: + assert callback in lists[0] + return callback + except ValueError: + callback: Final = created_callbacks[0] + assert all(callback not in callback_list for callback_list in lists) + raise + finally: + for callback_list, snapshot in zip(lists, snapshots): + callback_list[:] = snapshot + registry_module.guardrail_initializer_registry.pop(guardrail_type, None) + + def test_scope_without_logging_only_mode_is_ignored_at_load(self) -> None: + callback: Final = self._initialize(mode="pre_call", scope="input", assert_registered=True) + + assert callback.logging_only_scope is None + assert callback.should_run_guardrail(data={}, event_type=GuardrailEventHooks.pre_call) is True + + def test_scope_without_logging_only_mode_is_rejected_for_api_writes(self) -> None: + with pytest.raises(ValueError, match="logging_only_scope is set") as exc_info: + self._initialize(mode="pre_call", scope="input", reject_invalid_logging_only_scope=True) + + assert str(exc_info.value) == ( + "Guardrail logging-only-scope-guardrail: logging_only_scope is set, but mode does not include " + "logging_only, so it would never apply. Add logging_only to mode or remove logging_only_scope." + ) + + @pytest.mark.parametrize( + "mode", + ( + "logging_only", + ["pre_call", "logging_only"], + Mode(tags={"audit": "logging_only"}, default="pre_call"), + ), + ) + def test_scope_accepts_logging_only_in_supported_mode_forms(self, mode: str | list[str] | Mode) -> None: + callback: Final = self._initialize(mode=mode, scope="input") + + assert callback.logging_only_scope == "input" + + def test_directional_scope_is_ignored_at_load_when_guardrail_owns_logging_hook(self) -> None: + callback: Final = self._initialize( + mode="logging_only", + scope="input", + callback_type=_LoggingOnlyScopeUnsupportedGuardrail, + assert_registered=True, + ) + + assert callback.logging_only_scope is None + + def test_directional_scope_rejected_for_api_writes_when_guardrail_owns_logging_hook(self) -> None: + with pytest.raises(ValueError, match="logging_only_scope='input' is not supported") as exc_info: + self._initialize( + mode="logging_only", + scope="input", + callback_type=_LoggingOnlyScopeUnsupportedGuardrail, + reject_invalid_logging_only_scope=True, + ) + + assert str(exc_info.value) == ( + "Guardrail logging-only-scope-guardrail: logging_only_scope='input' is not supported by this " + "guardrail, whose logging_only hook scans on its own. Remove logging_only_scope." + ) + + def test_both_scope_accepted_when_guardrail_owns_logging_hook(self) -> None: + callback: Final = self._initialize( + mode="logging_only", + scope="both", + callback_type=_LoggingOnlyScopeUnsupportedGuardrail, + ) + + assert callback.logging_only_scope == "both" + + def test_output_scope_accepted_for_native_lifecycle_guardrail(self) -> None: + callback: Final = self._initialize( + mode="logging_only", + scope="output", + callback_type=_LoggingOnlyScopeNativeGuardrail, + ) + + assert callback.logging_only_scope == "output" + + def test_invalid_scope_fails_litellm_params_validation(self) -> None: + with pytest.raises(ValidationError): + LitellmParams(guardrail="test", mode="logging_only", logging_only_scope="request") + + def test_invalid_scope_literal_keeps_content_filter_registered_and_blocking(self) -> None: + import litellm + from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import ( + ContentFilterGuardrail, + ) + + handler: Final = InMemoryGuardrailHandler() + callback_lists: Final = _all_callback_lists() + callback_snapshots: Final = [list(callback_list) for callback_list in callback_lists] + guardrail: Final = _invalid_scope_content_filter_guardrail() + + try: + result: Final = handler.initialize_guardrail(guardrail=guardrail, source="config") + assert result is not None + callback: Final = handler.guardrail_id_to_custom_guardrail[result["guardrail_id"]] + assert isinstance(callback, ContentFilterGuardrail) + assert callback in litellm.callbacks + assert callback.logging_only_scope is None + assert callback.event_hook == GuardrailEventHooks.pre_call + assert callback._check_blocked_words("pineapple") is not None + finally: + handler.delete_in_memory_guardrail(guardrail["guardrail_id"]) + for callback_list, snapshot in zip(callback_lists, callback_snapshots): + callback_list[:] = snapshot + + def test_invalid_scope_literal_is_rejected_for_strict_initialization_without_callback_leakage(self) -> None: + handler: Final = InMemoryGuardrailHandler() + callback_lists: Final = _all_callback_lists() + callback_snapshots: Final = [list(callback_list) for callback_list in callback_lists] + + with pytest.raises(ValueError, match="logging_only_scope"): + handler.initialize_guardrail( + guardrail=_invalid_scope_content_filter_guardrail(), + source="config", + reject_invalid_logging_only_scope=True, + ) + + assert all(callback_list == snapshot for callback_list, snapshot in zip(callback_lists, callback_snapshots)) + + def test_invalid_scope_literal_does_not_tolerate_other_litellm_params_errors(self) -> None: + with pytest.raises(ValidationError): + parse_tolerant_litellm_params( + { + "guardrail": "litellm_content_filter", + "mode": "pre_call", + "logging_only_scope": "Input", + "default_on": "not-a-bool", + }, + "invalid-scope-content-filter", + ) + + @pytest.mark.asyncio async def test_update_guardrail_in_db_raises_when_row_missing(): prisma_client = MagicMock() @@ -1044,6 +1433,41 @@ async def test_update_guardrail_in_db_raises_when_row_missing(): ) +@pytest.mark.asyncio +async def test_update_guardrail_in_db_persists_raw_sparse_params_verbatim(): + """ + After a rejected PATCH, the endpoint rolls back by writing the stored row's + raw litellm_params through update_guardrail_in_db. A raw dict must be + persisted exactly as stored — a legacy 4-key row stays a 4-key row — instead + of being round-tripped through LitellmParams.model_dump(), which materializes + every field default and rewrites a row the admin never wrote. + """ + prisma_client = MagicMock() + prisma_client.db.litellm_guardrailstable.update = AsyncMock( + return_value={"guardrail_id": "legacy-row", "guardrail_name": "legacy-one"} + ) + legacy_params: Final = { + "guardrail": "litellm_content_filter", + "mode": "pre_call", + "guardrail_name": "legacy-one", + "blocked_words": [{"keyword": "x", "action": "BLOCK"}], + } + + await GuardrailRegistry().update_guardrail_in_db( + guardrail_id="legacy-row", + guardrail=Guardrail( + guardrail_id="legacy-row", + guardrail_name="legacy-one", + litellm_params=legacy_params, + guardrail_info={}, + ), + prisma_client=prisma_client, + ) + + persisted: Final = prisma_client.db.litellm_guardrailstable.update.call_args.kwargs["data"] + assert json.loads(persisted["litellm_params"]) == legacy_params + + def test_reinitialize_guardrail_restores_previous_on_failure(): """A reinitialization whose new params make the guardrail constructor raise must restore the previous instance instead of leaving the guardrail silently removed: @@ -1185,7 +1609,6 @@ _ENCRYPTED_PREFIX = "litellm_enc::" class _Row(dict[str, object]): - def __getattr__(self, name: str) -> object: return self[name] diff --git a/tests/unit/proxy/hooks/test_autorouter_baseline_cache.py b/tests/unit/proxy/hooks/test_autorouter_baseline_cache.py index 0f4fd2ff5cb..1be02d2f269 100644 --- a/tests/unit/proxy/hooks/test_autorouter_baseline_cache.py +++ b/tests/unit/proxy/hooks/test_autorouter_baseline_cache.py @@ -3,6 +3,7 @@ import json from collections.abc import AsyncIterator, Callable, Generator, Mapping from contextlib import contextmanager from datetime import datetime +from itertools import product from types import MappingProxyType from typing import Final, cast from uuid import uuid4 @@ -42,7 +43,7 @@ _MESSAGES_JSON: Final = """[{"role":"user","content":[ _MODELS: Final = _MESSAGES.validate_json("""[ {"model_name":"test-router","litellm_params":{"model":"auto_router/complexity_router", "complexity_router_config":{"tiers":{"SIMPLE":"sonnet","MEDIUM":"sonnet","COMPLEX":"sonnet", - "REASONING":"opus"},"session_affinity":false, + "REASONING":{"model_name":"opus","litellm_params":{"max_tokens":16}}},"session_affinity":false, "keyword_tier_rules":[{"keywords":["USE_OPUS"],"tier":"REASONING"}]}}}, {"model_name":"sonnet","litellm_params":{"model":"anthropic/claude-sonnet-5","api_key":"test-selected"}, "model_info":{"id":"selected"}}, @@ -82,7 +83,9 @@ class _CallContext(TypedDict): def _kwargs(logging_obj: Logging, trusted: bool = True, *, explicit_logging: bool = True) -> _CallContext: - context: Final = _OBJECTS.validate_json('{"litellm_metadata":{"user_api_key_hash":"test-caller-hash"}}') + context: Final = _OBJECTS.validate_json( + '{"max_tokens":16,"litellm_metadata":{"user_api_key_hash":"test-caller-hash"}}' + ) Router._record_routing_decision( # pyright: ignore[reportUnknownMemberType, reportPrivateUsage] # production trusted stamp owner context, StandardLoggingRoutingDecision( @@ -133,7 +136,10 @@ def _upstream(request: httpx.Request) -> httpx.Response: assert isinstance(model, str) stream: Final = body.get("stream") is True content: Final = b"".join(_sse(model=model)) if stream else json.dumps(_message(True, model)).encode() - return httpx.Response(200, content=content, request=request, + return httpx.Response( + 200, + content=content, + request=request, headers=MappingProxyType({"content-type": "text/event-stream" if stream else "application/json"}), ) @@ -182,6 +188,7 @@ async def _call( stream: Final = cast(AsyncIterator[object], response) # cast-ok: iterator checked; all items satisfy object assert tuple([chunk async for chunk in stream]) + class _Capture(CustomLogger): def __init__(self, call_id: str) -> None: self.call_id: Final = call_id @@ -199,9 +206,20 @@ class _Capture(CustomLogger): class _Rig: - def __init__(self, monkeypatch: pytest.MonkeyPatch, *, retries: int = 0, count: TokenCounter = _count) -> None: - self.router: Final = Router(model_list=_MODELS, num_retries=retries, - retry_policy=RetryPolicy(RateLimitErrorRetries=retries), disable_cooldowns=True) + def __init__( + self, + monkeypatch: pytest.MonkeyPatch, + *, + retries: int = 0, + count: TokenCounter = _count, + models: list[dict[str, JsonValue]] = _MODELS, + ) -> None: + self.router: Final = Router( + model_list=models, + num_retries=retries, + retry_policy=RetryPolicy(RateLimitErrorRetries=retries), + disable_cooldowns=True, + ) def router() -> Router: return self.router @@ -218,9 +236,16 @@ class _Rig: monkeypatch.setattr(litellm, "_async_success_callback", [self.capture]) def logging(self, stream: bool = False) -> Logging: - return Logging(model="anthropic/claude-sonnet-5", messages=_MESSAGES.validate_json(_MESSAGES_JSON), - stream=stream, call_type=CallTypes.anthropic_messages.value, start_time=datetime.now(), - litellm_call_id=self.call_id, function_id=self.call_id, kwargs={"litellm_session_id":"baseline-session"}) + return Logging( + model="anthropic/claude-sonnet-5", + messages=_MESSAGES.validate_json(_MESSAGES_JSON), + stream=stream, + call_type=CallTypes.anthropic_messages.value, + start_time=datetime.now(), + litellm_call_id=self.call_id, + function_id=self.call_id, + kwargs={"litellm_session_id": "baseline-session"}, + ) def _observation(payload: Mapping[str, object]) -> CapturedBaselineObservation: @@ -232,7 +257,9 @@ def _observation(payload: Mapping[str, object]) -> CapturedBaselineObservation: @pytest.mark.parametrize("stream,baseline", ((False, False), (True, False), (False, True), (True, True))) async def test_native_logging_captures_usage_without_publishing_hypothetical_savings( - monkeypatch: pytest.MonkeyPatch, stream: bool, baseline: bool, + monkeypatch: pytest.MonkeyPatch, + stream: bool, + baseline: bool, ) -> None: rig: Final = _Rig(monkeypatch) messages: Final = _MESSAGES_JSON.replace("question", "question USE_OPUS") if baseline else _MESSAGES_JSON @@ -283,11 +310,14 @@ async def test_caller_cannot_forge_an_observation_scope(monkeypatch: pytest.Monk assert payload["autorouter_savings"] is None -@pytest.mark.parametrize("model,key,endpoint", ( - ("claude-sonnet-5", "test-first", None), - ("claude-opus-5", "test-second", None), - ("claude-opus-5", "test-first", "https://example.test"), -)) +@pytest.mark.parametrize( + "model,key,endpoint", + ( + ("claude-sonnet-5", "test-first", None), + ("claude-opus-5", "test-second", None), + ("claude-opus-5", "test-first", "https://example.test"), + ), +) async def test_count_memo_is_scoped_to_provider_recipient(model: str, key: str, endpoint: str | None) -> None: counts: Final = iter((5000, 6000)) @@ -304,7 +334,8 @@ async def test_count_memo_is_scoped_to_provider_recipient(model: str, key: str, @pytest.mark.parametrize("stream", (False, True)) async def test_provider_counting_does_not_hold_the_inference_response( - monkeypatch: pytest.MonkeyPatch, stream: bool, + monkeypatch: pytest.MonkeyPatch, + stream: bool, ) -> None: counting: Final = asyncio.Event() release: Final = asyncio.Event() @@ -324,3 +355,744 @@ async def test_provider_counting_does_not_hold_the_inference_response( assert _observation(await rig.capture.payload()).observation.plan is not None finally: release.set() + + +@pytest.mark.parametrize("baseline_effort", (None, "medium")) +@pytest.mark.parametrize( + "automatic_system, caching", + ( + (None, "explicit"), + ("stable system", "request"), + ([{"type": "text", "text": "stable system"}], "request"), + ("stable system", "global"), + ([{"type": "text", "text": "stable system"}], "configured"), + ), +) +async def test_native_tier_switch_uses_baseline_settings_and_preserves_history( + monkeypatch: pytest.MonkeyPatch, + baseline_effort: str | None, + automatic_system: str | list[dict[str, str]] | None, + caching: str, +) -> None: + from litellm.proxy.spend_tracking.baseline_accounting import BaselineHistory, advance_baseline_history + + models: Final = _MESSAGES.validate_python( + [ + { + "model_name": "test-router", + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": { + "tiers": { + "SIMPLE": {"model_name": "sonnet", "litellm_params": {"reasoning_effort": "low"}}, + "MEDIUM": {"model_name": "sonnet", "litellm_params": {"reasoning_effort": "low"}}, + "COMPLEX": {"model_name": "sonnet", "litellm_params": {"reasoning_effort": "high"}}, + "REASONING": "opus", + }, + "session_affinity": False, + "keyword_tier_rules": [{"keywords": ["ESCALATE"], "tier": "COMPLEX"}], + }, + }, + }, + _MODELS[1], + { + "model_name": "opus", + "model_info": {"id": "baseline"}, + "litellm_params": { + "model": "anthropic/claude-opus-5", + "api_key": "test-selected", + **({"reasoning_effort": baseline_effort} if baseline_effort else {}), + }, + }, + ] + ) + rig: Final = _Rig(monkeypatch, models=models) + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", caching == "global") + controls: Final = ( + { + "cache_control_injection_points": [ + {"location": "message", "role": "system", "control": {"type": "ephemeral", "ttl": "1h"}}, + {"location": "message", "index": -1, "control": {"type": "ephemeral", "ttl": "1h"}}, + ] + } + if caching == "configured" + else {"enable_prompt_caching": caching == "request"} + ) + captures: Final[asyncio.Queue[CapturedBaselineObservation]] = asyncio.Queue() + with _transport(_upstream) as route: + for suffix in ("", " ESCALATE"): + log: Final = rig.logging() + await rig.router.anthropic_messages( + model="test-router", + max_tokens=4096, + messages=( + [{"role": "user", "content": "question" + suffix}] + if automatic_system is not None + else _MESSAGES.validate_json(_MESSAGES_JSON.replace("question", "question" + suffix)) + ), + system=automatic_system, + **controls, + litellm_logging_obj=log, + litellm_call_id=rig.call_id, + litellm_metadata={"user_api_key_hash": "test-caller-hash"}, + litellm_session_id="native-tiers", + ) + captures.put_nowait(_observation(await rig.capture.payload())) + first_wire, second_wire = (_JSON_OBJECT.validate_json(call.request.content) for call in route.calls) + assert (first_wire.get("thinking"), first_wire.get("output_config")) != ( + second_wire.get("thinking"), + second_wire.get("output_config"), + ) + first, second = (captures.get_nowait() for _ in range(2)) + assert first.scope == second.scope + assert first.observation.plan is not None and second.observation.plan is not None + assert first.observation.plan.breakpoints[0] == second.observation.plan.breakpoints[0] + assert len(first.observation.plan.breakpoints) == (2 if automatic_system is not None else 1) + history, _ = advance_baseline_history( + BaselineHistory(first_at=0.0), + (first.observation.model_copy(update={"request_id": "first", "started_at": 10000.0, "available_at": 10001.0}),), + ) + _, result = advance_baseline_history( + history, + ( + second.observation.model_copy( + update={"request_id": "second", "started_at": 10020.0, "available_at": 10021.0} + ), + ), + ) + assert result[0].usage is not None and result[0].usage.prompt_tokens_details.cached_tokens == 5000 + + +@pytest.mark.parametrize("call_type", (CallTypes.acompletion, CallTypes.aresponses, CallTypes.anthropic_messages)) +async def test_plain_requests_do_not_initialize_or_warn( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, + call_type: CallTypes, +) -> None: + rig: Final = _Rig(monkeypatch) + logging: Final = rig.logging() + await rig.hook.async_pre_call_deployment_hook( + { + "litellm_logging_obj": logging, + "litellm_metadata": {"session_id": "ordinary"}, + }, + call_type, + ) + assert logging.baseline_cache_context is None + assert "baseline observation could not be initialized" not in caplog.text + assert not rig.hook.counts + + +async def test_plain_fallback_invalidates_existing_autorouter_capture(monkeypatch: pytest.MonkeyPatch) -> None: + rig: Final = _Rig(monkeypatch) + logging: Final = rig.logging() + await rig.hook.async_pre_call_deployment_hook(_kwargs(logging), CallTypes.anthropic_messages) + assert logging.baseline_cache_context is not None + await rig.hook.async_pre_call_deployment_hook({"litellm_logging_obj": logging}, CallTypes.anthropic_messages) + assert logging.baseline_observation is not None + assert logging.baseline_observation.observation.reason == "retried_request" + + +async def test_native_count_finishing_after_quarter_worker_budget_keeps_plan_and_spend( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.litellm_core_utils import logging_worker + from litellm.litellm_core_utils.logging_worker import LoggingWorker + + release: Final = asyncio.Event() + + async def count(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int: + if not release.is_set(): + asyncio.get_running_loop().call_later(2.3, release.set) + await release.wait() + return await _count(model, api_key, body) + + worker: Final = LoggingWorker(timeout=8.0) + monkeypatch.setattr(logging_worker, "GLOBAL_LOGGING_WORKER", worker) + rig: Final = _Rig(monkeypatch, count=count) + try: + with _transport(_upstream): + await _call(rig.router, rig.logging()) + payload: Final = await rig.capture.payload() + observed: Final = _observation(payload).observation + assert observed.outcome == "complete" and observed.reason is None + assert observed.plan is not None and observed.plan.breakpoints[0].prefix_tokens == 5000 + assert payload["response_cost"] is not None and worker._timeout_total == 0 + finally: + release.set() + await worker.stop() + + +@pytest.mark.parametrize( + "options,on_deployment", + ( + ({"thinking": {"type": "enabled", "budget_tokens": 2048}}, False), + ({"extra_body": {"speed": "fast", "output_config": {"effort": "high"}}}, False), + *product( + ( + {"container": {"id": "container_test"}}, + {"mcp_servers": [{"type": "url", "name": "test", "url": "https://example.com/mcp"}]}, + {"inference_geo": "us"}, + {"safeguards": [{"type": "default"}]}, + ), + (False, True), + ), + ), +) +async def test_native_baseline_identity_keeps_the_actual_transformed_body( + monkeypatch: pytest.MonkeyPatch, options: dict[str, JsonValue], on_deployment: bool +) -> None: + models: Final = _MESSAGES.validate_python( + [ + *_MODELS[:2], + { + **_MODELS[2], + "litellm_params": { + **_JSON_OBJECT.validate_python(_MODELS[2]["litellm_params"]), + **(options if on_deployment else {}), + }, + }, + ] + ) + rig: Final = _Rig(monkeypatch, models=models) + log: Final = rig.logging() + with _transport(_upstream): + await rig.router.anthropic_messages( + model="test-router", + max_tokens=16, + messages=_MESSAGES.validate_json(_MESSAGES_JSON.replace("question", "question USE_OPUS")), + litellm_logging_obj=log, + litellm_call_id=rig.call_id, + litellm_metadata={"user_api_key_hash": "test-caller-hash"}, + litellm_session_id="native-identical", + **({} if on_deployment else options), + ) + observed: Final = _observation(await rig.capture.payload()).observation + assert observed.outcome == "complete" + assert log.baseline_cache_context is not None + assert observed.baseline_equivalent and observed.usage is not None, ( + log.baseline_cache_context.baseline_body, + log.baseline_cache_context.selected_body_digest, + ) + + from litellm.proxy.spend_tracking.baseline_accounting import BaselineHistory, advance_baseline_history + + _, estimates = advance_baseline_history(BaselineHistory(), (observed,)) + assert estimates[0].provenance == "observed_identical" and estimates[0].usage == observed.usage + + +@pytest.mark.parametrize("tier_limit", (8, 16)) +@pytest.mark.parametrize("extra", ({}, {"max_tokens": 8})) +async def test_native_baseline_identity_respects_caller_limit_and_tier_override( + monkeypatch: pytest.MonkeyPatch, tier_limit: int, extra: dict[str, int] +) -> None: + models: Final = _MESSAGES.validate_python( + [ + { + "model_name": "test-router", + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": { + "tiers": { + "SIMPLE": {"model_name": "opus", "litellm_params": {"max_tokens": tier_limit}}, + "MEDIUM": {"model_name": "opus", "litellm_params": {"max_tokens": tier_limit}}, + "COMPLEX": "opus", + "REASONING": "opus", + }, + "session_affinity": False, + }, + }, + }, + {**_MODELS[2], "litellm_params": {**_MODELS[2]["litellm_params"], "max_tokens": 64}}, + ] + ) + rig: Final = _Rig(monkeypatch, models=models) + with _transport(_upstream) as route: + await rig.router.anthropic_messages( + model="test-router", + max_tokens=8, + messages=_MESSAGES.validate_json(_MESSAGES_JSON), + litellm_logging_obj=rig.logging(), + litellm_call_id=rig.call_id, + litellm_metadata={"user_api_key_hash": "test-caller-hash"}, + litellm_session_id="native-limits", + extra_body=extra, + ) + observed: Final = _observation(await rig.capture.payload()).observation + wire: Final = _JSON_OBJECT.validate_json(route.calls.last.request.content) + assert wire["max_tokens"] == tier_limit + assert observed.baseline_equivalent == (tier_limit == 8) + + +@pytest.mark.parametrize("nested", (False, True)) +async def test_native_baseline_projection_matches_wire_parameter_placement( + monkeypatch: pytest.MonkeyPatch, + nested: bool, +) -> None: + counted: Final[asyncio.Queue[Mapping[str, JsonValue]]] = asyncio.Queue() + + async def count(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int: + counted.put_nowait(body) + return await _count(model, api_key, body) + + models: Final = _MESSAGES.validate_python( + [ + { + **entry, + "litellm_params": { + **_JSON_OBJECT.validate_python(entry["litellm_params"]), + "model": "anthropic/claude-opus-5", + }, + } + if entry["model_name"] == "sonnet" + else entry + for entry in _MODELS + ] + ) + rig: Final = _Rig(monkeypatch, count=count, models=models) + settings: Final = {"speed": "standard", "thinking": {"type": "adaptive"}, "output_config": {"effort": "medium"}} + with _transport(_upstream) as route: + await rig.router.anthropic_messages( + model="test-router", + max_tokens=4096, + messages=_MESSAGES.validate_json(_MESSAGES_JSON), + litellm_logging_obj=rig.logging(), + litellm_call_id=rig.call_id, + litellm_metadata={"user_api_key_hash": "test-caller-hash"}, + litellm_session_id="native-placement", + **({"extra_body": settings} if nested else settings), + ) + captured: Final = _observation(await rig.capture.payload()) + wire: Final = _JSON_OBJECT.validate_json(route.calls.last.request.content) + assert captured.observation.plan is not None and not captured.observation.baseline_equivalent + projected: Final = counted.get_nowait() + assert {key: projected[key] for key in settings if key in projected} == { + key: wire[key] for key in settings if key in wire + } + assert {key: wire[key] for key in settings if key in wire} == ({} if nested else settings) + + +_PARITY_TOOL: Final = {"name": "custom", "input_schema": {"type": "object"}, "cache_control": {"type": "ephemeral"}} +_PARITY_SYSTEM: Final = [{"type": "text", "text": "stable system", "cache_control": {"type": "ephemeral"}}] +_PARITY_POINTS: Final = [{"location": "message", "role": "system"}, {"location": "message", "index": -1}] + + +@pytest.mark.parametrize( + "caller,selected,baseline,summary", + ( + pytest.param({"extra_body": {"cache_control": {"type": "ephemeral"}}}, {}, {}, False, id="envelope-control"), + pytest.param({"extra_body": {"system": _PARITY_SYSTEM}}, {}, {}, False, id="envelope-system"), + pytest.param( + {"extra_body": {"messages": _MESSAGES.validate_json(_MESSAGES_JSON)}}, {}, {}, False, id="envelope-messages" + ), + pytest.param({}, {"tools": [_PARITY_TOOL]}, {}, False, id="selected-tool-mark"), + pytest.param({}, {}, {"tools": [_PARITY_TOOL]}, False, id="baseline-tool-mark"), + pytest.param({}, {"system": _PARITY_SYSTEM}, {"system": "baseline system"}, False, id="selected-system-mark"), + pytest.param({}, {"system": "selected system"}, {"system": _PARITY_SYSTEM}, False, id="baseline-system-mark"), + pytest.param({"system": None}, {}, {"system": "configured system"}, False, id="null-system"), + pytest.param({"thinking": None}, {}, {"thinking": {"type": "adaptive"}}, False, id="null-thinking"), + pytest.param({"tools": None}, {}, {"tools": [_PARITY_TOOL]}, False, id="null-tools"), + pytest.param({"verbosity": "low", "instructions": "ignored"}, {}, {}, False, id="ignored-native-options"), + pytest.param( + {"messages": _MESSAGES.validate_json(_MESSAGES_JSON.replace("stable", " "))}, + {}, + {}, + False, + id="empty-marked-block", + ), + pytest.param({"thinking": {"type": "adaptive"}}, {}, {}, True, id="reasoning-summary"), + pytest.param( + {"thinking": {"type": "adaptive"}, "additional_drop_params": ["thinking.display"]}, + {}, + {}, + True, + id="drop-nested-option", + ), + pytest.param( + {"cache_control_injection_points": _PARITY_POINTS}, + {"tools": [{**_PARITY_TOOL, "name": f"custom_{index}"} for index in range(4)]}, + {}, + False, + id="configured-cap", + ), + ), +) +async def test_native_baseline_projection_matches_direct_baseline_request( + monkeypatch: pytest.MonkeyPatch, + caller: dict[str, JsonValue], + selected: dict[str, JsonValue], + baseline: dict[str, JsonValue], + summary: bool, +) -> None: + counted: Final[asyncio.Queue[Mapping[str, JsonValue]]] = asyncio.Queue() + + async def count(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int: + counted.put_nowait(body) + return await _count(model, api_key, body) + + def upstream(request: httpx.Request) -> httpx.Response: + body: Final = _JSON_OBJECT.validate_json(request.content) + model: Final = body.get("model") + assert isinstance(model, str) + return httpx.Response( + 200, + request=request, + json={ + **_message(True, model), + "usage": {"input_tokens": 6000, "output_tokens": 10}, + }, + ) + + models: Final = _MESSAGES.validate_python( + [ + _MODELS[0], + { + **_MODELS[1], + "litellm_params": {**_JSON_OBJECT.validate_python(_MODELS[1]["litellm_params"]), **selected}, + }, + { + **_MODELS[2], + "litellm_params": {**_JSON_OBJECT.validate_python(_MODELS[2]["litellm_params"]), **baseline}, + }, + ] + ) + rig: Final = _Rig(monkeypatch, models=models, count=count) + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", False) + monkeypatch.setattr(litellm, "reasoning_auto_summary", summary) + monkeypatch.delenv("LITELLM_REASONING_AUTO_SUMMARY", raising=False) + request: Final = { + "messages": [{"role": "user", "content": "question"}], + **({"system": "stable system"} if "system" not in selected and "system" not in baseline else {}), + "max_tokens": 4096, + "enable_prompt_caching": True, + **caller, + } + with _transport(upstream) as route: + await rig.router.anthropic_messages(model="opus", **_JSON_OBJECT.validate_python(request)) + direct: Final = _JSON_OBJECT.validate_json(route.calls.last.request.content) + await rig.router.anthropic_messages( + model="test-router", + litellm_logging_obj=rig.logging(), + litellm_call_id=rig.call_id, + litellm_metadata={"user_api_key_hash": "test-caller-hash"}, + litellm_session_id="native-projection-parity", + **_JSON_OBJECT.validate_python(request), + ) + captured: Final = _observation(await rig.capture.payload()) + assert captured.observation.plan is not None, captured.observation.reason + projected: Final = counted.get_nowait() + assert {key: value for key, value in projected.items() if key not in ("metadata", "stream")} == { + key: value for key, value in direct.items() if key not in ("metadata", "stream") + } + + +@pytest.mark.parametrize( + "selected,baseline,usage_field,observed_value,multiplier", + ( + ({}, {"speed": "fast"}, "speed", "standard", 3.0), + ({"inference_geo": "us"}, {}, "inference_geo", "us", 1.0), + ), +) +async def test_native_baseline_prices_projected_settings_without_changing_actual_spend( + monkeypatch: pytest.MonkeyPatch, + selected: dict[str, JsonValue], + baseline: dict[str, JsonValue], + usage_field: str, + observed_value: str, + multiplier: float, +) -> None: + from litellm.proxy.spend_tracking.baseline_accounting import BaselineHistory, advance_baseline_history + from litellm.proxy.spend_tracking.savings import baseline_cost_snapshot, price_baseline_comparison + from litellm.types.utils import ModelInfo + + def upstream(request: httpx.Request) -> httpx.Response: + body: Final = _JSON_OBJECT.validate_json(request.content) + model: Final = body.get("model") + assert isinstance(model, str) + assert body.get(usage_field) == selected.get(usage_field) + return httpx.Response( + 200, + request=request, + json={ + **_message(True, model), + "usage": { + "input_tokens": 6000, + "output_tokens": 10, + usage_field: observed_value, + "cache_creation": {"ephemeral_5m_input_tokens": 0, "ephemeral_1h_input_tokens": 0}, + }, + }, + ) + + rig: Final = _Rig( + monkeypatch, + models=_MESSAGES.validate_python( + [ + _MODELS[0], + { + **_MODELS[1], + "litellm_params": { + **_JSON_OBJECT.validate_python(_MODELS[1]["litellm_params"]), + **selected, + }, + }, + { + **_MODELS[2], + "litellm_params": { + **_JSON_OBJECT.validate_python(_MODELS[2]["litellm_params"]), + **baseline, + }, + }, + ] + ), + ) + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", False) + with _transport(upstream): + response: Final = _JSON_OBJECT.validate_python( + await rig.router.anthropic_messages( + model="test-router", + max_tokens=16, + messages=[{"role": "user", "content": "question"}], + litellm_logging_obj=rig.logging(), + litellm_call_id=rig.call_id, + litellm_metadata={"user_api_key_hash": "test-caller-hash"}, + litellm_session_id="baseline-speed", + ) + ) + payload: Final = await rig.capture.payload() + captured: Final = _observation(payload) + _, estimates = advance_baseline_history(BaselineHistory(), (captured.observation,)) + estimate: Final = estimates[0] + assert captured.prices is not None + prices: Final[ModelInfo] = { + **captured.prices, + "input_cost_per_token": 1e-6, + "output_cost_per_token": 2e-6, + "provider_specific_entry": {"fast": 3.0, "us": 2.0}, + } + actual: Final = payload["response_cost"] + assert isinstance(actual, float) + snapshot: Final = baseline_cost_snapshot( + captured.model, + prices, + actual, + _OBJECTS.validate_python(payload["cost_breakdown"]), + None, + ) + comparison: Final = price_baseline_comparison(snapshot, estimate.usage, estimate.provenance) + assert comparison is not None and snapshot.actual_token_cost is not None, estimate.reason + assert comparison.baseline == pytest.approx( + actual + (6000 * 1e-6 + 10 * 2e-6) * multiplier - snapshot.actual_token_cost + ) + assert comparison.actual == actual + assert _JSON_OBJECT.validate_python(response["usage"])[usage_field] == observed_value + + +async def test_native_request_rewritten_after_capture_preserves_spend_without_guessing_baseline( + monkeypatch: pytest.MonkeyPatch, +) -> None: + class RewriteSystem(CustomLogger): + async def async_pre_call_deployment_hook( + self, kwargs: Mapping[str, object], call_type: CallTypes | None + ) -> dict[str, object]: + return {**kwargs, "system": "hook system"} + + rig: Final = _Rig(monkeypatch) + monkeypatch.setattr(litellm, "callbacks", [rig.hook, RewriteSystem()]) + with _transport(_upstream) as route: + await _call(rig.router, rig.logging()) + payload: Final = await rig.capture.payload() + wire: Final = _JSON_OBJECT.validate_json(route.calls.last.request.content) + observed: Final = _observation(payload).observation + assert wire["system"] == "hook system" + assert observed.reason == "unsupported_request_transformation" and observed.plan is None + assert observed.usage is not None + actual: Final = payload["response_cost"] + assert isinstance(actual, float) and actual > 0 + + +@pytest.mark.parametrize("history", ("long_session", "non_ascii")) +async def test_native_baseline_models_long_and_non_ascii_history(monkeypatch: pytest.MonkeyPatch, history: str) -> None: + rounds: Final = tuple( + message + for index in range(1200) + for message in ( + {"role": "assistant", "content": [{"type": "tool_use", "id": f"t{index}", "name": "Read", "input": {}}]}, + {"role": "user", "content": [{"type": "tool_result", "tool_use_id": f"t{index}", "content": "ok"}]}, + ) + ) + prefix: Final = ( + [{"role": "user", "content": "start"}, *rounds] + if history == "long_session" + else [{"role": "user", "content": "a" * 400_000 + "é"}, {"role": "assistant", "content": "ok"}] + ) + messages: Final = json.dumps([*prefix, *_MESSAGES.validate_json(_MESSAGES_JSON)]) + rig: Final = _Rig(monkeypatch) + with _transport(_upstream): + await _call(rig.router, rig.logging(), messages=messages) + observed: Final = _observation(await rig.capture.payload()).observation + assert observed.outcome == "complete" and observed.plan is not None + + +async def test_native_baseline_abstains_after_selected_tier_compaction(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.router_strategy.complexity_router.context_compaction import compaction_executor + + monkeypatch.setitem( + litellm.model_cost, + "summary-fixture", + { + "litellm_provider": "anthropic", + "mode": "chat", + "max_input_tokens": 32000, + "max_output_tokens": 4096, + "supports_anthropic_compaction": True, + }, + ) + models: Final = _MESSAGES.validate_python( + [ + { + "model_name": "test-router", + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": { + "tiers": {"SIMPLE": "sonnet", "MEDIUM": "opus", "COMPLEX": "opus", "REASONING": "opus"}, + "keyword_tier_rules": [{"keywords": ["answer"], "tier": "SIMPLE"}], + "session_affinity": False, + "enable_context_window_escalation": False, + "max_tokens_from_tier_model": False, + "context_compaction": {"model": "compactor", "max_tokens": 512}, + }, + }, + }, + { + "model_name": "sonnet", + "litellm_params": {"model": "anthropic/claude-sonnet-5", "api_key": "test-selected"}, + "model_info": {"id": "selected", "max_input_tokens": 512, "max_output_tokens": 64}, + }, + { + "model_name": "opus", + "litellm_params": {"model": "anthropic/claude-opus-5", "api_key": "test-selected"}, + "model_info": {"id": "baseline", "max_input_tokens": 200000, "max_output_tokens": 4096}, + }, + { + "model_name": "compactor", + "litellm_params": {"model": "anthropic/summary-fixture", "api_key": "test-compactor"}, + "model_info": {"id": "compactor"}, + }, + ] + ) + + async def summarize(protocol: object, request: object, parent_model: object = None) -> Mapping[str, object]: + return { + "stop_reason": "compaction", + "content": [{"type": "compaction", "content": "compacted", "signature": "s"}], + "usage": {"input_tokens": 0, "output_tokens": 0}, + } + + messages: Final = json.dumps( + [ + {"role": "user", "content": "Background detail. " * 300}, + {"role": "assistant", "content": "Recorded"}, + {"role": "user", "content": "Answer briefly"}, + ] + ) + rig: Final = _Rig(monkeypatch, models=models) + token: Final = compaction_executor.set(summarize) + try: + with _transport(_upstream) as route: + await _call(rig.router, rig.logging(), messages=messages) + observed: Final = _observation(await rig.capture.payload()).observation + wire: Final = route.calls.last.request.content.decode() + finally: + compaction_executor.reset(token) + assert "compacted" in wire and "Background detail" not in wire + assert observed.reason == "unsupported_request_transformation" and observed.plan is None + assert observed.usage is not None + + +async def test_selected_tier_cache_markers_do_not_hide_an_unmarked_baseline_plan( + monkeypatch: pytest.MonkeyPatch, +) -> None: + models: Final = _MESSAGES.validate_python( + [ + _MODELS[0], + { + "model_name": "sonnet", + "litellm_params": { + "model": "anthropic/claude-sonnet-5", + "api_key": "test-selected", + "cache_control_injection_points": [{"location": "message", "role": "user", "index": -1}], + }, + "model_info": {"id": "selected"}, + }, + _MODELS[2], + ] + ) + rig: Final = _Rig(monkeypatch, models=models) + with _transport(_upstream) as route: + await _call( + rig.router, + rig.logging(), + messages='[{"role":"user","content":[{"type":"text","text":"stable"},{"type":"text","text":"question"}]}]', + ) + observed: Final = _observation(await rig.capture.payload()).observation + wire: Final = route.calls.last.request.content.decode() + assert "cache_control" in wire + assert observed.outcome == "complete" and not observed.baseline_equivalent + assert observed.reason is None and observed.plan is not None and not observed.plan.breakpoints + + +@pytest.mark.parametrize("recovery", ("retry", "fallback")) +async def test_tier_pins_never_enter_the_caller_snapshot_on_later_routing_passes( + monkeypatch: pytest.MonkeyPatch, recovery: str +) -> None: + pinned: Final = {"model_name": "first", "litellm_params": {"reasoning_effort": "high", "max_tokens": 777}} + models: Final = _MESSAGES.validate_python( + [ + { + "model_name": "test-router", + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": { + "tiers": {"SIMPLE": pinned, "MEDIUM": pinned, "COMPLEX": pinned, "REASONING": "opus"}, + "session_affinity": False, + }, + }, + }, + { + "model_name": "fallback-router", + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": { + "tiers": {"SIMPLE": "sonnet", "MEDIUM": "sonnet", "COMPLEX": "sonnet", "REASONING": "opus"}, + "session_affinity": False, + }, + }, + }, + { + "model_name": "first", + "litellm_params": { + "model": "anthropic/claude-sonnet-5" if recovery == "retry" else "anthropic/claude-haiku-5", + "api_key": "test-selected", + }, + "model_info": {"id": "first"}, + }, + *_MODELS[1:], + ] + ) + rig: Final = _Rig(monkeypatch, models=models, retries=1 if recovery == "retry" else 0) + rig.router.fallbacks = [{"test-router": ["fallback-router"]}] + + def upstream(request: httpx.Request) -> httpx.Response: + return _upstream(request) if route.call_count else _error(request, 429, "first attempt") + + log: Final = rig.logging() + with _transport(upstream) as route: + await _call(rig.router, log) + await rig.capture.payload() + first_wire: Final = _JSON_OBJECT.validate_json(route.calls[0].request.content) + assert first_wire.get("output_config") == {"effort": "high"} and first_wire.get("max_tokens") == 777 + context: Final = log.baseline_cache_context + assert context is not None and context.baseline_body is not None + assert context.baseline_body.get("max_tokens") == 16 + assert "output_config" not in context.baseline_body and "thinking" not in context.baseline_body diff --git a/tests/unit/proxy/lens/test_endpoints.py b/tests/unit/proxy/lens/test_endpoints.py index 34d51110761..e1c04e3d6f4 100644 --- a/tests/unit/proxy/lens/test_endpoints.py +++ b/tests/unit/proxy/lens/test_endpoints.py @@ -1,27 +1,37 @@ +import asyncio +from collections.abc import Callable, Mapping from datetime import datetime, timedelta, timezone from types import SimpleNamespace from typing import Final import pytest from fastapi import HTTPException -from pydantic import ValidationError +from pydantic import TypeAdapter, ValidationError import litellm from litellm import Router from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.lens.endpoints import ( + claim_due, + get_signals, list_agents, + put_signals, read_reviews, result, run_settings, run_window, trace_findings, + trace_signal_statuses, user_scope, validate_model, + validate_signal_model, watchable, watching, worker_supports_model, ) +from litellm.proxy.lens.endpoints import ( + sample as worker_sample, +) from litellm.proxy.lens.models import ( ActivitySelection, Coverage, @@ -35,9 +45,12 @@ 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.signals import SignalConfig, StoredTraceSignal from litellm.proxy.lens.state import claim_job, queue_job, replace_job +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 @@ -64,6 +77,114 @@ class ResultDatabase: return len(self.completed) +class SignalStatusDatabase: + def __init__(self, config: SignalConfig, rows: Mapping[str, StoredTraceSignal]) -> None: + self.config: Final = config + self.rows: Final = rows + self.saved: Final[asyncio.Queue[tuple[object, ...]]] = asyncio.Queue() + + async def query_raw(self, query: str, *args: object) -> object: + if '"LiteLLM_LensSignalConfig"' in query: + return ({"data": self.config.model_dump(mode="json")},) + payload: Final = args[0] + assert isinstance(payload, str) + requested: Final = TypeAdapter(tuple[TraceIdentity, ...]).validate_json(payload) + return tuple( + {"data": row.model_dump(mode="json")} + for identity in requested + if (row := self.rows.get(identity.trace_id)) is not None + ) + + async def execute_raw(self, query: str, *args: object) -> int: + await self.saved.put(args) + return 1 + + +def signal_router() -> Router: + return Router( + model_list=[ + { + "model_name": "decision", + "litellm_params": {"model": "openai/test-decision", "api_key": "test-key"}, + "model_info": {"mode": "evaluation"}, + }, + { + "model_name": "chat", + "litellm_params": {"model": "openai/test-chat", "api_key": "test-key"}, + "model_info": {"mode": "chat"}, + }, + ] + ) + + +@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", @@ -394,6 +515,144 @@ async def test_trace_finding_counts_require_investigation_read_access() -> None: assert error.value.status_code == 403 +@pytest.mark.asyncio +async def test_signal_endpoints_return_statuses_in_request_order_for_admin_viewers( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.proxy import proxy_server + + config: Final = SignalConfig(model="decision") + rows: Final = { + "pending": StoredTraceSignal( + trace_id="pending", + config_key=config.key(), + span_count=1, + claimed_until=NOW + timedelta(minutes=1), + data={"status": "pending", "scores": {}, "model": "decision", "error": ""}, + ), + "classified": StoredTraceSignal( + trace_id="classified", + config_key=config.key(), + span_count=1, + classified_at=NOW, + data={ + "status": "classified", + "scores": {"user_frustration": 0.7, "missing_capability": 0.8}, + "model": "decision", + "error": "", + }, + ), + "failed": StoredTraceSignal( + trace_id="failed", + config_key=config.key(), + span_count=1, + classified_at=NOW, + data={"status": "failed", "scores": {}, "model": "decision", "error": "classification failed"}, + ), + "stale": StoredTraceSignal( + trace_id="stale", + config_key="old-config", + span_count=1, + classified_at=NOW, + data={ + "status": "classified", + "scores": {"user_frustration": 1.0}, + "model": "old", + "error": "", + }, + ), + } + database: Final = SignalStatusDatabase(config, rows) + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=database)) + viewer: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY) + request: Final = TraceFindingsRequest( + traces=tuple( + TraceIdentity(trace_id=trace_id) for trace_id in ("failed", "classified", "missing", "pending", "stale") + ) + ) + + assert await get_signals(viewer) == config + results: Final = await trace_signal_statuses(request, viewer) + + assert tuple((result.trace_id, result.status) for result in results) == ( + ("failed", "failed"), + ("classified", "classified"), + ("missing", "unclassified"), + ("pending", "pending"), + ("stale", "unclassified"), + ) + assert tuple((flag.signal_id, flag.name, flag.score) for flag in results[1].flags) == ( + ("missing_capability", "Missing capability", 0.8), + ("user_frustration", "User frustration", 0.7), + ) + + +@pytest.mark.asyncio +async def test_signal_endpoints_require_connected_postgres(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "prisma_client", None) + auth: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY) + + with pytest.raises(HTTPException) as error: + await get_signals(auth) + + assert error.value.status_code == 503 + + +@pytest.mark.asyncio +async def test_put_signals_saves_config_for_admin(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy import proxy_server + + database: Final = SignalStatusDatabase(SignalConfig(), {}) + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=database)) + monkeypatch.setattr(proxy_server, "llm_router", signal_router()) + auth: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + body: Final = SignalConfig(model="decision", threshold=0.7) + + assert await put_signals(body, auth) == body + + saved: Final = await database.saved.get() + assert saved[0] == "global" + assert isinstance(saved[1], str) + assert SignalConfig.model_validate_json(saved[1]) == body + + +def test_signal_model_requires_a_ready_router() -> None: + with pytest.raises(HTTPException) as error: + validate_signal_model(SignalConfig(model="decision"), None) + + assert error.value.status_code == 400 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("role", (LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY)) +async def test_put_signals_rejects_non_admin_roles(role: LitellmUserRoles) -> None: + auth: Final = UserAPIKeyAuth(user_role=role) + with pytest.raises(HTTPException) as error: + await put_signals(SignalConfig(model="decision"), auth) + assert error.value.status_code == 403 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("model", ("chat", "unconfigured")) +async def test_put_signals_rejects_chat_and_unknown_model_groups(monkeypatch: pytest.MonkeyPatch, model: str) -> None: + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "llm_router", signal_router()) + auth: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + + with pytest.raises(HTTPException) as error: + await put_signals(SignalConfig(model=model), auth) + + assert error.value.status_code == 400 + assert error.value.detail == "Choose a System 1 model (evaluation mode) configured on this proxy" + + +def test_signal_model_accepts_only_evaluation_mode_groups() -> None: + assert validate_signal_model(SignalConfig(model="decision"), signal_router()) is None + + @pytest.mark.parametrize( "role", (LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, LitellmUserRoles.TEAM), @@ -624,3 +883,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_signals.py b/tests/unit/proxy/lens/test_signals.py new file mode 100644 index 00000000000..d2c4eda3966 --- /dev/null +++ b/tests/unit/proxy/lens/test_signals.py @@ -0,0 +1,1002 @@ +import asyncio +import json +from collections.abc import AsyncGenerator, Mapping, Sequence +from contextlib import asynccontextmanager +from datetime import datetime, timedelta, timezone +from itertools import chain +from types import MappingProxyType, SimpleNamespace +from typing import Final + +import pytest +from pydantic import JsonValue, TypeAdapter, ValidationError + +from litellm.proxy.lens.models import Execution, Scope, TraceIdentity +from litellm.proxy.lens.repository import Database, Row +from litellm.proxy.lens.signal_repository import SignalRepository +from litellm.proxy.lens.signals import ( + DEFAULT_SIGNALS, + SIGNAL_CLAIM_LEASE, + SIGNAL_MAX_SCAN_PAGES, + SIGNAL_TASK, + DecisionQuestions, + DecisionState, + Signal, + SignalAttempt, + SignalClassifier, + SignalConfig, + SignalData, + SignalStep, + StoredTraceSignal, + candidate, + run_signal_loop, + run_signal_tick, + signal_state, + trace_signals, +) +from litellm.proxy.lens.sources import SourceReader +from litellm.rust_bridge.trace.generated.models import ( + ActivityAvailability, + AgentRow, + CountRow, + ExecutionRow, + LensAccessParams, + LensContentParams, + LensEvidenceParams, + LensSampleParams, + PartRow, +) +from litellm.types.decisions import DecisionsResponse +from litellm.types.decisions import NoulAnswer as DecisionsNoulAnswer + +NOW: Final = datetime(2026, 10, 7, 12, tzinfo=timezone.utc) +CURRENT_CONFIG_KEY: Final = SignalConfig(model="decision").key() +_SIGNAL_STEPS: Final[TypeAdapter[tuple[SignalStep, ...]]] = TypeAdapter(tuple[SignalStep, ...]) +_STORED_DATA: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) + + +def execution(identity: str, span_count: int = 1) -> Execution: + return Execution( + id=identity, + source="traces", + trace_id=identity, + team_id="", + name=identity, + start_time="", + span_count=span_count, + ) + + +def part(identity: str, content: str) -> PartRow: + return PartRow( + span_id=identity, + parent_span_id="", + name=identity, + kind="agent", + start_time="", + end_time="", + content=content, + truncated=0, + ) + + +def stored_trace( + config_key: str, + *, + trace_id: str = "trace", + status: str = "classified", + span_count: int = 1, + claimed_until: datetime | None = None, + classified_at: datetime | None = NOW - timedelta(minutes=10), + scores: dict[str, float] | None = None, + error: str = "", +) -> StoredTraceSignal: + return StoredTraceSignal( + trace_id=trace_id, + trace_ref="", + config_key=config_key, + span_count=span_count, + claimed_until=claimed_until, + classified_at=classified_at, + data=_STORED_DATA.validate_python( + { + "status": status, + "scores": scores or {}, + "model": "decision", + "error": error, + } + ), + ) + + +class SignalStorage: + def __init__( + self, + executions: tuple[ExecutionRow, ...] = (), + parts: tuple[PartRow, ...] = (), + ) -> None: + self.executions: Final = executions + self.parts: Final = parts + + async def lens_availability(self, parameters: LensAccessParams) -> Sequence[ActivityAvailability]: + return () + + async def lens_agents(self, parameters: LensAccessParams) -> Sequence[AgentRow]: + return () + + async def lens_sample(self, parameters: LensSampleParams) -> Sequence[ExecutionRow]: + return self.executions + + async def lens_content(self, parameters: LensContentParams) -> Sequence[PartRow]: + return self.parts or (part(parameters.id, parameters.id),) + + async def lens_evidence(self, parameters: LensEvidenceParams) -> Sequence[CountRow]: + return () + + +class PagedSignalStorage(SignalStorage): + def __init__(self, pages: tuple[tuple[PartRow, ...], ...]) -> None: + super().__init__() + self.pages: Final = pages + + async def lens_content(self, parameters: LensContentParams) -> Sequence[PartRow]: + index: Final = int(parameters.cursor) if parameters.cursor else 0 + return self.pages[index] + + +class PagedSampleStorage(SignalStorage): + def __init__(self, pages: tuple[tuple[ExecutionRow, ...], ...], initial_cursor: str = "") -> None: + super().__init__() + self.pages: Final = pages + self.cursors: Final[asyncio.Queue[str]] = asyncio.Queue() + self.page_by_cursor: Final = MappingProxyType( + { + initial_cursor: 0, + **{page[-1].selection_key: index + 1 for index, page in enumerate(pages[:-1])}, + } + ) + + async def lens_sample(self, parameters: LensSampleParams) -> Sequence[ExecutionRow]: + await self.cursors.put(parameters.after) + index: Final = self.page_by_cursor[parameters.after] + return self.pages[index] + + +class SignalDatabase: + def __init__( + self, + config: SignalConfig | None, + *, + stored_rows: tuple[StoredTraceSignal, ...] = (), + claim_result: bool = True, + ) -> None: + self.config: Final = config + self.stored_rows: Final = stored_rows + self.claim_result: Final = claim_result + self.calls: Final[asyncio.Queue[str]] = asyncio.Queue() + self.claims: Final[asyncio.Queue[str]] = asyncio.Queue() + self.claim_args: Final[asyncio.Queue[tuple[object, ...]]] = asyncio.Queue() + self.saved: Final[asyncio.Queue[tuple[object, ...]]] = asyncio.Queue() + + async def query_raw(self, query: str, *args: object) -> object: + if '"LiteLLM_LensSignalConfig"' in query: + return () if self.config is None else (Row(data=self.config.model_dump(mode="json")),) + if query.startswith("SELECT jsonb_build_object"): + payload: Final = args[0] + assert isinstance(payload, str) + requested: Final = TypeAdapter(tuple[TraceIdentity, ...]).validate_json(payload) + identities: Final = tuple((trace.trace_id, trace.trace_ref) for trace in requested) + return tuple( + Row(data=stored.model_dump(mode="json")) + for stored in self.stored_rows + if (stored.trace_id, stored.trace_ref) in identities + ) + if query.startswith('INSERT INTO "LiteLLM_LensTraceSignal"'): + await self.claim_args.put(args) + if not self.claim_result: + return () + trace_id: Final = args[0] + assert isinstance(trace_id, str) + await self.claims.put(trace_id) + return (Row(data={"trace_id": trace_id}),) + raise AssertionError(f"Unexpected query: {query}") + + async def execute_raw(self, query: str, *args: object) -> int: + await self.saved.put(args) + return 1 + + @asynccontextmanager + async def transaction(self) -> AsyncGenerator[Database, None]: + yield self + + +def saved_result(args: tuple[object, ...]) -> SignalData: + payload: Final = args[1] + assert isinstance(payload, str) + return SignalData.model_validate_json(payload) + + +@pytest.mark.asyncio +async def test_signal_repository_reads_defaults_and_saves_the_global_config() -> None: + database: Final = SignalDatabase(None) + repository: Final = SignalRepository(database) + updated: Final = SignalConfig(model="decision", threshold=0.7) + + assert await repository.get_config() == SignalConfig() + await repository.save_config(updated) + + saved: Final = await database.saved.get() + assert saved[0] == "global" + assert isinstance(saved[1], str) + assert SignalConfig.model_validate_json(saved[1]) == updated + + +@pytest.mark.asyncio +async def test_signal_repository_reads_rows_and_reports_a_lost_claim() -> None: + config: Final = SignalConfig(model="decision") + row: Final = stored_trace(config.key()) + database: Final = SignalDatabase(config, stored_rows=(row,), claim_result=False) + repository: Final = SignalRepository(database) + + assert await repository.traces(()) == () + assert await repository.traces((TraceIdentity(trace_id="trace"),)) == (row,) + assert not await repository.claim(execution("trace"), config, NOW + timedelta(minutes=5), NOW) + + +def test_signal_config_hashes_questions_but_not_threshold_or_display_name() -> None: + config: Final = SignalConfig(model="decision") + different_threshold: Final = config.model_copy(update={"threshold": 0.9}) + renamed: Final = config.model_copy( + update={ + "signals": ( + config.signals[0].model_copy(update={"name": "Frustration"}), + *config.signals[1:], + ) + } + ) + changed_question: Final = config.model_copy( + update={ + "signals": ( + config.signals[0].model_copy(update={"question": "Does this user sound upset?"}), + *config.signals[1:], + ) + } + ) + + assert config.key() == different_threshold.key() == renamed.key() + assert config.key() != changed_question.key() + assert DEFAULT_SIGNALS == config.signals + + +def test_signal_config_rejects_duplicate_ids_and_non_finite_thresholds() -> None: + duplicate: Final = Signal(id="same", name="First", question="Question one") + with pytest.raises(ValidationError): + SignalConfig(signals=(duplicate, duplicate)) + with pytest.raises(ValidationError): + SignalConfig(threshold=float("nan")) + + +@pytest.mark.parametrize( + "stored,trace_count,expected", + ( + (None, 1, True), + (stored_trace("old"), 1, True), + ( + stored_trace(CURRENT_CONFIG_KEY, span_count=1, classified_at=NOW - timedelta(minutes=6)), + 2, + True, + ), + ( + stored_trace( + CURRENT_CONFIG_KEY, + status="failed", + classified_at=(NOW - timedelta(minutes=31)).replace(tzinfo=None), + ), + 1, + True, + ), + (stored_trace(CURRENT_CONFIG_KEY), 1, False), + ( + stored_trace(CURRENT_CONFIG_KEY, span_count=1, classified_at=NOW - timedelta(minutes=2)), + 2, + False, + ), + ( + stored_trace("old", claimed_until=NOW + timedelta(minutes=1)), + 1, + False, + ), + ( + stored_trace( + CURRENT_CONFIG_KEY, + status="pending", + claimed_until=(NOW + timedelta(minutes=1)).replace(tzinfo=None), + classified_at=None, + ), + 1, + False, + ), + ( + stored_trace( + CURRENT_CONFIG_KEY, + status="pending", + claimed_until=(NOW - timedelta(minutes=1)).replace(tzinfo=None), + classified_at=None, + ), + 1, + True, + ), + (stored_trace(CURRENT_CONFIG_KEY, span_count=2), 1, False), + (stored_trace(CURRENT_CONFIG_KEY, span_count=1, classified_at=None), 2, False), + ), +) +def test_candidate_selection_respects_config_span_age_failure_age_and_claims( + stored: StoredTraceSignal | None, trace_count: int, expected: bool +) -> None: + config: Final = SignalConfig(model="decision") + assert candidate(execution("trace", trace_count), stored, config.key(), NOW) is expected + + +@pytest.mark.asyncio +async def test_classifier_sends_noul_questions_and_keeps_every_signal_score() -> None: + config: Final = SignalConfig(model="decision") + run: Final = execution("trace") + storage: Final = SignalStorage(parts=(part("agent", "user asks for a result"),)) + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + assert model == "decision" + assert state == { + "task": SIGNAL_TASK, + "steps": ({"kind": "agent", "name": "agent", "content": "user asks for a result"},), + } + assert _STORED_DATA.validate_json(json.dumps(state)) == { + "task": SIGNAL_TASK, + "steps": [{"kind": "agent", "name": "agent", "content": "user asks for a result"}], + } + expected_questions: Final = { + signal.id: {"type": "noul", "instructions": signal.question} for signal in config.signals + } + assert questions == expected_questions + assert _STORED_DATA.validate_json(json.dumps(questions)) == expected_questions + assert timeout == 60 + assert metadata == {"tags": ["litellm-lens-signals"]} + return DecisionsResponse( + answers={ + "user_frustration": DecisionsNoulAnswer(type="noul", noul=0.9), + "missing_capability": DecisionsNoulAnswer(type="noul", noul=0.6), + "repeated_request": DecisionsNoulAnswer(type="noul", noul=0.2), + "unknown": DecisionsNoulAnswer(type="noul", noul=1.0), + } + ) + + attempt: Final = await SignalClassifier(SourceReader(storage), decide, lambda: NOW).classify( + Scope(all_teams=True), run, config + ) + + assert attempt == SignalAttempt( + status="classified", + scores={"user_frustration": 0.9, "missing_capability": 0.6, "repeated_request": 0.2}, + model="decision", + ) + + +@pytest.mark.asyncio +async def test_missing_noul_answer_fails_while_unknown_and_non_noul_answers_are_ignored() -> None: + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return { + "answers": { + "user_frustration": {"type": "noul", "noul": 0.9}, + "missing_capability": {"type": "choice", "choice": "yes"}, + "unknown": {"type": "noul", "noul": 1.0}, + } + } + + attempt: Final = await SignalClassifier( + SourceReader(SignalStorage(parts=(part("agent", "content"),))), + decide, + lambda: NOW, + ).classify(Scope(all_teams=True), execution("trace"), SignalConfig(model="decision")) + + assert attempt.status == "failed" + assert attempt.scores == {"user_frustration": 0.9} + assert attempt.error == "Decisions response omitted a configured noul answer" + + +@pytest.mark.asyncio +async def test_classifier_turns_decisions_errors_into_failed_attempts() -> None: + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + raise RuntimeError("decisions unavailable") + + attempt: Final = await SignalClassifier( + SourceReader(SignalStorage(parts=(part("agent", "content"),))), + decide, + lambda: NOW, + ).classify(Scope(all_teams=True), execution("trace"), SignalConfig(model="decision")) + + assert attempt == SignalAttempt(status="failed", model="decision", error="decisions unavailable") + + +def test_signal_flags_use_current_threshold_and_current_display_name() -> None: + config: Final = SignalConfig(model="decision") + row: Final = stored_trace( + config.key(), + scores={"user_frustration": 0.91, "missing_capability": 0.67, "repeated_request": 0.49}, + ) + high_threshold: Final = config.model_copy( + update={ + "threshold": 0.9, + "signals": ( + config.signals[0].model_copy(update={"name": "Frustrated user"}), + *config.signals[1:], + ), + } + ) + trace: Final = TraceIdentity(trace_id="trace") + lower: Final = trace_signals(trace, row, config) + higher: Final = trace_signals(trace, row, high_threshold) + + assert config.key() == high_threshold.key() + assert tuple((flag.signal_id, flag.score) for flag in lower.flags) == ( + ("user_frustration", 0.91), + ("missing_capability", 0.67), + ) + assert tuple((flag.signal_id, flag.name, flag.score) for flag in higher.flags) == ( + ("user_frustration", "Frustrated user", 0.91), + ) + assert not candidate(execution("trace"), row, high_threshold.key(), NOW) + + +def test_signal_flags_report_stored_errors_even_when_the_status_is_classified() -> None: + config: Final = SignalConfig(model="decision") + row: Final = stored_trace(config.key(), status="classified", error="classification failed") + + result: Final = trace_signals(TraceIdentity(trace_id="trace"), row, config) + + assert result.status == "failed" + assert result.model == "decision" + assert result.classified_at == row.classified_at + + +@pytest.mark.asyncio +async def test_signal_state_caps_content_to_head_and_tail_with_omitted_step() -> None: + parts: Final = tuple(part(str(index), chr(97 + index) * 2000) for index in range(30)) + state: Final = await signal_state( + SourceReader(SignalStorage(parts=parts)), + Scope(all_teams=True), + execution("trace"), + ) + steps_value: Final = state["steps"] + assert isinstance(steps_value, tuple) + steps: Final = _SIGNAL_STEPS.validate_python(steps_value) + head: Final = steps[:8] + marker: Final = steps[8] + tail: Final = steps[9:] + + assert state["task"] == SIGNAL_TASK + assert sum(len(step.content) for step in head) == 15000 + assert sum(len(step.content) for step in tail) == 25000 + assert head[0].content == "a" * 2000 + assert head[-1].content == "h" * 1000 + assert marker == SignalStep(kind="omitted", name="", content="9 steps omitted") + assert tail[0].content == "r" * 1000 + assert tail[-1].content == "~" * 2000 + + +@pytest.mark.asyncio +async def test_signal_state_limits_content_pages_and_part_sizes() -> None: + pages: Final = tuple( + tuple(part(f"page-{page}-{index}", "x" * 2501 if index == 0 else "x") for index in range(39)) + + (part(str(page + 1), "x"),) + for page in range(4) + ) + state: Final = await signal_state( + SourceReader(PagedSignalStorage(pages)), + Scope(all_teams=True), + execution("trace"), + ) + steps: Final = _SIGNAL_STEPS.validate_python(state["steps"]) + + assert len(steps) == 120 + assert steps[0].content.startswith("x" * 800) + assert "[... 501 characters omitted ...]" in steps[0].content + assert steps[0].content.endswith("x" * 1200) + assert steps[-1].name == "3" + assert all(not step.name.startswith("page-3-") for step in steps) + + small_state: Final = await signal_state( + SourceReader(SignalStorage(parts=(part("small", "ok"),))), + Scope(all_teams=True), + execution("trace"), + ) + small_steps: Final = _SIGNAL_STEPS.validate_python(small_state["steps"]) + assert small_steps == (SignalStep(kind="agent", name="small", content="ok"),) + + +@pytest.mark.asyncio +async def test_signal_state_part_excerpt_preserves_the_output_tail() -> None: + content: Final = "I" * 5000 + "OUTPUT: refused" + state: Final = await signal_state( + SourceReader(SignalStorage(parts=(part("result", content),))), + Scope(all_teams=True), + execution("trace"), + ) + steps: Final = _SIGNAL_STEPS.validate_python(state["steps"]) + excerpt: Final = steps[0].content + marker: Final = "\n[... 3015 characters omitted ...]\n" + + assert marker in excerpt + assert excerpt.endswith("OUTPUT: refused") + assert len(excerpt) == 800 + len(marker) + 1200 + + +@pytest.mark.asyncio +async def test_signal_tick_classifies_at_most_50_traces_and_persists_scores() -> None: + config: Final = SignalConfig(model="decision") + executions: Final = tuple( + ExecutionRow( + source="traces", + trace_id=f"trace-{index}", + team_id="", + name=f"trace-{index}", + start_time="", + span_count=1, + root_seen=1, + eligible=60, + selected=60, + selection_key=f"cursor-{index}", + ) + for index in range(60) + ) + storage: Final = SignalStorage(executions=executions) + database: Final = SignalDatabase(config) + repository: Final = SignalRepository(database) + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + steps: Final = TypeAdapter(tuple[SignalStep, ...]).validate_python(state["steps"]) + await database.calls.put(steps[0].name) + return { + "answers": { + "user_frustration": {"type": "noul", "noul": 0.9}, + "missing_capability": {"type": "noul", "noul": 0.6}, + "repeated_request": {"type": "noul", "noul": 0.2}, + } + } + + await run_signal_tick(storage, repository, decide, lambda: NOW) + classified: Final = tuple(database.saved.get_nowait() for _ in range(database.saved.qsize())) + traces: Final = tuple(database.calls.get_nowait() for _ in range(database.calls.qsize())) + saved_data: Final = tuple(saved_result(args) for args in classified) + + assert len(classified) == 50 + assert frozenset(traces) == frozenset(f"trace-{index}" for index in range(50)) + assert ( + saved_data + == ( + SignalData( + status="classified", + scores={ + "user_frustration": 0.9, + "missing_capability": 0.6, + "repeated_request": 0.2, + }, + model="decision", + error="", + ), + ) + * 50 + ) + + +@pytest.mark.asyncio +async def test_signal_tick_claims_with_worker_start_time_and_skips_lost_claims() -> None: + config: Final = SignalConfig(model="decision") + executions: Final = tuple( + ExecutionRow( + source="traces", + trace_id=f"trace-{index}", + team_id="", + name=f"trace-{index}", + start_time="", + span_count=1, + root_seen=1, + eligible=2, + selected=2, + selection_key=f"cursor-{index}", + ) + for index in range(2) + ) + database: Final = SignalDatabase(config, claim_result=False) + repository: Final = SignalRepository(database) + + class AdvancingClock: + def __init__(self) -> None: + self.values: Final = tuple(NOW + timedelta(minutes=index) for index in range(3)) + self.index: int = 0 + + def __call__(self) -> datetime: + value: Final = self.values[self.index] + self.index += 1 + return value + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return {"answers": {}} + + await run_signal_tick(SignalStorage(executions=executions), repository, decide, AdvancingClock()) + + claims: Final = tuple(database.claim_args.get_nowait() for _ in range(database.claim_args.qsize())) + + def claim_times(args: tuple[object, ...]) -> tuple[datetime, datetime]: + claimed_until: Final = args[4] + claimed_at: Final = args[6] + assert isinstance(claimed_until, datetime) + assert isinstance(claimed_at, datetime) + return claimed_until, claimed_at + + times: Final = tuple(claim_times(claim) for claim in claims) + assert database.calls.empty() + assert database.saved.empty() + assert all(claimed_until == claimed_at + SIGNAL_CLAIM_LEASE for claimed_until, claimed_at in times) + assert all(claimed_at != NOW for _, claimed_at in times) + + +@pytest.mark.asyncio +async def test_signal_tick_resumes_after_ten_pages_and_resets_after_a_short_page() -> None: + config: Final = SignalConfig(model="decision") + + def sample_page(page: int) -> tuple[ExecutionRow, ...]: + return tuple( + ExecutionRow( + source="traces", + trace_id=f"trace-{page}-{index}", + team_id="", + name=f"trace-{page}-{index}", + start_time="", + span_count=1, + root_seen=1, + eligible=2500, + selected=2500, + selection_key=f"page-{page}-{index}", + ) + for index in range(100) + ) + + pages: Final = tuple(sample_page(page) for page in range(25)) + all_rows: Final = tuple(chain.from_iterable(pages)) + stored_rows: Final = tuple(stored_trace(CURRENT_CONFIG_KEY, trace_id=row.trace_id) for row in all_rows) + storage: Final = PagedSampleStorage(pages) + database: Final = SignalDatabase(config, stored_rows=stored_rows) + repository: Final = SignalRepository(database) + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return {"answers": {}} + + first_cursor: Final = await run_signal_tick(storage, repository, decide, lambda: NOW) + first_calls: Final = tuple(storage.cursors.get_nowait() for _ in range(storage.cursors.qsize())) + second_cursor: Final = await run_signal_tick( + storage, + repository, + decide, + lambda: NOW, + cursor=first_cursor, + ) + second_calls: Final = tuple(storage.cursors.get_nowait() for _ in range(storage.cursors.qsize())) + + assert len(first_calls) == SIGNAL_MAX_SCAN_PAGES + assert first_cursor + assert len(second_calls) == SIGNAL_MAX_SCAN_PAGES + assert second_calls[0] == first_cursor + assert second_cursor + + short_storage: Final = PagedSampleStorage((pages[0][:50],)) + short_database: Final = SignalDatabase(config, stored_rows=stored_rows[:50]) + short_cursor: Final = await run_signal_tick( + short_storage, + SignalRepository(short_database), + decide, + lambda: NOW, + ) + assert short_cursor == "" + + +@pytest.mark.asyncio +async def test_signal_tick_resumes_a_partially_consumed_page() -> None: + config: Final = SignalConfig(model="decision") + page: Final = tuple( + ExecutionRow( + source="traces", + trace_id=f"trace-{index}", + team_id="", + name=f"trace-{index}", + start_time="", + span_count=1, + root_seen=1, + eligible=100, + selected=100, + selection_key=f"cursor-{index}", + ) + for index in range(100) + ) + initial_rows: Final = tuple(stored_trace(CURRENT_CONFIG_KEY, trace_id=f"trace-{index}") for index in range(20)) + resume_cursor: Final = "resume-page" + storage: Final = PagedSampleStorage((page, ()), initial_cursor=resume_cursor) + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return { + "answers": { + "user_frustration": {"type": "noul", "noul": 0.9}, + "missing_capability": {"type": "noul", "noul": 0.6}, + "repeated_request": {"type": "noul", "noul": 0.2}, + } + } + + first_database: Final = SignalDatabase(config, stored_rows=initial_rows) + first_cursor: Final = await run_signal_tick( + storage, + SignalRepository(first_database), + decide, + lambda: NOW, + cursor=resume_cursor, + ) + first_claims: Final = tuple(first_database.claims.get_nowait() for _ in range(first_database.claims.qsize())) + + classified_first_rows: Final = tuple( + stored_trace(CURRENT_CONFIG_KEY, trace_id=trace_id) for trace_id in first_claims + ) + second_database: Final = SignalDatabase(config, stored_rows=(*initial_rows, *classified_first_rows)) + second_cursor: Final = await run_signal_tick( + storage, + SignalRepository(second_database), + decide, + lambda: NOW, + cursor=first_cursor, + ) + second_claims: Final = tuple(second_database.claims.get_nowait() for _ in range(second_database.claims.qsize())) + sample_cursors: Final = tuple(storage.cursors.get_nowait() for _ in range(storage.cursors.qsize())) + expected_eligible: Final = frozenset(f"trace-{index}" for index in range(20, 100)) + + assert first_cursor == resume_cursor + assert second_cursor == "" + assert len(first_claims) == 50 + assert len(second_claims) == 30 + assert frozenset(first_claims).isdisjoint(second_claims) + assert frozenset(first_claims) | frozenset(second_claims) == expected_eligible + assert sample_cursors == (resume_cursor, resume_cursor, page[-1].selection_key) + + +@pytest.mark.asyncio +async def test_signal_tick_skips_claims_and_writes_when_router_is_not_ready() -> None: + config: Final = SignalConfig(model="decision") + storage: Final = SignalStorage( + executions=( + ExecutionRow( + source="traces", + trace_id="trace", + team_id="", + name="trace", + start_time="", + span_count=1, + root_seen=1, + eligible=1, + selected=1, + selection_key="cursor", + ), + ) + ) + database: Final = SignalDatabase(config) + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return {"answers": {}} + + await run_signal_tick( + storage, + SignalRepository(database), + decide, + lambda: NOW, + router_ready=lambda: False, + ) + + assert database.claims.empty() + assert database.saved.empty() + + +@pytest.mark.asyncio +async def test_signal_tick_skips_missing_dependencies_and_disabled_configs() -> None: + storage: Final = SignalStorage() + + await run_signal_tick(storage, None, None, lambda: NOW) + + database: Final = SignalDatabase(SignalConfig()) + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + raise AssertionError("disabled signal config should not call Decisions") + + await run_signal_tick(storage, SignalRepository(database), decide, lambda: NOW) + assert database.claims.empty() + assert database.saved.empty() + + +class FailingStoreDatabase(SignalDatabase): + async def execute_raw(self, query: str, *args: object) -> int: + raise RuntimeError("store unavailable") + + +@pytest.mark.asyncio +async def test_signal_tick_continues_when_storing_a_result_fails() -> None: + config: Final = SignalConfig(model="decision") + execution_row: Final = ExecutionRow( + source="traces", + trace_id="trace", + team_id="", + name="trace", + start_time="", + span_count=1, + root_seen=1, + eligible=1, + selected=1, + selection_key="cursor", + ) + database: Final = FailingStoreDatabase(config) + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return { + "answers": { + "user_frustration": {"type": "noul", "noul": 0.9}, + "missing_capability": {"type": "noul", "noul": 0.6}, + "repeated_request": {"type": "noul", "noul": 0.2}, + } + } + + await run_signal_tick( + SignalStorage(executions=(execution_row,)), + SignalRepository(database), + decide, + lambda: NOW, + ) + + assert await database.claims.get() == "trace" + assert database.saved.empty() + + +class FailingSignalRepository: + def __init__(self) -> None: + self.started: Final = asyncio.Event() + + async def get_config(self) -> SignalConfig: + self.started.set() + await asyncio.sleep(0) + raise RuntimeError("tick failed") + + +@pytest.mark.asyncio +async def test_signal_loop_continues_after_a_tick_error() -> None: + repository: Final = FailingSignalRepository() + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return {"answers": {}} + + task: Final = asyncio.create_task(run_signal_loop(SignalStorage(), repository, decide, lambda: NOW)) + await repository.started.wait() + await asyncio.sleep(0) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + +@pytest.mark.asyncio +async def test_proxy_signal_call_resolves_the_current_router(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy import proxy_server + + async def first_decisions( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return "first" + + async def second_decisions( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return "second" + + async def call_current_router() -> object: + return await proxy_server._call_current_lens_signal_router( + model="decision", + state={"task": "task"}, + questions={}, + timeout=60, + metadata={"tags": ["test"]}, + ) + + monkeypatch.setattr(proxy_server, "llm_router", SimpleNamespace(adecisions=first_decisions)) + assert await call_current_router() == "first" + + monkeypatch.setattr(proxy_server, "llm_router", SimpleNamespace(adecisions=second_decisions)) + assert await call_current_router() == "second" + + monkeypatch.setattr(proxy_server, "llm_router", None) + with pytest.raises(RuntimeError, match="router is not initialized"): + await call_current_router() diff --git a/tests/unit/proxy/lens/test_sources.py b/tests/unit/proxy/lens/test_sources.py index 8512db04eff..ad1f90f97ce 100644 --- a/tests/unit/proxy/lens/test_sources.py +++ b/tests/unit/proxy/lens/test_sources.py @@ -10,8 +10,10 @@ from litellm.proxy.lens.sources import SourceReader, execution_id, parse_executi from litellm.rust_bridge.trace.generated.models import ( ActivityAvailability, AgentRow, + CountRow, ExecutionRow, LensContentParams, + LensEvidenceParams, PartRow, ) from tests.unit.proxy.lens.test_agent_workspace import python_data @@ -183,9 +185,17 @@ async def test_recorded_times_survive_source_catalog_reads_search_and_python( class ContentStorage: async def lens_content(self, parameters: LensContentParams) -> tuple[PartRow, ...]: - assert parameters.source == source and parameters.record_team == "team" + assert ( + parameters.source == source + and parameters.record_team == "team" + and parameters.start_time == run.start_time + ) return rows + async def lens_evidence(self, parameters: LensEvidenceParams) -> tuple[CountRow, ...]: + assert parameters.start_time == run.start_time + return (CountRow(count=1),) + reader: Final = SourceReader(ContentStorage()) async def read(identity: str, cursor: str, offset: int) -> ExecutionContent: @@ -217,6 +227,11 @@ async def test_recorded_times_survive_source_catalog_reads_search_and_python( assert computed.sessions[0].parts == expected assert min(computed.sessions[0].parts, key=lambda part: part.start_time).span_id == rows[-1].span_id assert await workspace.valid(Evidence(execution_id=run.id, span_id=rows[0].span_id, quote=rows[0].content)) + assert await reader.verify_evidence( + Scope(team_id="team"), + run, + Evidence(execution_id=run.id, span_id=rows[0].span_id, quote=rows[0].content), + ) @pytest.mark.asyncio 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/proxy/spend_tracking/test_baseline_accounting.py b/tests/unit/proxy/spend_tracking/test_baseline_accounting.py index a188d65502d..368704e7a75 100644 --- a/tests/unit/proxy/spend_tracking/test_baseline_accounting.py +++ b/tests/unit/proxy/spend_tracking/test_baseline_accounting.py @@ -111,7 +111,13 @@ def test_prefix_match_expiry_and_usage_pricing_fields(ttl: int) -> None: assert warm.usage.prompt_tokens_details.cached_tokens == 6000 assert cold.usage.prompt_tokens_details.cached_tokens == 0 assert cold.usage.prompt_tokens_details.cache_creation_tokens == 6000 - unaffected: Final = {"prompt_tokens", "total_tokens", "prompt_tokens_details", "cache_read_input_tokens", "cache_creation_input_tokens"} + unaffected: Final = { + "prompt_tokens", + "total_tokens", + "prompt_tokens_details", + "cache_read_input_tokens", + "cache_creation_input_tokens", + } assert warm.usage.model_dump(exclude=unaffected) == first.usage.model_dump(exclude=unaffected) assert cold.usage.model_dump(exclude=unaffected) == first.usage.model_dump(exclude=unaffected) @@ -124,7 +130,10 @@ def test_growth_lookback_and_mixed_ttl_keep_distinct_read_write_buckets(warm_tai second: Final = _replay(first, _observation("second", 10001.0, plan=grown))[-1] assert second.reason == "history_unavailable" history: Final = BaselineHistory( - first_at=1.0, last_at=10000.0, equivalent=False, uncertain_before=1.0, + first_at=1.0, + last_at=10000.0, + equivalent=False, + uncertain_before=1.0, entries=(CacheEntry("tail:300", "tail", 7000, 300, 10000.0, 10300.0),) if warm_tail else (), ) _, estimates = advance_baseline_history(history, (_observation("mixed", 10001.0, plan=grown),)) @@ -134,8 +143,12 @@ def test_growth_lookback_and_mixed_ttl_keep_distinct_read_write_buckets(warm_tai # Anthropic billing locations: B is the highest 1h breakpoint AFTER the highest hit A. # https://platform.claude.com/docs/en/build-with-claude/prompt-caching#mixing-different-ttls (2026-09-15) assert usage.prompt_tokens_details.cached_tokens == (7000 if warm_tail else 0) - assert usage.prompt_tokens_details.cache_creation_token_details.ephemeral_1h_input_tokens == (0 if warm_tail else 6500) - assert usage.prompt_tokens_details.cache_creation_token_details.ephemeral_5m_input_tokens == (0 if warm_tail else 500) + assert usage.prompt_tokens_details.cache_creation_token_details.ephemeral_1h_input_tokens == ( + 0 if warm_tail else 6500 + ) + assert usage.prompt_tokens_details.cache_creation_token_details.ephemeral_5m_input_tokens == ( + 0 if warm_tail else 500 + ) @pytest.mark.parametrize("change", ["prefix", "ttl", "unavailable", "failed", "response_cache"]) @@ -197,3 +210,62 @@ def test_modeled_read_cannot_recharge_the_original_private_write_count() -> None } input_cost, output_cost = cost_per_token("claude-opus-5", warm.usage, model_info=prices) assert input_cost + output_cost == pytest.approx((200 * 1e-6 + 6000 * 1e-7 + 30 * 2e-6) * 2.0 * 1.1) + + +def test_mixed_lifetime_lookback_preserves_a_compatible_native_hit() -> None: + marker: Final = _marker("prefix", 3600, 6000) + history: Final = BaselineHistory( + first_at=1.0, + last_at=10000.0, + equivalent=False, + uncertain_before=1.0, + entries=(CacheEntry(marker.fingerprint, marker.content_fingerprint, 6000, 3600, 10000.0, 13600.0),), + ) + plan: Final = CountedPromptCachePlan( + 7100, (_marker("grown", 3600, 6500, ("prefix",)), _marker("tail", 300, 7000, ("prefix",))) + ) + _, estimates = advance_baseline_history(history, (_observation("next", 10001.0, plan=plan),)) + usage: Final = estimates[0].usage + assert usage is not None, estimates[0].reason + assert usage.prompt_tokens_details.cached_tokens == 6000 + assert usage.prompt_tokens_details.text_tokens == 100 + assert usage.prompt_tokens_details.cache_creation_token_details == CacheCreationTokenDetails( + ephemeral_5m_input_tokens=500, ephemeral_1h_input_tokens=500 + ) + + +def test_short_lifetime_hit_cannot_seed_an_unpaid_long_lifetime_entry() -> None: + first: Final = _observation("initial", plan=CountedPromptCachePlan(6200, (_marker("5", 3600, 6000),))) + short: Final = _observation("short", 13700.0, plan=CountedPromptCachePlan(6200, (_marker("3", 300, 4600),))) + mixed: Final = _observation( + "mixed", + 13710.0, + plan=CountedPromptCachePlan(6200, (_marker("3", 3600, 4600), _marker("4", 300, 5500, ("3",)))), + ) + later: Final = _observation("later", 14710.0, plan=CountedPromptCachePlan(6200, (_marker("3", 3600, 4600),))) + _, _, upgrade, after_expiry = _replay(first, short, mixed, later) + assert upgrade.usage is not None and after_expiry.usage is not None + assert upgrade.usage.prompt_tokens_details.cached_tokens == 4600 + assert upgrade.usage.prompt_tokens_details.cache_creation_token_details.ephemeral_1h_input_tokens == 0 + assert after_expiry.usage.prompt_tokens_details.cached_tokens == 0 + assert after_expiry.usage.prompt_tokens_details.cache_creation_token_details.ephemeral_1h_input_tokens == 4600 + + +@pytest.mark.parametrize("writes", (0, 50, None)) +def test_cache_creation_split_is_optional_only_without_writes(writes: int | None) -> None: + usage: Final = Usage( + prompt_tokens=100, + completion_tokens=10, + total_tokens=110, + prompt_tokens_details=PromptTokensDetailsWrapper( + text_tokens=100 - (writes or 0), + cached_tokens=0, + cache_creation_tokens=writes, + ), + ) + observed: Final = _observation("no-split", usage=usage, plan=CountedPromptCachePlan(100, ())) + restored: Final = BaselineObservation.model_validate_json(observed.model_dump_json()) + estimate: Final = _replay(restored)[0] + assert (estimate.usage is not None) is (writes == 0) + if estimate.usage is not None: + assert estimate.usage.prompt_tokens == usage.prompt_tokens diff --git a/tests/unit/router_utils/test_baseline_request.py b/tests/unit/router_utils/test_baseline_request.py new file mode 100644 index 00000000000..37d21597752 --- /dev/null +++ b/tests/unit/router_utils/test_baseline_request.py @@ -0,0 +1,59 @@ +from typing import Final + +from litellm.router_utils.baseline_request import baseline_request, capture_baseline_parameters + + +def test_baseline_snapshot_owns_nested_caller_settings_and_overrides_routed_settings() -> None: + reasoning: Final = {"effort": "medium"} + snapshot: Final = capture_baseline_parameters({"reasoning": reasoning, "verbosity": "low"}) + assert snapshot is not None + reasoning["effort"] = "high" + projected: Final = baseline_request( + {"messages": [{"role": "user", "content": "hello"}], "reasoning": reasoning, "verbosity": "high"}, + snapshot, + {"verbosity": "medium"}, + ) + assert projected == { + "messages": [{"role": "user", "content": "hello"}], + "reasoning": {"effort": "medium"}, + "verbosity": "low", + } + + +def test_oversized_snapshot_fails_closed_before_json_validation() -> None: + assert capture_baseline_parameters({"output_config": {"format": "x" * 5_000_000}}) is None + + +def test_snapshot_retains_extra_body_settings_but_no_credentials() -> None: + assert capture_baseline_parameters({"api_key": "private", "extra_body": {"verbosity": "low"}}) == { + "extra_body": {"verbosity": "low"} + } + + +def test_chat_projection_applies_extra_body_after_top_level_parameters() -> None: + snapshot: Final = capture_baseline_parameters({"verbosity": "high", "extra_body": {"verbosity": "low"}}) + assert snapshot is not None + assert baseline_request({}, snapshot, {}) == {"verbosity": "low"} + + +def test_baseline_projection_keeps_caller_tools_and_request_parameter_precedence() -> None: + from litellm.router import Router + + deployment: Final = { + "tools": [{"type": "function", "function": {"name": "configured"}}], + "tool_choice": "required", + "max_tokens": 64, + } + caller: Final = { + "tools": [{"type": "function", "function": {"name": "caller"}}], + "tool_choice": "auto", + "max_tokens": 128, + } + actual_request: Final = dict(caller) + Router._merge_tools_from_deployment({"litellm_params": deployment}, actual_request) + snapshot: Final = capture_baseline_parameters(caller) + assert snapshot is not None + assert baseline_request({"tools": [{"name": "routed-only"}], "max_tokens": 4}, snapshot, deployment) == { + **deployment, + **actual_request, + } diff --git a/tests/unit/rust_bridge/trace/test_queries.py b/tests/unit/rust_bridge/trace/test_queries.py index 3c0556b1d97..9fd2964af20 100644 --- a/tests/unit/rust_bridge/trace/test_queries.py +++ b/tests/unit/rust_bridge/trace/test_queries.py @@ -18,6 +18,7 @@ def test_named_query_rejects_offsets_outside_the_native_integer_range(offset: in "source": "traces", "id": "trace", "record_team": "team", + "start_time": "", "trace_ref": "ref", "cursor": "", "offset": offset, @@ -34,6 +35,7 @@ def test_named_query_rejects_parameters_for_a_different_query() -> None: source="traces", id="trace", record_team="team", + start_time="", trace_ref="ref", cursor="", offset=0, 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"}, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailFormField.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailFormField.tsx index 0afd9aab0cf..ef0c8d074a7 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailFormField.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailFormField.tsx @@ -1,11 +1,17 @@ "use client"; import { CircleHelp } from "lucide-react"; -import React, { useId } from "react"; +import React, { useEffect, useId } from "react"; import { useController, type Control, type ControllerRenderProps, type RegisterOptions } from "react-hook-form"; import { Field, FieldDescription, FieldError, FieldLabel } from "@/components/ui/field"; import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; import { Tooltip, TooltipContent, TooltipTrigger } from "@/components/ui/tooltip"; +import { + getLoggingOnlyScopeOptions, + modeIncludesLoggingOnly, + normalizeLoggingOnlyScopeChoice, + type LoggingOnlyScopeChoice, +} from "./guardrail_info_helpers"; export interface GuardrailCriterion { name: string; @@ -15,6 +21,7 @@ export interface GuardrailCriterion { export interface GuardrailFormValues extends Record { criteria?: GuardrailCriterion[]; + logging_only_scope_choice?: LoggingOnlyScopeChoice; } export type GuardrailFormControl = Control; export type GuardrailFieldRules = Pick, "validate">; @@ -38,6 +45,13 @@ export const asText = (value: unknown): string => { return ""; }; +const LOGGING_ONLY_SCOPE_CHOICES: ReadonlySet = new Set( + getLoggingOnlyScopeOptions(true).map(({ value }) => value), +); + +const isLoggingOnlyScopeChoice = (value: unknown): value is LoggingOnlyScopeChoice => + typeof value === "string" && LOGGING_ONLY_SCOPE_CHOICES.has(value); + export const asStringArray = (value: unknown): string[] => { if (Array.isArray(value)) return value.filter((entry): entry is string => typeof entry === "string"); if (typeof value === "string" && value !== "") return [value]; @@ -123,3 +137,55 @@ export const SkipMessageSelect: React.FC<{ control: GuardrailFieldControlProps } ); }; + +export const LoggingOnlyScopeSelect: React.FC<{ + control: GuardrailFieldControlProps; + directionalScopeSupported: boolean; +}> = ({ control, directionalScopeSupported }) => { + const { id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy } = control; + const items = getLoggingOnlyScopeOptions(directionalScopeSupported); + + useEffect(() => { + const currentChoice = isLoggingOnlyScopeChoice(value) ? value : "default"; + const choice = normalizeLoggingOnlyScopeChoice(currentChoice, directionalScopeSupported); + if (choice !== value) onChange(choice); + }, [value, directionalScopeSupported, onChange]); + + return ( + + ); +}; + +export const LoggingOnlyScopeField: React.FC<{ + control: GuardrailFormControl; + mode: unknown; + directionalScopeSupported: boolean; +}> = ({ control, mode, directionalScopeSupported }) => { + if (!modeIncludesLoggingOnly(mode)) return null; + + return ( + + {(fieldControl) => ( + + )} + + ); +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailModeDisplay.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailModeDisplay.tsx new file mode 100644 index 00000000000..e5fa1669a44 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailModeDisplay.tsx @@ -0,0 +1,43 @@ +import React from "react"; +import { Badge } from "@/components/ui/badge"; +import { Card } from "@/components/ui/card"; +import { formatGuardrailMode, formatLoggingOnlyScope, modeIncludesLoggingOnly } from "./guardrail_info_helpers"; + +type GuardrailModeParams = { + mode?: unknown; + default_on?: boolean; + logging_only_scope?: string | null; +}; + +export const GuardrailModeCard: React.FC<{ litellmParams: GuardrailModeParams }> = ({ litellmParams }) => ( + +

Mode

+
+

{formatGuardrailMode(litellmParams.mode) || "-"}

+ + {litellmParams.default_on ? "Default On" : "Default Off"} + +
+ {modeIncludesLoggingOnly(litellmParams.mode) && ( +
+

Logging only scope

+

{formatLoggingOnlyScope(litellmParams.logging_only_scope)}

+
+ )} +
+); + +export const GuardrailModeRows: React.FC<{ litellmParams: GuardrailModeParams }> = ({ litellmParams }) => ( + <> +
+

Mode

+
{formatGuardrailMode(litellmParams.mode) || "-"}
+
+ {modeIncludesLoggingOnly(litellmParams.mode) && ( +
+

Logging only scope

+
{formatLoggingOnlyScope(litellmParams.logging_only_scope)}
+
+ )} + +); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.characterization.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.characterization.test.tsx index aadc9ec0213..230be13ef33 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.characterization.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.characterization.test.tsx @@ -37,6 +37,7 @@ const uiSettings = { supported_entities: [], supported_actions: [], supported_modes: ["pre_call", "post_call"], + providers_without_directional_logging_only_scope: [], pii_entity_categories: [], }; @@ -105,6 +106,67 @@ describe("AddGuardrailForm create payload characterization", () => { expect(payload()).toMatchObject({ litellm_params: { mode: ["pre_call", "post_call"] } }); }); + it("sends the selected output logging-only scope", async () => { + vi.mocked(networking.getGuardrailUISettings).mockResolvedValue({ + ...uiSettings, + supported_modes: ["pre_call", "logging_only"], + }); + const user = userEvent.setup({ delay: null }); + renderForm(); + + await user.type(await screen.findByLabelText("Guardrail Name"), "my-bedrock"); + await pickProvider(user, "Bedrock Guardrail"); + await user.click(screen.getByLabelText("Mode")); + await user.click((await screen.findAllByText("logging_only")).at(-1) as HTMLElement); + await chooseSelectOption(user, await screen.findByLabelText("Logging only scope"), "Output only (response)"); + + await user.type(await screen.findByPlaceholderText("The guardrail id on Bedrock"), "gr-123"); + await user.click(screen.getByRole("button", { name: "Next" })); + await user.click(await screen.findByRole("button", { name: "Create Guardrail" })); + + await waitFor(() => expect(networking.createGuardrailCall).toHaveBeenCalledTimes(1)); + expect(payload()?.litellm_params.logging_only_scope).toBe("output"); + }); + + it("hides directional scope choices for providers that do not support them", async () => { + vi.mocked(networking.getGuardrailUISettings).mockResolvedValue({ + ...uiSettings, + supported_modes: ["pre_call", "logging_only"], + providers_without_directional_logging_only_scope: ["xecguard"], + }); + vi.mocked(networking.getGuardrailProviderSpecificParams).mockResolvedValue({ + ...providerParams, + xecguard: { ui_friendly_name: "XecGuard" }, + }); + const user = userEvent.setup({ delay: null }); + renderForm(); + + await pickProvider(user, "XecGuard"); + await user.click(screen.getByLabelText("Mode")); + await user.click((await screen.findAllByText("logging_only")).at(-1) as HTMLElement); + await user.click(await screen.findByLabelText("Logging only scope")); + + expect(screen.queryByRole("option", { name: "Input only (request)" })).not.toBeInTheDocument(); + expect(screen.queryByRole("option", { name: "Output only (response)" })).not.toBeInTheDocument(); + expect(screen.getByRole("option", { name: "Both (request and response)" })).toBeInTheDocument(); + }); + + it("hides logging-only scope and omits it from a pre-call payload", async () => { + const user = userEvent.setup({ delay: null }); + renderForm(); + + await user.type(await screen.findByLabelText("Guardrail Name"), "my-bedrock"); + await pickProvider(user, "Bedrock Guardrail"); + expect(screen.queryByLabelText("Logging only scope")).not.toBeInTheDocument(); + + await user.type(await screen.findByPlaceholderText("The guardrail id on Bedrock"), "gr-123"); + await user.click(screen.getByRole("button", { name: "Next" })); + await user.click(await screen.findByRole("button", { name: "Create Guardrail" })); + + await waitFor(() => expect(networking.createGuardrailCall).toHaveBeenCalledTimes(1)); + expect(payload()?.litellm_params).not.toHaveProperty("logging_only_scope"); + }); + it("blocks Next when the user deselects every mode", async () => { const user = userEvent.setup({ delay: null }); renderForm(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx index a02cae097a7..43a542289f6 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx @@ -1,5 +1,5 @@ import React, { useEffect, useMemo, useState } from "react"; -import { useForm, type UseFormReturn } from "react-hook-form"; +import { useForm, useWatch, type UseFormReturn } from "react-hook-form"; import { toast } from "@/lib/toast"; import { createGuardrailCall, @@ -10,18 +10,23 @@ import { import ContentFilterConfiguration from "./content_filter/ContentFilterConfiguration"; import { type CompetitorIntentConfig } from "./content_filter/CompetitorIntentConfiguration"; import { + choiceToLoggingOnlyScope, choiceToSkipSystemForCreate, choiceToSkipToolForCreate, getGuardrailLogo, getGuardrailProviders, getSupportedModesForProvider, guardrail_provider_map, + modeIncludesLoggingOnly, populateGuardrailProviderMap, populateGuardrailProviders, shouldRenderContentFilterConfigSettings, shouldRenderLLMJudgeFields, shouldRenderPIIConfigSettings, + supportsDirectionalLoggingOnlyScope, toModeArray, + type LoggingOnlyScope, + type LoggingOnlyScopeChoice, } from "./guardrail_info_helpers"; import { Logo } from "@/components/molecules/logo/Logo"; import { MultiSelect } from "@/components/shared/MultiSelect"; @@ -49,6 +54,7 @@ import { requiredRule, type GuardrailCriterion, type GuardrailFormValues, + LoggingOnlyScopeField, SkipMessageSelect, } from "./GuardrailFormField"; import GuardrailOptionalParams from "./guardrail_optional_params"; @@ -92,6 +98,7 @@ interface GuardrailSettings { supported_actions: string[]; supported_modes: string[]; supported_modes_by_provider?: Record; + providers_without_directional_logging_only_scope?: string[]; pii_entity_categories: Array<{ category: string; entities: string[]; @@ -162,6 +169,7 @@ type SkipMessageChoice = "inherit" | "yes" | "no"; const INITIAL_VALUES: GuardrailFormValues = { mode: "pre_call", default_on: false, + logging_only_scope_choice: "default", skip_system_message_choice: "inherit", skip_tool_message_choice: "inherit", }; @@ -201,6 +209,7 @@ interface ProviderParamsResponse { const AddGuardrailForm: React.FC = ({ visible, onClose, accessToken, onSuccess, preset }) => { const form = useForm({ defaultValues: INITIAL_VALUES }); + const watchedMode = useWatch({ control: form.control, name: "mode" }); const [loading, setLoading] = useState(false); const [selectedProvider, setSelectedProvider] = useState(null); const [guardrailSettings, setGuardrailSettings] = useState(null); @@ -236,6 +245,7 @@ const AddGuardrailForm: React.FC = ({ visible, onClose, a const providerValue = guardrail_provider_map[selectedProvider]; return (providerValue || "").toLowerCase() === "tool_permission"; }, [selectedProvider]); + const directionalScopeSupported = supportsDirectionalLoggingOnlyScope(guardrailSettings, selectedProvider); // Fetch guardrail UI settings + provider params on mount / accessToken change useEffect(() => { @@ -279,6 +289,7 @@ const AddGuardrailForm: React.FC = ({ visible, onClose, a guardrail_name: preset.guardrailNameSuggestion, mode: preset.mode, default_on: preset.defaultOn, + logging_only_scope_choice: "default", skip_system_message_choice: "inherit", skip_tool_message_choice: "inherit", }; @@ -441,6 +452,7 @@ const AddGuardrailForm: React.FC = ({ visible, onClose, a guardrail_name: string; litellm_params: { guardrail: string; + logging_only_scope?: LoggingOnlyScope | null; [key: string]: unknown; // Allow dynamic properties }; guardrail_info: Record; @@ -464,6 +476,13 @@ const AddGuardrailForm: React.FC = ({ visible, onClose, a guardrailData.litellm_params.skip_tool_message_in_guardrail = skipToolForCreate; } + const loggingOnlyScope = choiceToLoggingOnlyScope( + values.logging_only_scope_choice as LoggingOnlyScopeChoice | undefined, + ); + if (modeIncludesLoggingOnly(values.mode) && loggingOnlyScope !== null) { + guardrailData.litellm_params.logging_only_scope = loggingOnlyScope; + } + // For Presidio PII, add the entity and action configurations if (providerKey === "PresidioPII" && selectedEntities.length > 0) { const piiEntitiesConfig: { [key: string]: string } = {}; @@ -798,6 +817,12 @@ const AddGuardrailForm: React.FC = ({ visible, onClose, a {(fieldControl) => } + + {/* Use the GuardrailProviderFields component to render provider-specific fields */} {showProviderFields && ( = new Set(LOGGING_ONLY_SCOPE_ITEMS.map(({ value }) => value)); + +const isLoggingOnlyScopeChoice = (value: unknown): value is LoggingOnlyScopeChoice => + typeof value === "string" && LOGGING_ONLY_SCOPE_CHOICES.has(value); + +export const CustomCodeLoggingOnlyScopeSelect: React.FC<{ + value: LoggingOnlyScopeChoice; + onChange: (choice: LoggingOnlyScopeChoice) => void; +}> = ({ value, onChange }) => ( +
+ + +
+); + +export const getCustomCodeLoggingOnlyScopeCreate = ( + mode: string[], + choice: LoggingOnlyScopeChoice, +): { logging_only_scope?: LoggingOnlyScope } => { + if (!mode.includes("logging_only")) return {}; + const scope = choiceToLoggingOnlyScope(choice); + return scope === null ? {} : { logging_only_scope: scope }; +}; + +export const getCustomCodeLoggingOnlyScopeUpdate = ( + mode: string[], + litellmParams: { logging_only_scope?: string | null } | null | undefined, + choice: LoggingOnlyScopeChoice, +): { logging_only_scope?: LoggingOnlyScope | null } => + mode.includes("logging_only") ? getLoggingOnlyScopeUpdate(litellmParams, choice) : {}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/CustomCodeModal.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/CustomCodeModal.test.tsx index fc19ca2a88b..a8b07a30329 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/CustomCodeModal.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/CustomCodeModal.test.tsx @@ -1,5 +1,5 @@ import { describe, it, expect, vi, beforeEach } from "vitest"; -import { render, screen, waitFor } from "@testing-library/react"; +import { fireEvent, render, screen, waitFor } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import CustomCodeModal from "./CustomCodeModal"; import { createGuardrailCall, updateGuardrailCall, testCustomCodeGuardrail } from "@/components/networking"; @@ -130,6 +130,127 @@ describe("CustomCodeModal", () => { expect(screen.queryByText("logging_only")).not.toBeInTheDocument(); }); + it("should update the logging-only scope in edit mode", async () => { + const user = userEvent.setup(); + renderModal({ + editData: { + guardrail_id: "g-1", + guardrail_name: "existing-guardrail", + litellm_params: { + mode: ["logging_only"], + logging_only_scope: "input", + custom_code: "def apply_guardrail(): pass", + }, + }, + }); + + const scopeSelect = await screen.findByRole("combobox", { name: "Logging only scope" }); + expect(scopeSelect).toHaveTextContent("Input only (request)"); + + await user.click(scopeSelect); + await user.click(await screen.findByRole("option", { name: "Output only (response)" })); + await user.click(screen.getByRole("button", { name: /update guardrail/i })); + + await waitFor(() => expect(mockUpdate).toHaveBeenCalledTimes(1)); + expect(mockUpdate.mock.calls[0][2]).toMatchObject({ + litellm_params: { logging_only_scope: "output" }, + }); + }); + + it("should clear the logging-only scope when Default is selected in edit mode", async () => { + const user = userEvent.setup(); + renderModal({ + editData: { + guardrail_id: "g-1", + guardrail_name: "existing-guardrail", + litellm_params: { + mode: ["logging_only"], + logging_only_scope: "input", + custom_code: "def apply_guardrail(): pass", + }, + }, + }); + + await user.click(await screen.findByRole("combobox", { name: "Logging only scope" })); + await user.click(await screen.findByRole("option", { name: "Default (request and response)" })); + await user.click(screen.getByRole("button", { name: /update guardrail/i })); + + await waitFor(() => expect(mockUpdate).toHaveBeenCalledTimes(1)); + expect(mockUpdate.mock.calls[0][2]).toHaveProperty("litellm_params.logging_only_scope", null); + }); + + it("should omit an unchanged logging-only scope from the edit payload", async () => { + renderModal({ + editData: { + guardrail_id: "g-1", + guardrail_name: "existing-guardrail", + litellm_params: { + mode: ["logging_only"], + logging_only_scope: "input", + custom_code: "def apply_guardrail(): pass", + }, + }, + }); + + fireEvent.change(await screen.findByPlaceholderText("e.g., block-pii-custom"), { + target: { value: "renamed-guardrail" }, + }); + await userEvent.setup().click(screen.getByRole("button", { name: /update guardrail/i })); + + await waitFor(() => expect(mockUpdate).toHaveBeenCalledTimes(1)); + expect(mockUpdate.mock.calls[0][2]).not.toHaveProperty("litellm_params.logging_only_scope"); + }); + + it("should not show logging-only scope for other modes in edit mode", async () => { + renderModal({ + editData: { + guardrail_id: "g-1", + guardrail_name: "existing-guardrail", + litellm_params: { mode: "pre_call", custom_code: "def apply_guardrail(): pass" }, + }, + }); + + expect(await screen.findByText("Edit Custom Guardrail")).toBeInTheDocument(); + expect(screen.queryByRole("combobox", { name: "Logging only scope" })).not.toBeInTheDocument(); + }); + + it("should create a logging-only guardrail with its selected scope", async () => { + const user = userEvent.setup(); + renderModal(); + + await user.click(screen.getAllByRole("combobox")[0]); + await user.keyboard("logging_only"); + await user.click(await screen.findByRole("option", { name: "logging_only" })); + await user.click(await screen.findByRole("combobox", { name: "Logging only scope" })); + await user.click(await screen.findByRole("option", { name: "Input only (request)" })); + fireEvent.change(screen.getByPlaceholderText("e.g., block-pii-custom"), { + target: { value: "logging-only-guardrail" }, + }); + await user.click(screen.getByRole("button", { name: /save guardrail/i })); + + await waitFor(() => expect(mockCreate).toHaveBeenCalledTimes(1)); + expect(mockCreate.mock.calls[0][1]).toMatchObject({ + litellm_params: { logging_only_scope: "input" }, + }); + }); + + it("should omit the default logging-only scope when creating a guardrail", async () => { + const user = userEvent.setup(); + renderModal(); + + await user.click(screen.getAllByRole("combobox")[0]); + await user.keyboard("logging_only"); + await user.click(await screen.findByRole("option", { name: "logging_only" })); + await screen.findByRole("combobox", { name: "Logging only scope" }); + fireEvent.change(screen.getByPlaceholderText("e.g., block-pii-custom"), { + target: { value: "logging-only-guardrail" }, + }); + await user.click(screen.getByRole("button", { name: /save guardrail/i })); + + await waitFor(() => expect(mockCreate).toHaveBeenCalledTimes(1)); + expect(mockCreate.mock.calls[0][1]).not.toHaveProperty("litellm_params.logging_only_scope"); + }); + it("should expand the test section and run a test against the backend", async () => { const user = userEvent.setup(); mockTest.mockResolvedValue({ success: true, result: { action: "allow" } } as never); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/CustomCodeModal.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/CustomCodeModal.tsx index 05a48598859..9deef4aeee3 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/CustomCodeModal.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/CustomCodeModal.tsx @@ -2,6 +2,13 @@ import React, { useState, useRef, useEffect } from "react"; import { CheckCircle2, ChevronRight, Code, ExternalLink, PlayCircle, Save, Users, XCircle } from "lucide-react"; import { createGuardrailCall, updateGuardrailCall, testCustomCodeGuardrail } from "@/components/networking"; import { toast } from "@/lib/toast"; +import { loggingOnlyScopeToChoice } from "../guardrail_info_helpers"; +import type { LoggingOnlyScope, LoggingOnlyScopeChoice } from "../guardrail_info_helpers"; +import { + CustomCodeLoggingOnlyScopeSelect, + getCustomCodeLoggingOnlyScopeCreate, + getCustomCodeLoggingOnlyScopeUpdate, +} from "./CustomCodeLoggingOnlyScope"; import { Button } from "@/components/ui/button"; import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible"; import { @@ -178,6 +185,7 @@ export interface EditGuardrailData { mode?: string | string[]; default_on?: boolean; custom_code?: string; + logging_only_scope?: LoggingOnlyScope | null; [key: string]: any; }; } @@ -196,6 +204,7 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS const isEditMode = !!editData; const [guardrailName, setGuardrailName] = useState(""); const [mode, setMode] = useState(["pre_call"]); + const [loggingOnlyScopeChoice, setLoggingOnlyScopeChoice] = useState("default"); const [defaultOn, setDefaultOn] = useState(false); const [selectedTemplate, setSelectedTemplate] = useState("empty"); const [code, setCode] = useState(CODE_TEMPLATES.empty.code); @@ -320,6 +329,7 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS // Edit mode: populate with existing data setGuardrailName(editData.guardrail_name || ""); setMode(normalizeMode(editData.litellm_params?.mode)); + setLoggingOnlyScopeChoice(loggingOnlyScopeToChoice(editData.litellm_params?.logging_only_scope)); setDefaultOn(editData.litellm_params?.default_on || false); setCode(editData.litellm_params?.custom_code || CODE_TEMPLATES.empty.code); setSelectedTemplate(""); // No template selected in edit mode @@ -327,6 +337,7 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS // Create mode: reset to defaults setGuardrailName(""); setMode(["pre_call"]); + setLoggingOnlyScopeChoice("default"); setDefaultOn(false); setSelectedTemplate("empty"); setCode(CODE_TEMPLATES.empty.code); @@ -384,6 +395,7 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS const updateData: any = { litellm_params: { custom_code: code, + ...getCustomCodeLoggingOnlyScopeUpdate(mode, editData.litellm_params, loggingOnlyScopeChoice), }, }; @@ -411,6 +423,7 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS mode: mode, default_on: defaultOn, custom_code: code, + ...getCustomCodeLoggingOnlyScopeCreate(mode, loggingOnlyScopeChoice), }, guardrail_info: {}, }; @@ -547,6 +560,9 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS + {mode.includes("logging_only") && ( + + )}