Skip to content
Closed
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
18 changes: 11 additions & 7 deletions packet.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
45 changes: 45 additions & 0 deletions packet_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
})
}