Cybersecurity-Projects/PROJECTS/simple-vulnerability-scanner/internal/cli/update.go

423 lines
8.4 KiB
Go

// ©AngelaMos | 2026
// update.go
package cli
import (
"context"
"fmt"
"os"
"path/filepath"
"sort"
"time"
"github.com/CarterPerez-dev/angela/internal/osv"
"github.com/CarterPerez-dev/angela/internal/pypi"
"github.com/CarterPerez-dev/angela/internal/pyproject"
"github.com/CarterPerez-dev/angela/pkg/types"
"github.com/spf13/cobra"
)
type updateFlags struct {
file string
safe bool
vulns bool
includePrerelease bool
}
func defaultCacheDir() string {
home, err := os.UserHomeDir()
if err != nil {
home = "."
}
return filepath.Join(home, ".angela", "cache")
}
// Execute sets up the CLI and runs the root command
func Execute() {
root := &cobra.Command{
Use: "angela",
Short: "Python dependency updater and vulnerability scanner",
Long: `angela scans your pyproject.toml, updates dependencies to their
latest stable versions, and checks for known CVEs using OSV.dev.`,
SilenceUsage: true,
SilenceErrors: true,
}
root.AddCommand(
newUpdateCmd(),
newCheckCmd(),
newScanCmd(),
newCacheCmd(),
)
if err := root.Execute(); err != nil {
PrintError(err.Error())
os.Exit(1)
}
}
func newUpdateCmd() *cobra.Command {
f := &updateFlags{}
cmd := &cobra.Command{
Use: "update",
Aliases: []string{"u"},
Short: "Update dependencies to latest stable versions",
RunE: func(cmd *cobra.Command, _ []string) error {
return runUpdate(cmd.Context(), f, false)
},
}
cmd.Flags().StringVarP(
&f.file, "file", "f", "pyproject.toml",
"path to pyproject.toml",
)
cmd.Flags().BoolVar(
&f.safe, "safe", false,
"skip major version bumps",
)
cmd.Flags().BoolVar(
&f.vulns, "vulns", false,
"also scan for vulnerabilities",
)
cmd.Flags().BoolVar(
&f.includePrerelease, "include-prerelease", false,
"include pre-release versions",
)
return cmd
}
func newCheckCmd() *cobra.Command {
f := &updateFlags{}
cmd := &cobra.Command{
Use: "check",
Aliases: []string{"c"},
Short: "Show available updates without modifying files",
RunE: func(cmd *cobra.Command, _ []string) error {
return runUpdate(cmd.Context(), f, true)
},
}
cmd.Flags().StringVarP(
&f.file, "file", "f", "pyproject.toml",
"path to pyproject.toml",
)
cmd.Flags().BoolVar(
&f.safe, "safe", false,
"skip major version bumps",
)
cmd.Flags().BoolVar(
&f.vulns, "vulns", false,
"also scan for vulnerabilities",
)
cmd.Flags().BoolVar(
&f.includePrerelease, "include-prerelease", false,
"include pre-release versions",
)
return cmd
}
func newScanCmd() *cobra.Command {
var file string
cmd := &cobra.Command{
Use: "scan",
Aliases: []string{"s"},
Short: "Scan dependencies for known vulnerabilities",
RunE: func(cmd *cobra.Command, _ []string) error {
return runScan(cmd.Context(), file)
},
}
cmd.Flags().StringVarP(
&file, "file", "f", "pyproject.toml",
"path to pyproject.toml",
)
return cmd
}
func newCacheCmd() *cobra.Command {
cmd := &cobra.Command{
Use: "cache",
Short: "Manage the local response cache",
}
cmd.AddCommand(&cobra.Command{
Use: "clear",
Short: "Remove all cached PyPI responses",
RunE: func(_ *cobra.Command, _ []string) error {
client, err := pypi.NewClient(defaultCacheDir())
if err != nil {
return err
}
if err := client.ClearCache(); err != nil {
return err
}
fmt.Println(" Cache cleared.")
return nil
},
})
return cmd
}
func runUpdate(
ctx context.Context,
f *updateFlags,
dryRun bool,
) error {
start := time.Now()
deps, err := pyproject.ParseFile(f.file)
if err != nil {
return err
}
PrintScanning(len(deps))
client, err := pypi.NewClient(defaultCacheDir())
if err != nil {
return err
}
names := make([]string, len(deps))
for i, d := range deps {
names[i] = d.Name
}
fetched := client.FetchAllVersions(ctx, names)
versionMap := make(map[string][]string, len(fetched))
for _, r := range fetched {
if r.Err == nil {
versionMap[r.Name] = r.Versions
}
}
updates, updateSpecs := resolveUpdates(
deps, versionMap, f.safe, f.includePrerelease,
)
sortUpdates(updates)
var vulns map[string][]types.Vulnerability
if f.vulns {
vulns = scanForVulns(ctx, deps)
}
if !dryRun && len(updateSpecs) > 0 {
if err := pyproject.UpdateFile(f.file, updateSpecs); err != nil {
return fmt.Errorf("write updates: %w", err)
}
}
PrintUpdates(updates)
PrintSkipped(updates)
if vulns != nil {
PrintVulnerabilities(vulns)
}
totalVulns := 0
for _, vl := range vulns {
totalVulns += len(vl)
}
PrintSummary(types.ScanResult{
Updates: updates,
Vulnerabilities: vulns,
TotalPackages: len(deps),
TotalUpdated: len(updateSpecs),
TotalVulns: totalVulns,
VulnsScanned: f.vulns,
Duration: time.Since(start),
}, !dryRun && len(updateSpecs) > 0)
return nil
}
func runScan(ctx context.Context, file string) error {
start := time.Now()
deps, err := pyproject.ParseFile(file)
if err != nil {
return err
}
PrintScanning(len(deps))
vulns := scanForVulns(ctx, deps)
PrintVulnerabilities(vulns)
totalVulns := 0
for _, vl := range vulns {
totalVulns += len(vl)
}
PrintSummary(types.ScanResult{
TotalPackages: len(deps),
TotalVulns: totalVulns,
VulnsScanned: true,
Duration: time.Since(start),
}, false)
return nil
}
func resolveUpdates(
deps []types.Dependency,
versionMap map[string][]string,
safe bool,
includePrerelease bool,
) ([]types.UpdateResult, map[string]string) {
var results []types.UpdateResult
specs := make(map[string]string)
for _, dep := range deps {
versions, ok := versionMap[dep.Name]
if !ok {
results = append(results, types.UpdateResult{
Name: dep.Name,
Skipped: true,
Reason: "not found on PyPI",
})
continue
}
currentStr := pyproject.ExtractMinVersion(dep.Spec)
if currentStr == "" {
results = append(results, types.UpdateResult{
Name: dep.Name,
Skipped: true,
Reason: "no version specifier",
})
continue
}
current, err := pypi.ParseVersion(currentStr)
if err != nil {
results = append(results, types.UpdateResult{
Name: dep.Name,
Skipped: true,
Reason: "unparseable version",
})
continue
}
var latest pypi.Version
if includePrerelease {
latest, err = latestAny(versions)
} else {
latest, err = pypi.LatestStable(versions)
}
if err != nil {
results = append(results, types.UpdateResult{
Name: dep.Name,
Skipped: true,
Reason: err.Error(),
})
continue
}
if latest.Compare(current) <= 0 {
continue
}
change := pypi.ClassifyChange(current, latest)
if safe && change == pypi.Major {
results = append(results, types.UpdateResult{
Name: dep.Name,
OldVer: current.String(),
NewVer: latest.String(),
Change: change.String(),
Skipped: true,
Reason: "major bump (use --all to include)",
})
continue
}
newSpec := ">=" + latest.String()
results = append(results, types.UpdateResult{
Name: dep.Name,
OldVer: current.String(),
NewVer: latest.String(),
OldSpec: dep.Spec,
NewSpec: newSpec,
Change: change.String(),
})
specs[dep.Name] = newSpec
}
return results, specs
}
func scanForVulns(
ctx context.Context,
deps []types.Dependency,
) map[string][]types.Vulnerability {
var queries []osv.PackageQuery
for _, dep := range deps {
ver := pyproject.ExtractMinVersion(dep.Spec)
if ver == "" {
continue
}
queries = append(queries, osv.PackageQuery{
Name: dep.Name,
Version: ver,
})
}
if len(queries) == 0 {
return nil
}
client := osv.NewClient()
vulns, err := client.ScanPackages(ctx, queries)
if err != nil {
PrintError(fmt.Sprintf("vulnerability scan: %v", err))
return nil
}
return vulns
}
func latestAny(versions []string) (pypi.Version, error) {
var latest pypi.Version
var found bool
for _, raw := range versions {
v, err := pypi.ParseVersion(raw)
if err != nil {
continue
}
if !found || v.Compare(latest) > 0 {
latest = v
found = true
}
}
if !found {
return pypi.Version{}, fmt.Errorf("no versions found")
}
return latest, nil
}
func sortUpdates(updates []types.UpdateResult) {
order := map[string]int{
pypi.Major.String(): 0,
pypi.Minor.String(): 1,
pypi.Patch.String(): 2,
}
sort.Slice(updates, func(i, j int) bool {
if updates[i].Skipped != updates[j].Skipped {
return !updates[i].Skipped
}
oi := order[updates[i].Change]
oj := order[updates[j].Change]
if oi != oj {
return oi < oj
}
return updates[i].Name < updates[j].Name
})
}