diff --git a/Cargo.lock b/Cargo.lock index 5c803231a..5bb46296d 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -417,6 +417,7 @@ dependencies = [ "ipnet", "jsonwebtoken", "lofty", + "percent-encoding", "rand 0.8.7", "regex", "reqwest 0.12.28", diff --git a/crates/aisix-core/src/models/passthrough_route.rs b/crates/aisix-core/src/models/passthrough_route.rs index 05db9d145..2413f29dc 100644 --- a/crates/aisix-core/src/models/passthrough_route.rs +++ b/crates/aisix-core/src/models/passthrough_route.rs @@ -170,10 +170,11 @@ pub struct PassthroughRoute { pub forward_client_headers: Vec, /// Maximum time, in milliseconds, for the upstream exchange. Bounds - /// the response-header phase and any non-SSE body read, but never a - /// healthy SSE relay (which ends with the upstream stream or the - /// client hanging up). When omitted, the gateway default request - /// timeout applies the same way. + /// the response-header phase and any non-SSE body read. For SSE it + /// bounds the wait for the first byte and every later no-byte gap, but + /// not the total duration of a healthy relay (which ends with the + /// upstream stream or the client hanging up). When omitted, the gateway + /// default request timeout applies the same way. #[serde(default, skip_serializing_if = "Option::is_none")] #[schemars(range(min = 1))] pub timeout_ms: Option, diff --git a/crates/aisix-gateway/src/upstream_tls.rs b/crates/aisix-gateway/src/upstream_tls.rs index 38f69984a..c547854be 100644 --- a/crates/aisix-gateway/src/upstream_tls.rs +++ b/crates/aisix-gateway/src/upstream_tls.rs @@ -211,6 +211,12 @@ pub fn reqwest_material() -> &'static ReqwestTlsMaterial { static PK_CLIENTS: OnceLock> = OnceLock::new(); +/// The raw-relay counterpart to [`PK_CLIENTS`]. Reqwest's content decoders +/// are client settings, so a passthrough relay that promises provider bytes +/// needs a separate pool even when it has the same TLS/resolve override. +static RAW_PK_CLIENTS: OnceLock> = + OnceLock::new(); + /// Clients built for Provider Key connection overrides so far. pub fn provider_key_client_count() -> u64 { PK_CLIENTS.get().map_or(0, |clients| clients.len() as u64) @@ -226,6 +232,11 @@ thread_local! { /// a build failure has been reported for this thread. static WORKER_CLIENT: std::cell::OnceCell> = const { std::cell::OnceCell::new() }; + + /// The worker-owned raw-relay pool. It has the same TLS and connection + /// settings as `WORKER_CLIENT`, but no automatic content decoding. + static RAW_WORKER_CLIENT: std::cell::OnceCell> = + const { std::cell::OnceCell::new() }; } /// Declares the calling thread a proxy worker with its own runtime, so @@ -270,6 +281,54 @@ fn worker_client() -> Option { }) } +/// Build a client that relays response bytes exactly as received. In +/// particular, reqwest otherwise transparently decodes `gzip`, `br`, +/// `deflate`, and `zstd`, while removing the matching response headers. +fn raw_client_builder() -> reqwest::ClientBuilder { + // This workspace enables only reqwest's `gzip` decoder (see the + // workspace dependency declaration). The remaining decoder features are + // not compiled in, so disabling gzip is sufficient to preserve every + // representation this binary could otherwise transform. + crate::upstream_http::client_builder() + .no_gzip() + // The passthrough target is validated before dispatch. Following an + // upstream Location would make a second, unchecked request and could + // carry the injected provider credential beyond that boundary. + .redirect(reqwest::redirect::Policy::none()) +} + +fn raw_client() -> &'static reqwest::Client { + static RAW_CLIENT: OnceLock = OnceLock::new(); + // The same builder's TLS material is validated at boot. Do not fall back + // to a bare client here: that could make a raw relay silently trust less + // than the deployment configured. + RAW_CLIENT.get_or_init(|| { + raw_client_builder() + .build() + .expect("configured raw passthrough HTTP client builds") + }) +} + +fn raw_worker_client() -> Option { + if !IS_WORKER_THREAD.get() { + return None; + } + RAW_WORKER_CLIENT.with(|cell| { + cell.get_or_init(|| match raw_client_builder().build() { + Ok(client) => Some(client), + Err(e) => { + tracing::error!( + error = %e, + "per-worker raw passthrough upstream pool could not be built; this worker \ + dispatches on the shared raw pool" + ); + None + } + }) + .clone() + }) +} + /// The client to dispatch this Provider Key's request on. /// /// Returns `shared` unchanged whenever the key sets no override, which @@ -294,11 +353,33 @@ pub fn client_for_provider_key( // the shared pool is used, as it always was. return worker_client().unwrap_or_else(|| shared.clone()); }; - let cache = PK_CLIENTS.get_or_init(dashmap::DashMap::new); + client_for_provider_key_override(shared, conn, &PK_CLIENTS, build_provider_key_client) +} + +/// The byte-preserving client for a passthrough raw relay. It applies the +/// same deployment TLS, per-ProviderKey CA, certificate-verification, and +/// address-resolution rules as [`client_for_provider_key`], but opts out of +/// every reqwest content decoder so a relay may retain both encoded bytes and +/// `Content-Encoding` / `Content-Length` headers. +pub fn raw_client_for_provider_key(conn: Option<&UpstreamConnection>) -> reqwest::Client { + let shared = raw_client(); + let Some(conn) = conn.filter(|c| !c.is_noop()) else { + return raw_worker_client().unwrap_or_else(|| shared.clone()); + }; + client_for_provider_key_override(shared, conn, &RAW_PK_CLIENTS, build_raw_provider_key_client) +} + +fn client_for_provider_key_override( + shared: &reqwest::Client, + conn: &UpstreamConnection, + cache: &OnceLock>, + build: fn(&UpstreamConnection) -> Result, +) -> reqwest::Client { + let cache = cache.get_or_init(dashmap::DashMap::new); if let Some(existing) = cache.get(conn) { return existing.clone(); } - match build_provider_key_client(conn) { + match build(conn) { Ok(client) => cache.entry(conn.clone()).or_insert(client).clone(), Err(e) if conn.resolve.is_empty() => { tracing::error!( @@ -333,7 +414,7 @@ pub fn client_for_provider_key( // `resolve_to_addrs` cannot fail, so the only way this second // build fails is a TLS backend that would not initialise — // which `shared` could not have been built over either. - build_provider_key_client(&resolution_only) + build(&resolution_only) .map(|client| cache.entry(resolution_only).or_insert(client).clone()) .unwrap_or_else(|_| shared.clone()) } @@ -341,11 +422,21 @@ pub fn client_for_provider_key( } fn build_provider_key_client(conn: &UpstreamConnection) -> Result { + build_provider_key_client_with(conn, crate::upstream_http::client_builder()) +} + +fn build_raw_provider_key_client(conn: &UpstreamConnection) -> Result { + build_provider_key_client_with(conn, raw_client_builder()) +} + +fn build_provider_key_client_with( + conn: &UpstreamConnection, + mut builder: reqwest::ClientBuilder, +) -> Result { // Layer the key's override ON TOP of the deployment settings rather // than replacing them: a deployment CA and a per-key CA are both // trust roots, and a client presenting the deployment's mTLS // identity must keep presenting it. - let mut builder = crate::upstream_http::client_builder(); if let Some(tls) = conn.tls.as_ref() { if let Some(pem) = tls.ca_cert.as_ref().filter(|p| !p.trim().is_empty()) { let roots = reqwest::Certificate::from_pem_bundle(pem.as_bytes()) @@ -804,6 +895,12 @@ mod tests { .is_some_and(|cache| cache.contains_key(conn)) } + fn raw_cached(conn: &UpstreamConnection) -> bool { + RAW_PK_CLIENTS + .get() + .is_some_and(|cache| cache.contains_key(conn)) + } + /// A key carrying only `tls`, in the shape the dispatch sites derive. fn tls_conn(tls: ProviderKeyTls) -> UpstreamConnection { UpstreamConnection { @@ -878,6 +975,16 @@ mod tests { assert!(cached(&conn)); } + #[test] + fn raw_relay_keeps_provider_key_connection_overrides() { + let conn = resolve_conn("vendor-raw-relay.invalid", &["192.0.2.12"]); + let _ = raw_client_for_provider_key(Some(&conn)); + assert!( + raw_cached(&conn), + "the raw relay must not discard a ProviderKey's address override" + ); + } + /// Two keys pointing the same hostname at different addresses must not /// share a client: the resolution lives on the client, so one pool /// could only ever dial one of the two. diff --git a/crates/aisix-proxy/Cargo.toml b/crates/aisix-proxy/Cargo.toml index b5c4a1376..d53be2ca9 100644 --- a/crates/aisix-proxy/Cargo.toml +++ b/crates/aisix-proxy/Cargo.toml @@ -48,7 +48,7 @@ http-body-util.workspace = true bytes.workspace = true regex.workspace = true serde.workspace = true -serde_json.workspace = true +serde_json = { workspace = true, features = ["raw_value"] } # Token-estimation fallback for usage telemetry (token_estimate.rs). tiktoken-rs.workspace = true futures.workspace = true @@ -61,6 +61,7 @@ base64.workspace = true # Scheme validation of provider-supplied redirect targets on # `GET /v1/videos/:id/content` (already in the tree via reqwest). url.workspace = true +percent-encoding.workspace = true # Per-request weighted-random target selection in `routing::weighted_pick`. # `thread_rng()` gives proper per-request entropy that converges to the # configured weights over a finite sample (fix for #197 — the prior diff --git a/crates/aisix-proxy/src/held_content.rs b/crates/aisix-proxy/src/held_content.rs index 1b9ae0286..f329b4bcf 100644 --- a/crates/aisix-proxy/src/held_content.rs +++ b/crates/aisix-proxy/src/held_content.rs @@ -1,9 +1,9 @@ //! What a streamed output guardrail's `max_buffer_bytes` measures (#513). //! -//! While a stream is held back for output inspection -//! ([`aisix_guardrails::StreamOutputPolicy::BufferFull`]), the cap bounds -//! the model-generated content held: assistant text, reasoning, and -//! tool-call arguments. SSE and JSON framing — event names, ids, indexes, +//! While a stream is held back for output inspection (a hold-back +//! [`aisix_guardrails::StreamOutputPolicy`]), the cap bounds +//! the model-generated content held: assistant text, refusals, reasoning, +//! and tool-call arguments. SSE and JSON framing — event names, ids, indexes, //! the envelope around each delta — is never counted, so the same response //! trips the cap at the same point whichever route and wire protocol //! carries it. @@ -18,10 +18,14 @@ //! [`RAW_HOLD_FACTOR`] times the same cap, through [`HeldBuffer`]. Crossing //! either bound is the same buffer-exceeded event. -use std::sync::atomic::{AtomicUsize, Ordering}; +use std::{ + collections::HashMap, + sync::atomic::{AtomicUsize, Ordering}, +}; use aisix_gateway::ChatDelta; use serde_json::Value; +use sha2::{Digest, Sha256}; /// Raw bytes a hold-back may keep, as a multiple of `max_buffer_bytes` /// (32 MiB at the 256 KiB default). Above the framing an ordinary token @@ -44,6 +48,21 @@ impl HeldBuffer { self.raw = self.raw.saturating_add(raw); } + /// Whether admitting one more held frame would cross either bound. + /// + /// The relay uses this before decoding an unterminated terminal frame so + /// the configured raw-buffer cap wins over its otherwise fail-closed + /// malformed-frame handling. + pub(crate) fn would_exceed_after( + &self, + content: usize, + raw: usize, + max_buffer_bytes: usize, + ) -> bool { + self.content.saturating_add(content) > max_buffer_bytes + || self.raw.saturating_add(raw) > max_buffer_bytes.saturating_mul(RAW_HOLD_FACTOR) + } + /// Past either bound: the content cap, or the raw-byte guard derived /// from it. pub(crate) fn exceeds(&self, max_buffer_bytes: usize) -> bool { @@ -178,8 +197,9 @@ pub(crate) fn chat_delta(delta: &ChatDelta) -> usize { } /// One stream event split the way the output guardrails read it: `scan` -/// is the generated text they inspect (assistant text and tool-call -/// arguments), `reasoning` the generated reasoning they do not inspect. +/// is the generated text they inspect (assistant text, refusals, and +/// tool-call arguments), `reasoning` the generated reasoning they do not +/// inspect. /// Both count toward the hold-back cap, so [`Parts::held`] is the cap's /// measure and `scan` the scanner's input — one extraction for both. #[derive(Debug, Default, PartialEq, Eq)] @@ -204,6 +224,521 @@ impl Parts { } } +#[derive(Clone, Copy)] +enum ResponsesCarrierMode { + Delta, + Snapshot, +} + +const MAX_RESPONSES_HELD_CARRIERS: usize = 1_024; +const MAX_RESPONSES_HELD_ID_BYTES: usize = 512; + +struct ResponsesHeldCarrier { + bytes: usize, + hash: Sha256, +} + +/// Logical generated content currently retained by a Responses SSE +/// hold-back buffer. OpenAI emits the same text as deltas, direct `.done` +/// events, content-part/item snapshots, and the terminal response object. +/// The raw-byte cap still counts every frame, while this state prevents those +/// equivalent source carriers from consuming the content cap repeatedly. +#[derive(Default)] +pub(crate) struct ResponsesHeldContent { + carriers: HashMap, + saturated: bool, +} + +impl ResponsesHeldContent { + /// Returns `None` when this is not one parseable Responses SSE frame; the + /// caller then uses the conservative generic held-content extraction. + pub(crate) fn observe_sse_frame(&mut self, frame: &[u8]) -> Option { + if self.saturated { + return None; + } + let payload = crate::redact::frame_payload(frame)?; + let payload = payload.trim(); + if payload.is_empty() || payload == "[DONE]" { + return Some(0); + } + let event = serde_json::from_str::(payload).ok()?; + self.observe_event(&event) + } + + /// Records one parsed event. `None` asks the caller to use the generic + /// content counter: either the ledger has reached its fixed key budget, + /// or this event was the one that reached it. + pub(crate) fn observe_event(&mut self, event: &Value) -> Option { + if self.saturated { + return None; + } + let held = self.observe_event_inner(event); + (!self.saturated).then_some(held) + } + + fn observe_event_inner(&mut self, event: &Value) -> usize { + match event.get("type").and_then(Value::as_str) { + Some("response.output_text.delta") => self.observe_event_content( + event, + "text", + event.get("delta").and_then(Value::as_str), + ResponsesCarrierMode::Delta, + ), + Some("response.refusal.delta") => self.observe_event_content( + event, + "refusal", + event.get("delta").and_then(Value::as_str), + ResponsesCarrierMode::Delta, + ), + Some("response.output_text.done") => self.observe_event_content( + event, + "text", + event.get("text").and_then(Value::as_str), + ResponsesCarrierMode::Snapshot, + ), + Some("response.refusal.done") => self.observe_event_content( + event, + "refusal", + event.get("refusal").and_then(Value::as_str), + ResponsesCarrierMode::Snapshot, + ), + Some("response.function_call_arguments.delta") => self.observe_event_tool( + event, + "function_call", + "arguments", + event.get("delta").and_then(Value::as_str), + ResponsesCarrierMode::Delta, + ), + Some("response.mcp_call_arguments.delta") => self.observe_event_tool( + event, + "mcp_call", + "arguments", + event.get("delta").and_then(Value::as_str), + ResponsesCarrierMode::Delta, + ), + Some("response.custom_tool_call_input.delta") => self.observe_event_tool( + event, + "custom_tool_call", + "input", + event.get("delta").and_then(Value::as_str), + ResponsesCarrierMode::Delta, + ), + Some("response.function_call_arguments.done") => { + self.observe_event_tool( + event, + "function_call", + "name", + event.get("name").and_then(Value::as_str), + ResponsesCarrierMode::Snapshot, + ) + self.observe_event_tool( + event, + "function_call", + "arguments", + event.get("arguments").and_then(Value::as_str), + ResponsesCarrierMode::Snapshot, + ) + } + Some("response.mcp_call_arguments.done") => { + self.observe_event_tool( + event, + "mcp_call", + "name", + event.get("name").and_then(Value::as_str), + ResponsesCarrierMode::Snapshot, + ) + self.observe_event_tool( + event, + "mcp_call", + "arguments", + event.get("arguments").and_then(Value::as_str), + ResponsesCarrierMode::Snapshot, + ) + } + Some("response.custom_tool_call_input.done") => { + self.observe_event_tool( + event, + "custom_tool_call", + "name", + event.get("name").and_then(Value::as_str), + ResponsesCarrierMode::Snapshot, + ) + self.observe_event_tool( + event, + "custom_tool_call", + "input", + event.get("input").and_then(Value::as_str), + ResponsesCarrierMode::Snapshot, + ) + } + Some("response.content_part.added" | "response.content_part.done") => event + .get("part") + .map(|part| self.observe_event_part(event, part, ResponsesCarrierMode::Snapshot)) + .unwrap_or(0), + Some("response.output_item.added" | "response.output_item.done") => event + .get("item") + .map(|item| { + self.observe_item( + item, + response_item_base(event, item), + ResponsesCarrierMode::Snapshot, + ) + }) + .unwrap_or(0), + Some("response.completed" | "response.incomplete" | "response.failed") => event + .get("response") + .and_then(|response| response.get("output")) + .and_then(Value::as_array) + .map(|items| { + items + .iter() + .enumerate() + .map(|(index, item)| { + self.observe_item( + item, + response_item_base_at(item, index), + ResponsesCarrierMode::Snapshot, + ) + }) + .sum() + }) + .unwrap_or(0), + Some("response.reasoning_text.delta") => self.observe_event_reasoning( + event, + "content", + event.get("content_index").and_then(Value::as_u64), + event.get("delta").and_then(Value::as_str), + ResponsesCarrierMode::Delta, + ), + Some("response.reasoning_summary_text.delta") => self.observe_event_reasoning( + event, + "summary", + event.get("summary_index").and_then(Value::as_u64), + event.get("delta").and_then(Value::as_str), + ResponsesCarrierMode::Delta, + ), + Some("response.reasoning_text.done") => self.observe_event_reasoning( + event, + "content", + event.get("content_index").and_then(Value::as_u64), + event.get("text").and_then(Value::as_str), + ResponsesCarrierMode::Snapshot, + ), + Some("response.reasoning_summary_text.done") => self.observe_event_reasoning( + event, + "summary", + event.get("summary_index").and_then(Value::as_u64), + event.get("text").and_then(Value::as_str), + ResponsesCarrierMode::Snapshot, + ), + Some( + "response.reasoning_summary_part.added" | "response.reasoning_summary_part.done", + ) => event + .get("part") + .and_then(|part| part.get("text")) + .and_then(Value::as_str) + .map(|text| { + self.observe_event_reasoning( + event, + "summary", + event.get("summary_index").and_then(Value::as_u64), + Some(text), + ResponsesCarrierMode::Snapshot, + ) + }) + .unwrap_or(0), + _ => 0, + } + } + + fn observe_event_content( + &mut self, + event: &Value, + field: &str, + text: Option<&str>, + mode: ResponsesCarrierMode, + ) -> usize { + let key = response_event_base(event).and_then(|base| { + event + .get("content_index") + .and_then(Value::as_u64) + .map(|index| format!("{base}/content/{index}/{field}")) + }); + self.observe_text(key, text, mode) + } + + fn observe_event_part( + &mut self, + event: &Value, + part: &Value, + mode: ResponsesCarrierMode, + ) -> usize { + match part.get("type").and_then(Value::as_str) { + Some("reasoning_text") => self.observe_event_reasoning( + event, + "content", + event.get("content_index").and_then(Value::as_u64), + part.get("text").and_then(Value::as_str), + mode, + ), + Some("summary_text") => self.observe_event_reasoning( + event, + "summary", + event.get("summary_index").and_then(Value::as_u64), + part.get("text").and_then(Value::as_str), + mode, + ), + _ => { + let Some((field, text)) = responses_part_text(part) else { + return 0; + }; + self.observe_event_content(event, field, Some(text), mode) + } + } + } + + fn observe_event_tool( + &mut self, + event: &Value, + tool_type: &str, + field: &str, + text: Option<&str>, + mode: ResponsesCarrierMode, + ) -> usize { + let key = response_event_base(event).map(|base| format!("{base}/tool/{tool_type}/{field}")); + self.observe_text(key, text, mode) + } + + fn observe_event_reasoning( + &mut self, + event: &Value, + group: &str, + index: Option, + text: Option<&str>, + mode: ResponsesCarrierMode, + ) -> usize { + let key = response_event_base(event) + .zip(index) + .map(|(base, index)| format!("{base}/reasoning/{group}/{index}")); + self.observe_text(key, text, mode) + } + + fn observe_item( + &mut self, + item: &Value, + base: Option, + mode: ResponsesCarrierMode, + ) -> usize { + match item.get("type").and_then(Value::as_str) { + Some("message") => item + .get("content") + .and_then(Value::as_array) + .map(|parts| { + parts + .iter() + .enumerate() + .map(|(index, part)| { + let Some((field, text)) = responses_part_text(part) else { + return 0; + }; + self.observe_text( + base.as_ref() + .map(|base| format!("{base}/content/{index}/{field}")), + Some(text), + mode, + ) + }) + .sum() + }) + .unwrap_or(0), + Some("reasoning") => ["content", "summary"] + .into_iter() + .map(|group| { + item.get(group) + .and_then(Value::as_array) + .map(|parts| { + parts + .iter() + .enumerate() + .map(|(index, part)| { + self.observe_text( + base.as_ref().map(|base| { + format!("{base}/reasoning/{group}/{index}") + }), + part.get("text").and_then(Value::as_str), + mode, + ) + }) + .sum::() + }) + .unwrap_or(0) + }) + .sum(), + Some("function_call" | "mcp_call") => { + let tool_type = item.get("type").and_then(Value::as_str).unwrap_or_default(); + self.observe_text( + base.as_ref() + .map(|base| format!("{base}/tool/{tool_type}/name")), + item.get("name").and_then(Value::as_str), + mode, + ) + self.observe_text( + base.map(|base| format!("{base}/tool/{tool_type}/arguments")), + item.get("arguments").and_then(Value::as_str), + mode, + ) + } + Some("custom_tool_call") => { + self.observe_text( + base.as_ref() + .map(|base| format!("{base}/tool/custom_tool_call/name")), + item.get("name").and_then(Value::as_str), + mode, + ) + self.observe_text( + base.map(|base| format!("{base}/tool/custom_tool_call/input")), + item.get("input").and_then(Value::as_str), + mode, + ) + } + _ => 0, + } + } + + fn observe_text( + &mut self, + key: Option, + text: Option<&str>, + mode: ResponsesCarrierMode, + ) -> usize { + let Some(text) = text.filter(|text| !text.is_empty()) else { + return 0; + }; + let Some(key) = key else { + // A missing/conflicting coordinate must never borrow another + // carrier's ledger entry. Charge it in full without retaining an + // unbounded anonymous key. + return text.len(); + }; + match mode { + ResponsesCarrierMode::Delta => { + if !self.carriers.contains_key(&key) + && self.carriers.len() >= MAX_RESPONSES_HELD_CARRIERS + { + self.saturated = true; + return text.len(); + } + match self.carriers.entry(key) { + std::collections::hash_map::Entry::Occupied(mut entry) => { + let entry = entry.get_mut(); + entry.bytes = entry.bytes.saturating_add(text.len()); + entry.hash.update(text.as_bytes()); + text.len() + } + std::collections::hash_map::Entry::Vacant(entry) => { + let mut hash = Sha256::new(); + hash.update(text.as_bytes()); + entry.insert(ResponsesHeldCarrier { + bytes: text.len(), + hash, + }); + text.len() + } + } + } + ResponsesCarrierMode::Snapshot => { + if !self.carriers.contains_key(&key) { + if self.carriers.len() >= MAX_RESPONSES_HELD_CARRIERS { + self.saturated = true; + return text.len(); + } + let mut hash = Sha256::new(); + hash.update(text.as_bytes()); + self.carriers.insert( + key, + ResponsesHeldCarrier { + bytes: text.len(), + hash, + }, + ); + return text.len(); + } + let previous = self + .carriers + .get_mut(&key) + .expect("carrier was present immediately before lookup"); + if text.len() == previous.bytes + && Sha256::digest(text.as_bytes()) == previous.hash.clone().finalize() + { + 0 + } else { + let prefix_matches = text.get(..previous.bytes).is_some_and(|prefix| { + Sha256::digest(prefix.as_bytes()) == previous.hash.clone().finalize() + }); + if prefix_matches { + let suffix = text + .get(previous.bytes..) + .expect("validated UTF-8 prefix boundary"); + previous.bytes = text.len(); + previous.hash.update(suffix.as_bytes()); + suffix.len() + } else { + let mut hash = Sha256::new(); + hash.update(text.as_bytes()); + *previous = ResponsesHeldCarrier { + bytes: text.len(), + hash, + }; + text.len() + } + } + } + } + } +} + +fn response_event_base(event: &Value) -> Option { + response_base( + event.get("item_id").and_then(Value::as_str), + event.get("output_index").and_then(Value::as_u64), + ) +} + +fn response_item_base(event: &Value, item: &Value) -> Option { + let top_level_id = event.get("item_id").and_then(Value::as_str); + let item_id = item.get("id").and_then(Value::as_str)?; + if top_level_id.is_some_and(|top_level_id| top_level_id != item_id) { + return None; + } + response_base( + Some(item_id), + event.get("output_index").and_then(Value::as_u64), + ) +} + +fn response_item_base_at(item: &Value, output_index: usize) -> Option { + response_base( + item.get("id").and_then(Value::as_str), + u64::try_from(output_index).ok(), + ) +} + +fn response_base(item_id: Option<&str>, output_index: Option) -> Option { + let item_id = item_id?; + let output_index = output_index?; + (!item_id.is_empty() && item_id.len() <= MAX_RESPONSES_HELD_ID_BYTES) + .then(|| format!("{}:{item_id}:{output_index}", item_id.len())) +} + +fn responses_part_text(part: &Value) -> Option<(&'static str, &str)> { + match part.get("type").and_then(Value::as_str) { + Some("output_text" | "text" | "input_text") => part + .get("text") + .and_then(Value::as_str) + .map(|text| ("text", text)), + Some("refusal") => part + .get("refusal") + .and_then(Value::as_str) + .map(|text| ("refusal", text)), + _ => None, + } +} + /// An Anthropic Messages stream event: `text` and `partial_json` deltas and /// the text or tool input a `content_block_start` already carries are /// scanned; `thinking` is reasoning. @@ -232,29 +767,126 @@ pub(crate) fn anthropic_event_parts(v: &Value) -> Parts { p } -/// An OpenAI Responses stream event. Only delta events count: the `.done` -/// events and the terminal `response.*` snapshot repeat content already -/// counted from its deltas. +fn responses_part_parts(parts: &mut Parts, part: &Value, reasoning: bool) { + let text = match part.get("type").and_then(Value::as_str) { + Some("output_text" | "text" | "input_text") => part.get("text"), + Some("refusal") => part.get("refusal"), + Some("reasoning_text" | "summary_text") => { + parts.reasoning_str(part.get("text")); + return; + } + _ => return, + }; + if reasoning { + parts.reasoning_str(text); + } else { + parts.scan_str(text); + } +} + +fn responses_item_parts(parts: &mut Parts, item: &Value) { + match item.get("type").and_then(Value::as_str) { + Some("reasoning") => { + for key in ["content", "summary"] { + for part in item + .get(key) + .and_then(Value::as_array) + .into_iter() + .flatten() + { + responses_part_parts(parts, part, true); + } + } + } + Some("message") => { + for part in item + .get("content") + .and_then(Value::as_array) + .into_iter() + .flatten() + { + responses_part_parts(parts, part, false); + } + } + Some("function_call" | "mcp_call") => { + parts.scan_str(item.get("name")); + parts.scan_str(item.get("arguments")); + } + Some("custom_tool_call") => { + parts.scan_str(item.get("name")); + parts.scan_str(item.get("input")); + } + _ => {} + } +} + +/// An OpenAI Responses stream event. A stream can legally end on a direct +/// `.done`, content-part, output-item, or terminal snapshot event without a +/// preceding delta, so every client-visible carrier contributes to the output +/// scan. The generic held-content counter below is deliberately conservative; +/// held Responses routes use [`ResponsesHeldContent`] to avoid charging the +/// same identified carrier again. Reasoning remains held but outside the +/// output-guardrail scan. pub(crate) fn responses_event_parts(v: &Value) -> Parts { let mut p = Parts::default(); match v.get("type").and_then(Value::as_str) { Some( "response.output_text.delta" + | "response.refusal.delta" | "response.function_call_arguments.delta" | "response.mcp_call_arguments.delta" | "response.custom_tool_call_input.delta", ) => p.scan_str(v.get("delta")), + Some("response.output_text.done") => p.scan_str(v.get("text")), + Some("response.refusal.done") => p.scan_str(v.get("refusal")), + Some("response.function_call_arguments.done" | "response.mcp_call_arguments.done") => { + p.scan_str(v.get("name")); + p.scan_str(v.get("arguments")); + } + Some("response.custom_tool_call_input.done") => { + p.scan_str(v.get("name")); + p.scan_str(v.get("input")); + } + Some("response.content_part.added" | "response.content_part.done") => { + if let Some(part) = v.get("part") { + responses_part_parts(&mut p, part, false); + } + } + Some("response.output_item.added" | "response.output_item.done") => { + if let Some(item) = v.get("item") { + responses_item_parts(&mut p, item); + } + } + Some("response.completed" | "response.incomplete" | "response.failed") => { + for item in v + .get("response") + .and_then(|response| response.get("output")) + .and_then(Value::as_array) + .into_iter() + .flatten() + { + responses_item_parts(&mut p, item); + } + } Some("response.reasoning_text.delta" | "response.reasoning_summary_text.delta") => { p.reasoning_str(v.get("delta")) } + Some("response.reasoning_text.done" | "response.reasoning_summary_text.done") => { + p.reasoning_str(v.get("text")) + } + Some("response.reasoning_summary_part.added" | "response.reasoning_summary_part.done") => { + if let Some(part) = v.get("part") { + responses_part_parts(&mut p, part, true); + } + } _ => {} } p } /// An OpenAI chat-completions stream chunk: every choice's `delta.content` -/// and tool-call arguments are scanned; `reasoning_content` (or the -/// `reasoning` spelling some relays use) is reasoning. +/// refusals, and tool-call arguments are scanned; `reasoning_content` (or +/// the `reasoning` spelling some relays use) is reasoning. pub(crate) fn chat_chunk_parts(v: &Value) -> Parts { let mut p = Parts::default(); for delta in v @@ -269,10 +901,14 @@ pub(crate) fn chat_chunk_parts(v: &Value) -> Parts { Some(Value::Array(parts)) => { for part in parts { p.scan_str(part.get("text")); + if part.get("type").and_then(Value::as_str) == Some("refusal") { + p.scan_str(part.get("refusal")); + } } } _ => {} } + p.scan_str(delta.get("refusal")); for tc in delta .get("tool_calls") .and_then(Value::as_array) @@ -282,6 +918,7 @@ pub(crate) fn chat_chunk_parts(v: &Value) -> Parts { p.scan_str(tc.get("function").and_then(|f| f.get("arguments"))); p.scan_str(tc.get("custom").and_then(|c| c.get("input"))); } + p.scan_str(delta.get("function_call").and_then(|f| f.get("arguments"))); p.reasoning_str(delta.get("reasoning_content")); p.reasoning_str(delta.get("reasoning")); } @@ -330,6 +967,28 @@ pub(crate) fn sse_frames(frames: &[u8], per_event: fn(&Value) -> usize) -> usize .sum() } +/// Held Responses content across every SSE frame in `frames`, with one +/// bounded source ledger shared by the complete frames and the final tail. +/// Once the ledger cannot safely identify more source carriers, fall back to +/// the generic counter rather than treating later content as free. +pub(crate) fn responses_sse_held_frames(ledger: &mut ResponsesHeldContent, frames: &[u8]) -> usize { + crate::redact::sse_frame_payloads(frames) + .iter() + .map(|payload| { + let payload = payload.trim(); + if payload.is_empty() || payload == "[DONE]" { + return 0; + } + match serde_json::from_str::(payload) { + Ok(event) => ledger + .observe_event(&event) + .unwrap_or_else(|| responses_event(&event)), + Err(_) => payload.len(), + } + }) + .sum() +} + #[cfg(test)] mod tests { use super::*; @@ -348,6 +1007,36 @@ mod tests { assert_eq!(chat_delta(&delta), 3 + 2 + 7); } + #[test] + fn chat_chunk_counts_refusals_and_legacy_tool_arguments_as_scanned_content() { + let direct = "direct streamed refusal"; + let direct_parts = chat_chunk_parts(&json!({ + "choices": [{"delta": {"refusal": direct}}], + })); + assert_eq!(direct_parts.scan, direct); + assert_eq!(direct_parts.held(), direct.len()); + + let typed = "typed streamed refusal"; + let typed_parts = chat_chunk_parts(&json!({ + "choices": [{"delta": {"content": [{"type": "refusal", "refusal": typed}]}}], + })); + assert_eq!(typed_parts.scan, typed); + assert_eq!(typed_parts.held(), typed.len()); + + let opaque_parts = chat_chunk_parts(&json!({ + "choices": [{"delta": {"content": [{"type": "future_media", "refusal": typed}]}}], + })); + assert!(opaque_parts.scan.is_empty()); + assert_eq!(opaque_parts.held(), 0); + + let arguments = "legacy streamed tool arguments"; + let legacy_parts = chat_chunk_parts(&json!({ + "choices": [{"delta": {"function_call": {"arguments": arguments}}}], + })); + assert_eq!(legacy_parts.scan, arguments); + assert_eq!(legacy_parts.held(), arguments.len()); + } + #[test] fn anthropic_event_ignores_envelope_and_empty_tool_input() { let start = json!({"type":"content_block_start","index":1,"content_block":{"type":"tool_use","id":"t","name":"lookup","input":{}}}); @@ -361,17 +1050,154 @@ mod tests { } #[test] - fn responses_frames_count_deltas_not_snapshots() { + fn responses_frames_count_every_authoritative_output_carrier() { let frames = concat!( "event: response.reasoning_summary_text.delta\n", "data: {\"type\":\"response.reasoning_summary_text.delta\",\"delta\":\"think\"}\n\n", "event: response.output_text.delta\n", "data: {\"type\":\"response.output_text.delta\",\"delta\":\"hi\"}\n\n", + "event: response.refusal.delta\n", + "data: {\"type\":\"response.refusal.delta\",\"delta\":\"no\"}\n\n", "event: response.output_text.done\n", "data: {\"type\":\"response.output_text.done\",\"text\":\"hi\"}\n\n", "data: [DONE]\n\n", ); - assert_eq!(sse_frames(frames.as_bytes(), responses_event), 5 + 2); + assert_eq!( + sse_frames(frames.as_bytes(), responses_event), + 5 + 2 + 2 + 2 + ); + } + + fn responses_frame(event: Value) -> Vec { + format!("data: {event}\n\n").into_bytes() + } + + #[test] + fn responses_held_content_counts_repeated_output_snapshots_once() { + let text = "x".repeat(300); + let frames = [ + responses_frame(json!({ + "type": "response.output_text.delta", + "item_id": "message_1", + "output_index": 0, + "content_index": 0, + "delta": text.as_str(), + })), + responses_frame(json!({ + "type": "response.output_text.done", + "item_id": "message_1", + "output_index": 0, + "content_index": 0, + "text": text.as_str(), + })), + responses_frame(json!({ + "type": "response.content_part.done", + "item_id": "message_1", + "output_index": 0, + "content_index": 0, + "part": { "type": "output_text", "text": text.as_str() }, + })), + responses_frame(json!({ + "type": "response.output_item.done", + "item_id": "message_1", + "output_index": 0, + "item": { + "id": "message_1", + "type": "message", + "content": [{ "type": "output_text", "text": text.as_str() }], + }, + })), + responses_frame(json!({ + "type": "response.completed", + "response": { + "output": [{ + "id": "message_1", + "type": "message", + "content": [{ "type": "output_text", "text": text.as_str() }], + }], + }, + })), + ]; + let mut ledger = ResponsesHeldContent::default(); + + assert_eq!( + responses_sse_held_frames(&mut ledger, &frames[0]), + text.len(), + "the initial delta establishes the logical carrier" + ); + assert_eq!( + responses_sse_held_frames(&mut ledger, &frames[1..].concat()), + 0, + "done, part, item, and terminal forms share that carrier across reads" + ); + } + + #[test] + fn responses_held_content_deduplicates_reasoning_part_snapshots() { + let text = "think"; + let frames = [ + responses_frame(json!({ + "type": "response.reasoning_text.delta", + "item_id": "reasoning_1", + "output_index": 0, + "content_index": 0, + "delta": text, + })), + responses_frame(json!({ + "type": "response.content_part.done", + "item_id": "reasoning_1", + "output_index": 0, + "content_index": 0, + "part": { "type": "reasoning_text", "text": text }, + })), + responses_frame(json!({ + "type": "response.output_item.done", + "item_id": "reasoning_1", + "output_index": 0, + "item": { + "id": "reasoning_1", + "type": "reasoning", + "content": [{ "type": "reasoning_text", "text": text }], + }, + })), + ]; + let mut ledger = ResponsesHeldContent::default(); + + assert_eq!( + frames + .iter() + .map(|frame| ledger.observe_sse_frame(frame).unwrap()) + .sum::(), + text.len() + ); + } + + #[test] + fn responses_held_content_counts_extensions_and_unidentified_snapshots() { + let mut ledger = ResponsesHeldContent::default(); + let delta = responses_frame(json!({ + "type": "response.output_text.delta", + "item_id": "message_1", + "output_index": 0, + "content_index": 0, + "delta": "abc", + })); + let extended = responses_frame(json!({ + "type": "response.output_text.done", + "item_id": "message_1", + "output_index": 0, + "content_index": 0, + "text": "abcdef", + })); + let unkeyed = responses_frame(json!({ + "type": "response.output_text.done", + "text": "abc", + })); + + assert_eq!(ledger.observe_sse_frame(&delta), Some(3)); + assert_eq!(ledger.observe_sse_frame(&extended), Some(3)); + assert_eq!(ledger.observe_sse_frame(&unkeyed), Some(3)); + assert_eq!(ledger.observe_sse_frame(&unkeyed), Some(3)); } #[test] diff --git a/crates/aisix-proxy/src/http_client.rs b/crates/aisix-proxy/src/http_client.rs index 56dcf59dd..e25e7397e 100644 --- a/crates/aisix-proxy/src/http_client.rs +++ b/crates/aisix-proxy/src/http_client.rs @@ -27,16 +27,23 @@ pub fn client() -> &'static Client { /// CA nor a resolution address — so the ordinary path keeps sharing one /// connection pool. /// -/// Every passthrough surface goes through here rather than [`client`]: -/// a key configured with a private CA, or with an upstream reachable only -/// at a fixed address, has to reach its endpoint on `/v1/messages`, -/// `/v1/responses`, `/v1/audio/*`, `/v1/videos/*`, the jobs surface and -/// the raw tunnel, not only on the endpoints that run through a provider -/// bridge. +/// Every typed passthrough surface goes through here rather than [`client`]. +/// The raw tunnel uses [`raw_relay_client_for`], which applies the same +/// per-key connection rules while preserving encoded response bytes. A key +/// configured with a private CA, or with an upstream reachable only at a +/// fixed address, must reach every endpoint rather than only the surfaces +/// that run through a provider bridge. pub fn client_for(conn: Option<&UpstreamConnection>) -> Client { aisix_gateway::upstream_tls::client_for_provider_key(client(), conn) } +/// The client for a raw passthrough relay. Unlike [`client_for`], this opts +/// out of reqwest's transparent content decoders so provider error and binary +/// responses retain their encoded bytes and representation headers. +pub fn raw_relay_client_for(conn: Option<&UpstreamConnection>) -> Client { + aisix_gateway::upstream_tls::raw_client_for_provider_key(conn) +} + #[cfg(test)] mod tests { /// The passthrough surfaces are a family — `/v1/messages`, diff --git a/crates/aisix-proxy/src/json_splice.rs b/crates/aisix-proxy/src/json_splice.rs index 2d9464515..6f15c7810 100644 --- a/crates/aisix-proxy/src/json_splice.rs +++ b/crates/aisix-proxy/src/json_splice.rs @@ -14,11 +14,11 @@ //! data — same rule as `collect_string_leaves` in the MCP scan path), //! but they ARE decoded to build the path handed to the predicate. //! -//! The scanner assumes syntactically valid JSON (callers run it on -//! bytes `serde_json` has already parsed) and still fails safe: any -//! unexpected byte, overrun, or depth blow-up returns an error rather -//! than a partially rewritten document. Callers decide the failure -//! policy (the MCP output hook fails closed). +//! The scanner is iterative, so deeply nested JSON does not consume the +//! Rust call stack. [`MAX_JSON_DEPTH`] bounds its per-request traversal +//! allocations. It still fails safe: any unexpected byte, overrun, or +//! depth excess returns an error rather than a partially rewritten document. +//! Callers decide the failure policy (the MCP output hook fails closed). use std::ops::Range; @@ -40,16 +40,112 @@ impl PathSeg { /// Scanner failure. Carries no document content (the byte offset only), /// so an error can be logged without leaking the payload. -#[derive(Debug, thiserror::Error)] +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum SpliceErrorKind { + Invalid, + DepthExceeded, + Unevaluable, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)] #[error("json splice scan failed at byte {at}")] pub struct SpliceError { at: usize, + kind: SpliceErrorKind, +} + +impl SpliceError { + /// Whether the scanner stopped at its bounded traversal limit rather than + /// because the input was malformed. + pub(crate) fn is_depth_exceeded(self) -> bool { + self.kind == SpliceErrorKind::DepthExceeded + } + + /// Whether source selection could not safely establish the fields a + /// guardrail is allowed to inspect. Unlike an invalid raw JSON document, + /// this must not fall back to scanning the whole body: that could expose + /// an opaque media carrier to an external guardrail. + pub(crate) fn is_unevaluable(self) -> bool { + matches!( + self.kind, + SpliceErrorKind::DepthExceeded | SpliceErrorKind::Unevaluable + ) + } + + /// Report an exhausted JSON traversal budget from a selector which uses + /// the same bounded raw-JSON policy as this scanner. Selectors retain raw + /// fragments rather than source offsets, so the synthetic error uses the + /// start of that fragment as its safe, non-payload-bearing location. + pub(crate) fn depth_exceeded() -> Self { + Self { + at: 0, + kind: SpliceErrorKind::DepthExceeded, + } + } + + /// Report a typed carrier that cannot be safely selected without falling + /// back to raw source text. + pub(crate) fn unevaluable() -> Self { + Self { + at: 0, + kind: SpliceErrorKind::Unevaluable, + } + } } -/// Depth cap. `serde_json` refuses documents deeper than 128, so bytes -/// that reached a splice call can never hit this; it bounds the scanner -/// on its own anyway. -const MAX_DEPTH: usize = 256; +/// Traversal depth cap. This keeps the iterative scanner stack-safe without +/// allowing an unbounded JSON path/frame allocation from an unbounded raw +/// passthrough body. It remains well beyond serde_json's usual recursion cap. +pub(crate) const MAX_JSON_DEPTH: usize = 4_096; + +/// Bounded collection policy for decoded JSON text passed to guardrails. +/// Source selectors use the same limits so a wide document cannot shift its +/// allocation from selector bookkeeping into decoded scan text. +pub(crate) const MAX_JSON_SCAN_VALUES: usize = 1_024; +pub(crate) const MAX_JSON_SCAN_TEXT_BYTES: usize = 256 * 1024; + +/// Validate JSON-number syntax without materializing the value. serde_json +/// can reject a syntactically valid number outside its runtime numeric range. +pub(crate) fn is_json_number(token: &[u8]) -> bool { + let mut pos = 0; + if token.get(pos) == Some(&b'-') { + pos += 1; + } + match token.get(pos) { + Some(b'0') => pos += 1, + Some(b'1'..=b'9') => { + pos += 1; + while token.get(pos).is_some_and(|byte| byte.is_ascii_digit()) { + pos += 1; + } + } + _ => return false, + } + if token.get(pos) == Some(&b'.') { + pos += 1; + let fraction_start = pos; + while token.get(pos).is_some_and(|byte| byte.is_ascii_digit()) { + pos += 1; + } + if pos == fraction_start { + return false; + } + } + if matches!(token.get(pos), Some(b'e' | b'E')) { + pos += 1; + if matches!(token.get(pos), Some(b'+' | b'-')) { + pos += 1; + } + let exponent_start = pos; + while token.get(pos).is_some_and(|byte| byte.is_ascii_digit()) { + pos += 1; + } + if pos == exponent_start { + return false; + } + } + pos == token.len() +} /// Rewrite the string values of `input` selected by `should_rewrite`, /// leaving every other byte untouched. @@ -71,7 +167,11 @@ pub fn rewrite_string_values( Array, } - let err = |at: usize| SpliceError { at }; + let err = |at: usize| SpliceError { + at, + kind: SpliceErrorKind::Invalid, + }; + std::str::from_utf8(input).map_err(|error| err(error.valid_up_to()))?; let mut splices: Vec<(Range, String)> = Vec::new(); let mut path: Vec = Vec::new(); let mut frames: Vec = Vec::new(); @@ -88,16 +188,32 @@ pub fn rewrite_string_values( let mut i = start + 1; while i < input.len() { match input[i] { - b'\\' => i += 2, // skips the escaped byte; `\uXXXX` needs no care (hex only) + b'\\' => match input.get(i + 1).copied() { + Some(b'"' | b'\\' | b'/' | b'b' | b'f' | b'n' | b'r' | b't') => i += 2, + Some(b'u') => { + let Some(hex) = input.get(i + 2..i + 6) else { + return Err(err(i)); + }; + if !hex.iter().all(|byte| byte.is_ascii_hexdigit()) { + return Err(err(i)); + } + i += 6; + } + _ => return Err(err(i)), + }, b'"' => return Ok(i + 1), + 0..=0x1f => return Err(err(i)), _ => i += 1, } } - Err(SpliceError { at: start }) + Err(err(start)) }; let decode_str = |range: Range| -> Result { let at = range.start; - serde_json::from_slice::(&input[range]).map_err(|_| SpliceError { at }) + serde_json::from_slice::(&input[range]).map_err(|_| SpliceError { + at, + kind: SpliceErrorKind::Invalid, + }) }; // `true` → the loop continues at a VALUE position; `false` → the @@ -108,8 +224,11 @@ pub fn rewrite_string_values( match b { b'{' => { frames.push(Frame::Object); - if frames.len() > MAX_DEPTH { - return Err(err(pos)); + if frames.len() > MAX_JSON_DEPTH { + return Err(SpliceError { + at: pos, + kind: SpliceErrorKind::DepthExceeded, + }); } pos += 1; skip_ws(&mut pos); @@ -135,8 +254,11 @@ pub fn rewrite_string_values( } b'[' => { frames.push(Frame::Array); - if frames.len() > MAX_DEPTH { - return Err(err(pos)); + if frames.len() > MAX_JSON_DEPTH { + return Err(SpliceError { + at: pos, + kind: SpliceErrorKind::DepthExceeded, + }); } pos += 1; skip_ws(&mut pos); @@ -161,16 +283,23 @@ pub fn rewrite_string_values( } pos = end; } - // Number / true / false / null. The scanner does not - // re-validate the token — the bytes already parsed upstream — - // it only needs the token's extent. + // Number / true / false / null. b'-' | b'0'..=b'9' | b't' | b'f' | b'n' => { + let start = pos; while pos < input.len() && matches!(input[pos], b'-' | b'+' | b'.' | b'0'..=b'9' | b'a'..=b'z' | b'A'..=b'Z') { pos += 1; } + let token = &input[start..pos]; + if token != b"true" + && token != b"false" + && token != b"null" + && !is_json_number(token) + { + return Err(err(start)); + } } _ => return Err(err(pos)), } @@ -242,6 +371,108 @@ pub fn rewrite_string_values( Ok(Some(out)) } +/// Decode and collect every JSON string **value** in source order. +/// +/// This reuses the iterative splice scanner with a no-op rewrite, so it +/// preserves duplicate keys and works beyond serde_json's default +/// container-recursion limit, up to [`MAX_JSON_DEPTH`]. Object keys are +/// decoded only to maintain the scanner's structure and are never included in +/// the returned text. +pub fn collect_string_values(input: &[u8]) -> Result { + collect_string_values_where(input, |_| true) +} + +/// Validate one UTF-8 JSON document with the same iterative depth limit as +/// the guardrail selector, without retaining any of its string values. +pub(crate) fn validate_json(input: &[u8]) -> Result<(), SpliceError> { + rewrite_string_values(input, |_| false, |_| None).map(|_| ()) +} + +/// Decode and collect selected JSON string **values** in source order. +/// +/// Like [`collect_string_values`], this preserves duplicate keys and stays +/// stack-safe within the bounded nesting limit. The predicate sees the +/// decoded path of each string value, never an object key. +pub fn collect_string_values_where( + input: &[u8], + mut include: impl FnMut(&[PathSeg]) -> bool, +) -> Result { + let mut out = String::new(); + let mut collect_error = None; + rewrite_string_values( + input, + |path| include(path), + |value| { + if collect_error.is_none() { + let separator = if out.is_empty() { 0 } else { 1 }; + let Some(next_len) = out + .len() + .checked_add(separator) + .and_then(|len| len.checked_add(value.len())) + else { + collect_error = Some(SpliceError::unevaluable()); + return None; + }; + if next_len > MAX_JSON_SCAN_TEXT_BYTES + || out.try_reserve(next_len - out.len()).is_err() + { + collect_error = Some(SpliceError::unevaluable()); + return None; + } + if separator != 0 { + out.push('\n'); + } + out.push_str(value); + } + None + }, + )?; + collect_error.map_or(Ok(out), Err) +} + +/// Decode selected JSON string values as separate source-order entries. +/// +/// Stream guardrails use this form to keep separate repeated carrier fields +/// in independent continuation channels rather than inserting separators +/// into a literal split across frames. +pub fn collect_string_values_where_vec( + input: &[u8], + mut include: impl FnMut(&[PathSeg]) -> bool, +) -> Result, SpliceError> { + let mut out = Vec::new(); + let mut source_bytes = 0usize; + let mut collect_error = None; + rewrite_string_values( + input, + |path| include(path), + |value| { + if collect_error.is_none() { + let Some(next_bytes) = source_bytes.checked_add(value.len()) else { + collect_error = Some(SpliceError::unevaluable()); + return None; + }; + if out.len() >= MAX_JSON_SCAN_VALUES + || next_bytes > MAX_JSON_SCAN_TEXT_BYTES + || out.try_reserve(1).is_err() + { + collect_error = Some(SpliceError::unevaluable()); + return None; + } + let mut decoded = String::new(); + if decoded.try_reserve(value.len()).is_err() { + collect_error = Some(SpliceError::unevaluable()); + return None; + } + decoded.push_str(value); + source_bytes = next_bytes; + out.push(decoded); + } + None + }, + )?; + collect_error.map_or(Ok(out), Err) +} + #[cfg(test)] mod tests { use super::*; @@ -272,6 +503,11 @@ mod tests { .is_none()); } + #[test] + fn validates_large_exponent_without_materializing_a_number() { + assert!(validate_json(br#"{"n":1e400}"#).is_ok()); + } + #[test] fn keys_are_never_offered_but_shape_the_path() { let doc = r#"{"secret": {"inner": "value"}}"#; @@ -320,6 +556,59 @@ mod tests { assert_eq!(seen, vec!["a", "b", "c"]); } + #[test] + fn collects_deep_string_values_without_recursion() { + let depth = 512; + let mut doc = "{\"v\":".repeat(depth); + doc.push_str(r#""\u0042LOCKME""#); + doc.push_str(&"}".repeat(depth)); + + assert_eq!(collect_string_values(doc.as_bytes()).unwrap(), "BLOCKME"); + } + + #[test] + fn collects_selected_paths_with_duplicate_keys_and_nested_values() { + let doc = r#"{"model":"routing-only","messages":[{"content":"first","metadata":{"note":"nested"}}],"messages":[{"content":"second"}]}"#; + assert_eq!( + collect_string_values_where(doc.as_bytes(), |path| { + !path.first().is_some_and(|segment| segment.is_key("model")) + }) + .unwrap(), + "first\nnested\nsecond" + ); + } + + #[test] + fn collects_selected_values_as_separate_source_ordered_entries() { + let doc = r#"{"type":"response.output_text.delta","delta":"FOR","delta":"ok"}"#; + assert_eq!( + collect_string_values_where_vec(doc.as_bytes(), |path| { + path.first().is_some_and(|segment| segment.is_key("delta")) + }) + .unwrap(), + vec!["FOR", "ok"] + ); + } + + #[test] + fn collected_json_text_over_the_shared_cap_is_unevaluable() { + let doc = format!( + r#"{{"text":"{}"}}"#, + "x".repeat(MAX_JSON_SCAN_TEXT_BYTES + 1) + ); + let error = collect_string_values(doc.as_bytes()) + .expect_err("a JSON text collection must stay bounded"); + assert!(error.is_unevaluable(), "{error}"); + let error = collect_string_values_where_vec(doc.as_bytes(), |_| true) + .expect_err("vector collection shares the text cap"); + assert!(error.is_unevaluable(), "{error}"); + + let values = format!("[{}]", "\"\",".repeat(MAX_JSON_SCAN_VALUES) + "\"\"",); + let error = collect_string_values_where_vec(values.as_bytes(), |_| true) + .expect_err("vector collection also bounds empty values"); + assert!(error.is_unevaluable(), "{error}"); + } + #[test] fn escaped_key_decodes_for_the_predicate() { // `param\u0073` decodes to "params" — the predicate must see the @@ -399,7 +688,7 @@ mod tests { } #[test] - fn depth_cap_errors() { + fn deep_nesting_is_stack_safe() { let mut doc = String::new(); for _ in 0..300 { doc.push('['); @@ -407,6 +696,16 @@ mod tests { for _ in 0..300 { doc.push(']'); } - assert!(rewrite_string_values(doc.as_bytes(), |_| true, |_| None).is_err()); + assert!(rewrite_string_values(doc.as_bytes(), |_| true, |_| None) + .expect("valid deep JSON") + .is_none()); + } + + #[test] + fn nesting_beyond_the_cap_errors() { + let depth = MAX_JSON_DEPTH + 1; + let doc = format!("{}\"value\"{}", "[".repeat(depth), "]".repeat(depth)); + let err = rewrite_string_values(doc.as_bytes(), |_| true, |_| None).unwrap_err(); + assert!(err.is_depth_exceeded()); } } diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index 608a5e752..d951eca8a 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -75,6 +75,8 @@ use axum::extract::{Request, State}; use axum::http::{header, HeaderMap, HeaderValue, Method, StatusCode}; use axum::middleware::Next; use axum::response::{IntoResponse, Response}; +use percent_encoding::percent_decode; +use url::Url; use aisix_core::resource::ResourceEntry; use aisix_core::{PassthroughAuthMode, PassthroughCredentialMode, PassthroughRoute}; @@ -401,6 +403,46 @@ impl RouteError { } } +/// `true` only for the SSE media type itself. Parameters are permitted, but a +/// prefix such as `text/event-streaming` is an ordinary buffered response and +/// must remain subject to the normal output-guardrail path. +fn response_is_sse(headers: &HeaderMap) -> bool { + let mut values = headers.get_all(header::CONTENT_TYPE).iter(); + let Some(value) = values.next() else { + return false; + }; + values.next().is_none() + && value.to_str().ok().is_some_and(|value| { + value.split(';').next().is_some_and(|media_type| { + media_type.trim().eq_ignore_ascii_case("text/event-stream") + }) + }) +} + +/// A raw-relay client must never hand encoded success bytes to an output +/// selector. `identity` is harmless; every other (or malformed) coding is +/// unavailable for inspection unless the upstream honored our identity-only +/// request negotiation. +fn response_has_non_identity_content_encoding(headers: &HeaderMap) -> bool { + headers + .get_all(header::CONTENT_ENCODING) + .iter() + .any(|value| { + let Ok(value) = value.to_str() else { + return true; + }; + let mut saw_coding = false; + for coding in value.split(',') { + let coding = coding.trim(); + if coding.is_empty() || !coding.eq_ignore_ascii_case("identity") { + return true; + } + saw_coding = true; + } + !saw_coding + }) +} + // --------------------------------------------------------------------------- // The pipeline // --------------------------------------------------------------------------- @@ -521,15 +563,8 @@ async fn dispatch( } else { rest_raw }; - let url = if rest.is_empty() { - base.clone() - } else { - format!("{base}/{rest}") - }; - let url = match &query { - Some(q) => format!("{url}?{q}"), - None => url, - }; + let url = + join_target_url(&base, rest, query.as_deref()).map_err(|e| RouteError::of(e, &auth))?; // End-user identity injected by the upstream device, captured before // the strip pass and recorded on the usage event. @@ -577,44 +612,90 @@ async fn dispatch( let resolved_chain = state.guardrail_index.resolve(&guardrail_ctx); *audit_out = resolved_chain.audit_log(); let mut monitor_hits: Vec = Vec::new(); + let output_guardrail_active = aisix_guardrails::Guardrail::runs_on_output(&resolved_chain); // Envelope detection: once per exchange, from the request body's - // top-level keys; the response and stream frames reuse it. - let protocol = detect_protocol(&body_bytes); + // top-level keys; the response and stream frames reuse it. Keep a + // malformed JSON-like root distinct from an ordinary opaque body: a + // failed selector must never silently become a whole-body raw scan. + let (protocol, probe_error) = match probe_json(&body_bytes) { + JsonProbe::Raw => (PassthroughProtocol::Raw, None), + JsonProbe::Protocol(protocol) => (protocol, None), + JsonProbe::Unevaluable(error) => (PassthroughProtocol::Raw, Some(error)), + }; let raw_shape = detect_raw_usage_shape(protocol, &body_bytes); // INPUT guardrails on the (envelope-extracted) request text. if !resolved_chain.is_empty() { - let text = request_guardrail_text(protocol, &body_bytes); - let chat = aisix_gateway::ChatFormat::new( - route.name.clone(), - vec![aisix_gateway::ChatMessage::user(text)], - ); - let (verdict, hits) = - aisix_guardrails::Guardrail::check_input_unmaskable_observed(&resolved_chain, &chat) - .await; - monitor_hits.extend(hits); - if let aisix_guardrails::GuardrailVerdict::Block { - reason, - guardrail_name, - unavailable, - } = verdict - { - // Per #153 the matched-pattern detail stays in ops logs only. - tracing::warn!( - guardrail_hook = "input", - route = %route.name, - reason = %reason, - "guardrail blocked passthrough-route request", + let request_text = + probe_error.map_or_else(|| try_request_guardrail_text(protocol, &body_bytes), Err); + let text = match request_text { + Ok(text) => Some(text), + Err(err) if !err.is_unevaluable() => { + Some(request_guardrail_text(protocol, &body_bytes)) + } + Err(err) + if !aisix_guardrails::Guardrail::refuses_unevaluable_input(&resolved_chain) => + { + tracing::debug!( + guardrail_hook = "input", + route = %route.name, + error = %err, + "cannot safely select passthrough-route request text; resolved chain does not fail closed", + ); + resolved_chain.record_unevaluable_input_bypass(crate::error::TAG_UNSCANNABLE_BODY); + None + } + Err(err) => { + tracing::warn!( + guardrail_hook = "input", + route = %route.name, + error = %err, + "cannot safely select passthrough-route request text; blocking", + ); + return Err(RouteError::of( + crate::error::guardrail_block_error( + "request", + None, + Some(crate::error::TAG_UNSCANNABLE_BODY), + ), + &auth, + )); + } + }; + if let Some(text) = text { + let chat = aisix_gateway::ChatFormat::new( + route.name.clone(), + vec![aisix_gateway::ChatMessage::user(text)], ); - return Err(RouteError::of( - crate::error::guardrail_block_error( - "request", - guardrail_name.as_deref(), - unavailable.as_deref(), - ), - &auth, - )); + let (verdict, hits) = aisix_guardrails::Guardrail::check_input_unmaskable_observed( + &resolved_chain, + &chat, + ) + .await; + monitor_hits.extend(hits); + if let aisix_guardrails::GuardrailVerdict::Block { + reason, + guardrail_name, + unavailable, + } = verdict + { + // Per #153 the matched-pattern detail stays in ops logs only. + tracing::warn!( + guardrail_hook = "input", + route = %route.name, + reason = %reason, + "guardrail blocked passthrough-route request", + ); + return Err(RouteError::of( + crate::error::guardrail_block_error( + "request", + guardrail_name.as_deref(), + unavailable.as_deref(), + ), + &auth, + )); + } } } @@ -649,7 +730,7 @@ async fn dispatch( .map(|pk| pk.value.provider.to_ascii_lowercase()) .filter(|prov| !prov.is_empty()) .and_then(|prov| body_model_rate_limit(snapshot, &prov, &body_bytes)); - let _reservation = crate::quota::enforce(state, snapshot, &auth, model_rl.as_ref()) + let reservation = crate::quota::enforce(state, snapshot, &auth, model_rl.as_ref()) .await .map_err(|e| RouteError::of(e, &auth))?; @@ -658,7 +739,11 @@ async fn dispatch( let conn = pk_entry .as_ref() .and_then(|pk| pk.value.upstream_connection()); - let http_client = crate::http_client::client_for(conn.as_ref()); + // A passthrough route promises provider response bytes verbatim. Use the + // no-decode client even for buffered responses: otherwise reqwest can + // transparently inflate a non-success SSE body while stripping its + // representation headers before `stream_opaque_response` relays it. + let http_client = crate::http_client::raw_relay_client_for(conn.as_ref()); // Strip set: protocol metadata always; per-mode credential handling. let mut strip: std::collections::HashSet = @@ -771,6 +856,13 @@ async fn dispatch( if aisix_core::header_forward_blocked(&lower) { continue; } + // The raw relay disables reqwest decompression to preserve an + // upstream's non-success SSE bytes. When output guardrails need to + // inspect successful replies, negotiate an inspectable response + // rather than forwarding a client request for a compressed one. + if output_guardrail_active && lower == "accept-encoding" { + continue; + } if strip.contains(&lower) { if !forwards(&lower) { continue; @@ -780,6 +872,10 @@ async fn dispatch( builder = builder.header(name, value); } + if output_guardrail_active { + builder = builder.header(header::ACCEPT_ENCODING, "identity"); + } + // Inject the gateway-held upstream credential (inject mode only). // Strip ran first, so this never adds a second value to a slot the // caller's own header already took (#411 ordering). That is a @@ -824,6 +920,15 @@ async fn dispatch( .timeout_ms .map(Duration::from_millis) .or(state.default_timeouts.request); + // A healthy SSE relay can be arbitrarily long, but no single silence + // gap may keep its concurrency reservation forever. Route-level timeout + // is the most specific bound; otherwise mirror the deployment stream → + // request fallback used by typed streaming routes. + let stream_read_timeout = route + .timeout_ms + .map(Duration::from_millis) + .or(state.default_timeouts.stream) + .or(state.default_timeouts.request); let bridge_timeout = |d: Duration| aisix_gateway::BridgeError::Timeout { elapsed_ms: d.as_millis().min(u64::MAX as u128) as u64, @@ -859,15 +964,12 @@ async fn dispatch( // same guess would 502 an upstream that merely mislabels itself. Not // drift: see that function's doc comment for why the two populations // take opposite defaults. - let is_sse = resp_headers - .get(header::CONTENT_TYPE) - .and_then(|v| v.to_str().ok()) - .map(|v| { - v.trim_start() - .to_ascii_lowercase() - .starts_with("text/event-stream") - }) - .unwrap_or(false); + let is_sse = response_is_sse(&resp_headers); + let success_response_is_encoded = status.is_success() + && output_guardrail_active + && response_has_non_identity_content_encoding(&resp_headers); + let bypass_uninspectable_output = success_response_is_encoded + && !aisix_guardrails::Guardrail::refuses_unevaluable_output(&resolved_chain); let mut telemetry = RouteTelemetry { state: state.clone(), @@ -914,8 +1016,57 @@ async fn dispatch( emitted: false, }; + // The raw relay keeps error representations untouched, including their + // `Content-Encoding` and byte length. A successful reply that ignores + // our identity-only negotiation cannot be parsed safely by either the + // buffered or SSE output selector, so refuse it before any compressed + // bytes can become a guardrail input. A configured fail-open output + // policy still forwards the original representation and records the + // bypass; fail-closed keeps the refusal contract. + if success_response_is_encoded { + if bypass_uninspectable_output { + tracing::debug!( + guardrail_hook = "output", + route = %route.name, + "cannot inspect an encoded successful passthrough-route response; resolved chain does not fail closed", + ); + resolved_chain.record_unevaluable_output_bypass(crate::error::TAG_UNSCANNABLE_BODY); + } else { + tracing::warn!( + guardrail_hook = "output", + route = %route.name, + "cannot inspect an encoded successful passthrough-route response; blocking", + ); + telemetry.guardrail_blocked = true; + telemetry.emitted = true; + return Err(RouteError::of( + crate::error::guardrail_block_error( + "response", + None, + Some(crate::error::TAG_UNSCANNABLE_BODY), + ), + &auth, + )); + } + } + if is_sse { telemetry.streaming = true; + let stream_hold = reservation.into_stream_hold(); + // Non-success replies, and encoded fail-open success replies, are + // opaque relay contracts. They must remain byte-for-byte data: no + // output guardrail, heartbeat, or SSE parsing may rewrite them. + if !status.is_success() || bypass_uninspectable_output { + return Ok(stream_opaque_response( + upstream_resp, + resp_headers, + status, + telemetry, + &client.request_id, + stream_hold, + stream_read_timeout, + )); + } return Ok(stream_response( protocol, resolved_chain, @@ -924,6 +1075,8 @@ async fn dispatch( status, telemetry, &client.request_id, + stream_hold, + stream_read_timeout, )); } @@ -951,52 +1104,90 @@ async fn dispatch( ) })?; - // OUTPUT guardrails on the (envelope-extracted) response text. - if !resolved_chain.is_empty() { - let text = response_guardrail_text(protocol, &resp_body); - let synth = aisix_gateway::ChatResponse { - id: String::new(), - model: route.name.clone(), - message: aisix_gateway::ChatMessage::assistant(text), - finish_reason: aisix_gateway::FinishReason::Stop, - usage: aisix_gateway::UsageStats::default(), + // Output guardrails govern generated successful answers. A provider's + // non-success body is its error contract, so preserve its status, headers, + // and bytes instead of replacing a 4xx/5xx with a local guardrail 422. + if status.is_success() && !bypass_uninspectable_output && output_guardrail_active { + let text = match try_buffered_response_guardrail_text(protocol, &resp_body) { + Ok(text) => Some(text), + Err(err) + if !aisix_guardrails::Guardrail::refuses_unevaluable_output(&resolved_chain) => + { + tracing::debug!( + guardrail_hook = "output", + route = %route.name, + error = %err, + "cannot safely select passthrough-route response text; resolved chain does not fail closed", + ); + resolved_chain.record_unevaluable_output_bypass(crate::error::TAG_UNSCANNABLE_BODY); + None + } + Err(err) => { + tracing::warn!( + guardrail_hook = "output", + route = %route.name, + error = %err, + "cannot safely select passthrough-route response text; blocking", + ); + telemetry.guardrail_blocked = true; + telemetry.emitted = true; + return Err(RouteError::of( + crate::error::guardrail_block_error( + "response", + None, + Some(crate::error::TAG_UNSCANNABLE_BODY), + ), + &auth, + )); + } }; - let (verdict, hits) = - aisix_guardrails::Guardrail::check_output_unmaskable_observed(&resolved_chain, &synth) - .await; - telemetry.monitor_hits.extend(hits); - if let aisix_guardrails::GuardrailVerdict::Block { - reason, - guardrail_name, - unavailable, - } = verdict - { - tracing::warn!( - guardrail_hook = "output", - route = %route.name, - reason = %reason, - "guardrail blocked passthrough-route response", - ); - telemetry.guardrail_blocked = true; - // The telemetry guard has not emitted yet; drop it silently and - // let the shared error path report the 422. - telemetry.emitted = true; - return Err(RouteError::of( - crate::error::guardrail_block_error( - "response", - guardrail_name.as_deref(), - unavailable.as_deref(), - ), - &auth, - )); + if let Some(text) = text { + let synth = aisix_gateway::ChatResponse { + id: String::new(), + model: route.name.clone(), + message: aisix_gateway::ChatMessage::assistant(text), + finish_reason: aisix_gateway::FinishReason::Stop, + usage: aisix_gateway::UsageStats::default(), + }; + let (verdict, hits) = aisix_guardrails::Guardrail::check_output_unmaskable_observed( + &resolved_chain, + &synth, + ) + .await; + telemetry.monitor_hits.extend(hits); + if let aisix_guardrails::GuardrailVerdict::Block { + reason, + guardrail_name, + unavailable, + } = verdict + { + tracing::warn!( + guardrail_hook = "output", + route = %route.name, + reason = %reason, + "guardrail blocked passthrough-route response", + ); + telemetry.guardrail_blocked = true; + // The telemetry guard has not emitted yet; drop it silently and + // let the shared error path report the 422. + telemetry.emitted = true; + return Err(RouteError::of( + crate::error::guardrail_block_error( + "response", + guardrail_name.as_deref(), + unavailable.as_deref(), + ), + &auth, + )); + } } } if let Some(u) = response_usage(protocol, raw_shape, &resp_body) { merge_usage(&mut telemetry.usage, u); } - if telemetry.content_cap.is_some() { - telemetry.response_text = response_guardrail_text(protocol, &resp_body); + if telemetry.content_cap.is_some() && !bypass_uninspectable_output { + telemetry.response_text = response_capture_text(protocol, &resp_body); } let mut response = Response::builder() @@ -1219,14 +1410,16 @@ fn anthropic_message_output_text(v: &serde_json::Value) -> String { /// The body envelope detected for one exchange. Not configuration: /// detected per request from the body's top-level keys -/// ([`detect_protocol`]) and sticky for the exchange — the buffered -/// response and every stream frame are read with the same detection. It -/// drives extraction (guardrail text, capture, usage) only; the relay +/// ([`detect_protocol`]) and sticky for request extraction, stream frames, +/// capture, and usage. Buffered output guardrails independently classify the +/// response envelope, because a provider can return a different compatible +/// envelope than the request. Detection drives extraction only; the relay /// forwards bytes verbatim regardless. #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum PassthroughProtocol { - /// No recognized envelope: bodies are opaque (guardrails scan them as - /// one lossy-UTF-8 text; buffered responses are not probed for usage). + /// No recognized envelope: guardrails scan every decoded JSON string + /// value, falling back to one lossy-UTF-8 text when the body is not JSON; + /// buffered responses are not probed for usage. /// A streamed opaque response reports usage from an explicit `usage` /// object, or — for the flat token shape agent backends use — only /// from a frame the server itself labels one (`event: token_usage`), @@ -1249,6 +1442,17 @@ enum PassthroughProtocol { OpenaiResponses, } +/// Detection preserves the safety outcome separately from the protocol hint. +/// A malformed root that merely resembles JSON cannot be treated as an opaque +/// raw request: doing so would turn a failed typed selector into a whole-body +/// scan and expose media fields to an external guardrail. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum JsonProbe { + Raw, + Protocol(PassthroughProtocol), + Unevaluable(crate::json_splice::SpliceError), +} + /// Detect the request envelope from the body's top-level keys. The three /// LLM envelopes are structurally exclusive — `messages`, `input` (a string /// or an array) and `prompt` are each the required carrier field of exactly @@ -1260,26 +1464,193 @@ enum PassthroughProtocol { /// and extraction degrades to the whole body when the detected shape /// yields no text. fn detect_protocol(body: &[u8]) -> PassthroughProtocol { - let Ok(v) = serde_json::from_slice::(body) else { - return PassthroughProtocol::Raw; - }; - if v.get("messages").is_some_and(serde_json::Value::is_array) { + // Do not materialize the entire document just to inspect its envelope: + // a valid request can exceed serde_json::Value's nesting limit in an + // unrelated forwarded field. The shallow source selector keeps the + // chosen top-level carrier bounded while preserving the last-key behavior + // of a JSON map. + if raw_top_level_last_has_shape(body, "messages", false) { PassthroughProtocol::OpenaiChat - } else if v - .get("input") - .is_some_and(|i| i.is_string() || i.is_array()) - { + } else if raw_top_level_last_has_shape(body, "input", true) { PassthroughProtocol::OpenaiResponses - } else if v - .get("prompt") - .is_some_and(|p| p.is_string() || p.is_array()) - { + } else if raw_top_level_last_has_shape(body, "prompt", true) { PassthroughProtocol::OpenaiCompletions } else { PassthroughProtocol::Raw } } +/// The concrete envelope carried by a buffered provider response. This is +/// intentionally distinct from [`detect_protocol`]: a Chat request may +/// receive a Responses body (and vice versa), so an output guardrail must +/// select visible text from what the provider actually returned. +fn resolve_buffered_response_protocol( + request_protocol: PassthroughProtocol, + body: &[u8], +) -> Result { + if matches!(request_protocol, PassthroughProtocol::Raw) { + // Raw traffic deliberately retains its broad decoded-string scan. + return Ok(PassthroughProtocol::Raw); + } + + detect_buffered_response_protocol(body) + .map_err(|_| crate::json_splice::SpliceError::unevaluable())? + .ok_or_else(crate::json_splice::SpliceError::unevaluable) +} + +/// Classify only providers' known response carriers. This uses source slices +/// rather than `Value`, so duplicate response carriers remain visible to the +/// existing selectors and opaque nested data never needs full deserialization. +fn detect_buffered_response_protocol(body: &[u8]) -> Result, ()> { + let object_label = raw_top_level_unique_string(body, "object")?; + let output = response_top_level_array_values(body, "output")? + .map(|_| PassthroughProtocol::OpenaiResponses); + let (has_choices, choices) = buffered_choices_response_protocol(body, object_label.as_deref())?; + if output.is_some() && has_choices { + // The two carriers name incompatible envelope families. Neither a + // request hint nor an empty `choices` array may choose between them. + return Err(()); + } + let object = match object_label.as_deref() { + // An object label only disambiguates an otherwise empty concrete + // carrier. A bare label could be an opaque provider extension and + // must not turn a typed fail-closed scan into an empty selector. + Some("response") if output.is_some() => Some(PassthroughProtocol::OpenaiResponses), + Some("chat.completion") if has_choices => Some(PassthroughProtocol::OpenaiChat), + Some("text_completion") if has_choices => Some(PassthroughProtocol::OpenaiCompletions), + Some("response" | "chat.completion" | "text_completion") => return Err(()), + Some(_) | None => None, + }; + let anthropic = match raw_top_level_unique_type(body)?.as_deref() { + Some("message") => response_top_level_array_values(body, "content")? + .is_some() + .then_some(PassthroughProtocol::OpenaiChat) + .ok_or(())?, + Some(_) | None => return select_buffered_response_protocol([output, choices, object]), + }; + + select_buffered_response_protocol([output, choices, object, Some(anthropic)]) +} + +fn select_buffered_response_protocol( + candidates: [Option; N], +) -> Result, ()> { + let mut selected = None; + for candidate in candidates.into_iter().flatten() { + if let Some(previous) = selected { + if previous != candidate { + return Err(()); + } + } else { + selected = Some(candidate); + } + } + Ok(selected) +} + +/// Return every occurrence of a response carrier only when all values are +/// arrays. Repeated Responses and Chat carriers are deliberately retained: +/// the downstream selector scans every client-visible occurrence. +fn response_top_level_array_values<'a>( + body: &'a [u8], + key: &str, +) -> Result>>, ()> { + let values = raw_top_level_value_refs(body, key).ok_or(())?; + if values.is_empty() { + return Ok(None); + } + if values + .iter() + .all(|value| value.get().trim_start().starts_with('[')) + { + Ok(Some(values)) + } else { + Err(()) + } +} + +/// `choices[].message` and `choices[].text` are mutually exclusive response +/// envelopes. Empty choices need an explicit root `object` hint; they cannot +/// safely inherit the request's protocol. +fn buffered_choices_response_protocol( + body: &[u8], + object_label: Option<&str>, +) -> Result<(bool, Option), ()> { + let Some(arrays) = response_top_level_array_values(body, "choices")? else { + return Ok((false, None)); + }; + let mut selected = None; + for array in &arrays { + for choice in raw_array_item_refs(array).ok_or(())? { + if !raw_is_object(&choice) { + return Err(()); + } + let choice_body = choice.get().as_bytes(); + let message = raw_top_level_unique_object(choice_body, "message")?.is_some(); + let text = raw_top_level_unique_string(choice_body, "text")?.is_some(); + let candidate = match (message, text) { + (true, false) => Some(PassthroughProtocol::OpenaiChat), + (false, true) => Some(PassthroughProtocol::OpenaiCompletions), + (false, false) => None, + // Some compatibility Chat replies retain their legacy + // `text` alongside the structured message. The response + // label is the only safe disambiguator; a bare mixed choice + // remains unevaluable rather than inheriting the request. + (true, true) => match object_label { + Some("chat.completion") => Some(PassthroughProtocol::OpenaiChat), + Some("text_completion") => Some(PassthroughProtocol::OpenaiCompletions), + _ => return Err(()), + }, + }; + if let Some(candidate) = candidate { + if let Some(previous) = selected { + if previous != candidate { + return Err(()); + } + } else { + selected = Some(candidate); + } + } + } + } + if matches!(selected, Some(PassthroughProtocol::OpenaiCompletions)) && arrays.len() != 1 { + return Err(()); + } + Ok((true, selected)) +} + +fn raw_json_container_like(body: &[u8]) -> bool { + body.iter() + .copied() + .find(|byte| !matches!(*byte, b' ' | b'\t' | b'\n' | b'\r')) + .is_some_and(|byte| matches!(byte, b'{' | b'[')) +} + +fn raw_json_scan_error( + body: &[u8], + error: crate::json_splice::SpliceError, +) -> crate::json_splice::SpliceError { + if raw_json_container_like(body) && !error.is_unevaluable() { + crate::json_splice::SpliceError::unevaluable() + } else { + error + } +} + +fn probe_json(body: &[u8]) -> JsonProbe { + let protocol = detect_protocol(body); + if !matches!(protocol, PassthroughProtocol::Raw) { + return JsonProbe::Protocol(protocol); + } + if !raw_json_container_like(body) { + return JsonProbe::Raw; + } + match crate::json_splice::validate_json(body) { + Ok(()) => JsonProbe::Raw, + Err(error) => JsonProbe::Unevaluable(raw_json_scan_error(body, error)), + } +} + /// A `Raw` request whose `model` and unary JSON response usage the gateway /// can still read. Consulted for model attribution and the BUFFERED /// response's usage only — guardrail text, capture and every stream frame @@ -1351,2357 +1722,7755 @@ fn body_model_name( .unwrap_or_default() } -/// The request text a guardrail scans, per the detected envelope. -/// Extraction is best-effort: a shape that yields no text degrades to the -/// raw lossy-UTF-8 body, so detection never loses audit coverage. -fn request_guardrail_text(protocol: PassthroughProtocol, body: &[u8]) -> String { - let raw = || String::from_utf8_lossy(body).into_owned(); - let Ok(v) = serde_json::from_slice::(body) else { - return raw(); - }; - let extracted = match protocol { - PassthroughProtocol::Raw => return raw(), - // An Anthropic Messages body carries its system prompt top-level. - PassthroughProtocol::OpenaiChat => v - .get("system") - .map(request_content_text) - .into_iter() - .chain( - v.get("messages") - .and_then(|m| m.as_array()) - .into_iter() - .flatten() - .map(|m| message_scan_text(m, true)), - ) - .filter(|t| !t.is_empty()) - .collect::>() - .join("\n"), - // Responses API: `input` is either a bare string or an array of - // items, read exactly as the typed route reads them - // (`responses::responses_item_text`) — message content, tool - // results, approval reasons, replayed reasoning summaries, and a - // replayed tool call's name, arguments and input. Every item, not - // only the common one: the raw-body fallback below fires only when - // the WHOLE extraction came back empty, so a slot left out here is - // never scanned while `/v1/responses` blocks the same body. - PassthroughProtocol::OpenaiResponses => match v.get("input") { - Some(serde_json::Value::String(t)) => t.clone(), - Some(serde_json::Value::Array(items)) => items - .iter() - .map(crate::responses::responses_item_text) - .filter(|t| !t.is_empty()) - .collect::>() - .join("\n"), - _ => String::new(), - }, - PassthroughProtocol::OpenaiCompletions => { - let prompt = v.get("prompt").map(|p| match p { - serde_json::Value::Array(items) => items - .iter() - .filter_map(|i| i.as_str()) - .collect::>() - .join("\n"), - other => content_text(other), - }); - let suffix = v.get("suffix").and_then(|s| s.as_str()); - let mut out = prompt.unwrap_or_default(); - if let Some(s) = suffix { - if !out.is_empty() { - out.push('\n'); - } - out.push_str(s); +fn append_scan_text(out: &mut String, text: &str) -> Option<()> { + if text.is_empty() { + return Some(()); + } + let separator = if out.is_empty() { 0 } else { 1 }; + let next_len = out.len().checked_add(separator)?.checked_add(text.len())?; + if next_len > MAX_RAW_SELECTOR_BYTES { + return None; + } + out.try_reserve(next_len - out.len()).ok()?; + if separator != 0 { + out.push('\n'); + } + out.push_str(text); + Some(()) +} + +fn decoded_json_string_values(body: &[u8]) -> Option { + crate::json_splice::collect_string_values(body) + .ok() + .filter(|out| !out.is_empty()) +} + +fn try_raw_json_string_values(body: &[u8]) -> Result { + crate::json_splice::collect_string_values(body) + .map_err(|error| raw_json_scan_error(body, error)) +} + +fn decoded_json_string_values_vec_where( + body: &[u8], + include: impl FnMut(&[crate::json_splice::PathSeg]) -> bool, +) -> Option> { + crate::json_splice::collect_string_values_where_vec(body, include) + .ok() + .filter(|out| !out.is_empty()) +} + +fn decoded_json_string_values_including_empty( + body: &[u8], + scan_error: &mut Option, +) -> Option { + match crate::json_splice::collect_string_values(body) { + Ok(values) => Some(values), + Err(error) => { + if scan_error.is_none() { + *scan_error = Some(error); } - out + None } - }; - if extracted.is_empty() { - raw() - } else { - extracted } } -/// The response text a guardrail scans / the capture records, per the -/// route's protocol hint. Best-effort like the request side. -fn response_guardrail_text(protocol: PassthroughProtocol, body: &[u8]) -> String { - let raw = || String::from_utf8_lossy(body).into_owned(); - let Ok(v) = serde_json::from_slice::(body) else { - return raw(); - }; - // Responses answers with `output` items, not `choices`: read them as - // the typed route does (`responses::responses_output_text`) — message - // text plus each tool call's name, arguments and input, with generated - // reasoning items left out of the output scope. - if matches!(protocol, PassthroughProtocol::OpenaiResponses) { - let joined = crate::responses::responses_output_text(&v); - return if joined.is_empty() { raw() } else { joined }; +fn mark_unevaluable(scan_error: &mut Option) { + // A malformed typed carrier is not safe for a raw-body fallback either. + // Preserve a depth error for observability, but turn every other selector + // failure into the fail-closed class used at the dispatch boundary. + if !scan_error.is_some_and(crate::json_splice::SpliceError::is_depth_exceeded) { + *scan_error = Some(crate::json_splice::SpliceError::unevaluable()); } - // An Anthropic Messages response on the chat envelope carries - // `content` blocks rather than `choices`. - if matches!(protocol, PassthroughProtocol::OpenaiChat) && v.get("choices").is_none() { - let text = anthropic_message_output_text(&v); - if !text.is_empty() { - return text; +} + +fn decoded_json_string_values_except_root_keys( + body: &[u8], + excluded: &[&str], + scan_error: &mut Option, +) -> Option { + let mut out = String::new(); + for value in raw_top_level_values_except(body, excluded)? { + match crate::json_splice::collect_string_values(value.get().as_bytes()) { + Ok(values) => append_scan_text(&mut out, &values)?, + Err(error) => { + if scan_error.is_none() { + *scan_error = Some(error); + } + return None; + } } } - let choices = match protocol { - PassthroughProtocol::Raw => return raw(), - _ => v.get("choices").and_then(|c| c.as_array()), - }; - let Some(choices) = choices else { return raw() }; - let texts: Vec = choices - .iter() - .filter_map(|c| match protocol { - PassthroughProtocol::OpenaiChat => { - c.get("message").map(|m| message_scan_text(m, false)) - } - PassthroughProtocol::OpenaiCompletions => { - c.get("text").and_then(|t| t.as_str()).map(str::to_string) - } - // Unreachable: handled above by the `output` branch. - PassthroughProtocol::OpenaiResponses | PassthroughProtocol::Raw => None, - }) - .filter(|t| !t.is_empty()) - .collect(); - if texts.is_empty() { - raw() - } else { - texts.join("\n") + Some(out) +} + +/// A source fragment selected without asking serde to recursively skip an +/// arbitrary caller-controlled value. The shallow parser below uses bounded +/// stack space while crossing opaque image/document/reasoning payloads; a +/// fragment is only walked with `json_splice` once it is eligible for a +/// guardrail scan, which is where [`crate::json_splice::MAX_JSON_DEPTH`] +/// applies. +#[derive(Clone, Copy)] +struct RawJson<'a> { + source: &'a str, +} + +impl<'a> RawJson<'a> { + fn get(self) -> &'a str { + self.source } } -/// Every token dimension a passthrough exchange can report, mirroring the -/// token fields of [`aisix_obs::UsageEvent`] 1:1 so a route reports what -/// the typed endpoint serving the same envelope would. -/// -/// Populated from the union of spellings the relayed APIs use — OpenAI's -/// nested `*_tokens_details`, the Responses API's `input`/`output` -/// spelling, Anthropic's separate cache counters, DeepSeek's native -/// `prompt_cache_hit_tokens`, and the flat token object agent backends -/// report on their own SSE event. -#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] -struct PassthroughUsage { - prompt_tokens: u32, - completion_tokens: u32, - cached_prompt_tokens: u32, - cache_write_tokens: Option, - reasoning_tokens: u32, - cache_creation_tokens: u32, - cache_read_tokens: u32, - /// The upstream's own `total_tokens`, verbatim; `None` when a report - /// carried none. Never a sum computed here — see - /// `UsageStats::upstream_total_tokens`. - upstream_total_tokens: Option, +// Source selectors retain only borrowed spans, but a wide array or repeated +// top-level carrier can still make their bookkeeping unbounded. Keep one +// shared budget for every selector collection and the nested Chat work list. +const MAX_RAW_SELECTOR_ITEMS: usize = crate::json_splice::MAX_JSON_SCAN_VALUES; +const MAX_RAW_SELECTOR_BYTES: usize = crate::json_splice::MAX_JSON_SCAN_TEXT_BYTES; +// Nested Anthropic tool-result carriers are selected from source-preserved +// spans. Cap their cumulative structural walk, not just the current work +// queue, so deeply nested valid JSON cannot make the selector re-scan every +// remaining suffix quadratically. +const MAX_NESTED_CONTENT_SELECTOR_WORK_BYTES: usize = MAX_RAW_SELECTOR_BYTES * 8; + +fn raw_selector_push<'a>( + values: &mut Vec>, + source_bytes: &mut usize, + value: RawJson<'a>, +) -> Option<()> { + if values.len() >= MAX_RAW_SELECTOR_ITEMS { + return None; + } + let next_bytes = source_bytes.checked_add(value.get().len())?; + if next_bytes > MAX_RAW_SELECTOR_BYTES { + return None; + } + values.try_reserve(1).ok()?; + *source_bytes = next_bytes; + values.push(value); + Some(()) } -impl PassthroughUsage { - /// Field-wise max, the accumulation the typed streaming paths use. - /// - /// One stream reports usage across several frames — Anthropic's - /// `message_start` carries the input and cache counters while its - /// terminal `message_delta` carries only the output ones — so a later - /// partial report must EXTEND the record rather than replace it. Max - /// also makes a provider that repeats a cumulative usage object - /// harmless. - fn merge(&mut self, other: Self) { - self.prompt_tokens = self.prompt_tokens.max(other.prompt_tokens); - self.completion_tokens = self.completion_tokens.max(other.completion_tokens); - self.cached_prompt_tokens = self.cached_prompt_tokens.max(other.cached_prompt_tokens); - self.cache_write_tokens = self.cache_write_tokens.max(other.cache_write_tokens); - self.reasoning_tokens = self.reasoning_tokens.max(other.reasoning_tokens); - self.cache_creation_tokens = self.cache_creation_tokens.max(other.cache_creation_tokens); - self.cache_read_tokens = self.cache_read_tokens.max(other.cache_read_tokens); - // A total stands only while every merged report carried one. - self.upstream_total_tokens = self - .upstream_total_tokens - .zip(other.upstream_total_tokens) - .map(|(a, b)| a.max(b)); +/// Keep a structural source reference without charging the span's bytes to a +/// text-selector budget. Callers must only use this for a carrier that they +/// subsequently narrow to a selected text field; opaque siblings must not +/// make an otherwise valid carrier unevaluable merely because they are large. +fn raw_selector_push_ref<'a>(values: &mut Vec>, value: RawJson<'a>) -> Option<()> { + if values.len() >= MAX_RAW_SELECTOR_ITEMS { + return None; } + values.try_reserve(1).ok()?; + values.push(value); + Some(()) } -/// Merge one usage report into an exchange's accumulated usage. The first -/// report is adopted whole rather than merged into zeros, so a total it -/// carried survives until a report without one arrives. -fn merge_usage(acc: &mut Option, report: PassthroughUsage) { - match acc { - Some(acc) => acc.merge(report), - None => *acc = Some(report), +fn charge_nested_content_selector_work(total: &mut usize, bytes: usize) -> bool { + let Some(next) = total.checked_add(bytes) else { + return false; + }; + if next > MAX_NESTED_CONTENT_SELECTOR_WORK_BYTES { + return false; } + *total = next; + true } -/// `usage` figures from a buffered protocol-aware response body. -fn response_usage( - protocol: PassthroughProtocol, - raw_shape: Option, - body: &[u8], -) -> Option { - if matches!(protocol, PassthroughProtocol::Raw) && raw_shape.is_none() { - return None; +fn raw_skip_ws(bytes: &[u8], pos: &mut usize) { + while bytes + .get(*pos) + .is_some_and(|byte| matches!(*byte, b' ' | b'\t' | b'\n' | b'\r')) + { + *pos += 1; } - let v = serde_json::from_slice::(body).ok()?; - match raw_shape { - Some(RawUsageShape::Rerank) => { - crate::rerank::rerank_prompt_tokens(&v).map(|prompt_tokens| PassthroughUsage { - prompt_tokens, - ..PassthroughUsage::default() - }) +} + +fn raw_string_end(bytes: &[u8], start: usize) -> Option { + (bytes.get(start) == Some(&b'"')).then_some(())?; + let mut pos = start + 1; + while let Some(&byte) = bytes.get(pos) { + match byte { + b'"' => return Some(pos + 1), + b'\\' => match bytes.get(pos + 1).copied()? { + b'"' | b'\\' | b'/' | b'b' | b'f' | b'n' | b'r' | b't' => pos += 2, + b'u' => { + let hex = bytes.get(pos + 2..pos + 6)?; + hex.iter().all(u8::is_ascii_hexdigit).then_some(())?; + pos += 6; + } + _ => return None, + }, + 0..=0x1f => return None, + _ => pos += 1, } - Some(RawUsageShape::DashscopeNative) => dashscope_native_usage(v.get("usage")?), - None => usage_of(v.get("usage")?), } + None } -/// Token counts from a DashScope native response's top-level `usage`. -/// -/// Every dimension the generic reader knows (cache hit, reasoning, …) is -/// read by [`usage_of`]; only prompt and completion follow DashScope's own -/// arithmetic. The flat `image_tokens` means different things per service: -/// multimodal generation counts it inside `input_tokens` -/// (`{input_tokens: 79, image_tokens: 66, output_tokens: 14, -/// total_tokens: 93}`), while the multimodal embedding and rerank services -/// count it beside (`{input_tokens: 44, image_tokens: 64, -/// total_tokens: 108}`). So `total_tokens` less the completion is the -/// prompt whenever a total is reported; without one (only some embedding -/// models omit it, and those report images beside the text), prompt is -/// `input_tokens` (or `prompt_tokens`) plus the flat `image_tokens`. A -/// nested `input_tokens_details.image_tokens` is a breakdown already inside -/// `input_tokens` and is never added. -fn dashscope_native_usage(usage: &serde_json::Value) -> Option { - let num = |k: &str| { - usage - .get(k) - .and_then(serde_json::Value::as_u64) - .map(|n| n.min(u32::MAX as u64) as u32) - }; - let completion = num("output_tokens").or_else(|| num("completion_tokens")); - let prompt = match num("total_tokens") { - Some(total) => Some(total.saturating_sub(completion.unwrap_or(0))), - None => { - let input = num("input_tokens").or_else(|| num("prompt_tokens")); - let image = num("image_tokens"); - (input.is_some() || image.is_some()) - .then(|| input.unwrap_or(0).saturating_add(image.unwrap_or(0))) +/// Find one raw JSON value's end without recursively deserializing it. The +/// caller enforces object/array separators around the returned span. The +/// pairing stack is capped at the guardrail traversal bound; farther opaque +/// descendants keep only a depth count so an image/document can remain +/// source-preserved without allocating one selector frame per nested value. +fn raw_value_end(bytes: &[u8], start: usize) -> Option { + match bytes.get(start).copied()? { + b'"' => raw_string_end(bytes, start), + b'{' | b'[' => { + let mut pos = start; + let mut frames = Vec::new(); + let mut opaque_depth = 0usize; + while let Some(&byte) = bytes.get(pos) { + match byte { + b'"' => pos = raw_string_end(bytes, pos)?, + b'{' | b'[' => { + if frames.len() < crate::json_splice::MAX_JSON_DEPTH { + frames.push(byte); + } else { + opaque_depth = opaque_depth.checked_add(1)?; + } + pos += 1; + } + b'}' | b']' => { + if opaque_depth > 0 { + opaque_depth -= 1; + } else { + let opener = frames.pop()?; + if !matches!((opener, byte), (b'{', b'}') | (b'[', b']')) { + return None; + } + } + pos += 1; + if frames.is_empty() && opaque_depth == 0 { + return Some(pos); + } + } + _ => pos += 1, + } + } + None + } + _ => { + let mut pos = start; + while bytes.get(pos).is_some_and(|byte| { + !matches!(*byte, b',' | b']' | b'}' | b' ' | b'\t' | b'\n' | b'\r') + }) { + pos += 1; + } + let token = bytes.get(start..pos)?; + if token == b"true" || token == b"false" || token == b"null" { + Some(pos) + } else { + crate::json_splice::is_json_number(token).then_some(pos) + } } - }; - let generic = usage_of(usage); - if prompt.is_none() && completion.is_none() && generic.is_none() { - return None; } - Some(PassthroughUsage { - prompt_tokens: prompt.unwrap_or(0), - completion_tokens: completion.unwrap_or(0), - upstream_total_tokens: num("total_tokens"), - ..generic.unwrap_or_default() +} + +/// Visit a root object's source members without `RawValue` / `IgnoredAny`. +/// The caller-facing relay keeps opaque payloads raw, and this parser must do +/// the same rather than making their nesting a guardrail traversal. +fn raw_object_members<'a>(body: &'a [u8], mut visit: impl FnMut(&str, RawJson<'a>)) -> Option<()> { + let source = std::str::from_utf8(body).ok()?; + let bytes = source.as_bytes(); + let mut pos = 0; + raw_skip_ws(bytes, &mut pos); + (bytes.get(pos) == Some(&b'{')).then_some(())?; + pos += 1; + loop { + raw_skip_ws(bytes, &mut pos); + if bytes.get(pos) == Some(&b'}') { + pos += 1; + raw_skip_ws(bytes, &mut pos); + return (pos == bytes.len()).then_some(()); + } + let key_start = pos; + let key_end = raw_string_end(bytes, key_start)?; + let key = serde_json::from_slice::(&bytes[key_start..key_end]).ok()?; + pos = key_end; + raw_skip_ws(bytes, &mut pos); + (bytes.get(pos) == Some(&b':')).then_some(())?; + pos += 1; + raw_skip_ws(bytes, &mut pos); + let value_start = pos; + let value_end = raw_value_end(bytes, value_start)?; + visit( + &key, + RawJson { + source: &source[value_start..value_end], + }, + ); + pos = value_end; + raw_skip_ws(bytes, &mut pos); + match bytes.get(pos) { + Some(b',') => pos += 1, + Some(b'}') => { + pos += 1; + raw_skip_ws(bytes, &mut pos); + return (pos == bytes.len()).then_some(()); + } + _ => return None, + } + } +} + +/// Source values of all occurrences of one top-level key. The shallow source +/// selector keeps repeated keys separate without recursively deserializing +/// unrelated forwarded fields. +fn raw_top_level_values<'a>(body: &'a [u8], wanted_key: &str) -> Option>> { + let mut values = Vec::new(); + let mut source_bytes = 0; + let mut within_cap = true; + raw_object_members(body, |key, value| { + if within_cap && key == wanted_key { + within_cap = raw_selector_push(&mut values, &mut source_bytes, value).is_some(); + } + })?; + within_cap.then_some(values) +} + +/// Retaining a source fragment is a pointer copy, so nested carrier walks do +/// not copy their remaining `tool_result.content` suffixes. +fn raw_top_level_value_refs<'a>(body: &'a [u8], wanted_key: &str) -> Option>> { + let mut values = Vec::new(); + let mut within_cap = true; + raw_object_members(body, |key, value| { + if within_cap && key == wanted_key { + within_cap = raw_selector_push_ref(&mut values, value).is_some(); + } + })?; + within_cap.then_some(values) +} + +/// Source values of top-level keys other than `excluded`. Known opaque +/// carriers are skipped before their internals are traversed for a scan. +fn raw_top_level_values_except<'a>(body: &'a [u8], excluded: &[&str]) -> Option>> { + let mut values = Vec::new(); + let mut source_bytes = 0; + let mut within_cap = true; + raw_object_members(body, |key, value| { + if within_cap && !excluded.contains(&key) { + within_cap = raw_selector_push(&mut values, &mut source_bytes, value).is_some(); + } + })?; + within_cap.then_some(values) +} + +/// Match the last source occurrence, the same duplicate-key convention a +/// materialized JSON map used before protocol detection became shallow. +/// `allow_string` is for the Responses and Completions bare-string forms; +/// Chat requires an array of messages. +fn raw_top_level_last_has_shape(body: &[u8], key: &str, allow_string: bool) -> bool { + let mut last = None; + raw_object_members(body, |candidate, value| { + if candidate == key { + last = Some(value); + } + }) + .and(last) + .is_some_and(|value| match value.get().trim_start().as_bytes().first() { + Some(b'[') => true, + Some(b'"') => allow_string, + _ => false, }) } -/// Read every token dimension out of one `usage` object (or, for the -/// labelled frame of an opaque stream, a flat token object). -/// -/// The spellings are read as a union rather than per protocol because a -/// passthrough route relays whichever API the caller addressed: the same -/// route carries an OpenAI chat envelope, an Anthropic one, and an agent -/// backend's private shape. They do not collide — each name belongs to -/// exactly one API — so reading them all costs nothing and a detected -/// envelope reports what its typed endpoint would. -/// -/// `None` when the object carries no recognised counter at all, which is -/// what keeps a `usage`-shaped object that is not a usage report from -/// minting zeros. -fn usage_of(usage: &serde_json::Value) -> Option { - let num = |v: Option<&serde_json::Value>| { - v.and_then(serde_json::Value::as_u64) - .map(|n| n.min(u32::MAX as u64) as u32) +fn raw_is_object(raw: &RawJson<'_>) -> bool { + raw.get().trim_start().starts_with('{') +} + +fn raw_array_items<'a>(raw: &RawJson<'a>) -> Option>> { + let source = raw.get(); + let bytes = source.as_bytes(); + let mut pos = 0; + raw_skip_ws(bytes, &mut pos); + (bytes.get(pos) == Some(&b'[')).then_some(())?; + pos += 1; + let mut values = Vec::new(); + let mut source_bytes = 0; + loop { + raw_skip_ws(bytes, &mut pos); + if bytes.get(pos) == Some(&b']') { + pos += 1; + raw_skip_ws(bytes, &mut pos); + return (pos == bytes.len()).then_some(values); + } + let value_start = pos; + let value_end = raw_value_end(bytes, value_start)?; + raw_selector_push( + &mut values, + &mut source_bytes, + RawJson { + source: &source[value_start..value_end], + }, + )?; + pos = value_end; + raw_skip_ws(bytes, &mut pos); + match bytes.get(pos) { + Some(b',') => pos += 1, + Some(b']') => { + pos += 1; + raw_skip_ws(bytes, &mut pos); + return (pos == bytes.len()).then_some(values); + } + _ => return None, + } + } +} + +fn raw_array_item_refs<'a>(raw: &RawJson<'a>) -> Option>> { + let source = raw.get(); + let bytes = source.as_bytes(); + let mut pos = 0; + raw_skip_ws(bytes, &mut pos); + (bytes.get(pos) == Some(&b'[')).then_some(())?; + pos += 1; + let mut values = Vec::new(); + loop { + raw_skip_ws(bytes, &mut pos); + if bytes.get(pos) == Some(&b']') { + pos += 1; + raw_skip_ws(bytes, &mut pos); + return (pos == bytes.len()).then_some(values); + } + let value_start = pos; + let value_end = raw_value_end(bytes, value_start)?; + raw_selector_push_ref( + &mut values, + RawJson { + source: &source[value_start..value_end], + }, + )?; + pos = value_end; + raw_skip_ws(bytes, &mut pos); + match bytes.get(pos) { + Some(b',') => pos += 1, + Some(b']') => { + pos += 1; + raw_skip_ws(bytes, &mut pos); + return (pos == bytes.len()).then_some(values); + } + _ => return None, + } + } +} + +/// `true` only for an unambiguous typed item. Conflicting or non-string +/// duplicate `type` fields stay in the output scan rather than becoming a +/// way to hide content. +fn raw_object_has_only_types(raw: &RawJson<'_>, allowed: &[&str]) -> bool { + let Some(values) = raw_top_level_values(raw.get().as_bytes(), "type") else { + return false; }; - // Flat counter under any of `names`, first hit wins. - let flat = |names: &[&str]| names.iter().find_map(|n| num(usage.get(*n))); - // `parent.child` counter, e.g. `prompt_tokens_details.cached_tokens`. - let nested = |parent: &str, child: &str| num(usage.get(parent).and_then(|d| d.get(child))); + let mut types = values + .into_iter() + .map(|value| serde_json::from_str::(value.get()).ok()); + let Some(Some(first)) = types.next() else { + return false; + }; + allowed.iter().any(|allowed| first == *allowed) + && types.all(|kind| kind.as_deref() == Some(first.as_str())) +} - let prompt = flat(&["prompt_tokens", "input_tokens"]); - let completion = flat(&["completion_tokens", "output_tokens"]); - // OpenAI nests the cache hit under `prompt_tokens_details`, the - // Responses API under `input_tokens_details`, DeepSeek reports it flat - // as `prompt_cache_hit_tokens`. A nested ZERO must not mask a real - // native count (the typed OpenAI bridge takes the same precedence). - let cached_prompt = nested("prompt_tokens_details", "cached_tokens") - .filter(|&n| n > 0) - .or_else(|| nested("input_tokens_details", "cached_tokens").filter(|&n| n > 0)) - .or_else(|| flat(&["prompt_cache_hit_tokens", "cached_tokens"])); - let cache_write = nested("prompt_tokens_details", "cache_write_tokens") - .or_else(|| nested("input_tokens_details", "cache_write_tokens")); - let reasoning = nested("completion_tokens_details", "reasoning_tokens") - .filter(|&n| n > 0) - .or_else(|| nested("output_tokens_details", "reasoning_tokens").filter(|&n| n > 0)) - .or_else(|| flat(&["reasoning_tokens"])); - // Anthropic's two cache counters sit beside `input_tokens`, and are - // ADDITIVE to it rather than a subset. - let cache_creation = flat(&["cache_creation_input_tokens", "cache_creation_tokens"]); - let cache_read = flat(&["cache_read_input_tokens", "cache_read_tokens"]); +fn raw_top_level_unique_type(body: &[u8]) -> Result, ()> { + let mut types = raw_top_level_values(body, "type") + .ok_or(())? + .into_iter() + .map(|value| serde_json::from_str::(value.get()).ok()); + let Some(Some(first)) = types.next() else { + return Ok(None); + }; + Ok(types + .all(|kind| kind.as_deref() == Some(first.as_str())) + .then_some(first)) +} - let dims = [ - prompt, - completion, - cached_prompt, - cache_write, - reasoning, - cache_creation, - cache_read, - ]; - if dims.iter().all(Option::is_none) { - return None; +fn raw_top_level_items_have_only_types(body: &[u8], key: &str, allowed: &[&str]) -> Option { + let values = raw_top_level_values(body, key)?; + Some( + !values.is_empty() + && values + .iter() + .all(|value| raw_object_has_only_types(value, allowed)), + ) +} + +fn append_raw_string_value(out: &mut String, raw: &RawJson<'_>) -> Option<()> { + append_scan_text(out, &serde_json::from_str::(raw.get()).ok()?)?; + Some(()) +} + +fn append_raw_top_level_strings(out: &mut String, body: &[u8], key: &str) -> Option<()> { + for value in raw_top_level_values(body, key)? { + // A valid but wrongly typed nominal text field must not make its + // opaque object/array sibling content eligible for a whole-body raw + // fallback at the guardrail boundary. + if let Ok(value) = serde_json::from_str::(value.get()) { + append_scan_text(out, &value)?; + } } - Some(PassthroughUsage { - prompt_tokens: prompt.unwrap_or(0), - completion_tokens: completion.unwrap_or(0), - cached_prompt_tokens: cached_prompt.unwrap_or(0), - cache_write_tokens: cache_write, - reasoning_tokens: reasoning.unwrap_or(0), - cache_creation_tokens: cache_creation.unwrap_or(0), - cache_read_tokens: cache_read.unwrap_or(0), - upstream_total_tokens: flat(&["total_tokens"]), - }) + Some(()) } -/// Model-level rate-limit identity from the JSON body's top-level `model` -/// field, scoped to `provider_lower` — the #805 contract carried over from -/// the removed implicit tunnel: `display_name` exact hit first, then the -/// provider-native `model_name` (deterministic on ties, wildcards -/// excluded), with the reservation keyed by `display_name` so route and -/// typed traffic to the same Model draw from one bucket. `None` for -/// non-JSON bodies, absent/unregistered names, or cross-provider names — -/// the request then reserves only the caller-level layers. -fn body_model_rate_limit( - snapshot: &aisix_core::AisixSnapshot, - provider_lower: &str, +fn append_typed_top_level_strings( + out: &mut String, body: &[u8], -) -> Option { - #[derive(serde::Deserialize)] - struct BodyModelProbe { - model: Option, + key: &str, + scan_error: &mut Option, +) -> Option<()> { + let Some(values) = raw_top_level_values(body, key) else { + mark_unevaluable(scan_error); + return None; + }; + for value in values { + let Ok(value) = serde_json::from_str::(value.get()) else { + mark_unevaluable(scan_error); + return None; + }; + append_scan_text(out, &value)?; } - let name = serde_json::from_slice::(body).ok()?.model?; - let matches_provider = |m: &aisix_core::Model| { - m.provider - .as_deref() - .is_some_and(|p| p.eq_ignore_ascii_case(provider_lower)) - }; - let entry = snapshot - .models - .get_by_name(&name) - .filter(|e| matches_provider(&e.value)) - .or_else(|| { - snapshot - .models - .entries() - .into_iter() - .filter(|e| { - matches_provider(&e.value) - && e.value.model_name.as_deref() == Some(name.as_str()) - && !e.value.display_name.contains('*') - }) - .min_by_key(|e| e.id.clone()) - })?; - Some(crate::quota::ModelRateLimit::from_model( - &entry.value.display_name, - &entry.id, - &entry.value, - )) + Some(()) } -/// `true` if `seg` is a strict api-version path component matching `v\d+`. -fn is_api_version_segment(seg: &str) -> bool { - seg.starts_with('v') && seg.len() > 1 && seg[1..].chars().all(|c| c.is_ascii_digit()) +fn raw_top_level_string_values(body: &[u8], key: &str) -> Option> { + let values = raw_top_level_values(body, key)?; + let mut strings = Vec::new(); + strings.try_reserve(values.len()).ok()?; + for value in values { + strings.push(serde_json::from_str::(value.get()).ok()?); + } + Some(strings) } -/// Strip one leading api-version segment from `rest` when it exactly -/// matches the trailing version segment of `base` (#164): an operator's -/// `target_url` ending in `/v1` joined with a caller path starting `v1/` -/// would otherwise produce `/v1/v1/...`. -fn strip_redundant_version_segment<'a>(base: &str, rest: &'a str) -> &'a str { - let base_tail = base.rsplit('/').next().unwrap_or(""); - if !is_api_version_segment(base_tail) { - return rest; +fn raw_top_level_unique_string(body: &[u8], key: &str) -> Result, ()> { + let mut values = raw_top_level_values(body, key).ok_or(())?; + match values.len() { + 0 => Ok(None), + 1 => serde_json::from_str(values.pop().expect("one value").get()) + .map(Some) + .map_err(|_| ()), + _ => Err(()), } - if let Some(remainder) = rest.strip_prefix(base_tail) { - if remainder.is_empty() { - return remainder; - } - if let Some(after_slash) = remainder.strip_prefix('/') { - return after_slash; +} + +fn raw_top_level_unique_index(body: &[u8], key: &str) -> Result, ()> { + let mut values = raw_top_level_values(body, key).ok_or(())?; + match values.len() { + 0 => Ok(None), + 1 => serde_json::from_str(values.pop().expect("one value").get()) + .map(Some) + .map_err(|_| ()), + _ => Err(()), + } +} + +fn raw_top_level_unique_object<'a>(body: &'a [u8], key: &str) -> Result>, ()> { + let mut values = raw_top_level_values(body, key).ok_or(())?; + match values.len() { + 0 => Ok(None), + 1 => { + let value = values.pop().expect("one value"); + raw_is_object(&value).then_some(value).map(Some).ok_or(()) } + _ => Err(()), } - rest } -// --------------------------------------------------------------------------- -// Streaming relay -// --------------------------------------------------------------------------- +fn raw_top_level_unique_array<'a>(body: &'a [u8], key: &str) -> Result>, ()> { + let mut values = raw_top_level_values(body, key).ok_or(())?; + match values.len() { + 0 => Ok(None), + 1 => { + let value = values.pop().expect("one value"); + value + .get() + .trim_start() + .starts_with('[') + .then_some(value) + .map(Some) + .ok_or(()) + } + _ => Err(()), + } +} -/// Cap on bytes buffered while waiting for one SSE frame terminator, and on -/// bytes held back by the `Window` policy while its char threshold has not -/// been reached. Both accumulators would otherwise grow without bound on an -/// upstream that never terminates a frame (or streams only delta-free -/// frames) — and a streaming route carries no reqwest-level timeout to end -/// the read. On overflow the oversized run is handed on as if it were a -/// complete frame (splitter) or force-scanned (window), so memory stays -/// bounded while the policy semantics degrade gracefully. -const MAX_HELD_STREAM_BYTES: usize = 1024 * 1024; +/// Like [`raw_top_level_unique_array`], but retains the carrier by reference +/// so a large opaque extension inside it cannot exhaust the selected-text +/// budget before the selector reaches the text field. +fn raw_top_level_unique_array_ref<'a>( + body: &'a [u8], + key: &str, +) -> Result>, ()> { + let mut values = raw_top_level_value_refs(body, key).ok_or(())?; + match values.len() { + 0 => Ok(None), + 1 => { + let value = values.pop().expect("one value"); + value + .get() + .trim_start() + .starts_with('[') + .then_some(value) + .map(Some) + .ok_or(()) + } + _ => Err(()), + } +} -/// The gateway's shared frame splitter, with this relay's overflow policy: -/// an oversized unterminated run is handed on as a frame. Bytes after the -/// last complete frame stay buffered until more arrive; `take_rest` drains -/// them at end-of-stream. -struct SseFrameSplitter(aisix_gateway::sse::SseFrameSplitter); +/// The typed content extractors inspect a bare string or the direct `text` +/// field of typed parts. Keep that boundary when walking raw source, so image +/// and document payloads never reach external guardrails as text. +fn append_raw_text_value( + out: &mut String, + raw: &RawJson<'_>, + strict_content_parts: bool, + scan_error: &mut Option, +) -> Option<()> { + let value = raw.get().trim_start(); + if value.starts_with('"') { + return append_raw_string_value(out, raw); + } + if !value.starts_with('[') { + return Some(()); + } + for part in raw_array_items(raw)? { + if !raw_is_object(&part) { + if strict_content_parts { + mark_unevaluable(scan_error); + return None; + } + continue; + } + let part_body = part.get().as_bytes(); + let kind = match raw_top_level_unique_type(part_body) { + Ok(kind) => kind, + Err(()) => { + mark_unevaluable(scan_error); + return None; + } + }; + if strict_content_parts && matches!(kind.as_deref(), Some("input_text" | "text")) { + append_typed_top_level_strings(out, part_body, "text", scan_error)?; + } else { + append_raw_top_level_strings(out, part_body, "text")?; + } + } + Some(()) +} -impl SseFrameSplitter { - fn new() -> Self { - Self(aisix_gateway::sse::SseFrameSplitter::new( - MAX_HELD_STREAM_BYTES, - )) +/// Source-preserving request text from Anthropic-compatible content blocks. +/// It mirrors [`request_content_text`]: `text`, nested `tool_result`, a +/// `tool_use` input, and plaintext `thinking` are input; image/document and +/// signed `redacted_thinking` payloads are deliberately opaque. An ambiguous +/// duplicate `type` is scanned as source rather than becoming a bypass. +fn append_chat_request_content_strings( + out: &mut String, + content: &RawJson<'_>, + scan_error: &mut Option, +) -> Option<()> { + enum Work<'a> { + Content { value: RawJson<'a>, depth: usize }, + Block { value: RawJson<'a>, depth: usize }, } - fn push(&mut self, chunk: &[u8]) -> Vec> { - self.0.push(chunk); - let mut frames = Vec::new(); - loop { - match self.0.next_frame() { - Ok(Some(frame)) => frames.push(frame), - Ok(None) => break, - // Frame-terminator starvation: hand the oversized run on as-is - // rather than buffering without bound. - Err(_) => { - frames.push(self.0.take_rest()); - break; + // `tool_result.content` can itself contain another `tool_result`. Keep + // that caller-controlled nesting off the Rust call stack. This measures + // the content-carrier depth, while a wide valid array stays valid just as + // it does for the byte scanner's frame stack. A carrier at the cap still + // scans; only one more nested carrier is unevaluable. + let mut work = Vec::new(); + let mut selector_work_bytes = 0; + let initial_bytes = content.get().len(); + if initial_bytes > MAX_RAW_SELECTOR_BYTES || work.try_reserve(1).is_err() { + mark_unevaluable(scan_error); + return None; + } + work.push(( + Work::Content { + value: *content, + depth: 0, + }, + initial_bytes, + )); + let mut work_bytes = initial_bytes; + while let Some((work_item, item_bytes)) = work.pop() { + work_bytes = work_bytes.saturating_sub(item_bytes); + match work_item { + Work::Content { value, depth } => { + let value_text = value.get().trim_start(); + if value_text.starts_with('"') { + append_raw_string_value(out, &value)?; + continue; + } + if !value_text.starts_with('[') { + // `null` is the normal empty assistant-content shape; + // every other non-string/non-array carrier cannot be + // selected without treating arbitrary source as text. + if value_text != "null" { + mark_unevaluable(scan_error); + return None; + } + continue; + } + // Push backwards so the LIFO work stack preserves the + // previous depth-first, source-order traversal. + if !charge_nested_content_selector_work(&mut selector_work_bytes, value.get().len()) + { + mark_unevaluable(scan_error); + return None; + } + let Some(blocks) = raw_array_item_refs(&value) else { + mark_unevaluable(scan_error); + return None; + }; + for block in blocks.into_iter().rev() { + let block_bytes = block.get().len(); + let Some(next_bytes) = work_bytes.checked_add(block_bytes) else { + mark_unevaluable(scan_error); + return None; + }; + if work.len() >= MAX_RAW_SELECTOR_ITEMS + || next_bytes > MAX_RAW_SELECTOR_BYTES + || work.try_reserve(1).is_err() + { + mark_unevaluable(scan_error); + return None; + } + work_bytes = next_bytes; + work.push(( + Work::Block { + value: block, + depth, + }, + block_bytes, + )); + } + } + Work::Block { + value: block, + depth, + } => { + if !raw_is_object(&block) { + mark_unevaluable(scan_error); + return None; + } + let block_body = block.get().as_bytes(); + if !charge_nested_content_selector_work(&mut selector_work_bytes, block_body.len()) + { + mark_unevaluable(scan_error); + return None; + } + let Some(types) = raw_top_level_values(block_body, "type") else { + mark_unevaluable(scan_error); + return None; + }; + if !charge_nested_content_selector_work(&mut selector_work_bytes, block_body.len()) + { + mark_unevaluable(scan_error); + return None; + } + let kind = match raw_top_level_unique_type(block_body) { + Ok(kind) => kind, + Err(()) => { + mark_unevaluable(scan_error); + return None; + } + }; + if !types.is_empty() && kind.is_none() { + append_scan_text( + out, + &decoded_json_string_values_including_empty(block_body, scan_error)?, + )?; + continue; + } + match kind.as_deref() { + Some("redacted_thinking") => {} + Some("tool_result") => { + if !charge_nested_content_selector_work( + &mut selector_work_bytes, + block_body.len(), + ) { + mark_unevaluable(scan_error); + return None; + } + let Some(nested) = raw_top_level_value_refs(block_body, "content") else { + mark_unevaluable(scan_error); + return None; + }; + if nested.is_empty() { + continue; + } + let Some(depth) = depth + .checked_add(1) + .filter(|depth| *depth <= crate::json_splice::MAX_JSON_DEPTH) + else { + if scan_error.is_none() { + *scan_error = + Some(crate::json_splice::SpliceError::depth_exceeded()); + } + return None; + }; + for nested in nested.into_iter().rev() { + let nested_bytes = nested.get().len(); + let Some(next_bytes) = work_bytes.checked_add(nested_bytes) else { + mark_unevaluable(scan_error); + return None; + }; + if work.len() >= MAX_RAW_SELECTOR_ITEMS + || next_bytes > MAX_RAW_SELECTOR_BYTES + || work.try_reserve(1).is_err() + { + mark_unevaluable(scan_error); + return None; + } + work_bytes = next_bytes; + work.push(( + Work::Content { + value: nested, + depth, + }, + nested_bytes, + )); + } + } + Some("tool_use") => { + if !charge_nested_content_selector_work( + &mut selector_work_bytes, + block_body.len(), + ) { + mark_unevaluable(scan_error); + return None; + } + let Some(inputs) = raw_top_level_values(block_body, "input") else { + mark_unevaluable(scan_error); + return None; + }; + for input in inputs { + append_scan_text( + out, + &decoded_json_string_values_including_empty( + input.get().as_bytes(), + scan_error, + )?, + )?; + } + } + Some("thinking") => { + append_typed_top_level_strings(out, block_body, "thinking", scan_error)? + } + Some("text") => { + append_typed_top_level_strings(out, block_body, "text", scan_error)? + } + _ => append_raw_top_level_strings(out, block_body, "text")?, } } } - frames } + Some(()) +} - fn take_rest(&mut self) -> Vec { - self.0.take_rest() +fn append_chat_request_message_strings( + out: &mut String, + message: &RawJson<'_>, + scan_error: &mut Option, +) -> Option<()> { + if !raw_is_object(message) { + mark_unevaluable(scan_error); + return None; + } + let message_body = message.get().as_bytes(); + append_scan_text( + out, + &decoded_json_string_values_except_root_keys( + message_body, + &["content", "tool_calls", "reasoning_content", "reasoning"], + scan_error, + )?, + )?; + let Some(contents) = raw_top_level_values(message_body, "content") else { + mark_unevaluable(scan_error); + return None; + }; + for content in contents { + append_chat_request_content_strings(out, &content, scan_error)?; + } + let Some(tool_calls) = raw_top_level_values(message_body, "tool_calls") else { + mark_unevaluable(scan_error); + return None; + }; + for tool_calls in tool_calls { + append_scan_text( + out, + &decoded_json_string_values_including_empty(tool_calls.get().as_bytes(), scan_error)?, + )?; } + append_raw_top_level_strings(out, message_body, "reasoning_content")?; + Some(()) } -/// `true` when the frame carries an `event:` line the server itself names -/// a usage report. The only evidence an opaque stream offers that a -/// token-shaped payload IS usage — see [`frame_delta`]. -fn is_usage_labelled_frame(frame: &[u8]) -> bool { - aisix_gateway::sse::lines(frame).any(|line| { - aisix_gateway::sse::parse_field(&frame[line.start..line.end]).is_some_and( - |(name, value)| { - let value = value.trim_ascii(); - name == b"event" - && (value.eq_ignore_ascii_case(b"token_usage") - || value.eq_ignore_ascii_case(b"usage")) - }, - ) - }) +fn decoded_chat_request_string_values( + body: &[u8], + scan_error: &mut Option, +) -> Option { + let mut out = decoded_json_string_values_except_root_keys( + body, + &["model", "system", "messages"], + scan_error, + )?; + let Some(system_values) = raw_top_level_values(body, "system") else { + mark_unevaluable(scan_error); + return None; + }; + for system in system_values { + append_chat_request_content_strings(&mut out, &system, scan_error)?; + } + let Some(message_values) = raw_top_level_values(body, "messages") else { + mark_unevaluable(scan_error); + return None; + }; + for array in message_values { + // The selected (last) carrier made this a Chat envelope. Preserve + // other duplicate source values without turning a malformed earlier + // carrier into a whole-body fallback that exposes opaque media. + let Some(messages) = raw_array_items(&array) else { + mark_unevaluable(scan_error); + return None; + }; + for message in messages { + append_chat_request_message_strings(&mut out, &message, scan_error)?; + } + } + Some(out) } -/// Text a guardrail scans from one SSE frame, per the protocol hint, plus -/// a usage probe on the same parsed payload. -/// -/// Usage is read from, in order of specificity: -/// -/// - the payload's own top-level `usage` object (every OpenAI-shape -/// stream, and Anthropic's terminal `message_delta`); -/// - `message.usage` on an Anthropic `message_start`, which is where the -/// input and cache counters arrive — its `message_delta` reports only -/// the output ones, so reading just the top level loses the prompt side -/// of every Anthropic stream; -/// - `response.usage` on a Responses stream's terminal event; -/// - for an OPAQUE (`Raw`) stream only, a FLAT token object on a frame the -/// server labelled a usage report (`event: token_usage`). An opaque -/// stream has no envelope to authenticate a payload against, so the -/// server's own label is the evidence — a payload that merely happens to -/// carry token-shaped fields must never mint billed tokens. -/// -/// Frames accumulate field-wise (see [`PassthroughUsage::merge`]) at the -/// call site, so a partial report never truncates an earlier one. -/// The upstream failure an SSE frame reports in-band, read with the same -/// mappings the typed endpoints use for the protocol the route carries. An -/// opaque (`Raw`) stream has no error envelope the gateway could recognise. -fn frame_in_band_error( - protocol: PassthroughProtocol, - frame: &[u8], -) -> Option { - if matches!(protocol, PassthroughProtocol::Raw) { +fn append_responses_item_strings( + out: &mut String, + item: &RawJson<'_>, + scan_error: &mut Option, +) -> Option<()> { + if !raw_is_object(item) { + mark_unevaluable(scan_error); return None; } - let payload = crate::redact::frame_payload(frame)?; - let payload = payload.trim(); - let value = serde_json::from_str::(payload).ok()?; + let item_body = item.get().as_bytes(); + let text_keys = [ + "content", + "output", + "reason", + "summary", + "name", + "arguments", + "input", + ]; + append_scan_text( + out, + &decoded_json_string_values_except_root_keys(item_body, &text_keys, scan_error)?, + )?; + for key in text_keys { + for value in raw_top_level_values(item_body, key)? { + append_raw_text_value(out, &value, key == "content", scan_error)?; + } + } + Some(()) +} + +fn decoded_responses_request_string_values( + body: &[u8], + scan_error: &mut Option, +) -> Option { + let mut out = + decoded_json_string_values_except_root_keys(body, &["model", "input"], scan_error)?; + for input in raw_top_level_values(body, "input")? { + let value = input.get().trim_start(); + if value.starts_with('"') { + append_raw_string_value(&mut out, &input)?; + } else if value.starts_with('[') { + for item in raw_array_items(&input)? { + append_responses_item_strings(&mut out, &item, scan_error)?; + } + } else { + mark_unevaluable(scan_error); + return None; + } + } + Some(out) +} + +fn decoded_completions_request_string_values( + body: &[u8], + scan_error: &mut Option, +) -> Option { + let mut out = decoded_json_string_values_except_root_keys( + body, + &["model", "prompt", "suffix"], + scan_error, + )?; + for prompt in raw_top_level_values(body, "prompt")? { + let value = prompt.get().trim_start(); + if value.starts_with('"') { + append_raw_string_value(&mut out, &prompt)?; + } else if value.starts_with('[') { + for part in raw_array_items(&prompt)? { + if part.get().trim_start().starts_with('"') { + append_raw_string_value(&mut out, &part)?; + } + } + } + } + append_raw_top_level_strings(&mut out, body, "suffix")?; + Some(out) +} + +/// The request text a guardrail scans, per the detected envelope. +/// +/// The route relays source bytes verbatim, while `serde_json::Value` drops +/// duplicate keys and stops at its default nesting limit. Scan decoded source +/// values while preserving typed opaque boundaries: signed Anthropic +/// `redacted_thinking`, image, and document payloads are not caller text. +fn request_guardrail_text_with_scan_error( + protocol: PassthroughProtocol, + body: &[u8], + scan_error: &mut Option, +) -> String { + let raw = || String::from_utf8_lossy(body).into_owned(); match protocol { - PassthroughProtocol::Raw => None, - PassthroughProtocol::OpenaiResponses => crate::responses::responses_in_band_error(&value), - // The chat envelope carries Anthropic Messages traffic too, whose - // in-band failure is a `type: "error"` event. - PassthroughProtocol::OpenaiChat | PassthroughProtocol::OpenaiCompletions => { - if value.get("type").and_then(|t| t.as_str()) == Some("error") { - if let Some(body) = value.get("error").and_then(|e| { - serde_json::from_value::< - aisix_provider_anthropic::wire::AnthropicStreamErrorBody, - >(e.clone()) - .ok() - }) { - return Some( - aisix_provider_anthropic::wire::stream_error_into_bridge_error(&body), - ); + PassthroughProtocol::Raw => decoded_json_string_values(body).unwrap_or_else(raw), + PassthroughProtocol::OpenaiChat => { + // A malformed Chat carrier is unevaluable rather than a reason + // to send its opaque media source through an input guardrail. + match decoded_chat_request_string_values(body, scan_error) { + Some(text) => text, + None => { + mark_unevaluable(scan_error); + String::new() + } + } + } + PassthroughProtocol::OpenaiCompletions => { + match decoded_completions_request_string_values(body, scan_error) { + Some(text) => text, + None => { + mark_unevaluable(scan_error); + String::new() + } + } + } + PassthroughProtocol::OpenaiResponses => { + match decoded_responses_request_string_values(body, scan_error) { + Some(text) => text, + None => { + mark_unevaluable(scan_error); + String::new() } } - aisix_gateway::capture_in_band_error(payload, aisix_gateway::UpstreamWire::OpenAI) } } } -#[cfg(test)] -fn frame_delta(protocol: PassthroughProtocol, frame: &[u8]) -> (String, Option) { - let (parts, usage) = frame_parts(protocol, frame); - (parts.scan, usage) +fn request_guardrail_text(protocol: PassthroughProtocol, body: &[u8]) -> String { + let mut ignored_scan_error = None; + request_guardrail_text_with_scan_error(protocol, body, &mut ignored_scan_error) } -/// One frame's generated content, split by [`crate::held_content::Parts`] -/// into what the output guardrails scan and what only counts toward the -/// hold-back cap, plus any usage it reports. The scan and the cap read the -/// same extraction, so they cannot disagree about a frame (#513). -fn frame_parts( +/// Like [`request_guardrail_text`], but preserves scanner failures from the +/// exact typed source selectors that read a value for input inspection. +fn try_request_guardrail_text( protocol: PassthroughProtocol, - frame: &[u8], -) -> (crate::held_content::Parts, Option) { - let usage_labelled = - matches!(protocol, PassthroughProtocol::Raw) && is_usage_labelled_frame(frame); - let mut parts = crate::held_content::Parts::default(); - let mut usage: Option = None; - let mut merge = |found: PassthroughUsage| { - merge_usage(&mut usage, found); - }; - // ONE read and ONE parse per frame: a payload spread over several - // `data:` lines is one document joined with `\n`, so parsing each line - // independently produced N unparseable fragments — no usage read, and - // on a `Raw` stream the JSON source text pushed into the guardrail - // scan instead of the values (#1100). `frame_payload` also strips the - // per-line `\r` a CRLF-framed upstream leaves behind, and returns - // `None` for a comment-only frame (`: OPENROUTER PROCESSING`). - 'payload: { - let Some(payload) = crate::redact::frame_payload(frame) else { - break 'payload; - }; - let payload = payload.trim(); - if payload.is_empty() || payload == "[DONE]" { - break 'payload; + body: &[u8], +) -> Result { + if matches!(protocol, PassthroughProtocol::Raw) { + return try_raw_json_string_values(body); + } + let mut scan_error = None; + let text = request_guardrail_text_with_scan_error(protocol, body, &mut scan_error); + scan_error.map_or(Ok(text), Err) +} + +/// Return the string field which an explicitly typed Chat content part +/// exposes to the client. Image, audio, file, and future part types stay +/// opaque at the external output-guardrail boundary. +fn chat_visible_content_part_field(kind: &str) -> Option<&'static str> { + match kind { + "text" => Some("text"), + "refusal" => Some("refusal"), + _ => None, + } +} + +/// Append a Chat `content` value that is known to be client-visible output. +/// A bare string is the Chat response's ordinary text shape; array entries +/// need one unambiguous known discriminator before their text or refusal +/// field may cross the output-guardrail boundary. +fn append_chat_visible_content_strings(out: &mut String, content: &RawJson<'_>) -> Option<()> { + let value = content.get().trim_start(); + if value.starts_with('"') { + return append_raw_string_value(out, content); + } + if !value.starts_with('[') { + return Some(()); + } + for part in raw_array_items(content)? { + if !raw_is_object(&part) { + continue; } - let Ok(v) = serde_json::from_str::(payload) else { - // Unparseable joined payload — a non-conformant upstream that - // put two independent JSON documents on two `data:` lines, say. - // The frame is still FORWARDED, so scanning nothing here is a - // way past an output block rule. Fall back to the raw payload - // text on every protocol, not just `Raw`: over-scanning can only - // produce a false positive, while under-scanning a frame the - // client receives is the bypass. (Per-line parsing used to catch - // the two-document case incidentally; this covers it and every - // other shape that does not parse.) - parts.scan.push_str(payload); - break 'payload; - }; - if let Some(u) = v.get("usage").and_then(usage_of) { - merge(u); + let part_body = part.get().as_bytes(); + if let Some(field) = raw_top_level_unique_type(part_body) + .ok()? + .as_deref() + .and_then(chat_visible_content_part_field) + { + append_raw_top_level_strings(out, part_body, field)?; } - // Anthropic opens its stream with the prompt + cache counters - // nested on `message_start`. Gated on the event type so no other - // envelope's `message` object can be read as usage. - if v.get("type").and_then(|t| t.as_str()) == Some("message_start") { - if let Some(u) = v - .get("message") - .and_then(|m| m.get("usage")) - .and_then(usage_of) - { - merge(u); + } + Some(()) +} + +/// The stream can omit a tool-call's discriminator after its first delta. +/// Accept that continuation shape, but never let an explicit unknown or +/// conflicting type borrow a function/custom field as visible tool text. +fn chat_tool_continuation_fields(body: &[u8]) -> Option<&'static [(&'static str, &'static str)]> { + let types = raw_top_level_values(body, "type")?; + match raw_top_level_unique_type(body).ok()?.as_deref() { + Some("function") => Some(&[("function", "arguments")]), + Some("custom") => Some(&[("custom", "input")]), + None if types.is_empty() => Some(&[("function", "arguments"), ("custom", "input")]), + _ => Some(&[]), + } +} + +fn append_chat_tool_call_strings(out: &mut String, tool_calls: &RawJson<'_>) -> Option<()> { + if !tool_calls.get().trim_start().starts_with('[') { + return Some(()); + } + for tool_call in raw_array_items(tool_calls)? { + if !raw_is_object(&tool_call) { + continue; + } + let tool_body = tool_call.get().as_bytes(); + for (container, field) in chat_tool_continuation_fields(tool_body)? { + for payload in raw_top_level_values(tool_body, container)? { + if !raw_is_object(&payload) { + continue; + } + let payload_body = payload.get().as_bytes(); + append_raw_top_level_strings(out, payload_body, "name")?; + append_raw_top_level_strings(out, payload_body, field)?; } } - if matches!(protocol, PassthroughProtocol::OpenaiResponses) { - // Responses streams carry usage on the terminal - // `response.completed` event's embedded response object. Read - // that shape ONLY here: another protocol's frame that happens - // to nest `response.usage` must not be read as usage. - if let Some(u) = v - .get("response") - .and_then(|r| r.get("usage")) - .and_then(usage_of) - { - merge(u); + } + Some(()) +} + +/// Preserve the pre-`tool_calls` Chat tool shape without widening the +/// response walk beyond its explicit `name` and `arguments` fields. +fn append_chat_legacy_function_call_strings( + out: &mut String, + function_call: &RawJson<'_>, +) -> Option<()> { + if !raw_is_object(function_call) { + return Some(()); + } + let function_body = function_call.get().as_bytes(); + append_raw_top_level_strings(out, function_body, "name")?; + append_raw_top_level_strings(out, function_body, "arguments")?; + Some(()) +} + +fn append_chat_output_message_strings(out: &mut String, message: &RawJson<'_>) -> Option<()> { + if !raw_is_object(message) { + return Some(()); + } + let message_body = message.get().as_bytes(); + for content in raw_top_level_values(message_body, "content")? { + append_chat_visible_content_strings(out, &content)?; + } + // OpenAI Chat also exposes a refusal as a direct message member rather + // than a typed content part. + append_raw_top_level_strings(out, message_body, "refusal")?; + for tool_calls in raw_top_level_values(message_body, "tool_calls")? { + append_chat_tool_call_strings(out, &tool_calls)?; + } + for function_call in raw_top_level_values(message_body, "function_call")? { + append_chat_legacy_function_call_strings(out, &function_call)?; + } + Some(()) +} + +/// Anthropic Messages replies can travel through a Chat passthrough route. +/// Their generated text and tool-use payloads are visible output; all other +/// content-block kinds remain opaque. +fn append_anthropic_output_content_strings( + out: &mut String, + content: &RawJson<'_>, + scan_error: &mut Option, +) -> Option<()> { + if !content.get().trim_start().starts_with('[') { + return Some(()); + } + for block in raw_array_items(content)? { + if !raw_is_object(&block) { + continue; + } + let block_body = block.get().as_bytes(); + match raw_top_level_unique_type(block_body).ok()?.as_deref() { + Some("text") => append_raw_top_level_strings(out, block_body, "text")?, + Some("tool_use") => { + append_raw_top_level_strings(out, block_body, "name")?; + for input in raw_top_level_values(block_body, "input")? { + append_scan_text( + out, + &decoded_json_string_values_including_empty( + input.get().as_bytes(), + scan_error, + )?, + )?; + } } + Some(_) | None => {} } - if usage_labelled { - // The agent-backend shape: a flat token object on the - // server's own usage event, with no `usage` wrapper. - if let Some(u) = usage_of(&v) { - merge(u); + } + Some(()) +} + +/// OpenAI Completions exposes generated text only through a top-level +/// `choices` array. Do not use a broad JSON-string collection here: an +/// otherwise-valid provider extension may contain image, audio, or other +/// opaque data that must never cross an external output-guardrail boundary. +/// +/// Unlike the Chat selector, a detected Completions response has no other +/// compatible response envelope. Its `choices` carrier is therefore required +/// and must be unambiguous; a malformed, repeated, or over-cap selector is +/// unevaluable rather than a reason to scan arbitrary source fields. +fn decoded_completions_response_string_values( + body: &[u8], + scan_error: &mut Option, +) -> Option { + let choices = match raw_top_level_unique_array_ref(body, "choices") { + Ok(Some(choices)) => choices, + Ok(None) | Err(()) => { + mark_unevaluable(scan_error); + return None; + } + }; + let choices = match raw_array_item_refs(&choices) { + Some(choices) => choices, + None => { + mark_unevaluable(scan_error); + return None; + } + }; + let mut out = String::new(); + for choice in choices { + if !raw_is_object(&choice) { + mark_unevaluable(scan_error); + return None; + } + match raw_top_level_unique_string(choice.get().as_bytes(), "text") { + Ok(Some(text)) => append_scan_text(&mut out, &text)?, + // An empty choice has no generated text to inspect. It is not a + // reason to widen the output selector to another field. + Ok(None) => {} + Err(()) => { + mark_unevaluable(scan_error); + return None; } } - parts = match protocol { - PassthroughProtocol::Raw => crate::held_content::Parts { - scan: payload.to_string(), - reasoning: 0, - }, - // The chat envelope also carries Anthropic Messages streams; the - // two event shapes are disjoint, so reading both is exact. - PassthroughProtocol::OpenaiChat => { - let mut p = crate::held_content::chat_chunk_parts(&v); - let a = crate::held_content::anthropic_event_parts(&v); - p.scan.push_str(&a.scan); - p.reasoning += a.reasoning; - p + } + Some(out) +} + +fn decoded_chat_response_string_values( + body: &[u8], + scan_error: &mut Option, +) -> Option { + let mut out = String::new(); + let choices = raw_top_level_values(body, "choices")?; + let has_chat_choices = choices + .iter() + .any(|choices| choices.get().trim_start().starts_with('[')); + for array in choices { + let choices = raw_array_items(&array)?; + for choice in choices { + if !raw_is_object(&choice) { + continue; } - PassthroughProtocol::OpenaiCompletions => { - crate::held_content::completions_chunk_parts(&v) + for message in raw_top_level_values(choice.get().as_bytes(), "message")? { + append_chat_output_message_strings(&mut out, &message)?; } - PassthroughProtocol::OpenaiResponses => crate::held_content::responses_event_parts(&v), - }; + } } - (parts, usage) + if !has_chat_choices && raw_top_level_unique_type(body).ok()?.as_deref() == Some("message") { + for content in raw_top_level_values(body, "content")? { + append_anthropic_output_content_strings(&mut out, &content, scan_error)?; + } + } + Some(out) } -/// The SSE error frame appended when an output guardrail blocks mid-relay, -/// in the protocol of the stream it ends: the frame `/v1/messages` emits for -/// the same refusal on an Anthropic Messages stream, the one -/// `/v1/chat/completions` emits on every other. -fn guardrail_error_frame( - anthropic: bool, - guardrail_name: Option<&str>, - unavailable: Option<&str>, -) -> Bytes { - if anthropic { - return Bytes::from(crate::messages::guardrail_block_frame( - guardrail_name, - unavailable, - )); - } - Bytes::from(format!( - "event: error\ndata: {}\n\n", - crate::error::guardrail_block_frame_payload(guardrail_name, unavailable) - )) +/// One direct Responses stream carrier that exposes client-visible output. +/// `content_index` separates text/refusal parts; tool calls are item-scoped. +#[derive(Clone, Copy)] +struct ResponsesDirectStreamEvent { + fields: &'static [&'static str], + content_index: bool, } -/// Whether a relayed chat-envelope frame is an Anthropic Messages event -/// (`Some(true)`) or an OpenAI chat chunk (`Some(false)`). The two event -/// shapes are disjoint; a frame that is neither (a comment, `[DONE]`, an -/// in-band error) decides nothing. -fn anthropic_stream_frame(frame: &[u8]) -> Option { - let payload = crate::redact::frame_payload(frame)?; - let v: serde_json::Value = serde_json::from_str(payload.trim()).ok()?; - if v.get("choices").is_some() { - return Some(false); +fn responses_direct_stream_event(kind: &str) -> Option { + match kind { + "response.output_text.delta" | "response.refusal.delta" => { + Some(ResponsesDirectStreamEvent { + fields: &["delta"], + content_index: true, + }) + } + "response.output_text.done" => Some(ResponsesDirectStreamEvent { + fields: &["text"], + content_index: true, + }), + "response.refusal.done" => Some(ResponsesDirectStreamEvent { + fields: &["refusal"], + content_index: true, + }), + "response.function_call_arguments.delta" | "response.mcp_call_arguments.delta" => { + Some(ResponsesDirectStreamEvent { + fields: &["delta"], + content_index: false, + }) + } + "response.function_call_arguments.done" => Some(ResponsesDirectStreamEvent { + fields: &["name", "arguments"], + content_index: false, + }), + "response.mcp_call_arguments.done" => Some(ResponsesDirectStreamEvent { + fields: &["arguments"], + content_index: false, + }), + "response.custom_tool_call_input.delta" => Some(ResponsesDirectStreamEvent { + fields: &["delta"], + content_index: false, + }), + "response.custom_tool_call_input.done" => Some(ResponsesDirectStreamEvent { + fields: &["input"], + content_index: false, + }), + _ => None, } - match v.get("type").and_then(serde_json::Value::as_str) { - Some( - "message_start" - | "message_delta" - | "message_stop" - | "content_block_start" - | "content_block_delta" - | "content_block_stop" - | "ping", - ) => Some(true), +} + +fn responses_stream_event_is_visible(kind: &str) -> bool { + responses_direct_stream_event(kind).is_some() + || matches!( + kind, + "response.content_part.added" + | "response.content_part.done" + | "response.output_item.added" + | "response.output_item.done" + | "response.completed" + | "response.incomplete" + | "response.failed" + ) +} + +/// The only Responses content-part fields the typed output guardrail reads. +/// Other part types can carry image, audio, file, or reasoning data. +fn responses_visible_content_part_field(kind: &str) -> Option<&'static str> { + match kind { + "output_text" | "text" | "input_text" => Some("text"), + "refusal" => Some("refusal"), _ => None, } } -/// Build the streamed relay response: upstream SSE frames are forwarded -/// incrementally, tee'd through the chain's [`StreamOutputPolicy`] -/// (window / full-buffer hold-back, end-of-stream check otherwise), while -/// usage and capture accumulate for the end-of-stream telemetry emit. The -/// telemetry guard also fires from `Drop` when the client disconnects -/// mid-relay. -#[allow(clippy::too_many_arguments)] -fn stream_response( - protocol: PassthroughProtocol, - chain: aisix_guardrails::GuardrailChain, - upstream_resp: reqwest::Response, - resp_headers: HeaderMap, - status: reqwest::StatusCode, - mut telemetry: RouteTelemetry, - request_id: &str, -) -> Response { - use aisix_guardrails::{Guardrail as _, GuardrailVerdict, StreamOutputPolicy}; - use futures::StreamExt; +/// Source-preserving counterpart to the typed Responses output scanner's +/// content-part walk. A missing or conflicting discriminator is opaque: a +/// media item can use any string-shaped field, so only a unique known text +/// part may cross the external guardrail boundary. +fn append_responses_visible_part_strings(out: &mut String, part: &RawJson<'_>) -> Option<()> { + let part_body = part.get().as_bytes(); + if let Some(field) = raw_top_level_unique_type(part_body) + .ok()? + .as_deref() + .and_then(responses_visible_content_part_field) + { + append_raw_top_level_strings(out, part_body, field)?; + } + Some(()) +} - let policy = if chain.is_empty() { - StreamOutputPolicy::EndOfStreamCheck - } else { - chain.stream_output_policy() - }; - let route_name = telemetry.route_name.clone(); - let capture_cap = telemetry.content_cap; +fn append_responses_visible_content_strings(out: &mut String, content: &RawJson<'_>) -> Option<()> { + let value = content.get().trim_start(); + if value.starts_with('"') { + return append_raw_string_value(out, content); + } + if !value.starts_with('[') { + return Some(()); + } + let parts = raw_array_items(content)?; + for part in parts { + append_responses_visible_part_strings(out, &part)?; + } + Some(()) +} - let stream = async_stream::stream! { - let mut upstream = upstream_resp.bytes_stream(); - let mut splitter = SseFrameSplitter::new(); - // Held-back frames (Window / BufferFull) not yet released. - let mut pending: Vec = Vec::new(); - let mut pending_held = crate::held_content::HeldBytes::default(); - // Frame bytes held under Window (a memory bound for delta-free runs). - let mut held_bytes: usize = 0; - // What BufferFull holds (#513): content, which `max_buffer_bytes` - // caps (the SSE framing is not counted), and the raw frame bytes it - // bounds too. - let mut held_content = crate::held_content::HeldBuffer::default(); - // Unscanned delta text for the CURRENT window / buffer. - let mut scan_buf = String::new(); - // Overlap carried between Window scans. - let mut overlap_tail = String::new(); - // Degrades BufferFull to live forwarding after a fail-open cap hit. - let mut fail_opened = false; - let mut blocked = false; - // The chat envelope also carries Anthropic Messages streams; the - // first frame that says which one decides the refusal frame's shape. - let mut anthropic: Option = None; +/// Source-preserving counterpart to `responses::responses_output_text`. +/// Restrict the walk to client-visible message text and the tool payloads the +/// typed output guardrail already reads. A missing or conflicting item type +/// is opaque rather than a generic raw fallback: without a unique item kind, +/// `text`, `arguments`, and `input` could be an image/audio/file payload. +fn append_responses_output_item_strings(out: &mut String, item: &RawJson<'_>) -> Option<()> { + let item_body = item.get().as_bytes(); + match raw_top_level_unique_type(item_body).ok()?.as_deref() { + Some("reasoning") => {} + Some("message") => { + for content in raw_top_level_values(item_body, "content")? { + append_responses_visible_content_strings(out, &content)?; + } + } + Some("function_call" | "mcp_call") => { + for key in ["name", "arguments"] { + append_raw_top_level_strings(out, item_body, key)?; + } + } + Some("custom_tool_call") => { + for key in ["name", "input"] { + append_raw_top_level_strings(out, item_body, key)?; + } + } + Some(_) | None => {} + } + Some(()) +} - 'outer: loop { - let chunk = match upstream.next().await { - Some(Ok(c)) => c, - Some(Err(err)) => { - // The response head is already on the wire, so there is - // no status left to carry the failure — record it on the - // event instead of ending as a silent success. - let bridge = crate::dispatch::reqwest_error_to_bridge(&err, telemetry.started); - telemetry.record_failure(&bridge); - tracing::warn!( - route = %route_name, - error = %telemetry.error_message, - "passthrough-route upstream stream failed mid-relay", - ); - break; - } - None => break, - }; - // TTFT on the first upstream chunk of any type — the same - // convention the typed streaming endpoints stamp. - if telemetry.upstream_ttft_ms == 0 { - telemetry.upstream_ttft_ms = telemetry - .attempt_started - .elapsed() - .as_millis() - .min(u32::MAX as u128) as u32; - } - for frame in splitter.push(&chunk) { - if anthropic.is_none() && matches!(protocol, PassthroughProtocol::OpenaiChat) { - anthropic = anthropic_stream_frame(&frame); - } - if let Some(err) = frame_in_band_error(protocol, &frame) { - telemetry.record_failure(&err); - } - let (parts, usage) = frame_parts(protocol, &frame); - let held = parts.held(); - let delta = parts.scan; - if let Some(u) = usage { - merge_usage(&mut telemetry.usage, u); - } - if capture_cap.is_some() { - push_capped(&mut telemetry.response_text, &delta, capture_cap); - } - let frame = Bytes::from(frame); - match &policy { - _ if fail_opened => { - telemetry.mark_first_delivery(); - yield Ok::<_, std::convert::Infallible>(frame); - } - StreamOutputPolicy::EndOfStreamCheck => { - scan_buf.push_str(&delta); - telemetry.mark_first_delivery(); - yield Ok(frame); - } - StreamOutputPolicy::Window { size_chars, overlap_chars, .. } => { - scan_buf.push_str(&delta); - held_bytes += frame.len(); - pending_held.add(frame.len()); - pending.push(frame); - // The char threshold only advances on extracted delta - // text, so a run of delta-free frames (role-only, - // keep-alives, usage-only) would hold frames without - // bound — force the scan once the held BYTES cross - // the cap, mirroring BufferFull's self-bound. - if scan_buf.chars().count() >= *size_chars - || held_bytes > MAX_HELD_STREAM_BYTES - { - let text = format!("{overlap_tail}{scan_buf}"); - match scan_output(&chain, &route_name, &text, &mut telemetry).await { - GuardrailVerdict::Block { - reason, - guardrail_name, - unavailable, - } => { - tracing::warn!( - guardrail_hook = "output", - route = %route_name, - reason = %reason, - "guardrail blocked passthrough-route stream (window)", - ); - blocked = true; - yield Ok(guardrail_error_frame(anthropic.unwrap_or(false), guardrail_name.as_deref(), unavailable.as_deref())); - break 'outer; - } - _ => { - for f in pending.drain(..) { - telemetry.mark_first_delivery(); - yield Ok(f); - } - pending_held.clear(); - held_bytes = 0; - let combined = format!("{overlap_tail}{scan_buf}"); - overlap_tail = tail_chars(&combined, *overlap_chars); - scan_buf.clear(); - } - } - } - } - StreamOutputPolicy::BufferFull { max_buffer_bytes, on_exceeded_fail_open } => { - scan_buf.push_str(&delta); - held_content.hold(held, frame.len()); - pending_held.add(frame.len()); - pending.push(frame); - if held_content.exceeds(*max_buffer_bytes) { - if *on_exceeded_fail_open { - for f in pending.drain(..) { - telemetry.mark_first_delivery(); - yield Ok(f); - } - pending_held.clear(); - fail_opened = true; - chain.record_bypass(crate::error::TAG_OUTPUT_BUFFER_EXCEEDED); - } else { - tracing::warn!( - route = %route_name, - "passthrough-route stream exceeded the guardrail buffer cap (fail-closed)", - ); - blocked = true; - chain.record_output_buffer_exceeded(); - yield Ok(guardrail_error_frame(anthropic.unwrap_or(false), None, Some(crate::error::TAG_OUTPUT_BUFFER_EXCEEDED))); - break 'outer; - } - } - } +fn append_responses_output_strings(out: &mut String, body: &[u8]) -> Option<()> { + for output in raw_top_level_values(body, "output")? { + let items = raw_array_items(&output)?; + for item in items { + append_responses_output_item_strings(out, &item)?; + } + } + Some(()) +} + +fn decoded_responses_response_string_values(body: &[u8]) -> Option { + let mut out = String::new(); + append_responses_output_strings(&mut out, body)?; + Some(out) +} + +/// The response text a guardrail scans, per the route's protocol hint. +/// +/// This deliberately reads raw source values rather than `Value`, retaining +/// duplicate visible text and tool carriers which the client receives +/// verbatim. Generated reasoning and opaque media remain out of scope. +fn response_guardrail_text_with_scan_error( + protocol: PassthroughProtocol, + body: &[u8], + scan_error: &mut Option, +) -> String { + let raw = || String::from_utf8_lossy(body).into_owned(); + match protocol { + PassthroughProtocol::Raw => decoded_json_string_values(body).unwrap_or_else(raw), + PassthroughProtocol::OpenaiChat => { + // A detected Chat response can carry opaque multimodal values. + // Without a successful type-aware selection, relay it but do not + // send a raw fallback to an external output guardrail. + match decoded_chat_response_string_values(body, scan_error) { + Some(text) => text, + None => { + mark_unevaluable(scan_error); + String::new() } } } - - if !blocked { - // Trailing bytes with no frame terminator, plus the final scan - // of whatever the policy has not cleared yet. - let rest = splitter.take_rest(); - if !rest.is_empty() { - if anthropic.is_none() && matches!(protocol, PassthroughProtocol::OpenaiChat) { - anthropic = anthropic_stream_frame(&rest); - } - let (parts, usage) = frame_parts(protocol, &rest); - let held = parts.held(); - let delta = parts.scan; - if let Some(u) = usage { - merge_usage(&mut telemetry.usage, u); - } - if capture_cap.is_some() { - push_capped(&mut telemetry.response_text, &delta, capture_cap); - } - scan_buf.push_str(&delta); - let rest = Bytes::from(rest); - // The tail is held like any frame, under the same cap. - let tripped = match &policy { - StreamOutputPolicy::BufferFull { max_buffer_bytes, on_exceeded_fail_open } - if !fail_opened => - { - held_content.hold(held, rest.len()); - held_content - .exceeds(*max_buffer_bytes) - .then_some(*on_exceeded_fail_open) - } - _ => None, - }; - match tripped { - Some(false) => { - tracing::warn!( - route = %route_name, - "passthrough-route stream exceeded the guardrail buffer cap (fail-closed)", - ); - chain.record_output_buffer_exceeded(); - pending.clear(); - pending_held.clear(); - yield Ok(guardrail_error_frame(anthropic.unwrap_or(false), None, Some(crate::error::TAG_OUTPUT_BUFFER_EXCEEDED))); - telemetry.guardrail_blocked = true; - telemetry.stream_reached_end = true; - telemetry.emit(); - return; - } - Some(true) => { - for f in pending.drain(..) { - telemetry.mark_first_delivery(); - yield Ok(f); - } - pending_held.clear(); - chain.record_bypass(crate::error::TAG_OUTPUT_BUFFER_EXCEEDED); - telemetry.mark_first_delivery(); - yield Ok(rest); - } - None if policy.holds_back() && !fail_opened => { - pending_held.add(rest.len()); - pending.push(rest) - } - None => { - telemetry.mark_first_delivery(); - yield Ok(rest); - } - } - } - let text = format!("{overlap_tail}{scan_buf}"); - if !chain.is_empty() && !text.is_empty() { - if let GuardrailVerdict::Block { - reason, - guardrail_name, - unavailable, - } = - scan_output(&chain, &route_name, &text, &mut telemetry).await - { - tracing::warn!( - guardrail_hook = "output", - route = %route_name, - reason = %reason, - "guardrail blocked passthrough-route stream (end)", - ); - // Held frames are dropped (fail closed); content already - // forwarded under EndOfStreamCheck cannot be unsent — - // the error frame is the caller-visible signal either way. - pending.clear(); - pending_held.clear(); - yield Ok(guardrail_error_frame(anthropic.unwrap_or(false), guardrail_name.as_deref(), unavailable.as_deref())); - telemetry.guardrail_blocked = true; - telemetry.stream_reached_end = true; - telemetry.emit(); - return; + PassthroughProtocol::OpenaiCompletions => { + // A detected Completions response may carry opaque provider + // extensions. Only `choices[].text` may leave this process for + // output inspection; a bad selector is unevaluable, never a + // whole-body fallback. + match decoded_completions_response_string_values(body, scan_error) { + Some(text) => text, + None => { + mark_unevaluable(scan_error); + String::new() } } - for f in pending.drain(..) { - telemetry.mark_first_delivery(); - yield Ok(f); - } - pending_held.clear(); - } else { - telemetry.guardrail_blocked = true; } - // The generator ran to its own end (upstream EOF, upstream error, or - // a guardrail block); only a client that went away first leaves this - // unset, and the emit turns that into a 499. - telemetry.stream_reached_end = true; - telemetry.emit(); - }; - - // Re-attach the request span (the body is polled after the request-id - // middleware returns, so end-of-stream telemetry would otherwise log - // without a request_id) and heartbeat silence gaps — this branch is - // SSE-only, where a comment frame is protocol-legal and identical to - // what the typed endpoints emit; relayed frames are untouched. - let mut response = Response::builder() - .status(status) - .body(Body::from_stream(crate::sse_keepalive::with_heartbeat( - crate::request_id::in_request_span(stream), - crate::sse_keepalive::interval(), - ))) - .unwrap(); - copy_safe_headers(&resp_headers, response.headers_mut()); - // The relay re-chunks the body; a stale upstream length must not ride - // along (SSE normally has none, but a lying upstream shouldn't wedge - // the client). - response.headers_mut().remove(header::CONTENT_LENGTH); - if let Ok(hv) = HeaderValue::from_str(request_id) { - response - .headers_mut() - .insert(header::HeaderName::from_static("x-aisix-request-id"), hv); + PassthroughProtocol::OpenaiResponses => { + // Without a safely decoded Responses envelope, no discriminator + // can establish that a string is visible text rather than opaque + // image/audio/file data. Privacy wins over a raw fallback here. + decoded_responses_response_string_values(body).unwrap_or_default() + } } - response } -/// One output scan over `text`, folding monitor hits into the telemetry. -async fn scan_output( - chain: &aisix_guardrails::GuardrailChain, - route_name: &str, - text: &str, - telemetry: &mut RouteTelemetry, -) -> aisix_guardrails::GuardrailVerdict { - use aisix_guardrails::Guardrail as _; - let synth = aisix_gateway::ChatResponse { - id: String::new(), - model: route_name.to_string(), - message: aisix_gateway::ChatMessage::assistant(text.to_string()), - finish_reason: aisix_gateway::FinishReason::Stop, - usage: aisix_gateway::UsageStats::default(), - }; - let (verdict, hits) = chain.check_output_unmaskable_observed(&synth).await; - telemetry.monitor_hits.extend(hits); - verdict +fn response_guardrail_text(protocol: PassthroughProtocol, body: &[u8]) -> String { + let mut ignored_scan_error = None; + response_guardrail_text_with_scan_error(protocol, body, &mut ignored_scan_error) } -/// The last `n` chars of `s` (whole string when shorter). -fn tail_chars(s: &str, n: usize) -> String { - let count = s.chars().count(); - if count <= n { - return s.to_string(); +/// Like [`response_guardrail_text`], but preserves scanner failures from the +/// exact typed source selectors that read a value for output inspection. +fn try_response_guardrail_text( + protocol: PassthroughProtocol, + body: &[u8], +) -> Result { + match protocol { + PassthroughProtocol::Raw => try_raw_json_string_values(body), + PassthroughProtocol::OpenaiCompletions => { + let mut scan_error = None; + let text = response_guardrail_text_with_scan_error(protocol, body, &mut scan_error); + scan_error.map_or(Ok(text), Err) + } + PassthroughProtocol::OpenaiChat => { + let mut scan_error = None; + let text = response_guardrail_text_with_scan_error(protocol, body, &mut scan_error); + scan_error.map_or(Ok(text), Err) + } + PassthroughProtocol::OpenaiResponses => decoded_responses_response_string_values(body) + .ok_or_else(crate::json_splice::SpliceError::unevaluable), } - s.chars().skip(count - n).collect() } -/// Append `delta` to `buf`, bounded by `cap` bytes (capture accumulation -/// must not grow with an unbounded stream). Char boundaries are respected. -fn push_capped(buf: &mut String, delta: &str, cap: Option) { - let Some(cap) = cap else { return }; - if buf.len() >= cap { - return; - } - if buf.len() + delta.len() <= cap { - buf.push_str(delta); - return; - } - for c in delta.chars() { - if buf.len() + c.len_utf8() > cap { - break; +/// Select buffered output from the response envelope itself before scanning. +/// Request detection stays authoritative for request processing and streams; +/// this one boundary prevents a compatible provider response from becoming an +/// empty typed selector under the request's different protocol hint. +fn try_buffered_response_guardrail_text( + request_protocol: PassthroughProtocol, + body: &[u8], +) -> Result { + let response_protocol = resolve_buffered_response_protocol(request_protocol, body)?; + match try_response_guardrail_text(response_protocol, body) { + Err(error) if !error.is_unevaluable() => { + Ok(response_guardrail_text(response_protocol, body)) } - buf.push(c); + result => result, } } -// --------------------------------------------------------------------------- -// Telemetry -// --------------------------------------------------------------------------- - -/// End-of-request telemetry for a passthrough-route exchange: one -/// UsageEvent (CP sink + exporter fan-out, with captured content on the -/// exporter path only), the request metric, and the access log line. The -/// buffered path calls [`RouteTelemetry::emit`] inline; the streaming path -/// calls it at end-of-stream, with `Drop` covering client disconnects. -struct RouteTelemetry { - state: ProxyState, - route_name: String, - provider_label: String, - /// The request's trace bundle (AISIX-Cloud#1279) — the Drop emit is - /// the request's terminal emission, so it carries the terminal spans. - trace: Option>, - pk_id: String, - method: Method, - path: String, - request_id: String, - api_key_id: String, - /// Org member the authenticating key belongs to (AISIX-Cloud#1389), - /// and that member's display name for the `user_name` metric label - /// (AISIX-Cloud#1455). Both `None` for a key bound to no member — - /// including the anonymous route key, which belongs to the route - /// rather than to a person. - user_id: Option, - user_name: Option, - jwt: Option>, - /// Whether the caller reached this route through `auth_mode: - /// anonymous` rather than a credential of its own. Stamped onto the - /// usage event so anonymous traffic stays distinguishable from the - /// bound key's own (see `usage_attr::apply_auth_type`). - anonymous: bool, - client_identity: String, - client_source_ip: String, - client_user_agent: String, - started: Instant, - /// When the upstream call itself began — the scope the two `upstream_*` - /// figures are measured in, distinct from `started` (request receipt). - attempt_started: Instant, - status: u16, - /// Every token dimension the exchange reported, accumulated field-wise - /// across the response (buffered) or its frames (streamed). - usage: Option, - /// The model alias the caller addressed, read from a DETECTED - /// envelope's own `model` field — the same value the typed endpoint - /// serving that envelope records. Empty for an opaque body, whose - /// `model`-shaped key means nothing the gateway can trust. - requested_model: String, - /// Time from the START OF THE ATTEMPT to the upstream's first streamed - /// frame. Zero on the buffered path, where there is none. - upstream_ttft_ms: u32, - /// What the caller waited for on a streamed relay: the moment the first - /// relayed frame was handed downstream, measured from `started`. `None` - /// until one is, so a stream that delivered nothing reports no - /// caller-wait at all rather than an invented one. - downstream_first_ms: Option, - /// `true` once the relay generator reached its own end — upstream EOF, - /// upstream error, or a guardrail block. It stays `false` only when the - /// CLIENT went away first, which is what the emit turns into a 499 - /// (same signal the typed streaming endpoints record). - stream_reached_end: bool, - /// Set on a streamed relay so the `Drop` emit can tell an abandoned - /// stream from the buffered path, which never streams at all. - streaming: bool, - /// Bounded error class + message for a failure the relay could not - /// answer with a status code — an upstream that dies mid-stream, after - /// the response head is already on the wire. - error_class: String, - error_message: String, - /// The status that same failure gets before the response head - /// ([`aisix_gateway::BridgeError::http_status`]). The emit records it in - /// place of the upstream's `200`: the caller's response line cannot - /// change any more, but the record of what happened can. - failure_status: Option, - monitor_hits: Vec, - /// The request's ENFORCE-mode audit handle (AISIX-Cloud#1330). Held - /// rather than snapshotted at construction: this struct's emit runs - /// from the relay's `Drop`, long after the output hook has recorded - /// whatever it masked. - audit: crate::usage_attr::GuardrailAudit, - captured_prompt: Option, - content_cap: Option, - response_text: String, - guardrail_blocked: bool, - emitted: bool, +/// The typed visible-response extraction used for telemetry capture. It is +/// intentionally separate from the broader guardrail source scan above. +fn response_visible_text(protocol: PassthroughProtocol, body: &[u8]) -> String { + let raw = || String::from_utf8_lossy(body).into_owned(); + if matches!(protocol, PassthroughProtocol::Raw) { + return raw(); + } + let Ok(v) = serde_json::from_slice::(body) else { + return decoded_json_string_values(body).unwrap_or_else(raw); + }; + if matches!(protocol, PassthroughProtocol::OpenaiResponses) { + let joined = crate::responses::responses_output_text(&v); + return if joined.is_empty() { raw() } else { joined }; + } + if matches!(protocol, PassthroughProtocol::OpenaiChat) && v.get("choices").is_none() { + let text = anthropic_message_output_text(&v); + if !text.is_empty() { + return text; + } + } + let choices = v.get("choices").and_then(|c| c.as_array()); + let Some(choices) = choices else { return raw() }; + let texts: Vec = choices + .iter() + .filter_map(|c| match protocol { + PassthroughProtocol::OpenaiChat => { + c.get("message").map(|m| message_scan_text(m, false)) + } + PassthroughProtocol::OpenaiCompletions => { + c.get("text").and_then(|t| t.as_str()).map(str::to_string) + } + PassthroughProtocol::OpenaiResponses | PassthroughProtocol::Raw => None, + }) + .filter(|t| !t.is_empty()) + .collect(); + if texts.is_empty() { + raw() + } else { + texts.join("\n") + } } -impl RouteTelemetry { - /// Record the upstream failure that ended a streamed relay after its - /// head went out. The first one is the cause; later ones do not replace - /// it. - fn record_failure(&mut self, err: &aisix_gateway::BridgeError) { - if self.failure_status.is_some() { - return; +/// Captures preserve the existing typed visible-content contract; guardrail +/// scanning may include additional raw fields that are forwarded downstream. +fn response_capture_text(protocol: PassthroughProtocol, body: &[u8]) -> String { + response_visible_text(protocol, body) +} + +/// Every token dimension a passthrough exchange can report, mirroring the +/// token fields of [`aisix_obs::UsageEvent`] 1:1 so a route reports what +/// the typed endpoint serving the same envelope would. +/// +/// Populated from the union of spellings the relayed APIs use — OpenAI's +/// nested `*_tokens_details`, the Responses API's `input`/`output` +/// spelling, Anthropic's separate cache counters, DeepSeek's native +/// `prompt_cache_hit_tokens`, and the flat token object agent backends +/// report on their own SSE event. +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] +struct PassthroughUsage { + prompt_tokens: u32, + completion_tokens: u32, + cached_prompt_tokens: u32, + cache_write_tokens: Option, + reasoning_tokens: u32, + cache_creation_tokens: u32, + cache_read_tokens: u32, + /// The upstream's own `total_tokens`, verbatim; `None` when a report + /// carried none. Never a sum computed here — see + /// `UsageStats::upstream_total_tokens`. + upstream_total_tokens: Option, +} + +impl PassthroughUsage { + /// Field-wise max, the accumulation the typed streaming paths use. + /// + /// One stream reports usage across several frames — Anthropic's + /// `message_start` carries the input and cache counters while its + /// terminal `message_delta` carries only the output ones — so a later + /// partial report must EXTEND the record rather than replace it. Max + /// also makes a provider that repeats a cumulative usage object + /// harmless. + fn merge(&mut self, other: Self) { + self.prompt_tokens = self.prompt_tokens.max(other.prompt_tokens); + self.completion_tokens = self.completion_tokens.max(other.completion_tokens); + self.cached_prompt_tokens = self.cached_prompt_tokens.max(other.cached_prompt_tokens); + self.cache_write_tokens = self.cache_write_tokens.max(other.cache_write_tokens); + self.reasoning_tokens = self.reasoning_tokens.max(other.reasoning_tokens); + self.cache_creation_tokens = self.cache_creation_tokens.max(other.cache_creation_tokens); + self.cache_read_tokens = self.cache_read_tokens.max(other.cache_read_tokens); + // A total stands only while every merged report carried one. + self.upstream_total_tokens = self + .upstream_total_tokens + .zip(other.upstream_total_tokens) + .map(|(a, b)| a.max(b)); + } +} + +/// Merge one usage report into an exchange's accumulated usage. The first +/// report is adopted whole rather than merged into zeros, so a total it +/// carried survives until a report without one arrives. +fn merge_usage(acc: &mut Option, report: PassthroughUsage) { + match acc { + Some(acc) => acc.merge(report), + None => *acc = Some(report), + } +} + +/// `usage` figures from a buffered protocol-aware response body. +fn response_usage( + protocol: PassthroughProtocol, + raw_shape: Option, + body: &[u8], +) -> Option { + if matches!(protocol, PassthroughProtocol::Raw) && raw_shape.is_none() { + return None; + } + let v = serde_json::from_slice::(body).ok()?; + match raw_shape { + Some(RawUsageShape::Rerank) => { + crate::rerank::rerank_prompt_tokens(&v).map(|prompt_tokens| PassthroughUsage { + prompt_tokens, + ..PassthroughUsage::default() + }) } - let failure = crate::attempt::StreamFailure::from_bridge(err); - self.error_class = failure.error_class.to_string(); - self.error_message = failure.error_message; - self.failure_status = Some(failure.status); + Some(RawUsageShape::DashscopeNative) => dashscope_native_usage(v.get("usage")?), + None => usage_of(v.get("usage")?), } +} - /// Stamp the caller's wait at the first RELAYED frame handed - /// downstream. - /// - /// Deliberately here and not where the frame was read off the upstream: - /// a hold-back guardrail policy sits between the two, and - /// `UsageEvent::downstream_latency_ms` counts that hold-back as part of - /// what the caller waited for. Called on both the live-forward and the - /// hold-back release paths, so it catches the first frame either way. - /// - /// A synthetic frame (a guardrail block's error event) deliberately - /// does NOT stamp: nothing the caller asked for was delivered. Same - /// rule as the typed streaming endpoints, which stamp only in the - /// chunk renderer. - fn mark_first_delivery(&mut self) { - if self.downstream_first_ms.is_none() { - self.downstream_first_ms = - Some(self.started.elapsed().as_millis().min(u32::MAX as u128) as u32); +/// Token counts from a DashScope native response's top-level `usage`. +/// +/// Every dimension the generic reader knows (cache hit, reasoning, …) is +/// read by [`usage_of`]; only prompt and completion follow DashScope's own +/// arithmetic. The flat `image_tokens` means different things per service: +/// multimodal generation counts it inside `input_tokens` +/// (`{input_tokens: 79, image_tokens: 66, output_tokens: 14, +/// total_tokens: 93}`), while the multimodal embedding and rerank services +/// count it beside (`{input_tokens: 44, image_tokens: 64, +/// total_tokens: 108}`). So `total_tokens` less the completion is the +/// prompt whenever a total is reported; without one (only some embedding +/// models omit it, and those report images beside the text), prompt is +/// `input_tokens` (or `prompt_tokens`) plus the flat `image_tokens`. A +/// nested `input_tokens_details.image_tokens` is a breakdown already inside +/// `input_tokens` and is never added. +fn dashscope_native_usage(usage: &serde_json::Value) -> Option { + let num = |k: &str| { + usage + .get(k) + .and_then(serde_json::Value::as_u64) + .map(|n| n.min(u32::MAX as u64) as u32) + }; + let completion = num("output_tokens").or_else(|| num("completion_tokens")); + let prompt = match num("total_tokens") { + Some(total) => Some(total.saturating_sub(completion.unwrap_or(0))), + None => { + let input = num("input_tokens").or_else(|| num("prompt_tokens")); + let image = num("image_tokens"); + (input.is_some() || image.is_some()) + .then(|| input.unwrap_or(0).saturating_add(image.unwrap_or(0))) } + }; + let generic = usage_of(usage); + if prompt.is_none() && completion.is_none() && generic.is_none() { + return None; + } + Some(PassthroughUsage { + prompt_tokens: prompt.unwrap_or(0), + completion_tokens: completion.unwrap_or(0), + upstream_total_tokens: num("total_tokens"), + ..generic.unwrap_or_default() + }) +} + +/// Read every token dimension out of one `usage` object (or, for the +/// labelled frame of an opaque stream, a flat token object). +/// +/// The spellings are read as a union rather than per protocol because a +/// passthrough route relays whichever API the caller addressed: the same +/// route carries an OpenAI chat envelope, an Anthropic one, and an agent +/// backend's private shape. They do not collide — each name belongs to +/// exactly one API — so reading them all costs nothing and a detected +/// envelope reports what its typed endpoint would. +/// +/// `None` when the object carries no recognised counter at all, which is +/// what keeps a `usage`-shaped object that is not a usage report from +/// minting zeros. +fn usage_of(usage: &serde_json::Value) -> Option { + let num = |v: Option<&serde_json::Value>| { + v.and_then(serde_json::Value::as_u64) + .map(|n| n.min(u32::MAX as u64) as u32) + }; + // Flat counter under any of `names`, first hit wins. + let flat = |names: &[&str]| names.iter().find_map(|n| num(usage.get(*n))); + // `parent.child` counter, e.g. `prompt_tokens_details.cached_tokens`. + let nested = |parent: &str, child: &str| num(usage.get(parent).and_then(|d| d.get(child))); + + let prompt = flat(&["prompt_tokens", "input_tokens"]); + let completion = flat(&["completion_tokens", "output_tokens"]); + // OpenAI nests the cache hit under `prompt_tokens_details`, the + // Responses API under `input_tokens_details`, DeepSeek reports it flat + // as `prompt_cache_hit_tokens`. A nested ZERO must not mask a real + // native count (the typed OpenAI bridge takes the same precedence). + let cached_prompt = nested("prompt_tokens_details", "cached_tokens") + .filter(|&n| n > 0) + .or_else(|| nested("input_tokens_details", "cached_tokens").filter(|&n| n > 0)) + .or_else(|| flat(&["prompt_cache_hit_tokens", "cached_tokens"])); + let cache_write = nested("prompt_tokens_details", "cache_write_tokens") + .or_else(|| nested("input_tokens_details", "cache_write_tokens")); + let reasoning = nested("completion_tokens_details", "reasoning_tokens") + .filter(|&n| n > 0) + .or_else(|| nested("output_tokens_details", "reasoning_tokens").filter(|&n| n > 0)) + .or_else(|| flat(&["reasoning_tokens"])); + // Anthropic's two cache counters sit beside `input_tokens`, and are + // ADDITIVE to it rather than a subset. + let cache_creation = flat(&["cache_creation_input_tokens", "cache_creation_tokens"]); + let cache_read = flat(&["cache_read_input_tokens", "cache_read_tokens"]); + + let dims = [ + prompt, + completion, + cached_prompt, + cache_write, + reasoning, + cache_creation, + cache_read, + ]; + if dims.iter().all(Option::is_none) { + return None; + } + Some(PassthroughUsage { + prompt_tokens: prompt.unwrap_or(0), + completion_tokens: completion.unwrap_or(0), + cached_prompt_tokens: cached_prompt.unwrap_or(0), + cache_write_tokens: cache_write, + reasoning_tokens: reasoning.unwrap_or(0), + cache_creation_tokens: cache_creation.unwrap_or(0), + cache_read_tokens: cache_read.unwrap_or(0), + upstream_total_tokens: flat(&["total_tokens"]), + }) +} + +/// Model-level rate-limit identity from the JSON body's top-level `model` +/// field, scoped to `provider_lower` — the #805 contract carried over from +/// the removed implicit tunnel: `display_name` exact hit first, then the +/// provider-native `model_name` (deterministic on ties, wildcards +/// excluded), with the reservation keyed by `display_name` so route and +/// typed traffic to the same Model draw from one bucket. `None` for +/// non-JSON bodies, absent/unregistered names, or cross-provider names — +/// the request then reserves only the caller-level layers. +fn body_model_rate_limit( + snapshot: &aisix_core::AisixSnapshot, + provider_lower: &str, + body: &[u8], +) -> Option { + #[derive(serde::Deserialize)] + struct BodyModelProbe { + model: Option, + } + let name = serde_json::from_slice::(body).ok()?.model?; + let matches_provider = |m: &aisix_core::Model| { + m.provider + .as_deref() + .is_some_and(|p| p.eq_ignore_ascii_case(provider_lower)) + }; + let entry = snapshot + .models + .get_by_name(&name) + .filter(|e| matches_provider(&e.value)) + .or_else(|| { + snapshot + .models + .entries() + .into_iter() + .filter(|e| { + matches_provider(&e.value) + && e.value.model_name.as_deref() == Some(name.as_str()) + && !e.value.display_name.contains('*') + }) + .min_by_key(|e| e.id.clone()) + })?; + Some(crate::quota::ModelRateLimit::from_model( + &entry.value.display_name, + &entry.id, + &entry.value, + )) +} + +/// `true` if `seg` is a strict api-version path component matching `v\d+`. +fn is_api_version_segment(seg: &str) -> bool { + seg.starts_with('v') && seg.len() > 1 && seg[1..].chars().all(|c| c.is_ascii_digit()) +} + +/// Strip one leading api-version segment from `rest` when it exactly +/// matches the trailing version segment of `base` (#164): an operator's +/// `target_url` ending in `/v1` joined with a caller path starting `v1/` +/// would otherwise produce `/v1/v1/...`. +fn strip_redundant_version_segment<'a>(base: &str, rest: &'a str) -> &'a str { + let base_path = Url::parse(base) + .map(|url| url.path().to_owned()) + .unwrap_or_else(|_| base.to_owned()); + let base_tail = base_path + .trim_end_matches('/') + .rsplit('/') + .next() + .unwrap_or(""); + if !is_api_version_segment(base_tail) { + return rest; + } + if let Some(remainder) = rest.strip_prefix(base_tail) { + if remainder.is_empty() { + return remainder; + } + if let Some(after_slash) = remainder.strip_prefix('/') { + return after_slash; + } + } + rest +} + +/// Join a caller-controlled route remainder below the operator-configured +/// target URL. Parsing the candidate before the prefix check is deliberate: +/// the URL implementation canonicalises both literal and percent-encoded dot +/// segments, so the check observes the path reqwest will actually send. +fn join_target_url(base: &str, rest: &str, query: Option<&str>) -> Result { + if has_path_traversal_segment(rest) { + return Err(ProxyError::InvalidRequest( + "passthrough path is outside the configured target URL".into(), + )); + } + + let mut base_url = Url::parse(base).map_err(|_| { + ProxyError::InvalidRequest("passthrough route has an invalid target URL".into()) + })?; + // A configured query is part of the operator-owned target. Preserve it + // and append a non-conflicting caller query after it, while keeping it + // out of the URL string used for path joining. A caller must not be able + // to override an operator-owned key through a last-value query parser. + let base_query = base_url + .query() + .filter(|query| !query.is_empty()) + .map(str::to_owned); + base_url.set_query(None); + base_url.set_fragment(None); + let base_for_join = base_url.as_str().trim_end_matches('/'); + let joined = if rest.is_empty() { + base_for_join.to_string() + } else { + format!("{base_for_join}/{rest}") + }; + let mut target = Url::parse(&joined).map_err(|_| { + ProxyError::InvalidRequest("passthrough path is outside the configured target URL".into()) + })?; + let inbound_query = query.filter(|query| !query.is_empty()); + if let (Some(base), Some(inbound)) = (base_query.as_deref(), inbound_query) { + if query_keys_overlap(base, inbound) { + return Err(ProxyError::InvalidRequest( + "passthrough query conflicts with the configured target URL".into(), + )); + } + } + let merged_query = match (base_query.as_deref(), inbound_query) { + (Some(base), Some(inbound)) => Some(format!("{base}&{inbound}")), + (Some(base), _) => Some(base.to_owned()), + (_, Some(inbound)) => Some(inbound.to_owned()), + (None, None) => None, + }; + target.set_query(merged_query.as_deref()); + + if target.origin() != base_url.origin() || !path_is_within_base(target.path(), base_url.path()) + { + return Err(ProxyError::InvalidRequest( + "passthrough path is outside the configured target URL".into(), + )); + } + Ok(target.into()) +} + +/// Query key comparison examines each bounded whole-query decoded form. A +/// backend may decode before splitting on `&` or `;`, so checking only keys +/// from the original form would miss `safe=1%26tenant%3Dcaller` becoming a +/// `tenant` key downstream. Some form parsers also canonicalize a key's +/// bracketed suffix and ASCII dot/space spelling, so compare that normalized +/// form at every decode level. Query sets here are tiny, so a simple vector +/// keeps the parser behavior explicit without adding a dependency. +fn query_keys_overlap(base: &str, inbound: &str) -> bool { + let Some(base_keys) = query_keys_at_all_decode_levels(base) else { + return true; + }; + let Some(inbound_keys) = query_keys_at_all_decode_levels(inbound) else { + return true; + }; + inbound_keys + .iter() + .any(|key| base_keys.iter().any(|base_key| base_key == key)) +} + +fn query_keys_at_all_decode_levels(query: &str) -> Option>> { + let mut decoded = query.as_bytes().to_vec(); + let mut keys = Vec::new(); + for _ in 0..=MAX_PERCENT_DECODE_PASSES { + for field in decoded.split(|byte| matches!(*byte, b'&' | b';')) { + for (key, _) in url::form_urlencoded::parse(field) { + let key = normalize_form_query_key(&key); + if !keys.iter().any(|existing| existing == &key) { + keys.push(key); + } + } + } + let next = percent_decode(&decoded).collect::>(); + if next == decoded { + return Some(keys); + } + decoded = next; + } + None +} + +/// Match PHP-style form-key registration after form decoding: leading ASCII +/// spaces are ignored, a NUL or bracketed suffix terminates the base key, and +/// ASCII dots/spaces in that base key are aliases for underscores. +fn normalize_form_query_key(key: &str) -> Vec { + key.as_bytes() + .iter() + .skip_while(|byte| **byte == b' ') + .take_while(|byte| !matches!(**byte, b'\0' | b'[')) + .map(|byte| match *byte { + b'.' | b' ' => b'_', + byte => byte, + }) + .collect() +} + +/// A backend may decode percent escapes before routing. Decode a bounded +/// number of times, so nested encoding cannot turn a harmless-looking +/// segment into `..` downstream. Inputs that keep changing beyond the bound +/// are rejected rather than delegated to an upstream with unknown decoding. +const MAX_PERCENT_DECODE_PASSES: usize = 4; + +fn has_path_traversal_segment(path: &str) -> bool { + let mut decoded = path.as_bytes().to_vec(); + for _ in 0..=MAX_PERCENT_DECODE_PASSES { + if decoded + .split(|byte| matches!(*byte, b'/' | b'\\')) + .any(|segment| { + let path_part = segment + .split(|byte| *byte == b';') + .next() + .unwrap_or_default(); + path_part == b"." || path_part == b".." + }) + { + return true; + } + + let next = percent_decode(&decoded).collect::>(); + if next == decoded { + return false; + } + decoded = next; + } + + true +} + +/// `candidate` must remain at `base` itself or below it on a path-segment +/// boundary. A root target intentionally permits every absolute path. +fn path_is_within_base(candidate: &str, base: &str) -> bool { + let base = base.trim_end_matches('/'); + if base.is_empty() { + return candidate.starts_with('/'); + } + candidate == base + || candidate + .strip_prefix(base) + .is_some_and(|remainder| remainder.starts_with('/')) +} + +// --------------------------------------------------------------------------- +// Streaming relay +// --------------------------------------------------------------------------- + +/// Default cap on bytes buffered while waiting for one SSE frame terminator. +/// A hold-back policy instead derives its splitter cap from its own raw-byte +/// limit, so an unterminated frame cannot outgrow the bytes it may hold. +const MAX_HELD_STREAM_BYTES: usize = 1024 * 1024; + +/// Bound independently scanned output candidates even when an upstream never +/// emits its terminal item event. Two source branches (first/last) are kept +/// for each identity; supplemental values consume the same budget so one +/// frame cannot fan out into an unbounded number of guardrail calls. +const MAX_STREAM_GUARDRAIL_CHANNELS: usize = 64; + +/// Empty epochs still occupy bookkeeping and eventually trigger a custom +/// guardrail scan, so cap them separately from their text candidates. +const MAX_STREAM_GUARDRAIL_EPOCHS: usize = 64; + +/// Source identities are bookkeeping only, never output text. Bound them +/// separately so a few large provider item ids cannot dominate the relay's +/// memory while the content-channel cap still sees only 64 branches. +const MAX_STREAM_GUARDRAIL_SOURCE_ID_BYTES: usize = 256; + +/// The gateway's shared frame splitter, with this relay's overflow policy: +/// an oversized unterminated run is handed on as a frame. Bytes after the +/// last complete frame stay buffered until more arrive; `take_rest` drains +/// them at end-of-stream. +struct SseFrameSplitter(aisix_gateway::sse::SseFrameSplitter); + +struct SseFrame { + bytes: Vec, + /// The upstream did not terminate the frame before the configured + /// hold-back raw-byte bound, so `bytes` cannot safely be decoded as an + /// SSE payload. + overflowed: bool, +} + +impl SseFrameSplitter { + fn with_max_frame_bytes(max_frame_bytes: usize) -> Self { + Self(aisix_gateway::sse::SseFrameSplitter::new(max_frame_bytes)) + } + + fn push(&mut self, chunk: &[u8]) -> Vec { + self.0.push(chunk); + let mut frames = Vec::new(); + loop { + match self.0.next_frame() { + Ok(Some(bytes)) => frames.push(SseFrame { + bytes, + overflowed: false, + }), + Ok(None) => break, + // Frame-terminator starvation: hand the oversized run on as-is + // rather than buffering without bound. + Err(_) => { + frames.push(SseFrame { + bytes: self.0.take_rest(), + overflowed: true, + }); + break; + } + } + } + frames + } + + fn take_rest(&mut self) -> Vec { + self.0.take_rest() + } +} + +/// `true` when the frame carries an `event:` line the server itself names +/// a usage report. The only evidence an opaque stream offers that a +/// token-shaped payload IS usage — see [`frame_delta`]. +fn is_usage_labelled_frame(frame: &[u8]) -> bool { + aisix_gateway::sse::lines(frame).any(|line| { + aisix_gateway::sse::parse_field(&frame[line.start..line.end]).is_some_and( + |(name, value)| { + let value = value.trim_ascii(); + name == b"event" + && (value.eq_ignore_ascii_case(b"token_usage") + || value.eq_ignore_ascii_case(b"usage")) + }, + ) + }) +} + +/// Text a guardrail scans from one SSE frame, per the protocol hint, plus +/// a usage probe on the same parsed payload. +/// +/// Usage is read from, in order of specificity: +/// +/// - the payload's own top-level `usage` object (every OpenAI-shape +/// stream, and Anthropic's terminal `message_delta`); +/// - `message.usage` on an Anthropic `message_start`, which is where the +/// input and cache counters arrive — its `message_delta` reports only +/// the output ones, so reading just the top level loses the prompt side +/// of every Anthropic stream; +/// - `response.usage` on a Responses stream's terminal event; +/// - for an OPAQUE (`Raw`) stream only, a FLAT token object on a frame the +/// server labelled a usage report (`event: token_usage`). An opaque +/// stream has no envelope to authenticate a payload against, so the +/// server's own label is the evidence — a payload that merely happens to +/// carry token-shaped fields must never mint billed tokens. +/// +/// Frames accumulate field-wise (see [`PassthroughUsage::merge`]) at the +/// call site, so a partial report never truncates an earlier one. +/// The upstream failure an SSE frame reports in-band, read with the same +/// mappings the typed endpoints use for the protocol the route carries. An +/// opaque (`Raw`) stream has no error envelope the gateway could recognise. +fn frame_in_band_error( + protocol: PassthroughProtocol, + frame: &[u8], +) -> Option { + if matches!(protocol, PassthroughProtocol::Raw) { + return None; + } + let payload = crate::redact::frame_payload(frame)?; + let payload = payload.trim(); + let value = serde_json::from_str::(payload).ok()?; + match protocol { + PassthroughProtocol::Raw => None, + PassthroughProtocol::OpenaiResponses => crate::responses::responses_in_band_error(&value), + // The chat envelope carries Anthropic Messages traffic too, whose + // in-band failure is a `type: "error"` event. + PassthroughProtocol::OpenaiChat | PassthroughProtocol::OpenaiCompletions => { + if value.get("type").and_then(|t| t.as_str()) == Some("error") { + if let Some(body) = value.get("error").and_then(|e| { + serde_json::from_value::< + aisix_provider_anthropic::wire::AnthropicStreamErrorBody, + >(e.clone()) + .ok() + }) { + return Some( + aisix_provider_anthropic::wire::stream_error_into_bridge_error(&body), + ); + } + } + aisix_gateway::capture_in_band_error(payload, aisix_gateway::UpstreamWire::OpenAI) + } + } +} + +#[cfg(test)] +fn frame_delta(protocol: PassthroughProtocol, frame: &[u8]) -> (String, Option) { + let (parts, usage) = frame_parts(protocol, frame); + (parts.scan, usage) +} + +/// Whether one unambiguous Anthropic stream event carries only generated +/// reasoning. `None` means the source shape was not safely inspectable, so +/// callers must preserve it rather than treating it as hidden. +fn hidden_chat_stream_reasoning_frame(body: &[u8]) -> Option { + match raw_top_level_unique_type(body).ok()?.as_deref() { + Some("content_block_delta") => raw_top_level_items_have_only_types( + body, + "delta", + &["thinking_delta", "signature_delta"], + ), + Some("content_block_start") => raw_top_level_items_have_only_types( + body, + "content_block", + &["thinking", "redacted_thinking"], + ), + _ => Some(false), + } +} + +#[cfg(test)] +fn decoded_chat_frame_string_values(body: &[u8]) -> Option { + let mut out = String::new(); + for value in decoded_chat_frame_continuations(body)? { + append_scan_text(&mut out, &value)?; + } + for value in decoded_chat_frame_supplemental_values(body)? { + append_scan_text(&mut out, &value)?; + } + Some(out) +} + +/// The raw stream selector accepts every typed, client-visible Responses +/// carrier. A provider may legally end after an authoritative `.done`, item, +/// part, or terminal snapshot without having sent a delta first. +#[cfg(test)] +fn decoded_responses_frame_string_values(body: &[u8]) -> Option { + let mut out = String::new(); + for value in responses_stream_selected_values(body).ok()? { + append_scan_text(&mut out, &value)?; + } + Some(out) +} + +fn decoded_chat_frame_continuations(body: &[u8]) -> Option> { + match raw_top_level_unique_type(body).ok()?.as_deref() { + Some("content_block_delta") => { + let delta = raw_top_level_unique_object(body, "delta").ok()??; + match raw_top_level_unique_type(delta.get().as_bytes()) + .ok()? + .as_deref() + { + Some("text_delta") => raw_top_level_string_values(delta.get().as_bytes(), "text"), + Some("input_json_delta") => { + raw_top_level_string_values(delta.get().as_bytes(), "partial_json") + } + Some(_) | None => Some(Vec::new()), + } + } + Some("content_block_start") => { + let block = raw_top_level_unique_object(body, "content_block").ok()??; + let block_body = block.get().as_bytes(); + match raw_top_level_unique_type(block_body).ok()?.as_deref() { + Some("text") => raw_top_level_string_values(block_body, "text"), + Some("tool_use") => { + let mut out = Vec::new(); + for input in raw_top_level_values(block_body, "input")? { + out.extend( + crate::json_splice::collect_string_values_where_vec( + input.get().as_bytes(), + |_| true, + ) + .ok()?, + ); + } + Some(out) + } + Some(_) | None => Some(Vec::new()), + } + } + _ => { + let choices = raw_top_level_unique_array(body, "choices").ok()??; + let choices = raw_array_items(&choices)?; + let mut out = Vec::new(); + for choice in choices { + if !raw_is_object(&choice) { + return None; + } + let delta = match raw_top_level_unique_object(choice.get().as_bytes(), "delta") { + Ok(Some(delta)) => delta, + Ok(None) => continue, + Err(()) => return None, + }; + let delta_body = delta.get().as_bytes(); + let mut content = raw_top_level_values(delta_body, "content")?; + if content.len() > 1 { + return None; + } + if let Some(content) = content.pop() { + match content.get().trim_start().as_bytes().first() { + Some(b'"') => out.push(serde_json::from_str(content.get()).ok()?), + Some(b'[') => { + for part in raw_array_items(&content)? { + if !raw_is_object(&part) { + return None; + } + let part_body = part.get().as_bytes(); + if let Some(field) = raw_top_level_unique_type(part_body) + .ok()? + .as_deref() + .and_then(chat_visible_content_part_field) + { + out.extend(raw_top_level_string_values(part_body, field)?); + } + } + } + Some(b'n') => {} + _ => return None, + } + } + out.extend(raw_top_level_string_values(delta_body, "refusal")?); + let tool_calls = raw_top_level_unique_array(delta_body, "tool_calls").ok()?; + if let Some(tool_calls) = tool_calls { + for tool_call in raw_array_items(&tool_calls)? { + if !raw_is_object(&tool_call) { + return None; + } + let tool_body = tool_call.get().as_bytes(); + for (container, field) in chat_tool_continuation_fields(tool_body)? { + for payload in raw_top_level_values(tool_body, container)? { + if !raw_is_object(&payload) { + return None; + } + out.extend(raw_top_level_string_values( + payload.get().as_bytes(), + field, + )?); + } + } + } + } + let function_call = + raw_top_level_unique_object(delta_body, "function_call").ok()?; + if let Some(function_call) = function_call { + out.extend(raw_top_level_string_values( + function_call.get().as_bytes(), + "arguments", + )?); + } + } + Some(out) + } + } +} + +/// The raw source continuations preserve every occurrence of visible carrier +/// fields. Supplementary scan text therefore excludes those same paths: a +/// normal frame must not send one visible value to a guardrail as typed, +/// source, and supplementary text at once. +fn decoded_chat_frame_supplemental_values(body: &[u8]) -> Option> { + if raw_top_level_unique_type(body).ok()?.as_deref() == Some("content_block_delta") { + return Some(Vec::new()); + } + if raw_top_level_unique_type(body).ok()?.as_deref() == Some("content_block_start") { + let block = raw_top_level_unique_object(body, "content_block").ok()??; + return match raw_top_level_unique_type(block.get().as_bytes()) + .ok()? + .as_deref() + { + Some("tool_use") => raw_top_level_string_values(block.get().as_bytes(), "name"), + Some(_) | None => Some(Vec::new()), + }; + } + + let choices = raw_top_level_unique_array(body, "choices").ok()??; + let choices = raw_array_items(&choices)?; + let mut out = Vec::new(); + for choice in choices { + if !raw_is_object(&choice) { + return None; + } + let delta = match raw_top_level_unique_object(choice.get().as_bytes(), "delta") { + Ok(Some(delta)) => delta, + Ok(None) => continue, + Err(()) => return None, + }; + let tool_calls = raw_top_level_unique_array(delta.get().as_bytes(), "tool_calls").ok()?; + if let Some(tool_calls) = tool_calls { + for tool_call in raw_array_items(&tool_calls)? { + if !raw_is_object(&tool_call) { + return None; + } + let tool_body = tool_call.get().as_bytes(); + for (container, _) in chat_tool_continuation_fields(tool_body)? { + for payload in raw_top_level_values(tool_body, container)? { + if !raw_is_object(&payload) { + return None; + } + out.extend(raw_top_level_string_values( + payload.get().as_bytes(), + "name", + )?); + } + } + } + } + let function_call = + raw_top_level_unique_object(delta.get().as_bytes(), "function_call").ok()?; + if let Some(function_call) = function_call { + out.extend(raw_top_level_string_values( + function_call.get().as_bytes(), + "name", + )?); + } + } + Some(out) +} + +#[derive(Clone, Debug, PartialEq, Eq)] +struct StreamContinuation { + key: String, + /// A semantic carrier family. A source form that has no durable carrier + /// identity is never allowed to merge with a different member of this + /// family on a later frame. + family: String, + identity: String, + identity_is_ambiguous: bool, + text: String, +} + +enum SourceContinuations { + Absent, + Ready(Vec), + Unevaluable, +} + +fn valid_stream_source_id(id: &str) -> bool { + !id.is_empty() && id.len() <= MAX_STREAM_GUARDRAIL_SOURCE_ID_BYTES +} + +fn bounded_stream_source_id(id: String) -> Result { + if !valid_stream_source_id(&id) { + Err(()) + } else { + Ok(id) + } +} + +fn append_source_branches( + out: &mut Vec, + seen_keys: &mut std::collections::HashSet, + source_values: &mut Vec, + family: String, + identity: String, + identity_is_ambiguous: bool, + values: Vec, +) -> Result<(), ()> { + let (first, last) = match values.as_slice() { + [] => return Ok(()), + [only] => (only, only), + [first, last] => (first, last), + _ => return Err(()), + }; + source_values.extend(values.iter().cloned()); + for (branch, value) in [("first", first), ("last", last)] { + let key = format!("{family}:{identity}:{branch}"); + if !seen_keys.insert(key.clone()) { + return Err(()); + } + out.push(StreamContinuation { + key, + family: family.clone(), + identity: identity.clone(), + identity_is_ambiguous, + text: value.clone(), + }); + } + Ok(()) +} + +fn source_values_match_expected(mut source_values: Vec, mut expected: Vec) -> bool { + // Object member order is not semantic. Keep duplicate counts, but do not + // reject a valid source frame merely because a provider serialised its + // tool fields before its text field. + source_values.sort_unstable(); + expected.sort_unstable(); + source_values == expected +} + +fn responses_stream_coordinates( + payload: &[u8], + needs_content_index: bool, +) -> Result<(String, String, String), ()> { + let item_id = raw_top_level_unique_string(payload, "item_id")? + .ok_or(()) + .and_then(bounded_stream_source_id)?; + let output_index = raw_top_level_unique_index(payload, "output_index")? + .ok_or(())? + .to_string(); + let content_index = if needs_content_index { + raw_top_level_unique_index(payload, "content_index")? + .ok_or(())? + .to_string() + } else { + String::new() + }; + Ok((item_id, output_index, content_index)) +} + +fn responses_output_item_coordinates( + payload: &[u8], + item: &RawJson<'_>, +) -> Result<(String, String), ()> { + let output_index = raw_top_level_unique_index(payload, "output_index")? + .ok_or(())? + .to_string(); + let nested_id = raw_top_level_unique_string(item.get().as_bytes(), "id")? + .ok_or(()) + .and_then(bounded_stream_source_id)?; + match raw_top_level_unique_string(payload, "item_id")? { + Some(top_level_id) if top_level_id != nested_id => Err(()), + Some(top_level_id) => Ok((bounded_stream_source_id(top_level_id)?, output_index)), + None => Ok((nested_id, output_index)), + } +} + +fn responses_content_part_selected_values(payload: &[u8]) -> Result, ()> { + let part = raw_top_level_unique_object(payload, "part")?.ok_or(())?; + let Some(field) = raw_top_level_unique_type(part.get().as_bytes())? + .as_deref() + .and_then(responses_visible_content_part_field) + else { + return Ok(Vec::new()); + }; + raw_top_level_string_values(part.get().as_bytes(), field).ok_or(()) +} + +fn responses_output_item_selected_values(item: &RawJson<'_>) -> Result, ()> { + let mut text = String::new(); + append_responses_output_item_strings(&mut text, item).ok_or(())?; + Ok((!text.is_empty()).then_some(text).into_iter().collect()) +} + +fn responses_terminal_selected_values(payload: &[u8]) -> Result, ()> { + let response = raw_top_level_unique_object(payload, "response")?.ok_or(())?; + let mut text = String::new(); + append_responses_output_strings(&mut text, response.get().as_bytes()).ok_or(())?; + Ok((!text.is_empty()).then_some(text).into_iter().collect()) +} + +/// The client-visible source values on one Responses event. Every arm is +/// type-aware: media/reasoning carriers remain opaque even when an upstream +/// sends the event without any preceding delta. +fn responses_stream_selected_values(payload: &[u8]) -> Result, ()> { + let Some(kind) = raw_top_level_unique_string(payload, "type")? else { + return Ok(Vec::new()); + }; + if let Some(event) = responses_direct_stream_event(&kind) { + let mut values = Vec::new(); + for field in event.fields { + values.extend(raw_top_level_string_values(payload, field).ok_or(())?); + } + return Ok(values); + } + match kind.as_str() { + "response.content_part.added" | "response.content_part.done" => { + responses_content_part_selected_values(payload) + } + "response.output_item.added" | "response.output_item.done" => { + let item = raw_top_level_unique_object(payload, "item")?.ok_or(())?; + responses_output_item_selected_values(&item) + } + "response.completed" | "response.incomplete" | "response.failed" => { + responses_terminal_selected_values(payload) + } + _ => Ok(Vec::new()), + } +} + +fn append_responses_source_continuation( + out: &mut Vec, + keys: &mut std::collections::HashSet, + family: String, + identity: String, + values: Vec, +) -> Result<(), ()> { + let mut source_values = Vec::new(); + append_source_branches( + out, + keys, + &mut source_values, + family, + identity, + false, + values, + ) +} + +fn responses_source_continuations(payload: &[u8]) -> SourceContinuations { + let kind = match raw_top_level_unique_string(payload, "type") { + Ok(Some(kind)) => kind, + Ok(None) => return SourceContinuations::Absent, + Err(()) => return SourceContinuations::Unevaluable, + }; + if !responses_stream_event_is_visible(&kind) { + return SourceContinuations::Absent; + } + + let mut out = Vec::new(); + let mut keys = std::collections::HashSet::new(); + let result = (|| -> Result<(), ()> { + if let Some(event) = responses_direct_stream_event(&kind) { + let fields = event + .fields + .iter() + .map(|field| { + Ok(( + *field, + raw_top_level_string_values(payload, field).ok_or(())?, + )) + }) + .collect::, ()>>()?; + if fields.iter().all(|(_, values)| values.is_empty()) { + return Ok(()); + } + let (item_id, output_index, content_index) = + responses_stream_coordinates(payload, event.content_index)?; + for (field, values) in fields { + append_responses_source_continuation( + &mut out, + &mut keys, + format!("responses:{item_id:?}"), + format!("{kind:?}:{output_index}:{content_index}:{field}"), + values, + )?; + } + return Ok(()); + } + match kind.as_str() { + "response.content_part.added" | "response.content_part.done" => { + let values = responses_content_part_selected_values(payload)?; + if values.is_empty() { + return Ok(()); + } + let (item_id, output_index, content_index) = + responses_stream_coordinates(payload, true)?; + append_responses_source_continuation( + &mut out, + &mut keys, + format!("responses:{item_id:?}"), + format!("{kind:?}:{output_index}:{content_index}:part"), + values, + ) + } + "response.output_item.added" | "response.output_item.done" => { + let item = raw_top_level_unique_object(payload, "item")?.ok_or(())?; + let values = responses_output_item_selected_values(&item)?; + if values.is_empty() { + return Ok(()); + } + let (item_id, output_index) = responses_output_item_coordinates(payload, &item)?; + append_responses_source_continuation( + &mut out, + &mut keys, + format!("responses:{item_id:?}"), + format!("{kind:?}:{output_index}:item"), + values, + ) + } + "response.completed" | "response.incomplete" | "response.failed" => { + let values = responses_terminal_selected_values(payload)?; + if values.is_empty() { + return Ok(()); + } + append_responses_source_continuation( + &mut out, + &mut keys, + "responses:terminal".to_owned(), + format!("{kind:?}:response"), + values, + ) + } + _ => Ok(()), + } + })(); + if result.is_err() { + SourceContinuations::Unevaluable + } else if out.is_empty() { + SourceContinuations::Absent + } else { + SourceContinuations::Ready(out) + } +} + +struct SourceBranchIdentity { + family: String, + identity: String, + identity_is_ambiguous: bool, +} + +fn append_raw_string_carrier( + out: &mut Vec, + keys: &mut std::collections::HashSet, + source_values: &mut Vec, + source: SourceBranchIdentity, + body: &[u8], + key: &str, +) -> Result<(), ()> { + append_source_branches( + out, + keys, + source_values, + source.family, + source.identity, + source.identity_is_ambiguous, + raw_top_level_string_values(body, key).ok_or(())?, + ) +} + +fn raw_part_identity(body: &[u8]) -> Result { + // Content arrays can repeat a visible text or refusal field. Their numeric index + // is the only canonical identity that remains stable when a provider + // later adds an optional `id`; id-only arrays use the unevaluable policy + // rather than silently switching source channels. + raw_top_level_unique_index(body, "index")? + .map(|index| format!("index:{index}")) + .ok_or(()) +} + +fn anthropic_source_continuations( + payload: &[u8], + kind: &str, + expected: Vec, +) -> SourceContinuations { + let index = match raw_top_level_unique_index(payload, "index") { + Ok(Some(index)) => index.to_string(), + _ => return SourceContinuations::Unevaluable, + }; + let family = format!("anthropic:{index}"); + let mut out = Vec::new(); + let mut keys = std::collections::HashSet::new(); + let mut source_values = Vec::new(); + let result = match kind { + "content_block_delta" => { + let delta = match raw_top_level_unique_object(payload, "delta") { + Ok(Some(delta)) => delta, + _ => return SourceContinuations::Unevaluable, + }; + let delta_body = delta.get().as_bytes(); + let kind = match raw_top_level_unique_type(delta_body) { + Ok(kind) => kind, + Err(()) => return SourceContinuations::Unevaluable, + }; + match kind.as_deref() { + Some("text_delta") => append_raw_string_carrier( + &mut out, + &mut keys, + &mut source_values, + SourceBranchIdentity { + family: family.clone(), + identity: "text".to_owned(), + identity_is_ambiguous: false, + }, + delta_body, + "text", + ), + Some("input_json_delta") => append_raw_string_carrier( + &mut out, + &mut keys, + &mut source_values, + SourceBranchIdentity { + family: family.clone(), + identity: "partial_json".to_owned(), + identity_is_ambiguous: false, + }, + delta_body, + "partial_json", + ), + Some(_) | None => Ok(()), + } + } + "content_block_start" => { + let block = match raw_top_level_unique_object(payload, "content_block") { + Ok(Some(block)) => block, + _ => return SourceContinuations::Unevaluable, + }; + let block_body = block.get().as_bytes(); + let kind = match raw_top_level_unique_type(block_body) { + Ok(kind) => kind, + Err(()) => return SourceContinuations::Unevaluable, + }; + match kind.as_deref() { + Some("text") => append_raw_string_carrier( + &mut out, + &mut keys, + &mut source_values, + SourceBranchIdentity { + family: family.clone(), + identity: "text".to_owned(), + identity_is_ambiguous: false, + }, + block_body, + "text", + ), + Some("tool_use") => { + let mut inputs = match raw_top_level_values(block_body, "input") { + Some(inputs) => inputs, + None => return SourceContinuations::Unevaluable, + }; + let Some(input) = inputs.pop() else { + return SourceContinuations::Absent; + }; + if !inputs.is_empty() { + return SourceContinuations::Unevaluable; + } + if input.get().trim_start().starts_with('"') { + append_raw_string_carrier( + &mut out, + &mut keys, + &mut source_values, + SourceBranchIdentity { + family: family.clone(), + identity: "input".to_owned(), + identity_is_ambiguous: false, + }, + block_body, + "input", + ) + } else { + // A nested tool input can have many source leaves but no + // durable leaf identity on this envelope. An empty object is + // harmless; a visible value must use the bounded policy. + let text = + match crate::json_splice::collect_string_values(input.get().as_bytes()) + { + Ok(text) => text, + Err(_) => return SourceContinuations::Unevaluable, + }; + if text.is_empty() { + Ok(()) + } else { + Err(()) + } + } + } + Some(_) | None => Ok(()), + } + } + _ => return SourceContinuations::Absent, + }; + if result.is_err() || !source_values_match_expected(source_values, expected) { + SourceContinuations::Unevaluable + } else if out.is_empty() { + SourceContinuations::Absent + } else { + SourceContinuations::Ready(out) + } +} + +fn chat_choice_source_continuations(payload: &[u8]) -> SourceContinuations { + let Some(expected) = decoded_chat_frame_continuations(payload) else { + return SourceContinuations::Absent; + }; + let kind = match raw_top_level_unique_string(payload, "type") { + Ok(Some(kind)) => Some(kind), + Ok(None) => None, + Err(()) => return SourceContinuations::Unevaluable, + }; + if let Some(kind @ ("content_block_delta" | "content_block_start")) = kind.as_deref() { + return anthropic_source_continuations(payload, kind, expected); + } + let choices = match raw_top_level_unique_array(payload, "choices") { + Ok(Some(choices)) => match raw_array_items(&choices) { + Some(choices) => choices, + None => return SourceContinuations::Unevaluable, + }, + Ok(None) | Err(()) => return SourceContinuations::Unevaluable, + }; + let mut out = Vec::new(); + let mut keys = std::collections::HashSet::new(); + let mut source_values = Vec::new(); + let mut choice_indexes = std::collections::HashSet::new(); + for choice in choices { + if !raw_is_object(&choice) { + return SourceContinuations::Unevaluable; + } + let choice_body = choice.get().as_bytes(); + let choice_index = match raw_top_level_unique_index(choice_body, "index") { + Ok(Some(index)) if choice_indexes.insert(index) => index.to_string(), + _ => return SourceContinuations::Unevaluable, + }; + let delta = match raw_top_level_unique_object(choice_body, "delta") { + Ok(Some(delta)) => delta, + Ok(None) => continue, + Err(()) => return SourceContinuations::Unevaluable, + }; + let delta_body = delta.get().as_bytes(); + let mut content = match raw_top_level_values(delta_body, "content") { + Some(content) => content, + None => return SourceContinuations::Unevaluable, + }; + if content.len() > 1 { + return SourceContinuations::Unevaluable; + } + if let Some(content) = content.pop() { + match content.get().trim_start().as_bytes().first() { + Some(b'"') => { + let value = match serde_json::from_str::(content.get()) { + Ok(value) => value, + Err(_) => return SourceContinuations::Unevaluable, + }; + if append_source_branches( + &mut out, + &mut keys, + &mut source_values, + format!("chat:{choice_index}:content"), + "scalar".to_owned(), + true, + vec![value], + ) + .is_err() + { + return SourceContinuations::Unevaluable; + } + } + Some(b'[') => { + let Some(parts) = raw_array_items(&content) else { + return SourceContinuations::Unevaluable; + }; + let mut part_ids = std::collections::HashSet::new(); + for part in parts { + if !raw_is_object(&part) { + return SourceContinuations::Unevaluable; + } + let part_body = part.get().as_bytes(); + let kind = match raw_top_level_unique_type(part_body) { + Ok(kind) => kind, + Err(()) => return SourceContinuations::Unevaluable, + }; + let Some(field) = kind.as_deref().and_then(chat_visible_content_part_field) + else { + continue; + }; + let visible_values = match raw_top_level_string_values(part_body, field) { + Some(visible_values) => visible_values, + None => return SourceContinuations::Unevaluable, + }; + if visible_values.is_empty() { + continue; + } + let part_id = match raw_part_identity(part_body) { + Ok(part_id) if part_ids.insert(part_id.clone()) => part_id, + _ => return SourceContinuations::Unevaluable, + }; + if append_source_branches( + &mut out, + &mut keys, + &mut source_values, + format!("chat:{choice_index}:content"), + format!("part:{part_id}:{field}"), + false, + visible_values, + ) + .is_err() + { + return SourceContinuations::Unevaluable; + } + } + } + Some(b'n') => {} + _ => return SourceContinuations::Unevaluable, + } + } + if append_raw_string_carrier( + &mut out, + &mut keys, + &mut source_values, + SourceBranchIdentity { + family: format!("chat:{choice_index}:refusal"), + identity: "refusal".to_owned(), + identity_is_ambiguous: false, + }, + delta_body, + "refusal", + ) + .is_err() + { + return SourceContinuations::Unevaluable; + } + let tool_calls = match raw_top_level_unique_array(delta_body, "tool_calls") { + Ok(tool_calls) => tool_calls, + Err(()) => return SourceContinuations::Unevaluable, + }; + if let Some(tool_calls) = tool_calls { + let Some(tool_calls) = raw_array_items(&tool_calls) else { + return SourceContinuations::Unevaluable; + }; + let mut tool_indexes = std::collections::HashSet::new(); + for tool_call in tool_calls { + if !raw_is_object(&tool_call) { + return SourceContinuations::Unevaluable; + } + let tool_body = tool_call.get().as_bytes(); + let tool_index = match raw_top_level_unique_index(tool_body, "index") { + Ok(Some(index)) if tool_indexes.insert(index) => index.to_string(), + _ => return SourceContinuations::Unevaluable, + }; + for (container, field) in match chat_tool_continuation_fields(tool_body) { + Some(fields) => fields, + None => return SourceContinuations::Unevaluable, + } { + let nested = match raw_top_level_unique_object(tool_body, container) { + Ok(nested) => nested, + Err(()) => return SourceContinuations::Unevaluable, + }; + if let Some(nested) = nested { + if append_raw_string_carrier( + &mut out, + &mut keys, + &mut source_values, + SourceBranchIdentity { + family: format!("chat:{choice_index}:tool:{tool_index}"), + identity: format!("{container}:{field}"), + identity_is_ambiguous: false, + }, + nested.get().as_bytes(), + field, + ) + .is_err() + { + return SourceContinuations::Unevaluable; + } + } + } + } + } + let function_call = match raw_top_level_unique_object(delta_body, "function_call") { + Ok(function_call) => function_call, + Err(()) => return SourceContinuations::Unevaluable, + }; + if let Some(function_call) = function_call { + if append_raw_string_carrier( + &mut out, + &mut keys, + &mut source_values, + SourceBranchIdentity { + family: format!("chat:{choice_index}:legacy_function"), + identity: "arguments".to_owned(), + identity_is_ambiguous: false, + }, + function_call.get().as_bytes(), + "arguments", + ) + .is_err() + { + return SourceContinuations::Unevaluable; + } + } + } + if source_values_match_expected(source_values, expected) { + SourceContinuations::Ready(out) + } else { + SourceContinuations::Unevaluable + } +} + +fn completions_source_continuations(payload: &[u8]) -> SourceContinuations { + let choices = match raw_top_level_unique_array_ref(payload, "choices") { + Ok(Some(choices)) => match raw_array_item_refs(&choices) { + Some(choices) => choices, + None => return SourceContinuations::Unevaluable, + }, + Ok(None) | Err(()) => return SourceContinuations::Unevaluable, + }; + let mut out = Vec::new(); + let mut keys = std::collections::HashSet::new(); + let mut source_values = Vec::new(); + let mut choice_indexes = std::collections::HashSet::new(); + for choice in choices { + if !raw_is_object(&choice) { + return SourceContinuations::Unevaluable; + } + let choice_body = choice.get().as_bytes(); + let choice_index = match raw_top_level_unique_index(choice_body, "index") { + Ok(Some(index)) if choice_indexes.insert(index) => index.to_string(), + _ => return SourceContinuations::Unevaluable, + }; + let text = match raw_top_level_unique_string(choice_body, "text") { + Ok(Some(text)) => vec![text], + Ok(None) => Vec::new(), + Err(()) => return SourceContinuations::Unevaluable, + }; + if append_source_branches( + &mut out, + &mut keys, + &mut source_values, + format!("completions:{choice_index}"), + "text".to_owned(), + false, + text, + ) + .is_err() + { + return SourceContinuations::Unevaluable; + } + } + if out.is_empty() { + SourceContinuations::Absent + } else { + SourceContinuations::Ready(out) + } +} + +fn stream_source_continuations( + protocol: PassthroughProtocol, + payload: &[u8], +) -> SourceContinuations { + let payload = payload.trim_ascii(); + if payload.is_empty() || payload == b"[DONE]" { + return SourceContinuations::Absent; + } + match protocol { + // An opaque protocol offers no carrier identity inside a JSON object + // or array. A bare JSON string is the one unambiguous source carrier; + // every broader Raw shape follows the configured unevaluable policy. + PassthroughProtocol::Raw => { + if payload.first() != Some(&b'"') { + return SourceContinuations::Unevaluable; + } + let Ok(mut values) = + crate::json_splice::collect_string_values_where_vec(payload, |_| true) + else { + return SourceContinuations::Unevaluable; + }; + let Some(text) = values.pop() else { + return SourceContinuations::Unevaluable; + }; + if !values.is_empty() { + return SourceContinuations::Unevaluable; + } + if text.is_empty() { + SourceContinuations::Absent + } else { + SourceContinuations::Ready(vec![StreamContinuation { + key: "raw:payload:first".to_owned(), + family: "raw".to_owned(), + identity: "payload".to_owned(), + identity_is_ambiguous: false, + text, + }]) + } + } + PassthroughProtocol::OpenaiChat => chat_choice_source_continuations(payload), + PassthroughProtocol::OpenaiCompletions => completions_source_continuations(payload), + PassthroughProtocol::OpenaiResponses => responses_source_continuations(payload), + } +} + +fn responses_terminal_continuation_prefix(payload: &[u8]) -> Result, ()> { + let payload = payload.trim_ascii(); + // `[DONE]` terminates an SSE stream but is not a JSON Responses event. + // Treat it like the source-continuation path does: no carrier closes and + // no malformed payload is introduced. Other malformed payloads still + // reach the fail-closed path below. + if payload.is_empty() || payload == b"[DONE]" { + return Ok(None); + } + let kind = match raw_top_level_unique_string(payload, "type") { + Ok(Some(kind)) => kind, + Ok(None) => return Ok(None), + Err(()) => return Err(()), + }; + if kind != "response.output_item.done" { + return Ok(None); + } + let top_level_id = raw_top_level_unique_string(payload, "item_id")?; + let nested_id = match raw_top_level_unique_object(payload, "item")? { + Some(item) => raw_top_level_unique_string(item.get().as_bytes(), "id")?, + None => None, + }; + if top_level_id + .as_ref() + .is_some_and(|item_id| !valid_stream_source_id(item_id)) + || nested_id + .as_ref() + .is_some_and(|item_id| !valid_stream_source_id(item_id)) + { + return Err(()); + } + let item_id = match (top_level_id, nested_id) { + (Some(top_level_id), Some(nested_id)) if top_level_id == nested_id => Some(top_level_id), + (Some(top_level_id), None) => Some(top_level_id), + (None, Some(nested_id)) => Some(nested_id), + (None, None) => None, + (Some(_), Some(_)) => return Err(()), + }; + Ok(item_id.map(|item_id| format!("responses:{item_id:?}:"))) +} + +fn frame_guardrail_supplemental_values( + protocol: PassthroughProtocol, + frame: &[u8], + has_source_continuations: bool, + has_typed_continuation: bool, +) -> Result, ()> { + // A detected Completions frame exposes no output carrier beyond + // `choices[].text`, which `completions_source_continuations` reads + // directly. Provider extensions are opaque, including when this frame + // has no text delta, so they never become a generic supplemental scan. + if matches!(protocol, PassthroughProtocol::OpenaiCompletions) { + return Ok(Vec::new()); + } + // A typed continuation without source proof becomes unevaluable; do not + // add a second, generic supplemental scan for that same frame. + if has_typed_continuation && !has_source_continuations { + return Ok(Vec::new()); + } + if !has_source_continuations { + return Ok(frame_guardrail_values(protocol, frame)); + } + + let Some(payload) = crate::redact::frame_payload(frame) else { + return Ok(Vec::new()); + }; + let payload = payload.trim(); + if payload.is_empty() || payload == "[DONE]" { + return Ok(Vec::new()); + } + match protocol { + // The raw source continuation is the complete decoded payload. + PassthroughProtocol::Raw => Ok(Vec::new()), + PassthroughProtocol::OpenaiChat => { + decoded_chat_frame_supplemental_values(payload.as_bytes()).ok_or(()) + } + PassthroughProtocol::OpenaiCompletions => Ok(Vec::new()), + // Responses source continuations cover every type-aware visible + // carrier, so no generic raw field is supplemental. + PassthroughProtocol::OpenaiResponses => Ok(Vec::new()), + } +} + +fn decoded_chat_frame_values(body: &[u8]) -> Option> { + decoded_chat_frame_supplemental_values(body) +} + +fn decoded_responses_frame_values(body: &[u8]) -> Option> { + responses_stream_selected_values(body).ok() +} + +fn frame_guardrail_values(protocol: PassthroughProtocol, frame: &[u8]) -> Vec { + let Some(payload) = crate::redact::frame_payload(frame) else { + return Vec::new(); + }; + let payload = payload.trim(); + if payload.is_empty() || payload == "[DONE]" { + return Vec::new(); + } + match protocol { + PassthroughProtocol::Raw => { + decoded_json_string_values_vec_where(payload.as_bytes(), |_| true) + .unwrap_or_else(|| vec![payload.to_string()]) + } + PassthroughProtocol::OpenaiChat => { + decoded_chat_frame_values(payload.as_bytes()).unwrap_or_default() + } + PassthroughProtocol::OpenaiCompletions => Vec::new(), + PassthroughProtocol::OpenaiResponses => { + decoded_responses_frame_values(payload.as_bytes()).unwrap_or_default() + } + } +} + +/// Guardrail-only text for a streamed frame. Capture and hold-back retain +/// their typed visible-content extraction in [`frame_parts`], while this +/// source-preserving pass retains duplicate selected carriers. Chat and +/// Responses media fields remain opaque even when forwarded verbatim. +#[cfg(test)] +fn frame_guardrail_text(protocol: PassthroughProtocol, frame: &[u8]) -> String { + let Some(payload) = crate::redact::frame_payload(frame) else { + return String::new(); + }; + let payload = payload.trim(); + if payload.is_empty() || payload == "[DONE]" { + return String::new(); + } + match protocol { + PassthroughProtocol::Raw => { + decoded_json_string_values(payload.as_bytes()).unwrap_or_else(|| payload.to_string()) + } + PassthroughProtocol::OpenaiChat => { + // A detected Chat frame can carry opaque multimodal values. + // Without a successful type-aware selection, relay it but do not + // send a raw fallback to an external output guardrail. + decoded_chat_frame_string_values(payload.as_bytes()).unwrap_or_default() + } + PassthroughProtocol::OpenaiCompletions => { + let mut scan_error = None; + decoded_completions_response_string_values(payload.as_bytes(), &mut scan_error) + .unwrap_or_default() + } + PassthroughProtocol::OpenaiResponses => { + // As with buffered Responses output, only a successful + // source-aware selection may cross the external guardrail + // boundary. A malformed frame stays forwarded but opaque. + decoded_responses_frame_string_values(payload.as_bytes()).unwrap_or_default() + } + } +} + +/// The independent source-identified channels scanned for one stream frame. +/// `supplemental` preserves selected non-carrier source values. Keeping them +/// separate prevents frame metadata from interrupting a sensitive literal +/// split across output deltas. +struct StreamGuardrailText { + continuations: Vec, + /// Values with no continuation carrier. They are checked separately so + /// unrelated JSON fields cannot become one regex/remote-model segment. + supplemental: Vec, + unevaluable: bool, + /// A supplementary selector failed after a source carrier was proven. + /// This is a local selection failure, not an upstream guardrail outage: + /// output `fail_open` must not turn it into an unscanned provider field. + supplemental_unevaluable: bool, + /// Responses item closures wait until their already-buffered text has + /// passed a guardrail scan. Removing them on the terminal event would + /// erase a short delta before an end-of-stream or full-buffer scan. + closed_prefixes: Vec, +} + +fn stream_guardrail_text( + protocol: PassthroughProtocol, + frame: &[u8], + continuation: String, +) -> StreamGuardrailText { + let payload = crate::redact::frame_payload(frame); + let hidden_reasoning = matches!(protocol, PassthroughProtocol::OpenaiChat) + && payload.as_ref().is_some_and(|payload| { + hidden_chat_stream_reasoning_frame(payload.trim().as_bytes()) == Some(true) + }); + let responses_visible_carrier = matches!(protocol, PassthroughProtocol::OpenaiResponses) + && payload.as_ref().is_some_and(|payload| { + matches!( + raw_top_level_unique_string(payload.trim().as_bytes(), "type"), + Ok(Some(kind)) if responses_stream_event_is_visible(&kind) + ) + }); + let typed_continuation = if hidden_reasoning + || (matches!(protocol, PassthroughProtocol::OpenaiResponses) && !responses_visible_carrier) + { + String::new() + } else if matches!(protocol, PassthroughProtocol::OpenaiChat) { + // `frame_parts` retains text-shaped fields for capture and hold-back, + // including ones on opaque multimodal parts. Its raw scan text must + // not make such a part an external-guardrail candidate or prevent a + // sibling, type-allowed text/tool value from being scanned. + payload + .as_ref() + .and_then(|payload| { + decoded_chat_frame_continuations(payload.trim().as_bytes()) + .map(|values| values.join("\n")) + }) + .unwrap_or(continuation) + } else { + continuation + }; + let has_typed_continuation = !typed_continuation.is_empty(); + let source = if hidden_reasoning { + SourceContinuations::Absent + } else { + payload + .as_ref() + .map_or(SourceContinuations::Absent, |payload| { + stream_source_continuations(protocol, payload.trim().as_bytes()) + }) + }; + let (closed_prefixes, terminal_unevaluable) = + if matches!(protocol, PassthroughProtocol::OpenaiResponses) { + match payload + .as_ref() + .map(|payload| responses_terminal_continuation_prefix(payload.trim().as_bytes())) + { + Some(Ok(Some(prefix))) => (vec![prefix], false), + Some(Ok(None)) | None => (Vec::new(), false), + Some(Err(())) => (Vec::new(), true), + } + } else { + (Vec::new(), false) + }; + let (continuations, has_source_continuations, source_unevaluable) = match source { + SourceContinuations::Ready(source) => { + let has_source_continuations = !source.is_empty(); + (source, has_source_continuations, false) + } + // A typed visible delta without a source carrier proof cannot be + // continued safely into the next frame. Do not create a generic + // positional fallback channel for it. + SourceContinuations::Absent if has_typed_continuation => (Vec::new(), false, true), + SourceContinuations::Absent => (Vec::new(), false, false), + // Do not assign an ordinal to an unkeyable source channel. The + // caller applies the configured fail-open/fail-closed policy. + SourceContinuations::Unevaluable => (Vec::new(), false, true), + }; + let mut unevaluable = source_unevaluable || terminal_unevaluable; + let mut supplemental_unevaluable = false; + let supplemental = if unevaluable { + Vec::new() + } else { + // A source-proofed frame can still carry supplemental selected + // strings (for example completions metadata). If that bounded walk + // fails, dropping it would turn a cap or malformed JSON into an + // unscanned bypass; let the caller apply its unevaluable policy. + match frame_guardrail_supplemental_values( + protocol, + frame, + has_source_continuations, + has_typed_continuation, + ) { + Ok(values) => values, + Err(()) => { + unevaluable = true; + supplemental_unevaluable = true; + Vec::new() + } + } + }; + StreamGuardrailText { + continuations, + supplemental, + unevaluable, + supplemental_unevaluable, + closed_prefixes, + } +} + +fn append_stream_guardrail_text( + continuations: &mut Vec, + _continuation_tails: &mut Vec, + supplemental: &mut Vec, + closed_prefixes: &mut Vec, + text: &StreamGuardrailText, +) { + for prefix in &text.closed_prefixes { + // Do not retain a terminal with no outstanding carrier: otherwise a + // stream of empty item.done events could grow the close set without + // contributing any text to the channel cap. + if continuations + .iter() + .chain(text.continuations.iter()) + .any(|continuation| continuation.key.starts_with(prefix)) + && !closed_prefixes.contains(prefix) + { + closed_prefixes.push(prefix.clone()); + } + } + for incoming in &text.continuations { + if let Some(existing) = continuations + .iter_mut() + .find(|continuation| continuation.key == incoming.key) + { + existing.text.push_str(&incoming.text); + } else { + continuations.push(incoming.clone()); + } + } + supplemental.extend( + text.supplemental + .iter() + .filter(|value| !value.is_empty()) + .cloned(), + ); +} + +fn stream_continuation_would_exceed_cap( + continuations: &[StreamContinuation], + supplemental: &[String], + _closed_prefixes: &[String], + text: &StreamGuardrailText, +) -> bool { + let mut keys: std::collections::HashSet<_> = continuations + .iter() + .map(|continuation| continuation.key.as_str()) + .collect(); + for continuation in &text.continuations { + keys.insert(continuation.key.as_str()); + } + let mut supplemental_candidates = supplemental + .iter() + .filter(|value| !value.is_empty()) + .map(String::as_str) + .collect::>(); + supplemental_candidates.extend( + text.supplemental + .iter() + .filter(|value| !value.is_empty()) + .map(String::as_str), + ); + keys.len() + supplemental_candidates.len() > MAX_STREAM_GUARDRAIL_CHANNELS +} + +fn stream_continuation_identity_conflicts( + continuations: &[StreamContinuation], + closed_prefixes: &[String], + text: &StreamGuardrailText, +) -> bool { + // A just-arrived `response.output_item.done` both carries an + // authoritative item snapshot and closes older deltas for that item. + // Its own prefix must not make the snapshot look like a post-close + // continuation; only prefixes established by an earlier frame do that. + let was_closed = |key: &str| closed_prefixes.iter().any(|prefix| key.starts_with(prefix)); + if text + .continuations + .iter() + .any(|continuation| was_closed(&continuation.key)) + { + return true; + } + let is_closed = |key: &str| { + was_closed(key) + || text + .closed_prefixes + .iter() + .any(|prefix| key.starts_with(prefix)) + }; + let mut all = continuations + .iter() + .filter(|continuation| !is_closed(&continuation.key)) + .collect::>(); + for incoming in &text.continuations { + if all.iter().any(|existing| { + existing.family == incoming.family + && existing.identity != incoming.identity + && (existing.identity_is_ambiguous || incoming.identity_is_ambiguous) + }) { + return true; + } + all.push(incoming); + } + false +} + +fn retire_scanned_stream_continuations( + continuations: &mut Vec, + continuation_tails: &mut Vec, + closed_prefixes: &mut Vec, +) { + continuations.retain(|continuation| { + !closed_prefixes + .iter() + .any(|prefix| continuation.key.starts_with(prefix)) + }); + continuation_tails.retain(|continuation| { + !closed_prefixes + .iter() + .any(|prefix| continuation.key.starts_with(prefix)) + }); + closed_prefixes.clear(); +} + +/// A fail-open frame with no provable source identity is a hard boundary: +/// nothing before it may be joined with a later keyed delta. The caller still +/// relays the frame, but starts a fresh guardrail scan epoch afterward. +fn reset_stream_guardrail_epoch( + continuations: &mut Vec, + continuation_tails: &mut Vec, + supplemental: &mut Vec, + closed_prefixes: &mut Vec, +) { + continuations.clear(); + continuation_tails.clear(); + supplemental.clear(); + closed_prefixes.clear(); +} + +/// Preserve the completed epoch for its own end-of-stream scan, then make an +/// unevaluable frame a hard continuity boundary for all following carriers. +/// Keeping epochs separate both retains monitor observations and prevents a +/// literal from joining across the unknown frame. +fn seal_stream_guardrail_epoch( + sealed_epochs: &mut Vec>, + queued_candidates: &mut usize, + continuations: &mut Vec, + continuation_tails: &mut Vec, + supplemental: &mut Vec, + closed_prefixes: &mut Vec, +) -> bool { + let candidates = stream_guardrail_scan_text(continuation_tails, continuations, supplemental); + let queued = candidates.is_empty() + || try_queue_stream_guardrail_epoch(sealed_epochs, queued_candidates, candidates); + reset_stream_guardrail_epoch( + continuations, + continuation_tails, + supplemental, + closed_prefixes, + ); + queued +} + +/// End-of-stream guardrails must keep epochs separate after an unevaluable +/// frame, but untrusted streams cannot queue arbitrarily many independent +/// external scans. An exhausted live fail-open stream records its bypass and +/// discards later candidates instead. +fn try_queue_stream_guardrail_epoch( + sealed_epochs: &mut Vec>, + queued_candidates: &mut usize, + candidates: Vec, +) -> bool { + if sealed_epochs.len() >= MAX_STREAM_GUARDRAIL_EPOCHS { + return false; + } + let Some(total) = queued_candidates.checked_add(candidates.len()) else { + return false; + }; + if total > MAX_STREAM_GUARDRAIL_CHANNELS { + return false; + } + *queued_candidates = total; + sealed_epochs.push(candidates); + true +} + +fn stream_guardrail_scan_text( + continuation_tails: &[StreamContinuation], + continuations: &[StreamContinuation], + supplemental: &[String], +) -> Vec { + let mut text = Vec::new(); + let mut scanned = std::collections::HashSet::new(); + for continuation in continuations { + let tail = continuation_tails + .iter() + .find(|candidate| candidate.key == continuation.key) + .map(|candidate| candidate.text.as_str()) + .unwrap_or_default(); + let candidate = format!("{tail}{}", continuation.text); + if !candidate.is_empty() && scanned.insert(candidate.clone()) { + text.push(candidate); + } + } + for value in supplemental { + if !value.is_empty() && scanned.insert(value.clone()) { + text.push(value.clone()); + } + } + text +} + +/// One frame's generated content, split by [`crate::held_content::Parts`] +/// into what the output guardrails scan and what only counts toward the +/// hold-back cap, plus any usage it reports. The scan and the cap read the +/// same extraction, so they cannot disagree about a frame (#513). +fn frame_parts( + protocol: PassthroughProtocol, + frame: &[u8], +) -> (crate::held_content::Parts, Option) { + let usage_labelled = + matches!(protocol, PassthroughProtocol::Raw) && is_usage_labelled_frame(frame); + let mut parts = crate::held_content::Parts::default(); + let mut usage: Option = None; + let mut merge = |found: PassthroughUsage| { + merge_usage(&mut usage, found); + }; + // This typed extraction reads and parses one complete payload per frame: a payload spread over several + // `data:` lines is one document joined with `\n`, so parsing each line + // independently produced N unparseable fragments — no usage read, and + // on a `Raw` stream the JSON source text pushed into the guardrail + // scan instead of the values (#1100). `frame_payload` also strips the + // per-line `\r` a CRLF-framed upstream leaves behind, and returns + // `None` for a comment-only frame (`: OPENROUTER PROCESSING`). + 'payload: { + let Some(payload) = crate::redact::frame_payload(frame) else { + break 'payload; + }; + let payload = payload.trim(); + if payload.is_empty() || payload == "[DONE]" { + break 'payload; + } + if matches!(protocol, PassthroughProtocol::Raw) { + // Raw payloads have no typed content envelope to extract. + // Scan them with the iterative value walker before touching + // serde_json::Value: its default recursion limit would otherwise + // turn a valid deeply nested escaped string into raw source text + // and let it bypass an output guardrail. + if let Ok(v) = serde_json::from_str::(payload) { + // An explicit `usage` object is self-describing even for an + // opaque stream. A server-labelled event additionally + // permits the flat agent-backend usage shape below. + if let Some(u) = v.get("usage").and_then(usage_of) { + merge(u); + } + if usage_labelled { + if let Some(u) = usage_of(&v) { + merge(u); + } + } + } + parts = crate::held_content::Parts { + scan: decoded_json_string_values(payload.as_bytes()) + .unwrap_or_else(|| payload.to_string()), + reasoning: 0, + }; + break 'payload; + } + let Ok(v) = serde_json::from_str::(payload) else { + // Unparseable joined payload — a non-conformant upstream that + // put two independent JSON documents on two `data:` lines, say. + // The frame is still FORWARDED, so scanning nothing here is a + // way past an output block rule. Fall back to the raw payload + // text on every protocol, not just `Raw`: over-scanning can only + // produce a false positive, while under-scanning a frame the + // client receives is the bypass. (Per-line parsing used to catch + // the two-document case incidentally; this covers it and every + // other shape that does not parse.) + parts.scan.push_str( + &decoded_json_string_values(payload.as_bytes()) + .unwrap_or_else(|| payload.to_string()), + ); + break 'payload; + }; + if let Some(u) = v.get("usage").and_then(usage_of) { + merge(u); + } + // Anthropic opens its stream with the prompt + cache counters + // nested on `message_start`. Gated on the event type so no other + // envelope's `message` object can be read as usage. + if v.get("type").and_then(|t| t.as_str()) == Some("message_start") { + if let Some(u) = v + .get("message") + .and_then(|m| m.get("usage")) + .and_then(usage_of) + { + merge(u); + } + } + if matches!(protocol, PassthroughProtocol::OpenaiResponses) { + // Responses streams carry usage on the terminal + // `response.completed` event's embedded response object. Read + // that shape ONLY here: another protocol's frame that happens + // to nest `response.usage` must not be read as usage. + if let Some(u) = v + .get("response") + .and_then(|r| r.get("usage")) + .and_then(usage_of) + { + merge(u); + } + } + if usage_labelled { + // The agent-backend shape: a flat token object on the + // server's own usage event, with no `usage` wrapper. + if let Some(u) = usage_of(&v) { + merge(u); + } + } + parts = match protocol { + PassthroughProtocol::Raw => unreachable!("handled before typed frame parsing"), + // The chat envelope also carries Anthropic Messages streams; the + // two event shapes are disjoint, so reading both is exact. + PassthroughProtocol::OpenaiChat => { + let mut p = crate::held_content::chat_chunk_parts(&v); + let a = crate::held_content::anthropic_event_parts(&v); + p.scan.push_str(&a.scan); + p.reasoning += a.reasoning; + p + } + PassthroughProtocol::OpenaiCompletions => { + crate::held_content::completions_chunk_parts(&v) + } + PassthroughProtocol::OpenaiResponses => crate::held_content::responses_event_parts(&v), + }; + } + (parts, usage) +} + +/// The text persisted for a streamed response. Raw passthrough capture keeps +/// the provider's JSON source representation even though the guardrail scans +/// decoded string values from that same frame. +fn frame_capture_text(protocol: PassthroughProtocol, frame: &[u8], scan: &str) -> String { + if !matches!(protocol, PassthroughProtocol::Raw) { + return scan.to_string(); + } + crate::redact::frame_payload(frame) + .map(|payload| payload.trim().to_owned()) + .filter(|payload| !payload.is_empty() && payload != "[DONE]") + .unwrap_or_else(|| scan.to_string()) +} + +/// The SSE error frame appended when an output guardrail blocks mid-relay, +/// in the protocol of the stream it ends: the frame `/v1/messages` emits for +/// the same refusal on an Anthropic Messages stream, the one +/// `/v1/chat/completions` emits on every other. +fn guardrail_error_frame( + anthropic: bool, + guardrail_name: Option<&str>, + unavailable: Option<&str>, +) -> Bytes { + if anthropic { + return Bytes::from(crate::messages::guardrail_block_frame( + guardrail_name, + unavailable, + )); + } + Bytes::from(format!( + "event: error\ndata: {}\n\n", + crate::error::guardrail_block_frame_payload(guardrail_name, unavailable) + )) +} + +/// Whether a relayed chat-envelope frame is an Anthropic Messages event +/// (`Some(true)`) or an OpenAI chat chunk (`Some(false)`). The two event +/// shapes are disjoint; a frame that is neither (a comment, `[DONE]`, an +/// in-band error) decides nothing. +fn anthropic_stream_frame(frame: &[u8]) -> Option { + let payload = crate::redact::frame_payload(frame)?; + let v: serde_json::Value = serde_json::from_str(payload.trim()).ok()?; + if v.get("choices").is_some() { + return Some(false); + } + match v.get("type").and_then(serde_json::Value::as_str) { + Some( + "message_start" + | "message_delta" + | "message_stop" + | "content_block_start" + | "content_block_delta" + | "content_block_stop" + | "ping", + ) => Some(true), + _ => None, + } +} + +/// Wait until a distributed stream-concurrency renewal proves that this +/// stream's member has disappeared from Redis. A backend outage is +/// deliberately not a signal: rate limiting remains fail-open when Redis +/// cannot establish whether the lease exists. +async fn wait_for_stream_lease_loss(receiver: &mut tokio::sync::watch::Receiver) { + loop { + if *receiver.borrow_and_update() { + return; + } + if receiver.changed().await.is_err() { + std::future::pending::<()>().await; + } + } +} + +/// Relay an opaque upstream SSE representation without parsing or mutating +/// its frames. This covers non-success error contracts and encoded successful +/// replies intentionally bypassed by an output fail-open policy. The telemetry +/// guard still fires from `Drop` when the client disconnects mid-relay. +#[allow(clippy::too_many_arguments)] +fn stream_opaque_response( + upstream_resp: reqwest::Response, + resp_headers: HeaderMap, + status: reqwest::StatusCode, + mut telemetry: RouteTelemetry, + request_id: &str, + stream_hold: aisix_ratelimit::StreamConcurrencyGuard, + stream_read_timeout: Option, +) -> Response { + use futures::StreamExt; + + let route_name = telemetry.route_name.clone(); + let mut lease_loss = stream_hold.lease_loss_receiver(); + let stream = async_stream::stream! { + let _stream_hold = stream_hold; + let read_timeout = crate::stream_timeout::ReadTimeoutSignal::default(); + let mut upstream = Box::pin(crate::stream_timeout::with_read_timeout_bytes_signalled( + upstream_resp.bytes_stream(), + stream_read_timeout, + read_timeout.clone(), + )); + loop { + let chunk = tokio::select! { + biased; + _ = wait_for_stream_lease_loss(&mut lease_loss) => { + telemetry.record_rate_limit_lease_lost(); + telemetry.stream_reached_end = true; + telemetry.emit(); + return; + } + chunk = upstream.next() => chunk, + }; + let Some(chunk) = chunk else { + break; + }; + match chunk { + Ok(chunk) => { + if telemetry.upstream_ttft_ms == 0 { + telemetry.upstream_ttft_ms = telemetry + .attempt_started + .elapsed() + .as_millis() + .min(u32::MAX as u128) as u32; + } + telemetry.mark_first_delivery(); + yield Ok::<_, std::convert::Infallible>(chunk); + } + Err(err) => { + let bridge = crate::dispatch::reqwest_error_to_bridge(&err, telemetry.started); + telemetry.record_failure(&bridge); + tracing::warn!( + route = %route_name, + error = %telemetry.error_message, + "passthrough-route opaque SSE relay failed mid-stream", + ); + break; + } + } + } + if let Some(err) = read_timeout.fired() { + telemetry.record_failure(&err); + tracing::warn!( + route = %route_name, + error = %telemetry.error_message, + "passthrough-route opaque SSE relay timed out mid-stream", + ); + } + telemetry.stream_reached_end = true; + telemetry.emit(); + }; + + let mut response = Response::builder() + .status(status) + .body(Body::from_stream(crate::request_id::in_request_span( + stream, + ))) + .unwrap(); + copy_safe_headers(&resp_headers, response.headers_mut()); + if let Ok(hv) = HeaderValue::from_str(request_id) { + response + .headers_mut() + .insert(header::HeaderName::from_static("x-aisix-request-id"), hv); + } + response +} + +/// Build a successful streamed relay response: upstream SSE frames are +/// forwarded incrementally, tee'd through the chain's [`StreamOutputPolicy`] +/// (window / full-buffer hold-back, end-of-stream check otherwise), while +/// usage and capture accumulate for the end-of-stream telemetry emit. The +/// telemetry guard also fires from `Drop` when the client disconnects +/// mid-relay. +#[allow(clippy::too_many_arguments)] +fn stream_response( + protocol: PassthroughProtocol, + chain: aisix_guardrails::GuardrailChain, + upstream_resp: reqwest::Response, + resp_headers: HeaderMap, + status: reqwest::StatusCode, + mut telemetry: RouteTelemetry, + request_id: &str, + stream_hold: aisix_ratelimit::StreamConcurrencyGuard, + stream_read_timeout: Option, +) -> Response { + use aisix_guardrails::{Guardrail as _, GuardrailVerdict, StreamOutputPolicy}; + use futures::StreamExt; + + let output_guardrail_active = aisix_guardrails::Guardrail::runs_on_output(&chain); + let policy = if !output_guardrail_active { + StreamOutputPolicy::EndOfStreamCheck + } else { + chain.stream_output_policy() + }; + let route_name = telemetry.route_name.clone(); + let capture_cap = telemetry.content_cap; + let splitter_cap = policy + .hold_cap() + .map(|(cap, _)| cap.saturating_mul(crate::held_content::RAW_HOLD_FACTOR)) + .unwrap_or(MAX_HELD_STREAM_BYTES); + let mut lease_loss = stream_hold.lease_loss_receiver(); + + let stream = async_stream::stream! { + // The rate limiter's reservation becomes an owned hold at the handoff + // from handler to body. It drops only when this body completes or the + // client cancels it, rather than when the response headers are built. + let _stream_hold = stream_hold; + let read_timeout = crate::stream_timeout::ReadTimeoutSignal::default(); + let mut upstream = Box::pin(crate::stream_timeout::with_read_timeout_bytes_signalled( + upstream_resp.bytes_stream(), + stream_read_timeout, + read_timeout.clone(), + )); + let mut splitter = SseFrameSplitter::with_max_frame_bytes(splitter_cap); + // Held-back frames (Window / BufferFull) not yet released. + let mut pending: Vec = Vec::new(); + let mut pending_held = crate::held_content::HeldBytes::default(); + // What Window and BufferFull hold (#513): content, which `max_buffer_bytes` + // caps (the SSE framing is not counted), and the raw frame bytes it + // bounds too. + let mut held_content = crate::held_content::HeldBuffer::default(); + // Responses repeats a logical carrier in delta, done, part, item, + // and terminal events. Keep a bounded identity ledger for the whole + // response so those representations do not consume the content cap + // repeatedly, including across Window releases. + let mut responses_held_content = crate::held_content::ResponsesHeldContent::default(); + // Each source-identified semantic delta stays contiguous across + // frames. Supplementary values remain individual scan candidates so + // unrelated fields cannot form one guardrail input. + let mut continuation_bufs: Vec = Vec::new(); + let mut supplemental_buf: Vec = Vec::new(); + // Overlap carried between Window scans, one per source channel. + let mut continuation_tails: Vec = Vec::new(); + // Item-done frames close a Responses carrier, but its buffered text + // remains until a successful scan has covered it. + let mut closed_continuation_prefixes: Vec = Vec::new(); + // Each unevaluable live frame seals the previous source epoch. Those + // epochs still need their own terminal scan, but must not concatenate + // with text that follows the unkeyable frame. + let mut sealed_guardrail_epochs: Vec> = Vec::new(); + let mut queued_guardrail_candidates = 0; + let mut scan_budget_exhausted = false; + // Degrades a hold-back policy to live forwarding after a fail-open cap hit. + let mut fail_opened = false; + let mut blocked = false; + // The chat envelope also carries Anthropic Messages streams; the + // first frame that says which one decides the refusal frame's shape. + let mut anthropic: Option = None; + + 'outer: loop { + let chunk = match tokio::select! { + biased; + _ = wait_for_stream_lease_loss(&mut lease_loss) => { + telemetry.record_rate_limit_lease_lost(); + telemetry.stream_reached_end = true; + telemetry.emit(); + return; + } + chunk = upstream.next() => chunk, + } { + Some(Ok(c)) => c, + Some(Err(err)) => { + // The response head is already on the wire, so there is + // no status left to carry the failure — record it on the + // event instead of ending as a silent success. + let bridge = crate::dispatch::reqwest_error_to_bridge(&err, telemetry.started); + telemetry.record_failure(&bridge); + tracing::warn!( + route = %route_name, + error = %telemetry.error_message, + "passthrough-route upstream stream failed mid-relay", + ); + break; + } + None => break, + }; + // TTFT on the first upstream chunk of any type — the same + // convention the typed streaming endpoints stamp. + if telemetry.upstream_ttft_ms == 0 { + telemetry.upstream_ttft_ms = telemetry + .attempt_started + .elapsed() + .as_millis() + .min(u32::MAX as u128) as u32; + } + for frame in splitter.push(&chunk) { + // Both completed and splitter-overflowed frames take the + // same raw-cap preflight below. + let _overflowed = frame.overflowed; + let frame = frame.bytes; + if !fail_opened { + if let Some((max_buffer_bytes, on_exceeded_fail_open)) = policy.hold_cap() { + // Apply the raw-byte bound before parsing every frame. + // The splitter marks an unterminated frame as overflowed, + // but a complete frame can be oversized as well. + if held_content.would_exceed_after(0, frame.len(), max_buffer_bytes) { + if on_exceeded_fail_open { + fail_opened = true; + chain.record_bypass(crate::error::TAG_OUTPUT_BUFFER_EXCEEDED); + for pending_frame in pending.drain(..) { + telemetry.mark_first_delivery(); + yield Ok(pending_frame); + } + pending_held.clear(); + held_content = crate::held_content::HeldBuffer::default(); + telemetry.mark_first_delivery(); + yield Ok(Bytes::from(frame)); + continue; + } + + tracing::warn!( + route = %route_name, + "passthrough-route stream exceeded the guardrail buffer cap (fail-closed)", + ); + blocked = true; + chain.record_output_buffer_exceeded(); + pending.clear(); + pending_held.clear(); + yield Ok(guardrail_error_frame(anthropic.unwrap_or(false), None, Some(crate::error::TAG_OUTPUT_BUFFER_EXCEEDED))); + break 'outer; + } + } + } + if anthropic.is_none() && matches!(protocol, PassthroughProtocol::OpenaiChat) { + anthropic = anthropic_stream_frame(&frame); + } + if let Some(err) = frame_in_band_error(protocol, &frame) { + telemetry.record_failure(&err); + } + let (parts, usage) = frame_parts(protocol, &frame); + let held = if matches!(protocol, PassthroughProtocol::OpenaiResponses) + && policy.hold_cap().is_some() + && !fail_opened + { + responses_held_content + .observe_sse_frame(&frame) + .unwrap_or_else(|| parts.held()) + } else { + parts.held() + }; + let delta = parts.scan; + if let Some(u) = usage { + merge_usage(&mut telemetry.usage, u); + } + if capture_cap.is_some() { + push_capped( + &mut telemetry.response_text, + &frame_capture_text(protocol, &frame, &delta), + capture_cap, + ); + } + // Source-preserving selectors impose their own bounded + // collection cap. A visible delta that already exceeds the + // configured hold cap is not an unscannable carrier: the + // established streaming contract reports + // `output_buffer_exceeded` before trying that selector. + if !fail_opened { + if let Some((max_buffer_bytes, on_exceeded_fail_open)) = policy.hold_cap() { + if held_content.would_exceed_after(held, frame.len(), max_buffer_bytes) { + if on_exceeded_fail_open { + fail_opened = true; + chain.record_bypass(crate::error::TAG_OUTPUT_BUFFER_EXCEEDED); + for pending_frame in pending.drain(..) { + telemetry.mark_first_delivery(); + yield Ok(pending_frame); + } + pending_held.clear(); + held_content = crate::held_content::HeldBuffer::default(); + telemetry.mark_first_delivery(); + yield Ok(Bytes::from(frame)); + continue; + } + + tracing::warn!( + route = %route_name, + "passthrough-route stream exceeded the guardrail buffer cap (fail-closed)", + ); + blocked = true; + chain.record_output_buffer_exceeded(); + pending.clear(); + pending_held.clear(); + yield Ok(guardrail_error_frame(anthropic.unwrap_or(false), None, Some(crate::error::TAG_OUTPUT_BUFFER_EXCEEDED))); + break 'outer; + } + } + } + let guardrail_text = (output_guardrail_active && !scan_budget_exhausted && !fail_opened) + .then(|| stream_guardrail_text(protocol, &frame, delta.clone())); + let unevaluable_output = guardrail_text.as_ref().is_some_and(|text| { + text.unevaluable + || stream_continuation_would_exceed_cap( + &continuation_bufs, + &supplemental_buf, + &closed_continuation_prefixes, + text, + ) + || stream_continuation_identity_conflicts( + &continuation_bufs, + &closed_continuation_prefixes, + text, + ) + }); + if unevaluable_output { + let supplemental_unevaluable = guardrail_text + .as_ref() + .is_some_and(|text| text.supplemental_unevaluable); + // A holding policy has already promised not to release a + // frame until it scans clean. `fail_open` can bypass an + // unevaluable live stream, but it cannot release the + // held prefix (or this frame) without a scan. + let must_refuse = supplemental_unevaluable + || (policy.holds_back() && !fail_opened) + || aisix_guardrails::Guardrail::refuses_unevaluable_output(&chain); + if must_refuse { + tracing::warn!( + guardrail_hook = "output", + route = %route_name, + "cannot preserve passthrough stream source continuity for guardrails; blocking", + ); + pending.clear(); + pending_held.clear(); + blocked = true; + yield Ok(guardrail_error_frame( + anthropic.unwrap_or(false), + None, + Some(crate::error::TAG_UNSCANNABLE_BODY), + )); + break 'outer; + } + chain.record_unevaluable_output_bypass(crate::error::TAG_UNSCANNABLE_BODY); + if !seal_stream_guardrail_epoch( + &mut sealed_guardrail_epochs, + &mut queued_guardrail_candidates, + &mut continuation_bufs, + &mut continuation_tails, + &mut supplemental_buf, + &mut closed_continuation_prefixes, + ) { + scan_budget_exhausted = true; + } + } + let frame = Bytes::from(frame); + match &policy { + _ if fail_opened => { + telemetry.mark_first_delivery(); + yield Ok::<_, std::convert::Infallible>(frame); + } + StreamOutputPolicy::EndOfStreamCheck => { + if let Some(text) = guardrail_text.as_ref().filter(|_| !unevaluable_output) { + append_stream_guardrail_text( + &mut continuation_bufs, + &mut continuation_tails, + &mut supplemental_buf, + &mut closed_continuation_prefixes, + text, + ); + } + telemetry.mark_first_delivery(); + yield Ok(frame); + } + StreamOutputPolicy::Window { + size_chars, + overlap_chars, + max_buffer_bytes, + on_exceeded_fail_open, + } => { + if let Some(text) = guardrail_text.as_ref().filter(|_| !unevaluable_output) { + append_stream_guardrail_text( + &mut continuation_bufs, + &mut continuation_tails, + &mut supplemental_buf, + &mut closed_continuation_prefixes, + text, + ); + } + held_content.hold(held, frame.len()); + pending_held.add(frame.len()); + pending.push(frame); + if held_content.exceeds(*max_buffer_bytes) { + if *on_exceeded_fail_open { + fail_opened = true; + chain.record_bypass(crate::error::TAG_OUTPUT_BUFFER_EXCEEDED); + for f in pending.drain(..) { + telemetry.mark_first_delivery(); + yield Ok(f); + } + pending_held.clear(); + held_content = crate::held_content::HeldBuffer::default(); + } else { + tracing::warn!( + route = %route_name, + "passthrough-route stream exceeded the guardrail buffer cap (fail-closed)", + ); + blocked = true; + chain.record_output_buffer_exceeded(); + pending.clear(); + pending_held.clear(); + yield Ok(guardrail_error_frame(anthropic.unwrap_or(false), None, Some(crate::error::TAG_OUTPUT_BUFFER_EXCEEDED))); + break 'outer; + } + } else if continuation_bufs + .iter() + .any(|continuation| continuation.text.chars().count() >= *size_chars) + || supplemental_buf + .iter() + .map(|value| value.chars().count()) + .sum::() + >= *size_chars + { + let candidates = stream_guardrail_scan_text( + &continuation_tails, + &continuation_bufs, + &supplemental_buf, + ); + match scan_output_candidates( + &chain, + &route_name, + &candidates, + &mut telemetry, + ) + .await + { + GuardrailVerdict::Block { + reason, + guardrail_name, + unavailable, + } => { + tracing::warn!( + guardrail_hook = "output", + route = %route_name, + reason = %reason, + "guardrail blocked passthrough-route stream (window)", + ); + blocked = true; + yield Ok(guardrail_error_frame(anthropic.unwrap_or(false), guardrail_name.as_deref(), unavailable.as_deref())); + break 'outer; + } + _ => { + for f in pending.drain(..) { + telemetry.mark_first_delivery(); + yield Ok(f); + } + pending_held.clear(); + held_content = crate::held_content::HeldBuffer::default(); + for continuation in &mut continuation_bufs { + let tail = continuation_tails + .iter() + .find(|tail| tail.key == continuation.key) + .map(|tail| tail.text.as_str()) + .unwrap_or_default(); + let combined = format!("{tail}{}", continuation.text); + let next_tail = tail_chars(&combined, *overlap_chars); + if let Some(tail) = continuation_tails + .iter_mut() + .find(|tail| tail.key == continuation.key) + { + tail.text = next_tail; + } else { + continuation_tails.push(StreamContinuation { + key: continuation.key.clone(), + family: continuation.family.clone(), + identity: continuation.identity.clone(), + identity_is_ambiguous: continuation + .identity_is_ambiguous, + text: next_tail, + }); + } + continuation.text.clear(); + } + supplemental_buf.clear(); + retire_scanned_stream_continuations( + &mut continuation_bufs, + &mut continuation_tails, + &mut closed_continuation_prefixes, + ); + } + } + } + } + StreamOutputPolicy::BufferFull { max_buffer_bytes, on_exceeded_fail_open } => { + if let Some(text) = guardrail_text.as_ref().filter(|_| !unevaluable_output) { + append_stream_guardrail_text( + &mut continuation_bufs, + &mut continuation_tails, + &mut supplemental_buf, + &mut closed_continuation_prefixes, + text, + ); + } + held_content.hold(held, frame.len()); + pending_held.add(frame.len()); + pending.push(frame); + if held_content.exceeds(*max_buffer_bytes) { + if *on_exceeded_fail_open { + fail_opened = true; + chain.record_bypass(crate::error::TAG_OUTPUT_BUFFER_EXCEEDED); + for f in pending.drain(..) { + telemetry.mark_first_delivery(); + yield Ok(f); + } + pending_held.clear(); + held_content = crate::held_content::HeldBuffer::default(); + } else { + tracing::warn!( + route = %route_name, + "passthrough-route stream exceeded the guardrail buffer cap (fail-closed)", + ); + blocked = true; + chain.record_output_buffer_exceeded(); + yield Ok(guardrail_error_frame(anthropic.unwrap_or(false), None, Some(crate::error::TAG_OUTPUT_BUFFER_EXCEEDED))); + break 'outer; + } + } + } + } + } + } + + if let Some(err) = read_timeout.fired() { + telemetry.record_failure(&err); + tracing::warn!( + route = %route_name, + error = %telemetry.error_message, + "passthrough-route upstream stream timed out mid-relay", + ); + } + + if !blocked { + // Trailing bytes with no frame terminator, plus the final scan + // of whatever the policy has not cleared yet. + let rest = splitter.take_rest(); + if !rest.is_empty() { + // The raw half of the hold-back cap is knowable without + // decoding the unterminated tail. Decide it before a parser + // can turn malformed JSON into an unrelated unscannable-body + // refusal or capture it into telemetry. + let tail_raw_cap_hit = match &policy { + StreamOutputPolicy::Window { + max_buffer_bytes, + on_exceeded_fail_open, + .. + } + | StreamOutputPolicy::BufferFull { + max_buffer_bytes, + on_exceeded_fail_open, + } if !fail_opened => held_content + .would_exceed_after(0, rest.len(), *max_buffer_bytes) + .then_some(*on_exceeded_fail_open), + _ => None, + }; + match tail_raw_cap_hit { + Some(false) => { + tracing::warn!( + route = %route_name, + "passthrough-route stream exceeded the guardrail buffer cap (fail-closed)", + ); + chain.record_output_buffer_exceeded(); + pending.clear(); + pending_held.clear(); + yield Ok(guardrail_error_frame(anthropic.unwrap_or(false), None, Some(crate::error::TAG_OUTPUT_BUFFER_EXCEEDED))); + telemetry.guardrail_blocked = true; + telemetry.stream_reached_end = true; + telemetry.emit(); + return; + } + Some(true) => { + fail_opened = true; + chain.record_bypass(crate::error::TAG_OUTPUT_BUFFER_EXCEEDED); + for frame in pending.drain(..) { + telemetry.mark_first_delivery(); + yield Ok(frame); + } + pending_held.clear(); + telemetry.mark_first_delivery(); + yield Ok(Bytes::from(rest)); + } + None => { + if anthropic.is_none() && matches!(protocol, PassthroughProtocol::OpenaiChat) { + anthropic = anthropic_stream_frame(&rest); + } + let (parts, usage) = frame_parts(protocol, &rest); + let held = if matches!(protocol, PassthroughProtocol::OpenaiResponses) + && policy.hold_cap().is_some() + && !fail_opened + { + responses_held_content + .observe_sse_frame(&rest) + .unwrap_or_else(|| parts.held()) + } else { + parts.held() + }; + let delta = parts.scan; + if let Some(u) = usage { + merge_usage(&mut telemetry.usage, u); + } + if capture_cap.is_some() { + push_capped( + &mut telemetry.response_text, + &frame_capture_text(protocol, &rest, &delta), + capture_cap, + ); + } + + // The raw cap was checked above before decoding. The regular + // held-frame path below additionally applies the decoded + // content cap to this tail. + let tail_cap_hit = match &policy { + StreamOutputPolicy::Window { + max_buffer_bytes, + on_exceeded_fail_open, + .. + } + | StreamOutputPolicy::BufferFull { + max_buffer_bytes, + on_exceeded_fail_open, + } if !fail_opened => held_content + .would_exceed_after(held, rest.len(), *max_buffer_bytes) + .then_some(*on_exceeded_fail_open), + _ => None, + }; + match tail_cap_hit { + Some(false) => { + tracing::warn!( + route = %route_name, + "passthrough-route stream exceeded the guardrail buffer cap (fail-closed)", + ); + chain.record_output_buffer_exceeded(); + pending.clear(); + pending_held.clear(); + yield Ok(guardrail_error_frame(anthropic.unwrap_or(false), None, Some(crate::error::TAG_OUTPUT_BUFFER_EXCEEDED))); + telemetry.guardrail_blocked = true; + telemetry.stream_reached_end = true; + telemetry.emit(); + return; + } + Some(true) => { + fail_opened = true; + chain.record_bypass(crate::error::TAG_OUTPUT_BUFFER_EXCEEDED); + for frame in pending.drain(..) { + telemetry.mark_first_delivery(); + yield Ok(frame); + } + pending_held.clear(); + telemetry.mark_first_delivery(); + yield Ok(Bytes::from(rest)); + } + None => { + let guardrail_text = + (output_guardrail_active && !scan_budget_exhausted && !fail_opened) + .then(|| stream_guardrail_text(protocol, &rest, delta.clone())); + let unevaluable_output = guardrail_text.as_ref().is_some_and(|text| { + text.unevaluable + || stream_continuation_would_exceed_cap( + &continuation_bufs, + &supplemental_buf, + &closed_continuation_prefixes, + text, + ) + || stream_continuation_identity_conflicts( + &continuation_bufs, + &closed_continuation_prefixes, + text, + ) + }); + if unevaluable_output { + let supplemental_unevaluable = guardrail_text + .as_ref() + .is_some_and(|text| text.supplemental_unevaluable); + // See the matching frame path above: a holding policy + // must not release its pending prefix unscanned just + // because this terminal fragment is unevaluable. + let must_refuse = supplemental_unevaluable + || (policy.holds_back() && !fail_opened) + || aisix_guardrails::Guardrail::refuses_unevaluable_output(&chain); + if must_refuse { + tracing::warn!( + guardrail_hook = "output", + route = %route_name, + "cannot preserve passthrough stream source continuity for guardrails; blocking", + ); + pending.clear(); + pending_held.clear(); + yield Ok(guardrail_error_frame( + anthropic.unwrap_or(false), + None, + Some(crate::error::TAG_UNSCANNABLE_BODY), + )); + telemetry.guardrail_blocked = true; + telemetry.stream_reached_end = true; + telemetry.emit(); + return; + } + chain.record_unevaluable_output_bypass(crate::error::TAG_UNSCANNABLE_BODY); + if !seal_stream_guardrail_epoch( + &mut sealed_guardrail_epochs, + &mut queued_guardrail_candidates, + &mut continuation_bufs, + &mut continuation_tails, + &mut supplemental_buf, + &mut closed_continuation_prefixes, + ) { + scan_budget_exhausted = true; + } + } + if let Some(text) = guardrail_text.as_ref().filter(|_| !unevaluable_output) { + append_stream_guardrail_text( + &mut continuation_bufs, + &mut continuation_tails, + &mut supplemental_buf, + &mut closed_continuation_prefixes, + text, + ); + } + let rest = Bytes::from(rest); + // The tail is held like any frame, under the same cap. + let tripped = match &policy { + StreamOutputPolicy::Window { + max_buffer_bytes, + on_exceeded_fail_open, + .. + } + | StreamOutputPolicy::BufferFull { + max_buffer_bytes, + on_exceeded_fail_open, + } if !fail_opened => { + held_content.hold(held, rest.len()); + held_content + .exceeds(*max_buffer_bytes) + .then_some(*on_exceeded_fail_open) + } + _ => None, + }; + match tripped { + Some(false) => { + tracing::warn!( + route = %route_name, + "passthrough-route stream exceeded the guardrail buffer cap (fail-closed)", + ); + chain.record_output_buffer_exceeded(); + pending.clear(); + pending_held.clear(); + yield Ok(guardrail_error_frame(anthropic.unwrap_or(false), None, Some(crate::error::TAG_OUTPUT_BUFFER_EXCEEDED))); + telemetry.guardrail_blocked = true; + telemetry.stream_reached_end = true; + telemetry.emit(); + return; + } + Some(true) => { + fail_opened = true; + chain.record_bypass(crate::error::TAG_OUTPUT_BUFFER_EXCEEDED); + for f in pending.drain(..) { + telemetry.mark_first_delivery(); + yield Ok(f); + } + pending_held.clear(); + telemetry.mark_first_delivery(); + yield Ok(rest); + } + None if policy.holds_back() && !fail_opened => { + pending_held.add(rest.len()); + pending.push(rest) + } + None => { + telemetry.mark_first_delivery(); + yield Ok(rest); + } + } + } + } + } + } + } + if output_guardrail_active && !fail_opened { + if !scan_budget_exhausted { + let candidates = stream_guardrail_scan_text( + &continuation_tails, + &continuation_bufs, + &supplemental_buf, + ); + if (!candidates.is_empty() || sealed_guardrail_epochs.is_empty()) + && !try_queue_stream_guardrail_epoch( + &mut sealed_guardrail_epochs, + &mut queued_guardrail_candidates, + candidates, + ) + { + chain.record_unevaluable_output_bypass(crate::error::TAG_UNSCANNABLE_BODY); + } + } + for candidates in sealed_guardrail_epochs { + if output_guardrail_active { + if let GuardrailVerdict::Block { + reason, + guardrail_name, + unavailable, + } = scan_output_candidates(&chain, &route_name, &candidates, &mut telemetry).await + { + tracing::warn!( + guardrail_hook = "output", + route = %route_name, + reason = %reason, + "guardrail blocked passthrough-route stream (end)", + ); + // Held frames are dropped (fail closed); content already + // forwarded under EndOfStreamCheck cannot be unsent — + // the error frame is the caller-visible signal either way. + pending.clear(); + pending_held.clear(); + yield Ok(guardrail_error_frame(anthropic.unwrap_or(false), guardrail_name.as_deref(), unavailable.as_deref())); + telemetry.guardrail_blocked = true; + telemetry.stream_reached_end = true; + telemetry.emit(); + return; + } + } + } + } + for f in pending.drain(..) { + telemetry.mark_first_delivery(); + yield Ok(f); + } + pending_held.clear(); + } else { + telemetry.guardrail_blocked = true; + } + // The generator ran to its own end (upstream EOF, upstream error, or + // a guardrail block); only a client that went away first leaves this + // unset, and the emit turns that into a 499. + telemetry.stream_reached_end = true; + telemetry.emit(); + }; + + // Re-attach the request span (the body is polled after the request-id + // middleware returns, so end-of-stream telemetry would otherwise log + // without a request_id) and heartbeat silence gaps — this branch is + // SSE-only, where a comment frame is protocol-legal and identical to + // what the typed endpoints emit; relayed frames are untouched. + let mut response = Response::builder() + .status(status) + .body(Body::from_stream(crate::sse_keepalive::with_heartbeat( + crate::request_id::in_request_span(stream), + crate::sse_keepalive::interval(), + ))) + .unwrap(); + copy_safe_headers(&resp_headers, response.headers_mut()); + // The relay re-chunks the body; a stale upstream length must not ride + // along (SSE normally has none, but a lying upstream shouldn't wedge + // the client). + response.headers_mut().remove(header::CONTENT_LENGTH); + if let Ok(hv) = HeaderValue::from_str(request_id) { + response + .headers_mut() + .insert(header::HeaderName::from_static("x-aisix-request-id"), hv); + } + response +} + +/// One output scan over `text`, folding monitor hits into the telemetry. +async fn scan_output( + chain: &aisix_guardrails::GuardrailChain, + route_name: &str, + text: &str, + telemetry: &mut RouteTelemetry, +) -> aisix_guardrails::GuardrailVerdict { + use aisix_guardrails::Guardrail as _; + let synth = aisix_gateway::ChatResponse { + id: String::new(), + model: route_name.to_string(), + message: aisix_gateway::ChatMessage::assistant(text.to_string()), + finish_reason: aisix_gateway::FinishReason::Stop, + usage: aisix_gateway::UsageStats::default(), + }; + let (verdict, hits) = chain.check_output_unmaskable_observed(&synth).await; + telemetry.monitor_hits.extend(hits); + verdict +} + +/// Scan each independently sourced candidate without letting a fail-open +/// result for one candidate skip a later candidate that another rule blocks. +async fn scan_output_candidates( + chain: &aisix_guardrails::GuardrailChain, + route_name: &str, + candidates: &[String], + telemetry: &mut RouteTelemetry, +) -> aisix_guardrails::GuardrailVerdict { + if candidates.is_empty() { + return scan_output(chain, route_name, "", telemetry).await; + } + for candidate in candidates { + if let verdict @ aisix_guardrails::GuardrailVerdict::Block { .. } = + scan_output(chain, route_name, candidate, telemetry).await + { + return verdict; + } + } + aisix_guardrails::GuardrailVerdict::Allow +} + +/// The last `n` chars of `s` (whole string when shorter). +fn tail_chars(s: &str, n: usize) -> String { + let count = s.chars().count(); + if count <= n { + return s.to_string(); + } + s.chars().skip(count - n).collect() +} + +/// Append `delta` to `buf`, bounded by `cap` bytes (capture accumulation +/// must not grow with an unbounded stream). Char boundaries are respected. +fn push_capped(buf: &mut String, delta: &str, cap: Option) { + let Some(cap) = cap else { return }; + if buf.len() >= cap { + return; + } + if buf.len() + delta.len() <= cap { + buf.push_str(delta); + return; + } + for c in delta.chars() { + if buf.len() + c.len_utf8() > cap { + break; + } + buf.push(c); + } +} + +// --------------------------------------------------------------------------- +// Telemetry +// --------------------------------------------------------------------------- + +/// End-of-request telemetry for a passthrough-route exchange: one +/// UsageEvent (CP sink + exporter fan-out, with captured content on the +/// exporter path only), the request metric, and the access log line. The +/// buffered path calls [`RouteTelemetry::emit`] inline; the streaming path +/// calls it at end-of-stream, with `Drop` covering client disconnects. +struct RouteTelemetry { + state: ProxyState, + route_name: String, + provider_label: String, + /// The request's trace bundle (AISIX-Cloud#1279) — the Drop emit is + /// the request's terminal emission, so it carries the terminal spans. + trace: Option>, + pk_id: String, + method: Method, + path: String, + request_id: String, + api_key_id: String, + /// Org member the authenticating key belongs to (AISIX-Cloud#1389), + /// and that member's display name for the `user_name` metric label + /// (AISIX-Cloud#1455). Both `None` for a key bound to no member — + /// including the anonymous route key, which belongs to the route + /// rather than to a person. + user_id: Option, + user_name: Option, + jwt: Option>, + /// Whether the caller reached this route through `auth_mode: + /// anonymous` rather than a credential of its own. Stamped onto the + /// usage event so anonymous traffic stays distinguishable from the + /// bound key's own (see `usage_attr::apply_auth_type`). + anonymous: bool, + client_identity: String, + client_source_ip: String, + client_user_agent: String, + started: Instant, + /// When the upstream call itself began — the scope the two `upstream_*` + /// figures are measured in, distinct from `started` (request receipt). + attempt_started: Instant, + status: u16, + /// Every token dimension the exchange reported, accumulated field-wise + /// across the response (buffered) or its frames (streamed). + usage: Option, + /// The model alias the caller addressed, read from a DETECTED + /// envelope's own `model` field — the same value the typed endpoint + /// serving that envelope records. Empty for an opaque body, whose + /// `model`-shaped key means nothing the gateway can trust. + requested_model: String, + /// Time from the START OF THE ATTEMPT to the upstream's first streamed + /// frame. Zero on the buffered path, where there is none. + upstream_ttft_ms: u32, + /// What the caller waited for on a streamed relay: the moment the first + /// relayed frame was handed downstream, measured from `started`. `None` + /// until one is, so a stream that delivered nothing reports no + /// caller-wait at all rather than an invented one. + downstream_first_ms: Option, + /// `true` once the relay generator reached its own end — upstream EOF, + /// upstream error, or a guardrail block. It stays `false` only when the + /// CLIENT went away first, which is what the emit turns into a 499 + /// (same signal the typed streaming endpoints record). + stream_reached_end: bool, + /// Set on a streamed relay so the `Drop` emit can tell an abandoned + /// stream from the buffered path, which never streams at all. + streaming: bool, + /// Bounded error class + message for a failure the relay could not + /// answer with a status code — an upstream that dies mid-stream, after + /// the response head is already on the wire. + error_class: String, + error_message: String, + /// The status that same failure gets before the response head + /// ([`aisix_gateway::BridgeError::http_status`]). The emit records it in + /// place of the upstream's `200`: the caller's response line cannot + /// change any more, but the record of what happened can. + failure_status: Option, + monitor_hits: Vec, + /// The request's ENFORCE-mode audit handle (AISIX-Cloud#1330). Held + /// rather than snapshotted at construction: this struct's emit runs + /// from the relay's `Drop`, long after the output hook has recorded + /// whatever it masked. + audit: crate::usage_attr::GuardrailAudit, + captured_prompt: Option, + content_cap: Option, + response_text: String, + guardrail_blocked: bool, + emitted: bool, +} + +impl RouteTelemetry { + /// Record the upstream failure that ended a streamed relay after its + /// head went out. The first one is the cause; later ones do not replace + /// it. + fn record_failure(&mut self, err: &aisix_gateway::BridgeError) { + if self.failure_status.is_some() { + return; + } + let failure = crate::attempt::StreamFailure::from_bridge(err); + self.error_class = failure.error_class.to_string(); + self.error_message = failure.error_message; + self.failure_status = Some(failure.status); + } + + /// A successful renewal can prove a stream still owns its shared + /// concurrency member; a missing member proves the opposite. The HTTP + /// response head has already been sent, so terminate at EOF and record a + /// distinct rate-limit outcome instead of misclassifying it as client + /// cancellation or an upstream transport failure. + fn record_rate_limit_lease_lost(&mut self) { + if self.failure_status.is_some() { + return; + } + self.error_class = "rate_limit_lease_lost".to_string(); + self.error_message = "distributed concurrency lease was lost".to_string(); + self.failure_status = Some(429); + } + + /// Stamp the caller's wait at the first RELAYED frame handed + /// downstream. + /// + /// Deliberately here and not where the frame was read off the upstream: + /// a hold-back guardrail policy sits between the two, and + /// `UsageEvent::downstream_latency_ms` counts that hold-back as part of + /// what the caller waited for. Called on both the live-forward and the + /// hold-back release paths, so it catches the first frame either way. + /// + /// A synthetic frame (a guardrail block's error event) deliberately + /// does NOT stamp: nothing the caller asked for was delivered. Same + /// rule as the typed streaming endpoints, which stamp only in the + /// chunk renderer. + fn mark_first_delivery(&mut self) { + if self.downstream_first_ms.is_none() { + self.downstream_first_ms = + Some(self.started.elapsed().as_millis().min(u32::MAX as u128) as u32); + } + } + + fn emit(&mut self) { + if self.emitted { + return; + } + self.emitted = true; + // A streamed relay the CLIENT abandoned never reached the + // generator's end. The upstream status is then not what happened + // to the request, so record the same 499 the typed streaming + // endpoints do rather than a success the caller never received. + // One an upstream failure ended records that failure's status + // instead, unless a guardrail refused it. + if self.streaming { + match self.failure_status.filter(|_| !self.guardrail_blocked) { + Some(status) => self.status = status, + None if !self.stream_reached_end => self.status = crate::CLIENT_CLOSED_REQUEST, + None => {} + } + } + let elapsed = self.started.elapsed(); + let snapshot = self.state.snapshot.load(); + let usage = self.usage.unwrap_or_default(); + + emit_access_log( + &self.method, + &self.path, + &self.route_name, + &self.api_key_id, + self.status, + // Same rule as the typed streaming endpoints, and the same + // figure this emit puts on the usage event below: a streamed + // relay reports the wait to its first relayed frame, a buffered + // one the whole response. A relay that delivered nothing waited + // the whole request for nothing, which is what `elapsed` says. + if self.streaming { + self.downstream_first_ms + .map(|ms| Duration::from_millis(u64::from(ms))) + .unwrap_or(elapsed) + } else { + elapsed + }, + elapsed, + &self.request_id, + Some(AccessLogTokens { + prompt: usage.prompt_tokens, + completion: usage.completion_tokens, + }), + None, + ); + + let pk = crate::usage_attr::ResolvedPk::resolve(&snapshot, &self.pk_id); + let caller = crate::request_metrics::Caller::from_api_key_id(&snapshot, &self.api_key_id); + crate::request_metrics::record( + &self.state, + ENDPOINT_LABEL, + caller.as_caller(), + crate::request_metrics::Upstream { + provider: &self.provider_label, + model: PASSTHROUGH_MODEL_LABEL, + pk: pk.labels(), + ..Default::default() + }, + self.status, + elapsed, + ); + + let mut event = aisix_obs::UsageEvent { + request_id: self.request_id.clone(), + occurred_at: aisix_obs::UsageEvent::occurred_at_now(), + api_key_id: self.api_key_id.clone(), + status_code: self.status, + requested_model: self.requested_model.clone(), + prompt_tokens: usage.prompt_tokens, + completion_tokens: usage.completion_tokens, + cached_prompt_tokens: usage.cached_prompt_tokens, + cache_write_tokens: usage.cache_write_tokens, + reasoning_tokens: usage.reasoning_tokens, + total_tokens: usage.upstream_total_tokens.unwrap_or(0), + cache_creation_tokens: usage.cache_creation_tokens, + cache_read_tokens: usage.cache_read_tokens, + upstream_latency_ms: self + .attempt_started + .elapsed() + .as_millis() + .min(u32::MAX as u128) as u32, + upstream_ttft_ms: self.upstream_ttft_ms, + // Streaming reports the caller's wait to the FIRST relayed + // frame, per the field's contract — a relay is a delivery + // mechanism for a response, not the response itself, so it is + // not the `/a2a` exception. Absent when a stream delivered + // nothing. The buffered path has no first frame: there the + // caller waited for the whole response to be written. + downstream_latency_ms: if self.streaming { + self.downstream_first_ms.unwrap_or(0) + } else { + elapsed.as_millis().min(u32::MAX as u128) as u32 + }, + error_class: std::mem::take(&mut self.error_class), + error_message: std::mem::take(&mut self.error_message), + inbound_protocol: "passthrough".to_string(), + passthrough_route_name: self.route_name.clone(), + client_identity: self.client_identity.clone(), + client_source_ip: self.client_source_ip.clone(), + client_user_agent: self.client_user_agent.clone(), + guardrail_blocked: self.guardrail_blocked, + guardrail_monitor_hits: std::mem::take(&mut self.monitor_hits), + applied_guardrails: crate::usage_attr::applied_guardrails(&self.audit), + guardrail_enforced_hits: crate::usage_attr::enforced_hits(&self.audit), + guardrail_scores: crate::usage_attr::guardrail_scores(&self.audit), + guardrail_bypassed_reason: crate::usage_attr::bypass_reason(&self.audit), + ..Default::default() + }; + crate::usage_attr::apply_pk_telemetry(&mut event, &pk); + crate::usage_attr::apply_caller_identity( + &mut event, + self.jwt.as_ref(), + self.user_id.as_deref(), + self.user_name.as_deref(), + ); + if self.anonymous { + event.auth_type = "anonymous".to_string(); + } + let usage_model = crate::usage_attr::usage_event_model_label( + // The snapshot loaded above: a config swap between two loads + // would make this label disagree with the emit's attribution. + &snapshot, + &event.requested_model, + ) + .into_owned(); + + // Captured content rides ONLY on the exporter fan-out, per the + // content_mode invariant (never the CP telemetry path). + let content = match (&self.captured_prompt, self.content_cap) { + (Some(prompt), Some(cap)) => Some(aisix_obs::CapturedContent::new( + prompt, + &self.response_text, + cap, + )), + _ => None, + }; + crate::usage_attr::emit_usage( + &self.state, + &snapshot, + crate::operation::PASSTHROUGH, + event, + crate::usage_attr::usage_event_labels(&usage_model, &pk), + content.as_ref(), + self.trace.as_ref(), + // The Drop emit is the request's end — body EOF or client drop. + /* terminal */ + true, + /* dispatched */ true, + ); + } +} + +impl Drop for RouteTelemetry { + fn drop(&mut self) { + if std::thread::panicking() { + return; + } + self.emit(); + } +} + +/// Copy response headers that are safe to relay to the downstream caller. +/// `append`, not `insert`: `HeaderMap` iteration yields one entry per +/// value, and a header the upstream sent several times (`Set-Cookie`, +/// `WWW-Authenticate`, `Vary`) must keep every value on a relay. +fn copy_safe_headers(src: &HeaderMap, dst: &mut HeaderMap) { + for (name, value) in src { + let n = name.as_str().to_lowercase(); + if matches!( + n.as_str(), + "transfer-encoding" + | "connection" + | "keep-alive" + | "proxy-authenticate" + | "proxy-authorization" + | "te" + | "trailer" + | "trailers" + | "upgrade" + ) { + continue; + } + dst.append(name.clone(), value.clone()); + } +} + +/// Token counts for one access-log line. `None` on the paths that never +/// reached an upstream, which is what keeps a rejected request out of the +/// token columns instead of logging it as a zero-token success. +struct AccessLogTokens { + prompt: u32, + completion: u32, +} + +#[allow(clippy::too_many_arguments)] +fn emit_access_log( + method: &Method, + path: &str, + route: &str, + api_key_id: &str, + status: u16, + // What the caller waited for: the first relayed frame on a streamed + // relay, the whole response otherwise — the same figure the usage + // event reports as `downstream_latency_ms`. + latency: Duration, + // How long the relay held the gateway, arrival to last byte out. On a + // streamed relay the two differ by the length of the stream. + duration: Duration, + request_id: &str, + tokens: Option, + error: Option<&ProxyError>, +) { + let (error_kind, error) = match error { + Some(e) => { + let (kind, msg) = crate::attempt::access_log_error(e); + (Some(kind), Some(msg)) + } + None => (None, None), + }; + let target = crate::attribution::AccessLogTarget::current(); + crate::attribution::emit_access_log(AccessLog { + method: method.as_str(), + path, + status, + latency, + duration, + provider: Some(route), + model: None, + upstream_model: target.upstream_model(), + provider_key_id: target.provider_key_id(), + api_key_id: Some(api_key_id), + prompt_tokens: tokens.as_ref().map(|t| u64::from(t.prompt)), + completion_tokens: tokens.as_ref().map(|t| u64::from(t.completion)), + total_tokens: tokens + .as_ref() + .map(|t| u64::from(t.prompt) + u64::from(t.completion)), + request_id, + provider_request_id: None, + served_by_model: None, + routing_attempt_count: None, + routing_fallback_count: None, + error_kind, + error: error.as_deref(), + mcp: None, + cache: None, + request_body_bytes: None, + response_body_bytes: None, + }); +} + +#[cfg(test)] +mod tests { + use super::*; + use aisix_core::resource::ResourceEntry; + use aisix_core::snapshot::SnapshotHandle; + use aisix_core::{AisixSnapshot, ApiKey, ProviderKey, ProxyConfig}; + use aisix_gateway::Hub; + use axum::body::to_bytes; + use axum::http::{Request, StatusCode}; + use std::sync::Arc; + use tower::ServiceExt; + use wiremock::matchers::{method as wm_method, path as wm_path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + fn scan_candidates_contain(candidates: &[String], expected: &str) -> bool { + candidates + .iter() + .any(|candidate| candidate.contains(expected)) + } + + fn cfg() -> ProxyConfig { + ProxyConfig { + addr: "127.0.0.1:0".into(), + request_body_limit_bytes: 1_048_576, + real_ip: Default::default(), + request_id: Default::default(), + url_rewrites: Vec::new(), + tls: None, + listeners: Vec::new(), + thread_per_core: None, + workers: None, + } + } + + const PK_ID: &str = "11111111-1111-1111-1111-111111111111"; + + /// A usage record carrying only the two canonical counters. + fn usage_dims(prompt: u32, completion: u32) -> PassthroughUsage { + PassthroughUsage { + prompt_tokens: prompt, + completion_tokens: completion, + ..Default::default() + } + } + + /// More than serde_json's default container recursion limit, while still + /// small enough to fit comfortably under the request body limit. + fn deeply_nested_escaped_block_json() -> Vec { + let mut json = "{\"v\":".repeat(160); + json.push_str(r#""\u0042LOCKME""#); + json.push_str(&"}".repeat(160)); + json.into_bytes() + } + + fn deeply_nested_responses_request_with_opaque_media() -> Vec { + let mut json = r#"{"input":[{"type":"message","content":[{"type":"input_image","image_url":"\u0042LOCKME"},{"type":"input_text","text":"clean"}]}],"metadata":"#.to_owned(); + json.push_str(&"{\"next\":".repeat(160)); + json.push_str(r#""deep""#); + json.push_str(&"}".repeat(160)); + json.push('}'); + json.into_bytes() + } + + fn nested_anthropic_tool_result_request(depth: usize, text: &str) -> Vec { + let text = serde_json::to_string(text).expect("test text serializes"); + let content = format!( + "{}[{{\"type\":\"text\",\"text\":{text}}}]{}", + r#"[{"type":"tool_result","tool_use_id":"t","content":"#.repeat(depth), + "}]".repeat(depth), + ); + format!(r#"{{"model":"claude","messages":[{{"role":"user","content":{content}}}]}}"#) + .into_bytes() + } + + fn nested_chat_messages_payload(depth: usize) -> Vec { + let metadata = format!( + "{}\"safe\"{}", + r#"{"next":"#.repeat(depth), + "}".repeat(depth), + ); + format!( + r#"{{"model":"chat","messages":[{{"role":"user","content":"safe","metadata":{metadata}}}]}}"# + ) + .into_bytes() + } + + fn provider_key_entry(api_base_unused: &str) -> ResourceEntry { + let json = format!( + r#"{{"display_name":"openai-up","secret":"sk-upstream","api_base":"{api_base_unused}","provider":"openai","adapter":"openai"}}"# + ); + let pk: ProviderKey = serde_json::from_str(&json).unwrap(); + ResourceEntry::new(PK_ID, pk, 1) + } + + fn apikey_entry(plaintext: &str, allowed_routes: Option<&[&str]>) -> ResourceEntry { + let routes = match allowed_routes { + Some(r) => format!( + r#", "allowed_routes": {}"#, + serde_json::to_string(r).unwrap() + ), + None => String::new(), + }; + let json = format!( + r#"{{"key_hash":"{}","allowed_models":["*"]{routes}}}"#, + ApiKey::hash_bearer(plaintext) + ); + let k: ApiKey = serde_json::from_str(&json).unwrap(); + ResourceEntry::new("k-1", k, 1) + } + + fn route_entry(id: &str, json: serde_json::Value) -> ResourceEntry { + let r: PassthroughRoute = serde_json::from_value(json).unwrap(); + ResourceEntry::new(id, r, 1) + } + + fn build_app(snap: AisixSnapshot) -> axum::Router { + let hub = Arc::new(Hub::new()); + let handle = SnapshotHandle::new(snap); + crate::build_router(crate::ProxyState::new(handle, hub, &cfg()).without_cache()) + } + + /// The `/passthrough/*` namespace carries no special case: with no + /// route claiming the path it is an ordinary router miss — a bare 404 + /// with an empty body, like any other unmatched path. + #[tokio::test] + async fn unclaimed_passthrough_path_takes_the_plain_404() { + let app = build_app(AisixSnapshot::new()); + let req = Request::builder() + .method("POST") + .uri("/passthrough/openai/v1/chat/completions") + .header("authorization", "Bearer whatever") + .body(axum::body::Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::NOT_FOUND); + let bytes = to_bytes(resp.into_body(), 65536).await.unwrap(); + assert!( + bytes.is_empty(), + "the miss path carries no error envelope, got {:?}", + String::from_utf8_lossy(&bytes) + ); + } + + #[tokio::test] + async fn unmatched_paths_keep_the_plain_404() { + let app = build_app(AisixSnapshot::new()); + let req = Request::builder() + .method("GET") + .uri("/definitely/not/a/route") + .body(axum::body::Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::NOT_FOUND); + } + + fn inject_route(target: &str) -> ResourceEntry { + route_entry( + "route-1", + serde_json::json!({ + "name": "openai-tunnel", + "path_prefix": "/passthrough/openai", + "target_url": target, + "provider_key_id": PK_ID + }), + ) + } + + #[tokio::test] + async fn inject_route_replaces_caller_auth_with_provider_key() { + let upstream = MockServer::start().await; + Mock::given(wm_method("GET")) + .and(wm_path("/v1/models")) + .and(wiremock::matchers::header( + "authorization", + "Bearer sk-upstream", + )) + .respond_with( + ResponseTemplate::new(200) + .set_body_json(serde_json::json!({"object": "list", "data": []})), + ) + .mount(&upstream) + .await; + + let snap = AisixSnapshot::new(); + snap.provider_keys + .insert(provider_key_entry("http://unused")); + snap.apikeys.insert(apikey_entry("sk-caller", Some(&["*"]))); + snap.passthrough_routes + .insert(inject_route(&upstream.uri())); + let app = build_app(snap); + + let req = Request::builder() + .method("GET") + .uri("/passthrough/openai/v1/models") + .header("authorization", "Bearer sk-caller") + .body(axum::body::Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + + // The caller's own Authorization must not have reached upstream. + let received = &upstream.received_requests().await.unwrap()[0]; + let auth_values: Vec<_> = received.headers.get_all("authorization").iter().collect(); + assert_eq!(auth_values.len(), 1); + } + + #[tokio::test] + async fn key_without_route_grant_is_403() { + let upstream = MockServer::start().await; + let snap = AisixSnapshot::new(); + snap.provider_keys + .insert(provider_key_entry("http://unused")); + snap.apikeys.insert(apikey_entry("sk-caller", None)); + snap.passthrough_routes + .insert(inject_route(&upstream.uri())); + let app = build_app(snap); + + let req = Request::builder() + .method("GET") + .uri("/passthrough/openai/v1/models") + .header("authorization", "Bearer sk-caller") + .body(axum::body::Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::FORBIDDEN); + let bytes = to_bytes(resp.into_body(), 65536).await.unwrap(); + let v: serde_json::Value = serde_json::from_slice(&bytes).unwrap(); + assert_eq!(v["error"]["type"], "permission_denied"); + } + + #[tokio::test] + async fn unauthenticated_route_request_is_401() { + let upstream = MockServer::start().await; + let snap = AisixSnapshot::new(); + snap.provider_keys + .insert(provider_key_entry("http://unused")); + snap.passthrough_routes + .insert(inject_route(&upstream.uri())); + let app = build_app(snap); + + let req = Request::builder() + .method("GET") + .uri("/passthrough/openai/v1/models") + .body(axum::body::Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } + + /// The forward-proxy shadowing case: a host-matched request whose path + /// collides with a typed gateway route must be served by the + /// passthrough route, not the typed handler. + #[tokio::test] + async fn host_match_wins_over_typed_route_on_colliding_path() { + let upstream = MockServer::start().await; + Mock::given(wm_method("POST")) + .and(wm_path("/v1/chat/completions")) + .respond_with( + ResponseTemplate::new(200).set_body_json(serde_json::json!({"routed": "byo"})), + ) + .mount(&upstream) + .await; + + let snap = AisixSnapshot::new(); + snap.apikeys.insert(apikey_entry("sk-caller", Some(&["*"]))); + // forward_client + header_key: Authorization belongs to the caller + // and must reach upstream verbatim. + snap.passthrough_routes.insert(route_entry( + "route-h", + serde_json::json!({ + "name": "byo-host", + "hosts": ["ai.example.com"], + "target_url": upstream.uri(), + "auth_mode": "header_key", + "auth_header_name": "x-aisix-api-key", + "credential_mode": "forward_client" + }), + )); + let app = build_app(snap); + + let req = Request::builder() + .method("POST") + .uri("/v1/chat/completions") + .header("host", "ai.example.com") + .header("authorization", "Bearer employee-official-token") + .header("x-aisix-api-key", "sk-caller") + .header("content-type", "application/json") + .body(axum::body::Body::from(r#"{"model":"gpt-4o"}"#)) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + let bytes = to_bytes(resp.into_body(), 65536).await.unwrap(); + let v: serde_json::Value = serde_json::from_slice(&bytes).unwrap(); + assert_eq!(v["routed"], "byo", "typed chat handler must not serve this"); + + // BYO: the employee credential reached upstream verbatim; the + // gateway's side-channel header did not. + let received = &upstream.received_requests().await.unwrap()[0]; + assert_eq!( + received.headers.get("authorization").unwrap(), + "Bearer employee-official-token" + ); + assert!(received.headers.get("x-aisix-api-key").is_none()); + } + + #[tokio::test] + async fn disabled_route_does_not_match() { + let upstream = MockServer::start().await; + let snap = AisixSnapshot::new(); + snap.provider_keys + .insert(provider_key_entry("http://unused")); + snap.apikeys.insert(apikey_entry("sk-caller", Some(&["*"]))); + let mut json = serde_json::json!({ + "name": "openai-tunnel", + "path_prefix": "/passthrough/openai", + "target_url": upstream.uri(), + "provider_key_id": PK_ID, + "enabled": false + }); + json["enabled"] = serde_json::Value::Bool(false); + snap.passthrough_routes.insert(route_entry("route-1", json)); + let app = build_app(snap); + + let req = Request::builder() + .method("GET") + .uri("/passthrough/openai/v1/models") + .header("authorization", "Bearer sk-caller") + .body(axum::body::Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + // Disabled → no match → the ordinary router miss. + assert_eq!(resp.status(), StatusCode::NOT_FOUND); + } + + #[tokio::test] + async fn anonymous_route_fails_closed_when_source_ip_is_unresolvable() { + let upstream = MockServer::start().await; + let snap = AisixSnapshot::new(); + snap.apikeys.insert(apikey_entry("sk-anon", Some(&["*"]))); + snap.passthrough_routes.insert(route_entry( + "route-a", + serde_json::json!({ + "name": "anon", + "path_prefix": "/anon", + "target_url": upstream.uri(), + "auth_mode": "anonymous", + "anonymous_key_id": "k-1", + "source_cidrs": ["0.0.0.0/0"], + "credential_mode": "forward_client" + }), + )); + let app = build_app(snap); + + // In-process requests resolve no client socket; an unparseable + // source must never satisfy the CIDR gate. + let req = Request::builder() + .method("GET") + .uri("/anon/x") + .body(axum::body::Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::FORBIDDEN); + } + + // ---- pure helpers ---- + + #[test] + fn path_prefix_matches_on_segment_boundary_only() { + assert!(path_under_prefix("/copilot", "/copilot")); + assert!(path_under_prefix("/copilot/chat", "/copilot")); + assert!(!path_under_prefix("/copilotx", "/copilot")); + } + + #[test] + fn joined_target_url_stays_under_its_configured_path() { + let joined = join_target_url( + "https://upstream.example/provider/v1", + "models", + Some("limit=3"), + ) + .unwrap(); + assert_eq!( + joined, + "https://upstream.example/provider/v1/models?limit=3" + ); + + let joined = join_target_url( + "https://upstream.example/provider/v1?fixed=1", + "models", + Some("limit=3"), + ) + .unwrap(); + assert_eq!( + joined, + "https://upstream.example/provider/v1/models?fixed=1&limit=3" + ); + assert_eq!( + strip_redundant_version_segment( + "https://upstream.example/provider/v1?fixed=1", + "v1/models", + ), + "models" + ); + + assert!( + join_target_url( + "https://upstream.example/provider/v1?tenant=operator", + "models", + Some("tenant=caller"), + ) + .is_err(), + "a caller must not override an operator-owned query key" + ); + assert!( + join_target_url( + "https://upstream.example/provider/v1?tenant=operator", + "models", + Some("%74enant=caller"), + ) + .is_err(), + "encoded query keys must not bypass the operator-owned key" + ); + assert!( + join_target_url( + "https://upstream.example/provider/v1?tenant=operator", + "models", + Some("%2574enant=caller"), + ) + .is_err(), + "nested-encoded query keys must not bypass the operator-owned key" + ); + assert!( + join_target_url( + "https://upstream.example/provider/v1?tenant=operator", + "models", + Some("safe=1%26tenant%3Dcaller"), + ) + .is_err(), + "an encoded query delimiter must not recreate an operator-owned key" + ); + assert!( + join_target_url( + "https://upstream.example/provider/v1?tenant=operator", + "models", + Some("safe=1%2526tenant%253Dcaller"), + ) + .is_err(), + "a nested-encoded query delimiter must not recreate an operator-owned key" + ); + for query in [ + "+tenant=caller", + "%20tenant=caller", + "%2520tenant=caller", + "tenant%00suffix=caller", + "tenant%2500suffix=caller", + ] { + assert!( + join_target_url( + "https://upstream.example/provider/v1?tenant=operator", + "models", + Some(query), + ) + .is_err(), + "{query} must not bypass an operator-owned key through PHP-style form-key registration" + ); + } + for query in [ + "safe=1;tenant=caller", + "safe=1%3Btenant%3Dcaller", + "safe=1%253Btenant%253Dcaller", + ] { + assert!( + join_target_url( + "https://upstream.example/provider/v1?tenant=operator", + "models", + Some(query), + ) + .is_err(), + "{query} must not recreate an operator-owned key through a semicolon delimiter" + ); + } + for query in [ + "tenant.id=caller", + "tenant%2Eid=caller", + "tenant%252Eid=caller", + "tenant+id=caller", + "tenant%20id=caller", + ] { + assert!( + join_target_url( + "https://upstream.example/provider/v1?tenant_id=operator", + "models", + Some(query), + ) + .is_err(), + "{query} must not bypass an operator-owned key through form-key normalization" + ); + } + assert!( + join_target_url( + "https://upstream.example/provider/v1?tenant=operator", + "models", + Some("tenant%5Brole%5D=caller"), + ) + .is_err(), + "a bracketed caller key must not bypass its operator-owned base key" + ); + + // Test raw, percent-encoded, and encoded-separator spellings. The + // gateway must reject them before a ProviderKey can be sent outside + // the route's configured target path. + for remainder in [ + "../models", + "%2e%2e/models", + "%2E%2E/models", + ".%2e/models", + "%2e./models", + "%2e%2e%2fmodels", + "%2e%2e%5cmodels", + "%252e%252e/models", + "%252e%252e%252fmodels", + "%252e%252e%255cmodels", + "..;ignored/models", + "%2e%2e%3bignored/models", + "%252e%252e%253bignored/models", + "%2e%2e%3bignored/%2e%2e%3bignored/admin", + ] { + assert!( + join_target_url("https://upstream.example/provider/v1", remainder, None).is_err(), + "{remainder} must not escape the configured target URL path" + ); + } + } + + #[tokio::test] + async fn unsafe_remainder_is_rejected_before_contacting_the_upstream() { + let upstream = MockServer::start().await; + let snap = AisixSnapshot::new(); + snap.provider_keys + .insert(provider_key_entry("http://unused")); + snap.apikeys.insert(apikey_entry("sk-caller", Some(&["*"]))); + snap.passthrough_routes + .insert(inject_route(&format!("{}/provider/v1", upstream.uri()))); + let app = build_app(snap); + + for remainder in [ + "../models", + "%2e%2e/models", + "%2e%2e%2fmodels", + "%252e%252e/models", + "%252e%252e%252fmodels", + ] { + let req = Request::builder() + .method("GET") + .uri(format!("/passthrough/openai/{remainder}")) + .header("authorization", "Bearer sk-caller") + .body(axum::body::Body::empty()) + .unwrap(); + let resp = app.clone().oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::BAD_REQUEST, "{remainder}"); + } + assert!( + upstream.received_requests().await.unwrap().is_empty(), + "an invalid joined path must not send the ProviderKey upstream" + ); + } + + #[test] + fn inbound_host_strips_port_and_lowercases() { + let req = Request::builder() + .uri("/x") + .header("host", "API.Example.COM:8443") + .body(axum::body::Body::empty()) + .unwrap(); + assert_eq!(inbound_host(&req).as_deref(), Some("api.example.com")); + } + + #[test] + fn longest_prefix_and_host_specificity_win() { + let snap = AisixSnapshot::new(); + let mk = |id: &str, json: serde_json::Value| { + snap.passthrough_routes.insert(route_entry(id, json)) + }; + mk( + "r-short", + serde_json::json!({"name":"short","path_prefix":"/p","target_url":"http://a","provider_key_id":"pk"}), + ); + mk( + "r-long", + serde_json::json!({"name":"long","path_prefix":"/p/deep","target_url":"http://b","provider_key_id":"pk"}), + ); + mk( + "r-host", + serde_json::json!({"name":"hosty","hosts":["h.example"],"target_url":"http://c","provider_key_id":"pk"}), + ); + + let m = match_route(&snap, None, "/p/deep/x").unwrap(); + assert_eq!(m.entry.value.name, "long"); + assert_eq!(m.remainder, "/x"); + assert!(m.prefix_matched); + + // Host match beats any path-only match. + let m = match_route(&snap, Some("h.example"), "/p/deep/x").unwrap(); + assert_eq!(m.entry.value.name, "hosty"); + assert_eq!(m.remainder, "/p/deep/x"); + assert!(!m.prefix_matched); + + // A preserve_host route narrowed by a prefix relays the WHOLE path: + // the prefix is a match condition on an upstream that owns its own + // path space, not a gateway mount point. GitHub Copilot's CLI needs + // this — its MCP server answers on /mcp/readonly of the same host it + // serves chat from, and a stripped "/readonly" 404s. + mk( + "r-mirror", + serde_json::json!({ + "name":"mirror","hosts":["m.example"],"path_prefix":"/mcp", + "preserve_host":true,"credential_mode":"forward_client" + }), + ); + let m = match_route(&snap, Some("m.example"), "/mcp/readonly").unwrap(); + assert_eq!(m.entry.value.name, "mirror"); + assert_eq!(m.remainder, "/mcp/readonly"); + assert!( + !m.prefix_matched, + "a mirrored path is never version-deduped" + ); + } + + #[test] + fn sse_splitter_emits_complete_frames_and_keeps_partials() { + let mut s = SseFrameSplitter::with_max_frame_bytes(MAX_HELD_STREAM_BYTES); + let frames = s.push(b"data: a\n\ndata: b\n\ndata: par"); + assert_eq!(frames.len(), 2); + assert_eq!(frames[0].bytes, b"data: a\n\n"); + assert!(!frames[0].overflowed); + let frames = s.push(b"tial\n\n"); + assert_eq!(frames.len(), 1); + assert_eq!(frames[0].bytes, b"data: partial\n\n"); + assert!(!frames[0].overflowed); + assert!(s.take_rest().is_empty()); + // CRLF boundaries too. + let mut s = SseFrameSplitter::with_max_frame_bytes(MAX_HELD_STREAM_BYTES); + let frames = s.push(b"data: x\r\n\r\nrest"); + assert_eq!(frames.len(), 1); + assert_eq!(s.take_rest(), b"rest"); + } + + #[test] + fn sse_splitter_and_usage_label_read_cr_framing() { + let frame = b"event: token_usage\rdata: {\"input_tokens\":3,\"output_tokens\":4}\r\r"; + let mut s = SseFrameSplitter::with_max_frame_bytes(MAX_HELD_STREAM_BYTES); + let frames = s.push(&[&frame[..], b"data: next"].concat()); + assert_eq!(frames.len(), 1); + assert_eq!(frames[0].bytes, frame); + assert!(!frames[0].overflowed); + assert_eq!(s.take_rest(), b"data: next"); + assert!(is_usage_labelled_frame(frame)); + let (_, usage) = frame_delta(PassthroughProtocol::Raw, frame); + assert!(usage.is_some(), "a CR-framed usage report is read"); + } + + #[test] + fn frame_in_band_error_reads_the_protocol_s_own_failure_events() { + let anthropic = b"event: error\ndata: {\"type\":\"error\",\"error\":{\"type\":\"overloaded_error\",\"message\":\"busy\"}}\n\n"; + let err = frame_in_band_error(PassthroughProtocol::OpenaiChat, anthropic).unwrap(); + // Anthropic documents 529 for overloaded; not a 4xx, so it maps to 502. + assert_eq!(err.http_status(), 502); + let openai = + br#"data: {"error":{"message":"slow down","type":"rate_limit_error","code":429}} + +"#; + let err = frame_in_band_error(PassthroughProtocol::OpenaiCompletions, openai).unwrap(); + assert_eq!(err.http_status(), 429); + let responses = br#"data: {"type":"response.failed","response":{"error":{"code":"server_error","message":"x"}}} + +"#; + assert!(frame_in_band_error(PassthroughProtocol::OpenaiResponses, responses).is_some()); + // An opaque stream is never read for one, and ordinary frames are not one. + assert!(frame_in_band_error(PassthroughProtocol::Raw, openai).is_none()); + let delta = br#"data: {"choices":[{"delta":{"content":"hel"}}]} + +"#; + assert!(frame_in_band_error(PassthroughProtocol::OpenaiChat, delta).is_none()); + } + + #[test] + fn frame_delta_extracts_chat_content_and_usage() { + let frame = br#"data: {"choices":[{"delta":{"content":"hel"}}]} + +"#; + let (text, usage) = frame_delta(PassthroughProtocol::OpenaiChat, frame); + assert_eq!(text, "hel"); + assert!(usage.is_none()); + + let done = br#"data: {"choices":[],"usage":{"prompt_tokens":7,"completion_tokens":3}} + +"#; + let (text, usage) = frame_delta(PassthroughProtocol::OpenaiChat, done); + assert_eq!(text, ""); + assert_eq!(usage, Some(usage_dims(7, 3))); + + let fim = br#"data: {"choices":[{"text":"def "}]} + +"#; + let (text, _) = frame_delta(PassthroughProtocol::OpenaiCompletions, fim); + assert_eq!(text, "def "); + } + + /// The recorded total is the upstream's own: adopted from the first + /// report, kept while every later report carries one, and dropped the + /// moment one does not — a surviving partial value would be a number + /// no upstream stated. + #[test] + fn merged_usage_keeps_a_total_only_while_every_report_carries_one() { + let with = usage_of( + &serde_json::json!({"prompt_tokens": 7, "completion_tokens": 3, "total_tokens": 10}), + ) + .unwrap(); + let without = usage_of(&serde_json::json!({"output_tokens": 5})).unwrap(); + assert_eq!(with.upstream_total_tokens, Some(10)); + assert_eq!(without.upstream_total_tokens, None); + + let mut acc = None; + merge_usage(&mut acc, with); + assert_eq!(acc.unwrap().upstream_total_tokens, Some(10)); + merge_usage(&mut acc, with); + assert_eq!(acc.unwrap().upstream_total_tokens, Some(10)); + merge_usage(&mut acc, without); + assert_eq!(acc.unwrap().upstream_total_tokens, None); + } + + #[test] + fn usage_of_reads_every_dimension_in_every_spelling() { + // OpenAI chat: the cache hit is nested under `prompt_tokens_details` + // and the reasoning count under `completion_tokens_details`. + let openai = serde_json::json!({ + "prompt_tokens": 100, + "completion_tokens": 20, + "prompt_tokens_details": {"cached_tokens": 80}, + "completion_tokens_details": {"reasoning_tokens": 12}, + }); + assert_eq!( + usage_of(&openai), + Some(PassthroughUsage { + prompt_tokens: 100, + completion_tokens: 20, + cached_prompt_tokens: 80, + reasoning_tokens: 12, + ..Default::default() + }) + ); + + // Responses API: the `input`/`output` spelling, details nested under + // the matching names. + let responses = serde_json::json!({ + "input_tokens": 30, + "output_tokens": 9, + "input_tokens_details": {"cached_tokens": 25}, + "output_tokens_details": {"reasoning_tokens": 4}, + }); + assert_eq!( + usage_of(&responses), + Some(PassthroughUsage { + prompt_tokens: 30, + completion_tokens: 9, + cached_prompt_tokens: 25, + reasoning_tokens: 4, + ..Default::default() + }) + ); + + // Anthropic: cache counters are separate, additive fields. + let anthropic = serde_json::json!({ + "input_tokens": 11, + "output_tokens": 5, + "cache_creation_input_tokens": 300, + "cache_read_input_tokens": 1200, + }); + assert_eq!( + usage_of(&anthropic), + Some(PassthroughUsage { + prompt_tokens: 11, + completion_tokens: 5, + cache_creation_tokens: 300, + cache_read_tokens: 1200, + ..Default::default() + }) + ); + + // DeepSeek reports the cache hit flat, and a ZEROED nested detail + // must not mask it (same precedence the typed OpenAI bridge uses). + let deepseek = serde_json::json!({ + "prompt_tokens": 40, + "completion_tokens": 6, + "prompt_tokens_details": {"cached_tokens": 0}, + "prompt_cache_hit_tokens": 32, + }); + assert_eq!(usage_of(&deepseek).unwrap().cached_prompt_tokens, 32); + + // The flat agent-backend shape, all five dimensions at the root. + let flat = serde_json::json!({ + "prompt_tokens": 14603, + "completion_tokens": 8, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 14272, + "reasoning_tokens": 8, + }); + assert_eq!( + usage_of(&flat), + Some(PassthroughUsage { + prompt_tokens: 14603, + completion_tokens: 8, + cache_read_tokens: 14272, + reasoning_tokens: 8, + ..Default::default() + }) + ); + + // An object with no recognised counter mints nothing. + assert_eq!(usage_of(&serde_json::json!({"disk": "80%"})), None); + assert_eq!(usage_of(&serde_json::Value::Null), None); } - fn emit(&mut self) { - if self.emitted { - return; - } - self.emitted = true; - // A streamed relay the CLIENT abandoned never reached the - // generator's end. The upstream status is then not what happened - // to the request, so record the same 499 the typed streaming - // endpoints do rather than a success the caller never received. - // One an upstream failure ended records that failure's status - // instead, unless a guardrail refused it. - if self.streaming { - match self.failure_status.filter(|_| !self.guardrail_blocked) { - Some(status) => self.status = status, - None if !self.stream_reached_end => self.status = crate::CLIENT_CLOSED_REQUEST, - None => {} - } - } - let elapsed = self.started.elapsed(); - let snapshot = self.state.snapshot.load(); - let usage = self.usage.unwrap_or_default(); + #[test] + fn anthropic_stream_reports_the_prompt_side_from_message_start() { + // Anthropic splits usage across two frames: `message_start` carries + // the input + cache counters, the terminal `message_delta` only the + // output ones. Reading the top level alone loses the prompt side. + let start = br#"data: {"type":"message_start","message":{"id":"msg_1","usage":{"input_tokens":12,"cache_creation_input_tokens":300,"cache_read_input_tokens":1200}}} - emit_access_log( - &self.method, - &self.path, - &self.route_name, - &self.api_key_id, - self.status, - // Same rule as the typed streaming endpoints, and the same - // figure this emit puts on the usage event below: a streamed - // relay reports the wait to its first relayed frame, a buffered - // one the whole response. A relay that delivered nothing waited - // the whole request for nothing, which is what `elapsed` says. - if self.streaming { - self.downstream_first_ms - .map(|ms| Duration::from_millis(u64::from(ms))) - .unwrap_or(elapsed) - } else { - elapsed - }, - elapsed, - &self.request_id, - Some(AccessLogTokens { - prompt: usage.prompt_tokens, - completion: usage.completion_tokens, - }), - None, - ); +"#; + let (_, start_usage) = frame_delta(PassthroughProtocol::OpenaiChat, start); + let start_usage = start_usage.expect("message_start must report usage"); + assert_eq!(start_usage.prompt_tokens, 12); + assert_eq!(start_usage.cache_creation_tokens, 300); + assert_eq!(start_usage.cache_read_tokens, 1200); - let pk = crate::usage_attr::ResolvedPk::resolve(&snapshot, &self.pk_id); - let caller = crate::request_metrics::Caller::from_api_key_id(&snapshot, &self.api_key_id); - crate::request_metrics::record( - &self.state, - ENDPOINT_LABEL, - caller.as_caller(), - crate::request_metrics::Upstream { - provider: &self.provider_label, - model: PASSTHROUGH_MODEL_LABEL, - pk: pk.labels(), + let delta = br#"data: {"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":7}} + +"#; + let (_, delta_usage) = frame_delta(PassthroughProtocol::OpenaiChat, delta); + let mut acc = start_usage; + acc.merge(delta_usage.expect("message_delta must report usage")); + // The later, partial report EXTENDS the record instead of + // truncating the prompt side to zero. + assert_eq!( + acc, + PassthroughUsage { + prompt_tokens: 12, + completion_tokens: 7, + cache_creation_tokens: 300, + cache_read_tokens: 1200, ..Default::default() - }, - self.status, - elapsed, + } ); - let mut event = aisix_obs::UsageEvent { - request_id: self.request_id.clone(), - occurred_at: aisix_obs::UsageEvent::occurred_at_now(), - api_key_id: self.api_key_id.clone(), - status_code: self.status, - requested_model: self.requested_model.clone(), - prompt_tokens: usage.prompt_tokens, - completion_tokens: usage.completion_tokens, - cached_prompt_tokens: usage.cached_prompt_tokens, - cache_write_tokens: usage.cache_write_tokens, - reasoning_tokens: usage.reasoning_tokens, - total_tokens: usage.upstream_total_tokens.unwrap_or(0), - cache_creation_tokens: usage.cache_creation_tokens, - cache_read_tokens: usage.cache_read_tokens, - upstream_latency_ms: self - .attempt_started - .elapsed() - .as_millis() - .min(u32::MAX as u128) as u32, - upstream_ttft_ms: self.upstream_ttft_ms, - // Streaming reports the caller's wait to the FIRST relayed - // frame, per the field's contract — a relay is a delivery - // mechanism for a response, not the response itself, so it is - // not the `/a2a` exception. Absent when a stream delivered - // nothing. The buffered path has no first frame: there the - // caller waited for the whole response to be written. - downstream_latency_ms: if self.streaming { - self.downstream_first_ms.unwrap_or(0) - } else { - elapsed.as_millis().min(u32::MAX as u128) as u32 - }, - error_class: std::mem::take(&mut self.error_class), - error_message: std::mem::take(&mut self.error_message), - inbound_protocol: "passthrough".to_string(), - passthrough_route_name: self.route_name.clone(), - client_identity: self.client_identity.clone(), - client_source_ip: self.client_source_ip.clone(), - client_user_agent: self.client_user_agent.clone(), - guardrail_blocked: self.guardrail_blocked, - guardrail_monitor_hits: std::mem::take(&mut self.monitor_hits), - applied_guardrails: crate::usage_attr::applied_guardrails(&self.audit), - guardrail_enforced_hits: crate::usage_attr::enforced_hits(&self.audit), - guardrail_scores: crate::usage_attr::guardrail_scores(&self.audit), - guardrail_bypassed_reason: crate::usage_attr::bypass_reason(&self.audit), - ..Default::default() - }; - crate::usage_attr::apply_pk_telemetry(&mut event, &pk); - crate::usage_attr::apply_caller_identity( - &mut event, - self.jwt.as_ref(), - self.user_id.as_deref(), - self.user_name.as_deref(), - ); - if self.anonymous { - event.auth_type = "anonymous".to_string(); - } - let usage_model = crate::usage_attr::usage_event_model_label( - // The snapshot loaded above: a config swap between two loads - // would make this label disagree with the emit's attribution. - &snapshot, - &event.requested_model, - ) - .into_owned(); + // The nested read is gated on the event type: a `message` object on + // any other frame is not a usage report. + let other = br#"data: {"type":"conversation","message":{"usage":{"input_tokens":999}}} - // Captured content rides ONLY on the exporter fan-out, per the - // content_mode invariant (never the CP telemetry path). - let content = match (&self.captured_prompt, self.content_cap) { - (Some(prompt), Some(cap)) => Some(aisix_obs::CapturedContent::new( - prompt, - &self.response_text, - cap, - )), - _ => None, - }; - crate::usage_attr::emit_usage( - &self.state, - &snapshot, - crate::operation::PASSTHROUGH, - event, - crate::usage_attr::usage_event_labels(&usage_model, &pk), - content.as_ref(), - self.trace.as_ref(), - // The Drop emit is the request's end — body EOF or client drop. - /* terminal */ - true, - /* dispatched */ true, +"#; + assert_eq!(frame_delta(PassthroughProtocol::OpenaiChat, other).1, None); + } + + /// A frame's payload is ALL of its `data:` lines joined with `\n` + /// (WHATWG SSE). Parsing each line on its own turns one document into + /// N unparseable fragments, so the frame's usage went unread and — on + /// a `Raw` stream — its JSON source text was pushed into the guardrail + /// scan instead of its values. + #[test] + fn a_payload_spread_over_several_data_lines_is_read_as_one_document() { + let frame = b"event: message_delta\ndata: {\"type\":\"message_delta\",\ndata: \"usage\":{\"output_tokens\":7,\"input_tokens\":12}}\n\n"; + let (_, usage) = frame_delta(PassthroughProtocol::OpenaiChat, frame); + assert_eq!( + usage, + Some(PassthroughUsage { + prompt_tokens: 12, + completion_tokens: 7, + ..Default::default() + }), ); + + // The `Raw` scan text is the payload's VALUES for a document that + // parses — never the raw JSON source, which is what a per-line read + // fell back to for each fragment. + let (text, _) = frame_delta(PassthroughProtocol::Raw, frame); + assert_eq!(text, "message_delta"); } -} -impl Drop for RouteTelemetry { - fn drop(&mut self) { - if std::thread::panicking() { - return; - } - self.emit(); + /// Framing varies per ENDPOINT, not per vendor: on one host + /// `/v1/audio/transcriptions` streams pure CRLF with `\r\n\r\n` + /// separators and no `event:` lines while `/v1/responses` on the same + /// host is pure LF. The `\r` belongs to the framing and must reach + /// neither the parser nor the scan text. + #[test] + fn a_crlf_framed_frame_reads_the_same_as_its_lf_twin() { + let crlf = b"data: {\"usage\":{\"prompt_tokens\":26,\"completion_tokens\":4}}\r\n\r\n"; + let lf = b"data: {\"usage\":{\"prompt_tokens\":26,\"completion_tokens\":4}}\n\n"; + assert_eq!( + frame_delta(PassthroughProtocol::Raw, crlf), + frame_delta(PassthroughProtocol::Raw, lf), + ); + assert_eq!( + frame_delta(PassthroughProtocol::Raw, crlf).1, + Some(usage_dims(26, 4)), + ); + // …and the frame splitter agrees about where such a frame ends. + let mut splitter = SseFrameSplitter::with_max_frame_bytes(MAX_HELD_STREAM_BYTES); + let frames = splitter.push(crlf); + assert_eq!(frames.len(), 1); + assert_eq!(frames[0].bytes, crlf); + assert!(!frames[0].overflowed); } -} -/// Copy response headers that are safe to relay to the downstream caller. -/// `append`, not `insert`: `HeaderMap` iteration yields one entry per -/// value, and a header the upstream sent several times (`Set-Cookie`, -/// `WWW-Authenticate`, `Vary`) must keep every value on a relay. -fn copy_safe_headers(src: &HeaderMap, dst: &mut HeaderMap) { - for (name, value) in src { - let n = name.as_str().to_lowercase(); - if matches!( - n.as_str(), - "transfer-encoding" - | "connection" - | "keep-alive" - | "proxy-authenticate" - | "proxy-authorization" - | "te" - | "trailers" - | "upgrade" - ) { - continue; + /// A comment-only frame — the keepalive some relays emit while the + /// upstream thinks — carries no `data:` line at all. It must contribute + /// no usage and no scan text on every protocol, rather than being read + /// as an empty or unparseable payload. + #[test] + fn a_comment_only_frame_contributes_nothing() { + for frame in [ + &b": OPENROUTER PROCESSING\n\n"[..], + &b": OPENROUTER PROCESSING\r\n\r\n"[..], + &b": keep-alive\nevent: ping\n\n"[..], + ] { + for protocol in [ + PassthroughProtocol::Raw, + PassthroughProtocol::OpenaiChat, + PassthroughProtocol::OpenaiCompletions, + PassthroughProtocol::OpenaiResponses, + ] { + assert_eq!( + frame_delta(protocol, frame), + (String::new(), None), + "{protocol:?} on {:?}", + String::from_utf8_lossy(frame), + ); + } } - dst.append(name.clone(), value.clone()); } -} - -/// Token counts for one access-log line. `None` on the paths that never -/// reached an upstream, which is what keeps a rejected request out of the -/// token columns instead of logging it as a zero-token success. -struct AccessLogTokens { - prompt: u32, - completion: u32, -} -#[allow(clippy::too_many_arguments)] -fn emit_access_log( - method: &Method, - path: &str, - route: &str, - api_key_id: &str, - status: u16, - // What the caller waited for: the first relayed frame on a streamed - // relay, the whole response otherwise — the same figure the usage - // event reports as `downstream_latency_ms`. - latency: Duration, - // How long the relay held the gateway, arrival to last byte out. On a - // streamed relay the two differ by the length of the stream. - duration: Duration, - request_id: &str, - tokens: Option, - error: Option<&ProxyError>, -) { - let (error_kind, error) = match error { - Some(e) => { - let (kind, msg) = crate::attempt::access_log_error(e); - (Some(kind), Some(msg)) + /// A frame whose joined payload does not parse is still FORWARDED to + /// the client, so producing no scan text for it is a way past an output + /// block rule. Every protocol falls back to the raw payload text — the + /// worst case is a false positive, while the alternative is a bypass. + #[test] + fn an_unparseable_payload_still_yields_scan_text_on_every_protocol() { + // Two independent JSON documents on two `data:` lines: joined per + // the SSE spec this is one unparseable payload, and per-line parsing + // used to catch it only incidentally. + let frame = b"data: {\"choices\":[{\"delta\":{\"content\":\"BLOCKME\"}}]}\ndata: {\"choices\":[]}\n\n"; + for protocol in [ + PassthroughProtocol::Raw, + PassthroughProtocol::OpenaiChat, + PassthroughProtocol::OpenaiCompletions, + PassthroughProtocol::OpenaiResponses, + ] { + let (text, _) = frame_delta(protocol, frame); + assert!( + text.contains("BLOCKME"), + "{protocol:?} must still offer the forwarded bytes to the scan, got {text:?}", + ); } - None => (None, None), - }; - let target = crate::attribution::AccessLogTarget::current(); - crate::attribution::emit_access_log(AccessLog { - method: method.as_str(), - path, - status, - latency, - duration, - provider: Some(route), - model: None, - upstream_model: target.upstream_model(), - provider_key_id: target.provider_key_id(), - api_key_id: Some(api_key_id), - prompt_tokens: tokens.as_ref().map(|t| u64::from(t.prompt)), - completion_tokens: tokens.as_ref().map(|t| u64::from(t.completion)), - total_tokens: tokens - .as_ref() - .map(|t| u64::from(t.prompt) + u64::from(t.completion)), - request_id, - provider_request_id: None, - served_by_model: None, - routing_attempt_count: None, - routing_fallback_count: None, - error_kind, - error: error.as_deref(), - mcp: None, - cache: None, - request_body_bytes: None, - response_body_bytes: None, - }); -} -#[cfg(test)] -mod tests { - use super::*; - use aisix_core::resource::ResourceEntry; - use aisix_core::snapshot::SnapshotHandle; - use aisix_core::{AisixSnapshot, ApiKey, ProviderKey, ProxyConfig}; - use aisix_gateway::Hub; - use axum::body::to_bytes; - use axum::http::{Request, StatusCode}; - use std::sync::Arc; - use tower::ServiceExt; - use wiremock::matchers::{method as wm_method, path as wm_path}; - use wiremock::{Mock, MockServer, ResponseTemplate}; + let anthropic = b"data: {\"type\":\"content_block_delta\",\"delta\":{\"type\":\"thinking_delta\",\"thinking\":\"\\u0042LOCKME\"}}\n\n"; + assert!( + !frame_guardrail_text(PassthroughProtocol::OpenaiChat, anthropic).contains("BLOCKME") + ); + } - fn cfg() -> ProxyConfig { - ProxyConfig { - addr: "127.0.0.1:0".into(), - request_body_limit_bytes: 1_048_576, - real_ip: Default::default(), - request_id: Default::default(), - url_rewrites: Vec::new(), - tls: None, - listeners: Vec::new(), - thread_per_core: None, - workers: None, + /// The `[DONE]` sentinel is not content, on either framing. A stream + /// that omits it entirely — OpenAI's Responses API sends none — is the + /// ordinary case, so nothing may depend on having seen one. + #[test] + fn the_done_sentinel_contributes_nothing_on_either_framing() { + for frame in [&b"data: [DONE]\n\n"[..], &b"data: [DONE]\r\n\r\n"[..]] { + assert_eq!( + frame_delta(PassthroughProtocol::Raw, frame), + (String::new(), None), + ); } } - const PK_ID: &str = "11111111-1111-1111-1111-111111111111"; - - /// A usage record carrying only the two canonical counters. - fn usage_dims(prompt: u32, completion: u32) -> PassthroughUsage { - PassthroughUsage { - prompt_tokens: prompt, - completion_tokens: completion, - ..Default::default() - } + /// Reasoning replayed by the caller is REQUEST text and is scanned; the + /// same field on a buffered RESPONSE is generated reasoning and is out + /// of the output-guardrail scope. One helper, two answers. + #[test] + fn replayed_reasoning_is_request_scan_text_and_not_response_scan_text() { + let msg = serde_json::json!({ + "role": "assistant", + "content": "visible", + "reasoning_content": "hidden reasoning payload", + }); + assert!(message_scan_text(&msg, true).contains("hidden reasoning payload")); + assert!(!message_scan_text(&msg, false).contains("hidden reasoning payload")); + assert!(message_scan_text(&msg, false).contains("visible")); } - fn provider_key_entry(api_base_unused: &str) -> ResourceEntry { - let json = format!( - r#"{{"display_name":"openai-up","secret":"sk-upstream","api_base":"{api_base_unused}","provider":"openai","adapter":"openai"}}"# + #[test] + fn opaque_stream_reads_flat_usage_only_from_a_labelled_frame() { + // An agent backend reached through a forward-proxy route has no + // recognisable envelope, and reports usage on its own event as a + // flat token object with no `usage` wrapper. + let labelled = b"event:token_usage\ndata:{\"name\":\"\",\"prompt_tokens\":14603,\"completion_tokens\":8,\"cache_read_input_tokens\":14272,\"reasoning_tokens\":8}\n\n"; + let usage = frame_delta(PassthroughProtocol::Raw, labelled) + .1 + .expect("a server-labelled usage frame must report usage"); + assert_eq!(usage.prompt_tokens, 14603); + assert_eq!(usage.completion_tokens, 8); + assert_eq!(usage.cache_read_tokens, 14272); + assert_eq!(usage.reasoning_tokens, 8); + + // The same flat shape on a frame the server did NOT name a usage + // report mints nothing: an opaque stream has no envelope to + // authenticate token-shaped fields against. + let unlabelled = b"event:history\ndata:{\"prompt_tokens\":99,\"completion_tokens\":9}\n\n"; + assert_eq!(frame_delta(PassthroughProtocol::Raw, unlabelled).1, None); + + // The flat allowance is Raw-only — a detected envelope keeps + // reading usage from its own shape. + assert_eq!( + frame_delta(PassthroughProtocol::OpenaiChat, labelled).1, + None ); - let pk: ProviderKey = serde_json::from_str(&json).unwrap(); - ResourceEntry::new(PK_ID, pk, 1) - } - fn apikey_entry(plaintext: &str, allowed_routes: Option<&[&str]>) -> ResourceEntry { - let routes = match allowed_routes { - Some(r) => format!( - r#", "allowed_routes": {}"#, - serde_json::to_string(r).unwrap() - ), - None => String::new(), - }; - let json = format!( - r#"{{"key_hash":"{}","allowed_models":["*"]{routes}}}"#, - ApiKey::hash_bearer(plaintext) + // An explicit `usage` OBJECT is still read from any opaque frame: + // it is self-describing, and this is the pre-existing behaviour. + let wrapped = + b"event:done\ndata:{\"usage\":{\"prompt_tokens\":4,\"completion_tokens\":2}}\n\n"; + assert_eq!( + frame_delta(PassthroughProtocol::Raw, wrapped).1, + Some(usage_dims(4, 2)) ); - let k: ApiKey = serde_json::from_str(&json).unwrap(); - ResourceEntry::new("k-1", k, 1) } - fn route_entry(id: &str, json: serde_json::Value) -> ResourceEntry { - let r: PassthroughRoute = serde_json::from_value(json).unwrap(); - ResourceEntry::new(id, r, 1) + #[test] + fn buffered_opaque_responses_are_never_probed_for_usage() { + // The no-phantom-tokens guarantee: a REST body that happens to carry + // a usage-shaped object is not a usage report. + let rpc = br#"{"jsonrpc":"2.0","id":1,"result":{"usage":{"prompt_tokens":99}}}"#; + assert_eq!(response_usage(PassthroughProtocol::Raw, None, rpc), None); + let top_level = br#"{"usage":{"prompt_tokens":99,"completion_tokens":9}}"#; + assert_eq!( + response_usage(PassthroughProtocol::Raw, None, top_level), + None + ); } - fn build_app(snap: AisixSnapshot) -> axum::Router { - let hub = Arc::new(Hub::new()); - let handle = SnapshotHandle::new(snap); - crate::build_router(crate::ProxyState::new(handle, hub, &cfg()).without_cache()) + #[test] + fn usage_merge_is_field_wise_max() { + let mut acc = PassthroughUsage { + prompt_tokens: 10, + completion_tokens: 4, + cache_read_tokens: 100, + ..Default::default() + }; + // A repeat of a cumulative usage object, and a partial one, are both + // harmless: no dimension ever regresses. + acc.merge(PassthroughUsage { + prompt_tokens: 10, + completion_tokens: 9, + reasoning_tokens: 3, + ..Default::default() + }); + assert_eq!( + acc, + PassthroughUsage { + prompt_tokens: 10, + completion_tokens: 9, + reasoning_tokens: 3, + cache_read_tokens: 100, + ..Default::default() + } + ); } - /// The `/passthrough/*` namespace carries no special case: with no - /// route claiming the path it is an ordinary router miss — a bare 404 - /// with an empty body, like any other unmatched path. - #[tokio::test] - async fn unclaimed_passthrough_path_takes_the_plain_404() { - let app = build_app(AisixSnapshot::new()); - let req = Request::builder() - .method("POST") - .uri("/passthrough/openai/v1/chat/completions") - .header("authorization", "Bearer whatever") - .body(axum::body::Body::empty()) - .unwrap(); - let resp = app.oneshot(req).await.unwrap(); - assert_eq!(resp.status(), StatusCode::NOT_FOUND); - let bytes = to_bytes(resp.into_body(), 65536).await.unwrap(); + #[test] + fn guardrail_text_covers_tool_calls_on_both_hooks() { + // A deny-listed string hidden in a tool call's arguments must be + // scanned — a benign `content` beside it would otherwise make the + // extraction non-empty and skip the raw-body fallback, letting the + // request pass a check the typed endpoint enforces. + let req = br#"{"model":"m","messages":[{"role":"assistant","content":"ok","tool_calls":[{"function":{"name":"run","arguments":"{\"cmd\":\"SECRET\"}"}}]}]}"#; + let text = request_guardrail_text(PassthroughProtocol::OpenaiChat, req); + assert!(text.contains("ok"), "content still scanned: {text}"); assert!( - bytes.is_empty(), - "the miss path carries no error envelope, got {:?}", - String::from_utf8_lossy(&bytes) + text.contains("SECRET"), + "tool-call arguments scanned: {text}" ); - } - #[tokio::test] - async fn unmatched_paths_keep_the_plain_404() { - let app = build_app(AisixSnapshot::new()); - let req = Request::builder() - .method("GET") - .uri("/definitely/not/a/route") - .body(axum::body::Body::empty()) - .unwrap(); - let resp = app.oneshot(req).await.unwrap(); - assert_eq!(resp.status(), StatusCode::NOT_FOUND); + let resp = br#"{"choices":[{"message":{"content":"sure","tool_calls":[{"function":{"name":"run","arguments":"{\"cmd\":\"SECRET\"}"}}]}}]}"#; + let text = response_guardrail_text(PassthroughProtocol::OpenaiChat, resp); + assert!(text.contains("sure")); + assert!(text.contains("SECRET"), "tool-call output scanned: {text}"); } - fn inject_route(target: &str) -> ResourceEntry { - route_entry( - "route-1", - serde_json::json!({ - "name": "openai-tunnel", - "path_prefix": "/passthrough/openai", - "target_url": target, - "provider_key_id": PK_ID - }), - ) + #[test] + fn buffered_output_guardrail_uses_the_response_envelope_not_the_request_hint() { + let cases = [ + ( + PassthroughProtocol::OpenaiChat, + br#"{"output":[{"type":"message","content":[{"type":"output_text","text":"BLOCKME"},{"type":"output_image","image_url":"MEDIA_SENTINEL"}]}]}"# + .as_slice(), + ), + ( + PassthroughProtocol::OpenaiCompletions, + br#"{"choices":[{"message":{"content":"BLOCKME","metadata":{"image":"MEDIA_SENTINEL"}}}]}"# + .as_slice(), + ), + ( + PassthroughProtocol::OpenaiResponses, + br#"{"choices":[{"text":"BLOCKME","metadata":{"image":"MEDIA_SENTINEL"}}]}"# + .as_slice(), + ), + ]; + for (request_protocol, body) in cases { + let text = try_buffered_response_guardrail_text(request_protocol, body) + .expect("a known response envelope must select its visible output"); + assert!(text.contains("BLOCKME"), "{request_protocol:?}: {text:?}"); + assert!( + !text.contains("MEDIA_SENTINEL"), + "{request_protocol:?} must not fall back to opaque media: {text:?}" + ); + } } - #[tokio::test] - async fn inject_route_replaces_caller_auth_with_provider_key() { - let upstream = MockServer::start().await; - Mock::given(wm_method("GET")) - .and(wm_path("/v1/models")) - .and(wiremock::matchers::header( - "authorization", - "Bearer sk-upstream", - )) - .respond_with( - ResponseTemplate::new(200) - .set_body_json(serde_json::json!({"object": "list", "data": []})), - ) - .mount(&upstream) - .await; - - let snap = AisixSnapshot::new(); - snap.provider_keys - .insert(provider_key_entry("http://unused")); - snap.apikeys.insert(apikey_entry("sk-caller", Some(&["*"]))); - snap.passthrough_routes - .insert(inject_route(&upstream.uri())); - let app = build_app(snap); - - let req = Request::builder() - .method("GET") - .uri("/passthrough/openai/v1/models") - .header("authorization", "Bearer sk-caller") - .body(axum::body::Body::empty()) - .unwrap(); - let resp = app.oneshot(req).await.unwrap(); - assert_eq!(resp.status(), StatusCode::OK); + #[test] + fn buffered_chat_response_uses_object_to_resolve_legacy_text() { + let text = try_buffered_response_guardrail_text( + PassthroughProtocol::OpenaiCompletions, + br#"{"object":"chat.completion","choices":[{"message":{"content":"BLOCKME"},"text":"legacy"}]}"#, + ) + .expect("a Chat response label resolves its structured message over legacy text"); + assert!(text.contains("BLOCKME"), "{text:?}"); + assert!(!text.contains("legacy"), "{text:?}"); - // The caller's own Authorization must not have reached upstream. - let received = &upstream.received_requests().await.unwrap()[0]; - let auth_values: Vec<_> = received.headers.get_all("authorization").iter().collect(); - assert_eq!(auth_values.len(), 1); + let error = try_buffered_response_guardrail_text( + PassthroughProtocol::OpenaiChat, + br#"{"choices":[{"message":{"content":"safe"},"text":"BLOCKME"}]}"#, + ) + .expect_err("mixed response carriers without a recognized label stay unevaluable"); + assert!(error.is_unevaluable(), "{error}"); } - #[tokio::test] - async fn key_without_route_grant_is_403() { - let upstream = MockServer::start().await; - let snap = AisixSnapshot::new(); - snap.provider_keys - .insert(provider_key_entry("http://unused")); - snap.apikeys.insert(apikey_entry("sk-caller", None)); - snap.passthrough_routes - .insert(inject_route(&upstream.uri())); - let app = build_app(snap); + #[test] + fn buffered_typed_output_without_one_response_envelope_is_unevaluable() { + for (request_protocol, body) in [ + ( + PassthroughProtocol::OpenaiChat, + br#"{"result":{"text":"BLOCKME"}}"#.as_slice(), + ), + ( + PassthroughProtocol::OpenaiCompletions, + br#"{"object":"chat.completion","result":{"text":"BLOCKME"}}"#.as_slice(), + ), + ( + PassthroughProtocol::OpenaiChat, + br#"{"output":[],"choices":[]}"#.as_slice(), + ), + ( + PassthroughProtocol::OpenaiChat, + br#"{"choices":[{"message":{"content":"safe"},"text":"BLOCKME"}]}"#.as_slice(), + ), + ] { + let error = try_buffered_response_guardrail_text(request_protocol, body).expect_err( + "a typed request must not choose an unknown or ambiguous response envelope", + ); + assert!(error.is_unevaluable(), "{error}"); + } - let req = Request::builder() - .method("GET") - .uri("/passthrough/openai/v1/models") - .header("authorization", "Bearer sk-caller") - .body(axum::body::Body::empty()) - .unwrap(); - let resp = app.oneshot(req).await.unwrap(); - assert_eq!(resp.status(), StatusCode::FORBIDDEN); - let bytes = to_bytes(resp.into_body(), 65536).await.unwrap(); - let v: serde_json::Value = serde_json::from_slice(&bytes).unwrap(); - assert_eq!(v["error"]["type"], "permission_denied"); + let raw = try_buffered_response_guardrail_text( + PassthroughProtocol::Raw, + br#"{"result":{"text":"BLOCKME"}}"#, + ) + .expect("Raw requests retain their broad decoded-string scan"); + assert!(raw.contains("BLOCKME"), "{raw:?}"); } - #[tokio::test] - async fn unauthenticated_route_request_is_401() { - let upstream = MockServer::start().await; - let snap = AisixSnapshot::new(); - snap.provider_keys - .insert(provider_key_entry("http://unused")); - snap.passthrough_routes - .insert(inject_route(&upstream.uri())); - let app = build_app(snap); - - let req = Request::builder() - .method("GET") - .uri("/passthrough/openai/v1/models") - .body(axum::body::Body::empty()) - .unwrap(); - let resp = app.oneshot(req).await.unwrap(); - assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + #[test] + fn buffered_response_envelope_ignores_opaque_sibling_size() { + let opaque = "x".repeat(MAX_RAW_SELECTOR_BYTES + 1); + let body = format!(r#"{{"choices":[{{"text":"BLOCKME","opaque":"{opaque}"}}]}}"#); + let text = + try_buffered_response_guardrail_text(PassthroughProtocol::OpenaiChat, body.as_bytes()) + .expect("a small visible completion stays scannable beside opaque data"); + assert_eq!(text, "BLOCKME"); } - /// The forward-proxy shadowing case: a host-matched request whose path - /// collides with a typed gateway route must be served by the - /// passthrough route, not the typed handler. - #[tokio::test] - async fn host_match_wins_over_typed_route_on_colliding_path() { - let upstream = MockServer::start().await; - Mock::given(wm_method("POST")) - .and(wm_path("/v1/chat/completions")) - .respond_with( - ResponseTemplate::new(200).set_body_json(serde_json::json!({"routed": "byo"})), - ) - .mount(&upstream) - .await; + #[test] + fn requested_model_comes_only_from_a_detected_envelope() { + let chat = br#"{"model":"gpt-4o","messages":[{"role":"user","content":"hi"}]}"#; + assert_eq!( + body_model_name(PassthroughProtocol::OpenaiChat, None, chat), + "gpt-4o" + ); + // An opaque body's `model`-shaped key belongs to some other API. + let opaque = br#"{"model":"whatever","config_name":"x"}"#; + assert_eq!(body_model_name(PassthroughProtocol::Raw, None, opaque), ""); + // Caller-supplied, so bounded and control-char free before it + // reaches telemetry. + let hostile = format!( + r#"{{"input":"x","model":"a\u0000b{}"}}"#, + "z".repeat(REQUESTED_MODEL_CAP * 2) + ); + let name = body_model_name( + PassthroughProtocol::OpenaiResponses, + None, + hostile.as_bytes(), + ); + assert_eq!(name.chars().count(), REQUESTED_MODEL_CAP); + assert!(!name.contains('\0')); + } - let snap = AisixSnapshot::new(); - snap.apikeys.insert(apikey_entry("sk-caller", Some(&["*"]))); - // forward_client + header_key: Authorization belongs to the caller - // and must reach upstream verbatim. - snap.passthrough_routes.insert(route_entry( - "route-h", - serde_json::json!({ - "name": "byo-host", - "hosts": ["ai.example.com"], - "target_url": upstream.uri(), - "auth_mode": "header_key", - "auth_header_name": "x-aisix-api-key", - "credential_mode": "forward_client" - }), - )); - let app = build_app(snap); + /// The passthrough route reads the SAME Responses shapes the typed + /// `/v1/responses` handler does, in the same directions. Request: + /// a replayed `reasoning` item's `content` AND `summary` are + /// caller-supplied text and are scanned. Response: a generated + /// `reasoning` item is out of the output scope and must not be — + /// the walk reads `content` off every item regardless of type, so + /// without an explicit skip a block rule matching only inside + /// reasoning refuses a response the typed route allows. + #[test] + fn responses_passthrough_scans_replayed_reasoning_but_not_generated_reasoning() { + let request = serde_json::json!({ + "input": [ + {"role": "user", "content": [{"type": "input_text", "text": "VISIBLE"}]}, + { + "type": "reasoning", + "summary": [{"type": "summary_text", "text": "SUMMARYSECRET"}], + "content": [{"type": "reasoning_text", "text": "REASONINGSECRET"}] + }, + {"type": "function_call_output", "call_id": "c1", "output": "TOOLRESULTSECRET"}, + {"type": "mcp_approval_response", "approve": true, "reason": "APPROVALSECRET"} + ] + }) + .to_string(); + let scanned = + request_guardrail_text(PassthroughProtocol::OpenaiResponses, request.as_bytes()); + assert!(scanned.contains("VISIBLE"), "got {scanned:?}"); + assert!(scanned.contains("REASONINGSECRET"), "got {scanned:?}"); + assert!(scanned.contains("SUMMARYSECRET"), "got {scanned:?}"); + // The tool-result and approval slots too. These matter precisely + // because the items beside them yield text: the raw-body fallback + // fires only on a WHOLLY empty extraction, so a mixed body would + // otherwise carry them past the scan while `/v1/responses` blocks + // the same envelope. + assert!(scanned.contains("TOOLRESULTSECRET"), "got {scanned:?}"); + assert!(scanned.contains("APPROVALSECRET"), "got {scanned:?}"); - let req = Request::builder() - .method("POST") - .uri("/v1/chat/completions") - .header("host", "ai.example.com") - .header("authorization", "Bearer employee-official-token") - .header("x-aisix-api-key", "sk-caller") - .header("content-type", "application/json") - .body(axum::body::Body::from(r#"{"model":"gpt-4o"}"#)) - .unwrap(); - let resp = app.oneshot(req).await.unwrap(); - assert_eq!(resp.status(), StatusCode::OK); - let bytes = to_bytes(resp.into_body(), 65536).await.unwrap(); - let v: serde_json::Value = serde_json::from_slice(&bytes).unwrap(); - assert_eq!(v["routed"], "byo", "typed chat handler must not serve this"); + let response = serde_json::json!({ + "output": [ + { + "type": "reasoning", + "summary": [{"type": "summary_text", "text": "SUMMARYSECRET"}], + "content": [{"type": "reasoning_text", "text": "REASONINGSECRET"}] + }, + { + "type": "message", + "content": [{"type": "output_text", "text": "the visible answer"}] + } + ] + }) + .to_string(); + let scanned = + response_guardrail_text(PassthroughProtocol::OpenaiResponses, response.as_bytes()); + assert!(scanned.contains("the visible answer"), "got {scanned:?}"); + assert!(!scanned.contains("REASONINGSECRET"), "got {scanned:?}"); + // Not a raw-body fallback: the message item yielded text, so a + // green above means the reasoning item was skipped rather than the + // whole walk having come back empty. + assert!(!scanned.contains("\"output\""), "got {scanned:?}"); + } - // BYO: the employee credential reached upstream verbatim; the - // gateway's side-channel header did not. - let received = &upstream.received_requests().await.unwrap()[0]; + #[test] + fn request_text_extraction_per_protocol() { + let chat = br#"{"model":"routing-model-only","messages":[{"role":"system","content":"s"},{"role":"user","content":[{"type":"text","text":"part"}]}],"forwarded_extra":"supplement"}"#; + let scanned = request_guardrail_text(PassthroughProtocol::OpenaiChat, chat); + for text in ["s", "part", "supplement"] { + assert!(scanned.contains(text), "{text} missing from {scanned:?}"); + } + // `model` selects the protocol target rather than supplying caller + // content. Arbitrary forwarded extras must still be scanned. + assert!( + !scanned.contains("routing-model-only"), + "protocol routing metadata leaked into guardrail text: {scanned:?}" + ); + let fim = br#"{"prompt":"def f(","suffix":"return"}"#; assert_eq!( - received.headers.get("authorization").unwrap(), - "Bearer employee-official-token" + request_guardrail_text(PassthroughProtocol::OpenaiCompletions, fim), + "def f(\nreturn" ); - assert!(received.headers.get("x-aisix-api-key").is_none()); + // Shape mismatch degrades to every decoded JSON string value. + let not_chat = br#"{"input":"x"}"#; + assert_eq!( + request_guardrail_text(PassthroughProtocol::OpenaiChat, not_chat), + "x" + ); + // Envelope metadata alone is not input text, but a forwarded + // supplementary field still is scanned. + let empty_chat = br#"{"messages":[{"role":"tool","tool_call_id":"1"}],"state":"fallback"}"#; + let scanned = request_guardrail_text(PassthroughProtocol::OpenaiChat, empty_chat); + assert!(scanned.contains("fallback"), "{scanned:?}"); } - #[tokio::test] - async fn disabled_route_does_not_match() { - let upstream = MockServer::start().await; - let snap = AisixSnapshot::new(); - snap.provider_keys - .insert(provider_key_entry("http://unused")); - snap.apikeys.insert(apikey_entry("sk-caller", Some(&["*"]))); - let mut json = serde_json::json!({ - "name": "openai-tunnel", - "path_prefix": "/passthrough/openai", - "target_url": upstream.uri(), - "provider_key_id": PK_ID, - "enabled": false - }); - json["enabled"] = serde_json::Value::Bool(false); - snap.passthrough_routes.insert(route_entry("route-1", json)); - let app = build_app(snap); + #[test] + fn request_source_scan_keeps_opaque_media_out_of_guardrail_text() { + let chat = br#"{"messages":[{"role":"user","content":[{"type":"image","source":{"data":"\u0042LOCKME"}},{"type":"document","source":{"data":"\u0042LOCKME"}},{"type":"text","text":"clean"}]}]}"#; + let scanned = request_guardrail_text(PassthroughProtocol::OpenaiChat, chat); + assert!(scanned.contains("clean"), "{scanned:?}"); + assert!(!scanned.contains("BLOCKME"), "{scanned:?}"); + + let responses = br#"{"input":[{"type":"message","content":[{"type":"input_image","image_url":"\u0042LOCKME"},{"type":"input_text","text":"clean"}]}]}"#; + let scanned = request_guardrail_text(PassthroughProtocol::OpenaiResponses, responses); + assert!(scanned.contains("clean"), "{scanned:?}"); + assert!(!scanned.contains("BLOCKME"), "{scanned:?}"); + + let completions = br#"{"prompt":[{"type":"image","data":"\u0042LOCKME"},"clean"]}"#; + let scanned = request_guardrail_text(PassthroughProtocol::OpenaiCompletions, completions); + assert!(scanned.contains("clean"), "{scanned:?}"); + assert!(!scanned.contains("BLOCKME"), "{scanned:?}"); + } - let req = Request::builder() - .method("GET") - .uri("/passthrough/openai/v1/models") - .header("authorization", "Bearer sk-caller") - .body(axum::body::Body::empty()) - .unwrap(); - let resp = app.oneshot(req).await.unwrap(); - // Disabled → no match → the ordinary router miss. - assert_eq!(resp.status(), StatusCode::NOT_FOUND); + #[test] + fn deep_responses_request_stays_typed_and_keeps_media_opaque() { + let request = deeply_nested_responses_request_with_opaque_media(); + assert_eq!( + detect_protocol(&request), + PassthroughProtocol::OpenaiResponses + ); + let scanned = request_guardrail_text(PassthroughProtocol::OpenaiResponses, &request); + assert!(scanned.contains("clean"), "{scanned:?}"); + assert!( + !scanned.contains("BLOCKME"), + "a deep valid envelope must not fall back to raw media: {scanned:?}" + ); } - #[tokio::test] - async fn anonymous_route_fails_closed_when_source_ip_is_unresolvable() { - let upstream = MockServer::start().await; - let snap = AisixSnapshot::new(); - snap.apikeys.insert(apikey_entry("sk-anon", Some(&["*"]))); - snap.passthrough_routes.insert(route_entry( - "route-a", - serde_json::json!({ - "name": "anon", - "path_prefix": "/anon", - "target_url": upstream.uri(), - "auth_mode": "anonymous", - "anonymous_key_id": "k-1", - "source_cidrs": ["0.0.0.0/0"], - "credential_mode": "forward_client" - }), - )); - let app = build_app(snap); + #[test] + fn duplicate_chat_carrier_is_unevaluable_without_leaking_media() { + let request = br#"{"messages":{"content":[{"type":"image","source":{"data":"\u0042LOCKME"}}]},"messages":[{"role":"user","content":"clean"}]}"#; + assert_eq!(detect_protocol(request), PassthroughProtocol::OpenaiChat); + let error = try_request_guardrail_text(PassthroughProtocol::OpenaiChat, request) + .expect_err("a duplicate malformed carrier must fail closed"); + assert!(error.is_unevaluable(), "{error}"); + assert!( + !request_guardrail_text(PassthroughProtocol::OpenaiChat, request).contains("BLOCKME"), + "a malformed duplicate must not trigger raw fallback" + ); + } - // In-process requests resolve no client socket; an unparseable - // source must never satisfy the CIDR gate. - let req = Request::builder() - .method("GET") - .uri("/anon/x") - .body(axum::body::Body::empty()) - .unwrap(); - let resp = app.oneshot(req).await.unwrap(); - assert_eq!(resp.status(), StatusCode::FORBIDDEN); + #[test] + fn malformed_typed_array_items_are_unevaluable_without_leaking_media() { + let cases = [ + ( + PassthroughProtocol::OpenaiChat, + br#"{"messages":["junk",{"role":"user","content":[{"type":"image","source":{"data":"\u0042LOCKME"}},{"type":"text","text":"clean"}]}]}"#.as_slice(), + ), + ( + PassthroughProtocol::OpenaiChat, + br#"{"messages":[{"role":"user","content":["junk",{"type":"image","source":{"data":"\u0042LOCKME"}},{"type":"text","text":"clean"}]}]}"#.as_slice(), + ), + ( + PassthroughProtocol::OpenaiChat, + br#"{"messages":[{"role":"user","content":[{"type":"text","text":{"image_url":"\u0042LOCKME"}},{"type":"text","text":"clean"}]}]}"#.as_slice(), + ), + ( + PassthroughProtocol::OpenaiResponses, + br#"{"input":["junk",{"type":"message","content":[{"type":"input_image","image_url":"\u0042LOCKME"},{"type":"input_text","text":"clean"}]}]}"#.as_slice(), + ), + ( + PassthroughProtocol::OpenaiResponses, + br#"{"input":[{"type":"message","content":[{"type":"input_text","text":{"image_url":"\u0042LOCKME"}},{"type":"input_text","text":"clean"}]}]}"#.as_slice(), + ), + ]; + for (protocol, request) in cases { + let error = try_request_guardrail_text(protocol, request) + .expect_err("a malformed typed array item must fail closed"); + assert!(error.is_unevaluable(), "{protocol:?}: {error}"); + assert!( + !request_guardrail_text(protocol, request).contains("BLOCKME"), + "a malformed array item must not trigger raw fallback for {protocol:?}" + ); + } } - // ---- pure helpers ---- + #[test] + fn guardrail_text_scans_decoded_forwarded_json_strings() { + let raw = br#"{"state":"\u0042LOCKME","nested":{"query":"\u4e2d\u6587"}}"#; + let scanned = request_guardrail_text(PassthroughProtocol::Raw, raw); + assert!(scanned.contains("BLOCKME"), "got {scanned:?}"); + assert!(scanned.contains("中文"), "got {scanned:?}"); + assert!(!scanned.contains(r#"\u0042"#), "got {scanned:?}"); + + let duplicate = br#"{"state":"\u0042LOCKME","state":"clean"}"#; + let scanned = request_guardrail_text(PassthroughProtocol::Raw, duplicate); + for text in ["BLOCKME", "clean"] { + assert!(scanned.contains(text), "{text} missing from {scanned:?}"); + } + + let chat = br#"{"messages":[{"role":"user","content":"clean"}],"state":{"query":"BLOCKME"},"state":"also-clean","documents":["\u4e2d\u6587"]}"#; + let scanned = request_guardrail_text(PassthroughProtocol::OpenaiChat, chat); + for text in ["clean", "BLOCKME", "also-clean", "中文"] { + assert!(scanned.contains(text), "{text} missing from {scanned:?}"); + } + + let response = br#"{"state":"\u0042LOCKME","state":"clean"}"#; + let scanned = response_guardrail_text(PassthroughProtocol::Raw, response); + for text in ["BLOCKME", "clean"] { + assert!(scanned.contains(text), "{text} missing from {scanned:?}"); + } + assert_eq!( + response_capture_text(PassthroughProtocol::Raw, response), + r#"{"state":"\u0042LOCKME","state":"clean"}"# + ); + + let deep = deeply_nested_escaped_block_json(); + assert!( + request_guardrail_text(PassthroughProtocol::Raw, &deep).contains("BLOCKME"), + "a valid deep Raw request must not fall back to escaped source" + ); + assert!( + response_guardrail_text(PassthroughProtocol::Raw, &deep).contains("BLOCKME"), + "a valid deep Raw response must not fall back to escaped source" + ); + } #[test] - fn path_prefix_matches_on_segment_boundary_only() { - assert!(path_under_prefix("/copilot", "/copilot")); - assert!(path_under_prefix("/copilot/chat", "/copilot")); - assert!(!path_under_prefix("/copilotx", "/copilot")); + fn known_request_envelopes_scan_duplicate_and_nested_source_strings() { + let cases = [ + ( + PassthroughProtocol::OpenaiChat, + br#"{"model":"routing-only","messages":[{"role":"user","content":"\u0069nputleakliteral"}],"messages":[{"role":"user","content":"clean","metadata":{"note":"NESTED"}}]}"# + .as_slice(), + ), + ( + PassthroughProtocol::OpenaiResponses, + br#"{"model":"routing-only","input":"\u0069nputleakliteral","input":"clean","metadata":{"note":"NESTED"}}"# + .as_slice(), + ), + ( + PassthroughProtocol::OpenaiCompletions, + br#"{"model":"routing-only","prompt":"\u0069nputleakliteral","prompt":"clean","metadata":{"note":"NESTED"}}"# + .as_slice(), + ), + ]; + for (protocol, body) in cases { + let scanned = request_guardrail_text(protocol, body); + for expected in ["inputleakliteral", "clean", "NESTED"] { + assert!(scanned.contains(expected), "{protocol:?}: {scanned:?}"); + } + assert!( + !scanned.contains("routing-only"), + "{protocol:?}: {scanned:?}" + ); + } } #[test] - fn inbound_host_strips_port_and_lowercases() { - let req = Request::builder() - .uri("/x") - .header("host", "API.Example.COM:8443") - .body(axum::body::Body::empty()) - .unwrap(); - assert_eq!(inbound_host(&req).as_deref(), Some("api.example.com")); + fn signed_redacted_thinking_is_not_request_guardrail_text() { + let redacted = br#"{"messages":[{"role":"assistant","content":[{"type":"redacted_thinking","data":"\u0042LOCKME","signature":"signed"},{"type":"text","text":"clean"}]}]}"#; + let scanned = request_guardrail_text(PassthroughProtocol::OpenaiChat, redacted); + assert!(scanned.contains("clean"), "{scanned:?}"); + assert!(!scanned.contains("BLOCKME"), "{scanned:?}"); + + let ambiguous = br#"{"messages":[{"role":"assistant","content":[{"type":"redacted_thinking","type":"text","data":"\u0042LOCKME"}]}]}"#; + assert!( + request_guardrail_text(PassthroughProtocol::OpenaiChat, ambiguous).contains("BLOCKME") + ); } #[test] - fn longest_prefix_and_host_specificity_win() { - let snap = AisixSnapshot::new(); - let mk = |id: &str, json: serde_json::Value| { - snap.passthrough_routes.insert(route_entry(id, json)) - }; - mk( - "r-short", - serde_json::json!({"name":"short","path_prefix":"/p","target_url":"http://a","provider_key_id":"pk"}), + fn known_response_envelopes_scan_duplicate_selected_source_strings() { + let chat = br#"{"model":"routing-only","choices":[{"message":{"content":"\u0042LOCKME","metadata":{"note":"NESTED"}}}],"choices":[{"message":{"content":"clean"}}]}"#; + let scanned = response_guardrail_text(PassthroughProtocol::OpenaiChat, chat); + for expected in ["BLOCKME", "clean"] { + assert!(scanned.contains(expected), "{scanned:?}"); + } + assert!(!scanned.contains("NESTED"), "{scanned:?}"); + assert!(!scanned.contains("routing-only"), "{scanned:?}"); + + let responses = br#"{"output":[{"type":"message","content":[{"type":"output_text","text":"\u0042LOCKME","metadata":{"note":"NESTED"}}]}],"output":[{"type":"message","content":[{"type":"output_text","text":"clean"}]}]}"#; + let scanned = response_guardrail_text(PassthroughProtocol::OpenaiResponses, responses); + for expected in ["BLOCKME", "clean"] { + assert!(scanned.contains(expected), "{scanned:?}"); + } + assert!(!scanned.contains("NESTED"), "{scanned:?}"); + + let refusal = br#"{"output":[{"type":"message","content":[{"type":"refusal","refusal":"\u0042LOCKREFUSAL","metadata":{"note":"NESTED"}}]}]}"#; + let scanned = response_guardrail_text(PassthroughProtocol::OpenaiResponses, refusal); + assert!(scanned.contains("BLOCKREFUSAL"), "{scanned:?}"); + assert!(!scanned.contains("NESTED"), "{scanned:?}"); + + let conflicting_type = br#"{"output":[{"type":"reasoning","type":"message","content":[{"text":"\u0042LOCKME"}]}]}"#; + assert!( + !response_guardrail_text(PassthroughProtocol::OpenaiResponses, conflicting_type) + .contains("BLOCKME") ); - mk( - "r-long", - serde_json::json!({"name":"long","path_prefix":"/p/deep","target_url":"http://b","provider_key_id":"pk"}), + + let same_hidden_type = br#"{"output":[{"type":"reasoning","type":"reasoning","summary":[{"text":"\u0042LOCKME"}]}]}"#; + assert!( + !response_guardrail_text(PassthroughProtocol::OpenaiResponses, same_hidden_type) + .contains("BLOCKME") ); - mk( - "r-host", - serde_json::json!({"name":"hosty","hosts":["h.example"],"target_url":"http://c","provider_key_id":"pk"}), + + let duplicate_text = br#"{"output":[{"type":"message","content":[{"type":"output_text","text":"\u0042LOCKME","text":"clean"}]}]}"#; + assert!( + response_guardrail_text(PassthroughProtocol::OpenaiResponses, duplicate_text) + .contains("BLOCKME") + ); + assert!( + response_guardrail_text(PassthroughProtocol::OpenaiResponses, duplicate_text) + .contains("clean") ); + } - let m = match_route(&snap, None, "/p/deep/x").unwrap(); - assert_eq!(m.entry.value.name, "long"); - assert_eq!(m.remainder, "/x"); - assert!(m.prefix_matched); + #[test] + fn completions_output_selects_only_choice_text_and_fails_closed_on_bad_carriers() { + let buffered = br#"{"choices":[{"text":"\u0042LOCKME","opaque":"OPAQUE"},{"text":"clean"}],"opaque":{"data":"OPAQUE"}}"#; + let scanned = response_guardrail_text(PassthroughProtocol::OpenaiCompletions, buffered); + assert!(scanned.contains("BLOCKME"), "{scanned:?}"); + assert!(scanned.contains("clean"), "{scanned:?}"); + assert!(!scanned.contains("OPAQUE"), "{scanned:?}"); + + let cases: [&[u8]; 6] = [ + br#"{"opaque":"BLOCKME"}"#, + br#"{"choices":{"text":"clean"},"opaque":"BLOCKME"}"#, + br#"{"choices":[{"text":"clean"}],"choices":[{"text":"other"}],"opaque":"BLOCKME"}"#, + br#"{"choices":["not-a-choice"],"opaque":"BLOCKME"}"#, + br#"{"choices":[{"text":{"opaque":"BLOCKME"}}]}"#, + br#"{"choices":[{"text":"clean","text":"other"}],"opaque":"BLOCKME"}"#, + ]; + for body in cases { + let error = try_response_guardrail_text(PassthroughProtocol::OpenaiCompletions, body) + .expect_err("a malformed completions carrier must fail closed"); + assert!(error.is_unevaluable(), "{error}"); + assert!( + !response_guardrail_text(PassthroughProtocol::OpenaiCompletions, body) + .contains("BLOCKME"), + "a malformed completions response must not fall back to opaque source text" + ); + } - // Host match beats any path-only match. - let m = match_route(&snap, Some("h.example"), "/p/deep/x").unwrap(); - assert_eq!(m.entry.value.name, "hosty"); - assert_eq!(m.remainder, "/p/deep/x"); - assert!(!m.prefix_matched); + let oversized = format!( + r#"{{"choices":[{{"text":"{}"}}],"opaque":"BLOCKME"}}"#, + "x".repeat(MAX_RAW_SELECTOR_BYTES + 1), + ); + assert!(try_response_guardrail_text( + PassthroughProtocol::OpenaiCompletions, + oversized.as_bytes() + ) + .expect_err("an oversized selected completions value must fail closed") + .is_unevaluable()); + } - // A preserve_host route narrowed by a prefix relays the WHOLE path: - // the prefix is a match condition on an upstream that owns its own - // path space, not a gateway mount point. GitHub Copilot's CLI needs - // this — its MCP server answers on /mcp/readonly of the same host it - // serves chat from, and a stripped "/readonly" 404s. - mk( - "r-mirror", - serde_json::json!({ - "name":"mirror","hosts":["m.example"],"path_prefix":"/mcp", - "preserve_host":true,"credential_mode":"forward_client" - }), + #[test] + fn streamed_completions_select_only_choice_text_and_reject_bad_carriers() { + let frame = format!( + "data: {{\"choices\":[{{\"index\":0,\"text\":\"clean\",\"opaque\":\"{}\"}}],\"opaque\":\"{}\"}}\n\n", + "BLOCKME".repeat(MAX_RAW_SELECTOR_BYTES / "BLOCKME".len() + 1), + "BLOCKME", ); - let m = match_route(&snap, Some("m.example"), "/mcp/readonly").unwrap(); - assert_eq!(m.entry.value.name, "mirror"); - assert_eq!(m.remainder, "/mcp/readonly"); + let typed = frame_parts(PassthroughProtocol::OpenaiCompletions, frame.as_bytes()) + .0 + .scan; + let text = stream_guardrail_text( + PassthroughProtocol::OpenaiCompletions, + frame.as_bytes(), + typed, + ); + assert!(!text.unevaluable); + let scanned = stream_guardrail_scan_text(&[], &text.continuations, &text.supplemental); + assert!(scan_candidates_contain(&scanned, "clean"), "{scanned:?}"); assert!( - !m.prefix_matched, - "a mirrored path is never version-deduped" + !scan_candidates_contain(&scanned, "BLOCKME"), + "opaque completions extensions must not become supplemental guardrail text: {scanned:?}" + ); + assert!( + !frame_guardrail_text(PassthroughProtocol::OpenaiCompletions, frame.as_bytes()) + .contains("BLOCKME"), + "a completions frame must not fall back to opaque source text" + ); + + for payload in [ + br#"{"choices":{"index":0,"text":"clean"}}"#.as_slice(), + br#"{"choices":[{"index":0,"text":"clean"}],"choices":[{"index":1,"text":"other"}]}"# + .as_slice(), + br#"{"choices":["not-a-choice"]}"#.as_slice(), + br#"{"choices":[{"index":0,"text":"clean","text":"other"}]}"#.as_slice(), + ] { + assert!(matches!( + stream_source_continuations(PassthroughProtocol::OpenaiCompletions, payload), + SourceContinuations::Unevaluable + )); + } + + let oversized = format!( + r#"{{"choices":[{{"index":0,"text":"{}"}}]}}"#, + "x".repeat(MAX_RAW_SELECTOR_BYTES + 1), ); + assert!(matches!( + stream_source_continuations( + PassthroughProtocol::OpenaiCompletions, + oversized.as_bytes() + ), + SourceContinuations::Unevaluable + )); } #[test] - fn sse_splitter_emits_complete_frames_and_keeps_partials() { - let mut s = SseFrameSplitter::new(); - let frames = s.push(b"data: a\n\ndata: b\n\ndata: par"); - assert_eq!(frames.len(), 2); - assert_eq!(frames[0], b"data: a\n\n"); - let frames = s.push(b"tial\n\n"); - assert_eq!(frames.len(), 1); - assert_eq!(frames[0], b"data: partial\n\n"); - assert!(s.take_rest().is_empty()); - // CRLF boundaries too. - let mut s = SseFrameSplitter::new(); - let frames = s.push(b"data: x\r\n\r\nrest"); - assert_eq!(frames.len(), 1); - assert_eq!(s.take_rest(), b"rest"); + fn malformed_or_capped_stream_supplemental_is_marked_separately() { + let supplemental_failure = |name_fields: String| { + format!( + "data: {{\"type\":\"content_block_start\",\"index\":0,\"content_block\":{{\"type\":\"tool_use\",\"input\":\"clean\",{name_fields}}}}}\n\n" + ) + }; + let malformed = supplemental_failure("\"name\":1".to_owned()); + let capped = supplemental_failure(format!( + "{}\"name\":\"last\"", + "\"name\":\"n\",".repeat(MAX_RAW_SELECTOR_ITEMS), + )); + for frame in [&malformed, &capped] { + let typed = frame_parts(PassthroughProtocol::OpenaiChat, frame.as_bytes()) + .0 + .scan; + let text = + stream_guardrail_text(PassthroughProtocol::OpenaiChat, frame.as_bytes(), typed); + assert!(text.unevaluable, "{frame}"); + assert!( + text.supplemental_unevaluable, + "the selected supplemental failure must remain distinguishable: {frame}" + ); + } } #[test] - fn sse_splitter_and_usage_label_read_cr_framing() { - let frame = b"event: token_usage\rdata: {\"input_tokens\":3,\"output_tokens\":4}\r\r"; - let mut s = SseFrameSplitter::new(); - assert_eq!( - s.push(&[&frame[..], b"data: next"].concat()), - vec![frame.to_vec()] - ); - assert_eq!(s.take_rest(), b"data: next"); - assert!(is_usage_labelled_frame(frame)); - let (_, usage) = frame_delta(PassthroughProtocol::Raw, frame); - assert!(usage.is_some(), "a CR-framed usage report is read"); + fn chat_output_guardrail_keeps_media_and_unknown_parts_opaque() { + let buffered = br#"{ + "choices": [{ + "message": { + "content": [ + {"type":"image_url","image_url":{"url":"BUFFERED_IMAGE_SENTINEL"}}, + {"type":"input_audio","input_audio":{"data":"BUFFERED_AUDIO_SENTINEL"}}, + {"type":"file","file":{"file_data":"BUFFERED_FILE_SENTINEL"}}, + {"type":"future_media","text":"BUFFERED_OPAQUE_PART_SENTINEL"}, + {"type":"text","text":"BUFFERED_VISIBLE_SENTINEL"} + ], + "reasoning_content": "BUFFERED_REASONING_SENTINEL", + "metadata": {"note":"BUFFERED_METADATA_SENTINEL"}, + "tool_calls": [ + {"type":"function","function":{"name":"lookup","arguments":"BUFFERED_TOOL_ARGUMENT_SENTINEL"}}, + {"type":"custom","custom":{"name":"custom","input":"BUFFERED_CUSTOM_INPUT_SENTINEL"}} + ], + "function_call": {"name":"legacy","arguments":"BUFFERED_LEGACY_ARGUMENT_SENTINEL"} + } + }] + }"#; + let scanned = response_guardrail_text(PassthroughProtocol::OpenaiChat, buffered); + for expected in [ + "BUFFERED_VISIBLE_SENTINEL", + "lookup", + "BUFFERED_TOOL_ARGUMENT_SENTINEL", + "custom", + "BUFFERED_CUSTOM_INPUT_SENTINEL", + "legacy", + "BUFFERED_LEGACY_ARGUMENT_SENTINEL", + ] { + assert!(scanned.contains(expected), "{scanned:?}"); + } + for opaque in [ + "BUFFERED_IMAGE_SENTINEL", + "BUFFERED_AUDIO_SENTINEL", + "BUFFERED_FILE_SENTINEL", + "BUFFERED_OPAQUE_PART_SENTINEL", + "BUFFERED_REASONING_SENTINEL", + "BUFFERED_METADATA_SENTINEL", + ] { + assert!(!scanned.contains(opaque), "{scanned:?}"); + } + + let frame = b"data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":[{\"index\":0,\"type\":\"image_url\",\"image_url\":{\"url\":\"STREAM_IMAGE_SENTINEL\"}},{\"index\":1,\"type\":\"input_audio\",\"input_audio\":{\"data\":\"STREAM_AUDIO_SENTINEL\"}},{\"index\":2,\"type\":\"file\",\"file\":{\"file_data\":\"STREAM_FILE_SENTINEL\"}},{\"index\":3,\"type\":\"future_media\",\"text\":\"STREAM_OPAQUE_PART_SENTINEL\"},{\"index\":4,\"type\":\"text\",\"text\":\"STREAM_VISIBLE_SENTINEL\"}],\"reasoning_content\":\"STREAM_REASONING_SENTINEL\",\"metadata\":{\"note\":\"STREAM_METADATA_SENTINEL\"},\"tool_calls\":[{\"index\":0,\"type\":\"function\",\"function\":{\"name\":\"lookup\",\"arguments\":\"STREAM_TOOL_ARGUMENT_SENTINEL\"}},{\"index\":1,\"type\":\"custom\",\"custom\":{\"name\":\"custom\",\"input\":\"STREAM_CUSTOM_INPUT_SENTINEL\"}}],\"function_call\":{\"name\":\"legacy\",\"arguments\":\"STREAM_LEGACY_ARGUMENT_SENTINEL\"}}}]}\n\n"; + let typed = frame_parts(PassthroughProtocol::OpenaiChat, frame).0.scan; + assert!(typed.contains("STREAM_OPAQUE_PART_SENTINEL"), "{typed:?}"); + let text = stream_guardrail_text(PassthroughProtocol::OpenaiChat, frame, typed); + assert!(!text.unevaluable); + let scanned = stream_guardrail_scan_text(&[], &text.continuations, &text.supplemental); + for expected in [ + "STREAM_VISIBLE_SENTINEL", + "lookup", + "STREAM_TOOL_ARGUMENT_SENTINEL", + "custom", + "STREAM_CUSTOM_INPUT_SENTINEL", + "legacy", + "STREAM_LEGACY_ARGUMENT_SENTINEL", + ] { + assert!(scan_candidates_contain(&scanned, expected), "{scanned:?}"); + } + for opaque in [ + "STREAM_IMAGE_SENTINEL", + "STREAM_AUDIO_SENTINEL", + "STREAM_FILE_SENTINEL", + "STREAM_OPAQUE_PART_SENTINEL", + "STREAM_REASONING_SENTINEL", + "STREAM_METADATA_SENTINEL", + ] { + assert!(!scan_candidates_contain(&scanned, opaque), "{scanned:?}"); + } } #[test] - fn frame_in_band_error_reads_the_protocol_s_own_failure_events() { - let anthropic = b"event: error\ndata: {\"type\":\"error\",\"error\":{\"type\":\"overloaded_error\",\"message\":\"busy\"}}\n\n"; - let err = frame_in_band_error(PassthroughProtocol::OpenaiChat, anthropic).unwrap(); - // Anthropic documents 529 for overloaded; not a 4xx, so it maps to 502. - assert_eq!(err.http_status(), 502); - let openai = - br#"data: {"error":{"message":"slow down","type":"rate_limit_error","code":429}} - -"#; - let err = frame_in_band_error(PassthroughProtocol::OpenaiCompletions, openai).unwrap(); - assert_eq!(err.http_status(), 429); - let responses = br#"data: {"type":"response.failed","response":{"error":{"code":"server_error","message":"x"}}} - -"#; - assert!(frame_in_band_error(PassthroughProtocol::OpenaiResponses, responses).is_some()); - // An opaque stream is never read for one, and ordinary frames are not one. - assert!(frame_in_band_error(PassthroughProtocol::Raw, openai).is_none()); - let delta = br#"data: {"choices":[{"delta":{"content":"hel"}}]} + fn chat_output_guardrail_scans_client_visible_refusals() { + let buffered = br#"{ + "choices": [{ + "message": { + "content": [ + {"type":"image_url","image_url":{"url":"BUFFERED_MEDIA_SENTINEL"}}, + {"type":"refusal","refusal":"CONTENT_REFUSAL_BLOCKME"} + ], + "refusal":"MESSAGE_REFUSAL_BLOCKME", + "reasoning_content":"BUFFERED_REASONING_SENTINEL" + } + }] + }"#; + let scanned = response_guardrail_text(PassthroughProtocol::OpenaiChat, buffered); + for refusal in ["CONTENT_REFUSAL_BLOCKME", "MESSAGE_REFUSAL_BLOCKME"] { + assert!(scanned.contains(refusal), "{scanned:?}"); + } + for opaque in ["BUFFERED_MEDIA_SENTINEL", "BUFFERED_REASONING_SENTINEL"] { + assert!(!scanned.contains(opaque), "{scanned:?}"); + } -"#; - assert!(frame_in_band_error(PassthroughProtocol::OpenaiChat, delta).is_none()); + let first_frame = + b"data: {\"choices\":[{\"index\":0,\"delta\":{\"refusal\":\"BLOC\"}}]}\n\n"; + let second_frame = + b"data: {\"choices\":[{\"index\":0,\"delta\":{\"refusal\":\"KME\"}}]}\n\n"; + let first = stream_guardrail_text( + PassthroughProtocol::OpenaiChat, + first_frame, + frame_parts(PassthroughProtocol::OpenaiChat, first_frame) + .0 + .scan, + ); + let second = stream_guardrail_text( + PassthroughProtocol::OpenaiChat, + second_frame, + frame_parts(PassthroughProtocol::OpenaiChat, second_frame) + .0 + .scan, + ); + assert!(!first.unevaluable); + assert!(!second.unevaluable); + + let mut continuations = Vec::new(); + let mut tails = Vec::new(); + let mut supplemental = Vec::new(); + let mut closed_prefixes = Vec::new(); + append_stream_guardrail_text( + &mut continuations, + &mut tails, + &mut supplemental, + &mut closed_prefixes, + &first, + ); + append_stream_guardrail_text( + &mut continuations, + &mut tails, + &mut supplemental, + &mut closed_prefixes, + &second, + ); + let scanned = stream_guardrail_scan_text(&tails, &continuations, &supplemental); + assert!(scan_candidates_contain(&scanned, "BLOCKME"), "{scanned:?}"); } #[test] - fn frame_delta_extracts_chat_content_and_usage() { - let frame = br#"data: {"choices":[{"delta":{"content":"hel"}}]} + fn responses_output_guardrail_keeps_generated_media_opaque() { + let buffered = br#"{"output":[{"type":"image_generation_call","result":"BUFFERED_MEDIA_SENTINEL"},{"type":"message","content":[{"type":"output_text","text":"VISIBLE_TEXT_SENTINEL"}]},{"type":"function_call","name":"lookup","arguments":"{\"query\":\"TOOL_ARGUMENT_SENTINEL\"}"},{"type":"mcp_call","name":"mcp","arguments":"MCP_ARGUMENT_SENTINEL"},{"type":"custom_tool_call","name":"custom","input":"CUSTOM_INPUT_SENTINEL"}]}"#; + let scanned = response_guardrail_text(PassthroughProtocol::OpenaiResponses, buffered); + for expected in [ + "VISIBLE_TEXT_SENTINEL", + "TOOL_ARGUMENT_SENTINEL", + "MCP_ARGUMENT_SENTINEL", + "CUSTOM_INPUT_SENTINEL", + ] { + assert!(scanned.contains(expected), "{scanned:?}"); + } + assert!(!scanned.contains("BUFFERED_MEDIA_SENTINEL"), "{scanned:?}"); + + // A conflicting item discriminator is opaque as a whole. Even + // safe-named fields could be an image/audio/file payload attached to + // the other discriminator. + let ambiguous = br#"{"output":[{"type":"image_generation_call","type":"message","result":"AMBIGUOUS_MEDIA_SENTINEL","content":[{"type":"output_text","text":"AMBIGUOUS_VISIBLE_SENTINEL"}],"arguments":"AMBIGUOUS_TOOL_SENTINEL"}]}"#; + let scanned = response_guardrail_text(PassthroughProtocol::OpenaiResponses, ambiguous); + for opaque in [ + "AMBIGUOUS_MEDIA_SENTINEL", + "AMBIGUOUS_VISIBLE_SENTINEL", + "AMBIGUOUS_TOOL_SENTINEL", + ] { + assert!(!scanned.contains(opaque), "{scanned:?}"); + } -"#; - let (text, usage) = frame_delta(PassthroughProtocol::OpenaiChat, frame); - assert_eq!(text, "hel"); - assert!(usage.is_none()); + let unknown = + br#"{"output":[{"text":"UNKNOWN_TEXT_SENTINEL","input":"UNKNOWN_MEDIA_SENTINEL"}]}"#; + let scanned = response_guardrail_text(PassthroughProtocol::OpenaiResponses, unknown); + assert!(!scanned.contains("UNKNOWN_TEXT_SENTINEL"), "{scanned:?}"); + assert!(!scanned.contains("UNKNOWN_MEDIA_SENTINEL"), "{scanned:?}"); - let done = br#"data: {"choices":[],"usage":{"prompt_tokens":7,"completion_tokens":3}} + let media_frame = b"data: {\"type\":\"response.image_generation_call.partial_image\",\"partial_image_b64\":\"STREAM_MEDIA_SENTINEL\"}\n\n"; + assert!( + !frame_guardrail_text(PassthroughProtocol::OpenaiResponses, media_frame) + .contains("STREAM_MEDIA_SENTINEL") + ); + let typed = frame_parts(PassthroughProtocol::OpenaiResponses, media_frame) + .0 + .scan; + let media = stream_guardrail_text(PassthroughProtocol::OpenaiResponses, media_frame, typed); + let scanned = stream_guardrail_scan_text(&[], &media.continuations, &media.supplemental); + assert!( + !scan_candidates_contain(&scanned, "STREAM_MEDIA_SENTINEL"), + "{scanned:?}" + ); -"#; - let (text, usage) = frame_delta(PassthroughProtocol::OpenaiChat, done); - assert_eq!(text, ""); - assert_eq!(usage, Some(usage_dims(7, 3))); + // A terminal response can be the only authoritative output carrier. + // Its message text is scanned, while opaque image-generation data is + // still excluded. + let terminal = b"data: {\"type\":\"response.completed\",\"response\":{\"output\":[{\"type\":\"image_generation_call\",\"result\":\"TERMINAL_MEDIA_SENTINEL\"},{\"type\":\"message\",\"content\":[{\"type\":\"output_text\",\"text\":\"TERMINAL_VISIBLE_SENTINEL\"}]}]}}\n\n"; + let scanned = frame_guardrail_text(PassthroughProtocol::OpenaiResponses, terminal); + assert!(!scanned.contains("TERMINAL_MEDIA_SENTINEL"), "{scanned:?}"); + assert!(scanned.contains("TERMINAL_VISIBLE_SENTINEL"), "{scanned:?}"); + + // `delta` is not a universally textual field: on a conflicting text + // and audio discriminator it is opaque, rather than a path for audio + // base64 to reach a text guardrail. + let conflicting_delta = b"data: {\"type\":\"response.output_audio.delta\",\"type\":\"response.output_text.delta\",\"delta\":\"CONFLICTING_MEDIA_SENTINEL\"}\n\n"; + let typed = frame_parts(PassthroughProtocol::OpenaiResponses, conflicting_delta) + .0 + .scan; + assert!(typed.contains("CONFLICTING_MEDIA_SENTINEL"), "{typed:?}"); + let text = stream_guardrail_text( + PassthroughProtocol::OpenaiResponses, + conflicting_delta, + typed, + ); + let scanned = stream_guardrail_scan_text(&[], &text.continuations, &text.supplemental); + assert!( + !scan_candidates_contain(&scanned, "CONFLICTING_MEDIA_SENTINEL"), + "{scanned:?}" + ); - let fim = br#"data: {"choices":[{"text":"def "}]} + let ambiguous_event = b"data: {\"type\":\"response.output_text.done\",\"type\":\"response.output_audio.done\",\"text\":\"AMBIGUOUS_EVENT_TEXT_SENTINEL\",\"input\":\"AMBIGUOUS_EVENT_MEDIA_SENTINEL\"}\n\n"; + let scanned = frame_guardrail_text(PassthroughProtocol::OpenaiResponses, ambiguous_event); + assert!( + !scanned.contains("AMBIGUOUS_EVENT_TEXT_SENTINEL"), + "{scanned:?}" + ); + assert!( + !scanned.contains("AMBIGUOUS_EVENT_MEDIA_SENTINEL"), + "{scanned:?}" + ); -"#; - let (text, _) = frame_delta(PassthroughProtocol::OpenaiCompletions, fim); - assert_eq!(text, "def "); + let text_frame = b"data: {\"type\":\"response.output_text.delta\",\"delta\":\"VISIBLE_STREAM_SENTINEL\"}\n\n"; + assert!( + frame_guardrail_text(PassthroughProtocol::OpenaiResponses, text_frame) + .contains("VISIBLE_STREAM_SENTINEL") + ); } - /// The recorded total is the upstream's own: adopted from the first - /// report, kept while every later report carries one, and dropped the - /// moment one does not — a surviving partial value would be a number - /// no upstream stated. #[test] - fn merged_usage_keeps_a_total_only_while_every_report_carries_one() { - let with = usage_of( - &serde_json::json!({"prompt_tokens": 7, "completion_tokens": 3, "total_tokens": 10}), - ) - .unwrap(); - let without = usage_of(&serde_json::json!({"output_tokens": 5})).unwrap(); - assert_eq!(with.upstream_total_tokens, Some(10)); - assert_eq!(without.upstream_total_tokens, None); + fn generated_reasoning_stays_out_of_the_source_preserving_output_scan() { + let buffered = br#"{"output":[{"type":"reasoning","summary":[{"type":"summary_text","text":"\u0042LOCKME"}]}]}"#; + assert!( + !response_guardrail_text(PassthroughProtocol::OpenaiResponses, buffered) + .contains("BLOCKME") + ); - let mut acc = None; - merge_usage(&mut acc, with); - assert_eq!(acc.unwrap().upstream_total_tokens, Some(10)); - merge_usage(&mut acc, with); - assert_eq!(acc.unwrap().upstream_total_tokens, Some(10)); - merge_usage(&mut acc, without); - assert_eq!(acc.unwrap().upstream_total_tokens, None); + for frame in [ + b"data: {\"type\":\"response.reasoning_text.delta\",\"delta\":\"\\u0042LOCKME\"}\n\n".as_slice(), + b"data: {\"type\":\"response.reasoning_text.done\",\"text\":\"\\u0042LOCKME\"}\n\n".as_slice(), + b"data: {\"type\":\"response.reasoning_summary_part.done\",\"part\":{\"type\":\"summary_text\",\"text\":\"\\u0042LOCKME\"}}\n\n".as_slice(), + b"data: {\"type\":\"response.content_part.done\",\"part\":{\"type\":\"reasoning_text\",\"text\":\"\\u0042LOCKME\"}}\n\n".as_slice(), + b"data: {\"type\":\"response.output_item.done\",\"item\":{\"type\":\"reasoning\",\"summary\":[{\"text\":\"\\u0042LOCKME\"}]}}\n\n".as_slice(), + ] { + assert!( + !frame_guardrail_text(PassthroughProtocol::OpenaiResponses, frame) + .contains("BLOCKME") + ); + } } #[test] - fn usage_of_reads_every_dimension_in_every_spelling() { - // OpenAI chat: the cache hit is nested under `prompt_tokens_details` - // and the reasoning count under `completion_tokens_details`. - let openai = serde_json::json!({ - "prompt_tokens": 100, - "completion_tokens": 20, - "prompt_tokens_details": {"cached_tokens": 80}, - "completion_tokens_details": {"reasoning_tokens": 12}, - }); - assert_eq!( - usage_of(&openai), - Some(PassthroughUsage { - prompt_tokens: 100, - completion_tokens: 20, - cached_prompt_tokens: 80, - reasoning_tokens: 12, - ..Default::default() - }) - ); + fn known_sse_envelopes_scan_selected_source_strings() { + let chat = b"data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":\"\\u0042LOCKME\",\"metadata\":{\"note\":\"NESTED\"}}},{\"index\":1,\"delta\":{\"content\":\"clean\"}}]}\n\n"; + let scanned = frame_guardrail_text(PassthroughProtocol::OpenaiChat, chat); + for expected in ["BLOCKME", "clean"] { + assert!(scanned.contains(expected), "{scanned:?}"); + } + assert!(!scanned.contains("NESTED"), "{scanned:?}"); - // Responses API: the `input`/`output` spelling, details nested under - // the matching names. - let responses = serde_json::json!({ - "input_tokens": 30, - "output_tokens": 9, - "input_tokens_details": {"cached_tokens": 25}, - "output_tokens_details": {"reasoning_tokens": 4}, - }); - assert_eq!( - usage_of(&responses), - Some(PassthroughUsage { - prompt_tokens: 30, - completion_tokens: 9, - cached_prompt_tokens: 25, - reasoning_tokens: 4, - ..Default::default() - }) - ); + let responses = b"data: {\"type\":\"response.output_text.delta\",\"delta\":\"\\u0042LOCKME\",\"delta\":\"clean\",\"metadata\":{\"note\":\"NESTED\"}}\n\n"; + let scanned = frame_guardrail_text(PassthroughProtocol::OpenaiResponses, responses); + for expected in ["BLOCKME", "clean"] { + assert!(scanned.contains(expected), "{scanned:?}"); + } + assert!(!scanned.contains("NESTED"), "{scanned:?}"); - // Anthropic: cache counters are separate, additive fields. - let anthropic = serde_json::json!({ - "input_tokens": 11, - "output_tokens": 5, - "cache_creation_input_tokens": 300, - "cache_read_input_tokens": 1200, - }); - assert_eq!( - usage_of(&anthropic), - Some(PassthroughUsage { - prompt_tokens: 11, - completion_tokens: 5, - cache_creation_tokens: 300, - cache_read_tokens: 1200, - ..Default::default() - }) - ); + let refusal = b"data: {\"type\":\"response.refusal.delta\",\"delta\":\"\\u0042LOCKREFUSAL\",\"metadata\":{\"note\":\"NESTED\"}}\n\n"; + let scanned = frame_guardrail_text(PassthroughProtocol::OpenaiResponses, refusal); + assert!(scanned.contains("BLOCKREFUSAL"), "{scanned:?}"); + assert!(!scanned.contains("NESTED"), "{scanned:?}"); - // DeepSeek reports the cache hit flat, and a ZEROED nested detail - // must not mask it (same precedence the typed OpenAI bridge uses). - let deepseek = serde_json::json!({ - "prompt_tokens": 40, - "completion_tokens": 6, - "prompt_tokens_details": {"cached_tokens": 0}, - "prompt_cache_hit_tokens": 32, - }); - assert_eq!(usage_of(&deepseek).unwrap().cached_prompt_tokens, 32); + let conflicting = b"data: {\"type\":\"response.content_part.done\",\"part\":{\"type\":\"reasoning_text\",\"type\":\"output_text\",\"text\":\"\\u0042LOCKME\"}}\n\n"; + assert!( + !frame_guardrail_text(PassthroughProtocol::OpenaiResponses, conflicting) + .contains("BLOCKME") + ); - // The flat agent-backend shape, all five dimensions at the root. - let flat = serde_json::json!({ - "prompt_tokens": 14603, - "completion_tokens": 8, - "cache_creation_input_tokens": 0, - "cache_read_input_tokens": 14272, - "reasoning_tokens": 8, - }); - assert_eq!( - usage_of(&flat), - Some(PassthroughUsage { - prompt_tokens: 14603, - completion_tokens: 8, - cache_read_tokens: 14272, - reasoning_tokens: 8, - ..Default::default() - }) + let conflicting_item = b"data: {\"type\":\"response.output_item.done\",\"item\":{\"type\":\"reasoning\",\"type\":\"message\",\"content\":[{\"text\":\"\\u0042LOCKME\"}]}}\n\n"; + assert!( + !frame_guardrail_text(PassthroughProtocol::OpenaiResponses, conflicting_item) + .contains("BLOCKME") + ); + let hidden_item = b"data: {\"type\":\"response.output_item.done\",\"item\":{\"type\":\"reasoning\",\"type\":\"reasoning\",\"summary\":[{\"text\":\"\\u0042LOCKME\"}]}}\n\n"; + assert!( + !frame_guardrail_text(PassthroughProtocol::OpenaiResponses, hidden_item) + .contains("BLOCKME") ); - // An object with no recognised counter mints nothing. - assert_eq!(usage_of(&serde_json::json!({"disk": "80%"})), None); - assert_eq!(usage_of(&serde_json::Value::Null), None); + let conflicting_anthropic = b"data: {\"type\":\"content_block_start\",\"content_block\":{\"type\":\"thinking\",\"type\":\"text\",\"text\":\"\\u0042LOCKME\"}}\n\n"; + assert!( + !frame_guardrail_text(PassthroughProtocol::OpenaiChat, conflicting_anthropic) + .contains("BLOCKME") + ); + let hidden_anthropic = b"data: {\"type\":\"content_block_start\",\"content_block\":{\"type\":\"thinking\",\"type\":\"thinking\",\"thinking\":\"\\u0042LOCKME\"}}\n\n"; + assert!( + !frame_guardrail_text(PassthroughProtocol::OpenaiChat, hidden_anthropic) + .contains("BLOCKME") + ); } #[test] - fn anthropic_stream_reports_the_prompt_side_from_message_start() { - // Anthropic splits usage across two frames: `message_start` carries - // the input + cache counters, the terminal `message_delta` only the - // output ones. Reading the top level alone loses the prompt side. - let start = br#"data: {"type":"message_start","message":{"id":"msg_1","usage":{"input_tokens":12,"cache_creation_input_tokens":300,"cache_read_input_tokens":1200}}} - -"#; - let (_, start_usage) = frame_delta(PassthroughProtocol::OpenaiChat, start); - let start_usage = start_usage.expect("message_start must report usage"); - assert_eq!(start_usage.prompt_tokens, 12); - assert_eq!(start_usage.cache_creation_tokens, 300); - assert_eq!(start_usage.cache_read_tokens, 1200); - - let delta = br#"data: {"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":7}} - -"#; - let (_, delta_usage) = frame_delta(PassthroughProtocol::OpenaiChat, delta); - let mut acc = start_usage; - acc.merge(delta_usage.expect("message_delta must report usage")); - // The later, partial report EXTENDS the record instead of - // truncating the prompt side to zero. - assert_eq!( - acc, - PassthroughUsage { - prompt_tokens: 12, - completion_tokens: 7, - cache_creation_tokens: 300, - cache_read_tokens: 1200, - ..Default::default() - } + fn stream_guardrail_text_preserves_response_duplicate_carrier_branches() { + let first = b"data: {\"type\":\"response.output_text.delta\",\"item_id\":\"one\",\"output_index\":0,\"content_index\":0,\"delta\":\"FOR\",\"delta\":\"ok\"}\n\n"; + let second = b"data: {\"type\":\"response.output_text.delta\",\"item_id\":\"one\",\"output_index\":0,\"content_index\":0,\"delta\":\"BIDDEN\"}\n\n"; + let first_parts = frame_parts(PassthroughProtocol::OpenaiResponses, first).0; + let second_parts = frame_parts(PassthroughProtocol::OpenaiResponses, second).0; + let first = stream_guardrail_text( + PassthroughProtocol::OpenaiResponses, + first, + first_parts.scan, + ); + let second = stream_guardrail_text( + PassthroughProtocol::OpenaiResponses, + second, + second_parts.scan, + ); + let mut continuations = Vec::new(); + let mut continuation_tails = Vec::new(); + let mut supplemental = Vec::new(); + let mut closed_prefixes = Vec::new(); + append_stream_guardrail_text( + &mut continuations, + &mut continuation_tails, + &mut supplemental, + &mut closed_prefixes, + &first, + ); + append_stream_guardrail_text( + &mut continuations, + &mut continuation_tails, + &mut supplemental, + &mut closed_prefixes, + &second, + ); + let scanned = + stream_guardrail_scan_text(&continuation_tails, &continuations, &supplemental); + assert!( + scan_candidates_contain(&scanned, "FORBIDDEN"), + "{scanned:?}" ); - - // The nested read is gated on the event type: a `message` object on - // any other frame is not a usage report. - let other = br#"data: {"type":"conversation","message":{"usage":{"input_tokens":999}}} - -"#; - assert_eq!(frame_delta(PassthroughProtocol::OpenaiChat, other).1, None); + assert!(scan_candidates_contain(&scanned, "okBIDDEN"), "{scanned:?}"); } - /// A frame's payload is ALL of its `data:` lines joined with `\n` - /// (WHATWG SSE). Parsing each line on its own turns one document into - /// N unparseable fragments, so the frame's usage went unread and — on - /// a `Raw` stream — its JSON source text was pushed into the guardrail - /// scan instead of its values. #[test] - fn a_payload_spread_over_several_data_lines_is_read_as_one_document() { - let frame = b"event: message_delta\ndata: {\"type\":\"message_delta\",\ndata: \"usage\":{\"output_tokens\":7,\"input_tokens\":12}}\n\n"; - let (_, usage) = frame_delta(PassthroughProtocol::OpenaiChat, frame); - assert_eq!( - usage, - Some(PassthroughUsage { - prompt_tokens: 12, - completion_tokens: 7, - ..Default::default() - }), + fn stream_guardrail_text_keys_reordered_chat_choices_by_source_index() { + let first = b"data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":\"FOR\"}},{\"index\":1,\"delta\":{\"content\":\"noise\"}}]}\n\n"; + let second = b"data: {\"choices\":[{\"index\":1,\"delta\":{\"content\":\"CLEAN\"}},{\"index\":0,\"delta\":{\"content\":\"BIDDEN\"}}]}\n\n"; + let first_parts = frame_parts(PassthroughProtocol::OpenaiChat, first).0; + let second_parts = frame_parts(PassthroughProtocol::OpenaiChat, second).0; + let first = stream_guardrail_text(PassthroughProtocol::OpenaiChat, first, first_parts.scan); + let second = + stream_guardrail_text(PassthroughProtocol::OpenaiChat, second, second_parts.scan); + assert!(!first.unevaluable); + assert!(!second.unevaluable); + let mut continuations = Vec::new(); + let mut continuation_tails = Vec::new(); + let mut supplemental = Vec::new(); + let mut closed_prefixes = Vec::new(); + append_stream_guardrail_text( + &mut continuations, + &mut continuation_tails, + &mut supplemental, + &mut closed_prefixes, + &first, ); - - // The `Raw` scan text is the payload's VALUES for a document that - // parses — never the raw JSON source, which is what a per-line read - // fell back to for each fragment. - let (text, _) = frame_delta(PassthroughProtocol::Raw, frame); - assert_eq!( - text, - "{\"type\":\"message_delta\",\n\"usage\":{\"output_tokens\":7,\"input_tokens\":12}}", + append_stream_guardrail_text( + &mut continuations, + &mut continuation_tails, + &mut supplemental, + &mut closed_prefixes, + &second, + ); + let scanned = + stream_guardrail_scan_text(&continuation_tails, &continuations, &supplemental); + assert!( + scan_candidates_contain(&scanned, "FORBIDDEN"), + "{scanned:?}" + ); + assert!( + !scan_candidates_contain(&scanned, "noiseBIDDEN"), + "{scanned:?}" ); } - /// Framing varies per ENDPOINT, not per vendor: on one host - /// `/v1/audio/transcriptions` streams pure CRLF with `\r\n\r\n` - /// separators and no `event:` lines while `/v1/responses` on the same - /// host is pure LF. The `\r` belongs to the framing and must reach - /// neither the parser nor the scan text. #[test] - fn a_crlf_framed_frame_reads_the_same_as_its_lf_twin() { - let crlf = b"data: {\"usage\":{\"prompt_tokens\":26,\"completion_tokens\":4}}\r\n\r\n"; - let lf = b"data: {\"usage\":{\"prompt_tokens\":26,\"completion_tokens\":4}}\n\n"; - assert_eq!( - frame_delta(PassthroughProtocol::Raw, crlf), - frame_delta(PassthroughProtocol::Raw, lf), + fn source_continuity_refuses_unkeyable_or_over_cap_response_branches() { + let over_cap = br#"{"type":"response.output_text.delta","item_id":"one","output_index":0,"content_index":0,"delta":"a","delta":"b","delta":"c"}"#; + assert!(matches!( + stream_source_continuations(PassthroughProtocol::OpenaiResponses, over_cap), + SourceContinuations::Unevaluable + )); + let missing_id = br#"{"type":"response.output_text.delta","output_index":0,"content_index":0,"delta":"FOR"}"#; + assert!(matches!( + stream_source_continuations(PassthroughProtocol::OpenaiResponses, missing_id), + SourceContinuations::Unevaluable + )); + let refusal = br#"{"type":"response.refusal.delta","item_id":"one","output_index":0,"content_index":0,"delta":"FOR"}"#; + assert!(matches!( + stream_source_continuations(PassthroughProtocol::OpenaiResponses, refusal), + SourceContinuations::Ready(_) + )); + let refusal_missing_content_index = + br#"{"type":"response.refusal.delta","item_id":"one","output_index":0,"delta":"FOR"}"#; + assert!(matches!( + stream_source_continuations( + PassthroughProtocol::OpenaiResponses, + refusal_missing_content_index + ), + SourceContinuations::Unevaluable + )); + let oversized_id = "x".repeat(MAX_STREAM_GUARDRAIL_SOURCE_ID_BYTES + 1); + let oversized_delta = format!( + "{{\"type\":\"response.output_text.delta\",\"item_id\":\"{oversized_id}\",\"output_index\":0,\"content_index\":0,\"delta\":\"safe\"}}" ); - assert_eq!( - frame_delta(PassthroughProtocol::Raw, crlf).1, - Some(usage_dims(26, 4)), + assert!(matches!( + stream_source_continuations( + PassthroughProtocol::OpenaiResponses, + oversized_delta.as_bytes() + ), + SourceContinuations::Unevaluable + )); + let oversized_done = format!( + "data: {{\"type\":\"response.output_item.done\",\"item\":{{\"id\":\"{oversized_id}\"}}}}\n\n" ); - // …and the frame splitter agrees about where such a frame ends. - let mut splitter = SseFrameSplitter::new(); - assert_eq!(splitter.push(crlf), vec![crlf.to_vec()]); + assert!( + stream_guardrail_text( + PassthroughProtocol::OpenaiResponses, + oversized_done.as_bytes(), + String::new(), + ) + .unevaluable + ); + assert!(matches!( + stream_source_continuations(PassthroughProtocol::Raw, br#"{"state":"FOR"}"#), + SourceContinuations::Unevaluable + )); + assert!(matches!( + stream_source_continuations(PassthroughProtocol::Raw, br#""FOR""#), + SourceContinuations::Ready(_) + )); + + let continuations = (0..MAX_STREAM_GUARDRAIL_CHANNELS) + .map(|index| StreamContinuation { + key: format!("responses:\"{index}\":text:first"), + family: format!("responses:\"{index}\""), + identity: "text".to_owned(), + identity_is_ambiguous: false, + text: "safe".to_owned(), + }) + .collect::>(); + let incoming = StreamGuardrailText { + continuations: vec![StreamContinuation { + key: "responses:\"next\":text:first".to_owned(), + family: "responses:\"next\"".to_owned(), + identity: "text".to_owned(), + identity_is_ambiguous: false, + text: "safe".to_owned(), + }], + supplemental: Vec::new(), + unevaluable: false, + supplemental_unevaluable: false, + closed_prefixes: Vec::new(), + }; + assert!(stream_continuation_would_exceed_cap( + &continuations, + &[], + &[], + &incoming, + )); + + let incoming_supplemental = StreamGuardrailText { + continuations: vec![StreamContinuation { + key: "chat:0:content:first".to_owned(), + family: "chat:0".to_owned(), + identity: "content".to_owned(), + identity_is_ambiguous: false, + text: "safe".to_owned(), + }], + supplemental: (0..MAX_STREAM_GUARDRAIL_CHANNELS) + .map(|index| format!("metadata-{index}")) + .collect(), + unevaluable: false, + supplemental_unevaluable: false, + closed_prefixes: Vec::new(), + }; + assert!(stream_continuation_would_exceed_cap( + &[], + &[], + &[], + &incoming_supplemental, + )); } - /// A comment-only frame — the keepalive some relays emit while the - /// upstream thinks — carries no `data:` line at all. It must contribute - /// no usage and no scan text on every protocol, rather than being read - /// as an empty or unparseable payload. #[test] - fn a_comment_only_frame_contributes_nothing() { - for frame in [ - &b": OPENROUTER PROCESSING\n\n"[..], - &b": OPENROUTER PROCESSING\r\n\r\n"[..], - &b": keep-alive\nevent: ping\n\n"[..], - ] { - for protocol in [ - PassthroughProtocol::Raw, - PassthroughProtocol::OpenaiChat, - PassthroughProtocol::OpenaiCompletions, - PassthroughProtocol::OpenaiResponses, - ] { - assert_eq!( - frame_delta(protocol, frame), - (String::new(), None), - "{protocol:?} on {:?}", - String::from_utf8_lossy(frame), - ); - } + fn stream_guardrail_epoch_queue_is_bounded() { + let mut candidate_epochs = Vec::new(); + let mut queued_candidates = 0; + for index in 0..MAX_STREAM_GUARDRAIL_CHANNELS { + assert!(try_queue_stream_guardrail_epoch( + &mut candidate_epochs, + &mut queued_candidates, + vec![format!("candidate-{index}")], + )); + } + assert!(!try_queue_stream_guardrail_epoch( + &mut candidate_epochs, + &mut queued_candidates, + vec!["one-too-many".to_owned()], + )); + + let mut empty_epochs = Vec::new(); + let mut queued_candidates = 0; + for _ in 0..MAX_STREAM_GUARDRAIL_EPOCHS { + assert!(try_queue_stream_guardrail_epoch( + &mut empty_epochs, + &mut queued_candidates, + Vec::new(), + )); } + assert!(!try_queue_stream_guardrail_epoch( + &mut empty_epochs, + &mut queued_candidates, + Vec::new(), + )); } - /// A frame whose joined payload does not parse is still FORWARDED to - /// the client, so producing no scan text for it is a way past an output - /// block rule. Every protocol falls back to the raw payload text — the - /// worst case is a false positive, while the alternative is a bypass. #[test] - fn an_unparseable_payload_still_yields_scan_text_on_every_protocol() { - // Two independent JSON documents on two `data:` lines: joined per - // the SSE spec this is one unparseable payload, and per-line parsing - // used to catch it only incidentally. - let frame = b"data: {\"choices\":[{\"delta\":{\"content\":\"BLOCKME\"}}]}\ndata: {\"choices\":[]}\n\n"; - for protocol in [ - PassthroughProtocol::Raw, - PassthroughProtocol::OpenaiChat, - PassthroughProtocol::OpenaiCompletions, - PassthroughProtocol::OpenaiResponses, - ] { - let (text, _) = frame_delta(protocol, frame); - assert!( - text.contains("BLOCKME"), - "{protocol:?} must still offer the forwarded bytes to the scan, got {text:?}", + fn closed_response_items_still_count_until_a_successful_window_scan() { + let mut continuations = Vec::new(); + let mut tails = Vec::new(); + let mut supplemental = Vec::new(); + let mut closed_prefixes = Vec::new(); + for index in 0..(MAX_STREAM_GUARDRAIL_CHANNELS / 2) { + let delta = format!( + "data: {{\"type\":\"response.output_text.delta\",\"item_id\":\"{index}\",\"output_index\":0,\"content_index\":0,\"delta\":\"safe\"}}\n\n" + ); + let delta = stream_guardrail_text( + PassthroughProtocol::OpenaiResponses, + delta.as_bytes(), + frame_parts(PassthroughProtocol::OpenaiResponses, delta.as_bytes()) + .0 + .scan, + ); + append_stream_guardrail_text( + &mut continuations, + &mut tails, + &mut supplemental, + &mut closed_prefixes, + &delta, + ); + let done = format!( + "data: {{\"type\":\"response.output_item.done\",\"item\":{{\"id\":\"{index}\"}}}}\n\n" + ); + let done = stream_guardrail_text( + PassthroughProtocol::OpenaiResponses, + done.as_bytes(), + String::new(), + ); + append_stream_guardrail_text( + &mut continuations, + &mut tails, + &mut supplemental, + &mut closed_prefixes, + &done, ); } + assert_eq!(continuations.len(), MAX_STREAM_GUARDRAIL_CHANNELS); + assert_eq!(closed_prefixes.len(), MAX_STREAM_GUARDRAIL_CHANNELS / 2); + let next = stream_guardrail_text( + PassthroughProtocol::OpenaiResponses, + b"data: {\"type\":\"response.output_text.delta\",\"item_id\":\"next\",\"output_index\":0,\"content_index\":0,\"delta\":\"safe\"}\n\n", + "safe".to_owned(), + ); + assert!(stream_continuation_would_exceed_cap( + &continuations, + &supplemental, + &closed_prefixes, + &next, + )); } - /// The `[DONE]` sentinel is not content, on either framing. A stream - /// that omits it entirely — OpenAI's Responses API sends none — is the - /// ordinary case, so nothing may depend on having seen one. #[test] - fn the_done_sentinel_contributes_nothing_on_either_framing() { - for frame in [&b"data: [DONE]\n\n"[..], &b"data: [DONE]\r\n\r\n"[..]] { - assert_eq!( - frame_delta(PassthroughProtocol::Raw, frame), - (String::new(), None), - ); + fn stream_guardrail_text_keeps_keyed_normal_forms_evaluable() { + let cases = [ + ( + PassthroughProtocol::OpenaiChat, + b"data: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"ANTHROPIC_TEXT\"}}\n\n".as_slice(), + "ANTHROPIC_TEXT", + ), + ( + PassthroughProtocol::OpenaiChat, + b"data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":[{\"index\":0,\"type\":\"text\",\"text\":\"PART_TEXT\"}]}}]}\n\n".as_slice(), + "PART_TEXT", + ), + ( + PassthroughProtocol::OpenaiChat, + b"data: {\"choices\":[{\"index\":0,\"delta\":{\"tool_calls\":[{\"index\":0,\"function\":{\"arguments\":\"TOOL_ARGS\"}}],\"content\":\"CHAT_TEXT\"}}]}\n\n".as_slice(), + "TOOL_ARGS", + ), + ]; + for (protocol, frame, expected) in cases { + let typed = frame_parts(protocol, frame).0.scan; + let text = stream_guardrail_text(protocol, frame, typed); + assert!(!text.unevaluable, "{frame:?}"); + let scanned = stream_guardrail_scan_text(&[], &text.continuations, &text.supplemental); + assert!(scan_candidates_contain(&scanned, expected), "{scanned:?}"); } + + let identityless_part = b"data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":[{\"type\":\"text\",\"text\":\"NO_ID\"}]}}]}\n\n"; + let typed = frame_parts(PassthroughProtocol::OpenaiChat, identityless_part) + .0 + .scan; + assert!( + stream_guardrail_text(PassthroughProtocol::OpenaiChat, identityless_part, typed) + .unevaluable + ); } - /// Reasoning replayed by the caller is REQUEST text and is scanned; the - /// same field on a buffered RESPONSE is generated reasoning and is out - /// of the output-guardrail scope. One helper, two answers. #[test] - fn replayed_reasoning_is_request_scan_text_and_not_response_scan_text() { - let msg = serde_json::json!({ - "role": "assistant", - "content": "visible", - "reasoning_content": "hidden reasoning payload", - }); - assert!(message_scan_text(&msg, true).contains("hidden reasoning payload")); - assert!(!message_scan_text(&msg, false).contains("hidden reasoning payload")); - assert!(message_scan_text(&msg, false).contains("visible")); + fn chat_content_part_index_survives_optional_id_changes() { + for (first_frame, second_frame) in [ + ( + b"data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":[{\"index\":0,\"type\":\"text\",\"text\":\"FOR\"}]}}]}\n\n".as_slice(), + b"data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":[{\"index\":0,\"id\":\"later\",\"type\":\"text\",\"text\":\"BIDDEN\"}]}}]}\n\n".as_slice(), + ), + ( + b"data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":[{\"index\":0,\"id\":\"first\",\"type\":\"text\",\"text\":\"FOR\"}]}}]}\n\n".as_slice(), + b"data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":[{\"index\":0,\"type\":\"text\",\"text\":\"BIDDEN\"}]}}]}\n\n".as_slice(), + ), + ] { + let first = stream_guardrail_text( + PassthroughProtocol::OpenaiChat, + first_frame, + frame_parts(PassthroughProtocol::OpenaiChat, first_frame).0.scan, + ); + let second = stream_guardrail_text( + PassthroughProtocol::OpenaiChat, + second_frame, + frame_parts(PassthroughProtocol::OpenaiChat, second_frame).0.scan, + ); + assert!(!first.unevaluable); + assert!(!second.unevaluable); + let mut continuations = Vec::new(); + let mut tails = Vec::new(); + let mut supplemental = Vec::new(); + let mut closed_prefixes = Vec::new(); + append_stream_guardrail_text( + &mut continuations, + &mut tails, + &mut supplemental, + &mut closed_prefixes, + &first, + ); + append_stream_guardrail_text( + &mut continuations, + &mut tails, + &mut supplemental, + &mut closed_prefixes, + &second, + ); + assert!(scan_candidates_contain( + &stream_guardrail_scan_text(&tails, &continuations, &supplemental), + "FORBIDDEN", + )); + } } #[test] - fn opaque_stream_reads_flat_usage_only_from_a_labelled_frame() { - // An agent backend reached through a forward-proxy route has no - // recognisable envelope, and reports usage on its own event as a - // flat token object with no `usage` wrapper. - let labelled = b"event:token_usage\ndata:{\"name\":\"\",\"prompt_tokens\":14603,\"completion_tokens\":8,\"cache_read_input_tokens\":14272,\"reasoning_tokens\":8}\n\n"; - let usage = frame_delta(PassthroughProtocol::Raw, labelled) - .1 - .expect("a server-labelled usage frame must report usage"); - assert_eq!(usage.prompt_tokens, 14603); - assert_eq!(usage.completion_tokens, 8); - assert_eq!(usage.cache_read_tokens, 14272); - assert_eq!(usage.reasoning_tokens, 8); - - // The same flat shape on a frame the server did NOT name a usage - // report mints nothing: an opaque stream has no envelope to - // authenticate token-shaped fields against. - let unlabelled = b"event:history\ndata:{\"prompt_tokens\":99,\"completion_tokens\":9}\n\n"; - assert_eq!(frame_delta(PassthroughProtocol::Raw, unlabelled).1, None); - - // The flat allowance is Raw-only — a detected envelope keeps - // reading usage from its own shape. - assert_eq!( - frame_delta(PassthroughProtocol::OpenaiChat, labelled).1, - None + fn stream_continuity_never_joins_distinct_carriers() { + let first = stream_guardrail_text( + PassthroughProtocol::OpenaiResponses, + b"data: {\"type\":\"response.output_text.delta\",\"item_id\":\"one\",\"output_index\":0,\"content_index\":0,\"delta\":\"FOR\"}\n\n", + "FOR".to_owned(), ); - - // An explicit `usage` OBJECT is still read from any opaque frame: - // it is self-describing, and this is the pre-existing behaviour. - let wrapped = - b"event:done\ndata:{\"usage\":{\"prompt_tokens\":4,\"completion_tokens\":2}}\n\n"; - assert_eq!( - frame_delta(PassthroughProtocol::Raw, wrapped).1, - Some(usage_dims(4, 2)) + let second = stream_guardrail_text( + PassthroughProtocol::OpenaiResponses, + b"data: {\"type\":\"response.output_text.delta\",\"item_id\":\"two\",\"output_index\":0,\"content_index\":0,\"delta\":\"BIDDEN\"}\n\n", + "BIDDEN".to_owned(), ); - } - - #[test] - fn buffered_opaque_responses_are_never_probed_for_usage() { - // The no-phantom-tokens guarantee: a REST body that happens to carry - // a usage-shaped object is not a usage report. - let rpc = br#"{"jsonrpc":"2.0","id":1,"result":{"usage":{"prompt_tokens":99}}}"#; - assert_eq!(response_usage(PassthroughProtocol::Raw, None, rpc), None); - let top_level = br#"{"usage":{"prompt_tokens":99,"completion_tokens":9}}"#; - assert_eq!( - response_usage(PassthroughProtocol::Raw, None, top_level), - None + let mut continuations = Vec::new(); + let mut tails = Vec::new(); + let mut supplemental = Vec::new(); + let mut closed_prefixes = Vec::new(); + append_stream_guardrail_text( + &mut continuations, + &mut tails, + &mut supplemental, + &mut closed_prefixes, + &first, ); - } + append_stream_guardrail_text( + &mut continuations, + &mut tails, + &mut supplemental, + &mut closed_prefixes, + &second, + ); + let scanned = stream_guardrail_scan_text(&tails, &continuations, &supplemental); + assert_eq!(scanned, vec!["FOR".to_owned(), "BIDDEN".to_owned()]); - #[test] - fn usage_merge_is_field_wise_max() { - let mut acc = PassthroughUsage { - prompt_tokens: 10, - completion_tokens: 4, - cache_read_tokens: 100, - ..Default::default() - }; - // A repeat of a cumulative usage object, and a partial one, are both - // harmless: no dimension ever regresses. - acc.merge(PassthroughUsage { - prompt_tokens: 10, - completion_tokens: 9, - reasoning_tokens: 3, - ..Default::default() - }); - assert_eq!( - acc, - PassthroughUsage { - prompt_tokens: 10, - completion_tokens: 9, - reasoning_tokens: 3, - cache_read_tokens: 100, - ..Default::default() - } + let scalar = stream_guardrail_text( + PassthroughProtocol::OpenaiChat, + b"data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":\"FOR\"}}]}\n\n", + "FOR".to_owned(), ); + let indexed_part = stream_guardrail_text( + PassthroughProtocol::OpenaiChat, + b"data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":[{\"index\":0,\"type\":\"text\",\"text\":\"BIDDEN\"}]}}]}\n\n", + "BIDDEN".to_owned(), + ); + let mut scalar_continuations = Vec::new(); + let mut scalar_tails = Vec::new(); + let mut scalar_supplemental = Vec::new(); + let mut scalar_closed_prefixes = Vec::new(); + append_stream_guardrail_text( + &mut scalar_continuations, + &mut scalar_tails, + &mut scalar_supplemental, + &mut scalar_closed_prefixes, + &scalar, + ); + assert!(stream_continuation_identity_conflicts( + &scalar_continuations, + &scalar_closed_prefixes, + &indexed_part, + )); } #[test] - fn guardrail_text_covers_tool_calls_on_both_hooks() { - // A deny-listed string hidden in a tool call's arguments must be - // scanned — a benign `content` beside it would otherwise make the - // extraction non-empty and skip the raw-body fallback, letting the - // request pass a check the typed endpoint enforces. - let req = br#"{"model":"m","messages":[{"role":"assistant","content":"ok","tool_calls":[{"function":{"name":"run","arguments":"{\"cmd\":\"SECRET\"}"}}]}]}"#; - let text = request_guardrail_text(PassthroughProtocol::OpenaiChat, req); - assert!(text.contains("ok"), "content still scanned: {text}"); - assert!( - text.contains("SECRET"), - "tool-call arguments scanned: {text}" - ); - - let resp = br#"{"choices":[{"message":{"content":"sure","tool_calls":[{"function":{"name":"run","arguments":"{\"cmd\":\"SECRET\"}"}}]}}]}"#; - let text = response_guardrail_text(PassthroughProtocol::OpenaiChat, resp); - assert!(text.contains("sure")); - assert!(text.contains("SECRET"), "tool-call output scanned: {text}"); + fn responses_index_identity_cannot_switch_mid_stream() { + let indexed = b"data: {\"type\":\"response.output_text.delta\",\"item_id\":\"one\",\"output_index\":0,\"content_index\":0,\"delta\":\"FOR\"}\n\n"; + let missing_indexes = b"data: {\"type\":\"response.output_text.delta\",\"item_id\":\"one\",\"delta\":\"BIDDEN\"}\n\n"; + for (first, second) in [ + (&indexed[..], &missing_indexes[..]), + (&missing_indexes[..], &indexed[..]), + ] { + let first = stream_guardrail_text( + PassthroughProtocol::OpenaiResponses, + first, + frame_parts(PassthroughProtocol::OpenaiResponses, first) + .0 + .scan, + ); + let second = stream_guardrail_text( + PassthroughProtocol::OpenaiResponses, + second, + frame_parts(PassthroughProtocol::OpenaiResponses, second) + .0 + .scan, + ); + assert!(first.unevaluable || second.unevaluable); + } } #[test] - fn requested_model_comes_only_from_a_detected_envelope() { - let chat = br#"{"model":"gpt-4o","messages":[{"role":"user","content":"hi"}]}"#; - assert_eq!( - body_model_name(PassthroughProtocol::OpenaiChat, None, chat), - "gpt-4o" + fn fail_open_unscannable_frame_seals_then_resets_guardrail_continuity() { + let first_frame = b"data: \"FOR\"\n\n"; + let opaque_frame = b"data: {\"state\":\"safe\"}\n\n"; + let second_frame = b"data: \"BIDDEN\"\n\n"; + let first = stream_guardrail_text( + PassthroughProtocol::Raw, + first_frame, + frame_parts(PassthroughProtocol::Raw, first_frame).0.scan, ); - // An opaque body's `model`-shaped key belongs to some other API. - let opaque = br#"{"model":"whatever","config_name":"x"}"#; - assert_eq!(body_model_name(PassthroughProtocol::Raw, None, opaque), ""); - // Caller-supplied, so bounded and control-char free before it - // reaches telemetry. - let hostile = format!( - r#"{{"input":"x","model":"a\u0000b{}"}}"#, - "z".repeat(REQUESTED_MODEL_CAP * 2) + let opaque = stream_guardrail_text( + PassthroughProtocol::Raw, + opaque_frame, + frame_parts(PassthroughProtocol::Raw, opaque_frame).0.scan, ); - let name = body_model_name( - PassthroughProtocol::OpenaiResponses, - None, - hostile.as_bytes(), + let second = stream_guardrail_text( + PassthroughProtocol::Raw, + second_frame, + frame_parts(PassthroughProtocol::Raw, second_frame).0.scan, + ); + assert!(!first.unevaluable); + assert!(opaque.unevaluable); + assert!(!second.unevaluable); + + let mut continuations = Vec::new(); + let mut tails = Vec::new(); + let mut supplemental = Vec::new(); + let mut closed_prefixes = Vec::new(); + let mut sealed_epochs = Vec::new(); + let mut queued_candidates = 0; + append_stream_guardrail_text( + &mut continuations, + &mut tails, + &mut supplemental, + &mut closed_prefixes, + &first, + ); + assert!(seal_stream_guardrail_epoch( + &mut sealed_epochs, + &mut queued_candidates, + &mut continuations, + &mut tails, + &mut supplemental, + &mut closed_prefixes, + )); + append_stream_guardrail_text( + &mut continuations, + &mut tails, + &mut supplemental, + &mut closed_prefixes, + &second, + ); + let scanned = stream_guardrail_scan_text(&tails, &continuations, &supplemental); + assert_eq!(sealed_epochs, vec![vec!["FOR".to_owned()]]); + assert!(scan_candidates_contain(&scanned, "BIDDEN"), "{scanned:?}"); + assert!( + !scan_candidates_contain(&scanned, "FORBIDDEN"), + "{scanned:?}" ); - assert_eq!(name.chars().count(), REQUESTED_MODEL_CAP); - assert!(!name.contains('\0')); } - /// The passthrough route reads the SAME Responses shapes the typed - /// `/v1/responses` handler does, in the same directions. Request: - /// a replayed `reasoning` item's `content` AND `summary` are - /// caller-supplied text and are scanned. Response: a generated - /// `reasoning` item is out of the output scope and must not be — - /// the walk reads `content` off every item regardless of type, so - /// without an explicit skip a block rule matching only inside - /// reasoning refuses a response the typed route allows. #[test] - fn responses_passthrough_scans_replayed_reasoning_but_not_generated_reasoning() { - let request = serde_json::json!({ - "input": [ - {"role": "user", "content": [{"type": "input_text", "text": "VISIBLE"}]}, - { - "type": "reasoning", - "summary": [{"type": "summary_text", "text": "SUMMARYSECRET"}], - "content": [{"type": "reasoning_text", "text": "REASONINGSECRET"}] - }, - {"type": "function_call_output", "call_id": "c1", "output": "TOOLRESULTSECRET"}, - {"type": "mcp_approval_response", "approve": true, "reason": "APPROVALSECRET"} - ] - }) - .to_string(); - let scanned = - request_guardrail_text(PassthroughProtocol::OpenaiResponses, request.as_bytes()); - assert!(scanned.contains("VISIBLE"), "got {scanned:?}"); - assert!(scanned.contains("REASONINGSECRET"), "got {scanned:?}"); - assert!(scanned.contains("SUMMARYSECRET"), "got {scanned:?}"); - // The tool-result and approval slots too. These matter precisely - // because the items beside them yield text: the raw-body fallback - // fires only on a WHOLLY empty extraction, so a mixed body would - // otherwise carry them past the scan while `/v1/responses` blocks - // the same envelope. - assert!(scanned.contains("TOOLRESULTSECRET"), "got {scanned:?}"); - assert!(scanned.contains("APPROVALSECRET"), "got {scanned:?}"); + fn responses_item_done_waits_for_a_successful_scan_before_retiring() { + let delta = b"data: {\"type\":\"response.output_text.delta\",\"item_id\":\"one\",\"output_index\":0,\"content_index\":0,\"delta\":\"FORBIDDEN\"}\n\n"; + let done = b"data: {\"type\":\"response.output_item.done\",\"output_index\":0,\"item\":{\"id\":\"one\",\"type\":\"message\"}}\n\n"; + let delta = stream_guardrail_text( + PassthroughProtocol::OpenaiResponses, + delta, + frame_parts(PassthroughProtocol::OpenaiResponses, delta) + .0 + .scan, + ); + let done = stream_guardrail_text(PassthroughProtocol::OpenaiResponses, done, String::new()); + assert!(!done.unevaluable); + assert_eq!(done.closed_prefixes, vec!["responses:\"one\":".to_owned()]); + + let mut continuations = Vec::new(); + let mut tails = Vec::new(); + let mut supplemental = Vec::new(); + let mut closed_prefixes = Vec::new(); + append_stream_guardrail_text( + &mut continuations, + &mut tails, + &mut supplemental, + &mut closed_prefixes, + &delta, + ); + append_stream_guardrail_text( + &mut continuations, + &mut tails, + &mut supplemental, + &mut closed_prefixes, + &done, + ); + assert!(scan_candidates_contain( + &stream_guardrail_scan_text(&tails, &continuations, &supplemental), + "FORBIDDEN", + )); + retire_scanned_stream_continuations(&mut continuations, &mut tails, &mut closed_prefixes); + assert!(continuations.is_empty()); + assert!(tails.is_empty()); + assert!(closed_prefixes.is_empty()); + + let unknown_done = + b"data: {\"type\":\"response.output_item.done\",\"item\":{\"type\":\"message\"}}\n\n"; + let unknown = stream_guardrail_text( + PassthroughProtocol::OpenaiResponses, + unknown_done, + String::new(), + ); + assert!(!unknown.unevaluable); + assert!(unknown.closed_prefixes.is_empty()); - let response = serde_json::json!({ - "output": [ - { - "type": "reasoning", - "summary": [{"type": "summary_text", "text": "SUMMARYSECRET"}], - "content": [{"type": "reasoning_text", "text": "REASONINGSECRET"}] - }, - { - "type": "message", - "content": [{"type": "output_text", "text": "the visible answer"}] - } - ] - }) - .to_string(); - let scanned = - response_guardrail_text(PassthroughProtocol::OpenaiResponses, response.as_bytes()); - assert!(scanned.contains("the visible answer"), "got {scanned:?}"); - assert!(!scanned.contains("REASONINGSECRET"), "got {scanned:?}"); - // Not a raw-body fallback: the message item yielded text, so a - // green above means the reasoning item was skipped rather than the - // whole walk having come back empty. - assert!(!scanned.contains("\"output\""), "got {scanned:?}"); - } + let conflicting_done = b"data: {\"type\":\"response.output_item.done\",\"item_id\":\"one\",\"item\":{\"id\":\"two\",\"type\":\"message\"}}\n\n"; + assert!( + stream_guardrail_text( + PassthroughProtocol::OpenaiResponses, + conflicting_done, + String::new(), + ) + .unevaluable + ); - #[test] - fn request_text_extraction_per_protocol() { - let chat = br#"{"model":"m","messages":[{"role":"system","content":"s"},{"role":"user","content":[{"type":"text","text":"part"}]}]}"#; - assert_eq!( - request_guardrail_text(PassthroughProtocol::OpenaiChat, chat), - "s\npart" + // A legal done-only item carries its complete visible message and + // closes itself after the current scan. It must not look like a + // post-close continuation or leave one channel behind forever. + let standalone_done = b"data: {\"type\":\"response.output_item.done\",\"output_index\":0,\"item\":{\"id\":\"standalone\",\"type\":\"message\",\"content\":[{\"type\":\"output_text\",\"text\":\"VISIBLE\"}]}}\n\n"; + let standalone = stream_guardrail_text( + PassthroughProtocol::OpenaiResponses, + standalone_done, + frame_parts(PassthroughProtocol::OpenaiResponses, standalone_done) + .0 + .scan, ); - let fim = br#"{"prompt":"def f(","suffix":"return"}"#; + assert!(!standalone.unevaluable); assert_eq!( - request_guardrail_text(PassthroughProtocol::OpenaiCompletions, fim), - "def f(\nreturn" + standalone.closed_prefixes, + vec!["responses:\"standalone\":".to_owned()] ); - // Shape mismatch degrades to the raw body. - let not_chat = br#"{"input":"x"}"#; - assert_eq!( - request_guardrail_text(PassthroughProtocol::OpenaiChat, not_chat), - r#"{"input":"x"}"# + assert!(!stream_continuation_identity_conflicts( + &[], + &[], + &standalone, + )); + let mut standalone_continuations = Vec::new(); + let mut standalone_tails = Vec::new(); + let mut standalone_supplemental = Vec::new(); + let mut standalone_closures = Vec::new(); + append_stream_guardrail_text( + &mut standalone_continuations, + &mut standalone_tails, + &mut standalone_supplemental, + &mut standalone_closures, + &standalone, + ); + assert!(scan_candidates_contain( + &stream_guardrail_scan_text( + &standalone_tails, + &standalone_continuations, + &standalone_supplemental, + ), + "VISIBLE", + )); + retire_scanned_stream_continuations( + &mut standalone_continuations, + &mut standalone_tails, + &mut standalone_closures, ); - // A detected envelope whose items carry no text ALSO degrades to - // the raw body — detection must never scan less than raw would. - let empty_chat = br#"{"messages":[{"role":"tool","tool_call_id":"1"}]}"#; + assert!(standalone_continuations.is_empty()); + assert!(standalone_tails.is_empty()); + assert!(standalone_closures.is_empty()); + } + + #[test] + fn done_sentinel_is_not_an_unevaluable_stream_carrier() { + for protocol in [ + PassthroughProtocol::Raw, + PassthroughProtocol::OpenaiChat, + PassthroughProtocol::OpenaiCompletions, + PassthroughProtocol::OpenaiResponses, + ] { + let text = stream_guardrail_text(protocol, b"data: [DONE]\n\n", String::new()); + assert!(!text.unevaluable, "{protocol:?}"); + } + } + + #[test] + fn stream_guardrail_text_scans_a_visible_carrier_once() { + let email = "carol@example.com"; + let frame = b"data: {\"id\":\"chatcmpl-once\",\"object\":\"chat.completion.chunk\",\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"ask carol@example.com\"}}]}\n\n"; + let typed = frame_parts(PassthroughProtocol::OpenaiChat, frame).0.scan; + let text = stream_guardrail_text(PassthroughProtocol::OpenaiChat, frame, typed); + let scanned = stream_guardrail_scan_text(&[], &text.continuations, &text.supplemental); assert_eq!( - request_guardrail_text(PassthroughProtocol::OpenaiChat, empty_chat), - String::from_utf8_lossy(empty_chat) + scanned + .iter() + .map(|candidate| candidate.matches(email).count()) + .sum::(), + 1, + "typed, raw source, and supplemental channels must not multiply one carrier: {scanned:?}" ); } + #[test] + fn stream_guardrail_text_excludes_unambiguous_anthropic_reasoning() { + let frame = b"data: {\"type\":\"content_block_delta\",\"delta\":{\"type\":\"thinking_delta\",\"text\":\"BLOCKME\"}}\n\n"; + let typed = frame_parts(PassthroughProtocol::OpenaiChat, frame).0.scan; + assert!(typed.contains("BLOCKME"), "{typed:?}"); + let text = stream_guardrail_text(PassthroughProtocol::OpenaiChat, frame, typed); + let scanned = stream_guardrail_scan_text(&[], &text.continuations, &text.supplemental); + assert!(!scan_candidates_contain(&scanned, "BLOCKME"), "{scanned:?}"); + + let signature = b"data: {\"type\":\"content_block_delta\",\"delta\":{\"type\":\"signature_delta\",\"signature\":\"BLOCKME\"}}\n\n"; + let typed = frame_parts(PassthroughProtocol::OpenaiChat, signature) + .0 + .scan; + let text = stream_guardrail_text(PassthroughProtocol::OpenaiChat, signature, typed); + let scanned = stream_guardrail_scan_text(&[], &text.continuations, &text.supplemental); + assert!(!scan_candidates_contain(&scanned, "BLOCKME"), "{scanned:?}"); + } + #[test] fn detect_protocol_from_request_envelope() { // The real Copilot CLI surface, one shape per endpoint family. @@ -3748,6 +9517,63 @@ mod tests { } } + #[test] + fn passthrough_sse_detection_requires_the_exact_media_type() { + for (value, expected) in [ + ("text/event-stream", true), + (" TEXT/EVENT-STREAM ; charset=utf-8", true), + ("text/event-streaming", false), + ("application/json", false), + ] { + let mut headers = HeaderMap::new(); + headers.insert(header::CONTENT_TYPE, HeaderValue::from_str(value).unwrap()); + assert_eq!(response_is_sse(&headers), expected, "{value}"); + } + + let mut repeated = HeaderMap::new(); + repeated.append( + header::CONTENT_TYPE, + HeaderValue::from_static("text/event-stream"), + ); + repeated.append( + header::CONTENT_TYPE, + HeaderValue::from_static("text/event-streaming"), + ); + assert!(!response_is_sse(&repeated)); + } + + #[test] + fn encoded_success_response_is_not_an_inspectable_guardrail_input() { + for (value, encoded) in [ + (None, false), + (Some("identity"), false), + (Some(" IDENTITY , identity "), false), + (Some("gzip"), true), + (Some("gzip, identity"), true), + (Some(""), true), + ] { + let mut headers = HeaderMap::new(); + if let Some(value) = value { + headers.insert( + header::CONTENT_ENCODING, + HeaderValue::from_str(value).unwrap(), + ); + } + assert_eq!( + response_has_non_identity_content_encoding(&headers), + encoded, + "{value:?}" + ); + } + + let mut invalid = HeaderMap::new(); + invalid.insert( + header::CONTENT_ENCODING, + HeaderValue::from_bytes(b"\xff").unwrap(), + ); + assert!(response_has_non_identity_content_encoding(&invalid)); + } + #[test] fn responses_protocol_extracts_prompt_completion_and_usage() { // GitHub's Copilot CLI sends every inference turn to POST @@ -3756,9 +9582,9 @@ mod tests { let req = br#"{"model":"gpt-5","input":[ {"role":"user","content":[{"type":"input_text","text":"list the files"}]} ]}"#; - assert_eq!( - request_guardrail_text(PassthroughProtocol::OpenaiResponses, req), - "list the files" + assert!( + request_guardrail_text(PassthroughProtocol::OpenaiResponses, req) + .contains("list the files") ); // A bare-string input is equally valid. let req_str = br#"{"model":"gpt-5","input":"hello there"}"#; @@ -3770,9 +9596,8 @@ mod tests { let resp = br#"{"output":[ {"type":"message","content":[{"type":"output_text","text":"done"}]} ],"usage":{"input_tokens":11,"output_tokens":3}}"#; - assert_eq!( - response_guardrail_text(PassthroughProtocol::OpenaiResponses, resp), - "done" + assert!( + response_guardrail_text(PassthroughProtocol::OpenaiResponses, resp).contains("done") ); assert_eq!( response_usage(PassthroughProtocol::OpenaiResponses, None, resp), @@ -4195,22 +10020,31 @@ mod tests { src.append("set-cookie", HeaderValue::from_static("a=1")); src.append("set-cookie", HeaderValue::from_static("b=2")); src.append("vary", HeaderValue::from_static("accept")); + src.append("trailer", HeaderValue::from_static("x-upstream-checksum")); let mut dst = HeaderMap::new(); copy_safe_headers(&src, &mut dst); let cookies: Vec<_> = dst.get_all("set-cookie").iter().collect(); assert_eq!(cookies.len(), 2, "both Set-Cookie values must relay"); + assert!( + dst.get("trailer").is_none(), + "a relay that does not forward trailers must not advertise them" + ); } #[test] fn sse_splitter_bounds_an_unterminated_frame() { - let mut s = SseFrameSplitter::new(); + let mut s = SseFrameSplitter::with_max_frame_bytes(MAX_HELD_STREAM_BYTES); // Feed > MAX_HELD_STREAM_BYTES without a frame terminator: the // splitter must hand the oversized run on instead of buffering // without bound. let chunk = vec![b'x'; 256 * 1024]; let mut emitted = 0usize; for _ in 0..8 { - emitted += s.push(&chunk).iter().map(Vec::len).sum::(); + emitted += s + .push(&chunk) + .iter() + .map(|frame| frame.bytes.len()) + .sum::(); } assert!( emitted >= MAX_HELD_STREAM_BYTES, @@ -4219,6 +10053,16 @@ mod tests { assert!(s.take_rest().len() <= MAX_HELD_STREAM_BYTES); } + #[test] + fn sse_splitter_honors_a_route_specific_frame_cap() { + let mut s = SseFrameSplitter::with_max_frame_bytes(4); + let frames = s.push(b"12345"); + assert_eq!(frames.len(), 1); + assert_eq!(frames[0].bytes, b"12345"); + assert!(frames[0].overflowed); + assert!(s.take_rest().is_empty()); + } + #[test] fn push_capped_respects_byte_cap_on_char_boundaries() { let mut buf = String::new(); @@ -4326,9 +10170,25 @@ mod tests { "data: {\"type\":\"response.reasoning_summary_text.delta\",\"delta\":\"plan\"}\n\n", ); assert_eq!((r.scan.as_str(), r.reasoning), ("", 4)); - // Raw counts and scans the whole payload, envelope included. + // A Raw payload without strings falls back to its full source text. let raw = parts(PassthroughProtocol::Raw, "data: {\"x\":1}\n\n"); assert_eq!((raw.scan.as_str(), raw.held()), ("{\"x\":1}", 7)); + let frame = b"data: {\"state\":\"\\u0042LOCKME\",\"state\":\"clean\"}\n\n"; + let raw = frame_parts(PassthroughProtocol::Raw, frame).0; + assert_eq!(raw.scan, "BLOCKME\nclean"); + assert_eq!( + frame_capture_text(PassthroughProtocol::Raw, frame, &raw.scan), + r#"{"state":"\u0042LOCKME","state":"clean"}"# + ); + let deep = deeply_nested_escaped_block_json(); + let frame = format!("data: {}\n\n", String::from_utf8_lossy(&deep)); + assert!( + frame_parts(PassthroughProtocol::Raw, frame.as_bytes()) + .0 + .scan + .contains("BLOCKME"), + "a valid deep Raw SSE payload must not fall back to escaped source" + ); } /// An Anthropic Messages body on the chat envelope is scanned in every @@ -4355,6 +10215,212 @@ mod tests { } } + #[test] + fn nested_anthropic_tool_result_content_keeps_duplicate_source_order() { + let body = br#"{"model":"claude","messages":[{"role":"user","content":[{"type":"tool_result","tool_use_id":"t","content":[{"type":"text","text":"FIRST"}],"content":[{"type":"text","text":"SECOND"}]}]}]}"#; + let text = request_guardrail_text(PassthroughProtocol::OpenaiChat, body); + let first = text.find("FIRST").expect("first content field is scanned"); + let second = text + .find("SECOND") + .expect("second content field is scanned"); + assert!(first < second, "{text}"); + } + + #[test] + fn nested_anthropic_tool_result_content_within_work_cap_is_scanned() { + let body = nested_anthropic_tool_result_request(32, "BLOCKME"); + let text = try_request_guardrail_text(PassthroughProtocol::OpenaiChat, &body) + .expect("nested tool results within the structural-work cap remain evaluable"); + assert!(text.contains("BLOCKME"), "{text}"); + } + + #[test] + fn nested_anthropic_tool_result_content_beyond_work_cap_is_unevaluable() { + let body = nested_anthropic_tool_result_request(crate::json_splice::MAX_JSON_DEPTH, "safe"); + let error = try_request_guardrail_text(PassthroughProtocol::OpenaiChat, &body) + .expect_err("nested tool results beyond the structural-work cap must not recurse"); + assert!(error.is_unevaluable(), "{error}"); + } + + #[test] + fn over_depth_chat_messages_payload_is_unevaluable_before_source_selection() { + let body = nested_chat_messages_payload(crate::json_splice::MAX_JSON_DEPTH + 1); + assert_eq!(detect_protocol(&body), PassthroughProtocol::OpenaiChat); + let error = try_request_guardrail_text(PassthroughProtocol::OpenaiChat, &body) + .expect_err("a messages payload beyond the shared depth cap must not be selected"); + assert!(error.is_depth_exceeded(), "{error}"); + } + + #[test] + fn malformed_chat_carriers_are_unevaluable_without_raw_fallback() { + let cases: [&[u8]; 2] = [ + br#"{"messages":["forbidden"]}"#, + br#"{"messages":[{"role":"user","content":["forbidden"]}]}"#, + ]; + for body in cases { + let error = try_request_guardrail_text(PassthroughProtocol::OpenaiChat, body) + .expect_err("a malformed Chat carrier must not be treated as an empty scan"); + assert!(error.is_unevaluable(), "{error}"); + assert!(!error.is_depth_exceeded(), "{error}"); + assert!( + !request_guardrail_text(PassthroughProtocol::OpenaiChat, body) + .contains("forbidden"), + "opaque carrier source must not become guardrail text" + ); + } + } + + #[test] + fn shallow_source_selector_makes_mismatched_or_invalid_json_unevaluable() { + let cases: [&[u8]; 2] = [br#"{"messages":[}]"#, br#"{"messages":[forbidden]}"#]; + for body in cases { + let error = try_request_guardrail_text(PassthroughProtocol::OpenaiChat, body) + .expect_err("a malformed source selector must not produce an empty scan"); + assert!(error.is_unevaluable(), "{error}"); + } + } + + #[test] + fn shallow_selector_accepts_valid_large_json_numbers() { + let body = br#"{"messages":[{"role":"user","content":"clean"}],"metadata":1e400}"#; + assert_eq!(detect_protocol(body), PassthroughProtocol::OpenaiChat); + assert!(try_request_guardrail_text(PassthroughProtocol::OpenaiChat, body).is_ok()); + } + + #[test] + fn malformed_json_like_raw_probe_cannot_fallback_after_typed_media() { + let malformed = br#"{"messages":[{"role":"user","content":[{"type":"image","source":{"data":"BLOCKME"}}]}],"broken":"#; + assert_eq!(detect_protocol(malformed), PassthroughProtocol::Raw); + let JsonProbe::Unevaluable(error) = probe_json(malformed) else { + panic!("a malformed JSON-like root must retain its unsafe selector state"); + }; + assert!(error.is_unevaluable(), "{error}"); + assert!( + try_request_guardrail_text(PassthroughProtocol::Raw, malformed) + .expect_err("a malformed JSON-like Raw body cannot be source-scanned") + .is_unevaluable() + ); + + let plain = b"plain raw BLOCKME"; + assert!(matches!(probe_json(plain), JsonProbe::Raw)); + let error = try_request_guardrail_text(PassthroughProtocol::Raw, plain) + .expect_err("plain text is not JSON"); + assert!(!error.is_unevaluable(), "{error}"); + assert!( + request_guardrail_text(PassthroughProtocol::Raw, plain).contains("BLOCKME"), + "ordinary non-JSON Raw bodies retain their text fallback" + ); + } + + #[test] + fn wide_typed_carriers_exceeding_selector_budgets_are_unevaluable() { + let assert_unevaluable = |body: &[u8]| { + assert_eq!(detect_protocol(body), PassthroughProtocol::OpenaiChat); + let error = try_request_guardrail_text(PassthroughProtocol::OpenaiChat, body) + .expect_err("a bounded typed selector must fail closed"); + assert!(error.is_unevaluable(), "{error}"); + }; + + let repeated_messages = format!( + "{{{}\"messages\":[]}}", + "\"messages\":[],".repeat(MAX_RAW_SELECTOR_ITEMS), + ); + assert_unevaluable(repeated_messages.as_bytes()); + + let message = r#"{"role":"user","content":"safe"}"#; + let wide_messages = format!( + r#"{{"messages":[{}]}}"#, + format!("{message},").repeat(MAX_RAW_SELECTOR_ITEMS) + message, + ); + assert_unevaluable(wide_messages.as_bytes()); + + let part = r#"{"type":"text","text":"safe"}"#; + let wide_content = format!( + r#"{{"messages":[{{"role":"user","content":[{}]}}]}}"#, + format!("{part},").repeat(MAX_RAW_SELECTOR_ITEMS) + part, + ); + assert_unevaluable(wide_content.as_bytes()); + + let sibling = r#"{"type":"text","text":"safe"}"#; + let nested_work = format!( + r#"{{"messages":[{{"role":"user","content":[{{"type":"tool_result","content":"safe","content":"safe"}},{}]}}]}}"#, + format!("{sibling},").repeat(MAX_RAW_SELECTOR_ITEMS - 2) + sibling, + ); + assert_unevaluable(nested_work.as_bytes()); + + let large_text = "x".repeat(MAX_RAW_SELECTOR_BYTES + 1); + let wide_bytes = format!(r#"{{"messages":[{{"role":"user","content":"{large_text}"}}]}}"#); + assert_unevaluable(wide_bytes.as_bytes()); + } + + #[test] + fn json_text_and_typed_output_selector_caps_are_unevaluable() { + let raw = format!( + r#"{{"state":"{}"}}"#, + "x".repeat(MAX_RAW_SELECTOR_BYTES + 1), + ); + assert!(matches!(probe_json(raw.as_bytes()), JsonProbe::Raw)); + let error = try_request_guardrail_text(PassthroughProtocol::Raw, raw.as_bytes()) + .expect_err("JSON-like Raw text over the cap must not fall back to source bytes"); + assert!(error.is_unevaluable(), "{error}"); + let raw_stream = format!(r#""{}""#, "x".repeat(MAX_RAW_SELECTOR_BYTES + 1)); + assert!(matches!( + stream_source_continuations(PassthroughProtocol::Raw, raw_stream.as_bytes()), + SourceContinuations::Unevaluable + )); + + let assert_output_unevaluable = |protocol, body: &str| { + let error = try_response_guardrail_text(protocol, body.as_bytes()) + .expect_err("a capped typed output selector must fail closed"); + assert!(error.is_unevaluable(), "{protocol:?}: {error}"); + assert!( + !response_guardrail_text(protocol, body.as_bytes()).contains("BLOCKME"), + "{protocol:?} output must not fall back to opaque source", + ); + }; + + let choice = r#"{"message":{"content":"safe"}}"#; + let wide_choices = format!( + r#"{{"choices":[{}]}}"#, + format!("{choice},").repeat(MAX_RAW_SELECTOR_ITEMS) + choice, + ); + assert_output_unevaluable(PassthroughProtocol::OpenaiChat, &wide_choices); + + let part = r#"{"type":"text","text":"safe"}"#; + let wide_content = format!( + r#"{{"choices":[{{"message":{{"content":[{}]}}}}]}}"#, + format!("{part},").repeat(MAX_RAW_SELECTOR_ITEMS) + part, + ); + assert_output_unevaluable(PassthroughProtocol::OpenaiChat, &wide_content); + + let item = r#"{"type":"message","content":"safe"}"#; + let wide_output = format!( + r#"{{"output":[{}]}}"#, + format!("{item},").repeat(MAX_RAW_SELECTOR_ITEMS) + item, + ); + assert_output_unevaluable(PassthroughProtocol::OpenaiResponses, &wide_output); + + let overflowing_type_part = format!( + r#"{{{}"text":"BLOCKME"}}"#, + r#""type":"text","#.repeat(MAX_RAW_SELECTOR_ITEMS + 1), + ); + let chat = + format!(r#"{{"choices":[{{"message":{{"content":[{overflowing_type_part}]}}}}]}}"#); + assert_output_unevaluable(PassthroughProtocol::OpenaiChat, &chat); + let responses = + format!(r#"{{"output":[{{"type":"message","content":[{overflowing_type_part}]}}]}}"#); + assert_output_unevaluable(PassthroughProtocol::OpenaiResponses, &responses); + + let request = + format!(r#"{{"input":[{{"type":"message","content":[{overflowing_type_part}]}}]}}"#); + let error = + try_request_guardrail_text(PassthroughProtocol::OpenaiResponses, request.as_bytes()) + .expect_err( + "a strict typed input part cannot treat a capped type selector as absent", + ); + assert!(error.is_unevaluable(), "{error}"); + } + /// Buffered Anthropic and Responses replies are read slot by slot: /// text and tool input in, generated reasoning out. #[test] @@ -4431,6 +10497,7 @@ mod tests { const OUTPUT_KEYWORD_BLOCK: &str = r#"{"name":"out-block","enabled":true,"hook_point":"output","kind":"keyword","patterns":[{"kind":"literal","value":"FORBIDDEN"}]}"#; const OUTPUT_CAP_FAIL_CLOSED: &str = r#"{"name":"out-cap","enabled":true,"hook_point":"output","kind":"azure_content_safety_text_moderation","endpoint":"http://127.0.0.1:1","api_key":"k","stream_processing_mode":"buffer_full","max_buffer_bytes":4,"on_buffer_exceeded":"fail_closed"}"#; + const OUTPUT_FAIL_OPEN: &str = r#"{"name":"out-open","enabled":true,"hook_point":"output","kind":"openai_moderation","endpoint":"http://127.0.0.1:1","api_key":"k","output_fail_open":true}"#; const ANTHROPIC_SSE: &str = "event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"id\":\"m\",\"type\":\"message\",\"role\":\"assistant\",\"content\":[],\"model\":\"claude\",\"usage\":{\"input_tokens\":3,\"output_tokens\":0}}}\n\n\ event: content_block_start\ndata: {\"type\":\"content_block_start\",\"index\":0,\"content_block\":{\"type\":\"text\",\"text\":\"\"}}\n\n\ @@ -4443,6 +10510,8 @@ event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"; data: {\"id\":\"c\",\"object\":\"chat.completion.chunk\",\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}]}\n\n\ data: [DONE]\n\n"; + const MALFORMED_SUPPLEMENTAL_SSE: &str = "event: content_block_start\ndata: {\"type\":\"content_block_start\",\"index\":0,\"content_block\":{\"type\":\"tool_use\",\"input\":\"clean\",\"name\":1}}\n\n"; + /// A refusal ending a relayed Anthropic Messages stream is the frame /// `/v1/messages` emits for it: an SDK-legal `error.type` and the /// refusal named on `error.code`. @@ -4476,4 +10545,11 @@ data: [DONE]\n\n"; assert_eq!(v["error"]["type"], "content_filter", "{v}"); assert_eq!(v["error"]["code"], "guardrail_unavailable", "{v}"); } + + #[tokio::test] + async fn a_malformed_supplemental_stream_selector_ignores_output_fail_open() { + let v = relayed_refusal_frame(MALFORMED_SUPPLEMENTAL_SSE, OUTPUT_FAIL_OPEN).await; + assert_eq!(v["error"]["type"], "invalid_request_error", "{v}"); + assert_eq!(v["error"]["code"], "guardrail_unavailable", "{v}"); + } } diff --git a/crates/aisix-proxy/src/redact.rs b/crates/aisix-proxy/src/redact.rs index 993bb8325..e143c8bf6 100644 --- a/crates/aisix-proxy/src/redact.rs +++ b/crates/aisix-proxy/src/redact.rs @@ -986,7 +986,8 @@ fn responses_item_is_system(item: &Value) -> bool { } /// One `/v1/responses` input/output item. `message` items carry -/// string-or-parts content (`input_text` / `output_text` / plain `text`); +/// string-or-parts content (`input_text` / `output_text` / plain `text` / +/// `refusal`); /// `function_call` carries JSON-encoded `arguments`; /// `function_call_output` carries a string `output`. fn redact_responses_item( @@ -1001,13 +1002,17 @@ fn redact_responses_item( Some(v @ Value::String(_)) => apply_to_value_string(chain, dir, v, counts), Some(Value::Array(parts)) => { for part in parts { - if matches!( - part.get("type").and_then(Value::as_str), - Some("input_text") | Some("output_text") | Some("text") - ) { - if let Some(text) = part.get_mut("text") { - apply_to_value_string(chain, dir, text, counts); + let field = match part.get("type").and_then(Value::as_str) { + Some("input_text" | "output_text" | "text") => Some("text"), + Some("refusal") + if matches!(dir, Direction::Output | Direction::OutputEcho) => + { + Some("refusal") } + _ => None, + }; + if let Some(field) = field.and_then(|field| part.get_mut(field)) { + apply_to_value_string(chain, dir, field, counts); } } } @@ -1146,12 +1151,11 @@ fn apply_to_text_slot( } /// Mask a `/v1/responses` non-streaming RESPONSE body in place — the same -/// surface the output check scans: message `output_text`, `text` and -/// `input_text` parts, and each item's tool-call `name` (scan-only) / -/// `arguments` / `input`. Hosted-tool results (`mcp_call.output`, -/// `file_search_call` results, `code_interpreter_call` logs), `refusal` -/// parts, and hosted-tool code or action text are neither scanned nor -/// masked. +/// surface the output check scans: message `output_text`, `text`, +/// `input_text`, and `refusal` parts, and each item's tool-call `name` +/// (scan-only) / `arguments` / `input`. Hosted-tool results +/// (`mcp_call.output`, `file_search_call` results, `code_interpreter_call` +/// logs) and hosted-tool code or action text are neither scanned nor masked. pub fn redact_responses_response(chain: &dyn Guardrail, body: &mut Value) -> RedactionCounts { let mut counts = RedactionCounts::new(); if !chain.redacts_output() { @@ -2026,9 +2030,9 @@ pub fn anthropic_sse_text(raw: &[u8]) -> String { channels.into_values().collect() } -/// The concatenated `output_text` delta content of a buffered -/// `/v1/responses` SSE stream (channel order). Same capture-rebuild role -/// as [`anthropic_sse_text`]. +/// The concatenated client-visible delta content of a buffered `/v1/responses` +/// SSE stream (channel order). Same capture-rebuild role as +/// [`anthropic_sse_text`]. pub fn responses_sse_text(raw: &[u8]) -> String { let (frames, _) = split_sse_frames(raw); // First-seen channel order (NOT key order): the rebuilt capture must @@ -2038,21 +2042,22 @@ pub fn responses_sse_text(raw: &[u8]) -> String { let Some(data) = frame.data.as_ref() else { continue; }; - if data.get("type").and_then(Value::as_str) != Some("response.output_text.delta") { - continue; - } + let kind = match data.get("type").and_then(Value::as_str) { + Some(kind @ ("response.output_text.delta" | "response.refusal.delta")) => kind, + _ => continue, + }; let Some(t) = data.get("delta").and_then(Value::as_str) else { continue; }; let key = match data.get("item_id").and_then(Value::as_str) { Some(id) => format!( - "{id}/{}", + "{kind}/{id}/{}", data.get("content_index") .and_then(Value::as_u64) .unwrap_or(0) ), None => format!( - "{}/{}", + "{kind}/{}/{}", data.get("output_index") .and_then(Value::as_u64) .unwrap_or(0), @@ -2073,7 +2078,7 @@ pub fn responses_sse_text(raw: &[u8]) -> String { /// Mask a fully-buffered Responses-API SSE byte stream (the `/v1/responses` /// verbatim hold-back and the cross-provider bridge release). Delta events -/// are reassembled per channel (`output_text.delta`, +/// are reassembled per channel (`output_text.delta`, `refusal.delta`, /// `function_call_arguments.delta`, `mcp_call_arguments.delta`, /// `custom_tool_call_input.delta`, each by item), masked once, and /// re-emitted on the channel's first frame; the aggregate events carry @@ -2081,8 +2086,8 @@ pub fn responses_sse_text(raw: &[u8]) -> String { /// event by its own arm, and `output_item.done` / `response.completed` /// (`.incomplete`, `.failed`) through [`redact_responses_item`], so the /// latter cover the same slots as [`redact_responses_response`] and no -/// more. Deterministic masking keeps them consistent with the delta -/// channels. +/// more. Both `*.added` and `*.done` aggregate forms are authoritative; +/// deterministic masking keeps them consistent with the delta channels. /// `None` = nothing matched, forward the original bytes byte-identical. pub fn redact_responses_sse( chain: &dyn Guardrail, @@ -2125,13 +2130,16 @@ pub fn redact_responses_sse( continue; }; match data.get("type").and_then(Value::as_str) { - Some("response.output_text.delta") => { + Some(ty @ ("response.output_text.delta" | "response.refusal.delta")) => { if data .get("delta") .and_then(Value::as_str) .is_some_and(|t| !t.is_empty()) { - text_channels.entry(channel_key(data)).or_default().push(fi); + text_channels + .entry(format!("{ty}/{}", channel_key(data))) + .or_default() + .push(fi); } } // MCP tool calls stream their JSON-encoded arguments on their @@ -2248,19 +2256,30 @@ pub fn redact_responses_sse( apply_to_value_string(chain, aggregate, text, &mut local); } } + "response.refusal.done" => { + if let Some(refusal) = data.get_mut("refusal") { + apply_to_value_string(chain, aggregate, refusal, &mut local); + } + } // The part union also carries `reasoning_text` (a reasoning // item's content), which is generated reasoning and out of the // output scope — the same part types `redact_responses_item` // walks on a message, and no others. - "response.content_part.done" => { - let part = data.get_mut("part").filter(|p| { - matches!( - p.get("type").and_then(Value::as_str), - Some("output_text" | "text" | "input_text") - ) - }); - if let Some(text) = part.and_then(|p| p.get_mut("text")) { - apply_to_value_string(chain, aggregate, text, &mut local); + "response.content_part.added" | "response.content_part.done" => { + let part = data.get_mut("part"); + let field = part + .as_deref() + .and_then(|part| part.get("type")) + .and_then(Value::as_str) + .and_then(|kind| match kind { + "output_text" | "text" | "input_text" => Some("text"), + "refusal" => Some("refusal"), + _ => None, + }); + if let Some(slot) = + field.and_then(|field| part.and_then(|part| part.get_mut(field))) + { + apply_to_value_string(chain, aggregate, slot, &mut local); } } "response.function_call_arguments.done" | "response.mcp_call_arguments.done" => { @@ -2275,7 +2294,7 @@ pub fn redact_responses_sse( apply_to_value_string(chain, aggregate, input, &mut local); } } - "response.output_item.done" => { + "response.output_item.added" | "response.output_item.done" => { if let Some(item) = data.get_mut("item") { redact_responses_item(chain, aggregate, item, &mut local); } @@ -3226,6 +3245,34 @@ mod tests { assert_eq!(counts.get("email"), Some(&1)); } + #[test] + fn responses_sse_masks_refusal_delta_and_aggregate_events() { + let chain = both(); + let raw = concat!( + "event: response.refusal.delta\ndata: {\"type\":\"response.refusal.delta\",\"item_id\":\"msg_1\",\"output_index\":0,\"content_index\":0,\"delta\":\"mail a@\"}\n\n", + "event: response.refusal.delta\ndata: {\"type\":\"response.refusal.delta\",\"item_id\":\"msg_1\",\"output_index\":0,\"content_index\":0,\"delta\":\"x.com now\"}\n\n", + "event: response.refusal.done\ndata: {\"type\":\"response.refusal.done\",\"item_id\":\"msg_1\",\"output_index\":0,\"content_index\":0,\"refusal\":\"mail a@x.com now\"}\n\n", + "event: response.content_part.done\ndata: {\"type\":\"response.content_part.done\",\"part\":{\"type\":\"refusal\",\"refusal\":\"mail a@x.com now\"}}\n\n", + "event: response.completed\ndata: {\"type\":\"response.completed\",\"response\":{\"output\":[{\"type\":\"message\",\"content\":[{\"type\":\"refusal\",\"refusal\":\"mail a@x.com now\"}]}]}}\n\n", + ); + let (out, counts) = redact_responses_sse(chain.as_ref(), raw.as_bytes()).unwrap(); + let out = String::from_utf8(out).unwrap(); + assert!(!out.contains("a@x.com"), "original must be gone: {out}"); + assert!( + out.contains("\"delta\":\"mail [EMAIL_REDACTED] now\""), + "out: {out}" + ); + assert!( + out.contains("\"refusal\":\"mail [EMAIL_REDACTED] now\""), + "out: {out}" + ); + assert_eq!( + responses_sse_text(out.as_bytes()), + "mail [EMAIL_REDACTED] now" + ); + assert_eq!(counts.get("email"), Some(&1)); + } + /// #1027: a Responses stream restates each delta channel on its /// aggregate events. The mask pass counts a span once, so a collector /// must tell the restatements apart — or a monitor preview counts the @@ -4004,7 +4051,8 @@ mod tests { "id": "resp_1", "output": [ {"type": "message", "role": "assistant", "content": [ - {"type": "output_text", "text": "mail a@x.com"} + {"type": "output_text", "text": "mail a@x.com"}, + {"type": "refusal", "refusal": "mail c@z.io"} ]}, {"type": "function_call", "call_id": "c", "name": "send", "arguments": "{\"to\":\"b@y.org\"}"} @@ -4015,11 +4063,15 @@ mod tests { body["output"][0]["content"][0]["text"], "mail [EMAIL_REDACTED]" ); + assert_eq!( + body["output"][0]["content"][1]["refusal"], + "mail [EMAIL_REDACTED]" + ); assert_eq!( body["output"][1]["arguments"], "{\"to\":\"[EMAIL_REDACTED]\"}" ); - assert_eq!(counts.get("email"), Some(&2)); + assert_eq!(counts.get("email"), Some(&3)); } #[test] diff --git a/crates/aisix-proxy/src/responses.rs b/crates/aisix-proxy/src/responses.rs index 0ad4c5a9d..d29eae555 100644 --- a/crates/aisix-proxy/src/responses.rs +++ b/crates/aisix-proxy/src/responses.rs @@ -1461,6 +1461,10 @@ async fn responses_to_target( // which `max_buffer_bytes` caps (the SSE/JSON envelope is not // counted), and every raw byte, unterminated tail included. let mut held_content = crate::held_content::HeldBuffer::default(); + // A Responses stream repeats one logical output carrier in + // delta, done, part, item, and terminal events. Keep that + // bounded identity ledger across the entire held response. + let mut responses_held_content = crate::held_content::ResponsesHeldContent::default(); let mut counted_upto = 0usize; let mut exceeded = false; loop { @@ -1520,9 +1524,9 @@ async fn responses_to_target( } let end = counted_upto + crate::redact::last_frame_end(&buf[counted_upto..]); held_content.hold( - crate::held_content::sse_frames( + crate::held_content::responses_sse_held_frames( + &mut responses_held_content, &buf[counted_upto..end], - crate::held_content::responses_event, ), chunk.len(), ); @@ -1535,9 +1539,9 @@ async fn responses_to_target( // A final frame the upstream never terminated is held content too. if !exceeded && counted_upto < buf.len() { held_content.hold( - crate::held_content::sse_frames( + crate::held_content::responses_sse_held_frames( + &mut responses_held_content, &buf[counted_upto..], - crate::held_content::responses_event, ), 0, ); @@ -1691,7 +1695,7 @@ async fn responses_to_target( usage.usage_estimated = true; } } - let mut out_text = responses_sse_output_text(&buf); + let mut out_text = responses_sse_guardrail_text(&buf); // #1100: an excised frame is not released, but a forbidden // literal inside it must still block the response — the block // pass reads raw text, so it can scan a payload nothing could @@ -3178,6 +3182,11 @@ struct SseTextCapture { /// the end-of-stream scan rebuilds from it the slots the buffered /// branch's mask walker would rewrite (#1027). terminal_response: Option, + /// A standalone visible event occurred before (or without) a terminal + /// snapshot. Its text must join the monitor scan even when the terminal + /// response is present, because an upstream can send a mismatching + /// authoritative aggregate after clean deltas. + saw_non_terminal_visible: bool, } impl SseTextCapture { @@ -3187,11 +3196,15 @@ impl SseTextCapture { deltas: String::new(), terminal: None, terminal_response: None, + saw_non_terminal_visible: false, } } /// The local-kind segments of the terminal response, when it was kept. fn terminal_segments(&self) -> Option> { + if self.saw_non_terminal_visible { + return None; + } let mut resp = self.terminal_response.clone()?; Some(crate::redact::collect_segments(|g| { let _ = crate::redact::redact_responses_response(g, &mut resp); @@ -3210,19 +3223,15 @@ impl SseTextCapture { } } } - Some( - "response.output_text.delta" - | "response.function_call_arguments.delta" - | "response.mcp_call_arguments.delta" - | "response.custom_tool_call_input.delta", - ) => { - if let Some(d) = json.get("delta").and_then(|d| d.as_str()) { + _ => { + let text = responses_stream_event_text(json); + if !text.is_empty() { + self.saw_non_terminal_visible = true; if self.deltas.len() < self.cap { - self.deltas.push_str(d); + self.deltas.push_str(&text); } } } - _ => {} } } @@ -3232,12 +3241,19 @@ impl SseTextCapture { self.terminal.unwrap_or(self.deltas) } - /// [`Self::into_text`] without consuming — the end-of-stream scan reads - /// the text while the completion guard stays armed (a client disconnect - /// mid-scan must still fire the guard's Drop emit with the captured - /// text), so it clones instead of taking. - fn text(&self) -> String { - self.terminal.clone().unwrap_or_else(|| self.deltas.clone()) + /// All authoritative output carriers for the end-of-stream monitor. + /// Capture retains its terminal preference, but a monitor must observe a + /// visible standalone `.done`/item/part event even if the terminal + /// snapshot later disagrees with it. + fn scan_text(&self) -> String { + let mut text = self.deltas.clone(); + if let Some(terminal) = &self.terminal { + if !text.is_empty() && !terminal.is_empty() { + text.push('\n'); + } + text.push_str(terminal); + } + text } } @@ -3493,7 +3509,10 @@ where let hits = match eos_scan { Some(scan) => { let capture = guard.parts().1; - let text = capture.as_ref().map(|c| c.text()).unwrap_or_default(); + let text = capture + .as_ref() + .map(|c| c.scan_text()) + .unwrap_or_default(); let segments = capture.and_then(|c| c.terminal_segments()); scan.observe(&text, segments).await } @@ -3531,24 +3550,30 @@ pub(crate) fn responses_output_text(resp: &Value) -> String { let Some(items) = resp.get("output").and_then(|v| v.as_array()) else { return String::new(); }; + items + .iter() + .map(responses_output_item_text) + .filter(|text| !text.is_empty()) + .collect::>() + .join("\n") +} + +/// The client-visible text on one Responses output item. This is shared by +/// terminal `response.*` snapshots and standalone `response.output_item.*` +/// stream events so either wire form receives the same output protection. +fn responses_output_item_text(item: &Value) -> String { + if item.get("type").and_then(|t| t.as_str()) == Some("reasoning") { + return String::new(); + } let mut parts: Vec<&str> = Vec::new(); - for it in items { - if it.get("type").and_then(|t| t.as_str()) == Some("reasoning") { - continue; - } - if let Some(content) = it.get("content").and_then(|c| c.as_array()) { - parts.extend( - content - .iter() - .filter_map(|p| p.get("text").and_then(|t| t.as_str())), - ); - } - // Tool-call items carry caller-visible model output under top-level - // `name`/`arguments` (function_call) or `name`/`input` (custom tool). - for key in ["name", "arguments", "input"] { - if let Some(s) = it.get(key).and_then(|v| v.as_str()) { - parts.push(s); - } + if let Some(content) = item.get("content").and_then(|c| c.as_array()) { + parts.extend(content.iter().filter_map(responses_visible_part_text)); + } + // Tool-call items carry caller-visible model output under top-level + // `name`/`arguments` (function_call) or `name`/`input` (custom tool). + for key in ["name", "arguments", "input"] { + if let Some(s) = item.get(key).and_then(|v| v.as_str()) { + parts.push(s); } } parts @@ -3558,6 +3583,60 @@ pub(crate) fn responses_output_text(resp: &Value) -> String { .join("\n") } +/// One client-visible Responses message part. A refusal has its own field, +/// unlike the ordinary text part types. +fn responses_visible_part_text(part: &Value) -> Option<&str> { + match part.get("type").and_then(Value::as_str) { + Some("output_text" | "text" | "input_text") => part.get("text").and_then(Value::as_str), + Some("refusal") => part.get("refusal").and_then(Value::as_str), + _ => None, + } +} + +/// Client-visible text from one official Responses streaming event. A +/// standalone aggregate is authoritative: providers may omit deltas and end +/// after this frame, so it cannot rely on a later terminal snapshot for +/// output enforcement. +fn responses_stream_event_text(event: &Value) -> String { + let fields = |fields: &[&str]| { + fields + .iter() + .filter_map(|field| event.get(*field).and_then(Value::as_str)) + .filter(|text| !text.is_empty()) + .collect::>() + .join("\n") + }; + match event.get("type").and_then(Value::as_str) { + Some( + "response.output_text.delta" + | "response.refusal.delta" + | "response.function_call_arguments.delta" + | "response.mcp_call_arguments.delta" + | "response.custom_tool_call_input.delta", + ) => fields(&["delta"]), + Some("response.output_text.done") => fields(&["text"]), + Some("response.refusal.done") => fields(&["refusal"]), + Some("response.function_call_arguments.done" | "response.mcp_call_arguments.done") => { + fields(&["name", "arguments"]) + } + Some("response.custom_tool_call_input.done") => fields(&["name", "input"]), + Some("response.content_part.added" | "response.content_part.done") => event + .get("part") + .and_then(responses_visible_part_text) + .unwrap_or_default() + .to_owned(), + Some("response.output_item.added" | "response.output_item.done") => event + .get("item") + .map(responses_output_item_text) + .unwrap_or_default(), + Some("response.completed" | "response.incomplete" | "response.failed") => event + .get("response") + .map(responses_output_text) + .unwrap_or_default(), + _ => String::new(), + } +} + /// Collect the assistant's streamed output text from a buffered /// Responses-API SSE response (#719/#546). Prefers the authoritative full /// output carried on a terminal `response` event — `response.completed`, @@ -3565,7 +3644,8 @@ pub(crate) fn responses_output_text(resp: &Value) -> String { /// same full `output[]` (incl. tool-call items) and fire routinely (e.g. /// `max_output_tokens` truncation). Falls back to concatenating the streamed /// deltas when no terminal `response` object is present (truncated/aborted): -/// both `response.output_text.delta` (assistant text) and +/// both `response.output_text.delta` / `response.refusal.delta` (assistant +/// text or refusal) and /// `response.function_call_arguments.delta` (tool-call args stream via their /// own event, NOT output_text) — otherwise blocked tool-call args would leak /// on a stream that never reaches a terminal object. The `type` field on each @@ -3595,7 +3675,7 @@ fn responses_sse_output_text(bytes: &[u8]) -> String { } } } - Some("response.output_text.delta") => { + Some("response.output_text.delta" | "response.refusal.delta") => { if let Some(d) = json.get("delta").and_then(|d| d.as_str()) { deltas.push_str(d); } @@ -3623,6 +3703,46 @@ fn responses_sse_output_text(bytes: &[u8]) -> String { deltas } +/// The text the buffered output guardrail evaluates. Unlike content capture +/// and token estimation, it deliberately retains every authoritative stream +/// carrier: a clean delta followed by a malicious `.done`, item, part, or +/// terminal snapshot must still block rather than being hidden by terminal +/// precedence. +fn responses_sse_guardrail_text(bytes: &[u8]) -> String { + let mut text = String::new(); + let mut previous_was_delta = false; + for payload in crate::redact::sse_frame_payloads(bytes) { + let data = payload.trim(); + if data.is_empty() || data == "[DONE]" { + continue; + } + let Ok(event) = serde_json::from_str::(data) else { + continue; + }; + let kind = event.get("type").and_then(Value::as_str); + let is_delta = matches!( + kind, + Some( + "response.output_text.delta" + | "response.refusal.delta" + | "response.function_call_arguments.delta" + | "response.mcp_call_arguments.delta" + | "response.custom_tool_call_input.delta" + ) + ); + let event_text = responses_stream_event_text(&event); + if event_text.is_empty() { + continue; + } + if !(text.is_empty() || previous_was_delta && is_delta) { + text.push('\n'); + } + text.push_str(&event_text); + previous_was_delta = is_delta; + } + text +} + /// Copy the upstream `content-type` onto the client response and stamp the /// `x-aisix-request-id` header. Shared by the streaming verbatim-passthrough /// and buffered hold-back paths. @@ -4900,6 +5020,49 @@ mod tests { ); } + /// A Responses refusal is client-visible assistant output, but it uses a + /// `refusal` field instead of a `text` field. It must not bypass output + /// guardrails on the non-streaming route. + #[tokio::test] + async fn output_guardrail_blocks_non_streaming_refusal_response() { + let upstream = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/v1/responses")) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "id": "resp_refusal", + "object": "response", + "output": [{"type":"message","role":"assistant","content":[{"type":"refusal","refusal":"sure: BLOCKME here"}]}], + "usage": {"input_tokens": 5, "output_tokens": 4} + }))) + .expect(1) + .mount(&upstream) + .await; + + let snap = new_snap_openai(&upstream.uri()); + snap.models.insert(openai_model("gpt-4o-resp")); + snap.apikeys.insert(apikey_entry(&["*"])); + crate::seed_env_scoped_guardrail(&snap, keyword_output_guardrail("BLOCKME")); + let app = build_app(snap); + + let resp = app + .oneshot(make_req( + serde_json::json!({"model":"gpt-4o-resp","input":"hi"}), + )) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::UNPROCESSABLE_ENTITY); + let bytes = to_bytes(resp.into_body(), 65536).await.unwrap(); + let v: serde_json::Value = serde_json::from_slice(&bytes).unwrap(); + assert_eq!(v["error"]["type"], "content_filter"); + assert!( + !v["error"]["message"] + .as_str() + .unwrap_or_default() + .contains("BLOCKME"), + "model refusal leaked in error: {v}" + ); + } + /// #719 companion: a clean non-streaming response with an output /// guardrail configured passes through unchanged → 200 with body. #[tokio::test] @@ -4980,6 +5143,93 @@ mod tests { assert_eq!(v["error"]["type"], "content_filter"); } + /// Streaming refusals use their own delta and aggregate events. The + /// held-back stream must block before any refusal text reaches the wire. + #[tokio::test] + async fn output_guardrail_blocks_streaming_refusal_response_holds_back() { + let upstream = MockServer::start().await; + let sse = "event: response.refusal.delta\n\ + data: {\"type\":\"response.refusal.delta\",\"item_id\":\"msg_refusal\",\"output_index\":0,\"content_index\":0,\"delta\":\"sure: BLOCKME\"}\n\n\ + event: response.refusal.done\n\ + data: {\"type\":\"response.refusal.done\",\"item_id\":\"msg_refusal\",\"output_index\":0,\"content_index\":0,\"refusal\":\"sure: BLOCKME\"}\n\n\ + event: response.completed\n\ + data: {\"type\":\"response.completed\",\"response\":{\"output\":[{\"type\":\"message\",\"content\":[{\"type\":\"refusal\",\"refusal\":\"sure: BLOCKME\"}]}]}}\n\n\ + data: [DONE]\n\n"; + Mock::given(method("POST")) + .and(path("/v1/responses")) + .respond_with( + ResponseTemplate::new(200) + .insert_header("content-type", "text/event-stream") + .set_body_string(sse), + ) + .expect(1) + .mount(&upstream) + .await; + + let snap = new_snap_openai(&upstream.uri()); + snap.models.insert(openai_model("gpt-4o-resp")); + snap.apikeys.insert(apikey_entry(&["*"])); + crate::seed_env_scoped_guardrail(&snap, keyword_output_guardrail("BLOCKME")); + let app = build_app(snap); + + let resp = app + .oneshot(make_req( + serde_json::json!({"model":"gpt-4o-resp","input":"hi","stream":true}), + )) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::UNPROCESSABLE_ENTITY); + let bytes = to_bytes(resp.into_body(), 65536).await.unwrap(); + assert!( + !String::from_utf8_lossy(&bytes).contains("BLOCKME"), + "streamed refusal leaked despite output block", + ); + let v: serde_json::Value = serde_json::from_slice(&bytes).unwrap(); + assert_eq!(v["error"]["type"], "content_filter"); + } + + /// A Responses provider may emit only the final refusal event, without a + /// delta or terminal response snapshot. The typed `/v1/responses` path + /// must still hold and block it before it reaches the client. + #[tokio::test] + async fn output_guardrail_blocks_done_only_streaming_refusal() { + let upstream = MockServer::start().await; + let sse = "event: response.refusal.done\n\ + data: {\"type\":\"response.refusal.done\",\"item_id\":\"msg_refusal\",\"output_index\":0,\"content_index\":0,\"refusal\":\"sure: BLOCKME\"}\n\n\ + data: [DONE]\n\n"; + Mock::given(method("POST")) + .and(path("/v1/responses")) + .respond_with( + ResponseTemplate::new(200) + .insert_header("content-type", "text/event-stream") + .set_body_string(sse), + ) + .expect(1) + .mount(&upstream) + .await; + + let snap = new_snap_openai(&upstream.uri()); + snap.models.insert(openai_model("gpt-4o-resp")); + snap.apikeys.insert(apikey_entry(&["*"])); + crate::seed_env_scoped_guardrail(&snap, keyword_output_guardrail("BLOCKME")); + let app = build_app(snap); + + let resp = app + .oneshot(make_req( + serde_json::json!({"model":"gpt-4o-resp","input":"hi","stream":true}), + )) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::UNPROCESSABLE_ENTITY); + let bytes = to_bytes(resp.into_body(), 65536).await.unwrap(); + assert!( + !String::from_utf8_lossy(&bytes).contains("BLOCKME"), + "done-only refusal leaked despite output block", + ); + let v: serde_json::Value = serde_json::from_slice(&bytes).unwrap(); + assert_eq!(v["error"]["type"], "content_filter"); + } + /// #719 companion: a clean streaming response with an output guardrail /// is scanned then released in full → 200 + the SSE body. #[tokio::test] @@ -6902,6 +7152,50 @@ data: [DONE]\n\n"; assert!(text.len() <= 20, "delta buffer must stay near the cap"); } + #[test] + fn responses_stream_event_text_reads_every_authoritative_output_carrier() { + let cases = [ + ( + serde_json::json!({"type":"response.output_text.done","text":"OUTPUT_DONE"}), + "OUTPUT_DONE", + ), + ( + serde_json::json!({"type":"response.refusal.done","refusal":"REFUSAL_DONE"}), + "REFUSAL_DONE", + ), + ( + serde_json::json!({"type":"response.function_call_arguments.done","name":"lookup","arguments":"FUNCTION_DONE"}), + "FUNCTION_DONE", + ), + ( + serde_json::json!({"type":"response.mcp_call_arguments.done","arguments":"MCP_DONE"}), + "MCP_DONE", + ), + ( + serde_json::json!({"type":"response.custom_tool_call_input.done","input":"CUSTOM_DONE"}), + "CUSTOM_DONE", + ), + ( + serde_json::json!({"type":"response.content_part.added","part":{"type":"refusal","refusal":"PART_ADDED"}}), + "PART_ADDED", + ), + ( + serde_json::json!({"type":"response.output_item.done","item":{"type":"message","content":[{"type":"refusal","refusal":"ITEM_DONE"}]}}), + "ITEM_DONE", + ), + ( + serde_json::json!({"type":"response.completed","response":{"output":[{"type":"message","content":[{"type":"refusal","refusal":"TERMINAL"}]}]}}), + "TERMINAL", + ), + ]; + for (event, expected) in cases { + assert!( + super::responses_stream_event_text(&event).contains(expected), + "{event}", + ); + } + } + // ───────────────────────────────────────────────────────────────── // AISIX-Cloud#1330 / #1024 — `/v1/responses` is the other family the // handler-family rule names by hand: it carries Codex traffic, and a diff --git a/crates/aisix-ratelimit/Cargo.toml b/crates/aisix-ratelimit/Cargo.toml index acea01629..4b336c0fe 100644 --- a/crates/aisix-ratelimit/Cargo.toml +++ b/crates/aisix-ratelimit/Cargo.toml @@ -20,3 +20,8 @@ tracing.workspace = true async-trait.workspace = true redis.workspace = true uuid.workspace = true + +[dev-dependencies] +# Paused Tokio time makes the unavailable-lease fail-open regression +# deterministic without introducing a wall-clock scheduling race in CI. +tokio = { workspace = true, features = ["test-util"] } diff --git a/crates/aisix-ratelimit/src/lib.rs b/crates/aisix-ratelimit/src/lib.rs index c91e52449..513424890 100644 --- a/crates/aisix-ratelimit/src/lib.rs +++ b/crates/aisix-ratelimit/src/lib.rs @@ -26,5 +26,5 @@ pub use limiter::{ }; pub use store::local::LocalStore; pub use store::redis::RedisStore; -pub use store::RateStore; +pub use store::{RateStore, StreamLeaseRefresh}; pub use window::{FixedWindowCounter, WindowCheck}; diff --git a/crates/aisix-ratelimit/src/limiter.rs b/crates/aisix-ratelimit/src/limiter.rs index 41315e52c..8ae0be220 100644 --- a/crates/aisix-ratelimit/src/limiter.rs +++ b/crates/aisix-ratelimit/src/limiter.rs @@ -29,7 +29,7 @@ use aisix_core::RateLimit; use crate::clock::Clock; use crate::error::RateLimitError; use crate::store::local::LocalStore; -use crate::store::RateStore; +use crate::store::{RateStore, StreamLeaseRefresh}; /// Current window state for a single key, returned by [`Limiter::peek`]. /// Used by the proxy handlers to inject the `x-ratelimit-*` response @@ -106,10 +106,25 @@ impl Limiter { ) -> Result { let member = self.next_member(); self.store.acquire(key, limits, &member).await?; + let has_concurrency_slot = limits.concurrency.is_some(); + // Keep this state from acquire onward. A slow upstream-header phase + // can discover a missing Redis member before the response is known + // to be SSE; `into_stream_hold` forwards the stored value later. + let (lease_loss, _) = tokio::sync::watch::channel(false); + let refresh_task = spawn_lease_refresh( + Arc::clone(&self.store), + key.to_string(), + member.clone(), + has_concurrency_slot, + Some(lease_loss.clone()), + ); Ok(Reservation { store: Arc::clone(&self.store), key: key.to_string(), member, + has_concurrency_slot, + refresh_task, + lease_loss, committed: false, }) } @@ -135,6 +150,71 @@ impl Limiter { } } +/// Keep a distributed concurrency member alive from acquisition until the +/// request either commits, drops, or transfers the member to a stream hold. +/// +/// A request can spend longer than the Redis concurrency TTL waiting for +/// upstream response headers. Starting at acquire, rather than only after +/// the SSE handoff, keeps that still-live request from being pruned first. +/// Local stores opt out, and callers outside a Tokio runtime retain the +/// backend's normal stale-lease recovery behavior. +fn spawn_lease_refresh( + store: Arc, + key: String, + member: String, + has_concurrency_slot: bool, + lease_loss: Option>, +) -> Option> { + if !has_concurrency_slot { + return None; + } + let interval = store.stream_lease_refresh_interval()?; + tokio::runtime::Handle::try_current().ok().map(|handle| { + handle.spawn(async move { + loop { + tokio::time::sleep(interval).await; + if matches!( + store.refresh_stream_lease_outcome(&key, &member).await, + StreamLeaseRefresh::Missing + ) { + if let Some(lease_loss) = &lease_loss { + // The response body may not have been polled yet, so + // no receiver may be subscribed at the exact moment + // Redis proves the member disappeared. Preserve this + // one-way state for the later body handoff instead of + // dropping a `watch::send` notification with zero + // receivers. + lease_loss.send_replace(true); + } + return; + } + } + }) + }) +} + +/// Relay each reservation's durable lease-loss state into the one receiver a +/// streamed response selects on. A multi-layer reservation must stop when any +/// of its shared concurrency members disappears. +fn spawn_lease_loss_forwarder( + mut source: tokio::sync::watch::Receiver, + target: tokio::sync::watch::Sender, +) -> Option> { + tokio::runtime::Handle::try_current().ok().map(|handle| { + handle.spawn(async move { + loop { + if *source.borrow_and_update() { + target.send_replace(true); + return; + } + if source.changed().await.is_err() { + return; + } + } + }) + }) +} + impl Default for Limiter { fn default() -> Self { Self::new() @@ -147,6 +227,14 @@ pub struct Reservation { store: Arc, key: String, member: String, + has_concurrency_slot: bool, + /// Starts at successful acquire, rather than at the later SSE handoff: + /// a slow header phase is still an active request that owns this slot. + refresh_task: Option>, + /// Durable across the response-header / streaming-body handoff. A + /// `watch::send` would drop a pre-body notification with no subscriber, + /// while `send_replace` keeps the terminal state for the later hold. + lease_loss: tokio::sync::watch::Sender, committed: bool, } @@ -160,9 +248,16 @@ impl std::fmt::Debug for Reservation { } impl Reservation { + fn stop_refresh(&mut self) { + if let Some(task) = self.refresh_task.take() { + task.abort(); + } + } + /// Post-deduct phase. Records the actual token cost against TPM/TPD /// and releases the concurrency slot. pub async fn commit_tokens(mut self, tokens: u64) { + self.stop_refresh(); self.store.commit(&self.key, tokens, &self.member).await; self.committed = true; } @@ -170,6 +265,7 @@ impl Reservation { impl Drop for Reservation { fn drop(&mut self) { + self.stop_refresh(); if self.committed { return; } @@ -213,8 +309,9 @@ impl MultiReservation { /// Convert into an owned [`StreamConcurrencyGuard`] for the streaming /// path. The per-layer concurrency slots stay held — they are NOT /// released here — and are released only when the returned guard drops, - /// i.e. at stream completion or cancellation. Token accounting still - /// happens via [`Limiter::add_tokens_post_stream`]. + /// i.e. at stream completion or cancellation. Shared stores renew their + /// live leases while the guard exists. Token accounting still happens via + /// [`Limiter::add_tokens_post_stream`]. /// /// A borrow-based reservation couldn't outlive the request handler, so /// the pre-fix streaming path dropped it at handler return; that @@ -223,6 +320,9 @@ impl MultiReservation { #[must_use = "dropping the returned guard immediately releases the concurrency \ slot, recreating the early-release bug this fixes"] pub fn into_stream_hold(mut self) -> StreamConcurrencyGuard { + let mut refresh_tasks = Vec::new(); + let mut lease_loss_tasks = Vec::new(); + let (lease_loss, _) = tokio::sync::watch::channel(false); let holds = self .reservations .iter_mut() @@ -230,11 +330,36 @@ impl MultiReservation { // Defuse each reservation's Drop so it doesn't release the // slot now; the returned guard owns release from here on. r.committed = true; - (Arc::clone(&r.store), r.key.clone(), r.member.clone()) + let store = Arc::clone(&r.store); + // Keep the acquisition-time worker: it may already have + // observed a missing lease while upstream response headers + // were pending. A per-layer watcher forwards that durable + // state into the one response-body receiver below. + if let Some(task) = r.refresh_task.take().or_else(|| { + spawn_lease_refresh( + Arc::clone(&store), + r.key.clone(), + r.member.clone(), + r.has_concurrency_slot, + Some(r.lease_loss.clone()), + ) + }) { + refresh_tasks.push(task); + } + let mut source = r.lease_loss.subscribe(); + if *source.borrow_and_update() { + lease_loss.send_replace(true); + } else if let Some(task) = spawn_lease_loss_forwarder(source, lease_loss.clone()) { + lease_loss_tasks.push(task); + } + (store, r.key.clone(), r.member.clone()) }) .collect(); StreamConcurrencyGuard { holds, + refresh_tasks, + lease_loss_tasks, + lease_loss, released: false, } } @@ -255,15 +380,38 @@ impl std::fmt::Debug for MultiReservation { pub struct StreamConcurrencyGuard { /// `(store, key, member)` per held layer. holds: Vec<(Arc, String, String)>, + /// Renews the lease used by shared stores while this guard owns a live + /// stream. These tasks began at reservation acquisition and are stopped + /// before terminal release to prevent a late renewal from racing teardown. + refresh_tasks: Vec>, + /// Bridges each held layer's pre-handoff watch into `lease_loss`. + lease_loss_tasks: Vec>, + /// One-way notification that a Redis refresh found one of this stream's + /// members missing. Passthrough SSE response bodies use it to stop before + /// another replica can reuse the now-free shared concurrency slot. + lease_loss: tokio::sync::watch::Sender, released: bool, } impl StreamConcurrencyGuard { + /// Subscribe from a passthrough SSE response body before moving this hold + /// into it. Local stores keep the channel at `false`, so the receiver + /// remains pending without changing their historical behavior. + pub fn lease_loss_receiver(&self) -> tokio::sync::watch::Receiver { + self.lease_loss.subscribe() + } + fn release_now(&mut self) { if self.released { return; } self.released = true; + for task in self.refresh_tasks.drain(..) { + task.abort(); + } + for task in self.lease_loss_tasks.drain(..) { + task.abort(); + } for (store, key, member) in &self.holds { store.release(key, member); } @@ -274,6 +422,8 @@ impl std::fmt::Debug for StreamConcurrencyGuard { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("StreamConcurrencyGuard") .field("layers", &self.holds.len()) + .field("refreshing", &!self.refresh_tasks.is_empty()) + .field("lease_loss_watchers", &self.lease_loss_tasks.len()) .field("released", &self.released) .finish() } @@ -289,6 +439,66 @@ impl Drop for StreamConcurrencyGuard { mod tests { use super::*; use crate::clock::TestClock; + use async_trait::async_trait; + use std::sync::atomic::{AtomicU64, Ordering}; + use std::time::Duration; + + struct LeaseRefreshStore { + inner: LocalStore, + refreshes: AtomicU64, + outcome: StreamLeaseRefresh, + } + + impl LeaseRefreshStore { + fn new(outcome: StreamLeaseRefresh) -> Self { + Self { + inner: LocalStore::with_clock(TestClock::new(100)), + refreshes: AtomicU64::new(0), + outcome, + } + } + } + + #[async_trait] + impl RateStore for LeaseRefreshStore { + async fn acquire( + &self, + key: &str, + limits: &RateLimit, + member: &str, + ) -> Result<(), RateLimitError> { + self.inner.acquire(key, limits, member).await + } + + async fn commit(&self, key: &str, tokens: u64, member: &str) { + self.inner.commit(key, tokens, member).await; + } + + fn release(&self, key: &str, member: &str) { + self.inner.release(key, member); + } + + fn add_tokens(&self, key: &str, tokens: u64) { + self.inner.add_tokens(key, tokens); + } + + fn stream_lease_refresh_interval(&self) -> Option { + Some(Duration::from_millis(5)) + } + + async fn refresh_stream_lease_outcome( + &self, + _key: &str, + _member: &str, + ) -> StreamLeaseRefresh { + self.refreshes.fetch_add(1, Ordering::Relaxed); + self.outcome + } + + async fn peek(&self, key: &str, limits: &RateLimit) -> Option { + self.inner.peek(key, limits).await + } + } fn limits(rpm: Option, tpm: Option, concurrency: Option) -> RateLimit { RateLimit { @@ -801,6 +1011,62 @@ mod tests { assert!(limiter.pre_commit("k", &l).await.is_ok()); } + #[tokio::test(start_paused = true)] + async fn unavailable_stream_lease_refresh_keeps_the_stream_hold_open() { + let store = Arc::new(LeaseRefreshStore::new(StreamLeaseRefresh::Unavailable)); + let limiter = Limiter::with_store(store.clone()); + let limits = limits(None, None, Some(1)); + let hold = MultiReservation::new(vec![limiter.pre_commit("k", &limits).await.unwrap()]) + .into_stream_hold(); + let mut lost = hold.lease_loss_receiver(); + + // Let the replacement stream worker register its first timer before + // advancing paused time. Two unavailable results prove it keeps + // retrying, rather than treating an outage as a confirmed loss. + tokio::task::yield_now().await; + for expected in 1..=2 { + tokio::time::advance(Duration::from_millis(5)).await; + tokio::task::yield_now().await; + assert_eq!(store.refreshes.load(Ordering::Relaxed), expected); + } + assert!( + !*lost.borrow_and_update(), + "an unavailable backend is fail-open, not a definitive lost lease" + ); + assert!( + !lost + .has_changed() + .expect("the hold retains the watch sender"), + "unavailable refreshes must not terminate a live stream" + ); + + drop(hold); + } + + #[tokio::test(start_paused = true)] + async fn missing_lease_before_stream_handoff_is_preserved() { + let store = Arc::new(LeaseRefreshStore::new(StreamLeaseRefresh::Missing)); + let limiter = Limiter::with_store(store.clone()); + let limits = limits(None, None, Some(1)); + let reservation = limiter.pre_commit("k", &limits).await.unwrap(); + + // The worker observes Missing before anyone knows the upstream is + // streaming. The later handoff must see that stored state without + // waiting for another refresh interval. + tokio::task::yield_now().await; + tokio::time::advance(Duration::from_millis(5)).await; + tokio::task::yield_now().await; + assert_eq!(store.refreshes.load(Ordering::Relaxed), 1); + + let hold = MultiReservation::new(vec![reservation]).into_stream_hold(); + let mut lost = hold.lease_loss_receiver(); + assert!( + *lost.borrow_and_update(), + "the response-body handoff must retain a pre-header Missing result" + ); + drop(hold); + } + #[tokio::test] async fn multi_reservation_keys_returns_all_held_keys() { let clock = TestClock::new(100); diff --git a/crates/aisix-ratelimit/src/store/mod.rs b/crates/aisix-ratelimit/src/store/mod.rs index 841059cd8..44c20dcf8 100644 --- a/crates/aisix-ratelimit/src/store/mod.rs +++ b/crates/aisix-ratelimit/src/store/mod.rs @@ -20,6 +20,8 @@ //! sync because they run from `Drop` and from the synchronous SSE //! completion callback; the Redis impl makes them fire-and-forget. +use std::time::Duration; + use aisix_core::RateLimit; use async_trait::async_trait; @@ -45,6 +47,21 @@ pub(crate) const DIM_RPD: &str = "rpd"; pub(crate) const DIM_TPM: &str = "tpm"; pub(crate) const DIM_TPD: &str = "tpd"; +/// Result of trying to renew a distributed streaming concurrency lease. +/// +/// Only [`StreamLeaseRefresh::Missing`] is definitive: the member was no +/// longer present in the shared semaphore. The passthrough SSE relay uses +/// that signal to end an already-headed response rather than continue outside +/// its configured concurrency limit. `Unavailable` includes backends that +/// cannot report an outcome and Redis transport failures; those retain the +/// rate limiter's established fail-open behavior. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum StreamLeaseRefresh { + Renewed, + Missing, + Unavailable, +} + /// A windowed request/token dimension active on a [`RateLimit`]: /// `(name, window_secs, limit)`. Shared by both stores so the Redis key /// layout and the local counter set never drift. @@ -127,6 +144,26 @@ pub trait RateStore: Send + Sync + 'static { /// completion callback; the Redis impl makes it fire-and-forget. fn add_tokens(&self, key: &str, tokens: u64); + /// How often a live streaming reservation should renew its distributed + /// concurrency lease. Local counters have no lease, so they opt out. + fn stream_lease_refresh_interval(&self) -> Option { + None + } + + /// Refresh a live streaming reservation's distributed concurrency lease. + /// Implementations must never recreate a member that has already been + /// released: a final refresh racing with stream teardown must be a no-op. + async fn refresh_stream_lease(&self, _key: &str, _member: &str) {} + + /// Refresh a streaming lease while reporting whether its member still + /// exists. This preserves the legacy [`RateStore::refresh_stream_lease`] + /// hook for external stores: an implementation that only provides that + /// hook remains fail-open because it cannot prove that its member vanished. + async fn refresh_stream_lease_outcome(&self, key: &str, member: &str) -> StreamLeaseRefresh { + self.refresh_stream_lease(key, member).await; + StreamLeaseRefresh::Unavailable + } + /// Read-only snapshot for the `x-ratelimit-*` headers. Returns `None` /// when there is nothing meaningful to report for the bucket. async fn peek(&self, key: &str, limits: &RateLimit) -> Option; diff --git a/crates/aisix-ratelimit/src/store/redis.rs b/crates/aisix-ratelimit/src/store/redis.rs index a866653a7..bc89acd47 100644 --- a/crates/aisix-ratelimit/src/store/redis.rs +++ b/crates/aisix-ratelimit/src/store/redis.rs @@ -11,13 +11,13 @@ //! Cluster slot, keeping the per-bucket Lua atomic): //! - `aisix:rl:{}::` — plain //! `INCR`/`GET` counters, `EXPIRE = window + grace`. -//! - `aisix:rl:{}:conc` — a ZSET (`member → score=now`) acting as a -//! crash-safe distributed semaphore: acquire prunes entries older than -//! `conc_ttl` then counts, so a slot leaked by a crashed/hung replica is -//! reclaimed within `conc_ttl`. (LiteLLM's latest tracks parallel -//! requests as a window-TTL counter; we use a ZSET with a request- -//! lifetime ttl because our streaming requests can outlive a 60s window -//! — the same reason `StreamConcurrencyGuard`/#450 exists.) +//! - `aisix:rl:{}:conc` — a ZSET (`member → last-refresh time`) acting +//! as a crash-safe distributed semaphore: acquire prunes stale entries then +//! counts, so a slot leaked by a crashed/hung replica is reclaimed after its +//! `conc_ttl` (with at most one extra second for rolling-upgrade safety). +//! Live streams renew their member while their `StreamConcurrencyGuard` +//! exists, so a deliberately long SSE response is not mistaken for a +//! crashed replica. //! //! `now` is read from `redis.call('TIME')` inside every script so window //! boundaries are identical across replicas regardless of host clock skew. @@ -43,9 +43,10 @@ use aisix_core::{RateLimit, RedisConnConfig}; use aisix_obs::metrics::Metrics; use aisix_redis::{ConnSlot, FailurePolicy}; use async_trait::async_trait; +use dashmap::DashMap; use redis::Script; -use super::{local::LocalStore, token_dims, Dim, RateStore}; +use super::{local::LocalStore, token_dims, Dim, RateStore, StreamLeaseRefresh}; use crate::error::{LimitDetail, RateLimitError}; use crate::limiter::RateLimitStatus; @@ -77,6 +78,17 @@ const CODE_CONCURRENCY: i64 = 1; const CODE_TOKENS: i64 = 2; const CODE_REQUESTS: i64 = 3; +/// What this process knows about one live concurrency member. A Redis command +/// can fail after the server applies its Lua mutation but before the client +/// receives a reply, so "not confirmed" is not always the same as "local +/// only". +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum LeaseOwnership { + ConfirmedRedis, + LocalOnly, + Unknown, +} + /// Atomic per-bucket acquire: concurrency gate + token check-only + /// request check-and-increment, all-or-nothing. See module docs for the /// key layout. Returns `{code, retry_after}`. @@ -85,20 +97,29 @@ local prefix = ARGV[1] local member = ARGV[2] local conc_max = tonumber(ARGV[3]) local conc_ttl = tonumber(ARGV[4]) -local grace = tonumber(ARGV[5]) +local conc_key_ttl = tonumber(ARGV[5]) +local grace = tonumber(ARGV[6]) local t = redis.call('TIME') local now = tonumber(t[1]) +-- Keep the score in seconds so rolling upgrades still understand existing +-- integer-second members, but retain Redis TIME's sub-second precision for +-- a live lease. +local conc_now = now + tonumber(t[2]) / 1000000 +-- Old nodes wrote integer-second scores. Prune at a whole-second boundary so +-- a new node cannot mistake an old member acquired late in that second for a +-- stale member; the safe cost is at most one extra second of retention. +local conc_prune_before = math.floor(conc_now) - conc_ttl local conc_key = prefix .. ':conc' if conc_max >= 0 then - redis.call('ZREMRANGEBYSCORE', conc_key, 0, now - conc_ttl) + redis.call('ZREMRANGEBYSCORE', conc_key, 0, '(' .. conc_prune_before) local in_flight = redis.call('ZCARD', conc_key) if in_flight >= conc_max then return {1, 0, 0, conc_max, in_flight} end end -local idx = 6 +local idx = 7 local nreq = tonumber(ARGV[idx]); idx = idx + 1 local req = {} for i = 1, nreq do @@ -143,12 +164,30 @@ for i = 1, nreq do end end if conc_max >= 0 then - redis.call('ZADD', conc_key, now, member) - redis.call('EXPIRE', conc_key, conc_ttl) + redis.call('ZADD', conc_key, conc_now, member) + redis.call('EXPIRE', conc_key, conc_key_ttl) end return {0, 0, 0, 0, 0} "#; +/// Refresh one live streaming member's concurrency lease. The member must +/// already exist: refresh racing stream teardown must not resurrect a member +/// after `release` removed it. ARGV: prefix, member, conc_key_ttl. +const REFRESH_CONCURRENCY_LUA: &str = r#" +local prefix = ARGV[1] +local member = ARGV[2] +local conc_key_ttl = tonumber(ARGV[3]) +local conc_key = prefix .. ':conc' +if not redis.call('ZSCORE', conc_key, member) then + return 0 +end +local t = redis.call('TIME') +local conc_now = tonumber(t[1]) + tonumber(t[2]) / 1000000 +redis.call('ZADD', conc_key, 'XX', conc_now, member) +redis.call('EXPIRE', conc_key, conc_key_ttl) +return 1 +"#; + /// Post-deduct: add `tokens` to the tpm/tpd windows AND release the /// concurrency slot held by `member`. Both token windows are always /// touched (matching the local backend); an unread tpd counter just @@ -201,10 +240,14 @@ local prefix = ARGV[1] local conc_ttl = tonumber(ARGV[2]) local t = redis.call('TIME') local now = tonumber(t[1]) +local conc_now = now + tonumber(t[2]) / 1000000 +-- See ACQUIRE_LUA: use the same conservative boundary so header reads do not +-- prune a live member while a rolling upgrade is in progress. +local conc_prune_before = math.floor(conc_now) - conc_ttl local ws = now - (now % 60) local rpm = tonumber(redis.call('GET', prefix .. ':rpm:' .. ws) or '0') local tpm = tonumber(redis.call('GET', prefix .. ':tpm:' .. ws) or '0') -redis.call('ZREMRANGEBYSCORE', prefix .. ':conc', 0, now - conc_ttl) +redis.call('ZREMRANGEBYSCORE', prefix .. ':conc', 0, '(' .. conc_prune_before) local inflight = redis.call('ZCARD', prefix .. ':conc') return {rpm, tpm, inflight, 60 - (now % 60)} "#; @@ -216,6 +259,10 @@ pub struct RedisStore { grace: u64, /// Per-process fallback used when Redis is unreachable (fail-open). local: Arc, + /// Per-reservation ownership provenance for stream lease refreshes. + /// Only a `ConfirmedRedis` member may turn a healthy ZSET miss into a + /// terminal signal; a timeout after dispatch is deliberately `Unknown`. + lease_members: DashMap<(String, String), LeaseOwnership>, /// One-shot guard so the degradation warning is logged once, not per /// request, while Redis stays down. Shared with the boot attach task, /// which starts it already set (the boot WARN said it) and clears it @@ -309,6 +356,7 @@ impl RedisStore { conc_ttl: DEFAULT_CONC_TTL_SECS, grace: DEFAULT_GRACE_SECS, local: Arc::new(LocalStore::new()), + lease_members: DashMap::new(), degraded_logged, metrics: None, } @@ -502,6 +550,9 @@ impl RateStore for RedisStore { member.to_string(), limits.concurrency.map(i64::from).unwrap_or(-1).to_string(), self.conc_ttl.to_string(), + // Pruning retains a pre-upgrade integer-second member through its + // boundary, so the key must survive for the same extra second. + self.conc_ttl.saturating_add(1).to_string(), self.grace.to_string(), ]; push_dims(&mut args, &super::request_dims(limits)); @@ -519,7 +570,14 @@ impl RateStore for RedisStore { Ok(c) => c, Err(e) => { self.note_failure("acquire", &e); - return self.local.acquire(key, limits, member).await; + let acquired = self.local.acquire(key, limits, member).await; + if acquired.is_ok() && limits.concurrency.is_some() { + self.lease_members.insert( + (key.to_string(), member.to_string()), + LeaseOwnership::LocalOnly, + ); + } + return acquired; } }; match invocation.invoke_async::>(&mut conn).await { @@ -544,7 +602,15 @@ impl RateStore for RedisStore { LimitDetail::window(name, limit, used, retry) }; match code { - CODE_OK => Ok(()), + CODE_OK => { + if limits.concurrency.is_some() { + self.lease_members.insert( + (key.to_string(), member.to_string()), + LeaseOwnership::ConfirmedRedis, + ); + } + Ok(()) + } CODE_CONCURRENCY => { let limit = reply.get(3).copied().unwrap_or(0).max(0) as u64; let in_flight = reply.get(4).copied().unwrap_or(0).max(0) as u64; @@ -566,12 +632,24 @@ impl RateStore for RedisStore { Err(e) => { self.note_failure("acquire", &e); self.conn.note_error().await; - self.local.acquire(key, limits, member).await + let acquired = self.local.acquire(key, limits, member).await; + if acquired.is_ok() && limits.concurrency.is_some() { + // A reply can be lost after Redis has already applied the + // Lua acquire. Probe on refresh rather than assuming the + // local fallback was the only reservation. + self.lease_members.insert( + (key.to_string(), member.to_string()), + LeaseOwnership::Unknown, + ); + } + acquired } } } async fn commit(&self, key: &str, tokens: u64, member: &str) { + self.lease_members + .remove(&(key.to_string(), member.to_string())); let prefix = self.bucket_prefix(key); let mut conn = match self.conn.acquire().await { Ok(c) => c, @@ -624,6 +702,8 @@ impl RateStore for RedisStore { } fn release(&self, key: &str, member: &str) { + self.lease_members + .remove(&(key.to_string(), member.to_string())); // Drop the local slot first (a cheap no-op when the bucket was // never acquired locally); covers the degraded-acquire case. self.local.release(key, member); @@ -692,6 +772,87 @@ impl RateStore for RedisStore { } } + fn stream_lease_refresh_interval(&self) -> Option { + let millis = self + .conc_ttl + .saturating_mul(1_000) + .saturating_div(3) + .clamp(100, 60_000); + Some(std::time::Duration::from_millis(millis)) + } + + async fn refresh_stream_lease(&self, key: &str, member: &str) { + let _ = self.refresh_stream_lease_outcome(key, member).await; + } + + async fn refresh_stream_lease_outcome(&self, key: &str, member: &str) -> StreamLeaseRefresh { + let ownership = self + .lease_members + .get(&(key.to_string(), member.to_string())) + .map(|entry| *entry); + let Some(ownership) = ownership else { + // The member was already released, or this store never admitted + // it. It is not a live shared lease that can be conclusively lost. + return StreamLeaseRefresh::Unavailable; + }; + if matches!(ownership, LeaseOwnership::LocalOnly) { + // `conn.acquire` failed before a command could be dispatched. + // This member exists solely in the local fail-open fallback. + return StreamLeaseRefresh::Unavailable; + } + let prefix = self.bucket_prefix(key); + let mut conn = match self.conn.acquire().await { + Ok(c) => c, + Err(e) => { + self.note_failure("refresh", &e); + // A post-dispatch I/O failure opens the shared Redis breaker. + // While it short-circuits, Unknown stays fail-open; only a + // later successful refresh can promote it to confirmed Redis + // ownership. + return StreamLeaseRefresh::Unavailable; + } + }; + let res: Result = Script::new(REFRESH_CONCURRENCY_LUA) + .key(&prefix) + .arg(&prefix) + .arg(member) + .arg(self.conc_ttl.saturating_add(1)) + .invoke_async(&mut conn) + .await; + match res { + Ok(1) => { + self.mark_ok(); + if matches!(ownership, LeaseOwnership::Unknown) { + self.lease_members.insert( + (key.to_string(), member.to_string()), + LeaseOwnership::ConfirmedRedis, + ); + } + StreamLeaseRefresh::Renewed + } + Ok(0) => { + // Redis did answer: the member itself is gone. This is not a + // backend outage, so clear any prior degradation marker while + // returning the definitive loss to the stream hold. + self.mark_ok(); + if matches!(ownership, LeaseOwnership::ConfirmedRedis) { + StreamLeaseRefresh::Missing + } else { + // A command may have failed before Redis ever saw it, so + // an Unknown member's absence is fail-open, not proof of + // an externally released live lease. + StreamLeaseRefresh::Unavailable + } + } + Ok(_) => StreamLeaseRefresh::Unavailable, + Err(e) => { + self.note_failure("refresh", &e); + self.conn.note_error().await; + StreamLeaseRefresh::Unavailable + } + } + } + async fn peek(&self, key: &str, limits: &RateLimit) -> Option { if limits.is_unrestricted() { return None; diff --git a/crates/aisix-ratelimit/tests/redis_integration.rs b/crates/aisix-ratelimit/tests/redis_integration.rs index 4a0d7dd65..75acb82bc 100644 --- a/crates/aisix-ratelimit/tests/redis_integration.rs +++ b/crates/aisix-ratelimit/tests/redis_integration.rs @@ -6,11 +6,15 @@ //! replicas pointed at one Redis — the exact api7/AISIX-Cloud#798 shape: //! a limit hit on one replica must already be hit on the other. +use std::sync::Arc; use std::time::Duration; use aisix_core::{RateLimit, RateLimitScope, RedisConnConfig, RedisMode}; use aisix_obs::metrics::Metrics; -use aisix_ratelimit::{RateStore, RedisStore}; +use aisix_ratelimit::{ + store::redis::DEFAULT_PREFIX, Limiter, MultiReservation, RateStore, RedisStore, + StreamLeaseRefresh, +}; fn redis_url() -> Option { std::env::var("RATELIMIT_TEST_REDIS_URL").ok() @@ -372,7 +376,8 @@ async fn stale_concurrency_slot_is_reclaimed_after_ttl() { eprintln!("skipping: RATELIMIT_TEST_REDIS_URL not set"); return; }; - // 1s slot lifetime: a never-released slot (crashed replica) is pruned. + // A 1s slot lifetime plus the one-second rolling-upgrade compatibility + // margin: a never-released slot (crashed replica) is eventually pruned. let a = store(&url).await.with_conc_ttl(1); let b = store(&url).await.with_conc_ttl(1); let key = unique_key("conc-ttl"); @@ -390,12 +395,260 @@ async fn stale_concurrency_slot_is_reclaimed_after_ttl() { "slot held while fresh" ); - tokio::time::sleep(Duration::from_millis(1_300)).await; + tokio::time::sleep(Duration::from_millis(2_200)).await; b.acquire(&key, &limits, "b-2") .await .expect("stale slot reclaimed after conc_ttl"); } +#[tokio::test] +async fn refreshing_released_member_does_not_recreate_slot() { + let Some(url) = redis_url() else { + eprintln!("skipping: RATELIMIT_TEST_REDIS_URL not set"); + return; + }; + let a = store(&url).await.with_conc_ttl(1); + let b = store(&url).await.with_conc_ttl(1); + let key = unique_key("conc-refresh-release"); + let limits = RateLimit { + concurrency: Some(1), + ..rl() + }; + + a.acquire(&key, &limits, "a-stream") + .await + .expect("first stream allowed"); + // `commit` removes the member synchronously in Redis and clears the + // store's ownership record. A queued refresh must neither recreate it + // nor convert this normal teardown into a terminal stream signal. + a.commit(&key, 0, "a-stream").await; + assert_eq!( + a.refresh_stream_lease_outcome(&key, "a-stream").await, + StreamLeaseRefresh::Unavailable, + "a released member is no longer a live shared lease" + ); + + b.acquire(&key, &limits, "b-stream") + .await + .expect("refresh after release must not recreate the slot"); +} + +#[tokio::test] +async fn legacy_integer_member_survives_the_upgrade_prune_boundary() { + let Some(url) = redis_url() else { + eprintln!("skipping: RATELIMIT_TEST_REDIS_URL not set"); + return; + }; + let b = store(&url).await.with_conc_ttl(1); + let key = unique_key("conc-legacy-upgrade"); + let limits = RateLimit { + concurrency: Some(1), + ..rl() + }; + let client = redis::Client::open(url.as_str()).expect("raw Redis client"); + let mut raw = client + .get_multiplexed_async_connection() + .await + .expect("raw Redis connection"); + + // A pre-renewal gateway wrote one integer-second score. Install a score + // that is exactly on the new build's stale boundary while still early in + // the current server second. An inclusive or fractional prune erases it + // too early; the conservative exclusive boundary keeps it. + let legacy_second = tokio::time::timeout(Duration::from_secs(2), async { + loop { + let time: Vec = redis::cmd("TIME") + .query_async(&mut raw) + .await + .expect("Redis TIME"); + let seconds: u64 = time[0].parse().expect("TIME seconds"); + let micros: u64 = time[1].parse().expect("TIME microseconds"); + if micros < 500_000 { + return seconds; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("reach an early Redis second"); + let conc_key = format!("{DEFAULT_PREFIX}:{{{key}}}:conc"); + let _: i64 = redis::cmd("ZADD") + .arg(&conc_key) + .arg(legacy_second.saturating_sub(1)) + .arg("legacy-stream") + .query_async(&mut raw) + .await + .expect("install legacy member"); + let _: i64 = redis::cmd("EXPIRE") + .arg(&conc_key) + .arg(5) + .query_async(&mut raw) + .await + .expect("keep legacy member for the boundary check"); + + assert!( + matches!( + b.acquire(&key, &limits, "new-stream").await, + Err(aisix_ratelimit::RateLimitError::Concurrency { .. }) + ), + "new code must not prune a live legacy member at the upgrade boundary" + ); +} + +#[tokio::test] +async fn stream_hold_renews_redis_lease_until_drop() { + let Some(url) = redis_url() else { + eprintln!("skipping: RATELIMIT_TEST_REDIS_URL not set"); + return; + }; + // Keep this short to prove a live stream is not mistaken for a crashed + // replica. The stores model two DPs sharing the same Redis backend. + let a = Limiter::with_store(Arc::new(store(&url).await.with_conc_ttl(1))); + let b = Limiter::with_store(Arc::new(store(&url).await.with_conc_ttl(1))); + let key = unique_key("conc-stream-lease"); + let limits = RateLimit { + concurrency: Some(1), + ..rl() + }; + + let hold = + MultiReservation::new(vec![a.pre_commit(&key, &limits).await.unwrap()]).into_stream_hold(); + + // This is longer than conc_ttl plus the one-second rolling-upgrade + // compatibility slack. Without the stream lease heartbeat, B's acquire + // prunes A and incorrectly admits a second live stream. + tokio::time::sleep(Duration::from_millis(2_200)).await; + assert!( + matches!( + b.pre_commit(&key, &limits).await, + Err(aisix_ratelimit::RateLimitError::Concurrency { .. }) + ), + "a live stream must keep its shared concurrency slot beyond conc_ttl" + ); + + drop(hold); + // The release is detached, so bound the wait instead of assuming a + // propagation delay on a slow CI runner. + let mut acquired = false; + for _ in 0..50 { + if b.pre_commit(&key, &limits).await.is_ok() { + acquired = true; + break; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + assert!(acquired, "slot must free cluster-wide when the stream ends"); +} + +#[tokio::test] +async fn stream_hold_reports_a_definitively_missing_redis_lease() { + let Some(url) = redis_url() else { + eprintln!("skipping: RATELIMIT_TEST_REDIS_URL not set"); + return; + }; + let a = Limiter::with_store(Arc::new(store(&url).await.with_conc_ttl(1))); + let key = unique_key("conc-stream-lease-loss"); + let limits = RateLimit { + concurrency: Some(1), + ..rl() + }; + let hold = + MultiReservation::new(vec![a.pre_commit(&key, &limits).await.unwrap()]).into_stream_hold(); + + let client = redis::Client::open(url.as_str()).expect("raw Redis client"); + let mut raw = client + .get_multiplexed_async_connection() + .await + .expect("raw Redis connection"); + let conc_key = format!("{DEFAULT_PREFIX}:{{{key}}}:conc"); + let members: Vec = redis::cmd("ZRANGE") + .arg(&conc_key) + .arg(0) + .arg(-1) + .query_async(&mut raw) + .await + .expect("read live stream member"); + assert_eq!(members.len(), 1, "one stream hold owns one Redis member"); + let removed: i64 = redis::cmd("ZREM") + .arg(&conc_key) + .arg(&members[0]) + .query_async(&mut raw) + .await + .expect("remove live stream member"); + assert_eq!(removed, 1, "remove the member the hold is renewing"); + + // The response body normally subscribes only after headers are sent. Let + // the heartbeat observe `Missing` before this test subscribes, then prove + // the one-way state survives that handoff with no receiver present. + tokio::time::sleep(Duration::from_secs(1)).await; + let mut lost = hold.lease_loss_receiver(); + tokio::time::timeout(Duration::from_secs(2), async { + loop { + if *lost.borrow_and_update() { + return; + } + lost.changed() + .await + .expect("stream hold retains the lease-loss sender"); + } + }) + .await + .expect("next Redis renewal must report the missing member"); + + drop(hold); +} + +/// The response-header phase can outlive Redis's crash-recovery TTL before +/// the gateway knows that the response is an SSE stream. That live request +/// must keep its slot through the later handoff to `StreamConcurrencyGuard`. +#[tokio::test] +async fn reservation_renews_redis_lease_before_stream_handoff() { + let Some(url) = redis_url() else { + eprintln!("skipping: RATELIMIT_TEST_REDIS_URL not set"); + return; + }; + let a = Limiter::with_store(Arc::new(store(&url).await.with_conc_ttl(1))); + let b = Limiter::with_store(Arc::new(store(&url).await.with_conc_ttl(1))); + let key = unique_key("conc-pre-stream-lease"); + let limits = RateLimit { + concurrency: Some(1), + ..rl() + }; + + // Simulate a slow upstream header phase. Before this fix, the member + // expired here because renewal did not start until `into_stream_hold`. + let reservation = MultiReservation::new(vec![a.pre_commit(&key, &limits).await.unwrap()]); + tokio::time::sleep(Duration::from_millis(2_200)).await; + assert!( + matches!( + b.pre_commit(&key, &limits).await, + Err(aisix_ratelimit::RateLimitError::Concurrency { .. }) + ), + "an active pre-handoff request must keep its shared concurrency slot" + ); + + let hold = reservation.into_stream_hold(); + tokio::time::sleep(Duration::from_millis(2_200)).await; + assert!( + matches!( + b.pre_commit(&key, &limits).await, + Err(aisix_ratelimit::RateLimitError::Concurrency { .. }) + ), + "the same lease must remain held after the stream handoff" + ); + + drop(hold); + let mut acquired = false; + for _ in 0..50 { + if b.pre_commit(&key, &limits).await.is_ok() { + acquired = true; + break; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + assert!(acquired, "slot must free cluster-wide when the stream ends"); +} + /// Redis Cluster: the multi-key acquire/commit Lua must route to the slot /// owning the `{bucket}` hash tag and enforce one shared window. A wrong /// (or missing) routing key would surface as a CROSSSLOT/MOVED error here. @@ -540,10 +793,16 @@ async fn env_namespace_isolates_model_alias_bucket() { /// container, a downed host or a partitioned network looks like from /// the client end. Nothing arrives, nothing is refused, and without a /// command budget the caller waits on TCP retransmission for minutes. +/// - [`fail_next_reply_after_apply`](RedisCutoff::fail_next_reply_after_apply) +/// — forward one command to Redis, replace its successful reply with a +/// local Redis error, then keep the TCP session usable. This models the +/// post-dispatch ambiguity without opening the client's connectivity +/// breaker, so the next refresh can prove ownership deterministically. struct RedisCutoff { port: u16, cut: std::sync::Arc, hole: std::sync::Arc, + fail_reply_after_apply: std::sync::Arc, } impl RedisCutoff { @@ -557,8 +816,10 @@ impl RedisCutoff { let port = listener.local_addr().unwrap().port(); let cut = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)); let hole = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)); + let fail_reply_after_apply = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)); let flag = cut.clone(); let hole_flag = hole.clone(); + let fail_reply_after_apply_flag = fail_reply_after_apply.clone(); tokio::spawn(async move { loop { let Ok((mut client, _)) = listener.accept().await else { @@ -569,6 +830,7 @@ impl RedisCutoff { }; let flag = flag.clone(); let hole_flag = hole_flag.clone(); + let fail_reply_after_apply_flag = fail_reply_after_apply_flag.clone(); tokio::spawn(async move { let mut from_client = [0u8; 8192]; let mut from_server = [0u8; 8192]; @@ -605,6 +867,24 @@ impl RedisCutoff { if hole_flag.load(std::sync::atomic::Ordering::Relaxed) { continue; } + if fail_reply_after_apply_flag + .swap(false, std::sync::atomic::Ordering::Relaxed) + { + // Redis applied this command, but the + // client receives an error instead of + // its success reply. Keeping the socket + // alive avoids tripping the 30s network + // breaker and leaves a deterministic + // post-dispatch ambiguity to resolve. + if client + .write_all(b"-ERR simulated lost redis reply\r\n") + .await + .is_err() + { + return; + } + continue; + } if client.write_all(&from_server[..n]).await.is_err() { return; } @@ -614,7 +894,12 @@ impl RedisCutoff { }); } }); - Self { port, cut, hole } + Self { + port, + cut, + hole, + fail_reply_after_apply, + } } fn url(&self) -> String { @@ -629,11 +914,17 @@ impl RedisCutoff { self.hole.store(true, std::sync::atomic::Ordering::Relaxed); } + fn fail_next_reply_after_apply(&self) { + self.fail_reply_after_apply + .store(true, std::sync::atomic::Ordering::Relaxed); + } + /// Forward again. Connections opened from now on work; the ones the /// blackhole swallowed mid-handshake were abandoned by the client /// when its own budget expired, which is what a Redis coming back up /// looks like from the caller's end. fn heal(&self) { + self.cut.store(false, std::sync::atomic::Ordering::Relaxed); self.hole.store(false, std::sync::atomic::Ordering::Relaxed); } } @@ -702,6 +993,145 @@ async fn a_redis_outage_counts_every_failed_operation() { ); } +/// A stream admitted by the local fail-open fallback never had a Redis ZSET +/// member. When Redis comes back, that absence must remain `Unavailable`, not +/// become a terminal `Missing` signal that truncates the still-live stream. +#[tokio::test] +async fn a_fallback_stream_lease_stays_open_after_redis_recovers() { + let Some(url) = redis_url() else { + eprintln!("skipping: RATELIMIT_TEST_REDIS_URL not set"); + return; + }; + let relay = RedisCutoff::start(&url).await; + let store = Arc::new( + RedisStore::connect(&single(&relay.url())) + .await + .expect("connect through relay") + .with_conc_ttl(1), + ); + let limiter = Limiter::with_store(store.clone()); + let key = unique_key("fallback-stream-recovery"); + let fallback_key = unique_key("fallback-stream-recovery-direct"); + let fallback_member = "fallback-stream-member"; + let limits = RateLimit { + concurrency: Some(1), + ..rl() + }; + + relay.cut(); + store + .acquire(&fallback_key, &limits, fallback_member) + .await + .expect("the local fallback admits while Redis is down"); + let hold = MultiReservation::new(vec![limiter.pre_commit(&key, &limits).await.unwrap()]) + .into_stream_hold(); + let mut lost = hold.lease_loss_receiver(); + + // A fresh connection through the healed relay proves Redis is serving + // again. The first stream still has only its local fallback reservation. + relay.heal(); + let peer = RedisStore::connect(&single(&relay.url())) + .await + .expect("Redis recovers through the relay"); + let probe_key = unique_key("fallback-stream-recovery-probe"); + peer.acquire(&probe_key, &limits, "recovered-peer") + .await + .expect("the recovered Redis accepts a shared reservation"); + peer.release(&probe_key, "recovered-peer"); + + assert_eq!( + store + .refresh_stream_lease_outcome(&fallback_key, fallback_member) + .await, + StreamLeaseRefresh::Unavailable, + "a local-only member is never a definitive Redis lease loss" + ); + tokio::time::sleep(Duration::from_secs(1)).await; + assert!( + !*lost.borrow_and_update(), + "Redis recovery must not truncate a fallback-admitted stream" + ); + assert!( + !lost + .has_changed() + .expect("the stream hold retains its lease-loss sender"), + "fallback refreshes remain fail-open after recovery" + ); + + drop(hold); + store.release(&fallback_key, fallback_member); +} + +/// A lost Redis reply leaves acquire provenance ambiguous: the server may +/// have added the member even though the client fell back locally. Refresh +/// must probe and promote that member when Redis confirms it, not let it age +/// out as if the fallback were certainly local-only. +#[tokio::test] +async fn an_ambiguous_acquire_promotes_when_redis_confirms_the_member() { + let Some(url) = redis_url() else { + eprintln!("skipping: RATELIMIT_TEST_REDIS_URL not set"); + return; + }; + let relay = RedisCutoff::start(&url).await; + let store = RedisStore::connect(&single(&relay.url())) + .await + .expect("connect through relay") + .with_conc_ttl(1); + let limits = RateLimit { + concurrency: Some(1), + ..rl() + }; + + // `Script::invoke_async` can load this Lua script with a NOSCRIPT retry + // before its first real execution. Warm it before replacing a reply so + // the injected error follows a Redis-applied ZADD on every clean CI DB. + let warmup_key = unique_key("ambiguous-stream-acquire-warmup"); + store + .acquire(&warmup_key, &limits, "warmup") + .await + .expect("warm acquire loads the Lua script"); + store.commit(&warmup_key, 0, "warmup").await; + + let key = unique_key("ambiguous-stream-acquire"); + let member = "ambiguous-stream-member"; + + relay.fail_next_reply_after_apply(); + store + .acquire(&key, &limits, member) + .await + .expect("a lost acquire reply falls back locally"); + + // The reply disappeared only after Redis processed it, so a different + // replica sees the shared slot and proves the member exists remotely. + let peer = RedisStore::connect(&single(&relay.url())) + .await + .expect("peer connects through relay"); + assert!( + matches!( + peer.acquire(&key, &limits, "peer-member").await, + Err(aisix_ratelimit::RateLimitError::Concurrency { .. }) + ), + "the ambiguous member must be present in Redis" + ); + + let refreshed = tokio::time::timeout(Duration::from_secs(2), async { + loop { + match store.refresh_stream_lease_outcome(&key, member).await { + StreamLeaseRefresh::Renewed => return StreamLeaseRefresh::Renewed, + StreamLeaseRefresh::Unavailable => { + tokio::time::sleep(Duration::from_millis(20)).await + } + StreamLeaseRefresh::Missing => return StreamLeaseRefresh::Missing, + } + } + }) + .await + .expect("the recovered connection must confirm the ambiguous member"); + assert_eq!(refreshed, StreamLeaseRefresh::Renewed); + + store.release(&key, member); +} + /// A Redis that stops answering without closing the socket must degrade /// the request, not hold it. /// diff --git a/schemas/resources-lenient/passthrough_route.schema.json b/schemas/resources-lenient/passthrough_route.schema.json index a4fe43c8f..84441ea62 100644 --- a/schemas/resources-lenient/passthrough_route.schema.json +++ b/schemas/resources-lenient/passthrough_route.schema.json @@ -554,7 +554,7 @@ ] }, "timeout_ms": { - "description": "Maximum time, in milliseconds, for the upstream exchange. Bounds the response-header phase and any non-SSE body read, but never a healthy SSE relay (which ends with the upstream stream or the client hanging up). When omitted, the gateway default request timeout applies the same way.", + "description": "Maximum time, in milliseconds, for the upstream exchange. Bounds the response-header phase and any non-SSE body read. For SSE it bounds the wait for the first byte and every later no-byte gap, but not the total duration of a healthy relay (which ends with the upstream stream or the client hanging up). When omitted, the gateway default request timeout applies the same way.", "format": "uint64", "minimum": 1.0, "type": [ diff --git a/schemas/resources/passthrough_route.schema.json b/schemas/resources/passthrough_route.schema.json index 3c2521836..01eee57da 100644 --- a/schemas/resources/passthrough_route.schema.json +++ b/schemas/resources/passthrough_route.schema.json @@ -555,7 +555,7 @@ ] }, "timeout_ms": { - "description": "Maximum time, in milliseconds, for the upstream exchange. Bounds the response-header phase and any non-SSE body read, but never a healthy SSE relay (which ends with the upstream stream or the client hanging up). When omitted, the gateway default request timeout applies the same way.", + "description": "Maximum time, in milliseconds, for the upstream exchange. Bounds the response-header phase and any non-SSE body read. For SSE it bounds the wait for the first byte and every later no-byte gap, but not the total duration of a healthy relay (which ends with the upstream stream or the client hanging up). When omitted, the gateway default request timeout applies the same way.", "format": "uint64", "minimum": 1.0, "type": [ diff --git a/tests/e2e/src/cases/guardrail-buffer-cap-enforced-hit-e2e.test.ts b/tests/e2e/src/cases/guardrail-buffer-cap-enforced-hit-e2e.test.ts index 621897dbb..5ddc867d3 100644 --- a/tests/e2e/src/cases/guardrail-buffer-cap-enforced-hit-e2e.test.ts +++ b/tests/e2e/src/cases/guardrail-buffer-cap-enforced-hit-e2e.test.ts @@ -50,6 +50,9 @@ const OPEN_ROW = "cap-fail-open"; const OPEN_ROUTE = "cap-open-passthrough"; const TAIL_ROUTE = "cap-tail-passthrough"; const OPEN_TAIL_ROUTE = "cap-open-tail-passthrough"; +const CUMULATIVE_TAIL_ROUTE = "cap-cumulative-tail-passthrough"; +const OPEN_CUMULATIVE_TAIL_ROUTE = "cap-open-cumulative-tail-passthrough"; +const COMPLETE_FRAME_ROUTE = "cap-complete-frame-passthrough"; // 30 pieces of 100 bytes: three times the tight cap, far under the loose one. const TIGHT_CAP = 1_000; @@ -121,14 +124,40 @@ const RESPONSES_STREAM = [ }), ]; -// A passthrough stream whose last frame never ends: a short answer, then a -// keep-alive the upstream leaves unterminated, carrying no content but past -// the raw bound (128 × the tight cap) the hold-back also keeps. +// A passthrough stream whose last frame never ends: a short answer, then an +// oversized JSON frame deliberately split before its closing bytes. The +// first piece exceeds the raw bound (128 × the tight cap), so the relay must +// make its buffer-cap decision without decoding an incomplete SSE payload. const TAIL_MARKER = "tail-marker"; +const UNTERMINATED_TAIL_JSON = JSON.stringify({ + id: `${TAIL_MARKER}-${"k".repeat(200_000)}`, + object: "chat.completion.chunk", + choices: [], +}); const UNTERMINATED_TAIL_STREAM = [ `data: ${chatChunk({ role: "assistant" })}\n\n`, `data: ${chatChunk({ content: PIECES[0] })}\n\n`, - `data: ${JSON.stringify({ id: `${TAIL_MARKER}-${"k".repeat(200_000)}`, object: "chat.completion.chunk", choices: [] })}`, + `data: ${UNTERMINATED_TAIL_JSON.slice(0, 150_000)}`, + UNTERMINATED_TAIL_JSON.slice(150_000), +]; + +// This malformed frame is fully terminated, so the SSE splitter returns it +// without an overflow marker. Its raw bytes still exceed the hold-back cap +// and must be rejected before any JSON parsing or guardrail extraction. +const COMPLETE_FRAME_MARKER = "complete-frame-marker"; +const COMPLETE_FRAME_STREAM = [ + `data: {"id":"${COMPLETE_FRAME_MARKER}-${"m".repeat(200_000)}\n\n`, +]; + +// These complete comment frames are each below the splitter's raw bound. +// Together they remain held just below it; the malformed EOF tail is also +// below that bound by itself, but makes the cumulative raw hold exceed it. +// A preflight must therefore pick `output_buffer_exceeded` before attempting +// to parse the incomplete JSON as a guardrail source carrier. +const CUMULATIVE_TAIL_MARKER = "cumulative-tail-marker"; +const CUMULATIVE_TAIL_STREAM = [ + ...Array.from({ length: 3 }, () => `: ${"c".repeat(40_000)}\n\n`), + `data: {"id":"${CUMULATIVE_TAIL_MARKER}-${"e".repeat(12_000)}`, ]; interface EnforcedHit { @@ -261,7 +290,12 @@ describe("a stream refused by the hold-back cap names the row whose cap it outgr }); await attachBoth("passthrough_route", route.id); - const tailUp = await startOpenAiUpstream({ rawStreamFrames: UNTERMINATED_TAIL_STREAM }); + const tailUp = await startOpenAiUpstream({ + rawStreamFrames: UNTERMINATED_TAIL_STREAM, + // Ensure the prefix is available to the relay before the closing + // bytes; it reproduces a live upstream that stalls mid-frame. + eventDelayMs: 25, + }); upstreams.push(tailUp); const tailPk = await seed.createProviderKey({ display_name: "cap-tail-backing-pk", @@ -287,6 +321,46 @@ describe("a stream refused by the hold-back cap names the row whose cap it outgr scope_id: openTailRoute.id, priority: 100, }); + const completeFrameUp = await startOpenAiUpstream({ rawStreamFrames: COMPLETE_FRAME_STREAM }); + upstreams.push(completeFrameUp); + const completeFramePk = await seed.createProviderKey({ + display_name: "cap-complete-frame-backing-pk", + secret: "sk-mock", + api_base: `${completeFrameUp.baseUrl}/v1`, + }); + const completeFrameRoute = await seed.createPassthroughRoute({ + name: COMPLETE_FRAME_ROUTE, + path_prefix: "/passthrough/complete-frame", + target_url: `${completeFrameUp.baseUrl}/v1`, + provider_key_id: completeFramePk.id, + }); + await attachBoth("passthrough_route", completeFrameRoute.id); + const cumulativeTailUp = await startOpenAiUpstream({ rawStreamFrames: CUMULATIVE_TAIL_STREAM }); + upstreams.push(cumulativeTailUp); + const cumulativeTailPk = await seed.createProviderKey({ + display_name: "cap-cumulative-tail-backing-pk", + secret: "sk-mock", + api_base: `${cumulativeTailUp.baseUrl}/v1`, + }); + const cumulativeTailRoute = await seed.createPassthroughRoute({ + name: CUMULATIVE_TAIL_ROUTE, + path_prefix: "/passthrough/cumulative-tail", + target_url: `${cumulativeTailUp.baseUrl}/v1`, + provider_key_id: cumulativeTailPk.id, + }); + await attachBoth("passthrough_route", cumulativeTailRoute.id); + const openCumulativeTailRoute = await seed.createPassthroughRoute({ + name: OPEN_CUMULATIVE_TAIL_ROUTE, + path_prefix: "/passthrough/open-cumulative-tail", + target_url: `${cumulativeTailUp.baseUrl}/v1`, + provider_key_id: cumulativeTailPk.id, + }); + await seed.update("guardrail_attachments", randomUUID(), { + guardrail_id: open.id, + scope_type: "passthrough_route", + scope_id: openCumulativeTailRoute.id, + priority: 100, + }); const openRoute = await seed.createPassthroughRoute({ name: OPEN_ROUTE, path_prefix: "/passthrough/cap-open", @@ -421,6 +495,52 @@ describe("a stream refused by the hold-back cap names the row whose cap it outgr await expectBypass("passthrough tail", (l) => l.get("passthrough_route_name") === OPEN_TAIL_ROUTE); }); + test("passthrough route: a complete oversized frame hits the raw cap before parsing", async (ctx) => { + if (!etcdReachable || !app || !sls) return ctx.skip(); + const body = await post("/passthrough/complete-frame/chat/completions", { + model: "gpt-4o-mini", + messages: [{ role: "user", content: "go" }], + }); + expect(body).toContain("output_buffer_exceeded"); + expect(body).not.toContain("unscannable_body"); + expect(body).not.toContain(COMPLETE_FRAME_MARKER); + await expectCapHit( + "passthrough complete frame", + (l) => l.get("passthrough_route_name") === COMPLETE_FRAME_ROUTE, + TIGHT, + ); + }); + + test("passthrough route: a cumulative unterminated EOF tail hits the raw cap before parsing", async (ctx) => { + if (!etcdReachable || !app || !sls) return ctx.skip(); + const body = await post("/passthrough/cumulative-tail/chat/completions", { + model: "gpt-4o-mini", + messages: [{ role: "user", content: "go" }], + }); + expect(body).toContain("output_buffer_exceeded"); + expect(body).not.toContain("unscannable_body"); + expect(body).not.toContain(CUMULATIVE_TAIL_MARKER); + await expectCapHit( + "passthrough cumulative tail", + (l) => l.get("passthrough_route_name") === CUMULATIVE_TAIL_ROUTE, + TIGHT, + ); + }); + + test("passthrough route: a cumulative unterminated EOF tail fails open as a bypass", async (ctx) => { + if (!etcdReachable || !app || !sls) return ctx.skip(); + const body = await post("/passthrough/open-cumulative-tail/chat/completions", { + model: "gpt-4o-mini", + messages: [{ role: "user", content: "go" }], + }); + expect(body, "fail_open releases the malformed EOF tail").toContain(CUMULATIVE_TAIL_MARKER); + expect(body).not.toContain("unscannable_body"); + await expectBypass( + "passthrough cumulative tail", + (l) => l.get("passthrough_route_name") === OPEN_CUMULATIVE_TAIL_ROUTE, + ); + }); + test("a kind with no max_buffer_bytes of its own is named when the default cap trips", async (ctx) => { if (!etcdReachable || !app || !sls) return ctx.skip(); const body = await chat("cap-default"); diff --git a/tests/e2e/src/cases/guardrail-window-stream-hold-e2e.test.ts b/tests/e2e/src/cases/guardrail-window-stream-hold-e2e.test.ts index d2aacb634..25aa1b314 100644 --- a/tests/e2e/src/cases/guardrail-window-stream-hold-e2e.test.ts +++ b/tests/e2e/src/cases/guardrail-window-stream-hold-e2e.test.ts @@ -1,4 +1,4 @@ -import { createHash } from "node:crypto"; +import { createHash, randomUUID } from "node:crypto"; import { createServer, type Server } from "node:http"; import type { AddressInfo } from "node:net"; import { afterAll, beforeAll, describe, expect, test } from "vitest"; @@ -18,7 +18,8 @@ import { // E2E: an Azure text-moderation OUTPUT guardrail in its default `window` // streaming mode, on the routes that hold a streamed response whole — -// /v1/messages and /v1/responses, each native and through the Chat bridge. +// /v1/messages and /v1/responses, each native and through the Chat bridge, +// plus passthrough Chat routes. // // A block-capable output guardrail has to judge the content before any of it // reaches the caller, so on each route: @@ -54,9 +55,18 @@ const FOLD_OPEN_ROW = "window-azure-fold-open"; const FOLD_CLOSED_ROW = "window-azure-fold-closed"; const MIXED_FULL_ROW = "mixed-azure-buffer-full"; const MIXED_WINDOW_ROW = "mixed-azure-window-default"; +const PASSTHROUGH_AT_CAP_ROUTE = "window-pt-cap-equal"; +const PASSTHROUGH_CLOSED_ROUTE = "window-pt-cap-closed"; +const PASSTHROUGH_OPEN_ROUTE = "window-pt-cap-open"; +const PASSTHROUGH_OPEN_TAIL_ROUTE = "window-pt-cap-open-tail"; +const PASSTHROUGH_CLOSED_TAIL_ROUTE = "window-pt-cap-closed-tail"; +const PASSTHROUGH_RAW_CLOSED_ROUTE = "window-pt-raw-cap-closed"; +const PASSTHROUGH_RAW_OPEN_ROUTE = "window-pt-raw-cap-open"; +const PASSTHROUGH_FRAGMENT_CLOSED_ROUTE = "window-pt-fragment-cap-closed"; // 30 pieces of 100 bytes: three times the rows' cap, far under the window. const CAP = 1_000; +const RAW_HOLD_FACTOR = 128; const BIG = Array.from({ length: 30 }, (_, i) => `${String(i).padStart(2, "0")}${"w".repeat(98)}`); // 300 pieces of 1,000 bytes: past the window default (256 KiB), under the // buffer_full row's 1 MiB. @@ -75,6 +85,25 @@ const chatEvents = (pieces: string[]) => [ chatChunk({}, "stop"), "[DONE]", ]; +const WINDOW_OPEN_MARKER = "window-open-marker"; +const WINDOW_OPEN_EVENTS = chatEvents([ + ...BIG.slice(0, 5), + `${WINDOW_OPEN_MARKER}${FLAGGED}${BIG[5]}`, + ...BIG.slice(6), +]); +const WINDOW_AT_CAP_MARKER = "window-at-cap-marker"; +const WINDOW_AT_CAP_EVENTS = chatEvents([ + ...BIG.slice(0, 9), + `${WINDOW_AT_CAP_MARKER}${"e".repeat(BIG[9]!.length - WINDOW_AT_CAP_MARKER.length)}`, +]); +const WINDOW_TAIL_MARKER = "window-tail-marker"; +const WINDOW_TAIL_FRAME = `data: ${chatChunk({ content: `${WINDOW_TAIL_MARKER}${FLAGGED}${"z".repeat(CAP)}` })}`; +const WINDOW_CLOSED_TAIL_MARKER = "window-closed-tail-marker"; +const WINDOW_CLOSED_TAIL_FRAME = `data: ${chatChunk({ content: `${WINDOW_CLOSED_TAIL_MARKER}${"z".repeat(CAP)}` })}`; +const WINDOW_RAW_MARKER = "window-raw-keepalive-marker"; +const WINDOW_RAW_FRAME = `:${WINDOW_RAW_MARKER}${"k".repeat(RAW_HOLD_FACTOR * CAP)}\n\n`; +const WINDOW_FRAGMENT_MARKER = "window-fragment-marker"; +const WINDOW_UNTERMINATED_FRAME = `:${WINDOW_FRAGMENT_MARKER}${"u".repeat(RAW_HOLD_FACTOR * CAP)}`; const anthropicEvents = (pieces: string[]) => [ JSON.stringify({ @@ -169,6 +198,7 @@ describe("a window-mode output guardrail holds a streamed response it cannot rel let sls: MockSls | undefined; let azure: { url: string; close: () => Promise } | undefined; const upstreams: OpenAiUpstream[] = []; + const passthroughUpstreams = new Map(); let etcdReachable = false; const routes = [ @@ -264,8 +294,86 @@ describe("a window-mode output guardrail holds a streamed response it cannot rel } } + const addPassthrough = async ( + name: string, + path_prefix: string, + upstream: OpenAiUpstream, + guardrail_id: string, + ) => { + upstreams.push(upstream); + passthroughUpstreams.set(name, upstream); + const providerKey = await seed.createProviderKey({ + display_name: `${name}-pk`, + secret: "sk-mock", + api_base: `${upstream.baseUrl}/v1`, + }); + const route = await seed.createPassthroughRoute({ + name, + path_prefix, + target_url: `${upstream.baseUrl}/v1`, + provider_key_id: providerKey.id, + }); + await seed.update("guardrail_attachments", randomUUID(), { + guardrail_id, + scope_type: "passthrough_route", + scope_id: route.id, + priority: 100, + }); + }; + await addPassthrough( + PASSTHROUGH_AT_CAP_ROUTE, + "/passthrough/window-cap-equal", + await startOpenAiUpstream({ streamEvents: WINDOW_AT_CAP_EVENTS }), + rows.closed[0]!.id, + ); + await addPassthrough( + PASSTHROUGH_CLOSED_ROUTE, + "/passthrough/window-cap-closed", + await startOpenAiUpstream({ streamEvents: chatEvents(BIG) }), + rows.closed[0]!.id, + ); + await addPassthrough( + PASSTHROUGH_OPEN_ROUTE, + "/passthrough/window-cap-open", + await startOpenAiUpstream({ streamEvents: WINDOW_OPEN_EVENTS }), + rows.open[0]!.id, + ); + await addPassthrough( + PASSTHROUGH_OPEN_TAIL_ROUTE, + "/passthrough/window-cap-open-tail", + await startOpenAiUpstream({ rawStreamFrames: [WINDOW_TAIL_FRAME] }), + rows.open[0]!.id, + ); + await addPassthrough( + PASSTHROUGH_CLOSED_TAIL_ROUTE, + "/passthrough/window-cap-closed-tail", + await startOpenAiUpstream({ rawStreamFrames: [WINDOW_CLOSED_TAIL_FRAME] }), + rows.closed[0]!.id, + ); + await addPassthrough( + PASSTHROUGH_RAW_CLOSED_ROUTE, + "/passthrough/window-raw-cap-closed", + await startOpenAiUpstream({ rawStreamFrames: [WINDOW_RAW_FRAME] }), + rows.closed[0]!.id, + ); + await addPassthrough( + PASSTHROUGH_RAW_OPEN_ROUTE, + "/passthrough/window-raw-cap-open", + await startOpenAiUpstream({ rawStreamFrames: [WINDOW_RAW_FRAME] }), + rows.open[0]!.id, + ); + await addPassthrough( + PASSTHROUGH_FRAGMENT_CLOSED_ROUTE, + "/passthrough/window-fragment-cap-closed", + await startOpenAiUpstream({ + rawStreamFrames: [WINDOW_UNTERMINATED_FRAME, "\n\n"], + eventDelayMs: 3_000, + }), + rows.closed[0]!.id, + ); + // Caller key LAST: it authenticating implies every row above is live. - await seed.createApiKey({ key_hash: hash(CALLER), allowed_models: ["*"] }); + await seed.createApiKey({ key_hash: hash(CALLER), allowed_models: ["*"], allowed_routes: ["*"] }); const proxy = new ProxyClient(app.proxyUrl, CALLER); await waitConfigPropagation(async () => (await proxy.listModels()).status === 200); }, 120_000); @@ -297,6 +405,72 @@ describe("a window-mode output guardrail holds a streamed response it cannot rel return res.text(); }; const byModel = (model: string) => (l: Map) => l.get("requested_model") === model; + const sendPassthrough = async (path: string) => { + const res = await fetch(`${app!.proxyUrl}${path}/chat/completions`, { + method: "POST", + headers: { + "content-type": "application/json", + authorization: `Bearer ${CALLER}`, + "x-api-key": CALLER, + }, + body: JSON.stringify({ + model: "gpt-4o-mini", + messages: [{ role: "user", content: "go" }], + stream: true, + }), + }); + return res.text(); + }; + const beforeTimeout = (promise: Promise, timeoutMs: number) => + new Promise((resolve, reject) => { + const timer = setTimeout( + () => reject(new Error(`passthrough response did not arrive within ${timeoutMs}ms`)), + timeoutMs, + ); + void promise.then( + (value) => { + clearTimeout(timer); + resolve(value); + }, + (error: unknown) => { + clearTimeout(timer); + reject(error); + }, + ); + }); + const passthroughUpstream = (name: string) => { + const upstream = passthroughUpstreams.get(name); + if (!upstream) throw new Error(`missing passthrough upstream for ${name}`); + return upstream; + }; + const waitForPassthroughRequest = async (name: string, timeoutMs = 5_000) => { + const upstream = passthroughUpstream(name); + await new Promise((resolve, reject) => { + const deadline = Date.now() + timeoutMs; + const poll = () => { + if (upstream.receivedRequests.length > 0) { + resolve(); + return; + } + if (Date.now() >= deadline) { + reject(new Error(`upstream ${name} did not receive a request`)); + return; + } + setTimeout(poll, 10); + }; + poll(); + }); + }; + const expectPassthroughRequest = (name: string) => { + const request = passthroughUpstream(name).receivedRequests[0]; + expect(passthroughUpstream(name).receivedRequests).toHaveLength(1); + expect(request).toMatchObject({ method: "POST", path: "/v1/chat/completions" }); + expect(JSON.parse(request!.body)).toMatchObject({ + model: "gpt-4o-mini", + messages: [{ role: "user", content: "go" }], + stream: true, + }); + }; for (const r of routes) { test(`${r.route}: a default window row does not tighten a 1 MiB fail_open buffer_full sibling`, async (ctx) => { @@ -370,4 +544,151 @@ describe("a window-mode output guardrail holds a streamed response it cannot rel ); }); } + + test("passthrough route: a Window cap does not trip at exactly its configured content limit", async (ctx) => { + if (!etcdReachable || !app || !sls) return ctx.skip(); + const body = await sendPassthrough("/passthrough/window-cap-equal"); + expectPassthroughRequest(PASSTHROUGH_AT_CAP_ROUTE); + expect(body).toContain(WINDOW_AT_CAP_MARKER); + expect(body).not.toContain("output_buffer_exceeded"); + const log = await waitForSlsLog( + sls, + LOGSTORE, + (l) => l.get("passthrough_route_name") === PASSTHROUGH_AT_CAP_ROUTE, + "passthrough Window at-cap control", + ); + expect(log.get("guardrail_blocked") ?? "false").not.toBe("true"); + expect(log.get("guardrail_bypassed_reason") ?? "").toBe(""); + expect(JSON.parse(log.get("guardrail_enforced_hits") ?? "[]")).toEqual([]); + }); + + test("passthrough route: a Window cap refuses regular SSE frames before the window closes", async (ctx) => { + if (!etcdReachable || !app || !sls) return ctx.skip(); + const body = await sendPassthrough("/passthrough/window-cap-closed"); + expectPassthroughRequest(PASSTHROUGH_CLOSED_ROUTE); + expect(body).toContain("output_buffer_exceeded"); + expect(body, "a Window cap must not release its held prefix before refusing").not.toContain(BIG[0]); + expect(body).not.toContain(BIG[29]); + const log = await waitForSlsLog( + sls, + LOGSTORE, + (l) => l.get("passthrough_route_name") === PASSTHROUGH_CLOSED_ROUTE && l.get("guardrail_blocked") === "true", + "passthrough Window fail-closed cap", + ); + const hits = JSON.parse(log.get("guardrail_enforced_hits") ?? "[]") as EnforcedHit[]; + expect(hits.map(({ guardrail_name, hook, action }) => ({ guardrail_name, hook, action }))).toEqual([ + { guardrail_name: CLOSED_ROW, hook: "output", action: "blocked_buffer_exceeded" }, + ]); + }); + + test("passthrough route: a Window fail-open cap does not rescan regular frames at EOF", async (ctx) => { + if (!etcdReachable || !app || !sls) return ctx.skip(); + const body = await sendPassthrough("/passthrough/window-cap-open"); + expectPassthroughRequest(PASSTHROUGH_OPEN_ROUTE); + expect(body, "fail_open releases the held regular frames").toContain(WINDOW_OPEN_MARKER); + expect(body, "a later EOF scan must not retract fail-open content").toContain(FLAGGED); + expect(body).not.toContain("output_buffer_exceeded"); + const log = await waitForSlsLog( + sls, + LOGSTORE, + (l) => l.get("passthrough_route_name") === PASSTHROUGH_OPEN_ROUTE, + "passthrough Window regular-frame fail-open cap", + ); + expect(log.get("guardrail_blocked") ?? "false").not.toBe("true"); + expect(log.get("guardrail_bypassed_reason")).toBe("output_buffer_exceeded"); + expect(JSON.parse(log.get("guardrail_enforced_hits") ?? "[]")).toEqual([]); + }); + + test("passthrough route: a Window cap fails open for an unterminated final SSE frame", async (ctx) => { + if (!etcdReachable || !app || !sls) return ctx.skip(); + const body = await sendPassthrough("/passthrough/window-cap-open-tail"); + expectPassthroughRequest(PASSTHROUGH_OPEN_TAIL_ROUTE); + expect(body, "fail_open releases the held final frame").toContain(WINDOW_TAIL_MARKER); + expect(body, "the final frame is not rescanned after it is released").toContain(FLAGGED); + expect(body).not.toContain("output_buffer_exceeded"); + const log = await waitForSlsLog( + sls, + LOGSTORE, + (l) => l.get("passthrough_route_name") === PASSTHROUGH_OPEN_TAIL_ROUTE, + "passthrough Window unterminated-tail fail-open cap", + ); + expect(log.get("guardrail_blocked") ?? "false").not.toBe("true"); + expect(log.get("guardrail_bypassed_reason")).toBe("output_buffer_exceeded"); + expect(JSON.parse(log.get("guardrail_enforced_hits") ?? "[]")).toEqual([]); + }); + + test("passthrough route: a Window cap refuses an unterminated final SSE frame under fail_closed", async (ctx) => { + if (!etcdReachable || !app || !sls) return ctx.skip(); + const body = await sendPassthrough("/passthrough/window-cap-closed-tail"); + expectPassthroughRequest(PASSTHROUGH_CLOSED_TAIL_ROUTE); + expect(body).toContain("output_buffer_exceeded"); + expect(body, "the terminal frame remains held when the cap fails closed").not.toContain(WINDOW_CLOSED_TAIL_MARKER); + const log = await waitForSlsLog( + sls, + LOGSTORE, + (l) => l.get("passthrough_route_name") === PASSTHROUGH_CLOSED_TAIL_ROUTE && l.get("guardrail_blocked") === "true", + "passthrough Window unterminated-tail fail-closed cap", + ); + const hits = JSON.parse(log.get("guardrail_enforced_hits") ?? "[]") as EnforcedHit[]; + expect(hits.map(({ guardrail_name, hook, action }) => ({ guardrail_name, hook, action }))).toEqual([ + { guardrail_name: CLOSED_ROW, hook: "output", action: "blocked_buffer_exceeded" }, + ]); + }); + + test("passthrough route: a Window raw-byte cap refuses a delta-free keepalive frame", async (ctx) => { + if (!etcdReachable || !app || !sls) return ctx.skip(); + const body = await sendPassthrough("/passthrough/window-raw-cap-closed"); + expectPassthroughRequest(PASSTHROUGH_RAW_CLOSED_ROUTE); + expect(body).toContain("output_buffer_exceeded"); + expect(body, "a delta-free frame still counts toward the raw held-byte cap").not.toContain(WINDOW_RAW_MARKER); + const log = await waitForSlsLog( + sls, + LOGSTORE, + (l) => l.get("passthrough_route_name") === PASSTHROUGH_RAW_CLOSED_ROUTE && l.get("guardrail_blocked") === "true", + "passthrough Window raw-byte fail-closed cap", + ); + const hits = JSON.parse(log.get("guardrail_enforced_hits") ?? "[]") as EnforcedHit[]; + expect(hits.map(({ guardrail_name, hook, action }) => ({ guardrail_name, hook, action }))).toEqual([ + { guardrail_name: CLOSED_ROW, hook: "output", action: "blocked_buffer_exceeded" }, + ]); + }); + + test("passthrough route: a Window raw-byte cap releases a delta-free keepalive frame under fail_open", async (ctx) => { + if (!etcdReachable || !app || !sls) return ctx.skip(); + const body = await sendPassthrough("/passthrough/window-raw-cap-open"); + expectPassthroughRequest(PASSTHROUGH_RAW_OPEN_ROUTE); + expect(body, "fail_open releases the raw keepalive frame").toContain(WINDOW_RAW_MARKER); + expect(body).not.toContain("output_buffer_exceeded"); + const log = await waitForSlsLog( + sls, + LOGSTORE, + (l) => l.get("passthrough_route_name") === PASSTHROUGH_RAW_OPEN_ROUTE, + "passthrough Window raw-byte fail-open cap", + ); + expect(log.get("guardrail_blocked") ?? "false").not.toBe("true"); + expect(log.get("guardrail_bypassed_reason")).toBe("output_buffer_exceeded"); + expect(JSON.parse(log.get("guardrail_enforced_hits") ?? "[]")).toEqual([]); + }); + + test("passthrough route: a Window frame cap applies before an unterminated upstream frame reaches EOF", async (ctx) => { + if (!etcdReachable || !app || !sls) return ctx.skip(); + const response = sendPassthrough("/passthrough/window-fragment-cap-closed"); + await waitForPassthroughRequest(PASSTHROUGH_FRAGMENT_CLOSED_ROUTE); + const body = await beforeTimeout(response, 1_500); + expectPassthroughRequest(PASSTHROUGH_FRAGMENT_CLOSED_ROUTE); + expect(body).toContain("output_buffer_exceeded"); + expect(body, "the oversized partial frame remains held when the cap fails closed").not.toContain( + WINDOW_FRAGMENT_MARKER, + ); + const log = await waitForSlsLog( + sls, + LOGSTORE, + (l) => l.get("passthrough_route_name") === PASSTHROUGH_FRAGMENT_CLOSED_ROUTE && l.get("guardrail_blocked") === "true", + "passthrough Window unterminated-frame fail-closed cap", + ); + const hits = JSON.parse(log.get("guardrail_enforced_hits") ?? "[]") as EnforcedHit[]; + expect(hits.map(({ guardrail_name, hook, action }) => ({ guardrail_name, hook, action }))).toEqual([ + { guardrail_name: CLOSED_ROW, hook: "output", action: "blocked_buffer_exceeded" }, + ]); + }); }); diff --git a/tests/e2e/src/cases/passthrough-chat-media-guardrail-e2e.test.ts b/tests/e2e/src/cases/passthrough-chat-media-guardrail-e2e.test.ts new file mode 100644 index 000000000..5a4e14a70 --- /dev/null +++ b/tests/e2e/src/cases/passthrough-chat-media-guardrail-e2e.test.ts @@ -0,0 +1,677 @@ +import { createHash } from "node:crypto"; +import { createServer, type Server } from "node:http"; +import { afterAll, beforeAll, describe, expect, test } from "vitest"; +import { + EtcdClient, + ProxyClient, + SeedClient, + pickFreePort, + spawnApp, + startOpenAiUpstream, + waitConfigPropagation, + type OpenAiUpstream, + type SpawnedApp, +} from "../harness/index.js"; + +// A real AISIX gateway relays Chat output to a local upstream and sends its +// selected output text to a separate OpenAI Moderation-compatible peer. The +// peer records exactly what AISIX sends across the external guardrail boundary. + +const CALLER = "sk-passthrough-chat-media"; +const CALLER_HASH = createHash("sha256").update(CALLER).digest("hex"); +const BUFFERED_IMAGE = "buffered-image-media-sentinel"; +const BUFFERED_AUDIO = "buffered-audio-media-sentinel"; +const BUFFERED_FILE = "buffered-file-media-sentinel"; +const BUFFERED_OPAQUE = "buffered-opaque-part-sentinel"; +const BUFFERED_REASONING = "buffered-reasoning-sentinel"; +const BUFFERED_VISIBLE = "buffered-visible-text-sentinel"; +const BUFFERED_TOOL = "buffered-tool-arguments-sentinel"; +const STREAM_IMAGE = "stream-image-media-sentinel"; +const STREAM_AUDIO = "stream-audio-media-sentinel"; +const STREAM_FILE = "stream-file-media-sentinel"; +const STREAM_OPAQUE = "stream-opaque-part-sentinel"; +const STREAM_REASONING = "stream-reasoning-sentinel"; +const STREAM_VISIBLE = "stream-visible-text-sentinel"; +const STREAM_TOOL = "stream-tool-arguments-sentinel"; +const BUFFERED_MESSAGE_REFUSAL = "buffered-message-refusal-BLOCKME"; +const BUFFERED_CONTENT_REFUSAL = "buffered-content-refusal-BLOCKME"; +const STREAM_REFUSAL = "stream-refusal-BLOCKME"; +// `openai_moderation` uses the default whole-stream hold cap (256 KiB). +const STREAM_REFUSAL_CAP = 262_144; +const STREAM_OVERSIZED_REFUSAL = `stream-oversized-refusal-${"x".repeat(STREAM_REFUSAL_CAP + 1)}`; +const COMPLETIONS_VISIBLE = "completions-visible-text-sentinel"; +const COMPLETIONS_OPAQUE = "completions-opaque-extension-sentinel"; + +interface ModerationSink { + baseUrl: string; + inputs: string[]; + close(): Promise; +} + +async function startModerationSink(): Promise { + const inputs: string[] = []; + const server: Server = createServer((req, res) => { + let raw = ""; + req.on("data", (chunk: Buffer) => (raw += chunk.toString("utf8"))); + req.on("end", () => { + let input: string | undefined; + try { + const body = JSON.parse(raw) as { input?: unknown }; + if (typeof body.input === "string") { + input = body.input; + inputs.push(input); + } + } catch { + // Reply normally so a malformed moderation request remains visible + // through the gateway's own output-policy behavior. + } + const flagged = input?.includes("BLOCKME") ?? false; + res.statusCode = 200; + res.setHeader("content-type", "application/json"); + res.end( + JSON.stringify({ + results: [{ flagged, categories: flagged ? { refusal: true } : {} }], + }), + ); + }); + }); + const port = await pickFreePort(); + await new Promise((resolve, reject) => { + server.once("error", reject); + server.listen(port, "127.0.0.1", resolve); + }); + return { + baseUrl: `http://127.0.0.1:${port}`, + inputs, + async close() { + await new Promise((resolve, reject) => { + server.close((err) => (err ? reject(err) : resolve())); + }); + }, + }; +} + +const bufferedResponse = { + id: "chat_media_buffered", + object: "chat.completion", + model: "gpt-4o-mini", + choices: [ + { + index: 0, + message: { + role: "assistant", + content: [ + { type: "image_url", image_url: { url: BUFFERED_IMAGE } }, + { type: "input_audio", input_audio: { data: BUFFERED_AUDIO } }, + { type: "file", file: { file_data: BUFFERED_FILE } }, + { type: "future_media", text: BUFFERED_OPAQUE }, + { type: "text", text: BUFFERED_VISIBLE }, + ], + reasoning_content: BUFFERED_REASONING, + tool_calls: [ + { + id: "call_media_buffered", + type: "function", + function: { name: "lookup", arguments: BUFFERED_TOOL }, + }, + ], + }, + finish_reason: "tool_calls", + }, + ], +}; + +const streamedResponse = [ + `data: ${JSON.stringify({ + id: "chat_media_stream", + object: "chat.completion.chunk", + model: "gpt-4o-mini", + choices: [ + { + index: 0, + delta: { + role: "assistant", + content: [ + { index: 0, type: "image_url", image_url: { url: STREAM_IMAGE } }, + { index: 1, type: "input_audio", input_audio: { data: STREAM_AUDIO } }, + { index: 2, type: "file", file: { file_data: STREAM_FILE } }, + { index: 3, type: "future_media", text: STREAM_OPAQUE }, + { index: 4, type: "text", text: STREAM_VISIBLE }, + ], + reasoning_content: STREAM_REASONING, + tool_calls: [ + { + index: 0, + id: "call_media_stream", + type: "function", + function: { name: "lookup", arguments: STREAM_TOOL }, + }, + ], + }, + }, + ], + })}\n\n`, + "data: [DONE]\n\n", +]; + +const bufferedMessageRefusalResponse = { + id: "chat_message_refusal_buffered", + object: "chat.completion", + model: "gpt-4o-mini", + choices: [ + { + index: 0, + message: { role: "assistant", content: null, refusal: BUFFERED_MESSAGE_REFUSAL }, + finish_reason: "content_filter", + }, + ], +}; + +const bufferedContentRefusalResponse = { + id: "chat_content_refusal_buffered", + object: "chat.completion", + model: "gpt-4o-mini", + choices: [ + { + index: 0, + message: { + role: "assistant", + content: [{ type: "refusal", refusal: BUFFERED_CONTENT_REFUSAL }], + }, + finish_reason: "content_filter", + }, + ], +}; + +const streamedRefusalResponse = [ + `data: ${JSON.stringify({ + id: "chat_refusal_stream", + object: "chat.completion.chunk", + model: "gpt-4o-mini", + choices: [{ index: 0, delta: { role: "assistant", refusal: STREAM_REFUSAL } }], + })}\n\n`, + "data: [DONE]\n\n", +]; + +const streamedOversizedRefusalResponse = [ + `data: ${JSON.stringify({ + id: "chat_oversized_refusal_stream", + object: "chat.completion.chunk", + model: "gpt-4o-mini", + choices: [{ index: 0, delta: { role: "assistant", refusal: STREAM_OVERSIZED_REFUSAL } }], + })}\n\n`, + "data: [DONE]\n\n", +]; + +const streamedOversizedContentRefusalResponse = [ + `data: ${JSON.stringify({ + id: "chat_oversized_content_refusal_stream", + object: "chat.completion.chunk", + model: "gpt-4o-mini", + choices: [ + { + index: 0, + delta: { + role: "assistant", + content: [{ index: 0, type: "refusal", refusal: STREAM_OVERSIZED_REFUSAL }], + }, + }, + ], + })}\n\n`, + "data: [DONE]\n\n", +]; + +const streamedOversizedLegacyToolResponse = [ + `data: ${JSON.stringify({ + id: "chat_oversized_legacy_tool_stream", + object: "chat.completion.chunk", + model: "gpt-4o-mini", + choices: [ + { + index: 0, + delta: { + role: "assistant", + function_call: { name: "lookup", arguments: STREAM_OVERSIZED_REFUSAL }, + }, + }, + ], + })}\n\n`, + "data: [DONE]\n\n", +]; + +const bufferedCompletionsResponse = { + choices: [{ text: COMPLETIONS_VISIBLE, opaque: COMPLETIONS_OPAQUE }], + opaque: { data: COMPLETIONS_OPAQUE }, +}; + +const malformedBufferedCompletionsResponse = { + choices: { text: COMPLETIONS_VISIBLE }, + opaque: { data: COMPLETIONS_OPAQUE }, +}; + +const malformedStreamedCompletionsResponse = [ + `data: ${JSON.stringify({ + choices: { index: 0, text: COMPLETIONS_VISIBLE }, + opaque: { data: COMPLETIONS_OPAQUE }, + })}\n\n`, + "data: [DONE]\n\n", +]; + +describe("Chat passthrough keeps media out of external output guardrails", () => { + let app: SpawnedApp | undefined; + let seed: SeedClient | undefined; + let bufferedUpstream: OpenAiUpstream | undefined; + let streamUpstream: OpenAiUpstream | undefined; + let bufferedMessageRefusalUpstream: OpenAiUpstream | undefined; + let bufferedContentRefusalUpstream: OpenAiUpstream | undefined; + let streamRefusalUpstream: OpenAiUpstream | undefined; + let streamOversizedRefusalUpstream: OpenAiUpstream | undefined; + let streamOversizedContentRefusalUpstream: OpenAiUpstream | undefined; + let streamOversizedLegacyToolUpstream: OpenAiUpstream | undefined; + let bufferedCompletionsUpstream: OpenAiUpstream | undefined; + let malformedBufferedCompletionsUpstream: OpenAiUpstream | undefined; + let malformedStreamedCompletionsUpstream: OpenAiUpstream | undefined; + let moderation: ModerationSink | undefined; + let etcdReachable = false; + + beforeAll(async () => { + const etcd = new EtcdClient(); + etcdReachable = await etcd.ping(); + if (!etcdReachable) return; + + moderation = await startModerationSink(); + bufferedUpstream = await startOpenAiUpstream({ nonStreamBody: bufferedResponse }); + streamUpstream = await startOpenAiUpstream({ rawStreamFrames: streamedResponse }); + bufferedMessageRefusalUpstream = await startOpenAiUpstream({ + nonStreamBody: bufferedMessageRefusalResponse, + }); + bufferedContentRefusalUpstream = await startOpenAiUpstream({ + nonStreamBody: bufferedContentRefusalResponse, + }); + streamRefusalUpstream = await startOpenAiUpstream({ rawStreamFrames: streamedRefusalResponse }); + streamOversizedRefusalUpstream = await startOpenAiUpstream({ + rawStreamFrames: streamedOversizedRefusalResponse, + }); + streamOversizedContentRefusalUpstream = await startOpenAiUpstream({ + rawStreamFrames: streamedOversizedContentRefusalResponse, + }); + streamOversizedLegacyToolUpstream = await startOpenAiUpstream({ + rawStreamFrames: streamedOversizedLegacyToolResponse, + }); + bufferedCompletionsUpstream = await startOpenAiUpstream({ + nonStreamBody: bufferedCompletionsResponse, + }); + malformedBufferedCompletionsUpstream = await startOpenAiUpstream({ + nonStreamBody: malformedBufferedCompletionsResponse, + }); + malformedStreamedCompletionsUpstream = await startOpenAiUpstream({ + rawStreamFrames: malformedStreamedCompletionsResponse, + }); + app = await spawnApp(); + seed = new SeedClient(etcd, app.etcdPrefix); + + const providerKey = await seed.createProviderKey({ + display_name: "passthrough-chat-media-pk", + secret: "sk-mock", + api_base: bufferedUpstream.baseUrl, + }); + await seed.createPassthroughRoute({ + name: "passthrough-chat-media-buffered", + path_prefix: "/chat-media-buffered", + target_url: bufferedUpstream.baseUrl, + provider_key_id: providerKey.id, + }); + await seed.createPassthroughRoute({ + name: "passthrough-chat-media-stream", + path_prefix: "/chat-media-stream", + target_url: streamUpstream.baseUrl, + provider_key_id: providerKey.id, + }); + await seed.createPassthroughRoute({ + name: "passthrough-chat-refusal-message", + path_prefix: "/chat-refusal-message", + target_url: bufferedMessageRefusalUpstream.baseUrl, + provider_key_id: providerKey.id, + }); + await seed.createPassthroughRoute({ + name: "passthrough-chat-refusal-content", + path_prefix: "/chat-refusal-content", + target_url: bufferedContentRefusalUpstream.baseUrl, + provider_key_id: providerKey.id, + }); + await seed.createPassthroughRoute({ + name: "passthrough-chat-refusal-stream", + path_prefix: "/chat-refusal-stream", + target_url: streamRefusalUpstream.baseUrl, + provider_key_id: providerKey.id, + }); + await seed.createPassthroughRoute({ + name: "passthrough-chat-refusal-stream-oversized", + path_prefix: "/chat-refusal-stream-oversized", + target_url: streamOversizedRefusalUpstream.baseUrl, + provider_key_id: providerKey.id, + }); + await seed.createPassthroughRoute({ + name: "passthrough-chat-refusal-stream-oversized-content", + path_prefix: "/chat-refusal-stream-oversized-content", + target_url: streamOversizedContentRefusalUpstream.baseUrl, + provider_key_id: providerKey.id, + }); + await seed.createPassthroughRoute({ + name: "passthrough-chat-legacy-tool-stream-oversized", + path_prefix: "/chat-legacy-tool-stream-oversized", + target_url: streamOversizedLegacyToolUpstream.baseUrl, + provider_key_id: providerKey.id, + }); + await seed.createPassthroughRoute({ + name: "passthrough-completions-buffered", + path_prefix: "/completions-buffered", + target_url: bufferedCompletionsUpstream.baseUrl, + provider_key_id: providerKey.id, + }); + await seed.createPassthroughRoute({ + name: "passthrough-completions-malformed-buffered", + path_prefix: "/completions-malformed-buffered", + target_url: malformedBufferedCompletionsUpstream.baseUrl, + provider_key_id: providerKey.id, + }); + await seed.createPassthroughRoute({ + name: "passthrough-completions-malformed-stream", + path_prefix: "/completions-malformed-stream", + target_url: malformedStreamedCompletionsUpstream.baseUrl, + provider_key_id: providerKey.id, + }); + await seed.createGuardrail({ + name: "passthrough-chat-media-output", + enabled: true, + hook_point: "output", + kind: "openai_moderation", + api_key: "sk-local-moderation", + endpoint: moderation.baseUrl, + output_fail_open: false, + }); + // Seed the authentication gate last, then wait for it through the real + // models endpoint so the routes and output guardrail are already active. + await seed.createApiKey({ + key_hash: CALLER_HASH, + allowed_models: [], + allowed_routes: ["*"], + }); + const proxy = new ProxyClient(app.proxyUrl, CALLER); + await waitConfigPropagation(async () => (await proxy.listModels()).status === 200); + }, 120_000); + + afterAll(async () => { + await app?.exit(); + await bufferedUpstream?.close(); + await streamUpstream?.close(); + await bufferedMessageRefusalUpstream?.close(); + await bufferedContentRefusalUpstream?.close(); + await streamRefusalUpstream?.close(); + await streamOversizedRefusalUpstream?.close(); + await streamOversizedContentRefusalUpstream?.close(); + await streamOversizedLegacyToolUpstream?.close(); + await bufferedCompletionsUpstream?.close(); + await malformedBufferedCompletionsUpstream?.close(); + await malformedStreamedCompletionsUpstream?.close(); + await moderation?.close(); + }); + + const request = (route: string, stream: boolean) => + fetch(`${app!.proxyUrl}/${route}/v1/chat/completions`, { + method: "POST", + headers: { + authorization: `Bearer ${CALLER}`, + "content-type": "application/json", + }, + body: JSON.stringify({ + model: "gpt-4o-mini", + messages: [{ role: "user", content: "go" }], + stream, + }), + }); + + const completionsRequest = (route: string, stream: boolean) => + fetch(`${app!.proxyUrl}/${route}/v1/completions`, { + method: "POST", + headers: { + authorization: `Bearer ${CALLER}`, + "content-type": "application/json", + }, + body: JSON.stringify({ model: "gpt-4o-mini", prompt: "go", stream }), + }); + + const expectExternalGuardrailText = ( + inputs: string[], + visible: string, + tool: string, + opaque: string[], + ) => { + expect(inputs.length, "the external output guardrail was invoked").toBeGreaterThan(0); + expect(inputs.some((input) => input.includes(visible)), inputs.join("\n")).toBe(true); + expect(inputs.some((input) => input.includes(tool)), inputs.join("\n")).toBe(true); + const visibleOccurrences = inputs.reduce( + (count, input) => count + input.split(visible).length - 1, + 0, + ); + expect(visibleOccurrences, `external guardrail input: ${inputs.join("\n")}`).toBe(1); + for (const value of opaque) { + expect(inputs.every((input) => !input.includes(value)), inputs.join("\n")).toBe(true); + } + }; + + const expectBlockedRefusal = async ( + route: string, + stream: boolean, + upstream: OpenAiUpstream, + refusal: string, + ) => { + const upstreamBefore = upstream.receivedRequests.length; + const moderationBefore = moderation!.inputs.length; + const response = await request(route, stream); + const body = await response.text(); + expect(response.status, body).toBe(stream ? 200 : 422); + expect(body).toContain("content_filter"); + if (stream) expect(body).toContain("event: error"); + expect(body).not.toContain(refusal); + expect(upstream.receivedRequests.length).toBe(upstreamBefore + 1); + const inputs = moderation!.inputs.slice(moderationBefore); + expect(inputs.some((input) => input.includes(refusal)), inputs.join("\n")).toBe(true); + }; + + test("buffered Chat media stays out of the external guardrail while text and tools are scanned", async (ctx) => { + if (!etcdReachable || !app || !bufferedUpstream || !moderation) { + ctx.skip(); + return; + } + const upstreamBefore = bufferedUpstream.receivedRequests.length; + const moderationBefore = moderation.inputs.length; + const response = await request("chat-media-buffered", false); + const body = await response.text(); + expect(response.status, body).toBe(200); + for (const value of [ + BUFFERED_IMAGE, + BUFFERED_AUDIO, + BUFFERED_FILE, + BUFFERED_OPAQUE, + BUFFERED_REASONING, + BUFFERED_VISIBLE, + BUFFERED_TOOL, + ]) { + expect(body).toContain(value); + } + expect(bufferedUpstream.receivedRequests.length).toBe(upstreamBefore + 1); + expectExternalGuardrailText( + moderation.inputs.slice(moderationBefore), + BUFFERED_VISIBLE, + BUFFERED_TOOL, + [BUFFERED_IMAGE, BUFFERED_AUDIO, BUFFERED_FILE, BUFFERED_OPAQUE, BUFFERED_REASONING], + ); + }); + + test("streamed Chat media stays out of the external guardrail while text and tools are scanned", async (ctx) => { + if (!etcdReachable || !app || !streamUpstream || !moderation) { + ctx.skip(); + return; + } + const upstreamBefore = streamUpstream.receivedRequests.length; + const moderationBefore = moderation.inputs.length; + const response = await request("chat-media-stream", true); + const body = await response.text(); + expect(response.status, body).toBe(200); + for (const value of [ + STREAM_IMAGE, + STREAM_AUDIO, + STREAM_FILE, + STREAM_OPAQUE, + STREAM_REASONING, + STREAM_VISIBLE, + STREAM_TOOL, + ]) { + expect(body).toContain(value); + } + expect(streamUpstream.receivedRequests.length).toBe(upstreamBefore + 1); + expectExternalGuardrailText( + moderation.inputs.slice(moderationBefore), + STREAM_VISIBLE, + STREAM_TOOL, + [STREAM_IMAGE, STREAM_AUDIO, STREAM_FILE, STREAM_OPAQUE, STREAM_REASONING], + ); + }); + + test("a Chat request with a buffered Completions response sends only choices text to the external guardrail", async (ctx) => { + if (!etcdReachable || !app || !bufferedCompletionsUpstream || !moderation) { + ctx.skip(); + return; + } + const upstreamBefore = bufferedCompletionsUpstream.receivedRequests.length; + const moderationBefore = moderation.inputs.length; + const response = await request("completions-buffered", false); + const body = await response.text(); + expect(response.status, body).toBe(200); + expect(body).toContain(COMPLETIONS_VISIBLE); + expect(body).toContain(COMPLETIONS_OPAQUE); + expect(bufferedCompletionsUpstream.receivedRequests.length).toBe(upstreamBefore + 1); + const inputs = moderation.inputs.slice(moderationBefore); + expect(inputs.some((input) => input.includes(COMPLETIONS_VISIBLE)), inputs.join("\n")).toBe(true); + expect(inputs.every((input) => !input.includes(COMPLETIONS_OPAQUE)), inputs.join("\n")).toBe(true); + }); + + test("malformed Completions carriers fail closed without calling the external guardrail", async (ctx) => { + if ( + !etcdReachable || + !app || + !malformedBufferedCompletionsUpstream || + !malformedStreamedCompletionsUpstream || + !moderation + ) { + ctx.skip(); + return; + } + + const bufferedBefore = malformedBufferedCompletionsUpstream.receivedRequests.length; + const moderationBeforeBuffered = moderation.inputs.length; + const buffered = await completionsRequest("completions-malformed-buffered", false); + const bufferedBody = await buffered.text(); + expect(buffered.status, bufferedBody).toBe(422); + expect(bufferedBody).toContain("guardrail_unavailable"); + expect(bufferedBody).not.toContain(COMPLETIONS_OPAQUE); + expect(malformedBufferedCompletionsUpstream.receivedRequests.length).toBe(bufferedBefore + 1); + expect(moderation.inputs.slice(moderationBeforeBuffered)).toHaveLength(0); + + const streamBefore = malformedStreamedCompletionsUpstream.receivedRequests.length; + const moderationBeforeStream = moderation.inputs.length; + const streamed = await completionsRequest("completions-malformed-stream", true); + const streamedBody = await streamed.text(); + expect(streamed.status, streamedBody).toBe(200); + expect(streamedBody).toContain("event: error"); + expect(streamedBody).toContain("guardrail_unavailable"); + expect(streamedBody).not.toContain(COMPLETIONS_OPAQUE); + expect(malformedStreamedCompletionsUpstream.receivedRequests.length).toBe(streamBefore + 1); + expect(moderation.inputs.slice(moderationBeforeStream)).toHaveLength(0); + }); + + test("buffered Chat refusals are blocked by the external output guardrail", async (ctx) => { + if ( + !etcdReachable || + !app || + !bufferedMessageRefusalUpstream || + !bufferedContentRefusalUpstream || + !moderation + ) { + ctx.skip(); + return; + } + await expectBlockedRefusal( + "chat-refusal-message", + false, + bufferedMessageRefusalUpstream, + BUFFERED_MESSAGE_REFUSAL, + ); + await expectBlockedRefusal( + "chat-refusal-content", + false, + bufferedContentRefusalUpstream, + BUFFERED_CONTENT_REFUSAL, + ); + }); + + test("streamed Chat refusal deltas are blocked by the external output guardrail", async (ctx) => { + if (!etcdReachable || !app || !streamRefusalUpstream || !moderation) { + ctx.skip(); + return; + } + await expectBlockedRefusal("chat-refusal-stream", true, streamRefusalUpstream, STREAM_REFUSAL); + }); + + test("streamed Chat refusal deltas count toward the output hold cap", async (ctx) => { + if (!etcdReachable || !app || !streamOversizedRefusalUpstream || !moderation) { + ctx.skip(); + return; + } + const upstreamBefore = streamOversizedRefusalUpstream.receivedRequests.length; + const moderationBefore = moderation.inputs.length; + const response = await request("chat-refusal-stream-oversized", true); + const body = await response.text(); + const bodyExcerpt = body.slice(0, 512); + expect(response.status, bodyExcerpt).toBe(200); + expect(body.includes("output_buffer_exceeded"), bodyExcerpt).toBe(true); + expect(body.includes(STREAM_OVERSIZED_REFUSAL), bodyExcerpt).toBe(false); + expect(streamOversizedRefusalUpstream.receivedRequests.length).toBe(upstreamBefore + 1); + expect(moderation.inputs.slice(moderationBefore)).toHaveLength(0); + }); + + test("streamed typed Chat refusals count toward the output hold cap", async (ctx) => { + if (!etcdReachable || !app || !streamOversizedContentRefusalUpstream || !moderation) { + ctx.skip(); + return; + } + const upstreamBefore = streamOversizedContentRefusalUpstream.receivedRequests.length; + const moderationBefore = moderation.inputs.length; + const response = await request("chat-refusal-stream-oversized-content", true); + const body = await response.text(); + const bodyExcerpt = body.slice(0, 512); + expect(response.status, bodyExcerpt).toBe(200); + expect(body.includes("output_buffer_exceeded"), bodyExcerpt).toBe(true); + expect(body.includes(STREAM_OVERSIZED_REFUSAL), bodyExcerpt).toBe(false); + expect(streamOversizedContentRefusalUpstream.receivedRequests.length).toBe(upstreamBefore + 1); + expect(moderation.inputs.slice(moderationBefore)).toHaveLength(0); + }); + + test("streamed legacy Chat tool arguments count toward the output hold cap", async (ctx) => { + if (!etcdReachable || !app || !streamOversizedLegacyToolUpstream || !moderation) { + ctx.skip(); + return; + } + const upstreamBefore = streamOversizedLegacyToolUpstream.receivedRequests.length; + const moderationBefore = moderation.inputs.length; + const response = await request("chat-legacy-tool-stream-oversized", true); + const body = await response.text(); + const bodyExcerpt = body.slice(0, 512); + expect(response.status, bodyExcerpt).toBe(200); + expect(body.includes("output_buffer_exceeded"), bodyExcerpt).toBe(true); + expect(body.includes(STREAM_OVERSIZED_REFUSAL), bodyExcerpt).toBe(false); + expect(streamOversizedLegacyToolUpstream.receivedRequests.length).toBe(upstreamBefore + 1); + expect(moderation.inputs.slice(moderationBefore)).toHaveLength(0); + }); +}); diff --git a/tests/e2e/src/cases/passthrough-guardrail-e2e.test.ts b/tests/e2e/src/cases/passthrough-guardrail-e2e.test.ts index 23bcaafc2..801838089 100644 --- a/tests/e2e/src/cases/passthrough-guardrail-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-guardrail-e2e.test.ts @@ -102,8 +102,8 @@ describe("passthrough guardrail (#911 [6])", () => { sseUpstream = await startOpenAiUpstream({ streamEvents: [ - JSON.stringify({ choices: [{ delta: { content: "prelude " } }] }), - JSON.stringify({ choices: [{ delta: { content: FORBIDDEN_OUTPUT } }] }), + JSON.stringify({ choices: [{ index: 0, delta: { content: "prelude " } }] }), + JSON.stringify({ choices: [{ index: 0, delta: { content: FORBIDDEN_OUTPUT } }] }), "[DONE]", ], }); diff --git a/tests/e2e/src/cases/passthrough-redirect-boundary-e2e.test.ts b/tests/e2e/src/cases/passthrough-redirect-boundary-e2e.test.ts new file mode 100644 index 000000000..1afd111cb --- /dev/null +++ b/tests/e2e/src/cases/passthrough-redirect-boundary-e2e.test.ts @@ -0,0 +1,148 @@ +import { createHash } from "node:crypto"; +import { afterAll, beforeAll, describe, expect, test } from "vitest"; +import { + EtcdClient, + ProxyClient, + SeedClient, + spawnApp, + startOpenAiUpstream, + waitConfigPropagation, + type OpenAiUpstream, + type SpawnedApp, +} from "../harness/index.js"; + +// E2E: a passthrough target is an outbound boundary. An upstream redirect +// must be relayed to the caller, not followed by the gateway: the next hop +// is neither path-checked against target_url nor entitled to its credential. + +const CALLER_PLAINTEXT = "sk-passthrough-redirect-boundary-caller"; +const CALLER_KEY_HASH = createHash("sha256") + .update(CALLER_PLAINTEXT) + .digest("hex"); +const PROVIDER_SECRET = "sk-passthrough-redirect-boundary-provider"; + +const SAME_ORIGIN_ROUTE = "/passthrough/redirect-same-origin"; +const CROSS_ORIGIN_ROUTE = "/passthrough/redirect-cross-origin"; +const SAME_ORIGIN_LOCATION = "/outside-target/same-origin"; + +describe("passthrough redirect boundary e2e", () => { + let app: SpawnedApp | undefined; + let sameOriginUpstream: OpenAiUpstream | undefined; + let crossOriginUpstream: OpenAiUpstream | undefined; + let crossOriginDestination: OpenAiUpstream | undefined; + let etcdReachable = false; + + const post = (path: string) => + fetch(`${app!.proxyUrl}${path}/attempt`, { + method: "POST", + redirect: "manual", + headers: { + authorization: `Bearer ${CALLER_PLAINTEXT}`, + "content-type": "text/plain", + }, + body: "redirect boundary probe", + }); + + beforeAll(async () => { + const etcd = new EtcdClient(); + etcdReachable = await etcd.ping(); + if (!etcdReachable) return; + + sameOriginUpstream = await startOpenAiUpstream({ + status: 307, + responseHeaders: { location: SAME_ORIGIN_LOCATION }, + }); + crossOriginDestination = await startOpenAiUpstream(); + crossOriginUpstream = await startOpenAiUpstream({ + status: 308, + responseHeaders: { + location: `${crossOriginDestination.baseUrl}/outside-target/cross-origin`, + }, + }); + + app = await spawnApp(); + const seed = new SeedClient(etcd, app.etcdPrefix); + const providerKey = await seed.createProviderKey({ + display_name: "passthrough-redirect-boundary-pk", + secret: PROVIDER_SECRET, + api_base: sameOriginUpstream.baseUrl, + }); + await seed.createPassthroughRoute({ + name: "passthrough-redirect-same-origin", + path_prefix: SAME_ORIGIN_ROUTE, + target_url: `${sameOriginUpstream.baseUrl}/bounded`, + auth_mode: "gateway_key", + credential_mode: "inject", + provider_key_id: providerKey.id, + }); + await seed.createPassthroughRoute({ + name: "passthrough-redirect-cross-origin", + path_prefix: CROSS_ORIGIN_ROUTE, + target_url: `${crossOriginUpstream.baseUrl}/bounded`, + auth_mode: "gateway_key", + credential_mode: "inject", + provider_key_id: providerKey.id, + }); + await seed.createApiKey({ + key_hash: CALLER_KEY_HASH, + allowed_models: [], + allowed_routes: ["*"], + }); + + // The API key is seeded last, so a successful auth-only request proves + // both redirect routes have reached the gateway snapshot. + const proxy = new ProxyClient(app.proxyUrl, CALLER_PLAINTEXT); + await waitConfigPropagation(async () => (await proxy.listModels()).status === 200); + }, 120_000); + + afterAll(async () => { + await app?.exit(); + await sameOriginUpstream?.close(); + await crossOriginUpstream?.close(); + await crossOriginDestination?.close(); + }); + + test("relays a same-origin 307 outside target_url without requesting it", async (ctx) => { + if (!etcdReachable || !sameOriginUpstream) { + ctx.skip(); + return; + } + + const before = sameOriginUpstream.receivedRequests.length; + const response = await post(SAME_ORIGIN_ROUTE); + + expect(response.status).toBe(307); + expect(response.headers.get("location")).toBe(SAME_ORIGIN_LOCATION); + await response.text(); + + // The initial request stayed under the configured /bounded mount. A + // redirect-following client would make an additional /outside-target + // request on this same origin, escaping that mount unchecked. + const requests = sameOriginUpstream.receivedRequests.slice(before); + expect(requests).toHaveLength(1); + expect(requests[0]?.path).toBe("/bounded/attempt"); + expect(requests[0]?.headers.authorization).toBe(`Bearer ${PROVIDER_SECRET}`); + }); + + test("relays a cross-origin 308 without handing its credential to the target", async (ctx) => { + if (!etcdReachable || !crossOriginUpstream || !crossOriginDestination) { + ctx.skip(); + return; + } + + const sourceBefore = crossOriginUpstream.receivedRequests.length; + const destinationBefore = crossOriginDestination.receivedRequests.length; + const location = `${crossOriginDestination.baseUrl}/outside-target/cross-origin`; + const response = await post(CROSS_ORIGIN_ROUTE); + + expect(response.status).toBe(308); + expect(response.headers.get("location")).toBe(location); + await response.text(); + + const sourceRequests = crossOriginUpstream.receivedRequests.slice(sourceBefore); + expect(sourceRequests).toHaveLength(1); + expect(sourceRequests[0]?.path).toBe("/bounded/attempt"); + expect(sourceRequests[0]?.headers.authorization).toBe(`Bearer ${PROVIDER_SECRET}`); + expect(crossOriginDestination.receivedRequests.length).toBe(destinationBefore); + }); +}); diff --git a/tests/e2e/src/cases/passthrough-responses-media-guardrail-e2e.test.ts b/tests/e2e/src/cases/passthrough-responses-media-guardrail-e2e.test.ts new file mode 100644 index 000000000..50a5bb8f1 --- /dev/null +++ b/tests/e2e/src/cases/passthrough-responses-media-guardrail-e2e.test.ts @@ -0,0 +1,240 @@ +import { createHash } from "node:crypto"; +import { createServer, type Server } from "node:http"; +import { afterAll, beforeAll, describe, expect, test } from "vitest"; +import { + EtcdClient, + ProxyClient, + SeedClient, + pickFreePort, + spawnApp, + startOpenAiUpstream, + waitConfigPropagation, + type OpenAiUpstream, + type SpawnedApp, +} from "../harness/index.js"; + +// E2E (AISIX-Cloud#1262): a detected Responses passthrough may relay image +// bytes verbatim, but it must not copy them to an external output guardrail. +// This starts the real AISIX binary and etcd plus two real local HTTP peers: +// the passthrough upstream and an OpenAI Moderation-compatible guardrail +// endpoint. The latter records the body AISIX actually sends it. + +const CALLER = "sk-passthrough-responses-media"; +const CALLER_HASH = createHash("sha256").update(CALLER).digest("hex"); +const BUFFERED_MEDIA = "buffered-media-sentinel-not-for-guardrail"; +const BUFFERED_VISIBLE = "buffered-visible-text-for-guardrail"; +const STREAM_MEDIA = "stream-media-sentinel-not-for-guardrail"; +const STREAM_VISIBLE = "stream-visible-text-for-guardrail"; + +interface ModerationSink { + baseUrl: string; + inputs: string[]; + close(): Promise; +} + +async function startModerationSink(): Promise { + const inputs: string[] = []; + const server: Server = createServer((req, res) => { + let raw = ""; + req.on("data", (chunk: Buffer) => (raw += chunk.toString("utf8"))); + req.on("end", () => { + try { + const body = JSON.parse(raw) as { input?: unknown }; + if (typeof body.input === "string") inputs.push(body.input); + } catch { + // The test asserts only requests that follow the moderation wire + // contract. A malformed request still receives a well-formed answer + // so the gateway's own failure path remains observable in CI. + } + res.statusCode = 200; + res.setHeader("content-type", "application/json"); + res.end(JSON.stringify({ results: [{ flagged: false }] })); + }); + }); + const port = await pickFreePort(); + await new Promise((resolve, reject) => { + server.once("error", reject); + server.listen(port, "127.0.0.1", resolve); + }); + return { + baseUrl: `http://127.0.0.1:${port}`, + inputs, + async close() { + await new Promise((resolve, reject) => { + server.close((err) => (err ? reject(err) : resolve())); + }); + }, + }; +} + +const bufferedResponse = { + id: "resp_media_buffered", + object: "response", + output: [ + { + type: "image_generation_call", + id: "ig_media_buffered", + status: "completed", + result: BUFFERED_MEDIA, + }, + { + type: "message", + content: [{ type: "output_text", text: BUFFERED_VISIBLE }], + }, + ], +}; + +const streamedResponse = [ + `data: ${JSON.stringify({ + type: "response.image_generation_call.partial_image", + item_id: "ig_media_stream", + output_index: 0, + partial_image_index: 0, + partial_image_b64: STREAM_MEDIA, + })}\n\n`, + `data: ${JSON.stringify({ + type: "response.output_text.delta", + item_id: "msg_media_stream", + output_index: 1, + content_index: 0, + delta: STREAM_VISIBLE, + })}\n\n`, + `data: ${JSON.stringify({ + type: "response.completed", + response: { + id: "resp_media_stream", + object: "response", + status: "completed", + output: [ + { type: "image_generation_call", id: "ig_media_stream", result: STREAM_MEDIA }, + { type: "message", content: [{ type: "output_text", text: STREAM_VISIBLE }] }, + ], + }, + })}\n\n`, + "data: [DONE]\n\n", +]; + +describe("Responses passthrough keeps generated media out of output guardrails", () => { + let app: SpawnedApp | undefined; + let seed: SeedClient | undefined; + let bufferedUpstream: OpenAiUpstream | undefined; + let streamUpstream: OpenAiUpstream | undefined; + let moderation: ModerationSink | undefined; + let etcdReachable = false; + + beforeAll(async () => { + const etcd = new EtcdClient(); + etcdReachable = await etcd.ping(); + if (!etcdReachable) return; + + moderation = await startModerationSink(); + bufferedUpstream = await startOpenAiUpstream({ nonStreamBody: bufferedResponse }); + streamUpstream = await startOpenAiUpstream({ rawStreamFrames: streamedResponse }); + app = await spawnApp(); + seed = new SeedClient(etcd, app.etcdPrefix); + + const providerKey = await seed.createProviderKey({ + display_name: "passthrough-responses-media-pk", + secret: "sk-mock", + api_base: bufferedUpstream.baseUrl, + }); + await seed.createPassthroughRoute({ + name: "passthrough-responses-media-buffered", + path_prefix: "/responses-media-buffered", + target_url: bufferedUpstream.baseUrl, + provider_key_id: providerKey.id, + }); + await seed.createPassthroughRoute({ + name: "passthrough-responses-media-stream", + path_prefix: "/responses-media-stream", + target_url: streamUpstream.baseUrl, + provider_key_id: providerKey.id, + }); + await seed.createGuardrail({ + name: "passthrough-responses-media-output", + enabled: true, + hook_point: "output", + kind: "openai_moderation", + api_key: "sk-local-moderation", + endpoint: moderation.baseUrl, + output_fail_open: false, + }); + // Seeded last: the successful auth gate proves every resource above has + // crossed this app's etcd watch before either privacy assertion runs. + await seed.createApiKey({ + key_hash: CALLER_HASH, + allowed_models: [], + allowed_routes: ["*"], + }); + const proxy = new ProxyClient(app.proxyUrl, CALLER); + await waitConfigPropagation(async () => (await proxy.listModels()).status === 200); + }, 120_000); + + afterAll(async () => { + await app?.exit(); + await bufferedUpstream?.close(); + await streamUpstream?.close(); + await moderation?.close(); + }); + + const request = (route: string, stream: boolean) => + fetch(`${app!.proxyUrl}/${route}/v1/responses`, { + method: "POST", + headers: { + authorization: `Bearer ${CALLER}`, + "content-type": "application/json", + }, + body: JSON.stringify({ model: "gpt-4o-mini", input: "go", stream }), + }); + + const expectExternalGuardrailText = (inputs: string[], visible: string, media: string) => { + expect(inputs.length, "the external output guardrail was invoked").toBeGreaterThan(0); + expect(inputs.every((input) => !input.includes(media)), inputs.join("\n")).toBe(true); + expect(inputs.some((input) => input.includes(visible)), inputs.join("\n")).toBe(true); + const visibleOccurrences = inputs.reduce( + (count, input) => count + input.split(visible).length - 1, + 0, + ); + expect(visibleOccurrences, `external guardrail input: ${inputs.join("\n")}`).toBe(1); + }; + + test("buffered Responses media stays out of the external guardrail while output text is scanned", async (ctx) => { + if (!etcdReachable || !app || !bufferedUpstream || !moderation) { + ctx.skip(); + return; + } + const upstreamBefore = bufferedUpstream.receivedRequests.length; + const moderationBefore = moderation.inputs.length; + const response = await request("responses-media-buffered", false); + const body = await response.text(); + expect(response.status, body).toBe(200); + expect(body).toContain(BUFFERED_MEDIA); + expect(body).toContain(BUFFERED_VISIBLE); + expect(bufferedUpstream.receivedRequests.length).toBe(upstreamBefore + 1); + expectExternalGuardrailText( + moderation.inputs.slice(moderationBefore), + BUFFERED_VISIBLE, + BUFFERED_MEDIA, + ); + }); + + test("streamed Responses partial-image media stays out of the external guardrail while text is scanned", async (ctx) => { + if (!etcdReachable || !app || !streamUpstream || !moderation) { + ctx.skip(); + return; + } + const upstreamBefore = streamUpstream.receivedRequests.length; + const moderationBefore = moderation.inputs.length; + const response = await request("responses-media-stream", true); + const body = await response.text(); + expect(response.status, body).toBe(200); + expect(body).toContain(STREAM_MEDIA); + expect(body).toContain(STREAM_VISIBLE); + expect(streamUpstream.receivedRequests.length).toBe(upstreamBefore + 1); + expectExternalGuardrailText( + moderation.inputs.slice(moderationBefore), + STREAM_VISIBLE, + STREAM_MEDIA, + ); + }); +}); diff --git a/tests/e2e/src/cases/passthrough-route-e2e.test.ts b/tests/e2e/src/cases/passthrough-route-e2e.test.ts index 33b0141d0..0c35a7b86 100644 --- a/tests/e2e/src/cases/passthrough-route-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-route-e2e.test.ts @@ -1,8 +1,10 @@ import { createHash } from "node:crypto"; +import { connect } from "node:net"; import { afterAll, beforeAll, describe, expect, test } from "vitest"; import { harnessRequest } from "../harness/http.js"; import { EtcdClient, + ProxyClient, SeedClient, spawnApp, startOpenAiUpstream, @@ -46,6 +48,65 @@ const CALLER_PLAINTEXT = "sk-ptr-e2e-caller"; const CALLER_KEY_HASH = createHash("sha256") .update(CALLER_PLAINTEXT) .digest("hex"); +const STREAM_LIMITED_PLAINTEXT = "sk-ptr-e2e-stream-limit"; +const STREAM_LIMITED_KEY_HASH = createHash("sha256") + .update(STREAM_LIMITED_PLAINTEXT) + .digest("hex"); +const STREAM_TIMEOUT_PLAINTEXT = "sk-ptr-e2e-stream-timeout"; +const STREAM_TIMEOUT_KEY_HASH = createHash("sha256") + .update(STREAM_TIMEOUT_PLAINTEXT) + .digest("hex"); +const STREAM_RAW_CHUNK_PLAINTEXT = "sk-ptr-e2e-stream-raw-chunk"; +const STREAM_RAW_CHUNK_KEY_HASH = createHash("sha256") + .update(STREAM_RAW_CHUNK_PLAINTEXT) + .digest("hex"); + +/** + * Send an HTTP/1.1 request line without a URL client parsing or normalizing + * its path. This pins the route boundary against the bytes a proxy receives. + */ +function rawHttpStatus(proxyUrl: string, path: string): Promise { + const target = new URL(proxyUrl); + const port = Number(target.port || "80"); + + return new Promise((resolve, reject) => { + const socket = connect({ host: target.hostname, port }); + let response = ""; + let settled = false; + const timeout = setTimeout( + () => finish(new Error(`timed out waiting for raw HTTP response to ${path}`)), + 5_000, + ); + + function finish(result: number | Error): void { + if (settled) return; + settled = true; + clearTimeout(timeout); + socket.destroy(); + if (result instanceof Error) reject(result); + else resolve(result); + } + + socket.once("connect", () => { + socket.write( + [ + `GET ${path} HTTP/1.1`, + `Host: ${target.host}`, + `Authorization: Bearer ${CALLER_PLAINTEXT}`, + "Connection: close", + "", + "", + ].join("\r\n"), + ); + }); + socket.on("data", (chunk: Buffer) => { + response += chunk.toString("utf8"); + const status = /^HTTP\/1\.[01] (\d{3})\b/.exec(response); + if (status) finish(Number(status[1])); + }); + socket.once("error", (error) => finish(error)); + }); +} describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed paths", () => { let app: SpawnedApp | undefined; @@ -221,6 +282,182 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat expect(await res.text()).toBe(""); }); + test("route boundary traversal and query conflicts never reach the configured upstream", async (ctx) => { + if (!etcdReachable || !app || !seed) { + ctx.skip(); + return; + } + + const upstream = await startOpenAiUpstream({ + nonStreamBody: { object: "safe-route-target" }, + }); + upstreams.push(upstream); + + const pk = await seed.createProviderKey({ + display_name: "ptr-boundary-pk", + secret: "sk-mock", + api_base: "http://unused-on-routes", + }); + await seed.createPassthroughRoute({ + name: "ptr-boundary", + path_prefix: "/ptr-boundary", + target_url: `${upstream.baseUrl}/provider/v1?tenant=operator`, + provider_key_id: pk.id, + }); + await seed.createPassthroughRoute({ + name: "ptr-form-key-boundary", + path_prefix: "/ptr-form-key-boundary", + target_url: `${upstream.baseUrl}/provider/v1?tenant_id=operator`, + provider_key_id: pk.id, + }); + + const headers = { authorization: `Bearer ${CALLER_PLAINTEXT}` }; + await waitConfigPropagation(async () => { + try { + const ready = await fetch(`${app!.proxyUrl}/ptr-form-key-boundary/models`, { + headers, + }); + await ready.text(); + return ready.status === 200; + } catch { + return false; + } + }); + + expect(upstream.receivedRequests.at(-1)?.path).toBe( + "/provider/v1/models?tenant_id=operator", + ); + const allowedQuery = await harnessRequest( + `${app.proxyUrl}/ptr-boundary/models?limit=3`, + { headers }, + ); + expect(allowedQuery.statusCode).toBe(200); + await allowedQuery.body.text(); + expect(upstream.receivedRequests.at(-1)?.path).toBe( + "/provider/v1/models?tenant=operator&limit=3", + ); + + const baseline = upstream.receivedRequests.length; + // URL clients may normalize these before the gateway sees the request. + // Write the request line ourselves to preserve the traversal bytes. + for (const remainder of [ + "../models", + "%2e%2e/models", + "%2E%2E/models", + "%2E./models", + "..\\models", + ]) { + const status = await rawHttpStatus( + app.proxyUrl, + `/ptr-boundary/${remainder}`, + ); + expect(status, remainder).toBe(400); + expect(upstream.receivedRequests, remainder).toHaveLength(baseline); + } + + for (const remainder of [ + "%252e%252e%252fmodels", + "..;ignored/models", + "%2e%2e%3bignored/models", + "%252e%252e%253bignored/models", + "%2e%2e%3bignored/%2e%2e%3bignored/admin", + ]) { + const rejected = await harnessRequest( + `${app.proxyUrl}/ptr-boundary/${remainder}`, + { headers }, + ); + expect(rejected.statusCode, remainder).toBe(400); + await rejected.body.text(); + expect(upstream.receivedRequests, remainder).toHaveLength(baseline); + } + + const conflictingQuery = await harnessRequest( + `${app.proxyUrl}/ptr-boundary/models?tenant=caller`, + { headers }, + ); + expect(conflictingQuery.statusCode).toBe(400); + await conflictingQuery.body.text(); + expect(upstream.receivedRequests).toHaveLength(baseline); + + const nestedConflictingQuery = await harnessRequest( + `${app.proxyUrl}/ptr-boundary/models?%2574enant=caller`, + { headers }, + ); + expect(nestedConflictingQuery.statusCode).toBe(400); + await nestedConflictingQuery.body.text(); + expect(upstream.receivedRequests).toHaveLength(baseline); + + const encodedQueryDelimiter = await harnessRequest( + `${app.proxyUrl}/ptr-boundary/models?safe=1%26tenant%3Dcaller`, + { headers }, + ); + expect(encodedQueryDelimiter.statusCode).toBe(400); + await encodedQueryDelimiter.body.text(); + expect(upstream.receivedRequests).toHaveLength(baseline); + + const nestedEncodedQueryDelimiter = await harnessRequest( + `${app.proxyUrl}/ptr-boundary/models?safe=1%2526tenant%253Dcaller`, + { headers }, + ); + expect(nestedEncodedQueryDelimiter.statusCode).toBe(400); + await nestedEncodedQueryDelimiter.body.text(); + expect(upstream.receivedRequests).toHaveLength(baseline); + + for (const query of [ + "+tenant=caller", + "%20tenant=caller", + "%2520tenant=caller", + "tenant%00suffix=caller", + "tenant%2500suffix=caller", + ]) { + const phpFormKeyConflict = await harnessRequest( + `${app.proxyUrl}/ptr-boundary/models?${query}`, + { headers }, + ); + expect(phpFormKeyConflict.statusCode, query).toBe(400); + await phpFormKeyConflict.body.text(); + expect(upstream.receivedRequests, query).toHaveLength(baseline); + } + + for (const query of [ + "safe=1;tenant=caller", + "safe=1%3Btenant%3Dcaller", + "safe=1%253Btenant%253Dcaller", + ]) { + const semicolonQueryDelimiter = await harnessRequest( + `${app.proxyUrl}/ptr-boundary/models?${query}`, + { headers }, + ); + expect(semicolonQueryDelimiter.statusCode, query).toBe(400); + await semicolonQueryDelimiter.body.text(); + expect(upstream.receivedRequests, query).toHaveLength(baseline); + } + + for (const query of [ + "tenant.id=caller", + "tenant%2Eid=caller", + "tenant%252Eid=caller", + "tenant+id=caller", + "tenant%20id=caller", + ]) { + const normalizedKeyConflict = await harnessRequest( + `${app.proxyUrl}/ptr-form-key-boundary/models?${query}`, + { headers }, + ); + expect(normalizedKeyConflict.statusCode, query).toBe(400); + await normalizedKeyConflict.body.text(); + expect(upstream.receivedRequests, query).toHaveLength(baseline); + } + + const bracketedKeyConflict = await harnessRequest( + `${app.proxyUrl}/ptr-boundary/models?tenant%5Brole%5D=caller`, + { headers }, + ); + expect(bracketedKeyConflict.statusCode).toBe(400); + await bracketedKeyConflict.body.text(); + expect(upstream.receivedRequests).toHaveLength(baseline); + }); + test("forward-proxy BYO: host match beats typed routes; Authorization forwarded verbatim", async (ctx) => { if (!etcdReachable || !app || !seed) { ctx.skip(); @@ -431,15 +668,22 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat return; } + const streamEvents = [ + JSON.stringify({ choices: [{ delta: { content: "hel" } }] }), + JSON.stringify({ choices: [{ delta: { content: "lo" } }] }), + JSON.stringify({ + choices: [], + usage: { prompt_tokens: 5, completion_tokens: 2 }, + }), + "[DONE]", + ]; const upstream = await startOpenAiUpstream({ - streamEvents: [ - JSON.stringify({ choices: [{ delta: { content: "hel" } }] }), - JSON.stringify({ choices: [{ delta: { content: "lo" } }] }), - JSON.stringify({ - choices: [], - usage: { prompt_tokens: 5, completion_tokens: 2 }, - }), - "[DONE]", + // The propagation probe below reaches the same passthrough route but + // receives this complete unary response, leaving the first stream for + // the journey asserted by the test. + scriptedResponses: [ + { nonStreamBody: { ready: true } }, + { streamEvents }, ], }); upstreams.push(upstream); @@ -473,12 +717,12 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat await waitConfigPropagation(async () => { try { - const r = await call(); - const ok = - r.status === 200 && - (r.headers.get("content-type") ?? "").includes("text/event-stream"); - await r.text(); - return ok; + const probe = await call(); + const ready = + probe.status === 200 && + !(probe.headers.get("content-type") ?? "").includes("text/event-stream"); + await probe.text(); + return ready; } catch { return false; } @@ -494,6 +738,266 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat expect(text).toContain("[DONE]"); }); + test("an SSE passthrough holds its concurrency slot until the body ends or cancels", async (ctx) => { + if (!etcdReachable || !app || !seed) { + ctx.skip(); + return; + } + + // Headers arrive immediately but the body remains open long enough to + // make the second request observe the in-flight stream, rather than a + // short unary exchange that happened to have already finished. + const streamEvents = [ + JSON.stringify({ choices: [{ delta: { content: "one" } }] }), + JSON.stringify({ choices: [{ delta: { content: "two" } }] }), + "[DONE]", + ]; + const upstream = await startOpenAiUpstream({ + scriptedResponses: [ + // A complete probe establishes propagation on this route without + // retaining a streaming concurrency slot. The first real stream + // then stalls until the client cancels it; the second ends naturally. + { nonStreamBody: { ready: true } }, + { streamEvents, firstEventDelayMs: 10_000 }, + { streamEvents, firstEventDelayMs: 25, eventDelayMs: 25 }, + { streamEvents }, + ], + }); + upstreams.push(upstream); + + const pk = await seed.createProviderKey({ + display_name: "ptr-sse-concurrency-pk", + secret: "sk-mock", + api_base: "http://unused-on-routes", + }); + await seed.createPassthroughRoute({ + name: "ptr-sse-concurrency", + path_prefix: "/sse-concurrency", + target_url: upstream.baseUrl, + provider_key_id: pk.id, + }); + // Write the constrained principal last. A completed unary probe proves + // the preceding route has reached the same snapshot and released its + // slot before this test starts the first real stream. + await seed.createApiKey({ + key_hash: STREAM_LIMITED_KEY_HASH, + allowed_models: ["*"], + allowed_routes: ["ptr-sse-concurrency"], + rate_limit: { concurrency: 1 }, + }); + + const headers = { + authorization: `Bearer ${STREAM_LIMITED_PLAINTEXT}`, + "content-type": "application/json", + }; + const call = () => + fetch(`${app!.proxyUrl}/sse-concurrency/chat/completions`, { + method: "POST", + headers, + body: JSON.stringify({ + model: "gpt-4o", + messages: [{ role: "user", content: "hi" }], + stream: true, + }), + }); + + await waitConfigPropagation(async () => { + try { + const probe = await call(); + const ready = + probe.status === 200 && + !(probe.headers.get("content-type") ?? "").includes("text/event-stream"); + await probe.text(); + return ready; + } catch { + return false; + } + }); + + // Fetch resolves as soon as the upstream headers are relayed. Keep this + // body unread while issuing the second request: it is the real caller + // journey that used to release the slot at handler return. + const first = await call(); + expect(first.status).toBe(200); + const upstreamCallsWhileStreaming = upstream.receivedRequests.length; + expect(upstreamCallsWhileStreaming).toBe(2); + + const second = await call(); + expect(second.status).toBe(429); + expect(second.headers.get("x-ratelimit-scope")).toBe("concurrency"); + await second.text(); + expect(upstream.receivedRequests).toHaveLength(upstreamCallsWhileStreaming); + + // A cancelled client must release the same hold. The admitted follow-up + // is deliberately a separate, naturally ending stream. + expect(first.body).not.toBeNull(); + await first.body!.cancel(); + let naturallyEnding: Response | undefined; + await waitConfigPropagation(async () => { + const afterCancel = await call(); + const admitted = afterCancel.status === 200; + if (admitted) naturallyEnding = afterCancel; + else await afterCancel.text(); + return admitted; + }, 3_000); + expect(naturallyEnding).toBeDefined(); + await naturallyEnding!.text(); + + const afterEnd = await call(); + expect(afterEnd.status).toBe(200); + await afterEnd.text(); + + expect(upstream.receivedRequests).toHaveLength(upstreamCallsWhileStreaming + 2); + }); + + test("a silent SSE gap terminates and releases the concurrency slot", async (ctx) => { + if (!etcdReachable || !app || !seed) { + ctx.skip(); + return; + } + + const upstream = await startOpenAiUpstream({ + scriptedResponses: [ + { + streamEvents: [ + JSON.stringify({ choices: [{ delta: { content: "first" } }] }), + JSON.stringify({ choices: [{ delta: { content: "late" } }] }), + "[DONE]", + ], + // The first event is immediate; the next one violates the route's + // 100 ms per-read budget after headers have already been relayed. + eventDelayMs: 1_500, + }, + { nonStreamBody: { recovered: true } }, + ], + }); + upstreams.push(upstream); + + const pk = await seed.createProviderKey({ + display_name: "ptr-sse-timeout-pk", + secret: "sk-mock", + api_base: "http://unused-on-routes", + }); + await seed.createPassthroughRoute({ + name: "ptr-sse-timeout", + path_prefix: "/sse-timeout", + target_url: upstream.baseUrl, + provider_key_id: pk.id, + timeout_ms: 100, + }); + await seed.createApiKey({ + key_hash: STREAM_TIMEOUT_KEY_HASH, + allowed_models: ["*"], + allowed_routes: ["ptr-sse-timeout"], + rate_limit: { concurrency: 1 }, + }); + + const headers = { + authorization: `Bearer ${STREAM_TIMEOUT_PLAINTEXT}`, + "content-type": "application/json", + }; + const call = () => + fetch(`${app!.proxyUrl}/sse-timeout/chat/completions`, { + method: "POST", + headers, + body: JSON.stringify({ + model: "gpt-4o", + messages: [{ role: "user", content: "hi" }], + stream: true, + }), + }); + + await waitConfigPropagation(async () => { + try { + const probe = await fetch(`${app!.proxyUrl}/v1/models`, { + headers: { authorization: `Bearer ${STREAM_TIMEOUT_PLAINTEXT}` }, + }); + await probe.text(); + return probe.status === 200; + } catch { + return false; + } + }); + + const started = Date.now(); + const stalled = await call(); + expect(stalled.status).toBe(200); + const body = await stalled.text(); + const elapsed = Date.now() - started; + expect(body).toContain("first"); + expect(body).not.toContain("late"); + expect(elapsed).toBeLessThan(1_000); + + // A new request can enter after the timeout. A leaked concurrency hold + // would instead be a gateway 429 and would never consume step three. + const recovered = await call(); + expect(recovered.status).toBe(200); + expect(await recovered.json()).toEqual({ recovered: true }); + }, 10_000); + + test("an SSE frame may span raw upstream chunks without timing out", async (ctx) => { + if (!etcdReachable || !app || !seed) { + ctx.skip(); + return; + } + + // Each raw-body gap is below the 100 ms route timeout, but assembling this + // one SSE frame takes longer than 100 ms. A frame-based timeout would fail + // before the third write; the relay's raw-chunk timeout must not. + const upstream = await startOpenAiUpstream({ + rawBodyChunks: [ + 'data: {"choices":[{"delta":{"content":"raw-', + 'chunk-', + 'timeout"}}]}\n\n', + "data: [DONE]\n\n", + ], + rawContentType: "text/event-stream", + eventDelayMs: 60, + }); + upstreams.push(upstream); + + const pk = await seed.createProviderKey({ + display_name: "ptr-sse-raw-chunk-timeout-pk", + secret: "sk-mock", + api_base: "http://unused-on-routes", + }); + await seed.createPassthroughRoute({ + name: "ptr-sse-raw-chunk-timeout", + path_prefix: "/sse-raw-chunk-timeout", + target_url: upstream.baseUrl, + provider_key_id: pk.id, + timeout_ms: 100, + }); + await seed.createApiKey({ + key_hash: STREAM_RAW_CHUNK_KEY_HASH, + allowed_models: ["*"], + allowed_routes: ["ptr-sse-raw-chunk-timeout"], + }); + + const call = () => + fetch(`${app!.proxyUrl}/sse-raw-chunk-timeout/chat/completions`, { + method: "POST", + headers: { + authorization: `Bearer ${STREAM_RAW_CHUNK_PLAINTEXT}`, + "content-type": "application/json", + }, + body: JSON.stringify({ + model: "gpt-4o", + messages: [{ role: "user", content: "hi" }], + stream: true, + }), + }); + + const probe = new ProxyClient(app.proxyUrl, STREAM_RAW_CHUNK_PLAINTEXT); + await waitConfigPropagation(async () => (await probe.listModels()).status === 200); + + const response = await call(); + const body = await response.text(); + expect(response.status, body).toBe(200); + expect(body).toContain("raw-chunk-timeout"); + expect(body).toContain("[DONE]"); + }, 10_000); + test("envelope auto-detection: usage follows the request body, never the config", async (ctx) => { if (!etcdReachable || !app || !seed) { ctx.skip(); diff --git a/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts b/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts index d13a3c824..7b859285c 100644 --- a/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts @@ -1,23 +1,29 @@ import { createHash } from "node:crypto"; +import { request, type IncomingMessage } from "node:http"; +import { gzipSync } from "node:zlib"; import { afterAll, beforeAll, describe, expect, test } from "vitest"; import { EtcdClient, ProxyClient, SeedClient, spawnApp, + startMockSls, startOpenAiUpstream, waitConfigPropagation, + waitForSlsLog, + type MockSls, type OpenAiUpstream, type SpawnedApp, } from "../harness/index.js"; +import { startMockOtlp, type MockOtlp } from "../harness/otlp-mock.js"; // E2E: what a passthrough route's guardrails read, per detected envelope. // // Output: a streamed reply is scanned for its generated text and tool-call // arguments on every envelope — Anthropic Messages text and tool input // deltas (carried on the chat envelope), chat tool-call arguments, and -// Responses function-call argument deltas. Generated reasoning is not -// scanned. The same extraction feeds the hold-back cap (#513): a stream +// Responses function-call argument/refusal deltas. Generated reasoning is +// not scanned. The same extraction feeds the hold-back cap (#513): a stream // whose frames outweigh `max_buffer_bytes` while its content does not is // released. // @@ -28,9 +34,113 @@ import { const CALLER = "sk-pt-scan-coverage"; const CALLER_HASH = createHash("sha256").update(CALLER).digest("hex"); +const ERROR_STREAM_CALLER = "sk-pt-scan-error-stream"; +const ERROR_STREAM_CALLER_HASH = createHash("sha256").update(ERROR_STREAM_CALLER).digest("hex"); +const SUPPLEMENTAL_CALLER = "sk-pt-scan-supplemental"; +const SUPPLEMENTAL_CALLER_HASH = createHash("sha256").update(SUPPLEMENTAL_CALLER).digest("hex"); const OUT_LIT = "outputleakliteral"; const IN_LIT = "inputleakliteral"; +const ESCAPED_BLOCK = "BLOCKME"; +const ESCAPED_CJK = "中文"; +const ESCAPED_BLOCK_JSON = String.raw`{"state":"\u0042LOCKME","state":"clean"}`; +const ESCAPED_CJK_JSON = String.raw`{"query":"\u4e2d\u6587"}`; +const SAFE_ESCAPED_JSON = String.raw`{"state":"\u0063lean","state":"safe"}`; +const RAW_BLOCK_SSE = String.raw`"\u0042LOCKME"`; +const RAW_SAFE_SSE = String.raw`"\u0063lean"`; +const RAW_PREFIX_SSE = `"FOR"`; +const RAW_UNEVALUABLE_SSE = String.raw`{"state":"safe"}`; +const RAW_SUFFIX_SSE = `"BIDDEN"`; +const RAW_HELD_BLOCK_SSE = `"FORBIDDEN"`; +const UPSTREAM_502_HTML = `${OUT_LIT}`; +const UPSTREAM_502_SSE = `event: upstream_error\ndata: {"message":"${OUT_LIT}"}\n\n`; +const UPSTREAM_502_GZIP = gzipSync(Buffer.from(UPSTREAM_502_SSE)); +const UPSTREAM_200_GZIP_JSON = gzipSync( + Buffer.from(`{"choices":[{"text":"${ESCAPED_BLOCK}"}]}`), +); +const UPSTREAM_200_GZIP_SSE = gzipSync( + Buffer.from( + `data: {"id":"c","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"content":"${ESCAPED_BLOCK}"}}]}\n\n`, + ), +); +const MALFORMED_SUPPLEMENTAL_SSE = + 'event: content_block_start\ndata: {"type":"content_block_start","index":0,"content_block":{"type":"tool_use","input":"clean","name":1}}\n\n'; +const deepEscapedBlockJSON = (depth: number) => + `${'{"v":'.repeat(depth)}"${String.raw`\u0042LOCKME`}"${'}'.repeat(depth)}`; +const deepLiteralBlockJSON = (depth: number) => + `${'{"v":'.repeat(depth)}"${ESCAPED_BLOCK}"${'}'.repeat(depth)}`; +const ANTHROPIC_TOOL_RESULT_PREFIX = '[{"type":"tool_result","tool_use_id":"t","content":'; +const ANTHROPIC_TOOL_RESULT_TEXT = `[{"type":"text","text":"${ESCAPED_BLOCK}"}]`; +const ANTHROPIC_TOOL_RESULT_SUFFIX = "}]"; +const deeplyNestedAnthropicToolResultRequest = (depth: number, model: string) => { + const content = [ + ANTHROPIC_TOOL_RESULT_PREFIX.repeat(depth), + ANTHROPIC_TOOL_RESULT_TEXT, + ANTHROPIC_TOOL_RESULT_SUFFIX.repeat(depth), + ].join(""); + return `{"model":"${model}","max_tokens":64,"messages":[{"role":"user","content":${content}}]}`; +}; +// Above serde_json's default recursion limit. It remains valid JSON and the +// provider receives it verbatim, so Raw guardrails must still decode the leaf. +const DEEP_ESCAPED_BLOCK_JSON = deepEscapedBlockJSON(160); +// Mirrors `json_splice::MAX_JSON_DEPTH`: one deeper is unscannable and must +// use the resolved guardrail failure policy instead of falling back to escapes. +const JSON_DEPTH_CAP = 4_096; +const NESTED_TOOL_RESULT_WORK_SAFE_DEPTH = 32; +const OVER_DEPTH_ESCAPED_BLOCK_JSON = deepEscapedBlockJSON(JSON_DEPTH_CAP + 1); +const OVER_DEPTH_OPAQUE_BLOCK_JSON = deepLiteralBlockJSON(JSON_DEPTH_CAP + 1); +const AT_DEPTH_ANTHROPIC_TOOL_RESULT_INPUT = deeplyNestedAnthropicToolResultRequest( + NESTED_TOOL_RESULT_WORK_SAFE_DEPTH, + "nested-tool-result-within-work-cap", +); +const OVER_DEPTH_ANTHROPIC_TOOL_RESULT_INPUT = deeplyNestedAnthropicToolResultRequest( + JSON_DEPTH_CAP, + "nested-tool-result-work-cap-fail-closed", +); +const OVER_DEPTH_ANTHROPIC_TOOL_RESULT_FAIL_OPEN_INPUT = deeplyNestedAnthropicToolResultRequest( + JSON_DEPTH_CAP, + "nested-tool-result-work-cap-fail-open", +); +const OVER_DEPTH_CHAT_INPUT = `{"model":"gpt-4o-mini","messages":[{"role":"user","content":"go","metadata":${OVER_DEPTH_ESCAPED_BLOCK_JSON}}]}`; +const OVER_DEPTH_RESPONSES_INPUT = `{"model":"gpt-4o-mini","input":[{"role":"user","metadata":${OVER_DEPTH_ESCAPED_BLOCK_JSON},"content":[{"type":"input_text","text":"go"}]}]}`; +const OVER_DEPTH_CHAT_OPAQUE_INPUT = `{"model":"gpt-4o-mini","messages":[{"role":"user","content":[{"type":"image","source":{"data":"image","metadata":${OVER_DEPTH_OPAQUE_BLOCK_JSON}}},{"type":"text","text":"go"}]}]}`; +const OVER_DEPTH_RESPONSES_OPAQUE_INPUT = `{"model":"gpt-4o-mini","input":[{"role":"user","content":[{"type":"input_image","image_url":{"url":"https://example.invalid/image","metadata":${OVER_DEPTH_OPAQUE_BLOCK_JSON}}},{"type":"input_text","text":"go"}]}]}`; +const OVER_DEPTH_ANTHROPIC_TOOL_OUTPUT = `{"type":"message","content":[{"type":"tool_use","id":"tool_1","name":"lookup","input":${OVER_DEPTH_ESCAPED_BLOCK_JSON}}]}`; const CAP = 1_000; +const RESPONSES_SNAPSHOT_TEXT = "x".repeat(300); +const SPLIT_BLOCK = "FORBIDDEN"; +const SPLIT_BLOCK_REGEX = String.raw`FOR\s*BIDDEN`; +const KNOWN_CHAT_OUTPUT = String.raw`{"model":"routing-only","choices":[{"message":{"content":"\u0042LOCKME","metadata":{"note":"${OUT_LIT}"}}}],"choices":[{"message":{"content":"clean"}}]}`; +const KNOWN_RESPONSES_OUTPUT = String.raw`{"output":[{"type":"message","content":[{"type":"output_text","text":"\u0042LOCKME","metadata":{"note":"${OUT_LIT}"}}]}],"output":[{"type":"message","content":[{"type":"output_text","text":"clean"}]}]}`; +const UNKNOWN_TYPED_BUFFERED_OUTPUT = String.raw`{"object":"chat.completion","result":{"text":"\u0042LOCKME"}}`; +const RESPONSES_REFUSAL_OUTPUT = JSON.stringify({ + output: [ + { + type: "message", + content: [{ type: "refusal", refusal: OUT_LIT, metadata: { note: ESCAPED_BLOCK } }], + }, + ], +}); +const KNOWN_CHAT_STREAM = `${String.raw`data: {"choices":[{"index":0,"delta":{"content":"\u0042LOCKME"}},{"index":1,"delta":{"content":"clean"}}]}`}\n\n`; +const KNOWN_RESPONSES_STREAM = `${String.raw`data: {"type":"response.output_text.delta","item_id":"known","output_index":0,"content_index":0,"delta":"\u0042LOCKME","delta":"clean","metadata":{"note":"${OUT_LIT}"}}`}\n\n`; +const SPLIT_RESPONSES_STREAM = [ + `data: {"type":"response.output_text.delta","item_id":"same","output_index":0,"content_index":0,"delta":"FOR"}\n\n`, + `data: {"type":"response.output_text.delta","item_id":"same","output_index":0,"content_index":0,"delta":"noise","delta":"BIDDEN"}\n\n`, + `data: {"type":"response.output_item.done","output_index":0,"item":{"id":"same","type":"message"}}\n\n`, + "data: [DONE]\n\n", +]; +const DISTINCT_ITEMS_RESPONSES_STREAM = [ + `data: {"type":"response.output_text.delta","item_id":"first","output_index":0,"content_index":0,"delta":"FOR"}\n\n`, + `data: {"type":"response.output_text.delta","item_id":"second","output_index":1,"content_index":0,"delta":"BIDDEN"}\n\n`, + "data: [DONE]\n\n", +]; +const RESPONSES_REASONING_STREAM = [ + `data: ${JSON.stringify({ type: "response.reasoning_text.done", text: OUT_LIT })}\n\n`, + `data: ${JSON.stringify({ type: "response.content_part.done", item_id: "reasoning", output_index: 0, content_index: 0, part: { type: "reasoning_text", text: OUT_LIT } })}\n\n`, + `data: ${JSON.stringify({ type: "response.output_item.done", output_index: 0, item: { id: "reasoning", type: "reasoning", summary: [{ type: "summary_text", text: OUT_LIT }] } })}\n\n`, + `data: ${JSON.stringify({ type: "response.output_text.delta", item_id: "message", output_index: 1, content_index: 0, delta: "clean" })}\n\n`, + `data: ${JSON.stringify({ type: "response.completed", response: { output: [{ type: "reasoning", summary: [{ type: "summary_text", text: OUT_LIT }] }, { type: "message", content: [{ type: "output_text", text: "clean" }] }] } })}\n\n`, + "data: [DONE]\n\n", +]; const anthropicEvents = (blocks: Array>) => [ JSON.stringify({ @@ -71,6 +181,59 @@ const STREAMS: Record = { JSON.stringify({ type: "response.function_call_arguments.delta", item_id: "fc", output_index: 0, delta: `{"q":"${OUT_LIT}"}` }), JSON.stringify({ type: "response.completed", response: { id: "r", status: "completed", output: [] } }), ], + "responses-refusal": [ + JSON.stringify({ type: "response.refusal.delta", item_id: "msg_refusal", output_index: 0, content_index: 0, delta: OUT_LIT }), + JSON.stringify({ type: "response.refusal.done", item_id: "msg_refusal", output_index: 0, content_index: 0, refusal: OUT_LIT }), + JSON.stringify({ type: "response.completed", response: { id: "r_refusal", status: "completed", output: [{ type: "message", content: [{ type: "refusal", refusal: OUT_LIT }] }] } }), + ], + "responses-refusal-cap": [ + JSON.stringify({ type: "response.refusal.delta", item_id: "msg_refusal_cap", output_index: 0, content_index: 0, delta: "x".repeat(CAP + 1) }), + JSON.stringify({ type: "response.completed", response: { id: "r_refusal_cap", status: "completed", output: [{ type: "message", content: [{ type: "refusal", refusal: "x".repeat(CAP + 1) }] }] } }), + ], + "responses-refusal-done-cap": [ + JSON.stringify({ type: "response.refusal.done", item_id: "msg_refusal_done_cap", output_index: 0, content_index: 0, refusal: "x".repeat(CAP + 1) }), + ], + "responses-repeated-snapshots-under-cap": [ + JSON.stringify({ type: "response.output_text.delta", item_id: "snapshot", output_index: 0, content_index: 0, delta: RESPONSES_SNAPSHOT_TEXT }), + JSON.stringify({ type: "response.output_text.done", item_id: "snapshot", output_index: 0, content_index: 0, text: RESPONSES_SNAPSHOT_TEXT }), + JSON.stringify({ type: "response.content_part.done", item_id: "snapshot", output_index: 0, content_index: 0, part: { type: "output_text", text: RESPONSES_SNAPSHOT_TEXT } }), + JSON.stringify({ type: "response.output_item.done", item_id: "snapshot", output_index: 0, item: { id: "snapshot", type: "message", content: [{ type: "output_text", text: RESPONSES_SNAPSHOT_TEXT }] } }), + JSON.stringify({ type: "response.completed", response: { output: [{ id: "snapshot", type: "message", content: [{ type: "output_text", text: RESPONSES_SNAPSHOT_TEXT }] }] } }), + ], + "responses-output-text-done": [ + JSON.stringify({ type: "response.output_text.done", item_id: "text_done", output_index: 0, content_index: 0, text: OUT_LIT }), + ], + "responses-refusal-done-only": [ + JSON.stringify({ type: "response.refusal.done", item_id: "refusal_done", output_index: 0, content_index: 0, refusal: OUT_LIT }), + ], + "responses-function-done": [ + JSON.stringify({ type: "response.function_call_arguments.done", item_id: "function_done", output_index: 0, name: "lookup", arguments: JSON.stringify({ query: OUT_LIT }) }), + ], + "responses-mcp-done": [ + JSON.stringify({ type: "response.mcp_call_arguments.done", item_id: "mcp_done", output_index: 0, arguments: JSON.stringify({ query: OUT_LIT }) }), + ], + "responses-custom-done": [ + JSON.stringify({ type: "response.custom_tool_call_input.done", item_id: "custom_done", output_index: 0, input: OUT_LIT }), + ], + "responses-content-part-added": [ + JSON.stringify({ type: "response.content_part.added", item_id: "part_added", output_index: 0, content_index: 0, part: { type: "refusal", refusal: OUT_LIT } }), + ], + "responses-content-part-done": [ + JSON.stringify({ type: "response.content_part.done", item_id: "part_done", output_index: 0, content_index: 0, part: { type: "refusal", refusal: OUT_LIT } }), + ], + "responses-output-item-added": [ + JSON.stringify({ type: "response.output_item.added", output_index: 0, item: { id: "item_added", type: "message", content: [{ type: "refusal", refusal: OUT_LIT }] } }), + ], + "responses-output-item-done": [ + JSON.stringify({ type: "response.output_item.done", output_index: 0, item: { id: "item_done", type: "message", content: [{ type: "refusal", refusal: OUT_LIT }] } }), + ], + "responses-terminal-only": [ + JSON.stringify({ type: "response.completed", response: { id: "terminal_only", status: "completed", output: [{ type: "message", content: [{ type: "refusal", refusal: OUT_LIT }] }] } }), + ], + "responses-clean-delta-malicious-done": [ + JSON.stringify({ type: "response.output_text.delta", item_id: "mismatch", output_index: 0, content_index: 0, delta: "clean" }), + JSON.stringify({ type: "response.refusal.done", item_id: "mismatch", output_index: 0, content_index: 0, refusal: OUT_LIT }), + ], // ~600 bytes of text over 60 frames: the frames outweigh the cap, the // text they carry does not. "chat-many-frames": [ @@ -80,9 +243,32 @@ const STREAMS: Record = { ], }; +function openRawHttpRequest( + url: string, + headers: Record, + body: string, +): Promise { + return new Promise((resolve, reject) => { + const req = request(url, { method: "POST", headers }, resolve); + req.once("error", reject); + req.end(body); + }); +} + +function readRawHttpBody(response: IncomingMessage): Promise { + return new Promise((resolve, reject) => { + const chunks: Buffer[] = []; + response.on("data", (chunk: Buffer) => chunks.push(chunk)); + response.once("error", reject); + response.once("end", () => resolve(Buffer.concat(chunks))); + }); +} + describe("passthrough guardrail scan coverage", () => { let app: SpawnedApp | undefined; + let seed: SeedClient | undefined; const upstreams: Record = {}; + const otlps: MockOtlp[] = []; let etcdReachable = false; beforeAll(async () => { @@ -90,13 +276,127 @@ describe("passthrough guardrail scan coverage", () => { etcdReachable = await etcd.ping(); if (!etcdReachable) return; app = await spawnApp(); - const seed = new SeedClient(etcd, app.etcdPrefix); + seed = new SeedClient(etcd, app.etcdPrefix); for (const [name, streamEvents] of Object.entries(STREAMS)) { upstreams[name] = await startOpenAiUpstream({ streamEvents }); } upstreams.input = await startOpenAiUpstream({ nonStreamBody: { id: "c", object: "chat.completion", choices: [] }, }); + upstreams["html-502"] = await startOpenAiUpstream({ + status: 502, + rawErrorBody: UPSTREAM_502_HTML, + responseHeaders: { + "content-type": "text/html; charset=utf-8", + "x-upstream-error": "edge-502", + }, + }); + upstreams["sse-502"] = await startOpenAiUpstream({ + status: 502, + rawErrorBody: UPSTREAM_502_SSE, + responseHeaders: { + "content-type": "text/event-stream; charset=utf-8", + "cache-control": "no-cache", + "x-upstream-error": "edge-sse-502", + }, + }); + upstreams["gzip-sse-502"] = await startOpenAiUpstream({ + status: 502, + rawErrorBody: UPSTREAM_502_GZIP, + responseHeaders: { + "content-type": "text/event-stream; charset=utf-8", + "content-encoding": "gzip", + "content-length": String(UPSTREAM_502_GZIP.byteLength), + "x-upstream-error": "edge-gzip-sse-502", + }, + }); + upstreams["gzip-buffered-200"] = await startOpenAiUpstream({ + rawBody: UPSTREAM_200_GZIP_JSON, + rawContentType: "application/json", + responseHeaders: { "content-encoding": "gzip" }, + }); + upstreams["gzip-sse-200"] = await startOpenAiUpstream({ + rawBody: UPSTREAM_200_GZIP_SSE, + rawContentType: "text/event-stream; charset=utf-8", + responseHeaders: { "content-encoding": "gzip" }, + }); + upstreams["delayed-sse-502"] = await startOpenAiUpstream({ + status: 502, + rawErrorBodyChunks: [UPSTREAM_502_SSE, ""], + eventDelayMs: 1_500, + responseHeaders: { + "content-type": "text/event-stream; charset=utf-8", + "x-upstream-error": "edge-delayed-sse-502", + }, + }); + upstreams["event-streaming"] = await startOpenAiUpstream({ + rawBody: ESCAPED_BLOCK_JSON, + rawContentType: "text/event-streaming; charset=utf-8", + }); + upstreams["raw-output"] = await startOpenAiUpstream({ + rawBody: ESCAPED_BLOCK_JSON, + rawContentType: "application/json", + }); + upstreams["raw-stream"] = await startOpenAiUpstream({ + rawStreamFrames: [`data: ${RAW_BLOCK_SSE}\n\n`, "data: [DONE]\n\n"], + }); + upstreams["raw-deep-output"] = await startOpenAiUpstream({ + rawBody: DEEP_ESCAPED_BLOCK_JSON, + rawContentType: "application/json", + }); + upstreams["raw-over-depth-output"] = await startOpenAiUpstream({ + rawBody: OVER_DEPTH_ESCAPED_BLOCK_JSON, + rawContentType: "application/json", + }); + upstreams["completions-over-depth-output"] = await startOpenAiUpstream({ + rawBody: OVER_DEPTH_ESCAPED_BLOCK_JSON, + rawContentType: "application/json", + }); + upstreams["anthropic-over-depth-tool-output"] = await startOpenAiUpstream({ + rawBody: OVER_DEPTH_ANTHROPIC_TOOL_OUTPUT, + rawContentType: "application/json", + }); + upstreams["raw-deep-stream"] = await startOpenAiUpstream({ + rawStreamFrames: [`data: ${DEEP_ESCAPED_BLOCK_JSON}\n\n`, "data: [DONE]\n\n"], + }); + upstreams["raw-safe-output"] = await startOpenAiUpstream({ + rawBody: SAFE_ESCAPED_JSON, + rawContentType: "application/json", + }); + upstreams["raw-safe-stream"] = await startOpenAiUpstream({ + rawStreamFrames: [`data: ${RAW_SAFE_SSE}\n\n`, "data: [DONE]\n\n"], + }); + upstreams["known-chat-buffered-output"] = await startOpenAiUpstream({ + rawBody: KNOWN_CHAT_OUTPUT, + rawContentType: "application/json", + }); + upstreams["known-responses-buffered-output"] = await startOpenAiUpstream({ + rawBody: KNOWN_RESPONSES_OUTPUT, + rawContentType: "application/json", + }); + upstreams["unknown-typed-buffered-output"] = await startOpenAiUpstream({ + rawBody: UNKNOWN_TYPED_BUFFERED_OUTPUT, + rawContentType: "application/json", + }); + upstreams["responses-refusal-buffered"] = await startOpenAiUpstream({ + rawBody: RESPONSES_REFUSAL_OUTPUT, + rawContentType: "application/json", + }); + upstreams["known-chat-stream-output"] = await startOpenAiUpstream({ + rawStreamFrames: [KNOWN_CHAT_STREAM], + }); + upstreams["known-responses-stream-output"] = await startOpenAiUpstream({ + rawStreamFrames: [KNOWN_RESPONSES_STREAM], + }); + upstreams["split-responses-stream-output"] = await startOpenAiUpstream({ + rawStreamFrames: SPLIT_RESPONSES_STREAM, + }); + upstreams["distinct-responses-stream-output"] = await startOpenAiUpstream({ + rawStreamFrames: DISTINCT_ITEMS_RESPONSES_STREAM, + }); + upstreams["known-responses-reasoning-stream"] = await startOpenAiUpstream({ + rawStreamFrames: RESPONSES_REASONING_STREAM, + }); const pk = await seed.createProviderKey({ display_name: "pt-scan-pk", secret: "sk-mock", @@ -115,14 +415,22 @@ describe("passthrough guardrail scan coverage", () => { enabled: true, hook_point: "output", kind: "keyword", - patterns: [{ kind: "literal", value: OUT_LIT }], + patterns: [ + { kind: "literal", value: OUT_LIT }, + { kind: "literal", value: ESCAPED_BLOCK }, + { kind: "regex", value: SPLIT_BLOCK_REGEX }, + ], }); await seed.createGuardrail({ name: "pt-scan-input", enabled: true, hook_point: "input", kind: "keyword", - patterns: [{ kind: "literal", value: IN_LIT }], + patterns: [ + { kind: "literal", value: IN_LIT }, + { kind: "literal", value: ESCAPED_BLOCK }, + { kind: "literal", value: ESCAPED_CJK }, + ], }); // Folds the output chain's hold-back cap down to CAP, fail-closed. await seed.createGuardrail({ @@ -134,6 +442,12 @@ describe("passthrough guardrail scan coverage", () => { max_buffer_bytes: CAP, on_buffer_exceeded: "fail_closed", }); + await seed.createApiKey({ + key_hash: ERROR_STREAM_CALLER_HASH, + allowed_models: [], + allowed_routes: ["pt-scan-delayed-sse-502"], + rate_limit: { concurrency: 1 }, + }); await seed.createApiKey({ key_hash: CALLER_HASH, allowed_models: [], allowed_routes: ["*"] }); const proxy = new ProxyClient(app.proxyUrl, CALLER); await waitConfigPropagation(async () => (await proxy.listModels()).status === 200); @@ -142,6 +456,7 @@ describe("passthrough guardrail scan coverage", () => { afterAll(async () => { await app?.exit(); await Promise.all(Object.values(upstreams).map((u) => u.close())); + await Promise.all(otlps.map((o) => o.close())); }); const call = (route: string, path: string, body: Record) => @@ -150,12 +465,23 @@ describe("passthrough guardrail scan coverage", () => { headers: { authorization: `Bearer ${CALLER}`, "content-type": "application/json" }, body: JSON.stringify(body), }); + const callRaw = (route: string, path: string, body: string) => + fetch(`${app!.proxyUrl}/pt-scan-${route}${path}`, { + method: "POST", + headers: { authorization: `Bearer ${CALLER}`, "content-type": "application/json" }, + body, + }); + const callRawHttp = (route: string, path: string, body: string, caller = CALLER) => + openRawHttpRequest(`${app!.proxyUrl}/pt-scan-${route}${path}`, { + authorization: `Bearer ${caller}`, + "content-type": "application/json", + }, body); const anthropicBody = { model: "claude-3-5-haiku-20241022", max_tokens: 64, stream: true, messages: [{ role: "user", content: "go" }] }; const chatBody = { model: "gpt-4o-mini", stream: true, messages: [{ role: "user", content: "go" }] }; const responsesBody = { model: "gpt-4o-mini", stream: true, input: "go" }; const ready = (ctx: { skip: () => void }) => { - if (!etcdReachable || !app) { + if (!etcdReachable || !app || !seed) { ctx.skip(); return false; } @@ -177,9 +503,268 @@ describe("passthrough guardrail scan coverage", () => { ["responses-tool", "/v1/responses", responsesBody], ] as const)("output: %s is scanned", async ([route, path, body], ctx) => { if (!ready(ctx)) return; + const upstream = upstreams[route]; + if (!upstream) throw new Error(`missing ${route} upstream`); + const before = upstream.receivedRequests.length; await expectBlocked(await call(route, path, body)); + expect(upstream.receivedRequests.length).toBe(before + 1); }); + test("output: a buffered Responses refusal is source-scanned", async (ctx) => { + if (!ready(ctx)) return; + const route = "responses-refusal-buffered"; + const upstream = upstreams[route]; + if (!upstream) throw new Error(`missing ${route} upstream`); + const before = upstream.receivedRequests.length; + const res = await callRaw(route, "/v1/any", `{"model":"gpt-4o-mini","input":"go"}`); + const body = await res.text(); + expect(res.status, body).toBe(422); + expect(body).toContain("pt-scan-output"); + expect(body).not.toContain(OUT_LIT); + expect(upstream.receivedRequests.length).toBe(before + 1); + }); + + test("output: a streamed Responses refusal is source-scanned", async (ctx) => { + if (!ready(ctx)) return; + const route = "responses-refusal"; + const upstream = upstreams[route]; + if (!upstream) throw new Error(`missing ${route} upstream`); + const before = upstream.receivedRequests.length; + await expectBlocked( + await callRaw(route, "/v1/any", `{"model":"gpt-4o-mini","stream":true,"input":"go"}`), + ); + expect(upstream.receivedRequests.length).toBe(before + 1); + }); + + test.for([ + "responses-output-text-done", + "responses-refusal-done-only", + "responses-function-done", + "responses-mcp-done", + "responses-custom-done", + "responses-content-part-added", + "responses-content-part-done", + "responses-output-item-added", + "responses-output-item-done", + "responses-terminal-only", + "responses-clean-delta-malicious-done", + ] as const)("output: standalone Responses carrier %s is source-scanned", async (route, ctx) => { + if (!ready(ctx)) return; + const upstream = upstreams[route]; + if (!upstream) throw new Error(`missing ${route} upstream`); + const before = upstream.receivedRequests.length; + const res = await callRaw(route, "/v1/any", `{"model":"gpt-4o-mini","stream":true,"input":"go"}`); + expect(res.status).toBe(200); + const body = await res.text(); + expect(body).toContain("event: error"); + expect(body).toContain("content_filter"); + expect(body).toContain("pt-scan-output"); + expect(body).not.toContain("guardrail_unavailable"); + expect(body).not.toContain("unscannable_body"); + expect(body).not.toContain(OUT_LIT); + expect(upstream.receivedRequests.length).toBe(before + 1); + }); + + test("output: a streamed Responses refusal obeys the hold-back cap", async (ctx) => { + if (!ready(ctx)) return; + const route = "responses-refusal-cap"; + const upstream = upstreams[route]; + if (!upstream) throw new Error(`missing ${route} upstream`); + const before = upstream.receivedRequests.length; + const res = await callRaw( + route, + "/v1/any", + `{"model":"gpt-4o-mini","stream":true,"input":"go"}`, + ); + expect(res.status).toBe(200); + const body = await res.text(); + expect(body).toContain("event: error"); + expect(body).toContain("guardrail_unavailable"); + expect(body).toContain("output_buffer_exceeded"); + expect(body).not.toContain("x".repeat(CAP + 1)); + expect(upstream.receivedRequests.length).toBe(before + 1); + }); + + test("output: a done-only Responses refusal obeys the hold-back cap", async (ctx) => { + if (!ready(ctx)) return; + const route = "responses-refusal-done-cap"; + const upstream = upstreams[route]; + if (!upstream) throw new Error(`missing ${route} upstream`); + const before = upstream.receivedRequests.length; + const res = await callRaw( + route, + "/v1/any", + `{"model":"gpt-4o-mini","stream":true,"input":"go"}`, + ); + expect(res.status).toBe(200); + const body = await res.text(); + expect(body).toContain("event: error"); + expect(body).toContain("guardrail_unavailable"); + expect(body).toContain("output_buffer_exceeded"); + expect(body).not.toContain("x".repeat(CAP + 1)); + expect(upstream.receivedRequests.length).toBe(before + 1); + }); + + test("output: repeated Responses snapshots count as one logical carrier toward the hold-back cap", async (ctx) => { + if (!ready(ctx)) return; + const route = "responses-repeated-snapshots-under-cap"; + const upstream = upstreams[route]; + if (!upstream) throw new Error(`missing ${route} upstream`); + const before = upstream.receivedRequests.length; + const res = await callRaw( + route, + "/v1/any", + `{"model":"gpt-4o-mini","stream":true,"input":"go"}`, + ); + const body = await res.text(); + expect(res.status, body).toBe(200); + expect(body).toContain(RESPONSES_SNAPSHOT_TEXT); + expect(body).not.toContain("event: error"); + expect(body).not.toContain("output_buffer_exceeded"); + expect(upstream.receivedRequests.length).toBe(before + 1); + }); + + test("output: upstream 502 HTML bypasses fail-closed output guardrails", async (ctx) => { + if (!ready(ctx)) return; + const upstream = upstreams["html-502"]; + if (!upstream) throw new Error("missing html-502 upstream"); + const before = upstream.receivedRequests.length; + const res = await callRaw( + "html-502", + "/v1/any", + `{"model":"gpt-4o-mini","messages":[{"role":"user","content":"clean"}]}`, + ); + expect(res.status).toBe(502); + expect(res.headers.get("content-type")).toBe("text/html; charset=utf-8"); + expect(res.headers.get("x-upstream-error")).toBe("edge-502"); + expect(await res.text()).toBe(UPSTREAM_502_HTML); + expect(upstream.receivedRequests.length).toBe(before + 1); + }); + + test("output: upstream 502 SSE bypasses guardrails and preserves the error stream", async (ctx) => { + if (!ready(ctx)) return; + const upstream = upstreams["sse-502"]; + if (!upstream) throw new Error("missing sse-502 upstream"); + const before = upstream.receivedRequests.length; + const res = await callRaw( + "sse-502", + "/v1/any", + `{"state":"clean"}`, + ); + expect(res.status).toBe(502); + expect(res.headers.get("content-type")).toBe("text/event-stream; charset=utf-8"); + expect(res.headers.get("cache-control")).toBe("no-cache"); + expect(res.headers.get("x-upstream-error")).toBe("edge-sse-502"); + expect(await res.text()).toBe(UPSTREAM_502_SSE); + expect(upstream.receivedRequests.length).toBe(before + 1); + }); + + test("output: a non-2xx SSE relay preserves encoded bytes and representation headers", async (ctx) => { + if (!ready(ctx)) return; + const upstream = upstreams["gzip-sse-502"]; + if (!upstream) throw new Error("missing gzip-sse-502 upstream"); + const before = upstream.receivedRequests.length; + const res = await callRawHttp("gzip-sse-502", "/v1/any", `{"state":"clean"}`); + expect(res.statusCode).toBe(502); + expect(res.headers["content-type"]).toBe("text/event-stream; charset=utf-8"); + expect(res.headers["content-encoding"]).toBe("gzip"); + expect(res.headers["content-length"]).toBe(String(UPSTREAM_502_GZIP.byteLength)); + expect(await readRawHttpBody(res)).toEqual(UPSTREAM_502_GZIP); + expect(upstream.receivedRequests.length).toBe(before + 1); + }); + + test.for([ + [ + "gzip-buffered-200", + "/v1/completions", + `{"model":"gpt-4o-mini","prompt":"clean"}`, + ], + [ + "gzip-sse-200", + "/v1/chat/completions", + `{"model":"gpt-4o-mini","stream":true,"messages":[{"role":"user","content":"clean"}]}`, + ], + ] as const)( + "output: encoded successful %s response is never sent to a guardrail selector", + async ([route, path, requestBody], ctx) => { + if (!ready(ctx)) return; + const upstream = upstreams[route]; + if (!upstream) throw new Error(`missing ${route} upstream`); + const before = upstream.receivedRequests.length; + const res = await callRaw(route, path, requestBody); + const body = await res.text(); + expect(res.status, body).toBe(422); + expect(body).toContain("guardrail_unavailable"); + expect(body).toContain("unscannable_body"); + expect(body).not.toContain(ESCAPED_BLOCK); + expect(upstream.receivedRequests.length).toBe(before + 1); + expect(upstream.receivedRequests.at(-1)?.headers["accept-encoding"]).toBe("identity"); + }, + ); + + test("output: text/event-streaming is buffered and scanned as a non-SSE response", async (ctx) => { + if (!ready(ctx)) return; + const upstream = upstreams["event-streaming"]; + if (!upstream) throw new Error("missing event-streaming upstream"); + const before = upstream.receivedRequests.length; + const res = await callRaw("event-streaming", "/v1/any", `{"state":"clean"}`); + const body = await res.text(); + expect(res.status, body).toBe(422); + expect(body).toContain("pt-scan-output"); + expect(body).not.toContain(ESCAPED_BLOCK); + expect(upstream.receivedRequests.length).toBe(before + 1); + }); + + test( + "output: a delayed 502 SSE holds its concurrency slot until EOF and relays bytes unchanged", + { timeout: 20_000 }, + async (ctx) => { + if (!ready(ctx)) return; + const upstream = upstreams["delayed-sse-502"]; + if (!upstream) throw new Error("missing delayed-sse-502 upstream"); + const headers = { authorization: `Bearer ${ERROR_STREAM_CALLER}` }; + await waitConfigPropagation(async () => { + const probe = await fetch(`${app!.proxyUrl}/v1/models`, { headers }); + await probe.text(); + return probe.status === 200; + }); + + const before = upstream.receivedRequests.length; + const first = await callRawHttp( + "delayed-sse-502", + "/v1/any", + `{"state":"clean"}`, + ERROR_STREAM_CALLER, + ); + expect(first.statusCode).toBe(502); + expect(first.headers["content-type"]).toBe("text/event-stream; charset=utf-8"); + + // Do not attach a reader before this request. The caller has headers + // but has not consumed EOF, which is the lifecycle that must retain + // the gateway's concurrency reservation. + const second = await fetch(`${app!.proxyUrl}/pt-scan-delayed-sse-502/v1/any`, { + method: "POST", + headers: { ...headers, "content-type": "application/json" }, + body: `{"state":"clean"}`, + }); + expect(second.status).toBe(429); + expect(second.headers.get("x-ratelimit-scope")).toBe("concurrency"); + await second.text(); + expect(upstream.receivedRequests.length).toBe(before + 1); + + expect(await readRawHttpBody(first)).toEqual(Buffer.from(UPSTREAM_502_SSE)); + const after = await callRawHttp( + "delayed-sse-502", + "/v1/any", + `{"state":"clean"}`, + ERROR_STREAM_CALLER, + ); + expect(after.statusCode).toBe(502); + await readRawHttpBody(after); + expect(upstream.receivedRequests.length).toBe(before + 2); + }, + ); + test("output: generated thinking is not scanned", async (ctx) => { if (!ready(ctx)) return; const res = await call("anthropic-thinking", "/v1/messages", anthropicBody); @@ -188,6 +773,136 @@ describe("passthrough guardrail scan coverage", () => { expect(body).not.toContain("event: error"); expect(body).toContain("visible answer"); }); + test.for([ + [ + "Chat response to a Responses request", + "known-chat-buffered-output", + `{"model":"gpt-4o-mini","input":"go"}`, + ], + [ + "Responses response to a Completions request", + "known-responses-buffered-output", + `{"model":"gpt-4o-mini","prompt":"go"}`, + ], + ] as const)("output: known %s is source-scanned when buffered", async ([, route, body], ctx) => { + if (!ready(ctx)) return; + const upstream = upstreams[route]; + if (!upstream) throw new Error(`missing ${route} upstream`); + const before = upstream.receivedRequests.length; + const res = await callRaw(route, "/v1/any", body); + expect(res.status).toBe(422); + const response = await res.text(); + expect(response).toContain("pt-scan-output"); + expect(response).not.toContain(ESCAPED_BLOCK); + expect(response).not.toContain(String.raw`\u0042LOCKME`); + expect(upstream.receivedRequests.length).toBe(before + 1); + }); + + test("output: an unknown typed buffered response fails closed instead of using the request hint", async (ctx) => { + if (!ready(ctx)) return; + const route = "unknown-typed-buffered-output"; + const upstream = upstreams[route]; + if (!upstream) throw new Error(`missing ${route} upstream`); + const before = upstream.receivedRequests.length; + const res = await callRaw( + route, + "/v1/completions", + `{"model":"gpt-4o-mini","prompt":"go"}`, + ); + const body = await res.text(); + expect(res.status, body).toBe(422); + expect(body).toContain("guardrail_unavailable"); + expect(body).toContain("unscannable_body"); + expect(body).not.toContain(ESCAPED_BLOCK); + expect(body).not.toContain(String.raw`\u0042LOCKME`); + expect(upstream.receivedRequests.length).toBe(before + 1); + }); + + test.for([ + [ + "chat", + "known-chat-stream-output", + `{"model":"gpt-4o-mini","stream":true,"messages":[{"role":"user","content":"go"}]}`, + ], + [ + "Responses", + "known-responses-stream-output", + `{"model":"gpt-4o-mini","stream":true,"input":"go"}`, + ], + ] as const)("output: known %s envelope is source-scanned when streamed", async ([, route, body], ctx) => { + if (!ready(ctx)) return; + const upstream = upstreams[route]; + if (!upstream) throw new Error(`missing ${route} upstream`); + const before = upstream.receivedRequests.length; + const res = await callRaw(route, "/v1/any", body); + expect(res.status).toBe(200); + const response = await res.text(); + expect(response).toContain("event: error"); + expect(response).toContain("content_filter"); + expect(response).not.toContain(ESCAPED_BLOCK); + expect(response).not.toContain(String.raw`\u0042LOCKME`); + expect(upstream.receivedRequests.length).toBe(before + 1); + }); + + test("output: visible deltas remain contiguous across stream metadata", async (ctx) => { + if (!ready(ctx)) return; + const route = "split-responses-stream-output"; + const upstream = upstreams[route]; + if (!upstream) throw new Error(`missing ${route} upstream`); + const before = upstream.receivedRequests.length; + const res = await callRaw( + route, + "/v1/any", + `{"model":"gpt-4o-mini","stream":true,"input":"go"}`, + ); + expect(res.status).toBe(200); + const response = await res.text(); + expect(response).toContain("event: error"); + expect(response).toContain("content_filter"); + expect(response).toContain("pt-scan-output"); + expect(response).not.toContain("guardrail_unavailable"); + expect(response).not.toContain("unscannable_body"); + expect(response).not.toContain(SPLIT_BLOCK); + expect(upstream.receivedRequests.length).toBe(before + 1); + }); + + test("output: visible deltas from distinct Responses items never concatenate", async (ctx) => { + if (!ready(ctx)) return; + const route = "distinct-responses-stream-output"; + const upstream = upstreams[route]; + if (!upstream) throw new Error(`missing ${route} upstream`); + const before = upstream.receivedRequests.length; + const res = await callRaw( + route, + "/v1/any", + `{"model":"gpt-4o-mini","stream":true,"input":"go"}`, + ); + expect(res.status).toBe(200); + const response = await res.text(); + expect(response).not.toContain("event: error"); + expect(response).toContain("FOR"); + expect(response).toContain("BIDDEN"); + expect(upstream.receivedRequests.length).toBe(before + 1); + }); + + test("output: generated Responses reasoning frames stay out of scope", async (ctx) => { + if (!ready(ctx)) return; + const route = "known-responses-reasoning-stream"; + const upstream = upstreams[route]; + if (!upstream) throw new Error(`missing ${route} upstream`); + const before = upstream.receivedRequests.length; + const res = await callRaw( + route, + "/v1/any", + `{"model":"gpt-4o-mini","stream":true,"input":"go"}`, + ); + expect(res.status).toBe(200); + const response = await res.text(); + expect(response).not.toContain("event: error"); + expect(response).toContain(OUT_LIT); + expect(response).toContain("clean"); + expect(upstream.receivedRequests.length).toBe(before + 1); + }); test("hold-back cap: frames over the cap, content under it, is released", async (ctx) => { if (!ready(ctx)) return; @@ -236,4 +951,898 @@ describe("passthrough guardrail scan coverage", () => { expect(await res.text()).toContain("pt-scan-input"); expect(upstreams.input!.receivedRequests.length).toBe(before); }); + + test.for([ + [ + "system-one state", + { + messages: [{ role: "user", content: "clean" }], + state: { query: IN_LIT }, + }, + ], + [ + "rerank query and documents", + { + messages: [{ role: "user", content: "clean" }], + query: IN_LIT, + documents: ["clean"], + }, + ], + ] as const)("input: forwarded %s is scanned", async ([, body], ctx) => { + if (!ready(ctx)) return; + const before = upstreams.input!.receivedRequests.length; + const res = await call("input", "/v1/any", body); + expect(res.status).toBe(422); + expect(await res.text()).toContain("pt-scan-input"); + expect(upstreams.input!.receivedRequests.length).toBe(before); + }); + + test("input: a duplicate forwarded field is scanned", async (ctx) => { + if (!ready(ctx)) return; + const before = upstreams.input!.receivedRequests.length; + const body = String.raw`{"messages":[{"role":"user","content":"clean"}],"state":{"query":"${IN_LIT}"},"state":"clean"}`; + const res = await callRaw("input", "/v1/any", body); + expect(res.status).toBe(422); + expect(await res.text()).toContain("pt-scan-input"); + expect(upstreams.input!.receivedRequests.length).toBe(before); + }); + test.for([ + [ + "duplicate chat messages", + String.raw`{"model":"gpt-4o-mini","messages":[{"role":"user","content":"\u0069nputleakliteral","metadata":{"note":"${IN_LIT}"}}],"messages":[{"role":"user","content":"clean"}]}`, + ], + [ + "duplicate Responses input", + String.raw`{"model":"gpt-4o-mini","input":"\u0069nputleakliteral","input":"clean","metadata":{"note":"${IN_LIT}"}}`, + ], + ] as const)("input: known envelope %s is source-scanned", async ([, body], ctx) => { + if (!ready(ctx)) return; + const before = upstreams.input!.receivedRequests.length; + const res = await callRaw("input", "/v1/any", body); + expect(res.status).toBe(422); + expect(await res.text()).toContain("pt-scan-input"); + expect(upstreams.input!.receivedRequests.length).toBe(before); + }); + + test.for([ + ["messages", `{"model":"gpt-4o-mini","messages":["forbidden"]}`], + [ + "content", + `{"model":"gpt-4o-mini","messages":[{"role":"user","content":["forbidden"]}]}`, + ], + ] as const)( + "input: malformed Chat %s carrier fails closed without raw-scanning it", + async ([, body], ctx) => { + if (!ready(ctx)) return; + const before = upstreams.input!.receivedRequests.length; + const res = await callRaw("input", "/v1/any", body); + expect(res.status).toBe(422); + const response = await res.text(); + expect(response).toContain("guardrail_unavailable"); + expect(response).toContain("unscannable_body"); + expect(response).not.toContain("forbidden"); + expect(upstreams.input!.receivedRequests.length).toBe(before); + }, + ); + + test("input: malformed JSON after typed media fails closed instead of becoming Raw", async (ctx) => { + if (!ready(ctx)) return; + const before = upstreams.input!.receivedRequests.length; + const body = `{"model":"gpt-4o-mini","messages":[{"role":"user","content":[{"type":"image","source":{"data":"${ESCAPED_BLOCK}"}}]}],"broken":`; + const res = await callRaw("input", "/v1/any", body); + expect(res.status).toBe(422); + const response = await res.text(); + expect(response).toContain("guardrail_unavailable"); + expect(response).toContain("unscannable_body"); + expect(response).not.toContain(ESCAPED_BLOCK); + expect(upstreams.input!.receivedRequests.length).toBe(before); + }); + + test.for([ + ["ASCII", ESCAPED_BLOCK_JSON], + ["CJK", ESCAPED_CJK_JSON], + ] as const)("input: raw JSON %s escapes are decoded before scanning", async ([, body], ctx) => { + if (!ready(ctx)) return; + const before = upstreams.input!.receivedRequests.length; + const res = await callRaw("input", "/v1/any", body); + expect(res.status).toBe(422); + expect(await res.text()).toContain("pt-scan-input"); + expect(upstreams.input!.receivedRequests.length).toBe(before); + }); + + test("input: deep raw JSON escapes are decoded before scanning", async (ctx) => { + if (!ready(ctx)) return; + const before = upstreams.input!.receivedRequests.length; + const res = await callRaw("input", "/v1/any", DEEP_ESCAPED_BLOCK_JSON); + expect(res.status).toBe(422); + expect(await res.text()).toContain("pt-scan-input"); + expect(upstreams.input!.receivedRequests.length).toBe(before); + }); + + test("input: Raw JSON beyond the scanner depth cap fails closed", async (ctx) => { + if (!ready(ctx)) return; + const before = upstreams.input!.receivedRequests.length; + const res = await callRaw("input", "/v1/any", OVER_DEPTH_ESCAPED_BLOCK_JSON); + expect(res.status).toBe(422); + const body = await res.text(); + expect(body).toContain("guardrail_unavailable"); + expect(body).toContain("unscannable_body"); + expect(body).not.toContain(ESCAPED_BLOCK); + expect(upstreams.input!.receivedRequests.length).toBe(before); + }); + + test("input: nested Anthropic tool results within the structural-work cap are blocked", async (ctx) => { + if (!ready(ctx)) return; + const before = upstreams.input!.receivedRequests.length; + const res = await callRaw("input", "/v1/any", AT_DEPTH_ANTHROPIC_TOOL_RESULT_INPUT); + expect(res.status).toBe(422); + const body = await res.text(); + expect(body).toContain("pt-scan-input"); + expect(body).not.toContain("guardrail_unavailable"); + expect(body).not.toContain("unscannable_body"); + expect(body).not.toContain(ESCAPED_BLOCK); + expect(upstreams.input!.receivedRequests.length).toBe(before); + }); + + test("input: nested Anthropic tool results beyond the structural-work cap fail closed", async (ctx) => { + if (!ready(ctx)) return; + const before = upstreams.input!.receivedRequests.length; + const res = await callRaw("input", "/v1/any", OVER_DEPTH_ANTHROPIC_TOOL_RESULT_INPUT); + expect(res.status).toBe(422); + const body = await res.text(); + expect(body).toContain("guardrail_unavailable"); + expect(body).toContain("unscannable_body"); + expect(body).not.toContain(ESCAPED_BLOCK); + expect(upstreams.input!.receivedRequests.length).toBe(before); + }); + + test.for([ + ["Chat", OVER_DEPTH_CHAT_INPUT], + ["Responses", OVER_DEPTH_RESPONSES_INPUT], + ] as const)("input: %s envelope beyond the scanner depth cap fails closed", async ([, body], ctx) => { + if (!ready(ctx)) return; + const before = upstreams.input!.receivedRequests.length; + const res = await callRaw("input", "/v1/any", body); + expect(res.status).toBe(422); + const response = await res.text(); + expect(response).toContain("guardrail_unavailable"); + expect(response).toContain("unscannable_body"); + expect(response).not.toContain(ESCAPED_BLOCK); + expect(upstreams.input!.receivedRequests.length).toBe(before); + }); + + test.for([ + ["Chat", OVER_DEPTH_CHAT_OPAQUE_INPUT], + ["Responses", OVER_DEPTH_RESPONSES_OPAQUE_INPUT], + ] as const)("input: %s opaque media beyond the scanner depth cap stays out of scope", async ([, body], ctx) => { + if (!ready(ctx)) return; + const before = upstreams.input!.receivedRequests.length; + const res = await callRaw("input", "/v1/any", body); + expect(res.status).toBe(200); + await res.text(); + expect(upstreams.input!.receivedRequests.length).toBe(before + 1); + expect(upstreams.input!.receivedRequests.at(-1)!.body).toBe(body); + }); + + test("input: safe raw JSON keeps its original bytes upstream", async (ctx) => { + if (!ready(ctx)) return; + const before = upstreams.input!.receivedRequests.length; + const res = await callRaw("input", "/v1/any", SAFE_ESCAPED_JSON); + expect(res.status).toBe(200); + await res.text(); + expect(upstreams.input!.receivedRequests.length).toBe(before + 1); + expect(upstreams.input!.receivedRequests.at(-1)!.body).toBe(SAFE_ESCAPED_JSON); + }); + + test("output: raw JSON escapes are decoded before scanning", async (ctx) => { + if (!ready(ctx)) return; + const before = upstreams["raw-output"]!.receivedRequests.length; + const res = await callRaw("raw-output", "/v1/any", String.raw`{"state":"clean"}`); + expect(res.status).toBe(422); + const body = await res.text(); + expect(body).toContain("pt-scan-output"); + expect(body).not.toContain(ESCAPED_BLOCK); + expect(upstreams["raw-output"]!.receivedRequests.length).toBe(before + 1); + }); + + test("output: safe raw JSON keeps its original bytes downstream", async (ctx) => { + if (!ready(ctx)) return; + const before = upstreams["raw-safe-output"]!.receivedRequests.length; + const res = await callRaw("raw-safe-output", "/v1/any", SAFE_ESCAPED_JSON); + expect(res.status).toBe(200); + expect(await res.text()).toBe(SAFE_ESCAPED_JSON); + expect(upstreams["raw-safe-output"]!.receivedRequests.length).toBe(before + 1); + }); + + test("output: deep raw JSON escapes are decoded before scanning", async (ctx) => { + if (!ready(ctx)) return; + const before = upstreams["raw-deep-output"]!.receivedRequests.length; + const res = await callRaw("raw-deep-output", "/v1/any", String.raw`{"state":"clean"}`); + expect(res.status).toBe(422); + const body = await res.text(); + expect(body).toContain("pt-scan-output"); + expect(body).not.toContain(ESCAPED_BLOCK); + expect(upstreams["raw-deep-output"]!.receivedRequests.length).toBe(before + 1); + }); + + test("output: Raw JSON beyond the scanner depth cap fails closed", async (ctx) => { + if (!ready(ctx)) return; + const before = upstreams["raw-over-depth-output"]!.receivedRequests.length; + const res = await callRaw("raw-over-depth-output", "/v1/any", String.raw`{"state":"clean"}`); + expect(res.status).toBe(422); + const body = await res.text(); + expect(body).toContain("guardrail_unavailable"); + expect(body).toContain("unscannable_body"); + expect(body).not.toContain(ESCAPED_BLOCK); + expect(upstreams["raw-over-depth-output"]!.receivedRequests.length).toBe(before + 1); + }); + + test("output: completions JSON beyond the scanner depth cap fails closed", async (ctx) => { + if (!ready(ctx)) return; + const before = upstreams["completions-over-depth-output"]!.receivedRequests.length; + const res = await callRaw( + "completions-over-depth-output", + "/v1/completions", + `{"model":"gpt-4o-mini","prompt":"go"}`, + ); + expect(res.status).toBe(422); + const body = await res.text(); + expect(body).toContain("guardrail_unavailable"); + expect(body).toContain("unscannable_body"); + expect(body).not.toContain(ESCAPED_BLOCK); + expect(upstreams["completions-over-depth-output"]!.receivedRequests.length).toBe(before + 1); + }); + + test("output: Anthropic tool input beyond the scanner depth cap fails closed", async (ctx) => { + if (!ready(ctx)) return; + const route = "anthropic-over-depth-tool-output"; + const upstream = upstreams[route]; + if (!upstream) throw new Error(`missing ${route} upstream`); + const before = upstream.receivedRequests.length; + const res = await callRaw( + route, + "/v1/any", + `{"model":"gpt-4o-mini","messages":[{"role":"user","content":"go"}]}`, + ); + expect(res.status).toBe(422); + const response = await res.text(); + expect(response).toContain("guardrail_unavailable"); + expect(response).toContain("unscannable_body"); + expect(response).not.toContain(ESCAPED_BLOCK); + expect(upstream.receivedRequests.length).toBe(before + 1); + }); + + test("output: raw SSE JSON escapes are decoded before scanning", async (ctx) => { + if (!ready(ctx)) return; + const before = upstreams["raw-stream"]!.receivedRequests.length; + const res = await callRaw("raw-stream", "/v1/any", String.raw`{"state":"clean"}`); + expect(res.status).toBe(200); + const body = await res.text(); + expect(body).toContain("event: error"); + expect(body).toContain("content_filter"); + expect(body).not.toContain(ESCAPED_BLOCK); + expect(upstreams["raw-stream"]!.receivedRequests.length).toBe(before + 1); + }); + + test("output: safe bare-string Raw SSE keeps its original bytes downstream", async (ctx) => { + if (!ready(ctx)) return; + const before = upstreams["raw-safe-stream"]!.receivedRequests.length; + const res = await callRaw("raw-safe-stream", "/v1/any", SAFE_ESCAPED_JSON); + expect(res.status).toBe(200); + expect(await res.text()).toBe(`data: ${RAW_SAFE_SSE}\n\ndata: [DONE]\n\n`); + expect(upstreams["raw-safe-stream"]!.receivedRequests.length).toBe(before + 1); + }); + + test("telemetry: Raw buffered and SSE responses retain their original JSON source", async (ctx) => { + if (!ready(ctx)) return; + + const otlp = await startMockOtlp(); + otlps.push(otlp); + await seed!.createObservabilityExporter({ + name: "pt-scan-raw-source-capture", + enabled: true, + kind: "otlp_http", + endpoint: otlp.url, + content_mode: "full", + content_max_bytes: 4_096, + }); + const requestUntilCaptured = async ( + route: string, + expectedBody: string, + expectedCompletion: string, + ) => { + // A healthy `/v1/models` reply proves only caller authentication. It + // does not prove this newly added exporter has reached the snapshot, so + // retry the real route until its real completion reaches the OTLP + // receiver. Each attempt must still preserve the exact client bytes. + const deadline = Date.now() + 30_000; + let last = "no response"; + while (Date.now() < deadline) { + const response = await callRaw(route, "/v1/any", SAFE_ESCAPED_JSON); + const body = await response.text(); + last = `${response.status}: ${body}`; + if (response.status !== 200 || body !== expectedBody) { + await new Promise((resolve) => setTimeout(resolve, 50)); + continue; + } + const span = otlp.spans.find( + (candidate) => + candidate.attributes["aisix.passthrough.route_name"] === `pt-scan-${route}` && + candidate.attributes["gen_ai.completion"] === expectedCompletion, + ); + if (span) return span; + await new Promise((resolve) => setTimeout(resolve, 50)); + } + throw new Error(`no raw-source OTLP completion for ${route}; last response ${last}`); + }; + + await requestUntilCaptured("raw-safe-output", SAFE_ESCAPED_JSON, SAFE_ESCAPED_JSON); + await requestUntilCaptured( + "raw-safe-stream", + `data: ${RAW_SAFE_SSE}\n\ndata: [DONE]\n\n`, + RAW_SAFE_SSE, + ); + expect(otlp.parseFailures).toEqual([]); + }); + + test("output: deep Raw JSON SSE is refused as an unevaluable carrier", async (ctx) => { + if (!ready(ctx)) return; + const before = upstreams["raw-deep-stream"]!.receivedRequests.length; + const res = await callRaw("raw-deep-stream", "/v1/any", String.raw`{"state":"clean"}`); + expect(res.status).toBe(200); + const body = await res.text(); + expect(body).toContain("event: error"); + expect(body).toContain("guardrail_unavailable"); + expect(body).toContain("unscannable_body"); + expect(body).not.toContain(ESCAPED_BLOCK); + expect(upstreams["raw-deep-stream"]!.receivedRequests.length).toBe(before + 1); + }); +}); + +// An output fail-open policy must relay an encoded successful representation +// unchanged. It cannot send compressed bytes to a text moderator, but it must +// record that the response bypassed inspection rather than turning a 200 into +// a local 422. +describe("passthrough encoded successful output fail-open", () => { + const caller = "sk-pt-encoded-fail-open"; + const callerHash = createHash("sha256").update(caller).digest("hex"); + const bufferedRoute = "pt-encoded-buffered-fail-open"; + const sseRoute = "pt-encoded-sse-fail-open"; + const logstore = "pt-encoded-fail-open"; + const credentialRef = "pt_encoded_open"; + let app: SpawnedApp | undefined; + let bufferedUpstream: OpenAiUpstream | undefined; + let sseUpstream: OpenAiUpstream | undefined; + let moderationUpstream: OpenAiUpstream | undefined; + let sls: MockSls | undefined; + let etcdReachable = false; + + beforeAll(async () => { + const etcd = new EtcdClient(); + etcdReachable = await etcd.ping(); + if (!etcdReachable) return; + + sls = await startMockSls(); + bufferedUpstream = await startOpenAiUpstream({ + rawBody: UPSTREAM_200_GZIP_JSON, + rawContentType: "application/json", + responseHeaders: { + "content-encoding": "gzip", + "content-length": String(UPSTREAM_200_GZIP_JSON.byteLength), + }, + }); + sseUpstream = await startOpenAiUpstream({ + rawBody: UPSTREAM_200_GZIP_SSE, + rawContentType: "text/event-stream; charset=utf-8", + responseHeaders: { + "content-encoding": "gzip", + "content-length": String(UPSTREAM_200_GZIP_SSE.byteLength), + }, + }); + moderationUpstream = await startOpenAiUpstream({ + nonStreamBody: { id: "unused", results: [] }, + }); + app = await spawnApp({ + extraEnv: { + [`SLS_CRED_${credentialRef.toUpperCase()}_AK_ID`]: "mock-akid", + [`SLS_CRED_${credentialRef.toUpperCase()}_AK_SECRET`]: "mock-secret", + }, + }); + const seed = new SeedClient(etcd, app.etcdPrefix); + await seed.createObservabilityExporter({ + name: "pt-encoded-fail-open-sls", + enabled: true, + kind: "aliyun_sls", + endpoint: sls.url, + project: "aisix-e2e-obs", + logstore, + credential_ref: credentialRef, + }); + const providerKey = await seed.createProviderKey({ + display_name: "pt-encoded-fail-open-pk", + secret: "sk-mock", + api_base: bufferedUpstream.baseUrl, + }); + await seed.createPassthroughRoute({ + name: bufferedRoute, + path_prefix: `/${bufferedRoute}`, + target_url: bufferedUpstream.baseUrl, + provider_key_id: providerKey.id, + }); + await seed.createPassthroughRoute({ + name: sseRoute, + path_prefix: `/${sseRoute}`, + target_url: sseUpstream.baseUrl, + provider_key_id: providerKey.id, + }); + await seed.createGuardrail({ + name: "pt-encoded-fail-open-output", + enabled: true, + hook_point: "output", + kind: "openai_moderation", + endpoint: moderationUpstream.baseUrl, + api_key: "sk-local-moderation", + output_fail_open: true, + }); + await seed.createApiKey({ key_hash: callerHash, allowed_models: [], allowed_routes: ["*"] }); + const proxy = new ProxyClient(app.proxyUrl, caller); + await waitConfigPropagation(async () => (await proxy.listModels()).status === 200); + }, 90_000); + + afterAll(async () => { + await app?.exit(); + await bufferedUpstream?.close(); + await sseUpstream?.close(); + await moderationUpstream?.close(); + await sls?.close(); + }); + + const assertBypassAudit = async (route: string, requestId: string) => { + const log = await waitForSlsLog( + sls!, + logstore, + (entry) => entry.get("passthrough_route_name") === route && entry.get("request_id") === requestId, + "encoded fail-open passthrough usage event", + ); + expect(log.get("guardrail_blocked") ?? "false").not.toBe("true"); + expect(log.get("guardrail_bypassed_reason")).toBe("unscannable_body"); + }; + + test("relays an encoded buffered response without calling moderation", async (ctx) => { + if (!etcdReachable || !app || !bufferedUpstream || !moderationUpstream || !sls) return ctx.skip(); + + const before = bufferedUpstream.receivedRequests.length; + const res = await openRawHttpRequest( + `${app.proxyUrl}/${bufferedRoute}/v1/completions`, + { authorization: `Bearer ${caller}`, "content-type": "application/json" }, + `{"model":"gpt-4o-mini","prompt":"clean"}`, + ); + expect(res.statusCode).toBe(200); + expect(res.headers["content-type"]).toBe("application/json"); + expect(res.headers["content-encoding"]).toBe("gzip"); + expect(res.headers["content-length"]).toBe(String(UPSTREAM_200_GZIP_JSON.byteLength)); + const requestId = String(res.headers["x-aisix-request-id"] ?? ""); + expect(requestId).toBeTruthy(); + expect(await readRawHttpBody(res)).toEqual(UPSTREAM_200_GZIP_JSON); + expect(bufferedUpstream.receivedRequests.length).toBe(before + 1); + expect(moderationUpstream.receivedRequests).toHaveLength(0); + await assertBypassAudit(bufferedRoute, requestId); + }); + + test("relays an encoded SSE response without calling moderation", async (ctx) => { + if (!etcdReachable || !app || !sseUpstream || !moderationUpstream || !sls) return ctx.skip(); + + const before = sseUpstream.receivedRequests.length; + const res = await openRawHttpRequest( + `${app.proxyUrl}/${sseRoute}/v1/chat/completions`, + { authorization: `Bearer ${caller}`, "content-type": "application/json" }, + `{"model":"gpt-4o-mini","stream":true,"messages":[{"role":"user","content":"clean"}]}`, + ); + expect(res.statusCode).toBe(200); + expect(res.headers["content-type"]).toBe("text/event-stream; charset=utf-8"); + expect(res.headers["content-encoding"]).toBe("gzip"); + expect(res.headers["content-length"]).toBe(String(UPSTREAM_200_GZIP_SSE.byteLength)); + const requestId = String(res.headers["x-aisix-request-id"] ?? ""); + expect(requestId).toBeTruthy(); + expect(await readRawHttpBody(res)).toEqual(UPSTREAM_200_GZIP_SSE); + expect(sseUpstream.receivedRequests.length).toBe(before + 1); + expect(moderationUpstream.receivedRequests).toHaveLength(0); + await assertBypassAudit(sseRoute, requestId); + }); +}); + +// A source carrier can be valid while a separate selected supplemental field +// is malformed. Even an explicitly output-fail-open remote guardrail may not +// turn that local selector failure into an unscanned provider frame. +describe("passthrough supplemental selector fail-closed", () => { + const route = "pt-scan-supplemental"; + let app: SpawnedApp | undefined; + let upstream: OpenAiUpstream | undefined; + let etcdReachable = false; + + beforeAll(async () => { + const etcd = new EtcdClient(); + etcdReachable = await etcd.ping(); + if (!etcdReachable) return; + + upstream = await startOpenAiUpstream({ rawStreamFrames: [MALFORMED_SUPPLEMENTAL_SSE] }); + app = await spawnApp(); + const seed = new SeedClient(etcd, app.etcdPrefix); + const providerKey = await seed.createProviderKey({ + display_name: "pt-scan-supplemental-pk", + secret: "sk-mock", + api_base: upstream.baseUrl, + }); + await seed.createPassthroughRoute({ + name: route, + path_prefix: `/${route}`, + target_url: upstream.baseUrl, + provider_key_id: providerKey.id, + }); + await seed.createGuardrail({ + name: "pt-scan-supplemental-output-open", + enabled: true, + hook_point: "output", + kind: "openai_moderation", + api_key: "sk-local-moderation", + endpoint: upstream.baseUrl, + output_fail_open: true, + }); + await seed.createApiKey({ + key_hash: SUPPLEMENTAL_CALLER_HASH, + allowed_models: [], + allowed_routes: [route], + }); + const proxy = new ProxyClient(app.proxyUrl, SUPPLEMENTAL_CALLER); + await waitConfigPropagation(async () => (await proxy.listModels()).status === 200); + }, 90_000); + + afterAll(async () => { + await app?.exit(); + await upstream?.close(); + }); + + test("malformed selected supplemental data is refused before an output-fail-open sink runs", async (ctx) => { + if (!etcdReachable || !app || !upstream) { + ctx.skip(); + return; + } + const before = upstream.receivedRequests.length; + const response = await fetch(`${app.proxyUrl}/${route}/v1/chat/completions`, { + method: "POST", + headers: { + authorization: `Bearer ${SUPPLEMENTAL_CALLER}`, + "content-type": "application/json", + }, + body: JSON.stringify({ + model: "gpt-4o-mini", + stream: true, + messages: [{ role: "user", content: "go" }], + }), + }); + const body = await response.text(); + expect(response.status, body).toBe(200); + expect(body).toContain("event: error"); + expect(body).toContain("guardrail_unavailable"); + expect(upstream.receivedRequests.length).toBe(before + 1); + expect(upstream.receivedRequests.at(-1)?.path).toBe("/v1/chat/completions"); + }); +}); + +// This is intentionally a separate DP: the main suite has an env-scoped +// blocking row, so it cannot demonstrate the live fail-open policy on either +// an unevaluable Raw input or a typed output response. +describe("passthrough unevaluable-output fail-open", () => { + const caller = "sk-pt-scan-fail-open"; + const callerHash = createHash("sha256").update(caller).digest("hex"); + const route = "pt-scan-fail-open"; + const depthInputRoute = "pt-scan-depth-fail-open-input"; + const unknownTypedRoute = "pt-scan-unknown-typed-fail-open"; + const logstore = "pt-scan-fail-open"; + const credentialRef = "pt_scan_open"; + let app: SpawnedApp | undefined; + let upstream: OpenAiUpstream | undefined; + let depthInputUpstream: OpenAiUpstream | undefined; + let unknownTypedUpstream: OpenAiUpstream | undefined; + let sls: MockSls | undefined; + let etcdReachable = false; + + beforeAll(async () => { + const etcd = new EtcdClient(); + etcdReachable = await etcd.ping(); + if (!etcdReachable) return; + + sls = await startMockSls(); + upstream = await startOpenAiUpstream({ + rawStreamFrames: [ + `data: ${RAW_PREFIX_SSE}\n\n`, + `data: ${RAW_UNEVALUABLE_SSE}\n\n`, + `data: ${RAW_SUFFIX_SSE}\n\n`, + "data: [DONE]\n\n", + ], + }); + depthInputUpstream = await startOpenAiUpstream({ + rawBody: SAFE_ESCAPED_JSON, + rawContentType: "application/json", + }); + unknownTypedUpstream = await startOpenAiUpstream({ + rawBody: UNKNOWN_TYPED_BUFFERED_OUTPUT, + rawContentType: "application/json", + }); + app = await spawnApp({ + extraEnv: { + [`SLS_CRED_${credentialRef.toUpperCase()}_AK_ID`]: "mock-akid", + [`SLS_CRED_${credentialRef.toUpperCase()}_AK_SECRET`]: "mock-secret", + }, + }); + const seed = new SeedClient(etcd, app.etcdPrefix); + await seed.createObservabilityExporter({ + name: "pt-scan-fail-open-sls", + enabled: true, + kind: "aliyun_sls", + endpoint: sls.url, + project: "aisix-e2e-obs", + logstore, + credential_ref: credentialRef, + }); + const providerKey = await seed.createProviderKey({ + display_name: "pt-scan-fail-open-pk", + secret: "sk-mock", + api_base: upstream.baseUrl, + }); + await seed.createPassthroughRoute({ + name: route, + path_prefix: "/pt-scan-fail-open", + target_url: upstream.baseUrl, + provider_key_id: providerKey.id, + }); + await seed.createPassthroughRoute({ + name: depthInputRoute, + path_prefix: `/${depthInputRoute}`, + target_url: depthInputUpstream.baseUrl, + provider_key_id: providerKey.id, + }); + await seed.createPassthroughRoute({ + name: unknownTypedRoute, + path_prefix: `/${unknownTypedRoute}`, + target_url: unknownTypedUpstream.baseUrl, + provider_key_id: providerKey.id, + }); + await seed.createGuardrail({ + name: "pt-scan-depth-fail-open-input", + enabled: true, + hook_point: "input", + fail_open: true, + kind: "keyword", + patterns: [{ kind: "literal", value: ESCAPED_BLOCK }], + }); + await seed.createGuardrail({ + name: "pt-scan-fail-open-output", + enabled: true, + enforcement_mode: "monitor", + hook_point: "output", + fail_open: true, + kind: "keyword", + patterns: [{ kind: "literal", value: "FOR" }], + }); + await seed.createGuardrail({ + name: "pt-scan-fail-open-boundary", + enabled: true, + enforcement_mode: "monitor", + hook_point: "output", + fail_open: true, + kind: "keyword", + patterns: [{ kind: "regex", value: SPLIT_BLOCK_REGEX }], + }); + await seed.createApiKey({ key_hash: callerHash, allowed_models: [], allowed_routes: ["*"] }); + const proxy = new ProxyClient(app.proxyUrl, caller); + await waitConfigPropagation(async () => (await proxy.listModels()).status === 200); + }, 90_000); + + afterAll(async () => { + await app?.exit(); + await upstream?.close(); + await depthInputUpstream?.close(); + await unknownTypedUpstream?.close(); + await sls?.close(); + }); + + test("forwards Raw JSON beyond the depth cap only under input fail_open", async (ctx) => { + if (!etcdReachable || !app || !sls || !depthInputUpstream) return ctx.skip(); + + const before = depthInputUpstream.receivedRequests.length; + const res = await fetch(`${app.proxyUrl}/${depthInputRoute}/v1/any`, { + method: "POST", + headers: { authorization: `Bearer ${caller}`, "content-type": "application/json" }, + body: OVER_DEPTH_ESCAPED_BLOCK_JSON, + }); + expect(res.status).toBe(200); + expect(await res.text()).toBe(SAFE_ESCAPED_JSON); + expect(depthInputUpstream.receivedRequests.length).toBe(before + 1); + + const log = await waitForSlsLog( + sls, + logstore, + (entry) => entry.get("passthrough_route_name") === depthInputRoute, + "fail-open depth-capped passthrough usage event", + ); + expect(log.get("guardrail_blocked") ?? "false").not.toBe("true"); + expect(log.get("guardrail_bypassed_reason")).toBe("unscannable_body"); + }); + + test("forwards an unknown typed buffered response only under output fail_open", async (ctx) => { + if (!etcdReachable || !app || !sls || !unknownTypedUpstream) return ctx.skip(); + + const before = unknownTypedUpstream.receivedRequests.length; + const res = await fetch(`${app.proxyUrl}/${unknownTypedRoute}/v1/completions`, { + method: "POST", + headers: { authorization: `Bearer ${caller}`, "content-type": "application/json" }, + body: JSON.stringify({ model: "gpt-4o-mini", prompt: "go" }), + }); + const requestId = res.headers.get("x-aisix-request-id") ?? ""; + expect(requestId).toBeTruthy(); + expect(res.status).toBe(200); + expect(await res.text()).toBe(UNKNOWN_TYPED_BUFFERED_OUTPUT); + expect(unknownTypedUpstream.receivedRequests.length).toBe(before + 1); + + const log = await waitForSlsLog( + sls, + logstore, + (entry) => + entry.get("passthrough_route_name") === unknownTypedRoute && + entry.get("request_id") === requestId, + "fail-open unknown typed passthrough usage event", + ); + expect(log.get("guardrail_blocked") ?? "false").not.toBe("true"); + expect(log.get("guardrail_bypassed_reason")).toBe("unscannable_body"); + }); + + test("forwards nested Anthropic tool results beyond the work cap only under input fail_open", async (ctx) => { + if (!etcdReachable || !app || !sls || !depthInputUpstream) return ctx.skip(); + + const before = depthInputUpstream.receivedRequests.length; + const res = await fetch(`${app.proxyUrl}/${depthInputRoute}/v1/any`, { + method: "POST", + headers: { authorization: `Bearer ${caller}`, "content-type": "application/json" }, + body: OVER_DEPTH_ANTHROPIC_TOOL_RESULT_FAIL_OPEN_INPUT, + }); + expect(res.status).toBe(200); + const requestId = res.headers.get("x-aisix-request-id") ?? ""; + expect(requestId).toBeTruthy(); + expect(await res.text()).toBe(SAFE_ESCAPED_JSON); + expect(depthInputUpstream.receivedRequests.length).toBe(before + 1); + expect(depthInputUpstream.receivedRequests.at(-1)!.body).toBe( + OVER_DEPTH_ANTHROPIC_TOOL_RESULT_FAIL_OPEN_INPUT, + ); + + const log = await waitForSlsLog( + sls, + logstore, + (entry) => + entry.get("passthrough_route_name") === depthInputRoute && + entry.get("request_id") === requestId, + "fail-open nested-tool-result passthrough usage event", + ); + expect(log.get("guardrail_blocked") ?? "false").not.toBe("true"); + expect(log.get("guardrail_bypassed_reason")).toBe("unscannable_body"); + }); + + test("starts a new live scan epoch after an unkeyable Raw SSE object", async (ctx) => { + if (!etcdReachable || !app || !sls || !upstream) return ctx.skip(); + + const before = upstream.receivedRequests.length; + const res = await fetch(`${app.proxyUrl}/pt-scan-fail-open/v1/any`, { + method: "POST", + headers: { authorization: `Bearer ${caller}`, "content-type": "application/json" }, + body: `{"model":"fail-open-raw","stream":true,"state":"go"}`, + }); + expect(res.status).toBe(200); + const body = await res.text(); + expect(body).not.toContain("event: error"); + expect(body).toBe( + `data: ${RAW_PREFIX_SSE}\n\ndata: ${RAW_UNEVALUABLE_SSE}\n\ndata: ${RAW_SUFFIX_SSE}\n\ndata: [DONE]\n\n`, + ); + expect(upstream.receivedRequests.length).toBe(before + 1); + + const log = await waitForSlsLog( + sls, + logstore, + (entry) => entry.get("passthrough_route_name") === route, + "fail-open passthrough usage event", + ); + expect(log.get("guardrail_blocked") ?? "false").not.toBe("true"); + expect(log.get("guardrail_bypassed_reason")).toBe("unscannable_body"); + const hits = JSON.parse(log.get("guardrail_monitor_hits") ?? "[]") as Array<{ + action: string; + hook: string; + guardrail_name: string; + }>; + expect(hits).toContainEqual( + expect.objectContaining({ + action: "would_block", + hook: "output", + guardrail_name: "pt-scan-fail-open-output", + }), + ); + expect(hits).not.toContainEqual( + expect.objectContaining({ + action: "would_block", + hook: "output", + guardrail_name: "pt-scan-fail-open-boundary", + }), + ); + }); +}); + +// `fail_open` only permits an unevaluable frame to pass on a live stream. A +// blocking guardrail still holds its prefix until it scans clean, so that +// prefix must never be released just because the next frame is unevaluable. +describe("passthrough Raw stream held fail-open continuity", () => { + const caller = "sk-pt-scan-held-fail-open"; + const callerHash = createHash("sha256").update(caller).digest("hex"); + const route = "pt-scan-held-fail-open"; + let app: SpawnedApp | undefined; + let upstream: OpenAiUpstream | undefined; + let etcdReachable = false; + + beforeAll(async () => { + const etcd = new EtcdClient(); + etcdReachable = await etcd.ping(); + if (!etcdReachable) return; + + upstream = await startOpenAiUpstream({ + rawStreamFrames: [ + `data: ${RAW_HELD_BLOCK_SSE}\n\n`, + `data: ${RAW_UNEVALUABLE_SSE}\n\n`, + "data: [DONE]\n\n", + ], + }); + app = await spawnApp(); + const seed = new SeedClient(etcd, app.etcdPrefix); + const providerKey = await seed.createProviderKey({ + display_name: "pt-scan-held-fail-open-pk", + secret: "sk-mock", + api_base: upstream.baseUrl, + }); + await seed.createPassthroughRoute({ + name: route, + path_prefix: "/pt-scan-held-fail-open", + target_url: upstream.baseUrl, + provider_key_id: providerKey.id, + }); + await seed.createGuardrail({ + name: "pt-scan-held-fail-open-output", + enabled: true, + hook_point: "output", + fail_open: true, + kind: "keyword", + patterns: [{ kind: "literal", value: "FORBIDDEN" }], + }); + await seed.createApiKey({ key_hash: callerHash, allowed_models: [], allowed_routes: ["*"] }); + const proxy = new ProxyClient(app.proxyUrl, caller); + await waitConfigPropagation(async () => (await proxy.listModels()).status === 200); + }, 90_000); + + afterAll(async () => { + await app?.exit(); + await upstream?.close(); + }); + + test("refuses a held prefix before an unevaluable Raw SSE object can release it", async (ctx) => { + if (!etcdReachable || !app || !upstream) return ctx.skip(); + + const before = upstream.receivedRequests.length; + const res = await fetch(`${app.proxyUrl}/pt-scan-held-fail-open/v1/any`, { + method: "POST", + headers: { authorization: `Bearer ${caller}`, "content-type": "application/json" }, + body: `{"model":"held-fail-open-raw","stream":true,"state":"go"}`, + }); + expect(res.status).toBe(200); + const body = await res.text(); + expect(body).toContain("event: error"); + expect(body).toContain("guardrail_unavailable"); + expect(body).toContain("unscannable_body"); + expect(body).not.toContain("FORBIDDEN"); + expect(upstream.receivedRequests.length).toBe(before + 1); + }); }); diff --git a/tests/e2e/src/cases/ratelimit-cluster-e2e.test.ts b/tests/e2e/src/cases/ratelimit-cluster-e2e.test.ts index e5f0a1524..fa1de105e 100644 --- a/tests/e2e/src/cases/ratelimit-cluster-e2e.test.ts +++ b/tests/e2e/src/cases/ratelimit-cluster-e2e.test.ts @@ -1,4 +1,5 @@ import { createHash, randomUUID } from "node:crypto"; +import { request, type IncomingMessage } from "node:http"; import { connect, createServer, type Server, type Socket } from "node:net"; import { afterAll, beforeAll, describe, expect, test } from "vitest"; import { @@ -7,9 +8,12 @@ import { SeedClient, ProxyClient, spawnApp, + startMockSls, startOpenAiUpstream, awaitWindowHeadroom, waitConfigPropagation, + waitForSlsLog, + type MockSls, type OpenAiUpstream, type SpawnedApp, } from "../harness/index.js"; @@ -33,6 +37,23 @@ const CALLER_PLAINTEXT = "sk-rl-cluster-e2e-caller"; const CALLER_KEY_HASH = createHash("sha256") .update(CALLER_PLAINTEXT) .digest("hex"); +const PASSTHROUGH_CALLER_PLAINTEXT = "sk-rl-cluster-passthrough-caller"; +const PASSTHROUGH_CALLER_KEY_HASH = createHash("sha256") + .update(PASSTHROUGH_CALLER_PLAINTEXT) + .digest("hex"); +const PASSTHROUGH_ROUTE = "rl-cluster-passthrough"; +const PASSTHROUGH_PREFIX = "/rl-cluster-passthrough"; +// Short enough to prove a live stream renews its lease, while leaving a +// generous interval for CI scheduling around the three-second assertion. +const PASSTHROUGH_CONCURRENCY_TTL_SECS = 1; +const PASSTHROUGH_WAIT_BEYOND_TTL_MS = 3_000; +const PASSTHROUGH_502_SSE = 'event: upstream_error\ndata: {"message":"still-live"}\n\n'; +// Keep a full CI-scheduling cushion after the cross-TTL probe, then wait for +// this explicit EOF before asserting that the shared slot is released. +const PASSTHROUGH_502_EOF_DELAY_MS = PASSTHROUGH_WAIT_BEYOND_TTL_MS + 3_000; +const LEASE_LOSS_SLS_CREDENTIAL_REF = "rate_limit_lease_loss"; +const LEASE_LOSS_SLS_PROJECT = "aisix-e2e-obs"; +const LEASE_LOSS_SLS_LOGSTORE = "rate-limit-lease-loss"; const ETCD_ENDPOINT = etcdEndpoint(); const REDIS_URL = process.env.AISIX_E2E_REDIS ?? "redis://127.0.0.1:6379"; @@ -59,10 +80,9 @@ async function redisPing(url: string): Promise { /** * One command over a fresh RESP connection, as an array of bulk strings. * - * Only used to manage an ACL user: the refusal case needs a credential - * the server rejects and then accepts, and `requirepass` is server-wide - * while every other file in this suite shares the same Redis. An ACL - * user is scoped to itself and named per run. + * Used only for isolated Redis state in this suite. The ACL case needs a + * credential the server rejects and then accepts, and the lease-loss cases + * remove a member from their unique per-run concurrency bucket. */ async function redisCommand(url: string, args: string[]): Promise { const m = /^redis:\/\/(?:[^@/]*@)?([^:/]+)(?::(\d+))?/.exec(url); @@ -114,6 +134,46 @@ function chatRequest(proxyUrl: string, model: string): Promise { }); } +function rawPost( + url: string, + headers: Record, + body: string, +): Promise { + return new Promise((resolve, reject) => { + const req = request(url, { method: "POST", headers }, resolve); + req.once("error", reject); + req.end(body); + }); +} + +function readRawBody(response: IncomingMessage): Promise { + return new Promise((resolve, reject) => { + const chunks: Buffer[] = []; + response.on("data", (chunk: Buffer) => chunks.push(chunk)); + response.once("error", reject); + response.once("end", () => resolve(Buffer.concat(chunks))); + }); +} + +function beforeTimeout(promise: Promise, timeoutMs: number): Promise { + return new Promise((resolve, reject) => { + const timer = setTimeout( + () => reject(new Error(`stream did not terminate within ${timeoutMs}ms`)), + timeoutMs, + ); + void promise.then( + (value) => { + clearTimeout(timer); + resolve(value); + }, + (error: unknown) => { + clearTimeout(timer); + reject(error); + }, + ); + }); +} + /** Seed one model + an RPM=1 ApiKey into the SHARED config namespace — * both replicas pick it up over the same etcd watch. */ async function seed(etcdRoot: string, upstreamBase: string, model: string) { @@ -205,6 +265,384 @@ describe("rate limit is shared across replicas with backend=redis (#798)", () => }); }); +// E2E for #1737: a route's SSE response has already returned headers when +// its body remains live. The shared Redis semaphore must therefore stay held +// beyond its short crash-recovery TTL, across a different gateway process. +describe("passthrough SSE concurrency is shared and renewed across Redis replicas (#1737)", () => { + let appA: SpawnedApp | undefined; + let appB: SpawnedApp | undefined; + let upstream: OpenAiUpstream | undefined; + let sls: MockSls | undefined; + let infraReady = false; + let apiKeyId = ""; + const prefix = `/aisix-e2e-rl-passthrough-${randomUUID()}`; + + const headers = { + authorization: `Bearer ${PASSTHROUGH_CALLER_PLAINTEXT}`, + "content-type": "application/json", + }; + const call = (proxyUrl: string) => + fetch(`${proxyUrl}${PASSTHROUGH_PREFIX}/v1/chat/completions`, { + method: "POST", + headers, + body: JSON.stringify({ + model: "gpt-4o-mini", + messages: [{ role: "user", content: "hold this stream open" }], + stream: true, + }), + }); + + beforeAll(async () => { + infraReady = (await new EtcdClient().ping()) && (await redisPing(REDIS_URL)); + if (!infraReady) return; + + sls = await startMockSls(); + const streamEvents = [ + JSON.stringify({ choices: [{ delta: { content: "released" } }] }), + "[DONE]", + ]; + upstream = await startOpenAiUpstream({ + scriptedResponses: [ + // This stream flushes headers then stays open through the TTL + // boundary. The post-cancel request ends normally. + { streamEvents, firstEventDelayMs: 10_000 }, + { streamEvents }, + // The third request stays idle until the test removes its live Redis + // member. A fourth request proves another replica can reuse the slot + // only after the first body has been terminated. + { streamEvents, firstEventDelayMs: 10_000 }, + { streamEvents }, + ], + }); + const extra = { + etcd: sharedEtcd(prefix), + ratelimit: { + backend: "redis", + redis: { url: REDIS_URL }, + concurrency_ttl_secs: PASSTHROUGH_CONCURRENCY_TTL_SECS, + }, + }; + const extraEnv = { + [`SLS_CRED_${LEASE_LOSS_SLS_CREDENTIAL_REF.toUpperCase()}_AK_ID`]: "mock-akid", + [`SLS_CRED_${LEASE_LOSS_SLS_CREDENTIAL_REF.toUpperCase()}_AK_SECRET`]: "mock-secret", + }; + appA = await spawnApp({ extra, extraEnv }); + appB = await spawnApp({ extra, extraEnv }); + + const seed = new SeedClient(new EtcdClient(), prefix); + const providerKey = await seed.createProviderKey({ + display_name: "rl-cluster-passthrough-pk", + secret: "sk-mock", + api_base: "http://unused-on-passthrough-route", + }); + await seed.createPassthroughRoute({ + name: PASSTHROUGH_ROUTE, + path_prefix: PASSTHROUGH_PREFIX, + target_url: upstream.baseUrl, + provider_key_id: providerKey.id, + }); + await seed.createObservabilityExporter({ + name: "rate-limit-lease-loss-sls", + enabled: true, + kind: "aliyun_sls", + endpoint: sls.url, + project: LEASE_LOSS_SLS_PROJECT, + logstore: LEASE_LOSS_SLS_LOGSTORE, + credential_ref: LEASE_LOSS_SLS_CREDENTIAL_REF, + content_mode: "metadata_only", + }); + // Write the caller last, then wait for its local models surface on each + // replica. That proves the route and its ProviderKey reached the same + // snapshot without consuming either scripted stream. + const apiKey = await seed.createApiKey({ + key_hash: PASSTHROUGH_CALLER_KEY_HASH, + allowed_models: ["*"], + allowed_routes: [PASSTHROUGH_ROUTE], + rate_limit: { concurrency: 1 }, + }); + apiKeyId = apiKey.id; + + for (const app of [appA!, appB!]) { + const probe = new ProxyClient(app.proxyUrl, PASSTHROUGH_CALLER_PLAINTEXT); + await waitConfigPropagation(async () => (await probe.listModels()).status === 200); + } + }); + + afterAll(async () => { + await appA?.exit(); + await appB?.exit(); + await upstream?.close(); + await sls?.close(); + if (infraReady) await new EtcdClient().deletePrefix(prefix); + }); + + test( + "a live stream blocks the other replica past the TTL, then cancellation frees it", + async (ctx) => { + if (!infraReady || !appA || !appB || !upstream || !sls) { + ctx.skip(); + return; + } + + // The first response has headers but no event yet, so leaving its body + // unread precisely models a client consuming a still-live SSE stream. + const first = await call(appA.proxyUrl); + expect(first.status).toBe(200); + expect(first.headers.get("content-type") ?? "").toContain("text/event-stream"); + const upstreamCallsWhileHeld = upstream.receivedRequests.length; + expect(upstreamCallsWhileHeld).toBe(1); + + // A stale lease would be reclaimed after one second. Keep the stream + // alive much longer, then prove B still sees the same global cap. + await new Promise((resolve) => + setTimeout(resolve, PASSTHROUGH_WAIT_BEYOND_TTL_MS), + ); + const blocked = await call(appB.proxyUrl); + expect(blocked.status).toBe(429); + expect(blocked.headers.get("x-ratelimit-scope")).toBe("concurrency"); + await blocked.text(); + expect(upstream.receivedRequests).toHaveLength(upstreamCallsWhileHeld); + + // Client cancellation releases the remote lease. Polling B rules out a + // locally released A-only hold and waits for the Redis release to land. + expect(first.body).not.toBeNull(); + await first.body!.cancel(); + let admitted: Response | undefined; + await waitConfigPropagation(async () => { + const response = await call(appB!.proxyUrl); + if (response.status !== 200) { + await response.text(); + return false; + } + admitted = response; + return true; + }, 5_000); + expect(admitted).toBeDefined(); + expect(admitted!.headers.get("content-type") ?? "").toContain("text/event-stream"); + expect(await admitted!.text()).toContain("[DONE]"); + expect(upstream.receivedRequests).toHaveLength(upstreamCallsWhileHeld + 1); + + // A successful Redis response that says the member is absent is not a + // backend outage. Stop this idle body rather than let a second replica + // use the freed semaphore slot while it still forwards provider bytes. + let leaseLost: Response | undefined; + await waitConfigPropagation(async () => { + const response = await call(appA!.proxyUrl); + if (response.status !== 200) { + await response.text(); + return false; + } + leaseLost = response; + return true; + }, 5_000); + expect(leaseLost).toBeDefined(); + const leaseLossRequestId = leaseLost!.headers.get("x-aisix-request-id") ?? ""; + expect(leaseLossRequestId).not.toBe(""); + expect( + await redisCommand(REDIS_URL, ["DEL", `aisix:rl:{${apiKeyId}}:conc`]), + ).toBe(":1\r\n"); + expect(await beforeTimeout(leaseLost!.text(), 2_000)).toBe(""); + const leaseLossLog = await waitForSlsLog( + sls, + LEASE_LOSS_SLS_LOGSTORE, + (log) => log.get("request_id") === leaseLossRequestId, + `lease-loss telemetry for ${leaseLossRequestId}`, + 20_000, + ); + expect(leaseLossLog.get("status_code")).toBe("429"); + expect(leaseLossLog.get("error_class")).toBe("rate_limit_lease_lost"); + + let admittedAfterLeaseLoss: Response | undefined; + await waitConfigPropagation(async () => { + const response = await call(appB!.proxyUrl); + if (response.status !== 200) { + await response.text(); + return false; + } + admittedAfterLeaseLoss = response; + return true; + }, 5_000); + expect(admittedAfterLeaseLoss).toBeDefined(); + expect(await admittedAfterLeaseLoss!.text()).toContain("[DONE]"); + expect(upstream.receivedRequests).toHaveLength(upstreamCallsWhileHeld + 3); + }, + 30_000, + ); +}); + +// Non-success SSE bodies use the same streaming handoff as successful ones. +// This needs its own E2E because a 502 must keep the distributed slot until +// the upstream error body's EOF without rewriting its bytes. +describe("passthrough 502 SSE concurrency is shared through delayed EOF (#1737)", () => { + let appA: SpawnedApp | undefined; + let appB: SpawnedApp | undefined; + let upstream: OpenAiUpstream | undefined; + let infraReady = false; + let apiKeyId = ""; + const prefix = `/aisix-e2e-rl-passthrough-502-${randomUUID()}`; + const route = "rl-cluster-passthrough-502"; + const routePrefix = `/${route}`; + const caller = "sk-rl-cluster-passthrough-502"; + const callerHash = createHash("sha256").update(caller).digest("hex"); + const headers = { + authorization: `Bearer ${caller}`, + "content-type": "application/json", + }; + const body = JSON.stringify({ model: "gpt-4o-mini", stream: true }); + const call = (proxyUrl: string) => rawPost(`${proxyUrl}${routePrefix}/v1/any`, headers, body); + + beforeAll(async () => { + infraReady = (await new EtcdClient().ping()) && (await redisPing(REDIS_URL)); + if (!infraReady) return; + + upstream = await startOpenAiUpstream({ + scriptedResponses: [ + { + status: 502, + rawErrorBodyChunks: [PASSTHROUGH_502_SSE], + eventDelayMs: PASSTHROUGH_502_EOF_DELAY_MS, + responseHeaders: { "content-type": "text/event-stream; charset=utf-8" }, + }, + { + status: 502, + rawErrorBody: PASSTHROUGH_502_SSE, + responseHeaders: { "content-type": "text/event-stream; charset=utf-8" }, + }, + { + status: 502, + rawErrorBodyChunks: [PASSTHROUGH_502_SSE], + eventDelayMs: PASSTHROUGH_502_EOF_DELAY_MS, + responseHeaders: { "content-type": "text/event-stream; charset=utf-8" }, + }, + { + status: 502, + rawErrorBody: PASSTHROUGH_502_SSE, + responseHeaders: { "content-type": "text/event-stream; charset=utf-8" }, + }, + ], + }); + const extra = { + etcd: sharedEtcd(prefix), + ratelimit: { + backend: "redis", + redis: { url: REDIS_URL }, + concurrency_ttl_secs: PASSTHROUGH_CONCURRENCY_TTL_SECS, + }, + }; + appA = await spawnApp({ extra }); + appB = await spawnApp({ extra }); + + const seed = new SeedClient(new EtcdClient(), prefix); + const providerKey = await seed.createProviderKey({ + display_name: "rl-cluster-passthrough-502-pk", + secret: "sk-mock", + api_base: "http://unused-on-passthrough-route", + }); + await seed.createPassthroughRoute({ + name: route, + path_prefix: routePrefix, + target_url: upstream.baseUrl, + provider_key_id: providerKey.id, + }); + const apiKey = await seed.createApiKey({ + key_hash: callerHash, + allowed_models: ["*"], + allowed_routes: [route], + rate_limit: { concurrency: 1 }, + }); + apiKeyId = apiKey.id; + for (const app of [appA!, appB!]) { + const probe = new ProxyClient(app.proxyUrl, caller); + await waitConfigPropagation(async () => (await probe.listModels()).status === 200); + } + }); + + afterAll(async () => { + await appA?.exit(); + await appB?.exit(); + await upstream?.close(); + if (infraReady) await new EtcdClient().deletePrefix(prefix); + }); + + test( + "an unread 502 SSE holds the shared slot beyond TTL, then EOF releases it without changing bytes", + async (ctx) => { + if (!infraReady || !appA || !appB || !upstream) { + ctx.skip(); + return; + } + + const first = await call(appA.proxyUrl); + expect(first.statusCode).toBe(502); + expect(first.headers["content-type"]).toBe("text/event-stream; charset=utf-8"); + const upstreamCallsWhileHeld = upstream.receivedRequests.length; + expect(upstreamCallsWhileHeld).toBe(1); + + // The client has response headers but no body reader. The source's + // delayed EOF crosses the one-second distributed lease TTL. + await new Promise((resolve) => setTimeout(resolve, PASSTHROUGH_WAIT_BEYOND_TTL_MS)); + const blocked = await call(appB.proxyUrl); + expect(blocked.statusCode).toBe(429); + expect(blocked.headers["x-ratelimit-scope"]).toBe("concurrency"); + await readRawBody(blocked); + expect(upstream.receivedRequests).toHaveLength(upstreamCallsWhileHeld); + + expect(await readRawBody(first)).toEqual(Buffer.from(PASSTHROUGH_502_SSE)); + let released: IncomingMessage | undefined; + await waitConfigPropagation(async () => { + const response = await call(appB!.proxyUrl); + if (response.statusCode !== 502) { + await readRawBody(response); + return false; + } + released = response; + return true; + }, 5_000); + expect(released).toBeDefined(); + expect(await readRawBody(released!)).toEqual(Buffer.from(PASSTHROUGH_502_SSE)); + expect(upstream.receivedRequests).toHaveLength(upstreamCallsWhileHeld + 1); + + // Opaque error representations must remain byte-for-byte upstream + // data. A definitive loss ends at EOF; it must not inject an SSE or + // JSON error into the provider's 502 body. + let leaseLost: IncomingMessage | undefined; + await waitConfigPropagation(async () => { + const response = await call(appA!.proxyUrl); + if (response.statusCode !== 502) { + await readRawBody(response); + return false; + } + leaseLost = response; + return true; + }, 5_000); + expect(leaseLost).toBeDefined(); + expect( + await redisCommand(REDIS_URL, ["DEL", `aisix:rl:{${apiKeyId}}:conc`]), + ).toBe(":1\r\n"); + expect(await beforeTimeout(readRawBody(leaseLost!), 2_000)).toEqual( + Buffer.from(PASSTHROUGH_502_SSE), + ); + + let admittedAfterLeaseLoss: IncomingMessage | undefined; + await waitConfigPropagation(async () => { + const response = await call(appB!.proxyUrl); + if (response.statusCode !== 502) { + await readRawBody(response); + return false; + } + admittedAfterLeaseLoss = response; + return true; + }, 5_000); + expect(admittedAfterLeaseLoss).toBeDefined(); + expect(await readRawBody(admittedAfterLeaseLoss!)).toEqual( + Buffer.from(PASSTHROUGH_502_SSE), + ); + expect(upstream.receivedRequests).toHaveLength(upstreamCallsWhileHeld + 3); + }, + 30_000, + ); +}); + describe("rate limit is NOT shared with backend=memory (per-replica, the #798 bug)", () => { let appA: SpawnedApp | undefined; let appB: SpawnedApp | undefined; diff --git a/tests/e2e/src/harness/upstream-openai.ts b/tests/e2e/src/harness/upstream-openai.ts index 441b9e6fa..7b4b74688 100644 --- a/tests/e2e/src/harness/upstream-openai.ts +++ b/tests/e2e/src/harness/upstream-openai.ts @@ -34,7 +34,12 @@ export interface OpenAiUpstreamOptions { * Error body written VERBATIM instead of JSON-encoding `errorBody` — an * empty string reproduces an upstream that answers with no body at all. */ - rawErrorBody?: string; + rawErrorBody?: string | Buffer; + /** + * Error body chunks written verbatim after headers. Together with + * `eventDelayMs`, this models a non-2xx SSE response whose EOF is delayed. + */ + rawErrorBodyChunks?: Array; /** * Content-Type for the error body (default `application/json`). Lets * tests reproduce upstreams / edge layers that return a JSON error @@ -52,7 +57,7 @@ export interface OpenAiUpstreamOptions { * gateway streams provider bytes back and injected the provider bearer on * the content GET. */ - rawBody?: string; + rawBody?: string | Buffer; /** * `rawBody` split into chunks written one at a time, `eventDelayMs` * apart, so a spec can tell a relayed body from a buffered one: an @@ -60,7 +65,7 @@ export interface OpenAiUpstreamOptions { * downstream if the gateway forwards them as they arrive. Takes * precedence over `rawBody`; no Content-Length is sent. */ - rawBodyChunks?: string[]; + rawBodyChunks?: Array; /** Content-Type for `rawBody` (default `application/octet-stream`). */ rawContentType?: string; /** Per-request response script; used in order before static opts. */ @@ -92,16 +97,18 @@ export interface OpenAiUpstreamStep { status?: number; errorBody?: unknown; /** See `OpenAiUpstreamOptions.rawErrorBody`. */ - rawErrorBody?: string; + rawErrorBody?: string | Buffer; + /** See `OpenAiUpstreamOptions.rawErrorBodyChunks`. */ + rawErrorBodyChunks?: Array; /** Content-Type for the error body (default `application/json`). See #543. */ errorContentType?: string; disconnectAfterEvents?: number; /** Extra response headers, same semantics as on the top-level options. */ responseHeaders?: Record; /** Raw (non-JSON) 200 body — see `OpenAiUpstreamOptions.rawBody`. */ - rawBody?: string; + rawBody?: string | Buffer; /** See `OpenAiUpstreamOptions.rawBodyChunks`. */ - rawBodyChunks?: string[]; + rawBodyChunks?: Array; /** Content-Type for `rawBody` (default `application/octet-stream`). */ rawContentType?: string; } @@ -181,6 +188,22 @@ export async function startOpenAiUpstream( const status = step.status ?? 200; if (status >= 300) { res.statusCode = status; + if (step.rawErrorBodyChunks !== undefined) { + if (!res.hasHeader("content-type")) { + res.setHeader( + "content-type", + step.errorContentType ?? opts.errorContentType ?? "application/json", + ); + } + res.flushHeaders(); + for (const chunk of step.rawErrorBodyChunks) { + if (res.writableEnded || res.destroyed) return; + res.write(Buffer.isBuffer(chunk) ? chunk : Buffer.from(chunk)); + if (step.eventDelayMs) await sleep(step.eventDelayMs); + } + if (!res.writableEnded && !res.destroyed) res.end(); + return; + } if (step.rawErrorBody !== undefined) { res.end(step.rawErrorBody); return;