mirror of
https://github.com/safedep/pmg.git
synced 2026-08-03 07:24:09 +02:00
Add support for defaulting to latest version and improve code structure
This commit is contained in:
@@ -39,7 +39,8 @@ go install github.com/safedep/pmg
|
||||
|
||||
- `SAFEDEP_API_KEY`: Your SafeDep API key
|
||||
- `SAFEDEP_TENANT_ID`: Your SafeDep tenant ID
|
||||
- `NPM_AUTH_TOKEN`: (Optional) NPM authentication token for private packages
|
||||
|
||||
Visit https://docs.safedep.io/cloud/quickstart for instructions on obtaining your API Key and Tenant ID.
|
||||
|
||||
## Usage
|
||||
|
||||
|
||||
+27
-4
@@ -4,12 +4,14 @@ import (
|
||||
"context"
|
||||
_ "embed"
|
||||
"fmt"
|
||||
"os"
|
||||
"time"
|
||||
|
||||
"github.com/safedep/dry/log"
|
||||
"github.com/safedep/pmg/pkg/analyser"
|
||||
"github.com/safedep/pmg/pkg/common/utils"
|
||||
"github.com/safedep/pmg/pkg/models"
|
||||
"github.com/safedep/pmg/pkg/registry"
|
||||
vetUtils "github.com/safedep/vet/pkg/common/utils"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
@@ -33,7 +35,7 @@ func NewNpmCommand() *cobra.Command {
|
||||
err := wrapNpm()
|
||||
if err != nil {
|
||||
log.Errorf("Failed to wrap npm: %v", err)
|
||||
return err
|
||||
os.Exit(1)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -59,10 +61,31 @@ func wrapNpm() error {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Minute)
|
||||
defer cancel()
|
||||
|
||||
deps, err := analyser.FetchDependencies(packageName)
|
||||
factory := registry.NewFetcherFactory(10 * time.Second)
|
||||
|
||||
// Get an NPM fetcher
|
||||
npmFetcher, err := factory.CreateFetcher(registry.RegistryNPM)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
name, version, err := utils.ParsePackageInfo(packageName)
|
||||
|
||||
// If version is empty, get the latest version
|
||||
if version == "" {
|
||||
log.Infof("No version specified for %s, fetching latest version...", name)
|
||||
version, err = npmFetcher.(*registry.NpmFetcher).ResolveVersion(ctx, name, version)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
log.Infof("Latest version of %s is %s", name, version)
|
||||
// Update packageName with resolved version for npm installation
|
||||
packageName = fmt.Sprintf("%s@%s", name, version)
|
||||
}
|
||||
deps, err := npmFetcher.GetFlattenedDependencies(ctx, name, version)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
maliciousPkgs := make(map[string]string)
|
||||
|
||||
client, err := analyser.GetMalwareAnalysisClient()
|
||||
@@ -73,7 +96,7 @@ func wrapNpm() error {
|
||||
handler := analyser.AnalysePackage(maliciousPkgs, client, ctx)
|
||||
|
||||
// Create work queue with appropriate buffer size and concurrency
|
||||
queue := vetUtils.NewWorkQueue[models.PackageAnalysisItem](100, 10, handler)
|
||||
queue := vetUtils.NewWorkQueue[models.Package](100, 10, handler)
|
||||
queue.Start()
|
||||
defer queue.Stop()
|
||||
|
||||
@@ -84,7 +107,7 @@ func wrapNpm() error {
|
||||
log.Errorf("Error while parsing info of package %s", name)
|
||||
continue
|
||||
}
|
||||
queue.Add(models.PackageAnalysisItem{
|
||||
queue.Add(models.Package{
|
||||
Name: name,
|
||||
Version: version,
|
||||
})
|
||||
|
||||
@@ -10,29 +10,12 @@ 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 {
|
||||
func AnalysePackage(maliciousPkgs map[string]string, client malysisv1grpc.MalwareAnalysisServiceClient, ctx context.Context) vetUtils.WorkQueueFn[models.Package] {
|
||||
return func(q *vetUtils.WorkQueue[models.Package], item models.Package) error {
|
||||
var maliciousPkgsMutex sync.Mutex
|
||||
resp, err := SubmitPackageForAnalysis(ctx, client,
|
||||
packagev1.Ecosystem_ECOSYSTEM_NPM, item.Name, item.Version)
|
||||
|
||||
@@ -20,19 +20,16 @@ func ParsePackageInfo(input string) (packageName, version string, err error) {
|
||||
}
|
||||
|
||||
pkg := strings.Split(input, "@")
|
||||
if len(pkg) != 2 {
|
||||
return "", "", fmt.Errorf("invalid format: expected 'package@version', got '%s'", input)
|
||||
if len(pkg) == 2 {
|
||||
packageName = strings.TrimSpace(pkg[0])
|
||||
version = strings.TrimSpace(pkg[1])
|
||||
return packageName, version, nil
|
||||
}
|
||||
|
||||
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")
|
||||
if len(pkg) == 1 {
|
||||
packageName = strings.TrimSpace(pkg[0])
|
||||
return packageName, "", nil
|
||||
}
|
||||
|
||||
return packageName, version, nil
|
||||
return "", "", fmt.Errorf("invalid format: expected 'package' OR 'package@version' , got '%s'", input)
|
||||
}
|
||||
|
||||
+2
-131
@@ -1,23 +1,15 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/safedep/dry/log"
|
||||
"github.com/safedep/pmg/pkg/common/utils"
|
||||
)
|
||||
|
||||
type PackageAnalysisItem struct {
|
||||
type Package struct {
|
||||
Name string
|
||||
Version string
|
||||
}
|
||||
|
||||
func (p PackageAnalysisItem) Id() string {
|
||||
func (p Package) Id() string {
|
||||
return fmt.Sprintf("%s@%s", p.Name, p.Version)
|
||||
}
|
||||
|
||||
@@ -32,124 +24,3 @@ type DependencyNode struct {
|
||||
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
|
||||
}
|
||||
|
||||
@@ -0,0 +1,110 @@
|
||||
package registry
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/safedep/pmg/pkg/models"
|
||||
)
|
||||
|
||||
// RegistryClient defines the interface for making requests to a registry
|
||||
type RegistryClient interface {
|
||||
// FetchPackageInfo fetches metadata for a specific package version
|
||||
FetchPackageInfo(ctx context.Context, pkg models.Package) (*models.PackageInfo, error)
|
||||
// GetLatestVersion fetches the latest version for a package
|
||||
GetLatestVersion(ctx context.Context, packageName string) (string, error)
|
||||
}
|
||||
|
||||
// HttpRegistryClient is a basic HTTP client for registry APIs
|
||||
type HttpRegistryClient struct {
|
||||
httpClient *http.Client
|
||||
urlFormat string
|
||||
parser func([]byte) (*models.PackageInfo, error)
|
||||
}
|
||||
|
||||
// NewHttpRegistryClient creates a new HTTP registry client
|
||||
func NewHttpRegistryClient(
|
||||
timeout time.Duration,
|
||||
urlFormat string,
|
||||
parser func([]byte) (*models.PackageInfo, error),
|
||||
) *HttpRegistryClient {
|
||||
return &HttpRegistryClient{
|
||||
httpClient: &http.Client{Timeout: timeout},
|
||||
urlFormat: urlFormat,
|
||||
parser: parser,
|
||||
}
|
||||
}
|
||||
|
||||
// FetchPackageInfo fetches package metadata from the registry
|
||||
func (c *HttpRegistryClient) FetchPackageInfo(ctx context.Context, pkg models.Package) (*models.PackageInfo, error) {
|
||||
url := fmt.Sprintf(c.urlFormat, pkg.Name, pkg.Version)
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("creating request: %w", err)
|
||||
}
|
||||
|
||||
resp, err := c.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("making request: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("registry returned status: %s", resp.Status)
|
||||
}
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("reading response body: %w", err)
|
||||
}
|
||||
|
||||
return c.parser(body)
|
||||
}
|
||||
|
||||
// GetLatestVersion fetches the latest version for an NPM package
|
||||
func (c *HttpRegistryClient) GetLatestVersion(ctx context.Context, packageName string) (string, error) {
|
||||
// For NPM, we can get latest version by querying the base package URL
|
||||
url := fmt.Sprintf("https://registry.npmjs.org/%s", packageName)
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("creating request: %w", err)
|
||||
}
|
||||
|
||||
resp, err := c.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("making request: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", fmt.Errorf("registry returned status: %s", resp.Status)
|
||||
}
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("reading response body: %w", err)
|
||||
}
|
||||
|
||||
// Parse the response to get the latest version
|
||||
var pkgData struct {
|
||||
DistTags struct {
|
||||
Latest string `json:"latest"`
|
||||
} `json:"dist-tags"`
|
||||
}
|
||||
|
||||
if err := json.Unmarshal(body, &pkgData); err != nil {
|
||||
return "", fmt.Errorf("parsing package info: %w", err)
|
||||
}
|
||||
|
||||
if pkgData.DistTags.Latest == "" {
|
||||
return "", fmt.Errorf("no latest version found for package %s", packageName)
|
||||
}
|
||||
|
||||
return pkgData.DistTags.Latest, nil
|
||||
}
|
||||
@@ -0,0 +1,37 @@
|
||||
package registry
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
// RegistryType represents different package registries
|
||||
type RegistryType string
|
||||
|
||||
const (
|
||||
RegistryNPM RegistryType = "npm"
|
||||
RegistryPyPI RegistryType = "pypi"
|
||||
RegistryGo RegistryType = "go"
|
||||
)
|
||||
|
||||
// FetcherFactory creates appropriate fetchers based on registry type
|
||||
type FetcherFactory struct {
|
||||
timeout time.Duration
|
||||
}
|
||||
|
||||
// NewFetcherFactory creates a new factory for registry fetchers
|
||||
func NewFetcherFactory(timeout time.Duration) *FetcherFactory {
|
||||
return &FetcherFactory{
|
||||
timeout: timeout,
|
||||
}
|
||||
}
|
||||
|
||||
// CreateFetcher returns a fetcher for the specified registry type
|
||||
func (ff *FetcherFactory) CreateFetcher(registryType RegistryType) (Fetcher, error) {
|
||||
switch registryType {
|
||||
case RegistryNPM:
|
||||
return NewNpmFetcher(ff.timeout), nil
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported registry type: %s", registryType)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,75 @@
|
||||
// Package registry provides interfaces and implementations for fetching dependencies
|
||||
// from various package registries (npm, pypi, go, etc.)
|
||||
package registry
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sync"
|
||||
|
||||
"github.com/safedep/pmg/pkg/models"
|
||||
)
|
||||
|
||||
// Fetcher defines the interface for registry dependency fetchers
|
||||
type Fetcher interface {
|
||||
// GetDependencyTree returns the complete dependency tree for a package
|
||||
GetDependencyTree(ctx context.Context, pkg models.Package) (*models.DependencyNode, error)
|
||||
|
||||
// GetFlattenedDependencies returns a list of all dependencies as package@version strings
|
||||
GetFlattenedDependencies(ctx context.Context, packageName, version string) ([]string, error)
|
||||
}
|
||||
|
||||
// BaseFetcher implements common functionality for all registry fetchers
|
||||
type BaseFetcher struct {
|
||||
visitedMu sync.RWMutex
|
||||
visited map[string]bool
|
||||
client RegistryClient
|
||||
}
|
||||
|
||||
// NewBaseFetcher creates a new BaseFetcher with the specified registry client
|
||||
func NewBaseFetcher(client RegistryClient) *BaseFetcher {
|
||||
return &BaseFetcher{
|
||||
visited: make(map[string]bool),
|
||||
client: client,
|
||||
}
|
||||
}
|
||||
|
||||
// isVisited checks if a package has already been visited
|
||||
func (bf *BaseFetcher) isVisited(key string) bool {
|
||||
bf.visitedMu.RLock()
|
||||
defer bf.visitedMu.RUnlock()
|
||||
return bf.visited[key]
|
||||
}
|
||||
|
||||
// markVisited marks a package as visited
|
||||
func (bf *BaseFetcher) markVisited(key string) {
|
||||
bf.visitedMu.Lock()
|
||||
defer bf.visitedMu.Unlock()
|
||||
bf.visited[key] = true
|
||||
}
|
||||
|
||||
// cacheKey generates a unique key for a package
|
||||
func cacheKey(pkg models.Package) string {
|
||||
return fmt.Sprintf("%s@%s", pkg.Name, pkg.Version)
|
||||
}
|
||||
|
||||
// flattenDependencyTree recursively converts a dependency tree to a flat list of strings
|
||||
func flattenDependencyTree(node *models.DependencyNode, result *[]string) {
|
||||
if node == nil {
|
||||
return
|
||||
}
|
||||
|
||||
depString := fmt.Sprintf("%s@%s", node.Name, node.Version)
|
||||
*result = append(*result, depString)
|
||||
|
||||
for _, dep := range node.Dependencies {
|
||||
flattenDependencyTree(dep, result)
|
||||
}
|
||||
}
|
||||
|
||||
// resetVisited resets the visited packages map
|
||||
func (bf *BaseFetcher) resetVisited() {
|
||||
bf.visitedMu.Lock()
|
||||
defer bf.visitedMu.Unlock()
|
||||
bf.visited = make(map[string]bool)
|
||||
}
|
||||
@@ -0,0 +1,163 @@
|
||||
package registry
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/safedep/dry/log"
|
||||
"github.com/safedep/pmg/pkg/common/utils"
|
||||
"github.com/safedep/pmg/pkg/models"
|
||||
)
|
||||
|
||||
// NpmFetcher fetches dependencies from NPM registry
|
||||
type NpmFetcher struct {
|
||||
*BaseFetcher
|
||||
}
|
||||
|
||||
// NewNpmFetcher creates a new NPM registry fetcher
|
||||
func NewNpmFetcher(timeout time.Duration) *NpmFetcher {
|
||||
client := NewHttpRegistryClient(
|
||||
timeout,
|
||||
"https://registry.npmjs.org/%s/%s",
|
||||
parseNpmPackageInfo,
|
||||
)
|
||||
return &NpmFetcher{
|
||||
BaseFetcher: NewBaseFetcher(client),
|
||||
}
|
||||
}
|
||||
|
||||
// parseNpmPackageInfo parses NPM package information from JSON
|
||||
func parseNpmPackageInfo(data []byte) (*models.PackageInfo, error) {
|
||||
var packageInfo models.PackageInfo
|
||||
if err := json.Unmarshal(data, &packageInfo); err != nil {
|
||||
return nil, fmt.Errorf("parsing package info: %w", err)
|
||||
}
|
||||
return &packageInfo, nil
|
||||
}
|
||||
|
||||
// GetDependencyTree fetches the complete dependency tree for an NPM package
|
||||
func (nf *NpmFetcher) GetDependencyTree(ctx context.Context, pkg models.Package) (*models.DependencyNode, error) {
|
||||
return nf.fetchDependenciesConcurrent(ctx, pkg)
|
||||
}
|
||||
|
||||
// GetFlattenedDependencies returns a flat list of all dependencies as strings
|
||||
func (nf *NpmFetcher) GetFlattenedDependencies(ctx context.Context, packageName, version string) ([]string, error) {
|
||||
// Reset the visited map to ensure we get a complete tree
|
||||
nf.resetVisited()
|
||||
|
||||
// Get the complete dependency tree
|
||||
tree, err := nf.GetDependencyTree(ctx, models.Package{
|
||||
Name: packageName,
|
||||
Version: version,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to fetch dependency tree: %w", err)
|
||||
}
|
||||
|
||||
// Convert tree to flat list
|
||||
var dependencies []string
|
||||
flattenDependencyTree(tree, &dependencies)
|
||||
|
||||
// Remove duplicates if needed
|
||||
uniqueDeps := make(map[string]bool)
|
||||
var result []string
|
||||
|
||||
for _, dep := range dependencies {
|
||||
if !uniqueDeps[dep] {
|
||||
uniqueDeps[dep] = true
|
||||
result = append(result, dep)
|
||||
}
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// fetchDependenciesConcurrent recursively fetches package dependencies concurrently
|
||||
func (nf *NpmFetcher) fetchDependenciesConcurrent(ctx context.Context, pkg models.Package) (*models.DependencyNode, error) {
|
||||
key := cacheKey(pkg)
|
||||
if nf.isVisited(key) {
|
||||
return &models.DependencyNode{
|
||||
Name: pkg.Name,
|
||||
Version: pkg.Version,
|
||||
}, nil
|
||||
}
|
||||
nf.markVisited(key)
|
||||
|
||||
packageInfo, err := nf.client.FetchPackageInfo(ctx, pkg)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to fetch package info for %s: %w", pkg.Name, err)
|
||||
}
|
||||
|
||||
dependencies := packageInfo.Dependencies
|
||||
node := &models.DependencyNode{
|
||||
Name: pkg.Name,
|
||||
Version: pkg.Version,
|
||||
Dependencies: make(map[string]*models.DependencyNode),
|
||||
}
|
||||
|
||||
if len(dependencies) == 0 {
|
||||
return node, nil
|
||||
}
|
||||
|
||||
// Process dependencies concurrently
|
||||
type result struct {
|
||||
name string
|
||||
node *models.DependencyNode
|
||||
err error
|
||||
}
|
||||
|
||||
var wg sync.WaitGroup
|
||||
resultChan := make(chan result, len(dependencies))
|
||||
|
||||
for depName, depVersion := range dependencies {
|
||||
// Check if context is canceled
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
default:
|
||||
// Continue processing
|
||||
}
|
||||
|
||||
wg.Add(1)
|
||||
go func(name, version string) {
|
||||
defer wg.Done()
|
||||
version = utils.CleanVersion(version)
|
||||
depNode, err := nf.fetchDependenciesConcurrent(ctx, models.Package{Name: name, Version: 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", res.name, res.err)
|
||||
continue
|
||||
}
|
||||
node.Dependencies[res.name] = res.node
|
||||
}
|
||||
|
||||
return node, nil
|
||||
}
|
||||
|
||||
func (nf *NpmFetcher) ResolveVersion(ctx context.Context, packageName, version string) (string, error) {
|
||||
if version != "" {
|
||||
return version, nil
|
||||
}
|
||||
|
||||
latestVersion, err := nf.client.GetLatestVersion(ctx, packageName)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to get latest version for %s: %w", packageName, err)
|
||||
}
|
||||
|
||||
return latestVersion, nil
|
||||
}
|
||||
Reference in New Issue
Block a user