Skip to content

Commit ec0941a

Browse files
committed
Fix potential issue with assume role and regions
1 parent ec77640 commit ec0941a

2 files changed

Lines changed: 63 additions & 0 deletions

File tree

src/opentaskpy/addons/aws/remotehandlers/creds.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -96,6 +96,8 @@ def get_aws_client( # pylint: disable=too-many-positional-arguments
9696
kwargs = {}
9797
if os.environ.get("AWS_ENDPOINT_URL"):
9898
kwargs["endpoint_url"] = os.environ.get("AWS_ENDPOINT_URL")
99+
if credentials.get("region_name"):
100+
kwargs["region_name"] = credentials["region_name"]
99101

100102
if assume_role_arn:
101103
logger.info(f"Assuming role: {assume_role_arn}")

tests/test_remotehandler_s3_transfer.py

Lines changed: 61 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717
from pytest_shell import fs
1818

1919
from opentaskpy import exceptions
20+
from opentaskpy.addons.aws.remotehandlers.creds import get_aws_client
2021
from opentaskpy.addons.aws.remotehandlers.s3 import S3Transfer
2122
from tests.fixtures.localstack import *
2223

@@ -611,6 +612,66 @@ def fake_get_aws_client(
611612
assert handler.s3_client is None
612613

613614

615+
def test_get_aws_client_passes_region_to_sts(monkeypatch):
616+
captured = {}
617+
618+
class DummySTSClient:
619+
def assume_role(self, **kwargs):
620+
captured["assume_role_kwargs"] = kwargs
621+
return {
622+
"Credentials": {
623+
"AccessKeyId": "assumed-key",
624+
"SecretAccessKey": "assumed-secret",
625+
"SessionToken": "assumed-token",
626+
}
627+
}
628+
629+
class DummySession:
630+
def __init__(self, **kwargs):
631+
captured["session_kwargs"] = kwargs
632+
633+
def client(self, client_type, **kwargs):
634+
captured["session_client_type"] = client_type
635+
captured["session_client_kwargs"] = kwargs
636+
return object()
637+
638+
def fake_boto3_client(service_name, **kwargs):
639+
captured["sts_service_name"] = service_name
640+
captured["sts_client_kwargs"] = kwargs
641+
return DummySTSClient()
642+
643+
monkeypatch.setenv("AWS_ENDPOINT_URL", "http://floci.test")
644+
monkeypatch.setattr(
645+
"opentaskpy.addons.aws.remotehandlers.creds.boto3.client",
646+
fake_boto3_client,
647+
)
648+
monkeypatch.setattr(
649+
"opentaskpy.addons.aws.remotehandlers.creds.boto3.session.Session",
650+
DummySession,
651+
)
652+
653+
result = get_aws_client(
654+
"s3",
655+
{
656+
"AccessKeyId": "test-key",
657+
"SecretAccessKey": "test-secret",
658+
"region_name": "eu-west-1",
659+
},
660+
assume_role_arn="arn:aws:iam::012345678900:role/dummy-role",
661+
)
662+
663+
assert captured["sts_service_name"] == "sts"
664+
assert captured["sts_client_kwargs"]["endpoint_url"] == "http://floci.test"
665+
assert captured["sts_client_kwargs"]["region_name"] == "eu-west-1"
666+
assert captured["assume_role_kwargs"]["RoleArn"] == (
667+
"arn:aws:iam::012345678900:role/dummy-role"
668+
)
669+
assert captured["session_kwargs"]["aws_access_key_id"] == "assumed-key"
670+
assert captured["session_client_type"] == "s3"
671+
assert captured["session_client_kwargs"]["endpoint_url"] == "http://floci.test"
672+
assert result["temporary_creds"]["AccessKeyId"] == "assumed-key"
673+
674+
614675
def test_s3_file_watch(s3_client, setup_bucket, tmp_path):
615676
transfer_obj = transfer.Transfer(
616677
None, "s3-file-watch", s3_file_watch_task_definition

0 commit comments

Comments
 (0)