Skip to content
Open
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
29 changes: 21 additions & 8 deletions decoder.go
Original file line number Diff line number Diff line change
Expand Up @@ -58,16 +58,18 @@ func (m DecoderMode) String() string {

// Decoder implements a binary Protobuf Decoder by sequentially reading from a provided []byte.
type Decoder struct {
p []byte
offset int
mode DecoderMode
p []byte
offset int
mode DecoderMode
maxFieldLen uint64
}

// NewDecoder initializes a new Protobuf decoder to read the provided buffer.
func NewDecoder(p []byte) *Decoder {
return &Decoder{
p: p,
offset: 0,
p: p,
offset: 0,
maxFieldLen: maxFieldLen,
}
}

Expand All @@ -81,6 +83,17 @@ func (d *Decoder) SetMode(m DecoderMode) {
d.mode = m
}

// SetMaxFieldLength sets the maximum size a field can have.
// This applies to bytes/string fields and messages.
// By default, the official protobuf limit of 2GB is used.
//
// You can set a different limit here, including a higher limit, but
// by doing so you deviate from the protobuf specified limits,
// and your interoperability with other implementations may be affected.
func (d *Decoder) SetMaxFieldLength(length uint64) {
d.maxFieldLen = length
}

// Seek sets the position of the next read operation to [offset], interpreted according to [whence]:
// [io.SeekStart] means relative to the start of the data, [io.SeekCurrent] means relative to the
// current offset, and [io.SeekEnd] means relative to the end.
Expand Down Expand Up @@ -195,7 +208,7 @@ func (d *Decoder) DecodeBytes() ([]byte, error) {
return nil, fmt.Errorf("invalid data at byte %d: %w", d.offset, err)
case n == 0:
return nil, fmt.Errorf("invalid data at byte %d: %w", d.offset, ErrInvalidVarintData)
case l > maxFieldLen:
case l > d.maxFieldLen:
return nil, fmt.Errorf("invalid length (%d) for length-delimited field at byte %d: %w", l, d.offset, ErrLenOverflow)
default:
// length is good
Expand Down Expand Up @@ -883,7 +896,7 @@ func (d *Decoder) DecodeNested(m interface{}) error {
return fmt.Errorf("invalid data at byte %d: %w", d.offset, err)
case n == 0:
return fmt.Errorf("invalid data at byte %d: %w", d.offset, ErrInvalidVarintData)
case l > maxFieldLen:
case l > d.maxFieldLen:
return fmt.Errorf("invalid length (%d) for length-delimited field at byte %d: %w", l, d.offset, ErrLenOverflow)
default:
// length is good
Expand Down Expand Up @@ -964,7 +977,7 @@ func (d *Decoder) Skip(tag int, wt WireType) ([]byte, error) {
return nil, fmt.Errorf("invalid data at byte %d: %w", d.offset, err)
case n == 0:
return nil, fmt.Errorf("invalid data at byte %d: %w", d.offset, ErrInvalidVarintData)
case l > maxFieldLen:
case l > d.maxFieldLen:
return nil, fmt.Errorf("invalid length (%d) for length-delimited field at byte %d: %w", l, d.offset, ErrLenOverflow)
default:
// length is good
Expand Down
35 changes: 35 additions & 0 deletions decoder_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -1348,6 +1348,41 @@ func TestDecodeTag(t *testing.T) {
}
}

func TestDecoder_SetMaxFieldLength(t *testing.T) {
payload := []byte{0x12, 0xE, 0x74, 0x68, 0x69, 0x73, 0x20, 0x69, 0x73, 0x20, 0x61, 0x20, 0x74, 0x65, 0x73, 0x74}
message := "this is a test"

dec := csproto.NewDecoder(payload)
dec.SetMaxFieldLength(5) // too small

tag, wt, err := dec.DecodeTag()
assert.Equal(t, 2, tag, "tag should match")
assert.Equal(t, csproto.WireTypeLengthDelimited, wt, "wire type should match")
assert.NoError(t, err, "should not fail")

// enforced for strings
_, err = dec.DecodeString()
assert.ErrorContains(t, err, "invalid length")

// enforced for bytes
_, err = dec.DecodeBytes()
assert.ErrorContains(t, err, "invalid length")

// enforced for nested messages (will never reach actual decoding into m)
err = dec.DecodeNested(nil)
assert.ErrorContains(t, err, "invalid length")

// enforced for skip
_, err = dec.Skip(2, csproto.WireTypeLengthDelimited)
assert.ErrorContains(t, err, "invalid length")

// now increase the limit to succeed
dec.SetMaxFieldLength(uint64(len(message))) // just large enough
got, err := dec.DecodeString()
assert.NoError(t, err)
assert.Equal(t, message, got)
}

func FuzzDecodeTag(f *testing.F) {
seedData := [][]byte{
{(1 << 3)}, // tag=1, wire type=0
Expand Down