Skip to content

Commit 4f46ebe

Browse files
authored
Merge branch 'main' into fix/bigquery_null_equals_incremental
2 parents a9bab6d + fe308ee commit 4f46ebe

5 files changed

Lines changed: 475 additions & 6 deletions

File tree

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,6 @@
1+
kind: Features
2+
body: Add STS AssumeRole support for cross-account access
3+
time: 2026-02-18T11:07:31.136966+09:00
4+
custom:
5+
Author: dtaniwaki
6+
Issue: "1657"

dbt-athena/src/dbt/adapters/athena/connections.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -61,6 +61,10 @@ class AthenaCredentials(Credentials):
6161
aws_access_key_id: Optional[str] = None
6262
aws_secret_access_key: Optional[str] = None
6363
aws_session_token: Optional[str] = None
64+
assume_role_arn: Optional[str] = None
65+
assume_role_external_id: Optional[str] = None
66+
assume_role_session_name: str = "dbt-athena"
67+
assume_role_duration_seconds: int = 3600
6468
poll_interval: float = 1.0
6569
debug_query_state: bool = False
6670
_ALIASES = {"catalog": "database"}
@@ -99,6 +103,12 @@ def _connection_keys(self) -> Tuple[str, ...]:
99103
"poll_interval",
100104
"aws_profile_name",
101105
"aws_access_key_id",
106+
"assume_role_arn",
107+
# external_id is not a secret; it is a shared condition value to prevent confused deputy attacks.
108+
# See: https://docs.aws.amazon.com/IAM/latest/UserGuide/id_roles_create_for-user_externalid.html
109+
"assume_role_external_id",
110+
"assume_role_session_name",
111+
"assume_role_duration_seconds",
102112
"endpoint_url",
103113
"s3_data_dir",
104114
"s3_data_naming",

dbt-athena/src/dbt/adapters/athena/session.py

Lines changed: 79 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,13 +1,14 @@
11
import json
22
import threading
33
import time
4-
from functools import cached_property
4+
from functools import cached_property, lru_cache
55
from hashlib import md5
6-
from typing import Any, Dict
6+
from typing import Any, Dict, NamedTuple, Optional
77
from uuid import UUID
88

99
import boto3
1010
import boto3.session
11+
from botocore.exceptions import ClientError
1112
from dbt_common.exceptions import DbtRuntimeError
1213
from dbt_common.invocation import get_invocation_id
1314

@@ -23,25 +24,99 @@
2324
spark_session_list: Dict[UUID, str] = {}
2425
spark_session_load: Dict[UUID, int] = {}
2526

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+
2695

2796
def get_boto3_session(connection: Connection) -> boto3.session.Session:
28-
return boto3.session.Session(
97+
base_session = boto3.session.Session(
2998
aws_access_key_id=connection.credentials.aws_access_key_id,
3099
aws_secret_access_key=connection.credentials.aws_secret_access_key,
31100
aws_session_token=connection.credentials.aws_session_token,
32101
region_name=connection.credentials.region_name,
33102
profile_name=connection.credentials.aws_profile_name,
34103
)
104+
if connection.credentials.assume_role_arn:
105+
return _assume_role_session(base_session, connection.credentials)
106+
return base_session
35107

36108

37109
def get_boto3_session_from_credentials(credentials: Any) -> boto3.session.Session:
38-
return boto3.session.Session(
110+
base_session = boto3.session.Session(
39111
aws_access_key_id=credentials.aws_access_key_id,
40112
aws_secret_access_key=credentials.aws_secret_access_key,
41113
aws_session_token=credentials.aws_session_token,
42114
region_name=credentials.region_name,
43115
profile_name=credentials.aws_profile_name,
44116
)
117+
if credentials.assume_role_arn:
118+
return _assume_role_session(base_session, credentials)
119+
return base_session
45120

46121

47122
class AthenaSparkSessionManager:

dbt-athena/tests/functional/conftest.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -22,5 +22,7 @@ def dbt_profile_target():
2222
"poll_interval": float(os.getenv("DBT_TEST_ATHENA_POLL_INTERVAL", "1.0")),
2323
"num_retries": int(os.getenv("DBT_TEST_ATHENA_NUM_RETRIES", "2")),
2424
"aws_profile_name": os.getenv("DBT_TEST_ATHENA_AWS_PROFILE_NAME") or None,
25+
"assume_role_arn": os.getenv("DBT_TEST_ATHENA_ASSUME_ROLE_ARN") or None,
26+
"assume_role_external_id": os.getenv("DBT_TEST_ATHENA_ASSUME_ROLE_EXTERNAL_ID") or None,
2527
"spark_work_group": os.getenv("DBT_TEST_ATHENA_SPARK_WORK_GROUP"),
2628
}

0 commit comments

Comments
 (0)