From 536870896329411aceb3090158d881a033b03729 Mon Sep 17 00:00:00 2001 From: Thomas Legris Date: Sat, 28 Mar 2026 20:45:11 +0900 Subject: [PATCH 1/5] Add optimizations for connection & handshake --- conn.go | 34 +++++++++++-------- errors.go | 2 ++ handshaker.go | 19 ++++++++++- pkg/protocol/application_data.go | 2 +- pkg/protocol/handshake/header.go | 12 ++++++- pkg/protocol/handshake/message_certificate.go | 6 ++++ pkg/protocol/recordlayer/header.go | 20 +++++++++-- pkg/protocol/recordlayer/recordlayer.go | 8 +++-- 8 files changed, 82 insertions(+), 21 deletions(-) diff --git a/conn.go b/conn.go index f97d3abbf..3eb3cf687 100644 --- a/conn.go +++ b/conn.go @@ -716,7 +716,12 @@ func (c *Conn) compactRawPackets(rawPackets [][]byte) [][]byte { currentCombinedRawPacket := make([]byte, 0) 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{} } @@ -795,7 +800,7 @@ 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) + var rawPackets [][]byte handshakeFragments, err := c.fragmentHandshake(dtlsHandshake) if err != nil { @@ -831,12 +836,15 @@ func (c *Conn) processHandshakePacket(pkt *packet, dtlsHandshake *handshake.Hand ConnectionID: c.state.remoteConnectionID, SequenceNumber: pkt.record.Header.SequenceNumber, } - rawPacket, err = cidHeader.Marshal() + + hs := recordlayer.FixedHeaderSize + len(cidHeader.ConnectionID) + rawPacket = make([]byte, hs+len(rawInner)) + err = cidHeader.MarshalInto(rawPacket) if err != nil { return nil, err } pkt.record.Header = *cidHeader - rawPacket = append(rawPacket, rawInner...) + copy(rawPacket[hs:], rawInner) } else { recordlayerHeader := &recordlayer.Header{ Version: pkt.record.Header.Version, @@ -846,13 +854,15 @@ func (c *Conn) processHandshakePacket(pkt *packet, dtlsHandshake *handshake.Hand SequenceNumber: seq, } - rawPacket, err = recordlayerHeader.Marshal() + hs := recordlayer.FixedHeaderSize + len(recordlayerHeader.ConnectionID) + rawPacket = make([]byte, hs+len(handshakeFragment)) + err = recordlayerHeader.MarshalInto(rawPacket) if err != nil { return nil, err } pkt.record.Header = *recordlayerHeader - rawPacket = append(rawPacket, handshakeFragment...) + copy(rawPacket[hs:], handshakeFragment) } if pkt.shouldEncrypt { @@ -875,14 +885,9 @@ func (c *Conn) fragmentHandshake(dtlsHandshake *handshake.Handshake) ([][]byte, return nil, err } - fragmentedHandshakes := make([][]byte, 0) + var fragmentedHandshakes [][]byte contentFragments := splitBytes(content, c.maximumTransmissionUnit) - if len(contentFragments) == 0 { - contentFragments = [][]byte{ - {}, - } - } offset := 0 for _, contentFragment := range contentFragments { @@ -898,12 +903,13 @@ func (c *Conn) fragmentHandshake(dtlsHandshake *handshake.Handshake) ([][]byte, offset += contentFragmentLen - fragmentedHandshake, err := headerFragment.Marshal() + fragmentedHandshake := make([]byte, handshake.HeaderLength+len(contentFragment)) + err := headerFragment.MarshalInto(fragmentedHandshake) if err != nil { return nil, err } - fragmentedHandshake = append(fragmentedHandshake, contentFragment...) + copy(fragmentedHandshake[handshake.HeaderLength:], contentFragment) fragmentedHandshakes = append(fragmentedHandshakes, fragmentedHandshake) } diff --git a/errors.go b/errors.go index 0db0de679..616d434cc 100644 --- a/errors.go +++ b/errors.go @@ -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 diff --git a/handshaker.go b/handshaker.go index 425934d47..f9ef37ddf 100644 --- a/handshaker.go +++ b/handshaker.go @@ -279,6 +279,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 { @@ -289,7 +298,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(): diff --git a/pkg/protocol/application_data.go b/pkg/protocol/application_data.go index f5d8153df..40dfd2360 100644 --- a/pkg/protocol/application_data.go +++ b/pkg/protocol/application_data.go @@ -19,7 +19,7 @@ func (a ApplicationData) ContentType() ContentType { // Marshal encodes the ApplicationData to binary. func (a *ApplicationData) Marshal() ([]byte, error) { - return append([]byte{}, a.Data...), nil + return a.Data, nil } // Unmarshal populates the ApplicationData from binary. diff --git a/pkg/protocol/handshake/header.go b/pkg/protocol/handshake/header.go index befeed10d..d19d384cd 100644 --- a/pkg/protocol/handshake/header.go +++ b/pkg/protocol/handshake/header.go @@ -29,6 +29,16 @@ type Header struct { // Marshal encodes the Header. func (h *Header) Marshal() ([]byte, error) { out := make([]byte, HeaderLength) + err := h.MarshalInto(out) + + return out, err +} + +// MarshalInto encodes the Header using a pre-allocated buffer. +func (h *Header) MarshalInto(out []byte) error { + if len(out) < HeaderLength { + return errBufferTooSmall + } out[0] = byte(h.Type) util.PutBigEndianUint24(out[1:], h.Length) @@ -36,7 +46,7 @@ func (h *Header) Marshal() ([]byte, error) { util.PutBigEndianUint24(out[6:], h.FragmentOffset) util.PutBigEndianUint24(out[9:], h.FragmentLength) - return out, nil + return nil } // Unmarshal populates the header from encoded data. diff --git a/pkg/protocol/handshake/message_certificate.go b/pkg/protocol/handshake/message_certificate.go index 1ac82ac83..615c34f1a 100644 --- a/pkg/protocol/handshake/message_certificate.go +++ b/pkg/protocol/handshake/message_certificate.go @@ -13,6 +13,7 @@ import ( // https://tools.ietf.org/html/rfc5246#section-7.4.2 type MessageCertificate struct { Certificate [][]byte + cache []byte } // Type returns the Handshake Type. @@ -26,6 +27,9 @@ const ( // Marshal encodes the Handshake. func (m *MessageCertificate) Marshal() ([]byte, error) { + if m.cache != nil { + return m.cache, nil + } total := handshakeMessageCertificateLengthFieldSize for _, cert := range m.Certificate { @@ -50,6 +54,8 @@ func (m *MessageCertificate) Marshal() ([]byte, error) { offset += len(cert) } + m.cache = out + return out, nil } diff --git a/pkg/protocol/recordlayer/header.go b/pkg/protocol/recordlayer/header.go index d1255f89f..e87d37748 100644 --- a/pkg/protocol/recordlayer/header.go +++ b/pkg/protocol/recordlayer/header.go @@ -37,8 +37,24 @@ func (h *Header) Marshal() ([]byte, error) { } hs := FixedHeaderSize + len(h.ConnectionID) - out := make([]byte, hs) + err := h.MarshalInto(out) + + return out, err +} + +// MarshalInto encodes a TLS RecordLayer Header to binary using pre-allocated buffer. +func (h *Header) MarshalInto(out []byte) error { + if h.SequenceNumber > MaxSequenceNumber { + return errSequenceNumberOverflow + } + + hs := FixedHeaderSize + len(h.ConnectionID) + + if len(out) < hs { + return errBufferTooSmall + } + out[0] = byte(h.ContentType) out[1] = h.Version.Major out[2] = h.Version.Minor @@ -47,7 +63,7 @@ func (h *Header) Marshal() ([]byte, error) { copy(out[11:11+len(h.ConnectionID)], h.ConnectionID) binary.BigEndian.PutUint16(out[hs-2:], h.ContentLen) - return out, nil + return nil } // Unmarshal populates a TLS RecordLayer Header from binary. diff --git a/pkg/protocol/recordlayer/recordlayer.go b/pkg/protocol/recordlayer/recordlayer.go index a9a456fc9..9a0042cd3 100644 --- a/pkg/protocol/recordlayer/recordlayer.go +++ b/pkg/protocol/recordlayer/recordlayer.go @@ -58,12 +58,16 @@ func (r *RecordLayer) Marshal() ([]byte, error) { r.Header.ContentLen = uint16(len(contentRaw)) //nolint:gosec // G115 r.Header.ContentType = r.Content.ContentType() - headerRaw, err := r.Header.Marshal() + out := make([]byte, len(contentRaw)+r.Header.Size()) + + err = r.Header.MarshalInto(out) if err != nil { return nil, err } - return append(headerRaw, contentRaw...), nil + copy(out[r.Header.Size():], contentRaw) + + return out, nil } // Unmarshal populates the RecordLayer from binary. From b08dbfffb96d01fc9ca1ceceeed714df682a51b6 Mon Sep 17 00:00:00 2001 From: Thomas Legris Date: Mon, 30 Mar 2026 11:52:20 +0900 Subject: [PATCH 2/5] Add MarhsalInto & Size across all content --- conn.go | 12 +-- handshaker.go | 3 +- pkg/protocol/alert/alert.go | 21 ++++- pkg/protocol/application_data.go | 17 +++- pkg/protocol/change_cipher_spec.go | 20 +++- pkg/protocol/content.go | 2 + pkg/protocol/extension/extension.go | 14 +++ pkg/protocol/handshake/handshake.go | 38 ++++++-- pkg/protocol/handshake/message_certificate.go | 27 +++--- .../handshake/message_certificate_13.go | 77 +++++++++++----- .../handshake/message_certificate_request.go | 68 ++++++++++---- .../message_certificate_request_13.go | 92 +++++++++++++++---- .../handshake/message_certificate_verify.go | 26 +++++- .../handshake/message_client_hello.go | 74 +++++++++++---- .../handshake/message_client_key_exchange.go | 55 +++++++++-- pkg/protocol/handshake/message_finished.go | 20 +++- .../handshake/message_hello_verify_request.go | 23 ++++- .../handshake/message_server_hello.go | 67 ++++++++++---- .../handshake/message_server_hello_done.go | 15 ++- .../handshake/message_server_key_exchange.go | 90 ++++++++++++++---- pkg/protocol/recordlayer/recordlayer.go | 16 ++-- 21 files changed, 615 insertions(+), 162 deletions(-) diff --git a/conn.go b/conn.go index 3eb3cf687..4987ae783 100644 --- a/conn.go +++ b/conn.go @@ -588,7 +588,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) @@ -712,8 +712,8 @@ 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(rawPacket) >= c.maximumTransmissionUnit { @@ -800,8 +800,6 @@ func (c *Conn) processPacket(pkt *packet) ([]byte, error) { //nolint:cyclop //nolint:cyclop func (c *Conn) processHandshakePacket(pkt *packet, dtlsHandshake *handshake.Handshake) ([][]byte, error) { - var rawPackets [][]byte - handshakeFragments, err := c.fragmentHandshake(dtlsHandshake) if err != nil { return nil, err @@ -811,6 +809,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 { @@ -885,11 +884,10 @@ func (c *Conn) fragmentHandshake(dtlsHandshake *handshake.Handshake) ([][]byte, return nil, err } - var fragmentedHandshakes [][]byte - contentFragments := splitBytes(content, c.maximumTransmissionUnit) offset := 0 + fragmentedHandshakes := make([][]byte, 0, len(contentFragments)) for _, contentFragment := range contentFragments { contentFragmentLen := len(contentFragment) diff --git a/handshaker.go b/handshaker.go index f9ef37ddf..2ffcb12c3 100644 --- a/handshaker.go +++ b/handshaker.go @@ -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) } diff --git a/pkg/protocol/alert/alert.go b/pkg/protocol/alert/alert.go index c317567ab..d85d72f27 100644 --- a/pkg/protocol/alert/alert.go +++ b/pkg/protocol/alert/alert.go @@ -145,9 +145,28 @@ func (a Alert) ContentType() protocol.ContentType { return protocol.ContentTypeAlert } +// Size returns the minimal buffer size required for MarshalInto. +func (a Alert) Size() int { + return 2 +} + // Marshal returns the encoded alert. func (a *Alert) Marshal() ([]byte, error) { - return []byte{byte(a.Level), byte(a.Description)}, nil + out := make([]byte, a.Size()) + err := a.MarshalInto(out) + + return out, err +} + +// MarshalInto returns the encoded alert. +func (a *Alert) MarshalInto(out []byte) error { + if len(out) < a.Size() { + return errBufferTooSmall + } + out[0] = byte(a.Level) + out[1] = byte(a.Description) + + return nil } // Unmarshal populates the alert from binary data. diff --git a/pkg/protocol/application_data.go b/pkg/protocol/application_data.go index 40dfd2360..1e290ee0e 100644 --- a/pkg/protocol/application_data.go +++ b/pkg/protocol/application_data.go @@ -19,7 +19,22 @@ func (a ApplicationData) ContentType() ContentType { // Marshal encodes the ApplicationData to binary. func (a *ApplicationData) Marshal() ([]byte, error) { - return a.Data, nil + out := make([]byte, len(a.Data)) + err := a.MarshalInto(out) + + return out, err +} + +// MarshalInto encodes the ApplicationData to binary into a pre-allocated buffer. +func (a *ApplicationData) MarshalInto(out []byte) error { + copy(out, a.Data) + + return nil +} + +// Size returns the size required for MarshalInto. +func (a ApplicationData) Size() int { + return len(a.Data) } // Unmarshal populates the ApplicationData from binary. diff --git a/pkg/protocol/change_cipher_spec.go b/pkg/protocol/change_cipher_spec.go index e8b18de4e..c1bc357b7 100644 --- a/pkg/protocol/change_cipher_spec.go +++ b/pkg/protocol/change_cipher_spec.go @@ -15,9 +15,27 @@ func (c ChangeCipherSpec) ContentType() ContentType { return ContentTypeChangeCipherSpec } +// Size returns the minimal buffer size required for MarshalInto. +func (c ChangeCipherSpec) Size() int { + return 1 +} + // Marshal encodes the ChangeCipherSpec to binary. func (c *ChangeCipherSpec) Marshal() ([]byte, error) { - return []byte{0x01}, nil + out := make([]byte, 1) + err := c.MarshalInto(out) + + return out, err +} + +// MarshalInto encodes the ChangeCipherSpec to binary into a pre-allocated buffer. +func (c *ChangeCipherSpec) MarshalInto(out []byte) error { + if len(out) < c.Size() { + return errBufferTooSmall + } + out[0] = 0x01 + + return nil } // Unmarshal populates the ChangeCipherSpec from binary. diff --git a/pkg/protocol/content.go b/pkg/protocol/content.go index 58bbdc5b9..56ccaccc5 100644 --- a/pkg/protocol/content.go +++ b/pkg/protocol/content.go @@ -21,5 +21,7 @@ const ( type Content interface { ContentType() ContentType Marshal() ([]byte, error) + MarshalInto([]byte) error Unmarshal(data []byte) error + Size() int } diff --git a/pkg/protocol/extension/extension.go b/pkg/protocol/extension/extension.go index ba7c7d836..ca9db5ae1 100644 --- a/pkg/protocol/extension/extension.go +++ b/pkg/protocol/extension/extension.go @@ -134,3 +134,17 @@ func Marshal(e []Extension) ([]byte, error) { return append(out, extensions...), nil } + +// Size returns the length of extensions marshal. +func Size(e []Extension) int { + total := 2 + for _, e := range e { + raw, err := e.Marshal() + if err != nil { + return 0 + } + total += len(raw) + } + + return total +} diff --git a/pkg/protocol/handshake/handshake.go b/pkg/protocol/handshake/handshake.go index 6d40a5a7e..3ef0c8a5f 100644 --- a/pkg/protocol/handshake/handshake.go +++ b/pkg/protocol/handshake/handshake.go @@ -62,6 +62,8 @@ func (t Type) String() string { //nolint:cyclop // Message is the body of a Handshake datagram. type Message interface { Marshal() ([]byte, error) + MarshalInto([]byte) error + Size() int Unmarshal(data []byte) error Type() Type } @@ -84,6 +86,11 @@ func (h Handshake) ContentType() protocol.ContentType { return protocol.ContentTypeHandshake } +// Size returns the minimal buffer size required for MarshalInto. +func (h *Handshake) Size() int { + return HeaderLength + h.Message.Size() +} + // Marshal encodes a handshake into a binary message. func (h *Handshake) Marshal() ([]byte, error) { if h.Message == nil { @@ -91,21 +98,38 @@ func (h *Handshake) Marshal() ([]byte, error) { } else if h.Header.FragmentOffset != 0 { return nil, errUnableToMarshalFragmented } + out := make([]byte, h.Size()) + err := h.MarshalInto(out) - msg, err := h.Message.Marshal() - if err != nil { - return nil, err + return out, err +} + +// MarshalInto encodes a handshake into a binary message into a pre-allocated buffer. +func (h *Handshake) MarshalInto(out []byte) error { + if h.Message == nil { + return errHandshakeMessageUnset + } else if h.Header.FragmentOffset != 0 { + return errUnableToMarshalFragmented } - h.Header.Length = uint32(len(msg)) //nolint:gosec // G115 + if len(out) < h.Size() { + return errBufferTooSmall + } + + h.Header.Length = uint32(h.Message.Size()) //nolint:gosec // G115 h.Header.FragmentLength = h.Header.Length h.Header.Type = h.Message.Type() - header, err := h.Header.Marshal() + err := h.Header.MarshalInto(out) if err != nil { - return nil, err + return err + } + + err = h.Message.MarshalInto(out[HeaderLength:]) + if err != nil { + return err } - return append(header, msg...), nil + return nil } // Unmarshal decodes a handshake from a binary message. diff --git a/pkg/protocol/handshake/message_certificate.go b/pkg/protocol/handshake/message_certificate.go index 615c34f1a..2ccef174b 100644 --- a/pkg/protocol/handshake/message_certificate.go +++ b/pkg/protocol/handshake/message_certificate.go @@ -13,7 +13,6 @@ import ( // https://tools.ietf.org/html/rfc5246#section-7.4.2 type MessageCertificate struct { Certificate [][]byte - cache []byte } // Type returns the Handshake Type. @@ -25,22 +24,30 @@ const ( handshakeMessageCertificateLengthFieldSize = 3 ) -// Marshal encodes the Handshake. -func (m *MessageCertificate) Marshal() ([]byte, error) { - if m.cache != nil { - return m.cache, nil - } +// Size returns the minimal size required for MarshalInto. +func (m *MessageCertificate) Size() int { total := handshakeMessageCertificateLengthFieldSize for _, cert := range m.Certificate { total += handshakeMessageCertificateLengthFieldSize + len(cert) } - out := make([]byte, total) + return total +} +// Marshal encodes the Handshake. +func (m *MessageCertificate) Marshal() ([]byte, error) { + out := make([]byte, m.Size()) + err := m.MarshalInto(out) + + return out, err +} + +// MarshalInto encodes the Handshake into a pre-allocated buffer. +func (m *MessageCertificate) MarshalInto(out []byte) error { // Total Payload Size //nolint:gosec // G115 - util.PutBigEndianUint24(out, uint32(total-handshakeMessageCertificateLengthFieldSize)) + util.PutBigEndianUint24(out, uint32(m.Size()-handshakeMessageCertificateLengthFieldSize)) offset := handshakeMessageCertificateLengthFieldSize for _, cert := range m.Certificate { @@ -54,9 +61,7 @@ func (m *MessageCertificate) Marshal() ([]byte, error) { offset += len(cert) } - m.cache = out - - return out, nil + return nil } // Unmarshal populates the message from encoded data. diff --git a/pkg/protocol/handshake/message_certificate_13.go b/pkg/protocol/handshake/message_certificate_13.go index 9e071a413..4d6da4bd9 100644 --- a/pkg/protocol/handshake/message_certificate_13.go +++ b/pkg/protocol/handshake/message_certificate_13.go @@ -71,44 +71,77 @@ func (m *MessageCertificate13) Marshal() ([]byte, error) { return nil, errCertificateRequestContextTooLong } + out := make([]byte, m.Size()) + err := m.MarshalInto(out) + + return out, err +} + +func (m *MessageCertificate13) Size() int { + return 1 + len(m.CertificateRequestContext) + cert13CertLengthFieldSize + m.certsSize() +} + +func (m *MessageCertificate13) certsSize() int { + certificateListSize := 0 + for _, entry := range m.CertificateList { + certificateListSize += cert13CertLengthFieldSize + certificateListSize += len(entry.CertificateData) + certificateListSize += extension.Size(entry.Extensions) + } + + return certificateListSize +} + +// MarshalInto is same as Marshal but uses a pre-allocated buffer. +func (m *MessageCertificate13) MarshalInto(out []byte) error { + // Validate certificate_request_context length + if len(m.CertificateRequestContext) > cert13ContextMaxLength { + return errCertificateRequestContextTooLong + } + + if len(out) < m.Size() { + return errBufferTooSmall + } + + // Check size of certificate_list is still within bounds + if m.certsSize() > maxUint24 { + return errCertificateListTooLong + } + // Start with certificate_request_context (1-byte length prefix) //nolint:gosec // G115: certificate_request_context length is validated to be <= 255 above. - out := []byte{byte(len(m.CertificateRequestContext))} - out = append(out, m.CertificateRequestContext...) + offset := 0 + out[0] = byte(len(m.CertificateRequestContext)) //nolint:gosec // G115 + offset += 1 + n := copy(out[offset:], m.CertificateRequestContext) //nolint:gosec // G115 + offset += n + + // Add certificate_list with 3-byte length prefix + util.PutBigEndianUint24(out[offset:], uint32(m.certsSize())) //nolint:gosec // G115 + offset += 3 // Build certificate_list - certificateList := []byte{} for _, entry := range m.CertificateList { // Add cert_data as a 3-byte length prefix certDataLen := len(entry.CertificateData) if certDataLen == 0 || certDataLen > maxUint24 { - return nil, errInvalidCertificateEntry + return errInvalidCertificateEntry } - certDataLenBytes := make([]byte, cert13CertLengthFieldSize) - util.PutBigEndianUint24(certDataLenBytes, uint32(certDataLen)) //nolint:gosec // G115 - certificateList = append(certificateList, certDataLenBytes...) - certificateList = append(certificateList, entry.CertificateData...) + util.PutBigEndianUint24(out[offset:], uint32(certDataLen)) //nolint:gosec // G115 + offset += 3 + n = copy(out[offset:], entry.CertificateData) + offset += n // Marshal extensions (includes a 2-byte length prefix) extensionsData, err := extension.Marshal(entry.Extensions) if err != nil { - return nil, err - } - certificateList = append(certificateList, extensionsData...) - - // Check size of certificate_list is still within bounds - if len(certificateList) > maxUint24 { - return nil, errCertificateListTooLong + return err } + n = copy(out[offset:], extensionsData) + offset += n } - // Add certificate_list with 3-byte length prefix - certificateListLenBytes := make([]byte, cert13CertLengthFieldSize) - util.PutBigEndianUint24(certificateListLenBytes, uint32(len(certificateList))) //nolint:gosec // G115 - out = append(out, certificateListLenBytes...) - out = append(out, certificateList...) - - return out, nil + return nil } // parseCertificate13Entry parses a single certificate entry from the cryptobyte string. diff --git a/pkg/protocol/handshake/message_certificate_request.go b/pkg/protocol/handshake/message_certificate_request.go index a46e39645..8520e31e6 100644 --- a/pkg/protocol/handshake/message_certificate_request.go +++ b/pkg/protocol/handshake/message_certificate_request.go @@ -35,40 +35,74 @@ func (m MessageCertificateRequest) Type() Type { return TypeCertificateRequest } +// Size returns the minimal size required for MarshalInto. +func (m *MessageCertificateRequest) Size() int { + return 1 + + len(m.CertificateTypes) + + 2 + // number of SignatureHashAlgorithms + 2*len(m.SignatureHashAlgorithms) + // SignatureHashAlgorithms size + 2 + // casLength + m.casLength() +} + +func (m *MessageCertificateRequest) casLength() int { + casLength := 0 + for _, ca := range m.CertificateAuthoritiesNames { + casLength += 2 + len(ca) + } + + return casLength +} + // Marshal encodes the Handshake. func (m *MessageCertificateRequest) Marshal() ([]byte, error) { + out := make([]byte, m.Size()) + err := m.MarshalInto(out) + + return out, err +} + +// MarshalInto encodes the Handshake into a pre-allocated buffer. +func (m *MessageCertificateRequest) MarshalInto(out []byte) error { if len(m.CertificateTypes) > 255 { - return nil, errCertificateTypesTooLong + return errCertificateTypesTooLong + } + + if len(out) < m.Size() { + return errBufferTooSmall } //nolint:gosec // G115: certificate types count is validated to be <= 255 above. - out := []byte{byte(len(m.CertificateTypes))} + offset := 0 + out[offset] = byte(len(m.CertificateTypes)) //nolint:gosec // G115 + offset += 1 for _, v := range m.CertificateTypes { - out = append(out, byte(v)) + out[offset] = byte(v) + offset += 1 } - out = append(out, []byte{0x00, 0x00}...) - binary.BigEndian.PutUint16(out[len(out)-2:], uint16(len(m.SignatureHashAlgorithms)*2)) //nolint:gosec //G115 + binary.BigEndian.PutUint16(out[offset:], uint16(len(m.SignatureHashAlgorithms)*2)) //nolint:gosec //G115 + offset += 2 + for _, v := range m.SignatureHashAlgorithms { - out = append(out, v.Marshal()...) + tmp := v.Marshal() + n := copy(out[offset:], tmp) + offset += n } // Distinguished Names - casLength := 0 - for _, ca := range m.CertificateAuthoritiesNames { - casLength += len(ca) + 2 - } - out = append(out, []byte{0x00, 0x00}...) - binary.BigEndian.PutUint16(out[len(out)-2:], uint16(casLength)) //nolint:gosec //G115 - if casLength > 0 { + binary.BigEndian.PutUint16(out[offset:], uint16(m.casLength())) //nolint:gosec //G115 + offset += 2 + if m.casLength() > 0 { for _, ca := range m.CertificateAuthoritiesNames { - out = append(out, []byte{0x00, 0x00}...) - binary.BigEndian.PutUint16(out[len(out)-2:], uint16(len(ca))) //nolint:gosec //G115 - out = append(out, ca...) + binary.BigEndian.PutUint16(out[offset:], uint16(len(ca))) //nolint:gosec //G115 + offset += 2 + n := copy(out[offset:], ca) + offset += n } } - return out, nil + return nil } // Unmarshal populates the message from encoded data. diff --git a/pkg/protocol/handshake/message_certificate_request_13.go b/pkg/protocol/handshake/message_certificate_request_13.go index 3563b0ea5..1a015cc73 100644 --- a/pkg/protocol/handshake/message_certificate_request_13.go +++ b/pkg/protocol/handshake/message_certificate_request_13.go @@ -22,6 +22,10 @@ type MessageCertificateRequest13 struct { // Extensions contains the list of extensions. // The signature_algorithms extension is REQUIRED per RFC 8446. Extensions []extension.Extension + + // Cache the marshal result + marshalCache []byte + marshalCacheErr error } // Type returns the handshake message type. @@ -44,11 +48,9 @@ const ( // [2 bytes] extensions length (from extension.Marshal) // [variable] extensions data func (m *MessageCertificateRequest13) Marshal() ([]byte, error) { - // Validate certificate_request_context length if len(m.CertificateRequestContext) > certReq13ContextMaxLength { return nil, errCertificateRequestContextTooLong } - // Validate that signature_algorithms extension is present (required by RFC 8446) hasSignatureAlgorithms := false for _, ext := range m.Extensions { @@ -61,26 +63,84 @@ func (m *MessageCertificateRequest13) Marshal() ([]byte, error) { if !hasSignatureAlgorithms { return nil, errMissingSignatureAlgorithmsExtension } + out := make([]byte, m.Size()) + err := m.MarshalInto(out) - var builder cryptobyte.Builder - - // Add certificate_request_context (1-byte length prefix) - builder.AddUint8LengthPrefixed(func(b *cryptobyte.Builder) { - b.AddBytes(m.CertificateRequestContext) - }) + return out, err +} - // Marshal extensions (includes 2-byte length prefix, like in TLS 1.2) - extensionsData, err := extension.Marshal(m.Extensions) +// Size returns the size needed for MarshalInto. +func (m *MessageCertificateRequest13) Size() int { + cache, err := m.innerMarshal() if err != nil { - return nil, err + return 0 + } + + return len(cache) +} + +func (m *MessageCertificateRequest13) innerMarshal() ([]byte, error) { + if m.marshalCache == nil && m.marshalCacheErr == nil { + var builder cryptobyte.Builder + + // Add certificate_request_context (1-byte length prefix) + builder.AddUint8LengthPrefixed(func(b *cryptobyte.Builder) { + b.AddBytes(m.CertificateRequestContext) + }) + + // Marshal extensions (includes 2-byte length prefix, like in TLS 1.2) + extensionsData, err := extension.Marshal(m.Extensions) + if err != nil { + m.marshalCacheErr = err + + return nil, err + } + // Validate extensions length is in valid range <2..2^16-1> + if len(extensionsData) < 2 || len(extensionsData) > maxUint16 { + m.marshalCacheErr = err + + return nil, errInvalidExtensionsLength + } + builder.AddBytes(extensionsData) + + m.marshalCache, m.marshalCacheErr = builder.Bytes() } - // Validate extensions length is in valid range <2..2^16-1> - if len(extensionsData) < 2 || len(extensionsData) > maxUint16 { - return nil, errInvalidExtensionsLength + + return m.marshalCache, m.marshalCacheErr +} + +// MarshalInto encodes like Marshal but in a pre-allocate buffer. +func (m *MessageCertificateRequest13) MarshalInto(out []byte) error { + // Validate certificate_request_context length + if len(m.CertificateRequestContext) > certReq13ContextMaxLength { + return errCertificateRequestContextTooLong } - builder.AddBytes(extensionsData) - return builder.Bytes() + // Validate that signature_algorithms extension is present (required by RFC 8446) + hasSignatureAlgorithms := false + for _, ext := range m.Extensions { + if ext.TypeValue() == extension.SupportedSignatureAlgorithmsTypeValue { + hasSignatureAlgorithms = true + + break + } + } + if !hasSignatureAlgorithms { + return errMissingSignatureAlgorithmsExtension + } + + if len(out) < m.Size() { + return errBufferTooSmall + } + + cache, err := m.innerMarshal() + if err != nil { + return err + } + + copy(out, cache) + + return nil } // Unmarshal decodes the MessageCertificateRequest13 from its wire format. diff --git a/pkg/protocol/handshake/message_certificate_verify.go b/pkg/protocol/handshake/message_certificate_verify.go index c6626bb6c..5207c9b06 100644 --- a/pkg/protocol/handshake/message_certificate_verify.go +++ b/pkg/protocol/handshake/message_certificate_verify.go @@ -29,6 +29,11 @@ func (m MessageCertificateVerify) Type() Type { return TypeCertificateVerify } +// Size returns the minimal required size for Marshalnto. +func (m MessageCertificateVerify) Size() int { + return 1 + 1 + 2 + len(m.Signature) +} + // Marshal encodes the Handshake. func (m *MessageCertificateVerify) Marshal() ([]byte, error) { if m.HashAlgorithm > 0xFF || m.SignatureAlgorithm > 0xFF { @@ -42,14 +47,31 @@ func (m *MessageCertificateVerify) Marshal() ([]byte, error) { return nil, errInvalidSignHashAlgorithm } - out := make([]byte, 1+1+2+len(m.Signature)) + out := make([]byte, m.Size()) + err := m.MarshalInto(out) + + return out, err +} + +// MarshalInto encodes the Handshake into a pre-allocated buffer. +func (m *MessageCertificateVerify) MarshalInto(out []byte) error { + if m.HashAlgorithm > 0xFF || m.SignatureAlgorithm > 0xFF { + return errInvalidSignHashAlgorithm + } + + // CertificateVerify in DTLS 1.2 encodes hash/signature as 1 byte each. + scheme := tls.SignatureScheme(uint16(m.HashAlgorithm)<<8 | uint16(m.SignatureAlgorithm)) + var alg signaturehash.Algorithm + if err := alg.Unmarshal(scheme); err != nil { + return errInvalidSignHashAlgorithm + } out[0] = byte(m.HashAlgorithm) out[1] = byte(m.SignatureAlgorithm) binary.BigEndian.PutUint16(out[2:], uint16(len(m.Signature))) //nolint:gosec // G115 copy(out[4:], m.Signature) - return out, nil + return nil } // Unmarshal populates the message from encoded data. diff --git a/pkg/protocol/handshake/message_client_hello.go b/pkg/protocol/handshake/message_client_hello.go index 5f777d65d..3d73892dd 100644 --- a/pkg/protocol/handshake/message_client_hello.go +++ b/pkg/protocol/handshake/message_client_hello.go @@ -36,46 +36,82 @@ func (m MessageClientHello) Type() Type { return TypeClientHello } +// Size returns the size needed for MarshalInto. +func (m *MessageClientHello) Size() int { + encodedCipherSuiteIDs := encodeCipherSuiteIDs(m.CipherSuiteIDs) + encodedCompressionMethods := protocol.EncodeCompressionMethods(m.CompressionMethods) + + return handshakeMessageClientHelloVariableWidthStart + + 1 + + len(m.SessionID) + + 1 + + len(m.Cookie) + + len(encodedCipherSuiteIDs) + + len(encodedCompressionMethods) + + extension.Size(m.Extensions) +} + // Marshal encodes the Handshake. func (m *MessageClientHello) Marshal() ([]byte, error) { + out := make([]byte, m.Size()) + err := m.MarshalInto(out) + + return out, err +} + +// MarshalInto encodes the Handshake into a pre-allocated buffer. +func (m *MessageClientHello) MarshalInto(out []byte) error { if len(m.Cookie) > 255 { - return nil, errCookieTooLong + return errCookieTooLong } if len(m.SessionID) > 255 { - return nil, errSessionIDTooLong + return errSessionIDTooLong } if len(m.CompressionMethods) > 255 { - return nil, errCompressionMethodsTooLong + return errCompressionMethodsTooLong } extensions, err := extension.Marshal(m.Extensions) if err != nil { - return nil, err + return err + } + + if len(out) < m.Size() { + return errBufferTooSmall } encodedCipherSuiteIDs := encodeCipherSuiteIDs(m.CipherSuiteIDs) encodedCompressionMethods := protocol.EncodeCompressionMethods(m.CompressionMethods) - out := make( - []byte, - 0, - handshakeMessageClientHelloVariableWidthStart+1+len(m.SessionID)+1+ - len(m.Cookie)+len(encodedCipherSuiteIDs)+len(encodedCompressionMethods)+len(extensions), - ) - out = append(out, m.Version.Major, m.Version.Minor) + offset := 0 + out[0] = m.Version.Major + out[1] = m.Version.Minor + offset += 2 rand := m.Random.MarshalFixed() - out = append(out, rand[:]...) + n := copy(out[offset:], rand[:]) + offset += n + out[offset] = byte(len(m.SessionID)) //nolint:gosec // G115: session ID length is validated to be <= 255 above. + offset += 1 + + n = copy(out[offset:], m.SessionID) + offset += n - out = append(out, byte(len(m.SessionID))) //nolint:gosec // G115: session ID length is validated to be <= 255 above. - out = append(out, m.SessionID...) + out[offset] = byte(len(m.Cookie)) //nolint:gosec // G115: cookie length is validated to be <= 255 above. + offset += 1 - out = append(out, byte(len(m.Cookie))) //nolint:gosec // G115: cookie length is validated to be <= 255 above. - out = append(out, m.Cookie...) - out = append(out, encodedCipherSuiteIDs...) - out = append(out, encodedCompressionMethods...) + n = copy(out[offset:], m.Cookie) + offset += n - return append(out, extensions...), nil + n = copy(out[offset:], encodedCipherSuiteIDs) + offset += n + + n = copy(out[offset:], encodedCompressionMethods) + offset += n + + copy(out[offset:], extensions) + + return nil } // Unmarshal populates the message from encoded data. diff --git a/pkg/protocol/handshake/message_client_key_exchange.go b/pkg/protocol/handshake/message_client_key_exchange.go index aa679b49d..627830d2c 100644 --- a/pkg/protocol/handshake/message_client_key_exchange.go +++ b/pkg/protocol/handshake/message_client_key_exchange.go @@ -30,25 +30,66 @@ func (m MessageClientKeyExchange) Type() Type { } // Marshal encodes the Handshake. -func (m *MessageClientKeyExchange) Marshal() (out []byte, err error) { +func (m *MessageClientKeyExchange) Marshal() ([]byte, error) { if m.IdentityHint == nil && m.PublicKey == nil { return nil, errInvalidClientKeyExchange } + if m.PublicKey != nil { + if len(m.PublicKey) > 255 { + return nil, errPublicKeyTooLong + } + } + + out := make([]byte, m.Size()) + err := m.MarshalInto(out) + + return out, err +} + +// Size returns the size required for MarshalInto. +func (m *MessageClientKeyExchange) Size() int { + total := 0 + if m.IdentityHint != nil { + total += 2 + } + + if m.PublicKey != nil { + total += 1 + total += len(m.PublicKey) + } + + return total +} + +// MarshalInto encodes the Handshake into a pre-allocated buffer. +func (m *MessageClientKeyExchange) MarshalInto(out []byte) error { + if m.IdentityHint == nil && m.PublicKey == nil { + return errInvalidClientKeyExchange + } + + if len(out) < m.Size() { + return errBufferTooSmall + } + + offset := 0 if m.IdentityHint != nil { - out = append([]byte{0x00, 0x00}, m.IdentityHint...) - binary.BigEndian.PutUint16(out, uint16(len(out)-2)) //nolint:gosec // G115 + binary.BigEndian.PutUint16(out[offset:], uint16(len(m.IdentityHint))) //nolint:gosec // G115 + offset += 2 + n := copy(out[offset:], m.IdentityHint) + offset += n } if m.PublicKey != nil { if len(m.PublicKey) > 255 { - return nil, errPublicKeyTooLong + return errPublicKeyTooLong } - out = append(out, byte(len(m.PublicKey))) //nolint:gosec // G115: public key length is validated to be <= 255 above. - out = append(out, m.PublicKey...) + out[offset] = byte(len(m.PublicKey)) //nolint:gosec // G115: public key length is validated to be <= 255 above. + offset += 1 + copy(out[offset:], m.PublicKey) } - return out, nil + return nil } // Unmarshal populates the message from encoded data. diff --git a/pkg/protocol/handshake/message_finished.go b/pkg/protocol/handshake/message_finished.go index 6362eaccc..4662a9250 100644 --- a/pkg/protocol/handshake/message_finished.go +++ b/pkg/protocol/handshake/message_finished.go @@ -18,9 +18,27 @@ func (m MessageFinished) Type() Type { return TypeFinished } +// Size returns the size required for MarshalInto. +func (m *MessageFinished) Size() int { + return len(m.VerifyData) +} + // Marshal encodes the Handshake. func (m *MessageFinished) Marshal() ([]byte, error) { - return append([]byte{}, m.VerifyData...), nil + out := make([]byte, m.Size()) + err := m.MarshalInto(out) + + return out, err +} + +// MarshalInto encodes the Handshake into a pre-allocated buffer. +func (m *MessageFinished) MarshalInto(out []byte) error { + if len(out) < m.Size() { + return errBufferTooSmall + } + copy(out, m.VerifyData) + + return nil } // Unmarshal populates the message from encoded data. diff --git a/pkg/protocol/handshake/message_hello_verify_request.go b/pkg/protocol/handshake/message_hello_verify_request.go index e1c7a4e1f..41123fc21 100644 --- a/pkg/protocol/handshake/message_hello_verify_request.go +++ b/pkg/protocol/handshake/message_hello_verify_request.go @@ -37,14 +37,33 @@ func (m *MessageHelloVerifyRequest) Marshal() ([]byte, error) { if len(m.Cookie) > 255 { return nil, errCookieTooLong } + out := make([]byte, m.Size()) + err := m.MarshalInto(out) + + return out, err +} + +// Size returns the size required for MarshalInto. +func (m *MessageHelloVerifyRequest) Size() int { + return 3 + len(m.Cookie) +} + +// MarshalInto encodes the Handshake into a pre-allocated buffer. +func (m *MessageHelloVerifyRequest) MarshalInto(out []byte) error { + if len(m.Cookie) > 255 { + return errCookieTooLong + } + + if len(out) < m.Size() { + return errBufferTooSmall + } - out := make([]byte, 3+len(m.Cookie)) out[0] = m.Version.Major out[1] = m.Version.Minor out[2] = byte(len(m.Cookie)) //nolint:gosec // G115: cookie length is validated to be <= 255 above. copy(out[3:], m.Cookie) - return out, nil + return nil } // Unmarshal populates the message from encoded data. diff --git a/pkg/protocol/handshake/message_server_hello.go b/pkg/protocol/handshake/message_server_hello.go index c62304de4..4d9a05aea 100644 --- a/pkg/protocol/handshake/message_server_hello.go +++ b/pkg/protocol/handshake/message_server_hello.go @@ -34,6 +34,51 @@ func (m MessageServerHello) Type() Type { return TypeServerHello } +// Size returns the size required by MarshalInto. +func (m *MessageServerHello) Size() int { + total := 0 + total += extension.Size(m.Extensions) + total += messageServerHelloVariableWidthStart + 1 + len(m.SessionID) + 2 + 1 + + return total +} + +// MarshalInto encodes the Handshake into a pre-allocated buffer. +func (m *MessageServerHello) MarshalInto(out []byte) error { + extensions, err := extension.Marshal(m.Extensions) + if err != nil { + return err + } + + if len(out) < m.Size() { + return errBufferTooSmall + } + + offset := 0 + out[0] = m.Version.Major + out[1] = m.Version.Minor + offset += 2 + + rand := m.Random.MarshalFixed() + n := copy(out[offset:], rand[:]) + offset += n + + out[offset] = byte(len(m.SessionID)) //nolint:gosec // G115 + offset += 1 + n = copy(out[offset:], m.SessionID) + offset += n + + binary.BigEndian.PutUint16(out[offset:], *m.CipherSuiteID) + offset += 2 + + out[offset] = byte(m.CompressionMethod.ID) + offset += 1 + + copy(out[offset:], extensions) + + return nil +} + // Marshal encodes the Handshake. func (m *MessageServerHello) Marshal() ([]byte, error) { switch { @@ -45,26 +90,10 @@ func (m *MessageServerHello) Marshal() ([]byte, error) { return nil, errSessionIDTooLong } - extensions, err := extension.Marshal(m.Extensions) - if err != nil { - return nil, err - } - - out := make([]byte, 0, messageServerHelloVariableWidthStart+1+len(m.SessionID)+2+1+len(extensions)) - out = append(out, m.Version.Major, m.Version.Minor) - - rand := m.Random.MarshalFixed() - out = append(out, rand[:]...) - - out = append(out, byte(len(m.SessionID))) //nolint:gosec // G115: session ID length is validated to be <= 255 above. - out = append(out, m.SessionID...) - - out = append(out, 0x00, 0x00) - binary.BigEndian.PutUint16(out[len(out)-2:], *m.CipherSuiteID) - - out = append(out, byte(m.CompressionMethod.ID)) + out := make([]byte, m.Size()) + err := m.MarshalInto(out) - return append(out, extensions...), nil + return out, err } // Unmarshal populates the message from encoded data. diff --git a/pkg/protocol/handshake/message_server_hello_done.go b/pkg/protocol/handshake/message_server_hello_done.go index 87cc9ad75..be9319962 100644 --- a/pkg/protocol/handshake/message_server_hello_done.go +++ b/pkg/protocol/handshake/message_server_hello_done.go @@ -15,7 +15,20 @@ func (m MessageServerHelloDone) Type() Type { // Marshal encodes the Handshake. func (m *MessageServerHelloDone) Marshal() ([]byte, error) { - return []byte{}, nil + out := []byte{} + err := m.MarshalInto(out) + + return out, err +} + +// Size returns the size for MarshalInto. +func (m *MessageServerHelloDone) Size() int { + return 0 +} + +// MarshalInto encodes the Handshake. +func (m *MessageServerHelloDone) MarshalInto(out []byte) error { + return nil } // Unmarshal populates the message from encoded data. diff --git a/pkg/protocol/handshake/message_server_key_exchange.go b/pkg/protocol/handshake/message_server_key_exchange.go index 4c0ed64a7..411c833e3 100644 --- a/pkg/protocol/handshake/message_server_key_exchange.go +++ b/pkg/protocol/handshake/message_server_key_exchange.go @@ -34,40 +34,94 @@ func (m MessageServerKeyExchange) Type() Type { return TypeServerKeyExchange } -// Marshal encodes the Handshake. -func (m *MessageServerKeyExchange) Marshal() ([]byte, error) { //nolint:cyclop - var out []byte +// Size returns the size required for MarshalInto. +func (m *MessageServerKeyExchange) Size() int { //nolint:cyclop + total := 0 if m.IdentityHint != nil { - out = append([]byte{0x00, 0x00}, m.IdentityHint...) - binary.BigEndian.PutUint16(out, uint16(len(out)-2)) //nolint:gosec //G115 + total += 2 + len(m.IdentityHint) } if m.EllipticCurveType == 0 || len(m.PublicKey) == 0 { - return out, nil + return total + } + + total += 3 + total += 1 + total += len(m.PublicKey) + + switch { + case m.HashAlgorithm != hash.None && len(m.Signature) == 0: + return 0 + case m.HashAlgorithm == hash.None && len(m.Signature) > 0: + return 0 + case m.SignatureAlgorithm == signature.Anonymous && (m.HashAlgorithm != hash.None || len(m.Signature) > 0): + return 0 + case m.SignatureAlgorithm == signature.Anonymous: + return total } - out = append(out, byte(m.EllipticCurveType), 0x00, 0x00) - binary.BigEndian.PutUint16(out[len(out)-2:], uint16(m.NamedCurve)) + + total += 2 // signature hash length + total += 2 + total += len(m.Signature) + + return total +} + +// MarshalInto encodes the Handshake into a pre-allocated buffer. +func (m *MessageServerKeyExchange) MarshalInto(out []byte) error { //nolint:cyclop + if len(out) < m.Size() { + return errBufferTooSmall + } + + offset := 0 + if m.IdentityHint != nil { + binary.BigEndian.PutUint16(out[offset:], uint16(len(m.IdentityHint))) //nolint:gosec //G115 + offset += 2 + n := copy(out[offset:], m.IdentityHint) + offset += n + } + + if m.EllipticCurveType == 0 || len(m.PublicKey) == 0 { + return nil + } + out[offset] = byte(m.EllipticCurveType) + offset += 1 + binary.BigEndian.PutUint16(out[offset:], uint16(m.NamedCurve)) + offset += 2 //nolint:gosec // G115, no risk of overflow, the biggest supported curve is 97 bytes. - out = append(out, byte(len(m.PublicKey))) - out = append(out, m.PublicKey...) + out[offset] = byte(len(m.PublicKey)) + offset += 1 + n := copy(out[offset:], m.PublicKey) + offset += n switch { case m.HashAlgorithm != hash.None && len(m.Signature) == 0: - return nil, errInvalidSignHashAlgorithm + return errInvalidSignHashAlgorithm case m.HashAlgorithm == hash.None && len(m.Signature) > 0: - return nil, errInvalidSignHashAlgorithm + return errInvalidSignHashAlgorithm case m.SignatureAlgorithm == signature.Anonymous && (m.HashAlgorithm != hash.None || len(m.Signature) > 0): - return nil, errInvalidSignHashAlgorithm + return errInvalidSignHashAlgorithm case m.SignatureAlgorithm == signature.Anonymous: - return out, nil + return nil } alg := signaturehash.Algorithm{Hash: m.HashAlgorithm, Signature: m.SignatureAlgorithm} - out = append(out, append(alg.Marshal(), []byte{0x00, 0x00}...)...) - binary.BigEndian.PutUint16(out[len(out)-2:], uint16(len(m.Signature))) //nolint:gosec // G115 - out = append(out, m.Signature...) + tmp := alg.Marshal() + n = copy(out[offset:], tmp) + offset += n + binary.BigEndian.PutUint16(out[offset:], uint16(len(m.Signature))) //nolint:gosec // G115 + offset += 2 + copy(out[offset:], m.Signature) + + return nil +} + +// Marshal encodes the Handshake. +func (m *MessageServerKeyExchange) Marshal() ([]byte, error) { + out := make([]byte, m.Size()) + err := m.MarshalInto(out) - return out, nil + return out, err } // Unmarshal populates the message from encoded data. diff --git a/pkg/protocol/recordlayer/recordlayer.go b/pkg/protocol/recordlayer/recordlayer.go index 9a0042cd3..1bcdb96d4 100644 --- a/pkg/protocol/recordlayer/recordlayer.go +++ b/pkg/protocol/recordlayer/recordlayer.go @@ -50,22 +50,20 @@ type RecordLayer struct { // Marshal encodes the RecordLayer to binary. func (r *RecordLayer) Marshal() ([]byte, error) { - contentRaw, err := r.Content.Marshal() - if err != nil { - return nil, err - } + out := make([]byte, r.Content.Size()+r.Header.Size()) - r.Header.ContentLen = uint16(len(contentRaw)) //nolint:gosec // G115 + r.Header.ContentLen = uint16(r.Content.Size()) //nolint:gosec // G115 r.Header.ContentType = r.Content.ContentType() - out := make([]byte, len(contentRaw)+r.Header.Size()) - - err = r.Header.MarshalInto(out) + err := r.Header.MarshalInto(out) if err != nil { return nil, err } - copy(out[r.Header.Size():], contentRaw) + err = r.Content.MarshalInto(out[r.Header.Size():]) + if err != nil { + return nil, err + } return out, nil } From 2c089680c70ea253e2efc2b06b0e8003059573de Mon Sep 17 00:00:00 2001 From: Thomas Legris Date: Mon, 30 Mar 2026 19:45:19 +0900 Subject: [PATCH 3/5] Add extensions caching --- pkg/protocol/extension/extension.go | 14 -------- .../handshake/message_certificate_13.go | 33 +++++++++++++++++-- .../handshake/message_client_hello.go | 26 +++++++++++---- .../handshake/message_server_hello.go | 28 ++++++++++++---- 4 files changed, 72 insertions(+), 29 deletions(-) diff --git a/pkg/protocol/extension/extension.go b/pkg/protocol/extension/extension.go index ca9db5ae1..ba7c7d836 100644 --- a/pkg/protocol/extension/extension.go +++ b/pkg/protocol/extension/extension.go @@ -134,17 +134,3 @@ func Marshal(e []Extension) ([]byte, error) { return append(out, extensions...), nil } - -// Size returns the length of extensions marshal. -func Size(e []Extension) int { - total := 2 - for _, e := range e { - raw, err := e.Marshal() - if err != nil { - return 0 - } - total += len(raw) - } - - return total -} diff --git a/pkg/protocol/handshake/message_certificate_13.go b/pkg/protocol/handshake/message_certificate_13.go index 4d6da4bd9..e3a2be6b9 100644 --- a/pkg/protocol/handshake/message_certificate_13.go +++ b/pkg/protocol/handshake/message_certificate_13.go @@ -37,7 +37,9 @@ type MessageCertificate13 struct { // CertificateList contains the certificate chain with each entry having // optional per-certificate extensions. - CertificateList []CertificateEntry13 + CertificateList []CertificateEntry13 + marchalledExtensions [][]byte + marchalledExtensionsErr error } // Type returns the handshake message type. @@ -81,12 +83,32 @@ func (m *MessageCertificate13) Size() int { return 1 + len(m.CertificateRequestContext) + cert13CertLengthFieldSize + m.certsSize() } +func (m *MessageCertificate13) cacheMarshalExtensions() error { + if m.marchalledExtensions == nil && m.marchalledExtensionsErr == nil { + for _, entry := range m.CertificateList { + for _, ext := range entry.Extensions { + var c []byte + c, m.marchalledExtensionsErr = ext.Marshal() + if m.marchalledExtensionsErr != nil { + return m.marchalledExtensionsErr + } + m.marchalledExtensions = append(m.marchalledExtensions, c) + } + } + } + return m.marchalledExtensionsErr +} + func (m *MessageCertificate13) certsSize() int { + err := m.cacheMarshalExtensions() + if err != nil { + return 0 + } certificateListSize := 0 - for _, entry := range m.CertificateList { + for i, entry := range m.CertificateList { certificateListSize += cert13CertLengthFieldSize certificateListSize += len(entry.CertificateData) - certificateListSize += extension.Size(entry.Extensions) + certificateListSize += len(m.marchalledExtensions[i]) } return certificateListSize @@ -103,6 +125,11 @@ func (m *MessageCertificate13) MarshalInto(out []byte) error { return errBufferTooSmall } + err := m.cacheMarshalExtensions() + if err != nil { + return err + } + // Check size of certificate_list is still within bounds if m.certsSize() > maxUint24 { return errCertificateListTooLong diff --git a/pkg/protocol/handshake/message_client_hello.go b/pkg/protocol/handshake/message_client_hello.go index 3d73892dd..476bb326a 100644 --- a/pkg/protocol/handshake/message_client_hello.go +++ b/pkg/protocol/handshake/message_client_hello.go @@ -24,9 +24,11 @@ type MessageClientHello struct { SessionID []byte - CipherSuiteIDs []uint16 - CompressionMethods []*protocol.CompressionMethod - Extensions []extension.Extension + CipherSuiteIDs []uint16 + CompressionMethods []*protocol.CompressionMethod + Extensions []extension.Extension + marchalledExtensions []byte + marchalledExtensionsErr error } const handshakeMessageClientHelloVariableWidthStart = 34 @@ -36,11 +38,23 @@ func (m MessageClientHello) Type() Type { return TypeClientHello } +func (m *MessageClientHello) cacheMarshalExtensions() error { + if m.marchalledExtensions == nil && m.marchalledExtensionsErr == nil { + m.marchalledExtensions, m.marchalledExtensionsErr = extension.Marshal(m.Extensions) + } + return m.marchalledExtensionsErr +} + // Size returns the size needed for MarshalInto. func (m *MessageClientHello) Size() int { encodedCipherSuiteIDs := encodeCipherSuiteIDs(m.CipherSuiteIDs) encodedCompressionMethods := protocol.EncodeCompressionMethods(m.CompressionMethods) + err := m.cacheMarshalExtensions() + if err != nil { + return 0 + } + return handshakeMessageClientHelloVariableWidthStart + 1 + len(m.SessionID) + @@ -48,7 +62,7 @@ func (m *MessageClientHello) Size() int { len(m.Cookie) + len(encodedCipherSuiteIDs) + len(encodedCompressionMethods) + - extension.Size(m.Extensions) + len(m.marchalledExtensions) } // Marshal encodes the Handshake. @@ -71,7 +85,7 @@ func (m *MessageClientHello) MarshalInto(out []byte) error { return errCompressionMethodsTooLong } - extensions, err := extension.Marshal(m.Extensions) + err := m.cacheMarshalExtensions() if err != nil { return err } @@ -109,7 +123,7 @@ func (m *MessageClientHello) MarshalInto(out []byte) error { n = copy(out[offset:], encodedCompressionMethods) offset += n - copy(out[offset:], extensions) + copy(out[offset:], m.marchalledExtensions) return nil } diff --git a/pkg/protocol/handshake/message_server_hello.go b/pkg/protocol/handshake/message_server_hello.go index 4d9a05aea..9c822563c 100644 --- a/pkg/protocol/handshake/message_server_hello.go +++ b/pkg/protocol/handshake/message_server_hello.go @@ -22,9 +22,11 @@ type MessageServerHello struct { SessionID []byte - CipherSuiteID *uint16 - CompressionMethod *protocol.CompressionMethod - Extensions []extension.Extension + CipherSuiteID *uint16 + CompressionMethod *protocol.CompressionMethod + Extensions []extension.Extension + marchalledExtensions []byte + marchalledExtensionsErr error } const messageServerHelloVariableWidthStart = 2 + RandomLength @@ -34,10 +36,23 @@ func (m MessageServerHello) Type() Type { return TypeServerHello } +func (m *MessageServerHello) cacheMarshalExtensions() error { + if m.marchalledExtensions == nil && m.marchalledExtensionsErr == nil { + m.marchalledExtensions, m.marchalledExtensionsErr = extension.Marshal(m.Extensions) + } + return m.marchalledExtensionsErr +} + // Size returns the size required by MarshalInto. func (m *MessageServerHello) Size() int { + + err := m.cacheMarshalExtensions() + if err != nil { + return 0 + } + total := 0 - total += extension.Size(m.Extensions) + total += len(m.marchalledExtensions) total += messageServerHelloVariableWidthStart + 1 + len(m.SessionID) + 2 + 1 return total @@ -45,7 +60,8 @@ func (m *MessageServerHello) Size() int { // MarshalInto encodes the Handshake into a pre-allocated buffer. func (m *MessageServerHello) MarshalInto(out []byte) error { - extensions, err := extension.Marshal(m.Extensions) + + err := m.cacheMarshalExtensions() if err != nil { return err } @@ -74,7 +90,7 @@ func (m *MessageServerHello) MarshalInto(out []byte) error { out[offset] = byte(m.CompressionMethod.ID) offset += 1 - copy(out[offset:], extensions) + copy(out[offset:], m.marchalledExtensions) return nil } From aedd328b923e43a3f729f34458275c04ea404f52 Mon Sep 17 00:00:00 2001 From: Thomas Legris Date: Wed, 1 Apr 2026 22:56:27 +0900 Subject: [PATCH 4/5] fix tests --- conn.go | 5 +++++ pkg/protocol/handshake/message_certificate_13.go | 14 ++++++-------- pkg/protocol/handshake/message_client_hello.go | 1 + .../handshake/message_client_key_exchange.go | 1 + pkg/protocol/handshake/message_server_hello.go | 3 +-- 5 files changed, 14 insertions(+), 10 deletions(-) diff --git a/conn.go b/conn.go index 4987ae783..37425cdf4 100644 --- a/conn.go +++ b/conn.go @@ -885,6 +885,11 @@ func (c *Conn) fragmentHandshake(dtlsHandshake *handshake.Handshake) ([][]byte, } contentFragments := splitBytes(content, c.maximumTransmissionUnit) + if len(contentFragments) == 0 { + contentFragments = [][]byte{ + {}, + } + } offset := 0 fragmentedHandshakes := make([][]byte, 0, len(contentFragments)) diff --git a/pkg/protocol/handshake/message_certificate_13.go b/pkg/protocol/handshake/message_certificate_13.go index e3a2be6b9..907667c84 100644 --- a/pkg/protocol/handshake/message_certificate_13.go +++ b/pkg/protocol/handshake/message_certificate_13.go @@ -85,17 +85,15 @@ func (m *MessageCertificate13) Size() int { func (m *MessageCertificate13) cacheMarshalExtensions() error { if m.marchalledExtensions == nil && m.marchalledExtensionsErr == nil { - for _, entry := range m.CertificateList { - for _, ext := range entry.Extensions { - var c []byte - c, m.marchalledExtensionsErr = ext.Marshal() - if m.marchalledExtensionsErr != nil { - return m.marchalledExtensionsErr - } - m.marchalledExtensions = append(m.marchalledExtensions, c) + m.marchalledExtensions = make([][]byte, len(m.CertificateList)) + for i, entry := range m.CertificateList { + m.marchalledExtensions[i], m.marchalledExtensionsErr = extension.Marshal(entry.Extensions) + if m.marchalledExtensionsErr != nil { + return m.marchalledExtensionsErr } } } + return m.marchalledExtensionsErr } diff --git a/pkg/protocol/handshake/message_client_hello.go b/pkg/protocol/handshake/message_client_hello.go index 476bb326a..78e45a192 100644 --- a/pkg/protocol/handshake/message_client_hello.go +++ b/pkg/protocol/handshake/message_client_hello.go @@ -42,6 +42,7 @@ func (m *MessageClientHello) cacheMarshalExtensions() error { if m.marchalledExtensions == nil && m.marchalledExtensionsErr == nil { m.marchalledExtensions, m.marchalledExtensionsErr = extension.Marshal(m.Extensions) } + return m.marchalledExtensionsErr } diff --git a/pkg/protocol/handshake/message_client_key_exchange.go b/pkg/protocol/handshake/message_client_key_exchange.go index 627830d2c..df15cafb9 100644 --- a/pkg/protocol/handshake/message_client_key_exchange.go +++ b/pkg/protocol/handshake/message_client_key_exchange.go @@ -52,6 +52,7 @@ func (m *MessageClientKeyExchange) Size() int { total := 0 if m.IdentityHint != nil { total += 2 + total += len(m.IdentityHint) } if m.PublicKey != nil { diff --git a/pkg/protocol/handshake/message_server_hello.go b/pkg/protocol/handshake/message_server_hello.go index 9c822563c..0a10fb63e 100644 --- a/pkg/protocol/handshake/message_server_hello.go +++ b/pkg/protocol/handshake/message_server_hello.go @@ -40,12 +40,12 @@ func (m *MessageServerHello) cacheMarshalExtensions() error { if m.marchalledExtensions == nil && m.marchalledExtensionsErr == nil { m.marchalledExtensions, m.marchalledExtensionsErr = extension.Marshal(m.Extensions) } + return m.marchalledExtensionsErr } // Size returns the size required by MarshalInto. func (m *MessageServerHello) Size() int { - err := m.cacheMarshalExtensions() if err != nil { return 0 @@ -60,7 +60,6 @@ func (m *MessageServerHello) Size() int { // MarshalInto encodes the Handshake into a pre-allocated buffer. func (m *MessageServerHello) MarshalInto(out []byte) error { - err := m.cacheMarshalExtensions() if err != nil { return err From fa869ce30e5f53d2c06a65d2c631090314b10579 Mon Sep 17 00:00:00 2001 From: Thomas Legris Date: Mon, 6 Apr 2026 11:56:37 +0900 Subject: [PATCH 5/5] adjust API --- conn.go | 47 ++++-- go.mod | 8 +- go.sum | 12 +- internal/net/udp/packet_conn.go | 156 +++++++++--------- internal/net/udp/packet_conn_test.go | 4 +- listener.go | 3 +- pkg/crypto/ciphersuite/cbc.go | 10 +- pkg/crypto/ciphersuite/chacha20poly1305.go | 10 +- pkg/crypto/ciphersuite/ciphersuite.go | 40 +---- pkg/protocol/alert/alert.go | 18 +- pkg/protocol/application_data.go | 12 +- pkg/protocol/change_cipher_spec.go | 16 +- pkg/protocol/content.go | 4 +- pkg/protocol/handshake/handshake.go | 38 ++--- pkg/protocol/handshake/header.go | 10 +- pkg/protocol/handshake/message_certificate.go | 18 +- .../handshake/message_certificate_13.go | 26 +-- .../handshake/message_certificate_request.go | 20 +-- .../message_certificate_request_13.go | 24 +-- .../handshake/message_certificate_verify.go | 18 +- .../handshake/message_client_hello.go | 26 +-- .../handshake/message_client_key_exchange.go | 22 +-- pkg/protocol/handshake/message_finished.go | 18 +- .../handshake/message_hello_verify_request.go | 20 +-- .../handshake/message_server_hello.go | 20 +-- .../handshake/message_server_hello_done.go | 12 +- .../handshake/message_server_key_exchange.go | 28 ++-- pkg/protocol/recordlayer/header.go | 16 +- pkg/protocol/recordlayer/recordlayer.go | 27 ++- pkg/protocol/recordlayer/recordlayer_test.go | 2 +- 30 files changed, 346 insertions(+), 339 deletions(-) diff --git a/conn.go b/conn.go index 37425cdf4..52ad67828 100644 --- a/conn.go +++ b/conn.go @@ -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(), @@ -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. // @@ -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) } @@ -836,14 +851,13 @@ func (c *Conn) processHandshakePacket(pkt *packet, dtlsHandshake *handshake.Hand SequenceNumber: pkt.record.Header.SequenceNumber, } - hs := recordlayer.FixedHeaderSize + len(cidHeader.ConnectionID) - rawPacket = make([]byte, hs+len(rawInner)) - err = cidHeader.MarshalInto(rawPacket) + rawPacket = make([]byte, cidHeader.MarshalSize()+len(rawInner)) + _, err = cidHeader.MarshalTo(rawPacket) if err != nil { return nil, err } pkt.record.Header = *cidHeader - copy(rawPacket[hs:], rawInner) + copy(rawPacket[cidHeader.MarshalSize():], rawInner) } else { recordlayerHeader := &recordlayer.Header{ Version: pkt.record.Header.Version, @@ -853,15 +867,14 @@ func (c *Conn) processHandshakePacket(pkt *packet, dtlsHandshake *handshake.Hand SequenceNumber: seq, } - hs := recordlayer.FixedHeaderSize + len(recordlayerHeader.ConnectionID) - rawPacket = make([]byte, hs+len(handshakeFragment)) - err = recordlayerHeader.MarshalInto(rawPacket) + rawPacket = make([]byte, recordlayerHeader.MarshalSize()+len(handshakeFragment)) + _, err = recordlayerHeader.MarshalTo(rawPacket) if err != nil { return nil, err } pkt.record.Header = *recordlayerHeader - copy(rawPacket[hs:], handshakeFragment) + copy(rawPacket[recordlayerHeader.MarshalSize():], handshakeFragment) } if pkt.shouldEncrypt { @@ -907,7 +920,7 @@ func (c *Conn) fragmentHandshake(dtlsHandshake *handshake.Handshake) ([][]byte, offset += contentFragmentLen fragmentedHandshake := make([]byte, handshake.HeaderLength+len(contentFragment)) - err := headerFragment.MarshalInto(fragmentedHandshake) + _, err := headerFragment.MarshalTo(fragmentedHandshake) if err != nil { return nil, err } @@ -990,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 { @@ -1121,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 diff --git a/go.mod b/go.mod index 01b71803a..34a80ac6b 100644 --- a/go.mod +++ b/go.mod @@ -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 diff --git a/go.sum b/go.sum index c34341904..f646ec5bb 100644 --- a/go.sum +++ b/go.sum @@ -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= diff --git a/internal/net/udp/packet_conn.go b/internal/net/udp/packet_conn.go index 0ea82b660..6febc8365 100644 --- a/internal/net/udp/packet_conn.go +++ b/internal/net/udp/packet_conn.go @@ -20,6 +20,7 @@ import ( "context" "errors" "net" + "runtime" "sync" "sync/atomic" "time" @@ -31,7 +32,7 @@ import ( const ( receiveMTU = 8192 - defaultListenBacklog = 128 // same as Linux default + defaultListenBacklog = 8192 // same as Linux default ) // Typed errors. @@ -42,7 +43,8 @@ var ( // listener augments a connection-oriented Listener over a UDP PacketConn. type listener struct { - pConn *net.UDPConn + pConns []*net.UDPConn + iConns atomic.Uint32 accepting atomic.Value // bool acceptCh chan *PacketConn @@ -51,10 +53,11 @@ type listener struct { acceptFilter func([]byte) bool datagramRouter func([]byte) (string, bool) connIdentifier func([]byte) (string, bool) + onConnAttempt func(net.Addr) error - connLock sync.Mutex - conns map[string]*PacketConn - connWG sync.WaitGroup + conns sync.Map //map[string]*PacketConn + nConns atomic.Int64 + connWG sync.WaitGroup readWG sync.WaitGroup errClose atomic.Value // error @@ -63,6 +66,10 @@ type listener struct { errRead atomic.Value // error } +func (l *listener) nextSock() *net.UDPConn { + return l.pConns[l.iConns.Add(1)%uint32(len(l.pConns))] +} + // Accept waits for and returns the next connection to the listener. func (l *listener) Accept() (net.PacketConn, net.Addr, error) { select { @@ -89,7 +96,6 @@ func (l *listener) Close() error { l.accepting.Store(false) close(l.doneCh) - l.connLock.Lock() // Close unaccepted connections lclose: for { @@ -99,31 +105,23 @@ func (l *listener) Close() error { // If we have an alternate identifier, remove it from the connection // map. if id := c.id.Load(); id != nil { - delete(l.conns, id.(string)) //nolint:forcetypeassert - } - // If we haven't already removed the remote address, remove it - // from the connection map. - if c.rmraddr.Load() == nil { - delete(l.conns, c.raddr.String()) - c.rmraddr.Store(true) + l.conns.Delete(id) + l.nConns.Add(-1) } + l.conns.Delete(c.raddr.String()) + l.nConns.Add(-1) default: break lclose } } - nConns := len(l.conns) - l.connLock.Unlock() + l.conns.Clear() l.connWG.Done() - if nConns == 0 { - // Wait if this is the final connection. - l.readWG.Wait() - if errClose, ok := l.errClose.Load().(error); ok { - err = errClose - } - } else { - err = nil + // Wait if this is the final connection. + l.readWG.Wait() + if errClose, ok := l.errClose.Load().(error); ok { + err = errClose } }) @@ -132,7 +130,7 @@ func (l *listener) Close() error { // Addr returns the listener's network address. func (l *listener) Addr() net.Addr { - return l.pConn.LocalAddr() + return l.pConns[0].LocalAddr() } // ListenConfig stores options for listening to an address. @@ -149,6 +147,10 @@ type ListenConfig struct { // the incoming packet. If not set, any packet creates new conn. AcceptFilter func([]byte) bool + // OnConnAttempt determines whether the new conn should be made for + // the incoming packet. If not set, any packet creates new conn. + OnConnAttempt func(net.Addr) error + // DatagramRouter routes an incoming datagram to a connection by extracting // an identifier from the its paylod DatagramRouter func([]byte) (string, bool) @@ -174,21 +176,25 @@ func (lc *ListenConfig) Listen(network string, laddr *net.UDPAddr) (dtlsnet.Pack if laddr != nil { laddrStr = laddr.String() } - innerConn, err := lc.ListenConfig.ListenPacket(context.Background(), network, laddrStr) - if err != nil { - return nil, err - } - conn, ok := innerConn.(*net.UDPConn) - if !ok { - return nil, errors.New("listen packet not a *net.UDPConn") //nolint:err113 + var conns []*net.UDPConn + for range runtime.NumCPU() { + innerConn, err := lc.ListenConfig.ListenPacket(context.Background(), network, laddrStr) + if err != nil { + return nil, err + } + conn, ok := innerConn.(*net.UDPConn) + if !ok { + return nil, errors.New("listen packet not a *net.UDPConn") //nolint:err113 + } + conns = append(conns, conn) } packetListener := &listener{ - pConn: conn, + pConns: conns, acceptCh: make(chan *PacketConn, lc.Backlog), - conns: make(map[string]*PacketConn), doneCh: make(chan struct{}), acceptFilter: lc.AcceptFilter, + onConnAttempt: lc.OnConnAttempt, datagramRouter: lc.DatagramRouter, connIdentifier: lc.ConnectionIdentifier, readDoneCh: make(chan struct{}), @@ -198,11 +204,15 @@ func (lc *ListenConfig) Listen(network string, laddr *net.UDPAddr) (dtlsnet.Pack packetListener.connWG.Add(1) packetListener.readWG.Add(2) // wait readLoop and Close execution routine - go packetListener.readLoop() + for _, conn := range packetListener.pConns { + go packetListener.readLoop(conn) + } go func() { packetListener.connWG.Wait() - if err := packetListener.pConn.Close(); err != nil { - packetListener.errClose.Store(err) + for _, conn := range packetListener.pConns { + if err := conn.Close(); err != nil { + packetListener.errClose.Store(err) + } } packetListener.readWG.Done() }() @@ -217,14 +227,13 @@ func Listen(network string, laddr *net.UDPAddr) (dtlsnet.PacketListener, error) // readLoop dispatches packets to the proper connection, creating a new one if // necessary, until all connections are closed. -func (l *listener) readLoop() { +func (l *listener) readLoop(conn *net.UDPConn) { defer l.readWG.Done() defer close(l.readDoneCh) - buf := make([]byte, receiveMTU) - + var buf [64 * 1024]byte for { - n, raddr, err := l.pConn.ReadFrom(buf) + n, raddr, err := conn.ReadFrom(buf[:]) if err != nil { l.errRead.Store(err) @@ -242,21 +251,24 @@ func (l *listener) readLoop() { // getConn gets an existing connection or creates a new one. func (l *listener) getConn(raddr net.Addr, buf []byte) (*PacketConn, bool, error) { //nolint:cyclop - l.connLock.Lock() - defer l.connLock.Unlock() // If we have a custom resolver, use it. if l.datagramRouter != nil { if id, ok := l.datagramRouter(buf); ok { - if conn, ok := l.conns[id]; ok { - return conn, true, nil + if conn, ok := l.conns.Load(id); ok { + return conn.(*PacketConn), true, nil } } } // If we don't have a custom resolver, or we were unable to find an // associated connection, fall back to remote address. - conn, ok := l.conns[raddr.String()] + conn, ok := l.conns.Load(raddr.String()) if !ok { + if l.onConnAttempt != nil { + if err := l.onConnAttempt(raddr); err != nil { + return nil, false, nil + } + } if isAccepting, ok := l.accepting.Load().(bool); !isAccepting || !ok { return nil, false, ErrClosedListener } @@ -265,16 +277,19 @@ func (l *listener) getConn(raddr net.Addr, buf []byte) (*PacketConn, bool, error return nil, false, nil } } - conn = l.newPacketConn(raddr) - select { - case l.acceptCh <- conn: - l.conns[raddr.String()] = conn - default: - return nil, false, ErrListenQueueExceeded + conn, ok = l.conns.LoadOrStore(raddr.String(), l.newPacketConn(raddr)) + if !ok { + select { + case l.acceptCh <- conn.(*PacketConn): + l.nConns.Add(1) + default: + l.conns.Delete(raddr.String()) + return nil, false, ErrListenQueueExceeded + } } } - return conn, true, nil + return conn.(*PacketConn), true, nil } // PacketConn is a net.PacketConn implementation that is able to dictate its @@ -284,9 +299,8 @@ func (l *listener) getConn(raddr net.Addr, buf []byte) (*PacketConn, bool, error type PacketConn struct { listener *listener - raddr net.Addr - rmraddr atomic.Value // bool - id atomic.Value // string + raddr net.Addr + id atomic.Value // string buffer *idtlsnet.PacketBuffer @@ -324,9 +338,7 @@ func (c *PacketConn) WriteTo(payload []byte, addr net.Addr) (n int, err error) { candidate, ok := c.listener.connIdentifier(payload) // If we have an identifier, add entry to connection map. if ok { - c.listener.connLock.Lock() - c.listener.conns[candidate] = c - c.listener.connLock.Unlock() + c.listener.conns.Store(candidate, c) c.id.Store(candidate) } } @@ -344,11 +356,8 @@ func (c *PacketConn) WriteTo(payload []byte, addr net.Addr) (n int, err error) { // resulting in the remote address entry being dropped prior to the // "real" client transitioning to sending using the alternate // identifier. - if id != nil && c.rmraddr.Load() == nil && addr.String() != c.raddr.String() { - c.listener.connLock.Lock() - delete(c.listener.conns, c.raddr.String()) - c.rmraddr.Store(true) - c.listener.connLock.Unlock() + if id != nil && addr.String() != c.raddr.String() { + c.listener.conns.Delete(c.raddr.String()) } } @@ -358,7 +367,7 @@ func (c *PacketConn) WriteTo(payload []byte, addr net.Addr) (n int, err error) { default: } - return c.listener.pConn.WriteTo(payload, addr) + return c.listener.nextSock().WriteTo(payload, addr) } // Close closes the conn and releases any Read calls. @@ -367,22 +376,17 @@ func (c *PacketConn) Close() error { c.doneOnce.Do(func() { c.listener.connWG.Done() close(c.doneCh) - c.listener.connLock.Lock() // If we have an alternate identifier, remove it from the connection // map. if id := c.id.Load(); id != nil { - delete(c.listener.conns, id.(string)) //nolint:forcetypeassert - } - // If we haven't already removed the remote address, remove it from the - // connection map. - if c.rmraddr.Load() == nil { - delete(c.listener.conns, c.raddr.String()) - c.rmraddr.Store(true) + c.listener.conns.Delete(id) //nolint:forcetypeassert + c.listener.nConns.Add(-1) } - nConns := len(c.listener.conns) - c.listener.connLock.Unlock() + c.listener.conns.Delete(c.raddr.String()) + c.listener.nConns.Add(-1) - if isAccepting, ok := c.listener.accepting.Load().(bool); nConns == 0 && !isAccepting && ok { + if isAccepting, ok := c.listener.accepting.Load().(bool); c.listener.nConns.Load() == 0 && + !isAccepting && ok { // Wait if this is the final connection c.listener.readWG.Wait() if errClose, ok := c.listener.errClose.Load().(error); ok { @@ -402,7 +406,7 @@ func (c *PacketConn) Close() error { // LocalAddr implements net.PacketConn.LocalAddr. func (c *PacketConn) LocalAddr() net.Addr { - return c.listener.pConn.LocalAddr() + return c.listener.nextSock().LocalAddr() } // SetDeadline implements net.PacketConn.SetDeadline. diff --git a/internal/net/udp/packet_conn_test.go b/internal/net/udp/packet_conn_test.go index 7264e6481..a74d17460 100644 --- a/internal/net/udp/packet_conn_test.go +++ b/internal/net/udp/packet_conn_test.go @@ -354,7 +354,7 @@ func TestConnClose(t *testing.T) { //nolint:cyclop // Close l.pConn to inject error. listener, ok := udpListener.(*listener) assert.True(t, ok) - assert.NoError(t, listener.pConn.Close()) + assert.NoError(t, listener.nextSock().Close()) assert.NoError(t, cb.Close()) assert.NoError(t, ca.Close()) assert.Error(t, udpListener.Close()) @@ -370,7 +370,7 @@ func TestConnClose(t *testing.T) { //nolint:cyclop // Close l.pConn to inject error. listener, ok := l.(*listener) assert.True(t, ok) - assert.NoError(t, listener.pConn.Close()) + assert.NoError(t, listener.nextSock().Close()) assert.NoError(t, cb.Close()) assert.NoError(t, l.Close()) assert.Error(t, ca.Close()) diff --git a/listener.go b/listener.go index 0ad26bbed..ba4296375 100644 --- a/listener.go +++ b/listener.go @@ -33,7 +33,8 @@ func Listen(network string, laddr *net.UDPAddr, config *Config) (net.Listener, e return h.ContentType == protocol.ContentTypeHandshake }, - ListenConfig: config.listenConfig, + OnConnAttempt: config.OnConnectionAttempt, + ListenConfig: config.listenConfig, } // If connection ID support is enabled, then they must be supported in // routing. diff --git a/pkg/crypto/ciphersuite/cbc.go b/pkg/crypto/ciphersuite/cbc.go index 4f33d6ac1..d1777d197 100644 --- a/pkg/crypto/ciphersuite/cbc.go +++ b/pkg/crypto/ciphersuite/cbc.go @@ -68,8 +68,8 @@ func NewCBC( // Encrypt encrypt a DTLS RecordLayer message. func (c *CBC) Encrypt(pkt *recordlayer.RecordLayer, raw []byte) ([]byte, error) { - payload := raw[pkt.Header.Size():] - raw = raw[:pkt.Header.Size()] + payload := raw[pkt.Header.MarshalSize():] + raw = raw[:pkt.Header.MarshalSize()] blockSize := c.writeCBC.BlockSize() // Generate + Append MAC @@ -110,7 +110,7 @@ func (c *CBC) Encrypt(pkt *recordlayer.RecordLayer, raw []byte) ([]byte, error) raw = append(raw, payload...) // Update recordLayer size to include IV+MAC+Padding - binary.BigEndian.PutUint16(raw[pkt.Header.Size()-2:], uint16(len(raw)-pkt.Header.Size())) //nolint:gosec //G115 + binary.BigEndian.PutUint16(raw[pkt.Header.MarshalSize()-2:], uint16(len(raw)-pkt.Header.MarshalSize())) //nolint:gosec //G115 return raw, nil } @@ -123,7 +123,7 @@ func (c *CBC) Decrypt(header recordlayer.Header, in []byte) ([]byte, error) { if err := header.Unmarshal(in); err != nil { return nil, err } - body := in[header.Size():] + body := in[header.MarshalSize():] switch { case header.ContentType == protocol.ContentTypeChangeCipherSpec: @@ -171,7 +171,7 @@ func (c *CBC) Decrypt(header recordlayer.Header, in []byte) ([]byte, error) { return nil, errInvalidMAC } - return append(in[:header.Size()], body[:dataEnd]...), nil + return append(in[:header.MarshalSize()], body[:dataEnd]...), nil } func (c *CBC) hmac( diff --git a/pkg/crypto/ciphersuite/chacha20poly1305.go b/pkg/crypto/ciphersuite/chacha20poly1305.go index 05e7cf5f0..8ae55c695 100644 --- a/pkg/crypto/ciphersuite/chacha20poly1305.go +++ b/pkg/crypto/ciphersuite/chacha20poly1305.go @@ -51,8 +51,8 @@ func NewChaCha20Poly1305(localKey, localWriteIV, remoteKey, remoteWriteIV []byte // Encrypt encrypts a DTLS RecordLayer message. func (c *ChaCha20Poly1305) Encrypt(pkt *recordlayer.RecordLayer, raw []byte) ([]byte, error) { - payload := raw[pkt.Header.Size():] - raw = raw[:pkt.Header.Size()] + payload := raw[pkt.Header.MarshalSize():] + raw = raw[:pkt.Header.MarshalSize()] var nonce [chachaNonceLength]byte copy(nonce[:], c.localWriteIV) @@ -80,7 +80,7 @@ func (c *ChaCha20Poly1305) Encrypt(pkt *recordlayer.RecordLayer, raw []byte) ([] copy(result, raw) copy(result[len(raw):], encrypted) - binary.BigEndian.PutUint16(result[pkt.Header.Size()-2:], uint16(len(encrypted))) //nolint:gosec + binary.BigEndian.PutUint16(result[pkt.Header.MarshalSize()-2:], uint16(len(encrypted))) //nolint:gosec return result, nil } @@ -108,7 +108,7 @@ func (c *ChaCha20Poly1305) Decrypt(header recordlayer.Header, in []byte) ([]byte } // NOTE: ChaCha20-Poly1305 has NO explicit nonce in the record - ciphertext := in[header.Size():] + ciphertext := in[header.MarshalSize():] var additionalData []byte if header.ContentType == protocol.ContentTypeConnectionID { @@ -122,5 +122,5 @@ func (c *ChaCha20Poly1305) Decrypt(header recordlayer.Header, in []byte) ([]byte return nil, fmt.Errorf("%w: %v", errDecryptPacket, err) //nolint:errorlint } - return append(in[:header.Size()], plaintext...), nil + return append(in[:header.MarshalSize()], plaintext...), nil } diff --git a/pkg/crypto/ciphersuite/ciphersuite.go b/pkg/crypto/ciphersuite/ciphersuite.go index e917bb7b1..fbd760e54 100644 --- a/pkg/crypto/ciphersuite/ciphersuite.go +++ b/pkg/crypto/ciphersuite/ciphersuite.go @@ -9,7 +9,6 @@ import ( "encoding/binary" "errors" "fmt" - "sync" "github.com/pion/dtls/v3/pkg/protocol" "github.com/pion/dtls/v3/pkg/protocol/recordlayer" @@ -41,9 +40,6 @@ type aead struct { remoteWriteIV []byte nonceLength int tagLength int - - // buffer pool for (fixed-size) nonces. - nonceBufferPool sync.Pool } // newAEAD creates a generic DTLS AEAD-based Cipher. @@ -62,23 +58,16 @@ func newAEAD( remoteWriteIV: remoteWriteIV, nonceLength: nonceLength, tagLength: tagLength, - nonceBufferPool: sync.Pool{ - New: func() any { - b := make([]byte, nonceLength) - return &b // nolint:nlreturn - }, - }, } } // encrypt encrypts a DTLS RecordLayer message. func (a *aead) encrypt(pkt *recordlayer.RecordLayer, raw []byte) ([]byte, error) { - payload := raw[pkt.Header.Size():] - raw = raw[:pkt.Header.Size()] + payload := raw[pkt.Header.MarshalSize():] + raw = raw[:pkt.Header.MarshalSize()] // Get nonce buffer from pool - noncePtr := a.nonceBufferPool.Get().(*[]byte) // nolint:forcetypeassert - nonce := *noncePtr + nonce := make([]byte, a.nonceLength) copy(nonce, a.localWriteIV[:4]) @@ -100,10 +89,7 @@ func (a *aead) encrypt(pkt *recordlayer.RecordLayer, raw []byte) ([]byte, error) a.localAEAD.Seal(r[len(raw)+8:len(raw)+8], nonce, payload, additionalData) // Update recordLayer size to include explicit nonce - binary.BigEndian.PutUint16(r[pkt.Header.Size()-2:], uint16(len(r)-pkt.Header.Size())) //nolint:gosec //G115 - - // Return nonce buffer to pool - a.nonceBufferPool.Put(noncePtr) + binary.BigEndian.PutUint16(r[pkt.Header.MarshalSize()-2:], uint16(len(r)-pkt.Header.MarshalSize())) //nolint:gosec //G115 return r, nil } @@ -117,17 +103,15 @@ func (a *aead) decrypt(header recordlayer.Header, in []byte) ([]byte, error) { case header.ContentType == protocol.ContentTypeChangeCipherSpec: // Nothing to encrypt with ChangeCipherSpec return in, nil - case len(in) <= (8 + header.Size()): + case len(in) <= (8 + header.MarshalSize()): return nil, errNotEnoughRoomForNonce } - // Get nonce buffer from pool - noncePtr := a.nonceBufferPool.Get().(*[]byte) // nolint:forcetypeassert - nonce := *noncePtr + nonce := make([]byte, a.nonceLength) copy(nonce[:4], a.remoteWriteIV[:4]) - copy(nonce[4:], in[header.Size():header.Size()+8]) - out := in[header.Size()+8:] + copy(nonce[4:], in[header.MarshalSize():header.MarshalSize()+8]) + out := in[header.MarshalSize()+8:] var additionalData []byte if header.ContentType == protocol.ContentTypeConnectionID { @@ -137,16 +121,10 @@ func (a *aead) decrypt(header recordlayer.Header, in []byte) ([]byte, error) { } out, err = a.remoteAEAD.Open(out[:0], nonce, out, additionalData) if err != nil { - // Return nonce buffer to pool - a.nonceBufferPool.Put(noncePtr) - return nil, fmt.Errorf("%w: %v", errDecryptPacket, err) //nolint:errorlint } - // Return nonce buffer to pool - a.nonceBufferPool.Put(noncePtr) - - return append(in[:header.Size()], out...), nil + return append(in[:header.MarshalSize()], out...), nil } func generateAEADAdditionalData(h *recordlayer.Header, payloadLen int) []byte { diff --git a/pkg/protocol/alert/alert.go b/pkg/protocol/alert/alert.go index d85d72f27..76d1afddb 100644 --- a/pkg/protocol/alert/alert.go +++ b/pkg/protocol/alert/alert.go @@ -145,28 +145,28 @@ func (a Alert) ContentType() protocol.ContentType { return protocol.ContentTypeAlert } -// Size returns the minimal buffer size required for MarshalInto. -func (a Alert) Size() int { +// MarshalSize returns the minimal buffer size required for MarshalTo. +func (a Alert) MarshalSize() int { return 2 } // Marshal returns the encoded alert. func (a *Alert) Marshal() ([]byte, error) { - out := make([]byte, a.Size()) - err := a.MarshalInto(out) + out := make([]byte, a.MarshalSize()) + _, err := a.MarshalTo(out) return out, err } -// MarshalInto returns the encoded alert. -func (a *Alert) MarshalInto(out []byte) error { - if len(out) < a.Size() { - return errBufferTooSmall +// MarshalTo returns the encoded alert. +func (a *Alert) MarshalTo(out []byte) (int, error) { + if len(out) < a.MarshalSize() { + return 0, errBufferTooSmall } out[0] = byte(a.Level) out[1] = byte(a.Description) - return nil + return 2, nil } // Unmarshal populates the alert from binary data. diff --git a/pkg/protocol/application_data.go b/pkg/protocol/application_data.go index 1e290ee0e..5467d73bb 100644 --- a/pkg/protocol/application_data.go +++ b/pkg/protocol/application_data.go @@ -20,20 +20,20 @@ func (a ApplicationData) ContentType() ContentType { // Marshal encodes the ApplicationData to binary. func (a *ApplicationData) Marshal() ([]byte, error) { out := make([]byte, len(a.Data)) - err := a.MarshalInto(out) + _, err := a.MarshalTo(out) return out, err } -// MarshalInto encodes the ApplicationData to binary into a pre-allocated buffer. -func (a *ApplicationData) MarshalInto(out []byte) error { +// MarshalTo encodes the ApplicationData to binary into a pre-allocated buffer. +func (a *ApplicationData) MarshalTo(out []byte) (int, error) { copy(out, a.Data) - return nil + return len(a.Data), nil } -// Size returns the size required for MarshalInto. -func (a ApplicationData) Size() int { +// MarshalSize returns the size required for MarshalTo. +func (a ApplicationData) MarshalSize() int { return len(a.Data) } diff --git a/pkg/protocol/change_cipher_spec.go b/pkg/protocol/change_cipher_spec.go index c1bc357b7..f0d210e11 100644 --- a/pkg/protocol/change_cipher_spec.go +++ b/pkg/protocol/change_cipher_spec.go @@ -15,27 +15,27 @@ func (c ChangeCipherSpec) ContentType() ContentType { return ContentTypeChangeCipherSpec } -// Size returns the minimal buffer size required for MarshalInto. -func (c ChangeCipherSpec) Size() int { +// MarshalSize returns the minimal buffer size required for MarshalTo. +func (c ChangeCipherSpec) MarshalSize() int { return 1 } // Marshal encodes the ChangeCipherSpec to binary. func (c *ChangeCipherSpec) Marshal() ([]byte, error) { out := make([]byte, 1) - err := c.MarshalInto(out) + _, err := c.MarshalTo(out) return out, err } -// MarshalInto encodes the ChangeCipherSpec to binary into a pre-allocated buffer. -func (c *ChangeCipherSpec) MarshalInto(out []byte) error { - if len(out) < c.Size() { - return errBufferTooSmall +// MarshalTo encodes the ChangeCipherSpec to binary into a pre-allocated buffer. +func (c *ChangeCipherSpec) MarshalTo(out []byte) (int, error) { + if len(out) < c.MarshalSize() { + return 0, errBufferTooSmall } out[0] = 0x01 - return nil + return 1, nil } // Unmarshal populates the ChangeCipherSpec from binary. diff --git a/pkg/protocol/content.go b/pkg/protocol/content.go index 56ccaccc5..26c9f46fb 100644 --- a/pkg/protocol/content.go +++ b/pkg/protocol/content.go @@ -21,7 +21,7 @@ const ( type Content interface { ContentType() ContentType Marshal() ([]byte, error) - MarshalInto([]byte) error + MarshalTo([]byte) (int, error) Unmarshal(data []byte) error - Size() int + MarshalSize() int } diff --git a/pkg/protocol/handshake/handshake.go b/pkg/protocol/handshake/handshake.go index 3ef0c8a5f..fdb3f1667 100644 --- a/pkg/protocol/handshake/handshake.go +++ b/pkg/protocol/handshake/handshake.go @@ -62,8 +62,8 @@ func (t Type) String() string { //nolint:cyclop // Message is the body of a Handshake datagram. type Message interface { Marshal() ([]byte, error) - MarshalInto([]byte) error - Size() int + MarshalTo([]byte) (int, error) + MarshalSize() int Unmarshal(data []byte) error Type() Type } @@ -86,9 +86,9 @@ func (h Handshake) ContentType() protocol.ContentType { return protocol.ContentTypeHandshake } -// Size returns the minimal buffer size required for MarshalInto. -func (h *Handshake) Size() int { - return HeaderLength + h.Message.Size() +// MarshalSize returns the minimal buffer size required for MarshalTo. +func (h *Handshake) MarshalSize() int { + return HeaderLength + h.Message.MarshalSize() } // Marshal encodes a handshake into a binary message. @@ -98,38 +98,38 @@ func (h *Handshake) Marshal() ([]byte, error) { } else if h.Header.FragmentOffset != 0 { return nil, errUnableToMarshalFragmented } - out := make([]byte, h.Size()) - err := h.MarshalInto(out) + out := make([]byte, h.MarshalSize()) + _, err := h.MarshalTo(out) return out, err } -// MarshalInto encodes a handshake into a binary message into a pre-allocated buffer. -func (h *Handshake) MarshalInto(out []byte) error { +// MarshalTo encodes a handshake into a binary message into a pre-allocated buffer. +func (h *Handshake) MarshalTo(out []byte) (int, error) { if h.Message == nil { - return errHandshakeMessageUnset + return 0, errHandshakeMessageUnset } else if h.Header.FragmentOffset != 0 { - return errUnableToMarshalFragmented + return 0, errUnableToMarshalFragmented } - if len(out) < h.Size() { - return errBufferTooSmall + if len(out) < h.MarshalSize() { + return 0, errBufferTooSmall } - h.Header.Length = uint32(h.Message.Size()) //nolint:gosec // G115 + h.Header.Length = uint32(h.Message.MarshalSize()) //nolint:gosec // G115 h.Header.FragmentLength = h.Header.Length h.Header.Type = h.Message.Type() - err := h.Header.MarshalInto(out) + _, err := h.Header.MarshalTo(out) if err != nil { - return err + return 0, err } - err = h.Message.MarshalInto(out[HeaderLength:]) + _, err = h.Message.MarshalTo(out[HeaderLength:]) if err != nil { - return err + return 0, err } - return nil + return h.MarshalSize(), nil } // Unmarshal decodes a handshake from a binary message. diff --git a/pkg/protocol/handshake/header.go b/pkg/protocol/handshake/header.go index d19d384cd..b093b8826 100644 --- a/pkg/protocol/handshake/header.go +++ b/pkg/protocol/handshake/header.go @@ -29,15 +29,15 @@ type Header struct { // Marshal encodes the Header. func (h *Header) Marshal() ([]byte, error) { out := make([]byte, HeaderLength) - err := h.MarshalInto(out) + _, err := h.MarshalTo(out) return out, err } -// MarshalInto encodes the Header using a pre-allocated buffer. -func (h *Header) MarshalInto(out []byte) error { +// MarshalTo encodes the Header using a pre-allocated buffer. +func (h *Header) MarshalTo(out []byte) (int, error) { if len(out) < HeaderLength { - return errBufferTooSmall + return 0, errBufferTooSmall } out[0] = byte(h.Type) @@ -46,7 +46,7 @@ func (h *Header) MarshalInto(out []byte) error { util.PutBigEndianUint24(out[6:], h.FragmentOffset) util.PutBigEndianUint24(out[9:], h.FragmentLength) - return nil + return HeaderLength, nil } // Unmarshal populates the header from encoded data. diff --git a/pkg/protocol/handshake/message_certificate.go b/pkg/protocol/handshake/message_certificate.go index 2ccef174b..a455e91b3 100644 --- a/pkg/protocol/handshake/message_certificate.go +++ b/pkg/protocol/handshake/message_certificate.go @@ -24,8 +24,8 @@ const ( handshakeMessageCertificateLengthFieldSize = 3 ) -// Size returns the minimal size required for MarshalInto. -func (m *MessageCertificate) Size() int { +// MarshalSize returns the minimal size required for MarshalTo. +func (m *MessageCertificate) MarshalSize() int { total := handshakeMessageCertificateLengthFieldSize for _, cert := range m.Certificate { @@ -37,17 +37,17 @@ func (m *MessageCertificate) Size() int { // Marshal encodes the Handshake. func (m *MessageCertificate) Marshal() ([]byte, error) { - out := make([]byte, m.Size()) - err := m.MarshalInto(out) + out := make([]byte, m.MarshalSize()) + _, err := m.MarshalTo(out) return out, err } -// MarshalInto encodes the Handshake into a pre-allocated buffer. -func (m *MessageCertificate) MarshalInto(out []byte) error { - // Total Payload Size +// MarshalTo encodes the Handshake into a pre-allocated buffer. +func (m *MessageCertificate) MarshalTo(out []byte) (int, error) { + // Total Payload MarshalSize //nolint:gosec // G115 - util.PutBigEndianUint24(out, uint32(m.Size()-handshakeMessageCertificateLengthFieldSize)) + util.PutBigEndianUint24(out, uint32(m.MarshalSize()-handshakeMessageCertificateLengthFieldSize)) offset := handshakeMessageCertificateLengthFieldSize for _, cert := range m.Certificate { @@ -61,7 +61,7 @@ func (m *MessageCertificate) MarshalInto(out []byte) error { offset += len(cert) } - return nil + return m.MarshalSize(), nil } // Unmarshal populates the message from encoded data. diff --git a/pkg/protocol/handshake/message_certificate_13.go b/pkg/protocol/handshake/message_certificate_13.go index 907667c84..bbdfb015c 100644 --- a/pkg/protocol/handshake/message_certificate_13.go +++ b/pkg/protocol/handshake/message_certificate_13.go @@ -73,13 +73,13 @@ func (m *MessageCertificate13) Marshal() ([]byte, error) { return nil, errCertificateRequestContextTooLong } - out := make([]byte, m.Size()) - err := m.MarshalInto(out) + out := make([]byte, m.MarshalSize()) + _, err := m.MarshalTo(out) return out, err } -func (m *MessageCertificate13) Size() int { +func (m *MessageCertificate13) MarshalSize() int { return 1 + len(m.CertificateRequestContext) + cert13CertLengthFieldSize + m.certsSize() } @@ -112,25 +112,25 @@ func (m *MessageCertificate13) certsSize() int { return certificateListSize } -// MarshalInto is same as Marshal but uses a pre-allocated buffer. -func (m *MessageCertificate13) MarshalInto(out []byte) error { +// MarshalTo is same as Marshal but uses a pre-allocated buffer. +func (m *MessageCertificate13) MarshalTo(out []byte) (int, error) { // Validate certificate_request_context length if len(m.CertificateRequestContext) > cert13ContextMaxLength { - return errCertificateRequestContextTooLong + return 0, errCertificateRequestContextTooLong } - if len(out) < m.Size() { - return errBufferTooSmall + if len(out) < m.MarshalSize() { + return 0, errBufferTooSmall } err := m.cacheMarshalExtensions() if err != nil { - return err + return 0, err } // Check size of certificate_list is still within bounds if m.certsSize() > maxUint24 { - return errCertificateListTooLong + return 0, errCertificateListTooLong } // Start with certificate_request_context (1-byte length prefix) @@ -150,7 +150,7 @@ func (m *MessageCertificate13) MarshalInto(out []byte) error { // Add cert_data as a 3-byte length prefix certDataLen := len(entry.CertificateData) if certDataLen == 0 || certDataLen > maxUint24 { - return errInvalidCertificateEntry + return 0, errInvalidCertificateEntry } util.PutBigEndianUint24(out[offset:], uint32(certDataLen)) //nolint:gosec // G115 offset += 3 @@ -160,13 +160,13 @@ func (m *MessageCertificate13) MarshalInto(out []byte) error { // Marshal extensions (includes a 2-byte length prefix) extensionsData, err := extension.Marshal(entry.Extensions) if err != nil { - return err + return 0, err } n = copy(out[offset:], extensionsData) offset += n } - return nil + return m.MarshalSize(), nil } // parseCertificate13Entry parses a single certificate entry from the cryptobyte string. diff --git a/pkg/protocol/handshake/message_certificate_request.go b/pkg/protocol/handshake/message_certificate_request.go index 8520e31e6..ad6448311 100644 --- a/pkg/protocol/handshake/message_certificate_request.go +++ b/pkg/protocol/handshake/message_certificate_request.go @@ -35,8 +35,8 @@ func (m MessageCertificateRequest) Type() Type { return TypeCertificateRequest } -// Size returns the minimal size required for MarshalInto. -func (m *MessageCertificateRequest) Size() int { +// MarshalSize returns the minimal size required for MarshalTo. +func (m *MessageCertificateRequest) MarshalSize() int { return 1 + len(m.CertificateTypes) + 2 + // number of SignatureHashAlgorithms @@ -56,20 +56,20 @@ func (m *MessageCertificateRequest) casLength() int { // Marshal encodes the Handshake. func (m *MessageCertificateRequest) Marshal() ([]byte, error) { - out := make([]byte, m.Size()) - err := m.MarshalInto(out) + out := make([]byte, m.MarshalSize()) + _, err := m.MarshalTo(out) return out, err } -// MarshalInto encodes the Handshake into a pre-allocated buffer. -func (m *MessageCertificateRequest) MarshalInto(out []byte) error { +// MarshalTo encodes the Handshake into a pre-allocated buffer. +func (m *MessageCertificateRequest) MarshalTo(out []byte) (int, error) { if len(m.CertificateTypes) > 255 { - return errCertificateTypesTooLong + return 0, errCertificateTypesTooLong } - if len(out) < m.Size() { - return errBufferTooSmall + if len(out) < m.MarshalSize() { + return 0, errBufferTooSmall } //nolint:gosec // G115: certificate types count is validated to be <= 255 above. @@ -102,7 +102,7 @@ func (m *MessageCertificateRequest) MarshalInto(out []byte) error { } } - return nil + return m.MarshalSize(), nil } // Unmarshal populates the message from encoded data. diff --git a/pkg/protocol/handshake/message_certificate_request_13.go b/pkg/protocol/handshake/message_certificate_request_13.go index 1a015cc73..4abe49bae 100644 --- a/pkg/protocol/handshake/message_certificate_request_13.go +++ b/pkg/protocol/handshake/message_certificate_request_13.go @@ -63,14 +63,14 @@ func (m *MessageCertificateRequest13) Marshal() ([]byte, error) { if !hasSignatureAlgorithms { return nil, errMissingSignatureAlgorithmsExtension } - out := make([]byte, m.Size()) - err := m.MarshalInto(out) + out := make([]byte, m.MarshalSize()) + _, err := m.MarshalTo(out) return out, err } -// Size returns the size needed for MarshalInto. -func (m *MessageCertificateRequest13) Size() int { +// MarshalSize returns the size needed for MarshalTo. +func (m *MessageCertificateRequest13) MarshalSize() int { cache, err := m.innerMarshal() if err != nil { return 0 @@ -109,11 +109,11 @@ func (m *MessageCertificateRequest13) innerMarshal() ([]byte, error) { return m.marshalCache, m.marshalCacheErr } -// MarshalInto encodes like Marshal but in a pre-allocate buffer. -func (m *MessageCertificateRequest13) MarshalInto(out []byte) error { +// MarshalTo encodes like Marshal but in a pre-allocate buffer. +func (m *MessageCertificateRequest13) MarshalTo(out []byte) (int, error) { // Validate certificate_request_context length if len(m.CertificateRequestContext) > certReq13ContextMaxLength { - return errCertificateRequestContextTooLong + return 0, errCertificateRequestContextTooLong } // Validate that signature_algorithms extension is present (required by RFC 8446) @@ -126,21 +126,21 @@ func (m *MessageCertificateRequest13) MarshalInto(out []byte) error { } } if !hasSignatureAlgorithms { - return errMissingSignatureAlgorithmsExtension + return 0, errMissingSignatureAlgorithmsExtension } - if len(out) < m.Size() { - return errBufferTooSmall + if len(out) < m.MarshalSize() { + return 0, errBufferTooSmall } cache, err := m.innerMarshal() if err != nil { - return err + return 0, err } copy(out, cache) - return nil + return m.MarshalSize(), nil } // Unmarshal decodes the MessageCertificateRequest13 from its wire format. diff --git a/pkg/protocol/handshake/message_certificate_verify.go b/pkg/protocol/handshake/message_certificate_verify.go index 5207c9b06..6e35ee13f 100644 --- a/pkg/protocol/handshake/message_certificate_verify.go +++ b/pkg/protocol/handshake/message_certificate_verify.go @@ -29,8 +29,8 @@ func (m MessageCertificateVerify) Type() Type { return TypeCertificateVerify } -// Size returns the minimal required size for Marshalnto. -func (m MessageCertificateVerify) Size() int { +// MarshalSize returns the minimal required size for Marshalnto. +func (m MessageCertificateVerify) MarshalSize() int { return 1 + 1 + 2 + len(m.Signature) } @@ -47,23 +47,23 @@ func (m *MessageCertificateVerify) Marshal() ([]byte, error) { return nil, errInvalidSignHashAlgorithm } - out := make([]byte, m.Size()) - err := m.MarshalInto(out) + out := make([]byte, m.MarshalSize()) + _, err := m.MarshalTo(out) return out, err } -// MarshalInto encodes the Handshake into a pre-allocated buffer. -func (m *MessageCertificateVerify) MarshalInto(out []byte) error { +// MarshalTo encodes the Handshake into a pre-allocated buffer. +func (m *MessageCertificateVerify) MarshalTo(out []byte) (int, error) { if m.HashAlgorithm > 0xFF || m.SignatureAlgorithm > 0xFF { - return errInvalidSignHashAlgorithm + return 0, errInvalidSignHashAlgorithm } // CertificateVerify in DTLS 1.2 encodes hash/signature as 1 byte each. scheme := tls.SignatureScheme(uint16(m.HashAlgorithm)<<8 | uint16(m.SignatureAlgorithm)) var alg signaturehash.Algorithm if err := alg.Unmarshal(scheme); err != nil { - return errInvalidSignHashAlgorithm + return 0, errInvalidSignHashAlgorithm } out[0] = byte(m.HashAlgorithm) @@ -71,7 +71,7 @@ func (m *MessageCertificateVerify) MarshalInto(out []byte) error { binary.BigEndian.PutUint16(out[2:], uint16(len(m.Signature))) //nolint:gosec // G115 copy(out[4:], m.Signature) - return nil + return m.MarshalSize(), nil } // Unmarshal populates the message from encoded data. diff --git a/pkg/protocol/handshake/message_client_hello.go b/pkg/protocol/handshake/message_client_hello.go index 78e45a192..9519549d1 100644 --- a/pkg/protocol/handshake/message_client_hello.go +++ b/pkg/protocol/handshake/message_client_hello.go @@ -46,8 +46,8 @@ func (m *MessageClientHello) cacheMarshalExtensions() error { return m.marchalledExtensionsErr } -// Size returns the size needed for MarshalInto. -func (m *MessageClientHello) Size() int { +// MarshalSize returns the size needed for MarshalTo. +func (m *MessageClientHello) MarshalSize() int { encodedCipherSuiteIDs := encodeCipherSuiteIDs(m.CipherSuiteIDs) encodedCompressionMethods := protocol.EncodeCompressionMethods(m.CompressionMethods) @@ -68,31 +68,31 @@ func (m *MessageClientHello) Size() int { // Marshal encodes the Handshake. func (m *MessageClientHello) Marshal() ([]byte, error) { - out := make([]byte, m.Size()) - err := m.MarshalInto(out) + out := make([]byte, m.MarshalSize()) + _, err := m.MarshalTo(out) return out, err } -// MarshalInto encodes the Handshake into a pre-allocated buffer. -func (m *MessageClientHello) MarshalInto(out []byte) error { +// MarshalTo encodes the Handshake into a pre-allocated buffer. +func (m *MessageClientHello) MarshalTo(out []byte) (int, error) { if len(m.Cookie) > 255 { - return errCookieTooLong + return 0, errCookieTooLong } if len(m.SessionID) > 255 { - return errSessionIDTooLong + return 0, errSessionIDTooLong } if len(m.CompressionMethods) > 255 { - return errCompressionMethodsTooLong + return 0, errCompressionMethodsTooLong } err := m.cacheMarshalExtensions() if err != nil { - return err + return 0, err } - if len(out) < m.Size() { - return errBufferTooSmall + if len(out) < m.MarshalSize() { + return 0, errBufferTooSmall } encodedCipherSuiteIDs := encodeCipherSuiteIDs(m.CipherSuiteIDs) @@ -126,7 +126,7 @@ func (m *MessageClientHello) MarshalInto(out []byte) error { copy(out[offset:], m.marchalledExtensions) - return nil + return m.MarshalSize(), nil } // Unmarshal populates the message from encoded data. diff --git a/pkg/protocol/handshake/message_client_key_exchange.go b/pkg/protocol/handshake/message_client_key_exchange.go index df15cafb9..9879aa0bf 100644 --- a/pkg/protocol/handshake/message_client_key_exchange.go +++ b/pkg/protocol/handshake/message_client_key_exchange.go @@ -41,14 +41,14 @@ func (m *MessageClientKeyExchange) Marshal() ([]byte, error) { } } - out := make([]byte, m.Size()) - err := m.MarshalInto(out) + out := make([]byte, m.MarshalSize()) + _, err := m.MarshalTo(out) return out, err } -// Size returns the size required for MarshalInto. -func (m *MessageClientKeyExchange) Size() int { +// MarshalSize returns the size required for MarshalTo. +func (m *MessageClientKeyExchange) MarshalSize() int { total := 0 if m.IdentityHint != nil { total += 2 @@ -63,14 +63,14 @@ func (m *MessageClientKeyExchange) Size() int { return total } -// MarshalInto encodes the Handshake into a pre-allocated buffer. -func (m *MessageClientKeyExchange) MarshalInto(out []byte) error { +// MarshalTo encodes the Handshake into a pre-allocated buffer. +func (m *MessageClientKeyExchange) MarshalTo(out []byte) (int, error) { if m.IdentityHint == nil && m.PublicKey == nil { - return errInvalidClientKeyExchange + return 0, errInvalidClientKeyExchange } - if len(out) < m.Size() { - return errBufferTooSmall + if len(out) < m.MarshalSize() { + return 0, errBufferTooSmall } offset := 0 @@ -83,14 +83,14 @@ func (m *MessageClientKeyExchange) MarshalInto(out []byte) error { if m.PublicKey != nil { if len(m.PublicKey) > 255 { - return errPublicKeyTooLong + return 0, errPublicKeyTooLong } out[offset] = byte(len(m.PublicKey)) //nolint:gosec // G115: public key length is validated to be <= 255 above. offset += 1 copy(out[offset:], m.PublicKey) } - return nil + return m.MarshalSize(), nil } // Unmarshal populates the message from encoded data. diff --git a/pkg/protocol/handshake/message_finished.go b/pkg/protocol/handshake/message_finished.go index 4662a9250..37c7de52f 100644 --- a/pkg/protocol/handshake/message_finished.go +++ b/pkg/protocol/handshake/message_finished.go @@ -18,27 +18,27 @@ func (m MessageFinished) Type() Type { return TypeFinished } -// Size returns the size required for MarshalInto. -func (m *MessageFinished) Size() int { +// MarshalSize returns the size required for MarshalTo. +func (m *MessageFinished) MarshalSize() int { return len(m.VerifyData) } // Marshal encodes the Handshake. func (m *MessageFinished) Marshal() ([]byte, error) { - out := make([]byte, m.Size()) - err := m.MarshalInto(out) + out := make([]byte, m.MarshalSize()) + _, err := m.MarshalTo(out) return out, err } -// MarshalInto encodes the Handshake into a pre-allocated buffer. -func (m *MessageFinished) MarshalInto(out []byte) error { - if len(out) < m.Size() { - return errBufferTooSmall +// MarshalTo encodes the Handshake into a pre-allocated buffer. +func (m *MessageFinished) MarshalTo(out []byte) (int, error) { + if len(out) < m.MarshalSize() { + return 0, errBufferTooSmall } copy(out, m.VerifyData) - return nil + return m.MarshalSize(), nil } // Unmarshal populates the message from encoded data. diff --git a/pkg/protocol/handshake/message_hello_verify_request.go b/pkg/protocol/handshake/message_hello_verify_request.go index 41123fc21..e157f2f4e 100644 --- a/pkg/protocol/handshake/message_hello_verify_request.go +++ b/pkg/protocol/handshake/message_hello_verify_request.go @@ -37,25 +37,25 @@ func (m *MessageHelloVerifyRequest) Marshal() ([]byte, error) { if len(m.Cookie) > 255 { return nil, errCookieTooLong } - out := make([]byte, m.Size()) - err := m.MarshalInto(out) + out := make([]byte, m.MarshalSize()) + _, err := m.MarshalTo(out) return out, err } -// Size returns the size required for MarshalInto. -func (m *MessageHelloVerifyRequest) Size() int { +// MarshalSize returns the size required for MarshalTo. +func (m *MessageHelloVerifyRequest) MarshalSize() int { return 3 + len(m.Cookie) } -// MarshalInto encodes the Handshake into a pre-allocated buffer. -func (m *MessageHelloVerifyRequest) MarshalInto(out []byte) error { +// MarshalTo encodes the Handshake into a pre-allocated buffer. +func (m *MessageHelloVerifyRequest) MarshalTo(out []byte) (int, error) { if len(m.Cookie) > 255 { - return errCookieTooLong + return 0, errCookieTooLong } - if len(out) < m.Size() { - return errBufferTooSmall + if len(out) < m.MarshalSize() { + return 0, errBufferTooSmall } out[0] = m.Version.Major @@ -63,7 +63,7 @@ func (m *MessageHelloVerifyRequest) MarshalInto(out []byte) error { out[2] = byte(len(m.Cookie)) //nolint:gosec // G115: cookie length is validated to be <= 255 above. copy(out[3:], m.Cookie) - return nil + return m.MarshalSize(), nil } // Unmarshal populates the message from encoded data. diff --git a/pkg/protocol/handshake/message_server_hello.go b/pkg/protocol/handshake/message_server_hello.go index 0a10fb63e..70328bc9d 100644 --- a/pkg/protocol/handshake/message_server_hello.go +++ b/pkg/protocol/handshake/message_server_hello.go @@ -44,8 +44,8 @@ func (m *MessageServerHello) cacheMarshalExtensions() error { return m.marchalledExtensionsErr } -// Size returns the size required by MarshalInto. -func (m *MessageServerHello) Size() int { +// MarshalSize returns the size required by MarshalTo. +func (m *MessageServerHello) MarshalSize() int { err := m.cacheMarshalExtensions() if err != nil { return 0 @@ -58,15 +58,15 @@ func (m *MessageServerHello) Size() int { return total } -// MarshalInto encodes the Handshake into a pre-allocated buffer. -func (m *MessageServerHello) MarshalInto(out []byte) error { +// MarshalTo encodes the Handshake into a pre-allocated buffer. +func (m *MessageServerHello) MarshalTo(out []byte) (int, error) { err := m.cacheMarshalExtensions() if err != nil { - return err + return 0, err } - if len(out) < m.Size() { - return errBufferTooSmall + if len(out) < m.MarshalSize() { + return 0, errBufferTooSmall } offset := 0 @@ -91,7 +91,7 @@ func (m *MessageServerHello) MarshalInto(out []byte) error { copy(out[offset:], m.marchalledExtensions) - return nil + return m.MarshalSize(), nil } // Marshal encodes the Handshake. @@ -105,8 +105,8 @@ func (m *MessageServerHello) Marshal() ([]byte, error) { return nil, errSessionIDTooLong } - out := make([]byte, m.Size()) - err := m.MarshalInto(out) + out := make([]byte, m.MarshalSize()) + _, err := m.MarshalTo(out) return out, err } diff --git a/pkg/protocol/handshake/message_server_hello_done.go b/pkg/protocol/handshake/message_server_hello_done.go index be9319962..5450fbc3b 100644 --- a/pkg/protocol/handshake/message_server_hello_done.go +++ b/pkg/protocol/handshake/message_server_hello_done.go @@ -16,19 +16,19 @@ func (m MessageServerHelloDone) Type() Type { // Marshal encodes the Handshake. func (m *MessageServerHelloDone) Marshal() ([]byte, error) { out := []byte{} - err := m.MarshalInto(out) + _, err := m.MarshalTo(out) return out, err } -// Size returns the size for MarshalInto. -func (m *MessageServerHelloDone) Size() int { +// MarshalSize returns the size for MarshalTo. +func (m *MessageServerHelloDone) MarshalSize() int { return 0 } -// MarshalInto encodes the Handshake. -func (m *MessageServerHelloDone) MarshalInto(out []byte) error { - return nil +// MarshalTo encodes the Handshake. +func (m *MessageServerHelloDone) MarshalTo(out []byte) (int, error) { + return 0, nil } // Unmarshal populates the message from encoded data. diff --git a/pkg/protocol/handshake/message_server_key_exchange.go b/pkg/protocol/handshake/message_server_key_exchange.go index 411c833e3..9e3d45bca 100644 --- a/pkg/protocol/handshake/message_server_key_exchange.go +++ b/pkg/protocol/handshake/message_server_key_exchange.go @@ -34,8 +34,8 @@ func (m MessageServerKeyExchange) Type() Type { return TypeServerKeyExchange } -// Size returns the size required for MarshalInto. -func (m *MessageServerKeyExchange) Size() int { //nolint:cyclop +// MarshalSize returns the size required for MarshalTo. +func (m *MessageServerKeyExchange) MarshalSize() int { //nolint:cyclop total := 0 if m.IdentityHint != nil { total += 2 + len(m.IdentityHint) @@ -67,10 +67,10 @@ func (m *MessageServerKeyExchange) Size() int { //nolint:cyclop return total } -// MarshalInto encodes the Handshake into a pre-allocated buffer. -func (m *MessageServerKeyExchange) MarshalInto(out []byte) error { //nolint:cyclop - if len(out) < m.Size() { - return errBufferTooSmall +// MarshalTo encodes the Handshake into a pre-allocated buffer. +func (m *MessageServerKeyExchange) MarshalTo(out []byte) (int, error) { //nolint:cyclop + if len(out) < m.MarshalSize() { + return 0, errBufferTooSmall } offset := 0 @@ -82,7 +82,7 @@ func (m *MessageServerKeyExchange) MarshalInto(out []byte) error { //nolint:cycl } if m.EllipticCurveType == 0 || len(m.PublicKey) == 0 { - return nil + return 0, nil } out[offset] = byte(m.EllipticCurveType) offset += 1 @@ -96,13 +96,13 @@ func (m *MessageServerKeyExchange) MarshalInto(out []byte) error { //nolint:cycl offset += n switch { case m.HashAlgorithm != hash.None && len(m.Signature) == 0: - return errInvalidSignHashAlgorithm + return 0, errInvalidSignHashAlgorithm case m.HashAlgorithm == hash.None && len(m.Signature) > 0: - return errInvalidSignHashAlgorithm + return 0, errInvalidSignHashAlgorithm case m.SignatureAlgorithm == signature.Anonymous && (m.HashAlgorithm != hash.None || len(m.Signature) > 0): - return errInvalidSignHashAlgorithm + return 0, errInvalidSignHashAlgorithm case m.SignatureAlgorithm == signature.Anonymous: - return nil + return 0, nil } alg := signaturehash.Algorithm{Hash: m.HashAlgorithm, Signature: m.SignatureAlgorithm} @@ -113,13 +113,13 @@ func (m *MessageServerKeyExchange) MarshalInto(out []byte) error { //nolint:cycl offset += 2 copy(out[offset:], m.Signature) - return nil + return m.MarshalSize(), nil } // Marshal encodes the Handshake. func (m *MessageServerKeyExchange) Marshal() ([]byte, error) { - out := make([]byte, m.Size()) - err := m.MarshalInto(out) + out := make([]byte, m.MarshalSize()) + _, err := m.MarshalTo(out) return out, err } diff --git a/pkg/protocol/recordlayer/header.go b/pkg/protocol/recordlayer/header.go index e87d37748..df4a19081 100644 --- a/pkg/protocol/recordlayer/header.go +++ b/pkg/protocol/recordlayer/header.go @@ -38,21 +38,21 @@ func (h *Header) Marshal() ([]byte, error) { hs := FixedHeaderSize + len(h.ConnectionID) out := make([]byte, hs) - err := h.MarshalInto(out) + _, err := h.MarshalTo(out) return out, err } -// MarshalInto encodes a TLS RecordLayer Header to binary using pre-allocated buffer. -func (h *Header) MarshalInto(out []byte) error { +// MarshalTo encodes a TLS RecordLayer Header to binary using pre-allocated buffer. +func (h *Header) MarshalTo(out []byte) (int, error) { if h.SequenceNumber > MaxSequenceNumber { - return errSequenceNumberOverflow + return 0, errSequenceNumberOverflow } hs := FixedHeaderSize + len(h.ConnectionID) if len(out) < hs { - return errBufferTooSmall + return 0, errBufferTooSmall } out[0] = byte(h.ContentType) @@ -63,7 +63,7 @@ func (h *Header) MarshalInto(out []byte) error { copy(out[11:11+len(h.ConnectionID)], h.ConnectionID) binary.BigEndian.PutUint16(out[hs-2:], h.ContentLen) - return nil + return h.MarshalSize(), nil } // Unmarshal populates a TLS RecordLayer Header from binary. @@ -96,7 +96,7 @@ func (h *Header) Unmarshal(data []byte) error { return nil } -// Size returns the total size of the header. -func (h *Header) Size() int { +// MarshalSize returns the total size of the header. +func (h *Header) MarshalSize() int { return FixedHeaderSize + len(h.ConnectionID) } diff --git a/pkg/protocol/recordlayer/recordlayer.go b/pkg/protocol/recordlayer/recordlayer.go index 1bcdb96d4..2afdf7cd3 100644 --- a/pkg/protocol/recordlayer/recordlayer.go +++ b/pkg/protocol/recordlayer/recordlayer.go @@ -50,22 +50,33 @@ type RecordLayer struct { // Marshal encodes the RecordLayer to binary. func (r *RecordLayer) Marshal() ([]byte, error) { - out := make([]byte, r.Content.Size()+r.Header.Size()) + out := make([]byte, r.MarshalSize()) - r.Header.ContentLen = uint16(r.Content.Size()) //nolint:gosec // G115 + _, err := r.MarshalTo(out) + + return out, err +} + +func (r *RecordLayer) MarshalSize() int { + return r.Content.MarshalSize() + r.Header.MarshalSize() +} + +// MarshalTo encodes the RecordLayer to binary. +func (r *RecordLayer) MarshalTo(out []byte) (int, error) { + r.Header.ContentLen = uint16(r.Content.MarshalSize()) //nolint:gosec // G115 r.Header.ContentType = r.Content.ContentType() - err := r.Header.MarshalInto(out) + _, err := r.Header.MarshalTo(out) if err != nil { - return nil, err + return 0, err } - err = r.Content.MarshalInto(out[r.Header.Size():]) + _, err = r.Content.MarshalTo(out[r.Header.MarshalSize():]) if err != nil { - return nil, err + return 0, err } - return out, nil + return r.MarshalSize(), nil } // Unmarshal populates the RecordLayer from binary. @@ -87,7 +98,7 @@ func (r *RecordLayer) Unmarshal(data []byte) error { return errInvalidContentType } - return r.Content.Unmarshal(data[r.Header.Size()+len(r.Header.ConnectionID):]) + return r.Content.Unmarshal(data[r.Header.MarshalSize()+len(r.Header.ConnectionID):]) } // UnpackDatagram extracts all RecordLayer messages from a single datagram. diff --git a/pkg/protocol/recordlayer/recordlayer_test.go b/pkg/protocol/recordlayer/recordlayer_test.go index 555464c12..0c2285591 100644 --- a/pkg/protocol/recordlayer/recordlayer_test.go +++ b/pkg/protocol/recordlayer/recordlayer_test.go @@ -143,7 +143,7 @@ func FuzzRecordLayer_MarshalUnmarshal_RoundTrip(f *testing.F) { require.Equal(t, recordLayer.Header.Epoch, back.Header.Epoch) require.Equal(t, recordLayer.Header.SequenceNumber, back.Header.SequenceNumber) - bodyLen := len(raw) - back.Header.Size() + bodyLen := len(raw) - back.Header.MarshalSize() appData, ok := back.Content.(*protocol.ApplicationData) require.True(t, ok) require.Equal(t, bodyLen, len(appData.Data))