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
44 changes: 36 additions & 8 deletions checks/cli.go
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,38 @@ import (

const maxCLIOutputBytesPerStream = 1024 * 1024

type commandShell struct {
path string
commandFlag string
}

func defaultShell() commandShell {
if runtime.GOOS == "windows" {
return commandShell{path: "powershell", commandFlag: "-Command"}
}
return commandShell{path: "sh", commandFlag: "-c"}
}

func resolveShell(name string) (commandShell, error) {
if name == "" {
return defaultShell(), nil
}
var flag string
switch name {
case "sh":
flag = "-c"
case "pwsh":
flag = "-Command"
default:
return commandShell{}, fmt.Errorf("unsupported shell %q: choose sh or pwsh", name)
}
path, err := exec.LookPath(name)
if err != nil {
return commandShell{}, fmt.Errorf("shell %q is unavailable: %w", name, err)
}
return commandShell{path: path, commandFlag: flag}, nil
}

type boundedBuffer struct {
buffer bytes.Buffer
limit int
Expand Down Expand Up @@ -44,25 +76,21 @@ func (b *boundedBuffer) String() string {
return b.buffer.String()
}

func runCLICommand(command api.CLIStepCLICommand, variables map[string]string) (result api.CLICommandResult) {
return runCLICommandWithOutputLimit(command, variables, maxCLIOutputBytesPerStream)
func runCLICommand(command api.CLIStepCLICommand, variables map[string]string, shell commandShell) (result api.CLICommandResult) {
return runCLICommandWithOutputLimit(command, variables, maxCLIOutputBytesPerStream, shell)
}

func runCLICommandWithOutputLimit(
command api.CLIStepCLICommand,
variables map[string]string,
maxOutputBytesPerStream int,
shell commandShell,
) (result api.CLICommandResult) {
finalCommand := InterpolateVariables(command.Command, variables)
result.FinalCommand = finalCommand
result.Command = command

var cmd *exec.Cmd
if runtime.GOOS == "windows" {
cmd = exec.Command("powershell", "-Command", finalCommand)
} else {
cmd = exec.Command("sh", "-c", finalCommand)
}
cmd := exec.Command(shell.path, shell.commandFlag, finalCommand)

cmd.Env = append(os.Environ(), "LANG=en_US.UTF-8")
stdout := newBoundedBuffer(maxOutputBytesPerStream)
Expand Down
14 changes: 6 additions & 8 deletions checks/cli_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@ func TestRunCLICommandCapsOutput(t *testing.T) {
},
variables,
4,
defaultShell(),
)

if !strings.Contains(result.Err, "per-stream limit") {
Expand All @@ -49,7 +50,7 @@ func TestRunCLICommandCapturesStdoutVariables(t *testing.T) {
Name: "goos",
Regex: `([a-z0-9]+)`,
}},
}, variables)
}, variables, defaultShell())

if result.Err != "" {
t.Fatalf("unexpected command error: %s", result.Err)
Expand Down Expand Up @@ -80,7 +81,7 @@ func TestRunCLICommandKeepsStderrSeparateFromStdoutChecks(t *testing.T) {
}},
}

result := runCLICommand(step, variables)
result := runCLICommand(step, variables, defaultShell())

if result.Stdout != "stdout-value" {
t.Fatalf("stdout = %q, want stdout-value", result.Stdout)
Expand Down Expand Up @@ -108,20 +109,20 @@ func TestRunCLICommandInterpolatesCapturedStdoutVariables(t *testing.T) {
Name: "goenv",
Regex: `"([A-Z]+)"`,
}},
}, variables)
}, variables, defaultShell())
if first.Err != "" {
t.Fatalf("unexpected first command error: %s", first.Err)
}

second := runCLICommand(api.CLIStepCLICommand{
Command: `go env ${goenv}`,
}, variables)
}, variables, defaultShell())
if second.Stdout != runtime.GOOS {
t.Fatalf("second stdout = %q, want %q", second.Stdout, runtime.GOOS)
}
}

func TestParseStdoutVariablesUsesGenericConfigurationError(t *testing.T) {
func TestParseStdoutVariablesRejectsInvalidConfiguration(t *testing.T) {
tests := []struct {
name string
vardef api.CLICommandStdoutVariable
Expand Down Expand Up @@ -151,9 +152,6 @@ func TestParseStdoutVariablesUsesGenericConfigurationError(t *testing.T) {
if err == nil {
t.Fatal("expected parse error")
}
if err.Error() != "invalid stdout variable configuration" {
t.Fatalf("error = %q, want invalid stdout variable configuration", err.Error())
}
})
}
}
123 changes: 92 additions & 31 deletions checks/local.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,11 +3,13 @@ package checks
import (
"fmt"
"math"
"math/big"
"reflect"
"strconv"
"strings"

api "github.com/bootdotdev/bootdev/client"
"github.com/goccy/go-json"
)

func LocalSubmissionEvent(cliData api.CLIData, results []api.CLIStepResult) api.LessonSubmissionEvent {
Expand Down Expand Up @@ -262,68 +264,127 @@ func evaluateStdoutJq(stdout string, test api.StdoutJqTest, variables map[string
if err != nil {
return err
}
if len(results) != len(test.ExpectedResults) {
return fmt.Errorf("expected jq query %q to return %d result(s), got %d", queryText, len(test.ExpectedResults), len(results))
if len(results) == 0 {
return fmt.Errorf("jq query returned no results")
}

for i, expected := range test.ExpectedResults {
want, err := jqExpectedValue(expected, variables)
if err != nil {
return err
outer:
for _, expected := range test.ExpectedResults {
if value, ok := expected.Value.(string); ok {
expected.Value = InterpolateVariables(value, variables)
}
if !compareValues(results[i], api.OperatorType(expected.Operator), want) {
return fmt.Errorf("expected jq result %d to be %s %v, got %v", i+1, expected.Operator, want, results[i])
for _, actual := range results {
if jqResultMatches(actual, expected) {
continue outer
}
}
return fmt.Errorf("expected jq results to contain %v", expected)
}

return nil
}

func jqExpectedValue(expected api.JqExpectedResult, variables map[string]string) (any, error) {
func jqResultMatches(actual any, expected api.JqExpectedResult) bool {
switch expected.Type {
case api.JqTypeString:
if str, ok := expected.Value.(string); ok {
return InterpolateVariables(str, variables), nil
}
return expected.Value, nil
got, gotOK := actual.(string)
want, wantOK := expected.Value.(string)
return gotOK && wantOK && expected.Operator == "==" && got == want
case api.JqTypeBool:
got, gotOK := coerceJqBool(actual)
want, wantOK := coerceJqBool(expected.Value)
return gotOK && wantOK && expected.Operator == "==" && got == want
case api.JqTypeInt:
if str, ok := expected.Value.(string); ok {
parsed, err := strconv.Atoi(InterpolateVariables(str, variables))
if err != nil {
return nil, err
}
return parsed, nil
got, gotOK := coerceJqInt(actual)
want, wantOK := coerceJqInt(expected.Value)
if !gotOK || !wantOK {
return false
}
return expected.Value, nil
case api.JqTypeBool:
if str, ok := expected.Value.(string); ok {
parsed, err := strconv.ParseBool(InterpolateVariables(str, variables))
if err != nil {
return nil, err
}
return parsed, nil
switch expected.Operator {
case "==":
return got == want
case ">":
return got > want
case ">=":
return got >= want
case "<":
return got < want
case "<=":
return got <= want
}
return expected.Value, nil
}
return false
}

func coerceJqBool(value any) (bool, bool) {
switch v := value.(type) {
case bool:
return v, true
case string:
parsed, err := strconv.ParseBool(v)
return parsed, err == nil
default:
return nil, fmt.Errorf("unsupported jq expected result type %q", expected.Type)
return false, false
}
}

func coerceJqInt(value any) (int, bool) {
switch v := value.(type) {
case int:
return v, true
case int64:
if v < math.MinInt || v > math.MaxInt {
return 0, false
}
return int(v), true
case float64:
// MaxInt rounds up as float64 on 64-bit hosts; use an exclusive upper bound.
if math.IsNaN(v) || math.Trunc(v) != v || v < float64(math.MinInt) || v >= -float64(math.MinInt) {
return 0, false
}
return int(v), true
case json.Number:
parsed, ok := new(big.Rat).SetString(v.String())
if !ok || !parsed.IsInt() || !parsed.Num().IsInt64() {
return 0, false
}
return coerceJqInt(parsed.Num().Int64())
case string:
parsed, err := strconv.Atoi(v)
return parsed, err == nil
default:
return 0, false
}
}

func compareValues(got any, operator api.OperatorType, want any) bool {
switch operator {
case api.OpEquals, "==":
return valuesEqual(got, want)
case api.OpGreaterThan, ">":
case api.OpGreaterThan, ">", ">=", "<", "<=":
gotNum, gotOK := numberValue(got)
wantNum, wantOK := numberValue(want)
return gotOK && wantOK && gotNum > wantNum
if !gotOK || !wantOK {
return false
}
switch operator {
case api.OpGreaterThan, ">":
return gotNum > wantNum
case ">=":
return gotNum >= wantNum
case "<":
return gotNum < wantNum
case "<=":
return gotNum <= wantNum
}
case api.OpContains:
return strings.Contains(fmt.Sprintf("%v", got), fmt.Sprintf("%v", want))
case api.OpNotContains:
return !strings.Contains(fmt.Sprintf("%v", got), fmt.Sprintf("%v", want))
default:
return false
}
return false
}

func valuesEqual(got any, want any) bool {
Expand Down
Loading
Loading