AkurAI Build
Menu

AkurAI-Build

public

Latest change 1c2159692a31765cd66ed709791ba11468054873 - Initial commit: bunfork v0.1.0 source tree by Olafur Bui

#![cfg(unix)]

use std::{
    collections::BTreeMap,
    env, fs,
    io::{Read, Write},
    net::{TcpListener, TcpStream},
    path::{Path, PathBuf},
    process::{Child, Command, Stdio},
    sync::atomic::{AtomicU16, AtomicU64, Ordering},
    thread,
    time::{Duration, Instant},
};

static NEXT_PORT: AtomicU16 = AtomicU16::new(31_100);
static NEXT_SCRATCH: AtomicU64 = AtomicU64::new(0);

#[test]
fn static_artifact_http_matrix() {
    mpa_contract();
    spa_contract();
    base_path_contract();
}

fn mpa_contract() {
    let mut server = StaticServer::start("mpa");
    server.assert_no_runtime_state();

    let home = server.request("GET", "/");
    home.assert_status(200);
    home.assert_body_contains("hand-authored MPA home");
    home.assert_header_contains("content-type", "text/html");
    home.assert_header_contains("cache-control", "no-cache");
    home.assert_security_headers();

    let head = server.request("HEAD", "/");
    assert_head_matches(&home, &head);

    let etag = home.headers.get("etag").expect("GET omitted ETag").clone();
    assert!(
        etag.starts_with('"') && etag.ends_with('"') && !etag.starts_with("W/"),
        "static artifact ETag must be strong: {etag}"
    );
    let not_modified = server.request_with_headers("GET", "/", &[("If-None-Match", &etag)]);
    not_modified.assert_status(304);
    assert!(not_modified.body.is_empty(), "304 returned a body");
    not_modified.assert_header_contains("etag", &etag);
    not_modified.assert_header_contains("cache-control", "no-cache");
    not_modified.assert_header_contains("vary", "Accept-Encoding");
    not_modified.assert_security_headers();

    let nested = server.request("GET", "/guide/");
    nested.assert_status(200);
    nested.assert_body_contains("hand-authored nested MPA page");

    let missing = server.request("GET", "/missing");
    missing.assert_status(404);
    missing.assert_body_contains("hand-authored MPA 404");
    missing.assert_header_contains("cache-control", "no-cache");

    let asset = server.request("GET", "/assets/app-a1b2c3d4.css");
    asset.assert_status(200);
    asset.assert_body_contains("rgb(12 34 56)");
    asset.assert_header_contains("content-type", "text/css");
    asset.assert_header_contains("cache-control", "max-age=0");
    asset.assert_header_contains("cache-control", "must-revalidate");
    asset.assert_security_headers();

    let missing_asset = server.request("GET", "/assets/missing-a1b2c3d4.css");
    missing_asset.assert_status(404);

    assert_method_not_allowed(&server, "/");
    assert_traversal_rejected(&server);
    server.assert_no_runtime_state();
    server.shutdown_with_sigterm();
}

fn spa_contract() {
    let mut server = StaticServer::start("spa");
    server.assert_no_runtime_state();

    let index = server.request("GET", "/");
    index.assert_status(200);
    index.assert_body_contains("hand-authored SPA fallback shell");
    index.assert_body_excludes("hand-authored SPA 200 fallback");

    let deep_link = server.request("GET", "/dashboard/settings?tab=one");
    deep_link.assert_status(200);
    deep_link.assert_body_contains("hand-authored SPA 200 fallback");
    deep_link.assert_header_contains("cache-control", "no-cache");
    deep_link.assert_security_headers();

    let asset = server.request("GET", "/assets/app-deadbeef.js");
    asset.assert_status(200);
    asset.assert_body_contains("static-spa");
    asset.assert_header_contains("content-type", "javascript");
    asset.assert_header_contains("cache-control", "must-revalidate");

    for excluded in [
        "/assets/missing.js",
        "/assets/missing",
        "/static/missing",
        "/_next/missing",
        "/_nuxt/missing",
        "/missing.js",
        "/api/private",
    ] {
        let response = server.request("GET", excluded);
        response.assert_status(404);
        response.assert_body_excludes("hand-authored SPA 200 fallback");
    }

    assert_method_not_allowed(&server, "/dashboard");
    assert_traversal_rejected(&server);
    server.shutdown_with_sigterm();
}

fn base_path_contract() {
    let mut server = StaticServer::start("base_path");

    for outside in ["/", "/guide/", "/assets/app-c0ffee12.css"] {
        server.request("GET", outside).assert_status(404);
    }

    let home = server.request("GET", "/docs/");
    home.assert_status(200);
    home.assert_body_contains("hand-authored /docs base-path home");

    let redirect = server.request("GET", "/docs?x=1");
    redirect.assert_status(308);
    redirect.assert_header_contains("location", "/docs/?x=1");

    let nested = server.request("GET", "/docs/guide/");
    nested.assert_status(200);
    nested.assert_body_contains("hand-authored /docs/guide page");

    let asset = server.request("GET", "/docs/assets/app-c0ffee12.css");
    asset.assert_status(200);
    asset.assert_header_contains("cache-control", "must-revalidate");

    server
        .request("GET", "/docs/assets/missing-c0ffee12.css")
        .assert_status(404);
    assert_method_not_allowed(&server, "/docs/");
    server.shutdown_with_sigterm();
}

fn assert_method_not_allowed(server: &StaticServer, path: &str) {
    let response = server.request("POST", path);
    response.assert_status(405);
    response.assert_header_contains("allow", "GET");
    response.assert_header_contains("allow", "HEAD");
    response.assert_security_headers();
}

fn assert_traversal_rejected(server: &StaticServer) {
    for path in ["/%2e%2e/secret.txt", "/assets/%2e%2e/index.html"] {
        let response = server.request("GET", path);
        assert!(
            matches!(response.status, 400 | 404),
            "traversal {path} returned {} instead of 400/404; response={response:?}",
            response.status
        );
        response.assert_body_excludes("hand-authored SPA 200 fallback");
        response.assert_body_excludes("hand-authored MPA home");
    }
}

fn assert_head_matches(get: &HttpResponse, head: &HttpResponse) {
    assert_eq!(head.status, get.status, "HEAD status differs from GET");
    assert!(head.body.is_empty(), "HEAD returned a response body");
    for name in ["content-type", "cache-control", "etag"] {
        assert_eq!(
            head.headers.get(name),
            get.headers.get(name),
            "HEAD {name} differs from GET"
        );
    }
    if let Some(length) = get.headers.get("content-length") {
        assert_eq!(
            head.headers.get("content-length"),
            Some(length),
            "HEAD content-length differs from GET"
        );
    }
}

struct StaticServer {
    child: Option<Child>,
    address: String,
    scratch: ScratchDir,
}

impl StaticServer {
    fn start(name: &str) -> Self {
        let fixture = fixture(name);
        let site = fixture.join("site");
        let manifest = fixture.join("manifest.json");
        assert!(site.is_dir(), "missing fixture site {}", site.display());
        assert!(
            manifest.is_file(),
            "missing fixture manifest {}",
            manifest.display()
        );

        let scratch = ScratchDir::new(name);
        let port = free_port();
        let address = format!("127.0.0.1:{port}");
        let mut command = Command::new(env!("CARGO_BIN_EXE_bunfork"));
        command
            .args(["serve", "--static"])
            .arg(&site)
            .arg("--manifest")
            .arg(&manifest)
            .arg("--address")
            .arg(&address)
            .current_dir(scratch.path())
            .stdin(Stdio::null())
            .stdout(Stdio::piped())
            .stderr(Stdio::piped());
        for variable in [
            "BUNFORK_DB",
            "BUNFORK_DB_KEY",
            "BUNFORK_KEY_FILE",
            "BUNFORK_API_TOKEN",
            "BUNFORK_TOKEN_FILE",
            "BUNFORK_TENANT",
            "BUNFORK_MODEL",
        ] {
            command.env_remove(variable);
        }

        let mut child = command.spawn().expect("spawn bunfork static server");
        let deadline = Instant::now() + Duration::from_secs(10);
        loop {
            if TcpStream::connect(&address).is_ok() {
                break;
            }
            if let Some(status) = child.try_wait().expect("poll bunfork static server") {
                let output = take_output(&mut child);
                panic!("bunfork exited before listening ({status}):\n{output}");
            }
            if Instant::now() >= deadline {
                let _ = child.kill();
                let _ = child.wait();
                let output = take_output(&mut child);
                panic!("bunfork did not listen on {address}:\n{output}");
            }
            thread::sleep(Duration::from_millis(20));
        }

        Self {
            child: Some(child),
            address,
            scratch,
        }
    }

    fn request(&self, method: &str, target: &str) -> HttpResponse {
        self.request_with_headers(method, target, &[])
    }

    fn request_with_headers(
        &self,
        method: &str,
        target: &str,
        headers: &[(&str, &str)],
    ) -> HttpResponse {
        let mut stream = TcpStream::connect(&self.address).expect("connect to static server");
        stream
            .set_read_timeout(Some(Duration::from_secs(5)))
            .expect("set HTTP read timeout");
        stream
            .set_write_timeout(Some(Duration::from_secs(5)))
            .expect("set HTTP write timeout");
        let mut request = format!("{method} {target} HTTP/1.1\r\nHost: {}\r\n", self.address);
        for (name, value) in headers {
            request.push_str(name);
            request.push_str(": ");
            request.push_str(value);
            request.push_str("\r\n");
        }
        request.push_str("Content-Length: 0\r\nConnection: close\r\n\r\n");
        stream
            .write_all(request.as_bytes())
            .expect("write HTTP request");
        let mut bytes = Vec::new();
        stream.read_to_end(&mut bytes).expect("read HTTP response");
        HttpResponse::parse(&bytes)
    }

    fn assert_no_runtime_state(&self) {
        for relative in [".bunfork.key", ".bunfork.token", "data", "backups"] {
            assert!(
                !self.scratch.path().join(relative).exists(),
                "static mode created runtime state: {relative}"
            );
        }
    }

    fn shutdown_with_sigterm(&mut self) {
        let child = self.child.as_mut().expect("server process missing");
        let signal = Command::new("kill")
            .args(["-TERM", &child.id().to_string()])
            .status()
            .expect("send SIGTERM");
        assert!(signal.success(), "kill -TERM failed: {signal}");

        let deadline = Instant::now() + Duration::from_secs(5);
        loop {
            if let Some(status) = child.try_wait().expect("wait for SIGTERM shutdown") {
                let output = take_output(child);
                assert!(
                    status.success(),
                    "bunfork did not exit cleanly after SIGTERM ({status}):\n{output}"
                );
                self.child = None;
                return;
            }
            if Instant::now() >= deadline {
                let _ = child.kill();
                let _ = child.wait();
                let output = take_output(child);
                panic!("bunfork did not stop after SIGTERM:\n{output}");
            }
            thread::sleep(Duration::from_millis(20));
        }
    }
}

impl Drop for StaticServer {
    fn drop(&mut self) {
        if let Some(child) = &mut self.child {
            let _ = child.kill();
            let _ = child.wait();
        }
    }
}

#[derive(Debug)]
struct HttpResponse {
    status: u16,
    headers: BTreeMap<String, String>,
    body: Vec<u8>,
}

impl HttpResponse {
    fn parse(bytes: &[u8]) -> Self {
        let separator = bytes
            .windows(4)
            .position(|window| window == b"\r\n\r\n")
            .expect("HTTP response omitted header terminator");
        let head = std::str::from_utf8(&bytes[..separator]).expect("HTTP headers are not UTF-8");
        let mut lines = head.split("\r\n");
        let status = lines
            .next()
            .and_then(|line| line.split_whitespace().nth(1))
            .and_then(|value| value.parse().ok())
            .expect("HTTP response omitted status");
        let mut headers = BTreeMap::<String, String>::new();
        for line in lines {
            let (name, value) = line.split_once(':').expect("malformed HTTP header");
            headers
                .entry(name.trim().to_ascii_lowercase())
                .and_modify(|existing| {
                    existing.push_str(", ");
                    existing.push_str(value.trim());
                })
                .or_insert_with(|| value.trim().to_owned());
        }
        let raw_body = &bytes[separator + 4..];
        let body = if raw_body.is_empty() {
            Vec::new()
        } else if headers
            .get("transfer-encoding")
            .is_some_and(|value| value.eq_ignore_ascii_case("chunked"))
        {
            decode_chunked(raw_body)
        } else {
            raw_body.to_vec()
        };
        Self {
            status,
            headers,
            body,
        }
    }

    fn assert_status(&self, expected: u16) {
        assert_eq!(self.status, expected, "unexpected HTTP response: {self:?}");
    }

    fn assert_header_contains(&self, name: &str, expected: &str) {
        let value = self
            .headers
            .get(name)
            .unwrap_or_else(|| panic!("missing {name} header: {self:?}"));
        assert!(
            value
                .to_ascii_lowercase()
                .contains(&expected.to_ascii_lowercase()),
            "{name}={value:?} does not contain {expected:?}"
        );
    }

    fn assert_body_contains(&self, expected: &str) {
        let body = String::from_utf8_lossy(&self.body);
        assert!(
            body.contains(expected),
            "body does not contain {expected:?}: {body:?}"
        );
    }

    fn assert_body_excludes(&self, unexpected: &str) {
        let body = String::from_utf8_lossy(&self.body);
        assert!(
            !body.contains(unexpected),
            "body unexpectedly contains {unexpected:?}: {body:?}"
        );
    }

    fn assert_security_headers(&self) {
        self.assert_header_contains("x-content-type-options", "nosniff");
        self.assert_header_contains("x-frame-options", "deny");
        self.assert_header_contains("referrer-policy", "no-referrer");
        self.assert_header_contains("content-security-policy", "default-src 'self'");
        self.assert_header_contains("permissions-policy", "camera=()");
    }
}

fn decode_chunked(bytes: &[u8]) -> Vec<u8> {
    let mut decoded = Vec::new();
    let mut cursor = 0;
    loop {
        let line_end = bytes[cursor..]
            .windows(2)
            .position(|window| window == b"\r\n")
            .map(|offset| cursor + offset)
            .expect("invalid chunk size line");
        let size_text = std::str::from_utf8(&bytes[cursor..line_end])
            .expect("chunk size is not ASCII")
            .split(';')
            .next()
            .expect("empty chunk size");
        let size = usize::from_str_radix(size_text, 16).expect("invalid chunk size");
        cursor = line_end + 2;
        if size == 0 {
            break;
        }
        let end = cursor.checked_add(size).expect("chunk size overflow");
        assert!(end + 2 <= bytes.len(), "truncated chunked response");
        decoded.extend_from_slice(&bytes[cursor..end]);
        assert_eq!(&bytes[end..end + 2], b"\r\n", "invalid chunk terminator");
        cursor = end + 2;
    }
    decoded
}

fn fixture(name: &str) -> PathBuf {
    Path::new(env!("CARGO_MANIFEST_DIR"))
        .join("tests/fixtures/static")
        .join(name)
}

fn free_port() -> u16 {
    for _ in 0..20_000 {
        let port = NEXT_PORT.fetch_add(1, Ordering::Relaxed);
        if port >= 3_100 && TcpListener::bind(("127.0.0.1", port)).is_ok() {
            return port;
        }
    }
    panic!("no free loopback port at or above 3100");
}

fn take_output(child: &mut Child) -> String {
    let mut output = String::new();
    if let Some(mut stdout) = child.stdout.take() {
        let _ = stdout.read_to_string(&mut output);
    }
    if let Some(mut stderr) = child.stderr.take() {
        let mut value = String::new();
        let _ = stderr.read_to_string(&mut value);
        output.push_str(&value);
    }
    output
}

struct ScratchDir(PathBuf);

impl ScratchDir {
    fn new(label: &str) -> Self {
        let sequence = NEXT_SCRATCH.fetch_add(1, Ordering::Relaxed);
        let path = env::temp_dir().join(format!(
            "bunfork-static-http-{label}-{}-{sequence}",
            std::process::id()
        ));
        if path.exists() {
            fs::remove_dir_all(&path).expect("remove stale test scratch directory");
        }
        fs::create_dir(&path).expect("create test scratch directory");
        Self(path)
    }

    fn path(&self) -> &Path {
        &self.0
    }
}

impl Drop for ScratchDir {
    fn drop(&mut self) {
        let _ = fs::remove_dir_all(&self.0);
    }
}