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
3 changes: 2 additions & 1 deletion samcli/commands/deploy/command.py
Original file line number Diff line number Diff line change
Expand Up @@ -341,6 +341,7 @@ def do_cli(
config_file=config_file,
disable_rollback=disable_rollback,
language_extensions_enabled=language_extensions_enabled,
role_arn=role_arn,
)
guided_context.run()
else:
Expand All @@ -362,7 +363,7 @@ def do_cli(
# after we figure out how to enable resolve-images-repos in package
if resolve_image_repos:
image_repositories = sync_ecr_stack(
template_file, stack_name, region, s3_bucket, s3_prefix, image_repositories
template_file, stack_name, region, s3_bucket, s3_prefix, image_repositories, role_arn
)
with osutils.tempfile_platform_independent() as output_template_file:
if guided:
Expand Down
12 changes: 10 additions & 2 deletions samcli/commands/deploy/guided_context.py
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,7 @@ def __init__(
config_file=None,
disable_rollback=None,
language_extensions_enabled: bool = False,
role_arn: Optional[str] = None,
):
self.template_file = template_file
self.stack_name = stack_name
Expand Down Expand Up @@ -97,6 +98,7 @@ def __init__(
self.function_provider: Optional[SamFunctionProvider] = None
self.disable_rollback = disable_rollback
self._language_extensions_enabled = language_extensions_enabled
self.role_arn = role_arn

@property
def guided_capabilities(self):
Expand Down Expand Up @@ -189,7 +191,13 @@ def guided_prompts(self, parameter_override_keys):

image_repositories = (
sync_ecr_stack(
self.template_file, stack_name, region, managed_s3_bucket, self.s3_prefix, self.image_repositories
self.template_file,
stack_name,
region,
managed_s3_bucket,
self.s3_prefix,
self.image_repositories,
self.role_arn,
)
if self.resolve_image_repositories
else self.prompt_image_repository(
Expand Down Expand Up @@ -359,7 +367,7 @@ def prompt_image_repository(
if repo_full_path:
updated_repositories[repo_full_path] = image_repo_uri
self.function_provider = SamFunctionProvider(stacks, ignore_code_extraction_warnings=True)
manager = CompanionStackManager(stack_name, region, s3_bucket, s3_prefix)
manager = CompanionStackManager(stack_name, region, s3_bucket, s3_prefix, self.role_arn)

function_logical_ids = [
function.full_path for function in self.function_provider.get_all() if function.packagetype == IMAGE
Expand Down
45 changes: 38 additions & 7 deletions samcli/lib/bootstrap/companion_stack/companion_stack_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
"""

import logging
from typing import Dict, List, Optional
from typing import Any, Dict, List, Optional

import boto3
from botocore.config import Config
Expand Down Expand Up @@ -38,17 +38,19 @@ class CompanionStackManager:
_delete_stack_waiter_config: WaiterConfigTypeDef
_s3_bucket: str
_s3_prefix: str
_role_arn: Optional[str]
_cfn_client: CloudFormationClient
_s3_client: S3Client

def __init__(self, stack_name, region, s3_bucket, s3_prefix):
def __init__(self, stack_name, region, s3_bucket, s3_prefix, role_arn: Optional[str] = None):
self._companion_stack = CompanionStack(stack_name)
self._builder = CompanionStackBuilder(self._companion_stack)
self._boto_config = Config(region_name=region if region else None)
self._update_stack_waiter_config = {"Delay": 10, "MaxAttempts": 120}
self._delete_stack_waiter_config = {"Delay": 10, "MaxAttempts": 120}
self._s3_bucket = s3_bucket
self._s3_prefix = s3_prefix
self._role_arn = role_arn
try:
self._cfn_client = boto3.client("cloudformation", config=self._boto_config)
self._ecr_client = boto3.client("ecr", config=self._boto_config)
Expand Down Expand Up @@ -116,16 +118,24 @@ def update_companion_stack(self) -> None:

template_url = s3_uploader.to_path_style_s3_url(parts["Key"], parts.get("Version", None))

extra_args: Dict[str, Any] = {"RoleARN": self._role_arn} if self._role_arn else {}

exists = self.does_companion_stack_exist()
if exists:
self._cfn_client.update_stack(
StackName=stack_name, TemplateURL=template_url, Capabilities=["CAPABILITY_AUTO_EXPAND"]
StackName=stack_name,
TemplateURL=template_url,
Capabilities=["CAPABILITY_AUTO_EXPAND"],
**extra_args,
)
update_waiter = self._cfn_client.get_waiter("stack_update_complete")
update_waiter.wait(StackName=stack_name, WaiterConfig=self._update_stack_waiter_config)
else:
self._cfn_client.create_stack(
StackName=stack_name, TemplateURL=template_url, Capabilities=["CAPABILITY_AUTO_EXPAND"]
StackName=stack_name,
TemplateURL=template_url,
Capabilities=["CAPABILITY_AUTO_EXPAND"],
**extra_args,
)
create_waiter = self._cfn_client.get_waiter("stack_create_complete")
create_waiter.wait(StackName=stack_name, WaiterConfig=self._update_stack_waiter_config)
Expand All @@ -135,8 +145,9 @@ def _delete_companion_stack(self) -> None:
Blocking call to delete the companion stack
"""
stack_name = self._companion_stack.stack_name
extra_args: Dict[str, Any] = {"RoleARN": self._role_arn} if self._role_arn else {}
waiter = self._cfn_client.get_waiter("stack_delete_complete")
self._cfn_client.delete_stack(StackName=stack_name)
self._cfn_client.delete_stack(StackName=stack_name, **extra_args)
waiter.wait(StackName=stack_name, WaiterConfig=self._delete_stack_waiter_config)

def list_deployed_repos(self) -> List[ECRRepo]:
Expand Down Expand Up @@ -189,13 +200,25 @@ def delete_unreferenced_repos(self) -> None:
"""
Blocking call to delete all deployed ECR repos that are unreferenced by a function
If repo does not exist, this will simply skip it.

This always deletes using the caller's own credentials, not role_arn: role_arn is a
CloudFormation execution role passed to create_stack/update_stack/delete_stack, not a
role the CLI itself assumes for direct service calls like ecr:DeleteRepository.
"""
repos = self.get_unreferenced_repos()
for repo in repos:
try:
self._ecr_client.delete_repository(repositoryName=repo.physical_id, force=True)
except self._ecr_client.exceptions.RepositoryNotFoundException:
LOG.debug("Image repo [%s] not found in companion stack. Skipping deletion.", repo.physical_id)
except ClientError as ex:
if ex.response.get("Error", {}).get("Code") == "AccessDeniedException":
raise AWSServiceClientError(
f"Insufficient permissions to delete ECR repo [{repo.physical_id}]. "
"The caller's own credentials need the ecr:DeleteRepository permission; "
"--role-arn only applies to CloudFormation stack operations."
) from ex
raise

def sync_repos(self) -> None:
"""
Expand Down Expand Up @@ -279,7 +302,13 @@ def is_repo_uri(self, repo_uri: Optional[str], function_logical_id: str) -> bool


def sync_ecr_stack(
template_file: str, stack_name: str, region: str, s3_bucket: str, s3_prefix: str, image_repositories: Dict[str, str]
template_file: str,
stack_name: str,
region: str,
s3_bucket: str,
s3_prefix: str,
image_repositories: Dict[str, str],
role_arn: Optional[str] = None,
) -> Dict[str, str]:
"""Blocking call to sync local functions with ECR Companion Stack

Expand All @@ -297,6 +326,8 @@ def sync_ecr_stack(
S3 prefix for the bucket
image_repositories : Dict[str, str]
Mapping between function logical ID and ECR URI
role_arn : Optional[str]
IAM role ARN used when creating/updating the companion stack

Returns
-------
Expand All @@ -305,7 +336,7 @@ def sync_ecr_stack(
for Functions without a repo specified.
"""
image_repositories = image_repositories.copy() if image_repositories else {}
manager = CompanionStackManager(stack_name, region, s3_bucket, s3_prefix)
manager = CompanionStackManager(stack_name, region, s3_bucket, s3_prefix, role_arn)

stacks = SamLocalStackProvider.get_stacks(template_file, language_extensions_enabled=False)[0]
function_provider = SamFunctionProvider(stacks, ignore_code_extraction_warnings=True)
Expand Down
14 changes: 14 additions & 0 deletions tests/unit/commands/deploy/test_command.py
Original file line number Diff line number Diff line change
Expand Up @@ -455,6 +455,16 @@ def test_all_args_guided_use_defaults(
express=self.express,
)

mock_sync_ecr_stack.assert_called_with(
self.template_file,
"sam-app",
"us-east-1",
"managed-s3-bucket",
self.s3_prefix,
None,
self.role_arn,
)

mock_deploy_context.assert_called_with(
template_file=ANY,
stack_name="sam-app",
Expand Down Expand Up @@ -1284,6 +1294,10 @@ def test_all_args_resolve_image_repos(
output="text",
)

mock_sync_ecr_stack.assert_called_with(
self.template_file, self.stack_name, self.region, self.s3_bucket, self.s3_prefix, None, self.role_arn
)

mock_deploy_context.assert_called_with(
template_file=ANY,
stack_name=self.stack_name,
Expand Down
5 changes: 5 additions & 0 deletions tests/unit/commands/deploy/test_guided_context.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@ def setUp(self):
image_repository=None,
image_repositories={"RandomFunction": "image-repo"},
disable_rollback=False,
role_arn="role_arn",
)

self.unreferenced_repo_mock = MagicMock()
Expand Down Expand Up @@ -237,6 +238,10 @@ def test_guided_prompts_check_defaults_public_resources_images(
]
self.assertEqual(expected_click_secho_calls, patched_click_secho.call_args_list)

self.companion_stack_manager_mock.assert_called_once_with(
"sam-app", "region", "managed_s3_stack", self.gc.s3_prefix, "role_arn"
)

@patch("samcli.commands.deploy.guided_context.get_resource_full_path_by_id")
@patch("samcli.commands.deploy.guided_context.prompt")
@patch("samcli.commands.deploy.guided_context.confirm")
Expand Down
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
from botocore.exceptions import ClientError
from samcli.commands.exceptions import AWSServiceClientError
from samcli.lib.bootstrap.companion_stack.companion_stack_manager import CompanionStackManager, sync_ecr_stack
from unittest import TestCase
from unittest.mock import ANY, MagicMock, Mock, patch
Expand Down Expand Up @@ -70,6 +71,27 @@ def test_create_companion_stack(
self.cfn_client.get_waiter.assert_called_once_with("stack_create_complete")
cfn_waiter.wait.assert_called_once_with(StackName=self.companion_stack_name, WaiterConfig=ANY)

@patch("samcli.lib.bootstrap.companion_stack.companion_stack_manager.mktempfile")
@patch("samcli.lib.bootstrap.companion_stack.companion_stack_manager.S3Uploader")
@patch("samcli.lib.bootstrap.companion_stack.companion_stack_manager.parse_s3_url")
def test_create_companion_stack_with_role_arn(
self,
parse_s3_url_mock,
s3_uploader_mock,
mktempfile_mock,
):
cfn_waiter = Mock()
self.cfn_client.get_waiter.return_value = cfn_waiter

self.manager._role_arn = "role-arn"
self.manager.does_companion_stack_exist = lambda: False

self.manager.update_companion_stack()

self.cfn_client.create_stack.assert_called_once_with(
StackName=self.companion_stack_name, TemplateURL=ANY, Capabilities=ANY, RoleARN="role-arn"
)

@patch("samcli.lib.bootstrap.companion_stack.companion_stack_manager.mktempfile")
@patch("samcli.lib.bootstrap.companion_stack.companion_stack_manager.S3Uploader")
@patch("samcli.lib.bootstrap.companion_stack.companion_stack_manager.parse_s3_url")
Expand All @@ -94,6 +116,27 @@ def test_update_companion_stack(
self.cfn_client.get_waiter.assert_called_once_with("stack_update_complete")
cfn_waiter.wait.assert_called_once_with(StackName=self.companion_stack_name, WaiterConfig=ANY)

@patch("samcli.lib.bootstrap.companion_stack.companion_stack_manager.mktempfile")
@patch("samcli.lib.bootstrap.companion_stack.companion_stack_manager.S3Uploader")
@patch("samcli.lib.bootstrap.companion_stack.companion_stack_manager.parse_s3_url")
def test_update_companion_stack_with_role_arn(
self,
parse_s3_url_mock,
s3_uploader_mock,
mktempfile_mock,
):
cfn_waiter = Mock()
self.cfn_client.get_waiter.return_value = cfn_waiter

self.manager._role_arn = "role-arn"
self.manager.does_companion_stack_exist = lambda: True

self.manager.update_companion_stack()

self.cfn_client.update_stack.assert_called_once_with(
StackName=self.companion_stack_name, TemplateURL=ANY, Capabilities=ANY, RoleARN="role-arn"
)

def test_delete_companion_stack(self):
cfn_waiter = Mock()
self.cfn_client.get_waiter.return_value = cfn_waiter
Expand All @@ -104,6 +147,17 @@ def test_delete_companion_stack(self):
self.cfn_client.get_waiter.assert_called_once_with("stack_delete_complete")
cfn_waiter.wait.assert_called_once_with(StackName=self.companion_stack_name, WaiterConfig=ANY)

def test_delete_companion_stack_with_role_arn(self):
cfn_waiter = Mock()
self.cfn_client.get_waiter.return_value = cfn_waiter

self.manager._role_arn = "role-arn"
self.manager._delete_companion_stack()

self.cfn_client.delete_stack.assert_called_once_with(
StackName=self.companion_stack_name, RoleARN="role-arn"
)

@patch("samcli.lib.bootstrap.companion_stack.companion_stack_manager.ECRRepo")
@patch("samcli.lib.bootstrap.companion_stack.companion_stack_manager.boto3.resource")
def test_list_deployed_repos(self, boto3_resource_mock, ecr_repo_mock):
Expand Down Expand Up @@ -212,6 +266,32 @@ def test_delete_unreferenced_repos(self):
self.ecr_client.delete_repository.assert_any_call(repositoryName=repo_a_id, force=True)
self.ecr_client.delete_repository.assert_any_call(repositoryName=repo_b_id, force=True)

def test_delete_unreferenced_repos_access_denied(self):
repo_a = Mock()
repo_a.physical_id = "ECRRepoA"

self.ecr_client.exceptions.RepositoryNotFoundException = type("RepositoryNotFoundException", (Exception,), {})
error = ClientError({"Error": {"Code": "AccessDeniedException"}}, "DeleteRepository")
self.ecr_client.delete_repository.side_effect = error

self.manager.get_unreferenced_repos = lambda: [repo_a]

with self.assertRaises(AWSServiceClientError):
self.manager.delete_unreferenced_repos()

def test_delete_unreferenced_repos_other_client_error(self):
repo_a = Mock()
repo_a.physical_id = "ECRRepoA"

self.ecr_client.exceptions.RepositoryNotFoundException = type("RepositoryNotFoundException", (Exception,), {})
error = ClientError({"Error": {"Code": "ThrottlingException"}}, "DeleteRepository")
self.ecr_client.delete_repository.side_effect = error

self.manager.get_unreferenced_repos = lambda: [repo_a]

with self.assertRaises(ClientError):
self.manager.delete_unreferenced_repos()

def test_sync_repos_exists(self):
self.manager.does_companion_stack_exist = lambda: True
self.manager.get_repository_mapping = lambda: {"a": ""}
Expand Down Expand Up @@ -276,8 +356,23 @@ def test_sync_ecr_stack(self, function_provider_mock, stack_provider_mock, manag

result = sync_ecr_stack("template.yaml", "stack-name", "region", "s3-bucket", "s3-prefix", image_repositories)

manager_mock.assert_called_once_with("stack-name", "region", "s3-bucket", "s3-prefix")
manager_mock.assert_called_once_with("stack-name", "region", "s3-bucket", "s3-prefix", None)
function_provider_mock.assert_called_once_with(stacks, ignore_code_extraction_warnings=True)
manager_mock.return_value.sync_repos.assert_called_once_with()

self.assertEqual(result, {"Function1": "uri1", "Function2": "uri2"})

@patch("samcli.lib.bootstrap.companion_stack.companion_stack_manager.CompanionStackManager")
@patch("samcli.lib.bootstrap.companion_stack.companion_stack_manager.SamLocalStackProvider")
@patch("samcli.lib.bootstrap.companion_stack.companion_stack_manager.SamFunctionProvider")
def test_sync_ecr_stack_with_role_arn(self, function_provider_mock, stack_provider_mock, manager_mock):
image_repositories = {"Function1": "uri1"}
stacks = MagicMock()
stack_provider_mock.get_stacks.return_value = (stacks, None)
manager_mock.return_value.get_repository_mapping.return_value = {"Function2": "uri2"}

sync_ecr_stack(
"template.yaml", "stack-name", "region", "s3-bucket", "s3-prefix", image_repositories, "role-arn"
)

manager_mock.assert_called_once_with("stack-name", "region", "s3-bucket", "s3-prefix", "role-arn")