fix: Npm dependency resolver

This commit is contained in:
abhisek
2025-05-14 09:20:44 +05:30
parent e294936b0f
commit 0883140732
3 changed files with 144 additions and 75 deletions
+10 -4
View File
@@ -2,15 +2,20 @@ package packagemanager
import ( import (
"fmt" "fmt"
"slices"
"strings" "strings"
packagev1 "buf.build/gen/go/safedep/api/protocolbuffers/go/safedep/messages/package/v1" packagev1 "buf.build/gen/go/safedep/api/protocolbuffers/go/safedep/messages/package/v1"
) )
type NpmPackageManagerConfig struct{} type NpmPackageManagerConfig struct {
InstallCommands []string
}
func DefaultNpmPackageManagerConfig() NpmPackageManagerConfig { func DefaultNpmPackageManagerConfig() NpmPackageManagerConfig {
return NpmPackageManagerConfig{} return NpmPackageManagerConfig{
InstallCommands: []string{"install", "i", "add"},
}
} }
type npmPackageManager struct { type npmPackageManager struct {
@@ -89,7 +94,7 @@ func (npm *npmPackageManager) ParseCommand(args []string) (*ParsedCommand, error
} }
func (npm *npmPackageManager) isInstallCommand(cmd string) bool { func (npm *npmPackageManager) isInstallCommand(cmd string) bool {
return cmd == "install" || cmd == "i" || cmd == "add" return slices.Contains(npm.Config.InstallCommands, cmd)
} }
func npmParsePackageInfo(input string) (packageName, version string, err error) { func npmParsePackageInfo(input string) (packageName, version string, err error) {
@@ -130,7 +135,8 @@ func npmParsePackageInfo(input string) (packageName, version string, err error)
func npmCleanVersion(version string) string { func npmCleanVersion(version string) string {
version = strings.TrimPrefix(version, "^") version = strings.TrimPrefix(version, "^")
version = strings.TrimPrefix(version, "~") version = strings.TrimPrefix(version, "~")
if version == "*" {
if version == "*" || version == "" {
return "latest" return "latest"
} }
+89 -68
View File
@@ -22,7 +22,7 @@ func NewDefaultNpmDependencyResolverConfig() NpmDependencyResolverConfig {
return NpmDependencyResolverConfig{ return NpmDependencyResolverConfig{
IncludeDevDependencies: true, IncludeDevDependencies: true,
IncludeTransitiveDependencies: true, IncludeTransitiveDependencies: true,
TransitiveDepth: 100, TransitiveDepth: 5,
FailFast: false, FailFast: false,
} }
} }
@@ -66,6 +66,8 @@ func (r *npmDependencyResolver) ResolveLatestVersion(ctx context.Context,
}, nil }, nil
} }
// TODO: Refactor this into a generic dependency resolver that depends on
// package registry client
func (r *npmDependencyResolver) ResolveDependencies(ctx context.Context, func (r *npmDependencyResolver) ResolveDependencies(ctx context.Context,
packageVersion *packagev1.PackageVersion) ([]*packagev1.PackageVersion, error) { packageVersion *packagev1.PackageVersion) ([]*packagev1.PackageVersion, error) {
pd, err := r.registry.PackageDiscovery() pd, err := r.registry.PackageDiscovery()
@@ -73,76 +75,95 @@ func (r *npmDependencyResolver) ResolveDependencies(ctx context.Context,
return nil, fmt.Errorf("failed to get package discovery: %w", err) return nil, fmt.Errorf("failed to get package discovery: %w", err)
} }
// inputQueue is the queue of package versions to resolve // Track visited packages to avoid cycles
inputQueue := make([]*packagev1.PackageVersion, 0)
// outputQueue is the queue of resolved package versions
outputQueue := make([]*packagev1.PackageVersion, 0)
// Start with the initial package version to resolve
inputQueue = append(inputQueue, packageVersion)
resolutionDepth := 0
visitedPackages := make(map[string]bool) visitedPackages := make(map[string]bool)
for { // Result collection
if len(inputQueue) == 0 { dependencies := make([]*packagev1.PackageVersion, 0)
break
}
if resolutionDepth > r.config.TransitiveDepth { // Start recursive resolution
return nil, fmt.Errorf("exceeded maximum transitive depth of %d", r.config.TransitiveDepth) err = r.resolvePackageDependenciesRecursive(ctx, pd, packageVersion, 0, visitedPackages, &dependencies)
} if err != nil {
return nil, fmt.Errorf("failed to resolve dependencies: %w", err)
packageVersion = inputQueue[0]
inputQueue = inputQueue[1:]
// Skip if we've already visited this package version
// Like npm, we will reuse an existing version instead of considering
// every single version of a dependency. This is a heuristic. We are not
// actually checking for compatibility here like npm does
packageKey := fmt.Sprintf("%s@*", packageVersion.Package.Name)
if visitedPackages[packageKey] {
continue
}
// Mark the package version as visited
visitedPackages[packageKey] = true
dependencyList, err := pd.GetPackageDependencies(packageVersion.Package.Name, packageVersion.Version)
if err != nil {
if r.config.FailFast {
return nil, fmt.Errorf("failed to get package dependencies: %w", err)
}
log.Warnf("failed to get package dependencies: %s", err)
continue
}
dependencies := dependencyList.Dependencies
if r.config.IncludeDevDependencies {
dependencies = append(dependencies, dependencyList.DevDependencies...)
}
resolvedPackageVersionDependencies := make([]*packagev1.PackageVersion, 0)
for _, dependency := range dependencies {
resolvedPackageVersionDependencies = append(resolvedPackageVersionDependencies, &packagev1.PackageVersion{
Package: &packagev1.Package{
Name: dependency.Name,
Ecosystem: packagev1.Ecosystem_ECOSYSTEM_NPM,
},
Version: npmCleanVersion(dependency.VersionSpec),
})
}
outputQueue = append(outputQueue, resolvedPackageVersionDependencies...)
if r.config.IncludeTransitiveDependencies {
inputQueue = append(inputQueue, resolvedPackageVersionDependencies...)
}
resolutionDepth++
} }
return outputQueue, nil return dependencies, nil
}
// resolvePackageDependenciesRecursive resolves dependencies for a package version recursively
func (r *npmDependencyResolver) resolvePackageDependenciesRecursive(
ctx context.Context,
pd packageregistry.PackageDiscovery,
packageVersion *packagev1.PackageVersion,
depth int,
visitedPackages map[string]bool,
result *[]*packagev1.PackageVersion) error {
ff := func(err error) error {
if r.config.FailFast {
return err
}
log.Warnf("error resolving package dependencies: %w", err)
return nil
}
// Check depth limit
if depth > r.config.TransitiveDepth {
return ff(fmt.Errorf("exceeded maximum transitive depth of %d", r.config.TransitiveDepth))
}
// Skip if already visited
packageKey := r.packageKey(packageVersion)
if visitedPackages[packageKey] {
return nil
}
// Mark as visited
visitedPackages[packageKey] = true
log.Debugf("resolving dependencies for %s@%s", packageVersion.Package.Name, packageVersion.Version)
// Get dependencies for the current package
dependencyList, err := pd.GetPackageDependencies(packageVersion.Package.Name, packageVersion.Version)
if err != nil {
return ff(fmt.Errorf("failed to get package dependencies: %w", err))
}
// Collect all dependencies (and optionally dev dependencies)
dependencies := dependencyList.Dependencies
if r.config.IncludeDevDependencies {
dependencies = append(dependencies, dependencyList.DevDependencies...)
}
// Create package version objects for all dependencies
resolvedDependencies := make([]*packagev1.PackageVersion, 0, len(dependencies))
for _, dependency := range dependencies {
resolvedDependencies = append(resolvedDependencies, &packagev1.PackageVersion{
Package: &packagev1.Package{
Ecosystem: packagev1.Ecosystem_ECOSYSTEM_NPM,
Name: dependency.Name,
},
Version: npmCleanVersion(dependency.VersionSpec),
})
}
// Add all resolved dependencies to the result
*result = append(*result, resolvedDependencies...)
// Process transitive dependencies if enabled
if r.config.IncludeTransitiveDependencies {
for _, dependency := range resolvedDependencies {
err := r.resolvePackageDependenciesRecursive(ctx, pd, dependency, depth+1, visitedPackages, result)
if err != nil {
return ff(fmt.Errorf("failed to resolve transitive dependency: %w", err))
}
}
}
return nil
}
func (r *npmDependencyResolver) packageKey(pkg *packagev1.PackageVersion) string {
return fmt.Sprintf("%s@%s", pkg.Package.Name, pkg.Version)
} }
+45 -3
View File
@@ -5,9 +5,51 @@ import (
"testing" "testing"
packagev1 "buf.build/gen/go/safedep/api/protocolbuffers/go/safedep/messages/package/v1" packagev1 "buf.build/gen/go/safedep/api/protocolbuffers/go/safedep/messages/package/v1"
"github.com/safedep/dry/semver"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
func TestNpmDependencyResolver_ResolveLatestVersion(t *testing.T) {
cases := []struct {
name string
pkg *packagev1.Package
assertFn func(t *testing.T, pv *packagev1.PackageVersion, err error)
}{
{
name: "should resolve latest version for a package",
pkg: &packagev1.Package{
Name: "react",
Ecosystem: packagev1.Ecosystem_ECOSYSTEM_NPM,
},
assertFn: func(t *testing.T, pv *packagev1.PackageVersion, err error) {
require.NoError(t, err)
require.True(t, semver.IsAhead("19.0.0", pv.Version))
},
},
{
name: "should return an error if the package is not found",
pkg: &packagev1.Package{
Name: "nonexistent",
Ecosystem: packagev1.Ecosystem_ECOSYSTEM_NPM,
},
assertFn: func(t *testing.T, pv *packagev1.PackageVersion, err error) {
require.Error(t, err)
require.Nil(t, pv)
},
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
resolver, err := NewNpmDependencyResolver(NewDefaultNpmDependencyResolverConfig())
require.NoError(t, err)
pv, err := resolver.ResolveLatestVersion(context.Background(), tc.pkg)
tc.assertFn(t, pv, err)
})
}
}
func TestNpmDependencyResolver_ResolveDependencies(t *testing.T) { func TestNpmDependencyResolver_ResolveDependencies(t *testing.T) {
cases := []struct { cases := []struct {
name string name string
@@ -38,13 +80,13 @@ func TestNpmDependencyResolver_ResolveDependencies(t *testing.T) {
name: "should resolve all dependencies for a package when transitive dependencies are included", name: "should resolve all dependencies for a package when transitive dependencies are included",
pkg: &packagev1.PackageVersion{ pkg: &packagev1.PackageVersion{
Package: &packagev1.Package{ Package: &packagev1.Package{
Name: "react", Name: "express",
Ecosystem: packagev1.Ecosystem_ECOSYSTEM_NPM, Ecosystem: packagev1.Ecosystem_ECOSYSTEM_NPM,
}, },
Version: "18.2.0", Version: "4.18.2",
}, },
includeTransitiveDependencies: true, includeTransitiveDependencies: true,
transitiveDepth: 250, transitiveDepth: 5,
assertFn: func(t *testing.T, dependencies []*packagev1.PackageVersion, err error) { assertFn: func(t *testing.T, dependencies []*packagev1.PackageVersion, err error) {
require.NoError(t, err) require.NoError(t, err)
require.Equal(t, 1, len(dependencies)) require.Equal(t, 1, len(dependencies))