29 changed files with 1391 additions and 69 deletions
+3 -1
View File
@@ -335,7 +335,9 @@ func runServe(ctx context.Context, cfg config.Config) error {
drainCancel() drainCancel()
shutCtx, shutCancel := context.WithTimeout(context.Background(), 5*time.Second) shutCtx, shutCancel := context.WithTimeout(context.Background(), 5*time.Second)
_ = brk.Shutdown(shutCtx) if shutErr := brk.Shutdown(shutCtx); shutErr != nil && !errors.Is(shutErr, context.DeadlineExceeded) && !errors.Is(shutErr, context.Canceled) {
slog.Error("broker shutdown", "err", shutErr)
}
shutCancel() shutCancel()
secondDrain := drainBudget - time.Since(drainStart) secondDrain := drainBudget - time.Since(drainStart)
+80
View File
@@ -1718,3 +1718,83 @@ issue #3 未关闭,`feat/fix-3-downlink-deadlock` 未合入 `main`。下面是
- 原因:原先 12 项备注写着未穷尽仍标通过,交付说明写成「通过 23」。 - 原因:原先 12 项备注写着未穷尽仍标通过,交付说明写成「通过 23」。
- 备选方案:为每个未测子项补验收用例(本波不做,避免为变绿放松断言)。 - 备选方案:为每个未测子项补验收用例(本波不做,避免为变绿放松断言)。
- 影响:汇总改为通过 19、部分通过 4(F03/F08/F21/F22)、失败 0。F19 仍引用仓库内 SDK 清单、本波不重跑。 - 影响:汇总改为通过 19、部分通过 4(F03/F08/F21/F22)、失败 0。F19 仍引用仓库内 SDK 清单、本波不重跑。
### 复审修复 R3-01
1. **推送 worker 漏唤醒:定时分发回执、超限后续、撤回腾窗**
- 日期:2026-09-30
- 原条款:Gitea #65;推送 worker 只在 `WakePush` 时跑一轮 `PushPending`。
- 实际做法:`dispatchDueBatch` 本条已分发时除 pending 接收方外始终 `WakePush` 发送方;`PushPending` 遇 `too_large` 等跳过且本轮未占满窗口时再 `WakePush` 当前接收方(不在同一次递归扫表);`Recall` 成功改成 `recalled` 的接收方各 `WakePush` 一次,已推送撤回仍走 `flushRevokes`。不改 `PublishDown` 签名。
- 原因:终态回执、超限后的后续 pending、撤回腾出窗口后都依赖再唤醒,否则空闲在线端收不到。
- 备选方案:在同一次 `PushPending` 里循环扫完整 pending 表(否决,指令要求合并唤醒下一轮)。
- 影响:仅 `internal/app/message/`;相关单测见 `review_r3_65_test.go`。
### 复审修复 R3-03
1. **写队列满时 Close 不再与 enqueue 死锁**
- 日期:2026-09-30
- 原条款:issue #67;停机需 Drain/Close 在超时内返回;已入队任务仍由写 goroutine 处理完。
- 实际做法:Queue 增加 closing 信号。enqueue 在持 sendMu.RLock 时 select 同时等待 q.ch、closing 与 ctx.Done(),通道满时不再无限阻塞。Close 先 Swap(closed) 并 close(closing) 唤醒在途发送方释放读锁,再在写锁内 close(q.ch),最后等 loop 退出。不向已关闭 channel 发送。
- 原因:通道满时并发 Do 占着读锁堵在发送上,Close 的写锁拿不到,停机卡死;未提交的在途写也会丢。
- 备选方案:入队改为非阻塞,满则立即 ErrBusy(否决:改变背压语义,正常高峰会误伤提交)。
- 影响:仅 internal/store/queue.go 与测试;不改迁移编号。
### 复审修复 R3-04
- 日期:2026-09-30
- 原条款:Gitea #68;`dispatchSend` 在 `PublishUp` 失败且未停止重连时清 inflight 后静默 return,`Send` 永久挂起。
- 实际做法:失败且未 `stopReconnect` 时保留 id/body/send_at_ms,经 `regenerateSendLocked` 换新 rid,清本次 inflight 并 `drainSendQueue` 继续泵;已停止重连时仍 `finishSend` 返回错误。
- 原因:条目已不在途,重连时的 `requeueInflightLocked` 不会捡回,调用方一直等 `result`。
- 备选方案:失败即 `finishSend` 报错(否决,与断线/限速重交语义不一致);只换 rid 不唤醒队列(否决,等同 JS #69)。
- 影响:仅 `sdk/go`;假传输可模拟 send 发布失败一次。
### 复审修复 R3-05
1. **JS SDK 传输失败后须再泵发送队列**
- 日期:2026-09-30
- 原条款:DEVELOPMENT 第 5 节 / 附录 SDK 行为约定:发送在途失败后换新 `rid` 并重交;`id` / 正文 / `send_at_ms` 不变。Gitea #69。
- 实际做法:`dispatchSend` 中 `publishUp` 抛出非协议 `APIError` 时,若未 `stopReconnect` 则清 inflight、换新 `rid` 后立刻 `drainSendQueue`;若已 `stopReconnect` 则 `finishSendErr` 结束本次 `send`。不额外再乘一次抖动。
- 原因:原先只换 `rid` 并 `return`,连接未断时 `hello` 不再走,队列停泵,`send()` Promise 永不结束。
- 备选方案:按 `rate_limited` 同一条的退避再泵(否决本波,连接仍在线时立即重交更贴切,且避免与已有等待叠乘抖动)。
- 影响:仅 `sdk/js`;假传输失败一次后会再次上行且 `rid` 已变。
### 复审修复 R3-06
1. **Python/Java 首次握手失败应停止重连;Java rate_limited 在途计数只减一次**
- 日期:2026-09-30
- 原条款:DEVELOPMENT 第 9 节附录「首次连接超时或握手失败应停止重连,并向 connect() 返回未连接」;issue #70。
- 实际做法:`sdk/python` 的 `connect()` 在 `_handshake_error` 或非 ONLINE 终态时置 `_stop_reconnect`/`_want_connected=False`;内部连接/握手超时改用 `not_connected`(不再用 `busy`)。`sdk/java` 的 `connectSync` 在 `handshakeError` 时同样置 `stopReconnect`;`rate_limited` 路径去掉第二次 `inflightSends--`,只换新 rid 再排队。Python `transport.py` 将 paho 改为惰性导入,便于本机无 paho 时仍跑 FakeTransport 单测。不改 Go/JS 重连公式。
- 原因:原先抛错后工作线程仍按退避重连;Java 限速把在途计数减了两次。
- 备选方案:在后台 `_attempt_connect`/`attemptConnect` 内首次失败即停(否决:应用再次 `connect()` 才应恢复,标志应在对外 `connect` 失败路径统一置位)。
- 影响:仅 `sdk/python`、`sdk/java`;应用需再次调用 `connect()` 才会重连。
### 复审修复 R3-07
1. **客户端建群事务内复核创建者**
- 日期:2026-09-30
- 原条款:Gitea #71。`group.Create` 在 `Queue.Do` 里直接 `insertMemberTx` 群主,不跑 `endpointCheckTx`;停用/删除与建群抢在事务外对话密码窗口时,可插入已停用甚至刚删除的 `owner_id`。
- 实际做法:`Create` 同一写事务里,插入群与群主成员之前对 `actorID` 做与 `createAdmin` 相同的 `endpointCheckTx`;已停用 `endpoint_disabled`,不存在 `invalid_target`,且不插入群。不改 `emit`、不改迁移编号。
- 原因:成员侧已有事务内复核,群主侧缺对称校验,竞态会留下非法群主。
- 备选方案:事务外再查一次创建者(否决,无法覆盖密码窗口与写事务之间的竞态)。
- 影响:创建者在进入写事务前被停用/删除时建群失败且无新群行;与后台建群错误码对齐。
### 复审修复 R3-08
1. **后台加人审计按实际插入数,全跳过不记 ok**
- 日期:2026-09-30
- 原条款:Gitea #72;`AdminAddMembers` 对已是成员与请求内重复编号静默跳过;审计用 `len(member_ids)-len(failed)` 当成功数。
- 实际做法:`handleGroupAddMembers` 用加人前后 `group_members` 行数差得到 `added`;审计 `result` 用 `added` 与 `failed`(全跳过且无失败记 `noop`,不记 `ok`);响应增加 `added`,`failed` 仍只含真正失败项。前端按 `added` 提示,已是成员不改成错误码。
- 原因:全是已有成员时旧逻辑审计成 ok、界面提示已加人,实际插入 0 行。
- 备选方案:把已是成员写入 `failed`(否决,会改变客户端错误语义);扩展 `group.AddResult` 返回插入列表(可后续做,本波不改身份线接口)。
- 影响:管理 API 加人成功体多 `added` 字段;契约文档示例仍写 `{"failed":[]}`,以本偏差为准。`added` 取写事务内实际插入数(`AddResult.Added`,不进端协议 JSON),不用事务外两次 COUNT 的差,避免并发加人/踢人把别人的变更算进本请求。
### 复审修复 R3-02
1. **积压时 fatal/logout 与停机 0x8B 须等本帧写出**
- 日期:2026-09-30
- 原条款:Gitea #66;B-04 / B-08。
- 实际做法:带断开的下行帧用本帧 `OnPacketSent` 完成信号(优先 packet id,否则按载荷匹配),不再用连接级 `sentPub` 总数。`Shutdown` 在 ctx 未取消时先等下行队列与 `wirePending` 排空,再 `DisconnectClient` 发 `0x8B` 并在截止前等连接拆掉;ctx 已取消则发完即 `Close`。`serve` 仍给 5 秒预算并记录非超时错误。未合 `feat/fix-3-downlink-deadlock`。
- 原因:前面 PUBLISH 的 `OnPacketSent` 会让总数等待提前返回;`Shutdown` 对 ctx 非阻塞 select 使 5 秒预算用不上,有 outbound 积压时 `0x8B` 只进 outbuf 随 `Stop` 丢掉。
- 备选方案:恢复固定 `Sleep`(否决);改 `PublishDown` 签名(否决)。
- 影响:队列/outbound 有积压时 fatal、logout 先到客户端再断开;停机在预算内尽量发出 `0x8B`,超时返回 ctx 错误而非空等。
- 补强:`sendOne` 从取出帧到返回前记 `inSend`,排空等待把它算上(含等大帧名额、尚未 `wirePending++` 的窗口)。有截止时间时排空最多用到截止前 1 秒,剩下的时间留给 `0x8B` 写出;排空没完成也不会因此跳过这段等待。
+12 -5
View File
@@ -325,15 +325,22 @@ func (h *Handler) handleGroupAddMembers(w http.ResponseWriter, r *http.Request)
if failed == nil { if failed == nil {
failed = []group.MemberFail{} failed = []group.MemberFail{}
} }
okN := len(req.MemberIDs) - len(failed) added := res.Added
if okN < 0 { if added < 0 {
okN = 0 added = 0
} }
h.auditPD(p, "group_add_members", id, batchAuditResult(okN, len(failed)), ip, map[string]any{
// 已是成员/请求内重复编号会静默跳过,不算失败;审计用实际插入数,避免全跳过写成 ok。
result := batchAuditResult(added, len(failed))
if added == 0 && len(failed) == 0 {
result = "noop"
}
h.auditPD(p, "group_add_members", id, result, ip, map[string]any{
"members": req.MemberIDs, "members": req.MemberIDs,
"failed": failed, "failed": failed,
"added": added,
}) })
httpx.WriteOK(w, map[string]any{"failed": failed}) httpx.WriteOK(w, map[string]any{"failed": failed, "added": added})
} }
func (h *Handler) handleGroupRemoveMember(w http.ResponseWriter, r *http.Request) { func (h *Handler) handleGroupRemoveMember(w http.ResponseWriter, r *http.Request) {
+87 -3
View File
@@ -262,6 +262,7 @@ func TestH02BatchPartialImportAndGroups(t *testing.T) {
insertEndpoint(t, db, "keep-1", "admin", true, false) insertEndpoint(t, db, "keep-1", "admin", true, false)
insertEndpoint(t, db, "alice", "admin", true, false) insertEndpoint(t, db, "alice", "admin", true, false)
insertEndpoint(t, db, "bob", "admin", true, false) insertEndpoint(t, db, "bob", "admin", true, false)
insertEndpoint(t, db, "carol", "admin", true, false)
res := doReq(t, client, http.MethodPost, srv.URL+"/api/admin/endpoints/batch", res := doReq(t, client, http.MethodPost, srv.URL+"/api/admin/endpoints/batch",
`{"ids":["keep-1","missing-ep"],"action":"disable"}`, `{"ids":["keep-1","missing-ep"],"action":"disable"}`,
@@ -304,11 +305,24 @@ func TestH02BatchPartialImportAndGroups(t *testing.T) {
} }
res = doReq(t, client, http.MethodPost, srv.URL+"/api/admin/groups/"+created.ID+"/members", res = doReq(t, client, http.MethodPost, srv.URL+"/api/admin/groups/"+created.ID+"/members",
`{"member_ids":["keep-1","ghost-ep"]}`, csrf()) `{"member_ids":["carol","ghost-ep"]}`, csrf())
env = decodeEnv(t, res) env = decodeEnv(t, res)
if res.StatusCode != http.StatusOK || !env.OK { if res.StatusCode != http.StatusOK || !env.OK {
t.Fatalf("add members: %d %+v", res.StatusCode, env) t.Fatalf("add members: %d %+v", res.StatusCode, env)
} }
var addBody struct {
Failed []any `json:"failed"`
Added int `json:"added"`
}
if err := json.Unmarshal(env.Data, &addBody); err != nil {
t.Fatal(err)
}
if addBody.Added != 1 {
t.Fatalf("add response added=%d want 1", addBody.Added)
}
if len(addBody.Failed) != 1 {
t.Fatalf("add response failed=%v want 1", addBody.Failed)
}
res = doReq(t, client, http.MethodPost, srv.URL+"/api/admin/groups/"+created.ID+"/transfer", res = doReq(t, client, http.MethodPost, srv.URL+"/api/admin/groups/"+created.ID+"/transfer",
`{"endpoint_id":"bob"}`, csrf()) `{"endpoint_id":"bob"}`, csrf())
@@ -346,8 +360,8 @@ func TestH02BatchPartialImportAndGroups(t *testing.T) {
if ad == nil { if ad == nil {
t.Fatalf("add members missing detail: %v", add) t.Fatalf("add members missing detail: %v", add)
} }
if add["result"] != "partial" && add["result"] != "failed" && add["result"] != "ok" { if add["result"] != "partial" {
t.Fatalf("add members result=%v", add["result"]) t.Fatalf("add members result=%v want partial", add["result"])
} }
if _, ok := ad["members"]; !ok { if _, ok := ad["members"]; !ok {
t.Fatalf("add members missing members: %v", ad) t.Fatalf("add members missing members: %v", ad)
@@ -355,6 +369,9 @@ func TestH02BatchPartialImportAndGroups(t *testing.T) {
if _, ok := ad["failed"]; !ok { if _, ok := ad["failed"]; !ok {
t.Fatalf("add members missing failed: %v", ad) t.Fatalf("add members missing failed: %v", ad)
} }
if ad["added"] != float64(1) {
t.Fatalf("add members added=%v want 1", ad["added"])
}
tr := lastAuditByAction(t, recs, "group_transfer") tr := lastAuditByAction(t, recs, "group_transfer")
td, _ := tr["detail"].(map[string]any) td, _ := tr["detail"].(map[string]any)
@@ -364,3 +381,70 @@ func TestH02BatchPartialImportAndGroups(t *testing.T) {
assertNoSecrets(t, auditBuf.String(), "csv-pass-secret-1", "csv-pass-secret-2", testPassword) assertNoSecrets(t, auditBuf.String(), "csv-pass-secret-1", "csv-pass-secret-2", testPassword)
} }
func TestH02AddMembersAllSkippedAuditNotOK(t *testing.T) {
_, auditBuf, db, srv, client := setupH02(t)
login(t, client, srv.URL)
insertEndpoint(t, db, "alice", "admin", true, false)
insertEndpoint(t, db, "bob", "admin", true, false)
res := doReq(t, client, http.MethodPost, srv.URL+"/api/admin/groups",
`{"name":"一组","owner_id":"alice","member_ids":["bob"]}`, csrf())
env := decodeEnv(t, res)
if res.StatusCode != http.StatusOK || !env.OK {
t.Fatalf("group create: %d %+v", res.StatusCode, env)
}
var created struct {
ID string `json:"id"`
}
if err := json.Unmarshal(env.Data, &created); err != nil {
t.Fatal(err)
}
auditBuf.Reset()
var before int
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM group_members WHERE group_id = ?`, created.ID).Scan(&before); err != nil {
t.Fatal(err)
}
res = doReq(t, client, http.MethodPost, srv.URL+"/api/admin/groups/"+created.ID+"/members",
`{"member_ids":["bob","bob","alice"]}`, csrf())
env = decodeEnv(t, res)
if res.StatusCode != http.StatusOK || !env.OK {
t.Fatalf("add members: %d %+v", res.StatusCode, env)
}
var body struct {
Failed []any `json:"failed"`
Added int `json:"added"`
}
if err := json.Unmarshal(env.Data, &body); err != nil {
t.Fatal(err)
}
if len(body.Failed) != 0 {
t.Fatalf("failed=%v want empty (already members are not errors)", body.Failed)
}
if body.Added != 0 {
t.Fatalf("added=%d want 0", body.Added)
}
var after int
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM group_members WHERE group_id = ?`, created.ID).Scan(&after); err != nil {
t.Fatal(err)
}
if after != before {
t.Fatalf("member rows before=%d after=%d want unchanged", before, after)
}
recs := parseSlogJSON(t, auditBuf)
add := lastAuditByAction(t, recs, "group_add_members")
if add["result"] == "ok" {
t.Fatalf("audit result must not be ok when all skipped: %v", add)
}
ad, _ := add["detail"].(map[string]any)
if ad == nil {
t.Fatalf("missing detail: %v", add)
}
if ad["added"] != float64(0) {
t.Fatalf("detail.added=%v want 0", ad["added"])
}
}
+11 -2
View File
@@ -147,6 +147,15 @@ func (a *App) Create(ctx context.Context, actorID string, req *protocol.GroupCre
var inserted []string var inserted []string
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
if e := endpointCheckTx(tx, actorID); e != nil {
if protoCode(e) == protocol.CodeInvalidTarget {
return errCode(protocol.CodeInvalidTarget, "owner not found")
}
if protoCode(e) == protocol.CodeEndpointDisabled {
return errCode(protocol.CodeEndpointDisabled, "owner disabled")
}
return e
}
var exists int var exists int
qErr := tx.QueryRow(`SELECT 1 FROM groups WHERE id = ?`, gid).Scan(&exists) qErr := tx.QueryRow(`SELECT 1 FROM groups WHERE id = ?`, gid).Scan(&exists)
if qErr == nil { if qErr == nil {
@@ -611,7 +620,7 @@ func (a *App) AdminAddMembers(ctx context.Context, groupID string, memberIDs []s
failed = append(failed, MemberFail{ID: id, Code: protocol.CodeGroupFull}) failed = append(failed, MemberFail{ID: id, Code: protocol.CodeGroupFull})
} }
if len(toAdd) == 0 { if len(toAdd) == 0 {
return AddResult{Failed: failed}, nil return AddResult{Failed: failed, Added: 0}, nil
} }
now := a.nowMs() now := a.nowMs()
var inserted []string var inserted []string
@@ -655,7 +664,7 @@ func (a *App) AdminAddMembers(ctx context.Context, groupID string, memberIDs []s
for _, id := range inserted { for _, id := range inserted {
a.emit(ctx, notify, groupID, eventMemberAdded, id, now) a.emit(ctx, notify, groupID, eventMemberAdded, id, now)
} }
return AddResult{Failed: failed}, nil return AddResult{Failed: failed, Added: len(inserted)}, nil
} }
func (a *App) createAdmin(ctx context.Context, ownerID, name, gid string, members []protocol.GroupMemberIn) (CreateResult, error) { func (a *App) createAdmin(ctx context.Context, ownerID, name, gid string, members []protocol.GroupMemberIn) (CreateResult, error) {
+100
View File
@@ -604,6 +604,106 @@ func TestU02CreateDedupesMembers(t *testing.T) {
} }
} }
// raceTalk 在对话密码校验成功后执行一次 mutate,模拟事务外密码窗口里创建者被停用/删除。
type raceTalk struct {
inner group.TalkGate
mutate func()
once sync.Once
}
func (r *raceTalk) CheckTalkPasswordForJoin(ctx context.Context, actorID, targetID, talkPassword, remoteIP string) error {
err := r.inner.CheckTalkPasswordForJoin(ctx, actorID, targetID, talkPassword, remoteIP)
if err == nil && r.mutate != nil {
r.once.Do(r.mutate)
}
return err
}
func TestR307CreateRejectsActorDisabledOrDeletedBeforeWrite(t *testing.T) {
t.Parallel()
cases := []struct {
name string
code string
mutate func(ctx context.Context, t *testing.T, idApp *identity.App, db *store.DB)
}{
{
name: "disabled",
code: protocol.CodeEndpointDisabled,
mutate: func(ctx context.Context, t *testing.T, idApp *identity.App, _ *store.DB) {
t.Helper()
if err := idApp.Disable(ctx, "alice"); err != nil {
t.Fatal(err)
}
},
},
{
name: "deleted",
code: protocol.CodeInvalidTarget,
mutate: func(ctx context.Context, t *testing.T, idApp *identity.App, _ *store.DB) {
t.Helper()
if err := idApp.Delete(ctx, "alice"); err != nil {
t.Fatal(err)
}
},
},
}
for _, tc := range cases {
tc := tc
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
db, err := store.Open(filepath.Join(t.TempDir(), "data"), "FULL")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
fixed := time.UnixMilli(1_700_000_000_000)
locks := auth.NewLoginLocks()
idApp := identity.New(identity.Config{
DB: db, Hash: auth.NewStubHashPool(), Locks: locks,
Sessions: auth.NewSessionTokens(),
Now: func() time.Time { return fixed },
})
gate := &raceTalk{inner: idApp}
gApp := group.New(group.Config{
DB: db, Talk: gate, MaxGroupMembers: 1000,
Now: func() time.Time { return fixed }, DefaultRemoteIP: "1.1.1.1",
})
ctx := context.Background()
insertEP(t, db, "alice", 1)
insertEP(t, db, "bob", 1)
if setErr := idApp.SelfSetTalkPassword(ctx, "bob", "secret"); setErr != nil {
t.Fatal(setErr)
}
gate.mutate = func() { tc.mutate(ctx, t, idApp, db) }
var before int
if qErr := db.Read.QueryRow(`SELECT COUNT(*) FROM groups`).Scan(&before); qErr != nil {
t.Fatal(qErr)
}
_, err = gApp.Create(ctx, "alice", &protocol.GroupCreate{
V: protocol.Version, Type: protocol.TypeGroupCreate, RID: "1",
ID: "g_r307_" + tc.name, Name: "竞态群",
Members: []protocol.GroupMemberIn{{ID: "bob", TalkPassword: "secret"}},
})
if protoCode(err) != tc.code {
t.Fatalf("want %s got %v", tc.code, err)
}
var after int
if qErr := db.Read.QueryRow(`SELECT COUNT(*) FROM groups`).Scan(&after); qErr != nil {
t.Fatal(qErr)
}
if after != before {
t.Fatalf("groups leaked: before=%d after=%d", before, after)
}
var exists int
qErr := db.Read.QueryRow(`SELECT 1 FROM groups WHERE id = ?`, "g_r307_"+tc.name).Scan(&exists)
if !errors.Is(qErr, sql.ErrNoRows) {
t.Fatalf("expected no group row, got exists=%d err=%v", exists, qErr)
}
})
}
}
func TestU02AdminCreateOwnerMustExistAndEnabled(t *testing.T) { func TestU02AdminCreateOwnerMustExistAndEnabled(t *testing.T) {
t.Parallel() t.Parallel()
gApp, _, _, db, down := setup(t) gApp, _, _, db, down := setup(t)
+2
View File
@@ -28,6 +28,8 @@ type CreateResult struct {
// AddResult 是加人结果。 // AddResult 是加人结果。
type AddResult struct { type AddResult struct {
Failed []MemberFail `json:"failed,omitempty"` Failed []MemberFail `json:"failed,omitempty"`
// Added 是本请求在写事务内实际插入的人数。不进协议 JSON,供后台审计使用。
Added int `json:"-"`
} }
// ListItem 是 group.list 一项。 // ListItem 是 group.list 一项。
+5
View File
@@ -90,6 +90,7 @@ func (a *App) Recall(ctx context.Context, senderID string, req *protocol.Recall)
nowMs := a.now().UnixMilli() nowMs := a.now().UnixMilli()
var data protocol.RecallData var data protocol.RecallData
var revokes []revokeJob var revokes []revokeJob
var wakeReceivers []string
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
var seq int64 var seq int64
var state string var state string
@@ -156,6 +157,7 @@ WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
continue continue
} }
recalled++ recalled++
wakeReceivers = append(wakeReceivers, p.ep)
if p.pushed.Valid && p.pushed.String != "" { if p.pushed.Valid && p.pushed.String != "" {
revokes = append(revokes, revokeJob{ revokes = append(revokes, revokeJob{
endpointID: p.ep, endpointID: p.ep,
@@ -190,6 +192,9 @@ SELECT COUNT(*) FROM deliveries WHERE seq = ? AND state IN ('expired','dropped',
a.pendingRevoke = append(a.pendingRevoke, revokes...) a.pendingRevoke = append(a.pendingRevoke, revokes...)
a.mu.Unlock() a.mu.Unlock()
a.flushRevokes(ctx) a.flushRevokes(ctx)
for _, ep := range wakeReceivers {
a.WakePush(ep)
}
return data, nil return data, nil
} }
+8
View File
@@ -127,6 +127,7 @@ LIMIT ?`, nowMs, limit)
} }
mu.Lock() mu.Lock()
n++ n++
wake[d.senderID] = struct{}{}
mu.Unlock() mu.Unlock()
rows2, qErr := a.db.Read.QueryContext(ctx, ` rows2, qErr := a.db.Read.QueryContext(ctx, `
SELECT DISTINCT endpoint_id FROM deliveries WHERE seq = ? AND state = 'pending'`, d.seq) SELECT DISTINCT endpoint_id FROM deliveries WHERE seq = ? AND state = 'pending'`, d.seq)
@@ -203,8 +204,10 @@ LIMIT ?`, endpointID, room)
} }
var toClaim []pushItem var toClaim []pushItem
skipped := false
for _, it := range items { for _, it := range items {
if it.body == nil { if it.body == nil {
skipped = true
continue continue
} }
msg := protocol.Msg{ msg := protocol.Msg{
@@ -227,6 +230,7 @@ LIMIT ?`, endpointID, room)
return rejErr return rejErr
} }
a.WakePush(it.senderID) a.WakePush(it.senderID)
skipped = true
continue continue
} }
it.payload = payload it.payload = payload
@@ -252,6 +256,10 @@ LIMIT ?`, endpointID, room)
a.observeDispatchToPush(it.sendAt, nowMs) a.observeDispatchToPush(it.sendAt, nowMs)
} }
} }
// 超限等跳过未占满窗口时再唤醒本端,让 worker 下一轮取后续 pending(不在本轮递归扫表)。
if skipped && len(claimed) < room {
a.WakePush(endpointID)
}
return a.pushReceipts(ctx, endpointID, connID, nowMs) return a.pushReceipts(ctx, endpointID, connID, nowMs)
} }
+180
View File
@@ -0,0 +1,180 @@
package message
import (
"context"
"database/sql"
"encoding/json"
"testing"
"time"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
"git.asio.asia/nixevol/NixMsg/internal/protocol"
)
// 定时到点且接收方离线成终态时,在线发送方经 WakePush 拿到回执。
func TestR365DispatchDueWakesSenderForReceipt(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, nil)
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
_ = e.db.Queue.Do(context.Background(), func(tx *sql.Tx) error {
_, err := tx.Exec(`UPDATE endpoints SET offline_since = ? WHERE id=?`, e.nowMs-120_000, "bob")
return err
})
ctx := context.Background()
aliceConn := LiveConn{ConnID: "c-alice", Ready: true}
e.conns.Set("alice", aliceConn)
if err := e.app.OnHandshakeComplete(ctx, "alice", aliceConn); err != nil {
t.Fatal(err)
}
t.Cleanup(func() { e.app.stopPushWorker("alice", "") })
delay := int64(5000)
req := baseSend("due-rcpt", "bob")
req.DelayMs = &delay
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil {
t.Fatal(err)
}
st, _ := e.msgState("alice", "due-rcpt")
if st != StateScheduled {
t.Fatalf("want scheduled got %s", st)
}
if e.down.FilterType(protocol.TypeReceipt) != 0 {
t.Fatal("receipt before due")
}
e.setNow(e.nowMs + 5000)
if _, err := e.app.DispatchDue(ctx, e.nowMs, 10); err != nil {
t.Fatal(err)
}
deadline := time.Now().Add(2 * time.Second)
for time.Now().Before(deadline) {
if e.down.FilterType(protocol.TypeReceipt) > 0 {
return
}
time.Sleep(20 * time.Millisecond)
}
t.Fatalf("sender got no receipt after scheduled dispatch, receipts=%d", e.down.FilterType(protocol.TypeReceipt))
}
// 窗口内首条超限拒收后,同连接后续小消息仍被 worker 推送。
func TestR365TooLargeThenSmallContinuesPush(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, func(l *Limits) { l.DeliveryWindow = 1 })
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
ctx := context.Background()
// 整帧上限:小消息能过,大正文整帧超限被拒。
bobConn := LiveConn{ConnID: "c-bob", Ready: true, MaxReceiveBytes: 200}
e.conns.Set("bob", bobConn)
if err := e.app.OnHandshakeComplete(ctx, "bob", bobConn); err != nil {
t.Fatal(err)
}
t.Cleanup(func() { e.app.stopPushWorker("bob", "") })
big := baseSend("big-skip", "bob")
b := make([]byte, 400)
for i := range b {
b[i] = 'A'
}
big.Body.Data = string(b)
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, big); err != nil {
t.Fatal(err)
}
small := baseSend("small-ok", "bob")
small.Body.Data = "hi"
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, small); err != nil {
t.Fatal(err)
}
deadline := time.Now().Add(2 * time.Second)
for time.Now().Before(deadline) {
foundSmall := false
for _, p := range e.down.Snapshots() {
if payloadType(p.Payload) != protocol.TypeMsg {
continue
}
var m struct {
ID string `json:"id"`
}
_ = json.Unmarshal(p.Payload, &m)
if m.ID == "small-ok" {
foundSmall = true
}
}
if foundSmall {
seqBig := e.seqOf("alice", "big-skip")
st, reason := e.deliveryState(seqBig, "bob")
if st != DeliveryRejected || reason != ReasonTooLarge {
t.Fatalf("big: %s/%s", st, reason)
}
return
}
time.Sleep(20 * time.Millisecond)
}
t.Fatalf("small msg not pushed after too_large; msgs=%d snapshots=%d",
e.down.FilterType(protocol.TypeMsg), len(e.down.Snapshots()))
}
// 撤回已推在途消息后,同连接更晚的 pending 继续推。
func TestR365RecallFreesWindowForLaterPending(t *testing.T) {
t.Parallel()
e := openDeliveryEnv(t, func(l *Limits) { l.DeliveryWindow = 1 })
insertEndpoint(t, e.db, "alice", "", 1, 0)
insertEndpoint(t, e.db, "bob", "", 1, 0)
ctx := context.Background()
bobConn := LiveConn{ConnID: "c-bob", Ready: true}
e.conns.Set("bob", bobConn)
if err := e.app.OnHandshakeComplete(ctx, "bob", bobConn); err != nil {
t.Fatal(err)
}
t.Cleanup(func() { e.app.stopPushWorker("bob", "") })
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("first-in-flight", "bob")); err != nil {
t.Fatal(err)
}
if _, err := e.app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("second-wait", "bob")); err != nil {
t.Fatal(err)
}
deadline := time.Now().Add(2 * time.Second)
for time.Now().Before(deadline) {
if e.down.FilterType(protocol.TypeMsg) >= 1 {
break
}
time.Sleep(10 * time.Millisecond)
}
if e.down.FilterType(protocol.TypeMsg) < 1 {
t.Fatal("first msg not pushed")
}
seq1 := e.seqOf("alice", "first-in-flight")
var pushed sql.NullString
_ = e.db.Read.QueryRow(`SELECT pushed_conn FROM deliveries WHERE seq=? AND endpoint_id='bob'`, seq1).Scan(&pushed)
if !pushed.Valid {
t.Fatal("first not in-flight")
}
if _, err := e.app.Recall(ctx, "alice", &protocol.Recall{
V: protocol.Version, Type: protocol.TypeRecall, RID: "r1", ID: "first-in-flight",
}); err != nil {
t.Fatal(err)
}
deadline = time.Now().Add(2 * time.Second)
for time.Now().Before(deadline) {
for _, p := range e.down.Snapshots() {
if payloadType(p.Payload) != protocol.TypeMsg {
continue
}
var m struct {
ID string `json:"id"`
}
_ = json.Unmarshal(p.Payload, &m)
if m.ID == "second-wait" {
return
}
}
time.Sleep(20 * time.Millisecond)
}
t.Fatalf("second pending not pushed after recall; msgs=%d", e.down.FilterType(protocol.TypeMsg))
}
+179
View File
@@ -87,6 +87,50 @@ func TestPublishThenDisconnectWritesThenCloses(t *testing.T) {
} }
} }
func TestPublishThenDisconnectAfterQueuedFrame(t *testing.T) {
b, w, done := startTCPClient(t, "ep-ptd-q")
defer func() { _ = b.Close() }()
defer func() {
_ = w.Close()
select {
case <-done:
case <-time.After(3 * time.Second):
}
}()
writeConnect(t, w, "ep-ptd-q", 30, 0)
readExactPacket(t, w, packets.Connack, 3*time.Second)
writeSubscribe(t, w, downTopic("ep-ptd-q"))
readExactPacket(t, w, packets.Suback, 3*time.Second)
waitSession(t, b, "ep-ptd-q")
first := []byte(`{"v":1,"type":"msg","id":"queued-ahead"}`)
fatal := []byte(`{"v":1,"type":"fatal","reason":"disabled"}`)
if err := b.PublishDown(context.Background(), "ep-ptd-q", "", first, port.PublishOpts{QoS: 1}); err != nil {
t.Fatal(err)
}
if err := b.PublishThenDisconnect(context.Background(), "ep-ptd-q", "", fatal, 1, port.DisconnectFatal); err != nil {
t.Fatal(err)
}
gotFirst := readDownPayload(t, w, 3*time.Second)
if !bytes.Equal(gotFirst, first) {
t.Fatalf("first got %s", gotFirst)
}
gotFatal := readDownPayload(t, w, 3*time.Second)
if !bytes.Equal(gotFatal, fatal) {
t.Fatalf("fatal got %s want %s (disconnected before fatal frame)", gotFatal, fatal)
}
_ = w.SetReadDeadline(time.Now().Add(3 * time.Second))
buf := make([]byte, 64)
n, err := io.ReadAtLeast(w, buf, 2)
if err != nil && n == 0 {
return
}
if n > 0 && buf[0]>>4 == packets.Disconnect {
return
}
}
func TestShutdownUsesServerShuttingDown(t *testing.T) { func TestShutdownUsesServerShuttingDown(t *testing.T) {
b, w, done := startTCPClient(t, "ep-shut") b, w, done := startTCPClient(t, "ep-shut")
defer func() { defer func() {
@@ -116,6 +160,141 @@ func TestShutdownUsesServerShuttingDown(t *testing.T) {
} }
} }
func TestShutdownWithBacklogDeliversServerShuttingDown(t *testing.T) {
b, w, done := startTCPClient(t, "ep-shut-bl")
defer func() {
_ = w.Close()
select {
case <-done:
case <-time.After(3 * time.Second):
}
}()
writeConnect(t, w, "ep-shut-bl", 30, 0)
readExactPacket(t, w, packets.Connack, 3*time.Second)
writeSubscribe(t, w, downTopic("ep-shut-bl"))
readExactPacket(t, w, packets.Suback, 3*time.Second)
waitSession(t, b, "ep-shut-bl")
payload := bytes.Repeat([]byte("b"), 1024)
for i := 0; i < 8; i++ {
if err := b.PublishDown(context.Background(), "ep-shut-bl", "", payload, port.PublishOpts{QoS: 1}); err != nil {
t.Fatalf("publish %d: %v", i, err)
}
}
saw8B := make(chan bool, 1)
go func() {
deadline := time.Now().Add(3 * time.Second)
for time.Now().Before(deadline) {
_ = w.SetReadDeadline(time.Now().Add(200 * time.Millisecond))
hdr := make([]byte, 1)
if _, err := io.ReadFull(w, hdr); err != nil {
continue
}
rem, err := readRemainingLengthConn(w)
if err != nil {
continue
}
body := make([]byte, rem)
if _, err := io.ReadFull(w, body); err != nil {
continue
}
switch hdr[0] >> 4 {
case packets.Publish:
qos := (hdr[0] >> 1) & 0x3
if qos > 0 {
pk := new(packets.Packet)
pk.ProtocolVersion = 5
pk.FixedHeader = packets.FixedHeader{Type: packets.Publish, Remaining: rem, Qos: qos}
if decErr := pk.PublishDecode(body); decErr == nil {
ack := packets.Packet{
FixedHeader: packets.FixedHeader{Type: packets.Puback},
ProtocolVersion: 5,
PacketID: pk.PacketID,
}
var ab bytes.Buffer
_ = ack.PubackEncode(&ab)
_, _ = w.Write(ab.Bytes())
}
}
case packets.Disconnect:
if rem >= 1 && body[0] == packets.ErrServerShuttingDown.Code {
saw8B <- true
return
}
}
}
saw8B <- false
}()
time.Sleep(20 * time.Millisecond)
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
err := b.Shutdown(ctx)
got := false
select {
case got = <-saw8B:
case <-time.After(4 * time.Second):
t.Fatal("reader hung")
}
if got {
if err != nil && !errors.Is(err, context.DeadlineExceeded) {
t.Fatalf("shutdown after 0x8B: %v", err)
}
return
}
if err == nil {
t.Fatal("expected 0x8B or shutdown deadline error, got neither")
}
if !errors.Is(err, context.DeadlineExceeded) && !errors.Is(err, context.Canceled) {
t.Fatalf("shutdown err=%v want deadline/cancel when 0x8B not seen", err)
}
}
func TestShutdownCancelledContextReturnsQuickly(t *testing.T) {
b, w, done := startTCPClient(t, "ep-shut-cancel")
defer func() {
_ = w.Close()
select {
case <-done:
case <-time.After(3 * time.Second):
}
}()
writeConnect(t, w, "ep-shut-cancel", 30, 0)
readExactPacket(t, w, packets.Connack, 3*time.Second)
waitSession(t, b, "ep-shut-cancel")
ctx, cancel := context.WithCancel(context.Background())
cancel()
start := time.Now()
_ = b.Shutdown(ctx)
if time.Since(start) > 500*time.Millisecond {
t.Fatalf("cancelled shutdown took %s", time.Since(start))
}
}
func TestWaitConnsQuietSeesInSend(t *testing.T) {
b, err := New(Options{})
if err != nil {
t.Fatal(err)
}
defer func() { _ = b.Close() }()
st := &connState{downCh: make(chan downItem, 1)}
st.inSend.Store(1)
ctx, cancel := context.WithTimeout(context.Background(), 40*time.Millisecond)
defer cancel()
if b.waitConnsQuiet(ctx, []*connState{st}) {
t.Fatal("inSend should keep shutdown from treating the conn as quiet")
}
st.inSend.Store(0)
ctx2, cancel2 := context.WithTimeout(context.Background(), 200*time.Millisecond)
defer cancel2()
if !b.waitConnsQuiet(ctx2, []*connState{st}) {
t.Fatal("quiet when inSend is 0 and queues are empty")
}
}
func TestEffectivePayloadLimitSubtractsOverhead(t *testing.T) { func TestEffectivePayloadLimitSubtractsOverhead(t *testing.T) {
got := EffectivePayloadLimit(200, 0) got := EffectivePayloadLimit(200, 0)
if got != 200-packetOverheadBudget { if got != 200-packetOverheadBudget {
+97 -5
View File
@@ -141,9 +141,15 @@ type connState struct {
downStop chan struct{} downStop chan struct{}
downDone chan struct{} downDone chan struct{}
downBytes atomic.Int64 downBytes atomic.Int64
sentPub atomic.Int64 wirePending atomic.Int64 // Publish 入 mochi outbound 后、OnPacketSent 前
inSend atomic.Int32 // downLoop 已取出帧、尚未从 sendOne 返回
mu sync.Mutex mu sync.Mutex
// 带断开的下行帧:只等本帧 OnPacketSent,不用连接级计数。
writeWaitCh chan struct{}
writeWaitPayload []byte
writeWaitPID uint16 // 非 0 时优先按 packet id 匹配
handshakeTimer *time.Timer handshakeTimer *time.Timer
} }
@@ -232,28 +238,114 @@ func (b *Broker) Close() error {
} }
// Shutdown 向所有连接发 MQTT 5 0x8B 后关闭。完整 HTTP 停机顺序见 L-03。 // Shutdown 向所有连接发 MQTT 5 0x8B 后关闭。完整 HTTP 停机顺序见 L-03。
// ctx 未取消时先等下行队列、正在 sendOne 的帧与 wirePending 排空(若有截止时间则预留约 1 秒),
// 再 DisconnectClient,并在截止前等连接拆掉;ctx 已取消则发完即 Close,不等待。
func (b *Broker) Shutdown(ctx context.Context) error { func (b *Broker) Shutdown(ctx context.Context) error {
if b.closed.Load() { if b.closed.Load() {
return nil return nil
} }
if ctx == nil {
ctx = context.Background()
}
b.connsMu.RLock() b.connsMu.RLock()
states := make([]*connState, 0, len(b.byClient))
clients := make([]*mqtt.Client, 0, len(b.byClient)) clients := make([]*mqtt.Client, 0, len(b.byClient))
for cl := range b.byClient { for cl, st := range b.byClient {
if cl != nil { if cl != nil {
clients = append(clients, cl) clients = append(clients, cl)
} }
if st != nil {
states = append(states, st)
}
} }
b.connsMu.RUnlock() b.connsMu.RUnlock()
for _, st := range states {
st.mu.Lock()
st.closing = true
st.mu.Unlock()
}
alreadyCancelled := false
select {
case <-ctx.Done():
alreadyCancelled = true
default:
}
var waitErr error
if !alreadyCancelled {
// 留出约 1 秒给 0x8B 写出,避免排空把整个截止时间用完后立刻 Close。
quietCtx := ctx
cancelQuiet := func() {}
if dl, ok := ctx.Deadline(); ok && time.Until(dl) > time.Second {
quietCtx, cancelQuiet = context.WithDeadline(ctx, dl.Add(-time.Second))
}
if !b.waitConnsQuiet(quietCtx, states) && ctx.Err() != nil {
waitErr = ctx.Err()
}
cancelQuiet()
}
for _, cl := range clients { for _, cl := range clients {
_ = b.server.DisconnectClient(cl, packets.ErrServerShuttingDown) _ = b.server.DisconnectClient(cl, packets.ErrServerShuttingDown)
} }
if ctx != nil {
if !alreadyCancelled && ctx.Err() == nil {
if err := b.waitConnsGone(ctx); err != nil {
waitErr = err
}
}
closeErr := b.Close()
if waitErr != nil {
return waitErr
}
return closeErr
}
func (b *Broker) waitConnsQuiet(ctx context.Context, states []*connState) bool {
for {
quiet := true
for _, st := range states {
if st.inSend.Load() > 0 || st.wirePending.Load() > 0 {
quiet = false
break
}
st.mu.Lock()
ch := st.downCh
st.mu.Unlock()
if len(ch) > 0 {
quiet = false
break
}
}
if quiet {
return true
}
select { select {
case <-ctx.Done(): case <-ctx.Done():
default: return false
case <-time.After(2 * time.Millisecond):
}
}
}
func (b *Broker) waitConnsGone(ctx context.Context) error {
for {
b.connsMu.RLock()
n := len(b.byClient)
b.connsMu.RUnlock()
if n == 0 {
return nil
}
select {
case <-ctx.Done():
return ctx.Err()
case <-time.After(2 * time.Millisecond):
} }
} }
return b.Close()
} }
// AttachTCP 把裸 TCP/TLS 连接交给 mochi;阻塞到连接结束。 // AttachTCP 把裸 TCP/TLS 连接交给 mochi;阻塞到连接结束。
+117 -12
View File
@@ -1,10 +1,12 @@
package broker package broker
import ( import (
"bytes"
"context" "context"
"time" "time"
"git.asio.asia/nixevol/NixMsg/internal/app/port" "git.asio.asia/nixevol/NixMsg/internal/app/port"
"github.com/mochi-mqtt/server/v2/packets"
) )
type downItem struct { type downItem struct {
@@ -116,6 +118,8 @@ func (st *connState) sendOne(b *Broker, item downItem) {
st.signalSent(item) st.signalSent(item)
return return
} }
st.inSend.Add(1)
defer st.inSend.Add(-1)
large := len(item.payload) > largeFrameBytes large := len(item.payload) > largeFrameBytes
if large { if large {
if err := b.acquireLarge(context.Background()); err != nil { if err := b.acquireLarge(context.Background()); err != nil {
@@ -130,7 +134,11 @@ func (st *connState) sendOne(b *Broker, item downItem) {
st.mu.Unlock() st.mu.Unlock()
} }
topic := downTopic(st.endpointID) topic := downTopic(st.endpointID)
before := st.sentPub.Load() var waitCh chan struct{}
if item.disconnect != "" {
waitCh = st.armWriteWait(item.payload)
}
st.wirePending.Add(1)
var err error var err error
for !b.closed.Load() { for !b.closed.Load() {
select { select {
@@ -150,6 +158,14 @@ func (st *connState) sendOne(b *Broker, item downItem) {
} }
break break
} }
if err != nil {
st.wirePending.Add(-1)
st.clearWriteWait(waitCh)
} else if waitCh != nil {
if pid, ok := st.lookupInflightPID(item.payload); ok {
st.setWriteWaitPID(waitCh, pid)
}
}
if large { if large {
b.finishLargePublish(st) b.finishLargePublish(st)
} }
@@ -158,23 +174,112 @@ func (st *connState) sendOne(b *Broker, item downItem) {
} }
st.signalSent(item) st.signalSent(item)
if err == nil && item.disconnect != "" { if err == nil && item.disconnect != "" {
st.waitPacketWritten(before) st.waitWriteDone(waitCh)
_ = b.Disconnect(context.Background(), st.endpointID, st.connID, item.disconnect) _ = b.Disconnect(context.Background(), st.endpointID, st.connID, item.disconnect)
} }
} }
func (st *connState) waitPacketWritten(before int64) { func (st *connState) armWriteWait(payload []byte) chan struct{} {
deadline := time.Now().Add(2 * time.Second) ch := make(chan struct{})
for time.Now().Before(deadline) { st.mu.Lock()
if st.sentPub.Load() > before { st.writeWaitCh = ch
return st.writeWaitPayload = payload
} st.writeWaitPID = 0
select { st.mu.Unlock()
case <-st.downStop: return ch
return }
case <-time.After(2 * time.Millisecond):
func (st *connState) setWriteWaitPID(ch chan struct{}, pid uint16) {
st.mu.Lock()
if st.writeWaitCh == ch {
st.writeWaitPID = pid
}
st.mu.Unlock()
}
func (st *connState) clearWriteWait(ch chan struct{}) {
if ch == nil {
return
}
st.mu.Lock()
if st.writeWaitCh == ch {
st.writeWaitCh = nil
st.writeWaitPayload = nil
st.writeWaitPID = 0
}
st.mu.Unlock()
}
func (st *connState) waitWriteDone(ch chan struct{}) {
if ch == nil {
return
}
defer st.clearWriteWait(ch)
deadline := time.NewTimer(2 * time.Second)
defer deadline.Stop()
select {
case <-ch:
case <-st.downStop:
case <-deadline.C:
}
}
func (st *connState) notePacketSent(pk packets.Packet) {
if pk.FixedHeader.Type == packets.Publish {
for {
cur := st.wirePending.Load()
if cur <= 0 {
break
}
if st.wirePending.CompareAndSwap(cur, cur-1) {
break
}
} }
} }
st.mu.Lock()
ch := st.writeWaitCh
pid := st.writeWaitPID
want := st.writeWaitPayload
st.mu.Unlock()
if ch == nil || pk.FixedHeader.Type != packets.Publish {
return
}
match := false
if pid != 0 {
match = pk.PacketID == pid
} else if want != nil {
match = bytes.Equal(pk.Payload, want)
}
if !match {
return
}
st.mu.Lock()
if st.writeWaitCh == ch {
st.writeWaitCh = nil
st.writeWaitPayload = nil
st.writeWaitPID = 0
}
st.mu.Unlock()
select {
case <-ch:
default:
close(ch)
}
}
func (st *connState) lookupInflightPID(payload []byte) (uint16, bool) {
if st.client == nil || st.client.State.Inflight == nil {
return 0, false
}
for _, pk := range st.client.State.Inflight.GetAll(false) {
if pk.FixedHeader.Type != packets.Publish {
continue
}
if bytes.Equal(pk.Payload, payload) {
return pk.PacketID, pk.PacketID != 0
}
}
return 0, false
} }
func (st *connState) signalSent(item downItem) { func (st *connState) signalSent(item downItem) {
+2 -4
View File
@@ -175,14 +175,11 @@ func (h *nixHook) OnSubscribed(cl *mqtt.Client, pk packets.Packet, reasonCodes [
} }
func (h *nixHook) OnPacketSent(cl *mqtt.Client, pk packets.Packet, _ []byte) { func (h *nixHook) OnPacketSent(cl *mqtt.Client, pk packets.Packet, _ []byte) {
if pk.FixedHeader.Type != packets.Publish {
return
}
h.b.connsMu.RLock() h.b.connsMu.RLock()
st := h.b.byClient[cl] st := h.b.byClient[cl]
h.b.connsMu.RUnlock() h.b.connsMu.RUnlock()
if st != nil { if st != nil {
st.sentPub.Add(1) st.notePacketSent(pk)
} }
} }
@@ -254,6 +251,7 @@ func (h *nixHook) OnDisconnect(cl *mqtt.Client, err error, _ bool) {
lk.Lock() lk.Lock()
st.stopDownLoop() st.stopDownLoop()
lk.Unlock() lk.Unlock()
st.wirePending.Store(0)
h.b.releaseAllLarge(st) h.b.releaseAllLarge(st)
h.b.cancelHandshakeDeadline(st.endpointID, st.connID) h.b.cancelHandshakeDeadline(st.endpointID, st.connID)
+22 -9
View File
@@ -36,10 +36,11 @@ type writeJob struct {
type Queue struct { type Queue struct {
db *sql.DB db *sql.DB
ch chan writeJob ch chan writeJob
done chan struct{} done chan struct{}
closed atomic.Bool closing chan struct{} // Close 时关闭,唤醒持读锁阻塞在发送上的 enqueue
sendMu sync.RWMutex closed atomic.Bool
sendMu sync.RWMutex
mu sync.Mutex mu sync.Mutex
ready bool ready bool
@@ -52,11 +53,16 @@ type Queue struct {
// NewQueue 创建合并写入队列并启动写 goroutine。 // NewQueue 创建合并写入队列并启动写 goroutine。
func NewQueue(db *sql.DB) *Queue { func NewQueue(db *sql.DB) *Queue {
return newQueue(db, queueBuffSize)
}
func newQueue(db *sql.DB, buffSize int) *Queue {
q := &Queue{ q := &Queue{
db: db, db: db,
ch: make(chan writeJob, queueBuffSize), ch: make(chan writeJob, buffSize),
done: make(chan struct{}), done: make(chan struct{}),
ready: true, closing: make(chan struct{}),
ready: true,
} }
go q.loop() go q.loop()
return q return q
@@ -115,10 +121,15 @@ func (q *Queue) enqueue(job writeJob) error {
return ErrQueueClosed return ErrQueueClosed
} }
q.addPending(1) q.addPending(1)
// 通道满时不得只堵在发送上持有读锁:Close 需要写锁关闭 q.ch。
select { select {
case q.ch <- job: case q.ch <- job:
q.sendMu.RUnlock() q.sendMu.RUnlock()
return nil return nil
case <-q.closing:
q.addPending(-1)
q.sendMu.RUnlock()
return ErrQueueClosed
case <-job.ctx.Done(): case <-job.ctx.Done():
q.addPending(-1) q.addPending(-1)
q.sendMu.RUnlock() q.sendMu.RUnlock()
@@ -376,11 +387,13 @@ func (q *Queue) Drain(ctx context.Context) error {
} }
// Close 关闭队列:不再接受新任务,并等待写 goroutine 处理完已入队任务后退出。 // Close 关闭队列:不再接受新任务,并等待写 goroutine 处理完已入队任务后退出。
// 在写锁内关闭数据通道,避免并发 Do 向已关闭 channel 发送而 panic。 // 先关闭 closing 唤醒因通道满而阻塞的发送方并释放读锁,再在无发送者时关闭数据通道,
// 避免向已关闭 channel 发送而 panic,也避免与持读锁的 enqueue 死锁。
func (q *Queue) Close() error { func (q *Queue) Close() error {
if q.closed.Swap(true) { if q.closed.Swap(true) {
return nil return nil
} }
close(q.closing)
q.sendMu.Lock() q.sendMu.Lock()
close(q.ch) close(q.ch)
q.sendMu.Unlock() q.sendMu.Unlock()
+134
View File
@@ -289,3 +289,137 @@ ON CONFLICT(key) DO UPDATE SET value = excluded.value, updated_at = excluded.upd
} }
} }
} }
// openSmallQueue 用小缓冲队列替换默认队列,便于测满通道时的 Close/Drain。
func openSmallQueue(t *testing.T, buf int) *DB {
t.Helper()
dir := t.TempDir()
db, err := Open(dir, "FULL")
if err != nil {
t.Fatal(err)
}
if err := db.Queue.Close(); err != nil {
t.Fatal(err)
}
db.Queue = newQueue(db.Write, buf)
return db
}
func TestQueueCloseUnblocksFullChannel(t *testing.T) {
t.Parallel()
const buf = 4
db := openSmallQueue(t, buf)
defer func() { _ = db.Close() }()
q := db.Queue
hold := make(chan struct{})
blockerStarted := make(chan struct{})
blockerErr := make(chan error, 1)
go func() {
blockerErr <- q.Do(context.Background(), func(tx *sql.Tx) error {
close(blockerStarted)
<-hold
return nil
})
}()
<-blockerStarted
ctx := context.Background()
var fillWG sync.WaitGroup
for i := 0; i < buf; i++ {
fillWG.Add(1)
go func() {
defer fillWG.Done()
_ = q.Do(ctx, func(tx *sql.Tx) error { return nil })
}()
}
deadline := time.Now().Add(2 * time.Second)
for q.Len() < buf+1 && time.Now().Before(deadline) {
time.Sleep(2 * time.Millisecond)
}
if q.Len() < buf+1 {
t.Fatalf("channel not full: len=%d", q.Len())
}
blockedErr := make(chan error, 1)
go func() {
blockedErr <- q.Do(ctx, func(tx *sql.Tx) error { return nil })
}()
// 等额外 Do 堵在 enqueue(pending 超过通道容量+正在执行的一条)。
deadline = time.Now().Add(2 * time.Second)
for q.Len() < buf+2 && time.Now().Before(deadline) {
time.Sleep(2 * time.Millisecond)
}
closeDone := make(chan error, 1)
go func() { closeDone <- q.Close() }()
select {
case err := <-blockedErr:
if !errors.Is(err, ErrQueueClosed) {
t.Fatalf("blocked Do: %v", err)
}
case <-time.After(2 * time.Second):
t.Fatal("Close did not unblock full-channel enqueue")
}
close(hold)
select {
case err := <-closeDone:
if err != nil {
t.Fatal(err)
}
case <-time.After(2 * time.Second):
t.Fatal("Close hung after writer released")
}
<-blockerErr
fillWG.Wait()
}
func TestQueueDrainTimeoutWhileWriterBlocked(t *testing.T) {
t.Parallel()
const buf = 4
db := openSmallQueue(t, buf)
defer func() { _ = db.Close() }()
q := db.Queue
hold := make(chan struct{})
started := make(chan struct{})
go func() {
_ = q.Do(context.Background(), func(tx *sql.Tx) error {
close(started)
<-hold
return nil
})
}()
<-started
ctx := context.Background()
var wg sync.WaitGroup
for i := 0; i < buf; i++ {
wg.Add(1)
go func() {
defer wg.Done()
_ = q.Do(ctx, func(tx *sql.Tx) error { return nil })
}()
}
deadline := time.Now().Add(2 * time.Second)
for q.Len() < buf+1 && time.Now().Before(deadline) {
time.Sleep(2 * time.Millisecond)
}
drainCtx, cancel := context.WithTimeout(context.Background(), 80*time.Millisecond)
defer cancel()
start := time.Now()
err := q.Drain(drainCtx)
elapsed := time.Since(start)
if !errors.Is(err, context.DeadlineExceeded) {
t.Fatalf("Drain err=%v want deadline", err)
}
if elapsed > 500*time.Millisecond {
t.Fatalf("Drain took %s, should return on timeout", elapsed)
}
close(hold)
wg.Wait()
}
+60
View File
@@ -156,6 +156,66 @@ func TestFrameTooLargeLocal(t *testing.T) {
} }
} }
func TestPublishUpFailRetriesWithNewRID(t *testing.T) {
fake := NewFakeTransport()
fake.FailSendPublishN(1)
c := connectFake(t, fake)
defer c.Close()
at := time.UnixMilli(1_700_000_000_000)
fixedID := "msg-publish-retry"
done := make(chan struct{})
var firstRID, secondRID string
go func() {
defer close(done)
deadline := time.Now().Add(3 * time.Second)
for time.Now().Before(deadline) {
sends := fake.FindUp("send")
if len(sends) < 2 {
time.Sleep(5 * time.Millisecond)
continue
}
first, second := sends[0], sends[1]
firstRID, _ = first["rid"].(string)
secondRID, _ = second["rid"].(string)
if first["id"] != fixedID || second["id"] != fixedID {
t.Errorf("id changed: %v -> %v", first["id"], second["id"])
}
if first["send_at_ms"] != second["send_at_ms"] {
t.Errorf("send_at_ms changed: %v -> %v", first["send_at_ms"], second["send_at_ms"])
}
body1, _ := json.Marshal(first["body"])
body2, _ := json.Marshal(second["body"])
if string(body1) != string(body2) {
t.Errorf("body changed: %s -> %s", body1, body2)
}
if firstRID == "" || firstRID == secondRID {
t.Errorf("rid not regenerated: %q -> %q", firstRID, secondRID)
}
fake.ReplyOK(secondRID, map[string]any{
"id": fixedID, "send_at_ms": second["send_at_ms"], "state": "scheduled",
})
return
}
t.Error("timed out waiting for send retry")
}()
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
res, err := c.Send(ctx, Target{Kind: "endpoint", ID: "b"}, Body{Enc: "utf8", Data: "hi"}, SendOptions{ID: fixedID, SendAt: &at})
<-done
if err != nil {
t.Fatalf("Send hung or failed: %v", err)
}
if res.ID != fixedID {
t.Fatalf("result id=%q want %q", res.ID, fixedID)
}
if firstRID == "" || secondRID == "" || firstRID == secondRID {
t.Fatalf("rids=%q/%q", firstRID, secondRID)
}
}
func TestResendKeepsIDAndSendAt(t *testing.T) { func TestResendKeepsIDAndSendAt(t *testing.T) {
fake := NewFakeTransport() fake := NewFakeTransport()
c := connectFake(t, fake) c := connectFake(t, fake)
+21
View File
@@ -3,6 +3,7 @@ package nixmsg
import ( import (
"context" "context"
"encoding/json" "encoding/json"
"errors"
"sync" "sync"
"sync/atomic" "sync/atomic"
"time" "time"
@@ -29,6 +30,8 @@ type FakeTransport struct {
MaxFrameBytes int MaxFrameBytes int
HelloDelay time.Duration HelloDelay time.Duration
ReceiveMaximumSet bool ReceiveMaximumSet bool
// failSendLeft 接下来若干次 type=send 的 PublishUp 返回错误(仍记入 up)。
failSendLeft int
} }
type fakeConnect struct { type fakeConnect struct {
@@ -65,6 +68,13 @@ func (f *FakeTransport) Start(ctx context.Context, cfg transportConfig) error {
return nil return nil
} }
// FailSendPublishN 让接下来 n 次 send 帧 PublishUp 失败(hello 等其它类型不受影响)。
func (f *FakeTransport) FailSendPublishN(n int) {
f.mu.Lock()
f.failSendLeft = n
f.mu.Unlock()
}
func (f *FakeTransport) PublishUp(payload []byte) error { func (f *FakeTransport) PublishUp(payload []byte) error {
f.mu.Lock() f.mu.Lock()
cp := append([]byte(nil), payload...) cp := append([]byte(nil), payload...)
@@ -80,6 +90,17 @@ func (f *FakeTransport) PublishUp(payload []byte) error {
if auto && head.Type == "hello" && head.RID != "" { if auto && head.Type == "hello" && head.RID != "" {
f.replyHello(head.RID) f.replyHello(head.RID)
} }
if head.Type == "send" {
f.mu.Lock()
fail := f.failSendLeft > 0
if fail {
f.failSendLeft--
}
f.mu.Unlock()
if fail {
return errors.New("publish failed")
}
}
return nil return nil
} }
+16 -11
View File
@@ -202,19 +202,24 @@ func (c *Client) dispatchSend(tr transport, item *sendItem, rid string, payload
if err := tr.PublishUp(payload); err != nil { if err := tr.PublishUp(payload); err != nil {
c.mu.Lock() c.mu.Lock()
delete(c.pending, rid) delete(c.pending, rid)
if item.epoch == epoch && item.inflight { if item.epoch != epoch || !item.inflight {
item.inflight = false c.mu.Unlock()
if c.inflight > 0 { return
c.inflight--
}
if c.stopReconnect {
errStop := c.stopErrLocked()
c.mu.Unlock()
c.finishSend(item, SendResult{}, errStop)
return
}
} }
item.inflight = false
if c.inflight > 0 {
c.inflight--
}
if c.stopReconnect {
errStop := c.stopErrLocked()
c.mu.Unlock()
c.finishSend(item, SendResult{}, errStop)
return
}
// 保留 id/body/send_at_ms,换新 rid 后继续泵,避免 Send 永久挂起。
c.regenerateSendLocked(item)
c.mu.Unlock() c.mu.Unlock()
c.drainSendQueue()
return return
} }
@@ -143,6 +143,12 @@ public final class Client {
return lastStopCode == null ? "" : lastStopCode; return lastStopCode == null ? "" : lastStopCode;
} }
int inflightSendsForTest() {
synchronized (lock) {
return inflightSends;
}
}
void setSendQueueLimitForTest(int n) { void setSendQueueLimitForTest(int n) {
sendQueueLimit = n; sendQueueLimit = n;
} }
@@ -211,6 +217,19 @@ public final class Client {
} }
} }
if (handshakeError != null) { if (handshakeError != null) {
synchronized (lock) {
stopReconnect = true;
wantConnected = false;
if (lastStopCode == null || lastStopCode.isEmpty()) {
if (handshakeError instanceof NixMsgException) {
lastStopCode = ((NixMsgException) handshakeError).getCode();
lastStopErr = (NixMsgException) handshakeError;
} else {
lastStopCode = "not_connected";
lastStopErr = new NixMsgException("not_connected", handshakeError.getMessage());
}
}
}
if (handshakeError instanceof NixMsgException) { if (handshakeError instanceof NixMsgException) {
throw (NixMsgException) handshakeError; throw (NixMsgException) handshakeError;
} }
@@ -229,6 +248,8 @@ public final class Client {
synchronized (lock) { synchronized (lock) {
stopReconnect = true; stopReconnect = true;
wantConnected = false; wantConnected = false;
lastStopCode = "not_connected";
lastStopErr = new NixMsgException("not_connected", "连接未成功: " + state);
} }
throw new NixMsgException("not_connected", "连接未成功: " + state); throw new NixMsgException("not_connected", "连接未成功: " + state);
} }
@@ -704,7 +725,9 @@ public final class Client {
transport.publish(Protocol.upTopic(endpointId), Protocol.dumps(hello)); transport.publish(Protocol.upTopic(endpointId), Protocol.dumps(hello));
await(p.future, connectTimeoutMs); await(p.future, connectTimeoutMs);
if (p.error != null) { if (p.error != null) {
throw p.error instanceof RuntimeException ? (RuntimeException) p.error : new NixMsgException("busy", p.error.getMessage()); throw p.error instanceof RuntimeException
? (RuntimeException) p.error
: new NixMsgException("not_connected", p.error.getMessage());
} }
if (p.response == null) { if (p.response == null) {
throw new NixMsgException("not_connected", "握手超时"); throw new NixMsgException("not_connected", "握手超时");
@@ -883,9 +906,9 @@ public final class Client {
if (!Boolean.TRUE.equals(frame.get("ok"))) { if (!Boolean.TRUE.equals(frame.get("ok"))) {
Map<String, Object> err = asMap(frame.get("error")); Map<String, Object> err = asMap(frame.get("error"));
if ("rate_limited".equals(str(err.get("code"), ""))) { if ("rate_limited".equals(str(err.get("code"), ""))) {
// 在途计数已在上方减过一次;只换新 rid 再排队,勿再减。
synchronized (lock) { synchronized (lock) {
pending.remove(rid); pending.remove(rid);
inflightSends = Math.max(0, inflightSends - 1);
p.rid = ""; p.rid = "";
p.response = null; p.response = null;
p.error = null; p.error = null;
@@ -41,7 +41,7 @@ public class K00Test {
} }
@Test @Test
public void testK00FirstConnectTimeout() { public void testK00FirstConnectTimeout() throws Exception {
FakeTransport tr = new FakeTransport(); FakeTransport tr = new FakeTransport();
tr.autoHello = null; tr.autoHello = null;
client = new Client(tr, true, Types.DEFAULT_MAX_FRAME, Types.CLIENT_NAME, 150); client = new Client(tr, true, Types.DEFAULT_MAX_FRAME, Types.CLIENT_NAME, 150);
@@ -51,6 +51,9 @@ public class K00Test {
} catch (NixMsgException e) { } catch (NixMsgException e) {
assertEquals("not_connected", e.getCode()); assertEquals("not_connected", e.getCode());
} }
int n = tr.connects.size();
Thread.sleep(350);
assertEquals("首次握手失败后不得自动重连", n, tr.connects.size());
tr.autoHello = new LinkedHashMap<String, Object>(); tr.autoHello = new LinkedHashMap<String, Object>();
tr.autoHello.put("server_time_ms", 1750000000000L); tr.autoHello.put("server_time_ms", 1750000000000L);
tr.autoHello.put("server_version", "0.1.0"); tr.autoHello.put("server_version", "0.1.0");
@@ -63,6 +66,108 @@ public class K00Test {
tr.autoHello.put("session_token", "nst_test_token"); tr.autoHello.put("session_token", "nst_test_token");
client.connectSync("ws://example.test/mqtt", "ep1", "p", null, false); client.connectSync("ws://example.test/mqtt", "ep1", "p", null, false);
assertEquals(ConnectionState.ONLINE, client.getState()); assertEquals(ConnectionState.ONLINE, client.getState());
assertTrue(tr.connects.size() > n);
}
@Test
public void testK00RateLimitedInflightOnce() throws Exception {
Types.ReconnectBackoff.disableJitterForTest();
FakeTransport tr = new FakeTransport();
tr.autoSendOk = false;
final List<String> holdRids = new CopyOnWriteArrayList<String>();
final List<String> limitedRids = new CopyOnWriteArrayList<String>();
tr.onUp(new java.util.function.Function<Map<String, Object>, Map<String, Object>>() {
@Override
public Map<String, Object> apply(Map<String, Object> frame) {
if (!"send".equals(String.valueOf(frame.get("type")))) {
return null;
}
String rid = String.valueOf(frame.get("rid"));
String id = String.valueOf(frame.get("id"));
if ("hold".equals(id)) {
holdRids.add(rid);
return null; // 保持在途
}
limitedRids.add(rid);
if (limitedRids.size() == 1) {
Map<String, Object> err = new LinkedHashMap<String, Object>();
err.put("code", "rate_limited");
err.put("message", "slow");
Map<String, Object> resp = new LinkedHashMap<String, Object>();
resp.put("v", 1);
resp.put("type", "resp");
resp.put("rid", rid);
resp.put("ok", false);
resp.put("error", err);
return resp;
}
Map<String, Object> data = new LinkedHashMap<String, Object>();
data.put("id", frame.get("id"));
data.put("send_at_ms", frame.get("send_at_ms"));
data.put("state", "scheduled");
Map<String, Object> resp = new LinkedHashMap<String, Object>();
resp.put("v", 1);
resp.put("type", "resp");
resp.put("rid", rid);
resp.put("ok", true);
resp.put("data", data);
return resp;
}
});
connectOnline(tr);
Thread holder = new Thread(new Runnable() {
@Override
public void run() {
try {
SendOptions opt = new SendOptions();
opt.messageId = "hold";
opt.sendAtMs = 1700000000001L;
client.sendSync(new Target("endpoint", "ep2"), new Body("hold"), opt);
} catch (Exception ignored) {
}
}
});
holder.setDaemon(true);
holder.start();
long deadline = System.currentTimeMillis() + 2000;
while (holdRids.isEmpty() && System.currentTimeMillis() < deadline) {
Thread.sleep(10);
}
assertFalse(holdRids.isEmpty());
assertEquals(1, client.inflightSendsForTest());
final SendOptions opt = new SendOptions();
opt.messageId = "lim";
opt.sendAtMs = 1700000000002L;
Thread sender = new Thread(new Runnable() {
@Override
public void run() {
try {
client.sendSync(new Target("endpoint", "ep2"), new Body("lim"), opt);
} catch (Exception ignored) {
}
}
});
sender.setDaemon(true);
sender.start();
deadline = System.currentTimeMillis() + 3000;
while (limitedRids.size() < 1 && System.currentTimeMillis() < deadline) {
Thread.sleep(10);
}
assertTrue(limitedRids.size() >= 1);
// rate_limited 后只应减 1:仍剩 hold 那一条在途
deadline = System.currentTimeMillis() + 500;
int seen = -1;
while (System.currentTimeMillis() < deadline) {
seen = client.inflightSendsForTest();
if (seen == 1) {
break;
}
Thread.sleep(10);
}
assertEquals("rate_limited 后在途应只减 1", 1, seen);
client.close();
holder.join(1000);
sender.join(1000);
} }
@Test @Test
+2
View File
@@ -739,10 +739,12 @@ export class Client {
this.finishSendErr(item, this.stopErr()); this.finishSendErr(item, this.stopErr());
return; return;
} }
// 传输失败但连接未断:换 rid 后立刻再泵,否则 hello 不再走、send 会挂起
item.inflight = false; item.inflight = false;
this.inflight = Math.max(0, this.inflight - 1); this.inflight = Math.max(0, this.inflight - 1);
this.pending.delete(rid); this.pending.delete(rid);
this.regenerateSend(item); this.regenerateSend(item);
void this.drainSendQueue();
return; return;
} }
this.finishSendErr(item, e); this.finishSendErr(item, e);
+40
View File
@@ -126,6 +126,46 @@ describe("nixmsg sdk", () => {
await c.close(); await c.close();
}, 10000); }, 10000);
it("publishUp transport fail retries with new rid", async () => {
const fake = new FakeTransport();
const c = await connectFake(fake);
const at = new Date(1_700_000_000_000);
const attempts: Array<Record<string, unknown>> = [];
let failOnce = true;
fake.publishUpImpl = async (s) => {
const m = JSON.parse(s) as Record<string, unknown>;
if (m.type !== "send") return;
attempts.push(m);
if (failOnce) {
failOnce = false;
throw new Error("transient publish");
}
};
const timer = setInterval(() => {
const sends = fake.findUp("send");
if (sends.length < 1) return;
const last = sends[sends.length - 1]!;
fake.replyOK(String(last.rid), {
id: last.id,
send_at_ms: last.send_at_ms,
state: "scheduled",
});
}, 5);
const res = await c.send(
{ kind: "endpoint", id: "b" },
{ enc: "utf8", data: "hi" },
{ sendAt: at },
);
clearInterval(timer);
expect(attempts.length).toBeGreaterThanOrEqual(2);
expect(String(attempts[1]!.rid)).not.toBe(String(attempts[0]!.rid));
expect(attempts[1]!.id).toBe(attempts[0]!.id);
expect(attempts[1]!.send_at_ms).toBe(attempts[0]!.send_at_ms);
expect((attempts[1]!.body as { data: string }).data).toBe("hi");
expect(res.id).toBe(String(attempts[0]!.id));
await c.close();
}, 5000);
it("register HTTP from ws url", async () => { it("register HTTP from ws url", async () => {
const srv = createServer((req, res) => { const srv = createServer((req, res) => {
expect(req.url).toBe("/api/client/register"); expect(req.url).toBe("/api/client/register");
+16 -2
View File
@@ -227,6 +227,15 @@ class Client:
raise NixMsgError("not_connected", "连接超时") raise NixMsgError("not_connected", "连接超时")
err = self._handshake_error err = self._handshake_error
if err: if err:
with self._lock:
self._stop_reconnect = True
self._want_connected = False
if not self._last_stop_code:
code = getattr(err, "code", None) or "not_connected"
self._last_stop_code = str(code)
self._last_stop_err = err if isinstance(err, NixMsgError) else NixMsgError(
"not_connected", str(err)
)
raise err raise err
if self._state not in (ConnectionState.ONLINE,): if self._state not in (ConnectionState.ONLINE,):
if self._state == ConnectionState.AUTH_FAILED: if self._state == ConnectionState.AUTH_FAILED:
@@ -236,6 +245,11 @@ class Client:
) )
if self._state == ConnectionState.KICKED: if self._state == ConnectionState.KICKED:
raise NixMsgError("taken_over", "会话被接管") raise NixMsgError("taken_over", "会话被接管")
with self._lock:
self._stop_reconnect = True
self._want_connected = False
self._last_stop_code = "not_connected"
self._last_stop_err = NixMsgError("not_connected", f"连接未成功: {self._state.value}")
raise NixMsgError("not_connected", f"连接未成功: {self._state.value}") raise NixMsgError("not_connected", f"连接未成功: {self._state.value}")
def close(self) -> None: def close(self) -> None:
@@ -576,7 +590,7 @@ class Client:
except Exception: except Exception:
pass pass
with self._lock: with self._lock:
self._handshake_error = NixMsgError("busy", "连接超时") self._handshake_error = NixMsgError("not_connected", "连接超时")
return return
def _on_transport_connected(self) -> None: def _on_transport_connected(self) -> None:
@@ -598,7 +612,7 @@ class Client:
self._pending[rid] = pending self._pending[rid] = pending
self._transport.publish(up_topic(self._endpoint_id), dumps(hello)) self._transport.publish(up_topic(self._endpoint_id), dumps(hello))
if not pending.event.wait(self._connect_timeout_s): if not pending.event.wait(self._connect_timeout_s):
raise NixMsgError("busy", "握手超时") raise NixMsgError("not_connected", "握手超时")
if pending.error: if pending.error:
raise pending.error raise pending.error
assert pending.response is not None assert pending.response is not None
+22 -7
View File
@@ -8,16 +8,23 @@ from dataclasses import dataclass, field
from typing import Any, Callable, Optional, Protocol from typing import Any, Callable, Optional, Protocol
from urllib.parse import urlparse from urllib.parse import urlparse
from paho.mqtt.client import CallbackAPIVersion, Client as PahoClient, MQTT_ERR_SUCCESS, MQTTv5
from paho.mqtt.enums import MQTTErrorCode
from paho.mqtt.reasoncodes import ReasonCode
DownHandler = Callable[[bytes], None] DownHandler = Callable[[bytes], None]
ConnHandler = Callable[[], None] ConnHandler = Callable[[], None]
DiscHandler = Callable[[Optional[str], bool], None] DiscHandler = Callable[[Optional[str], bool], None]
# reason_code_str, stop_reconnect # reason_code_str, stop_reconnect
def _require_paho():
"""真实 MQTT 路径才加载 paho;假传输单测不依赖。"""
try:
from paho.mqtt.client import CallbackAPIVersion, Client as PahoClient, MQTT_ERR_SUCCESS, MQTTv5
from paho.mqtt.enums import MQTTErrorCode
from paho.mqtt.reasoncodes import ReasonCode
except ImportError as e:
raise ImportError("需要 paho-mqtt>=2.0(真实 MQTT 连接)") from e
return CallbackAPIVersion, PahoClient, MQTT_ERR_SUCCESS, MQTTv5, MQTTErrorCode, ReasonCode
@dataclass @dataclass
class ConnectParams: class ConnectParams:
url: str url: str
@@ -180,7 +187,7 @@ class PahoTransport:
"""paho-mqtt 2.x CallbackAPIVersion.VERSION2。""" """paho-mqtt 2.x CallbackAPIVersion.VERSION2。"""
def __init__(self) -> None: def __init__(self) -> None:
self._client: Optional[PahoClient] = None self._client: Any = None
self._on_connected: Optional[ConnHandler] = None self._on_connected: Optional[ConnHandler] = None
self._on_disconnected: Optional[DiscHandler] = None self._on_disconnected: Optional[DiscHandler] = None
self._on_down: Optional[DownHandler] = None self._on_down: Optional[DownHandler] = None
@@ -200,6 +207,7 @@ class PahoTransport:
self._on_down = on_down self._on_down = on_down
def connect(self, params: ConnectParams) -> None: def connect(self, params: ConnectParams) -> None:
CallbackAPIVersion, PahoClient, _, MQTTv5, _, _ = _require_paho()
self.disconnect() self.disconnect()
url = params.url url = params.url
u = urlparse(url if "://" in url else "ws://" + url) u = urlparse(url if "://" in url else "ws://" + url)
@@ -261,6 +269,7 @@ class PahoTransport:
# 等待连接结果由回调驱动;超时由 Client 层处理 # 等待连接结果由回调驱动;超时由 Client 层处理
def subscribe(self, topic: str) -> None: def subscribe(self, topic: str) -> None:
_, _, MQTT_ERR_SUCCESS, _, _, _ = _require_paho()
self._down_topic = topic self._down_topic = topic
if not self._client: if not self._client:
return return
@@ -273,6 +282,7 @@ class PahoTransport:
raise RuntimeError("subscribe timeout") raise RuntimeError("subscribe timeout")
def publish(self, topic: str, payload: bytes) -> None: def publish(self, topic: str, payload: bytes) -> None:
_, _, MQTT_ERR_SUCCESS, _, _, _ = _require_paho()
if not self._client: if not self._client:
raise RuntimeError("not connected") raise RuntimeError("not connected")
info = self._client.publish(topic, payload, qos=1) info = self._client.publish(topic, payload, qos=1)
@@ -338,9 +348,14 @@ def _reason_to_int(reason_code) -> Optional[int]:
return None return None
if isinstance(reason_code, int): if isinstance(reason_code, int):
return reason_code return reason_code
if isinstance(reason_code, ReasonCode): try:
_, _, _, _, MQTTErrorCode, ReasonCode = _require_paho()
except ImportError:
MQTTErrorCode = () # type: ignore[assignment,misc]
ReasonCode = () # type: ignore[assignment,misc]
if ReasonCode and isinstance(reason_code, ReasonCode):
return int(reason_code.value) return int(reason_code.value)
if isinstance(reason_code, MQTTErrorCode): if MQTTErrorCode and isinstance(reason_code, MQTTErrorCode):
return int(reason_code) return int(reason_code)
# paho 偶发其它包装 # paho 偶发其它包装
val = getattr(reason_code, "value", None) val = getattr(reason_code, "value", None)
+31
View File
@@ -30,11 +30,42 @@ class K00Tests(unittest.TestCase):
with self.assertRaises(NixMsgError) as cm: with self.assertRaises(NixMsgError) as cm:
c.connect("ws://example.test/mqtt", "ep1", password="p", wait=True) c.connect("ws://example.test/mqtt", "ep1", password="p", wait=True)
self.assertEqual(cm.exception.code, "not_connected") self.assertEqual(cm.exception.code, "not_connected")
n = len(tr.connects)
time.sleep(0.35)
self.assertEqual(len(tr.connects), n, "首次失败后不得自动重连")
tr.auto_accept = True tr.auto_accept = True
c.connect("ws://example.test/mqtt", "ep1", password="p", wait=True) c.connect("ws://example.test/mqtt", "ep1", password="p", wait=True)
self.assertEqual(c.state, ConnectionState.ONLINE) self.assertEqual(c.state, ConnectionState.ONLINE)
c.close() c.close()
def test_k00_first_hello_fail_stops_reconnect(self) -> None:
"""MQTT 已通但 hello 失败:connect 返回未连接,且后台不再连。"""
tr = FakeTransport()
tr.auto_hello = None
c = Client(transport=tr, connect_timeout_s=0.2)
with self.assertRaises(NixMsgError) as cm:
c.connect("ws://example.test/mqtt", "ep1", password="p", wait=True)
self.assertEqual(cm.exception.code, "not_connected")
n = len(tr.connects)
self.assertGreaterEqual(n, 1)
time.sleep(0.45)
self.assertEqual(len(tr.connects), n, "hello 失败后不得自动重连")
tr.auto_hello = {
"server_time_ms": 1_750_000_000_000,
"server_version": "0.1.0",
"max_body_bytes": 262144,
"max_meta_bytes": 4096,
"max_frame_bytes": 786432,
"max_ttl_seconds": 2592000,
"max_schedule_seconds": 31536000,
"ack_timeout_seconds": 300,
"session_token": "nst_retry",
}
c.connect("ws://example.test/mqtt", "ep1", password="p", wait=True)
self.assertEqual(c.state, ConnectionState.ONLINE)
self.assertGreater(len(tr.connects), n)
c.close()
def test_k00_taken_over_reason(self) -> None: def test_k00_taken_over_reason(self) -> None:
c, tr = self._connect() c, tr = self._connect()
got = [] got = []
+1 -1
View File
@@ -296,7 +296,7 @@ export function deleteGroup(id: string) {
export function addGroupMembers(id: string, memberIds: string[]) { export function addGroupMembers(id: string, memberIds: string[]) {
return run(() => return run(() =>
requestAdmin<{ failed: GroupFailed[] }>( requestAdmin<{ failed: GroupFailed[]; added: number }>(
`/api/admin/groups/${encodeURIComponent(id)}/members`, `/api/admin/groups/${encodeURIComponent(id)}/members`,
{ {
method: "POST", method: "POST",
+4 -2
View File
@@ -691,13 +691,14 @@ export const mockApi = {
return {}; return {};
}, },
async addGroupMembers(id: string, memberIds: string[]): Promise<{ failed: GroupFailed[] }> { async addGroupMembers(id: string, memberIds: string[]): Promise<{ failed: GroupFailed[]; added: number }> {
requireSession(); requireSession();
const g = groups.find((x) => x.id === id); const g = groups.find((x) => x.id === id);
if (!g) { if (!g) {
throw new ApiError("not_found", "群不存在", 404); throw new ApiError("not_found", "群不存在", 404);
} }
const failed: GroupFailed[] = []; const failed: GroupFailed[] = [];
let added = 0;
for (const mid of memberIds) { for (const mid of memberIds) {
if (!endpoints.some((e) => e.id === mid)) { if (!endpoints.some((e) => e.id === mid)) {
failed.push({ id: mid, code: "not_found" }); failed.push({ id: mid, code: "not_found" });
@@ -705,9 +706,10 @@ export const mockApi = {
} }
if (!g.members.some((m) => m.id === mid)) { if (!g.members.some((m) => m.id === mid)) {
g.members.push({ id: mid, joined_at_ms: now() }); g.members.push({ id: mid, joined_at_ms: now() });
added++;
} }
} }
return { failed }; return { failed, added };
}, },
async removeGroupMember(id: string, endpointId: string): Promise<Record<string, never>> { async removeGroupMember(id: string, endpointId: string): Promise<Record<string, never>> {
+8 -2
View File
@@ -233,9 +233,15 @@ async function onAddMembers() {
if (!ids.length) return; if (!ids.length) return;
const res = await addGroupMembers(detail.value.id, ids); const res = await addGroupMembers(detail.value.id, ids);
if (res.failed.length) { if (res.failed.length) {
message.warning(`部分失败:${formatFailed(res.failed)}`); if (res.added > 0) {
message.warning(`已加入 ${res.added} 人,部分失败:${formatFailed(res.failed)}`);
} else {
message.warning(`部分失败:${formatFailed(res.failed)}`);
}
} else if (res.added === 0) {
message.info("没有新成员加入");
} else { } else {
message.success("已加人"); message.success(`已加人 ${res.added} 人`);
} }
await loadDetail(detail.value.id); await loadDetail(detail.value.id);
await load(); await load();