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 <bansalsahil315@gmail.com>

* 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 <omkarphansopkar@gmail.com>
Signed-off-by: Sahil Bansal <bansalsahil315@gmail.com>

* Update cmd/npm/pnpm.go

Co-authored-by: Omkar Phansopkar <omkarphansopkar@gmail.com>
Signed-off-by: Sahil Bansal <bansalsahil315@gmail.com>

* Update cmd/pypi/pip.go

Co-authored-by: Omkar Phansopkar <omkarphansopkar@gmail.com>
Signed-off-by: Sahil Bansal <bansalsahil315@gmail.com>

---------

Signed-off-by: Sahil Bansal <bansalsahil315@gmail.com>
Co-authored-by: Omkar Phansopkar <omkarphansopkar@gmail.com>
This commit is contained in:
Sahil Bansal
2025-06-09 17:37:08 +05:30
committed by GitHub
co-authored by Omkar Phansopkar
parent 4031219375
commit 53783c6604
17 changed files with 1111 additions and 57 deletions
+23 -2
View File
@@ -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)
}
+20 -1
View File
@@ -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)
}
+60
View File
@@ -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)
}
+1
View File
@@ -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
+2
View File
@@ -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=
+1 -6
View File
@@ -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)
+12 -23
View File
@@ -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)
}
+2
View File
@@ -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 {
+52 -24
View File
@@ -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)
}
+15
View File
@@ -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")
)
+2
View File
@@ -37,6 +37,8 @@ func NewNpmPackageManager(config NpmPackageManagerConfig) (*npmPackageManager, e
}, nil
}
var _ PackageManager = &npmPackageManager{}
func (npm *npmPackageManager) Name() string {
return "npm"
}
+6 -1
View File
@@ -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)
}
+5
View File
@@ -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 {
+204
View File
@@ -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)
}
}
+346
View File
@@ -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
}
+202
View File
@@ -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)
}
})
}
}
+158
View File
@@ -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)
})
}
}