Skip to content

Commit 65829fd

Browse files
authored
EC2: improve UserData handling (#9510)
1 parent bf84647 commit 65829fd

11 files changed

Lines changed: 100 additions & 23 deletions

File tree

moto/autoscaling/responses.py

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
11
from moto.core.common_types import TYPE_RESPONSE
22
from moto.core.responses import ActionResult, BaseResponse, EmptyResult
3-
from moto.core.types import Base64EncodedString
3+
from moto.ec2.utils import parse_user_data
44
from moto.utilities.aws_headers import amz_crc32
55

66
from .models import AutoScalingBackend, autoscaling_backends
@@ -21,9 +21,7 @@ def call_action(self) -> TYPE_RESPONSE:
2121

2222
def create_launch_configuration(self) -> ActionResult:
2323
params = self._get_params()
24-
user_data = params.get("UserData")
25-
if user_data is not None:
26-
user_data = Base64EncodedString(user_data)
24+
user_data = parse_user_data(params.get("UserData"))
2725
self.autoscaling_backend.create_launch_configuration(
2826
name=params.get("LaunchConfigurationName"), # type: ignore[arg-type]
2927
image_id=params.get("ImageId"), # type: ignore[arg-type]

moto/ec2/exceptions.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -913,3 +913,11 @@ def __init__(self) -> None:
913913
"AuthFailure",
914914
"Unauthorized attempt to access restricted resource",
915915
)
916+
917+
918+
class InvalidUserDataError(EC2ClientError):
919+
def __init__(self, message: str):
920+
super().__init__(
921+
"InvalidUserData.Malformed",
922+
message,
923+
)

moto/ec2/models/launch_templates.py

Lines changed: 7 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@
1313
from ..utils import (
1414
convert_tag_spec,
1515
generic_filter,
16+
parse_user_data,
1617
random_launch_template_id,
1718
random_launch_template_name,
1819
utc_date_and_time,
@@ -51,9 +52,10 @@ def security_groups(self) -> list[str]:
5152

5253
@property
5354
def user_data(self) -> Optional[Base64EncodedString]:
54-
user_data = self.data.get("UserData", None)
55-
if user_data is not None:
56-
user_data = Base64EncodedString(user_data)
55+
user_data = self.data.get("UserData")
56+
# UserData can be specified via multiple services/api endpoints,
57+
# so we make an assertion here that it's in the format we expect.
58+
assert user_data is None or isinstance(user_data, Base64EncodedString)
5759
return user_data
5860

5961

@@ -143,6 +145,7 @@ def create_from_cloudformation_json( # type: ignore[misc]
143145
properties = cloudformation_json["Properties"]
144146
name = properties.get("LaunchTemplateName")
145147
data = properties.get("LaunchTemplateData")
148+
data["UserData"] = parse_user_data(data.get("UserData"))
146149
description = properties.get("VersionDescription")
147150
tag_spec = convert_tag_spec(
148151
properties.get("TagSpecifications", {}), tag_key="Tags"
@@ -173,6 +176,7 @@ def update_from_cloudformation_json( # type: ignore[misc]
173176
properties = cloudformation_json["Properties"]
174177

175178
data = properties.get("LaunchTemplateData")
179+
data["UserData"] = parse_user_data(data.get("UserData"))
176180
description = properties.get("VersionDescription")
177181

178182
launch_template = backend.get_launch_template(original_resource.id)

moto/ec2/models/spot_requests.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
from typing import TYPE_CHECKING, Any, Optional
44

55
from moto.core.common_models import BaseModel, CloudFormationModel
6+
from moto.core.types import Base64EncodedString
67
from moto.ec2.exceptions import InvalidParameterValueErrorTagSpotFleetRequest
78

89
if TYPE_CHECKING:
@@ -56,7 +57,7 @@ def __init__(
5657
availability_zone_group: Optional[str],
5758
key_name: str,
5859
security_groups: list[str],
59-
user_data: dict[str, Any],
60+
user_data: Optional[Base64EncodedString],
6061
instance_type: str,
6162
placement: Optional[str],
6263
kernel_id: Optional[str],
@@ -411,7 +412,7 @@ def request_spot_instances(
411412
availability_zone_group: Optional[str],
412413
key_name: str,
413414
security_groups: list[str],
414-
user_data: dict[str, Any],
415+
user_data: Optional[Base64EncodedString],
415416
instance_type: str,
416417
placement: Optional[str],
417418
kernel_id: Optional[str],

moto/ec2/responses/instances.py

Lines changed: 3 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -2,14 +2,13 @@
22
from typing import Any
33

44
from moto.core.responses import ActionResult, EmptyResult
5-
from moto.core.types import Base64EncodedString
65
from moto.core.utils import camelcase_to_underscores
76
from moto.ec2.exceptions import (
87
InvalidParameterCombination,
98
InvalidRequest,
109
MissingParameterError,
1110
)
12-
from moto.ec2.utils import filter_iam_instance_profiles
11+
from moto.ec2.utils import filter_iam_instance_profiles, parse_user_data
1312

1413
from ._base_response import EC2BaseResponse
1514

@@ -49,9 +48,7 @@ def describe_instances(self) -> ActionResult:
4948
def run_instances(self) -> ActionResult:
5049
min_count = int(self._get_param("MinCount", if_none="1"))
5150
image_id = self._get_param("ImageId")
52-
user_data = self._get_param("UserData")
53-
if user_data is not None:
54-
user_data = Base64EncodedString(user_data)
51+
user_data = parse_user_data(self._get_param("UserData"))
5552
security_group_names = self._get_param("SecurityGroups", [])
5653
kwargs = {
5754
"instance_type": self._get_param("InstanceType", "m1.small"),
@@ -340,7 +337,7 @@ def _dot_value_instance_attribute_handler(self) -> bool:
340337
attr_name = camelcase_to_underscores(attribute)
341338
attr_value = self._get_param(f"{attribute}.Value")
342339
if attribute == "UserData" and attr_value:
343-
attr_value = Base64EncodedString.from_encoded_bytes(attr_value)
340+
attr_value = parse_user_data(attr_value)
344341
self.ec2_backend.modify_instance_attribute(
345342
instance_id, attr_name, attr_value
346343
)

moto/ec2/responses/launch_templates.py

Lines changed: 7 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44

55
from moto.core.types import Base64EncodedString
66
from moto.ec2.exceptions import FilterNotImplementedError
7+
from moto.ec2.utils import parse_user_data
78
from moto.moto_api._internal import mock_random
89

910
from ._base_response import EC2BaseResponse
@@ -20,6 +21,8 @@ def xml_root(name: str) -> ElementTree.Element:
2021

2122

2223
def xml_serialize(tree: ElementTree.Element, key: str, value: Any) -> None:
24+
if value is None:
25+
return
2326
name = key[0].lower() + key[1:]
2427
if isinstance(value, list):
2528
if name[-1] == "s":
@@ -29,7 +32,7 @@ def xml_serialize(tree: ElementTree.Element, key: str, value: Any) -> None:
2932

3033
node = ElementTree.SubElement(tree, name)
3134

32-
if isinstance(value, (str, int, float, str)):
35+
if isinstance(value, (str, int, float, str, Base64EncodedString)):
3336
node.text = str(value)
3437
elif isinstance(value, dict):
3538
for dictkey, dictvalue in value.items():
@@ -58,10 +61,9 @@ def create_launch_template(self) -> str:
5861
tag_spec = self._parse_tag_specification()
5962

6063
parsed_template_data = self._get_param("LaunchTemplateData", {})
61-
if parsed_template_data.get("UserData"):
62-
parsed_template_data["UserData"] = Base64EncodedString(
63-
parsed_template_data["UserData"]
64-
)
64+
parsed_template_data["UserData"] = parse_user_data(
65+
self._get_param("LaunchTemplateData.UserData")
66+
)
6567
self.error_on_dryrun()
6668

6769
if tag_spec:

moto/ec2/responses/spot_instances.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,4 @@
1+
from ..utils import parse_user_data
12
from ._base_response import EC2BaseResponse
23

34

@@ -59,7 +60,7 @@ def request_spot_instances(self) -> str:
5960
availability_zone_group = self._get_param("AvailabilityZoneGroup")
6061
key_name = self._get_param("LaunchSpecification.KeyName")
6162
security_groups = self._get_param("LaunchSpecification.SecurityGroups", [])
62-
user_data = self._get_param("LaunchSpecification.UserData")
63+
user_data = parse_user_data(self._get_param("LaunchSpecification.UserData"))
6364
instance_type = self._get_param("LaunchSpecification.InstanceType", "m1.small")
6465
placement = self._get_param("LaunchSpecification.Placement.AvailabilityZone")
6566
kernel_id = self._get_param("LaunchSpecification.KernelId")

moto/ec2/utils.py

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,9 @@
2222
)
2323
from cryptography.hazmat.primitives.asymmetric.rsa import RSAPublicKey
2424

25+
from moto.core.types import Base64EncodedString
2526
from moto.core.utils import utcnow
27+
from moto.ec2.exceptions import InvalidUserDataError
2628
from moto.iam import iam_backends
2729
from moto.moto_api._internal import mock_random as random
2830
from moto.utilities.utils import md5_hash
@@ -931,3 +933,16 @@ def convert_tag_spec(
931933
{tag["Key"]: tag["Value"] for tag in tag_spec[tag_key]}
932934
)
933935
return tags
936+
937+
938+
def parse_user_data(value: Any) -> Optional[Base64EncodedString]:
939+
if value is None:
940+
return None
941+
try:
942+
if isinstance(value, bytes):
943+
user_data = Base64EncodedString.from_encoded_bytes(value)
944+
else:
945+
user_data = Base64EncodedString(value)
946+
except ValueError:
947+
raise InvalidUserDataError("Invalid BASE64 encoding of user data.")
948+
return user_data

tests/test_autoscaling/test_autoscaling_groups.py

Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44

55
from moto import mock_aws
66
from moto.core import DEFAULT_ACCOUNT_ID as ACCOUNT_ID
7+
from moto.core.types import Base64EncodedString
78
from tests import EXAMPLE_AMI_ID
89

910
from .utils import setup_networking
@@ -363,3 +364,35 @@ def test_launch_template_with_tags():
363364
tags = instances["Reservations"][0]["Instances"][0]["Tags"]
364365
assert {"Value": "TestTagValue1", "Key": "TestTagKey1"} in tags
365366
assert {"Key": "from_lt", "Value": "val"} in tags
367+
368+
369+
@mock_aws
370+
def test_launch_template_with_user_data():
371+
mocked_networking = setup_networking()
372+
user_data = Base64EncodedString.from_raw_string("test user data")
373+
ec2_client = boto3.client("ec2", region_name="us-east-1")
374+
template = ec2_client.create_launch_template(
375+
LaunchTemplateName="test_launch_template",
376+
LaunchTemplateData={
377+
"ImageId": EXAMPLE_AMI_ID,
378+
"InstanceType": "t2.micro",
379+
"UserData": str(user_data),
380+
},
381+
)["LaunchTemplate"]
382+
as_client = boto3.client("autoscaling", region_name="us-east-1")
383+
as_client.create_auto_scaling_group(
384+
AutoScalingGroupName="myasgroup",
385+
MinSize=2,
386+
MaxSize=3,
387+
LaunchTemplate={"LaunchTemplateId": template["LaunchTemplateId"]},
388+
VPCZoneIdentifier=mocked_networking["subnet1"],
389+
)
390+
resp = ec2_client.describe_instances(
391+
Filters=[{"Name": "tag:aws:autoscaling:groupName", "Values": ["myasgroup"]}]
392+
)
393+
instances = resp["Reservations"][0]["Instances"]
394+
for instance in instances:
395+
attr_resp = ec2_client.describe_instance_attribute(
396+
InstanceId=instance["InstanceId"], Attribute="userData"
397+
)
398+
assert attr_resp["UserData"]["Value"] == str(user_data)

tests/test_ec2/test_launch_templates.py

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -763,3 +763,20 @@ def test_modify_launch_template_by_name():
763763
assert (
764764
version["LaunchTemplateId"] == create_resp["LaunchTemplate"]["LaunchTemplateId"]
765765
)
766+
767+
768+
@mock_aws
769+
def test_create_launch_template_with_non_base64_encoded_user_data_fails():
770+
client = boto3.client("ec2", region_name="us-east-1")
771+
with pytest.raises(ClientError) as exc:
772+
client.create_launch_template(
773+
LaunchTemplateName="test-template",
774+
LaunchTemplateData={
775+
"ImageId": "ami-12345678",
776+
"InstanceType": "t2.nano",
777+
"UserData": "not base64 encoded",
778+
},
779+
)
780+
error = exc.value.response["Error"]
781+
assert error["Code"] == "InvalidUserData.Malformed"
782+
assert error["Message"] == "Invalid BASE64 encoding of user data."

0 commit comments

Comments
 (0)