AkurAI Build
Menu

BifrOSt-Apps

public

Latest change f35619cfe0419c596645328bf0e1b682ad3a4473 - Harden RÚV for the 0.1.1 candidate by Ólafur Búi Ólafsson

//! Loopback mock-HTTP behavior tests: byte caps, deadlines, redirect policy,
//! and the four-download image admission budget. No traffic leaves 127.0.0.1.

use std::io::{Read, Write};
use std::net::{TcpListener, TcpStream};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::mpsc::Sender;
use std::thread;
use std::time::{Duration, Instant};

#[allow(dead_code)]
#[path = "../src/model.rs"]
mod model;

/// Serves exactly one connection with `handler`, returning the `http://` URL.
fn spawn_server<F>(handler: F) -> String
where
    F: FnOnce(TcpStream) + Send + 'static,
{
    let listener = TcpListener::bind("127.0.0.1:0").expect("bind loopback");
    let port = listener.local_addr().expect("local addr").port();
    thread::spawn(move || {
        if let Ok((stream, _)) = listener.accept() {
            handler(stream);
        }
    });
    format!("http://127.0.0.1:{port}/")
}

/// A listener that must never be reached; every accepted connection is counted.
fn counting_sink() -> (String, Arc<AtomicUsize>) {
    let listener = TcpListener::bind("127.0.0.1:0").expect("bind loopback");
    let port = listener.local_addr().expect("local addr").port();
    let hits = Arc::new(AtomicUsize::new(0));
    let counter = hits.clone();
    thread::spawn(move || {
        while let Ok((stream, _)) = listener.accept() {
            counter.fetch_add(1, Ordering::SeqCst);
            drop(stream);
        }
    });
    (format!("http://127.0.0.1:{port}/"), hits)
}

fn read_request_head(stream: &mut TcpStream) {
    let _ = stream.set_read_timeout(Some(Duration::from_secs(10)));
    let mut head = Vec::new();
    let mut buffer = [0_u8; 1024];
    while !head.windows(4).any(|window| window == b"\r\n\r\n") {
        match stream.read(&mut buffer) {
            Ok(0) | Err(_) => break,
            Ok(count) => head.extend_from_slice(&buffer[..count]),
        }
        if head.len() > 16 * 1024 {
            break;
        }
    }
}

fn respond(stream: &mut TcpStream, payload: &str) {
    let _ = stream.set_write_timeout(Some(Duration::from_secs(10)));
    let _ = stream.write_all(payload.as_bytes());
}

/// Streams `total` body bytes as HTTP/1.1 chunks until the peer hangs up.
fn flood_chunks(stream: &mut TcpStream, total: usize) {
    let chunk = [b'a'; 4096];
    let header = format!("{:x}\r\n", chunk.len());
    let mut sent = 0;
    while sent < total {
        let written = stream
            .write_all(header.as_bytes())
            .and_then(|()| stream.write_all(&chunk))
            .and_then(|()| stream.write_all(b"\r\n"));
        if written.is_err() {
            return;
        }
        sent += chunk.len();
    }
    let _ = stream.write_all(b"0\r\n\r\n");
}

/// Streams `total` raw body bytes until the peer hangs up.
fn flood_raw(stream: &mut TcpStream, total: usize) {
    let chunk = [b'a'; 4096];
    let mut sent = 0;
    while sent < total {
        if stream.write_all(&chunk).is_err() {
            return;
        }
        sent += chunk.len();
    }
}

/// Counts concurrently connected downloads and holds every response until
/// `gate` opens. Returns the bare `host:port` address.
fn gated_counting_server(
    gate: Arc<AtomicBool>,
    active: Arc<AtomicUsize>,
    peak: Arc<AtomicUsize>,
) -> String {
    let listener = TcpListener::bind("127.0.0.1:0").expect("bind loopback");
    let port = listener.local_addr().expect("local addr").port();
    thread::spawn(move || {
        while let Ok((mut stream, _)) = listener.accept() {
            let gate = gate.clone();
            let active = active.clone();
            let peak = peak.clone();
            thread::spawn(move || {
                let now = active.fetch_add(1, Ordering::SeqCst) + 1;
                peak.fetch_max(now, Ordering::SeqCst);
                read_request_head(&mut stream);
                let started = Instant::now();
                while !gate.load(Ordering::SeqCst) && started.elapsed() < Duration::from_secs(10) {
                    thread::sleep(Duration::from_millis(5));
                }
                respond(
                    &mut stream,
                    "HTTP/1.1 200 OK\r\nContent-Type: image/png\r\nContent-Length: 2\r\n\r\nok",
                );
                active.fetch_sub(1, Ordering::SeqCst);
            });
        }
    });
    format!("127.0.0.1:{port}")
}

fn spawn_download(address: String, path: String, done: Sender<()>) {
    thread::spawn(move || {
        if let Ok(mut stream) = TcpStream::connect(&address) {
            let _ = stream.set_read_timeout(Some(Duration::from_secs(15)));
            let _ = write!(
                stream,
                "GET /{path} HTTP/1.1\r\nHost: {address}\r\nConnection: close\r\n\r\n"
            );
            let mut sink = Vec::new();
            let _ = stream.read_to_end(&mut sink);
        }
        let _ = done.send(());
    });
}

/// The image admission budget in `ImagePump` must never let more than four
/// downloads run at once, observed as server-side connection concurrency.
#[test]
fn image_pump_admits_at_most_four_concurrent_downloads() {
    use ruv::app::{IMAGE_CONCURRENCY, ImagePump};

    assert_eq!(IMAGE_CONCURRENCY, 4, "README promises at most four");

    let gate = Arc::new(AtomicBool::new(false));
    let active = Arc::new(AtomicUsize::new(0));
    let peak = Arc::new(AtomicUsize::new(0));
    let address = gated_counting_server(gate.clone(), active.clone(), peak.clone());

    const TOTAL: usize = 12;
    let mut pump = ImagePump::default();
    for index in 0..TOTAL {
        pump.enqueue(format!("mynd-{index}.png"));
    }

    let (done_tx, done_rx) = std::sync::mpsc::channel();
    let mut launched = 0_usize;

    for path in pump.admit() {
        launched += 1;
        spawn_download(address.clone(), path, done_tx.clone());
    }
    assert_eq!(launched, IMAGE_CONCURRENCY);
    assert_eq!(pump.inflight(), IMAGE_CONCURRENCY);
    assert!(
        pump.admit().is_empty(),
        "a saturated pump must not admit more downloads"
    );

    // All four admitted downloads must be connected simultaneously before the
    // gate opens; only then is the concurrency observation meaningful.
    let started = Instant::now();
    while active.load(Ordering::SeqCst) < IMAGE_CONCURRENCY
        && started.elapsed() < Duration::from_secs(10)
    {
        thread::sleep(Duration::from_millis(5));
    }
    assert_eq!(active.load(Ordering::SeqCst), IMAGE_CONCURRENCY);
    assert!(pump.admit().is_empty());

    gate.store(true, Ordering::SeqCst);
    let mut finished = 0_usize;
    while finished < TOTAL {
        done_rx
            .recv_timeout(Duration::from_secs(10))
            .expect("every admitted download should finish");
        finished += 1;
        pump.complete();
        for path in pump.admit() {
            launched += 1;
            spawn_download(address.clone(), path, done_tx.clone());
        }
        assert!(pump.inflight() <= IMAGE_CONCURRENCY);
    }

    assert_eq!(launched, TOTAL, "every queued image must be fetched");
    assert_eq!(
        peak.load(Ordering::SeqCst),
        IMAGE_CONCURRENCY,
        "the server must never observe more than four concurrent downloads"
    );
}

mod behavior {
    #![allow(dead_code)]
    include!("../src/api.rs");

    use super::{counting_sink, flood_chunks, flood_raw, read_request_head, respond, spawn_server};
    use std::sync::atomic::Ordering;
    use std::time::Instant;

    #[tokio::test]
    async fn chunked_body_exceeding_the_cap_is_rejected_mid_stream() {
        let url = spawn_server(|mut stream| {
            read_request_head(&mut stream);
            respond(
                &mut stream,
                "HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nTransfer-Encoding: chunked\r\n\r\n",
            );
            flood_chunks(&mut stream, 4 * 1024 * 1024);
        });
        let ruv = RuvClient::new().expect("client");
        let response = ruv
            .client
            .get(&url)
            .send()
            .await
            .expect("headers should arrive");
        let result = read_response(response, 256 * 1024).await;
        assert_eq!(result.err(), Some(AppError::ResponseTooLarge));
    }

    #[tokio::test]
    async fn close_delimited_body_without_content_length_is_capped_mid_stream() {
        let url = spawn_server(|mut stream| {
            read_request_head(&mut stream);
            respond(
                &mut stream,
                "HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nConnection: close\r\n\r\n",
            );
            flood_raw(&mut stream, 4 * 1024 * 1024);
        });
        let ruv = RuvClient::new().expect("client");
        let response = ruv
            .client
            .get(&url)
            .send()
            .await
            .expect("headers should arrive");
        let result = read_response(response, 256 * 1024).await;
        assert_eq!(result.err(), Some(AppError::ResponseTooLarge));
    }

    #[tokio::test]
    async fn oversized_declared_content_length_is_rejected_before_the_body() {
        let url = spawn_server(|mut stream| {
            read_request_head(&mut stream);
            // Announce far more than the cap, then stall: the client must
            // reject on the declared size instead of waiting for bytes.
            respond(
                &mut stream,
                "HTTP/1.1 200 OK\r\nContent-Length: 10485760\r\n\r\n",
            );
            std::thread::sleep(Duration::from_secs(8));
        });
        let ruv = RuvClient::new().expect("client");
        let started = Instant::now();
        let response = ruv
            .client
            .get(&url)
            .send()
            .await
            .expect("headers should arrive");
        let result = read_response(response, 1024).await;
        assert_eq!(result.err(), Some(AppError::ResponseTooLarge));
        assert!(
            started.elapsed() < Duration::from_secs(5),
            "rejection must come from the header, not from downloading the body"
        );
    }

    #[tokio::test]
    async fn lying_small_content_length_never_yields_more_than_declared() {
        let url = spawn_server(|mut stream| {
            read_request_head(&mut stream);
            respond(&mut stream, "HTTP/1.1 200 OK\r\nContent-Length: 16\r\n\r\n");
            flood_raw(&mut stream, 1024 * 1024);
        });
        let ruv = RuvClient::new().expect("client");
        let response = ruv
            .client
            .get(&url)
            .send()
            .await
            .expect("headers should arrive");
        let body = read_response(response, 1024)
            .await
            .expect("a body within the cap is accepted");
        assert_eq!(body.len(), 16, "framing must stop at the declared length");
    }

    #[tokio::test]
    async fn stalled_server_is_cut_off_by_the_client_deadline() {
        let url = spawn_server(|mut stream| {
            read_request_head(&mut stream);
            // Longer than the 15-second client deadline; the connection is
            // dropped afterwards so the test is bounded even on regression.
            std::thread::sleep(Duration::from_secs(25));
        });
        let ruv = RuvClient::new().expect("client");
        let started = Instant::now();
        let error = ruv
            .client
            .get(&url)
            .send()
            .await
            .expect_err("a stalled response must not succeed");
        assert!(error.is_timeout(), "expected a deadline error: {error}");
        assert!(
            started.elapsed() < Duration::from_secs(20),
            "the deadline must fire before the server gives up"
        );
    }

    #[tokio::test]
    async fn redirect_to_an_untrusted_http_origin_is_stopped_unvisited() {
        let (leak_url, hits) = counting_sink();
        let location = format!("{leak_url}secret.json");
        let url = spawn_server(move |mut stream| {
            read_request_head(&mut stream);
            respond(
                &mut stream,
                &format!("HTTP/1.1 302 Found\r\nLocation: {location}\r\nContent-Length: 0\r\n\r\n"),
            );
        });
        let ruv = RuvClient::new().expect("client");
        let response = ruv
            .client
            .get(&url)
            .send()
            .await
            .expect("the stopped redirect is returned, not followed");
        assert_eq!(response.status(), StatusCode::FOUND);
        assert_eq!(
            read_response(response, 1024).await.err(),
            Some(AppError::InvalidResponse)
        );
        std::thread::sleep(Duration::from_millis(250));
        assert_eq!(
            hits.load(Ordering::SeqCst),
            0,
            "the untrusted redirect target must never be contacted"
        );
    }

    #[tokio::test]
    async fn redirect_to_an_untrusted_https_origin_is_stopped_unvisited() {
        let (leak_url, hits) = counting_sink();
        // HTTPS scheme but an IP host on a non-443 port: every part of the
        // redirect allowlist must be re-validated, not just the scheme.
        let location = leak_url.replacen("http://", "https://", 1);
        let url = spawn_server(move |mut stream| {
            read_request_head(&mut stream);
            respond(
                &mut stream,
                &format!("HTTP/1.1 302 Found\r\nLocation: {location}\r\nContent-Length: 0\r\n\r\n"),
            );
        });
        let ruv = RuvClient::new().expect("client");
        let response = ruv
            .client
            .get(&url)
            .send()
            .await
            .expect("the stopped redirect is returned, not followed");
        assert_eq!(response.status(), StatusCode::FOUND);
        std::thread::sleep(Duration::from_millis(250));
        assert_eq!(
            hits.load(Ordering::SeqCst),
            0,
            "the untrusted redirect target must never be contacted"
        );
    }
}