diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index 9cd009d..cd84643 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -416,6 +416,25 @@ - 备选方案:放到 `internal/admin`(超出 N 目录)。 - 影响: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 ### M1 2026-09-30 diff --git a/internal/broker/ws.go b/internal/broker/ws.go index 7aa827b..1272954 100644 --- a/internal/broker/ws.go +++ b/internal/broker/ws.go @@ -4,6 +4,7 @@ import ( "context" "net" "net/http" + "time" "git.asio.asia/nixevol/NixMsg/internal/listener" "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.SetReadDeadline(time.Now().Add(10 * time.Second)) _ = b.AttachWS(nc) }) } diff --git a/internal/listener/accept_test.go b/internal/listener/accept_test.go new file mode 100644 index 0000000..a549696 --- /dev/null +++ b/internal/listener/accept_test.go @@ -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 +} diff --git a/internal/listener/conn.go b/internal/listener/conn.go index 0f565f2..b2eebfb 100644 --- a/internal/listener/conn.go +++ b/internal/listener/conn.go @@ -2,12 +2,16 @@ package listener import ( "bufio" + "errors" "net" "strconv" "time" ) -const firstByteTimeout = 10 * time.Second +const ( + firstByteTimeout = 10 * time.Second + maxMQTTConnectRemaining = 64 << 10 +) // bufferedConn 把已读字节放回连接,供后续 TLS/HTTP/MQTT 继续读。 type bufferedConn struct { @@ -41,6 +45,31 @@ func peekFirstByte(c net.Conn) (net.Conn, byte, error) { 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。 type addrConn struct { net.Conn diff --git a/internal/listener/handshake_test.go b/internal/listener/handshake_test.go new file mode 100644 index 0000000..be845d0 --- /dev/null +++ b/internal/listener/handshake_test.go @@ -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") + } +} diff --git a/internal/listener/listener_test.go b/internal/listener/listener_test.go index 40938b1..4ab0dc0 100644 --- a/internal/listener/listener_test.go +++ b/internal/listener/listener_test.go @@ -138,7 +138,7 @@ func TestIdentifyTLSHTTPAndMQTT(t *testing.T) { if err != nil { t.Fatal(err) } - _, _ = raw.Write([]byte{0x10}) + _, _ = raw.Write([]byte{0x10, 0x00}) select { case mc := <-gotMQTT: _ = mc.Close() diff --git a/internal/listener/server.go b/internal/listener/server.go index bf92e27..4a3d436 100644 --- a/internal/listener/server.go +++ b/internal/listener/server.go @@ -15,6 +15,12 @@ import ( "time" ) +const ( + defaultPreHandshakeLimit = 1024 + defaultIdleTimeout = 120 * time.Second + maxHTTPHeaderBytes = 64 << 10 +) + // Kind 识别结果。 type Kind int @@ -39,6 +45,12 @@ type Options struct { // OnMQTT 在 listen 上识别到裸 MQTT(含 TLS 后)时调用;应阻塞到连接结束。 OnMQTT func(conn net.Conn) 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 监听上的协议识别与分流。 @@ -64,6 +76,7 @@ type Server struct { wg sync.WaitGroup closed chan struct{} closeOnce sync.Once + hsSem chan struct{} } // New 校验选项并准备证书;不开始监听。 @@ -82,11 +95,16 @@ func New(opts Options) (*Server, error) { if err != nil { return nil, fmt.Errorf("trusted_proxies: %w", err) } + hsLimit := opts.PreHandshakeLimit + if hsLimit <= 0 { + hsLimit = defaultPreHandshakeLimit + } s := &Server{ opts: opts, log: log, proxies: ps, closed: make(chan struct{}), + hsSem: make(chan struct{}, hsLimit), } hasCert := opts.CertFile != "" && opts.KeyFile != "" if hasCert { @@ -119,10 +137,7 @@ func (s *Server) Start(ctx context.Context) error { } s.clientHTTP = NewChanListener(ln.Addr(), 128) - s.clientSrv = &http.Server{ - Handler: s.opts.ClientHandler, - ReadHeaderTimeout: 10 * time.Second, - } + s.clientSrv = s.newHTTPServer(s.opts.ClientHandler) s.wg.Add(1) go func() { defer s.wg.Done() @@ -153,10 +168,7 @@ func (s *Server) Start(ctx context.Context) error { adminHandler = http.NotFoundHandler() } s.adminHTTP = NewChanListener(aln.Addr(), 64) - s.adminSrv = &http.Server{ - Handler: adminHandler, - ReadHeaderTimeout: 10 * time.Second, - } + s.adminSrv = s.newHTTPServer(adminHandler) s.wg.Add(1) go func() { defer s.wg.Done() @@ -232,15 +244,36 @@ func (s *Server) Close() error { func (s *Server) acceptLoop(ln net.Listener, isAdmin bool) { defer s.wg.Done() + var delay time.Duration for { c, err := ln.Accept() if err != nil { - select { - case <-s.closed: - return - default: + if s.shuttingDown() || errors.Is(err, net.ErrClosed) { 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) 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) { + var once sync.Once + release := func() { once.Do(s.releaseHandshake) } + defer release() + kind, out, err := s.classify(conn, isAdmin, false) if err != nil || kind == KindClosed { _ = conn.Close() return } + release() switch kind { case KindHTTP: httpLn := s.clientHTTP @@ -269,6 +361,7 @@ func (s *Server) handleConn(conn net.Conn, isAdmin bool) { httpLn.Enqueue(out) case KindMQTT: if s.opts.OnMQTT != nil { + _ = out.SetReadDeadline(time.Now().Add(s.handshakeTimeout())) s.opts.OnMQTT(out) } else { _ = 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") } tlsConn := tls.Server(c, s.certs.TLSConfig()) + _ = c.SetDeadline(time.Now().Add(s.handshakeTimeout())) if err := tlsConn.Handshake(); err != nil { return KindClosed, nil, err } + _ = tlsConn.SetDeadline(time.Time{}) 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 { 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 KindClosed, nil, errors.New("unknown first byte")