mirror of
https://github.com/Buriburizaem0n/nezha_domains.git
synced 2026-09-21 10:40:13 +00:00
feat: server transfer rotation
This commit is contained in:
@@ -149,7 +149,7 @@ func checkStatus() {
|
||||
role = u.Role
|
||||
}
|
||||
UserLock.RUnlock()
|
||||
if alert.UserID != server.UserID && !role.IsAdmin() {
|
||||
if alert.UserID != server.GetUserID() && !role.IsAdmin() {
|
||||
continue
|
||||
}
|
||||
alertsStore[alert.ID][server.ID] = append(alertsStore[alert.
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package singleton
|
||||
|
||||
import (
|
||||
"log"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
@@ -26,6 +27,13 @@ func InitConfigFromPath(path string) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
rotated, err := Conf.RotateJWTSecretKeyIfNeeded(Version)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if rotated {
|
||||
log.Printf("NEZHA>> Rotated jwt_secret_key for dashboard version %s", Version)
|
||||
}
|
||||
|
||||
Conf.updateIgnoredIPNotificationID()
|
||||
Conf.Oauth2Providers = utils.MapKeysToSlice(Conf.Oauth2)
|
||||
|
||||
@@ -0,0 +1,54 @@
|
||||
package singleton
|
||||
|
||||
import (
|
||||
"os"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/nezhahq/nezha/model"
|
||||
)
|
||||
|
||||
func TestInitConfigFromPathRotatesJWTSecretKey(t *testing.T) {
|
||||
file, err := os.CreateTemp(t.TempDir(), "nezha-config-*.yaml")
|
||||
if err != nil {
|
||||
t.Fatalf("create temp config: %v", err)
|
||||
}
|
||||
if _, err := file.WriteString("jwt_secret_key: leaked-secret\nagent_secret_key: agent-secret\njwt_secret_key_last_rotated_version: v2.0.12\n"); err != nil {
|
||||
t.Fatalf("write temp config: %v", err)
|
||||
}
|
||||
if err := file.Close(); err != nil {
|
||||
t.Fatalf("close temp config: %v", err)
|
||||
}
|
||||
|
||||
originalConf := Conf
|
||||
originalVersion := Version
|
||||
originalTemplates := FrontendTemplates
|
||||
Version = "v2.0.13"
|
||||
FrontendTemplates = nil
|
||||
t.Cleanup(func() {
|
||||
Conf = originalConf
|
||||
Version = originalVersion
|
||||
FrontendTemplates = originalTemplates
|
||||
})
|
||||
|
||||
if err := InitConfigFromPath(file.Name()); err != nil {
|
||||
t.Fatalf("init config: %v", err)
|
||||
}
|
||||
if Conf.JWTSecretKey == "leaked-secret" {
|
||||
t.Fatal("jwt_secret_key was not rotated")
|
||||
}
|
||||
if Conf.JWTSecretKeyLastRotatedVersion != model.JWTSecretKeyRotationBaselineVersion {
|
||||
t.Fatalf("jwt secret key marker = %q, want %q", Conf.JWTSecretKeyLastRotatedVersion, model.JWTSecretKeyRotationBaselineVersion)
|
||||
}
|
||||
|
||||
saved, err := os.ReadFile(file.Name())
|
||||
if err != nil {
|
||||
t.Fatalf("read saved config: %v", err)
|
||||
}
|
||||
if strings.Contains(string(saved), "leaked-secret") {
|
||||
t.Fatalf("saved config still contains leaked jwt_secret_key: %s", saved)
|
||||
}
|
||||
if !strings.Contains(string(saved), "jwt_secret_key_last_rotated_version: v2.0.13") {
|
||||
t.Fatalf("saved config did not persist jwt secret key marker: %s", saved)
|
||||
}
|
||||
}
|
||||
@@ -264,12 +264,13 @@ func CronTrigger(cr *model.Cron, triggerServer ...uint64) func() {
|
||||
if !cronCanSendToServer(cr, s) {
|
||||
return
|
||||
}
|
||||
if s.TaskStream != nil {
|
||||
stream := s.GetTaskStream()
|
||||
if stream != nil {
|
||||
cronShared := CronShared
|
||||
if cronShared != nil {
|
||||
cronShared.reserveAlertTriggerCronResult(cr.ID, s.ID)
|
||||
}
|
||||
if err := s.TaskStream.Send(&pb.Task{
|
||||
if err := stream.Send(&pb.Task{
|
||||
Id: cr.ID,
|
||||
Data: cr.Command,
|
||||
Type: model.TaskTypeCommand,
|
||||
@@ -296,8 +297,8 @@ func CronTrigger(cr *model.Cron, triggerServer ...uint64) func() {
|
||||
if cr.Cover == model.CronCoverIgnoreAll && !crIgnoreMap[s.ID] {
|
||||
continue
|
||||
}
|
||||
if s.TaskStream != nil {
|
||||
s.TaskStream.Send(&pb.Task{
|
||||
if stream := s.GetTaskStream(); stream != nil {
|
||||
stream.Send(&pb.Task{
|
||||
Id: cr.ID,
|
||||
Data: cr.Command,
|
||||
Type: model.TaskTypeCommand,
|
||||
@@ -313,7 +314,7 @@ func CronTrigger(cr *model.Cron, triggerServer ...uint64) func() {
|
||||
}
|
||||
|
||||
func cronCanSendToServer(cr *model.Cron, server *model.Server) bool {
|
||||
return cr.UserID == server.UserID || userIsAdmin(cr.UserID)
|
||||
return cr.UserID == server.GetUserID() || userIsAdmin(cr.UserID)
|
||||
}
|
||||
|
||||
func userIsAdmin(userID uint64) bool {
|
||||
|
||||
@@ -39,6 +39,16 @@ func (s *capturedTaskStream) Context() context.Context { return context.Bac
|
||||
func (s *capturedTaskStream) SendMsg(any) error { return nil }
|
||||
func (s *capturedTaskStream) RecvMsg(any) error { return context.Canceled }
|
||||
|
||||
// withTaskStream attaches a TaskStream to a freshly constructed Server using the
|
||||
// new atomic accessor. The field itself is unexported (see Fix #12) precisely
|
||||
// because direct struct-literal access invited torn interface reads on hot
|
||||
// paths — tests use this helper rather than reaching in, mirroring production
|
||||
// callsites.
|
||||
func withTaskStream(s *model.Server, stream pb.NezhaService_RequestTaskServer) *model.Server {
|
||||
s.SetTaskStream(stream)
|
||||
return s
|
||||
}
|
||||
|
||||
func replaceServerSharedForSecurityTest(t *testing.T, servers ...*model.Server) {
|
||||
t.Helper()
|
||||
|
||||
@@ -75,8 +85,8 @@ func TestCronTriggerSkipsServersOwnedByOtherUsers(t *testing.T) {
|
||||
firstStream := newCapturedTaskStream()
|
||||
secondStream := newCapturedTaskStream()
|
||||
replaceServerSharedForSecurityTest(t,
|
||||
&model.Server{Common: model.Common{ID: 1, UserID: 100}, Name: "member-server", TaskStream: firstStream},
|
||||
&model.Server{Common: model.Common{ID: 2, UserID: 200}, Name: "admin-server", TaskStream: secondStream},
|
||||
withTaskStream(&model.Server{Common: model.Common{ID: 1, UserID: 100}, Name: "member-server"}, firstStream),
|
||||
withTaskStream(&model.Server{Common: model.Common{ID: 2, UserID: 200}, Name: "admin-server"}, secondStream),
|
||||
)
|
||||
|
||||
cronTask := &model.Cron{
|
||||
@@ -95,7 +105,7 @@ func TestCronTriggerSkipsServersOwnedByOtherUsers(t *testing.T) {
|
||||
func TestSendTriggerTasksSkipsCronOwnedByAnotherUser(t *testing.T) {
|
||||
attackerStream := newCapturedTaskStream()
|
||||
replaceServerSharedForSecurityTest(t,
|
||||
&model.Server{Common: model.Common{ID: 7, UserID: 200}, Name: "attacker-server", TaskStream: attackerStream},
|
||||
withTaskStream(&model.Server{Common: model.Common{ID: 7, UserID: 200}, Name: "attacker-server"}, attackerStream),
|
||||
)
|
||||
replaceUserInfoMapForSecurityTest(t, map[uint64]model.UserInfo{
|
||||
1: {Role: model.RoleAdmin},
|
||||
@@ -147,7 +157,7 @@ func assertNoTask(t *testing.T, stream *capturedTaskStream) {
|
||||
func TestCronTriggerSendsToMemberOwnedServer(t *testing.T) {
|
||||
memberStream := newCapturedTaskStream()
|
||||
replaceServerSharedForSecurityTest(t,
|
||||
&model.Server{Common: model.Common{ID: 1, UserID: 100}, Name: "member-server", TaskStream: memberStream},
|
||||
withTaskStream(&model.Server{Common: model.Common{ID: 1, UserID: 100}, Name: "member-server"}, memberStream),
|
||||
)
|
||||
replaceUserInfoMapForSecurityTest(t, map[uint64]model.UserInfo{
|
||||
100: {Role: model.RoleMember},
|
||||
@@ -168,8 +178,8 @@ func TestCronTriggerAdminCronFansOutAcrossOwners(t *testing.T) {
|
||||
first := newCapturedTaskStream()
|
||||
second := newCapturedTaskStream()
|
||||
replaceServerSharedForSecurityTest(t,
|
||||
&model.Server{Common: model.Common{ID: 1, UserID: 100}, Name: "member-server", TaskStream: first},
|
||||
&model.Server{Common: model.Common{ID: 2, UserID: 200}, Name: "admin-server", TaskStream: second},
|
||||
withTaskStream(&model.Server{Common: model.Common{ID: 1, UserID: 100}, Name: "member-server"}, first),
|
||||
withTaskStream(&model.Server{Common: model.Common{ID: 2, UserID: 200}, Name: "admin-server"}, second),
|
||||
)
|
||||
replaceUserInfoMapForSecurityTest(t, map[uint64]model.UserInfo{
|
||||
1: {Role: model.RoleAdmin},
|
||||
@@ -192,7 +202,7 @@ func TestCronTriggerAdminCronFansOutAcrossOwners(t *testing.T) {
|
||||
func TestCronTriggerLegacyZeroOwnerFansOut(t *testing.T) {
|
||||
first := newCapturedTaskStream()
|
||||
replaceServerSharedForSecurityTest(t,
|
||||
&model.Server{Common: model.Common{ID: 1, UserID: 100}, Name: "member-server", TaskStream: first},
|
||||
withTaskStream(&model.Server{Common: model.Common{ID: 1, UserID: 100}, Name: "member-server"}, first),
|
||||
)
|
||||
replaceUserInfoMapForSecurityTest(t, map[uint64]model.UserInfo{
|
||||
100: {Role: model.RoleMember},
|
||||
@@ -212,7 +222,7 @@ func TestCronTriggerLegacyZeroOwnerFansOut(t *testing.T) {
|
||||
func TestCronTriggerSkipsServersWhenOwnerNotKnown(t *testing.T) {
|
||||
stream := newCapturedTaskStream()
|
||||
replaceServerSharedForSecurityTest(t,
|
||||
&model.Server{Common: model.Common{ID: 1, UserID: 100}, Name: "member-server", TaskStream: stream},
|
||||
withTaskStream(&model.Server{Common: model.Common{ID: 1, UserID: 100}, Name: "member-server"}, stream),
|
||||
)
|
||||
replaceUserInfoMapForSecurityTest(t, map[uint64]model.UserInfo{
|
||||
100: {Role: model.RoleMember},
|
||||
@@ -232,7 +242,7 @@ func TestCronTriggerSkipsServersWhenOwnerNotKnown(t *testing.T) {
|
||||
func TestSendTriggerTasksAllowsSelfOwnedCron(t *testing.T) {
|
||||
stream := newCapturedTaskStream()
|
||||
replaceServerSharedForSecurityTest(t,
|
||||
&model.Server{Common: model.Common{ID: 7, UserID: 200}, Name: "member-server", TaskStream: stream},
|
||||
withTaskStream(&model.Server{Common: model.Common{ID: 7, UserID: 200}, Name: "member-server"}, stream),
|
||||
)
|
||||
replaceUserInfoMapForSecurityTest(t, map[uint64]model.UserInfo{
|
||||
200: {Role: model.RoleMember},
|
||||
@@ -257,7 +267,7 @@ func TestSendTriggerTasksAllowsSelfOwnedCron(t *testing.T) {
|
||||
func TestSendTriggerTasksAllowsAdminCallerToTriggerAny(t *testing.T) {
|
||||
stream := newCapturedTaskStream()
|
||||
replaceServerSharedForSecurityTest(t,
|
||||
&model.Server{Common: model.Common{ID: 9, UserID: 100}, Name: "any-server", TaskStream: stream},
|
||||
withTaskStream(&model.Server{Common: model.Common{ID: 9, UserID: 100}, Name: "any-server"}, stream),
|
||||
)
|
||||
replaceUserInfoMapForSecurityTest(t, map[uint64]model.UserInfo{
|
||||
1: {Role: model.RoleAdmin},
|
||||
@@ -283,7 +293,7 @@ func TestSendTriggerTasksAllowsAdminCallerToTriggerAny(t *testing.T) {
|
||||
func TestSendTriggerTasksIgnoresUnknownTaskIDs(t *testing.T) {
|
||||
stream := newCapturedTaskStream()
|
||||
replaceServerSharedForSecurityTest(t,
|
||||
&model.Server{Common: model.Common{ID: 7, UserID: 200}, Name: "member-server", TaskStream: stream},
|
||||
withTaskStream(&model.Server{Common: model.Common{ID: 7, UserID: 200}, Name: "member-server"}, stream),
|
||||
)
|
||||
replaceUserInfoMapForSecurityTest(t, map[uint64]model.UserInfo{
|
||||
200: {Role: model.RoleMember},
|
||||
@@ -304,7 +314,7 @@ func TestSendTriggerTasksIgnoresUnknownTaskIDs(t *testing.T) {
|
||||
func TestSendTriggerTasksMixedCronIDsOnlyFiresAllowed(t *testing.T) {
|
||||
stream := newCapturedTaskStream()
|
||||
replaceServerSharedForSecurityTest(t,
|
||||
&model.Server{Common: model.Common{ID: 7, UserID: 200}, Name: "member-server", TaskStream: stream},
|
||||
withTaskStream(&model.Server{Common: model.Common{ID: 7, UserID: 200}, Name: "member-server"}, stream),
|
||||
)
|
||||
replaceUserInfoMapForSecurityTest(t, map[uint64]model.UserInfo{
|
||||
1: {Role: model.RoleAdmin},
|
||||
@@ -540,7 +550,7 @@ func (s *failingTaskStream) Send(task *pb.Task) error {
|
||||
func TestCronTriggerRevokesAlertTriggerAuthorizationOnSendFailure(t *testing.T) {
|
||||
failing := newFailingTaskStream(context.Canceled)
|
||||
replaceServerSharedForSecurityTest(t,
|
||||
&model.Server{Common: model.Common{ID: 7, UserID: 100}, Name: "broken-server", TaskStream: failing},
|
||||
withTaskStream(&model.Server{Common: model.Common{ID: 7, UserID: 100}, Name: "broken-server"}, failing),
|
||||
)
|
||||
replaceUserInfoMapForSecurityTest(t, map[uint64]model.UserInfo{
|
||||
100: {Role: model.RoleMember},
|
||||
@@ -828,6 +838,12 @@ func newServiceMonitorSecurityHarness(t *testing.T, servers ...*model.Server) *S
|
||||
t.Fatal(err)
|
||||
}
|
||||
ServiceSentinelShared = ss
|
||||
// LIFO Cleanup ordering: this Close() runs BEFORE the earlier t.Cleanup that
|
||||
// restores Conf/Cache/CronShared/NotificationShared/TSDBShared, so the
|
||||
// worker has fully exited before we swap those globals out. Skipping this
|
||||
// step causes `go test -race` to flag the write-vs-read between the
|
||||
// teardown and the still-running worker.
|
||||
t.Cleanup(func() { ss.Close() })
|
||||
return ss
|
||||
}
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -85,6 +85,16 @@ type ServiceSentinel struct {
|
||||
// 30天数据缓存
|
||||
monthlyStatusLock sync.Mutex
|
||||
monthlyStatus map[uint64]*serviceResponseItem
|
||||
|
||||
// closeOnce + workerWG together let Close() wait for the worker goroutine
|
||||
// to fully exit. Without this, a test that swaps ServiceSentinelShared back
|
||||
// to its original value in t.Cleanup races against the still-running
|
||||
// worker, which keeps reading globals like Conf/CronShared/NotificationShared.
|
||||
// Production never calls Close() — the process exits while the worker is
|
||||
// still running and that is fine — but tests must drain the worker before
|
||||
// restoring globals.
|
||||
closeOnce sync.Once
|
||||
workerWG sync.WaitGroup
|
||||
}
|
||||
|
||||
// NewServiceSentinel 创建服务监控器
|
||||
@@ -113,7 +123,11 @@ func NewServiceSentinel(serviceSentinelDispatchBus chan<- *model.Service) (*Serv
|
||||
ss.loadTodayStats(today)
|
||||
|
||||
// 启动服务监控器
|
||||
go ss.worker()
|
||||
ss.workerWG.Add(1)
|
||||
go func() {
|
||||
defer ss.workerWG.Done()
|
||||
ss.worker()
|
||||
}()
|
||||
|
||||
// 每日将游标往后推一天
|
||||
_, err = CronShared.AddFunc("0 0 0 * * *", ss.refreshMonthlyServiceStatus)
|
||||
@@ -489,10 +503,34 @@ func canReportServiceResult(service *model.Service, reporter *model.Server, task
|
||||
return false
|
||||
}
|
||||
|
||||
return service.UserID == reporter.UserID || userIsAdmin(service.UserID)
|
||||
return service.UserID == reporter.GetUserID() || userIsAdmin(service.UserID)
|
||||
}
|
||||
|
||||
// Close shuts down the ServiceSentinel worker goroutine and waits for it to
|
||||
// exit. It is idempotent and safe to call more than once.
|
||||
//
|
||||
// Why this exists: the worker reads multiple package-level globals during
|
||||
// each report (Conf, CronShared via notifyCheck, NotificationShared via
|
||||
// UnMuteNotification, ServerShared, TSDBShared). A test fixture that swaps
|
||||
// those globals out in t.Cleanup MUST first call Close() — otherwise the
|
||||
// cleanup write races the still-running worker's read and `go test -race`
|
||||
// fires (see security_regression_test.go newServiceMonitorSecurityHarness).
|
||||
// Production never calls Close because the process exits with the worker
|
||||
// still running, which is fine.
|
||||
func (ss *ServiceSentinel) Close() {
|
||||
ss.closeOnce.Do(func() {
|
||||
close(ss.serviceReportChannel)
|
||||
ss.workerWG.Wait()
|
||||
})
|
||||
}
|
||||
|
||||
// worker 服务监控的实际工作流程
|
||||
//
|
||||
// IMPORTANT: this loop reads several package-level globals (Conf, CronShared,
|
||||
// NotificationShared, ServerShared, TSDBShared). Any test that replaces those
|
||||
// globals via t.Cleanup must first call ServiceSentinel.Close() so the worker
|
||||
// drains and exits before the swap, otherwise the race detector trips. See
|
||||
// the Close() comment above for the full rationale.
|
||||
func (ss *ServiceSentinel) worker() {
|
||||
// 从服务状态汇报管道获取汇报的服务数据
|
||||
for r := range ss.serviceReportChannel {
|
||||
|
||||
@@ -34,6 +34,10 @@ var (
|
||||
NotificationShared *NotificationClass
|
||||
NATShared *NATClass
|
||||
CronShared *CronClass
|
||||
// ServerTransferShared is initialized in LoadSingleton AFTER ServerShared
|
||||
// (so the in-memory pending index can write back into ServerShared.UserID
|
||||
// on transitions) and AFTER initUser (so PushIfOnline can read secrets
|
||||
// from UserInfoMap).
|
||||
)
|
||||
|
||||
//go:embed frontend-templates.yaml
|
||||
@@ -59,6 +63,7 @@ func LoadSingleton(bus chan<- *model.Service) (err error) {
|
||||
NotificationShared = NewNotificationClass()
|
||||
ServerShared = NewServerClass()
|
||||
CronShared = NewCronClass()
|
||||
ServerTransferShared = NewServerTransferClass()
|
||||
// 最后初始化 ServiceSentinel
|
||||
ServiceSentinelShared, err = NewServiceSentinel(bus)
|
||||
return
|
||||
@@ -89,7 +94,7 @@ func InitDBFromPath(path string) error {
|
||||
model.Notification{}, model.AlertRule{}, model.Service{}, model.NotificationGroupNotification{},
|
||||
model.Cron{}, model.Transfer{}, model.ServerGroupServer{},
|
||||
model.NAT{}, model.DDNSProfile{}, model.NotificationGroupNotification{},
|
||||
model.WAF{}, model.Oauth2Bind{})
|
||||
model.WAF{}, model.Oauth2Bind{}, model.ServerTransfer{})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -40,10 +40,34 @@ func initUser() {
|
||||
|
||||
UserInfoMap[u.ID] = model.UserInfo{
|
||||
Role: u.Role,
|
||||
Username: u.Username,
|
||||
AgentSecret: u.AgentSecret,
|
||||
}
|
||||
AgentSecretToUserId[u.AgentSecret] = u.ID
|
||||
}
|
||||
|
||||
model.ServerOwnerLookup = lookupServerOwner
|
||||
}
|
||||
|
||||
// lookupServerOwner resolves Server.UserID into a display-ready owner
|
||||
// record for model.Server.MarshalJSON. uid=0 is the legacy global agent
|
||||
// secret (a pseudo-owner with no User row) and intentionally returns
|
||||
// ok=false with no username; the frontend renders it as "Global". Other
|
||||
// uids return ok=false when the user has been deleted, so the JSON still
|
||||
// carries the bare id and the frontend can render an "Unknown (#<uid>)"
|
||||
// placeholder. RLock is required because OnUserUpdate / OnUserDelete may
|
||||
// mutate UserInfoMap concurrently with serialization.
|
||||
func lookupServerOwner(uid uint64) (model.ServerOwnerInfo, bool) {
|
||||
if uid == 0 {
|
||||
return model.ServerOwnerInfo{}, false
|
||||
}
|
||||
UserLock.RLock()
|
||||
info, ok := UserInfoMap[uid]
|
||||
UserLock.RUnlock()
|
||||
if !ok {
|
||||
return model.ServerOwnerInfo{}, false
|
||||
}
|
||||
return model.ServerOwnerInfo{ID: uid, Username: info.Username}, true
|
||||
}
|
||||
|
||||
func OnUserUpdate(u *model.User) {
|
||||
@@ -56,6 +80,7 @@ func OnUserUpdate(u *model.User) {
|
||||
|
||||
UserInfoMap[u.ID] = model.UserInfo{
|
||||
Role: u.Role,
|
||||
Username: u.Username,
|
||||
AgentSecret: u.AgentSecret,
|
||||
}
|
||||
AgentSecretToUserId[u.AgentSecret] = u.ID
|
||||
@@ -69,6 +94,10 @@ func OnUserDelete(id []uint64, errorFunc func(string, ...any) error) error {
|
||||
return Localizer.ErrorT("user id not specified")
|
||||
}
|
||||
|
||||
if ServerTransferShared != nil {
|
||||
ServerTransferShared.OnUsersDeleted(id)
|
||||
}
|
||||
|
||||
var (
|
||||
cron, server bool
|
||||
crons, servers []uint64
|
||||
@@ -127,6 +156,11 @@ func OnUserDelete(id []uint64, errorFunc func(string, ...any) error) error {
|
||||
}
|
||||
}
|
||||
AlertsLock.Unlock()
|
||||
// Cancel pending transfers before ServerShared drops the
|
||||
// in-memory entry: same ordering rationale as batchDeleteServer.
|
||||
if ServerTransferShared != nil {
|
||||
ServerTransferShared.OnServersDeleted(servers)
|
||||
}
|
||||
ServerShared.Delete(servers)
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,90 @@
|
||||
package singleton
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/nezhahq/nezha/model"
|
||||
"github.com/nezhahq/nezha/pkg/i18n"
|
||||
)
|
||||
|
||||
func setupOnUserDeleteFixture(t *testing.T) (*ServerTransferClass, func()) {
|
||||
t.Helper()
|
||||
c, transferCleanup := setupTransferFixture(t)
|
||||
|
||||
require.NoError(t, DB.AutoMigrate(&model.Cron{}, &model.Transfer{}, &model.ServerGroupServer{}))
|
||||
originalCronShared := CronShared
|
||||
CronShared = &CronClass{
|
||||
class: class[uint64, *model.Cron]{list: map[uint64]*model.Cron{}},
|
||||
}
|
||||
|
||||
originalLocalizer := Localizer
|
||||
Localizer = i18n.NewLocalizer("zh_CN", domain, "translations", i18n.Translations)
|
||||
|
||||
cleanup := func() {
|
||||
Localizer = originalLocalizer
|
||||
CronShared = originalCronShared
|
||||
transferCleanup()
|
||||
}
|
||||
return c, cleanup
|
||||
}
|
||||
|
||||
func TestOnUserDeleteCancelsPendingTransfersAwayFromDeletedUser(t *testing.T) {
|
||||
c, cleanup := setupOnUserDeleteFixture(t)
|
||||
defer cleanup()
|
||||
|
||||
const fromUser = uint64(100)
|
||||
const toUser = uint64(200)
|
||||
const serverID = uint64(1)
|
||||
|
||||
seedServerForTransfer(t, serverID, fromUser)
|
||||
|
||||
require.NoError(t, DB.AutoMigrate(&model.User{}))
|
||||
require.NoError(t, DB.Create(&model.User{
|
||||
Common: model.Common{ID: fromUser},
|
||||
Username: "alice",
|
||||
AgentSecret: "alice-secret",
|
||||
}).Error)
|
||||
require.NoError(t, DB.Create(&model.User{
|
||||
Common: model.Common{ID: toUser},
|
||||
Username: "bob",
|
||||
AgentSecret: "bob-secret",
|
||||
}).Error)
|
||||
UserLock.Lock()
|
||||
UserInfoMap[fromUser] = model.UserInfo{Role: model.RoleMember, AgentSecret: "alice-secret"}
|
||||
UserInfoMap[toUser] = model.UserInfo{Role: model.RoleMember, AgentSecret: "bob-secret"}
|
||||
UserLock.Unlock()
|
||||
|
||||
tr := initiateAndRegister(t, c, serverID, fromUser, toUser, fromUser)
|
||||
require.True(t, c.HasPending(serverID), "precondition: pending transfer published")
|
||||
|
||||
srv, ok := ServerShared.Get(serverID)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, toUser, srv.GetUserID(), "precondition: pending transfer flipped owner to ToUserID")
|
||||
|
||||
require.NoError(t, OnUserDelete([]uint64{fromUser}, func(format string, args ...any) error {
|
||||
return nil
|
||||
}))
|
||||
|
||||
if c.HasPending(serverID) {
|
||||
t.Fatal("OnUserDelete on the transfer FromUserID must terminate the pending transfer so a later Cancel/Fail/Timeout cannot revert ownership to the deleted user")
|
||||
}
|
||||
|
||||
if srv, ok := ServerShared.Get(serverID); ok {
|
||||
require.NotEqual(t, fromUser, srv.GetUserID(),
|
||||
"server owner must not be reverted to the deleted FromUserID; got owner=%d", srv.GetUserID())
|
||||
}
|
||||
|
||||
if _, err := c.Cancel(tr.ID); err == nil {
|
||||
var refreshed model.ServerTransfer
|
||||
if err := DB.First(&refreshed, tr.ID).Error; err == nil {
|
||||
require.NotEqual(t, model.ServerTransferStatusPending, refreshed.Status,
|
||||
"after OnUserDelete a subsequent Cancel must not leave the transfer Pending")
|
||||
if srv, ok := ServerShared.Get(serverID); ok {
|
||||
require.NotEqual(t, fromUser, srv.GetUserID(),
|
||||
"a late Cancel against the terminated transfer must not revert server.UserID to the deleted FromUserID; got owner=%d", srv.GetUserID())
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user