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
62 changes: 62 additions & 0 deletions helpers/media_types.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,62 @@
// Copyright 2023-2026 Princess Beef Heavy Industries, LLC / Dave Shanley
// SPDX-License-Identifier: MIT

package helpers

import (
"strings"

"github.com/pb33f/libopenapi/orderedmap"

v3 "github.com/pb33f/libopenapi/datamodel/high/v3"
)

// FindMediaType returns the content entry that applies to contentType, a Content-Type header value.
//
// Content keys may be media ranges, and OpenAPI applies only the most specific key that matches:
// an exact media type, then a structured syntax range (application/*+json), then a type range
// (application/*), then a subtype range (*/json), then */*. Keys are compared without their
// parameters, and case is ignored. Keys of equal specificity keep document order.
func FindMediaType(content *orderedmap.Map[string, *v3.MediaType], contentType string) (*v3.MediaType, bool) {
if content == nil {
return nil, false
}
mediaType, _, _ := ExtractContentType(contentType)
if found, ok := content.Get(mediaType); ok {
return found, true
}

typ, subtype, _ := strings.Cut(mediaType, "/")
suffix := ""
if plus := strings.LastIndexByte(subtype, '+'); plus >= 0 {
suffix = subtype[plus:]
}

var found *v3.MediaType
foundRank := 0
for pair := content.First(); pair != nil; pair = pair.Next() {
var rank int
switch key := normalizeMediaRange(pair.Key()); key {
case mediaType:
rank = 5
case typ + "/*" + suffix:
rank = 4 // the same key as the type range below when there is no suffix
case typ + "/*":
rank = 3
case "*/" + subtype:
rank = 2
case "*/*":
rank = 1
}
if rank > foundRank {
found, foundRank = pair.Value(), rank
}
}
return found, foundRank > 0
}

// normalizeMediaRange lowercases a media type or range and drops its parameters.
func normalizeMediaRange(value string) string {
base, _, _ := strings.Cut(value, ";")
return strings.ToLower(strings.TrimSpace(base))
}
67 changes: 67 additions & 0 deletions helpers/media_types_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,67 @@
// Copyright 2023-2026 Princess Beef Heavy Industries, LLC / Dave Shanley
// SPDX-License-Identifier: MIT

package helpers

import (
"testing"

"github.com/pb33f/libopenapi/orderedmap"
"github.com/pb33f/testify/assert"

v3 "github.com/pb33f/libopenapi/datamodel/high/v3"
)

func mediaTypeContent(keys ...string) *orderedmap.Map[string, *v3.MediaType] {
content := orderedmap.New[string, *v3.MediaType]()
for _, key := range keys {
content.Set(key, &v3.MediaType{})
}
return content
}

func TestFindMediaType(t *testing.T) {
all := mediaTypeContent("*/*", "*/json", "application/*", "application/*+json", "application/json")

for _, test := range []struct {
name string
content *orderedmap.Map[string, *v3.MediaType]
contentType string
expected string
}{
{"exact beats every range", all, "application/json; charset=utf-8", "application/json"},
{"structured syntax range beats type range", all, "application/problem+json", "application/*+json"},
{"type range beats subtype range", all, "application/xml", "application/*"},
{"subtype range beats */*", all, "text/json", "*/json"},
{"*/* matches anything", all, "image/png", "*/*"},
{"type range without suffix range", mediaTypeContent("*/*", "application/*"), "application/problem+json", "application/*"},
{"keys ignore case and parameters", mediaTypeContent("Text/Plain; charset=utf-8"), "text/plain", "Text/Plain; charset=utf-8"},
{"content type ignores case", mediaTypeContent("application/json"), "Application/JSON", "application/json"},
{"equal specificity keeps document order", mediaTypeContent("text/*", "TEXT/*"), "text/csv", "text/*"},
{"no slash only matches ranges it fits", mediaTypeContent("application/json", "application/*"), "application", "application/*"},
{"unparseable content type only matches */*", mediaTypeContent("application/json", "*/*"), "application/", "*/*"},
} {
t.Run(test.name, func(t *testing.T) {
found, ok := FindMediaType(test.content, test.contentType)
assert.True(t, ok)
assert.Same(t, test.content.GetOrZero(test.expected), found)
})
}

for _, test := range []struct {
name string
content *orderedmap.Map[string, *v3.MediaType]
contentType string
}{
{"nil content", nil, "application/json"},
{"no matching key", mediaTypeContent("application/json", "text/*"), "image/png"},
{"no slash and no range", mediaTypeContent("application/json"), "application"},
{"ranges are not matched in reverse", mediaTypeContent("application/json"), "application/*"},
} {
t.Run(test.name, func(t *testing.T) {
found, ok := FindMediaType(test.content, test.contentType)
assert.False(t, ok)
assert.Nil(t, found)
})
}
}
42 changes: 42 additions & 0 deletions helpers/number_utilities.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,42 @@
// Copyright 2023-2026 Princess Beef Heavy Industries, LLC / Dave Shanley
// SPDX-License-Identifier: MIT

package helpers

import (
"errors"
"math"
"strconv"
"strings"
)

// ParseInteger parses a parameter value as a JSON Schema integer: any number with a zero
// fractional part, so "1.0" and "1e3" are integers as well as "1". Values outside the int64
// range, and anything that is not a decimal number, return an error. A value written with a
// fraction or exponent is read as a float64, so it is exact up to 2^53.
func ParseInteger(value string) (int64, error) {
parsed, err := strconv.ParseInt(value, 10, 64)
if err == nil || errors.Is(err, strconv.ErrRange) {
return parsed, err
}
f, floatErr := ParseNumber(value)
if floatErr != nil || f != math.Trunc(f) || math.Abs(f) >= math.MaxInt64 {
return 0, err
}
return int64(f), nil
}

// ParseNumber parses a parameter value as a JSON Schema number, which is always finite.
// strconv.ParseFloat alone also reads NaN, Inf, hex and underscores, none of which are JSON numbers.
func ParseNumber(value string) (float64, error) {
if strings.ContainsFunc(value, notDecimalNumberRune) {
return 0, &strconv.NumError{Func: "ParseFloat", Num: value, Err: strconv.ErrSyntax}
}
// a decimal number too large for a float64 returns an ErrRange error, never an infinity
return strconv.ParseFloat(value, 64)
}

// notDecimalNumberRune reports whether r cannot appear in a decimal number such as "-1.5e3".
func notDecimalNumberRune(r rune) bool {
return (r < '0' || r > '9') && !strings.ContainsRune("+-.eE", r)
}
52 changes: 52 additions & 0 deletions helpers/number_utilities_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,52 @@
// Copyright 2023-2026 Princess Beef Heavy Industries, LLC / Dave Shanley
// SPDX-License-Identifier: MIT

package helpers

import (
"testing"

"github.com/pb33f/testify/assert"
)

func TestParseInteger(t *testing.T) {
for value, expected := range map[string]int64{
"1": 1,
"-7": -7,
"+3": 3,
"007": 7,
"1.0": 1,
"-2.000": -2,
"1e3": 1000,
"2.5E1": 25,
"-0.0": 0,
"9223372036854775807": 9223372036854775807,
} {
parsed, err := ParseInteger(value)
assert.NoError(t, err, value)
assert.Equal(t, expected, parsed, value)
}

for _, value := range []string{
"", "abc", "1.5", "1e-3", "0x10", "0x1p3", "1_000", "NaN", "Inf", "-infinity",
"9223372036854775808", "-9223372036854775809", "1e19", "-1e19", "1e400",
} {
_, err := ParseInteger(value)
assert.Error(t, err, value)
}
}

func TestParseNumber(t *testing.T) {
for value, expected := range map[string]float64{
"1": 1, "-2.5": -2.5, "+3": 3, "1e3": 1000, "2.5E-1": 0.25, "1e-400": 0,
} {
parsed, err := ParseNumber(value)
assert.NoError(t, err, value)
assert.Equal(t, expected, parsed, value)
}

for _, value := range []string{"", "abc", "NaN", "Inf", "+Inf", "-infinity", "0x1p3", "1_000", "1e400", "-1e400", "1.2.3"} {
_, err := ParseNumber(value)
assert.Error(t, err, value)
}
}
10 changes: 4 additions & 6 deletions helpers/parameter_utilities.go
Original file line number Diff line number Diff line change
Expand Up @@ -205,14 +205,12 @@ func cast(v string) any {
b, _ := strconv.ParseBool(v)
return b
}
if i, err := strconv.ParseFloat(v, 64); err == nil {
// check if this is an int or not
if !strings.Contains(v, Period) {
iv, _ := strconv.ParseInt(v, 10, 64)
return iv
}
if i, err := strconv.ParseInt(v, 10, 64); err == nil {
return i
}
if f, err := ParseNumber(v); err == nil {
return f
}
return v
}

Expand Down
49 changes: 49 additions & 0 deletions helpers/schema_compiler.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,8 +3,11 @@ package helpers
import (
"bytes"
"encoding/json"
"errors"
"fmt"
"math"
"sort"
"strconv"

"github.com/santhosh-tekuri/jsonschema/v6"

Expand All @@ -24,12 +27,16 @@ func ConfigureCompiler(c *jsonschema.Compiler, o *config.ValidationOptions) {

if o.FormatAssertions {
c.AssertFormat()
for _, format := range openAPIFormats {
c.RegisterFormat(format)
}
}

if o.ContentAssertions {
c.AssertContent()
}

// custom formats are registered last, so they replace built-in formats of the same name.
for n, v := range o.Formats {
c.RegisterFormat(&jsonschema.Format{
Name: n,
Expand All @@ -38,6 +45,48 @@ func ConfigureCompiler(c *jsonschema.Compiler, o *config.ValidationOptions) {
}
}

// openAPIFormats are the integer formats in the OpenAPI format registry, which JSON Schema
// does not define. See https://spec.openapis.org/registry/format/
var openAPIFormats = []*jsonschema.Format{
{Name: "int32", Validate: integerFormat(math.MinInt32, math.MaxInt32)},
{Name: "int64", Validate: integerFormat(math.MinInt64, math.MaxInt64)},
}

// integerFormat returns a format validator that requires numbers to be whole and within the
// given bounds. Values of other types are left to the rest of the schema.
func integerFormat(minimum, maximum int64) func(any) error {
return func(v any) error {
var valid bool
switch n := v.(type) {
case int64:
valid = n >= minimum && n <= maximum
case float64:
valid = wholeNumberInRange(n, minimum, maximum)
case json.Number:
i, err := n.Int64()
if err == nil {
valid = i >= minimum && i <= maximum
} else if !errors.Is(err, strconv.ErrRange) {
// not an integer literal, but "1.0" and "1e3" are whole numbers
f, floatErr := n.Float64()
valid = floatErr == nil && wholeNumberInRange(f, minimum, maximum)
}
default:
return nil
}
if !valid {
return fmt.Errorf("must be a whole number from %d to %d", minimum, maximum)
}
return nil
}
}

// wholeNumberInRange reports whether f is a whole number from minimum to maximum. The upper bound
// is compared as maximum+1, exclusive, which stays exact when float64(maximum) rounds up (int64).
func wholeNumberInRange(f float64, minimum, maximum int64) bool {
return f == math.Trunc(f) && f >= float64(minimum) && f < float64(maximum)+1
}

// NewCompilerWithOptions mints a new JSON schema compiler with custom configuration.
func NewCompilerWithOptions(o *config.ValidationOptions) *jsonschema.Compiler {
c := jsonschema.NewCompiler()
Expand Down
49 changes: 49 additions & 0 deletions helpers/schema_compiler_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ package helpers
import (
"encoding/json"
"fmt"
"math"
"testing"
"unicode"

Expand Down Expand Up @@ -981,3 +982,51 @@ func TestTransformNullableSchema_EnumWithNull(t *testing.T) {
}
assert.Equal(t, 1, nullCount, "enum should contain exactly one null value")
}

func TestNewCompiledSchema_OpenAPIIntegerFormats(t *testing.T) {
schema := []byte(`{"type": "object", "properties": {"small": {"format": "int32"}, "large": {"format": "int64"}}}`)

asserting, err := NewCompiledSchema("formats", schema, config.NewValidationOptions(config.WithFormatAssertions()))
require.NoError(t, err)
annotating, err := NewCompiledSchema("formats", schema, config.NewValidationOptions())
require.NoError(t, err)

for _, test := range []struct {
property string
value any
valid bool
}{
{"small", int64(math.MaxInt32), true},
{"small", int64(math.MaxInt32) + 1, false},
{"small", float64(math.MinInt32), true},
{"small", float64(math.MinInt32) - 1, false},
{"small", 1.5, false},
{"small", json.Number("7.0"), true},
{"small", json.Number("2147483648"), false},
{"small", "not a number", true},
{"large", int64(math.MaxInt64), true},
{"large", json.Number("9223372036854775807"), true},
{"large", json.Number("9223372036854775808"), false},
{"large", json.Number("9223372036854775808.0"), false},
{"large", float64(1 << 63), false},
{"large", float64(math.MinInt64), true},
{"large", json.Number("1e19"), false},
{"large", json.Number("not-a-number"), false},
} {
t.Run(fmt.Sprintf("%s=%v", test.property, test.value), func(t *testing.T) {
instance := map[string]any{test.property: test.value}
assert.Equal(t, test.valid, asserting.Validate(instance) == nil)
assert.NoError(t, annotating.Validate(instance), "formats are annotations unless asserted")
})
}
}

func TestNewCompiledSchema_CustomFormatReplacesOpenAPIFormat(t *testing.T) {
schema := []byte(`{"format": "int32"}`)
options := config.NewValidationOptions(config.WithFormatAssertions(),
config.WithCustomFormat("int32", func(any) error { return nil }))

compiled, err := NewCompiledSchema("custom", schema, options)
require.NoError(t, err)
assert.NoError(t, compiled.Validate(int64(math.MaxInt64)))
}
5 changes: 5 additions & 0 deletions internal/requeststate/route.go
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,11 @@ func Route(request *http.Request) *router.Route {
return route
}

// WithRoute returns a shallow copy of request that carries route. The request itself is not changed.
func WithRoute(request *http.Request, route *router.Route) *http.Request {
return request.WithContext(context.WithValue(request.Context(), routeContextKey{}, route))
}

// AttachRoute scopes a resolved route to a request and returns an idempotent restoration function.
func AttachRoute(request *http.Request, route *router.Route) func() {
if request == nil || route == nil {
Expand Down
Loading
Loading