From f7bf472c7395ea555d9c6ce02c2b48a923ce51b3 Mon Sep 17 00:00:00 2001 From: Wei Lim Date: Mon, 6 Apr 2026 14:49:29 -0700 Subject: [PATCH 1/6] Use resource tenant ID where appropriate Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- cli/azd/cmd/container.go | 3 + cli/azd/cmd/deeper_coverage3_test.go | 182 ++++++++++++++++++ cli/azd/cmd/env.go | 13 +- cli/azd/pkg/account/subscriptions_manager.go | 6 + .../current_principal_id_provider.go | 11 +- .../current_principal_id_provider_test.go | 83 ++++++++ 6 files changed, 287 insertions(+), 11 deletions(-) create mode 100644 cli/azd/pkg/infra/provisioning/current_principal_id_provider_test.go diff --git a/cli/azd/cmd/container.go b/cli/azd/cmd/container.go index 0f1f5d467a3..693e6bb7a77 100644 --- a/cli/azd/cmd/container.go +++ b/cli/azd/cmd/container.go @@ -703,6 +703,9 @@ func registerCommonDependencies(container *ioc.NestedContainer) { container.MustRegisterSingleton(func(subManager *account.SubscriptionsManager) account.SubscriptionTenantResolver { return subManager }) + container.MustRegisterSingleton(func(subManager *account.SubscriptionsManager) account.SubscriptionResolver { + return subManager + }) // Tools container.MustRegisterSingleton(azapi.NewAzureClient) diff --git a/cli/azd/cmd/deeper_coverage3_test.go b/cli/azd/cmd/deeper_coverage3_test.go index f5110c2a675..d1a6f4076dc 100644 --- a/cli/azd/cmd/deeper_coverage3_test.go +++ b/cli/azd/cmd/deeper_coverage3_test.go @@ -9,10 +9,17 @@ import ( "errors" "fmt" "testing" + "time" + "github.com/Azure/azure-sdk-for-go/sdk/azcore" + "github.com/Azure/azure-sdk-for-go/sdk/azcore/policy" "github.com/azure/azure-dev/cli/azd/internal" + "github.com/azure/azure-dev/cli/azd/pkg/account" "github.com/azure/azure-dev/cli/azd/pkg/alpha" + "github.com/azure/azure-dev/cli/azd/pkg/azapi" + "github.com/azure/azure-dev/cli/azd/pkg/cloud" "github.com/azure/azure-dev/cli/azd/pkg/config" + "github.com/azure/azure-dev/cli/azd/pkg/entraid" "github.com/azure/azure-dev/cli/azd/pkg/environment" "github.com/azure/azure-dev/cli/azd/pkg/environment/azdcontext" "github.com/azure/azure-dev/cli/azd/pkg/input" @@ -20,6 +27,7 @@ import ( "github.com/azure/azure-dev/cli/azd/pkg/output" "github.com/azure/azure-dev/cli/azd/pkg/project" "github.com/azure/azure-dev/cli/azd/pkg/prompt" + "github.com/azure/azure-dev/cli/azd/test/mocks" "github.com/azure/azure-dev/cli/azd/test/mocks/mockenv" "github.com/azure/azure-dev/cli/azd/test/mocks/mockinput" "github.com/stretchr/testify/assert" @@ -138,6 +146,49 @@ func (m *mockSubTenantResolver) LookupTenant(ctx context.Context, subscriptionId return args.String(0), args.Error(1) } +func (m *mockSubTenantResolver) GetSubscription(ctx context.Context, subscriptionId string) (*account.Subscription, error) { + tenantId, err := m.LookupTenant(ctx, subscriptionId) + if err != nil { + return nil, err + } + + return &account.Subscription{ + Id: subscriptionId, + TenantId: tenantId, + UserAccessTenantId: tenantId, + }, nil +} + +type staticSubscriptionResolver struct { + subscription *account.Subscription +} + +func (s *staticSubscriptionResolver) LookupTenant(ctx context.Context, subscriptionId string) (string, error) { + return s.subscription.UserAccessTenantId, nil +} + +func (s *staticSubscriptionResolver) GetSubscription(ctx context.Context, subscriptionId string) (*account.Subscription, error) { + return s.subscription, nil +} + +type mockEnvSetSecretEntraIdService struct { + entraid.EntraIdService + subscriptionId string + scope string + roleId string + principalId string +} + +func (m *mockEnvSetSecretEntraIdService) CreateRbac( + ctx context.Context, subscriptionId string, scope, roleId, principalId string, +) error { + m.subscriptionId = subscriptionId + m.scope = scope + m.roleId = roleId + m.principalId = principalId + return nil +} + // ==================== envSetSecretAction Tests ==================== func newTestEnvSetSecretAction( @@ -874,6 +925,137 @@ func Test_EnvSetSecretConstructor(t *testing.T) { require.NotNil(t, action) } +func Test_EnvSetSecretAction_UsesResourceTenantForKeyVaultAndPrincipalId(t *testing.T) { + t.Parallel() + + console := mockinput.NewMockConsole() + selectCount := 0 + console.WhenSelect(func(options input.ConsoleOptions) bool { + return true + }).RespondFn(func(options input.ConsoleOptions) (any, error) { + selectCount++ + if selectCount > 2 { + return nil, fmt.Errorf("unexpected select: %s", options.Message) + } + return 0, nil + }) + + promptCount := 0 + console.WhenPrompt(func(options input.ConsoleOptions) bool { + return true + }).RespondFn(func(options input.ConsoleOptions) (any, error) { + promptCount++ + switch promptCount { + case 1: + return "kv-name", nil + case 2: + return "my-secret-kv", nil + case 3: + return "secret-value", nil + default: + return nil, fmt.Errorf("unexpected prompt: %s", options.Message) + } + }) + + env := environment.NewWithValues("test", map[string]string{}) + envManager := &mockenv.MockEnvManager{} + envManager.On("Save", mock.Anything, env).Return(nil) + + prompter := &mockPrompter{} + prompter.On("PromptSubscription", + mock.Anything, + "Select the subscription where you want to create the Key Vault secret", + ).Return("sub-123", nil) + prompter.On("PromptLocation", + mock.Anything, + "sub-123", + "Select the location to create the Key Vault", + mock.Anything, + mock.Anything, + ).Return("westus", nil) + prompter.On("PromptResourceGroupFrom", + mock.Anything, + "sub-123", + "westus", + prompt.PromptResourceGroupFromOptions{ + DefaultName: "rg-for-my-key-vault", + NewResourceGroupHelp: "The name of the new resource group where the Key Vault will be created.", + }, + ).Return("rg-name", nil) + + kvSvc := &mockKeyVaultService{} + kvSvc.On("ListSubscriptionVaults", mock.Anything, "sub-123").Return([]keyvault.Vault{}, nil) + kvSvc.On("CreateVault", + mock.Anything, + "resource-tenant", + "sub-123", + "rg-name", + "westus", + "kv-name", + ).Return(keyvault.Vault{ + Id: "/subscriptions/sub-123/resourceGroups/rg-name/providers/Microsoft.KeyVault/vaults/kv-name", + Name: "kv-name", + }, nil) + kvSvc.On("CreateKeyVaultSecret", + mock.Anything, + "sub-123", + "kv-name", + "my-secret-kv", + "secret-value", + ).Return(nil) + + mockContext := mocks.NewMockContext(context.Background()) + userProfileService := azapi.NewUserProfileService( + &mocks.MockMultiTenantCredentialProvider{ + TokenMap: map[string]mocks.MockCredentials{ + "resource-tenant": { + GetTokenFn: func(ctx context.Context, options policy.TokenRequestOptions) (azcore.AccessToken, error) { + return azcore.AccessToken{ + // cspell:disable-next-line + Token: "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJvaWQiOiJ0aGlzLWlzLWEtdGVzdCJ9.vrKZx2J7-hsydI4rzdFVHqU1S6lHqLT95VSPx2RfQ04", + ExpiresOn: time.Now().Add(time.Hour), + }, nil + }, + }, + }, + }, + &azcore.ClientOptions{ + Transport: mockContext.HttpClient, + }, + cloud.AzurePublic(), + ) + + entraIdService := &mockEnvSetSecretEntraIdService{} + action := &envSetSecretAction{ + console: console, + env: env, + envManager: envManager, + flags: &envSetFlags{}, + args: []string{"MY_SECRET"}, + prompter: prompter, + kvService: kvSvc, + entraIdService: entraIdService, + subResolver: &staticSubscriptionResolver{ + subscription: &account.Subscription{ + Id: "sub-123", + TenantId: "resource-tenant", + UserAccessTenantId: "home-tenant", + }, + }, + userProfileService: userProfileService, + alphaFeatureManager: alpha.NewFeaturesManagerWithConfig(config.NewEmptyConfig()), + projectConfig: &project.ProjectConfig{ + Resources: map[string]*project.ResourceConfig{}, + }, + } + + _, err := action.Run(t.Context()) + require.NoError(t, err) + require.Equal(t, "sub-123", entraIdService.subscriptionId) + require.Equal(t, keyvault.RoleIdKeyVaultAdministrator, entraIdService.roleId) + require.Equal(t, "this-is-a-test", entraIdService.principalId) +} + // ==================== Suppressed errors.Is / errors.AsType coverage ==================== func Test_ErrorWithSuggestion_Type(t *testing.T) { diff --git a/cli/azd/cmd/env.go b/cli/azd/cmd/env.go index 5f299e8bb91..40bd177862f 100644 --- a/cli/azd/cmd/env.go +++ b/cli/azd/cmd/env.go @@ -359,7 +359,7 @@ type envSetSecretAction struct { prompter prompt.Prompter kvService keyvault.KeyVaultService entraIdService entraid.EntraIdService - subResolver account.SubscriptionTenantResolver + subResolver account.SubscriptionResolver userProfileService *azapi.UserProfileService alphaFeatureManager *alpha.FeatureManager projectConfig *project.ProjectConfig @@ -498,10 +498,11 @@ func (e *envSetSecretAction) Run(ctx context.Context) (*actions.ActionResult, er if err != nil { return nil, fmt.Errorf("prompting for subscription: %w", err) } - tenantId, err := e.subResolver.LookupTenant(ctx, subId) + subscription, err := e.subResolver.GetSubscription(ctx, subId) if err != nil { - return nil, fmt.Errorf("looking up tenant for subscription: %w", err) + return nil, fmt.Errorf("getting subscription %s: %w", subId, err) } + resourceTenantId := subscription.TenantId e.console.ShowSpinner(ctx, "Finding Key Vaults from the selected subscription", input.Step) vaultsList, err := e.kvService.ListSubscriptionVaults(ctx, subId) @@ -602,7 +603,7 @@ func (e *envSetSecretAction) Run(ctx context.Context) (*actions.ActionResult, er } e.console.ShowSpinner(ctx, "Creating Key Vault", input.Step) - vault, err := e.kvService.CreateVault(ctx, tenantId, subId, rg, location, kvAccountName) + vault, err := e.kvService.CreateVault(ctx, resourceTenantId, subId, rg, location, kvAccountName) e.console.StopSpinner(ctx, "", input.Step) if err != nil { return nil, fmt.Errorf("error creating Key Vault: %w", err) @@ -611,7 +612,7 @@ func (e *envSetSecretAction) Run(ctx context.Context) (*actions.ActionResult, er // RBAC role assignment e.console.ShowSpinner(ctx, "Adding Administrator Role", input.Step) - principalId, err := azureutil.GetCurrentPrincipalId(ctx, e.userProfileService, tenantId) + principalId, err := azureutil.GetCurrentPrincipalId(ctx, e.userProfileService, resourceTenantId) if err != nil { return nil, fmt.Errorf("getting current principal ID: %w", err) } @@ -734,7 +735,7 @@ func newEnvSetSecretAction( prompter prompt.Prompter, kvService keyvault.KeyVaultService, entraIdService entraid.EntraIdService, - subResolver account.SubscriptionTenantResolver, + subResolver account.SubscriptionResolver, userProfileService *azapi.UserProfileService, alphaFeatureManager *alpha.FeatureManager, projectConfig *project.ProjectConfig, diff --git a/cli/azd/pkg/account/subscriptions_manager.go b/cli/azd/pkg/account/subscriptions_manager.go index 9d819cf9c48..447bedf6438 100644 --- a/cli/azd/pkg/account/subscriptions_manager.go +++ b/cli/azd/pkg/account/subscriptions_manager.go @@ -28,6 +28,12 @@ type SubscriptionTenantResolver interface { LookupTenant(ctx context.Context, subscriptionId string) (tenantId string, err error) } +// SubscriptionResolver allows resolving both the access tenant and subscription details. +type SubscriptionResolver interface { + SubscriptionTenantResolver + GetSubscription(ctx context.Context, subscriptionId string) (*Subscription, error) +} + // Typically auth.Manager type principalInfoProvider interface { GetLoggedInServicePrincipalTenantID(ctx context.Context) (*string, error) diff --git a/cli/azd/pkg/infra/provisioning/current_principal_id_provider.go b/cli/azd/pkg/infra/provisioning/current_principal_id_provider.go index 6fed01e918e..493bca7cd09 100644 --- a/cli/azd/pkg/infra/provisioning/current_principal_id_provider.go +++ b/cli/azd/pkg/infra/provisioning/current_principal_id_provider.go @@ -24,7 +24,7 @@ type CurrentPrincipalIdProvider interface { func NewPrincipalIdProvider( env *environment.Environment, userProfileService *azapi.UserProfileService, - subResolver account.SubscriptionTenantResolver, + subResolver account.SubscriptionResolver, authManager *auth.Manager, ) CurrentPrincipalIdProvider { return &principalIDProvider{ @@ -38,17 +38,18 @@ func NewPrincipalIdProvider( type principalIDProvider struct { env *environment.Environment userProfileService *azapi.UserProfileService - subResolver account.SubscriptionTenantResolver + subResolver account.SubscriptionResolver authManager *auth.Manager } func (p *principalIDProvider) CurrentPrincipalId(ctx context.Context) (string, error) { - tenantId, err := p.subResolver.LookupTenant(ctx, p.env.GetSubscriptionId()) + subscriptionId := p.env.GetSubscriptionId() + sub, err := p.subResolver.GetSubscription(ctx, subscriptionId) if err != nil { - return "", fmt.Errorf("getting tenant id for subscription %s. Error: %w", p.env.GetSubscriptionId(), err) + return "", fmt.Errorf("getting subscription %s: %w", subscriptionId, err) } - principalId, err := azureutil.GetCurrentPrincipalId(ctx, p.userProfileService, tenantId) + principalId, err := azureutil.GetCurrentPrincipalId(ctx, p.userProfileService, sub.TenantId) if err != nil { return "", fmt.Errorf("fetching current user information: %w", err) } diff --git a/cli/azd/pkg/infra/provisioning/current_principal_id_provider_test.go b/cli/azd/pkg/infra/provisioning/current_principal_id_provider_test.go new file mode 100644 index 00000000000..62837ba1ac9 --- /dev/null +++ b/cli/azd/pkg/infra/provisioning/current_principal_id_provider_test.go @@ -0,0 +1,83 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +package provisioning + +import ( + "context" + "testing" + "time" + + "github.com/Azure/azure-sdk-for-go/sdk/azcore" + "github.com/Azure/azure-sdk-for-go/sdk/azcore/policy" + "github.com/azure/azure-dev/cli/azd/pkg/account" + "github.com/azure/azure-dev/cli/azd/pkg/azapi" + "github.com/azure/azure-dev/cli/azd/pkg/cloud" + "github.com/azure/azure-dev/cli/azd/pkg/environment" + "github.com/azure/azure-dev/cli/azd/test/mocks" + "github.com/stretchr/testify/require" +) + +type fakeSubscriptionResolver struct { + subscription *account.Subscription + lookupTenantCalls int + getSubscriptionCalls int +} + +func (f *fakeSubscriptionResolver) LookupTenant(ctx context.Context, subscriptionId string) (string, error) { + f.lookupTenantCalls++ + return "home-tenant", nil +} + +func (f *fakeSubscriptionResolver) GetSubscription(ctx context.Context, subscriptionId string) (*account.Subscription, error) { + f.getSubscriptionCalls++ + return f.subscription, nil +} + +func TestPrincipalIDProvider_CurrentPrincipalIdUsesSubscriptionTenant(t *testing.T) { + t.Parallel() + + mockContext := mocks.NewMockContext(context.Background()) + userProfileService := azapi.NewUserProfileService( + &mocks.MockMultiTenantCredentialProvider{ + TokenMap: map[string]mocks.MockCredentials{ + "resource-tenant": { + GetTokenFn: func(ctx context.Context, options policy.TokenRequestOptions) (azcore.AccessToken, error) { + return azcore.AccessToken{ + // cspell:disable-next-line + Token: "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJvaWQiOiJ0aGlzLWlzLWEtdGVzdCJ9.vrKZx2J7-hsydI4rzdFVHqU1S6lHqLT95VSPx2RfQ04", + ExpiresOn: time.Now().Add(time.Hour), + }, nil + }, + }, + }, + }, + &azcore.ClientOptions{ + Transport: mockContext.HttpClient, + }, + cloud.AzurePublic(), + ) + + resolver := &fakeSubscriptionResolver{ + subscription: &account.Subscription{ + Id: "sub-123", + TenantId: "resource-tenant", + UserAccessTenantId: "home-tenant", + }, + } + + provider := NewPrincipalIdProvider( + environment.NewWithValues("test", map[string]string{ + environment.SubscriptionIdEnvVarName: "sub-123", + }), + userProfileService, + resolver, + nil, + ) + + principalId, err := provider.CurrentPrincipalId(t.Context()) + require.NoError(t, err) + require.Equal(t, "this-is-a-test", principalId) + require.Equal(t, 1, resolver.getSubscriptionCalls) + require.Zero(t, resolver.lookupTenantCalls) +} From 0f1bb24c47c8b92f5e897f10cf216a7da7c414aa Mon Sep 17 00:00:00 2001 From: Wei Lim Date: Mon, 6 Apr 2026 14:59:22 -0700 Subject: [PATCH 2/6] Prefer ARM token over Graph /me Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- cli/azd/pkg/azureutil/principal.go | 34 +++++----- cli/azd/pkg/azureutil/principal_test.go | 86 +++++++++++++++++++++++++ 2 files changed, 105 insertions(+), 15 deletions(-) create mode 100644 cli/azd/pkg/azureutil/principal_test.go diff --git a/cli/azd/pkg/azureutil/principal.go b/cli/azd/pkg/azureutil/principal.go index 15a9ead1199..0e8c3d62e1c 100644 --- a/cli/azd/pkg/azureutil/principal.go +++ b/cli/azd/pkg/azureutil/principal.go @@ -5,32 +5,36 @@ package azureutil import ( "context" + "errors" "fmt" "github.com/azure/azure-dev/cli/azd/pkg/auth" "github.com/azure/azure-dev/cli/azd/pkg/azapi" ) -// GetCurrentPrincipalId returns the object id of the current -// principal authenticated with the CLI -// (via ad sp signed-in-user), falling back to extracting the -// `oid` claim from an access token a principal can not be -// obtained in this way. +// GetCurrentPrincipalId returns the object ID of the current principal authenticated with the CLI. +// It prefers the oid claim from an ARM access token, falling back to Graph /me when the token does +// not include a usable oid. func GetCurrentPrincipalId(ctx context.Context, userProfile *azapi.UserProfileService, tenantId string) (string, error) { - principalId, err := userProfile.GetSignedInUserId(ctx, tenantId) + token, err := userProfile.GetAccessToken(ctx, tenantId) if err == nil { - return principalId, nil - } + oid, oidErr := auth.GetOidFromAccessToken(token.AccessToken) + if oidErr == nil { + return oid, nil + } - token, err := userProfile.GetAccessToken(ctx, tenantId) - if err != nil { - return "", fmt.Errorf("getting access token: %w", err) + err = fmt.Errorf("getting oid from token: %w", oidErr) + } else { + err = fmt.Errorf("getting access token: %w", err) } - oid, err := auth.GetOidFromAccessToken(token.AccessToken) - if err != nil { - return "", fmt.Errorf("getting oid from token: %w", err) + principalId, graphErr := userProfile.GetSignedInUserId(ctx, tenantId) + if graphErr == nil { + return principalId, nil } - return oid, nil + return "", fmt.Errorf( + "resolving current principal ID from token oid and Graph fallback: %w", + errors.Join(err, fmt.Errorf("getting signed-in user id: %w", graphErr)), + ) } diff --git a/cli/azd/pkg/azureutil/principal_test.go b/cli/azd/pkg/azureutil/principal_test.go new file mode 100644 index 00000000000..d3730b8d188 --- /dev/null +++ b/cli/azd/pkg/azureutil/principal_test.go @@ -0,0 +1,86 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +package azureutil + +import ( + "context" + "net/http" + "strings" + "testing" + "time" + + "github.com/Azure/azure-sdk-for-go/sdk/azcore" + "github.com/Azure/azure-sdk-for-go/sdk/azcore/policy" + "github.com/azure/azure-dev/cli/azd/pkg/azapi" + "github.com/azure/azure-dev/cli/azd/pkg/cloud" + "github.com/azure/azure-dev/cli/azd/pkg/graphsdk" + "github.com/azure/azure-dev/cli/azd/test/mocks" + "github.com/stretchr/testify/require" +) + +func TestGetCurrentPrincipalId_PrefersOidFromAccessToken(t *testing.T) { + t.Parallel() + + mockContext := mocks.NewMockContext(context.Background()) + userProfile := azapi.NewUserProfileService( + &mocks.MockMultiTenantCredentialProvider{ + TokenMap: map[string]mocks.MockCredentials{ + "resource-tenant": { + GetTokenFn: func(ctx context.Context, options policy.TokenRequestOptions) (azcore.AccessToken, error) { + return azcore.AccessToken{ + // cspell:disable-next-line + Token: "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJvaWQiOiJ0aGlzLWlzLWEtdGVzdCJ9.vrKZx2J7-hsydI4rzdFVHqU1S6lHqLT95VSPx2RfQ04", + ExpiresOn: time.Now().Add(time.Hour), + }, nil + }, + }, + }, + }, + &azcore.ClientOptions{ + Transport: mockContext.HttpClient, + }, + cloud.AzurePublic(), + ) + + principalId, err := GetCurrentPrincipalId(*mockContext.Context, userProfile, "resource-tenant") + require.NoError(t, err) + require.Equal(t, "this-is-a-test", principalId) +} + +func TestGetCurrentPrincipalId_FallsBackToGraphWhenOidMissing(t *testing.T) { + t.Parallel() + + mockContext := mocks.NewMockContext(context.Background()) + mockContext.HttpClient.When(func(request *http.Request) bool { + return request.Method == http.MethodGet && strings.Contains(request.URL.Path, "/me") + }).RespondFn(func(request *http.Request) (*http.Response, error) { + return mocks.CreateHttpResponseWithBody(request, http.StatusOK, &graphsdk.UserProfile{ + Id: "graph-user-id", + }) + }) + + userProfile := azapi.NewUserProfileService( + &mocks.MockMultiTenantCredentialProvider{ + TokenMap: map[string]mocks.MockCredentials{ + "resource-tenant": { + GetTokenFn: func(ctx context.Context, options policy.TokenRequestOptions) (azcore.AccessToken, error) { + return azcore.AccessToken{ + // cspell:disable-next-line + Token: "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJ0ZXN0IjoiZmFpbCJ9.0Stzv5ZHG96ss-0_AnANqZfVLoULCtivJCE8AVWFZi8", + ExpiresOn: time.Now().Add(time.Hour), + }, nil + }, + }, + }, + }, + &azcore.ClientOptions{ + Transport: mockContext.HttpClient, + }, + cloud.AzurePublic(), + ) + + principalId, err := GetCurrentPrincipalId(*mockContext.Context, userProfile, "resource-tenant") + require.NoError(t, err) + require.Equal(t, "graph-user-id", principalId) +} From 779eaa1f2cf1047f1c370dbae23a1aa36002776e Mon Sep 17 00:00:00 2001 From: Wei Lim Date: Mon, 6 Apr 2026 15:18:52 -0700 Subject: [PATCH 3/6] Refactor to SubscriptionResolver Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- cli/azd/cmd/auth_token.go | 18 +++---- cli/azd/cmd/auth_token_test.go | 38 ++++++++------- cli/azd/cmd/container.go | 3 -- cli/azd/cmd/deeper_coverage3_test.go | 14 +++--- cli/azd/cmd/monitor.go | 7 +-- .../docs/style-guidelines/new-azd-command.md | 2 +- cli/azd/pkg/account/credentials.go | 12 ++--- cli/azd/pkg/account/credentials_test.go | 39 ++++++++++----- cli/azd/pkg/account/subscriptions_manager.go | 13 ++--- .../pkg/account/subscriptions_manager_test.go | 48 +++++++++++++++++++ .../pkg/azsdk/storage/storage_blob_client.go | 11 ++--- .../storage_blob_client_coverage_test.go | 8 ++-- .../azsdk/storage/storage_blob_client_test.go | 35 ++++++++++---- cli/azd/pkg/environment/manager_test.go | 23 +++++---- .../current_principal_id_provider.go | 4 +- .../current_principal_id_provider_test.go | 19 +++----- cli/azd/test/functional/remote_state_test.go | 21 ++++---- .../test/mocks/mockaccount/mock_manager.go | 2 +- 18 files changed, 199 insertions(+), 118 deletions(-) diff --git a/cli/azd/cmd/auth_token.go b/cli/azd/cmd/auth_token.go index 481ab3f6ae6..20e279d5812 100644 --- a/cli/azd/cmd/auth_token.go +++ b/cli/azd/cmd/auth_token.go @@ -59,7 +59,7 @@ type authTokenAction struct { formatter output.Formatter writer io.Writer envResolver environment.EnvironmentResolver - subResolver account.SubscriptionTenantResolver + subResolver account.SubscriptionResolver flags *authTokenFlags cloud *cloud.Cloud } @@ -70,7 +70,7 @@ func newAuthTokenAction( writer io.Writer, flags *authTokenFlags, envResolver environment.EnvironmentResolver, - subResolver account.SubscriptionTenantResolver, + subResolver account.SubscriptionResolver, cloud *cloud.Cloud, ) actions.Action { return &authTokenAction{ @@ -87,7 +87,7 @@ func newAuthTokenAction( func getTenantIdFromAzdEnv( ctx context.Context, envResolver environment.EnvironmentResolver, - subResolver account.SubscriptionTenantResolver) (tenantId string, err error) { + subResolver account.SubscriptionResolver) (tenantId string, err error) { azdEnv, err := envResolver(ctx) if err != nil { // No azd env, return empty tenantId @@ -100,20 +100,21 @@ func getTenantIdFromAzdEnv( return tenantId, nil } - tenantId, err = subResolver.LookupTenant(ctx, subIdAtAzdEnv) + subscription, err := subResolver.GetSubscription(ctx, subIdAtAzdEnv) if err != nil { return tenantId, fmt.Errorf( - "resolving the Azure Directory from azd environment (%s): %w", + "getting the subscription from azd environment (%s): %w", azdEnv.Name(), err) } + tenantId = subscription.UserAccessTenantId return tenantId, nil } func getTenantIdFromEnv( ctx context.Context, - subResolver account.SubscriptionTenantResolver) (tenantId string, err error) { + subResolver account.SubscriptionResolver) (tenantId string, err error) { subIdAtSysEnv, found := os.LookupEnv(environment.SubscriptionIdEnvVarName) if !found { @@ -121,11 +122,12 @@ func getTenantIdFromEnv( return tenantId, nil } - tenantId, err = subResolver.LookupTenant(ctx, subIdAtSysEnv) + subscription, err := subResolver.GetSubscription(ctx, subIdAtSysEnv) if err != nil { return tenantId, fmt.Errorf( - "resolving the Azure Directory from system environment (%s): %w", environment.SubscriptionIdEnvVarName, err) + "getting the subscription from system environment (%s): %w", environment.SubscriptionIdEnvVarName, err) } + tenantId = subscription.UserAccessTenantId return tenantId, nil } diff --git a/cli/azd/cmd/auth_token_test.go b/cli/azd/cmd/auth_token_test.go index 215ce1d1512..3a2a9633df4 100644 --- a/cli/azd/cmd/auth_token_test.go +++ b/cli/azd/cmd/auth_token_test.go @@ -16,6 +16,7 @@ import ( "github.com/Azure/azure-sdk-for-go/sdk/azcore" "github.com/Azure/azure-sdk-for-go/sdk/azcore/policy" "github.com/azure/azure-dev/cli/azd/internal" + "github.com/azure/azure-dev/cli/azd/pkg/account" "github.com/azure/azure-dev/cli/azd/pkg/auth" "github.com/azure/azure-dev/cli/azd/pkg/cloud" "github.com/azure/azure-dev/cli/azd/pkg/contracts" @@ -50,7 +51,7 @@ func TestAuthToken(t *testing.T) { func(ctx context.Context) (*environment.Environment, error) { return nil, fmt.Errorf("not an azd env directory") }, - &mockSubscriptionTenantResolver{}, + &mockSubscriptionResolver{}, cloud.AzurePublic(), ) @@ -86,7 +87,7 @@ func TestAuthToken_DefaultUnformattedOutput(t *testing.T) { func(ctx context.Context) (*environment.Environment, error) { return nil, fmt.Errorf("not an azd env directory") }, - &mockSubscriptionTenantResolver{}, + &mockSubscriptionResolver{}, cloud.AzurePublic(), ) @@ -120,7 +121,7 @@ func TestAuthTokenSysEnv(t *testing.T) { func(ctx context.Context) (*environment.Environment, error) { return nil, fmt.Errorf("not an azd env directory") }, - &mockSubscriptionTenantResolver{ + &mockSubscriptionResolver{ TenantId: expectedTenant, }, cloud.AzurePublic(), @@ -168,7 +169,7 @@ func TestAuthTokenSysEnvError(t *testing.T) { func(ctx context.Context) (*environment.Environment, error) { return nil, fmt.Errorf("not an azd env directory") }, - &mockSubscriptionTenantResolver{ + &mockSubscriptionResolver{ Err: errors.New(expectedError), }, cloud.AzurePublic(), @@ -179,7 +180,7 @@ func TestAuthTokenSysEnvError(t *testing.T) { t, err, fmt.Sprintf( - "resolving the Azure Directory from system environment (%s): %s", + "getting the subscription from system environment (%s): %s", environment.SubscriptionIdEnvVarName, expectedError), ) @@ -212,7 +213,7 @@ func TestAuthTokenAzdEnvError(t *testing.T) { environment.SubscriptionIdEnvVarName: expectedSubId, }), nil }, - &mockSubscriptionTenantResolver{ + &mockSubscriptionResolver{ Err: errors.New(expectedError), }, cloud.AzurePublic(), @@ -223,7 +224,7 @@ func TestAuthTokenAzdEnvError(t *testing.T) { t, err, fmt.Sprintf( - "resolving the Azure Directory from azd environment (%s): %s", + "getting the subscription from azd environment (%s): %s", expectedEnvName, expectedError), ) @@ -253,7 +254,7 @@ func TestAuthTokenAzdEnv(t *testing.T) { environment.SubscriptionIdEnvVarName: "sub-id", }), nil }, - &mockSubscriptionTenantResolver{ + &mockSubscriptionResolver{ TenantId: expectedTenant, }, cloud.AzurePublic(), @@ -294,7 +295,7 @@ func TestAuthTokenAzdEnvWithEmpty(t *testing.T) { environment.SubscriptionIdEnvVarName: "", }), nil }, - &mockSubscriptionTenantResolver{ + &mockSubscriptionResolver{ TenantId: expectedTenant, }, cloud.AzurePublic(), @@ -333,7 +334,7 @@ func TestAuthTokenCustomScopes(t *testing.T) { func(ctx context.Context) (*environment.Environment, error) { return nil, fmt.Errorf("not an azd env directory") }, - &mockSubscriptionTenantResolver{}, + &mockSubscriptionResolver{}, cloud.AzurePublic(), ) @@ -355,7 +356,7 @@ func TestAuthTokenFailure(t *testing.T) { func(ctx context.Context) (*environment.Environment, error) { return nil, fmt.Errorf("not an azd env directory") }, - &mockSubscriptionTenantResolver{}, + &mockSubscriptionResolver{}, cloud.AzurePublic(), ) @@ -380,16 +381,21 @@ func credentialProviderForTokenFn( } -type mockSubscriptionTenantResolver struct { +type mockSubscriptionResolver struct { TenantId string Err error } -func (m *mockSubscriptionTenantResolver) LookupTenant( - ctx context.Context, subscriptionId string) (tenantId string, err error) { +func (m *mockSubscriptionResolver) GetSubscription( + ctx context.Context, subscriptionId string, +) (*account.Subscription, error) { if m.Err != nil { - return "", m.Err + return nil, m.Err } - return m.TenantId, nil + return &account.Subscription{ + Id: subscriptionId, + TenantId: "resource-" + m.TenantId, + UserAccessTenantId: m.TenantId, + }, nil } diff --git a/cli/azd/cmd/container.go b/cli/azd/cmd/container.go index 693e6bb7a77..da236a99e6f 100644 --- a/cli/azd/cmd/container.go +++ b/cli/azd/cmd/container.go @@ -700,9 +700,6 @@ func registerCommonDependencies(container *ioc.NestedContainer) { }) }) - container.MustRegisterSingleton(func(subManager *account.SubscriptionsManager) account.SubscriptionTenantResolver { - return subManager - }) container.MustRegisterSingleton(func(subManager *account.SubscriptionsManager) account.SubscriptionResolver { return subManager }) diff --git a/cli/azd/cmd/deeper_coverage3_test.go b/cli/azd/cmd/deeper_coverage3_test.go index d1a6f4076dc..853725562b9 100644 --- a/cli/azd/cmd/deeper_coverage3_test.go +++ b/cli/azd/cmd/deeper_coverage3_test.go @@ -146,7 +146,9 @@ func (m *mockSubTenantResolver) LookupTenant(ctx context.Context, subscriptionId return args.String(0), args.Error(1) } -func (m *mockSubTenantResolver) GetSubscription(ctx context.Context, subscriptionId string) (*account.Subscription, error) { +func (m *mockSubTenantResolver) GetSubscription( + ctx context.Context, subscriptionId string, +) (*account.Subscription, error) { tenantId, err := m.LookupTenant(ctx, subscriptionId) if err != nil { return nil, err @@ -163,11 +165,9 @@ type staticSubscriptionResolver struct { subscription *account.Subscription } -func (s *staticSubscriptionResolver) LookupTenant(ctx context.Context, subscriptionId string) (string, error) { - return s.subscription.UserAccessTenantId, nil -} - -func (s *staticSubscriptionResolver) GetSubscription(ctx context.Context, subscriptionId string) (*account.Subscription, error) { +func (s *staticSubscriptionResolver) GetSubscription( + ctx context.Context, subscriptionId string, +) (*account.Subscription, error) { return s.subscription, nil } @@ -443,7 +443,7 @@ func Test_EnvSetSecretAction_LookupTenantError(t *testing.T) { _, err := action.Run(t.Context()) require.Error(t, err) - assert.Contains(t, err.Error(), "looking up tenant for subscription") + assert.Contains(t, err.Error(), "getting subscription") } func Test_EnvSetSecretAction_ListVaultsError(t *testing.T) { diff --git a/cli/azd/cmd/monitor.go b/cli/azd/cmd/monitor.go index a1eea96ff58..7192d441f65 100644 --- a/cli/azd/cmd/monitor.go +++ b/cli/azd/cmd/monitor.go @@ -62,7 +62,7 @@ func newMonitorCmd() *cobra.Command { type monitorAction struct { azdCtx *azdcontext.AzdContext env *environment.Environment - subResolver account.SubscriptionTenantResolver + subResolver account.SubscriptionResolver resourceManager infra.ResourceManager resourceService *azapi.ResourceService console input.Console @@ -74,7 +74,7 @@ type monitorAction struct { func newMonitorAction( azdCtx *azdcontext.AzdContext, env *environment.Environment, - subResolver account.SubscriptionTenantResolver, + subResolver account.SubscriptionResolver, resourceManager infra.ResourceManager, resourceService *azapi.ResourceService, console input.Console, @@ -152,10 +152,11 @@ func (m *monitorAction) Run(ctx context.Context) (*actions.ActionResult, error) } } - tenantId, err := m.subResolver.LookupTenant(ctx, m.env.GetSubscriptionId()) + subscription, err := m.subResolver.GetSubscription(ctx, m.env.GetSubscriptionId()) if err != nil { return nil, err } + tenantId := subscription.UserAccessTenantId for _, insightsResource := range insightsResources { if m.flags.monitorLive { diff --git a/cli/azd/docs/style-guidelines/new-azd-command.md b/cli/azd/docs/style-guidelines/new-azd-command.md index fa9e025b213..5bff9de53e3 100644 --- a/cli/azd/docs/style-guidelines/new-azd-command.md +++ b/cli/azd/docs/style-guidelines/new-azd-command.md @@ -894,7 +894,7 @@ env *environment.Environment // Azure services accountManager account.Manager -subscriptionResolver account.SubscriptionTenantResolver +subscriptionResolver account.SubscriptionResolver resourceManager infra.ResourceManager resourceService *azapi.ResourceService diff --git a/cli/azd/pkg/account/credentials.go b/cli/azd/pkg/account/credentials.go index 311f9e0c312..6baf2b5d5c5 100644 --- a/cli/azd/pkg/account/credentials.go +++ b/cli/azd/pkg/account/credentials.go @@ -20,19 +20,18 @@ var ( ) // SubscriptionCredentialProvider provides an [azcore.TokenCredential] configured -// to use the tenant id that corresponds to the tenant the given subscription -// is located in. +// to use the access tenant required by the current account for the given subscription. type SubscriptionCredentialProvider interface { CredentialForSubscription(ctx context.Context, subscriptionId string) (azcore.TokenCredential, error) } type subscriptionCredentialProvider struct { credProvider auth.MultiTenantCredentialProvider - subResolver SubscriptionTenantResolver + subResolver SubscriptionResolver } func NewSubscriptionCredentialProvider( - subResolver SubscriptionTenantResolver, + subResolver SubscriptionResolver, credProvider auth.MultiTenantCredentialProvider, ) SubscriptionCredentialProvider { return &subscriptionCredentialProvider{ @@ -45,9 +44,9 @@ func (p *subscriptionCredentialProvider) CredentialForSubscription( ctx context.Context, subscriptionId string, ) (azcore.TokenCredential, error) { - tenantId, err := p.subResolver.LookupTenant(ctx, subscriptionId) + subscription, err := p.subResolver.GetSubscription(ctx, subscriptionId) if err != nil { - // If we can't resolve the tenant for this subscription, it might be because: + // If we can't resolve the subscription for this ID, it might be because: // 1. User manually set AZURE_SUBSCRIPTION_ID in .env // 2. User called `azd env set AZURE_SUBSCRIPTION_ID` instead of selecting from azd's cache // In these cases, suggest they also set AZURE_TENANT_ID @@ -60,6 +59,7 @@ func (p *subscriptionCredentialProvider) CredentialForSubscription( err, ) } + tenantId := subscription.UserAccessTenantId cred, err := p.credProvider.GetTokenCredential(ctx, tenantId) if err != nil { diff --git a/cli/azd/pkg/account/credentials_test.go b/cli/azd/pkg/account/credentials_test.go index 3a7bab4f35f..c99eefdedf3 100644 --- a/cli/azd/pkg/account/credentials_test.go +++ b/cli/azd/pkg/account/credentials_test.go @@ -35,12 +35,15 @@ func TestSubscriptionCredentialProvider(t *testing.T) { } provider := NewSubscriptionCredentialProvider( - subscriptionTenantResolverFunc(func(ctx context.Context, subscriptionId string) (string, error) { + subscriptionResolverFunc(func(ctx context.Context, subscriptionId string) (*Subscription, error) { if tenantId, has := subToTenant[subscriptionId]; has { - return tenantId, nil - } else { - return "", errors.New("unknown subscription") + return &Subscription{ + Id: subscriptionId, + TenantId: "resource-" + tenantId, + UserAccessTenantId: tenantId, + }, nil } + return nil, errors.New("unknown subscription") }), multiTenantCredentialProviderFunc(func(ctx context.Context, tenantId string) (azcore.TokenCredential, error) { if credential, has := tenantToCred[tenantId]; has { @@ -75,8 +78,12 @@ func TestSubscriptionCredentialProvider_AADSTSErrors(t *testing.T) { t.Run("AADSTS70043_WithoutExistingSuggestion", func(t *testing.T) { provider := NewSubscriptionCredentialProvider( - subscriptionTenantResolverFunc(func(ctx context.Context, subId string) (string, error) { - return tenantId, nil + subscriptionResolverFunc(func(ctx context.Context, subId string) (*Subscription, error) { + return &Subscription{ + Id: subId, + TenantId: "resource-" + tenantId, + UserAccessTenantId: tenantId, + }, nil }), multiTenantCredentialProviderFunc(func(ctx context.Context, tid string) (azcore.TokenCredential, error) { return nil, errors.New("AADSTS70043: The refresh token has expired") @@ -100,8 +107,12 @@ func TestSubscriptionCredentialProvider_AADSTSErrors(t *testing.T) { t.Run("AADSTS700082_RefreshTokenExpired", func(t *testing.T) { provider := NewSubscriptionCredentialProvider( - subscriptionTenantResolverFunc(func(ctx context.Context, subId string) (string, error) { - return tenantId, nil + subscriptionResolverFunc(func(ctx context.Context, subId string) (*Subscription, error) { + return &Subscription{ + Id: subId, + TenantId: "resource-" + tenantId, + UserAccessTenantId: tenantId, + }, nil }), multiTenantCredentialProviderFunc(func(ctx context.Context, tid string) (azcore.TokenCredential, error) { return nil, errors.New("AADSTS700082: The refresh token has expired") @@ -125,8 +136,8 @@ func TestSubscriptionCredentialProvider_AADSTSErrors(t *testing.T) { t.Run("TenantLookupFailure_EnhancedError", func(t *testing.T) { provider := NewSubscriptionCredentialProvider( - subscriptionTenantResolverFunc(func(ctx context.Context, subId string) (string, error) { - return "", errors.New("failed to resolve tenant") + subscriptionResolverFunc(func(ctx context.Context, subId string) (*Subscription, error) { + return nil, errors.New("failed to resolve tenant") }), multiTenantCredentialProviderFunc(func(ctx context.Context, tid string) (azcore.TokenCredential, error) { return &dummyCredential{}, nil @@ -140,10 +151,12 @@ func TestSubscriptionCredentialProvider_AADSTSErrors(t *testing.T) { }) } -// subscriptionTenantResolverFunc implements [SubscriptionTenantResolver] using a provided function. -type subscriptionTenantResolverFunc func(ctx context.Context, subscriptionId string) (string, error) +// subscriptionResolverFunc implements [SubscriptionResolver] using a provided function. +type subscriptionResolverFunc func(ctx context.Context, subscriptionId string) (*Subscription, error) -func (r subscriptionTenantResolverFunc) LookupTenant(ctx context.Context, subscriptionId string) (string, error) { +func (r subscriptionResolverFunc) GetSubscription( + ctx context.Context, subscriptionId string, +) (*Subscription, error) { return r(ctx, subscriptionId) } diff --git a/cli/azd/pkg/account/subscriptions_manager.go b/cli/azd/pkg/account/subscriptions_manager.go index 447bedf6438..de5685d0a9f 100644 --- a/cli/azd/pkg/account/subscriptions_manager.go +++ b/cli/azd/pkg/account/subscriptions_manager.go @@ -21,16 +21,9 @@ import ( "go.uber.org/multierr" ) -// SubscriptionTenantResolver allows resolving the correct tenant ID -// that allows the current account access to a given subscription. -type SubscriptionTenantResolver interface { - // Resolve the tenant ID required by the current account to access the given subscription. - LookupTenant(ctx context.Context, subscriptionId string) (tenantId string, err error) -} - -// SubscriptionResolver allows resolving both the access tenant and subscription details. +// SubscriptionResolver resolves subscription metadata for the current account, including both the resource tenant +// and the user access tenant. type SubscriptionResolver interface { - SubscriptionTenantResolver GetSubscription(ctx context.Context, subscriptionId string) (*Subscription, error) } @@ -214,6 +207,8 @@ func (m *SubscriptionsManager) getSubscriptions(ctx context.Context) (getSubscri }, nil } +// GetSubscription retrieves subscription metadata for the current account, including both the resource tenant +// and the tenant through which the current user can access it. func (m *SubscriptionsManager) GetSubscription(ctx context.Context, subscriptionId string) (*Subscription, error) { subscriptions, err := m.GetSubscriptions(ctx) if err != nil { diff --git a/cli/azd/pkg/account/subscriptions_manager_test.go b/cli/azd/pkg/account/subscriptions_manager_test.go index ba9de09bb17..b347ebcf763 100644 --- a/cli/azd/pkg/account/subscriptions_manager_test.go +++ b/cli/azd/pkg/account/subscriptions_manager_test.go @@ -177,6 +177,54 @@ func TestSubscriptionsManager_ListSubscriptions(t *testing.T) { } } +type staticSubCache struct { + subscriptions []Subscription +} + +func (c *staticSubCache) Load(ctx context.Context, key string) ([]Subscription, error) { + return c.subscriptions, nil +} + +func (c *staticSubCache) Save(ctx context.Context, key string, save []Subscription) error { + return nil +} + +func (c *staticSubCache) Merge(ctx context.Context, key string, save []Subscription) error { + return nil +} + +func (c *staticSubCache) Clear(ctx context.Context) error { + return nil +} + +func TestSubscriptionsManager_GetSubscription_PreservesTenantFields(t *testing.T) { + t.Parallel() + + subManager := &SubscriptionsManager{ + cache: &staticSubCache{ + subscriptions: []Subscription{ + { + Id: "sub-123", + Name: "Subscription 123", + TenantId: "resource-tenant", + UserAccessTenantId: "access-tenant", + }, + }, + }, + principalInfo: &principalInfoProviderMock{}, + console: mockinput.NewMockConsole(), + } + + subscription, err := subManager.GetSubscription(t.Context(), "sub-123") + require.NoError(t, err) + require.Equal(t, &Subscription{ + Id: "sub-123", + Name: "Subscription 123", + TenantId: "resource-tenant", + UserAccessTenantId: "access-tenant", + }, subscription) +} + func generateTenants(total int) []*armsubscriptions.TenantIDDescription { results := make([]*armsubscriptions.TenantIDDescription, 0, total) for i := 1; i <= total; i++ { diff --git a/cli/azd/pkg/azsdk/storage/storage_blob_client.go b/cli/azd/pkg/azsdk/storage/storage_blob_client.go index d20098d40a2..12244af3082 100644 --- a/cli/azd/pkg/azsdk/storage/storage_blob_client.go +++ b/cli/azd/pkg/azsdk/storage/storage_blob_client.go @@ -277,7 +277,7 @@ func NewBlobSdkClient( userConfigManager config.UserConfigManager, coreClientOptions *azcore.ClientOptions, cloud *cloud.Cloud, - tenantResolver account.SubscriptionTenantResolver, + subscriptionResolver account.SubscriptionResolver, ) (*azblob.Client, error) { blobOptions := &azblob.ClientOptions{ ClientOptions: *coreClientOptions, @@ -301,13 +301,12 @@ func NewBlobSdkClient( tenantId := "" if subscriptionId != "" { - // If a subscription ID is configured, resolve the tenant ID for that subscription - resolvedTenantId, err := tenantResolver.LookupTenant(context.Background(), subscriptionId) + // If a subscription ID is configured, resolve the subscription for the current account. + subscription, err := subscriptionResolver.GetSubscription(context.Background(), subscriptionId) if err != nil { - return nil, fmt.Errorf( - "failed to resolve tenant for subscription '%s': %w", subscriptionId, err) + return nil, fmt.Errorf("failed to get subscription '%s': %w", subscriptionId, err) } - tenantId = resolvedTenantId + tenantId = subscription.UserAccessTenantId } // Otherwise, use home tenant ID (empty string) diff --git a/cli/azd/pkg/azsdk/storage/storage_blob_client_coverage_test.go b/cli/azd/pkg/azsdk/storage/storage_blob_client_coverage_test.go index 94d6b8be972..f7b278770df 100644 --- a/cli/azd/pkg/azsdk/storage/storage_blob_client_coverage_test.go +++ b/cli/azd/pkg/azsdk/storage/storage_blob_client_coverage_test.go @@ -16,7 +16,7 @@ import ( func Test_NewBlobSdkClient_UsesCustomEndpoint(t *testing.T) { mockCredProvider := &mockMultiTenantCredentialProvider{} - mockTenantResolver := &mockSubscriptionTenantResolver{} + mockTenantResolver := &mockSubscriptionResolver{} mockCred := &mockTokenCredential{} mockConfigMgr := &mockUserConfigManager{} @@ -51,7 +51,7 @@ func Test_NewBlobSdkClient_UsesCustomEndpoint(t *testing.T) { func Test_NewBlobSdkClient_DefaultEndpointFromCloud(t *testing.T) { mockCredProvider := &mockMultiTenantCredentialProvider{} - mockTenantResolver := &mockSubscriptionTenantResolver{} + mockTenantResolver := &mockSubscriptionResolver{} mockCred := &mockTokenCredential{} mockConfigMgr := &mockUserConfigManager{} @@ -86,7 +86,7 @@ func Test_NewBlobSdkClient_DefaultEndpointFromCloud(t *testing.T) { func Test_NewBlobSdkClient_CredentialProviderError(t *testing.T) { mockCredProvider := &mockMultiTenantCredentialProvider{} - mockTenantResolver := &mockSubscriptionTenantResolver{} + mockTenantResolver := &mockSubscriptionResolver{} mockConfigMgr := &mockUserConfigManager{} accountCfg := &AccountConfig{ @@ -117,7 +117,7 @@ func Test_NewBlobSdkClient_CredentialProviderError(t *testing.T) { func Test_NewBlobSdkClient_EmptyDefaultSubscriptionIgnored(t *testing.T) { mockCredProvider := &mockMultiTenantCredentialProvider{} - mockTenantResolver := &mockSubscriptionTenantResolver{} + mockTenantResolver := &mockSubscriptionResolver{} mockCred := &mockTokenCredential{} mockConfigMgr := &mockUserConfigManager{} diff --git a/cli/azd/pkg/azsdk/storage/storage_blob_client_test.go b/cli/azd/pkg/azsdk/storage/storage_blob_client_test.go index 766e76b37e3..8611a8da056 100644 --- a/cli/azd/pkg/azsdk/storage/storage_blob_client_test.go +++ b/cli/azd/pkg/azsdk/storage/storage_blob_client_test.go @@ -31,19 +31,34 @@ func (m *mockMultiTenantCredentialProvider) GetTokenCredential( return args.Get(0).(azcore.TokenCredential), args.Error(1) } -// mockSubscriptionTenantResolver is a mock implementation for testing -type mockSubscriptionTenantResolver struct { +// mockSubscriptionResolver is a mock implementation for testing. +type mockSubscriptionResolver struct { mock.Mock } -var _ account.SubscriptionTenantResolver = (*mockSubscriptionTenantResolver)(nil) +var _ account.SubscriptionResolver = (*mockSubscriptionResolver)(nil) -func (m *mockSubscriptionTenantResolver) LookupTenant( +func (m *mockSubscriptionResolver) LookupTenant( ctx context.Context, subscriptionId string) (string, error) { args := m.Called(ctx, subscriptionId) return args.String(0), args.Error(1) } +func (m *mockSubscriptionResolver) GetSubscription( + ctx context.Context, subscriptionId string, +) (*account.Subscription, error) { + tenantId, err := m.LookupTenant(ctx, subscriptionId) + if err != nil { + return nil, err + } + + return &account.Subscription{ + Id: subscriptionId, + TenantId: "resource-" + tenantId, + UserAccessTenantId: tenantId, + }, nil +} + // mockTokenCredential is a minimal mock implementation for testing type mockTokenCredential struct { mock.Mock @@ -77,7 +92,7 @@ func (m *mockUserConfigManager) Load() (config.Config, error) { func Test_NewBlobSdkClient_UsesHomeTenantWhenNoSubscriptionId(t *testing.T) { mockCredProvider := &mockMultiTenantCredentialProvider{} - mockTenantResolver := &mockSubscriptionTenantResolver{} + mockTenantResolver := &mockSubscriptionResolver{} mockCred := &mockTokenCredential{} mockConfigMgr := &mockUserConfigManager{} @@ -114,7 +129,7 @@ func Test_NewBlobSdkClient_UsesHomeTenantWhenNoSubscriptionId(t *testing.T) { func Test_NewBlobSdkClient_ResolvesTenantWhenSubscriptionIdProvided(t *testing.T) { mockCredProvider := &mockMultiTenantCredentialProvider{} - mockTenantResolver := &mockSubscriptionTenantResolver{} + mockTenantResolver := &mockSubscriptionResolver{} mockCred := &mockTokenCredential{} mockConfigMgr := &mockUserConfigManager{} @@ -155,7 +170,7 @@ func Test_NewBlobSdkClient_ResolvesTenantWhenSubscriptionIdProvided(t *testing.T func Test_NewBlobSdkClient_ReturnsErrorWhenTenantResolutionFails(t *testing.T) { mockCredProvider := &mockMultiTenantCredentialProvider{} - mockTenantResolver := &mockSubscriptionTenantResolver{} + mockTenantResolver := &mockSubscriptionResolver{} mockConfigMgr := &mockUserConfigManager{} testSubscriptionId := "test-subscription-id" @@ -183,13 +198,13 @@ func Test_NewBlobSdkClient_ReturnsErrorWhenTenantResolutionFails(t *testing.T) { require.Error(t, err) require.Nil(t, client) - require.Contains(t, err.Error(), "failed to resolve tenant for subscription") + require.Contains(t, err.Error(), "failed to get subscription") mockTenantResolver.AssertExpectations(t) } func Test_NewBlobSdkClient_FallsBackToDefaultSubscriptionFromUserConfig(t *testing.T) { mockCredProvider := &mockMultiTenantCredentialProvider{} - mockTenantResolver := &mockSubscriptionTenantResolver{} + mockTenantResolver := &mockSubscriptionResolver{} mockCred := &mockTokenCredential{} mockConfigMgr := &mockUserConfigManager{} @@ -236,7 +251,7 @@ func Test_NewBlobSdkClient_FallsBackToDefaultSubscriptionFromUserConfig(t *testi func Test_NewBlobSdkClient_UsesHomeTenantWhenUserConfigLoadFails(t *testing.T) { mockCredProvider := &mockMultiTenantCredentialProvider{} - mockTenantResolver := &mockSubscriptionTenantResolver{} + mockTenantResolver := &mockSubscriptionResolver{} mockCred := &mockTokenCredential{} mockConfigMgr := &mockUserConfigManager{} diff --git a/cli/azd/pkg/environment/manager_test.go b/cli/azd/pkg/environment/manager_test.go index 9f256e33774..a2436261559 100644 --- a/cli/azd/pkg/environment/manager_test.go +++ b/cli/azd/pkg/environment/manager_test.go @@ -384,9 +384,9 @@ func registerContainerComponents(t *testing.T, mockContext *mocks.MockContext) { return mockContext.CoreClientOptions }) - // Register a mock SubscriptionTenantResolver for tests - mockContext.Container.MustRegisterSingleton(func() account.SubscriptionTenantResolver { - return &mockSubscriptionTenantResolver{} + // Register a mock subscription resolver for tests. + mockContext.Container.MustRegisterSingleton(func() account.SubscriptionResolver { + return &mockSubscriptionResolver{} }) mockContext.Container.MustRegisterSingleton(func() config.UserConfigManager { @@ -419,14 +419,19 @@ func registerContainerComponents(t *testing.T, mockContext *mocks.MockContext) { }) } -// mockSubscriptionTenantResolver is a simple mock for testing -type mockSubscriptionTenantResolver struct{} +// mockSubscriptionResolver is a simple mock for testing. +type mockSubscriptionResolver struct{} -var _ account.SubscriptionTenantResolver = (*mockSubscriptionTenantResolver)(nil) +var _ account.SubscriptionResolver = (*mockSubscriptionResolver)(nil) -func (m *mockSubscriptionTenantResolver) LookupTenant(ctx context.Context, subscriptionId string) (string, error) { - // For tests, just return empty string (home tenant) - return "", nil +func (m *mockSubscriptionResolver) GetSubscription( + ctx context.Context, subscriptionId string, +) (*account.Subscription, error) { + return &account.Subscription{ + Id: subscriptionId, + TenantId: "", + UserAccessTenantId: "", + }, nil } type MockDataStore struct { diff --git a/cli/azd/pkg/infra/provisioning/current_principal_id_provider.go b/cli/azd/pkg/infra/provisioning/current_principal_id_provider.go index 493bca7cd09..2b20c850581 100644 --- a/cli/azd/pkg/infra/provisioning/current_principal_id_provider.go +++ b/cli/azd/pkg/infra/provisioning/current_principal_id_provider.go @@ -44,12 +44,12 @@ type principalIDProvider struct { func (p *principalIDProvider) CurrentPrincipalId(ctx context.Context) (string, error) { subscriptionId := p.env.GetSubscriptionId() - sub, err := p.subResolver.GetSubscription(ctx, subscriptionId) + subscription, err := p.subResolver.GetSubscription(ctx, subscriptionId) if err != nil { return "", fmt.Errorf("getting subscription %s: %w", subscriptionId, err) } - principalId, err := azureutil.GetCurrentPrincipalId(ctx, p.userProfileService, sub.TenantId) + principalId, err := azureutil.GetCurrentPrincipalId(ctx, p.userProfileService, subscription.TenantId) if err != nil { return "", fmt.Errorf("fetching current user information: %w", err) } diff --git a/cli/azd/pkg/infra/provisioning/current_principal_id_provider_test.go b/cli/azd/pkg/infra/provisioning/current_principal_id_provider_test.go index 62837ba1ac9..1f9c53e2c20 100644 --- a/cli/azd/pkg/infra/provisioning/current_principal_id_provider_test.go +++ b/cli/azd/pkg/infra/provisioning/current_principal_id_provider_test.go @@ -19,18 +19,14 @@ import ( ) type fakeSubscriptionResolver struct { - subscription *account.Subscription - lookupTenantCalls int - getSubscriptionCalls int + subscription *account.Subscription + getCalls int } -func (f *fakeSubscriptionResolver) LookupTenant(ctx context.Context, subscriptionId string) (string, error) { - f.lookupTenantCalls++ - return "home-tenant", nil -} - -func (f *fakeSubscriptionResolver) GetSubscription(ctx context.Context, subscriptionId string) (*account.Subscription, error) { - f.getSubscriptionCalls++ +func (f *fakeSubscriptionResolver) GetSubscription( + ctx context.Context, subscriptionId string, +) (*account.Subscription, error) { + f.getCalls++ return f.subscription, nil } @@ -78,6 +74,5 @@ func TestPrincipalIDProvider_CurrentPrincipalIdUsesSubscriptionTenant(t *testing principalId, err := provider.CurrentPrincipalId(t.Context()) require.NoError(t, err) require.Equal(t, "this-is-a-test", principalId) - require.Equal(t, 1, resolver.getSubscriptionCalls) - require.Zero(t, resolver.lookupTenantCalls) + require.Equal(t, 1, resolver.getCalls) } diff --git a/cli/azd/test/functional/remote_state_test.go b/cli/azd/test/functional/remote_state_test.go index 58f0dc3f4df..45cbc3abf8b 100644 --- a/cli/azd/test/functional/remote_state_test.go +++ b/cli/azd/test/functional/remote_state_test.go @@ -93,8 +93,8 @@ func createBlobClient( ) require.NoError(t, err) - // Create a mock SubscriptionTenantResolver that returns empty tenant (home tenant) - tenantResolver := &mockTenantResolver{} + // Create a mock SubscriptionResolver that returns an empty user-access tenant. + tenantResolver := &mockSubscriptionResolver{} userConfigManager := config.NewUserConfigManager(fileConfigManager) sdkClient, err := storage.NewBlobSdkClient( @@ -111,14 +111,19 @@ func createBlobClient( return storage.NewBlobClient(storageConfig, sdkClient) } -// mockTenantResolver is a simple mock implementation for testing -type mockTenantResolver struct{} +// mockSubscriptionResolver is a simple mock implementation for testing. +type mockSubscriptionResolver struct{} -var _ account.SubscriptionTenantResolver = (*mockTenantResolver)(nil) +var _ account.SubscriptionResolver = (*mockSubscriptionResolver)(nil) -func (m *mockTenantResolver) LookupTenant(ctx context.Context, subscriptionId string) (string, error) { - // For tests, just return empty string (home tenant) - return "", nil +func (m *mockSubscriptionResolver) GetSubscription( + ctx context.Context, subscriptionId string, +) (*account.Subscription, error) { + return &account.Subscription{ + Id: subscriptionId, + TenantId: "", + UserAccessTenantId: "", + }, nil } type remoteStateTestFunc func(storageConfig *storage.AccountConfig) diff --git a/cli/azd/test/mocks/mockaccount/mock_manager.go b/cli/azd/test/mocks/mockaccount/mock_manager.go index fcbf8e79a87..1b1f21a9c6a 100644 --- a/cli/azd/test/mocks/mockaccount/mock_manager.go +++ b/cli/azd/test/mocks/mockaccount/mock_manager.go @@ -100,7 +100,7 @@ func (a *MockAccountManager) SetDefaultLocation( return nil, nil } -// SubscriptionTenantResolverFunc implements [account.SubscriptionCredentialProvider] using the provided function. +// SubscriptionCredentialProviderFunc implements [account.SubscriptionCredentialProvider] using the provided function. type SubscriptionCredentialProviderFunc func(ctx context.Context, subscriptionId string) (azcore.TokenCredential, error) func (f SubscriptionCredentialProviderFunc) CredentialForSubscription( From 5da52788d5228228e65a9497bd278c4ff67a4f62 Mon Sep 17 00:00:00 2001 From: Wei Lim Date: Mon, 6 Apr 2026 16:48:40 -0700 Subject: [PATCH 4/6] Fix lint in test JWT fixtures Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- cli/azd/cmd/deeper_coverage3_test.go | 5 +-- cli/azd/pkg/azureutil/principal_test.go | 10 +++--- .../current_principal_id_provider_test.go | 5 +-- cli/azd/test/mocks/jwt.go | 35 +++++++++++++++++++ 4 files changed, 47 insertions(+), 8 deletions(-) create mode 100644 cli/azd/test/mocks/jwt.go diff --git a/cli/azd/cmd/deeper_coverage3_test.go b/cli/azd/cmd/deeper_coverage3_test.go index 853725562b9..21cbc5b3ad7 100644 --- a/cli/azd/cmd/deeper_coverage3_test.go +++ b/cli/azd/cmd/deeper_coverage3_test.go @@ -1011,8 +1011,9 @@ func Test_EnvSetSecretAction_UsesResourceTenantForKeyVaultAndPrincipalId(t *test "resource-tenant": { GetTokenFn: func(ctx context.Context, options policy.TokenRequestOptions) (azcore.AccessToken, error) { return azcore.AccessToken{ - // cspell:disable-next-line - Token: "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJvaWQiOiJ0aGlzLWlzLWEtdGVzdCJ9.vrKZx2J7-hsydI4rzdFVHqU1S6lHqLT95VSPx2RfQ04", + Token: mocks.CreateJwtToken(t, map[string]string{ + "oid": "this-is-a-test", + }), ExpiresOn: time.Now().Add(time.Hour), }, nil }, diff --git a/cli/azd/pkg/azureutil/principal_test.go b/cli/azd/pkg/azureutil/principal_test.go index d3730b8d188..3a3043952e4 100644 --- a/cli/azd/pkg/azureutil/principal_test.go +++ b/cli/azd/pkg/azureutil/principal_test.go @@ -29,8 +29,9 @@ func TestGetCurrentPrincipalId_PrefersOidFromAccessToken(t *testing.T) { "resource-tenant": { GetTokenFn: func(ctx context.Context, options policy.TokenRequestOptions) (azcore.AccessToken, error) { return azcore.AccessToken{ - // cspell:disable-next-line - Token: "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJvaWQiOiJ0aGlzLWlzLWEtdGVzdCJ9.vrKZx2J7-hsydI4rzdFVHqU1S6lHqLT95VSPx2RfQ04", + Token: mocks.CreateJwtToken(t, map[string]string{ + "oid": "this-is-a-test", + }), ExpiresOn: time.Now().Add(time.Hour), }, nil }, @@ -66,8 +67,9 @@ func TestGetCurrentPrincipalId_FallsBackToGraphWhenOidMissing(t *testing.T) { "resource-tenant": { GetTokenFn: func(ctx context.Context, options policy.TokenRequestOptions) (azcore.AccessToken, error) { return azcore.AccessToken{ - // cspell:disable-next-line - Token: "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJ0ZXN0IjoiZmFpbCJ9.0Stzv5ZHG96ss-0_AnANqZfVLoULCtivJCE8AVWFZi8", + Token: mocks.CreateJwtToken(t, map[string]string{ + "test": "fail", + }), ExpiresOn: time.Now().Add(time.Hour), }, nil }, diff --git a/cli/azd/pkg/infra/provisioning/current_principal_id_provider_test.go b/cli/azd/pkg/infra/provisioning/current_principal_id_provider_test.go index 1f9c53e2c20..2476470474c 100644 --- a/cli/azd/pkg/infra/provisioning/current_principal_id_provider_test.go +++ b/cli/azd/pkg/infra/provisioning/current_principal_id_provider_test.go @@ -40,8 +40,9 @@ func TestPrincipalIDProvider_CurrentPrincipalIdUsesSubscriptionTenant(t *testing "resource-tenant": { GetTokenFn: func(ctx context.Context, options policy.TokenRequestOptions) (azcore.AccessToken, error) { return azcore.AccessToken{ - // cspell:disable-next-line - Token: "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJvaWQiOiJ0aGlzLWlzLWEtdGVzdCJ9.vrKZx2J7-hsydI4rzdFVHqU1S6lHqLT95VSPx2RfQ04", + Token: mocks.CreateJwtToken(t, map[string]string{ + "oid": "this-is-a-test", + }), ExpiresOn: time.Now().Add(time.Hour), }, nil }, diff --git a/cli/azd/test/mocks/jwt.go b/cli/azd/test/mocks/jwt.go new file mode 100644 index 00000000000..d0b63804b36 --- /dev/null +++ b/cli/azd/test/mocks/jwt.go @@ -0,0 +1,35 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +package mocks + +import ( + "encoding/base64" + "encoding/json" + "strings" + "testing" +) + +// CreateJwtToken creates a JWT-like token for tests without using hardcoded token literals. +func CreateJwtToken(t testing.TB, claims any) string { + t.Helper() + + header, err := json.Marshal(map[string]string{ + "alg": "none", + "typ": "JWT", + }) + if err != nil { + t.Fatalf("marshaling JWT header: %v", err) + } + + payload, err := json.Marshal(claims) + if err != nil { + t.Fatalf("marshaling JWT claims: %v", err) + } + + return strings.Join([]string{ + base64.RawURLEncoding.EncodeToString(header), + base64.RawURLEncoding.EncodeToString(payload), + "signature", + }, ".") +} From b717622d4b44cbb0c6932b9d494dec895ee83e01 Mon Sep 17 00:00:00 2001 From: Wei Lim Date: Mon, 6 Apr 2026 16:51:23 -0700 Subject: [PATCH 5/6] Address PR review feedback Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- cli/azd/cmd/deeper_coverage3_test.go | 60 ++++++++++++++-------------- cli/azd/pkg/azureutil/principal.go | 4 +- 2 files changed, 32 insertions(+), 32 deletions(-) diff --git a/cli/azd/cmd/deeper_coverage3_test.go b/cli/azd/cmd/deeper_coverage3_test.go index 21cbc5b3ad7..35a80405a6b 100644 --- a/cli/azd/cmd/deeper_coverage3_test.go +++ b/cli/azd/cmd/deeper_coverage3_test.go @@ -137,28 +137,28 @@ func (m *mockPrompter) PromptResourceGroupFrom( return args.String(0), args.Error(1) } -type mockSubTenantResolver struct { +type mockEnvSetSecretSubscriptionResolver struct { mock.Mock } -func (m *mockSubTenantResolver) LookupTenant(ctx context.Context, subscriptionId string) (string, error) { - args := m.Called(ctx, subscriptionId) - return args.String(0), args.Error(1) -} - -func (m *mockSubTenantResolver) GetSubscription( +func (m *mockEnvSetSecretSubscriptionResolver) GetSubscription( ctx context.Context, subscriptionId string, ) (*account.Subscription, error) { - tenantId, err := m.LookupTenant(ctx, subscriptionId) - if err != nil { - return nil, err + args := m.Called(ctx, subscriptionId) + + if subscription, ok := args.Get(0).(*account.Subscription); ok { + return subscription, args.Error(1) + } + + if tenantId, ok := args.Get(0).(string); ok { + return &account.Subscription{ + Id: subscriptionId, + TenantId: tenantId, + UserAccessTenantId: tenantId, + }, args.Error(1) } - return &account.Subscription{ - Id: subscriptionId, - TenantId: tenantId, - UserAccessTenantId: tenantId, - }, nil + return nil, args.Error(1) } type staticSubscriptionResolver struct { @@ -199,7 +199,7 @@ func newTestEnvSetSecretAction( projectConfig *project.ProjectConfig, kvService keyvault.KeyVaultService, prompter *mockPrompter, - subResolver *mockSubTenantResolver, + subResolver *mockEnvSetSecretSubscriptionResolver, ) *envSetSecretAction { if projectConfig == nil { projectConfig = &project.ProjectConfig{ @@ -435,8 +435,8 @@ func Test_EnvSetSecretAction_LookupTenantError(t *testing.T) { prompter.On("PromptSubscription", mock.Anything, mock.Anything). Return("sub-123", nil) - resolver := &mockSubTenantResolver{} - resolver.On("LookupTenant", mock.Anything, "sub-123"). + resolver := &mockEnvSetSecretSubscriptionResolver{} + resolver.On("GetSubscription", mock.Anything, "sub-123"). Return("", fmt.Errorf("tenant not found")) action := newTestEnvSetSecretAction(console, env, nil, []string{"mySecret"}, nil, nil, prompter, resolver) @@ -459,8 +459,8 @@ func Test_EnvSetSecretAction_ListVaultsError(t *testing.T) { prompter.On("PromptSubscription", mock.Anything, mock.Anything). Return("sub-123", nil) - resolver := &mockSubTenantResolver{} - resolver.On("LookupTenant", mock.Anything, "sub-123"). + resolver := &mockEnvSetSecretSubscriptionResolver{} + resolver.On("GetSubscription", mock.Anything, "sub-123"). Return("tenant-123", nil) kvSvc := &mockKeyVaultService{} @@ -495,8 +495,8 @@ func Test_EnvSetSecretAction_SelectExisting_NoVaults(t *testing.T) { prompter.On("PromptSubscription", mock.Anything, mock.Anything). Return("sub-123", nil) - resolver := &mockSubTenantResolver{} - resolver.On("LookupTenant", mock.Anything, "sub-123"). + resolver := &mockEnvSetSecretSubscriptionResolver{} + resolver.On("GetSubscription", mock.Anything, "sub-123"). Return("tenant-123", nil) kvSvc := &mockKeyVaultService{} @@ -530,8 +530,8 @@ func Test_EnvSetSecretAction_SelectKVError(t *testing.T) { prompter.On("PromptSubscription", mock.Anything, mock.Anything). Return("sub-123", nil) - resolver := &mockSubTenantResolver{} - resolver.On("LookupTenant", mock.Anything, "sub-123"). + resolver := &mockEnvSetSecretSubscriptionResolver{} + resolver.On("GetSubscription", mock.Anything, "sub-123"). Return("tenant-123", nil) kvSvc := &mockKeyVaultService{} @@ -565,8 +565,8 @@ func Test_EnvSetSecretAction_CreateNewKV_LocationError(t *testing.T) { prompter.On("PromptLocation", mock.Anything, "sub-123", mock.Anything, mock.Anything, mock.Anything). Return("", fmt.Errorf("location error")) - resolver := &mockSubTenantResolver{} - resolver.On("LookupTenant", mock.Anything, "sub-123"). + resolver := &mockEnvSetSecretSubscriptionResolver{} + resolver.On("GetSubscription", mock.Anything, "sub-123"). Return("tenant-123", nil) kvSvc := &mockKeyVaultService{} @@ -853,8 +853,8 @@ func Test_EnvSetSecretAction_SelectExisting_VaultListError(t *testing.T) { prompter.On("PromptSubscription", mock.Anything, mock.Anything). Return("sub-123", nil) - resolver := &mockSubTenantResolver{} - resolver.On("LookupTenant", mock.Anything, "sub-123"). + resolver := &mockEnvSetSecretSubscriptionResolver{} + resolver.On("GetSubscription", mock.Anything, "sub-123"). Return("tenant-123", nil) kvSvc := &mockKeyVaultService{} @@ -888,8 +888,8 @@ func Test_EnvSetSecretAction_CreateNew_ExistingVault_ListSecretsError(t *testing prompter.On("PromptSubscription", mock.Anything, mock.Anything). Return("sub-123", nil) - resolver := &mockSubTenantResolver{} - resolver.On("LookupTenant", mock.Anything, "sub-123"). + resolver := &mockEnvSetSecretSubscriptionResolver{} + resolver.On("GetSubscription", mock.Anything, "sub-123"). Return("tenant-123", nil) kvSvc := &mockKeyVaultService{} diff --git a/cli/azd/pkg/azureutil/principal.go b/cli/azd/pkg/azureutil/principal.go index 0e8c3d62e1c..39d6d7329cc 100644 --- a/cli/azd/pkg/azureutil/principal.go +++ b/cli/azd/pkg/azureutil/principal.go @@ -13,8 +13,8 @@ import ( ) // GetCurrentPrincipalId returns the object ID of the current principal authenticated with the CLI. -// It prefers the oid claim from an ARM access token, falling back to Graph /me when the token does -// not include a usable oid. +// It prefers the oid claim from an ARM access token, falling back to Graph /me when acquiring the +// token fails or when the token does not include a usable oid. func GetCurrentPrincipalId(ctx context.Context, userProfile *azapi.UserProfileService, tenantId string) (string, error) { token, err := userProfile.GetAccessToken(ctx, tenantId) if err == nil { From 6c927cb219fb87f933be13c9a58a7a601500b51e Mon Sep 17 00:00:00 2001 From: Wei Lim Date: Tue, 7 Apr 2026 11:31:13 -0700 Subject: [PATCH 6/6] Address Jon review comments Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- .../azsdk/storage/storage_blob_client_test.go | 57 +++++++++++-------- cli/azd/pkg/azureutil/principal_test.go | 39 +++++++++++++ 2 files changed, 72 insertions(+), 24 deletions(-) diff --git a/cli/azd/pkg/azsdk/storage/storage_blob_client_test.go b/cli/azd/pkg/azsdk/storage/storage_blob_client_test.go index 8611a8da056..5293398ed2f 100644 --- a/cli/azd/pkg/azsdk/storage/storage_blob_client_test.go +++ b/cli/azd/pkg/azsdk/storage/storage_blob_client_test.go @@ -38,25 +38,24 @@ type mockSubscriptionResolver struct { var _ account.SubscriptionResolver = (*mockSubscriptionResolver)(nil) -func (m *mockSubscriptionResolver) LookupTenant( - ctx context.Context, subscriptionId string) (string, error) { - args := m.Called(ctx, subscriptionId) - return args.String(0), args.Error(1) -} - func (m *mockSubscriptionResolver) GetSubscription( ctx context.Context, subscriptionId string, ) (*account.Subscription, error) { - tenantId, err := m.LookupTenant(ctx, subscriptionId) - if err != nil { - return nil, err + args := m.Called(ctx, subscriptionId) + + if subscription, ok := args.Get(0).(*account.Subscription); ok { + return subscription, args.Error(1) + } + + if tenantId, ok := args.Get(0).(string); ok { + return &account.Subscription{ + Id: subscriptionId, + TenantId: "resource-" + tenantId, + UserAccessTenantId: tenantId, + }, args.Error(1) } - return &account.Subscription{ - Id: subscriptionId, - TenantId: "resource-" + tenantId, - UserAccessTenantId: tenantId, - }, nil + return nil, args.Error(1) } // mockTokenCredential is a minimal mock implementation for testing @@ -123,8 +122,8 @@ func Test_NewBlobSdkClient_UsesHomeTenantWhenNoSubscriptionId(t *testing.T) { require.NotNil(t, client) mockCredProvider.AssertExpectations(t) - // TenantResolver should NOT be called when no subscription ID is provided - mockTenantResolver.AssertNotCalled(t, "LookupTenant", mock.Anything, mock.Anything) + // GetSubscription should NOT be called when no subscription ID is provided. + mockTenantResolver.AssertNotCalled(t, "GetSubscription", mock.Anything, mock.Anything) } func Test_NewBlobSdkClient_ResolvesTenantWhenSubscriptionIdProvided(t *testing.T) { @@ -144,8 +143,13 @@ func Test_NewBlobSdkClient_ResolvesTenantWhenSubscriptionIdProvided(t *testing.T coreClientOptions := &azcore.ClientOptions{} - // Expect tenant resolver to be called with the subscription ID - mockTenantResolver.On("LookupTenant", mock.Anything, testSubscriptionId).Return(testTenantId, nil) + // Expect the subscription resolver to be called with the subscription ID. + mockTenantResolver.On("GetSubscription", mock.Anything, testSubscriptionId). + Return(&account.Subscription{ + Id: testSubscriptionId, + TenantId: "resource-" + testTenantId, + UserAccessTenantId: testTenantId, + }, nil) // Expect credential provider to be called with resolved tenant ID mockCredProvider.On("GetTokenCredential", mock.Anything, testTenantId).Return(mockCred, nil) @@ -183,9 +187,9 @@ func Test_NewBlobSdkClient_ReturnsErrorWhenTenantResolutionFails(t *testing.T) { coreClientOptions := &azcore.ClientOptions{} - // Simulate tenant resolution failure - mockTenantResolver.On("LookupTenant", mock.Anything, testSubscriptionId). - Return("", errors.New("subscription not found")) + // Simulate subscription resolution failure. + mockTenantResolver.On("GetSubscription", mock.Anything, testSubscriptionId). + Return((*account.Subscription)(nil), errors.New("subscription not found")) client, err := NewBlobSdkClient( mockCredProvider, @@ -227,8 +231,13 @@ func Test_NewBlobSdkClient_FallsBackToDefaultSubscriptionFromUserConfig(t *testi }) mockConfigMgr.On("Load").Return(userCfg, nil) - // Expect tenant resolver to be called with the default subscription - mockTenantResolver.On("LookupTenant", mock.Anything, defaultSubscriptionId).Return(resolvedTenantId, nil) + // Expect the subscription resolver to be called with the default subscription. + mockTenantResolver.On("GetSubscription", mock.Anything, defaultSubscriptionId). + Return(&account.Subscription{ + Id: defaultSubscriptionId, + TenantId: "resource-" + resolvedTenantId, + UserAccessTenantId: resolvedTenantId, + }, nil) // Expect credential provider to be called with resolved tenant ID mockCredProvider.On("GetTokenCredential", mock.Anything, resolvedTenantId).Return(mockCred, nil) @@ -280,5 +289,5 @@ func Test_NewBlobSdkClient_UsesHomeTenantWhenUserConfigLoadFails(t *testing.T) { require.NoError(t, err) require.NotNil(t, client) mockCredProvider.AssertExpectations(t) - mockTenantResolver.AssertNotCalled(t, "LookupTenant", mock.Anything, mock.Anything) + mockTenantResolver.AssertNotCalled(t, "GetSubscription", mock.Anything, mock.Anything) } diff --git a/cli/azd/pkg/azureutil/principal_test.go b/cli/azd/pkg/azureutil/principal_test.go index 3a3043952e4..2970c6eb67f 100644 --- a/cli/azd/pkg/azureutil/principal_test.go +++ b/cli/azd/pkg/azureutil/principal_test.go @@ -86,3 +86,42 @@ func TestGetCurrentPrincipalId_FallsBackToGraphWhenOidMissing(t *testing.T) { require.NoError(t, err) require.Equal(t, "graph-user-id", principalId) } + +func TestGetCurrentPrincipalId_ReturnsJoinedErrorWhenTokenAndGraphFail(t *testing.T) { + t.Parallel() + + mockContext := mocks.NewMockContext(context.Background()) + mockContext.HttpClient.When(func(request *http.Request) bool { + return request.Method == http.MethodGet && strings.Contains(request.URL.Path, "/me") + }).RespondFn(func(request *http.Request) (*http.Response, error) { + return mocks.CreateEmptyHttpResponse(request, http.StatusBadRequest) + }) + + userProfile := azapi.NewUserProfileService( + &mocks.MockMultiTenantCredentialProvider{ + TokenMap: map[string]mocks.MockCredentials{ + "resource-tenant": { + GetTokenFn: func(ctx context.Context, options policy.TokenRequestOptions) (azcore.AccessToken, error) { + return azcore.AccessToken{ + Token: mocks.CreateJwtToken(t, map[string]string{ + "test": "fail", + }), + ExpiresOn: time.Now().Add(time.Hour), + }, nil + }, + }, + }, + }, + &azcore.ClientOptions{ + Transport: mockContext.HttpClient, + }, + cloud.AzurePublic(), + ) + + principalId, err := GetCurrentPrincipalId(*mockContext.Context, userProfile, "resource-tenant") + require.Error(t, err) + require.Empty(t, principalId) + require.ErrorContains(t, err, "resolving current principal ID from token oid and Graph fallback") + require.ErrorContains(t, err, "getting oid from token: no oid claim") + require.ErrorContains(t, err, "getting signed-in user id: failed retrieving current user profile") +}