diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 91b9860..8b18a55 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -1,18 +1,49 @@ -on: [push, pull_request] name: Test + +on: + push: + branches: + - master + pull_request: + branches: + - master + jobs: test: strategy: + fail-fast: false matrix: - go-version: [1.13.x, 1.14.x] + go-version: ["1.23", "stable"] platform: [ubuntu-latest, macos-latest, windows-latest] runs-on: ${{ matrix.platform }} + steps: - - name: Install Go - uses: actions/setup-go@v2 + - name: Checkout code + uses: actions/checkout@v4 + + - name: Setup Go + uses: actions/setup-go@v5 with: go-version: ${{ matrix.go-version }} - - name: Checkout code - uses: actions/checkout@v2 + - name: Test - run: go test ./... \ No newline at end of file + run: go test -race ./... + + lint: + runs-on: ubuntu-latest + + steps: + - name: Checkout code + uses: actions/checkout@v4 + + # The linter is built against a specific Go release and cannot type-check a newer + # standard library, so pin its toolchain to the module's Go version rather than "stable". + - name: Setup Go + uses: actions/setup-go@v5 + with: + go-version-file: go.mod + + - name: Run golangci-lint + uses: golangci/golangci-lint-action@v8 + with: + version: v2.10 diff --git a/.golangci.yml b/.golangci.yml new file mode 100644 index 0000000..4a584cb --- /dev/null +++ b/.golangci.yml @@ -0,0 +1,22 @@ +version: "2" + +run: + timeout: 5m + +linters: + default: standard + enable: + - copyloopvar + - errorlint + - misspell + - modernize + - nilerr + - unconvert + - unparam + - usestdlibvars + - wastedassign + +formatters: + enable: + - gofmt + - goimports diff --git a/Makefile b/Makefile index 0ce7ec8..06cd15e 100644 --- a/Makefile +++ b/Makefile @@ -1,2 +1,8 @@ +.PHONY: test lint + test: - go test -v . + go vet ./... + go test -race ./... + +lint: + golangci-lint run ./... diff --git a/README.md b/README.md index 4122fc7..30cc111 100644 --- a/README.md +++ b/README.md @@ -1,4 +1,4 @@ -![Test](https://github.com/steinfletcher/apitest-jsonpath/workflows/Test/badge.svg) +[![Test](https://github.com/steinfletcher/apitest-jsonpath/actions/workflows/ci.yml/badge.svg)](https://github.com/steinfletcher/apitest-jsonpath/actions/workflows/ci.yml) # apitest-jsonpath @@ -7,7 +7,7 @@ This library provides jsonpath assertions for [apitest](https://github.com/stein # Installation ```bash -go get -u github.com/steinfletcher/apitest-jsonpath +go get github.com/steinfletcher/apitest-jsonpath ``` ## Examples @@ -17,7 +17,7 @@ go get -u github.com/steinfletcher/apitest-jsonpath `Equal` checks for value equality when the json path expression returns a single result. Given the response is `{"id": 12345}` ```go -apitest.New(handler). +apitest.Handler(handler). Get("/hello"). Expect(t). Assert(jsonpath.Equal(`$.id`, float64(12345))). @@ -31,7 +31,7 @@ apitest.New(). Handler(handler). Get("/hello"). Expect(t). - Assert(jsonpath.Equal(`$`, map[string]interface{}{"message": "hello", "id": float64(12345)})). + Assert(jsonpath.Equal(`$`, map[string]any{"message": "hello", "id": float64(12345)})). End() ``` @@ -40,7 +40,7 @@ apitest.New(). `NotEqual` checks that the json path expression value is not equal to given value ```go -apitest.New(handler). +apitest.Handler(handler). Get("/hello"). Expect(t). Assert(jsonpath.NotEqual(`$.a`, float64(56789))). @@ -54,7 +54,7 @@ apitest.New(). Handler(handler). Get("/hello"). Expect(t). - Assert(jsonpath.NotEqual(`$`, map[string]interface{}{"a": "hello", "b": float64(56789)})). + Assert(jsonpath.NotEqual(`$`, map[string]any{"a": "hello", "b": float64(56789)})). End() ``` @@ -193,4 +193,35 @@ Assert( Equal("f", "c"). End(), ). -``` \ No newline at end of file +``` + +A dot is inserted between the root and each sub-expression unless the sub-expression starts with a bracket, so array +elements under the root can be addressed directly: + +```go +Assert( + jsonpath.Root("$.items"). + Equal("[0].id", float64(1)). + Equal("[1].id", float64(2)). + End(), +). +``` + +### Matching mock request bodies + +The `mocks` package provides the same assertions as `apitest.Matcher` functions, for matching the JSON body of a +request made to an [apitest mock](https://github.com/steinfletcher/apitest#mocking-external-http-calls). + +```go +import "github.com/steinfletcher/apitest-jsonpath/mocks" + +var createUser = apitest.NewMock(). + Post("/user-api"). + AddMatcher(mocks.Equal(`$.name`, "jon")). + AddMatcher(mocks.Contains(`$.roles`, "admin")). + RespondWith(). + Status(http.StatusCreated). + End() +``` + +`mocks.Equal`, `mocks.NotEqual`, `mocks.Contains`, `mocks.Len` and `mocks.GreaterThan` are available. \ No newline at end of file diff --git a/go.mod b/go.mod index 494bccd..115bfba 100644 --- a/go.mod +++ b/go.mod @@ -1,9 +1,15 @@ module github.com/steinfletcher/apitest-jsonpath -go 1.13 +go 1.23 require ( github.com/PaesslerAG/jsonpath v0.1.1 - github.com/steinfletcher/apitest v1.5.10 - github.com/stretchr/testify v1.7.0 + github.com/steinfletcher/apitest v1.6.1 + github.com/stretchr/testify v1.12.1 +) + +require ( + github.com/PaesslerAG/gval v1.0.0 // indirect + github.com/davecgh/go-spew v1.1.1 // indirect + go.yaml.in/yaml/v3 v3.0.5 // indirect ) diff --git a/go.sum b/go.sum index 28bc60e..c7d6da1 100644 --- a/go.sum +++ b/go.sum @@ -1,22 +1,13 @@ github.com/PaesslerAG/gval v1.0.0 h1:GEKnRwkWDdf9dOmKcNrar9EA1bz1z9DqPIO1+iLzhd8= github.com/PaesslerAG/gval v1.0.0/go.mod h1:y/nm5yEyTeX6av0OfKJNp9rBNj2XrGhAf5+v24IBN1I= -github.com/PaesslerAG/jsonpath v0.1.0 h1:gADYeifvlqK3R3i2cR5B4DGgxLXIPb3TRTH1mGi0jPI= github.com/PaesslerAG/jsonpath v0.1.0/go.mod h1:4BzmtoM/PI8fPO4aQGIusjGxGir2BzcV0grWtFzq1Y8= github.com/PaesslerAG/jsonpath v0.1.1 h1:c1/AToHQMVsduPAa4Vh6xp2U0evy4t8SWp8imEsylIk= github.com/PaesslerAG/jsonpath v0.1.1/go.mod h1:lVboNxFGal/VwW6d9JzIy56bUsYAP6tH/x80vjnCseY= -github.com/davecgh/go-spew v1.1.0 h1:ZDRjVQ15GmhC3fiQ8ni8+OwkZQO4DARzQgrnXU1Liz8= -github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= -github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= -github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= -github.com/steinfletcher/apitest v1.5.10 h1:uxEm/boegmZI9csm1fLVywB5b07ijcrcHo3PZO6sfns= -github.com/steinfletcher/apitest v1.5.10/go.mod h1:cf7Bneo52IIAgpqhP8xaLlzWgAiQ9fHtsDMjeDnZ3so= -github.com/stretchr/objx v0.1.0 h1:4G4v2dO3VZwixGIRoQ5Lfboy6nUhCyYzaqnIAPPhYs4= -github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= -github.com/stretchr/testify v1.7.0 h1:nwc3DEeHmmLAfoZucVR881uASk0Mfjw8xYJ99tb5CcY= -github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= -gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= -gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= -gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c h1:dUUwHk2QECo/6vqA44rthZ8ie2QXMNeKRTHCNY2nXvo= -gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= +github.com/steinfletcher/apitest v1.6.1 h1:gZRLz/Q4sVl2QVa2NYGIGZ1/UOV0rn15Xzd1SO51pvI= +github.com/steinfletcher/apitest v1.6.1/go.mod h1:ToZZZP3/cTb6+gTkDUYrlcgFX9C2Z4965/PtmCQ9Onw= +github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE= +github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg= +go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw= +go.yaml.in/yaml/v3 v3.0.5/go.mod h1:HVTZu1O7/Vkt2N+BFy8Zza+lnLsABggaTM2ZpNIGuKg= diff --git a/http/http.go b/http/http.go index 6cf5865..553c765 100644 --- a/http/http.go +++ b/http/http.go @@ -2,7 +2,7 @@ package http import ( "bytes" - "io/ioutil" + "io" "net/http" "net/url" ) @@ -14,15 +14,15 @@ func CopyResponse(response *http.Response) *http.Response { var resBodyBytes []byte if response.Body != nil { - resBodyBytes, _ = ioutil.ReadAll(response.Body) - response.Body = ioutil.NopCloser(bytes.NewBuffer(resBodyBytes)) + resBodyBytes, _ = io.ReadAll(response.Body) + response.Body = io.NopCloser(bytes.NewBuffer(resBodyBytes)) } resCopy := &http.Response{ Header: map[string][]string{}, StatusCode: response.StatusCode, Status: response.Status, - Body: ioutil.NopCloser(bytes.NewBuffer(resBodyBytes)), + Body: io.NopCloser(bytes.NewBuffer(resBodyBytes)), Proto: response.Proto, ProtoMinor: response.ProtoMinor, ProtoMajor: response.ProtoMajor, @@ -30,13 +30,17 @@ func CopyResponse(response *http.Response) *http.Response { } for name, values := range response.Header { - resCopy.Header[name] = values + resCopy.Header[name] = append([]string(nil), values...) } return resCopy } func CopyRequest(request *http.Request) *http.Request { + if request == nil { + return nil + } + resCopy := &http.Request{ Method: request.Method, Host: request.Host, @@ -49,9 +53,9 @@ func CopyRequest(request *http.Request) *http.Request { resCopy = resCopy.WithContext(request.Context()) if request.Body != nil { - bodyBytes, _ := ioutil.ReadAll(request.Body) - resCopy.Body = ioutil.NopCloser(bytes.NewBuffer(bodyBytes)) - request.Body = ioutil.NopCloser(bytes.NewBuffer(bodyBytes)) + bodyBytes, _ := io.ReadAll(request.Body) + resCopy.Body = io.NopCloser(bytes.NewBuffer(bodyBytes)) + request.Body = io.NopCloser(bytes.NewBuffer(bodyBytes)) } if request.URL != nil { diff --git a/http/http_test.go b/http/http_test.go new file mode 100644 index 0000000..7fd9744 --- /dev/null +++ b/http/http_test.go @@ -0,0 +1,56 @@ +package http + +import ( + "bytes" + "io" + nethttp "net/http" + "net/http/httptest" + "testing" +) + +func TestCopyResponse(t *testing.T) { + original := &nethttp.Response{ + StatusCode: nethttp.StatusOK, + Header: nethttp.Header{"X-Custom": {"a"}}, + Body: io.NopCloser(bytes.NewBufferString("body")), + } + + copied := CopyResponse(original) + copied.Header["X-Custom"][0] = "changed" + copied.Header.Add("X-Custom", "b") + + if got := original.Header["X-Custom"]; len(got) != 1 || got[0] != "a" { + t.Fatalf("expected the original headers to be untouched, got %v", got) + } + for name, body := range map[string]*nethttp.Response{"copy": copied, "original": original} { + b, _ := io.ReadAll(body.Body) + if string(b) != "body" { + t.Fatalf("expected %s body to be readable, got %q", name, b) + } + } + if CopyResponse(nil) != nil { + t.Fatal("expected a nil response to copy as nil") + } +} + +func TestCopyRequest(t *testing.T) { + original := httptest.NewRequest(nethttp.MethodPost, "/path?a=1", bytes.NewBufferString("body")) + original.Header.Set("X-Custom", "a") + + copied := CopyRequest(original) + copied.Header.Set("X-Custom", "changed") + copied.URL.Path = "/other" + + if original.Header.Get("X-Custom") != "a" || original.URL.Path != "/path" { + t.Fatalf("expected the original request to be untouched, got %v %s", original.Header, original.URL) + } + for name, req := range map[string]*nethttp.Request{"copy": copied, "original": original} { + b, _ := io.ReadAll(req.Body) + if string(b) != "body" { + t.Fatalf("expected %s body to be readable, got %q", name, b) + } + } + if CopyRequest(nil) != nil { + t.Fatal("expected a nil request to copy as nil") + } +} diff --git a/jsonpath.go b/jsonpath.go index bc8d0da..581b1d3 100644 --- a/jsonpath.go +++ b/jsonpath.go @@ -5,27 +5,28 @@ import ( "net/http" "reflect" regex "regexp" + "strings" httputil "github.com/steinfletcher/apitest-jsonpath/http" "github.com/steinfletcher/apitest-jsonpath/jsonpath" ) // Contains is a convenience function to assert that a jsonpath expression extracts a value in an array -func Contains(expression string, expected interface{}) func(*http.Response, *http.Request) error { +func Contains(expression string, expected any) func(*http.Response, *http.Request) error { return func(res *http.Response, req *http.Request) error { return jsonpath.Contains(expression, expected, res.Body) } } // Equal is a convenience function to assert that a jsonpath expression extracts a value -func Equal(expression string, expected interface{}) func(*http.Response, *http.Request) error { +func Equal(expression string, expected any) func(*http.Response, *http.Request) error { return func(res *http.Response, req *http.Request) error { return jsonpath.Equal(expression, expected, res.Body) } } // NotEqual is a function to check json path expression value is not equal to given value -func NotEqual(expression string, expected interface{}) func(*http.Response, *http.Request) error { +func NotEqual(expression string, expected any) func(*http.Response, *http.Request) error { return func(res *http.Response, req *http.Request) error { return jsonpath.NotEqual(expression, expected, res.Body) } @@ -73,7 +74,10 @@ func Matches(expression string, regexp string) func(*http.Response, *http.Reques if err != nil { return fmt.Errorf("invalid pattern: '%s'", regexp) } - value, _ := jsonpath.JsonPath(res.Body, expression) + value, err := jsonpath.JsonPath(res.Body, expression) + if err != nil { + return err + } if value == nil { return fmt.Errorf("no match for pattern: '%s'", expression) } @@ -94,7 +98,7 @@ func Matches(expression string, regexp string) func(*http.Response, *http.Reques reflect.Float32, reflect.Float64, reflect.String: - if !pattern.Match([]byte(fmt.Sprintf("%v", value))) { + if !pattern.MatchString(fmt.Sprintf("%v", value)) { return fmt.Errorf("value '%v' does not match pattern '%v'", value, regexp) } return nil @@ -111,7 +115,7 @@ func Chain() *AssertionChain { // Root creates a new assertion chain prefixed with the given expression func Root(expression string) *AssertionChain { - return &AssertionChain{rootExpression: expression + "."} + return &AssertionChain{rootExpression: strings.TrimSuffix(expression, ".")} } // AssertionChain supports chaining assertions and root expressions @@ -121,41 +125,54 @@ type AssertionChain struct { } // Equal adds an Equal assertion to the chain -func (r *AssertionChain) Equal(expression string, expected interface{}) *AssertionChain { - r.assertions = append(r.assertions, Equal(r.rootExpression+expression, expected)) +func (r *AssertionChain) Equal(expression string, expected any) *AssertionChain { + r.assertions = append(r.assertions, Equal(r.path(expression), expected)) return r } // NotEqual adds an NotEqual assertion to the chain -func (r *AssertionChain) NotEqual(expression string, expected interface{}) *AssertionChain { - r.assertions = append(r.assertions, NotEqual(r.rootExpression+expression, expected)) +func (r *AssertionChain) NotEqual(expression string, expected any) *AssertionChain { + r.assertions = append(r.assertions, NotEqual(r.path(expression), expected)) return r } // Contains adds an Contains assertion to the chain -func (r *AssertionChain) Contains(expression string, expected interface{}) *AssertionChain { - r.assertions = append(r.assertions, Contains(r.rootExpression+expression, expected)) +func (r *AssertionChain) Contains(expression string, expected any) *AssertionChain { + r.assertions = append(r.assertions, Contains(r.path(expression), expected)) return r } // Present adds an Present assertion to the chain func (r *AssertionChain) Present(expression string) *AssertionChain { - r.assertions = append(r.assertions, Present(r.rootExpression+expression)) + r.assertions = append(r.assertions, Present(r.path(expression))) return r } // NotPresent adds an NotPresent assertion to the chain func (r *AssertionChain) NotPresent(expression string) *AssertionChain { - r.assertions = append(r.assertions, NotPresent(r.rootExpression+expression)) + r.assertions = append(r.assertions, NotPresent(r.path(expression))) return r } // Matches adds an Matches assertion to the chain func (r *AssertionChain) Matches(expression, regexp string) *AssertionChain { - r.assertions = append(r.assertions, Matches(r.rootExpression+expression, regexp)) + r.assertions = append(r.assertions, Matches(r.path(expression), regexp)) return r } +// path joins the root expression and the given expression. A dot is inserted between them unless +// the expression already starts with one or with a bracket, so that Root("$.items").Equal("[0].id", 1) +// evaluates "$.items[0].id". +func (r *AssertionChain) path(expression string) string { + if r.rootExpression == "" { + return expression + } + if strings.HasPrefix(expression, "[") || strings.HasPrefix(expression, ".") { + return r.rootExpression + expression + } + return r.rootExpression + "." + expression +} + // End returns an func(*http.Response, *http.Request) error which is a combination of the registered assertions func (r *AssertionChain) End() func(*http.Response, *http.Request) error { return func(res *http.Response, req *http.Request) error { diff --git a/jsonpath/jsonpath.go b/jsonpath/jsonpath.go index f835f14..ee54e3e 100644 --- a/jsonpath/jsonpath.go +++ b/jsonpath/jsonpath.go @@ -6,100 +6,106 @@ import ( "errors" "fmt" "io" - "io/ioutil" "reflect" "strings" "github.com/PaesslerAG/jsonpath" ) -func Contains(expression string, expected interface{}, data io.Reader) error { +func Contains(expression string, expected any, data io.Reader) error { value, err := JsonPath(data, expression) if err != nil { return err } ok, found := IncludesElement(value, expected) if !ok { - return fmt.Errorf("\"%s\" could not be applied builtin len()", expected) + return fmt.Errorf("\"%v\" could not be applied builtin len()", expected) } if !found { - return fmt.Errorf("\"%s\" does not contain \"%s\"", value, expected) + return fmt.Errorf("\"%v\" does not contain \"%v\"", value, expected) } return nil } -func Equal(expression string, expected interface{}, data io.Reader) error { +func Equal(expression string, expected any, data io.Reader) error { value, err := JsonPath(data, expression) if err != nil { return err } if !ObjectsAreEqual(value, expected) { - return fmt.Errorf("\"%s\" not equal to \"%s\"", value, expected) + return fmt.Errorf("\"%v\" not equal to \"%v\"", value, expected) } return nil } -func NotEqual(expression string, expected interface{}, data io.Reader) error { +func NotEqual(expression string, expected any, data io.Reader) error { value, err := JsonPath(data, expression) if err != nil { return err } if ObjectsAreEqual(value, expected) { - return fmt.Errorf("\"%s\" value is equal to \"%s\"", expression, expected) + return fmt.Errorf("\"%s\" value is equal to \"%v\"", expression, expected) } return nil } func Length(expression string, expectedLength int, data io.Reader) error { - value, err := JsonPath(data, expression) + length, err := lengthOf(expression, data) if err != nil { return err } - if value == nil { - return errors.New("value is null") - } - - v := reflect.ValueOf(value) - if v.Len() != expectedLength { - return fmt.Errorf("\"%d\" not equal to \"%d\"", v.Len(), expectedLength) + if length != expectedLength { + return fmt.Errorf("\"%d\" not equal to \"%d\"", length, expectedLength) } return nil } func GreaterThan(expression string, minimumLength int, data io.Reader) error { - value, err := JsonPath(data, expression) + length, err := lengthOf(expression, data) if err != nil { return err } - if value == nil { - return fmt.Errorf("value is null") + if length < minimumLength { + return fmt.Errorf("\"%d\" is less than \"%d\"", length, minimumLength) } + return nil +} - v := reflect.ValueOf(value) - if v.Len() < minimumLength { - return fmt.Errorf("\"%d\" is greater than \"%d\"", v.Len(), minimumLength) +func LessThan(expression string, maximumLength int, data io.Reader) error { + length, err := lengthOf(expression, data) + if err != nil { + return err + } + + if length > maximumLength { + return fmt.Errorf("\"%d\" is greater than \"%d\"", length, maximumLength) } return nil } -func LessThan(expression string, maximumLength int, data io.Reader) error { +// lengthOf evaluates the expression and returns the length of the result. Only arrays, slices, +// maps and strings have a length; a null result or a result of any other type is an error rather +// than a panic, so that a failed assertion is reported normally. +func lengthOf(expression string, data io.Reader) (int, error) { value, err := JsonPath(data, expression) if err != nil { - return err + return 0, err } if value == nil { - return fmt.Errorf("value is null") + return 0, errors.New("value is null") } v := reflect.ValueOf(value) - if v.Len() > maximumLength { - return fmt.Errorf("\"%d\" is less than \"%d\"", v.Len(), maximumLength) + switch v.Kind() { + case reflect.Array, reflect.Chan, reflect.Map, reflect.Slice, reflect.String: + return v.Len(), nil + default: + return 0, fmt.Errorf("value of type %s has no length", v.Kind()) } - return nil } func Present(expression string, data io.Reader) error { @@ -118,9 +124,9 @@ func NotPresent(expression string, data io.Reader) error { return nil } -func JsonPath(reader io.Reader, expression string) (interface{}, error) { - v := interface{}(nil) - b, err := ioutil.ReadAll(reader) +func JsonPath(reader io.Reader, expression string) (any, error) { + v := any(nil) + b, err := io.ReadAll(reader) if err != nil { return nil, err } @@ -132,13 +138,13 @@ func JsonPath(reader io.Reader, expression string) (interface{}, error) { value, err := jsonpath.Get(expression, v) if err != nil { - return nil, fmt.Errorf("evaluating '%s' resulted in error: '%s'", expression, err) + return nil, fmt.Errorf("evaluating '%s' resulted in error: '%w'", expression, err) } return value, nil } // courtesy of github.com/stretchr/testify -func IncludesElement(list interface{}, element interface{}) (ok, found bool) { +func IncludesElement(list any, element any) (ok, found bool) { listValue := reflect.ValueOf(list) elementValue := reflect.ValueOf(element) defer func() { @@ -154,7 +160,7 @@ func IncludesElement(list interface{}, element interface{}) (ok, found bool) { if reflect.TypeOf(list).Kind() == reflect.Map { mapKeys := listValue.MapKeys() - for i := 0; i < len(mapKeys); i++ { + for i := range mapKeys { if ObjectsAreEqual(mapKeys[i].Interface(), element) { return true, true } @@ -170,7 +176,7 @@ func IncludesElement(list interface{}, element interface{}) (ok, found bool) { return true, false } -func ObjectsAreEqual(expected, actual interface{}) bool { +func ObjectsAreEqual(expected, actual any) bool { if expected == nil || actual == nil { return expected == actual } @@ -190,7 +196,7 @@ func ObjectsAreEqual(expected, actual interface{}) bool { return bytes.Equal(exp, act) } -func isEmpty(object interface{}) bool { +func isEmpty(object any) bool { if object == nil { return true } diff --git a/jsonpath/jsonpath_test.go b/jsonpath/jsonpath_test.go new file mode 100644 index 0000000..127fee8 --- /dev/null +++ b/jsonpath/jsonpath_test.go @@ -0,0 +1,114 @@ +package jsonpath + +import ( + "strings" + "testing" +) + +func assertError(t *testing.T, err error, expected string) { + t.Helper() + if expected == "" { + if err != nil { + t.Fatalf("expected no error, got %q", err) + } + return + } + if err == nil || err.Error() != expected { + t.Fatalf("expected error %q, got %v", expected, err) + } +} + +func TestLength(t *testing.T) { + tests := map[string]struct { + body string + expression string + length int + expected string + }{ + "array": {`{"items": [1, 2, 3]}`, `$.items`, 3, ""}, + "string": {`{"name": "jan"}`, `$.name`, 3, ""}, + "map": {`{"user": {"a": 1, "b": 2}}`, `$.user`, 2, ""}, + "wrong": {`{"items": [1, 2, 3]}`, `$.items`, 2, `"3" not equal to "2"`}, + "null": {`{"items": null}`, `$.items`, 0, "value is null"}, + "number": {`{"count": 3}`, `$.count`, 3, "value of type float64 has no length"}, + "bool": {`{"ok": true}`, `$.ok`, 1, "value of type bool has no length"}, + "missing key": {`{"a": 1}`, `$.items`, 0, "evaluating '$.items' resulted in error: 'unknown key items'"}, + } + + for name, test := range tests { + t.Run(name, func(t *testing.T) { + assertError(t, Length(test.expression, test.length, strings.NewReader(test.body)), test.expected) + }) + } +} + +func TestGreaterThan(t *testing.T) { + tests := map[string]struct { + body string + minimum int + expected string + }{ + "longer": {`{"items": [1, 2, 3]}`, 2, ""}, + "equal": {`{"items": [1, 2]}`, 2, ""}, + "shorter": {`{"items": [1]}`, 2, `"1" is less than "2"`}, + "null": {`{"items": null}`, 0, "value is null"}, + "number": {`{"items": 5}`, 0, "value of type float64 has no length"}, + } + + for name, test := range tests { + t.Run(name, func(t *testing.T) { + assertError(t, GreaterThan(`$.items`, test.minimum, strings.NewReader(test.body)), test.expected) + }) + } +} + +func TestLessThan(t *testing.T) { + tests := map[string]struct { + body string + maximum int + expected string + }{ + "shorter": {`{"items": [1]}`, 2, ""}, + "equal": {`{"items": [1, 2]}`, 2, ""}, + "longer": {`{"items": [1, 2, 3]}`, 2, `"3" is greater than "2"`}, + "null": {`{"items": null}`, 0, "value is null"}, + "bool": {`{"items": false}`, 0, "value of type bool has no length"}, + } + + for name, test := range tests { + t.Run(name, func(t *testing.T) { + assertError(t, LessThan(`$.items`, test.maximum, strings.NewReader(test.body)), test.expected) + }) + } +} + +func TestContains(t *testing.T) { + tests := map[string]struct { + body string + expected any + err string + }{ + "number in array": {`{"items": [1, 2]}`, float64(2), ""}, + "string in array": {`{"items": ["a", "b"]}`, "b", ""}, + "substring": {`{"items": "abc"}`, "b", ""}, + "key in map": {`{"items": {"a": 1}}`, "a", ""}, + "number not in array": {`{"items": [1, 2]}`, float64(5), `"[1 2]" does not contain "5"`}, + "value without length": {`{"items": 5}`, float64(5), `"5" could not be applied builtin len()`}, + } + + for name, test := range tests { + t.Run(name, func(t *testing.T) { + assertError(t, Contains(`$.items`, test.expected, strings.NewReader(test.body)), test.err) + }) + } +} + +func TestEqualAndNotEqual(t *testing.T) { + body := `{"id": 12345, "name": "jan"}` + + assertError(t, Equal(`$.id`, float64(12345), strings.NewReader(body)), "") + assertError(t, Equal(`$.id`, float64(1), strings.NewReader(body)), `"12345" not equal to "1"`) + assertError(t, Equal(`$.missing`, "x", strings.NewReader(body)), "evaluating '$.missing' resulted in error: 'unknown key missing'") + assertError(t, NotEqual(`$.name`, "jon", strings.NewReader(body)), "") + assertError(t, NotEqual(`$.name`, "jan", strings.NewReader(body)), `"$.name" value is equal to "jan"`) +} diff --git a/jsonpath_test.go b/jsonpath_test.go index 09e85e2..f6a9e7f 100644 --- a/jsonpath_test.go +++ b/jsonpath_test.go @@ -4,8 +4,9 @@ import ( "bytes" "errors" "fmt" - "io/ioutil" + "io" "net/http" + "net/http/httptest" "testing" "github.com/steinfletcher/apitest" @@ -86,7 +87,7 @@ func TestApiTest_Equal_Map(t *testing.T) { Handler(handler). Get("/hello"). Expect(t). - Assert(jsonpath.Equal(`$`, map[string]interface{}{"a": "hello", "b": float64(12345)})). + Assert(jsonpath.Equal(`$`, map[string]any{"a": "hello", "b": float64(12345)})). End() } @@ -143,7 +144,7 @@ func TestApiTest_NotEqual_Map(t *testing.T) { Handler(handler). Get("/hello"). Expect(t). - Assert(jsonpath.NotEqual(`$`, map[string]interface{}{"a": "hello", "b": float64(1)})). + Assert(jsonpath.NotEqual(`$`, map[string]any{"a": "hello", "b": float64(1)})). End() } @@ -350,7 +351,7 @@ func TestApiTest_Matches_FailForObject(t *testing.T) { matcher := jsonpath.Matches(`$.anObject`, `.+`) err := matcher(&http.Response{ - Body: ioutil.NopCloser(bytes.NewBuffer([]byte(`{"anObject":{"aString":"lol"}}`))), + Body: io.NopCloser(bytes.NewBuffer([]byte(`{"anObject":{"aString":"lol"}}`))), }, nil) assert.EqualError(t, err, "unable to match using type: map") @@ -360,7 +361,7 @@ func TestApiTest_Matches_FailForArray(t *testing.T) { matcher := jsonpath.Matches(`$.aSlice`, `.+`) err := matcher(&http.Response{ - Body: ioutil.NopCloser(bytes.NewBuffer([]byte(`{"aSlice":[1,2,3]}`))), + Body: io.NopCloser(bytes.NewBuffer([]byte(`{"aSlice":[1,2,3]}`))), }, nil) assert.EqualError(t, err, "unable to match using type: slice") @@ -370,8 +371,53 @@ func TestApiTest_Matches_FailForNilValue(t *testing.T) { matcher := jsonpath.Matches(`$.nothingHere`, `.+`) err := matcher(&http.Response{ - Body: ioutil.NopCloser(bytes.NewBuffer([]byte(`{"aSlice":[1,2,3]}`))), + Body: io.NopCloser(bytes.NewBuffer([]byte(`{"aSlice":[1,2,3]}`))), + }, nil) + + assert.EqualError(t, err, "evaluating '$.nothingHere' resulted in error: 'unknown key nothingHere'") +} + +func TestApiTest_Matches_FailForNullValue(t *testing.T) { + matcher := jsonpath.Matches(`$.nothingHere`, `.+`) + + err := matcher(&http.Response{ + Body: io.NopCloser(bytes.NewBuffer([]byte(`{"nothingHere": null}`))), }, nil) assert.EqualError(t, err, "no match for pattern: '$.nothingHere'") } + +func TestApiTest_Matches_ReportsInvalidJSON(t *testing.T) { + matcher := jsonpath.Matches(`$.a`, `.+`) + + err := matcher(&http.Response{ + Body: io.NopCloser(bytes.NewBuffer([]byte(`not json`))), + }, nil) + + assert.EqualError(t, err, "invalid character 'o' in literal null (expecting 'u')") +} + +func TestApiTest_Matches_ReportsInvalidExpressions(t *testing.T) { + matcher := jsonpath.Matches(`$[`, `.+`) + + err := matcher(&http.Response{ + Body: io.NopCloser(bytes.NewBuffer([]byte(`{"a": 1}`))), + }, nil) + + assert.Error(t, err) + assert.Contains(t, err.Error(), "evaluating '$[' resulted in error") +} + +func TestApiTest_Root_JoinsSubExpressions(t *testing.T) { + body := `{"items": [{"id": 1, "name": "jan"}], "count": 1}` + req := httptest.NewRequest(http.MethodGet, "/", nil) + response := func() *http.Response { + return &http.Response{Body: io.NopCloser(bytes.NewBufferString(body))} + } + + assert.NoError(t, jsonpath.Root(`$.items`).Equal(`[0].id`, float64(1)).End()(response(), req)) + assert.NoError(t, jsonpath.Root(`$.items[0]`).Equal(`id`, float64(1)).Equal(`.name`, "jan").End()(response(), req)) + assert.NoError(t, jsonpath.Root(`$.items[0].`).Equal(`name`, "jan").End()(response(), req)) + assert.NoError(t, jsonpath.Chain().Equal(`$.count`, float64(1)).End()(response(), req)) + assert.EqualError(t, jsonpath.Root(`$.items[0]`).Equal(`id`, float64(2)).End()(response(), req), `"1" not equal to "2"`) +} diff --git a/jwt.go b/jwt.go index 584ab1e..0a6d275 100644 --- a/jwt.go +++ b/jwt.go @@ -16,15 +16,15 @@ const ( jwtPayloadIndex = 1 ) -func JWTHeaderEqual(tokenSelector func(*http.Response) (string, error), expression string, expected interface{}) func(*http.Response, *http.Request) error { +func JWTHeaderEqual(tokenSelector func(*http.Response) (string, error), expression string, expected any) func(*http.Response, *http.Request) error { return jwtEqual(tokenSelector, expression, expected, jwtHeaderIndex) } -func JWTPayloadEqual(tokenSelector func(*http.Response) (string, error), expression string, expected interface{}) func(*http.Response, *http.Request) error { +func JWTPayloadEqual(tokenSelector func(*http.Response) (string, error), expression string, expected any) func(*http.Response, *http.Request) error { return jwtEqual(tokenSelector, expression, expected, jwtPayloadIndex) } -func jwtEqual(tokenSelector func(*http.Response) (string, error), expression string, expected interface{}, index int) func(*http.Response, *http.Request) error { +func jwtEqual(tokenSelector func(*http.Response) (string, error), expression string, expected any, index int) func(*http.Response, *http.Request) error { return func(response *http.Response, request *http.Request) error { token, err := tokenSelector(response) if err != nil { @@ -33,22 +33,21 @@ func jwtEqual(tokenSelector func(*http.Response) (string, error), expression str parts := strings.Split(token, ".") if len(parts) != 3 { - splitErr := errors.New("invalid token: token should contain header, payload and secret") - return splitErr + return errors.New("invalid token: token should contain header, payload and signature") } - decodedPayload, PayloadErr := base64Decode(parts[index]) - if PayloadErr != nil { - return fmt.Errorf("invalid jwt: %s", PayloadErr.Error()) + decoded, err := base64Decode(parts[index]) + if err != nil { + return fmt.Errorf("invalid jwt: %w", err) } - value, err := jsonpath.JsonPath(bytes.NewReader(decodedPayload), expression) + value, err := jsonpath.JsonPath(bytes.NewReader(decoded), expression) if err != nil { return err } if !jsonpath.ObjectsAreEqual(value, expected) { - return errors.New(fmt.Sprintf("\"%s\" not equal to \"%s\"", value, expected)) + return fmt.Errorf("\"%v\" not equal to \"%v\"", value, expected) } return nil @@ -62,8 +61,7 @@ func base64Decode(src string) ([]byte, error) { decoded, err := base64.URLEncoding.DecodeString(src) if err != nil { - errMsg := fmt.Errorf("decoding Error %s", err) - return nil, errMsg + return nil, fmt.Errorf("decoding error: %w", err) } return decoded, nil } diff --git a/jwt_test.go b/jwt_test.go index afddc83..28e5edf 100644 --- a/jwt_test.go +++ b/jwt_test.go @@ -34,3 +34,26 @@ func TestApiTest_JWT(t *testing.T) { func fromAuthHeader(response *http.Response) (string, error) { return response.Header.Get("Authorization"), nil } + +func TestApiTest_JWT_Errors(t *testing.T) { + selector := func(token string) func(*http.Response) (string, error) { + return func(*http.Response) (string, error) { return token, nil } + } + tests := map[string]struct { + token string + expected string + }{ + "wrong number of parts": {"a.b", "invalid token: token should contain header, payload and signature"}, + "payload not base64": {"a.!!!.c", "invalid jwt: decoding error: illegal base64 data at input byte 0"}, + "payload not json": {"a.bm90IGpzb24.c", "invalid character 'o' in literal null (expecting 'u')"}, + } + + for name, test := range tests { + t.Run(name, func(t *testing.T) { + err := jsonpath.JWTPayloadEqual(selector(test.token), `$.sub`, "x")(nil, nil) + if err == nil || err.Error() != test.expected { + t.Fatalf("expected error %q, got %v", test.expected, err) + } + }) + } +} diff --git a/mocks/mocks.go b/mocks/mocks.go index 009a0fd..9fbb92e 100644 --- a/mocks/mocks.go +++ b/mocks/mocks.go @@ -9,21 +9,21 @@ import ( ) // Contains is a convenience function to assert that a jsonpath expression extracts a value in an array -func Contains(expression string, expected interface{}) apitest.Matcher { +func Contains(expression string, expected any) apitest.Matcher { return func(req *http.Request, mockReq *apitest.MockRequest) error { return jsonpath.Contains(expression, expected, httputil.CopyRequest(req).Body) } } // Equal is a convenience function to assert that a jsonpath expression matches the given value -func Equal(expression string, expected interface{}) apitest.Matcher { +func Equal(expression string, expected any) apitest.Matcher { return func(req *http.Request, mockReq *apitest.MockRequest) error { return jsonpath.Equal(expression, expected, httputil.CopyRequest(req).Body) } } // NotEqual is a function to check json path expression value is not equal to given value -func NotEqual(expression string, expected interface{}) apitest.Matcher { +func NotEqual(expression string, expected any) apitest.Matcher { return func(req *http.Request, mockReq *apitest.MockRequest) error { return jsonpath.NotEqual(expression, expected, httputil.CopyRequest(req).Body) } diff --git a/mocks/mocks_test.go b/mocks/mocks_test.go index 8306ca1..b2a1a3c 100644 --- a/mocks/mocks_test.go +++ b/mocks/mocks_test.go @@ -3,7 +3,7 @@ package mocks_test import ( "encoding/json" "fmt" - "io/ioutil" + "io" "net/http" "strings" "testing" @@ -86,13 +86,13 @@ type userResponse struct { IsContactable bool `json:"is_contactable"` } -func httpGet(path string, response interface{}) error { +func httpGet(path string, response any) error { res, err := http.DefaultClient.Get(fmt.Sprintf("http://localhost:8080%s", path)) if err != nil { return err } - bytes, err := ioutil.ReadAll(res.Body) + bytes, err := io.ReadAll(res.Body) if err != nil { return err } @@ -105,13 +105,13 @@ func httpGet(path string, response interface{}) error { return nil } -func httpPost(path string, requestBody string, response interface{}) error { +func httpPost(path string, requestBody string, response any) error { res, err := http.DefaultClient.Post(fmt.Sprintf("http://localhost:8080%s", path), "application/json", strings.NewReader(requestBody)) if err != nil { return err } - bytes, err := ioutil.ReadAll(res.Body) + bytes, err := io.ReadAll(res.Body) if err != nil { return err }