merge: broker b01-b12
This commit is contained in:
+5
-1
@@ -153,7 +153,7 @@ func runServe(ctx context.Context, cfg config.Config) error {
|
||||
Sessions: sessionTokens,
|
||||
MaxScheduleSeconds: int64(cfg.Limits.MaxScheduleSeconds),
|
||||
Logger: slog.Default(),
|
||||
ConnControl: brk,
|
||||
ConnControl: nil, // B-04:踢线走 Session 钩子,避免 identity 20ms 异步 Disconnect
|
||||
Downlink: brk,
|
||||
ClientIP: func(r *http.Request) string {
|
||||
return httpx.ClientIP(r, trustedNets)
|
||||
@@ -310,6 +310,10 @@ func runServe(ctx context.Context, cfg config.Config) error {
|
||||
|
||||
<-ctx.Done()
|
||||
loopCancel()
|
||||
// B-08:先对 MQTT 连接发 0x8B。HTTP Shutdown 与监听器完整停机顺序见 L-03。
|
||||
shutCtx, shutCancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
_ = brk.Shutdown(shutCtx)
|
||||
shutCancel()
|
||||
_ = lnSrv.Close()
|
||||
drainCtx, drainCancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer drainCancel()
|
||||
|
||||
+42
-16
@@ -5,12 +5,14 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"sync"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/group"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/identity"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/message"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"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/protocol"
|
||||
)
|
||||
@@ -25,9 +27,31 @@ type appUplink struct {
|
||||
down port.Downlink
|
||||
log *slog.Logger
|
||||
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 {
|
||||
lk := u.epLife(conn.EndpointID)
|
||||
lk.Lock()
|
||||
defer lk.Unlock()
|
||||
u.conns.Set(conn.EndpointID, message.LiveConn{
|
||||
ConnID: conn.ConnID,
|
||||
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 {
|
||||
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{
|
||||
ConnID: hs.ConnID,
|
||||
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) {
|
||||
lk := u.epLife(conn.EndpointID)
|
||||
lk.Lock()
|
||||
defer lk.Unlock()
|
||||
if u.presence != nil {
|
||||
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)
|
||||
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 {
|
||||
u.log.Error("message disconnect", "endpoint", conn.EndpointID, "err", err)
|
||||
}
|
||||
@@ -293,7 +333,7 @@ func (u *appUplink) publishResp(ctx context.Context, conn port.ConnInfo, resp pr
|
||||
return
|
||||
}
|
||||
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 {
|
||||
tooLarge := protocol.Resp{
|
||||
V: protocol.Version,
|
||||
@@ -313,20 +353,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 {
|
||||
var peek struct {
|
||||
RID string `json:"rid"`
|
||||
|
||||
@@ -1207,3 +1207,111 @@ issue #3 未关闭,`feat/fix-3-downlink-deadlock` 未合入 `main`。下面是
|
||||
5. **建议的正确方向**
|
||||
- 在 broker 把对本连接的下行 `InjectPacket` 与上行 worker 解耦:上行读循环先写完 PUBACK,处理 `HandleUplink` 期间不要同步向本连接注入;handler 返回后再发 `resp` 和 `group_event`。不要靠固定 `Sleep`。`InlineClient: true` 保持,`OnPublish` 对 InlineClient 继续放行。
|
||||
- 覆盖 presence 等其他同步 `PublishDown`,而不只包一层 `emit`。
|
||||
|
||||
### 复审修复 B-01
|
||||
|
||||
- 日期:2026-09-30
|
||||
- 原条款:DEVELOPMENT 第 5 节装配 mochi;未写客户端 Receive Maximum。Gitea #8。
|
||||
- 实际做法:`OnConnect` 在心跳校正后调用 `cl.State.Inflight.ResetSendQuota(0)`,不 fork mochi。CONNECT 声明的 Receive Maximum 小于 256 时打 warn,连接仍接受。应用层窗口(推送 32、回执 64、在途 resp 等)约束未确认的 QoS 1。
|
||||
- 原因:mochi v2.7.9 在 `sendQuota>0` 时走 `NextImmediate` 递归读锁,并可因补发后删除 inflight 泄漏配额;已验证置 0 绕开整条路径。
|
||||
- 备选方案:fork 修补 mochi(只修递归读锁仍观察到停滞)。
|
||||
- 影响:服务端不再执行客户端 Receive Maximum;裸设备若带过小的 Receive Maximum,实际在途可能超过该值。
|
||||
|
||||
### 复审修复 B-02
|
||||
|
||||
- 日期:2026-09-30
|
||||
- 原条款:DEVELOPMENT 7.5 / DEVIATIONS N1/N2 第 4 条:大帧名额在 PUBACK、丢弃、断线时归还。Gitea #9。
|
||||
- 实际做法:`OnQosPublish` 按 PacketID 记下超过 64KiB 的出站包;`OnQosComplete`/`OnQosDropped`/断线按 ID 归还。获取名额最多等 5 秒,超时返回 `ErrLargeFrameTimeout`。Publish 未产生 inflight(无订阅者、队列丢弃)时立即归还。不采用「发布完成即归还」。
|
||||
- 原因:mochi 传给 `OnQosComplete` 的是 PUBACK,没有载荷,旧实现从未归还。
|
||||
- 备选方案:发布后立即归还(会把卡死点挪到消息包那份名额)。
|
||||
- 影响:只在 broker 保留一份全局 64 名额;确认超时仍由消息线踢线/清标记触发断线归还。
|
||||
|
||||
### 复审修复 B-05
|
||||
|
||||
- 日期:2026-09-30
|
||||
- 原条款:DEVELOPMENT 第 5 节连接表;Gitea #12。
|
||||
- 实际做法:`OnConnect` 只在认证通过时写入 `byClient`/`byConnID`;拒绝与内部错误不登记。`connState` 增加 `established` 与 `createdAt`,每分钟清扫未建立且已关闭超过 1 分钟的条目。按连接代号查找改为 O(1)。
|
||||
- 原因:mochi 在认证失败路径不调用 `OnDisconnect`,旧实现会永久泄漏。
|
||||
- 备选方案:失败路径也登记再在 Authenticate 返回 false 时删除(仍覆盖不了 CONNACK 失败)。
|
||||
- 影响:失败连接不再占用查找路径;行为对客户端不变(仍回 0x86 或不回 CONNACK)。
|
||||
|
||||
### 复审修复 B-07
|
||||
|
||||
- 日期:2026-09-30
|
||||
- 原条款:PRD §8 日志无正文、无密码、无令牌。Gitea #14。
|
||||
- 实际做法:`broker.New` 给 mochi 包一层 slog.Handler,把 `packets.Packet` / `*packets.Packet` 换成类型、QoS、包号、主题、正文长度。
|
||||
- 原因:默认 info 下第二个 CONNECT、3.1.1 发到错误主题等会把整包写入 JSON 日志。
|
||||
- 备选方案:改 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 优雅停机仍未做。
|
||||
|
||||
+118
-10
@@ -32,6 +32,9 @@ type Login struct {
|
||||
// 内存中的 session_used_at(毫秒)与上次落库时间。
|
||||
usedAt map[string]int64
|
||||
lastFlush map[string]int64
|
||||
|
||||
verifyMu sync.Mutex
|
||||
verifySem map[string]chan struct{}
|
||||
}
|
||||
|
||||
// LoginOptions 装配 Login。
|
||||
@@ -67,9 +70,15 @@ func NewLogin(opts LoginOptions) *Login {
|
||||
Now: now,
|
||||
usedAt: 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。
|
||||
func (l *Login) Authenticate(ctx context.Context, endpointID string, password []byte, remoteIP string) (AuthResult, error) {
|
||||
if l == nil || l.DB == nil {
|
||||
@@ -78,6 +87,11 @@ func (l *Login) Authenticate(ctx context.Context, endpointID string, password []
|
||||
if endpointID == "" {
|
||||
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)
|
||||
if err != nil {
|
||||
@@ -107,6 +121,8 @@ type endpointAuthRow struct {
|
||||
loginHash string
|
||||
sessionHash []byte // 原始 32 字节;无令牌时 nil
|
||||
sessionUsedAt int64 // 毫秒;无则 0
|
||||
onlineSince int64
|
||||
offlineSince int64
|
||||
}
|
||||
|
||||
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
|
||||
sessHex sql.NullString
|
||||
usedAt sql.NullInt64
|
||||
online sql.NullInt64
|
||||
offline sql.NullInt64
|
||||
)
|
||||
err := l.DB.Read.QueryRowContext(ctx, `
|
||||
SELECT login_hash, enabled, session_hash, session_used_at
|
||||
FROM endpoints WHERE id = ?`, id).Scan(&loginHash, &enabled, &sessHex, &usedAt)
|
||||
SELECT login_hash, enabled, session_hash, session_used_at, online_since, offline_since
|
||||
FROM endpoints WHERE id = ?`, id).Scan(&loginHash, &enabled, &sessHex, &usedAt, &online, &offline)
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return endpointAuthRow{}, ErrEndpointNotFound
|
||||
@@ -132,6 +150,12 @@ FROM endpoints WHERE id = ?`, id).Scan(&loginHash, &enabled, &sessHex, &usedAt)
|
||||
if usedAt.Valid {
|
||||
row.sessionUsedAt = usedAt.Int64
|
||||
}
|
||||
if online.Valid {
|
||||
row.onlineSince = online.Int64
|
||||
}
|
||||
if offline.Valid {
|
||||
row.offlineSince = offline.Int64
|
||||
}
|
||||
if sessHex.Valid && sessHex.String != "" {
|
||||
raw, decErr := hex.DecodeString(sessHex.String)
|
||||
if decErr != nil || len(raw) != 32 {
|
||||
@@ -160,11 +184,8 @@ func (l *Login) authSession(ctx context.Context, endpointID, token string, row e
|
||||
usedAt = mem
|
||||
}
|
||||
l.usedMu.Unlock()
|
||||
if l.IdleDays > 0 {
|
||||
idle := time.Duration(l.IdleDays) * 24 * time.Hour
|
||||
if usedAt <= 0 || now.Sub(time.UnixMilli(usedAt)) > idle {
|
||||
return false, nil
|
||||
}
|
||||
if l.IdleDays > 0 && !sessionIdleOK(now, usedAt, row.onlineSince, row.offlineSince, l.IdleDays) {
|
||||
return false, nil
|
||||
}
|
||||
if err := l.touchSessionUsed(ctx, endpointID, nowMs); err != nil {
|
||||
return false, err
|
||||
@@ -204,7 +225,11 @@ func (l *Login) authPassword(ctx context.Context, endpointID, password, remoteIP
|
||||
if l.Pool == nil {
|
||||
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)
|
||||
l.releaseVerify(endpointID)
|
||||
if verErr != nil {
|
||||
return false, "", verErr
|
||||
}
|
||||
@@ -220,12 +245,28 @@ func (l *Login) authPassword(ctx context.Context, endpointID, password, remoteIP
|
||||
}
|
||||
nowMs := l.Now().UnixMilli()
|
||||
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 {
|
||||
_, e := tx.Exec(`
|
||||
res, e := tx.Exec(`
|
||||
UPDATE endpoints
|
||||
SET session_hash = ?, session_issued_at = ?, session_used_at = ?
|
||||
WHERE id = ?`, hashHex, nowMs, nowMs, endpointID)
|
||||
return e
|
||||
WHERE id = ? AND COALESCE(session_hash, '') = ?`, hashHex, nowMs, nowMs, endpointID, oldHex)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
n, nErr := res.RowsAffected()
|
||||
if nErr != nil {
|
||||
return nErr
|
||||
}
|
||||
if n == 0 {
|
||||
return ErrSessionWriteConflict
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if writeErr != nil {
|
||||
return false, "", writeErr
|
||||
@@ -288,6 +329,73 @@ func (l *Login) SessionHashOf(ctx context.Context, endpointID string) ([]byte, e
|
||||
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 暴露给测试。
|
||||
func (l *Login) LooksLikeSessionToken(s string) bool {
|
||||
return strings.HasPrefix(s, "nst_")
|
||||
|
||||
@@ -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 }
|
||||
@@ -0,0 +1,225 @@
|
||||
package broker
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"io"
|
||||
"net"
|
||||
"runtime"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"github.com/mochi-mqtt/server/v2/packets"
|
||||
)
|
||||
|
||||
func TestReceiveMaximumDoesNotDeadlockPublish(t *testing.T) {
|
||||
b, err := New(Options{Authenticator: AllowAuthenticator{}})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = b.Close() }()
|
||||
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = ln.Close() }()
|
||||
|
||||
clientDone := make(chan struct{})
|
||||
go func() {
|
||||
defer close(clientDone)
|
||||
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)
|
||||
}
|
||||
defer func() { _ = w.Close() }()
|
||||
|
||||
endpoint := "ep-rm-quota"
|
||||
writeConnectFull(t, w, endpoint, 30, 0, 20)
|
||||
readExactPacket(t, w, packets.Connack, 3*time.Second)
|
||||
writeSubscribe(t, w, downTopic(endpoint))
|
||||
readExactPacket(t, w, packets.Suback, 3*time.Second)
|
||||
|
||||
deadline := time.Now().Add(3 * time.Second)
|
||||
for {
|
||||
if _, ok := b.ConnInfoOf(endpoint); ok {
|
||||
break
|
||||
}
|
||||
if time.Now().After(deadline) {
|
||||
t.Fatal("session not established")
|
||||
}
|
||||
time.Sleep(5 * time.Millisecond)
|
||||
}
|
||||
|
||||
var received atomic.Int64
|
||||
var writeMu sync.Mutex
|
||||
stop := make(chan struct{})
|
||||
var stopOnce sync.Once
|
||||
halt := func() { stopOnce.Do(func() { close(stop) }) }
|
||||
defer halt()
|
||||
|
||||
go func() {
|
||||
for {
|
||||
select {
|
||||
case <-stop:
|
||||
return
|
||||
default:
|
||||
}
|
||||
_ = 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
|
||||
}
|
||||
typ := hdr[0] >> 4
|
||||
if typ != packets.Publish {
|
||||
continue
|
||||
}
|
||||
qos := (hdr[0] >> 1) & 0x3
|
||||
received.Add(1)
|
||||
if qos == 0 {
|
||||
continue
|
||||
}
|
||||
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 {
|
||||
continue
|
||||
}
|
||||
ack := packets.Packet{
|
||||
FixedHeader: packets.FixedHeader{Type: packets.Puback},
|
||||
ProtocolVersion: 5,
|
||||
PacketID: pk.PacketID,
|
||||
}
|
||||
var ab bytes.Buffer
|
||||
_ = ack.PubackEncode(&ab)
|
||||
writeMu.Lock()
|
||||
_, _ = w.Write(ab.Bytes())
|
||||
writeMu.Unlock()
|
||||
}
|
||||
}()
|
||||
|
||||
go func() {
|
||||
tick := time.NewTicker(2 * time.Millisecond)
|
||||
defer tick.Stop()
|
||||
pk := packets.Packet{
|
||||
FixedHeader: packets.FixedHeader{Type: packets.Pingreq},
|
||||
ProtocolVersion: 5,
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
_ = pk.PingreqEncode(&buf)
|
||||
ping := append([]byte(nil), buf.Bytes()...)
|
||||
for {
|
||||
select {
|
||||
case <-stop:
|
||||
return
|
||||
case <-tick.C:
|
||||
writeMu.Lock()
|
||||
_, _ = w.Write(ping)
|
||||
writeMu.Unlock()
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
payload := []byte(`{"v":1,"type":"resp","rid":"x"}`)
|
||||
pubCtx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < 8; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for pubCtx.Err() == nil {
|
||||
_ = b.PublishDown(pubCtx, endpoint, "", payload, port.PublishOpts{QoS: 1})
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
runFor := 5 * time.Second
|
||||
watch := 2 * time.Second
|
||||
start := time.Now()
|
||||
last := received.Load()
|
||||
lastChange := time.Now()
|
||||
for time.Since(start) < runFor {
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
n := received.Load()
|
||||
if n > last {
|
||||
last = n
|
||||
lastChange = time.Now()
|
||||
}
|
||||
if time.Since(lastChange) > watch {
|
||||
buf := make([]byte, 1<<20)
|
||||
nstack := runtime.Stack(buf, true)
|
||||
halt()
|
||||
cancel()
|
||||
_ = w.Close()
|
||||
t.Fatalf("progress stalled at %d after %s\n%s", last, time.Since(lastChange), buf[:nstack])
|
||||
}
|
||||
}
|
||||
cancel()
|
||||
wg.Wait()
|
||||
halt()
|
||||
if last < 100 {
|
||||
t.Fatalf("too few publishes delivered: %d", last)
|
||||
}
|
||||
|
||||
_ = w.Close()
|
||||
select {
|
||||
case <-clientDone:
|
||||
case <-time.After(3 * time.Second):
|
||||
}
|
||||
}
|
||||
|
||||
func readExactPacket(t *testing.T, conn net.Conn, wantType byte, timeout time.Duration) {
|
||||
t.Helper()
|
||||
_ = conn.SetReadDeadline(time.Now().Add(timeout))
|
||||
hdr := make([]byte, 1)
|
||||
if _, err := io.ReadFull(conn, hdr); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if hdr[0]>>4 != wantType {
|
||||
t.Fatalf("want packet type %d got %d", wantType, hdr[0]>>4)
|
||||
}
|
||||
rem, err := readRemainingLengthConn(conn)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
body := make([]byte, rem)
|
||||
if _, err := io.ReadFull(conn, body); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func readRemainingLengthConn(r io.Reader) (int, error) {
|
||||
var mul uint32 = 1
|
||||
var value uint32
|
||||
for i := 0; i < 4; i++ {
|
||||
var b [1]byte
|
||||
if _, err := io.ReadFull(r, b[:]); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
value += uint32(b[0]&127) * mul
|
||||
if b[0]&128 == 0 {
|
||||
return int(value), nil
|
||||
}
|
||||
mul *= 128
|
||||
}
|
||||
return 0, io.ErrUnexpectedEOF
|
||||
}
|
||||
@@ -0,0 +1,223 @@
|
||||
package broker
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"github.com/mochi-mqtt/server/v2/packets"
|
||||
)
|
||||
|
||||
func TestLargeFrameQuotaReleasedOnPuback(t *testing.T) {
|
||||
b, w, done := startTCPClient(t, "ep-large-ack")
|
||||
defer func() { _ = b.Close() }()
|
||||
defer func() {
|
||||
_ = w.Close()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(3 * time.Second):
|
||||
}
|
||||
}()
|
||||
|
||||
writeConnect(t, w, "ep-large-ack", 30, 0)
|
||||
readExactPacket(t, w, packets.Connack, 3*time.Second)
|
||||
writeSubscribe(t, w, downTopic("ep-large-ack"))
|
||||
readExactPacket(t, w, packets.Suback, 3*time.Second)
|
||||
waitSession(t, b, "ep-large-ack")
|
||||
|
||||
stop := make(chan struct{})
|
||||
defer close(stop)
|
||||
var writeMu sync.Mutex
|
||||
go autoPuback(w, stop, &writeMu)
|
||||
|
||||
payload := bytes.Repeat([]byte("x"), 70*1024)
|
||||
for i := 0; i < 65; i++ {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||
err := b.PublishDown(ctx, "ep-large-ack", "", payload, port.PublishOpts{QoS: 1})
|
||||
cancel()
|
||||
if err != nil {
|
||||
t.Fatalf("publish %d: %v", i+1, err)
|
||||
}
|
||||
}
|
||||
deadline := time.Now().Add(2 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
if len(b.largeSem) == 0 {
|
||||
return
|
||||
}
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
t.Fatalf("slots still held: %d", len(b.largeSem))
|
||||
}
|
||||
|
||||
func TestLargeFrameQuotaReleasedOnDisconnect(t *testing.T) {
|
||||
b, w, done := startTCPClient(t, "ep-large-disc")
|
||||
defer func() { _ = b.Close() }()
|
||||
|
||||
writeConnect(t, w, "ep-large-disc", 30, 0)
|
||||
readExactPacket(t, w, packets.Connack, 3*time.Second)
|
||||
writeSubscribe(t, w, downTopic("ep-large-disc"))
|
||||
readExactPacket(t, w, packets.Suback, 3*time.Second)
|
||||
waitSession(t, b, "ep-large-disc")
|
||||
|
||||
go func() {
|
||||
buf := make([]byte, 32*1024)
|
||||
for {
|
||||
_, err := w.Read(buf)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
payload := bytes.Repeat([]byte("y"), 70*1024)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||
if err := b.PublishDown(ctx, "ep-large-disc", "", payload, port.PublishOpts{QoS: 1}); err != nil {
|
||||
cancel()
|
||||
t.Fatal(err)
|
||||
}
|
||||
cancel()
|
||||
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")
|
||||
}
|
||||
_ = w.Close()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("client attach did not return")
|
||||
}
|
||||
deadline = time.Now().Add(2 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
if len(b.largeSem) == 0 {
|
||||
return
|
||||
}
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
t.Fatalf("slots after disconnect: %d", len(b.largeSem))
|
||||
}
|
||||
|
||||
func TestLargeFrameQuotaReleasedWithoutSubscriber(t *testing.T) {
|
||||
b, w, done := startTCPClient(t, "ep-large-nosub")
|
||||
defer func() { _ = b.Close() }()
|
||||
defer func() {
|
||||
_ = w.Close()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(3 * time.Second):
|
||||
}
|
||||
}()
|
||||
|
||||
writeConnect(t, w, "ep-large-nosub", 30, 0)
|
||||
readExactPacket(t, w, packets.Connack, 3*time.Second)
|
||||
waitSession(t, b, "ep-large-nosub")
|
||||
|
||||
payload := bytes.Repeat([]byte("z"), 70*1024)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||
err := b.PublishDown(ctx, "ep-large-nosub", "", payload, port.PublishOpts{QoS: 1})
|
||||
cancel()
|
||||
if !errors.Is(err, ErrNotSubscribed) {
|
||||
t.Fatalf("err=%v want ErrNotSubscribed", err)
|
||||
}
|
||||
if len(b.largeSem) != 0 {
|
||||
t.Fatalf("held slots without subscriber: %d", len(b.largeSem))
|
||||
}
|
||||
}
|
||||
|
||||
func startTCPClient(t *testing.T, _ string) (*Broker, net.Conn, chan struct{}) {
|
||||
t.Helper()
|
||||
b, err := New(Options{Authenticator: AllowAuthenticator{}})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
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 b, w, done
|
||||
}
|
||||
|
||||
func waitSession(t *testing.T, b *Broker, endpoint string) {
|
||||
t.Helper()
|
||||
deadline := time.Now().Add(3 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
if _, ok := b.ConnInfoOf(endpoint); ok {
|
||||
return
|
||||
}
|
||||
time.Sleep(5 * time.Millisecond)
|
||||
}
|
||||
t.Fatal("session not established")
|
||||
}
|
||||
|
||||
func autoPuback(w net.Conn, stop <-chan struct{}, writeMu *sync.Mutex) {
|
||||
for {
|
||||
select {
|
||||
case <-stop:
|
||||
return
|
||||
default:
|
||||
}
|
||||
_ = 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
|
||||
}
|
||||
if hdr[0]>>4 != packets.Publish {
|
||||
continue
|
||||
}
|
||||
qos := (hdr[0] >> 1) & 0x3
|
||||
if qos == 0 {
|
||||
continue
|
||||
}
|
||||
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 {
|
||||
continue
|
||||
}
|
||||
ack := packets.Packet{
|
||||
FixedHeader: packets.FixedHeader{Type: packets.Puback},
|
||||
ProtocolVersion: 5,
|
||||
PacketID: pk.PacketID,
|
||||
}
|
||||
var ab bytes.Buffer
|
||||
_ = ack.PubackEncode(&ab)
|
||||
writeMu.Lock()
|
||||
_, _ = w.Write(ab.Bytes())
|
||||
writeMu.Unlock()
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,270 @@
|
||||
package broker
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/mochi-mqtt/server/v2/packets"
|
||||
)
|
||||
|
||||
func TestFailedAuthDoesNotLeakConnTable(t *testing.T) {
|
||||
secret := "s3cret-token-xyz"
|
||||
var logBuf bytes.Buffer
|
||||
log := slog.New(slog.NewTextHandler(&logBuf, &slog.HandlerOptions{Level: slog.LevelDebug}))
|
||||
|
||||
b, err := New(Options{Authenticator: RejectAuthenticator{}, Logger: log})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = b.Close() }()
|
||||
|
||||
const n = 200
|
||||
for i := 0; i < n; i++ {
|
||||
dialFailedCONNECT(t, b, func(w net.Conn) {
|
||||
writeConnect(t, w, "ep-rej", 30, 0)
|
||||
})
|
||||
}
|
||||
for i := 0; i < n; i++ {
|
||||
dialFailedCONNECT(t, b, func(w net.Conn) {
|
||||
writeConnectMismatch(t, w)
|
||||
})
|
||||
}
|
||||
b2, err := New(Options{Authenticator: &errAuthenticator{err: context.DeadlineExceeded}})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = b2.Close() }()
|
||||
for i := 0; i < n; i++ {
|
||||
dialFailedCONNECT(t, b2, func(w net.Conn) {
|
||||
writeConnect(t, w, "ep-err", 30, 0)
|
||||
})
|
||||
}
|
||||
if got := len(b.byClient); got != 0 {
|
||||
t.Fatalf("reject/mismatch leaked %d", got)
|
||||
}
|
||||
if got := len(b2.byClient); got != 0 {
|
||||
t.Fatalf("internal error leaked %d", got)
|
||||
}
|
||||
|
||||
// B-07:拒绝路径的 mochi 日志不能带密码
|
||||
b3, err := New(Options{Authenticator: AllowAuthenticator{}, Logger: log})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = b3.Close() }()
|
||||
r, w := net.Pipe()
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
defer close(done)
|
||||
_ = b3.AttachTCP(r)
|
||||
}()
|
||||
writeConnectWithPassword(t, w, "ep-log", secret)
|
||||
readExactPacket(t, w, packets.Connack, 3*time.Second)
|
||||
writeConnectWithPassword(t, w, "ep-log", secret) // 同一连接第二个 CONNECT
|
||||
_ = w.Close()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(2 * time.Second):
|
||||
}
|
||||
out := logBuf.String()
|
||||
if strings.Contains(out, secret) {
|
||||
t.Fatalf("log contains password: %s", out)
|
||||
}
|
||||
if strings.Contains(out, base64.StdEncoding.EncodeToString([]byte(secret))) {
|
||||
t.Fatalf("log contains password base64: %s", out)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSweepUnestablishedClosedConn(t *testing.T) {
|
||||
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)
|
||||
}()
|
||||
writeConnect(t, w, "ep-sweep", 30, 0)
|
||||
_ = w.Close()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(3 * time.Second):
|
||||
}
|
||||
b.sweepUnestablished(0)
|
||||
if got := len(b.byClient); got != 0 {
|
||||
t.Fatalf("after sweep byClient=%d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLookupByConnIDIndependentOfFailedConns(t *testing.T) {
|
||||
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)
|
||||
}()
|
||||
connectAndSubscribe(t, w, "ep-ok", 0)
|
||||
waitSession(t, b, "ep-ok")
|
||||
info, ok := b.ConnInfoOf("ep-ok")
|
||||
if !ok {
|
||||
t.Fatal("missing session")
|
||||
}
|
||||
st := b.lookupConn("ep-ok", info.ConnID)
|
||||
if st == nil {
|
||||
t.Fatal("lookup by conn id")
|
||||
}
|
||||
_ = w.Close()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(3 * time.Second):
|
||||
}
|
||||
}
|
||||
|
||||
func dialFailedCONNECT(t *testing.T, b *Broker, write func(net.Conn)) {
|
||||
t.Helper()
|
||||
r, w := net.Pipe()
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
defer close(done)
|
||||
_ = b.AttachTCP(r)
|
||||
}()
|
||||
write(w)
|
||||
_ = w.Close()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("attach did not return")
|
||||
}
|
||||
}
|
||||
|
||||
func writeConnectMismatch(t *testing.T, w net.Conn) {
|
||||
t.Helper()
|
||||
pk := packets.Packet{
|
||||
FixedHeader: packets.FixedHeader{Type: packets.Connect},
|
||||
ProtocolVersion: 5,
|
||||
Connect: packets.ConnectParams{
|
||||
ProtocolName: []byte("MQTT"),
|
||||
Clean: true,
|
||||
ClientIdentifier: "id-a",
|
||||
Keepalive: 30,
|
||||
UsernameFlag: true,
|
||||
Username: []byte("id-b"),
|
||||
PasswordFlag: true,
|
||||
Password: []byte("nope"),
|
||||
},
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
if err := pk.ConnectEncode(&buf); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := w.Write(buf.Bytes()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func writeConnectWithPassword(t *testing.T, w net.Conn, endpoint, password string) {
|
||||
t.Helper()
|
||||
pk := packets.Packet{
|
||||
FixedHeader: packets.FixedHeader{Type: packets.Connect},
|
||||
ProtocolVersion: 5,
|
||||
Connect: packets.ConnectParams{
|
||||
ProtocolName: []byte("MQTT"),
|
||||
Clean: true,
|
||||
ClientIdentifier: endpoint,
|
||||
Keepalive: 30,
|
||||
UsernameFlag: true,
|
||||
Username: []byte(endpoint),
|
||||
PasswordFlag: true,
|
||||
Password: []byte(password),
|
||||
},
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
if err := pk.ConnectEncode(&buf); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := w.Write(buf.Bytes()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMQTT311UnauthorizedPublishOmitsPayloadInLogs(t *testing.T) {
|
||||
var logBuf bytes.Buffer
|
||||
log := slog.New(slog.NewTextHandler(&logBuf, &slog.HandlerOptions{Level: slog.LevelDebug}))
|
||||
b, err := New(Options{Authenticator: AllowAuthenticator{}, Logger: log})
|
||||
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)
|
||||
}()
|
||||
pk := packets.Packet{
|
||||
FixedHeader: packets.FixedHeader{Type: packets.Connect},
|
||||
ProtocolVersion: 4,
|
||||
Connect: packets.ConnectParams{
|
||||
ProtocolName: []byte("MQTT"),
|
||||
Clean: true,
|
||||
ClientIdentifier: "ep311",
|
||||
Keepalive: 30,
|
||||
UsernameFlag: true,
|
||||
Username: []byte("ep311"),
|
||||
PasswordFlag: true,
|
||||
Password: []byte("test"),
|
||||
},
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
if err := pk.ConnectEncode(&buf); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := w.Write(buf.Bytes()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_ = w.SetReadDeadline(time.Now().Add(3 * time.Second))
|
||||
raw := make([]byte, 256)
|
||||
if _, err := io.ReadAtLeast(w, raw, 2); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
body := []byte(`{"talk_password":"super-secret-body"}`)
|
||||
pub := packets.Packet{
|
||||
FixedHeader: packets.FixedHeader{Type: packets.Publish, Qos: 1},
|
||||
ProtocolVersion: 4,
|
||||
TopicName: "nix/c/other/up",
|
||||
PacketID: 7,
|
||||
Payload: body,
|
||||
}
|
||||
buf.Reset()
|
||||
if err := pub.PublishEncode(&buf); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, _ = w.Write(buf.Bytes())
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
_ = w.Close()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(2 * time.Second):
|
||||
}
|
||||
out := logBuf.String()
|
||||
if strings.Contains(out, "super-secret-body") {
|
||||
t.Fatalf("log contains publish payload: %s", out)
|
||||
}
|
||||
}
|
||||
+258
-71
@@ -34,6 +34,24 @@ var ErrPayloadTooLarge = errors.New("broker: payload exceeds client limit")
|
||||
// ErrNoConnection 目标端没有当前连接。
|
||||
var ErrNoConnection = errors.New("broker: no active connection")
|
||||
|
||||
// ErrLargeFrameTimeout 全局大帧名额在有界等待内拿不到。
|
||||
var ErrLargeFrameTimeout = errors.New("broker: large frame quota timeout")
|
||||
|
||||
// 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 默认拒绝)。
|
||||
type AuthResult struct {
|
||||
OK bool
|
||||
@@ -87,12 +105,17 @@ type Broker struct {
|
||||
connsMu sync.RWMutex
|
||||
current map[string]*connState
|
||||
byClient map[*mqtt.Client]*connState
|
||||
byConnID map[port.ConnID]*connState
|
||||
closedCh chan struct{}
|
||||
|
||||
queuesMu sync.Mutex
|
||||
queues map[string]*uplinkQueue
|
||||
|
||||
largeSem chan struct{}
|
||||
closed atomic.Bool
|
||||
|
||||
lifeMu sync.Mutex
|
||||
lifeLocks map[string]*sync.Mutex
|
||||
}
|
||||
|
||||
type connState struct {
|
||||
@@ -108,8 +131,18 @@ type connState struct {
|
||||
sessionToken string
|
||||
handshook bool
|
||||
subscribedDown bool
|
||||
largeHeld int
|
||||
largePIDs map[uint16]struct{}
|
||||
largePending int
|
||||
metricsCounted bool
|
||||
established bool
|
||||
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
|
||||
|
||||
handshakeTimer *time.Timer
|
||||
@@ -129,6 +162,7 @@ func New(opts Options) (*Broker, error) {
|
||||
if log == nil {
|
||||
log = slog.Default()
|
||||
}
|
||||
log = slog.New(newRedactHandler(log.Handler()))
|
||||
|
||||
caps := mqtt.NewDefaultServerCapabilities()
|
||||
caps.MaximumClients = maxClients
|
||||
@@ -151,16 +185,19 @@ func New(opts Options) (*Broker, error) {
|
||||
})
|
||||
|
||||
b := &Broker{
|
||||
server: srv,
|
||||
auth: auth,
|
||||
uplink: uplink,
|
||||
log: log,
|
||||
onDrop: opts.OnPublishDropped,
|
||||
metrics: opts.Metrics,
|
||||
current: make(map[string]*connState),
|
||||
byClient: make(map[*mqtt.Client]*connState),
|
||||
queues: make(map[string]*uplinkQueue),
|
||||
largeSem: make(chan struct{}, largeFrameSlots),
|
||||
server: srv,
|
||||
auth: auth,
|
||||
uplink: uplink,
|
||||
log: log,
|
||||
onDrop: opts.OnPublishDropped,
|
||||
metrics: opts.Metrics,
|
||||
current: make(map[string]*connState),
|
||||
byClient: make(map[*mqtt.Client]*connState),
|
||||
byConnID: make(map[port.ConnID]*connState),
|
||||
closedCh: make(chan struct{}),
|
||||
queues: make(map[string]*uplinkQueue),
|
||||
largeSem: make(chan struct{}, largeFrameSlots),
|
||||
lifeLocks: make(map[string]*sync.Mutex),
|
||||
}
|
||||
b.hook = &nixHook{b: b}
|
||||
if err := srv.AddHook(b.hook, nil); err != nil {
|
||||
@@ -169,6 +206,7 @@ func New(opts Options) (*Broker, error) {
|
||||
if err := srv.Serve(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
go b.sweepLoop()
|
||||
return b, nil
|
||||
}
|
||||
|
||||
@@ -180,14 +218,45 @@ func (b *Broker) Close() error {
|
||||
if b.closed.Swap(true) {
|
||||
return nil
|
||||
}
|
||||
select {
|
||||
case <-b.closedCh:
|
||||
default:
|
||||
close(b.closedCh)
|
||||
}
|
||||
b.queuesMu.Lock()
|
||||
for _, q := range b.queues {
|
||||
q.close()
|
||||
}
|
||||
b.queues = make(map[string]*uplinkQueue)
|
||||
b.queuesMu.Unlock()
|
||||
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;阻塞到连接结束。
|
||||
func (b *Broker) AttachTCP(conn net.Conn) error {
|
||||
return b.server.EstablishConnection("tcp", conn)
|
||||
@@ -200,73 +269,128 @@ func (b *Broker) AttachWS(conn net.Conn) error {
|
||||
|
||||
// PublishDown 实现 port.Downlink。
|
||||
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
|
||||
if qos > 1 {
|
||||
qos = 1
|
||||
}
|
||||
topic := downTopic(endpointID)
|
||||
large := len(payload) > largeFrameBytes
|
||||
|
||||
if large {
|
||||
select {
|
||||
case b.largeSem <- struct{}{}:
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
}
|
||||
st.mu.Lock()
|
||||
st.largeHeld++
|
||||
st.mu.Unlock()
|
||||
}
|
||||
|
||||
if err := b.server.Publish(topic, payload, false, qos); err != nil {
|
||||
if large {
|
||||
b.releaseOneLarge(st)
|
||||
}
|
||||
return err
|
||||
}
|
||||
if large && qos == 0 {
|
||||
b.releaseOneLarge(st)
|
||||
}
|
||||
return nil
|
||||
return b.enqueueDownlink(endpointID, connID, payload, qos, "")
|
||||
}
|
||||
|
||||
func (b *Broker) releaseOneLarge(st *connState) {
|
||||
// PublishThenDisconnect 把一帧写入该连接下行队列,写出后再断开(无固定 sleep)。
|
||||
func (b *Broker) PublishThenDisconnect(_ context.Context, endpointID string, connID port.ConnID, payload []byte, qos byte, reason port.DisconnectReason) error {
|
||||
if qos > 1 {
|
||||
qos = 1
|
||||
}
|
||||
if reason == "" {
|
||||
reason = port.DisconnectNormal
|
||||
}
|
||||
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
|
||||
}
|
||||
|
||||
st.mu.Lock()
|
||||
if st.largeHeld > 0 {
|
||||
st.largeHeld--
|
||||
st.mu.Unlock()
|
||||
select {
|
||||
case <-b.largeSem:
|
||||
default:
|
||||
}
|
||||
return
|
||||
maxRecv := st.maxRecvBytes
|
||||
closing := st.closing
|
||||
superseded := st.superseded
|
||||
st.mu.Unlock()
|
||||
if closing || superseded {
|
||||
return ErrNoConnection
|
||||
}
|
||||
if !b.hasDownSub(st) {
|
||||
return ErrNotSubscribed
|
||||
}
|
||||
|
||||
limit := EffectivePayloadLimit(st.maxPacketSize, maxRecv)
|
||||
if limit > 0 && len(payload) > limit {
|
||||
return ErrPayloadTooLarge
|
||||
}
|
||||
|
||||
return st.enqueueDown(downItem{
|
||||
payload: append([]byte(nil), payload...),
|
||||
qos: qos,
|
||||
disconnect: disconnect,
|
||||
})
|
||||
}
|
||||
|
||||
func (b *Broker) acquireLarge(ctx context.Context) error {
|
||||
timer := time.NewTimer(largeAcquireWait)
|
||||
defer timer.Stop()
|
||||
select {
|
||||
case b.largeSem <- struct{}{}:
|
||||
return nil
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
case <-timer.C:
|
||||
return ErrLargeFrameTimeout
|
||||
}
|
||||
}
|
||||
|
||||
func (b *Broker) releaseLargeSlot() {
|
||||
select {
|
||||
case <-b.largeSem:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func (b *Broker) finishLargePublish(st *connState) {
|
||||
b.reconcileLargeInflight(st)
|
||||
st.mu.Lock()
|
||||
n := st.largePending
|
||||
st.largePending = 0
|
||||
st.mu.Unlock()
|
||||
for i := 0; i < n; i++ {
|
||||
b.releaseLargeSlot()
|
||||
}
|
||||
}
|
||||
|
||||
func (b *Broker) releaseLargePID(st *connState, id uint16) {
|
||||
st.mu.Lock()
|
||||
_, ok := st.largePIDs[id]
|
||||
if ok {
|
||||
delete(st.largePIDs, id)
|
||||
}
|
||||
st.mu.Unlock()
|
||||
if ok {
|
||||
b.releaseLargeSlot()
|
||||
}
|
||||
}
|
||||
|
||||
func (b *Broker) reconcileLargeInflight(st *connState) {
|
||||
if st == nil {
|
||||
return
|
||||
}
|
||||
st.mu.Lock()
|
||||
ids := make([]uint16, 0, len(st.largePIDs))
|
||||
for id := range st.largePIDs {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
st.mu.Unlock()
|
||||
for _, id := range ids {
|
||||
if st.client != nil {
|
||||
if _, ok := st.client.State.Inflight.Get(id); ok {
|
||||
continue
|
||||
}
|
||||
}
|
||||
b.releaseLargePID(st, id)
|
||||
}
|
||||
}
|
||||
|
||||
func (b *Broker) releaseAllLarge(st *connState) {
|
||||
st.mu.Lock()
|
||||
n := st.largeHeld
|
||||
st.largeHeld = 0
|
||||
n := len(st.largePIDs) + st.largePending
|
||||
st.largePIDs = nil
|
||||
st.largePending = 0
|
||||
st.mu.Unlock()
|
||||
for i := 0; i < n; i++ {
|
||||
select {
|
||||
case <-b.largeSem:
|
||||
default:
|
||||
}
|
||||
b.releaseLargeSlot()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -295,16 +419,40 @@ func (b *Broker) lookupConn(endpointID string, connID port.ConnID) *connState {
|
||||
b.connsMu.RLock()
|
||||
defer b.connsMu.RUnlock()
|
||||
if connID != "" {
|
||||
for _, st := range b.byClient {
|
||||
if st.endpointID == endpointID && st.connID == connID {
|
||||
return st
|
||||
}
|
||||
st := b.byConnID[connID]
|
||||
if st != nil && st.endpointID == endpointID {
|
||||
return st
|
||||
}
|
||||
return nil
|
||||
}
|
||||
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 {
|
||||
return "nix/c/" + endpointID + "/down"
|
||||
}
|
||||
@@ -313,6 +461,11 @@ func upTopic(endpointID string) string {
|
||||
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 {
|
||||
limit := 0
|
||||
if maxPacketSize > 0 {
|
||||
@@ -412,12 +565,46 @@ func (b *Broker) CurrentConnID(endpointID string) (port.ConnID, bool) {
|
||||
func (b *Broker) connStateOf(endpointID string, connID port.ConnID) *connState {
|
||||
b.connsMu.RLock()
|
||||
defer b.connsMu.RUnlock()
|
||||
for _, st := range b.byClient {
|
||||
if st.endpointID == endpointID && st.connID == connID {
|
||||
return st
|
||||
st := b.byConnID[connID]
|
||||
if st == nil || st.endpointID != endpointID {
|
||||
return nil
|
||||
}
|
||||
return st
|
||||
}
|
||||
|
||||
func (b *Broker) sweepLoop() {
|
||||
tick := time.NewTicker(time.Minute)
|
||||
defer tick.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-tick.C:
|
||||
b.sweepUnestablished(time.Minute)
|
||||
case <-b.closedCh:
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (b *Broker) sweepUnestablished(minAge time.Duration) {
|
||||
now := time.Now()
|
||||
b.connsMu.Lock()
|
||||
defer b.connsMu.Unlock()
|
||||
for cl, st := range b.byClient {
|
||||
if st.established {
|
||||
continue
|
||||
}
|
||||
if cl != nil && !cl.Closed() {
|
||||
continue
|
||||
}
|
||||
if minAge > 0 && now.Sub(st.createdAt) < minAge {
|
||||
continue
|
||||
}
|
||||
delete(b.byClient, cl)
|
||||
delete(b.byConnID, st.connID)
|
||||
if b.current[st.endpointID] == st {
|
||||
delete(b.current, st.endpointID)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (b *Broker) hasDownSub(st *connState) bool {
|
||||
|
||||
@@ -212,6 +212,11 @@ func connectAndSubscribe(t *testing.T, w net.Conn, endpoint string, maxPacket ui
|
||||
}
|
||||
|
||||
func writeConnect(t *testing.T, w net.Conn, endpoint string, keepalive uint16, maxPacket uint32) {
|
||||
t.Helper()
|
||||
writeConnectFull(t, w, endpoint, keepalive, maxPacket, 0)
|
||||
}
|
||||
|
||||
func writeConnectFull(t *testing.T, w net.Conn, endpoint string, keepalive uint16, maxPacket uint32, receiveMax uint16) {
|
||||
t.Helper()
|
||||
pk := packets.Packet{
|
||||
FixedHeader: packets.FixedHeader{Type: packets.Connect},
|
||||
@@ -228,6 +233,7 @@ func writeConnect(t *testing.T, w net.Conn, endpoint string, keepalive uint16, m
|
||||
},
|
||||
Properties: packets.Properties{
|
||||
MaximumPacketSize: maxPacket,
|
||||
ReceiveMaximum: receiveMax,
|
||||
},
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
+88
-14
@@ -3,6 +3,7 @@ package broker
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
mqtt "github.com/mochi-mqtt/server/v2"
|
||||
@@ -25,8 +26,11 @@ func (h *nixHook) Provides(b byte) bool {
|
||||
mqtt.OnPublishDropped,
|
||||
mqtt.OnSessionEstablished,
|
||||
mqtt.OnDisconnect,
|
||||
mqtt.OnQosPublish,
|
||||
mqtt.OnQosComplete,
|
||||
mqtt.OnQosDropped,
|
||||
mqtt.OnSubscribed,
|
||||
mqtt.OnPacketSent,
|
||||
}, []byte{b})
|
||||
}
|
||||
|
||||
@@ -45,12 +49,11 @@ func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error {
|
||||
remoteIP: remoteIP,
|
||||
client: cl,
|
||||
maxPacketSize: pk.Properties.MaximumPacketSize,
|
||||
createdAt: time.Now(),
|
||||
}
|
||||
|
||||
// ClientID、Username 都必须等于端编号
|
||||
if clientID == "" || endpointID == "" || clientID != endpointID {
|
||||
st.authOK = false
|
||||
h.rememberPending(cl, st)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -67,13 +70,26 @@ func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error {
|
||||
cl.State.ServerKeepalive = true
|
||||
}
|
||||
|
||||
res, err := h.b.auth.Authenticate(context.Background(), endpointID, pk.Connect.Password, remoteIP)
|
||||
if err != nil {
|
||||
st.authErr = err
|
||||
h.rememberPending(cl, st)
|
||||
return err // mochi 不回 CONNACK,直接断开
|
||||
// B-01:绕开 mochi 发送配额路径(NextImmediate 递归读锁 + PUBACK 配额泄漏)。
|
||||
// ParseConnect 已按客户端 Receive Maximum 设过 sendQuota;此处一律置 0。
|
||||
if cl.State.Inflight != nil {
|
||||
cl.State.Inflight.ResetSendQuota(0)
|
||||
}
|
||||
st.authOK = res.OK
|
||||
if rm := pk.Properties.ReceiveMaximum; rm > 0 && rm < 256 {
|
||||
h.b.log.Warn("client receive maximum below 256; server ignores MQTT send quota",
|
||||
"endpoint", endpointID, "receive_maximum", rm)
|
||||
}
|
||||
|
||||
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 {
|
||||
return err // mochi 不回 CONNACK,直接断开;不登记连接表
|
||||
}
|
||||
if !res.OK {
|
||||
return nil
|
||||
}
|
||||
st.authOK = true
|
||||
st.sessionToken = res.SessionToken
|
||||
h.rememberPending(cl, st)
|
||||
return nil
|
||||
@@ -82,6 +98,7 @@ func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error {
|
||||
func (h *nixHook) rememberPending(cl *mqtt.Client, st *connState) {
|
||||
h.b.connsMu.Lock()
|
||||
h.b.byClient[cl] = st
|
||||
h.b.byConnID[st.connID] = st
|
||||
h.b.connsMu.Unlock()
|
||||
}
|
||||
|
||||
@@ -157,15 +174,27 @@ func (h *nixHook) OnSubscribed(cl *mqtt.Client, pk packets.Packet, reasonCodes [
|
||||
}
|
||||
}
|
||||
|
||||
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))
|
||||
if h.b.onDrop == nil {
|
||||
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 {
|
||||
if st != nil {
|
||||
st.sentPub.Add(1)
|
||||
}
|
||||
}
|
||||
|
||||
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.connsMu.RLock()
|
||||
st := h.b.byClient[cl]
|
||||
h.b.connsMu.RUnlock()
|
||||
if st != nil {
|
||||
h.b.reconcileLargeInflight(st)
|
||||
}
|
||||
if h.b.onDrop == nil || st == nil {
|
||||
return
|
||||
}
|
||||
h.b.onDrop(context.Background(), st.endpointID, st.connID, append([]byte(nil), pk.Payload...))
|
||||
@@ -174,13 +203,25 @@ func (h *nixHook) OnPublishDropped(cl *mqtt.Client, pk packets.Packet) {
|
||||
func (h *nixHook) OnSessionEstablished(cl *mqtt.Client, _ packets.Packet) {
|
||||
h.b.connsMu.Lock()
|
||||
st := h.b.byClient[cl]
|
||||
var old *connState
|
||||
if st != nil {
|
||||
old = h.b.current[st.endpointID]
|
||||
h.b.current[st.endpointID] = st
|
||||
st.established = true
|
||||
}
|
||||
h.b.connsMu.Unlock()
|
||||
if st == nil {
|
||||
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{
|
||||
ConnID: st.connID,
|
||||
EndpointID: st.endpointID,
|
||||
@@ -197,6 +238,9 @@ func (h *nixHook) OnDisconnect(cl *mqtt.Client, err error, _ bool) {
|
||||
h.b.connsMu.Lock()
|
||||
st := h.b.byClient[cl]
|
||||
delete(h.b.byClient, cl)
|
||||
if st != nil {
|
||||
delete(h.b.byConnID, st.connID)
|
||||
}
|
||||
isCurrent := false
|
||||
if st != nil && h.b.current[st.endpointID] == st {
|
||||
delete(h.b.current, st.endpointID)
|
||||
@@ -206,6 +250,10 @@ func (h *nixHook) OnDisconnect(cl *mqtt.Client, err error, _ bool) {
|
||||
if st == nil {
|
||||
return
|
||||
}
|
||||
lk := h.b.endpointLife(st.endpointID)
|
||||
lk.Lock()
|
||||
st.stopDownLoop()
|
||||
lk.Unlock()
|
||||
h.b.releaseAllLarge(st)
|
||||
h.b.cancelHandshakeDeadline(st.endpointID, st.connID)
|
||||
|
||||
@@ -253,7 +301,7 @@ func (h *nixHook) noteConnectionClose(st *connState) {
|
||||
st.metricsCounted = false
|
||||
}
|
||||
|
||||
func (h *nixHook) OnQosComplete(cl *mqtt.Client, pk packets.Packet) {
|
||||
func (h *nixHook) OnQosPublish(cl *mqtt.Client, pk packets.Packet, _ int64, _ int) {
|
||||
if len(pk.Payload) <= largeFrameBytes {
|
||||
return
|
||||
}
|
||||
@@ -263,5 +311,31 @@ func (h *nixHook) OnQosComplete(cl *mqtt.Client, pk packets.Packet) {
|
||||
if st == nil {
|
||||
return
|
||||
}
|
||||
h.b.releaseOneLarge(st)
|
||||
st.mu.Lock()
|
||||
if st.largePending > 0 {
|
||||
st.largePending--
|
||||
}
|
||||
if st.largePIDs == nil {
|
||||
st.largePIDs = make(map[uint16]struct{})
|
||||
}
|
||||
st.largePIDs[pk.PacketID] = struct{}{}
|
||||
st.mu.Unlock()
|
||||
}
|
||||
|
||||
func (h *nixHook) OnQosComplete(cl *mqtt.Client, pk packets.Packet) {
|
||||
h.releaseLargeByPacketID(cl, pk.PacketID)
|
||||
}
|
||||
|
||||
func (h *nixHook) OnQosDropped(cl *mqtt.Client, pk packets.Packet) {
|
||||
h.releaseLargeByPacketID(cl, pk.PacketID)
|
||||
}
|
||||
|
||||
func (h *nixHook) releaseLargeByPacketID(cl *mqtt.Client, id uint16) {
|
||||
h.b.connsMu.RLock()
|
||||
st := h.b.byClient[cl]
|
||||
h.b.connsMu.RUnlock()
|
||||
if st == nil {
|
||||
return
|
||||
}
|
||||
h.b.releaseLargePID(st, id)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,88 @@
|
||||
package broker
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
|
||||
"github.com/mochi-mqtt/server/v2/packets"
|
||||
)
|
||||
|
||||
type redactHandler struct {
|
||||
inner slog.Handler
|
||||
}
|
||||
|
||||
func newRedactHandler(inner slog.Handler) slog.Handler {
|
||||
if inner == nil {
|
||||
inner = slog.Default().Handler()
|
||||
}
|
||||
return &redactHandler{inner: inner}
|
||||
}
|
||||
|
||||
func (h *redactHandler) Enabled(ctx context.Context, level slog.Level) bool {
|
||||
return h.inner.Enabled(ctx, level)
|
||||
}
|
||||
|
||||
func (h *redactHandler) Handle(ctx context.Context, r slog.Record) error {
|
||||
rec := slog.NewRecord(r.Time, r.Level, r.Message, r.PC)
|
||||
r.Attrs(func(a slog.Attr) bool {
|
||||
rec.AddAttrs(redactSlogAttr(a))
|
||||
return true
|
||||
})
|
||||
return h.inner.Handle(ctx, rec)
|
||||
}
|
||||
|
||||
func (h *redactHandler) WithAttrs(attrs []slog.Attr) slog.Handler {
|
||||
out := make([]slog.Attr, len(attrs))
|
||||
for i, a := range attrs {
|
||||
out[i] = redactSlogAttr(a)
|
||||
}
|
||||
return &redactHandler{inner: h.inner.WithAttrs(out)}
|
||||
}
|
||||
|
||||
func (h *redactHandler) WithGroup(name string) slog.Handler {
|
||||
return &redactHandler{inner: h.inner.WithGroup(name)}
|
||||
}
|
||||
|
||||
func redactSlogAttr(a slog.Attr) slog.Attr {
|
||||
a.Value = a.Value.Resolve()
|
||||
switch v := a.Value.Any().(type) {
|
||||
case packets.Packet:
|
||||
return slog.Any(a.Key, summarizePacket(v))
|
||||
case *packets.Packet:
|
||||
if v == nil {
|
||||
return a
|
||||
}
|
||||
return slog.Any(a.Key, summarizePacket(*v))
|
||||
}
|
||||
if a.Value.Kind() == slog.KindGroup {
|
||||
group := a.Value.Group()
|
||||
out := make([]slog.Attr, len(group))
|
||||
for i, g := range group {
|
||||
out[i] = redactSlogAttr(g)
|
||||
}
|
||||
return slog.Attr{Key: a.Key, Value: slog.GroupValue(out...)}
|
||||
}
|
||||
return a
|
||||
}
|
||||
|
||||
type mqttPacketLog struct {
|
||||
Type string `json:"type"`
|
||||
QoS byte `json:"qos"`
|
||||
PacketID uint16 `json:"packet_id"`
|
||||
Topic string `json:"topic,omitempty"`
|
||||
PayloadLen int `json:"payload_len"`
|
||||
}
|
||||
|
||||
func summarizePacket(pk packets.Packet) mqttPacketLog {
|
||||
name := packets.PacketNames[pk.FixedHeader.Type]
|
||||
if name == "" {
|
||||
name = "unknown"
|
||||
}
|
||||
return mqttPacketLog{
|
||||
Type: name,
|
||||
QoS: pk.FixedHeader.Qos,
|
||||
PacketID: pk.PacketID,
|
||||
Topic: pk.TopicName,
|
||||
PayloadLen: len(pk.Payload),
|
||||
}
|
||||
}
|
||||
+26
-12
@@ -210,6 +210,13 @@ func (s *Session) handleHello(ctx context.Context, conn port.ConnInfo, st *connS
|
||||
}
|
||||
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{
|
||||
ServerTimeMs: s.now().UnixMilli(),
|
||||
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}
|
||||
if err := s.publishJSON(ctx, conn, resp, 1); err != nil {
|
||||
s.log.Error("logout resp", "endpoint", conn.EndpointID, "err", err)
|
||||
raw, err := protocol.Marshal(resp)
|
||||
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() {
|
||||
_ = s.b.Disconnect(context.Background(), conn.EndpointID, conn.ConnID, port.DisconnectNormal)
|
||||
}()
|
||||
}
|
||||
go func() {
|
||||
// 稍等让 QoS1 resp 写入连接,再断开
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
_ = s.b.Disconnect(context.Background(), conn.EndpointID, conn.ConnID, port.DisconnectNormal)
|
||||
}()
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -327,11 +337,15 @@ func (s *Session) fatalKick(ctx context.Context, endpointID, reason string) erro
|
||||
return nil
|
||||
}
|
||||
fatal := protocol.Fatal{V: protocol.Version, Type: protocol.TypeFatal, Reason: reason}
|
||||
_ = s.publishJSON(ctx, info, fatal, 1)
|
||||
go func() {
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
_ = s.b.Disconnect(context.Background(), endpointID, info.ConnID, port.DisconnectFatal)
|
||||
}()
|
||||
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() {
|
||||
_ = s.b.Disconnect(context.Background(), endpointID, info.ConnID, port.DisconnectFatal)
|
||||
}()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user