fix: 重置密码踢线失败时回退断开并记 ok_kick_failed

This commit is contained in:
Nixevol
2026-09-30 16:24:04 +08:00
parent e294b5db71
commit 8f47167128
3 changed files with 203 additions and 8 deletions
+33 -8
View File
@@ -104,28 +104,49 @@ func (h *Handler) kickEndpoint(ctx context.Context, id string) (bool, error) {
func (h *Handler) passwordResetKick(ctx context.Context, id string) (bool, error) {
if h.resetKick != nil {
return h.resetKick(ctx, id)
kicked, err := h.resetKick(ctx, id)
if err == nil {
return kicked, nil
}
h.log.Error("password reset kick failed", "endpoint", id, "err", err)
fallbackKicked, fallbackErr := h.kickEndpoint(ctx, id)
if fallbackErr != nil {
h.log.Error("password reset kick fallback failed", "endpoint", id, "err", fallbackErr)
}
return fallbackKicked, err
}
return h.kickEndpoint(ctx, id)
kicked, err := h.kickEndpoint(ctx, id)
if err != nil {
h.log.Error("password reset kick failed", "endpoint", id, "err", err)
}
return kicked, err
}
func (h *Handler) afterDisableKick(ctx context.Context, id string) {
if h.disableKick != nil {
_, _ = h.disableKick(ctx, id)
if _, err := h.disableKick(ctx, id); err != nil {
h.log.Error("disable kick failed", "endpoint", id, "err", err)
}
return
}
if h.identity == nil {
_, _ = h.kickEndpoint(ctx, id)
if _, err := h.kickEndpoint(ctx, id); err != nil {
h.log.Error("disable kick failed", "endpoint", id, "err", err)
}
}
}
func (h *Handler) afterDeleteKick(ctx context.Context, id string) {
if h.deleteKick != nil {
_, _ = h.deleteKick(ctx, id)
if _, err := h.deleteKick(ctx, id); err != nil {
h.log.Error("delete kick failed", "endpoint", id, "err", err)
}
return
}
if h.identity == nil {
_, _ = h.kickEndpoint(ctx, id)
if _, err := h.kickEndpoint(ctx, id); err != nil {
h.log.Error("delete kick failed", "endpoint", id, "err", err)
}
}
}
@@ -557,8 +578,12 @@ func (h *Handler) handleEndpointResetLoginPassword(w http.ResponseWriter, r *htt
httpx.WriteError(w, http.StatusNotFound, "not_found", "端不存在")
return
}
_, _ = h.passwordResetKick(r.Context(), id)
h.auditP(p, "endpoint_reset_login_password", id, "ok", ip)
_, kickErr := h.passwordResetKick(r.Context(), id)
result := "ok"
if kickErr != nil {
result = "ok_kick_failed"
}
h.auditP(p, "endpoint_reset_login_password", id, result, ip)
httpx.WriteOK(w, map[string]any{loginPasswordOnceKey: pw})
}
+161
View File
@@ -5,7 +5,9 @@ import (
"context"
"database/sql"
"encoding/json"
"errors"
"io"
"log/slog"
"net/http"
"net/http/cookiejar"
"net/http/httptest"
@@ -369,6 +371,165 @@ func TestEndpointImportBOMAndCreate(t *testing.T) {
}
}
func TestEndpointResetPasswordKickFailureFallsBack(t *testing.T) {
dir := t.TempDir()
db, err := store.Open(filepath.Join(dir, "data"), "FULL")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
hash := auth.NewStubHashPool()
if seedErr := admin.SeedAdminPassword(context.Background(), db, hash, testPassword); seedErr != nil {
t.Fatal(seedErr)
}
kick := &kickRecorder{}
var logBuf bytes.Buffer
logger := slog.New(slog.NewTextHandler(&logBuf, nil))
h := admin.New(admin.Deps{
DB: db,
Hash: hash,
Tokens: admin.NewRandomAPITokens(),
Locks: admin.NewMemoryLoginLocks(),
Logger: logger,
KickEndpoint: kick.Kick,
PasswordResetKick: func(_ context.Context, _ string) (bool, error) {
return false, errors.New("fatal kick failed")
},
})
srv := httptest.NewServer(h)
t.Cleanup(srv.Close)
jar, err := cookiejar.New(nil)
if err != nil {
t.Fatal(err)
}
client := &http.Client{Jar: jar}
res := postJSON(t, client, srv.URL+"/api/admin/login",
`{"username":"admin","password":"`+testPassword+`"}`, nil)
env := decodeEnv(t, res)
if res.StatusCode != http.StatusOK || !env.OK {
t.Fatalf("login failed: %d %+v", res.StatusCode, env)
}
res = postJSON(t, client, srv.URL+"/api/admin/endpoints",
`{"id":"kick-fail","name":"踢失败","login_password":"oldpass12"}`,
csrfHeaders())
env = decodeEnv(t, res)
if res.StatusCode != 200 || !env.OK {
t.Fatalf("create: %d %+v", res.StatusCode, env)
}
res = postJSON(t, client, srv.URL+"/api/admin/endpoints/kick-fail/reset-login-password",
`{}`, csrfHeaders())
body, _ := io.ReadAll(res.Body)
_ = res.Body.Close()
if res.StatusCode != 200 {
t.Fatalf("reset status=%d body=%s", res.StatusCode, body)
}
var resetEnv struct {
OK bool `json:"ok"`
Data struct {
LoginPassword string `json:"login_password"`
} `json:"data"`
}
if err := json.Unmarshal(body, &resetEnv); err != nil {
t.Fatal(err)
}
if !resetEnv.OK || resetEnv.Data.LoginPassword == "" {
t.Fatalf("want one-time password after kick failure, got %s", body)
}
if kick.count() != 1 {
t.Fatalf("want KickEndpoint fallback once, got %d", kick.count())
}
logs := logBuf.String()
if !strings.Contains(logs, "action=endpoint_reset_login_password") {
t.Fatalf("want reset password audit, logs=%s", logs)
}
if !strings.Contains(logs, "result=ok_kick_failed") {
t.Fatalf("audit want ok_kick_failed, logs=%s", logs)
}
if !strings.Contains(logs, "password reset kick failed") || !strings.Contains(logs, "fatal kick failed") {
t.Fatalf("want error log for reset kick, logs=%s", logs)
}
}
func TestDisableDeleteKickLogsError(t *testing.T) {
dir := t.TempDir()
db, err := store.Open(filepath.Join(dir, "data"), "FULL")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
hash := auth.NewStubHashPool()
if seedErr := admin.SeedAdminPassword(context.Background(), db, hash, testPassword); seedErr != nil {
t.Fatal(seedErr)
}
kick := &kickRecorder{}
var logBuf bytes.Buffer
logger := slog.New(slog.NewTextHandler(&logBuf, nil))
kickErr := errors.New("hook kick failed")
h := admin.New(admin.Deps{
DB: db,
Hash: hash,
Tokens: admin.NewRandomAPITokens(),
Locks: admin.NewMemoryLoginLocks(),
Logger: logger,
KickEndpoint: kick.Kick,
DisableKick: func(_ context.Context, _ string) (bool, error) {
return false, kickErr
},
DeleteKick: func(_ context.Context, _ string) (bool, error) {
return false, kickErr
},
})
srv := httptest.NewServer(h)
t.Cleanup(srv.Close)
jar, err := cookiejar.New(nil)
if err != nil {
t.Fatal(err)
}
client := &http.Client{Jar: jar}
res := postJSON(t, client, srv.URL+"/api/admin/login",
`{"username":"admin","password":"`+testPassword+`"}`, nil)
env := decodeEnv(t, res)
if res.StatusCode != http.StatusOK || !env.OK {
t.Fatalf("login failed: %d %+v", res.StatusCode, env)
}
for _, id := range []string{"dis-fail", "del-fail"} {
res = postJSON(t, client, srv.URL+"/api/admin/endpoints",
`{"id":"`+id+`","name":"`+id+`","login_password":"password1"}`,
csrfHeaders())
env = decodeEnv(t, res)
if res.StatusCode != 200 || !env.OK {
t.Fatalf("create %s: %d %+v", id, res.StatusCode, env)
}
}
res = postJSON(t, client, srv.URL+"/api/admin/endpoints/batch",
`{"ids":["dis-fail"],"action":"disable"}`, csrfHeaders())
env = decodeEnv(t, res)
if res.StatusCode != 200 || !env.OK {
t.Fatalf("disable: %d %+v", res.StatusCode, env)
}
res = doReq(t, client, http.MethodDelete, srv.URL+"/api/admin/endpoints/del-fail", "", csrfHeaders())
env = decodeEnv(t, res)
if res.StatusCode != 200 || !env.OK {
t.Fatalf("delete: %d %+v", res.StatusCode, env)
}
if kick.count() != 0 {
t.Fatalf("disable/delete hook errors must not fall back to KickEndpoint, got %d", kick.count())
}
logs := logBuf.String()
if !strings.Contains(logs, "disable kick failed") || !strings.Contains(logs, "delete kick failed") {
t.Fatalf("want disable/delete kick error logs, logs=%s", logs)
}
}
func csrfHeaders() map[string]string {
return map[string]string{"X-Nixmsg-Request": "1"}
}