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
18 changes: 10 additions & 8 deletions cli/azd/cmd/auth_token.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
Expand All @@ -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{
Expand All @@ -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
Expand All @@ -100,32 +100,34 @@ 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 {
// no env var from system
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
}
Expand Down
38 changes: 22 additions & 16 deletions cli/azd/cmd/auth_token_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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(),
)

Expand Down Expand Up @@ -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(),
)

Expand Down Expand Up @@ -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(),
Expand Down Expand Up @@ -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(),
Expand All @@ -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),
)
Expand Down Expand Up @@ -212,7 +213,7 @@ func TestAuthTokenAzdEnvError(t *testing.T) {
environment.SubscriptionIdEnvVarName: expectedSubId,
}), nil
},
&mockSubscriptionTenantResolver{
&mockSubscriptionResolver{
Err: errors.New(expectedError),
},
cloud.AzurePublic(),
Expand All @@ -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),
)
Expand Down Expand Up @@ -253,7 +254,7 @@ func TestAuthTokenAzdEnv(t *testing.T) {
environment.SubscriptionIdEnvVarName: "sub-id",
}), nil
},
&mockSubscriptionTenantResolver{
&mockSubscriptionResolver{
TenantId: expectedTenant,
},
cloud.AzurePublic(),
Expand Down Expand Up @@ -294,7 +295,7 @@ func TestAuthTokenAzdEnvWithEmpty(t *testing.T) {
environment.SubscriptionIdEnvVarName: "",
}), nil
},
&mockSubscriptionTenantResolver{
&mockSubscriptionResolver{
TenantId: expectedTenant,
},
cloud.AzurePublic(),
Expand Down Expand Up @@ -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(),
)

Expand All @@ -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(),
)

Expand All @@ -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
}
2 changes: 1 addition & 1 deletion cli/azd/cmd/container.go
Original file line number Diff line number Diff line change
Expand Up @@ -700,7 +700,7 @@ func registerCommonDependencies(container *ioc.NestedContainer) {
})
})

container.MustRegisterSingleton(func(subManager *account.SubscriptionsManager) account.SubscriptionTenantResolver {
container.MustRegisterSingleton(func(subManager *account.SubscriptionsManager) account.SubscriptionResolver {
return subManager
})

Expand Down
Loading
Loading