Skip to content

Commit 1870385

Browse files
committed
Optimize allocations for handshake and compaction
1 parent 2a3670e commit 1870385

5 files changed

Lines changed: 145 additions & 61 deletions

File tree

conn.go

Lines changed: 115 additions & 54 deletions
Original file line numberDiff line numberDiff line change
@@ -573,7 +573,7 @@ func (c *Conn) Write(payload []byte) (int, error) {
573573
return 0, err
574574
}
575575

576-
ctx, cancel := c.contextWithClose(c.writeDeadline)
576+
ctx, cancel := c.contextWithClose()
577577
defer cancel()
578578

579579
err := c.writeApplicationData(ctx, []*dtlsflight.Outbound{
@@ -829,62 +829,94 @@ func (c *Conn) cacheHandshake(outbound *dtlsflight.Outbound, dtlsHandshake *hand
829829
return nil
830830
}
831831

832-
func (c *Conn) contextWithClose(ctx context.Context) (context.Context, context.CancelFunc) {
833-
closeCtx, cancel := context.WithCancelCause(context.WithoutCancel(ctx))
834-
go func() {
835-
select {
836-
case <-c.closed.Done():
837-
cancel(context.Canceled)
838-
case <-ctx.Done():
839-
err := ctx.Err()
840-
if err == nil {
841-
err = context.DeadlineExceeded
842-
}
843-
cancel(err)
844-
case <-closeCtx.Done():
845-
}
846-
}()
832+
func (c *Conn) contextWithClose() (context.Context, context.CancelFunc) {
833+
ctx := context.Background()
834+
835+
var cancelDeadline context.CancelFunc = func() {}
836+
if deadline, ok := c.readDeadline.Deadline(); ok {
837+
ctx, cancelDeadline = context.WithDeadline(ctx, deadline)
838+
}
839+
840+
closeCtx, cancelClose := context.WithCancelCause(ctx)
841+
detachLifetime := context.AfterFunc(c.closed, func() {
842+
cancelClose(c.closed.Err())
843+
})
847844

848845
return closeCtx, func() {
849-
cancel(context.Canceled)
846+
detachLifetime()
847+
cancelDeadline()
848+
cancelClose(context.Canceled)
850849
}
851850
}
852851

853852
func (c *Conn) contextWithCloseAndWriteDeadline(ctx context.Context) (context.Context, context.CancelFunc) {
854-
operationCtx, cancel := context.WithCancelCause(context.Background())
855-
go func() {
856-
select {
857-
case <-c.closed.Done():
858-
cancel(context.Canceled)
859-
case <-c.writeDeadline.Done():
860-
cancel(context.DeadlineExceeded)
861-
case <-ctx.Done():
862-
cancel(ctx.Err())
863-
case <-operationCtx.Done():
853+
if ctx == nil {
854+
ctx = context.Background()
855+
}
856+
857+
var cancelDeadline context.CancelFunc = func() {}
858+
if deadline, ok := c.writeDeadline.Deadline(); ok {
859+
ctx, cancelDeadline = context.WithDeadline(ctx, deadline)
860+
}
861+
862+
operationCtx, cancelClose := context.WithCancelCause(ctx)
863+
detachLifetime := context.AfterFunc(c.closed, func() {
864+
err := c.closed.Err()
865+
if err == nil {
866+
err = context.Canceled
864867
}
865-
}()
868+
cancelClose(err)
869+
})
866870

867871
return operationCtx, func() {
868-
cancel(context.Canceled)
872+
detachLifetime()
873+
cancelDeadline()
874+
cancelClose(context.Canceled)
869875
}
870876
}
871877

872878
func (c *Conn) compactPreparedRecords(records []preparedRecord) []preparedDatagram {
873-
datagrams := make([]preparedDatagram, 0, len(records))
874-
current := preparedDatagram{}
879+
if len(records) == 0 {
880+
return []preparedDatagram{}
881+
}
882+
883+
totalSize := 0
875884
for _, record := range records {
876-
if len(current.raw) > 0 && len(current.raw)+len(record.raw) >= c.maximumTransmissionUnit {
877-
datagrams = append(datagrams, current)
878-
current = preparedDatagram{}
885+
totalSize += len(record.raw)
886+
}
887+
888+
datagrams := make([]preparedDatagram, len(records))
889+
flatRaw := make([]byte, 0, totalSize)
890+
891+
datagramIndex := 0
892+
currentSize := 0
893+
offset := 0
894+
895+
for _, record := range records {
896+
recordSize := len(record.raw)
897+
898+
flatRaw = append(flatRaw, record.raw...)
899+
900+
if currentSize > 0 && currentSize+recordSize >= c.maximumTransmissionUnit {
901+
datagrams[datagramIndex].raw = flatRaw[offset : offset+currentSize]
902+
datagramIndex++
903+
offset += currentSize
904+
currentSize = 0
879905
}
880-
current.raw = append(current.raw, record.raw...)
906+
907+
currentSize += recordSize
908+
881909
if record.tracked != nil {
882-
current.tracked = append(current.tracked, *record.tracked)
910+
datagrams[datagramIndex].tracked = append(
911+
datagrams[datagramIndex].tracked,
912+
*record.tracked,
913+
)
883914
}
884915
}
885-
datagrams = append(datagrams, current)
886916

887-
return datagrams
917+
datagrams[datagramIndex].raw = flatRaw[offset : offset+currentSize]
918+
919+
return datagrams[:datagramIndex+1]
888920
}
889921

890922
func (c *Conn) prepareRecord(outbound *dtlsflight.Outbound) ([]byte, error) {
@@ -1214,42 +1246,76 @@ func selectHandshakeFragment(offsets map[uint32]uint32, raw []byte) (bool, error
12141246
return ok && length == header.FragmentLength, nil
12151247
}
12161248

1249+
var noContentFragments = [][]byte{ //nolint:gochecknoglobals
1250+
{},
1251+
}
1252+
12171253
func (c *Conn) fragmentHandshake(dtlsHandshake *handshake.Handshake) ([][]byte, error) {
1254+
messageSize := dtlsHandshake.Message.MarshalSize()
1255+
numFragments := (messageSize-1)/c.maximumTransmissionUnit + 1
1256+
1257+
fragmentedHandshakes := make([][]byte, numFragments)
1258+
1259+
if numFragments == 1 {
1260+
fragmentedHandshake := make([]byte, handshake.HeaderLength+messageSize)
1261+
1262+
headerFragment := handshake.Header{
1263+
Type: dtlsHandshake.Header.Type,
1264+
Length: dtlsHandshake.Header.Length,
1265+
MessageSequence: dtlsHandshake.Header.MessageSequence,
1266+
FragmentOffset: uint32(0),
1267+
FragmentLength: uint32(messageSize), //nolint:gosec // G115
1268+
}
1269+
1270+
_, err := headerFragment.MarshalTo(fragmentedHandshake)
1271+
if err != nil {
1272+
return nil, err
1273+
}
1274+
1275+
_, err = dtlsHandshake.Message.MarshalTo(fragmentedHandshake[handshake.HeaderLength:])
1276+
if err != nil {
1277+
return nil, err
1278+
}
1279+
1280+
fragmentedHandshakes[0] = fragmentedHandshake
1281+
1282+
return fragmentedHandshakes, nil
1283+
}
1284+
12181285
content, err := dtlsHandshake.Message.Marshal()
12191286
if err != nil {
12201287
return nil, err
12211288
}
12221289

1223-
fragmentedHandshakes := make([][]byte, 0)
1224-
12251290
contentFragments := util.SplitBytes(content, c.maximumTransmissionUnit)
12261291
if len(contentFragments) == 0 {
1227-
contentFragments = [][]byte{
1228-
{},
1229-
}
1292+
contentFragments = noContentFragments
12301293
}
12311294

12321295
offset := 0
1233-
for _, contentFragment := range contentFragments {
1296+
for i, contentFragment := range contentFragments {
12341297
contentFragmentLen := len(contentFragment)
12351298

1236-
headerFragment := &handshake.Header{
1299+
headerFragment := handshake.Header{
12371300
Type: dtlsHandshake.Header.Type,
12381301
Length: dtlsHandshake.Header.Length,
12391302
MessageSequence: dtlsHandshake.Header.MessageSequence,
12401303
FragmentOffset: uint32(offset),
12411304
FragmentLength: uint32(contentFragmentLen), //nolint:gosec // G115
12421305
}
12431306

1307+
fragmentedHandshake := make([]byte, handshake.HeaderLength+contentFragmentLen)
1308+
12441309
offset += contentFragmentLen
12451310

1246-
fragmentedHandshake, err := headerFragment.Marshal()
1311+
_, err := headerFragment.MarshalTo(fragmentedHandshake)
12471312
if err != nil {
12481313
return nil, err
12491314
}
12501315

1251-
fragmentedHandshake = append(fragmentedHandshake, contentFragment...)
1252-
fragmentedHandshakes = append(fragmentedHandshakes, fragmentedHandshake)
1316+
copy(fragmentedHandshake[handshake.HeaderLength:], contentFragment)
1317+
1318+
fragmentedHandshakes[i] = fragmentedHandshake
12531319
}
12541320

12551321
return fragmentedHandshakes, nil
@@ -2756,12 +2822,7 @@ func (c *Conn) close(byUser bool) error {
27562822
}
27572823

27582824
func (c *Conn) isConnectionClosed() bool {
2759-
select {
2760-
case <-c.closed.Done():
2761-
return true
2762-
default:
2763-
return false
2764-
}
2825+
return c.closed.Err() != nil
27652826
}
27662827

27672828
func (c *Conn) setLocalEpoch(epoch uint16) {

internal/closer/closer.go

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@ package closer
66

77
import (
88
"context"
9+
"time"
910
)
1011

1112
// Closer allows for each signaling a channel for shutdown.
@@ -48,3 +49,11 @@ func (c *Closer) Err() error {
4849
func (c *Closer) Close() {
4950
c.closeFunc()
5051
}
52+
53+
func (c *Closer) Deadline() (deadline time.Time, ok bool) {
54+
return c.ctx.Deadline()
55+
}
56+
57+
func (c *Closer) Value(key any) any {
58+
return c.ctx.Value(key)
59+
}

internal/flight/cache.go

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,6 @@
55
package flight
66

77
import (
8-
"bytes"
98
"fmt"
109
"sync"
1110

@@ -45,7 +44,7 @@ func (h *Cache) Push(data []byte, epoch, messageSequence uint16, typ handshake.T
4544
h.mu.Lock()
4645
defer h.mu.Unlock()
4746

48-
h.cache = append(h.cache, &HandshakeCacheItem{Data: bytes.Clone(data), Epoch: epoch, MessageSequence: messageSequence, Typ: typ, IsClient: isClient})
47+
h.cache = append(h.cache, &HandshakeCacheItem{Data: data, Epoch: epoch, MessageSequence: messageSequence, Typ: typ, IsClient: isClient})
4948
}
5049

5150
// PullExact returns the handshake message with the requested sequence and

internal/handshake/fsm12.go

Lines changed: 17 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@ package dtlshandshake
55

66
import (
77
"context"
8+
"sync"
89
"time"
910

1011
dtlsconfig "github.com/pion/dtls/v3/internal/config"
@@ -147,8 +148,23 @@ func (s *fsm12) send(ctx context.Context, c Conn) (State, error) {
147148
return StateWaiting, nil
148149
}
149150

151+
var timerPool = sync.Pool{ //nolint:gochecknoglobals
152+
New: func() any {
153+
timer := time.NewTimer(1 * time.Hour)
154+
timer.Stop()
155+
156+
return timer
157+
},
158+
}
159+
150160
func (s *fsm12) wait(ctx context.Context, conn Conn) (State, error) { //nolint:gocognit,cyclop
151-
retransmitTimer := time.NewTimer(s.retransmitInterval)
161+
retransmitTimer := timerPool.Get().(*time.Timer) //nolint:forcetypeassert
162+
retransmitTimer.Reset(s.retransmitInterval)
163+
defer func() {
164+
retransmitTimer.Stop()
165+
timerPool.Put(retransmitTimer)
166+
}()
167+
152168
for {
153169
select {
154170
case state := <-conn.RecvHandshake():

pkg/protocol/handshake/handshake.go

Lines changed: 3 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -131,19 +131,18 @@ func (h *Handshake) Marshal() ([]byte, error) {
131131
return nil, dtlserrors.ErrUnableToMarshalFragmented
132132
}
133133

134-
message, err := h.Message.Marshal()
134+
out := make([]byte, HeaderLength+h.Message.MarshalSize())
135+
n, err := h.Message.MarshalTo(out[HeaderLength:])
135136
if err != nil {
136137
return nil, err
137138
}
138139

139-
out := make([]byte, HeaderLength+len(message))
140-
h.Header.Length = uint32(len(message)) //nolint:gosec // handshake messages are bounded to uint24 on the wire.
140+
h.Header.Length = uint32(n) //nolint:gosec // handshake messages are bounded to uint24 on the wire.
141141
h.Header.FragmentLength = h.Header.Length
142142
h.Header.Type = h.Message.Type()
143143
if _, err = h.Header.MarshalTo(out); err != nil {
144144
return nil, err
145145
}
146-
copy(out[HeaderLength:], message)
147146

148147
return out, nil
149148
}

0 commit comments

Comments
 (0)