From 216657afef61a6107a41b9e38961ce2b5811e9e8 Mon Sep 17 00:00:00 2001 From: Sahilb315 Date: Thu, 22 May 2025 02:05:34 +0530 Subject: [PATCH] feat: add Python dependency parsing with extras support --- diff.txt | 362 ------------------------------- packagemanager/errors.go | 15 ++ packagemanager/packagemanager.go | 5 + packagemanager/pypi.go | 47 ++-- packagemanager/pypi_resolver.go | 111 ++++++++++ packagemanager/pypi_test.go | 21 +- 6 files changed, 178 insertions(+), 383 deletions(-) delete mode 100644 diff.txt create mode 100644 packagemanager/errors.go diff --git a/diff.txt b/diff.txt deleted file mode 100644 index c508c7e..0000000 --- a/diff.txt +++ /dev/null @@ -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 - } 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/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 index 08a79cf..71de10f 100644 --- a/packagemanager/pypi.go +++ b/packagemanager/pypi.go @@ -64,13 +64,13 @@ func (pip *pipPackageManager) ParseCommand(args []string) (*ParsedCommand, error var installTargets []*PackageInstallTarget for _, pkg := range packages { - packageName, version, err := pipParsePackageInfo(pkg) + packageName, version, extras, err := pipParsePackageInfo(pkg) if err != nil { 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 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, "==") @@ -91,6 +91,7 @@ func (pip *pipPackageManager) ParseCommand(args []string) (*ParsedCommand, error }, Version: version, }, + Extras: extras, }) } @@ -100,20 +101,38 @@ func (pip *pipPackageManager) ParseCommand(args []string) (*ParsedCommand, error }, nil } -// pipParsePackageInfo parses python package strings like: -// "fastapi", "fastapi==0.115.7", "requests>=2.0,<3.0", "pydantic!=1.8,!=1.8.1" -// Returns packageName and version (empty if none specified). -func pipParsePackageInfo(input string) (packageName, version string, err error) { +// 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 "", "", fmt.Errorf("package info cannot be empty") + 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 - // We'll find the first occurrence of these operators for splitting. - operators := []string{"==", ">=", "<=", "!=", ">", "<", "~="} index := -1 @@ -127,21 +146,17 @@ func pipParsePackageInfo(input string) (packageName, version string, err error) if index == -1 { // No operator found, whole input is package name, no version - return input, "", nil + return strings.TrimSpace(input), "", extras, nil } packageName = 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 == "" { - 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 { diff --git a/packagemanager/pypi_resolver.go b/packagemanager/pypi_resolver.go index 5e952bb..eded678 100644 --- a/packagemanager/pypi_resolver.go +++ b/packagemanager/pypi_resolver.go @@ -2,7 +2,10 @@ package packagemanager import ( "context" + "encoding/json" "fmt" + "net/http" + "regexp" "strings" 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 } +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) { // Already a exact version if strings.HasPrefix(versionConstraint, "==") { diff --git a/packagemanager/pypi_test.go b/packagemanager/pypi_test.go index 6b68c8f..928a33e 100644 --- a/packagemanager/pypi_test.go +++ b/packagemanager/pypi_test.go @@ -12,6 +12,7 @@ func TestPipParsePackageInfo(t *testing.T) { input string pkgName string version string + extras []string wantErr bool }{ { @@ -19,13 +20,15 @@ func TestPipParsePackageInfo(t *testing.T) { input: "fastapi", pkgName: "fastapi", version: "", + extras: nil, wantErr: false, }, { - name: "package with exact version", - input: "fastapi==0.115.7", + name: "package with exact version & extra", + input: "fastapi[all]==0.115.7", pkgName: "fastapi", version: "==0.115.7", + extras: []string{"all"}, wantErr: false, }, { @@ -33,6 +36,7 @@ func TestPipParsePackageInfo(t *testing.T) { input: "requests>=2.0,<3.0", pkgName: "requests", version: ">=2.0,<3.0", + extras: nil, wantErr: false, }, { @@ -47,13 +51,15 @@ func TestPipParsePackageInfo(t *testing.T) { input: "django~=3.1.0", pkgName: "django", version: "~=3.1.0", + extras: nil, wantErr: false, }, { - name: "package with greater than", - input: "numpy>1.20.0", + name: "package with greater than with empty extra", + input: "numpy[]>1.20.0", pkgName: "numpy", version: ">1.20.0", + extras: nil, wantErr: false, }, { @@ -61,6 +67,7 @@ func TestPipParsePackageInfo(t *testing.T) { input: "pandas<2.0.0", pkgName: "pandas", version: "<2.0.0", + extras: nil, wantErr: false, }, { @@ -68,6 +75,7 @@ func TestPipParsePackageInfo(t *testing.T) { input: "", pkgName: "", version: "", + extras: nil, wantErr: true, }, { @@ -75,6 +83,7 @@ func TestPipParsePackageInfo(t *testing.T) { input: "==1.0.0", pkgName: "", version: "", + extras: nil, wantErr: true, }, { @@ -82,19 +91,21 @@ func TestPipParsePackageInfo(t *testing.T) { 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, err := pipParsePackageInfo(tc.input) + 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) } }) }