From 53783c66046a7c253e632944c2f60b2ca35d0e79 Mon Sep 17 00:00:00 2001 From: Sahil Bansal Date: Mon, 9 Jun 2025 17:37:08 +0530 Subject: [PATCH] feat: Add support for pip Package Manager (#33) * feat/init-pip-cmd * feat: Add PyPi resolver * test: Add tests for pypi and pypi_resolver * refactor: unify package dependency resolution and improve PyPI version handling using registry adapter * chore: remove unused file Signed-off-by: Sahil Bansal * feat: add Python dependency parsing with extras support * refactor(deps): Improve PyPI dependency resolution and add custom resolver support * chore: remove extra file * fix: improve dependency resolution and package deduplication * chore: remove extra print statement * chore: typo fix * feat: support PyPi package extras * test: add test for pypi dependency parse function * fix: remove overwritten of parsedCmd * chore: remove extra print statement * chore: typo fix * feat: support PyPi package extras * test: add test for pypi dependency parse function * fix: remove overwritten of parsedCmd * refactor: enhance code readability & remove extra code * chore: remove extra code * Update cmd/npm/npm.go Co-authored-by: Omkar Phansopkar Signed-off-by: Sahil Bansal * Update cmd/npm/pnpm.go Co-authored-by: Omkar Phansopkar Signed-off-by: Sahil Bansal * Update cmd/pypi/pip.go Co-authored-by: Omkar Phansopkar Signed-off-by: Sahil Bansal --------- Signed-off-by: Sahil Bansal Co-authored-by: Omkar Phansopkar --- cmd/npm/npm.go | 25 +- cmd/npm/pnpm.go | 21 +- cmd/pypi/pip.go | 60 +++++ go.mod | 1 + go.sum | 2 + guard/guard.go | 7 +- internal/flows/common.go | 35 +-- main.go | 2 + packagemanager/dependency_resolver.go | 76 ++++-- packagemanager/errors.go | 15 ++ packagemanager/npm.go | 2 + packagemanager/npm_resolver.go | 7 +- packagemanager/packagemanager.go | 5 + packagemanager/pypi.go | 204 +++++++++++++++ packagemanager/pypi_resolver.go | 346 ++++++++++++++++++++++++++ packagemanager/pypi_resolver_test.go | 202 +++++++++++++++ packagemanager/pypi_test.go | 158 ++++++++++++ 17 files changed, 1111 insertions(+), 57 deletions(-) create mode 100644 cmd/pypi/pip.go create mode 100644 packagemanager/errors.go create mode 100644 packagemanager/pypi.go create mode 100644 packagemanager/pypi_resolver.go create mode 100644 packagemanager/pypi_resolver_test.go create mode 100644 packagemanager/pypi_test.go diff --git a/cmd/npm/npm.go b/cmd/npm/npm.go index 33ccaa2..ab7c6ad 100644 --- a/cmd/npm/npm.go +++ b/cmd/npm/npm.go @@ -2,9 +2,10 @@ package npm import ( "context" - _ "embed" + "fmt" "github.com/safedep/dry/log" + "github.com/safedep/pmg/config" "github.com/safedep/pmg/internal/flows" "github.com/safedep/pmg/internal/ui" "github.com/safedep/pmg/packagemanager" @@ -33,5 +34,25 @@ func executeNpmFlow(ctx context.Context, args []string) error { ui.Fatalf("Failed to create npm package manager proxy: %s", err) } - return flows.Common(packageManager).Run(ctx, args) + config, err := config.FromContext(ctx) + 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) + } + + packageResolverConfig := packagemanager.NewDefaultNpmDependencyResolverConfig() + packageResolverConfig.IncludeTransitiveDependencies = config.Transitive + packageResolverConfig.TransitiveDepth = config.TransitiveDepth + packageResolverConfig.IncludeDevDependencies = config.IncludeDevDependencies + + packageResolver, err := packagemanager.NewNpmDependencyResolver(packageResolverConfig) + if err != nil { + ui.Fatalf("Failed to create dependency resolver: %s", err) + } + + return flows.Common(packageManager, packageResolver, config).Run(ctx, args, parsedCommand) } diff --git a/cmd/npm/pnpm.go b/cmd/npm/pnpm.go index c161c0d..ab017a9 100644 --- a/cmd/npm/pnpm.go +++ b/cmd/npm/pnpm.go @@ -3,8 +3,10 @@ package npm import ( "context" _ "embed" + "fmt" "github.com/safedep/dry/log" + "github.com/safedep/pmg/config" "github.com/safedep/pmg/internal/flows" "github.com/safedep/pmg/internal/ui" "github.com/safedep/pmg/packagemanager" @@ -33,5 +35,22 @@ func executePnpmFlow(ctx context.Context, args []string) error { ui.Fatalf("Failed to create pnpm package manager proxy: %s", err) } - return flows.Common(packageManager).Run(ctx, args) + config, err := config.FromContext(ctx) + 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) + } + + packageResolverConfig := packagemanager.NewDefaultNpmDependencyResolverConfig() + packageResolverConfig.IncludeTransitiveDependencies = config.Transitive + packageResolverConfig.TransitiveDepth = config.TransitiveDepth + packageResolverConfig.IncludeDevDependencies = config.IncludeDevDependencies + + packageResolver, err := packagemanager.NewNpmDependencyResolver(packageResolverConfig) + + return flows.Common(packageManager, packageResolver, config).Run(ctx, args, parsedCommand) } diff --git a/cmd/pypi/pip.go b/cmd/pypi/pip.go new file mode 100644 index 0000000..e7e3c57 --- /dev/null +++ b/cmd/pypi/pip.go @@ -0,0 +1,60 @@ +package pypi + +import ( + "context" + "fmt" + + "github.com/safedep/dry/log" + "github.com/safedep/pmg/config" + "github.com/safedep/pmg/internal/flows" + "github.com/safedep/pmg/internal/ui" + "github.com/safedep/pmg/packagemanager" + "github.com/spf13/cobra" +) + +func NewPipCommand() *cobra.Command { + return &cobra.Command{ + Use: "pip [action] [package]", + Short: "Guard pip package manager", + DisableFlagParsing: true, + RunE: func(cmd *cobra.Command, args []string) error { + err := executePipFlow(cmd.Context(), args) + if err != nil { + log.Errorf("Failed to execute pip flow: %s", err) + } + + return nil + }, + } +} + +func executePipFlow(ctx context.Context, args []string) error { + packageManager, err := packagemanager.NewPipPackageManager(packagemanager.DefaultPipPackageManagerConfig()) + if err != nil { + return fmt.Errorf("failed to create pip package manager: %w", err) + } + + config, err := config.FromContext(ctx) + 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, parsedCommand) +} diff --git a/go.mod b/go.mod index 262f534..185ad76 100644 --- a/go.mod +++ b/go.mod @@ -7,6 +7,7 @@ tool github.com/golangci/golangci-lint/cmd/golangci-lint require ( buf.build/gen/go/safedep/api/grpc/go v1.5.1-20250418165058-162f6b0cc319.2 buf.build/gen/go/safedep/api/protocolbuffers/go v1.36.6-20250418165058-162f6b0cc319.1 + github.com/Masterminds/semver v1.5.0 github.com/fatih/color v1.18.0 github.com/jedib0t/go-pretty/v6 v6.6.7 github.com/safedep/dry v0.0.0-20250514080944-bb77f30c7175 diff --git a/go.sum b/go.sum index 0f30795..8486a82 100644 --- a/go.sum +++ b/go.sum @@ -28,6 +28,8 @@ github.com/Djarvur/go-err113 v0.0.0-20210108212216-aea10b59be24 h1:sHglBQTwgx+rW github.com/Djarvur/go-err113 v0.0.0-20210108212216-aea10b59be24/go.mod h1:4UJr5HIiMZrwgkSPdsjy2uOQExX/WEILpIrO9UPGuXs= github.com/GaijinEntertainment/go-exhaustruct/v3 v3.3.1 h1:Sz1JIXEcSfhz7fUi7xHnhpIE0thVASYjvosApmHuD2k= github.com/GaijinEntertainment/go-exhaustruct/v3 v3.3.1/go.mod h1:n/LSCXNuIYqVfBlVXyHfMQkZDdp1/mmxfSjADd3z1Zg= +github.com/Masterminds/semver v1.5.0 h1:H65muMkzWKEuNDnfl9d70GUjFniHKHRbFPGBuZ3QEww= +github.com/Masterminds/semver v1.5.0/go.mod h1:MB6lktGJrhw8PrUyiEoblNEGEQ+RzHPF078ddwwvV3Y= github.com/Masterminds/semver/v3 v3.3.1 h1:QtNSWtVZ3nBfk8mAOu/B6v7FMJ+NHTIgUPi7rj+4nv4= github.com/Masterminds/semver/v3 v3.3.1/go.mod h1:4V+yj/TJE1HU9XfppCwVMZq3I84lprf4nC11bSS5beM= github.com/OpenPeeDeeP/depguard/v2 v2.2.1 h1:vckeWVESWp6Qog7UZSARNqfu/cZqvki8zsuj3piCMx4= diff --git a/guard/guard.go b/guard/guard.go index 8795951..8e05384 100644 --- a/guard/guard.go +++ b/guard/guard.go @@ -69,14 +69,9 @@ 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) - if err != nil { - return fmt.Errorf("failed to parse command: %w", err) - } - if !parsedCommand.HasInstallTarget() { log.Debugf("No install target found, continuing execution") return g.continueExecution(ctx, parsedCommand) diff --git a/internal/flows/common.go b/internal/flows/common.go index a8cbaae..a49ac90 100644 --- a/internal/flows/common.go +++ b/internal/flows/common.go @@ -11,37 +11,26 @@ import ( ) type commonFlow struct { - pm packagemanager.PackageManager + pm packagemanager.PackageManager + packageResolver packagemanager.PackageResolver + config config.Config } // Creates a common flow of execution for all package managers. This should work for most // of the cases unless a package manager has its own unique requirements. Configuration // should be passed through the context (Global Config) -func Common(pm packagemanager.PackageManager) *commonFlow { +func Common(pm packagemanager.PackageManager, pkgResolver packagemanager.PackageResolver, config config.Config) *commonFlow { return &commonFlow{ - pm: pm, + pm: pm, + packageResolver: pkgResolver, + config: config, } } -func (f *commonFlow) Run(ctx context.Context, args []string) error { - config, err := config.FromContext(ctx) - if err != nil { - ui.Fatalf("Failed to get config: %s", err) - } - - packageResolverConfig := packagemanager.NewDefaultNpmDependencyResolverConfig() - packageResolverConfig.IncludeTransitiveDependencies = config.Transitive - packageResolverConfig.TransitiveDepth = config.TransitiveDepth - packageResolverConfig.IncludeDevDependencies = config.IncludeDevDependencies - - packageResolver, err := packagemanager.NewNpmDependencyResolver(packageResolverConfig) - if err != nil { - ui.Fatalf("Failed to create dependency resolver: %s", err) - } - +func (f *commonFlow) Run(ctx context.Context, args []string, parsedCmd *packagemanager.ParsedCommand) error { var analyzers []analyzer.PackageVersionAnalyzer - if config.Paranoid { + if f.config.Paranoid { malysisActiveScanAnalyzer, err := analyzer.NewMalysisActiveScanAnalyzer(analyzer.DefaultMalysisActiveScanAnalyzerConfig()) if err != nil { ui.Fatalf("Failed to create malware analyzer: %s", err) @@ -65,14 +54,14 @@ func (f *commonFlow) Run(ctx context.Context, args []string) error { } guardConfig := guard.DefaultPackageManagerGuardConfig() - guardConfig.DryRun = config.DryRun + guardConfig.DryRun = f.config.DryRun - proxy, err := guard.NewPackageManagerGuard(guardConfig, f.pm, packageResolver, analyzers, interaction) + proxy, err := guard.NewPackageManagerGuard(guardConfig, f.pm, f.packageResolver, analyzers, interaction) if err != nil { 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/main.go b/main.go index 2196265..a241a8f 100644 --- a/main.go +++ b/main.go @@ -6,6 +6,7 @@ import ( "github.com/safedep/dry/log" "github.com/safedep/pmg/cmd/npm" + "github.com/safedep/pmg/cmd/pypi" "github.com/safedep/pmg/cmd/version" "github.com/safedep/pmg/config" "github.com/safedep/pmg/internal/ui" @@ -79,6 +80,7 @@ func main() { cmd.AddCommand(npm.NewNpmCommand()) cmd.AddCommand(npm.NewPnpmCommand()) + cmd.AddCommand(pypi.NewPipCommand()) cmd.AddCommand(version.NewVersionCommand()) if err := cmd.Execute(); err != nil { diff --git a/packagemanager/dependency_resolver.go b/packagemanager/dependency_resolver.go index 1f3c7a1..a9a1e5e 100644 --- a/packagemanager/dependency_resolver.go +++ b/packagemanager/dependency_resolver.go @@ -3,7 +3,6 @@ package packagemanager import ( "context" "fmt" - "slices" "sync" packagev1 "buf.build/gen/go/safedep/api/protocolbuffers/go/safedep/messages/package/v1" @@ -13,7 +12,11 @@ import ( // Contract for a function that implements ecosystem specific version // resolver from a version range specification. -type versionSpecResolver func(version string) string +type versionSpecResolverFn func(packageName, version string) string + +type dependencyResolverFn func(packageName, version string) (*packageregistry.PackageDependencyList, error) + +type packageIdentifierFn func(pkg *packagev1.PackageVersion) string type dependencyResolverConfig struct { IncludeDevDependencies bool @@ -24,29 +27,35 @@ type dependencyResolverConfig struct { } type dependencyResolver struct { - client packageregistry.Client - config dependencyResolverConfig - mutex sync.Mutex - versionSpecResolver versionSpecResolver + client packageregistry.Client + config dependencyResolverConfig + mutex sync.Mutex + versionSpecResolver versionSpecResolverFn + packageDependencyResolver dependencyResolverFn + packageIdentifierFn packageIdentifierFn + resultSet map[string]bool } func newDependencyResolver(client packageregistry.Client, config dependencyResolverConfig, - versionSpecResolver versionSpecResolver) *dependencyResolver { + versionSpecResolver versionSpecResolverFn, packageDependencyResolver dependencyResolverFn, packageKeyFn packageIdentifierFn) *dependencyResolver { if config.MaxConcurrency <= 0 { config.MaxConcurrency = 10 } if versionSpecResolver == nil { // Default version spec resolver - versionSpecResolver = func(version string) string { + versionSpecResolver = func(packageName, version string) string { return version } } return &dependencyResolver{ - client: client, - config: config, - versionSpecResolver: versionSpecResolver, + client: client, + config: config, + versionSpecResolver: versionSpecResolver, + packageDependencyResolver: packageDependencyResolver, + packageIdentifierFn: packageKeyFn, + resultSet: make(map[string]bool), } } @@ -63,6 +72,7 @@ func (r *dependencyResolver) resolveDependencies(ctx context.Context, // Result collection dependencies := make([]*packagev1.PackageVersion, 0) + r.resultSet = make(map[string]bool) // Reset // Start concurrent resolution err = r.resolvePackageDependenciesConcurrent(ctx, pd, packageVersion, 0, visitedPackages, &dependencies) if err != nil { @@ -102,27 +112,42 @@ func (r *dependencyResolver) resolvePackageDependenciesConcurrent( return ff(fmt.Errorf("exceeded maximum transitive depth of %d", r.config.TransitiveDepth)) } + var packageKey string + var packageKeyFn packageIdentifierFn + // 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() { - 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 } - // Mark the current package as visited - r.synchronize(func() { - 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) + var dependencyList *packageregistry.PackageDependencyList + var err error + if r.packageDependencyResolver != nil { + dependencyList, err = r.packageDependencyResolver(packageVersion.Package.Name, packageVersion.Version) + } else { + dependencyList, err = pd.GetPackageDependencies(packageVersion.Package.Name, packageVersion.Version) + } + if err != nil { return ff(fmt.Errorf("failed to get package dependencies: %w", err)) } @@ -141,14 +166,17 @@ func (r *dependencyResolver) resolvePackageDependenciesConcurrent( Ecosystem: packageVersion.GetPackage().GetEcosystem(), Name: dependency.Name, }, - Version: r.versionSpecResolver(dependency.VersionSpec), + Version: r.versionSpecResolver(dependency.Name, dependency.VersionSpec), }) } // Add resolved dependencies to the result r.synchronize(func() { for _, dependency := range resolvedDependencies { - if !slices.Contains(*result, dependency) { + dependencyKey := packageKeyFn(dependency) + + if !r.resultSet[dependencyKey] { + r.resultSet[dependencyKey] = true *result = append(*result, dependency) } } @@ -190,7 +218,7 @@ func (r *dependencyResolver) resolvePackageDependenciesConcurrent( 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) } diff --git a/packagemanager/errors.go b/packagemanager/errors.go new file mode 100644 index 0000000..beb9127 --- /dev/null +++ b/packagemanager/errors.go @@ -0,0 +1,15 @@ +package packagemanager + +import ( + "errors" +) + +var ( + ErrPackageNotFound = errors.New("package not found") + ErrFailedToFetchPackage = errors.New("failed to fetch package") + ErrFailedToParsePackage = errors.New("failed to parse package") + ErrNoPackagesFound = errors.New("no packages found") + ErrAuthorNotFound = errors.New("author not found") + + ErrGitHubRateLimitExceeded = errors.New("github api rate limit exceeded") +) diff --git a/packagemanager/npm.go b/packagemanager/npm.go index ffed095..60154d0 100644 --- a/packagemanager/npm.go +++ b/packagemanager/npm.go @@ -37,6 +37,8 @@ func NewNpmPackageManager(config NpmPackageManagerConfig) (*npmPackageManager, e }, nil } +var _ PackageManager = &npmPackageManager{} + func (npm *npmPackageManager) Name() string { return "npm" } diff --git a/packagemanager/npm_resolver.go b/packagemanager/npm_resolver.go index 68293a4..9cc6c7c 100644 --- a/packagemanager/npm_resolver.go +++ b/packagemanager/npm_resolver.go @@ -72,13 +72,18 @@ func (r *npmDependencyResolver) ResolveLatestVersion(ctx context.Context, func (r *npmDependencyResolver) ResolveDependencies(ctx context.Context, packageVersion *packagev1.PackageVersion) ([]*packagev1.PackageVersion, error) { + + npmVersionSpecResolverFn := func(packageName, version string) string { + return npmCleanVersion(version) + } + resolver := newDependencyResolver(r.registry, dependencyResolverConfig{ IncludeDevDependencies: r.config.IncludeDevDependencies, IncludeTransitiveDependencies: r.config.IncludeTransitiveDependencies, TransitiveDepth: r.config.TransitiveDepth, FailFast: r.config.FailFast, MaxConcurrency: r.config.MaxConcurrency, - }, npmCleanVersion) + }, npmVersionSpecResolverFn, nil, nil) return resolver.resolveDependencies(ctx, packageVersion) } diff --git a/packagemanager/packagemanager.go b/packagemanager/packagemanager.go index 5ac1d9d..fbedd87 100644 --- a/packagemanager/packagemanager.go +++ b/packagemanager/packagemanager.go @@ -13,6 +13,11 @@ type Command struct { type PackageInstallTarget struct { PackageVersion *packagev1.PackageVersion + + // Extras specifies additional features to be installed with a Python package + // Example: "django[mysql,redis]" has Extras as ["mysql", "redis"] + // Currently only specific to Python packages + Extras []string } func (pit *PackageInstallTarget) HasVersion() bool { diff --git a/packagemanager/pypi.go b/packagemanager/pypi.go new file mode 100644 index 0000000..b437355 --- /dev/null +++ b/packagemanager/pypi.go @@ -0,0 +1,204 @@ +package packagemanager + +import ( + "fmt" + "slices" + "strconv" + "strings" + + packagev1 "buf.build/gen/go/safedep/api/protocolbuffers/go/safedep/messages/package/v1" +) + +type PipPackageManagerConfig struct { + InstallCommands []string + CommandName string +} + +func DefaultPipPackageManagerConfig() PipPackageManagerConfig { + return PipPackageManagerConfig{ + InstallCommands: []string{"install"}, + CommandName: "pip", + } +} + +type pipPackageManager struct { + Config PipPackageManagerConfig +} + +func NewPipPackageManager(config PipPackageManagerConfig) (*pipPackageManager, error) { + return &pipPackageManager{ + Config: config, + }, nil +} + +var _ PackageManager = &pipPackageManager{} + +func (pip *pipPackageManager) Name() string { + return "pip" +} + +func (pip *pipPackageManager) ParseCommand(args []string) (*ParsedCommand, error) { + if len(args) > 0 && args[0] == "pip" { + args = args[1:] + } + command := Command{Exe: pip.Config.CommandName, Args: args} + + if len(args) < 2 { + return &ParsedCommand{ + Command: command, + }, nil + } + + var packages []string + for idx, arg := range args { + if slices.Contains(pip.Config.InstallCommands, arg) { + for i := idx + 1; i < len(args); i++ { + if strings.HasPrefix(args[i], "-") { + continue + } + packages = append(packages, args[i]) + } + break + } + } + var installTargets []*PackageInstallTarget + + for _, pkg := range packages { + packageName, version, extras, err := pipParsePackageInfo(pkg) + if err != nil { + return nil, fmt.Errorf("failed to parse package info: %w", err) + } + + if version != "" { + // If exact version provided just trim it. If not get a version that satisfies a given version specifier + if strings.HasPrefix(version, "==") { + // Exact version, just trim + version = strings.TrimPrefix(version, "==") + } else { + // Version range, resolve from PyPI + version, err = pipGetMatchingVersion(packageName, version) + if err != nil { + return nil, fmt.Errorf("error resolving version for %s: %s", packageName, err.Error()) + } + } + } + + installTargets = append(installTargets, &PackageInstallTarget{ + PackageVersion: &packagev1.PackageVersion{ + Package: &packagev1.Package{ + Ecosystem: packagev1.Ecosystem_ECOSYSTEM_PYPI, + Name: packageName, + }, + Version: version, + }, + Extras: extras, + }) + } + + return &ParsedCommand{ + Command: command, + InstallTargets: installTargets, + }, nil +} + +// pipParsePackageInfo parses a pip install package specification, separating the package name, +// version constraints, and any extras (additional features) to be installed. +// Example: "django[mysql,redis]>=3.0" returns ("django", ">=3.0", ["mysql", "redis"], nil) +func pipParsePackageInfo(input string) (packageName, version string, extras []string, err error) { + if input == "" { + return "", "", nil, fmt.Errorf("package info cannot be empty") + } + + input = strings.TrimSpace(input) + + // First extract any extras if present + openBracket := strings.Index(input, "[") + closeBracket := strings.Index(input, "]") + + if openBracket != -1 && closeBracket != -1 && openBracket < closeBracket { + extrasStr := strings.TrimSpace(input[openBracket+1 : closeBracket]) + if extrasStr != "" { + // Split extras by comma and trim each extra + for _, extra := range strings.Split(extrasStr, ",") { + if trimmedExtra := strings.TrimSpace(extra); trimmedExtra != "" { + extras = append(extras, trimmedExtra) + } + } + } + // Remove the extra part from input for further processing + input = input[:openBracket] + input[closeBracket+1:] + } else if (openBracket != -1 && closeBracket == -1) || (openBracket == -1 && closeBracket != -1) { + return "", "", nil, fmt.Errorf("mismatched brackets in input '%s'", input) + } + + // Python package version specifiers are typically separated by one of: + // '==', '>=', '<=', '!=', '>', '<', '~=', or direct comma separated list + operators := []string{"==", ">=", "<=", "!=", ">", "<", "~="} + index := -1 + + // Find the earliest operator occurrence + for _, op := range operators { + i := strings.Index(input, op) + if i != -1 && (index == -1 || i < index) { + index = i + } + } + + if index == -1 { + // No operator found, whole input is package name, no version + return strings.TrimSpace(input), "", extras, nil + } + + packageName = strings.TrimSpace(input[:index]) + version = strings.TrimSpace(input[index:]) + + if packageName == "" { + return "", "", nil, fmt.Errorf("invalid package name in input '%s'", input) + } + + return packageName, version, extras, nil +} + +func pipConvertCompatibleRelease(version string) string { + if !strings.HasPrefix(version, "~=") { + return version + } + + version = strings.TrimPrefix(version, "~=") + parts := strings.Split(version, ".") + if len(parts) < 2 { + return "" // invalid + } + + switch len(parts) { + case 2: + // ~=X.Y case, increment major version: ~=2.1 -> >=2.1,<3.0 + major := parts[0] + nextMajor, _ := strconv.Atoi(major) + nextMajor += 1 + return fmt.Sprintf(">=%s,<%d.0", version, nextMajor) + + case 3: + // ~=X.Y.Z case, increment minor version: ~=2.1.5 -> >=2.1.5,<2.2.0 + major := parts[0] + minor := parts[1] + nextMinor, _ := strconv.Atoi(minor) + nextMinor += 1 + return fmt.Sprintf(">=%s,<%s.%d.0", version, major, nextMinor) + + default: + // ~=X.Y.Z.W[.more] case, increment second-to-last component + // ~=2.1.5.2 -> >=2.1.5.2,<2.1.6 + incIndex := len(parts) - 2 + upperBoundParts := make([]string, incIndex+1) + copy(upperBoundParts, parts[:incIndex+1]) + + increment, _ := strconv.Atoi(upperBoundParts[incIndex]) + increment++ + upperBoundParts[incIndex] = strconv.Itoa(increment) + + upperBound := strings.Join(upperBoundParts, ".") + + return fmt.Sprintf(">=%s,<%s", version, upperBound) + } +} diff --git a/packagemanager/pypi_resolver.go b/packagemanager/pypi_resolver.go new file mode 100644 index 0000000..417046f --- /dev/null +++ b/packagemanager/pypi_resolver.go @@ -0,0 +1,346 @@ +package packagemanager + +import ( + "context" + "encoding/json" + "fmt" + "net/http" + "regexp" + "slices" + "strings" + + packagev1 "buf.build/gen/go/safedep/api/protocolbuffers/go/safedep/messages/package/v1" + "github.com/Masterminds/semver" + "github.com/safedep/dry/log" + "github.com/safedep/dry/packageregistry" +) + +type PyPiDependencyResolverConfig struct { + IncludeDevDependencies bool + IncludeTransitiveDependencies bool + TransitiveDepth int + + // FailFast will stop resolving dependencies after the first error + FailFast bool + + // MaxConcurrency limits the number of concurrent goroutines used for dependency resolution + MaxConcurrency int + + PackageInstallTargets []*PackageInstallTarget +} + +func NewDefaultPypiDependencyResolverConfig() PyPiDependencyResolverConfig { + return PyPiDependencyResolverConfig{ + IncludeDevDependencies: false, + IncludeTransitiveDependencies: true, + TransitiveDepth: 5, + FailFast: false, + MaxConcurrency: 10, + PackageInstallTargets: []*PackageInstallTarget{}, + } +} + +type pypiDependencyResolver struct { + registry packageregistry.Client + config PyPiDependencyResolverConfig +} + +var _ PackageResolver = &pypiDependencyResolver{} + +func NewPypiDependencyResolver(config PyPiDependencyResolverConfig) (*pypiDependencyResolver, error) { + client, err := packageregistry.NewPypiAdapter() + if err != nil { + return nil, fmt.Errorf("failed to create pypi adapter: %w", err) + } + + return &pypiDependencyResolver{ + config: config, + registry: client, + }, nil +} + +func (p *pypiDependencyResolver) ResolveDependencies(ctx context.Context, pkg *packagev1.PackageVersion) ([]*packagev1.PackageVersion, error) { + pypiVersionSpecResolverFn := func(packageName, version string) string { + ver, err := pipGetMatchingVersion(packageName, version) + if err != nil { + log.Debugf("error getting matching version for %s@%s", packageName, version) + return "" + } + return ver + } + + pypiDependencyResolverFn := func(packageName, version string) (*packageregistry.PackageDependencyList, error) { + resolvedDependencies, err := getPypiPackageDependencies(packageName, version, p.config.PackageInstallTargets) + if err != nil { + return nil, err + } + dependencies := make([]packageregistry.PackageDependencyInfo, 0) + for _, dep := range resolvedDependencies { + dependencies = append(dependencies, packageregistry.PackageDependencyInfo{ + Name: dep.PackageNameExtra, + VersionSpec: dep.VersionSpec, + }) + } + + return &packageregistry.PackageDependencyList{ + Dependencies: dependencies, + }, 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{ + IncludeDevDependencies: p.config.IncludeDevDependencies, + IncludeTransitiveDependencies: p.config.IncludeTransitiveDependencies, + TransitiveDepth: p.config.TransitiveDepth, + FailFast: p.config.FailFast, + MaxConcurrency: p.config.MaxConcurrency, + }, pypiVersionSpecResolverFn, pypiDependencyResolverFn, packageKeyFn) + + return resolver.resolveDependencies(ctx, pkg) +} + +func (p *pypiDependencyResolver) ResolveLatestVersion(ctx context.Context, pkg *packagev1.Package) (*packagev1.PackageVersion, error) { + pd, err := p.registry.PackageDiscovery() + if err != nil { + return nil, fmt.Errorf("failed to get package discovery: %w", err) + } + + pkgInfo, err := pd.GetPackage(pkg.Name) + if err != nil { + return nil, fmt.Errorf("failed to get package: %w", err) + } + log.Debugf("Resolved pypi/%s to latest version %s", pkg.Name, pkgInfo.LatestVersion) + + return &packagev1.PackageVersion{ + Package: pkg, + Version: pkgInfo.LatestVersion, + }, nil +} + +type pypiPackage struct { + Info pypiPackageInfo `json:"info"` + Releases map[string]any `json:"releases"` +} + +type PyPIDependencySpec struct { + // PackageNameExtra is the package name including any direct extras in brackets + // Example: "uvicorn[standard]" + PackageNameExtra string + + // VersionSpec is the version constraint for the package + // Example: ">=0.12.0", "==1.0.0", ">=2.0,<3.0" + VersionSpec string + + // Extra is the conditional extra marker that defines when this dependency applies + // Example: "all" from "; extra == \"all\"" + Extra string +} + +type pypiPackageInfo struct { + Name string `json:"name"` + Description string `json:"summary"` + LatestVersion string `json:"version"` + PackageURL string `json:"package_url"` + Author string `json:"author"` + AuthorEmail string `json:"author_email"` + Maintainer string `json:"maintainer"` + MaintainerEmail string `json:"maintainer_email"` + RequiresDist []string `json:"requires_dist"` +} + +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) + if err != nil { + return nil, ErrFailedToFetchPackage + } + + if res.StatusCode == 404 { + return nil, ErrPackageNotFound + } + + if res.StatusCode != 200 { + return nil, ErrFailedToFetchPackage + } + defer res.Body.Close() + + var pypipkg pypiPackage + err = json.NewDecoder(res.Body).Decode(&pypipkg) + if err != nil { + 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) + + // 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 && slices.Contains(requestedExtras, extra)) { + pkgDeps = append(pkgDeps, PyPIDependencySpec{ + PackageNameExtra: name, + VersionSpec: version, + Extra: extra, + }) + } + } + + return pkgDeps, nil +} + +// pypiParseDependency parses a PyPI dependency specification, handling both package extras +// and conditional dependencies. Keeps extras as part of the package name. +// Example: "uvicorn[standard]>=0.12.0; extra == \"all\"" returns ("uvicorn[standard]", ">=0.12.0", "all") +func pypiParseDependency(input string) (string, string, string) { + // Split line by ';' to separate version and markers + parts := strings.SplitN(input, ";", 2) + mainPart := strings.TrimSpace(parts[0]) + + // Regex to match the first occurrence of version operators + // Using lookahead to ensure we match standalone operators + operatorRegex := regexp.MustCompile(`(==|>=|<=|!=|>|<|~=)(?:\d|$)`) + match := operatorRegex.FindStringIndex(mainPart) + + var name, version string + if match != nil { + // Everything before the operator is the name + name = strings.TrimSpace(mainPart[:match[0]]) + // Remove trailing parentheses from name if present + name = strings.TrimRight(name, " (") + + // Everything from the operator onwards is the version spec + version = strings.TrimSpace(mainPart[match[0]:]) + // Remove parentheses from version spec if present + version = strings.Trim(version, "()") + } else { + // No version operator found + name = mainPart + version = "" + } + + // Extract extra marker if present + var extra string + if len(parts) == 2 { + extraRe := regexp.MustCompile(`extra\s*==\s*["']([^"']+)["']`) + if match := extraRe.FindStringSubmatch(parts[1]); len(match) == 2 { + extra = match[1] + } + } + + return name, version, extra +} + +func pipGetMatchingVersion(packageName, versionConstraint string) (string, error) { + // Already a exact version + if strings.HasPrefix(versionConstraint, "==") { + return versionConstraint, nil + } + + // Handle compatible release operator + if strings.HasPrefix(versionConstraint, "~=") { + versionConstraint = pipConvertCompatibleRelease(versionConstraint) + } + // Handle empty version constraint + if versionConstraint == "" { + // Get latest version + registry, err := packageregistry.NewPypiAdapter() + if err != nil { + return "", fmt.Errorf("failed to create pypi adapter: %w", err) + } + + pd, err := registry.PackageDiscovery() + if err != nil { + return "", fmt.Errorf("failed to get package discovery: %w", err) + } + + pkg, err := pd.GetPackage(packageName) + if err != nil { + return "", err + } + + return pkg.LatestVersion, nil + } + + registry, err := packageregistry.NewPypiAdapter() + if err != nil { + return "", fmt.Errorf("failed to create pypi adapter: %w", err) + } + + pd, err := registry.PackageDiscovery() + if err != nil { + return "", fmt.Errorf("failed to get package discovery: %w", err) + } + + // Get package info from PyPI + pkg, err := pd.GetPackage(packageName) + if err != nil { + return "", err + } + + // Parse version constraint + constraint, err := semver.NewConstraint(versionConstraint) + if err != nil { + return "", fmt.Errorf("invalid version constraint: %w", err) + } + + // Get valid versions and find best match + bestMatch, err := findBestMatchingVersion(pkg.Versions, constraint) + if err != nil { + return "", fmt.Errorf("no version matches constraint %q: %w", versionConstraint, err) + } + + return bestMatch.Original(), nil +} + +func findBestMatchingVersion(releases []packageregistry.PackageVersionInfo, constraint *semver.Constraints) (*semver.Version, error) { + if len(releases) == 0 { + return nil, fmt.Errorf("no versions available") + } + + var bestMatch *semver.Version + // We'll iterate once through all versions + for _, v := range releases { + ver, err := semver.NewVersion(v.Version) + if err != nil { + continue // Skip invalid versions + } + + // Update bestMatch if this version is higher and matches constraint + if constraint.Check(ver) && (bestMatch == nil || ver.GreaterThan(bestMatch)) { + bestMatch = ver + } + } + + if bestMatch == nil { + return nil, fmt.Errorf("no version matches constraint") + } + 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 +} diff --git a/packagemanager/pypi_resolver_test.go b/packagemanager/pypi_resolver_test.go new file mode 100644 index 0000000..1f6a347 --- /dev/null +++ b/packagemanager/pypi_resolver_test.go @@ -0,0 +1,202 @@ +package packagemanager + +import ( + "context" + "fmt" + "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 TestPypiDependencyResolver_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: "requests", + Ecosystem: packagev1.Ecosystem_ECOSYSTEM_PYPI, + }, + assertFn: func(t *testing.T, pv *packagev1.PackageVersion, err error) { + require.NoError(t, err) + require.True(t, semver.IsAhead("2.30.0", pv.Version)) + }, + }, + { + name: "should return an error if the package is not found", + pkg: &packagev1.Package{ + Name: "nonexistent-package-12345", + Ecosystem: packagev1.Ecosystem_ECOSYSTEM_PYPI, + }, + 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 := NewPypiDependencyResolver(NewDefaultPypiDependencyResolverConfig()) + require.NoError(t, err) + + pv, err := resolver.ResolveLatestVersion(context.Background(), tc.pkg) + tc.assertFn(t, pv, err) + }) + } +} + +func TestPipGetLatestMatchingVersion(t *testing.T) { + cases := []struct { + name string + packageName string + versionConstraint string + assertFn func(t *testing.T, version string, err error) + }{ + { + name: "should resolve exact version", + packageName: "requests", + versionConstraint: "==2.28.0", + assertFn: func(t *testing.T, version string, err error) { + require.NoError(t, err) + require.Equal(t, "==2.28.0", version) + }, + }, + { + name: "should resolve compatible version", + packageName: "requests", + versionConstraint: "~=2.26.0", + assertFn: func(t *testing.T, version string, err error) { + require.NoError(t, err) + require.NotEmpty(t, version) + fmt.Println("Version: ", version) + require.True(t, semver.IsAheadOrEqual("2.26.0", version) && !semver.IsAhead("2.27.0", version)) + }, + }, + { + name: "should return error for nonexistent package", + packageName: "nonexistent-package-12345", + versionConstraint: ">=1.0.0", + assertFn: func(t *testing.T, version string, err error) { + require.Error(t, err) + require.Empty(t, version) + }, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + version, err := pipGetMatchingVersion(tc.packageName, tc.versionConstraint) + tc.assertFn(t, version, err) + }) + } +} + +func TestPypiParseDependency(t *testing.T) { + tests := []struct { + name string + input string + wantName string + wantVersion string + wantExtra string + }{ + { + name: "simple package without version", + input: "requests", + wantName: "requests", + wantVersion: "", + wantExtra: "", + }, + { + name: "package with exact version", + input: "requests==2.28.1", + wantName: "requests", + wantVersion: "==2.28.1", + wantExtra: "", + }, + { + name: "package with greater than version", + input: "django>=4.2.0", + wantName: "django", + wantVersion: ">=4.2.0", + wantExtra: "", + }, + { + name: "package with less than version", + input: "pylint<3.0.0", + wantName: "pylint", + wantVersion: "<3.0.0", + wantExtra: "", + }, + { + name: "package with not equal version", + input: "pytest!=3.0.0", + wantName: "pytest", + wantVersion: "!=3.0.0", + wantExtra: "", + }, + { + name: "package with compatible release version", + input: "sphinx~=4.0.0", + wantName: "sphinx", + wantVersion: "~=4.0.0", + wantExtra: "", + }, + { + name: "package with extra", + input: "requests;extra=='security'", + wantName: "requests", + wantVersion: "", + wantExtra: "security", + }, + { + name: "package with version and extra", + input: "requests>=2.28.1;extra=='security'", + wantName: "requests", + wantVersion: ">=2.28.1", + wantExtra: "security", + }, + { + name: "package with single quotes in extra", + input: "django>=4.2.0;extra=='testing'", + wantName: "django", + wantVersion: ">=4.2.0", + wantExtra: "testing", + }, + { + name: "package with double quotes in extra", + input: "django>=4.2.0;extra==\"testing\"", + wantName: "django", + wantVersion: ">=4.2.0", + wantExtra: "testing", + }, + { + name: "package with multiple version constraints", + input: "requests>=2.28.1,<3.0.0", + wantName: "requests", + wantVersion: ">=2.28.1,<3.0.0", + wantExtra: "", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + gotName, gotVersion, gotExtra := pypiParseDependency(tt.input) + + if gotName != tt.wantName { + t.Errorf("pypiParseDependency() gotName = %v, want %v", gotName, tt.wantName) + } + if gotVersion != tt.wantVersion { + t.Errorf("pypiParseDependency() gotVersion = %v, want %v", gotVersion, tt.wantVersion) + } + if gotExtra != tt.wantExtra { + t.Errorf("pypiParseDependency() gotExtra = %v, want %v", gotExtra, tt.wantExtra) + } + }) + } +} diff --git a/packagemanager/pypi_test.go b/packagemanager/pypi_test.go new file mode 100644 index 0000000..928a33e --- /dev/null +++ b/packagemanager/pypi_test.go @@ -0,0 +1,158 @@ +package packagemanager + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestPipParsePackageInfo(t *testing.T) { + cases := []struct { + name string + input string + pkgName string + version string + extras []string + wantErr bool + }{ + { + name: "simple package name", + input: "fastapi", + pkgName: "fastapi", + version: "", + extras: nil, + wantErr: false, + }, + { + name: "package with exact version & extra", + input: "fastapi[all]==0.115.7", + pkgName: "fastapi", + version: "==0.115.7", + extras: []string{"all"}, + wantErr: false, + }, + { + name: "package with version range", + input: "requests>=2.0,<3.0", + pkgName: "requests", + version: ">=2.0,<3.0", + extras: nil, + wantErr: false, + }, + { + name: "package with exclusion", + input: "pydantic!=1.8,!=1.8.1", + pkgName: "pydantic", + version: "!=1.8,!=1.8.1", + wantErr: false, + }, + { + name: "package with compatible release", + input: "django~=3.1.0", + pkgName: "django", + version: "~=3.1.0", + extras: nil, + wantErr: false, + }, + { + name: "package with greater than with empty extra", + input: "numpy[]>1.20.0", + pkgName: "numpy", + version: ">1.20.0", + extras: nil, + wantErr: false, + }, + { + name: "package with less than", + input: "pandas<2.0.0", + pkgName: "pandas", + version: "<2.0.0", + extras: nil, + wantErr: false, + }, + { + name: "empty input", + input: "", + pkgName: "", + version: "", + extras: nil, + wantErr: true, + }, + { + name: "only version specifier", + input: "==1.0.0", + pkgName: "", + version: "", + extras: nil, + wantErr: true, + }, + { + name: "package with whitespace", + input: " requests == 2.0.0 ", + pkgName: "requests", + version: "== 2.0.0", + extras: nil, + wantErr: false, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + pkgName, version, extras, err := pipParsePackageInfo(tc.input) + if tc.wantErr { + assert.Error(t, err) + } else { + assert.NoError(t, err) + assert.Equal(t, tc.pkgName, pkgName) + assert.Equal(t, tc.version, version) + assert.Equal(t, tc.extras, extras) + } + }) + } +} + +func TestPipConvertCompatibleRelease(t *testing.T) { + cases := []struct { + name string + input string + expected string + }{ + { + name: "standard version", + input: "~=3.1.0", + expected: ">=3.1.0,<3.2.0", + }, + { + name: "single digit minor", + input: "~=2.1.5", + expected: ">=2.1.5,<2.2.0", + }, + { + name: "double digit minor", + input: "~=1.10.0", + expected: ">=1.10.0,<1.11.0", + }, + { + name: "invalid format", + input: "~=1", + expected: "", + }, + { + name: "missing prefix", + input: "3.1.0", + expected: "3.1.0", + }, + { + name: "extra segments", + input: "~=2.1.5.2", + expected: ">=2.1.5.2,<2.1.6", + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + result := pipConvertCompatibleRelease(tc.input) + assert.Equal(t, tc.expected, result) + }) + } +}