Skip to content
Open
Show file tree
Hide file tree
Changes from 2 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
135 changes: 94 additions & 41 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 @@ -410,16 +431,23 @@ def _get_azure_client(self, type):
if type in self._clients:
return self._clients[type]

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

try:
sdk = _azure_sdk()
except ImportError:
return None

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]
Expand All @@ -429,9 +457,15 @@ def check_cluster_name_availability(self, cluster_name):
Checks whether given server name is available or not
:param cluster_name
"""
try:
sdk = _azure_sdk()
except ImportError as e:
return False, str(e)
postgresql_client = self._get_azure_client('postgresql')
if not postgresql_client:
return False, 'Failed to initialize Azure client.'
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,18 @@ 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 = self._get_azure_client('subscription')
if not subscription_client:
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:
return subscriptions_list
return subscriptions_list

def list_resource_groups(self, subscription_id):
Expand All @@ -458,14 +497,19 @@ 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 = self._get_azure_client('resource')
if not resource_client:
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:
return resource_groups_list
return resource_groups_list

def list_regions(self, subscription_id):
Expand All @@ -475,13 +519,18 @@ 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 = self._get_azure_client('subscription')
if not subscription_client:
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:
return locations_list
return locations_list

def is_zone_redundant_ha_supported(self, region):
Expand All @@ -492,8 +541,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 @@ -597,6 +648,8 @@ def _get_available_capabilities_object(self, region):
:return: azure capabilities object
"""
postgresql_client = self._get_azure_client('postgresql')
if not postgresql_client:
return []
return postgresql_client.location_based_capabilities.execute(
location_name=region)

Expand Down
35 changes: 35 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,38 @@ 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 = a._get_azure_client('postgresql')
self.assertIsNone(client)

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)
Loading
Loading