Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 18 additions & 4 deletions cli/azd/pkg/azsdk/storage/storage_blob_client.go
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@ import (
"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/config"
)

// AccountConfig contains the configuration for connecting to a storage account
Expand Down Expand Up @@ -274,6 +275,7 @@ func (bc *blobClient) ensureContainerExists(ctx context.Context) error {
func NewBlobSdkClient(
credentialProvider auth.MultiTenantCredentialProvider,
accountConfig *AccountConfig,
userConfigManager config.UserConfigManager,
coreClientOptions *azcore.ClientOptions,
cloud *cloud.Cloud,
tenantResolver account.SubscriptionTenantResolver,
Expand All @@ -286,17 +288,29 @@ func NewBlobSdkClient(
accountConfig.Endpoint = cloud.StorageEndpointSuffix
}

// Determine which tenant to use for authentication
// Determine if we have a subscriptionId either from the config, or the default subscription
subscriptionId := accountConfig.SubscriptionId
if subscriptionId == "" {
userConfig, err := userConfigManager.Load()
if err == nil {
userSubscription, exists := userConfig.GetString("defaults.subscription")
if exists && userSubscription != "" {
subscriptionId = userSubscription
}
}
}

tenantId := ""
if accountConfig.SubscriptionId != "" {
if subscriptionId != "" {
// If a subscription ID is configured, resolve the tenant ID for that subscription
resolvedTenantId, err := tenantResolver.LookupTenant(context.Background(), accountConfig.SubscriptionId)
resolvedTenantId, err := tenantResolver.LookupTenant(context.Background(), subscriptionId)
if err != nil {
return nil, fmt.Errorf(
"failed to resolve tenant for subscription '%s': %w", accountConfig.SubscriptionId, err)
"failed to resolve tenant for subscription '%s': %w", subscriptionId, err)
}
tenantId = resolvedTenantId
}

// Otherwise, use home tenant ID (empty string)

credential, err := credentialProvider.GetTokenCredential(context.Background(), tenantId)
Expand Down
126 changes: 120 additions & 6 deletions cli/azd/pkg/azsdk/storage/storage_blob_client_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@ import (
"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/cloud"
"github.com/azure/azure-dev/cli/azd/pkg/config"
"github.com/stretchr/testify/mock"
"github.com/stretchr/testify/require"
)
Expand Down Expand Up @@ -54,25 +55,50 @@ func (m *mockTokenCredential) GetToken(
return args.Get(0).(azcore.AccessToken), args.Error(1)
}

// mockUserConfigManager is a mock implementation for testing
type mockUserConfigManager struct {
mock.Mock
}

var _ config.UserConfigManager = (*mockUserConfigManager)(nil)

func (m *mockUserConfigManager) Save(c config.Config) error {
args := m.Called(c)
return args.Error(0)
}

func (m *mockUserConfigManager) Load() (config.Config, error) {
args := m.Called()
if args.Get(0) == nil {
return nil, args.Error(1)
}
return args.Get(0).(config.Config), args.Error(1)
}

func Test_NewBlobSdkClient_UsesHomeTenantWhenNoSubscriptionId(t *testing.T) {
mockCredProvider := &mockMultiTenantCredentialProvider{}
mockTenantResolver := &mockSubscriptionTenantResolver{}
mockCred := &mockTokenCredential{}
mockConfigMgr := &mockUserConfigManager{}

accountConfig := &AccountConfig{
accountCfg := &AccountConfig{
AccountName: "testaccount",
ContainerName: "testcontainer",
// No SubscriptionId set - should use home tenant
}

coreClientOptions := &azcore.ClientOptions{}

// User config has no default subscription either
mockConfigMgr.On("Load").Return(config.NewEmptyConfig(), nil)

// Expect credential provider to be called with empty tenant ID (home tenant)
mockCredProvider.On("GetTokenCredential", mock.Anything, "").Return(mockCred, nil)

client, err := NewBlobSdkClient(
mockCredProvider,
accountConfig,
accountCfg,
mockConfigMgr,
coreClientOptions,
cloud.AzurePublic(),
mockTenantResolver,
Expand All @@ -90,11 +116,12 @@ func Test_NewBlobSdkClient_ResolvesTenantWhenSubscriptionIdProvided(t *testing.T
mockCredProvider := &mockMultiTenantCredentialProvider{}
mockTenantResolver := &mockSubscriptionTenantResolver{}
mockCred := &mockTokenCredential{}
mockConfigMgr := &mockUserConfigManager{}

testSubscriptionId := "test-subscription-id"
testTenantId := "test-tenant-id"

accountConfig := &AccountConfig{
accountCfg := &AccountConfig{
AccountName: "testaccount",
ContainerName: "testcontainer",
SubscriptionId: testSubscriptionId,
Expand All @@ -110,7 +137,8 @@ func Test_NewBlobSdkClient_ResolvesTenantWhenSubscriptionIdProvided(t *testing.T

client, err := NewBlobSdkClient(
mockCredProvider,
accountConfig,
accountCfg,
mockConfigMgr,
coreClientOptions,
cloud.AzurePublic(),
mockTenantResolver,
Expand All @@ -120,15 +148,19 @@ func Test_NewBlobSdkClient_ResolvesTenantWhenSubscriptionIdProvided(t *testing.T
require.NotNil(t, client)
mockCredProvider.AssertExpectations(t)
mockTenantResolver.AssertExpectations(t)

// UserConfigManager should NOT be called when SubscriptionId is already set
mockConfigMgr.AssertNotCalled(t, "Load")
}

func Test_NewBlobSdkClient_ReturnsErrorWhenTenantResolutionFails(t *testing.T) {
mockCredProvider := &mockMultiTenantCredentialProvider{}
mockTenantResolver := &mockSubscriptionTenantResolver{}
mockConfigMgr := &mockUserConfigManager{}

testSubscriptionId := "test-subscription-id"

accountConfig := &AccountConfig{
accountCfg := &AccountConfig{
AccountName: "testaccount",
ContainerName: "testcontainer",
SubscriptionId: testSubscriptionId,
Expand All @@ -142,7 +174,8 @@ func Test_NewBlobSdkClient_ReturnsErrorWhenTenantResolutionFails(t *testing.T) {

client, err := NewBlobSdkClient(
mockCredProvider,
accountConfig,
accountCfg,
mockConfigMgr,
coreClientOptions,
cloud.AzurePublic(),
mockTenantResolver,
Expand All @@ -153,3 +186,84 @@ func Test_NewBlobSdkClient_ReturnsErrorWhenTenantResolutionFails(t *testing.T) {
require.Contains(t, err.Error(), "failed to resolve tenant for subscription")
mockTenantResolver.AssertExpectations(t)
}

func Test_NewBlobSdkClient_FallsBackToDefaultSubscriptionFromUserConfig(t *testing.T) {
mockCredProvider := &mockMultiTenantCredentialProvider{}
mockTenantResolver := &mockSubscriptionTenantResolver{}
mockCred := &mockTokenCredential{}
mockConfigMgr := &mockUserConfigManager{}

defaultSubscriptionId := "default-sub-id"
resolvedTenantId := "resolved-tenant-id"

accountCfg := &AccountConfig{
AccountName: "testaccount",
ContainerName: "testcontainer",
// No SubscriptionId - should fall back to user config default
}

coreClientOptions := &azcore.ClientOptions{}

// User config returns a default subscription
userCfg := config.NewConfig(map[string]any{
"defaults": map[string]any{
"subscription": defaultSubscriptionId,
},
})
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 credential provider to be called with resolved tenant ID
mockCredProvider.On("GetTokenCredential", mock.Anything, resolvedTenantId).Return(mockCred, nil)

client, err := NewBlobSdkClient(
mockCredProvider,
accountCfg,
mockConfigMgr,
coreClientOptions,
cloud.AzurePublic(),
mockTenantResolver,
)

require.NoError(t, err)
require.NotNil(t, client)
mockConfigMgr.AssertExpectations(t)
mockTenantResolver.AssertExpectations(t)
mockCredProvider.AssertExpectations(t)
}

func Test_NewBlobSdkClient_UsesHomeTenantWhenUserConfigLoadFails(t *testing.T) {
mockCredProvider := &mockMultiTenantCredentialProvider{}
mockTenantResolver := &mockSubscriptionTenantResolver{}
mockCred := &mockTokenCredential{}
mockConfigMgr := &mockUserConfigManager{}

accountCfg := &AccountConfig{
AccountName: "testaccount",
ContainerName: "testcontainer",
}

coreClientOptions := &azcore.ClientOptions{}

// User config fails to load - should gracefully fall back to home tenant
mockConfigMgr.On("Load").Return(nil, errors.New("config not found"))

// Expect credential provider to be called with empty tenant ID (home tenant)
mockCredProvider.On("GetTokenCredential", mock.Anything, "").Return(mockCred, nil)

client, err := NewBlobSdkClient(
mockCredProvider,
accountCfg,
mockConfigMgr,
coreClientOptions,
cloud.AzurePublic(),
mockTenantResolver,
)

require.NoError(t, err)
require.NotNil(t, client)
mockCredProvider.AssertExpectations(t)
mockTenantResolver.AssertNotCalled(t, "LookupTenant", mock.Anything, mock.Anything)
}
5 changes: 5 additions & 0 deletions cli/azd/pkg/environment/manager_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -389,6 +389,11 @@ func registerContainerComponents(t *testing.T, mockContext *mocks.MockContext) {
return &mockSubscriptionTenantResolver{}
})

mockContext.Container.MustRegisterSingleton(func() config.UserConfigManager {
fileConfigManager := config.NewFileConfigManager(config.NewManager())
return config.NewUserConfigManager(fileConfigManager)
Comment thread
weikanglim marked this conversation as resolved.
})

mockContext.Container.MustRegisterSingleton(storage.NewBlobSdkClient)
mockContext.Container.MustRegisterSingleton(config.NewManager)
mockContext.Container.MustRegisterSingleton(storage.NewBlobClient)
Expand Down
2 changes: 2 additions & 0 deletions cli/azd/test/functional/remote_state_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -96,9 +96,11 @@ func createBlobClient(
// Create a mock SubscriptionTenantResolver that returns empty tenant (home tenant)
tenantResolver := &mockTenantResolver{}

userConfigManager := config.NewUserConfigManager(fileConfigManager)
sdkClient, err := storage.NewBlobSdkClient(
auth.NewMultiTenantCredentialProvider(authManager),
storageConfig,
userConfigManager,
coreClientOptions,
cloud.AzurePublic(),
tenantResolver,
Expand Down
Loading