1use 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
40const STRICT_PATH_ENCODE_SET: percent_encoding::AsciiSet = STRICT_ENCODE_SET.remove(b'/');
42
43type StdError = Box<dyn std::error::Error + Send + Sync>;
44
45static EMPTY_SHA256_HASH: &str = "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855";
47
48#[derive(Eq, PartialEq)]
50pub struct AwsCredential {
51 pub key_id: String,
53 pub secret_key: String,
55 pub token: Option<String>,
57}
58
59impl 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 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#[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 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 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 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 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
265fn 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
296fn canonicalize_headers(header_map: &HeaderMap) -> (String, String) {
300 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#[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#[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
443async 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") .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#[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 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#[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#[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
653async 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#[derive(Debug)]
716pub(crate) struct TaskCredentialProvider {
717 pub url: String,
718 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
747async 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 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 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 key_id: "AKIAIOSFODNN7EXAMPLE".to_string(), secret_key: "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY".to_string(), 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" )
841 }
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 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 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 let _mock7 = server
998 .mock("PUT", "/latest/api/token")
999 .with_status(403)
1000 .with_body("")
1001 .create_async()
1002 .await;
1003
1004 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>"#; 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>"#; #[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 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"); 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 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"); assert!(token.token.token.is_some());
1187 assert!(token.expiry.is_some());
1188
1189 _mock.assert_async().await;
1190 }
1191}