From 5ed3401fee0b7d765335cd5a5d7c45dcd099d50c Mon Sep 17 00:00:00 2001 From: Igor Lazic Date: Mon, 18 May 2026 11:57:06 +0200 Subject: [PATCH] Make SCTPListener.Close race-safe; return ErrListenerClosed from Accept MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit SCTPListener.Close currently calls syscall.Shutdown and syscall.Close on ln.fd without any lock or atomic guard. Concurrent AcceptSCTP and Addr calls read ln.fd directly. This produces three observable bugs: 1. A goroutine blocked inside accept4() returns a bare syscall.EINVAL or syscall.EBADF when Close runs concurrently, instead of a sentinel that callers can detect (cf. net.TCPListener which returns net.ErrClosed since Go 1.16). 2. AcceptSCTP after Close calls accept4 on the now-dead fd integer. If the OS has reassigned that fd to an unrelated resource between Close and Accept, accept4 either fails or, worse, returns a fresh accepted-fd from someone else's socket. 3. AcceptSCTP always returns a non-nil *SCTPConn (wrapping fd=-1 on the error path), violating the (net.Conn, error) API contract. Callers that defensively call methods on the returned conn end up with another EBADF from setsockopt etc. This change mirrors the pattern already used by SCTPConn.Close: - Rename SCTPListener.fd to _fd int32, accessed via atomic load/swap. - Close uses atomic.SwapInt32 so only one caller observes a real fd; subsequent Close returns syscall.EBADF without re-issuing syscall.Close on a possibly-reused fd number. - AcceptSCTP and Addr take an atomic snapshot of _fd and bail to ErrListenerClosed / nil when the listener has been closed. If accept4 was already blocked inside the kernel when Close ran, the post-call atomic re-check normalises the resulting EBADF/EINVAL to ErrListenerClosed so callers can detect graceful shutdown with a plain == comparison instead of inspecting raw errnos. - AcceptSCTP now returns (nil, err) on error, matching the contract of every other net.Listener. ErrListenerClosed is a new exported sentinel (defined with errors.New to stay compatible with the go 1.12 directive — errors.Is and net.ErrClosed are Go 1.13+ / 1.16+). Tests: sctp_listener_close_test.go covers Close-during-Accept, Accept-after-Close, double-Close, concurrent-Close, and Addr-after-Close. The first two tests fail on master with the exact errnos this patch fixes (invalid argument / bad file descriptor); all five pass with this change. Existing tests remain green under -race. --- sctp.go | 20 ++++- sctp_linux.go | 45 ++++++++-- sctp_listener_close_test.go | 165 ++++++++++++++++++++++++++++++++++++ 3 files changed, 219 insertions(+), 11 deletions(-) create mode 100644 sctp_listener_close_test.go diff --git a/sctp.go b/sctp.go index 7b0b68e..59a6bd9 100644 --- a/sctp.go +++ b/sctp.go @@ -18,6 +18,7 @@ package sctp import ( "bytes" "encoding/binary" + "errors" "fmt" "net" "strconv" @@ -29,6 +30,11 @@ import ( "unsafe" ) +// ErrListenerClosed is returned by SCTPListener.Accept and AcceptSCTP when the +// listener has been closed, either before the call or while it was blocked in +// the kernel. Compare with == (errors.Is is not available on Go 1.12). +var ErrListenerClosed = errors.New("sctp: use of closed listener") + const ( SOL_SCTP = 132 @@ -736,13 +742,23 @@ func (c *SCTPConn) SetWriteDeadline(t time.Time) error { } type SCTPListener struct { - fd int + _fd int32 // -1 once Close has run; read/written via sync/atomic m sync.Mutex notificationHandler NotificationHandler } +// fd returns the current listener file descriptor or -1 if the listener has +// been closed. Reads are atomic so callers can detect concurrent Close. +func (ln *SCTPListener) fd() int { + return int(atomic.LoadInt32(&ln._fd)) +} + func (ln *SCTPListener) Addr() net.Addr { - laddr, err := sctpGetAddrs(ln.fd, 0, SCTP_GET_LOCAL_ADDRS) + fd := ln.fd() + if fd < 0 { + return nil + } + laddr, err := sctpGetAddrs(fd, 0, SCTP_GET_LOCAL_ADDRS) if err != nil { return nil } diff --git a/sctp_linux.go b/sctp_linux.go index 778d3cf..e2ee41e 100644 --- a/sctp_linux.go +++ b/sctp_linux.go @@ -248,7 +248,7 @@ func listenSCTPExtConfig(network string, laddr *SCTPAddr, options InitMsg, contr return nil, err } return &SCTPListener{ - fd: sock, + _fd: int32(sock), notificationHandler: notificationHandler, }, nil @@ -271,33 +271,60 @@ func FileListener(file *os.File) (*SCTPListener, error) { } return &SCTPListener{ - fd: int(r1), + _fd: int32(r1), notificationHandler: nil, }, nil } // AcceptSCTP waits for and returns the next SCTP connection to the listener. +// +// If the listener has been closed (either before AcceptSCTP was called or +// concurrently while it was blocked in accept4), AcceptSCTP returns +// (nil, ErrListenerClosed) instead of leaking the bare syscall.EBADF / EINVAL +// returned by the kernel on a torn-down fd. func (ln *SCTPListener) AcceptSCTP() (*SCTPConn, error) { - fd, _, err := syscall.Accept4(ln.fd, 0) - return NewSCTPConn(fd, ln.notificationHandler), err + fd := atomic.LoadInt32(&ln._fd) + if fd < 0 { + return nil, ErrListenerClosed + } + nfd, _, err := syscall.Accept4(int(fd), 0) + if err != nil { + // If Close ran concurrently, surface a clean ErrListenerClosed. + if atomic.LoadInt32(&ln._fd) < 0 { + return nil, ErrListenerClosed + } + return nil, err + } + return NewSCTPConn(nfd, ln.notificationHandler), nil } // Accept waits for and returns the next connection connection to the listener. func (ln *SCTPListener) Accept() (net.Conn, error) { - return ln.AcceptSCTP() + c, err := ln.AcceptSCTP() + if err != nil { + return nil, err + } + return c, nil } +// Close releases the listener fd. The first call returns the result of the +// underlying syscall.Close; subsequent calls return syscall.EBADF without +// touching the fd, so a reused fd number cannot be closed by mistake. func (ln *SCTPListener) Close() error { - syscall.Shutdown(ln.fd, syscall.SHUT_RDWR) - return syscall.Close(ln.fd) + fd := atomic.SwapInt32(&ln._fd, -1) + if fd < 0 { + return syscall.EBADF + } + syscall.Shutdown(int(fd), syscall.SHUT_RDWR) + return syscall.Close(int(fd)) } func (ln *SCTPListener) SyscallConn() (syscall.RawConn, error) { - fd := ln.fd + fd := ln.fd() if fd < 0 { return nil, syscall.EINVAL } - return &rawConn{sockfd: int(fd)}, nil + return &rawConn{sockfd: fd}, nil } // DialSCTP - bind socket to laddr (if given) and connect to raddr diff --git a/sctp_listener_close_test.go b/sctp_listener_close_test.go new file mode 100644 index 0000000..553dafe --- /dev/null +++ b/sctp_listener_close_test.go @@ -0,0 +1,165 @@ +//go:build linux && !386 +// +build linux,!386 + +// Copyright 2024 Wataru Ishida. All rights reserved. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or +// implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package sctp + +import ( + "net" + "sync" + "syscall" + "testing" + "time" +) + +// newLoopbackListener returns a fresh SCTPListener on an ephemeral 127.0.0.1 port. +func newLoopbackListener(t *testing.T) *SCTPListener { + t.Helper() + addr := &SCTPAddr{IPAddrs: []net.IPAddr{{IP: net.IPv4(127, 0, 0, 1)}}} + ln, err := ListenSCTP("sctp", addr) + if err != nil { + t.Fatalf("ListenSCTP failed: %v", err) + } + return ln +} + +// TestListenerCloseUnblocksAccept verifies that closing a listener while +// AcceptSCTP is blocked unblocks Accept and returns ErrListenerClosed (not a +// bare syscall.EBADF / syscall.EINVAL leaked from accept4). +func TestListenerCloseUnblocksAccept(t *testing.T) { + ln := newLoopbackListener(t) + + type acceptResult struct { + conn *SCTPConn + err error + } + resCh := make(chan acceptResult, 1) + go func() { + c, err := ln.AcceptSCTP() + resCh <- acceptResult{c, err} + }() + + // Give the goroutine a moment to block inside accept4. + time.Sleep(50 * time.Millisecond) + + if err := ln.Close(); err != nil { + t.Fatalf("first Close returned %v, want nil", err) + } + + select { + case res := <-resCh: + if res.err != ErrListenerClosed { + t.Errorf("AcceptSCTP after Close returned err=%v, want ErrListenerClosed", res.err) + } + if res.conn != nil { + t.Errorf("AcceptSCTP after Close returned non-nil conn=%+v, want nil", res.conn) + } + case <-time.After(2 * time.Second): + t.Fatal("AcceptSCTP did not return within 2s after Close") + } +} + +// TestAcceptAfterCloseReturnsErrClosed verifies that calling AcceptSCTP on an +// already-closed listener returns ErrListenerClosed with a nil conn, instead +// of invoking accept4 on a stale (possibly reused) fd. +func TestAcceptAfterCloseReturnsErrClosed(t *testing.T) { + ln := newLoopbackListener(t) + if err := ln.Close(); err != nil { + t.Fatalf("Close returned %v, want nil", err) + } + + conn, err := ln.AcceptSCTP() + if err != ErrListenerClosed { + t.Errorf("AcceptSCTP on closed listener err=%v, want ErrListenerClosed", err) + } + if conn != nil { + t.Errorf("AcceptSCTP on closed listener conn=%+v, want nil", conn) + } +} + +// TestDoubleCloseSafe verifies that calling Close twice does not attempt to +// close the underlying fd twice (which could close a different, reused fd). +// The first call returns nil; the second returns syscall.EBADF. +func TestDoubleCloseSafe(t *testing.T) { + ln := newLoopbackListener(t) + + if err := ln.Close(); err != nil { + t.Fatalf("first Close returned %v, want nil", err) + } + + if err := ln.Close(); err != syscall.EBADF { + t.Errorf("second Close returned %v, want syscall.EBADF", err) + } +} + +// TestConcurrentCloseSingleCloser verifies that many concurrent Close calls +// result in exactly one successful close (returns nil) and the rest return +// syscall.EBADF — guarding against the OS reassigning the fd between calls. +func TestConcurrentCloseSingleCloser(t *testing.T) { + ln := newLoopbackListener(t) + + const n = 16 + var ( + wg sync.WaitGroup + mu sync.Mutex + successes int + others []error + ) + wg.Add(n) + for i := 0; i < n; i++ { + go func() { + defer wg.Done() + err := ln.Close() + mu.Lock() + defer mu.Unlock() + if err == nil { + successes++ + } else { + others = append(others, err) + } + }() + } + wg.Wait() + + if successes != 1 { + t.Errorf("got %d successful Close calls, want exactly 1 (others=%v)", successes, others) + } + for _, err := range others { + if err != syscall.EBADF { + t.Errorf("concurrent Close returned %v, want syscall.EBADF", err) + } + } +} + +// TestAddrAfterCloseReturnsNil verifies that Addr() on a closed listener does +// not call sctpGetAddrs on a stale fd number — it should return nil because +// our fd accessor reports -1. +func TestAddrAfterCloseReturnsNil(t *testing.T) { + ln := newLoopbackListener(t) + + if a := ln.Addr(); a == nil { + t.Fatalf("Addr() before Close returned nil, want a real address") + } + + if err := ln.Close(); err != nil { + t.Fatalf("Close returned %v, want nil", err) + } + + if a := ln.Addr(); a != nil { + t.Errorf("Addr() after Close returned %v, want nil", a) + } +}