Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 11 additions & 1 deletion config.go
Original file line number Diff line number Diff line change
Expand Up @@ -184,7 +184,7 @@ type Config struct { //nolint:dupl

// InsecureSkipVerifyHello, if true and when acting as server, allow client to
// skip hello verify phase and receive ServerHello after initial ClientHello.
// This have implication on DoS attack resistance.
// This has implications on DoS attack resistance.
InsecureSkipVerifyHello bool

// ConnectionIDGenerator generates connection identifiers that should be
Expand Down Expand Up @@ -227,6 +227,16 @@ type Config struct { //nolint:dupl
// message is sent from a server. The returned handshake message replaces the original message.
CertificateRequestMessageHook func(handshake.MessageCertificateRequest) handshake.Message

// outboundHandshakePacketInterceptor is an optional callback that can be set to
// intercept outgoing raw handshake packets. It is called with the raw packet bytes
// and a boolean flag specifying if this is the last packet of a flight.
// The interceptor can decide to drop the packet by returning true.
outboundHandshakePacketInterceptor func(packet []byte, end bool) bool
// inboundHandshakePacketNotifier is an optional callback that can be set to
// receive notifications about incoming raw handshake packets. It is called with
// the raw packet bytes after the packet has been processed.
inboundHandshakePacketNotifier func(packet []byte)

// OnConnectionAttempt is fired Whenever a connection attempt is made,
// the server or application can call this callback function.
// The callback function can then implement logic to handle the connection attempt, such as logging the attempt,
Expand Down
114 changes: 102 additions & 12 deletions conn.go
Original file line number Diff line number Diff line change
Expand Up @@ -89,13 +89,19 @@ type Conn struct {

reading chan struct{}
handshakeRecv chan recvHandshakeState
inboundPacketInject chan addrPkt
cancelHandshaker func()
cancelHandshakeReader func()

fsm *handshakeFSM

replayProtectionWindow uint

// Allows intercepting and rerouting outgoing handshake packets.
outboundHandshakePacketInterceptor func(packet []byte, end bool) bool
// Allows getting notified about incoming handshake packets.
inboundHandshakePacketNotifier func(packet []byte)

handshakeConfig *handshakeConfig
}

Expand Down Expand Up @@ -234,12 +240,16 @@ func createConn(

reading: make(chan struct{}, 1),
handshakeRecv: make(chan recvHandshakeState),
inboundPacketInject: make(chan addrPkt),
closed: closer.NewCloser(),
cancelHandshaker: func() {},
cancelHandshakeReader: func() {},

replayProtectionWindow: uint(replayProtectionWindow), //nolint:gosec // G115

outboundHandshakePacketInterceptor: config.outboundHandshakePacketInterceptor,
inboundHandshakePacketNotifier: config.inboundHandshakePacketNotifier,

state: State{
isClient: isClient,
},
Expand Down Expand Up @@ -533,14 +543,15 @@ func (c *Conn) RemoteSRTPMasterKeyIdentifier() ([]byte, bool) {
return c.state.remoteSRTPMasterKeyIdentifier, true
}

func (c *Conn) writePackets(ctx context.Context, pkts []*packet) error {
//nolint:cyclop
func (c *Conn) writeHandshakePackets(ctx context.Context, pkts []*packet) error {
c.lock.Lock()
defer c.lock.Unlock()

var rawPackets [][]byte

for _, pkt := range pkts {
if dtlsHandshake, ok := pkt.record.Content.(*handshake.Handshake); ok {
dtlsHandshake, ok := pkt.record.Content.(*handshake.Handshake)
if ok { // Not true for change cipher spec.
handshakeRaw, err := pkt.record.Marshal()
if err != nil {
return err
Expand Down Expand Up @@ -571,6 +582,39 @@ func (c *Conn) writePackets(ctx context.Context, pkts []*packet) error {
rawPackets = append(rawPackets, rawPacket)
}
}

if len(rawPackets) == 0 {
return nil
}
compactedRawPackets := c.compactRawPackets(rawPackets)

for idx, compactedRawPacket := range compactedRawPackets {
if c.outboundHandshakePacketInterceptor != nil {
if c.outboundHandshakePacketInterceptor(compactedRawPacket, idx == len(compactedRawPackets)-1) {
continue
}
}
if _, err := c.nextConn.WriteToContext(ctx, compactedRawPacket, c.rAddr); err != nil {
return netError(err)
}
}

return nil
}

func (c *Conn) writePackets(ctx context.Context, pkts []*packet) error {
c.lock.Lock()
defer c.lock.Unlock()

var rawPackets [][]byte

for _, pkt := range pkts {
rawPacket, err := c.processPacket(pkt)
if err != nil {
return err
}
rawPackets = append(rawPackets, rawPacket)
}
if len(rawPackets) == 0 {
return nil
}
Expand Down Expand Up @@ -797,20 +841,61 @@ var poolReadBuffer = sync.Pool{ //nolint:gochecknoglobals
},
}

func (c *Conn) readAndBuffer(ctx context.Context) error { //nolint:cyclop
bufptr, ok := poolReadBuffer.Get().(*[]byte)
if !ok {
return errFailedToAccessPoolReadBuffer
func (c *Conn) InjectInboundPacket(p []byte, rAddr net.Addr) {
c.inboundPacketInject <- addrPkt{rAddr, p}
}

func (c *Conn) nextPacket(ctx context.Context) ([]byte, net.Addr, error) {
type readResult struct {
data []byte
rAddr net.Addr
err error
}
defer poolReadBuffer.Put(bufptr)
readCh := make(chan readResult, 1)

go func() {
bufptr, ok := poolReadBuffer.Get().(*[]byte)
if !ok {
readCh <- readResult{err: errFailedToAccessPoolReadBuffer}

b := *bufptr
i, rAddr, err := c.nextConn.ReadFromContext(ctx, b)
return
}
b := *bufptr

i, rAddr, err := c.nextConn.ReadFromContext(ctx, b)
if err != nil {
readCh <- readResult{err: err}
poolReadBuffer.Put(bufptr)

return
}

data := make([]byte, i)
copy(data, b[:i])
poolReadBuffer.Put(bufptr)

readCh <- readResult{
data: data,
rAddr: rAddr,
}
}()
select {
case p := <-c.inboundPacketInject:
return p.data, p.rAddr, nil
case p := <-readCh:
return p.data, p.rAddr, p.err
case <-ctx.Done():
return nil, nil, ctx.Err()
}
}

func (c *Conn) readAndBuffer(ctx context.Context) error { //nolint:cyclop
data, rAddr, err := c.nextPacket(ctx)
if err != nil {
return netError(err)
}

pkts, err := recordlayer.ContentAwareUnpackDatagram(b[:i], len(c.state.getLocalConnectionID()))
pkts, err := recordlayer.ContentAwareUnpackDatagram(data, len(c.state.getLocalConnectionID()))
if err != nil {
return err
}
Expand Down Expand Up @@ -841,14 +926,19 @@ func (c *Conn) readAndBuffer(ctx context.Context) error { //nolint:cyclop
}
}
if hasHandshake {
if c.inboundHandshakePacketNotifier != nil {
// Would it be useful to know this was injected?
// Should this work on a copy?

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

flagging

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I would perfer that we do not implement such logic here, thus the users of this functionality should handle it themselves. I think it's reasonable to assume that it's the same library that would both use inject and the notifier.

c.inboundHandshakePacketNotifier(data)
}
s := recvHandshakeState{
done: make(chan struct{}),
isRetransmit: isRetransmit,
}
select {
case c.handshakeRecv <- s:
// If the other party may retransmit the flight,
// we should respond even if it not a new message.
// we should respond even if it is not a new message.
<-s.done
case <-c.fsm.Done():
}
Expand Down
161 changes: 160 additions & 1 deletion conn_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -168,7 +168,7 @@ func TestSequenceNumberOverflow(t *testing.T) {
atomic.StoreUint64(&ca.state.localSequenceNumber[0], recordlayer.MaxSequenceNumber+1)

// Try to send handshake packet.
werr := ca.writePackets(ctx, []*packet{
werr := ca.writeHandshakePackets(ctx, []*packet{
{
record: &recordlayer.RecordLayer{
Header: recordlayer.Header{
Expand Down Expand Up @@ -3500,3 +3500,162 @@ func TestCloseWithoutHandshake(t *testing.T) {
assert.NoError(t, err)
assert.NoError(t, server.Close())
}

func TestOutboundInterceptor(t *testing.T) {
defer test.CheckRoutines(t)()
defer test.TimeOut(time.Second * 10).Stop()

ca, cb := dpipe.Pipe()
serverCert, err := selfsign.GenerateSelfSigned()
assert.NoError(t, err)

var client *Conn
server, err := Server(dtlsnet.PacketConnFromConn(cb), cb.RemoteAddr(), &Config{
Certificates: []tls.Certificate{serverCert},
outboundHandshakePacketInterceptor: func(packet []byte, end bool) bool {
client.InjectInboundPacket(packet, ca.RemoteAddr())

return true
},
InsecureSkipVerify: true,
InsecureSkipVerifyHello: true,
})
assert.NoError(t, err)

go func() {
_ = server.Handshake()
}()

clientCert, err := selfsign.GenerateSelfSigned()
assert.NoError(t, err)

client, err = Client(dtlsnet.PacketConnFromConn(ca), ca.RemoteAddr(), &Config{
Certificates: []tls.Certificate{clientCert},
outboundHandshakePacketInterceptor: func(packet []byte, end bool) bool {
server.InjectInboundPacket(packet, cb.RemoteAddr())

return true
},
InsecureSkipVerify: true,
InsecureSkipVerifyHello: true,
})
assert.NoError(t, err)

assert.NoError(t, client.Handshake())
assert.NoError(t, server.Close())
assert.NoError(t, client.Close())
}

func TestOutboundInterceptorSmallMtuFlush(t *testing.T) {
defer test.CheckRoutines(t)()
defer test.TimeOut(time.Second * 10).Stop()

ca, cb := dpipe.Pipe()
serverCert, err := selfsign.GenerateSelfSigned()
assert.NoError(t, err)

var client *Conn
serverPackets, serverFlights := 0, 0
server, err := Server(dtlsnet.PacketConnFromConn(cb), cb.RemoteAddr(), &Config{
Certificates: []tls.Certificate{serverCert},
outboundHandshakePacketInterceptor: func(packet []byte, end bool) bool {
serverPackets++
if end {
serverFlights++
}
client.InjectInboundPacket(packet, ca.RemoteAddr())

return true
},
InsecureSkipVerify: true,
InsecureSkipVerifyHello: true,
MTU: 400,
})
assert.NoError(t, err)

go func() {
_ = server.Handshake()
}()

clientCert, err := selfsign.GenerateSelfSigned()
assert.NoError(t, err)

clientPackets, clientFlights := 0, 0
client, err = Client(dtlsnet.PacketConnFromConn(ca), ca.RemoteAddr(), &Config{
Certificates: []tls.Certificate{clientCert},
outboundHandshakePacketInterceptor: func(packet []byte, end bool) bool {
clientPackets++
if end {
clientFlights++
}
server.InjectInboundPacket(packet, cb.RemoteAddr())

return true
},
InsecureSkipVerify: true,
InsecureSkipVerifyHello: true,
MTU: 500,
})
assert.NoError(t, err)

assert.NoError(t, client.Handshake())
assert.NoError(t, server.Close())
assert.NoError(t, client.Close())
assert.Equal(t, 2, clientPackets)
assert.Equal(t, 2, clientFlights)
assert.Equal(t, 4, serverPackets)
assert.Equal(t, 2, serverFlights)
}

func TestInboundNotifier(t *testing.T) {
defer test.CheckRoutines(t)()
defer test.TimeOut(time.Second * 10).Stop()

ca, cb := dpipe.Pipe()
serverCert, err := selfsign.GenerateSelfSigned()
assert.NoError(t, err)

var inboundHandshakePackets [][]byte
server, err := Server(dtlsnet.PacketConnFromConn(cb), cb.RemoteAddr(), &Config{
Certificates: []tls.Certificate{serverCert},
inboundHandshakePacketNotifier: func(packet []byte) {
data := make([]byte, len(packet))
copy(data, packet)
inboundHandshakePackets = append(inboundHandshakePackets, data)
},
InsecureSkipVerify: true,
InsecureSkipVerifyHello: true,
})
assert.NoError(t, err)

go func() {
_ = server.Handshake()
}()

clientCert, err := selfsign.GenerateSelfSigned()
assert.NoError(t, err)

var outboundHandshakePackets [][]byte
client, err := Client(dtlsnet.PacketConnFromConn(ca), ca.RemoteAddr(), &Config{
Certificates: []tls.Certificate{clientCert},
outboundHandshakePacketInterceptor: func(packet []byte, end bool) bool {
data := make([]byte, len(packet))
copy(data, packet)
outboundHandshakePackets = append(outboundHandshakePackets, data)

return false
},
InsecureSkipVerify: true,
InsecureSkipVerifyHello: true,
})
assert.NoError(t, err)

assert.NoError(t, client.Handshake())
assert.NoError(t, server.Close())
assert.NoError(t, client.Close())
assert.Equal(t, len(inboundHandshakePackets), len(outboundHandshakePackets))

for i := range inboundHandshakePackets {
assert.Equal(t, inboundHandshakePackets[i], outboundHandshakePackets[i])
}
}
8 changes: 8 additions & 0 deletions errors.go
Original file line number Diff line number Diff line change
Expand Up @@ -202,6 +202,14 @@ var (
Err: errors.New("client hello message hook option requires a non-nil function"),
}
//nolint:err113
errNilOutboundHandshakePacketInterceptor = &FatalError{
Err: errors.New("outbound handshake packet interceptor option requires a non-nil function"),
}
//nolint:err113
errNilInboundHandshakePacketNotifier = &FatalError{
Err: errors.New("inbound handshake packet notifier option requires a non-nil function"),
}
//nolint:err113
errNilGetCertificate = &FatalError{Err: errors.New("get certificate option requires a non-nil callback")}
//nolint:err113
errNilServerHelloMessageHook = &FatalError{
Expand Down
Loading
Loading