refactor: improve progress bar logic and update display (#14)

* fix: Revise progress tracking mechanism

* chore: remove unused getDependencies func

* chore: removed unused property & add fetcher check

* Update pkg/wrapper/npm_base.go

Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com>
Signed-off-by: Sahil Bansal <bansalsahil315@gmail.com>

* refactor: make SetProgressTracker common for all fetchers

---------

Signed-off-by: Sahil Bansal <bansalsahil315@gmail.com>
Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com>
This commit is contained in:
Sahil Bansal
2025-05-06 16:35:59 +05:30
committed by GitHub
co-authored by Copilot
parent df754ccc82
commit 6a28fb16a1
3 changed files with 31 additions and 15 deletions
+12 -3
View File
@@ -6,7 +6,9 @@ import (
"context" "context"
"fmt" "fmt"
"sync" "sync"
"sync/atomic"
"github.com/safedep/pmg/internal/ui"
"github.com/safedep/pmg/pkg/models" "github.com/safedep/pmg/pkg/models"
) )
@@ -21,9 +23,11 @@ type Fetcher interface {
// BaseFetcher implements common functionality for all registry fetchers // BaseFetcher implements common functionality for all registry fetchers
type BaseFetcher struct { type BaseFetcher struct {
visitedMu sync.RWMutex visitedMu sync.RWMutex
visited map[string]bool visited map[string]bool
client RegistryClient client RegistryClient
progressTracker ui.ProgressTracker
fetchedDeps int32
} }
// NewBaseFetcher creates a new BaseFetcher with the specified registry client // NewBaseFetcher creates a new BaseFetcher with the specified registry client
@@ -34,6 +38,11 @@ func NewBaseFetcher(client RegistryClient) *BaseFetcher {
} }
} }
func (bf *BaseFetcher) SetProgressTracker(tracker ui.ProgressTracker) {
bf.progressTracker = tracker
atomic.StoreInt32(&bf.fetchedDeps, 0)
}
// isVisited checks if a package has already been visited // isVisited checks if a package has already been visited
func (bf *BaseFetcher) isVisited(key string) bool { func (bf *BaseFetcher) isVisited(key string) bool {
bf.visitedMu.RLock() bf.visitedMu.RLock()
+12
View File
@@ -5,9 +5,11 @@ import (
"encoding/json" "encoding/json"
"fmt" "fmt"
"sync" "sync"
"sync/atomic"
"time" "time"
"github.com/safedep/dry/log" "github.com/safedep/dry/log"
"github.com/safedep/pmg/internal/ui"
"github.com/safedep/pmg/pkg/common/utils" "github.com/safedep/pmg/pkg/common/utils"
"github.com/safedep/pmg/pkg/models" "github.com/safedep/pmg/pkg/models"
) )
@@ -17,6 +19,14 @@ type NpmFetcher struct {
*BaseFetcher *BaseFetcher
} }
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 // NewNpmFetcher creates a new NPM registry fetcher
func NewNpmFetcher(timeout time.Duration) *NpmFetcher { func NewNpmFetcher(timeout time.Duration) *NpmFetcher {
client := NewHttpRegistryClient( client := NewHttpRegistryClient(
@@ -92,6 +102,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) return nil, fmt.Errorf("failed to fetch package info for %s: %w", pkg.Name, err)
} }
nf.incrementProgress()
dependencies := packageInfo.Dependencies dependencies := packageInfo.Dependencies
node := &models.DependencyNode{ node := &models.DependencyNode{
Name: pkg.Name, Name: pkg.Name,
+7 -12
View File
@@ -66,11 +66,17 @@ func (pmw *PackageManagerWrapper) scanAndInstall(ctx context.Context, progressTr
pmw.PackageName = fmt.Sprintf("%s@%s", name, version) 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)
npmFetcher.SetProgressTracker(progressTracker)
deps, err := npmFetcher.GetFlattenedDependencies(ctx, name, version)
if err != nil { if err != nil {
return err 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 { if err := pmw.analyzeDependencies(ctx, deps, progressTracker); err != nil {
return err return err
} }
@@ -88,17 +94,6 @@ func (pmw *PackageManagerWrapper) resolveLatestVersion(ctx context.Context, fetc
return version, nil return version, nil
} }
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)
if err != nil {
return nil, err
}
ui.SetTrackerTotal(progressTracker, int64(len(deps)))
return deps, nil
}
func (pmw *PackageManagerWrapper) analyzeDependencies(ctx context.Context, deps []string, progressTracker ui.ProgressTracker) error { func (pmw *PackageManagerWrapper) analyzeDependencies(ctx context.Context, deps []string, progressTracker ui.ProgressTracker) error {
client, err := analyser.GetMalwareAnalysisClient() client, err := analyser.GetMalwareAnalysisClient()
if err != nil { if err != nil {