mirror of
https://github.com/Buriburizaem0n/nezha_domains.git
synced 2026-09-19 17:50:12 +00:00
feat: server transfer rotation
This commit is contained in:
@@ -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")))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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", "<html>admin index</html>")
|
||||
writeFrontendFallbackTestFile(t, "admin-dist/assets/app.js", "console.log('admin asset')")
|
||||
writeFrontendFallbackTestFile(t, "user-dist/index.html", "<html>user index</html>")
|
||||
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())
|
||||
}
|
||||
}
|
||||
@@ -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}`))
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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"`
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user