Skip to content

Commit d8a79b4

Browse files
committed
fix(athena): install assumed session via boto3.DEFAULT_SESSION to avoid s3 client double-injection
1 parent dbd5e2c commit d8a79b4

4 files changed

Lines changed: 49 additions & 4 deletions

File tree

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

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -38,7 +38,7 @@ def spark_engine_version(self) -> str:
3838

3939
@property
4040
def is_spark_connect(self) -> bool:
41-
"""True when the model requests Apache Spark 3.5+ via Spark Connect."""
41+
"""True when the model requests Apache Spark 3.5, which runs via Spark Connect."""
4242
return self.spark_engine_version == "3.5"
4343

4444
def set_timeout(self) -> int:

dbt-athena/src/dbt/adapters/athena/spark_connect/job.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -254,7 +254,10 @@ def _install_assumed_default_session(self) -> None:
254254
if not self.credentials.assume_role_arn:
255255
return
256256
assumed = get_boto3_session_from_credentials(self.credentials)
257-
boto3.setup_default_session(botocore_session=assumed._session)
257+
# Assign directly; setup_default_session(botocore_session=assumed._session)
258+
# re-registers the creating-client-class.s3 handler and breaks the model's
259+
# first boto3.client("s3") with a duplicate upload_file injection error.
260+
boto3.DEFAULT_SESSION = assumed
258261

259262
def _is_transient_failure(self, e: BaseException) -> bool:
260263
return is_transient_spark_error(e)

dbt-athena/tests/functional/adapter/test_spark_connect_python_submissions.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -138,6 +138,7 @@ def model(dbt, spark):
138138
dbt.config(materialized='table', spark_engine_version='3.5')
139139
import boto3
140140
141+
boto3.client("s3")
141142
arn = boto3.client("sts").get_caller_identity()["Arn"]
142143
return spark.createDataFrame([(arn,)], ["caller_arn"])
143144
"""

dbt-athena/tests/unit/spark_connect/test_job.py

Lines changed: 43 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,9 @@
44
import time
55
from unittest.mock import MagicMock, Mock, patch
66

7+
import boto3
78
import botocore.exceptions
9+
import botocore.session
810
import pytest
911
from dbt_common.exceptions import DbtRuntimeError
1012

@@ -200,7 +202,6 @@ def test_assume_role_installs_assumed_default_session(
200202
self._stub_endpoint_and_channel(submitter, monkeypatch)
201203

202204
assumed_session = Mock()
203-
assumed_session._session = Mock()
204205
get_session = Mock(return_value=assumed_session)
205206
setup_default = Mock()
206207
monkeypatch.setattr(
@@ -210,11 +211,46 @@ def test_assume_role_installs_assumed_default_session(
210211
monkeypatch.setattr(
211212
"dbt.adapters.athena.spark_connect.job.boto3.setup_default_session", setup_default
212213
)
214+
monkeypatch.setattr(
215+
"dbt.adapters.athena.spark_connect.job.boto3.DEFAULT_SESSION", None, raising=False
216+
)
213217

214218
submitter.submit("x = 1")
215219

216220
get_session.assert_called_once_with(mock_credentials)
217-
setup_default.assert_called_once_with(botocore_session=assumed_session._session)
221+
assert boto3.DEFAULT_SESSION is assumed_session
222+
setup_default.assert_not_called()
223+
224+
def test_assume_role_default_session_survives_s3_client_creation(
225+
self, mock_credentials, spark_connect_parsed_model, monkeypatch
226+
):
227+
# Regression: re-wrapping the assumed session as the default double-
228+
# registers boto3's s3 client-class handler, so boto3.client("s3")
229+
# raises a duplicate upload_file injection.
230+
mock_credentials.assume_role_arn = "arn:aws:iam::123456789012:role/dbt"
231+
mock_pool = Mock()
232+
mock_pool.acquire.return_value = "sid-1"
233+
submitter = self._make_submitter(spark_connect_parsed_model, mock_credentials, mock_pool)
234+
self._stub_endpoint_and_channel(submitter, monkeypatch)
235+
236+
base = botocore.session.Session()
237+
base.set_config_variable("region", "us-east-1")
238+
assumed_session = boto3.session.Session(
239+
botocore_session=base,
240+
aws_access_key_id="AKIA_ASSUMED",
241+
aws_secret_access_key="SECRET",
242+
aws_session_token="TOKEN",
243+
)
244+
monkeypatch.setattr(
245+
"dbt.adapters.athena.spark_connect.job.get_boto3_session_from_credentials",
246+
Mock(return_value=assumed_session),
247+
)
248+
monkeypatch.setattr(boto3, "DEFAULT_SESSION", None, raising=False)
249+
250+
submitter.submit("x = 1")
251+
252+
boto3.client("s3", region_name="us-east-1")
253+
assert boto3.DEFAULT_SESSION is assumed_session
218254

219255
def test_no_assume_role_leaves_default_session_untouched(
220256
self, mock_credentials, spark_connect_parsed_model, monkeypatch
@@ -227,18 +263,23 @@ def test_no_assume_role_leaves_default_session_untouched(
227263

228264
get_session = Mock()
229265
setup_default = Mock()
266+
sentinel = Mock()
230267
monkeypatch.setattr(
231268
"dbt.adapters.athena.spark_connect.job.get_boto3_session_from_credentials",
232269
get_session,
233270
)
234271
monkeypatch.setattr(
235272
"dbt.adapters.athena.spark_connect.job.boto3.setup_default_session", setup_default
236273
)
274+
monkeypatch.setattr(
275+
"dbt.adapters.athena.spark_connect.job.boto3.DEFAULT_SESSION", sentinel, raising=False
276+
)
237277

238278
submitter.submit("x = 1")
239279

240280
get_session.assert_not_called()
241281
setup_default.assert_not_called()
282+
assert boto3.DEFAULT_SESSION is sentinel
242283

243284
def test_transient_error_retries_with_new_session(
244285
self, mock_credentials, spark_connect_parsed_model, monkeypatch

0 commit comments

Comments
 (0)