acorn-lib 0.3.2

ACORN library
use super::*;
use crate::test::server::TestServer;
use crate::util::constants::env::{HTTP_MAX_RETRIES, HTTP_STALL_TIMEOUT};
use acorn_core::prelude::Arc;
use alloc::sync::Arc as StdArc;
use axum::body::{Body, Bytes};
use axum::http::StatusCode;
use axum::response::Redirect;
use axum::routing::{get, post};
use axum::Router;
use core::sync::atomic::{AtomicUsize, Ordering};
use futures::{stream, StreamExt};
use jiff::SignedDuration;
use std::env;
use std::io::{self, Write};
use std::sync::{Mutex, MutexGuard};

static ENV_LOCK: Mutex<()> = Mutex::new(());
#[derive(Clone)]
struct CaptureWriter(StdArc<Mutex<Vec<u8>>>);
struct RestoreEnv {
    _lock: MutexGuard<'static, ()>,
    max_retries: Option<String>,
    stall_timeout: Option<String>,
}
impl Write for CaptureWriter {
    fn write(&mut self, bytes: &[u8]) -> io::Result<usize> {
        self.0.lock().unwrap().extend_from_slice(bytes);
        Ok(bytes.len())
    }
    fn flush(&mut self) -> io::Result<()> {
        Ok(())
    }
}
impl Drop for RestoreEnv {
    fn drop(&mut self) {
        match &self.max_retries {
            | Some(value) => env::set_var(HTTP_MAX_RETRIES, value),
            | None => env::remove_var(HTTP_MAX_RETRIES),
        }
        match &self.stall_timeout {
            | Some(value) => env::set_var(HTTP_STALL_TIMEOUT, value),
            | None => env::remove_var(HTTP_STALL_TIMEOUT),
        }
    }
}
impl RestoreEnv {
    fn apply(max_retries: Option<&str>, stall_timeout: Option<&str>) -> Self {
        let lock = ENV_LOCK.lock().unwrap();
        let backup = Self {
            _lock: lock,
            max_retries: env::var(HTTP_MAX_RETRIES).ok(),
            stall_timeout: env::var(HTTP_STALL_TIMEOUT).ok(),
        };
        match max_retries {
            | Some(value) => env::set_var(HTTP_MAX_RETRIES, value),
            | None => env::remove_var(HTTP_MAX_RETRIES),
        }
        match stall_timeout {
            | Some(value) => env::set_var(HTTP_STALL_TIMEOUT, value),
            | None => env::remove_var(HTTP_STALL_TIMEOUT),
        }
        backup
    }
}

#[tokio::test]
async fn test_http_get_retries_anonymously_when_configured_credentials_are_rejected() {
    let requests = Arc::new(AtomicUsize::new(0));
    let state = Arc::clone(&requests);
    let router = Router::new().route(
        "/public",
        get(move |headers: HeaderMap| {
            let state = Arc::clone(&state);
            async move {
                state.fetch_add(1, Ordering::SeqCst);
                match headers.contains_key("private-token") {
                    | true => (StatusCode::UNAUTHORIZED, "bad credential"),
                    | false => (StatusCode::OK, "public content"),
                }
            }
        }),
    );
    let server = TestServer::start(router).await.unwrap();
    let response = super::get(format!("{}/public", server.base_url))
        .headers(super::headers([("PRIVATE-TOKEN", "invalid")]))
        .send()
        .await
        .expect("anonymous fallback should succeed");
    assert_eq!(response.status_code, StatusCode::OK.as_u16());
    assert_eq!(response.text().await.unwrap(), "public content");
    assert_eq!(requests.load(Ordering::SeqCst), 2);
    server.stop().await.unwrap();
}
#[test]
fn test_http_response_success_accepts_2xx_and_reports_other_statuses() {
    let response = HttpResponse {
        body: Vec::new(),
        headers: HeaderMap::new(),
        status_code: StatusCode::NO_CONTENT.as_u16(),
    };
    assert_eq!(response.success("read resource").expect("successful response").status_code, 204);
    let error = HttpResponse {
        body: Vec::new(),
        headers: HeaderMap::new(),
        status_code: StatusCode::BAD_GATEWAY.as_u16(),
    }
    .success("read resource")
    .expect_err("failed response");
    assert_eq!(error.to_string(), "Failed to read resource — HTTP 502");
}
#[tokio::test]
async fn test_restricted_request_refuses_redirects() {
    let destination_calls = Arc::new(AtomicUsize::new(0));
    let state = Arc::clone(&destination_calls);
    let router = Router::new()
        .route("/callback", post(|| async { Redirect::temporary("/destination") }))
        .route(
            "/destination",
            post(move || {
                let state = Arc::clone(&state);
                async move {
                    state.fetch_add(1, Ordering::SeqCst);
                    StatusCode::OK
                }
            }),
        );
    let server = TestServer::start(router).await.unwrap();
    let request = HttpRequest::init()
        .allow_anonymous_fallback(false)
        .follow_redirects(false)
        .json_body(serde_json::json!({"event": "created"}))
        .max_response_bytes(1024)
        .method(HttpMethod::Post)
        .sensitive_url(true)
        .url(format!("{}/callback?sig=secret", server.base_url))
        .build();
    let response = ReqwestHttpService::default().execute(request).await.unwrap();
    assert_eq!(response.status_code, StatusCode::TEMPORARY_REDIRECT.as_u16());
    assert_eq!(destination_calls.load(Ordering::SeqCst), 0);
    server.stop().await.unwrap();
}
#[tokio::test(flavor = "current_thread")]
async fn test_restricted_request_rejects_oversized_response_without_leaking_url() {
    let captured = StdArc::new(Mutex::new(Vec::new()));
    let output = StdArc::clone(&captured);
    let subscriber = tracing_subscriber::fmt()
        .with_ansi(false)
        .without_time()
        .with_max_level(tracing::Level::TRACE)
        .with_writer(move || CaptureWriter(StdArc::clone(&output)))
        .finish();
    let guard = tracing::subscriber::set_default(subscriber);
    let router = Router::new().route("/callback", post(|| async { "x".repeat(128) }));
    let server = TestServer::start(router).await.unwrap();
    let request = HttpRequest::init()
        .allow_anonymous_fallback(false)
        .follow_redirects(false)
        .max_response_bytes(16)
        .method(HttpMethod::Post)
        .sensitive_url(true)
        .url(format!("{}/callback?sig=secret", server.base_url))
        .build();
    let diagnostic = format!("{request:?}");
    assert!(!diagnostic.contains("sig=secret"));
    let error = ReqwestHttpService::default().execute(request).await.unwrap_err();
    assert!(error.to_string().contains("configured 16-byte limit"));
    assert!(!error.to_string().contains("sig=secret"));
    server.stop().await.unwrap();
    drop(guard);
    let logs = String::from_utf8(captured.lock().unwrap().clone()).unwrap();
    assert!(logs.contains("REDACTED WEBHOOK URL"));
    assert!(!logs.contains("sig=secret"));
}
#[test]
fn test_shared_http_policy_defaults() {
    let _backup = RestoreEnv::apply(None, None);
    let policy = policy::shared_http_policy();
    assert_eq!(policy.timeout.as_secs(), 30);
    assert_eq!(policy.max_retries, 2);
    assert_eq!(policy.max_attempts(), 3);
}
#[test]
fn test_shared_http_policy_reads_env_overrides_with_min_stall_clamp() {
    let _backup = RestoreEnv::apply(Some("4"), Some("0"));
    let policy = policy::shared_http_policy();
    assert_eq!(policy.max_retries, 4);
    assert_eq!(policy.stall_timeout, SignedDuration::from_secs(1));
}
#[tokio::test(flavor = "current_thread")]
async fn test_streaming_download_aborts_when_server_stalls_between_bytes() {
    let _backup = RestoreEnv::apply(Some("0"), Some("1"));
    let router = Router::new().route(
        "/stall",
        get(|| async {
            Body::from_stream(stream::once(async { Ok::<_, std::io::Error>(Bytes::from_static(b"incomplete")) }).chain(stream::pending()))
        }),
    );
    let server = TestServer::start(router).await.unwrap();
    let output = std::env::temp_dir().join("acorn-stall-download-test.bin");
    let error = download_with_progress(&format!("{}/stall", server.base_url), &output, |_, _| {}, None, None, None, None)
        .await
        .expect_err("stalled download should fail");
    assert!(error.to_string().contains("stalled for 1 seconds"));
    let _ = std::fs::remove_file(output);
    server.stop().await.unwrap();
}
#[tokio::test]
async fn test_streaming_download_reports_credential_and_anonymous_failures() {
    let requests = Arc::new(AtomicUsize::new(0));
    let state = Arc::clone(&requests);
    let router = Router::new().route(
        "/private",
        get(move || {
            let state = Arc::clone(&state);
            async move {
                state.fetch_add(1, Ordering::SeqCst);
                (StatusCode::UNAUTHORIZED, "authentication required")
            }
        }),
    );
    let server = TestServer::start(router).await.unwrap();
    let output = std::env::temp_dir().join(format!("acorn-auth-fallback-failure-{}.txt", std::process::id()));
    let why = download_with_progress(
        &format!("{}/private", server.base_url),
        &output,
        |_, _| {},
        Some(super::headers([("PRIVATE-TOKEN", "invalid")])),
        None,
        Some("configured credential rejected"),
        None,
    )
    .await
    .expect_err("private content should reject both attempts");
    let message = why.to_string();
    assert!(message.contains("configured credential rejected"));
    assert!(message.contains("anonymous fallback also failed"));
    assert!(!message.contains("invalid"));
    assert_eq!(requests.load(Ordering::SeqCst), 2);
    let _ = std::fs::remove_file(output);
    server.stop().await.unwrap();
}
#[tokio::test]
async fn test_streaming_download_retries_anonymously_when_configured_credentials_are_rejected() {
    let requests = Arc::new(AtomicUsize::new(0));
    let state = Arc::clone(&requests);
    let router = Router::new().route(
        "/public",
        get(move |headers: HeaderMap| {
            let state = Arc::clone(&state);
            async move {
                state.fetch_add(1, Ordering::SeqCst);
                match headers.contains_key("private-token") {
                    | true => (StatusCode::FORBIDDEN, "bad credential"),
                    | false => (StatusCode::OK, "public content"),
                }
            }
        }),
    );
    let server = TestServer::start(router).await.unwrap();
    let output = std::env::temp_dir().join(format!("acorn-auth-fallback-{}.txt", std::process::id()));
    download_with_progress(
        &format!("{}/public", server.base_url),
        &output,
        |_, _| {},
        Some(super::headers([("PRIVATE-TOKEN", "invalid")])),
        None,
        None,
        None,
    )
    .await
    .expect("anonymous streaming fallback should succeed");
    assert_eq!(std::fs::read_to_string(&output).unwrap(), "public content");
    assert_eq!(requests.load(Ordering::SeqCst), 2);
    let _ = std::fs::remove_file(output);
    server.stop().await.unwrap();
}