diff --git a/packet.go b/packet.go index 6d2c7f3..a4c10f8 100644 --- a/packet.go +++ b/packet.go @@ -423,16 +423,20 @@ func (h *Header) SetExtension(id uint8, payload []byte) error { //nolint:gocogni return nil } - // No existing header extensions - h.Extension = true + // No existing header extensions. The one byte profile can only carry + // ids 1-14 and payloads of 1-16 bytes, everything else needs two bytes. + var profile uint16 = ExtensionProfileOneByte + if id > 14 || len(payload) == 0 || len(payload) > 16 { + profile = ExtensionProfileTwoByte + } - switch payloadLen := len(payload); { - case payloadLen <= 16: - h.ExtensionProfile = ExtensionProfileOneByte - case payloadLen > 16 && payloadLen < 256: - h.ExtensionProfile = ExtensionProfileTwoByte + // Don't mutate the header if Set is going to fail anyway + if err := headerExtensionCheck(profile, id, payload); err != nil { + return err } + h.Extension = true + h.ExtensionProfile = profile h.Extensions = append(h.Extensions, Extension{id: id, payload: payload}) return nil diff --git a/packet_test.go b/packet_test.go index a607449..11586c9 100644 --- a/packet_test.go +++ b/packet_test.go @@ -1708,3 +1708,48 @@ func BenchmarkUnmarshalHeader(b *testing.B) { } }) } + +func TestSetExtensionFirstExtension(t *testing.T) { + t.Run("selects profile that can carry the extension", func(t *testing.T) { + cases := map[string]struct { + id uint8 + payload []byte + profile uint16 + }{ + "short payload": {1, []byte{0xAA}, ExtensionProfileOneByte}, + "16 byte": {14, make([]byte, 16), ExtensionProfileOneByte}, + "17 byte": {1, make([]byte, 17), ExtensionProfileTwoByte}, + "255 byte": {1, make([]byte, 255), ExtensionProfileTwoByte}, + "id above 14": {15, []byte{0xAA}, ExtensionProfileTwoByte}, + "empty payload": {1, []byte{}, ExtensionProfileTwoByte}, + "max id": {255, []byte{0xAA}, ExtensionProfileTwoByte}, + } + for name, testCase := range cases { + t.Run(name, func(t *testing.T) { + header := Header{Version: 2} + assert.NoError(t, header.SetExtension(testCase.id, testCase.payload)) + assert.Equal(t, testCase.profile, header.ExtensionProfile) + + raw, err := header.Marshal() + assert.NoError(t, err) + + var parsed Header + _, err = parsed.Unmarshal(raw) + assert.NoError(t, err) + assert.Equal(t, []uint8{testCase.id}, parsed.GetExtensionIDs()) + assert.Len(t, parsed.GetExtension(testCase.id), len(testCase.payload)) + }) + } + }) + + t.Run("rejects extensions that cannot be encoded", func(t *testing.T) { + header := Header{Version: 2} + assert.Error(t, header.SetExtension(1, make([]byte, 256))) + assert.Error(t, header.SetExtension(0, []byte{0xAA})) + + // The failed calls must not leave a partial extension behind. + assert.False(t, header.Extension) + assert.Zero(t, header.ExtensionProfile) + assert.Empty(t, header.Extensions) + }) +}