From 0883140732aba394aacb09ae992613f611af406b Mon Sep 17 00:00:00 2001 From: abhisek Date: Wed, 14 May 2025 09:20:44 +0530 Subject: [PATCH] fix: Npm dependency resolver --- packagemanager/npm.go | 14 ++- packagemanager/npm_resolver.go | 157 ++++++++++++++++------------ packagemanager/npm_resolver_test.go | 48 ++++++++- 3 files changed, 144 insertions(+), 75 deletions(-) diff --git a/packagemanager/npm.go b/packagemanager/npm.go index aafe7d9..92da955 100644 --- a/packagemanager/npm.go +++ b/packagemanager/npm.go @@ -2,15 +2,20 @@ package packagemanager import ( "fmt" + "slices" "strings" packagev1 "buf.build/gen/go/safedep/api/protocolbuffers/go/safedep/messages/package/v1" ) -type NpmPackageManagerConfig struct{} +type NpmPackageManagerConfig struct { + InstallCommands []string +} func DefaultNpmPackageManagerConfig() NpmPackageManagerConfig { - return NpmPackageManagerConfig{} + return NpmPackageManagerConfig{ + InstallCommands: []string{"install", "i", "add"}, + } } type npmPackageManager struct { @@ -89,7 +94,7 @@ func (npm *npmPackageManager) ParseCommand(args []string) (*ParsedCommand, error } 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) { @@ -130,7 +135,8 @@ func npmParsePackageInfo(input string) (packageName, version string, err error) func npmCleanVersion(version string) string { version = strings.TrimPrefix(version, "^") version = strings.TrimPrefix(version, "~") - if version == "*" { + + if version == "*" || version == "" { return "latest" } diff --git a/packagemanager/npm_resolver.go b/packagemanager/npm_resolver.go index e3173d0..e7eb6a2 100644 --- a/packagemanager/npm_resolver.go +++ b/packagemanager/npm_resolver.go @@ -22,7 +22,7 @@ func NewDefaultNpmDependencyResolverConfig() NpmDependencyResolverConfig { return NpmDependencyResolverConfig{ IncludeDevDependencies: true, IncludeTransitiveDependencies: true, - TransitiveDepth: 100, + TransitiveDepth: 5, FailFast: false, } } @@ -66,6 +66,8 @@ func (r *npmDependencyResolver) ResolveLatestVersion(ctx context.Context, }, nil } +// TODO: Refactor this into a generic dependency resolver that depends on +// package registry client func (r *npmDependencyResolver) ResolveDependencies(ctx context.Context, packageVersion *packagev1.PackageVersion) ([]*packagev1.PackageVersion, error) { 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) } - // inputQueue is the queue of package versions to resolve - 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 + // Track visited packages to avoid cycles visitedPackages := make(map[string]bool) - for { - if len(inputQueue) == 0 { - break - } + // Result collection + dependencies := make([]*packagev1.PackageVersion, 0) - if resolutionDepth > r.config.TransitiveDepth { - return nil, fmt.Errorf("exceeded maximum transitive depth of %d", r.config.TransitiveDepth) - } - - 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++ + // Start recursive resolution + err = r.resolvePackageDependenciesRecursive(ctx, pd, packageVersion, 0, visitedPackages, &dependencies) + if err != nil { + return nil, fmt.Errorf("failed to resolve dependencies: %w", err) } - 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) } diff --git a/packagemanager/npm_resolver_test.go b/packagemanager/npm_resolver_test.go index b4689c7..d2b25bf 100644 --- a/packagemanager/npm_resolver_test.go +++ b/packagemanager/npm_resolver_test.go @@ -5,9 +5,51 @@ import ( "testing" packagev1 "buf.build/gen/go/safedep/api/protocolbuffers/go/safedep/messages/package/v1" + "github.com/safedep/dry/semver" "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) { cases := []struct { 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", pkg: &packagev1.PackageVersion{ Package: &packagev1.Package{ - Name: "react", + Name: "express", Ecosystem: packagev1.Ecosystem_ECOSYSTEM_NPM, }, - Version: "18.2.0", + Version: "4.18.2", }, includeTransitiveDependencies: true, - transitiveDepth: 250, + transitiveDepth: 5, assertFn: func(t *testing.T, dependencies []*packagev1.PackageVersion, err error) { require.NoError(t, err) require.Equal(t, 1, len(dependencies))