merge: listener l01-l07
This commit is contained in:
@@ -26,12 +26,13 @@ func cmdHealthcheck(_ []string) error {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
client := &http.Client{Timeout: 3 * time.Second}
|
|
||||||
scheme := "http"
|
scheme := "http"
|
||||||
if healthcheckUseHTTPS(cfg) {
|
client := &http.Client{Timeout: 3 * time.Second}
|
||||||
|
if strings.TrimSpace(cfg.TLS.CertFile) != "" && !cfg.TLS.AllowPlaintext {
|
||||||
scheme = "https"
|
scheme = "https"
|
||||||
|
// 只连 127.0.0.1 做存活探测,不涉及证书身份校验。
|
||||||
client.Transport = &http.Transport{
|
client.Transport = &http.Transport{
|
||||||
TLSClientConfig: &tls.Config{InsecureSkipVerify: true}, // 本机 HEALTHCHECK,自签证书可接受
|
TLSClientConfig: &tls.Config{InsecureSkipVerify: true}, //nolint:gosec
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
url := scheme + "://" + addr + "/healthz"
|
url := scheme + "://" + addr + "/healthz"
|
||||||
@@ -47,12 +48,6 @@ func cmdHealthcheck(_ []string) error {
|
|||||||
return nil
|
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) {
|
func resolveHealthAddr(dataDir, listen string) (string, error) {
|
||||||
path := filepath.Join(dataDir, "listen.addr")
|
path := filepath.Join(dataDir, "listen.addr")
|
||||||
if b, err := os.ReadFile(path); err == nil {
|
if b, err := os.ReadFile(path); err == nil {
|
||||||
|
|||||||
@@ -1,11 +1,24 @@
|
|||||||
package main
|
package main
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/ecdsa"
|
||||||
|
"crypto/elliptic"
|
||||||
|
"crypto/rand"
|
||||||
|
"crypto/x509"
|
||||||
|
"crypto/x509/pkix"
|
||||||
|
"encoding/pem"
|
||||||
|
"math/big"
|
||||||
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/config"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestResolveHealthAddrFromListenAddrFile(t *testing.T) {
|
func TestResolveHealthAddrFromListenAddrFile(t *testing.T) {
|
||||||
@@ -40,7 +53,6 @@ func TestCmdHealthcheckOK(t *testing.T) {
|
|||||||
t.Cleanup(srv.Close)
|
t.Cleanup(srv.Close)
|
||||||
|
|
||||||
dir := t.TempDir()
|
dir := t.TempDir()
|
||||||
// httptest URL is like http://127.0.0.1:port — write host:port into listen.addr
|
|
||||||
hostPort := srv.Listener.Addr().String()
|
hostPort := srv.Listener.Addr().String()
|
||||||
if err := os.WriteFile(filepath.Join(dir, "listen.addr"), []byte(hostPort+"\n"), 0o644); err != nil {
|
if err := os.WriteFile(filepath.Join(dir, "listen.addr"), []byte(hostPort+"\n"), 0o644); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
@@ -55,33 +67,90 @@ func TestCmdHealthcheckOK(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestCmdHealthcheckHTTPS(t *testing.T) {
|
func TestCmdHealthcheckWithTLS(t *testing.T) {
|
||||||
srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
dataDir := t.TempDir()
|
||||||
w.WriteHeader(http.StatusOK)
|
certPath, keyPath := writeServeTestCert(t, dataDir, "hc")
|
||||||
_, _ = w.Write([]byte("ok"))
|
initAdminForTest(t, dataDir)
|
||||||
}))
|
cfgYAML := "listen: \"127.0.0.1:0\"\ndata_dir: \"" + filepath.ToSlash(dataDir) + "\"\n" +
|
||||||
t.Cleanup(srv.Close)
|
"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()
|
deadline := time.Now().Add(10 * time.Second)
|
||||||
hostPort := srv.Listener.Addr().String()
|
for time.Now().Before(deadline) {
|
||||||
if err := os.WriteFile(filepath.Join(dir, "listen.addr"), []byte(hostPort+"\n"), 0o644); err != nil {
|
b, readErr := os.ReadFile(filepath.Join(dataDir, "listen.addr"))
|
||||||
t.Fatal(err)
|
if readErr == nil && strings.TrimSpace(string(b)) != "" {
|
||||||
|
break
|
||||||
}
|
}
|
||||||
cert := filepath.Join(dir, "cert.pem")
|
time.Sleep(20 * time.Millisecond)
|
||||||
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)
|
|
||||||
}
|
}
|
||||||
t.Setenv("NIXMSG_CONFIG", cfgPath)
|
t.Setenv("NIXMSG_CONFIG", cfgPath)
|
||||||
if err := cmdHealthcheck(nil); err != nil {
|
if err := cmdHealthcheck(nil); err != nil {
|
||||||
t.Fatal(err)
|
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
@@ -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 == ""
|
shared := cfg.AdminListen == ""
|
||||||
var clientHandler, adminHTTP http.Handler
|
var clientHandler, adminHTTP http.Handler
|
||||||
if shared {
|
if shared {
|
||||||
@@ -278,17 +282,6 @@ func runServe(ctx context.Context, cfg config.Config) error {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("listener: %w", err)
|
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 {
|
if err := lnSrv.Start(ctx); err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -355,10 +348,10 @@ func staticFileHandler() http.Handler {
|
|||||||
}
|
}
|
||||||
if f, err := sub.Open("index.html"); err == nil {
|
if f, err := sub.Open("index.html"); err == nil {
|
||||||
_ = f.Close()
|
_ = 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 {
|
func writeListenAddr(dataDir, addr string) error {
|
||||||
|
|||||||
@@ -102,6 +102,20 @@ func TestServeHealthzAndListenAddr(t *testing.T) {
|
|||||||
t.Fatalf("readyz status=%d", ready.StatusCode)
|
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 {
|
if _, err := os.Stat(filepath.Join(dataDir, "nixmsg.db")); err != nil {
|
||||||
t.Fatalf("db missing: %v", err)
|
t.Fatalf("db missing: %v", err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ admin_listen: ""
|
|||||||
tls:
|
tls:
|
||||||
cert_file: ""
|
cert_file: ""
|
||||||
key_file: ""
|
key_file: ""
|
||||||
allow_plaintext: true
|
allow_plaintext: false
|
||||||
data_dir: /data
|
data_dir: /data
|
||||||
log:
|
log:
|
||||||
level: info
|
level: info
|
||||||
|
|||||||
@@ -416,6 +416,61 @@
|
|||||||
- 备选方案:放到 `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 测试(范围限制)。
|
||||||
|
|
||||||
|
### 复审修复 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
|
## 消息 M
|
||||||
|
|
||||||
### M1 2026-09-30
|
### M1 2026-09-30
|
||||||
|
|||||||
@@ -4,7 +4,9 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"net"
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/httpx"
|
||||||
"git.asio.asia/nixevol/NixMsg/internal/listener"
|
"git.asio.asia/nixevol/NixMsg/internal/listener"
|
||||||
"github.com/coder/websocket"
|
"github.com/coder/websocket"
|
||||||
)
|
)
|
||||||
@@ -30,12 +32,15 @@ func (b *Broker) WSHandler(proxies *listener.ProxySet) http.Handler {
|
|||||||
defer cancel()
|
defer cancel()
|
||||||
|
|
||||||
nc := websocket.NetConn(ctx, c, websocket.MessageBinary)
|
nc := websocket.NetConn(ctx, c, websocket.MessageBinary)
|
||||||
|
var nets []*net.IPNet
|
||||||
if proxies != nil {
|
if proxies != nil {
|
||||||
ip := proxies.ClientIP(r)
|
nets = proxies.IPNets()
|
||||||
if ip != "" {
|
|
||||||
nc = listener.WithRemoteAddr(nc, &net.TCPAddr{IP: net.ParseIP(ip)})
|
|
||||||
}
|
}
|
||||||
|
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)
|
_ = b.AttachWS(nc)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -7,7 +7,8 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// ClientIP 按 DEVELOPMENT 4.5:仅当对端在 trusted 网段内时采信
|
// ClientIP 按 DEVELOPMENT 4.5:仅当对端在 trusted 网段内时采信
|
||||||
// X-Forwarded-For(从右往左第一个不在 trusted 内的地址)。
|
// X-Forwarded-For(合并多行,从右往左第一个不在 trusted 内的地址)。
|
||||||
|
// 遇到无法解析的项立即停下,回退到对端地址。
|
||||||
func ClientIP(r *http.Request, trusted []*net.IPNet) string {
|
func ClientIP(r *http.Request, trusted []*net.IPNet) string {
|
||||||
host, _, err := net.SplitHostPort(r.RemoteAddr)
|
host, _, err := net.SplitHostPort(r.RemoteAddr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -20,17 +21,20 @@ func ClientIP(r *http.Request, trusted []*net.IPNet) string {
|
|||||||
if !ipInNets(ip, trusted) {
|
if !ipInNets(ip, trusted) {
|
||||||
return ip.String()
|
return ip.String()
|
||||||
}
|
}
|
||||||
xff := r.Header.Get("X-Forwarded-For")
|
xff := strings.Join(r.Header.Values("X-Forwarded-For"), ",")
|
||||||
if xff == "" {
|
if xff == "" {
|
||||||
return ip.String()
|
return ip.String()
|
||||||
}
|
}
|
||||||
parts := strings.Split(xff, ",")
|
parts := strings.Split(xff, ",")
|
||||||
for i := len(parts) - 1; i >= 0; i-- {
|
for i := len(parts) - 1; i >= 0; i-- {
|
||||||
cand := strings.TrimSpace(parts[i])
|
cand := strings.TrimSpace(parts[i])
|
||||||
parsed := net.ParseIP(cand)
|
if cand == "" {
|
||||||
if parsed == nil {
|
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
parsed := parseForwardedIP(cand)
|
||||||
|
if parsed == nil {
|
||||||
|
return ip.String()
|
||||||
|
}
|
||||||
if !ipInNets(parsed, trusted) {
|
if !ipInNets(parsed, trusted) {
|
||||||
return parsed.String()
|
return parsed.String()
|
||||||
}
|
}
|
||||||
@@ -38,6 +42,20 @@ func ClientIP(r *http.Request, trusted []*net.IPNet) string {
|
|||||||
return ip.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)。
|
// IsHTTPS 判定请求是否视为 HTTPS(直连 TLS 或受信任代理的 X-Forwarded-Proto)。
|
||||||
func IsHTTPS(r *http.Request, trusted []*net.IPNet) bool {
|
func IsHTTPS(r *http.Request, trusted []*net.IPNet) bool {
|
||||||
if r.TLS != nil {
|
if r.TLS != nil {
|
||||||
|
|||||||
@@ -32,3 +32,43 @@ func TestParseCIDRsAndTrustedXFF(t *testing.T) {
|
|||||||
t.Fatal("expected https via proxy")
|
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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,6 +1,8 @@
|
|||||||
package httpx
|
package httpx
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"crypto/sha256"
|
||||||
|
"crypto/subtle"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strings"
|
"strings"
|
||||||
)
|
)
|
||||||
@@ -16,7 +18,7 @@ func MetricsGate(token string, next http.Handler) http.Handler {
|
|||||||
}
|
}
|
||||||
auth := r.Header.Get("Authorization")
|
auth := r.Header.Get("Authorization")
|
||||||
const prefix = "Bearer "
|
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)
|
w.WriteHeader(http.StatusUnauthorized)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -27,3 +29,9 @@ func MetricsGate(token string, next http.Handler) http.Handler {
|
|||||||
next.ServeHTTP(w, r)
|
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
|
||||||
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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())
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -2,12 +2,17 @@ package listener
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"bufio"
|
"bufio"
|
||||||
|
"crypto/tls"
|
||||||
|
"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 {
|
||||||
@@ -20,10 +25,49 @@ func (c *bufferedConn) Read(p []byte) (int, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func wrapBuffered(c net.Conn) *bufferedConn {
|
func wrapBuffered(c net.Conn) *bufferedConn {
|
||||||
if bc, ok := c.(*bufferedConn); ok {
|
switch v := c.(type) {
|
||||||
return bc
|
case *bufferedConn:
|
||||||
}
|
return v
|
||||||
|
case *tlsBufferedConn:
|
||||||
|
return v.bufferedConn
|
||||||
|
default:
|
||||||
return &bufferedConn{Conn: c, r: bufio.NewReader(c)}
|
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
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// peekFirstByte 在超时内读首字节并 Unread,返回仍可读完整流的连接。
|
// peekFirstByte 在超时内读首字节并 Unread,返回仍可读完整流的连接。
|
||||||
@@ -41,6 +85,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
|
||||||
|
|||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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()
|
||||||
|
|||||||
+12
-25
@@ -4,6 +4,8 @@ import (
|
|||||||
"net"
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/httpx"
|
||||||
)
|
)
|
||||||
|
|
||||||
// ProxySet 保存受信任代理地址段。
|
// ProxySet 保存受信任代理地址段。
|
||||||
@@ -52,32 +54,17 @@ func (ps *ProxySet) Contains(ip net.IP) bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
// ClientIP 按 DEVELOPMENT 4.5:来自受信任代理时,取 X-Forwarded-For 从右往左第一个不在段内的 IP。
|
// IPNets 返回受信任网段,供 httpx.ClientIP 使用。
|
||||||
// 非代理来源忽略转发头,返回 RemoteAddr 的 IP。
|
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 {
|
func (ps *ProxySet) ClientIP(r *http.Request) string {
|
||||||
remoteIP := ipFromAddr(r.RemoteAddr)
|
return httpx.ClientIP(r, ps.IPNets())
|
||||||
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()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// IsHTTPS 来自受信任代理时按 X-Forwarded-Proto 判断。
|
// IsHTTPS 来自受信任代理时按 X-Forwarded-Proto 判断。
|
||||||
|
|||||||
+113
-13
@@ -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)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -308,12 +403,17 @@ func (s *Server) classify(conn net.Conn, isAdmin, afterTLS bool) (Kind, net.Conn
|
|||||||
}
|
}
|
||||||
|
|
||||||
if b >= 'A' && b <= 'Z' {
|
if b >= 'A' && b <= 'Z' {
|
||||||
return KindHTTP, c, nil
|
return KindHTTP, asHTTPConn(c, afterTLS), nil
|
||||||
}
|
}
|
||||||
if b == 0x10 {
|
if b == 0x10 {
|
||||||
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")
|
||||||
|
|||||||
@@ -18,6 +18,7 @@ type CertReloader struct {
|
|||||||
cert *tls.Certificate
|
cert *tls.Certificate
|
||||||
certMod time.Time
|
certMod time.Time
|
||||||
keyMod time.Time
|
keyMod time.Time
|
||||||
|
cfg *tls.Config
|
||||||
stop chan struct{}
|
stop chan struct{}
|
||||||
stopOnce sync.Once
|
stopOnce sync.Once
|
||||||
}
|
}
|
||||||
@@ -36,6 +37,11 @@ func NewCertReloader(certFile, keyFile string, log *slog.Logger) (*CertReloader,
|
|||||||
if err := r.reload(true); err != nil {
|
if err := r.reload(true); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
r.cfg = &tls.Config{
|
||||||
|
GetCertificate: r.GetCertificate,
|
||||||
|
MinVersion: tls.VersionTLS12,
|
||||||
|
// 故意不设 NextProtos,以便声明 ALPN mqtt 的客户端仍能握手。
|
||||||
|
}
|
||||||
go r.loop()
|
go r.loop()
|
||||||
return r, nil
|
return r, nil
|
||||||
}
|
}
|
||||||
@@ -59,8 +65,10 @@ func (r *CertReloader) Close() {
|
|||||||
r.stopOnce.Do(func() { close(r.stop) })
|
r.stopOnce.Do(func() { close(r.stop) })
|
||||||
}
|
}
|
||||||
|
|
||||||
|
const certCheckInterval = 2 * time.Minute
|
||||||
|
|
||||||
func (r *CertReloader) loop() {
|
func (r *CertReloader) loop() {
|
||||||
t := time.NewTicker(time.Hour)
|
t := time.NewTicker(certCheckInterval)
|
||||||
defer t.Stop()
|
defer t.Stop()
|
||||||
for {
|
for {
|
||||||
select {
|
select {
|
||||||
@@ -107,11 +115,7 @@ func (r *CertReloader) reload(force bool) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// TLSConfig 构造不设 NextProtos 的服务端 TLS 配置。
|
// TLSConfig 返回进程内复用的 TLS 配置(会话票据才能跨连接恢复)。
|
||||||
func (r *CertReloader) TLSConfig() *tls.Config {
|
func (r *CertReloader) TLSConfig() *tls.Config {
|
||||||
return &tls.Config{
|
return r.cfg
|
||||||
GetCertificate: r.GetCertificate,
|
|
||||||
MinVersion: tls.VersionTLS12,
|
|
||||||
// 故意不设 NextProtos,以便声明 ALPN mqtt 的客户端仍能握手。
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user