From 0e013c7984a2f1b324b60e4c820fe6295932f14c Mon Sep 17 00:00:00 2001 From: Konrad Wojas Date: Wed, 7 Jan 2026 16:23:43 +0800 Subject: [PATCH] feat: add Decoder.SetMaxFieldLength (#207) add Decoder.SetMaxFieldLength method to allow overriding the default field limit of 2GB. --- decoder.go | 29 +++++++++++++++++++++-------- decoder_test.go | 35 +++++++++++++++++++++++++++++++++++ 2 files changed, 56 insertions(+), 8 deletions(-) diff --git a/decoder.go b/decoder.go index 07a1d56..4de34d6 100644 --- a/decoder.go +++ b/decoder.go @@ -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, } } @@ -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. @@ -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 @@ -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 @@ -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 diff --git a/decoder_test.go b/decoder_test.go index 1ef8d6e..6df1096 100644 --- a/decoder_test.go +++ b/decoder_test.go @@ -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