harn-vm 0.10.42

Async bytecode virtual machine for the Harn programming language
Documentation
use std::collections::BTreeMap;

use aws_credential_types::Credentials;
use aws_sigv4::http_request::{
    sign as sign_http_request, PayloadChecksumKind, SignableBody, SignableRequest, SigningSettings,
};
use aws_sigv4::sign::v4;
use chrono::{DateTime, Utc};
use url::{Host, Url};

#[derive(Clone, PartialEq, Eq)]
pub(crate) struct AwsSigV4Credentials {
    pub access_key_id: String,
    pub secret_access_key: String,
    pub session_token: Option<String>,
}

impl std::fmt::Debug for AwsSigV4Credentials {
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
        f.debug_struct("AwsSigV4Credentials")
            .field("access_key_id", &"<redacted>")
            .field("secret_access_key", &"<redacted>")
            .field(
                "session_token",
                &self.session_token.as_ref().map(|_| "<redacted>"),
            )
            .finish()
    }
}

pub(crate) struct AwsSigV4Input<'a> {
    pub credentials: &'a AwsSigV4Credentials,
    pub method: &'a str,
    pub url: &'a str,
    pub service: &'a str,
    pub region: &'a str,
    pub headers: &'a BTreeMap<String, String>,
    pub body: &'a [u8],
    pub timestamp: DateTime<Utc>,
}

#[derive(Clone, PartialEq, Eq)]
pub(crate) struct AwsSigV4SignedRequest {
    pub headers: BTreeMap<String, String>,
    pub authorization: String,
    pub amz_date: String,
    pub content_sha256: String,
    pub security_token: Option<String>,
    pub signed_headers: String,
}

impl std::fmt::Debug for AwsSigV4SignedRequest {
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
        let mut headers = self.headers.clone();
        for (name, value) in headers.iter_mut() {
            let lower = name.to_ascii_lowercase();
            if lower == "authorization" || lower.contains("token") {
                *value = "<redacted>".to_string();
            }
        }
        f.debug_struct("AwsSigV4SignedRequest")
            .field("headers", &headers)
            .field("authorization", &"<redacted>")
            .field("amz_date", &self.amz_date)
            .field("content_sha256", &self.content_sha256)
            .field(
                "security_token",
                &self.security_token.as_ref().map(|_| "<redacted>"),
            )
            .field("signed_headers", &self.signed_headers)
            .finish()
    }
}

pub(crate) fn sign(input: AwsSigV4Input<'_>) -> Result<AwsSigV4SignedRequest, String> {
    validate_required("access_key_id", &input.credentials.access_key_id)?;
    validate_required("secret_access_key", &input.credentials.secret_access_key)?;
    let method = input.method.trim().to_ascii_uppercase();
    validate_required("method", &method)?;
    validate_required("service", input.service)?;
    validate_required("region", input.region)?;

    let parsed = Url::parse(input.url)
        .map_err(|_| "url must be an absolute URL with a scheme and host".to_string())?;
    let host = host_header(&parsed)?;
    for (name, value) in input.headers {
        let parsed_name = reqwest::header::HeaderName::from_bytes(name.as_bytes())
            .map_err(|_| format!("invalid header name `{name}`"))?;
        reqwest::header::HeaderValue::from_str(value)
            .map_err(|_| format!("headers.{name} contains an invalid value"))?;
        let lower = parsed_name.as_str();
        match lower {
            "authorization" => {
                return Err("headers.Authorization is generated by aws_sigv4_headers".to_string());
            }
            "x-amz-date" | "x-amz-content-sha256" | "x-amz-security-token" => {
                return Err(format!("headers.{name} is generated by aws_sigv4_headers"));
            }
            "host" if value.trim() != host => {
                return Err("headers.Host must match the URL host".to_string());
            }
            _ => {}
        }
    }

    let identity = Credentials::new(
        input.credentials.access_key_id.clone(),
        input.credentials.secret_access_key.clone(),
        input
            .credentials
            .session_token
            .clone()
            .filter(|token| !token.trim().is_empty()),
        None,
        "harn",
    )
    .into();
    let mut settings = SigningSettings::default();
    settings.payload_checksum_kind = PayloadChecksumKind::XAmzSha256;
    let params = v4::SigningParams::builder()
        .identity(&identity)
        .region(input.region.trim())
        .name(input.service.trim())
        .time(input.timestamp.into())
        .settings(settings)
        .build()
        .map_err(|error| format!("invalid signing parameters: {error}"))?
        .into();
    let request = SignableRequest::new(
        &method,
        input.url,
        input
            .headers
            .iter()
            .map(|(name, value)| (name.as_str(), value.as_str())),
        SignableBody::Bytes(input.body),
    )
    .map_err(|error| format!("request cannot be signed: {error}"))?;
    let (instructions, _) = sign_http_request(request, &params)
        .map_err(|error| format!("request signing failed: {error}"))?
        .into_parts();

    let mut headers = input
        .headers
        .iter()
        .filter(|(name, _)| !name.eq_ignore_ascii_case("host"))
        .map(|(name, value)| (name.clone(), value.clone()))
        .collect::<BTreeMap<_, _>>();
    headers.insert("Host".to_string(), host);
    for (name, value) in instructions.headers() {
        headers.insert(output_header_name(name).to_string(), value.to_string());
    }
    let generated = |name: &str| {
        headers
            .iter()
            .find(|(candidate, _)| candidate.eq_ignore_ascii_case(name))
            .map(|(_, value)| value.clone())
            .ok_or_else(|| format!("AWS signer did not generate {name}"))
    };
    let authorization = generated("authorization")?;
    let amz_date = generated("x-amz-date")?;
    let content_sha256 = generated("x-amz-content-sha256")?;
    let security_token = headers
        .iter()
        .find(|(name, _)| name.eq_ignore_ascii_case("x-amz-security-token"))
        .map(|(_, value)| value.clone());
    let signed_headers = authorization
        .split("SignedHeaders=")
        .nth(1)
        .and_then(|tail| tail.split(',').next())
        .ok_or_else(|| "AWS signer returned malformed authorization metadata".to_string())?
        .to_string();

    Ok(AwsSigV4SignedRequest {
        headers,
        authorization,
        amz_date,
        content_sha256,
        security_token,
        signed_headers,
    })
}

fn validate_required(label: &str, value: &str) -> Result<(), String> {
    if value.trim().is_empty() {
        Err(format!("{label} is required"))
    } else {
        Ok(())
    }
}

fn host_header(url: &Url) -> Result<String, String> {
    let mut host = match url.host() {
        Some(Host::Domain(domain)) => domain.to_string(),
        Some(Host::Ipv4(addr)) => addr.to_string(),
        Some(Host::Ipv6(addr)) => format!("[{addr}]"),
        None => return Err("url must include a host".to_string()),
    };
    if let Some(port) = url.port() {
        host.push(':');
        host.push_str(&port.to_string());
    }
    Ok(host)
}

fn output_header_name(name: &str) -> &str {
    match name {
        "authorization" => "Authorization",
        "x-amz-content-sha256" => "X-Amz-Content-Sha256",
        "x-amz-date" => "X-Amz-Date",
        "x-amz-security-token" => "X-Amz-Security-Token",
        other => other,
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use chrono::TimeZone;

    fn credentials() -> AwsSigV4Credentials {
        AwsSigV4Credentials {
            access_key_id: "AKIDEXAMPLE".to_string(),
            secret_access_key: "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY".to_string(),
            session_token: None,
        }
    }

    fn fixed_time() -> DateTime<Utc> {
        Utc.with_ymd_and_hms(2015, 8, 30, 12, 36, 0).unwrap()
    }

    #[test]
    fn signs_standard_headers_with_aws_signer() {
        let headers = BTreeMap::from([
            (
                "Content-Type".to_string(),
                "application/x-amz-json-1.1".to_string(),
            ),
            (
                "host".to_string(),
                "service.us-east-1.amazonaws.com".to_string(),
            ),
        ]);
        let signed = sign(AwsSigV4Input {
            credentials: &credentials(),
            method: "POST",
            url: "https://service.us-east-1.amazonaws.com/",
            service: "service",
            region: "us-east-1",
            headers: &headers,
            body: br#"{"hello":"world"}"#,
            timestamp: fixed_time(),
        })
        .expect("signed request");

        assert_eq!(signed.amz_date, "20150830T123600Z");
        assert_eq!(
            signed.signed_headers,
            "content-type;host;x-amz-content-sha256;x-amz-date"
        );
        assert_eq!(
            signed.content_sha256,
            "93a23971a914e5eacbf0a8d25154cda309c3c1c72fbb9914d47c60f3cb681588"
        );
        assert!(signed
            .authorization
            .contains("Credential=AKIDEXAMPLE/20150830/us-east-1/service/aws4_request"));
        assert_eq!(
            signed
                .headers
                .keys()
                .filter(|name| name.eq_ignore_ascii_case("host"))
                .count(),
            1
        );
    }

    #[test]
    fn signs_session_token() {
        let mut credentials = credentials();
        credentials.session_token = Some("session-token".to_string());
        let signed = sign(AwsSigV4Input {
            credentials: &credentials,
            method: "GET",
            url: "https://service.us-east-1.amazonaws.com/",
            service: "service",
            region: "us-east-1",
            headers: &BTreeMap::new(),
            body: b"",
            timestamp: fixed_time(),
        })
        .expect("signed request");

        assert_eq!(signed.security_token.as_deref(), Some("session-token"));
        assert!(signed.signed_headers.contains("x-amz-security-token"));
    }

    #[test]
    fn validation_errors_do_not_include_credentials() {
        let credentials = AwsSigV4Credentials {
            access_key_id: "AKIAIOSFODNN7EXAMPLE".to_string(),
            secret_access_key: "secret-that-must-not-leak".to_string(),
            session_token: Some("session-that-must-not-leak".to_string()),
        };
        let error = sign(AwsSigV4Input {
            credentials: &credentials,
            method: "POST",
            url: "not a url",
            service: "service",
            region: "us-east-1",
            headers: &BTreeMap::new(),
            body: b"",
            timestamp: fixed_time(),
        })
        .expect_err("invalid url should fail");

        for secret in [
            "AKIAIOSFODNN7EXAMPLE",
            "secret-that-must-not-leak",
            "session-that-must-not-leak",
        ] {
            assert!(!error.contains(secret));
        }

        let invalid_headers =
            BTreeMap::from([("X-Test".to_string(), "safe\r\nInjected: yes".to_string())]);
        let error = sign(AwsSigV4Input {
            credentials: &credentials,
            method: "POST",
            url: "https://example.amazonaws.com/",
            service: "service",
            region: "us-east-1",
            headers: &invalid_headers,
            body: b"",
            timestamp: fixed_time(),
        })
        .expect_err("header injection must fail");
        assert!(error.contains("invalid value"));
    }
}