From 6b88cdb0122c5c8968a24d58ccabf22b15c02f67 Mon Sep 17 00:00:00 2001 From: naiba Date: Mon, 25 May 2026 10:15:21 +0000 Subject: [PATCH] feat: server transfer rotation --- cmd/dashboard/controller/controller.go | 76 +- cmd/dashboard/controller/fm.go | 8 +- .../controller/force_update_ownership_test.go | 2 +- .../controller/frontend_fallback_test.go | 133 + .../controller/permission_matrix_test.go | 15 + cmd/dashboard/controller/server.go | 138 +- .../controller/stream_ownership_test.go | 33 + cmd/dashboard/controller/terminal.go | 8 +- cmd/dashboard/controller/transfer.go | 217 ++ .../controller/transfer_list_filter_test.go | 231 ++ .../controller/transfer_retry_authz_test.go | 219 ++ cmd/dashboard/controller/ws.go | 2 +- cmd/dashboard/rpc/rpc.go | 40 +- go.mod | 6 +- go.sum | 12 +- model/common.go | 25 +- model/config.go | 74 +- model/config_test.go | 105 + model/notification_test.go | 1 - model/server.go | 108 +- model/server_owner_json_test.go | 129 + model/server_owner_race_test.go | 86 + model/server_taskstream_race_test.go | 91 + model/server_transfer.go | 128 + model/service.go | 7 +- model/user.go | 1 + service/rpc/apply_config_authz_test.go | 565 +++++ service/rpc/auth.go | 182 +- service/rpc/auth_test.go | 507 +++- service/rpc/io_stream.go | 31 + service/rpc/nezha.go | 34 +- service/rpc/request_task_security_test.go | 68 +- service/singleton/alertsentinel.go | 2 +- service/singleton/config.go | 8 + service/singleton/config_test.go | 54 + service/singleton/crontask.go | 11 +- service/singleton/security_regression_test.go | 42 +- service/singleton/server_transfer.go | 1440 +++++++++++ service/singleton/server_transfer_test.go | 2194 +++++++++++++++++ service/singleton/servicesentinel.go | 42 +- service/singleton/singleton.go | 7 +- service/singleton/user.go | 34 + service/singleton/user_test.go | 90 + 43 files changed, 7072 insertions(+), 134 deletions(-) create mode 100644 cmd/dashboard/controller/frontend_fallback_test.go create mode 100644 cmd/dashboard/controller/transfer.go create mode 100644 cmd/dashboard/controller/transfer_list_filter_test.go create mode 100644 cmd/dashboard/controller/transfer_retry_authz_test.go create mode 100644 model/server_owner_json_test.go create mode 100644 model/server_owner_race_test.go create mode 100644 model/server_taskstream_race_test.go create mode 100644 model/server_transfer.go create mode 100644 service/rpc/apply_config_authz_test.go create mode 100644 service/singleton/config_test.go create mode 100644 service/singleton/server_transfer.go create mode 100644 service/singleton/server_transfer_test.go create mode 100644 service/singleton/user_test.go diff --git a/cmd/dashboard/controller/controller.go b/cmd/dashboard/controller/controller.go index 13518b8c..b771d861 100644 --- a/cmd/dashboard/controller/controller.go +++ b/cmd/dashboard/controller/controller.go @@ -8,7 +8,6 @@ import ( "log" "net/http" "os" - "path" "regexp" "slices" "strings" @@ -118,6 +117,11 @@ func routers(r *gin.Engine, frontendDist fs.FS) { auth.POST("/batch-move/server", commonHandler(batchMoveServer)) auth.POST("/force-update/server", commonHandler(forceUpdateServer)) + auth.GET("/transfer", listHandler(listServerTransfer)) + auth.POST("/transfer/:id/cancel", commonHandler(cancelServerTransfer)) + auth.POST("/transfer/:id/retry", commonHandler(retryServerTransfer)) + auth.GET("/ws/transfer", commonHandler(transferStream)) + auth.GET("/notification", listHandler(listNotification)) auth.POST("/notification", commonHandler(createNotification)) auth.PATCH("/notification/:id", commonHandler(updateNotification)) @@ -320,27 +324,52 @@ func getUid(c *gin.Context) uint64 { } func fallbackToFrontend(frontendDist fs.FS) func(*gin.Context) { - checkLocalFileOrFs := func(c *gin.Context, fs fs.FS, path string, customStatusCode int) bool { - if _, err := os.Stat(path); err == nil { - http.ServeFile(utils.NewGinCustomWriter(c, customStatusCode), c.Request, path) - return true - } - f, err := fs.Open(path) - if err != nil { - return false - } - defer f.Close() - fileStat, err := f.Stat() + serveFile := func(c *gin.Context, name string, file fs.File, customStatusCode int) bool { + defer file.Close() + fileStat, err := file.Stat() if err != nil { return false } if fileStat.IsDir() { return false } - http.ServeContent(utils.NewGinCustomWriter(c, customStatusCode), c.Request, path, fileStat.ModTime(), f.(io.ReadSeeker)) + readSeeker, ok := file.(io.ReadSeeker) + if !ok { + return false + } + http.ServeContent(utils.NewGinCustomWriter(c, customStatusCode), c.Request, name, fileStat.ModTime(), readSeeker) return true } + checkLocalFileOrFs := func(c *gin.Context, frontendFS fs.FS, templateRoot, filePath string, customStatusCode int) bool { + if filePath != "" { + localRoot, err := os.OpenRoot(templateRoot) + if err == nil { + defer localRoot.Close() + // URL paths must stay inside the selected template root; never join them against the process cwd. + if file, err := localRoot.Open(filePath); err == nil && serveFile(c, filePath, file, customStatusCode) { + return true + } + } + } + + if !fs.ValidPath(filePath) { + return false + } + templateFS, err := fs.Sub(frontendFS, templateRoot) + if err != nil { + return false + } + file, err := templateFS.Open(filePath) + if err != nil { + return false + } + if serveFile(c, filePath, file, customStatusCode) { + return true + } + return false + } + frontendPageUrlRegistry := []*regexp.Regexp{ // official user frontend regexp.MustCompile(`^/$`), @@ -361,6 +390,11 @@ func fallbackToFrontend(frontendDist fs.FS) func(*gin.Context) { regexp.MustCompile(`^/dashboard/settings/user$`), regexp.MustCompile(`^/dashboard/settings/online-user$`), regexp.MustCompile(`^/dashboard/settings/waf$`), + // 注意:这里的白名单决定哪些 URL 走 index.html fallback;漏一条就会把 + // 直接刷新该页面变成 404(HTTP 状态码层面,body 仍是 index.html,所以 + // 浏览器内 SPA 看起来正常,但 monitoring / 链接预览会以为站点挂了)。 + // 新增前端路由时必须在 admin-frontend/src/main.tsx 与这里同步加。 + regexp.MustCompile(`^/dashboard/transfer$`), } getFallbackStatusCode := func(path string) int { @@ -385,22 +419,22 @@ func fallbackToFrontend(frontendDist fs.FS) func(*gin.Context) { } fallbackStatusCode := getFallbackStatusCode(c.Request.URL.Path) - if strings.HasPrefix(c.Request.URL.Path, "/dashboard") { - stripPath := strings.TrimPrefix(c.Request.URL.Path, "/dashboard") - localFilePath := path.Join(singleton.Conf.AdminTemplate, stripPath) - if checkLocalFileOrFs(c, frontendDist, localFilePath, http.StatusOK) { + // Only /dashboard/ belongs to the admin frontend; /dashboard.. must not be trimmed into ../. + if strings.HasPrefix(c.Request.URL.Path, "/dashboard/") { + stripPath := strings.TrimPrefix(c.Request.URL.Path, "/dashboard/") + if checkLocalFileOrFs(c, frontendDist, singleton.Conf.AdminTemplate, stripPath, http.StatusOK) { return } - if !checkLocalFileOrFs(c, frontendDist, singleton.Conf.AdminTemplate+"/index.html", fallbackStatusCode) { + if !checkLocalFileOrFs(c, frontendDist, singleton.Conf.AdminTemplate, "index.html", fallbackStatusCode) { c.JSON(http.StatusNotFound, newErrorResponse(errors.New("404 Not Found"))) } return } - localFilePath := path.Join(singleton.Conf.UserTemplate, c.Request.URL.Path) - if checkLocalFileOrFs(c, frontendDist, localFilePath, http.StatusOK) { + stripPath := strings.TrimPrefix(c.Request.URL.Path, "/") + if checkLocalFileOrFs(c, frontendDist, singleton.Conf.UserTemplate, stripPath, http.StatusOK) { return } - if !checkLocalFileOrFs(c, frontendDist, singleton.Conf.UserTemplate+"/index.html", fallbackStatusCode) { + if !checkLocalFileOrFs(c, frontendDist, singleton.Conf.UserTemplate, "index.html", fallbackStatusCode) { c.JSON(http.StatusNotFound, newErrorResponse(errors.New("404 Not Found"))) } } diff --git a/cmd/dashboard/controller/fm.go b/cmd/dashboard/controller/fm.go index 32d53783..638dcb7b 100644 --- a/cmd/dashboard/controller/fm.go +++ b/cmd/dashboard/controller/fm.go @@ -33,7 +33,11 @@ func createFM(c *gin.Context) (*model.CreateFMResponse, error) { } server, _ := singleton.ServerShared.Get(id) - if server == nil || server.TaskStream == nil { + if server == nil { + return nil, singleton.Localizer.ErrorT("server not found or not connected") + } + stream := server.GetTaskStream() + if stream == nil { return nil, singleton.Localizer.ErrorT("server not found or not connected") } @@ -51,7 +55,7 @@ func createFM(c *gin.Context) (*model.CreateFMResponse, error) { fmData, _ := json.Marshal(&model.TaskFM{ StreamID: streamId, }) - if err := server.TaskStream.Send(&proto.Task{ + if err := stream.Send(&proto.Task{ Type: model.TaskTypeFM, Data: string(fmData), }); err != nil { diff --git a/cmd/dashboard/controller/force_update_ownership_test.go b/cmd/dashboard/controller/force_update_ownership_test.go index d501b8e8..76f4be15 100644 --- a/cmd/dashboard/controller/force_update_ownership_test.go +++ b/cmd/dashboard/controller/force_update_ownership_test.go @@ -56,7 +56,7 @@ func setupServerOwnershipFixture(t *testing.T) (stream *fakeTaskStream, reset fu alice, _ := singleton.ServerShared.Get(1) stream = &fakeTaskStream{} - alice.TaskStream = stream + alice.SetTaskStream(stream) return stream, func() { singleton.DB = originalDB diff --git a/cmd/dashboard/controller/frontend_fallback_test.go b/cmd/dashboard/controller/frontend_fallback_test.go new file mode 100644 index 00000000..f4ca8744 --- /dev/null +++ b/cmd/dashboard/controller/frontend_fallback_test.go @@ -0,0 +1,133 @@ +package controller + +import ( + "io/fs" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" + + "github.com/gin-gonic/gin" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/service/singleton" +) + +func newFrontendFallbackTestRouter(t *testing.T) *gin.Engine { + t.Helper() + gin.SetMode(gin.TestMode) + + originalConf := singleton.Conf + singleton.Conf = &singleton.ConfigClass{Config: &model.Config{ + ConfigDashboard: model.ConfigDashboard{ + AdminTemplate: "admin-dist", + UserTemplate: "user-dist", + }, + }} + t.Cleanup(func() { singleton.Conf = originalConf }) + + writeFrontendFallbackTestFile(t, "admin-dist/index.html", "admin index") + writeFrontendFallbackTestFile(t, "admin-dist/assets/app.js", "console.log('admin asset')") + writeFrontendFallbackTestFile(t, "user-dist/index.html", "user index") + writeFrontendFallbackTestFile(t, "data/config.yaml", "jwt_secret_key: traversal-secret") + + r := gin.New() + r.NoRoute(fallbackToFrontend(testFrontendDist{})) + return r +} + +func writeFrontendFallbackTestFile(t *testing.T, name, content string) { + t.Helper() + if err := os.MkdirAll(filepath.Dir(name), 0o755); err != nil { + t.Fatalf("create fixture directory: %v", err) + } + if err := os.WriteFile(name, []byte(content), 0o644); err != nil { + t.Fatalf("write fixture file: %v", err) + } +} + +type testFrontendDist struct{} + +func (testFrontendDist) Open(string) (fs.File, error) { + return nil, fs.ErrNotExist +} + +func performFrontendFallbackRequest(t *testing.T, router *gin.Engine, target string) *httptest.ResponseRecorder { + t.Helper() + w := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, target, nil) + router.ServeHTTP(w, req) + return w +} + +func TestFallbackToFrontendBlocksDashboardTraversal(t *testing.T) { + t.Chdir(t.TempDir()) + router := newFrontendFallbackTestRouter(t) + + tests := []string{ + "/dashboard../data/config.yaml", + "/dashboard%2e%2e/data/config.yaml", + "/dashboard%2e%2e%2fdata%2fconfig.yaml", + "/dashboard/../data/config.yaml", + "/dashboard/%2e%2e/data/config.yaml", + "/dashboard../assets/app.js", + } + + for _, target := range tests { + t.Run(target, func(t *testing.T) { + w := performFrontendFallbackRequest(t, router, target) + body := w.Body.String() + if strings.Contains(body, "traversal-secret") || strings.Contains(body, "jwt_secret_key") || strings.Contains(body, "admin asset") { + t.Fatalf("%s leaked protected content with status %d: %q", target, w.Code, body) + } + }) + } +} + +func TestFallbackToFrontendBlocksUserTraversal(t *testing.T) { + t.Chdir(t.TempDir()) + router := newFrontendFallbackTestRouter(t) + + tests := []string{ + "/../data/config.yaml", + "/%2e%2e/data/config.yaml", + "/%2e%2e%2fdata%2fconfig.yaml", + "/../admin-dist/assets/app.js", + "/%2e%2e/admin-dist/assets/app.js", + } + + for _, target := range tests { + t.Run(target, func(t *testing.T) { + w := performFrontendFallbackRequest(t, router, target) + body := w.Body.String() + if strings.Contains(body, "traversal-secret") || strings.Contains(body, "jwt_secret_key") || strings.Contains(body, "admin asset") { + t.Fatalf("%s leaked protected content with status %d: %q", target, w.Code, body) + } + }) + } +} + +func TestFallbackToFrontendPreservesDashboardRoutes(t *testing.T) { + t.Chdir(t.TempDir()) + router := newFrontendFallbackTestRouter(t) + + w := performFrontendFallbackRequest(t, router, "/dashboard") + if w.Code != http.StatusMovedPermanently { + t.Fatalf("/dashboard status = %d, want %d", w.Code, http.StatusMovedPermanently) + } + if location := w.Header().Get("Location"); location != "/dashboard/" { + t.Fatalf("/dashboard Location = %q, want /dashboard/", location) + } + + w = performFrontendFallbackRequest(t, router, "/dashboard/") + if w.Code != http.StatusOK || !strings.Contains(w.Body.String(), "admin index") { + t.Fatalf("/dashboard/ status = %d body = %q, want admin index", w.Code, w.Body.String()) + } + + w = performFrontendFallbackRequest(t, router, "/dashboard/assets/app.js") + if w.Code != http.StatusOK || !strings.Contains(w.Body.String(), "admin asset") { + t.Fatalf("/dashboard/assets/app.js status = %d body = %q, want admin asset", w.Code, w.Body.String()) + } +} diff --git a/cmd/dashboard/controller/permission_matrix_test.go b/cmd/dashboard/controller/permission_matrix_test.go index a3c180e4..68f0f047 100644 --- a/cmd/dashboard/controller/permission_matrix_test.go +++ b/cmd/dashboard/controller/permission_matrix_test.go @@ -478,6 +478,21 @@ func TestBatchMoveServerAllowsAdminCrossUser(t *testing.T) { assert.NoError(t, err) } +func TestBatchMoveServerMasksForeignServerIDsForMembers(t *testing.T) { + ctx := newMemberValidationContext(t) + assert.NoError(t, singleton.DB.Create(&model.Server{Common: model.Common{ID: 1, UserID: 1}, Name: "foreign", UUID: "foreign-server"}).Error) + singleton.ServerShared = singleton.NewServerClass() + + ctx.Request = httptest.NewRequest(http.MethodPost, "/batch-move/server", strings.NewReader(`{"ids":[1,9999],"to_user":200}`)) + ctx.Request.Header.Set("Content-Type", "application/json") + got, err := batchMoveServer(ctx) + + assert.NoError(t, err) + assert.Len(t, got, 2) + assert.Equal(t, model.BatchMoveServerResultServerNotFound, got[0].Status) + assert.Equal(t, model.BatchMoveServerResultServerNotFound, got[1].Status) +} + func TestNATRejectsUnknownServerID(t *testing.T) { ctx := newMemberValidationContext(t) ctx.Request = httptest.NewRequest(http.MethodPost, "/nat", strings.NewReader(`{"name":"x","domain":"x.example","host":"127.0.0.1:80","server_id":9999}`)) diff --git a/cmd/dashboard/controller/server.go b/cmd/dashboard/controller/server.go index 4f3e015a..a56cb586 100644 --- a/cmd/dashboard/controller/server.go +++ b/cmd/dashboard/controller/server.go @@ -1,6 +1,7 @@ package controller import ( + "errors" "slices" "strconv" "sync" @@ -153,6 +154,14 @@ func batchDeleteServer(c *gin.Context) (any, error) { singleton.DB.Unscoped().Delete(&model.Transfer{}, "server_id in (?)", servers) singleton.AlertsLock.Unlock() + // Cancel any in-flight transfers BEFORE the in-memory ServerShared + // entry is dropped: the order shortens the window in which a + // concurrent Retry/Register could install a fresh pending entry for + // the same serverID and have it wiped by the cleanup. The + // transferID-guarded delete inside OnServersDeleted is the + // authoritative protection against that race; the ordering here is + // belt and braces. + singleton.ServerTransferShared.OnServersDeleted(servers) singleton.ServerShared.Delete(servers) return nil, nil } @@ -187,8 +196,8 @@ func forceUpdateServer(c *gin.Context) (*model.ServerTaskResponse, error) { forceUpdateResp.Offline = append(forceUpdateResp.Offline, sid) continue } - if server.TaskStream != nil { - if err := server.TaskStream.Send(&pb.Task{ + if stream := server.GetTaskStream(); stream != nil { + if err := stream.Send(&pb.Task{ Type: model.TaskTypeUpgrade, }); err != nil { forceUpdateResp.Failure = append(forceUpdateResp.Failure, sid) @@ -220,7 +229,11 @@ func getServerConfig(c *gin.Context) (string, error) { } s, ok := singleton.ServerShared.Get(id) - if !ok || s.TaskStream == nil { + if !ok { + return "", nil + } + stream := s.GetTaskStream() + if stream == nil { return "", nil } @@ -228,7 +241,7 @@ func getServerConfig(c *gin.Context) (string, error) { return "", singleton.Localizer.ErrorT("permission denied") } - if err := s.TaskStream.Send(&pb.Task{ + if err := stream.Send(&pb.Task{ Type: model.TaskTypeReportConfig, }); err != nil { return "", err @@ -276,7 +289,7 @@ func setServerConfig(c *gin.Context) (*model.ServerTaskResponse, error) { if !s.HasPermission(c) { return nil, singleton.Localizer.ErrorT("permission denied") } - if s.TaskStream == nil { + if s.GetTaskStream() == nil { resp.Offline = append(resp.Offline, s.ID) continue } @@ -300,7 +313,14 @@ func setServerConfig(c *gin.Context) (*model.ServerTaskResponse, error) { Type: model.TaskTypeApplyConfig, Data: configForm.Config, } - if err := s.TaskStream.Send(task); err != nil { + stream := s.GetTaskStream() + if stream == nil { + respMu.Lock() + resp.Offline = append(resp.Offline, s.ID) + respMu.Unlock() + continue + } + if err := stream.Send(task); err != nil { respMu.Lock() resp.Failure = append(resp.Failure, s.ID) respMu.Unlock() @@ -321,23 +341,29 @@ func setServerConfig(c *gin.Context) (*model.ServerTaskResponse, error) { // @Summary Batch move servers to other user // @Security BearerAuth // @Schemes -// @Description Batch move servers to other user +// @Description Initiates one ServerTransfer per requested server and returns a +// @Description per-server result. The old behaviour flipped Server.UserID in +// @Description a single SQL UPDATE without telling the agent, so the agent +// @Description kept presenting its old AgentSecret — which now belonged to a +// @Description different user — and authorizeAgentForUUID dropped it. The +// @Description current flow writes a Pending ServerTransfer row and flips +// @Description Server.UserID to the target owner immediately; that row keeps +// @Description the old owner's AgentSecret acceptable for this UUID until the +// @Description agent reconnects under the new secret (MarkVerified clears the +// @Description pending window) or the transfer Cancel/Fail/Timeout-out and +// @Description reverts Server.UserID to the source owner. // @Tags auth required // @Accept json // @Param request body model.BatchMoveServerForm true "BatchMoveServerForm" // @Produce json -// @Success 200 {object} model.CommonResponse[any] +// @Success 200 {object} model.CommonResponse[[]model.BatchMoveServerResult] // @Router /batch-move/server [post] -func batchMoveServer(c *gin.Context) (any, error) { +func batchMoveServer(c *gin.Context) ([]model.BatchMoveServerResult, error) { var moveForm model.BatchMoveServerForm if err := c.ShouldBindJSON(&moveForm); err != nil { return nil, err } - if !singleton.ServerShared.CheckPermission(c, slices.Values(moveForm.Ids)) { - return nil, singleton.Localizer.ErrorT("permission denied") - } - if moveForm.ToUser == 0 { return nil, singleton.Localizer.ErrorT("user id is required") } @@ -347,35 +373,81 @@ func batchMoveServer(c *gin.Context) (any, error) { } singleton.UserLock.RLock() - defer singleton.UserLock.RUnlock() - if _, ok := singleton.UserInfoMap[moveForm.ToUser]; !ok { + _, toUserExists := singleton.UserInfoMap[moveForm.ToUser] + singleton.UserLock.RUnlock() + if !toUserExists { return nil, singleton.Localizer.ErrorT("user id %d does not exist", moveForm.ToUser) } - err := singleton.DB.Transaction(func(tx *gorm.DB) error { - if err := tx.Model(&model.Server{}).Where("id in (?)", moveForm.Ids).Update("user_id", moveForm.ToUser).Error; err != nil { - return err - } - return nil - }) + results := make([]model.BatchMoveServerResult, 0, len(moveForm.Ids)) + uid := getUid(c) + isAdmin := callerIsAdmin(c) - if err != nil { - return nil, newGormError("%v", err) - } + for _, sid := range moveForm.Ids { + res := model.BatchMoveServerResult{ServerID: sid} - idsMap := make(map[uint64]bool) - for _, id := range moveForm.Ids { - idsMap[id] = true - } - - for _, s := range singleton.ServerShared.Range { - if s == nil || !idsMap[s.ID] { + srv, ok := singleton.ServerShared.Get(sid) + if !ok || srv == nil { + res.Status = model.BatchMoveServerResultServerNotFound + results = append(results, res) continue } - s.UserID = moveForm.ToUser + + // Per-server permission: admin or current owner. We do NOT use the + // bulk CheckPermission because we want a partial-success response + // rather than rejecting the whole batch on the first unauthorized id. + // + // 必须走 GetUserID() 而不是裸读 srv.UserID — ServerTransfer.Register + // 和 revertTransition 会通过 atomic.StoreUint64 改写当前 Server.UserID + // 以反映新所有者。batchMoveServer 与 transfer 流程并发时(典型场景:两 + // 个 operator 几乎同时发起 move),裸读会与 SetUserID 形成 data race, + // 且可能在 transfer 切换瞬间读到过期值并据此做权限/同所有者/fromUser + // 判断。 + currentOwner := srv.GetUserID() + if !isAdmin && currentOwner != uid { + // Match the unknown-id response for members. A distinct + // permission_denied result lets callers enumerate foreign server IDs. + res.Status = model.BatchMoveServerResultServerNotFound + results = append(results, res) + continue + } + + if currentOwner == moveForm.ToUser { + res.Status = model.BatchMoveServerResultSameOwner + results = append(results, res) + continue + } + + // One active ServerTransfer per server. InitiateExclusive serializes + // the HasPending guard, the DB transaction, and the in-memory + // Register under a per-server claim so two concurrent operators + // can't both observe "no pending", both insert, and silently end up + // with two Pending rows for the same server. + fromUser := currentOwner + created, err := singleton.ServerTransferShared.InitiateExclusive(sid, fromUser, moveForm.ToUser, uid) + if err != nil { + switch { + case errors.Is(err, singleton.ErrServerAlreadyTransferring): + res.Status = model.BatchMoveServerResultAlreadyTransferring + case errors.Is(err, singleton.ErrAgentTooOldForTransfer): + res.Status = model.BatchMoveServerResultAgentTooOld + res.Error = err.Error() + default: + res.Status = model.BatchMoveServerResultServerNotFound + res.Error = err.Error() + } + results = append(results, res) + continue + } + + singleton.ServerTransferShared.PushIfOnline(created) + + res.Status = model.BatchMoveServerResultPending + res.TransferID = created.ID + results = append(results, res) } - return nil, nil + return results, nil } var serverMetricMap = map[string]tsdb.MetricType{ diff --git a/cmd/dashboard/controller/stream_ownership_test.go b/cmd/dashboard/controller/stream_ownership_test.go index 9592742a..55291fe0 100644 --- a/cmd/dashboard/controller/stream_ownership_test.go +++ b/cmd/dashboard/controller/stream_ownership_test.go @@ -141,6 +141,39 @@ func TestJWTInitParamsPinsAlgorithmAndSameSite(t *testing.T) { } } +// MEDIUM security: an IOStream session created by the old owner (terminal, +// file-manager, NAT) must be torn down when the server's ownership rotates +// — Register on Initiate, revertTransition on Cancel/Fail/Timeout, and +// OnServersDeleted on delete. Otherwise the old owner keeps an open +// websocket attached to a server they no longer own, which is effectively +// post-transfer RCE / file-read. +func TestServerTransferTransitionRevokesActiveIOStreams(t *testing.T) { + gin.SetMode(gin.TestMode) + ensureLocalizerForStreamTests(t) + rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler() + originalHook := singleton.ServerTransferStreamRevocationHook + singleton.ServerTransferStreamRevocationHook = rpc.NezhaHandlerSingleton.RevokeStreamsForServer + defer func() { + singleton.ServerTransferStreamRevocationHook = originalHook + }() + + rpc.NezhaHandlerSingleton.CreateStream("term-server-1", 100, 1) + rpc.NezhaHandlerSingleton.CreateStream("fm-server-1", 100, 1) + rpc.NezhaHandlerSingleton.CreateStream("term-server-2", 100, 2) + + singleton.ServerTransferRevokeStreamsForServer(1) + + if _, exists := rpc.NezhaHandlerSingleton.StreamOwnership("term-server-1"); exists { + t.Fatal("terminal stream for transferred server 1 must be revoked on ownership rotation") + } + if _, exists := rpc.NezhaHandlerSingleton.StreamOwnership("fm-server-1"); exists { + t.Fatal("file-manager stream for transferred server 1 must be revoked on ownership rotation") + } + if _, exists := rpc.NezhaHandlerSingleton.StreamOwnership("term-server-2"); !exists { + t.Fatal("unrelated server's stream must NOT be revoked") + } +} + // nz-o2s carries the OAuth2 state binding that authenticates the callback. // The frontend never reads it, so HttpOnly is safe to enable and shuts the // door on XSS attempting to steal the state. diff --git a/cmd/dashboard/controller/terminal.go b/cmd/dashboard/controller/terminal.go index 728cbee6..521795db 100644 --- a/cmd/dashboard/controller/terminal.go +++ b/cmd/dashboard/controller/terminal.go @@ -31,7 +31,11 @@ func createTerminal(c *gin.Context) (*model.CreateTerminalResponse, error) { } server, _ := singleton.ServerShared.Get(createTerminalReq.ServerID) - if server == nil || server.TaskStream == nil { + if server == nil { + return nil, singleton.Localizer.ErrorT("server not found or not connected") + } + stream := server.GetTaskStream() + if stream == nil { return nil, singleton.Localizer.ErrorT("server not found or not connected") } @@ -49,7 +53,7 @@ func createTerminal(c *gin.Context) (*model.CreateTerminalResponse, error) { terminalData, _ := json.Marshal(&model.TerminalTask{ StreamID: streamId, }) - if err := server.TaskStream.Send(&proto.Task{ + if err := stream.Send(&proto.Task{ Type: model.TaskTypeTerminalGRPC, Data: string(terminalData), }); err != nil { diff --git a/cmd/dashboard/controller/transfer.go b/cmd/dashboard/controller/transfer.go new file mode 100644 index 00000000..b3fbc6e5 --- /dev/null +++ b/cmd/dashboard/controller/transfer.go @@ -0,0 +1,217 @@ +package controller + +import ( + "strconv" + "time" + + "github.com/gin-gonic/gin" + "github.com/goccy/go-json" + "github.com/gorilla/websocket" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/service/singleton" +) + +// List server transfers +// @Summary List server transfers +// @Security BearerAuth +// @Schemes +// @Description Returns transfers visible to the caller. Admin sees all; a +// @Description member sees rows where they are FromUserID, ToUserID, or +// @Description InitiatorID. The same predicate is enforced both at the SQL +// @Description level (this handler) and by the listHandler post-filter +// @Description (ServerTransfer.HasPermission) — defence in depth. +// @Tags auth required +// @Produce json +// @Success 200 {object} model.CommonResponse[[]model.ServerTransfer] +// @Router /transfer [get] +func listServerTransfer(c *gin.Context) ([]*model.ServerTransfer, error) { + q := singleton.DB.Order("id DESC") + // ServerTransfer is the only listX endpoint that hits the DB — the others + // all serve in-memory caches — and it is an append-only audit log. Without + // this SQL-side filter, every member's page load scans the entire historical + // population only to have the listHandler post-filter throw most of it + // away. As the table grows (a single transfer per server-move adds a row + // forever) this degrades from cheap to dashboard-blocking. Mirror the + // HasPermission predicate at the WHERE clause for non-admins. The post-filter + // still runs unconditionally as a defence-in-depth guard. + if !callerIsAdmin(c) { + uid := getUid(c) + q = q.Where("from_user_id = ? OR to_user_id = ? OR initiator_id = ?", uid, uid, uid) + } + var transfers []*model.ServerTransfer + if err := q.Find(&transfers).Error; err != nil { + return nil, newGormError("%v", err) + } + return transfers, nil +} + +// Cancel server transfer +// @Summary Cancel server transfer +// @Security BearerAuth +// @Schemes +// @Description Cancels a Pending transfer and reverts Server.UserID back to +// @Description FromUserID. Only admin or the original FromUserID may cancel +// @Description (the new owner cannot — that would be a denial primitive +// @Description against a server they don't own yet). No-op if the transfer is +// @Description already terminal. +// @Tags auth required +// @Param id path uint true "Transfer ID" +// @Produce json +// @Success 200 {object} model.CommonResponse[model.ServerTransfer] +// @Router /transfer/{id}/cancel [post] +func cancelServerTransfer(c *gin.Context) (*model.ServerTransfer, error) { + tid, err := strconv.ParseUint(c.Param("id"), 10, 64) + if err != nil { + return nil, err + } + + // Avoid leaking transfer-row existence via response shape. Admin can + // look up any row; a member can only look up rows where they are the + // FromUserID. Both "row does not exist" and "row exists but caller is + // not FromUserID" must surface identically as permission denied. + q := singleton.DB + if !callerIsAdmin(c) { + q = q.Where("from_user_id = ?", getUid(c)) + } + var t model.ServerTransfer + if err := q.First(&t, tid).Error; err != nil { + return nil, singleton.Localizer.ErrorT("permission denied") + } + + updated, err := singleton.ServerTransferShared.Cancel(tid) + if err != nil { + return nil, err + } + if updated == nil { + // Already terminal — return current state so the UI can refresh. + return &t, nil + } + return updated, nil +} + +// Retry server transfer +// @Summary Retry server transfer +// @Security BearerAuth +// @Schemes +// @Description Creates a fresh Pending transfer with the same From/To as a +// @Description terminal (Failed/Timeout/Cancelled) transfer. The previous row +// @Description is left intact for audit; this returns the new row. Admin-only: +// @Description non-admin transfer semantics are enforced by batchMoveServer's +// @Description "ToUser == self" rule, so allowing any historical From/To/ +// @Description Initiator to retry would reintroduce the give-away path. Non- +// @Description admins receive permission denied before the transfer row is +// @Description read so the response cannot enumerate transfer ids. +// @Tags auth required +// @Param id path uint true "Transfer ID" +// @Produce json +// @Success 200 {object} model.CommonResponse[model.ServerTransfer] +// @Router /transfer/{id}/retry [post] +func retryServerTransfer(c *gin.Context) (*model.ServerTransfer, error) { + // Retry is admin-only. Non-admin transfer semantics are enforced by + // batchMoveServer's "ToUser == self" rule, which means a member can only + // receive a server, never give one away. Allowing a member to retry a + // historical row would reintroduce the give-away path: any prev.ToUserID + // on file becomes a one-click bypass of that policy. Members who want + // the server moved elsewhere ask an admin or use batch-move to pull it + // onto themselves. + // + // Refuse non-admins before reading the row so the response cannot be + // used to enumerate which transfer ids exist. + if !callerIsAdmin(c) { + return nil, singleton.Localizer.ErrorT("permission denied") + } + + tid, err := strconv.ParseUint(c.Param("id"), 10, 64) + if err != nil { + return nil, err + } + + var prev model.ServerTransfer + if err := singleton.DB.First(&prev, tid).Error; err != nil { + return nil, newGormError("%v", err) + } + + return singleton.ServerTransferShared.Retry(&prev, getUid(c)) +} + +// transferStreamWriteTimeout caps a single WriteMessage so a stuck or +// half-open client cannot block the broker fan-out forever — once exceeded +// the connection is considered dead and dropped. Matches the cadence of the +// keepalive ping below. +const transferStreamWriteTimeout = 10 * time.Second + +// transferStreamPingInterval is how often we send a ping to keep the +// connection alive through aggressive proxies. Independent of the event +// stream, so silent transfers still keep the socket warm. +const transferStreamPingInterval = 30 * time.Second + +// Websocket server transfer stream +// @Summary Websocket server transfer stream +// @Security BearerAuth +// @Schemes +// @Description Pushes ServerTransfer state transitions (Pending → Verified / +// @Description Failed / Timeout / Cancelled) to the dashboard so the UI can +// @Description react without polling. Each frame is a single JSON-encoded +// @Description ServerTransfer. Subscribers see only transfers visible to +// @Description them (ServerTransfer.HasPermission). +// @tags common +// @Produce json +// @Success 200 {object} model.ServerTransfer +// @Router /ws/transfer [get] +func transferStream(c *gin.Context) (any, error) { + conn, err := upgrader.Upgrade(c.Writer, c.Request, nil) + if err != nil { + return nil, newWsError("%v", err) + } + defer conn.Close() + + subID, ch := singleton.ServerTransferShared.Subscribe() + defer singleton.ServerTransferShared.Unsubscribe(subID) + + // Pings keep the socket warm even when the broker is quiet. Without this + // a long idle period followed by a transfer event would race against + // upstream proxy idle-timeouts that may have already closed the conn. + ping := time.NewTicker(transferStreamPingInterval) + defer ping.Stop() + + // Reader goroutine: needed only to surface client disconnects through + // SetReadDeadline / ReadMessage. We never expect inbound payloads. + closed := make(chan struct{}) + go func() { + defer close(closed) + for { + if _, _, err := conn.NextReader(); err != nil { + return + } + } + }() + + for { + select { + case <-closed: + return nil, newWsError("") + case <-ping.C: + if err := conn.WriteControl(websocket.PingMessage, nil, time.Now().Add(transferStreamWriteTimeout)); err != nil { + return nil, newWsError("%v", err) + } + case t, ok := <-ch: + if !ok { + return nil, newWsError("") + } + if !t.HasPermission(c) { + continue + } + payload, err := json.Marshal(t) + if err != nil { + continue + } + if err := conn.SetWriteDeadline(time.Now().Add(transferStreamWriteTimeout)); err != nil { + return nil, newWsError("%v", err) + } + if err := conn.WriteMessage(websocket.TextMessage, payload); err != nil { + return nil, newWsError("%v", err) + } + } + } +} diff --git a/cmd/dashboard/controller/transfer_list_filter_test.go b/cmd/dashboard/controller/transfer_list_filter_test.go new file mode 100644 index 00000000..a3433733 --- /dev/null +++ b/cmd/dashboard/controller/transfer_list_filter_test.go @@ -0,0 +1,231 @@ +package controller + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/driver/sqlite" + "gorm.io/gorm" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/pkg/i18n" + "github.com/nezhahq/nezha/service/singleton" +) + +// listServerTransfer originally did `SELECT * FROM server_transfers ORDER BY +// id DESC` and relied on the listHandler post-filter (HasPermission) to drop +// rows the caller can't see. Functionally correct, but ServerTransfer is the +// only listX endpoint that hits the DB (the others serve in-memory caches), +// and it's an append-only audit table — every page load by a member who has +// participated in two transfers triggers a full table scan over the entire +// historical population. That cost is silent until the table is big and the +// dashboard slows for everyone at once. The fix pushes the same predicate +// HasPermission encodes down into the WHERE clause for non-admin callers. +// +// This test pins down the behavioural contract: regardless of the optimisation, +// a member must see only their own rows. It is intentionally written against +// the same response shape as the production handler so a regression in either +// the SQL filter OR the post-filter would fail it. +func TestListServerTransferReturnsOnlyCallerVisibleRowsForMember(t *testing.T) { + cleanup := setupListServerTransferFixture(t) + defer cleanup() + + // alice=100, bob=200, charlie=300. Seed five rows covering every position + // a member could occupy plus one row the member must NEVER see. + aliceFrom := seedTerminalTransfer(t, 1, 100, 200, 100) + aliceTo := seedTerminalTransfer(t, 2, 300, 100, 300) + aliceInitiated := seedTerminalTransfer(t, 3, 200, 300, 100) + bobAlone := seedTerminalTransfer(t, 4, 200, 300, 200) + charlieAlone := seedTerminalTransfer(t, 5, 300, 200, 300) + + ids := callListServerTransfer(t, 100, model.RoleMember) + + assert.ElementsMatch(t, + []uint64{aliceFrom, aliceTo, aliceInitiated}, + ids, + "member must see exactly the rows where they are From/To/Initiator", + ) + for _, forbidden := range []uint64{bobAlone, charlieAlone} { + assert.NotContains(t, ids, forbidden, "member must never see a row they do not participate in") + } +} + +// Admin sees every row. This pins down that the SQL-level filter is gated on +// role and is not accidentally applied to admins (which would be a regression +// in the other direction). +func TestListServerTransferReturnsEveryRowForAdmin(t *testing.T) { + cleanup := setupListServerTransferFixture(t) + defer cleanup() + + a := seedTerminalTransfer(t, 1, 100, 200, 100) + b := seedTerminalTransfer(t, 2, 200, 300, 200) + c := seedTerminalTransfer(t, 3, 300, 100, 300) + + ids := callListServerTransfer(t, 999, model.RoleAdmin) + + assert.ElementsMatch(t, []uint64{a, b, c}, ids, "admin sees every row") +} + +// Pin down the optimisation directly: the SELECT that hits the audit table +// for a non-admin caller MUST include a per-user WHERE clause. The behavioural +// tests above would pass even if the SQL stayed `SELECT *` (the post-filter +// hides forbidden rows), so they cannot regress-detect the perf fix going +// away. This test captures the executed SQL and asserts the filter is pushed +// down to the database. +// +// Why pinning the optimisation matters: ServerTransfer is the only listX +// endpoint that hits the DB (the rest serve in-memory caches) and it grows +// unbounded as an audit log. Without SQL-side filtering, every member's page +// load scans the entire historical table. A future refactor that drops the +// per-caller WHERE clause would not be caught by any behavioural assertion — +// hence this explicit guard. +func TestListServerTransferPushesPerCallerFilterIntoSQLForMember(t *testing.T) { + cleanup, captured := setupListServerTransferFixtureWithSQLCapture(t) + defer cleanup() + + seedTerminalTransfer(t, 1, 100, 200, 100) + seedTerminalTransfer(t, 2, 200, 300, 200) + + _ = callListServerTransfer(t, 100, model.RoleMember) + + stmt := findSelectAgainstServerTransfers(captured.Snapshot()) + require.NotEmpty(t, stmt, "expected a SELECT against server_transfers to be issued") + low := strings.ToLower(stmt) + require.Contains(t, low, "where", "non-admin list must apply a per-caller WHERE filter at the SQL level") + for _, col := range []string{"from_user_id", "to_user_id", "initiator_id"} { + require.Contains(t, low, col, "WHERE clause must filter on %s", col) + } +} + +// The admin path must NOT push the per-user filter — admins see everything. +// Without this assertion a refactor that always applies the filter would +// silently hide cross-tenant rows from admins (an availability regression of +// the admin observability surface). +func TestListServerTransferOmitsPerCallerFilterForAdmin(t *testing.T) { + cleanup, captured := setupListServerTransferFixtureWithSQLCapture(t) + defer cleanup() + + seedTerminalTransfer(t, 1, 100, 200, 100) + + _ = callListServerTransfer(t, 999, model.RoleAdmin) + + stmt := findSelectAgainstServerTransfers(captured.Snapshot()) + require.NotEmpty(t, stmt, "expected a SELECT against server_transfers to be issued") + low := strings.ToLower(stmt) + for _, col := range []string{"from_user_id", "to_user_id", "initiator_id"} { + require.NotContains(t, low, col, "admin list must NOT filter by %s — admins observe all transfers", col) + } +} + +func setupListServerTransferFixture(t *testing.T) func() { + t.Helper() + if singleton.Localizer == nil { + singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations) + } + originalDB := singleton.DB + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + assert.NoError(t, err) + assert.NoError(t, db.AutoMigrate(&model.Server{}, &model.ServerTransfer{})) + singleton.DB = db + return func() { singleton.DB = originalDB } +} + +// sqlCapture records every SQL statement gorm executes against the test DB. +// Used by the optimisation guards above to assert that listServerTransfer +// actually pushes its per-caller filter into the WHERE clause. +type sqlCapture struct { + mu sync.Mutex + stmts []string +} + +func (s *sqlCapture) record(stmt string) { + s.mu.Lock() + defer s.mu.Unlock() + s.stmts = append(s.stmts, stmt) +} + +func (s *sqlCapture) Snapshot() []string { + s.mu.Lock() + defer s.mu.Unlock() + out := make([]string, len(s.stmts)) + copy(out, s.stmts) + return out +} + +func setupListServerTransferFixtureWithSQLCapture(t *testing.T) (func(), *sqlCapture) { + t.Helper() + cleanup := setupListServerTransferFixture(t) + + cap := &sqlCapture{} + err := singleton.DB.Callback().Query().After("gorm:query").Register("test:capture_sql", func(tx *gorm.DB) { + cap.record(tx.Statement.SQL.String()) + }) + require.NoError(t, err) + + return func() { + _ = singleton.DB.Callback().Query().Remove("test:capture_sql") + cleanup() + }, cap +} + +// findSelectAgainstServerTransfers returns the first captured SELECT whose +// FROM clause is `server_transfers`. We don't care about ordering callbacks +// or AutoMigrate scaffolding queries — only the handler's own SELECT. +func findSelectAgainstServerTransfers(stmts []string) string { + for _, s := range stmts { + low := strings.ToLower(s) + if strings.HasPrefix(strings.TrimSpace(low), "select") && strings.Contains(low, "server_transfers") { + return s + } + } + return "" +} + +func seedTerminalTransfer(t *testing.T, serverID, fromUserID, toUserID, initiatorID uint64) uint64 { + t.Helper() + tr := &model.ServerTransfer{ + ServerID: serverID, + FromUserID: fromUserID, + ToUserID: toUserID, + InitiatorID: initiatorID, + Status: model.ServerTransferStatusVerified, + } + assert.NoError(t, singleton.DB.Create(tr).Error) + return tr.ID +} + +func callListServerTransfer(t *testing.T, callerID uint64, role model.Role) []uint64 { + t.Helper() + r := gin.New() + r.Use(func(c *gin.Context) { + setAuthUser(c, callerID, role) + c.Next() + }) + r.GET("/transfer", listHandler(listServerTransfer)) + + w := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/transfer", nil) + r.ServeHTTP(w, req) + assert.Equal(t, http.StatusOK, w.Code) + + var resp struct { + Success bool `json:"success"` + Data []*model.ServerTransfer `json:"data"` + Error string `json:"error"` + } + assert.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp)) + assert.True(t, resp.Success, "list call must succeed: %s", resp.Error) + + ids := make([]uint64, 0, len(resp.Data)) + for _, tr := range resp.Data { + ids = append(ids, tr.ID) + } + return ids +} diff --git a/cmd/dashboard/controller/transfer_retry_authz_test.go b/cmd/dashboard/controller/transfer_retry_authz_test.go new file mode 100644 index 00000000..33bd299f --- /dev/null +++ b/cmd/dashboard/controller/transfer_retry_authz_test.go @@ -0,0 +1,219 @@ +package controller + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "strconv" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" + "gorm.io/driver/sqlite" + "gorm.io/gorm" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/pkg/i18n" + "github.com/nezhahq/nezha/service/singleton" +) + +// retryServerTransfer previously gated on prev.HasPermission(c), which honours +// the historical transfer row (FromUserID, ToUserID, InitiatorID). That lets +// any of those original parties re-initiate a transfer of the server long +// after ownership has moved on. Concretely: a stale "alice -> bob" Failed row +// stays visible to alice forever — even after she's transferred the server +// off to charlie — and the original endpoint would happily move it from +// charlie to bob without ever consulting the current owner. +// +// Authorization for an action that mutates the live server must use the +// live server, not a historical audit row. +func TestRetryServerTransferRejectsCallerWhoNoLongerOwnsServer(t *testing.T) { + cleanup := setupRetryServerTransferFixture(t) + defer cleanup() + + // Seed: server originally owned by user 100 (alice). Failed transfer to + // 200 (bob) is recorded but server ownership has since moved to 300 + // (charlie) — e.g. alice transferred elsewhere afterwards. Alice is no + // longer the owner, so retrying the stale row would be an unauthorized + // grab. + seedServer(t, 1, 300) + staleID := seedFailedTransfer(t, 1, 100 /*from*/, 200 /*to*/, 100 /*initiator*/) + + resp, status := callRetryServerTransfer(t, staleID, 100, model.RoleMember) + + assert.Equal(t, http.StatusOK, status) + assert.False(t, resp.Success, "alice no longer owns server 1; retry must be rejected") + assert.Contains(t, resp.Error, "permission denied") + + var s model.Server + assert.NoError(t, singleton.DB.First(&s, 1).Error) + assert.Equal(t, uint64(300), s.UserID, "rejected retry must not flip ownership") + + var count int64 + assert.NoError(t, singleton.DB.Model(&model.ServerTransfer{}).Where("status = ?", model.ServerTransferStatusPending).Count(&count).Error) + assert.Equal(t, int64(0), count, "rejected retry must not create a Pending row") +} + +// The historical ToUserID must also not be able to grab the server back via +// the stale row. Same root cause; this is the explicit assertion that the +// fix covers the To side, not just the From side. +func TestRetryServerTransferRejectsHistoricalTargetWhoNeverOwnedServer(t *testing.T) { + cleanup := setupRetryServerTransferFixture(t) + defer cleanup() + + seedServer(t, 1, 300) + staleID := seedFailedTransfer(t, 1, 100, 200, 100) + + resp, status := callRetryServerTransfer(t, staleID, 200, model.RoleMember) + + assert.Equal(t, http.StatusOK, status) + assert.False(t, resp.Success, "bob was the failed transfer's target; he never owned server 1") + assert.Contains(t, resp.Error, "permission denied") +} + +// Members never retry: batchMoveServer's "ToUser == self" policy means a +// member can only RECEIVE a server, not give one away. Retry of a failed +// "alice -> bob" by alice (member, current owner) is exactly the give-away +// case batch-move would refuse. Retry is admin-only. +func TestRetryServerTransferRejectsCurrentOwnerWhoIsMember(t *testing.T) { + cleanup := setupRetryServerTransferFixture(t) + defer cleanup() + + seedServer(t, 1, 100) + failedID := seedFailedTransfer(t, 1, 100, 200, 100) + + resp, status := callRetryServerTransfer(t, failedID, 100, model.RoleMember) + + assert.Equal(t, http.StatusOK, status) + assert.False(t, resp.Success, "member retry is forbidden — the give-away semantics bypass batchMoveServer's ToUser==self policy") + assert.Contains(t, resp.Error, "permission denied") +} + +// batchMoveServer enforces "non-admin caller may only move a server TO +// themselves" (controller/server.go: ToUser != getUid(c) returns permission +// denied). retryServerTransfer historically only checked the live owner +// and not the transfer's ToUserID, which let the current owner re-push the +// server to ANY historical ToUserID — bypassing the batch-move policy. +// +// Concretely: alice (member) currently owns server 1; she finds a Failed +// transfer whose ToUserID is bob and retries it. The server lands on bob +// even though batch-move would have refused "alice -> bob" from her. +func TestRetryServerTransferRejectsNonAdminPushingToForeignToUserID(t *testing.T) { + cleanup := setupRetryServerTransferFixture(t) + defer cleanup() + + seedServer(t, 1, 100) + staleID := seedFailedTransfer(t, 1, 100 /*from*/, 200 /*to*/, 100 /*initiator*/) + + resp, status := callRetryServerTransfer(t, staleID, 100, model.RoleMember) + + assert.Equal(t, http.StatusOK, status) + assert.False(t, resp.Success, "non-admin owner cannot push their server to a historical foreign ToUserID — that would bypass batchMoveServer's ToUser==self policy") + assert.Contains(t, resp.Error, "permission denied") + + var s model.Server + assert.NoError(t, singleton.DB.First(&s, 1).Error) + assert.Equal(t, uint64(100), s.UserID, "rejected retry must not flip ownership") + + var count int64 + assert.NoError(t, singleton.DB.Model(&model.ServerTransfer{}).Where("status = ?", model.ServerTransferStatusPending).Count(&count).Error) + assert.Equal(t, int64(0), count, "rejected retry must not create a Pending row") +} + +// Admins must always be able to retry — they are the last-resort recovery +// path when an operator-cancelled transfer needs to be re-pushed. +func TestRetryServerTransferAllowsAdmin(t *testing.T) { + cleanup := setupRetryServerTransferFixture(t) + defer cleanup() + + seedServer(t, 1, 300) + failedID := seedFailedTransfer(t, 1, 100, 200, 100) + + resp, status := callRetryServerTransfer(t, failedID, 999, model.RoleAdmin) + + assert.Equal(t, http.StatusOK, status) + assert.True(t, resp.Success, "admin must be able to retry any transfer: error=%s", resp.Error) +} + +func setupRetryServerTransferFixture(t *testing.T) func() { + t.Helper() + if singleton.Localizer == nil { + singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations) + } + originalDB := singleton.DB + originalShared := singleton.ServerShared + originalTransferShared := singleton.ServerTransferShared + originalUserMap := singleton.UserInfoMap + + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + assert.NoError(t, err) + assert.NoError(t, db.AutoMigrate(&model.Server{}, &model.ServerTransfer{})) + singleton.DB = db + singleton.ServerShared = singleton.NewServerClass() + singleton.UserInfoMap = map[uint64]model.UserInfo{ + 100: {Role: model.RoleMember, AgentSecret: "alice-secret"}, + 200: {Role: model.RoleMember, AgentSecret: "bob-secret"}, + 300: {Role: model.RoleMember, AgentSecret: "charlie-secret"}, + } + singleton.ServerTransferShared = singleton.NewServerTransferClass() + + return func() { + if singleton.ServerTransferShared != nil { + singleton.ServerTransferShared.Stop() + } + singleton.DB = originalDB + singleton.ServerShared = originalShared + singleton.ServerTransferShared = originalTransferShared + singleton.UserInfoMap = originalUserMap + } +} + +func seedServer(t *testing.T, id, ownerID uint64) { + t.Helper() + s := &model.Server{ + Common: model.Common{ID: id, UserID: ownerID}, + UUID: "uuid-" + strconv.FormatUint(id, 10), + Name: "seeded", + } + assert.NoError(t, singleton.DB.Create(s).Error) + model.InitServer(s) + singleton.ServerShared.Update(s, s.UUID) +} + +func seedFailedTransfer(t *testing.T, serverID, fromUserID, toUserID, initiatorID uint64) uint64 { + t.Helper() + tr := &model.ServerTransfer{ + ServerID: serverID, + FromUserID: fromUserID, + ToUserID: toUserID, + InitiatorID: initiatorID, + Status: model.ServerTransferStatusFailed, + LastError: "seeded", + } + assert.NoError(t, singleton.DB.Create(tr).Error) + return tr.ID +} + +func callRetryServerTransfer(t *testing.T, transferID, callerID uint64, role model.Role) (commonResponseShape, int) { + t.Helper() + r := gin.New() + r.Use(func(c *gin.Context) { + setAuthUser(c, callerID, role) + c.Next() + }) + r.POST("/transfer/:id/retry", commonHandler(retryServerTransfer)) + + w := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/transfer/"+strconv.FormatUint(transferID, 10)+"/retry", bytes.NewReader(nil)) + r.ServeHTTP(w, req) + + var resp commonResponseShape + assert.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp)) + return resp, w.Code +} + +type commonResponseShape struct { + Success bool `json:"success"` + Error string `json:"error"` +} diff --git a/cmd/dashboard/controller/ws.go b/cmd/dashboard/controller/ws.go index e8c52213..ffcc15d8 100644 --- a/cmd/dashboard/controller/ws.go +++ b/cmd/dashboard/controller/ws.go @@ -195,7 +195,7 @@ func getServerStat(withPublicNote bool, viewerUserID uint64, viewerIsAdmin bool) func filterServersForViewer(servers []*model.Server, viewerUserID uint64, viewerIsAdmin bool, withPublicNote bool) []model.StreamServer { out := make([]model.StreamServer, 0, len(servers)) for _, server := range servers { - isOwnerOrAdmin := viewerIsAdmin || (viewerUserID != 0 && server.UserID == viewerUserID) + isOwnerOrAdmin := viewerIsAdmin || (viewerUserID != 0 && server.GetUserID() == viewerUserID) if server.HideForGuest && !isOwnerOrAdmin { continue } diff --git a/cmd/dashboard/rpc/rpc.go b/cmd/dashboard/rpc/rpc.go index ab9491b6..bd8e73cf 100644 --- a/cmd/dashboard/rpc/rpc.go +++ b/cmd/dashboard/rpc/rpc.go @@ -24,6 +24,10 @@ import ( func ServeRPC() *grpc.Server { server := grpc.NewServer(grpc.ChainUnaryInterceptor(getRealIp, waf)) rpcService.NezhaHandlerSingleton = rpcService.NewNezhaHandler() + // Install the IOStream revocation hook so ServerTransferShared can tear + // down terminal/FM/NAT sessions held by the previous owner on every + // ownership rotation (Register/revertTransition/OnServersDeleted). + singleton.ServerTransferStreamRevocationHook = rpcService.NezhaHandlerSingleton.RevokeStreamsForServer proto.RegisterNezhaServiceServer(server, rpcService.NezhaHandlerSingleton) return server } @@ -89,22 +93,30 @@ func DispatchTask(serviceSentinelDispatchBus <-chan *model.Service) { } server, _ := singleton.ServerShared.Get(id) - if server == nil || server.TaskStream == nil { + if server == nil { + continue + } + stream := server.GetTaskStream() + if stream == nil { continue } if canSendTaskToServer(task, server) { - server.TaskStream.Send(task.PB()) + stream.Send(task.PB()) } } case model.ServiceCoverAll: for id, server := range singleton.ServerShared.Range { - if server == nil || server.TaskStream == nil || task.SkipServers[id] { + if server == nil || task.SkipServers[id] { + continue + } + stream := server.GetTaskStream() + if stream == nil { continue } if canSendTaskToServer(task, server) { - server.TaskStream.Send(task.PB()) + stream.Send(task.PB()) } } } @@ -115,17 +127,27 @@ func DispatchKeepalive() { singleton.CronShared.AddFunc("@every 20s", func() { list := singleton.ServerShared.GetSortedList() for _, s := range list { - if s == nil || s.TaskStream == nil { + if s == nil { continue } - s.TaskStream.Send(&proto.Task{Type: model.TaskTypeKeepalive}) + stream := s.GetTaskStream() + if stream == nil { + continue + } + stream.Send(&proto.Task{Type: model.TaskTypeKeepalive}) } }) } func ServeNAT(w http.ResponseWriter, r *http.Request, natConfig *model.NAT) { server, _ := singleton.ServerShared.Get(natConfig.ServerID) - if server == nil || server.TaskStream == nil { + if server == nil { + w.WriteHeader(http.StatusServiceUnavailable) + w.Write([]byte("server not found or not connected")) + return + } + stream := server.GetTaskStream() + if stream == nil { w.WriteHeader(http.StatusServiceUnavailable) w.Write([]byte("server not found or not connected")) return @@ -157,7 +179,7 @@ func ServeNAT(w http.ResponseWriter, r *http.Request, natConfig *model.NAT) { return } - if err := server.TaskStream.Send(&proto.Task{ + if err := stream.Send(&proto.Task{ Type: model.TaskTypeNAT, Data: string(taskData), }); err != nil { @@ -192,5 +214,5 @@ func canSendTaskToServer(task *model.Service, server *model.Server) bool { } singleton.UserLock.RUnlock() - return task.UserID == server.UserID || role.IsAdmin() + return task.UserID == server.GetUserID() || role.IsAdmin() } diff --git a/go.mod b/go.mod index 12d8e3bd..deaf8a0f 100644 --- a/go.mod +++ b/go.mod @@ -32,9 +32,9 @@ require ( github.com/swaggo/gin-swagger v1.6.1 github.com/swaggo/swag v1.16.6 github.com/tidwall/gjson v1.19.0 - golang.org/x/crypto v0.51.0 + golang.org/x/crypto v0.52.0 golang.org/x/exp v0.0.0-20260508232706-74f9aab9d74a - golang.org/x/net v0.54.0 + golang.org/x/net v0.55.0 golang.org/x/oauth2 v0.36.0 golang.org/x/sync v0.20.0 google.golang.org/grpc v1.81.1 @@ -109,7 +109,7 @@ require ( go.yaml.in/yaml/v3 v3.0.4 // indirect golang.org/x/arch v0.27.0 // indirect golang.org/x/mod v0.36.0 // indirect - golang.org/x/sys v0.44.0 // indirect + golang.org/x/sys v0.45.0 // indirect golang.org/x/text v0.37.0 // indirect golang.org/x/time v0.15.0 // indirect golang.org/x/tools v0.45.0 // indirect diff --git a/go.sum b/go.sum index bb50ddac..167a2ef6 100644 --- a/go.sum +++ b/go.sum @@ -246,8 +246,8 @@ golang.org/x/arch v0.27.0 h1:0WNVcR8u9yFz8j5FvdHpgwNp3FS5U4guYdzHwEiGjoU= golang.org/x/arch v0.27.0/go.mod h1:0X+GdSIP+kL5wPmpK7sdkEVTt2XoYP0cSjQSbZBwOi8= golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= golang.org/x/crypto v0.0.0-20210921155107-089bfa567519/go.mod h1:GvvjBRRGRdwPK5ydBHafDWAxML/pGHZbMvKqRZ5+Abc= -golang.org/x/crypto v0.51.0 h1:IBPXwPfKxY7cWQZ38ZCIRPI50YLeevDLlLnyC5wRGTI= -golang.org/x/crypto v0.51.0/go.mod h1:8AdwkbraGNABw2kOX6YFPs3WM22XqI4EXEd8g+x7Oc8= +golang.org/x/crypto v0.52.0 h1:RMs7fP2rXdep0CftQlK8Uf+kibLm7qkCcradZWYz988= +golang.org/x/crypto v0.52.0/go.mod h1:1QgfPxDqh0T2M/elOJtp9RvuR95kVjir0e6/BvEmGbc= golang.org/x/exp v0.0.0-20260508232706-74f9aab9d74a h1:+3jdDGGB8NGb1Zktc737jlt3/A5f6UlwSzmvqUuufxw= golang.org/x/exp v0.0.0-20260508232706-74f9aab9d74a/go.mod h1:d2fgXJLVs4dYDHUk5lwMIfzRzSrWCfGZb0ZqeLa/Vcw= golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4= @@ -257,8 +257,8 @@ golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLL golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= golang.org/x/net v0.0.0-20220722155237-a158d28d115b/go.mod h1:XRhObCWvk6IyKnWLug+ECip1KBveYUHfp+8e9klMJ9c= golang.org/x/net v0.7.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs= -golang.org/x/net v0.54.0 h1:2zJIZAxAHV/OHCDTCOHAYehQzLfSXuf/5SoL/Dv6w/w= -golang.org/x/net v0.54.0/go.mod h1:Sj4oj8jK6XmHpBZU/zWHw3BV3abl4Kvi+Ut7cQcY+cQ= +golang.org/x/net v0.55.0 h1:bcvxaJn3e1U6InsFWt1JUq1aSjnRxLzT2rtD2KfkDF8= +golang.org/x/net v0.55.0/go.mod h1:L5U2KuzuOe1lY7Z+aWVIKK6qEeJXnXV9yzGA+WCHJww= golang.org/x/oauth2 v0.36.0 h1:peZ/1z27fi9hUOFCAZaHyrpWG5lwe0RJEEEeH0ThlIs= golang.org/x/oauth2 v0.36.0/go.mod h1:YDBUJMTkDnJS+A4BP4eZBjCqtokkg1hODuPjwiGPO7Q= golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= @@ -271,8 +271,8 @@ golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBc golang.org/x/sys v0.0.0-20220520151302-bc2c85ada10a/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220722155257-8c9f86f7a55f/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.5.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.44.0 h1:ildZl3J4uzeKP07r2F++Op7E9B29JRUy+a27EibtBTQ= -golang.org/x/sys v0.44.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/sys v0.45.0 h1:dO4czNzziLiiXplLQgBCEpCvXQ3dnkn0SdaZSYdQ+FY= +golang.org/x/sys v0.45.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8= golang.org/x/term v0.5.0/go.mod h1:jMB1sMXY+tzblOD4FWmEbocvup2/aLOaQEp7JmGp78k= diff --git a/model/common.go b/model/common.go index 321f815b..3d165f94 100644 --- a/model/common.go +++ b/model/common.go @@ -6,6 +6,7 @@ import ( "slices" "strconv" "strings" + "sync/atomic" "time" "github.com/gin-gonic/gin" @@ -37,8 +38,21 @@ func (c *Common) GetID() uint64 { return c.ID } +// GetUserID 原子读取所属用户 ID。Server.UserID 会在 ServerTransfer 的 +// Register/revertTransition 流程里被实时改写以反映新所有者,同时 auth +// 热路径在每次 agent RPC 都会读它。任何并发读必须走 atomic,否则与 SetUserID +// 一起会被 go race detector 识别为 data race(见 +// TestServerUserIDConcurrentAccessIsRaceFree)。 func (c *Common) GetUserID() uint64 { - return c.UserID + return atomic.LoadUint64(&c.UserID) +} + +// SetUserID 原子改写所属用户 ID。仅在「server 已经在 in-memory cache 里」 +// 的写入路径(ServerTransfer.Register / revertTransition)需要用 atomic +// 保证可见性;普通 GORM AfterFind / Create 因为没有并发读所以可以直接赋 +// 值。配合 GetUserID 形成 atomic-only 的并发访问协议。 +func (c *Common) SetUserID(uid uint64) { + atomic.StoreUint64(&c.UserID, uid) } func (c *Common) HasPermission(ctx *gin.Context) bool { @@ -52,7 +66,14 @@ func (c *Common) HasPermission(ctx *gin.Context) bool { return true } - return user.ID == c.UserID + // 必须走 GetUserID 而不是裸读 c.UserID — Server.UserID 在 + // ServerTransfer.Register / revertTransition 里会被 atomic.StoreUint64 + // 改写,dashboard 各 controller 在 listHandler post-filter 这条热路径上 + // 高频对同一 *Server 调 HasPermission。裸读会与 SetUserID 形成 data + // race(TestCommonHasPermissionConcurrentWithSetUserIDIsRaceFree 在 + // -race 下钉死该不变量),并且在 transfer 切换瞬间可能给出错误的权限 + // 判断。 + return user.ID == c.GetUserID() } type CommonInterface interface { diff --git a/model/config.go b/model/config.go index 9e549319..5f74f2f3 100644 --- a/model/config.go +++ b/model/config.go @@ -3,6 +3,7 @@ package model import ( "os" "path/filepath" + "strconv" "strings" "github.com/go-viper/mapstructure/v2" @@ -16,8 +17,12 @@ import ( ) const ( - ConfigUsePeerIP = "NZ::Use-Peer-IP" - ConfigCoverAll = iota + ConfigUsePeerIP = "NZ::Use-Peer-IP" + JWTSecretKeyRotationBaselineVersion = "v2.0.13" +) + +const ( + ConfigCoverAll = iota + 1 ConfigCoverIgnoreAll ) @@ -60,9 +65,10 @@ type Config struct { AgentSecretKey string `koanf:"agent_secret_key" json:"agent_secret_key,omitempty"` JWTTimeout int `koanf:"jwt_timeout" json:"jwt_timeout,omitempty"` // JWT token过期时间(小时) - JWTSecretKey string `koanf:"jwt_secret_key" json:"jwt_secret_key,omitempty"` - ListenPort uint16 `koanf:"listen_port" json:"listen_port,omitempty"` - ListenHost string `koanf:"listen_host" json:"listen_host,omitempty"` + JWTSecretKey string `koanf:"jwt_secret_key" json:"jwt_secret_key,omitempty"` + JWTSecretKeyLastRotatedVersion string `koanf:"jwt_secret_key_last_rotated_version" json:"jwt_secret_key_last_rotated_version,omitempty"` + ListenPort uint16 `koanf:"listen_port" json:"listen_port,omitempty"` + ListenHost string `koanf:"listen_host" json:"listen_host,omitempty"` // oauth2 配置 Oauth2 map[string]*Oauth2Config `koanf:"oauth2" json:"oauth2,omitempty"` @@ -193,6 +199,30 @@ func (c *Config) Save() error { return c.save() } +func (c *Config) RotateJWTSecretKeyIfNeeded(currentVersion string) (bool, error) { + currentVersion = strings.TrimSpace(currentVersion) + if compareVersion(currentVersion, JWTSecretKeyRotationBaselineVersion) < 0 { + return false, nil + } + + initialMarker := c.JWTSecretKeyLastRotatedVersion + shouldRotate := c.JWTSecretKeyLastRotatedVersion == "" || compareVersion(c.JWTSecretKeyLastRotatedVersion, JWTSecretKeyRotationBaselineVersion) < 0 + if shouldRotate { + secret, err := utils.GenerateRandomString(1024) + if err != nil { + return false, err + } + c.JWTSecretKey = secret + } + + c.JWTSecretKeyLastRotatedVersion = currentVersion + + if !shouldRotate && c.JWTSecretKeyLastRotatedVersion == initialMarker { + return false, nil + } + return shouldRotate, c.Save() +} + func (c *Config) save() error { data, err := yaml.Marshal(c) if err != nil { @@ -211,6 +241,40 @@ func (c *Config) write(data []byte) error { return os.WriteFile(c.filePath, data, 0600) } +func compareVersion(left, right string) int { + leftParts, leftOK := parseVersion(left) + rightParts, rightOK := parseVersion(right) + if !leftOK || !rightOK { + return -1 + } + for i := range leftParts { + if leftParts[i] < rightParts[i] { + return -1 + } + if leftParts[i] > rightParts[i] { + return 1 + } + } + return 0 +} + +func parseVersion(version string) ([3]int, bool) { + version = strings.TrimPrefix(strings.TrimSpace(version), "v") + parts := strings.Split(version, ".") + if len(parts) != 3 { + return [3]int{}, false + } + var parsed [3]int + for i, part := range parts { + value, err := strconv.Atoi(part) + if err != nil { + return [3]int{}, false + } + parsed[i] = value + } + return parsed, true +} + func koanfConf(c any) koanf.UnmarshalConf { return koanf.UnmarshalConf{ DecoderConfig: &mapstructure.DecoderConfig{ diff --git a/model/config_test.go b/model/config_test.go index bce288bc..4daa73ad 100644 --- a/model/config_test.go +++ b/model/config_test.go @@ -151,6 +151,111 @@ func TestReadConfig(t *testing.T) { }) } +func TestRotateJWTSecretKeyIfNeeded(t *testing.T) { + tests := []struct { + name string + initialMarker string + currentVersion string + wantRotated bool + wantStoredVersion string + wantSecretChanged bool + wantSavedConfigKey bool + }{ + { + name: "empty marker rotates leaked secret", + currentVersion: "v2.0.13", + wantRotated: true, + wantStoredVersion: "v2.0.13", + wantSecretChanged: true, + wantSavedConfigKey: true, + }, + { + name: "old marker rotates leaked secret", + initialMarker: "v2.0.12", + currentVersion: "v2.0.14", + wantRotated: true, + wantStoredVersion: "v2.0.14", + wantSecretChanged: true, + wantSavedConfigKey: true, + }, + { + name: "threshold marker keeps secret and advances marker", + initialMarker: "v2.0.13", + currentVersion: "v2.0.14", + wantStoredVersion: "v2.0.14", + wantSavedConfigKey: true, + }, + { + name: "current marker keeps secret", + initialMarker: "v2.0.14", + currentVersion: "v2.0.14", + wantStoredVersion: "v2.0.14", + }, + { + name: "debug version skips rotation and marker update", + currentVersion: "debug", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + file := newTempConfig(t, "") + t.Cleanup(func() { os.Remove(file) }) + + c := &Config{ + JWTSecretKey: "leaked-secret", + JWTSecretKeyLastRotatedVersion: tt.initialMarker, + filePath: file, + } + + rotated, err := c.RotateJWTSecretKeyIfNeeded(tt.currentVersion) + if err != nil { + t.Fatalf("rotate jwt secret key failed: %v", err) + } + if rotated != tt.wantRotated { + t.Fatalf("rotated = %v, want %v", rotated, tt.wantRotated) + } + if c.JWTSecretKeyLastRotatedVersion != tt.wantStoredVersion { + t.Fatalf("jwt secret key marker = %q, want %q", c.JWTSecretKeyLastRotatedVersion, tt.wantStoredVersion) + } + secretChanged := c.JWTSecretKey != "leaked-secret" + if secretChanged != tt.wantSecretChanged { + t.Fatalf("secret changed = %v, want %v", secretChanged, tt.wantSecretChanged) + } + + saved, err := os.ReadFile(file) + if err != nil { + t.Fatalf("read saved config: %v", err) + } + hasMarker := strings.Contains(string(saved), "jwt_secret_key_last_rotated_version") + if hasMarker != tt.wantSavedConfigKey { + t.Fatalf("saved marker present = %v, want %v, config = %s", hasMarker, tt.wantSavedConfigKey, saved) + } + }) + } +} + +// Mirrors the upstream single-block declaration so iota lines up exactly: +// ConfigUsePeerIP occupies iota=0 (as a typed string), ConfigCoverAll=1, +// ConfigCoverIgnoreAll=2. Pins persisted `cover` semantics. +const ( + originalConfigUsePeerIP = "NZ::Use-Peer-IP" + originalConfigCoverAll = iota + originalConfigCoverIgnoreAll +) + +func TestConfigCoverConstantValues(t *testing.T) { + if ConfigUsePeerIP != originalConfigUsePeerIP { + t.Fatalf("ConfigUsePeerIP = %q, want %q", ConfigUsePeerIP, originalConfigUsePeerIP) + } + if ConfigCoverAll != originalConfigCoverAll { + t.Fatalf("ConfigCoverAll = %d, want original value %d", ConfigCoverAll, originalConfigCoverAll) + } + if ConfigCoverIgnoreAll != originalConfigCoverIgnoreAll { + t.Fatalf("ConfigCoverIgnoreAll = %d, want original value %d", ConfigCoverIgnoreAll, originalConfigCoverIgnoreAll) + } +} + func newTempConfig(t *testing.T, cfg string) string { t.Helper() diff --git a/model/notification_test.go b/model/notification_test.go index 25e6599a..44e39f34 100644 --- a/model/notification_test.go +++ b/model/notification_test.go @@ -80,7 +80,6 @@ func execCase(t *testing.T, item testSt) { CountryCode: "", }, LastActive: time.Time{}, - TaskStream: nil, PrevTransferInSnapshot: 0, PrevTransferOutSnapshot: 0, } diff --git a/model/server.go b/model/server.go index 165d7acb..3d25c027 100644 --- a/model/server.go +++ b/model/server.go @@ -3,6 +3,7 @@ package model import ( "log" "slices" + "sync/atomic" "time" "github.com/goccy/go-json" @@ -32,13 +33,68 @@ type Server struct { GeoIP *GeoIP `gorm:"-" json:"geoip,omitempty"` LastActive time.Time `gorm:"-" json:"last_active,omitempty"` - TaskStream pb.NezhaService_RequestTaskServer `gorm:"-" json:"-"` - ConfigCache chan any `gorm:"-" json:"-"` + // taskStream MUST be accessed only via SetTaskStream / GetTaskStream. Direct + // field access from outside this file races with the gRPC RequestTask + // handler that reassigns the stream on every reconnect — a torn read of the + // two-word interface header would panic on a subsequent .Send call. The + // atomic.Pointer + holder struct lets us swap the stream lock-free while + // every reader observes a single, consistent value. + taskStream atomic.Pointer[taskStreamHolder] + ConfigCache chan any `gorm:"-" json:"-"` PrevTransferInSnapshot uint64 `gorm:"-" json:"-"` // 上次数据点时的入站使用量 PrevTransferOutSnapshot uint64 `gorm:"-" json:"-"` // 上次数据点时的出站使用量 } +// taskStreamHolder wraps the interface so atomic.Pointer (which requires a +// concrete pointed-to type) can publish it atomically. The previous bare +// field `TaskStream pb.NezhaService_RequestTaskServer` was a plain interface +// value: two words on the heap (type ptr + data ptr). Concurrent assignment +// produced torn reads detectable by `go test -race` and crashable in production. +type taskStreamHolder struct { + s pb.NezhaService_RequestTaskServer +} + +// SetTaskStream publishes the agent's RequestTask stream so other goroutines +// can deliver tasks to the agent. Pass nil to detach (e.g. on disconnect). +func (s *Server) SetTaskStream(stream pb.NezhaService_RequestTaskServer) { + if stream == nil { + s.taskStream.Store(nil) + return + } + s.taskStream.Store(&taskStreamHolder{s: stream}) +} + +// ClearTaskStreamIfCurrent detaches stream only if it is still the published +// RequestTask stream. Disconnect cleanup uses this guard so an old stream +// returning after a reconnect cannot erase the newer live stream. +func (s *Server) ClearTaskStreamIfCurrent(stream pb.NezhaService_RequestTaskServer) bool { + if stream == nil { + return false + } + for { + h := s.taskStream.Load() + if h == nil || h.s != stream { + return false + } + if s.taskStream.CompareAndSwap(h, nil) { + return true + } + } +} + +// GetTaskStream returns the currently-published stream, or nil if the agent +// is offline. Callers MUST capture the return into a local variable before +// using it — re-reading via GetTaskStream() across a Send call reopens the +// race we're trying to close. +func (s *Server) GetTaskStream() pb.NezhaService_RequestTaskServer { + h := s.taskStream.Load() + if h == nil { + return nil + } + return h.s +} + func InitServer(s *Server) { s.Host = &Host{} s.State = &HostState{} @@ -51,7 +107,9 @@ func (s *Server) CopyFromRunningServer(old *Server) { s.State = old.State s.GeoIP = old.GeoIP s.LastActive = old.LastActive - s.TaskStream = old.TaskStream + // taskStream is an atomic.Pointer; copy the published value rather than + // the field itself (atomic.Pointer is not safe to copy by value). + s.SetTaskStream(old.GetTaskStream()) s.ConfigCache = old.ConfigCache s.PrevTransferInSnapshot = old.PrevTransferInSnapshot s.PrevTransferOutSnapshot = old.PrevTransferOutSnapshot @@ -73,6 +131,50 @@ func (s *Server) AfterFind(tx *gorm.DB) error { return nil } +// ServerOwnerInfo carries the user-facing identity for Server.UserID. It is +// returned by the lookup function installed by the singleton layer; model +// must not import singleton (cycle), so the dependency flows through a +// package-level function variable instead. +type ServerOwnerInfo struct { + ID uint64 `json:"id"` + Username string `json:"username,omitempty"` +} + +// ServerOwnerLookup is installed by singleton at startup to resolve a +// Server.UserID into a display-ready owner record. Returns ok=false when +// the uid does not map to a known user; the caller renders that as an +// "unknown user" placeholder so deleted-user rows stay debuggable. Left nil +// in tests / headless contexts so the JSON simply omits the owner field. +var ServerOwnerLookup func(uid uint64) (ServerOwnerInfo, bool) + +type serverJSON Server + +type serverWithOwner struct { + *serverJSON + Owner *ServerOwnerInfo `json:"owner,omitempty"` +} + +// MarshalJSON projects Server.UserID into a structured owner field on the +// wire. Server.UserID itself stays `json:"-"` (set on Common) so callers +// that do not need owner info pay nothing and members do not accidentally +// receive raw uid integers. The lookup function is consulted only when +// installed; if absent we still emit a minimal {id} record so clients can +// at least distinguish ownership, except for uid=0 which is the legacy +// global-secret pseudo-owner and is best surfaced as such by the caller's +// translation table on the frontend. +func (s *Server) MarshalJSON() ([]byte, error) { + owner := &ServerOwnerInfo{ID: s.GetUserID()} + if ServerOwnerLookup != nil { + if info, ok := ServerOwnerLookup(owner.ID); ok { + owner.Username = info.Username + } + } + return json.Marshal(serverWithOwner{ + serverJSON: (*serverJSON)(s), + Owner: owner, + }) +} + func (s *Server) SplitList(x []*Server) ([]*Server, []*Server) { pri := func(s *Server) bool { return s.DisplayIndex == 0 diff --git a/model/server_owner_json_test.go b/model/server_owner_json_test.go new file mode 100644 index 00000000..4c0e7119 --- /dev/null +++ b/model/server_owner_json_test.go @@ -0,0 +1,129 @@ +package model + +import ( + "encoding/json" + "testing" +) + +// Server.MarshalJSON projects Server.UserID into a public owner field while +// keeping the raw UserID json-hidden. The lookup function is package-level +// and shared across tests; each subtest installs its own stub and restores +// the original to avoid leaking state. +func TestServerMarshalJSONOwnerProjection(t *testing.T) { + original := ServerOwnerLookup + t.Cleanup(func() { ServerOwnerLookup = original }) + + tests := []struct { + name string + uid uint64 + lookup func(uid uint64) (ServerOwnerInfo, bool) + wantID uint64 + wantHasName bool + wantName string + }{ + { + // uid=0 is the legacy global agent secret pseudo-owner. The + // lookup deliberately returns ok=false so the frontend can + // render it as "Global Agent" instead of a real username. + name: "uid_zero_has_no_username", + uid: 0, + lookup: func(uint64) (ServerOwnerInfo, bool) { + return ServerOwnerInfo{}, false + }, + wantID: 0, + wantHasName: false, + }, + { + // Known user → username flows through to the wire so the + // admin frontend can show it without a separate /user fetch + // (which members cannot call anyway). + name: "known_user_has_username", + uid: 42, + lookup: func(uid uint64) (ServerOwnerInfo, bool) { + return ServerOwnerInfo{ID: uid, Username: "alice"}, true + }, + wantID: 42, + wantHasName: true, + wantName: "alice", + }, + { + // Deleted user → lookup returns ok=false. The wire still + // carries owner.id so the frontend can render an "Unknown + // user (#id)" placeholder; otherwise the row would silently + // appear ownerless and ops would lose the audit trail. + name: "deleted_user_keeps_id_without_username", + uid: 999, + lookup: func(uint64) (ServerOwnerInfo, bool) { + return ServerOwnerInfo{}, false + }, + wantID: 999, + wantHasName: false, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + ServerOwnerLookup = tc.lookup + s := &Server{Common: Common{ID: 7, UserID: tc.uid}, Name: "srv"} + + raw, err := json.Marshal(s) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got struct { + Owner *ServerOwnerInfo `json:"owner"` + // Owner must never appear as the raw uid via Common.UserID; + // the Common.UserID json tag is "-" and a regression that + // flips it to "user_id" would expose internal owner ids + // to the wire bypassing the lookup-controlled rendering. + UserID *uint64 `json:"user_id,omitempty"` + } + if err := json.Unmarshal(raw, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.UserID != nil { + t.Fatalf("Common.UserID must not appear on the wire as user_id, got %d", *got.UserID) + } + if got.Owner == nil { + t.Fatalf("owner field must always be present, raw=%s", raw) + } + if got.Owner.ID != tc.wantID { + t.Fatalf("owner.id=%d, want %d", got.Owner.ID, tc.wantID) + } + if tc.wantHasName { + if got.Owner.Username != tc.wantName { + t.Fatalf("owner.username=%q, want %q", got.Owner.Username, tc.wantName) + } + } else if got.Owner.Username != "" { + t.Fatalf("owner.username must be omitted for uid=%d, got %q", tc.uid, got.Owner.Username) + } + }) + } +} + +// When no lookup is installed (tests / headless tools), MarshalJSON must +// still emit a minimal owner record so consumers do not crash on missing +// fields. Without this guard a future refactor could silently drop the +// owner key entirely whenever the hook is nil. +func TestServerMarshalJSONEmitsOwnerWithoutLookup(t *testing.T) { + original := ServerOwnerLookup + t.Cleanup(func() { ServerOwnerLookup = original }) + ServerOwnerLookup = nil + + raw, err := json.Marshal(&Server{Common: Common{ID: 1, UserID: 17}, Name: "srv"}) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var got struct { + Owner *ServerOwnerInfo `json:"owner"` + } + if err := json.Unmarshal(raw, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if got.Owner == nil || got.Owner.ID != 17 || got.Owner.Username != "" { + t.Fatalf("expected bare owner record {id:17}, got %+v", got.Owner) + } +} diff --git a/model/server_owner_race_test.go b/model/server_owner_race_test.go new file mode 100644 index 00000000..aa93d83f --- /dev/null +++ b/model/server_owner_race_test.go @@ -0,0 +1,86 @@ +package model + +import ( + "net/http/httptest" + "sync" + "testing" + + "github.com/gin-gonic/gin" +) + +// Server.UserID 在 server-transfer rotation 流程里会被 ServerTransfer 的 +// Register/revertTransition 改写以反映新所有者,同时 authorizeAgentForUUID +// 在每次 agent RPC 里读取它。原实现两处都是裸字段访问,race detector 会 +// 报告 data race;这是 review 评分 75 的真实问题。 +// +// 修复后所有并发读写都走 SetUserID/GetUserID 的 atomic 包装,本测试在 +// `go test -race` 下应该完全跑干净。 +func TestServerUserIDConcurrentAccessIsRaceFree(t *testing.T) { + s := &Server{} + + const ( + writers = 4 + readers = 8 + rounds = 500 + ) + var wg sync.WaitGroup + wg.Add(writers + readers) + + for i := 0; i < writers; i++ { + uid := uint64(i + 1) + go func() { + defer wg.Done() + for j := 0; j < rounds; j++ { + s.SetUserID(uid) + } + }() + } + for i := 0; i < readers; i++ { + go func() { + defer wg.Done() + for j := 0; j < rounds; j++ { + _ = s.GetUserID() + } + }() + } + wg.Wait() +} + +// Common.HasPermission 是 server-transfer 旋转下与 SetUserID 并发的主要读者 +// 之一:dashboard 各 controller 的 listHandler post-filter 在 transfer 窗口 +// 内不断对同一 *Server 调用 HasPermission,而 Register/revertTransition 同 +// 时通过 SetUserID 改写所属用户。原实现的 `user.ID == c.UserID` 是裸读,会 +// 与 atomic.StoreUint64 形成 data race(go test -race 必爆)。修复后改成走 +// GetUserID() 走 atomic 协议。这个测试就是用来在 -race 下钉死该不变量的。 +func TestCommonHasPermissionConcurrentWithSetUserIDIsRaceFree(t *testing.T) { + s := &Server{Common: Common{ID: 1}} + + const ( + writers = 4 + readers = 8 + rounds = 500 + ) + var wg sync.WaitGroup + wg.Add(writers + readers) + + for i := 0; i < writers; i++ { + uid := uint64(i + 1) + go func() { + defer wg.Done() + for j := 0; j < rounds; j++ { + s.SetUserID(uid) + } + }() + } + for i := 0; i < readers; i++ { + go func() { + defer wg.Done() + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 2}, Role: RoleMember}) + for j := 0; j < rounds; j++ { + _ = s.HasPermission(ctx) + } + }() + } + wg.Wait() +} diff --git a/model/server_taskstream_race_test.go b/model/server_taskstream_race_test.go new file mode 100644 index 00000000..8ea53008 --- /dev/null +++ b/model/server_taskstream_race_test.go @@ -0,0 +1,91 @@ +package model + +import ( + "context" + "sync" + "testing" + + pb "github.com/nezhahq/nezha/proto" +) + +// raceProbeStream is the smallest fake of pb.NezhaService_RequestTaskServer +// the race probe needs. We only call Send on it from the test; the embedded +// interface satisfies the rest of the contract with nil-panicking methods we +// never invoke. +type raceProbeStream struct { + pb.NezhaService_RequestTaskServer +} + +func (raceProbeStream) Send(*pb.Task) error { return nil } +func (raceProbeStream) Context() context.Context { return context.Background() } + +// model.Server.TaskStream is read from many goroutines (singleton cron pushes, +// transfer ApplyConfig pushes, terminal/fm proxies, dashboard rpc keepalives, +// per-server batch pushes) and written from exactly one (the gRPC RequestTask +// goroutine on every fresh agent connection). The bare-field access pattern +// `if s.TaskStream != nil { s.TaskStream.Send(...) }` is a data race on the +// interface header (two-word value) and can torn-read into a panic on a +// reconnect. This test pins down "concurrent set + send must be race-free" +// using the Go race detector — without the fix, `go test -race` reports a +// data race on TaskStream; with the fix the field is encapsulated behind +// atomic methods and the test runs clean. Without `-race` both versions are +// indistinguishable, so this test is only meaningful under the race flag — +// run it from CI as `go test -race ./model/`. +func TestServerTaskStreamConcurrentAccessIsRaceFree(t *testing.T) { + s := &Server{} + InitServer(s) + + const ( + writers = 4 + readers = 8 + rounds = 200 + ) + var wg sync.WaitGroup + wg.Add(writers + readers) + + for i := 0; i < writers; i++ { + go func() { + defer wg.Done() + for j := 0; j < rounds; j++ { + s.SetTaskStream(raceProbeStream{}) + s.SetTaskStream(nil) + } + }() + } + for i := 0; i < readers; i++ { + go func() { + defer wg.Done() + for j := 0; j < rounds; j++ { + if stream := s.GetTaskStream(); stream != nil { + _ = stream.Send(nil) + } + } + }() + } + wg.Wait() +} + +func TestServerClearTaskStreamIfCurrentClearsOnlyMatchingStream(t *testing.T) { + s := &Server{} + InitServer(s) + + first := &raceProbeStream{} + second := &raceProbeStream{} + + s.SetTaskStream(first) + if !s.ClearTaskStreamIfCurrent(first) { + t.Fatal("matching current stream must be cleared") + } + if got := s.GetTaskStream(); got != nil { + t.Fatalf("expected cleared task stream, got %T", got) + } + + s.SetTaskStream(first) + s.SetTaskStream(second) + if s.ClearTaskStreamIfCurrent(first) { + t.Fatal("stale stream cleanup must not clear a newer stream") + } + if got := s.GetTaskStream(); got != second { + t.Fatalf("expected newer stream to remain published, got %T", got) + } +} diff --git a/model/server_transfer.go b/model/server_transfer.go new file mode 100644 index 00000000..90e49410 --- /dev/null +++ b/model/server_transfer.go @@ -0,0 +1,128 @@ +package model + +import ( + "time" + + "github.com/gin-gonic/gin" +) + +// ServerTransferStatus represents the lifecycle state of a server ownership +// transfer. A transfer's life starts at Pending (server.user_id has been +// flipped to the new owner; agent still authenticates with the old owner's +// AgentSecret) and ends in exactly one of the terminal states. +type ServerTransferStatus uint8 + +const ( + // ServerTransferStatusPending means the dashboard has flipped Server.UserID + // to the new owner and queued an ApplyConfig task to swap the agent's + // client_secret. Auth still accepts the old owner's AgentSecret for this + // UUID until verification arrives or the transfer times out. + ServerTransferStatusPending ServerTransferStatus = iota + // ServerTransferStatusVerified means the agent successfully reconnected + // using the new owner's AgentSecret. Auth no longer tolerates the old + // owner's secret on this UUID. + ServerTransferStatusVerified + // ServerTransferStatusFailed means the agent explicitly reported the + // ApplyConfig task as unsuccessful (e.g. DisableCommandExecute). The + // dashboard has rolled Server.UserID back to FromUserID. + ServerTransferStatusFailed + // ServerTransferStatusTimeout means the verification window expired + // without the agent reconnecting under the new secret. The dashboard has + // rolled Server.UserID back to FromUserID. + ServerTransferStatusTimeout + // ServerTransferStatusCancelled means an administrator cancelled the + // transfer before any verification event was observed. The dashboard has + // rolled Server.UserID back to FromUserID. + ServerTransferStatusCancelled +) + +// IsTerminal reports whether the status represents a settled transfer. Only +// terminal transfers are eligible for retry and they will never be in the +// pending index. +func (s ServerTransferStatus) IsTerminal() bool { + return s != ServerTransferStatusPending +} + +// ServerTransfer records a single attempt to transfer ownership of one server +// to another user. It is the source of truth for the auth-tolerance window +// during a transfer — service/rpc.authorizeAgentForUUID consults the pending +// index built from this table to decide whether to accept the old owner's +// AgentSecret on the affected UUID. +// +// Naming note: the existing model.Transfer records hourly traffic snapshots +// and is unrelated. This entity is named ServerTransfer to disambiguate. +type ServerTransfer struct { + Common + ServerID uint64 `json:"server_id" gorm:"index"` + FromUserID uint64 `json:"from_user_id"` + ToUserID uint64 `json:"to_user_id"` + InitiatorID uint64 `json:"initiator_id"` + Status ServerTransferStatus `json:"status" gorm:"index"` + LastError string `json:"last_error,omitempty"` + AckedAt *time.Time `json:"acked_at,omitempty"` + // HandshakeSecret is a per-transfer random credential that PushIfOnline + // delivers in place of the destination user's global AgentSecret. The + // agent treats it as a temporary handshake token: it rotates to this + // secret on the 10s reload, reconnects, and the dashboard's auth path + // recognises it as proof of transfer delivery (MarkVerified). It is + // scoped to this single transfer and to this single UUID — leaking it + // to the previous owner who hijacks the stream still does NOT expose + // the destination user's other agents. Never returned to API clients. + HandshakeSecret string `json:"-" gorm:"type:char(32)"` + // RevertHandshakeSecret is the same idea for the rollback path: when + // the dashboard pushes a revert ApplyConfig over a stream now held by + // the destination user, we must not embed the source user's global + // AgentSecret. Instead the agent rotates back through this token, which + // is recognised by the auth path during the revert window only. + RevertHandshakeSecret string `json:"-" gorm:"type:char(32)"` +} + +// HasPermission overrides Common.HasPermission so a transfer is visible to +// admins, the source user, the destination user, and the initiator. Listing +// uses this to filter what the caller can see; mutating endpoints (cancel, +// retry) layer additional checks on top. +func (t *ServerTransfer) HasPermission(ctx *gin.Context) bool { + auth, ok := ctx.Get(CtxKeyAuthorizedUser) + if !ok { + return false + } + user := *auth.(*User) + if user.Role == RoleAdmin { + return true + } + return user.ID == t.FromUserID || user.ID == t.ToUserID || user.ID == t.InitiatorID +} + +// BatchMoveServerResultStatus is the per-server outcome returned by the +// batch-move endpoint. It maps to TransferStatus for transfers that were +// successfully created, plus extra synchronous-failure modes (permission, +// duplicate active transfer, missing server) that never produce a row. +type BatchMoveServerResultStatus string + +const ( + // BatchMoveServerResultPending: ServerTransfer row created, agent push + // in progress. Callers should watch the WS for terminal status. + BatchMoveServerResultPending BatchMoveServerResultStatus = "pending" + // BatchMoveServerResultPermissionDenied: caller cannot move this server. + BatchMoveServerResultPermissionDenied BatchMoveServerResultStatus = "permission_denied" + // BatchMoveServerResultAlreadyTransferring: server already has an in-flight + // ServerTransfer row, cancel or wait first. + BatchMoveServerResultAlreadyTransferring BatchMoveServerResultStatus = "already_transferring" + // BatchMoveServerResultServerNotFound: server id does not exist. + BatchMoveServerResultServerNotFound BatchMoveServerResultStatus = "server_not_found" + // BatchMoveServerResultSameOwner: target user already owns this server. + BatchMoveServerResultSameOwner BatchMoveServerResultStatus = "same_owner" + // BatchMoveServerResultAgentTooOld: agent build does not understand + // TaskTypeServerTransferApply, so the rotation would never complete and + // dashboard refuses to start it. Operator must upgrade the agent. + BatchMoveServerResultAgentTooOld BatchMoveServerResultStatus = "agent_too_old" +) + +// BatchMoveServerResult is one entry in the batchMoveServer response, one +// per requested server id, in the same order. +type BatchMoveServerResult struct { + ServerID uint64 `json:"server_id"` + Status BatchMoveServerResultStatus `json:"status"` + TransferID uint64 `json:"transfer_id,omitempty"` + Error string `json:"error,omitempty"` +} diff --git a/model/service.go b/model/service.go index 72153bc1..26627205 100644 --- a/model/service.go +++ b/model/service.go @@ -26,6 +26,10 @@ const ( TaskTypeFM TaskTypeReportConfig TaskTypeApplyConfig + // TaskTypeServerTransferApply: per-transfer credential rotation. + // Pre-transfer agents do not recognise this type — dashboard MUST gate + // transfers on agent capability before pushing. + TaskTypeServerTransferApply ) type TerminalTask struct { @@ -133,7 +137,8 @@ func IsServiceSentinelNeeded(t uint64) bool { switch t { case TaskTypeCommand, TaskTypeTerminalGRPC, TaskTypeUpgrade, TaskTypeKeepalive, TaskTypeNAT, TaskTypeFM, - TaskTypeReportConfig, TaskTypeApplyConfig: + TaskTypeReportConfig, TaskTypeApplyConfig, + TaskTypeServerTransferApply: return false default: return true diff --git a/model/user.go b/model/user.go index 2c40563b..6a3ffcd0 100644 --- a/model/user.go +++ b/model/user.go @@ -32,6 +32,7 @@ type User struct { type UserInfo struct { Role Role + Username string AgentSecret string } diff --git a/service/rpc/apply_config_authz_test.go b/service/rpc/apply_config_authz_test.go new file mode 100644 index 00000000..5f290a54 --- /dev/null +++ b/service/rpc/apply_config_authz_test.go @@ -0,0 +1,565 @@ +package rpc + +import ( + "context" + "errors" + "strings" + "testing" + "time" + + "google.golang.org/grpc/metadata" + "gorm.io/driver/sqlite" + "gorm.io/gorm" + + "github.com/nezhahq/nezha/model" + pb "github.com/nezhahq/nezha/proto" + "github.com/nezhahq/nezha/service/singleton" +) + +// A malicious or buggy agent owning server A must NOT be able to fail a +// ServerTransfer row belonging to server B by reporting a TaskResult whose +// Id is set to B's transfer ID. The agent-task-result authorization +// invariant (commit 02129f1) requires the dashboard to verify the result's +// addressed object actually belongs to the reporting agent before acting +// on it. Without the cross-check, any compromised agent could cancel/fail +// every in-flight transfer in the system. +func TestRequestTaskApplyConfigIgnoresForeignTransferFailure(t *testing.T) { + // Two distinct servers with different owners. attackerSrv reports the + // failure; victimSrv is the one a pending transfer points at. + attackerSrv := &model.Server{ + Common: model.Common{ID: 7, UserID: 100}, + UUID: "cccccccc-cccc-cccc-cccc-cccccccccccc", + Name: "attacker", + } + victimSrv := &model.Server{ + Common: model.Common{ID: 8, UserID: 200}, + UUID: "dddddddd-dddd-dddd-dddd-dddddddddddd", + Name: "victim", + } + users := map[uint64]model.UserInfo{ + 100: {Role: model.RoleMember}, + 200: {Role: model.RoleMember}, + 300: {Role: model.RoleMember, AgentSecret: "to-user-secret"}, + } + secrets := map[string]uint64{ + "attacker-secret": 100, + "to-user-secret": 300, + } + setupApplyConfigAuthzFixture(t, []*model.Server{attackerSrv, victimSrv}, users, secrets) + + // Pending transfer for victimSrv (200 -> 300). attackerSrv is unrelated. + tr := initiateAndRegisterPendingTransfer(t, victimSrv.ID, 200, 300, 1) + + // Attacker reports a failed ApplyConfig carrying the victim's transfer ID. + runApplyConfigAuthzResult(t, "attacker-secret", attackerSrv.UUID, &pb.TaskResult{ + Id: tr.ID, + Type: model.TaskTypeServerTransferApply, + Successful: false, + Data: "spoofed failure", + }) + + var refreshed model.ServerTransfer + if err := singleton.DB.First(&refreshed, tr.ID).Error; err != nil { + t.Fatalf("re-read transfer: %v", err) + } + if refreshed.Status != model.ServerTransferStatusPending { + t.Fatalf("foreign-server ApplyConfig failure must leave transfer Pending, got status=%d last_error=%q", + refreshed.Status, refreshed.LastError) + } + + var vs model.Server + if err := singleton.DB.First(&vs, victimSrv.ID).Error; err != nil { + t.Fatalf("re-read victim server: %v", err) + } + if vs.UserID != 300 { + t.Fatalf("victim server ownership must remain at ToUserID, got %d", vs.UserID) + } +} + +// The legitimate path must still mark the transfer Failed: the reporter is +// the actual transfer subject. This guards against an over-tight ownership +// check that would also break the working flow. +func TestRequestTaskApplyConfigAcceptsOwnTransferFailure(t *testing.T) { + srv := &model.Server{ + Common: model.Common{ID: 9, UserID: 200}, + UUID: "eeeeeeee-eeee-eeee-eeee-eeeeeeeeeeee", + Name: "subject", + } + users := map[uint64]model.UserInfo{ + 200: {Role: model.RoleMember, AgentSecret: "from-user-secret"}, + 300: {Role: model.RoleMember, AgentSecret: "to-user-secret"}, + } + secrets := map[string]uint64{ + // During Pending the agent still authenticates with the previous + // owner's secret — that's exactly the auth-tolerance window the + // transfer feature exists for. + "from-user-secret": 200, + "to-user-secret": 300, + } + setupApplyConfigAuthzFixture(t, []*model.Server{srv}, users, secrets) + + tr := initiateAndRegisterPendingTransfer(t, srv.ID, 200, 300, 1) + + runApplyConfigAuthzResult(t, "from-user-secret", srv.UUID, &pb.TaskResult{ + Id: tr.ID, + Type: model.TaskTypeServerTransferApply, + Successful: false, + Data: "DisableCommandExecute=true", + }) + + var refreshed model.ServerTransfer + if err := singleton.DB.First(&refreshed, tr.ID).Error; err != nil { + t.Fatalf("re-read transfer: %v", err) + } + if refreshed.Status != model.ServerTransferStatusFailed { + t.Fatalf("own-server ApplyConfig failure must mark transfer Failed, got status=%d", refreshed.Status) + } +} + +func TestRequestTaskCancelledTransferAllowsForwardHandshakeReconnectForRevert(t *testing.T) { + srv := &model.Server{ + Common: model.Common{ID: 12, UserID: 200}, + UUID: "12121212-1212-1212-1212-121212121212", + Name: "cancelled-revert", + } + users := map[uint64]model.UserInfo{ + 200: {Role: model.RoleMember, AgentSecret: "cancel-from-secret"}, + 300: {Role: model.RoleMember, AgentSecret: "cancel-to-secret"}, + } + secrets := map[string]uint64{ + "cancel-from-secret": 200, + "cancel-to-secret": 300, + } + setupApplyConfigAuthzFixture(t, []*model.Server{srv}, users, secrets) + + tr := initiateAndRegisterPendingTransfer(t, srv.ID, 200, 300, 1) + forward := tr.HandshakeSecret + if forward == "" { + t.Fatal("precondition: pending transfer must carry a forward HandshakeSecret") + } + if _, err := singleton.ServerTransferShared.Cancel(tr.ID); err != nil { + t.Fatalf("cancel transfer: %v", err) + } + + sent := runApplyConfigAuthzReconnect(t, forward, srv.UUID) + if len(sent) != 1 { + t.Fatalf("expected one revert ApplyConfig task, got %d", len(sent)) + } + if sent[0].Type != model.TaskTypeServerTransferApply { + t.Fatalf("expected ApplyConfig task, got type=%d", sent[0].Type) + } + var settled model.ServerTransfer + if err := singleton.DB.First(&settled, tr.ID).Error; err != nil { + t.Fatalf("reload transfer: %v", err) + } + if !strings.Contains(sent[0].Data, settled.RevertHandshakeSecret) { + t.Fatalf("cancelled transfer rollback must push the per-transfer RevertHandshakeSecret, got payload %q", sent[0].Data) + } + if strings.Contains(sent[0].Data, "cancel-from-secret") || strings.Contains(sent[0].Data, "cancel-to-secret") { + t.Fatalf("user-global AgentSecrets must never appear in transfer payloads, got %q", sent[0].Data) + } +} + +func TestRequestTaskTimedOutTransferAllowsForwardHandshakeReconnectForRevert(t *testing.T) { + srv := &model.Server{ + Common: model.Common{ID: 17, UserID: 200}, + UUID: "17171717-1717-1717-1717-171717171717", + Name: "timeout-revert", + } + users := map[uint64]model.UserInfo{ + 200: {Role: model.RoleMember, AgentSecret: "timeout-from-secret"}, + 300: {Role: model.RoleMember, AgentSecret: "timeout-to-secret"}, + } + secrets := map[string]uint64{ + "timeout-from-secret": 200, + "timeout-to-secret": 300, + } + setupApplyConfigAuthzFixture(t, []*model.Server{srv}, users, secrets) + + tr := initiateAndRegisterPendingTransfer(t, srv.ID, 200, 300, 1) + forward := tr.HandshakeSecret + if forward == "" { + t.Fatal("precondition: pending transfer must carry a forward HandshakeSecret") + } + staleUpdatedAt := time.Now().Add(-25 * time.Hour) + if err := singleton.DB.Model(&model.ServerTransfer{}). + Where("id = ?", tr.ID). + UpdateColumn("updated_at", staleUpdatedAt).Error; err != nil { + t.Fatalf("stale transfer update: %v", err) + } + if _, err := singleton.ServerTransferShared.MarkTimeout(tr.ID); err != nil { + t.Fatalf("timeout transfer: %v", err) + } + + sent := runApplyConfigAuthzReconnect(t, forward, srv.UUID) + if len(sent) != 1 { + t.Fatalf("expected one timeout revert ApplyConfig task, got %d", len(sent)) + } + var settled model.ServerTransfer + if err := singleton.DB.First(&settled, tr.ID).Error; err != nil { + t.Fatalf("reload transfer: %v", err) + } + if !strings.Contains(sent[0].Data, settled.RevertHandshakeSecret) { + t.Fatalf("timeout rollback must push the per-transfer RevertHandshakeSecret, got payload %q", sent[0].Data) + } + if strings.Contains(sent[0].Data, "timeout-from-secret") || strings.Contains(sent[0].Data, "timeout-to-secret") { + t.Fatalf("user-global AgentSecrets must never appear in transfer payloads, got %q", sent[0].Data) + } +} + +func TestRequestTaskRejectsToUserGlobalSecretEvenWithLiveRevertDelivery(t *testing.T) { + srv := &model.Server{ + Common: model.Common{ID: 16, UserID: 200}, + UUID: "16161616-1616-1616-1616-161616161616", + Name: "to-user-global-rejected", + } + users := map[uint64]model.UserInfo{ + 200: {Role: model.RoleMember, AgentSecret: "rejected-from-secret"}, + 300: {Role: model.RoleMember, AgentSecret: "rejected-to-secret"}, + } + secrets := map[string]uint64{ + "rejected-from-secret": 200, + "rejected-to-secret": 300, + } + setupApplyConfigAuthzFixture(t, []*model.Server{srv}, users, secrets) + + tr := initiateAndRegisterPendingTransfer(t, srv.ID, 200, 300, 1) + if _, err := singleton.ServerTransferShared.Cancel(tr.ID); err != nil { + t.Fatalf("cancel transfer: %v", err) + } + if _, ok := singleton.ServerTransferShared.LookupRevertDelivery(srv.ID); !ok { + t.Fatal("precondition: cancel must register a revert delivery") + } + + sent := 0 + stream := &requestTaskSecurityStream{ + ctx: metadata.NewIncomingContext(context.Background(), metadata.Pairs( + "client_secret", "rejected-to-secret", + "client_uuid", srv.UUID, + )), + onSend: func(*pb.Task) { + sent++ + }, + } + if err := NewNezhaHandler().RequestTask(stream); err == nil || errors.Is(err, context.Canceled) { + t.Fatal("ToUserID global AgentSecret must never authenticate via revert recovery; PushIfOnline only delivers per-transfer secrets to the real agent") + } + if sent != 0 { + t.Fatalf("rejected ToUserID auth must not trigger any ApplyConfig push, got %d sends", sent) + } + if _, ok := singleton.ServerTransferShared.LookupRevertDelivery(srv.ID); !ok { + t.Fatal("rejected ToUserID auth must not consume the revert delivery — the real agent still needs it for the eventual per-transfer recovery") + } +} + +// Whether or not a revert delivery is still in flight, the destination +// user's global AgentSecret must be rejected on every auth path — +// PushIfOnline never sends that secret to the agent so a reconnect under +// it cannot come from the real agent. This pins the post-fix invariant. +func TestReportSystemInfoRejectsCancelledTransferToUserSecret(t *testing.T) { + srv := &model.Server{ + Common: model.Common{ID: 15, UserID: 200}, + UUID: "15151515-1515-1515-1515-151515151515", + Name: "cancelled-report", + } + users := map[uint64]model.UserInfo{ + 200: {Role: model.RoleMember, AgentSecret: "report-from-secret"}, + 300: {Role: model.RoleMember, AgentSecret: "report-to-secret"}, + } + secrets := map[string]uint64{ + "report-from-secret": 200, + "report-to-secret": 300, + } + setupApplyConfigAuthzFixture(t, []*model.Server{srv}, users, secrets) + + tr := initiateAndRegisterPendingTransfer(t, srv.ID, 200, 300, 1) + if _, err := singleton.ServerTransferShared.Cancel(tr.ID); err != nil { + t.Fatalf("cancel transfer: %v", err) + } + + ctx := metadata.NewIncomingContext(context.Background(), metadata.Pairs( + "client_secret", "report-to-secret", + "client_uuid", srv.UUID, + )) + if _, err := NewNezhaHandler().ReportSystemInfo(ctx, &pb.Host{}); err == nil { + t.Fatal("ReportSystemInfo must reject the destination user's global AgentSecret during revert recovery; PushIfOnline never delivers that credential to the real agent") + } + if _, ok := singleton.ServerTransferShared.LookupRevertDelivery(srv.ID); !ok { + t.Fatal("rejected non-RequestTask auth must not consume the revert delivery") + } +} + +func setupApplyConfigAuthzFixture(t *testing.T, servers []*model.Server, users map[uint64]model.UserInfo, agentSecrets map[string]uint64) { + t.Helper() + + originalDB := singleton.DB + originalConf := singleton.Conf + originalLoc := singleton.Loc + originalServerShared := singleton.ServerShared + originalUserInfoMap := singleton.UserInfoMap + originalAgentSecretToUserID := singleton.AgentSecretToUserId + originalServerTransferShared := singleton.ServerTransferShared + + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + if err != nil { + t.Fatal(err) + } + sqlDB, err := db.DB() + if err != nil { + t.Fatal(err) + } + sqlDB.SetMaxOpenConns(1) + + singleton.DB = db + singleton.Conf = &singleton.ConfigClass{Config: &model.Config{}} + singleton.Loc = time.UTC + if err := singleton.DB.AutoMigrate(model.Server{}, model.ServerTransfer{}); err != nil { + t.Fatal(err) + } + for _, server := range servers { + if err := singleton.DB.Create(server).Error; err != nil { + t.Fatal(err) + } + } + + singleton.UserLock.Lock() + singleton.UserInfoMap = users + singleton.AgentSecretToUserId = agentSecrets + singleton.UserLock.Unlock() + singleton.ServerShared = singleton.NewServerClass() + for _, server := range servers { + model.InitServer(server) + singleton.ServerShared.Update(server, server.UUID) + } + singleton.ServerTransferShared = singleton.NewServerTransferClass() + + t.Cleanup(func() { + if singleton.ServerTransferShared != nil { + singleton.ServerTransferShared.Stop() + } + sqlDB.Close() + singleton.DB = originalDB + singleton.Conf = originalConf + singleton.Loc = originalLoc + singleton.ServerShared = originalServerShared + singleton.ServerTransferShared = originalServerTransferShared + singleton.UserLock.Lock() + singleton.UserInfoMap = originalUserInfoMap + singleton.AgentSecretToUserId = originalAgentSecretToUserID + singleton.UserLock.Unlock() + }) +} + +func initiateAndRegisterPendingTransfer(t *testing.T, serverID, fromUserID, toUserID, initiatorID uint64) *model.ServerTransfer { + t.Helper() + var created *model.ServerTransfer + err := singleton.DB.Transaction(func(tx *gorm.DB) error { + var err error + created, err = singleton.ServerTransferShared.Initiate(tx, serverID, fromUserID, toUserID, initiatorID) + return err + }) + if err != nil { + t.Fatalf("initiate transfer: %v", err) + } + singleton.ServerTransferShared.Register(created) + return created +} + +func runApplyConfigAuthzResult(t *testing.T, secret, uuid string, result *pb.TaskResult) { + t.Helper() + stream := &requestTaskSecurityStream{ + ctx: metadata.NewIncomingContext(context.Background(), metadata.Pairs( + "client_secret", secret, + "client_uuid", uuid, + )), + results: []*pb.TaskResult{result}, + } + err := NewNezhaHandler().RequestTask(stream) + if !errors.Is(err, context.Canceled) { + t.Fatalf("expected RequestTask to finish after test result, got %v", err) + } +} + +func runApplyConfigAuthzReconnect(t *testing.T, secret, uuid string) []*pb.Task { + t.Helper() + var sent []*pb.Task + stream := &requestTaskSecurityStream{ + ctx: metadata.NewIncomingContext(context.Background(), metadata.Pairs( + "client_secret", secret, + "client_uuid", uuid, + )), + onSend: func(task *pb.Task) { + sent = append(sent, task) + }, + } + err := NewNezhaHandler().RequestTask(stream) + if !errors.Is(err, context.Canceled) { + t.Fatalf("expected RequestTask to finish after reconnect probe, got %v", err) + } + return sent +} + +// Finding B regression: during the agent's 10s delayed ApplyConfig swap +// window, the agent still talks to the dashboard with the OLD (FromUserID) +// secret. After a cancel/fail/timeout, a registered revert delivery is the +// only signal that lets the eventually-arriving new-secret reconnect +// recover. The previous implementation cleared revertDeliveries from ANY +// successful old-secret authentication — including ReportSystemInfo2 from +// the periodic reportHost path — so a single old-secret RPC during the +// timer window could destroy the rollback record before the agent ever +// actually swapped secrets. Clearing the delivery is only safe when the +// auth call also gets a chance to consume it by pushing the rollback, +// which only the RequestTask handler does via OnAgentReconnect. +func TestReportSystemInfoDoesNotClearRevertDeliveryForOldSecret(t *testing.T) { + srv := &model.Server{ + Common: model.Common{ID: 23, UserID: 200}, + UUID: "23232323-2323-2323-2323-232323232323", + Name: "preserve-revert-delivery", + } + users := map[uint64]model.UserInfo{ + 200: {Role: model.RoleMember, AgentSecret: "from-secret-23"}, + 300: {Role: model.RoleMember, AgentSecret: "to-secret-23"}, + } + secrets := map[string]uint64{ + "from-secret-23": 200, + "to-secret-23": 300, + } + setupApplyConfigAuthzFixture(t, []*model.Server{srv}, users, secrets) + + tr := initiateAndRegisterPendingTransfer(t, srv.ID, 200, 300, 1) + if _, err := singleton.ServerTransferShared.Cancel(tr.ID); err != nil { + t.Fatalf("cancel transfer: %v", err) + } + if _, ok := singleton.ServerTransferShared.LookupRevertDelivery(srv.ID); !ok { + t.Fatal("precondition: cancel must have registered a revert delivery") + } + + // Simulate the agent's periodic reportHost calling ReportSystemInfo2 + // with the still-current (FromUserID) secret during the 10s pending + // ApplyConfig window. Must succeed (server already reverted to + // FromUserID) but must NOT clear the revert delivery — the agent has + // not yet swapped secrets, and destroying the only recovery record + // now would lock the agent out once its timer fires. + ctx := metadata.NewIncomingContext(context.Background(), metadata.Pairs( + "client_secret", "from-secret-23", + "client_uuid", srv.UUID, + )) + if _, err := NewNezhaHandler().ReportSystemInfo2(ctx, &pb.Host{}); err != nil { + t.Fatalf("ReportSystemInfo2 with old (FromUserID) secret must succeed after revert, got %v", err) + } + + if _, ok := singleton.ServerTransferShared.LookupRevertDelivery(srv.ID); !ok { + t.Fatal("non-RequestTask auth with old secret must NOT clear revert delivery; it cannot push the rollback, so destroying the record locks out the eventually-switched agent") + } +} + +// Regression: when cancel/fail/timeout happens while the agent is offline +// (its only TaskStream is gone), pushRevertIfOnline is a no-op and the +// revertDelivery is the only signal we have left. The agent will reconnect +// *with the original FromUserID secret* (its in-memory liveCredentials still +// points at the secret it had before the swap), and that very reconnect must +// be the one that delivers the rollback ApplyConfig — otherwise the agent's +// 10s reload timer eventually commits the new secret and the dashboard, which +// already restored ownership to FromUserID, rejects every subsequent connect. +// +// The previous implementation cleared the revertDelivery inside +// authorizeAgentForUUIDWithRevertRecovery *before* RequestTask reached +// OnAgentReconnect, so the rollback push that OnAgentReconnect relies on +// (LookupRevertDelivery → pushRevertIfOnline) found nothing and the agent +// got no rollback at all. +func TestRequestTaskCancelledTransferDeliversRollbackOnOldSecretReconnect(t *testing.T) { + srv := &model.Server{ + Common: model.Common{ID: 24, UserID: 200}, + UUID: "24242424-2424-2424-2424-242424242424", + Name: "old-secret-rollback", + } + users := map[uint64]model.UserInfo{ + 200: {Role: model.RoleMember, AgentSecret: "rollback-from-secret"}, + 300: {Role: model.RoleMember, AgentSecret: "rollback-to-secret"}, + } + secrets := map[string]uint64{ + "rollback-from-secret": 200, + "rollback-to-secret": 300, + } + setupApplyConfigAuthzFixture(t, []*model.Server{srv}, users, secrets) + + tr := initiateAndRegisterPendingTransfer(t, srv.ID, 200, 300, 1) + // Cancel while the agent is offline — the in-memory TaskStream is nil + // (we never attached one), so pushRevertIfOnline silently no-ops. + if _, err := singleton.ServerTransferShared.Cancel(tr.ID); err != nil { + t.Fatalf("cancel transfer: %v", err) + } + if _, ok := singleton.ServerTransferShared.LookupRevertDelivery(srv.ID); !ok { + t.Fatal("precondition: cancel while offline must leave a revert delivery for the eventual reconnect") + } + + // Agent now reconnects with its original FromUserID secret (it never + // received the new-secret ApplyConfig because it was offline). This + // RequestTask must deliver the rollback so the agent's reload timer + // supersedes onto the correct credential. + sent := runApplyConfigAuthzReconnect(t, "rollback-from-secret", srv.UUID) + if len(sent) != 1 { + t.Fatalf("expected one rollback ApplyConfig task on old-secret reconnect, got %d", len(sent)) + } + var settled model.ServerTransfer + if err := singleton.DB.First(&settled, tr.ID).Error; err != nil { + t.Fatalf("reload transfer: %v", err) + } + if !strings.Contains(sent[0].Data, settled.RevertHandshakeSecret) { + t.Fatalf("old-secret reconnect rollback must carry the per-transfer RevertHandshakeSecret, got %q", sent[0].Data) + } + if strings.Contains(sent[0].Data, "rollback-from-secret") || strings.Contains(sent[0].Data, "rollback-to-secret") { + t.Fatalf("user-global AgentSecrets must never appear in transfer payloads, got %q", sent[0].Data) + } +} + +// FORWARD-RECOVERY end-to-end: the exact production scenario the fix +// targets. PushIfOnline only ever delivers t.HandshakeSecret, the agent's +// 10s timer commits it to disk, the operator Cancels in that 10s window +// (revert push misses because the stream had no agent yet, or arrived +// before the forward apply finished). The agent reconnects with the +// forward HandshakeSecret it has on disk. RequestTask MUST accept that +// auth and then deliver one rollback ApplyConfig carrying the per-transfer +// RevertHandshakeSecret so the agent's next reload rotates onto the correct +// credential. Without this, the agent has no path back into the dashboard. +func TestRequestTaskForwardHandshakeSecretReconnectAfterCancelDeliversRollback(t *testing.T) { + srv := &model.Server{ + Common: model.Common{ID: 31, UserID: 200}, + UUID: "31313131-3131-3131-3131-313131313131", + Name: "forward-recovery", + } + users := map[uint64]model.UserInfo{ + 200: {Role: model.RoleMember, AgentSecret: "fr-from-secret"}, + 300: {Role: model.RoleMember, AgentSecret: "fr-to-secret"}, + } + secrets := map[string]uint64{ + "fr-from-secret": 200, + "fr-to-secret": 300, + } + setupApplyConfigAuthzFixture(t, []*model.Server{srv}, users, secrets) + + tr := initiateAndRegisterPendingTransfer(t, srv.ID, 200, 300, 1) + forward := tr.HandshakeSecret + if forward == "" { + t.Fatal("precondition: pending transfer must carry a forward HandshakeSecret") + } + + if _, err := singleton.ServerTransferShared.Cancel(tr.ID); err != nil { + t.Fatalf("dashboard Cancel must succeed: %v", err) + } + + sent := runApplyConfigAuthzReconnect(t, forward, srv.UUID) + if len(sent) != 1 { + t.Fatalf("forward-secret reconnect after Cancel must deliver one rollback ApplyConfig task, got %d", len(sent)) + } + var settled model.ServerTransfer + if err := singleton.DB.First(&settled, tr.ID).Error; err != nil { + t.Fatalf("reload transfer: %v", err) + } + if !strings.Contains(sent[0].Data, settled.RevertHandshakeSecret) { + t.Fatalf("rollback delivered after forward-secret recovery must carry the per-transfer RevertHandshakeSecret, got %q", sent[0].Data) + } + if strings.Contains(sent[0].Data, "fr-from-secret") || strings.Contains(sent[0].Data, "fr-to-secret") { + t.Fatalf("user-global AgentSecrets must never appear in transfer payloads, got %q", sent[0].Data) + } +} diff --git a/service/rpc/auth.go b/service/rpc/auth.go index 02f1418b..9f8fae12 100644 --- a/service/rpc/auth.go +++ b/service/rpc/auth.go @@ -3,6 +3,7 @@ package rpc import ( "context" "fmt" + "log" "strings" petname "github.com/dustinkirkland/golang-petname" @@ -21,6 +22,18 @@ type authHandler struct { } func (a *authHandler) Check(ctx context.Context) (uint64, error) { + return a.check(ctx) +} + +func (a *authHandler) CheckRequestTask(ctx context.Context) (uint64, error) { + return a.check(ctx) +} + +// 所有 auth caller 走完全相同的 ServerTransfer dual-secret 容忍策略。 +// revertDelivery 不在 auth 阶段消费 —— 真正派发 rollback ApplyConfig 的 +// pushRevertIfOnline 才有资格清理它,否则 auth 提前清就会让 OnAgentReconnect +// 找不到 recovery 记录,agent 10s timer 一到就锁死在被拒绝的新 secret 上。 +func (a *authHandler) check(ctx context.Context) (uint64, error) { md, ok := metadata.FromIncomingContext(ctx) if !ok { return 0, status.Errorf(codes.Unauthenticated, "获取 metaData 失败") @@ -37,6 +50,107 @@ func (a *authHandler) Check(ctx context.Context) (uint64, error) { ip, _ := ctx.Value(model.CtxKeyRealIP{}).(string) + var clientUUID string + if value, ok := md["client_uuid"]; ok { + clientUUID = value[0] + } + + if _, err := uuid.ParseUUID(clientUUID); err != nil { + // Keep this counter on the same trigger surface as the + // unknown-secret path below: an attacker who pairs a bad secret + // with a malformed/missing UUID otherwise bypasses + // WAFBlockReasonTypeAgentAuthFail entirely and gets unbounded + // retries (TestAuthBadSecret*InvalidUUIDStillIncrementsAgentAuthFailWAF). + model.BlockIP(singleton.DB, ip, model.WAFBlockReasonTypeAgentAuthFail, model.BlockIDgRPC) + return 0, status.Error(codes.Unauthenticated, "客户端 UUID 不合法") + } + + // Per-transfer handshake secret path: ApplyConfig delivers a random + // per-transfer token instead of the destination user's global AgentSecret + // (see PushIfOnline). When the agent reconnects under that token the auth + // layer recognises it here, scoped to the matching server UUID, and + // promotes the transfer to Verified. The user-global secret lookup below + // continues to handle every non-transfer agent, plus the still-tolerated + // previous-owner secret during the Pending window. Checked before the + // global lookup so the handshake-secret token can never collide with + // some other user's accidental match. + if singleton.ServerTransferShared != nil { + if t, ok := singleton.ServerTransferShared.LookupByHandshakeSecret(clientSecret); ok { + cid, found := singleton.ServerShared.UUIDToID(clientUUID) + if !found || cid != t.ServerID { + return 0, status.Error(codes.Unauthenticated, "transfer handshake secret bound to a different server") + } + // Auth via per-transfer HandshakeSecret succeeds only when + // MarkVerified actually performs the Pending → Verified + // transition. A lost CAS (concurrent Cancel/Fail/Timeout) + // means the credential is stale; the verifiedHandshakes + // fallthrough below will still admit it if it had been + // promoted by a successful previous reconnect, otherwise it + // is rejected. + verified, _, err := singleton.ServerTransferShared.MarkVerified(t.ServerID, t.ID) + if err != nil { + log.Printf("NEZHA>> ServerTransfer MarkVerified(cid=%d) via handshake secret failed: %v", t.ServerID, err) + return 0, status.Error(codes.Unauthenticated, "transfer handshake verification failed") + } + if verified { + model.UnblockIP(singleton.DB, ip, model.BlockIDgRPC) + return t.ServerID, nil + } + } + // Bounded terminal-recovery window: a transfer was Cancel/Fail/ + // Timeout-ed and the agent may still be presenting either of its + // per-transfer secrets. Single lookup + kind switch: + // + // forward — agent committed t.HandshakeSecret to disk before + // the dashboard observed MarkVerified. Admit so + // RequestTask → OnAgentReconnect can deliver the + // rollback ApplyConfig. DO NOT call MarkVerified + // (transfer is terminal) and DO NOT promote into + // verifiedHandshakes (the agent's stable post-rollback + // credential will be the revert secret, not this one). + // + // revert — agent has applied the rollback and presented + // t.RevertHandshakeSecret. Promote via + // MarkRevertDelivered so the credential survives + // past the recovery window (~24h sweep). + // + // SECURITY: terminalSecretRecovery is only populated by + // revertTransition. A stolen per-transfer secret on a transfer + // whose terminal status was forged in the DB never reaches this + // table — TestAuthHandshakeSecretRejectedAfterTransferTerminated + // pins that path closed. + if t, kind, ok := singleton.ServerTransferShared.LookupByTerminalSecretRecovery(clientSecret); ok { + cid, found := singleton.ServerShared.UUIDToID(clientUUID) + if !found || cid != t.ServerID { + return 0, status.Error(codes.Unauthenticated, "transfer terminal-recovery secret bound to a different server") + } + model.UnblockIP(singleton.DB, ip, model.BlockIDgRPC) + if kind == singleton.TerminalRecoveryRevert { + if err := singleton.ServerTransferShared.MarkRevertDelivered(t.ServerID, t.ID); err != nil { + log.Printf("NEZHA>> ServerTransfer MarkRevertDelivered(server=%d transfer=%d) failed: %v", t.ServerID, t.ID, err) + } + } + return t.ServerID, nil + } + // Post-MarkVerified path: the agent's persisted client_secret is + // the per-transfer HandshakeSecret (PushIfOnline never delivers a + // user-global secret), and no follow-up ApplyConfig swaps it back + // out. So every reconnect after the first one — stream drop, agent + // restart, etc. — must still match this credential, bound strictly + // to (serverID, UUID). The match is constrained to a single server + // because the handshake secret was generated per-transfer; it does + // not unlock any other agent. A new transfer for the same server + // invalidates the entry inside Register, closing this acceptance + // window before the next HandshakeSecret takes over. + if cid, ok := singleton.ServerTransferShared.LookupServerByVerifiedHandshakeSecret(clientSecret); ok { + if uuidCID, found := singleton.ServerShared.UUIDToID(clientUUID); found && uuidCID == cid { + model.UnblockIP(singleton.DB, ip, model.BlockIDgRPC) + return cid, nil + } + return 0, status.Error(codes.Unauthenticated, "transfer verified handshake secret bound to a different server") + } + } + singleton.UserLock.RLock() userId, ok := singleton.AgentSecretToUserId[clientSecret] if !ok { @@ -48,15 +162,6 @@ func (a *authHandler) Check(ctx context.Context) (uint64, error) { model.UnblockIP(singleton.DB, ip, model.BlockIDgRPC) - var clientUUID string - if value, ok := md["client_uuid"]; ok { - clientUUID = value[0] - } - - if _, err := uuid.ParseUUID(clientUUID); err != nil { - return 0, status.Error(codes.Unauthenticated, "客户端 UUID 不合法") - } - clientID, hasID, err := authorizeAgentForUUID(userId, clientUUID) if err != nil { return 0, status.Error(codes.Unauthenticated, err.Error()) @@ -90,6 +195,16 @@ func (a *authHandler) Check(ctx context.Context) (uint64, error) { // an agent persistently fails with "client UUID does not belong to the // agent secret owner", it pins down which user's secret has been reused // against a server they don't own. +// +// Server transfer interaction: while a ServerTransfer is Pending for this +// server, the agent is still authenticating with the previous owner's +// AgentSecret (the new secret has not yet propagated). To keep that agent +// online during the rollover, accept userId==FromUserID for the duration of +// the pending window. The dual-secret tolerance is narrowly scoped to the +// affected server only — every other agent of either user is unaffected. +// Once the agent reconnects under the new owner's secret (userId==ToUserID +// matching server.UserID), MarkVerified promotes the transfer and closes +// the tolerance window. func authorizeAgentForUUID(userId uint64, clientUUID string) (clientID uint64, hasID bool, err error) { cid, found := singleton.ServerShared.UUIDToID(clientUUID) if !found { @@ -106,8 +221,51 @@ func authorizeAgentForUUID(userId uint64, clientUUID string) (clientID uint64, h // agent secrets, so keep it compatible by allowing any existing UUID. return cid, true, nil } - if server.UserID != userId { - return 0, false, fmt.Errorf("client UUID does not belong to the agent secret owner") + if server.GetUserID() == userId { + // SECURITY: while a transfer is Pending, Server.UserID has already + // been flipped to ToUserID by Register, so userId==Server.UserID + // here also matches the destination user's user-global AgentSecret. + // PushIfOnline only delivers the per-transfer HandshakeSecret on + // the wire; the destination user's global AgentSecret is never + // pushed to the agent, so a reconnect under that secret is not + // proof of agent rotation. Admitting it would let the destination + // user — who can see Server.UUID — authenticate as the agent + // during the Pending window. Reject the user-global secret until + // the transfer settles; the HandshakeSecret path in check() is + // the only valid promotion route. + if singleton.ServerTransferShared != nil { + if _, ok := singleton.ServerTransferShared.LookupPending(cid); ok { + return 0, false, fmt.Errorf("destination user's global AgentSecret cannot authenticate during a pending transfer; agent must rotate to per-transfer HandshakeSecret") + } + } + return cid, true, nil } - return cid, true, nil + // server.UserID != userId — normally an impersonation attempt. Allow it + // only when a ServerTransfer for this server is Pending AND the secret in + // hand is the previous owner's (FromUserID), OR when a recently terminated + // transfer left a revert-delivery for FromUserID and the agent is still + // presenting its pre-transfer global secret. + // + // SECURITY: we deliberately do NOT accept the destination user's global + // AgentSecret on the LookupRevertDelivery path. PushIfOnline only ever + // delivers per-transfer HandshakeSecret / RevertHandshakeSecret to the + // agent — the ToUserID global secret never travels over the wire — so a + // reconnect under that credential is not proof of agent rotation; it can + // only come from the destination user themselves, who can see Server.UUID + // once Register flips Server.UserID. Admitting it would let that user + // impersonate the agent during the rollback window, trigger + // pushRevertIfOnline to leak RevertHandshakeSecret, and then be promoted + // into verifiedHandshakes via MarkRevertDelivered. The legitimate recovery + // paths are: FromUserID global secret (handled below), forward + // HandshakeSecret and RevertHandshakeSecret (handled by the + // terminalSecretRecovery / verifiedHandshakes lookups in check()). + if singleton.ServerTransferShared != nil { + if t, ok := singleton.ServerTransferShared.LookupRevertDelivery(cid); ok && t.FromUserID == userId { + return cid, true, nil + } + if t, ok := singleton.ServerTransferShared.LookupPending(cid); ok && t.FromUserID == userId { + return cid, true, nil + } + } + return 0, false, fmt.Errorf("client UUID does not belong to the agent secret owner") } diff --git a/service/rpc/auth_test.go b/service/rpc/auth_test.go index d61afec4..797da940 100644 --- a/service/rpc/auth_test.go +++ b/service/rpc/auth_test.go @@ -1,15 +1,93 @@ package rpc import ( + "context" + "errors" "testing" + "google.golang.org/grpc/metadata" "gorm.io/driver/sqlite" "gorm.io/gorm" "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/pkg/utils" "github.com/nezhahq/nezha/service/singleton" ) +// authCheckWithSecret drives (*authHandler).check end-to-end via the same +// gRPC metadata path the real RPC handler uses. Tests rely on it to assert +// what a real reconnect — secret + UUID supplied on the wire — would do. +func authCheckWithSecret(secret, uuid string) (uint64, error) { + ctx := metadata.NewIncomingContext(context.Background(), metadata.Pairs( + "client_secret", secret, + "client_uuid", uuid, + )) + return (&authHandler{}).Check(ctx) +} + +// authHandshakeUUID is RFC4122-shaped so it survives the uuid.ParseUUID gate +// at the top of check(); setupAuthAgentFixture's "uuid-alice" / "uuid-bob" +// only work for callers that bypass check() and exercise the inner helpers. +const authHandshakeUUID = "11111111-1111-1111-1111-111111111111" + +// setupAuthHandshakeFixture seeds a single server (id=11, owner=user 100, +// real UUID) plus the user-secret tables so the global-secret fall-through +// in check() has something to match. Mirrors setupAuthAgentFixture's reset +// discipline but additionally restores AgentSecretToUserId / UserInfoMap. +func setupAuthHandshakeFixture(t *testing.T) func() { + t.Helper() + originalDB := singleton.DB + originalServerShared := singleton.ServerShared + originalServerTransferShared := singleton.ServerTransferShared + originalUserInfoMap := singleton.UserInfoMap + originalAgentSecretToUserId := singleton.AgentSecretToUserId + + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + if err != nil { + t.Fatalf("open db: %v", err) + } + if err := db.AutoMigrate(&model.Server{}, &model.ServerTransfer{}, &model.WAF{}); err != nil { + t.Fatalf("migrate: %v", err) + } + if err := db.Create(&model.Server{ + Common: model.Common{ID: 11, UserID: 100}, + UUID: authHandshakeUUID, + Name: "handshake-srv", + }).Error; err != nil { + t.Fatalf("create handshake server: %v", err) + } + singleton.DB = db + singleton.ServerShared = singleton.NewServerClass() + srv := &model.Server{Common: model.Common{ID: 11, UserID: 100}, UUID: authHandshakeUUID, Name: "handshake-srv"} + model.InitServer(srv) + singleton.ServerShared.Update(srv, authHandshakeUUID) + singleton.ServerTransferShared = singleton.NewServerTransferClass() + + singleton.UserLock.Lock() + singleton.UserInfoMap = map[uint64]model.UserInfo{ + 100: {Role: model.RoleMember, AgentSecret: "alice-global"}, + 200: {Role: model.RoleMember, AgentSecret: "bob-global"}, + } + singleton.AgentSecretToUserId = map[string]uint64{ + "alice-global": 100, + "bob-global": 200, + } + singleton.UserLock.Unlock() + + return func() { + if singleton.ServerTransferShared != nil { + singleton.ServerTransferShared.Stop() + } + singleton.DB = originalDB + singleton.ServerShared = originalServerShared + singleton.ServerTransferShared = originalServerTransferShared + singleton.UserLock.Lock() + singleton.UserInfoMap = originalUserInfoMap + singleton.AgentSecretToUserId = originalAgentSecretToUserId + singleton.UserLock.Unlock() + } +} + // setupAuthAgentFixture seeds an in-memory DB and ServerShared with two // servers belonging to different users so we can assert that a secret bound // to user A cannot resolve a server UUID owned by user B. @@ -17,12 +95,13 @@ func setupAuthAgentFixture(t *testing.T) func() { t.Helper() originalDB := singleton.DB originalServerShared := singleton.ServerShared + originalServerTransferShared := singleton.ServerTransferShared db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) if err != nil { t.Fatalf("open db: %v", err) } - if err := db.AutoMigrate(&model.Server{}); err != nil { + if err := db.AutoMigrate(&model.Server{}, &model.ServerTransfer{}); err != nil { t.Fatalf("migrate: %v", err) } if err := db.Create(&model.Server{ @@ -41,10 +120,15 @@ func setupAuthAgentFixture(t *testing.T) func() { } singleton.DB = db singleton.ServerShared = singleton.NewServerClass() + singleton.ServerTransferShared = singleton.NewServerTransferClass() return func() { + if singleton.ServerTransferShared != nil { + singleton.ServerTransferShared.Stop() + } singleton.DB = originalDB singleton.ServerShared = originalServerShared + singleton.ServerTransferShared = originalServerTransferShared } } @@ -99,3 +183,424 @@ func TestAuthorizeAgentForUUIDPermitsUnknownUUIDForRegistration(t *testing.T) { t.Fatalf("hasID must be false for unknown UUID, got cid=%d", cid) } } + +// initiatePendingTransfer mirrors the controller flow used by the batch-move +// endpoint to drive ownership through ServerTransferShared. Tests use it to +// set up the auth-tolerance window with the Server row already flipped to +// ToUserID. Returns nothing; callers use ServerTransferShared.LookupPending +// to fetch the row if they need it. +func initiatePendingTransfer(t *testing.T, serverID, fromUserID, toUserID uint64) { + t.Helper() + var created *model.ServerTransfer + err := singleton.DB.Transaction(func(tx *gorm.DB) error { + var err error + created, err = singleton.ServerTransferShared.Initiate(tx, serverID, fromUserID, toUserID, fromUserID) + return err + }) + if err != nil { + t.Fatalf("initiate pending transfer: %v", err) + } + singleton.ServerTransferShared.Register(created) +} + +// The auth-tolerance window: while a Pending transfer exists for this server, +// the old owner's AgentSecret must still authenticate this UUID — the agent +// hasn't received the new secret yet via ApplyConfig. Without this, every +// in-flight transfer would knock the affected agent offline immediately. +func TestAuthorizeAgentForUUIDAcceptsFromUserDuringPendingTransfer(t *testing.T) { + defer setupAuthAgentFixture(t)() + + // Alice initiates: server 1 moves from alice (100) to bob (200). + // Server.UserID is now 200; alice's agent still presents secret==100. + initiatePendingTransfer(t, 1, 100, 200) + + cid, hasID, err := authorizeAgentForUUID(100, "uuid-alice") + if err != nil { + t.Fatalf("FromUserID secret must be accepted during pending window, got %v", err) + } + if !hasID || cid != 1 { + t.Fatalf("expected (cid=1, hasID=true), got (cid=%d, hasID=%v)", cid, hasID) + } +} + +// Tolerance is narrowly scoped: an unrelated user's secret must NOT be +// accepted just because *some* transfer is in flight. Specifically, only +// secrets matching FromUserID or ToUserID get through. +func TestAuthorizeAgentForUUIDRejectsThirdPartyDuringPendingTransfer(t *testing.T) { + defer setupAuthAgentFixture(t)() + initiatePendingTransfer(t, 1, 100, 200) + + // userId=999 has nothing to do with this transfer. + _, _, err := authorizeAgentForUUID(999, "uuid-alice") + if err == nil { + t.Fatalf("third-party secret must be rejected even while a transfer is pending") + } +} + +// SECURITY: during a Pending transfer the destination user's user-global +// AgentSecret must NOT close the pending window. PushIfOnline only delivers +// the per-transfer HandshakeSecret on the wire, so a reconnect under the +// destination user's global AgentSecret is not proof of agent rotation — +// it could just be the destination user authenticating with their own +// secret + the now-visible Server.UUID. Reject it; only the per-transfer +// HandshakeSecret path may promote to Verified. +func TestAuthorizeAgentForUUIDRejectsToUserGlobalSecretDuringPendingTransfer(t *testing.T) { + defer setupAuthAgentFixture(t)() + initiatePendingTransfer(t, 1, 100, 200) + + if _, _, err := authorizeAgentForUUID(200, "uuid-alice"); err == nil { + t.Fatal("destination user's global AgentSecret must NOT authenticate during pending transfer; only per-transfer HandshakeSecret may close the window") + } + if !singleton.ServerTransferShared.HasPending(1) { + t.Fatal("pending transfer must survive a destination-user global AgentSecret reconnect") + } + + if _, _, err := authorizeAgentForUUID(100, "uuid-alice"); err != nil { + t.Fatalf("FromUser tolerance window must remain open while transfer is still Pending, got %v", err) + } +} + +// During the revert recovery window the destination user's global +// AgentSecret must NOT be accepted by authorizeAgentForUUID. PushIfOnline +// only delivers per-transfer HandshakeSecret / RevertHandshakeSecret on +// the wire, so a reconnect under the ToUserID global secret cannot come +// from the real agent — it can only come from the destination user +// themselves, who can see Server.UUID and would otherwise impersonate the +// agent during rollback, trigger pushRevertIfOnline to leak +// RevertHandshakeSecret, and get promoted via MarkRevertDelivered. +// Legitimate recovery goes through FromUserID's global secret, the +// forward HandshakeSecret, or the RevertHandshakeSecret. +func TestAuthorizeAgentForUUIDRejectsToUserGlobalSecretDuringRevertRecovery(t *testing.T) { + defer setupAuthAgentFixture(t)() + initiatePendingTransfer(t, 1, 100, 200) + pending, ok := singleton.ServerTransferShared.LookupPending(1) + if !ok { + t.Fatal("expected pending transfer") + } + if _, err := singleton.ServerTransferShared.Cancel(pending.ID); err != nil { + t.Fatalf("cancel transfer: %v", err) + } + + if _, _, err := authorizeAgentForUUID(200, "uuid-alice"); err == nil { + t.Fatal("destination user's global AgentSecret must NOT authenticate during revert recovery; only per-transfer HandshakeSecret / RevertHandshakeSecret may close the window") + } + if _, ok := singleton.ServerTransferShared.LookupRevertDelivery(1); !ok { + t.Fatal("rejected ToUserID auth must not consume the revert delivery — the real agent still needs it for the eventual per-transfer recovery") + } +} + +// Regression for finding A: after MarkVerified deletes the pending entry, +// the agent's persisted ClientSecret is still the per-transfer +// HandshakeSecret (PushIfOnline only ever delivered that value). The very +// next reconnect — gRPC stream drop, agent restart, network blip — must +// keep authenticating, otherwise the agent silently locks itself out on +// the now-orphaned handshake token. There is no follow-up ApplyConfig +// path that swaps the agent over to the destination user's stable +// AgentSecret, so auth itself has to keep treating the post-Verified +// HandshakeSecret as a valid credential for that server. +func TestAuthHandshakeSecretStillAuthenticatesAfterMarkVerified(t *testing.T) { + defer setupAuthHandshakeFixture(t)() + + initiatePendingTransfer(t, 11, 100, 200) + pending, ok := singleton.ServerTransferShared.LookupPending(11) + if !ok { + t.Fatal("expected pending transfer") + } + handshakeSecret := pending.HandshakeSecret + if handshakeSecret == "" { + t.Fatal("precondition: pending transfer must carry a HandshakeSecret") + } + + cid, err := authCheckWithSecret(handshakeSecret, authHandshakeUUID) + if err != nil { + t.Fatalf("first reconnect with HandshakeSecret must promote the transfer, got %v", err) + } + if cid != 11 { + t.Fatalf("first reconnect must resolve to server 11, got %d", cid) + } + if singleton.ServerTransferShared.HasPending(11) { + t.Fatal("MarkVerified must have cleared the pending index after the handshake reconnect") + } + + if _, err := authCheckWithSecret(handshakeSecret, authHandshakeUUID); err != nil { + t.Fatalf("second reconnect with the same HandshakeSecret must still authenticate (the agent has no other credential to present until a final hand-off completes); got %v", err) + } +} + +// First successful auth with RevertHandshakeSecret proves the agent has +// applied the rollback (10s reload + applyPendingReload have committed +// the secret to disk). At that point the auth path must promote the +// secret into the long-term verifiedHandshakes map and consume the +// temporary revertDeliveries entry — otherwise the only acceptance path +// is LookupByRevertHandshakeSecret, which prunes after +// defaultRevertDeliveryRecoveryWindow and leaves the agent locked out +// ~24h later. See ServerTransferClass.MarkRevertDelivered. +func TestAuthRevertHandshakeSecretPromotesToVerifiedAndKeepsAuthenticating(t *testing.T) { + defer setupAuthHandshakeFixture(t)() + + initiatePendingTransfer(t, 11, 100, 200) + pending, ok := singleton.ServerTransferShared.LookupPending(11) + if !ok { + t.Fatal("expected pending transfer") + } + if _, err := singleton.ServerTransferShared.Cancel(pending.ID); err != nil { + t.Fatalf("cancel transfer to register a revert delivery: %v", err) + } + revert, ok := singleton.ServerTransferShared.LookupRevertDelivery(11) + if !ok { + t.Fatal("precondition: cancel must have registered a revert delivery") + } + revertHandshake := revert.RevertHandshakeSecret + if revertHandshake == "" { + t.Fatal("precondition: revert delivery must carry a RevertHandshakeSecret") + } + + if _, err := authCheckWithSecret(revertHandshake, authHandshakeUUID); err != nil { + t.Fatalf("first auth with RevertHandshakeSecret must succeed, got %v", err) + } + + if _, ok := singleton.ServerTransferShared.LookupRevertDelivery(11); ok { + t.Fatal("first successful auth must consume the temporary revertDelivery — the credential is now promoted to the long-term map") + } + + sid, ok := singleton.ServerTransferShared.LookupServerByVerifiedHandshakeSecret(revertHandshake) + if !ok || sid != 11 { + t.Fatalf("RevertHandshakeSecret must be promoted into verifiedHandshakes; lookup got (sid=%d, ok=%v)", sid, ok) + } + + if _, err := authCheckWithSecret(revertHandshake, authHandshakeUUID); err != nil { + t.Fatalf("second auth via the promoted verifiedHandshakes path must still succeed, got %v", err) + } +} + +// HIGH security regression: if a transfer has already been Cancelled/Failed/ +// Timed out, its HandshakeSecret must NEVER authenticate. Today auth.check +// calls MarkVerified on the lookup result and treats RowsAffected==0 as +// success, so an attacker who learned the per-transfer HandshakeSecret +// (e.g. previous owner whose stream was hijacked during Pending) can +// authenticate inside the narrow race window where revertTransition has +// changed DB status but not yet deleted the in-memory pending entry, or +// after that window simply because the swallowed return is `return +// t.ServerID, nil`. +// +// Expected: when the transfer row is no longer Pending, auth must reject +// the HandshakeSecret entirely. +func TestAuthHandshakeSecretRejectedAfterTransferTerminated(t *testing.T) { + defer setupAuthHandshakeFixture(t)() + + initiatePendingTransfer(t, 11, 100, 200) + pending, ok := singleton.ServerTransferShared.LookupPending(11) + if !ok { + t.Fatal("expected pending transfer") + } + handshakeSecret := pending.HandshakeSecret + + // Settle the DB row to Cancelled WITHOUT touching the in-memory + // pending entry. This reproduces the race window in revertTransition + // between the DB CAS and the c.mu.Lock that deletes the pending + // entry; LookupByHandshakeSecret still hits. + if err := singleton.DB.Model(&model.ServerTransfer{}). + Where("id = ?", pending.ID). + Update("status", model.ServerTransferStatusCancelled).Error; err != nil { + t.Fatalf("simulate concurrent cancel: %v", err) + } + + _, err := authCheckWithSecret(handshakeSecret, authHandshakeUUID) + if err == nil { + t.Fatal("HandshakeSecret on a terminated transfer must be rejected — auth swallowed MarkVerified RowsAffected==0 and returned success, enabling auth bypass with a stale per-transfer secret") + } +} + +// HIGH security regression: the auth tolerance window for the old owner's +// global AgentSecret must close in lockstep with MarkVerified. Holding c.mu +// across the DB CAS, the c.pending delete and the verifiedHandshakes write +// inside MarkVerified makes those three steps a single observable event for +// any auth-path lookup taking c.mu.RLock; once MarkVerified returns +// verified=true, no later authorizeAgentForUUID can still see the pending +// entry that previously admitted FromUserID. +func TestAuthOldOwnerSecretRejectedOnceTransferIsVerifiedInDB(t *testing.T) { + defer setupAuthHandshakeFixture(t)() + + initiatePendingTransfer(t, 11, 100, 200) + pending, ok := singleton.ServerTransferShared.LookupPending(11) + if !ok { + t.Fatal("expected pending transfer") + } + + verified, _, err := singleton.ServerTransferShared.MarkVerified(11, pending.ID) + if err != nil { + t.Fatalf("MarkVerified must succeed for a fresh pending: %v", err) + } + if !verified { + t.Fatal("MarkVerified must report verified=true for a fresh pending") + } + + if _, _, err := authorizeAgentForUUID(100, authHandshakeUUID); err == nil { + t.Fatal("old owner's global AgentSecret must be rejected once MarkVerified has returned — the auth tolerance window must not outlive the verified transition") + } +} + +// FORWARD-RECOVERY (HIGH): symmetric to TestAuthHandshakeSecretRejectedAfter +// TransferTerminated. That test pokes the DB directly to simulate an +// attacker who learned the per-transfer forward HandshakeSecret outside +// of any dashboard-driven cancellation; auth must reject. This test +// exercises the OTHER scenario: a legitimate agent that already wrote the +// forward HandshakeSecret to disk via the 10s reload timer, and the +// dashboard cancels the transfer via the normal Cancel API (which goes +// through revertTransition). The agent's next reconnect presents the +// forward HandshakeSecret. Auth must authenticate it so RequestTask can +// run OnAgentReconnect and push the RevertHandshakeSecret rollback — +// otherwise the agent is permanently locked out and the operator has to +// SSH in and edit the config by hand. +// +// The distinguishing signal is whether revertTransition was the one that +// settled the row: it populates terminalForwardRecovery; a direct DB +// poke does not. +func TestAuthForwardHandshakeSecretAcceptedAfterDashboardCancel(t *testing.T) { + defer setupAuthHandshakeFixture(t)() + + initiatePendingTransfer(t, 11, 100, 200) + pending, ok := singleton.ServerTransferShared.LookupPending(11) + if !ok { + t.Fatal("expected pending transfer") + } + forward := pending.HandshakeSecret + + if _, err := singleton.ServerTransferShared.Cancel(pending.ID); err != nil { + t.Fatalf("dashboard Cancel must succeed: %v", err) + } + + cid, err := authCheckWithSecret(forward, authHandshakeUUID) + if err != nil { + t.Fatalf("forward HandshakeSecret must authenticate after dashboard Cancel so RequestTask can deliver the rollback; got %v", err) + } + if cid != 11 { + t.Fatalf("forward HandshakeSecret must resolve to its bound server, got cid=%d", cid) + } + + if _, ok := singleton.ServerTransferShared.LookupServerByVerifiedHandshakeSecret(forward); ok { + t.Fatal("forward HandshakeSecret on a terminated transfer must NOT be promoted into verifiedHandshakes — promotion would outlive the bounded recovery window and turn a cancelled credential into a permanent one") + } +} + +// wafAgentAuthFailCount returns the recorded WAF count for the given IP + +// gRPC block identifier. Used by the bad-credential WAF tests to assert +// FirstOrCreate / UPDATE actually fired. +func wafAgentAuthFailCount(t *testing.T, ip string) uint64 { + t.Helper() + bin, err := utils.IPStringToBinary(ip) + if err != nil { + t.Fatalf("ip parse: %v", err) + } + var w model.WAF + res := singleton.DB.Where("ip = ? AND block_identifier = ?", bin, model.BlockIDgRPC).First(&w) + if res.Error != nil { + if errors.Is(res.Error, gorm.ErrRecordNotFound) { + return 0 + } + t.Fatalf("query waf: %v", res.Error) + } + return w.Count +} + +// authCheckFromIP feeds an attacker IP through the real Check entry point +// so the WAF BlockIP path observes a non-empty CtxKeyRealIP. authCheckWithSecret +// uses a bare context.Background which keeps the IP empty and short-circuits +// BlockIP(ip == ""), masking the very regression these tests want to pin. +func authCheckFromIP(secret, uuid, ip string) (uint64, error) { + ctx := metadata.NewIncomingContext(context.Background(), metadata.Pairs( + "client_secret", secret, + "client_uuid", uuid, + )) + ctx = context.WithValue(ctx, model.CtxKeyRealIP{}, ip) + return (&authHandler{}).Check(ctx) +} + +// REGRESSION: the new per-transfer handshake path moved the client_uuid +// validation in front of the global AgentSecretToUserId lookup. A bad +// secret paired with a malformed/missing UUID now short-circuits to +// "客户端 UUID 不合法" and skips the BlockIP(WAFBlockReasonTypeAgentAuthFail) +// counter the previous implementation incremented. That counter is the +// only thing throttling brute-force on agent secrets — losing it lets an +// attacker enumerate secrets indefinitely just by also corrupting the +// UUID metadata. Both the missing-secret and bad-secret cases must still +// count toward AgentAuthFail when the UUID is unusable. +func TestAuthBadSecretInvalidUUIDStillIncrementsAgentAuthFailWAF(t *testing.T) { + defer setupAuthHandshakeFixture(t)() + const attackerIP = "203.0.113.7" + + if _, err := authCheckFromIP("definitely-not-a-real-secret", "not-a-uuid", attackerIP); err == nil { + t.Fatal("Check must reject bogus credentials") + } + + if got := wafAgentAuthFailCount(t, attackerIP); got == 0 { + t.Fatalf("bad client_secret + invalid client_uuid must still count toward WAFBlockReasonTypeAgentAuthFail; got count=%d", got) + } +} + +// Mirror of the above for the empty-UUID metadata path. uuid.ParseUUID("") +// also errors out, so the same auth-fail counting must apply — otherwise an +// attacker can just omit the metadata key entirely. +func TestAuthBadSecretEmptyUUIDStillIncrementsAgentAuthFailWAF(t *testing.T) { + defer setupAuthHandshakeFixture(t)() + const attackerIP = "203.0.113.8" + + if _, err := authCheckFromIP("another-bad-secret", "", attackerIP); err == nil { + t.Fatal("Check must reject bogus credentials") + } + + if got := wafAgentAuthFailCount(t, attackerIP); got == 0 { + t.Fatalf("bad client_secret + empty client_uuid must still count toward WAFBlockReasonTypeAgentAuthFail; got count=%d", got) + } +} + +// FORWARD-RECOVERY: forward secret bound to server A must not authenticate +// when presented with server B's UUID. Defence against an attacker who +// learns one server's forward secret and tries to attach it to a different +// agent during the recovery window. +func TestAuthForwardHandshakeSecretRejectedForDifferentUUID(t *testing.T) { + defer setupAuthHandshakeFixture(t)() + + initiatePendingTransfer(t, 11, 100, 200) + pending, ok := singleton.ServerTransferShared.LookupPending(11) + if !ok { + t.Fatal("expected pending transfer") + } + forward := pending.HandshakeSecret + + if _, err := singleton.ServerTransferShared.Cancel(pending.ID); err != nil { + t.Fatalf("dashboard Cancel must succeed: %v", err) + } + + const otherUUID = "22222222-2222-2222-2222-222222222222" + if _, err := authCheckWithSecret(forward, otherUUID); err == nil { + t.Fatal("forward HandshakeSecret must be rejected when paired with a different server UUID even during recovery — token is per-(server, transfer)") + } +} + +// SECURITY (P1): PushIfOnline only ever delivers the per-transfer +// HandshakeSecret to the real agent; the destination user's global +// AgentSecret is never sent on the wire and is therefore not proof of +// agent rotation. Server.UUID is visible to the destination user once +// Register flips Server.UserID, so admitting (ToUser global secret, real +// UUID) and calling MarkVerified would let the destination user clear +// the auth tolerance window for FromUser's secret (locking the real +// agent out) and flip transfer state to Verified without the agent ever +// applying the new credential. Only LookupByHandshakeSecret may promote. +func TestAuthDestinationUserGlobalSecretDoesNotVerifyPendingTransfer(t *testing.T) { + defer setupAuthHandshakeFixture(t)() + + initiatePendingTransfer(t, 11, 100, 200) + if !singleton.ServerTransferShared.HasPending(11) { + t.Fatal("precondition: pending transfer must be registered") + } + + cid, err := authCheckWithSecret("bob-global", authHandshakeUUID) + if err == nil { + t.Fatalf("destination user's global AgentSecret must not close the transfer's pending window; got cid=%d", cid) + } + + if !singleton.ServerTransferShared.HasPending(11) { + t.Fatal("pending transfer must survive a destination-user global AgentSecret reconnect; only the per-transfer HandshakeSecret may promote to Verified") + } +} diff --git a/service/rpc/io_stream.go b/service/rpc/io_stream.go index 23ad28f1..6a61a298 100644 --- a/service/rpc/io_stream.go +++ b/service/rpc/io_stream.go @@ -119,6 +119,35 @@ func (s *NezhaHandler) GetStream(streamId string) (*ioStreamContext, error) { return nil, errors.New("stream not found") } +// RevokeStreamsForServer tears down every IOStream whose targetServerID +// matches serverID. Called by the singleton package via the +// ServerTransferStreamRevocationHook on every transfer ownership +// transition — a stream the previous owner had open against this server +// must not survive into the new tenant, otherwise terminal/file-manager/NAT +// sessions become post-transfer hijack channels (effectively RCE). +// +// Underlying IO pipes are closed inline so the dashboard websocket loop +// sees EOF immediately rather than at the next idle-timeout. +func (s *NezhaHandler) RevokeStreamsForServer(serverID uint64) { + if serverID == 0 { + return + } + s.ioStreamMutex.Lock() + defer s.ioStreamMutex.Unlock() + for streamId, ctx := range s.ioStreams { + if ctx.targetServerID != serverID { + continue + } + if ctx.userIo != nil { + ctx.userIo.Close() + } + if ctx.agentIo != nil { + ctx.agentIo.Close() + } + delete(s.ioStreams, streamId) + } +} + func (s *NezhaHandler) CloseStream(streamId string) error { s.ioStreamMutex.Lock() defer s.ioStreamMutex.Unlock() @@ -136,6 +165,8 @@ func (s *NezhaHandler) CloseStream(streamId string) error { return nil } + + func (s *NezhaHandler) UserConnected(streamId string, userIo io.ReadWriteCloser) error { stream, err := s.GetStream(streamId) if err != nil { diff --git a/service/rpc/nezha.go b/service/rpc/nezha.go index 77c41c04..97325418 100644 --- a/service/rpc/nezha.go +++ b/service/rpc/nezha.go @@ -40,12 +40,21 @@ func NewNezhaHandler() *NezhaHandler { func (s *NezhaHandler) RequestTask(stream pb.NezhaService_RequestTaskServer) error { var clientID uint64 var err error - if clientID, err = s.Auth.Check(stream.Context()); err != nil { + if clientID, err = s.Auth.CheckRequestTask(stream.Context()); err != nil { return err } server, _ := singleton.ServerShared.Get(clientID) - server.TaskStream = stream + server.SetTaskStream(stream) + defer server.ClearTaskStreamIfCurrent(stream) + // If a transfer is mid-flight for this server, the agent has just brought + // up a fresh bidi stream — this is the moment to (re)deliver the + // ApplyConfig task carrying the new owner's AgentSecret. Pushes from + // dashboard mutation time are best-effort; this hook is the reliable + // re-delivery point that closes the offline-during-transfer gap. + if singleton.ServerTransferShared != nil { + singleton.ServerTransferShared.OnAgentReconnect(clientID) + } var result *pb.TaskResult for { result, err = stream.Recv() @@ -83,6 +92,27 @@ func (s *NezhaHandler) RequestTask(stream pb.NezhaService_RequestTaskServer) err } server.ConfigCache <- result.Data } + case model.TaskTypeServerTransferApply: + // Authorization: TaskResult.Id is attacker-controlled. Without + // the pending.ID == result.Id check below, agent A could cancel + // server B's in-flight transfer by spoofing B's transfer ID — + // same class of bug as commit 02129f1 in the cron path. + // Successful=true here is best-effort only; the authoritative + // verification is the agent's reconnect under the new secret. + if singleton.ServerTransferShared == nil { + continue + } + pending, ok := singleton.ServerTransferShared.LookupPending(clientID) + if !ok || pending.ID != result.GetId() { + log.Printf("NEZHA>> ServerTransferApply result ignored: clientID=%d reported transferID=%d but no matching pending transfer", clientID, result.GetId()) + continue + } + if result.GetSuccessful() { + continue + } + if _, err := singleton.ServerTransferShared.MarkFailed(result.GetId(), result.GetData()); err != nil { + log.Printf("NEZHA>> ServerTransfer MarkFailed(%d) failed: %v", result.GetId(), err) + } default: if model.IsServiceSentinelNeeded(result.GetType()) { singleton.ServiceSentinelShared.Dispatch(singleton.ReportData{ diff --git a/service/rpc/request_task_security_test.go b/service/rpc/request_task_security_test.go index 94217d20..2ae3d997 100644 --- a/service/rpc/request_task_security_test.go +++ b/service/rpc/request_task_security_test.go @@ -18,6 +18,7 @@ import ( type requestTaskSecurityStream struct { ctx context.Context results []*pb.TaskResult + onRecv func() onSend func(*pb.Task) sendErr error } @@ -31,6 +32,9 @@ func (s *requestTaskSecurityStream) Send(task *pb.Task) error { func (s *requestTaskSecurityStream) Recv() (*pb.TaskResult, error) { if len(s.results) == 0 { + if s.onRecv != nil { + s.onRecv() + } return nil, context.Canceled } result := s.results[0] @@ -201,6 +205,52 @@ func TestRequestTaskSkipsAlertTriggerCronResultAfterSendFailure(t *testing.T) { } } +func TestRequestTaskClearsTaskStreamOnRecvError(t *testing.T) { + reporter := requestTaskSecurityServer(7, 200, "cccccccc-cccc-cccc-cccc-cccccccccccc") + setupRequestTaskSecurityFixture(t, []*model.Server{reporter}, nil, map[uint64]model.UserInfo{ + 200: {Role: model.RoleMember}, + }, map[string]uint64{"reporter-secret": 200}) + + stream := requestTaskSecurityAuthedStream("reporter-secret", reporter.UUID) + err := NewNezhaHandler().RequestTask(stream) + if !errors.Is(err, context.Canceled) { + t.Fatalf("expected RequestTask to finish after Recv error, got %v", err) + } + + server, ok := singleton.ServerShared.Get(reporter.ID) + if !ok { + t.Fatalf("server %d not found", reporter.ID) + } + if got := server.GetTaskStream(); got != nil { + t.Fatalf("dead RequestTask stream must be cleared, got %T", got) + } +} + +func TestRequestTaskKeepsNewerTaskStreamOnOldRecvError(t *testing.T) { + reporter := requestTaskSecurityServer(7, 200, "dddddddd-dddd-dddd-dddd-dddddddddddd") + setupRequestTaskSecurityFixture(t, []*model.Server{reporter}, nil, map[uint64]model.UserInfo{ + 200: {Role: model.RoleMember}, + }, map[string]uint64{"reporter-secret": 200}) + + server, ok := singleton.ServerShared.Get(reporter.ID) + if !ok { + t.Fatalf("server %d not found", reporter.ID) + } + newer := &requestTaskSecurityStream{ctx: context.Background()} + old := requestTaskSecurityAuthedStream("reporter-secret", reporter.UUID) + old.onRecv = func() { + server.SetTaskStream(newer) + } + + err := NewNezhaHandler().RequestTask(old) + if !errors.Is(err, context.Canceled) { + t.Fatalf("expected RequestTask to finish after Recv error, got %v", err) + } + if got := server.GetTaskStream(); got != newer { + t.Fatalf("old stream cleanup must keep newer stream, got %T", got) + } +} + func setupRequestTaskSecurityFixture(t *testing.T, servers []*model.Server, crons []*model.Cron, users map[uint64]model.UserInfo, agentSecrets map[string]uint64) { t.Helper() @@ -305,22 +355,26 @@ func connectRequestTaskSecurityTaskStreamWithSendHook(t *testing.T, serverID uin if !ok { t.Fatalf("server %d not found", serverID) } - server.TaskStream = &requestTaskSecurityStream{ctx: context.Background(), sendErr: sendErr, onSend: onSend} + server.SetTaskStream(&requestTaskSecurityStream{ctx: context.Background(), sendErr: sendErr, onSend: onSend}) } func runRequestTaskSecurityResult(t *testing.T, secret string, uuid string, result *pb.TaskResult) { t.Helper() - stream := &requestTaskSecurityStream{ + stream := requestTaskSecurityAuthedStream(secret, uuid) + stream.results = []*pb.TaskResult{result} + err := NewNezhaHandler().RequestTask(stream) + if !errors.Is(err, context.Canceled) { + t.Fatalf("expected RequestTask to finish after test result, got %v", err) + } +} + +func requestTaskSecurityAuthedStream(secret string, uuid string) *requestTaskSecurityStream { + return &requestTaskSecurityStream{ ctx: metadata.NewIncomingContext(context.Background(), metadata.Pairs( "client_secret", secret, "client_uuid", uuid, )), - results: []*pb.TaskResult{result}, - } - err := NewNezhaHandler().RequestTask(stream) - if !errors.Is(err, context.Canceled) { - t.Fatalf("expected RequestTask to finish after test result, got %v", err) } } diff --git a/service/singleton/alertsentinel.go b/service/singleton/alertsentinel.go index cb3cbe91..e2c1d0de 100644 --- a/service/singleton/alertsentinel.go +++ b/service/singleton/alertsentinel.go @@ -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. diff --git a/service/singleton/config.go b/service/singleton/config.go index 4b3d3ace..65f19687 100644 --- a/service/singleton/config.go +++ b/service/singleton/config.go @@ -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) diff --git a/service/singleton/config_test.go b/service/singleton/config_test.go new file mode 100644 index 00000000..a486212d --- /dev/null +++ b/service/singleton/config_test.go @@ -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) + } +} diff --git a/service/singleton/crontask.go b/service/singleton/crontask.go index ab0bdade..85418c2c 100644 --- a/service/singleton/crontask.go +++ b/service/singleton/crontask.go @@ -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 { diff --git a/service/singleton/security_regression_test.go b/service/singleton/security_regression_test.go index 44a8f6c1..32305334 100644 --- a/service/singleton/security_regression_test.go +++ b/service/singleton/security_regression_test.go @@ -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 } diff --git a/service/singleton/server_transfer.go b/service/singleton/server_transfer.go new file mode 100644 index 00000000..bdd47f5d --- /dev/null +++ b/service/singleton/server_transfer.go @@ -0,0 +1,1440 @@ +package singleton + +import ( + "errors" + "fmt" + "log" + "sort" + "strings" + "sync" + "sync/atomic" + "time" + + "github.com/goccy/go-json" + "golang.org/x/mod/semver" + "gorm.io/gorm" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/pkg/utils" + pb "github.com/nezhahq/nezha/proto" +) + +// transferHandshakeSecretLength matches model.DefaultAgentSecretLength so +// agent-side validation that expects char(32) accepts handshake secrets +// without a special case. +const transferHandshakeSecretLength = 32 + +// MinServerTransferAgentVersion is the minimum agent build version that +// recognises TaskTypeServerTransferApply. Pre-transfer agents see the type +// fall through their `switch task.GetType()` default and never reply, so +// dashboard would wait the full 24h timeout sweep. Refuse the transfer +// up-front instead, with a clear operator-facing reason. +const MinServerTransferAgentVersion = "v1.18.0" + +// ServerTransferShared owns the lifecycle of in-flight ServerTransfer rows: +// in-memory pending index used by auth tolerance, state-machine transitions +// (verified / failed / timeout / cancelled), best-effort ApplyConfig push to +// the affected agent, and a fan-out broker for the dashboard WebSocket. +var ServerTransferShared *ServerTransferClass + +// ServerTransferStreamRevocationHook is installed by the rpc service at +// startup. It is invoked whenever a transfer transition (Register on +// Initiate, revertTransition on Cancel/Fail/Timeout, OnServersDeleted) +// rotates a server's effective ownership; the rpc package closes every +// IOStream whose targetServerID matches, so a terminal/file-manager/NAT +// session opened by the old owner cannot survive into the new owner's +// tenancy. The dashboard package leaves this nil when running without +// the rpc service (tests). +// +// Singleton can't import rpc directly without a cycle, so we expose the +// hook as a package-level function variable and let cmd/dashboard/rpc +// wire it in ServeRPC. +var ServerTransferStreamRevocationHook func(serverID uint64) + +// ServerTransferRevokeStreamsForServer is the dispatch entry the +// state-machine calls. It is safe to call when no hook is installed +// (tests, headless dashboard); revocation simply becomes a no-op. +func ServerTransferRevokeStreamsForServer(serverID uint64) { + hook := ServerTransferStreamRevocationHook + if hook == nil { + return + } + hook(serverID) +} + +// defaultServerTransferTimeout is the upper bound a Pending transfer may live +// before being auto-failed. Chosen at 24h so an agent that's offline at the +// time of transfer still has a generous window to come back online and pick +// up its new credentials. Cancellable mid-window. +const defaultServerTransferTimeout = 24 * time.Hour + +// serverTransferTimeoutTickInterval governs how often the timeout sweeper +// runs. 30s gives near-instant detection on the (rare) timeout cases without +// hammering the DB on a system that's idle most of the time. +const serverTransferTimeoutTickInterval = 30 * time.Second + +const defaultRevertDeliveryRecoveryWindow = defaultServerTransferTimeout + +// ServerTransferClass is the singleton holding pending transfers and their +// subscribers. All mutating operations go through methods so DB and in-memory +// state stay in sync. +type ServerTransferClass struct { + mu sync.RWMutex + pending map[uint64]*model.ServerTransfer + revertDeliveries map[uint64]*model.ServerTransfer + // revertRecovery holds RevertHandshakeSecrets the dashboard has pushed + // but the agent has not yet acknowledged, in the window between Cancel/ + // Fail/Timeout and either the agent's reconnect (which MarkRevertDelivered + // promotes) or expiry. It is consulted by auth via LookupByRevertHandshakeSecret + // alongside revertDeliveries, but unlike revertDeliveries it is not used + // to drive new ApplyConfig pushes — that distinction is what lets + // Register clear revertDeliveries (so a stale pushRevertIfOnline cannot + // overwrite a freshly-applied new HandshakeSecret on the agent) while + // still keeping the auth recovery channel open for the agent that may + // still hold the old RevertHandshakeSecret on disk. + // terminalSecretRecovery holds the just-terminated transfer for each + // server so the agent can authenticate during the bounded recovery + // window even after Cancel/Fail/Timeout. One slot per server covers + // BOTH per-transfer secrets simultaneously: + // + // forward (t.HandshakeSecret) — agent committed it to disk via + // the 10s reload timer before the + // dashboard observed MarkVerified. + // Auth admits it but does NOT + // promote, so RequestTask runs + // OnAgentReconnect and the + // rollback ApplyConfig swaps the + // agent onto the revert secret. + // + // revert (t.RevertHandshakeSecret) — dashboard pushed the rollback; + // the agent has 10s before its + // reload commits. Auth admits the + // revert secret and on success + // promotes it (MarkRevertDelivered) + // into verifiedHandshakes — that + // is the agent's stable credential + // from there on. + // + // One slot, two kinds, same TTL (defaultRevertDeliveryRecoveryWindow), + // same eviction triggers (Register on a NEW transfer for this server + // for the forward kind only — see below — / MarkRevertDelivered / + // MarkVerified / OnServersDeleted). Register-on-Retry intentionally + // preserves the slot so the agent's still-in-flight rollback can + // recover even while a fresh pending row is being set up. + // + // SECURITY: only revertTransition populates this map. A direct DB poke + // to a terminal status (the attacker-reuse model exercised by + // TestAuthHandshakeSecretRejectedAfterTransferTerminated) never reaches + // this code path, so a stolen per-transfer secret cannot authenticate + // even if the attacker can forge a terminal row in the DB. + terminalSecretRecovery map[uint64]*model.ServerTransfer + // verifiedHandshakes maps serverID -> the HandshakeSecret of the most + // recent Verified transfer that landed on this server. PushIfOnline + // delivers ONLY the per-transfer HandshakeSecret to the agent, never a + // long-term user-global AgentSecret, so once MarkVerified completes the + // agent's persistent on-disk credential for this server IS the handshake + // secret. Auth has to keep accepting it for that (serverID, secret) pair + // on every subsequent reconnect, or the agent silently locks itself out + // the next time the gRPC stream drops. Invalidated when a new transfer + // is initiated for the same server (Initiate / Register). + verifiedHandshakes map[uint64]string + // initiating tracks servers whose InitiateExclusive call is currently + // running the DB transaction. It exists separately from `pending` + // because the row hasn't been Registered yet — without this set, two + // concurrent callers could both pass the HasPending guard, both run + // their transactions, and both succeed in creating Pending rows. + initiating map[uint64]bool + // applyConfigSendLocks orders ApplyConfig sends per server transfer lifecycle. + // Do not use c.mu for this: stream.Send may block, but stale new-secret + // pushes and cancel/fail/timeout revert pushes must not overtake each other + // for the same server because the agent applies the last task it receives. + applyConfigSendLocks sync.Map + + subMu sync.Mutex + subs map[uint64]chan *model.ServerTransfer + nextSubID uint64 + + timeout time.Duration + stopOnce sync.Once + stopCh chan struct{} +} + +// ErrServerAlreadyTransferring is returned by InitiateExclusive when a +// concurrent caller has already claimed the server for a new transfer (or a +// Pending row already exists). Callers that surface a structured outcome +// (batch-move, retry) should detect it with errors.Is and translate to +// their domain-specific status. +var ErrServerAlreadyTransferring = errors.New("server already has an in-flight transfer") + +// ErrAgentTooOldForTransfer is returned by InitiateExclusive when the agent's +// reported build version is older than MinServerTransferAgentVersion and +// therefore does not understand TaskTypeServerTransferApply. Refusing the +// transfer up-front avoids a 24h timeout sweep on an agent that will never +// reply. If the agent has never connected (Server.Host == nil) the check is +// deferred to OnAgentReconnect / PushIfOnline. +var ErrAgentTooOldForTransfer = fmt.Errorf("agent build older than %s does not support server transfer (TaskTypeServerTransferApply)", MinServerTransferAgentVersion) + +// agentSupportsTransfer reports whether s has reported a build version >= +// MinServerTransferAgentVersion. Returns true when version is unknown (agent +// never reported) so callers can defer the decision; PushIfOnline re-checks +// at push time. +func agentSupportsTransfer(s *model.Server) bool { + if s == nil || s.Host == nil { + return true + } + v := strings.TrimSpace(s.Host.Version) + if v == "" { + return true + } + if !strings.HasPrefix(v, "v") { + v = "v" + v + } + if !semver.IsValid(v) { + return true + } + return semver.Compare(v, MinServerTransferAgentVersion) >= 0 +} + +// NewServerTransferClass loads any persisted Pending transfers from the DB +// into the in-memory index and starts the timeout sweeper. Called from +// LoadSingleton. +func NewServerTransferClass() *ServerTransferClass { + c := &ServerTransferClass{ + pending: make(map[uint64]*model.ServerTransfer), + revertDeliveries: make(map[uint64]*model.ServerTransfer), + terminalSecretRecovery: make(map[uint64]*model.ServerTransfer), + verifiedHandshakes: make(map[uint64]string), + initiating: make(map[uint64]bool), + subs: make(map[uint64]chan *model.ServerTransfer), + timeout: defaultServerTransferTimeout, + stopCh: make(chan struct{}), + } + + var pending []model.ServerTransfer + // 不要再吞掉这个错误:旧代码直接 DB.Where(...).Find(&pending) 忽略 + // res.Error,schema 损坏 / 表丢失 / DB 锁等情况下 pending 会被静默 + // 留空,所有进行中的 transfer 在 dashboard 重启后就丢失了 auth 容忍窗口, + // 对应 agent 会在重连时被拒绝。GORM 默认 logger 也会打这条 SQL,但混在 + // SQL 日志里很难被注意到;这里显式发一条 NEZHA>> 前缀让运维能立刻看到。 + if res := DB.Where("status = ?", model.ServerTransferStatusPending).Find(&pending); res.Error != nil { + log.Printf("NEZHA>> ServerTransferClass: failed to load pending transfers from DB: %v", res.Error) + } + for i := range pending { + t := pending[i] + // Ghost guard: the server may have been hard-deleted while the + // dashboard was down. Skipping orphans keeps HasPending honest + // (no false positives blocking new transfers) and prevents the + // timeout sweeper from looping forever on a row whose server + // row no longer exists. + if server, ok := ServerShared.Get(t.ServerID); !ok || server == nil { + log.Printf("NEZHA>> ServerTransferClass: dropping pending transfer %d for missing server %d", t.ID, t.ServerID) + continue + } + c.pending[t.ServerID] = &t + } + for i := range pending { + t := pending[i] + // Skip ghost rows whose server has been deleted out from under + // the transfer (e.g. before OnServersDeleted existed, or because + // the row predates this branch). Loading them would resurrect a + // HasPending state that no longer corresponds to a real server + // and the timeout sweeper would log errors every 30s without + // being able to settle the row. + if s, ok := ServerShared.Get(t.ServerID); !ok || s == nil { + log.Printf("NEZHA>> ServerTransferClass: ignoring pending transfer %d for missing server %d (likely a leftover from before OnServersDeleted was wired)", t.ID, t.ServerID) + continue + } + c.pending[t.ServerID] = &t + } + + var reverted []model.ServerTransfer + // acked_at IS NULL is non-negotiable: MarkRevertDelivered persists + // acked_at the moment the agent has provably rotated to the rollback + // credential and intentionally clears the in-memory delivery + recovery + // slots to close the auth tolerance window. Without filtering on + // acked_at here, every dashboard restart within + // defaultRevertDeliveryRecoveryWindow rehydrates the consumed rollback + // into revertDeliveries / terminalSecretRecovery and reopens the + // LookupRevertDelivery + LookupByTerminalSecretRecovery paths in + // service/rpc/auth.go — readmitting the rolled-back ToUserID's global + // AgentSecret long after the rollback has been delivered. ACKed rows + // are rebuilt below into verifiedHandshakes from the same acked_at, + // so the long-term credential the agent actually holds on disk still + // authenticates. + if res := DB. + Where("status IN ? AND updated_at >= ? AND acked_at IS NULL", []model.ServerTransferStatus{ + model.ServerTransferStatusFailed, + model.ServerTransferStatusTimeout, + model.ServerTransferStatusCancelled, + }, time.Now().Add(-defaultRevertDeliveryRecoveryWindow)). + Order("updated_at ASC"). + Find(&reverted); res.Error != nil { + log.Printf("NEZHA>> ServerTransferClass: failed to load reverted transfer deliveries from DB: %v", res.Error) + } + for i := range reverted { + t := reverted[i] + server, ok := ServerShared.Get(t.ServerID) + if !ok || server == nil || server.GetUserID() != t.FromUserID { + continue + } + c.revertDeliveries[t.ServerID] = &t + c.terminalSecretRecovery[t.ServerID] = &t + } + + // Rebuild verifiedHandshakes by merging Verified rows and acked rollback + // rows and picking, per server, the credential whose AckedAt is the + // newest. That AckedAt is the moment the agent provably rotated to that + // secret — so the newest one is the one currently on disk. The old + // two-pass "Verified first, rollback only fills empty slots" approach + // stranded the agent in the chained transfer+rollback case where the + // rollback credential is newer than the older Verified credential. + var verified []model.ServerTransfer + if res := DB. + Where("status = ? AND acked_at IS NOT NULL", model.ServerTransferStatusVerified). + Find(&verified); res.Error != nil { + log.Printf("NEZHA>> ServerTransferClass: failed to load verified transfers from DB: %v", res.Error) + } + var rollbackAcked []model.ServerTransfer + if res := DB. + Where("status IN ? AND acked_at IS NOT NULL", []model.ServerTransferStatus{ + model.ServerTransferStatusFailed, + model.ServerTransferStatusTimeout, + model.ServerTransferStatusCancelled, + }). + Find(&rollbackAcked); res.Error != nil { + log.Printf("NEZHA>> ServerTransferClass: failed to load acked rollback transfers from DB: %v", res.Error) + } + + type credCandidate struct { + serverID uint64 + secret string + ackedAt time.Time + isRevert bool + toUserID uint64 + } + candidates := make([]credCandidate, 0, len(verified)+len(rollbackAcked)) + for i := range verified { + t := verified[i] + if t.HandshakeSecret == "" || t.AckedAt == nil { + continue + } + candidates = append(candidates, credCandidate{ + serverID: t.ServerID, + secret: t.HandshakeSecret, + ackedAt: *t.AckedAt, + toUserID: t.ToUserID, + }) + } + for i := range rollbackAcked { + t := rollbackAcked[i] + if t.RevertHandshakeSecret == "" || t.AckedAt == nil { + continue + } + candidates = append(candidates, credCandidate{ + serverID: t.ServerID, + secret: t.RevertHandshakeSecret, + ackedAt: *t.AckedAt, + isRevert: true, + toUserID: t.FromUserID, + }) + } + sort.SliceStable(candidates, func(i, j int) bool { + return candidates[i].ackedAt.After(candidates[j].ackedAt) + }) + + for _, cand := range candidates { + if _, alreadySeen := c.verifiedHandshakes[cand.serverID]; alreadySeen { + continue + } + server, ok := ServerShared.Get(cand.serverID) + if !ok || server == nil { + continue + } + // Forward Verified credential is accepted when either the server + // still belongs to ToUserID (steady state) or a subsequent transfer + // is Pending whose FromUserID equals this ToUserID (chained-transfer + // rollover window — agent on disk still holds the previous + // HandshakeSecret until MarkVerified on the new transfer). + // Rollback credential is accepted only when current owner still + // equals the original FromUserID (the rollback target). + if cand.isRevert { + if server.GetUserID() != cand.toUserID { + continue + } + } else { + if server.GetUserID() != cand.toUserID { + if pending, hasPending := c.pending[cand.serverID]; !hasPending || pending.FromUserID != cand.toUserID { + continue + } + } + } + c.verifiedHandshakes[cand.serverID] = cand.secret + } + + go c.timeoutSweepLoop() + return c +} + +// Stop terminates the background timeout sweeper. Intended for tests; in +// production the singleton lives for the lifetime of the process. +func (c *ServerTransferClass) Stop() { + c.stopOnce.Do(func() { + close(c.stopCh) + }) +} + +// LookupPending returns the pending transfer for a server if one exists. +// Hot path: called from authorizeAgentForUUID on every agent RPC, so it +// uses an RWMutex and a map lookup only. +func (c *ServerTransferClass) LookupPending(serverID uint64) (*model.ServerTransfer, bool) { + c.mu.RLock() + defer c.mu.RUnlock() + t, ok := c.pending[serverID] + return t, ok +} + +// HasPending reports whether the given server has an in-flight transfer. +// Used by Initiate to enforce the "one active transfer per server" invariant. +func (c *ServerTransferClass) HasPending(serverID uint64) bool { + c.mu.RLock() + defer c.mu.RUnlock() + _, ok := c.pending[serverID] + return ok +} + +func (c *ServerTransferClass) LookupRevertDelivery(serverID uint64) (*model.ServerTransfer, bool) { + c.mu.Lock() + defer c.mu.Unlock() + t, ok := c.revertDeliveries[serverID] + if ok && t.UpdatedAt.Before(time.Now().Add(-defaultRevertDeliveryRecoveryWindow)) { + delete(c.revertDeliveries, serverID) + return nil, false + } + return t, ok +} + +func (c *ServerTransferClass) ClearRevertDelivery(serverID, transferID uint64) { + c.mu.Lock() + if t, ok := c.revertDeliveries[serverID]; ok && t.ID == transferID { + delete(c.revertDeliveries, serverID) + } + c.mu.Unlock() +} + +// MarkRevertDelivered is called from the auth path the first time the agent +// authenticates with a transfer's RevertHandshakeSecret. The agent has now +// persisted that secret as its on-disk credential (handleApplyConfigTask's +// 10s timer has fired and applyPendingReload has saved + published it), so +// it is the long-term credential for this server until another transfer +// rotates it again. Promote it into verifiedHandshakes — the auth-path +// long-term map — and persist AckedAt so dashboard restart can rebuild. +// Without this, the only acceptance path is LookupByRevertHandshakeSecret, +// which prunes after defaultRevertDeliveryRecoveryWindow and leaves the +// agent permanently locked out. +func (c *ServerTransferClass) MarkRevertDelivered(serverID, transferID uint64) error { + now := time.Now() + res := DB.Model(&model.ServerTransfer{}). + Where("id = ? AND status IN ? AND acked_at IS NULL", transferID, []model.ServerTransferStatus{ + model.ServerTransferStatusFailed, + model.ServerTransferStatusTimeout, + model.ServerTransferStatusCancelled, + }). + Update("acked_at", &now) + if res.Error != nil { + return res.Error + } + + c.mu.Lock() + defer c.mu.Unlock() + // Agent has rotated to the revert secret on disk, so the entire + // per-server terminal-recovery slot (covering both forward and revert + // kinds for THIS transfer) is now stale. Promote the revert secret + // into verifiedHandshakes first so the long-term credential is in + // place before we drop the bounded recovery entry. + if t, ok := c.terminalSecretRecovery[serverID]; ok && t.ID == transferID && t.RevertHandshakeSecret != "" { + c.verifiedHandshakes[serverID] = t.RevertHandshakeSecret + t.AckedAt = &now + delete(c.terminalSecretRecovery, serverID) + } + // revertDeliveries is the push queue (drives pushRevertIfOnline); it + // can lag terminalSecretRecovery when Register-on-Retry already + // dropped the push entry. Clear by id only — a newer transfer's push + // entry must survive. + if t, ok := c.revertDeliveries[serverID]; ok && t.ID == transferID { + delete(c.revertDeliveries, serverID) + } + return nil +} + +// LookupByHandshakeSecret returns the Pending transfer whose per-transfer +// HandshakeSecret matches secret, or (nil, false). Called from the gRPC auth +// path so an agent that received the ApplyConfig and reconnected under the +// handshake secret can be authenticated without exposing the destination +// user's global AgentSecret. O(n) over the pending map: n is bounded by the +// count of in-flight transfers, in practice tiny. +func (c *ServerTransferClass) LookupByHandshakeSecret(secret string) (*model.ServerTransfer, bool) { + if secret == "" { + return nil, false + } + c.mu.RLock() + defer c.mu.RUnlock() + for _, t := range c.pending { + if t.HandshakeSecret == secret { + return t, true + } + } + return nil, false +} + +// LookupServerByVerifiedHandshakeSecret returns the server ID whose most +// recent Verified transfer's HandshakeSecret equals secret. Called from the +// auth path on every reconnect that misses the pending-handshake and +// revert-handshake lookups, so a Verified agent — whose persisted +// credential is the per-transfer handshake secret because no final-rotation +// ApplyConfig ever swaps it out — keeps authenticating across stream drops +// and restarts. O(n) over the verifiedHandshakes map, n is bounded by the +// number of distinct servers that have ever completed a transfer in this +// process's lifetime; in practice tiny relative to total auth traffic, and +// only consulted when the global secret lookup is about to fail. +func (c *ServerTransferClass) LookupServerByVerifiedHandshakeSecret(secret string) (uint64, bool) { + if secret == "" { + return 0, false + } + c.mu.RLock() + defer c.mu.RUnlock() + for serverID, s := range c.verifiedHandshakes { + if s == secret { + return serverID, true + } + } + return 0, false +} + +// TerminalRecoveryKind distinguishes which per-transfer secret matched +// inside terminalSecretRecovery so auth can pick the right post-match +// behaviour: forward → admit but do NOT promote (rollback delivery still +// has to happen); revert → admit and trigger MarkRevertDelivered to +// promote into verifiedHandshakes. +type TerminalRecoveryKind uint8 + +const ( + TerminalRecoveryNone TerminalRecoveryKind = iota + TerminalRecoveryForward + TerminalRecoveryRevert +) + +// LookupByTerminalSecretRecovery is the single auth-facing entry into +// terminalSecretRecovery. Both per-kind wrappers delegate here so there is +// exactly one TTL-prune + secret-match site to audit. Returns the matched +// transfer and which secret matched. +func (c *ServerTransferClass) LookupByTerminalSecretRecovery(secret string) (*model.ServerTransfer, TerminalRecoveryKind, bool) { + if secret == "" { + return nil, TerminalRecoveryNone, false + } + c.mu.Lock() + defer c.mu.Unlock() + cutoff := time.Now().Add(-defaultRevertDeliveryRecoveryWindow) + for serverID, t := range c.terminalSecretRecovery { + if t.UpdatedAt.Before(cutoff) { + delete(c.terminalSecretRecovery, serverID) + continue + } + if t.HandshakeSecret == secret { + return t, TerminalRecoveryForward, true + } + if t.RevertHandshakeSecret == secret { + return t, TerminalRecoveryRevert, true + } + } + return nil, TerminalRecoveryNone, false +} + +// LookupByRevertHandshakeSecret keeps the prior per-kind signature so +// callers outside the singleton (auth.go's promote-on-success path) do +// not need to know about the unified table. Only returns matches with +// kind=revert. +func (c *ServerTransferClass) LookupByRevertHandshakeSecret(secret string) (*model.ServerTransfer, bool) { + t, kind, ok := c.LookupByTerminalSecretRecovery(secret) + if !ok || kind != TerminalRecoveryRevert { + return nil, false + } + return t, true +} + +// LookupByForwardHandshakeSecretInTerminalRecovery is the symmetric +// per-kind wrapper for the forward secret. Only returns matches with +// kind=forward. +func (c *ServerTransferClass) LookupByForwardHandshakeSecretInTerminalRecovery(secret string) (*model.ServerTransfer, bool) { + t, kind, ok := c.LookupByTerminalSecretRecovery(secret) + if !ok || kind != TerminalRecoveryForward { + return nil, false + } + return t, true +} + +func (c *ServerTransferClass) registerRevertDelivery(t *model.ServerTransfer) { + c.mu.Lock() + c.revertDeliveries[t.ServerID] = t + c.mu.Unlock() +} + +// registerTerminalSecretRecovery records the just-terminated transfer so +// auth can recognise either of its per-transfer secrets during the bounded +// recovery window. One call per revertTransition; the per-server slot is +// overwritten by a later terminal transition, mirroring the behaviour +// agents experience on disk (last credential applied wins). +func (c *ServerTransferClass) registerTerminalSecretRecovery(t *model.ServerTransfer) { + if t.HandshakeSecret == "" && t.RevertHandshakeSecret == "" { + return + } + c.mu.Lock() + c.terminalSecretRecovery[t.ServerID] = t + c.mu.Unlock() +} + +func (c *ServerTransferClass) applyConfigSendLock(serverID uint64) *sync.Mutex { + lock, _ := c.applyConfigSendLocks.LoadOrStore(serverID, &sync.Mutex{}) + return lock.(*sync.Mutex) +} + +// Initiate runs inside the given transaction and: +// - creates the ServerTransfer row with Status=Pending +// - flips Server.UserID to toUserID +// +// Caller is responsible for ensuring no concurrent transfer exists for +// serverID (HasPending check earlier in the same critical section) and for +// invoking Register + PushIfOnline after the transaction commits. +func (c *ServerTransferClass) Initiate(tx *gorm.DB, serverID, fromUserID, toUserID, initiatorID uint64) (*model.ServerTransfer, error) { + // Generate both handshake secrets up-front. PushIfOnline embeds + // HandshakeSecret in the agent ApplyConfig instead of the destination + // user's global AgentSecret; the rollback path mirrors with + // RevertHandshakeSecret. Per-transfer scope: a leak to a hijacked stream + // gives the attacker only this one server's rotation token, never the + // user's global secret. Generation must succeed — falling back to the + // global secret here would silently reintroduce the cross-user leak. + handshake, err := utils.GenerateRandomString(transferHandshakeSecretLength) + if err != nil { + return nil, fmt.Errorf("generate transfer handshake secret: %w", err) + } + revertHandshake, err := utils.GenerateRandomString(transferHandshakeSecretLength) + if err != nil { + return nil, fmt.Errorf("generate transfer revert handshake secret: %w", err) + } + t := &model.ServerTransfer{ + ServerID: serverID, + FromUserID: fromUserID, + ToUserID: toUserID, + InitiatorID: initiatorID, + Status: model.ServerTransferStatusPending, + HandshakeSecret: handshake, + RevertHandshakeSecret: revertHandshake, + } + if err := tx.Create(t).Error; err != nil { + return nil, err + } + // RowsAffected==1 is the only signal that a real server row was mutated: + // if the row was deleted between the caller's pre-check and this UPDATE, + // returning success would let Register publish a ghost pending entry and + // auth.go would then keep accepting the previous owner's secret for a + // server that doesn't exist. Surface the divergence so the surrounding + // transaction rolls back the orphan ServerTransfer. + res := tx.Model(&model.Server{}).Where("id = ?", serverID).Update("user_id", toUserID) + if res.Error != nil { + return nil, res.Error + } + if res.RowsAffected != 1 { + return nil, fmt.Errorf("server %d: ownership update affected %d rows (want 1) — row likely deleted concurrently", serverID, res.RowsAffected) + } + return t, nil +} + +// Register makes a freshly-persisted Pending transfer visible to the auth +// tolerance path. Must be called only after the Initiate transaction has +// committed, otherwise authorizeAgentForUUID could observe a transfer that +// doesn't yet exist in the DB. +// +// Ordering invariant: the in-memory Server.UserID is updated BEFORE the +// pending entry is published. Inverting these two would leave a window +// where authorizeAgentForUUID still sees the old owner via ServerShared +// and admits the old AgentSecret on the happy "owner match" path — +// bypassing the bounded pending-tolerance contract. +func (c *ServerTransferClass) Register(t *model.ServerTransfer) { + if s, ok := ServerShared.Get(t.ServerID); ok && s != nil { + // SetUserID over atomic write — auth.go hot path concurrently + // reads this field; a plain assignment would be a data race. + s.SetUserID(t.ToUserID) + } + + c.mu.Lock() + c.pending[t.ServerID] = t + // Drop only the push queue entry: pushRevertIfOnline must not re-send + // the prior rollback now that a new transfer is taking over the agent's + // credential. The auth-side recovery for the prior transfer's secrets + // stays alive in terminalSecretRecovery — the agent's 10s reload may + // not have committed the rollback yet and we still need to admit either + // the previous forward HandshakeSecret (last-completed Verified) or + // the previous RevertHandshakeSecret (uncommitted rollback) until + // MarkVerified on this fresh transfer supersedes both. + delete(c.revertDeliveries, t.ServerID) + // Do NOT delete verifiedHandshakes[t.ServerID] here. The agent's on-disk + // credential is the previous HandshakeSecret (PushIfOnline never + // delivers a user-global secret), and that secret must keep + // authenticating for the entire rollover: Register precedes PushIfOnline, + // the agent's reload timer adds another ~10s delay, and PushIfOnline is + // best-effort against stream loss. MarkVerified replaces the entry with + // the new HandshakeSecret once the agent has provably rotated; + // Cancel/Fail/Timeout leave it in place so the agent stays online while + // ownership rolls back. + c.mu.Unlock() + + // Ownership has rotated to ToUserID — tear down any IOStream the + // previous owner had open against this server so it cannot survive + // into the new tenancy. + ServerTransferRevokeStreamsForServer(t.ServerID) + + c.broadcast(t) +} + +// InitiateExclusive runs the full create-and-publish flow for a new +// ServerTransfer with mutual exclusion on serverID. The HasPending check, +// the DB transaction, and the Register call are serialized via a per-server +// claim so two concurrent callers (e.g. two operators submitting batch-move +// at the same instant) cannot both pass the guard and end up creating two +// Pending rows for the same server. Without this, the older HasPending + +// Initiate + Register sequence had a TOCTOU window — both callers would +// observe "no pending", both run their tx, both Register, with the second +// Register silently overwriting the first in the in-memory index while two +// rows remained Pending in the DB. +// +// Returns ErrServerAlreadyTransferring when a Pending row already exists or +// another caller currently holds the claim. The caller is responsible for +// PushIfOnline after a successful return. +func (c *ServerTransferClass) InitiateExclusive(serverID, fromUserID, toUserID, initiatorID uint64) (*model.ServerTransfer, error) { + if s, ok := ServerShared.Get(serverID); ok && !agentSupportsTransfer(s) { + return nil, ErrAgentTooOldForTransfer + } + c.mu.Lock() + if _, hasPending := c.pending[serverID]; hasPending { + c.mu.Unlock() + return nil, ErrServerAlreadyTransferring + } + if c.initiating[serverID] { + c.mu.Unlock() + return nil, ErrServerAlreadyTransferring + } + c.initiating[serverID] = true + c.mu.Unlock() + + defer func() { + c.mu.Lock() + delete(c.initiating, serverID) + c.mu.Unlock() + }() + + var created *model.ServerTransfer + err := DB.Transaction(func(tx *gorm.DB) error { + t, err := c.Initiate(tx, serverID, fromUserID, toUserID, initiatorID) + if err != nil { + return err + } + created = t + return nil + }) + if err != nil { + return nil, err + } + + c.Register(created) + return created, nil +} + +// PushIfOnline best-effort sends an ApplyConfig task carrying the transfer's +// per-transfer HandshakeSecret to the affected agent. The destination user's +// global AgentSecret is intentionally NOT embedded: during Pending the agent +// stream is still authenticated by the OLD owner's secret (auth tolerance), +// and a malicious previous owner who hijacks the stream would otherwise +// recover a secret that grants access to every agent that destination user +// owns. HandshakeSecret is scoped to this single transfer and UUID; even if +// it leaks, the blast radius is one server. If the agent is offline +// (TaskStream nil), the push is skipped — OnAgentReconnect will retry when +// the agent returns. Errors are not surfaced; agent failure to apply is +// detected via the explicit TaskResult or the timeout sweeper. +// +// Stale-transfer guard: callers such as OnAgentReconnect look up the pending +// transfer and then call PushIfOnline, but a concurrent Cancel/MarkFailed/ +// MarkTimeout can settle the row between those two steps. The agent treats +// later ApplyConfig tasks as supersedes (last arrival wins inside the 10s +// reload window), so a stale push that races past pushRevertIfOnline would +// commit the rejected secret and lock the agent out. Re-check pending state +// right before Send to keep the push consistent with the dashboard's +// authoritative view. +func (c *ServerTransferClass) PushIfOnline(t *model.ServerTransfer) { + s, ok := ServerShared.Get(t.ServerID) + if !ok || s == nil { + return + } + stream := s.GetTaskStream() + if stream == nil { + return + } + + if !agentSupportsTransfer(s) { + if _, err := c.MarkFailed(t.ID, ErrAgentTooOldForTransfer.Error()); err != nil { + log.Printf("NEZHA>> ServerTransfer PushIfOnline: MarkFailed for too-old agent %d failed: %v", t.ServerID, err) + } + return + } + + if current, ok := c.LookupPending(t.ServerID); !ok || current.ID != t.ID { + return + } + + if t.HandshakeSecret == "" { + // Defence against a legacy Pending row loaded from a pre-fix DB + // snapshot. Without a handshake secret we have nothing safe to send; + // the operator must cancel and re-initiate the transfer. + log.Printf("NEZHA>> ServerTransfer PushIfOnline: transfer %d has empty HandshakeSecret; refusing to fall back to user-global AgentSecret", t.ID) + return + } + + payload, err := json.Marshal(map[string]string{ + "client_secret": t.HandshakeSecret, + }) + if err != nil { + return + } + + task := &pb.Task{ + Id: t.ID, + Type: model.TaskTypeServerTransferApply, + Data: string(payload), + } + lock := c.applyConfigSendLock(t.ServerID) + lock.Lock() + defer lock.Unlock() + if current, ok := c.LookupPending(t.ServerID); !ok || current.ID != t.ID { + return + } + c.sendApplyConfigTask(s, stream, task) +} + +// OnAgentReconnect is invoked by the gRPC RequestTask handler right after +// the new TaskStream is attached. If a Pending transfer exists for this +// server, push the ApplyConfig task — the agent reconnected with the old +// secret (the only secret it knows so far), so this is the moment to deliver +// the new one. +func (c *ServerTransferClass) OnAgentReconnect(serverID uint64) { + t, ok := c.LookupPending(serverID) + if ok { + c.PushIfOnline(t) + return + } + if t, ok := c.LookupRevertDelivery(serverID); ok { + c.pushRevertIfOnline(t) + } +} + +// pushRevertIfOnline best-effort sends an ApplyConfig task carrying the +// transfer's per-transfer RevertHandshakeSecret, instructing the agent to +// either skip or overwrite the swap it was about to perform. The source +// user's global AgentSecret is intentionally NOT embedded: after a Verified +// rollover the stream is authenticated by the NEW owner, and revealing the +// previous owner's user-global secret would compromise every agent that +// user owns. Used by revertTransition (Cancel / MarkFailed / MarkTimeout) +// to keep the agent's view of the credential in sync with the dashboard's +// reverted Server.UserID. +// +// Without this counter-push, an operator who cancels within the agent's 10s +// reload window leaves a permanent split-brain: the agent commits the swap to +// the rejected new secret and immediately fails auth because the dashboard +// has already restored ownership to FromUserID. The agent's ApplyConfig +// supersede behaviour relies on this counter-push to actually be delivered +// during the 10s window — that's the entire reason supersede exists. +// +// Best-effort: agent offline is fine if it never received the original task. +// If it already switched secrets before the revert landed, the reverted +// transfer is kept as a reconnect-delivery until the old secret is restored. +func (c *ServerTransferClass) pushRevertIfOnline(t *model.ServerTransfer) { + s, ok := ServerShared.Get(t.ServerID) + if !ok || s == nil { + return + } + stream := s.GetTaskStream() + if stream == nil { + return + } + + if t.RevertHandshakeSecret == "" { + log.Printf("NEZHA>> ServerTransfer pushRevertIfOnline: transfer %d has empty RevertHandshakeSecret; refusing to fall back to user-global AgentSecret", t.ID) + return + } + + payload, err := json.Marshal(map[string]string{ + "client_secret": t.RevertHandshakeSecret, + }) + if err != nil { + return + } + + task := &pb.Task{ + Id: t.ID, + Type: model.TaskTypeServerTransferApply, + Data: string(payload), + } + lock := c.applyConfigSendLock(t.ServerID) + lock.Lock() + defer lock.Unlock() + // Re-check revertDelivery currency inside the send lock. Without this, a + // concurrent Retry can install a new pending transfer (clearing + // revertDeliveries[serverID]) and have PushIfOnline win the lock first to + // deliver the new-owner secret; pushRevertIfOnline then acquires the lock + // next and Sends the old-owner rollback, which the agent's last-arrival + // supersede commits — silently rolling back the just-applied new secret + // and leaving the fresh transfer Pending until the 24h timeout sweep. + // Mirrors the in-lock LookupPending guard PushIfOnline uses at line ~338. + if current, ok := c.LookupRevertDelivery(t.ServerID); !ok || current.ID != t.ID { + return + } + // Send-success does NOT mean the agent has rotated yet: handleApplyConfigTask + // schedules the credential swap on a 10s time.AfterFunc, so the agent only + // reconnects under RevertHandshakeSecret well after Send returns. Clearing + // the recovery record here would close LookupByRevertHandshakeSecret before + // that reconnect arrives, falling through to the global-secret table that + // doesn't know the per-transfer token — and the agent ends up permanently + // locked out. Leave the record in place; it will be cleared on one of: + // (a) auth.go observing a successful reconnect under RevertHandshakeSecret + // (the agent has provably finished applying the rollback), + // (b) a Retry/Register installing a newer transfer for this server, + // (c) the natural defaultRevertDeliveryRecoveryWindow expiry sweep. + _ = c.sendApplyConfigTask(s, stream, task) +} + +func (c *ServerTransferClass) sendApplyConfigTask(s *model.Server, stream pb.NezhaService_RequestTaskServer, task *pb.Task) error { + // Keep Send synchronous under the per-server lock. A goroutine+timeout cannot + // cancel grpc.ServerStream.Send; returning early would let a stale new-secret + // ApplyConfig complete after a cancel/fail revert and overwrite the rollback. + if err := stream.Send(task); err != nil { + log.Printf("NEZHA>> ServerTransfer ApplyConfig send failed: serverID=%d transferID=%d: %v", s.ID, task.Id, err) + s.ClearTaskStreamIfCurrent(stream) + return err + } + return nil +} + +// MarkVerified finalizes a pending transfer after the agent has successfully +// reconnected under the new owner's secret. +// +// Return tuple: +// - (t, nil) — this call transitioned the row to Verified +// - (nil, nil) — idempotent no-op (no pending entry, or a concurrent caller +// already settled the row out of Pending so RowsAffected=0) +// - (nil, err) — DB-level failure during the CAS UPDATE; caller MUST log +// it or the auth-tolerance window stays open silently for this server +// +// The old signature returned (*ServerTransfer, bool) which conflated the +// idempotent no-op and the DB-error cases, so a broken DB looked identical to +// "already verified" and operators got no signal. The auth path now logs the +// error path explicitly; do not collapse the three return shapes back into a +// bool. +// +// The status update is gated by a WHERE clause so concurrent Cancel or +// timeout sweep cannot race past it: if status is no longer Pending in the +// DB, the UPDATE affects zero rows and the in-memory state is left alone. +// MarkVerified atomically transitions a Pending transfer to Verified. +// +// Invariant: c.mu is held across the DB CAS, the in-memory pending delete, +// and the verifiedHandshakes write. Callers reading c.pending under c.mu +// (auth.go's tolerance window) therefore can never observe a state where +// the DB row is Verified but c.pending still flags the transfer as Pending — +// the auth-bypass window that would otherwise let either a stale +// HandshakeSecret or the old owner's global AgentSecret authenticate +// between the two updates. +// +// Returns verified=true exactly when this call performed the Pending → +// Verified transition for the supplied (serverID, transferID). All other +// outcomes (no pending, transfer id mismatch, lost CAS, DB error) return +// verified=false and the caller (auth) must reject the credential. +func (c *ServerTransferClass) MarkVerified(serverID, transferID uint64) (verified bool, transfer *model.ServerTransfer, err error) { + c.mu.Lock() + + t, ok := c.pending[serverID] + if !ok || t.ID != transferID { + c.mu.Unlock() + return false, nil, nil + } + + now := time.Now() + res := DB.Model(&model.ServerTransfer{}). + Where("id = ? AND status = ?", t.ID, model.ServerTransferStatusPending). + Updates(map[string]any{ + "status": model.ServerTransferStatusVerified, + "acked_at": &now, + }) + if res.Error != nil { + c.mu.Unlock() + return false, nil, res.Error + } + if res.RowsAffected == 0 { + // Concurrent caller settled the row to a terminal status. The + // in-memory pending entry is now stale; drop it so the next auth + // call cannot read it. Do NOT promote any handshake secret. + delete(c.pending, t.ServerID) + c.mu.Unlock() + return false, nil, nil + } + t.Status = model.ServerTransferStatusVerified + t.AckedAt = &now + + delete(c.pending, t.ServerID) + // Promote the handshake secret to this server's long-term credential: + // PushIfOnline delivered ONLY HandshakeSecret to the agent, so the + // agent's persisted on-disk client_secret is exactly this string and + // every future reconnect presents it. Auth's verified-handshake lookup + // uses this map to keep accepting the credential after the pending + // entry has been removed. + if t.HandshakeSecret != "" { + c.verifiedHandshakes[t.ServerID] = t.HandshakeSecret + } + delete(c.terminalSecretRecovery, t.ServerID) + c.mu.Unlock() + + c.broadcast(t) + return true, t, nil +} + +// MarkFailed transitions a pending transfer to Failed with the supplied +// reason and reverts Server.UserID back to FromUserID. Used by the RPC +// handler when an agent reports an explicit failure via TaskResult. +func (c *ServerTransferClass) MarkFailed(transferID uint64, reason string) (*model.ServerTransfer, error) { + return c.revertTransition(transferID, model.ServerTransferStatusFailed, reason) +} + +// MarkTimeout transitions a pending transfer to Timeout and reverts +// Server.UserID. Invoked by the timeout sweeper. +func (c *ServerTransferClass) MarkTimeout(transferID uint64) (*model.ServerTransfer, error) { + return c.revertTransition(transferID, model.ServerTransferStatusTimeout, "") +} + +// Cancel transitions a pending transfer to Cancelled and reverts +// Server.UserID. Permission filtering happens at the HTTP layer; this method +// trusts the caller and only enforces "still Pending" via CAS. +func (c *ServerTransferClass) Cancel(transferID uint64) (*model.ServerTransfer, error) { + return c.revertTransition(transferID, model.ServerTransferStatusCancelled, "") +} + +// Retry creates a new Pending transfer with the same From/To as an existing +// terminal transfer. Used by the dashboard to re-issue after a timeout or +// failure without forcing the operator to retype the target user. Concurrent +// safety against another in-flight transfer is delegated to +// InitiateExclusive (same TOCTOU-free contract batch-move relies on). +// +// 必须校验 s.UserID == prev.FromUserID:操作员在 dashboard 上看到的是 +// "prev.FromUserID → prev.ToUserID" 这条记录,如果在 retry 之前有别的并发 +// transfer 把 server 划到了第三个用户,旧逻辑会用「当前 owner」当 +// FromUserID,悄悄发出一条语义完全不同的 transfer("new_owner → prev.To")。 +// 强制要求当前 owner 仍是 prev.FromUserID,否则报错让操作员重新发起。 +// Retry creates a new Pending transfer with the same From/To as an existing +// terminal transfer. Used by the dashboard to re-issue after a timeout or +// failure without forcing the operator to retype the target user. Concurrent +// safety against another in-flight transfer is delegated to +// InitiateExclusive (same TOCTOU-free contract batch-move relies on). +// +// 不在这里对比 s.UserID == prev.FromUserID: +// - 非 admin 调用方在 controller 层已经被强制为「current.UserID == caller」, +// 所以走到这里时 s.UserID 必然是 caller 自己,不存在静默漂移; +// - admin 调用方是 last-resort 回收路径,UX 的契约就是「不管 server 现在归 +// 谁,把它推给 prev.ToUserID」,被 TestRetryServerTransferAllowsAdmin 钉死。 +// 再加一次 FromUserID 校验会把这条 admin 路径拒掉。 +func (c *ServerTransferClass) Retry(prev *model.ServerTransfer, initiatorID uint64) (*model.ServerTransfer, error) { + if !prev.Status.IsTerminal() { + return nil, fmt.Errorf("cannot retry a non-terminal transfer (status=%d)", prev.Status) + } + var s model.Server + if err := DB.First(&s, prev.ServerID).Error; err != nil { + return nil, err + } + // This must happen before InitiateExclusive because that call flips ownership + // in the DB; if the target user was deleted, we must fail before any mutation. + UserLock.RLock() + _, ok := UserInfoMap[prev.ToUserID] + UserLock.RUnlock() + if !ok { + return nil, fmt.Errorf("target user %d not found", prev.ToUserID) + } + if s.UserID == prev.ToUserID { + return nil, fmt.Errorf("server already belongs to the target user") + } + created, err := c.InitiateExclusive(prev.ServerID, s.UserID, prev.ToUserID, initiatorID) + if err != nil { + return nil, err + } + c.PushIfOnline(created) + return created, nil +} + +// revertTransition is the shared body of MarkFailed/MarkTimeout/Cancel: +// CAS the status, revert Server.UserID, drop from pending index, broadcast. +// Returns the transfer in its post-transition state, or nil if it was no +// longer Pending (silent no-op for idempotency). +// +// In-memory cleanup runs regardless of whether THIS call performed the +// transition. The CAS UPDATE can return RowsAffected=0 because a concurrent +// caller (MarkVerified on the auth path, another revert, the timeout sweep) +// already settled the row between our tx.First and our UPDATE; in that case +// the in-memory pending entry is stale and the auth tolerance window for +// this server has already closed in the DB sense — letting the cache lag +// would keep accepting the old owner's secret for a server that has moved +// on. Self-heal by dropping the in-memory entry whenever the DB shows the +// row as non-Pending. +func (c *ServerTransferClass) revertTransition(transferID uint64, newStatus model.ServerTransferStatus, reason string) (*model.ServerTransfer, error) { + var t model.ServerTransfer + // transitionedByThisCall distinguishes "this call performed the CAS" + // from "row was already terminal before we got here". Both early-return + // branches MUST leave it false so the post-tx SetUserID(FromUserID) + // step below is gated on a real Pending → newStatus transition. Without + // this, OnUsersDeleted + a late Cancel would re-write Server.UserID + // back to a possibly-deleted FromUserID (regression pinned by + // TestOnUserDeleteCancelsPendingTransfersAwayFromDeletedUser). + var transitionedByThisCall bool + err := DB.Transaction(func(tx *gorm.DB) error { + if err := tx.First(&t, transferID).Error; err != nil { + return err + } + if t.Status != model.ServerTransferStatusPending { + return nil + } + now := time.Now() + updates := map[string]any{ + "status": newStatus, + "updated_at": now, + } + if reason != "" { + updates["last_error"] = reason + } + res := tx.Model(&model.ServerTransfer{}). + Where("id = ? AND status = ?", t.ID, model.ServerTransferStatusPending). + Updates(updates) + if res.Error != nil { + return res.Error + } + if res.RowsAffected == 0 { + // Concurrent caller won the CAS. Re-read so the outer cleanup + // observes the authoritative status — otherwise t still holds + // the Pending snapshot we read at the top and the self-heal + // below would falsely treat the entry as still Pending. + return tx.First(&t, transferID).Error + } + // As in Initiate: require RowsAffected==1 so a vanished server row + // aborts the revert instead of silently flipping in-memory state to + // FromUserID for a row that no longer exists. + revertRes := tx.Model(&model.Server{}). + Where("id = ?", t.ServerID). + Update("user_id", t.FromUserID) + if revertRes.Error != nil { + return revertRes.Error + } + if revertRes.RowsAffected != 1 { + return fmt.Errorf("server %d: revert ownership update affected %d rows (want 1) — row likely deleted concurrently", t.ServerID, revertRes.RowsAffected) + } + t.Status = newStatus + t.LastError = reason + t.UpdatedAt = now + transitionedByThisCall = true + return nil + }) + if err != nil { + return nil, err + } + + // Ordering invariant: when this call actually performed the revert + // (newStatus reached), the in-memory Server.UserID must be reverted + // to FromUserID BEFORE any other state becomes observable, so auth + // no longer admits the destination user's global AgentSecret via + // ServerShared.GetUserID() == userId on the happy "owner match" path. + if transitionedByThisCall { + if s, ok := ServerShared.Get(t.ServerID); ok && s != nil { + s.SetUserID(t.FromUserID) + } + } + + // Self-heal: any non-Pending DB status invalidates the in-memory entry — + // but only if that entry is THIS transfer. Without the id check a stale + // terminal id (e.g. Cancel against a transfer that already failed and + // has been superseded via Retry by a new Pending row for the same server) + // would silently wipe the new entry's auth-tolerance window and re-open + // `HasPending` so a duplicate Initiate could land. cancelServerTransfer + // does not gate on `t.Status == Pending`, so the stale-id path is + // reachable from operator UI and replayed API calls; the id match is the + // only thing keeping the in-memory pending index honest here. The DB row + // we read (t) is never the live entry's row in that case, so converging + // the cache to t.Status would be the wrong direction anyway. + if t.Status != model.ServerTransferStatusPending { + c.mu.Lock() + if existing, ok := c.pending[t.ServerID]; ok && existing.ID == t.ID { + delete(c.pending, t.ServerID) + } + c.mu.Unlock() + } + + // Gate ALL post-tx side effects on transitionedByThisCall, not on + // `t.Status == newStatus`. A stale terminal-id Cancel against a row + // that is already Cancelled has t.Status == newStatus too, so the + // older `!= newStatus` gate let the fall-through re-register the OLD + // transfer's revertDelivery / terminalSecretRecovery and re-push its + // RevertHandshakeSecret. After a Retry installed a NEW Pending + // transfer and delivered its forward HandshakeSecret, the stale + // rollback supersedes the new credential inside the agent's 10s + // reload window and strands the new transfer until the 24h timeout + // sweep. Only the call that actually performed Pending -> newStatus + // is allowed to drive rollback delivery, recovery registration, stream + // revocation, broadcast, and push. + if !transitionedByThisCall { + return nil, nil + } + c.registerRevertDelivery(&t) + c.registerTerminalSecretRecovery(&t) + + // Ownership rotated back to FromUserID — close any IOStream the + // destination user opened while they briefly held the server, so the + // rolled-back FromUserID is not exposed to live sessions from the + // would-be ToUserID. + ServerTransferRevokeStreamsForServer(t.ServerID) + + c.broadcast(&t) + c.pushRevertIfOnline(&t) + return &t, nil +} + +// OnServersDeleted finalizes any in-flight transfers for servers that have +// just been deleted. Without this hook, revertTransition cannot complete +// (its UPDATE on the gone server row fails the RowsAffected==1 invariant +// and aborts), so a Pending row would stay Pending forever, HasPending +// would keep returning true for the doomed server id, and the timeout +// sweeper would log errors every 30s without making progress. +// +// We must NOT touch model.Server here — it is already gone. We CAS each +// Pending row that the listing returned and only invalidate in-memory map +// slots whose (serverID, transferID) match a row we authoritatively +// terminated, so a concurrent Retry that landed a brand-new pending +// transfer in the same slot is not collateral damage. +func (c *ServerTransferClass) OnServersDeleted(serverIDs []uint64) { + if len(serverIDs) == 0 { + return + } + + const reason = "server deleted" + terminated := make([]model.ServerTransfer, 0, len(serverIDs)) + for _, sid := range serverIDs { + var pending []model.ServerTransfer + if err := DB.Where("server_id = ? AND status = ?", sid, model.ServerTransferStatusPending).Find(&pending).Error; err != nil { + log.Printf("NEZHA>> ServerTransfer OnServersDeleted: list pending for server %d: %v", sid, err) + continue + } + now := time.Now() + for i := range pending { + t := pending[i] + res := DB.Model(&model.ServerTransfer{}). + Where("id = ? AND status = ?", t.ID, model.ServerTransferStatusPending). + Updates(map[string]any{ + "status": model.ServerTransferStatusCancelled, + "updated_at": now, + "last_error": reason, + }) + if res.Error != nil { + log.Printf("NEZHA>> ServerTransfer OnServersDeleted: cancel transfer %d: %v", t.ID, res.Error) + continue + } + if res.RowsAffected == 0 { + continue + } + t.Status = model.ServerTransferStatusCancelled + t.LastError = reason + t.UpdatedAt = now + terminated = append(terminated, t) + } + } + + c.mu.Lock() + for i := range terminated { + t := &terminated[i] + if existing, ok := c.pending[t.ServerID]; ok && existing.ID == t.ID { + delete(c.pending, t.ServerID) + } + if existing, ok := c.revertDeliveries[t.ServerID]; ok && existing.ID == t.ID { + delete(c.revertDeliveries, t.ServerID) + } + } + // terminalSecretRecovery is keyed by serverID and can outlive an + // id-matched Cancel/Fail/Timeout that ran before this point — the + // server itself is gone, so drop unconditionally to prevent a recycled + // id from inheriting a stale per-transfer credential. + for _, sid := range serverIDs { + delete(c.terminalSecretRecovery, sid) + } + c.mu.Unlock() + + for _, sid := range serverIDs { + ServerTransferRevokeStreamsForServer(sid) + } + + for i := range terminated { + c.broadcast(&terminated[i]) + } +} + +// OnUsersDeleted terminates any Pending transfer whose FromUserID or +// ToUserID is in userIDs, BEFORE the caller drops the corresponding User +// rows. revertTransition's Cancel/Fail/Timeout paths blindly write +// Server.UserID back to FromUserID; if a pending A→B transfer outlives the +// deletion of A, a later timeout sweep (or any Cancel) would silently +// resurrect the deleted user as the server's owner. The same hazard exists +// symmetrically when B is deleted while pending: MarkVerified would promote +// to a nonexistent ToUserID. Settle the row up-front instead, mirroring +// OnServersDeleted's CAS + in-memory cleanup pattern. +// +// We deliberately do NOT touch model.Server here — the live owner may be +// a third party (chained transfers) or the surviving counterparty, and the +// caller's own delete loop (singleton.OnUserDelete) is responsible for any +// servers still attributed to the deleted user. +func (c *ServerTransferClass) OnUsersDeleted(userIDs []uint64) { + if len(userIDs) == 0 { + return + } + + const reason = "user deleted" + terminated := make([]model.ServerTransfer, 0) + var pending []model.ServerTransfer + if err := DB.Where("status = ? AND (from_user_id IN ? OR to_user_id IN ?)", + model.ServerTransferStatusPending, userIDs, userIDs).Find(&pending).Error; err != nil { + log.Printf("NEZHA>> ServerTransfer OnUsersDeleted: list pending for users %v: %v", userIDs, err) + return + } + now := time.Now() + for i := range pending { + t := pending[i] + res := DB.Model(&model.ServerTransfer{}). + Where("id = ? AND status = ?", t.ID, model.ServerTransferStatusPending). + Updates(map[string]any{ + "status": model.ServerTransferStatusCancelled, + "updated_at": now, + "last_error": reason, + }) + if res.Error != nil { + log.Printf("NEZHA>> ServerTransfer OnUsersDeleted: cancel transfer %d: %v", t.ID, res.Error) + continue + } + if res.RowsAffected == 0 { + continue + } + t.Status = model.ServerTransferStatusCancelled + t.LastError = reason + t.UpdatedAt = now + terminated = append(terminated, t) + } + + c.mu.Lock() + for i := range terminated { + t := &terminated[i] + if existing, ok := c.pending[t.ServerID]; ok && existing.ID == t.ID { + delete(c.pending, t.ServerID) + } + if existing, ok := c.revertDeliveries[t.ServerID]; ok && existing.ID == t.ID { + delete(c.revertDeliveries, t.ServerID) + } + if existing, ok := c.terminalSecretRecovery[t.ServerID]; ok && existing.ID == t.ID { + delete(c.terminalSecretRecovery, t.ServerID) + } + } + c.mu.Unlock() + + for i := range terminated { + c.broadcast(&terminated[i]) + } +} + +// timeoutSweepLoop is the goroutine started in NewServerTransferClass. It +// wakes every serverTransferTimeoutTickInterval, snapshots the pending index, +// and times out anything older than c.timeout. +func (c *ServerTransferClass) timeoutSweepLoop() { + ticker := time.NewTicker(serverTransferTimeoutTickInterval) + defer ticker.Stop() + for { + select { + case <-c.stopCh: + return + case <-ticker.C: + c.sweepTimeouts() + } + } +} + +func (c *ServerTransferClass) sweepTimeouts() { + deadline := time.Now().Add(-c.timeout) + + c.mu.RLock() + candidates := make([]uint64, 0, len(c.pending)) + for _, t := range c.pending { + if t.CreatedAt.Before(deadline) { + candidates = append(candidates, t.ID) + } + } + c.mu.RUnlock() + + // Fan out per-candidate: MarkTimeout's pushRevertIfOnline does a + // synchronous grpc.ServerStream.Send under the per-server + // applyConfigSendLock. A single wedged agent would otherwise stall every + // later candidate in this tick — and because the ticker drops on a busy + // channel, every subsequent tick too — freezing timeout detection across + // all tenants. Per-server send ordering is preserved by + // applyConfigSendLocks; cross-server parallelism is safe. We Wait so the + // sweep is a synchronous unit, which keeps tests deterministic. + var wg sync.WaitGroup + wg.Add(len(candidates)) + for _, id := range candidates { + id := id + go func() { + defer wg.Done() + _, _ = c.MarkTimeout(id) + }() + } + wg.Wait() +} + +// Subscribe registers a channel that will receive every transfer transition +// event from this point forward. The caller MUST Unsubscribe when done or +// the broker will block forever if the channel is unbuffered or full. +func (c *ServerTransferClass) Subscribe() (uint64, <-chan *model.ServerTransfer) { + c.subMu.Lock() + defer c.subMu.Unlock() + + id := atomic.AddUint64(&c.nextSubID, 1) + ch := make(chan *model.ServerTransfer, 16) + c.subs[id] = ch + return id, ch +} + +func (c *ServerTransferClass) Unsubscribe(id uint64) { + c.subMu.Lock() + ch, ok := c.subs[id] + delete(c.subs, id) + c.subMu.Unlock() + if ok { + close(ch) + } +} + +// broadcast fans the given event out to all subscribers without blocking. +// A subscriber whose buffer is full silently drops the event — the WS layer +// is expected to re-sync via REST when the user revisits a stale view. +func (c *ServerTransferClass) broadcast(t *model.ServerTransfer) { + snapshot := *t + + c.subMu.Lock() + defer c.subMu.Unlock() + for _, ch := range c.subs { + select { + case ch <- &snapshot: + default: + } + } +} diff --git a/service/singleton/server_transfer_test.go b/service/singleton/server_transfer_test.go new file mode 100644 index 00000000..2be8eb05 --- /dev/null +++ b/service/singleton/server_transfer_test.go @@ -0,0 +1,2194 @@ +package singleton + +import ( + "bytes" + "errors" + "fmt" + "log" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/require" + "gorm.io/driver/sqlite" + "gorm.io/gorm" + + "github.com/nezhahq/nezha/model" + pb "github.com/nezhahq/nezha/proto" +) + +// fakeTaskStream is the smallest stub of pb.NezhaService_RequestTaskServer +// PushIfOnline needs: a Send that captures dispatched tasks. We only call Send +// from the production code under test, so the embedded interface satisfies the +// rest of the contract with nil-panicking methods we never invoke. +type fakeTaskStream struct { + pb.NezhaService_RequestTaskServer + mu sync.Mutex + sent []*pb.Task +} + +func newFakeTaskStream() *fakeTaskStream { return &fakeTaskStream{} } + +func (f *fakeTaskStream) Send(t *pb.Task) error { + f.mu.Lock() + defer f.mu.Unlock() + f.sent = append(f.sent, t) + return nil +} + +func (f *fakeTaskStream) reset() { + f.mu.Lock() + defer f.mu.Unlock() + f.sent = nil +} + +func (f *fakeTaskStream) sendCount() int { + f.mu.Lock() + defer f.mu.Unlock() + return len(f.sent) +} + +// setupTransferFixture wires up an in-memory DB, ServerShared, and a fresh +// ServerTransferClass with the timeout sweeper stopped (each test that needs +// timeout behavior overrides c.timeout and calls c.sweepTimeouts directly). +func setupTransferFixture(t *testing.T) (*ServerTransferClass, func()) { + t.Helper() + originalDB := DB + originalServerShared := ServerShared + originalServerTransfer := ServerTransferShared + originalUserInfoMap := UserInfoMap + + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + require.NoError(t, err) + // Pin the connection pool to 1: ":memory:" creates a NEW database per + // connection, so a concurrent goroutine that the pool routes to a fresh + // connection sees an empty DB ("no such table"). Tests using + // sweepTimeouts's per-server fan-out goroutines hit this. + if sqlDB, errInner := db.DB(); errInner == nil { + sqlDB.SetMaxOpenConns(1) + } + require.NoError(t, db.AutoMigrate(&model.Server{}, &model.ServerTransfer{})) + DB = db + + ServerShared = NewServerClass() + UserInfoMap = make(map[uint64]model.UserInfo) + + c := NewServerTransferClass() + ServerTransferShared = c + + cleanup := func() { + c.Stop() + DB = originalDB + ServerShared = originalServerShared + ServerTransferShared = originalServerTransfer + UserInfoMap = originalUserInfoMap + } + return c, cleanup +} + +func seedServerForTransfer(t *testing.T, id, userID uint64) { + t.Helper() + s := &model.Server{ + Common: model.Common{ID: id, UserID: userID}, + UUID: fmt.Sprintf("uuid-%s-%d", t.Name(), id), + Name: "test-srv", + } + require.NoError(t, DB.Create(s).Error) + model.InitServer(s) + ServerShared.Update(s, s.UUID) +} + +// initiateAndRegister mirrors the controller flow: open a transaction, call +// Initiate, commit, then Register. Tests use it to set up a Pending transfer. +func initiateAndRegister(t *testing.T, c *ServerTransferClass, serverID, fromUserID, toUserID, initiatorID uint64) *model.ServerTransfer { + t.Helper() + var created *model.ServerTransfer + err := DB.Transaction(func(tx *gorm.DB) error { + var err error + created, err = c.Initiate(tx, serverID, fromUserID, toUserID, initiatorID) + return err + }) + require.NoError(t, err) + c.Register(created) + return created +} + +// markPendingVerified is a test convenience that resolves the current +// pending transfer for serverID and drives the new MarkVerified(serverID, +// transferID) signature. Tests that simulate the auth-path call don't care +// about the transferID lookup detail. +func markPendingVerified(t *testing.T, c *ServerTransferClass, serverID uint64) (verified bool, transfer *model.ServerTransfer, err error) { + t.Helper() + pending, ok := c.LookupPending(serverID) + if !ok { + return c.MarkVerified(serverID, 0) + } + return c.MarkVerified(serverID, pending.ID) +} + +func TestServerTransferInitiateFlipsServerUserID(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + require.Equal(t, model.ServerTransferStatusPending, tr.Status) + + var s model.Server + require.NoError(t, DB.First(&s, 1).Error) + require.Equal(t, uint64(200), s.UserID, "Server.UserID must be flipped to ToUserID inside the transaction") + + cached, ok := ServerShared.Get(1) + require.True(t, ok) + require.Equal(t, uint64(200), cached.UserID, "in-memory ServerShared must also reflect the new owner") +} + +func TestServerTransferLookupPendingDuringWindow(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + require.False(t, c.HasPending(1)) + + initiateAndRegister(t, c, 1, 100, 200, 1) + + require.True(t, c.HasPending(1), "HasPending must report the freshly-registered transfer") + got, ok := c.LookupPending(1) + require.True(t, ok) + require.Equal(t, uint64(100), got.FromUserID) + require.Equal(t, uint64(200), got.ToUserID) +} + +func TestServerTransferMarkVerifiedClearsPending(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + initiateAndRegister(t, c, 1, 100, 200, 1) + + ok, verified, err := markPendingVerified(t, c, 1) + require.NoError(t, err) + require.True(t, ok, "first call on a fresh pending must return verified=true") + require.NotNil(t, verified) + require.Equal(t, model.ServerTransferStatusVerified, verified.Status) + require.NotNil(t, verified.AckedAt) + + require.False(t, c.HasPending(1), "Pending index must drop the row after MarkVerified") + + var s model.Server + require.NoError(t, DB.First(&s, 1).Error) + require.Equal(t, uint64(200), s.UserID, "Server.UserID stays at ToUserID after Verified") +} + +// MarkVerified must be idempotent — a second call must not flip the row back +// or panic. The auth path calls MarkVerified opportunistically on every RPC +// authenticated as the new owner. +func TestServerTransferMarkVerifiedIsIdempotent(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + initiateAndRegister(t, c, 1, 100, 200, 1) + + ok, verified, err := markPendingVerified(t, c, 1) + require.NoError(t, err) + require.True(t, ok, "first call must perform the transition") + require.NotNil(t, verified, "first call must transition the row to Verified") + + ok, verified, err = markPendingVerified(t, c, 1) + require.NoError(t, err) + require.False(t, ok, "second call must report verified=false") + require.Nil(t, verified, "second call must be a silent no-op (RowsAffected=0)") +} + +func TestServerTransferMarkFailedRevertsOwnership(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + + failed, err := c.MarkFailed(tr.ID, "disable_command_execute") + require.NoError(t, err) + require.Equal(t, model.ServerTransferStatusFailed, failed.Status) + require.Equal(t, "disable_command_execute", failed.LastError) + + require.False(t, c.HasPending(1)) + + var s model.Server + require.NoError(t, DB.First(&s, 1).Error) + require.Equal(t, uint64(100), s.UserID, "Failed must revert Server.UserID to FromUserID") +} + +func TestServerTransferCancelRevertsOwnership(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + + cancelled, err := c.Cancel(tr.ID) + require.NoError(t, err) + require.Equal(t, model.ServerTransferStatusCancelled, cancelled.Status) + + var s model.Server + require.NoError(t, DB.First(&s, 1).Error) + require.Equal(t, uint64(100), s.UserID, "Cancel must revert Server.UserID to FromUserID") +} + +func TestServerTransferTimeoutRevertsOwnership(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + + // Force the row to look ancient so the sweeper catches it. + require.NoError(t, DB.Model(&model.ServerTransfer{}). + Where("id = ?", tr.ID). + Update("created_at", time.Now().Add(-48*time.Hour)).Error) + c.mu.Lock() + if pending, ok := c.pending[1]; ok { + pending.CreatedAt = time.Now().Add(-48 * time.Hour) + } + c.mu.Unlock() + + c.sweepTimeouts() + + var refreshed model.ServerTransfer + require.NoError(t, DB.First(&refreshed, tr.ID).Error) + require.Equal(t, model.ServerTransferStatusTimeout, refreshed.Status) + + var s model.Server + require.NoError(t, DB.First(&s, 1).Error) + require.Equal(t, uint64(100), s.UserID, "Timeout must revert Server.UserID to FromUserID") +} + +// Cancel after MarkVerified must be a no-op. The CAS guard (WHERE status = +// Pending) is the only thing preventing the auth-tolerance path and the +// timeout sweeper from racing past each other in production. +func TestServerTransferCancelAfterVerifiedIsNoOp(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + + _, _, err := markPendingVerified(t, c, 1) + require.NoError(t, err) + + result, err := c.Cancel(tr.ID) + require.NoError(t, err) + require.Nil(t, result, "Cancel on a non-Pending row returns (nil, nil)") + + var refreshed model.ServerTransfer + require.NoError(t, DB.First(&refreshed, tr.ID).Error) + require.Equal(t, model.ServerTransferStatusVerified, refreshed.Status) + + var s model.Server + require.NoError(t, DB.First(&s, 1).Error) + require.Equal(t, uint64(200), s.UserID, "Server.UserID must remain at ToUserID") +} + +// Retry guards two distinct conditions and we need a test per condition, +// otherwise a regression in one guard hides behind the other. +// +// Guard 1 (this test): the previous row must be terminal — passing a Pending +// row trips IsTerminal() before HasPending() is even consulted. +func TestServerTransferRetryRefusesOnNonTerminalStatus(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + UserLock.Lock() + UserInfoMap[200] = model.UserInfo{AgentSecret: "target-secret"} + UserLock.Unlock() + prev := initiateAndRegister(t, c, 1, 100, 200, 1) + + _, err := c.Retry(prev, 1) + require.Error(t, err) + require.Contains(t, err.Error(), "non-terminal", "must fail on the IsTerminal guard, not on HasPending") +} + +// Guard 2: even with a properly terminal prev row, Retry must still refuse +// when the server has acquired a new in-flight transfer in the meantime — +// otherwise the operator could double-book the same server. The original +// test for this guard was a copy-paste of the non-terminal test and never +// actually exercised HasPending; this version drives it directly. +func TestServerTransferRetryRefusesWhenServerHasAnotherInflight(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + UserLock.Lock() + UserInfoMap[200] = model.UserInfo{AgentSecret: "target-secret"} + UserLock.Unlock() + + failed := initiateAndRegister(t, c, 1, 100, 200, 1) + _, err := c.MarkFailed(failed.ID, "boom") + require.NoError(t, err) + var refreshed model.ServerTransfer + require.NoError(t, DB.First(&refreshed, failed.ID).Error) + require.True(t, refreshed.Status.IsTerminal(), "precondition: prev row must be terminal") + + // A different operator kicks off a new transfer right after the failure — + // server now has an active Pending row again. + initiateAndRegister(t, c, 1, 100, 300, 2) + require.True(t, c.HasPending(1)) + + _, err = c.Retry(&refreshed, 2) + require.Error(t, err) + require.Contains(t, err.Error(), "in-flight", "must fail on the HasPending guard specifically") +} + +func TestServerTransferRetryRecreatesPending(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + UserLock.Lock() + UserInfoMap[200] = model.UserInfo{AgentSecret: "target-secret"} + UserLock.Unlock() + prev := initiateAndRegister(t, c, 1, 100, 200, 1) + + _, err := c.MarkFailed(prev.ID, "boom") + require.NoError(t, err) + + // Refresh `prev` so IsTerminal sees Failed and Retry proceeds. + var refreshed model.ServerTransfer + require.NoError(t, DB.First(&refreshed, prev.ID).Error) + + created, err := c.Retry(&refreshed, 1) + require.NoError(t, err) + require.Equal(t, model.ServerTransferStatusPending, created.Status) + require.NotEqual(t, prev.ID, created.ID) + require.Equal(t, uint64(100), created.FromUserID, "Retry uses the current Server.UserID as FromUserID after the revert") + require.Equal(t, uint64(200), created.ToUserID) + + require.True(t, c.HasPending(1)) +} + +func TestServerTransferRetryRejectsMissingTargetUser(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + UserLock.Lock() + UserInfoMap[200] = model.UserInfo{AgentSecret: "target-secret"} + UserLock.Unlock() + + prev := initiateAndRegister(t, c, 1, 100, 200, 1) + _, err := c.MarkFailed(prev.ID, "boom") + require.NoError(t, err) + UserLock.Lock() + delete(UserInfoMap, 200) + UserLock.Unlock() + + var refreshed model.ServerTransfer + require.NoError(t, DB.First(&refreshed, prev.ID).Error) + + created, err := c.Retry(&refreshed, 1) + require.Error(t, err) + require.Contains(t, err.Error(), "target user") + require.Nil(t, created) + require.False(t, c.HasPending(1)) + + var s model.Server + require.NoError(t, DB.First(&s, 1).Error) + require.Equal(t, uint64(100), s.UserID) +} + +// Guard 3 (anti-regression): Retry must NOT compare s.UserID against +// prev.FromUserID. The non-admin path is already forced by the controller's +// authz check (current.UserID == caller); for the admin path, the design +// contract — pinned down by TestRetryServerTransferAllowsAdmin — is "issue +// a transfer to prev.ToUserID using whatever current owner exists, regardless +// of drift". Adding a FromUserID-must-match check inside Retry would silently +// break the admin recovery path. This test exists so future cleanups don't +// reintroduce that check. +func TestServerTransferRetryDoesNotEnforceFromUserIDMatch(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + UserLock.Lock() + UserInfoMap[200] = model.UserInfo{AgentSecret: "target-secret"} + UserLock.Unlock() + + prev := initiateAndRegister(t, c, 1, 100, 200, 1) + _, err := c.MarkFailed(prev.ID, "boom") + require.NoError(t, err) + var refreshed model.ServerTransfer + require.NoError(t, DB.First(&refreshed, prev.ID).Error) + + // Ownership drifts to user 300 (e.g. via an out-of-band transfer or admin + // override). The historical "from=100" no longer matches the live owner. + require.NoError(t, DB.Model(&model.Server{}).Where("id = ?", uint64(1)).Update("user_id", uint64(300)).Error) + if s, ok := ServerShared.Get(1); ok { + s.SetUserID(300) + } + + created, err := c.Retry(&refreshed, 999) + require.NoError(t, err, "Retry must still issue against the current owner — drift is not an error here") + require.Equal(t, uint64(300), created.FromUserID, "FromUserID tracks the live owner, not prev.FromUserID") + require.Equal(t, uint64(200), created.ToUserID) +} + +// On dashboard restart, persisted Pending rows must rehydrate the in-memory +// pending index — otherwise the auth-tolerance window evaporates after every +// restart and in-flight agents start failing authentication. +func TestServerTransferLoadsPendingFromDBOnConstruction(t *testing.T) { + _, cleanup := setupTransferFixture(t) + defer cleanup() + + seedServerForTransfer(t, 1, 200) + require.NoError(t, DB.Create(&model.ServerTransfer{ + Common: model.Common{ID: 42}, + ServerID: 1, + FromUserID: 100, + ToUserID: 200, + Status: model.ServerTransferStatusPending, + }).Error) + + reborn := NewServerTransferClass() + defer reborn.Stop() + + require.True(t, reborn.HasPending(1), "Pending row must be rehydrated from DB on construction") +} + +// Subscribe must observe every transition broadcast, in order. WS clients +// rely on this to keep their cache fresh without polling. +func TestServerTransferBroadcastReachesSubscribers(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + id, ch := c.Subscribe() + defer c.Unsubscribe(id) + + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + + // Register broadcasts; expect one event for the Pending registration. + select { + case ev := <-ch: + require.Equal(t, tr.ID, ev.ID) + require.Equal(t, model.ServerTransferStatusPending, ev.Status) + case <-time.After(time.Second): + t.Fatal("expected Pending broadcast within 1s") + } + + _, _, err := markPendingVerified(t, c, 1) + require.NoError(t, err) + select { + case ev := <-ch: + require.Equal(t, model.ServerTransferStatusVerified, ev.Status) + case <-time.After(time.Second): + t.Fatal("expected Verified broadcast within 1s") + } +} + +// The "one active transfer per server" invariant must hold under concurrent +// callers (two operators batch-moving the same server, or batch-move racing +// retry). The old flow had a TOCTOU between HasPending() and Initiate() that +// allowed two Pending rows to be created for the same server; this test pins +// down the contract that InitiateExclusive serializes the check + the +// transaction + the registration atomically. +func TestServerTransferInitiateExclusiveSerializesConcurrentCallers(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + const callers = 32 + var ( + started sync.WaitGroup + release = make(chan struct{}) + successes atomic.Int64 + conflicts atomic.Int64 + otherErrs atomic.Int64 + ) + started.Add(callers) + + for i := 0; i < callers; i++ { + go func() { + started.Done() + <-release + _, err := c.InitiateExclusive(1, 100, 200, 1) + switch { + case err == nil: + successes.Add(1) + case errors.Is(err, ErrServerAlreadyTransferring): + conflicts.Add(1) + default: + otherErrs.Add(1) + } + }() + } + + started.Wait() + close(release) + + require.Eventually(t, func() bool { + return successes.Load()+conflicts.Load()+otherErrs.Load() == callers + }, time.Second, 10*time.Millisecond, "expected all callers to settle") + + require.Equal(t, int64(0), otherErrs.Load(), "no caller should error with anything other than ErrServerAlreadyTransferring") + require.Equal(t, int64(1), successes.Load(), "exactly one InitiateExclusive may win") + require.Equal(t, int64(callers-1), conflicts.Load(), "all losers must observe ErrServerAlreadyTransferring") + + var pendingCount int64 + require.NoError(t, DB.Model(&model.ServerTransfer{}). + Where("server_id = ? AND status = ?", uint64(1), model.ServerTransferStatusPending). + Count(&pendingCount).Error) + require.Equal(t, int64(1), pendingCount, "DB must contain exactly one Pending row") +} + +// A failure inside the DB transaction must release the per-server claim so a +// later caller can retry. Without this the first failed initiation would +// permanently mark the server as "in flight" in memory and every subsequent +// move would mysteriously return ErrServerAlreadyTransferring. +func TestServerTransferInitiateExclusiveReleasesClaimOnFailure(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + // No server seeded — Initiate's UPDATE will affect zero rows but the + // INSERT still succeeds in SQLite. Force a failure by closing the DB + // briefly via a sub-test that uses an invalid server id; instead, do + // the simpler thing: seed the server, run a successful initiation, + // fail the second (HasPending conflict), then release the first via + // MarkFailed and confirm a fresh initiation succeeds — this exercises + // the release path for both the conflict and the post-terminal recovery. + seedServerForTransfer(t, 1, 100) + + first, err := c.InitiateExclusive(1, 100, 200, 1) + require.NoError(t, err) + require.NotNil(t, first) + + _, err = c.InitiateExclusive(1, 100, 200, 1) + require.ErrorIs(t, err, ErrServerAlreadyTransferring) + + _, err = c.MarkFailed(first.ID, "boom") + require.NoError(t, err) + require.False(t, c.HasPending(1), "MarkFailed must release the pending claim") + + second, err := c.InitiateExclusive(1, 100, 300, 1) + require.NoError(t, err, "after MarkFailed the server must be eligible for a new transfer") + require.NotEqual(t, first.ID, second.ID) +} + +// revertTransition's CAS UPDATE returns RowsAffected=0 whenever a concurrent +// caller (the auth path's MarkVerified, another revert, the timeout sweep) +// has already transitioned the row out of Pending between our tx.First and +// our UPDATE. The old code silently returned (nil, nil) but left the +// in-memory pending entry behind, so the affected server kept enjoying the +// auth tolerance window long after the transfer was settled — a stale +// FromUserID secret would continue to authenticate against a server that +// had moved on. This test pins down "revertTransition must converge the +// in-memory cache to whatever the DB now shows, even on its no-op path." +func TestServerTransferRevertTransitionDropsStaleMemoryOnConcurrentWin(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + require.True(t, c.HasPending(1), "precondition: in-memory pending must hold the row") + + // Simulate "another caller already won the CAS" by transitioning the DB + // row directly. The in-memory pending entry is intentionally left intact + // — we are emulating the race window between two callers. + require.NoError(t, DB.Model(&model.ServerTransfer{}). + Where("id = ?", tr.ID). + Update("status", model.ServerTransferStatusVerified).Error) + + // Cancel's CAS will see RowsAffected=0 and return (nil, nil). With the + // fix in place, in-memory pending must converge to the DB state. + result, err := c.Cancel(tr.ID) + require.NoError(t, err) + require.Nil(t, result, "Cancel against a non-Pending row is a no-op result") + + require.False(t, c.HasPending(1), "in-memory pending must be cleaned when the DB row is no longer Pending") +} + +// OnAgentReconnect is invoked from the gRPC stream handler on every fresh +// agent connection. It looks up the pending transfer and hands it to +// PushIfOnline. A concurrent Cancel can settle the transfer between those +// two steps; if PushIfOnline trusts its parameter blindly and sends the +// ApplyConfig anyway, the new secret races past the cancel's counter-push +// (pushRevertIfOnline). The agent's supersede behaviour gives the last +// arrival priority — so if our stale push arrives last, the agent commits +// the cancelled credential and locks itself out. This test pins down the +// re-check contract: PushIfOnline must verify the transfer is still pending +// for its server right before sending, and become a no-op otherwise. +func TestServerTransferPushIfOnlineSkipsStaleTransferAfterCancel(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + UserLock.Lock() + UserInfoMap[200] = model.UserInfo{AgentSecret: "new-owner-secret"} + UserLock.Unlock() + + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + stream := newFakeTaskStream() + s, _ := ServerShared.Get(1) + s.SetTaskStream(stream) + + // Cancel wins the race against the reconnect-triggered push. + _, err := c.Cancel(tr.ID) + require.NoError(t, err) + require.False(t, c.HasPending(1), "precondition: pending must be cleared by Cancel") + + // Drain Cancel's revert push so the next inspection sees only the stale + // PushIfOnline (or its absence). + stream.reset() + + // Simulate the OnAgentReconnect call that captured `tr` BEFORE Cancel + // landed and is only now reaching PushIfOnline. + c.PushIfOnline(tr) + + require.Equal(t, 0, stream.sendCount(), "PushIfOnline must skip a transfer that is no longer pending — otherwise it races past the cancel's counter-push and the agent commits the rejected secret") +} + +type cancelRaceApplyConfigStream struct { + pb.NezhaService_RequestTaskServer + + firstSendBlocked chan struct{} + releaseFirstSend chan struct{} + firstSendClaimed atomic.Bool + releaseOnce sync.Once + + mu sync.Mutex + sent []*pb.Task +} + +type neverReturningTaskStream struct { + pb.NezhaService_RequestTaskServer + reachedSend chan struct{} + release chan struct{} + reachOnce sync.Once + releaseOnce sync.Once +} + +func newNeverReturningTaskStream() *neverReturningTaskStream { + return &neverReturningTaskStream{ + reachedSend: make(chan struct{}), + release: make(chan struct{}), + } +} + +func (s *neverReturningTaskStream) Send(*pb.Task) error { + s.reachOnce.Do(func() { close(s.reachedSend) }) + <-s.release + return nil +} + +func (s *neverReturningTaskStream) releaseAll() { + s.releaseOnce.Do(func() { close(s.release) }) +} + +func newCancelRaceApplyConfigStream() *cancelRaceApplyConfigStream { + return &cancelRaceApplyConfigStream{ + firstSendBlocked: make(chan struct{}), + releaseFirstSend: make(chan struct{}), + } +} + +func (s *cancelRaceApplyConfigStream) Send(task *pb.Task) error { + if s.firstSendClaimed.CompareAndSwap(false, true) { + close(s.firstSendBlocked) + <-s.releaseFirstSend + } + + s.mu.Lock() + defer s.mu.Unlock() + s.sent = append(s.sent, task) + return nil +} + +func (s *cancelRaceApplyConfigStream) releaseBlockedFirstSend() { + s.releaseOnce.Do(func() { + close(s.releaseFirstSend) + }) +} + +func (s *cancelRaceApplyConfigStream) sentTasksSnapshot() []*pb.Task { + s.mu.Lock() + defer s.mu.Unlock() + snapshot := make([]*pb.Task, len(s.sent)) + copy(snapshot, s.sent) + return snapshot +} + +func TestServerTransferCancelRevertWinsWhenPushIfOnlineSendWasAlreadyInFlight(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + UserLock.Lock() + UserInfoMap[100] = model.UserInfo{AgentSecret: "old-owner-secret"} + UserInfoMap[200] = model.UserInfo{AgentSecret: "new-owner-secret"} + UserLock.Unlock() + + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + stream := newCancelRaceApplyConfigStream() + defer stream.releaseBlockedFirstSend() + s, ok := ServerShared.Get(1) + require.True(t, ok) + s.SetTaskStream(stream) + + pushDone := make(chan struct{}) + go func() { + defer close(pushDone) + c.PushIfOnline(tr) + }() + + select { + case <-stream.firstSendBlocked: + case <-time.After(time.Second): + t.Fatal("expected PushIfOnline to reach Send before cancelling") + } + + cancelDone := make(chan error, 1) + go func() { + _, err := c.Cancel(tr.ID) + cancelDone <- err + }() + + require.Eventually(t, func() bool { + return !c.HasPending(1) + }, time.Second, 10*time.Millisecond, "Cancel must clear pending while the stale push is blocked") + + stream.releaseBlockedFirstSend() + select { + case <-pushDone: + case <-time.After(time.Second): + t.Fatal("expected blocked PushIfOnline Send to finish") + } + select { + case err := <-cancelDone: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("expected Cancel to finish after the stale push is released") + } + + var refreshed model.ServerTransfer + require.NoError(t, DB.First(&refreshed, tr.ID).Error) + + sentAfterRelease := stream.sentTasksSnapshot() + require.NotEmpty(t, sentAfterRelease, "expected at least one delivered ApplyConfig task") + finalApplyConfig := sentAfterRelease[len(sentAfterRelease)-1] + require.Equal(t, uint64(model.TaskTypeServerTransferApply), finalApplyConfig.Type) + require.Contains(t, finalApplyConfig.Data, refreshed.RevertHandshakeSecret, "Cancel revert (RevertHandshakeSecret) must remain the final delivered ApplyConfig") + require.NotContains(t, finalApplyConfig.Data, refreshed.HandshakeSecret, "stale forward HandshakeSecret push must not arrive after the cancel revert") + require.NotContains(t, finalApplyConfig.Data, "old-owner-secret", "user-global AgentSecret must never appear in transfer ApplyConfig payloads") + require.NotContains(t, finalApplyConfig.Data, "new-owner-secret", "user-global AgentSecret must never appear in transfer ApplyConfig payloads") +} + +func TestServerTransferBlockedApplyConfigSendDoesNotBlockUnrelatedRevert(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + seedServerForTransfer(t, 2, 300) + + UserLock.Lock() + UserInfoMap[100] = model.UserInfo{AgentSecret: "server-a-old-secret"} + UserInfoMap[200] = model.UserInfo{AgentSecret: "server-a-new-secret"} + UserInfoMap[300] = model.UserInfo{AgentSecret: "server-b-old-secret"} + UserInfoMap[400] = model.UserInfo{AgentSecret: "server-b-new-secret"} + UserLock.Unlock() + + transferA := initiateAndRegister(t, c, 1, 100, 200, 1) + transferB := initiateAndRegister(t, c, 2, 300, 400, 1) + + blockedStream := newCancelRaceApplyConfigStream() + defer blockedStream.releaseBlockedFirstSend() + serverA, ok := ServerShared.Get(1) + require.True(t, ok) + serverA.SetTaskStream(blockedStream) + + serverBStream := newFakeTaskStream() + serverB, ok := ServerShared.Get(2) + require.True(t, ok) + serverB.SetTaskStream(serverBStream) + + pushDone := make(chan struct{}) + go func() { + defer close(pushDone) + c.PushIfOnline(transferA) + }() + select { + case <-blockedStream.firstSendBlocked: + case <-time.After(time.Second): + t.Fatal("expected server A PushIfOnline to block inside Send") + } + + cancelDone := make(chan error, 1) + go func() { + _, err := c.Cancel(transferB.ID) + cancelDone <- err + }() + + select { + case err := <-cancelDone: + require.NoError(t, err) + case <-time.After(200 * time.Millisecond): + t.Fatal("blocked Send for server A must not block server B cancel/revert delivery") + } + + require.Equal(t, 1, serverBStream.sendCount(), "server B revert ApplyConfig must be delivered while server A is blocked") + var refreshedB model.ServerTransfer + require.NoError(t, DB.First(&refreshedB, transferB.ID).Error) + require.Contains(t, serverBStream.sent[0].Data, refreshedB.RevertHandshakeSecret, "server B revert must carry its per-transfer RevertHandshakeSecret") + require.NotContains(t, serverBStream.sent[0].Data, "server-b-old-secret", "user-global AgentSecret must never appear in transfer payloads") + require.NotContains(t, serverBStream.sent[0].Data, "server-b-new-secret", "user-global AgentSecret must never appear in transfer payloads") + + blockedStream.releaseBlockedFirstSend() + select { + case <-pushDone: + case <-time.After(time.Second): + t.Fatal("expected blocked server A send to finish after release") + } +} + +func TestServerTransferApplyConfigSendDoesNotReturnBeforeSendCompletes(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + UserLock.Lock() + UserInfoMap[200] = model.UserInfo{AgentSecret: "timeout-new-secret"} + UserLock.Unlock() + + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + stream := newCancelRaceApplyConfigStream() + defer stream.releaseBlockedFirstSend() + server, ok := ServerShared.Get(1) + require.True(t, ok) + server.SetTaskStream(stream) + + done := make(chan struct{}) + go func() { + defer close(done) + c.PushIfOnline(tr) + }() + + select { + case <-stream.firstSendBlocked: + case <-time.After(time.Second): + t.Fatal("expected PushIfOnline to enter stream.Send") + } + select { + case <-done: + t.Fatal("PushIfOnline must not return while stream.Send is still blocked; stale ApplyConfig could arrive after a revert") + case <-time.After(200 * time.Millisecond): + } + stream.releaseBlockedFirstSend() + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("expected PushIfOnline to finish after stream.Send unblocks") + } +} + +func TestServerTransferRestartRestoresRevertDeliveryWindow(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + UserLock.Lock() + UserInfoMap[100] = model.UserInfo{AgentSecret: "restart-old-secret"} + UserInfoMap[200] = model.UserInfo{AgentSecret: "restart-new-secret"} + UserLock.Unlock() + + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + _, err := c.Cancel(tr.ID) + require.NoError(t, err) + c.Stop() + + reborn := NewServerTransferClass() + defer reborn.Stop() + ServerTransferShared = reborn + + if got, ok := reborn.LookupRevertDelivery(1); !ok || got.ID != tr.ID { + t.Fatalf("restart must preserve reverted transfer delivery window, got transfer=%v ok=%v", got, ok) + } + + stream := newFakeTaskStream() + server, ok := ServerShared.Get(1) + require.True(t, ok) + server.SetTaskStream(stream) + + reborn.OnAgentReconnect(1) + + require.Equal(t, 1, stream.sendCount(), "new-secret reconnect after dashboard restart must receive the rollback ApplyConfig") + var rebornTr model.ServerTransfer + require.NoError(t, DB.First(&rebornTr, tr.ID).Error) + require.Contains(t, stream.sent[0].Data, rebornTr.RevertHandshakeSecret, "rollback must carry the per-transfer RevertHandshakeSecret") + require.NotContains(t, stream.sent[0].Data, "restart-old-secret", "user-global AgentSecret must never appear in transfer payloads") + require.NotContains(t, stream.sent[0].Data, "restart-new-secret", "user-global AgentSecret must never appear in transfer payloads") +} + +// MarkVerified is called from the auth hot path on every agent RPC. The old +// signature conflated "no pending entry" (the expected idempotent case) with +// "DB UPDATE failed" by both returning the same (nil, false) tuple, so a real +// DB error during the transition was silently dropped. That left the auth +// tolerance window open indefinitely for the affected server (the in-memory +// pending entry was never cleared because the transition appeared to +// succeed-but-no-op) and gave operators no signal that the dashboard couldn't +// finalize transfers. This test pins down: a genuine DB failure during the +// CAS UPDATE must surface as a non-nil error so callers can log it. +func TestServerTransferMarkVerifiedSurfacesDBError(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + initiateAndRegister(t, c, 1, 100, 200, 1) + + // Dropping the table makes any UPDATE against server_transfers return + // "no such table". This is the same shape of failure a corrupt schema, + // closed connection, or runaway lock would produce in production. + require.NoError(t, DB.Migrator().DropTable(&model.ServerTransfer{})) + + _, transfer, err := markPendingVerified(t, c, 1) + require.Error(t, err, "DB-level failures must propagate up to the caller") + require.Nil(t, transfer) +} + +// The idempotent no-op cases (no pending entry OR concurrent caller already +// settled the row) must return (nil, nil) — distinguishable from a real DB +// error by the absence of an error. Without this contract the auth path +// cannot tell "already verified, all good" from "DB is broken, bail". +func TestServerTransferMarkVerifiedNoOpReturnsNilError(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + // Case 1: no pending entry at all. + ok, transfer, err := markPendingVerified(t, c, 1) + require.NoError(t, err, "no pending entry must be a silent no-op, not an error") + require.False(t, ok) + require.Nil(t, transfer) + + // Case 2: pending entry exists but the DB row was concurrently transitioned + // out of Pending — RowsAffected=0 is still an idempotent no-op. + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + require.NoError(t, DB.Model(&model.ServerTransfer{}). + Where("id = ?", tr.ID). + Update("status", model.ServerTransferStatusCancelled).Error) + + ok, transfer, err = markPendingVerified(t, c, 1) + require.NoError(t, err, "concurrent CAS loser must be a silent no-op, not an error") + require.False(t, ok, "lost CAS must report verified=false so auth rejects the credential") + require.Nil(t, transfer) +} + +// NewServerTransferClass must surface DB load failures via the standard logger +// so operators see the failure. The original implementation discarded the +// error from DB.Where(...).Find(&pending); a corrupted schema or transient +// query failure on startup would silently leave the in-memory pending index +// empty, evaporating the auth-tolerance window for every in-flight transfer +// without any operator-visible signal. This test pins the contract: the error +// must be logged with a NEZHA prefix. +func TestNewServerTransferClassLogsDBLoadError(t *testing.T) { + originalDB := DB + originalServerShared := ServerShared + originalServerTransfer := ServerTransferShared + defer func() { + DB = originalDB + ServerShared = originalServerShared + ServerTransferShared = originalServerTransfer + }() + + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + require.NoError(t, err) + DB = db + ServerShared = NewServerClass() + + // Don't migrate ServerTransfer — the Find below will fail with + // "no such table", which is the same failure shape a corrupted DB + // would surface in production. + + var buf bytes.Buffer + originalOutput := log.Writer() + log.SetOutput(&buf) + defer log.SetOutput(originalOutput) + + c := NewServerTransferClass() + defer c.Stop() + + logged := buf.String() + require.True(t, + strings.Contains(logged, "NEZHA") && strings.Contains(logged, "transfer"), + "NewServerTransferClass must log DB load failures so operators notice; got %q", logged) +} + +// revertTransition's self-heal step previously deleted c.pending[t.ServerID] +// whenever the DB row was non-Pending, regardless of which transfer ID the +// in-memory entry was pointing at. After Retry creates a new Pending row for +// the same server, the in-memory pending entry holds the NEW transfer — but +// `cancelServerTransfer` still accepts the historical (terminal) transfer ID +// and routes it through revertTransition. The stale-id Cancel then wiped the +// fresh Pending entry's auth-tolerance window and re-opened the door for +// double initiation, even though the actual DB row it operated on never +// changed status. This test pins the contract: revertTransition's self-heal +// must only drop the in-memory entry it actually owns (same transfer ID). +func TestServerTransferCancelOnStaleTerminalKeepsNewPendingIntact(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + UserLock.Lock() + UserInfoMap[200] = model.UserInfo{AgentSecret: "target-secret"} + UserLock.Unlock() + + first := initiateAndRegister(t, c, 1, 100, 200, 1) + _, err := c.MarkFailed(first.ID, "boom") + require.NoError(t, err) + require.False(t, c.HasPending(1), "precondition: first transfer must be released") + + // Operator (or admin) re-issues the transfer. Retry uses the live owner + // (still user 100 because MarkFailed reverted it) and emits a fresh + // Pending row for the same server. + var refreshed model.ServerTransfer + require.NoError(t, DB.First(&refreshed, first.ID).Error) + second, err := c.Retry(&refreshed, 1) + require.NoError(t, err) + require.NotEqual(t, first.ID, second.ID, "Retry must create a new transfer id") + require.True(t, c.HasPending(1), "precondition: Retry must register the new pending row") + + // Now the buggy path: someone (UI, replayed API call, automation script) + // calls Cancel against the OLD terminal transfer's id. cancelServerTransfer + // doesn't gate on t.Status == Pending — only on permission — so the call + // reaches revertTransition. + result, err := c.Cancel(first.ID) + require.NoError(t, err) + require.Nil(t, result, "Cancel on a terminal row must be a silent no-op") + + require.True(t, c.HasPending(1), + "the fresh pending transfer must survive Cancel against the stale terminal id") + got, ok := c.LookupPending(1) + require.True(t, ok) + require.Equal(t, second.ID, got.ID, + "in-memory pending must still point at the new transfer, not be wiped by a stale id") +} + +// Cancel -> revertTransition synchronously calls pushRevertIfOnline at the +// end. That call captures the terminal transfer `tr` and races for the +// per-server applyConfigSendLock against any concurrent PushIfOnline (e.g. +// because the operator immediately Retried). If the new transfer's +// PushIfOnline acquires the lock FIRST and delivers the new owner's secret, +// then the still-queued pushRevertIfOnline for the OLD transfer must NOT +// send its rollback — doing so overwrites the new transfer's secret on the +// agent (supersede is last-arrival-wins), the new transfer never reconnects +// under its target secret, and it sits Pending until the 24h timeout sweep. +// +// Mirror of TestServerTransferPushIfOnlineSkipsStaleTransferAfterCancel +// (which pins down the same invariant on the `pending` index). Pin it on +// the `revertDeliveries` index as well: pushRevertIfOnline must re-check +// revertDelivery currency immediately before Send. +func TestServerTransferPushRevertIfOnlineSkipsStaleDeliveryAfterRetry(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + UserLock.Lock() + UserInfoMap[100] = model.UserInfo{AgentSecret: "old-owner-secret"} + UserInfoMap[200] = model.UserInfo{AgentSecret: "new-owner-secret"} + UserLock.Unlock() + + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + + stream := newFakeTaskStream() + s, ok := ServerShared.Get(1) + require.True(t, ok) + s.SetTaskStream(stream) + + // Cancel registers a revertDelivery for tr and synchronously invokes + // pushRevertIfOnline, which delivers the rollback (RevertHandshakeSecret). + _, err := c.Cancel(tr.ID) + require.NoError(t, err) + require.Equal(t, 1, stream.sendCount(), "precondition: Cancel must deliver the rollback ApplyConfig once") + var cancelled model.ServerTransfer + require.NoError(t, DB.First(&cancelled, tr.ID).Error) + require.Contains(t, stream.sent[0].Data, cancelled.RevertHandshakeSecret) + require.NotContains(t, stream.sent[0].Data, "old-owner-secret") + + // Operator immediately Retries — this clears the revertDelivery, installs + // a fresh Pending row, and pushes the new transfer's HandshakeSecret. + var refreshed model.ServerTransfer + require.NoError(t, DB.First(&refreshed, tr.ID).Error) + retried, err := c.Retry(&refreshed, 1) + require.NoError(t, err) + require.True(t, c.HasPending(1), "precondition: Retry must register the new pending row") + require.Equal(t, 2, stream.sendCount(), "precondition: Retry must deliver the new-pending ApplyConfig") + require.Contains(t, stream.sent[1].Data, retried.HandshakeSecret) + require.NotContains(t, stream.sent[1].Data, "new-owner-secret") + + // Simulate the bug window: Cancel's pushRevertIfOnline was scheduled but + // only now reaches the per-server lock — long after Retry has already + // landed and delivered the new secret. Replay it with the stale tr. + stream.reset() + c.pushRevertIfOnline(tr) + + require.Equal(t, 0, stream.sendCount(), + "pushRevertIfOnline must skip a transfer whose revertDelivery has been superseded by a Retry — "+ + "otherwise the agent's last-arrival ApplyConfig supersede commits the rejected old-owner secret "+ + "and the new transfer sits Pending until the 24h timeout") +} + +// SECURITY: the ApplyConfig payload PushIfOnline writes to the agent stream +// must NEVER contain another user's global AgentSecret. During a Pending +// transfer the agent stream is still authenticated by the OLD owner's secret +// (auth tolerance). A malicious old owner who knows their own user-global +// AgentSecret can run a fake agent process under the server's UUID, hold the +// RequestTask stream open, and intercept whatever the dashboard sends. If we +// embed the destination user's global AgentSecret in the payload, the +// attacker recovers a secret that grants access to EVERY agent that +// destination user owns. The transfer credential must therefore be scoped to +// this transfer only — a one-time, per-transfer token that gates the +// agent's reconnect under the new owner's identity and grants no further +// access if leaked. +func TestServerTransferPushDoesNotLeakDestinationUserGlobalSecret(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + UserLock.Lock() + UserInfoMap[100] = model.UserInfo{AgentSecret: "old-owner-global-secret"} + UserInfoMap[200] = model.UserInfo{AgentSecret: "DESTINATION-USER-GLOBAL-SECRET-MUST-NOT-LEAK"} + UserLock.Unlock() + + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + + stream := newFakeTaskStream() + s, ok := ServerShared.Get(1) + require.True(t, ok) + s.SetTaskStream(stream) + + c.PushIfOnline(tr) + require.GreaterOrEqual(t, stream.sendCount(), 1, "PushIfOnline must dispatch a transfer ApplyConfig") + + for i, sent := range stream.sent { + require.NotContains(t, sent.Data, "DESTINATION-USER-GLOBAL-SECRET-MUST-NOT-LEAK", + "task[%d] embeds the destination user's global AgentSecret; a malicious previous owner holding the stream can recover it", i) + } +} + +// Symmetric coverage for the revert path: pushRevertIfOnline must not embed +// the FROM-user's global AgentSecret when delivering the rollback over a +// stream that — by definition of the revert window — is now authenticated by +// the NEW owner. Otherwise the new owner (legitimate or compromised) can +// recover the previous owner's secret. +func TestServerTransferRevertPushDoesNotLeakFromUserGlobalSecret(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + UserLock.Lock() + UserInfoMap[100] = model.UserInfo{AgentSecret: "FROM-USER-GLOBAL-SECRET-MUST-NOT-LEAK"} + UserInfoMap[200] = model.UserInfo{AgentSecret: "new-owner-global-secret"} + UserLock.Unlock() + + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + + stream := newFakeTaskStream() + s, ok := ServerShared.Get(1) + require.True(t, ok) + s.SetTaskStream(stream) + + _, err := c.Cancel(tr.ID) + require.NoError(t, err) + require.GreaterOrEqual(t, stream.sendCount(), 1, "Cancel must dispatch a revert ApplyConfig") + + for i, sent := range stream.sent { + require.NotContains(t, sent.Data, "FROM-USER-GLOBAL-SECRET-MUST-NOT-LEAK", + "revert task[%d] embeds the source user's global AgentSecret; the now-authenticated destination owner can recover it", i) + } +} + +// sweepTimeouts iterates Pending candidates and synchronously revertTransitions +// each one. If MarkTimeout's pushRevertIfOnline blocks indefinitely on a stuck +// stream.Send, every Pending transfer after it in the sweep would otherwise +// wait — the dashboard's timeout detection would freeze for all other tenants. +// The invariant: the sweeper must process every expired candidate's +// state-transition + revert delivery within a bounded time window regardless +// of how long any single agent's Send takes. +func TestServerTransferSweepTimeoutsNotBlockedByStuckSend(t *testing.T) { + stuck := newNeverReturningTaskStream() + // Release the stuck stream BEFORE the singleton cleanup runs (cleanup is + // deferred first, releaseAll second, so LIFO unblocks the send first). + // Without this, the Wait inside sweepTimeouts holds the fan-out goroutine + // open and the test would hang on DB teardown. + + c, cleanup := setupTransferFixture(t) + defer cleanup() + defer stuck.releaseAll() + + seedServerForTransfer(t, 1, 100) + seedServerForTransfer(t, 2, 300) + + UserLock.Lock() + UserInfoMap[100] = model.UserInfo{AgentSecret: "a-old"} + UserInfoMap[200] = model.UserInfo{AgentSecret: "a-new"} + UserInfoMap[300] = model.UserInfo{AgentSecret: "b-old"} + UserInfoMap[400] = model.UserInfo{AgentSecret: "b-new"} + UserLock.Unlock() + + trA := initiateAndRegister(t, c, 1, 100, 200, 1) + trB := initiateAndRegister(t, c, 2, 300, 400, 1) + + serverA, ok := ServerShared.Get(1) + require.True(t, ok) + serverA.SetTaskStream(stuck) + + healthy := newFakeTaskStream() + serverB, ok := ServerShared.Get(2) + require.True(t, ok) + serverB.SetTaskStream(healthy) + + expired := time.Now().Add(-2 * c.timeout) + require.NoError(t, DB.Model(&model.ServerTransfer{}). + Where("id IN ?", []uint64{trA.ID, trB.ID}). + Update("created_at", expired).Error) + c.mu.Lock() + if entry, ok := c.pending[1]; ok { + entry.CreatedAt = expired + } + if entry, ok := c.pending[2]; ok { + entry.CreatedAt = expired + } + c.mu.Unlock() + + sweepDone := make(chan struct{}) + go func() { + c.sweepTimeouts() + close(sweepDone) + }() + + deadline := time.After(2 * time.Second) + for { + if healthy.sendCount() > 0 { + break + } + select { + case <-deadline: + stuck.releaseAll() + <-sweepDone + t.Fatal("sweepTimeouts blocked on server A's stuck Send and never reached server B's rollback delivery") + case <-time.After(20 * time.Millisecond): + } + } + + var savedB model.ServerTransfer + require.NoError(t, DB.First(&savedB, trB.ID).Error) + require.Equal(t, model.ServerTransferStatusTimeout, savedB.Status, + "server B's transfer must be marked Timeout while server A's stuck delivery is in flight") + + stuck.releaseAll() + <-sweepDone +} + +// Initiate must refuse to register a Pending transfer when the targeted server +// row no longer exists. Without a RowsAffected==1 check the UPDATE silently +// succeeds with 0 rows touched, Register flips an in-memory ghost entry, and +// auth.go's tolerance window then accepts the previous owner's secret for a +// server that was never actually mutated. The whole transfer must roll back. +func TestServerTransferInitiateAbortsWhenServerRowMissing(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + + const ghostServerID uint64 = 4242 + + var created *model.ServerTransfer + err := DB.Transaction(func(tx *gorm.DB) error { + var err error + created, err = c.Initiate(tx, ghostServerID, 100, 200, 1) + return err + }) + require.Error(t, err, "Initiate must error when servers.id is missing") + require.Nil(t, created, "no transfer row may be returned on a missing server") + + var rows []model.ServerTransfer + require.NoError(t, DB.Where("server_id = ?", ghostServerID).Find(&rows).Error) + require.Empty(t, rows, "the failed Initiate transaction must leave NO ServerTransfer row behind") + + require.False(t, c.HasPending(ghostServerID), "no in-memory pending entry may exist for the ghost server") +} + +// revertTransition must refuse to advance to a terminal state if the +// underlying server row has vanished since the transfer was created. Updating +// servers.user_id with 0 rows affected was silently succeeding and the +// in-memory ServerShared cache was still being flipped back to FromUserID, +// leaving DB and cache divergent on a row that nobody owns. +func TestServerTransferRevertTransitionAbortsWhenServerRowMissing(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + + require.NoError(t, DB.Delete(&model.Server{}, 1).Error) + + _, err := c.MarkFailed(tr.ID, "agent-rejected") + require.Error(t, err, "MarkFailed must error when servers.id has vanished") + + var saved model.ServerTransfer + require.NoError(t, DB.First(&saved, tr.ID).Error) + require.Equal(t, model.ServerTransferStatusPending, saved.Status, + "transfer must remain Pending — a partial revert with no server row would leave DB/cache divergent") +} + +// Regression: pushRevertIfOnline must NOT clear the revertDelivery as soon as +// stream.Send returns success. The agent's handleApplyConfigTask delays the +// actual credential swap by 10s (time.AfterFunc), so by the time the agent +// reconnects under RevertHandshakeSecret, LookupByRevertHandshakeSecret has +// to still find the record — otherwise auth falls through to the global-secret +// table which doesn't know the handshake token, and the agent is permanently +// locked out. The recovery record may only be cleared after the agent has +// actually proven it received and applied the rollback (i.e. after it +// authenticates with RevertHandshakeSecret), or via the natural +// defaultRevertDeliveryRecoveryWindow expiry sweep. +func TestServerTransferPushRevertIfOnlineKeepsRevertDeliveryUntilAgentRotates(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + UserLock.Lock() + UserInfoMap[200] = model.UserInfo{AgentSecret: "new-owner-secret"} + UserLock.Unlock() + + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + stream := newFakeTaskStream() + s, ok := ServerShared.Get(1) + require.True(t, ok) + s.SetTaskStream(stream) + + _, err := c.Cancel(tr.ID) + require.NoError(t, err) + + require.GreaterOrEqual(t, stream.sendCount(), 1, "Cancel must push the rollback ApplyConfig down the live stream") + + revert, ok := c.LookupRevertDelivery(1) + require.True(t, ok, "revertDelivery must persist after the rollback ApplyConfig has been sent — agent applies the new client_secret after a 10s timer and only then reconnects under RevertHandshakeSecret") + require.Equal(t, tr.ID, revert.ID) + + found, ok := c.LookupByRevertHandshakeSecret(revert.RevertHandshakeSecret) + require.True(t, ok, "LookupByRevertHandshakeSecret must succeed after the send — clearing on send strands the agent on a credential the dashboard no longer accepts") + require.Equal(t, tr.ID, found.ID) +} + +// BUG-1 regression: Register() must NOT drop the previous Verified +// HandshakeSecret. A server that has already completed a transfer (A→B) +// holds the per-transfer HandshakeSecret H1 on disk, NOT a user-global +// AgentSecret. When B→C is initiated, the only auth path for H1 is +// LookupServerByVerifiedHandshakeSecret. If Register() deletes the entry +// before the agent has actually rotated to H2 (which only happens ~10s +// after PushIfOnline returns, due to the agent's reload timer, and +// requires Send success in the first place), the agent cannot reconnect +// during the rollover window and may be permanently locked out if the +// process restarts before applying H2. +func TestRegisterMustNotDropPreviousVerifiedHandshakeForChainedTransfer(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + // Round 1: A=100 → B=200, then MarkVerified so the agent's persisted + // credential becomes H1. + t1 := initiateAndRegister(t, c, 1, 100, 200, 1) + require.NotEmpty(t, t1.HandshakeSecret) + h1 := t1.HandshakeSecret + _, _, err := markPendingVerified(t, c, 1) + require.NoError(t, err) + + sid, ok := c.LookupServerByVerifiedHandshakeSecret(h1) + require.True(t, ok, "after Round 1 MarkVerified, H1 must authenticate") + require.Equal(t, uint64(1), sid) + + // Round 2: B=200 → C=300. The agent has NOT yet received the new + // HandshakeSecret H2 (Register runs before PushIfOnline, and even after + // Send the agent has a 10s reload delay). H1 is still the credential + // on disk and must keep authenticating. + initiateAndRegister(t, c, 1, 200, 300, 1) + + sid, ok = c.LookupServerByVerifiedHandshakeSecret(h1) + require.True(t, ok, + "Register() must NOT delete the previous Verified HandshakeSecret — the agent still holds H1 on disk and has no other credential path during the new transfer's reload window") + require.Equal(t, uint64(1), sid) +} + +// BUG-1 regression (restart path): if dashboard restarts while a chained +// transfer is Pending, the previous Verified row must still be rebuilt into +// verifiedHandshakes. The default rebuild gate (server.UserID == ToUserID) +// would reject H1 because Server.UserID has already been flipped to C by +// the pending B→C transfer. Rebuild must additionally accept the case where +// the server has a Pending transfer whose FromUserID equals the previous +// Verified row's ToUserID. +func TestNewServerTransferClassRebuildsPreviousVerifiedDuringChainedPending(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + // Round 1: A=100 → B=200, MarkVerified. + t1 := initiateAndRegister(t, c, 1, 100, 200, 1) + h1 := t1.HandshakeSecret + _, _, err := markPendingVerified(t, c, 1) + require.NoError(t, err) + + // Round 2: B=200 → C=300, leave Pending (no MarkVerified). + initiateAndRegister(t, c, 1, 200, 300, 1) + + // Simulate dashboard restart by constructing a fresh class against the + // same DB & ServerShared. + c.Stop() + c2 := NewServerTransferClass() + ServerTransferShared = c2 + defer c2.Stop() + + sid, ok := c2.LookupServerByVerifiedHandshakeSecret(h1) + require.True(t, ok, + "after restart with Pending B→C, the previous Verified A→B HandshakeSecret must still be rebuilt — agent on disk has H1 and reconnects must succeed until the new transfer completes") + require.Equal(t, uint64(1), sid) +} + +// BUG-2 regression: MarkRevertDelivered must promote RevertHandshakeSecret +// into verifiedHandshakes so the agent — which has now persisted that secret +// as its long-term on-disk credential after the 10s reload — keeps +// authenticating after defaultRevertDeliveryRecoveryWindow expires. Without +// this promotion, the auth path only finds the secret via the temporary +// revertDeliveries window; once that 24h window sweeps the entry, the agent +// has no auth path left and is permanently locked out. +func TestMarkRevertDeliveredPromotesRevertHandshakeSecret(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + require.NotEmpty(t, tr.RevertHandshakeSecret) + revertSecret := tr.RevertHandshakeSecret + + _, err := c.Cancel(tr.ID) + require.NoError(t, err) + require.True(t, hasRevertDeliveryFor(c, 1, tr.ID), + "Cancel must register the rollback for delivery") + + require.NoError(t, c.MarkRevertDelivered(1, tr.ID)) + + sid, ok := c.LookupServerByVerifiedHandshakeSecret(revertSecret) + require.True(t, ok, + "after the agent has authenticated with RevertHandshakeSecret, that secret must be promoted to verifiedHandshakes so it survives the 24h recovery sweep") + require.Equal(t, uint64(1), sid) + + require.False(t, hasRevertDeliveryFor(c, 1, tr.ID), + "promotion must consume the delivery record — keeping both would let a stale revert overwrite a later transfer") + + var saved model.ServerTransfer + require.NoError(t, DB.First(&saved, tr.ID).Error) + require.NotNil(t, saved.AckedAt, + "AckedAt must be persisted so dashboard restart can rebuild the promoted credential") +} + +// BUG-2 regression (restart path): after MarkRevertDelivered persists +// AckedAt on a terminal row, dashboard restart must rebuild +// verifiedHandshakes[serverID] = RevertHandshakeSecret. +func TestNewServerTransferClassRebuildsAckedRollbackCredential(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + revertSecret := tr.RevertHandshakeSecret + + _, err := c.Cancel(tr.ID) + require.NoError(t, err) + require.NoError(t, c.MarkRevertDelivered(1, tr.ID)) + + c.Stop() + c2 := NewServerTransferClass() + ServerTransferShared = c2 + defer c2.Stop() + + sid, ok := c2.LookupServerByVerifiedHandshakeSecret(revertSecret) + require.True(t, ok, + "restart must rebuild acked rollback credentials from terminal rows with acked_at set") + require.Equal(t, uint64(1), sid) +} + +// BUG: NewServerTransferClass loads Verified rows first and then skips any +// rollback-acked row whose serverID already appears in verifiedHandshakes. +// In a chained "transfer-then-rollback" history the most recent credential +// the agent actually rotated to on disk is the RevertHandshakeSecret of the +// later, rolled-back transfer — not the HandshakeSecret of the earlier +// Verified transfer. The two-pass alreadySeen check therefore rebuilds the +// wrong credential and the agent is locked out on the first post-restart +// reconnect. +// +// Scenario reproduced here: +// 1. Server S owned by A=100. Transfer t1 A→B, MarkVerified. Server.UserID=B, +// agent on disk = H1, verifiedHandshakes[S]=H1, t1.AckedAt set. +// 2. Transfer t2 B→A initiated and Cancelled. Server.UserID reverts to B +// (FromUserID), MarkRevertDelivered → agent on disk = R2, +// verifiedHandshakes[S]=R2, t2.AckedAt set, t2.UpdatedAt > t1.AckedAt. +// 3. Dashboard restart. +// +// After restart the agent presents R2 — that is what is actually persisted +// on disk after step 2's reload. Auth must accept it. Today the loader +// rebuilds H1 instead and the agent is locked out forever. +func TestNewServerTransferClassPrefersNewerRollbackCredentialOverOlderVerified(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + t1 := initiateAndRegister(t, c, 1, 100, 200, 1) + h1 := t1.HandshakeSecret + _, _, err := markPendingVerified(t, c, 1) + require.NoError(t, err) + + t2 := initiateAndRegister(t, c, 1, 200, 100, 1) + r2 := t2.RevertHandshakeSecret + require.NotEmpty(t, r2) + _, err = c.Cancel(t2.ID) + require.NoError(t, err) + require.NoError(t, c.MarkRevertDelivered(1, t2.ID)) + + require.NotEqual(t, h1, r2, + "sanity: round 2's revert secret must differ from round 1's handshake secret") + + c.Stop() + c2 := NewServerTransferClass() + ServerTransferShared = c2 + defer c2.Stop() + + sid, ok := c2.LookupServerByVerifiedHandshakeSecret(r2) + require.True(t, ok, + "restart must rebuild the newer rollback credential R2 — that is what the agent has on disk after step 2. Picking the older Verified H1 locks the agent out.") + require.Equal(t, uint64(1), sid) + + _, h1StillAccepted := c2.LookupServerByVerifiedHandshakeSecret(h1) + require.False(t, h1StillAccepted, + "the stale H1 from an earlier Verified row must NOT be accepted after a newer rollback has been acked — the agent no longer holds it") +} + +func hasRevertDeliveryFor(c *ServerTransferClass, serverID, transferID uint64) bool { + t, ok := c.LookupRevertDelivery(serverID) + return ok && t.ID == transferID +} + +// forceForwardRecoveryAge back-dates a recovery entry's UpdatedAt so TTL +// tests don't have to sleep through defaultRevertDeliveryRecoveryWindow. +func forceForwardRecoveryAge(c *ServerTransferClass, serverID uint64, age time.Duration) { + c.mu.Lock() + defer c.mu.Unlock() + if t, ok := c.terminalSecretRecovery[serverID]; ok { + t.UpdatedAt = time.Now().Add(-age) + } +} + +// BUG-3 regression + HIGH-7 hardening: a Retry that runs before the agent +// has actually authenticated with the in-flight RevertHandshakeSecret must +// NOT strand auth recovery for that secret. The agent's 10s reload timer +// means rollback application is lazy. Register clears revertDeliveries +// (so a late pushRevertIfOnline does not re-deliver the now-stale rollback +// secret and overwrite the freshly applied new HandshakeSecret on the +// agent), but it must move the secret into the bounded revertRecovery +// slot so authentication keeps working until either the agent +// reconnects (MarkRevertDelivered promotes), MarkVerified on the new +// transfer supersedes, or the recovery window expires. +// +// Crucially, the secret must NOT be promoted to the permanent +// verifiedHandshakes map: that would keep an unacknowledged credential +// alive indefinitely. +func TestRegisterPreservesInflightRollbackSecretAcrossRetry(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + t1 := initiateAndRegister(t, c, 1, 100, 200, 1) + revertSecret := t1.RevertHandshakeSecret + require.NotEmpty(t, revertSecret) + + _, err := c.Cancel(t1.ID) + require.NoError(t, err) + require.True(t, hasRevertDeliveryFor(c, 1, t1.ID), + "Cancel must register a rollback delivery carrying RevertHandshakeSecret") + + initiateAndRegister(t, c, 1, 100, 200, 1) + + if _, stillVerified := c.LookupServerByVerifiedHandshakeSecret(revertSecret); stillVerified { + t.Fatal("Register on Retry must NOT promote the unacknowledged RevertHandshakeSecret into the permanent verifiedHandshakes map — that bypasses the bounded recovery window") + } + + rec, ok := c.LookupByRevertHandshakeSecret(revertSecret) + require.True(t, ok, + "Register on Retry must keep the in-flight RevertHandshakeSecret reachable via the bounded recovery lookup (revertRecovery)") + require.Equal(t, uint64(1), rec.ServerID) + require.Equal(t, t1.ID, rec.ID) +} + +// BUG: batchDeleteServer removes Server rows but never notifies +// ServerTransferShared. Any Pending transfer for that server is left in the +// DB as Pending forever, the in-memory `pending` map still holds it (so +// HasPending/InitiateExclusive still see it), and the timeout sweeper later +// tries to revert Server.UserID on a row that no longer exists. The UI shows +// the row as Pending indefinitely. +// +// OnServersDeleted must transition such Pending rows to a terminal status +// without touching the (now-gone) Server row, clear the in-memory indexes, +// and broadcast so subscribers can update. +func TestOnServersDeletedTerminatesPendingTransfersAndClearsIndexes(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + require.True(t, c.HasPending(1)) + + subID, ch := c.Subscribe() + defer c.Unsubscribe(subID) + + require.NoError(t, DB.Unscoped().Delete(&model.Server{}, tr.ServerID).Error) + c.OnServersDeleted([]uint64{tr.ServerID}) + + require.False(t, c.HasPending(1), + "OnServersDeleted must drop the in-memory pending entry so a future server with the same id cannot inherit a stale transfer") + + var saved model.ServerTransfer + require.NoError(t, DB.First(&saved, tr.ID).Error) + require.True(t, saved.Status.IsTerminal(), + "OnServersDeleted must transition the DB row to a terminal status; got status=%d", saved.Status) + + select { + case got, ok := <-ch: + require.True(t, ok) + require.Equal(t, tr.ID, got.ID) + require.True(t, got.Status.IsTerminal()) + case <-time.After(time.Second): + t.Fatal("OnServersDeleted must broadcast the terminal transition to subscribers") + } +} + +// Defence in depth: after OnServersDeleted runs, a subsequent timeout sweep +// must not blow up trying to revert Server.UserID on the deleted server, and +// must not log spurious errors. Today revertTransition would touch the +// (gone) Server row via res.RowsAffected==0 and self-heal — but it also +// performs an UPDATE on the server table that touches no rows, which +// MarkTimeout will treat as "concurrent caller won the CAS" and return nil. +// Just make sure the sweep is a no-op after deletion. +func TestSweepTimeoutsAfterServerDeletedIsNoOp(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + c.timeout = time.Nanosecond + seedServerForTransfer(t, 1, 100) + + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + _ = tr + + require.NoError(t, DB.Unscoped().Delete(&model.Server{}, uint64(1)).Error) + c.OnServersDeleted([]uint64{1}) + + require.NotPanics(t, func() { c.sweepTimeouts() }, + "sweepTimeouts after OnServersDeleted must be a no-op even though the Server row is gone") +} + +// HIGH security regression: Register must publish the new in-memory +// Server.UserID (which auth.go reads to enforce ownership) BEFORE other +// state becomes observable. Otherwise the gap between Initiate (DB +// already says ToUserID) and Register's SetUserID is a window where +// authorizeAgentForUUID sees the OLD owner == userId via the in-memory +// cache and admits the old owner via the happy "owner match" path +// rather than the bounded pending-tolerance path. +// +// The contract we lock down: after Register returns, ServerShared +// reports the new owner. There is no atomic test for "during" Register, +// so we assert the post-condition. +func TestRegisterPublishesNewOwnerBeforeReturning(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + + s, ok := ServerShared.Get(tr.ServerID) + require.True(t, ok) + require.Equal(t, uint64(200), s.GetUserID(), + "after Register returns, ServerShared must report the new owner so auth observes a consistent (DB, in-memory) snapshot") +} + +// HIGH security regression: revertTransition must restore the in-memory +// Server.UserID (FromUserID) immediately after the DB transaction commits. +// During the window where DB says reverted but in-memory still says +// ToUserID, an auth call from the destination user's global AgentSecret +// would be admitted via the happy path and obtain a long-lived stream +// for a server that no longer belongs to them. +func TestCancelPublishesRevertedOwnerBeforeReturning(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + + s, ok := ServerShared.Get(tr.ServerID) + require.True(t, ok) + require.Equal(t, uint64(200), s.GetUserID(), "precondition: pending flipped ownership") + + cancelled, err := c.Cancel(tr.ID) + require.NoError(t, err) + require.NotNil(t, cancelled) + + require.Equal(t, uint64(100), s.GetUserID(), + "after Cancel returns, ServerShared must report the rolled-back FromUserID so auth no longer admits the destination user's global AgentSecret") +} + +// HIGH security regression: OnServersDeleted must guard map deletions by +// transferID. Between the moment the deletion enumeration captured the +// pending rows for server S and the moment the in-memory map deletion +// runs, a concurrent path can install a brand-new pending transfer for +// the same serverID (either via ID reuse after delete, or in the more +// common case via a Retry whose Register lands in the slot). A naive +// `delete(c.pending, serverID)` would wipe that fresh entry. +// +// Asserted contract: if the in-memory pending entry for S no longer +// matches any transferID OnServersDeleted authoritatively terminated, +// the entry survives. +func TestOnServersDeletedGuardsByTransferIDAgainstUnrelatedNewTransfer(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + staleTransferID := uint64(99999) + require.NoError(t, DB.Create(&model.ServerTransfer{ + Common: model.Common{ID: staleTransferID}, + ServerID: 1, + FromUserID: 100, + ToUserID: 200, + Status: model.ServerTransferStatusCancelled, + LastError: "test-seeded terminal row", + }).Error) + + fresh := initiateAndRegister(t, c, 1, 100, 200, 1) + require.NotEqual(t, staleTransferID, fresh.ID, "fresh transfer must be a distinct row") + + c.OnServersDeleted([]uint64{1}) + + current, stillPending := c.LookupPending(1) + require.False(t, stillPending && current.ID != fresh.ID, + "OnServersDeleted must not silently replace the pending entry") + if stillPending { + require.Equal(t, fresh.ID, current.ID, + "pending entry left intact must still point at fresh transfer") + } +} + +// FORWARD-RECOVERY (HIGH): regression for the recovery gap where an agent +// has already committed the per-transfer forward HandshakeSecret to disk +// but the transfer is Cancel/Fail/Timeout-ed before the dashboard observed +// the MarkVerified-via-handshake reconnect. PushIfOnline only ever +// delivers t.HandshakeSecret (never a user-global secret), so the only +// credential the agent now holds for this server is the forward +// HandshakeSecret of a transfer the dashboard has already settled. +// +// The fix introduces a bounded terminalForwardRecovery slot, populated +// from revertTransition, that lets auth admit the forward HandshakeSecret +// long enough for RequestTask → OnAgentReconnect → pushRevertIfOnline to +// push the RevertHandshakeSecret rollback. Without this, the agent is +// permanently locked out — TestAuthHandshakeSecretRejectedAfterTransferTerminated +// keeps the attacker-reuse path closed (it bypasses revertTransition by +// poking the DB directly, so terminalForwardRecovery never sees it). + +// Cancel must register the just-terminated transfer's forward +// HandshakeSecret in the bounded terminalForwardRecovery slot AND keep +// LookupByRevertHandshakeSecret working for the RevertHandshakeSecret as +// today. The two recovery channels are separate maps because they have +// distinct lifecycles: revert recovery is consumed by MarkRevertDelivered +// (which promotes); forward recovery is consumed by the rollback delivery +// itself completing (handled via MarkRevertDelivered on the next loop). +func TestCancelRegistersForwardHandshakeSecretRecovery(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + forward := tr.HandshakeSecret + revert := tr.RevertHandshakeSecret + require.NotEmpty(t, forward) + require.NotEmpty(t, revert) + + _, err := c.Cancel(tr.ID) + require.NoError(t, err) + + got, ok := c.LookupByForwardHandshakeSecretInTerminalRecovery(forward) + require.True(t, ok, + "Cancel must register forward HandshakeSecret into terminalForwardRecovery so an agent that already applied it on disk can still authenticate long enough to receive the rollback") + require.Equal(t, tr.ID, got.ID) + require.Equal(t, uint64(1), got.ServerID) + + // revert recovery channel still works as before — fix must not regress it. + _, ok = c.LookupByRevertHandshakeSecret(revert) + require.True(t, ok, "RevertHandshakeSecret recovery path must remain available alongside the new forward path") +} + +// MarkFailed (agent reports failure via TaskResult) and MarkTimeout (sweeper) +// take the same revertTransition path as Cancel, so the forward recovery +// must also be registered on those. +func TestMarkFailedRegistersForwardHandshakeSecretRecovery(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + forward := tr.HandshakeSecret + + _, err := c.MarkFailed(tr.ID, "agent-rejected") + require.NoError(t, err) + + got, ok := c.LookupByForwardHandshakeSecretInTerminalRecovery(forward) + require.True(t, ok, "MarkFailed must register forward HandshakeSecret recovery") + require.Equal(t, tr.ID, got.ID) +} + +func TestMarkTimeoutRegistersForwardHandshakeSecretRecovery(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + forward := tr.HandshakeSecret + + _, err := c.MarkTimeout(tr.ID) + require.NoError(t, err) + + got, ok := c.LookupByForwardHandshakeSecretInTerminalRecovery(forward) + require.True(t, ok, "MarkTimeout must register forward HandshakeSecret recovery") + require.Equal(t, tr.ID, got.ID) +} + +// MarkVerified happens when the agent reconnects under the forward +// HandshakeSecret and the transfer is still Pending. The terminal recovery +// slot must not survive into the Verified lifecycle: once Verified, the +// forward secret is promoted into verifiedHandshakes (the long-term map) and +// keeping a stale terminal-recovery copy around could collide later if the +// same server transfers again. +func TestMarkVerifiedClearsForwardHandshakeRecoveryIfAny(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + forward := tr.HandshakeSecret + + // Simulate a prior failed cycle on the same server to seed a recovery + // entry; then a fresh transfer is verified. The fresh transfer's + // forward secret is unrelated to the prior terminal entry, but the + // per-server slot must be cleared so verifiedHandshakes is the single + // source of truth post-Verified. + _, err := c.Cancel(tr.ID) + require.NoError(t, err) + _, ok := c.LookupByForwardHandshakeSecretInTerminalRecovery(forward) + require.True(t, ok, "precondition: cancel populated forward recovery") + + tr2 := initiateAndRegister(t, c, 1, 100, 200, 1) + require.NotEqual(t, tr.ID, tr2.ID) + verified, _, err := c.MarkVerified(1, tr2.ID) + require.NoError(t, err) + require.True(t, verified) + + _, stillRecovered := c.LookupByForwardHandshakeSecretInTerminalRecovery(forward) + require.False(t, stillRecovered, + "MarkVerified on a newer transfer for this server must purge any stale forward-secret terminal recovery entry — verifiedHandshakes is now the canonical credential") +} + +// MarkRevertDelivered means the agent has authenticated with the +// RevertHandshakeSecret, which proves the rollback ApplyConfig was applied +// and the on-disk credential is now the revert secret — not the forward +// secret. The terminal-recovery entry for the forward secret is therefore +// stale and must be cleared so a leaked forward token cannot re-enter via +// recovery later in the window. +func TestMarkRevertDeliveredClearsForwardHandshakeRecovery(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + forward := tr.HandshakeSecret + + _, err := c.Cancel(tr.ID) + require.NoError(t, err) + _, ok := c.LookupByForwardHandshakeSecretInTerminalRecovery(forward) + require.True(t, ok) + + require.NoError(t, c.MarkRevertDelivered(1, tr.ID)) + + _, stillRecovered := c.LookupByForwardHandshakeSecretInTerminalRecovery(forward) + require.False(t, stillRecovered, + "MarkRevertDelivered proves the agent rotated off the forward secret; recovery slot must be cleared so a leaked forward token cannot recover later") +} + +// OnServersDeleted must also clear any forward-recovery entries so a +// future server with a recycled id cannot inherit a stale credential. +func TestOnServersDeletedClearsForwardHandshakeRecovery(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + forward := tr.HandshakeSecret + + _, err := c.Cancel(tr.ID) + require.NoError(t, err) + _, ok := c.LookupByForwardHandshakeSecretInTerminalRecovery(forward) + require.True(t, ok) + + require.NoError(t, DB.Unscoped().Delete(&model.Server{}, uint64(1)).Error) + c.OnServersDeleted([]uint64{1}) + + _, stillRecovered := c.LookupByForwardHandshakeSecretInTerminalRecovery(forward) + require.False(t, stillRecovered, + "OnServersDeleted must clear forward-recovery so a recycled server id cannot inherit the credential") +} + +// TTL: a recovery entry older than defaultRevertDeliveryRecoveryWindow must +// be pruned on read. Same bound as the revert recovery channel so operators +// only have one window to reason about. +func TestForwardHandshakeRecoveryExpiresAfterWindow(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + forward := tr.HandshakeSecret + + _, err := c.Cancel(tr.ID) + require.NoError(t, err) + _, ok := c.LookupByForwardHandshakeSecretInTerminalRecovery(forward) + require.True(t, ok, "precondition: cancel populated forward recovery") + + // Back-date the in-memory entry past the recovery window. We poke the + // private slot via a helper so the test does not depend on time.Now() + // monkey-patching. + forceForwardRecoveryAge(c, 1, defaultRevertDeliveryRecoveryWindow+time.Minute) + + _, stillRecovered := c.LookupByForwardHandshakeSecretInTerminalRecovery(forward) + require.False(t, stillRecovered, + "forward-recovery lookup must prune entries past defaultRevertDeliveryRecoveryWindow on read") +} + +// UNIFIED TERMINAL RECOVERY (HIGH): the bounded "transfer terminated but +// agent may still hold one of its per-transfer secrets" window is one +// concept, not two. Both the forward HandshakeSecret (committed via the +// agent's 10s reload before Cancel landed) and the RevertHandshakeSecret +// (rollback ApplyConfig pushed, agent hasn't ACKed yet) need the same +// bounded acceptance — same TTL, same eviction triggers (Register on a +// new transfer for the same server, MarkRevertDelivered, MarkVerified on +// a newer transfer, OnServersDeleted). They differ only in which secret +// field on the same model.ServerTransfer is being presented. Express +// that in one table with a kind tag, not two parallel tables. +// +// This batch of tests pins down the unified surface: +// - LookupByTerminalSecretRecovery dispatches by which secret matched +// - one revertTransition call registers BOTH kinds in one slot +// - Register-on-Retry preserves the slot (a fresh transfer for the same +// server does NOT wipe rollback recovery the agent may still need) +// - MarkRevertDelivered / MarkVerified / OnServersDeleted clear it +// - the existing per-kind lookups remain as thin wrappers so callers +// outside the singleton don't have to know about kind + +type terminalSecretRecoveryMatch struct { + transfer *model.ServerTransfer + kind TerminalRecoveryKind +} + +func lookupTerminalRecoveryForTest(c *ServerTransferClass, secret string) (terminalSecretRecoveryMatch, bool) { + transfer, kind, ok := c.LookupByTerminalSecretRecovery(secret) + if !ok { + return terminalSecretRecoveryMatch{}, false + } + return terminalSecretRecoveryMatch{transfer: transfer, kind: kind}, true +} + +func TestTerminalSecretRecoveryRegistersBothKindsOnRevertTransition(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + forward := tr.HandshakeSecret + revert := tr.RevertHandshakeSecret + + _, err := c.Cancel(tr.ID) + require.NoError(t, err) + + gotF, okF := lookupTerminalRecoveryForTest(c, forward) + require.True(t, okF, "forward HandshakeSecret must resolve from terminalSecretRecovery after Cancel") + require.Equal(t, TerminalRecoveryForward, gotF.kind, "lookup must report the kind so auth can decide whether to promote") + require.Equal(t, tr.ID, gotF.transfer.ID) + + gotR, okR := lookupTerminalRecoveryForTest(c, revert) + require.True(t, okR, "RevertHandshakeSecret must resolve from the SAME terminalSecretRecovery slot") + require.Equal(t, TerminalRecoveryRevert, gotR.kind) + require.Equal(t, tr.ID, gotR.transfer.ID) +} + +// Per-kind wrappers must continue to work — they are the public-facing +// API existing call sites (and the auth layer) use. +func TestTerminalSecretRecoveryPerKindWrappersStayConsistent(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + forward := tr.HandshakeSecret + revert := tr.RevertHandshakeSecret + + _, err := c.Cancel(tr.ID) + require.NoError(t, err) + + gotF, okF := c.LookupByForwardHandshakeSecretInTerminalRecovery(forward) + require.True(t, okF) + require.Equal(t, tr.ID, gotF.ID) + + gotR, okR := c.LookupByRevertHandshakeSecret(revert) + require.True(t, okR) + require.Equal(t, tr.ID, gotR.ID) +} + +// Register-on-Retry: a fresh pending transfer for the same server MUST +// NOT evict the prior transfer's rollback recovery — the agent's reload +// timer is still running and the agent may not have rotated off the +// previous RevertHandshakeSecret yet. The forward recovery for the prior +// transfer is moot once a new transfer starts pushing a new +// HandshakeSecret, but the revert recovery must survive. +// +// This is the exact invariant TestRegisterPreservesInflightRollbackSecret +// AcrossRetry pins down today via revertRecovery; it must still hold after +// the unified-table refactor. +func TestTerminalSecretRecoveryPreservesRollbackAcrossRetry(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + t1 := initiateAndRegister(t, c, 1, 100, 200, 1) + revertSecret := t1.RevertHandshakeSecret + require.NotEmpty(t, revertSecret) + + _, err := c.Cancel(t1.ID) + require.NoError(t, err) + + initiateAndRegister(t, c, 1, 100, 200, 1) + + got, ok := c.LookupByRevertHandshakeSecret(revertSecret) + require.True(t, ok, + "unified terminalSecretRecovery must preserve the previous transfer's RevertHandshakeSecret across a Retry — the agent's on-disk credential may still be the prior revert secret during the 10s reload") + require.Equal(t, t1.ID, got.ID) + + if _, stillVerified := c.LookupServerByVerifiedHandshakeSecret(revertSecret); stillVerified { + t.Fatal("recovery slot must NOT promote into verifiedHandshakes — that bypasses the bounded window") + } +} + +// REGRESSION: re-invoking Cancel against an already-Cancelled transfer must +// be a true no-op. The cancelServerTransfer HTTP handler does not gate on +// `t.Status == Pending`, so a stale terminal id can reach revertTransition +// via UI replay / lingering tabs / scripted retries. revertTransition's +// transaction returns early for non-Pending rows (transitionedByThisCall +// stays false), but the post-transaction code historically only suppressed +// side effects via `if t.Status != newStatus` — which is FALSE when both +// sides are Cancelled. The fall-through re-registered the OLD transfer's +// revertDelivery and pushed its RevertHandshakeSecret, after a Retry had +// already installed a NEW Pending transfer and delivered its forward +// HandshakeSecret. The agent's ApplyConfig is last-arrival-wins inside the +// 10s reload window, so the stale rollback overwrites the new credential +// and the new transfer is stranded until the 24h timeout sweep. The fix +// gates ALL post-tx side effects on transitionedByThisCall. +func TestServerTransferRepeatedCancelOnTerminalDoesNotResendStaleRollback(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + UserLock.Lock() + UserInfoMap[100] = model.UserInfo{AgentSecret: "old-owner-secret"} + UserInfoMap[200] = model.UserInfo{AgentSecret: "new-owner-secret"} + UserLock.Unlock() + + first := initiateAndRegister(t, c, 1, 100, 200, 1) + _, err := c.Cancel(first.ID) + require.NoError(t, err) + + // Admin Retries the failed transfer. Register clears the old + // revertDeliveries entry and pushes the NEW transfer's HandshakeSecret; + // the agent is now committed to the new credential. + var refreshedFirst model.ServerTransfer + require.NoError(t, DB.First(&refreshedFirst, first.ID).Error) + second, err := c.Retry(&refreshedFirst, 1) + require.NoError(t, err) + require.True(t, c.HasPending(1), "precondition: Retry must register a fresh Pending transfer") + + stream := newFakeTaskStream() + s, ok := ServerShared.Get(1) + require.True(t, ok) + s.SetTaskStream(stream) + // Push the new transfer's ApplyConfig so the agent is on the new + // HandshakeSecret. After this, the stream must NOT see another + // per-transfer secret unless something authoritative changes. + c.PushIfOnline(second) + require.Equal(t, 1, stream.sendCount(), "precondition: new transfer's HandshakeSecret must be the latest ApplyConfig on the wire") + require.Contains(t, stream.sent[0].Data, second.HandshakeSecret) + stream.reset() + + // Stale terminal-id Cancel arrives (UI replay / lingering session / etc). + // `cancelServerTransfer` does not pre-gate on status, so it reaches + // revertTransition with the historical terminal row. + result, err := c.Cancel(first.ID) + require.NoError(t, err) + require.Nil(t, result, + "Cancel on an already-Cancelled row must be a silent no-op — no rollback re-delivery, no recovery re-registration") + + require.Equal(t, 0, stream.sendCount(), + "a stale terminal Cancel must NOT push the OLD transfer's RevertHandshakeSecret — doing so supersedes the new transfer's just-applied HandshakeSecret and strands the new transfer until the 24h timeout") + + // The new transfer's runtime state must be intact: its pending entry, + // its revertDelivery absence, and the agent's last-known credential + // (still the new HandshakeSecret) must all be unchanged. + require.True(t, c.HasPending(1), "fresh Pending transfer must survive a stale terminal Cancel") + got, ok := c.LookupPending(1) + require.True(t, ok) + require.Equal(t, second.ID, got.ID, "in-memory pending must still point at the new transfer") + + if existing, ok := c.LookupRevertDelivery(1); ok { + require.NotEqual(t, first.ID, existing.ID, + "stale Cancel must NOT re-install the OLD transfer's revertDelivery and overwrite the fresh push queue state") + } +} + +// REGRESSION: dashboard restart must NOT rehydrate the +// revertDelivery / terminalSecretRecovery slots for transfers whose rollback +// has already been ACKed via MarkRevertDelivered. The auth path treats an +// entry in revertDeliveries as proof that the rollback window is still open +// and admits the rolled-back ToUserID's global AgentSecret accordingly +// (service/rpc/auth.go authorizeAgentForUUID's LookupRevertDelivery branch). +// MarkRevertDelivered persists acked_at and clears the in-memory delivery +// precisely to close that window — but NewServerTransferClass loaded all +// terminal rows within the recovery window without filtering acked_at, +// reopening it after every restart. Loading must skip acked rows; the +// acked credential is already rebuilt into verifiedHandshakes via the +// existing acked-row pass. +func TestNewServerTransferClassSkipsAckedRollbackRecovery(t *testing.T) { + c, cleanup := setupTransferFixture(t) + defer cleanup() + seedServerForTransfer(t, 1, 100) + + UserLock.Lock() + UserInfoMap[100] = model.UserInfo{AgentSecret: "from-secret"} + UserInfoMap[200] = model.UserInfo{AgentSecret: "to-secret"} + UserLock.Unlock() + + tr := initiateAndRegister(t, c, 1, 100, 200, 1) + _, err := c.Cancel(tr.ID) + require.NoError(t, err) + require.NoError(t, c.MarkRevertDelivered(1, tr.ID), + "precondition: rollback must be ACKed so the in-memory delivery is consumed") + require.False(t, hasRevertDeliveryFor(c, 1, tr.ID), + "precondition: MarkRevertDelivered must clear the in-memory delivery") + + // Simulate dashboard restart against the same DB + ServerShared. + c.Stop() + reborn := NewServerTransferClass() + defer reborn.Stop() + ServerTransferShared = reborn + + if _, ok := reborn.LookupRevertDelivery(1); ok { + t.Fatal("restart must NOT rehydrate an already-ACKed rollback into revertDeliveries — reopening the auth tolerance window for the ToUserID global AgentSecret contradicts MarkRevertDelivered's contract") + } + + // terminalSecretRecovery must also be empty for the ACKed rollback — + // auth's terminal-recovery lookups would otherwise readmit the per- + // transfer secrets the agent has already rotated past. + if got, _, ok := reborn.LookupByTerminalSecretRecovery(tr.RevertHandshakeSecret); ok { + t.Fatalf("restart must NOT rehydrate ACKed RevertHandshakeSecret into terminalSecretRecovery; got=%v", got) + } + if got, _, ok := reborn.LookupByTerminalSecretRecovery(tr.HandshakeSecret); ok { + t.Fatalf("restart must NOT rehydrate ACKed forward HandshakeSecret into terminalSecretRecovery; got=%v", got) + } + + // Sanity: the long-term verifiedHandshakes credential must still be + // rebuilt from the same row's acked_at, so the agent can keep + // authenticating with the rotated RevertHandshakeSecret. + sid, ok := reborn.LookupServerByVerifiedHandshakeSecret(tr.RevertHandshakeSecret) + require.True(t, ok, "ACKed rollback secret must still be rebuilt into verifiedHandshakes — the agent on disk holds exactly this credential") + require.Equal(t, uint64(1), sid) +} diff --git a/service/singleton/servicesentinel.go b/service/singleton/servicesentinel.go index 08b18313..ef1b4b87 100644 --- a/service/singleton/servicesentinel.go +++ b/service/singleton/servicesentinel.go @@ -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 { diff --git a/service/singleton/singleton.go b/service/singleton/singleton.go index b04be1a2..ed808e30 100644 --- a/service/singleton/singleton.go +++ b/service/singleton/singleton.go @@ -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 } diff --git a/service/singleton/user.go b/service/singleton/user.go index aef06260..69fdce02 100644 --- a/service/singleton/user.go +++ b/service/singleton/user.go @@ -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 (#)" +// 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) } diff --git a/service/singleton/user_test.go b/service/singleton/user_test.go new file mode 100644 index 00000000..35f6ff9f --- /dev/null +++ b/service/singleton/user_test.go @@ -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()) + } + } + } +}