Watch
1
0
Fork
You've already forked pkg-proxy
1
mirror of https://github.com/git-pkgs/proxy.git synced 2026-08-23 20:34:56 -04:00
pkg-proxy/internal/enrichment/enrichment_test.go
2026-08-20 09:01:33 +01:00

189 lines
5.1 KiB
Go

package enrichment
import (
"context"
"log/slog"
"os"
"testing"
"github.com/git-pkgs/purl"
"github.com/git-pkgs/vulns"
)
type recordingVulnerabilitySource struct {
purls []*purl.PURL
}
func (s *recordingVulnerabilitySource) Name() string {
return "recording"
}
func (s *recordingVulnerabilitySource) Query(context.Context, *purl.PURL) ([]vulns.Vulnerability, error) {
return nil, nil
}
func (s *recordingVulnerabilitySource) QueryBatch(_ context.Context, purls []*purl.PURL) ([][]vulns.Vulnerability, error) {
s.purls = purls
results := make([][]vulns.Vulnerability, len(purls))
for i := range results {
results[i] = []vulns.Vulnerability{{ID: "TEST-1"}}
}
return results, nil
}
func (s *recordingVulnerabilitySource) Get(context.Context, string) (*vulns.Vulnerability, error) {
return nil, nil
}
func TestNew(t *testing.T) {
logger := slog.New(slog.NewTextHandler(os.Stdout, nil))
svc := New(logger)
if svc == nil {
t.Fatal("New() returned nil")
}
if svc.regClient == nil {
t.Error("regClient is nil")
}
if svc.vulnSource == nil {
t.Error("vulnSource is nil")
}
}
func TestSwiftRegistryIdentitySkipsPURLDependentLookups(t *testing.T) {
svc := New(slog.New(slog.NewTextHandler(os.Stdout, nil)))
ctx := context.Background()
packageInfo, err := svc.EnrichPackage(ctx, "swift", "apple/example")
if err != nil || packageInfo != nil {
t.Errorf("EnrichPackage() = %#v, %v; want nil, nil", packageInfo, err)
}
versionInfo, err := svc.EnrichVersion(ctx, "swift", "apple/example", "1.2.3")
if err != nil || versionInfo != nil {
t.Errorf("EnrichVersion() = %#v, %v; want nil, nil", versionInfo, err)
}
vulnerabilities, err := svc.CheckVulnerabilities(ctx, "swift", "apple/example", "1.2.3")
if err != nil || vulnerabilities != nil {
t.Errorf("CheckVulnerabilities() = %#v, %v; want nil, nil", vulnerabilities, err)
}
latest, err := svc.GetLatestVersion(ctx, "swift", "apple/example")
if err != nil || latest != "" {
t.Errorf("GetLatestVersion() = %q, %v; want empty string, nil", latest, err)
}
packages := []struct{ Ecosystem, Name string }{{Ecosystem: "swift", Name: "apple/example"}}
if got := svc.BulkEnrichPackages(ctx, packages); len(got) != 0 {
t.Errorf("BulkEnrichPackages() = %#v, want empty result", got)
}
versions := []struct{ Ecosystem, Name, Version string }{
{Ecosystem: "swift", Name: "apple/example", Version: "1.2.3"},
}
got, err := svc.BulkCheckVulnerabilities(ctx, versions)
if err != nil || len(got) != 0 {
t.Errorf("BulkCheckVulnerabilities() = %#v, %v; want empty result, nil", got, err)
}
}
func TestBulkCheckVulnerabilitiesFiltersUnsupportedPackageIdentities(t *testing.T) {
source := &recordingVulnerabilitySource{}
svc := New(slog.New(slog.NewTextHandler(os.Stdout, nil)))
svc.vulnSource = source
packages := []struct{ Ecosystem, Name, Version string }{
{Ecosystem: "swift", Name: "apple/example", Version: "1.2.3"},
{Ecosystem: "npm", Name: "lodash", Version: "4.17.21"},
}
got, err := svc.BulkCheckVulnerabilities(context.Background(), packages)
if err != nil {
t.Fatalf("BulkCheckVulnerabilities() error = %v", err)
}
if len(source.purls) != 1 || source.purls[0].String() != "pkg:npm/lodash@4.17.21" {
t.Fatalf("queried PURLs = %#v, want only lodash", source.purls)
}
if len(got["pkg:npm/lodash@4.17.21"]) != 1 {
t.Errorf("result = %#v, want lodash vulnerability", got)
}
if _, exists := got[""]; exists {
t.Error("result contains an empty PURL key")
}
}
func TestIsOutdated(t *testing.T) {
logger := slog.New(slog.NewTextHandler(os.Stdout, nil))
svc := New(logger)
tests := []struct {
current string
latest string
expected bool
}{
{"1.0.0", "2.0.0", true},
{"2.0.0", "2.0.0", false},
{"2.0.0", "1.0.0", false},
{"1.0.0", "", false},
{"", "2.0.0", false},
{"1.2.3", "1.2.4", true},
{"1.2.4", "1.2.3", false},
}
for _, tc := range tests {
result := svc.IsOutdated(tc.current, tc.latest)
if result != tc.expected {
t.Errorf("IsOutdated(%q, %q) = %v, want %v", tc.current, tc.latest, result, tc.expected)
}
}
}
func TestCategorizeLicense(t *testing.T) {
logger := slog.New(slog.NewTextHandler(os.Stdout, nil))
svc := New(logger)
tests := []struct {
license string
expected LicenseCategory
}{
{"MIT", LicensePermissive},
{"Apache-2.0", LicensePermissive},
{"BSD-3-Clause", LicensePermissive},
{"GPL-3.0", LicenseCopyleft},
{"AGPL-3.0", LicenseCopyleft},
{"LGPL-2.1", LicenseCopyleft},
{"", LicenseUnknown},
{"Unknown", LicenseUnknown},
}
for _, tc := range tests {
result := svc.CategorizeLicense(tc.license)
if result != tc.expected {
t.Errorf("CategorizeLicense(%q) = %v, want %v", tc.license, result, tc.expected)
}
}
}
func TestNormalizeLicense(t *testing.T) {
logger := slog.New(slog.NewTextHandler(os.Stdout, nil))
svc := New(logger)
tests := []struct {
input string
expected string
}{
{"MIT", "MIT"},
{"Apache 2", "Apache-2.0"},
{"Apache-2.0", "Apache-2.0"},
{"", ""},
}
for _, tc := range tests {
result := svc.NormalizeLicense(tc.input)
if result != tc.expected {
t.Errorf("NormalizeLicense(%q) = %q, want %q", tc.input, result, tc.expected)
}
}
}