fix: 合并停用删除重置密码发 fatal 与 revoked

This commit is contained in:
Nixevol
2026-09-30 12:08:51 +08:00
7 changed files with 396 additions and 36 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)
}
}
+30
View File
@@ -146,6 +146,7 @@ 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 { ClientIP: func(r *http.Request) string {
return httpx.ClientIP(r, trustedNets) return httpx.ClientIP(r, trustedNets)
}, },
@@ -176,6 +177,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
@@ -185,6 +187,34 @@ 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 err := sess.Disable(kickCtx, endpointID); err != nil {
return false, err
}
return true, nil
},
DeleteKick: func(kickCtx context.Context, endpointID string) (bool, error) {
if _, found := brk.ConnInfoOf(endpointID); !found {
return false, nil
}
if err := sess.Deleted(kickCtx, endpointID); err != nil {
return false, err
}
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() metricsReg := metrics.New()
+9
View File
@@ -1119,3 +1119,12 @@
- 原因:L-WIRE 已挂注册 Handler,管理与 WS 已接 `trusted_proxies`,唯独注册漏接,反向代理后会把安全码锁定计到代理 IP。 - 原因:L-WIRE 已挂注册 Handler,管理与 WS 已接 `trusted_proxies`,唯独注册漏接,反向代理后会把安全码锁定计到代理 IP。
- 备选方案:在 listener 层统一改写 `RemoteAddr` 后再交给注册 Handler。 - 备选方案:在 listener 层统一改写 `RemoteAddr` 后再交给注册 Handler。
- 影响:经受信代理开放注册时,输错安全码按真实客户端 IP 锁定。 - 影响:经受信代理开放注册时,输错安全码按真实客户端 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。
+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 -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)