diff --git a/cmd/ecosystems/npm.go b/cmd/ecosystems/npm.go index 3736e60..d5665cd 100644 --- a/cmd/ecosystems/npm.go +++ b/cmd/ecosystems/npm.go @@ -4,13 +4,10 @@ import ( "context" _ "embed" "fmt" - "os" - "strings" "time" "github.com/safedep/dry/log" "github.com/safedep/pmg/pkg/analyser" - "github.com/safedep/pmg/pkg/common" "github.com/safedep/pmg/pkg/common/utils" "github.com/safedep/pmg/pkg/models" vetUtils "github.com/safedep/vet/pkg/common/utils" @@ -22,9 +19,6 @@ var ( action string ) -//go:embed tree/arborist-bundle.js -var arboristJs string - func NewNpmCommand() *cobra.Command { cmd := &cobra.Command{ Use: "npm [action] [package]", @@ -65,29 +59,10 @@ func wrapNpm() error { ctx, cancel := context.WithTimeout(context.Background(), 3*time.Minute) defer cancel() - // Extract package information - outputFile, err := common.RunPkgExtractor(common.ExtractorOptions{ - PackageName: packageName, - ScriptContent: arboristJs, - Interpreter: "node", - ScriptType: "js", - Args: []string{}, - Env: map[string]string{ - "NPM_AUTH_TOKEN": utils.NpmAuthToken(), - }, - }) - + deps, err := analyser.FetchDependencies(packageName) if err != nil { - return fmt.Errorf("failed to extract package info: %w", err) + return err } - // Clean up the temporary file when done - defer os.Remove(outputFile) - - data, err := os.ReadFile(outputFile) - if err != nil { - return fmt.Errorf("error while reading package output file: %w", err) - } - maliciousPkgs := make(map[string]string) client, err := analyser.GetMalwareAnalysisClient() @@ -103,22 +78,12 @@ func wrapNpm() error { defer queue.Stop() // Add packages to the queue - lines := strings.SplitSeq(string(data), "\n") - for line := range lines { - line = strings.TrimSpace(line) - if line == "" || !strings.Contains(line, "@") { + for _, dep := range deps { + name, version, err := utils.ParsePackageInfo(dep) + if err != nil { + log.Errorf("Error while parsing info of package %s", name) continue } - - idx := strings.LastIndex(line, "@") - if idx <= 0 { - log.Debugf("Invalid package line: %s", line) - continue - } - - name := line[:idx] - version := line[idx+1:] - queue.Add(models.PackageAnalysisItem{ Name: name, Version: version, diff --git a/pkg/analyser/analysis.go b/pkg/analyser/analysis.go index 14be470..1ea632d 100644 --- a/pkg/analyser/analysis.go +++ b/pkg/analyser/analysis.go @@ -10,10 +10,27 @@ import ( packagev1 "buf.build/gen/go/safedep/api/protocolbuffers/go/safedep/messages/package/v1" malysisv1 "buf.build/gen/go/safedep/api/protocolbuffers/go/safedep/services/malysis/v1" "github.com/safedep/dry/log" + "github.com/safedep/pmg/pkg/common" + "github.com/safedep/pmg/pkg/common/utils" "github.com/safedep/pmg/pkg/models" vetUtils "github.com/safedep/vet/pkg/common/utils" ) +func FetchDependencies(packageName string) ([]string, error) { + fetcher := models.NewFlatteningFetcher() + name, version, err := utils.ParsePackageInfo(packageName) + if err != nil { + return nil, err + } + + deps, err := fetcher.GetPackageDependencies(name, version) + flatDeps := common.FlattenDependencyTree(deps) + if err != nil { + return nil, fmt.Errorf("error while fetching package dependencies: %v\n", err) + } + return flatDeps, nil +} + func AnalysePackage(maliciousPkgs map[string]string, client malysisv1grpc.MalwareAnalysisServiceClient, ctx context.Context) vetUtils.WorkQueueFn[models.PackageAnalysisItem] { return func(q *vetUtils.WorkQueue[models.PackageAnalysisItem], item models.PackageAnalysisItem) error { var maliciousPkgsMutex sync.Mutex diff --git a/pkg/common/extractor.go b/pkg/common/extractor.go index c085efe..f24e298 100644 --- a/pkg/common/extractor.go +++ b/pkg/common/extractor.go @@ -7,6 +7,7 @@ import ( "github.com/safedep/dry/crypto" "github.com/safedep/pmg/pkg/common/utils" + "github.com/safedep/pmg/pkg/models" ) // ExtractorOptions holds configuration for running an extractor script @@ -19,6 +20,29 @@ type ExtractorOptions struct { Env map[string]string // Environment variables to pass to the script } +func FlattenDependencyTree(node *models.DependencyNode) []string { + result := make([]string, 0) + seen := make(map[string]bool) + + var flatten func(*models.DependencyNode) + flatten = func(n *models.DependencyNode) { + key := fmt.Sprintf("%s@%s", n.Name, n.Version) + if seen[key] { + return + } + seen[key] = true + + result = append(result, fmt.Sprintf("%s@%s", n.Name, n.Version)) + + for _, dep := range n.Dependencies { + flatten(dep) + } + } + + flatten(node) + return result +} + // RunExtractor extracts an embedded script to a temp file and executes it func RunPkgExtractor(opts ExtractorOptions) (string, error) { interpreterPath, err := utils.GetExecutablePath(opts.Interpreter) diff --git a/pkg/common/utils/utils.go b/pkg/common/utils/utils.go new file mode 100644 index 0000000..7ac0bab --- /dev/null +++ b/pkg/common/utils/utils.go @@ -0,0 +1,38 @@ +package utils + +import ( + "fmt" + "strings" +) + +func CleanVersion(version string) string { + version = strings.TrimPrefix(version, "^") + version = strings.TrimPrefix(version, "~") + if version == "*" { + return "latest" + } + return version +} + +func ParsePackageInfo(input string) (packageName, version string, err error) { + if input == "" { + return "", "", fmt.Errorf("package info cannot be empty") + } + + pkg := strings.Split(input, "@") + if len(pkg) != 2 { + return "", "", fmt.Errorf("invalid format: expected 'package@version', got '%s'", input) + } + + packageName = strings.TrimSpace(pkg[0]) + version = strings.TrimSpace(pkg[1]) + + if packageName == "" { + return "", "", fmt.Errorf("package name cannot be empty") + } + if version == "" { + return "", "", fmt.Errorf("version cannot be empty") + } + + return packageName, version, nil +} diff --git a/pkg/models/models.go b/pkg/models/models.go index d64244c..39ddb96 100644 --- a/pkg/models/models.go +++ b/pkg/models/models.go @@ -1,6 +1,16 @@ package models -import "fmt" +import ( + "encoding/json" + "fmt" + "io" + "net/http" + "sync" + "time" + + "github.com/safedep/dry/log" + "github.com/safedep/pmg/pkg/common/utils" +) type PackageAnalysisItem struct { Name string @@ -10,3 +20,136 @@ type PackageAnalysisItem struct { func (p PackageAnalysisItem) Id() string { return fmt.Sprintf("%s@%s", p.Name, p.Version) } + +type PackageInfo struct { + Name string `json:"name"` + Version string `json:"version"` + Dependencies map[string]string `json:"dependencies"` +} + +type DependencyNode struct { + Name string + Version string + Dependencies map[string]*DependencyNode +} + +// FlatteningFetcher fetches package dependencies and builds a dependency tree +type FlatteningFetcher struct { + visitedMu sync.RWMutex + visited map[string]bool + client *http.Client +} + +func NewFlatteningFetcher() *FlatteningFetcher { + return &FlatteningFetcher{ + visited: make(map[string]bool), + client: &http.Client{Timeout: 10 * time.Second}, + } +} + +func (ff *FlatteningFetcher) GetPackageDependencies(packageName, version string) (*DependencyNode, error) { + return ff.fetchDependenciesConcurrent(packageName, version) +} + +func (ff *FlatteningFetcher) isVisited(key string) bool { + ff.visitedMu.RLock() + defer ff.visitedMu.RUnlock() + return ff.visited[key] +} + +func (ff *FlatteningFetcher) markVisited(key string) { + ff.visitedMu.Lock() + defer ff.visitedMu.Unlock() + ff.visited[key] = true +} + +func (ff *FlatteningFetcher) fetchDependenciesConcurrent(packageName, version string) (*DependencyNode, error) { + cacheKey := packageName + "@" + version + if ff.isVisited(cacheKey) { + return &DependencyNode{ + Name: packageName, + Version: version, + }, nil + } + ff.markVisited(cacheKey) + + packageInfo, err := ff.fetchPackageInfo(packageName, version) + if err != nil { + return nil, fmt.Errorf("failed to fetch package info for %s: %v", packageName, err) + } + + dependencies := packageInfo.Dependencies + + node := &DependencyNode{ + Name: packageName, + Version: version, + Dependencies: make(map[string]*DependencyNode), + } + + if len(dependencies) == 0 { + return node, nil + } + + // Create channels for concurrent processing + type result struct { + name string + node *DependencyNode + err error + } + + var wg sync.WaitGroup + resultChan := make(chan result, len(dependencies)) + + // Process dependencies concurrently + for depName, depVersion := range dependencies { + wg.Add(1) + go func(name, version string) { + defer wg.Done() + version = utils.CleanVersion(version) + depNode, err := ff.fetchDependenciesConcurrent(name, version) + resultChan <- result{name, depNode, err} + }(depName, depVersion) + } + + // Wait for all goroutines to complete in a separate goroutine + go func() { + wg.Wait() + close(resultChan) + }() + + // Collect results + for res := range resultChan { + if res.err != nil { + log.Warnf("Failed to fetch dependency %s: %v\n", res.name, res.err) + continue + } + node.Dependencies[res.name] = res.node + } + + return node, nil +} + +func (ff *FlatteningFetcher) fetchPackageInfo(packageName, version string) (*PackageInfo, error) { + url := fmt.Sprintf("https://registry.npmjs.org/%s/%s", packageName, version) + resp, err := ff.client.Get(url) + if err != nil { + return nil, err + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("failed to fetch package info: %s", resp.Status) + } + + body, err := io.ReadAll(resp.Body) + if err != nil { + return nil, err + } + + var packageInfo PackageInfo + if err := json.Unmarshal(body, &packageInfo); err != nil { + return nil, err + } + + return &packageInfo, nil +}