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
66 changes: 44 additions & 22 deletions conn.go
Original file line number Diff line number Diff line change
Expand Up @@ -250,7 +250,7 @@ func createConn(
maximumTransmissionUnit: mtu,
paddingLengthGenerator: paddingLengthGenerator,

decrypted: make(chan any, 1),
decrypted: make(chan any, 16),
log: logger,

readDeadline: deadline.New(),
Expand All @@ -275,6 +275,26 @@ func createConn(
return conn, nil
}

func (c *Conn) Discard() int {
n := 0
for {
select {
case out, ok := <-c.decrypted:
if !ok {
return 0
}
switch out.(type) {
case (error):
return 0
default:
n++
}
default:
return n
}
}
}

// Handshake runs the client or server DTLS handshake
// protocol if it has not yet been run.
//
Expand Down Expand Up @@ -410,11 +430,6 @@ func serverWithConfig(conn net.PacketConn, rAddr net.Addr, config *Config) (*Con
if config == nil {
return nil, errNoConfigProvided
}
if config.OnConnectionAttempt != nil {
if err := config.OnConnectionAttempt(rAddr); err != nil {
return nil, err
}
}

return createConn(conn, rAddr, config, false, nil)
}
Expand Down Expand Up @@ -588,7 +603,7 @@ func (c *Conn) prepareRawPackets(pkts []*packet) ([][]byte, net.Addr, error) {
c.lock.Lock()
defer c.lock.Unlock()

var rawPackets [][]byte
rawPackets := make([][]byte, 0, len(pkts))

for _, pkt := range pkts {
pktRawPackets, err := c.prepareRawPacket(pkt)
Expand Down Expand Up @@ -712,11 +727,16 @@ func (c *Conn) compactRawPackets(rawPackets [][]byte) [][]byte {
return rawPackets
}

combinedRawPackets := make([][]byte, 0)
currentCombinedRawPacket := make([]byte, 0)
combinedRawPackets := make([][]byte, 0, len(rawPackets))
var currentCombinedRawPacket []byte

for _, rawPacket := range rawPackets {
if len(currentCombinedRawPacket) > 0 && len(currentCombinedRawPacket)+len(rawPacket) >= c.maximumTransmissionUnit {
if len(currentCombinedRawPacket) == 0 && len(rawPacket) >= c.maximumTransmissionUnit {
combinedRawPackets = append(combinedRawPackets, rawPacket)

continue
} else if len(currentCombinedRawPacket) > 0 &&
len(currentCombinedRawPacket)+len(rawPacket) >= c.maximumTransmissionUnit {
combinedRawPackets = append(combinedRawPackets, currentCombinedRawPacket)
currentCombinedRawPacket = []byte{}
}
Expand Down Expand Up @@ -795,8 +815,6 @@ func (c *Conn) processPacket(pkt *packet) ([]byte, error) { //nolint:cyclop

//nolint:cyclop
func (c *Conn) processHandshakePacket(pkt *packet, dtlsHandshake *handshake.Handshake) ([][]byte, error) {
rawPackets := make([][]byte, 0)

handshakeFragments, err := c.fragmentHandshake(dtlsHandshake)
if err != nil {
return nil, err
Expand All @@ -806,6 +824,7 @@ func (c *Conn) processHandshakePacket(pkt *packet, dtlsHandshake *handshake.Hand
c.state.localSequenceNumber = append(c.state.localSequenceNumber, uint64(0))
}

rawPackets := make([][]byte, 0, len(handshakeFragments))
for _, handshakeFragment := range handshakeFragments {
seq := atomic.AddUint64(&c.state.localSequenceNumber[epoch], 1) - 1
if seq > recordlayer.MaxSequenceNumber {
Expand All @@ -831,12 +850,14 @@ func (c *Conn) processHandshakePacket(pkt *packet, dtlsHandshake *handshake.Hand
ConnectionID: c.state.remoteConnectionID,
SequenceNumber: pkt.record.Header.SequenceNumber,
}
rawPacket, err = cidHeader.Marshal()

rawPacket = make([]byte, cidHeader.MarshalSize()+len(rawInner))
_, err = cidHeader.MarshalTo(rawPacket)
if err != nil {
return nil, err
}
pkt.record.Header = *cidHeader
rawPacket = append(rawPacket, rawInner...)
copy(rawPacket[cidHeader.MarshalSize():], rawInner)
} else {
recordlayerHeader := &recordlayer.Header{
Version: pkt.record.Header.Version,
Expand All @@ -846,13 +867,14 @@ func (c *Conn) processHandshakePacket(pkt *packet, dtlsHandshake *handshake.Hand
SequenceNumber: seq,
}

rawPacket, err = recordlayerHeader.Marshal()
rawPacket = make([]byte, recordlayerHeader.MarshalSize()+len(handshakeFragment))
_, err = recordlayerHeader.MarshalTo(rawPacket)
if err != nil {
return nil, err
}

pkt.record.Header = *recordlayerHeader
rawPacket = append(rawPacket, handshakeFragment...)
copy(rawPacket[recordlayerHeader.MarshalSize():], handshakeFragment)
}

if pkt.shouldEncrypt {
Expand All @@ -875,8 +897,6 @@ func (c *Conn) fragmentHandshake(dtlsHandshake *handshake.Handshake) ([][]byte,
return nil, err
}

fragmentedHandshakes := make([][]byte, 0)

contentFragments := splitBytes(content, c.maximumTransmissionUnit)
if len(contentFragments) == 0 {
contentFragments = [][]byte{
Expand All @@ -885,6 +905,7 @@ func (c *Conn) fragmentHandshake(dtlsHandshake *handshake.Handshake) ([][]byte,
}

offset := 0
fragmentedHandshakes := make([][]byte, 0, len(contentFragments))
for _, contentFragment := range contentFragments {
contentFragmentLen := len(contentFragment)

Expand All @@ -898,12 +919,13 @@ func (c *Conn) fragmentHandshake(dtlsHandshake *handshake.Handshake) ([][]byte,

offset += contentFragmentLen

fragmentedHandshake, err := headerFragment.Marshal()
fragmentedHandshake := make([]byte, handshake.HeaderLength+len(contentFragment))
_, err := headerFragment.MarshalTo(fragmentedHandshake)
if err != nil {
return nil, err
}

fragmentedHandshake = append(fragmentedHandshake, contentFragment...)
copy(fragmentedHandshake[handshake.HeaderLength:], contentFragment)
fragmentedHandshakes = append(fragmentedHandshakes, fragmentedHandshake)
}

Expand Down Expand Up @@ -981,7 +1003,7 @@ func (c *Conn) readAndBuffer(ctx context.Context) error { //nolint:cyclop
func (c *Conn) handleQueuedPackets(ctx context.Context) error {
c.lock.Lock()
pkts := c.encryptedPackets
c.encryptedPackets = nil
c.encryptedPackets = c.encryptedPackets[:0]
c.lock.Unlock()

for _, p := range pkts {
Expand Down Expand Up @@ -1112,7 +1134,7 @@ func (c *Conn) handleIncomingPacket(
if header.ContentType == protocol.ContentTypeConnectionID {
originalCID = true
ip := &recordlayer.InnerPlaintext{}
if err := ip.Unmarshal(buf[header.Size():]); err != nil { //nolint:govet
if err := ip.Unmarshal(buf[header.MarshalSize():]); err != nil { //nolint:govet
c.log.Debugf("unpacking inner plaintext failed: %s", err)

return false, false, nil, nil
Expand Down
2 changes: 2 additions & 0 deletions errors.go
Original file line number Diff line number Diff line change
Expand Up @@ -143,6 +143,8 @@ var (
//nolint:err113
errFailedToAccessPoolReadBuffer = &InternalError{Err: errors.New("failed to access pool read buffer")}
//nolint:err113
errFailedToAccessPoolTimer = &InternalError{Err: errors.New("failed to access pool timer")}
//nolint:err113
errFragmentBufferOverflow = &InternalError{Err: errors.New("fragment buffer overflow")}

//nolint:err113
Expand Down
8 changes: 4 additions & 4 deletions go.mod
Original file line number Diff line number Diff line change
Expand Up @@ -4,18 +4,18 @@ require (
github.com/pion/logging v0.2.4
github.com/pion/transport/v4 v4.0.2
github.com/stretchr/testify v1.11.1
golang.org/x/crypto v0.48.0
golang.org/x/net v0.49.0
golang.org/x/crypto v0.50.0
golang.org/x/net v0.53.0
)

require (
github.com/davecgh/go-spew v1.1.1 // indirect
github.com/pmezard/go-difflib v1.0.0 // indirect
golang.org/x/sys v0.41.0 // indirect
golang.org/x/sys v0.43.0 // indirect
gopkg.in/yaml.v3 v3.0.1 // indirect
)

go 1.24.0
go 1.25.0

// Retract version with broken RSA interop with OpenSSL DTLS 1.2.
retract v3.1.0
Expand Down
12 changes: 6 additions & 6 deletions go.sum
Original file line number Diff line number Diff line change
Expand Up @@ -8,12 +8,12 @@ github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZb
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
golang.org/x/crypto v0.48.0 h1:/VRzVqiRSggnhY7gNRxPauEQ5Drw9haKdM0jqfcCFts=
golang.org/x/crypto v0.48.0/go.mod h1:r0kV5h3qnFPlQnBSrULhlsRfryS2pmewsg+XfMgkVos=
golang.org/x/net v0.49.0 h1:eeHFmOGUTtaaPSGNmjBKpbng9MulQsJURQUAfUwY++o=
golang.org/x/net v0.49.0/go.mod h1:/ysNB2EvaqvesRkuLAyjI1ycPZlQHM3q01F02UY/MV8=
golang.org/x/sys v0.41.0 h1:Ivj+2Cp/ylzLiEU89QhWblYnOE9zerudt9Ftecq2C6k=
golang.org/x/sys v0.41.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
golang.org/x/crypto v0.50.0 h1:zO47/JPrL6vsNkINmLoo/PH1gcxpls50DNogFvB5ZGI=
golang.org/x/crypto v0.50.0/go.mod h1:3muZ7vA7PBCE6xgPX7nkzzjiUq87kRItoJQM1Yo8S+Q=
golang.org/x/net v0.53.0 h1:d+qAbo5L0orcWAr0a9JweQpjXF19LMXJE8Ey7hwOdUA=
golang.org/x/net v0.53.0/go.mod h1:JvMuJH7rrdiCfbeHoo3fCQU24Lf5JJwT9W3sJFulfgs=
golang.org/x/sys v0.43.0 h1:Rlag2XtaFTxp19wS8MXlJwTvoh8ArU6ezoyFsMyCTNI=
golang.org/x/sys v0.43.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
Expand Down
22 changes: 20 additions & 2 deletions handshaker.go
Original file line number Diff line number Diff line change
Expand Up @@ -191,7 +191,8 @@ func (s *handshakeFSM) Run(ctx context.Context, conn flightConn, initialState ha
close(s.closed)
}()
for {
s.cfg.log.Tracef("[handshake:%s] %s: %s", srvCliStr(s.state.isClient), s.currentFlight.String(), state.String())
s.cfg.log.Tracef("[handshake:%s] %s: %s",
srvCliStr(s.state.isClient), s.currentFlight.String(), state.String())
if s.cfg.onFlightState != nil {
s.cfg.onFlightState(s.currentFlight, state)
}
Expand Down Expand Up @@ -279,6 +280,15 @@ func (s *handshakeFSM) send(ctx context.Context, c flightConn) (handshakeState,
return handshakeWaiting, nil
}

var timerPool = sync.Pool{ //nolint:gochecknoglobals
New: func() any {
t := time.NewTimer(time.Millisecond)
t.Stop()

return t
},
}

func (s *handshakeFSM) wait(ctx context.Context, conn flightConn) (handshakeState, error) { //nolint:gocognit,cyclop
parse, errFlight := s.currentFlight.getFlightParser()
if errFlight != nil {
Expand All @@ -289,7 +299,15 @@ func (s *handshakeFSM) wait(ctx context.Context, conn flightConn) (handshakeStat
return handshakeErrored, errFlight
}

retransmitTimer := time.NewTimer(s.retransmitInterval)
retransmitTimer, ok := timerPool.Get().(*time.Timer)
if !ok {
return handshakeErrored, errFailedToAccessPoolTimer
}
defer func() {
retransmitTimer.Stop()
timerPool.Put(retransmitTimer)
}()
retransmitTimer.Reset(s.retransmitInterval)
for {
select {
case state := <-conn.recvHandshake():
Expand Down
Loading
Loading