diff --git a/cli/azd/cmd/auto_install.go b/cli/azd/cmd/auto_install.go index 19c0eba16d9..064164e4d50 100644 --- a/cli/azd/cmd/auto_install.go +++ b/cli/azd/cmd/auto_install.go @@ -10,6 +10,7 @@ import ( "io" "log" "os" + "path/filepath" "slices" "strconv" "strings" @@ -338,14 +339,37 @@ func tryAutoInstallForPartialNamespace( func tryAutoInstallExtension( ctx context.Context, console input.Console, - extensionManager *extensions.Manager, + extensionManager extensionAutoInstallManager, extension extensions.ExtensionMetadata) (bool, error) { + return tryAutoInstallExtensionVersion(ctx, console, extensionManager, extension, "") +} + +type extensionAutoInstallManager interface { + FindExtensions(ctx context.Context, options *extensions.FilterOptions) ([]*extensions.ExtensionMetadata, error) + GetInstalled(options extensions.FilterOptions) (*extensions.Extension, error) + Install( + ctx context.Context, + extension *extensions.ExtensionMetadata, + versionPreference string, + ) (*extensions.ExtensionVersion, error) + ListInstalled() (map[string]*extensions.Extension, error) +} +func tryAutoInstallExtensionVersion( + ctx context.Context, + console input.Console, + extensionManager extensionAutoInstallManager, + extension extensions.ExtensionMetadata, + versionPreference string, +) (bool, error) { // Check if the extension is already installed - _, err := extensionManager.GetInstalled(extensions.FilterOptions{ + installedExtension, err := extensionManager.GetInstalled(extensions.FilterOptions{ Id: extension.Id, }) if err == nil { + if err := validateInstalledExtensionVersion(installedExtension, versionPreference); err != nil { + return false, err + } return false, nil } @@ -372,7 +396,7 @@ func tryAutoInstallExtension( Message: "Confirm installation", }) if err != nil { - return false, nil + return false, err } if !shouldInstall { @@ -381,7 +405,7 @@ func tryAutoInstallExtension( // Install the extension console.Message(ctx, fmt.Sprintf("Installing extension '%s'...\n", extension.Id)) - _, err = extensionManager.Install(ctx, &extension, "") + _, err = extensionManager.Install(ctx, &extension, versionPreference) if err != nil { return false, fmt.Errorf("failed to install extension: %w", err) } @@ -452,6 +476,42 @@ type ExecuteResult struct { LatestVersion <-chan *update.VersionInfo } +func newRootCmdForExecution( + rootContainer *ioc.NestedContainer, + globalOpts *internal.GlobalCommandOptions, +) (*cobra.Command, error) { + if globalOpts.Cwd == "" { + return NewRootCmd(false, nil, rootContainer), nil + } + + absoluteCwd, err := filepath.Abs(globalOpts.Cwd) + if err != nil { + return nil, fmt.Errorf("resolving cwd: %w", err) + } + globalOpts.Cwd = absoluteCwd + + if _, err := os.Stat(absoluteCwd); os.IsNotExist(err) { + // PersistentPreRunE owns prompting for and creating a missing --cwd directory. + return NewRootCmd(false, nil, rootContainer), nil + } else if err != nil { + return nil, fmt.Errorf("checking cwd: %w", err) + } + + previousCwd, err := os.Getwd() + if err != nil { + return nil, fmt.Errorf("getting current directory: %w", err) + } + if err := os.Chdir(absoluteCwd); err != nil { + return nil, fmt.Errorf("changing directory to %s: %w", absoluteCwd, err) + } + + rootCmd := NewRootCmd(false, nil, rootContainer) + if err := os.Chdir(previousCwd); err != nil { + return nil, fmt.Errorf("restoring current directory: %w", err) + } + return rootCmd, nil +} + // ExecuteWithAutoInstall executes the command and handles auto-installation of extensions for unknown commands. func ExecuteWithAutoInstall(ctx context.Context, rootContainer *ioc.NestedContainer) *ExecuteResult { result := &ExecuteResult{} @@ -473,7 +533,12 @@ func ExecuteWithAutoInstall(ctx context.Context, rootContainer *ioc.NestedContai // Creating the RootCmd takes care of registering common dependencies in rootContainer. // The command tree will retrieve globalOpts from the container via its FlagsResolver. - rootCmd := NewRootCmd(false, nil, rootContainer) + rootCmd, err := newRootCmdForExecution(rootContainer, globalOpts) + if err != nil { + fmt.Fprintln(os.Stderr, output.WithErrorFormat("ERROR: %s", err.Error())) + result.Err = err + return result + } var extensionManager *extensions.Manager var console input.Console @@ -483,6 +548,8 @@ func ExecuteWithAutoInstall(ctx context.Context, rootContainer *ioc.NestedContai // This allows us to determine if a subcommand was provided or not or if the command is unknown. foundCmd, originalArgs, err := rootCmd.Find(os.Args[1:]) if err == nil { + projectExtensionsHandled := false + // Detect lightspeed commands from the cobra annotation set by CobraBuilder. result.IsLightspeed = foundCmd.Annotations[actions.AnnotationLightspeed] == "true" @@ -495,18 +562,37 @@ func ExecuteWithAutoInstall(ctx context.Context, rootContainer *ioc.NestedContai result.LatestVersion = startUpdateCheck(ctx) } + handled, installed, err := tryAutoInstallProjectExtensions(ctx, rootContainer, foundCmd) + projectExtensionsHandled = handled + if err != nil { + if resolveErr := rootContainer.Resolve(&console); resolveErr != nil { + fmt.Fprintln(os.Stderr, output.WithErrorFormat("ERROR: %s", err.Error())) + } else { + displayAutoInstallError(ctx, console, err) + } + result.Err = err + return result + } else if installed { + rootCmd = newRootCmdWithoutRegistration(rootContainer) + foundCmd, originalArgs, err = rootCmd.Find(os.Args[1:]) + if err != nil { + result.Err = err + return result + } + } + // Check for partial namespace match (e.g., "ai" found but "ai.agent" not installed) if installed := tryAutoInstallForPartialNamespace( ctx, rootContainer, foundCmd, originalArgs, ); installed { // Extension was installed, rebuild command tree and execute - rootCmd = NewRootCmd(false, nil, rootContainer) + rootCmd = newRootCmdWithoutRegistration(rootContainer) result.Err = rootCmd.ExecuteContext(ctx) return result } // Known command, proceed with normal execution - err := rootCmd.ExecuteContext(ctx) + err = rootCmd.ExecuteContext(ctx) // Only attempt service-host auto-install when the command failed with that specific error. // Other command errors (for example, unsupported output formats) should be returned directly. @@ -515,6 +601,15 @@ func ExecuteWithAutoInstall(ctx context.Context, rootContainer *ioc.NestedContai result.Err = err return result } + if projectExtensionsHandled { + if resolveErr := rootContainer.Resolve(&console); resolveErr != nil { + fmt.Fprintln(os.Stderr, unsupportedErr.ErrorMessage) + } else { + console.Message(ctx, unsupportedErr.ErrorMessage) + } + result.Err = err + return result + } if err := rootContainer.Resolve(&extensionManager); err != nil { log.Panic("failed to resolve extension manager for auto-install:", err) @@ -534,9 +629,18 @@ func ExecuteWithAutoInstall(ctx context.Context, rootContainer *ioc.NestedContai console.Message(ctx, unsupportedErr.ErrorMessage) return result } - // Note: We don't need to filter or check which extensions are installed. - // If any of these extensions would be installed, the auto-install wouldn't have been triggered because - // there would be at least one extensions providing the capability and provider. + installedExtensions, err := extensionManager.ListInstalled() + if err != nil { + log.Println("Error: list installed extensions. Skipping auto-install:", err) + console.Message(ctx, unsupportedErr.ErrorMessage) + return result + } + availableExtensionsForHost = filterExtensionsForProvider( + availableExtensionsForHost, + extensions.ServiceTargetProviderCapability, + requiredHost, + ) + availableExtensionsForHost = uninstalledExtensionMatches(availableExtensionsForHost, installedExtensions) if len(availableExtensionsForHost) == 0 { // did not find an extension with the capability, just print the original error message console.Message(ctx, unsupportedErr.ErrorMessage) @@ -572,7 +676,7 @@ func ExecuteWithAutoInstall(ctx context.Context, rootContainer *ioc.NestedContai if installed { // Extension was installed, build command tree and execute - rootCmd := NewRootCmd(false, nil, rootContainer) + rootCmd := newRootCmdWithoutRegistration(rootContainer) result.Err = rootCmd.ExecuteContext(ctx) return result } @@ -699,7 +803,7 @@ func ExecuteWithAutoInstall(ctx context.Context, rootContainer *ioc.NestedContai if installed { // Extension was installed, build command tree and execute - rootCmd := NewRootCmd(false, nil, rootContainer) + rootCmd := newRootCmdWithoutRegistration(rootContainer) result.Err = rootCmd.ExecuteContext(ctx) return result } diff --git a/cli/azd/cmd/auto_install_test.go b/cli/azd/cmd/auto_install_test.go index 2b738c66096..2a9fe6c6d72 100644 --- a/cli/azd/cmd/auto_install_test.go +++ b/cli/azd/cmd/auto_install_test.go @@ -4,7 +4,10 @@ package cmd import ( + "context" "fmt" + "path/filepath" + "slices" "strings" "testing" @@ -15,11 +18,763 @@ import ( "github.com/azure/azure-dev/cli/azd/internal" "github.com/azure/azure-dev/cli/azd/internal/runcontext/agentdetect" "github.com/azure/azure-dev/cli/azd/pkg/extensions" + "github.com/azure/azure-dev/cli/azd/pkg/infra/provisioning" "github.com/azure/azure-dev/cli/azd/pkg/input" "github.com/azure/azure-dev/cli/azd/pkg/ioc" + "github.com/azure/azure-dev/cli/azd/pkg/project" "github.com/azure/azure-dev/cli/azd/test/mocks/mockinput" ) +type fakeExtensionAutoInstallManager struct { + available []*extensions.ExtensionMetadata + installed map[string]*extensions.Extension + findErr error +} + +func (m *fakeExtensionAutoInstallManager) FindExtensions( + _ context.Context, + options *extensions.FilterOptions, +) ([]*extensions.ExtensionMetadata, error) { + if m.findErr != nil { + return nil, m.findErr + } + + var matches []*extensions.ExtensionMetadata + for _, extension := range m.available { + if options.Id != "" && extension.Id != options.Id { + continue + } + if options.Source != "" && !strings.EqualFold(extension.Source, options.Source) { + continue + } + if options.Version != "" { + if _, err := extensions.ResolveExtensionVersion(extension, options.Version, nil); err != nil { + continue + } + } + hasCapability := slices.ContainsFunc(extension.Versions, func(version extensions.ExtensionVersion) bool { + return slices.Contains(version.Capabilities, options.Capability) + }) + if options.Capability != "" && !hasCapability { + continue + } + hasProvider := slices.ContainsFunc(extension.Versions, func(version extensions.ExtensionVersion) bool { + return slices.ContainsFunc(version.Providers, func(provider extensions.Provider) bool { + return provider.Name == options.Provider + }) + }) + if options.Provider != "" && !hasProvider { + continue + } + matches = append(matches, extension) + } + return matches, nil +} + +func (m *fakeExtensionAutoInstallManager) GetInstalled( + options extensions.FilterOptions, +) (*extensions.Extension, error) { + if extension, ok := m.installed[options.Id]; ok { + return extension, nil + } + return nil, fmt.Errorf("extension not installed") +} + +func (m *fakeExtensionAutoInstallManager) Install( + _ context.Context, + extension *extensions.ExtensionMetadata, + _ string, +) (*extensions.ExtensionVersion, error) { + version := &extension.Versions[0] + m.installed[extension.Id] = &extensions.Extension{ + Id: extension.Id, + Version: version.Version, + } + return version, nil +} + +func (m *fakeExtensionAutoInstallManager) ListInstalled() (map[string]*extensions.Extension, error) { + return m.installed, nil +} + +func TestMissingProjectExtensions(t *testing.T) { + versionConstraint := ">=1.0.0-beta.4" + manager := &fakeExtensionAutoInstallManager{ + available: []*extensions.ExtensionMetadata{ + { + Id: "azure.ai.projects", + Versions: []extensions.ExtensionVersion{ + {Version: "2.0.0"}, + { + Version: "1.0.0", + Capabilities: []extensions.CapabilityType{extensions.ServiceTargetProviderCapability}, + Providers: []extensions.Provider{{ + Name: "azure.ai.project", + Type: extensions.ServiceTargetProviderType, + }}, + }, + }, + }, + { + Id: "azure.ai.agents", + Versions: []extensions.ExtensionVersion{{ + Version: "1.0.0", + Capabilities: []extensions.CapabilityType{extensions.ServiceTargetProviderCapability}, + Providers: []extensions.Provider{{ + Name: "azure.ai.agent", + Type: extensions.ServiceTargetProviderType, + }}, + }}, + }, + { + Id: "microsoft.foundry", + Versions: []extensions.ExtensionVersion{{ + Version: "1.0.0", + Capabilities: []extensions.CapabilityType{extensions.ProvisioningProviderCapability}, + Providers: []extensions.Provider{{ + Name: "microsoft.foundry", + Type: extensions.ProvisioningProviderType, + }}, + }}, + }, + }, + installed: map[string]*extensions.Extension{}, + } + projectConfig := &project.ProjectConfig{ + RequiredVersions: &project.RequiredVersions{ + Extensions: map[string]*string{ + "microsoft.foundry": new(versionConstraint), + }, + }, + Services: map[string]*project.ServiceConfig{ + "project": {Host: "azure.ai.project"}, + "agent": {Host: "azure.ai.agent"}, + }, + Infra: provisioning.Options{Provider: "microsoft.foundry"}, + } + + requirements, err := missingProjectExtensions( + t.Context(), + mockinput.NewMockConsole(), + manager, + projectConfig, + ) + require.NoError(t, err) + require.Len(t, requirements, 3) + assert.Equal(t, "microsoft.foundry", requirements[0].extension.Id) + assert.Equal(t, versionConstraint, requirements[0].versionPreference) + assert.Equal(t, "azure.ai.agents", requirements[1].extension.Id) + assert.Equal(t, "azure.ai.projects", requirements[2].extension.Id) + require.Len(t, requirements[2].extension.Versions, 1) + assert.Equal(t, "1.0.0", requirements[2].extension.Versions[0].Version) +} + +func TestMissingProjectExtensionsSkipsInstalledProviderAcrossSources(t *testing.T) { + manager := &fakeExtensionAutoInstallManager{ + available: []*extensions.ExtensionMetadata{ + { + Id: "microsoft.azd.demo", + Source: "azd", + Versions: []extensions.ExtensionVersion{{ + Version: "0.7.0", + Capabilities: []extensions.CapabilityType{extensions.ServiceTargetProviderCapability}, + Providers: []extensions.Provider{{ + Name: "demo", + Type: extensions.ServiceTargetProviderType, + }}, + }}, + }, + { + Id: "microsoft.azd.demo", + Source: "local", + Versions: []extensions.ExtensionVersion{{ + Version: "0.7.0", + Capabilities: []extensions.CapabilityType{extensions.ServiceTargetProviderCapability}, + Providers: []extensions.Provider{{ + Name: "demo", + Type: extensions.ServiceTargetProviderType, + }}, + }}, + }, + }, + installed: map[string]*extensions.Extension{ + "microsoft.azd.demo": { + Id: "microsoft.azd.demo", + Version: "0.3.0", + Source: "azd", + }, + }, + } + projectConfig := &project.ProjectConfig{ + Services: map[string]*project.ServiceConfig{ + "demo": {Host: "demo"}, + }, + } + + // The mock has no Select response. The test panics if source selection is prompted. + requirements, err := missingProjectExtensions( + t.Context(), + mockinput.NewMockConsole(), + manager, + projectConfig, + ) + + require.NoError(t, err) + require.Empty(t, requirements) +} + +func TestMissingProjectExtensionsReusesSourceChoiceAcrossProviders(t *testing.T) { + providerVersion := extensions.ExtensionVersion{ + Version: "0.7.0", + Capabilities: []extensions.CapabilityType{ + extensions.ServiceTargetProviderCapability, + extensions.ProvisioningProviderCapability, + }, + Providers: []extensions.Provider{ + {Name: "demo", Type: extensions.ServiceTargetProviderType}, + {Name: "demo", Type: extensions.ProvisioningProviderType}, + }, + } + manager := &fakeExtensionAutoInstallManager{ + available: []*extensions.ExtensionMetadata{ + { + Id: "microsoft.azd.demo", + Source: "azd", + Versions: []extensions.ExtensionVersion{providerVersion}, + }, + { + Id: "microsoft.azd.demo", + Source: "local", + Versions: []extensions.ExtensionVersion{providerVersion}, + }, + }, + installed: map[string]*extensions.Extension{}, + } + projectConfig := &project.ProjectConfig{ + Services: map[string]*project.ServiceConfig{ + "demo": {Host: "demo"}, + }, + Infra: provisioning.Options{Provider: "demo"}, + } + selectCount := 0 + console := mockinput.NewMockConsole() + console.WhenSelect(func(options input.ConsoleOptions) bool { + selectCount++ + return true + }).Respond(0) + + requirements, err := missingProjectExtensions(t.Context(), console, manager, projectConfig) + + require.NoError(t, err) + require.Len(t, requirements, 1) + require.Equal(t, "azd", requirements[0].extension.Source) + require.Equal(t, 1, selectCount) +} + +func TestMissingProjectExtensionsSkipsExtensionPackDependencies(t *testing.T) { + manager := &fakeExtensionAutoInstallManager{ + available: []*extensions.ExtensionMetadata{ + { + Id: "microsoft.foundry", + Source: "azd", + Versions: []extensions.ExtensionVersion{{ + Version: "1.0.0", + Dependencies: []extensions.ExtensionDependency{ + {Id: "microsoft.foundry.bundle"}, + }, + }}, + }, + { + Id: "microsoft.foundry.bundle", + Source: "azd", + Versions: []extensions.ExtensionVersion{{ + Version: "1.0.0", + Dependencies: []extensions.ExtensionDependency{ + {Id: "azure.ai.agents"}, + {Id: "azure.ai.projects"}, + }, + }}, + }, + { + Id: "azure.ai.agents", + Source: "azd", + Versions: []extensions.ExtensionVersion{{ + Version: "1.0.0", + Capabilities: []extensions.CapabilityType{ + extensions.ServiceTargetProviderCapability, + extensions.ProvisioningProviderCapability, + }, + Providers: []extensions.Provider{ + {Name: "azure.ai.agent", Type: extensions.ServiceTargetProviderType}, + {Name: "microsoft.foundry", Type: extensions.ProvisioningProviderType}, + }, + }}, + }, + { + Id: "azure.ai.agents", + Source: "local", + Versions: []extensions.ExtensionVersion{{ + Version: "1.0.0", + Capabilities: []extensions.CapabilityType{ + extensions.ServiceTargetProviderCapability, + extensions.ProvisioningProviderCapability, + }, + Providers: []extensions.Provider{ + {Name: "azure.ai.agent", Type: extensions.ServiceTargetProviderType}, + {Name: "microsoft.foundry", Type: extensions.ProvisioningProviderType}, + }, + }}, + }, + { + Id: "azure.ai.projects", + Source: "azd", + Versions: []extensions.ExtensionVersion{{ + Version: "1.0.0", + Capabilities: []extensions.CapabilityType{extensions.ServiceTargetProviderCapability}, + Providers: []extensions.Provider{{ + Name: "azure.ai.project", + Type: extensions.ServiceTargetProviderType, + }}, + }}, + }, + }, + installed: map[string]*extensions.Extension{}, + } + projectConfig := &project.ProjectConfig{ + RequiredVersions: &project.RequiredVersions{ + Extensions: map[string]*string{ + "microsoft.foundry": new("1.0.0"), + }, + }, + Services: map[string]*project.ServiceConfig{ + "agent": {Host: "azure.ai.agent"}, + "project": {Host: "azure.ai.project"}, + }, + Infra: provisioning.Options{Provider: "microsoft.foundry"}, + } + + // The mock has no Select response. The test panics if a pack dependency prompts for a source. + requirements, err := missingProjectExtensions( + t.Context(), + mockinput.NewMockConsole(), + manager, + projectConfig, + ) + + require.NoError(t, err) + require.Len(t, requirements, 1) + require.Equal(t, "microsoft.foundry", requirements[0].extension.Id) +} + +func TestMissingProjectExtensionsRejectsDependencyVersionWithoutProvider(t *testing.T) { + manager := &fakeExtensionAutoInstallManager{ + available: []*extensions.ExtensionMetadata{ + { + Id: "test.pack", + Source: "azd", + Versions: []extensions.ExtensionVersion{{ + Version: "1.0.0", + Dependencies: []extensions.ExtensionDependency{ + {Id: "test.provider", Version: "1.0.0"}, + }, + }}, + }, + { + Id: "test.provider", + Source: "azd", + Versions: []extensions.ExtensionVersion{ + {Version: "1.0.0"}, + { + Version: "2.0.0", + Capabilities: []extensions.CapabilityType{extensions.ServiceTargetProviderCapability}, + Providers: []extensions.Provider{{ + Name: "demo", + Type: extensions.ServiceTargetProviderType, + }}, + }, + }, + }, + }, + installed: map[string]*extensions.Extension{}, + } + projectConfig := &project.ProjectConfig{ + RequiredVersions: &project.RequiredVersions{ + Extensions: map[string]*string{ + "test.pack": new("1.0.0"), + }, + }, + Services: map[string]*project.ServiceConfig{ + "demo": {Host: "demo"}, + }, + } + + // The mock has no Select response. The dependency conflict must fail before prompting. + _, err := missingProjectExtensions( + t.Context(), + mockinput.NewMockConsole(), + manager, + projectConfig, + ) + + require.ErrorContains(t, err, "test.pack requires dependency test.provider version 1.0.0") + require.ErrorContains(t, err, `does not provide service-target-provider "demo"`) +} + +func TestMissingProjectExtensionsRejectsInstalledDependencyWithoutProvider(t *testing.T) { + manager := &fakeExtensionAutoInstallManager{ + available: []*extensions.ExtensionMetadata{ + { + Id: "test.pack", + Source: "azd", + Versions: []extensions.ExtensionVersion{{ + Version: "1.0.0", + Dependencies: []extensions.ExtensionDependency{ + {Id: "test.provider"}, + }, + }}, + }, + { + Id: "test.provider", + Source: "azd", + Versions: []extensions.ExtensionVersion{ + {Version: "1.0.0"}, + { + Version: "2.0.0", + Capabilities: []extensions.CapabilityType{extensions.ServiceTargetProviderCapability}, + Providers: []extensions.Provider{{ + Name: "demo", + Type: extensions.ServiceTargetProviderType, + }}, + }, + }, + }, + }, + installed: map[string]*extensions.Extension{ + "test.provider": { + Id: "test.provider", + Version: "1.0.0", + }, + }, + } + projectConfig := &project.ProjectConfig{ + RequiredVersions: &project.RequiredVersions{ + Extensions: map[string]*string{ + "test.pack": new("1.0.0"), + }, + }, + Services: map[string]*project.ServiceConfig{ + "demo": {Host: "demo"}, + }, + } + + _, err := missingProjectExtensions( + t.Context(), + mockinput.NewMockConsole(), + manager, + projectConfig, + ) + + require.ErrorContains(t, err, "test.pack requires dependency test.provider version 1.0.0") + require.ErrorContains(t, err, `does not provide service-target-provider "demo"`) +} + +func TestMissingProjectExtensionsIgnoresSplitProviderMetadata(t *testing.T) { + manager := &fakeExtensionAutoInstallManager{ + available: []*extensions.ExtensionMetadata{ + { + Id: "test.provider", + Source: "azd", + Versions: []extensions.ExtensionVersion{ + { + Version: "1.0.0", + Capabilities: []extensions.CapabilityType{extensions.ServiceTargetProviderCapability}, + }, + { + Version: "2.0.0", + Providers: []extensions.Provider{{ + Name: "demo", + Type: extensions.ServiceTargetProviderType, + }}, + }, + }, + }, + }, + installed: map[string]*extensions.Extension{}, + } + projectConfig := &project.ProjectConfig{ + Services: map[string]*project.ServiceConfig{ + "demo": {Host: "demo"}, + }, + } + + // The mock has no Select response. No single version provides both the capability and provider. + requirements, err := missingProjectExtensions( + t.Context(), + mockinput.NewMockConsole(), + manager, + projectConfig, + ) + + require.NoError(t, err) + require.Empty(t, requirements) +} + +func TestMissingProjectExtensionsInstalledIdIsCaseInsensitive(t *testing.T) { + manager := &fakeExtensionAutoInstallManager{ + installed: map[string]*extensions.Extension{ + "microsoft.foundry": { + Id: "microsoft.foundry", + Version: "1.0.0", + }, + }, + } + projectConfig := &project.ProjectConfig{ + RequiredVersions: &project.RequiredVersions{ + Extensions: map[string]*string{ + "Microsoft.Foundry": new("1.0.0"), + }, + }, + } + + requirements, err := missingProjectExtensions( + t.Context(), + mockinput.NewMockConsole(), + manager, + projectConfig, + ) + + require.NoError(t, err) + require.Empty(t, requirements) +} + +func TestMissingProjectExtensionsRejectsInstalledVersionConstraint(t *testing.T) { + manager := &fakeExtensionAutoInstallManager{ + installed: map[string]*extensions.Extension{ + "test.extension": { + Id: "test.extension", + Version: "1.0.0", + }, + }, + } + projectConfig := &project.ProjectConfig{ + RequiredVersions: &project.RequiredVersions{ + Extensions: map[string]*string{ + "test.extension": new(">=2.0.0"), + }, + }, + } + + _, err := missingProjectExtensions( + t.Context(), + mockinput.NewMockConsole(), + manager, + projectConfig, + ) + + require.EqualError( + t, + err, + `installed extension test.extension version 1.0.0 does not satisfy constraint ">=2.0.0"`, + ) +} + +func TestMissingProjectExtensionsRejectsExplicitVersionWithoutProvider(t *testing.T) { + manager := &fakeExtensionAutoInstallManager{ + available: []*extensions.ExtensionMetadata{ + { + Id: "test.extension", + Source: "azd", + Versions: []extensions.ExtensionVersion{ + {Version: "1.0.0"}, + { + Version: "2.0.0", + Capabilities: []extensions.CapabilityType{extensions.ServiceTargetProviderCapability}, + Providers: []extensions.Provider{{ + Name: "demo", + Type: extensions.ServiceTargetProviderType, + }}, + }, + }, + }, + }, + installed: map[string]*extensions.Extension{}, + } + projectConfig := &project.ProjectConfig{ + RequiredVersions: &project.RequiredVersions{ + Extensions: map[string]*string{ + "test.extension": new("1.0.0"), + }, + }, + Services: map[string]*project.ServiceConfig{ + "demo": {Host: "demo"}, + }, + } + + _, err := missingProjectExtensions( + t.Context(), + mockinput.NewMockConsole(), + manager, + projectConfig, + ) + + require.EqualError( + t, + err, + `required extension test.extension version 1.0.0 does not provide service-target-provider "demo"`, + ) +} + +func TestMissingProjectExtensionsPropagatesProviderLookupError(t *testing.T) { + manager := &fakeExtensionAutoInstallManager{ + installed: map[string]*extensions.Extension{}, + findErr: fmt.Errorf("registry unavailable"), + } + projectConfig := &project.ProjectConfig{ + Services: map[string]*project.ServiceConfig{ + "demo": {Host: "demo"}, + }, + } + + _, err := missingProjectExtensions( + t.Context(), + mockinput.NewMockConsole(), + manager, + projectConfig, + ) + + require.ErrorContains(t, err, `finding extension for provider "demo": registry unavailable`) +} + +func TestNewRootCmdForExecutionUsesCwd(t *testing.T) { + currentDir := t.TempDir() + targetDir := t.TempDir() + require.NoError(t, project.Save( + t.Context(), + &project.ProjectConfig{Name: "current-project"}, + filepath.Join(currentDir, "azure.yaml"), + )) + require.NoError(t, project.Save( + t.Context(), + &project.ProjectConfig{Name: "target-project"}, + filepath.Join(targetDir, "azure.yaml"), + )) + t.Chdir(currentDir) + + container := ioc.NewNestedContainer(nil) + ioc.RegisterInstance(container, context.WithoutCancel(t.Context())) + globalOpts := &internal.GlobalCommandOptions{Cwd: targetDir} + ioc.RegisterInstance(container, globalOpts) + _, err := newRootCmdForExecution(container, globalOpts) + require.NoError(t, err) + + var projectConfig *project.ProjectConfig + require.NoError(t, container.Resolve(&projectConfig)) + require.Equal(t, "target-project", projectConfig.Name) +} + +func TestExtensionVersionProvidesProviderMatchesType(t *testing.T) { + version := &extensions.ExtensionVersion{ + Capabilities: []extensions.CapabilityType{ + extensions.ServiceTargetProviderCapability, + extensions.ProvisioningProviderCapability, + }, + Providers: []extensions.Provider{ + {Name: "service", Type: extensions.ServiceTargetProviderType}, + {Name: "infra", Type: extensions.ProvisioningProviderType}, + }, + } + + require.True(t, extensionVersionProvidesProvider( + version, + extensions.ServiceTargetProviderCapability, + "service", + )) + require.True(t, extensionVersionProvidesProvider( + version, + extensions.ProvisioningProviderCapability, + "infra", + )) + require.False(t, extensionVersionProvidesProvider( + version, + extensions.ServiceTargetProviderCapability, + "infra", + )) + require.False(t, extensionVersionProvidesProvider( + version, + extensions.ProvisioningProviderCapability, + "service", + )) +} + +func TestTryAutoInstallExtensionVersionRejectsInstalledVersionConstraint(t *testing.T) { + manager := &fakeExtensionAutoInstallManager{ + installed: map[string]*extensions.Extension{ + "test.extension": { + Id: "test.extension", + Version: "1.0.0", + }, + }, + } + + installed, err := tryAutoInstallExtensionVersion( + t.Context(), + mockinput.NewMockConsole(), + manager, + extensions.ExtensionMetadata{Id: "test.extension"}, + ">=2.0.0", + ) + + require.False(t, installed) + require.EqualError( + t, + err, + `installed extension test.extension version 1.0.0 does not satisfy constraint ">=2.0.0"`, + ) +} + +func TestDisplayAutoInstallError(t *testing.T) { + t.Run("RegularError", func(t *testing.T) { + console := mockinput.NewMockConsole() + + displayAutoInstallError(t.Context(), console, fmt.Errorf("install failed")) + + require.Contains(t, strings.Join(console.Output(), "\n"), "ERROR: install failed") + }) + + t.Run("ErrorWithSuggestion", func(t *testing.T) { + console := mockinput.NewMockConsole() + + displayAutoInstallError(t.Context(), console, &internal.ErrorWithSuggestion{ + Err: fmt.Errorf("install failed"), + Message: "The required extension could not be installed.", + Suggestion: "Check the extension version and retry.", + }) + + output := strings.Join(console.Output(), "\n") + require.Contains(t, output, "ERROR: The required extension could not be installed.") + require.Contains(t, output, "Suggestion: Check the extension version and retry.") + }) +} + +func TestProjectCommandSupportsExtensionAutoInstall(t *testing.T) { + root := &cobra.Command{Use: "azd"} + up := &cobra.Command{Use: "up"} + extension := &cobra.Command{Use: "agent", Annotations: map[string]string{"extension.id": "azure.ai.agents"}} + env := &cobra.Command{Use: "env"} + refresh := &cobra.Command{Use: "refresh"} + infra := &cobra.Command{Use: "infra"} + generate := &cobra.Command{Use: "generate", Aliases: []string{"gen", "synth"}} + env.AddCommand(refresh) + infra.AddCommand(generate) + root.AddCommand(up, extension, env, infra) + + assert.True(t, projectCommandSupportsExtensionAutoInstall(up)) + assert.True(t, projectCommandSupportsExtensionAutoInstall(refresh)) + assert.True(t, projectCommandSupportsExtensionAutoInstall(generate)) + assert.False(t, projectCommandSupportsExtensionAutoInstall(extension)) + assert.False(t, projectCommandSupportsExtensionAutoInstall(env)) +} + func TestFindFirstNonFlagArg(t *testing.T) { t.Parallel() // Mock flags that take values for testing diff --git a/cli/azd/cmd/project_extension_auto_install.go b/cli/azd/cmd/project_extension_auto_install.go new file mode 100644 index 00000000000..980495c8923 --- /dev/null +++ b/cli/azd/cmd/project_extension_auto_install.go @@ -0,0 +1,604 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +package cmd + +import ( + "cmp" + "context" + "errors" + "fmt" + "log" + "maps" + "slices" + "strings" + + "github.com/azure/azure-dev/cli/azd/internal" + "github.com/azure/azure-dev/cli/azd/pkg/extensions" + "github.com/azure/azure-dev/cli/azd/pkg/input" + "github.com/azure/azure-dev/cli/azd/pkg/ioc" + "github.com/azure/azure-dev/cli/azd/pkg/output" + "github.com/azure/azure-dev/cli/azd/pkg/output/ux" + "github.com/azure/azure-dev/cli/azd/pkg/project" + "github.com/spf13/cobra" +) + +type projectExtensionRequirement struct { + extension *extensions.ExtensionMetadata + versionPreference string + explicit bool +} + +type resolvedExtensionDependency struct { + parentId string + version string + capabilities []extensions.CapabilityType + providers []extensions.Provider +} + +func projectCommandSupportsExtensionAutoInstall(cmd *cobra.Command) bool { + if _, isExtensionCommand := cmd.Annotations["extension.id"]; isExtensionCommand { + return false + } + + path := getCommandPath(cmd) + if len(path) == 0 { + return false + } + + switch path[0] { + case "up", "provision", "deploy", "package", "restore", "down", "show", "monitor": + return true + case "infra": + return len(path) > 1 && path[1] == "generate" + case "env": + return len(path) > 1 && path[1] == "refresh" + default: + return false + } +} + +func findExtensionForProvider( + ctx context.Context, + console input.Console, + extensionManager extensionAutoInstallManager, + installed map[string]*extensions.Extension, + resolvedDependencies map[string]resolvedExtensionDependency, + requirementConflicts map[string]error, + capability extensions.CapabilityType, + provider string, +) (*extensions.ExtensionMetadata, error) { + matches, err := extensionManager.FindExtensions(ctx, &extensions.FilterOptions{ + Capability: capability, + Provider: provider, + }) + if err != nil { + return nil, fmt.Errorf("finding extension for provider %q: %w", provider, err) + } + matches = filterExtensionsForProvider(matches, capability, provider) + matchedRequirementConflicts := map[string]error{} + matches = slices.DeleteFunc(matches, func(extension *extensions.ExtensionMetadata) bool { + conflict, hasConflict := requirementConflicts[strings.ToLower(extension.Id)] + if hasConflict { + matchedRequirementConflicts[extension.Id] = conflict + } + return hasConflict + }) + dependencyConflicts := map[string]resolvedExtensionDependency{} + matches = slices.DeleteFunc(matches, func(extension *extensions.ExtensionMetadata) bool { + dependency, isDependency := resolvedDependencies[strings.ToLower(extension.Id)] + if isDependency { + dependencyConflicts[extension.Id] = dependency + } + return isDependency + }) + matches = uninstalledExtensionMatches(matches, installed) + if len(matches) == 0 { + if len(matchedRequirementConflicts) > 0 { + extensionId := slices.Sorted(maps.Keys(matchedRequirementConflicts))[0] + return nil, matchedRequirementConflicts[extensionId] + } + if len(dependencyConflicts) > 0 { + extensionId := slices.Sorted(maps.Keys(dependencyConflicts))[0] + dependency := dependencyConflicts[extensionId] + return nil, fmt.Errorf( + "extension %s requires dependency %s version %s, which does not provide %s %q", + dependency.parentId, + extensionId, + dependency.version, + capability, + provider, + ) + } + return nil, nil + } + + return promptForExtensionChoice(ctx, console, matches) +} + +func uninstalledExtensionMatches( + matches []*extensions.ExtensionMetadata, + installed map[string]*extensions.Extension, +) []*extensions.ExtensionMetadata { + return slices.DeleteFunc(slices.Clone(matches), func(extension *extensions.ExtensionMetadata) bool { + _, isInstalled := installedExtensionById(installed, extension.Id) + return isInstalled + }) +} + +func installedExtensionById( + installed map[string]*extensions.Extension, + extensionId string, +) (*extensions.Extension, bool) { + for installedId, extension := range installed { + if strings.EqualFold(installedId, extensionId) { + return extension, true + } + } + return nil, false +} + +func validateInstalledExtensionVersion( + installed *extensions.Extension, + versionPreference string, +) error { + if versionPreference == "" { + return nil + } + + installedMetadata := &extensions.ExtensionMetadata{ + Id: installed.Id, + Versions: []extensions.ExtensionVersion{{ + Version: installed.Version, + }}, + } + if _, err := extensions.ResolveExtensionVersion(installedMetadata, versionPreference, nil); err != nil { + return fmt.Errorf( + "installed extension %s version %s does not satisfy constraint %q", + installed.Id, + installed.Version, + versionPreference, + ) + } + return nil +} + +func resolveExtensionRequirementDependencies( + ctx context.Context, + extensionManager extensionAutoInstallManager, + requirements map[string]projectExtensionRequirement, +) (map[string]resolvedExtensionDependency, error) { + resolved := map[string]resolvedExtensionDependency{} + resolving := map[string]struct{}{} + + for _, requirement := range sortedProjectExtensionRequirements(requirements) { + version, err := extensions.ResolveExtensionVersion( + requirement.extension, + requirement.versionPreference, + nil, + ) + if err != nil { + return nil, fmt.Errorf("resolving required extension %s: %w", requirement.extension.Id, err) + } + + key := strings.ToLower(requirement.extension.Source + "\x00" + requirement.extension.Id) + resolving[key] = struct{}{} + err = resolveExtensionDependencies( + ctx, + extensionManager, + requirement.extension, + version.Dependencies, + resolved, + resolving, + ) + delete(resolving, key) + if err != nil { + return nil, err + } + } + + return resolved, nil +} + +func resolveExtensionDependencies( + ctx context.Context, + extensionManager extensionAutoInstallManager, + parent *extensions.ExtensionMetadata, + dependencies []extensions.ExtensionDependency, + resolved map[string]resolvedExtensionDependency, + resolving map[string]struct{}, +) error { + for _, dependency := range dependencies { + key := strings.ToLower(parent.Source + "\x00" + dependency.Id) + if _, isResolving := resolving[key]; isResolving { + return fmt.Errorf("dependency cycle detected involving extension %s", dependency.Id) + } + dependencyId := strings.ToLower(dependency.Id) + if _, isResolved := resolved[dependencyId]; isResolved { + continue + } + + // Installation reuses a compatible installed dependency instead of replacing it with the registry selection. + installedDependency, err := extensionManager.GetInstalled(extensions.FilterOptions{Id: dependency.Id}) + if err == nil && installedDependency != nil { + if dependency.Version != "" { + installedMetadata := &extensions.ExtensionMetadata{ + Id: dependency.Id, + Versions: []extensions.ExtensionVersion{{ + Version: installedDependency.Version, + }}, + } + if _, err := extensions.ResolveExtensionVersion(installedMetadata, dependency.Version, nil); err != nil { + return fmt.Errorf( + "installed dependency %s version %s does not satisfy constraint %q", + dependency.Id, + installedDependency.Version, + dependency.Version, + ) + } + } + + resolved[dependencyId] = resolvedExtensionDependency{ + parentId: parent.Id, + version: installedDependency.Version, + capabilities: installedDependency.Capabilities, + providers: installedDependency.Providers, + } + continue + } + + matches, err := extensionManager.FindExtensions(ctx, &extensions.FilterOptions{ + Id: dependency.Id, + Version: dependency.Version, + Source: parent.Source, + }) + if err != nil { + return fmt.Errorf("finding dependency %s: %w", dependency.Id, err) + } + if len(matches) == 0 { + return &extensions.DependencyNotFoundError{ + DependencyId: dependency.Id, + ParentId: parent.Id, + } + } + if len(matches) > 1 { + sources := make([]string, 0, len(matches)) + for _, match := range matches { + sources = append(sources, match.Source) + } + slices.Sort(sources) + sources = slices.Compact(sources) + return &extensions.DependencyAmbiguousSourceError{ + DependencyId: dependency.Id, + ParentId: parent.Id, + Sources: sources, + } + } + + dependencyExtension := matches[0] + version, err := extensions.ResolveExtensionVersion(dependencyExtension, dependency.Version, nil) + if err != nil { + return fmt.Errorf("resolving dependency %s: %w", dependency.Id, err) + } + resolved[dependencyId] = resolvedExtensionDependency{ + parentId: parent.Id, + version: version.Version, + capabilities: version.Capabilities, + providers: version.Providers, + } + + resolving[key] = struct{}{} + err = resolveExtensionDependencies( + ctx, + extensionManager, + dependencyExtension, + version.Dependencies, + resolved, + resolving, + ) + delete(resolving, key) + if err != nil { + return err + } + } + + return nil +} + +func extensionProvidesProvider( + capabilities []extensions.CapabilityType, + providers []extensions.Provider, + capability extensions.CapabilityType, + providerName string, +) bool { + expectedType, hasProviderType := providerTypeForCapability(capability) + if !hasProviderType || !slices.Contains(capabilities, capability) { + return false + } + + return slices.ContainsFunc(providers, func(provider extensions.Provider) bool { + return provider.Type == expectedType && strings.EqualFold(provider.Name, providerName) + }) +} + +func providerTypeForCapability(capability extensions.CapabilityType) (extensions.ProviderType, bool) { + switch capability { + case extensions.ServiceTargetProviderCapability: + return extensions.ServiceTargetProviderType, true + case extensions.ProvisioningProviderCapability: + return extensions.ProvisioningProviderType, true + default: + return "", false + } +} + +func filterExtensionsForProvider( + matches []*extensions.ExtensionMetadata, + capability extensions.CapabilityType, + providerName string, +) []*extensions.ExtensionMetadata { + filtered := make([]*extensions.ExtensionMetadata, 0, len(matches)) + for _, extension := range matches { + providerExtension := extensionForProvider(extension, capability, providerName) + if len(providerExtension.Versions) > 0 { + filtered = append(filtered, providerExtension) + } + } + return filtered +} + +func extensionVersionProvidesProvider( + version *extensions.ExtensionVersion, + capability extensions.CapabilityType, + providerName string, +) bool { + return extensionProvidesProvider(version.Capabilities, version.Providers, capability, providerName) +} + +func resolvedDependencyProvidesProvider( + dependency resolvedExtensionDependency, + capability extensions.CapabilityType, + providerName string, +) bool { + return extensionProvidesProvider( + dependency.capabilities, + dependency.providers, + capability, + providerName, + ) +} + +func extensionForProvider( + extension *extensions.ExtensionMetadata, + capability extensions.CapabilityType, + providerName string, +) *extensions.ExtensionMetadata { + filtered := *extension + filtered.Versions = slices.DeleteFunc(slices.Clone(extension.Versions), func(version extensions.ExtensionVersion) bool { + return !extensionVersionProvidesProvider(&version, capability, providerName) + }) + return &filtered +} + +func missingProjectExtensions( + ctx context.Context, + console input.Console, + extensionManager extensionAutoInstallManager, + projectConfig *project.ProjectConfig, +) ([]projectExtensionRequirement, error) { + installed, err := extensionManager.ListInstalled() + if err != nil { + return nil, fmt.Errorf("listing installed extensions: %w", err) + } + + requirements := map[string]projectExtensionRequirement{} + if projectConfig.RequiredVersions != nil { + for _, extensionId := range slices.Sorted(maps.Keys(projectConfig.RequiredVersions.Extensions)) { + versionPreference := "" + if constraint := projectConfig.RequiredVersions.Extensions[extensionId]; constraint != nil { + versionPreference = *constraint + } + if installedExtension, isInstalled := installedExtensionById(installed, extensionId); isInstalled { + if err := validateInstalledExtensionVersion(installedExtension, versionPreference); err != nil { + return nil, err + } + continue + } + + matches, err := extensionManager.FindExtensions(ctx, &extensions.FilterOptions{ + Id: extensionId, + Version: versionPreference, + }) + if err != nil { + return nil, fmt.Errorf("finding required extension %s: %w", extensionId, err) + } + if len(matches) == 0 { + return nil, fmt.Errorf("required extension %s not found", extensionId) + } + + extension, err := promptForExtensionChoice(ctx, console, matches) + if err != nil { + return nil, fmt.Errorf("selecting required extension %s: %w", extensionId, err) + } + + requirements[extension.Id] = projectExtensionRequirement{ + extension: extension, + versionPreference: versionPreference, + explicit: true, + } + } + } + + addProvider := func(capability extensions.CapabilityType, provider string) error { + if provider == "" { + return nil + } + + requirementConflicts := map[string]error{} + for _, extensionId := range slices.Sorted(maps.Keys(requirements)) { + requirement := requirements[extensionId] + selectedVersion, err := extensions.ResolveExtensionVersion( + requirement.extension, + requirement.versionPreference, + nil, + ) + if err != nil { + return fmt.Errorf("resolving required extension %s: %w", extensionId, err) + } + if extensionVersionProvidesProvider(selectedVersion, capability, provider) { + return nil + } + + if len(extensionForProvider(requirement.extension, capability, provider).Versions) == 0 { + continue + } + requirementConflicts[strings.ToLower(extensionId)] = fmt.Errorf( + "required extension %s version %s does not provide %s %q", + extensionId, + selectedVersion.Version, + capability, + provider, + ) + } + + resolvedDependencies, err := resolveExtensionRequirementDependencies(ctx, extensionManager, requirements) + if err != nil { + return err + } + for dependency := range maps.Values(resolvedDependencies) { + if resolvedDependencyProvidesProvider(dependency, capability, provider) { + return nil + } + } + + extension, err := findExtensionForProvider( + ctx, + console, + extensionManager, + installed, + resolvedDependencies, + requirementConflicts, + capability, + provider, + ) + if err != nil || extension == nil { + return err + } + if requirement, alreadyRequired := requirements[extension.Id]; alreadyRequired { + requirement.extension = extensionForProvider(requirement.extension, capability, provider) + if len(requirement.extension.Versions) == 0 { + return fmt.Errorf( + "required extension %s does not provide %s %q", + extension.Id, + capability, + provider, + ) + } + requirements[extension.Id] = requirement + } else { + requirements[extension.Id] = projectExtensionRequirement{ + extension: extension, + } + } + return nil + } + + for _, serviceName := range slices.Sorted(maps.Keys(projectConfig.Services)) { + if err := addProvider( + extensions.ServiceTargetProviderCapability, + string(projectConfig.Services[serviceName].Host), + ); err != nil { + return nil, err + } + } + + for _, infra := range projectConfig.Infra.GetLayers() { + if err := addProvider(extensions.ProvisioningProviderCapability, string(infra.Provider)); err != nil { + return nil, err + } + } + + return sortedProjectExtensionRequirements(requirements), nil +} + +func sortedProjectExtensionRequirements( + requirements map[string]projectExtensionRequirement, +) []projectExtensionRequirement { + result := slices.Collect(maps.Values(requirements)) + slices.SortFunc(result, func(a, b projectExtensionRequirement) int { + if a.explicit != b.explicit { + if a.explicit { + return -1 + } + return 1 + } + return cmp.Compare(a.extension.Id, b.extension.Id) + }) + + return result +} + +func tryAutoInstallProjectExtensions( + ctx context.Context, + rootContainer *ioc.NestedContainer, + foundCmd *cobra.Command, +) (handled bool, installed bool, err error) { + if !projectCommandSupportsExtensionAutoInstall(foundCmd) { + return false, false, nil + } + + var projectConfig *project.ProjectConfig + if err := rootContainer.Resolve(&projectConfig); err != nil { + log.Printf("skipping project extension auto-install: %v", err) + return false, false, nil + } + + var extensionManager *extensions.Manager + if err := rootContainer.Resolve(&extensionManager); err != nil { + return false, false, fmt.Errorf("resolving extension manager: %w", err) + } + var console input.Console + if err := rootContainer.Resolve(&console); err != nil { + return false, false, fmt.Errorf("resolving console: %w", err) + } + + requirements, err := missingProjectExtensions(ctx, console, extensionManager, projectConfig) + if err != nil { + return false, false, err + } + if len(requirements) == 0 { + return false, false, nil + } + + installedAny := false + for _, requirement := range requirements { + installed, err := tryAutoInstallExtensionVersion( + ctx, + console, + extensionManager, + *requirement.extension, + requirement.versionPreference, + ) + if err != nil { + return true, installedAny, err + } + installedAny = installedAny || installed + } + + return true, installedAny, nil +} + +func displayAutoInstallError(ctx context.Context, console input.Console, err error) { + if suggestionErr, ok := errors.AsType[*internal.ErrorWithSuggestion](err); ok { + console.Message(ctx, "") + console.MessageUxItem(ctx, &ux.ErrorWithSuggestion{ + Err: suggestionErr.Err, + Message: suggestionErr.Message, + Suggestion: suggestionErr.Suggestion, + Links: suggestionErr.Links, + }) + return + } + + console.Message(ctx, output.WithErrorFormat("\nERROR: %s", err.Error())) +} diff --git a/cli/azd/pkg/extensions/manager.go b/cli/azd/pkg/extensions/manager.go index 65fae091afc..53b6658c182 100644 --- a/cli/azd/pkg/extensions/manager.go +++ b/cli/azd/pkg/extensions/manager.go @@ -220,13 +220,17 @@ func bestSatisfyingVersionForAzd( return bestSatisfyingVersion(expr, compatible) } -// resolveExtensionVersion selects the best published version of extension that satisfies +// ResolveExtensionVersion selects the best published version of extension that satisfies // versionPreference and is compatible with azdVersion, or returns a descriptive error. -func resolveExtensionVersion( +func ResolveExtensionVersion( extension *ExtensionMetadata, versionPreference string, azdVersion *semver.Version, ) (*ExtensionVersion, error) { + if extension == nil { + return nil, fmt.Errorf("extension metadata cannot be nil") + } + selected := bestSatisfyingVersionForAzd(versionPreference, extension.Versions, azdVersion) if selected != nil { return selected, nil @@ -575,7 +579,7 @@ func (m *Manager) installInternal( } // Resolve to the latest published version that satisfies the preference. - selectedVersion, err := resolveExtensionVersion(extension, opts.VersionPreference, opts.AzdVersion) + selectedVersion, err := ResolveExtensionVersion(extension, opts.VersionPreference, opts.AzdVersion) if err != nil { return nil, err } @@ -871,7 +875,7 @@ func (m *Manager) ReconcileDependencies( return nil, nil, fmt.Errorf("extension metadata cannot be nil") } - selectedVersion, err := resolveExtensionVersion(extension, opts.VersionPreference, opts.AzdVersion) + selectedVersion, err := ResolveExtensionVersion(extension, opts.VersionPreference, opts.AzdVersion) if err != nil { return nil, nil, err } diff --git a/cli/azd/pkg/extensions/manager_test.go b/cli/azd/pkg/extensions/manager_test.go index 8ea9d337809..0b16a95d134 100644 --- a/cli/azd/pkg/extensions/manager_test.go +++ b/cli/azd/pkg/extensions/manager_test.go @@ -381,6 +381,13 @@ func Test_MatchesVersionConstraint(t *testing.T) { } } +func TestResolveExtensionVersionNil(t *testing.T) { + version, err := ResolveExtensionVersion(nil, "", nil) + + require.Nil(t, version) + require.EqualError(t, err, "extension metadata cannot be nil") +} + func Test_CreateExtensionFilter_VersionConstraints(t *testing.T) { ext := &ExtensionMetadata{ Id: "test.constraints",