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
127 changes: 50 additions & 77 deletions netctx/packetconn.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ package netctx

import (
"context"
"errors"
"io"
"net"
"sync"
Expand Down Expand Up @@ -32,18 +33,16 @@ type PacketConn interface {
}

type packetConn struct {
nextConn net.PacketConn
closed chan struct{}
closeOnce sync.Once
readMu sync.Mutex
writeMu sync.Mutex
nextConn net.PacketConn
closed atomic.Bool
readMu sync.Mutex
writeMu sync.Mutex
}

// NewPacketConn creates a new PacketConn wrapping the given net.PacketConn.
func NewPacketConn(pconn net.PacketConn) PacketConn {
p := &packetConn{
nextConn: pconn,
closed: make(chan struct{}),
}

return p
Expand All @@ -58,116 +57,90 @@ func NewPacketConn(pconn net.PacketConn) PacketConn {
// the n > 0 bytes returned before considering the error err.
// Unlike net.PacketConn.ReadFrom(), the provided context is
// used to control timeout.
func (p *packetConn) ReadFromContext(ctx context.Context, b []byte) (int, net.Addr, error) { //nolint:cyclop
func (p *packetConn) ReadFromContext(ctx context.Context, b []byte) (int, net.Addr, error) {
p.readMu.Lock()
defer p.readMu.Unlock()

select {
case <-p.closed:
if p.closed.Load() {
return 0, nil, net.ErrClosed
default:
}
if ctx.Err() != nil {
return 0, nil, ctx.Err()
}

done := make(chan struct{})
var wg sync.WaitGroup
var errSetDeadline atomic.Value
wg.Add(1)
go func() {
defer wg.Done()
select {
case <-ctx.Done():
// context canceled
if err := p.nextConn.SetReadDeadline(veryOld); err != nil {
errSetDeadline.Store(err)

return
}
<-done
if err := p.nextConn.SetReadDeadline(time.Time{}); err != nil {
errSetDeadline.Store(err)
}
case <-done:
if deadline, ok := ctx.Deadline(); ok {
if err := p.nextConn.SetReadDeadline(deadline); err != nil {
return 0, nil, err
}
}

detachDeadline := context.AfterFunc(ctx, func() {
if err := p.nextConn.SetReadDeadline(veryOld); err != nil {
_ = p.nextConn.Close()
}
}()
})

n, raddr, err := p.nextConn.ReadFrom(b)

close(done)
wg.Wait()
if e := ctx.Err(); e != nil && n == 0 {
err = e
}
if err2, ok := errSetDeadline.Load().(error); ok && err == nil && err2 != nil {
err = err2
detachDeadline()

var setDeadlineErr error
if !p.closed.Load() {
setDeadlineErr = p.nextConn.SetReadDeadline(time.Time{})
}

return n, raddr, err
return n, raddr, errors.Join(err, ctx.Err(), setDeadlineErr)
}

// WriteToContext writes a packet with payload p to addr.
// Unlike net.PacketConn.WriteTo(), the provided context
// is used to control timeout.
// On packet-oriented connections, write timeouts are rare.
func (p *packetConn) WriteToContext(ctx context.Context, b []byte, raddr net.Addr) (int, error) { //nolint:cyclop
func (p *packetConn) WriteToContext(ctx context.Context, b []byte, raddr net.Addr) (int, error) {
p.writeMu.Lock()
defer p.writeMu.Unlock()

select {
case <-p.closed:
if p.closed.Load() {
return 0, ErrClosing
default:
}
if ctx.Err() != nil {
return 0, ctx.Err()
}

done := make(chan struct{})
var wg sync.WaitGroup
var errSetDeadline atomic.Value
wg.Add(1)
go func() {
defer wg.Done()
select {
case <-ctx.Done():
// context canceled
if err := p.nextConn.SetWriteDeadline(veryOld); err != nil {
errSetDeadline.Store(err)
if deadline, ok := ctx.Deadline(); ok {
if err := p.nextConn.SetWriteDeadline(deadline); err != nil {
return 0, err
}
}

return
}
<-done
if err := p.nextConn.SetWriteDeadline(time.Time{}); err != nil {
errSetDeadline.Store(err)
detachDeadline := context.AfterFunc(ctx, func() {
if errors.Is(ctx.Err(), context.Canceled) {
if err := p.nextConn.SetWriteDeadline(veryOld); err != nil {
_ = p.nextConn.Close()
}
case <-done:
}
}()
})

n, err := p.nextConn.WriteTo(b, raddr)

close(done)
wg.Wait()
if e := ctx.Err(); e != nil && n == 0 {
err = e
}
if err2, ok := errSetDeadline.Load().(error); ok && err == nil && err2 != nil {
err = err2
detachDeadline()
var setDeadlineErr error
if !p.closed.Load() {
setDeadlineErr = p.nextConn.SetWriteDeadline(time.Time{})
}

return n, err
return n, errors.Join(ctx.Err(), setDeadlineErr, err)
}

// Close closes the connection.
// Any blocked ReadFromContext or WriteToContext operations will be unblocked
// and return errors.
func (p *packetConn) Close() error {
err := p.nextConn.Close()
p.closeOnce.Do(func() {
p.writeMu.Lock()
p.readMu.Lock()
close(p.closed)
p.readMu.Unlock()
p.writeMu.Unlock()
})
if !p.closed.Swap(true) {
return p.nextConn.Close()
}

return err
return nil
}

// LocalAddr returns the local network address, if known.
Expand Down
110 changes: 109 additions & 1 deletion netctx/packetconn_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -99,6 +99,22 @@ func TestReadFromTimeout(t *testing.T) {
assert.Empty(t, n, "Wrong data length")
}

func TestReadFromAlreadyTimeout(t *testing.T) {
ca, _ := pipe()
defer func() {
_ = ca.Close()
}()

ctx, cancel := context.WithTimeout(context.Background(), 0)
defer cancel()

c := NewPacketConn(ca)
b := make([]byte, 100)
n, _, err := c.ReadFromContext(ctx, b)
assert.Error(t, err)
assert.Empty(t, n, "Wrong data length")
}

func TestReadFromCancel(t *testing.T) {
ca, _ := pipe()
defer func() {
Expand Down Expand Up @@ -130,6 +146,49 @@ func TestReadFromClosed(t *testing.T) {
assert.Empty(t, n, "Wrong data length")
}

func TestReadFromClosedAfter(t *testing.T) {
ca, _ := pipe()

c := NewPacketConn(ca)
go func() {
<-time.After(time.Second)
_ = c.Close()
}()

b := make([]byte, 100)
n, _, err := c.ReadFromContext(context.Background(), b)
assert.ErrorIs(t, err, io.ErrClosedPipe)
assert.Empty(t, n, "Wrong data length")
}

var errSetDeadline = errors.New("set read deadline failed")

type badDeadlineConn struct {
net.PacketConn
}

func (c *badDeadlineConn) SetReadDeadline(time.Time) error {
return errSetDeadline
}

func TestReadFromSetDeadlineErr(t *testing.T) {
ca, _ := pipe()

b := NewPacketConn(&badDeadlineConn{
PacketConn: ca,
})

ctx, cancel := context.WithCancel(context.Background())
go func() {
<-time.After(time.Second)
cancel()
}()

packet := make([]byte, 100)
_, _, err := b.ReadFromContext(ctx, packet)
assert.ErrorIs(t, err, context.Canceled)
}

func TestWriteTo(t *testing.T) {
ca, cb := pipe()
defer func() {
Expand Down Expand Up @@ -174,6 +233,22 @@ func TestWriteToTimeout(t *testing.T) {
assert.Empty(t, n, "Wrong data length")
}

func TestWriteToAlreadyTimeout(t *testing.T) {
ca, _ := pipe()
defer func() {
_ = ca.Close()
}()

ctx, cancel := context.WithTimeout(context.Background(), 0)
defer cancel()

c := NewPacketConn(ca)
b := make([]byte, 100)
n, err := c.WriteToContext(ctx, b, nil)
assert.Error(t, err)
assert.Empty(t, n, "Wrong data length")
}

func TestWriteToCancel(t *testing.T) {
ca, _ := pipe()
defer func() {
Expand Down Expand Up @@ -205,6 +280,39 @@ func TestWriteToClosed(t *testing.T) {
assert.Empty(t, n, "Wrong data length")
}

func TestWriteToClosedAfter(t *testing.T) {
ca, _ := pipe()

c := NewPacketConn(ca)
go func() {
<-time.After(time.Second)
_ = c.Close()
}()

b := make([]byte, 100)
n, err := c.WriteToContext(context.Background(), b, nil)
assert.ErrorIs(t, err, io.ErrClosedPipe)
assert.Empty(t, n, "Wrong data length")
}

func TestWriteToSetDeadlineErr(t *testing.T) {
ca, _ := pipe()

b := NewPacketConn(&badDeadlineConn{
PacketConn: ca,
})

ctx, cancel := context.WithCancel(context.Background())
go func() {
<-time.After(time.Second)
cancel()
}()

packet := make([]byte, 100)
_, err := b.WriteToContext(ctx, packet, nil)
assert.ErrorIs(t, err, context.Canceled)
}

type packetConnAddrMock struct{}

func (*packetConnAddrMock) LocalAddr() net.Addr { return stringAddr{"local_net", "local_addr"} }
Expand Down Expand Up @@ -333,7 +441,7 @@ func BenchmarkReadFrom(b *testing.B) {
count := 0
for {
n, _, err := c.ReadFromContext(context.Background(), buf)
if err != nil {
if n == 0 && err != nil {
if !errors.Is(err, io.EOF) {
b.Fatal(err)
}
Expand Down
Loading