feat: add Python dependency parsing with extras support

This commit is contained in:
Sahilb315
2025-05-22 02:05:34 +05:30
parent 0b4729f18b
commit 216657afef
6 changed files with 178 additions and 383 deletions
-362
View File
@@ -1,362 +0,0 @@
diff --git a/cmd/npm/npm.go b/cmd/npm/npm.go
index e89dc36..5e8f86f 100644
--- a/cmd/npm/npm.go
+++ b/cmd/npm/npm.go
@@ -4,6 +4,7 @@ import (
"context"
"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"
@@ -31,6 +32,20 @@ func executeNpmFlow(ctx context.Context, args []string) error {
if err != nil {
ui.Fatalf("Failed to create npm package manager proxy: %s", err)
}
+ 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)
+ }
- return flows.Common(packageManager).Run(ctx, args)
+ return flows.Common(packageManager, packageResolver, config).Run(ctx, args)
}
diff --git a/cmd/npm/pnpm.go b/cmd/npm/pnpm.go
index c161c0d..a010244 100644
--- a/cmd/npm/pnpm.go
+++ b/cmd/npm/pnpm.go
@@ -5,6 +5,7 @@ import (
_ "embed"
"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"
@@ -32,6 +33,17 @@ func executePnpmFlow(ctx context.Context, args []string) error {
if err != nil {
ui.Fatalf("Failed to create pnpm package manager proxy: %s", err)
}
+ 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)
- return flows.Common(packageManager).Run(ctx, args)
+ return flows.Common(packageManager, packageResolver, config).Run(ctx, args)
}
diff --git a/cmd/pypi/pip.go b/cmd/pypi/pip.go
index 56d03f7..4cc9a65 100644
--- a/cmd/pypi/pip.go
+++ b/cmd/pypi/pip.go
@@ -6,6 +6,7 @@ import (
"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"
@@ -17,12 +18,7 @@ func NewPipCommand() *cobra.Command {
Short: "Guard pip package manager",
DisableFlagParsing: true,
RunE: func(cmd *cobra.Command, args []string) error {
- config, err := config.FromContext(cmd.Context())
- if err != nil {
- ui.Fatalf("Failed to get config: %s", err)
- }
-
- err = executePipFlow(cmd.Context(), config, args)
+ err := executePipFlow(cmd.Context(), args)
if err != nil {
log.Errorf("Failed to execute pip flow: %s", err)
}
@@ -32,12 +28,24 @@ func NewPipCommand() *cobra.Command {
}
}
-func executePipFlow(context context.Context, config config.Config, args []string) error {
+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)
}
- cmd, _ := packageManager.ParseCommand(args)
- fmt.Println("Cmd: ", cmd.InstallTargets[0])
- return nil
+ config, err := config.FromContext(ctx)
+ if err != nil {
+ ui.Fatalf("Failed to get config: %s", err)
+ }
+ packageResolverConfig := packagemanager.NewDefaultPypiDependencyResolverConfig()
+ packageResolverConfig.IncludeTransitiveDependencies = config.Transitive
+ packageResolverConfig.TransitiveDepth = config.TransitiveDepth
+ packageResolverConfig.IncludeDevDependencies = config.IncludeDevDependencies
+
+ 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)
}
diff --git a/internal/flows/common.go b/internal/flows/common.go
index a8cbaae..dbc82ca 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)
- }
-
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,9 +54,9 @@ 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)
}
diff --git a/packagemanager/dependency_resolver.go b/packagemanager/dependency_resolver.go
index 1f3c7a1..2d48a8f 100644
--- a/packagemanager/dependency_resolver.go
+++ b/packagemanager/dependency_resolver.go
@@ -13,7 +13,7 @@ import (
// Contract for a function that implements ecosystem specific version
// resolver from a version range specification.
-type versionSpecResolver func(version string) string
+type versionSpecResolver func(packageName, version string) string
type dependencyResolverConfig struct {
IncludeDevDependencies bool
@@ -38,7 +38,7 @@ func newDependencyResolver(client packageregistry.Client, config dependencyResol
if versionSpecResolver == nil {
// Default version spec resolver
- versionSpecResolver = func(version string) string {
+ versionSpecResolver = func(packageName, version string) string {
return version
}
}
@@ -141,7 +141,7 @@ func (r *dependencyResolver) resolvePackageDependenciesConcurrent(
Ecosystem: packageVersion.GetPackage().GetEcosystem(),
Name: dependency.Name,
},
- Version: r.versionSpecResolver(dependency.VersionSpec),
+ Version: r.versionSpecResolver(dependency.Name, dependency.VersionSpec),
})
}
diff --git a/packagemanager/npm_resolver.go b/packagemanager/npm_resolver.go
index 68293a4..dd46c2c 100644
--- a/packagemanager/npm_resolver.go
+++ b/packagemanager/npm_resolver.go
@@ -78,7 +78,9 @@ func (r *npmDependencyResolver) ResolveDependencies(ctx context.Context,
TransitiveDepth: r.config.TransitiveDepth,
FailFast: r.config.FailFast,
MaxConcurrency: r.config.MaxConcurrency,
- }, npmCleanVersion)
+ }, func(packageName, version string) string {
+ return npmCleanVersion(version)
+ })
return resolver.resolveDependencies(ctx, packageVersion)
}
diff --git a/packagemanager/pypi.go b/packagemanager/pypi.go
index 615a664..08a79cf 100644
--- a/packagemanager/pypi.go
+++ b/packagemanager/pypi.go
@@ -83,8 +83,6 @@ func (pip *pipPackageManager) ParseCommand(args []string) (*ParsedCommand, error
}
}
- fmt.Printf("Package Name: %s Version: %s\n", packageName, version)
-
installTargets = append(installTargets, &PackageInstallTarget{
PackageVersion: &packagev1.PackageVersion{
Package: &packagev1.Package{
diff --git a/packagemanager/pypi_resolver.go b/packagemanager/pypi_resolver.go
index dd6d22a..5e952bb 100644
--- a/packagemanager/pypi_resolver.go
+++ b/packagemanager/pypi_resolver.go
@@ -2,11 +2,8 @@ package packagemanager
import (
"context"
- "encoding/json"
"fmt"
- "net/http"
"strings"
- "time"
packagev1 "buf.build/gen/go/safedep/api/protocolbuffers/go/safedep/messages/package/v1"
"github.com/Masterminds/semver"
@@ -62,6 +59,13 @@ func (p *pypiDependencyResolver) ResolveDependencies(ctx context.Context, pkg *p
TransitiveDepth: p.config.TransitiveDepth,
FailFast: p.config.FailFast,
MaxConcurrency: p.config.MaxConcurrency,
+ }, 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
})
return resolver.resolveDependencies(ctx, pkg)
@@ -74,6 +78,7 @@ func (p *pypiDependencyResolver) ResolveLatestVersion(ctx context.Context, pkg *
}
pkgInfo, err := pd.GetPackage(pkg.Name)
+ fmt.Println("Package Info version: ", pkgInfo.LatestVersion, " Error: ", err)
if err != nil {
return nil, fmt.Errorf("failed to get package: %w", err)
}
@@ -85,13 +90,6 @@ func (p *pypiDependencyResolver) ResolveLatestVersion(ctx context.Context, pkg *
}, nil
}
-// PyPIPackage represents the package information from PyPI
-type PyPIPackage struct {
- Releases map[string]any `json:"releases"`
-}
-
-var httpClient = &http.Client{Timeout: 10 * time.Second}
-
func pipGetMatchingVersion(packageName, versionConstraint string) (string, error) {
// Already a exact version
if strings.HasPrefix(versionConstraint, "==") {
@@ -103,8 +101,18 @@ func pipGetMatchingVersion(packageName, versionConstraint string) (string, error
versionConstraint = pipConvertCompatibleRelease(versionConstraint)
}
+ 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 := pipFetchPackageVersionsInfo(packageName)
+ pkg, err := pd.GetPackage(packageName)
if err != nil {
return "", err
}
@@ -116,7 +124,7 @@ func pipGetMatchingVersion(packageName, versionConstraint string) (string, error
}
// Get valid versions and find best match
- bestMatch, err := findBestMatchingVersion(pkg.Releases, constraint)
+ bestMatch, err := findBestMatchingVersion(pkg.Versions, constraint)
if err != nil {
return "", fmt.Errorf("no version matches constraint %q: %w", versionConstraint, err)
}
@@ -124,36 +132,15 @@ func pipGetMatchingVersion(packageName, versionConstraint string) (string, error
return bestMatch.Original(), nil
}
-// pipFetchPackageVersionsInfo retrieves package information from PyPI
-func pipFetchPackageVersionsInfo(packageName string) (*PyPIPackage, error) {
- url := fmt.Sprintf("https://pypi.org/pypi/%s/json", packageName)
- resp, err := httpClient.Get(url)
- if err != nil {
- return nil, fmt.Errorf("failed to fetch package info: %w", err)
- }
- defer resp.Body.Close()
-
- if resp.StatusCode != 200 {
- return nil, fmt.Errorf("package not found or HTTP error: %d", resp.StatusCode)
- }
-
- var pypiPkg PyPIPackage
- if err := json.NewDecoder(resp.Body).Decode(&pypiPkg); err != nil {
- return nil, fmt.Errorf("failed to parse JSON: %w", err)
- }
-
- return &pypiPkg, nil
-}
-
-func findBestMatchingVersion(releases map[string]any, constraint *semver.Constraints) (*semver.Version, error) {
+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)
+ for _, v := range releases {
+ ver, err := semver.NewVersion(v.Version)
if err != nil {
continue // Skip invalid versions
}
+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")
)
+5
View File
@@ -13,6 +13,11 @@ type Command struct {
type PackageInstallTarget struct { type PackageInstallTarget struct {
PackageVersion *packagev1.PackageVersion 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 { func (pit *PackageInstallTarget) HasVersion() bool {
+31 -16
View File
@@ -64,13 +64,13 @@ func (pip *pipPackageManager) ParseCommand(args []string) (*ParsedCommand, error
var installTargets []*PackageInstallTarget var installTargets []*PackageInstallTarget
for _, pkg := range packages { for _, pkg := range packages {
packageName, version, err := pipParsePackageInfo(pkg) packageName, version, extras, err := pipParsePackageInfo(pkg)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to parse package info: %w", err) return nil, fmt.Errorf("failed to parse package info: %w", err)
} }
// If exact version provided just trim it. If not get a version that satisfies a given version specifier
if version != "" { if version != "" {
// If exact version provided just trim it. If not get a version that satisfies a given version specifier
if strings.HasPrefix(version, "==") { if strings.HasPrefix(version, "==") {
// Exact version, just trim // Exact version, just trim
version = strings.TrimPrefix(version, "==") version = strings.TrimPrefix(version, "==")
@@ -91,6 +91,7 @@ func (pip *pipPackageManager) ParseCommand(args []string) (*ParsedCommand, error
}, },
Version: version, Version: version,
}, },
Extras: extras,
}) })
} }
@@ -100,20 +101,38 @@ func (pip *pipPackageManager) ParseCommand(args []string) (*ParsedCommand, error
}, nil }, nil
} }
// pipParsePackageInfo parses python package strings like: // pipParsePackageInfo parses a pip install package specification, separating the package name,
// "fastapi", "fastapi==0.115.7", "requests>=2.0,<3.0", "pydantic!=1.8,!=1.8.1" // version constraints, and any extras (additional features) to be installed.
// Returns packageName and version (empty if none specified). // Example: "django[mysql,redis]>=3.0" returns ("django", ">=3.0", ["mysql", "redis"], nil)
func pipParsePackageInfo(input string) (packageName, version string, err error) { func pipParsePackageInfo(input string) (packageName, version string, extras []string, err error) {
if input == "" { if input == "" {
return "", "", fmt.Errorf("package info cannot be empty") return "", "", nil, fmt.Errorf("package info cannot be empty")
} }
input = strings.TrimSpace(input) 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: // Python package version specifiers are typically separated by one of:
// '==', '>=', '<=', '!=', '>', '<', '~=', or direct comma separated list // '==', '>=', '<=', '!=', '>', '<', '~=', or direct comma separated list
// We'll find the first occurrence of these operators for splitting.
operators := []string{"==", ">=", "<=", "!=", ">", "<", "~="} operators := []string{"==", ">=", "<=", "!=", ">", "<", "~="}
index := -1 index := -1
@@ -127,21 +146,17 @@ func pipParsePackageInfo(input string) (packageName, version string, err error)
if index == -1 { if index == -1 {
// No operator found, whole input is package name, no version // No operator found, whole input is package name, no version
return input, "", nil return strings.TrimSpace(input), "", extras, nil
} }
packageName = strings.TrimSpace(input[:index]) packageName = strings.TrimSpace(input[:index])
version = strings.TrimSpace(input[index:]) version = strings.TrimSpace(input[index:])
// Some version specs can have multiple constraints separated by commas
// Example: "requests>=2.0,<3.0"
// So keep version as is
if packageName == "" { if packageName == "" {
return "", "", fmt.Errorf("invalid package name in input '%s'", input) return "", "", nil, fmt.Errorf("invalid package name in input '%s'", input)
} }
return packageName, version, nil return packageName, version, extras, nil
} }
func pipConvertCompatibleRelease(version string) string { func pipConvertCompatibleRelease(version string) string {
+111
View File
@@ -2,7 +2,10 @@ package packagemanager
import ( import (
"context" "context"
"encoding/json"
"fmt" "fmt"
"net/http"
"regexp"
"strings" "strings"
packagev1 "buf.build/gen/go/safedep/api/protocolbuffers/go/safedep/messages/package/v1" packagev1 "buf.build/gen/go/safedep/api/protocolbuffers/go/safedep/messages/package/v1"
@@ -90,6 +93,114 @@ func (p *pypiDependencyResolver) ResolveLatestVersion(ctx context.Context, pkg *
}, nil }, 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 GetPackageDependencies(packageName, version string) ([]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
}
pkgDeps := make([]PyPIDependencySpec, len(pypipkg.Info.RequiresDist))
for _, dep := range pypipkg.Info.RequiresDist {
name, version, extra := pypiParseDependency(dep)
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) {
var name string
var version string
// Split line by ';' to separate version and markers
parts := strings.SplitN(input, ";", 2)
mainPart := strings.TrimSpace(parts[0])
// Find last occurrence of version operators
operators := []string{"==", ">=", "<=", "!=", ">", "<", "~="}
versionIndex := -1
for _, op := range operators {
if idx := strings.LastIndex(mainPart, op); idx != -1 {
if idx > versionIndex {
versionIndex = idx
}
}
}
if versionIndex != -1 {
name = strings.TrimSpace(mainPart[:versionIndex])
version = strings.TrimSpace(mainPart[versionIndex:])
} else {
name = mainPart
}
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) { func pipGetMatchingVersion(packageName, versionConstraint string) (string, error) {
// Already a exact version // Already a exact version
if strings.HasPrefix(versionConstraint, "==") { if strings.HasPrefix(versionConstraint, "==") {
+16 -5
View File
@@ -12,6 +12,7 @@ func TestPipParsePackageInfo(t *testing.T) {
input string input string
pkgName string pkgName string
version string version string
extras []string
wantErr bool wantErr bool
}{ }{
{ {
@@ -19,13 +20,15 @@ func TestPipParsePackageInfo(t *testing.T) {
input: "fastapi", input: "fastapi",
pkgName: "fastapi", pkgName: "fastapi",
version: "", version: "",
extras: nil,
wantErr: false, wantErr: false,
}, },
{ {
name: "package with exact version", name: "package with exact version & extra",
input: "fastapi==0.115.7", input: "fastapi[all]==0.115.7",
pkgName: "fastapi", pkgName: "fastapi",
version: "==0.115.7", version: "==0.115.7",
extras: []string{"all"},
wantErr: false, wantErr: false,
}, },
{ {
@@ -33,6 +36,7 @@ func TestPipParsePackageInfo(t *testing.T) {
input: "requests>=2.0,<3.0", input: "requests>=2.0,<3.0",
pkgName: "requests", pkgName: "requests",
version: ">=2.0,<3.0", version: ">=2.0,<3.0",
extras: nil,
wantErr: false, wantErr: false,
}, },
{ {
@@ -47,13 +51,15 @@ func TestPipParsePackageInfo(t *testing.T) {
input: "django~=3.1.0", input: "django~=3.1.0",
pkgName: "django", pkgName: "django",
version: "~=3.1.0", version: "~=3.1.0",
extras: nil,
wantErr: false, wantErr: false,
}, },
{ {
name: "package with greater than", name: "package with greater than with empty extra",
input: "numpy>1.20.0", input: "numpy[]>1.20.0",
pkgName: "numpy", pkgName: "numpy",
version: ">1.20.0", version: ">1.20.0",
extras: nil,
wantErr: false, wantErr: false,
}, },
{ {
@@ -61,6 +67,7 @@ func TestPipParsePackageInfo(t *testing.T) {
input: "pandas<2.0.0", input: "pandas<2.0.0",
pkgName: "pandas", pkgName: "pandas",
version: "<2.0.0", version: "<2.0.0",
extras: nil,
wantErr: false, wantErr: false,
}, },
{ {
@@ -68,6 +75,7 @@ func TestPipParsePackageInfo(t *testing.T) {
input: "", input: "",
pkgName: "", pkgName: "",
version: "", version: "",
extras: nil,
wantErr: true, wantErr: true,
}, },
{ {
@@ -75,6 +83,7 @@ func TestPipParsePackageInfo(t *testing.T) {
input: "==1.0.0", input: "==1.0.0",
pkgName: "", pkgName: "",
version: "", version: "",
extras: nil,
wantErr: true, wantErr: true,
}, },
{ {
@@ -82,19 +91,21 @@ func TestPipParsePackageInfo(t *testing.T) {
input: " requests == 2.0.0 ", input: " requests == 2.0.0 ",
pkgName: "requests", pkgName: "requests",
version: "== 2.0.0", version: "== 2.0.0",
extras: nil,
wantErr: false, wantErr: false,
}, },
} }
for _, tc := range cases { for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
pkgName, version, err := pipParsePackageInfo(tc.input) pkgName, version, extras, err := pipParsePackageInfo(tc.input)
if tc.wantErr { if tc.wantErr {
assert.Error(t, err) assert.Error(t, err)
} else { } else {
assert.NoError(t, err) assert.NoError(t, err)
assert.Equal(t, tc.pkgName, pkgName) assert.Equal(t, tc.pkgName, pkgName)
assert.Equal(t, tc.version, version) assert.Equal(t, tc.version, version)
assert.Equal(t, tc.extras, extras)
} }
}) })
} }