diff --git a/src/lib.rs b/src/lib.rs index e9b751c5..a730ae80 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,6 +1,6 @@ use axum::{ Json, Router, - body::Bytes, + body::{Body, Bytes}, extract::{DefaultBodyLimit, Path as PathParam, Query, State}, http::{HeaderMap, HeaderValue, Method, StatusCode, Uri}, response::{Html, IntoResponse, Response}, @@ -2622,11 +2622,8 @@ async fn proxy_request_with_headers( let admitted_response_headers = gateway_mediation::admit_response_headers(response.headers()) .map_err(|message| format!("upstream response rejected: {message}"))?; - let bytes = response - .bytes() - .await - .map_err(|error| format!("upstream body read failed: {error}"))?; - let mut response = (status, bytes).into_response(); + let body = Body::from_stream(response.bytes_stream()); + let mut response = (status, body).into_response(); *response.headers_mut() = admitted_response_headers; Ok(response) } @@ -6962,7 +6959,13 @@ mod tests { let truncated_response = app_request(&app, empty_request(Method::GET, "/gateway/truncated")).await; - assert_eq!(truncated_response.status(), StatusCode::BAD_GATEWAY); + assert_eq!(truncated_response.status(), StatusCode::OK); + let truncated_body = + axum::body::to_bytes(truncated_response.into_body(), usize::MAX).await; + assert!( + truncated_body.is_err(), + "truncated upstream streaming must surface a downstream body error" + ); raw_task.join().unwrap(); upstream_task.abort(); diff --git a/tests/gateway_a_streaming_backpressure.rs b/tests/gateway_a_streaming_backpressure.rs new file mode 100644 index 00000000..39e5ae30 --- /dev/null +++ b/tests/gateway_a_streaming_backpressure.rs @@ -0,0 +1,317 @@ +//! Hostile buyer acceptance for #442: generic gateway streaming must preserve +//! downstream backpressure and remain usable under concurrent held streams. +//! +//! This file is test-only. Protected production source is intentionally left +//! unchanged until the serialized gateway writer lane is clear. + +use std::{ + convert::Infallible, + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, + time::Duration, +}; + +use axum::{ + Router, + body::{Body, Bytes}, + extract::State, + http::{Method, Request, StatusCode, header::CONTENT_TYPE}, + response::Response, + routing::any, +}; +use futures_util::{StreamExt, future::join_all, stream}; +use tokio::sync::Notify; +use tower::ServiceExt; +use waf_ids_ai_soc::{AppState, build_app}; + +const BACKPRESSURE_CHUNK_BYTES: usize = 64 * 1024; +const BACKPRESSURE_TOTAL_CHUNKS: usize = 1024; // 64 MiB if fully drained. +const MAX_UNPOLLED_CHUNKS: usize = 128; // 8 MiB bounded read-ahead witness. +const CONCURRENT_SLOW_REQUESTS: usize = 4; + +#[derive(Clone, Default)] +struct BackpressureProbe { + request_seen: Arc, + chunks_emitted: Arc, +} + +async fn backpressure_upstream(State(probe): State) -> Response { + probe.request_seen.notify_one(); + + let body_stream = stream::unfold(0usize, move |emitted| { + let probe = probe.clone(); + async move { + if emitted >= BACKPRESSURE_TOTAL_CHUNKS { + return None; + } + + probe.chunks_emitted.fetch_add(1, Ordering::Release); + Some(( + Ok::(Bytes::from(vec![0x5A; BACKPRESSURE_CHUNK_BYTES])), + emitted + 1, + )) + } + }); + + Response::builder() + .status(StatusCode::OK) + .header(CONTENT_TYPE, "application/octet-stream") + .body(Body::from_stream(body_stream)) + .expect("valid loopback backpressure response") +} + +async fn gateway_with_backpressure_upstream() +-> (Router, BackpressureProbe, tokio::task::JoinHandle<()>) { + let probe = BackpressureProbe::default(); + let upstream_app = Router::new() + .route("/v1/backpressure", any(backpressure_upstream)) + .with_state(probe.clone()); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("loopback upstream listener"); + let upstream_addr = listener.local_addr().expect("loopback upstream address"); + let upstream_task = tokio::spawn(async move { + axum::serve(listener, upstream_app) + .await + .expect("loopback upstream must serve until test cleanup"); + }); + + let app = build_app(AppState::seeded(Some("secret".to_string()))); + let route = serde_json::json!({ + "id": "streaming-backpressure-red", + "path_prefix": "/stream", + "upstream": format!("http://{upstream_addr}"), + "mode": "monitor", + "enabled": true, + "block_threshold": null + }); + let created = app + .clone() + .oneshot( + Request::builder() + .method(Method::POST) + .uri("/api/routes") + .header(CONTENT_TYPE, "application/json") + .header("x-admin-token", "secret") + .body(Body::from(route.to_string())) + .expect("valid route registration request"), + ) + .await + .expect("Wardnet must answer route registration"); + assert_eq!(created.status(), StatusCode::CREATED); + + (app, probe, upstream_task) +} + +#[tokio::test] +async fn unpolled_buyer_body_does_not_drain_the_entire_large_upstream() { + assert!( + MAX_UNPOLLED_CHUNKS < BACKPRESSURE_TOTAL_CHUNKS, + "the bounded read-ahead witness must be smaller than the hostile body" + ); + + let (app, probe, upstream_task) = gateway_with_backpressure_upstream().await; + let request = Request::builder() + .method(Method::GET) + .uri("/gateway/stream/v1/backpressure") + .body(Body::empty()) + .expect("valid buyer request"); + + let mut gateway_task = tokio::spawn(async move { app.oneshot(request).await }); + tokio::time::timeout(Duration::from_secs(2), probe.request_seen.notified()) + .await + .expect("loopback upstream must receive the request before the backpressure assertion"); + + let response = match tokio::time::timeout(Duration::from_secs(2), &mut gateway_task).await { + Ok(joined) => joined + .expect("gateway task must not panic") + .expect("gateway service must answer"), + Err(_) => { + upstream_task.abort(); + panic!( + "Wardnet did not expose a downstream response while a finite 64 MiB upstream was being drained; the relay must return a streaming body instead of waiting for whole-body materialization" + ); + } + }; + assert_eq!(response.status(), StatusCode::OK); + + let mut downstream = response.into_body().into_data_stream(); + tokio::time::sleep(Duration::from_millis(250)).await; + let emitted_while_buyer_idle = probe.chunks_emitted.load(Ordering::Acquire); + assert!( + emitted_while_buyer_idle <= MAX_UNPOLLED_CHUNKS, + "an unpolled buyer body caused Wardnet/upstream read-ahead of {emitted_while_buyer_idle} x 64 KiB chunks; bounded relay backpressure must not drain the 64 MiB body in the background" + ); + + let first = tokio::time::timeout(Duration::from_secs(2), downstream.next()) + .await + .expect("polling the buyer body must make one admitted chunk available") + .expect("the downstream stream must produce data") + .expect("the first admitted chunk must not be an error"); + assert!(!first.is_empty()); + + drop(downstream); + upstream_task.abort(); +} + +#[derive(Clone, Default)] +struct ConcurrentProbe { + slow_requests_seen: Arc, + release_tails: Arc, + fast_request_seen: Arc, +} + +async fn concurrent_slow_upstream(State(probe): State) -> Response { + probe.slow_requests_seen.fetch_add(1, Ordering::Release); + + let release_tails = probe.release_tails.clone(); + let body_stream = stream::unfold(0u8, move |stage| { + let release_tails = release_tails.clone(); + async move { + match stage { + 0 => Some((Ok::(Bytes::from_static(b"first-")), 1)), + 1 => { + release_tails.notified().await; + Some((Ok(Bytes::from_static(b"tail")), 2)) + } + _ => None, + } + } + }); + + Response::builder() + .status(StatusCode::OK) + .header(CONTENT_TYPE, "application/octet-stream") + .body(Body::from_stream(body_stream)) + .expect("valid held loopback response") +} + +async fn concurrent_fast_upstream(State(probe): State) -> Response { + probe.fast_request_seen.notify_one(); + Response::builder() + .status(StatusCode::OK) + .header(CONTENT_TYPE, "text/plain") + .body(Body::from("fast")) + .expect("valid fast loopback response") +} + +async fn gateway_with_concurrent_upstream() -> (Router, ConcurrentProbe, tokio::task::JoinHandle<()>) +{ + let probe = ConcurrentProbe::default(); + let upstream_app = Router::new() + .route("/v1/slow", any(concurrent_slow_upstream)) + .route("/v1/fast", any(concurrent_fast_upstream)) + .with_state(probe.clone()); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("loopback upstream listener"); + let upstream_addr = listener.local_addr().expect("loopback upstream address"); + let upstream_task = tokio::spawn(async move { + axum::serve(listener, upstream_app) + .await + .expect("loopback upstream must serve until test cleanup"); + }); + + let app = build_app(AppState::seeded(Some("secret".to_string()))); + let route = serde_json::json!({ + "id": "streaming-concurrency-red", + "path_prefix": "/stream", + "upstream": format!("http://{upstream_addr}"), + "mode": "monitor", + "enabled": true, + "block_threshold": null + }); + let created = app + .clone() + .oneshot( + Request::builder() + .method(Method::POST) + .uri("/api/routes") + .header(CONTENT_TYPE, "application/json") + .header("x-admin-token", "secret") + .body(Body::from(route.to_string())) + .expect("valid route registration request"), + ) + .await + .expect("Wardnet must answer route registration"); + assert_eq!(created.status(), StatusCode::CREATED); + + (app, probe, upstream_task) +} + +#[tokio::test] +async fn concurrent_held_streams_expose_prefixes_and_do_not_block_fast_buyer_traffic() { + let (app, probe, upstream_task) = gateway_with_concurrent_upstream().await; + + let mut slow_tasks = Vec::with_capacity(CONCURRENT_SLOW_REQUESTS); + for _ in 0..CONCURRENT_SLOW_REQUESTS { + let request = Request::builder() + .method(Method::GET) + .uri("/gateway/stream/v1/slow") + .body(Body::empty()) + .expect("valid slow buyer request"); + let app = app.clone(); + slow_tasks.push(tokio::spawn(async move { app.oneshot(request).await })); + } + + tokio::time::timeout(Duration::from_secs(2), async { + while probe.slow_requests_seen.load(Ordering::Acquire) < CONCURRENT_SLOW_REQUESTS { + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("all concurrent slow requests must reach the loopback upstream"); + + let joined = match tokio::time::timeout(Duration::from_millis(500), join_all(slow_tasks)).await + { + Ok(joined) => joined, + Err(_) => { + probe.release_tails.notify_waiters(); + upstream_task.abort(); + panic!( + "concurrent held upstreams kept buyer response heads behind their unreleased tails; Wardnet must expose each admitted stream without whole-body buffering" + ); + } + }; + + let mut slow_bodies = Vec::with_capacity(CONCURRENT_SLOW_REQUESTS); + for joined in joined { + let response = joined + .expect("slow gateway task must not panic") + .expect("slow gateway request must produce a response"); + assert_eq!(response.status(), StatusCode::OK); + slow_bodies.push(response.into_body().into_data_stream()); + } + + for body in &mut slow_bodies { + let first = tokio::time::timeout(Duration::from_millis(500), body.next()) + .await + .expect("each concurrent buyer must receive its admitted prefix promptly") + .expect("each held stream must produce a prefix") + .expect("each admitted prefix must not be an error"); + assert_eq!(first, Bytes::from_static(b"first-")); + } + + let fast_request = Request::builder() + .method(Method::GET) + .uri("/gateway/stream/v1/fast") + .body(Body::empty()) + .expect("valid fast buyer request"); + let fast_response = tokio::time::timeout( + Duration::from_millis(500), + app.clone().oneshot(fast_request), + ) + .await + .expect("unrelated fast buyer traffic must remain responsive while slow streams are held") + .expect("fast gateway request must produce a response"); + assert_eq!(fast_response.status(), StatusCode::OK); + tokio::time::timeout(Duration::from_secs(2), probe.fast_request_seen.notified()) + .await + .expect("the fast request must reach the real loopback upstream"); + + drop(slow_bodies); + probe.release_tails.notify_waiters(); + upstream_task.abort(); +} diff --git a/tests/gateway_streaming.rs b/tests/gateway_streaming.rs new file mode 100644 index 00000000..bf77bd01 --- /dev/null +++ b/tests/gateway_streaming.rs @@ -0,0 +1,391 @@ +//! Hostile buyer acceptance for #442: the generic gateway must not withhold an +//! admitted downstream response until the upstream body has completed. +//! +//! This test is intentionally RED against protected `main`. The current +//! `proxy_request()` materializes the complete reqwest body before constructing +//! an Axum response, so a long-lived or delayed upstream tail blocks the buyer +//! from receiving even the response head and first body chunk. Production +//! source is intentionally unchanged in this lane. + +use std::{convert::Infallible, io, sync::Arc, time::Duration}; + +use axum::{ + Router, + body::{Body, Bytes}, + extract::State, + http::{ + Method, Request, StatusCode, + header::{CONTENT_LENGTH, CONTENT_TYPE}, + }, + response::Response, + routing::any, +}; +use futures_util::{StreamExt, stream}; +use tokio::sync::Notify; +use tower::ServiceExt; +use waf_ids_ai_soc::{AppState, build_app}; + +const DISHONEST_CONTENT_LENGTH: &str = "8388608"; + +#[derive(Clone, Default)] +struct StreamProbe { + request_seen: Arc, + release_tail: Arc, + release_failure: Arc, +} + +async fn delayed_chunk_upstream(State(probe): State) -> Response { + probe.request_seen.notify_one(); + + let release_tail = probe.release_tail.clone(); + let body_stream = stream::unfold(0u8, move |stage| { + let release_tail = release_tail.clone(); + async move { + match stage { + 0 => Some((Ok::(Bytes::from_static(b"first-")), 1)), + 1 => { + release_tail.notified().await; + Some((Ok(Bytes::from_static(b"second")), 2)) + } + _ => None, + } + } + }); + + Response::builder() + .status(StatusCode::OK) + .header(CONTENT_TYPE, "application/octet-stream") + .body(Body::from_stream(body_stream)) + .expect("valid loopback streaming response") +} + +async fn dishonest_content_length_upstream(State(probe): State) -> Response { + probe.request_seen.notify_one(); + + let release_tail = probe.release_tail.clone(); + let body_stream = stream::unfold(0u8, move |stage| { + let release_tail = release_tail.clone(); + async move { + match stage { + 0 => Some((Ok::(Bytes::from_static(b"first-")), 1)), + 1 => { + release_tail.notified().await; + Some((Ok(Bytes::from_static(b"short-tail")), 2)) + } + _ => None, + } + } + }); + + Response::builder() + .status(StatusCode::OK) + .header(CONTENT_TYPE, "application/octet-stream") + .header(CONTENT_LENGTH, DISHONEST_CONTENT_LENGTH) + .body(Body::from_stream(body_stream)) + .expect("valid loopback response with hostile content length") +} + +async fn partial_failure_upstream(State(probe): State) -> Response { + probe.request_seen.notify_one(); + + let release_failure = probe.release_failure.clone(); + let body_stream = stream::unfold(0u8, move |stage| { + let release_failure = release_failure.clone(); + async move { + match stage { + 0 => Some((Ok::(Bytes::from_static(b"first-")), 1)), + 1 => { + release_failure.notified().await; + Some(( + Err(io::Error::new( + io::ErrorKind::ConnectionReset, + "hostile upstream reset", + )), + 2, + )) + } + _ => None, + } + } + }); + + Response::builder() + .status(StatusCode::OK) + .header(CONTENT_TYPE, "application/octet-stream") + .body(Body::from_stream(body_stream)) + .expect("valid loopback partial-failure response") +} + +async fn gateway_with_delayed_upstream() -> (Router, StreamProbe, tokio::task::JoinHandle<()>) { + let probe = StreamProbe::default(); + let upstream_app = Router::new() + .route("/v1/stream", any(delayed_chunk_upstream)) + .with_state(probe.clone()); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("loopback upstream listener"); + let upstream_addr = listener.local_addr().expect("loopback upstream address"); + let upstream_task = tokio::spawn(async move { + axum::serve(listener, upstream_app) + .await + .expect("loopback upstream must serve until test cleanup"); + }); + + let app = build_app(AppState::seeded(Some("secret".to_string()))); + let route = serde_json::json!({ + "id": "streaming-red", + "path_prefix": "/stream", + "upstream": format!("http://{upstream_addr}"), + "mode": "monitor", + "enabled": true, + "block_threshold": null + }); + let created = app + .clone() + .oneshot( + Request::builder() + .method(Method::POST) + .uri("/api/routes") + .header(CONTENT_TYPE, "application/json") + .header("x-admin-token", "secret") + .body(Body::from(route.to_string())) + .expect("valid route registration request"), + ) + .await + .expect("Wardnet must answer route registration"); + assert_eq!(created.status(), StatusCode::CREATED); + + (app, probe, upstream_task) +} + +async fn gateway_with_dishonest_length_upstream() +-> (Router, StreamProbe, tokio::task::JoinHandle<()>) { + let probe = StreamProbe::default(); + let upstream_app = Router::new() + .route( + "/v1/dishonest-length", + any(dishonest_content_length_upstream), + ) + .with_state(probe.clone()); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("loopback upstream listener"); + let upstream_addr = listener.local_addr().expect("loopback upstream address"); + let upstream_task = tokio::spawn(async move { + axum::serve(listener, upstream_app) + .await + .expect("loopback upstream must serve until test cleanup"); + }); + + let app = build_app(AppState::seeded(Some("secret".to_string()))); + let route = serde_json::json!({ + "id": "streaming-dishonest-length-red", + "path_prefix": "/stream", + "upstream": format!("http://{upstream_addr}"), + "mode": "monitor", + "enabled": true, + "block_threshold": null + }); + let created = app + .clone() + .oneshot( + Request::builder() + .method(Method::POST) + .uri("/api/routes") + .header(CONTENT_TYPE, "application/json") + .header("x-admin-token", "secret") + .body(Body::from(route.to_string())) + .expect("valid route registration request"), + ) + .await + .expect("Wardnet must answer route registration"); + assert_eq!(created.status(), StatusCode::CREATED); + + (app, probe, upstream_task) +} + +async fn gateway_with_partial_failure_upstream() +-> (Router, StreamProbe, tokio::task::JoinHandle<()>) { + let probe = StreamProbe::default(); + let upstream_app = Router::new() + .route("/v1/partial", any(partial_failure_upstream)) + .with_state(probe.clone()); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("loopback upstream listener"); + let upstream_addr = listener.local_addr().expect("loopback upstream address"); + let upstream_task = tokio::spawn(async move { + axum::serve(listener, upstream_app) + .await + .expect("loopback upstream must serve until test cleanup"); + }); + + let app = build_app(AppState::seeded(Some("secret".to_string()))); + let route = serde_json::json!({ + "id": "streaming-partial-failure-red", + "path_prefix": "/stream", + "upstream": format!("http://{upstream_addr}"), + "mode": "monitor", + "enabled": true, + "block_threshold": null + }); + let created = app + .clone() + .oneshot( + Request::builder() + .method(Method::POST) + .uri("/api/routes") + .header(CONTENT_TYPE, "application/json") + .header("x-admin-token", "secret") + .body(Body::from(route.to_string())) + .expect("valid route registration request"), + ) + .await + .expect("Wardnet must answer route registration"); + assert_eq!(created.status(), StatusCode::CREATED); + + (app, probe, upstream_task) +} + +#[tokio::test] +async fn gateway_exposes_first_upstream_chunk_without_waiting_for_delayed_tail() { + let (app, probe, upstream_task) = gateway_with_delayed_upstream().await; + let request = Request::builder() + .method(Method::GET) + .uri("/gateway/stream/v1/stream") + .body(Body::empty()) + .expect("valid buyer request"); + + let mut gateway_task = tokio::spawn(async move { app.oneshot(request).await }); + + tokio::time::timeout(Duration::from_secs(2), probe.request_seen.notified()) + .await + .expect("loopback upstream must receive the admitted request before the RED assertion"); + + let response = match tokio::time::timeout(Duration::from_millis(500), &mut gateway_task).await { + Ok(joined) => joined + .expect("gateway task must not panic") + .expect("gateway service must answer"), + Err(_) => { + probe.release_tail.notify_waiters(); + let _ = tokio::time::timeout(Duration::from_secs(2), &mut gateway_task).await; + upstream_task.abort(); + panic!( + "Wardnet withheld the downstream response until the delayed upstream tail was released; generic proxying must return a streaming response without whole-body materialization" + ); + } + }; + + assert_eq!(response.status(), StatusCode::OK); + let mut downstream = response.into_body().into_data_stream(); + let first = tokio::time::timeout(Duration::from_millis(500), downstream.next()) + .await + .expect("the first admitted downstream body chunk must be available promptly") + .expect("the downstream stream must produce a first chunk") + .expect("the first downstream chunk must not be an error"); + assert_eq!(first, Bytes::from_static(b"first-")); + + probe.release_tail.notify_waiters(); + let second = tokio::time::timeout(Duration::from_secs(2), downstream.next()) + .await + .expect("the released tail must arrive") + .expect("the downstream stream must produce the tail") + .expect("the downstream tail must not be an error"); + assert_eq!(second, Bytes::from_static(b"second")); + + upstream_task.abort(); +} + +#[tokio::test] +async fn gateway_streams_prefix_before_trusting_dishonest_content_length() { + let (app, probe, upstream_task) = gateway_with_dishonest_length_upstream().await; + let request = Request::builder() + .method(Method::GET) + .uri("/gateway/stream/v1/dishonest-length") + .body(Body::empty()) + .expect("valid buyer request"); + + let mut gateway_task = tokio::spawn(async move { app.oneshot(request).await }); + + tokio::time::timeout(Duration::from_secs(2), probe.request_seen.notified()) + .await + .expect("loopback upstream must receive the admitted request before the RED assertion"); + + let response = match tokio::time::timeout(Duration::from_millis(500), &mut gateway_task).await { + Ok(joined) => joined + .expect("gateway task must not panic") + .expect("gateway service must answer"), + Err(_) => { + probe.release_tail.notify_waiters(); + let _ = tokio::time::timeout(Duration::from_secs(2), &mut gateway_task).await; + upstream_task.abort(); + panic!( + "Wardnet trusted a hostile large Content-Length enough to withhold the admitted response prefix; relay admission must not require whole-body materialization" + ); + } + }; + + assert_eq!(response.status(), StatusCode::OK); + let mut downstream = response.into_body().into_data_stream(); + let first = tokio::time::timeout(Duration::from_millis(500), downstream.next()) + .await + .expect("the admitted prefix must stream before the declared body length is satisfied") + .expect("the downstream stream must produce the admitted prefix") + .expect("the admitted prefix must not be an error"); + assert_eq!(first, Bytes::from_static(b"first-")); + + probe.release_tail.notify_waiters(); + drop(downstream); + upstream_task.abort(); +} + +#[tokio::test] +async fn gateway_exposes_admitted_prefix_before_partial_upstream_failure() { + let (app, probe, upstream_task) = gateway_with_partial_failure_upstream().await; + let request = Request::builder() + .method(Method::GET) + .uri("/gateway/stream/v1/partial") + .body(Body::empty()) + .expect("valid buyer request"); + + let mut gateway_task = tokio::spawn(async move { app.oneshot(request).await }); + + tokio::time::timeout(Duration::from_secs(2), probe.request_seen.notified()) + .await + .expect("loopback upstream must receive the admitted request before the RED assertion"); + + let response = match tokio::time::timeout(Duration::from_millis(500), &mut gateway_task).await { + Ok(joined) => joined + .expect("gateway task must not panic") + .expect("gateway service must answer"), + Err(_) => { + probe.release_failure.notify_waiters(); + let _ = tokio::time::timeout(Duration::from_secs(2), &mut gateway_task).await; + upstream_task.abort(); + panic!( + "Wardnet withheld the downstream response until the upstream stream failed; an admitted prefix must be relayed before a later upstream body failure" + ); + } + }; + + assert_eq!(response.status(), StatusCode::OK); + let mut downstream = response.into_body().into_data_stream(); + let first = tokio::time::timeout(Duration::from_millis(500), downstream.next()) + .await + .expect("the admitted prefix must be available before the upstream failure") + .expect("the downstream stream must produce the admitted prefix") + .expect("the admitted prefix must not be an error"); + assert_eq!(first, Bytes::from_static(b"first-")); + + probe.release_failure.notify_waiters(); + let failure = tokio::time::timeout(Duration::from_secs(2), downstream.next()) + .await + .expect("the upstream failure must become observable promptly") + .expect("the downstream stream must surface the upstream failure"); + assert!( + failure.is_err(), + "a partial upstream failure must not be converted into successful downstream body completion" + ); + + upstream_task.abort(); +} diff --git a/tests/gateway_streaming_bounded_memory.rs b/tests/gateway_streaming_bounded_memory.rs new file mode 100644 index 00000000..b7b34b5b --- /dev/null +++ b/tests/gateway_streaming_bounded_memory.rs @@ -0,0 +1,175 @@ +//! Hostile buyer acceptance for #442: a real upstream response larger than the +//! relay-memory budget must start reaching the buyer before the upstream tail +//! completes. +//! +//! This fixture is intentionally RED against protected `main`. The current +//! `proxy_request()` materializes the complete reqwest body, so even after the +//! upstream has made more than the test relay-memory budget available, Wardnet +//! withholds the downstream response until the held tail is released. + +use std::{convert::Infallible, sync::Arc, time::Duration}; + +use axum::{ + Router, + body::{Body, Bytes}, + extract::State, + http::{Method, Request, StatusCode, header::CONTENT_TYPE}, + response::Response, + routing::any, +}; +use futures_util::{StreamExt, stream}; +use tokio::sync::Notify; +use tower::ServiceExt; +use waf_ids_ai_soc::{AppState, build_app}; + +const RELAY_MEMORY_BUDGET_BYTES: usize = 1024 * 1024; +const LARGE_PREFIX_BYTES: usize = RELAY_MEMORY_BUDGET_BYTES + 64 * 1024; +const UPSTREAM_CHUNK_BYTES: usize = 16 * 1024; + +#[derive(Clone, Default)] +struct LargeResponseProbe { + request_seen: Arc, + release_tail: Arc, +} + +async fn large_prefix_upstream(State(probe): State) -> Response { + probe.request_seen.notify_one(); + + let release_tail = probe.release_tail.clone(); + let body_stream = stream::unfold(0usize, move |emitted| { + let release_tail = release_tail.clone(); + async move { + if emitted < LARGE_PREFIX_BYTES { + let remaining = LARGE_PREFIX_BYTES - emitted; + let chunk_len = remaining.min(UPSTREAM_CHUNK_BYTES); + return Some(( + Ok::(Bytes::from(vec![0xA5; chunk_len])), + emitted + chunk_len, + )); + } + + if emitted == LARGE_PREFIX_BYTES { + release_tail.notified().await; + return Some((Ok(Bytes::from_static(b"tail")), LARGE_PREFIX_BYTES + 1)); + } + + None + } + }); + + Response::builder() + .status(StatusCode::OK) + .header(CONTENT_TYPE, "application/octet-stream") + .body(Body::from_stream(body_stream)) + .expect("valid loopback large streaming response") +} + +async fn gateway_with_large_upstream() -> (Router, LargeResponseProbe, tokio::task::JoinHandle<()>) +{ + let probe = LargeResponseProbe::default(); + let upstream_app = Router::new() + .route("/v1/large", any(large_prefix_upstream)) + .with_state(probe.clone()); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("loopback upstream listener"); + let upstream_addr = listener.local_addr().expect("loopback upstream address"); + let upstream_task = tokio::spawn(async move { + axum::serve(listener, upstream_app) + .await + .expect("loopback upstream must serve until test cleanup"); + }); + + let app = build_app(AppState::seeded(Some("secret".to_string()))); + let route = serde_json::json!({ + "id": "streaming-large-response-red", + "path_prefix": "/stream", + "upstream": format!("http://{upstream_addr}"), + "mode": "monitor", + "enabled": true, + "block_threshold": null + }); + let created = app + .clone() + .oneshot( + Request::builder() + .method(Method::POST) + .uri("/api/routes") + .header(CONTENT_TYPE, "application/json") + .header("x-admin-token", "secret") + .body(Body::from(route.to_string())) + .expect("valid route registration request"), + ) + .await + .expect("Wardnet must answer route registration"); + assert_eq!(created.status(), StatusCode::CREATED); + + (app, probe, upstream_task) +} + +#[tokio::test] +async fn gateway_releases_large_prefix_before_held_tail() { + assert!( + UPSTREAM_CHUNK_BYTES < RELAY_MEMORY_BUDGET_BYTES, + "the hostile fixture must not manufacture one upstream chunk larger than the relay budget" + ); + + let (app, probe, upstream_task) = gateway_with_large_upstream().await; + let request = Request::builder() + .method(Method::GET) + .uri("/gateway/stream/v1/large") + .body(Body::empty()) + .expect("valid buyer request"); + + let mut gateway_task = tokio::spawn(async move { app.oneshot(request).await }); + + tokio::time::timeout(Duration::from_secs(2), probe.request_seen.notified()) + .await + .expect("loopback upstream must receive the admitted request before the RED assertion"); + + let response = match tokio::time::timeout(Duration::from_millis(500), &mut gateway_task).await { + Ok(joined) => joined + .expect("gateway task must not panic") + .expect("gateway service must answer"), + Err(_) => { + probe.release_tail.notify_waiters(); + let _ = tokio::time::timeout(Duration::from_secs(2), &mut gateway_task).await; + upstream_task.abort(); + panic!( + "Wardnet withheld a response whose available prefix exceeds the bounded relay-memory budget until the upstream tail completed" + ); + } + }; + + assert_eq!(response.status(), StatusCode::OK); + let mut downstream = response.into_body().into_data_stream(); + let mut admitted_prefix_bytes = 0usize; + + while admitted_prefix_bytes < LARGE_PREFIX_BYTES { + let next = tokio::time::timeout(Duration::from_secs(2), downstream.next()) + .await + .expect("the large admitted prefix must remain readable while the tail is held") + .expect("the downstream body must not complete before the held tail") + .expect("the admitted prefix must not become an error"); + admitted_prefix_bytes += next.len(); + } + + assert_eq!( + admitted_prefix_bytes, LARGE_PREFIX_BYTES, + "the bounded fixture must relay exactly the admitted prefix before the held tail" + ); + assert!( + admitted_prefix_bytes > RELAY_MEMORY_BUDGET_BYTES, + "the buyer must receive more than the relay-memory budget before the upstream tail is released" + ); + + probe.release_tail.notify_waiters(); + let tail = tokio::time::timeout(Duration::from_secs(2), downstream.next()) + .await + .expect("the released tail must arrive") + .expect("the downstream body must produce the released tail") + .expect("the released tail must not be an error"); + assert_eq!(tail, Bytes::from_static(b"tail")); + + upstream_task.abort(); +} diff --git a/tests/gateway_streaming_cancellation.rs b/tests/gateway_streaming_cancellation.rs new file mode 100644 index 00000000..e092547e --- /dev/null +++ b/tests/gateway_streaming_cancellation.rs @@ -0,0 +1,180 @@ +//! Hostile buyer acceptance for #442: once Wardnet has admitted and exposed a +//! streaming response, downstream cancellation must stop the upstream body +//! instead of continuing to read it to completion in the background. +//! +//! This fixture is intentionally RED against protected `main`: current +//! `proxy_request()` collects the whole reqwest response before it can return an +//! Axum response, so the buyer cannot cancel the admitted body while the +//! upstream tail is still held. + +use std::{ + convert::Infallible, + sync::{ + Arc, + atomic::{AtomicBool, Ordering}, + }, + time::Duration, +}; + +use axum::{ + Router, + body::{Body, Bytes}, + extract::State, + http::{Method, Request, StatusCode, header::CONTENT_TYPE}, + response::Response, + routing::any, +}; +use futures_util::{StreamExt, stream}; +use tokio::sync::Notify; +use tower::ServiceExt; +use waf_ids_ai_soc::{AppState, build_app}; + +#[derive(Clone, Default)] +struct CancellationProbe { + request_seen: Arc, + release_tail: Arc, + upstream_body_dropped: Arc, +} + +struct UpstreamBodyDropSignal(Arc); + +impl Drop for UpstreamBodyDropSignal { + fn drop(&mut self) { + self.0.store(true, Ordering::Release); + } +} + +struct CancellationStreamState { + stage: u8, + release_tail: Arc, + _drop_signal: UpstreamBodyDropSignal, +} + +async fn cancellation_upstream(State(probe): State) -> Response { + probe.request_seen.notify_one(); + + let body_stream = stream::unfold( + CancellationStreamState { + stage: 0, + release_tail: probe.release_tail.clone(), + _drop_signal: UpstreamBodyDropSignal(probe.upstream_body_dropped.clone()), + }, + |mut state| async move { + match state.stage { + 0 => { + state.stage = 1; + Some(( + Ok::(Bytes::from_static(b"first-")), + state, + )) + } + 1 => { + state.release_tail.notified().await; + state.stage = 2; + Some((Ok(Bytes::from_static(b"tail")), state)) + } + _ => None, + } + }, + ); + + Response::builder() + .status(StatusCode::OK) + .header(CONTENT_TYPE, "application/octet-stream") + .body(Body::from_stream(body_stream)) + .expect("valid loopback cancellation response") +} + +async fn gateway_with_cancellation_upstream() +-> (Router, CancellationProbe, tokio::task::JoinHandle<()>) { + let probe = CancellationProbe::default(); + let upstream_app = Router::new() + .route("/v1/cancel", any(cancellation_upstream)) + .with_state(probe.clone()); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("loopback upstream listener"); + let upstream_addr = listener.local_addr().expect("loopback upstream address"); + let upstream_task = tokio::spawn(async move { + axum::serve(listener, upstream_app) + .await + .expect("loopback upstream must serve until test cleanup"); + }); + + let app = build_app(AppState::seeded(Some("secret".to_string()))); + let route = serde_json::json!({ + "id": "streaming-cancellation-red", + "path_prefix": "/stream", + "upstream": format!("http://{upstream_addr}"), + "mode": "monitor", + "enabled": true, + "block_threshold": null + }); + let created = app + .clone() + .oneshot( + Request::builder() + .method(Method::POST) + .uri("/api/routes") + .header(CONTENT_TYPE, "application/json") + .header("x-admin-token", "secret") + .body(Body::from(route.to_string())) + .expect("valid route registration request"), + ) + .await + .expect("Wardnet must answer route registration"); + assert_eq!(created.status(), StatusCode::CREATED); + + (app, probe, upstream_task) +} + +#[tokio::test] +async fn downstream_cancellation_drops_held_upstream_body() { + let (app, probe, upstream_task) = gateway_with_cancellation_upstream().await; + let request = Request::builder() + .method(Method::GET) + .uri("/gateway/stream/v1/cancel") + .body(Body::empty()) + .expect("valid buyer request"); + + let mut gateway_task = tokio::spawn(async move { app.oneshot(request).await }); + + tokio::time::timeout(Duration::from_secs(2), probe.request_seen.notified()) + .await + .expect("loopback upstream must receive the admitted request before the RED assertion"); + + let response = match tokio::time::timeout(Duration::from_millis(500), &mut gateway_task).await { + Ok(joined) => joined + .expect("gateway task must not panic") + .expect("gateway service must answer"), + Err(_) => { + probe.release_tail.notify_waiters(); + let _ = tokio::time::timeout(Duration::from_secs(2), &mut gateway_task).await; + upstream_task.abort(); + panic!( + "Wardnet withheld the admitted response until the held upstream tail completed; downstream cancellation cannot propagate while the whole body is materialized" + ); + } + }; + + assert_eq!(response.status(), StatusCode::OK); + let mut downstream = response.into_body().into_data_stream(); + let first = tokio::time::timeout(Duration::from_millis(500), downstream.next()) + .await + .expect("the admitted prefix must be available before cancellation") + .expect("the downstream stream must produce the admitted prefix") + .expect("the admitted prefix must not be an error"); + assert_eq!(first, Bytes::from_static(b"first-")); + + drop(downstream); + + tokio::time::timeout(Duration::from_secs(2), async { + while !probe.upstream_body_dropped.load(Ordering::Acquire) { + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("dropping the admitted downstream body must promptly cancel the held upstream body"); + + upstream_task.abort(); +}