feat: 实现端口识别、TLS 证书重载与路由骨架
This commit is contained in:
@@ -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:
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -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"
|
||||||
|
}
|
||||||
@@ -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 的客户端仍能握手。
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user