diff --git a/ledger/common/utils/utils.go b/ledger/common/utils/utils.go index 2889317f4db..4a6da783f3b 100644 --- a/ledger/common/utils/utils.go +++ b/ledger/common/utils/utils.go @@ -75,10 +75,13 @@ func AppendLongData(input []byte, data []byte) []byte { return input } -// ReadSlice reads `size` bytes from the input -func ReadSlice(input []byte, size int) (value []byte, rest []byte, err error) { - if len(input) < size { - return nil, input, fmt.Errorf("input size is too small to be splited %d < %d ", len(input), size) +// ReadSlice reads `size` bytes from the input and returns the slice and the rest. +// +// Expected error returns during normal operation: +// - Generic error: if the input has fewer than `size` bytes remaining. +func ReadSlice(input []byte, size uint64) (value []byte, rest []byte, err error) { + if uint64(len(input)) < size { + return nil, input, fmt.Errorf("input size is too small to be split: %d < %d", len(input), size) } return input[:size], input[size:], nil } @@ -115,15 +118,20 @@ func ReadUint64(input []byte) (value uint64, rest []byte, err error) { return binary.BigEndian.Uint64(input[:8]), input[8:], nil } -// ReadShortData read data shorter than 16kB and return the rest of bytes +// ReadShortData read data shorter than 16kB and return the rest of bytes. +// +// Expected error returns during normal operation: +// - Generic error: if the input has fewer than the declared size bytes remaining. func ReadShortData(input []byte) (data []byte, rest []byte, err error) { var size uint16 size, rest, err = ReadUint16(input) if err != nil { return nil, rest, err } - data = rest[:size] - rest = rest[size:] + data, rest, err = ReadSlice(rest, uint64(size)) + if err != nil { + return nil, rest, fmt.Errorf("short data length exceeds input: %w", err) + } return } diff --git a/ledger/trie_encoder.go b/ledger/trie_encoder.go index d7bc6f98438..5204dd40bf6 100644 --- a/ledger/trie_encoder.go +++ b/ledger/trie_encoder.go @@ -288,7 +288,7 @@ func decodeKeyPartWithEncodedSizeInfo( // Read encoded key part var kpEnc []byte - kpEnc, rest, err = utils.ReadSlice(rest, int(kpEncSize)) + kpEnc, rest, err = utils.ReadSlice(rest, uint64(kpEncSize)) if err != nil { return 0, nil, nil, fmt.Errorf("error decoding key part: %w", err) } @@ -511,7 +511,7 @@ func decodePayload(inp []byte, zeroCopy bool, version uint16) (*Payload, error) } // read encoded key - ek, rest, err := utils.ReadSlice(rest, int(encKeySize)) + ek, rest, err := utils.ReadSlice(rest, uint64(encKeySize)) if err != nil { return nil, fmt.Errorf("error decoding payload: %w", err) } @@ -537,7 +537,7 @@ func decodePayload(inp []byte, zeroCopy bool, version uint16) (*Payload, error) } // read encoded value - encValue, _, err := utils.ReadSlice(rest, encValueSize) + encValue, _, err := utils.ReadSlice(rest, uint64(encValueSize)) if err != nil { return nil, fmt.Errorf("error decoding payload: %w", err) } @@ -635,7 +635,7 @@ func decodeTrieUpdate(inp []byte, version uint16) (*TrieUpdate, error) { return nil, fmt.Errorf("error decoding trie update: %w", err) } - rhBytes, rest, err := utils.ReadSlice(rest, int(rhSize)) + rhBytes, rest, err := utils.ReadSlice(rest, uint64(rhSize)) if err != nil { return nil, fmt.Errorf("error decoding trie update: %w", err) } @@ -662,7 +662,7 @@ func decodeTrieUpdate(inp []byte, version uint16) (*TrieUpdate, error) { var path Path var encPath []byte for i := 0; i < int(numOfPaths); i++ { - encPath, rest, err = utils.ReadSlice(rest, int(pathSize)) + encPath, rest, err = utils.ReadSlice(rest, uint64(pathSize)) if err != nil { return nil, fmt.Errorf("error decoding trie update: %w", err) } @@ -682,7 +682,7 @@ func decodeTrieUpdate(inp []byte, version uint16) (*TrieUpdate, error) { if err != nil { return nil, fmt.Errorf("error decoding trie update: %w", err) } - encPayload, rest, err = utils.ReadSlice(rest, int(payloadSize)) + encPayload, rest, err = utils.ReadSlice(rest, uint64(payloadSize)) if err != nil { return nil, fmt.Errorf("error decoding trie update: %w", err) } @@ -819,7 +819,7 @@ func decodeTrieProof(inp []byte, version uint16) (*TrieProof, error) { if err != nil { return nil, fmt.Errorf("error decoding proof: %w", err) } - flags, rest, err := utils.ReadSlice(rest, int(flagsSize)) + flags, rest, err := utils.ReadSlice(rest, uint64(flagsSize)) if err != nil { return nil, fmt.Errorf("error decoding proof: %w", err) } @@ -830,7 +830,7 @@ func decodeTrieProof(inp []byte, version uint16) (*TrieProof, error) { if err != nil { return nil, fmt.Errorf("error decoding proof: %w", err) } - path, rest, err := utils.ReadSlice(rest, int(pathSize)) + path, rest, err := utils.ReadSlice(rest, uint64(pathSize)) if err != nil { return nil, fmt.Errorf("error decoding proof: %w", err) } @@ -844,7 +844,7 @@ func decodeTrieProof(inp []byte, version uint16) (*TrieProof, error) { if err != nil { return nil, fmt.Errorf("error decoding proof: %w", err) } - encPayload, rest, err := utils.ReadSlice(rest, int(encPayloadSize)) + encPayload, rest, err := utils.ReadSlice(rest, uint64(encPayloadSize)) if err != nil { return nil, fmt.Errorf("error decoding proof: %w", err) } @@ -873,7 +873,7 @@ func decodeTrieProof(inp []byte, version uint16) (*TrieProof, error) { return nil, fmt.Errorf("error decoding proof: %w", err) } - interimBytes, rest, err = utils.ReadSlice(rest, int(interimSize)) + interimBytes, rest, err = utils.ReadSlice(rest, uint64(interimSize)) if err != nil { return nil, fmt.Errorf("error decoding proof: %w", err) } @@ -961,7 +961,7 @@ func decodeTrieBatchProof(inp []byte, version uint16) (*TrieBatchProof, error) { } // read encoded proof - encProof, rest, err = utils.ReadSlice(rest, int(encProofSize)) + encProof, rest, err = utils.ReadSlice(rest, uint64(encProofSize)) if err != nil { return nil, fmt.Errorf("error decoding batch proof (content): %w", err) } diff --git a/ledger/trie_encoder_test.go b/ledger/trie_encoder_test.go index f1094ffb383..271b3fea2b9 100644 --- a/ledger/trie_encoder_test.go +++ b/ledger/trie_encoder_test.go @@ -881,3 +881,42 @@ func TestTrieUpdateEncodingMethodsPreservesValueTypes(t *testing.T) { require.Equal(t, decoded2, decoded3, "EncodeTrieUpdateCBOR and EncodeTrieUpdateProtoBuf should produce same TrieUpdate") }) } + +// TestDecodeTrieBatchProofRejectsOversizedProof verifies that DecodeTrieBatchProof +// returns an error instead of panicking when a proof declares a length larger than +// the remaining input. This exercises the ReadSlice uint64 bounds check on the +// proof decoding path. +func TestDecodeTrieBatchProofRejectsOversizedProof(t *testing.T) { + t.Parallel() + + encodedBatchProofHead := []byte{ + 0x00, 0x00, // version 0 + 0x08, // type BatchProof + 0x00, 0x00, 0x00, 0x01, // number of proofs: 1 + } + + t.Run("proof size exceeds remaining input", func(t *testing.T) { + t.Parallel() + + // Declare a proof size of MaxUint64; only a few bytes follow. + encoded := append([]byte{}, encodedBatchProofHead...) + encoded = append(encoded, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff) // proof size + encoded = append(encoded, 0x00, 0x00) // incomplete proof payload + + _, err := ledger.DecodeTrieBatchProof(encoded) + require.Error(t, err) + }) + + t.Run("proof size wraps to negative on int conversion", func(t *testing.T) { + t.Parallel() + + // Declare a proof size of 1<<63, which becomes negative if cast to int on a 64-bit platform. + encoded := append([]byte{}, encodedBatchProofHead...) + encoded = append(encoded, + 0x80, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, // proof size = 1<<63 + ) + + _, err := ledger.DecodeTrieBatchProof(encoded) + require.Error(t, err) + }) +}