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
176 changes: 176 additions & 0 deletions close_native_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,176 @@
//go:build linux && !baremetal && !tinygo.wasm

package net

import (
"errors"
"sync"
"syscall"
"testing"
"time"
)

type closeCountNetdev struct {
nopNetdev
calls int
}

func (d *closeCountNetdev) Close(int) error {
d.calls++
return nil
}

func TestCloseConcurrent(t *testing.T) {
previous := netdev
defer func() { netdev = previous }()
d := &closeCountNetdev{}
netdev = d
c := &TCPConn{fd: 42}
var wg sync.WaitGroup
results := make(chan error, 32)
for i := 0; i < 32; i++ {
wg.Add(1)
go func() { defer wg.Done(); results <- c.Close() }()
}
wg.Wait()
close(results)
success := 0
for err := range results {
if err == nil {
success++
} else if !errors.Is(err, ErrClosed) {
t.Fatal(err)
}
}
if success != 1 || d.calls != 1 {
t.Fatalf("success=%d device calls=%d", success, d.calls)
}
}

func TestCloseDescriptorReuse(t *testing.T) {
for _, kind := range []string{"listener", "tcp", "udp"} {
t.Run(kind, func(t *testing.T) {
fd, err := syscall.Socket(syscall.AF_INET, syscall.SOCK_STREAM, 0)
if err != nil {
t.Fatal(err)
}
var closeSocket func() error
switch kind {
case "listener":
closeSocket = (&listener{fd: fd}).Close
case "tcp":
closeSocket = (&TCPConn{fd: fd}).Close
case "udp":
closeSocket = (&UDPConn{fd: fd}).Close
}
if err := closeSocket(); err != nil {
t.Fatal(err)
}
replacement, err := syscall.Socket(syscall.AF_INET, syscall.SOCK_STREAM, 0)
if err != nil {
t.Fatal(err)
}
if replacement != fd {
if err := syscall.Dup3(replacement, fd, 0); err != nil {
syscall.Close(replacement)
t.Fatal(err)
}
syscall.Close(replacement)
}
defer syscall.Close(fd)
err = closeSocket()
if !errors.Is(err, ErrClosed) {
t.Errorf("second Close = %v, want ErrClosed", err)
}
if _, err := syscall.Getsockname(fd); err != nil {
t.Fatalf("second Close damaged replacement socket: %v", err)
}
})
}
}

func TestCloseBlockedAccept(t *testing.T) {
l, err := Listen("tcp4", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
done := make(chan error, 1)
go func() {
c, err := l.Accept()
if c != nil {
c.Close()
}
done <- err
}()
time.Sleep(25 * time.Millisecond)
if err := l.Close(); err != nil {
t.Fatal(err)
}
select {
case err := <-done:
if err == nil {
t.Fatal("Accept succeeded after close")
}
case <-time.After(time.Second):
t.Fatal("Accept did not stop after close")
}
}

func TestCloseBlockedRead(t *testing.T) {
pair, err := syscall.Socketpair(syscall.AF_UNIX, syscall.SOCK_STREAM, 0)
if err != nil {
t.Fatal(err)
}
defer syscall.Close(pair[1])
// Every socket this netdev creates is SOCK_NONBLOCK; a hand-made fd must
// match, or Read blocks in the kernel where the poller cannot wake it.
if err := syscall.SetNonblock(pair[0], true); err != nil {
t.Fatal(err)
}
c := &TCPConn{fd: pair[0], net: "tcp"}
done := make(chan error, 1)
go func() { _, err := c.Read(make([]byte, 1)); done <- err }()
time.Sleep(25 * time.Millisecond)
if err := c.Close(); err != nil {
t.Fatal(err)
}
select {
case err := <-done:
if err == nil {
t.Fatal("Read succeeded after close")
}
case <-time.After(time.Second):
t.Fatal("Read did not stop after close")
}
}

func TestCloseBlockedWrite(t *testing.T) {
pair, err := syscall.Socketpair(syscall.AF_UNIX, syscall.SOCK_STREAM, 0)
if err != nil {
t.Fatal(err)
}
defer syscall.Close(pair[1])
if err := syscall.SetsockoptInt(pair[0], syscall.SOL_SOCKET, syscall.SO_SNDBUF, 4096); err != nil {
t.Fatal(err)
}
// See TestCloseBlockedRead: the fd must be non-blocking to park on the
// poller rather than in the kernel.
if err := syscall.SetNonblock(pair[0], true); err != nil {
t.Fatal(err)
}
c := &TCPConn{fd: pair[0], net: "tcp"}
done := make(chan error, 1)
go func() { _, err := c.Write(make([]byte, 1<<20)); done <- err }()
time.Sleep(25 * time.Millisecond)
if err := c.Close(); err != nil {
t.Fatal(err)
}
select {
case err := <-done:
if err == nil {
t.Fatal("Write succeeded after close")
}
case <-time.After(time.Second):
t.Fatal("Write did not stop after close")
}
}
16 changes: 16 additions & 0 deletions net.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,9 +9,25 @@ package net
import (
"errors"
"io"
"sync"
"time"
)

type closeGuard struct {
mu sync.Mutex
closed bool
}

func (c *closeGuard) close(fd int) error {
c.mu.Lock()
defer c.mu.Unlock()
if c.closed {
return ErrClosed
}
c.closed = true
return netdev.Close(fd)
}

// Addr represents a network end point address.
//
// The two methods [Addr.Network] and [Addr.String] conventionally return strings
Expand Down
15 changes: 15 additions & 0 deletions netdev.go
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,21 @@ const (
var netdev netdever = &nopNetdev{}

// (useNetdev is go:linkname'd from tinygo/drivers package)
// errPollInterrupted is returned from a netdev's Recv/Send when a concurrent
// deadline change interrupted a blocked operation. The net package retries the
// operation with the fresh deadline; the error never escapes to callers.
var errPollInterrupted = errors.New("net: I/O interrupted by deadline change")

// pollInterrupt asks a netdev that supports it to wake goroutines blocked in
// Recv (write=false) or Send (write=true) on sockfd so they re-evaluate a
// just-changed deadline. Netdevs without that ability ignore deadline changes
// on in-flight I/O, as before.
func pollInterrupt(sockfd int, write bool) {
if p, ok := netdev.(interface{ PollInterrupt(sockfd int, write bool) }); ok {
p.PollInterrupt(sockfd, write)
}
}

func useNetdev(dev netdever) {
netdev = dev
}
Expand Down
94 changes: 72 additions & 22 deletions netdev_native.go
Original file line number Diff line number Diff line change
Expand Up @@ -90,7 +90,9 @@ func (*hostNetdev) Socket(domain, stype, protocol int) (int, error) {
protocol = syscall.IPPROTO_TCP
}

fd, err := syscall.Socket(domain, stype, protocol)
// Non-blocking so the poller (netpoll_native.go) can park goroutines on
// EAGAIN instead of pinning a thread in a blocking syscall.
fd, err := syscall.Socket(domain, stype|syscall.SOCK_NONBLOCK|syscall.SOCK_CLOEXEC, protocol)
if err != nil {
return -1, err
}
Expand Down Expand Up @@ -121,30 +123,70 @@ func (n *hostNetdev) Connect(sockfd int, host string, ip netip.AddrPort) error {
sa := sockaddrFromParts(addr, ip.Port())
for {
err := syscall.Connect(sockfd, sa)
if err == syscall.EINTR {
switch err {
case nil:
return nil
case syscall.EINTR:
continue
case syscall.EINPROGRESS, syscall.EALREADY, syscall.EAGAIN:
// Non-blocking connect in progress: wait for the socket to become
// writable, then read the pending error via SO_ERROR.
if werr := poller.wait(sockfd, true, time.Time{}); werr != nil {
return werr
}
soErr, gerr := syscall.GetsockoptInt(sockfd, syscall.SOL_SOCKET, syscall.SO_ERROR)
if gerr != nil {
return gerr
}
if soErr != 0 {
return syscall.Errno(soErr)
}
return nil
default:
return err
}
return err
}
}

// PollInterrupt wakes any goroutine parked in Recv/Send on sockfd so it
// re-evaluates its deadline. The net package calls this (through an optional
// interface) when a deadline is changed on a connection with I/O in flight.
func (*hostNetdev) PollInterrupt(sockfd int, write bool) {
poller.interrupt(sockfd, write)
}

func (*hostNetdev) Listen(sockfd int, backlog int) error {
return syscall.Listen(sockfd, backlog)
}

func (*hostNetdev) Accept(sockfd int) (int, netip.AddrPort, error) {
nfd, sa, err := syscall.Accept(sockfd)
if err != nil {
return -1, netip.AddrPort{}, err
}
var raddr netip.AddrPort
switch s := sa.(type) {
case *syscall.SockaddrInet4:
raddr = netip.AddrPortFrom(netip.AddrFrom4(s.Addr), uint16(s.Port))
case *syscall.SockaddrInet6:
raddr = netip.AddrPortFrom(netip.AddrFrom16(s.Addr), uint16(s.Port))
for {
// Accept4 with SOCK_NONBLOCK|SOCK_CLOEXEC keeps accepted sockets
// non-blocking too, so their reads/writes also go through the poller.
nfd, sa, err := syscall.Accept4(sockfd, syscall.SOCK_NONBLOCK|syscall.SOCK_CLOEXEC)
switch err {
case nil:
case syscall.EINTR:
continue
case syscall.EAGAIN: // == EWOULDBLOCK on Linux
// No pending connection: park until the listener is readable or the
// listening fd is closed (which unblocks accept for shutdown).
if werr := poller.wait(sockfd, false, time.Time{}); werr != nil {
return -1, netip.AddrPort{}, werr
}
continue
default:
return -1, netip.AddrPort{}, err
}
var raddr netip.AddrPort
switch s := sa.(type) {
case *syscall.SockaddrInet4:
raddr = netip.AddrPortFrom(netip.AddrFrom4(s.Addr), uint16(s.Port))
case *syscall.SockaddrInet6:
raddr = netip.AddrPortFrom(netip.AddrFrom16(s.Addr), uint16(s.Port))
}
return nfd, raddr, nil
}
return nfd, raddr, nil
}

func (*hostNetdev) Send(sockfd int, buf []byte, flags int, deadline time.Time) (int, error) {
Expand All @@ -156,16 +198,18 @@ func (*hostNetdev) Send(sockfd int, buf []byte, flags int, deadline time.Time) (
if expired(deadline) {
return total, timeoutError{}
}
if err := setSockTimeout(sockfd, syscall.SO_SNDTIMEO, deadline); err != nil {
return total, err
}
n, err := syscall.Write(sockfd, buf[total:])
if err != nil {
if err == syscall.EINTR {
continue
}
if err == syscall.EAGAIN || err == syscall.EWOULDBLOCK {
return total, timeoutError{}
// Send buffer full: park until writable, the deadline expires,
// or the fd is closed.
if werr := poller.wait(sockfd, true, deadline); werr != nil {
return total, werr
}
continue
}
return total, err
}
Expand All @@ -181,9 +225,6 @@ func (*hostNetdev) Recv(sockfd int, buf []byte, flags int, deadline time.Time) (
if expired(deadline) {
return 0, timeoutError{}
}
if err := setSockTimeout(sockfd, syscall.SO_RCVTIMEO, deadline); err != nil {
return 0, err
}

for {
n, err := syscall.Read(sockfd, buf)
Expand All @@ -192,7 +233,12 @@ func (*hostNetdev) Recv(sockfd int, buf []byte, flags int, deadline time.Time) (
continue
}
if err == syscall.EAGAIN || err == syscall.EWOULDBLOCK {
return 0, timeoutError{}
// Nothing to read yet: park until readable, the deadline
// expires, or the fd is closed.
if werr := poller.wait(sockfd, false, deadline); werr != nil {
return 0, werr
}
continue
}
return n, err
}
Expand All @@ -208,6 +254,10 @@ func (*hostNetdev) Recv(sockfd int, buf []byte, flags int, deadline time.Time) (
}

func (*hostNetdev) Close(sockfd int) error {
// Wake any goroutines parked on this fd (with errPollClosed) before closing
// it, so a blocked Accept/Recv/Send returns promptly on shutdown instead of
// hanging — which is what lets graceful shutdown and Ctrl+C complete.
poller.close(sockfd)
return syscall.Close(sockfd)
}

Expand Down
Loading