refactor: Analyzer to generalise

This commit is contained in:
abhisek
2025-05-14 18:56:58 +05:30
parent 8eace0c7f6
commit 4dcf7cad0b
7 changed files with 114 additions and 51 deletions
+25 -13
View File
@@ -3,7 +3,6 @@ package analyzer
import ( import (
"context" "context"
malysisv1 "buf.build/gen/go/safedep/api/protocolbuffers/go/safedep/messages/malysis/v1"
packagev1 "buf.build/gen/go/safedep/api/protocolbuffers/go/safedep/messages/package/v1" packagev1 "buf.build/gen/go/safedep/api/protocolbuffers/go/safedep/messages/package/v1"
) )
@@ -12,21 +11,34 @@ type Analyzer interface {
Name() string Name() string
} }
type MalysisResult struct { type Action int
const (
ActionUnknown Action = iota
ActionAllow
ActionConfirm
ActionBlock
)
type PackageVersionAnalysisResult struct {
PackageVersion *packagev1.PackageVersion
// Analyser specific analysis ID
AnalysisID string AnalysisID string
Report *malysisv1.Report
// The action to take as recommended by the analyzer
Action Action
// Summary of the analysis
Summary string
// Analyzer specific data
Data any
} }
func (m *MalysisResult) IsMalware() bool { // Contract for implementing package version specific analyzers
return m.Report.GetInference().GetIsMalware() type PackageVersionAnalyzer interface {
}
func (m *MalysisResult) Summary() string {
return m.Report.GetInference().GetSummary()
}
type MalysisAnalyzer interface {
Analyzer Analyzer
Analyze(ctx context.Context, packageVersion *packagev1.PackageVersion) (*MalysisResult, error) Analyze(ctx context.Context, packageVersion *packagev1.PackageVersion) (*PackageVersionAnalysisResult, error)
} }
+22 -5
View File
@@ -21,7 +21,7 @@ type malysisQueryAnalyzer struct {
} }
var _ Analyzer = &malysisQueryAnalyzer{} var _ Analyzer = &malysisQueryAnalyzer{}
var _ MalysisAnalyzer = &malysisQueryAnalyzer{} var _ PackageVersionAnalyzer = &malysisQueryAnalyzer{}
func NewMalysisQueryAnalyzer(config MalysisQueryAnalyzerConfig) (*malysisQueryAnalyzer, error) { func NewMalysisQueryAnalyzer(config MalysisQueryAnalyzerConfig) (*malysisQueryAnalyzer, error) {
client, err := drygrpc.GrpcClient("pmg-malysis-query", client, err := drygrpc.GrpcClient("pmg-malysis-query",
@@ -41,7 +41,7 @@ func (a *malysisQueryAnalyzer) Name() string {
} }
func (a *malysisQueryAnalyzer) Analyze(ctx context.Context, func (a *malysisQueryAnalyzer) Analyze(ctx context.Context,
packageVersion *packagev1.PackageVersion) (*MalysisResult, error) { packageVersion *packagev1.PackageVersion) (*PackageVersionAnalysisResult, error) {
res, err := a.client.QueryPackageAnalysis(ctx, &malysisv1.QueryPackageAnalysisRequest{ res, err := a.client.QueryPackageAnalysis(ctx, &malysisv1.QueryPackageAnalysisRequest{
Target: &malysisv1pb.PackageAnalysisTarget{ Target: &malysisv1pb.PackageAnalysisTarget{
@@ -52,7 +52,24 @@ func (a *malysisQueryAnalyzer) Analyze(ctx context.Context,
return nil, fmt.Errorf("failed to query package analysis: %w", err) return nil, fmt.Errorf("failed to query package analysis: %w", err)
} }
return &MalysisResult{ // By default, the analyzer allows the package version
Report: res.GetReport(), analysisResult := &PackageVersionAnalysisResult{
}, nil PackageVersion: packageVersion,
Action: ActionAllow,
AnalysisID: res.GetAnalysisId(),
Summary: res.GetReport().GetInference().GetSummary(),
Data: res.GetReport(),
}
// Mark the package version to be confirmed if it is malicious (not confirmed)
if res.GetReport().GetInference().GetIsMalware() {
analysisResult.Action = ActionConfirm
}
// This is a confirmed malicious package, we must always block it
if res.GetVerificationRecord().GetIsMalware() {
analysisResult.Action = ActionBlock
}
return analysisResult, nil
} }
+1 -1
View File
@@ -34,7 +34,7 @@ func executeCommonFlow(ctx context.Context, config config.Config, pm packagemana
} }
proxy, err := guard.NewPackageManagerGuard(guard.DefaultPackageManagerGuardConfig(), proxy, err := guard.NewPackageManagerGuard(guard.DefaultPackageManagerGuardConfig(),
pm, packageResolver, []analyzer.MalysisAnalyzer{malysisQueryAnalyzer}, interaction) pm, packageResolver, []analyzer.PackageVersionAnalyzer{malysisQueryAnalyzer}, interaction)
if err != nil { if err != nil {
return fmt.Errorf("failed to create package manager guard: %w", err) return fmt.Errorf("failed to create package manager guard: %w", err)
} }
+18 -18
View File
@@ -17,7 +17,7 @@ import (
type PackageManagerGuardInteraction struct { type PackageManagerGuardInteraction struct {
SetStatus func(status string) SetStatus func(status string)
ClearStatus func() ClearStatus func()
GetConfirmationOnMalware func(malwarePackages []*packagev1.PackageVersion) (bool, error) GetConfirmationOnMalware func(malwarePackages []*analyzer.PackageVersionAnalysisResult) (bool, error)
Block func() error Block func() error
} }
@@ -38,7 +38,7 @@ func DefaultPackageManagerGuardConfig() PackageManagerGuardConfig {
type packageManagerGuard struct { type packageManagerGuard struct {
config PackageManagerGuardConfig config PackageManagerGuardConfig
interaction PackageManagerGuardInteraction interaction PackageManagerGuardInteraction
analyzers []analyzer.MalysisAnalyzer analyzers []analyzer.PackageVersionAnalyzer
packageManager packagemanager.PackageManager packageManager packagemanager.PackageManager
packageResolver packagemanager.PackageResolver packageResolver packagemanager.PackageResolver
} }
@@ -46,7 +46,7 @@ type packageManagerGuard struct {
func NewPackageManagerGuard(config PackageManagerGuardConfig, func NewPackageManagerGuard(config PackageManagerGuardConfig,
packageManager packagemanager.PackageManager, packageManager packagemanager.PackageManager,
packageResolver packagemanager.PackageResolver, packageResolver packagemanager.PackageResolver,
analyzers []analyzer.MalysisAnalyzer, analyzers []analyzer.PackageVersionAnalyzer,
interaction PackageManagerGuardInteraction, interaction PackageManagerGuardInteraction,
) (*packageManagerGuard, error) { ) (*packageManagerGuard, error) {
return &packageManagerGuard{ return &packageManagerGuard{
@@ -118,15 +118,20 @@ func (g *packageManagerGuard) Run(ctx context.Context, args []string) error {
return fmt.Errorf("failed to analyze packages: %w", err) return fmt.Errorf("failed to analyze packages: %w", err)
} }
maliciousPackages := []*packagev1.PackageVersion{} confirmableMalwarePackages := []*analyzer.PackageVersionAnalysisResult{}
for _, result := range analysisResults { for _, result := range analysisResults {
if result.result.IsMalware() { if result.Action == analyzer.ActionBlock {
maliciousPackages = append(maliciousPackages, result.pkg) _ = g.blockInstallation()
return fmt.Errorf("malicious packages detected, installation aborted")
}
if result.Action == analyzer.ActionConfirm {
confirmableMalwarePackages = append(confirmableMalwarePackages, result)
} }
} }
if len(maliciousPackages) > 0 { if len(confirmableMalwarePackages) > 0 {
confirmed, err := g.getConfirmationOnMalware(ctx, maliciousPackages) confirmed, err := g.getConfirmationOnMalware(ctx, confirmableMalwarePackages)
if err != nil { if err != nil {
return fmt.Errorf("failed to get confirmation on malware: %w", err) return fmt.Errorf("failed to get confirmation on malware: %w", err)
} }
@@ -156,20 +161,15 @@ func (g *packageManagerGuard) continueExecution(ctx context.Context, pc *package
return cmd.Run() return cmd.Run()
} }
type analyzePackageResult struct {
pkg *packagev1.PackageVersion
result *analyzer.MalysisResult
}
func (g *packageManagerGuard) concurrentAnalyzePackages(ctx context.Context, func (g *packageManagerGuard) concurrentAnalyzePackages(ctx context.Context,
packages []*packagev1.PackageVersion) ([]*analyzePackageResult, error) { packages []*packagev1.PackageVersion) ([]*analyzer.PackageVersionAnalysisResult, error) {
ctx, cancel := context.WithTimeout(ctx, g.config.AnalysisTimeout) ctx, cancel := context.WithTimeout(ctx, g.config.AnalysisTimeout)
defer cancel() defer cancel()
wg := sync.WaitGroup{} wg := sync.WaitGroup{}
jobs := make(chan *packagev1.PackageVersion, len(packages)) jobs := make(chan *packagev1.PackageVersion, len(packages))
results := make(chan *analyzePackageResult, len(packages)) results := make(chan *analyzer.PackageVersionAnalysisResult, len(packages))
for i := 0; i < g.config.MaxConcurrentAnalyzes; i++ { for i := 0; i < g.config.MaxConcurrentAnalyzes; i++ {
wg.Add(1) wg.Add(1)
@@ -184,7 +184,7 @@ func (g *packageManagerGuard) concurrentAnalyzePackages(ctx context.Context,
continue continue
} }
results <- &analyzePackageResult{pkg: pkg, result: analysisResult} results <- analysisResult
} }
} }
}() }()
@@ -195,7 +195,7 @@ func (g *packageManagerGuard) concurrentAnalyzePackages(ctx context.Context,
} }
close(jobs) close(jobs)
analysisResults := []*analyzePackageResult{} analysisResults := []*analyzer.PackageVersionAnalysisResult{}
go func() { go func() {
for result := range results { for result := range results {
analysisResults = append(analysisResults, result) analysisResults = append(analysisResults, result)
@@ -218,7 +218,7 @@ func (g *packageManagerGuard) concurrentAnalyzePackages(ctx context.Context,
return analysisResults, nil return analysisResults, nil
} }
func (g *packageManagerGuard) getConfirmationOnMalware(ctx context.Context, malwarePackages []*packagev1.PackageVersion) (bool, error) { func (g *packageManagerGuard) getConfirmationOnMalware(ctx context.Context, malwarePackages []*analyzer.PackageVersionAnalysisResult) (bool, error) {
if g.interaction.GetConfirmationOnMalware == nil { if g.interaction.GetConfirmationOnMalware == nil {
return false, nil return false, nil
} }
+8 -4
View File
@@ -2,14 +2,18 @@ package ui
import "github.com/fatih/color" import "github.com/fatih/color"
type ColorFn func(format string, a ...interface{}) string
type TerminalColors struct { type TerminalColors struct {
Red func(format string, a ...interface{}) string Normal ColorFn
Yellow func(format string, a ...interface{}) string Red ColorFn
Cyan func(format string, a ...interface{}) string Yellow ColorFn
Green func(format string, a ...interface{}) string Cyan ColorFn
Green ColorFn
} }
var Colors = TerminalColors{ var Colors = TerminalColors{
Normal: color.New().SprintfFunc(),
Red: color.New(color.FgRed, color.Bold).SprintfFunc(), Red: color.New(color.FgRed, color.Bold).SprintfFunc(),
Yellow: color.New(color.FgYellow).SprintfFunc(), Yellow: color.New(color.FgYellow).SprintfFunc(),
Cyan: color.New(color.FgCyan).SprintfFunc(), Cyan: color.New(color.FgCyan).SprintfFunc(),
+9 -1
View File
@@ -8,6 +8,14 @@ import (
var spinnerChan chan bool var spinnerChan chan bool
func StartSpinner(msg string) { func StartSpinner(msg string) {
StartSpinnerWithColor(msg, Colors.Normal)
}
func StartSpinnerWithColor(msg string, c ColorFn) {
if c == nil {
c = Colors.Normal
}
style := `⠋⠙⠹⠸⠼⠴⠦⠧⠇⠏` style := `⠋⠙⠹⠸⠼⠴⠦⠧⠇⠏`
frames := []rune(style) frames := []rune(style)
length := len(frames) length := len(frames)
@@ -24,7 +32,7 @@ func StartSpinner(msg string) {
ticker.Stop() ticker.Stop()
return return
case <-ticker.C: case <-ticker.C:
fmt.Printf("\r%s ... %s", msg, string(frames[pos%length])) fmt.Printf("\r%s ... %s", c(msg), string(frames[pos%length]))
pos += 1 pos += 1
} }
} }
+31 -9
View File
@@ -5,7 +5,7 @@ import (
"os" "os"
"strings" "strings"
packagev1 "buf.build/gen/go/safedep/api/protocolbuffers/go/safedep/messages/package/v1" "github.com/safedep/pmg/analyzer"
) )
// The UI is internal to PMG and opinionated for the CLI. // The UI is internal to PMG and opinionated for the CLI.
@@ -40,7 +40,9 @@ func ClearStatus() {
func Block() error { func Block() error {
StopSpinner() StopSpinner()
fmt.Println(Colors.Red("❌ Malicious packages detected, installation blocked!")) fmt.Println()
fmt.Println(Colors.Red("❌ Malicious package blocked!"))
os.Exit(1) os.Exit(1)
return nil return nil
@@ -52,17 +54,19 @@ func SetStatus(status string) {
} }
StopSpinner() StopSpinner()
StartSpinnerWithColor(fmt.Sprintf("️ %s", status), Colors.Green)
fmt.Print("\r", Colors.Green(status), " ")
StartSpinner(status)
} }
func GetConfirmationOnMalware(malwarePackages []*packagev1.PackageVersion) (bool, error) { func GetConfirmationOnMalware(malwarePackages []*analyzer.PackageVersionAnalysisResult) (bool, error) {
StopSpinner() StopSpinner()
fmt.Println(Colors.Red("🚨 Malicious packages detected:")) fmt.Println(Colors.Red(fmt.Sprintf("🚨 Malicious packages detected: %d", len(malwarePackages))))
fmt.Println()
for _, pkg := range malwarePackages { for _, mp := range malwarePackages {
fmt.Println(" ⚠️ ", Colors.Red(fmt.Sprintf("%s@%s", pkg.Package.Name, pkg.Version))) fmt.Println("⚠️ ", Colors.Red(fmt.Sprintf("%s@%s", mp.PackageVersion.GetPackage().GetName(),
mp.PackageVersion.GetVersion())))
fmt.Println(Colors.Yellow(termWidthFormatText(mp.Summary, 60)))
fmt.Println()
} }
fmt.Println() fmt.Println()
@@ -85,3 +89,21 @@ func GetConfirmationOnMalware(malwarePackages []*packagev1.PackageVersion) (bool
return false, nil return false, nil
} }
// Format the string to be maximum maxWidth. Use newlines to wrap the text.
func termWidthFormatText(text string, maxWidth int) string {
words := strings.Split(text, " ")
lines := []string{}
currentLine := ""
for _, word := range words {
if len(currentLine)+len(word) > maxWidth {
lines = append(lines, currentLine)
currentLine = word
} else {
currentLine += " " + word
}
}
return strings.Join(lines, "\n")
}