use aws_sdk_s3::config::Credentials;
use axum::extract::FromRequestParts;
use axum::http::request::Parts;
use std::convert::Infallible;
pub const HEADER_ACCESS_KEY_ID: &str = "x-dial9-aws-access-key-id";
pub const HEADER_SECRET_ACCESS_KEY: &str = "x-dial9-aws-secret-access-key";
pub const HEADER_SESSION_TOKEN: &str = "x-dial9-aws-session-token";
pub const HEADER_REGION: &str = "x-dial9-aws-region";
pub const HEADER_ROLE_ARN: &str = "x-dial9-aws-role-arn";
pub const QUERY_ROLE_ARN: &str = "aws_role_arn";
pub const QUERY_REGION: &str = "aws_region";
pub const ASSUME_ROLE_SESSION_NAME: &str = "dial9-viewer";
const PROVIDER_NAME: &str = "dial9-byo";
#[derive(Clone, Debug)]
pub struct TempCredentials {
pub credentials: Credentials,
pub region: Option<String>,
}
impl TempCredentials {
pub fn new(
access_key_id: impl Into<String>,
secret_access_key: impl Into<String>,
session_token: Option<String>,
region: Option<String>,
) -> Self {
Self {
credentials: Credentials::new(
access_key_id,
secret_access_key,
session_token,
None, PROVIDER_NAME,
),
region,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RoleArn(String);
impl RoleArn {
pub fn as_str(&self) -> &str {
&self.0
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum CredSource {
Default,
Static(TempCredentials),
AssumeRole {
role_arn: RoleArn,
region: Option<String>,
},
}
impl PartialEq for TempCredentials {
fn eq(&self, other: &Self) -> bool {
self.credentials.access_key_id() == other.credentials.access_key_id()
&& self.credentials.secret_access_key() == other.credentials.secret_access_key()
&& self.credentials.session_token() == other.credentials.session_token()
&& self.region == other.region
}
}
impl Eq for TempCredentials {}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum CredError {
Incomplete,
Malformed,
InvalidRegion,
ConflictingCredentials,
InvalidRoleArn,
}
impl CredError {
pub fn message(&self) -> &'static str {
match self {
CredError::Incomplete => {
"incomplete credentials: both access key id and secret access key are required"
}
CredError::Malformed => "malformed credential header",
CredError::InvalidRegion => "invalid region",
CredError::ConflictingCredentials => {
"supply either bring-your-own credentials or a role ARN, not both"
}
CredError::InvalidRoleArn => "invalid role ARN",
}
}
}
fn is_valid_role_arn(arn: &str) -> bool {
if arn.is_empty() || arn.len() > 2048 {
return false;
}
let parts: Vec<&str> = arn.splitn(6, ':').collect();
if parts.len() != 6 {
return false;
}
let [prefix, partition, service, region, account, resource] = parts[..] else {
return false;
};
if prefix != "arn" {
return false;
}
if !matches!(partition, "aws" | "aws-cn" | "aws-us-gov") {
return false;
}
if service != "iam" || !region.is_empty() {
return false;
}
if account.len() != 12 || !account.bytes().all(|b| b.is_ascii_digit()) {
return false;
}
match resource.strip_prefix("role/") {
Some(rest) => !rest.is_empty() && !rest.contains('*') && !rest.contains('?'),
None => false,
}
}
fn is_valid_region(region: &str) -> bool {
if region.is_empty() || region.len() > 40 {
return false;
}
let bytes = region.as_bytes();
if !bytes[0].is_ascii_lowercase() {
return false;
}
if !bytes[bytes.len() - 1].is_ascii_alphanumeric() {
return false;
}
let mut prev_hyphen = false;
for &b in bytes {
if b == b'-' {
if prev_hyphen {
return false;
}
prev_hyphen = true;
} else if b.is_ascii_lowercase() || b.is_ascii_digit() {
prev_hyphen = false;
} else {
return false;
}
}
true
}
pub struct MaybeCreds(pub Result<CredSource, CredError>);
impl<S: Send + Sync> FromRequestParts<S> for MaybeCreds {
type Rejection = Infallible;
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Infallible> {
Ok(MaybeCreds(parse_cred_inputs(
&parts.headers,
parts.uri.query(),
)))
}
}
struct QueryCreds {
role_arn: Option<String>,
region: Option<String>,
}
fn parse_query_creds(query: Option<&str>) -> QueryCreds {
let mut role_arn = None;
let mut region = None;
if let Some(q) = query {
for (k, v) in form_urlencoded::parse(q.as_bytes()) {
match k.as_ref() {
QUERY_ROLE_ARN if !v.is_empty() => role_arn = Some(v.into_owned()),
QUERY_REGION if !v.is_empty() => region = Some(v.into_owned()),
_ => {}
}
}
}
QueryCreds { role_arn, region }
}
pub fn parse_cred_inputs(
headers: &axum::http::HeaderMap,
query: Option<&str>,
) -> Result<CredSource, CredError> {
let get = |name: &str| -> Result<Option<String>, CredError> {
match headers.get(name) {
None => Ok(None),
Some(v) => v
.to_str()
.map(|s| Some(s.to_string()))
.map_err(|_| CredError::Malformed),
}
};
let access_key_id = get(HEADER_ACCESS_KEY_ID)?;
let secret_access_key = get(HEADER_SECRET_ACCESS_KEY)?;
let session_token = get(HEADER_SESSION_TOKEN)?;
let header_role_arn = get(HEADER_ROLE_ARN)?.filter(|s| !s.is_empty());
let header_region = get(HEADER_REGION)?.filter(|s| !s.is_empty());
let QueryCreds {
role_arn: query_role_arn,
region: query_region,
} = parse_query_creds(query);
let role_arn = match (header_role_arn, query_role_arn) {
(Some(_), Some(_)) => return Err(CredError::ConflictingCredentials),
(Some(a), None) | (None, Some(a)) => Some(a),
(None, None) => None,
};
let region = match header_region.or(query_region) {
Some(r) if !is_valid_region(&r) => return Err(CredError::InvalidRegion),
other => other,
};
let has_byoc = access_key_id.is_some() || secret_access_key.is_some();
if has_byoc && role_arn.is_some() {
return Err(CredError::ConflictingCredentials);
}
if let Some(arn) = role_arn {
if !is_valid_role_arn(&arn) {
return Err(CredError::InvalidRoleArn);
}
return Ok(CredSource::AssumeRole {
role_arn: RoleArn(arn),
region,
});
}
match (access_key_id, secret_access_key) {
(None, None) => Ok(CredSource::Default),
(Some(akid), Some(secret)) => Ok(CredSource::Static(TempCredentials::new(
akid,
secret,
session_token.filter(|s| !s.is_empty()),
region,
))),
_ => Err(CredError::Incomplete),
}
}
pub trait RoleAssumer: Send + Sync {
fn assume_role<'a>(
&'a self,
role_arn: &'a RoleArn,
region: Option<&'a str>,
) -> std::pin::Pin<
Box<dyn std::future::Future<Output = Result<TempCredentials, AssumeRoleError>> + Send + 'a>,
>;
}
#[derive(Debug)]
pub struct AssumeRoleError(pub String);
impl std::fmt::Display for AssumeRoleError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "assume-role failed: {}", self.0)
}
}
impl std::error::Error for AssumeRoleError {}
pub struct StsRoleAssumer {
config: aws_config::SdkConfig,
}
impl StsRoleAssumer {
pub async fn from_env() -> Self {
Self::from_config(aws_config::load_defaults(aws_config::BehaviorVersion::latest()).await)
}
pub fn from_config(config: aws_config::SdkConfig) -> Self {
Self { config }
}
}
impl RoleAssumer for StsRoleAssumer {
fn assume_role<'a>(
&'a self,
role_arn: &'a RoleArn,
region: Option<&'a str>,
) -> std::pin::Pin<
Box<dyn std::future::Future<Output = Result<TempCredentials, AssumeRoleError>> + Send + 'a>,
> {
use aws_sdk_s3::config::ProvideCredentials;
Box::pin(async move {
let mut builder = aws_config::sts::AssumeRoleProvider::builder(role_arn.as_str())
.configure(&self.config)
.session_name(ASSUME_ROLE_SESSION_NAME);
if let Some(region) = region {
builder = builder.region(aws_sdk_s3::config::Region::new(region.to_string()));
}
let provider = builder.build().await;
let creds = provider
.provide_credentials()
.await
.map_err(|e| AssumeRoleError(format!("{e}")))?;
Ok(TempCredentials::new(
creds.access_key_id(),
creds.secret_access_key(),
creds.session_token().map(str::to_string),
region.map(str::to_string),
))
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use axum::http::HeaderMap;
fn headers(pairs: &[(&'static str, &str)]) -> HeaderMap {
let mut h = HeaderMap::new();
for (k, v) in pairs {
h.insert(*k, v.parse().unwrap());
}
h
}
fn parse_cred_headers(headers: &HeaderMap) -> Result<CredSource, CredError> {
parse_cred_inputs(headers, None)
}
fn expect_static(src: CredSource) -> TempCredentials {
match src {
CredSource::Static(t) => t,
other => panic!("expected CredSource::Static, got {other:?}"),
}
}
#[test]
fn absent_headers_yield_default() {
let parsed = parse_cred_headers(&HeaderMap::new()).unwrap();
assert_eq!(parsed, CredSource::Default);
}
#[test]
fn full_credentials_parse() {
let h = headers(&[
(HEADER_ACCESS_KEY_ID, "AKIA"),
(HEADER_SECRET_ACCESS_KEY, "secret"),
(HEADER_SESSION_TOKEN, "token"),
(HEADER_REGION, "us-west-2"),
]);
let creds = expect_static(parse_cred_headers(&h).unwrap());
assert_eq!(creds.credentials.access_key_id(), "AKIA");
assert_eq!(creds.credentials.secret_access_key(), "secret");
assert_eq!(creds.credentials.session_token(), Some("token"));
assert_eq!(creds.region.as_deref(), Some("us-west-2"));
}
#[test]
fn long_lived_keys_without_token_or_region() {
let h = headers(&[
(HEADER_ACCESS_KEY_ID, "AKIA"),
(HEADER_SECRET_ACCESS_KEY, "secret"),
]);
let creds = expect_static(parse_cred_headers(&h).unwrap());
assert_eq!(creds.credentials.session_token(), None);
assert_eq!(creds.region, None);
}
#[test]
fn empty_token_and_region_treated_as_absent() {
let h = headers(&[
(HEADER_ACCESS_KEY_ID, "AKIA"),
(HEADER_SECRET_ACCESS_KEY, "secret"),
(HEADER_SESSION_TOKEN, ""),
(HEADER_REGION, ""),
]);
let creds = expect_static(parse_cred_headers(&h).unwrap());
assert_eq!(creds.credentials.session_token(), None);
assert_eq!(creds.region, None);
}
#[test]
fn akid_without_secret_is_incomplete() {
let h = headers(&[(HEADER_ACCESS_KEY_ID, "AKIA")]);
assert!(matches!(parse_cred_headers(&h), Err(CredError::Incomplete)));
}
#[test]
fn secret_without_akid_is_incomplete() {
let h = headers(&[(HEADER_SECRET_ACCESS_KEY, "secret")]);
assert!(matches!(parse_cred_headers(&h), Err(CredError::Incomplete)));
}
#[test]
fn valid_region_charset_accepted() {
for region in [
"us-east-1",
"ap-southeast-2",
"eu-central-1",
"us-gov-west-1",
] {
let h = headers(&[
(HEADER_ACCESS_KEY_ID, "AKIA"),
(HEADER_SECRET_ACCESS_KEY, "secret"),
(HEADER_REGION, region),
]);
let creds = expect_static(parse_cred_headers(&h).unwrap());
assert_eq!(creds.region.as_deref(), Some(region));
}
}
#[test]
fn invalid_region_rejected() {
for region in [
"US-EAST-1",
"evil.com",
"us-east-1/../foo",
"us east 1",
"us_east_1",
"-",
"-us-east-1",
"us-east-",
"us--east-1",
"1-east-1",
] {
let h = headers(&[
(HEADER_ACCESS_KEY_ID, "AKIA"),
(HEADER_SECRET_ACCESS_KEY, "secret"),
(HEADER_REGION, region),
]);
assert!(
matches!(parse_cred_headers(&h), Err(CredError::InvalidRegion)),
"expected {region:?} to be rejected"
);
}
}
#[test]
fn overlong_region_rejected() {
let h = headers(&[
(HEADER_ACCESS_KEY_ID, "AKIA"),
(HEADER_SECRET_ACCESS_KEY, "secret"),
(HEADER_REGION, &"a".repeat(41)),
]);
assert!(matches!(
parse_cred_headers(&h),
Err(CredError::InvalidRegion)
));
}
#[test]
fn role_arn_parses_to_assume_role() {
let h = headers(&[
(
HEADER_ROLE_ARN,
"arn:aws:iam::123456789012:role/dial9-reader",
),
(HEADER_REGION, "us-east-1"),
]);
match parse_cred_headers(&h).unwrap() {
CredSource::AssumeRole { role_arn, region } => {
assert_eq!(
role_arn.as_str(),
"arn:aws:iam::123456789012:role/dial9-reader"
);
assert_eq!(region.as_deref(), Some("us-east-1"));
}
other => panic!("expected AssumeRole, got {other:?}"),
}
}
#[test]
fn role_arn_without_region_is_allowed() {
let h = headers(&[(HEADER_ROLE_ARN, "arn:aws:iam::123456789012:role/r")]);
match parse_cred_headers(&h).unwrap() {
CredSource::AssumeRole { region, .. } => assert_eq!(region, None),
other => panic!("expected AssumeRole, got {other:?}"),
}
}
#[test]
fn empty_role_arn_treated_as_absent() {
let h = headers(&[(HEADER_ROLE_ARN, "")]);
assert_eq!(parse_cred_headers(&h).unwrap(), CredSource::Default);
}
#[test]
fn byoc_and_role_arn_together_conflict() {
let h = headers(&[
(HEADER_ACCESS_KEY_ID, "AKIA"),
(HEADER_SECRET_ACCESS_KEY, "secret"),
(HEADER_ROLE_ARN, "arn:aws:iam::123456789012:role/r"),
]);
assert!(matches!(
parse_cred_headers(&h),
Err(CredError::ConflictingCredentials)
));
}
#[test]
fn lone_akid_with_role_arn_still_conflicts() {
let h = headers(&[
(HEADER_ACCESS_KEY_ID, "AKIA"),
(HEADER_ROLE_ARN, "arn:aws:iam::123456789012:role/r"),
]);
assert!(matches!(
parse_cred_headers(&h),
Err(CredError::ConflictingCredentials)
));
}
#[test]
fn valid_role_arns_accepted() {
for arn in [
"arn:aws:iam::123456789012:role/dial9-reader",
"arn:aws:iam::123456789012:role/path/to/Reader_Role",
"arn:aws-cn:iam::123456789012:role/r",
"arn:aws-us-gov:iam::123456789012:role/r",
] {
assert!(is_valid_role_arn(arn), "expected {arn:?} to be valid");
}
}
#[test]
fn invalid_role_arns_rejected() {
for arn in [
"",
"not-an-arn",
"arn:aws:iam::123456789012:user/bob", "arn:aws:s3:::bucket/key", "arn:aws:iam:us-east-1:123456789012:role/r", "arn:aws:iam::12345:role/r", "arn:aws:iam::123456789012:role/", "arn:aws:iam::123456789012:role/*", "arn:evil:iam::123456789012:role/r", ] {
assert!(!is_valid_role_arn(arn), "expected {arn:?} to be rejected");
}
}
#[test]
fn overlong_role_arn_rejected() {
let arn = format!("arn:aws:iam::123456789012:role/{}", "a".repeat(2048));
assert!(!is_valid_role_arn(&arn));
}
#[test]
fn role_arn_from_query_parses_to_assume_role() {
let q = "aws_role_arn=arn%3Aaws%3Aiam%3A%3A123456789012%3Arole%2Fdial9-reader\
&aws_region=us-west-2&bucket=traces";
match parse_cred_inputs(&HeaderMap::new(), Some(q)).unwrap() {
CredSource::AssumeRole { role_arn, region } => {
assert_eq!(
role_arn.as_str(),
"arn:aws:iam::123456789012:role/dial9-reader"
);
assert_eq!(region.as_deref(), Some("us-west-2"));
}
other => panic!("expected AssumeRole, got {other:?}"),
}
}
#[test]
fn unrelated_query_params_ignored() {
let q = "bucket=traces&prefix=dial9-traces&tz=UTC";
assert_eq!(
parse_cred_inputs(&HeaderMap::new(), Some(q)).unwrap(),
CredSource::Default
);
}
#[test]
fn invalid_role_arn_from_query_rejected() {
let q = "aws_role_arn=not-an-arn";
assert!(matches!(
parse_cred_inputs(&HeaderMap::new(), Some(q)),
Err(CredError::InvalidRoleArn)
));
}
#[test]
fn invalid_region_from_query_rejected() {
let q = "aws_role_arn=arn%3Aaws%3Aiam%3A%3A123456789012%3Arole%2Fr&aws_region=US-EAST-1";
assert!(matches!(
parse_cred_inputs(&HeaderMap::new(), Some(q)),
Err(CredError::InvalidRegion)
));
}
#[test]
fn role_arn_in_both_header_and_query_conflicts() {
let h = headers(&[(HEADER_ROLE_ARN, "arn:aws:iam::123456789012:role/h")]);
let q = "aws_role_arn=arn%3Aaws%3Aiam%3A%3A123456789012%3Arole%2Fq";
assert!(matches!(
parse_cred_inputs(&h, Some(q)),
Err(CredError::ConflictingCredentials)
));
}
#[test]
fn byoc_headers_with_query_role_arn_conflict() {
let h = headers(&[
(HEADER_ACCESS_KEY_ID, "AKIA"),
(HEADER_SECRET_ACCESS_KEY, "secret"),
]);
let q = "aws_role_arn=arn%3Aaws%3Aiam%3A%3A123456789012%3Arole%2Fr";
assert!(matches!(
parse_cred_inputs(&h, Some(q)),
Err(CredError::ConflictingCredentials)
));
}
#[test]
fn header_role_arn_still_works_with_other_query_params() {
let h = headers(&[(HEADER_ROLE_ARN, "arn:aws:iam::123456789012:role/r")]);
match parse_cred_inputs(&h, Some("bucket=traces")).unwrap() {
CredSource::AssumeRole { role_arn, .. } => {
assert_eq!(role_arn.as_str(), "arn:aws:iam::123456789012:role/r");
}
other => panic!("expected AssumeRole, got {other:?}"),
}
}
#[test]
fn query_region_alone_is_default() {
assert_eq!(
parse_cred_inputs(&HeaderMap::new(), Some("aws_region=us-east-1")).unwrap(),
CredSource::Default
);
}
#[test]
fn header_region_wins_over_query_region() {
let h = headers(&[
(HEADER_ROLE_ARN, "arn:aws:iam::123456789012:role/r"),
(HEADER_REGION, "eu-west-1"),
]);
match parse_cred_inputs(&h, Some("aws_region=us-east-1")).unwrap() {
CredSource::AssumeRole { region, .. } => {
assert_eq!(region.as_deref(), Some("eu-west-1"));
}
other => panic!("expected AssumeRole, got {other:?}"),
}
}
#[tokio::test]
async fn sts_role_assumer_mints_assumed_credentials() {
use aws_smithy_http_client::test_util::{ReplayEvent, StaticReplayClient};
use aws_smithy_types::body::SdkBody;
let http_client = StaticReplayClient::new(vec![ReplayEvent::new(
http::Request::new(SdkBody::from("assume-role request")),
http::Response::builder()
.status(200)
.body(SdkBody::from(
"<AssumeRoleResponse xmlns=\"https://sts.amazonaws.com/doc/2011-06-15/\">\
<AssumeRoleResult><Credentials>\
<AccessKeyId>ASIAASSUMED</AccessKeyId>\
<SecretAccessKey>assumed-secret</SecretAccessKey>\
<SessionToken>assumed-token</SessionToken>\
<Expiration>2030-01-01T00:00:00Z</Expiration>\
</Credentials></AssumeRoleResult>\
</AssumeRoleResponse>",
))
.unwrap(),
)]);
let config = aws_config::SdkConfig::builder()
.behavior_version(aws_config::BehaviorVersion::latest())
.credentials_provider(aws_sdk_s3::config::SharedCredentialsProvider::new(
Credentials::new("base", "base", None, None, "test"),
))
.region(aws_sdk_s3::config::Region::new("us-east-1"))
.time_source(aws_smithy_async::time::SystemTimeSource::new())
.sleep_impl(aws_smithy_async::rt::sleep::TokioSleep::new())
.http_client(http_client)
.build();
let assumer = StsRoleAssumer::from_config(config);
let arn = RoleArn("arn:aws:iam::123456789012:role/dial9-reader".to_string());
let temp = assumer
.assume_role(&arn, Some("us-west-2"))
.await
.expect("assume-role should succeed against the replayed response");
assert_eq!(temp.credentials.access_key_id(), "ASIAASSUMED");
assert_eq!(temp.credentials.secret_access_key(), "assumed-secret");
assert_eq!(temp.credentials.session_token(), Some("assumed-token"));
assert_eq!(temp.region.as_deref(), Some("us-west-2"));
}
}