12 Commits
22 changed files with 964 additions and 55 deletions
@@ -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)
}
}
+51 -7
View File
@@ -71,6 +71,10 @@ func runServe(ctx context.Context, cfg config.Config) error {
loginLocks := auth.NewLoginLocks() loginLocks := auth.NewLoginLocks()
memConns := message.NewMemoryConns() memConns := message.NewMemoryConns()
msgLim := message.LimitsFromFullConfig(cfg) msgLim := message.LimitsFromFullConfig(cfg)
metricsReg := metrics.New()
db.Queue.OnBatchCommit = func(d time.Duration) {
metricsReg.WriteCommitSeconds.Observe(d.Seconds())
}
login := broker.NewLogin(broker.LoginOptions{ login := broker.NewLogin(broker.LoginOptions{
DB: db, DB: db,
@@ -84,11 +88,13 @@ func runServe(ctx context.Context, cfg config.Config) error {
msgApp := message.New(db, msgLim, hashPool, msgApp := message.New(db, msgLim, hashPool,
message.WithLocks(loginLocks), message.WithLocks(loginLocks),
message.WithConnRegistry(memConns), message.WithConnRegistry(memConns),
message.WithMetrics(metricsReg),
) )
uplink := &appUplink{ uplink := &appUplink{
msg: msgApp, msg: msgApp,
conns: memConns, conns: memConns,
log: slog.Default(), log: slog.Default(),
metrics: metricsReg,
} }
sess := broker.NewSession(broker.SessionOptions{ sess := broker.NewSession(broker.SessionOptions{
Login: login, Login: login,
@@ -109,6 +115,7 @@ func runServe(ctx context.Context, cfg config.Config) error {
Authenticator: login, Authenticator: login,
Uplink: sess, Uplink: sess,
Logger: slog.Default(), Logger: slog.Default(),
Metrics: metricsReg,
OnPublishDropped: func(dropCtx context.Context, endpointID string, connID port.ConnID, payload []byte) { OnPublishDropped: func(dropCtx context.Context, endpointID string, connID port.ConnID, payload []byte) {
if dropErr := msgApp.OnPublishDropped(dropCtx, endpointID, connID, payload); dropErr != nil { if dropErr := msgApp.OnPublishDropped(dropCtx, endpointID, connID, payload); dropErr != nil {
slog.Error("on publish dropped", "endpoint", endpointID, "err", dropErr) 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.WithLocks(loginLocks),
message.WithConnRegistry(memConns), message.WithConnRegistry(memConns),
message.WithDownlink(brk), message.WithDownlink(brk),
message.WithMetrics(metricsReg),
) )
uplink.msg = msgApp uplink.msg = msgApp
uplink.down = brk uplink.down = brk
@@ -137,6 +145,7 @@ func runServe(ctx context.Context, cfg config.Config) error {
sess.SetPresence(presApp) sess.SetPresence(presApp)
uplink.presence = presApp uplink.presence = presApp
trustedNets := httpx.ParseCIDRs(cfg.TrustedProxies)
idApp := identity.New(identity.Config{ idApp := identity.New(identity.Config{
DB: db, DB: db,
Hash: hashPool, Hash: hashPool,
@@ -145,6 +154,10 @@ func runServe(ctx context.Context, cfg config.Config) error {
MaxScheduleSeconds: int64(cfg.Limits.MaxScheduleSeconds), MaxScheduleSeconds: int64(cfg.Limits.MaxScheduleSeconds),
Logger: slog.Default(), Logger: slog.Default(),
ConnControl: brk, ConnControl: brk,
Downlink: brk,
ClientIP: func(r *http.Request) string {
return httpx.ClientIP(r, trustedNets)
},
}) })
uplink.identity = idApp uplink.identity = idApp
@@ -161,7 +174,6 @@ func runServe(ctx context.Context, cfg config.Config) error {
return fmt.Errorf("message recover: %w", recoverErr) return fmt.Errorf("message recover: %w", recoverErr)
} }
trustedNets := httpx.ParseCIDRs(cfg.TrustedProxies)
adminHandler := admin.New(admin.Deps{ adminHandler := admin.New(admin.Deps{
DB: db, DB: db,
Hash: hashPool, Hash: hashPool,
@@ -173,6 +185,7 @@ func runServe(ctx context.Context, cfg config.Config) error {
Groups: groupApp, Groups: groupApp,
Config: cfg, Config: cfg,
Version: Version, Version: Version,
// Kick:只断开,令牌不变,SDK 重连(PRD 踢下线)。
KickEndpoint: func(kickCtx context.Context, endpointID string) (bool, error) { KickEndpoint: func(kickCtx context.Context, endpointID string) (bool, error) {
if _, found := brk.ConnInfoOf(endpointID); !found { if _, found := brk.ConnInfoOf(endpointID); !found {
return false, nil return false, nil
@@ -182,9 +195,36 @@ func runServe(ctx context.Context, cfg config.Config) error {
} }
return true, nil 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 { buildHandlers := func(proxies *listener.ProxySet) listener.Handlers {
return listener.Handlers{ return listener.Handlers{
MQTT: brk.WSHandler(proxies), MQTT: brk.WSHandler(proxies),
@@ -266,7 +306,7 @@ func runServe(ctx context.Context, cfg config.Config) error {
loopCtx, loopCancel := context.WithCancel(ctx) loopCtx, loopCancel := context.WithCancel(ctx)
defer loopCancel() defer loopCancel()
go messageLoops(loopCtx, msgApp, memConns) go messageLoops(loopCtx, msgApp, memConns, db, hashPool, metricsReg)
<-ctx.Done() <-ctx.Done()
loopCancel() loopCancel()
@@ -279,7 +319,7 @@ func runServe(ctx context.Context, cfg config.Config) error {
return nil 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) t := time.NewTicker(time.Second)
defer t.Stop() defer t.Stop()
for { for {
@@ -299,6 +339,10 @@ func messageLoops(ctx context.Context, msgApp *message.App, conns *message.Memor
if err := msgApp.CleanupOnce(ctx, nowMs); err != nil { if err := msgApp.CleanupOnce(ctx, nowMs); err != nil {
slog.Error("cleanup once", "err", err) 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())
} }
} }
} }
+5
View File
@@ -11,6 +11,7 @@ import (
"git.asio.asia/nixevol/NixMsg/internal/app/message" "git.asio.asia/nixevol/NixMsg/internal/app/message"
"git.asio.asia/nixevol/NixMsg/internal/app/port" "git.asio.asia/nixevol/NixMsg/internal/app/port"
"git.asio.asia/nixevol/NixMsg/internal/app/presence" "git.asio.asia/nixevol/NixMsg/internal/app/presence"
"git.asio.asia/nixevol/NixMsg/internal/metrics"
"git.asio.asia/nixevol/NixMsg/internal/protocol" "git.asio.asia/nixevol/NixMsg/internal/protocol"
) )
@@ -23,6 +24,7 @@ type appUplink struct {
conns *message.MemoryConns conns *message.MemoryConns
down port.Downlink down port.Downlink
log *slog.Logger log *slog.Logger
metrics *metrics.Registry
} }
func (u *appUplink) OnSessionEstablished(ctx context.Context, conn port.ConnInfo) error { 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 == "" { if rid == "" {
rid = "0" rid = "0"
} }
if u.metrics != nil && code != "" {
u.metrics.ErrorsTotal.WithLabelValues(code).Inc()
}
resp := protocol.Resp{ resp := protocol.Resp{
V: protocol.Version, V: protocol.Version,
Type: protocol.TypeResp, Type: protocol.TypeResp,
+79
View File
@@ -1092,3 +1092,82 @@
- 原因:同一连接上 `group.create`/`group.add` 同步向本连接注入下行时,与 mochi InlineClient 互相等待,`resp` 回不去(`TestUplinkDMOfflineGroupRecall` 在清掉测试客户端 dial deadline 后稳定复现)。 - 原因:同一连接上 `group.create`/`group.add` 同步向本连接注入下行时,与 mochi InlineClient 互相等待,`resp` 回不去(`TestUplinkDMOfflineGroupRecall` 在清掉测试客户端 dial deadline 后稳定复现)。
- 备选方案:broker 层对 Inline 发布做无锁队列。 - 备选方案:broker 层对 Inline 发布做无锁队列。
- 影响:`group_event` 可能略晚于 `resp` 到达;业务结果仍以 `resp` 为准。 - 影响:`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**
- 原条款:PRD F23 / D18 注册安全码按来源 IP 锁定;DEVELOPMENT 4.5 来自受信代理时用 `X-Forwarded-For`;I1.4 曾写「经代理部署时接线方必须注入真实 IP」。
- 实际做法:`cmd/nixmsg/serve.go` 在 `identity.New` 注入与管理接口相同的 `httpx.ClientIP(r, trustedNets)`;不改锁定阈值与注册开关/安全码语义,不在 identity 内复制解析。
- 原因: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`。
+76
View File
@@ -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) { func TestBadAPITokenCountsTowardLock(t *testing.T) {
_, srv, _, _ := setup(t) _, srv, _, _ := setup(t)
base := srv.URL base := srv.URL
+2 -6
View File
@@ -43,12 +43,8 @@ func (h *Handler) auth(next http.HandlerFunc) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
ip := httpx.ClientIP(r, h.trusted) ip := httpx.ClientIP(r, h.trusted)
if locked, retry := h.locks.Check(auth.LockKey{Kind: auth.LockAdminIP, IP: ip}); locked { // 锁定只拦密码登录(login.go)与错误令牌试错累计;
w.Header().Set("Retry-After", formatRetryAfter(retry)) // 已认证的 Cookie / 合法 API 令牌在锁定期仍可用(对齐 PRD D18)。
httpx.WriteError(w, http.StatusTooManyRequests, "rate_limited", "登录已锁定,请稍后再试")
return
}
p, errCode, errMsg, status := h.authenticate(r, ip) p, errCode, errMsg, status := h.authenticate(r, ip)
if status != 0 { if status != 0 {
if status == http.StatusTooManyRequests { if status == http.StatusTooManyRequests {
+32 -5
View File
@@ -102,6 +102,33 @@ func (h *Handler) kickEndpoint(ctx context.Context, id string) (bool, error) {
return h.kick(ctx, id) 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) { func (h *Handler) handleEndpointList(w http.ResponseWriter, r *http.Request) {
q := r.URL.Query() q := r.URL.Query()
limit := defaultListLimit limit := defaultListLimit
@@ -350,7 +377,7 @@ func (h *Handler) handleEndpointPatch(w http.ResponseWriter, r *http.Request) {
return return
} }
if !*req.Enabled { if !*req.Enabled {
_, _ = h.kickEndpoint(r.Context(), id) h.afterDisableKick(r.Context(), id)
} }
} else if req.Enabled != nil && !*req.Enabled && wasEnabled { } else if req.Enabled != nil && !*req.Enabled && wasEnabled {
_, _ = h.kickEndpoint(r.Context(), id) _, _ = 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", "端不存在") httpx.WriteError(w, http.StatusNotFound, "not_found", "端不存在")
return return
} }
_, _ = h.kickEndpoint(r.Context(), id) h.afterDeleteKick(r.Context(), id)
h.audit(actorString(p), "endpoint_delete", id, "ok", ip) h.audit(actorString(p), "endpoint_delete", id, "ok", ip)
httpx.WriteOK(w, map[string]any{}) httpx.WriteOK(w, map[string]any{})
} }
@@ -416,14 +443,14 @@ func (h *Handler) handleEndpointBatch(w http.ResponseWriter, r *http.Request) {
case "disable": case "disable":
found, opErr = h.setEndpointEnabled(r.Context(), id, false) found, opErr = h.setEndpointEnabled(r.Context(), id, false)
if found && opErr == nil { if found && opErr == nil {
_, _ = h.kickEndpoint(r.Context(), id) h.afterDisableKick(r.Context(), id)
} }
case "enable": case "enable":
found, opErr = h.setEndpointEnabled(r.Context(), id, true) found, opErr = h.setEndpointEnabled(r.Context(), id, true)
case "delete": case "delete":
found, opErr = h.deleteEndpointBasic(r.Context(), id) found, opErr = h.deleteEndpointBasic(r.Context(), id)
if found && opErr == nil { if found && opErr == nil {
_, _ = h.kickEndpoint(r.Context(), id) h.afterDeleteKick(r.Context(), id)
} }
} }
if opErr != nil { if opErr != nil {
@@ -512,7 +539,7 @@ func (h *Handler) handleEndpointResetLoginPassword(w http.ResponseWriter, r *htt
httpx.WriteError(w, http.StatusNotFound, "not_found", "端不存在") httpx.WriteError(w, http.StatusNotFound, "not_found", "端不存在")
return return
} }
_, _ = h.kickEndpoint(r.Context(), id) _, _ = h.passwordResetKick(r.Context(), id)
h.audit(actorString(p), "endpoint_reset_login_password", id, "ok", ip) h.audit(actorString(p), "endpoint_reset_login_password", id, "ok", ip)
httpx.WriteOK(w, map[string]any{loginPasswordOnceKey: pw}) httpx.WriteOK(w, map[string]any{loginPasswordOnceKey: pw})
} }
+42 -30
View File
@@ -41,6 +41,12 @@ type Deps struct {
SecureCookies bool SecureCookies bool
// KickEndpoint 踢下线钩子(只断开连接);nil 时踢线为 no-op。 // KickEndpoint 踢下线钩子(只断开连接);nil 时踢线为 no-op。
KickEndpoint EndpointKickFunc 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 端停用/启用/删除级联(I5);nil 时回退为仅改 enabled/删行。
Identity identity.Service Identity identity.Service
@@ -54,20 +60,23 @@ type Deps struct {
// Handler 是可挂载的管理接口(路由前缀 /api/admin/)。 // Handler 是可挂载的管理接口(路由前缀 /api/admin/)。
type Handler struct { type Handler struct {
db *store.DB db *store.DB
hash auth.HashPool hash auth.HashPool
tokens auth.APITokens tokens auth.APITokens
locks auth.LoginLocks locks auth.LoginLocks
log *slog.Logger log *slog.Logger
trusted []*net.IPNet trusted []*net.IPNet
ttl time.Duration ttl time.Duration
forceSec bool forceSec bool
kick EndpointKickFunc kick EndpointKickFunc
identity identity.Service resetKick EndpointKickFunc
groups group.Service disableKick EndpointKickFunc
cfg config.Config deleteKick EndpointKickFunc
version string identity identity.Service
startedAt time.Time groups group.Service
cfg config.Config
version string
startedAt time.Time
mux *http.ServeMux mux *http.ServeMux
@@ -96,22 +105,25 @@ func New(d Deps) *Handler {
ver = "dev" ver = "dev"
} }
h := &Handler{ h := &Handler{
db: d.DB, db: d.DB,
hash: d.Hash, hash: d.Hash,
tokens: d.Tokens, tokens: d.Tokens,
locks: d.Locks, locks: d.Locks,
log: d.Logger, log: d.Logger,
trusted: d.TrustedProxies, trusted: d.TrustedProxies,
ttl: ttl, ttl: ttl,
forceSec: d.SecureCookies, forceSec: d.SecureCookies,
kick: d.KickEndpoint, kick: d.KickEndpoint,
identity: d.Identity, resetKick: d.PasswordResetKick,
groups: d.Groups, disableKick: d.DisableKick,
cfg: cfg, deleteKick: d.DeleteKick,
version: ver, identity: d.Identity,
startedAt: time.Now(), groups: d.Groups,
mux: http.NewServeMux(), cfg: cfg,
lastUsed: make(map[string]time.Time), version: ver,
startedAt: time.Now(),
mux: http.NewServeMux(),
lastUsed: make(map[string]time.Time),
} }
h.routes() h.routes()
return h return h
+10
View File
@@ -296,6 +296,16 @@ WHERE m.id='gm1' AND d.endpoint_id='bob'`).Scan(&reason)
if err != nil || state != "completed" || mreason != "group_dissolved" { if err != nil || state != "completed" || mreason != "group_dissolved" {
t.Fatalf("state=%s reason=%s err=%v", state, mreason, err) 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{ created2, err := gApp.Create(ctx, "alice", &protocol.GroupCreate{
V: protocol.Version, Type: protocol.TypeGroupCreate, RID: "6", V: protocol.Version, Type: protocol.TypeGroupCreate, RID: "6",
+2 -1
View File
@@ -139,10 +139,11 @@ UPDATE messages SET state = 'completed', reason = ? WHERE seq = ? AND state = 's
return execErr return execErr
} }
if r.receipt != 0 { if r.receipt != 0 {
// 消息级作废回执:endpoint_id 空,state=rejected(DEVELOPMENT 6.4);消息行仍为 completed
if _, execErr := tx.Exec(` if _, execErr := tx.Exec(`
INSERT INTO receipts(sender_id, msg_id, endpoint_id, state, reason, created_at, acked) INSERT INTO receipts(sender_id, msg_id, endpoint_id, state, reason, created_at, acked)
VALUES(?,?,?,?,?,?,0)`, VALUES(?,?,?,?,?,?,0)`,
r.senderID, r.msgID, "", "completed", reasonGroupDissolved, nowMs); execErr != nil { r.senderID, r.msgID, "", "rejected", reasonGroupDissolved, nowMs); execErr != nil {
return execErr return execErr
} }
} }
+10 -1
View File
@@ -5,6 +5,7 @@ import (
"context" "context"
"database/sql" "database/sql"
"errors" "errors"
"time"
"git.asio.asia/nixevol/NixMsg/internal/app/port" "git.asio.asia/nixevol/NixMsg/internal/app/port"
"git.asio.asia/nixevol/NixMsg/internal/protocol" "git.asio.asia/nixevol/NixMsg/internal/protocol"
@@ -22,6 +23,9 @@ const (
eventLeft = "left" eventLeft = "left"
eventMemberRemoved = "member_removed" eventMemberRemoved = "member_removed"
eventDissolved = "dissolved" eventDissolved = "dissolved"
// kickFlushDelay 给接线方 Session.Disable/Deleted 留出发 fatal 的窗口。
kickFlushDelay = 20 * time.Millisecond
) )
type revokeItem struct { type revokeItem struct {
@@ -124,8 +128,13 @@ WHERE id = ?`, endpointID); e != nil {
a.publishRevokes(ctx, revokes) a.publishRevokes(ctx, revokes)
a.publishGroupEvents(ctx, notifies) a.publishGroupEvents(ctx, notifies)
// fatal+断开由 admin DisableKick/DeleteKick(Session.Disable/Deleted)完成。
// 未接 Kick 钩子的单元测试仍可用 ConnControl 兜底断开。
if a.connCtrl != nil { 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 return nil
} }
+77
View File
@@ -3,6 +3,7 @@ package identity_test
import ( import (
"context" "context"
"database/sql" "database/sql"
"encoding/json"
"io" "io"
"net/http" "net/http"
"net/http/cookiejar" "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) { func TestF01DisableVoidsScheduledAndRejectsNew(t *testing.T) {
t.Parallel() t.Parallel()
idApp, msgApp, db := openLifecycle(t) idApp, msgApp, db := openLifecycle(t)
+78
View File
@@ -15,6 +15,7 @@ import (
"time" "time"
"git.asio.asia/nixevol/NixMsg/internal/auth" "git.asio.asia/nixevol/NixMsg/internal/auth"
"git.asio.asia/nixevol/NixMsg/internal/httpx"
"git.asio.asia/nixevol/NixMsg/internal/protocol" "git.asio.asia/nixevol/NixMsg/internal/protocol"
"git.asio.asia/nixevol/NixMsg/internal/store" "git.asio.asia/nixevol/NixMsg/internal/store"
) )
@@ -309,6 +310,83 @@ func TestRegisterF23_WrongCodeLock(t *testing.T) {
} }
} }
// TestRegisterTrustedProxyClientIPLock 验证与管理接口相同的 httpx.ClientIP:
// 受信代理的 X-Forwarded-For 按真实客户端 IP 计锁;非信任来源不采信转发头。
func TestRegisterTrustedProxyClientIPLock(t *testing.T) {
db, err := store.Open(t.TempDir(), "FULL")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
locks := newRegisterIPLocker(time.Now)
trusted := httpx.ParseCIDRs([]string{"127.0.0.1/32"})
handler := NewRegisterHandler(RegisterConfig{
DB: db,
Hash: auth.NewStubHashPool(),
Locks: locks,
ClientIP: func(r *http.Request) string {
return httpx.ClientIP(r, trusted)
},
})
env := &testEnv{db: db, hash: auth.NewStubHashPool(), locks: locks, handler: handler}
env.setRegistration(t, true, "proxy-lock-1")
post := func(remote, xff, body string) (int, registerResp) {
t.Helper()
req := httptest.NewRequest(http.MethodPost, "/api/client/register", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
req.RemoteAddr = remote
if xff != "" {
req.Header.Set("X-Forwarded-For", xff)
}
rr := httptest.NewRecorder()
handler.ServeHTTP(rr, req)
var resp registerResp
if err := json.Unmarshal(rr.Body.Bytes(), &resp); err != nil {
t.Fatalf("decode: %v body=%s", err, rr.Body.String())
}
return rr.Code, resp
}
wrong := `{"registration_code":"wrong-code","id":"ep_px","login_password":"password1"}`
good := `{"registration_code":"proxy-lock-1","id":"ep_px","login_password":"password1"}`
for i := 0; i < 10; i++ {
code, resp := post("127.0.0.1:9000", "198.51.100.7", wrong)
if code != http.StatusForbidden || resp.Error == nil || resp.Error.Code != protocol.CodeRegistrationCodeInvalid {
t.Fatalf("trusted fail #%d: status=%d resp=%+v", i+1, code, resp)
}
}
code, resp := post("127.0.0.1:9000", "198.51.100.7", good)
if code != http.StatusTooManyRequests || resp.Error == nil || resp.Error.Code != protocol.CodeRateLimited {
t.Fatalf("real client should be locked: status=%d resp=%+v", code, resp)
}
code, resp = post("127.0.0.1:9000", "198.51.100.8", good)
if code != http.StatusOK || !resp.OK || resp.Data.ID != "ep_px" {
t.Fatalf("other XFF client must not share lock: status=%d resp=%+v", code, resp)
}
// 非信任对端:忽略 XFF,按 RemoteAddr 计锁。
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)
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)
}
}
code, resp = post("203.0.113.50:4433", "198.51.100.7", `{"registration_code":"proxy-lock-1","id":"ep_px2","login_password":"password1"}`)
if code != http.StatusTooManyRequests || resp.Error == nil || resp.Error.Code != protocol.CodeRateLimited {
t.Fatalf("untrusted RemoteAddr should be locked: status=%d resp=%+v", code, resp)
}
// 若误采信 XFF,198.51.100.7 会已锁;直连该 IP 应仍可注册。
code, resp = post("198.51.100.7:5555", "", `{"registration_code":"proxy-lock-1","id":"ep_px3","login_password":"password1"}`)
if code != http.StatusOK || !resp.OK || resp.Data.ID != "ep_px3" {
t.Fatalf("spoofed XFF must not lock real client: status=%d resp=%+v", code, resp)
}
}
func TestRegisterF23_IDTakenKeepsOriginal(t *testing.T) { func TestRegisterF23_IDTakenKeepsOriginal(t *testing.T) {
env := openTestEnv(t) env := openTestEnv(t)
env.setRegistration(t, true, "taken-code") env.setRegistration(t, true, "taken-code")
+13
View File
@@ -17,6 +17,8 @@ func (a *App) Ack(ctx context.Context, endpointID string, req *protocol.Ack) (Ac
nowMs := a.now().UnixMilli() nowMs := a.now().UnixMilli()
var out AckResult var out AckResult
var seq int64 var seq int64
var ackLatencySec float64
var observeAck bool
err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { 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) err := tx.QueryRow(`SELECT seq FROM messages WHERE sender_id = ? AND id = ?`, req.From, req.ID).Scan(&seq)
if err == sql.ErrNoRows { if err == sql.ErrNoRows {
@@ -25,6 +27,10 @@ func (a *App) Ack(ctx context.Context, endpointID string, req *protocol.Ack) (Ac
if err != nil { if err != nil {
return err 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(` res, err := tx.Exec(`
UPDATE deliveries SET state = ?, reason = '', pushed_conn = NULL, updated_at = ? UPDATE deliveries SET state = ?, reason = '', pushed_conn = NULL, updated_at = ?
WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`, WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
@@ -35,6 +41,10 @@ WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`,
aff, _ := res.RowsAffected() aff, _ := res.RowsAffected()
if aff > 0 { if aff > 0 {
out.Result = DeliveryAccepted 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 { if e := insertReceiptTx(tx, req.From, seq, endpointID, DeliveryAccepted, "", nowMs); e != nil {
return e return e
} }
@@ -55,6 +65,9 @@ SELECT state FROM deliveries WHERE seq = ? AND endpoint_id = ?`, seq, endpointID
if err != nil { if err != nil {
return out, err return out, err
} }
if observeAck && a.met != nil {
a.met.AckSeconds.Observe(ackLatencySec)
}
if out.Result == DeliveryAccepted { if out.Result == DeliveryAccepted {
a.releaseLarge(seq, endpointID) a.releaseLarge(seq, endpointID)
} }
+7
View File
@@ -7,6 +7,7 @@ import (
"git.asio.asia/nixevol/NixMsg/internal/app/port" "git.asio.asia/nixevol/NixMsg/internal/app/port"
"git.asio.asia/nixevol/NixMsg/internal/auth" "git.asio.asia/nixevol/NixMsg/internal/auth"
"git.asio.asia/nixevol/NixMsg/internal/config" "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/protocol"
"git.asio.asia/nixevol/NixMsg/internal/store" "git.asio.asia/nixevol/NixMsg/internal/store"
) )
@@ -86,6 +87,7 @@ type App struct {
down port.Downlink down port.Downlink
conns ConnRegistry conns ConnRegistry
met *metrics.Registry
mu sync.Mutex mu sync.Mutex
largeSem chan struct{} largeSem chan struct{}
@@ -117,6 +119,11 @@ func WithConnRegistry(c ConnRegistry) Option {
return func(a *App) { a.conns = c } return func(a *App) { a.conns = c }
} }
// WithMetrics 注入 Prometheus 注册表(投递耗时直方图)。
func WithMetrics(m *metrics.Registry) Option {
return func(a *App) { a.met = m }
}
// New 创建消息服务实现。 // New 创建消息服务实现。
func New(db *store.DB, lim Limits, hash auth.HashPool, opts ...Option) *App { func New(db *store.DB, lim Limits, hash auth.HashPool, opts ...Option) *App {
if lim.RequestBurst <= 0 { if lim.RequestBurst <= 0 {
+9
View File
@@ -229,11 +229,20 @@ WHERE seq = ? AND endpoint_id = ? AND state = 'pending' AND pushed_conn IS NULL`
a.releaseLarge(it.seq, endpointID) a.releaseLarge(it.seq, endpointID)
} }
a.scheduleRepush(endpointID, time.Second) a.scheduleRepush(endpointID, time.Second)
} else {
a.observeDispatchToPush(it.sendAt, nowMs)
} }
} }
return a.pushReceipts(ctx, endpointID, connID, 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 { 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 { return a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
res, err := tx.Exec(` res, err := tx.Exec(`
+11 -5
View File
@@ -12,6 +12,7 @@ import (
"time" "time"
"git.asio.asia/nixevol/NixMsg/internal/app/port" "git.asio.asia/nixevol/NixMsg/internal/app/port"
"git.asio.asia/nixevol/NixMsg/internal/metrics"
mqtt "github.com/mochi-mqtt/server/v2" mqtt "github.com/mochi-mqtt/server/v2"
"github.com/mochi-mqtt/server/v2/packets" "github.com/mochi-mqtt/server/v2/packets"
) )
@@ -68,15 +69,18 @@ type Options struct {
Logger *slog.Logger Logger *slog.Logger
// OnPublishDropped 可选;nil 时仅打 debug 日志。 // OnPublishDropped 可选;nil 时仅打 debug 日志。
OnPublishDropped PublishDroppedFunc OnPublishDropped PublishDroppedFunc
// Metrics 可选;会话建立/断开时更新 nixmsg_connections。
Metrics *metrics.Registry
} }
// Broker 内置 mochi,不自带监听端口。 // Broker 内置 mochi,不自带监听端口。
type Broker struct { type Broker struct {
server *mqtt.Server server *mqtt.Server
auth Authenticator auth Authenticator
uplink port.UplinkHandler uplink port.UplinkHandler
log *slog.Logger log *slog.Logger
onDrop PublishDroppedFunc onDrop PublishDroppedFunc
metrics *metrics.Registry
hook *nixHook hook *nixHook
@@ -105,6 +109,7 @@ type connState struct {
handshook bool handshook bool
subscribedDown bool subscribedDown bool
largeHeld int largeHeld int
metricsCounted bool
mu sync.Mutex mu sync.Mutex
handshakeTimer *time.Timer handshakeTimer *time.Timer
@@ -151,6 +156,7 @@ func New(opts Options) (*Broker, error) {
uplink: uplink, uplink: uplink,
log: log, log: log,
onDrop: opts.OnPublishDropped, onDrop: opts.OnPublishDropped,
metrics: opts.Metrics,
current: make(map[string]*connState), current: make(map[string]*connState),
byClient: make(map[*mqtt.Client]*connState), byClient: make(map[*mqtt.Client]*connState),
queues: make(map[string]*uplinkQueue), queues: make(map[string]*uplinkQueue),
+19
View File
@@ -190,6 +190,7 @@ func (h *nixHook) OnSessionEstablished(cl *mqtt.Client, _ packets.Packet) {
MaxPacketSize: st.maxPacketSize, MaxPacketSize: st.maxPacketSize,
} }
_ = h.b.uplink.OnSessionEstablished(context.Background(), info) _ = h.b.uplink.OnSessionEstablished(context.Background(), info)
h.noteConnectionOpen(st)
} }
func (h *nixHook) OnDisconnect(cl *mqtt.Client, err error, _ bool) { 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 { if sess, ok := h.b.uplink.(*Session); ok {
sess.HandleDisconnect(context.Background(), info, reason, isCurrent) sess.HandleDisconnect(context.Background(), info, reason, isCurrent)
h.noteConnectionClose(st)
return return
} }
h.b.uplink.OnDisconnect(context.Background(), info, reason) 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) { func (h *nixHook) OnQosComplete(cl *mqtt.Client, pk packets.Packet) {
+118
View File
@@ -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
}
+36
View File
@@ -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))
}
+84
View File
@@ -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)
}
}
+7
View File
@@ -43,6 +43,9 @@ type Queue struct {
ready bool ready bool
lastWriteErr error lastWriteErr error
pending int pending int
// OnBatchCommit 可选;每次合并提交成功后回调耗时(秒级指标用)。
OnBatchCommit func(d time.Duration)
} }
// NewQueue 创建合并写入队列并启动写 goroutine。 // NewQueue 创建合并写入队列并启动写 goroutine。
@@ -129,6 +132,7 @@ func (q *Queue) loop() {
} }
func (q *Queue) runBatch(batch []writeJob) { func (q *Queue) runBatch(batch []writeJob) {
started := time.Now()
defer func() { defer func() {
q.mu.Lock() q.mu.Lock()
q.pending -= len(batch) q.pending -= len(batch)
@@ -229,6 +233,9 @@ func (q *Queue) runBatch(batch []writeJob) {
} }
return return
} }
if q.OnBatchCommit != nil {
q.OnBatchCommit(time.Since(started))
}
for _, o := range outcomes { for _, o := range outcomes {
if o.success { if o.success {
o.job.res <- nil o.job.res <- nil