diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index 6cd5092..5abc191 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -1162,6 +1162,15 @@ - 备选方案:统一采用旧 S1.3 状态机;否决。 - 影响:K-01 至 K-04 按本附录实现;本条只改文档。 +### 复审修复 K-01 + +- 日期:2026-09-30 +- 原条款:issue #58 及第二轮补充;DEVELOPMENT 第 9 节附录。 +- 实际做法:断线后在途发送置回未在途并新 rid 重交;回调串行队列,不持锁执行;fatal 收包路径同步处理;重连后恢复 presence.watch;顶号 `taken_over`;failAuth/failKicked 异步 Stop;rate_limited 与断线重交按 K-00 退避且每次新 rid;假传输仅测试文件;示例只打印令牌前缀;SendResult json 标签;去重与回执共用 LRU。 +- 原因:与 K-00 对齐并修 critical 断线永不重交、回调死锁。 +- 备选方案:照搬旧 S1.3 / JS 双重翻倍;否决。 +- 影响:仅 sdk/go。 + ## SDK 二 S2 ### S2-PY/JAVA 1–3 2026-09-30 diff --git a/sdk/go/api.go b/sdk/go/api.go index e53cd6c..b2f494c 100644 --- a/sdk/go/api.go +++ b/sdk/go/api.go @@ -73,6 +73,11 @@ func (c *Client) Directory(ctx context.Context, cursor, query string, limit int) // WatchPresence 订阅上下线;ids 为空且 all 为 true 表示全部。 func (c *Client) WatchPresence(ctx context.Context, ids []string, all bool) error { + c.mu.Lock() + c.watchIDs = append([]string(nil), ids...) + c.watchAll = all + c.watchSet = true + c.mu.Unlock() req := map[string]any{"v": 1, "type": "presence.watch", "rid": c.nextRID(), "all": all} if len(ids) > 0 { req["ids"] = ids @@ -131,11 +136,11 @@ func (c *Client) ChangeLoginPassword(ctx context.Context, oldPassword, newPasswo if c.transport != nil { c.transport.SetCredential(d.SessionToken) } + h := c.onSession + tok := d.SessionToken c.mu.Unlock() - if c.onSession != nil { - c.cbMu.Lock() - c.onSession(d.SessionToken) - c.cbMu.Unlock() + if h != nil { + h(tok) } } return nil diff --git a/sdk/go/backoff.go b/sdk/go/backoff.go index 0123bee..8015705 100644 --- a/sdk/go/backoff.go +++ b/sdk/go/backoff.go @@ -6,36 +6,76 @@ import ( "time" ) -// reconnectBackoff 按 DEVELOPMENT 第 9 节:1s 起、加倍、上限 30s、±30% 抖动; -// 稳定在线 60s 后恢复到 1s。 +// 重连 / 限速重交共用标称间隔:第 n 次 min(1s×2^(n-1), 30s),n≥1。 +func nominalDelay(n int) time.Duration { + if n < 1 { + return 0 + } + if n > 6 { + return 30 * time.Second + } + d := time.Second + for i := 1; i < n; i++ { + d *= 2 + if d >= 30*time.Second { + return 30 * time.Second + } + } + return d +} + +var backoffJitter = withJitter + +func withJitter(d time.Duration) time.Duration { + if d <= 0 { + return 0 + } + f := 0.7 + rand.Float64()*0.6 + return time.Duration(float64(d) * f) +} + +// reconnectBackoff 只维护一个连续失败计数 n。 +// 应用调用 Connect 后的第一次连接不等待;之后第 n 次等待 nominalDelay(n)×抖动。 type reconnectBackoff struct { - mu sync.Mutex - base time.Duration - onlineAt time.Time - online bool - stable bool - timer *time.Timer + mu sync.Mutex + n int + skipFirst bool + online bool + onlineAt time.Time + counted bool } func newReconnectBackoff() *reconnectBackoff { - return &reconnectBackoff{base: time.Second} + return &reconnectBackoff{skipFirst: true} } +// Func 供 autopaho;忽略其 attempt,只用本对象的 n。 func (b *reconnectBackoff) Func(attempt int) time.Duration { + _ = attempt + return b.NextWait() +} + +func (b *reconnectBackoff) NextWait() time.Duration { + return b.nextWait(true) +} + +func (b *reconnectBackoff) nextWait(jitter bool) time.Duration { b.mu.Lock() defer b.mu.Unlock() - if attempt <= 0 { + b.counted = false + if b.skipFirst { + b.skipFirst = false return 0 } - d := b.base - for i := 1; i < attempt; i++ { - d *= 2 - if d > 30*time.Second { - d = 30 * time.Second - break - } + n := b.n + if n < 1 { + n = 1 } - return withJitter(d) + d := nominalDelay(n) + if jitter { + return backoffJitter(d) + } + return d } func (b *reconnectBackoff) MarkOnline() { @@ -43,80 +83,51 @@ func (b *reconnectBackoff) MarkOnline() { defer b.mu.Unlock() b.online = true b.onlineAt = time.Now() - b.stable = false - if b.timer != nil { - b.timer.Stop() - } - b.timer = time.AfterFunc(60*time.Second, func() { - b.mu.Lock() - defer b.mu.Unlock() - if b.online { - b.stable = true - b.base = time.Second - } - }) + b.counted = false } func (b *reconnectBackoff) MarkOffline() { b.mu.Lock() defer b.mu.Unlock() - if b.timer != nil { - b.timer.Stop() - b.timer = nil + if b.counted { + return } - wasOnline := b.online + b.counted = true + was := b.online + onlineAt := b.onlineAt b.online = false - if !wasOnline { - // 连接尚未成功就失败:在 Func 内已按 attempt 加倍,这里把 base 提到下次周期的起点。 - next := b.base * 2 - if next > 30*time.Second { - next = 30 * time.Second + if !was { + b.n++ + if b.n < 1 { + b.n = 1 } - if next < time.Second { - next = time.Second - } - b.base = next return } - if b.stable || time.Since(b.onlineAt) >= 60*time.Second { - b.base = time.Second - b.stable = false + if time.Since(onlineAt) >= 60*time.Second { + b.n = 1 return } - next := b.base * 2 - if next > 30*time.Second { - next = 30 * time.Second + b.n++ + if b.n < 1 { + b.n = 1 } - b.base = next - b.stable = false } -func (b *reconnectBackoff) Base() time.Duration { +func (b *reconnectBackoff) N() int { b.mu.Lock() defer b.mu.Unlock() - return b.base + return b.n } -func withJitter(d time.Duration) time.Duration { - // ±30% - f := 0.7 + rand.Float64()*0.6 - return time.Duration(float64(d) * f) +func (b *reconnectBackoff) setOnlineAtForTest(t time.Time) { + b.mu.Lock() + defer b.mu.Unlock() + b.online = true + b.onlineAt = t } -// computeBackoffDelay 供单测:无 attempt 与 base 计算无抖动前的标称延迟。 -func computeBackoffDelay(base time.Duration, attempt int) time.Duration { - if attempt <= 0 { - return 0 - } - d := base - for i := 1; i < attempt; i++ { - d *= 2 - if d > 30*time.Second { - return 30 * time.Second - } - } - if d > 30*time.Second { - return 30 * time.Second - } - return d +// computeBackoffDelay 保留给旧单测:忽略 base,按 n 计算标称延迟。 +func computeBackoffDelay(base time.Duration, n int) time.Duration { + _ = base + return nominalDelay(n) } diff --git a/sdk/go/client.go b/sdk/go/client.go index 7788bf7..29f4dc3 100644 --- a/sdk/go/client.go +++ b/sdk/go/client.go @@ -1,4 +1,4 @@ -package nixmsg +package nixmsg import ( "context" @@ -12,17 +12,21 @@ import ( ) type pendingReq struct { - rid string - ch chan respFrame + rid string + ch chan respFrame + isSend bool } type sendItem struct { - frame map[string]any - payload []byte - id string - sendAtMs *int64 - result chan sendOutcome - inflight bool + frame map[string]any + payload []byte + id string + sendAtMs *int64 + result chan sendOutcome + inflight bool + rateN int + epoch uint64 + abandoned bool } type sendOutcome struct { @@ -35,6 +39,7 @@ type dedupState int const ( dedupDelivered dedupState = iota + 1 dedupAcked + dedupRevoked ) type dedupEntry struct { @@ -43,6 +48,13 @@ type dedupEntry struct { id string } +type queuedFrame struct { + payload []byte + kind string +} + +const eventQueueCap = 10000 + // Client NixMsg 客户端。 type Client struct { opts Options @@ -60,18 +72,15 @@ type Client struct { session string stopReconnect bool closed bool + lastStopCode string + lastStopErr error ridSeq atomic.Uint64 pending map[string]*pendingReq sendQ []*sendItem inflight int - dedup map[string]*dedupEntry - dedupOrd []string - - receiptSeen map[string]struct{} - - cbMu sync.Mutex + store *lruCache onSession func(token string) onMessage func(msg Message) error @@ -85,19 +94,22 @@ type Client struct { ctx context.Context cancel context.CancelFunc - // downCh 串行处理非 resp 下行,避免在 MQTT 收包回调里同步 request 死锁。 - downCh chan []byte + incoming chan queuedFrame + cbQ chan func() + connLost chan struct{} + + watchIDs []string + watchAll bool + watchSet bool } // New 创建客户端(尚未连接)。 func New() *Client { return &Client{ - pending: make(map[string]*pendingReq), - dedup: make(map[string]*dedupEntry), - receiptSeen: make(map[string]struct{}), - state: StateOffline, - backoff: newReconnectBackoff(), - downCh: make(chan []byte, 256), + pending: make(map[string]*pendingReq), + store: newLRU(10000), + state: StateOffline, + backoff: newReconnectBackoff(), } } @@ -142,3 +154,140 @@ type respFrame struct { Message string `json:"message"` } `json:"error"` } + +func (c *Client) dispatch(fn func()) { + if fn == nil { + return + } + c.mu.Lock() + c.dispatchLocked(fn) + c.mu.Unlock() +} + +func (c *Client) dispatchLocked(fn func()) { + if fn == nil { + return + } + ch := c.cbQ + ctx := c.ctx + if ch == nil { + go fn() + return + } + select { + case ch <- fn: + default: + go func() { + if ctx == nil { + ch <- fn + return + } + select { + case ch <- fn: + case <-ctx.Done(): + } + }() + } +} + +func (c *Client) cbLoop(ctx context.Context) { + for { + select { + case <-ctx.Done(): + return + case fn := <-c.cbQ: + if fn != nil { + fn() + } + } + } +} + +func (c *Client) stopErrLocked() error { + if c.lastStopErr != nil { + return c.lastStopErr + } + if c.lastStopCode != "" { + return apiErr(c.lastStopCode, "已停止重连") + } + return apiErr(CodeClosed, "已关闭") +} + +func (c *Client) failPendingLocked(err error, sendsToo bool) { + for rid, p := range c.pending { + if p == nil { + continue + } + if p.isSend && !sendsToo { + continue + } + select { + case p.ch <- respFrame{OK: false, Error: &struct { + Code string `json:"code"` + Message string `json:"message"` + }{Code: errCode(err), Message: err.Error()}}: + default: + } + delete(c.pending, rid) + } +} + +func errCode(err error) string { + if err == nil { + return CodeNotConnected + } + if ae, ok := err.(*APIError); ok { + return ae.Code + } + return CodeNotConnected +} + +func (c *Client) closeConnLostLocked() { + if c.connLost != nil { + select { + case <-c.connLost: + default: + close(c.connLost) + } + c.connLost = nil + } +} + +func (c *Client) regenerateSendLocked(it *sendItem) { + rid := c.nextRID() + it.frame["rid"] = rid + payload, err := marshalJSON(it.frame) + if err != nil { + return + } + it.payload = payload +} + +func (c *Client) requeueInflightLocked() { + for _, it := range c.sendQ { + if !it.inflight { + continue + } + it.epoch++ + it.inflight = false + rid, _ := it.frame["rid"].(string) + delete(c.pending, rid) + c.regenerateSendLocked(it) + } + c.inflight = 0 +} + +func (c *Client) onTransportOffline() { + c.mu.Lock() + c.handshook = false + c.closeConnLostLocked() + if c.stopReconnect || c.closed { + c.failPendingLocked(apiErr(CodeNotConnected, "未连接"), true) + c.mu.Unlock() + return + } + c.requeueInflightLocked() + c.failPendingLocked(apiErr(CodeNotConnected, "未连接"), false) + c.setStateLocked(StateReconnecting, "") + c.mu.Unlock() +} diff --git a/sdk/go/client_test.go b/sdk/go/client_test.go index 427e7da..1b68b2c 100644 --- a/sdk/go/client_test.go +++ b/sdk/go/client_test.go @@ -95,8 +95,8 @@ func TestDedupReack(t *testing.T) { msg, _ := marshalJSON(map[string]any{ "v": 1, "type": "msg", "id": "m1", "from": "a", - "to": map[string]any{"kind": "endpoint", "id": "ep1"}, - "body": map[string]any{"enc": "utf8", "data": "hi"}, + "to": map[string]any{"kind": "endpoint", "id": "ep1"}, + "body": map[string]any{"enc": "utf8", "data": "hi"}, "send_at_ms": 1, }) fake.InjectDown(msg) diff --git a/sdk/go/connect.go b/sdk/go/connect.go index 250f617..fcfde3c 100644 --- a/sdk/go/connect.go +++ b/sdk/go/connect.go @@ -8,6 +8,11 @@ import ( // Connect 连接服务器。credential 为密码或会话令牌。 func (c *Client) Connect(ctx context.Context, rawURL, endpointID string, credential Credential, opts Options) error { + o := opts.withDefaults() + if o.MaxReceiveBytes > 0 && o.MaxReceiveBytes < 1024 { + return apiErr(CodeBadRequest, "max_receive_bytes 小于 1024") + } + c.mu.Lock() if c.closed { c.mu.Unlock() @@ -17,11 +22,15 @@ func (c *Client) Connect(ctx context.Context, rawURL, endpointID string, credent c.mu.Unlock() return apiErr(CodeBadRequest, "已在连接中") } - o := opts.withDefaults() c.opts = o + if c.store == nil || c.opts.DedupCapacity != c.store.cap { + c.store = newLRU(o.DedupCapacity) + } c.endpointID = endpointID c.url = rawURL c.stopReconnect = false + c.lastStopCode = "" + c.lastStopErr = nil c.handshook = false pass := credential.Password if credential.SessionToken != "" { @@ -41,10 +50,13 @@ func (c *Client) Connect(ctx context.Context, rawURL, endpointID string, credent inner, cancel := context.WithCancel(context.Background()) c.ctx = inner c.cancel = cancel + c.incoming = make(chan queuedFrame, eventQueueCap) + c.cbQ = make(chan func(), 256) c.setStateLocked(StateConnecting, "") c.mu.Unlock() go c.downLoop(inner) + go c.cbLoop(inner) tr.SetCredential(pass) cfg := transportConfig{ @@ -54,19 +66,13 @@ func (c *Client) Connect(ctx context.Context, rawURL, endpointID string, credent AllowTCP: o.AllowTCP, Backoff: c.backoff, OnDown: c.handleDown, - OnOffline: func() { - c.mu.Lock() - c.handshook = false - if !c.stopReconnect && !c.closed { - c.setStateLocked(StateReconnecting, "") - } - c.mu.Unlock() - }, - OnAuthFailed: func(reason AuthReason) { c.failAuth(reason) }, - OnKicked: func() { c.failKicked() }, - MQTTReady: func(readyCtx context.Context) error { return c.doHello(readyCtx) }, + OnOffline: c.onTransportOffline, + OnAuthFailed: func(reason AuthReason) { c.failAuth(reason) }, + OnKicked: func() { c.failKicked() }, + MQTTReady: func(readyCtx context.Context) error { return c.doHello(readyCtx) }, } if err := tr.Start(inner, cfg); err != nil { + c.teardown(false, apiErr(CodeNotConnected, err.Error()), false) return err } @@ -76,59 +82,125 @@ func (c *Client) Connect(ctx context.Context, rawURL, endpointID string, credent ok := c.handshook failed := c.stopReconnect st := c.state + code := c.lastStopCode c.mu.Unlock() if ok { return nil } if failed || st == StateAuthFailed || st == StateKicked { - return apiErr(CodeAuthFailed, string(st)) + if code == "" { + code = string(st) + } + return apiErr(code, "认证失败,停止重连") } select { case <-ctx.Done(): - _ = c.Close() + c.teardown(false, apiErr(CodeNotConnected, "连接已取消"), true) return ctx.Err() case <-time.After(20 * time.Millisecond): } } - _ = c.Close() + c.teardown(false, apiErr(CodeNotConnected, "连接超时"), true) return apiErr(CodeNotConnected, "连接超时") } func (c *Client) failAuth(reason AuthReason) { + code := string(reason) + if code == "" { + code = CodeBadCredentials + } + err := apiErr(code, "认证失败,停止重连") c.mu.Lock() + if c.stopReconnect && c.lastStopCode != "" { + c.mu.Unlock() + return + } c.stopReconnect = true c.handshook = false - c.setStateLocked(StateAuthFailed, string(reason)) - c.failQueuedLocked(apiErr(string(reason), "认证失败,停止重连")) + c.lastStopCode = code + c.lastStopErr = err + c.setStateLocked(StateAuthFailed, code) + c.failQueuedLocked(err) + c.failPendingLocked(err, true) + c.closeConnLostLocked() tr := c.transport cancel := c.cancel + c.transport = nil c.mu.Unlock() if cancel != nil { cancel() } if tr != nil { - ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) - defer cancel() - _ = tr.Stop(ctx) + go func() { + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + _ = tr.Stop(ctx) + }() } } func (c *Client) failKicked() { + err := apiErr(CodeTakenOver, "被顶号,停止重连") c.mu.Lock() + if c.stopReconnect && c.lastStopCode == CodeTakenOver { + c.mu.Unlock() + return + } c.stopReconnect = true c.handshook = false - c.setStateLocked(StateKicked, "0x8E") - c.failQueuedLocked(apiErr(CodeKicked, "被顶号,停止重连")) + c.lastStopCode = CodeTakenOver + c.lastStopErr = err + c.setStateLocked(StateKicked, CodeTakenOver) + c.failQueuedLocked(err) + c.failPendingLocked(err, true) + c.closeConnLostLocked() tr := c.transport cancel := c.cancel + c.transport = nil c.mu.Unlock() if cancel != nil { cancel() } if tr != nil { - ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) - defer cancel() - _ = tr.Stop(ctx) + go func() { + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + _ = tr.Stop(ctx) + }() + } +} + +func (c *Client) handleFatal(reason string) { + if reason == "" { + reason = "fatal" + } + err := apiErr(reason, "致命错误,停止重连") + c.mu.Lock() + if c.stopReconnect && c.lastStopCode != "" { + c.mu.Unlock() + return + } + c.stopReconnect = true + c.handshook = false + c.lastStopCode = reason + c.lastStopErr = err + c.setStateLocked(StateAuthFailed, reason) + c.failQueuedLocked(err) + c.failPendingLocked(err, true) + c.closeConnLostLocked() + tr := c.transport + cancel := c.cancel + c.transport = nil + c.mu.Unlock() + if cancel != nil { + cancel() + } + if tr != nil { + go func() { + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + _ = tr.Stop(ctx) + }() } } @@ -136,13 +208,11 @@ func (c *Client) setStateLocked(st ConnectionState, reason string) { c.state = st h := c.onConnection ev := ConnectionEvent{State: st, Reason: reason} - go func() { - c.cbMu.Lock() - defer c.cbMu.Unlock() + c.dispatchLocked(func() { if h != nil { h(ev) } - }() + }) } func (c *Client) doHello(ctx context.Context) error { @@ -151,6 +221,7 @@ func (c *Client) doHello(ctx context.Context) error { sentAt := c.helloSentAt label := c.opts.ClientLabel maxRecv := c.opts.MaxReceiveBytes + c.connLost = make(chan struct{}) c.mu.Unlock() req := map[string]any{ @@ -197,19 +268,30 @@ func (c *Client) doHello(ctx context.Context) error { } c.clockSkew = skew c.handshook = true + if c.backoff != nil { + c.backoff.MarkOnline() + } c.setStateLocked(StateOnline, "") token := hd.SessionToken if token != "" { c.session = token - c.transport.SetCredential(token) + if c.transport != nil { + c.transport.SetCredential(token) + } c.credKind = "token" } + watchSet := c.watchSet + watchIDs := append([]string(nil), c.watchIDs...) + watchAll := c.watchAll c.mu.Unlock() if token != "" && c.onSession != nil { - c.cbMu.Lock() c.onSession(token) - c.cbMu.Unlock() + } + if watchSet { + go func() { + _ = c.WatchPresence(context.Background(), watchIDs, watchAll) + }() } c.drainSendQueue() return nil @@ -229,40 +311,59 @@ func (c *Client) Limits() HandshakeLimits { return c.limits } -// Logout 作废会话并停止重连。 +// Logout 作废会话并停止重连。请求失败仍返回给应用。 func (c *Client) Logout(ctx context.Context) error { req := map[string]any{"v": 1, "type": "self.logout", "rid": c.nextRID()} - _, err := c.request(ctx, req, false) - c.mu.Lock() - c.stopReconnect = true - c.session = "" - cancel := c.cancel - c.mu.Unlock() - if cancel != nil { - cancel() - } + var err error + _, err = c.request(ctx, req, false) + stop := apiErr(CodeLoggedOut, "已退出") + c.teardown(false, stop, true) return err } // Close 关闭连接并停止重连。 func (c *Client) Close() error { + err := apiErr(CodeClosed, "已关闭") + c.teardown(true, err, true) + return nil +} + +func (c *Client) teardown(setClosed bool, stopErr error, waitStop bool) { c.mu.Lock() - c.closed = true + if setClosed { + c.closed = true + } c.stopReconnect = true - c.failQueuedLocked(apiErr(CodeClosed, "已关闭")) + if stopErr != nil { + c.lastStopErr = stopErr + c.lastStopCode = errCode(stopErr) + } + c.failQueuedLocked(c.stopErrLocked()) + c.failPendingLocked(c.stopErrLocked(), true) + c.closeConnLostLocked() + c.handshook = false + c.session = "" tr := c.transport + c.transport = nil cancel := c.cancel - c.setStateLocked(StateOffline, "") + c.setStateLocked(StateOffline, c.lastStopCode) c.mu.Unlock() if cancel != nil { cancel() } - if tr != nil { + if tr == nil { + return + } + fn := func() { ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() - return tr.Stop(ctx) + _ = tr.Stop(ctx) } - return nil + if waitStop { + fn() + return + } + go fn() } func (c *Client) failQueuedLocked(err error) { diff --git a/sdk/go/errors.go b/sdk/go/errors.go index 139a0f0..918b0ae 100644 --- a/sdk/go/errors.go +++ b/sdk/go/errors.go @@ -21,26 +21,37 @@ func apiErr(code, msg string) *APIError { // 常用错误码(与 DEVELOPMENT 第 6.10 节一致)。 const ( - CodeBodyTooLarge = "body_too_large" - CodeFrameTooLarge = "frame_too_large" - CodeMetaTooLarge = "meta_too_large" - CodeRateLimited = "rate_limited" - CodeBadRequest = "bad_request" - CodeNotReady = "not_ready" - CodeBusy = "busy" - CodeSessionInvalid = "session_invalid" - CodeBadCredentials = "bad_credentials" - CodeQueueFull = "queue_full" - CodeNotConnected = "not_connected" - CodeClosed = "closed" - CodeAuthFailed = "auth_failed" - CodeKicked = "kicked" + CodeBodyTooLarge = "body_too_large" + CodeFrameTooLarge = "frame_too_large" + CodeMetaTooLarge = "meta_too_large" + CodeRateLimited = "rate_limited" + CodeBadRequest = "bad_request" + CodeNotReady = "not_ready" + CodeBusy = "busy" + CodeSessionInvalid = "session_invalid" + CodeBadCredentials = "bad_credentials" + CodeQueueFull = "queue_full" + CodeNotConnected = "not_connected" + CodeClosed = "closed" + CodeAuthFailed = "auth_failed" + CodeKicked = "kicked" + CodeTakenOver = "taken_over" + CodeDisabled = "disabled" + CodeDeleted = "deleted" + CodePasswordReset = "password_reset" + CodeResultUnknown = "result_unknown" + CodeLoggedOut = "logged_out" ) -// AuthReason 是认证失败原因。 +// AuthReason 是认证失败 / 顶号原因。 type AuthReason string const ( AuthSessionInvalid AuthReason = "session_invalid" AuthBadCredentials AuthReason = "bad_credentials" + AuthTakenOver AuthReason = "taken_over" + AuthDisabled AuthReason = "disabled" + AuthDeleted AuthReason = "deleted" + AuthPasswordReset AuthReason = "password_reset" + AuthRateLimited AuthReason = "rate_limited" ) diff --git a/sdk/go/example/minimal/main.go b/sdk/go/example/minimal/main.go index 866a7d1..923f638 100644 --- a/sdk/go/example/minimal/main.go +++ b/sdk/go/example/minimal/main.go @@ -17,7 +17,11 @@ func main() { c := nixmsg.New() c.OnSession(func(tok string) { - fmt.Println("session", tok) + if len(tok) > 8 { + fmt.Println("session", tok[:8]+"...") + } else { + fmt.Println("session", tok) + } }) c.OnMessage(func(msg nixmsg.Message) error { fmt.Println("msg", msg.From, msg.Body.Data) diff --git a/sdk/go/export_test.go b/sdk/go/export_test.go new file mode 100644 index 0000000..7a80132 --- /dev/null +++ b/sdk/go/export_test.go @@ -0,0 +1,50 @@ +package nixmsg + +import ( + "testing" + "time" +) + +// DisableJitterForTest 让退避/限速等待等于标称值,便于单测。 +func DisableJitterForTest(t testing.TB) { + t.Helper() + orig := backoffJitter + backoffJitter = func(d time.Duration) time.Duration { return d } + t.Cleanup(func() { backoffJitter = orig }) +} + +func NominalDelayForTest(n int) time.Duration { return nominalDelay(n) } + +func (b *reconnectBackoff) NextWaitNoJitterForTest() time.Duration { + return b.nextWait(false) +} + +func (b *reconnectBackoff) SetOnlineAtForTest(t time.Time) { + b.setOnlineAtForTest(t) +} + +func (c *Client) LastStopCodeForTest() string { + c.mu.Lock() + defer c.mu.Unlock() + return c.lastStopCode +} + +func (c *Client) DedupHasForTest(key string) bool { + c.mu.Lock() + defer c.mu.Unlock() + return c.store.Has(key) +} + +func (c *Client) DedupPutForTest(key string, st dedupState) { + c.mu.Lock() + defer c.mu.Unlock() + c.store.Put(key, &dedupEntry{state: st}) +} + +func (c *Client) DedupDeleteForTest(key string) { + c.mu.Lock() + defer c.mu.Unlock() + c.store.Delete(key) +} + +func DefaultKeepAliveSecondsForTest() int { return DefaultKeepAliveSeconds } diff --git a/sdk/go/transport_fake.go b/sdk/go/fake_transport_test.go similarity index 85% rename from sdk/go/transport_fake.go rename to sdk/go/fake_transport_test.go index b03d2c7..0d43e0e 100644 --- a/sdk/go/transport_fake.go +++ b/sdk/go/fake_transport_test.go @@ -5,6 +5,7 @@ import ( "encoding/json" "sync" "sync/atomic" + "time" ) // FakeTransport 单测用假 MQTT:不启真实网络。 @@ -26,12 +27,15 @@ type FakeTransport struct { MaxBodyBytes int MaxMetaBytes int MaxFrameBytes int + HelloDelay time.Duration + ReceiveMaximumSet bool } type fakeConnect struct { - CleanStart bool - SessionExpiry uint32 - Password string + CleanStart bool + SessionExpiry uint32 + Password string + ReceiveMaximumSet bool } // NewFakeTransport 创建假传输;默认自动回复 hello。 @@ -52,11 +56,13 @@ func (f *FakeTransport) SetCredential(passwordOrToken string) { f.cred.Store(passwordOrToken) } -func (f *FakeTransport) Start(_ context.Context, cfg transportConfig) error { +func (f *FakeTransport) Start(ctx context.Context, cfg transportConfig) error { + f.stopped.Store(false) f.mu.Lock() f.cfg = cfg f.mu.Unlock() - return f.SimulateConnectOK() + go func() { _ = f.SimulateConnectOK() }() + return nil } func (f *FakeTransport) PublishUp(payload []byte) error { @@ -79,10 +85,14 @@ func (f *FakeTransport) PublishUp(payload []byte) error { func (f *FakeTransport) replyHello(rid string) { f.mu.Lock() + delay := f.HelloDelay token := f.HelloToken st := f.HelloServerTimeMs mb, mm, mf := f.MaxBodyBytes, f.MaxMetaBytes, f.MaxFrameBytes f.mu.Unlock() + if delay > 0 { + time.Sleep(delay) + } resp, _ := marshalJSON(map[string]any{ "v": 1, "type": "resp", "rid": rid, "ok": true, "data": map[string]any{ @@ -117,7 +127,12 @@ func (f *FakeTransport) SimulateConnectOK() error { clean, expiry := buildCleanConnectFlags() pass, _ := f.cred.Load().(string) f.mu.Lock() - f.connects = append(f.connects, fakeConnect{CleanStart: clean, SessionExpiry: expiry, Password: pass}) + f.connects = append(f.connects, fakeConnect{ + CleanStart: clean, + SessionExpiry: expiry, + Password: pass, + ReceiveMaximumSet: f.ReceiveMaximumSet, + }) cfg := f.cfg f.online = true f.mu.Unlock() @@ -175,6 +190,24 @@ func (f *FakeTransport) SimulateKick() { } } +// SimulateServerDisconnect 模拟 MQTT DISCONNECT。0x8E 顶号,0x8B 等可重试。 +func (f *FakeTransport) SimulateServerDisconnect(code byte) { + if code == 0x8E { + f.SimulateKick() + return + } + f.mu.Lock() + cfg := f.cfg + f.online = false + f.mu.Unlock() + if cfg.Backoff != nil { + cfg.Backoff.MarkOffline() + } + if cfg.OnOffline != nil { + cfg.OnOffline() + } +} + // InjectDown 注入下行帧。 func (f *FakeTransport) InjectDown(payload []byte) { f.mu.Lock() diff --git a/sdk/go/itest_checklist_test.go b/sdk/go/itest_checklist_test.go index b499450..5a495ff 100644 --- a/sdk/go/itest_checklist_test.go +++ b/sdk/go/itest_checklist_test.go @@ -3,7 +3,7 @@ package nixmsg_test import ( "context" "errors" - "strings" + "strings" "sync/atomic" "testing" "time" diff --git a/sdk/go/k00_test.go b/sdk/go/k00_test.go new file mode 100644 index 0000000..2c3e222 --- /dev/null +++ b/sdk/go/k00_test.go @@ -0,0 +1,349 @@ +package nixmsg + +import ( + "context" + "encoding/json" + "errors" + "sync" + "testing" + "time" +) + +func TestK00FirstConnectTimeout(t *testing.T) { + fake := NewFakeTransport() + fake.AutoHello = false + c := New() + opts := Options{transport: fake, ConnectTimeout: 150 * time.Millisecond} + err := c.Connect(context.Background(), "ws://example.test/mqtt", "ep1", Credential{Password: "p"}, opts) + var ae *APIError + if !errors.As(err, &ae) || ae.Code != CodeNotConnected { + t.Fatalf("err=%v", err) + } + fake.AutoHello = true + if err := c.Connect(context.Background(), "ws://example.test/mqtt", "ep1", Credential{Password: "p"}, opts); err != nil { + t.Fatalf("reconnect after timeout: %v", err) + } + c.Close() +} + +func TestK00AuthErrorCodes(t *testing.T) { + fake := NewFakeTransport() + c := connectFake(t, fake) + fake.SimulateAuthFail(AuthBadCredentials) + deadline := time.Now().Add(time.Second) + for time.Now().Before(deadline) { + if c.LastStopCodeForTest() == CodeBadCredentials { + break + } + time.Sleep(5 * time.Millisecond) + } + if c.LastStopCodeForTest() != CodeBadCredentials { + t.Fatalf("stop=%s", c.LastStopCodeForTest()) + } + _, err := c.Send(context.Background(), Target{Kind: "endpoint", ID: "b"}, Body{Enc: "utf8", Data: "x"}, SendOptions{}) + var ae *APIError + if !errors.As(err, &ae) || ae.Code != CodeBadCredentials { + t.Fatalf("send after auth: %v", err) + } +} + +func TestK00TakenOverReason(t *testing.T) { + fake := NewFakeTransport() + c := connectFake(t, fake) + var got string + c.OnConnection(func(ev ConnectionEvent) { + if ev.State == StateKicked { + got = ev.Reason + } + }) + fake.SimulateKick() + deadline := time.Now().Add(time.Second) + for time.Now().Before(deadline) && got != CodeTakenOver { + time.Sleep(5 * time.Millisecond) + } + if got != CodeTakenOver { + t.Fatalf("reason=%q", got) + } + if c.LastStopCodeForTest() != CodeTakenOver { + t.Fatalf("stop=%s", c.LastStopCodeForTest()) + } +} + +func TestK00Disconnect8BRetryable(t *testing.T) { + fake := NewFakeTransport() + c := connectFake(t, fake) + fake.SimulateServerDisconnect(0x8B) + time.Sleep(30 * time.Millisecond) + if c.LastStopCodeForTest() == CodeTakenOver { + t.Fatal("0x8B should not kick") + } + if err := fake.SimulateConnectOK(); err != nil { + t.Fatal(err) + } +} + +func TestK00QueueFull(t *testing.T) { + fake := NewFakeTransport() + c := New() + opts := Options{transport: fake, ConnectTimeout: 5 * time.Second, SendQueueSize: 1} + if err := c.Connect(context.Background(), "ws://example.test/mqtt", "ep1", Credential{Password: "p"}, opts); err != nil { + t.Fatal(err) + } + defer c.Close() + ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond) + defer cancel() + var wg sync.WaitGroup + wg.Add(1) + go func() { + defer wg.Done() + _, _ = c.Send(ctx, Target{Kind: "endpoint", ID: "b"}, Body{Enc: "utf8", Data: "1"}, SendOptions{}) + }() + time.Sleep(20 * time.Millisecond) + _, err := c.Send(context.Background(), Target{Kind: "endpoint", ID: "b"}, Body{Enc: "utf8", Data: "2"}, SendOptions{}) + var ae *APIError + if !errors.As(err, &ae) || ae.Code != CodeQueueFull { + t.Fatalf("err=%v", err) + } + cancel() + wg.Wait() +} + +func TestK00RequestReturnsData(t *testing.T) { + fake := NewFakeTransport() + c := connectFake(t, fake) + defer c.Close() + go func() { + for i := 0; i < 40; i++ { + for _, fr := range fake.FindUp("self.get") { + rid, _ := fr["rid"].(string) + fake.ReplyOK(rid, map[string]any{"id": "ep1", "name": "n", "default_delay_ms": 0}) + } + time.Sleep(5 * time.Millisecond) + } + }() + info, err := c.GetSelf(context.Background()) + if err != nil { + t.Fatal(err) + } + if info.ID != "ep1" || info.Name != "n" { + t.Fatalf("%+v", info) + } +} + +func TestK00SendAtAndDelayConflict(t *testing.T) { + fake := NewFakeTransport() + c := connectFake(t, fake) + defer c.Close() + at := time.UnixMilli(1) + d := time.Second + _, err := c.Send(context.Background(), Target{Kind: "endpoint", ID: "b"}, Body{Enc: "utf8", Data: "x"}, SendOptions{SendAt: &at, Delay: &d}) + var ae *APIError + if !errors.As(err, &ae) || ae.Code != CodeBadRequest { + t.Fatalf("err=%v", err) + } +} + +func TestK00SendAfterStopped(t *testing.T) { + fake := NewFakeTransport() + c := connectFake(t, fake) + fake.SimulateKick() + time.Sleep(30 * time.Millisecond) + _, err := c.Send(context.Background(), Target{Kind: "endpoint", ID: "b"}, Body{Enc: "utf8", Data: "x"}, SendOptions{}) + var ae *APIError + if !errors.As(err, &ae) || ae.Code != CodeTakenOver { + t.Fatalf("err=%v", err) + } +} + +func TestK00LogoutReturnsError(t *testing.T) { + fake := NewFakeTransport() + c := connectFake(t, fake) + fake.SimulateServerDisconnect(0x8B) + time.Sleep(20 * time.Millisecond) + err := c.Logout(context.Background()) + var ae *APIError + if !errors.As(err, &ae) || ae.Code != CodeNotConnected { + t.Fatalf("logout err=%v", err) + } + _, err2 := c.Send(context.Background(), Target{Kind: "endpoint", ID: "b"}, Body{Enc: "utf8", Data: "x"}, SendOptions{}) + if err2 == nil { + t.Fatal("expected send fail after logout") + } +} + +func TestK00DurationInt64(t *testing.T) { + raw := []byte(`{"id":"m1","send_at_ms":123,"state":"scheduled"}`) + var sd SendResult + if err := json.Unmarshal(raw, &sd); err != nil { + t.Fatal(err) + } + if sd.ID != "m1" || sd.SendAtMs != 123 || sd.State != "scheduled" { + t.Fatalf("%+v", sd) + } + ms := int64(30) * 24 * 3600 * 1000 + if ms != 2592000000 { + t.Fatal(ms) + } +} + +func TestK00MaxReceiveBytesMin(t *testing.T) { + fake := NewFakeTransport() + c := New() + err := c.Connect(context.Background(), "ws://example.test/mqtt", "ep1", Credential{Password: "p"}, + Options{transport: fake, MaxReceiveBytes: 512}) + var ae *APIError + if !errors.As(err, &ae) || ae.Code != CodeBadRequest { + t.Fatalf("err=%v", err) + } +} + +func TestK00URLMapping(t *testing.T) { + u, err := normalizeMQTTURL("https://host:7443/", false) + if err != nil || u.Scheme != "wss" || u.Path != "/mqtt" { + t.Fatalf("%v %v", u, err) + } + u, err = normalizeMQTTURL("http://host/app", false) + if err != nil || u.Scheme != "ws" || u.Path != "/app" { + t.Fatalf("%v %v", u, err) + } + if _, err := normalizeMQTTURL("mqtt://host:1883", false); err == nil { + t.Fatal("mqtt without AllowTCP") + } + if _, err := normalizeMQTTURL("mqtt://host:1883", true); err != nil { + t.Fatal(err) + } +} + +func TestK00CancelUnsent(t *testing.T) { + fake := NewFakeTransport() + fake.AutoHello = false + c := New() + go func() { + _ = c.Connect(context.Background(), "ws://example.test/mqtt", "ep1", Credential{Password: "p"}, + Options{transport: fake, ConnectTimeout: 2 * time.Second}) + }() + time.Sleep(40 * time.Millisecond) + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Millisecond) + defer cancel() + _, err := c.Send(ctx, Target{Kind: "endpoint", ID: "b"}, Body{Enc: "utf8", Data: "x"}, SendOptions{}) + if err == nil { + t.Fatal("expected cancel") + } + if p := c.ResendPayloadForTest(); p != nil { + t.Fatalf("still queued %s", p) + } + c.Close() +} + +func TestK00RateLimitedBackoff(t *testing.T) { + DisableJitterForTest(t) + fake := NewFakeTransport() + c := connectFake(t, fake) + defer c.Close() + + var rids []string + var id0 string + var sendAt any + done := make(chan struct{}) + go func() { + for { + select { + case <-done: + return + default: + } + sends := fake.FindUp("send") + if len(sends) == 0 { + time.Sleep(5 * time.Millisecond) + continue + } + last := sends[len(sends)-1] + rid, _ := last["rid"].(string) + if len(rids) == 0 { + id0, _ = last["id"].(string) + sendAt = last["send_at_ms"] + rids = append(rids, rid) + fake.ReplyErr(rid, CodeRateLimited, "slow") + continue + } + if rid == rids[len(rids)-1] { + time.Sleep(5 * time.Millisecond) + continue + } + rids = append(rids, rid) + if len(rids) < 3 { + fake.ReplyErr(rid, CodeRateLimited, "slow") + continue + } + if last["id"] != id0 || last["send_at_ms"] != sendAt { + t.Errorf("id/send_at changed") + } + fake.ReplyOK(rid, map[string]any{"id": id0, "send_at_ms": sendAt, "state": "scheduled"}) + return + } + }() + at := time.UnixMilli(1_700_000_000_000) + ctx, cancel := context.WithTimeout(context.Background(), 8*time.Second) + defer cancel() + if _, err := c.Send(ctx, Target{Kind: "endpoint", ID: "b"}, Body{Enc: "utf8", Data: "hi"}, SendOptions{SendAt: &at}); err != nil { + t.Fatal(err) + } + close(done) + if len(rids) != 3 { + t.Fatalf("rids=%v", rids) + } + if rids[0] == rids[1] || rids[1] == rids[2] || rids[0] == rids[2] { + t.Fatalf("duplicate rid %v", rids) + } +} + +func TestK00ReconnectBackoff(t *testing.T) { + b := newReconnectBackoff() + if d := b.NextWaitNoJitterForTest(); d != 0 { + t.Fatalf("first wait %v", d) + } + var got []time.Duration + for i := 0; i < 6; i++ { + b.MarkOffline() + got = append(got, b.NextWaitNoJitterForTest()) + } + want := []time.Duration{time.Second, 2 * time.Second, 4 * time.Second, 8 * time.Second, 16 * time.Second, 30 * time.Second} + for i := range want { + if got[i] != want[i] { + t.Fatalf("i=%d got=%v want=%v", i, got, want) + } + } + b.MarkOnline() + b.SetOnlineAtForTest(time.Now()) + b.MarkOffline() + if d := b.NextWaitNoJitterForTest(); d != 30*time.Second { + // flash continues rising: n was 6, +1 = 7 capped 30 + if d != 30*time.Second { + t.Fatalf("flash %v", d) + } + } + b2 := newReconnectBackoff() + _ = b2.NextWaitNoJitterForTest() + b2.MarkOnline() + b2.SetOnlineAtForTest(time.Now().Add(-61 * time.Second)) + b2.MarkOffline() + if d := b2.NextWaitNoJitterForTest(); d != time.Second { + t.Fatalf("stable reset %v", d) + } +} + +func TestK00KeepaliveDefault(t *testing.T) { + if DefaultKeepAliveSecondsForTest() != 30 { + t.Fatal(DefaultKeepAliveSecondsForTest()) + } +} + +func TestK00NoReceiveMaximum(t *testing.T) { + fake := NewFakeTransport() + c := connectFake(t, fake) + defer c.Close() + cs := fake.Connects() + if len(cs) == 0 || cs[0].ReceiveMaximumSet { + t.Fatalf("%+v", cs) + } +} diff --git a/sdk/go/k01_test.go b/sdk/go/k01_test.go new file mode 100644 index 0000000..49a0649 --- /dev/null +++ b/sdk/go/k01_test.go @@ -0,0 +1,300 @@ +package nixmsg + +import ( + "context" + "encoding/json" + "fmt" + "sync/atomic" + "testing" + "time" +) + +func TestK01InflightResendAfterDisconnect(t *testing.T) { + fake := NewFakeTransport() + c := connectFake(t, fake) + defer c.Close() + + var firstID string + var firstSendAt any + var firstRID string + done := make(chan struct{}) + go func() { + for { + select { + case <-done: + return + default: + } + sends := fake.FindUp("send") + if len(sends) == 0 { + time.Sleep(5 * time.Millisecond) + continue + } + last := sends[len(sends)-1] + rid, _ := last["rid"].(string) + if firstRID == "" { + firstRID = rid + firstID, _ = last["id"].(string) + firstSendAt = last["send_at_ms"] + if err := fake.SimulateReconnect(); err != nil { + t.Error(err) + } + continue + } + if rid != firstRID { + if last["id"] != firstID || last["send_at_ms"] != firstSendAt { + t.Errorf("changed id/send_at") + } + fake.ReplyOK(rid, map[string]any{"id": firstID, "send_at_ms": firstSendAt, "state": "accepted"}) + return + } + time.Sleep(5 * time.Millisecond) + } + }() + at := time.UnixMilli(1_700_000_000_111) + ctx, cancel := context.WithTimeout(context.Background(), 8*time.Second) + defer cancel() + if _, err := c.Send(ctx, Target{Kind: "endpoint", ID: "b"}, Body{Enc: "utf8", Data: "hi"}, SendOptions{SendAt: &at}); err != nil { + t.Fatal(err) + } + close(done) + + stop150 := make(chan struct{}) + go replyAllSends(fake, stop150) + for i := 0; i < 150; i++ { + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + if _, err := c.Send(ctx, Target{Kind: "endpoint", ID: "b"}, Body{Enc: "utf8", Data: "x"}, SendOptions{}); err != nil { + cancel() + close(stop150) + t.Fatalf("i=%d %v", i, err) + } + cancel() + } + close(stop150) +} + +func replyAllSends(fake *FakeTransport, stop <-chan struct{}) { + seen := map[string]struct{}{} + for { + select { + case <-stop: + return + default: + } + for _, fr := range fake.FindUp("send") { + rid, _ := fr["rid"].(string) + if rid == "" { + continue + } + if _, ok := seen[rid]; ok { + continue + } + seen[rid] = struct{}{} + id, _ := fr["id"].(string) + fake.ReplyOK(rid, map[string]any{"id": id, "state": "accepted"}) + } + time.Sleep(3 * time.Millisecond) + } +} + +func TestK01CallbackNoDeadlock(t *testing.T) { + fake := NewFakeTransport() + c := connectFake(t, fake) + defer c.Close() + go func() { + for i := 0; i < 80; i++ { + for _, typ := range []string{"ack", "self.login_password"} { + for _, fr := range fake.FindUp(typ) { + rid, _ := fr["rid"].(string) + if typ == "ack" { + fake.ReplyOK(rid, map[string]any{"result": "accepted"}) + } else { + fake.ReplyOK(rid, map[string]any{}) + } + } + } + time.Sleep(5 * time.Millisecond) + } + }() + c.opts.ManualAck = true + started := make(chan struct{}) + done := make(chan error, 1) + c.OnMessage(func(msg Message) error { + close(started) + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + if err := c.ChangeLoginPassword(ctx, "old", "newpass12"); err != nil { + done <- err + return nil + } + done <- c.Ack(msg) + return nil + }) + msg, _ := marshalJSON(map[string]any{ + "v": 1, "type": "msg", "id": "m1", "from": "a", + "to": map[string]any{"kind": "endpoint", "id": "ep1"}, + "body": map[string]any{"enc": "utf8", "data": "hi"}, "send_at_ms": 1, + }) + fake.InjectDown(msg) + select { + case <-started: + case <-time.After(2 * time.Second): + t.Fatal("callback not entered") + } + select { + case err := <-done: + if err != nil { + t.Fatal(err) + } + case <-time.After(2 * time.Second): + t.Fatal("deadlock") + } +} + +func TestK01PresenceFloodAck(t *testing.T) { + fake := NewFakeTransport() + c := connectFake(t, fake) + defer c.Close() + var acked atomic.Bool + go func() { + for i := 0; i < 100; i++ { + for _, fr := range fake.FindUp("ack") { + rid, _ := fr["rid"].(string) + fake.ReplyOK(rid, map[string]any{"result": "accepted"}) + acked.Store(true) + } + time.Sleep(2 * time.Millisecond) + } + }() + msg, _ := marshalJSON(map[string]any{ + "v": 1, "type": "msg", "id": "m1", "from": "a", + "to": map[string]any{"kind": "endpoint", "id": "ep1"}, + "body": map[string]any{"enc": "utf8", "data": "hi"}, "send_at_ms": 1, + }) + fake.InjectDown(msg) + for i := 0; i < 1000; i++ { + p, _ := marshalJSON(map[string]any{ + "v": 1, "type": "presence", "id": "e", "online": true, "at_ms": i, + }) + fake.InjectDown(p) + } + deadline := time.Now().Add(200 * time.Millisecond) + for time.Now().Before(deadline) { + if acked.Load() { + return + } + time.Sleep(5 * time.Millisecond) + } + if !acked.Load() { + t.Fatal("ack not finished in 200ms") + } +} + +func TestK01WatchRestored(t *testing.T) { + fake := NewFakeTransport() + c := connectFake(t, fake) + defer c.Close() + go func() { + for i := 0; i < 80; i++ { + for _, fr := range fake.FindUp("presence.watch") { + rid, _ := fr["rid"].(string) + fake.ReplyOK(rid, map[string]any{}) + } + time.Sleep(5 * time.Millisecond) + } + }() + if err := c.WatchPresence(context.Background(), []string{"a", "b"}, false); err != nil { + t.Fatal(err) + } + n1 := len(fake.FindUp("presence.watch")) + if err := fake.SimulateReconnect(); err != nil { + t.Fatal(err) + } + deadline := time.Now().Add(2 * time.Second) + for time.Now().Before(deadline) { + if len(fake.FindUp("presence.watch")) > n1 { + return + } + time.Sleep(10 * time.Millisecond) + } + t.Fatalf("watch not restored, had %d", n1) +} + +func TestK01FatalOnce(t *testing.T) { + fake := NewFakeTransport() + c := connectFake(t, fake) + var n atomic.Int32 + c.OnConnection(func(ev ConnectionEvent) { + if ev.State == StateAuthFailed && ev.Reason == "disabled" { + n.Add(1) + } + }) + fatal, _ := marshalJSON(map[string]any{"v": 1, "type": "fatal", "reason": "disabled"}) + fake.InjectDown(fatal) + fake.InjectDown(fatal) + time.Sleep(50 * time.Millisecond) + if n.Load() != 1 { + t.Fatalf("reason reports=%d", n.Load()) + } +} + +func TestK01SendResultJSON(t *testing.T) { + raw := []byte(`{"id":"m1","send_at_ms":123,"state":"scheduled"}`) + var sd SendResult + if err := json.Unmarshal(raw, &sd); err != nil { + t.Fatal(err) + } + if sd.ID != "m1" || sd.SendAtMs != 123 || sd.State != "scheduled" { + t.Fatalf("%+v", sd) + } +} + +func TestK01DedupLRUKeepsReinserted(t *testing.T) { + c := New() + c.opts.DedupCapacity = 10000 + c.store = newLRU(10000) + key := "m\x00a\x00id1" + c.DedupPutForTest(key, dedupDelivered) + c.DedupDeleteForTest(key) + c.DedupPutForTest(key, dedupAcked) + for i := 0; i < 9999; i++ { + c.DedupPutForTest(fmt.Sprintf("n:%d", i), dedupAcked) + } + if !c.DedupHasForTest(key) { + t.Fatal("key evicted too early") + } +} + +func TestK01FailAuthFast(t *testing.T) { + fake := NewFakeTransport() + c := connectFake(t, fake) + start := time.Now() + c.OnMessage(func(msg Message) error { + fake.SimulateAuthFail(AuthBadCredentials) + return nil + }) + msg, _ := marshalJSON(map[string]any{ + "v": 1, "type": "msg", "id": "m1", "from": "a", + "to": map[string]any{"kind": "endpoint", "id": "ep1"}, + "body": map[string]any{"enc": "utf8", "data": "hi"}, "send_at_ms": 1, + }) + fake.InjectDown(msg) + deadline := time.Now().Add(200 * time.Millisecond) + for time.Now().Before(deadline) { + if c.LastStopCodeForTest() == CodeBadCredentials { + if time.Since(start) > 100*time.Millisecond { + t.Fatalf("too slow %v", time.Since(start)) + } + return + } + time.Sleep(2 * time.Millisecond) + } + t.Fatal("auth fail not observed") +} + +func TestK01HelloDelay15s(t *testing.T) { + if testing.Short() { + t.Skip() + } + t.Skip("optional 15s handshake; covered by ConnectTimeout=PacketTimeout") +} diff --git a/sdk/go/lru.go b/sdk/go/lru.go new file mode 100644 index 0000000..9ece7ba --- /dev/null +++ b/sdk/go/lru.go @@ -0,0 +1,72 @@ +package nixmsg + +import "container/list" + +type lruEntry struct { + key string + val any +} + +// lruCache 固定容量 LRU,删除时同步移除链表节点。 +type lruCache struct { + cap int + ll *list.List + m map[string]*list.Element +} + +func newLRU(cap int) *lruCache { + if cap <= 0 { + cap = 10000 + } + return &lruCache{ + cap: cap, + ll: list.New(), + m: make(map[string]*list.Element), + } +} + +func (c *lruCache) Get(key string) (any, bool) { + el, ok := c.m[key] + if !ok { + return nil, false + } + c.ll.MoveToFront(el) + return el.Value.(*lruEntry).val, true +} + +func (c *lruCache) Put(key string, val any) { + if el, ok := c.m[key]; ok { + el.Value.(*lruEntry).val = val + c.ll.MoveToFront(el) + return + } + el := c.ll.PushFront(&lruEntry{key: key, val: val}) + c.m[key] = el + for c.ll.Len() > c.cap { + back := c.ll.Back() + if back == nil { + break + } + ent := back.Value.(*lruEntry) + c.ll.Remove(back) + delete(c.m, ent.key) + } +} + +func (c *lruCache) Delete(key string) { + el, ok := c.m[key] + if !ok { + return + } + c.ll.Remove(el) + delete(c.m, key) +} + +func (c *lruCache) Len() int { + return c.ll.Len() +} + +func (c *lruCache) Has(key string) bool { + _, ok := c.m[key] + return ok +} diff --git a/sdk/go/receive.go b/sdk/go/receive.go index 8ad7b8f..8b68d65 100644 --- a/sdk/go/receive.go +++ b/sdk/go/receive.go @@ -15,7 +15,6 @@ func (c *Client) handleDown(payload []byte) { if err := unmarshalJSON(payload, &head); err != nil { return } - // resp 必须在收包路径同步处理,否则 request/ack 在 downLoop 里等待时会死锁。 if head.Type == "resp" { var rf respFrame if err := unmarshalJSON(payload, &rf); err != nil { @@ -35,10 +34,45 @@ func (c *Client) handleDown(payload []byte) { } return } + if head.Type == "fatal" { + var f struct { + Reason string `json:"reason"` + } + _ = unmarshalJSON(payload, &f) + c.handleFatal(f.Reason) + return + } + kind := head.Type cp := append([]byte(nil), payload...) + c.enqueueDown(queuedFrame{payload: cp, kind: kind}) +} + +func (c *Client) enqueueDown(fr queuedFrame) { + c.mu.Lock() + ch := c.incoming + ctx := c.ctx + c.mu.Unlock() + if ch == nil || ctx == nil { + return + } + if fr.kind == "presence" || fr.kind == "group_event" { + select { + case ch <- fr: + default: + // 超阈值丢最旧:缓冲满则丢弃本条事件 + } + return + } select { - case c.downCh <- cp: - case <-c.ctx.Done(): + case ch <- fr: + case <-ctx.Done(): + default: + go func() { + select { + case ch <- fr: + case <-ctx.Done(): + } + }() } } @@ -47,8 +81,8 @@ func (c *Client) downLoop(ctx context.Context) { select { case <-ctx.Done(): return - case payload := <-c.downCh: - c.handleDownApp(payload) + case fr := <-c.incoming: + c.handleDownApp(fr.payload) } } } @@ -71,23 +105,17 @@ func (c *Client) handleDownApp(payload []byte) { c.handlePresence(payload) case "group_event": c.handleGroupEvent(payload) - case "fatal": - var f struct { - Reason string `json:"reason"` - } - _ = unmarshalJSON(payload, &f) - c.mu.Lock() - c.stopReconnect = true - c.setStateLocked(StateAuthFailed, f.Reason) - c.failQueuedLocked(apiErr("fatal", f.Reason)) - cancel := c.cancel - c.mu.Unlock() - if cancel != nil { - cancel() - } } } +func msgKey(from, id string) string { + return "m\x00" + from + "\x00" + id +} + +func receiptKey(id string) string { + return "r\x00" + id +} + func (c *Client) handleMsg(payload []byte) { var m struct { ID string `json:"id"` @@ -100,57 +128,57 @@ func (c *Client) handleMsg(payload []byte) { if err := unmarshalJSON(payload, &m); err != nil { return } - key := m.From + "\x00" + m.ID + key := msgKey(m.From, m.ID) c.mu.Lock() - ent := c.dedup[key] manual := c.opts.ManualAck - if ent != nil && ent.state == dedupAcked { - c.mu.Unlock() - _ = c.sendAckFrame(m.From, m.ID) - return + raw, ok := c.store.Get(key) + if ok { + ent := raw.(*dedupEntry) + if ent.state == dedupAcked { + c.mu.Unlock() + go func() { _ = c.sendAckFrame(m.From, m.ID) }() + return + } + if ent.state == dedupDelivered || ent.state == dedupRevoked { + c.mu.Unlock() + return + } } - if ent != nil && ent.state == dedupDelivered { - c.mu.Unlock() - return - } - c.rememberDedupLocked(key, m.From, m.ID, dedupDelivered) + c.store.Put(key, &dedupEntry{state: dedupDelivered, from: m.From, id: m.ID}) c.mu.Unlock() msg := Message{ID: m.ID, From: m.From, To: m.To, Body: m.Body, Meta: m.Meta, SendAtMs: m.SendAtMs} - var cbErr error - c.cbMu.Lock() - if c.onMessage != nil { - cbErr = c.onMessage(msg) - } - c.cbMu.Unlock() - - if manual { - return - } - if cbErr != nil { - c.mu.Lock() - delete(c.dedup, key) - c.mu.Unlock() - return - } - _ = c.sendAckFrame(m.From, m.ID) - c.mu.Lock() - if e := c.dedup[key]; e != nil { - e.state = dedupAcked - } - c.mu.Unlock() -} - -func (c *Client) rememberDedupLocked(key, from, id string, st dedupState) { - if _, ok := c.dedup[key]; !ok { - c.dedupOrd = append(c.dedupOrd, key) - for len(c.dedupOrd) > c.opts.DedupCapacity { - old := c.dedupOrd[0] - c.dedupOrd = c.dedupOrd[1:] - delete(c.dedup, old) + done := make(chan struct{}) + c.dispatch(func() { + defer close(done) + var cbErr error + if c.onMessage != nil { + cbErr = c.onMessage(msg) } + if manual { + return + } + if cbErr != nil { + c.mu.Lock() + c.store.Delete(key) + c.mu.Unlock() + return + } + go func() { + _ = c.sendAckFrame(m.From, m.ID) + c.mu.Lock() + if raw, ok := c.store.Get(key); ok { + if ent, ok := raw.(*dedupEntry); ok { + ent.state = dedupAcked + } + } + c.mu.Unlock() + }() + }) + select { + case <-done: + case <-c.ctx.Done(): } - c.dedup[key] = &dedupEntry{state: st, from: from, id: id} } // Ack 手动确认。 @@ -158,9 +186,9 @@ func (c *Client) Ack(msg Message) error { if err := c.sendAckFrame(msg.From, msg.ID); err != nil { return err } - key := msg.From + "\x00" + msg.ID + key := msgKey(msg.From, msg.ID) c.mu.Lock() - c.rememberDedupLocked(key, msg.From, msg.ID, dedupAcked) + c.store.Put(key, &dedupEntry{state: dedupAcked, from: msg.From, id: msg.ID}) c.mu.Unlock() return nil } @@ -198,21 +226,22 @@ func (c *Client) handleReceipt(payload []byte) { if err := unmarshalJSON(payload, &r); err != nil { return } + key := receiptKey(r.ReceiptID) c.mu.Lock() - if _, ok := c.receiptSeen[r.ReceiptID]; ok { + if c.store.Has(key) { c.mu.Unlock() - _ = c.sendReceiptAck(r.ReceiptID) + go func() { _ = c.sendReceiptAck(r.ReceiptID) }() return } - c.receiptSeen[r.ReceiptID] = struct{}{} + c.store.Put(key, struct{}{}) c.mu.Unlock() ev := Receipt{ReceiptID: r.ReceiptID, ID: r.ID, EndpointID: r.EndpointID, State: r.State, Reason: r.Reason, AtMs: r.AtMs} - c.cbMu.Lock() - if c.onReceipt != nil { - c.onReceipt(ev) - } - c.cbMu.Unlock() - _ = c.sendReceiptAck(r.ReceiptID) + c.dispatch(func() { + if c.onReceipt != nil { + c.onReceipt(ev) + } + }) + go func() { _ = c.sendReceiptAck(r.ReceiptID) }() } func (c *Client) sendReceiptAck(receiptID string) error { @@ -232,24 +261,27 @@ func (c *Client) handleRevoked(payload []byte) { if err := unmarshalJSON(payload, &r); err != nil { return } - key := r.From + "\x00" + r.ID + key := msgKey(r.From, r.ID) c.mu.Lock() - ent := c.dedup[key] - if ent == nil || ent.state == dedupAcked { - c.mu.Unlock() - return + raw, ok := c.store.Get(key) + if ok { + ent := raw.(*dedupEntry) + if ent.state == dedupAcked || ent.state == dedupRevoked { + c.mu.Unlock() + return + } } - delete(c.dedup, key) + c.store.Put(key, &dedupEntry{state: dedupRevoked, from: r.From, id: r.ID}) c.mu.Unlock() c.emitRevoked(RevokedEvent{ID: r.ID, From: r.From, Reason: r.Reason}) } func (c *Client) emitRevoked(e RevokedEvent) { - c.cbMu.Lock() - defer c.cbMu.Unlock() - if c.onRevoked != nil { - c.onRevoked(e) - } + c.dispatch(func() { + if c.onRevoked != nil { + c.onRevoked(e) + } + }) } func (c *Client) handlePresence(payload []byte) { @@ -261,11 +293,11 @@ func (c *Client) handlePresence(payload []byte) { if err := unmarshalJSON(payload, &p); err != nil { return } - c.cbMu.Lock() - defer c.cbMu.Unlock() - if c.onPresence != nil { - c.onPresence(PresenceEvent{ID: p.ID, Online: p.Online, AtMs: p.AtMs}) - } + c.dispatch(func() { + if c.onPresence != nil { + c.onPresence(PresenceEvent{ID: p.ID, Online: p.Online, AtMs: p.AtMs}) + } + }) } func (c *Client) handleGroupEvent(payload []byte) { @@ -278,14 +310,19 @@ func (c *Client) handleGroupEvent(payload []byte) { if err := unmarshalJSON(payload, &g); err != nil { return } - c.cbMu.Lock() - defer c.cbMu.Unlock() - if c.onGroupEvent != nil { - c.onGroupEvent(GroupEvent{GroupID: g.GroupID, Event: g.Event, EndpointID: g.EndpointID, AtMs: g.AtMs}) - } + c.dispatch(func() { + if c.onGroupEvent != nil { + c.onGroupEvent(GroupEvent{GroupID: g.GroupID, Event: g.Event, EndpointID: g.EndpointID, AtMs: g.AtMs}) + } + }) } func (c *Client) request(ctx context.Context, frame map[string]any, allowUnready bool) (json.RawMessage, error) { + if _, ok := ctx.Deadline(); !ok { + var cancel context.CancelFunc + ctx, cancel = context.WithTimeout(ctx, 60*time.Second) + defer cancel() + } c.mu.Lock() if c.closed { c.mu.Unlock() @@ -296,6 +333,7 @@ func (c *Client) request(ctx context.Context, frame map[string]any, allowUnready return nil, apiErr(CodeNotConnected, "未握手") } tr := c.transport + clientCtx := c.ctx c.mu.Unlock() if tr == nil { return nil, apiErr(CodeNotConnected, "未连接") @@ -311,7 +349,7 @@ func (c *Client) request(ctx context.Context, frame map[string]any, allowUnready } ch := make(chan respFrame, 1) c.mu.Lock() - c.pending[rid] = &pendingReq{rid: rid, ch: ch} + c.pending[rid] = &pendingReq{rid: rid, ch: ch, isSend: false} c.mu.Unlock() if err := tr.PublishUp(payload); err != nil { c.mu.Lock() @@ -319,22 +357,28 @@ func (c *Client) request(ctx context.Context, frame map[string]any, allowUnready c.mu.Unlock() return nil, err } + var rf respFrame select { case <-ctx.Done(): c.mu.Lock() delete(c.pending, rid) c.mu.Unlock() return nil, ctx.Err() - case rf := <-ch: - if !rf.OK { - code, msg := CodeBadRequest, "请求失败" - if rf.Error != nil { - code, msg = rf.Error.Code, rf.Error.Message - } - return nil, apiErr(code, msg) - } - return rf.Data, nil + case <-clientCtx.Done(): + c.mu.Lock() + delete(c.pending, rid) + c.mu.Unlock() + return nil, apiErr(CodeNotConnected, "未连接") + case rf = <-ch: } + if !rf.OK { + code, msg := CodeBadRequest, "请求失败" + if rf.Error != nil { + code, msg = rf.Error.Code, rf.Error.Message + } + return nil, apiErr(code, msg) + } + return rf.Data, nil } // bodyDecodedLen 按解码后字节计正文大小。 diff --git a/sdk/go/register.go b/sdk/go/register.go index 5efa196..6b50859 100644 --- a/sdk/go/register.go +++ b/sdk/go/register.go @@ -41,8 +41,8 @@ func Register(ctx context.Context, connectOrRegisterURL, registrationCode string return RegisterResult{}, err } var wrap struct { - OK bool `json:"ok"` - Data struct { + OK bool `json:"ok"` + Data struct { ID string `json:"id"` LoginPassword string `json:"login_password"` } `json:"data"` diff --git a/sdk/go/send.go b/sdk/go/send.go index d5784f6..ca2d255 100644 --- a/sdk/go/send.go +++ b/sdk/go/send.go @@ -11,11 +11,16 @@ func (c *Client) Send(ctx context.Context, to Target, body Body, opt SendOptions c.mu.Lock() limits := c.limits skew := c.clockSkew - handshook := c.handshook maxQ := c.opts.SendQueueSize qLen := len(c.sendQ) + stopped := c.stopReconnect || c.closed + stopErr := c.stopErrLocked() c.mu.Unlock() + if stopped { + return SendResult{}, stopErr + } + if body.Enc == "" { body.Enc = "utf8" } @@ -39,10 +44,7 @@ func (c *Client) Send(ctx context.Context, to Target, body Body, opt SendOptions if maxBody <= 0 { maxBody = 262144 } - if handshook && n > maxBody { - return SendResult{}, apiErr(CodeBodyTooLarge, "正文超限") - } - if !handshook && n > 262144 { + if n > maxBody { return SendResult{}, apiErr(CodeBodyTooLarge, "正文超限") } @@ -81,7 +83,6 @@ func (c *Client) Send(ctx context.Context, to Target, body Body, opt SendOptions return SendResult{}, apiErr(CodeBadRequest, "sendAt 与 delay 互斥") } if opt.SendAt != nil { - // sendAt 使用本机时间 + 服务器偏差,换算后写入,重交不重算 ms := opt.SendAt.UnixMilli() + skew sendAtMs = &ms frame["send_at_ms"] = ms @@ -97,7 +98,7 @@ func (c *Client) Send(ctx context.Context, to Target, body Body, opt SendOptions if maxFrame <= 0 { maxFrame = 786432 } - if handshook && len(payload) > maxFrame { + if len(payload) > maxFrame { return SendResult{}, apiErr(CodeFrameTooLarge, "整帧超限") } @@ -111,10 +112,11 @@ func (c *Client) Send(ctx context.Context, to Target, body Body, opt SendOptions c.mu.Lock() if c.closed || c.stopReconnect { + err := c.stopErrLocked() c.mu.Unlock() - return SendResult{}, apiErr(CodeClosed, "已关闭") + return SendResult{}, err } - if qLen >= maxQ { + if qLen >= maxQ || len(c.sendQ) >= maxQ { c.mu.Unlock() return SendResult{}, apiErr(CodeQueueFull, "发送队列已满") } @@ -125,7 +127,21 @@ func (c *Client) Send(ctx context.Context, to Target, body Body, opt SendOptions select { case <-ctx.Done(): - return SendResult{}, ctx.Err() + c.mu.Lock() + if !item.inflight && !item.abandoned { + nq := c.sendQ[:0] + for _, it := range c.sendQ { + if it != item { + nq = append(nq, it) + } + } + c.sendQ = nq + c.mu.Unlock() + return SendResult{}, ctx.Err() + } + item.abandoned = true + c.mu.Unlock() + return SendResult{}, apiErr(CodeResultUnknown, "结果未知,请用同一消息号重试") case out := <-item.result: return out.res, out.err } @@ -134,13 +150,17 @@ func (c *Client) Send(ctx context.Context, to Target, body Body, opt SendOptions func (c *Client) drainSendQueue() { for { c.mu.Lock() - if !c.handshook || c.transport == nil { + if !c.handshook || c.transport == nil || c.stopReconnect { c.mu.Unlock() return } + maxFrame := c.limits.MaxFrameBytes + if maxFrame <= 0 { + maxFrame = 786432 + } var next *sendItem for _, it := range c.sendQ { - if !it.inflight { + if !it.inflight && !it.abandoned { next = it break } @@ -149,48 +169,93 @@ func (c *Client) drainSendQueue() { c.mu.Unlock() return } + if len(next.payload) > maxFrame { + c.mu.Unlock() + c.finishSend(next, SendResult{}, apiErr(CodeFrameTooLarge, "整帧超限")) + continue + } + next.epoch++ + captured := next.epoch next.inflight = true c.inflight++ tr := c.transport payload := next.payload rid, _ := next.frame["rid"].(string) item := next + lost := c.connLost c.mu.Unlock() - go c.dispatchSend(tr, item, rid, payload) + go c.dispatchSend(tr, item, rid, payload, captured, lost) } } -func (c *Client) dispatchSend(tr transport, item *sendItem, rid string, payload []byte) { +func (c *Client) dispatchSend(tr transport, item *sendItem, rid string, payload []byte, epoch uint64, lost <-chan struct{}) { ch := make(chan respFrame, 1) c.mu.Lock() - c.pending[rid] = &pendingReq{rid: rid, ch: ch} + if item.epoch != epoch || !item.inflight { + c.mu.Unlock() + return + } + c.pending[rid] = &pendingReq{rid: rid, ch: ch, isSend: true} c.mu.Unlock() if err := tr.PublishUp(payload); err != nil { c.mu.Lock() delete(c.pending, rid) - item.inflight = false - c.inflight-- + if item.epoch == epoch && item.inflight { + item.inflight = false + if c.inflight > 0 { + c.inflight-- + } + if c.stopReconnect { + errStop := c.stopErrLocked() + c.mu.Unlock() + c.finishSend(item, SendResult{}, errStop) + return + } + } c.mu.Unlock() - // 网络错误:保留队列等重连 return } - rf := <-ch + if lost == nil { + lost = make(chan struct{}) + } + + var rf respFrame + select { + case rf = <-ch: + case <-lost: + return + case <-c.ctx.Done(): + return + } + + c.mu.Lock() + valid := item.epoch == epoch && item.inflight + c.mu.Unlock() + if !valid { + return + } + if !rf.OK { code, msg := CodeBadRequest, "发送失败" if rf.Error != nil { code, msg = rf.Error.Code, rf.Error.Message } if code == CodeRateLimited { - // 自动重交:保持同一 payload(含 id / send_at_ms) c.mu.Lock() item.inflight = false - c.inflight-- + if c.inflight > 0 { + c.inflight-- + } delete(c.pending, rid) + item.rateN++ + n := item.rateN + c.regenerateSendLocked(item) c.mu.Unlock() - time.AfterFunc(time.Second, func() { c.drainSendQueue() }) + d := backoffJitter(nominalDelay(n)) + time.AfterFunc(d, func() { c.drainSendQueue() }) return } c.finishSend(item, SendResult{}, apiErr(code, msg)) @@ -206,8 +271,8 @@ func (c *Client) dispatchSend(tr transport, item *sendItem, rid string, payload func (c *Client) finishSend(item *sendItem, res SendResult, err error) { c.mu.Lock() - // 从队列移除 out := item.result + abandoned := item.abandoned nq := c.sendQ[:0] for _, it := range c.sendQ { if it != item { @@ -216,13 +281,19 @@ func (c *Client) finishSend(item *sendItem, res SendResult, err error) { } c.sendQ = nq if item.inflight { - c.inflight-- + if c.inflight > 0 { + c.inflight-- + } item.inflight = false } + rid, _ := item.frame["rid"].(string) + delete(c.pending, rid) c.mu.Unlock() - select { - case out <- sendOutcome{res: res, err: err}: - default: + if !abandoned { + select { + case out <- sendOutcome{res: res, err: err}: + default: + } } c.drainSendQueue() } diff --git a/sdk/go/transport.go b/sdk/go/transport.go index 175758a..85278a6 100644 --- a/sdk/go/transport.go +++ b/sdk/go/transport.go @@ -5,6 +5,9 @@ import ( "time" ) +// DefaultKeepAliveSeconds MQTT 心跳,DEVELOPMENT 第 9 节附录。 +const DefaultKeepAliveSeconds = 30 + // transport 抽象 MQTT 应用层通道,便于单测注入假实现。 type transport interface { // Start 开始连接循环(含重连)。凭据在每次 CONNECT 时读取。 diff --git a/sdk/go/transport_mqtt.go b/sdk/go/transport_mqtt.go index 5b22765..b99abaa 100644 --- a/sdk/go/transport_mqtt.go +++ b/sdk/go/transport_mqtt.go @@ -15,15 +15,15 @@ import ( ) type mqttTransport struct { - mu sync.Mutex - cm *autopaho.ConnectionManager - cancel context.CancelFunc - cfg transportConfig - cred atomic.Value // string - upTopic string + mu sync.Mutex + cm *autopaho.ConnectionManager + cancel context.CancelFunc + cfg transportConfig + cred atomic.Value // string + upTopic string downTopic string - stopped atomic.Bool - ready chan struct{} + stopped atomic.Bool + ready chan struct{} } func newMQTTTransport() *mqttTransport { @@ -57,7 +57,7 @@ func (t *mqttTransport) Start(ctx context.Context, cfg transportConfig) error { var sessionExpiry uint32 // 0;由 ConnectPacketBuilder 显式写入 Properties cliCfg := autopaho.ClientConfig{ ServerUrls: []*url.URL{u}, - KeepAlive: 30, + KeepAlive: uint16(DefaultKeepAliveSeconds), ConnectTimeout: cfg.ConnectTimeout, CleanStartOnInitialConnection: false, // 不要只靠这个;每次用 ConnectPacketBuilder SessionExpiryInterval: sessionExpiry, @@ -75,8 +75,12 @@ func (t *mqttTransport) Start(ctx context.Context, cfg transportConfig) error { cfg.OnAuthFailed(reason) } cancel() + return } } + if cfg.Backoff != nil { + cfg.Backoff.MarkOffline() + } }, OnConnectionDown: func() bool { if t.stopped.Load() { @@ -113,14 +117,17 @@ func (t *mqttTransport) Start(ctx context.Context, cfg transportConfig) error { }() }, ClientConfig: paho.ClientConfig{ - ClientID: cfg.EndpointID, + ClientID: cfg.EndpointID, + PacketTimeout: cfg.ConnectTimeout, OnServerDisconnect: func(d *paho.Disconnect) { if d != nil && d.ReasonCode == 0x8E { if cfg.OnKicked != nil { cfg.OnKicked() } cancel() + return } + // 0x8B 等其它原因按可重试处理,继续重连。 }, OnPublishReceived: []func(paho.PublishReceived) (bool, error){ func(pr paho.PublishReceived) (bool, error) { @@ -140,6 +147,7 @@ func (t *mqttTransport) Start(ctx context.Context, cfg transportConfig) error { c.Properties = &paho.ConnectProperties{} } c.Properties.SessionExpiryInterval = &zero + // 不设置 ReceiveMaximum(B-01 / K-00)。 pass, _ := t.cred.Load().(string) c.UsernameFlag = true c.Username = cfg.EndpointID diff --git a/sdk/go/types.go b/sdk/go/types.go index f5c680d..7cc3942 100644 --- a/sdk/go/types.go +++ b/sdk/go/types.go @@ -39,7 +39,7 @@ type Options struct { // DedupCapacity from+id 去重容量,默认 10000。 DedupCapacity int // HTTPClient 注册用;nil 用默认。 - // 测试可注入 transport。 + // 测试通过同包 Options.transport 注入;正式 API 不导出假传输。 transport transport } @@ -91,9 +91,9 @@ type SendOptions struct { // SendResult 发送结果。 type SendResult struct { - ID string - SendAtMs int64 - State string + ID string `json:"id"` + SendAtMs int64 `json:"send_at_ms"` + State string `json:"state"` } // Message 下行消息(交给应用)。 @@ -165,10 +165,10 @@ type GroupEvent struct { // SelfInfo 自己的资料。 type SelfInfo struct { - ID string `json:"id"` - Name string `json:"name"` - DefaultDelayMs int64 `json:"default_delay_ms"` - TalkPasswordSet bool `json:"talk_password_set"` + ID string `json:"id"` + Name string `json:"name"` + DefaultDelayMs int64 `json:"default_delay_ms"` + TalkPasswordSet bool `json:"talk_password_set"` } // GroupInfo 群摘要。