Skip to content
157 changes: 112 additions & 45 deletions web/pgadmin/misc/cloud/azure/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,15 +25,6 @@
import os


from azure.mgmt.rdbms.postgresql_flexibleservers import \
PostgreSQLManagementClient
from azure.identity import AzureCliCredential, DeviceCodeCredential,\
AuthenticationRecord
from azure.mgmt.resource import ResourceManagementClient
from azure.mgmt.subscription import SubscriptionClient
from azure.mgmt.rdbms.postgresql_flexibleservers.models import \
NameAvailabilityRequest

MODULE_NAME = 'azure'


Expand Down Expand Up @@ -265,6 +256,34 @@ def clear_session():
return make_json_response(success=1)


def _azure_sdk():
"""Defer heavy Azure SDK imports until required by user actions.
Repeat calls are cheap via sys.modules caching.
"""
from types import SimpleNamespace
from azure.identity import (
AzureCliCredential, DeviceCodeCredential, AuthenticationRecord
)
from azure.mgmt.rdbms.postgresql_flexibleservers import (
PostgreSQLManagementClient
)
from azure.mgmt.rdbms.postgresql_flexibleservers.models import (
NameAvailabilityRequest
)
from azure.mgmt.resource import ResourceManagementClient
from azure.mgmt.subscription import SubscriptionClient

return SimpleNamespace(
AzureCliCredential=AzureCliCredential,
DeviceCodeCredential=DeviceCodeCredential,
AuthenticationRecord=AuthenticationRecord,
PostgreSQLManagementClient=PostgreSQLManagementClient,
ResourceManagementClient=ResourceManagementClient,
SubscriptionClient=SubscriptionClient,
NameAvailabilityRequest=NameAvailabilityRequest,
)


class Azure:
def __init__(self, interactive_browser_credential, tenant_id=None,
session_token=None, region='eastus'):
Expand Down Expand Up @@ -367,7 +386,8 @@ def _get_azure_credentials(self):

def _azure_cli_auth(self):
if self._cli_credentials is None:
self._cli_credentials = AzureCliCredential()
sdk = _azure_sdk()
self._cli_credentials = sdk.AzureCliCredential()
self.list_subscriptions()
return self._cli_credentials

Expand All @@ -380,8 +400,9 @@ def _azure_interactive_auth_prompt_callback(
session['azure']['azure_auth_code'] = azure_auth_code

def _azure_interactive_auth(self):
sdk = _azure_sdk()
if self.authentication_record_json is None:
_interactive_credential = DeviceCodeCredential(
_interactive_credential = sdk.DeviceCodeCredential(
tenant_id=self._tenant_id,
timeout=180,
prompt_callback=self._azure_interactive_auth_prompt_callback,
Expand All @@ -392,9 +413,9 @@ def _azure_interactive_auth(self):
_auth_record = _interactive_credential.authenticate()
self.authentication_record_json = _auth_record.serialize()
else:
deserialized_auth_record = AuthenticationRecord.deserialize(
deserialized_auth_record = sdk.AuthenticationRecord.deserialize(
self.authentication_record_json)
_interactive_credential = DeviceCodeCredential(
_interactive_credential = sdk.DeviceCodeCredential(
tenant_id=self._tenant_id,
timeout=180,
prompt_callback=self._azure_interactive_auth_prompt_callback,
Expand All @@ -408,30 +429,43 @@ def _azure_interactive_auth(self):
def _get_azure_client(self, type):
""" Create/cache/return an Azure client object """
if type in self._clients:
return self._clients[type]
return self._clients[type], None

_, _credentials = self._get_azure_credentials()
status, _credentials = self._get_azure_credentials()
if not status:
return None, _credentials

try:
sdk = _azure_sdk()
except ImportError as e:
return None, str(e)

if type == 'postgresql':
client = PostgreSQLManagementClient(_credentials,
self.subscription_id)
client = sdk.PostgreSQLManagementClient(_credentials,
self.subscription_id)
elif type == 'resource':
client = ResourceManagementClient(_credentials,
self.subscription_id)
client = sdk.ResourceManagementClient(_credentials,
self.subscription_id)
elif type == 'subscription':
client = SubscriptionClient(_credentials)
client = sdk.SubscriptionClient(_credentials)

self._clients[type] = client
return self._clients[type]
return self._clients[type], None

def check_cluster_name_availability(self, cluster_name):
"""
Checks whether given server name is available or not
:param cluster_name
"""
postgresql_client = self._get_azure_client('postgresql')
try:
sdk = _azure_sdk()
except ImportError as e:
return False, str(e)
postgresql_client, error = self._get_azure_client('postgresql')
if not postgresql_client:
return False, error
res = postgresql_client.check_name_availability.execute(
NameAvailabilityRequest(
sdk.NameAvailabilityRequest(
name=cluster_name,
type='Microsoft.DBforPostgreSQL/flexibleServers'))
res = res.__dict__
Expand All @@ -442,13 +476,22 @@ def list_subscriptions(self):
List subscriptions
:return:
"""
subscription_client = self._get_azure_client('subscription')
sub_list = subscription_client.subscriptions.list()
subscriptions_list = []
for group in list(sub_list):
subscriptions_list.append(
{'subscription_id': group.subscription_id,
'subscription_name': group.display_name})
try:
subscription_client, error = self._get_azure_client('subscription')
if not subscription_client:
if current_app:
current_app.logger.error(error)
return subscriptions_list
sub_list = subscription_client.subscriptions.list()
for group in list(sub_list):
subscriptions_list.append(
{'subscription_id': group.subscription_id,
'subscription_name': group.display_name})
except ImportError as e:
if current_app:
current_app.logger.error(str(e))
return subscriptions_list
return subscriptions_list

def list_resource_groups(self, subscription_id):
Expand All @@ -458,14 +501,23 @@ def list_resource_groups(self, subscription_id):
:return:
"""
self.subscription_id = subscription_id
resource_client = self._get_azure_client('resource')
group_list = resource_client.resource_groups.list()
resource_groups_list = []
for group in list(group_list):
resource_groups_list.append(
{'label': group.name,
'value': group.name,
'region': group.location})
try:
resource_client, error = self._get_azure_client('resource')
if not resource_client:
if current_app:
current_app.logger.error(error)
return resource_groups_list
group_list = resource_client.resource_groups.list()
for group in list(group_list):
resource_groups_list.append(
{'label': group.name,
'value': group.name,
'region': group.location})
except ImportError as e:
if current_app:
current_app.logger.error(str(e))
return resource_groups_list
return resource_groups_list

def list_regions(self, subscription_id):
Expand All @@ -475,13 +527,22 @@ def list_regions(self, subscription_id):
:return:
"""
self.subscription_id = subscription_id
subscription_client = self._get_azure_client('subscription')
locations = subscription_client.subscriptions.list_locations(
subscription_id=self.subscription_id)
locations_list = []
for location in locations:
locations_list.append(
{'label': location.display_name, 'value': location.name})
try:
subscription_client, error = self._get_azure_client('subscription')
if not subscription_client:
if current_app:
current_app.logger.error(error)
return locations_list
locations = subscription_client.subscriptions.list_locations(
subscription_id=self.subscription_id)
for location in locations:
locations_list.append(
{'label': location.display_name, 'value': location.name})
except ImportError as e:
if current_app:
current_app.logger.error(str(e))
return locations_list
return locations_list

def is_zone_redundant_ha_supported(self, region):
Expand All @@ -492,8 +553,10 @@ def is_zone_redundant_ha_supported(self, region):
else:
self._available_capabilities_list = \
self._get_available_capabilities_list(region)
return self._available_capabilities_list[0][
'zone_redundant_ha_supported']
if self._available_capabilities_list:
return self._available_capabilities_list[0][
'zone_redundant_ha_supported']
return False

def list_azure_availability_zones(self, region):
"""
Expand Down Expand Up @@ -596,7 +659,11 @@ def _get_available_capabilities_object(self, region):
:param region:
:return: azure capabilities object
"""
postgresql_client = self._get_azure_client('postgresql')
postgresql_client, error = self._get_azure_client('postgresql')
if not postgresql_client:
if current_app:
current_app.logger.error(error)
return []
return postgresql_client.location_based_capabilities.execute(
location_name=region)

Expand Down
103 changes: 103 additions & 0 deletions web/pgadmin/misc/cloud/azure/tests/test_azure_session_state.py
Original file line number Diff line number Diff line change
Expand Up @@ -206,3 +206,106 @@ def runTest(self):
"cloud.azure must not import the unsafe deserializer")
self.assertNotIn(forbidden + '.dumps(', src)
self.assertNotIn(forbidden + '.loads(', src)


class TestAzureImportErrorHandling(
_SkipServerSetUpMixin, BaseTestGenerator):
"""Azure methods must handle ImportError cleanly without raising."""

scenarios = [('default', dict())]

def runTest(self):
from unittest.mock import patch
from pgadmin.misc.cloud.azure import Azure
import pgadmin.misc.cloud.azure as azure_mod

a = Azure.from_state({'tenant_id': 'tid'})
with patch.object(
azure_mod, '_azure_sdk',
side_effect=ImportError('No module named azure.mgmt')):
client, err = a._get_azure_client('postgresql')
self.assertIsNone(client)
self.assertIn('azure.mgmt', err)

avail, msg = a.check_cluster_name_availability('server')
self.assertFalse(avail)
self.assertIn('azure.mgmt', msg)

subs = a.list_subscriptions()
self.assertEqual(subs, [])

rgs = a.list_resource_groups('sub-1')
self.assertEqual(rgs, [])

regs = a.list_regions('sub-1')
self.assertEqual(regs, [])

ha = a.is_zone_redundant_ha_supported('eastus')
self.assertFalse(ha)

from flask import Flask, current_app
app = Flask(__name__)
app.secret_key = 'test'
with app.app_context():
with patch.object(current_app.logger, 'error') as mock_log:
self.assertEqual(a.list_subscriptions(), [])
mock_log.assert_called()
self.assertIn('azure.mgmt', str(mock_log.call_args[0][0]))

with patch.object(current_app.logger, 'error') as mock_log:
self.assertEqual(a.list_resource_groups('sub-1'), [])
mock_log.assert_called()
self.assertIn('azure.mgmt', str(mock_log.call_args[0][0]))

with patch.object(current_app.logger, 'error') as mock_log:
self.assertEqual(a.list_regions('sub-1'), [])
mock_log.assert_called()
self.assertIn('azure.mgmt', str(mock_log.call_args[0][0]))


class TestAzureCredentialFailureHandling(
_SkipServerSetUpMixin, BaseTestGenerator):
"""Azure methods must propagate or log credential errors cleanly."""

scenarios = [('default', dict())]

def runTest(self):
from flask import Flask, current_app
from unittest.mock import patch
from pgadmin.misc.cloud.azure import Azure

a = Azure.from_state({'tenant_id': 'tid'})
cred_error = 'Azure authentication failed: az login expired'
with patch.object(
a, '_get_azure_credentials',
return_value=(False, cred_error)):
client, err = a._get_azure_client('postgresql')
self.assertIsNone(client)
self.assertEqual(err, cred_error)

avail, msg = a.check_cluster_name_availability('server')
self.assertFalse(avail)
self.assertEqual(msg, cred_error)

app = Flask(__name__)
app.secret_key = 'test'
with app.app_context():
with patch.object(current_app.logger, 'error') as mock_log:
subs = a.list_subscriptions()
self.assertEqual(subs, [])
mock_log.assert_called_with(cred_error)

with patch.object(current_app.logger, 'error') as mock_log:
rgs = a.list_resource_groups('sub-1')
self.assertEqual(rgs, [])
mock_log.assert_called_with(cred_error)

with patch.object(current_app.logger, 'error') as mock_log:
regs = a.list_regions('sub-1')
self.assertEqual(regs, [])
mock_log.assert_called_with(cred_error)

with patch.object(current_app.logger, 'error') as mock_log:
caps = a._get_available_capabilities_object('eastus')
self.assertEqual(caps, [])
mock_log.assert_called_with(cred_error)
Loading
Loading