From 6e2dc34f74b7e426b02ed3d53fd12c2bbcdbb7bb Mon Sep 17 00:00:00 2001 From: Sahilb315 Date: Mon, 2 Jun 2025 02:42:55 +0530 Subject: [PATCH] feat: support PyPi package extras --- cmd/npm/npm.go | 8 +++++++- cmd/npm/pnpm.go | 8 +++++++- cmd/pypi/pip.go | 10 +++++++++- guard/guard.go | 2 +- internal/flows/common.go | 4 ++-- packagemanager/pypi_resolver.go | 24 +++++++++++++++++++----- pkg/utils/utils.go | 10 ++++++++++ 7 files changed, 55 insertions(+), 11 deletions(-) create mode 100644 pkg/utils/utils.go diff --git a/cmd/npm/npm.go b/cmd/npm/npm.go index 5e8f86f..d4826ed 100644 --- a/cmd/npm/npm.go +++ b/cmd/npm/npm.go @@ -2,6 +2,7 @@ package npm import ( "context" + "fmt" "github.com/safedep/dry/log" "github.com/safedep/pmg/config" @@ -37,6 +38,11 @@ func executeNpmFlow(ctx context.Context, args []string) error { ui.Fatalf("Failed to get config: %s", err) } + parsedCommand, err := packageManager.ParseCommand(args) + if err != nil { + return fmt.Errorf("failed to parse command: %w", err) + } + packageResolverConfig := packagemanager.NewDefaultNpmDependencyResolverConfig() packageResolverConfig.IncludeTransitiveDependencies = config.Transitive packageResolverConfig.TransitiveDepth = config.TransitiveDepth @@ -47,5 +53,5 @@ func executeNpmFlow(ctx context.Context, args []string) error { ui.Fatalf("Failed to create dependency resolver: %s", err) } - return flows.Common(packageManager, packageResolver, config).Run(ctx, args) + return flows.Common(packageManager, packageResolver, config).Run(ctx, args, parsedCommand) } diff --git a/cmd/npm/pnpm.go b/cmd/npm/pnpm.go index a010244..ae19fa6 100644 --- a/cmd/npm/pnpm.go +++ b/cmd/npm/pnpm.go @@ -3,6 +3,7 @@ package npm import ( "context" _ "embed" + "fmt" "github.com/safedep/dry/log" "github.com/safedep/pmg/config" @@ -38,6 +39,11 @@ func executePnpmFlow(ctx context.Context, args []string) error { ui.Fatalf("Failed to get config: %s", err) } + parsedCommand, err := packageManager.ParseCommand(args) + if err != nil { + return fmt.Errorf("failed to parse command: %w", err) + } + packageResolverConfig := packagemanager.NewDefaultNpmDependencyResolverConfig() packageResolverConfig.IncludeTransitiveDependencies = config.Transitive packageResolverConfig.TransitiveDepth = config.TransitiveDepth @@ -45,5 +51,5 @@ func executePnpmFlow(ctx context.Context, args []string) error { packageResolver, err := packagemanager.NewNpmDependencyResolver(packageResolverConfig) - return flows.Common(packageManager, packageResolver, config).Run(ctx, args) + return flows.Common(packageManager, packageResolver, config).Run(ctx, args, parsedCommand) } diff --git a/cmd/pypi/pip.go b/cmd/pypi/pip.go index 4cc9a65..b5fc36a 100644 --- a/cmd/pypi/pip.go +++ b/cmd/pypi/pip.go @@ -37,15 +37,23 @@ func executePipFlow(ctx context.Context, args []string) error { if err != nil { ui.Fatalf("Failed to get config: %s", err) } + + parsedCommand, err := packageManager.ParseCommand(args) + if err != nil { + return fmt.Errorf("failed to parse command: %w", err) + } + + // Parse the args right here packageResolverConfig := packagemanager.NewDefaultPypiDependencyResolverConfig() packageResolverConfig.IncludeTransitiveDependencies = config.Transitive packageResolverConfig.TransitiveDepth = config.TransitiveDepth packageResolverConfig.IncludeDevDependencies = config.IncludeDevDependencies + packageResolverConfig.PackageInstallTargets = parsedCommand.InstallTargets packageResolver, err := packagemanager.NewPypiDependencyResolver(packageResolverConfig) if err != nil { ui.Fatalf("Failed to create dependency resolver: %s", err) } - return flows.Common(packageManager, packageResolver, config).Run(ctx, args) + return flows.Common(packageManager, packageResolver, config).Run(ctx, args, parsedCommand) } diff --git a/guard/guard.go b/guard/guard.go index 8795951..03a666c 100644 --- a/guard/guard.go +++ b/guard/guard.go @@ -69,7 +69,7 @@ func NewPackageManagerGuard(config PackageManagerGuardConfig, }, nil } -func (g *packageManagerGuard) Run(ctx context.Context, args []string) error { +func (g *packageManagerGuard) Run(ctx context.Context, args []string, parsedCommand *packagemanager.ParsedCommand) error { log.Debugf("Running package manager guard with args: %v", args) parsedCommand, err := g.packageManager.ParseCommand(args) diff --git a/internal/flows/common.go b/internal/flows/common.go index dbc82ca..a49ac90 100644 --- a/internal/flows/common.go +++ b/internal/flows/common.go @@ -27,7 +27,7 @@ func Common(pm packagemanager.PackageManager, pkgResolver packagemanager.Package } } -func (f *commonFlow) Run(ctx context.Context, args []string) error { +func (f *commonFlow) Run(ctx context.Context, args []string, parsedCmd *packagemanager.ParsedCommand) error { var analyzers []analyzer.PackageVersionAnalyzer if f.config.Paranoid { @@ -61,7 +61,7 @@ func (f *commonFlow) Run(ctx context.Context, args []string) error { ui.Fatalf("Failed to create package manager guard: %s", err) } - err = proxy.Run(ctx, args) + err = proxy.Run(ctx, args, parsedCmd) if err != nil { ui.Fatalf("pmg: failed to execute command: %s", err) } diff --git a/packagemanager/pypi_resolver.go b/packagemanager/pypi_resolver.go index 37b0d24..b237f0a 100644 --- a/packagemanager/pypi_resolver.go +++ b/packagemanager/pypi_resolver.go @@ -12,6 +12,7 @@ import ( "github.com/Masterminds/semver" "github.com/safedep/dry/log" "github.com/safedep/dry/packageregistry" + "github.com/safedep/pmg/pkg/utils" ) type PyPiDependencyResolverConfig struct { @@ -24,6 +25,8 @@ type PyPiDependencyResolverConfig struct { // MaxConcurrency limits the number of concurrent goroutines used for dependency resolution MaxConcurrency int + + PackageInstallTargets []*PackageInstallTarget } func NewDefaultPypiDependencyResolverConfig() PyPiDependencyResolverConfig { @@ -33,6 +36,7 @@ func NewDefaultPypiDependencyResolverConfig() PyPiDependencyResolverConfig { TransitiveDepth: 5, FailFast: false, MaxConcurrency: 10, + PackageInstallTargets: []*PackageInstallTarget{}, } } @@ -58,7 +62,6 @@ func NewPypiDependencyResolver(config PyPiDependencyResolverConfig) (*pypiDepend func (p *pypiDependencyResolver) ResolveDependencies(ctx context.Context, pkg *packagev1.PackageVersion) ([]*packagev1.PackageVersion, error) { pypiVersionSpecResolverFn := func(packageName, version string) string { ver, err := pipGetMatchingVersion(packageName, version) - // fmt.Printf("Resolved %s for %s to %s\n", version, packageName, ver) if err != nil { log.Debugf("error getting matching version for %s@%s", packageName, version) return "" @@ -67,7 +70,7 @@ func (p *pypiDependencyResolver) ResolveDependencies(ctx context.Context, pkg *p } pypiDependencyResolverFn := func(packageName, version string) (*packageregistry.PackageDependencyList, error) { - resolvedDependencies, err := getPypiPackageDependencies(packageName, version) + resolvedDependencies, err := getPypiPackageDependencies(packageName, version, p.config.PackageInstallTargets) if err != nil { return nil, err } @@ -150,7 +153,7 @@ type pypiPackageInfo struct { RequiresDist []string `json:"requires_dist"` } -func getPypiPackageDependencies(packageName, version string) ([]PyPIDependencySpec, error) { +func getPypiPackageDependencies(packageName, version string, packageTargets []*PackageInstallTarget) ([]PyPIDependencySpec, error) { url := fmt.Sprintf("https://pypi.org/pypi/%s/%s/json", packageName, version) res, err := http.Get(url) @@ -173,13 +176,24 @@ func getPypiPackageDependencies(packageName, version string) ([]PyPIDependencySp return nil, ErrFailedToParsePackage } + // Find if this package has any specified extras in the install targets + var requestedExtras []string + for _, target := range packageTargets { + if target.PackageVersion.Package.Name == packageName { + requestedExtras = target.Extras + break + } + } + pkgDeps := make([]PyPIDependencySpec, 0, len(pypipkg.Info.RequiresDist)) for _, dep := range pypipkg.Info.RequiresDist { name, version, extra := pypiParseDependency(dep) - // Skip dependencies with extras/conditions to avoid resolution issues - if extra == "" { + // Include dependencies if they either: + // 1. Have no extras (base dependencies) + // 2. Have an extra that matches one of our requested extras + if extra == "" || (len(requestedExtras) > 0 && utils.Contains(requestedExtras, extra)) { pkgDeps = append(pkgDeps, PyPIDependencySpec{ PackageNameExtra: name, VersionSpec: version, diff --git a/pkg/utils/utils.go b/pkg/utils/utils.go new file mode 100644 index 0000000..386fe72 --- /dev/null +++ b/pkg/utils/utils.go @@ -0,0 +1,10 @@ +package utils + +func Contains(slice []string, item string) bool { + for _, s := range slice { + if s == item { + return true + } + } + return false +}