Skip to main content

olai_http/aws/
credential.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18use std::collections::BTreeMap;
19use std::sync::Arc;
20use std::time::Instant;
21
22use async_trait::async_trait;
23use bytes::Buf;
24use chrono::{DateTime, Utc};
25use percent_encoding::utf8_percent_encode;
26use reqwest::header::{AUTHORIZATION, HeaderMap, HeaderValue};
27use reqwest::{Client, Method, Request, RequestBuilder, StatusCode};
28use serde::Deserialize;
29use tracing::warn;
30use url::Url;
31
32use crate::retry::RetryExt;
33use crate::service::HttpService;
34use crate::token::{TemporaryToken, TokenCache};
35use crate::util::{hex_digest, hex_encode, hmac_sha256};
36use crate::{CredentialProvider, Result, RetryConfig, TokenProvider};
37
38use crate::util::STRICT_ENCODE_SET;
39
40/// This is used to maintain the URI path encoding
41const STRICT_PATH_ENCODE_SET: percent_encoding::AsciiSet = STRICT_ENCODE_SET.remove(b'/');
42
43type StdError = Box<dyn std::error::Error + Send + Sync>;
44
45/// SHA256 hash of empty string
46static EMPTY_SHA256_HASH: &str = "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855";
47
48/// A set of AWS security credentials
49#[derive(Eq, PartialEq)]
50pub struct AwsCredential {
51    /// AWS_ACCESS_KEY_ID
52    pub key_id: String,
53    /// AWS_SECRET_ACCESS_KEY
54    pub secret_key: String,
55    /// AWS_SESSION_TOKEN
56    pub token: Option<String>,
57}
58
59// Manual `Debug` impl: an `AwsCredential` holds long-lived secrets, so we never
60// expose the field values (including the access key id, which is itself a
61// sensitive identifier). See the credential-redaction convention in the README.
62impl std::fmt::Debug for AwsCredential {
63    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
64        f.debug_struct("AwsCredential")
65            .field("key_id", &"<redacted>")
66            .field("secret_key", &"<redacted>")
67            .field("token", &self.token.as_ref().map(|_| "<redacted>"))
68            .finish()
69    }
70}
71
72impl AwsCredential {
73    /// Signs a string
74    ///
75    /// <https://docs.aws.amazon.com/general/latest/gr/sigv4-calculate-signature.html>
76    fn sign(&self, to_sign: &str, date: DateTime<Utc>, region: &str, service: &str) -> String {
77        let date_string = date.format("%Y%m%d").to_string();
78        let date_hmac = hmac_sha256(format!("AWS4{}", self.secret_key), date_string);
79        let region_hmac = hmac_sha256(date_hmac, region);
80        let service_hmac = hmac_sha256(region_hmac, service);
81        let signing_hmac = hmac_sha256(service_hmac, b"aws4_request");
82        hex_encode(hmac_sha256(signing_hmac, to_sign).as_ref())
83    }
84}
85
86/// Authorize a [`Request`] with an [`AwsCredential`] using [AWS SigV4]
87///
88/// [AWS SigV4]: https://docs.aws.amazon.com/general/latest/gr/sigv4-calculate-signature.html
89#[derive(Debug)]
90pub struct AwsAuthorizer<'a> {
91    date: Option<DateTime<Utc>>,
92    credential: &'a AwsCredential,
93    service: &'a str,
94    region: &'a str,
95}
96
97static DATE_HEADER: hyper::header::HeaderName =
98    hyper::header::HeaderName::from_static("x-amz-date");
99static HASH_HEADER: hyper::header::HeaderName =
100    hyper::header::HeaderName::from_static("x-amz-content-sha256");
101static TOKEN_HEADER: hyper::header::HeaderName =
102    hyper::header::HeaderName::from_static("x-amz-security-token");
103const ALGORITHM: &str = "AWS4-HMAC-SHA256";
104
105impl<'a> AwsAuthorizer<'a> {
106    /// Create a new [`AwsAuthorizer`]
107    pub fn new(credential: &'a AwsCredential, service: &'a str, region: &'a str) -> Self {
108        Self {
109            credential,
110            service,
111            region,
112            date: None,
113        }
114    }
115
116    /// Authorize `request` with an optional pre-calculated SHA256 digest by attaching
117    /// the relevant [AWS SigV4] headers
118    ///
119    /// [AWS SigV4]: https://docs.aws.amazon.com/IAM/latest/UserGuide/create-signed-request.html
120    pub fn authorize(
121        &self,
122        request: &mut Request,
123        pre_calculated_digest: Option<&[u8]>,
124    ) -> crate::Result<()> {
125        if let Some(ref token) = self.credential.token {
126            let token_val = HeaderValue::from_str(token)?;
127            request.headers_mut().insert(&TOKEN_HEADER, token_val);
128        }
129
130        let host = &request.url()[url::Position::BeforeHost..url::Position::AfterPort];
131        let host_val = HeaderValue::from_str(host)?;
132        request.headers_mut().insert("host", host_val);
133
134        let date = self.date.unwrap_or_else(Utc::now);
135        let date_str = date.format("%Y%m%dT%H%M%SZ").to_string();
136        let date_val = HeaderValue::from_str(&date_str)?;
137        request.headers_mut().insert(&DATE_HEADER, date_val);
138
139        let digest = match pre_calculated_digest {
140            Some(digest) => hex_encode(digest),
141            None => match request.body() {
142                None => EMPTY_SHA256_HASH.to_string(),
143                Some(body) => match body.as_bytes() {
144                    Some(bytes) => hex_digest(bytes),
145                    None => EMPTY_SHA256_HASH.to_string(),
146                },
147            },
148        };
149
150        let header_digest = HeaderValue::from_str(&digest)?;
151        request.headers_mut().insert(&HASH_HEADER, header_digest);
152
153        let (signed_headers, canonical_headers) = canonicalize_headers(request.headers());
154
155        let scope = self.scope(date);
156
157        let string_to_sign = self.string_to_sign(
158            date,
159            &scope,
160            request.method(),
161            request.url(),
162            &canonical_headers,
163            &signed_headers,
164            &digest,
165        );
166
167        let signature = self
168            .credential
169            .sign(&string_to_sign, date, self.region, self.service);
170
171        let authorisation = format!(
172            "{} Credential={}/{}, SignedHeaders={}, Signature={}",
173            ALGORITHM, self.credential.key_id, scope, signed_headers, signature
174        );
175
176        let authorization_val = HeaderValue::from_str(&authorisation)?;
177        request
178            .headers_mut()
179            .insert(&AUTHORIZATION, authorization_val);
180        Ok(())
181    }
182
183    #[allow(clippy::too_many_arguments)]
184    fn string_to_sign(
185        &self,
186        date: DateTime<Utc>,
187        scope: &str,
188        request_method: &Method,
189        url: &Url,
190        canonical_headers: &str,
191        signed_headers: &str,
192        digest: &str,
193    ) -> String {
194        // Each path segment must be URI-encoded twice (except for Amazon S3 which only gets
195        // URI-encoded once).
196        // see https://docs.aws.amazon.com/general/latest/gr/sigv4-create-canonical-request.html
197        let canonical_uri = match self.service {
198            "s3" => url.path().to_string(),
199            _ => utf8_percent_encode(url.path(), &STRICT_PATH_ENCODE_SET).to_string(),
200        };
201
202        let canonical_query = canonicalize_query(url);
203
204        let canonical_request = format!(
205            "{}\n{}\n{}\n{}\n{}\n{}",
206            request_method.as_str(),
207            canonical_uri,
208            canonical_query,
209            canonical_headers,
210            signed_headers,
211            digest
212        );
213
214        let hashed_canonical_request = hex_digest(canonical_request.as_bytes());
215
216        format!(
217            "{}\n{}\n{}\n{}",
218            ALGORITHM,
219            date.format("%Y%m%dT%H%M%SZ"),
220            scope,
221            hashed_canonical_request
222        )
223    }
224
225    fn scope(&self, date: DateTime<Utc>) -> String {
226        format!(
227            "{}/{}/{}/aws4_request",
228            date.format("%Y%m%d"),
229            self.region,
230            self.service
231        )
232    }
233}
234
235pub(crate) trait CredentialExt {
236    /// Sign a request <https://docs.aws.amazon.com/general/latest/gr/sigv4_signing.html>
237    fn with_aws_sigv4(
238        self,
239        authorizer: Option<AwsAuthorizer<'_>>,
240        payload_sha256: Option<&[u8]>,
241    ) -> Self;
242}
243
244impl CredentialExt for RequestBuilder {
245    fn with_aws_sigv4(
246        self,
247        authorizer: Option<AwsAuthorizer<'_>>,
248        payload_sha256: Option<&[u8]>,
249    ) -> Self {
250        match authorizer {
251            Some(authorizer) => {
252                let (client, request) = self.build_split();
253                let mut request = request.expect("request valid");
254                authorizer
255                    .authorize(&mut request, payload_sha256)
256                    .expect("credential values must be valid header characters");
257
258                Self::from_parts(client, request)
259            }
260            None => self,
261        }
262    }
263}
264
265/// Canonicalizes query parameters into the AWS canonical form
266///
267/// <https://docs.aws.amazon.com/general/latest/gr/sigv4-create-canonical-request.html>
268fn canonicalize_query(url: &Url) -> String {
269    use std::fmt::Write;
270
271    let capacity = match url.query() {
272        Some(q) if !q.is_empty() => q.len(),
273        _ => return String::new(),
274    };
275    let mut encoded = String::with_capacity(capacity + 1);
276
277    let mut headers = url.query_pairs().collect::<Vec<_>>();
278    headers.sort_unstable_by(|(a, _), (b, _)| a.cmp(b));
279
280    let mut first = true;
281    for (k, v) in headers {
282        if !first {
283            encoded.push('&');
284        }
285        first = false;
286        let _ = write!(
287            encoded,
288            "{}={}",
289            utf8_percent_encode(k.as_ref(), &STRICT_ENCODE_SET),
290            utf8_percent_encode(v.as_ref(), &STRICT_ENCODE_SET)
291        );
292    }
293    encoded
294}
295
296/// Canonicalizes headers into the AWS Canonical Form.
297///
298/// <https://docs.aws.amazon.com/general/latest/gr/sigv4-create-canonical-request.html>
299fn canonicalize_headers(header_map: &HeaderMap) -> (String, String) {
300    // Use owned Strings for values since header bytes may not be valid UTF-8;
301    // we convert using lossy UTF-8 to avoid panicking on unusual header values.
302    let mut headers = BTreeMap::<&str, Vec<String>>::new();
303    let mut value_count = 0;
304    let mut value_bytes = 0;
305    let mut key_bytes = 0;
306
307    for (key, value) in header_map {
308        let key = key.as_str();
309        if ["authorization", "content-length", "user-agent"].contains(&key) {
310            continue;
311        }
312
313        let value = String::from_utf8_lossy(value.as_bytes()).into_owned();
314        key_bytes += key.len();
315        value_bytes += value.len();
316        value_count += 1;
317        headers.entry(key).or_default().push(value);
318    }
319
320    let mut signed_headers = String::with_capacity(key_bytes + headers.len());
321    let mut canonical_headers =
322        String::with_capacity(key_bytes + value_bytes + headers.len() + value_count);
323
324    for (header_idx, (name, values)) in headers.into_iter().enumerate() {
325        if header_idx != 0 {
326            signed_headers.push(';');
327        }
328
329        signed_headers.push_str(name);
330        canonical_headers.push_str(name);
331        canonical_headers.push(':');
332        for (value_idx, value) in values.into_iter().enumerate() {
333            if value_idx != 0 {
334                canonical_headers.push(',');
335            }
336            canonical_headers.push_str(value.trim());
337        }
338        canonical_headers.push('\n');
339    }
340
341    (signed_headers, canonical_headers)
342}
343
344/// Credentials sourced from the instance metadata service
345///
346/// Fetches short-lived `AwsCredential` values from the EC2 Instance Metadata
347/// Service (IMDS).  Supports both IMDSv2 (token-protected) and IMDSv1 as a
348/// fallback when `imdsv1_fallback` is set.
349///
350/// # References
351/// - <https://docs.aws.amazon.com/AWSEC2/latest/UserGuide/ec2-instance-metadata.html>
352/// - <https://docs.aws.amazon.com/AWSEC2/latest/UserGuide/configuring-instance-metadata-service.html>
353#[derive(Debug)]
354pub(crate) struct InstanceCredentialProvider {
355    pub imdsv1_fallback: bool,
356    pub metadata_endpoint: String,
357}
358
359#[async_trait]
360impl TokenProvider for InstanceCredentialProvider {
361    type Credential = AwsCredential;
362
363    async fn fetch_token(
364        &self,
365        client: &Client,
366        service: &Arc<dyn HttpService>,
367        retry: &RetryConfig,
368    ) -> Result<TemporaryToken<Arc<AwsCredential>>> {
369        instance_creds(
370            client,
371            service,
372            retry,
373            &self.metadata_endpoint,
374            self.imdsv1_fallback,
375        )
376        .await
377        .map_err(|source| crate::Error::Generic { source })
378    }
379}
380
381/// Credentials sourced using `AssumeRoleWithWebIdentity`.
382///
383/// Exchanges an OIDC token (e.g. a Kubernetes projected service-account token)
384/// for temporary AWS credentials by calling the STS
385/// `AssumeRoleWithWebIdentity` API.  Commonly used in EKS pod identity
386/// scenarios where the token file path is supplied via the
387/// `AWS_WEB_IDENTITY_TOKEN_FILE` environment variable.
388///
389/// # References
390/// - <https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRoleWithWebIdentity.html>
391/// - <https://docs.aws.amazon.com/eks/latest/userguide/iam-roles-for-service-accounts-technical-overview.html>
392#[derive(Debug)]
393pub(crate) struct WebIdentityProvider {
394    pub token_path: String,
395    pub role_arn: String,
396    pub session_name: String,
397    pub endpoint: String,
398}
399
400#[async_trait]
401impl TokenProvider for WebIdentityProvider {
402    type Credential = AwsCredential;
403
404    async fn fetch_token(
405        &self,
406        client: &Client,
407        service: &Arc<dyn HttpService>,
408        retry: &RetryConfig,
409    ) -> Result<TemporaryToken<Arc<AwsCredential>>> {
410        web_identity(
411            client,
412            service,
413            retry,
414            &self.token_path,
415            &self.role_arn,
416            &self.session_name,
417            &self.endpoint,
418        )
419        .await
420        .map_err(|source| crate::Error::Generic { source })
421    }
422}
423
424#[derive(Debug, Deserialize)]
425#[serde(rename_all = "PascalCase")]
426struct InstanceCredentials {
427    access_key_id: String,
428    secret_access_key: String,
429    token: String,
430    expiration: DateTime<Utc>,
431}
432
433impl From<InstanceCredentials> for AwsCredential {
434    fn from(s: InstanceCredentials) -> Self {
435        Self {
436            key_id: s.access_key_id,
437            secret_key: s.secret_access_key,
438            token: Some(s.token),
439        }
440    }
441}
442
443/// <https://docs.aws.amazon.com/AWSEC2/latest/UserGuide/iam-roles-for-amazon-ec2.html#instance-metadata-security-credentials>
444async fn instance_creds(
445    client: &Client,
446    service: &Arc<dyn HttpService>,
447    retry_config: &RetryConfig,
448    endpoint: &str,
449    imdsv1_fallback: bool,
450) -> Result<TemporaryToken<Arc<AwsCredential>>, StdError> {
451    const CREDENTIALS_PATH: &str = "latest/meta-data/iam/security-credentials";
452    const AWS_EC2_METADATA_TOKEN_HEADER: &str = "X-aws-ec2-metadata-token";
453
454    let token_url = format!("{endpoint}/latest/api/token");
455
456    let token_result = client
457        .request(Method::PUT, token_url)
458        .header("X-aws-ec2-metadata-token-ttl-seconds", "600") // 10 minute TTL
459        .retryable(retry_config, service.clone())
460        .idempotent(true)
461        .send()
462        .await;
463
464    let token = match token_result {
465        Ok(t) => Some(t.text().await?),
466        Err(e) if imdsv1_fallback && matches!(e.status(), Some(StatusCode::FORBIDDEN)) => {
467            warn!("received 403 from metadata endpoint, falling back to IMDSv1");
468            None
469        }
470        Err(e) => return Err(e.into()),
471    };
472
473    let role_url = format!("{endpoint}/{CREDENTIALS_PATH}/");
474    let mut role_request = client.request(Method::GET, role_url);
475
476    if let Some(token) = &token {
477        role_request = role_request.header(AWS_EC2_METADATA_TOKEN_HEADER, token);
478    }
479
480    let role = role_request
481        .send_retry(retry_config, service.clone())
482        .await?
483        .text()
484        .await?;
485
486    let creds_url = format!("{endpoint}/{CREDENTIALS_PATH}/{role}");
487    let mut creds_request = client.request(Method::GET, creds_url);
488    if let Some(token) = &token {
489        creds_request = creds_request.header(AWS_EC2_METADATA_TOKEN_HEADER, token);
490    }
491
492    let creds: InstanceCredentials = creds_request
493        .send_retry(retry_config, service.clone())
494        .await?
495        .json()
496        .await?;
497
498    let now = Utc::now();
499    let ttl = (creds.expiration - now).to_std().unwrap_or_default();
500    Ok(TemporaryToken {
501        token: Arc::new(creds.into()),
502        expiry: Some(Instant::now() + ttl),
503    })
504}
505
506/// Exchanges long-term or instance credentials for temporary credentials by
507/// calling the STS [`AssumeRole`] API.  The caller's base credentials are used
508/// to SigV4-sign the STS request.
509///
510/// Wire this provider via [`AmazonBuilder::with_role_arn`] or set the
511/// `AWS_ROLE_ARN` / `AWS_ROLE_SESSION_NAME` environment variables alongside a
512/// static or instance credential source.
513///
514/// # References
515/// - <https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRole.html>
516/// - <https://docs.aws.amazon.com/IAM/latest/UserGuide/id_roles_use_switch-role-api.html>
517#[derive(Debug)]
518pub(crate) struct AssumeRoleProvider {
519    pub role_arn: String,
520    pub session_name: String,
521    pub endpoint: String,
522    pub base_credentials: Arc<dyn CredentialProvider<Credential = AwsCredential>>,
523    pub region: String,
524    /// Optional inline session policy (URL-encoded JSON).
525    /// Intersected with the role's own policy to further restrict access.
526    pub policy: Option<String>,
527}
528
529#[async_trait]
530impl TokenProvider for AssumeRoleProvider {
531    type Credential = AwsCredential;
532
533    async fn fetch_token(
534        &self,
535        client: &Client,
536        service: &Arc<dyn HttpService>,
537        retry: &RetryConfig,
538    ) -> Result<TemporaryToken<Arc<AwsCredential>>> {
539        let base_cred = self
540            .base_credentials
541            .get_credential()
542            .await
543            .map_err(|source| crate::Error::Generic {
544                source: Box::new(source),
545            })?;
546
547        assume_role(
548            client,
549            service,
550            retry,
551            &base_cred,
552            &self.region,
553            &self.role_arn,
554            &self.session_name,
555            &self.endpoint,
556            self.policy.as_deref(),
557        )
558        .await
559        .map_err(|source| crate::Error::Generic { source })
560    }
561}
562
563/// Calls `STS:AssumeRole` and returns temporary credentials.
564///
565/// The request is SigV4-signed using the provided `base_cred`.
566///
567/// # References
568/// - <https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRole.html>
569#[allow(clippy::too_many_arguments)]
570async fn assume_role(
571    client: &Client,
572    service: &Arc<dyn HttpService>,
573    retry_config: &RetryConfig,
574    base_cred: &AwsCredential,
575    region: &str,
576    role_arn: &str,
577    session_name: &str,
578    endpoint: &str,
579    policy: Option<&str>,
580) -> Result<TemporaryToken<Arc<AwsCredential>>, StdError> {
581    let mut query: Vec<(&str, &str)> = vec![
582        ("Action", "AssumeRole"),
583        ("DurationSeconds", "3600"),
584        ("RoleArn", role_arn),
585        ("RoleSessionName", session_name),
586        ("Version", "2011-06-15"),
587    ];
588    if let Some(p) = policy {
589        query.push(("Policy", p));
590    }
591    let bytes = client
592        .request(Method::POST, endpoint)
593        .query(&query)
594        .with_aws_sigv4(Some(AwsAuthorizer::new(base_cred, "sts", region)), None)
595        .retryable(retry_config, service.clone())
596        .idempotent(true)
597        .send()
598        .await?
599        .bytes()
600        .await?;
601
602    let resp: AssumeRoleXmlResponse = quick_xml::de::from_reader(bytes.reader())
603        .map_err(|e| format!("Invalid AssumeRole response: {e}"))?;
604
605    let creds = resp.assume_role_result.credentials;
606    let now = Utc::now();
607    let ttl = (creds.expiration - now).to_std().unwrap_or_default();
608
609    Ok(TemporaryToken {
610        token: Arc::new(creds.into()),
611        expiry: Some(Instant::now() + ttl),
612    })
613}
614
615/// XML response for `AssumeRole`.
616#[derive(Debug, Deserialize)]
617#[serde(rename_all = "PascalCase")]
618struct AssumeRoleXmlResponse {
619    assume_role_result: AssumeRoleResult,
620}
621
622#[derive(Debug, Deserialize)]
623#[serde(rename_all = "PascalCase")]
624struct AssumeRoleResponse {
625    assume_role_with_web_identity_result: AssumeRoleResult,
626}
627
628#[derive(Debug, Deserialize)]
629#[serde(rename_all = "PascalCase")]
630struct AssumeRoleResult {
631    credentials: SessionCredentials,
632}
633
634#[derive(Debug, Deserialize)]
635#[serde(rename_all = "PascalCase")]
636struct SessionCredentials {
637    session_token: String,
638    secret_access_key: String,
639    access_key_id: String,
640    expiration: DateTime<Utc>,
641}
642
643impl From<SessionCredentials> for AwsCredential {
644    fn from(s: SessionCredentials) -> Self {
645        Self {
646            key_id: s.access_key_id,
647            secret_key: s.secret_access_key,
648            token: Some(s.session_token),
649        }
650    }
651}
652
653/// <https://docs.aws.amazon.com/eks/latest/userguide/iam-roles-for-service-accounts-technical-overview.html>
654async fn web_identity(
655    client: &Client,
656    service: &Arc<dyn HttpService>,
657    retry_config: &RetryConfig,
658    token_path: &str,
659    role_arn: &str,
660    session_name: &str,
661    endpoint: &str,
662) -> Result<TemporaryToken<Arc<AwsCredential>>, StdError> {
663    let token = std::fs::read_to_string(token_path)
664        .map_err(|e| format!("Failed to read token file '{token_path}': {e}"))?;
665
666    let bytes = client
667        .request(Method::POST, endpoint)
668        .query(&[
669            ("Action", "AssumeRoleWithWebIdentity"),
670            ("DurationSeconds", "3600"),
671            ("RoleArn", role_arn),
672            ("RoleSessionName", session_name),
673            ("Version", "2011-06-15"),
674            ("WebIdentityToken", &token),
675        ])
676        .retryable(retry_config, service.clone())
677        .idempotent(true)
678        .sensitive(true)
679        .send()
680        .await?
681        .bytes()
682        .await?;
683
684    let resp: AssumeRoleResponse = quick_xml::de::from_reader(bytes.reader())
685        .map_err(|e| format!("Invalid AssumeRoleWithWebIdentity response: {e}"))?;
686
687    let creds = resp.assume_role_with_web_identity_result.credentials;
688    let now = Utc::now();
689    let ttl = (creds.expiration - now).to_std().unwrap_or_default();
690
691    Ok(TemporaryToken {
692        token: Arc::new(creds.into()),
693        expiry: Some(Instant::now() + ttl),
694    })
695}
696
697/// Credentials sourced from a task IAM role.
698///
699/// Fetches short-lived `AwsCredential` values from the ECS task-metadata
700/// credential endpoint.  The URL is resolved from environment variables in
701/// this priority order:
702///
703/// 1. `AWS_CONTAINER_CREDENTIALS_FULL_URI` — absolute URL (used by EKS Pod
704///    Identity and Lambda)
705/// 2. `AWS_CONTAINER_CREDENTIALS_RELATIVE_URI` — path appended to
706///    `http://169.254.170.2` (classic ECS task role)
707///
708/// When `AWS_CONTAINER_AUTHORIZATION_TOKEN_FILE` is set the file contents are
709/// sent as the `Authorization` request header, enabling EKS Pod Identity.
710/// Results are cached in the embedded [`TokenCache`] until they approach expiry.
711///
712/// # References
713/// - <https://docs.aws.amazon.com/AmazonECS/latest/developerguide/task-iam-roles.html>
714/// - <https://docs.aws.amazon.com/eks/latest/userguide/pod-id-how-it-works.html>
715#[derive(Debug)]
716pub(crate) struct TaskCredentialProvider {
717    pub url: String,
718    /// Optional authorization token sent as the `Authorization` header.
719    /// Used by EKS Pod Identity (`AWS_CONTAINER_AUTHORIZATION_TOKEN_FILE`).
720    pub auth_token_file: Option<String>,
721    pub retry: RetryConfig,
722    pub client: Client,
723    pub service: Arc<dyn HttpService>,
724    pub cache: TokenCache<Arc<AwsCredential>>,
725}
726
727#[async_trait]
728impl CredentialProvider for TaskCredentialProvider {
729    type Credential = AwsCredential;
730
731    async fn get_credential(&self) -> Result<Arc<AwsCredential>> {
732        self.cache
733            .get_or_insert_with(|| {
734                task_credential(
735                    &self.client,
736                    &self.service,
737                    &self.retry,
738                    &self.url,
739                    self.auth_token_file.as_deref(),
740                )
741            })
742            .await
743            .map_err(|source| crate::Error::Generic { source })
744    }
745}
746
747/// <https://docs.aws.amazon.com/AmazonECS/latest/developerguide/task-iam-roles.html>
748async fn task_credential(
749    client: &Client,
750    service: &Arc<dyn HttpService>,
751    retry: &RetryConfig,
752    url: &str,
753    auth_token_file: Option<&str>,
754) -> Result<TemporaryToken<Arc<AwsCredential>>, StdError> {
755    let mut req = client.get(url);
756
757    // EKS Pod Identity requires an Authorization header read from a file
758    if let Some(token_file) = auth_token_file {
759        let token = std::fs::read_to_string(token_file)
760            .map_err(|e| format!("Failed to read auth token file '{token_file}': {e}"))?;
761        req = req.header(reqwest::header::AUTHORIZATION, token.trim());
762    }
763
764    let creds: InstanceCredentials = req.send_retry(retry, service.clone()).await?.json().await?;
765
766    let now = Utc::now();
767    let ttl = (creds.expiration - now).to_std().unwrap_or_default();
768    Ok(TemporaryToken {
769        token: Arc::new(creds.into()),
770        expiry: Some(Instant::now() + ttl),
771    })
772}
773
774#[cfg(test)]
775mod tests {
776    use super::*;
777    use crate::service::ReqwestService;
778    use reqwest::{Client, Method};
779    use std::env;
780
781    #[test]
782    fn test_debug_does_not_leak_secrets() {
783        // Deliberately non-secret-looking sentinel values so the secret
784        // scanner doesn't flag this test; we only care that the Debug output
785        // does not echo them back.
786        let credential = AwsCredential {
787            key_id: "fake-key-id-do-not-print".to_string(),
788            secret_key: "fake-secret-do-not-print".to_string(),
789            token: Some("fake-session-token-do-not-print".to_string()),
790        };
791        let rendered = format!("{credential:?}");
792        assert!(
793            !rendered.contains("fake-key-id-do-not-print"),
794            "key_id leaked: {rendered}"
795        );
796        assert!(
797            !rendered.contains("fake-secret-do-not-print"),
798            "secret_key leaked: {rendered}"
799        );
800        assert!(
801            !rendered.contains("fake-session-token-do-not-print"),
802            "session token leaked: {rendered}"
803        );
804        assert!(rendered.contains("<redacted>"));
805    }
806
807    #[test]
808    fn test_sign_with_signed_payload() {
809        let client = Client::new();
810
811        let credential = AwsCredential {
812            // AWS-documented SigV4 example vectors — the signature asserted
813            // below is computed from exactly these values, so they cannot be
814            // swapped for non-secret-looking placeholders.
815            key_id: "AKIAIOSFODNN7EXAMPLE".to_string(), // gitleaks:allow
816            secret_key: "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY".to_string(), // gitleaks:allow
817            token: None,
818        };
819
820        let date = DateTime::parse_from_rfc3339("2022-08-06T18:01:34Z")
821            .unwrap()
822            .with_timezone(&Utc);
823
824        let mut request = client
825            .request(Method::GET, "https://ec2.amazon.com/")
826            .build()
827            .unwrap();
828
829        let signer = AwsAuthorizer {
830            date: Some(date),
831            credential: &credential,
832            service: "ec2",
833            region: "us-east-1",
834        };
835
836        signer.authorize(&mut request, None).unwrap();
837        assert_eq!(
838            request.headers().get(&AUTHORIZATION).unwrap(),
839            "AWS4-HMAC-SHA256 Credential=AKIAIOSFODNN7EXAMPLE/20220806/us-east-1/ec2/aws4_request, SignedHeaders=host;x-amz-content-sha256;x-amz-date, Signature=a3c787a7ed37f7fdfbfd2d7056a3d7c9d85e6d52a2bfbec73793c0be6e7862d4" // gitleaks:allow
840        )
841        // gitleaks:allow
842    }
843
844    #[test]
845    fn test_sign_port() {
846        let client = Client::new();
847
848        let credential = AwsCredential {
849            key_id: "H20ABqCkLZID4rLe".to_string(),
850            secret_key: "jMqRDgxSsBqqznfmddGdu1TmmZOJQxdM".to_string(),
851            token: None,
852        };
853
854        let date = DateTime::parse_from_rfc3339("2022-08-09T13:05:25Z")
855            .unwrap()
856            .with_timezone(&Utc);
857
858        let mut request = client
859            .request(Method::GET, "http://localhost:9000/tsm-schemas")
860            .query(&[
861                ("delimiter", "/"),
862                ("encoding-type", "url"),
863                ("list-type", "2"),
864                ("prefix", ""),
865            ])
866            .build()
867            .unwrap();
868
869        let authorizer = AwsAuthorizer {
870            date: Some(date),
871            credential: &credential,
872            service: "s3",
873            region: "us-east-1",
874        };
875
876        authorizer.authorize(&mut request, None).unwrap();
877        assert_eq!(
878            request.headers().get(&AUTHORIZATION).unwrap(),
879            "AWS4-HMAC-SHA256 Credential=H20ABqCkLZID4rLe/20220809/us-east-1/s3/aws4_request, SignedHeaders=host;x-amz-content-sha256;x-amz-date, Signature=9ebf2f92872066c99ac94e573b4e1b80f4dbb8a32b1e8e23178318746e7d1b4d"
880        )
881    }
882
883    #[tokio::test]
884    async fn test_instance_metadata() {
885        if env::var("TEST_INTEGRATION").is_err() {
886            eprintln!("skipping AWS integration test");
887            return;
888        }
889
890        let endpoint = env::var("EC2_METADATA_ENDPOINT").unwrap();
891        let client = Client::new();
892        let service: Arc<dyn HttpService> = Arc::new(ReqwestService::new(client.clone()));
893        let retry_config = RetryConfig::default();
894
895        let resp = client
896            .request(Method::GET, format!("{endpoint}/latest/meta-data/ami-id"))
897            .send()
898            .await
899            .unwrap();
900
901        assert_eq!(
902            resp.status(),
903            StatusCode::UNAUTHORIZED,
904            "Ensure metadata endpoint is set to only allow IMDSv2"
905        );
906
907        let creds = instance_creds(&client, &service, &retry_config, &endpoint, false)
908            .await
909            .unwrap();
910
911        let id = &creds.token.key_id;
912        let secret = &creds.token.secret_key;
913        let token = creds.token.token.as_ref().unwrap();
914
915        assert!(!id.is_empty());
916        assert!(!secret.is_empty());
917        assert!(!token.is_empty())
918    }
919    #[tokio::test]
920    async fn test_mock() {
921        let mut server = mockito::Server::new_async().await;
922
923        const IMDSV2_HEADER: &str = "X-aws-ec2-metadata-token";
924
925        let secret_access_key = "SECRET";
926        let access_key_id = "KEYID";
927        let token = "TOKEN";
928
929        let endpoint = server.url();
930        let client = Client::new();
931        let service: Arc<dyn HttpService> = Arc::new(ReqwestService::new(client.clone()));
932        let retry_config = RetryConfig::default();
933
934        // Test IMDSv2
935        let _mock1 = server
936            .mock("PUT", "/latest/api/token")
937            .with_status(200)
938            .with_body("cupcakes")
939            .create_async()
940            .await;
941
942        let _mock2 = server
943            .mock("GET", "/latest/meta-data/iam/security-credentials/")
944            .match_header(IMDSV2_HEADER, "cupcakes")
945            .with_status(200)
946            .with_body("myrole")
947            .create_async()
948            .await;
949
950        let _mock3 = server
951            .mock("GET", "/latest/meta-data/iam/security-credentials/myrole")
952            .match_header(IMDSV2_HEADER, "cupcakes")
953            .with_status(200)
954            .with_body(r#"{"AccessKeyId":"KEYID","Code":"Success","Expiration":"2022-08-30T10:51:04Z","LastUpdated":"2022-08-30T10:21:04Z","SecretAccessKey":"SECRET","Token":"TOKEN","Type":"AWS-HMAC"}"#)
955            .create_async()
956            .await;
957
958        let creds = instance_creds(&client, &service, &retry_config, &endpoint, true)
959            .await
960            .unwrap();
961
962        assert_eq!(creds.token.token.as_deref().unwrap(), token);
963        assert_eq!(&creds.token.key_id, access_key_id);
964        assert_eq!(&creds.token.secret_key, secret_access_key);
965
966        // Test IMDSv1 fallback
967        let _mock4 = server
968            .mock("PUT", "/latest/api/token")
969            .with_status(403)
970            .with_body("")
971            .create_async()
972            .await;
973
974        let _mock5 = server
975            .mock("GET", "/latest/meta-data/iam/security-credentials/")
976            .with_status(200)
977            .with_body("myrole")
978            .create_async()
979            .await;
980
981        let _mock6 = server
982            .mock("GET", "/latest/meta-data/iam/security-credentials/myrole")
983            .with_status(200)
984            .with_body(r#"{"AccessKeyId":"KEYID","Code":"Success","Expiration":"2022-08-30T10:51:04Z","LastUpdated":"2022-08-30T10:21:04Z","SecretAccessKey":"SECRET","Token":"TOKEN","Type":"AWS-HMAC"}"#)
985            .create_async()
986            .await;
987
988        let creds = instance_creds(&client, &service, &retry_config, &endpoint, true)
989            .await
990            .unwrap();
991
992        assert_eq!(creds.token.token.as_deref().unwrap(), token);
993        assert_eq!(&creds.token.key_id, access_key_id);
994        assert_eq!(&creds.token.secret_key, secret_access_key);
995
996        // Test IMDSv1 fallback disabled
997        let _mock7 = server
998            .mock("PUT", "/latest/api/token")
999            .with_status(403)
1000            .with_body("")
1001            .create_async()
1002            .await;
1003
1004        // Should fail
1005        instance_creds(&client, &service, &retry_config, &endpoint, false)
1006            .await
1007            .unwrap_err();
1008    }
1009
1010    const STS_WEB_IDENTITY_XML: &str = r#"<AssumeRoleWithWebIdentityResponse xmlns="https://sts.amazonaws.com/doc/2011-06-15/"><AssumeRoleWithWebIdentityResult><Credentials><AccessKeyId>AKIAIOSFODNN7EXAMPLE</AccessKeyId><SecretAccessKey>wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY</SecretAccessKey><SessionToken>FwoGZXIvYXdzEJr//////////wEaDHer8</SessionToken><Expiration>2099-01-01T00:00:00Z</Expiration></Credentials></AssumeRoleWithWebIdentityResult></AssumeRoleWithWebIdentityResponse>"#; // gitleaks:allow
1011
1012    const STS_ASSUME_ROLE_XML: &str = r#"<AssumeRoleResponse xmlns="https://sts.amazonaws.com/doc/2011-06-15/"><AssumeRoleResult><Credentials><AccessKeyId>AKIAIOSFODNN7EXAMPLE</AccessKeyId><SecretAccessKey>wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY</SecretAccessKey><SessionToken>FwoGZXIvYXdzEJr//////////wEaDHer8</SessionToken><Expiration>2099-01-01T00:00:00Z</Expiration></Credentials></AssumeRoleResult></AssumeRoleResponse>"#; // gitleaks:allow
1013
1014    #[tokio::test]
1015    async fn test_web_identity_provider() {
1016        use std::io::Write as _;
1017        let mut server = mockito::Server::new_async().await;
1018
1019        // Write a fake JWT to a temp file
1020        let mut token_file = tempfile::NamedTempFile::new().unwrap();
1021        write!(token_file, "fake-jwt-token").unwrap();
1022
1023        let sts_url = server.url();
1024
1025        let _mock = server
1026            .mock("POST", "/")
1027            .match_query(mockito::Matcher::AllOf(vec![
1028                mockito::Matcher::UrlEncoded("Action".into(), "AssumeRoleWithWebIdentity".into()),
1029                mockito::Matcher::UrlEncoded(
1030                    "RoleArn".into(),
1031                    "arn:aws:iam::123456789012:role/TestRole".into(),
1032                ),
1033                mockito::Matcher::UrlEncoded("WebIdentityToken".into(), "fake-jwt-token".into()),
1034            ]))
1035            .with_status(200)
1036            .with_body(STS_WEB_IDENTITY_XML)
1037            .create_async()
1038            .await;
1039
1040        let client = Client::new();
1041        let service: Arc<dyn HttpService> = Arc::new(ReqwestService::new(client.clone()));
1042        let retry_config = RetryConfig::default();
1043
1044        let provider = WebIdentityProvider {
1045            token_path: token_file.path().to_str().unwrap().to_owned(),
1046            role_arn: "arn:aws:iam::123456789012:role/TestRole".into(),
1047            session_name: "test-session".into(),
1048            endpoint: sts_url,
1049        };
1050
1051        let creds = provider
1052            .fetch_token(&client, &service, &retry_config)
1053            .await
1054            .unwrap();
1055
1056        assert_eq!(creds.token.key_id, "AKIAIOSFODNN7EXAMPLE"); // gitleaks:allow
1057        assert!(creds.token.token.is_some());
1058        assert!(creds.expiry.is_some());
1059
1060        _mock.assert_async().await;
1061    }
1062
1063    #[tokio::test]
1064    async fn test_task_credential_provider() {
1065        let mut server = mockito::Server::new_async().await;
1066
1067        let endpoint = server.url();
1068
1069        let _mock = server
1070            .mock("GET", "/v2/credentials/test")
1071            .with_status(200)
1072            .with_body(r#"{"AccessKeyId":"TASKID","Code":"Success","Expiration":"2099-01-01T00:00:00Z","LastUpdated":"2022-08-30T10:21:04Z","SecretAccessKey":"TASKSECRET","Token":"TASKTOKEN","Type":"AWS-HMAC"}"#)
1073            .create_async()
1074            .await;
1075
1076        let client = Client::new();
1077        let service: Arc<dyn HttpService> = Arc::new(ReqwestService::new(client.clone()));
1078        let retry = RetryConfig::default();
1079
1080        let provider = TaskCredentialProvider {
1081            url: format!("{}/v2/credentials/test", endpoint),
1082            auth_token_file: None,
1083            retry: retry.clone(),
1084            client: client.clone(),
1085            service: service.clone(),
1086            cache: Default::default(),
1087        };
1088
1089        let creds = provider.get_credential().await.unwrap();
1090
1091        assert_eq!(creds.key_id, "TASKID");
1092        assert_eq!(creds.secret_key, "TASKSECRET");
1093        assert_eq!(creds.token.as_deref(), Some("TASKTOKEN"));
1094
1095        _mock.assert_async().await;
1096    }
1097
1098    #[tokio::test]
1099    async fn test_task_credential_with_auth_token() {
1100        use std::io::Write as _;
1101        let mut server = mockito::Server::new_async().await;
1102
1103        // Write auth token to a temp file
1104        let mut auth_file = tempfile::NamedTempFile::new().unwrap();
1105        write!(auth_file, "my-pod-identity-token").unwrap();
1106
1107        let endpoint = server.url();
1108
1109        let _mock = server
1110            .mock("GET", "/v2/credentials/eks")
1111            .match_header("Authorization", "my-pod-identity-token")
1112            .with_status(200)
1113            .with_body(r#"{"AccessKeyId":"EKSID","Code":"Success","Expiration":"2099-01-01T00:00:00Z","LastUpdated":"2022-08-30T10:21:04Z","SecretAccessKey":"EKSSECRET","Token":"EKSTOKEN","Type":"AWS-HMAC"}"#)
1114            .create_async()
1115            .await;
1116
1117        let client = Client::new();
1118        let service: Arc<dyn HttpService> = Arc::new(ReqwestService::new(client.clone()));
1119        let retry = RetryConfig::default();
1120
1121        let provider = TaskCredentialProvider {
1122            url: format!("{}/v2/credentials/eks", endpoint),
1123            auth_token_file: Some(auth_file.path().to_str().unwrap().to_owned()),
1124            retry: retry.clone(),
1125            client: client.clone(),
1126            service: service.clone(),
1127            cache: Default::default(),
1128        };
1129
1130        let creds = provider.get_credential().await.unwrap();
1131
1132        assert_eq!(creds.key_id, "EKSID");
1133        assert_eq!(creds.token.as_deref(), Some("EKSTOKEN"));
1134
1135        _mock.assert_async().await;
1136    }
1137
1138    #[tokio::test]
1139    async fn test_assume_role_provider() {
1140        use crate::StaticCredentialProvider;
1141        let mut server = mockito::Server::new_async().await;
1142
1143        let sts_endpoint = server.url();
1144
1145        let _mock = server
1146            .mock("POST", "/")
1147            .match_query(mockito::Matcher::AllOf(vec![
1148                mockito::Matcher::UrlEncoded("Action".into(), "AssumeRole".into()),
1149                mockito::Matcher::UrlEncoded(
1150                    "RoleArn".into(),
1151                    "arn:aws:iam::123456789012:role/AssumedRole".into(),
1152                ),
1153            ]))
1154            .with_status(200)
1155            .with_body(STS_ASSUME_ROLE_XML)
1156            .create_async()
1157            .await;
1158
1159        let base_cred = AwsCredential {
1160            key_id: "BASEKEYID".to_string(),
1161            secret_key: "BASESECRET".to_string(),
1162            token: None,
1163        };
1164        let base_provider: Arc<dyn CredentialProvider<Credential = AwsCredential>> =
1165            Arc::new(StaticCredentialProvider::new(base_cred));
1166
1167        let client = Client::new();
1168        let service: Arc<dyn HttpService> = Arc::new(ReqwestService::new(client.clone()));
1169        let retry_config = RetryConfig::default();
1170
1171        let provider = AssumeRoleProvider {
1172            role_arn: "arn:aws:iam::123456789012:role/AssumedRole".into(),
1173            session_name: "test-assume-session".into(),
1174            endpoint: format!("{}/", sts_endpoint),
1175            base_credentials: base_provider,
1176            region: "us-east-1".into(),
1177            policy: None,
1178        };
1179
1180        let token = provider
1181            .fetch_token(&client, &service, &retry_config)
1182            .await
1183            .unwrap();
1184
1185        assert_eq!(token.token.key_id, "AKIAIOSFODNN7EXAMPLE"); // gitleaks:allow
1186        assert!(token.token.token.is_some());
1187        assert!(token.expiry.is_some());
1188
1189        _mock.assert_async().await;
1190    }
1191}