Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions definitions/cargo.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,7 @@ commands:
args:
# cargo add supports package@version syntax
package: {position: 0, required: true, validate: cargo_crate}
version: {position: 0, suffix: "@", required: false}
flags:
dev: [--dev]
build: [--build]
Expand Down
1 change: 1 addition & 0 deletions definitions/gomod.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@ commands:
base: [get]
args:
package: {position: 0, required: true, validate: go_module}
version: {position: 0, suffix: "@", required: false}
flags:
# -t includes test dependencies
test: [-t]
Expand Down
41 changes: 41 additions & 0 deletions generic_manager_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,47 @@ func newTestManager(def *definitions.Definition, runner *MockRunner) *GenericMan
}
}

func embeddedDef(t *testing.T, name string) *definitions.Definition {
t.Helper()
defs, err := definitions.LoadEmbedded()
if err != nil {
t.Fatalf("LoadEmbedded: %v", err)
}
for _, d := range defs {
if d.Name == name {
return d
}
}
t.Fatalf("no embedded definition %q", name)
return nil
}

func TestGenericManager_Add_GomodVersion(t *testing.T) {
runner := NewMockRunner()
mgr := newTestManager(embeddedDef(t, "gomod"), runner)
_, err := mgr.Add(context.Background(), "github.com/pkg/errors", AddOptions{Version: "v0.9.1"})
if err != nil {
t.Fatalf("Add: %v", err)
}
want := []string{"go", "get", "github.com/pkg/errors@v0.9.1"}
if !slicesEqual(runner.Captured[0], want) {
t.Errorf("got %v, want %v", runner.Captured[0], want)
}
}

func TestGenericManager_Add_CargoVersion(t *testing.T) {
runner := NewMockRunner()
mgr := newTestManager(embeddedDef(t, "cargo"), runner)
_, err := mgr.Add(context.Background(), "serde", AddOptions{Version: "1.0.219"})
if err != nil {
t.Fatalf("Add: %v", err)
}
want := []string{"cargo", "add", "serde@1.0.219"}
if !slicesEqual(runner.Captured[0], want) {
t.Errorf("got %v, want %v", runner.Captured[0], want)
}
}

func TestGenericManager_Path_Raw(t *testing.T) {
def := &definitions.Definition{
Name: "testpkg",
Expand Down
35 changes: 27 additions & 8 deletions policy.go
Original file line number Diff line number Diff line change
Expand Up @@ -255,15 +255,34 @@ func (PackageBlocklistPolicy) Name() string { return "package-blocklist" }

func (p PackageBlocklistPolicy) Check(ctx context.Context, op *PolicyOperation) (*PolicyResult, error) {
for _, pkg := range op.Packages {
if reason, blocked := p.Blocked[pkg]; blocked {
return &PolicyResult{
Allowed: false,
Reason: reason,
Metadata: map[string]any{
"blocked_package": pkg,
},
}, nil
if r := p.match(pkg); r != nil {
return r, nil
}
}
return &PolicyResult{Allowed: true}, nil
}

// match checks pkg against the blocklist, both as given and with any
// trailing @version stripped so that a versioned add such as
// `go get example.com/foo@v1.2.3` or `npm install @scope/name@1.0.0`
// is caught by an entry keyed on the bare package name.
func (p PackageBlocklistPolicy) match(pkg string) *PolicyResult {
if reason, blocked := p.Blocked[pkg]; blocked {
return blockedResult(pkg, reason)
}
if i := strings.LastIndex(pkg, "@"); i > 0 {
bare := pkg[:i]
if reason, blocked := p.Blocked[bare]; blocked {
return blockedResult(bare, reason)
}
}
return nil
}

func blockedResult(pkg, reason string) *PolicyResult {
return &PolicyResult{
Allowed: false,
Reason: reason,
Metadata: map[string]any{"blocked_package": pkg},
}
}
33 changes: 33 additions & 0 deletions policy_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -153,6 +153,7 @@ func TestPackageBlocklistPolicy(t *testing.T) {
{"blocked package", []string{"evil-package"}, false},
{"mixed packages", []string{"lodash", "deprecated-lib"}, false},
{"empty packages", []string{}, true},
{"versioned blocked", []string{"evil-package@1.2.3"}, false},
}

for _, tt := range tests {
Expand All @@ -169,6 +170,32 @@ func TestPackageBlocklistPolicy(t *testing.T) {
}
}

func TestPackageBlocklistScopedAndVersioned(t *testing.T) {
policy := PackageBlocklistPolicy{
Blocked: map[string]string{
"@scope/name": "scoped bare",
"github.com/org/mod": "go module",
},
}
cases := []struct {
pkg string
allowed bool
}{
{"@scope/name", false},
{"@scope/name@7.0.0", false},
{"@scope/other", true},
{"github.com/org/mod@v1.2.3", false},
{"github.com/org/mod", false},
{"github.com/org/other@v1.0.0", true},
}
for _, tc := range cases {
res, _ := policy.Check(context.Background(), &PolicyOperation{Packages: []string{tc.pkg}})
if res.Allowed != tc.allowed {
t.Errorf("%s: allowed=%v, want %v", tc.pkg, res.Allowed, tc.allowed)
}
}
}

func TestPackageBlocklistViaRunner(t *testing.T) {
mock := NewMockRunner()
policy := PackageBlocklistPolicy{
Expand All @@ -194,6 +221,12 @@ func TestPackageBlocklistViaRunner(t *testing.T) {
if err != nil {
t.Fatalf("allowed package should pass: %v", err)
}

// A versioned form of a blocked package must also be denied.
_, err = pr.Run(context.Background(), "/tmp", "go", "get", "evil-package@v1.2.3")
if err == nil {
t.Fatal("expected blocklist policy to deny versioned package")
}
}

func TestPolicyRunnerWithContext(t *testing.T) {
Expand Down
43 changes: 26 additions & 17 deletions translator.go
Original file line number Diff line number Diff line change
Expand Up @@ -93,15 +93,17 @@ func (t *Translator) buildSingleCommand(binary string, cmd definitions.Command,

baseOverrideUsed := t.applyBaseOverrides(&args, cmd, input)

packageVal := input.Args["package"]

sortedArgs := t.sortArgs(cmd)

packageIdx := -1
for _, entry := range sortedArgs {
val, err := t.processArg(entry.name, entry.argDef, input, &args)
val, at, err := t.processArg(entry.name, entry.argDef, input, &args)
if err != nil {
return nil, err
}
if entry.name == "package" {
packageIdx = at
}
if val == "" {
continue
}
Expand All @@ -110,7 +112,7 @@ func (t *Translator) buildSingleCommand(binary string, cmd definitions.Command,
}
}

t.applyVersionSuffix(&args, cmd, input, packageVal)
t.applyVersionSuffix(&args, cmd, input, packageIdx)

if !suppressDefaultFlags {
args = append(args, cmd.DefaultFlags...)
Expand Down Expand Up @@ -158,40 +160,52 @@ func (t *Translator) sortArgs(cmd definitions.Command) []argEntry {
return sorted
}

func (t *Translator) processArg(name string, argDef definitions.Arg, input CommandInput, args *[]string) (string, error) {
// processArg appends the argument to args and returns the value, the index
// at which the value was appended (or -1 when nothing was appended or the
// value is not a standalone positional), and any validation error.
func (t *Translator) processArg(name string, argDef definitions.Arg, input CommandInput, args *[]string) (string, int, error) {
val, provided := input.Args[name]
if !provided {
if argDef.Required && !argDef.ExtractionOnly {
return "", ErrMissingArgument{Argument: name}
return "", -1, ErrMissingArgument{Argument: name}
}
return "", nil
return "", -1, nil
}

if argDef.ExtractionOnly {
return "", nil
return "", -1, nil
}

if argDef.Validate != "" {
if err := t.validate(argDef.Validate, val); err != nil {
return "", err
return "", -1, err
}
}

at := -1
switch {
case argDef.Flag != "":
*args = append(*args, argDef.Flag, val)
at = len(*args) - 1
case argDef.FixedSuffix != "":
*args = append(*args, val+argDef.FixedSuffix)
case argDef.Suffix != "" && name == "version":
// Handled in applyVersionSuffix
default:
*args = append(*args, val)
at = len(*args) - 1
}

return val, nil
return val, at, nil
}

func (t *Translator) applyVersionSuffix(args *[]string, cmd definitions.Command, input CommandInput, packageVal string) {
// applyVersionSuffix appends the version to the package argument in place.
// packageIdx is the index into args where the package positional was
// written; a negative value means no package positional was appended.
func (t *Translator) applyVersionSuffix(args *[]string, cmd definitions.Command, input CommandInput, packageIdx int) {
if packageIdx < 0 || packageIdx >= len(*args) {
return
}
versionDef, hasVersion := cmd.Args["version"]
if !hasVersion || versionDef.Suffix == "" {
return
Expand All @@ -200,12 +214,7 @@ func (t *Translator) applyVersionSuffix(args *[]string, cmd definitions.Command,
if !hasVersionVal {
return
}
for i, a := range *args {
if a == packageVal {
(*args)[i] = a + versionDef.Suffix + version
break
}
}
(*args)[packageIdx] += versionDef.Suffix + version
}

func (t *Translator) applyUserFlags(args *[]string, cmd definitions.Command, input CommandInput, baseOverrideUsed string) {
Expand Down
85 changes: 85 additions & 0 deletions translator_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -315,6 +315,77 @@ func TestCargoAdd(t *testing.T) {
}
}

func TestCargoAddVersion(t *testing.T) {
tr := loadTranslator(t)
cmd, err := tr.BuildCommand("cargo", "add", CommandInput{
Args: map[string]string{"package": "serde", "version": "1.0.219"},
})
if err != nil {
t.Fatalf("BuildCommand failed: %v", err)
}
expected := []string{"cargo", "add", "serde@1.0.219"}
if !reflect.DeepEqual(cmd, expected) {
t.Errorf("got %v, want %v", cmd, expected)
}
}

func TestVersionSuffixOnFlaggedPackage(t *testing.T) {
tr := NewTranslator()
tr.Register(&definitions.Definition{
Name: "flagpkg",
Binary: "tool",
Commands: map[string]definitions.Command{
"add": {
Base: []string{"add"},
Args: map[string]definitions.Arg{
"package": {Flag: "--package", Required: true},
"version": {Suffix: "@"},
},
},
},
})
got, err := tr.BuildCommand("flagpkg", "add", CommandInput{
Args: map[string]string{"package": "foo", "version": "1.0"},
})
if err != nil {
t.Fatalf("BuildCommand: %v", err)
}
want := []string{"tool", "add", "--package", "foo@1.0"}
if !reflect.DeepEqual(got, want) {
t.Errorf("got %v, want %v", got, want)
}
}

func TestVersionSuffixPackageNameCollision(t *testing.T) {
tr := loadTranslator(t)
cases := []struct {
manager string
pkg string
version string
want []string
}{
// crate name equal to the base subcommand
{"cargo", "add", "1.0", []string{"cargo", "add", "add@1.0"}},
// crate name equal to the binary
{"cargo", "cargo", "1.0", []string{"cargo", "add", "cargo@1.0"}},
// npm package name equal to the base subcommand
{"npm", "install", "2.0", []string{"npm", "install", "install@2.0"}},
// go module path equal to the base subcommand
{"gomod", "get", "v1.0.0", []string{"go", "get", "get@v1.0.0"}},
}
for _, tc := range cases {
got, err := tr.BuildCommand(tc.manager, "add", CommandInput{
Args: map[string]string{"package": tc.pkg, "version": tc.version},
})
if err != nil {
t.Fatalf("%s %s: %v", tc.manager, tc.pkg, err)
}
if !reflect.DeepEqual(got, tc.want) {
t.Errorf("%s %s: got %v, want %v", tc.manager, tc.pkg, got, tc.want)
}
}
}

func TestCargoAddDev(t *testing.T) {
tr := loadTranslator(t)
cmd, err := tr.BuildCommand("cargo", "add", CommandInput{
Expand Down Expand Up @@ -398,6 +469,20 @@ func TestGomodAdd(t *testing.T) {
}
}

func TestGomodAddVersion(t *testing.T) {
tr := loadTranslator(t)
cmd, err := tr.BuildCommand("gomod", "add", CommandInput{
Args: map[string]string{"package": "github.com/pkg/errors", "version": "v0.9.1"},
})
if err != nil {
t.Fatalf("BuildCommand failed: %v", err)
}
expected := []string{"go", "get", "github.com/pkg/errors@v0.9.1"}
if !reflect.DeepEqual(cmd, expected) {
t.Errorf("got %v, want %v", cmd, expected)
}
}

func TestGomodAddChain(t *testing.T) {
tr := loadTranslator(t)
cmds, err := tr.BuildCommands("gomod", "add", CommandInput{
Expand Down