fix: improve dependency resolution and package deduplication

This commit is contained in:
Sahilb315
2025-05-29 00:39:10 +05:30
parent dad32fd6d1
commit 04e7b97785
4 changed files with 52 additions and 16 deletions
+1
View File
@@ -120,6 +120,7 @@ func (g *packageManagerGuard) Run(ctx context.Context, args []string) error {
} }
} }
fmt.Println("Packages to analyse: ", packagesToAnalyze)
log.Debugf("Checking %d packages for malware", len(packagesToAnalyze)) log.Debugf("Checking %d packages for malware", len(packagesToAnalyze))
g.setStatus(fmt.Sprintf("Analyzing %d packages for malware", len(packagesToAnalyze))) g.setStatus(fmt.Sprintf("Analyzing %d packages for malware", len(packagesToAnalyze)))
+31 -13
View File
@@ -3,7 +3,6 @@ package packagemanager
import ( import (
"context" "context"
"fmt" "fmt"
"slices"
"sync" "sync"
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"
@@ -17,6 +16,8 @@ type versionSpecResolverFn func(packageName, version string) string
type dependencyResolverFn func(packageName, version string) (*packageregistry.PackageDependencyList, error) type dependencyResolverFn func(packageName, version string) (*packageregistry.PackageDependencyList, error)
type packageIndentifierFn func(pkg *packagev1.PackageVersion) string
type dependencyResolverConfig struct { type dependencyResolverConfig struct {
IncludeDevDependencies bool IncludeDevDependencies bool
IncludeTransitiveDependencies bool IncludeTransitiveDependencies bool
@@ -31,10 +32,12 @@ type dependencyResolver struct {
mutex sync.Mutex mutex sync.Mutex
versionSpecResolver versionSpecResolverFn versionSpecResolver versionSpecResolverFn
packageDependencyResolver dependencyResolverFn packageDependencyResolver dependencyResolverFn
packageIdentifierFn packageIndentifierFn
resultSet map[string]bool
} }
func newDependencyResolver(client packageregistry.Client, config dependencyResolverConfig, func newDependencyResolver(client packageregistry.Client, config dependencyResolverConfig,
versionSpecResolver versionSpecResolverFn, packageDependencyResolver dependencyResolverFn) *dependencyResolver { versionSpecResolver versionSpecResolverFn, packageDependencyResolver dependencyResolverFn, packageKeyFn packageIndentifierFn) *dependencyResolver {
if config.MaxConcurrency <= 0 { if config.MaxConcurrency <= 0 {
config.MaxConcurrency = 10 config.MaxConcurrency = 10
} }
@@ -51,6 +54,8 @@ func newDependencyResolver(client packageregistry.Client, config dependencyResol
config: config, config: config,
versionSpecResolver: versionSpecResolver, versionSpecResolver: versionSpecResolver,
packageDependencyResolver: packageDependencyResolver, packageDependencyResolver: packageDependencyResolver,
packageIdentifierFn: packageKeyFn,
resultSet: make(map[string]bool),
} }
} }
@@ -67,6 +72,7 @@ func (r *dependencyResolver) resolveDependencies(ctx context.Context,
// Result collection // Result collection
dependencies := make([]*packagev1.PackageVersion, 0) dependencies := make([]*packagev1.PackageVersion, 0)
r.resultSet = make(map[string]bool) // Reset
// Start concurrent resolution // Start concurrent resolution
err = r.resolvePackageDependenciesConcurrent(ctx, pd, packageVersion, 0, visitedPackages, &dependencies) err = r.resolvePackageDependenciesConcurrent(ctx, pd, packageVersion, 0, visitedPackages, &dependencies)
if err != nil { if err != nil {
@@ -106,23 +112,31 @@ func (r *dependencyResolver) resolvePackageDependenciesConcurrent(
return ff(fmt.Errorf("exceeded maximum transitive depth of %d", r.config.TransitiveDepth)) return ff(fmt.Errorf("exceeded maximum transitive depth of %d", r.config.TransitiveDepth))
} }
var packageKey string
var packageKeyFn packageIndentifierFn
// Skip if already visited // Skip if already visited
packageKey := r.packageKey(packageVersion) if r.packageIdentifierFn != nil {
packageKeyFn = r.packageIdentifierFn
} else {
packageKeyFn = createPackageKey
}
packageKey = packageKeyFn(packageVersion)
shouldProcess := false
alreadyVisited := false
r.synchronize(func() { r.synchronize(func() {
alreadyVisited = visitedPackages[packageKey] if !visitedPackages[packageKey] {
visitedPackages[packageKey] = true
shouldProcess = true
}
}) })
if alreadyVisited { // If another goroutine is already processing this package, skip
if !shouldProcess {
return nil return nil
} }
// Mark the current package as visited
r.synchronize(func() {
visitedPackages[packageKey] = true
})
log.Debugf("resolving dependencies for %s@%s", packageVersion.Package.Name, packageVersion.Version) log.Debugf("resolving dependencies for %s@%s", packageVersion.Package.Name, packageVersion.Version)
// Get dependencies for the current package // Get dependencies for the current package
@@ -159,7 +173,10 @@ func (r *dependencyResolver) resolvePackageDependenciesConcurrent(
// Add resolved dependencies to the result // Add resolved dependencies to the result
r.synchronize(func() { r.synchronize(func() {
for _, dependency := range resolvedDependencies { for _, dependency := range resolvedDependencies {
if !slices.Contains(*result, dependency) { dependencyKey := packageKeyFn(dependency)
if !r.resultSet[dependencyKey] {
r.resultSet[dependencyKey] = true
*result = append(*result, dependency) *result = append(*result, dependency)
} }
} }
@@ -197,10 +214,11 @@ func (r *dependencyResolver) resolvePackageDependenciesConcurrent(
return ff(fmt.Errorf("failed to resolve transitive dependency: %w", err)) return ff(fmt.Errorf("failed to resolve transitive dependency: %w", err))
} }
} }
return nil return nil
} }
func (r *dependencyResolver) packageKey(pkg *packagev1.PackageVersion) string { func createPackageKey(pkg *packagev1.PackageVersion) string {
return fmt.Sprintf("%s@%s", pkg.Package.Name, pkg.Version) return fmt.Sprintf("%s@%s", pkg.Package.Name, pkg.Version)
} }
+1 -1
View File
@@ -80,7 +80,7 @@ func (r *npmDependencyResolver) ResolveDependencies(ctx context.Context,
MaxConcurrency: r.config.MaxConcurrency, MaxConcurrency: r.config.MaxConcurrency,
}, func(packageName, version string) string { }, func(packageName, version string) string {
return npmCleanVersion(version) return npmCleanVersion(version)
}, nil) }, nil, nil)
return resolver.resolveDependencies(ctx, packageVersion) return resolver.resolveDependencies(ctx, packageVersion)
} }
+19 -2
View File
@@ -58,7 +58,7 @@ func NewPypiDependencyResolver(config PyPiDependencyResolverConfig) (*pypiDepend
func (p *pypiDependencyResolver) ResolveDependencies(ctx context.Context, pkg *packagev1.PackageVersion) ([]*packagev1.PackageVersion, error) { func (p *pypiDependencyResolver) ResolveDependencies(ctx context.Context, pkg *packagev1.PackageVersion) ([]*packagev1.PackageVersion, error) {
pypiVersionSpecResolverFn := func(packageName, version string) string { pypiVersionSpecResolverFn := func(packageName, version string) string {
ver, err := pipGetMatchingVersion(packageName, version) ver, err := pipGetMatchingVersion(packageName, version)
fmt.Printf("Resolved %s for %s to %s\n", version, packageName, ver) // fmt.Printf("Resolved %s for %s to %s\n", version, packageName, ver)
if err != nil { if err != nil {
log.Debugf("error getting matching version for %s@%s", packageName, version) log.Debugf("error getting matching version for %s@%s", packageName, version)
return "" return ""
@@ -84,13 +84,19 @@ func (p *pypiDependencyResolver) ResolveDependencies(ctx context.Context, pkg *p
}, nil }, nil
} }
// Python treats package names with '-' and '_' as equivalent (e.g., 'my-package' and 'my_package' refer to the same package)
packageKeyFn := func(pkg *packagev1.PackageVersion) string {
normalizedName := normalizePackageName(pkg.Package.Name)
return fmt.Sprintf("%s@%s", normalizedName, pkg.Version)
}
resolver := newDependencyResolver(p.registry, dependencyResolverConfig{ resolver := newDependencyResolver(p.registry, dependencyResolverConfig{
IncludeDevDependencies: p.config.IncludeDevDependencies, IncludeDevDependencies: p.config.IncludeDevDependencies,
IncludeTransitiveDependencies: p.config.IncludeTransitiveDependencies, IncludeTransitiveDependencies: p.config.IncludeTransitiveDependencies,
TransitiveDepth: p.config.TransitiveDepth, TransitiveDepth: p.config.TransitiveDepth,
FailFast: p.config.FailFast, FailFast: p.config.FailFast,
MaxConcurrency: p.config.MaxConcurrency, MaxConcurrency: p.config.MaxConcurrency,
}, pypiVersionSpecResolverFn, pypiDependencyResolverFn) }, pypiVersionSpecResolverFn, pypiDependencyResolverFn, packageKeyFn)
return resolver.resolveDependencies(ctx, pkg) return resolver.resolveDependencies(ctx, pkg)
} }
@@ -314,3 +320,14 @@ func findBestMatchingVersion(releases []packageregistry.PackageVersionInfo, cons
} }
return bestMatch, nil return bestMatch, nil
} }
func normalizePackageName(name string) string {
// Convert to lowercase
name = strings.ToLower(name)
// Replace any sequence of [-_.] with a single hyphen
re := regexp.MustCompile(`[-_.]+`)
name = re.ReplaceAllString(name, "-")
return name
}