fix: 按 K-00 约定修复 Go SDK 断线重交与退避
This commit is contained in:
+9
-4
@@ -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
|
||||
|
||||
+84
-73
@@ -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)
|
||||
}
|
||||
|
||||
+172
-23
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
+149
-48
@@ -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) {
|
||||
|
||||
+26
-15
@@ -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"
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
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()
|
||||
@@ -3,7 +3,7 @@ package nixmsg_test
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
+148
-104
@@ -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 按解码后字节计正文大小。
|
||||
|
||||
+2
-2
@@ -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"`
|
||||
|
||||
+98
-27
@@ -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()
|
||||
}
|
||||
|
||||
@@ -5,6 +5,9 @@ import (
|
||||
"time"
|
||||
)
|
||||
|
||||
// DefaultKeepAliveSeconds MQTT 心跳,DEVELOPMENT 第 9 节附录。
|
||||
const DefaultKeepAliveSeconds = 30
|
||||
|
||||
// transport 抽象 MQTT 应用层通道,便于单测注入假实现。
|
||||
type transport interface {
|
||||
// Start 开始连接循环(含重连)。凭据在每次 CONNECT 时读取。
|
||||
|
||||
+18
-10
@@ -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
|
||||
|
||||
+8
-8
@@ -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 群摘要。
|
||||
|
||||
Reference in New Issue
Block a user