feat: server transfer rotation

This commit is contained in:
naiba
2026-05-25 10:17:34 +00:00
parent 636f4a9716
commit e816e94fca
43 changed files with 7072 additions and 134 deletions
+55 -21
View File
@@ -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")))
}
}
+6 -2
View File
@@ -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}`))
+105 -33
View File
@@ -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.
+6 -2
View File
@@ -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 {
+217
View File
@@ -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"`
}
+1 -1
View File
@@ -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
}
+31 -9
View File
@@ -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()
}