fix: 修复 TLS ConnectionState、SPA 回退、健康检查与监听小问题

This commit is contained in:
Nixevol
2026-09-30 15:08:45 +08:00
parent 67a33ba9e5
commit 9e3769f740
17 changed files with 598 additions and 62 deletions
+10 -1
View File
@@ -1,6 +1,7 @@
package main
import (
"crypto/tls"
"fmt"
"io"
"net"
@@ -25,8 +26,16 @@ func cmdHealthcheck(_ []string) error {
if err != nil {
return err
}
url := "http://" + addr + "/healthz"
scheme := "http"
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}, //nolint:gosec
}
}
url := scheme + "://" + addr + "/healthz"
resp, err := client.Get(url)
if err != nil {
return fmt.Errorf("healthcheck %s: %w", url, err)
+101 -1
View File
@@ -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)
@@ -54,3 +66,91 @@ func TestCmdHealthcheckOK(t *testing.T) {
t.Fatal(err)
}
}
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) }()
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
}
+7 -14
View File
@@ -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
@@ -356,10 +349,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 {
+14
View File
@@ -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)
}
+1 -1
View File
@@ -4,7 +4,7 @@ admin_listen: ""
tls:
cert_file: ""
key_file: ""
allow_plaintext: true
allow_plaintext: false
data_dir: /data
log:
level: info
+36
View File
@@ -392,6 +392,42 @@
- 影响:不完整 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
+7 -4
View File
@@ -6,6 +6,7 @@ import (
"net/http"
"time"
"git.asio.asia/nixevol/NixMsg/internal/httpx"
"git.asio.asia/nixevol/NixMsg/internal/listener"
"github.com/coder/websocket"
)
@@ -31,11 +32,13 @@ 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)
+22 -4
View File
@@ -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 {
+40
View File
@@ -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)
}
})
}
+9 -1
View File
@@ -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
}
+57
View File
@@ -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)
}
+70
View File
@@ -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("<html>app</html>")},
"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())
}
})
}
+43 -3
View File
@@ -2,6 +2,7 @@ package listener
import (
"bufio"
"crypto/tls"
"errors"
"net"
"strconv"
@@ -24,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,返回仍可读完整流的连接。
+157
View File
@@ -2,8 +2,11 @@ package listener
import (
"context"
"crypto/tls"
"io"
"net"
"net/http"
"os"
"testing"
"time"
)
@@ -228,3 +231,157 @@ func TestPreHandshakeLimitDropsNewConn(t *testing.T) {
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")
}
}
+12 -25
View File
@@ -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 判断。
+1 -1
View File
@@ -403,7 +403,7 @@ 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 {
+11 -7
View File
@@ -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
}