feat: support PyPi package extras

This commit is contained in:
Sahilb315
2025-06-02 20:56:01 +05:30
parent 30c914d07e
commit 6e2dc34f74
7 changed files with 55 additions and 11 deletions
+7 -1
View File
@@ -2,6 +2,7 @@ package npm
import ( import (
"context" "context"
"fmt"
"github.com/safedep/dry/log" "github.com/safedep/dry/log"
"github.com/safedep/pmg/config" "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) 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 := packagemanager.NewDefaultNpmDependencyResolverConfig()
packageResolverConfig.IncludeTransitiveDependencies = config.Transitive packageResolverConfig.IncludeTransitiveDependencies = config.Transitive
packageResolverConfig.TransitiveDepth = config.TransitiveDepth 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) 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)
} }
+7 -1
View File
@@ -3,6 +3,7 @@ package npm
import ( import (
"context" "context"
_ "embed" _ "embed"
"fmt"
"github.com/safedep/dry/log" "github.com/safedep/dry/log"
"github.com/safedep/pmg/config" "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) 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 := packagemanager.NewDefaultNpmDependencyResolverConfig()
packageResolverConfig.IncludeTransitiveDependencies = config.Transitive packageResolverConfig.IncludeTransitiveDependencies = config.Transitive
packageResolverConfig.TransitiveDepth = config.TransitiveDepth packageResolverConfig.TransitiveDepth = config.TransitiveDepth
@@ -45,5 +51,5 @@ func executePnpmFlow(ctx context.Context, args []string) error {
packageResolver, err := packagemanager.NewNpmDependencyResolver(packageResolverConfig) packageResolver, err := packagemanager.NewNpmDependencyResolver(packageResolverConfig)
return flows.Common(packageManager, packageResolver, config).Run(ctx, args) return flows.Common(packageManager, packageResolver, config).Run(ctx, args, parsedCommand)
} }
+9 -1
View File
@@ -37,15 +37,23 @@ func executePipFlow(ctx context.Context, args []string) error {
if err != nil { if err != nil {
ui.Fatalf("Failed to get config: %s", err) 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 := packagemanager.NewDefaultPypiDependencyResolverConfig()
packageResolverConfig.IncludeTransitiveDependencies = config.Transitive packageResolverConfig.IncludeTransitiveDependencies = config.Transitive
packageResolverConfig.TransitiveDepth = config.TransitiveDepth packageResolverConfig.TransitiveDepth = config.TransitiveDepth
packageResolverConfig.IncludeDevDependencies = config.IncludeDevDependencies packageResolverConfig.IncludeDevDependencies = config.IncludeDevDependencies
packageResolverConfig.PackageInstallTargets = parsedCommand.InstallTargets
packageResolver, err := packagemanager.NewPypiDependencyResolver(packageResolverConfig) packageResolver, err := packagemanager.NewPypiDependencyResolver(packageResolverConfig)
if err != nil { if err != nil {
ui.Fatalf("Failed to create dependency resolver: %s", err) 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)
} }
+1 -1
View File
@@ -69,7 +69,7 @@ func NewPackageManagerGuard(config PackageManagerGuardConfig,
}, nil }, 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) log.Debugf("Running package manager guard with args: %v", args)
parsedCommand, err := g.packageManager.ParseCommand(args) parsedCommand, err := g.packageManager.ParseCommand(args)
+2 -2
View File
@@ -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 var analyzers []analyzer.PackageVersionAnalyzer
if f.config.Paranoid { 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) ui.Fatalf("Failed to create package manager guard: %s", err)
} }
err = proxy.Run(ctx, args) err = proxy.Run(ctx, args, parsedCmd)
if err != nil { if err != nil {
ui.Fatalf("pmg: failed to execute command: %s", err) ui.Fatalf("pmg: failed to execute command: %s", err)
} }
+19 -5
View File
@@ -12,6 +12,7 @@ import (
"github.com/Masterminds/semver" "github.com/Masterminds/semver"
"github.com/safedep/dry/log" "github.com/safedep/dry/log"
"github.com/safedep/dry/packageregistry" "github.com/safedep/dry/packageregistry"
"github.com/safedep/pmg/pkg/utils"
) )
type PyPiDependencyResolverConfig struct { type PyPiDependencyResolverConfig struct {
@@ -24,6 +25,8 @@ type PyPiDependencyResolverConfig struct {
// MaxConcurrency limits the number of concurrent goroutines used for dependency resolution // MaxConcurrency limits the number of concurrent goroutines used for dependency resolution
MaxConcurrency int MaxConcurrency int
PackageInstallTargets []*PackageInstallTarget
} }
func NewDefaultPypiDependencyResolverConfig() PyPiDependencyResolverConfig { func NewDefaultPypiDependencyResolverConfig() PyPiDependencyResolverConfig {
@@ -33,6 +36,7 @@ func NewDefaultPypiDependencyResolverConfig() PyPiDependencyResolverConfig {
TransitiveDepth: 5, TransitiveDepth: 5,
FailFast: false, FailFast: false,
MaxConcurrency: 10, 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) { 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)
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 ""
@@ -67,7 +70,7 @@ func (p *pypiDependencyResolver) ResolveDependencies(ctx context.Context, pkg *p
} }
pypiDependencyResolverFn := func(packageName, version string) (*packageregistry.PackageDependencyList, error) { pypiDependencyResolverFn := func(packageName, version string) (*packageregistry.PackageDependencyList, error) {
resolvedDependencies, err := getPypiPackageDependencies(packageName, version) resolvedDependencies, err := getPypiPackageDependencies(packageName, version, p.config.PackageInstallTargets)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -150,7 +153,7 @@ type pypiPackageInfo struct {
RequiresDist []string `json:"requires_dist"` 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) url := fmt.Sprintf("https://pypi.org/pypi/%s/%s/json", packageName, version)
res, err := http.Get(url) res, err := http.Get(url)
@@ -173,13 +176,24 @@ func getPypiPackageDependencies(packageName, version string) ([]PyPIDependencySp
return nil, ErrFailedToParsePackage 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)) pkgDeps := make([]PyPIDependencySpec, 0, len(pypipkg.Info.RequiresDist))
for _, dep := range pypipkg.Info.RequiresDist { for _, dep := range pypipkg.Info.RequiresDist {
name, version, extra := pypiParseDependency(dep) name, version, extra := pypiParseDependency(dep)
// Skip dependencies with extras/conditions to avoid resolution issues // Include dependencies if they either:
if extra == "" { // 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{ pkgDeps = append(pkgDeps, PyPIDependencySpec{
PackageNameExtra: name, PackageNameExtra: name,
VersionSpec: version, VersionSpec: version,
+10
View File
@@ -0,0 +1,10 @@
package utils
func Contains(slice []string, item string) bool {
for _, s := range slice {
if s == item {
return true
}
}
return false
}