diff --git a/src/azure-cli/azure/cli/command_modules/postgresql/_help.py b/src/azure-cli/azure/cli/command_modules/postgresql/_help.py index 10858602eeb..f99101a71a4 100644 --- a/src/azure-cli/azure/cli/command_modules/postgresql/_help.py +++ b/src/azure-cli/azure/cli/command_modules/postgresql/_help.py @@ -342,6 +342,10 @@ --source-server /subscriptions/{sourceSubscriptionId}/resourceGroups/{sourceResourceGroup}/providers/Microsoft.DBforPostgreSQL/flexibleServers/{sourceServer} - name: Restore 'testserver' to current point-in-time as a new server 'testservernew' using Premium SSD v2 Disks by setting storage type to "PremiumV2_LRS" text: az postgres flexible-server restore --resource-group testgroup --name testservernew --source-server testserver --storage-type PremiumV2_LRS + - name: Restore 'testserver' to current point-in-time as a new server 'testservernew' with a different compute size. + text: az postgres flexible-server restore --resource-group testgroup --name testservernew --source-server testserver --sku-name Standard_D4s_v3 + - name: Restore 'testserver' to current point-in-time as a new server 'testservernew' with a different compute tier and compute size. + text: az postgres flexible-server restore --resource-group testgroup --name testservernew --source-server testserver --tier MemoryOptimized --sku-name Standard_E2ds_v4 """ helps['postgres flexible-server maintenance-event'] = """ diff --git a/src/azure-cli/azure/cli/command_modules/postgresql/_params.py b/src/azure-cli/azure/cli/command_modules/postgresql/_params.py index 117daad37b2..8d92f81a614 100644 --- a/src/azure-cli/azure/cli/command_modules/postgresql/_params.py +++ b/src/azure-cli/azure/cli/command_modules/postgresql/_params.py @@ -461,6 +461,8 @@ def _flexible_server_params(command_group): with self.argument_context('{} flexible-server restore'.format(command_group)) as c: c.argument('restore_point_in_time', arg_type=restore_point_in_time_arg_type) c.argument('source_server', arg_type=source_server_arg_type) + c.argument('sku_name', arg_type=sku_name_arg_type) + c.argument('tier', arg_type=tier_arg_type) c.argument('vnet', arg_type=vnet_arg_type) c.argument('subnet', arg_type=subnet_arg_type) c.argument('private_dns_zone_arguments', private_dns_zone_arguments_arg_type) diff --git a/src/azure-cli/azure/cli/command_modules/postgresql/commands/custom_commands.py b/src/azure-cli/azure/cli/command_modules/postgresql/commands/custom_commands.py index 158124d9d54..5220b9892fa 100644 --- a/src/azure-cli/azure/cli/command_modules/postgresql/commands/custom_commands.py +++ b/src/azure-cli/azure/cli/command_modules/postgresql/commands/custom_commands.py @@ -41,6 +41,8 @@ check_resource_group, pg_arguments_validator, pg_byok_validator, + pg_restore_sku_validator, + pg_restore_tier_validator, pg_restore_validator, resolve_private_dns_zone_id, validate_and_format_restore_point_in_time, @@ -328,7 +330,7 @@ def flexible_server_restore(cmd, client, private_dns_zone_arguments=None, geo_redundant_backup=None, byok_identity=None, byok_key=None, backup_byok_identity=None, backup_byok_key=None, federated_client_id=None, backup_federated_client_id=None, - storage_type=None, yes=False): + storage_type=None, sku_name=None, tier=None, yes=False): server_name = server_name.lower() @@ -368,7 +370,23 @@ def flexible_server_restore(cmd, client, federated_client_id=federated_client_id, backup_federated_client_id=backup_federated_client_id) - pg_restore_validator(source_server_object.sku.tier, storage_type=storage_type) + if sku_name or tier: + sku_info = get_postgres_location_capability_info(cmd, location)['sku_info'] + if tier: + pg_restore_tier_validator(tier, source_server_object.sku.tier, sku_info) + else: + tier = source_server_object.sku.tier + if sku_name: + pg_restore_sku_validator(sku_name, sku_info, tier) + else: + sku_name = get_postgres_default_sku(sku_info, tier) + logger.warning('--sku-name was not specified. The restored server will use the default compute ' + 'size \'%s\' for the %s tier.', sku_name, tier) + else: + tier = source_server_object.sku.tier + sku_name = source_server_object.sku.name + + pg_restore_validator(tier, storage_type=storage_type) storage = postgresql_flexibleservers.models.Storage(type=storage_type if source_server_object.storage.type != "PremiumV2_LRS" else None) parameters = postgresql_flexibleservers.models.Server( @@ -377,6 +395,7 @@ def flexible_server_restore(cmd, client, source_server_resource_id=source_server_id, # this should be the source server name, not id create_mode="PointInTimeRestore", availability_zone=zone, + sku=postgresql_flexibleservers.models.Sku(name=sku_name, tier=tier), storage=storage ) @@ -404,6 +423,9 @@ def flexible_server_restore(cmd, client, federated_client_id=federated_client_id, backup_federated_client_id=backup_federated_client_id) + # Let argument validation errors surface as-is instead of being masked as "not found". + except CLIError: + raise except Exception as e: raise ResourceNotFoundError(e) diff --git a/src/azure-cli/azure/cli/command_modules/postgresql/tests/latest/test_postgres_flexible_restore_sku_params.py b/src/azure-cli/azure/cli/command_modules/postgresql/tests/latest/test_postgres_flexible_restore_sku_params.py new file mode 100644 index 00000000000..375cfee2aa6 --- /dev/null +++ b/src/azure-cli/azure/cli/command_modules/postgresql/tests/latest/test_postgres_flexible_restore_sku_params.py @@ -0,0 +1,83 @@ +# -------------------------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. See License.txt in the project root for license information. +# -------------------------------------------------------------------------------------------- + +import inspect +import unittest +from unittest import mock + +from azure.cli.core.azclierror import ValidationError +from azure.cli.core.mock import DummyCli +from azure.cli.command_modules.postgresql import PostgreSQLCommandsLoader +from azure.cli.command_modules.postgresql.commands.custom_commands import flexible_server_restore +from azure.cli.command_modules.postgresql.utils.validators import pg_restore_tier_validator + +RESTORE_COMMAND = 'postgres flexible-server restore' + +# Shape mirrors the 'sku_info' entry returned by get_postgres_location_capability_info. +SKU_INFO = { + 'Burstable': {'skus': {'Standard_B1ms', 'Standard_B2s'}}, + 'GeneralPurpose': {'skus': {'Standard_D2s_v3', 'Standard_D4s_v3'}}, + 'MemoryOptimized': {'skus': {'Standard_E2ds_v4', 'Standard_E4ds_v4'}}, +} + + +def _load_command_and_arguments(command): + cli = DummyCli(commands_loader_cls=PostgreSQLCommandsLoader) + loader = PostgreSQLCommandsLoader(cli) + cli.invocation = mock.MagicMock() + cli.invocation.commands_loader = loader + loader.command_name = command + loader.load_command_table(None) + loader.load_arguments(command) + loader._update_command_definitions() + return loader + + +class RestoreSkuArgumentsTest(unittest.TestCase): + """`az postgres flexible-server restore` must expose --sku-name and --tier.""" + + def test_sku_name_and_tier_are_registered(self): + loader = _load_command_and_arguments(RESTORE_COMMAND) + arguments = loader.argument_registry.arguments.get(RESTORE_COMMAND, {}) + + for dest, option in (('sku_name', '--sku-name'), ('tier', '--tier')): + arg = arguments.get(dest) + self.assertIsNotNone(arg, "'{}' not found in argument registry".format(dest)) + self.assertIn(option, arg.settings.get('options_list')) + + def test_registered_arguments_are_accepted_by_the_custom_command(self): + """An argument registered for a dest the custom command does not accept is silently dropped.""" + loader = _load_command_and_arguments(RESTORE_COMMAND) + accepted = set(inspect.signature(flexible_server_restore).parameters) + registered = set(loader.argument_registry.arguments.get(RESTORE_COMMAND, {})) + + self.assertEqual(registered - accepted, set()) + + +class RestoreTierValidatorTest(unittest.TestCase): + + def test_upgrading_tier_is_allowed(self): + pg_restore_tier_validator('MemoryOptimized', 'GeneralPurpose', SKU_INFO) + + def test_same_tier_is_allowed(self): + pg_restore_tier_validator('GeneralPurpose', 'GeneralPurpose', SKU_INFO) + + def test_downgrading_tier_is_rejected(self): + with self.assertRaises(ValidationError) as context: + pg_restore_tier_validator('GeneralPurpose', 'MemoryOptimized', SKU_INFO) + self.assertIn('must not go below the source server compute tier', str(context.exception)) + + def test_downgrading_to_burstable_is_rejected(self): + with self.assertRaises(ValidationError): + pg_restore_tier_validator('Burstable', 'GeneralPurpose', SKU_INFO) + + def test_unknown_tier_is_rejected(self): + with self.assertRaises(Exception) as context: + pg_restore_tier_validator('NotATier', 'GeneralPurpose', SKU_INFO) + self.assertIn('Invalid value for --tier', str(context.exception)) + + +if __name__ == '__main__': + unittest.main() diff --git a/src/azure-cli/azure/cli/command_modules/postgresql/utils/validators.py b/src/azure-cli/azure/cli/command_modules/postgresql/utils/validators.py index 56988dc4423..71547e9dd12 100644 --- a/src/azure-cli/azure/cli/command_modules/postgresql/utils/validators.py +++ b/src/azure-cli/azure/cli/command_modules/postgresql/utils/validators.py @@ -43,6 +43,9 @@ logger = get_logger(__name__) IP_ADDRESS_CHECKER = 'https://api.ipify.org' +# Compute tiers ordered from lowest to highest capability. +PG_TIER_RANK = {'burstable': 0, 'generalpurpose': 1, 'memoryoptimized': 2} + # pylint: disable=import-outside-toplevel, raise-missing-from, unbalanced-tuple-unpacking def _get_resource_group_from_server_name(cli_ctx, server_name): @@ -912,6 +915,19 @@ def pg_restore_validator(compute_tier, **args): '--storage-type set to "PremiumV2_LRS".') +def pg_restore_tier_validator(target_tier, source_tier, sku_info): + _pg_tier_validator(target_tier, sku_info) + target_rank = PG_TIER_RANK.get(target_tier.lower()) + source_rank = PG_TIER_RANK.get(source_tier.lower()) + if target_rank is not None and source_rank is not None and target_rank < source_rank: + raise ValidationError('Invalid value for --tier. The restored server must not go below the source server ' + 'compute tier. The source server compute tier is {}.'.format(source_tier)) + + +def pg_restore_sku_validator(sku_name, sku_info, tier): + _pg_sku_name_validator(sku_name, sku_info, tier, None) + + def _pg_authentication_validator(password_auth, is_microsoft_entra_auth_enabled, admin_name, admin_id, admin_type, instance): if instance is None: