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
4 changes: 4 additions & 0 deletions src/azure-cli/azure/cli/command_modules/postgresql/_help.py
Original file line number Diff line number Diff line change
Expand Up @@ -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'] = """
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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()

Expand Down Expand Up @@ -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(
Expand All @@ -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
)

Expand Down Expand Up @@ -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)

Expand Down
Original file line number Diff line number Diff line change
@@ -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):

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

Instead add new args to restore test already in project and re-record results

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()
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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:
Expand Down
Loading