From fcb9b7d4529ed9be2f495861c684cb41d835d6fc Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Mon, 20 Jul 2026 22:10:39 +0000 Subject: [PATCH 1/7] Initial plan From 02427a51923856abc96d2196dc8b2529884f8abb Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Mon, 20 Jul 2026 22:28:44 +0000 Subject: [PATCH 2/7] Auto-install project extension requirements Co-authored-by: JeffreyCA <9157833+JeffreyCA@users.noreply.github.com> --- cli/azd/cmd/auto_install.go | 228 ++++++++++++++++++++++++++++++- cli/azd/cmd/auto_install_test.go | 143 +++++++++++++++++++ 2 files changed, 368 insertions(+), 3 deletions(-) diff --git a/cli/azd/cmd/auto_install.go b/cli/azd/cmd/auto_install.go index 19c0eba16d9..07f7bfa2a79 100644 --- a/cli/azd/cmd/auto_install.go +++ b/cli/azd/cmd/auto_install.go @@ -4,11 +4,13 @@ package cmd import ( + "cmp" "context" "errors" "fmt" "io" "log" + "maps" "os" "slices" "strconv" @@ -338,9 +340,29 @@ 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{ Id: extension.Id, @@ -372,7 +394,7 @@ func tryAutoInstallExtension( Message: "Confirm installation", }) if err != nil { - return false, nil + return false, err } if !shouldInstall { @@ -381,7 +403,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) } @@ -390,6 +412,194 @@ func tryAutoInstallExtension( return true, nil } +type projectExtensionRequirement struct { + extension *extensions.ExtensionMetadata + versionPreference string + explicit bool +} + +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 "env": + return len(path) > 1 && path[1] == "refresh" + default: + return false + } +} + +func findExtensionForProvider( + ctx context.Context, + console input.Console, + extensionManager extensionAutoInstallManager, + capability extensions.CapabilityType, + provider string, +) (*extensions.ExtensionMetadata, error) { + matches, err := extensionManager.FindExtensions(ctx, &extensions.FilterOptions{ + Capability: capability, + Provider: provider, + }) + if err != nil { + log.Printf("failed to find an extension for provider %q: %v", provider, err) + return nil, nil + } + if len(matches) == 0 { + return nil, nil + } + + return promptForExtensionChoice(ctx, console, matches) +} + +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)) { + if _, isInstalled := installed[extensionId]; isInstalled { + continue + } + + versionPreference := "" + if constraint := projectConfig.RequiredVersions.Extensions[extensionId]; constraint != nil { + versionPreference = *constraint + } + 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 + } + + extension, err := findExtensionForProvider(ctx, console, extensionManager, capability, provider) + if err != nil || extension == nil { + return err + } + if _, isInstalled := installed[extension.Id]; isInstalled { + return nil + } + if _, alreadyRequired := requirements[extension.Id]; !alreadyRequired { + 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 + } + } + + 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, nil +} + +func tryAutoInstallProjectExtensions( + ctx context.Context, + rootContainer *ioc.NestedContainer, + foundCmd *cobra.Command, +) (bool, error) { + if !projectCommandSupportsExtensionAutoInstall(foundCmd) { + return false, nil + } + + var projectConfig *project.ProjectConfig + if err := rootContainer.Resolve(&projectConfig); err != nil { + log.Printf("skipping project extension auto-install: %v", err) + return false, nil + } + + var extensionManager *extensions.Manager + if err := rootContainer.Resolve(&extensionManager); err != nil { + return false, fmt.Errorf("resolving extension manager: %w", err) + } + var console input.Console + if err := rootContainer.Resolve(&console); err != nil { + return false, fmt.Errorf("resolving console: %w", err) + } + + requirements, err := missingProjectExtensions(ctx, console, extensionManager, projectConfig) + if err != nil { + return false, err + } + + installedAny := false + for _, requirement := range requirements { + installed, err := tryAutoInstallExtensionVersion( + ctx, + console, + extensionManager, + *requirement.extension, + requirement.versionPreference, + ) + if err != nil { + return installedAny, err + } + installedAny = installedAny || installed + } + + return installedAny, nil +} + // startUpdateCheck launches a background goroutine that checks for a newer // version of azd and returns a channel that will receive the result. // The caller should read from the returned channel after command execution. @@ -495,6 +705,18 @@ func ExecuteWithAutoInstall(ctx context.Context, rootContainer *ioc.NestedContai result.LatestVersion = startUpdateCheck(ctx) } + if installed, err := tryAutoInstallProjectExtensions(ctx, rootContainer, foundCmd); err != nil { + result.Err = err + return result + } else if installed { + rootCmd = NewRootCmd(false, nil, 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, diff --git a/cli/azd/cmd/auto_install_test.go b/cli/azd/cmd/auto_install_test.go index 2b738c66096..529ba7cf81a 100644 --- a/cli/azd/cmd/auto_install_test.go +++ b/cli/azd/cmd/auto_install_test.go @@ -4,7 +4,9 @@ package cmd import ( + "context" "fmt" + "slices" "strings" "testing" @@ -15,11 +17,152 @@ 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 +} + +func (m *fakeExtensionAutoInstallManager) FindExtensions( + _ context.Context, + options *extensions.FilterOptions, +) ([]*extensions.ExtensionMetadata, error) { + var matches []*extensions.ExtensionMetadata + for _, extension := range m.available { + if options.Id != "" && extension.Id != options.Id { + continue + } + if options.Capability != "" && !slices.ContainsFunc(extension.Versions, func(version extensions.ExtensionVersion) bool { + return slices.Contains(version.Capabilities, options.Capability) && + slices.ContainsFunc(version.Providers, func(provider extensions.Provider) bool { + return provider.Name == options.Provider + }) + }) { + 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: "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, + }}, + Dependencies: []extensions.ExtensionDependency{ + {Id: "azure.ai.projects"}, + {Id: "azure.ai.agents"}, + }, + }}, + }, + }, + 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) +} + +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"} + env.AddCommand(refresh) + root.AddCommand(up, extension, env) + + assert.True(t, projectCommandSupportsExtensionAutoInstall(up)) + assert.True(t, projectCommandSupportsExtensionAutoInstall(refresh)) + assert.False(t, projectCommandSupportsExtensionAutoInstall(extension)) + assert.False(t, projectCommandSupportsExtensionAutoInstall(env)) +} + func TestFindFirstNonFlagArg(t *testing.T) { t.Parallel() // Mock flags that take values for testing From 6f568aadd1642d225c6bcb895d8bf1107ce46a51 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Mon, 20 Jul 2026 22:38:32 +0000 Subject: [PATCH 3/7] Handle provider version compatibility Co-authored-by: JeffreyCA <9157833+JeffreyCA@users.noreply.github.com> --- cli/azd/cmd/auto_install.go | 34 +++++++++++++++++++++++-- cli/azd/cmd/auto_install_test.go | 43 +++++++++++++++++++++----------- 2 files changed, 60 insertions(+), 17 deletions(-) diff --git a/cli/azd/cmd/auto_install.go b/cli/azd/cmd/auto_install.go index 07f7bfa2a79..d7030e31417 100644 --- a/cli/azd/cmd/auto_install.go +++ b/cli/azd/cmd/auto_install.go @@ -431,6 +431,8 @@ func projectCommandSupportsExtensionAutoInstall(cmd *cobra.Command) bool { 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: @@ -460,6 +462,21 @@ func findExtensionForProvider( return promptForExtensionChoice(ctx, console, matches) } +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 !slices.Contains(version.Capabilities, capability) || + !slices.ContainsFunc(version.Providers, func(provider extensions.Provider) bool { + return strings.EqualFold(provider.Name, providerName) + }) + }) + return &filtered +} + func missingProjectExtensions( ctx context.Context, console input.Console, @@ -518,8 +535,21 @@ func missingProjectExtensions( if _, isInstalled := installed[extension.Id]; isInstalled { return nil } - if _, alreadyRequired := requirements[extension.Id]; !alreadyRequired { - requirements[extension.Id] = projectExtensionRequirement{extension: extension} + 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: extensionForProvider(extension, capability, provider), + } } return nil } diff --git a/cli/azd/cmd/auto_install_test.go b/cli/azd/cmd/auto_install_test.go index 529ba7cf81a..65e32b93950 100644 --- a/cli/azd/cmd/auto_install_test.go +++ b/cli/azd/cmd/auto_install_test.go @@ -38,12 +38,16 @@ func (m *fakeExtensionAutoInstallManager) FindExtensions( if options.Id != "" && extension.Id != options.Id { continue } - if options.Capability != "" && !slices.ContainsFunc(extension.Versions, func(version extensions.ExtensionVersion) bool { - return slices.Contains(version.Capabilities, options.Capability) && - slices.ContainsFunc(version.Providers, func(provider extensions.Provider) bool { - return provider.Name == options.Provider - }) - }) { + hasCapabilityAndProvider := slices.ContainsFunc( + extension.Versions, + func(version extensions.ExtensionVersion) bool { + return slices.Contains(version.Capabilities, options.Capability) && + slices.ContainsFunc(version.Providers, func(provider extensions.Provider) bool { + return provider.Name == options.Provider + }) + }, + ) + if options.Capability != "" && !hasCapabilityAndProvider { continue } matches = append(matches, extension) @@ -83,14 +87,17 @@ func TestMissingProjectExtensions(t *testing.T) { available: []*extensions.ExtensionMetadata{ { Id: "azure.ai.projects", - Versions: []extensions.ExtensionVersion{{ - Version: "1.0.0", - Capabilities: []extensions.CapabilityType{extensions.ServiceTargetProviderCapability}, - Providers: []extensions.Provider{{ - Name: "azure.ai.project", - Type: extensions.ServiceTargetProviderType, - }}, - }}, + 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", @@ -146,6 +153,8 @@ func TestMissingProjectExtensions(t *testing.T) { 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 TestProjectCommandSupportsExtensionAutoInstall(t *testing.T) { @@ -154,11 +163,15 @@ func TestProjectCommandSupportsExtensionAutoInstall(t *testing.T) { 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) - root.AddCommand(up, extension, env) + 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)) } From 1b7f158c3eb2a891ab4e3869aaf39dd5bf20627a Mon Sep 17 00:00:00 2001 From: Jeffrey Chen Date: Wed, 22 Jul 2026 21:26:25 +0000 Subject: [PATCH 4/7] Avoid redundant project extension prompts Resolve exact same-source dependency versions during preflight so extension packs cover inferred providers without extra source prompts. Skip installed extension IDs before selection, reuse source choices, and surface preflight failures. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: c802f692-befd-4e8b-aca8-152699d438a0 --- cli/azd/cmd/auto_install.go | 232 ++++++++++++++++++++- cli/azd/cmd/auto_install_test.go | 275 ++++++++++++++++++++++++- cli/azd/pkg/extensions/manager.go | 12 +- cli/azd/pkg/extensions/manager_test.go | 7 + 4 files changed, 510 insertions(+), 16 deletions(-) diff --git a/cli/azd/cmd/auto_install.go b/cli/azd/cmd/auto_install.go index d7030e31417..1e67359e46b 100644 --- a/cli/azd/cmd/auto_install.go +++ b/cli/azd/cmd/auto_install.go @@ -418,6 +418,11 @@ type projectExtensionRequirement struct { explicit bool } +type resolvedExtensionDependency struct { + parentId string + version *extensions.ExtensionVersion +} + func projectCommandSupportsExtensionAutoInstall(cmd *cobra.Command) bool { if _, isExtensionCommand := cmd.Annotations["extension.id"]; isExtensionCommand { return false @@ -444,6 +449,8 @@ func findExtensionForProvider( ctx context.Context, console input.Console, extensionManager extensionAutoInstallManager, + installed map[string]*extensions.Extension, + resolvedDependencies map[string]resolvedExtensionDependency, capability extensions.CapabilityType, provider string, ) (*extensions.ExtensionMetadata, error) { @@ -455,13 +462,165 @@ func findExtensionForProvider( log.Printf("failed to find an extension for provider %q: %v", provider, err) return nil, nil } + matches = uninstalledExtensionMatches(matches, installed) + dependencyConflicts := map[string]resolvedExtensionDependency{} + matches = slices.DeleteFunc(matches, func(extension *extensions.ExtensionMetadata) bool { + dependency, isDependency := resolvedDependencies[extension.Id] + if isDependency { + dependencyConflicts[extension.Id] = dependency + } + return isDependency + }) if len(matches) == 0 { + 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.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 := installed[extension.Id] + return isInstalled + }) +} + +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) + } + if _, isResolved := resolved[dependency.Id]; isResolved { + 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[dependency.Id] = resolvedExtensionDependency{ + parentId: parent.Id, + version: version, + } + + resolving[key] = struct{}{} + err = resolveExtensionDependencies( + ctx, + extensionManager, + dependencyExtension, + version.Dependencies, + resolved, + resolving, + ) + delete(resolving, key) + if err != nil { + return err + } + } + + return nil +} + +func extensionVersionProvidesProvider( + version *extensions.ExtensionVersion, + capability extensions.CapabilityType, + providerName string, +) bool { + return slices.Contains(version.Capabilities, capability) && + slices.ContainsFunc(version.Providers, func(provider extensions.Provider) bool { + return strings.EqualFold(provider.Name, providerName) + }) +} + func extensionForProvider( extension *extensions.ExtensionMetadata, capability extensions.CapabilityType, @@ -528,12 +687,39 @@ func missingProjectExtensions( return nil } - extension, err := findExtensionForProvider(ctx, console, extensionManager, capability, provider) - if err != nil || extension == nil { + for _, extensionId := range slices.Sorted(maps.Keys(requirements)) { + requirement := requirements[extensionId] + extension := extensionForProvider(requirement.extension, capability, provider) + if len(extension.Versions) == 0 { + continue + } + + requirement.extension = extension + requirements[extensionId] = requirement + return nil + } + + resolvedDependencies, err := resolveExtensionRequirementDependencies(ctx, extensionManager, requirements) + if err != nil { return err } - if _, isInstalled := installed[extension.Id]; isInstalled { - return nil + for dependency := range maps.Values(resolvedDependencies) { + if extensionVersionProvidesProvider(dependency.version, capability, provider) { + return nil + } + } + + extension, err := findExtensionForProvider( + ctx, + console, + extensionManager, + installed, + resolvedDependencies, + capability, + provider, + ) + if err != nil || extension == nil { + return err } if requirement, alreadyRequired := requirements[extension.Id]; alreadyRequired { requirement.extension = extensionForProvider(requirement.extension, capability, provider) @@ -569,6 +755,12 @@ func missingProjectExtensions( } } + 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 { @@ -580,7 +772,7 @@ func missingProjectExtensions( return cmp.Compare(a.extension.Id, b.extension.Id) }) - return result, nil + return result } func tryAutoInstallProjectExtensions( @@ -630,6 +822,21 @@ func tryAutoInstallProjectExtensions( return 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())) +} + // startUpdateCheck launches a background goroutine that checks for a newer // version of azd and returns a channel that will receive the result. // The caller should read from the returned channel after command execution. @@ -736,6 +943,11 @@ func ExecuteWithAutoInstall(ctx context.Context, rootContainer *ioc.NestedContai } if installed, err := tryAutoInstallProjectExtensions(ctx, rootContainer, foundCmd); 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 { @@ -786,9 +998,13 @@ 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 = 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) diff --git a/cli/azd/cmd/auto_install_test.go b/cli/azd/cmd/auto_install_test.go index 65e32b93950..2aa98edbb15 100644 --- a/cli/azd/cmd/auto_install_test.go +++ b/cli/azd/cmd/auto_install_test.go @@ -38,6 +38,14 @@ func (m *fakeExtensionAutoInstallManager) FindExtensions( 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 + } + } hasCapabilityAndProvider := slices.ContainsFunc( extension.Versions, func(version extensions.ExtensionVersion) bool { @@ -119,10 +127,6 @@ func TestMissingProjectExtensions(t *testing.T) { Name: "microsoft.foundry", Type: extensions.ProvisioningProviderType, }}, - Dependencies: []extensions.ExtensionDependency{ - {Id: "azure.ai.projects"}, - {Id: "azure.ai.agents"}, - }, }}, }, }, @@ -157,6 +161,269 @@ func TestMissingProjectExtensions(t *testing.T) { 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"}}, + }}, + }, + { + Id: "microsoft.azd.demo", + Source: "local", + Versions: []extensions.ExtensionVersion{{ + Version: "0.7.0", + Capabilities: []extensions.CapabilityType{extensions.ServiceTargetProviderCapability}, + Providers: []extensions.Provider{{Name: "demo"}}, + }}, + }, + }, + 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"}}, + }}, + }, + }, + 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"}}, + }, + }, + }, + }, + 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 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"} 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", From eefc3f017bbb8a28403ff25b5f56a0a7475d465f Mon Sep 17 00:00:00 2001 From: Jeffrey Chen Date: Wed, 22 Jul 2026 21:51:01 +0000 Subject: [PATCH 5/7] Address extension auto-install review feedback Require capability, provider name, and provider type on the same extension version; honor installed dependency metadata; normalize installed IDs; and preserve container singletons when rebuilding command bindings. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: c802f692-befd-4e8b-aca8-152699d438a0 --- cli/azd/cmd/auto_install.go | 151 +++++++++++++++++++---- cli/azd/cmd/auto_install_test.go | 203 ++++++++++++++++++++++++++++--- 2 files changed, 315 insertions(+), 39 deletions(-) diff --git a/cli/azd/cmd/auto_install.go b/cli/azd/cmd/auto_install.go index 1e67359e46b..033582bbcdd 100644 --- a/cli/azd/cmd/auto_install.go +++ b/cli/azd/cmd/auto_install.go @@ -419,8 +419,10 @@ type projectExtensionRequirement struct { } type resolvedExtensionDependency struct { - parentId string - version *extensions.ExtensionVersion + parentId string + version string + capabilities []extensions.CapabilityType + providers []extensions.Provider } func projectCommandSupportsExtensionAutoInstall(cmd *cobra.Command) bool { @@ -462,15 +464,16 @@ func findExtensionForProvider( log.Printf("failed to find an extension for provider %q: %v", provider, err) return nil, nil } - matches = uninstalledExtensionMatches(matches, installed) + matches = filterExtensionsForProvider(matches, capability, provider) dependencyConflicts := map[string]resolvedExtensionDependency{} matches = slices.DeleteFunc(matches, func(extension *extensions.ExtensionMetadata) bool { - dependency, isDependency := resolvedDependencies[extension.Id] + 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(dependencyConflicts) > 0 { extensionId := slices.Sorted(maps.Keys(dependencyConflicts))[0] @@ -479,7 +482,7 @@ func findExtensionForProvider( "extension %s requires dependency %s version %s, which does not provide %s %q", dependency.parentId, extensionId, - dependency.version.Version, + dependency.version, capability, provider, ) @@ -495,11 +498,23 @@ func uninstalledExtensionMatches( installed map[string]*extensions.Extension, ) []*extensions.ExtensionMetadata { return slices.DeleteFunc(slices.Clone(matches), func(extension *extensions.ExtensionMetadata) bool { - _, isInstalled := installed[extension.Id] + _, 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 resolveExtensionRequirementDependencies( ctx context.Context, extensionManager extensionAutoInstallManager, @@ -550,7 +565,37 @@ func resolveExtensionDependencies( if _, isResolving := resolving[key]; isResolving { return fmt.Errorf("dependency cycle detected involving extension %s", dependency.Id) } - if _, isResolved := resolved[dependency.Id]; isResolved { + 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 } @@ -587,9 +632,11 @@ func resolveExtensionDependencies( if err != nil { return fmt.Errorf("resolving dependency %s: %w", dependency.Id, err) } - resolved[dependency.Id] = resolvedExtensionDependency{ - parentId: parent.Id, - version: version, + resolved[dependencyId] = resolvedExtensionDependency{ + parentId: parent.Id, + version: version.Version, + capabilities: version.Capabilities, + providers: version.Providers, } resolving[key] = struct{}{} @@ -610,15 +657,67 @@ func resolveExtensionDependencies( 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 slices.Contains(version.Capabilities, capability) && - slices.ContainsFunc(version.Providers, func(provider extensions.Provider) bool { - return strings.EqualFold(provider.Name, providerName) - }) + 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( @@ -628,10 +727,7 @@ func extensionForProvider( ) *extensions.ExtensionMetadata { filtered := *extension filtered.Versions = slices.DeleteFunc(slices.Clone(extension.Versions), func(version extensions.ExtensionVersion) bool { - return !slices.Contains(version.Capabilities, capability) || - !slices.ContainsFunc(version.Providers, func(provider extensions.Provider) bool { - return strings.EqualFold(provider.Name, providerName) - }) + return !extensionVersionProvidesProvider(&version, capability, providerName) }) return &filtered } @@ -650,7 +746,7 @@ func missingProjectExtensions( requirements := map[string]projectExtensionRequirement{} if projectConfig.RequiredVersions != nil { for _, extensionId := range slices.Sorted(maps.Keys(projectConfig.RequiredVersions.Extensions)) { - if _, isInstalled := installed[extensionId]; isInstalled { + if _, isInstalled := installedExtensionById(installed, extensionId); isInstalled { continue } @@ -704,7 +800,7 @@ func missingProjectExtensions( return err } for dependency := range maps.Values(resolvedDependencies) { - if extensionVersionProvidesProvider(dependency.version, capability, provider) { + if resolvedDependencyProvidesProvider(dependency, capability, provider) { return nil } } @@ -734,7 +830,7 @@ func missingProjectExtensions( requirements[extension.Id] = requirement } else { requirements[extension.Id] = projectExtensionRequirement{ - extension: extensionForProvider(extension, capability, provider), + extension: extension, } } return nil @@ -951,7 +1047,7 @@ func ExecuteWithAutoInstall(ctx context.Context, rootContainer *ioc.NestedContai result.Err = err return result } else if installed { - rootCmd = NewRootCmd(false, nil, rootContainer) + rootCmd = newRootCmdWithoutRegistration(rootContainer) foundCmd, originalArgs, err = rootCmd.Find(os.Args[1:]) if err != nil { result.Err = err @@ -964,7 +1060,7 @@ func ExecuteWithAutoInstall(ctx context.Context, rootContainer *ioc.NestedContai 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 } @@ -1004,6 +1100,11 @@ func ExecuteWithAutoInstall(ctx context.Context, rootContainer *ioc.NestedContai 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 @@ -1040,7 +1141,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 } @@ -1167,7 +1268,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 2aa98edbb15..a3f0e982bd7 100644 --- a/cli/azd/cmd/auto_install_test.go +++ b/cli/azd/cmd/auto_install_test.go @@ -46,16 +46,18 @@ func (m *fakeExtensionAutoInstallManager) FindExtensions( continue } } - hasCapabilityAndProvider := slices.ContainsFunc( - extension.Versions, - func(version extensions.ExtensionVersion) bool { - return slices.Contains(version.Capabilities, options.Capability) && - slices.ContainsFunc(version.Providers, func(provider extensions.Provider) bool { - return provider.Name == options.Provider - }) - }, - ) - if options.Capability != "" && !hasCapabilityAndProvider { + 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) @@ -170,7 +172,10 @@ func TestMissingProjectExtensionsSkipsInstalledProviderAcrossSources(t *testing. Versions: []extensions.ExtensionVersion{{ Version: "0.7.0", Capabilities: []extensions.CapabilityType{extensions.ServiceTargetProviderCapability}, - Providers: []extensions.Provider{{Name: "demo"}}, + Providers: []extensions.Provider{{ + Name: "demo", + Type: extensions.ServiceTargetProviderType, + }}, }}, }, { @@ -179,7 +184,10 @@ func TestMissingProjectExtensionsSkipsInstalledProviderAcrossSources(t *testing. Versions: []extensions.ExtensionVersion{{ Version: "0.7.0", Capabilities: []extensions.CapabilityType{extensions.ServiceTargetProviderCapability}, - Providers: []extensions.Provider{{Name: "demo"}}, + Providers: []extensions.Provider{{ + Name: "demo", + Type: extensions.ServiceTargetProviderType, + }}, }}, }, }, @@ -317,7 +325,10 @@ func TestMissingProjectExtensionsSkipsExtensionPackDependencies(t *testing.T) { Versions: []extensions.ExtensionVersion{{ Version: "1.0.0", Capabilities: []extensions.CapabilityType{extensions.ServiceTargetProviderCapability}, - Providers: []extensions.Provider{{Name: "azure.ai.project"}}, + Providers: []extensions.Provider{{ + Name: "azure.ai.project", + Type: extensions.ServiceTargetProviderType, + }}, }}, }, }, @@ -370,7 +381,10 @@ func TestMissingProjectExtensionsRejectsDependencyVersionWithoutProvider(t *test { Version: "2.0.0", Capabilities: []extensions.CapabilityType{extensions.ServiceTargetProviderCapability}, - Providers: []extensions.Provider{{Name: "demo"}}, + Providers: []extensions.Provider{{ + Name: "demo", + Type: extensions.ServiceTargetProviderType, + }}, }, }, }, @@ -400,6 +414,167 @@ func TestMissingProjectExtensionsRejectsDependencyVersionWithoutProvider(t *test 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 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 TestDisplayAutoInstallError(t *testing.T) { t.Run("RegularError", func(t *testing.T) { console := mockinput.NewMockConsole() From 2e737b4f84e875d108a6a7edb077fb3754e6087d Mon Sep 17 00:00:00 2001 From: Jeffrey Chen Date: Wed, 22 Jul 2026 21:56:26 +0000 Subject: [PATCH 6/7] Split project extension preflight logic Move project-specific extension discovery, dependency resolution, and preflight error rendering into a dedicated command file without changing behavior. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: c802f692-befd-4e8b-aca8-152699d438a0 --- cli/azd/cmd/auto_install.go | 523 ----------------- cli/azd/cmd/project_extension_auto_install.go | 545 ++++++++++++++++++ 2 files changed, 545 insertions(+), 523 deletions(-) create mode 100644 cli/azd/cmd/project_extension_auto_install.go diff --git a/cli/azd/cmd/auto_install.go b/cli/azd/cmd/auto_install.go index 033582bbcdd..60c07352c55 100644 --- a/cli/azd/cmd/auto_install.go +++ b/cli/azd/cmd/auto_install.go @@ -4,13 +4,11 @@ package cmd import ( - "cmp" "context" "errors" "fmt" "io" "log" - "maps" "os" "slices" "strconv" @@ -412,527 +410,6 @@ func tryAutoInstallExtensionVersion( return true, nil } -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, - capability extensions.CapabilityType, - provider string, -) (*extensions.ExtensionMetadata, error) { - matches, err := extensionManager.FindExtensions(ctx, &extensions.FilterOptions{ - Capability: capability, - Provider: provider, - }) - if err != nil { - log.Printf("failed to find an extension for provider %q: %v", provider, err) - return nil, nil - } - matches = filterExtensionsForProvider(matches, capability, provider) - 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(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 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)) { - if _, isInstalled := installedExtensionById(installed, extensionId); isInstalled { - continue - } - - versionPreference := "" - if constraint := projectConfig.RequiredVersions.Extensions[extensionId]; constraint != nil { - versionPreference = *constraint - } - 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 - } - - for _, extensionId := range slices.Sorted(maps.Keys(requirements)) { - requirement := requirements[extensionId] - extension := extensionForProvider(requirement.extension, capability, provider) - if len(extension.Versions) == 0 { - continue - } - - requirement.extension = extension - requirements[extensionId] = requirement - return nil - } - - 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, - 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, -) (bool, error) { - if !projectCommandSupportsExtensionAutoInstall(foundCmd) { - return false, nil - } - - var projectConfig *project.ProjectConfig - if err := rootContainer.Resolve(&projectConfig); err != nil { - log.Printf("skipping project extension auto-install: %v", err) - return false, nil - } - - var extensionManager *extensions.Manager - if err := rootContainer.Resolve(&extensionManager); err != nil { - return false, fmt.Errorf("resolving extension manager: %w", err) - } - var console input.Console - if err := rootContainer.Resolve(&console); err != nil { - return false, fmt.Errorf("resolving console: %w", err) - } - - requirements, err := missingProjectExtensions(ctx, console, extensionManager, projectConfig) - if err != nil { - return false, err - } - - installedAny := false - for _, requirement := range requirements { - installed, err := tryAutoInstallExtensionVersion( - ctx, - console, - extensionManager, - *requirement.extension, - requirement.versionPreference, - ) - if err != nil { - return installedAny, err - } - installedAny = installedAny || installed - } - - return 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())) -} - // startUpdateCheck launches a background goroutine that checks for a newer // version of azd and returns a channel that will receive the result. // The caller should read from the returned channel after command execution. 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..f80fed2c7d3 --- /dev/null +++ b/cli/azd/cmd/project_extension_auto_install.go @@ -0,0 +1,545 @@ +// 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, + capability extensions.CapabilityType, + provider string, +) (*extensions.ExtensionMetadata, error) { + matches, err := extensionManager.FindExtensions(ctx, &extensions.FilterOptions{ + Capability: capability, + Provider: provider, + }) + if err != nil { + log.Printf("failed to find an extension for provider %q: %v", provider, err) + return nil, nil + } + matches = filterExtensionsForProvider(matches, capability, provider) + 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(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 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)) { + if _, isInstalled := installedExtensionById(installed, extensionId); isInstalled { + continue + } + + versionPreference := "" + if constraint := projectConfig.RequiredVersions.Extensions[extensionId]; constraint != nil { + versionPreference = *constraint + } + 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 + } + + for _, extensionId := range slices.Sorted(maps.Keys(requirements)) { + requirement := requirements[extensionId] + extension := extensionForProvider(requirement.extension, capability, provider) + if len(extension.Versions) == 0 { + continue + } + + requirement.extension = extension + requirements[extensionId] = requirement + return nil + } + + 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, + 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, +) (bool, error) { + if !projectCommandSupportsExtensionAutoInstall(foundCmd) { + return false, nil + } + + var projectConfig *project.ProjectConfig + if err := rootContainer.Resolve(&projectConfig); err != nil { + log.Printf("skipping project extension auto-install: %v", err) + return false, nil + } + + var extensionManager *extensions.Manager + if err := rootContainer.Resolve(&extensionManager); err != nil { + return false, fmt.Errorf("resolving extension manager: %w", err) + } + var console input.Console + if err := rootContainer.Resolve(&console); err != nil { + return false, fmt.Errorf("resolving console: %w", err) + } + + requirements, err := missingProjectExtensions(ctx, console, extensionManager, projectConfig) + if err != nil { + return false, err + } + + installedAny := false + for _, requirement := range requirements { + installed, err := tryAutoInstallExtensionVersion( + ctx, + console, + extensionManager, + *requirement.extension, + requirement.versionPreference, + ) + if err != nil { + return installedAny, err + } + installedAny = installedAny || installed + } + + return 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())) +} From 7251bc4071939c6789851c6cc33e34e0b2b253d5 Mon Sep 17 00:00:00 2001 From: Jeffrey Chen Date: Wed, 22 Jul 2026 22:16:41 +0000 Subject: [PATCH 7/7] Handle project extension preflight edge cases Honor --cwd during command-tree construction, propagate provider lookup failures, validate constrained and installed extension versions, and avoid repeating a declined preflight prompt in the legacy fallback. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: c802f692-befd-4e8b-aca8-152699d438a0 --- cli/azd/cmd/auto_install.go | 66 +++++++- cli/azd/cmd/auto_install_test.go | 157 ++++++++++++++++++ cli/azd/cmd/project_extension_auto_install.go | 99 ++++++++--- 3 files changed, 298 insertions(+), 24 deletions(-) diff --git a/cli/azd/cmd/auto_install.go b/cli/azd/cmd/auto_install.go index 60c07352c55..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" @@ -362,10 +363,13 @@ func tryAutoInstallExtensionVersion( 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 } @@ -472,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{} @@ -493,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 @@ -503,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" @@ -515,7 +562,9 @@ func ExecuteWithAutoInstall(ctx context.Context, rootContainer *ioc.NestedContai result.LatestVersion = startUpdateCheck(ctx) } - if installed, err := tryAutoInstallProjectExtensions(ctx, rootContainer, foundCmd); err != nil { + 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 { @@ -543,7 +592,7 @@ func ExecuteWithAutoInstall(ctx context.Context, rootContainer *ioc.NestedContai } // 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. @@ -552,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) diff --git a/cli/azd/cmd/auto_install_test.go b/cli/azd/cmd/auto_install_test.go index a3f0e982bd7..2a9fe6c6d72 100644 --- a/cli/azd/cmd/auto_install_test.go +++ b/cli/azd/cmd/auto_install_test.go @@ -6,6 +6,7 @@ package cmd import ( "context" "fmt" + "path/filepath" "slices" "strings" "testing" @@ -27,12 +28,17 @@ import ( 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 { @@ -541,6 +547,131 @@ func TestMissingProjectExtensionsInstalledIdIsCaseInsensitive(t *testing.T) { 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{ @@ -575,6 +706,32 @@ func TestExtensionVersionProvidesProviderMatchesType(t *testing.T) { )) } +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() diff --git a/cli/azd/cmd/project_extension_auto_install.go b/cli/azd/cmd/project_extension_auto_install.go index f80fed2c7d3..980495c8923 100644 --- a/cli/azd/cmd/project_extension_auto_install.go +++ b/cli/azd/cmd/project_extension_auto_install.go @@ -64,6 +64,7 @@ func findExtensionForProvider( extensionManager extensionAutoInstallManager, installed map[string]*extensions.Extension, resolvedDependencies map[string]resolvedExtensionDependency, + requirementConflicts map[string]error, capability extensions.CapabilityType, provider string, ) (*extensions.ExtensionMetadata, error) { @@ -72,10 +73,17 @@ func findExtensionForProvider( Provider: provider, }) if err != nil { - log.Printf("failed to find an extension for provider %q: %v", provider, err) - return nil, 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)] @@ -86,6 +94,10 @@ func findExtensionForProvider( }) 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] @@ -126,6 +138,31 @@ func installedExtensionById( 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, @@ -357,14 +394,17 @@ func missingProjectExtensions( requirements := map[string]projectExtensionRequirement{} if projectConfig.RequiredVersions != nil { for _, extensionId := range slices.Sorted(maps.Keys(projectConfig.RequiredVersions.Extensions)) { - if _, isInstalled := installedExtensionById(installed, extensionId); isInstalled { - continue - } - 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, @@ -394,16 +434,31 @@ func missingProjectExtensions( return nil } + requirementConflicts := map[string]error{} for _, extensionId := range slices.Sorted(maps.Keys(requirements)) { requirement := requirements[extensionId] - extension := extensionForProvider(requirement.extension, capability, provider) - if len(extension.Versions) == 0 { - continue + 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 } - requirement.extension = extension - requirements[extensionId] = requirement - 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) @@ -422,6 +477,7 @@ func missingProjectExtensions( extensionManager, installed, resolvedDependencies, + requirementConflicts, capability, provider, ) @@ -486,29 +542,32 @@ func tryAutoInstallProjectExtensions( ctx context.Context, rootContainer *ioc.NestedContainer, foundCmd *cobra.Command, -) (bool, error) { +) (handled bool, installed bool, err error) { if !projectCommandSupportsExtensionAutoInstall(foundCmd) { - return false, nil + 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, nil + return false, false, nil } var extensionManager *extensions.Manager if err := rootContainer.Resolve(&extensionManager); err != nil { - return false, fmt.Errorf("resolving extension manager: %w", err) + return false, false, fmt.Errorf("resolving extension manager: %w", err) } var console input.Console if err := rootContainer.Resolve(&console); err != nil { - return false, fmt.Errorf("resolving console: %w", err) + return false, false, fmt.Errorf("resolving console: %w", err) } requirements, err := missingProjectExtensions(ctx, console, extensionManager, projectConfig) if err != nil { - return false, err + return false, false, err + } + if len(requirements) == 0 { + return false, false, nil } installedAny := false @@ -521,12 +580,12 @@ func tryAutoInstallProjectExtensions( requirement.versionPreference, ) if err != nil { - return installedAny, err + return true, installedAny, err } installedAny = installedAny || installed } - return installedAny, nil + return true, installedAny, nil } func displayAutoInstallError(ctx context.Context, console input.Console, err error) {