|
1 | 1 | import json |
2 | 2 | import threading |
3 | 3 | import time |
4 | | -from functools import cached_property |
| 4 | +from functools import cached_property, lru_cache |
5 | 5 | from hashlib import md5 |
6 | | -from typing import Any, Dict |
| 6 | +from typing import Any, Dict, NamedTuple, Optional |
7 | 7 | from uuid import UUID |
8 | 8 |
|
9 | 9 | import boto3 |
10 | 10 | import boto3.session |
| 11 | +from botocore.exceptions import ClientError |
11 | 12 | from dbt_common.exceptions import DbtRuntimeError |
12 | 13 | from dbt_common.invocation import get_invocation_id |
13 | 14 |
|
|
23 | 24 | spark_session_list: Dict[UUID, str] = {} |
24 | 25 | spark_session_load: Dict[UUID, int] = {} |
25 | 26 |
|
| 27 | +# Refresh credentials this many seconds before actual expiration |
| 28 | +_EXPIRY_BUFFER_SECONDS = 300 |
| 29 | + |
| 30 | + |
| 31 | +class _AssumeRoleParams(NamedTuple): |
| 32 | + assume_role_arn: Optional[str] |
| 33 | + assume_role_external_id: Optional[str] |
| 34 | + assume_role_session_name: str |
| 35 | + assume_role_duration_seconds: int |
| 36 | + region_name: str |
| 37 | + num_retries: int |
| 38 | + |
| 39 | + |
| 40 | +def _assume_role_session( |
| 41 | + base_session: boto3.session.Session, |
| 42 | + credentials: Any, |
| 43 | +) -> boto3.session.Session: |
| 44 | + duration = credentials.assume_role_duration_seconds |
| 45 | + # Valid range is 900–43200 seconds per AWS STS docs: |
| 46 | + # https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRole.html#API_AssumeRole_RequestParameters |
| 47 | + if not (900 <= duration <= 43200): |
| 48 | + raise DbtRuntimeError( |
| 49 | + f"assume_role_duration_seconds must be between 900 and 43200, got {duration}" |
| 50 | + ) |
| 51 | + ttl = duration - _EXPIRY_BUFFER_SECONDS |
| 52 | + key = _AssumeRoleParams( |
| 53 | + assume_role_arn=credentials.assume_role_arn, |
| 54 | + assume_role_external_id=credentials.assume_role_external_id, |
| 55 | + assume_role_session_name=credentials.assume_role_session_name, |
| 56 | + assume_role_duration_seconds=duration, |
| 57 | + region_name=credentials.region_name, |
| 58 | + num_retries=credentials.effective_num_retries, |
| 59 | + ) |
| 60 | + # Increments every ttl seconds, causing lru_cache to treat each period as a distinct call |
| 61 | + ttl_bucket = int(time.time() / ttl) |
| 62 | + return _get_assume_role_session(base_session, key, ttl_bucket) |
| 63 | + |
| 64 | + |
| 65 | +@lru_cache(maxsize=1) |
| 66 | +def _get_assume_role_session( |
| 67 | + base_session: boto3.session.Session, |
| 68 | + key: _AssumeRoleParams, |
| 69 | + _ttl_hash: int, # artificial value that changes every ttl seconds to force cache invalidation |
| 70 | +) -> boto3.session.Session: |
| 71 | + LOGGER.debug(f"Assuming role: {key.assume_role_arn}") |
| 72 | + sts_client = base_session.client( |
| 73 | + "sts", |
| 74 | + config=get_boto3_config(key.num_retries), |
| 75 | + ) |
| 76 | + kwargs: Dict[str, Any] = { |
| 77 | + "RoleArn": key.assume_role_arn, |
| 78 | + "RoleSessionName": key.assume_role_session_name, |
| 79 | + } |
| 80 | + if key.assume_role_external_id: |
| 81 | + kwargs["ExternalId"] = key.assume_role_external_id |
| 82 | + if key.assume_role_duration_seconds: |
| 83 | + kwargs["DurationSeconds"] = key.assume_role_duration_seconds |
| 84 | + try: |
| 85 | + response = sts_client.assume_role(**kwargs) |
| 86 | + except ClientError as e: |
| 87 | + raise DbtRuntimeError(f"Failed to assume role {key.assume_role_arn}: {e}") from e |
| 88 | + return boto3.session.Session( |
| 89 | + aws_access_key_id=response["Credentials"]["AccessKeyId"], |
| 90 | + aws_secret_access_key=response["Credentials"]["SecretAccessKey"], |
| 91 | + aws_session_token=response["Credentials"]["SessionToken"], |
| 92 | + region_name=key.region_name, |
| 93 | + ) |
| 94 | + |
26 | 95 |
|
27 | 96 | def get_boto3_session(connection: Connection) -> boto3.session.Session: |
28 | | - return boto3.session.Session( |
| 97 | + base_session = boto3.session.Session( |
29 | 98 | aws_access_key_id=connection.credentials.aws_access_key_id, |
30 | 99 | aws_secret_access_key=connection.credentials.aws_secret_access_key, |
31 | 100 | aws_session_token=connection.credentials.aws_session_token, |
32 | 101 | region_name=connection.credentials.region_name, |
33 | 102 | profile_name=connection.credentials.aws_profile_name, |
34 | 103 | ) |
| 104 | + if connection.credentials.assume_role_arn: |
| 105 | + return _assume_role_session(base_session, connection.credentials) |
| 106 | + return base_session |
35 | 107 |
|
36 | 108 |
|
37 | 109 | def get_boto3_session_from_credentials(credentials: Any) -> boto3.session.Session: |
38 | | - return boto3.session.Session( |
| 110 | + base_session = boto3.session.Session( |
39 | 111 | aws_access_key_id=credentials.aws_access_key_id, |
40 | 112 | aws_secret_access_key=credentials.aws_secret_access_key, |
41 | 113 | aws_session_token=credentials.aws_session_token, |
42 | 114 | region_name=credentials.region_name, |
43 | 115 | profile_name=credentials.aws_profile_name, |
44 | 116 | ) |
| 117 | + if credentials.assume_role_arn: |
| 118 | + return _assume_role_session(base_session, credentials) |
| 119 | + return base_session |
45 | 120 |
|
46 | 121 |
|
47 | 122 | class AthenaSparkSessionManager: |
|
0 commit comments