diff --git a/pkg/registry/npm_fetcher.go b/pkg/registry/npm_fetcher.go index fa7ef46..dfd0823 100644 --- a/pkg/registry/npm_fetcher.go +++ b/pkg/registry/npm_fetcher.go @@ -5,9 +5,11 @@ import ( "encoding/json" "fmt" "sync" + "sync/atomic" "time" "github.com/safedep/dry/log" + "github.com/safedep/pmg/internal/ui" "github.com/safedep/pmg/pkg/common/utils" "github.com/safedep/pmg/pkg/models" ) @@ -15,6 +17,22 @@ import ( // NpmFetcher fetches dependencies from NPM registry type NpmFetcher struct { *BaseFetcher + progressTracker ui.ProgressTracker + mu sync.Mutex + fetchedDeps int32 +} + +func (nf *NpmFetcher) SetProgressTracker(tracker ui.ProgressTracker) { + nf.progressTracker = tracker + atomic.StoreInt32(&nf.fetchedDeps, 0) +} + +func (nf *NpmFetcher) incrementProgress() { + if nf.progressTracker != nil { + atomic.AddInt32(&nf.fetchedDeps, 1) + // Update progress message to show number of packages fetched + ui.SetPinnedMessageOnProgressWriter(fmt.Sprintf("Fetched %d packages", atomic.LoadInt32(&nf.fetchedDeps))) + } } // NewNpmFetcher creates a new NPM registry fetcher @@ -92,6 +110,8 @@ func (nf *NpmFetcher) fetchDependenciesConcurrent(ctx context.Context, pkg model return nil, fmt.Errorf("failed to fetch package info for %s: %w", pkg.Name, err) } + nf.incrementProgress() + dependencies := packageInfo.Dependencies node := &models.DependencyNode{ Name: pkg.Name, diff --git a/pkg/wrapper/npm_base.go b/pkg/wrapper/npm_base.go index 04e4efa..5408ca7 100644 --- a/pkg/wrapper/npm_base.go +++ b/pkg/wrapper/npm_base.go @@ -65,11 +65,18 @@ func (pmw *PackageManagerWrapper) scanAndInstall(ctx context.Context, progressTr pmw.PackageName = fmt.Sprintf("%s@%s", name, version) } - deps, err := pmw.getDependencies(ctx, fetcher, name, version, progressTracker) + // Get dependencies with progress tracking + npmFetcher := fetcher.(*registry.NpmFetcher) + ui.IncrementTrackerTotal(progressTracker, 0) + npmFetcher.SetProgressTracker(progressTracker) + + deps, err := npmFetcher.GetFlattenedDependencies(ctx, name, version) if err != nil { return err } + // We know the total deps, set progress for analysis phase + ui.IncrementTrackerTotal(progressTracker, int64(len(deps))) if err := pmw.analyzeDependencies(ctx, deps, progressTracker); err != nil { return err } @@ -88,9 +95,9 @@ func (pmw *PackageManagerWrapper) resolveLatestVersion(ctx context.Context, fetc } func (pmw *PackageManagerWrapper) getDependencies(ctx context.Context, fetcher registry.Fetcher, name, version string, progressTracker ui.ProgressTracker) ([]string, error) { - ui.IncrementProgress(progressTracker, 1) - deps, err := fetcher.GetFlattenedDependencies(ctx, name, version) - ui.IncrementProgress(progressTracker, 1) + npmFetcher := fetcher.(*registry.NpmFetcher) + npmFetcher.SetProgressTracker(progressTracker) + deps, err := npmFetcher.GetFlattenedDependencies(ctx, name, version) if err != nil { return nil, err }