use reqsign_core::{Error, Result};
use serde::Deserialize;
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) struct AwsPartition {
pub(crate) id: &'static str,
pub(crate) dns_suffix: &'static str,
}
pub(crate) fn partition_for_region(region: &str) -> Result<AwsPartition> {
if !(3..=64).contains(®ion.len())
|| !region
.bytes()
.all(|byte| byte.is_ascii_lowercase() || byte.is_ascii_digit() || byte == b'-')
|| !region
.as_bytes()
.first()
.is_some_and(u8::is_ascii_alphanumeric)
|| !region
.as_bytes()
.last()
.is_some_and(u8::is_ascii_alphanumeric)
{
return Err(Error::config_invalid("AWS STS signing region is invalid"));
}
let partition = if region.starts_with("cn-") {
AwsPartition {
id: "aws-cn",
dns_suffix: "amazonaws.com.cn",
}
} else if region.starts_with("eusc-") {
AwsPartition {
id: "aws-eusc",
dns_suffix: "amazonaws.eu",
}
} else if region.starts_with("us-isob-") {
AwsPartition {
id: "aws-iso-b",
dns_suffix: "sc2s.sgov.gov",
}
} else if region.starts_with("us-iso-") {
AwsPartition {
id: "aws-iso",
dns_suffix: "c2s.ic.gov",
}
} else if region.starts_with("eu-isoe-") {
AwsPartition {
id: "aws-iso-e",
dns_suffix: "cloud.adc-e.uk",
}
} else if region.starts_with("us-isof-") {
AwsPartition {
id: "aws-iso-f",
dns_suffix: "csp.hci.ic.gov",
}
} else if region.starts_with("us-gov-") {
AwsPartition {
id: "aws-us-gov",
dns_suffix: "amazonaws.com",
}
} else {
AwsPartition {
id: "aws",
dns_suffix: "amazonaws.com",
}
};
Ok(partition)
}
pub fn sts_endpoint(region: Option<&str>, use_regional: bool) -> Result<String> {
if use_regional {
let region =
region.ok_or_else(|| Error::config_invalid("regional STS endpoint requires region"))?;
let partition = partition_for_region(region)?;
return Ok(format!("sts.{region}.{}", partition.dns_suffix));
}
match region {
None => Ok("sts.amazonaws.com".to_string()),
Some(region) => match partition_for_region(region)?.id {
"aws" => Ok("sts.amazonaws.com".to_string()),
"aws-cn" => Ok("sts.amazonaws.com.cn".to_string()),
_ => Err(Error::config_invalid(
"legacy global STS endpoint is unavailable for this AWS partition",
)),
},
}
}
#[derive(Debug, Deserialize)]
pub struct AwsErrorResponse {
#[serde(rename = "Error")]
pub error: AwsError,
}
#[derive(Debug, Deserialize)]
pub struct AwsError {
#[serde(rename = "Code")]
pub code: String,
#[serde(rename = "Message")]
pub message: String,
}
pub fn parse_sts_error(operation: &str, status: http::StatusCode, body: &str) -> Error {
if let Ok(error_resp) = quick_xml::de::from_str::<AwsErrorResponse>(body) {
let code = &error_resp.error.code;
let message = &error_resp.error.message;
let detail = format!("{code}: {message}");
let error = match code.as_str() {
"AccessDenied" | "UnauthorizedAccess" | "Forbidden" => Error::permission_denied(detail),
"ExpiredToken"
| "TokenRefreshRequired"
| "InvalidToken"
| "InvalidIdentityToken"
| "IDPRejectedClaim"
| "IDPCommunicationError" => Error::credential_invalid(detail),
"InvalidParameterValue" | "MissingParameter" | "InvalidParameterCombination" => {
Error::config_invalid(detail)
}
"Throttling" | "RequestLimitExceeded" | "TooManyRequestsException" => {
Error::rate_limited(detail)
}
"ServiceUnavailable" | "InternalError" | "InternalFailure" => {
Error::unexpected(detail).set_retryable(true)
}
"InvalidRequest" | "MalformedQueryString" => Error::request_invalid(detail),
_ => Error::unexpected(detail),
};
error
.with_context(format!("operation: {operation}"))
.with_context(format!("error_code: {code}"))
} else {
let detail = format!("STS request failed with {status}: {body}");
let mut error = match status.as_u16() {
400..=499 if status == http::StatusCode::FORBIDDEN => Error::permission_denied(detail),
400..=499 if status == http::StatusCode::UNAUTHORIZED => {
Error::credential_invalid(detail)
}
429 => Error::rate_limited(detail),
400..=499 => Error::request_invalid(detail),
500..=599 => Error::unexpected(detail).set_retryable(true),
_ => Error::unexpected(detail),
};
error = error
.with_context(format!("operation: {operation}"))
.with_context(format!("http_status: {status}"));
error
}
}
pub fn parse_imds_error(operation: &str, status: http::StatusCode, body: &str) -> Error {
#[derive(Debug, Deserialize)]
struct ImdsError {
#[serde(rename = "Code")]
code: String,
#[serde(rename = "Message")]
message: String,
}
if let Ok(error) = serde_json::from_str::<ImdsError>(body) {
let err = match error.code.as_str() {
"AssumeRoleUnauthorizedAccess" => Error::permission_denied(format!(
"EC2 instance not authorized to assume role: {}",
error.message
))
.with_context("hint: check if the IAM role has a trust relationship with EC2"),
"InvalidUserData.Malformed" => {
Error::config_invalid(format!("malformed instance metadata: {}", error.message))
}
_ if error.code.contains("Expired") => {
Error::credential_invalid(format!("IMDS credentials expired: {}", error.message))
}
_ => Error::unexpected(format!("IMDS error [{}]: {}", error.code, error.message)),
};
err.with_context(format!("operation: {operation}"))
.with_context(format!("error_code: {}", error.code))
} else {
match status.as_u16() {
401 | 403 => Error::permission_denied(format!("IMDS access denied: {body}"))
.with_context(format!("operation: {operation}"))
.with_context("hint: check if IMDSv2 is required"),
404 => Error::config_invalid("instance metadata not found")
.with_context(format!("operation: {operation}"))
.with_context("hint: are you running on EC2?"),
500..=599 => Error::unexpected(format!("IMDS server error: {body}"))
.with_context(format!("operation: {operation}"))
.set_retryable(true),
_ => Error::unexpected(format!("IMDS request failed: {body}"))
.with_context(format!("operation: {operation}"))
.with_context(format!("http_status: {status}")),
}
}
}
#[cfg(test)]
mod tests {
use reqsign_core::ErrorKind;
use super::*;
#[test]
fn resolves_partition_aware_sts_endpoints() {
let cases = [
("us-east-1", "sts.us-east-1.amazonaws.com"),
("cn-north-1", "sts.cn-north-1.amazonaws.com.cn"),
("eusc-de-east-1", "sts.eusc-de-east-1.amazonaws.eu"),
("us-iso-east-1", "sts.us-iso-east-1.c2s.ic.gov"),
("us-isob-east-1", "sts.us-isob-east-1.sc2s.sgov.gov"),
("eu-isoe-west-1", "sts.eu-isoe-west-1.cloud.adc-e.uk"),
("us-isof-east-1", "sts.us-isof-east-1.csp.hci.ic.gov"),
("us-gov-west-1", "sts.us-gov-west-1.amazonaws.com"),
];
for (region, expected) in cases {
assert_eq!(
sts_endpoint(Some(region), true).expect("regional endpoint must resolve"),
expected
);
}
assert_eq!(
sts_endpoint(None, false).expect("commercial global endpoint must resolve"),
"sts.amazonaws.com"
);
assert_eq!(
sts_endpoint(Some("cn-north-1"), false).expect("China global endpoint must resolve"),
"sts.amazonaws.com.cn"
);
assert_eq!(
sts_endpoint(Some("us-iso-east-1"), false)
.expect_err("isolated partitions must reject the legacy global endpoint")
.kind(),
ErrorKind::ConfigInvalid
);
}
#[test]
fn preserves_sts_error_kinds_and_surfaces_message_text() {
let cases = [
("RegionDisabled", ErrorKind::Unexpected),
("MalformedPolicyDocument", ErrorKind::Unexpected),
("PackedPolicyTooLarge", ErrorKind::Unexpected),
("InvalidRequest", ErrorKind::RequestInvalid),
("InvalidIdentityToken", ErrorKind::CredentialInvalid),
];
for (code, expected_kind) in cases {
let body = format!(
"<ErrorResponse><Error><Code>{code}</Code>\
<Message>diagnostic detail for {code}</Message></Error></ErrorResponse>"
);
let error = parse_sts_error(
"AssumeRoleWithWebIdentity",
http::StatusCode::BAD_REQUEST,
&body,
);
assert_eq!(error.kind(), expected_kind);
let debug = format!("{error:?}");
assert!(debug.contains(&format!("error_code: {code}")));
assert!(debug.contains(&format!("diagnostic detail for {code}")));
}
let error = parse_sts_error(
"AssumeRole",
http::StatusCode::FORBIDDEN,
"raw diagnostic body",
);
let debug = format!("{error:?}");
assert_eq!(error.kind(), ErrorKind::PermissionDenied);
assert!(debug.contains("raw diagnostic body"));
}
}