fix: 修复监听 accept 退避与握手前超时上限

This commit is contained in:
Nixevol
2026-09-30 16:21:51 +08:00
parent 8ca545ae99
commit c0b2903ab7
7 changed files with 484 additions and 14 deletions
+19
View File
@@ -416,6 +416,25 @@
- 备选方案:放到 `internal/admin`(超出 N 目录)。 - 备选方案:放到 `internal/admin`(超出 N 目录)。
- 影响:A/I 接线时调用这些方法即可。 - 影响:A/I 接线时调用这些方法即可。
### 复审修复 L-02
1. **accept 临时错误退避重试**
- 原条款:无(issue #21);net/http `Serve` 对 Accept 临时错误退避。
- 实际做法:`acceptLoop` 仅在已关闭或 `errors.Is(err, net.ErrClosed)` 时退出;其余错误记日志后从 5ms 倍增到最多 1s 再 Accept。
- 原因:文件描述符短暂耗尽不应永久停收新连接。
- 备选方案:只对 `net.Error.Timeout` 重试(Go 已弃用 Temporary)。
- 影响:进程在瞬时 EMFILE/ENOBUFS 后可自行恢复。
### 复审修复 L-01
1. **握手前超时、CONNECT 预读上限、HTTP IdleTimeout、握手前信号量**
- 原条款:DEVELOPMENT 4.2 仅规定首字节 10s;issue #20。
- 实际做法:TLS `Handshake` 前 `SetDeadline(10s)`;MQTT CONNECT 剩余长度 >64KiB 立即关闭;预读 CONNECT 头超时同样 10s;`OnMQTT` / `AttachWS` 前再设读超时;`http.Server` 设 `IdleTimeout=120s`、`MaxHeaderBytes=64KiB`;accept 到分流完成占用容量 1024 的信号量,满则关新连接。测试用 `HandshakeTimeout`/`IdleTimeout`/`PreHandshakeLimit` 缩短等待。
- 原因:`MaximumClients` 只统计认证后会话,握手前可被慢连接占满。
- 备选方案:给 HTTP 再加 `ReadTimeout`(须在 WS Accept 前清掉,改动面更大)。
- 影响:不完整 TLS/CONNECT 约 10s 内断开;超长 remaining length 的 CONNECT 不再让 mochi 预分配近 768KiB。
- 未做:L-01 验收里「接真实 broker 只发 0x10」在 listener 层用 `OnMQTT` 回调等价覆盖(预读阶段即关闭,不进入 broker)。WebSocket 升级后超时只在 `ws.go` 设读 deadline,未另写 broker 测试(范围限制)。
## 消息 M ## 消息 M
### M1 2026-09-30 ### M1 2026-09-30
+2
View File
@@ -4,6 +4,7 @@ import (
"context" "context"
"net" "net"
"net/http" "net/http"
"time"
"git.asio.asia/nixevol/NixMsg/internal/listener" "git.asio.asia/nixevol/NixMsg/internal/listener"
"github.com/coder/websocket" "github.com/coder/websocket"
@@ -36,6 +37,7 @@ func (b *Broker) WSHandler(proxies *listener.ProxySet) http.Handler {
nc = listener.WithRemoteAddr(nc, &net.TCPAddr{IP: net.ParseIP(ip)}) nc = listener.WithRemoteAddr(nc, &net.TCPAddr{IP: net.ParseIP(ip)})
} }
} }
_ = nc.SetReadDeadline(time.Now().Add(10 * time.Second))
_ = b.AttachWS(nc) _ = b.AttachWS(nc)
}) })
} }
+90
View File
@@ -0,0 +1,90 @@
package listener
import (
"errors"
"log/slog"
"net"
"sync"
"testing"
"time"
)
// scriptedListener 先返回指定错误,再交出连接,用于测 accept 退避重试。
type scriptedListener struct {
mu sync.Mutex
n int
steps []acceptStep
addr net.Addr
closed chan struct{}
}
type acceptStep struct {
c net.Conn
err error
}
func (l *scriptedListener) Accept() (net.Conn, error) {
l.mu.Lock()
if l.n >= len(l.steps) {
l.mu.Unlock()
<-l.closed
return nil, net.ErrClosed
}
st := l.steps[l.n]
l.n++
l.mu.Unlock()
return st.c, st.err
}
func (l *scriptedListener) Close() error {
select {
case <-l.closed:
default:
close(l.closed)
}
return nil
}
func (l *scriptedListener) Addr() net.Addr { return l.addr }
func TestAcceptLoopRetriesThenAccepts(t *testing.T) {
client, server := net.Pipe()
defer func() { _ = client.Close() }()
writeDone := make(chan struct{})
go func() {
_, _ = client.Write([]byte{'G'})
close(writeDone)
}()
ln := &scriptedListener{
addr: &net.TCPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 9},
closed: make(chan struct{}),
steps: []acceptStep{
{err: &net.OpError{Op: "accept", Net: "tcp", Err: errors.New("too many open files")}},
{c: server},
},
}
httpLn := NewChanListener(ln.Addr(), 8)
s := &Server{
log: slog.Default(),
closed: make(chan struct{}),
clientHTTP: httpLn,
opts: Options{AllowPlaintext: true},
}
s.wg.Add(1)
go s.acceptLoop(ln, false)
select {
case got := <-httpLn.ch:
_ = got.Close()
case <-time.After(2 * time.Second):
t.Fatal("expected connection to be classified after temporary accept error")
}
close(s.closed)
_ = ln.Close()
s.wg.Wait()
<-writeDone
}
+30 -1
View File
@@ -2,12 +2,16 @@ package listener
import ( import (
"bufio" "bufio"
"errors"
"net" "net"
"strconv" "strconv"
"time" "time"
) )
const firstByteTimeout = 10 * time.Second const (
firstByteTimeout = 10 * time.Second
maxMQTTConnectRemaining = 64 << 10
)
// bufferedConn 把已读字节放回连接,供后续 TLS/HTTP/MQTT 继续读。 // bufferedConn 把已读字节放回连接,供后续 TLS/HTTP/MQTT 继续读。
type bufferedConn struct { type bufferedConn struct {
@@ -41,6 +45,31 @@ func peekFirstByte(c net.Conn) (net.Conn, byte, error) {
return bc, b, nil return bc, b, nil
} }
// mqttConnectRemaining 偷看 CONNECT 剩余长度(不消费缓冲)。超过 4 字节编码则失败。
func mqttConnectRemaining(c net.Conn) (int, error) {
bc := wrapBuffered(c)
value := 0
multiplier := 1
for i := 0; i < 4; i++ {
buf, err := bc.r.Peek(2 + i)
if err != nil {
return 0, err
}
encoded := int(buf[1+i])
value += (encoded & 127) * multiplier
if encoded&128 == 0 {
return value, nil
}
if i == 3 {
break
}
multiplier *= 128
}
return 0, errMQTTRemainOverflow
}
var errMQTTRemainOverflow = errors.New("mqtt remaining length overflow")
// addrConn 只改 RemoteAddr,用于受信任代理后的真实 IP。 // addrConn 只改 RemoteAddr,用于受信任代理后的真实 IP。
type addrConn struct { type addrConn struct {
net.Conn net.Conn
+230
View File
@@ -0,0 +1,230 @@
package listener
import (
"context"
"net"
"net/http"
"testing"
"time"
)
func startPlain(t *testing.T, opts Options) *Server {
t.Helper()
if opts.Listen == "" {
opts.Listen = "127.0.0.1:0"
}
if opts.DataDir == "" {
opts.DataDir = t.TempDir()
}
if opts.ClientHandler == nil {
opts.ClientHandler = NewMux(RoleShared, Handlers{
Healthz: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
_, _ = w.Write([]byte("ok"))
}),
})
}
opts.AllowPlaintext = true
s, err := New(opts)
if err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithCancel(context.Background())
t.Cleanup(cancel)
if err := s.Start(ctx); err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = s.Close() })
return s
}
func TestTLSHandshakeTimeout(t *testing.T) {
dir := t.TempDir()
certPath, keyPath := writeTestCert(t, dir, "hs")
mux := NewMux(RoleShared, Handlers{
Healthz: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
_, _ = w.Write([]byte("ok"))
}),
})
s, err := New(Options{
Listen: "127.0.0.1:0",
DataDir: dir,
CertFile: certPath,
KeyFile: keyPath,
AllowPlaintext: false,
ClientHandler: mux,
HandshakeTimeout: 300 * time.Millisecond,
})
if err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
if err := s.Start(ctx); err != nil {
t.Fatal(err)
}
defer func() { _ = s.Close() }()
c, err := net.DialTimeout("tcp", s.ListenAddr(), time.Second)
if err != nil {
t.Fatal(err)
}
defer func() { _ = c.Close() }()
if _, err := c.Write([]byte{0x16}); err != nil {
t.Fatal(err)
}
start := time.Now()
_ = c.SetReadDeadline(time.Now().Add(2 * time.Second))
buf := make([]byte, 1)
_, err = c.Read(buf)
if err == nil {
t.Fatal("expected TLS handshake timeout to close connection")
}
if d := time.Since(start); d > time.Second {
t.Fatalf("closed after %s, want around handshake timeout", d)
}
}
func TestMQTTConnectHeaderTimeout(t *testing.T) {
seen := make(chan struct{}, 1)
s := startPlain(t, Options{
HandshakeTimeout: 300 * time.Millisecond,
OnMQTT: func(c net.Conn) {
seen <- struct{}{}
_ = c.Close()
},
})
c, err := net.DialTimeout("tcp", s.ListenAddr(), time.Second)
if err != nil {
t.Fatal(err)
}
defer func() { _ = c.Close() }()
if _, err := c.Write([]byte{0x10}); err != nil {
t.Fatal(err)
}
start := time.Now()
_ = c.SetReadDeadline(time.Now().Add(2 * time.Second))
buf := make([]byte, 1)
_, err = c.Read(buf)
if err == nil {
t.Fatal("expected incomplete CONNECT to be closed")
}
if d := time.Since(start); d > time.Second {
t.Fatalf("closed after %s", d)
}
select {
case <-seen:
t.Fatal("OnMQTT should not run for incomplete CONNECT header")
default:
}
}
func TestMQTTConnectRemainingTooLarge(t *testing.T) {
seen := make(chan struct{}, 1)
s := startPlain(t, Options{
OnMQTT: func(c net.Conn) {
seen <- struct{}{}
_ = c.Close()
},
})
c, err := net.DialTimeout("tcp", s.ListenAddr(), time.Second)
if err != nil {
t.Fatal(err)
}
defer func() { _ = c.Close() }()
if _, err := c.Write([]byte{0x10, 0xFF, 0xFF, 0x2F}); err != nil {
t.Fatal(err)
}
_ = c.SetReadDeadline(time.Now().Add(2 * time.Second))
buf := make([]byte, 1)
_, err = c.Read(buf)
if err == nil {
t.Fatal("expected oversized CONNECT remaining length to close")
}
select {
case <-seen:
t.Fatal("OnMQTT should not run for oversized remaining length")
case <-time.After(50 * time.Millisecond):
}
}
func TestHTTPIdleTimeoutCloses(t *testing.T) {
s := startPlain(t, Options{IdleTimeout: 200 * time.Millisecond})
c, err := net.DialTimeout("tcp", s.ListenAddr(), time.Second)
if err != nil {
t.Fatal(err)
}
defer func() { _ = c.Close() }()
if _, err := c.Write([]byte("GET /healthz HTTP/1.1\r\nHost: x\r\n\r\n")); err != nil {
t.Fatal(err)
}
_ = c.SetReadDeadline(time.Now().Add(2 * time.Second))
buf := make([]byte, 4096)
n, err := c.Read(buf)
if err != nil || n == 0 {
t.Fatalf("first response n=%d err=%v", n, err)
}
_ = c.SetReadDeadline(time.Now().Add(2 * time.Second))
_, err = c.Read(buf)
if err == nil {
t.Fatal("expected idle timeout to close keep-alive connection")
}
}
func TestPreHandshakeLimitDropsNewConn(t *testing.T) {
dir := t.TempDir()
certPath, keyPath := writeTestCert(t, dir, "lim")
mux := NewMux(RoleShared, Handlers{})
s, err := New(Options{
Listen: "127.0.0.1:0",
DataDir: dir,
CertFile: certPath,
KeyFile: keyPath,
AllowPlaintext: false,
ClientHandler: mux,
HandshakeTimeout: 2 * time.Second,
PreHandshakeLimit: 1,
})
if err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
if err := s.Start(ctx); err != nil {
t.Fatal(err)
}
defer func() { _ = s.Close() }()
c1, err := net.DialTimeout("tcp", s.ListenAddr(), time.Second)
if err != nil {
t.Fatal(err)
}
defer func() { _ = c1.Close() }()
if _, err := c1.Write([]byte{0x16}); err != nil {
t.Fatal(err)
}
deadline := time.Now().Add(time.Second)
for time.Now().Before(deadline) {
if len(s.hsSem) > 0 {
break
}
time.Sleep(5 * time.Millisecond)
}
if len(s.hsSem) == 0 {
t.Fatal("handshake slot not acquired")
}
c2, err := net.DialTimeout("tcp", s.ListenAddr(), time.Second)
if err != nil {
t.Fatal(err)
}
defer func() { _ = c2.Close() }()
_ = c2.SetReadDeadline(time.Now().Add(400 * time.Millisecond))
buf := make([]byte, 1)
_, err = c2.Read(buf)
if ne, ok := err.(net.Error); ok && ne.Timeout() {
t.Fatal("second connection still open while handshake slots are full")
}
if err == nil {
t.Fatal("unexpected data on dropped connection")
}
}
+1 -1
View File
@@ -138,7 +138,7 @@ func TestIdentifyTLSHTTPAndMQTT(t *testing.T) {
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
_, _ = raw.Write([]byte{0x10}) _, _ = raw.Write([]byte{0x10, 0x00})
select { select {
case mc := <-gotMQTT: case mc := <-gotMQTT:
_ = mc.Close() _ = mc.Close()
+112 -12
View File
@@ -15,6 +15,12 @@ import (
"time" "time"
) )
const (
defaultPreHandshakeLimit = 1024
defaultIdleTimeout = 120 * time.Second
maxHTTPHeaderBytes = 64 << 10
)
// Kind 识别结果。 // Kind 识别结果。
type Kind int type Kind int
@@ -39,6 +45,12 @@ type Options struct {
// OnMQTT 在 listen 上识别到裸 MQTT(含 TLS 后)时调用;应阻塞到连接结束。 // OnMQTT 在 listen 上识别到裸 MQTT(含 TLS 后)时调用;应阻塞到连接结束。
OnMQTT func(conn net.Conn) OnMQTT func(conn net.Conn)
Logger *slog.Logger Logger *slog.Logger
// HandshakeTimeout 首字节之后 TLS 握手 / MQTT CONNECT 预读的时限;零值 10s。
HandshakeTimeout time.Duration
// IdleTimeout HTTP keep-alive 空闲超时;零值 120s。
IdleTimeout time.Duration
// PreHandshakeLimit accept 到分流完成的并发上限;零值 1024。
PreHandshakeLimit int
} }
// Server 一个或两个 TCP 监听上的协议识别与分流。 // Server 一个或两个 TCP 监听上的协议识别与分流。
@@ -64,6 +76,7 @@ type Server struct {
wg sync.WaitGroup wg sync.WaitGroup
closed chan struct{} closed chan struct{}
closeOnce sync.Once closeOnce sync.Once
hsSem chan struct{}
} }
// New 校验选项并准备证书;不开始监听。 // New 校验选项并准备证书;不开始监听。
@@ -82,11 +95,16 @@ func New(opts Options) (*Server, error) {
if err != nil { if err != nil {
return nil, fmt.Errorf("trusted_proxies: %w", err) return nil, fmt.Errorf("trusted_proxies: %w", err)
} }
hsLimit := opts.PreHandshakeLimit
if hsLimit <= 0 {
hsLimit = defaultPreHandshakeLimit
}
s := &Server{ s := &Server{
opts: opts, opts: opts,
log: log, log: log,
proxies: ps, proxies: ps,
closed: make(chan struct{}), closed: make(chan struct{}),
hsSem: make(chan struct{}, hsLimit),
} }
hasCert := opts.CertFile != "" && opts.KeyFile != "" hasCert := opts.CertFile != "" && opts.KeyFile != ""
if hasCert { if hasCert {
@@ -119,10 +137,7 @@ func (s *Server) Start(ctx context.Context) error {
} }
s.clientHTTP = NewChanListener(ln.Addr(), 128) s.clientHTTP = NewChanListener(ln.Addr(), 128)
s.clientSrv = &http.Server{ s.clientSrv = s.newHTTPServer(s.opts.ClientHandler)
Handler: s.opts.ClientHandler,
ReadHeaderTimeout: 10 * time.Second,
}
s.wg.Add(1) s.wg.Add(1)
go func() { go func() {
defer s.wg.Done() defer s.wg.Done()
@@ -153,10 +168,7 @@ func (s *Server) Start(ctx context.Context) error {
adminHandler = http.NotFoundHandler() adminHandler = http.NotFoundHandler()
} }
s.adminHTTP = NewChanListener(aln.Addr(), 64) s.adminHTTP = NewChanListener(aln.Addr(), 64)
s.adminSrv = &http.Server{ s.adminSrv = s.newHTTPServer(adminHandler)
Handler: adminHandler,
ReadHeaderTimeout: 10 * time.Second,
}
s.wg.Add(1) s.wg.Add(1)
go func() { go func() {
defer s.wg.Done() defer s.wg.Done()
@@ -232,15 +244,36 @@ func (s *Server) Close() error {
func (s *Server) acceptLoop(ln net.Listener, isAdmin bool) { func (s *Server) acceptLoop(ln net.Listener, isAdmin bool) {
defer s.wg.Done() defer s.wg.Done()
var delay time.Duration
for { for {
c, err := ln.Accept() c, err := ln.Accept()
if err != nil { if err != nil {
select { if s.shuttingDown() || errors.Is(err, net.ErrClosed) {
case <-s.closed:
return
default:
return return
} }
if delay == 0 {
delay = 5 * time.Millisecond
} else {
delay *= 2
}
if delay > time.Second {
delay = time.Second
}
s.log.Error("accept error, retrying", "err", err, "delay", delay)
timer := time.NewTimer(delay)
select {
case <-s.closed:
timer.Stop()
return
case <-timer.C:
}
continue
}
delay = 0
if !s.acquireHandshake() {
s.log.Debug("pre-handshake connections full, closing")
_ = c.Close()
continue
} }
s.wg.Add(1) s.wg.Add(1)
go func(conn net.Conn) { go func(conn net.Conn) {
@@ -250,12 +283,71 @@ func (s *Server) acceptLoop(ln net.Listener, isAdmin bool) {
} }
} }
func (s *Server) shuttingDown() bool {
select {
case <-s.closed:
return true
default:
return false
}
}
func (s *Server) acquireHandshake() bool {
if s.hsSem == nil {
return true
}
select {
case s.hsSem <- struct{}{}:
return true
default:
return false
}
}
func (s *Server) releaseHandshake() {
if s.hsSem == nil {
return
}
select {
case <-s.hsSem:
default:
}
}
func (s *Server) handshakeTimeout() time.Duration {
if s.opts.HandshakeTimeout > 0 {
return s.opts.HandshakeTimeout
}
return firstByteTimeout
}
func (s *Server) idleTimeout() time.Duration {
if s.opts.IdleTimeout > 0 {
return s.opts.IdleTimeout
}
return defaultIdleTimeout
}
func (s *Server) newHTTPServer(h http.Handler) *http.Server {
return &http.Server{
Handler: h,
ReadHeaderTimeout: 10 * time.Second,
IdleTimeout: s.idleTimeout(),
MaxHeaderBytes: maxHTTPHeaderBytes,
}
}
func (s *Server) handleConn(conn net.Conn, isAdmin bool) { func (s *Server) handleConn(conn net.Conn, isAdmin bool) {
var once sync.Once
release := func() { once.Do(s.releaseHandshake) }
defer release()
kind, out, err := s.classify(conn, isAdmin, false) kind, out, err := s.classify(conn, isAdmin, false)
if err != nil || kind == KindClosed { if err != nil || kind == KindClosed {
_ = conn.Close() _ = conn.Close()
return return
} }
release()
switch kind { switch kind {
case KindHTTP: case KindHTTP:
httpLn := s.clientHTTP httpLn := s.clientHTTP
@@ -269,6 +361,7 @@ func (s *Server) handleConn(conn net.Conn, isAdmin bool) {
httpLn.Enqueue(out) httpLn.Enqueue(out)
case KindMQTT: case KindMQTT:
if s.opts.OnMQTT != nil { if s.opts.OnMQTT != nil {
_ = out.SetReadDeadline(time.Now().Add(s.handshakeTimeout()))
s.opts.OnMQTT(out) s.opts.OnMQTT(out)
} else { } else {
_ = out.Close() _ = out.Close()
@@ -296,9 +389,11 @@ func (s *Server) classify(conn net.Conn, isAdmin, afterTLS bool) (Kind, net.Conn
return KindClosed, nil, errors.New("nested tls") return KindClosed, nil, errors.New("nested tls")
} }
tlsConn := tls.Server(c, s.certs.TLSConfig()) tlsConn := tls.Server(c, s.certs.TLSConfig())
_ = c.SetDeadline(time.Now().Add(s.handshakeTimeout()))
if err := tlsConn.Handshake(); err != nil { if err := tlsConn.Handshake(); err != nil {
return KindClosed, nil, err return KindClosed, nil, err
} }
_ = tlsConn.SetDeadline(time.Time{})
return s.classify(tlsConn, isAdmin, true) return s.classify(tlsConn, isAdmin, true)
} }
@@ -314,6 +409,11 @@ func (s *Server) classify(conn net.Conn, isAdmin, afterTLS bool) (Kind, net.Conn
if isAdmin { if isAdmin {
return KindClosed, nil, errors.New("mqtt not allowed on admin_listen") return KindClosed, nil, errors.New("mqtt not allowed on admin_listen")
} }
_ = c.SetReadDeadline(time.Now().Add(s.handshakeTimeout()))
remain, remErr := mqttConnectRemaining(c)
if remErr != nil || remain > maxMQTTConnectRemaining {
return KindClosed, nil, errors.New("mqtt connect remaining too large")
}
return KindMQTT, c, nil return KindMQTT, c, nil
} }
return KindClosed, nil, errors.New("unknown first byte") return KindClosed, nil, errors.New("unknown first byte")