diff --git a/cmd/nixmsg/healthcheck.go b/cmd/nixmsg/healthcheck.go index 72e8da6..ad09f9f 100644 --- a/cmd/nixmsg/healthcheck.go +++ b/cmd/nixmsg/healthcheck.go @@ -26,12 +26,13 @@ func cmdHealthcheck(_ []string) error { if err != nil { return err } - client := &http.Client{Timeout: 3 * time.Second} scheme := "http" - if healthcheckUseHTTPS(cfg) { + client := &http.Client{Timeout: 3 * time.Second} + if strings.TrimSpace(cfg.TLS.CertFile) != "" && !cfg.TLS.AllowPlaintext { scheme = "https" + // 只连 127.0.0.1 做存活探测,不涉及证书身份校验。 client.Transport = &http.Transport{ - TLSClientConfig: &tls.Config{InsecureSkipVerify: true}, // 本机 HEALTHCHECK,自签证书可接受 + TLSClientConfig: &tls.Config{InsecureSkipVerify: true}, //nolint:gosec } } url := scheme + "://" + addr + "/healthz" @@ -47,12 +48,6 @@ func cmdHealthcheck(_ []string) error { return nil } -func healthcheckUseHTTPS(cfg config.Config) bool { - cert := strings.TrimSpace(cfg.TLS.CertFile) - key := strings.TrimSpace(cfg.TLS.KeyFile) - return cert != "" && key != "" && !cfg.TLS.AllowPlaintext -} - func resolveHealthAddr(dataDir, listen string) (string, error) { path := filepath.Join(dataDir, "listen.addr") if b, err := os.ReadFile(path); err == nil { diff --git a/cmd/nixmsg/healthcheck_test.go b/cmd/nixmsg/healthcheck_test.go index b863d5f..8f940ff 100644 --- a/cmd/nixmsg/healthcheck_test.go +++ b/cmd/nixmsg/healthcheck_test.go @@ -1,11 +1,24 @@ package main import ( + "context" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "math/big" + "net" "net/http" "net/http/httptest" "os" "path/filepath" + "strings" "testing" + "time" + + "git.asio.asia/nixevol/NixMsg/internal/config" ) func TestResolveHealthAddrFromListenAddrFile(t *testing.T) { @@ -40,7 +53,6 @@ func TestCmdHealthcheckOK(t *testing.T) { t.Cleanup(srv.Close) dir := t.TempDir() - // httptest URL is like http://127.0.0.1:port — write host:port into listen.addr hostPort := srv.Listener.Addr().String() if err := os.WriteFile(filepath.Join(dir, "listen.addr"), []byte(hostPort+"\n"), 0o644); err != nil { t.Fatal(err) @@ -55,33 +67,90 @@ func TestCmdHealthcheckOK(t *testing.T) { } } -func TestCmdHealthcheckHTTPS(t *testing.T) { - srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte("ok")) - })) - t.Cleanup(srv.Close) +func TestCmdHealthcheckWithTLS(t *testing.T) { + dataDir := t.TempDir() + certPath, keyPath := writeServeTestCert(t, dataDir, "hc") + initAdminForTest(t, dataDir) + cfgYAML := "listen: \"127.0.0.1:0\"\ndata_dir: \"" + filepath.ToSlash(dataDir) + "\"\n" + + "tls:\n cert_file: \"" + filepath.ToSlash(certPath) + "\"\n" + + " key_file: \"" + filepath.ToSlash(keyPath) + "\"\n" + + " allow_plaintext: false\nlog:\n level: error\n" + cfgPath := filepath.Join(dataDir, "config.yaml") + if err := os.WriteFile(cfgPath, []byte(cfgYAML), 0o644); err != nil { + t.Fatal(err) + } + cfg, err := config.Load(cfgPath) + if err != nil { + t.Fatal(err) + } + if err := cfg.Validate(); err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + errCh := make(chan error, 1) + go func() { errCh <- runServe(ctx, cfg) }() - dir := t.TempDir() - hostPort := srv.Listener.Addr().String() - if err := os.WriteFile(filepath.Join(dir, "listen.addr"), []byte(hostPort+"\n"), 0o644); err != nil { - t.Fatal(err) - } - cert := filepath.Join(dir, "cert.pem") - key := filepath.Join(dir, "key.pem") - if err := os.WriteFile(cert, []byte("dummy"), 0o644); err != nil { - t.Fatal(err) - } - if err := os.WriteFile(key, []byte("dummy"), 0o644); err != nil { - t.Fatal(err) - } - cfgPath := filepath.Join(dir, "config.yaml") - body := "listen: \":0\"\ndata_dir: \"" + filepath.ToSlash(dir) + "\"\ntls:\n cert_file: \"" + filepath.ToSlash(cert) + "\"\n key_file: \"" + filepath.ToSlash(key) + "\"\n allow_plaintext: false\n" - if err := os.WriteFile(cfgPath, []byte(body), 0o644); err != nil { - t.Fatal(err) + deadline := time.Now().Add(10 * time.Second) + for time.Now().Before(deadline) { + b, readErr := os.ReadFile(filepath.Join(dataDir, "listen.addr")) + if readErr == nil && strings.TrimSpace(string(b)) != "" { + break + } + time.Sleep(20 * time.Millisecond) } t.Setenv("NIXMSG_CONFIG", cfgPath) if err := cmdHealthcheck(nil); err != nil { t.Fatal(err) } + cancel() + select { + case err := <-errCh: + if err != nil { + t.Fatalf("serve exit: %v", err) + } + case <-time.After(10 * time.Second): + t.Fatal("serve did not stop") + } +} + +func writeServeTestCert(t *testing.T, dir, cn string) (certPath, keyPath string) { + t.Helper() + key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + t.Fatal(err) + } + tmpl := &x509.Certificate{ + SerialNumber: big.NewInt(time.Now().UnixNano()), + Subject: pkix.Name{CommonName: cn}, + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(24 * time.Hour), + KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, + DNSNames: []string{"localhost"}, + IPAddresses: []net.IP{net.ParseIP("127.0.0.1")}, + } + der, err := x509.CreateCertificate(rand.Reader, tmpl, tmpl, &key.PublicKey, key) + if err != nil { + t.Fatal(err) + } + certPath = filepath.Join(dir, cn+".crt") + keyPath = filepath.Join(dir, cn+".key") + certOut, err := os.Create(certPath) + if err != nil { + t.Fatal(err) + } + _ = pem.Encode(certOut, &pem.Block{Type: "CERTIFICATE", Bytes: der}) + _ = certOut.Close() + keyOut, err := os.Create(keyPath) + if err != nil { + t.Fatal(err) + } + b, err := x509.MarshalECPrivateKey(key) + if err != nil { + t.Fatal(err) + } + _ = pem.Encode(keyOut, &pem.Block{Type: "EC PRIVATE KEY", Bytes: b}) + _ = keyOut.Close() + return certPath, keyPath } diff --git a/cmd/nixmsg/serve.go b/cmd/nixmsg/serve.go index 89bbcd5..9cd7725 100644 --- a/cmd/nixmsg/serve.go +++ b/cmd/nixmsg/serve.go @@ -249,7 +249,11 @@ func runServe(ctx context.Context, cfg config.Config) error { } } - handlers := buildHandlers(nil) + proxies, err := listener.ParseTrustedProxies(cfg.TrustedProxies) + if err != nil { + return fmt.Errorf("trusted_proxies: %w", err) + } + handlers := buildHandlers(proxies) shared := cfg.AdminListen == "" var clientHandler, adminHTTP http.Handler if shared { @@ -278,17 +282,6 @@ func runServe(ctx context.Context, cfg config.Config) error { if err != nil { return fmt.Errorf("listener: %w", err) } - handlers = buildHandlers(lnSrv.Proxies()) - if shared { - lnOpts.ClientHandler = listener.NewMux(listener.RoleShared, handlers) - } else { - lnOpts.ClientHandler = listener.NewMux(listener.RoleClient, handlers) - lnOpts.AdminHandler = listener.NewMux(listener.RoleAdmin, handlers) - } - lnSrv, err = listener.New(lnOpts) - if err != nil { - return fmt.Errorf("listener: %w", err) - } if err := lnSrv.Start(ctx); err != nil { return err @@ -355,10 +348,10 @@ func staticFileHandler() http.Handler { } if f, err := sub.Open("index.html"); err == nil { _ = f.Close() - return http.FileServer(http.FS(sub)) + return httpx.SPA(sub) } } - return http.FileServer(http.FS(root)) + return httpx.SPA(root) } func writeListenAddr(dataDir, addr string) error { diff --git a/cmd/nixmsg/serve_test.go b/cmd/nixmsg/serve_test.go index b91134c..009cab1 100644 --- a/cmd/nixmsg/serve_test.go +++ b/cmd/nixmsg/serve_test.go @@ -102,6 +102,20 @@ func TestServeHealthzAndListenAddr(t *testing.T) { t.Fatalf("readyz status=%d", ready.StatusCode) } + spa, err := http.Get("http://" + addr + "/endpoints") + if err != nil { + t.Fatalf("spa: %v", err) + } + defer func() { _ = spa.Body.Close() }() + spaBody, _ := io.ReadAll(spa.Body) + if spa.StatusCode != http.StatusOK { + t.Fatalf("spa status=%d body=%s", spa.StatusCode, spaBody) + } + ct := spa.Header.Get("Content-Type") + if !strings.Contains(ct, "text/html") { + t.Fatalf("spa content-type=%q", ct) + } + if _, err := os.Stat(filepath.Join(dataDir, "nixmsg.db")); err != nil { t.Fatalf("db missing: %v", err) } diff --git a/deploy/config.docker.yaml b/deploy/config.docker.yaml index ef764d1..6a5d602 100644 --- a/deploy/config.docker.yaml +++ b/deploy/config.docker.yaml @@ -4,7 +4,7 @@ admin_listen: "" tls: cert_file: "" key_file: "" - allow_plaintext: true + allow_plaintext: false data_dir: /data log: level: info diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index e1587f8..211c06a 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -416,6 +416,61 @@ - 备选方案:放到 `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 测试(范围限制)。 + +### 复审修复 L-04 + +1. **TLS 包装连接实现 ConnectionState** + - 原条款:DEVELOPMENT 第 8、12 节 HTTPS 时 Cookie 加 Secure;issue #23。 + - 实际做法:HTTP 且已 TLS 时返回 `tlsBufferedConn`,`ConnectionState()` 转发内层 `*tls.Conn`;明文仍用 `bufferedConn`,不加该方法。未改 `admin.Deps.SecureCookies`(serve.go 只允许改 SPA 与 listener 装配;`r.TLS` 恢复后 `IsHTTPS` 已足够)。未加 HSTS。 + - 原因:peek 把 `*tls.Conn` 包进 `bufferedConn` 后 net/http 不填 `r.TLS`。 + - 备选方案:ConnContext 回填(审查 A-01);或只靠 SecureCookies 兜底(本线明确不采用)。 + - 影响:直连 TLS 的管理员 Cookie 带 Secure。 + +### 复审修复 L-05 + +1. **后台静态页 SPA 回退** + - 原条款:DEVELOPMENT 4.3「未知前端路由回 index.html」;issue #24。 + - 实际做法:`httpx.SPA`;`staticFileHandler` 改为调用它。目录与无扩展名路径回 `index.html`(`Cache-Control: no-cache`);有扩展名且不存在返回 404。未加 embeddist 的 Playwright e2e(属网页线;本线用 Go 单测与 serve `/endpoints` 覆盖)。 + - 原因:`http.FileServer` 找不到路径即 404。 + - 备选方案:在 listener mux 里写回退。 + - 影响:刷新 `/endpoints` 等前端路由得到 HTML。 + +### 复审修复 L-06 + +1. **健康检查走 HTTPS** + - 原条款:PRD F21/F22;issue #25。 + - 实际做法:`cert_file` 非空且 `allow_plaintext=false` 时 `healthcheck` 请求 `https://`,`InsecureSkipVerify`(仅连 127.0.0.1 存活探测)。`deploy/config.docker.yaml` 改回 `allow_plaintext: false`。 + - 原因:配证书关明文后明文 `/healthz` 会被 listener 断开。 + - 备选方案:健康检查走独立明文端口。 + - 影响:推荐 TLS 部署下容器 HEALTHCHECK 可通过。issue 要求改 `healthcheck.go` 与 docker 示例,超出 serve.go 两段限制,按 issue 正文执行。 + +### 复审修复 L-07 + +1. **XFF / TLS 配置复用 / metrics 常量时间 / 单次 New** + - 原条款:DEVELOPMENT 4.5、12;issue #26。 + - 实际做法:XFF 合并多行,无法解析则停并回退对端;带端口与 IPv6 方括号可解析。`tls.Config` 只建一次,证书检查 2 分钟。`/metrics` 令牌 SHA-256 后 `ConstantTimeCompare`。serve 先 `ParseTrustedProxies` 再只 `listener.New` 一次。WS 调用 `httpx.ClientIP`。未把失败计入 `LockAdminIP`。 + - 原因:两份 XFF 实现可伪造;每连接新 Config 无法会话恢复;二次 New 泄漏证书重载 goroutine。 + - 备选方案:metrics 失败锁定管理员 IP。 + - 影响:锁定按真实客户端 IP;TLS 重连可 DidResume。 + ## 消息 M ### M1 2026-09-30 diff --git a/internal/broker/ws.go b/internal/broker/ws.go index 7aa827b..1a4df66 100644 --- a/internal/broker/ws.go +++ b/internal/broker/ws.go @@ -4,7 +4,9 @@ import ( "context" "net" "net/http" + "time" + "git.asio.asia/nixevol/NixMsg/internal/httpx" "git.asio.asia/nixevol/NixMsg/internal/listener" "github.com/coder/websocket" ) @@ -30,12 +32,15 @@ func (b *Broker) WSHandler(proxies *listener.ProxySet) http.Handler { defer cancel() nc := websocket.NetConn(ctx, c, websocket.MessageBinary) + var nets []*net.IPNet if proxies != nil { - ip := proxies.ClientIP(r) - if ip != "" { - nc = listener.WithRemoteAddr(nc, &net.TCPAddr{IP: net.ParseIP(ip)}) - } + nets = proxies.IPNets() } + ip := httpx.ClientIP(r, nets) + if parsed := net.ParseIP(ip); parsed != nil { + nc = listener.WithRemoteAddr(nc, &net.TCPAddr{IP: parsed}) + } + _ = nc.SetReadDeadline(time.Now().Add(10 * time.Second)) _ = b.AttachWS(nc) }) } diff --git a/internal/httpx/clientip.go b/internal/httpx/clientip.go index 35426fd..0ad7351 100644 --- a/internal/httpx/clientip.go +++ b/internal/httpx/clientip.go @@ -7,7 +7,8 @@ import ( ) // ClientIP 按 DEVELOPMENT 4.5:仅当对端在 trusted 网段内时采信 -// X-Forwarded-For(从右往左第一个不在 trusted 内的地址)。 +// X-Forwarded-For(合并多行,从右往左第一个不在 trusted 内的地址)。 +// 遇到无法解析的项立即停下,回退到对端地址。 func ClientIP(r *http.Request, trusted []*net.IPNet) string { host, _, err := net.SplitHostPort(r.RemoteAddr) if err != nil { @@ -20,17 +21,20 @@ func ClientIP(r *http.Request, trusted []*net.IPNet) string { if !ipInNets(ip, trusted) { return ip.String() } - xff := r.Header.Get("X-Forwarded-For") + xff := strings.Join(r.Header.Values("X-Forwarded-For"), ",") if xff == "" { return ip.String() } parts := strings.Split(xff, ",") for i := len(parts) - 1; i >= 0; i-- { cand := strings.TrimSpace(parts[i]) - parsed := net.ParseIP(cand) - if parsed == nil { + if cand == "" { continue } + parsed := parseForwardedIP(cand) + if parsed == nil { + return ip.String() + } if !ipInNets(parsed, trusted) { return parsed.String() } @@ -38,6 +42,20 @@ func ClientIP(r *http.Request, trusted []*net.IPNet) string { return ip.String() } +func parseForwardedIP(s string) net.IP { + if ip := net.ParseIP(s); ip != nil { + return ip + } + host, _, err := net.SplitHostPort(s) + if err == nil { + return net.ParseIP(host) + } + if strings.HasPrefix(s, "[") && strings.HasSuffix(s, "]") { + return net.ParseIP(s[1 : len(s)-1]) + } + return nil +} + // IsHTTPS 判定请求是否视为 HTTPS(直连 TLS 或受信任代理的 X-Forwarded-Proto)。 func IsHTTPS(r *http.Request, trusted []*net.IPNet) bool { if r.TLS != nil { diff --git a/internal/httpx/clientip_test.go b/internal/httpx/clientip_test.go index c4a5d35..5e52e50 100644 --- a/internal/httpx/clientip_test.go +++ b/internal/httpx/clientip_test.go @@ -32,3 +32,43 @@ func TestParseCIDRsAndTrustedXFF(t *testing.T) { t.Fatal("expected https via proxy") } } + +func TestClientIPMultiLinePortAndIPv6(t *testing.T) { + trusted := httpx.ParseCIDRs([]string{"10.0.0.0/8", "127.0.0.1/32"}) + + t.Run("two header lines", func(t *testing.T) { + r := httptest.NewRequest(http.MethodGet, "/", nil) + r.RemoteAddr = "10.1.2.3:9" + r.Header["X-Forwarded-For"] = []string{"198.51.100.1", "203.0.113.8, 10.9.9.9"} + if got := httpx.ClientIP(r, trusted); got != "203.0.113.8" { + t.Fatalf("got %q", got) + } + }) + + t.Run("port on rightmost client", func(t *testing.T) { + r := httptest.NewRequest(http.MethodGet, "/", nil) + r.RemoteAddr = "10.1.2.3:9" + r.Header.Set("X-Forwarded-For", "198.51.100.9, 203.0.113.10:1234, 10.9.9.9") + if got := httpx.ClientIP(r, trusted); got != "203.0.113.10" { + t.Fatalf("got %q", got) + } + }) + + t.Run("unparseable stops", func(t *testing.T) { + r := httptest.NewRequest(http.MethodGet, "/", nil) + r.RemoteAddr = "10.1.2.3:9" + r.Header.Set("X-Forwarded-For", "198.51.100.1, not-an-ip, 10.9.9.9") + if got := httpx.ClientIP(r, trusted); got != "10.1.2.3" { + t.Fatalf("got %q", got) + } + }) + + t.Run("ipv6 brackets", func(t *testing.T) { + r := httptest.NewRequest(http.MethodGet, "/", nil) + r.RemoteAddr = "10.1.2.3:9" + r.Header.Set("X-Forwarded-For", "[2001:db8::1], 10.9.9.9") + if got := httpx.ClientIP(r, trusted); got != "2001:db8::1" { + t.Fatalf("got %q", got) + } + }) +} diff --git a/internal/httpx/metrics.go b/internal/httpx/metrics.go index f22e3e3..1e08377 100644 --- a/internal/httpx/metrics.go +++ b/internal/httpx/metrics.go @@ -1,6 +1,8 @@ package httpx import ( + "crypto/sha256" + "crypto/subtle" "net/http" "strings" ) @@ -16,7 +18,7 @@ func MetricsGate(token string, next http.Handler) http.Handler { } auth := r.Header.Get("Authorization") const prefix = "Bearer " - if !strings.HasPrefix(auth, prefix) || auth[len(prefix):] != token { + if !strings.HasPrefix(auth, prefix) || !tokenEqual(auth[len(prefix):], token) { w.WriteHeader(http.StatusUnauthorized) return } @@ -27,3 +29,9 @@ func MetricsGate(token string, next http.Handler) http.Handler { next.ServeHTTP(w, r) }) } + +func tokenEqual(got, want string) bool { + sumGot := sha256.Sum256([]byte(got)) + sumWant := sha256.Sum256([]byte(want)) + return subtle.ConstantTimeCompare(sumGot[:], sumWant[:]) == 1 +} diff --git a/internal/httpx/spa.go b/internal/httpx/spa.go new file mode 100644 index 0000000..c6b1ef5 --- /dev/null +++ b/internal/httpx/spa.go @@ -0,0 +1,57 @@ +package httpx + +import ( + "io" + "io/fs" + "net/http" + "path" + "strings" +) + +// SPA 提供后台静态页:存在的文件原样返回;无扩展名的前端路由回退 index.html; +// 有扩展名但不存在返回 404;访问目录不列清单,同样回退 index.html。 +func SPA(fsys fs.FS) http.Handler { + files := http.FileServer(http.FS(fsys)) + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet && r.Method != http.MethodHead { + w.WriteHeader(http.StatusMethodNotAllowed) + return + } + cleaned := path.Clean("/" + r.URL.Path) + rel := strings.TrimPrefix(cleaned, "/") + if rel == "" { + rel = "." + } + fi, err := fs.Stat(fsys, rel) + if err == nil && !fi.IsDir() { + files.ServeHTTP(w, r) + return + } + if path.Ext(cleaned) != "" { + http.NotFound(w, r) + return + } + serveIndex(w, r, fsys) + }) +} + +func serveIndex(w http.ResponseWriter, r *http.Request, fsys fs.FS) { + f, err := fsys.Open("index.html") + if err != nil { + http.NotFound(w, r) + return + } + defer func() { _ = f.Close() }() + st, err := f.Stat() + if err != nil || st.IsDir() { + http.NotFound(w, r) + return + } + w.Header().Set("Cache-Control", "no-cache") + w.Header().Set("Content-Type", "text/html; charset=utf-8") + w.WriteHeader(http.StatusOK) + if r.Method == http.MethodHead { + return + } + _, _ = io.Copy(w, f) +} diff --git a/internal/httpx/spa_test.go b/internal/httpx/spa_test.go new file mode 100644 index 0000000..14d5a43 --- /dev/null +++ b/internal/httpx/spa_test.go @@ -0,0 +1,70 @@ +package httpx_test + +import ( + "net/http" + "net/http/httptest" + "strings" + "testing" + "testing/fstest" + + "git.asio.asia/nixevol/NixMsg/internal/httpx" +) + +func TestSPAFallbackAndMissingAsset(t *testing.T) { + fsys := fstest.MapFS{ + "index.html": &fstest.MapFile{Data: []byte("app")}, + "assets/app.js": &fstest.MapFile{Data: []byte("console.log(1)")}, + } + h := httpx.SPA(fsys) + + t.Run("frontend route", func(t *testing.T) { + res := httptest.NewRecorder() + h.ServeHTTP(res, httptest.NewRequest(http.MethodGet, "/endpoints", nil)) + if res.Code != http.StatusOK { + t.Fatalf("status=%d", res.Code) + } + ct := res.Header().Get("Content-Type") + if !strings.Contains(ct, "text/html") { + t.Fatalf("content-type=%q", ct) + } + if !strings.Contains(res.Body.String(), "app") { + t.Fatalf("body=%q", res.Body.String()) + } + if res.Header().Get("Cache-Control") != "no-cache" { + t.Fatalf("cache-control=%q", res.Header().Get("Cache-Control")) + } + }) + + t.Run("missing js", func(t *testing.T) { + res := httptest.NewRecorder() + h.ServeHTTP(res, httptest.NewRequest(http.MethodGet, "/assets/nope.js", nil)) + if res.Code != http.StatusNotFound { + t.Fatalf("status=%d", res.Code) + } + }) + + t.Run("directory no listing", func(t *testing.T) { + res := httptest.NewRecorder() + h.ServeHTTP(res, httptest.NewRequest(http.MethodGet, "/assets/", nil)) + if res.Code != http.StatusOK { + t.Fatalf("status=%d", res.Code) + } + if strings.Contains(res.Body.String(), "app.js") { + t.Fatal("directory listing leaked") + } + if !strings.Contains(res.Body.String(), "app") { + t.Fatalf("want spa index, body=%q", res.Body.String()) + } + }) + + t.Run("existing file", func(t *testing.T) { + res := httptest.NewRecorder() + h.ServeHTTP(res, httptest.NewRequest(http.MethodGet, "/assets/app.js", nil)) + if res.Code != http.StatusOK { + t.Fatalf("status=%d", res.Code) + } + if res.Body.String() != "console.log(1)" { + t.Fatalf("body=%q", res.Body.String()) + } + }) +} 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..d14704d 100644 --- a/internal/listener/conn.go +++ b/internal/listener/conn.go @@ -2,12 +2,17 @@ package listener import ( "bufio" + "crypto/tls" + "errors" "net" "strconv" "time" ) -const firstByteTimeout = 10 * time.Second +const ( + firstByteTimeout = 10 * time.Second + maxMQTTConnectRemaining = 64 << 10 +) // bufferedConn 把已读字节放回连接,供后续 TLS/HTTP/MQTT 继续读。 type bufferedConn struct { @@ -20,10 +25,49 @@ func (c *bufferedConn) Read(p []byte) (int, error) { } func wrapBuffered(c net.Conn) *bufferedConn { - if bc, ok := c.(*bufferedConn); ok { - return bc + switch v := c.(type) { + case *bufferedConn: + return v + case *tlsBufferedConn: + return v.bufferedConn + default: + return &bufferedConn{Conn: c, r: bufio.NewReader(c)} + } +} + +// tlsBufferedConn 把已 peek 的 TLS 连接交给 http.Server,并实现 ConnectionState, +// 让 net/http 填 r.TLS。明文连接不得使用此类型。 +type tlsBufferedConn struct { + *bufferedConn + tc *tls.Conn +} + +func (c *tlsBufferedConn) ConnectionState() tls.ConnectionState { + return c.tc.ConnectionState() +} + +func asHTTPConn(c net.Conn, afterTLS bool) net.Conn { + if !afterTLS { + return c + } + bc := wrapBuffered(c) + if tc := tlsConnOf(bc); tc != nil { + return &tlsBufferedConn{bufferedConn: bc, tc: tc} + } + return bc +} + +func tlsConnOf(c net.Conn) *tls.Conn { + switch v := c.(type) { + case *tls.Conn: + return v + case *tlsBufferedConn: + return v.tc + case *bufferedConn: + return tlsConnOf(v.Conn) + default: + return nil } - return &bufferedConn{Conn: c, r: bufio.NewReader(c)} } // peekFirstByte 在超时内读首字节并 Unread,返回仍可读完整流的连接。 @@ -41,6 +85,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..b961a1e --- /dev/null +++ b/internal/listener/handshake_test.go @@ -0,0 +1,387 @@ +package listener + +import ( + "context" + "crypto/tls" + "io" + "net" + "net/http" + "os" + "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") + } +} + +func TestHTTPRequestTLSState(t *testing.T) { + t.Run("tls", func(t *testing.T) { + dir := t.TempDir() + certPath, keyPath := writeTestCert(t, dir, "tls-http") + got := make(chan bool, 1) + mux := NewMux(RoleShared, Handlers{ + Healthz: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + got <- r.TLS != nil + _, _ = w.Write([]byte("ok")) + }), + }) + s, err := New(Options{ + Listen: "127.0.0.1:0", + DataDir: dir, + CertFile: certPath, + KeyFile: keyPath, + AllowPlaintext: false, + ClientHandler: mux, + }) + 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() }() + cl := &http.Client{ + Timeout: 5 * time.Second, + Transport: &http.Transport{TLSClientConfig: &tls.Config{InsecureSkipVerify: true}}, + } + resp, err := cl.Get("https://" + s.ListenAddr() + "/healthz") + if err != nil { + t.Fatal(err) + } + _ = resp.Body.Close() + select { + case ok := <-got: + if !ok { + t.Fatal("r.TLS == nil on direct TLS HTTP") + } + case <-time.After(2 * time.Second): + t.Fatal("handler not called") + } + }) + t.Run("plaintext", func(t *testing.T) { + got := make(chan bool, 1) + s := startPlain(t, Options{ + ClientHandler: NewMux(RoleShared, Handlers{ + Healthz: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + got <- r.TLS != nil + _, _ = w.Write([]byte("ok")) + }), + }), + }) + resp, err := http.Get("http://" + s.ListenAddr() + "/healthz") + if err != nil { + t.Fatal(err) + } + _ = resp.Body.Close() + select { + case ok := <-got: + if ok { + t.Fatal("r.TLS != nil on plaintext HTTP") + } + case <-time.After(2 * time.Second): + t.Fatal("handler not called") + } + }) +} + +func TestTLSSessionResume(t *testing.T) { + dir := t.TempDir() + certPath, keyPath := writeTestCert(t, dir, "resume") + s, err := New(Options{ + Listen: "127.0.0.1:0", + DataDir: dir, + CertFile: certPath, + KeyFile: keyPath, + AllowPlaintext: false, + ClientHandler: NewMux(RoleShared, Handlers{ + Healthz: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write([]byte("ok")) + }), + }), + }) + 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() }() + + cfg := &tls.Config{ + InsecureSkipVerify: true, + ClientSessionCache: tls.NewLRUClientSessionCache(8), + } + tr := &http.Transport{TLSClientConfig: cfg} + client := &http.Client{Transport: tr, Timeout: 5 * time.Second} + resp, err := client.Get("https://" + s.ListenAddr() + "/healthz") + if err != nil { + t.Fatal(err) + } + _, _ = io.Copy(io.Discard, resp.Body) + _ = resp.Body.Close() + + c2, err := tls.Dial("tcp", s.ListenAddr(), cfg) + if err != nil { + t.Fatal(err) + } + defer func() { _ = c2.Close() }() + if !c2.ConnectionState().DidResume { + t.Fatal("second handshake DidResume=false; tls.Config should be reused") + } +} + +func TestCertThenKeyReload(t *testing.T) { + dir := t.TempDir() + certPath, keyPath := writeTestCert(t, dir, "first") + cr, err := NewCertReloader(certPath, keyPath, nil) + if err != nil { + t.Fatal(err) + } + defer cr.Close() + old := cr.Certificate() + + time.Sleep(20 * time.Millisecond) + c2, k2 := writeTestCert(t, dir, "second") + data, _ := os.ReadFile(c2) + if err := os.WriteFile(certPath, data, 0o644); err != nil { + t.Fatal(err) + } + if err := cr.ReloadNow(); err == nil { + t.Fatal("expected reload to fail after replacing cert before key") + } + if cr.Certificate() != old { + t.Fatal("old cert should be kept when key does not match") + } + data, _ = os.ReadFile(k2) + if err := os.WriteFile(keyPath, data, 0o644); err != nil { + t.Fatal(err) + } + if err := cr.ReloadNow(); err != nil { + t.Fatal(err) + } + if cr.Certificate() == old { + t.Fatal("certificate not reloaded after key replaced") + } +} 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/proxy.go b/internal/listener/proxy.go index 1bf9682..a8cb5ba 100644 --- a/internal/listener/proxy.go +++ b/internal/listener/proxy.go @@ -4,6 +4,8 @@ import ( "net" "net/http" "strings" + + "git.asio.asia/nixevol/NixMsg/internal/httpx" ) // ProxySet 保存受信任代理地址段。 @@ -52,32 +54,17 @@ func (ps *ProxySet) Contains(ip net.IP) bool { return false } -// ClientIP 按 DEVELOPMENT 4.5:来自受信任代理时,取 X-Forwarded-For 从右往左第一个不在段内的 IP。 -// 非代理来源忽略转发头,返回 RemoteAddr 的 IP。 +// IPNets 返回受信任网段,供 httpx.ClientIP 使用。 +func (ps *ProxySet) IPNets() []*net.IPNet { + if ps == nil { + return nil + } + return ps.nets +} + +// ClientIP 委托 httpx.ClientIP,保证 WS 与 HTTP 同一套 XFF 规则。 func (ps *ProxySet) ClientIP(r *http.Request) string { - remoteIP := ipFromAddr(r.RemoteAddr) - if ps == nil || remoteIP == nil || !ps.Contains(remoteIP) { - if remoteIP != nil { - return remoteIP.String() - } - return "" - } - xff := r.Header.Get("X-Forwarded-For") - if xff == "" { - return remoteIP.String() - } - parts := strings.Split(xff, ",") - for i := len(parts) - 1; i >= 0; i-- { - ipStr := strings.TrimSpace(parts[i]) - ip := net.ParseIP(ipStr) - if ip == nil { - continue - } - if !ps.Contains(ip) { - return ip.String() - } - } - return remoteIP.String() + return httpx.ClientIP(r, ps.IPNets()) } // IsHTTPS 来自受信任代理时按 X-Forwarded-Proto 判断。 diff --git a/internal/listener/server.go b/internal/listener/server.go index bf92e27..c69dd72 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) } @@ -308,12 +403,17 @@ func (s *Server) classify(conn net.Conn, isAdmin, afterTLS bool) (Kind, net.Conn } if b >= 'A' && b <= 'Z' { - return KindHTTP, c, nil + return KindHTTP, asHTTPConn(c, afterTLS), nil } if b == 0x10 { 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") diff --git a/internal/listener/tls.go b/internal/listener/tls.go index aa112b5..03ba0f6 100644 --- a/internal/listener/tls.go +++ b/internal/listener/tls.go @@ -18,6 +18,7 @@ type CertReloader struct { cert *tls.Certificate certMod time.Time keyMod time.Time + cfg *tls.Config stop chan struct{} stopOnce sync.Once } @@ -36,6 +37,11 @@ func NewCertReloader(certFile, keyFile string, log *slog.Logger) (*CertReloader, if err := r.reload(true); err != nil { return nil, err } + r.cfg = &tls.Config{ + GetCertificate: r.GetCertificate, + MinVersion: tls.VersionTLS12, + // 故意不设 NextProtos,以便声明 ALPN mqtt 的客户端仍能握手。 + } go r.loop() return r, nil } @@ -59,8 +65,10 @@ func (r *CertReloader) Close() { r.stopOnce.Do(func() { close(r.stop) }) } +const certCheckInterval = 2 * time.Minute + func (r *CertReloader) loop() { - t := time.NewTicker(time.Hour) + t := time.NewTicker(certCheckInterval) defer t.Stop() for { select { @@ -107,11 +115,7 @@ func (r *CertReloader) reload(force bool) error { return nil } -// TLSConfig 构造不设 NextProtos 的服务端 TLS 配置。 +// TLSConfig 返回进程内复用的 TLS 配置(会话票据才能跨连接恢复)。 func (r *CertReloader) TLSConfig() *tls.Config { - return &tls.Config{ - GetCertificate: r.GetCertificate, - MinVersion: tls.VersionTLS12, - // 故意不设 NextProtos,以便声明 ALPN mqtt 的客户端仍能握手。 - } + return r.cfg }