From 88f39b56ff48797111057e8ac53177e95f3158da Mon Sep 17 00:00:00 2001 From: Sahil Bansal Date: Sun, 4 May 2025 18:50:13 +0530 Subject: [PATCH] fix: parsePackageInfo to handle pkg names with special character (#11) * fix: parsePackageInfo to handle pkg names with special character * test: add tests for CleanVersion and ParsePackageInfo util functions --- internal/ui/progress.go | 6 ++ pkg/common/utils/utils.go | 12 +++ pkg/common/utils/utils_test.go | 170 +++++++++++++++++++++++++++++++++ pkg/wrapper/npm_base.go | 4 +- 4 files changed, 190 insertions(+), 2 deletions(-) create mode 100644 pkg/common/utils/utils_test.go diff --git a/internal/ui/progress.go b/internal/ui/progress.go index ab652b9..432fa15 100644 --- a/internal/ui/progress.go +++ b/internal/ui/progress.go @@ -119,6 +119,12 @@ func IncrementTrackerTotal(i any, count int64) { } } +func SetTrackerTotal(i any, count int64) { + if tracker, ok := i.(ProgressTracker); ok { + tracker.UpdateTotal(count) + } +} + func IncrementProgress(i any, count int64) { if tracker, ok := i.(ProgressTracker); ok && (progressTrackerDelta(tracker) > count) { tracker.Increment(count) diff --git a/pkg/common/utils/utils.go b/pkg/common/utils/utils.go index 51e13e3..d9bdf09 100644 --- a/pkg/common/utils/utils.go +++ b/pkg/common/utils/utils.go @@ -18,6 +18,18 @@ func ParsePackageInfo(input string) (packageName, version string, err error) { if input == "" { return "", "", fmt.Errorf("package info cannot be empty") } + input = strings.TrimSpace(input) + + if strings.HasPrefix(input, "@") { + lastAtIndex := strings.LastIndex(input, "@") + if lastAtIndex > 0 { + packageName = strings.TrimSpace(input[:lastAtIndex]) + version = strings.TrimSpace(input[lastAtIndex+1:]) + return packageName, version, nil + } + // If no version specifier, return the whole input as package name + return strings.TrimSpace(input), "", nil + } pkg := strings.Split(input, "@") if len(pkg) == 2 { diff --git a/pkg/common/utils/utils_test.go b/pkg/common/utils/utils_test.go new file mode 100644 index 0000000..bb82cf7 --- /dev/null +++ b/pkg/common/utils/utils_test.go @@ -0,0 +1,170 @@ +package utils + +import ( + "strings" + "testing" +) + +func TestCleanVersion(t *testing.T) { + tests := []struct { + name string + version string + expected string + }{ + { + name: "caret version", + version: "^1.2.3", + expected: "1.2.3", + }, + { + name: "tilde version", + version: "~1.2.3", + expected: "1.2.3", + }, + { + name: "exact version", + version: "1.2.3", + expected: "1.2.3", + }, + { + name: "asterisk version", + version: "*", + expected: "latest", + }, + { + name: "empty version", + version: "", + expected: "", + }, + { + name: "both caret and tilde", + version: "^~1.2.3", + expected: "1.2.3", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := CleanVersion(tt.version) + if result != tt.expected { + t.Errorf("CleanVersion(%q) = %q, want %q", tt.version, result, tt.expected) + } + }) + } +} + +func TestParsePackageInfo(t *testing.T) { + tests := []struct { + name string + input string + wantPackage string + wantVersion string + wantErr bool + errorContains string + }{ + { + name: "simple package", + input: "express", + wantPackage: "express", + wantVersion: "", + wantErr: false, + }, + { + name: "package with version", + input: "express@4.17.1", + wantPackage: "express", + wantVersion: "4.17.1", + wantErr: false, + }, + { + name: "scoped package", + input: "@angular/core", + wantPackage: "@angular/core", + wantVersion: "", + wantErr: false, + }, + { + name: "scoped package with version", + input: "@angular/core@12.0.0", + wantPackage: "@angular/core", + wantVersion: "12.0.0", + wantErr: false, + }, + { + name: "package with caret version", + input: "react@^17.0.2", + wantPackage: "react", + wantVersion: "^17.0.2", + wantErr: false, + }, + { + name: "package with tilde version", + input: "lodash@~4.17.21", + wantPackage: "lodash", + wantVersion: "~4.17.21", + wantErr: false, + }, + { + name: "empty input", + input: "", + wantPackage: "", + wantVersion: "", + wantErr: true, + errorContains: "package info cannot be empty", + }, + { + name: "invalid format with multiple @", + input: "pkg@1.0.0@2.0.0", + wantPackage: "", + wantVersion: "", + wantErr: true, + errorContains: "invalid format", + }, + { + name: "package with spaces", + input: " express@4.17.1 ", + wantPackage: "express", + wantVersion: "4.17.1", + wantErr: false, + }, + { + name: "scoped package with spaces", + input: " @types/node@14.14.31 ", + wantPackage: "@types/node", + wantVersion: "14.14.31", + wantErr: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + packageName, version, err := ParsePackageInfo(tt.input) + + // Check error + if tt.wantErr { + if err == nil { + t.Errorf("ParsePackageInfo(%q) expected error, got nil", tt.input) + return + } + if tt.errorContains != "" && !strings.Contains(err.Error(), tt.errorContains) { + t.Errorf("ParsePackageInfo(%q) error = %v, want error containing %q", tt.input, err, tt.errorContains) + } + return + } + if err != nil { + t.Errorf("ParsePackageInfo(%q) unexpected error: %v", tt.input, err) + return + } + + // Check package name + if packageName != tt.wantPackage { + t.Errorf("ParsePackageInfo(%q) package = %q, want %q", tt.input, packageName, tt.wantPackage) + } + + // Check version + if version != tt.wantVersion { + t.Errorf("ParsePackageInfo(%q) version = %q, want %q", tt.input, version, tt.wantVersion) + } + }) + } +} diff --git a/pkg/wrapper/npm_base.go b/pkg/wrapper/npm_base.go index db96919..04e4efa 100644 --- a/pkg/wrapper/npm_base.go +++ b/pkg/wrapper/npm_base.go @@ -29,7 +29,7 @@ func NewPackageManagerWrapper(registryType registry.RegistryType) *PackageManage func (pmw *PackageManagerWrapper) Wrap() error { ui.StartProgressWriter() - progressTracker := ui.TrackProgress(fmt.Sprintf("Scanning %s ", pmw.PackageName), 1) + progressTracker := ui.TrackProgress(fmt.Sprintf("Scanning %s ", pmw.PackageName), 5) if pmw.PackageName == "" { return fmt.Errorf("package name cannot be empty") } @@ -94,7 +94,7 @@ func (pmw *PackageManagerWrapper) getDependencies(ctx context.Context, fetcher r if err != nil { return nil, err } - ui.IncrementTrackerTotal(progressTracker, int64(len(deps))) + ui.SetTrackerTotal(progressTracker, int64(len(deps))) return deps, nil }