From 71d4e05cd5ee16f94da8c2ef57e419a560c3858a Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Wed, 30 Sep 2026 08:04:15 +0800 Subject: [PATCH 01/37] fix(passthrough): preserve route and stream boundaries --- Cargo.lock | 1 + crates/aisix-proxy/Cargo.toml | 1 + crates/aisix-proxy/src/passthrough_route.rs | 273 +++++++++++++++++- .../src/cases/passthrough-route-e2e.test.ts | 181 ++++++++++++ 4 files changed, 445 insertions(+), 11 deletions(-) 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-proxy/Cargo.toml b/crates/aisix-proxy/Cargo.toml index b5c4a1376..a03c50c47 100644 --- a/crates/aisix-proxy/Cargo.toml +++ b/crates/aisix-proxy/Cargo.toml @@ -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/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index 608a5e752..fc756f2a3 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}; @@ -521,15 +523,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. @@ -649,7 +644,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))?; @@ -924,6 +919,7 @@ async fn dispatch( status, telemetry, &client.request_id, + reservation.into_stream_hold(), )); } @@ -1719,7 +1715,14 @@ fn is_api_version_segment(seg: &str) -> bool { /// `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(""); + 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; } @@ -1734,6 +1737,135 @@ fn strip_redundant_version_segment<'a>(base: &str, rest: &'a str) -> &'a str { 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 first uses form decoding, then a bounded number of +/// percent-decoding passes. That matches the path guard: a backend must not +/// be able to turn a nested encoding into an operator-owned key downstream. +/// 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 base_keys = url::form_urlencoded::parse(base.as_bytes()) + .map(|(key, _)| normalize_query_key(&key)) + .collect::>>(); + let Some(base_keys) = base_keys else { + return true; + }; + url::form_urlencoded::parse(inbound.as_bytes()).any(|(key, _)| { + let Some(key) = normalize_query_key(&key) else { + return true; + }; + base_keys.iter().any(|base_key| base_key == &key) + }) +} + +fn normalize_query_key(key: &str) -> Option> { + let mut decoded = key.as_bytes().to_vec(); + for _ in 0..=MAX_PERCENT_DECODE_PASSES { + let next = percent_decode(&decoded).collect::>(); + if next == decoded { + return Some(decoded); + } + decoded = next; + } + None +} + +/// 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| segment == b"." || segment == 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 // --------------------------------------------------------------------------- @@ -2023,6 +2155,7 @@ fn stream_response( status: reqwest::StatusCode, mut telemetry: RouteTelemetry, request_id: &str, + stream_hold: aisix_ratelimit::StreamConcurrencyGuard, ) -> Response { use aisix_guardrails::{Guardrail as _, GuardrailVerdict, StreamOutputPolicy}; use futures::StreamExt; @@ -2036,6 +2169,10 @@ fn stream_response( let capture_cap = telemetry.content_cap; 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 mut upstream = upstream_resp.bytes_stream(); let mut splitter = SseFrameSplitter::new(); // Held-back frames (Window / BufferFull) not yet released. @@ -3068,6 +3205,120 @@ mod tests { 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" + ); + + // 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", + ] { + 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() diff --git a/tests/e2e/src/cases/passthrough-route-e2e.test.ts b/tests/e2e/src/cases/passthrough-route-e2e.test.ts index 33b0141d0..d3fd159a3 100644 --- a/tests/e2e/src/cases/passthrough-route-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-route-e2e.test.ts @@ -46,6 +46,10 @@ 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"); describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed paths", () => { let app: SpawnedApp | undefined; @@ -221,6 +225,70 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat expect(await res.text()).toBe(""); }); + test("a nested-encoded route remainder never reaches 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, + }); + + const headers = { authorization: `Bearer ${CALLER_PLAINTEXT}` }; + await waitConfigPropagation(async () => { + try { + const ready = await fetch(`${app!.proxyUrl}/ptr-boundary/models`, { + headers, + }); + await ready.text(); + return ready.status === 200; + } catch { + return false; + } + }); + + const baseline = upstream.receivedRequests.length; + // Use undici's raw request helper: fetch implementations are allowed to + // normalize URL escapes before the gateway receives the wire path. + const rejected = await harnessRequest( + `${app.proxyUrl}/ptr-boundary/%252e%252e%252fmodels`, + { headers }, + ); + expect(rejected.statusCode).toBe(400); + await rejected.body.text(); + expect(upstream.receivedRequests).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); + }); + test("forward-proxy BYO: host match beats typed routes; Authorization forwarded verbatim", async (ctx) => { if (!etcdReachable || !app || !seed) { ctx.skip(); @@ -494,6 +562,119 @@ 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({ + // The propagation probe does not reach this upstream. The first real + // stream ends normally; the second stalls for far longer than the + // bounded cancellation assertion, so a natural upstream close cannot + // make a leaked concurrency reservation look released. + scriptedResponses: [ + { streamEvents }, + { streamEvents, firstEventDelayMs: 900, eventDelayMs: 600 }, + { streamEvents, firstEventDelayMs: 10_000 }, + { 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 successful authenticated + // readiness probe proves the preceding route has reached the same + // snapshot without spending its concurrency slot. + 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 ready = await fetch(`${app!.proxyUrl}/v1/models`, { headers }); + await ready.text(); + return ready.status === 200; + } 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(1); + + 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); + + await first.text(); + const afterEnd = await call(); + expect(afterEnd.status).toBe(200); + + // Keep a second stream open to exercise cancellation separately. + expect(afterEnd.body).not.toBeNull(); + const secondStream = afterEnd; + + // A cancelled client must release the same hold. The next caller should + // not wait for the upstream's delayed frames to finish naturally. + await secondStream.body!.cancel(); + await waitConfigPropagation(async () => { + try { + const afterCancel = await call(); + const admitted = afterCancel.status === 200; + await afterCancel.body?.cancel(); + return admitted; + } catch { + return false; + } + }, 3_000); + + expect(upstream.receivedRequests).toHaveLength(upstreamCallsWhileStreaming + 2); + }); + test("envelope auto-detection: usage follows the request body, never the config", async (ctx) => { if (!etcdReachable || !app || !seed) { ctx.skip(); From 7a2937803e521c77647d6e443ffcf24ea6a77bc1 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Wed, 30 Sep 2026 08:48:32 +0800 Subject: [PATCH 02/37] fix(passthrough): harden query and stream coverage --- crates/aisix-proxy/src/passthrough_route.rs | 58 +++++++++++------ .../src/cases/passthrough-route-e2e.test.ts | 62 ++++++++++++++----- 2 files changed, 85 insertions(+), 35 deletions(-) diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index fc756f2a3..3b074e1d1 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -1795,32 +1795,36 @@ fn join_target_url(base: &str, rest: &str, query: Option<&str>) -> Result bool { - let base_keys = url::form_urlencoded::parse(base.as_bytes()) - .map(|(key, _)| normalize_query_key(&key)) - .collect::>>(); - let Some(base_keys) = base_keys else { + let Some(base_keys) = query_keys_at_all_decode_levels(base) else { return true; }; - url::form_urlencoded::parse(inbound.as_bytes()).any(|(key, _)| { - let Some(key) = normalize_query_key(&key) else { - return true; - }; - base_keys.iter().any(|base_key| base_key == &key) - }) + 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 normalize_query_key(key: &str) -> Option> { - let mut decoded = key.as_bytes().to_vec(); +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 (key, _) in url::form_urlencoded::parse(&decoded) { + let key = key.into_owned().into_bytes(); + if !keys.iter().any(|existing| existing == &key) { + keys.push(key); + } + } let next = percent_decode(&decoded).collect::>(); if next == decoded { - return Some(decoded); + return Some(keys); } decoded = next; } @@ -3263,6 +3267,24 @@ mod tests { .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" + ); // Test raw, percent-encoded, and encoded-separator spellings. The // gateway must reject them before a ProviderKey can be sent outside diff --git a/tests/e2e/src/cases/passthrough-route-e2e.test.ts b/tests/e2e/src/cases/passthrough-route-e2e.test.ts index d3fd159a3..f888c049f 100644 --- a/tests/e2e/src/cases/passthrough-route-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-route-e2e.test.ts @@ -261,6 +261,19 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat } }); + expect(upstream.receivedRequests.at(-1)?.path).toBe( + "/provider/v1/models?tenant=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; // Use undici's raw request helper: fetch implementations are allowed to // normalize URL escapes before the gateway receives the wire path. @@ -287,6 +300,22 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat 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); }); test("forward-proxy BYO: host match beats typed routes; Authorization forwarded verbatim", async (ctx) => { @@ -578,13 +607,11 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat ]; const upstream = await startOpenAiUpstream({ // The propagation probe does not reach this upstream. The first real - // stream ends normally; the second stalls for far longer than the - // bounded cancellation assertion, so a natural upstream close cannot - // make a leaked concurrency reservation look released. + // stream stalls until the client cancels it. The second ends naturally, + // so this test covers both lifetime boundaries of the same reservation. scriptedResponses: [ - { streamEvents }, - { streamEvents, firstEventDelayMs: 900, eventDelayMs: 600 }, { streamEvents, firstEventDelayMs: 10_000 }, + { streamEvents, firstEventDelayMs: 25, eventDelayMs: 25 }, { streamEvents }, ], }); @@ -650,27 +677,28 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat await second.text(); expect(upstream.receivedRequests).toHaveLength(upstreamCallsWhileStreaming); - await first.text(); - const afterEnd = await call(); - expect(afterEnd.status).toBe(200); - - // Keep a second stream open to exercise cancellation separately. - expect(afterEnd.body).not.toBeNull(); - const secondStream = afterEnd; - - // A cancelled client must release the same hold. The next caller should - // not wait for the upstream's delayed frames to finish naturally. - await secondStream.body!.cancel(); + // 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 () => { try { const afterCancel = await call(); const admitted = afterCancel.status === 200; - await afterCancel.body?.cancel(); + if (admitted) naturallyEnding = afterCancel; + else await afterCancel.text(); return admitted; } catch { return false; } }, 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); }); From 268781a9b25ea0f7d530d014d519a91ed3219c66 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Wed, 30 Sep 2026 09:02:16 +0800 Subject: [PATCH 03/37] fix(passthrough): harden matrix route boundaries --- crates/aisix-proxy/src/passthrough_route.rs | 45 +++++++++++++++---- .../src/cases/passthrough-route-e2e.test.ts | 38 ++++++++++++---- 2 files changed, 66 insertions(+), 17 deletions(-) diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index 3b074e1d1..bec9f955e 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -1796,10 +1796,10 @@ fn join_target_url(base: &str, rest: &str, query: Option<&str>) -> Result bool { let Some(base_keys) = query_keys_at_all_decode_levels(base) else { return true; @@ -1816,10 +1816,12 @@ 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 (key, _) in url::form_urlencoded::parse(&decoded) { - let key = key.into_owned().into_bytes(); - if !keys.iter().any(|existing| existing == &key) { - keys.push(key); + for field in decoded.split(|byte| matches!(*byte, b'&' | b';')) { + for (key, _) in url::form_urlencoded::parse(field) { + let key = key.into_owned().into_bytes(); + if !keys.iter().any(|existing| existing == &key) { + keys.push(key); + } } } let next = percent_decode(&decoded).collect::>(); @@ -1842,7 +1844,13 @@ fn has_path_traversal_segment(path: &str) -> bool { for _ in 0..=MAX_PERCENT_DECODE_PASSES { if decoded .split(|byte| matches!(*byte, b'/' | b'\\')) - .any(|segment| segment == b"." || segment == b"..") + .any(|segment| { + let path_part = segment + .split(|byte| *byte == b';') + .next() + .unwrap_or_default(); + path_part == b"." || path_part == b".." + }) { return true; } @@ -3285,6 +3293,21 @@ mod tests { .is_err(), "a nested-encoded query delimiter must not recreate an operator-owned key" ); + 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" + ); + } // Test raw, percent-encoded, and encoded-separator spellings. The // gateway must reject them before a ProviderKey can be sent outside @@ -3300,6 +3323,10 @@ mod tests { "%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(), diff --git a/tests/e2e/src/cases/passthrough-route-e2e.test.ts b/tests/e2e/src/cases/passthrough-route-e2e.test.ts index f888c049f..93d7b631a 100644 --- a/tests/e2e/src/cases/passthrough-route-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-route-e2e.test.ts @@ -225,7 +225,7 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat expect(await res.text()).toBe(""); }); - test("a nested-encoded route remainder never reaches the configured upstream", async (ctx) => { + test("route boundary traversal and query conflicts never reach the configured upstream", async (ctx) => { if (!etcdReachable || !app || !seed) { ctx.skip(); return; @@ -277,13 +277,21 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat const baseline = upstream.receivedRequests.length; // Use undici's raw request helper: fetch implementations are allowed to // normalize URL escapes before the gateway receives the wire path. - const rejected = await harnessRequest( - `${app.proxyUrl}/ptr-boundary/%252e%252e%252fmodels`, - { headers }, - ); - expect(rejected.statusCode).toBe(400); - await rejected.body.text(); - expect(upstream.receivedRequests).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`, @@ -316,6 +324,20 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat expect(nestedEncodedQueryDelimiter.statusCode).toBe(400); await nestedEncodedQueryDelimiter.body.text(); expect(upstream.receivedRequests).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); + } }); test("forward-proxy BYO: host match beats typed routes; Authorization forwarded verbatim", async (ctx) => { From 331eb57f19440d5f2f339d53426b9eb9e00377c4 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Wed, 30 Sep 2026 09:18:00 +0800 Subject: [PATCH 04/37] fix(passthrough): normalize conflicting query keys --- crates/aisix-proxy/src/passthrough_route.rs | 46 ++++++++++++++++++- .../src/cases/passthrough-route-e2e.test.ts | 34 +++++++++++++- 2 files changed, 76 insertions(+), 4 deletions(-) diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index bec9f955e..ecc4beeea 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -1798,7 +1798,9 @@ fn join_target_url(base: &str, rest: &str, query: Option<&str>) -> Result bool { let Some(base_keys) = query_keys_at_all_decode_levels(base) else { @@ -1818,7 +1820,7 @@ fn query_keys_at_all_decode_levels(query: &str) -> Option>> { 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 = key.into_owned().into_bytes(); + let key = normalize_form_query_key(&key); if !keys.iter().any(|existing| existing == &key) { keys.push(key); } @@ -1833,6 +1835,20 @@ fn query_keys_at_all_decode_levels(query: &str) -> Option>> { None } +/// Match the key canonicalization used by common form parsers: a bracketed +/// suffix selects a nested value beneath the base key, and ASCII dots/spaces +/// are aliases for underscores. +fn normalize_form_query_key(key: &str) -> Vec { + key.as_bytes() + .iter() + .take_while(|byte| **byte != 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 @@ -3308,6 +3324,32 @@ mod tests { "{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 diff --git a/tests/e2e/src/cases/passthrough-route-e2e.test.ts b/tests/e2e/src/cases/passthrough-route-e2e.test.ts index 93d7b631a..0b09540fc 100644 --- a/tests/e2e/src/cases/passthrough-route-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-route-e2e.test.ts @@ -247,11 +247,17 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat 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-boundary/models`, { + const ready = await fetch(`${app!.proxyUrl}/ptr-form-key-boundary/models`, { headers, }); await ready.text(); @@ -262,7 +268,7 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat }); expect(upstream.receivedRequests.at(-1)?.path).toBe( - "/provider/v1/models?tenant=operator", + "/provider/v1/models?tenant_id=operator", ); const allowedQuery = await harnessRequest( `${app.proxyUrl}/ptr-boundary/models?limit=3`, @@ -338,6 +344,30 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat 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) => { From 9aee6af2d837253095af300d35e76f59dc778f18 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Wed, 30 Sep 2026 09:29:24 +0800 Subject: [PATCH 05/37] fix(passthrough): canonicalize form query keys --- crates/aisix-proxy/src/passthrough_route.rs | 26 ++++++++++++++++--- .../src/cases/passthrough-route-e2e.test.ts | 16 ++++++++++++ 2 files changed, 38 insertions(+), 4 deletions(-) diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index ecc4beeea..2aec1583a 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -1835,13 +1835,14 @@ fn query_keys_at_all_decode_levels(query: &str) -> Option>> { None } -/// Match the key canonicalization used by common form parsers: a bracketed -/// suffix selects a nested value beneath the base key, and ASCII dots/spaces -/// are aliases for underscores. +/// 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() - .take_while(|byte| **byte != b'[') + .skip_while(|byte| **byte == b' ') + .take_while(|byte| !matches!(**byte, b'\0' | b'[')) .map(|byte| match *byte { b'.' | b' ' => b'_', byte => byte, @@ -3309,6 +3310,23 @@ mod tests { .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", diff --git a/tests/e2e/src/cases/passthrough-route-e2e.test.ts b/tests/e2e/src/cases/passthrough-route-e2e.test.ts index 0b09540fc..794e3d270 100644 --- a/tests/e2e/src/cases/passthrough-route-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-route-e2e.test.ts @@ -331,6 +331,22 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat 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", From 96bdf1c654ace5914d7e2b92f12a053f4e28aa38 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Wed, 30 Sep 2026 10:35:07 +0800 Subject: [PATCH 06/37] fix(ratelimit): renew stream concurrency leases --- crates/aisix-ratelimit/src/limiter.rs | 44 +++++++++- crates/aisix-ratelimit/src/store/mod.rs | 13 +++ crates/aisix-ratelimit/src/store/redis.rs | 84 ++++++++++++++++--- .../tests/redis_integration.rs | 50 ++++++++++- 4 files changed, 176 insertions(+), 15 deletions(-) diff --git a/crates/aisix-ratelimit/src/limiter.rs b/crates/aisix-ratelimit/src/limiter.rs index 41315e52c..e367afaeb 100644 --- a/crates/aisix-ratelimit/src/limiter.rs +++ b/crates/aisix-ratelimit/src/limiter.rs @@ -110,6 +110,7 @@ impl Limiter { store: Arc::clone(&self.store), key: key.to_string(), member, + has_concurrency_slot: limits.concurrency.is_some(), committed: false, }) } @@ -147,6 +148,7 @@ pub struct Reservation { store: Arc, key: String, member: String, + has_concurrency_slot: bool, committed: bool, } @@ -213,8 +215,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 +226,8 @@ 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_interval = None; + let mut refresh_holds = Vec::new(); let holds = self .reservations .iter_mut() @@ -230,11 +235,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); + if r.has_concurrency_slot { + if let Some(interval) = store.stream_lease_refresh_interval() { + refresh_interval = Some( + refresh_interval.map_or(interval, |current| current.min(interval)), + ); + refresh_holds.push((Arc::clone(&store), r.key.clone(), r.member.clone())); + } + } + (store, r.key.clone(), r.member.clone()) }) .collect(); + // Proxy handlers create streaming guards inside Tokio. Keep ordinary + // callers that do not have a runtime from panicking; they retain the + // backend's normal stale-lease recovery behavior instead. + let refresh_task = refresh_interval.and_then(|interval| { + tokio::runtime::Handle::try_current().ok().map(|handle| { + handle.spawn(async move { + loop { + tokio::time::sleep(interval).await; + for (store, key, member) in &refresh_holds { + store.refresh_stream_lease(key, member).await; + } + } + }) + }) + }); StreamConcurrencyGuard { holds, + refresh_task, released: false, } } @@ -255,6 +285,10 @@ 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. It is stopped before the terminal release to prevent a late + /// renewal from racing stream teardown. + refresh_task: Option>, released: bool, } @@ -264,6 +298,9 @@ impl StreamConcurrencyGuard { return; } self.released = true; + if let Some(task) = self.refresh_task.take() { + task.abort(); + } for (store, key, member) in &self.holds { store.release(key, member); } @@ -274,6 +311,7 @@ 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_task.is_some()) .field("released", &self.released) .finish() } diff --git a/crates/aisix-ratelimit/src/store/mod.rs b/crates/aisix-ratelimit/src/store/mod.rs index 841059cd8..05aba0fce 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; @@ -127,6 +129,17 @@ 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) {} + /// 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..176d38aa0 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. @@ -88,10 +88,18 @@ local conc_ttl = tonumber(ARGV[4]) local grace = tonumber(ARGV[5]) 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} @@ -143,12 +151,30 @@ for i = 1, nreq do end end if conc_max >= 0 then - redis.call('ZADD', conc_key, now, member) + redis.call('ZADD', conc_key, conc_now, member) redis.call('EXPIRE', conc_key, conc_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_ttl. +const REFRESH_CONCURRENCY_LUA: &str = r#" +local prefix = ARGV[1] +local member = ARGV[2] +local conc_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_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 +227,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)} "#; @@ -692,6 +722,40 @@ 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 prefix = self.bucket_prefix(key); + let mut conn = match self.conn.acquire().await { + Ok(c) => c, + Err(e) => { + self.note_failure("refresh", &e); + return; + } + }; + let res: Result = Script::new(REFRESH_CONCURRENCY_LUA) + .key(&prefix) + .arg(&prefix) + .arg(member) + .arg(self.conc_ttl) + .invoke_async(&mut conn) + .await; + match res { + Ok(_) => self.mark_ok(), + Err(e) => { + self.note_failure("refresh", &e); + self.conn.note_error().await; + } + } + } + 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..6721e3671 100644 --- a/crates/aisix-ratelimit/tests/redis_integration.rs +++ b/crates/aisix-ratelimit/tests/redis_integration.rs @@ -6,11 +6,12 @@ //! 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::{Limiter, MultiReservation, RateStore, RedisStore}; fn redis_url() -> Option { std::env::var("RATELIMIT_TEST_REDIS_URL").ok() @@ -390,12 +391,57 @@ 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 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"); +} + /// 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. From 8ed8d6bab4e0836280caaf17fcb201694b2edd10 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Wed, 30 Sep 2026 10:38:29 +0800 Subject: [PATCH 07/37] fix(ratelimit): type lease refresh interval --- crates/aisix-ratelimit/src/limiter.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/crates/aisix-ratelimit/src/limiter.rs b/crates/aisix-ratelimit/src/limiter.rs index e367afaeb..30cd9bf86 100644 --- a/crates/aisix-ratelimit/src/limiter.rs +++ b/crates/aisix-ratelimit/src/limiter.rs @@ -226,7 +226,7 @@ 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_interval = None; + let mut refresh_interval: Option = None; let mut refresh_holds = Vec::new(); let holds = self .reservations From 64fa1b19273aa6488c0ab6657d03f12cb6c54f9e Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Wed, 30 Sep 2026 10:53:59 +0800 Subject: [PATCH 08/37] fix(ratelimit): protect stream leases during upgrades --- crates/aisix-ratelimit/src/store/redis.rs | 18 ++-- .../tests/redis_integration.rs | 96 ++++++++++++++++++- 2 files changed, 105 insertions(+), 9 deletions(-) diff --git a/crates/aisix-ratelimit/src/store/redis.rs b/crates/aisix-ratelimit/src/store/redis.rs index 176d38aa0..e8d8a8f16 100644 --- a/crates/aisix-ratelimit/src/store/redis.rs +++ b/crates/aisix-ratelimit/src/store/redis.rs @@ -85,7 +85,8 @@ 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 @@ -106,7 +107,7 @@ if conc_max >= 0 then end end -local idx = 6 +local idx = 7 local nreq = tonumber(ARGV[idx]); idx = idx + 1 local req = {} for i = 1, nreq do @@ -152,18 +153,18 @@ for i = 1, nreq do end if conc_max >= 0 then redis.call('ZADD', conc_key, conc_now, member) - redis.call('EXPIRE', conc_key, conc_ttl) + 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_ttl. +/// 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_ttl = tonumber(ARGV[3]) +local conc_key_ttl = tonumber(ARGV[3]) local conc_key = prefix .. ':conc' if not redis.call('ZSCORE', conc_key, member) then return 0 @@ -171,7 +172,7 @@ 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_ttl) +redis.call('EXPIRE', conc_key, conc_key_ttl) return 1 "#; @@ -532,6 +533,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)); @@ -744,7 +748,7 @@ impl RateStore for RedisStore { .key(&prefix) .arg(&prefix) .arg(member) - .arg(self.conc_ttl) + .arg(self.conc_ttl.saturating_add(1)) .invoke_async(&mut conn) .await; match res { diff --git a/crates/aisix-ratelimit/tests/redis_integration.rs b/crates/aisix-ratelimit/tests/redis_integration.rs index 6721e3671..ee574b63b 100644 --- a/crates/aisix-ratelimit/tests/redis_integration.rs +++ b/crates/aisix-ratelimit/tests/redis_integration.rs @@ -11,7 +11,9 @@ use std::time::Duration; use aisix_core::{RateLimit, RateLimitScope, RedisConnConfig, RedisMode}; use aisix_obs::metrics::Metrics; -use aisix_ratelimit::{Limiter, MultiReservation, RateStore, RedisStore}; +use aisix_ratelimit::{ + store::redis::DEFAULT_PREFIX, Limiter, MultiReservation, RateStore, RedisStore, +}; fn redis_url() -> Option { std::env::var("RATELIMIT_TEST_REDIS_URL").ok() @@ -373,7 +375,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"); @@ -397,6 +400,95 @@ async fn stale_concurrency_slot_is_reclaimed_after_ttl() { .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. A queued lease + // refresh that follows must see it missing rather than add it back. + a.commit(&key, 0, "a-stream").await; + a.refresh_stream_lease(&key, "a-stream").await; + + 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 { From 84f301598e05abbd794c59edc9f8dd07dbb777af Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Wed, 30 Sep 2026 12:11:31 +0800 Subject: [PATCH 09/37] fix(passthrough): cover forwarded JSON guardrail fields --- crates/aisix-proxy/src/passthrough_route.rs | 367 +++++++++++++++--- .../passthrough-scan-coverage-e2e.test.ts | 134 ++++++- 2 files changed, 444 insertions(+), 57 deletions(-) diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index 2aec1583a..040395a0f 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -992,7 +992,7 @@ async fn dispatch( merge_usage(&mut telemetry.usage, u); } if telemetry.content_cap.is_some() { - telemetry.response_text = response_guardrail_text(protocol, &resp_body); + telemetry.response_text = response_capture_text(protocol, &resp_body); } let mut response = Response::builder() @@ -1221,8 +1221,9 @@ fn anthropic_message_output_text(v: &serde_json::Value) -> String { /// 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`), @@ -1347,31 +1348,210 @@ fn body_model_name( .unwrap_or_default() } +fn append_scan_text(out: &mut String, text: &str) { + if text.is_empty() { + return; + } + if !out.is_empty() { + out.push('\n'); + } + out.push_str(text); +} + +/// A duplicate-preserving JSON value walker. `serde_json::Value` is right for +/// typed envelope extraction but keeps only the final value for a repeated +/// object key; passthrough forwards the original bytes, so guardrail scanning +/// must see every decoded string value the upstream can see. Object keys are +/// structural metadata and deliberately stay out of the guardrail text. +struct JsonStringCollector<'a> { + out: &'a mut String, +} + +struct JsonStringVisitor<'a> { + out: &'a mut String, +} + +struct OtherTopLevelJsonStringsVisitor<'out, 'excluded> { + out: &'out mut String, + excluded: &'excluded [&'excluded str], +} + +impl<'de, 'a, 'b> serde::de::DeserializeSeed<'de> for &'a mut JsonStringCollector<'b> { + type Value = (); + + fn deserialize(self, deserializer: D) -> Result + where + D: serde::Deserializer<'de>, + { + deserializer.deserialize_any(JsonStringVisitor { out: self.out }) + } +} + +impl<'de> serde::de::Visitor<'de> for JsonStringVisitor<'_> { + type Value = (); + + fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str("a JSON value") + } + + fn visit_bool(self, _: bool) -> Result + where + E: serde::de::Error, + { + Ok(()) + } + + fn visit_i64(self, _: i64) -> Result + where + E: serde::de::Error, + { + Ok(()) + } + + fn visit_u64(self, _: u64) -> Result + where + E: serde::de::Error, + { + Ok(()) + } + + fn visit_f64(self, _: f64) -> Result + where + E: serde::de::Error, + { + Ok(()) + } + + fn visit_str(self, text: &str) -> Result + where + E: serde::de::Error, + { + append_scan_text(self.out, text); + Ok(()) + } + + fn visit_none(self) -> Result + where + E: serde::de::Error, + { + Ok(()) + } + + fn visit_unit(self) -> Result + where + E: serde::de::Error, + { + Ok(()) + } + + fn visit_some(self, deserializer: D) -> Result + where + D: serde::Deserializer<'de>, + { + deserializer.deserialize_any(self) + } + + fn visit_seq(self, mut sequence: A) -> Result + where + A: serde::de::SeqAccess<'de>, + { + let mut collector = JsonStringCollector { out: self.out }; + while sequence.next_element_seed(&mut collector)?.is_some() {} + Ok(()) + } + + fn visit_map(self, mut map: A) -> Result + where + A: serde::de::MapAccess<'de>, + { + let mut collector = JsonStringCollector { out: self.out }; + while map.next_key::()?.is_some() { + map.next_value_seed(&mut collector)?; + } + Ok(()) + } +} + +impl<'de> serde::de::Visitor<'de> for OtherTopLevelJsonStringsVisitor<'_, '_> { + type Value = (); + + fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str("a JSON object") + } + + fn visit_map(self, mut map: A) -> Result + where + A: serde::de::MapAccess<'de>, + { + let mut collector = JsonStringCollector { out: self.out }; + while let Some(key) = map.next_key::()? { + if self.excluded.iter().any(|excluded| *excluded == key) { + map.next_value::()?; + } else { + map.next_value_seed(&mut collector)?; + } + } + Ok(()) + } +} + +fn decoded_json_string_values(body: &[u8]) -> Option { + let mut out = String::new(); + let mut collector = JsonStringCollector { out: &mut out }; + let mut deserializer = serde_json::Deserializer::from_slice(body); + serde::de::DeserializeSeed::deserialize(&mut collector, &mut deserializer).ok()?; + deserializer.end().ok()?; + (!out.is_empty()).then_some(out) +} + +/// All occurrences of every non-envelope top-level value, preserving +/// duplicate keys in the source document. The caller's known envelope fields +/// stay on the existing typed extraction path. +fn decoded_other_top_level_json_string_values(body: &[u8], excluded: &[&str]) -> Option { + let mut out = String::new(); + let mut deserializer = serde_json::Deserializer::from_slice(body); + serde::de::Deserializer::deserialize_map( + &mut deserializer, + OtherTopLevelJsonStringsVisitor { + out: &mut out, + excluded, + }, + ) + .ok()?; + deserializer.end().ok()?; + (!out.is_empty()).then_some(out) +} + /// 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. +/// Extraction is best-effort: a shape that yields no typed content falls back +/// to all decoded JSON strings, then to the raw lossy-UTF-8 body when parsing +/// is impossible, 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(), + let (mut extracted, envelope_keys): (String, &[&str]) = match protocol { + PassthroughProtocol::Raw => { + return decoded_json_string_values(body).unwrap_or_else(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"), + 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"), + &["system", "messages"], + ), // 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 @@ -1380,16 +1560,19 @@ fn request_guardrail_text(protocol: PassthroughProtocol, body: &[u8]) -> String // 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::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(), + }, + &["input"], + ), PassthroughProtocol::OpenaiCompletions => { let prompt = v.get("prompt").map(|p| match p { serde_json::Value::Array(items) => items @@ -1407,23 +1590,28 @@ fn request_guardrail_text(protocol: PassthroughProtocol, body: &[u8]) -> String } out.push_str(s); } - out + (out, &["prompt", "suffix"]) } }; if extracted.is_empty() { - raw() - } else { - extracted + return decoded_json_string_values(body).unwrap_or_else(raw); + } + if let Some(other) = decoded_other_top_level_json_string_values(body, envelope_keys) { + append_scan_text(&mut extracted, &other); } + extracted } -/// The response text a guardrail scans / the capture records, per the -/// route's protocol hint. Best-effort like the request side. +/// The response text a guardrail scans, 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(); }; + if matches!(protocol, PassthroughProtocol::Raw) { + return decoded_json_string_values(body).unwrap_or_else(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 @@ -1440,10 +1628,7 @@ fn response_guardrail_text(protocol: PassthroughProtocol, body: &[u8]) -> String return text; } } - let choices = match protocol { - PassthroughProtocol::Raw => return raw(), - _ => v.get("choices").and_then(|c| c.as_array()), - }; + let choices = v.get("choices").and_then(|c| c.as_array()); let Some(choices) = choices else { return raw() }; let texts: Vec = choices .iter() @@ -1466,6 +1651,16 @@ fn response_guardrail_text(protocol: PassthroughProtocol, body: &[u8]) -> String } } +/// Captures preserve raw passthrough responses for the existing telemetry +/// contract; guardrail scanning may use a decoded representation instead. +fn response_capture_text(protocol: PassthroughProtocol, body: &[u8]) -> String { + if matches!(protocol, PassthroughProtocol::Raw) { + String::from_utf8_lossy(body).into_owned() + } else { + response_guardrail_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. @@ -2103,7 +2298,8 @@ fn frame_parts( } parts = match protocol { PassthroughProtocol::Raw => crate::held_content::Parts { - scan: payload.to_string(), + scan: decoded_json_string_values(payload.as_bytes()) + .unwrap_or_else(|| payload.to_string()), reasoning: 0, }, // The chat envelope also carries Anthropic Messages streams; the @@ -2124,6 +2320,20 @@ fn frame_parts( (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()) + .filter(|payload| !payload.is_empty() && *payload != "[DONE]") + .map(str::to_string) + .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 @@ -2265,7 +2475,11 @@ fn stream_response( merge_usage(&mut telemetry.usage, u); } if capture_cap.is_some() { - push_capped(&mut telemetry.response_text, &delta, capture_cap); + push_capped( + &mut telemetry.response_text, + &frame_capture_text(protocol, &frame, &delta), + capture_cap, + ); } let frame = Bytes::from(frame); match &policy { @@ -2367,7 +2581,11 @@ fn stream_response( merge_usage(&mut telemetry.usage, u); } if capture_cap.is_some() { - push_capped(&mut telemetry.response_text, &delta, capture_cap); + push_capped( + &mut telemetry.response_text, + &frame_capture_text(protocol, &rest, &delta), + capture_cap, + ); } scan_buf.push_str(&delta); let rest = Bytes::from(rest); @@ -4038,27 +4256,59 @@ mod tests { #[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" - ); + let scanned = request_guardrail_text(PassthroughProtocol::OpenaiChat, chat); + for text in ["s", "part", "m"] { + assert!(scanned.contains(text), "{text} missing from {scanned:?}"); + } let fim = br#"{"prompt":"def f(","suffix":"return"}"#; assert_eq!( request_guardrail_text(PassthroughProtocol::OpenaiCompletions, fim), "def f(\nreturn" ); - // Shape mismatch degrades to the raw body. + // Shape mismatch degrades to every decoded JSON string value. let not_chat = br#"{"input":"x"}"#; assert_eq!( request_guardrail_text(PassthroughProtocol::OpenaiChat, not_chat), - r#"{"input":"x"}"# + "x" ); - // A detected envelope whose items carry no text ALSO degrades to - // the raw body — detection must never scan less than raw would. + // A detected envelope whose items carry no typed text ALSO degrades + // to every decoded JSON string value — detection must never scan + // less than the forwarded request carries. let empty_chat = br#"{"messages":[{"role":"tool","tool_call_id":"1"}]}"#; + let scanned = request_guardrail_text(PassthroughProtocol::OpenaiChat, empty_chat); + for text in ["tool", "1"] { + assert!(scanned.contains(text), "{text} missing from {scanned:?}"); + } + } + + #[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!( - request_guardrail_text(PassthroughProtocol::OpenaiChat, empty_chat), - String::from_utf8_lossy(empty_chat) + response_capture_text(PassthroughProtocol::Raw, response), + r#"{"state":"\u0042LOCKME","state":"clean"}"# ); } @@ -4686,9 +4936,16 @@ 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"}"# + ); } /// An Anthropic Messages body on the chat envelope is scanned in every 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..a575208a1 100644 --- a/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts @@ -30,6 +30,11 @@ const CALLER = "sk-pt-scan-coverage"; const CALLER_HASH = createHash("sha256").update(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"}`; const CAP = 1_000; const anthropicEvents = (blocks: Array>) => [ @@ -97,6 +102,20 @@ describe("passthrough guardrail scan coverage", () => { upstreams.input = await startOpenAiUpstream({ nonStreamBody: { id: "c", object: "chat.completion", choices: [] }, }); + upstreams["raw-output"] = await startOpenAiUpstream({ + rawBody: ESCAPED_BLOCK_JSON, + rawContentType: "application/json", + }); + upstreams["raw-stream"] = await startOpenAiUpstream({ + rawStreamFrames: [`data: ${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: ${SAFE_ESCAPED_JSON}\n\n`], + }); const pk = await seed.createProviderKey({ display_name: "pt-scan-pk", secret: "sk-mock", @@ -115,14 +134,21 @@ 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 }, + ], }); 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({ @@ -150,6 +176,12 @@ 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 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" }; @@ -236,4 +268,102 @@ 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([ + ["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: 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: 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 raw SSE JSON 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: ${SAFE_ESCAPED_JSON}\n\n`); + expect(upstreams["raw-safe-stream"]!.receivedRequests.length).toBe(before + 1); + }); }); From d0f7024ea4fbbb8eeb61e472b1f302bc88acf92c Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Wed, 30 Sep 2026 12:34:36 +0800 Subject: [PATCH 10/37] fix(passthrough): scan deeply nested JSON values --- crates/aisix-proxy/src/json_splice.rs | 52 ++++++---- crates/aisix-proxy/src/passthrough_route.rs | 88 ++++++++++++----- .../src/cases/passthrough-route-e2e.test.ts | 41 +++----- .../passthrough-scan-coverage-e2e.test.ts | 96 ++++++++++++++++++- 4 files changed, 206 insertions(+), 71 deletions(-) diff --git a/crates/aisix-proxy/src/json_splice.rs b/crates/aisix-proxy/src/json_splice.rs index 2d9464515..2d25ae36f 100644 --- a/crates/aisix-proxy/src/json_splice.rs +++ b/crates/aisix-proxy/src/json_splice.rs @@ -14,11 +14,10 @@ //! 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. It still fails safe: any unexpected byte or overrun +//! returns an error rather than a partially rewritten document. Callers +//! decide the failure policy (the MCP output hook fails closed). use std::ops::Range; @@ -46,11 +45,6 @@ pub struct SpliceError { at: usize, } -/// 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; - /// Rewrite the string values of `input` selected by `should_rewrite`, /// leaving every other byte untouched. /// @@ -108,9 +102,6 @@ pub fn rewrite_string_values( match b { b'{' => { frames.push(Frame::Object); - if frames.len() > MAX_DEPTH { - return Err(err(pos)); - } pos += 1; skip_ws(&mut pos); match input.get(pos) { @@ -135,9 +126,6 @@ pub fn rewrite_string_values( } b'[' => { frames.push(Frame::Array); - if frames.len() > MAX_DEPTH { - return Err(err(pos)); - } pos += 1; skip_ws(&mut pos); if input.get(pos) == Some(&b']') { @@ -242,6 +230,28 @@ 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 keeps working beyond serde_json's default +/// container-recursion limit. 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 { + let mut out = String::new(); + rewrite_string_values( + input, + |_| true, + |value| { + if !out.is_empty() { + out.push('\n'); + } + out.push_str(value); + None + }, + )?; + Ok(out) +} + #[cfg(test)] mod tests { use super::*; @@ -320,6 +330,16 @@ 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 escaped_key_decodes_for_the_predicate() { // `param\u0073` decodes to "params" — the predicate must see the diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index 040395a0f..aa661ba09 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -1496,12 +1496,9 @@ impl<'de> serde::de::Visitor<'de> for OtherTopLevelJsonStringsVisitor<'_, '_> { } fn decoded_json_string_values(body: &[u8]) -> Option { - let mut out = String::new(); - let mut collector = JsonStringCollector { out: &mut out }; - let mut deserializer = serde_json::Deserializer::from_slice(body); - serde::de::DeserializeSeed::deserialize(&mut collector, &mut deserializer).ok()?; - deserializer.end().ok()?; - (!out.is_empty()).then_some(out) + crate::json_splice::collect_string_values(body) + .ok() + .filter(|out| !out.is_empty()) } /// All occurrences of every non-envelope top-level value, preserving @@ -1528,13 +1525,14 @@ fn decoded_other_top_level_json_string_values(body: &[u8], excluded: &[&str]) -> /// is impossible, so detection never loses audit coverage. fn request_guardrail_text(protocol: PassthroughProtocol, body: &[u8]) -> String { let raw = || String::from_utf8_lossy(body).into_owned(); + if matches!(protocol, PassthroughProtocol::Raw) { + return decoded_json_string_values(body).unwrap_or_else(raw); + } let Ok(v) = serde_json::from_slice::(body) else { - return raw(); + return decoded_json_string_values(body).unwrap_or_else(raw); }; let (mut extracted, envelope_keys): (String, &[&str]) = match protocol { - PassthroughProtocol::Raw => { - return decoded_json_string_values(body).unwrap_or_else(raw); - } + PassthroughProtocol::Raw => unreachable!("handled before typed envelope parsing"), // An Anthropic Messages body carries its system prompt top-level. PassthroughProtocol::OpenaiChat => ( v.get("system") @@ -1606,12 +1604,12 @@ fn request_guardrail_text(protocol: PassthroughProtocol, body: &[u8]) -> String /// 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(); - }; if matches!(protocol, PassthroughProtocol::Raw) { return decoded_json_string_values(body).unwrap_or_else(raw); } + let Ok(v) = serde_json::from_slice::(body) else { + return decoded_json_string_values(body).unwrap_or_else(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 @@ -2248,6 +2246,26 @@ fn frame_parts( if payload.is_empty() || payload == "[DONE]" { break 'payload; } + if matches!(protocol, PassthroughProtocol::Raw) { + // Raw payloads have no typed content/usage 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 usage_labelled { + if let Ok(v) = serde_json::from_str::(payload) { + 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. @@ -2258,7 +2276,10 @@ fn frame_parts( // 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); + 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) { @@ -2297,11 +2318,7 @@ fn frame_parts( } } parts = match protocol { - PassthroughProtocol::Raw => crate::held_content::Parts { - scan: decoded_json_string_values(payload.as_bytes()) - .unwrap_or_else(|| payload.to_string()), - reasoning: 0, - }, + 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 => { @@ -2328,9 +2345,8 @@ fn frame_capture_text(protocol: PassthroughProtocol, frame: &[u8], scan: &str) - return scan.to_string(); } crate::redact::frame_payload(frame) - .map(|payload| payload.trim()) - .filter(|payload| !payload.is_empty() && *payload != "[DONE]") - .map(str::to_string) + .map(|payload| payload.trim().to_owned()) + .filter(|payload| !payload.is_empty() && payload != "[DONE]") .unwrap_or_else(|| scan.to_string()) } @@ -3165,6 +3181,15 @@ mod tests { } } + /// 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 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"}}"# @@ -4310,6 +4335,16 @@ mod tests { 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] @@ -4946,6 +4981,15 @@ mod tests { 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 diff --git a/tests/e2e/src/cases/passthrough-route-e2e.test.ts b/tests/e2e/src/cases/passthrough-route-e2e.test.ts index 794e3d270..531bb8c34 100644 --- a/tests/e2e/src/cases/passthrough-route-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-route-e2e.test.ts @@ -3,6 +3,7 @@ import { afterAll, beforeAll, describe, expect, test } from "vitest"; import { harnessRequest } from "../harness/http.js"; import { EtcdClient, + ProxyClient, SeedClient, spawnApp, startOpenAiUpstream, @@ -636,18 +637,9 @@ 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; - } catch { - return false; - } - }); + // Readiness must not exercise the streaming journey this test asserts. + const readiness = new ProxyClient(app.proxyUrl, CALLER_PLAINTEXT); + await waitConfigPropagation(async () => (await readiness.listModels()).status === 200); const res = await call(); expect(res.status).toBe(200); @@ -721,15 +713,8 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat }), }); - await waitConfigPropagation(async () => { - try { - const ready = await fetch(`${app!.proxyUrl}/v1/models`, { headers }); - await ready.text(); - return ready.status === 200; - } catch { - return false; - } - }); + const readiness = new ProxyClient(app.proxyUrl, STREAM_LIMITED_PLAINTEXT); + await waitConfigPropagation(async () => (await readiness.listModels()).status === 200); // Fetch resolves as soon as the upstream headers are relayed. Keep this // body unread while issuing the second request: it is the real caller @@ -751,15 +736,11 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat await first.body!.cancel(); let naturallyEnding: Response | undefined; await waitConfigPropagation(async () => { - try { - const afterCancel = await call(); - const admitted = afterCancel.status === 200; - if (admitted) naturallyEnding = afterCancel; - else await afterCancel.text(); - return admitted; - } catch { - return false; - } + 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(); 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 a575208a1..d51450e25 100644 --- a/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts @@ -10,6 +10,7 @@ import { 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. // @@ -34,7 +35,12 @@ 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"}`; +const SAFE_ESCAPED_JSON = String.raw`{"state":"\u0063lean","state":"safe"}`; +const deepEscapedBlockJSON = (depth: number) => + `${'{"v":'.repeat(depth)}"${String.raw`\u0042LOCKME`}"${'}'.repeat(depth)}`; +// 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); const CAP = 1_000; const anthropicEvents = (blocks: Array>) => [ @@ -87,7 +93,9 @@ const STREAMS: Record = { describe("passthrough guardrail scan coverage", () => { let app: SpawnedApp | undefined; + let seed: SeedClient | undefined; const upstreams: Record = {}; + const otlps: MockOtlp[] = []; let etcdReachable = false; beforeAll(async () => { @@ -95,7 +103,7 @@ 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 }); } @@ -109,6 +117,13 @@ describe("passthrough guardrail scan coverage", () => { upstreams["raw-stream"] = await startOpenAiUpstream({ rawStreamFrames: [`data: ${ESCAPED_BLOCK_JSON}\n\n`, "data: [DONE]\n\n"], }); + upstreams["raw-deep-output"] = await startOpenAiUpstream({ + rawBody: DEEP_ESCAPED_BLOCK_JSON, + 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", @@ -168,6 +183,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) => @@ -187,7 +203,7 @@ describe("passthrough guardrail scan coverage", () => { 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; } @@ -316,6 +332,15 @@ describe("passthrough guardrail scan coverage", () => { 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: safe raw JSON keeps its original bytes upstream", async (ctx) => { if (!ready(ctx)) return; const before = upstreams.input!.receivedRequests.length; @@ -346,6 +371,17 @@ describe("passthrough guardrail scan coverage", () => { 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 SSE JSON escapes are decoded before scanning", async (ctx) => { if (!ready(ctx)) return; const before = upstreams["raw-stream"]!.receivedRequests.length; @@ -366,4 +402,58 @@ describe("passthrough guardrail scan coverage", () => { expect(await res.text()).toBe(`data: ${SAFE_ESCAPED_JSON}\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 proxy = new ProxyClient(app!.proxyUrl, CALLER); + await waitConfigPropagation(async () => (await proxy.listModels()).status === 200); + + const completionFor = async (route: string) => { + const deadline = Date.now() + 10_000; + while (Date.now() < deadline) { + const span = otlp.spans.find( + (candidate) => + candidate.attributes["aisix.passthrough.route_name"] === `pt-scan-${route}` && + candidate.attributes["gen_ai.completion"] === SAFE_ESCAPED_JSON, + ); + if (span) return span; + await new Promise((resolve) => setTimeout(resolve, 50)); + } + throw new Error(`no raw-source OTLP completion for ${route}`); + }; + + const buffered = await callRaw("raw-safe-output", "/v1/any", SAFE_ESCAPED_JSON); + expect(buffered.status).toBe(200); + expect(await buffered.text()).toBe(SAFE_ESCAPED_JSON); + await completionFor("raw-safe-output"); + + const streamed = await callRaw("raw-safe-stream", "/v1/any", SAFE_ESCAPED_JSON); + expect(streamed.status).toBe(200); + expect(await streamed.text()).toBe(`data: ${SAFE_ESCAPED_JSON}\n\n`); + await completionFor("raw-safe-stream"); + expect(otlp.parseFailures).toEqual([]); + }); + + test("output: deep raw SSE JSON escapes are decoded before scanning", 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("content_filter"); + expect(body).not.toContain(ESCAPED_BLOCK); + expect(upstreams["raw-deep-stream"]!.receivedRequests.length).toBe(before + 1); + }); }); From 606094454d56f6002c0fe27bec5a9c7a00d31a2b Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Wed, 30 Sep 2026 12:54:59 +0800 Subject: [PATCH 11/37] fix(passthrough): retain raw stream usage --- crates/aisix-proxy/src/json_splice.rs | 6 +- crates/aisix-proxy/src/passthrough_route.rs | 23 ++++--- .../src/cases/passthrough-route-e2e.test.ts | 68 +++++++++++++------ .../passthrough-scan-coverage-e2e.test.ts | 32 +++++---- 4 files changed, 81 insertions(+), 48 deletions(-) diff --git a/crates/aisix-proxy/src/json_splice.rs b/crates/aisix-proxy/src/json_splice.rs index 2d25ae36f..d72efe175 100644 --- a/crates/aisix-proxy/src/json_splice.rs +++ b/crates/aisix-proxy/src/json_splice.rs @@ -419,7 +419,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('['); @@ -427,6 +427,8 @@ 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()); } } diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index aa661ba09..218b6f0c5 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -1548,7 +1548,7 @@ fn request_guardrail_text(protocol: PassthroughProtocol, body: &[u8]) -> String .filter(|t| !t.is_empty()) .collect::>() .join("\n"), - &["system", "messages"], + &["model", "system", "messages"], ), // Responses API: `input` is either a bare string or an array of // items, read exactly as the typed route reads them @@ -1569,7 +1569,7 @@ fn request_guardrail_text(protocol: PassthroughProtocol, body: &[u8]) -> String .join("\n"), _ => String::new(), }, - &["input"], + &["model", "input"], ), PassthroughProtocol::OpenaiCompletions => { let prompt = v.get("prompt").map(|p| match p { @@ -1588,7 +1588,7 @@ fn request_guardrail_text(protocol: PassthroughProtocol, body: &[u8]) -> String } out.push_str(s); } - (out, &["prompt", "suffix"]) + (out, &["model", "prompt", "suffix"]) } }; if extracted.is_empty() { @@ -2247,13 +2247,19 @@ fn frame_parts( break 'payload; } if matches!(protocol, PassthroughProtocol::Raw) { - // Raw payloads have no typed content/usage envelope to extract. + // 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 usage_labelled { - if let Ok(v) = serde_json::from_str::(payload) { + 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); } @@ -3988,10 +3994,7 @@ mod tests { // 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}}", - ); + assert_eq!(text, "message_delta"); } /// Framing varies per ENDPOINT, not per vendor: on one host diff --git a/tests/e2e/src/cases/passthrough-route-e2e.test.ts b/tests/e2e/src/cases/passthrough-route-e2e.test.ts index 531bb8c34..3fdb5abe3 100644 --- a/tests/e2e/src/cases/passthrough-route-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-route-e2e.test.ts @@ -3,7 +3,6 @@ import { afterAll, beforeAll, describe, expect, test } from "vitest"; import { harnessRequest } from "../harness/http.js"; import { EtcdClient, - ProxyClient, SeedClient, spawnApp, startOpenAiUpstream, @@ -597,15 +596,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); @@ -637,9 +643,18 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat }), }); - // Readiness must not exercise the streaming journey this test asserts. - const readiness = new ProxyClient(app.proxyUrl, CALLER_PLAINTEXT); - await waitConfigPropagation(async () => (await readiness.listModels()).status === 200); + 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; + } + }); const res = await call(); expect(res.status).toBe(200); @@ -666,10 +681,11 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat "[DONE]", ]; const upstream = await startOpenAiUpstream({ - // The propagation probe does not reach this upstream. The first real - // stream stalls until the client cancels it. The second ends naturally, - // so this test covers both lifetime boundaries of the same reservation. 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 }, @@ -688,9 +704,9 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat target_url: upstream.baseUrl, provider_key_id: pk.id, }); - // Write the constrained principal last. A successful authenticated - // readiness probe proves the preceding route has reached the same - // snapshot without spending its concurrency slot. + // 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: ["*"], @@ -713,8 +729,18 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat }), }); - const readiness = new ProxyClient(app.proxyUrl, STREAM_LIMITED_PLAINTEXT); - await waitConfigPropagation(async () => (await readiness.listModels()).status === 200); + 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 @@ -722,7 +748,7 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat const first = await call(); expect(first.status).toBe(200); const upstreamCallsWhileStreaming = upstream.receivedRequests.length; - expect(upstreamCallsWhileStreaming).toBe(1); + expect(upstreamCallsWhileStreaming).toBe(2); const second = await call(); expect(second.status).toBe(429); 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 d51450e25..97a3749d0 100644 --- a/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts @@ -416,12 +416,21 @@ describe("passthrough guardrail scan coverage", () => { content_mode: "full", content_max_bytes: 4_096, }); - const proxy = new ProxyClient(app!.proxyUrl, CALLER); - await waitConfigPropagation(async () => (await proxy.listModels()).status === 200); - - const completionFor = async (route: string) => { - const deadline = Date.now() + 10_000; + const requestUntilCaptured = async (route: string, expectedBody: 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}` && @@ -430,18 +439,11 @@ describe("passthrough guardrail scan coverage", () => { if (span) return span; await new Promise((resolve) => setTimeout(resolve, 50)); } - throw new Error(`no raw-source OTLP completion for ${route}`); + throw new Error(`no raw-source OTLP completion for ${route}; last response ${last}`); }; - const buffered = await callRaw("raw-safe-output", "/v1/any", SAFE_ESCAPED_JSON); - expect(buffered.status).toBe(200); - expect(await buffered.text()).toBe(SAFE_ESCAPED_JSON); - await completionFor("raw-safe-output"); - - const streamed = await callRaw("raw-safe-stream", "/v1/any", SAFE_ESCAPED_JSON); - expect(streamed.status).toBe(200); - expect(await streamed.text()).toBe(`data: ${SAFE_ESCAPED_JSON}\n\n`); - await completionFor("raw-safe-stream"); + await requestUntilCaptured("raw-safe-output", SAFE_ESCAPED_JSON); + await requestUntilCaptured("raw-safe-stream", `data: ${SAFE_ESCAPED_JSON}\n\n`); expect(otlp.parseFailures).toEqual([]); }); From 19607715a54ab595f55fb5deb6ceb43c941ad3ca Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Wed, 30 Sep 2026 13:02:40 +0800 Subject: [PATCH 12/37] test(proxy): cover supplemental passthrough fields --- crates/aisix-proxy/src/passthrough_route.rs | 10 ++++++++-- 1 file changed, 8 insertions(+), 2 deletions(-) diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index 218b6f0c5..359828a74 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -4283,11 +4283,17 @@ mod tests { #[test] fn request_text_extraction_per_protocol() { - let chat = br#"{"model":"m","messages":[{"role":"system","content":"s"},{"role":"user","content":[{"type":"text","text":"part"}]}]}"#; + 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", "m"] { + 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!( request_guardrail_text(PassthroughProtocol::OpenaiCompletions, fim), From 190a42725cb2fa05c1654edfc3f91c8fb7409f6e Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Wed, 30 Sep 2026 14:29:32 +0800 Subject: [PATCH 13/37] fix(guardrail): preserve source-safe passthrough scans --- crates/aisix-proxy/Cargo.toml | 2 +- crates/aisix-proxy/src/json_splice.rs | 59 +- crates/aisix-proxy/src/passthrough_route.rs | 1222 +++++++++++++---- .../passthrough-scan-coverage-e2e.test.ts | 140 ++ 4 files changed, 1184 insertions(+), 239 deletions(-) diff --git a/crates/aisix-proxy/Cargo.toml b/crates/aisix-proxy/Cargo.toml index a03c50c47..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 diff --git a/crates/aisix-proxy/src/json_splice.rs b/crates/aisix-proxy/src/json_splice.rs index d72efe175..467cca7b8 100644 --- a/crates/aisix-proxy/src/json_splice.rs +++ b/crates/aisix-proxy/src/json_splice.rs @@ -237,10 +237,22 @@ pub fn rewrite_string_values( /// container-recursion limit. 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) +} + +/// Decode and collect selected JSON string **values** in source order. +/// +/// Like [`collect_string_values`], this preserves duplicate keys and stays +/// stack-safe for deeply nested documents. 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(); rewrite_string_values( input, - |_| true, + |path| include(path), |value| { if !out.is_empty() { out.push('\n'); @@ -252,6 +264,27 @@ pub fn collect_string_values(input: &[u8]) -> Result { Ok(out) } +/// 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(); + rewrite_string_values( + input, + |path| include(path), + |value| { + out.push(value.to_string()); + None + }, + )?; + Ok(out) +} + #[cfg(test)] mod tests { use super::*; @@ -340,6 +373,30 @@ mod tests { 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 escaped_key_decodes_for_the_predicate() { // `param\u0073` decodes to "params" — the predicate must see the diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index 359828a74..8b18c751d 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -1358,268 +1358,450 @@ fn append_scan_text(out: &mut String, text: &str) { out.push_str(text); } -/// A duplicate-preserving JSON value walker. `serde_json::Value` is right for -/// typed envelope extraction but keeps only the final value for a repeated -/// object key; passthrough forwards the original bytes, so guardrail scanning -/// must see every decoded string value the upstream can see. Object keys are -/// structural metadata and deliberately stay out of the guardrail text. -struct JsonStringCollector<'a> { - out: &'a mut String, +fn decoded_json_string_values(body: &[u8]) -> Option { + crate::json_splice::collect_string_values(body) + .ok() + .filter(|out| !out.is_empty()) } -struct JsonStringVisitor<'a> { - out: &'a mut String, +fn decoded_json_string_values_where( + body: &[u8], + include: impl FnMut(&[crate::json_splice::PathSeg]) -> bool, +) -> Option { + crate::json_splice::collect_string_values_where(body, include) + .ok() + .filter(|out| !out.is_empty()) } -struct OtherTopLevelJsonStringsVisitor<'out, 'excluded> { - out: &'out mut String, - excluded: &'excluded [&'excluded str], +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()) } -impl<'de, 'a, 'b> serde::de::DeserializeSeed<'de> for &'a mut JsonStringCollector<'b> { - type Value = (); +fn is_root_key(path: &[crate::json_splice::PathSeg], key: &str) -> bool { + path.first().is_some_and(|segment| segment.is_key(key)) +} - fn deserialize(self, deserializer: D) -> Result - where - D: serde::Deserializer<'de>, - { - deserializer.deserialize_any(JsonStringVisitor { out: self.out }) - } +/// A detected envelope still forwards raw bytes, including duplicate keys and +/// arbitrary nested fields. Scan every decoded string the upstream can read; +/// the root `model` alone is routing metadata rather than caller content. +fn decoded_non_model_json_string_values(body: &[u8]) -> Option { + decoded_json_string_values_where(body, |path| !is_root_key(path, "model")) } -impl<'de> serde::de::Visitor<'de> for JsonStringVisitor<'_> { - type Value = (); +fn decoded_json_string_values_including_empty(body: &[u8]) -> Option { + crate::json_splice::collect_string_values(body).ok() +} - fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - formatter.write_str("a JSON value") - } +fn decoded_json_string_values_except_root_keys(body: &[u8], excluded: &[&str]) -> Option { + crate::json_splice::collect_string_values_where(body, |path| { + !excluded.iter().any(|key| is_root_key(path, key)) + }) + .ok() +} - fn visit_bool(self, _: bool) -> Result - where - E: serde::de::Error, - { - Ok(()) +/// Source values of all occurrences of one top-level key. `RawValue` keeps +/// repeated keys separate, unlike `serde_json::Value`. +fn raw_top_level_values( + body: &[u8], + wanted_key: &str, +) -> Option>> { + struct Values<'a> { + wanted_key: &'a str, } - fn visit_i64(self, _: i64) -> Result - where - E: serde::de::Error, - { - Ok(()) - } + impl<'de> serde::de::Visitor<'de> for Values<'_> { + type Value = Vec>; - fn visit_u64(self, _: u64) -> Result - where - E: serde::de::Error, - { - Ok(()) - } + fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str("a JSON object") + } - fn visit_f64(self, _: f64) -> Result - where - E: serde::de::Error, - { - Ok(()) + fn visit_map(self, mut map: A) -> Result + where + A: serde::de::MapAccess<'de>, + { + let mut values = Vec::new(); + while let Some(key) = map.next_key::()? { + if key == self.wanted_key { + values.push(map.next_value::>()?); + } else { + map.next_value::()?; + } + } + Ok(values) + } } - fn visit_str(self, text: &str) -> Result - where - E: serde::de::Error, - { - append_scan_text(self.out, text); - Ok(()) - } + let mut deserializer = serde_json::Deserializer::from_slice(body); + let values = + serde::de::Deserializer::deserialize_map(&mut deserializer, Values { wanted_key }).ok()?; + deserializer.end().ok()?; + Some(values) +} - fn visit_none(self) -> Result - where - E: serde::de::Error, - { - Ok(()) - } +fn raw_array_items( + raw: &serde_json::value::RawValue, +) -> Option>> { + serde_json::from_str(raw.get()).ok() +} - fn visit_unit(self) -> Result - where - E: serde::de::Error, - { - Ok(()) +/// `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: &serde_json::value::RawValue, allowed: &[&str]) -> bool { + let Some(values) = raw_top_level_values(raw.get().as_bytes(), "type") else { + return false; + }; + 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())) +} + +fn raw_top_level_unique_type(body: &[u8]) -> Option { + let mut types = raw_top_level_values(body, "type") + .into_iter() + .flatten() + .map(|value| serde_json::from_str::(value.get()).ok()); + let first = types.next()??; + types + .all(|kind| kind.as_deref() == Some(first.as_str())) + .then_some(first) +} + +fn raw_top_level_has_any_type(body: &[u8], wanted: &[&str]) -> bool { + raw_top_level_values(body, "type") + .into_iter() + .flatten() + .filter_map(|value| serde_json::from_str::(value.get()).ok()) + .any(|kind| wanted.iter().any(|wanted| kind == *wanted)) +} + +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_array_item_strings( + out: &mut String, + body: &[u8], + key: &str, + mut skip: impl FnMut(&serde_json::value::RawValue) -> bool, +) -> Option<()> { + for array in raw_top_level_values(body, key)? { + for item in raw_array_items(&array)? { + if !skip(&item) { + append_scan_text( + out, + &decoded_json_string_values_including_empty(item.get().as_bytes())?, + ); + } + } } + Some(()) +} - fn visit_some(self, deserializer: D) -> Result - where - D: serde::Deserializer<'de>, - { - deserializer.deserialize_any(self) +fn append_raw_string_value(out: &mut String, raw: &serde_json::value::RawValue) -> 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)? { + append_raw_string_value(out, &value)?; } + Some(()) +} - fn visit_seq(self, mut sequence: A) -> Result - where - A: serde::de::SeqAccess<'de>, - { - let mut collector = JsonStringCollector { out: self.out }; - while sequence.next_element_seed(&mut collector)?.is_some() {} - Ok(()) +fn raw_top_level_string_values(body: &[u8], key: &str) -> Option> { + raw_top_level_values(body, key)? + .into_iter() + .map(|value| serde_json::from_str::(value.get()).ok()) + .collect() +} + +/// 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: &serde_json::value::RawValue) -> 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)? { + append_raw_top_level_strings(out, part.get().as_bytes(), "text")?; + } + Some(()) +} - fn visit_map(self, mut map: A) -> Result - where - A: serde::de::MapAccess<'de>, - { - let mut collector = JsonStringCollector { out: self.out }; - while map.next_key::()?.is_some() { - map.next_value_seed(&mut collector)?; +/// 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: &serde_json::value::RawValue, +) -> 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 block in raw_array_items(content)? { + let block_body = block.get().as_bytes(); + let types = raw_top_level_values(block_body, "type")?; + let kind = raw_top_level_unique_type(block_body); + if !types.is_empty() && kind.is_none() { + append_scan_text( + out, + &decoded_json_string_values_including_empty(block_body)?, + ); + continue; + } + match kind.as_deref() { + Some("redacted_thinking") => {} + Some("tool_result") => { + for nested in raw_top_level_values(block_body, "content")? { + append_chat_request_content_strings(out, &nested)?; + } + } + Some("tool_use") => { + for input in raw_top_level_values(block_body, "input")? { + append_scan_text( + out, + &decoded_json_string_values_including_empty(input.get().as_bytes())?, + ); + } + } + Some("thinking") => append_raw_top_level_strings(out, block_body, "thinking")?, + _ => append_raw_top_level_strings(out, block_body, "text")?, } - Ok(()) } + Some(()) } -impl<'de> serde::de::Visitor<'de> for OtherTopLevelJsonStringsVisitor<'_, '_> { - type Value = (); +fn append_chat_request_message_strings( + out: &mut String, + message: &serde_json::value::RawValue, +) -> Option<()> { + 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"], + )?, + ); + for content in raw_top_level_values(message_body, "content")? { + append_chat_request_content_strings(out, &content)?; + } + for tool_calls in raw_top_level_values(message_body, "tool_calls")? { + append_scan_text( + out, + &decoded_json_string_values_including_empty(tool_calls.get().as_bytes())?, + ); + } + append_raw_top_level_strings(out, message_body, "reasoning_content")?; + Some(()) +} - fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - formatter.write_str("a JSON object") +fn decoded_chat_request_string_values(body: &[u8]) -> Option { + let mut out = + decoded_json_string_values_except_root_keys(body, &["model", "system", "messages"])?; + for system in raw_top_level_values(body, "system")? { + append_chat_request_content_strings(&mut out, &system)?; } + for array in raw_top_level_values(body, "messages")? { + for message in raw_array_items(&array)? { + append_chat_request_message_strings(&mut out, &message)?; + } + } + Some(out) +} - fn visit_map(self, mut map: A) -> Result - where - A: serde::de::MapAccess<'de>, - { - let mut collector = JsonStringCollector { out: self.out }; - while let Some(key) = map.next_key::()? { - if self.excluded.iter().any(|excluded| *excluded == key) { - map.next_value::()?; - } else { - map.next_value_seed(&mut collector)?; - } +fn append_responses_item_strings( + out: &mut String, + item: &serde_json::value::RawValue, +) -> Option<()> { + 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)?, + ); + for key in text_keys { + for value in raw_top_level_values(item_body, key)? { + append_raw_text_value(out, &value)?; } - Ok(()) } + Some(()) } -fn decoded_json_string_values(body: &[u8]) -> Option { - crate::json_splice::collect_string_values(body) - .ok() - .filter(|out| !out.is_empty()) +fn decoded_responses_request_string_values(body: &[u8]) -> Option { + let mut out = decoded_json_string_values_except_root_keys(body, &["model", "input"])?; + 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)?; + } + } + } + Some(out) } -/// All occurrences of every non-envelope top-level value, preserving -/// duplicate keys in the source document. The caller's known envelope fields -/// stay on the existing typed extraction path. -fn decoded_other_top_level_json_string_values(body: &[u8], excluded: &[&str]) -> Option { - let mut out = String::new(); - let mut deserializer = serde_json::Deserializer::from_slice(body); - serde::de::Deserializer::deserialize_map( - &mut deserializer, - OtherTopLevelJsonStringsVisitor { - out: &mut out, - excluded, - }, - ) - .ok()?; - deserializer.end().ok()?; - (!out.is_empty()).then_some(out) +fn decoded_completions_request_string_values(body: &[u8]) -> Option { + let mut out = + decoded_json_string_values_except_root_keys(body, &["model", "prompt", "suffix"])?; + 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. -/// Extraction is best-effort: a shape that yields no typed content falls back -/// to all decoded JSON strings, then to the raw lossy-UTF-8 body when parsing -/// is impossible, so detection never loses audit coverage. +/// +/// 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(protocol: PassthroughProtocol, body: &[u8]) -> String { let raw = || String::from_utf8_lossy(body).into_owned(); - if matches!(protocol, PassthroughProtocol::Raw) { - return decoded_json_string_values(body).unwrap_or_else(raw); - } - let Ok(v) = serde_json::from_slice::(body) else { - return decoded_json_string_values(body).unwrap_or_else(raw); - }; - let (mut extracted, envelope_keys): (String, &[&str]) = match protocol { - PassthroughProtocol::Raw => unreachable!("handled before typed envelope parsing"), - // 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"), - &["model", "system", "messages"], - ), - // 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(), - }, - &["model", "input"], - ), + match protocol { + PassthroughProtocol::Raw => decoded_json_string_values(body).unwrap_or_else(raw), + PassthroughProtocol::OpenaiChat => { + decoded_chat_request_string_values(body).unwrap_or_else(raw) + } 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); - } - (out, &["model", "prompt", "suffix"]) + decoded_completions_request_string_values(body).unwrap_or_else(raw) + } + PassthroughProtocol::OpenaiResponses => { + decoded_responses_request_string_values(body).unwrap_or_else(raw) } - }; - if extracted.is_empty() { - return decoded_json_string_values(body).unwrap_or_else(raw); } - if let Some(other) = decoded_other_top_level_json_string_values(body, envelope_keys) { - append_scan_text(&mut extracted, &other); +} + +fn is_hidden_chat_reasoning_path(path: &[crate::json_splice::PathSeg]) -> bool { + use crate::json_splice::PathSeg; + + let path = match path { + [PathSeg::Key(choices), PathSeg::Index(_), rest @ ..] if choices == "choices" => rest, + _ => path, + }; + matches!( + path, + [ + PathSeg::Key(message_or_delta), + PathSeg::Key(reasoning), + .. + ] if matches!(message_or_delta.as_str(), "message" | "delta") + && matches!(reasoning.as_str(), "reasoning_content" | "reasoning") + ) +} + +fn decoded_chat_response_string_values(body: &[u8]) -> Option { + let mut out = + decoded_json_string_values_except_root_keys(body, &["model", "content", "choices"])?; + append_raw_array_item_strings(&mut out, body, "content", |item| { + raw_object_has_only_types(item, &["thinking", "redacted_thinking"]) + })?; + for array in raw_top_level_values(body, "choices")? { + for item in raw_array_items(&array)? { + let text = + crate::json_splice::collect_string_values_where(item.get().as_bytes(), |path| { + !is_hidden_chat_reasoning_path(path) + }) + .ok()?; + append_scan_text(&mut out, &text); + } } - extracted + Some(out) +} + +fn decoded_responses_response_string_values(body: &[u8]) -> Option { + let mut out = decoded_json_string_values_except_root_keys(body, &["model", "output"])?; + append_raw_array_item_strings(&mut out, body, "output", |item| { + raw_object_has_only_types(item, &["reasoning"]) + })?; + Some(out) } /// The response text a guardrail scans, per the route's protocol hint. -/// Best-effort like the request side. +/// +/// This deliberately reads raw source values rather than `Value`: a +/// passthrough upstream may send duplicate or deeply nested fields which the +/// client receives verbatim. Generated reasoning remains out of scope only +/// for an unambiguous standard reasoning item. fn response_guardrail_text(protocol: PassthroughProtocol, body: &[u8]) -> 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 => { + decoded_chat_response_string_values(body).unwrap_or_else(raw) + } + PassthroughProtocol::OpenaiCompletions => { + decoded_non_model_json_string_values(body).unwrap_or_else(raw) + } + PassthroughProtocol::OpenaiResponses => { + decoded_responses_response_string_values(body).unwrap_or_else(raw) + } + } +} + +/// 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 decoded_json_string_values(body).unwrap_or_else(raw); + return raw(); } let Ok(v) = serde_json::from_slice::(body) else { return decoded_json_string_values(body).unwrap_or_else(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 }; } - // 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() { @@ -1637,7 +1819,6 @@ fn response_guardrail_text(protocol: PassthroughProtocol, body: &[u8]) -> String 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()) @@ -1649,14 +1830,10 @@ fn response_guardrail_text(protocol: PassthroughProtocol, body: &[u8]) -> String } } -/// Captures preserve raw passthrough responses for the existing telemetry -/// contract; guardrail scanning may use a decoded representation instead. +/// 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 { - if matches!(protocol, PassthroughProtocol::Raw) { - String::from_utf8_lossy(body).into_owned() - } else { - response_guardrail_text(protocol, body) - } + response_visible_text(protocol, body) } /// Every token dimension a passthrough exchange can report, mirroring the @@ -2216,6 +2393,301 @@ fn frame_delta(protocol: PassthroughProtocol, frame: &[u8]) -> (String, Option

bool { + is_hidden_chat_reasoning_path(path) +} + +/// 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).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), + } +} + +fn decoded_chat_frame_string_values(body: &[u8]) -> Option { + match raw_top_level_unique_type(body).as_deref() { + Some("content_block_delta") + if raw_top_level_items_have_only_types( + body, + "delta", + &["thinking_delta", "signature_delta"], + )? => + { + return decoded_json_string_values_except_root_keys(body, &["model", "delta"]); + } + Some("content_block_start") + if raw_top_level_items_have_only_types( + body, + "content_block", + &["thinking", "redacted_thinking"], + )? => + { + return decoded_json_string_values_except_root_keys(body, &["model", "content_block"]); + } + _ => {} + } + crate::json_splice::collect_string_values_where(body, |path| { + !is_root_key(path, "model") && !is_hidden_chat_stream_reasoning_path(path) + }) + .ok() +} + +fn decoded_responses_frame_string_values(body: &[u8]) -> Option { + match raw_top_level_unique_type(body).as_deref() { + Some("response.reasoning_text.delta" | "response.reasoning_summary_text.delta") => { + return decoded_json_string_values_except_root_keys(body, &["model", "delta"]); + } + Some("response.reasoning_text.done" | "response.reasoning_summary_text.done") => { + return decoded_json_string_values_except_root_keys(body, &["model", "text"]); + } + Some("response.reasoning_summary_part.added" | "response.reasoning_summary_part.done") => { + return decoded_json_string_values_except_root_keys(body, &["model", "part"]); + } + Some("response.output_item.added" | "response.output_item.done") + if raw_top_level_items_have_only_types(body, "item", &["reasoning"])? => + { + return decoded_json_string_values_except_root_keys(body, &["model", "item"]); + } + Some("response.content_part.added" | "response.content_part.done") + if raw_top_level_items_have_only_types( + body, + "part", + &["reasoning_text", "reasoning_summary"], + )? => + { + return decoded_json_string_values_except_root_keys(body, &["model", "part"]); + } + _ => {} + } + + let responses = raw_top_level_values(body, "response")?; + if responses.is_empty() { + return decoded_non_model_json_string_values(body); + } + let mut out = decoded_json_string_values_except_root_keys(body, &["model", "response"])?; + for response in responses { + append_scan_text( + &mut out, + &response_guardrail_text( + PassthroughProtocol::OpenaiResponses, + response.get().as_bytes(), + ), + ); + } + Some(out) +} + +fn decoded_chat_frame_continuations(body: &[u8]) -> Option> { + use crate::json_splice::PathSeg; + + if hidden_chat_stream_reasoning_frame(body) == Some(true) { + return None; + } + if raw_top_level_has_any_type(body, &["content_block_delta"]) { + return decoded_json_string_values_vec_where(body, |path| { + matches!( + path, + [PathSeg::Key(delta), PathSeg::Key(field)] + if delta == "delta" && matches!(field.as_str(), "text" | "partial_json") + ) + }); + } + if raw_top_level_has_any_type(body, &["content_block_start"]) { + let mut out = Vec::new(); + for block in raw_top_level_values(body, "content_block")? { + out.extend(raw_top_level_string_values(block.get().as_bytes(), "text")?); + for input in raw_top_level_values(block.get().as_bytes(), "input")? { + out.extend( + crate::json_splice::collect_string_values_where_vec( + input.get().as_bytes(), + |_| true, + ) + .ok()?, + ); + } + } + return (!out.is_empty()).then_some(out); + } + decoded_json_string_values_vec_where(body, |path| { + matches!( + path, + [PathSeg::Key(choices), PathSeg::Index(_), PathSeg::Key(delta), PathSeg::Key(field)] + if choices == "choices" + && delta == "delta" + && field == "content" + ) || matches!( + path, + [PathSeg::Key(choices), PathSeg::Index(_), PathSeg::Key(delta), PathSeg::Key(content), PathSeg::Index(_), PathSeg::Key(text)] + if choices == "choices" + && delta == "delta" + && content == "content" + && text == "text" + ) || matches!( + path, + [PathSeg::Key(choices), PathSeg::Index(_), PathSeg::Key(delta), PathSeg::Key(tool_calls), PathSeg::Index(_), PathSeg::Key(kind), PathSeg::Key(value)] + if choices == "choices" + && delta == "delta" + && tool_calls == "tool_calls" + && ((kind == "function" && value == "arguments") + || (kind == "custom" && value == "input")) + ) + }) +} + +fn decoded_completions_frame_continuations(body: &[u8]) -> Option> { + use crate::json_splice::PathSeg; + + decoded_json_string_values_vec_where(body, |path| { + matches!( + path, + [PathSeg::Key(choices), PathSeg::Index(_), PathSeg::Key(text)] + if choices == "choices" && text == "text" + ) + }) +} + +fn decoded_responses_frame_continuations(body: &[u8]) -> Option> { + const VISIBLE_DELTA_EVENTS: &[&str] = &[ + "response.output_text.delta", + "response.function_call_arguments.delta", + "response.mcp_call_arguments.delta", + "response.custom_tool_call_input.delta", + ]; + if !raw_top_level_has_any_type(body, VISIBLE_DELTA_EVENTS) { + return None; + } + raw_top_level_string_values(body, "delta").filter(|values| !values.is_empty()) +} + +fn frame_source_continuations( + protocol: PassthroughProtocol, + payload: &[u8], +) -> Option> { + match protocol { + PassthroughProtocol::Raw => decoded_json_string_values(payload).map(|text| vec![text]), + PassthroughProtocol::OpenaiChat => decoded_chat_frame_continuations(payload), + PassthroughProtocol::OpenaiCompletions => decoded_completions_frame_continuations(payload), + PassthroughProtocol::OpenaiResponses => decoded_responses_frame_continuations(payload), + } +} + +/// 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 also sees duplicate and arbitrary forwarded fields. +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(); + } + let raw = || payload.to_string(); + match protocol { + PassthroughProtocol::Raw => { + decoded_json_string_values(payload.as_bytes()).unwrap_or_else(raw) + } + PassthroughProtocol::OpenaiChat => { + decoded_chat_frame_string_values(payload.as_bytes()).unwrap_or_else(raw) + } + PassthroughProtocol::OpenaiCompletions => { + decoded_non_model_json_string_values(payload.as_bytes()).unwrap_or_else(raw) + } + PassthroughProtocol::OpenaiResponses => { + decoded_responses_frame_string_values(payload.as_bytes()).unwrap_or_else(raw) + } + } +} + +/// The independent channels scanned for one stream frame. The first +/// continuation is the typed visible-output sequence; later channels retain +/// source output-carrier occurrences across frames. `supplemental` preserves +/// every additional decoded source value (including duplicate or nested +/// fields). Keeping them separate prevents frame metadata from interrupting a +/// sensitive literal split across output deltas. +struct StreamGuardrailText { + continuations: Vec, + supplemental: String, +} + +fn stream_guardrail_text( + protocol: PassthroughProtocol, + frame: &[u8], + continuation: String, +) -> StreamGuardrailText { + // Keep the typed continuation in a stable first channel even when source + // extraction finds raw carriers. A provider can add or remove duplicate + // carrier fields between frames; only the typed last-wins sequence then + // follows what a normal JSON client sees across that boundary. + 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 mut continuations = vec![(!hidden_reasoning) + .then_some(continuation) + .unwrap_or_default()]; + if let Some(source) = + payload.and_then(|payload| frame_source_continuations(protocol, payload.trim().as_bytes())) + { + continuations.extend(source); + } + StreamGuardrailText { + continuations, + supplemental: frame_guardrail_text(protocol, frame), + } +} + +fn append_stream_guardrail_text( + continuations: &mut Vec, + supplemental: &mut String, + text: &StreamGuardrailText, +) { + if continuations.len() < text.continuations.len() { + continuations.resize(text.continuations.len(), String::new()); + } + for (index, continuation) in text.continuations.iter().enumerate() { + continuations[index].push_str(continuation); + } + append_scan_text(supplemental, &text.supplemental); +} + +fn stream_guardrail_scan_text( + continuation_tails: &[String], + continuations: &[String], + supplemental_tail: &str, + supplemental: &str, +) -> String { + let mut text = String::new(); + for (index, continuation) in continuations.iter().enumerate() { + let tail = continuation_tails + .get(index) + .map(String::as_str) + .unwrap_or_default(); + append_scan_text(&mut text, &format!("{tail}{continuation}")); + } + let supplemental = match (supplemental_tail.is_empty(), supplemental.is_empty()) { + (true, true) => String::new(), + (false, true) => supplemental_tail.to_string(), + (true, false) => supplemental.to_string(), + (false, false) => format!("{supplemental_tail}\n{supplemental}"), + }; + append_scan_text(&mut text, &supplemental); + 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 @@ -2231,7 +2703,7 @@ fn frame_parts( let mut merge = |found: PassthroughUsage| { merge_usage(&mut usage, found); }; - // ONE read and ONE parse per frame: a payload spread over several + // 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 @@ -2445,10 +2917,14 @@ fn stream_response( // 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(); + // The semantic delta channel stays contiguous across frames. Raw + // supplementary values are scanned separately so metadata cannot + // break a literal split over two output deltas. + let mut continuation_bufs: Vec = Vec::new(); + let mut supplemental_buf = String::new(); + // Overlap carried between Window scans, one per channel. + let mut continuation_tails: Vec = Vec::new(); + let mut supplemental_tail = String::new(); // Degrades BufferFull to live forwarding after a fail-open cap hit. let mut fail_opened = false; let mut blocked = false; @@ -2493,6 +2969,8 @@ fn stream_response( let (parts, usage) = frame_parts(protocol, &frame); let held = parts.held(); let delta = parts.scan; + let guardrail_text = (!chain.is_empty()) + .then(|| stream_guardrail_text(protocol, &frame, delta.clone())); if let Some(u) = usage { merge_usage(&mut telemetry.usage, u); } @@ -2510,12 +2988,24 @@ fn stream_response( yield Ok::<_, std::convert::Infallible>(frame); } StreamOutputPolicy::EndOfStreamCheck => { - scan_buf.push_str(&delta); + if let Some(text) = guardrail_text.as_ref() { + append_stream_guardrail_text( + &mut continuation_bufs, + &mut supplemental_buf, + text, + ); + } telemetry.mark_first_delivery(); yield Ok(frame); } StreamOutputPolicy::Window { size_chars, overlap_chars, .. } => { - scan_buf.push_str(&delta); + if let Some(text) = guardrail_text.as_ref() { + append_stream_guardrail_text( + &mut continuation_bufs, + &mut supplemental_buf, + text, + ); + } held_bytes += frame.len(); pending_held.add(frame.len()); pending.push(frame); @@ -2524,10 +3014,18 @@ fn stream_response( // 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 + if continuation_bufs + .iter() + .any(|continuation| continuation.chars().count() >= *size_chars) + || supplemental_buf.chars().count() >= *size_chars || held_bytes > MAX_HELD_STREAM_BYTES { - let text = format!("{overlap_tail}{scan_buf}"); + let text = stream_guardrail_scan_text( + &continuation_tails, + &continuation_bufs, + &supplemental_tail, + &supplemental_buf, + ); match scan_output(&chain, &route_name, &text, &mut telemetry).await { GuardrailVerdict::Block { reason, @@ -2551,15 +3049,44 @@ fn stream_response( } pending_held.clear(); held_bytes = 0; - let combined = format!("{overlap_tail}{scan_buf}"); - overlap_tail = tail_chars(&combined, *overlap_chars); - scan_buf.clear(); + if continuation_tails.len() < continuation_bufs.len() { + continuation_tails.resize( + continuation_bufs.len(), + String::new(), + ); + } + for (index, continuation) in continuation_bufs.iter_mut().enumerate() { + let combined = format!( + "{}{}", + continuation_tails[index], + continuation.as_str(), + ); + continuation_tails[index] = + tail_chars(&combined, *overlap_chars); + continuation.clear(); + } + let combined_supplemental = if supplemental_tail.is_empty() { + supplemental_buf.clone() + } else if supplemental_buf.is_empty() { + supplemental_tail.clone() + } else { + format!("{supplemental_tail}\n{supplemental_buf}") + }; + supplemental_tail = + tail_chars(&combined_supplemental, *overlap_chars); + supplemental_buf.clear(); } } } } StreamOutputPolicy::BufferFull { max_buffer_bytes, on_exceeded_fail_open } => { - scan_buf.push_str(&delta); + if let Some(text) = guardrail_text.as_ref() { + append_stream_guardrail_text( + &mut continuation_bufs, + &mut supplemental_buf, + text, + ); + } held_content.hold(held, frame.len()); pending_held.add(frame.len()); pending.push(frame); @@ -2599,6 +3126,8 @@ fn stream_response( let (parts, usage) = frame_parts(protocol, &rest); let held = parts.held(); let delta = parts.scan; + let guardrail_text = (!chain.is_empty()) + .then(|| stream_guardrail_text(protocol, &rest, delta.clone())); if let Some(u) = usage { merge_usage(&mut telemetry.usage, u); } @@ -2609,7 +3138,13 @@ fn stream_response( capture_cap, ); } - scan_buf.push_str(&delta); + if let Some(text) = guardrail_text.as_ref() { + append_stream_guardrail_text( + &mut continuation_bufs, + &mut supplemental_buf, + text, + ); + } let rest = Bytes::from(rest); // The tail is held like any frame, under the same cap. let tripped = match &policy { @@ -2658,7 +3193,12 @@ fn stream_response( } } } - let text = format!("{overlap_tail}{scan_buf}"); + let text = stream_guardrail_scan_text( + &continuation_tails, + &continuation_bufs, + &supplemental_tail, + &supplemental_buf, + ); if !chain.is_empty() && !text.is_empty() { if let GuardrailVerdict::Block { reason, @@ -4068,6 +4608,11 @@ mod tests { "{protocol:?} must still offer the forwarded bytes to the scan, got {text:?}", ); } + + 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") + ); } /// The `[DONE]` sentinel is not content, on either framing. A stream @@ -4305,14 +4850,29 @@ mod tests { request_guardrail_text(PassthroughProtocol::OpenaiChat, not_chat), "x" ); - // A detected envelope whose items carry no typed text ALSO degrades - // to every decoded JSON string value — detection must never scan - // less than the forwarded request carries. - let empty_chat = br#"{"messages":[{"role":"tool","tool_call_id":"1"}]}"#; + // 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); - for text in ["tool", "1"] { - assert!(scanned.contains(text), "{text} missing from {scanned:?}"); - } + assert!(scanned.contains("fallback"), "{scanned:?}"); + } + + #[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:?}"); } #[test] @@ -4356,6 +4916,195 @@ mod tests { ); } + #[test] + 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 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 known_response_envelopes_scan_duplicate_and_nested_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", "NESTED", "clean"] { + assert!(scanned.contains(expected), "{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", "NESTED", "clean"] { + assert!(scanned.contains(expected), "{scanned:?}"); + } + + let conflicting_type = br#"{"output":[{"type":"reasoning","type":"message","content":[{"text":"\u0042LOCKME"}]}]}"#; + assert!( + response_guardrail_text(PassthroughProtocol::OpenaiResponses, conflicting_type) + .contains("BLOCKME") + ); + + let same_hidden_type = br#"{"output":[{"type":"reasoning","type":"reasoning","summary":[{"text":"\u0042LOCKME"}]}]}"#; + assert!( + !response_guardrail_text(PassthroughProtocol::OpenaiResponses, same_hidden_type) + .contains("BLOCKME") + ); + + let deep = format!( + r#"{{"output":[{{"type":"message","content":[{{"type":"output_text","nested":{}}}]}}]}}"#, + String::from_utf8(deeply_nested_escaped_block_json()).expect("valid test JSON") + ); + assert!( + response_guardrail_text(PassthroughProtocol::OpenaiResponses, deep.as_bytes()) + .contains("BLOCKME"), + "a known deep response must retain its decoded source leaf" + ); + } + + #[test] + 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") + ); + + 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 known_sse_envelopes_scan_duplicate_and_nested_source_strings() { + let chat = b"data: {\"choices\":[{\"delta\":{\"content\":\"\\u0042LOCKME\",\"metadata\":{\"note\":\"NESTED\"}}}],\"choices\":[{\"delta\":{\"content\":\"clean\"}}]}\n\n"; + let scanned = frame_guardrail_text(PassthroughProtocol::OpenaiChat, chat); + for expected in ["BLOCKME", "NESTED", "clean"] { + assert!(scanned.contains(expected), "{scanned:?}"); + } + + 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", "NESTED", "clean"] { + assert!(scanned.contains(expected), "{scanned:?}"); + } + + 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") + ); + + 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") + ); + + 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 stream_guardrail_text_keeps_visible_deltas_contiguous() { + let first = b"data: {\"type\":\"response.output_text.delta\",\"item_id\":\"one\",\"delta\":\"FOR\"}\n\n"; + let second = b"data: {\"type\":\"response.output_text.delta\",\"item_id\":\"two\",\"delta\":\"noise\",\"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 supplemental = String::new(); + append_stream_guardrail_text(&mut continuations, &mut supplemental, &first); + append_stream_guardrail_text(&mut continuations, &mut supplemental, &second); + let scanned = stream_guardrail_scan_text(&[], &continuations, "", &supplemental); + assert!(scanned.starts_with("FORBIDDEN"), "{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!(!scanned.contains("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!(!scanned.contains("BLOCKME"), "{scanned:?}"); + } + #[test] fn detect_protocol_from_request_envelope() { // The real Copilot CLI surface, one shape per endpoint family. @@ -4410,9 +5159,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"}"#; @@ -4424,9 +5173,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), 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 97a3749d0..00cd48efe 100644 --- a/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts @@ -42,6 +42,24 @@ const deepEscapedBlockJSON = (depth: number) => // provider receives it verbatim, so Raw guardrails must still decode the leaf. const DEEP_ESCAPED_BLOCK_JSON = deepEscapedBlockJSON(160); const CAP = 1_000; +const SPLIT_BLOCK = "FORBIDDEN"; +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 KNOWN_CHAT_STREAM = `${String.raw`data: {"choices":[{"delta":{"content":"\u0042LOCKME","metadata":{"note":"${OUT_LIT}"}}}],"choices":[{"delta":{"content":"clean"}}]}`}\n\n`; +const KNOWN_RESPONSES_STREAM = `${String.raw`data: {"type":"response.output_text.delta","delta":"\u0042LOCKME","delta":"clean","metadata":{"note":"${OUT_LIT}"}}`}\n\n`; +const SPLIT_RESPONSES_STREAM = [ + `data: {"type":"response.output_text.delta","item_id":"first","delta":"FOR"}\n\n`, + `data: {"type":"response.output_text.delta","item_id":"second","delta":"noise","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", part: { type: "reasoning_text", text: OUT_LIT } })}\n\n`, + `data: ${JSON.stringify({ type: "response.output_item.done", item: { type: "reasoning", summary: [{ type: "summary_text", text: OUT_LIT }] } })}\n\n`, + `data: ${JSON.stringify({ type: "response.output_text.delta", 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({ @@ -131,6 +149,26 @@ describe("passthrough guardrail scan coverage", () => { upstreams["raw-safe-stream"] = await startOpenAiUpstream({ rawStreamFrames: [`data: ${SAFE_ESCAPED_JSON}\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["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["known-responses-reasoning-stream"] = await startOpenAiUpstream({ + rawStreamFrames: RESPONSES_REASONING_STREAM, + }); const pk = await seed.createProviderKey({ display_name: "pt-scan-pk", secret: "sk-mock", @@ -152,6 +190,7 @@ describe("passthrough guardrail scan coverage", () => { patterns: [ { kind: "literal", value: OUT_LIT }, { kind: "literal", value: ESCAPED_BLOCK }, + { kind: "literal", value: SPLIT_BLOCK }, ], }); await seed.createGuardrail({ @@ -236,6 +275,90 @@ describe("passthrough guardrail scan coverage", () => { expect(body).not.toContain("event: error"); expect(body).toContain("visible answer"); }); + test.for([ + [ + "chat", + "known-chat-buffered-output", + `{"model":"gpt-4o-mini","messages":[{"role":"user","content":"go"}]}`, + ], + [ + "Responses", + "known-responses-buffered-output", + `{"model":"gpt-4o-mini","input":"go"}`, + ], + ] as const)("output: known %s envelope 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.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).not.toContain(SPLIT_BLOCK); + 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; @@ -319,6 +442,23 @@ describe("passthrough guardrail scan coverage", () => { 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([ ["ASCII", ESCAPED_BLOCK_JSON], From aecf7ad91e4878252518e4c09b50ad8994bc39be Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Wed, 30 Sep 2026 14:37:36 +0800 Subject: [PATCH 14/37] fix(guardrail): satisfy stream scan lint --- crates/aisix-proxy/src/passthrough_route.rs | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index 8b18c751d..8afdc82a1 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -2636,9 +2636,12 @@ fn stream_guardrail_text( && payload.as_ref().is_some_and(|payload| { hidden_chat_stream_reasoning_frame(payload.trim().as_bytes()) == Some(true) }); - let mut continuations = vec![(!hidden_reasoning) - .then_some(continuation) - .unwrap_or_default()]; + let typed_continuation = if hidden_reasoning { + String::new() + } else { + continuation + }; + let mut continuations = vec![typed_continuation]; if let Some(source) = payload.and_then(|payload| frame_source_continuations(protocol, payload.trim().as_bytes())) { From a64ce6be565bd0cc7c0c112aca6c002a57bb6545 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Wed, 30 Sep 2026 15:19:12 +0800 Subject: [PATCH 15/37] fix(guardrail): keep Responses media opaque --- crates/aisix-proxy/src/passthrough_route.rs | 522 +++++++++++++++--- ...ough-responses-media-guardrail-e2e.test.ts | 240 ++++++++ 2 files changed, 676 insertions(+), 86 deletions(-) create mode 100644 tests/e2e/src/cases/passthrough-responses-media-guardrail-e2e.test.ts diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index 8afdc82a1..45cf59ee4 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -1486,6 +1486,21 @@ fn raw_top_level_has_any_type(body: &[u8], wanted: &[&str]) -> bool { .any(|kind| wanted.iter().any(|wanted| kind == *wanted)) } +/// `true` only when every source `type` value is one of `allowed`. This is +/// stricter than [`raw_top_level_has_any_type`]: an audio or image event must +/// not borrow a text event's carrier merely by repeating a conflicting type. +fn raw_top_level_has_only_types(body: &[u8], allowed: &[&str]) -> bool { + let Some(values) = raw_top_level_values(body, "type") else { + return false; + }; + !values.is_empty() + && values.into_iter().all(|value| { + serde_json::from_str::(value.get()) + .ok() + .is_some_and(|kind| allowed.iter().any(|allowed| kind == *allowed)) + }) +} + fn raw_top_level_items_have_only_types(body: &[u8], key: &str, allowed: &[&str]) -> Option { let values = raw_top_level_values(body, key)?; Some( @@ -1527,6 +1542,19 @@ fn append_raw_top_level_strings(out: &mut String, body: &[u8], key: &str) -> Opt Some(()) } +/// Append direct string values but do not turn an unexpected non-string into +/// a raw-body fallback. Response output uses this at the privacy boundary: a +/// malformed optional text field must not make opaque sibling fields readable +/// by an external guardrail. +fn append_raw_top_level_string_values(out: &mut String, body: &[u8], key: &str) -> Option<()> { + for value in raw_top_level_values(body, key)? { + if let Ok(value) = serde_json::from_str::(value.get()) { + append_scan_text(out, &value); + } + } + Some(()) +} + fn raw_top_level_string_values(body: &[u8], key: &str) -> Option> { raw_top_level_values(body, key)? .into_iter() @@ -1758,20 +1786,110 @@ fn decoded_chat_response_string_values(body: &[u8]) -> Option { Some(out) } +const RESPONSES_VISIBLE_DELTA_EVENTS: &[&str] = &[ + "response.output_text.delta", + "response.function_call_arguments.delta", + "response.mcp_call_arguments.delta", + "response.custom_tool_call_input.delta", +]; + +/// The only Responses content-part `text` fields the typed output guardrail +/// reads. Other part types can carry image, audio, file, or reasoning data. +const RESPONSES_VISIBLE_TEXT_PART_TYPES: &[&str] = &["output_text", "text", "input_text"]; + +/// 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: &serde_json::value::RawValue, +) -> Option<()> { + let part_body = part.get().as_bytes(); + match raw_top_level_unique_type(part_body).as_deref() { + Some(kind) if RESPONSES_VISIBLE_TEXT_PART_TYPES.contains(&kind) => { + append_raw_top_level_string_values(out, part_body, "text")? + } + Some(_) | None => {} + } + Some(()) +} + +fn append_responses_visible_content_strings( + out: &mut String, + content: &serde_json::value::RawValue, +) -> 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 Some(parts) = raw_array_items(content) else { + return Some(()); + }; + for part in parts { + append_responses_visible_part_strings(out, &part)?; + } + Some(()) +} + +/// 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: &serde_json::value::RawValue, +) -> Option<()> { + let item_body = item.get().as_bytes(); + match raw_top_level_unique_type(item_body).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_string_values(out, item_body, key)?; + } + } + Some("custom_tool_call") => { + for key in ["name", "input"] { + append_raw_top_level_string_values(out, item_body, key)?; + } + } + Some(_) | None => {} + } + Some(()) +} + +fn append_responses_output_strings(out: &mut String, body: &[u8]) -> Option<()> { + for output in raw_top_level_values(body, "output")? { + let Some(items) = raw_array_items(&output) else { + continue; + }; + for item in items { + append_responses_output_item_strings(out, &item)?; + } + } + Some(()) +} + fn decoded_responses_response_string_values(body: &[u8]) -> Option { - let mut out = decoded_json_string_values_except_root_keys(body, &["model", "output"])?; - append_raw_array_item_strings(&mut out, body, "output", |item| { - raw_object_has_only_types(item, &["reasoning"]) - })?; + 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`: a -/// passthrough upstream may send duplicate or deeply nested fields which the -/// client receives verbatim. Generated reasoning remains out of scope only -/// for an unambiguous standard reasoning item. +/// 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(protocol: PassthroughProtocol, body: &[u8]) -> String { let raw = || String::from_utf8_lossy(body).into_owned(); match protocol { @@ -1783,7 +1901,10 @@ fn response_guardrail_text(protocol: PassthroughProtocol, body: &[u8]) -> String decoded_non_model_json_string_values(body).unwrap_or_else(raw) } PassthroughProtocol::OpenaiResponses => { - decoded_responses_response_string_values(body).unwrap_or_else(raw) + // 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() } } } @@ -2444,47 +2565,14 @@ fn decoded_chat_frame_string_values(body: &[u8]) -> Option { .ok() } +/// The typed stream extractor deliberately takes only text/tool delta events: +/// `.done`, output-item, and terminal response snapshots repeat those +/// carriers. Keep the raw source pass on that same boundary, both to avoid +/// duplicate external moderation and to keep any opaque terminal media out. fn decoded_responses_frame_string_values(body: &[u8]) -> Option { - match raw_top_level_unique_type(body).as_deref() { - Some("response.reasoning_text.delta" | "response.reasoning_summary_text.delta") => { - return decoded_json_string_values_except_root_keys(body, &["model", "delta"]); - } - Some("response.reasoning_text.done" | "response.reasoning_summary_text.done") => { - return decoded_json_string_values_except_root_keys(body, &["model", "text"]); - } - Some("response.reasoning_summary_part.added" | "response.reasoning_summary_part.done") => { - return decoded_json_string_values_except_root_keys(body, &["model", "part"]); - } - Some("response.output_item.added" | "response.output_item.done") - if raw_top_level_items_have_only_types(body, "item", &["reasoning"])? => - { - return decoded_json_string_values_except_root_keys(body, &["model", "item"]); - } - Some("response.content_part.added" | "response.content_part.done") - if raw_top_level_items_have_only_types( - body, - "part", - &["reasoning_text", "reasoning_summary"], - )? => - { - return decoded_json_string_values_except_root_keys(body, &["model", "part"]); - } - _ => {} - } - - let responses = raw_top_level_values(body, "response")?; - if responses.is_empty() { - return decoded_non_model_json_string_values(body); - } - let mut out = decoded_json_string_values_except_root_keys(body, &["model", "response"])?; - for response in responses { - append_scan_text( - &mut out, - &response_guardrail_text( - PassthroughProtocol::OpenaiResponses, - response.get().as_bytes(), - ), - ); + let mut out = String::new(); + if raw_top_level_has_only_types(body, RESPONSES_VISIBLE_DELTA_EVENTS) { + append_raw_top_level_string_values(&mut out, body, "delta")?; } Some(out) } @@ -2558,14 +2646,94 @@ fn decoded_completions_frame_continuations(body: &[u8]) -> Option> { }) } +fn is_chat_choice_continuation_path(path: &[crate::json_splice::PathSeg]) -> bool { + use crate::json_splice::PathSeg; + + matches!( + path, + [PathSeg::Key(choices), PathSeg::Index(_), PathSeg::Key(delta), PathSeg::Key(field)] + if choices == "choices" && delta == "delta" && field == "content" + ) || matches!( + path, + [PathSeg::Key(choices), PathSeg::Index(_), PathSeg::Key(delta), PathSeg::Key(content), PathSeg::Index(_), PathSeg::Key(text)] + if choices == "choices" + && delta == "delta" + && content == "content" + && text == "text" + ) || matches!( + path, + [PathSeg::Key(choices), PathSeg::Index(_), PathSeg::Key(delta), PathSeg::Key(tool_calls), PathSeg::Index(_), PathSeg::Key(kind), PathSeg::Key(value)] + if choices == "choices" + && delta == "delta" + && tool_calls == "tool_calls" + && ((kind == "function" && value == "arguments") + || (kind == "custom" && value == "input")) + ) +} + +fn is_anthropic_delta_continuation_path(path: &[crate::json_splice::PathSeg]) -> bool { + use crate::json_splice::PathSeg; + + matches!( + path, + [PathSeg::Key(delta), PathSeg::Key(field)] + if delta == "delta" && matches!(field.as_str(), "text" | "partial_json") + ) +} + +fn is_anthropic_start_continuation_path(path: &[crate::json_splice::PathSeg]) -> bool { + use crate::json_splice::PathSeg; + + matches!( + path, + [PathSeg::Key(content_block), PathSeg::Key(field)] + if content_block == "content_block" && field == "text" + ) || matches!( + path, + [PathSeg::Key(content_block), PathSeg::Key(input), ..] + if content_block == "content_block" && input == "input" + ) +} + +fn is_completions_continuation_path(path: &[crate::json_splice::PathSeg]) -> bool { + use crate::json_splice::PathSeg; + + matches!( + path, + [PathSeg::Key(choices), PathSeg::Index(_), PathSeg::Key(text)] + if choices == "choices" && text == "text" + ) +} + +/// 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_string_values(body: &[u8]) -> Option { + let anthropic_delta = raw_top_level_has_any_type(body, &["content_block_delta"]); + let anthropic_start = + !anthropic_delta && raw_top_level_has_any_type(body, &["content_block_start"]); + decoded_json_string_values_where(body, |path| { + !is_root_key(path, "model") + && !is_hidden_chat_stream_reasoning_path(path) + && !(if anthropic_delta { + is_anthropic_delta_continuation_path(path) + } else if anthropic_start { + is_anthropic_start_continuation_path(path) + } else { + is_chat_choice_continuation_path(path) + }) + }) +} + +fn decoded_completions_frame_supplemental_string_values(body: &[u8]) -> Option { + decoded_json_string_values_where(body, |path| { + !is_root_key(path, "model") && !is_completions_continuation_path(path) + }) +} + fn decoded_responses_frame_continuations(body: &[u8]) -> Option> { - const VISIBLE_DELTA_EVENTS: &[&str] = &[ - "response.output_text.delta", - "response.function_call_arguments.delta", - "response.mcp_call_arguments.delta", - "response.custom_tool_call_input.delta", - ]; - if !raw_top_level_has_any_type(body, VISIBLE_DELTA_EVENTS) { + if !raw_top_level_has_only_types(body, RESPONSES_VISIBLE_DELTA_EVENTS) { return None; } raw_top_level_string_values(body, "delta").filter(|values| !values.is_empty()) @@ -2583,9 +2751,49 @@ fn frame_source_continuations( } } +fn frame_guardrail_supplemental_text( + protocol: PassthroughProtocol, + frame: &[u8], + has_source_continuations: bool, + has_typed_continuation: bool, +) -> String { + // On a malformed or unknown frame, typed extraction is the only + // available output carrier. It already contains the raw fallback, so a + // second generic scan would double-count it. + if has_typed_continuation && !has_source_continuations { + return String::new(); + } + if !has_source_continuations { + return frame_guardrail_text(protocol, frame); + } + + 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 { + // The raw source continuation is the complete decoded payload. + PassthroughProtocol::Raw => String::new(), + PassthroughProtocol::OpenaiChat => { + decoded_chat_frame_supplemental_string_values(payload.as_bytes()).unwrap_or_default() + } + PassthroughProtocol::OpenaiCompletions => { + decoded_completions_frame_supplemental_string_values(payload.as_bytes()) + .unwrap_or_default() + } + // Responses source continuations exist only for the explicitly safe + // text/tool delta events, whose sole output carrier is `delta`. + PassthroughProtocol::OpenaiResponses => String::new(), + } +} + /// 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 also sees duplicate and arbitrary forwarded fields. +/// source-preserving pass retains duplicate selected carriers. Responses +/// media fields remain opaque even when forwarded verbatim. fn frame_guardrail_text(protocol: PassthroughProtocol, frame: &[u8]) -> String { let Some(payload) = crate::redact::frame_payload(frame) else { return String::new(); @@ -2606,16 +2814,19 @@ fn frame_guardrail_text(protocol: PassthroughProtocol, frame: &[u8]) -> String { decoded_non_model_json_string_values(payload.as_bytes()).unwrap_or_else(raw) } PassthroughProtocol::OpenaiResponses => { - decoded_responses_frame_string_values(payload.as_bytes()).unwrap_or_else(raw) + // 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 channels scanned for one stream frame. The first -/// continuation is the typed visible-output sequence; later channels retain -/// source output-carrier occurrences across frames. `supplemental` preserves -/// every additional decoded source value (including duplicate or nested -/// fields). Keeping them separate prevents frame metadata from interrupting a +/// continuation is the typed visible-output sequence when it maps to one raw +/// carrier; later channels retain source-only duplicate occurrences across +/// frames. `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, @@ -2627,29 +2838,59 @@ fn stream_guardrail_text( frame: &[u8], continuation: String, ) -> StreamGuardrailText { - // Keep the typed continuation in a stable first channel even when source - // extraction finds raw carriers. A provider can add or remove duplicate - // carrier fields between frames; only the typed last-wins sequence then - // follows what a normal JSON client sees across that boundary. + // Keep the typed continuation in a stable first channel when it maps to + // one raw carrier. Source-only duplicate occurrences occupy later + // channels, preserving bytes a JSON client would discard without sending + // the canonical value to the guardrail twice. 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 typed_continuation = if hidden_reasoning { + let responses_visible_delta = matches!(protocol, PassthroughProtocol::OpenaiResponses) + && payload.as_ref().is_some_and(|payload| { + raw_top_level_has_only_types(payload.trim().as_bytes(), RESPONSES_VISIBLE_DELTA_EVENTS) + }); + let typed_continuation = if hidden_reasoning + || (matches!(protocol, PassthroughProtocol::OpenaiResponses) && !responses_visible_delta) + { String::new() } else { continuation }; - let mut continuations = vec![typed_continuation]; - if let Some(source) = - payload.and_then(|payload| frame_source_continuations(protocol, payload.trim().as_bytes())) - { - continuations.extend(source); - } + let has_typed_continuation = !typed_continuation.is_empty(); + let source = payload + .and_then(|payload| frame_source_continuations(protocol, payload.trim().as_bytes())) + .filter(|source| !source.is_empty()); + let has_source_continuations = source.is_some(); + let continuations = match source { + Some(mut source) if has_typed_continuation => { + if let Some(canonical) = source + .iter() + .rposition(|candidate| candidate == &typed_continuation) + { + source.remove(canonical); + let mut continuations = vec![typed_continuation]; + continuations.extend(source); + continuations + } else { + // Multiple source carriers can make a typed extractor join + // distinct values. Keep the source occurrences only rather + // than scanning their joined representation in addition. + source + } + } + Some(source) => source, + None => vec![typed_continuation], + }; StreamGuardrailText { continuations, - supplemental: frame_guardrail_text(protocol, frame), + supplemental: frame_guardrail_supplemental_text( + protocol, + frame, + has_source_continuations, + has_typed_continuation, + ), } } @@ -4964,7 +5205,7 @@ mod tests { } #[test] - fn known_response_envelopes_scan_duplicate_and_nested_source_strings() { + 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", "NESTED", "clean"] { @@ -4974,13 +5215,14 @@ mod tests { 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", "NESTED", "clean"] { + for expected in ["BLOCKME", "clean"] { assert!(scanned.contains(expected), "{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) + !response_guardrail_text(PassthroughProtocol::OpenaiResponses, conflicting_type) .contains("BLOCKME") ); @@ -4990,14 +5232,107 @@ mod tests { .contains("BLOCKME") ); - let deep = format!( - r#"{{"output":[{{"type":"message","content":[{{"type":"output_text","nested":{}}}]}}]}}"#, - String::from_utf8(deeply_nested_escaped_block_json()).expect("valid test JSON") + 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, deep.as_bytes()) - .contains("BLOCKME"), - "a known deep response must retain its decoded source leaf" + response_guardrail_text(PassthroughProtocol::OpenaiResponses, duplicate_text) + .contains("clean") + ); + } + + #[test] + 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 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 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!(!scanned.contains("STREAM_MEDIA_SENTINEL"), "{scanned:?}"); + + // The terminal response repeats prior delta content, including media + // from image-generation output. It is never a second scan carrier. + 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!( + !scanned.contains("CONFLICTING_MEDIA_SENTINEL"), + "{scanned:?}" + ); + + 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 = 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") ); } @@ -5033,19 +5368,20 @@ mod tests { 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", "NESTED", "clean"] { + for expected in ["BLOCKME", "clean"] { assert!(scanned.contains(expected), "{scanned:?}"); } + assert!(!scanned.contains("NESTED"), "{scanned:?}"); 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) + !frame_guardrail_text(PassthroughProtocol::OpenaiResponses, conflicting) .contains("BLOCKME") ); 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) + !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"; @@ -5090,6 +5426,20 @@ mod tests { assert!(scanned.starts_with("FORBIDDEN"), "{scanned:?}"); } + #[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\":[{\"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!( + scanned.matches(email).count(), + 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"; 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, + ); + }); +}); From 19b4bc5c86b6bc791d61b5c3f058ad1f7805ce19 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Wed, 30 Sep 2026 20:03:18 +0800 Subject: [PATCH 16/37] fix(guardrail): keep malformed envelopes private --- crates/aisix-proxy/src/passthrough_route.rs | 149 ++++++++++++++++---- 1 file changed, 121 insertions(+), 28 deletions(-) diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index 45cf59ee4..02aa5e8e4 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -1257,20 +1257,15 @@ 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. `RawValue` keeps the chosen top-level + // carrier shallow 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 @@ -1444,6 +1439,24 @@ fn raw_top_level_values( 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 { + raw_top_level_values(body, key) + .and_then(|values| values.into_iter().last()) + .is_some_and(|value| match value.get().trim_start().as_bytes().first() { + Some(b'[') => true, + Some(b'"') => allow_string, + _ => false, + }) +} + +fn raw_is_object(raw: &serde_json::value::RawValue) -> bool { + raw.get().trim_start().starts_with('{') +} + fn raw_array_items( raw: &serde_json::value::RawValue, ) -> Option>> { @@ -1537,17 +1550,9 @@ fn append_raw_string_value(out: &mut String, raw: &serde_json::value::RawValue) fn append_raw_top_level_strings(out: &mut String, body: &[u8], key: &str) -> Option<()> { for value in raw_top_level_values(body, key)? { - append_raw_string_value(out, &value)?; - } - Some(()) -} - -/// Append direct string values but do not turn an unexpected non-string into -/// a raw-body fallback. Response output uses this at the privacy boundary: a -/// malformed optional text field must not make opaque sibling fields readable -/// by an external guardrail. -fn append_raw_top_level_string_values(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); } @@ -1574,6 +1579,9 @@ fn append_raw_text_value(out: &mut String, raw: &serde_json::value::RawValue) -> return Some(()); } for part in raw_array_items(raw)? { + if !raw_is_object(&part) { + continue; + } append_raw_top_level_strings(out, part.get().as_bytes(), "text")?; } Some(()) @@ -1596,6 +1604,9 @@ fn append_chat_request_content_strings( return Some(()); } for block in raw_array_items(content)? { + if !raw_is_object(&block) { + continue; + } let block_body = block.get().as_bytes(); let types = raw_top_level_values(block_body, "type")?; let kind = raw_top_level_unique_type(block_body); @@ -1632,6 +1643,9 @@ fn append_chat_request_message_strings( out: &mut String, message: &serde_json::value::RawValue, ) -> Option<()> { + if !raw_is_object(message) { + return Some(()); + } let message_body = message.get().as_bytes(); append_scan_text( out, @@ -1660,7 +1674,13 @@ fn decoded_chat_request_string_values(body: &[u8]) -> Option { append_chat_request_content_strings(&mut out, &system)?; } for array in raw_top_level_values(body, "messages")? { - for message in raw_array_items(&array)? { + // 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 { + continue; + }; + for message in messages { append_chat_request_message_strings(&mut out, &message)?; } } @@ -1671,6 +1691,9 @@ fn append_responses_item_strings( out: &mut String, item: &serde_json::value::RawValue, ) -> Option<()> { + if !raw_is_object(item) { + return Some(()); + } let item_body = item.get().as_bytes(); let text_keys = [ "content", @@ -1808,7 +1831,7 @@ fn append_responses_visible_part_strings( let part_body = part.get().as_bytes(); match raw_top_level_unique_type(part_body).as_deref() { Some(kind) if RESPONSES_VISIBLE_TEXT_PART_TYPES.contains(&kind) => { - append_raw_top_level_string_values(out, part_body, "text")? + append_raw_top_level_strings(out, part_body, "text")? } Some(_) | None => {} } @@ -1854,12 +1877,12 @@ fn append_responses_output_item_strings( } Some("function_call" | "mcp_call") => { for key in ["name", "arguments"] { - append_raw_top_level_string_values(out, item_body, key)?; + append_raw_top_level_strings(out, item_body, key)?; } } Some("custom_tool_call") => { for key in ["name", "input"] { - append_raw_top_level_string_values(out, item_body, key)?; + append_raw_top_level_strings(out, item_body, key)?; } } Some(_) | None => {} @@ -2572,7 +2595,7 @@ fn decoded_chat_frame_string_values(body: &[u8]) -> Option { fn decoded_responses_frame_string_values(body: &[u8]) -> Option { let mut out = String::new(); if raw_top_level_has_only_types(body, RESPONSES_VISIBLE_DELTA_EVENTS) { - append_raw_top_level_string_values(&mut out, body, "delta")?; + append_raw_top_level_strings(&mut out, body, "delta")?; } Some(out) } @@ -3980,6 +4003,15 @@ mod tests { 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 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"}}"# @@ -5119,6 +5151,67 @@ mod tests { assert!(!scanned.contains("BLOCKME"), "{scanned:?}"); } + #[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:?}" + ); + } + + #[test] + fn duplicate_chat_carrier_skips_malformed_source_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 scanned = request_guardrail_text(PassthroughProtocol::OpenaiChat, request); + assert!(scanned.contains("clean"), "{scanned:?}"); + assert!( + !scanned.contains("BLOCKME"), + "a malformed duplicate must not trigger raw fallback: {scanned:?}" + ); + } + + #[test] + fn malformed_typed_array_items_keep_media_opaque() { + 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 scanned = request_guardrail_text(protocol, request); + assert!(scanned.contains("clean"), "{protocol:?}: {scanned:?}"); + assert!( + !scanned.contains("BLOCKME"), + "a malformed array item must not trigger raw fallback for {protocol:?}: {scanned:?}" + ); + } + } + #[test] fn guardrail_text_scans_decoded_forwarded_json_strings() { let raw = br#"{"state":"\u0042LOCKME","nested":{"query":"\u4e2d\u6587"}}"#; From 95fa14f8dea0a5b0d20c7d8fbb46b56acb6c9126 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Wed, 30 Sep 2026 21:44:19 +0800 Subject: [PATCH 17/37] fix(passthrough): preserve stream guardrail source boundaries --- crates/aisix-proxy/src/passthrough_route.rs | 1917 +++++++++++++++-- .../cases/passthrough-guardrail-e2e.test.ts | 4 +- .../passthrough-scan-coverage-e2e.test.ts | 291 ++- 3 files changed, 2013 insertions(+), 199 deletions(-) diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index 02aa5e8e4..1ba33865f 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -1399,6 +1399,15 @@ fn decoded_json_string_values_except_root_keys(body: &[u8], excluded: &[&str]) - .ok() } +fn decoded_json_string_values_except_root_keys_vec( + body: &[u8], + excluded: &[&str], +) -> Option> { + decoded_json_string_values_vec_where(body, |path| { + !excluded.iter().any(|key| is_root_key(path, key)) + }) +} + /// Source values of all occurrences of one top-level key. `RawValue` keeps /// repeated keys separate, unlike `serde_json::Value`. fn raw_top_level_values( @@ -1567,6 +1576,63 @@ fn raw_top_level_string_values(body: &[u8], key: &str) -> Option> { .collect() } +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(()), + } +} + +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( + body: &[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).ok_or(()) + } + _ => Err(()), + } +} + +fn raw_top_level_unique_array( + body: &[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) + .ok_or(()) + } + _ => Err(()), + } +} + /// 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. @@ -2423,6 +2489,21 @@ fn path_is_within_base(candidate: &str, base: &str) -> bool { /// bounded while the policy semantics degrade gracefully. 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 @@ -2560,6 +2641,7 @@ fn hidden_chat_stream_reasoning_frame(body: &[u8]) -> Option { } } +#[cfg(test)] fn decoded_chat_frame_string_values(body: &[u8]) -> Option { match raw_top_level_unique_type(body).as_deref() { Some("content_block_delta") @@ -2592,6 +2674,7 @@ fn decoded_chat_frame_string_values(body: &[u8]) -> Option { /// `.done`, output-item, and terminal response snapshots repeat those /// carriers. Keep the raw source pass on that same boundary, both to avoid /// duplicate external moderation and to keep any opaque terminal media out. +#[cfg(test)] fn decoded_responses_frame_string_values(body: &[u8]) -> Option { let mut out = String::new(); if raw_top_level_has_only_types(body, RESPONSES_VISIBLE_DELTA_EVENTS) { @@ -2732,11 +2815,11 @@ fn is_completions_continuation_path(path: &[crate::json_splice::PathSeg]) -> boo /// 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_string_values(body: &[u8]) -> Option { +fn decoded_chat_frame_supplemental_values(body: &[u8]) -> Option> { let anthropic_delta = raw_top_level_has_any_type(body, &["content_block_delta"]); let anthropic_start = !anthropic_delta && raw_top_level_has_any_type(body, &["content_block_start"]); - decoded_json_string_values_where(body, |path| { + decoded_json_string_values_vec_where(body, |path| { !is_root_key(path, "model") && !is_hidden_chat_stream_reasoning_path(path) && !(if anthropic_delta { @@ -2749,67 +2832,634 @@ fn decoded_chat_frame_supplemental_string_values(body: &[u8]) -> Option }) } -fn decoded_completions_frame_supplemental_string_values(body: &[u8]) -> Option { - decoded_json_string_values_where(body, |path| { +fn decoded_completions_frame_supplemental_values(body: &[u8]) -> Option> { + decoded_json_string_values_vec_where(body, |path| { !is_root_key(path, "model") && !is_completions_continuation_path(path) }) } -fn decoded_responses_frame_continuations(body: &[u8]) -> Option> { - if !raw_top_level_has_only_types(body, RESPONSES_VISIBLE_DELTA_EVENTS) { - return None; +#[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_source_continuations(payload: &[u8]) -> SourceContinuations { + let kind = match raw_top_level_unique_string(payload, "type") { + Ok(Some(kind)) if RESPONSES_VISIBLE_DELTA_EVENTS.contains(&kind.as_str()) => kind, + Ok(Some(_)) | Ok(None) => return SourceContinuations::Absent, + Err(()) => return SourceContinuations::Unevaluable, + }; + let item_id = match raw_top_level_unique_string(payload, "item_id") { + Ok(Some(item_id)) => match bounded_stream_source_id(item_id) { + Ok(item_id) => item_id, + Err(()) => return SourceContinuations::Unevaluable, + }, + Ok(None) | Err(()) => return SourceContinuations::Unevaluable, + }; + let output_index = match raw_top_level_unique_index(payload, "output_index") { + Ok(Some(index)) => index.to_string(), + Ok(None) | Err(()) => return SourceContinuations::Unevaluable, + }; + let content_index = if kind == "response.output_text.delta" { + match raw_top_level_unique_index(payload, "content_index") { + Ok(Some(index)) => index.to_string(), + Ok(None) | Err(()) => return SourceContinuations::Unevaluable, + } + } else { + // Tool-argument deltas have no content part. Their stable identity is + // the item plus output index and event kind; an incidental content + // field must not create a second source channel. + String::new() + }; + let values = match raw_top_level_string_values(payload, "delta") { + Some(values) => values, + None => return SourceContinuations::Unevaluable, + }; + if values.is_empty() { + return SourceContinuations::Absent; + } + let mut out = Vec::new(); + let mut keys = std::collections::HashSet::new(); + let mut source_values = Vec::new(); + let family = format!("responses:{item_id:?}"); + let identity = format!("{kind:?}:{output_index}:{content_index}:delta"); + if append_source_branches( + &mut out, + &mut keys, + &mut source_values, + family, + identity, + false, + values, + ) + .is_err() + { + return SourceContinuations::Unevaluable; + } + SourceContinuations::Ready(out) +} + +fn append_raw_string_carrier( + out: &mut Vec, + keys: &mut std::collections::HashSet, + source_values: &mut Vec, + family: String, + identity: String, + identity_is_ambiguous: bool, + body: &[u8], + key: &str, +) -> Result<(), ()> { + append_source_branches( + out, + keys, + source_values, + family, + identity, + 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` 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(); + append_raw_string_carrier( + &mut out, + &mut keys, + &mut source_values, + family.clone(), + "text".to_owned(), + false, + delta_body, + "text", + ) + .and_then(|()| { + append_raw_string_carrier( + &mut out, + &mut keys, + &mut source_values, + family.clone(), + "partial_json".to_owned(), + false, + delta_body, + "partial_json", + ) + }) + } + "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(); + append_raw_string_carrier( + &mut out, + &mut keys, + &mut source_values, + family.clone(), + "text".to_owned(), + false, + block_body, + "text", + ) + .and_then(|()| { + let mut inputs = raw_top_level_values(block_body, "input").ok_or(())?; + let Some(input) = inputs.pop() else { + return Ok(()); + }; + if !inputs.is_empty() { + return Err(()); + } + if input.get().trim_start().starts_with('"') { + return append_raw_string_carrier( + &mut out, + &mut keys, + &mut source_values, + family.clone(), + "input".to_owned(), + false, + block_body, + "input", + ); + } + // 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 = crate::json_splice::collect_string_values(input.get().as_bytes()) + .map_err(|_| ())?; + if text.is_empty() { + Ok(()) + } else { + Err(()) + } + }) + } + _ => return SourceContinuations::Absent, + }; + if result.is_err() { + SourceContinuations::Unevaluable + } else if !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 text_values = + match raw_top_level_string_values(part.get().as_bytes(), "text") { + Some(text_values) => text_values, + None => return SourceContinuations::Unevaluable, + }; + if text_values.is_empty() { + continue; + } + let part_id = match raw_part_identity(part.get().as_bytes()) { + 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}:text"), + false, + text_values, + ) + .is_err() + { + return SourceContinuations::Unevaluable; + } + } + } + Some(b'n') => {} + _ => 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 [("function", "arguments"), ("custom", "input")] { + 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, + format!("chat:{choice_index}:tool:{tool_index}"), + format!("{container}:{field}"), + false, + nested.get().as_bytes(), + field, + ) + .is_err() + { + return SourceContinuations::Unevaluable; + } + } + } + } + } + } + source_values_match_expected(source_values, expected) + .then_some(SourceContinuations::Ready(out)) + .unwrap_or(SourceContinuations::Unevaluable) +} + +fn completions_source_continuations(payload: &[u8]) -> SourceContinuations { + let Some(expected) = decoded_completions_frame_continuations(payload) else { + return SourceContinuations::Absent; + }; + 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, + }; + if append_raw_string_carrier( + &mut out, + &mut keys, + &mut source_values, + format!("completions:{choice_index}"), + "text".to_owned(), + false, + choice_body, + "text", + ) + .is_err() + { + return SourceContinuations::Unevaluable; + } } - raw_top_level_string_values(body, "delta").filter(|values| !values.is_empty()) + source_values_match_expected(source_values, expected) + .then_some(SourceContinuations::Ready(out)) + .unwrap_or(SourceContinuations::Unevaluable) } -fn frame_source_continuations( +fn stream_source_continuations( protocol: PassthroughProtocol, payload: &[u8], -) -> Option> { +) -> SourceContinuations { + let payload = payload.trim_ascii(); + if payload.is_empty() || payload == b"[DONE]" { + return SourceContinuations::Absent; + } match protocol { - PassthroughProtocol::Raw => decoded_json_string_values(payload).map(|text| vec![text]), - PassthroughProtocol::OpenaiChat => decoded_chat_frame_continuations(payload), - PassthroughProtocol::OpenaiCompletions => decoded_completions_frame_continuations(payload), - PassthroughProtocol::OpenaiResponses => decoded_responses_frame_continuations(payload), + // 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 => match serde_json::from_slice::(payload) { + Ok(text) if !text.is_empty() => SourceContinuations::Ready(vec![StreamContinuation { + key: "raw:payload:first".to_owned(), + family: "raw".to_owned(), + identity: "payload".to_owned(), + identity_is_ambiguous: false, + text, + }]), + Ok(_) => SourceContinuations::Absent, + Err(_) => SourceContinuations::Unevaluable, + }, + PassthroughProtocol::OpenaiChat => chat_choice_source_continuations(payload), + PassthroughProtocol::OpenaiCompletions => completions_source_continuations(payload), + PassthroughProtocol::OpenaiResponses => responses_source_continuations(payload), } } -fn frame_guardrail_supplemental_text( +fn responses_terminal_continuation_prefix(payload: &[u8]) -> Result, ()> { + 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, -) -> String { +) -> Vec { // On a malformed or unknown frame, typed extraction is the only // available output carrier. It already contains the raw fallback, so a // second generic scan would double-count it. if has_typed_continuation && !has_source_continuations { - return String::new(); + return Vec::new(); } if !has_source_continuations { - return frame_guardrail_text(protocol, frame); + return frame_guardrail_values(protocol, frame); } let Some(payload) = crate::redact::frame_payload(frame) else { - return String::new(); + return Vec::new(); }; let payload = payload.trim(); if payload.is_empty() || payload == "[DONE]" { - return String::new(); + return Vec::new(); } match protocol { // The raw source continuation is the complete decoded payload. - PassthroughProtocol::Raw => String::new(), + PassthroughProtocol::Raw => Vec::new(), PassthroughProtocol::OpenaiChat => { - decoded_chat_frame_supplemental_string_values(payload.as_bytes()).unwrap_or_default() + decoded_chat_frame_supplemental_values(payload.as_bytes()).unwrap_or_default() } PassthroughProtocol::OpenaiCompletions => { - decoded_completions_frame_supplemental_string_values(payload.as_bytes()) - .unwrap_or_default() + decoded_completions_frame_supplemental_values(payload.as_bytes()).unwrap_or_default() } // Responses source continuations exist only for the explicitly safe // text/tool delta events, whose sole output carrier is `delta`. - PassthroughProtocol::OpenaiResponses => String::new(), + PassthroughProtocol::OpenaiResponses => Vec::new(), + } +} + +fn decoded_chat_frame_values(body: &[u8]) -> Option> { + match raw_top_level_unique_type(body).as_deref() { + Some("content_block_delta") + if raw_top_level_items_have_only_types( + body, + "delta", + &["thinking_delta", "signature_delta"], + )? => + { + return decoded_json_string_values_except_root_keys_vec(body, &["model", "delta"]); + } + Some("content_block_start") + if raw_top_level_items_have_only_types( + body, + "content_block", + &["thinking", "redacted_thinking"], + )? => + { + return decoded_json_string_values_except_root_keys_vec( + body, + &["model", "content_block"], + ); + } + _ => {} + } + decoded_json_string_values_vec_where(body, |path| { + !is_root_key(path, "model") && !is_hidden_chat_stream_reasoning_path(path) + }) +} + +fn decoded_responses_frame_values(body: &[u8]) -> Option> { + raw_top_level_has_only_types(body, RESPONSES_VISIBLE_DELTA_EVENTS) + .then(|| raw_top_level_string_values(body, "delta")) + .flatten() + .or_else(|| Some(Vec::new())) +} + +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(); + } + let raw = || vec![payload.to_string()]; + match protocol { + PassthroughProtocol::Raw => { + decoded_json_string_values_vec_where(payload.as_bytes(), |_| true).unwrap_or_else(raw) + } + PassthroughProtocol::OpenaiChat => { + decoded_chat_frame_values(payload.as_bytes()).unwrap_or_else(raw) + } + PassthroughProtocol::OpenaiCompletions => { + decoded_json_string_values_vec_where(payload.as_bytes(), |path| { + !is_root_key(path, "model") + }) + .unwrap_or_else(raw) + } + PassthroughProtocol::OpenaiResponses => { + decoded_responses_frame_values(payload.as_bytes()).unwrap_or_default() + } } } @@ -2817,6 +3467,7 @@ fn frame_guardrail_supplemental_text( /// their typed visible-content extraction in [`frame_parts`], while this /// source-preserving pass retains duplicate selected carriers. 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(); @@ -2845,15 +3496,20 @@ fn frame_guardrail_text(protocol: PassthroughProtocol, frame: &[u8]) -> String { } } -/// The independent channels scanned for one stream frame. The first -/// continuation is the typed visible-output sequence when it maps to one raw -/// carrier; later channels retain source-only duplicate occurrences across -/// frames. `supplemental` preserves selected non-carrier source values. -/// Keeping them separate prevents frame metadata from interrupting a -/// sensitive literal split across output deltas. +/// 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, - supplemental: String, + 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, + /// 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( @@ -2861,10 +3517,6 @@ fn stream_guardrail_text( frame: &[u8], continuation: String, ) -> StreamGuardrailText { - // Keep the typed continuation in a stable first channel when it maps to - // one raw carrier. Source-only duplicate occurrences occupy later - // channels, preserving bytes a JSON client would discard without sending - // the canonical value to the guardrail twice. let payload = crate::redact::frame_payload(frame); let hidden_reasoning = matches!(protocol, PassthroughProtocol::OpenaiChat) && payload.as_ref().is_some_and(|payload| { @@ -2882,76 +3534,262 @@ fn stream_guardrail_text( continuation }; let has_typed_continuation = !typed_continuation.is_empty(); - let source = payload - .and_then(|payload| frame_source_continuations(protocol, payload.trim().as_bytes())) - .filter(|source| !source.is_empty()); - let has_source_continuations = source.is_some(); - let continuations = match source { - Some(mut source) if has_typed_continuation => { - if let Some(canonical) = source - .iter() - .rposition(|candidate| candidate == &typed_continuation) + 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())) { - source.remove(canonical); - let mut continuations = vec![typed_continuation]; - continuations.extend(source); - continuations - } else { - // Multiple source carriers can make a typed extractor join - // distinct values. Keep the source occurrences only rather - // than scanning their joined representation in addition. - source + 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) } - Some(source) => source, - None => vec![typed_continuation], + // 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 unevaluable = source_unevaluable || terminal_unevaluable; StreamGuardrailText { continuations, - supplemental: frame_guardrail_supplemental_text( - protocol, - frame, - has_source_continuations, - has_typed_continuation, - ), + supplemental: (!unevaluable) + .then(|| { + frame_guardrail_supplemental_values( + protocol, + frame, + has_source_continuations, + has_typed_continuation, + ) + }) + .unwrap_or_default(), + unevaluable, + closed_prefixes, } } fn append_stream_guardrail_text( - continuations: &mut Vec, - supplemental: &mut String, + 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() + .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 { + let is_closed = |key: &str| { + closed_prefixes + .iter() + .chain(text.closed_prefixes.iter()) + .any(|prefix| key.starts_with(prefix)) + }; + if text + .continuations + .iter() + .any(|continuation| is_closed(&continuation.key)) + { + return true; + } + 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, ) { - if continuations.len() < text.continuations.len() { - continuations.resize(text.continuations.len(), String::new()); + 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; } - for (index, continuation) in text.continuations.iter().enumerate() { - continuations[index].push_str(continuation); + let Some(total) = queued_candidates.checked_add(candidates.len()) else { + return false; + }; + if total > MAX_STREAM_GUARDRAIL_CHANNELS { + return false; } - append_scan_text(supplemental, &text.supplemental); + *queued_candidates = total; + sealed_epochs.push(candidates); + true } fn stream_guardrail_scan_text( - continuation_tails: &[String], - continuations: &[String], - supplemental_tail: &str, - supplemental: &str, -) -> String { - let mut text = String::new(); - for (index, continuation) in continuations.iter().enumerate() { + 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 - .get(index) - .map(String::as_str) + .iter() + .find(|candidate| candidate.key == continuation.key) + .map(|candidate| candidate.text.as_str()) .unwrap_or_default(); - append_scan_text(&mut text, &format!("{tail}{continuation}")); + 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()); + } } - let supplemental = match (supplemental_tail.is_empty(), supplemental.is_empty()) { - (true, true) => String::new(), - (false, true) => supplemental_tail.to_string(), - (true, false) => supplemental.to_string(), - (false, false) => format!("{supplemental_tail}\n{supplemental}"), - }; - append_scan_text(&mut text, &supplemental); text } @@ -3184,14 +4022,22 @@ fn stream_response( // caps (the SSE framing is not counted), and the raw frame bytes it // bounds too. let mut held_content = crate::held_content::HeldBuffer::default(); - // The semantic delta channel stays contiguous across frames. Raw - // supplementary values are scanned separately so metadata cannot - // break a literal split over two output deltas. - let mut continuation_bufs: Vec = Vec::new(); - let mut supplemental_buf = String::new(); - // Overlap carried between Window scans, one per channel. - let mut continuation_tails: Vec = Vec::new(); - let mut supplemental_tail = String::new(); + // 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 BufferFull to live forwarding after a fail-open cap hit. let mut fail_opened = false; let mut blocked = false; @@ -3236,7 +4082,7 @@ fn stream_response( let (parts, usage) = frame_parts(protocol, &frame); let held = parts.held(); let delta = parts.scan; - let guardrail_text = (!chain.is_empty()) + let guardrail_text = (!chain.is_empty() && !scan_budget_exhausted) .then(|| stream_guardrail_text(protocol, &frame, delta.clone())); if let Some(u) = usage { merge_usage(&mut telemetry.usage, u); @@ -3248,6 +4094,55 @@ fn stream_response( capture_cap, ); } + 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 { + // 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 = (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 => { @@ -3255,10 +4150,12 @@ fn stream_response( yield Ok::<_, std::convert::Infallible>(frame); } StreamOutputPolicy::EndOfStreamCheck => { - if let Some(text) = guardrail_text.as_ref() { + 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, ); } @@ -3266,10 +4163,12 @@ fn stream_response( yield Ok(frame); } StreamOutputPolicy::Window { size_chars, overlap_chars, .. } => { - if let Some(text) = guardrail_text.as_ref() { + 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, ); } @@ -3283,17 +4182,27 @@ fn stream_response( // the cap, mirroring BufferFull's self-bound. if continuation_bufs .iter() - .any(|continuation| continuation.chars().count() >= *size_chars) - || supplemental_buf.chars().count() >= *size_chars + .any(|continuation| continuation.text.chars().count() >= *size_chars) + || supplemental_buf + .iter() + .map(|value| value.chars().count()) + .sum::() + >= *size_chars || held_bytes > MAX_HELD_STREAM_BYTES { - let text = stream_guardrail_scan_text( + let candidates = stream_guardrail_scan_text( &continuation_tails, &continuation_bufs, - &supplemental_tail, &supplemental_buf, ); - match scan_output(&chain, &route_name, &text, &mut telemetry).await { + match scan_output_candidates( + &chain, + &route_name, + &candidates, + &mut telemetry, + ) + .await + { GuardrailVerdict::Block { reason, guardrail_name, @@ -3316,41 +4225,48 @@ fn stream_response( } pending_held.clear(); held_bytes = 0; - if continuation_tails.len() < continuation_bufs.len() { - continuation_tails.resize( - continuation_bufs.len(), - String::new(), - ); - } - for (index, continuation) in continuation_bufs.iter_mut().enumerate() { - let combined = format!( - "{}{}", - continuation_tails[index], - continuation.as_str(), - ); - continuation_tails[index] = - tail_chars(&combined, *overlap_chars); - continuation.clear(); + 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(); } - let combined_supplemental = if supplemental_tail.is_empty() { - supplemental_buf.clone() - } else if supplemental_buf.is_empty() { - supplemental_tail.clone() - } else { - format!("{supplemental_tail}\n{supplemental_buf}") - }; - supplemental_tail = - tail_chars(&combined_supplemental, *overlap_chars); 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() { + 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, ); } @@ -3393,7 +4309,7 @@ fn stream_response( let (parts, usage) = frame_parts(protocol, &rest); let held = parts.held(); let delta = parts.scan; - let guardrail_text = (!chain.is_empty()) + let guardrail_text = (!chain.is_empty() && !scan_budget_exhausted) .then(|| stream_guardrail_text(protocol, &rest, delta.clone())); if let Some(u) = usage { merge_usage(&mut telemetry.usage, u); @@ -3405,10 +4321,62 @@ fn stream_response( capture_cap, ); } - if let Some(text) = guardrail_text.as_ref() { + 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 { + // 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 = (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, ); } @@ -3460,36 +4428,47 @@ fn stream_response( } } } - let text = stream_guardrail_scan_text( - &continuation_tails, - &continuation_bufs, - &supplemental_tail, - &supplemental_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 + 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, + ) { - 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; + chain.record_unevaluable_output_bypass(crate::error::TAG_UNSCANNABLE_BODY); + } + } + for candidates in sealed_guardrail_epochs { + if !chain.is_empty() { + 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(..) { @@ -3552,6 +4531,27 @@ async fn scan_output( 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(); @@ -3969,6 +4969,12 @@ mod tests { 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(), @@ -5378,9 +6384,11 @@ mod tests { .0 .scan; let media = stream_guardrail_text(PassthroughProtocol::OpenaiResponses, media_frame, typed); - let scanned = - stream_guardrail_scan_text(&[], &media.continuations, "", &media.supplemental); - assert!(!scanned.contains("STREAM_MEDIA_SENTINEL"), "{scanned:?}"); + let scanned = stream_guardrail_scan_text(&[], &media.continuations, &media.supplemental); + assert!( + !scan_candidates_contain(&scanned, "STREAM_MEDIA_SENTINEL"), + "{scanned:?}" + ); // The terminal response repeats prior delta content, including media // from image-generation output. It is never a second scan carrier. @@ -5405,9 +6413,9 @@ mod tests { conflicting_delta, typed, ); - let scanned = stream_guardrail_scan_text(&[], &text.continuations, "", &text.supplemental); + let scanned = stream_guardrail_scan_text(&[], &text.continuations, &text.supplemental); assert!( - !scanned.contains("CONFLICTING_MEDIA_SENTINEL"), + !scan_candidates_contain(&scanned, "CONFLICTING_MEDIA_SENTINEL"), "{scanned:?}" ); @@ -5496,9 +6504,9 @@ mod tests { } #[test] - fn stream_guardrail_text_keeps_visible_deltas_contiguous() { - let first = b"data: {\"type\":\"response.output_text.delta\",\"item_id\":\"one\",\"delta\":\"FOR\"}\n\n"; - let second = b"data: {\"type\":\"response.output_text.delta\",\"item_id\":\"two\",\"delta\":\"noise\",\"delta\":\"BIDDEN\"}\n\n"; + 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( @@ -5512,22 +6520,573 @@ mod tests { second_parts.scan, ); let mut continuations = Vec::new(); - let mut supplemental = String::new(); - append_stream_guardrail_text(&mut continuations, &mut supplemental, &first); - append_stream_guardrail_text(&mut continuations, &mut supplemental, &second); - let scanned = stream_guardrail_scan_text(&[], &continuations, "", &supplemental); - assert!(scanned.starts_with("FORBIDDEN"), "{scanned:?}"); + 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:?}" + ); + assert!(scan_candidates_contain(&scanned, "okBIDDEN"), "{scanned:?}"); + } + + #[test] + 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, + ); + 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:?}" + ); + } + + #[test] + 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 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!(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" + ); + 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, + 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, + closed_prefixes: Vec::new(), + }; + assert!(stream_continuation_would_exceed_cap( + &[], + &[], + &[], + &incoming_supplemental, + )); + } + + #[test] + 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(), + )); + } + + #[test] + 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, + )); + } + + #[test] + 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 + ); + } + + #[test] + 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 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(), + ); + 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(), + ); + 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()]); + + 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 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 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, + ); + let opaque = stream_guardrail_text( + PassthroughProtocol::Raw, + opaque_frame, + frame_parts(PassthroughProtocol::Raw, opaque_frame).0.scan, + ); + 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:?}" + ); + } + + #[test] + 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 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 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\":[{\"delta\":{\"content\":\"ask carol@example.com\"}}]}\n\n"; + 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); + let scanned = stream_guardrail_scan_text(&[], &text.continuations, &text.supplemental); assert_eq!( - scanned.matches(email).count(), + scanned + .iter() + .map(|candidate| candidate.matches(email).count()) + .sum::(), 1, "typed, raw source, and supplemental channels must not multiply one carrier: {scanned:?}" ); @@ -5539,16 +7098,16 @@ mod tests { 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!(!scanned.contains("BLOCKME"), "{scanned:?}"); + 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!(!scanned.contains("BLOCKME"), "{scanned:?}"); + let scanned = stream_guardrail_scan_text(&[], &text.continuations, &text.supplemental); + assert!(!scan_candidates_contain(&scanned, "BLOCKME"), "{scanned:?}"); } #[test] 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-scan-coverage-e2e.test.ts b/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts index 00cd48efe..2f02e888c 100644 --- a/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts @@ -5,8 +5,11 @@ import { ProxyClient, SeedClient, spawnApp, + startMockSls, startOpenAiUpstream, waitConfigPropagation, + waitForSlsLog, + type MockSls, type OpenAiUpstream, type SpawnedApp, } from "../harness/index.js"; @@ -36,6 +39,12 @@ 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 deepEscapedBlockJSON = (depth: number) => `${'{"v":'.repeat(depth)}"${String.raw`\u0042LOCKME`}"${'}'.repeat(depth)}`; // Above serde_json's default recursion limit. It remains valid JSON and the @@ -43,20 +52,27 @@ const deepEscapedBlockJSON = (depth: number) => const DEEP_ESCAPED_BLOCK_JSON = deepEscapedBlockJSON(160); const CAP = 1_000; 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 KNOWN_CHAT_STREAM = `${String.raw`data: {"choices":[{"delta":{"content":"\u0042LOCKME","metadata":{"note":"${OUT_LIT}"}}}],"choices":[{"delta":{"content":"clean"}}]}`}\n\n`; -const KNOWN_RESPONSES_STREAM = `${String.raw`data: {"type":"response.output_text.delta","delta":"\u0042LOCKME","delta":"clean","metadata":{"note":"${OUT_LIT}"}}`}\n\n`; +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":"first","delta":"FOR"}\n\n`, - `data: {"type":"response.output_text.delta","item_id":"second","delta":"noise","delta":"BIDDEN"}\n\n`, + `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","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", part: { type: "reasoning_text", text: OUT_LIT } })}\n\n`, - `data: ${JSON.stringify({ type: "response.output_item.done", item: { type: "reasoning", summary: [{ type: "summary_text", text: OUT_LIT }] } })}\n\n`, - `data: ${JSON.stringify({ type: "response.output_text.delta", delta: "clean" })}\n\n`, + `data: ${JSON.stringify({ type: "response.output_item.done", 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: 0, 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", ]; @@ -133,7 +149,7 @@ describe("passthrough guardrail scan coverage", () => { rawContentType: "application/json", }); upstreams["raw-stream"] = await startOpenAiUpstream({ - rawStreamFrames: [`data: ${ESCAPED_BLOCK_JSON}\n\n`, "data: [DONE]\n\n"], + rawStreamFrames: [`data: ${RAW_BLOCK_SSE}\n\n`, "data: [DONE]\n\n"], }); upstreams["raw-deep-output"] = await startOpenAiUpstream({ rawBody: DEEP_ESCAPED_BLOCK_JSON, @@ -147,7 +163,7 @@ describe("passthrough guardrail scan coverage", () => { rawContentType: "application/json", }); upstreams["raw-safe-stream"] = await startOpenAiUpstream({ - rawStreamFrames: [`data: ${SAFE_ESCAPED_JSON}\n\n`], + rawStreamFrames: [`data: ${RAW_SAFE_SSE}\n\n`, "data: [DONE]\n\n"], }); upstreams["known-chat-buffered-output"] = await startOpenAiUpstream({ rawBody: KNOWN_CHAT_OUTPUT, @@ -166,6 +182,9 @@ describe("passthrough guardrail scan coverage", () => { 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, }); @@ -190,7 +209,7 @@ describe("passthrough guardrail scan coverage", () => { patterns: [ { kind: "literal", value: OUT_LIT }, { kind: "literal", value: ESCAPED_BLOCK }, - { kind: "literal", value: SPLIT_BLOCK }, + { kind: "regex", value: SPLIT_BLOCK_REGEX }, ], }); await seed.createGuardrail({ @@ -332,7 +351,11 @@ describe("passthrough guardrail scan coverage", () => { 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 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"); @@ -341,6 +364,25 @@ describe("passthrough guardrail scan coverage", () => { 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"; @@ -534,12 +576,12 @@ describe("passthrough guardrail scan coverage", () => { expect(upstreams["raw-stream"]!.receivedRequests.length).toBe(before + 1); }); - test("output: safe raw SSE JSON keeps its original bytes downstream", async (ctx) => { + 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: ${SAFE_ESCAPED_JSON}\n\n`); + expect(await res.text()).toBe(`data: ${RAW_SAFE_SSE}\n\ndata: [DONE]\n\n`); expect(upstreams["raw-safe-stream"]!.receivedRequests.length).toBe(before + 1); }); @@ -556,7 +598,11 @@ describe("passthrough guardrail scan coverage", () => { content_mode: "full", content_max_bytes: 4_096, }); - const requestUntilCaptured = async (route: string, expectedBody: string) => { + 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 @@ -574,7 +620,7 @@ describe("passthrough guardrail scan coverage", () => { const span = otlp.spans.find( (candidate) => candidate.attributes["aisix.passthrough.route_name"] === `pt-scan-${route}` && - candidate.attributes["gen_ai.completion"] === SAFE_ESCAPED_JSON, + candidate.attributes["gen_ai.completion"] === expectedCompletion, ); if (span) return span; await new Promise((resolve) => setTimeout(resolve, 50)); @@ -582,20 +628,229 @@ describe("passthrough guardrail scan coverage", () => { throw new Error(`no raw-source OTLP completion for ${route}; last response ${last}`); }; - await requestUntilCaptured("raw-safe-output", SAFE_ESCAPED_JSON); - await requestUntilCaptured("raw-safe-stream", `data: ${SAFE_ESCAPED_JSON}\n\n`); + 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 SSE JSON escapes are decoded before scanning", async (ctx) => { + 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("content_filter"); + 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); }); }); + +// This is intentionally a separate DP: the main suite has an env-scoped +// blocking row, so it cannot demonstrate the live, monitor-only policy of +// an output `fail_open: true` chain. +describe("passthrough Raw stream 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 logstore = "pt-scan-fail-open"; + const credentialRef = "pt_scan_open"; + let app: SpawnedApp | undefined; + let upstream: 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", + ], + }); + 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.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 sls?.close(); + }); + + 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); + }); +}); From 7908693c5fb08627e78c28c502e746c367dd6f76 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Thu, 1 Oct 2026 07:25:54 +0800 Subject: [PATCH 18/37] fix: enforce passthrough stream hold limits --- crates/aisix-proxy/src/held_content.rs | 4 +- crates/aisix-proxy/src/passthrough_route.rs | 191 ++++++---- .../guardrail-window-stream-hold-e2e.test.ts | 327 +++++++++++++++++- 3 files changed, 442 insertions(+), 80 deletions(-) diff --git a/crates/aisix-proxy/src/held_content.rs b/crates/aisix-proxy/src/held_content.rs index 1b9ae0286..dca28d461 100644 --- a/crates/aisix-proxy/src/held_content.rs +++ b/crates/aisix-proxy/src/held_content.rs @@ -1,7 +1,7 @@ //! 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 +//! 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, 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 diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index 1ba33865f..6c3bc33b3 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -1607,7 +1607,7 @@ fn raw_top_level_unique_object( 0 => Ok(None), 1 => { let value = values.pop().expect("one value"); - raw_is_object(&value).then_some(value).ok_or(()) + raw_is_object(&value).then_some(value).map(Some).ok_or(()) } _ => Err(()), } @@ -1627,6 +1627,7 @@ fn raw_top_level_unique_array( .trim_start() .starts_with('[') .then_some(value) + .map(Some) .ok_or(()) } _ => Err(()), @@ -2479,14 +2480,9 @@ fn path_is_within_base(candidate: &str, base: &str) -> bool { // Streaming relay // --------------------------------------------------------------------------- -/// 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. +/// 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 @@ -2512,9 +2508,11 @@ struct SseFrameSplitter(aisix_gateway::sse::SseFrameSplitter); impl SseFrameSplitter { fn new() -> Self { - Self(aisix_gateway::sse::SseFrameSplitter::new( - MAX_HELD_STREAM_BYTES, - )) + Self::with_max_frame_bytes(MAX_HELD_STREAM_BYTES) + } + + 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> { @@ -4005,6 +4003,10 @@ fn stream_response( }; 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 stream = async_stream::stream! { // The rate limiter's reservation becomes an owned hold at the handoff @@ -4012,13 +4014,11 @@ fn stream_response( // client cancels it, rather than when the response headers are built. let _stream_hold = stream_hold; let mut upstream = upstream_resp.bytes_stream(); - let mut splitter = SseFrameSplitter::new(); + 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(); - // 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` + // 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(); @@ -4038,7 +4038,7 @@ fn stream_response( let mut sealed_guardrail_epochs: Vec> = Vec::new(); let mut queued_guardrail_candidates = 0; let mut scan_budget_exhausted = false; - // Degrades BufferFull to live forwarding after a fail-open cap hit. + // 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 @@ -4082,7 +4082,7 @@ fn stream_response( let (parts, usage) = frame_parts(protocol, &frame); let held = parts.held(); let delta = parts.scan; - let guardrail_text = (!chain.is_empty() && !scan_budget_exhausted) + let guardrail_text = (!chain.is_empty() && !scan_budget_exhausted && !fail_opened) .then(|| stream_guardrail_text(protocol, &frame, delta.clone())); if let Some(u) = usage { merge_usage(&mut telemetry.usage, u); @@ -4162,7 +4162,12 @@ fn stream_response( telemetry.mark_first_delivery(); yield Ok(frame); } - StreamOutputPolicy::Window { size_chars, overlap_chars, .. } => { + 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, @@ -4172,15 +4177,32 @@ fn stream_response( text, ); } - held_bytes += frame.len(); + held_content.hold(held, 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 continuation_bufs + 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 @@ -4188,7 +4210,6 @@ fn stream_response( .map(|value| value.chars().count()) .sum::() >= *size_chars - || held_bytes > MAX_HELD_STREAM_BYTES { let candidates = stream_guardrail_scan_text( &continuation_tails, @@ -4224,7 +4245,7 @@ fn stream_response( yield Ok(f); } pending_held.clear(); - held_bytes = 0; + held_content = crate::held_content::HeldBuffer::default(); for continuation in &mut continuation_bufs { let tail = continuation_tails .iter() @@ -4275,13 +4296,14 @@ fn stream_response( 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(); - fail_opened = true; - chain.record_bypass(crate::error::TAG_OUTPUT_BUFFER_EXCEEDED); + held_content = crate::held_content::HeldBuffer::default(); } else { tracing::warn!( route = %route_name, @@ -4309,7 +4331,7 @@ fn stream_response( let (parts, usage) = frame_parts(protocol, &rest); let held = parts.held(); let delta = parts.scan; - let guardrail_text = (!chain.is_empty() && !scan_budget_exhausted) + let guardrail_text = (!chain.is_empty() && !scan_budget_exhausted && !fail_opened) .then(|| stream_guardrail_text(protocol, &rest, delta.clone())); if let Some(u) = usage { merge_usage(&mut telemetry.usage, u); @@ -4383,9 +4405,15 @@ fn stream_response( 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 => - { + 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) @@ -4409,12 +4437,13 @@ fn stream_response( 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(); - chain.record_bypass(crate::error::TAG_OUTPUT_BUFFER_EXCEEDED); telemetry.mark_first_delivery(); yield Ok(rest); } @@ -4428,46 +4457,48 @@ fn stream_response( } } } - 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 !chain.is_empty() { - if let GuardrailVerdict::Block { - reason, - guardrail_name, - unavailable, - } = scan_output_candidates(&chain, &route_name, &candidates, &mut telemetry).await + if !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, + ) { - 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; + chain.record_unevaluable_output_bypass(crate::error::TAG_UNSCANNABLE_BODY); + } + } + for candidates in sealed_guardrail_epochs { + if !chain.is_empty() { + 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; + } } } } @@ -6919,7 +6950,10 @@ mod tests { 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)] { + for (first, second) in [ + (&indexed[..], &missing_indexes[..]), + (&missing_indexes[..], &indexed[..]), + ] { let first = stream_guardrail_text( PassthroughProtocol::OpenaiResponses, first, @@ -7626,6 +7660,13 @@ 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); + assert_eq!(s.push(b"12345"), vec![b"12345".to_vec()]); + assert!(s.take_rest().is_empty()); + } + #[test] fn push_capped_respects_byte_cap_on_char_boundaries() { let mut buf = String::new(); 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" }, + ]); + }); }); From 924687eb27be890673122503610e2e27f62913d7 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Thu, 1 Oct 2026 08:35:11 +0800 Subject: [PATCH 19/37] fix: harden passthrough stream concurrency --- .../src/models/passthrough_route.rs | 9 +- crates/aisix-proxy/src/held_content.rs | 15 + crates/aisix-proxy/src/passthrough_route.rs | 343 ++++++++++++++---- crates/aisix-ratelimit/src/limiter.rs | 95 +++-- .../tests/redis_integration.rs | 51 +++ ...rdrail-buffer-cap-enforced-hit-e2e.test.ts | 91 ++++- .../src/cases/passthrough-route-e2e.test.ts | 89 +++++ 7 files changed, 583 insertions(+), 110 deletions(-) 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-proxy/src/held_content.rs b/crates/aisix-proxy/src/held_content.rs index dca28d461..487bf273b 100644 --- a/crates/aisix-proxy/src/held_content.rs +++ b/crates/aisix-proxy/src/held_content.rs @@ -44,6 +44,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 { diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index 6c3bc33b3..2596e1f3b 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -819,6 +819,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, @@ -920,6 +929,7 @@ async fn dispatch( telemetry, &client.request_id, reservation.into_stream_hold(), + stream_read_timeout, )); } @@ -2506,26 +2516,36 @@ const MAX_STREAM_GUARDRAIL_SOURCE_ID_BYTES: usize = 256; /// them at end-of-stream. struct SseFrameSplitter(aisix_gateway::sse::SseFrameSplitter); -impl SseFrameSplitter { - fn new() -> Self { - Self::with_max_frame_bytes(MAX_HELD_STREAM_BYTES) - } +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> { + 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(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(self.0.take_rest()); + frames.push(SseFrame { + bytes: self.0.take_rest(), + overflowed: true, + }); break; } } @@ -2963,13 +2983,17 @@ fn responses_source_continuations(payload: &[u8]) -> SourceContinuations { 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, - family: String, - identity: String, - identity_is_ambiguous: bool, + source: SourceBranchIdentity, body: &[u8], key: &str, ) -> Result<(), ()> { @@ -2977,9 +3001,9 @@ fn append_raw_string_carrier( out, keys, source_values, - family, - identity, - identity_is_ambiguous, + source.family, + source.identity, + source.identity_is_ambiguous, raw_top_level_string_values(body, key).ok_or(())?, ) } @@ -3018,9 +3042,11 @@ fn anthropic_source_continuations( &mut out, &mut keys, &mut source_values, - family.clone(), - "text".to_owned(), - false, + SourceBranchIdentity { + family: family.clone(), + identity: "text".to_owned(), + identity_is_ambiguous: false, + }, delta_body, "text", ) @@ -3029,9 +3055,11 @@ fn anthropic_source_continuations( &mut out, &mut keys, &mut source_values, - family.clone(), - "partial_json".to_owned(), - false, + SourceBranchIdentity { + family: family.clone(), + identity: "partial_json".to_owned(), + identity_is_ambiguous: false, + }, delta_body, "partial_json", ) @@ -3047,9 +3075,11 @@ fn anthropic_source_continuations( &mut out, &mut keys, &mut source_values, - family.clone(), - "text".to_owned(), - false, + SourceBranchIdentity { + family: family.clone(), + identity: "text".to_owned(), + identity_is_ambiguous: false, + }, block_body, "text", ) @@ -3066,9 +3096,11 @@ fn anthropic_source_continuations( &mut out, &mut keys, &mut source_values, - family.clone(), - "input".to_owned(), - false, + SourceBranchIdentity { + family: family.clone(), + identity: "input".to_owned(), + identity_is_ambiguous: false, + }, block_body, "input", ); @@ -3087,9 +3119,7 @@ fn anthropic_source_continuations( } _ => return SourceContinuations::Absent, }; - if result.is_err() { - SourceContinuations::Unevaluable - } else if !source_values_match_expected(source_values, expected) { + if result.is_err() || !source_values_match_expected(source_values, expected) { SourceContinuations::Unevaluable } else if out.is_empty() { SourceContinuations::Absent @@ -3232,9 +3262,11 @@ fn chat_choice_source_continuations(payload: &[u8]) -> SourceContinuations { &mut out, &mut keys, &mut source_values, - format!("chat:{choice_index}:tool:{tool_index}"), - format!("{container}:{field}"), - false, + SourceBranchIdentity { + family: format!("chat:{choice_index}:tool:{tool_index}"), + identity: format!("{container}:{field}"), + identity_is_ambiguous: false, + }, nested.get().as_bytes(), field, ) @@ -3247,9 +3279,11 @@ fn chat_choice_source_continuations(payload: &[u8]) -> SourceContinuations { } } } - source_values_match_expected(source_values, expected) - .then_some(SourceContinuations::Ready(out)) - .unwrap_or(SourceContinuations::Unevaluable) + if source_values_match_expected(source_values, expected) { + SourceContinuations::Ready(out) + } else { + SourceContinuations::Unevaluable + } } fn completions_source_continuations(payload: &[u8]) -> SourceContinuations { @@ -3280,9 +3314,11 @@ fn completions_source_continuations(payload: &[u8]) -> SourceContinuations { &mut out, &mut keys, &mut source_values, - format!("completions:{choice_index}"), - "text".to_owned(), - false, + SourceBranchIdentity { + family: format!("completions:{choice_index}"), + identity: "text".to_owned(), + identity_is_ambiguous: false, + }, choice_body, "text", ) @@ -3291,9 +3327,11 @@ fn completions_source_continuations(payload: &[u8]) -> SourceContinuations { return SourceContinuations::Unevaluable; } } - source_values_match_expected(source_values, expected) - .then_some(SourceContinuations::Ready(out)) - .unwrap_or(SourceContinuations::Unevaluable) + if source_values_match_expected(source_values, expected) { + SourceContinuations::Ready(out) + } else { + SourceContinuations::Unevaluable + } } fn stream_source_continuations( @@ -3326,6 +3364,14 @@ fn stream_source_continuations( } 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), @@ -3571,16 +3617,16 @@ fn stream_guardrail_text( let unevaluable = source_unevaluable || terminal_unevaluable; StreamGuardrailText { continuations, - supplemental: (!unevaluable) - .then(|| { - frame_guardrail_supplemental_values( - protocol, - frame, - has_source_continuations, - has_typed_continuation, - ) - }) - .unwrap_or_default(), + supplemental: if unevaluable { + Vec::new() + } else { + frame_guardrail_supplemental_values( + protocol, + frame, + has_source_continuations, + has_typed_continuation, + ) + }, unevaluable, closed_prefixes, } @@ -3992,6 +4038,7 @@ fn stream_response( 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; @@ -4013,7 +4060,12 @@ fn stream_response( // 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 mut upstream = upstream_resp.bytes_stream(); + 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(); @@ -4073,6 +4125,41 @@ fn stream_response( .min(u32::MAX as u128) as u32; } for frame in splitter.push(&chunk) { + let overflowed = frame.overflowed; + let frame = frame.bytes; + if overflowed && !fail_opened { + if let Some((max_buffer_bytes, on_exceeded_fail_open)) = policy.hold_cap() { + // `overflowed` means the splitter crossed this same + // raw-byte bound before it found a frame terminator. + // Do not hand that partial payload to any decoder. + 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); } @@ -4082,8 +4169,6 @@ fn stream_response( let (parts, usage) = frame_parts(protocol, &frame); let held = parts.held(); let delta = parts.scan; - let guardrail_text = (!chain.is_empty() && !scan_budget_exhausted && !fail_opened) - .then(|| stream_guardrail_text(protocol, &frame, delta.clone())); if let Some(u) = usage { merge_usage(&mut telemetry.usage, u); } @@ -4094,6 +4179,8 @@ fn stream_response( capture_cap, ); } + let guardrail_text = (!chain.is_empty() && !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( @@ -4320,19 +4407,71 @@ fn stream_response( } } + 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 = parts.held(); let delta = parts.scan; - let guardrail_text = (!chain.is_empty() && !scan_budget_exhausted && !fail_opened) - .then(|| stream_guardrail_text(protocol, &rest, delta.clone())); if let Some(u) = usage { merge_usage(&mut telemetry.usage, u); } @@ -4343,7 +4482,55 @@ fn stream_response( capture_cap, ); } - let unevaluable_output = guardrail_text.as_ref().is_some_and(|text| { + + // 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 = + (!chain.is_empty() && !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, @@ -4456,6 +4643,10 @@ fn stream_response( yield Ok(rest); } } + } + } + } + } } if !fail_opened { if !scan_budget_exhausted { @@ -5593,16 +5784,18 @@ mod tests { #[test] fn sse_splitter_emits_complete_frames_and_keeps_partials() { - let mut s = SseFrameSplitter::new(); + 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], b"data: a\n\n"); + 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], b"data: partial\n\n"); + 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::new(); + 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"); @@ -5611,11 +5804,11 @@ mod tests { #[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()] - ); + 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); @@ -5868,8 +6061,11 @@ mod tests { Some(usage_dims(26, 4)), ); // …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()]); + 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); } /// A comment-only frame — the keepalive some relays emit while the @@ -7644,14 +7840,18 @@ mod tests { #[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, @@ -7663,7 +7863,10 @@ mod tests { #[test] fn sse_splitter_honors_a_route_specific_frame_cap() { let mut s = SseFrameSplitter::with_max_frame_bytes(4); - assert_eq!(s.push(b"12345"), vec![b"12345".to_vec()]); + 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()); } diff --git a/crates/aisix-ratelimit/src/limiter.rs b/crates/aisix-ratelimit/src/limiter.rs index 30cd9bf86..419ff18d4 100644 --- a/crates/aisix-ratelimit/src/limiter.rs +++ b/crates/aisix-ratelimit/src/limiter.rs @@ -106,11 +106,19 @@ impl Limiter { ) -> Result { let member = self.next_member(); self.store.acquire(key, limits, &member).await?; + let has_concurrency_slot = limits.concurrency.is_some(); + let refresh_task = spawn_lease_refresh( + Arc::clone(&self.store), + key.to_string(), + member.clone(), + has_concurrency_slot, + ); Ok(Reservation { store: Arc::clone(&self.store), key: key.to_string(), member, - has_concurrency_slot: limits.concurrency.is_some(), + has_concurrency_slot, + refresh_task, committed: false, }) } @@ -136,6 +144,34 @@ 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 only after the SSE handoff lets a +/// second replica prune that still-live request before the handoff happens. +/// 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, +) -> 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; + store.refresh_stream_lease(&key, &member).await; + } + }) + }) +} + impl Default for Limiter { fn default() -> Self { Self::new() @@ -149,6 +185,9 @@ pub struct Reservation { 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>, committed: bool, } @@ -162,9 +201,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; } @@ -172,6 +218,7 @@ impl Reservation { impl Drop for Reservation { fn drop(&mut self) { + self.stop_refresh(); if self.committed { return; } @@ -226,8 +273,7 @@ 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_interval: Option = None; - let mut refresh_holds = Vec::new(); + let mut refresh_tasks = Vec::new(); let holds = self .reservations .iter_mut() @@ -236,35 +282,22 @@ impl MultiReservation { // slot now; the returned guard owns release from here on. r.committed = true; let store = Arc::clone(&r.store); - if r.has_concurrency_slot { - if let Some(interval) = store.stream_lease_refresh_interval() { - refresh_interval = Some( - refresh_interval.map_or(interval, |current| current.min(interval)), - ); - refresh_holds.push((Arc::clone(&store), r.key.clone(), r.member.clone())); - } + 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, + ) + }) { + refresh_tasks.push(task); } (store, r.key.clone(), r.member.clone()) }) .collect(); - // Proxy handlers create streaming guards inside Tokio. Keep ordinary - // callers that do not have a runtime from panicking; they retain the - // backend's normal stale-lease recovery behavior instead. - let refresh_task = refresh_interval.and_then(|interval| { - tokio::runtime::Handle::try_current().ok().map(|handle| { - handle.spawn(async move { - loop { - tokio::time::sleep(interval).await; - for (store, key, member) in &refresh_holds { - store.refresh_stream_lease(key, member).await; - } - } - }) - }) - }); StreamConcurrencyGuard { holds, - refresh_task, + refresh_tasks, released: false, } } @@ -286,9 +319,9 @@ 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. It is stopped before the terminal release to prevent a late - /// renewal from racing stream teardown. - refresh_task: Option>, + /// stream. These tasks began at reservation acquisition and are stopped + /// before terminal release to prevent a late renewal from racing teardown. + refresh_tasks: Vec>, released: bool, } @@ -298,7 +331,7 @@ impl StreamConcurrencyGuard { return; } self.released = true; - if let Some(task) = self.refresh_task.take() { + for task in self.refresh_tasks.drain(..) { task.abort(); } for (store, key, member) in &self.holds { @@ -311,7 +344,7 @@ 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_task.is_some()) + .field("refreshing", &!self.refresh_tasks.is_empty()) .field("released", &self.released) .finish() } diff --git a/crates/aisix-ratelimit/tests/redis_integration.rs b/crates/aisix-ratelimit/tests/redis_integration.rs index ee574b63b..06b6b1562 100644 --- a/crates/aisix-ratelimit/tests/redis_integration.rs +++ b/crates/aisix-ratelimit/tests/redis_integration.rs @@ -534,6 +534,57 @@ async fn stream_hold_renews_redis_lease_until_drop() { assert!(acquired, "slot must free cluster-wide when the stream ends"); } +/// 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. 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..5415c743c 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,8 @@ 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"; // 30 pieces of 100 bytes: three times the tight cap, far under the loose one. const TIGHT_CAP = 1_000; @@ -121,14 +123,32 @@ 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), +]; + +// 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 +281,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 +312,32 @@ describe("a stream refused by the hold-back cap names the row whose cap it outgr scope_id: openTailRoute.id, priority: 100, }); + 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 +472,36 @@ 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 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/passthrough-route-e2e.test.ts b/tests/e2e/src/cases/passthrough-route-e2e.test.ts index 3fdb5abe3..400277960 100644 --- a/tests/e2e/src/cases/passthrough-route-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-route-e2e.test.ts @@ -50,6 +50,10 @@ 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"); describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed paths", () => { let app: SpawnedApp | undefined; @@ -778,6 +782,91 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat 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("envelope auto-detection: usage follows the request body, never the config", async (ctx) => { if (!etcdReachable || !app || !seed) { ctx.skip(); From 4c9f33960ec81521bef78778ad927088f07ead27 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Thu, 1 Oct 2026 08:45:34 +0800 Subject: [PATCH 20/37] chore: sync passthrough route schemas --- schemas/resources-lenient/passthrough_route.schema.json | 2 +- schemas/resources/passthrough_route.schema.json | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) 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": [ From 6e1fdefb5545c1420b760c48c89b89cb28f6ddc0 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Thu, 1 Oct 2026 09:58:44 +0800 Subject: [PATCH 21/37] fix: harden passthrough stream boundaries --- crates/aisix-proxy/src/passthrough_route.rs | 958 ++++++++++++------ ...rdrail-buffer-cap-enforced-hit-e2e.test.ts | 39 + ...ssthrough-chat-media-guardrail-e2e.test.ts | 429 ++++++++ .../src/cases/passthrough-route-e2e.test.ts | 134 ++- .../src/cases/ratelimit-cluster-e2e.test.ts | 147 +++ 5 files changed, 1381 insertions(+), 326 deletions(-) create mode 100644 tests/e2e/src/cases/passthrough-chat-media-guardrail-e2e.test.ts diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index 2596e1f3b..68cc4d74d 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -1409,15 +1409,6 @@ fn decoded_json_string_values_except_root_keys(body: &[u8], excluded: &[&str]) - .ok() } -fn decoded_json_string_values_except_root_keys_vec( - body: &[u8], - excluded: &[&str], -) -> Option> { - decoded_json_string_values_vec_where(body, |path| { - !excluded.iter().any(|key| is_root_key(path, key)) - }) -} - /// Source values of all occurrences of one top-level key. `RawValue` keeps /// repeated keys separate, unlike `serde_json::Value`. fn raw_top_level_values( @@ -1510,17 +1501,9 @@ fn raw_top_level_unique_type(body: &[u8]) -> Option { .then_some(first) } -fn raw_top_level_has_any_type(body: &[u8], wanted: &[&str]) -> bool { - raw_top_level_values(body, "type") - .into_iter() - .flatten() - .filter_map(|value| serde_json::from_str::(value.get()).ok()) - .any(|kind| wanted.iter().any(|wanted| kind == *wanted)) -} - /// `true` only when every source `type` value is one of `allowed`. This is -/// stricter than [`raw_top_level_has_any_type`]: an audio or image event must -/// not borrow a text event's carrier merely by repeating a conflicting type. +/// deliberately strict: an audio or image event must not borrow a text +/// event's carrier merely by repeating a conflicting type. fn raw_top_level_has_only_types(body: &[u8], allowed: &[&str]) -> bool { let Some(values) = raw_top_level_values(body, "type") else { return false; @@ -1543,25 +1526,6 @@ fn raw_top_level_items_have_only_types(body: &[u8], key: &str, allowed: &[&str]) ) } -fn append_raw_array_item_strings( - out: &mut String, - body: &[u8], - key: &str, - mut skip: impl FnMut(&serde_json::value::RawValue) -> bool, -) -> Option<()> { - for array in raw_top_level_values(body, key)? { - for item in raw_array_items(&array)? { - if !skip(&item) { - append_scan_text( - out, - &decoded_json_string_values_including_empty(item.get().as_bytes())?, - ); - } - } - } - Some(()) -} - fn append_raw_string_value(out: &mut String, raw: &serde_json::value::RawValue) -> Option<()> { append_scan_text(out, &serde_json::from_str::(raw.get()).ok()?); Some(()) @@ -1849,38 +1813,178 @@ fn request_guardrail_text(protocol: PassthroughProtocol, body: &[u8]) -> String } } -fn is_hidden_chat_reasoning_path(path: &[crate::json_splice::PathSeg]) -> bool { - use crate::json_splice::PathSeg; +/// 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, + } +} - let path = match path { - [PathSeg::Key(choices), PathSeg::Index(_), rest @ ..] if choices == "choices" => rest, - _ => path, - }; - matches!( - path, - [ - PathSeg::Key(message_or_delta), - PathSeg::Key(reasoning), - .. - ] if matches!(message_or_delta.as_str(), "message" | "delta") - && matches!(reasoning.as_str(), "reasoning_content" | "reasoning") - ) +/// 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: &serde_json::value::RawValue, +) -> 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 part_body = part.get().as_bytes(); + if let Some(field) = raw_top_level_unique_type(part_body) + .as_deref() + .and_then(chat_visible_content_part_field) + { + append_raw_top_level_strings(out, part_body, field)?; + } + } + 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).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: &serde_json::value::RawValue, +) -> 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)?; + } + } + } + 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: &serde_json::value::RawValue, +) -> 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: &serde_json::value::RawValue, +) -> 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: &serde_json::value::RawValue, +) -> 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).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())?, + ); + } + } + Some(_) | None => {} + } + } + Some(()) } fn decoded_chat_response_string_values(body: &[u8]) -> Option { - let mut out = - decoded_json_string_values_except_root_keys(body, &["model", "content", "choices"])?; - append_raw_array_item_strings(&mut out, body, "content", |item| { - raw_object_has_only_types(item, &["thinking", "redacted_thinking"]) - })?; - for array in raw_top_level_values(body, "choices")? { - for item in raw_array_items(&array)? { - let text = - crate::json_splice::collect_string_values_where(item.get().as_bytes(), |path| { - !is_hidden_chat_reasoning_path(path) - }) - .ok()?; - append_scan_text(&mut out, &text); + 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 Some(choices) = raw_array_items(&array) else { + continue; + }; + for choice in choices { + if !raw_is_object(&choice) { + continue; + } + for message in raw_top_level_values(choice.get().as_bytes(), "message")? { + append_chat_output_message_strings(&mut out, &message)?; + } + } + } + if !has_chat_choices && raw_top_level_unique_type(body).as_deref() == Some("message") { + for content in raw_top_level_values(body, "content")? { + append_anthropic_output_content_strings(&mut out, &content)?; } } Some(out) @@ -1995,7 +2099,10 @@ fn response_guardrail_text(protocol: PassthroughProtocol, body: &[u8]) -> String match protocol { PassthroughProtocol::Raw => decoded_json_string_values(body).unwrap_or_else(raw), PassthroughProtocol::OpenaiChat => { - decoded_chat_response_string_values(body).unwrap_or_else(raw) + // 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. + decoded_chat_response_string_values(body).unwrap_or_default() } PassthroughProtocol::OpenaiCompletions => { decoded_non_model_json_string_values(body).unwrap_or_else(raw) @@ -2636,10 +2743,6 @@ fn frame_delta(protocol: PassthroughProtocol, frame: &[u8]) -> (String, Option

bool { - is_hidden_chat_reasoning_path(path) -} - /// 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. @@ -2661,31 +2764,14 @@ fn hidden_chat_stream_reasoning_frame(body: &[u8]) -> Option { #[cfg(test)] fn decoded_chat_frame_string_values(body: &[u8]) -> Option { - match raw_top_level_unique_type(body).as_deref() { - Some("content_block_delta") - if raw_top_level_items_have_only_types( - body, - "delta", - &["thinking_delta", "signature_delta"], - )? => - { - return decoded_json_string_values_except_root_keys(body, &["model", "delta"]); - } - Some("content_block_start") - if raw_top_level_items_have_only_types( - body, - "content_block", - &["thinking", "redacted_thinking"], - )? => - { - return decoded_json_string_values_except_root_keys(body, &["model", "content_block"]); - } - _ => {} + let mut out = String::new(); + for value in decoded_chat_frame_continuations(body)? { + append_scan_text(&mut out, &value); } - crate::json_splice::collect_string_values_where(body, |path| { - !is_root_key(path, "model") && !is_hidden_chat_stream_reasoning_path(path) - }) - .ok() + for value in decoded_chat_frame_supplemental_values(body)? { + append_scan_text(&mut out, &value); + } + Some(out) } /// The typed stream extractor deliberately takes only text/tool delta events: @@ -2702,60 +2788,110 @@ fn decoded_responses_frame_string_values(body: &[u8]) -> Option { } fn decoded_chat_frame_continuations(body: &[u8]) -> Option> { - use crate::json_splice::PathSeg; - - if hidden_chat_stream_reasoning_frame(body) == Some(true) { - return None; - } - if raw_top_level_has_any_type(body, &["content_block_delta"]) { - return decoded_json_string_values_vec_where(body, |path| { - matches!( - path, - [PathSeg::Key(delta), PathSeg::Key(field)] - if delta == "delta" && matches!(field.as_str(), "text" | "partial_json") - ) - }); - } - if raw_top_level_has_any_type(body, &["content_block_start"]) { - let mut out = Vec::new(); - for block in raw_top_level_values(body, "content_block")? { - out.extend(raw_top_level_string_values(block.get().as_bytes(), "text")?); - for input in raw_top_level_values(block.get().as_bytes(), "input")? { - out.extend( - crate::json_splice::collect_string_values_where_vec( - input.get().as_bytes(), - |_| true, - ) - .ok()?, - ); + match raw_top_level_unique_type(body).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()).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).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()), } } - return (!out.is_empty()).then_some(out); + _ => { + 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) + .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) + } } - decoded_json_string_values_vec_where(body, |path| { - matches!( - path, - [PathSeg::Key(choices), PathSeg::Index(_), PathSeg::Key(delta), PathSeg::Key(field)] - if choices == "choices" - && delta == "delta" - && field == "content" - ) || matches!( - path, - [PathSeg::Key(choices), PathSeg::Index(_), PathSeg::Key(delta), PathSeg::Key(content), PathSeg::Index(_), PathSeg::Key(text)] - if choices == "choices" - && delta == "delta" - && content == "content" - && text == "text" - ) || matches!( - path, - [PathSeg::Key(choices), PathSeg::Index(_), PathSeg::Key(delta), PathSeg::Key(tool_calls), PathSeg::Index(_), PathSeg::Key(kind), PathSeg::Key(value)] - if choices == "choices" - && delta == "delta" - && tool_calls == "tool_calls" - && ((kind == "function" && value == "arguments") - || (kind == "custom" && value == "input")) - ) - }) } fn decoded_completions_frame_continuations(body: &[u8]) -> Option> { @@ -2770,55 +2906,6 @@ fn decoded_completions_frame_continuations(body: &[u8]) -> Option> { }) } -fn is_chat_choice_continuation_path(path: &[crate::json_splice::PathSeg]) -> bool { - use crate::json_splice::PathSeg; - - matches!( - path, - [PathSeg::Key(choices), PathSeg::Index(_), PathSeg::Key(delta), PathSeg::Key(field)] - if choices == "choices" && delta == "delta" && field == "content" - ) || matches!( - path, - [PathSeg::Key(choices), PathSeg::Index(_), PathSeg::Key(delta), PathSeg::Key(content), PathSeg::Index(_), PathSeg::Key(text)] - if choices == "choices" - && delta == "delta" - && content == "content" - && text == "text" - ) || matches!( - path, - [PathSeg::Key(choices), PathSeg::Index(_), PathSeg::Key(delta), PathSeg::Key(tool_calls), PathSeg::Index(_), PathSeg::Key(kind), PathSeg::Key(value)] - if choices == "choices" - && delta == "delta" - && tool_calls == "tool_calls" - && ((kind == "function" && value == "arguments") - || (kind == "custom" && value == "input")) - ) -} - -fn is_anthropic_delta_continuation_path(path: &[crate::json_splice::PathSeg]) -> bool { - use crate::json_splice::PathSeg; - - matches!( - path, - [PathSeg::Key(delta), PathSeg::Key(field)] - if delta == "delta" && matches!(field.as_str(), "text" | "partial_json") - ) -} - -fn is_anthropic_start_continuation_path(path: &[crate::json_splice::PathSeg]) -> bool { - use crate::json_splice::PathSeg; - - matches!( - path, - [PathSeg::Key(content_block), PathSeg::Key(field)] - if content_block == "content_block" && field == "text" - ) || matches!( - path, - [PathSeg::Key(content_block), PathSeg::Key(input), ..] - if content_block == "content_block" && input == "input" - ) -} - fn is_completions_continuation_path(path: &[crate::json_splice::PathSeg]) -> bool { use crate::json_splice::PathSeg; @@ -2834,20 +2921,59 @@ fn is_completions_continuation_path(path: &[crate::json_splice::PathSeg]) -> boo /// 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> { - let anthropic_delta = raw_top_level_has_any_type(body, &["content_block_delta"]); - let anthropic_start = - !anthropic_delta && raw_top_level_has_any_type(body, &["content_block_start"]); - decoded_json_string_values_vec_where(body, |path| { - !is_root_key(path, "model") - && !is_hidden_chat_stream_reasoning_path(path) - && !(if anthropic_delta { - is_anthropic_delta_continuation_path(path) - } else if anthropic_start { - is_anthropic_start_continuation_path(path) - } else { - is_chat_choice_continuation_path(path) - }) - }) + if raw_top_level_unique_type(body).as_deref() == Some("content_block_delta") { + return Some(Vec::new()); + } + if raw_top_level_unique_type(body).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()).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) } fn decoded_completions_frame_supplemental_values(body: &[u8]) -> Option> { @@ -3009,7 +3135,7 @@ fn append_raw_string_carrier( } fn raw_part_identity(body: &[u8]) -> Result { - // Content arrays can repeat a visible `text` field. Their numeric index + // 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. @@ -3038,20 +3164,20 @@ fn anthropic_source_continuations( _ => return SourceContinuations::Unevaluable, }; let delta_body = delta.get().as_bytes(); - 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", - ) - .and_then(|()| { - append_raw_string_carrier( + match raw_top_level_unique_type(delta_body).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, @@ -3062,8 +3188,9 @@ fn anthropic_source_continuations( }, delta_body, "partial_json", - ) - }) + ), + Some(_) | None => Ok(()), + } } "content_block_start" => { let block = match raw_top_level_unique_object(payload, "content_block") { @@ -3071,51 +3198,56 @@ fn anthropic_source_continuations( _ => return SourceContinuations::Unevaluable, }; let block_body = block.get().as_bytes(); - 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", - ) - .and_then(|()| { - let mut inputs = raw_top_level_values(block_body, "input").ok_or(())?; - let Some(input) = inputs.pop() else { - return Ok(()); - }; - if !inputs.is_empty() { - return Err(()); - } - if input.get().trim_start().starts_with('"') { - return 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", - ); - } - // 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 = crate::json_splice::collect_string_values(input.get().as_bytes()) - .map_err(|_| ())?; - if text.is_empty() { - Ok(()) - } else { - Err(()) + match raw_top_level_unique_type(block_body).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 = raw_top_level_values(block_body, "input").ok_or(())?; + 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 = + crate::json_splice::collect_string_values(input.get().as_bytes()) + .map_err(|_| ())?; + if text.is_empty() { + Ok(()) + } else { + Err(()) + } + } } - }) + Some(_) | None => Ok(()), + } } _ => return SourceContinuations::Absent, }; @@ -3203,15 +3335,21 @@ fn chat_choice_source_continuations(payload: &[u8]) -> SourceContinuations { if !raw_is_object(&part) { return SourceContinuations::Unevaluable; } - let text_values = - match raw_top_level_string_values(part.get().as_bytes(), "text") { - Some(text_values) => text_values, - None => return SourceContinuations::Unevaluable, - }; - if text_values.is_empty() { + let part_body = part.get().as_bytes(); + let Some(field) = raw_top_level_unique_type(part_body) + .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.get().as_bytes()) { + let part_id = match raw_part_identity(part_body) { Ok(part_id) if part_ids.insert(part_id.clone()) => part_id, _ => return SourceContinuations::Unevaluable, }; @@ -3220,9 +3358,9 @@ fn chat_choice_source_continuations(payload: &[u8]) -> SourceContinuations { &mut keys, &mut source_values, format!("chat:{choice_index}:content"), - format!("part:{part_id}:text"), + format!("part:{part_id}:{field}"), false, - text_values, + visible_values, ) .is_err() { @@ -3234,6 +3372,22 @@ fn chat_choice_source_continuations(payload: &[u8]) -> SourceContinuations { _ => 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, @@ -3252,7 +3406,10 @@ fn chat_choice_source_continuations(payload: &[u8]) -> SourceContinuations { Ok(Some(index)) if tool_indexes.insert(index) => index.to_string(), _ => return SourceContinuations::Unevaluable, }; - for (container, field) in [("function", "arguments"), ("custom", "input")] { + 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, @@ -3278,6 +3435,28 @@ fn chat_choice_source_continuations(payload: &[u8]) -> SourceContinuations { } } } + 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) @@ -3410,9 +3589,8 @@ fn frame_guardrail_supplemental_values( has_source_continuations: bool, has_typed_continuation: bool, ) -> Vec { - // On a malformed or unknown frame, typed extraction is the only - // available output carrier. It already contains the raw fallback, so a - // second generic scan would double-count it. + // 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 Vec::new(); } @@ -3443,33 +3621,7 @@ fn frame_guardrail_supplemental_values( } fn decoded_chat_frame_values(body: &[u8]) -> Option> { - match raw_top_level_unique_type(body).as_deref() { - Some("content_block_delta") - if raw_top_level_items_have_only_types( - body, - "delta", - &["thinking_delta", "signature_delta"], - )? => - { - return decoded_json_string_values_except_root_keys_vec(body, &["model", "delta"]); - } - Some("content_block_start") - if raw_top_level_items_have_only_types( - body, - "content_block", - &["thinking", "redacted_thinking"], - )? => - { - return decoded_json_string_values_except_root_keys_vec( - body, - &["model", "content_block"], - ); - } - _ => {} - } - decoded_json_string_values_vec_where(body, |path| { - !is_root_key(path, "model") && !is_hidden_chat_stream_reasoning_path(path) - }) + decoded_chat_frame_supplemental_values(body) } fn decoded_responses_frame_values(body: &[u8]) -> Option> { @@ -3493,7 +3645,7 @@ fn frame_guardrail_values(protocol: PassthroughProtocol, frame: &[u8]) -> Vec { - decoded_chat_frame_values(payload.as_bytes()).unwrap_or_else(raw) + decoded_chat_frame_values(payload.as_bytes()).unwrap_or_default() } PassthroughProtocol::OpenaiCompletions => { decoded_json_string_values_vec_where(payload.as_bytes(), |path| { @@ -3509,8 +3661,8 @@ fn frame_guardrail_values(protocol: PassthroughProtocol, frame: &[u8]) -> Vec String { let Some(payload) = crate::redact::frame_payload(frame) else { @@ -3526,7 +3678,10 @@ fn frame_guardrail_text(protocol: PassthroughProtocol, frame: &[u8]) -> String { decoded_json_string_values(payload.as_bytes()).unwrap_or_else(raw) } PassthroughProtocol::OpenaiChat => { - decoded_chat_frame_string_values(payload.as_bytes()).unwrap_or_else(raw) + // 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 => { decoded_non_model_json_string_values(payload.as_bytes()).unwrap_or_else(raw) @@ -3574,6 +3729,18 @@ fn stream_guardrail_text( || (matches!(protocol, PassthroughProtocol::OpenaiResponses) && !responses_visible_delta) { 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 }; @@ -4125,13 +4292,15 @@ fn stream_response( .min(u32::MAX as u128) as u32; } for frame in splitter.push(&chunk) { - let overflowed = frame.overflowed; + // Both completed and splitter-overflowed frames take the + // same raw-cap preflight below. + let _overflowed = frame.overflowed; let frame = frame.bytes; - if overflowed && !fail_opened { + if !fail_opened { if let Some((max_buffer_bytes, on_exceeded_fail_open)) = policy.hold_cap() { - // `overflowed` means the splitter crossed this same - // raw-byte bound before it found a frame terminator. - // Do not hand that partial payload to any decoder. + // 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; @@ -6534,9 +6703,10 @@ mod tests { 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", "NESTED", "clean"] { + 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"}]}]}"#; @@ -6569,6 +6739,145 @@ mod tests { ); } + #[test] + 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 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:?}"); + } + + 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 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"}]}"#; @@ -6687,12 +6996,13 @@ mod tests { } #[test] - fn known_sse_envelopes_scan_duplicate_and_nested_source_strings() { - let chat = b"data: {\"choices\":[{\"delta\":{\"content\":\"\\u0042LOCKME\",\"metadata\":{\"note\":\"NESTED\"}}}],\"choices\":[{\"delta\":{\"content\":\"clean\"}}]}\n\n"; + 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", "NESTED", "clean"] { + for expected in ["BLOCKME", "clean"] { assert!(scanned.contains(expected), "{scanned:?}"); } + assert!(!scanned.contains("NESTED"), "{scanned:?}"); 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); @@ -6720,7 +7030,7 @@ mod tests { 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) + !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"; 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 5415c743c..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 @@ -52,6 +52,7 @@ 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; @@ -140,6 +141,14 @@ const UNTERMINATED_TAIL_STREAM = [ 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. @@ -312,6 +321,20 @@ 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({ @@ -472,6 +495,22 @@ 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", { 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..dfa69bbd7 --- /dev/null +++ b/tests/e2e/src/cases/passthrough-chat-media-guardrail-e2e.test.ts @@ -0,0 +1,429 @@ +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"; + +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", +]; + +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 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 }); + 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.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 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 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("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); + }); +}); diff --git a/tests/e2e/src/cases/passthrough-route-e2e.test.ts b/tests/e2e/src/cases/passthrough-route-e2e.test.ts index 400277960..70216ae83 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, @@ -54,6 +56,57 @@ 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; @@ -285,8 +338,22 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat ); const baseline = upstream.receivedRequests.length; - // Use undici's raw request helper: fetch implementations are allowed to - // normalize URL escapes before the gateway receives the wire path. + // 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", + "..\\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", @@ -867,6 +934,69 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat 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/ratelimit-cluster-e2e.test.ts b/tests/e2e/src/cases/ratelimit-cluster-e2e.test.ts index e5f0a1524..f7ace415f 100644 --- a/tests/e2e/src/cases/ratelimit-cluster-e2e.test.ts +++ b/tests/e2e/src/cases/ratelimit-cluster-e2e.test.ts @@ -33,6 +33,16 @@ 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 ETCD_ENDPOINT = etcdEndpoint(); const REDIS_URL = process.env.AISIX_E2E_REDIS ?? "redis://127.0.0.1:6379"; @@ -205,6 +215,143 @@ 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 infraReady = false; + 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; + + 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 }, + ], + }); + 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-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, + }); + // 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. + await seed.createApiKey({ + key_hash: PASSTHROUGH_CALLER_KEY_HASH, + allowed_models: ["*"], + allowed_routes: [PASSTHROUGH_ROUTE], + rate_limit: { concurrency: 1 }, + }); + + 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(); + 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) { + 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); + }, + 15_000, + ); +}); + describe("rate limit is NOT shared with backend=memory (per-replica, the #798 bug)", () => { let appA: SpawnedApp | undefined; let appB: SpawnedApp | undefined; From 9ca594fa30ada9964ec2ad2a5096b8ac02f0f7ef Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Thu, 1 Oct 2026 11:01:19 +0800 Subject: [PATCH 22/37] fix: harden passthrough guardrail boundaries --- crates/aisix-proxy/src/held_content.rs | 48 +- crates/aisix-proxy/src/json_splice.rs | 75 ++- crates/aisix-proxy/src/passthrough_route.rs | 435 ++++++++++++++---- ...ssthrough-chat-media-guardrail-e2e.test.ts | 133 ++++++ .../passthrough-scan-coverage-e2e.test.ts | 159 ++++++- 5 files changed, 730 insertions(+), 120 deletions(-) diff --git a/crates/aisix-proxy/src/held_content.rs b/crates/aisix-proxy/src/held_content.rs index 487bf273b..dad389974 100644 --- a/crates/aisix-proxy/src/held_content.rs +++ b/crates/aisix-proxy/src/held_content.rs @@ -2,8 +2,8 @@ //! //! 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, reasoning, and -//! tool-call arguments. SSE and JSON framing — event names, ids, indexes, +//! 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. @@ -193,8 +193,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)] @@ -268,8 +269,8 @@ pub(crate) fn responses_event_parts(v: &Value) -> Parts { } /// 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 @@ -284,10 +285,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) @@ -297,6 +302,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")); } @@ -363,6 +369,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":{}}}); diff --git a/crates/aisix-proxy/src/json_splice.rs b/crates/aisix-proxy/src/json_splice.rs index 467cca7b8..39415f1c3 100644 --- a/crates/aisix-proxy/src/json_splice.rs +++ b/crates/aisix-proxy/src/json_splice.rs @@ -15,9 +15,10 @@ //! but they ARE decoded to build the path handed to the predicate. //! //! The scanner is iterative, so deeply nested JSON does not consume the -//! Rust call stack. It still fails safe: any unexpected byte or overrun -//! returns an error rather than a partially rewritten document. Callers -//! decide the failure policy (the MCP output hook fails closed). +//! 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; @@ -39,12 +40,32 @@ 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, +} + +#[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 + } +} + +/// 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. +const MAX_JSON_DEPTH: usize = 4_096; + /// Rewrite the string values of `input` selected by `should_rewrite`, /// leaving every other byte untouched. /// @@ -65,7 +86,10 @@ pub fn rewrite_string_values( Array, } - let err = |at: usize| SpliceError { at }; + let err = |at: usize| SpliceError { + at, + kind: SpliceErrorKind::Invalid, + }; let mut splices: Vec<(Range, String)> = Vec::new(); let mut path: Vec = Vec::new(); let mut frames: Vec = Vec::new(); @@ -87,11 +111,17 @@ pub fn rewrite_string_values( _ => i += 1, } } - Err(SpliceError { at: start }) + Err(SpliceError { + at: start, + kind: SpliceErrorKind::Invalid, + }) }; 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 @@ -102,6 +132,12 @@ pub fn rewrite_string_values( match b { b'{' => { frames.push(Frame::Object); + if frames.len() > MAX_JSON_DEPTH { + return Err(SpliceError { + at: pos, + kind: SpliceErrorKind::DepthExceeded, + }); + } pos += 1; skip_ws(&mut pos); match input.get(pos) { @@ -126,6 +162,12 @@ pub fn rewrite_string_values( } b'[' => { frames.push(Frame::Array); + if frames.len() > MAX_JSON_DEPTH { + return Err(SpliceError { + at: pos, + kind: SpliceErrorKind::DepthExceeded, + }); + } pos += 1; skip_ws(&mut pos); if input.get(pos) == Some(&b']') { @@ -233,9 +275,10 @@ pub fn rewrite_string_values( /// 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 keeps working beyond serde_json's default -/// container-recursion limit. Object keys are decoded only to maintain the -/// scanner's structure and are never included in the returned text. +/// 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) } @@ -243,8 +286,8 @@ pub fn collect_string_values(input: &[u8]) -> Result { /// Decode and collect selected JSON string **values** in source order. /// /// Like [`collect_string_values`], this preserves duplicate keys and stays -/// stack-safe for deeply nested documents. The predicate sees the decoded -/// path of each string value, never an object key. +/// 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, @@ -488,4 +531,12 @@ mod tests { .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 68cc4d74d..a496b7fa7 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -580,36 +580,73 @@ async fn dispatch( // 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 text = match try_request_guardrail_text(protocol, &body_bytes) { + Ok(text) => Some(text), + Err(err) if !err.is_depth_exceeded() => { + 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 scan passthrough-route request to its bounded depth; nothing attached both reads the request and fails 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 scan passthrough-route request to its bounded depth; 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, + )); + } } } @@ -959,42 +996,81 @@ 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(), + let text = match try_response_guardrail_text(protocol, &resp_body) { + Ok(text) => Some(text), + Err(err) if !err.is_depth_exceeded() => { + Some(response_guardrail_text(protocol, &resp_body)) + } + Err(err) + if !aisix_guardrails::Guardrail::refuses_unevaluable_output(&resolved_chain) => + { + tracing::debug!( + guardrail_hook = "output", + route = %route.name, + error = %err, + "cannot scan passthrough-route response to its bounded depth; nothing attached both reads the response and fails 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 scan passthrough-route response to its bounded depth; 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, + )); + } } } @@ -1398,15 +1474,39 @@ fn decoded_non_model_json_string_values(body: &[u8]) -> Option { decoded_json_string_values_where(body, |path| !is_root_key(path, "model")) } -fn decoded_json_string_values_including_empty(body: &[u8]) -> Option { - crate::json_splice::collect_string_values(body).ok() +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); + } + None + } + } } -fn decoded_json_string_values_except_root_keys(body: &[u8], excluded: &[&str]) -> Option { - crate::json_splice::collect_string_values_where(body, |path| { - !excluded.iter().any(|key| is_root_key(path, key)) - }) - .ok() +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; + } + } + } + Some(out) } /// Source values of all occurrences of one top-level key. `RawValue` keeps @@ -1449,6 +1549,46 @@ fn raw_top_level_values( Some(values) } +/// Source values of top-level keys other than `excluded`. Values are captured +/// as raw JSON before filtering so a known opaque carrier can be skipped +/// without recursively deserializing its payload. +fn raw_top_level_values_except( + body: &[u8], + excluded: &[&str], +) -> Option>> { + struct Values<'a> { + excluded: &'a [&'a str], + } + + impl<'de> serde::de::Visitor<'de> for Values<'_> { + type Value = Vec>; + + fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str("a JSON object") + } + + fn visit_map(self, mut map: A) -> Result + where + A: serde::de::MapAccess<'de>, + { + let mut values = Vec::new(); + while let Some(key) = map.next_key::()? { + let value = map.next_value::>()?; + if !self.excluded.iter().any(|excluded| key == *excluded) { + values.push(value); + } + } + Ok(values) + } + } + + let mut deserializer = serde_json::Deserializer::from_slice(body); + let values = + serde::de::Deserializer::deserialize_map(&mut deserializer, Values { excluded }).ok()?; + deserializer.end().ok()?; + 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; @@ -1636,6 +1776,7 @@ fn append_raw_text_value(out: &mut String, raw: &serde_json::value::RawValue) -> fn append_chat_request_content_strings( out: &mut String, content: &serde_json::value::RawValue, + scan_error: &mut Option, ) -> Option<()> { let value = content.get().trim_start(); if value.starts_with('"') { @@ -1654,7 +1795,7 @@ fn append_chat_request_content_strings( if !types.is_empty() && kind.is_none() { append_scan_text( out, - &decoded_json_string_values_including_empty(block_body)?, + &decoded_json_string_values_including_empty(block_body, scan_error)?, ); continue; } @@ -1662,14 +1803,17 @@ fn append_chat_request_content_strings( Some("redacted_thinking") => {} Some("tool_result") => { for nested in raw_top_level_values(block_body, "content")? { - append_chat_request_content_strings(out, &nested)?; + append_chat_request_content_strings(out, &nested, scan_error)?; } } Some("tool_use") => { for input in raw_top_level_values(block_body, "input")? { append_scan_text( out, - &decoded_json_string_values_including_empty(input.get().as_bytes())?, + &decoded_json_string_values_including_empty( + input.get().as_bytes(), + scan_error, + )?, ); } } @@ -1683,6 +1827,7 @@ fn append_chat_request_content_strings( fn append_chat_request_message_strings( out: &mut String, message: &serde_json::value::RawValue, + scan_error: &mut Option, ) -> Option<()> { if !raw_is_object(message) { return Some(()); @@ -1693,26 +1838,33 @@ fn append_chat_request_message_strings( &decoded_json_string_values_except_root_keys( message_body, &["content", "tool_calls", "reasoning_content", "reasoning"], + scan_error, )?, ); for content in raw_top_level_values(message_body, "content")? { - append_chat_request_content_strings(out, &content)?; + append_chat_request_content_strings(out, &content, scan_error)?; } for tool_calls in raw_top_level_values(message_body, "tool_calls")? { append_scan_text( out, - &decoded_json_string_values_including_empty(tool_calls.get().as_bytes())?, + &decoded_json_string_values_including_empty(tool_calls.get().as_bytes(), scan_error)?, ); } append_raw_top_level_strings(out, message_body, "reasoning_content")?; Some(()) } -fn decoded_chat_request_string_values(body: &[u8]) -> Option { - let mut out = - decoded_json_string_values_except_root_keys(body, &["model", "system", "messages"])?; +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, + )?; for system in raw_top_level_values(body, "system")? { - append_chat_request_content_strings(&mut out, &system)?; + append_chat_request_content_strings(&mut out, &system, scan_error)?; } for array in raw_top_level_values(body, "messages")? { // The selected (last) carrier made this a Chat envelope. Preserve @@ -1722,7 +1874,7 @@ fn decoded_chat_request_string_values(body: &[u8]) -> Option { continue; }; for message in messages { - append_chat_request_message_strings(&mut out, &message)?; + append_chat_request_message_strings(&mut out, &message, scan_error)?; } } Some(out) @@ -1731,6 +1883,7 @@ fn decoded_chat_request_string_values(body: &[u8]) -> Option { fn append_responses_item_strings( out: &mut String, item: &serde_json::value::RawValue, + scan_error: &mut Option, ) -> Option<()> { if !raw_is_object(item) { return Some(()); @@ -1747,7 +1900,7 @@ fn append_responses_item_strings( ]; append_scan_text( out, - &decoded_json_string_values_except_root_keys(item_body, &text_keys)?, + &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)? { @@ -1757,24 +1910,34 @@ fn append_responses_item_strings( Some(()) } -fn decoded_responses_request_string_values(body: &[u8]) -> Option { - let mut out = decoded_json_string_values_except_root_keys(body, &["model", "input"])?; +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)?; + append_responses_item_strings(&mut out, &item, scan_error)?; } } } Some(out) } -fn decoded_completions_request_string_values(body: &[u8]) -> Option { - let mut out = - decoded_json_string_values_except_root_keys(body, &["model", "prompt", "suffix"])?; +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('"') { @@ -1797,22 +1960,45 @@ fn decoded_completions_request_string_values(body: &[u8]) -> Option { /// 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(protocol: PassthroughProtocol, body: &[u8]) -> String { +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 => decoded_json_string_values(body).unwrap_or_else(raw), PassthroughProtocol::OpenaiChat => { - decoded_chat_request_string_values(body).unwrap_or_else(raw) + decoded_chat_request_string_values(body, scan_error).unwrap_or_else(raw) } PassthroughProtocol::OpenaiCompletions => { - decoded_completions_request_string_values(body).unwrap_or_else(raw) + decoded_completions_request_string_values(body, scan_error).unwrap_or_else(raw) } PassthroughProtocol::OpenaiResponses => { - decoded_responses_request_string_values(body).unwrap_or_else(raw) + decoded_responses_request_string_values(body, scan_error).unwrap_or_else(raw) } } } +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) +} + +/// 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, + body: &[u8], +) -> Result { + if matches!(protocol, PassthroughProtocol::Raw) { + return crate::json_splice::collect_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. @@ -1937,6 +2123,7 @@ fn append_chat_output_message_strings( fn append_anthropic_output_content_strings( out: &mut String, content: &serde_json::value::RawValue, + scan_error: &mut Option, ) -> Option<()> { if !content.get().trim_start().starts_with('[') { return Some(()); @@ -1953,7 +2140,10 @@ fn append_anthropic_output_content_strings( for input in raw_top_level_values(block_body, "input")? { append_scan_text( out, - &decoded_json_string_values_including_empty(input.get().as_bytes())?, + &decoded_json_string_values_including_empty( + input.get().as_bytes(), + scan_error, + )?, ); } } @@ -1963,7 +2153,10 @@ fn append_anthropic_output_content_strings( Some(()) } -fn decoded_chat_response_string_values(body: &[u8]) -> Option { +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 @@ -1984,7 +2177,7 @@ fn decoded_chat_response_string_values(body: &[u8]) -> Option { } if !has_chat_choices && raw_top_level_unique_type(body).as_deref() == Some("message") { for content in raw_top_level_values(body, "content")? { - append_anthropic_output_content_strings(&mut out, &content)?; + append_anthropic_output_content_strings(&mut out, &content, scan_error)?; } } Some(out) @@ -2094,7 +2287,11 @@ fn decoded_responses_response_string_values(body: &[u8]) -> Option { /// 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(protocol: PassthroughProtocol, body: &[u8]) -> String { +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), @@ -2102,7 +2299,7 @@ fn response_guardrail_text(protocol: PassthroughProtocol, body: &[u8]) -> String // 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. - decoded_chat_response_string_values(body).unwrap_or_default() + decoded_chat_response_string_values(body, scan_error).unwrap_or_default() } PassthroughProtocol::OpenaiCompletions => { decoded_non_model_json_string_values(body).unwrap_or_else(raw) @@ -2116,6 +2313,38 @@ fn response_guardrail_text(protocol: PassthroughProtocol, body: &[u8]) -> String } } +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) +} + +/// 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 => crate::json_splice::collect_string_values(body), + PassthroughProtocol::OpenaiCompletions => { + let text = crate::json_splice::collect_string_values_where(body, |path| { + !is_root_key(path, "model") + })?; + Ok(if text.is_empty() { + response_guardrail_text(protocol, body) + } else { + text + }) + } + 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 => Ok(response_guardrail_text(protocol, body)), + } +} + /// 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 { @@ -3212,7 +3441,10 @@ fn anthropic_source_continuations( "text", ), Some("tool_use") => { - let mut inputs = raw_top_level_values(block_body, "input").ok_or(())?; + 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; }; @@ -3237,8 +3469,11 @@ fn anthropic_source_continuations( // durable leaf identity on this envelope. An empty object is // harmless; a visible value must use the bounded policy. let text = - crate::json_splice::collect_string_values(input.get().as_bytes()) - .map_err(|_| ())?; + match crate::json_splice::collect_string_values(input.get().as_bytes()) + { + Ok(text) => text, + Err(_) => return SourceContinuations::Unevaluable, + }; if text.is_empty() { Ok(()) } else { 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 index dfa69bbd7..e4d2a59bb 100644 --- a/tests/e2e/src/cases/passthrough-chat-media-guardrail-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-chat-media-guardrail-e2e.test.ts @@ -36,6 +36,9 @@ 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)}`; interface ModerationSink { baseUrl: string; @@ -188,6 +191,52 @@ const streamedRefusalResponse = [ "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", +]; + describe("Chat passthrough keeps media out of external output guardrails", () => { let app: SpawnedApp | undefined; let seed: SeedClient | undefined; @@ -196,6 +245,9 @@ describe("Chat passthrough keeps media out of external output guardrails", () => 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 moderation: ModerationSink | undefined; let etcdReachable = false; @@ -214,6 +266,15 @@ describe("Chat passthrough keeps media out of external output guardrails", () => nonStreamBody: bufferedContentRefusalResponse, }); streamRefusalUpstream = await startOpenAiUpstream({ rawStreamFrames: streamedRefusalResponse }); + streamOversizedRefusalUpstream = await startOpenAiUpstream({ + rawStreamFrames: streamedOversizedRefusalResponse, + }); + streamOversizedContentRefusalUpstream = await startOpenAiUpstream({ + rawStreamFrames: streamedOversizedContentRefusalResponse, + }); + streamOversizedLegacyToolUpstream = await startOpenAiUpstream({ + rawStreamFrames: streamedOversizedLegacyToolResponse, + }); app = await spawnApp(); seed = new SeedClient(etcd, app.etcdPrefix); @@ -252,6 +313,24 @@ describe("Chat passthrough keeps media out of external output guardrails", () => 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.createGuardrail({ name: "passthrough-chat-media-output", enabled: true, @@ -279,6 +358,9 @@ describe("Chat passthrough keeps media out of external output guardrails", () => await bufferedMessageRefusalUpstream?.close(); await bufferedContentRefusalUpstream?.close(); await streamRefusalUpstream?.close(); + await streamOversizedRefusalUpstream?.close(); + await streamOversizedContentRefusalUpstream?.close(); + await streamOversizedLegacyToolUpstream?.close(); await moderation?.close(); }); @@ -426,4 +508,55 @@ describe("Chat passthrough keeps media out of external output guardrails", () => } 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-scan-coverage-e2e.test.ts b/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts index 2f02e888c..d0719a4ad 100644 --- a/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts @@ -47,9 +47,21 @@ const RAW_SUFFIX_SSE = `"BIDDEN"`; const RAW_HELD_BLOCK_SSE = `"FORBIDDEN"`; const deepEscapedBlockJSON = (depth: number) => `${'{"v":'.repeat(depth)}"${String.raw`\u0042LOCKME`}"${'}'.repeat(depth)}`; +const deepLiteralBlockJSON = (depth: number) => + `${'{"v":'.repeat(depth)}"${ESCAPED_BLOCK}"${'}'.repeat(depth)}`; // 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 OVER_DEPTH_ESCAPED_BLOCK_JSON = deepEscapedBlockJSON(JSON_DEPTH_CAP + 1); +const OVER_DEPTH_OPAQUE_BLOCK_JSON = deepLiteralBlockJSON(JSON_DEPTH_CAP + 1); +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 SPLIT_BLOCK = "FORBIDDEN"; const SPLIT_BLOCK_REGEX = String.raw`FOR\s*BIDDEN`; @@ -155,6 +167,18 @@ describe("passthrough guardrail scan coverage", () => { 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"], }); @@ -523,6 +547,46 @@ describe("passthrough guardrail scan coverage", () => { 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.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; @@ -564,6 +628,53 @@ describe("passthrough guardrail scan coverage", () => { 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; @@ -652,16 +763,18 @@ describe("passthrough guardrail scan coverage", () => { }); // This is intentionally a separate DP: the main suite has an env-scoped -// blocking row, so it cannot demonstrate the live, monitor-only policy of -// an output `fail_open: true` chain. +// blocking row, so it cannot demonstrate the live fail-open policy on either +// an unevaluable Raw input or output stream. describe("passthrough Raw stream 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 logstore = "pt-scan-fail-open"; const credentialRef = "pt_scan_open"; let app: SpawnedApp | undefined; let upstream: OpenAiUpstream | undefined; + let depthInputUpstream: OpenAiUpstream | undefined; let sls: MockSls | undefined; let etcdReachable = false; @@ -679,6 +792,10 @@ describe("passthrough Raw stream unevaluable-output fail-open", () => { "data: [DONE]\n\n", ], }); + depthInputUpstream = await startOpenAiUpstream({ + rawBody: SAFE_ESCAPED_JSON, + rawContentType: "application/json", + }); app = await spawnApp({ extraEnv: { [`SLS_CRED_${credentialRef.toUpperCase()}_AK_ID`]: "mock-akid", @@ -706,6 +823,20 @@ describe("passthrough Raw stream unevaluable-output 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.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, @@ -732,9 +863,33 @@ describe("passthrough Raw stream unevaluable-output fail-open", () => { afterAll(async () => { await app?.exit(); await upstream?.close(); + await depthInputUpstream?.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("starts a new live scan epoch after an unkeyable Raw SSE object", async (ctx) => { if (!etcdReachable || !app || !sls || !upstream) return ctx.skip(); From 8f445b43457c9e037ce46a3e2f3b290df40cb8e6 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Thu, 1 Oct 2026 11:30:18 +0800 Subject: [PATCH 23/37] fix: bound nested passthrough tool-result scanning --- crates/aisix-proxy/src/json_splice.rs | 13 +- crates/aisix-proxy/src/passthrough_route.rs | 144 +++++++++++++----- .../passthrough-scan-coverage-e2e.test.ts | 59 +++++++ 3 files changed, 179 insertions(+), 37 deletions(-) diff --git a/crates/aisix-proxy/src/json_splice.rs b/crates/aisix-proxy/src/json_splice.rs index 39415f1c3..e6c86363b 100644 --- a/crates/aisix-proxy/src/json_splice.rs +++ b/crates/aisix-proxy/src/json_splice.rs @@ -59,12 +59,23 @@ impl SpliceError { pub(crate) fn is_depth_exceeded(self) -> bool { self.kind == SpliceErrorKind::DepthExceeded } + + /// 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, + } + } } /// 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. -const MAX_JSON_DEPTH: usize = 4_096; +pub(crate) const MAX_JSON_DEPTH: usize = 4_096; /// Rewrite the string values of `input` selected by `should_rewrite`, /// leaving every other byte untouched. diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index a496b7fa7..664c80ad4 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -1775,50 +1775,104 @@ fn append_raw_text_value(out: &mut String, raw: &serde_json::value::RawValue) -> /// duplicate `type` is scanned as source rather than becoming a bypass. fn append_chat_request_content_strings( out: &mut String, - content: &serde_json::value::RawValue, + content: Box, scan_error: &mut Option, ) -> Option<()> { - let value = content.get().trim_start(); - if value.starts_with('"') { - return append_raw_string_value(out, content); - } - if !value.starts_with('[') { - return Some(()); + enum Work { + Content { + value: Box, + depth: usize, + }, + Block { + value: Box, + depth: usize, + }, } - for block in raw_array_items(content)? { - if !raw_is_object(&block) { - continue; - } - let block_body = block.get().as_bytes(); - let types = raw_top_level_values(block_body, "type")?; - let kind = raw_top_level_unique_type(block_body); - 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") => { - for nested in raw_top_level_values(block_body, "content")? { - append_chat_request_content_strings(out, &nested, scan_error)?; + + // `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. + let mut work = vec![Work::Content { + value: content, + depth: 1, + }]; + while let Some(work_item) = work.pop() { + 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('[') { + continue; + } + // Push backwards so the LIFO work stack preserves the + // previous depth-first, source-order traversal. + for block in raw_array_items(&value)?.into_iter().rev() { + work.push(Work::Block { + value: block, + depth, + }); } } - Some("tool_use") => { - for input in raw_top_level_values(block_body, "input")? { + Work::Block { + value: block, + depth, + } => { + if !raw_is_object(&block) { + continue; + } + let block_body = block.get().as_bytes(); + let types = raw_top_level_values(block_body, "type")?; + let kind = raw_top_level_unique_type(block_body); + if !types.is_empty() && kind.is_none() { append_scan_text( out, - &decoded_json_string_values_including_empty( - input.get().as_bytes(), - scan_error, - )?, + &decoded_json_string_values_including_empty(block_body, scan_error)?, ); + continue; + } + match kind.as_deref() { + Some("redacted_thinking") => {} + Some("tool_result") => { + let nested = raw_top_level_values(block_body, "content")?; + 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() { + work.push(Work::Content { + value: nested, + depth, + }); + } + } + Some("tool_use") => { + 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("thinking") => append_raw_top_level_strings(out, block_body, "thinking")?, + _ => append_raw_top_level_strings(out, block_body, "text")?, } } - Some("thinking") => append_raw_top_level_strings(out, block_body, "thinking")?, - _ => append_raw_top_level_strings(out, block_body, "text")?, } } Some(()) @@ -1842,7 +1896,7 @@ fn append_chat_request_message_strings( )?, ); for content in raw_top_level_values(message_body, "content")? { - append_chat_request_content_strings(out, &content, scan_error)?; + append_chat_request_content_strings(out, content, scan_error)?; } for tool_calls in raw_top_level_values(message_body, "tool_calls")? { append_scan_text( @@ -1864,7 +1918,7 @@ fn decoded_chat_request_string_values( scan_error, )?; for system in raw_top_level_values(body, "system")? { - append_chat_request_content_strings(&mut out, &system, scan_error)?; + append_chat_request_content_strings(&mut out, system, scan_error)?; } for array in raw_top_level_values(body, "messages")? { // The selected (last) carrier made this a Chat envelope. Preserve @@ -5644,6 +5698,16 @@ mod tests { json.into_bytes() } + fn nested_anthropic_tool_result_request(depth: usize) -> Vec { + let content = format!( + "{}[{{\"type\":\"text\",\"text\":\"safe\"}}]{}", + 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 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"}}"# @@ -8567,6 +8631,14 @@ mod tests { } } + #[test] + fn nested_anthropic_tool_result_content_beyond_depth_cap_is_unevaluable() { + let body = nested_anthropic_tool_result_request(crate::json_splice::MAX_JSON_DEPTH + 1); + let error = try_request_guardrail_text(PassthroughProtocol::OpenaiChat, &body) + .expect_err("nested tool results beyond the shared JSON depth cap must not recurse"); + assert!(error.is_depth_exceeded(), "{error}"); + } + /// Buffered Anthropic and Responses replies are read slot by slot: /// text and tool input in, generated reasoning out. #[test] 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 d0719a4ad..7e7402498 100644 --- a/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts @@ -49,6 +49,17 @@ 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); @@ -57,6 +68,14 @@ const DEEP_ESCAPED_BLOCK_JSON = deepEscapedBlockJSON(160); const JSON_DEPTH_CAP = 4_096; const OVER_DEPTH_ESCAPED_BLOCK_JSON = deepEscapedBlockJSON(JSON_DEPTH_CAP + 1); const OVER_DEPTH_OPAQUE_BLOCK_JSON = deepLiteralBlockJSON(JSON_DEPTH_CAP + 1); +const OVER_DEPTH_ANTHROPIC_TOOL_RESULT_INPUT = deeplyNestedAnthropicToolResultRequest( + JSON_DEPTH_CAP + 1, + "nested-tool-result-fail-closed", +); +const OVER_DEPTH_ANTHROPIC_TOOL_RESULT_FAIL_OPEN_INPUT = deeplyNestedAnthropicToolResultRequest( + JSON_DEPTH_CAP + 1, + "nested-tool-result-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"}]}]}`; @@ -559,6 +578,18 @@ describe("passthrough guardrail scan coverage", () => { expect(upstreams.input!.receivedRequests.length).toBe(before); }); + test("input: nested Anthropic tool results beyond the scanner depth 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], @@ -890,6 +921,34 @@ describe("passthrough Raw stream unevaluable-output fail-open", () => { expect(log.get("guardrail_bypassed_reason")).toBe("unscannable_body"); }); + test("forwards nested Anthropic tool results 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_ANTHROPIC_TOOL_RESULT_FAIL_OPEN_INPUT, + }); + expect(res.status).toBe(200); + 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("requested_model") === "nested-tool-result-fail-open", + "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(); From 0344c1d5e8586ef2e005a7eb70a68912809b94c4 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Thu, 1 Oct 2026 11:43:30 +0800 Subject: [PATCH 24/37] test: join deep passthrough telemetry by request id --- tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) 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 7e7402498..e44497f6c 100644 --- a/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts @@ -931,6 +931,8 @@ describe("passthrough Raw stream unevaluable-output fail-open", () => { 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( @@ -942,7 +944,7 @@ describe("passthrough Raw stream unevaluable-output fail-open", () => { logstore, (entry) => entry.get("passthrough_route_name") === depthInputRoute && - entry.get("requested_model") === "nested-tool-result-fail-open", + entry.get("request_id") === requestId, "fail-open nested-tool-result passthrough usage event", ); expect(log.get("guardrail_blocked") ?? "false").not.toBe("true"); From 8ccdccd1e37616c1beefa0aa2509e696ada14a92 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Thu, 1 Oct 2026 12:04:43 +0800 Subject: [PATCH 25/37] fix: bound raw tool-result scan allocation --- crates/aisix-proxy/src/passthrough_route.rs | 96 ++++++++++++++++--- .../passthrough-scan-coverage-e2e.test.ts | 17 ++++ 2 files changed, 100 insertions(+), 13 deletions(-) diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index 664c80ad4..6f5ff2601 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -1549,6 +1549,47 @@ fn raw_top_level_values( Some(values) } +/// Borrowed source values of all occurrences of one top-level key. This is +/// for traversals which revisit nested carriers: retaining a reference avoids +/// copying every remaining `tool_result.content` suffix at each level. +fn raw_top_level_value_refs<'a>( + body: &'a [u8], + wanted_key: &str, +) -> Option> { + struct Values<'a> { + wanted_key: &'a str, + } + + impl<'de> serde::de::Visitor<'de> for Values<'_> { + type Value = Vec<&'de serde_json::value::RawValue>; + + fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str("a JSON object") + } + + fn visit_map(self, mut map: A) -> Result + where + A: serde::de::MapAccess<'de>, + { + let mut values = Vec::new(); + while let Some(key) = map.next_key::()? { + if key == self.wanted_key { + values.push(map.next_value::<&'de serde_json::value::RawValue>()?); + } else { + map.next_value::()?; + } + } + Ok(values) + } + } + + let mut deserializer = serde_json::Deserializer::from_slice(body); + let values = + serde::de::Deserializer::deserialize_map(&mut deserializer, Values { wanted_key }).ok()?; + deserializer.end().ok()?; + Some(values) +} + /// Source values of top-level keys other than `excluded`. Values are captured /// as raw JSON before filtering so a known opaque carrier can be skipped /// without recursively deserializing its payload. @@ -1613,6 +1654,12 @@ fn raw_array_items( serde_json::from_str(raw.get()).ok() } +fn raw_array_item_refs<'a>( + raw: &'a serde_json::value::RawValue, +) -> Option> { + serde_json::from_str(raw.get()).ok() +} + /// `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. @@ -1775,16 +1822,16 @@ fn append_raw_text_value(out: &mut String, raw: &serde_json::value::RawValue) -> /// duplicate `type` is scanned as source rather than becoming a bypass. fn append_chat_request_content_strings( out: &mut String, - content: Box, + content: &serde_json::value::RawValue, scan_error: &mut Option, ) -> Option<()> { - enum Work { + enum Work<'a> { Content { - value: Box, + value: &'a serde_json::value::RawValue, depth: usize, }, Block { - value: Box, + value: &'a serde_json::value::RawValue, depth: usize, }, } @@ -1792,10 +1839,11 @@ fn append_chat_request_content_strings( // `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. + // 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![Work::Content { value: content, - depth: 1, + depth: 0, }]; while let Some(work_item) = work.pop() { match work_item { @@ -1810,7 +1858,7 @@ fn append_chat_request_content_strings( } // Push backwards so the LIFO work stack preserves the // previous depth-first, source-order traversal. - for block in raw_array_items(&value)?.into_iter().rev() { + for block in raw_array_item_refs(value)?.into_iter().rev() { work.push(Work::Block { value: block, depth, @@ -1837,7 +1885,7 @@ fn append_chat_request_content_strings( match kind.as_deref() { Some("redacted_thinking") => {} Some("tool_result") => { - let nested = raw_top_level_values(block_body, "content")?; + let nested = raw_top_level_value_refs(block_body, "content")?; if nested.is_empty() { continue; } @@ -1896,7 +1944,7 @@ fn append_chat_request_message_strings( )?, ); for content in raw_top_level_values(message_body, "content")? { - append_chat_request_content_strings(out, content, scan_error)?; + append_chat_request_content_strings(out, &content, scan_error)?; } for tool_calls in raw_top_level_values(message_body, "tool_calls")? { append_scan_text( @@ -1918,7 +1966,7 @@ fn decoded_chat_request_string_values( scan_error, )?; for system in raw_top_level_values(body, "system")? { - append_chat_request_content_strings(&mut out, system, scan_error)?; + append_chat_request_content_strings(&mut out, &system, scan_error)?; } for array in raw_top_level_values(body, "messages")? { // The selected (last) carrier made this a Chat envelope. Preserve @@ -5698,9 +5746,10 @@ mod tests { json.into_bytes() } - fn nested_anthropic_tool_result_request(depth: usize) -> Vec { + 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\":\"safe\"}}]{}", + "{}[{{\"type\":\"text\",\"text\":{text}}}]{}", r#"[{"type":"tool_result","tool_use_id":"t","content":"#.repeat(depth), "}]".repeat(depth), ); @@ -8631,9 +8680,30 @@ 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_at_depth_cap_is_scanned() { + let body = + nested_anthropic_tool_result_request(crate::json_splice::MAX_JSON_DEPTH, "BLOCKME"); + let text = try_request_guardrail_text(PassthroughProtocol::OpenaiChat, &body) + .expect("nested tool results at the shared JSON depth cap remain evaluable"); + assert!(text.contains("BLOCKME"), "{text}"); + } + #[test] fn nested_anthropic_tool_result_content_beyond_depth_cap_is_unevaluable() { - let body = nested_anthropic_tool_result_request(crate::json_splice::MAX_JSON_DEPTH + 1); + let body = + nested_anthropic_tool_result_request(crate::json_splice::MAX_JSON_DEPTH + 1, "safe"); let error = try_request_guardrail_text(PassthroughProtocol::OpenaiChat, &body) .expect_err("nested tool results beyond the shared JSON depth cap must not recurse"); assert!(error.is_depth_exceeded(), "{error}"); 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 e44497f6c..dc345cb4d 100644 --- a/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts @@ -68,6 +68,10 @@ const DEEP_ESCAPED_BLOCK_JSON = deepEscapedBlockJSON(160); const JSON_DEPTH_CAP = 4_096; 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( + JSON_DEPTH_CAP, + "nested-tool-result-at-depth-cap", +); const OVER_DEPTH_ANTHROPIC_TOOL_RESULT_INPUT = deeplyNestedAnthropicToolResultRequest( JSON_DEPTH_CAP + 1, "nested-tool-result-fail-closed", @@ -578,6 +582,19 @@ describe("passthrough guardrail scan coverage", () => { expect(upstreams.input!.receivedRequests.length).toBe(before); }); + test("input: nested Anthropic tool results at the scanner depth 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 scanner depth cap fail closed", async (ctx) => { if (!ready(ctx)) return; const before = upstreams.input!.receivedRequests.length; From b5007f6cf8d6dd9e37dc7b3047281459d6c53ef3 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Thu, 1 Oct 2026 12:14:15 +0800 Subject: [PATCH 26/37] fix(passthrough): satisfy raw scan clippy --- crates/aisix-proxy/src/passthrough_route.rs | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index 6f5ff2601..b7c624f2a 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -1654,9 +1654,9 @@ fn raw_array_items( serde_json::from_str(raw.get()).ok() } -fn raw_array_item_refs<'a>( - raw: &'a serde_json::value::RawValue, -) -> Option> { +fn raw_array_item_refs( + raw: &serde_json::value::RawValue, +) -> Option> { serde_json::from_str(raw.get()).ok() } @@ -1850,7 +1850,7 @@ fn append_chat_request_content_strings( Work::Content { value, depth } => { let value_text = value.get().trim_start(); if value_text.starts_with('"') { - append_raw_string_value(out, &value)?; + append_raw_string_value(out, value)?; continue; } if !value_text.starts_with('[') { @@ -1869,7 +1869,7 @@ fn append_chat_request_content_strings( value: block, depth, } => { - if !raw_is_object(&block) { + if !raw_is_object(block) { continue; } let block_body = block.get().as_bytes(); From 49e8f7ae0c5df690007de993b7cffb4fbfd03ebb Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Thu, 1 Oct 2026 13:00:07 +0800 Subject: [PATCH 27/37] fix: fail closed for unevaluable passthrough scans --- crates/aisix-proxy/src/json_splice.rs | 21 + crates/aisix-proxy/src/passthrough_route.rs | 547 ++++++++++++------ .../passthrough-scan-coverage-e2e.test.ts | 21 + 3 files changed, 407 insertions(+), 182 deletions(-) diff --git a/crates/aisix-proxy/src/json_splice.rs b/crates/aisix-proxy/src/json_splice.rs index e6c86363b..86b6cd210 100644 --- a/crates/aisix-proxy/src/json_splice.rs +++ b/crates/aisix-proxy/src/json_splice.rs @@ -44,6 +44,7 @@ impl PathSeg { enum SpliceErrorKind { Invalid, DepthExceeded, + Unevaluable, } #[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)] @@ -60,6 +61,17 @@ impl SpliceError { 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 @@ -70,6 +82,15 @@ impl SpliceError { 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, + } + } } /// Traversal depth cap. This keeps the iterative scanner stack-safe without diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index b7c624f2a..4c3965743 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -582,7 +582,7 @@ async fn dispatch( if !resolved_chain.is_empty() { let text = match try_request_guardrail_text(protocol, &body_bytes) { Ok(text) => Some(text), - Err(err) if !err.is_depth_exceeded() => { + Err(err) if !err.is_unevaluable() => { Some(request_guardrail_text(protocol, &body_bytes)) } Err(err) @@ -592,7 +592,7 @@ async fn dispatch( guardrail_hook = "input", route = %route.name, error = %err, - "cannot scan passthrough-route request to its bounded depth; nothing attached both reads the request and fails closed", + "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 @@ -602,7 +602,7 @@ async fn dispatch( guardrail_hook = "input", route = %route.name, error = %err, - "cannot scan passthrough-route request to its bounded depth; blocking", + "cannot safely select passthrough-route request text; blocking", ); return Err(RouteError::of( crate::error::guardrail_block_error( @@ -998,7 +998,7 @@ async fn dispatch( if !resolved_chain.is_empty() { let text = match try_response_guardrail_text(protocol, &resp_body) { Ok(text) => Some(text), - Err(err) if !err.is_depth_exceeded() => { + Err(err) if !err.is_unevaluable() => { Some(response_guardrail_text(protocol, &resp_body)) } Err(err) @@ -1008,7 +1008,7 @@ async fn dispatch( guardrail_hook = "output", route = %route.name, error = %err, - "cannot scan passthrough-route response to its bounded depth; nothing attached both reads the response and fails closed", + "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 @@ -1018,7 +1018,7 @@ async fn dispatch( guardrail_hook = "output", route = %route.name, error = %err, - "cannot scan passthrough-route response to its bounded depth; blocking", + "cannot safely select passthrough-route response text; blocking", ); telemetry.guardrail_blocked = true; telemetry.emitted = true; @@ -1345,8 +1345,9 @@ enum PassthroughProtocol { fn detect_protocol(body: &[u8]) -> PassthroughProtocol { // 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. `RawValue` keeps the chosen top-level - // carrier shallow while preserving the last-key behavior of a JSON map. + // 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 raw_top_level_last_has_shape(body, "input", true) { @@ -1489,6 +1490,15 @@ fn decoded_json_string_values_including_empty( } } +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()); + } +} + fn decoded_json_string_values_except_root_keys( body: &[u8], excluded: &[&str], @@ -1509,124 +1519,190 @@ fn decoded_json_string_values_except_root_keys( Some(out) } -/// Source values of all occurrences of one top-level key. `RawValue` keeps -/// repeated keys separate, unlike `serde_json::Value`. -fn raw_top_level_values( - body: &[u8], - wanted_key: &str, -) -> Option>> { - struct Values<'a> { - wanted_key: &'a str, - } +/// 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<'de> serde::de::Visitor<'de> for Values<'_> { - type Value = Vec>; +impl<'a> RawJson<'a> { + fn get(self) -> &'a str { + self.source + } +} - fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - formatter.write_str("a JSON object") +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; + } +} + +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, } + } + None +} - fn visit_map(self, mut map: A) -> Result - where - A: serde::de::MapAccess<'de>, - { - let mut values = Vec::new(); - while let Some(key) = map.next_key::()? { - if key == self.wanted_key { - values.push(map.next_value::>()?); - } else { - map.next_value::()?; +/// 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, } } - Ok(values) + 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 { + serde_json::from_slice::(token) + .ok() + .map(|_| pos) + } } } - - let mut deserializer = serde_json::Deserializer::from_slice(body); - let values = - serde::de::Deserializer::deserialize_map(&mut deserializer, Values { wanted_key }).ok()?; - deserializer.end().ok()?; - Some(values) } -/// Borrowed source values of all occurrences of one top-level key. This is -/// for traversals which revisit nested carriers: retaining a reference avoids -/// copying every remaining `tool_result.content` suffix at each level. -fn raw_top_level_value_refs<'a>( - body: &'a [u8], - wanted_key: &str, -) -> Option> { - struct Values<'a> { - wanted_key: &'a str, - } - - impl<'de> serde::de::Visitor<'de> for Values<'_> { - type Value = Vec<&'de serde_json::value::RawValue>; - - fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - formatter.write_str("a JSON object") - } - - fn visit_map(self, mut map: A) -> Result - where - A: serde::de::MapAccess<'de>, - { - let mut values = Vec::new(); - while let Some(key) = map.next_key::()? { - if key == self.wanted_key { - values.push(map.next_value::<&'de serde_json::value::RawValue>()?); - } else { - map.next_value::()?; - } +/// 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(()); } - Ok(values) + _ => return None, } } +} - let mut deserializer = serde_json::Deserializer::from_slice(body); - let values = - serde::de::Deserializer::deserialize_map(&mut deserializer, Values { wanted_key }).ok()?; - deserializer.end().ok()?; +/// 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(); + raw_object_members(body, |key, value| { + if key == wanted_key { + values.push(value); + } + })?; Some(values) } -/// Source values of top-level keys other than `excluded`. Values are captured -/// as raw JSON before filtering so a known opaque carrier can be skipped -/// without recursively deserializing its payload. -fn raw_top_level_values_except( - body: &[u8], - excluded: &[&str], -) -> Option>> { - struct Values<'a> { - excluded: &'a [&'a str], - } - - impl<'de> serde::de::Visitor<'de> for Values<'_> { - type Value = Vec>; - - fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - formatter.write_str("a JSON object") - } +/// 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>> { + raw_top_level_values(body, wanted_key) +} - fn visit_map(self, mut map: A) -> Result - where - A: serde::de::MapAccess<'de>, - { - let mut values = Vec::new(); - while let Some(key) = map.next_key::()? { - let value = map.next_value::>()?; - if !self.excluded.iter().any(|excluded| key == *excluded) { - values.push(value); - } - } - Ok(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(); + raw_object_members(body, |key, value| { + if !excluded.iter().any(|excluded| key == *excluded) { + values.push(value); } - } - - let mut deserializer = serde_json::Deserializer::from_slice(body); - let values = - serde::de::Deserializer::deserialize_map(&mut deserializer, Values { excluded }).ok()?; - deserializer.end().ok()?; + })?; Some(values) } @@ -1644,26 +1720,52 @@ fn raw_top_level_last_has_shape(body: &[u8], key: &str, allow_string: bool) -> b }) } -fn raw_is_object(raw: &serde_json::value::RawValue) -> bool { +fn raw_is_object(raw: &RawJson<'_>) -> bool { raw.get().trim_start().starts_with('{') } -fn raw_array_items( - raw: &serde_json::value::RawValue, -) -> Option>> { - serde_json::from_str(raw.get()).ok() +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(); + 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)?; + values.push(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( - raw: &serde_json::value::RawValue, -) -> Option> { - serde_json::from_str(raw.get()).ok() +fn raw_array_item_refs<'a>(raw: &RawJson<'a>) -> Option>> { + raw_array_items(raw) } /// `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: &serde_json::value::RawValue, allowed: &[&str]) -> bool { +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; }; @@ -1713,7 +1815,7 @@ fn raw_top_level_items_have_only_types(body: &[u8], key: &str, allowed: &[&str]) ) } -fn append_raw_string_value(out: &mut String, raw: &serde_json::value::RawValue) -> Option<()> { +fn append_raw_string_value(out: &mut String, raw: &RawJson<'_>) -> Option<()> { append_scan_text(out, &serde_json::from_str::(raw.get()).ok()?); Some(()) } @@ -1759,10 +1861,7 @@ fn raw_top_level_unique_index(body: &[u8], key: &str) -> Result, ( } } -fn raw_top_level_unique_object( - body: &[u8], - key: &str, -) -> Result>, ()> { +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), @@ -1774,10 +1873,7 @@ fn raw_top_level_unique_object( } } -fn raw_top_level_unique_array( - body: &[u8], - key: &str, -) -> Result>, ()> { +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), @@ -1798,7 +1894,7 @@ fn raw_top_level_unique_array( /// 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: &serde_json::value::RawValue) -> Option<()> { +fn append_raw_text_value(out: &mut String, raw: &RawJson<'_>) -> Option<()> { let value = raw.get().trim_start(); if value.starts_with('"') { return append_raw_string_value(out, raw); @@ -1822,18 +1918,12 @@ fn append_raw_text_value(out: &mut String, raw: &serde_json::value::RawValue) -> /// duplicate `type` is scanned as source rather than becoming a bypass. fn append_chat_request_content_strings( out: &mut String, - content: &serde_json::value::RawValue, + content: &RawJson<'_>, scan_error: &mut Option, ) -> Option<()> { enum Work<'a> { - Content { - value: &'a serde_json::value::RawValue, - depth: usize, - }, - Block { - value: &'a serde_json::value::RawValue, - depth: usize, - }, + Content { value: RawJson<'a>, depth: usize }, + Block { value: RawJson<'a>, depth: usize }, } // `tool_result.content` can itself contain another `tool_result`. Keep @@ -1842,7 +1932,7 @@ fn append_chat_request_content_strings( // 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![Work::Content { - value: content, + value: *content, depth: 0, }]; while let Some(work_item) = work.pop() { @@ -1850,15 +1940,26 @@ fn append_chat_request_content_strings( Work::Content { value, depth } => { let value_text = value.get().trim_start(); if value_text.starts_with('"') { - append_raw_string_value(out, value)?; + 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. - for block in raw_array_item_refs(value)?.into_iter().rev() { + let Some(blocks) = raw_array_item_refs(&value) else { + mark_unevaluable(scan_error); + return None; + }; + for block in blocks.into_iter().rev() { work.push(Work::Block { value: block, depth, @@ -1869,11 +1970,15 @@ fn append_chat_request_content_strings( value: block, depth, } => { - if !raw_is_object(block) { - continue; + if !raw_is_object(&block) { + mark_unevaluable(scan_error); + return None; } let block_body = block.get().as_bytes(); - let types = raw_top_level_values(block_body, "type")?; + let Some(types) = raw_top_level_values(block_body, "type") else { + mark_unevaluable(scan_error); + return None; + }; let kind = raw_top_level_unique_type(block_body); if !types.is_empty() && kind.is_none() { append_scan_text( @@ -1885,7 +1990,10 @@ fn append_chat_request_content_strings( match kind.as_deref() { Some("redacted_thinking") => {} Some("tool_result") => { - let nested = raw_top_level_value_refs(block_body, "content")?; + let Some(nested) = raw_top_level_value_refs(block_body, "content") else { + mark_unevaluable(scan_error); + return None; + }; if nested.is_empty() { continue; } @@ -1907,7 +2015,11 @@ fn append_chat_request_content_strings( } } Some("tool_use") => { - for input in raw_top_level_values(block_body, "input")? { + 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( @@ -1928,11 +2040,12 @@ fn append_chat_request_content_strings( fn append_chat_request_message_strings( out: &mut String, - message: &serde_json::value::RawValue, + message: &RawJson<'_>, scan_error: &mut Option, ) -> Option<()> { if !raw_is_object(message) { - return Some(()); + mark_unevaluable(scan_error); + return None; } let message_body = message.get().as_bytes(); append_scan_text( @@ -1943,10 +2056,18 @@ fn append_chat_request_message_strings( scan_error, )?, ); - for content in raw_top_level_values(message_body, "content")? { + 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)?; } - for tool_calls in raw_top_level_values(message_body, "tool_calls")? { + 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)?, @@ -1965,15 +2086,24 @@ fn decoded_chat_request_string_values( &["model", "system", "messages"], scan_error, )?; - for system in raw_top_level_values(body, "system")? { + 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)?; } - for array in raw_top_level_values(body, "messages")? { + 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 { - continue; + mark_unevaluable(scan_error); + return None; }; for message in messages { append_chat_request_message_strings(&mut out, &message, scan_error)?; @@ -1984,7 +2114,7 @@ fn decoded_chat_request_string_values( fn append_responses_item_strings( out: &mut String, - item: &serde_json::value::RawValue, + item: &RawJson<'_>, scan_error: &mut Option, ) -> Option<()> { if !raw_is_object(item) { @@ -2071,13 +2201,27 @@ fn request_guardrail_text_with_scan_error( match protocol { PassthroughProtocol::Raw => decoded_json_string_values(body).unwrap_or_else(raw), PassthroughProtocol::OpenaiChat => { - decoded_chat_request_string_values(body, scan_error).unwrap_or_else(raw) + // 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 => { decoded_completions_request_string_values(body, scan_error).unwrap_or_else(raw) } PassthroughProtocol::OpenaiResponses => { - decoded_responses_request_string_values(body, scan_error).unwrap_or_else(raw) + match decoded_responses_request_string_values(body, scan_error) { + Some(text) => text, + None => { + mark_unevaluable(scan_error); + String::new() + } + } } } } @@ -2116,10 +2260,7 @@ fn chat_visible_content_part_field(kind: &str) -> Option<&'static str> { /// 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: &serde_json::value::RawValue, -) -> Option<()> { +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); @@ -2155,10 +2296,7 @@ fn chat_tool_continuation_fields(body: &[u8]) -> Option<&'static [(&'static str, } } -fn append_chat_tool_call_strings( - out: &mut String, - tool_calls: &serde_json::value::RawValue, -) -> Option<()> { +fn append_chat_tool_call_strings(out: &mut String, tool_calls: &RawJson<'_>) -> Option<()> { if !tool_calls.get().trim_start().starts_with('[') { return Some(()); } @@ -2185,7 +2323,7 @@ fn append_chat_tool_call_strings( /// response walk beyond its explicit `name` and `arguments` fields. fn append_chat_legacy_function_call_strings( out: &mut String, - function_call: &serde_json::value::RawValue, + function_call: &RawJson<'_>, ) -> Option<()> { if !raw_is_object(function_call) { return Some(()); @@ -2196,10 +2334,7 @@ fn append_chat_legacy_function_call_strings( Some(()) } -fn append_chat_output_message_strings( - out: &mut String, - message: &serde_json::value::RawValue, -) -> Option<()> { +fn append_chat_output_message_strings(out: &mut String, message: &RawJson<'_>) -> Option<()> { if !raw_is_object(message) { return Some(()); } @@ -2224,7 +2359,7 @@ fn append_chat_output_message_strings( /// content-block kinds remain opaque. fn append_anthropic_output_content_strings( out: &mut String, - content: &serde_json::value::RawValue, + content: &RawJson<'_>, scan_error: &mut Option, ) -> Option<()> { if !content.get().trim_start().starts_with('[') { @@ -2300,10 +2435,7 @@ const RESPONSES_VISIBLE_TEXT_PART_TYPES: &[&str] = &["output_text", "text", "inp /// 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: &serde_json::value::RawValue, -) -> Option<()> { +fn append_responses_visible_part_strings(out: &mut String, part: &RawJson<'_>) -> Option<()> { let part_body = part.get().as_bytes(); match raw_top_level_unique_type(part_body).as_deref() { Some(kind) if RESPONSES_VISIBLE_TEXT_PART_TYPES.contains(&kind) => { @@ -2314,10 +2446,7 @@ fn append_responses_visible_part_strings( Some(()) } -fn append_responses_visible_content_strings( - out: &mut String, - content: &serde_json::value::RawValue, -) -> Option<()> { +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); @@ -2339,10 +2468,7 @@ fn append_responses_visible_content_strings( /// 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: &serde_json::value::RawValue, -) -> Option<()> { +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).as_deref() { Some("reasoning") => {} @@ -2401,7 +2527,13 @@ fn response_guardrail_text_with_scan_error( // 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. - decoded_chat_response_string_values(body, scan_error).unwrap_or_default() + match decoded_chat_response_string_values(body, scan_error) { + Some(text) => text, + None => { + mark_unevaluable(scan_error); + String::new() + } + } } PassthroughProtocol::OpenaiCompletions => { decoded_non_model_json_string_values(body).unwrap_or_else(raw) @@ -2443,7 +2575,8 @@ fn try_response_guardrail_text( let text = response_guardrail_text_with_scan_error(protocol, body, &mut scan_error); scan_error.map_or(Ok(text), Err) } - PassthroughProtocol::OpenaiResponses => Ok(response_guardrail_text(protocol, body)), + PassthroughProtocol::OpenaiResponses => decoded_responses_response_string_values(body) + .ok_or_else(crate::json_splice::SpliceError::unevaluable), } } @@ -5757,6 +5890,18 @@ mod tests { .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"}}"# @@ -8709,6 +8854,44 @@ mod tests { assert!(error.is_depth_exceeded(), "{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}"); + } + } + /// Buffered Anthropic and Responses replies are read slot by slot: /// text and tool input in, generated reasoning out. #[test] 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 dc345cb4d..9fdbe48e4 100644 --- a/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts @@ -549,6 +549,27 @@ describe("passthrough guardrail scan coverage", () => { 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.for([ ["ASCII", ESCAPED_BLOCK_JSON], ["CJK", ESCAPED_CJK_JSON], From 01c7b44460f806473aa82e41cb15ccb906437223 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Thu, 1 Oct 2026 13:04:21 +0800 Subject: [PATCH 28/37] fix: satisfy passthrough selector lint --- crates/aisix-proxy/src/passthrough_route.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index 4c3965743..ddce42a66 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -1699,7 +1699,7 @@ fn raw_top_level_value_refs<'a>(body: &'a [u8], wanted_key: &str) -> Option(body: &'a [u8], excluded: &[&str]) -> Option>> { let mut values = Vec::new(); raw_object_members(body, |key, value| { - if !excluded.iter().any(|excluded| key == *excluded) { + if !excluded.contains(&key) { values.push(value); } })?; From c5234d77ba3a0014f2e66e194b82973fd70fbb11 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Thu, 1 Oct 2026 13:46:40 +0800 Subject: [PATCH 29/37] fix: harden passthrough guardrail boundaries --- crates/aisix-proxy/src/json_splice.rs | 165 +++- crates/aisix-proxy/src/passthrough_route.rs | 906 ++++++++++++++---- .../passthrough-scan-coverage-e2e.test.ts | 71 ++ 3 files changed, 960 insertions(+), 182 deletions(-) diff --git a/crates/aisix-proxy/src/json_splice.rs b/crates/aisix-proxy/src/json_splice.rs index 86b6cd210..6f15c7810 100644 --- a/crates/aisix-proxy/src/json_splice.rs +++ b/crates/aisix-proxy/src/json_splice.rs @@ -98,6 +98,55 @@ impl SpliceError { /// 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. /// @@ -122,6 +171,7 @@ pub fn rewrite_string_values( 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(); @@ -138,15 +188,25 @@ 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, - kind: SpliceErrorKind::Invalid, - }) + Err(err(start)) }; let decode_str = |range: Range| -> Result { let at = range.start; @@ -223,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)), } @@ -315,6 +382,12 @@ 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 @@ -325,18 +398,36 @@ pub fn collect_string_values_where( 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 !out.is_empty() { - out.push('\n'); + 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); } - out.push_str(value); None }, )?; - Ok(out) + collect_error.map_or(Ok(out), Err) } /// Decode selected JSON string values as separate source-order entries. @@ -349,15 +440,37 @@ pub fn collect_string_values_where_vec( 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| { - out.push(value.to_string()); + 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 }, )?; - Ok(out) + collect_error.map_or(Ok(out), Err) } #[cfg(test)] @@ -390,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"}}"#; @@ -472,6 +590,25 @@ mod tests { ); } + #[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 diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index ddce42a66..1db46470d 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -574,13 +574,21 @@ async fn dispatch( let mut monitor_hits: Vec = Vec::new(); // 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 = match try_request_guardrail_text(protocol, &body_bytes) { + 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)) @@ -957,6 +965,21 @@ async fn dispatch( if is_sse { telemetry.streaming = true; + let stream_hold = reservation.into_stream_hold(); + // A non-success SSE response is an upstream error contract, just + // like a buffered 4xx/5xx. It must remain byte-for-byte relay data: + // no output guardrail, heartbeat, or SSE parsing may rewrite it. + if !status.is_success() { + return Ok(stream_non_success_response( + upstream_resp, + resp_headers, + status, + telemetry, + &client.request_id, + stream_hold, + stream_read_timeout, + )); + } return Ok(stream_response( protocol, resolved_chain, @@ -965,7 +988,7 @@ async fn dispatch( status, telemetry, &client.request_id, - reservation.into_stream_hold(), + stream_hold, stream_read_timeout, )); } @@ -994,8 +1017,10 @@ async fn dispatch( ) })?; - // OUTPUT guardrails on the (envelope-extracted) response text. - if !resolved_chain.is_empty() { + // 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() && !resolved_chain.is_empty() { let text = match try_response_guardrail_text(protocol, &resp_body) { Ok(text) => Some(text), Err(err) if !err.is_unevaluable() => { @@ -1332,6 +1357,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 @@ -1359,6 +1395,38 @@ fn detect_protocol(body: &[u8]) -> PassthroughProtocol { } } +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 @@ -1430,14 +1498,21 @@ fn body_model_name( .unwrap_or_default() } -fn append_scan_text(out: &mut String, text: &str) { +fn append_scan_text(out: &mut String, text: &str) -> Option<()> { if text.is_empty() { - return; + 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; } - if !out.is_empty() { + 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 { @@ -1446,6 +1521,12 @@ fn decoded_json_string_values(body: &[u8]) -> Option { .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)) +} + +#[cfg(test)] fn decoded_json_string_values_where( body: &[u8], include: impl FnMut(&[crate::json_splice::PathSeg]) -> bool, @@ -1471,6 +1552,7 @@ fn is_root_key(path: &[crate::json_splice::PathSeg], key: &str) -> bool { /// A detected envelope still forwards raw bytes, including duplicate keys and /// arbitrary nested fields. Scan every decoded string the upstream can read; /// the root `model` alone is routing metadata rather than caller content. +#[cfg(test)] fn decoded_non_model_json_string_values(body: &[u8]) -> Option { decoded_json_string_values_where(body, |path| !is_root_key(path, "model")) } @@ -1507,7 +1589,7 @@ fn decoded_json_string_values_except_root_keys( 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), + Ok(values) => append_scan_text(&mut out, &values)?, Err(error) => { if scan_error.is_none() { *scan_error = Some(error); @@ -1536,6 +1618,30 @@ impl<'a> RawJson<'a> { } } +// 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; + +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(()) +} + fn raw_skip_ws(bytes: &[u8], pos: &mut usize) { while bytes .get(*pos) @@ -1620,9 +1726,7 @@ fn raw_value_end(bytes: &[u8], start: usize) -> Option { if token == b"true" || token == b"false" || token == b"null" { Some(pos) } else { - serde_json::from_slice::(token) - .ok() - .map(|_| pos) + crate::json_splice::is_json_number(token).then_some(pos) } } } @@ -1680,12 +1784,14 @@ fn raw_object_members<'a>(body: &'a [u8], mut visit: impl FnMut(&str, RawJson<'a /// 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 key == wanted_key { - values.push(value); + if within_cap && key == wanted_key { + within_cap = raw_selector_push(&mut values, &mut source_bytes, value).is_some(); } })?; - Some(values) + within_cap.then_some(values) } /// Retaining a source fragment is a pointer copy, so nested carrier walks do @@ -1698,12 +1804,14 @@ fn raw_top_level_value_refs<'a>(body: &'a [u8], wanted_key: &str) -> Option(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 !excluded.contains(&key) { - values.push(value); + if within_cap && !excluded.contains(&key) { + within_cap = raw_selector_push(&mut values, &mut source_bytes, value).is_some(); } })?; - Some(values) + within_cap.then_some(values) } /// Match the last source occurrence, the same duplicate-key convention a @@ -1711,13 +1819,18 @@ fn raw_top_level_values_except<'a>(body: &'a [u8], excluded: &[&str]) -> Option< /// `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 { - raw_top_level_values(body, key) - .and_then(|values| values.into_iter().last()) - .is_some_and(|value| match value.get().trim_start().as_bytes().first() { - Some(b'[') => true, - Some(b'"') => allow_string, - _ => false, - }) + 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, + }) } fn raw_is_object(raw: &RawJson<'_>) -> bool { @@ -1732,6 +1845,7 @@ fn raw_array_items<'a>(raw: &RawJson<'a>) -> Option>> { (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']') { @@ -1741,9 +1855,13 @@ fn raw_array_items<'a>(raw: &RawJson<'a>) -> Option>> { } let value_start = pos; let value_end = raw_value_end(bytes, value_start)?; - values.push(RawJson { - source: &source[value_start..value_end], - }); + 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) { @@ -1779,15 +1897,17 @@ fn raw_object_has_only_types(raw: &RawJson<'_>, allowed: &[&str]) -> bool { && types.all(|kind| kind.as_deref() == Some(first.as_str())) } -fn raw_top_level_unique_type(body: &[u8]) -> Option { +fn raw_top_level_unique_type(body: &[u8]) -> Result, ()> { let mut types = raw_top_level_values(body, "type") + .ok_or(())? .into_iter() - .flatten() .map(|value| serde_json::from_str::(value.get()).ok()); - let first = types.next()??; - types + let Some(Some(first)) = types.next() else { + return Ok(None); + }; + Ok(types .all(|kind| kind.as_deref() == Some(first.as_str())) - .then_some(first) + .then_some(first)) } /// `true` only when every source `type` value is one of `allowed`. This is @@ -1816,7 +1936,7 @@ fn raw_top_level_items_have_only_types(body: &[u8], key: &str, allowed: &[&str]) } fn append_raw_string_value(out: &mut String, raw: &RawJson<'_>) -> Option<()> { - append_scan_text(out, &serde_json::from_str::(raw.get()).ok()?); + append_scan_text(out, &serde_json::from_str::(raw.get()).ok()?)?; Some(()) } @@ -1826,17 +1946,40 @@ fn append_raw_top_level_strings(out: &mut String, body: &[u8], key: &str) -> Opt // 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); + append_scan_text(out, &value)?; } } Some(()) } +fn append_typed_top_level_strings( + out: &mut String, + body: &[u8], + 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)?; + } + Some(()) +} + fn raw_top_level_string_values(body: &[u8], key: &str) -> Option> { - raw_top_level_values(body, key)? - .into_iter() - .map(|value| serde_json::from_str::(value.get()).ok()) - .collect() + 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) } fn raw_top_level_unique_string(body: &[u8], key: &str) -> Result, ()> { @@ -1894,7 +2037,12 @@ fn raw_top_level_unique_array<'a>(body: &'a [u8], key: &str) -> Result) -> Option<()> { +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); @@ -1904,9 +2052,25 @@ fn append_raw_text_value(out: &mut String, raw: &RawJson<'_>) -> Option<()> { } for part in raw_array_items(raw)? { if !raw_is_object(&part) { + if strict_content_parts { + mark_unevaluable(scan_error); + return None; + } continue; } - append_raw_top_level_strings(out, part.get().as_bytes(), "text")?; + 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(()) } @@ -1931,11 +2095,22 @@ fn append_chat_request_content_strings( // 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![Work::Content { - value: *content, - depth: 0, - }]; - while let Some(work_item) = work.pop() { + let mut work = Vec::new(); + 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(); @@ -1960,10 +2135,26 @@ fn append_chat_request_content_strings( return None; }; for block in blocks.into_iter().rev() { - work.push(Work::Block { - value: block, - depth, - }); + 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 { @@ -1979,12 +2170,18 @@ fn append_chat_request_content_strings( mark_unevaluable(scan_error); return None; }; - let kind = raw_top_level_unique_type(block_body); + 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() { @@ -2008,10 +2205,26 @@ fn append_chat_request_content_strings( return None; }; for nested in nested.into_iter().rev() { - work.push(Work::Content { - value: nested, - depth, - }); + 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") => { @@ -2026,10 +2239,15 @@ fn append_chat_request_content_strings( input.get().as_bytes(), scan_error, )?, - ); + )?; } } - Some("thinking") => append_raw_top_level_strings(out, block_body, "thinking")?, + 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")?, } } @@ -2055,7 +2273,7 @@ fn append_chat_request_message_strings( &["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; @@ -2071,7 +2289,7 @@ fn append_chat_request_message_strings( 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(()) @@ -2118,7 +2336,8 @@ fn append_responses_item_strings( scan_error: &mut Option, ) -> Option<()> { if !raw_is_object(item) { - return Some(()); + mark_unevaluable(scan_error); + return None; } let item_body = item.get().as_bytes(); let text_keys = [ @@ -2133,10 +2352,10 @@ fn append_responses_item_strings( 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)?; + append_raw_text_value(out, &value, key == "content", scan_error)?; } } Some(()) @@ -2156,6 +2375,9 @@ fn decoded_responses_request_string_values( 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) @@ -2212,7 +2434,13 @@ fn request_guardrail_text_with_scan_error( } } PassthroughProtocol::OpenaiCompletions => { - decoded_completions_request_string_values(body, scan_error).unwrap_or_else(raw) + 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) { @@ -2238,7 +2466,7 @@ fn try_request_guardrail_text( body: &[u8], ) -> Result { if matches!(protocol, PassthroughProtocol::Raw) { - return crate::json_splice::collect_string_values(body); + 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); @@ -2274,6 +2502,7 @@ fn append_chat_visible_content_strings(out: &mut String, content: &RawJson<'_>) } 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) { @@ -2288,7 +2517,7 @@ fn append_chat_visible_content_strings(out: &mut String, content: &RawJson<'_>) /// 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).as_deref() { + 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")]), @@ -2370,7 +2599,7 @@ fn append_anthropic_output_content_strings( continue; } let block_body = block.get().as_bytes(); - match raw_top_level_unique_type(block_body).as_deref() { + 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")?; @@ -2381,7 +2610,7 @@ fn append_anthropic_output_content_strings( input.get().as_bytes(), scan_error, )?, - ); + )?; } } Some(_) | None => {} @@ -2400,9 +2629,7 @@ fn decoded_chat_response_string_values( .iter() .any(|choices| choices.get().trim_start().starts_with('[')); for array in choices { - let Some(choices) = raw_array_items(&array) else { - continue; - }; + let choices = raw_array_items(&array)?; for choice in choices { if !raw_is_object(&choice) { continue; @@ -2412,7 +2639,7 @@ fn decoded_chat_response_string_values( } } } - if !has_chat_choices && raw_top_level_unique_type(body).as_deref() == Some("message") { + 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)?; } @@ -2437,7 +2664,7 @@ const RESPONSES_VISIBLE_TEXT_PART_TYPES: &[&str] = &["output_text", "text", "inp /// 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(); - match raw_top_level_unique_type(part_body).as_deref() { + match raw_top_level_unique_type(part_body).ok()?.as_deref() { Some(kind) if RESPONSES_VISIBLE_TEXT_PART_TYPES.contains(&kind) => { append_raw_top_level_strings(out, part_body, "text")? } @@ -2454,9 +2681,7 @@ fn append_responses_visible_content_strings(out: &mut String, content: &RawJson< if !value.starts_with('[') { return Some(()); } - let Some(parts) = raw_array_items(content) else { - return Some(()); - }; + let parts = raw_array_items(content)?; for part in parts { append_responses_visible_part_strings(out, &part)?; } @@ -2470,7 +2695,7 @@ fn append_responses_visible_content_strings(out: &mut String, content: &RawJson< /// `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).as_deref() { + 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")? { @@ -2494,9 +2719,7 @@ fn append_responses_output_item_strings(out: &mut String, item: &RawJson<'_>) -> fn append_responses_output_strings(out: &mut String, body: &[u8]) -> Option<()> { for output in raw_top_level_values(body, "output")? { - let Some(items) = raw_array_items(&output) else { - continue; - }; + let items = raw_array_items(&output)?; for item in items { append_responses_output_item_strings(out, &item)?; } @@ -2536,7 +2759,21 @@ fn response_guardrail_text_with_scan_error( } } PassthroughProtocol::OpenaiCompletions => { - decoded_non_model_json_string_values(body).unwrap_or_else(raw) + // A detected completions response can also contain opaque + // provider fields. If its JSON cannot be selected safely, do + // not turn that failure into a raw whole-body guardrail scan. + match crate::json_splice::collect_string_values_where(body, |path| { + !is_root_key(path, "model") + }) { + Ok(text) => text, + Err(error) => { + if scan_error.is_none() { + *scan_error = Some(error); + } + mark_unevaluable(scan_error); + String::new() + } + } } PassthroughProtocol::OpenaiResponses => { // Without a safely decoded Responses envelope, no discriminator @@ -2559,16 +2796,11 @@ fn try_response_guardrail_text( body: &[u8], ) -> Result { match protocol { - PassthroughProtocol::Raw => crate::json_splice::collect_string_values(body), + PassthroughProtocol::Raw => try_raw_json_string_values(body), PassthroughProtocol::OpenaiCompletions => { - let text = crate::json_splice::collect_string_values_where(body, |path| { - !is_root_key(path, "model") - })?; - Ok(if text.is_empty() { - response_guardrail_text(protocol, body) - } else { - text - }) + 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; @@ -3211,7 +3443,7 @@ fn frame_delta(protocol: PassthroughProtocol, frame: &[u8]) -> (String, Option

Option { - match raw_top_level_unique_type(body).as_deref() { + match raw_top_level_unique_type(body).ok()?.as_deref() { Some("content_block_delta") => raw_top_level_items_have_only_types( body, "delta", @@ -3230,10 +3462,10 @@ fn hidden_chat_stream_reasoning_frame(body: &[u8]) -> Option { 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); + append_scan_text(&mut out, &value)?; } for value in decoded_chat_frame_supplemental_values(body)? { - append_scan_text(&mut out, &value); + append_scan_text(&mut out, &value)?; } Some(out) } @@ -3252,10 +3484,13 @@ fn decoded_responses_frame_string_values(body: &[u8]) -> Option { } fn decoded_chat_frame_continuations(body: &[u8]) -> Option> { - match raw_top_level_unique_type(body).as_deref() { + 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()).as_deref() { + 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") @@ -3266,7 +3501,7 @@ fn decoded_chat_frame_continuations(body: &[u8]) -> Option> { 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).as_deref() { + 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(); @@ -3312,6 +3547,7 @@ fn decoded_chat_frame_continuations(body: &[u8]) -> Option> { } 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) { @@ -3358,18 +3594,6 @@ fn decoded_chat_frame_continuations(body: &[u8]) -> Option> { } } -fn decoded_completions_frame_continuations(body: &[u8]) -> Option> { - use crate::json_splice::PathSeg; - - decoded_json_string_values_vec_where(body, |path| { - matches!( - path, - [PathSeg::Key(choices), PathSeg::Index(_), PathSeg::Key(text)] - if choices == "choices" && text == "text" - ) - }) -} - fn is_completions_continuation_path(path: &[crate::json_splice::PathSeg]) -> bool { use crate::json_splice::PathSeg; @@ -3385,12 +3609,15 @@ fn is_completions_continuation_path(path: &[crate::json_splice::PathSeg]) -> boo /// 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).as_deref() == Some("content_block_delta") { + 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).as_deref() == Some("content_block_start") { + 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()).as_deref() { + 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()), }; @@ -3628,7 +3855,11 @@ fn anthropic_source_continuations( _ => return SourceContinuations::Unevaluable, }; let delta_body = delta.get().as_bytes(); - match raw_top_level_unique_type(delta_body).as_deref() { + 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, @@ -3662,7 +3893,11 @@ fn anthropic_source_continuations( _ => return SourceContinuations::Unevaluable, }; let block_body = block.get().as_bytes(); - match raw_top_level_unique_type(block_body).as_deref() { + 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, @@ -3806,9 +4041,11 @@ fn chat_choice_source_continuations(payload: &[u8]) -> SourceContinuations { return SourceContinuations::Unevaluable; } let part_body = part.get().as_bytes(); - let Some(field) = raw_top_level_unique_type(part_body) - .as_deref() - .and_then(chat_visible_content_part_field) + 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; }; @@ -3936,8 +4173,12 @@ fn chat_choice_source_continuations(payload: &[u8]) -> SourceContinuations { } fn completions_source_continuations(payload: &[u8]) -> SourceContinuations { - let Some(expected) = decoded_completions_frame_continuations(payload) else { - return SourceContinuations::Absent; + let expected = match crate::json_splice::collect_string_values_where_vec(payload, |path| { + is_completions_continuation_path(path) + }) { + Ok(expected) if expected.is_empty() => return SourceContinuations::Absent, + Ok(expected) => expected, + Err(_) => return SourceContinuations::Unevaluable, }; let choices = match raw_top_level_unique_array(payload, "choices") { Ok(Some(choices)) => match raw_array_items(&choices) { @@ -3995,17 +4236,33 @@ fn stream_source_continuations( // 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 => match serde_json::from_slice::(payload) { - Ok(text) if !text.is_empty() => SourceContinuations::Ready(vec![StreamContinuation { - key: "raw:payload:first".to_owned(), - family: "raw".to_owned(), - identity: "payload".to_owned(), - identity_is_ambiguous: false, - text, - }]), - Ok(_) => SourceContinuations::Absent, - Err(_) => SourceContinuations::Unevaluable, - }, + 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), @@ -4058,35 +4315,35 @@ fn frame_guardrail_supplemental_values( frame: &[u8], has_source_continuations: bool, has_typed_continuation: bool, -) -> Vec { +) -> Result, ()> { // 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 Vec::new(); + return Ok(Vec::new()); } if !has_source_continuations { - return frame_guardrail_values(protocol, frame); + return Ok(frame_guardrail_values(protocol, frame)); } let Some(payload) = crate::redact::frame_payload(frame) else { - return Vec::new(); + return Ok(Vec::new()); }; let payload = payload.trim(); if payload.is_empty() || payload == "[DONE]" { - return Vec::new(); + return Ok(Vec::new()); } match protocol { // The raw source continuation is the complete decoded payload. - PassthroughProtocol::Raw => Vec::new(), + PassthroughProtocol::Raw => Ok(Vec::new()), PassthroughProtocol::OpenaiChat => { - decoded_chat_frame_supplemental_values(payload.as_bytes()).unwrap_or_default() + decoded_chat_frame_supplemental_values(payload.as_bytes()).ok_or(()) } PassthroughProtocol::OpenaiCompletions => { - decoded_completions_frame_supplemental_values(payload.as_bytes()).unwrap_or_default() + decoded_completions_frame_supplemental_values(payload.as_bytes()).ok_or(()) } // Responses source continuations exist only for the explicitly safe // text/tool delta events, whose sole output carrier is `delta`. - PassthroughProtocol::OpenaiResponses => Vec::new(), + PassthroughProtocol::OpenaiResponses => Ok(Vec::new()), } } @@ -4109,10 +4366,10 @@ fn frame_guardrail_values(protocol: PassthroughProtocol, frame: &[u8]) -> Vec { - decoded_json_string_values_vec_where(payload.as_bytes(), |_| true).unwrap_or_else(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() @@ -4121,7 +4378,7 @@ fn frame_guardrail_values(protocol: PassthroughProtocol, frame: &[u8]) -> Vec { decoded_responses_frame_values(payload.as_bytes()).unwrap_or_default() @@ -4142,10 +4399,9 @@ fn frame_guardrail_text(protocol: PassthroughProtocol, frame: &[u8]) -> String { if payload.is_empty() || payload == "[DONE]" { return String::new(); } - let raw = || payload.to_string(); match protocol { PassthroughProtocol::Raw => { - decoded_json_string_values(payload.as_bytes()).unwrap_or_else(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. @@ -4154,7 +4410,7 @@ fn frame_guardrail_text(protocol: PassthroughProtocol, frame: &[u8]) -> String { decoded_chat_frame_string_values(payload.as_bytes()).unwrap_or_default() } PassthroughProtocol::OpenaiCompletions => { - decoded_non_model_json_string_values(payload.as_bytes()).unwrap_or_else(raw) + decoded_non_model_json_string_values(payload.as_bytes()).unwrap_or_default() } PassthroughProtocol::OpenaiResponses => { // As with buffered Responses output, only a successful @@ -4251,19 +4507,30 @@ fn stream_guardrail_text( // caller applies the configured fail-open/fail-closed policy. SourceContinuations::Unevaluable => (Vec::new(), false, true), }; - let unevaluable = source_unevaluable || terminal_unevaluable; + let mut unevaluable = source_unevaluable || terminal_unevaluable; + 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; + Vec::new() + } + } + }; StreamGuardrailText { continuations, - supplemental: if unevaluable { - Vec::new() - } else { - frame_guardrail_supplemental_values( - protocol, - frame, - has_source_continuations, - has_typed_continuation, - ) - }, + supplemental, unevaluable, closed_prefixes, } @@ -4659,8 +4926,84 @@ fn anthropic_stream_frame(frame: &[u8]) -> Option { } } -/// Build the streamed relay response: upstream SSE frames are forwarded -/// incrementally, tee'd through the chain's [`StreamOutputPolicy`] +/// Relay a non-success upstream SSE error without parsing or mutating its +/// frames. The telemetry guard still fires from `Drop` when the client +/// disconnects mid-relay. +#[allow(clippy::too_many_arguments)] +fn stream_non_success_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 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(), + )); + while let Some(chunk) = upstream.next().await { + 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 non-success 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 non-success 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 @@ -4818,6 +5161,41 @@ fn stream_response( 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 = (!chain.is_empty() && !scan_budget_exhausted && !fail_opened) .then(|| stream_guardrail_text(protocol, &frame, delta.clone())); let unevaluable_output = guardrail_text.as_ref().is_some_and(|text| { @@ -7062,19 +7440,20 @@ mod tests { } #[test] - fn duplicate_chat_carrier_skips_malformed_source_without_leaking_media() { + 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 scanned = request_guardrail_text(PassthroughProtocol::OpenaiChat, request); - assert!(scanned.contains("clean"), "{scanned:?}"); + let error = try_request_guardrail_text(PassthroughProtocol::OpenaiChat, request) + .expect_err("a duplicate malformed carrier must fail closed"); + assert!(error.is_unevaluable(), "{error}"); assert!( - !scanned.contains("BLOCKME"), - "a malformed duplicate must not trigger raw fallback: {scanned:?}" + !request_guardrail_text(PassthroughProtocol::OpenaiChat, request).contains("BLOCKME"), + "a malformed duplicate must not trigger raw fallback" ); } #[test] - fn malformed_typed_array_items_keep_media_opaque() { + fn malformed_typed_array_items_are_unevaluable_without_leaking_media() { let cases = [ ( PassthroughProtocol::OpenaiChat, @@ -7098,11 +7477,12 @@ mod tests { ), ]; for (protocol, request) in cases { - let scanned = request_guardrail_text(protocol, request); - assert!(scanned.contains("clean"), "{protocol:?}: {scanned:?}"); + 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!( - !scanned.contains("BLOCKME"), - "a malformed array item must not trigger raw fallback for {protocol:?}: {scanned:?}" + !request_guardrail_text(protocol, request).contains("BLOCKME"), + "a malformed array item must not trigger raw fallback for {protocol:?}" ); } } @@ -7232,6 +7612,55 @@ mod tests { ); } + #[test] + fn malformed_completions_output_is_unevaluable_without_raw_fallback() { + let buffered = br#"{"choices":[{"text":"clean"}],"opaque":{"data":"BLOCKME"}"#; + let error = try_response_guardrail_text(PassthroughProtocol::OpenaiCompletions, buffered) + .expect_err("a malformed completions response must fail closed"); + assert!(error.is_unevaluable(), "{error}"); + assert!( + !response_guardrail_text(PassthroughProtocol::OpenaiCompletions, buffered) + .contains("BLOCKME"), + "a malformed completions response must not fall back to opaque source text" + ); + + let streamed = b"data: {\"choices\":[{\"index\":0,\"text\":\"clean\"}],\"opaque\":{\"data\":\"BLOCKME\"}\n\n"; + let typed = frame_parts(PassthroughProtocol::OpenaiCompletions, streamed) + .0 + .scan; + assert!( + stream_guardrail_text(PassthroughProtocol::OpenaiCompletions, streamed, typed) + .unevaluable + ); + assert!( + !frame_guardrail_text(PassthroughProtocol::OpenaiCompletions, streamed) + .contains("BLOCKME"), + "a malformed completions stream frame must not fall back to opaque source text" + ); + } + + #[test] + fn streamed_completions_supplemental_cap_is_unevaluable_without_raw_fallback() { + let frame = format!( + "data: {{\"choices\":[{{\"index\":0,\"text\":\"clean\"}}],\"opaque\":\"{}\"}}\n\n", + "BLOCKME".repeat(MAX_RAW_SELECTOR_BYTES / "BLOCKME".len() + 1), + ); + 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, "BLOCKME"), + "a capped supplemental selector must not raw-fallback opaque output: {scanned:?}" + ); + } + #[test] fn chat_output_guardrail_keeps_media_and_unknown_parts_opaque() { let buffered = br#"{ @@ -8892,6 +9321,147 @@ mod tests { } } + #[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] 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 9fdbe48e4..a02a3b809 100644 --- a/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts @@ -45,6 +45,8 @@ 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 deepEscapedBlockJSON = (depth: number) => `${'{"v":'.repeat(depth)}"${String.raw`\u0042LOCKME`}"${'}'.repeat(depth)}`; const deepLiteralBlockJSON = (depth: number) => @@ -179,6 +181,23 @@ describe("passthrough guardrail scan coverage", () => { 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["raw-output"] = await startOpenAiUpstream({ rawBody: ESCAPED_BLOCK_JSON, rawContentType: "application/json", @@ -330,7 +349,46 @@ 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: 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: generated thinking is not scanned", async (ctx) => { @@ -570,6 +628,19 @@ describe("passthrough guardrail scan coverage", () => { }, ); + 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], From d8a40cccae0bd779a28de103b3819dff1412b63b Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Thu, 1 Oct 2026 15:37:56 +0800 Subject: [PATCH 30/37] fix: harden passthrough guardrail boundaries --- crates/aisix-gateway/src/upstream_tls.rs | 110 ++++- crates/aisix-proxy/src/http_client.rs | 19 +- crates/aisix-proxy/src/passthrough_route.rs | 444 ++++++++++++++---- ...ssthrough-chat-media-guardrail-e2e.test.ts | 115 +++++ .../passthrough-scan-coverage-e2e.test.ts | 266 +++++++++++ .../src/cases/ratelimit-cluster-e2e.test.ts | 149 ++++++ tests/e2e/src/harness/upstream-openai.ts | 35 +- 7 files changed, 1021 insertions(+), 117 deletions(-) diff --git a/crates/aisix-gateway/src/upstream_tls.rs b/crates/aisix-gateway/src/upstream_tls.rs index 38f69984a..bfb667316 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,49 @@ 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() +} + +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 +348,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 +409,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 +417,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 +890,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 +970,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/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/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index 1db46470d..2d164a0dc 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -403,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 // --------------------------------------------------------------------------- @@ -572,6 +612,7 @@ 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. Keep a @@ -698,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_non_success_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 = @@ -811,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; @@ -820,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 @@ -908,15 +964,10 @@ 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 mut telemetry = RouteTelemetry { state: state.clone(), @@ -963,6 +1014,29 @@ 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 or reach the caller. + if success_response_is_encoded { + 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(); @@ -1549,14 +1623,6 @@ fn is_root_key(path: &[crate::json_splice::PathSeg], key: &str) -> bool { path.first().is_some_and(|segment| segment.is_key(key)) } -/// A detected envelope still forwards raw bytes, including duplicate keys and -/// arbitrary nested fields. Scan every decoded string the upstream can read; -/// the root `model` alone is routing metadata rather than caller content. -#[cfg(test)] -fn decoded_non_model_json_string_values(body: &[u8]) -> Option { - decoded_json_string_values_where(body, |path| !is_root_key(path, "model")) -} - fn decoded_json_string_values_including_empty( body: &[u8], scan_error: &mut Option, @@ -2619,6 +2685,53 @@ fn append_anthropic_output_content_strings( 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(body, "choices") { + Ok(Some(choices)) => choices, + Ok(None) | Err(()) => { + mark_unevaluable(scan_error); + return None; + } + }; + let choices = match raw_array_items(&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; + } + } + } + Some(out) +} + fn decoded_chat_response_string_values( body: &[u8], scan_error: &mut Option, @@ -2759,17 +2872,13 @@ fn response_guardrail_text_with_scan_error( } } PassthroughProtocol::OpenaiCompletions => { - // A detected completions response can also contain opaque - // provider fields. If its JSON cannot be selected safely, do - // not turn that failure into a raw whole-body guardrail scan. - match crate::json_splice::collect_string_values_where(body, |path| { - !is_root_key(path, "model") - }) { - Ok(text) => text, - Err(error) => { - if scan_error.is_none() { - *scan_error = Some(error); - } + // 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() } @@ -3594,16 +3703,6 @@ fn decoded_chat_frame_continuations(body: &[u8]) -> Option> { } } -fn is_completions_continuation_path(path: &[crate::json_splice::PathSeg]) -> bool { - use crate::json_splice::PathSeg; - - matches!( - path, - [PathSeg::Key(choices), PathSeg::Index(_), PathSeg::Key(text)] - if choices == "choices" && text == "text" - ) -} - /// 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, @@ -3667,12 +3766,6 @@ fn decoded_chat_frame_supplemental_values(body: &[u8]) -> Option> { Some(out) } -fn decoded_completions_frame_supplemental_values(body: &[u8]) -> Option> { - decoded_json_string_values_vec_where(body, |path| { - !is_root_key(path, "model") && !is_completions_continuation_path(path) - }) -} - #[derive(Clone, Debug, PartialEq, Eq)] struct StreamContinuation { key: String, @@ -4173,13 +4266,6 @@ fn chat_choice_source_continuations(payload: &[u8]) -> SourceContinuations { } fn completions_source_continuations(payload: &[u8]) -> SourceContinuations { - let expected = match crate::json_splice::collect_string_values_where_vec(payload, |path| { - is_completions_continuation_path(path) - }) { - Ok(expected) if expected.is_empty() => return SourceContinuations::Absent, - Ok(expected) => expected, - Err(_) => return SourceContinuations::Unevaluable, - }; let choices = match raw_top_level_unique_array(payload, "choices") { Ok(Some(choices)) => match raw_array_items(&choices) { Some(choices) => choices, @@ -4200,27 +4286,29 @@ fn completions_source_continuations(payload: &[u8]) -> SourceContinuations { Ok(Some(index)) if choice_indexes.insert(index) => index.to_string(), _ => return SourceContinuations::Unevaluable, }; - if append_raw_string_carrier( + 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, - SourceBranchIdentity { - family: format!("completions:{choice_index}"), - identity: "text".to_owned(), - identity_is_ambiguous: false, - }, - choice_body, - "text", + format!("completions:{choice_index}"), + "text".to_owned(), + false, + text, ) .is_err() { return SourceContinuations::Unevaluable; } } - if source_values_match_expected(source_values, expected) { - SourceContinuations::Ready(out) + if out.is_empty() { + SourceContinuations::Absent } else { - SourceContinuations::Unevaluable + SourceContinuations::Ready(out) } } @@ -4316,6 +4404,13 @@ fn frame_guardrail_supplemental_values( 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 { @@ -4338,9 +4433,7 @@ fn frame_guardrail_supplemental_values( PassthroughProtocol::OpenaiChat => { decoded_chat_frame_supplemental_values(payload.as_bytes()).ok_or(()) } - PassthroughProtocol::OpenaiCompletions => { - decoded_completions_frame_supplemental_values(payload.as_bytes()).ok_or(()) - } + PassthroughProtocol::OpenaiCompletions => Ok(Vec::new()), // Responses source continuations exist only for the explicitly safe // text/tool delta events, whose sole output carrier is `delta`. PassthroughProtocol::OpenaiResponses => Ok(Vec::new()), @@ -4374,12 +4467,7 @@ fn frame_guardrail_values(protocol: PassthroughProtocol, frame: &[u8]) -> Vec { decoded_chat_frame_values(payload.as_bytes()).unwrap_or_default() } - PassthroughProtocol::OpenaiCompletions => { - decoded_json_string_values_vec_where(payload.as_bytes(), |path| { - !is_root_key(path, "model") - }) - .unwrap_or_default() - } + PassthroughProtocol::OpenaiCompletions => Vec::new(), PassthroughProtocol::OpenaiResponses => { decoded_responses_frame_values(payload.as_bytes()).unwrap_or_default() } @@ -4410,7 +4498,9 @@ fn frame_guardrail_text(protocol: PassthroughProtocol, frame: &[u8]) -> String { decoded_chat_frame_string_values(payload.as_bytes()).unwrap_or_default() } PassthroughProtocol::OpenaiCompletions => { - decoded_non_model_json_string_values(payload.as_bytes()).unwrap_or_default() + 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 @@ -4431,6 +4521,10 @@ struct StreamGuardrailText { /// 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. @@ -4508,6 +4602,7 @@ fn stream_guardrail_text( 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 { @@ -4524,6 +4619,7 @@ fn stream_guardrail_text( Ok(values) => values, Err(()) => { unevaluable = true; + supplemental_unevaluable = true; Vec::new() } } @@ -4532,6 +4628,7 @@ fn stream_guardrail_text( continuations, supplemental, unevaluable, + supplemental_unevaluable, closed_prefixes, } } @@ -5213,11 +5310,15 @@ fn stream_response( ) }); 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 = (policy.holds_back() && !fail_opened) + let must_refuse = supplemental_unevaluable + || (policy.holds_back() && !fail_opened) || aisix_guardrails::Guardrail::refuses_unevaluable_output(&chain); if must_refuse { tracing::warn!( @@ -5562,10 +5663,14 @@ fn stream_response( ) }); 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 = (policy.holds_back() && !fail_opened) + let must_refuse = supplemental_unevaluable + || (policy.holds_back() && !fail_opened) || aisix_guardrails::Guardrail::refuses_unevaluable_output(&chain); if must_refuse { tracing::warn!( @@ -7613,37 +7718,50 @@ mod tests { } #[test] - fn malformed_completions_output_is_unevaluable_without_raw_fallback() { - let buffered = br#"{"choices":[{"text":"clean"}],"opaque":{"data":"BLOCKME"}"#; - let error = try_response_guardrail_text(PassthroughProtocol::OpenaiCompletions, buffered) - .expect_err("a malformed completions response must fail closed"); - assert!(error.is_unevaluable(), "{error}"); - assert!( - !response_guardrail_text(PassthroughProtocol::OpenaiCompletions, buffered) - .contains("BLOCKME"), - "a malformed completions response must not fall back to opaque source text" - ); + 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" + ); + } - let streamed = b"data: {\"choices\":[{\"index\":0,\"text\":\"clean\"}],\"opaque\":{\"data\":\"BLOCKME\"}\n\n"; - let typed = frame_parts(PassthroughProtocol::OpenaiCompletions, streamed) - .0 - .scan; - assert!( - stream_guardrail_text(PassthroughProtocol::OpenaiCompletions, streamed, typed) - .unevaluable - ); - assert!( - !frame_guardrail_text(PassthroughProtocol::OpenaiCompletions, streamed) - .contains("BLOCKME"), - "a malformed completions stream frame must not fall back to opaque source text" + 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()); } #[test] - fn streamed_completions_supplemental_cap_is_unevaluable_without_raw_fallback() { + fn streamed_completions_select_only_choice_text_and_reject_bad_carriers() { let frame = format!( - "data: {{\"choices\":[{{\"index\":0,\"text\":\"clean\"}}],\"opaque\":\"{}\"}}\n\n", + "data: {{\"choices\":[{{\"index\":0,\"text\":\"clean\",\"opaque\":\"{}\"}}],\"opaque\":\"{}\"}}\n\n", "BLOCKME".repeat(MAX_RAW_SELECTOR_BYTES / "BLOCKME".len() + 1), + "BLOCKME", ); let typed = frame_parts(PassthroughProtocol::OpenaiCompletions, frame.as_bytes()) .0 @@ -7653,12 +7771,69 @@ mod tests { frame.as_bytes(), typed, ); - assert!(text.unevaluable); + assert!(!text.unevaluable); let scanned = stream_guardrail_scan_text(&[], &text.continuations, &text.supplemental); + assert!(scan_candidates_contain(&scanned, "clean"), "{scanned:?}"); assert!( !scan_candidates_contain(&scanned, "BLOCKME"), - "a capped supplemental selector must not raw-fallback opaque output: {scanned:?}" + "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 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] @@ -8618,6 +8793,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 @@ -9538,6 +9770,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\ @@ -9550,6 +9783,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`. @@ -9583,4 +9818,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"], "content_filter", "{v}"); + assert_eq!(v["error"]["code"], "guardrail_unavailable", "{v}"); + } } 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 index e4d2a59bb..762f99bdb 100644 --- a/tests/e2e/src/cases/passthrough-chat-media-guardrail-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-chat-media-guardrail-e2e.test.ts @@ -39,6 +39,8 @@ 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; @@ -237,6 +239,24 @@ const streamedOversizedLegacyToolResponse = [ "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; @@ -248,6 +268,9 @@ describe("Chat passthrough keeps media out of external output guardrails", () => 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; @@ -275,6 +298,15 @@ describe("Chat passthrough keeps media out of external output guardrails", () => 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); @@ -331,6 +363,24 @@ describe("Chat passthrough keeps media out of external output guardrails", () => 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, @@ -361,6 +411,9 @@ describe("Chat passthrough keeps media out of external output guardrails", () => await streamOversizedRefusalUpstream?.close(); await streamOversizedContentRefusalUpstream?.close(); await streamOversizedLegacyToolUpstream?.close(); + await bufferedCompletionsUpstream?.close(); + await malformedBufferedCompletionsUpstream?.close(); + await malformedStreamedCompletionsUpstream?.close(); await moderation?.close(); }); @@ -378,6 +431,16 @@ describe("Chat passthrough keeps media out of external output guardrails", () => }), }); + 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, @@ -476,6 +539,58 @@ describe("Chat passthrough keeps media out of external output guardrails", () => ); }); + test("buffered Completions 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 completionsRequest("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 || 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 a02a3b809..414934ba0 100644 --- a/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts @@ -1,4 +1,6 @@ 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, @@ -32,6 +34,10 @@ import { startMockOtlp, type MockOtlp } from "../harness/otlp-mock.js"; 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"; @@ -47,6 +53,17 @@ 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) => @@ -162,6 +179,27 @@ 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; @@ -198,6 +236,39 @@ describe("passthrough guardrail scan coverage", () => { "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", @@ -299,6 +370,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); @@ -322,6 +399,11 @@ describe("passthrough guardrail scan coverage", () => { 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" }; @@ -391,6 +473,112 @@ describe("passthrough guardrail scan coverage", () => { 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); @@ -902,6 +1090,84 @@ describe("passthrough guardrail scan coverage", () => { }); }); +// 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 output stream. diff --git a/tests/e2e/src/cases/ratelimit-cluster-e2e.test.ts b/tests/e2e/src/cases/ratelimit-cluster-e2e.test.ts index f7ace415f..0459c9303 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 { @@ -43,6 +44,10 @@ const PASSTHROUGH_PREFIX = "/rl-cluster-passthrough"; // 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 ETCD_ENDPOINT = etcdEndpoint(); const REDIS_URL = process.env.AISIX_E2E_REDIS ?? "redis://127.0.0.1:6379"; @@ -124,6 +129,27 @@ 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))); + }); +} + /** 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) { @@ -352,6 +378,129 @@ describe("passthrough SSE concurrency is shared and renewed across Redis replica ); }); +// 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; + 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" }, + }, + ], + }); + 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, + }); + await seed.createApiKey({ + key_hash: callerHash, + allowed_models: ["*"], + allowed_routes: [route], + rate_limit: { concurrency: 1 }, + }); + 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); + }, + 15_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; From bf348ac0f93e8acb4beecbc3f5636768d6d1d2c1 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Thu, 1 Oct 2026 15:50:33 +0800 Subject: [PATCH 31/37] test: complete passthrough guardrail fixtures --- crates/aisix-proxy/src/passthrough_route.rs | 8 +++----- 1 file changed, 3 insertions(+), 5 deletions(-) diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index 2d164a0dc..08f4b7385 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -1619,10 +1619,6 @@ fn decoded_json_string_values_vec_where( .filter(|out| !out.is_empty()) } -fn is_root_key(path: &[crate::json_splice::PathSeg], key: &str) -> bool { - path.first().is_some_and(|segment| segment.is_key(key)) -} - fn decoded_json_string_values_including_empty( body: &[u8], scan_error: &mut Option, @@ -7814,7 +7810,7 @@ mod tests { 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" + "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()); @@ -8283,6 +8279,7 @@ mod tests { }], supplemental: Vec::new(), unevaluable: false, + supplemental_unevaluable: false, closed_prefixes: Vec::new(), }; assert!(stream_continuation_would_exceed_cap( @@ -8304,6 +8301,7 @@ mod tests { .map(|index| format!("metadata-{index}")) .collect(), unevaluable: false, + supplemental_unevaluable: false, closed_prefixes: Vec::new(), }; assert!(stream_continuation_would_exceed_cap( From 2cfa5db971835246b41f80693409ee48060ec7cb Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Thu, 1 Oct 2026 15:53:59 +0800 Subject: [PATCH 32/37] test: remove obsolete passthrough selector helper --- crates/aisix-proxy/src/passthrough_route.rs | 10 ---------- 1 file changed, 10 deletions(-) diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index 08f4b7385..c904c0a02 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -1600,16 +1600,6 @@ fn try_raw_json_string_values(body: &[u8]) -> Result bool, -) -> Option { - crate::json_splice::collect_string_values_where(body, include) - .ok() - .filter(|out| !out.is_empty()) -} - fn decoded_json_string_values_vec_where( body: &[u8], include: impl FnMut(&[crate::json_splice::PathSeg]) -> bool, From 2b674e87474178048696b2dd51101c4a402175ea Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Thu, 1 Oct 2026 17:22:02 +0800 Subject: [PATCH 33/37] fix: harden passthrough stream boundaries --- crates/aisix-gateway/src/upstream_tls.rs | 7 +- crates/aisix-proxy/src/held_content.rs | 787 +++++++++++++++++- crates/aisix-proxy/src/passthrough_route.rs | 731 ++++++++++++---- crates/aisix-proxy/src/redact.rs | 126 ++- crates/aisix-proxy/src/responses.rs | 374 ++++++++- .../passthrough-redirect-boundary-e2e.test.ts | 148 ++++ .../passthrough-scan-coverage-e2e.test.ts | 366 +++++++- 7 files changed, 2292 insertions(+), 247 deletions(-) create mode 100644 tests/e2e/src/cases/passthrough-redirect-boundary-e2e.test.ts diff --git a/crates/aisix-gateway/src/upstream_tls.rs b/crates/aisix-gateway/src/upstream_tls.rs index bfb667316..c547854be 100644 --- a/crates/aisix-gateway/src/upstream_tls.rs +++ b/crates/aisix-gateway/src/upstream_tls.rs @@ -289,7 +289,12 @@ fn raw_client_builder() -> reqwest::ClientBuilder { // 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() + 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 { diff --git a/crates/aisix-proxy/src/held_content.rs b/crates/aisix-proxy/src/held_content.rs index dad389974..f329b4bcf 100644 --- a/crates/aisix-proxy/src/held_content.rs +++ b/crates/aisix-proxy/src/held_content.rs @@ -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 @@ -220,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. @@ -248,21 +767,118 @@ 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 @@ -351,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::*; @@ -412,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/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index c904c0a02..d70e0e3d1 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -742,7 +742,7 @@ async fn dispatch( // 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_non_success_response` relays it. + // 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. @@ -968,6 +968,8 @@ async fn dispatch( 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(), @@ -1018,33 +1020,44 @@ async fn dispatch( // `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 or reach the caller. + // 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 { - 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 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(); - // A non-success SSE response is an upstream error contract, just - // like a buffered 4xx/5xx. It must remain byte-for-byte relay data: - // no output guardrail, heartbeat, or SSE parsing may rewrite it. - if !status.is_success() { - return Ok(stream_non_success_response( + // 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, @@ -1094,7 +1107,7 @@ async fn dispatch( // 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() && !resolved_chain.is_empty() { + if status.is_success() && !bypass_uninspectable_output && !resolved_chain.is_empty() { let text = match try_response_guardrail_text(protocol, &resp_body) { Ok(text) => Some(text), Err(err) if !err.is_unevaluable() => { @@ -1176,7 +1189,7 @@ async fn dispatch( if let Some(u) = response_usage(protocol, raw_shape, &resp_body) { merge_usage(&mut telemetry.usage, u); } - if telemetry.content_cap.is_some() { + if telemetry.content_cap.is_some() && !bypass_uninspectable_output { telemetry.response_text = response_capture_text(protocol, &resp_body); } @@ -1675,6 +1688,11 @@ impl<'a> RawJson<'a> { // 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>, @@ -1694,6 +1712,30 @@ fn raw_selector_push<'a>( Some(()) } +/// 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(()) +} + +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 +} + fn raw_skip_ws(bytes: &[u8], pos: &mut usize) { while bytes .get(*pos) @@ -1849,7 +1891,14 @@ fn raw_top_level_values<'a>(body: &'a [u8], wanted_key: &str) -> Option(body: &'a [u8], wanted_key: &str) -> Option>> { - raw_top_level_values(body, wanted_key) + 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 @@ -1929,7 +1978,40 @@ fn raw_array_items<'a>(raw: &RawJson<'a>) -> Option>> { } fn raw_array_item_refs<'a>(raw: &RawJson<'a>) -> Option>> { - raw_array_items(raw) + 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 @@ -1962,21 +2044,6 @@ fn raw_top_level_unique_type(body: &[u8]) -> Result, ()> { .then_some(first)) } -/// `true` only when every source `type` value is one of `allowed`. This is -/// deliberately strict: an audio or image event must not borrow a text -/// event's carrier merely by repeating a conflicting type. -fn raw_top_level_has_only_types(body: &[u8], allowed: &[&str]) -> bool { - let Some(values) = raw_top_level_values(body, "type") else { - return false; - }; - !values.is_empty() - && values.into_iter().all(|value| { - serde_json::from_str::(value.get()) - .ok() - .is_some_and(|kind| allowed.iter().any(|allowed| kind == *allowed)) - }) -} - fn raw_top_level_items_have_only_types(body: &[u8], key: &str, allowed: &[&str]) -> Option { let values = raw_top_level_values(body, key)?; Some( @@ -2086,6 +2153,30 @@ fn raw_top_level_unique_array<'a>(body: &'a [u8], key: &str) -> Result( + 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 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. @@ -2148,6 +2239,7 @@ fn append_chat_request_content_strings( // 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); @@ -2182,6 +2274,11 @@ fn append_chat_request_content_strings( } // 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; @@ -2218,10 +2315,20 @@ fn append_chat_request_content_strings( 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(()) => { @@ -2239,6 +2346,13 @@ fn append_chat_request_content_strings( 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; @@ -2280,6 +2394,13 @@ fn append_chat_request_content_strings( } } 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; @@ -2684,14 +2805,14 @@ fn decoded_completions_response_string_values( body: &[u8], scan_error: &mut Option, ) -> Option { - let choices = match raw_top_level_unique_array(body, "choices") { + 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_items(&choices) { + let choices = match raw_array_item_refs(&choices) { Some(choices) => choices, None => { mark_unevaluable(scan_error); @@ -2746,16 +2867,79 @@ fn decoded_chat_response_string_values( Some(out) } -const RESPONSES_VISIBLE_DELTA_EVENTS: &[&str] = &[ - "response.output_text.delta", - "response.function_call_arguments.delta", - "response.mcp_call_arguments.delta", - "response.custom_tool_call_input.delta", -]; +/// 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, +} + +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, + } +} + +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 `text` fields the typed output guardrail -/// reads. Other part types can carry image, audio, file, or reasoning data. -const RESPONSES_VISIBLE_TEXT_PART_TYPES: &[&str] = &["output_text", "text", "input_text"]; +/// 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, + } +} /// Source-preserving counterpart to the typed Responses output scanner's /// content-part walk. A missing or conflicting discriminator is opaque: a @@ -2763,11 +2947,12 @@ const RESPONSES_VISIBLE_TEXT_PART_TYPES: &[&str] = &["output_text", "text", "inp /// 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(); - match raw_top_level_unique_type(part_body).ok()?.as_deref() { - Some(kind) if RESPONSES_VISIBLE_TEXT_PART_TYPES.contains(&kind) => { - append_raw_top_level_strings(out, part_body, "text")? - } - Some(_) | None => {} + 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(()) } @@ -3565,15 +3750,14 @@ fn decoded_chat_frame_string_values(body: &[u8]) -> Option { Some(out) } -/// The typed stream extractor deliberately takes only text/tool delta events: -/// `.done`, output-item, and terminal response snapshots repeat those -/// carriers. Keep the raw source pass on that same boundary, both to avoid -/// duplicate external moderation and to keep any opaque terminal media 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(); - if raw_top_level_has_only_types(body, RESPONSES_VISIBLE_DELTA_EVENTS) { - append_raw_top_level_strings(&mut out, body, "delta")?; + for value in responses_stream_selected_values(body).ok()? { + append_scan_text(&mut out, &value)?; } Some(out) } @@ -3823,60 +4007,209 @@ fn source_values_match_expected(mut source_values: Vec, mut expected: Ve source_values == expected } -fn responses_source_continuations(payload: &[u8]) -> SourceContinuations { - let kind = match raw_top_level_unique_string(payload, "type") { - Ok(Some(kind)) if RESPONSES_VISIBLE_DELTA_EVENTS.contains(&kind.as_str()) => kind, - Ok(Some(_)) | Ok(None) => return SourceContinuations::Absent, - Err(()) => return SourceContinuations::Unevaluable, - }; - let item_id = match raw_top_level_unique_string(payload, "item_id") { - Ok(Some(item_id)) => match bounded_stream_source_id(item_id) { - Ok(item_id) => item_id, - Err(()) => return SourceContinuations::Unevaluable, - }, - Ok(None) | Err(()) => return SourceContinuations::Unevaluable, - }; - let output_index = match raw_top_level_unique_index(payload, "output_index") { - Ok(Some(index)) => index.to_string(), - Ok(None) | Err(()) => return SourceContinuations::Unevaluable, - }; - let content_index = if kind == "response.output_text.delta" { - match raw_top_level_unique_index(payload, "content_index") { - Ok(Some(index)) => index.to_string(), - Ok(None) | Err(()) => return SourceContinuations::Unevaluable, - } +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 { - // Tool-argument deltas have no content part. Their stable identity is - // the item plus output index and event kind; an incidental content - // field must not create a second source channel. String::new() }; - let values = match raw_top_level_string_values(payload, "delta") { - Some(values) => values, - None => return SourceContinuations::Unevaluable, + 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 values.is_empty() { - return SourceContinuations::Absent; + 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); } - let mut out = Vec::new(); - let mut keys = std::collections::HashSet::new(); + 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(); - let family = format!("responses:{item_id:?}"); - let identity = format!("{kind:?}:{output_index}:{content_index}:delta"); - if append_source_branches( - &mut out, - &mut keys, + append_source_branches( + out, + keys, &mut source_values, family, identity, false, values, ) - .is_err() - { - return SourceContinuations::Unevaluable; +} + +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) } - SourceContinuations::Ready(out) } struct SourceBranchIdentity { @@ -4252,8 +4585,8 @@ fn chat_choice_source_continuations(payload: &[u8]) -> SourceContinuations { } fn completions_source_continuations(payload: &[u8]) -> SourceContinuations { - let choices = match raw_top_level_unique_array(payload, "choices") { - Ok(Some(choices)) => match raw_array_items(&choices) { + 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, }, @@ -4420,8 +4753,8 @@ fn frame_guardrail_supplemental_values( decoded_chat_frame_supplemental_values(payload.as_bytes()).ok_or(()) } PassthroughProtocol::OpenaiCompletions => Ok(Vec::new()), - // Responses source continuations exist only for the explicitly safe - // text/tool delta events, whose sole output carrier is `delta`. + // Responses source continuations cover every type-aware visible + // carrier, so no generic raw field is supplemental. PassthroughProtocol::OpenaiResponses => Ok(Vec::new()), } } @@ -4431,10 +4764,7 @@ fn decoded_chat_frame_values(body: &[u8]) -> Option> { } fn decoded_responses_frame_values(body: &[u8]) -> Option> { - raw_top_level_has_only_types(body, RESPONSES_VISIBLE_DELTA_EVENTS) - .then(|| raw_top_level_string_values(body, "delta")) - .flatten() - .or_else(|| Some(Vec::new())) + responses_stream_selected_values(body).ok() } fn frame_guardrail_values(protocol: PassthroughProtocol, frame: &[u8]) -> Vec { @@ -4527,12 +4857,15 @@ fn stream_guardrail_text( && payload.as_ref().is_some_and(|payload| { hidden_chat_stream_reasoning_frame(payload.trim().as_bytes()) == Some(true) }); - let responses_visible_delta = matches!(protocol, PassthroughProtocol::OpenaiResponses) + let responses_visible_carrier = matches!(protocol, PassthroughProtocol::OpenaiResponses) && payload.as_ref().is_some_and(|payload| { - raw_top_level_has_only_types(payload.trim().as_bytes(), RESPONSES_VISIBLE_DELTA_EVENTS) + 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_delta) + || (matches!(protocol, PassthroughProtocol::OpenaiResponses) && !responses_visible_carrier) { String::new() } else if matches!(protocol, PassthroughProtocol::OpenaiChat) { @@ -4632,6 +4965,7 @@ fn append_stream_guardrail_text( // 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) { @@ -4688,19 +5022,25 @@ fn stream_continuation_identity_conflicts( closed_prefixes: &[String], text: &StreamGuardrailText, ) -> bool { - let is_closed = |key: &str| { - closed_prefixes - .iter() - .chain(text.closed_prefixes.iter()) - .any(|prefix| key.starts_with(prefix)) - }; + // 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| is_closed(&continuation.key)) + .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)) @@ -5009,11 +5349,12 @@ fn anthropic_stream_frame(frame: &[u8]) -> Option { } } -/// Relay a non-success upstream SSE error without parsing or mutating its -/// frames. The telemetry guard still fires from `Drop` when the client -/// disconnects mid-relay. +/// 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_non_success_response( +fn stream_opaque_response( upstream_resp: reqwest::Response, resp_headers: HeaderMap, status: reqwest::StatusCode, @@ -5052,7 +5393,7 @@ fn stream_non_success_response( tracing::warn!( route = %route_name, error = %telemetry.error_message, - "passthrough-route non-success SSE relay failed mid-stream", + "passthrough-route opaque SSE relay failed mid-stream", ); break; } @@ -5063,7 +5404,7 @@ fn stream_non_success_response( tracing::warn!( route = %route_name, error = %telemetry.error_message, - "passthrough-route non-success SSE relay timed out mid-stream", + "passthrough-route opaque SSE relay timed out mid-stream", ); } telemetry.stream_reached_end = true; @@ -5137,6 +5478,11 @@ fn stream_response( // 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. @@ -5232,7 +5578,16 @@ fn stream_response( telemetry.record_failure(&err); } let (parts, usage) = frame_parts(protocol, &frame); - let held = parts.held(); + 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); @@ -5574,7 +5929,16 @@ fn stream_response( anthropic = anthropic_stream_frame(&rest); } let (parts, usage) = frame_parts(protocol, &rest); - let held = parts.held(); + 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); @@ -6212,6 +6576,7 @@ fn copy_safe_headers(src: &HeaderMap, dst: &mut HeaderMap) { | "proxy-authenticate" | "proxy-authorization" | "te" + | "trailer" | "trailers" | "upgrade" ) { @@ -7680,6 +8045,11 @@ mod tests { } 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) @@ -8009,15 +8379,13 @@ mod tests { "{scanned:?}" ); - // The terminal response repeats prior delta content, including media - // from image-generation output. It is never a second scan carrier. + // 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:?}" - ); + 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 @@ -8094,6 +8462,11 @@ mod tests { } assert!(!scanned.contains("NESTED"), "{scanned:?}"); + 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:?}"); + 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) @@ -8219,6 +8592,20 @@ mod tests { 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\"}}" @@ -8685,6 +9072,55 @@ mod tests { ) .unevaluable ); + + // 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, + ); + assert!(!standalone.unevaluable); + assert_eq!( + standalone.closed_prefixes, + vec!["responses:\"standalone\":".to_owned()] + ); + 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, + ); + assert!(standalone_continuations.is_empty()); + assert!(standalone_tails.is_empty()); + assert!(standalone_closures.is_empty()); } #[test] @@ -9284,10 +9720,15 @@ 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] @@ -9486,21 +9927,19 @@ mod tests { } #[test] - fn nested_anthropic_tool_result_content_at_depth_cap_is_scanned() { - let body = - nested_anthropic_tool_result_request(crate::json_splice::MAX_JSON_DEPTH, "BLOCKME"); + 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 at the shared JSON depth cap remain evaluable"); + .expect("nested tool results within the structural-work cap remain evaluable"); assert!(text.contains("BLOCKME"), "{text}"); } #[test] - fn nested_anthropic_tool_result_content_beyond_depth_cap_is_unevaluable() { - let body = - nested_anthropic_tool_result_request(crate::json_splice::MAX_JSON_DEPTH + 1, "safe"); + 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 shared JSON depth cap must not recurse"); - assert!(error.is_depth_exceeded(), "{error}"); + .expect_err("nested tool results beyond the structural-work cap must not recurse"); + assert!(error.is_unevaluable(), "{error}"); } #[test] @@ -9810,7 +10249,7 @@ data: [DONE]\n\n"; #[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"], "content_filter", "{v}"); + 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..2469b3af0 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,13 @@ 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 dir.is_output() => 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 +1147,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 +2026,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 +2038,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 +2074,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 +2082,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 +2126,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 +2252,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 +2290,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 +3241,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 +4047,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 +4059,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..435ff6cf9 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/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-scan-coverage-e2e.test.ts b/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts index 414934ba0..58ab7105b 100644 --- a/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts @@ -22,8 +22,8 @@ import { startMockOtlp, type MockOtlp } from "../harness/otlp-mock.js"; // 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. // @@ -85,19 +85,20 @@ 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( - JSON_DEPTH_CAP, - "nested-tool-result-at-depth-cap", + NESTED_TOOL_RESULT_WORK_SAFE_DEPTH, + "nested-tool-result-within-work-cap", ); const OVER_DEPTH_ANTHROPIC_TOOL_RESULT_INPUT = deeplyNestedAnthropicToolResultRequest( - JSON_DEPTH_CAP + 1, - "nested-tool-result-fail-closed", + JSON_DEPTH_CAP, + "nested-tool-result-work-cap-fail-closed", ); const OVER_DEPTH_ANTHROPIC_TOOL_RESULT_FAIL_OPEN_INPUT = deeplyNestedAnthropicToolResultRequest( - JSON_DEPTH_CAP + 1, - "nested-tool-result-fail-open", + 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"}]}]}`; @@ -105,16 +106,25 @@ const OVER_DEPTH_CHAT_OPAQUE_INPUT = `{"model":"gpt-4o-mini","messages":[{"role" 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 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","item":{"id":"same","type":"message"}}\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 = [ @@ -124,9 +134,9 @@ const DISTINCT_ITEMS_RESPONSES_STREAM = [ ]; 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", part: { type: "reasoning_text", text: OUT_LIT } })}\n\n`, - `data: ${JSON.stringify({ type: "response.output_item.done", 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: 0, content_index: 0, delta: "clean" })}\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", ]; @@ -170,6 +180,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": [ @@ -310,6 +373,10 @@ describe("passthrough guardrail scan coverage", () => { rawBody: KNOWN_RESPONSES_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], }); @@ -438,6 +505,120 @@ describe("passthrough guardrail scan coverage", () => { 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"]; @@ -653,6 +834,9 @@ describe("passthrough guardrail scan coverage", () => { 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); }); @@ -862,7 +1046,7 @@ describe("passthrough guardrail scan coverage", () => { expect(upstreams.input!.receivedRequests.length).toBe(before); }); - test("input: nested Anthropic tool results at the scanner depth cap are blocked", async (ctx) => { + 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); @@ -875,7 +1059,7 @@ describe("passthrough guardrail scan coverage", () => { expect(upstreams.input!.receivedRequests.length).toBe(before); }); - test("input: nested Anthropic tool results beyond the scanner depth cap fail closed", async (ctx) => { + 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); @@ -1090,6 +1274,158 @@ describe("passthrough guardrail scan coverage", () => { }); }); +// 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. @@ -1296,7 +1632,7 @@ describe("passthrough Raw stream unevaluable-output fail-open", () => { expect(log.get("guardrail_bypassed_reason")).toBe("unscannable_body"); }); - test("forwards nested Anthropic tool results beyond the depth cap only under input fail_open", async (ctx) => { + 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; From 3bc5ac2b9c58c2db6423f82fd14bf6a0c39dcdc2 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Thu, 1 Oct 2026 17:29:39 +0800 Subject: [PATCH 34/37] fix: restore passthrough guardrail compilation --- crates/aisix-proxy/src/passthrough_route.rs | 6 +++--- crates/aisix-proxy/src/redact.rs | 6 +++++- 2 files changed, 8 insertions(+), 4 deletions(-) diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index d70e0e3d1..25c155347 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -4015,11 +4015,11 @@ fn responses_stream_coordinates( .ok_or(()) .and_then(bounded_stream_source_id)?; let output_index = raw_top_level_unique_index(payload, "output_index")? - .ok_or(()) + .ok_or(())? .to_string(); let content_index = if needs_content_index { raw_top_level_unique_index(payload, "content_index")? - .ok_or(()) + .ok_or(())? .to_string() } else { String::new() @@ -4032,7 +4032,7 @@ fn responses_output_item_coordinates( item: &RawJson<'_>, ) -> Result<(String, String), ()> { let output_index = raw_top_level_unique_index(payload, "output_index")? - .ok_or(()) + .ok_or(())? .to_string(); let nested_id = raw_top_level_unique_string(item.get().as_bytes(), "id")? .ok_or(()) diff --git a/crates/aisix-proxy/src/redact.rs b/crates/aisix-proxy/src/redact.rs index 2469b3af0..e143c8bf6 100644 --- a/crates/aisix-proxy/src/redact.rs +++ b/crates/aisix-proxy/src/redact.rs @@ -1004,7 +1004,11 @@ fn redact_responses_item( for part in parts { let field = match part.get("type").and_then(Value::as_str) { Some("input_text" | "output_text" | "text") => Some("text"), - Some("refusal") if dir.is_output() => Some("refusal"), + Some("refusal") + if matches!(dir, Direction::Output | Direction::OutputEcho) => + { + Some("refusal") + } _ => None, }; if let Some(field) = field.and_then(|field| part.get_mut(field)) { From 0225a9abb5dbb3bf8c142751567714d06f11dae2 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Thu, 1 Oct 2026 17:42:03 +0800 Subject: [PATCH 35/37] fix: satisfy responses guardrail lint --- crates/aisix-proxy/src/responses.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/crates/aisix-proxy/src/responses.rs b/crates/aisix-proxy/src/responses.rs index 435ff6cf9..d29eae555 100644 --- a/crates/aisix-proxy/src/responses.rs +++ b/crates/aisix-proxy/src/responses.rs @@ -3734,7 +3734,7 @@ fn responses_sse_guardrail_text(bytes: &[u8]) -> String { if event_text.is_empty() { continue; } - if !text.is_empty() && !(previous_was_delta && is_delta) { + if !(text.is_empty() || previous_was_delta && is_delta) { text.push('\n'); } text.push_str(&event_text); From 26f291d932532cabc406ff8dea736a5a7ababdb9 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Thu, 1 Oct 2026 19:07:10 +0800 Subject: [PATCH 36/37] fix: enforce passthrough stream lease boundaries --- crates/aisix-proxy/src/passthrough_route.rs | 327 +++++++++++++++++- crates/aisix-ratelimit/Cargo.toml | 5 + crates/aisix-ratelimit/src/lib.rs | 2 +- crates/aisix-ratelimit/src/limiter.rs | 203 ++++++++++- crates/aisix-ratelimit/src/store/mod.rs | 24 ++ crates/aisix-ratelimit/src/store/redis.rs | 105 +++++- .../tests/redis_integration.rs | 249 ++++++++++++- ...ssthrough-chat-media-guardrail-e2e.test.ts | 4 +- .../src/cases/passthrough-route-e2e.test.ts | 1 + .../passthrough-scan-coverage-e2e.test.ts | 79 ++++- .../src/cases/ratelimit-cluster-e2e.test.ts | 164 ++++++++- 11 files changed, 1113 insertions(+), 50 deletions(-) diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index 25c155347..1e0b3d369 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -1107,12 +1107,9 @@ async fn dispatch( // 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 && !resolved_chain.is_empty() { - let text = match try_response_guardrail_text(protocol, &resp_body) { + 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 !err.is_unevaluable() => { - Some(response_guardrail_text(protocol, &resp_body)) - } Err(err) if !aisix_guardrails::Guardrail::refuses_unevaluable_output(&resolved_chain) => { @@ -1413,9 +1410,10 @@ 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 { @@ -1482,6 +1480,142 @@ fn detect_protocol(body: &[u8]) -> PassthroughProtocol { } } +/// 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); + } + values + .iter() + .all(|value| value.get().trim_start().starts_with('[')) + .then_some(values) + .ok_or(()) +} + +/// `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() @@ -3092,6 +3226,23 @@ fn try_response_guardrail_text( } } +/// 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)) + } + result => result, + } +} + /// 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 { @@ -5349,6 +5500,21 @@ fn anthropic_stream_frame(frame: &[u8]) -> Option { } } +/// 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 @@ -5366,6 +5532,7 @@ fn stream_opaque_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(); @@ -5374,7 +5541,20 @@ fn stream_opaque_response( stream_read_timeout, read_timeout.clone(), )); - while let Some(chunk) = upstream.next().await { + 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 { @@ -5447,7 +5627,8 @@ fn stream_response( use aisix_guardrails::{Guardrail as _, GuardrailVerdict, StreamOutputPolicy}; use futures::StreamExt; - let policy = if chain.is_empty() { + let output_guardrail_active = aisix_guardrails::Guardrail::runs_on_output(&chain); + let policy = if !output_guardrail_active { StreamOutputPolicy::EndOfStreamCheck } else { chain.stream_output_policy() @@ -5458,6 +5639,7 @@ fn stream_response( .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 @@ -5507,7 +5689,16 @@ fn stream_response( let mut anthropic: Option = None; 'outer: loop { - let chunk = match upstream.next().await { + 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 @@ -5634,7 +5825,7 @@ fn stream_response( } } } - let guardrail_text = (!chain.is_empty() && !scan_budget_exhausted && !fail_opened) + 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 @@ -5996,7 +6187,7 @@ fn stream_response( } None => { let guardrail_text = - (!chain.is_empty() && !scan_budget_exhausted && !fail_opened) + (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 @@ -6120,7 +6311,7 @@ fn stream_response( } } } - if !fail_opened { + if output_guardrail_active && !fail_opened { if !scan_budget_exhausted { let candidates = stream_guardrail_scan_text( &continuation_tails, @@ -6138,7 +6329,7 @@ fn stream_response( } } for candidates in sealed_guardrail_epochs { - if !chain.is_empty() { + if output_guardrail_active { if let GuardrailVerdict::Block { reason, guardrail_name, @@ -6377,6 +6568,20 @@ impl RouteTelemetry { 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. /// @@ -7746,6 +7951,98 @@ mod tests { assert!(text.contains("SECRET"), "tool-call output scanned: {text}"); } + #[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:?}" + ); + } + } + + #[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:?}"); + + 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}"); + } + + #[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 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:?}"); + } + + #[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"); + } + #[test] fn requested_model_comes_only_from_a_detected_envelope() { let chat = br#"{"model":"gpt-4o","messages":[{"role":"user","content":"hi"}]}"#; 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 419ff18d4..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 @@ -107,11 +107,16 @@ impl Limiter { 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), @@ -119,6 +124,7 @@ impl Limiter { member, has_concurrency_slot, refresh_task, + lease_loss, committed: false, }) } @@ -148,8 +154,8 @@ impl Limiter { /// 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 only after the SSE handoff lets a -/// second replica prune that still-live request before the handoff happens. +/// 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( @@ -157,6 +163,7 @@ fn spawn_lease_refresh( key: String, member: String, has_concurrency_slot: bool, + lease_loss: Option>, ) -> Option> { if !has_concurrency_slot { return None; @@ -166,7 +173,43 @@ fn spawn_lease_refresh( handle.spawn(async move { loop { tokio::time::sleep(interval).await; - store.refresh_stream_lease(&key, &member).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; + } } }) }) @@ -188,6 +231,10 @@ pub struct Reservation { /// 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, } @@ -274,6 +321,8 @@ impl MultiReservation { 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() @@ -282,22 +331,35 @@ impl MultiReservation { // slot now; the returned guard owns release from here on. r.committed = true; 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, } } @@ -322,10 +384,23 @@ pub struct StreamConcurrencyGuard { /// 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; @@ -334,6 +409,9 @@ impl StreamConcurrencyGuard { 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); } @@ -345,6 +423,7 @@ impl std::fmt::Debug for StreamConcurrencyGuard { 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() } @@ -360,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 { @@ -872,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 05aba0fce..44c20dcf8 100644 --- a/crates/aisix-ratelimit/src/store/mod.rs +++ b/crates/aisix-ratelimit/src/store/mod.rs @@ -47,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. @@ -140,6 +155,15 @@ pub trait RateStore: Send + Sync + 'static { /// 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 e8d8a8f16..bc89acd47 100644 --- a/crates/aisix-ratelimit/src/store/redis.rs +++ b/crates/aisix-ratelimit/src/store/redis.rs @@ -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}`. @@ -247,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 @@ -340,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, } @@ -553,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 { @@ -578,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; @@ -600,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, @@ -658,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); @@ -736,12 +782,34 @@ impl RateStore for RedisStore { } 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); - return; + // 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) @@ -752,10 +820,35 @@ impl RateStore for RedisStore { .invoke_async(&mut conn) .await; match res { - Ok(_) => self.mark_ok(), + 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 } } } diff --git a/crates/aisix-ratelimit/tests/redis_integration.rs b/crates/aisix-ratelimit/tests/redis_integration.rs index 06b6b1562..75acb82bc 100644 --- a/crates/aisix-ratelimit/tests/redis_integration.rs +++ b/crates/aisix-ratelimit/tests/redis_integration.rs @@ -13,6 +13,7 @@ use aisix_core::{RateLimit, RateLimitScope, RedisConnConfig, RedisMode}; use aisix_obs::metrics::Metrics; use aisix_ratelimit::{ store::redis::DEFAULT_PREFIX, Limiter, MultiReservation, RateStore, RedisStore, + StreamLeaseRefresh, }; fn redis_url() -> Option { @@ -417,10 +418,15 @@ async fn refreshing_released_member_does_not_recreate_slot() { a.acquire(&key, &limits, "a-stream") .await .expect("first stream allowed"); - // `commit` removes the member synchronously in Redis. A queued lease - // refresh that follows must see it missing rather than add it back. + // `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; - a.refresh_stream_lease(&key, "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 @@ -534,6 +540,64 @@ async fn stream_hold_renews_redis_lease_until_drop() { 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`. @@ -729,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 { @@ -746,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 { @@ -758,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]; @@ -794,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; } @@ -803,7 +894,12 @@ impl RedisCutoff { }); } }); - Self { port, cut, hole } + Self { + port, + cut, + hole, + fail_reply_after_apply, + } } fn url(&self) -> String { @@ -818,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); } } @@ -891,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/tests/e2e/src/cases/passthrough-chat-media-guardrail-e2e.test.ts b/tests/e2e/src/cases/passthrough-chat-media-guardrail-e2e.test.ts index 762f99bdb..5a4e14a70 100644 --- a/tests/e2e/src/cases/passthrough-chat-media-guardrail-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-chat-media-guardrail-e2e.test.ts @@ -539,14 +539,14 @@ describe("Chat passthrough keeps media out of external output guardrails", () => ); }); - test("buffered Completions sends only choices text to the external guardrail", async (ctx) => { + 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 completionsRequest("completions-buffered", false); + const response = await request("completions-buffered", false); const body = await response.text(); expect(response.status, body).toBe(200); expect(body).toContain(COMPLETIONS_VISIBLE); diff --git a/tests/e2e/src/cases/passthrough-route-e2e.test.ts b/tests/e2e/src/cases/passthrough-route-e2e.test.ts index 70216ae83..0c35a7b86 100644 --- a/tests/e2e/src/cases/passthrough-route-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-route-e2e.test.ts @@ -344,6 +344,7 @@ describe("passthrough-route e2e: explicit routes, BYO credentials, unclaimed pat "../models", "%2e%2e/models", "%2E%2E/models", + "%2E./models", "..\\models", ]) { const status = await rawHttpStatus( 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 58ab7105b..7b859285c 100644 --- a/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts +++ b/tests/e2e/src/cases/passthrough-scan-coverage-e2e.test.ts @@ -111,6 +111,7 @@ 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: [ { @@ -373,6 +374,10 @@ describe("passthrough guardrail scan coverage", () => { 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", @@ -770,16 +775,16 @@ describe("passthrough guardrail scan coverage", () => { }); test.for([ [ - "chat", + "Chat response to a Responses request", "known-chat-buffered-output", - `{"model":"gpt-4o-mini","messages":[{"role":"user","content":"go"}]}`, + `{"model":"gpt-4o-mini","input":"go"}`, ], [ - "Responses", + "Responses response to a Completions request", "known-responses-buffered-output", - `{"model":"gpt-4o-mini","input":"go"}`, + `{"model":"gpt-4o-mini","prompt":"go"}`, ], - ] as const)("output: known %s envelope is source-scanned when buffered", async ([, route, body], ctx) => { + ] 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`); @@ -793,6 +798,26 @@ describe("passthrough guardrail scan coverage", () => { 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", @@ -1506,17 +1531,19 @@ describe("passthrough supplemental selector fail-closed", () => { // 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 output stream. -describe("passthrough Raw stream unevaluable-output fail-open", () => { +// 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; @@ -1538,6 +1565,10 @@ describe("passthrough Raw stream unevaluable-output fail-open", () => { 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", @@ -1571,6 +1602,12 @@ describe("passthrough Raw stream unevaluable-output fail-open", () => { 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, @@ -1606,6 +1643,7 @@ describe("passthrough Raw stream unevaluable-output fail-open", () => { await app?.exit(); await upstream?.close(); await depthInputUpstream?.close(); + await unknownTypedUpstream?.close(); await sls?.close(); }); @@ -1632,6 +1670,33 @@ describe("passthrough Raw stream unevaluable-output fail-open", () => { 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(); diff --git a/tests/e2e/src/cases/ratelimit-cluster-e2e.test.ts b/tests/e2e/src/cases/ratelimit-cluster-e2e.test.ts index 0459c9303..fa1de105e 100644 --- a/tests/e2e/src/cases/ratelimit-cluster-e2e.test.ts +++ b/tests/e2e/src/cases/ratelimit-cluster-e2e.test.ts @@ -8,9 +8,12 @@ import { SeedClient, ProxyClient, spawnApp, + startMockSls, startOpenAiUpstream, awaitWindowHeadroom, waitConfigPropagation, + waitForSlsLog, + type MockSls, type OpenAiUpstream, type SpawnedApp, } from "../harness/index.js"; @@ -48,6 +51,9 @@ const PASSTHROUGH_502_SSE = 'event: upstream_error\ndata: {"message":"still-live // 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"; @@ -74,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); @@ -150,6 +155,25 @@ function readRawBody(response: IncomingMessage): Promise { }); } +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) { @@ -248,7 +272,9 @@ describe("passthrough SSE concurrency is shared and renewed across Redis replica 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 = { @@ -270,6 +296,7 @@ describe("passthrough SSE concurrency is shared and renewed across Redis replica infraReady = (await new EtcdClient().ping()) && (await redisPing(REDIS_URL)); if (!infraReady) return; + sls = await startMockSls(); const streamEvents = [ JSON.stringify({ choices: [{ delta: { content: "released" } }] }), "[DONE]", @@ -280,6 +307,11 @@ describe("passthrough SSE concurrency is shared and renewed across Redis replica // 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 = { @@ -290,8 +322,12 @@ describe("passthrough SSE concurrency is shared and renewed across Redis replica concurrency_ttl_secs: PASSTHROUGH_CONCURRENCY_TTL_SECS, }, }; - appA = await spawnApp({ extra }); - appB = await spawnApp({ extra }); + 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({ @@ -305,15 +341,26 @@ describe("passthrough SSE concurrency is shared and renewed across Redis replica 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. - await seed.createApiKey({ + 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); @@ -325,13 +372,14 @@ describe("passthrough SSE concurrency is shared and renewed across Redis replica 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) { + if (!infraReady || !appA || !appB || !upstream || !sls) { ctx.skip(); return; } @@ -373,8 +421,52 @@ describe("passthrough SSE concurrency is shared and renewed across Redis replica 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); }, - 15_000, + 30_000, ); }); @@ -386,6 +478,7 @@ describe("passthrough 502 SSE concurrency is shared through delayed EOF (#1737)" 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}`; @@ -415,6 +508,17 @@ describe("passthrough 502 SSE concurrency is shared through delayed EOF (#1737)" 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 = { @@ -440,12 +544,13 @@ describe("passthrough 502 SSE concurrency is shared through delayed EOF (#1737)" target_url: upstream.baseUrl, provider_key_id: providerKey.id, }); - await seed.createApiKey({ + 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); @@ -496,8 +601,45 @@ describe("passthrough 502 SSE concurrency is shared through delayed EOF (#1737)" 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); }, - 15_000, + 30_000, ); }); From 3140277c156402a2d6652846804fde45edaf4193 Mon Sep 17 00:00:00 2001 From: Ming Wen Date: Thu, 1 Oct 2026 19:18:07 +0800 Subject: [PATCH 37/37] fix: return optional response carrier values --- crates/aisix-proxy/src/passthrough_route.rs | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/crates/aisix-proxy/src/passthrough_route.rs b/crates/aisix-proxy/src/passthrough_route.rs index 1e0b3d369..d951eca8a 100644 --- a/crates/aisix-proxy/src/passthrough_route.rs +++ b/crates/aisix-proxy/src/passthrough_route.rs @@ -1559,11 +1559,14 @@ fn response_top_level_array_values<'a>( if values.is_empty() { return Ok(None); } - values + if values .iter() .all(|value| value.get().trim_start().starts_with('[')) - .then_some(values) - .ok_or(()) + { + Ok(Some(values)) + } else { + Err(()) + } } /// `choices[].message` and `choices[].text` are mutually exclusive response