Skip to content

Commit aa2bfbd

Browse files
authored
Allow Azure Arc managed identity selectors
1 parent 1c008b8 commit aa2bfbd

2 files changed

Lines changed: 56 additions & 11 deletions

File tree

msal/managed_identity.py

Lines changed: 12 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -452,13 +452,8 @@ def _obtain_token(
452452
)
453453
arc_endpoint = _get_arc_endpoint()
454454
if arc_endpoint:
455-
if ManagedIdentity.is_user_assigned(managed_identity):
456-
raise ManagedIdentityError( # Note: Azure Identity for Python raised exception too
457-
"Invalid managed_identity parameter. "
458-
"Azure Arc supports only system-assigned managed identity, "
459-
"See also "
460-
"https://learn.microsoft.com/en-us/azure/service-fabric/configure-existing-cluster-enable-managed-identity-token-service")
461-
return _obtain_token_on_arc(http_client, arc_endpoint, resource)
455+
return _obtain_token_on_arc(
456+
http_client, arc_endpoint, resource, managed_identity)
462457
return _obtain_token_on_azure_vm(http_client, managed_identity, resource)
463458

464459

@@ -643,12 +638,19 @@ def _obtain_token_on_service_fabric(
643638
class ArcPlatformNotSupportedError(ManagedIdentityError):
644639
pass
645640

646-
def _obtain_token_on_arc(http_client, endpoint, resource):
641+
def _obtain_token_on_arc(http_client, endpoint, resource, managed_identity=None):
647642
# https://learn.microsoft.com/en-us/azure/azure-arc/servers/managed-identity-authentication
648643
logger.debug("Obtaining token via managed identity on Azure Arc")
644+
params = {"api-version": "2020-06-01", "resource": resource}
645+
if managed_identity:
646+
_adjust_param(params, managed_identity, types_mapping={
647+
ManagedIdentity.CLIENT_ID: "client_id",
648+
ManagedIdentity.RESOURCE_ID: "mi_res_id",
649+
ManagedIdentity.OBJECT_ID: "object_id",
650+
})
649651
resp = http_client.get(
650652
endpoint,
651-
params={"api-version": "2020-06-01", "resource": resource},
653+
params=params.copy(),
652654
headers={"Metadata": "true"},
653655
)
654656
www_auth = "www-authenticate" # Header in lower case
@@ -674,7 +676,7 @@ def _obtain_token_on_arc(http_client, endpoint, resource):
674676
secret = f.read()
675677
response = http_client.get(
676678
endpoint,
677-
params={"api-version": "2020-06-01", "resource": resource},
679+
params=params.copy(),
678680
headers={"Metadata": "true", "Authorization": "Basic {}".format(secret)},
679681
)
680682
try:

tests/test_mi.py

Lines changed: 44 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -477,6 +477,50 @@ def test_arc_error_should_be_normalized(self, mocked_stat):
477477
if sys.platform in _supported_arc_platforms_and_their_prefixes:
478478
self.fail("Should not raise ArcPlatformNotSupportedError")
479479

480+
def _assert_user_assigned_selector(self, managed_identity, selector_name, selector_value):
481+
app = ManagedIdentityClient(managed_identity, http_client=requests.Session())
482+
with patch.object(app._http_client, "get", side_effect=[
483+
self.challenge,
484+
MinimalResponse(
485+
status_code=200,
486+
text='{"access_token": "AT", "expires_in": "1234", "resource": "R"}',
487+
),
488+
]) as mocked_method:
489+
try:
490+
result = app.acquire_token_for_client(resource="R")
491+
self.assertEqual("AT", result["access_token"])
492+
expected_params = {
493+
"api-version": "2020-06-01",
494+
"resource": "R",
495+
selector_name: selector_value,
496+
}
497+
self.assertEqual(expected_params, mocked_method.call_args_list[0].kwargs["params"])
498+
self.assertEqual(expected_params, mocked_method.call_args_list[1].kwargs["params"])
499+
except ArcPlatformNotSupportedError:
500+
if sys.platform in _supported_arc_platforms_and_their_prefixes:
501+
self.fail("Should not raise ArcPlatformNotSupportedError")
502+
503+
def test_arc_user_assigned_client_id_should_be_forwarded(self, mocked_stat):
504+
self._assert_user_assigned_selector(
505+
UserAssignedManagedIdentity(client_id="client-id"),
506+
"client_id",
507+
"client-id",
508+
)
509+
510+
def test_arc_user_assigned_resource_id_should_be_forwarded_as_mi_res_id(self, mocked_stat):
511+
self._assert_user_assigned_selector(
512+
UserAssignedManagedIdentity(resource_id="resource-id"),
513+
"mi_res_id",
514+
"resource-id",
515+
)
516+
517+
def test_arc_user_assigned_object_id_should_be_forwarded(self, mocked_stat):
518+
self._assert_user_assigned_selector(
519+
UserAssignedManagedIdentity(object_id="object-id"),
520+
"object_id",
521+
"object-id",
522+
)
523+
480524

481525
class GetManagedIdentitySourceTestCase(unittest.TestCase):
482526

@@ -531,4 +575,3 @@ def test_cloud_shell(self):
531575

532576
def test_default_to_vm(self):
533577
self.assertEqual(get_managed_identity_source(), DEFAULT_TO_VM)
534-

0 commit comments

Comments
 (0)