Skip to main content

reqsign_aws_core/provide_credential/
ecs.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 crate::Credential;
19use http::{HeaderValue, Method, Request, StatusCode};
20use log::debug;
21use reqsign_core::{Context, Error, ProvideCredential, Result};
22use serde::Deserialize;
23
24const AWS_CONTAINER_CREDENTIALS_RELATIVE_URI: &str = "AWS_CONTAINER_CREDENTIALS_RELATIVE_URI";
25const AWS_CONTAINER_CREDENTIALS_FULL_URI: &str = "AWS_CONTAINER_CREDENTIALS_FULL_URI";
26const AWS_CONTAINER_AUTHORIZATION_TOKEN: &str = "AWS_CONTAINER_AUTHORIZATION_TOKEN";
27const AWS_CONTAINER_AUTHORIZATION_TOKEN_FILE: &str = "AWS_CONTAINER_AUTHORIZATION_TOKEN_FILE";
28const ECS_METADATA_ENDPOINT: &str = "http://169.254.170.2";
29
30/// ECS Task Role Credentials Provider
31///
32/// This provider fetches IAM credentials from the ECS Task IAM Roles endpoint.
33/// It supports both relative URI (ECS) and full URI (Fargate) modes.
34///
35/// # Important Note
36/// This provider fetches IAM **credentials**, not task metadata. The credentials
37/// endpoint is separate from the task metadata endpoints (v2/v3/v4). While metadata
38/// endpoints provide information about the task and container, the credentials
39/// endpoint provides IAM role credentials for authentication.
40///
41/// # Configuration
42///
43/// Configuration values can be provided directly via builder methods or through environment
44/// variables. Direct configuration takes precedence over environment variables.
45///
46/// ## Builder Methods
47/// - [`ECSCredentialProvider::with_relative_uri`]: Set the relative URI for ECS environments
48/// - [`ECSCredentialProvider::with_endpoint`]: Set a complete custom endpoint URL (for Fargate or custom setups)
49/// - [`ECSCredentialProvider::with_auth_token`]: Set the authorization token directly
50/// - [`ECSCredentialProvider::with_auth_token_file`]: Set the path to the authorization token file
51/// - [`ECSCredentialProvider::with_metadata_uri_override`]: Override the base metadata endpoint for relative URIs
52///
53/// ## Environment Variables (Fallback)
54/// - `AWS_CONTAINER_CREDENTIALS_RELATIVE_URI`: Relative URI to fetch credentials (ECS)
55/// - `AWS_CONTAINER_CREDENTIALS_FULL_URI`: Full URI to fetch credentials (Fargate)
56/// - `AWS_CONTAINER_AUTHORIZATION_TOKEN`: Authorization token for the request
57/// - `AWS_CONTAINER_AUTHORIZATION_TOKEN_FILE`: File containing the authorization token
58///
59/// # Examples
60///
61/// ```rust,no_run
62/// use reqsign_aws_core::ECSCredentialProvider;
63///
64/// // Configure for ECS with relative URI
65/// let provider = ECSCredentialProvider::new()
66///     .with_relative_uri("/v2/credentials/task-role")
67///     .with_auth_token("my-auth-token");
68///
69/// // Configure for Fargate with full endpoint
70/// let provider = ECSCredentialProvider::new()
71///     .with_endpoint("http://169.254.170.2/v2/credentials/task-role")
72///     .with_auth_token_file("/tmp/auth-token");
73/// ```
74#[derive(Debug, Clone)]
75pub struct ECSCredentialProvider {
76    endpoint: Option<String>,
77    auth_token: Option<String>,
78    auth_token_file: Option<String>,
79    relative_uri: Option<String>,
80    metadata_uri_override: Option<String>,
81}
82
83impl Default for ECSCredentialProvider {
84    fn default() -> Self {
85        Self::new()
86    }
87}
88
89impl ECSCredentialProvider {
90    /// Create a new ECS credential provider
91    pub fn new() -> Self {
92        Self {
93            endpoint: None,
94            auth_token: None,
95            auth_token_file: None,
96            relative_uri: None,
97            metadata_uri_override: None,
98        }
99    }
100
101    /// Create with custom endpoint
102    pub fn with_endpoint(mut self, endpoint: impl Into<String>) -> Self {
103        self.endpoint = Some(endpoint.into());
104        self
105    }
106
107    /// Create with custom auth token
108    pub fn with_auth_token(mut self, token: impl Into<String>) -> Self {
109        self.auth_token = Some(token.into());
110        self
111    }
112
113    /// Create with custom auth token file path
114    pub fn with_auth_token_file(mut self, file_path: impl Into<String>) -> Self {
115        self.auth_token_file = Some(file_path.into());
116        self
117    }
118
119    /// Create with container credentials relative URI
120    /// This is used in ECS environments where the base metadata URI is known
121    pub fn with_relative_uri(mut self, uri: impl Into<String>) -> Self {
122        self.relative_uri = Some(uri.into());
123        self
124    }
125
126    /// Override the metadata URI base endpoint (typically for testing)
127    /// Defaults to http://169.254.170.2 if not specified
128    pub fn with_metadata_uri_override(mut self, uri: impl Into<String>) -> Self {
129        self.metadata_uri_override = Some(uri.into());
130        self
131    }
132
133    async fn load_auth_token(&self, ctx: &Context) -> Result<Option<String>> {
134        // If auth token is already set, use it
135        if let Some(token) = &self.auth_token {
136            return Ok(Some(token.clone()));
137        }
138
139        // Try to get token from configured file first
140        if let Some(token_file) = &self.auth_token_file {
141            let token = ctx.file_read(token_file).await.map_err(|e| {
142                Error::config_invalid("failed to read ECS auth token file")
143                    .with_source(e)
144                    .with_context(format!("file: {token_file}"))
145            })?;
146            return Ok(Some(String::from_utf8_lossy(&token).trim().to_string()));
147        }
148
149        // Try to get token from environment
150        if let Some(token) = ctx.env_var(AWS_CONTAINER_AUTHORIZATION_TOKEN) {
151            return Ok(Some(token));
152        }
153
154        // Try to get token from environment file
155        if let Some(token_file) = ctx.env_var(AWS_CONTAINER_AUTHORIZATION_TOKEN_FILE) {
156            let token = ctx.file_read(&token_file).await.map_err(|e| {
157                Error::config_invalid("failed to read ECS auth token file")
158                    .with_source(e)
159                    .with_context(format!("file: {token_file}"))
160            })?;
161            return Ok(Some(String::from_utf8_lossy(&token).trim().to_string()));
162        }
163
164        Ok(None)
165    }
166
167    fn get_endpoint(&self, ctx: &Context) -> Result<String> {
168        // Use custom endpoint if provided (highest priority)
169        if let Some(endpoint) = &self.endpoint {
170            return Ok(endpoint.clone());
171        }
172
173        // Try configured relative URI (ECS)
174        if let Some(relative_uri) = &self.relative_uri {
175            let base_endpoint = self
176                .metadata_uri_override
177                .as_deref()
178                .unwrap_or(ECS_METADATA_ENDPOINT);
179            return Ok(format!("{base_endpoint}{relative_uri}"));
180        }
181
182        // Fall back to environment variables
183        // Try full URI from environment (Fargate)
184        if let Some(full_uri) = ctx.env_var(AWS_CONTAINER_CREDENTIALS_FULL_URI) {
185            return Ok(full_uri);
186        }
187
188        // Try relative URI from environment (ECS)
189        if let Some(relative_uri) = ctx.env_var(AWS_CONTAINER_CREDENTIALS_RELATIVE_URI) {
190            let base_endpoint = self
191                .metadata_uri_override
192                .as_deref()
193                .unwrap_or(ECS_METADATA_ENDPOINT);
194            return Ok(format!("{base_endpoint}{relative_uri}"));
195        }
196
197        Err(Error::config_invalid(
198            "ECS container credentials endpoint not configured"
199        )
200        .with_context("hint: use with_relative_uri(), with_endpoint(), or set AWS_CONTAINER_CREDENTIALS_RELATIVE_URI/AWS_CONTAINER_CREDENTIALS_FULL_URI")
201        .with_context("note: are you running on ECS or Fargate?"))
202    }
203}
204
205#[derive(Debug, Deserialize)]
206#[serde(rename_all = "PascalCase")]
207struct ECSCredentialResponse {
208    access_key_id: String,
209    secret_access_key: String,
210    token: String,
211    expiration: String,
212}
213impl ProvideCredential for ECSCredentialProvider {
214    type Credential = Credential;
215
216    async fn provide_credential(&self, ctx: &Context) -> Result<Option<Self::Credential>> {
217        let endpoint = match self.get_endpoint(ctx) {
218            Ok(ep) => ep,
219            Err(_) => {
220                debug!("ECS credential provider: no container credentials endpoint found");
221                return Ok(None);
222            }
223        };
224
225        debug!("ECS credential provider: fetching credentials from {endpoint}");
226
227        let mut req = Request::builder()
228            .method(Method::GET)
229            .uri(&endpoint)
230            .body(bytes::Bytes::new())
231            .map_err(|e| {
232                Error::request_invalid("failed to build ECS credentials request")
233                    .with_source(e)
234                    .with_context(format!("endpoint: {endpoint}"))
235            })?;
236
237        // Add authorization token if available
238        if let Some(token) = self.load_auth_token(ctx).await? {
239            req.headers_mut().insert(
240                "Authorization",
241                HeaderValue::from_str(&token).map_err(|e| {
242                    Error::config_invalid("invalid ECS authorization token")
243                        .with_source(e)
244                        .with_context("token_source: environment or file")
245                })?,
246            );
247        }
248
249        let resp = ctx.http_send(req).await.map_err(|e| {
250            Error::unexpected("failed to fetch ECS credentials")
251                .with_source(e)
252                .with_context(format!("endpoint: {endpoint}"))
253                .with_context("hint: check if running on ECS/Fargate with proper IAM role")
254                .set_retryable(true)
255        })?;
256
257        if resp.status() != StatusCode::OK {
258            let status = resp.status();
259            let body = String::from_utf8_lossy(resp.body());
260
261            let error = match status.as_u16() {
262                401 | 403 => Error::permission_denied(format!(
263                    "ECS task not authorized to fetch credentials: {body}"
264                ))
265                .with_context("hint: check if task has proper IAM role attached"),
266                404 => Error::config_invalid("ECS credentials endpoint not found")
267                    .with_context(format!("endpoint: {endpoint}"))
268                    .with_context("hint: verify the container credentials URI"),
269                500..=599 => Error::unexpected(format!("ECS metadata service error: {body}"))
270                    .set_retryable(true),
271                _ => Error::unexpected(format!(
272                    "ECS metadata endpoint returned unexpected status {status}: {body}"
273                )),
274            };
275
276            return Err(error
277                .with_context(format!("http_status: {status}"))
278                .with_context(format!("endpoint: {endpoint}")));
279        }
280
281        let body = resp.into_body();
282        let creds: ECSCredentialResponse = serde_json::from_slice(&body).map_err(|e| {
283            Error::unexpected("failed to parse ECS credentials response")
284                .with_source(e)
285                .with_context(format!("response_length: {}", body.len()))
286                .with_context(format!("endpoint: {endpoint}"))
287        })?;
288
289        let expires_in = creds.expiration.parse().map_err(|e| {
290            Error::unexpected("failed to parse ECS credential expiration")
291                .with_source(e)
292                .with_context(format!("expiration_value: {}", creds.expiration))
293        })?;
294
295        Ok(Some(Credential {
296            access_key_id: creds.access_key_id,
297            secret_access_key: creds.secret_access_key,
298            session_token: Some(creds.token),
299            expires_in: Some(expires_in),
300        }))
301    }
302}
303
304#[cfg(test)]
305mod tests {
306    use super::*;
307    use reqsign_core::StaticEnv;
308    use reqsign_file_read_tokio::TokioFileRead;
309    use reqsign_http_send_reqwest::ReqwestHttpSend;
310    use std::collections::HashMap;
311
312    #[tokio::test]
313    async fn test_ecs_provider_no_env() {
314        let ctx = Context::new()
315            .with_file_read(TokioFileRead)
316            .with_http_send(ReqwestHttpSend::default());
317        let ctx = ctx.with_env(StaticEnv {
318            home_dir: None,
319            envs: HashMap::new(),
320        });
321
322        let provider = ECSCredentialProvider::new();
323        let result = provider.provide_credential(&ctx).await.unwrap();
324        assert!(result.is_none());
325    }
326
327    #[tokio::test]
328    async fn test_get_endpoint_relative_uri() {
329        let ctx = Context::new()
330            .with_file_read(TokioFileRead)
331            .with_http_send(ReqwestHttpSend::default());
332        let ctx = ctx.with_env(StaticEnv {
333            home_dir: None,
334            envs: HashMap::from_iter([(
335                AWS_CONTAINER_CREDENTIALS_RELATIVE_URI.to_string(),
336                "/v2/credentials/task-role".to_string(),
337            )]),
338        });
339
340        let provider = ECSCredentialProvider::new();
341        let endpoint = provider.get_endpoint(&ctx).unwrap();
342        assert_eq!(endpoint, "http://169.254.170.2/v2/credentials/task-role");
343    }
344
345    #[tokio::test]
346    async fn test_fargate_metadata_uri_does_not_override_credentials_endpoint() {
347        let ctx = Context::new()
348            .with_file_read(TokioFileRead)
349            .with_http_send(ReqwestHttpSend::default());
350        let ctx = ctx.with_env(StaticEnv {
351            home_dir: None,
352            envs: HashMap::from_iter([
353                (
354                    AWS_CONTAINER_CREDENTIALS_RELATIVE_URI.to_string(),
355                    "/creds".to_string(),
356                ),
357                (
358                    "ECS_CONTAINER_METADATA_URI".to_string(),
359                    "http://169.254.170.2/v3/task-id".to_string(),
360                ),
361            ]),
362        });
363
364        let provider = ECSCredentialProvider::new();
365        let endpoint = provider.get_endpoint(&ctx).unwrap();
366        assert_eq!(endpoint, "http://169.254.170.2/creds");
367    }
368
369    #[tokio::test]
370    async fn test_get_endpoint_full_uri() {
371        let ctx = Context::new()
372            .with_file_read(TokioFileRead)
373            .with_http_send(ReqwestHttpSend::default());
374        let ctx = ctx.with_env(StaticEnv {
375            home_dir: None,
376            envs: HashMap::from_iter([(
377                AWS_CONTAINER_CREDENTIALS_FULL_URI.to_string(),
378                "http://localhost:8080/credentials".to_string(),
379            )]),
380        });
381
382        let provider = ECSCredentialProvider::new();
383        let endpoint = provider.get_endpoint(&ctx).unwrap();
384        assert_eq!(endpoint, "http://localhost:8080/credentials");
385    }
386
387    #[tokio::test]
388    async fn test_custom_endpoint() {
389        let ctx = Context::new()
390            .with_file_read(TokioFileRead)
391            .with_http_send(ReqwestHttpSend::default());
392        let provider = ECSCredentialProvider::new().with_endpoint("http://custom-endpoint/creds");
393
394        let endpoint = provider.get_endpoint(&ctx).unwrap();
395        assert_eq!(endpoint, "http://custom-endpoint/creds");
396    }
397
398    #[tokio::test]
399    async fn test_configured_relative_uri() {
400        let ctx = Context::new()
401            .with_file_read(TokioFileRead)
402            .with_http_send(ReqwestHttpSend::default())
403            .with_env(StaticEnv {
404                home_dir: None,
405                envs: HashMap::new(),
406            });
407
408        let provider = ECSCredentialProvider::new().with_relative_uri("/v2/credentials/task-role");
409
410        let endpoint = provider.get_endpoint(&ctx).unwrap();
411        assert_eq!(endpoint, "http://169.254.170.2/v2/credentials/task-role");
412    }
413
414    #[tokio::test]
415    async fn test_configured_relative_uri_with_custom_base() {
416        let ctx = Context::new()
417            .with_file_read(TokioFileRead)
418            .with_http_send(ReqwestHttpSend::default())
419            .with_env(StaticEnv {
420                home_dir: None,
421                envs: HashMap::new(),
422            });
423
424        let provider = ECSCredentialProvider::new()
425            .with_relative_uri("/creds")
426            .with_metadata_uri_override("http://localhost:51679");
427
428        let endpoint = provider.get_endpoint(&ctx).unwrap();
429        assert_eq!(endpoint, "http://localhost:51679/creds");
430    }
431
432    #[tokio::test]
433    async fn test_configured_values_override_env() {
434        let ctx = Context::new()
435            .with_file_read(TokioFileRead)
436            .with_http_send(ReqwestHttpSend::default())
437            .with_env(StaticEnv {
438                home_dir: None,
439                envs: HashMap::from_iter([
440                    (
441                        AWS_CONTAINER_CREDENTIALS_FULL_URI.to_string(),
442                        "http://env-endpoint/creds".to_string(),
443                    ),
444                    (
445                        AWS_CONTAINER_CREDENTIALS_RELATIVE_URI.to_string(),
446                        "/env-relative".to_string(),
447                    ),
448                ]),
449            });
450
451        let provider =
452            ECSCredentialProvider::new().with_endpoint("http://configured-endpoint/creds");
453
454        let endpoint = provider.get_endpoint(&ctx).unwrap();
455        // Configured value should override environment
456        assert_eq!(endpoint, "http://configured-endpoint/creds");
457    }
458
459    #[tokio::test]
460    async fn test_priority_order() {
461        let ctx = Context::new()
462            .with_file_read(TokioFileRead)
463            .with_http_send(ReqwestHttpSend::default())
464            .with_env(StaticEnv {
465                home_dir: None,
466                envs: HashMap::from_iter([(
467                    AWS_CONTAINER_CREDENTIALS_FULL_URI.to_string(),
468                    "http://env-full-uri/creds".to_string(),
469                )]),
470            });
471
472        // Test priority: custom endpoint > relative URI > env
473        let provider = ECSCredentialProvider::new()
474            .with_endpoint("http://custom/creds")
475            .with_relative_uri("/relative");
476
477        let endpoint = provider.get_endpoint(&ctx).unwrap();
478        assert_eq!(endpoint, "http://custom/creds");
479
480        // Test without custom endpoint
481        let provider = ECSCredentialProvider::new().with_relative_uri("/relative");
482
483        let endpoint = provider.get_endpoint(&ctx).unwrap();
484        assert_eq!(endpoint, "http://169.254.170.2/relative");
485    }
486
487    #[tokio::test]
488    async fn test_configured_auth_token() {
489        let ctx = Context::new()
490            .with_file_read(TokioFileRead)
491            .with_http_send(ReqwestHttpSend::default())
492            .with_env(StaticEnv {
493                home_dir: None,
494                envs: HashMap::from_iter([(
495                    AWS_CONTAINER_AUTHORIZATION_TOKEN.to_string(),
496                    "env-token".to_string(),
497                )]),
498            });
499
500        let provider = ECSCredentialProvider::new().with_auth_token("configured-token");
501
502        let token = provider.load_auth_token(&ctx).await.unwrap();
503        // Configured token should override environment
504        assert_eq!(token, Some("configured-token".to_string()));
505    }
506
507    #[tokio::test]
508    async fn test_configured_auth_token_file() {
509        use std::io::Write;
510        use tempfile::NamedTempFile;
511
512        let mut temp_file = NamedTempFile::new().unwrap();
513        writeln!(temp_file, "file-token").unwrap();
514        let temp_path = temp_file.path().to_str().unwrap();
515
516        let ctx = Context::new()
517            .with_file_read(TokioFileRead)
518            .with_http_send(ReqwestHttpSend::default())
519            .with_env(StaticEnv {
520                home_dir: None,
521                envs: HashMap::from_iter([
522                    (
523                        AWS_CONTAINER_AUTHORIZATION_TOKEN.to_string(),
524                        "env-token".to_string(),
525                    ),
526                    (
527                        AWS_CONTAINER_AUTHORIZATION_TOKEN_FILE.to_string(),
528                        "env-file.txt".to_string(),
529                    ),
530                ]),
531            });
532
533        let provider = ECSCredentialProvider::new().with_auth_token_file(temp_path);
534
535        let token = provider.load_auth_token(&ctx).await.unwrap();
536        // Configured file should override environment
537        assert_eq!(token, Some("file-token".to_string()));
538    }
539
540    #[tokio::test]
541    async fn test_auth_token_priority() {
542        use std::io::Write;
543        use tempfile::NamedTempFile;
544
545        let mut temp_file = NamedTempFile::new().unwrap();
546        writeln!(temp_file, "file-token").unwrap();
547        let temp_path = temp_file.path().to_str().unwrap();
548
549        let ctx = Context::new()
550            .with_file_read(TokioFileRead)
551            .with_http_send(ReqwestHttpSend::default())
552            .with_env(StaticEnv {
553                home_dir: None,
554                envs: HashMap::new(),
555            });
556
557        // Test priority: direct token > configured file > env token > env file
558        let provider = ECSCredentialProvider::new()
559            .with_auth_token("direct-token")
560            .with_auth_token_file(temp_path);
561
562        let token = provider.load_auth_token(&ctx).await.unwrap();
563        assert_eq!(token, Some("direct-token".to_string()));
564
565        // Test without direct token
566        let provider = ECSCredentialProvider::new().with_auth_token_file(temp_path);
567
568        let token = provider.load_auth_token(&ctx).await.unwrap();
569        assert_eq!(token, Some("file-token".to_string()));
570    }
571}