Author SHA1 Message Date
Nixevol a14ca0d08d fix: 积压时等本帧写出再断开并让 Shutdown 等待 0x8B 2026-09-30 19:20:24 +08:00
11 changed files with 384 additions and 229 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)
+7 -7
View File
@@ -1719,12 +1719,12 @@ issue #3 未关闭,`feat/fix-3-downlink-deadlock` 未合入 `main`。下面是
- 备选方案:为每个未测子项补验收用例(本波不做,避免为变绿放松断言)。 - 备选方案:为每个未测子项补验收用例(本波不做,避免为变绿放松断言)。
- 影响:汇总改为通过 19、部分通过 4(F03/F08/F21/F22)、失败 0。F19 仍引用仓库内 SDK 清单、本波不重跑。 - 影响:汇总改为通过 19、部分通过 4(F03/F08/F21/F22)、失败 0。F19 仍引用仓库内 SDK 清单、本波不重跑。
### 复审修复 R3-06 ### 复审修复 R3-02
1. **Python/Java 首次握手失败应停止重连;Java rate_limited 在途计数只减一次** 1. **积压时 fatal/logout 与停机 0x8B 须等本帧写出**
- 日期:2026-09-30 - 日期:2026-09-30
- 原条款:DEVELOPMENT 第 9 节附录「首次连接超时或握手失败应停止重连,并向 connect() 返回未连接」;issue #70。 - 原条款:Gitea #66;B-04 / B-08。
- 实际做法:`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 重连公式。 - 实际做法:带断开的下行帧用本帧 `OnPacketSent` 完成信号(优先 packet id,否则按载荷匹配),不再用连接级 `sentPub` 总数。`Shutdown` 在 ctx 未取消时先等下行队列与 `wirePending` 排空,再 `DisconnectClient` 发 `0x8B` 并在截止前等连接拆掉;ctx 已取消则发完即 `Close`。`serve` 仍给 5 秒预算并记录非超时错误。未合 `feat/fix-3-downlink-deadlock`。
- 原因:原先抛错后工作线程仍按退避重连;Java 限速把在途计数减了两次。 - 原因:前面 PUBLISH 的 `OnPacketSent` 会让总数等待提前返回;`Shutdown` 对 ctx 非阻塞 select 使 5 秒预算用不上,有 outbound 积压时 `0x8B` 只进 outbuf 随 `Stop` 丢掉。
- 备选方案:在后台 `_attempt_connect`/`attemptConnect` 内首次失败即停(否决:应用再次 `connect()` 才应恢复,标志应在对外 `connect` 失败路径统一置位)。 - 备选方案:恢复固定 `Sleep`(否决);改 `PublishDown` 签名(否决)。
- 影响:仅 `sdk/python`、`sdk/java`;应用需再次调用 `connect()` 才会重连。 - 影响:队列/outbound 有积压时 fatal、logout 先到客户端再断开;停机在预算内尽量发出 `0x8B`,超时返回 ctx 错误而非空等。
+158
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,120 @@ 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 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 {
+87 -5
View File
@@ -141,9 +141,14 @@ 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 前
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 +237,105 @@ func (b *Broker) Close() error {
} }
// Shutdown 向所有连接发 MQTT 5 0x8B 后关闭。完整 HTTP 停机顺序见 L-03。 // Shutdown 向所有连接发 MQTT 5 0x8B 后关闭。完整 HTTP 停机顺序见 L-03。
// ctx 未取消时先等下行队列与 wirePending 排空,再 DisconnectClient(此时 outbound 空,
// 0x8B 直写套接字),并在截止前等连接拆掉;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 {
if !b.waitConnsQuiet(ctx, states) {
waitErr = ctx.Err()
}
}
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 && waitErr == nil {
waitErr = b.waitConnsGone(ctx)
}
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.wirePending.Load() > 0 {
quiet = false
break
}
st.mu.Lock()
ch := st.downCh
st.mu.Unlock()
if ch != nil && 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;阻塞到连接结束。
+110 -7
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 {
@@ -130,7 +132,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 +156,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 +172,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
st.writeWaitPayload = payload
st.writeWaitPID = 0
st.mu.Unlock()
return ch
}
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 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 { select {
case <-ch:
case <-st.downStop: 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 return
case <-time.After(2 * time.Millisecond): }
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)
@@ -143,12 +143,6 @@ 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;
} }
@@ -217,19 +211,6 @@ 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;
} }
@@ -248,8 +229,6 @@ 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);
} }
@@ -725,9 +704,7 @@ 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 throw p.error instanceof RuntimeException ? (RuntimeException) p.error : new NixMsgException("busy", p.error.getMessage());
? (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", "握手超时");
@@ -906,9 +883,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() throws Exception { public void testK00FirstConnectTimeout() {
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,9 +51,6 @@ 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");
@@ -66,108 +63,6 @@ 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 -16
View File
@@ -227,15 +227,6 @@ 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:
@@ -245,11 +236,6 @@ 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:
@@ -590,7 +576,7 @@ class Client:
except Exception: except Exception:
pass pass
with self._lock: with self._lock:
self._handshake_error = NixMsgError("not_connected", "连接超时") self._handshake_error = NixMsgError("busy", "连接超时")
return return
def _on_transport_connected(self) -> None: def _on_transport_connected(self) -> None:
@@ -612,7 +598,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("not_connected", "握手超时") raise NixMsgError("busy", "握手超时")
if pending.error: if pending.error:
raise pending.error raise pending.error
assert pending.response is not None assert pending.response is not None
+7 -22
View File
@@ -8,23 +8,16 @@ 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
@@ -187,7 +180,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: Any = None self._client: Optional[PahoClient] = 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
@@ -207,7 +200,6 @@ 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)
@@ -269,7 +261,6 @@ 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
@@ -282,7 +273,6 @@ 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)
@@ -348,14 +338,9 @@ 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
try: if isinstance(reason_code, ReasonCode):
_, _, _, _, 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 MQTTErrorCode and isinstance(reason_code, MQTTErrorCode): if 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,42 +30,11 @@ 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 = []