Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
63 changes: 62 additions & 1 deletion tap_dynamodb/connectors/aws_boto_connector.py
Original file line number Diff line number Diff line change
Expand Up @@ -76,6 +76,41 @@
).to_dict()


def _config_or_env(config: dict, config_key: str, *env_keys: str) -> str | None:
value = config.get(config_key)
if value:
return value
for key in env_keys:
value = os.environ.get(key)
if value:
return value
return None


_TAP_DYNAMODB_AUTH_ENV_KEYS = (
"TAP_DYNAMODB_USE_AWS_ENV_VARS",
"TAP_DYNAMODB_AWS_ACCESS_KEY_ID",
"TAP_DYNAMODB_AWS_SECRET_ACCESS_KEY",
"TAP_DYNAMODB_AWS_SESSION_TOKEN",
"TAP_DYNAMODB_AWS_PROFILE",
"TAP_DYNAMODB_AWS_DEFAULT_REGION",
"TAP_DYNAMODB_AWS_ENDPOINT_URL",
"TAP_DYNAMODB_AWS_ASSUME_ROLE_ARN",
)

_AWS_AUTH_ENV_KEYS = (
"AWS_ACCESS_KEY_ID",
"AWS_SECRET_ACCESS_KEY",
"AWS_SESSION_TOKEN",
"AWS_PROFILE",
"AWS_DEFAULT_REGION",
)


def _env_presence(env_keys: tuple[str, ...]) -> dict[str, str]:
return {key: "set" if os.environ.get(key) else "unset" for key in env_keys}


_T = t.TypeVar("_T", bound=t.Union[ServiceResource, BaseClient])
_R = t.TypeVar("_R", bound=ServiceResource)
_C = t.TypeVar("_C", bound=BaseClient)
Expand Down Expand Up @@ -113,8 +148,34 @@ def __init__(
self.aws_profile = config.get("aws_profile")
self.aws_default_region = config.get("aws_default_region")

self.aws_endpoint_url = config.get("aws_endpoint_url")
self.aws_endpoint_url = _config_or_env(
config, "aws_endpoint_url", "TAP_DYNAMODB_AWS_ENDPOINT_URL"
)
self.aws_assume_role_arn = config.get("aws_assume_role_arn")
self._log_auth_env(config)

def _log_auth_env(self, config: dict) -> None:
self.logger.info(
"Incoming TAP_DYNAMODB auth env vars: %s",
_env_presence(_TAP_DYNAMODB_AUTH_ENV_KEYS),
)
self.logger.info(
"Incoming AWS auth env vars (used when use_aws_env_vars=true): %s",
_env_presence(_AWS_AUTH_ENV_KEYS),
)
self.logger.info(
"Resolved AWS auth settings: use_aws_env_vars=%s, region=%s, "
"endpoint_url=%s, profile=%s, assume_role_arn=%s, "
"access_key_id=%s, secret_access_key=%s, session_token=%s",
config.get("use_aws_env_vars"),
self.aws_default_region,
self.aws_endpoint_url,
self.aws_profile,
self.aws_assume_role_arn,
"set" if self.aws_access_key_id else "unset",
"set" if self.aws_secret_access_key else "unset",
"set" if self.aws_session_token else "unset",
)

@property
def config(self) -> dict:
Expand Down
66 changes: 65 additions & 1 deletion tests/test_boto_connector.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
from unittest.mock import patch
import os
from unittest.mock import MagicMock, patch

from moto import mock_aws

Expand Down Expand Up @@ -124,3 +125,66 @@ def test_get_resource():
)
session = auth.get_session()
auth.get_resource(session, "dynamodb")


@mock_aws
def test_get_client_without_endpoint_url():
auth = AWSBotoConnector({}, "dynamodb")
session = MagicMock()
session.client = MagicMock(return_value="mock_client")

client = auth.get_client(session, "dynamodb")

session.client.assert_called_once_with("dynamodb")
assert client == "mock_client"


@mock_aws
@patch.dict(os.environ, {"TAP_DYNAMODB_AWS_ENDPOINT_URL": "http://localhost:4566"})
def test_get_client_with_endpoint_url_from_env():
auth = AWSBotoConnector({}, "dynamodb")
session = MagicMock()
session.client = MagicMock(return_value="mock_client")

client = auth.get_client(session, "dynamodb")

session.client.assert_called_once_with(
"dynamodb",
endpoint_url="http://localhost:4566",
)
assert client == "mock_client"


@mock_aws
@patch.dict(os.environ, {"TAP_DYNAMODB_AWS_ENDPOINT_URL": "http://localhost:4566"})
def test_assume_role_sts_client_uses_endpoint_url():
auth = AWSBotoConnector(
{
"aws_access_key_id": "foo",
"aws_secret_access_key": "bar",
"aws_default_region": "baz",
"aws_assume_role_arn": "arn:aws:iam::123456778910:role/my-role-name",
},
"dynamodb",
)
session = MagicMock()
sts_client = MagicMock()
sts_client.assume_role.return_value = {
"Credentials": {
"AccessKeyId": "assumed-key",
"SecretAccessKey": "assumed-secret",
"SessionToken": "assumed-token",
}
}
session.client = MagicMock(return_value=sts_client)

with patch(
"tap_dynamodb.connectors.aws_boto_connector.boto3.Session",
return_value=session,
):
auth.get_session()

session.client.assert_called_once_with(
"sts",
endpoint_url="http://localhost:4566",
)