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, ¶ms)
.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"));
}
}