fix: 按 K-00 约定修复 Go SDK 断线重交与退避

This commit is contained in:
Nixevol
2026-09-30 16:24:13 +08:00
parent 3749b9bdf1
commit f6f8ccf269
20 changed files with 1544 additions and 324 deletions
+9
View File
@@ -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
+9 -4
View File
@@ -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
+81 -70
View File
@@ -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
n int
skipFirst bool
online bool
stable bool
timer *time.Timer
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
}
d := nominalDelay(n)
if jitter {
return backoffJitter(d)
}
return withJitter(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)
}
+161 -12
View File
@@ -1,4 +1,4 @@
package nixmsg
package nixmsg
import (
"context"
@@ -14,6 +14,7 @@ import (
type pendingReq struct {
rid string
ch chan respFrame
isSend bool
}
type sendItem struct {
@@ -23,6 +24,9 @@ type sendItem struct {
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{}),
store: newLRU(10000),
state: StateOffline,
backoff: newReconnectBackoff(),
downCh: make(chan []byte, 256),
}
}
@@ -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()
}
+140 -39
View File
@@ -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()
},
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)
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)
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
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()
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) {
+12 -1
View File
@@ -35,12 +35,23 @@ const (
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"
)
+4
View File
@@ -17,7 +17,11 @@ func main() {
c := nixmsg.New()
c.OnSession(func(tok string) {
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)
+50
View File
@@ -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 }
@@ -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
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()
+349
View File
@@ -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)
}
}
+300
View File
@@ -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")
}
+72
View File
@@ -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
}
+106 -62
View File
@@ -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 c.downCh <- cp:
case <-c.ctx.Done():
case ch <- fr:
default:
// 超阈值丢最旧:缓冲满则丢弃本条事件
}
return
}
select {
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,21 +105,15 @@ 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) {
@@ -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 {
raw, ok := c.store.Get(key)
if ok {
ent := raw.(*dedupEntry)
if ent.state == dedupAcked {
c.mu.Unlock()
_ = c.sendAckFrame(m.From, m.ID)
go func() { _ = c.sendAckFrame(m.From, m.ID) }()
return
}
if ent != nil && ent.state == dedupDelivered {
if ent.state == dedupDelivered || ent.state == dedupRevoked {
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}
done := make(chan struct{})
c.dispatch(func() {
defer close(done)
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.store.Delete(key)
c.mu.Unlock()
return
}
go func() {
_ = c.sendAckFrame(m.From, m.ID)
c.mu.Lock()
if e := c.dedup[key]; e != nil {
e.state = dedupAcked
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():
}
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)
}
}
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()
c.dispatch(func() {
if c.onReceipt != nil {
c.onReceipt(ev)
}
c.cbMu.Unlock()
_ = c.sendReceiptAck(r.ReceiptID)
})
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 {
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()
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()
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()
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,13 +357,20 @@ 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:
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 {
@@ -335,7 +380,6 @@ func (c *Client) request(ctx context.Context, frame map[string]any, allowUnready
}
return rf.Data, nil
}
}
// bodyDecodedLen 按解码后字节计正文大小。
func bodyDecodedLen(b Body) (int, error) {
+90 -19
View File
@@ -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():
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)
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
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,14 +281,20 @@ func (c *Client) finishSend(item *sendItem, res SendResult, err error) {
}
c.sendQ = nq
if item.inflight {
if c.inflight > 0 {
c.inflight--
}
item.inflight = false
}
rid, _ := item.frame["rid"].(string)
delete(c.pending, rid)
c.mu.Unlock()
if !abandoned {
select {
case out <- sendOutcome{res: res, err: err}:
default:
}
}
c.drainSendQueue()
}
+3
View File
@@ -5,6 +5,9 @@ import (
"time"
)
// DefaultKeepAliveSeconds MQTT 心跳,DEVELOPMENT 第 9 节附录。
const DefaultKeepAliveSeconds = 30
// transport 抽象 MQTT 应用层通道,便于单测注入假实现。
type transport interface {
// Start 开始连接循环(含重连)。凭据在每次 CONNECT 时读取。
+9 -1
View File
@@ -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() {
@@ -114,13 +118,16 @@ func (t *mqttTransport) Start(ctx context.Context, cfg transportConfig) error {
},
ClientConfig: paho.ClientConfig{
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
+4 -4
View File
@@ -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 下行消息(交给应用)。