Skip to content
Merged
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
99 changes: 79 additions & 20 deletions src/cloudformation_cli_python_lib/resource.py
Original file line number Diff line number Diff line change
@@ -1,28 +1,38 @@
import json
import logging
from functools import wraps
from time import sleep
from typing import Any, Callable, MutableMapping, Optional, Tuple, Type, Union

# boto3 doesn't have stub files
import boto3 # type: ignore

from .boto3_proxy import SessionProxy, _get_boto_session
from .exceptions import InvalidRequest, _HandlerError
from .exceptions import InternalFailure, InvalidRequest, _HandlerError
from .interface import (
Action,
BaseResourceHandlerRequest,
HandlerErrorCode,
OperationStatus,
ProgressEvent,
)
from .log_delivery import ProviderLogHandler
from .scheduler import CloudWatchScheduler
from .utils import (
BaseResourceModel,
Credentials,
HandlerRequest,
KitchenSinkEncoder,
LambdaContext,
TestEvent,
UnmodelledRequest,
)

LOG = logging.getLogger(__name__)

MUTATING_ACTIONS = (Action.CREATE, Action.UPDATE, Action.DELETE)
INVOCATION_TIMEOUT_MS = 60000

HandlerSignature = Callable[
[Optional[SessionProxy], Any, MutableMapping[str, Any]], ProgressEvent
]
Expand Down Expand Up @@ -60,6 +70,43 @@ def _add_handler(f: HandlerSignature) -> HandlerSignature:

return _add_handler

@staticmethod
def schedule_reinvocation(
handler_request: HandlerRequest,
handler_response: ProgressEvent,
context: LambdaContext,
session: boto3.Session,
) -> bool:
if handler_response.status != OperationStatus.IN_PROGRESS:
return False
# modify requestContext dict in-place, so that invoke count is bumped on local
# reinvoke too
reinvoke_context = handler_request.requestContext
reinvoke_context["invocation"] = reinvoke_context.get("invocation", 0) + 1
callback_delay_s = handler_response.callbackDelaySeconds
remaining_ms = context.get_remaining_time_in_millis()

# when a handler requests a sub-minute callback delay, and if the lambda
# invocation has enough runtime (with 20% buffer), we can re-run the handler
# locally otherwise we re-invoke through CloudWatchEvents
needed_ms_remaining = callback_delay_s * 1200 + INVOCATION_TIMEOUT_MS
Comment thread
jaymccon marked this conversation as resolved.
if callback_delay_s < 60 and remaining_ms > needed_ms_remaining:
LOG.info(
"Scheduling re-invoke locally after %s seconds, with Context {%s}",
callback_delay_s,
reinvoke_context,
)
sleep(callback_delay_s)
return True
LOG.info("Scheduling re-invoke with Context {%s}", reinvoke_context)
callback_delay_min = int(callback_delay_s / 60)
CloudWatchScheduler(boto3_session=session).reschedule_after_minutes(
function_arn=context.invoked_function_arn,
minutes_from_now=callback_delay_min,
handler_request=handler_request,
)
return False

def _invoke_handler(
self,
session: Optional[SessionProxy],
Expand All @@ -73,8 +120,12 @@ def _invoke_handler(
return ProgressEvent.failed(
HandlerErrorCode.InternalFailure, f"No handler for {action}"
)

return handler(session, request, callback_context)
progress = handler(session, request, callback_context)
is_in_progress = progress.status == OperationStatus.IN_PROGRESS
is_mutable = action in MUTATING_ACTIONS
if is_in_progress and not is_mutable:
raise InternalFailure("READ and LIST handlers must return synchronously.")
return progress

def _parse_test_request(
self, event_data: MutableMapping[str, Any]
Expand Down Expand Up @@ -121,52 +172,60 @@ def _parse_request(
self, event_data: MutableMapping[str, Any]
) -> Tuple[
Optional[SessionProxy],
boto3.Session,
BaseResourceHandlerRequest,
Action,
MutableMapping[str, Any],
HandlerRequest,
]:
try:
event = HandlerRequest.deserialize(event_data)
creds = event.requestData.callerCredentials
caller_creds = event.requestData.callerCredentials
platform_creds = event.requestData.platformCredentials
request: BaseResourceHandlerRequest = UnmodelledRequest(
clientRequestToken=event.bearerToken,
desiredResourceState=event.requestData.resourceProperties,
previousResourceState=event.requestData.previousResourceProperties,
logicalResourceIdentifier=event.requestData.logicalResourceId,
).to_modelled(self._model_cls)

session = _get_boto_session(creds, event.region)
caller_sess = _get_boto_session(caller_creds, event.region)
# No need to proxy as platform creds are required in the request
platform_sess = boto3.Session(
aws_access_key_id=platform_creds.accessKeyId,
aws_secret_access_key=platform_creds.secretAccessKey,
aws_session_token=platform_creds.sessionToken,
)
action = Action[event.action]
callback_context = event.requestContext.get("callbackContext", {})
Credentials(**event_data["requestData"]["platformCredentials"])
except Exception as e: # pylint: disable=broad-except
LOG.exception("Invalid request")
raise InvalidRequest(f"{e} ({type(e).__name__})") from e
return session, request, action, callback_context
return caller_sess, platform_sess, request, action, callback_context, event

@_ensure_serialize
def __call__(
self, event_data: MutableMapping[str, Any], _context: Any
self, event_data: MutableMapping[str, Any], context: LambdaContext
) -> MutableMapping[str, Any]:
try:
ProviderLogHandler.setup(event_data)
parsed = self._parse_request(event_data)
session, request, action, callback_context = parsed
progress_event = self._invoke_handler(
session, request, action, callback_context
)
caller_sess, platform_sess, request, action, callback, event = parsed
invoke = True
while invoke:
progress = self._invoke_handler(caller_sess, request, action, callback)
invoke = self.schedule_reinvocation(
event, progress, context, platform_sess
)
except _HandlerError as e:
LOG.exception("Handler error", exc_info=True)
progress_event = e.to_progress_event()
progress = e.to_progress_event()
except Exception as e: # pylint: disable=broad-except
LOG.exception("Exception caught", exc_info=True)
progress_event = ProgressEvent.failed(
HandlerErrorCode.InternalFailure, str(e)
)
progress = ProgressEvent.failed(HandlerErrorCode.InternalFailure, str(e))
except BaseException as e: # pylint: disable=broad-except
LOG.critical("Base exception caught (this is usually bad)", exc_info=True)
progress_event = ProgressEvent.failed(
HandlerErrorCode.InternalFailure, str(e)
)
return progress_event._serialize( # pylint: disable=protected-access
progress = ProgressEvent.failed(HandlerErrorCode.InternalFailure, str(e))
return progress._serialize( # pylint: disable=protected-access
to_response=True, bearer_token=event_data.get("bearerToken")
)
7 changes: 6 additions & 1 deletion src/cloudformation_cli_python_lib/utils.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
import json
from dataclasses import dataclass, field
from datetime import date, datetime, time
from typing import Any, Mapping, MutableMapping, Optional, Type
from typing import Any, Callable, Mapping, MutableMapping, Optional, Type

from .interface import Action, BaseResourceHandlerRequest, BaseResourceModel

Expand Down Expand Up @@ -118,3 +118,8 @@ def to_modelled(
logicalResourceIdentifier=self.logicalResourceIdentifier,
nextToken=self.nextToken,
)


class LambdaContext:
get_remaining_time_in_millis: Callable[["LambdaContext"], int]
invoked_function_arn: str
105 changes: 97 additions & 8 deletions tests/lib/resource_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -149,12 +149,16 @@ def test__parse_request_valid_request():

with patch(
"cloudformation_cli_python_lib.resource._get_boto_session"
) as mock_session:
) as mock_caller_session, patch(
"cloudformation_cli_python_lib.resource.boto3.Session"
) as mock_platform_session:
ret = resource._parse_request(ENTRYPOINT_PAYLOAD)
session, request, action, callback_context = ret
caller_sess, platform_sess, request, action, callback_context, _event = ret

mock_session.assert_called_once()
assert session is mock_session.return_value
mock_caller_session.assert_called_once()
assert caller_sess is mock_caller_session.return_value
mock_platform_session.assert_called_once()
assert platform_sess is mock_platform_session.return_value

mock_model._deserialize.assert_has_calls(
[call(sentinel.state_in1), call(sentinel.state_in2)]
Expand Down Expand Up @@ -218,17 +222,29 @@ def test__invoke_handler_not_found(resource):


def test__invoke_handler_was_found(resource):
mock_handler = resource.handler(Action.CREATE)(Mock(return_value=sentinel.response))
progress_event = ProgressEvent(status=OperationStatus.IN_PROGRESS)
mock_handler = resource.handler(Action.CREATE)(Mock(return_value=progress_event))

resp = resource._invoke_handler(
sentinel.session, sentinel.request, Action.CREATE, sentinel.context
)
assert resp is sentinel.response
assert resp is progress_event
mock_handler.assert_called_once_with(
sentinel.session, sentinel.request, sentinel.context
)


@pytest.mark.parametrize("action", [Action.LIST, Action.READ])
def test__invoke_handler_non_mutating_must_be_synchronous(resource, action):
progress_event = ProgressEvent(status=OperationStatus.IN_PROGRESS)
resource.handler(action)(Mock(return_value=progress_event))
with pytest.raises(Exception) as excinfo:
resource._invoke_handler(
sentinel.session, sentinel.request, action, sentinel.context
)
assert excinfo.value.args[0] == "READ and LIST handlers must return synchronously."


@pytest.mark.parametrize("event,messages", [({}, ("missing", "credentials"))])
def test__parse_test_request_invalid_request(resource, event, messages):
with pytest.raises(InvalidRequest) as excinfo:
Expand Down Expand Up @@ -301,7 +317,8 @@ def test_test_entrypoint_success():
mock_model._deserialize.side_effect = [None, None]

resource = Resource(mock_model)
mock_handler = resource.handler(Action.CREATE)(Mock(return_value=sentinel.response))
progress_event = ProgressEvent(status=OperationStatus.SUCCESS)
mock_handler = resource.handler(Action.CREATE)(Mock(return_value=progress_event))

payload = {
"credentials": {"accessKeyId": "", "secretAccessKey": "", "sessionToken": ""},
Expand All @@ -317,7 +334,79 @@ def test_test_entrypoint_success():
event = resource.test_entrypoint.__wrapped__( # pylint: disable=no-member
resource, payload, None
)
assert event is sentinel.response
assert event is progress_event

mock_model._deserialize.assert_has_calls([call(None), call(None)])
mock_handler.assert_called_once()


def test_schedule_reinvocation_not_in_progress():
progress = ProgressEvent(status=OperationStatus.SUCCESS)
with patch(
"cloudformation_cli_python_lib.resource.boto3.Session", autospec=True
) as mock_session, patch(
"cloudformation_cli_python_lib.resource.CloudWatchScheduler", autospec=True
) as mock_scheduler:
reinvoke = Resource.schedule_reinvocation(
sentinel.request, progress, sentinel.context, sentinel.session
)
assert reinvoke is False
mock_session.assert_not_called()
mock_scheduler.assert_not_called()


def test_schedule_reinvocation_local_callback():
progress = ProgressEvent(status=OperationStatus.IN_PROGRESS, callbackDelaySeconds=5)
mock_request = Mock(
"cloudformation_cli_python_lib.interface.HandlerRequest", autospec=True
)()
mock_request.requestContext = {}
mock_context = Mock(
"cloudformation_cli_python_lib.interface.LambdaContext", autospec=True
)()
mock_context.get_remaining_time_in_millis.return_value = 600000
with patch(
"cloudformation_cli_python_lib.resource.sleep", autospec=True
) as mock_sleep:
reinvoke = Resource.schedule_reinvocation(
mock_request, progress, mock_context, sentinel.session
)
assert reinvoke is True
mock_sleep.assert_called_once_with(5)
assert mock_request.requestContext.get("invocation") == 1


def test_schedule_reinvocation_cloudwatch_callback():

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

there are a lot of mocks here. i vastly prefer some decoupling, such as creating the platform session somewhere else

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

reduced the mock count somewhat, session is created elsewhere, still quite a few mocks left though

progress = ProgressEvent(
status=OperationStatus.IN_PROGRESS, callbackDelaySeconds=60
)
mock_request = Mock(
"cloudformation_cli_python_lib.interface.HandlerRequest", autospec=True
)()
mock_request.requestContext = {}
mock_context = Mock(
"cloudformation_cli_python_lib.interface.LambdaContext", autospec=True
)()
mock_context.get_remaining_time_in_millis.return_value = 6000
mock_context.invoked_function_arn = "arn:aaa:bbb:ccc"
with patch(
"cloudformation_cli_python_lib.resource.CloudWatchScheduler", autospec=True
) as mock_scheduler, patch(
"cloudformation_cli_python_lib.resource.sleep", autospec=True
) as mock_sleep:
reinvoke = Resource.schedule_reinvocation(
mock_request, progress, mock_context, Mock()
)
assert reinvoke is False
mock_scheduler.assert_called_once()
assert mock_scheduler.method_calls[0] == (
"().reschedule_after_minutes",
(),
{
"function_arn": "arn:aaa:bbb:ccc",
"minutes_from_now": 1,
"handler_request": mock_request,
},
)
mock_sleep.assert_not_called()
assert mock_request.requestContext.get("invocation") == 1