fix: 合并停用删除重置密码发 fatal 与 revoked
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user