2 Commits
15 changed files with 2186 additions and 1 deletions
+36 -1
View File
@@ -239,7 +239,42 @@
## 连接 N
暂无。
### N1 / N2 2026-09-30
1. **未接线 `cmd/nixmsg`**
- 原条款:serve 最终应挂上端口识别、broker、`/mqtt`。
- 实际做法:本任务只交付 `internal/listener`、`internal/broker`;按总控要求不改 `cmd/nixmsg`。
- 原因:避免与平台/总控并行改 wire 冲突;合并时再接线。
- 备选方案:本分支顺带改 `wire.go`(与指令冲突)。
- 影响:当前 `serve` 仍是 T0.4 的简单 `/healthz` 监听,不含 MQTT。
2. **`listen.addr` / `admin.addr` 仅端口为 0 时写入**
- 原条款:DEVELOPMENT 4.1「端口写 0 时」写地址文件;T0.1 偏差曾改为 always write。
- 实际做法:`listener.Server` 仅当配置地址端口为 `0` 时写 `listen.addr` / `admin.addr`。
- 原因:本任务说明与 DEVELOPMENT 4.1 字面一致;T0.1 的 always write 在 `cmd/nixmsg`,本线未改。
- 备选方案:接线时统一为 always write 以兼容 harness。
- 影响:固定端口场景下 harness 若只读地址文件会读不到;接线时建议沿用 T0.1 超集或改 harness。
3. **登录校验为可替换接口,默认拒绝**
- 原条款:第 5 节完整会话令牌/密码/锁定属 N3。
- 实际做法:`broker.Authenticator` 接口 + 默认 `RejectAuthenticator`;内部错误在 `OnConnect` 返回 error;测试提供 `AllowAuthenticator`。
- 原因:N3 范围;N2 需可跑通装配与钩子。
- 备选方案:N2 内做假登录表(超出范围)。
- 影响:真实端连不上直到 N3;总控接线时注入 Authenticator。
4. **大帧并发名额释放策略**
- 原条款:DEVELOPMENT 7.5 大于 64KiB 全局同时不超过 64;PUBACK / 超时 / 断线释放。
- 实际做法:发布前申请名额;QoS 0 发布成功立即释放;QoS 1 在 `OnQosComplete` 且 payload>64KiB 时释放,断线 `releaseAllLarge`;未单独做「确认超时」计时释放(确认超时属 M 线推送循环)。
- 原因:N2 无投递确认计时器;与 M 线推送超时释放衔接。
- 备选方案:broker 内对大帧自建超时(与 M 重复)。
- 影响:若客户端永不 PUBACK 且不断线,名额可能占满直到断开;M 线超时踢线或回调 Disconnect 可释放。
5. **`OnPublishDropped` 仅打日志**
- 原条款:清「已推送」标记并 1 秒后重推。
- 实际做法:钩子记录 debug 日志;清标记/重推留给消息 M。
- 原因:投递状态在 M/store,N2 无投递表。
- 备选方案:N2 暴露回调给 M 注册。
- 影响:接线后 M 需订阅或包装该钩子;当前接口可后续加 `OnPublishDropped` 回调字段。
## 消息 M
+4
View File
@@ -3,6 +3,8 @@ module git.asio.asia/nixevol/NixMsg
go 1.27
require (
github.com/coder/websocket v1.8.14
github.com/mochi-mqtt/server/v2 v2.7.9
github.com/prometheus/client_golang v1.24.1
go.yaml.in/yaml/v3 v3.0.5
golang.org/x/crypto v0.57.0
@@ -14,6 +16,7 @@ require (
github.com/cespare/xxhash/v2 v2.3.0 // indirect
github.com/dustin/go-humanize v1.0.1 // indirect
github.com/google/uuid v1.6.0 // indirect
github.com/gorilla/websocket v1.5.0 // indirect
github.com/mattn/go-isatty v0.0.24 // indirect
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect
github.com/ncruces/go-strftime v1.0.0 // indirect
@@ -21,6 +24,7 @@ require (
github.com/prometheus/common v0.70.1 // indirect
github.com/prometheus/procfs v0.21.1 // indirect
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
github.com/rs/xid v1.4.0 // indirect
golang.org/x/sys v0.48.0 // indirect
google.golang.org/protobuf v1.36.11 // indirect
modernc.org/libc v1.77.1 // indirect
+10
View File
@@ -2,6 +2,8 @@ github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM=
github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw=
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
github.com/coder/websocket v1.8.14 h1:9L0p0iKiNOibykf283eHkKUHHrpG7f65OE3BhhO7v9g=
github.com/coder/websocket v1.8.14/go.mod h1:NX3SzP+inril6yawo5CQXx8+fk145lPDC6pumgx0mVg=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY=
@@ -12,14 +14,20 @@ github.com/google/pprof v0.0.0-20260802141513-ef3492d7dac3 h1:LMLX+LgTNWpfvCBdFe
github.com/google/pprof v0.0.0-20260802141513-ef3492d7dac3/go.mod h1:jl5iWTm0/hd5PjEYEOuwAJ57L/CibdZfrqZ5XA5GrCk=
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
github.com/gorilla/websocket v1.5.0 h1:PPwGk2jz7EePpoHN/+ClbZu8SPxiqlu12wZP/3sWmnc=
github.com/gorilla/websocket v1.5.0/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k=
github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM=
github.com/jinzhu/copier v0.3.5 h1:GlvfUwHk62RokgqVNvYsku0TATCF7bAHVwEXoBh3iJg=
github.com/jinzhu/copier v0.3.5/go.mod h1:DfbEm0FYsaqBcKcFuvmOZb218JkPGtvSHsKg8S8hyyg=
github.com/klauspost/compress v1.19.1 h1:VsB4HPswih7mmZ8WleSFQ75c/Ui1M4trX5oAsJnhSlk=
github.com/klauspost/compress v1.19.1/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ=
github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc=
github.com/kylelemons/godebug v1.1.0/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw=
github.com/mattn/go-isatty v0.0.24 h1:tGZZoVgT/KiqK1c8ocVLeDS8BSWMRd47J3Lbz7vsReI=
github.com/mattn/go-isatty v0.0.24/go.mod h1:nMCL3Zebbrt45jsMDgnfIwz6ydEQApk5oEI3HqDio6A=
github.com/mochi-mqtt/server/v2 v2.7.9 h1:y0g4vrSLAag7T07l2oCzOa/+nKVLoazKEWAArwqBNYI=
github.com/mochi-mqtt/server/v2 v2.7.9/go.mod h1:lZD3j35AVNqJL5cezlnSkuG05c0FCHSsfAKSPBOSbqc=
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA=
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ=
github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w=
@@ -36,6 +44,8 @@ github.com/prometheus/procfs v0.21.1 h1:GljZCt+zSTS+NZq88cyQ1LjZ+RCHp3uVuabBWA5+
github.com/prometheus/procfs v0.21.1/go.mod h1:aB55Cww9pdSJVHk0hUf0inxWyyjPogFIjmHKYgMKmtY=
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE=
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=
github.com/rs/xid v1.4.0 h1:qd7wPTDkN6KQx2VmMBLrpHkiyQwgFXRnkOLacUiaSNY=
github.com/rs/xid v1.4.0/go.mod h1:trrq9SKmegXys3aeAKXMUTdJsYXVwGY3RLcfgqegfbg=
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
+379
View File
@@ -0,0 +1,379 @@
package broker
import (
"context"
"crypto/rand"
"encoding/hex"
"errors"
"log/slog"
"net"
"sync"
"sync/atomic"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
mqtt "github.com/mochi-mqtt/server/v2"
"github.com/mochi-mqtt/server/v2/packets"
)
const (
maxClients = 2000
maxPacketSize = 786432
uplinkQueueSize = 256
largeFrameBytes = 64 * 1024
largeFrameSlots = 64
packetOverheadBudget = 128 // 主题与 MQTT 包头预留
keepaliveMin = 10
keepaliveMax = 600
)
// ErrPayloadTooLarge 下行超过客户端 Maximum Packet Size(减包头预留)或 max_receive_bytes。
var ErrPayloadTooLarge = errors.New("broker: payload exceeds client limit")
// ErrNoConnection 目标端没有当前连接。
var ErrNoConnection = errors.New("broker: no active connection")
// AuthResult 是登录校验结论(N3 实现真实逻辑;N2 默认拒绝)。
type AuthResult struct {
OK bool
SessionToken string // 密码登录成功时由 N3 填写
}
// Authenticator 由 N3 实现;内部故障必须返回 error,不得当成密码错误。
type Authenticator interface {
Authenticate(ctx context.Context, endpointID string, password []byte, remoteIP string) (AuthResult, error)
}
// RejectAuthenticator 默认拒绝所有客户端(CONNACK 用户名密码错误)。
type RejectAuthenticator struct{}
func (RejectAuthenticator) Authenticate(context.Context, string, []byte, string) (AuthResult, error) {
return AuthResult{OK: false}, nil
}
// AllowAuthenticator 测试用:允许任意编号。
type AllowAuthenticator struct{}
func (AllowAuthenticator) Authenticate(context.Context, string, []byte, string) (AuthResult, error) {
return AuthResult{OK: true}, nil
}
// Options 装配 broker。
type Options struct {
Authenticator Authenticator
Uplink port.UplinkHandler
Logger *slog.Logger
}
// Broker 内置 mochi,不自带监听端口。
type Broker struct {
server *mqtt.Server
auth Authenticator
uplink port.UplinkHandler
log *slog.Logger
hook *nixHook
connsMu sync.RWMutex
current map[string]*connState
byClient map[*mqtt.Client]*connState
queuesMu sync.Mutex
queues map[string]*uplinkQueue
largeSem chan struct{}
closed atomic.Bool
}
type connState struct {
connID port.ConnID
endpointID string
transport port.Transport
remoteIP string
client *mqtt.Client
maxPacketSize uint32
maxRecvBytes int
authOK bool
authErr error
sessionToken string
largeHeld int
mu sync.Mutex
}
// New 创建并 Serve mochi(无监听器)。
func New(opts Options) (*Broker, error) {
auth := opts.Authenticator
if auth == nil {
auth = RejectAuthenticator{}
}
uplink := opts.Uplink
if uplink == nil {
uplink = port.StubUplinkHandler{}
}
log := opts.Logger
if log == nil {
log = slog.Default()
}
caps := mqtt.NewDefaultServerCapabilities()
caps.MaximumClients = maxClients
caps.MaximumQos = 1
caps.MaximumPacketSize = maxPacketSize
caps.MaximumSessionExpiryInterval = 0
caps.ReceiveMaximum = 1024
caps.MaximumInflight = 1024
caps.MaximumClientWritesPending = 1024
caps.RetainAvailable = 0
caps.WildcardSubAvailable = 0
caps.SharedSubAvailable = 0
caps.TopicAliasMaximum = 0
caps.Compatibilities.ObscureNotAuthorized = true
srv := mqtt.New(&mqtt.Options{
InlineClient: true,
Capabilities: caps,
Logger: log,
})
b := &Broker{
server: srv,
auth: auth,
uplink: uplink,
log: log,
current: make(map[string]*connState),
byClient: make(map[*mqtt.Client]*connState),
queues: make(map[string]*uplinkQueue),
largeSem: make(chan struct{}, largeFrameSlots),
}
b.hook = &nixHook{b: b}
if err := srv.AddHook(b.hook, nil); err != nil {
return nil, err
}
if err := srv.Serve(); err != nil {
return nil, err
}
return b, nil
}
// Server 返回底层 mochi(测试用)。
func (b *Broker) Server() *mqtt.Server { return b.server }
// Close 关闭 broker。
func (b *Broker) Close() error {
if b.closed.Swap(true) {
return nil
}
b.queuesMu.Lock()
for _, q := range b.queues {
q.close()
}
b.queuesMu.Unlock()
return b.server.Close()
}
// AttachTCP 把裸 TCP/TLS 连接交给 mochi;阻塞到连接结束。
func (b *Broker) AttachTCP(conn net.Conn) error {
return b.server.EstablishConnection("tcp", conn)
}
// AttachWS 把 WebSocket NetConn 交给 mochi;阻塞到连接结束。
func (b *Broker) AttachWS(conn net.Conn) error {
return b.server.EstablishConnection("ws", conn)
}
// PublishDown 实现 port.Downlink。
func (b *Broker) PublishDown(ctx context.Context, endpointID string, connID port.ConnID, payload []byte, opts port.PublishOpts) error {
if b.closed.Load() {
return errors.New("broker: closed")
}
st := b.lookupConn(endpointID, connID)
if st == nil {
return ErrNoConnection
}
limit := effectivePayloadLimit(st.maxPacketSize, st.maxRecvBytes)
if limit > 0 && len(payload) > limit {
return ErrPayloadTooLarge
}
qos := opts.QoS
if qos > 1 {
qos = 1
}
topic := downTopic(endpointID)
large := len(payload) > largeFrameBytes
if large {
select {
case b.largeSem <- struct{}{}:
case <-ctx.Done():
return ctx.Err()
}
st.mu.Lock()
st.largeHeld++
st.mu.Unlock()
}
if err := b.server.Publish(topic, payload, false, qos); err != nil {
if large {
b.releaseOneLarge(st)
}
return err
}
if large && qos == 0 {
b.releaseOneLarge(st)
}
return nil
}
func (b *Broker) releaseOneLarge(st *connState) {
st.mu.Lock()
if st.largeHeld > 0 {
st.largeHeld--
st.mu.Unlock()
select {
case <-b.largeSem:
default:
}
return
}
st.mu.Unlock()
}
func (b *Broker) releaseAllLarge(st *connState) {
st.mu.Lock()
n := st.largeHeld
st.largeHeld = 0
st.mu.Unlock()
for i := 0; i < n; i++ {
select {
case <-b.largeSem:
default:
}
}
}
// Disconnect 实现 port.ConnControl。
func (b *Broker) Disconnect(_ context.Context, endpointID string, connID port.ConnID, reason port.DisconnectReason) error {
st := b.lookupConn(endpointID, connID)
if st == nil {
return ErrNoConnection
}
code := packets.CodeDisconnect
switch reason {
case port.DisconnectTakenOver:
code = packets.ErrSessionTakenOver
case port.DisconnectKicked, port.DisconnectFatal:
code = packets.ErrAdministrativeAction
}
return b.server.DisconnectClient(st.client, code)
}
func (b *Broker) lookupConn(endpointID string, connID port.ConnID) *connState {
b.connsMu.RLock()
defer b.connsMu.RUnlock()
if connID != "" {
for _, st := range b.byClient {
if st.endpointID == endpointID && st.connID == connID {
return st
}
}
return nil
}
return b.current[endpointID]
}
func downTopic(endpointID string) string {
return "nix/c/" + endpointID + "/down"
}
func upTopic(endpointID string) string {
return "nix/c/" + endpointID + "/up"
}
func effectivePayloadLimit(maxPacketSize uint32, maxRecvBytes int) int {
limit := 0
if maxPacketSize > 0 {
if maxPacketSize > packetOverheadBudget {
limit = int(maxPacketSize) - packetOverheadBudget
}
}
if maxRecvBytes > 0 {
if limit == 0 || maxRecvBytes < limit {
limit = maxRecvBytes
}
}
return limit
}
func randomConnID() port.ConnID {
var b [16]byte
_, _ = rand.Read(b[:])
return port.ConnID(hex.EncodeToString(b[:]))
}
func transportOf(cl *mqtt.Client) port.Transport {
if cl != nil && cl.Net.Listener == "ws" {
return port.TransportWS
}
return port.TransportTCP
}
func remoteIPOf(cl *mqtt.Client) string {
if cl == nil {
return ""
}
addr := cl.Net.Remote
if addr == "" && cl.Net.Conn != nil && cl.Net.Conn.RemoteAddr() != nil {
addr = cl.Net.Conn.RemoteAddr().String()
}
host, _, err := net.SplitHostPort(addr)
if err != nil {
return addr
}
return host
}
// SetMaxReceiveBytes 供 N3 握手后设置;0 表示不限。
func (b *Broker) SetMaxReceiveBytes(endpointID string, connID port.ConnID, n int) {
st := b.lookupConn(endpointID, connID)
if st == nil {
return
}
st.mu.Lock()
st.maxRecvBytes = n
st.mu.Unlock()
}
// ConnInfoOf 返回连接信息(测试/N3)。
func (b *Broker) ConnInfoOf(endpointID string) (port.ConnInfo, bool) {
b.connsMu.RLock()
st := b.current[endpointID]
b.connsMu.RUnlock()
if st == nil {
return port.ConnInfo{}, false
}
return port.ConnInfo{
ConnID: st.connID,
EndpointID: st.endpointID,
Transport: st.transport,
RemoteIP: st.remoteIP,
SessionToken: st.sessionToken,
MaxPacketSize: st.maxPacketSize,
}, true
}
func (b *Broker) enqueueUplink(endpointID string, conn port.ConnInfo, payload []byte) {
b.queuesMu.Lock()
q, ok := b.queues[endpointID]
if !ok {
q = newUplinkQueue(b, endpointID)
b.queues[endpointID] = q
}
b.queuesMu.Unlock()
q.push(uplinkItem{conn: conn, payload: payload})
}
var (
_ port.Downlink = (*Broker)(nil)
_ port.ConnControl = (*Broker)(nil)
)
+259
View File
@@ -0,0 +1,259 @@
package broker
import (
"bytes"
"context"
"errors"
"io"
"net"
"net/http"
"net/http/httptest"
"sync"
"testing"
"time"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
"github.com/coder/websocket"
"github.com/mochi-mqtt/server/v2/packets"
)
func TestWSCrossOriginAllowed(t *testing.T) {
b, err := New(Options{Authenticator: AllowAuthenticator{}})
if err != nil {
t.Fatal(err)
}
defer func() { _ = b.Close() }()
mux := http.NewServeMux()
mux.Handle("/mqtt", b.WSHandler(nil))
srv := httptest.NewServer(mux)
defer srv.Close()
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
c, _, err := websocket.Dial(ctx, "ws"+srv.URL[len("http"):]+"/mqtt", &websocket.DialOptions{
HTTPHeader: http.Header{"Origin": []string{"https://other.example"}},
Subprotocols: []string{"mqtt"},
})
if err != nil {
t.Fatalf("cross-origin dial: %v", err)
}
defer func() { _ = c.Close(websocket.StatusNormalClosure, "") }()
if c.Subprotocol() != "mqtt" {
t.Fatalf("subprotocol=%q", c.Subprotocol())
}
}
func TestWSWrongSubprotocolClosed(t *testing.T) {
b, err := New(Options{Authenticator: AllowAuthenticator{}})
if err != nil {
t.Fatal(err)
}
defer func() { _ = b.Close() }()
mux := http.NewServeMux()
mux.Handle("/mqtt", b.WSHandler(nil))
srv := httptest.NewServer(mux)
defer srv.Close()
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
c, _, err := websocket.Dial(ctx, "ws"+srv.URL[len("http"):]+"/mqtt", &websocket.DialOptions{
HTTPHeader: http.Header{"Origin": []string{"https://other.example"}},
Subprotocols: []string{"not-mqtt"},
})
if err != nil {
// 有的实现在握手阶段就失败;也算关闭
return
}
defer func() { _ = c.Close(websocket.StatusNormalClosure, "") }()
// 服务端应立刻关掉;后续读写会失败
c.SetReadLimit(16)
_, _, readErr := c.Read(ctx)
if readErr == nil {
t.Fatal("expected connection closed for wrong subprotocol")
}
}
func TestPublishDownExceedsClientMax(t *testing.T) {
b, err := New(Options{Authenticator: AllowAuthenticator{}})
if err != nil {
t.Fatal(err)
}
defer func() { _ = b.Close() }()
clientDone := make(chan struct{})
r, w := net.Pipe()
go func() {
defer close(clientDone)
_ = b.AttachTCP(r)
}()
endpoint := "ep-limit"
connectAndSubscribe(t, w, endpoint, 200) // MaximumPacketSize=200 → payload limit 72
// 等会话建立
deadline := time.Now().Add(3 * time.Second)
for {
if _, ok := b.ConnInfoOf(endpoint); ok {
break
}
if time.Now().After(deadline) {
t.Fatal("session not established")
}
time.Sleep(10 * time.Millisecond)
}
big := bytes.Repeat([]byte("x"), 100) // > 200-128
pubErr := b.PublishDown(context.Background(), endpoint, "", big, port.PublishOpts{QoS: 1})
if !errors.Is(pubErr, ErrPayloadTooLarge) {
t.Fatalf("PublishDown err=%v want ErrPayloadTooLarge", pubErr)
}
// 合法大小应成功
small := []byte(`{"v":1,"type":"resp"}`)
if err := b.PublishDown(context.Background(), endpoint, "", small, port.PublishOpts{QoS: 0}); err != nil {
t.Fatalf("small publish: %v", err)
}
_ = w.Close()
select {
case <-clientDone:
case <-time.After(3 * time.Second):
}
}
func TestInternalAuthErrorDoesNotReturnBadPassword(t *testing.T) {
auth := &errAuthenticator{err: context.DeadlineExceeded}
b, err := New(Options{Authenticator: auth})
if err != nil {
t.Fatal(err)
}
defer func() { _ = b.Close() }()
r, w := net.Pipe()
errCh := make(chan error, 1)
go func() { errCh <- b.AttachTCP(r) }()
writeConnect(t, w, "ep-err", 30, 0)
// 不应收到 CONNACK(内部错误直接断开)
_ = w.SetReadDeadline(time.Now().Add(500 * time.Millisecond))
buf := make([]byte, 64)
n, readErr := w.Read(buf)
if readErr == nil && n > 0 {
// 若收到包,不能是 bad username/password CONNACK (reason 0x86)
if n >= 2 && buf[0]>>4 == packets.Connack {
t.Fatalf("unexpected connack on internal error: %x", buf[:n])
}
}
_ = w.Close()
select {
case <-errCh:
case <-time.After(2 * time.Second):
}
}
func TestRejectUnknownByDefault(t *testing.T) {
b, err := New(Options{}) // RejectAuthenticator
if err != nil {
t.Fatal(err)
}
defer func() { _ = b.Close() }()
r, w := net.Pipe()
go func() { _ = b.AttachTCP(r) }()
writeConnect(t, w, "ep-unknown", 30, 0)
_ = w.SetReadDeadline(time.Now().Add(2 * time.Second))
buf := make([]byte, 128)
n, err := io.ReadAtLeast(w, buf, 2)
if err != nil {
t.Fatal(err)
}
if buf[0]>>4 != packets.Connack {
t.Fatalf("want connack, got %x", buf[:n])
}
_ = w.Close()
}
type errAuthenticator struct {
err error
mu sync.Mutex
}
func (a *errAuthenticator) Authenticate(context.Context, string, []byte, string) (AuthResult, error) {
a.mu.Lock()
defer a.mu.Unlock()
return AuthResult{}, a.err
}
func connectAndSubscribe(t *testing.T, w net.Conn, endpoint string, maxPacket uint32) {
t.Helper()
writeConnect(t, w, endpoint, 30, maxPacket)
// read CONNACK
_ = w.SetReadDeadline(time.Now().Add(3 * time.Second))
buf := make([]byte, 256)
n, err := io.ReadAtLeast(w, buf, 2)
if err != nil {
t.Fatal(err)
}
if buf[0]>>4 != packets.Connack {
t.Fatalf("want connack got %x", buf[:n])
}
writeSubscribe(t, w, downTopic(endpoint))
// read SUBACK
n, err = io.ReadAtLeast(w, buf, 2)
if err != nil {
t.Fatal(err)
}
if buf[0]>>4 != packets.Suback {
t.Fatalf("want suback got %x", buf[:n])
}
}
func writeConnect(t *testing.T, w net.Conn, endpoint string, keepalive uint16, maxPacket uint32) {
t.Helper()
pk := packets.Packet{
FixedHeader: packets.FixedHeader{Type: packets.Connect},
ProtocolVersion: 5,
Connect: packets.ConnectParams{
ProtocolName: []byte("MQTT"),
Clean: true,
ClientIdentifier: endpoint,
Keepalive: keepalive,
UsernameFlag: true,
Username: []byte(endpoint),
PasswordFlag: true,
Password: []byte("test"),
},
Properties: packets.Properties{
MaximumPacketSize: maxPacket,
},
}
var buf bytes.Buffer
if err := pk.ConnectEncode(&buf); err != nil {
t.Fatal(err)
}
if _, err := w.Write(buf.Bytes()); err != nil {
t.Fatal(err)
}
}
func writeSubscribe(t *testing.T, w net.Conn, topic string) {
t.Helper()
pk := packets.Packet{
FixedHeader: packets.FixedHeader{Type: packets.Subscribe, Qos: 1},
ProtocolVersion: 5,
PacketID: 1,
Filters: packets.Subscriptions{
{Filter: topic, Qos: 1},
},
}
var buf bytes.Buffer
if err := pk.SubscribeEncode(&buf); err != nil {
t.Fatal(err)
}
if _, err := w.Write(buf.Bytes()); err != nil {
t.Fatal(err)
}
}
+196
View File
@@ -0,0 +1,196 @@
package broker
import (
"bytes"
"context"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
mqtt "github.com/mochi-mqtt/server/v2"
"github.com/mochi-mqtt/server/v2/packets"
)
type nixHook struct {
mqtt.HookBase
b *Broker
}
func (h *nixHook) ID() string { return "nixmsg" }
func (h *nixHook) Provides(b byte) bool {
return bytes.Contains([]byte{
mqtt.OnConnect,
mqtt.OnConnectAuthenticate,
mqtt.OnACLCheck,
mqtt.OnPublish,
mqtt.OnPublishDropped,
mqtt.OnSessionEstablished,
mqtt.OnDisconnect,
mqtt.OnQosComplete,
}, []byte{b})
}
func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error {
endpointID := string(pk.Connect.Username)
if endpointID == "" {
endpointID = pk.Connect.ClientIdentifier
}
remoteIP := remoteIPOf(cl)
st := &connState{
connID: randomConnID(),
endpointID: endpointID,
transport: transportOf(cl),
remoteIP: remoteIP,
client: cl,
maxPacketSize: pk.Properties.MaximumPacketSize,
}
// 心跳校正:超出 10–600 秒就改写 Keepalive 并设 ServerKeepalive
ka := pk.Connect.Keepalive
if ka < keepaliveMin || ka > keepaliveMax {
if ka < keepaliveMin {
ka = keepaliveMin
}
if ka > keepaliveMax {
ka = keepaliveMax
}
cl.State.Keepalive = ka
cl.State.ServerKeepalive = true
}
res, err := h.b.auth.Authenticate(context.Background(), endpointID, pk.Connect.Password, remoteIP)
if err != nil {
st.authErr = err
h.rememberPending(cl, st)
return err // mochi 不回 CONNACK,直接断开
}
st.authOK = res.OK
st.sessionToken = res.SessionToken
h.rememberPending(cl, st)
return nil
}
func (h *nixHook) rememberPending(cl *mqtt.Client, st *connState) {
h.b.connsMu.Lock()
h.b.byClient[cl] = st
h.b.connsMu.Unlock()
}
func (h *nixHook) OnConnectAuthenticate(cl *mqtt.Client, _ packets.Packet) bool {
h.b.connsMu.RLock()
st := h.b.byClient[cl]
h.b.connsMu.RUnlock()
if st == nil {
return false
}
// 内部故障已在 OnConnect 返回 error;此处只反映业务上的拒绝
return st.authOK
}
func (h *nixHook) OnACLCheck(cl *mqtt.Client, topic string, write bool) bool {
h.b.connsMu.RLock()
st := h.b.byClient[cl]
h.b.connsMu.RUnlock()
if st == nil || st.endpointID == "" {
return false
}
up := upTopic(st.endpointID)
down := downTopic(st.endpointID)
if write {
return topic == up
}
return topic == down
}
func (h *nixHook) OnPublish(cl *mqtt.Client, pk packets.Packet) (packets.Packet, error) {
h.b.connsMu.RLock()
st := h.b.byClient[cl]
h.b.connsMu.RUnlock()
if st == nil {
return pk, packets.CodeSuccessIgnore
}
payload := append([]byte(nil), pk.Payload...)
info := port.ConnInfo{
ConnID: st.connID,
EndpointID: st.endpointID,
Transport: st.transport,
RemoteIP: st.remoteIP,
SessionToken: st.sessionToken,
MaxPacketSize: st.maxPacketSize,
}
h.b.enqueueUplink(st.endpointID, info, payload)
return pk, packets.CodeSuccessIgnore
}
func (h *nixHook) OnPublishDropped(cl *mqtt.Client, pk packets.Packet) {
h.b.log.Debug("publish dropped", "client", cl.ID, "topic", pk.TopicName, "size", len(pk.Payload))
}
func (h *nixHook) OnSessionEstablished(cl *mqtt.Client, _ packets.Packet) {
h.b.connsMu.Lock()
st := h.b.byClient[cl]
if st != nil {
h.b.current[st.endpointID] = st
}
h.b.connsMu.Unlock()
if st == nil {
return
}
info := port.ConnInfo{
ConnID: st.connID,
EndpointID: st.endpointID,
Transport: st.transport,
RemoteIP: st.remoteIP,
SessionToken: st.sessionToken,
MaxPacketSize: st.maxPacketSize,
}
_ = h.b.uplink.OnSessionEstablished(context.Background(), info)
}
func (h *nixHook) OnDisconnect(cl *mqtt.Client, err error, _ bool) {
h.b.connsMu.Lock()
st := h.b.byClient[cl]
delete(h.b.byClient, cl)
if st != nil && h.b.current[st.endpointID] == st {
delete(h.b.current, st.endpointID)
}
h.b.connsMu.Unlock()
if st == nil {
return
}
h.b.releaseAllLarge(st)
reason := port.DisconnectNormal
if err != nil {
if code, ok := err.(packets.Code); ok {
switch code.Code {
case packets.ErrSessionTakenOver.Code:
reason = port.DisconnectTakenOver
case packets.ErrAdministrativeAction.Code:
reason = port.DisconnectKicked
}
}
}
info := port.ConnInfo{
ConnID: st.connID,
EndpointID: st.endpointID,
Transport: st.transport,
RemoteIP: st.remoteIP,
SessionToken: st.sessionToken,
MaxPacketSize: st.maxPacketSize,
}
h.b.uplink.OnDisconnect(context.Background(), info, reason)
}
func (h *nixHook) OnQosComplete(cl *mqtt.Client, pk packets.Packet) {
if len(pk.Payload) <= largeFrameBytes {
return
}
h.b.connsMu.RLock()
st := h.b.byClient[cl]
h.b.connsMu.RUnlock()
if st == nil {
return
}
h.b.releaseOneLarge(st)
}
+45
View File
@@ -0,0 +1,45 @@
package broker
import (
"context"
"sync"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
)
type uplinkItem struct {
conn port.ConnInfo
payload []byte
}
// uplinkQueue 每端串行队列,长度 256,满了堵住 OnPublish(背压)。
type uplinkQueue struct {
b *Broker
endpointID string
ch chan uplinkItem
once sync.Once
}
func newUplinkQueue(b *Broker, endpointID string) *uplinkQueue {
q := &uplinkQueue{
b: b,
endpointID: endpointID,
ch: make(chan uplinkItem, uplinkQueueSize),
}
go q.loop()
return q
}
func (q *uplinkQueue) push(item uplinkItem) {
q.ch <- item // 满则阻塞读循环,形成背压
}
func (q *uplinkQueue) close() {
q.once.Do(func() { close(q.ch) })
}
func (q *uplinkQueue) loop() {
for item := range q.ch {
_ = q.b.uplink.HandleUplink(context.Background(), item.conn, item.payload)
}
}
+41
View File
@@ -0,0 +1,41 @@
package broker
import (
"context"
"net"
"net/http"
"git.asio.asia/nixevol/NixMsg/internal/listener"
"github.com/coder/websocket"
)
// WSHandler 返回 /mqtt 的 WebSocket 升级处理。
// Accept 时 InsecureSkipVerify=true;之后检查 Subprotocol==mqtt。
// NetConn 使用 Background 派生的 context,不用请求 Context。
func (b *Broker) WSHandler(proxies *listener.ProxySet) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
c, err := websocket.Accept(w, r, &websocket.AcceptOptions{
Subprotocols: []string{"mqtt"},
InsecureSkipVerify: true,
})
if err != nil {
return
}
if c.Subprotocol() != "mqtt" {
_ = c.Close(websocket.StatusPolicyViolation, "subprotocol must be mqtt")
return
}
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
nc := websocket.NetConn(ctx, c, websocket.MessageBinary)
if proxies != nil {
ip := proxies.ClientIP(r)
if ip != "" {
nc = listener.WithRemoteAddr(nc, &net.TCPAddr{IP: net.ParseIP(ip)})
}
}
_ = b.AttachWS(nc)
})
}
+62
View File
@@ -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:
}
}
+76
View File
@@ -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}
}
+398
View File
@@ -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
}
+101
View File
@@ -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)
}
+124
View File
@@ -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)
})
}
+338
View File
@@ -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"
}
+117
View File
@@ -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 的客户端仍能握手。
}
}