use std::sync::{Arc, Mutex};
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use serde::Deserialize;
use crate::net::http::{self, Url};
use crate::sha::{hmac_sha256, sha256_hex, to_hex};
#[derive(Debug, Clone)]
pub struct AwsCreds {
pub access_key: String,
pub secret_key: String,
pub session_token: Option<String>,
}
pub fn creds_from_env() -> Result<AwsCreds, String> {
let access_key = std::env::var("AWS_ACCESS_KEY_ID")
.map_err(|_| "aws: AWS_ACCESS_KEY_ID is not set".to_string())?;
let secret_key = std::env::var("AWS_SECRET_ACCESS_KEY")
.map_err(|_| "aws: AWS_SECRET_ACCESS_KEY is not set".to_string())?;
Ok(AwsCreds {
access_key,
secret_key,
session_token: std::env::var("AWS_SESSION_TOKEN").ok(),
})
}
fn timestamps(now: SystemTime) -> (String, String) {
let secs = now.duration_since(UNIX_EPOCH).unwrap_or_default().as_secs() as i64;
let days = secs.div_euclid(86_400);
let sod = secs.rem_euclid(86_400);
let (y, m, d) = crate::obs::log::civil_from_days(days);
let (hh, mm, ss) = (sod / 3600, (sod % 3600) / 60, sod % 60);
(
format!("{y:04}{m:02}{d:02}T{hh:02}{mm:02}{ss:02}Z"),
format!("{y:04}{m:02}{d:02}"),
)
}
#[allow(clippy::too_many_arguments)]
pub fn sigv4_headers(
creds: &AwsCreds,
region: &str,
service: &str,
method: &str,
host: &str,
target: &str,
body: &[u8],
amz_date: &str,
date_stamp: &str,
) -> Vec<(String, String)> {
let (path, query) = match target.split_once('?') {
Some((p, q)) => (p, q),
None => (target, ""),
};
let canonical_uri = if path.is_empty() { "/" } else { path };
let canonical_query = canonical_query(query);
let payload_hash = sha256_hex(body);
let mut canonical_headers = format!("host:{host}\nx-amz-date:{amz_date}\n");
let mut signed_headers = String::from("host;x-amz-date");
if creds.session_token.is_some() {
let tok = creds.session_token.as_deref().unwrap_or("");
canonical_headers.push_str(&format!("x-amz-security-token:{tok}\n"));
signed_headers.push_str(";x-amz-security-token");
}
let canonical_request = format!(
"{method}\n{canonical_uri}\n{canonical_query}\n{canonical_headers}\n{signed_headers}\n{payload_hash}"
);
let scope = format!("{date_stamp}/{region}/{service}/aws4_request");
let string_to_sign = format!(
"AWS4-HMAC-SHA256\n{amz_date}\n{scope}\n{}",
sha256_hex(canonical_request.as_bytes())
);
let k_date = hmac_sha256(
format!("AWS4{}", creds.secret_key).as_bytes(),
date_stamp.as_bytes(),
);
let k_region = hmac_sha256(&k_date, region.as_bytes());
let k_service = hmac_sha256(&k_region, service.as_bytes());
let k_signing = hmac_sha256(&k_service, b"aws4_request");
let signature = to_hex(&hmac_sha256(&k_signing, string_to_sign.as_bytes()));
let authorization = format!(
"AWS4-HMAC-SHA256 Credential={}/{scope}, SignedHeaders={signed_headers}, Signature={signature}",
creds.access_key
);
let mut out = vec![
("X-Amz-Date".to_string(), amz_date.to_string()),
("Authorization".to_string(), authorization),
];
if let Some(tok) = &creds.session_token {
out.push(("X-Amz-Security-Token".to_string(), tok.clone()));
}
out
}
fn canonical_query(query: &str) -> String {
if query.is_empty() {
return String::new();
}
let mut pairs: Vec<(String, String)> = query
.split('&')
.filter(|p| !p.is_empty())
.map(|p| match p.split_once('=') {
Some((k, v)) => (uri_encode(k, false), uri_encode(v, false)),
None => (uri_encode(p, false), String::new()),
})
.collect();
pairs.sort();
pairs
.iter()
.map(|(k, v)| format!("{k}={v}"))
.collect::<Vec<_>>()
.join("&")
}
fn uri_encode(s: &str, keep_slash: bool) -> String {
let mut out = String::with_capacity(s.len());
for &b in s.as_bytes() {
match b {
b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'.' | b'_' | b'~' => {
out.push(b as char)
}
b'/' if keep_slash => out.push('/'),
_ => out.push_str(&format!("%{b:02X}")),
}
}
out
}
enum CredProvider {
Fixed(AwsCreds),
Sso(String),
Imds(Mutex<Option<(AwsCreds, u64)>>),
Irsa(Mutex<Option<(AwsCreds, u64)>>),
}
pub struct SigV4Signer {
provider: CredProvider,
region: String,
service: String,
timeout: Duration,
}
impl SigV4Signer {
pub fn new(creds: AwsCreds, region: String, service: String) -> SigV4Signer {
SigV4Signer {
provider: CredProvider::Fixed(creds),
region,
service,
timeout: Duration::from_secs(10),
}
}
pub fn from_spec(
auth: &crate::config::AuthSpec,
target: &str,
) -> Result<Arc<SigV4Signer>, String> {
let region = auth.region.clone().ok_or("aws: `region` is required")?;
let service = auth
.service
.clone()
.ok_or("aws: `service` is required (e.g. bedrock, execute-api)")?;
let provider = match auth.source.as_deref().unwrap_or("env") {
"env" | "static" => CredProvider::Fixed(creds_from_env()?),
"sso" => CredProvider::Sso(target.to_string()),
"imds" => CredProvider::Imds(Mutex::new(None)),
"irsa" => CredProvider::Irsa(Mutex::new(None)),
other => {
return Err(format!(
"aws: unknown source '{other}' (env|static|sso|imds|irsa)"
));
}
};
Ok(Arc::new(SigV4Signer {
provider,
region,
service,
timeout: Duration::from_secs(10),
}))
}
fn current_creds(&self) -> Option<AwsCreds> {
match &self.provider {
CredProvider::Fixed(c) => Some(c.clone()),
CredProvider::Sso(t) => crate::auth::aws_sso::cached_creds(t),
CredProvider::Imds(cache) => self.cached_or_fetch(cache, || fetch_imds(self.timeout)),
CredProvider::Irsa(cache) => {
self.cached_or_fetch(cache, || fetch_irsa(&self.region, self.timeout))
}
}
}
fn cached_or_fetch(
&self,
cache: &Mutex<Option<(AwsCreds, u64)>>,
fetch: impl FnOnce() -> Result<(AwsCreds, u64), String>,
) -> Option<AwsCreds> {
let now = crate::auth::cache::now_ms();
{
let g = cache.lock().unwrap_or_else(|e| e.into_inner());
if let Some((c, exp)) = g.as_ref()
&& (*exp == 0 || now + 60_000 < *exp)
{
return Some(c.clone());
}
}
match fetch() {
Ok((c, exp)) => {
*cache.lock().unwrap_or_else(|e| e.into_inner()) = Some((c.clone(), exp));
Some(c)
}
Err(_) => None,
}
}
}
fn fetch_imds(timeout: Duration) -> Result<(AwsCreds, u64), String> {
let base = std::env::var("AGENTD_IMDS_ENDPOINT")
.unwrap_or_else(|_| "http://169.254.169.254".to_string());
let token = imds_req(
&format!("{base}/latest/api/token"),
"PUT",
&[("X-aws-ec2-metadata-token-ttl-seconds", "21600")],
timeout,
)?;
let hdr = [("X-aws-ec2-metadata-token", token.trim())];
let role = imds_req(
&format!("{base}/latest/meta-data/iam/security-credentials/"),
"GET",
&hdr,
timeout,
)?;
let role = role.trim();
let body = imds_req(
&format!("{base}/latest/meta-data/iam/security-credentials/{role}"),
"GET",
&hdr,
timeout,
)?;
#[derive(Deserialize)]
#[serde(rename_all = "PascalCase")]
struct ImdsCreds {
access_key_id: String,
secret_access_key: String,
token: String,
}
let c: ImdsCreds =
serde_json::from_str(&body).map_err(|e| format!("imds: bad credentials json: {e}"))?;
Ok((
AwsCreds {
access_key: c.access_key_id,
secret_key: c.secret_access_key,
session_token: Some(c.token),
},
crate::auth::cache::now_ms() + 50 * 60_000,
))
}
fn fetch_irsa(region: &str, timeout: Duration) -> Result<(AwsCreds, u64), String> {
let token_file = std::env::var("AWS_WEB_IDENTITY_TOKEN_FILE")
.map_err(|_| "irsa: AWS_WEB_IDENTITY_TOKEN_FILE is not set".to_string())?;
let role_arn =
std::env::var("AWS_ROLE_ARN").map_err(|_| "irsa: AWS_ROLE_ARN is not set".to_string())?;
let token = std::fs::read_to_string(&token_file)
.map_err(|e| format!("irsa: token file: {e}"))?
.trim()
.to_string();
let sts = std::env::var("AGENTD_STS_ENDPOINT")
.unwrap_or_else(|_| format!("https://sts.{region}.amazonaws.com"));
let form = format!(
"Action=AssumeRoleWithWebIdentity&RoleArn={}&RoleSessionName=agentd&WebIdentityToken={}&Version=2011-06-15",
uri_encode(&role_arn, false),
uri_encode(&token, false)
);
let body = sts_post(&format!("{sts}/"), form.as_bytes(), timeout)?;
let ak = xml_field(&body, "AccessKeyId").ok_or("irsa: STS returned no AccessKeyId")?;
let sk = xml_field(&body, "SecretAccessKey").ok_or("irsa: STS returned no SecretAccessKey")?;
let st = xml_field(&body, "SessionToken");
Ok((
AwsCreds {
access_key: ak,
secret_key: sk,
session_token: st,
},
crate::auth::cache::now_ms() + 50 * 60_000,
))
}
fn imds_req(
url: &str,
method: &str,
headers: &[(&str, &str)],
timeout: Duration,
) -> Result<String, String> {
let u = Url::parse(url).map_err(|e| format!("imds: url {url}: {e}"))?;
let mut s =
http::connect_tcp(&u.host, u.port, timeout).map_err(|e| format!("imds: connect: {e}"))?;
let resp = http::send(&mut s, &u.host_header(), method, &u.path, headers, &[])
.map_err(|e| format!("imds: request failed: {e}"))?;
if !resp.is_success() {
return Err(format!("imds: {url} → HTTP {}", resp.status));
}
Ok(resp.body_str().to_string())
}
fn sts_post(url: &str, form: &[u8], timeout: Duration) -> Result<String, String> {
let u = Url::parse(url).map_err(|e| format!("sts: url {url}: {e}"))?;
let tcp =
http::connect_tcp(&u.host, u.port, timeout).map_err(|e| format!("sts: connect: {e}"))?;
let mut stream: Box<dyn http::Stream> = if u.is_tls() {
#[cfg(feature = "tls")]
{
Box::new(
crate::net::tls::connect(tcp, &u.host, None)
.map_err(|e| format!("sts: tls: {e}"))?,
)
}
#[cfg(not(feature = "tls"))]
{
return Err("sts: https requires --features tls".to_string());
}
} else {
Box::new(tcp)
};
let resp = http::send(
stream.as_mut(),
&u.host_header(),
"POST",
&u.path,
&[("Content-Type", "application/x-www-form-urlencoded")],
form,
)
.map_err(|e| format!("sts: request failed: {e}"))?;
if !resp.is_success() {
return Err(format!("sts: HTTP {}", resp.status));
}
Ok(resp.body_str().to_string())
}
fn xml_field(body: &str, tag: &str) -> Option<String> {
let open = format!("<{tag}>");
let close = format!("</{tag}>");
let start = body.find(&open)? + open.len();
let end = body[start..].find(&close)? + start;
Some(body[start..end].trim().to_string())
}
impl ::mcp::http::RequestSigner for SigV4Signer {
fn sign(
&self,
method: &str,
authority: &str,
path: &str,
body: &[u8],
) -> Vec<(String, String)> {
let Some(creds) = self.current_creds() else {
return Vec::new();
};
let (amz_date, date_stamp) = timestamps(SystemTime::now());
sigv4_headers(
&creds,
&self.region,
&self.service,
method,
authority,
path,
body,
&amz_date,
&date_stamp,
)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn sigv4_matches_the_aws_get_vanilla_vector() {
let creds = AwsCreds {
access_key: "AKIDEXAMPLE".into(),
secret_key: "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY".into(),
session_token: None,
};
let h = sigv4_headers(
&creds,
"us-east-1",
"service",
"GET",
"example.amazonaws.com",
"/",
b"",
"20150830T123600Z",
"20150830",
);
let auth = &h.iter().find(|(k, _)| k == "Authorization").unwrap().1;
assert_eq!(
auth,
"AWS4-HMAC-SHA256 Credential=AKIDEXAMPLE/20150830/us-east-1/service/aws4_request, \
SignedHeaders=host;x-amz-date, \
Signature=5fa00fa31553b73ebf1942676e86291e8372ff2a2260956d9b8aae1d763fbf31"
);
assert_eq!(
h.iter().find(|(k, _)| k == "X-Amz-Date").unwrap().1,
"20150830T123600Z"
);
}
#[test]
fn session_token_is_signed_and_present() {
let creds = AwsCreds {
access_key: "AKID".into(),
secret_key: "secret".into(),
session_token: Some("tok-123".into()),
};
let h = sigv4_headers(
&creds,
"us-east-1",
"bedrock",
"POST",
"h",
"/x",
b"{}",
"20200101T000000Z",
"20200101",
);
let auth = &h.iter().find(|(k, _)| k == "Authorization").unwrap().1;
assert!(
auth.contains("SignedHeaders=host;x-amz-date;x-amz-security-token"),
"the session token is a signed header: {auth}"
);
assert_eq!(
h.iter()
.find(|(k, _)| k == "X-Amz-Security-Token")
.unwrap()
.1,
"tok-123"
);
}
#[test]
fn canonical_query_sorts_and_encodes() {
assert_eq!(canonical_query(""), "");
assert_eq!(canonical_query("b=2&a=1"), "a=1&b=2");
assert_eq!(canonical_query("x=a b"), "x=a%20b");
}
#[test]
fn xml_field_extracts_sts_credentials() {
let body = "<AssumeRoleWithWebIdentityResult><Credentials>\
<AccessKeyId>ASIAIRSA</AccessKeyId>\
<SecretAccessKey>irsa-secret</SecretAccessKey>\
<SessionToken>irsa-sess</SessionToken>\
<Expiration>2026-01-01T00:00:00Z</Expiration>\
</Credentials></AssumeRoleWithWebIdentityResult>";
assert_eq!(xml_field(body, "AccessKeyId").as_deref(), Some("ASIAIRSA"));
assert_eq!(
xml_field(body, "SecretAccessKey").as_deref(),
Some("irsa-secret")
);
assert_eq!(
xml_field(body, "SessionToken").as_deref(),
Some("irsa-sess")
);
assert_eq!(xml_field(body, "Nope"), None);
}
}