From e6de06517285738ae0888ce2ead8c52b8b8437e7 Mon Sep 17 00:00:00 2001 From: skovranek <59619403+skovranek@users.noreply.github.com> Date: Tue, 15 Sep 2026 17:21:23 -0400 Subject: [PATCH] Align local grading and assertion descriptions with backend behavior --- checks/cli.go | 13 +- checks/http.go | 33 +-- checks/jq.go | 11 +- checks/jq_test.go | 12 +- checks/local.go | 592 +++++++++++++++++++++++------------------- checks/local_test.go | 65 +---- checks/runner.go | 3 + checks/runner_test.go | 2 +- checks/validation.go | 116 +++++++++ 9 files changed, 479 insertions(+), 368 deletions(-) create mode 100644 checks/validation.go diff --git a/checks/cli.go b/checks/cli.go index f9a2e8d..b9fbf50 100644 --- a/checks/cli.go +++ b/checks/cli.go @@ -148,12 +148,13 @@ func parseStdoutVariables(stdout string, vardefs []api.CLICommandStdoutVariable, } func prettyPrintCLICommand(test api.CLICommandTest, variables map[string]string) string { + var descriptions []string if test.ExitCode != nil { - return fmt.Sprintf("Expect exit code %d", *test.ExitCode) + descriptions = append(descriptions, fmt.Sprintf("Expect exit code %d", *test.ExitCode)) } if test.StdoutLinesGT != nil { - return fmt.Sprintf("Expect > %d lines on stdout", *test.StdoutLinesGT) + descriptions = append(descriptions, fmt.Sprintf("Expect > %d lines on stdout", *test.StdoutLinesGT)) } if test.StdoutContainsAll != nil { @@ -163,7 +164,7 @@ func prettyPrintCLICommand(test api.CLICommandTest, variables map[string]string) interpolatedContains := InterpolateVariables(contains, variables) fmt.Fprintf(&str, "\n - '%s'", interpolatedContains) } - return str.String() + descriptions = append(descriptions, str.String()) } if test.StdoutContainsNone != nil { @@ -173,12 +174,12 @@ func prettyPrintCLICommand(test api.CLICommandTest, variables map[string]string) interpolatedContainsNone := InterpolateVariables(containsNone, variables) fmt.Fprintf(&str, "\n - '%s'", interpolatedContainsNone) } - return str.String() + descriptions = append(descriptions, str.String()) } if test.StdoutJq != nil { - return prettyPrintStdoutJqTest(*test.StdoutJq, variables) + descriptions = append(descriptions, prettyPrintStdoutJqTest(*test.StdoutJq, variables)) } - return "" + return strings.Join(descriptions, "\n") } diff --git a/checks/http.go b/checks/http.go index c994292..dba42bf 100644 --- a/checks/http.go +++ b/checks/http.go @@ -146,36 +146,27 @@ func interpolateJSONStrings(value any, variables map[string]string) any { } func prettyPrintHTTPTest(test api.HTTPRequestTest, variables map[string]string) string { + var descriptions []string if test.StatusCode != nil { - return fmt.Sprintf("Expecting status code: %d", *test.StatusCode) + descriptions = append(descriptions, fmt.Sprintf("Expecting status code: %d", *test.StatusCode)) } if test.BodyContains != nil { - interpolated := InterpolateVariables(*test.BodyContains, variables) - return fmt.Sprintf("Expecting response body to contain: %s", interpolated) + descriptions = append(descriptions, fmt.Sprintf("Expecting response body to contain: %s", *test.BodyContains)) } if test.BodyContainsNone != nil { - interpolated := InterpolateVariables(*test.BodyContainsNone, variables) - return fmt.Sprintf("Expecting response body to not contain: %s", interpolated) + descriptions = append(descriptions, fmt.Sprintf("Expecting response body to not contain: %s", *test.BodyContainsNone)) } if test.HeadersEqual != nil { - interpolatedKey := InterpolateVariables(test.HeadersEqual.Key, variables) - interpolatedValue := InterpolateVariables(test.HeadersEqual.Value, variables) - return fmt.Sprintf("Expecting header to equal: '%s: %v'", interpolatedKey, interpolatedValue) + descriptions = append(descriptions, fmt.Sprintf("Expecting header to equal: '%s: %v'", test.HeadersEqual.Key, test.HeadersEqual.Value)) } if test.HeadersContain != nil { - interpolatedKey := InterpolateVariables(test.HeadersContain.Key, variables) - interpolatedValue := InterpolateVariables(test.HeadersContain.Value, variables) - return fmt.Sprintf("Expecting header to contain: '%s: %v'", interpolatedKey, interpolatedValue) + descriptions = append(descriptions, fmt.Sprintf("Expecting header to contain: '%s: %v'", test.HeadersContain.Key, test.HeadersContain.Value)) } if test.TrailersEqual != nil { - interpolatedKey := InterpolateVariables(test.TrailersEqual.Key, variables) - interpolatedValue := InterpolateVariables(test.TrailersEqual.Value, variables) - return fmt.Sprintf("Expecting trailer to equal: '%s: %v'", interpolatedKey, interpolatedValue) + descriptions = append(descriptions, fmt.Sprintf("Expecting trailer to equal: '%s: %v'", test.TrailersEqual.Key, test.TrailersEqual.Value)) } if test.TrailersContain != nil { - interpolatedKey := InterpolateVariables(test.TrailersContain.Key, variables) - interpolatedValue := InterpolateVariables(test.TrailersContain.Value, variables) - return fmt.Sprintf("Expecting trailer to contain: '%s: %v'", interpolatedKey, interpolatedValue) + descriptions = append(descriptions, fmt.Sprintf("Expecting trailer to contain: '%s: %v'", test.TrailersContain.Key, test.TrailersContain.Value)) } if test.JSONValue != nil { var val any @@ -183,7 +174,7 @@ func prettyPrintHTTPTest(test api.HTTPRequestTest, variables map[string]string) case test.JSONValue.IntValue != nil: val = *test.JSONValue.IntValue case test.JSONValue.StringValue != nil: - val = *test.JSONValue.StringValue + val = InterpolateVariables(*test.JSONValue.StringValue, variables) case test.JSONValue.BoolValue != nil: val = *test.JSONValue.BoolValue } @@ -201,9 +192,9 @@ func prettyPrintHTTPTest(test api.HTTPRequestTest, variables map[string]string) } expecting := fmt.Sprintf("Expecting JSON at %v %s %v", test.JSONValue.Path, op, val) - return InterpolateVariables(expecting, variables) + descriptions = append(descriptions, expecting) } - return "" + return strings.Join(descriptions, "\n") } // Return a capped string representation of the response body. @@ -267,7 +258,7 @@ func parseVariables(body []byte, vardefs []api.HTTPRequestResponseVariable, vari func parseHeaderVariables(headers map[string]string, vardefs []api.HTTPRequestResponseHeaderVariable, variables map[string]string) error { for _, vardef := range vardefs { headerValue, ok := findHeaderValue(headers, vardef.Header) - if !ok || headerValue == "" { + if !ok { continue } diff --git a/checks/jq.go b/checks/jq.go index 0688b97..d9ef4e4 100644 --- a/checks/jq.go +++ b/checks/jq.go @@ -2,19 +2,19 @@ package checks import ( "bytes" + "encoding/json" "errors" "fmt" "io" "strings" api "github.com/bootdotdev/bootdev/client" - "github.com/goccy/go-json" "github.com/itchyny/gojq" "github.com/tailscale/hujson" ) func prettyPrintStdoutJqTest(test api.StdoutJqTest, variables map[string]string) string { - queryText := InterpolateVariables(test.Query, variables) + queryText := test.Query var str strings.Builder fmt.Fprintf(&str, "Expect jq query '%s' to yield values satisfying:", queryText) if len(test.ExpectedResults) == 0 { @@ -30,11 +30,6 @@ func prettyPrintStdoutJqTest(test api.StdoutJqTest, variables map[string]string) func formatJqExpectedValue(expected api.JqExpectedResult, variables map[string]string) string { value := expected.Value - if expected.Type == api.JqTypeString { - if stringValue, ok := expected.Value.(string); ok { - value = InterpolateVariables(stringValue, variables) - } - } encoded, err := json.Marshal(value) if err != nil { return fmt.Sprintf("%v", value) @@ -54,7 +49,7 @@ func collectStdoutJqOutputs(cmd api.CLIStepCLICommand, result api.CLICommandResu } func runStdoutJqQuery(stdout string, test api.StdoutJqTest, variables map[string]string) api.CLICommandJqOutput { - queryText := InterpolateVariables(test.Query, variables) + queryText := test.Query input, err := parseJqInput(stdout, test.InputMode) if err != nil { return api.CLICommandJqOutput{Query: queryText, Error: err.Error()} diff --git a/checks/jq_test.go b/checks/jq_test.go index 57e6731..2fedd06 100644 --- a/checks/jq_test.go +++ b/checks/jq_test.go @@ -17,19 +17,19 @@ func TestRunStdoutJqQuery(t *testing.T) { wantError bool }{ { - name: "queries json with interpolated query", + name: "queries JSON with comments using a literal query", stdout: `{ - // Users to query - "users": [/* users */ {"name":"Lane"},{"name":"Theo",},], - }`, + // Users to query + "users": [/* users */ {"name":"Lane"},{"name":"Theo",},], + }`, test: api.StdoutJqTest{ InputMode: "json", Query: `.users[] | select(.name == "${name}") | .name`, }, variables: map[string]string{"name": "Theo"}, want: api.CLICommandJqOutput{ - Query: `.users[] | select(.name == "Theo") | .name`, - Results: []string{`"Theo"`}, + Query: `.users[] | select(.name == "${name}") | .name`, + Results: nil, }, }, { diff --git a/checks/local.go b/checks/local.go index 650c344..e76e2f9 100644 --- a/checks/local.go +++ b/checks/local.go @@ -1,17 +1,19 @@ package checks import ( + "encoding/json" + "errors" "fmt" "math" "math/big" - "reflect" + "regexp" "strconv" "strings" api "github.com/bootdotdev/bootdev/client" - "github.com/goccy/go-json" ) +// Local grading mirrors the backend; success is represented by nil. func LocalSubmissionEvent(cliData api.CLIData, results []api.CLIStepResult) api.LessonSubmissionEvent { failure := EvaluateCLIResults(cliData, results) slug := api.VerificationResultSlugSuccess @@ -32,400 +34,442 @@ func LocalSubmissionEvent(cliData api.CLIData, results []api.CLIStepResult) api. } func EvaluateCLIResults(cliData api.CLIData, results []api.CLIStepResult) *api.StructuredErrCLI { - for stepIndex, step := range cliData.Steps { - if stepIndex >= len(results) { - return localFailure(stepIndex, 0, "missing result for step") - } + if len(cliData.Steps) != len(results) { + return localFailure(-1, -1, "wrong number of steps") + } - switch { - case step.CLICommand != nil: - result := results[stepIndex].CLICommandResult - if result == nil { - return localFailure(stepIndex, 0, "missing CLI command result") - } - if failure := evaluateCLICommandTests(stepIndex, *step.CLICommand, *result); failure != nil { - return failure - } - case step.HTTPRequest != nil: - result := results[stepIndex].HTTPRequestResult - if result == nil { - return localFailure(stepIndex, 0, "missing HTTP request result") + for i, step := range cliData.Steps { + actual := results[i] + + if step.CLICommand != nil && actual.CLICommandResult != nil { + verificationErr := evaluateCLICommandTests(i, *step.CLICommand, *actual.CLICommandResult) + if verificationErr != nil { + return verificationErr } - if failure := evaluateHTTPRequestTests(stepIndex, *step.HTTPRequest, *result); failure != nil { - return failure + } else if step.HTTPRequest != nil && actual.HTTPRequestResult != nil { + verificationErr := evaluateHTTPRequestTests(i, *step.HTTPRequest, *actual.HTTPRequestResult) + if verificationErr != nil { + return verificationErr } - default: - return localFailure(stepIndex, 0, "missing step definition") + } else { + return localFailure(-1, -1, "invalid step") } } return nil } -func evaluateCLICommandTests(stepIndex int, cmd api.CLIStepCLICommand, result api.CLICommandResult) *api.StructuredErrCLI { - if result.Err != "" { - return localFailure(stepIndex, 0, result.Err) +func evaluateCLICommandTests(stepIndex int, expect api.CLIStepCLICommand, actual api.CLICommandResult) *api.StructuredErrCLI { + if err := validateCommandAssertions(expect); err != nil { + return &api.StructuredErrCLI{ErrorMessage: err.Error(), FailedStepIndex: stepIndex, FailedTestIndex: -1} + } + if actual.ExitCode < 0 { + return localFailure(stepIndex, -1, "failed to start command") } - for testIndex, test := range cmd.Tests { - var err error - - switch { - case test.ExitCode != nil: - if result.ExitCode != *test.ExitCode { - err = fmt.Errorf("expected exit code %d, got %d", *test.ExitCode, result.ExitCode) + for i, expectedTest := range expect.Tests { + if expectedTest.ExitCode != nil { + if *expectedTest.ExitCode != actual.ExitCode { + return localFailure(stepIndex, i, fmt.Sprintf("expected status code %v, got %v", *expectedTest.ExitCode, actual.ExitCode)) } - case len(test.StdoutContainsAll) > 0: - for _, contains := range test.StdoutContainsAll { - needle := InterpolateVariables(contains, result.Variables) - if !strings.Contains(result.Stdout, needle) { - err = fmt.Errorf("expected stdout to contain %q", needle) - break - } + } + if expectedTest.StdoutJq != nil { + jqInput, err := parseJqInput(actual.Stdout, expectedTest.StdoutJq.InputMode) + if err != nil { + return localFailure(stepIndex, i, fmt.Sprintf("failed to read jq input: %v", err)) } - case len(test.StdoutContainsNone) > 0: - for _, containsNone := range test.StdoutContainsNone { - needle := InterpolateVariables(containsNone, result.Variables) - if strings.Contains(result.Stdout, needle) { - err = fmt.Errorf("expected stdout to not contain %q", needle) - break + jqResults, err := executeJqQuery(expectedTest.StdoutJq.Query, jqInput) + if err != nil { + return localFailure(stepIndex, i, fmt.Sprintf("failed to run jq query: %v", err)) + } + if len(jqResults) == 0 { + return localFailure(stepIndex, i, "jq query returned no results") + } + outer: + for _, expectedResult := range expectedTest.StdoutJq.ExpectedResults { + for _, actualResult := range jqResults { + if jqResultMatches(actualResult, expectedResult) { + continue outer + } } + return localFailure(stepIndex, i, fmt.Sprintf("expected jq results to contain %v", expectedResult)) } - case test.StdoutLinesGT != nil: - lineCount := stdoutLineCount(result.Stdout) - if lineCount <= *test.StdoutLinesGT { - err = fmt.Errorf("expected stdout to have more than %d lines, got %d", *test.StdoutLinesGT, lineCount) + } + if expectedTest.StdoutLinesGT != nil { + count := strings.Count(actual.Stdout, "\n") + if actual.Stdout != "" { + count++ + } + if count <= *expectedTest.StdoutLinesGT { + return localFailure(stepIndex, i, fmt.Sprintf("expected more than %v lines, got %v", *expectedTest.StdoutLinesGT, count)) } - case test.StdoutJq != nil: - err = evaluateStdoutJq(result.Stdout, *test.StdoutJq, result.Variables) - default: - err = fmt.Errorf("unsupported CLI command test") } - - if err != nil { - return localFailure(stepIndex, testIndex, err.Error()) + if expectedTest.StdoutContainsAll != nil { + for _, expectedContains := range expectedTest.StdoutContainsAll { + interpolatedContains := InterpolateVariables(expectedContains, actual.Variables) + if !strings.Contains(actual.Stdout, interpolatedContains) { + return localFailure(stepIndex, i, fmt.Sprintf("expected stdout to contain %v", interpolatedContains)) + } + } + } + if expectedTest.StdoutContainsNone != nil { + for _, expectedContainsNone := range expectedTest.StdoutContainsNone { + interpolatedContainsNone := InterpolateVariables(expectedContainsNone, actual.Variables) + if strings.Contains(actual.Stdout, interpolatedContainsNone) { + return localFailure(stepIndex, i, fmt.Sprintf("expected stdout to not contain %v", interpolatedContainsNone)) + } + } } } return nil } -func evaluateHTTPRequestTests(stepIndex int, req api.CLIStepHTTPRequest, result api.HTTPRequestResult) *api.StructuredErrCLI { - if result.Err != "" { - return localFailure(stepIndex, 0, result.Err) +func evaluateHTTPRequestTests(stepIndex int, expect api.CLIStepHTTPRequest, actual api.HTTPRequestResult) *api.StructuredErrCLI { + if err := validateHTTPAssertions(expect); err != nil { + return &api.StructuredErrCLI{ErrorMessage: err.Error(), FailedStepIndex: stepIndex, FailedTestIndex: -1} + } + if actual.Err != "" { + return localFailure(stepIndex, -1, fmt.Sprintf("fetch error: %v", actual.Err)) } - for testIndex, test := range req.Tests { - var err error + for i, expectedTest := range expect.Tests { + if expectedTest.StatusCode != nil { + if *expectedTest.StatusCode != actual.StatusCode { + return localFailure(stepIndex, i, fmt.Sprintf("expected status code %v, got %v", *expectedTest.StatusCode, actual.StatusCode)) + } + } - switch { - case test.StatusCode != nil: - if result.StatusCode != *test.StatusCode { - err = fmt.Errorf("expected status code %d, got %d", *test.StatusCode, result.StatusCode) + if expectedTest.BodyContains != nil { + if !strings.Contains(actual.BodyString, *expectedTest.BodyContains) { + return localFailure(stepIndex, i, fmt.Sprintf("expected response body to contain '%v', but it did not", *expectedTest.BodyContains)) } - case test.BodyContains != nil: - needle := InterpolateVariables(*test.BodyContains, result.Variables) - if !strings.Contains(result.BodyString, needle) { - err = fmt.Errorf("expected response body to contain %q", needle) + } + + if expectedTest.BodyContainsNone != nil { + if strings.Contains(actual.BodyString, *expectedTest.BodyContainsNone) { + return localFailure(stepIndex, i, fmt.Sprintf("expected response body to not contain '%v', but it did", *expectedTest.BodyContainsNone)) } - case test.BodyContainsNone != nil: - needle := InterpolateVariables(*test.BodyContainsNone, result.Variables) - if strings.Contains(result.BodyString, needle) { - err = fmt.Errorf("expected response body to not contain %q", needle) + } + + if expectedTest.HeadersEqual != nil { + actualHeaderValue, ok := findHeaderValue(actual.ResponseHeaders, expectedTest.HeadersEqual.Key) + if !ok || actualHeaderValue != expectedTest.HeadersEqual.Value { + return localFailure(stepIndex, i, fmt.Sprintf("expected '%v' header to equal '%v', but it did not", expectedTest.HeadersEqual.Key, expectedTest.HeadersEqual.Value)) } - case test.HeadersEqual != nil: - err = evaluateHeaderEquals(result.ResponseHeaders, *test.HeadersEqual, result.Variables, "header") - case test.HeadersContain != nil: - err = evaluateHeaderContains(result.ResponseHeaders, *test.HeadersContain, result.Variables, "header") - case test.TrailersEqual != nil: - err = evaluateHeaderEquals(result.ResponseTrailers, *test.TrailersEqual, result.Variables, "trailer") - case test.TrailersContain != nil: - err = evaluateHeaderContains(result.ResponseTrailers, *test.TrailersContain, result.Variables, "trailer") - case test.JSONValue != nil: - err = evaluateHTTPJSONValue(result.BodyString, *test.JSONValue, result.Variables) - default: - err = fmt.Errorf("unsupported HTTP request test") } - if err != nil { - return localFailure(stepIndex, testIndex, err.Error()) + if expectedTest.HeadersContain != nil { + actualHeaderValue, ok := findHeaderValue(actual.ResponseHeaders, expectedTest.HeadersContain.Key) + if !ok || !strings.Contains(actualHeaderValue, expectedTest.HeadersContain.Value) { + return localFailure(stepIndex, i, fmt.Sprintf("expected '%v' header to contain '%v', but it did not", expectedTest.HeadersContain.Key, expectedTest.HeadersContain.Value)) + } } - } - captureIndex := len(req.Tests) - for _, vardef := range req.ResponseVariables { - expected := map[string]string{} - if err := parseVariables([]byte(result.BodyString), []api.HTTPRequestResponseVariable{vardef}, expected); err != nil { - return localFailure(stepIndex, captureIndex, err.Error()) + if expectedTest.TrailersEqual != nil { + actualTrailerValue, ok := findHeaderValue(actual.ResponseTrailers, expectedTest.TrailersEqual.Key) + if !ok || actualTrailerValue != expectedTest.TrailersEqual.Value { + return localFailure(stepIndex, i, fmt.Sprintf("expected '%v' trailer to equal '%v', but it did not", expectedTest.TrailersEqual.Key, expectedTest.TrailersEqual.Value)) + } } - want, found := expected[vardef.Name] - if !found { - return localFailure(stepIndex, captureIndex, fmt.Sprintf("missing value for response variable %q", vardef.Name)) + if expectedTest.TrailersContain != nil { + actualTrailerValue, ok := findHeaderValue(actual.ResponseTrailers, expectedTest.TrailersContain.Key) + if !ok || !strings.Contains(actualTrailerValue, expectedTest.TrailersContain.Value) { + return localFailure(stepIndex, i, fmt.Sprintf("expected '%v' trailer to contain '%v', but it did not", expectedTest.TrailersContain.Key, expectedTest.TrailersContain.Value)) + } } - got, captured := result.Variables[vardef.Name] - if !captured || got != want { - return localFailure(stepIndex, captureIndex, fmt.Sprintf("captured response variable %q did not match the response body", vardef.Name)) + + if expectedTest.JSONValue != nil { + err := jsonValOp(*expectedTest.JSONValue, actual.BodyString, actual.Variables) + if err != nil { + return localFailure(stepIndex, i, fmt.Sprintf("%v", err)) + } } } - if len(req.ResponseVariables) > 0 { - captureIndex++ + responseVariableTestIndex := len(expect.Tests) + 1 + responseHeaderVariableTestIndex := responseVariableTestIndex + if len(expect.ResponseVariables) > 0 { + responseHeaderVariableTestIndex++ } - for _, vardef := range req.ResponseHeaderVariables { - expected := map[string]string{} - if err := parseHeaderVariables(result.ResponseHeaders, []api.HTTPRequestResponseHeaderVariable{vardef}, expected); err != nil { - return localFailure(stepIndex, captureIndex, err.Error()) - } - want, found := expected[vardef.Name] - if !found { - return localFailure(stepIndex, captureIndex, fmt.Sprintf("missing value for response header variable %q", vardef.Name)) + for _, expectedVar := range expect.ResponseVariables { + expectedValue, ok := responseVariableValue(expectedVar, actual.BodyString) + if !ok { + return localFailure(stepIndex, responseVariableTestIndex, fmt.Sprintf("missing value for variable '%s'", expectedVar.Name)) } - got, captured := result.Variables[vardef.Name] - if !captured || got != want { - return localFailure(stepIndex, captureIndex, fmt.Sprintf("captured response header variable %q did not match the response header", vardef.Name)) + + if !capturedVariableMatches(actual.Variables, expectedVar.Name, expectedValue) { + return localFailure(stepIndex, responseVariableTestIndex, fmt.Sprintf("captured variable '%s' did not match expected response body value", expectedVar.Name)) } } - return nil -} - -func evaluateHeaderEquals(headers map[string]string, test api.HTTPRequestTestHeader, variables map[string]string, label string) error { - key := InterpolateVariables(test.Key, variables) - want := InterpolateVariables(test.Value, variables) + for _, expectedVar := range expect.ResponseHeaderVariables { + expectedValue, ok := responseHeaderVariableValue(expectedVar, actual.ResponseHeaders) + if !ok { + return localFailure(stepIndex, responseHeaderVariableTestIndex, fmt.Sprintf("missing value for variable '%s'", expectedVar.Name)) + } - got, ok := findHeaderValue(headers, key) - if !ok { - return fmt.Errorf("expected %s %q to exist", label, key) - } - if got != want { - return fmt.Errorf("expected %s %q to equal %q, got %q", label, key, want, got) + if !capturedVariableMatches(actual.Variables, expectedVar.Name, expectedValue) { + return localFailure(stepIndex, responseHeaderVariableTestIndex, fmt.Sprintf("captured variable '%s' did not match expected response header value", expectedVar.Name)) + } } return nil } -func evaluateHeaderContains(headers map[string]string, test api.HTTPRequestTestHeader, variables map[string]string, label string) error { - key := InterpolateVariables(test.Key, variables) - want := InterpolateVariables(test.Value, variables) - - got, ok := findHeaderValue(headers, key) - if !ok { - return fmt.Errorf("expected %s %q to exist", label, key) - } - if !strings.Contains(strings.ToLower(got), strings.ToLower(want)) { - return fmt.Errorf("expected %s %q to contain %q, got %q", label, key, want, got) - } - - return nil +func capturedVariableMatches(vars map[string]string, name, expectedValue string) bool { + actualValue, ok := vars[name] + return ok && actualValue == expectedValue } -func evaluateHTTPJSONValue(body string, test api.HTTPRequestTestJSONValue, variables map[string]string) error { - got, err := valFromJqPath(test.Path, body) - if err != nil { - return err +func responseVariableValue(expectedVar api.HTTPRequestResponseVariable, body string) (string, bool) { + if expectedVar.Path != "" { + val, err := valFromJqPath(expectedVar.Path, body) + if err != nil || val == nil { + return "", false + } + return fmt.Sprintf("%v", val), true } - want, err := httpJSONExpectedValue(test, variables) + re, err := regexp.Compile(expectedVar.BodyRegex) if err != nil { - return err + return "", false } - if !compareValues(got, test.Operator, want) { - return fmt.Errorf("expected JSON at %s %s %v, got %v", test.Path, test.Operator, want, got) + matches := re.FindStringSubmatch(body) + if len(matches) != 2 { + return "", false } - return nil + return matches[1], true } -func httpJSONExpectedValue(test api.HTTPRequestTestJSONValue, variables map[string]string) (any, error) { - switch { - case test.IntValue != nil: - return *test.IntValue, nil - case test.StringValue != nil: - return InterpolateVariables(*test.StringValue, variables), nil - case test.BoolValue != nil: - return *test.BoolValue, nil - default: - return nil, fmt.Errorf("missing expected JSON value") +func responseHeaderVariableValue(expectedVar api.HTTPRequestResponseHeaderVariable, headers map[string]string) (string, bool) { + headerValue, ok := findHeaderValue(headers, expectedVar.Header) + if !ok { + return "", false } -} -func evaluateStdoutJq(stdout string, test api.StdoutJqTest, variables map[string]string) error { - queryText := InterpolateVariables(test.Query, variables) + if expectedVar.Regex == "" { + return headerValue, true + } - input, err := parseJqInput(stdout, test.InputMode) + re, err := regexp.Compile(expectedVar.Regex) if err != nil { - return err + return "", false + } + + matches := re.FindStringSubmatch(headerValue) + if len(matches) != 2 { + return "", false } - results, err := executeJqQuery(queryText, input) + return matches[1], true +} + +func jsonValOp(test api.HTTPRequestTestJSONValue, jsn string, variables map[string]string) error { + val, err := valFromJqPath(test.Path, jsn) if err != nil { return err } - if len(results) == 0 { - return fmt.Errorf("jq query returned no results") + if test.BoolValue != nil { + vBool, ok := val.(bool) + if !ok { + return errors.New("expected boolean value") + } + if test.Operator == api.OpEquals { + if vBool != *test.BoolValue { + return errors.New("boolean value not equal") + } + return nil + } + return errors.New("operator not supported") } - -outer: - for _, expected := range test.ExpectedResults { - if value, ok := expected.Value.(string); ok { - expected.Value = InterpolateVariables(value, variables) + if test.IntValue != nil { + var v int + vInt, intOk := val.(int) + vFloat, floatOk := val.(float64) + switch { + case intOk: + v = vInt + case floatOk: + v = int(vFloat) + default: + return errors.New("expected int value") + } + if test.Operator == api.OpEquals { + if v != *test.IntValue { + return errors.New("int value not equal") + } + return nil + } + if test.Operator == api.OpGreaterThan { + if v <= *test.IntValue { + return errors.New("int value not greater than") + } + return nil } - for _, actual := range results { - if jqResultMatches(actual, expected) { - continue outer + return errors.New("operator not supported") + } + if test.StringValue != nil { + vStr, ok := val.(string) + if !ok { + return errors.New("expected string value") + } + if test.Operator == api.OpEquals { + interpolatedStr := InterpolateVariables(*test.StringValue, variables) + if vStr != interpolatedStr { + return errors.New("string value not equal") + } + return nil + } + if test.Operator == api.OpContains { + interpolatedStr := InterpolateVariables(*test.StringValue, variables) + if !strings.Contains(vStr, interpolatedStr) { + return fmt.Errorf("%s does not contain %s", vStr, interpolatedStr) } + return nil } - return fmt.Errorf("expected jq results to contain %v", expected) + if test.Operator == api.OpNotContains { + interpolatedStr := InterpolateVariables(*test.StringValue, variables) + if strings.Contains(vStr, interpolatedStr) { + return fmt.Errorf("%s contains %s", vStr, interpolatedStr) + } + return nil + } + return errors.New("operator not supported") } - return nil + return errors.New("no test value provided") } -func jqResultMatches(actual any, expected api.JqExpectedResult) bool { - switch expected.Type { - case api.JqTypeString: - got, gotOK := actual.(string) - want, wantOK := expected.Value.(string) - return gotOK && wantOK && expected.Operator == "==" && got == want +func jqResultMatches(actualResult any, expectedResult api.JqExpectedResult) bool { + switch expectedResult.Type { case api.JqTypeBool: - got, gotOK := coerceJqBool(actual) - want, wantOK := coerceJqBool(expected.Value) - return gotOK && wantOK && expected.Operator == "==" && got == want - case api.JqTypeInt: - got, gotOK := coerceJqInt(actual) - want, wantOK := coerceJqInt(expected.Value) - if !gotOK || !wantOK { + expected, expectedOk := coerceBool(expectedResult.Value) + actual, actualOk := coerceBool(actualResult) + if !expectedOk || !actualOk { + return false + } + return compareBool(actual, expected, expectedResult.Operator) + case api.JqTypeString: + expected, expectedOk := coerceString(expectedResult.Value) + actual, actualOk := coerceString(actualResult) + if !expectedOk || !actualOk { return false } - switch expected.Operator { - case "==": - return got == want - case ">": - return got > want - case ">=": - return got >= want - case "<": - return got < want - case "<=": - return got <= want + return compareString(actual, expected, expectedResult.Operator) + case api.JqTypeInt: + expected, expectedOk := coerceInt(expectedResult.Value) + actual, actualOk := coerceInt(actualResult) + if !expectedOk || !actualOk { + return false } + return compareInt(actual, expected, expectedResult.Operator) + default: + return false } - return false } -func coerceJqBool(value any) (bool, bool) { - switch v := value.(type) { +func coerceBool(value any) (bool, bool) { + switch typed := value.(type) { case bool: - return v, true + return typed, true case string: - parsed, err := strconv.ParseBool(v) - return parsed, err == nil + parsed, err := strconv.ParseBool(typed) + if err != nil { + return false, false + } + return parsed, true default: return false, false } } -func coerceJqInt(value any) (int, bool) { - switch v := value.(type) { +func coerceString(value any) (string, bool) { + switch typed := value.(type) { + case string: + return typed, true + default: + return "", false + } +} + +func coerceInt(value any) (int, bool) { + switch typed := value.(type) { case int: - return v, true + return typed, true case int64: - if v < math.MinInt || v > math.MaxInt { + if typed > math.MaxInt || typed < math.MinInt { return 0, false } - return int(v), true + return int(typed), true case float64: + if math.IsNaN(typed) || math.IsInf(typed, 0) { + return 0, false + } + if math.Trunc(typed) != typed { + return 0, false + } // 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) { + if typed >= -float64(math.MinInt) || typed < float64(math.MinInt) { return 0, false } - return int(v), true + return int(typed), true case json.Number: - parsed, ok := new(big.Rat).SetString(v.String()) + parsed, ok := new(big.Rat).SetString(typed.String()) if !ok || !parsed.IsInt() || !parsed.Num().IsInt64() { return 0, false } - return coerceJqInt(parsed.Num().Int64()) + return coerceInt(parsed.Num().Int64()) case string: - parsed, err := strconv.Atoi(v) - return parsed, err == nil + parsed, err := strconv.Atoi(typed) + if err != nil { + return 0, false + } + return parsed, true default: return 0, false } } -func compareValues(got any, operator api.OperatorType, want any) bool { +func compareBool(actual bool, expected bool, operator api.JqOperator) bool { switch operator { - case api.OpEquals, "==": - return valuesEqual(got, want) - case api.OpGreaterThan, ">", ">=", "<", "<=": - gotNum, gotOK := numberValue(got) - wantNum, wantOK := numberValue(want) - 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)) + case "==": + return actual == expected default: return false } - return false } -func valuesEqual(got any, want any) bool { - if gotNum, gotOK := numberValue(got); gotOK { - wantNum, wantOK := numberValue(want) - return wantOK && math.Abs(gotNum-wantNum) < 0.000000001 - } - return reflect.DeepEqual(got, want) -} - -func numberValue(value any) (float64, bool) { - switch v := value.(type) { - case int: - return float64(v), true - case int64: - return float64(v), true - case float64: - return v, true - case jsonNumber: - parsed, err := strconv.ParseFloat(v.String(), 64) - return parsed, err == nil +func compareString(actual string, expected string, operator api.JqOperator) bool { + switch operator { + case "==": + return actual == expected default: - return 0, false - } -} - -func stdoutLineCount(stdout string) int { - if stdout == "" { - return 0 + return false } - return strings.Count(stdout, "\n") + 1 } -func localFailure(stepIndex int, testIndex int, message string) *api.StructuredErrCLI { - return &api.StructuredErrCLI{ - ErrorMessage: message, - FailedStepIndex: stepIndex, - FailedTestIndex: testIndex, +func compareInt(actual int, expected int, operator api.JqOperator) bool { + switch operator { + case "==": + return actual == expected + case ">": + return actual > expected + case ">=": + return actual >= expected + case "<": + return actual < expected + case "<=": + return actual <= expected + default: + return false } } -type jsonNumber interface { - String() string +func localFailure(stepIndex, testIndex int, message string) *api.StructuredErrCLI { + return &api.StructuredErrCLI{ErrorMessage: message, FailedStepIndex: stepIndex, FailedTestIndex: testIndex} } diff --git a/checks/local_test.go b/checks/local_test.go index a869d12..074b80f 100644 --- a/checks/local_test.go +++ b/checks/local_test.go @@ -1,12 +1,12 @@ package checks import ( + "encoding/json" "math" "strconv" "testing" api "github.com/bootdotdev/bootdev/client" - "github.com/goccy/go-json" ) func TestLocalSubmissionEventPassesCLIAndHTTPResults(t *testing.T) { @@ -77,22 +77,6 @@ func TestLocalSubmissionEventReportsFirstFailure(t *testing.T) { } } -func TestEvaluateCLICommandReportsExecutionError(t *testing.T) { - const message = "invalid stdout variable configuration" - failure := evaluateCLICommandTests( - 0, - api.CLIStepCLICommand{}, - api.CLICommandResult{Err: message}, - ) - - if failure == nil { - t.Fatal("expected structured failure") - } - if failure.ErrorMessage != message { - t.Fatalf("ErrorMessage = %q, want %q", failure.ErrorMessage, message) - } -} - func TestEvaluateStdoutJqNumericComparisons(t *testing.T) { for _, tt := range []struct { operator api.JqOperator @@ -106,7 +90,7 @@ func TestEvaluateStdoutJqNumericComparisons(t *testing.T) { } { for i, stdout := range []string{"4", "5", "6"} { t.Run(stdout+string(tt.operator)+"5", func(t *testing.T) { - err := evaluateStdoutJq(stdout, api.StdoutJqTest{ + jqTest := api.StdoutJqTest{ InputMode: "json", Query: ".", ExpectedResults: []api.JqExpectedResult{{ @@ -114,7 +98,8 @@ func TestEvaluateStdoutJqNumericComparisons(t *testing.T) { Operator: tt.operator, Value: 5, }}, - }, nil) + } + err := evaluateCLICommandTests(0, api.CLIStepCLICommand{Tests: []api.CLICommandTest{{StdoutJq: &jqTest}}}, api.CLICommandResult{Stdout: stdout}) if (err == nil) != tt.pass[i] { t.Fatalf("comparison passed = %t, want %t; error: %v", err == nil, tt.pass[i], err) } @@ -232,8 +217,8 @@ func TestLocalSubmissionEventRejectsMissingHTTPResponseCaptures(t *testing.T) { if event.StructuredErrCLI == nil { t.Fatal("expected structured failure") } - if event.StructuredErrCLI.FailedStepIndex != 0 || event.StructuredErrCLI.FailedTestIndex != 1 { - t.Fatalf("failure = %#v, want step 0 capture test 1", event.StructuredErrCLI) + if event.StructuredErrCLI.FailedStepIndex != 0 || event.StructuredErrCLI.FailedTestIndex != 2 { + t.Fatalf("failure = %#v, want step 0 capture test 2", event.StructuredErrCLI) } }) } @@ -260,31 +245,6 @@ func TestLocalSubmissionEventAcceptsEmptyHTTPResponseCapture(t *testing.T) { } } -func TestValuesEqualPreservesTypes(t *testing.T) { - tests := []struct { - name string - got any - want any - ok bool - }{ - {name: "same strings", got: "1", want: "1", ok: true}, - {name: "string and int", got: "1", want: 1, ok: false}, - {name: "string and bool", got: "true", want: true, ok: false}, - {name: "same bools", got: true, want: true, ok: true}, - {name: "numeric int and float", got: 1, want: 1.0, ok: true}, - {name: "numeric json number and int", got: json.Number("1"), want: 1, ok: true}, - {name: "nil and string", got: nil, want: "", ok: false}, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - if got := valuesEqual(tt.got, tt.want); got != tt.ok { - t.Fatalf("valuesEqual(%#v, %#v) = %v, want %v", tt.got, tt.want, got, tt.ok) - } - }) - } -} - func intPtr(v int) *int { return &v } @@ -305,7 +265,7 @@ func TestEvaluateStdoutJqMatchesAnyResult(t *testing.T) { {"missing expected result", "[1, 3]", []int{1, 2}, false}, {"empty results", "[]", []int{1}, false}, {"empty results without expectations", "[]", nil, false}, - {"nonempty results without expectations", "[1]", nil, true}, + {"nonempty results without expectations", "[1]", nil, false}, } { t.Run(tt.name, func(t *testing.T) { test := api.StdoutJqTest{InputMode: "json", Query: ".[]"} @@ -314,7 +274,7 @@ func TestEvaluateStdoutJqMatchesAnyResult(t *testing.T) { Type: api.JqTypeInt, Operator: "==", Value: value, }) } - err := evaluateStdoutJq(tt.stdout, test, nil) + err := evaluateCLICommandTests(0, api.CLIStepCLICommand{Tests: []api.CLICommandTest{{StdoutJq: &test}}}, api.CLICommandResult{Stdout: tt.stdout}) if (err == nil) != tt.pass { t.Fatalf("passed = %t, want %t; error: %v", err == nil, tt.pass, err) } @@ -332,7 +292,7 @@ func TestEvaluateStdoutJqResultTypes(t *testing.T) { pass bool }{ {"numeric string", `"5"`, api.JqTypeInt, "==", 5, true}, - {"interpolated integer", "5", api.JqTypeInt, "==", "${value}", true}, + {"integer expectation is literal", "5", api.JqTypeInt, "==", "${value}", false}, {"integral expected float", "5", api.JqTypeInt, "==", 5.0, true}, {"decimal JSON number", "5.0", api.JqTypeInt, "==", 5, true}, {"exponent JSON number", "5e0", api.JqTypeInt, "==", 5, true}, @@ -346,20 +306,21 @@ func TestEvaluateStdoutJqResultTypes(t *testing.T) { {"boolean", "true", api.JqTypeBool, "==", true, true}, {"boolean strings", `"true"`, api.JqTypeBool, "==", "true", true}, {"invalid boolean", `"yes"`, api.JqTypeBool, "==", true, false}, - {"interpolated string", `"5"`, api.JqTypeString, "==", "${value}", true}, + {"string expectation is literal", `"5"`, api.JqTypeString, "==", "${value}", false}, {"string type rejects numbers", "5", api.JqTypeString, "==", 5, false}, {"boolean type rejects numbers", "1", api.JqTypeBool, "==", 1, false}, {"string ordering unsupported", `"b"`, api.JqTypeString, ">", "a", false}, {"unknown operator", "5", api.JqTypeInt, "!=", 4, false}, } { t.Run(tt.name, func(t *testing.T) { - err := evaluateStdoutJq(tt.stdout, api.StdoutJqTest{ + jqTest := api.StdoutJqTest{ InputMode: "json", Query: ".", ExpectedResults: []api.JqExpectedResult{{ Type: tt.kind, Operator: tt.operator, Value: tt.want, }}, - }, map[string]string{"value": "5"}) + } + err := evaluateCLICommandTests(0, api.CLIStepCLICommand{Tests: []api.CLICommandTest{{StdoutJq: &jqTest}}}, api.CLICommandResult{Stdout: tt.stdout, Variables: map[string]string{"value": "5"}}) if (err == nil) != tt.pass { t.Fatalf("passed = %t, want %t; error: %v", err == nil, tt.pass, err) } diff --git a/checks/runner.go b/checks/runner.go index 9f0632c..9dbcb1a 100644 --- a/checks/runner.go +++ b/checks/runner.go @@ -19,6 +19,9 @@ type RunOptions struct { } func CLIChecks(cliData api.CLIData, options RunOptions, send func(tea.Msg)) ([]api.CLIStepResult, error) { + if err := validateCLIAssertions(cliData); err != nil { + return nil, err + } shell, err := resolveShell(options.Shell) if err != nil { return nil, err diff --git a/checks/runner_test.go b/checks/runner_test.go index 8668f4d..809ef4d 100644 --- a/checks/runner_test.go +++ b/checks/runner_test.go @@ -118,7 +118,7 @@ func TestCLIChecksReturnsManifestErrors(t *testing.T) { { name: "missing step type", data: api.CLIData{Steps: []api.CLIStep{{}}}, - want: "unable to run lesson: missing step", + want: "must contain exactly one command or HTTP request", }, } diff --git a/checks/validation.go b/checks/validation.go new file mode 100644 index 0000000..c705448 --- /dev/null +++ b/checks/validation.go @@ -0,0 +1,116 @@ +package checks + +import ( + "errors" + "fmt" + "regexp" + + api "github.com/bootdotdev/bootdev/client" +) + +func validateCLIAssertions(cliData api.CLIData) error { + for stepIndex, step := range cliData.Steps { + if (step.CLICommand == nil) == (step.HTTPRequest == nil) { + return fmt.Errorf("step %d must contain exactly one command or HTTP request", stepIndex+1) + } + var err error + if step.CLICommand != nil { + err = validateCommandAssertions(*step.CLICommand) + } else { + err = validateHTTPAssertions(*step.HTTPRequest) + } + if err != nil { + return fmt.Errorf("step %d: %w", stepIndex+1, err) + } + } + return nil +} + +func validateCommandAssertions(command api.CLIStepCLICommand) error { + for testIndex, test := range command.Tests { + if test.ExitCode == nil && test.StdoutLinesGT == nil && test.StdoutJq == nil && len(test.StdoutContainsAll) == 0 && len(test.StdoutContainsNone) == 0 { + return fmt.Errorf("test %d contains no assertions", testIndex+1) + } + if test.StdoutJq == nil { + continue + } + if test.StdoutJq.Query == "" || len(test.StdoutJq.ExpectedResults) == 0 { + return fmt.Errorf("test %d requires a jq query and expected results", testIndex+1) + } + for _, expected := range test.StdoutJq.ExpectedResults { + if expected.Type != api.JqTypeBool && expected.Type != api.JqTypeString && expected.Type != api.JqTypeInt { + return errors.New("invalid jq expected type") + } + switch expected.Operator { + case "==", ">", ">=", "<", "<=": + default: + return errors.New("invalid jq operator") + } + if expected.Value == nil { + return errors.New("missing jq expected value") + } + } + } + return nil +} + +func validateHTTPAssertions(request api.CLIStepHTTPRequest) error { + for testIndex, test := range request.Tests { + if test.StatusCode == nil && test.BodyContains == nil && test.BodyContainsNone == nil && test.HeadersEqual == nil && test.HeadersContain == nil && test.TrailersEqual == nil && test.TrailersContain == nil && test.JSONValue == nil { + return fmt.Errorf("test %d contains no assertions", testIndex+1) + } + if test.JSONValue == nil { + continue + } + valueCount := 0 + if test.JSONValue.IntValue != nil { + valueCount++ + } + if test.JSONValue.StringValue != nil { + valueCount++ + } + if test.JSONValue.BoolValue != nil { + valueCount++ + } + if valueCount != 1 { + return errors.New("JSON assertion requires exactly one expected value type") + } + if test.JSONValue.Path == "" { + return errors.New("JSON assertion requires a path") + } + switch test.JSONValue.Operator { + case api.OpEquals, api.OpGreaterThan, api.OpContains, api.OpNotContains: + default: + return errors.New("invalid JSON assertion operator") + } + } + for _, capture := range request.ResponseVariables { + if (capture.Path == "") == (capture.BodyRegex == "") { + return errors.New("response variable requires exactly one of path or bodyRegex") + } + if capture.BodyRegex != "" { + if err := validateCaptureRegex(capture.BodyRegex); err != nil { + return err + } + } + } + for _, capture := range request.ResponseHeaderVariables { + if capture.Regex != "" { + if err := validateCaptureRegex(capture.Regex); err != nil { + return err + } + } + } + return nil +} + +func validateCaptureRegex(pattern string) error { + expression, err := regexp.Compile(pattern) + if err != nil { + return err + } + if expression.NumSubexp() != 1 { + return errors.New("capture regex requires exactly one capture group") + } + return nil +}