From 532ee44da36ce4e36923db5425a94d2153206571 Mon Sep 17 00:00:00 2001 From: Nixevol Date: Wed, 30 Sep 2026 06:51:12 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E5=AE=9E=E7=8E=B0=20WebSocket=E3=80=81?= =?UTF-8?q?mochi=20broker=20=E4=B8=8E=E4=B8=8B=E8=A1=8C=E5=8F=91=E5=B8=83?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/DEVIATIONS.md | 37 +++- go.mod | 4 + go.sum | 10 + internal/broker/broker.go | 379 +++++++++++++++++++++++++++++++++ internal/broker/broker_test.go | 259 ++++++++++++++++++++++ internal/broker/hooks.go | 196 +++++++++++++++++ internal/broker/queue.go | 45 ++++ internal/broker/ws.go | 41 ++++ 8 files changed, 970 insertions(+), 1 deletion(-) create mode 100644 internal/broker/broker.go create mode 100644 internal/broker/broker_test.go create mode 100644 internal/broker/hooks.go create mode 100644 internal/broker/queue.go create mode 100644 internal/broker/ws.go diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index ac733f3..edf2445 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -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 diff --git a/go.mod b/go.mod index c87fdec..67ef14f 100644 --- a/go.mod +++ b/go.mod @@ -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 diff --git a/go.sum b/go.sum index eceed2a..033d74f 100644 --- a/go.sum +++ b/go.sum @@ -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= diff --git a/internal/broker/broker.go b/internal/broker/broker.go new file mode 100644 index 0000000..41e103b --- /dev/null +++ b/internal/broker/broker.go @@ -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) +) diff --git a/internal/broker/broker_test.go b/internal/broker/broker_test.go new file mode 100644 index 0000000..69a66c2 --- /dev/null +++ b/internal/broker/broker_test.go @@ -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) + } +} diff --git a/internal/broker/hooks.go b/internal/broker/hooks.go new file mode 100644 index 0000000..1becf62 --- /dev/null +++ b/internal/broker/hooks.go @@ -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) +} diff --git a/internal/broker/queue.go b/internal/broker/queue.go new file mode 100644 index 0000000..3c1d7eb --- /dev/null +++ b/internal/broker/queue.go @@ -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) + } +} diff --git a/internal/broker/ws.go b/internal/broker/ws.go new file mode 100644 index 0000000..7aa827b --- /dev/null +++ b/internal/broker/ws.go @@ -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) + }) +}