fix: 完成 broker 复审 B-03 至 B-12

每连接异步下发与背压、写出后断开、校验当前连接与订阅、生命周期串行、登录条件更新、闲置按在线计、认证超时并发与 Shutdown 0x8B。
This commit is contained in:
Nixevol
2026-09-30 16:21:05 +08:00
parent 0b9ce0359a
commit 91e887ba46
12 changed files with 1099 additions and 89 deletions
+5 -1
View File
@@ -153,7 +153,7 @@ func runServe(ctx context.Context, cfg config.Config) error {
Sessions: sessionTokens, Sessions: sessionTokens,
MaxScheduleSeconds: int64(cfg.Limits.MaxScheduleSeconds), MaxScheduleSeconds: int64(cfg.Limits.MaxScheduleSeconds),
Logger: slog.Default(), Logger: slog.Default(),
ConnControl: brk, ConnControl: nil, // B-04:踢线走 Session 钩子,避免 identity 20ms 异步 Disconnect
Downlink: brk, Downlink: brk,
ClientIP: func(r *http.Request) string { ClientIP: func(r *http.Request) string {
return httpx.ClientIP(r, trustedNets) return httpx.ClientIP(r, trustedNets)
@@ -310,6 +310,10 @@ func runServe(ctx context.Context, cfg config.Config) error {
<-ctx.Done() <-ctx.Done()
loopCancel() loopCancel()
// B-08:先对 MQTT 连接发 0x8B。HTTP Shutdown 与监听器完整停机顺序见 L-03。
shutCtx, shutCancel := context.WithTimeout(context.Background(), 5*time.Second)
_ = brk.Shutdown(shutCtx)
shutCancel()
_ = lnSrv.Close() _ = lnSrv.Close()
drainCtx, drainCancel := context.WithTimeout(context.Background(), 10*time.Second) drainCtx, drainCancel := context.WithTimeout(context.Background(), 10*time.Second)
defer drainCancel() defer drainCancel()
+42 -16
View File
@@ -5,12 +5,14 @@ import (
"encoding/json" "encoding/json"
"errors" "errors"
"log/slog" "log/slog"
"sync"
"git.asio.asia/nixevol/NixMsg/internal/app/group" "git.asio.asia/nixevol/NixMsg/internal/app/group"
"git.asio.asia/nixevol/NixMsg/internal/app/identity" "git.asio.asia/nixevol/NixMsg/internal/app/identity"
"git.asio.asia/nixevol/NixMsg/internal/app/message" "git.asio.asia/nixevol/NixMsg/internal/app/message"
"git.asio.asia/nixevol/NixMsg/internal/app/port" "git.asio.asia/nixevol/NixMsg/internal/app/port"
"git.asio.asia/nixevol/NixMsg/internal/app/presence" "git.asio.asia/nixevol/NixMsg/internal/app/presence"
"git.asio.asia/nixevol/NixMsg/internal/broker"
"git.asio.asia/nixevol/NixMsg/internal/metrics" "git.asio.asia/nixevol/NixMsg/internal/metrics"
"git.asio.asia/nixevol/NixMsg/internal/protocol" "git.asio.asia/nixevol/NixMsg/internal/protocol"
) )
@@ -25,9 +27,31 @@ type appUplink struct {
down port.Downlink down port.Downlink
log *slog.Logger log *slog.Logger
metrics *metrics.Registry metrics *metrics.Registry
lifeMu sync.Mutex
lifeLocks map[string]*sync.Mutex
hsMu sync.Mutex
handshake map[port.ConnID]string // 已 hello 的连接代号 → 端编号
}
func (u *appUplink) epLife(endpointID string) *sync.Mutex {
u.lifeMu.Lock()
defer u.lifeMu.Unlock()
if u.lifeLocks == nil {
u.lifeLocks = make(map[string]*sync.Mutex)
}
m := u.lifeLocks[endpointID]
if m == nil {
m = &sync.Mutex{}
u.lifeLocks[endpointID] = m
}
return m
} }
func (u *appUplink) OnSessionEstablished(ctx context.Context, conn port.ConnInfo) error { func (u *appUplink) OnSessionEstablished(ctx context.Context, conn port.ConnInfo) error {
lk := u.epLife(conn.EndpointID)
lk.Lock()
defer lk.Unlock()
u.conns.Set(conn.EndpointID, message.LiveConn{ u.conns.Set(conn.EndpointID, message.LiveConn{
ConnID: conn.ConnID, ConnID: conn.ConnID,
MaxPacketSize: conn.MaxPacketSize, MaxPacketSize: conn.MaxPacketSize,
@@ -36,6 +60,15 @@ func (u *appUplink) OnSessionEstablished(ctx context.Context, conn port.ConnInfo
} }
func (u *appUplink) OnHandshakeComplete(ctx context.Context, hs port.HandshakeInfo) error { func (u *appUplink) OnHandshakeComplete(ctx context.Context, hs port.HandshakeInfo) error {
lk := u.epLife(hs.EndpointID)
lk.Lock()
defer lk.Unlock()
u.hsMu.Lock()
if u.handshake == nil {
u.handshake = make(map[port.ConnID]string)
}
u.handshake[hs.ConnID] = hs.EndpointID
u.hsMu.Unlock()
live := message.LiveConn{ live := message.LiveConn{
ConnID: hs.ConnID, ConnID: hs.ConnID,
MaxReceiveBytes: hs.MaxReceiveBytes, MaxReceiveBytes: hs.MaxReceiveBytes,
@@ -50,11 +83,18 @@ func (u *appUplink) OnHandshakeComplete(ctx context.Context, hs port.HandshakeIn
} }
func (u *appUplink) OnDisconnect(ctx context.Context, conn port.ConnInfo, reason port.DisconnectReason) { func (u *appUplink) OnDisconnect(ctx context.Context, conn port.ConnInfo, reason port.DisconnectReason) {
lk := u.epLife(conn.EndpointID)
lk.Lock()
defer lk.Unlock()
if u.presence != nil { if u.presence != nil {
u.presence.ClearWatch(conn.ConnID) u.presence.ClearWatch(conn.ConnID)
} }
u.hsMu.Lock()
_, handshook := u.handshake[conn.ConnID]
delete(u.handshake, conn.ConnID)
u.hsMu.Unlock()
live, ok := u.conns.Current(conn.EndpointID) live, ok := u.conns.Current(conn.EndpointID)
isCurrent := ok && live.ConnID == conn.ConnID isCurrent := ok && live.ConnID == conn.ConnID && handshook
if err := u.msg.OnDisconnect(ctx, conn.EndpointID, conn.ConnID, isCurrent); err != nil { if err := u.msg.OnDisconnect(ctx, conn.EndpointID, conn.ConnID, isCurrent); err != nil {
u.log.Error("message disconnect", "endpoint", conn.EndpointID, "err", err) u.log.Error("message disconnect", "endpoint", conn.EndpointID, "err", err)
} }
@@ -280,7 +320,7 @@ func (u *appUplink) publishResp(ctx context.Context, conn port.ConnInfo, resp pr
return return
} }
if live, ok := u.conns.Current(conn.EndpointID); ok && live.ConnID == conn.ConnID { if live, ok := u.conns.Current(conn.EndpointID); ok && live.ConnID == conn.ConnID {
limit := respPayloadLimit(live.MaxPacketSize, live.MaxReceiveBytes) limit := broker.EffectivePayloadLimit(live.MaxPacketSize, live.MaxReceiveBytes)
if limit > 0 && len(b) > limit { if limit > 0 && len(b) > limit {
tooLarge := protocol.Resp{ tooLarge := protocol.Resp{
V: protocol.Version, V: protocol.Version,
@@ -300,20 +340,6 @@ func (u *appUplink) publishResp(ctx context.Context, conn port.ConnInfo, resp pr
} }
} }
func respPayloadLimit(maxPacketSize uint32, maxRecvBytes int) int {
limit := 0
if maxRecvBytes > 0 {
limit = maxRecvBytes
}
if maxPacketSize > 0 {
n := int(maxPacketSize)
if limit == 0 || n < limit {
limit = n
}
}
return limit
}
func peekRID(payload []byte) string { func peekRID(payload []byte) string {
var peek struct { var peek struct {
RID string `json:"rid"` RID string `json:"rid"`
+72
View File
@@ -1250,3 +1250,75 @@ issue #3 未关闭,`feat/fix-3-downlink-deadlock` 未合入 `main`。下面是
- 原因:默认 info 下第二个 CONNECT、3.1.1 发到错误主题等会把整包写入 JSON 日志。 - 原因:默认 info 下第二个 CONNECT、3.1.1 发到错误主题等会把整包写入 JSON 日志。
- 备选方案:改 mochi 日志调用点(需 fork)。 - 备选方案:改 mochi 日志调用点(需 fork)。
- 影响:排障时看不到载荷与密码,只见摘要。 - 影响:排障时看不到载荷与密码,只见摘要。
### 复审修复 B-03
- 日期:2026-09-30
- 原条款:DEVELOPMENT 第 5 节每端串行队列;Gitea #10。不改 `PublishDown` 签名。
- 实际做法:每连接独立下行队列(256 帧 / 16MiB)和发送 goroutine。`PublishDown` 只入队;发送与上行读循环解耦。队列满返回 `ErrBackpressure`。
- 原因:同连接同步 `InjectPacket` 与读循环写 PUBACK 会互相等待。
- 备选方案:改 `PublishDown` 签名或继续用 20ms sleep。
- 影响:调用方入队即返回;慢客户端只挡住该连接的发送 goroutine。
### 复审修复 B-06
- 日期:2026-09-30
- 原条款:Gitea #13。`PublishDown` 校验当前连接与下行订阅;导出有效载荷上限。
- 实际做法:非空 `connID` 必须仍是当前连接。未订阅 down 返回 `ErrNotSubscribed`。导出 `EffectivePayloadLimit`(Maximum Packet Size 减 128 字节包头预留)。`uplink.publishResp` 改用该函数。新连接建立时把旧连接标为 `superseded`。
- 原因:旧连接或未订阅时写入会静默失败或写错连接。
- 备选方案:发送时再检查(入队后连接可能已换)。
- 影响:无订阅时下行立即失败,不再占用大帧名额。
### 复审修复 B-04
- 日期:2026-09-30
- 原条款:Gitea #11。写出后再断开,不用固定 sleep。`serve.go` 只改 `identity.New` 的 ConnControl。
- 实际做法:`PublishThenDisconnect` 把帧与断开原因一并入队,发送 goroutine 写完再 `Disconnect`。logout / fatalKick 改走该原语。`identity.New` 的 `ConnControl` 置 nil,踢线仍走 Session 钩子。
- 原因:固定 20ms/50ms sleep 在慢客户端上会先断开,在快路径上又多余等待。
- 备选方案:继续 sleep;或改 identity 生命周期(本线不改)。
- 影响:identity 未接 ConnControl 时不再自己 20ms 踢线,生产路径统一由 Session 写出后断开。
### 复审修复 B-09
- 日期:2026-09-30
- 原条款:Gitea #16。生命周期串行化。不改 presence/app.go。
- 实际做法:broker 与 `appUplink` 按端编号加锁串行 `OnSessionEstablished` / `OnDisconnect` / 握手。uplink 另记 hello 握手表,仅已握手连接的断开才按当前连接通知消息线。
- 原因:顶号时旧连接 `OnDisconnect` 可能和新连接登记交错。
- 备选方案:改 presence 在线表(超出本线允许文件)。
- 影响:未 hello 的断开不再把消息连接表当成已握手在线来清推送标记。
### 复审修复 B-10
- 日期:2026-09-30
- 原条款:Gitea #17。登录写库条件更新;hello 重读令牌。不改 identity/self.go。
- 实际做法:密码登录 `UPDATE ... WHERE COALESCE(session_hash,'') = 读到的旧值`,影响行数为 0 则 `ErrSessionWriteConflict`。hello 用 `TokenMatchesDB` 核对明文,库已被换则响应里不带回旧令牌。
- 原因:两处同时密码登录会互相覆盖;hello 可能把已作废明文交给客户端。
- 备选方案:写库后无条件返回本次签发明文。
- 影响:写冲突时 OnConnect 返回 error(不回 0x86),客户端按网络故障重连。
### 复审修复 B-11
- 日期:2026-09-30
- 原条款:Gitea #18。令牌闲置按在线计。
- 实际做法:闲置判断取 `session_used_at` / `online_since` / `offline_since` 的较新者;当前在线(`online_since >= offline_since`)视为未闲置。
- 原因:只看 `session_used_at` 会让长期在线却很少写库的令牌过期。
- 备选方案:在线时每小时强制刷新 used_at(已有 touch,但仍可能窗口不够)。
- 影响:在线设备不会因为闲置天数被踢;离线后从最后一次在线/离线时刻起算。
### 复审修复 B-12
- 日期:2026-09-30
- 原条款:Gitea #19。认证超时与每端校验并发。不改 auth 池/PHC。
- 实际做法:`Authenticate` 套 30 秒超时;argon2 `Verify` 前每端信号量 2。`OnConnect` 同样带 30 秒 ctx。
- 原因:慢哈希或卡住的校验会堵住 mochi 读循环;同一编号并发登录会打满全局哈希池。
- 备选方案:改全局 Pool 大小(超出允许文件)。
- 影响:超时表现为内部错误断开(不回 0x86)。
### 复审修复 B-08
- 日期:2026-09-30
- 原条款:Gitea #15。Shutdown API。完整 HTTP 停机依赖 L-03。
- 实际做法:`Broker.Shutdown` 对现有连接发 MQTT 5 `0x8B`,清空上行队列并 `Close`。`serve` 在 listener Close 之前调用。HTTP `Shutdown` 留给 L-03。
- 原因:只关 listener 时 MQTT 客户端看不到规范的停机原因码。
- 备选方案:等 L-03 一并做(本线仍提供 broker API,避免监听线无法调用)。
- 影响:进程退出时端会收到 server shutting down;监听器 HTTP 优雅停机仍未做。
+116 -8
View File
@@ -32,6 +32,9 @@ type Login struct {
// 内存中的 session_used_at(毫秒)与上次落库时间。 // 内存中的 session_used_at(毫秒)与上次落库时间。
usedAt map[string]int64 usedAt map[string]int64
lastFlush map[string]int64 lastFlush map[string]int64
verifyMu sync.Mutex
verifySem map[string]chan struct{}
} }
// LoginOptions 装配 Login。 // LoginOptions 装配 Login。
@@ -67,9 +70,15 @@ func NewLogin(opts LoginOptions) *Login {
Now: now, Now: now,
usedAt: make(map[string]int64), usedAt: make(map[string]int64),
lastFlush: make(map[string]int64), lastFlush: make(map[string]int64),
verifySem: make(map[string]chan struct{}),
} }
} }
const (
authTimeout = 30 * time.Second
verifyPerEndpoint = 2
)
// Authenticate 按 DEVELOPMENT 第 5 节校验;内部故障返回 error。 // Authenticate 按 DEVELOPMENT 第 5 节校验;内部故障返回 error。
func (l *Login) Authenticate(ctx context.Context, endpointID string, password []byte, remoteIP string) (AuthResult, error) { func (l *Login) Authenticate(ctx context.Context, endpointID string, password []byte, remoteIP string) (AuthResult, error) {
if l == nil || l.DB == nil { if l == nil || l.DB == nil {
@@ -78,6 +87,11 @@ func (l *Login) Authenticate(ctx context.Context, endpointID string, password []
if endpointID == "" { if endpointID == "" {
return AuthResult{OK: false}, nil return AuthResult{OK: false}, nil
} }
if ctx == nil {
ctx = context.Background()
}
ctx, cancel := context.WithTimeout(ctx, authTimeout)
defer cancel()
row, err := l.loadEndpoint(ctx, endpointID) row, err := l.loadEndpoint(ctx, endpointID)
if err != nil { if err != nil {
@@ -107,6 +121,8 @@ type endpointAuthRow struct {
loginHash string loginHash string
sessionHash []byte // 原始 32 字节;无令牌时 nil sessionHash []byte // 原始 32 字节;无令牌时 nil
sessionUsedAt int64 // 毫秒;无则 0 sessionUsedAt int64 // 毫秒;无则 0
onlineSince int64
offlineSince int64
} }
func (l *Login) loadEndpoint(ctx context.Context, id string) (endpointAuthRow, error) { func (l *Login) loadEndpoint(ctx context.Context, id string) (endpointAuthRow, error) {
@@ -115,10 +131,12 @@ func (l *Login) loadEndpoint(ctx context.Context, id string) (endpointAuthRow, e
enabled int enabled int
sessHex sql.NullString sessHex sql.NullString
usedAt sql.NullInt64 usedAt sql.NullInt64
online sql.NullInt64
offline sql.NullInt64
) )
err := l.DB.Read.QueryRowContext(ctx, ` err := l.DB.Read.QueryRowContext(ctx, `
SELECT login_hash, enabled, session_hash, session_used_at SELECT login_hash, enabled, session_hash, session_used_at, online_since, offline_since
FROM endpoints WHERE id = ?`, id).Scan(&loginHash, &enabled, &sessHex, &usedAt) FROM endpoints WHERE id = ?`, id).Scan(&loginHash, &enabled, &sessHex, &usedAt, &online, &offline)
if err != nil { if err != nil {
if errors.Is(err, sql.ErrNoRows) { if errors.Is(err, sql.ErrNoRows) {
return endpointAuthRow{}, ErrEndpointNotFound return endpointAuthRow{}, ErrEndpointNotFound
@@ -132,6 +150,12 @@ FROM endpoints WHERE id = ?`, id).Scan(&loginHash, &enabled, &sessHex, &usedAt)
if usedAt.Valid { if usedAt.Valid {
row.sessionUsedAt = usedAt.Int64 row.sessionUsedAt = usedAt.Int64
} }
if online.Valid {
row.onlineSince = online.Int64
}
if offline.Valid {
row.offlineSince = offline.Int64
}
if sessHex.Valid && sessHex.String != "" { if sessHex.Valid && sessHex.String != "" {
raw, decErr := hex.DecodeString(sessHex.String) raw, decErr := hex.DecodeString(sessHex.String)
if decErr != nil || len(raw) != 32 { if decErr != nil || len(raw) != 32 {
@@ -160,12 +184,9 @@ func (l *Login) authSession(ctx context.Context, endpointID, token string, row e
usedAt = mem usedAt = mem
} }
l.usedMu.Unlock() l.usedMu.Unlock()
if l.IdleDays > 0 { if l.IdleDays > 0 && !sessionIdleOK(now, usedAt, row.onlineSince, row.offlineSince, l.IdleDays) {
idle := time.Duration(l.IdleDays) * 24 * time.Hour
if usedAt <= 0 || now.Sub(time.UnixMilli(usedAt)) > idle {
return false, nil return false, nil
} }
}
if err := l.touchSessionUsed(ctx, endpointID, nowMs); err != nil { if err := l.touchSessionUsed(ctx, endpointID, nowMs); err != nil {
return false, err return false, err
} }
@@ -204,7 +225,11 @@ func (l *Login) authPassword(ctx context.Context, endpointID, password, remoteIP
if l.Pool == nil { if l.Pool == nil {
return false, "", errors.New("broker: password pool not configured") return false, "", errors.New("broker: password pool not configured")
} }
if err := l.acquireVerify(ctx, endpointID); err != nil {
return false, "", err
}
match, verErr := l.Pool.Verify(ctx, auth.PasswordLogin, password, row.loginHash) match, verErr := l.Pool.Verify(ctx, auth.PasswordLogin, password, row.loginHash)
l.releaseVerify(endpointID)
if verErr != nil { if verErr != nil {
return false, "", verErr return false, "", verErr
} }
@@ -220,12 +245,28 @@ func (l *Login) authPassword(ctx context.Context, endpointID, password, remoteIP
} }
nowMs := l.Now().UnixMilli() nowMs := l.Now().UnixMilli()
hashHex := hex.EncodeToString(hash) hashHex := hex.EncodeToString(hash)
var oldHex any
if len(row.sessionHash) == 0 {
oldHex = ""
} else {
oldHex = hex.EncodeToString(row.sessionHash)
}
writeErr := l.DB.Queue.Do(ctx, func(tx *sql.Tx) error { writeErr := l.DB.Queue.Do(ctx, func(tx *sql.Tx) error {
_, e := tx.Exec(` res, e := tx.Exec(`
UPDATE endpoints UPDATE endpoints
SET session_hash = ?, session_issued_at = ?, session_used_at = ? SET session_hash = ?, session_issued_at = ?, session_used_at = ?
WHERE id = ?`, hashHex, nowMs, nowMs, endpointID) WHERE id = ? AND COALESCE(session_hash, '') = ?`, hashHex, nowMs, nowMs, endpointID, oldHex)
if e != nil {
return e return e
}
n, nErr := res.RowsAffected()
if nErr != nil {
return nErr
}
if n == 0 {
return ErrSessionWriteConflict
}
return nil
}) })
if writeErr != nil { if writeErr != nil {
return false, "", writeErr return false, "", writeErr
@@ -288,6 +329,73 @@ func (l *Login) SessionHashOf(ctx context.Context, endpointID string) ([]byte, e
return hex.DecodeString(sessHex.String) return hex.DecodeString(sessHex.String)
} }
// TokenMatchesDB 握手时重读:明文令牌是否仍对应库中当前哈希。
func (l *Login) TokenMatchesDB(ctx context.Context, endpointID, token string) (bool, error) {
if l == nil || l.DB == nil || token == "" {
return false, nil
}
got := l.Tokens.HashToken(token)
dbHash, err := l.SessionHashOf(ctx, endpointID)
if err != nil {
return false, err
}
if len(dbHash) == 0 {
return false, nil
}
return auth.EqualHash(got, dbHash), nil
}
func sessionIdleOK(now time.Time, usedAt, onlineSince, offlineSince int64, idleDays int) bool {
if idleDays <= 0 {
return true
}
online := onlineSince > 0 && onlineSince >= offlineSince
if online {
return true
}
activity := usedAt
if onlineSince > activity {
activity = onlineSince
}
if offlineSince > activity {
activity = offlineSince
}
if activity <= 0 {
return false
}
idle := time.Duration(idleDays) * 24 * time.Hour
return now.Sub(time.UnixMilli(activity)) <= idle
}
func (l *Login) acquireVerify(ctx context.Context, endpointID string) error {
l.verifyMu.Lock()
sem := l.verifySem[endpointID]
if sem == nil {
sem = make(chan struct{}, verifyPerEndpoint)
l.verifySem[endpointID] = sem
}
l.verifyMu.Unlock()
select {
case sem <- struct{}{}:
return nil
case <-ctx.Done():
return ctx.Err()
}
}
func (l *Login) releaseVerify(endpointID string) {
l.verifyMu.Lock()
sem := l.verifySem[endpointID]
l.verifyMu.Unlock()
if sem == nil {
return
}
select {
case <-sem:
default:
}
}
// LooksLikeSessionToken 暴露给测试。 // LooksLikeSessionToken 暴露给测试。
func (l *Login) LooksLikeSessionToken(s string) bool { func (l *Login) LooksLikeSessionToken(s string) bool {
return strings.HasPrefix(s, "nst_") return strings.HasPrefix(s, "nst_")
+123
View File
@@ -0,0 +1,123 @@
package broker
import (
"context"
"database/sql"
"sync"
"testing"
"time"
"git.asio.asia/nixevol/NixMsg/internal/auth"
"git.asio.asia/nixevol/NixMsg/internal/store"
)
func TestSessionIdleOKUsesOnlineOffline(t *testing.T) {
now := time.UnixMilli(1_700_000_000_000)
idleDays := 1
old := now.Add(-48 * time.Hour).UnixMilli()
recentOffline := now.Add(-2 * time.Hour).UnixMilli()
if sessionIdleOK(now, old, 0, 0, idleDays) {
t.Fatal("stale used_at should expire")
}
if !sessionIdleOK(now, old, now.UnixMilli(), 0, idleDays) {
t.Fatal("currently online should not expire")
}
if !sessionIdleOK(now, old, 0, recentOffline, idleDays) {
t.Fatal("recent offline_since should keep token")
}
}
func TestPasswordLoginConditionalUpdate(t *testing.T) {
dir := t.TempDir()
db, err := store.Open(dir, "FULL")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
pool := auth.NewStubHashPool()
login := NewLogin(LoginOptions{DB: db, Pool: pool, Tokens: auth.NewSessionTokens(), Locks: auth.NewLoginLocks(), IdleDays: 30})
phc, _ := pool.Hash(context.Background(), auth.PasswordLogin, "password1")
_ = db.Queue.Do(context.Background(), func(tx *sql.Tx) error {
_, e := tx.Exec(`INSERT INTO endpoints(id, name, login_hash, talk_hash, talk_version, default_delay_ms, enabled, created_at)
VALUES ('ep-cond', '', ?, NULL, 0, 0, 1, ?)`, phc, time.Now().UnixMilli())
return e
})
res, err := login.Authenticate(context.Background(), "ep-cond", []byte("password1"), "1.1.1.1")
if err != nil || !res.OK || res.SessionToken == "" {
t.Fatalf("first login %+v err=%v", res, err)
}
ok, err := login.TokenMatchesDB(context.Background(), "ep-cond", res.SessionToken)
if err != nil || !ok {
t.Fatalf("match=%v err=%v", ok, err)
}
res2, err := login.Authenticate(context.Background(), "ep-cond", []byte("password1"), "1.1.1.1")
if err != nil || !res2.OK {
t.Fatalf("second login %+v err=%v", res2, err)
}
ok, _ = login.TokenMatchesDB(context.Background(), "ep-cond", res.SessionToken)
if ok {
t.Fatal("old token should not match after second login")
}
ok, _ = login.TokenMatchesDB(context.Background(), "ep-cond", res2.SessionToken)
if !ok {
t.Fatal("new token should match")
}
}
func TestAuthenticateRespectsCanceledContext(t *testing.T) {
dir := t.TempDir()
db, err := store.Open(dir, "FULL")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
pool := &blockingPool{ready: make(chan struct{}), release: make(chan struct{})}
login := NewLogin(LoginOptions{DB: db, Pool: pool, Tokens: auth.NewSessionTokens(), Locks: auth.NewLoginLocks()})
phc, _ := auth.NewStubHashPool().Hash(context.Background(), auth.PasswordLogin, "password1")
_ = db.Queue.Do(context.Background(), func(tx *sql.Tx) error {
_, e := tx.Exec(`INSERT INTO endpoints(id, name, login_hash, talk_hash, talk_version, default_delay_ms, enabled, created_at)
VALUES ('ep-to', '', ?, NULL, 0, 0, 1, ?)`, phc, time.Now().UnixMilli())
return e
})
ctx, cancel := context.WithCancel(context.Background())
var wg sync.WaitGroup
wg.Add(1)
var gotErr error
go func() {
defer wg.Done()
_, gotErr = login.Authenticate(ctx, "ep-to", []byte("password1"), "9.9.9.9")
}()
select {
case <-pool.ready:
case <-time.After(2 * time.Second):
t.Fatal("verify did not start")
}
cancel()
wg.Wait()
close(pool.release)
if gotErr == nil {
t.Fatal("expected canceled auth")
}
}
type blockingPool struct {
ready chan struct{}
release chan struct{}
once sync.Once
}
func (p *blockingPool) Hash(ctx context.Context, kind auth.PasswordKind, password string) (string, error) {
return auth.NewStubHashPool().Hash(ctx, kind, password)
}
func (p *blockingPool) Verify(ctx context.Context, _ auth.PasswordKind, _, _ string) (bool, error) {
p.once.Do(func() { close(p.ready) })
select {
case <-ctx.Done():
return false, ctx.Err()
case <-p.release:
return true, nil
}
}
func (p *blockingPool) QueueLen() int { return 0 }
+16 -6
View File
@@ -3,6 +3,7 @@ package broker
import ( import (
"bytes" "bytes"
"context" "context"
"errors"
"io" "io"
"net" "net"
"sync" "sync"
@@ -81,7 +82,16 @@ func TestLargeFrameQuotaReleasedOnDisconnect(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
cancel() cancel()
if len(b.largeSem) == 0 { held := false
deadline := time.Now().Add(2 * time.Second)
for time.Now().Before(deadline) {
if len(b.largeSem) > 0 {
held = true
break
}
time.Sleep(10 * time.Millisecond)
}
if !held {
t.Fatal("expected a held large slot before disconnect") t.Fatal("expected a held large slot before disconnect")
} }
_ = w.Close() _ = w.Close()
@@ -90,7 +100,7 @@ func TestLargeFrameQuotaReleasedOnDisconnect(t *testing.T) {
case <-time.After(3 * time.Second): case <-time.After(3 * time.Second):
t.Fatal("client attach did not return") t.Fatal("client attach did not return")
} }
deadline := time.Now().Add(2 * time.Second) deadline = time.Now().Add(2 * time.Second)
for time.Now().Before(deadline) { for time.Now().Before(deadline) {
if len(b.largeSem) == 0 { if len(b.largeSem) == 0 {
return return
@@ -119,11 +129,11 @@ func TestLargeFrameQuotaReleasedWithoutSubscriber(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), time.Second) ctx, cancel := context.WithTimeout(context.Background(), time.Second)
err := b.PublishDown(ctx, "ep-large-nosub", "", payload, port.PublishOpts{QoS: 1}) err := b.PublishDown(ctx, "ep-large-nosub", "", payload, port.PublishOpts{QoS: 1})
cancel() cancel()
if err != nil { if !errors.Is(err, ErrNotSubscribed) {
t.Fatal(err) t.Fatalf("err=%v want ErrNotSubscribed", err)
} }
if n := len(b.largeSem); n != 0 { if len(b.largeSem) != 0 {
t.Fatalf("held slots without subscriber: %d", n) t.Fatalf("held slots without subscriber: %d", len(b.largeSem))
} }
} }
+212
View File
@@ -0,0 +1,212 @@
package broker
import (
"bytes"
"context"
"errors"
"io"
"net"
"strconv"
"testing"
"time"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
"github.com/mochi-mqtt/server/v2/packets"
)
func TestPublishDownBackpressureWhenQueueFull(t *testing.T) {
// 无缓冲 pipe:发送 goroutine 在客户端不读时堵住,队列才能填满。
b, err := New(Options{Authenticator: AllowAuthenticator{}})
if err != nil {
t.Fatal(err)
}
defer func() { _ = b.Close() }()
r, w := net.Pipe()
done := make(chan struct{})
go func() {
defer close(done)
_ = b.AttachTCP(r)
}()
defer func() {
_ = w.Close()
select {
case <-done:
case <-time.After(3 * time.Second):
}
}()
connectAndSubscribe(t, w, "ep-bp", 0)
waitSession(t, b, "ep-bp")
payload := bytes.Repeat([]byte("q"), 1024)
var sawBP bool
start := time.Now()
// mochi outbound 缓冲 1024,发送 goroutine 要先填满它才会堵住,随后才轮到本地下行队列。
for i := 0; i < 1024+downQueueMax+16; i++ {
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
err := b.PublishDown(ctx, "ep-bp", "", payload, port.PublishOpts{QoS: 1})
cancel()
if errors.Is(err, ErrBackpressure) {
if time.Since(start) > 50*time.Millisecond && i == 0 {
t.Fatalf("first backpressure took %s", time.Since(start))
}
sawBP = true
break
}
if err != nil {
t.Fatalf("publish %d: %v", i, err)
}
}
if !sawBP {
t.Fatal("expected ErrBackpressure")
}
}
func TestSlowClientDoesNotBlockOtherPublishDown(t *testing.T) {
b, err := New(Options{Authenticator: AllowAuthenticator{}})
if err != nil {
t.Fatal(err)
}
defer func() { _ = b.Close() }()
slowW, slowDone := acceptAndDial(t, b)
fastW, fastDone := acceptAndDial(t, b)
defer func() {
_ = slowW.Close()
_ = fastW.Close()
select {
case <-slowDone:
case <-time.After(3 * time.Second):
}
select {
case <-fastDone:
case <-time.After(3 * time.Second):
}
}()
writeConnect(t, slowW, "ep-slow", 30, 0)
readExactPacket(t, slowW, packets.Connack, 3*time.Second)
writeSubscribe(t, slowW, downTopic("ep-slow"))
readExactPacket(t, slowW, packets.Suback, 3*time.Second)
writeConnect(t, fastW, "ep-fast", 30, 0)
readExactPacket(t, fastW, packets.Connack, 3*time.Second)
writeSubscribe(t, fastW, downTopic("ep-fast"))
readExactPacket(t, fastW, packets.Suback, 3*time.Second)
waitSession(t, b, "ep-slow")
waitSession(t, b, "ep-fast")
big := bytes.Repeat([]byte("s"), 64*1024)
for i := 0; i < 8; i++ {
_ = b.PublishDown(context.Background(), "ep-slow", "", big, port.PublishOpts{QoS: 1})
}
small := []byte(`{"v":1,"type":"resp"}`)
start := time.Now()
if err := b.PublishDown(context.Background(), "ep-fast", "", small, port.PublishOpts{QoS: 1}); err != nil {
t.Fatalf("fast publish: %v", err)
}
if time.Since(start) > 100*time.Millisecond {
t.Fatalf("fast PublishDown took %s", time.Since(start))
}
got := readDownPayload(t, fastW, 2*time.Second)
if !bytes.Equal(got, small) {
t.Fatalf("fast got %q", got)
}
}
func TestDownlinkFIFOOrder(t *testing.T) {
b, w, done := startTCPClient(t, "ep-ord")
defer func() { _ = b.Close() }()
defer func() {
_ = w.Close()
select {
case <-done:
case <-time.After(3 * time.Second):
}
}()
writeConnect(t, w, "ep-ord", 30, 0)
readExactPacket(t, w, packets.Connack, 3*time.Second)
writeSubscribe(t, w, downTopic("ep-ord"))
readExactPacket(t, w, packets.Suback, 3*time.Second)
waitSession(t, b, "ep-ord")
const n = 64
for i := 0; i < n; i++ {
p := []byte(strconv.Itoa(i))
if err := b.PublishDown(context.Background(), "ep-ord", "", p, port.PublishOpts{QoS: 1}); err != nil {
t.Fatal(err)
}
}
for i := 0; i < n; i++ {
got := readDownPayload(t, w, 3*time.Second)
want := []byte(strconv.Itoa(i))
if !bytes.Equal(got, want) {
t.Fatalf("order %d: got %s want %s", i, got, want)
}
}
}
func acceptAndDial(t *testing.T, b *Broker) (net.Conn, chan struct{}) {
t.Helper()
ln, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = ln.Close() })
done := make(chan struct{})
go func() {
defer close(done)
c, accErr := ln.Accept()
if accErr != nil {
return
}
_ = b.AttachTCP(c)
}()
w, err := net.Dial("tcp", ln.Addr().String())
if err != nil {
t.Fatal(err)
}
return w, done
}
func readDownPayload(t *testing.T, conn net.Conn, timeout time.Duration) []byte {
t.Helper()
deadline := time.Now().Add(timeout)
for time.Now().Before(deadline) {
_ = conn.SetReadDeadline(time.Now().Add(200 * time.Millisecond))
hdr := make([]byte, 1)
if _, err := io.ReadFull(conn, hdr); err != nil {
continue
}
rem, err := readRemainingLengthConn(conn)
if err != nil {
continue
}
body := make([]byte, rem)
if _, err := io.ReadFull(conn, body); err != nil {
continue
}
if hdr[0]>>4 != packets.Publish {
continue
}
qos := (hdr[0] >> 1) & 0x3
pk := new(packets.Packet)
pk.ProtocolVersion = 5
pk.FixedHeader = packets.FixedHeader{Type: packets.Publish, Remaining: rem, Qos: qos}
if err := pk.PublishDecode(body); err != nil {
t.Fatal(err)
}
if qos > 0 {
ack := packets.Packet{
FixedHeader: packets.FixedHeader{Type: packets.Puback},
ProtocolVersion: 5,
PacketID: pk.PacketID,
}
var ab bytes.Buffer
_ = ack.PubackEncode(&ab)
_, _ = conn.Write(ab.Bytes())
}
return pk.Payload
}
t.Fatal("timeout waiting publish")
return nil
}
+128
View File
@@ -0,0 +1,128 @@
package broker
import (
"bytes"
"context"
"errors"
"io"
"testing"
"time"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
"github.com/mochi-mqtt/server/v2/packets"
)
func TestPublishDownRequiresDownSubscription(t *testing.T) {
b, w, done := startTCPClient(t, "ep-nosub")
defer func() { _ = b.Close() }()
defer func() {
_ = w.Close()
select {
case <-done:
case <-time.After(3 * time.Second):
}
}()
writeConnect(t, w, "ep-nosub", 30, 0)
readExactPacket(t, w, packets.Connack, 3*time.Second)
waitSession(t, b, "ep-nosub")
err := b.PublishDown(context.Background(), "ep-nosub", "", []byte(`{"v":1}`), port.PublishOpts{QoS: 0})
if !errors.Is(err, ErrNotSubscribed) {
t.Fatalf("err=%v want ErrNotSubscribed", err)
}
}
func TestPublishDownRejectsStaleConnID(t *testing.T) {
b, w, done := startTCPClient(t, "ep-stale")
defer func() { _ = b.Close() }()
defer func() {
_ = w.Close()
select {
case <-done:
case <-time.After(3 * time.Second):
}
}()
writeConnect(t, w, "ep-stale", 30, 0)
readExactPacket(t, w, packets.Connack, 3*time.Second)
writeSubscribe(t, w, downTopic("ep-stale"))
readExactPacket(t, w, packets.Suback, 3*time.Second)
waitSession(t, b, "ep-stale")
err := b.PublishDown(context.Background(), "ep-stale", "dead-conn", []byte(`{"v":1}`), port.PublishOpts{QoS: 0})
if !errors.Is(err, ErrNoConnection) {
t.Fatalf("err=%v want ErrNoConnection", err)
}
}
func TestPublishThenDisconnectWritesThenCloses(t *testing.T) {
b, w, done := startTCPClient(t, "ep-ptd")
defer func() { _ = b.Close() }()
defer func() {
_ = w.Close()
select {
case <-done:
case <-time.After(3 * time.Second):
}
}()
writeConnect(t, w, "ep-ptd", 30, 0)
readExactPacket(t, w, packets.Connack, 3*time.Second)
writeSubscribe(t, w, downTopic("ep-ptd"))
readExactPacket(t, w, packets.Suback, 3*time.Second)
waitSession(t, b, "ep-ptd")
payload := []byte(`{"v":1,"type":"fatal","reason":"disabled"}`)
if err := b.PublishThenDisconnect(context.Background(), "ep-ptd", "", payload, 1, port.DisconnectFatal); err != nil {
t.Fatal(err)
}
got := readDownPayload(t, w, 3*time.Second)
if !bytes.Equal(got, payload) {
t.Fatalf("got %s", got)
}
_ = 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) {
b, w, done := startTCPClient(t, "ep-shut")
defer func() {
_ = w.Close()
select {
case <-done:
case <-time.After(3 * time.Second):
}
}()
writeConnect(t, w, "ep-shut", 30, 0)
readExactPacket(t, w, packets.Connack, 3*time.Second)
waitSession(t, b, "ep-shut")
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
if err := b.Shutdown(ctx); err != nil {
t.Fatal(err)
}
_ = w.SetReadDeadline(time.Now().Add(2 * time.Second))
buf := make([]byte, 32)
n, err := io.ReadAtLeast(w, buf, 2)
if err != nil && n == 0 {
return
}
if n > 0 && buf[0]>>4 != packets.Disconnect {
t.Fatalf("want disconnect got %x", buf[:n])
}
}
func TestEffectivePayloadLimitSubtractsOverhead(t *testing.T) {
got := EffectivePayloadLimit(200, 0)
if got != 200-packetOverheadBudget {
t.Fatalf("got %d", got)
}
got = EffectivePayloadLimit(200, 50)
if got != 50 {
t.Fatalf("got %d want 50", got)
}
}
+120 -29
View File
@@ -37,7 +37,20 @@ var ErrNoConnection = errors.New("broker: no active connection")
// ErrLargeFrameTimeout 全局大帧名额在有界等待内拿不到。 // ErrLargeFrameTimeout 全局大帧名额在有界等待内拿不到。
var ErrLargeFrameTimeout = errors.New("broker: large frame quota timeout") var ErrLargeFrameTimeout = errors.New("broker: large frame quota timeout")
const largeAcquireWait = 5 * time.Second // ErrBackpressure 该连接下行队列已满(帧数或字节数)。
var ErrBackpressure = errors.New("broker: downlink backpressure")
// ErrNotSubscribed 当前连接尚未订阅下行主题。
var ErrNotSubscribed = errors.New("broker: down topic not subscribed")
// ErrSessionWriteConflict 密码登录写令牌时发现库已被并发更新。
var ErrSessionWriteConflict = errors.New("broker: session token write conflict")
const (
largeAcquireWait = 5 * time.Second
downQueueMax = 256
downQueueBytes = 16 << 20
)
// AuthResult 是登录校验结论(N3 实现真实逻辑;N2 默认拒绝)。 // AuthResult 是登录校验结论(N3 实现真实逻辑;N2 默认拒绝)。
type AuthResult struct { type AuthResult struct {
@@ -100,6 +113,9 @@ type Broker struct {
largeSem chan struct{} largeSem chan struct{}
closed atomic.Bool closed atomic.Bool
lifeMu sync.Mutex
lifeLocks map[string]*sync.Mutex
} }
type connState struct { type connState struct {
@@ -120,6 +136,13 @@ type connState struct {
metricsCounted bool metricsCounted bool
established bool established bool
createdAt time.Time createdAt time.Time
closing bool
superseded bool
downCh chan downItem
downStop chan struct{}
downDone chan struct{}
downBytes atomic.Int64
sentPub atomic.Int64
mu sync.Mutex mu sync.Mutex
handshakeTimer *time.Timer handshakeTimer *time.Timer
@@ -174,6 +197,7 @@ func New(opts Options) (*Broker, error) {
closedCh: make(chan struct{}), closedCh: make(chan struct{}),
queues: make(map[string]*uplinkQueue), queues: make(map[string]*uplinkQueue),
largeSem: make(chan struct{}, largeFrameSlots), largeSem: make(chan struct{}, largeFrameSlots),
lifeLocks: make(map[string]*sync.Mutex),
} }
b.hook = &nixHook{b: b} b.hook = &nixHook{b: b}
if err := srv.AddHook(b.hook, nil); err != nil { if err := srv.AddHook(b.hook, nil); err != nil {
@@ -203,10 +227,36 @@ func (b *Broker) Close() error {
for _, q := range b.queues { for _, q := range b.queues {
q.close() q.close()
} }
b.queues = make(map[string]*uplinkQueue)
b.queuesMu.Unlock() b.queuesMu.Unlock()
return b.server.Close() return b.server.Close()
} }
// Shutdown 向所有连接发 MQTT 5 0x8B 后关闭。完整 HTTP 停机顺序见 L-03。
func (b *Broker) Shutdown(ctx context.Context) error {
if b.closed.Load() {
return nil
}
b.connsMu.RLock()
clients := make([]*mqtt.Client, 0, len(b.byClient))
for cl := range b.byClient {
if cl != nil {
clients = append(clients, cl)
}
}
b.connsMu.RUnlock()
for _, cl := range clients {
_ = b.server.DisconnectClient(cl, packets.ErrServerShuttingDown)
}
if ctx != nil {
select {
case <-ctx.Done():
default:
}
}
return b.Close()
}
// AttachTCP 把裸 TCP/TLS 连接交给 mochi;阻塞到连接结束。 // AttachTCP 把裸 TCP/TLS 连接交给 mochi;阻塞到连接结束。
func (b *Broker) AttachTCP(conn net.Conn) error { func (b *Broker) AttachTCP(conn net.Conn) error {
return b.server.EstablishConnection("tcp", conn) return b.server.EstablishConnection("tcp", conn)
@@ -219,44 +269,55 @@ func (b *Broker) AttachWS(conn net.Conn) error {
// PublishDown 实现 port.Downlink。 // PublishDown 实现 port.Downlink。
func (b *Broker) PublishDown(ctx context.Context, endpointID string, connID port.ConnID, payload []byte, opts port.PublishOpts) error { func (b *Broker) PublishDown(ctx context.Context, endpointID string, connID port.ConnID, payload []byte, opts port.PublishOpts) error {
if b.closed.Load() {
return errors.New("broker: closed")
}
st := b.lookupConn(endpointID, connID)
if st == nil {
return ErrNoConnection
}
limit := effectivePayloadLimit(st.maxPacketSize, st.maxRecvBytes)
if limit > 0 && len(payload) > limit {
return ErrPayloadTooLarge
}
qos := opts.QoS qos := opts.QoS
if qos > 1 { if qos > 1 {
qos = 1 qos = 1
} }
topic := downTopic(endpointID) return b.enqueueDownlink(endpointID, connID, payload, qos, "")
large := len(payload) > largeFrameBytes }
if large {
if err := b.acquireLarge(ctx); err != nil { // PublishThenDisconnect 把一帧写入该连接下行队列,写出后再断开(无固定 sleep)。
return err func (b *Broker) PublishThenDisconnect(_ context.Context, endpointID string, connID port.ConnID, payload []byte, qos byte, reason port.DisconnectReason) error {
if qos > 1 {
qos = 1
} }
st.mu.Lock() if reason == "" {
st.largePending++ reason = port.DisconnectNormal
st.mu.Unlock() }
return b.enqueueDownlink(endpointID, connID, payload, qos, reason)
}
func (b *Broker) enqueueDownlink(endpointID string, connID port.ConnID, payload []byte, qos byte, disconnect port.DisconnectReason) error {
if b.closed.Load() {
return errors.New("broker: closed")
}
st := b.lookupCurrent(endpointID, connID)
if st == nil {
return ErrNoConnection
} }
if err := b.server.Publish(topic, payload, false, qos); err != nil { st.mu.Lock()
if large { maxRecv := st.maxRecvBytes
b.finishLargePublish(st) closing := st.closing
superseded := st.superseded
st.mu.Unlock()
if closing || superseded {
return ErrNoConnection
} }
return err if !b.hasDownSub(st) {
return ErrNotSubscribed
} }
if large {
b.finishLargePublish(st) limit := EffectivePayloadLimit(st.maxPacketSize, maxRecv)
if limit > 0 && len(payload) > limit {
return ErrPayloadTooLarge
} }
return nil
return st.enqueueDown(downItem{
payload: append([]byte(nil), payload...),
qos: qos,
disconnect: disconnect,
})
} }
func (b *Broker) acquireLarge(ctx context.Context) error { func (b *Broker) acquireLarge(ctx context.Context) error {
@@ -367,6 +428,31 @@ func (b *Broker) lookupConn(endpointID string, connID port.ConnID) *connState {
return b.current[endpointID] return b.current[endpointID]
} }
// lookupCurrent 只返回该端当前连接;connID 非空时必须仍是当前连接。
func (b *Broker) lookupCurrent(endpointID string, connID port.ConnID) *connState {
b.connsMu.RLock()
defer b.connsMu.RUnlock()
cur := b.current[endpointID]
if cur == nil {
return nil
}
if connID != "" && cur.connID != connID {
return nil
}
return cur
}
func (b *Broker) endpointLife(endpointID string) *sync.Mutex {
b.lifeMu.Lock()
defer b.lifeMu.Unlock()
m := b.lifeLocks[endpointID]
if m == nil {
m = &sync.Mutex{}
b.lifeLocks[endpointID] = m
}
return m
}
func downTopic(endpointID string) string { func downTopic(endpointID string) string {
return "nix/c/" + endpointID + "/down" return "nix/c/" + endpointID + "/down"
} }
@@ -375,6 +461,11 @@ func upTopic(endpointID string) string {
return "nix/c/" + endpointID + "/up" return "nix/c/" + endpointID + "/up"
} }
// EffectivePayloadLimit 下行载荷上限:客户端 Maximum Packet Size 减包头预留,再与 max_receive_bytes 取更严者。
func EffectivePayloadLimit(maxPacketSize uint32, maxRecvBytes int) int {
return effectivePayloadLimit(maxPacketSize, maxRecvBytes)
}
func effectivePayloadLimit(maxPacketSize uint32, maxRecvBytes int) int { func effectivePayloadLimit(maxPacketSize uint32, maxRecvBytes int) int {
limit := 0 limit := 0
if maxPacketSize > 0 { if maxPacketSize > 0 {
+192
View File
@@ -0,0 +1,192 @@
package broker
import (
"context"
"time"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
)
type downItem struct {
payload []byte
qos byte
disconnect port.DisconnectReason // 非空表示该帧写出后断开(B-04)
sent chan struct{}
}
func (st *connState) startDownLoop(b *Broker) {
st.mu.Lock()
if st.downCh != nil {
st.mu.Unlock()
return
}
st.downCh = make(chan downItem, downQueueMax)
st.downStop = make(chan struct{})
st.downDone = make(chan struct{})
st.mu.Unlock()
go st.downLoop(b)
}
func (st *connState) stopDownLoop() {
st.mu.Lock()
stop := st.downStop
done := st.downDone
ch := st.downCh
st.mu.Unlock()
if stop == nil {
return
}
select {
case <-stop:
default:
close(stop)
}
if done != nil {
select {
case <-done:
case <-time.After(2 * time.Second):
}
}
if ch != nil {
for {
select {
case item := <-ch:
st.downBytes.Add(-int64(len(item.payload)))
default:
return
}
}
}
}
func (st *connState) enqueueDown(item downItem) error {
st.mu.Lock()
ch := st.downCh
stop := st.downStop
closing := st.closing
st.mu.Unlock()
if ch == nil || stop == nil {
return ErrNoConnection
}
select {
case <-stop:
return ErrNoConnection
default:
}
if closing && item.disconnect == "" {
return ErrNoConnection
}
n := int64(len(item.payload))
for {
cur := st.downBytes.Load()
if cur+n > downQueueBytes {
return ErrBackpressure
}
if st.downBytes.CompareAndSwap(cur, cur+n) {
break
}
}
select {
case ch <- item:
return nil
default:
st.downBytes.Add(-n)
return ErrBackpressure
}
}
func (st *connState) downLoop(b *Broker) {
defer close(st.downDone)
for {
select {
case <-st.downStop:
return
case item, ok := <-st.downCh:
if !ok {
return
}
st.downBytes.Add(-int64(len(item.payload)))
st.sendOne(b, item)
}
}
}
func (st *connState) sendOne(b *Broker, item downItem) {
if b.closed.Load() {
st.signalSent(item)
return
}
large := len(item.payload) > largeFrameBytes
if large {
if err := b.acquireLarge(context.Background()); err != nil {
if b.onDrop != nil {
b.onDrop(context.Background(), st.endpointID, st.connID, item.payload)
}
st.signalSent(item)
return
}
st.mu.Lock()
st.largePending++
st.mu.Unlock()
}
topic := downTopic(st.endpointID)
before := st.sentPub.Load()
var err error
for {
if b.closed.Load() {
break
}
select {
case <-st.downStop:
err = ErrNoConnection
default:
err = b.server.Publish(topic, item.payload, false, item.qos)
if err == nil {
break
}
select {
case <-st.downStop:
err = ErrNoConnection
case <-time.After(2 * time.Millisecond):
continue
}
}
break
}
if large {
b.finishLargePublish(st)
}
if err != nil && b.onDrop != nil {
b.onDrop(context.Background(), st.endpointID, st.connID, item.payload)
}
st.signalSent(item)
if err == nil && item.disconnect != "" {
st.waitPacketWritten(before)
_ = b.Disconnect(context.Background(), st.endpointID, st.connID, item.disconnect)
}
}
func (st *connState) waitPacketWritten(before int64) {
deadline := time.Now().Add(2 * time.Second)
for time.Now().Before(deadline) {
if st.sentPub.Load() > before {
return
}
select {
case <-st.downStop:
return
case <-time.After(2 * time.Millisecond):
}
}
}
func (st *connState) signalSent(item downItem) {
if item.sent == nil {
return
}
select {
case <-item.sent:
default:
close(item.sent)
}
}
+31 -1
View File
@@ -30,6 +30,7 @@ func (h *nixHook) Provides(b byte) bool {
mqtt.OnQosComplete, mqtt.OnQosComplete,
mqtt.OnQosDropped, mqtt.OnQosDropped,
mqtt.OnSubscribed, mqtt.OnSubscribed,
mqtt.OnPacketSent,
}, []byte{b}) }, []byte{b})
} }
@@ -79,7 +80,9 @@ func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error {
"endpoint", endpointID, "receive_maximum", rm) "endpoint", endpointID, "receive_maximum", rm)
} }
res, err := h.b.auth.Authenticate(context.Background(), endpointID, pk.Connect.Password, remoteIP) authCtx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
res, err := h.b.auth.Authenticate(authCtx, endpointID, pk.Connect.Password, remoteIP)
if err != nil { if err != nil {
return err // mochi 不回 CONNACK,直接断开;不登记连接表 return err // mochi 不回 CONNACK,直接断开;不登记连接表
} }
@@ -171,6 +174,18 @@ func (h *nixHook) OnSubscribed(cl *mqtt.Client, pk packets.Packet, reasonCodes [
} }
} }
func (h *nixHook) OnPacketSent(cl *mqtt.Client, pk packets.Packet, _ []byte) {
if pk.FixedHeader.Type != packets.Publish {
return
}
h.b.connsMu.RLock()
st := h.b.byClient[cl]
h.b.connsMu.RUnlock()
if st != nil {
st.sentPub.Add(1)
}
}
func (h *nixHook) OnPublishDropped(cl *mqtt.Client, pk packets.Packet) { func (h *nixHook) OnPublishDropped(cl *mqtt.Client, pk packets.Packet) {
h.b.log.Debug("publish dropped", "client", cl.ID, "topic", pk.TopicName, "size", len(pk.Payload)) h.b.log.Debug("publish dropped", "client", cl.ID, "topic", pk.TopicName, "size", len(pk.Payload))
h.b.connsMu.RLock() h.b.connsMu.RLock()
@@ -188,7 +203,9 @@ func (h *nixHook) OnPublishDropped(cl *mqtt.Client, pk packets.Packet) {
func (h *nixHook) OnSessionEstablished(cl *mqtt.Client, _ packets.Packet) { func (h *nixHook) OnSessionEstablished(cl *mqtt.Client, _ packets.Packet) {
h.b.connsMu.Lock() h.b.connsMu.Lock()
st := h.b.byClient[cl] st := h.b.byClient[cl]
var old *connState
if st != nil { if st != nil {
old = h.b.current[st.endpointID]
h.b.current[st.endpointID] = st h.b.current[st.endpointID] = st
st.established = true st.established = true
} }
@@ -196,6 +213,15 @@ func (h *nixHook) OnSessionEstablished(cl *mqtt.Client, _ packets.Packet) {
if st == nil { if st == nil {
return return
} }
lk := h.b.endpointLife(st.endpointID)
lk.Lock()
if old != nil && old != st {
old.mu.Lock()
old.superseded = true
old.mu.Unlock()
}
lk.Unlock()
st.startDownLoop(h.b)
info := port.ConnInfo{ info := port.ConnInfo{
ConnID: st.connID, ConnID: st.connID,
EndpointID: st.endpointID, EndpointID: st.endpointID,
@@ -224,6 +250,10 @@ func (h *nixHook) OnDisconnect(cl *mqtt.Client, err error, _ bool) {
if st == nil { if st == nil {
return return
} }
lk := h.b.endpointLife(st.endpointID)
lk.Lock()
st.stopDownLoop()
lk.Unlock()
h.b.releaseAllLarge(st) h.b.releaseAllLarge(st)
h.b.cancelHandshakeDeadline(st.endpointID, st.connID) h.b.cancelHandshakeDeadline(st.endpointID, st.connID)
+20 -6
View File
@@ -210,6 +210,13 @@ func (s *Session) handleHello(ctx context.Context, conn port.ConnInfo, st *connS
} }
s.b.SetMaxReceiveBytes(conn.EndpointID, conn.ConnID, maxRecv) s.b.SetMaxReceiveBytes(conn.EndpointID, conn.ConnID, maxRecv)
if conn.SessionToken != "" && s.login != nil {
if keep, chkErr := s.login.TokenMatchesDB(ctx, conn.EndpointID, conn.SessionToken); chkErr != nil {
s.log.Error("re-read session token", "endpoint", conn.EndpointID, "err", chkErr)
} else if !keep {
conn.SessionToken = ""
}
}
data := protocol.HelloData{ data := protocol.HelloData{
ServerTimeMs: s.now().UnixMilli(), ServerTimeMs: s.now().UnixMilli(),
ServerVersion: s.limits.ServerVersion, ServerVersion: s.limits.ServerVersion,
@@ -279,14 +286,17 @@ func (s *Session) handleLogout(ctx context.Context, conn port.ConnInfo, req *pro
} }
} }
resp := protocol.Resp{V: protocol.Version, Type: protocol.TypeResp, RID: req.RID, OK: true} resp := protocol.Resp{V: protocol.Version, Type: protocol.TypeResp, RID: req.RID, OK: true}
if err := s.publishJSON(ctx, conn, resp, 1); err != nil { raw, err := protocol.Marshal(resp)
s.log.Error("logout resp", "endpoint", conn.EndpointID, "err", err) if err != nil {
s.replyErr(ctx, conn, req.RID, protocol.CodeBusy, "marshal logout resp")
return nil
} }
if pubErr := s.b.PublishThenDisconnect(ctx, conn.EndpointID, conn.ConnID, raw, 1, port.DisconnectNormal); pubErr != nil {
s.log.Error("logout resp", "endpoint", conn.EndpointID, "err", pubErr)
go func() { go func() {
// 稍等让 QoS1 resp 写入连接,再断开
time.Sleep(50 * time.Millisecond)
_ = s.b.Disconnect(context.Background(), conn.EndpointID, conn.ConnID, port.DisconnectNormal) _ = s.b.Disconnect(context.Background(), conn.EndpointID, conn.ConnID, port.DisconnectNormal)
}() }()
}
return nil return nil
} }
@@ -327,11 +337,15 @@ func (s *Session) fatalKick(ctx context.Context, endpointID, reason string) erro
return nil return nil
} }
fatal := protocol.Fatal{V: protocol.Version, Type: protocol.TypeFatal, Reason: reason} fatal := protocol.Fatal{V: protocol.Version, Type: protocol.TypeFatal, Reason: reason}
_ = s.publishJSON(ctx, info, fatal, 1) raw, err := protocol.Marshal(fatal)
if err != nil {
return err
}
if pubErr := s.b.PublishThenDisconnect(ctx, info.EndpointID, info.ConnID, raw, 1, port.DisconnectFatal); pubErr != nil {
go func() { go func() {
time.Sleep(20 * time.Millisecond)
_ = s.b.Disconnect(context.Background(), endpointID, info.ConnID, port.DisconnectFatal) _ = s.b.Disconnect(context.Background(), endpointID, info.ConnID, port.DisconnectFatal)
}() }()
}
return nil return nil
} }