diff --git a/internal/listener/chanlistener.go b/internal/listener/chanlistener.go new file mode 100644 index 0000000..5bbf94e --- /dev/null +++ b/internal/listener/chanlistener.go @@ -0,0 +1,62 @@ +package listener + +import ( + "net" + "sync" +) + +// ChanListener 是从通道取连接的 net.Listener,交给同一个 http.Server。 +type ChanListener struct { + addr net.Addr + ch chan net.Conn + closed chan struct{} + once sync.Once +} + +// NewChanListener 创建缓冲通道监听器;addr 仅用于 Addr()。 +func NewChanListener(addr net.Addr, buf int) *ChanListener { + if buf < 1 { + buf = 64 + } + if addr == nil { + addr = &net.TCPAddr{IP: net.IPv4zero, Port: 0} + } + return &ChanListener{ + addr: addr, + ch: make(chan net.Conn, buf), + closed: make(chan struct{}), + } +} + +// Addr 返回构造时给出的地址。 +func (l *ChanListener) Addr() net.Addr { return l.addr } + +// Accept 阻塞直到有连接或关闭。 +func (l *ChanListener) Accept() (net.Conn, error) { + select { + case <-l.closed: + return nil, net.ErrClosed + case c, ok := <-l.ch: + if !ok { + return nil, net.ErrClosed + } + return c, nil + } +} + +// Close 关闭监听器并唤醒 Accept。 +func (l *ChanListener) Close() error { + l.once.Do(func() { + close(l.closed) + }) + return nil +} + +// Enqueue 把识别为 HTTP 的连接交给 http.Server;已关闭时丢弃并关闭连接。 +func (l *ChanListener) Enqueue(c net.Conn) { + select { + case <-l.closed: + _ = c.Close() + case l.ch <- c: + } +} diff --git a/internal/listener/conn.go b/internal/listener/conn.go new file mode 100644 index 0000000..0f565f2 --- /dev/null +++ b/internal/listener/conn.go @@ -0,0 +1,76 @@ +package listener + +import ( + "bufio" + "net" + "strconv" + "time" +) + +const firstByteTimeout = 10 * time.Second + +// bufferedConn 把已读字节放回连接,供后续 TLS/HTTP/MQTT 继续读。 +type bufferedConn struct { + net.Conn + r *bufio.Reader +} + +func (c *bufferedConn) Read(p []byte) (int, error) { + return c.r.Read(p) +} + +func wrapBuffered(c net.Conn) *bufferedConn { + if bc, ok := c.(*bufferedConn); ok { + return bc + } + return &bufferedConn{Conn: c, r: bufio.NewReader(c)} +} + +// peekFirstByte 在超时内读首字节并 Unread,返回仍可读完整流的连接。 +func peekFirstByte(c net.Conn) (net.Conn, byte, error) { + bc := wrapBuffered(c) + _ = bc.SetReadDeadline(time.Now().Add(firstByteTimeout)) + b, err := bc.r.ReadByte() + _ = bc.SetReadDeadline(time.Time{}) + if err != nil { + return nil, 0, err + } + if err := bc.r.UnreadByte(); err != nil { + return nil, 0, err + } + return bc, b, nil +} + +// addrConn 只改 RemoteAddr,用于受信任代理后的真实 IP。 +type addrConn struct { + net.Conn + remote net.Addr +} + +func (c *addrConn) RemoteAddr() net.Addr { + if c.remote != nil { + return c.remote + } + return c.Conn.RemoteAddr() +} + +// WithRemoteAddr 包装连接,使 RemoteAddr 返回指定地址(通常是解析出的客户端 IP)。 +func WithRemoteAddr(c net.Conn, remote net.Addr) net.Conn { + if remote == nil { + return c + } + return &addrConn{Conn: c, remote: remote} +} + +// TCPAddrFromIPPort 把 "ip:port" 或纯 IP 转成 *net.TCPAddr。 +func TCPAddrFromIPPort(ipPort string) *net.TCPAddr { + if ipPort == "" { + return &net.TCPAddr{} + } + host, portStr, err := net.SplitHostPort(ipPort) + if err != nil { + return &net.TCPAddr{IP: net.ParseIP(ipPort)} + } + port, _ := strconv.Atoi(portStr) + return &net.TCPAddr{IP: net.ParseIP(host), Port: port} +} diff --git a/internal/listener/listener_test.go b/internal/listener/listener_test.go new file mode 100644 index 0000000..40938b1 --- /dev/null +++ b/internal/listener/listener_test.go @@ -0,0 +1,398 @@ +package listener + +import ( + "context" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "io" + "math/big" + "net" + "net/http" + "os" + "path/filepath" + "testing" + "time" +) + +func TestIdentifyPlainHTTPAndMQTT(t *testing.T) { + dir := t.TempDir() + gotMQTT := make(chan net.Conn, 1) + mux := NewMux(RoleShared, Handlers{ + Healthz: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte("ok")) + }), + }) + s, err := New(Options{ + Listen: "127.0.0.1:0", + DataDir: dir, + ClientHandler: mux, + AllowPlaintext: true, + OnMQTT: func(c net.Conn) { + gotMQTT <- c + buf := make([]byte, 1) + _, _ = c.Read(buf) + _ = c.Close() + }, + }) + if err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + if startErr := s.Start(ctx); startErr != nil { + t.Fatal(startErr) + } + defer func() { _ = s.Close() }() + + addr := s.ListenAddr() + if addr == "" { + t.Fatal("empty listen addr") + } + b, err := os.ReadFile(filepath.Join(dir, "listen.addr")) + if err != nil { + t.Fatal(err) + } + if string(b) != addr+"\n" { + t.Fatalf("listen.addr=%q want %q", b, addr+"\n") + } + + resp, err := http.Get("http://" + addr + "/healthz") + if err != nil { + t.Fatal(err) + } + body, _ := io.ReadAll(resp.Body) + _ = resp.Body.Close() + if resp.StatusCode != 200 || string(body) != "ok" { + t.Fatalf("healthz: %d %q", resp.StatusCode, body) + } + + c, err := net.DialTimeout("tcp", addr, 2*time.Second) + if err != nil { + t.Fatal(err) + } + _, _ = c.Write([]byte{0x10, 0x00}) + select { + case mc := <-gotMQTT: + _ = mc.Close() + case <-time.After(3 * time.Second): + t.Fatal("mqtt not delivered") + } + _ = c.Close() +} + +func TestIdentifyTLSHTTPAndMQTT(t *testing.T) { + dir := t.TempDir() + certPath, keyPath := writeTestCert(t, dir, "old") + gotMQTT := make(chan net.Conn, 1) + mux := NewMux(RoleShared, Handlers{ + Healthz: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write([]byte("tls-ok")) + }), + }) + s, err := New(Options{ + Listen: "127.0.0.1:0", + DataDir: dir, + CertFile: certPath, + KeyFile: keyPath, + AllowPlaintext: false, + ClientHandler: mux, + OnMQTT: func(c net.Conn) { + gotMQTT <- c + _ = c.Close() + }, + }) + if err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + if startErr := s.Start(ctx); startErr != nil { + t.Fatal(startErr) + } + defer func() { _ = s.Close() }() + + addr := s.ListenAddr() + tlsCfg := &tls.Config{InsecureSkipVerify: true} + + // TLS + HTTP + tr := &http.Transport{TLSClientConfig: tlsCfg} + client := &http.Client{Transport: tr, Timeout: 5 * time.Second} + resp, err := client.Get("https://" + addr + "/healthz") + if err != nil { + t.Fatal(err) + } + body, _ := io.ReadAll(resp.Body) + _ = resp.Body.Close() + if string(body) != "tls-ok" { + t.Fatalf("body=%q", body) + } + + // TLS + MQTT (0x10 after handshake) + raw, err := tls.Dial("tcp", addr, tlsCfg) + if err != nil { + t.Fatal(err) + } + _, _ = raw.Write([]byte{0x10}) + select { + case mc := <-gotMQTT: + _ = mc.Close() + case <-time.After(3 * time.Second): + t.Fatal("tls mqtt not delivered") + } + _ = raw.Close() + + // ALPN mqtt 客户端仍能握手(服务端不设 NextProtos) + alpn, err := tls.Dial("tcp", addr, &tls.Config{ + InsecureSkipVerify: true, + NextProtos: []string{"mqtt"}, + }) + if err != nil { + t.Fatalf("alpn mqtt handshake: %v", err) + } + _ = alpn.Close() +} + +func TestCertReloadUsesNewCert(t *testing.T) { + dir := t.TempDir() + certPath, keyPath := writeTestCert(t, dir, "v1") + cr, err := NewCertReloader(certPath, keyPath, nil) + if err != nil { + t.Fatal(err) + } + defer cr.Close() + old := cr.Certificate() + if old == nil { + t.Fatal("nil cert") + } + + time.Sleep(20 * time.Millisecond) // 保证 mtime 变化 + certPath2, keyPath2 := writeTestCert(t, dir, "v2") + // 覆盖原路径 + data, _ := os.ReadFile(certPath2) + _ = os.WriteFile(certPath, data, 0o644) + data, _ = os.ReadFile(keyPath2) + _ = os.WriteFile(keyPath, data, 0o644) + + if reloadErr := cr.ReloadNow(); reloadErr != nil { + t.Fatal(reloadErr) + } + neu := cr.Certificate() + if neu == nil || neu == old { + t.Fatal("certificate not reloaded") + } + + // 完整服务:重载后新连接用新证书(用 Leaf CN 区分) + mux := NewMux(RoleShared, Handlers{}) + s, err := New(Options{ + Listen: "127.0.0.1:0", + DataDir: dir, + CertFile: certPath, + KeyFile: keyPath, + ClientHandler: mux, + }) + if err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + if startErr := s.Start(ctx); startErr != nil { + t.Fatal(startErr) + } + defer func() { _ = s.Close() }() + + // 再换一版证书 + time.Sleep(20 * time.Millisecond) + c3, k3 := writeTestCert(t, dir, "v3") + data, _ = os.ReadFile(c3) + _ = os.WriteFile(certPath, data, 0o644) + data, _ = os.ReadFile(k3) + _ = os.WriteFile(keyPath, data, 0o644) + if reloadErr := s.certs.ReloadNow(); reloadErr != nil { + t.Fatal(reloadErr) + } + + conn, err := tls.Dial("tcp", s.ListenAddr(), &tls.Config{InsecureSkipVerify: true}) + if err != nil { + t.Fatal(err) + } + defer func() { _ = conn.Close() }() + state := conn.ConnectionState() + if len(state.PeerCertificates) == 0 { + t.Fatal("no peer cert") + } + if cn := state.PeerCertificates[0].Subject.CommonName; cn != "v3" { + t.Fatalf("cn=%q want v3", cn) + } +} + +func TestAdminSeparateMQTTClosedAndAdmin404OnListen(t *testing.T) { + dir := t.TempDir() + clientMux := NewMux(RoleClient, Handlers{ + Healthz: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write([]byte("client")) + }), + }) + adminMux := NewMux(RoleAdmin, Handlers{ + AdminAPI: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write([]byte("admin")) + }), + Healthz: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write([]byte("admin-health")) + }), + }) + mqttSeen := make(chan struct{}, 1) + s, err := New(Options{ + Listen: "127.0.0.1:0", + AdminListen: "127.0.0.1:0", + DataDir: dir, + AllowPlaintext: true, + ClientHandler: clientMux, + AdminHandler: adminMux, + OnMQTT: func(c net.Conn) { + mqttSeen <- struct{}{} + _ = c.Close() + }, + }) + if err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + if startErr := s.Start(ctx); startErr != nil { + t.Fatal(startErr) + } + defer func() { _ = s.Close() }() + + // listen 上 /api/admin/ 404 + resp, err := http.Get("http://" + s.ListenAddr() + "/api/admin/x") + if err != nil { + t.Fatal(err) + } + _ = resp.Body.Close() + if resp.StatusCode != 404 { + t.Fatalf("listen admin status=%d", resp.StatusCode) + } + + // admin 上 API 可用 + resp, err = http.Get("http://" + s.AdminAddr() + "/api/admin/x") + if err != nil { + t.Fatal(err) + } + body, _ := io.ReadAll(resp.Body) + _ = resp.Body.Close() + if string(body) != "admin" { + t.Fatalf("admin body=%q", body) + } + + // admin_listen 上裸 MQTT 关闭,不回调 + ac, err := net.DialTimeout("tcp", s.AdminAddr(), 2*time.Second) + if err != nil { + t.Fatal(err) + } + _, _ = ac.Write([]byte{0x10}) + time.Sleep(200 * time.Millisecond) + buf := make([]byte, 1) + _ = ac.SetReadDeadline(time.Now().Add(500 * time.Millisecond)) + _, readErr := ac.Read(buf) + _ = ac.Close() + if readErr == nil { + t.Fatal("expected admin mqtt connection closed") + } + select { + case <-mqttSeen: + t.Fatal("mqtt should not be accepted on admin") + default: + } + + if _, err := os.ReadFile(filepath.Join(dir, "admin.addr")); err != nil { + t.Fatal(err) + } +} + +func TestTrustedProxyClientIP(t *testing.T) { + ps, err := ParseTrustedProxies([]string{"10.0.0.0/8", "192.168.1.1"}) + if err != nil { + t.Fatal(err) + } + r := &http.Request{ + RemoteAddr: "10.1.2.3:1234", + Header: http.Header{"X-Forwarded-For": []string{"1.1.1.1, 10.9.9.9"}}, + } + if got := ps.ClientIP(r); got != "1.1.1.1" { + t.Fatalf("got %q", got) + } + // 非代理来源忽略头 + r2 := &http.Request{ + RemoteAddr: "8.8.8.8:9", + Header: http.Header{"X-Forwarded-For": []string{"1.1.1.1"}}, + } + if got := ps.ClientIP(r2); got != "8.8.8.8" { + t.Fatalf("got %q", got) + } +} + +func TestFirstByteTimeout(t *testing.T) { + c1, c2 := net.Pipe() + defer func() { _ = c1.Close() }() + defer func() { _ = c2.Close() }() + done := make(chan error, 1) + go func() { + _, _, err := peekFirstByte(c2) + done <- err + }() + select { + case err := <-done: + if err == nil { + t.Fatal("expected timeout error") + } + case <-time.After(12 * time.Second): + t.Fatal("peek did not time out") + } +} + +func writeTestCert(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/internal/listener/proxy.go b/internal/listener/proxy.go new file mode 100644 index 0000000..1bf9682 --- /dev/null +++ b/internal/listener/proxy.go @@ -0,0 +1,101 @@ +package listener + +import ( + "net" + "net/http" + "strings" +) + +// ProxySet 保存受信任代理地址段。 +type ProxySet struct { + nets []*net.IPNet +} + +// ParseTrustedProxies 解析 CIDR 或单 IP 列表。 +func ParseTrustedProxies(cidrs []string) (*ProxySet, error) { + ps := &ProxySet{} + for _, s := range cidrs { + s = strings.TrimSpace(s) + if s == "" { + continue + } + if !strings.Contains(s, "/") { + ip := net.ParseIP(s) + if ip == nil { + continue + } + if ip.To4() != nil { + s += "/32" + } else { + s += "/128" + } + } + _, n, err := net.ParseCIDR(s) + if err != nil { + return nil, err + } + ps.nets = append(ps.nets, n) + } + return ps, nil +} + +// Contains 判断 IP 是否在受信任段内。 +func (ps *ProxySet) Contains(ip net.IP) bool { + if ps == nil || ip == nil { + return false + } + for _, n := range ps.nets { + if n.Contains(ip) { + return true + } + } + return false +} + +// ClientIP 按 DEVELOPMENT 4.5:来自受信任代理时,取 X-Forwarded-For 从右往左第一个不在段内的 IP。 +// 非代理来源忽略转发头,返回 RemoteAddr 的 IP。 +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() +} + +// IsHTTPS 来自受信任代理时按 X-Forwarded-Proto 判断。 +func (ps *ProxySet) IsHTTPS(r *http.Request) bool { + if r.TLS != nil { + return true + } + remoteIP := ipFromAddr(r.RemoteAddr) + if ps == nil || remoteIP == nil || !ps.Contains(remoteIP) { + return false + } + return strings.EqualFold(r.Header.Get("X-Forwarded-Proto"), "https") +} + +func ipFromAddr(remoteAddr string) net.IP { + host, _, err := net.SplitHostPort(remoteAddr) + if err != nil { + return net.ParseIP(remoteAddr) + } + return net.ParseIP(host) +} diff --git a/internal/listener/routes.go b/internal/listener/routes.go new file mode 100644 index 0000000..cd0f7bf --- /dev/null +++ b/internal/listener/routes.go @@ -0,0 +1,124 @@ +package listener + +import ( + "net/http" + "strings" +) + +// RouteRole 区分监听用途,决定哪些路径可用。 +type RouteRole int + +const ( + // RoleShared listen 与 admin 共用同一端口。 + RoleShared RouteRole = iota + // RoleClient 仅端接入(admin_listen 已分离)。 + RoleClient + // RoleAdmin 仅后台。 + RoleAdmin +) + +// Handlers 由上层注入各路径处理函数;未设置的路径返回 404。 +type Handlers struct { + MQTT http.Handler // /mqtt + ClientAPI http.Handler // /api/client/ + AdminAPI http.Handler // /api/admin/ + Metrics http.Handler // /metrics + Static http.Handler // 后台静态页 + Healthz http.Handler // /healthz + Readyz http.Handler // /readyz + // MetricsToken 共用端口时校验 Authorization: Bearer;空则 /metrics 404。 + MetricsToken string +} + +// NewMux 按角色装配 HTTP 路由骨架。 +func NewMux(role RouteRole, h Handlers) http.Handler { + mux := http.NewServeMux() + + healthz := h.Healthz + if healthz == nil { + healthz = http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte("ok")) + }) + } + readyz := h.Readyz + if readyz == nil { + readyz = http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte("ok")) + }) + } + + mux.Handle("GET /healthz", healthz) + mux.Handle("GET /readyz", readyz) + + switch role { + case RoleClient: + if h.MQTT != nil { + mux.Handle("/mqtt", h.MQTT) + } + if h.ClientAPI != nil { + mux.Handle("/api/client/", h.ClientAPI) + } + // 后台路径在端端口一律 404 + mux.Handle("/api/admin/", http.NotFoundHandler()) + mux.Handle("/metrics", http.NotFoundHandler()) + mux.Handle("/", http.NotFoundHandler()) + + case RoleAdmin: + if h.AdminAPI != nil { + mux.Handle("/api/admin/", h.AdminAPI) + } + if h.Metrics != nil { + mux.Handle("GET /metrics", h.Metrics) + } else { + mux.Handle("GET /metrics", http.NotFoundHandler()) + } + mux.Handle("/mqtt", http.NotFoundHandler()) + mux.Handle("/api/client/", http.NotFoundHandler()) + if h.Static != nil { + mux.Handle("/", h.Static) + } else { + mux.Handle("/", http.NotFoundHandler()) + } + + default: // RoleShared + if h.MQTT != nil { + mux.Handle("/mqtt", h.MQTT) + } + if h.ClientAPI != nil { + mux.Handle("/api/client/", h.ClientAPI) + } + if h.AdminAPI != nil { + mux.Handle("/api/admin/", h.AdminAPI) + } + mux.Handle("GET /metrics", metricsGate(h.MetricsToken, h.Metrics)) + if h.Static != nil { + mux.Handle("/", h.Static) + } else { + mux.Handle("/", http.NotFoundHandler()) + } + } + + return mux +} + +func metricsGate(token string, next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if token == "" { + http.NotFound(w, r) + return + } + auth := r.Header.Get("Authorization") + const prefix = "Bearer " + if !strings.HasPrefix(auth, prefix) || auth[len(prefix):] != token { + w.WriteHeader(http.StatusUnauthorized) + return + } + if next == nil { + http.NotFound(w, r) + return + } + next.ServeHTTP(w, r) + }) +} diff --git a/internal/listener/server.go b/internal/listener/server.go new file mode 100644 index 0000000..bf92e27 --- /dev/null +++ b/internal/listener/server.go @@ -0,0 +1,338 @@ +package listener + +import ( + "context" + "crypto/tls" + "errors" + "fmt" + "log/slog" + "net" + "net/http" + "os" + "path/filepath" + "strings" + "sync" + "time" +) + +// Kind 识别结果。 +type Kind int + +const ( + KindHTTP Kind = iota + KindMQTT + KindClosed +) + +// Options 控制双端口识别与分流。 +type Options struct { + Listen string + AdminListen string + DataDir string + CertFile string + KeyFile string + AllowPlaintext bool + TrustedProxies []string + // ClientHandler / AdminHandler 分别为端端口与后台端口的 HTTP 处理;AdminListen 为空时只用 ClientHandler。 + ClientHandler http.Handler + AdminHandler http.Handler + // OnMQTT 在 listen 上识别到裸 MQTT(含 TLS 后)时调用;应阻塞到连接结束。 + OnMQTT func(conn net.Conn) + Logger *slog.Logger +} + +// Server 一个或两个 TCP 监听上的协议识别与分流。 +type Server struct { + opts Options + log *slog.Logger + proxies *ProxySet + certs *CertReloader + + clientHTTP *ChanListener + adminHTTP *ChanListener + clientSrv *http.Server + adminSrv *http.Server + + clientLn net.Listener + adminLn net.Listener + + listenAddr string + adminAddr string + writeListenAddr bool + writeAdminAddr bool + + wg sync.WaitGroup + closed chan struct{} + closeOnce sync.Once +} + +// New 校验选项并准备证书;不开始监听。 +func New(opts Options) (*Server, error) { + if opts.Listen == "" { + return nil, errors.New("listener: listen is required") + } + if opts.ClientHandler == nil { + return nil, errors.New("listener: ClientHandler is required") + } + log := opts.Logger + if log == nil { + log = slog.Default() + } + ps, err := ParseTrustedProxies(opts.TrustedProxies) + if err != nil { + return nil, fmt.Errorf("trusted_proxies: %w", err) + } + s := &Server{ + opts: opts, + log: log, + proxies: ps, + closed: make(chan struct{}), + } + hasCert := opts.CertFile != "" && opts.KeyFile != "" + if hasCert { + cr, err := NewCertReloader(opts.CertFile, opts.KeyFile, log) + if err != nil { + return nil, fmt.Errorf("tls: %w", err) + } + s.certs = cr + } else { + log.Warn("tls not configured; plaintext only") + } + s.writeListenAddr = isPortZero(opts.Listen) + s.writeAdminAddr = opts.AdminListen != "" && isPortZero(opts.AdminListen) + return s, nil +} + +// Start 开始监听并分流;非阻塞,关闭用 Close。 +func (s *Server) Start(ctx context.Context) error { + ln, err := net.Listen("tcp", s.opts.Listen) + if err != nil { + return fmt.Errorf("listen %s: %w", s.opts.Listen, err) + } + s.clientLn = ln + s.listenAddr = ln.Addr().String() + if s.writeListenAddr { + if err := writeAddrFile(s.opts.DataDir, "listen.addr", s.listenAddr); err != nil { + _ = ln.Close() + return err + } + } + + s.clientHTTP = NewChanListener(ln.Addr(), 128) + s.clientSrv = &http.Server{ + Handler: s.opts.ClientHandler, + ReadHeaderTimeout: 10 * time.Second, + } + s.wg.Add(1) + go func() { + defer s.wg.Done() + err := s.clientSrv.Serve(s.clientHTTP) + if err != nil && !errors.Is(err, http.ErrServerClosed) { + s.log.Error("client http serve", "err", err) + } + }() + s.wg.Add(1) + go s.acceptLoop(ln, false) + + if s.opts.AdminListen != "" { + aln, err := net.Listen("tcp", s.opts.AdminListen) + if err != nil { + _ = s.Close() + return fmt.Errorf("admin_listen %s: %w", s.opts.AdminListen, err) + } + s.adminLn = aln + s.adminAddr = aln.Addr().String() + if s.writeAdminAddr { + if err := writeAddrFile(s.opts.DataDir, "admin.addr", s.adminAddr); err != nil { + _ = s.Close() + return err + } + } + adminHandler := s.opts.AdminHandler + if adminHandler == nil { + adminHandler = http.NotFoundHandler() + } + s.adminHTTP = NewChanListener(aln.Addr(), 64) + s.adminSrv = &http.Server{ + Handler: adminHandler, + ReadHeaderTimeout: 10 * time.Second, + } + s.wg.Add(1) + go func() { + defer s.wg.Done() + err := s.adminSrv.Serve(s.adminHTTP) + if err != nil && !errors.Is(err, http.ErrServerClosed) { + s.log.Error("admin http serve", "err", err) + } + }() + s.wg.Add(1) + go s.acceptLoop(aln, true) + } + + go func() { + <-ctx.Done() + _ = s.Close() + }() + return nil +} + +// ListenAddr 返回端监听实际地址。 +func (s *Server) ListenAddr() string { return s.listenAddr } + +// AdminAddr 返回后台监听实际地址(未分离时为空)。 +func (s *Server) AdminAddr() string { return s.adminAddr } + +// Proxies 返回受信任代理集合。 +func (s *Server) Proxies() *ProxySet { return s.proxies } + +// TLSConfig 返回当前 TLS 配置(未配置证书时为 nil)。 +func (s *Server) TLSConfig() *tls.Config { + if s.certs == nil { + return nil + } + return s.certs.TLSConfig() +} + +// Close 停止接受并关闭 HTTP。 +func (s *Server) Close() error { + var first error + s.closeOnce.Do(func() { + close(s.closed) + if s.clientLn != nil { + _ = s.clientLn.Close() + } + if s.adminLn != nil { + _ = s.adminLn.Close() + } + shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + if s.clientSrv != nil { + if err := s.clientSrv.Shutdown(shutdownCtx); err != nil && first == nil { + first = err + } + } + if s.adminSrv != nil { + if err := s.adminSrv.Shutdown(shutdownCtx); err != nil && first == nil { + first = err + } + } + if s.clientHTTP != nil { + _ = s.clientHTTP.Close() + } + if s.adminHTTP != nil { + _ = s.adminHTTP.Close() + } + if s.certs != nil { + s.certs.Close() + } + }) + s.wg.Wait() + return first +} + +func (s *Server) acceptLoop(ln net.Listener, isAdmin bool) { + defer s.wg.Done() + for { + c, err := ln.Accept() + if err != nil { + select { + case <-s.closed: + return + default: + return + } + } + s.wg.Add(1) + go func(conn net.Conn) { + defer s.wg.Done() + s.handleConn(conn, isAdmin) + }(c) + } +} + +func (s *Server) handleConn(conn net.Conn, isAdmin bool) { + kind, out, err := s.classify(conn, isAdmin, false) + if err != nil || kind == KindClosed { + _ = conn.Close() + return + } + switch kind { + case KindHTTP: + httpLn := s.clientHTTP + if isAdmin { + httpLn = s.adminHTTP + } + if httpLn == nil { + _ = out.Close() + return + } + httpLn.Enqueue(out) + case KindMQTT: + if s.opts.OnMQTT != nil { + s.opts.OnMQTT(out) + } else { + _ = out.Close() + } + default: + _ = out.Close() + } +} + +// classify 读首字节分流;afterTLS 表示已在 TLS 内层再识别。 +func (s *Server) classify(conn net.Conn, isAdmin, afterTLS bool) (Kind, net.Conn, error) { + c, b, err := peekFirstByte(conn) + if err != nil { + return KindClosed, nil, err + } + + hasCert := s.certs != nil + allowPlain := s.opts.AllowPlaintext || !hasCert + + if b == 0x16 { + if !hasCert { + return KindClosed, nil, errors.New("tls client hello but no certificate") + } + if afterTLS { + return KindClosed, nil, errors.New("nested tls") + } + tlsConn := tls.Server(c, s.certs.TLSConfig()) + if err := tlsConn.Handshake(); err != nil { + return KindClosed, nil, err + } + return s.classify(tlsConn, isAdmin, true) + } + + if !allowPlain && !afterTLS { + // 配了证书且未允许明文:非 TLS 首字节直接关 + return KindClosed, nil, errors.New("plaintext not allowed") + } + + if b >= 'A' && b <= 'Z' { + return KindHTTP, c, nil + } + if b == 0x10 { + if isAdmin { + return KindClosed, nil, errors.New("mqtt not allowed on admin_listen") + } + return KindMQTT, c, nil + } + return KindClosed, nil, errors.New("unknown first byte") +} + +func writeAddrFile(dataDir, name, addr string) error { + if dataDir == "" { + return errors.New("data_dir required to write addr file") + } + if err := os.MkdirAll(dataDir, 0o755); err != nil { + return err + } + return os.WriteFile(filepath.Join(dataDir, name), []byte(addr+"\n"), 0o644) +} + +func isPortZero(addr string) bool { + _, port, err := net.SplitHostPort(addr) + if err != nil { + return strings.HasSuffix(addr, ":0") || addr == ":0" + } + return port == "0" +} diff --git a/internal/listener/tls.go b/internal/listener/tls.go new file mode 100644 index 0000000..aa112b5 --- /dev/null +++ b/internal/listener/tls.go @@ -0,0 +1,117 @@ +package listener + +import ( + "crypto/tls" + "log/slog" + "os" + "sync" + "time" +) + +// CertReloader 按文件修改时间每小时重载证书;失败继续用旧证书。 +type CertReloader struct { + certFile string + keyFile string + log *slog.Logger + + mu sync.RWMutex + cert *tls.Certificate + certMod time.Time + keyMod time.Time + stop chan struct{} + stopOnce sync.Once +} + +// NewCertReloader 立即加载一次证书。 +func NewCertReloader(certFile, keyFile string, log *slog.Logger) (*CertReloader, error) { + if log == nil { + log = slog.Default() + } + r := &CertReloader{ + certFile: certFile, + keyFile: keyFile, + log: log, + stop: make(chan struct{}), + } + if err := r.reload(true); err != nil { + return nil, err + } + go r.loop() + return r, nil +} + +// GetCertificate 供 tls.Config.GetCertificate 使用。 +func (r *CertReloader) GetCertificate(*tls.ClientHelloInfo) (*tls.Certificate, error) { + r.mu.RLock() + defer r.mu.RUnlock() + return r.cert, nil +} + +// Certificate 返回当前证书(测试用)。 +func (r *CertReloader) Certificate() *tls.Certificate { + r.mu.RLock() + defer r.mu.RUnlock() + return r.cert +} + +// Close 停止重载循环。 +func (r *CertReloader) Close() { + r.stopOnce.Do(func() { close(r.stop) }) +} + +func (r *CertReloader) loop() { + t := time.NewTicker(time.Hour) + defer t.Stop() + for { + select { + case <-r.stop: + return + case <-t.C: + if err := r.reload(false); err != nil { + r.log.Error("tls cert reload failed, keeping old cert", "err", err) + } + } + } +} + +// ReloadNow 立即按 mtime 检查并重载(测试用)。 +func (r *CertReloader) ReloadNow() error { + return r.reload(false) +} + +func (r *CertReloader) reload(force bool) error { + certInfo, err := os.Stat(r.certFile) + if err != nil { + return err + } + keyInfo, err := os.Stat(r.keyFile) + if err != nil { + return err + } + r.mu.RLock() + same := !force && certInfo.ModTime().Equal(r.certMod) && keyInfo.ModTime().Equal(r.keyMod) && r.cert != nil + r.mu.RUnlock() + if same { + return nil + } + cert, err := tls.LoadX509KeyPair(r.certFile, r.keyFile) + if err != nil { + return err + } + r.mu.Lock() + r.cert = &cert + r.certMod = certInfo.ModTime() + r.keyMod = keyInfo.ModTime() + r.mu.Unlock() + r.log.Info("tls certificate loaded", "cert", r.certFile) + return nil +} + +// TLSConfig 构造不设 NextProtos 的服务端 TLS 配置。 +func (r *CertReloader) TLSConfig() *tls.Config { + return &tls.Config{ + GetCertificate: r.GetCertificate, + MinVersion: tls.VersionTLS12, + // 故意不设 NextProtos,以便声明 ALPN mqtt 的客户端仍能握手。 + } +}