src/encoding/asn1/asn1.go | 21 +++++++++++++++------ src/encoding/asn1/asn1_test.go | 87 +++++++++++++++++++++++++++++++++++++++++++++++++++++ diff --git a/src/encoding/asn1/asn1.go b/src/encoding/asn1/asn1.go index f4be515b98ef1cc46b8a57c633511257a4b4fb90..4e65924c0f1d6f62cc2668b6db74784e82b93a53 100644 --- a/src/encoding/asn1/asn1.go +++ b/src/encoding/asn1/asn1.go @@ -26,6 +26,7 @@ "internal/saferio" "math" "math/big" "reflect" + "runtime" "slices" "strconv" "strings" @@ -629,7 +630,7 @@ // parseSequenceOf is used for SEQUENCE OF and SET OF values. It tries to parse // a number of ASN.1 values from the given byte slice and returns them as a // slice of Go values of the given type. -func parseSequenceOf(bytes []byte, sliceType reflect.Type, elemType reflect.Type) (ret reflect.Value, err error) { +func parseSequenceOf(bytes []byte, sliceType reflect.Type, elemType reflect.Type, depth int) (ret reflect.Value, err error) { matchAny, expectedTag, compoundType, ok := getUniversalType(elemType) if !ok { err = StructuralError{"unknown Go type for slice"} @@ -678,7 +679,7 @@ params := fieldParameters{} offset := 0 for i := 0; i < numElements; i++ { ret = reflect.Append(ret, reflect.Zero(elemType)) - offset, err = parseField(ret.Index(i), bytes, offset, params) + offset, err = parseField(ret.Index(i), bytes, offset, params, depth) if err != nil { return } @@ -706,7 +707,15 @@ // parseField is the main parsing function. Given a byte slice and an offset // into the array, it will try to parse a suitable ASN.1 value out and store it // in the given Value. -func parseField(v reflect.Value, bytes []byte, initOffset int, params fieldParameters) (offset int, err error) { +func parseField(v reflect.Value, bytes []byte, initOffset int, params fieldParameters, depth int) (offset int, err error) { + depth++ + const ( + maxDecodeDepth = 10000 + maxDecodeDepthWasm = 5000 // go.dev/issue/56498 + ) + if depth > maxDecodeDepth || runtime.GOARCH == "wasm" && depth > maxDecodeDepthWasm { + return initOffset, StructuralError{"nesting depth exceeded"} + } offset = initOffset fieldType := v.Type() @@ -977,7 +986,7 @@ field := structType.Field(i) if i == 0 && field.Type == rawContentsType { continue } - innerOffset, err = parseField(val.Field(i), innerBytes, innerOffset, parseFieldParameters(field.Tag.Get("asn1"))) + innerOffset, err = parseField(val.Field(i), innerBytes, innerOffset, parseFieldParameters(field.Tag.Get("asn1")), depth) if err != nil { return } @@ -993,7 +1002,7 @@ val.Set(reflect.MakeSlice(sliceType, len(innerBytes), len(innerBytes))) reflect.Copy(val, reflect.ValueOf(innerBytes)) return } - newSlice, err1 := parseSequenceOf(innerBytes, sliceType, sliceType.Elem()) + newSlice, err1 := parseSequenceOf(innerBytes, sliceType, sliceType.Elem(), depth) if err1 == nil { val.Set(newSlice) } @@ -1165,7 +1174,7 @@ v := reflect.ValueOf(val) if v.Kind() != reflect.Pointer || v.IsNil() { return nil, &invalidUnmarshalError{reflect.TypeOf(val)} } - offset, err := parseField(v.Elem(), b, 0, parseFieldParameters(params)) + offset, err := parseField(v.Elem(), b, 0, parseFieldParameters(params), 0) if err != nil { return nil, err } diff --git a/src/encoding/asn1/asn1_test.go b/src/encoding/asn1/asn1_test.go index 41cc0ba50ec304fedf6b87acb506eb92a336f6ff..1aaab0eb9f7975d01944c683de5bc85a95f22050 100644 --- a/src/encoding/asn1/asn1_test.go +++ b/src/encoding/asn1/asn1_test.go @@ -1254,3 +1254,90 @@ if memDiff > 10<<21 { t.Errorf("Too much memory allocated while parsing DER: %v MiB", memDiff/1024/1024) } } + +func TestUnmarshalNestingLimitSlice(t *testing.T) { + type Recursive []Recursive + + limit := 10000 + if runtime.GOARCH == "wasm" { + limit = 5000 + } + + makeData := func(t *testing.T, depth int) []byte { + var r Recursive + for range depth - 1 { + r = Recursive{r} + } + data, err := Marshal(r) + if err != nil { + t.Fatalf("Marshal failed: %v", err) + } + return data + } + + t.Run("below limit", func(t *testing.T) { + data := makeData(t, limit) + var r Recursive + if _, err := Unmarshal(data, &r); err != nil { + t.Errorf("Unmarshal failed at depth %d: %v", limit, err) + } + }) + + t.Run("above limit", func(t *testing.T) { + data := makeData(t, limit+1) + var r Recursive + _, err := Unmarshal(data, &r) + if err == nil { + t.Fatalf("Unmarshal succeeded at depth %d, want error", limit+1) + } + if got, want := err.Error(), "asn1: structure error: nesting depth exceeded"; got != want { + t.Errorf("Unmarshal error mismatch\ngot: %q\nwant: %q", got, want) + } + }) +} + +// Note that recursive structs fail in half the normal limit because each level +// of nesting in a struct (with a slice field) involves two depth increments +// (one for the struct and one for the slice). +func TestUnmarshalNestingLimitStruct(t *testing.T) { + type Recursive struct { + Next []Recursive `asn1:"optional"` + } + + limit := 5000 + if runtime.GOARCH == "wasm" { + limit = 2500 + } + + makeData := func(t *testing.T, depth int) []byte { + var r Recursive + for range depth - 1 { + r = Recursive{Next: []Recursive{r}} + } + data, err := Marshal(r) + if err != nil { + t.Fatalf("Marshal failed: %v", err) + } + return data + } + + t.Run("below limit", func(t *testing.T) { + data := makeData(t, limit) + var r Recursive + if _, err := Unmarshal(data, &r); err != nil { + t.Errorf("Unmarshal failed at depth %d: %v", limit, err) + } + }) + + t.Run("above limit", func(t *testing.T) { + data := makeData(t, limit+1) + var r Recursive + _, err := Unmarshal(data, &r) + if err == nil { + t.Fatalf("Unmarshal succeeded at depth %d, want error", limit+1) + } + if got, want := err.Error(), "asn1: structure error: nesting depth exceeded"; got != want { + t.Errorf("Unmarshal error mismatch\ngot: %q\nwant: %q", got, want) + } + }) +}