Compare commits
11
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4059a1576b | ||
|
|
7b209ce3d6 | ||
|
|
bb8ce5f178 | ||
|
|
60bc873ef6 | ||
|
|
613f4bffa4 | ||
|
|
99c134b1ce | ||
|
|
b723ff13cf | ||
|
|
1d6be59652 | ||
|
|
b4789fc9e2 | ||
|
|
de2c64d111 | ||
|
|
8b4da3dc05 |
@@ -0,0 +1,196 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/cookiejar"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/config"
|
||||
)
|
||||
|
||||
// TestUplinkDisableFatalAndRevoked 验证停用在线端收到 fatal,已推送投递收到 revoked。
|
||||
func TestUplinkDisableFatalAndRevoked(t *testing.T) {
|
||||
dataDir := t.TempDir()
|
||||
cfgPath := writeTestConfig(t, dataDir)
|
||||
initAdminForTest(t, dataDir)
|
||||
enableRegistration(t, dataDir, "uplink-code")
|
||||
|
||||
cfg, err := config.Load(cfgPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if vErr := cfg.Validate(); vErr != nil {
|
||||
t.Fatal(vErr)
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
errCh := make(chan error, 1)
|
||||
go func() { errCh <- runServe(ctx, cfg) }()
|
||||
defer func() {
|
||||
cancel()
|
||||
select {
|
||||
case err := <-errCh:
|
||||
if err != nil {
|
||||
t.Errorf("serve exit: %v", err)
|
||||
}
|
||||
case <-time.After(15 * time.Second):
|
||||
t.Error("serve did not stop")
|
||||
}
|
||||
}()
|
||||
|
||||
addr := waitListenAddr(t, dataDir, 15*time.Second)
|
||||
base := "http://" + addr
|
||||
|
||||
registerEP(t, base, "alice", "password12", "Alice")
|
||||
registerEP(t, base, "bob", "password12", "Bob")
|
||||
|
||||
alice := mqttSessionLogin(t, base, "alice", "password12")
|
||||
defer alice.Close()
|
||||
bob := mqttSessionLogin(t, base, "bob", "password12")
|
||||
defer bob.Close()
|
||||
|
||||
delay0 := int64(0)
|
||||
sendResp := alice.Request(t, map[string]any{
|
||||
"v": 1, "type": "send", "rid": "s1", "id": "dm-fatal-1",
|
||||
"to": map[string]any{"kind": "endpoint", "id": "bob"},
|
||||
"body": map[string]any{"enc": "utf8", "data": "to-void"},
|
||||
"delay_ms": delay0,
|
||||
})
|
||||
if !sendResp.OK {
|
||||
t.Fatalf("send: %+v", sendResp)
|
||||
}
|
||||
msg := bob.WaitType(t, "msg", 8*time.Second)
|
||||
if msg["id"] != "dm-fatal-1" {
|
||||
t.Fatalf("bob msg=%v", msg)
|
||||
}
|
||||
|
||||
admin := adminHTTPClient(t, base)
|
||||
disableEP(t, admin, base, "bob")
|
||||
|
||||
fatal := bob.WaitType(t, "fatal", 8*time.Second)
|
||||
if fatal["reason"] != "disabled" {
|
||||
t.Fatalf("fatal=%v", fatal)
|
||||
}
|
||||
revoked := bob.WaitType(t, "revoked", 8*time.Second)
|
||||
if revoked["id"] != "dm-fatal-1" || revoked["reason"] != "endpoint_disabled" {
|
||||
t.Fatalf("revoked=%v", revoked)
|
||||
}
|
||||
}
|
||||
|
||||
// TestUplinkResetPasswordFatal 验证重置登录密码后在线端收到 fatal(password_reset)。
|
||||
func TestUplinkResetPasswordFatal(t *testing.T) {
|
||||
dataDir := t.TempDir()
|
||||
cfgPath := writeTestConfig(t, dataDir)
|
||||
initAdminForTest(t, dataDir)
|
||||
enableRegistration(t, dataDir, "uplink-code")
|
||||
|
||||
cfg, err := config.Load(cfgPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if vErr := cfg.Validate(); vErr != nil {
|
||||
t.Fatal(vErr)
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
errCh := make(chan error, 1)
|
||||
go func() { errCh <- runServe(ctx, cfg) }()
|
||||
defer func() {
|
||||
cancel()
|
||||
select {
|
||||
case err := <-errCh:
|
||||
if err != nil {
|
||||
t.Errorf("serve exit: %v", err)
|
||||
}
|
||||
case <-time.After(15 * time.Second):
|
||||
t.Error("serve did not stop")
|
||||
}
|
||||
}()
|
||||
|
||||
addr := waitListenAddr(t, dataDir, 15*time.Second)
|
||||
base := "http://" + addr
|
||||
|
||||
registerEP(t, base, "carol", "password12", "Carol")
|
||||
carol := mqttSessionLogin(t, base, "carol", "password12")
|
||||
defer carol.Close()
|
||||
|
||||
admin := adminHTTPClient(t, base)
|
||||
resetLoginPassword(t, admin, base, "carol", "password99xx")
|
||||
|
||||
fatal := carol.WaitType(t, "fatal", 8*time.Second)
|
||||
if fatal["reason"] != "password_reset" {
|
||||
t.Fatalf("fatal=%v", fatal)
|
||||
}
|
||||
}
|
||||
|
||||
func adminHTTPClient(t *testing.T, base string) *http.Client {
|
||||
t.Helper()
|
||||
jar, err := cookiejar.New(nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
client := &http.Client{Jar: jar, Timeout: 10 * time.Second}
|
||||
loginBody, _ := json.Marshal(map[string]string{
|
||||
"username": "admin",
|
||||
"password": "test-admin-password-xx",
|
||||
})
|
||||
resp, err := client.Post(base+"/api/admin/login", "application/json", bytes.NewReader(loginBody))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
raw, _ := io.ReadAll(resp.Body)
|
||||
_ = resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
t.Fatalf("admin login: %d %s", resp.StatusCode, raw)
|
||||
}
|
||||
return client
|
||||
}
|
||||
|
||||
func disableEP(t *testing.T, client *http.Client, base, id string) {
|
||||
t.Helper()
|
||||
req, err := http.NewRequest(http.MethodPatch, base+"/api/admin/endpoints/"+id,
|
||||
strings.NewReader(`{"enabled":false}`))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("X-Nixmsg-Request", "1")
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
raw, _ := io.ReadAll(resp.Body)
|
||||
_ = resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
t.Fatalf("disable %s: %d %s", id, resp.StatusCode, raw)
|
||||
}
|
||||
}
|
||||
|
||||
func resetLoginPassword(t *testing.T, client *http.Client, base, id, password string) {
|
||||
t.Helper()
|
||||
body := `{"login_password":"` + password + `"}`
|
||||
req, err := http.NewRequest(http.MethodPost, base+"/api/admin/endpoints/"+id+"/reset-login-password",
|
||||
strings.NewReader(body))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("X-Nixmsg-Request", "1")
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
raw, _ := io.ReadAll(resp.Body)
|
||||
_ = resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
t.Fatalf("reset password %s: %d %s", id, resp.StatusCode, raw)
|
||||
}
|
||||
}
|
||||
+47
-6
@@ -71,6 +71,10 @@ func runServe(ctx context.Context, cfg config.Config) error {
|
||||
loginLocks := auth.NewLoginLocks()
|
||||
memConns := message.NewMemoryConns()
|
||||
msgLim := message.LimitsFromFullConfig(cfg)
|
||||
metricsReg := metrics.New()
|
||||
db.Queue.OnBatchCommit = func(d time.Duration) {
|
||||
metricsReg.WriteCommitSeconds.Observe(d.Seconds())
|
||||
}
|
||||
|
||||
login := broker.NewLogin(broker.LoginOptions{
|
||||
DB: db,
|
||||
@@ -84,11 +88,13 @@ func runServe(ctx context.Context, cfg config.Config) error {
|
||||
msgApp := message.New(db, msgLim, hashPool,
|
||||
message.WithLocks(loginLocks),
|
||||
message.WithConnRegistry(memConns),
|
||||
message.WithMetrics(metricsReg),
|
||||
)
|
||||
uplink := &appUplink{
|
||||
msg: msgApp,
|
||||
conns: memConns,
|
||||
log: slog.Default(),
|
||||
msg: msgApp,
|
||||
conns: memConns,
|
||||
log: slog.Default(),
|
||||
metrics: metricsReg,
|
||||
}
|
||||
sess := broker.NewSession(broker.SessionOptions{
|
||||
Login: login,
|
||||
@@ -109,6 +115,7 @@ func runServe(ctx context.Context, cfg config.Config) error {
|
||||
Authenticator: login,
|
||||
Uplink: sess,
|
||||
Logger: slog.Default(),
|
||||
Metrics: metricsReg,
|
||||
OnPublishDropped: func(dropCtx context.Context, endpointID string, connID port.ConnID, payload []byte) {
|
||||
if dropErr := msgApp.OnPublishDropped(dropCtx, endpointID, connID, payload); dropErr != nil {
|
||||
slog.Error("on publish dropped", "endpoint", endpointID, "err", dropErr)
|
||||
@@ -125,6 +132,7 @@ func runServe(ctx context.Context, cfg config.Config) error {
|
||||
message.WithLocks(loginLocks),
|
||||
message.WithConnRegistry(memConns),
|
||||
message.WithDownlink(brk),
|
||||
message.WithMetrics(metricsReg),
|
||||
)
|
||||
uplink.msg = msgApp
|
||||
uplink.down = brk
|
||||
@@ -146,6 +154,7 @@ func runServe(ctx context.Context, cfg config.Config) error {
|
||||
MaxScheduleSeconds: int64(cfg.Limits.MaxScheduleSeconds),
|
||||
Logger: slog.Default(),
|
||||
ConnControl: brk,
|
||||
Downlink: brk,
|
||||
ClientIP: func(r *http.Request) string {
|
||||
return httpx.ClientIP(r, trustedNets)
|
||||
},
|
||||
@@ -176,6 +185,7 @@ func runServe(ctx context.Context, cfg config.Config) error {
|
||||
Groups: groupApp,
|
||||
Config: cfg,
|
||||
Version: Version,
|
||||
// Kick:只断开,令牌不变,SDK 重连(PRD 踢下线)。
|
||||
KickEndpoint: func(kickCtx context.Context, endpointID string) (bool, error) {
|
||||
if _, found := brk.ConnInfoOf(endpointID); !found {
|
||||
return false, nil
|
||||
@@ -185,9 +195,36 @@ func runServe(ctx context.Context, cfg config.Config) error {
|
||||
}
|
||||
return true, nil
|
||||
},
|
||||
// 停用/删除/重置:先 fatal 再断开(DEVELOPMENT 6.8)。
|
||||
DisableKick: func(kickCtx context.Context, endpointID string) (bool, error) {
|
||||
if _, found := brk.ConnInfoOf(endpointID); !found {
|
||||
return false, nil
|
||||
}
|
||||
if disableErr := sess.Disable(kickCtx, endpointID); disableErr != nil {
|
||||
return false, disableErr
|
||||
}
|
||||
return true, nil
|
||||
},
|
||||
DeleteKick: func(kickCtx context.Context, endpointID string) (bool, error) {
|
||||
if _, found := brk.ConnInfoOf(endpointID); !found {
|
||||
return false, nil
|
||||
}
|
||||
if deleteErr := sess.Deleted(kickCtx, endpointID); deleteErr != nil {
|
||||
return false, deleteErr
|
||||
}
|
||||
return true, nil
|
||||
},
|
||||
PasswordResetKick: func(kickCtx context.Context, endpointID string) (bool, error) {
|
||||
if _, found := brk.ConnInfoOf(endpointID); !found {
|
||||
return false, nil
|
||||
}
|
||||
if resetErr := sess.ResetPassword(kickCtx, endpointID); resetErr != nil {
|
||||
return false, resetErr
|
||||
}
|
||||
return true, nil
|
||||
},
|
||||
})
|
||||
|
||||
metricsReg := metrics.New()
|
||||
buildHandlers := func(proxies *listener.ProxySet) listener.Handlers {
|
||||
return listener.Handlers{
|
||||
MQTT: brk.WSHandler(proxies),
|
||||
@@ -269,7 +306,7 @@ func runServe(ctx context.Context, cfg config.Config) error {
|
||||
|
||||
loopCtx, loopCancel := context.WithCancel(ctx)
|
||||
defer loopCancel()
|
||||
go messageLoops(loopCtx, msgApp, memConns)
|
||||
go messageLoops(loopCtx, msgApp, memConns, db, hashPool, metricsReg)
|
||||
|
||||
<-ctx.Done()
|
||||
loopCancel()
|
||||
@@ -282,7 +319,7 @@ func runServe(ctx context.Context, cfg config.Config) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func messageLoops(ctx context.Context, msgApp *message.App, conns *message.MemoryConns) {
|
||||
func messageLoops(ctx context.Context, msgApp *message.App, conns *message.MemoryConns, db *store.DB, hashPool auth.HashPool, met *metrics.Registry) {
|
||||
t := time.NewTicker(time.Second)
|
||||
defer t.Stop()
|
||||
for {
|
||||
@@ -302,6 +339,10 @@ func messageLoops(ctx context.Context, msgApp *message.App, conns *message.Memor
|
||||
if err := msgApp.CleanupOnce(ctx, nowMs); err != nil {
|
||||
slog.Error("cleanup once", "err", err)
|
||||
}
|
||||
if err := metrics.SampleStoreGauges(ctx, met, db.Read); err != nil {
|
||||
slog.Debug("sample store gauges", "err", err)
|
||||
}
|
||||
metrics.SampleQueues(met, db.Queue.Len(), hashPool.QueueLen())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"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/metrics"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
)
|
||||
|
||||
@@ -23,6 +24,7 @@ type appUplink struct {
|
||||
conns *message.MemoryConns
|
||||
down port.Downlink
|
||||
log *slog.Logger
|
||||
metrics *metrics.Registry
|
||||
}
|
||||
|
||||
func (u *appUplink) OnSessionEstablished(ctx context.Context, conn port.ConnInfo) error {
|
||||
@@ -255,6 +257,9 @@ func (u *appUplink) replyErr(ctx context.Context, conn port.ConnInfo, rid, code,
|
||||
if rid == "" {
|
||||
rid = "0"
|
||||
}
|
||||
if u.metrics != nil && code != "" {
|
||||
u.metrics.ErrorsTotal.WithLabelValues(code).Inc()
|
||||
}
|
||||
resp := protocol.Resp{
|
||||
V: protocol.Version,
|
||||
Type: protocol.TypeResp,
|
||||
|
||||
@@ -1093,6 +1093,24 @@
|
||||
- 备选方案:broker 层对 Inline 发布做无锁队列。
|
||||
- 影响:`group_event` 可能略晚于 `resp` 到达;业务结果仍以 `resp` 为准。
|
||||
|
||||
### fix-issue-1
|
||||
|
||||
1. **管理员 IP 锁定不再阻断已认证会话**
|
||||
- 原条款:PRD D18 / F02(密码锁只拦密码登录,不拦已有会话令牌);DEVELOPMENT 第 5/8 节(管理员登录锁定、错误令牌按 IP 计入锁定);issue #1。
|
||||
- 实际做法:去掉 `internal/admin/auth.go` 的 `auth()` 鉴权前 `Check(LockAdminIP)`;登录入口仍 `Check`/`Fail`,错误或停用 API 令牌仍经 `authFail` 计入锁定。有效 Cookie 与合法 Bearer 在锁定期可继续调管理接口。
|
||||
- 原因:先前把「防暴力登录」扩成「封整个管理面」,同 NAT 下刷错误 Bearer 即可锁死已登录管理员,与端侧 nst_ 重连语义不一致。
|
||||
- 备选方案:锁定期对 Cookie 与令牌也拒绝(否决,违背 D18 对齐)。
|
||||
- 影响:仅管理后台鉴权中间件;端侧登录锁定未改。
|
||||
|
||||
### fix-issue-6
|
||||
|
||||
1. **解散群时 scheduled 消息级回执 state 改为 rejected**
|
||||
- 原条款:DEVELOPMENT 6.4 消息级作废写 `endpoint_id` 空、`state=rejected`;7.6 解散群将 `scheduled` 消息改为 `completed`/`group_dissolved` 并写消息级回执。I4 旧实现把回执 state 误写成消息状态 `completed`。
|
||||
- 实际做法:`internal/app/group/void.go` 的 `voidGroupAllTx` 插入回执时改用 `rejected`(与 I5.2 / identity lifecycle 一致);消息行仍为 `completed`。
|
||||
- 原因:`completed` 不在回执枚举(accepted|recalled|expired|dropped|rejected)内,会误导 SDK/后台。
|
||||
- 备选方案:沿用 `completed`(违反协议)。
|
||||
- 影响:仅修正解散路径回执字段;不改 emit / PublishDown。
|
||||
|
||||
### fix-issue-2
|
||||
|
||||
1. **自助注册接入 trusted_proxies 客户端 IP**
|
||||
@@ -1101,3 +1119,55 @@
|
||||
- 原因:L-WIRE 已挂注册 Handler,管理与 WS 已接 `trusted_proxies`,唯独注册漏接,反向代理后会把安全码锁定计到代理 IP。
|
||||
- 备选方案:在 listener 层统一改写 `RemoteAddr` 后再交给注册 Handler。
|
||||
- 影响:经受信代理开放注册时,输错安全码按真实客户端 IP 锁定。
|
||||
|
||||
### fix-issue-4
|
||||
|
||||
1. **接线补齐 Downlink 与停用/删除/重置密码 fatal**
|
||||
- 原条款:DEVELOPMENT 6.8 / 7.6:停用、删除、重置密码先发 `fatal` 再断开;已推送作废投递尽力发 `revoked`。
|
||||
- 实际做法:`serve` 给 `identity.New` 注入 `Downlink: brk`(作废后 `publishRevokes`);`DisableKick`/`DeleteKick`/`PasswordResetKick` 分别接到 `Session.Disable`/`Deleted`/`ResetPassword`;`KickEndpoint` 仍只 `Kick`。Identity 在未接 Kick 钩子时仍可用 `ConnControl` 异步断开兜底。
|
||||
- 原因:原先 Downlink 未注入导致 revoked 丢失;管理路径只 `Kick`/`Disconnect` 不发 fatal。
|
||||
- 备选方案:仅在 identity 内 `PublishDown(fatal)` 再断开;联调中该路径不如 Session.fatalKick 稳,故生产致命踢线统一走 Session。
|
||||
- 影响:管理「踢下线」语义不变;SDK 可按 fatal 停止重连;接收方能收到已推送消息的 revoked。
|
||||
|
||||
### fix-issue-5
|
||||
|
||||
1. **指标在真实事件点打点,不新造名字**
|
||||
- 原条款:PRD F22 / DEVELOPMENT 4.3 / issue #5;DEVIATIONS P4 已定名但从未接线。
|
||||
- 实际做法:`nixmsg_connections{transport}` 在 broker `OnSessionEstablished`/`OnDisconnect` 末尾 Inc/Dec;`endpoints`/`deliveries_pending`/`messages_scheduled` 与写队列、哈希排队在 `messageLoops` 每秒按库/队列真实长度采样;`dispatch_to_push`/`ack` 直方图在成功推送与确认路径 Observe;写批提交耗时经 `store.Queue.OnBatchCommit`;`errors_total` 仅在上行 `replyErr` 时按错误码递增。门禁不变。
|
||||
- 原因:空指标等于监控未交付;采样避免在每条写路径上改大段分发逻辑,并减小与 #3/#6 的合并面。
|
||||
- 备选方案:全部改为纯事件加减(pending 等需在每处状态迁移维护计数)。
|
||||
- 影响:仪表盘按既有名字即可看在线连接与待投递;无对应事件时计数保持 0,不做假数。
|
||||
|
||||
### 死锁未修(issue #3)
|
||||
|
||||
issue #3 未关闭,`feat/fix-3-downlink-deadlock` 未合入 `main`。下面是核对过的调用链、三次尝试和仍留在 `main` 上的绕过。不改产品行为。
|
||||
|
||||
1. **现象**
|
||||
- 清掉测试客户端 dial deadline 后,`cmd/nixmsg/uplink_integration_test.go` 的 `TestUplinkDMOfflineGroupRecall` 在 `group.create`(`rid=g1`)稳定超时,`resp` 回不去。
|
||||
- 行号以本次合入后的 `main` 为准。
|
||||
|
||||
2. **调用链**
|
||||
- 每端一条上行队列。`internal/broker/queue.go` 的 `loop`(约 41 行)同步调用 `HandleUplink`。
|
||||
- `cmd/nixmsg/uplink.go` 的 `HandleUplink`(约 64 行)先 `dispatch`,再在同一调用栈里 `replyOK` → `publishResp`(约 273 行)用 QoS 1 调 `PublishDown`。`group.create` / `group.add` 走 `groups.Create` / `Add`,在返回 `resp` 之前就 `emit`(`internal/app/group/app.go` 约 168、223 行)。
|
||||
- `Broker.PublishDown`(`internal/broker/broker.go` 约 202 行)进入 `server.Publish`(约 234 行)。`New` 设了 `InlineClient: true`(约 148 行,DEVELOPMENT 要求保持)。Inline 发布走到 `InjectPacket` → `OnPublish`。
|
||||
- `internal/broker/hooks.go` 的 `OnPublish`(约 114 行)对 `cl.Net.Inline` 必须直接放行,否则 `PublishDown` 送不到订阅者(见本文更早的 InlineClient 偏差)。
|
||||
- 群操作因此在同一次上行调用栈里,再向本连接 `PublishDown` `group_event`。`InjectPacket`(`NextPacketID` / 写路径)与读循环随后写 PUBACK 抢同一把 Client 锁,两边互等,`resp` 出不去。
|
||||
- `presence.notify`(`internal/app/presence/app.go` 约 317 行)仍在业务调用栈里同步 `PublishDown`(约 341 行),不在 20ms 绕过的覆盖范围内。
|
||||
|
||||
3. **已尝试**
|
||||
- 尝试 1:`emit` 改成立刻起 goroutine 做 `PublishDown`。仍死锁,因为 `group_event` 与同连接上的 `resp` 一起抢注入。
|
||||
- 尝试 2(已在 `main`,来自 `479a08e` 的 Q accept-rest 第 4 条):`emit` 起 goroutine 后 `time.Sleep(20ms)` 再下发(`internal/app/group/app.go` 约 672–685 行),让 `resp` 先出去。当时 `task check` 通过。这是时间差绕过,不是根因修复;presence 以及其他同步 `PublishDown` 仍可能卡。
|
||||
- 尝试 3(负责人叫停,未合入、未验证):工作树 `e:\code\NixMsg-wt\fix3`,分支 `feat/fix-3-downlink-deadlock` 停在 `479a08e`,与合入前的 `origin/main` 相同,**没有可合的提交**。未提交改动在 broker 层:
|
||||
- `hooks.go` `OnPublish`:客户端 QoS≥1 先在读循环里 `WritePacket(PUBACK)`,再把包降成 QoS 0 并 `Ignore`,然后入队,返回 `nil`(不再用 `CodeSuccessIgnore` 让 mochi 事后写 PUBACK)。注释写明:若先入队,worker 的 `PublishDown` → `InjectPacket` → `NextPacketID` 会与随后的 `WritePacket(PUBACK)` 争 Client 锁。
|
||||
- `queue.go` `loop`:`HandleUplink` 前后调用 `beginUplink` / `endUplink`。
|
||||
- `broker.go`:该端 `depth>0` 时,对本端的 `PublishDown` 只推进延后队列,handler 返回后由上行 worker 再 `server.Publish`;其他端仍同步下发。`InlineClient` 放行未改。
|
||||
- `group/app.go` `emit` 改回同步 `PublishDown`,去掉 20ms sleep。
|
||||
- 同目录 `docs/DEVIATIONS.md` 有一段未提交的 `### fix-issue-3` 草稿。
|
||||
- 本会话没有跑这套未提交代码的 `task check`,不把它们合进 `main`。工作树保留,给工程师看。
|
||||
|
||||
4. **仍在 main 上的做法**
|
||||
- 继续用尝试 2 的 20ms 绕过。群事件可能略晚于 `resp`;业务结果仍以 `resp` 为准。
|
||||
|
||||
5. **建议的正确方向**
|
||||
- 在 broker 把对本连接的下行 `InjectPacket` 与上行 worker 解耦:上行读循环先写完 PUBACK,处理 `HandleUplink` 期间不要同步向本连接注入;handler 返回后再发 `resp` 和 `group_event`。不要靠固定 `Sleep`。`InlineClient: true` 保持,`OnPublish` 对 InlineClient 继续放行。
|
||||
- 覆盖 presence 等其他同步 `PublishDown`,而不只包一层 `emit`。
|
||||
|
||||
@@ -280,6 +280,82 @@ func TestLoginLock(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestAdminLockDoesNotBlockAuthedSession:密码失败触发 IP 锁后,
|
||||
// 已有 Cookie 会话与合法 API 令牌仍可调管理接口;未认证密码登录仍被拒。
|
||||
func TestAdminLockDoesNotBlockAuthedSession(t *testing.T) {
|
||||
_, srv, cookieClient, _ := setup(t)
|
||||
base := srv.URL
|
||||
|
||||
login(t, cookieClient, base)
|
||||
|
||||
res := postJSON(t, cookieClient, base+"/api/admin/tokens",
|
||||
`{"name":"ops-lock"}`,
|
||||
map[string]string{"X-Nixmsg-Request": "1"})
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("create token: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
var created struct {
|
||||
Token string `json:"token"`
|
||||
}
|
||||
if err := json.Unmarshal(env.Data, &created); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
for i := 0; i < 10; i++ {
|
||||
bad := &http.Client{}
|
||||
res = postJSON(t, bad, base+"/api/admin/login",
|
||||
`{"username":"admin","password":"wrong-password!!"}`, nil)
|
||||
env = decodeEnv(t, res)
|
||||
if i < 9 {
|
||||
if res.StatusCode != 401 {
|
||||
t.Fatalf("fail %d: want 401 got %d %+v", i, res.StatusCode, env)
|
||||
}
|
||||
continue
|
||||
}
|
||||
if res.StatusCode != 429 || env.Error == nil || env.Error.Code != "rate_limited" {
|
||||
t.Fatalf("10th fail want 429 rate_limited got %d %+v", res.StatusCode, env)
|
||||
}
|
||||
}
|
||||
|
||||
res = doReq(t, cookieClient, http.MethodGet, base+"/api/admin/me", "", nil)
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("cookie me after lock: want 200 got %d %+v", res.StatusCode, env)
|
||||
}
|
||||
var me map[string]any
|
||||
_ = json.Unmarshal(env.Data, &me)
|
||||
if me["auth"] != "cookie" {
|
||||
t.Fatalf("cookie me auth=%v", me["auth"])
|
||||
}
|
||||
|
||||
tokClient := &http.Client{}
|
||||
hdr := map[string]string{"Authorization": "Bearer " + created.Token}
|
||||
res = doReq(t, tokClient, http.MethodGet, base+"/api/admin/me", "", hdr)
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("token me after lock: want 200 got %d %+v", res.StatusCode, env)
|
||||
}
|
||||
_ = json.Unmarshal(env.Data, &me)
|
||||
if me["auth"] != "token" {
|
||||
t.Fatalf("token me auth=%v", me["auth"])
|
||||
}
|
||||
|
||||
res = doReq(t, tokClient, http.MethodGet, base+"/api/admin/overview", "", hdr)
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("token overview after lock: want 200 got %d %+v", res.StatusCode, env)
|
||||
}
|
||||
|
||||
anon := &http.Client{}
|
||||
res = postJSON(t, anon, base+"/api/admin/login",
|
||||
`{"username":"admin","password":"`+testPassword+`"}`, nil)
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != 429 || env.Error == nil || env.Error.Code != "rate_limited" {
|
||||
t.Fatalf("password login while locked want 429 got %d %+v", res.StatusCode, env)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBadAPITokenCountsTowardLock(t *testing.T) {
|
||||
_, srv, _, _ := setup(t)
|
||||
base := srv.URL
|
||||
|
||||
@@ -43,12 +43,8 @@ func (h *Handler) auth(next http.HandlerFunc) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
ip := httpx.ClientIP(r, h.trusted)
|
||||
|
||||
if locked, retry := h.locks.Check(auth.LockKey{Kind: auth.LockAdminIP, IP: ip}); locked {
|
||||
w.Header().Set("Retry-After", formatRetryAfter(retry))
|
||||
httpx.WriteError(w, http.StatusTooManyRequests, "rate_limited", "登录已锁定,请稍后再试")
|
||||
return
|
||||
}
|
||||
|
||||
// 锁定只拦密码登录(login.go)与错误令牌试错累计;
|
||||
// 已认证的 Cookie / 合法 API 令牌在锁定期仍可用(对齐 PRD D18)。
|
||||
p, errCode, errMsg, status := h.authenticate(r, ip)
|
||||
if status != 0 {
|
||||
if status == http.StatusTooManyRequests {
|
||||
|
||||
@@ -102,6 +102,33 @@ func (h *Handler) kickEndpoint(ctx context.Context, id string) (bool, error) {
|
||||
return h.kick(ctx, id)
|
||||
}
|
||||
|
||||
func (h *Handler) passwordResetKick(ctx context.Context, id string) (bool, error) {
|
||||
if h.resetKick != nil {
|
||||
return h.resetKick(ctx, id)
|
||||
}
|
||||
return h.kickEndpoint(ctx, id)
|
||||
}
|
||||
|
||||
func (h *Handler) afterDisableKick(ctx context.Context, id string) {
|
||||
if h.disableKick != nil {
|
||||
_, _ = h.disableKick(ctx, id)
|
||||
return
|
||||
}
|
||||
if h.identity == nil {
|
||||
_, _ = h.kickEndpoint(ctx, id)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) afterDeleteKick(ctx context.Context, id string) {
|
||||
if h.deleteKick != nil {
|
||||
_, _ = h.deleteKick(ctx, id)
|
||||
return
|
||||
}
|
||||
if h.identity == nil {
|
||||
_, _ = h.kickEndpoint(ctx, id)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) handleEndpointList(w http.ResponseWriter, r *http.Request) {
|
||||
q := r.URL.Query()
|
||||
limit := defaultListLimit
|
||||
@@ -350,7 +377,7 @@ func (h *Handler) handleEndpointPatch(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
if !*req.Enabled {
|
||||
_, _ = h.kickEndpoint(r.Context(), id)
|
||||
h.afterDisableKick(r.Context(), id)
|
||||
}
|
||||
} else if req.Enabled != nil && !*req.Enabled && wasEnabled {
|
||||
_, _ = h.kickEndpoint(r.Context(), id)
|
||||
@@ -381,7 +408,7 @@ func (h *Handler) handleEndpointDelete(w http.ResponseWriter, r *http.Request) {
|
||||
httpx.WriteError(w, http.StatusNotFound, "not_found", "端不存在")
|
||||
return
|
||||
}
|
||||
_, _ = h.kickEndpoint(r.Context(), id)
|
||||
h.afterDeleteKick(r.Context(), id)
|
||||
h.audit(actorString(p), "endpoint_delete", id, "ok", ip)
|
||||
httpx.WriteOK(w, map[string]any{})
|
||||
}
|
||||
@@ -416,14 +443,14 @@ func (h *Handler) handleEndpointBatch(w http.ResponseWriter, r *http.Request) {
|
||||
case "disable":
|
||||
found, opErr = h.setEndpointEnabled(r.Context(), id, false)
|
||||
if found && opErr == nil {
|
||||
_, _ = h.kickEndpoint(r.Context(), id)
|
||||
h.afterDisableKick(r.Context(), id)
|
||||
}
|
||||
case "enable":
|
||||
found, opErr = h.setEndpointEnabled(r.Context(), id, true)
|
||||
case "delete":
|
||||
found, opErr = h.deleteEndpointBasic(r.Context(), id)
|
||||
if found && opErr == nil {
|
||||
_, _ = h.kickEndpoint(r.Context(), id)
|
||||
h.afterDeleteKick(r.Context(), id)
|
||||
}
|
||||
}
|
||||
if opErr != nil {
|
||||
@@ -512,7 +539,7 @@ func (h *Handler) handleEndpointResetLoginPassword(w http.ResponseWriter, r *htt
|
||||
httpx.WriteError(w, http.StatusNotFound, "not_found", "端不存在")
|
||||
return
|
||||
}
|
||||
_, _ = h.kickEndpoint(r.Context(), id)
|
||||
_, _ = h.passwordResetKick(r.Context(), id)
|
||||
h.audit(actorString(p), "endpoint_reset_login_password", id, "ok", ip)
|
||||
httpx.WriteOK(w, map[string]any{loginPasswordOnceKey: pw})
|
||||
}
|
||||
|
||||
+42
-30
@@ -41,6 +41,12 @@ type Deps struct {
|
||||
SecureCookies bool
|
||||
// KickEndpoint 踢下线钩子(只断开连接);nil 时踢线为 no-op。
|
||||
KickEndpoint EndpointKickFunc
|
||||
// PasswordResetKick 重置登录密码后踢线(应发 fatal);nil 时回退 KickEndpoint。
|
||||
PasswordResetKick EndpointKickFunc
|
||||
// DisableKick 停用后踢线(应发 fatal(disabled));nil 且已注入 Identity 时不再 Kick。
|
||||
DisableKick EndpointKickFunc
|
||||
// DeleteKick 删除后踢线(应发 fatal(deleted));nil 且已注入 Identity 时不再 Kick。
|
||||
DeleteKick EndpointKickFunc
|
||||
// Identity 端停用/启用/删除级联(I5);nil 时回退为仅改 enabled/删行。
|
||||
Identity identity.Service
|
||||
|
||||
@@ -54,20 +60,23 @@ type Deps struct {
|
||||
|
||||
// Handler 是可挂载的管理接口(路由前缀 /api/admin/)。
|
||||
type Handler struct {
|
||||
db *store.DB
|
||||
hash auth.HashPool
|
||||
tokens auth.APITokens
|
||||
locks auth.LoginLocks
|
||||
log *slog.Logger
|
||||
trusted []*net.IPNet
|
||||
ttl time.Duration
|
||||
forceSec bool
|
||||
kick EndpointKickFunc
|
||||
identity identity.Service
|
||||
groups group.Service
|
||||
cfg config.Config
|
||||
version string
|
||||
startedAt time.Time
|
||||
db *store.DB
|
||||
hash auth.HashPool
|
||||
tokens auth.APITokens
|
||||
locks auth.LoginLocks
|
||||
log *slog.Logger
|
||||
trusted []*net.IPNet
|
||||
ttl time.Duration
|
||||
forceSec bool
|
||||
kick EndpointKickFunc
|
||||
resetKick EndpointKickFunc
|
||||
disableKick EndpointKickFunc
|
||||
deleteKick EndpointKickFunc
|
||||
identity identity.Service
|
||||
groups group.Service
|
||||
cfg config.Config
|
||||
version string
|
||||
startedAt time.Time
|
||||
|
||||
mux *http.ServeMux
|
||||
|
||||
@@ -96,22 +105,25 @@ func New(d Deps) *Handler {
|
||||
ver = "dev"
|
||||
}
|
||||
h := &Handler{
|
||||
db: d.DB,
|
||||
hash: d.Hash,
|
||||
tokens: d.Tokens,
|
||||
locks: d.Locks,
|
||||
log: d.Logger,
|
||||
trusted: d.TrustedProxies,
|
||||
ttl: ttl,
|
||||
forceSec: d.SecureCookies,
|
||||
kick: d.KickEndpoint,
|
||||
identity: d.Identity,
|
||||
groups: d.Groups,
|
||||
cfg: cfg,
|
||||
version: ver,
|
||||
startedAt: time.Now(),
|
||||
mux: http.NewServeMux(),
|
||||
lastUsed: make(map[string]time.Time),
|
||||
db: d.DB,
|
||||
hash: d.Hash,
|
||||
tokens: d.Tokens,
|
||||
locks: d.Locks,
|
||||
log: d.Logger,
|
||||
trusted: d.TrustedProxies,
|
||||
ttl: ttl,
|
||||
forceSec: d.SecureCookies,
|
||||
kick: d.KickEndpoint,
|
||||
resetKick: d.PasswordResetKick,
|
||||
disableKick: d.DisableKick,
|
||||
deleteKick: d.DeleteKick,
|
||||
identity: d.Identity,
|
||||
groups: d.Groups,
|
||||
cfg: cfg,
|
||||
version: ver,
|
||||
startedAt: time.Now(),
|
||||
mux: http.NewServeMux(),
|
||||
lastUsed: make(map[string]time.Time),
|
||||
}
|
||||
h.routes()
|
||||
return h
|
||||
|
||||
@@ -296,6 +296,16 @@ WHERE m.id='gm1' AND d.endpoint_id='bob'`).Scan(&reason)
|
||||
if err != nil || state != "completed" || mreason != "group_dissolved" {
|
||||
t.Fatalf("state=%s reason=%s err=%v", state, mreason, err)
|
||||
}
|
||||
// 消息级回执 state 须为协议枚举 rejected,不得写成消息状态 completed(issue #6 / DEVELOPMENT 6.4)
|
||||
var rState, rReason, rEndpoint string
|
||||
err = db.Read.QueryRow(`
|
||||
SELECT state, reason, endpoint_id FROM receipts WHERE sender_id='alice' AND msg_id='gm2'`).Scan(&rState, &rReason, &rEndpoint)
|
||||
if err != nil {
|
||||
t.Fatalf("receipt for dissolved scheduled: %v", err)
|
||||
}
|
||||
if rState != "rejected" || rReason != "group_dissolved" || rEndpoint != "" {
|
||||
t.Fatalf("receipt state=%q reason=%q endpoint=%q want rejected/group_dissolved/empty", rState, rReason, rEndpoint)
|
||||
}
|
||||
// 同编号新建群
|
||||
created2, err := gApp.Create(ctx, "alice", &protocol.GroupCreate{
|
||||
V: protocol.Version, Type: protocol.TypeGroupCreate, RID: "6",
|
||||
|
||||
@@ -139,10 +139,11 @@ UPDATE messages SET state = 'completed', reason = ? WHERE seq = ? AND state = 's
|
||||
return execErr
|
||||
}
|
||||
if r.receipt != 0 {
|
||||
// 消息级作废回执:endpoint_id 空,state=rejected(DEVELOPMENT 6.4);消息行仍为 completed
|
||||
if _, execErr := tx.Exec(`
|
||||
INSERT INTO receipts(sender_id, msg_id, endpoint_id, state, reason, created_at, acked)
|
||||
VALUES(?,?,?,?,?,?,0)`,
|
||||
r.senderID, r.msgID, "", "completed", reasonGroupDissolved, nowMs); execErr != nil {
|
||||
r.senderID, r.msgID, "", "rejected", reasonGroupDissolved, nowMs); execErr != nil {
|
||||
return execErr
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
@@ -22,6 +23,9 @@ const (
|
||||
eventLeft = "left"
|
||||
eventMemberRemoved = "member_removed"
|
||||
eventDissolved = "dissolved"
|
||||
|
||||
// kickFlushDelay 给接线方 Session.Disable/Deleted 留出发 fatal 的窗口。
|
||||
kickFlushDelay = 20 * time.Millisecond
|
||||
)
|
||||
|
||||
type revokeItem struct {
|
||||
@@ -124,8 +128,13 @@ WHERE id = ?`, endpointID); e != nil {
|
||||
|
||||
a.publishRevokes(ctx, revokes)
|
||||
a.publishGroupEvents(ctx, notifies)
|
||||
// fatal+断开由 admin DisableKick/DeleteKick(Session.Disable/Deleted)完成。
|
||||
// 未接 Kick 钩子的单元测试仍可用 ConnControl 兜底断开。
|
||||
if a.connCtrl != nil {
|
||||
_ = a.connCtrl.Disconnect(ctx, endpointID, "", port.DisconnectFatal)
|
||||
go func() {
|
||||
time.Sleep(kickFlushDelay)
|
||||
_ = a.connCtrl.Disconnect(context.Background(), endpointID, "", port.DisconnectFatal)
|
||||
}()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ package identity_test
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/cookiejar"
|
||||
@@ -64,6 +65,82 @@ VALUES(?,?,?,?,0,0,1,?,?)`, id, id, "stub$login", nil, 1_700_000_000_000, 1_700_
|
||||
}
|
||||
}
|
||||
|
||||
func TestDisableEmitsRevokedForPushed(t *testing.T) {
|
||||
t.Parallel()
|
||||
dir := t.TempDir()
|
||||
db, err := store.Open(filepath.Join(dir, "data"), "FULL")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
fixed := time.UnixMilli(1_700_000_000_000)
|
||||
down := &message.RecordingDownlink{}
|
||||
ctrl := &port.StubConnControl{}
|
||||
idApp := identity.New(identity.Config{
|
||||
DB: db,
|
||||
Hash: auth.NewStubHashPool(),
|
||||
Locks: auth.NewStubLoginLocks(),
|
||||
Sessions: auth.NewSessionTokens(),
|
||||
MaxScheduleSeconds: int64(config.Default().Limits.MaxScheduleSeconds),
|
||||
Now: func() time.Time { return fixed },
|
||||
ConnControl: ctrl,
|
||||
Downlink: down,
|
||||
})
|
||||
ctx := context.Background()
|
||||
insertEPFull(t, db, "alice")
|
||||
insertEPFull(t, db, "bob")
|
||||
|
||||
err = db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
res, e := tx.Exec(`
|
||||
INSERT INTO messages(
|
||||
id, sender_id, dest_kind, dest_id, meta, content_type, body_enc,
|
||||
send_at, keep, ttl_seconds, receipt, state, reason, created_at)
|
||||
VALUES('pushed-1','alice','endpoint','bob','{}','text/plain','utf8',?,1,0,0,'dispatched','',?)`,
|
||||
fixed.UnixMilli(), fixed.UnixMilli())
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
seq, _ := res.LastInsertId()
|
||||
_, e = tx.Exec(`
|
||||
INSERT INTO deliveries(seq, endpoint_id, send_at, keep, state, reason, updated_at, pushed_at, pushed_conn)
|
||||
VALUES(?,?,?,1,'pending','',?,?,?)`,
|
||||
seq, "bob", fixed.UnixMilli(), fixed.UnixMilli(), fixed.UnixMilli(), "c-bob")
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
if err := idApp.Disable(ctx, "bob"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
if down.FilterType(protocol.TypeRevoked) != 1 {
|
||||
t.Fatalf("want 1 revoked, got snapshots=%v", down.Snapshots())
|
||||
}
|
||||
p := down.Snapshots()[0]
|
||||
var head struct {
|
||||
Type string `json:"type"`
|
||||
Reason string `json:"reason"`
|
||||
ID string `json:"id"`
|
||||
}
|
||||
_ = json.Unmarshal(p.Payload, &head)
|
||||
if head.Type != protocol.TypeRevoked || head.Reason != "endpoint_disabled" || head.ID != "pushed-1" {
|
||||
t.Fatalf("revoked=%+v", head)
|
||||
}
|
||||
if p.EndpointID != "bob" || p.QoS != 1 {
|
||||
t.Fatalf("publish=%+v", p)
|
||||
}
|
||||
deadline := time.Now().Add(2 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
if len(ctrl.Calls) == 1 && ctrl.Calls[0] == "bob" {
|
||||
return
|
||||
}
|
||||
time.Sleep(5 * time.Millisecond)
|
||||
}
|
||||
t.Fatalf("disconnect calls=%v", ctrl.Calls)
|
||||
}
|
||||
|
||||
func TestF01DisableVoidsScheduledAndRejectsNew(t *testing.T) {
|
||||
t.Parallel()
|
||||
idApp, msgApp, db := openLifecycle(t)
|
||||
|
||||
@@ -371,7 +371,7 @@ func TestRegisterTrustedProxyClientIPLock(t *testing.T) {
|
||||
locks.Clear(auth.LockKey{Kind: auth.LockRegisterIP, IP: "198.51.100.7"})
|
||||
locks.Clear(auth.LockKey{Kind: auth.LockRegisterIP, IP: "203.0.113.50"})
|
||||
for i := 0; i < 10; i++ {
|
||||
code, resp := post("203.0.113.50:4433", "198.51.100.7", wrong)
|
||||
code, resp = post("203.0.113.50:4433", "198.51.100.7", wrong)
|
||||
if code != http.StatusForbidden || resp.Error == nil || resp.Error.Code != protocol.CodeRegistrationCodeInvalid {
|
||||
t.Fatalf("untrusted fail #%d: status=%d resp=%+v", i+1, code, resp)
|
||||
}
|
||||
|
||||
@@ -17,6 +17,8 @@ func (a *App) Ack(ctx context.Context, endpointID string, req *protocol.Ack) (Ac
|
||||
nowMs := a.now().UnixMilli()
|
||||
var out AckResult
|
||||
var seq int64
|
||||
var ackLatencySec float64
|
||||
var observeAck bool
|
||||
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
err := tx.QueryRow(`SELECT seq FROM messages WHERE sender_id = ? AND id = ?`, req.From, req.ID).Scan(&seq)
|
||||
if err == sql.ErrNoRows {
|
||||
@@ -25,6 +27,10 @@ func (a *App) Ack(ctx context.Context, endpointID string, req *protocol.Ack) (Ac
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
var pushedAt sql.NullInt64
|
||||
_ = tx.QueryRow(`
|
||||
SELECT pushed_at FROM deliveries
|
||||
WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`, seq, endpointID).Scan(&pushedAt)
|
||||
res, err := tx.Exec(`
|
||||
UPDATE deliveries SET state = ?, reason = '', pushed_conn = NULL, updated_at = ?
|
||||
WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
|
||||
@@ -35,6 +41,10 @@ WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
|
||||
aff, _ := res.RowsAffected()
|
||||
if aff > 0 {
|
||||
out.Result = DeliveryAccepted
|
||||
if pushedAt.Valid && pushedAt.Int64 > 0 && nowMs >= pushedAt.Int64 {
|
||||
ackLatencySec = float64(nowMs-pushedAt.Int64) / 1000.0
|
||||
observeAck = true
|
||||
}
|
||||
if e := insertReceiptTx(tx, req.From, seq, endpointID, DeliveryAccepted, "", nowMs); e != nil {
|
||||
return e
|
||||
}
|
||||
@@ -55,6 +65,9 @@ SELECT state FROM deliveries WHERE seq = ? AND endpoint_id = ?`, seq, endpointID
|
||||
if err != nil {
|
||||
return out, err
|
||||
}
|
||||
if observeAck && a.met != nil {
|
||||
a.met.AckSeconds.Observe(ackLatencySec)
|
||||
}
|
||||
if out.Result == DeliveryAccepted {
|
||||
a.releaseLarge(seq, endpointID)
|
||||
}
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/config"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/metrics"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||
)
|
||||
@@ -86,6 +87,7 @@ type App struct {
|
||||
|
||||
down port.Downlink
|
||||
conns ConnRegistry
|
||||
met *metrics.Registry
|
||||
|
||||
mu sync.Mutex
|
||||
largeSem chan struct{}
|
||||
@@ -117,6 +119,11 @@ func WithConnRegistry(c ConnRegistry) Option {
|
||||
return func(a *App) { a.conns = c }
|
||||
}
|
||||
|
||||
// WithMetrics 注入 Prometheus 注册表(投递耗时直方图)。
|
||||
func WithMetrics(m *metrics.Registry) Option {
|
||||
return func(a *App) { a.met = m }
|
||||
}
|
||||
|
||||
// New 创建消息服务实现。
|
||||
func New(db *store.DB, lim Limits, hash auth.HashPool, opts ...Option) *App {
|
||||
if lim.RequestBurst <= 0 {
|
||||
|
||||
@@ -229,11 +229,20 @@ WHERE seq = ? AND endpoint_id = ? AND state = 'pending' AND pushed_conn IS NULL`
|
||||
a.releaseLarge(it.seq, endpointID)
|
||||
}
|
||||
a.scheduleRepush(endpointID, time.Second)
|
||||
} else {
|
||||
a.observeDispatchToPush(it.sendAt, nowMs)
|
||||
}
|
||||
}
|
||||
return a.pushReceipts(ctx, endpointID, connID, nowMs)
|
||||
}
|
||||
|
||||
func (a *App) observeDispatchToPush(sendAtMs, pushedAtMs int64) {
|
||||
if a.met == nil || pushedAtMs < sendAtMs {
|
||||
return
|
||||
}
|
||||
a.met.DispatchToPushSeconds.Observe(float64(pushedAtMs-sendAtMs) / 1000.0)
|
||||
}
|
||||
|
||||
func (a *App) rejectTooLarge(ctx context.Context, seq int64, endpointID, senderID string, nowMs int64) error {
|
||||
return a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
res, err := tx.Exec(`
|
||||
|
||||
@@ -12,6 +12,7 @@ import (
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/metrics"
|
||||
mqtt "github.com/mochi-mqtt/server/v2"
|
||||
"github.com/mochi-mqtt/server/v2/packets"
|
||||
)
|
||||
@@ -68,15 +69,18 @@ type Options struct {
|
||||
Logger *slog.Logger
|
||||
// OnPublishDropped 可选;nil 时仅打 debug 日志。
|
||||
OnPublishDropped PublishDroppedFunc
|
||||
// Metrics 可选;会话建立/断开时更新 nixmsg_connections。
|
||||
Metrics *metrics.Registry
|
||||
}
|
||||
|
||||
// Broker 内置 mochi,不自带监听端口。
|
||||
type Broker struct {
|
||||
server *mqtt.Server
|
||||
auth Authenticator
|
||||
uplink port.UplinkHandler
|
||||
log *slog.Logger
|
||||
onDrop PublishDroppedFunc
|
||||
server *mqtt.Server
|
||||
auth Authenticator
|
||||
uplink port.UplinkHandler
|
||||
log *slog.Logger
|
||||
onDrop PublishDroppedFunc
|
||||
metrics *metrics.Registry
|
||||
|
||||
hook *nixHook
|
||||
|
||||
@@ -105,6 +109,7 @@ type connState struct {
|
||||
handshook bool
|
||||
subscribedDown bool
|
||||
largeHeld int
|
||||
metricsCounted bool
|
||||
mu sync.Mutex
|
||||
|
||||
handshakeTimer *time.Timer
|
||||
@@ -151,6 +156,7 @@ func New(opts Options) (*Broker, error) {
|
||||
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),
|
||||
|
||||
@@ -190,6 +190,7 @@ func (h *nixHook) OnSessionEstablished(cl *mqtt.Client, _ packets.Packet) {
|
||||
MaxPacketSize: st.maxPacketSize,
|
||||
}
|
||||
_ = h.b.uplink.OnSessionEstablished(context.Background(), info)
|
||||
h.noteConnectionOpen(st)
|
||||
}
|
||||
|
||||
func (h *nixHook) OnDisconnect(cl *mqtt.Client, err error, _ bool) {
|
||||
@@ -229,9 +230,27 @@ func (h *nixHook) OnDisconnect(cl *mqtt.Client, err error, _ bool) {
|
||||
}
|
||||
if sess, ok := h.b.uplink.(*Session); ok {
|
||||
sess.HandleDisconnect(context.Background(), info, reason, isCurrent)
|
||||
h.noteConnectionClose(st)
|
||||
return
|
||||
}
|
||||
h.b.uplink.OnDisconnect(context.Background(), info, reason)
|
||||
h.noteConnectionClose(st)
|
||||
}
|
||||
|
||||
func (h *nixHook) noteConnectionOpen(st *connState) {
|
||||
if h.b.metrics == nil || st == nil || st.metricsCounted {
|
||||
return
|
||||
}
|
||||
h.b.metrics.Connections.WithLabelValues(string(st.transport)).Inc()
|
||||
st.metricsCounted = true
|
||||
}
|
||||
|
||||
func (h *nixHook) noteConnectionClose(st *connState) {
|
||||
if h.b.metrics == nil || st == nil || !st.metricsCounted {
|
||||
return
|
||||
}
|
||||
h.b.metrics.Connections.WithLabelValues(string(st.transport)).Dec()
|
||||
st.metricsCounted = false
|
||||
}
|
||||
|
||||
func (h *nixHook) OnQosComplete(cl *mqtt.Client, pk packets.Packet) {
|
||||
|
||||
@@ -0,0 +1,118 @@
|
||||
package broker
|
||||
|
||||
import (
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/metrics"
|
||||
dto "github.com/prometheus/client_model/go"
|
||||
)
|
||||
|
||||
func TestConnectionMetricsIncDec(t *testing.T) {
|
||||
reg := metrics.New()
|
||||
b, err := New(Options{Authenticator: AllowAuthenticator{}, Metrics: reg})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = b.Close() }()
|
||||
|
||||
if got := gaugeValue(t, reg, "nixmsg_connections", "tcp"); got != 0 {
|
||||
t.Fatalf("before connect tcp=%v", got)
|
||||
}
|
||||
|
||||
r, w := net.Pipe()
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
defer close(done)
|
||||
_ = b.AttachTCP(r)
|
||||
}()
|
||||
connectAndSubscribe(t, w, "ep-metrics", 0)
|
||||
|
||||
deadline := time.Now().Add(3 * time.Second)
|
||||
for {
|
||||
if _, ok := b.ConnInfoOf("ep-metrics"); ok {
|
||||
break
|
||||
}
|
||||
if time.Now().After(deadline) {
|
||||
t.Fatal("session not established")
|
||||
}
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
if got := gaugeValue(t, reg, "nixmsg_connections", "tcp"); got != 1 {
|
||||
t.Fatalf("after connect tcp=%v want 1", got)
|
||||
}
|
||||
|
||||
_ = w.Close()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("attach did not finish")
|
||||
}
|
||||
deadline = time.Now().Add(3 * time.Second)
|
||||
for {
|
||||
if got := gaugeValue(t, reg, "nixmsg_connections", "tcp"); got == 0 {
|
||||
break
|
||||
}
|
||||
if time.Now().After(deadline) {
|
||||
t.Fatalf("after disconnect tcp=%v want 0", gaugeValue(t, reg, "nixmsg_connections", "tcp"))
|
||||
}
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWSConnectionMetrics(t *testing.T) {
|
||||
reg := metrics.New()
|
||||
b, err := New(Options{Authenticator: AllowAuthenticator{}, Metrics: reg})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = b.Close() }()
|
||||
|
||||
mux := http.NewServeMux()
|
||||
mux.Handle("/mqtt", b.WSHandler(nil))
|
||||
srv := httptest.NewServer(mux)
|
||||
defer srv.Close()
|
||||
|
||||
// 仅验证 handler 暴露指标文本仍含初始标签;建连用 TCP 测即可。
|
||||
req := httptest.NewRequest(http.MethodGet, "/metrics", nil)
|
||||
rec := httptest.NewRecorder()
|
||||
reg.Handler().ServeHTTP(rec, req)
|
||||
body, _ := io.ReadAll(rec.Body)
|
||||
if !strings.Contains(string(body), `nixmsg_connections{transport="ws"} 0`) {
|
||||
t.Fatalf("missing ws series: %s", body)
|
||||
}
|
||||
}
|
||||
|
||||
func gaugeValue(t *testing.T, reg *metrics.Registry, name, transport string) float64 {
|
||||
t.Helper()
|
||||
mfs, err := reg.Gatherer().Gather()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, mf := range mfs {
|
||||
if mf.GetName() != name {
|
||||
continue
|
||||
}
|
||||
for _, m := range mf.GetMetric() {
|
||||
if matchLabel(m, "transport", transport) {
|
||||
return m.GetGauge().GetValue()
|
||||
}
|
||||
}
|
||||
}
|
||||
t.Fatalf("metric %s transport=%s not found", name, transport)
|
||||
return 0
|
||||
}
|
||||
|
||||
func matchLabel(m *dto.Metric, key, val string) bool {
|
||||
for _, lp := range m.GetLabel() {
|
||||
if lp.GetName() == key && lp.GetValue() == val {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,36 @@
|
||||
package metrics
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
)
|
||||
|
||||
// SampleStoreGauges 按库内真实计数刷新端总数、待投递、定时消息(无对应行则为 0)。
|
||||
func SampleStoreGauges(ctx context.Context, r *Registry, db *sql.DB) error {
|
||||
if r == nil || db == nil {
|
||||
return nil
|
||||
}
|
||||
var endpoints, pending, scheduled int
|
||||
if err := db.QueryRowContext(ctx, `SELECT COUNT(*) FROM endpoints`).Scan(&endpoints); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := db.QueryRowContext(ctx, `SELECT COUNT(*) FROM deliveries WHERE state = 'pending'`).Scan(&pending); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := db.QueryRowContext(ctx, `SELECT COUNT(*) FROM messages WHERE state = 'scheduled'`).Scan(&scheduled); err != nil {
|
||||
return err
|
||||
}
|
||||
r.EndpointsTotal.Set(float64(endpoints))
|
||||
r.DeliveriesPending.Set(float64(pending))
|
||||
r.MessagesScheduled.Set(float64(scheduled))
|
||||
return nil
|
||||
}
|
||||
|
||||
// SampleQueues 刷新写队列与密码哈希排队长度(传入当前真实长度,不做估算)。
|
||||
func SampleQueues(r *Registry, writeQueueLen, passwordHashQueueLen int) {
|
||||
if r == nil {
|
||||
return
|
||||
}
|
||||
r.WriteQueueLength.Set(float64(writeQueueLen))
|
||||
r.PasswordHashQueue.Set(float64(passwordHashQueueLen))
|
||||
}
|
||||
@@ -0,0 +1,84 @@
|
||||
package metrics
|
||||
|
||||
import (
|
||||
"context"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||
)
|
||||
|
||||
func TestSampleStoreGaugesPendingNonZero(t *testing.T) {
|
||||
t.Parallel()
|
||||
dir := t.TempDir()
|
||||
db, err := store.Open(filepath.Join(dir, "data"), "FULL")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
|
||||
ctx := context.Background()
|
||||
_, err = db.Write.ExecContext(ctx, `
|
||||
INSERT INTO endpoints(id, name, login_hash, enabled, source, created_at)
|
||||
VALUES ('alice', 'A', 'x', 1, 'admin', 1)`)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err = db.Write.ExecContext(ctx, `
|
||||
INSERT INTO messages(sender_id, id, dest_kind, dest_id, send_at, keep, ttl_seconds, receipt, state, reason, created_at, content_type, body_enc, meta)
|
||||
VALUES ('alice', 'm1', 'endpoint', 'bob', 100, 1, 0, 0, 'dispatched', '', 1, 'text/plain', 'utf8', '{}')`)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var seq int64
|
||||
err = db.Write.QueryRowContext(ctx, `SELECT seq FROM messages WHERE id='m1'`).Scan(&seq)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err = db.Write.ExecContext(ctx, `
|
||||
INSERT INTO deliveries(seq, endpoint_id, send_at, keep, state, reason, attempts, updated_at)
|
||||
VALUES (?, 'bob', 100, 1, 'pending', '', 0, 1)`, seq)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err = db.Write.ExecContext(ctx, `
|
||||
INSERT INTO messages(sender_id, id, dest_kind, dest_id, send_at, keep, ttl_seconds, receipt, state, reason, created_at, content_type, body_enc, meta)
|
||||
VALUES ('alice', 'm2', 'endpoint', 'bob', 999999, 0, 0, 0, 'scheduled', '', 1, 'text/plain', 'utf8', '{}')`)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
reg := New()
|
||||
err = SampleStoreGauges(ctx, reg, db.Read)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
SampleQueues(reg, 3, 1)
|
||||
|
||||
mfs, err := reg.Gatherer().Gather()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got := map[string]float64{}
|
||||
for _, mf := range mfs {
|
||||
switch mf.GetName() {
|
||||
case "nixmsg_endpoints", "nixmsg_deliveries_pending", "nixmsg_messages_scheduled",
|
||||
"nixmsg_write_queue_length", "nixmsg_password_hash_queue_length":
|
||||
if len(mf.GetMetric()) > 0 {
|
||||
got[mf.GetName()] = mf.GetMetric()[0].GetGauge().GetValue()
|
||||
}
|
||||
}
|
||||
}
|
||||
if got["nixmsg_endpoints"] != 1 {
|
||||
t.Fatalf("endpoints=%v", got["nixmsg_endpoints"])
|
||||
}
|
||||
if got["nixmsg_deliveries_pending"] != 1 {
|
||||
t.Fatalf("pending=%v", got["nixmsg_deliveries_pending"])
|
||||
}
|
||||
if got["nixmsg_messages_scheduled"] != 1 {
|
||||
t.Fatalf("scheduled=%v", got["nixmsg_messages_scheduled"])
|
||||
}
|
||||
if got["nixmsg_write_queue_length"] != 3 || got["nixmsg_password_hash_queue_length"] != 1 {
|
||||
t.Fatalf("queues=%v", got)
|
||||
}
|
||||
}
|
||||
@@ -43,6 +43,9 @@ type Queue struct {
|
||||
ready bool
|
||||
lastWriteErr error
|
||||
pending int
|
||||
|
||||
// OnBatchCommit 可选;每次合并提交成功后回调耗时(秒级指标用)。
|
||||
OnBatchCommit func(d time.Duration)
|
||||
}
|
||||
|
||||
// NewQueue 创建合并写入队列并启动写 goroutine。
|
||||
@@ -129,6 +132,7 @@ func (q *Queue) loop() {
|
||||
}
|
||||
|
||||
func (q *Queue) runBatch(batch []writeJob) {
|
||||
started := time.Now()
|
||||
defer func() {
|
||||
q.mu.Lock()
|
||||
q.pending -= len(batch)
|
||||
@@ -229,6 +233,9 @@ func (q *Queue) runBatch(batch []writeJob) {
|
||||
}
|
||||
return
|
||||
}
|
||||
if q.OnBatchCommit != nil {
|
||||
q.OnBatchCommit(time.Since(started))
|
||||
}
|
||||
for _, o := range outcomes {
|
||||
if o.success {
|
||||
o.job.res <- nil
|
||||
|
||||
Reference in New Issue
Block a user