From 7061651587d54bdf6b3a79b31f8c1d3389044da8 Mon Sep 17 00:00:00 2001 From: naiba Date: Fri, 15 May 2026 02:46:27 +0000 Subject: [PATCH] test: cover multi-user permission boundaries Co-authored-by: naiba/CloudCode --- cmd/dashboard/controller/jwt_test.go | 48 ++- .../controller/permission_matrix_test.go | 404 ++++++++++++++++++ model/common_test.go | 38 ++ model/notification_test.go | 84 +++- service/singleton/security_regression_test.go | 242 +++++++++++ 5 files changed, 801 insertions(+), 15 deletions(-) create mode 100644 cmd/dashboard/controller/permission_matrix_test.go diff --git a/cmd/dashboard/controller/jwt_test.go b/cmd/dashboard/controller/jwt_test.go index 9c21ac8b..b2392f29 100644 --- a/cmd/dashboard/controller/jwt_test.go +++ b/cmd/dashboard/controller/jwt_test.go @@ -115,16 +115,34 @@ func TestValidateServersRejectsForeignTriggerTasks(t *testing.T) { func newMemberValidationContext(t *testing.T) *gin.Context { t.Helper() + return newValidationContext(t, 200, model.RoleMember) +} + +func newAdminValidationContext(t *testing.T) *gin.Context { + t.Helper() + return newValidationContext(t, 1, model.RoleAdmin) +} + +func newValidationContext(t *testing.T, userID uint64, role model.Role) *gin.Context { + t.Helper() originalDB := singleton.DB originalLoc := singleton.Loc originalLocalizer := singleton.Localizer originalCronShared := singleton.CronShared originalServerShared := singleton.ServerShared + originalUserInfo := singleton.UserInfoMap db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) assert.NoError(t, err) - assert.NoError(t, db.AutoMigrate(&model.Cron{}, &model.Server{})) + assert.NoError(t, db.AutoMigrate( + &model.Cron{}, + &model.Server{}, + &model.NotificationGroup{}, + &model.NotificationGroupNotification{}, + &model.ServerGroup{}, + &model.ServerGroupServer{}, + )) assert.NoError(t, db.Create(&model.Cron{ Common: model.Common{ID: 42, UserID: 1}, Name: "foreign trigger task", @@ -132,25 +150,49 @@ func newMemberValidationContext(t *testing.T) *gin.Context { TaskType: model.CronTypeTriggerTask, Cover: model.CronCoverAlertTrigger, }).Error) + assert.NoError(t, db.Create(&model.Cron{ + Common: model.Common{ID: 43, UserID: 200}, + Name: "member trigger task", + Command: "member-task", + TaskType: model.CronTypeTriggerTask, + Cover: model.CronCoverAlertTrigger, + }).Error) + assert.NoError(t, db.Create(&model.NotificationGroup{ + Common: model.Common{ID: 7, UserID: 1}, + Name: "admin group", + }).Error) + assert.NoError(t, db.Create(&model.NotificationGroup{ + Common: model.Common{ID: 8, UserID: 200}, + Name: "member group", + }).Error) singleton.DB = db singleton.Loc = time.Local singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations) singleton.CronShared = singleton.NewCronClass() singleton.ServerShared = singleton.NewServerClass() + singleton.UserLock.Lock() + singleton.UserInfoMap = map[uint64]model.UserInfo{ + 1: {Role: model.RoleAdmin}, + 200: {Role: model.RoleMember}, + } + singleton.UserLock.Unlock() t.Cleanup(func() { singleton.DB = originalDB singleton.Loc = originalLoc singleton.Localizer = originalLocalizer singleton.CronShared = originalCronShared singleton.ServerShared = originalServerShared + singleton.UserLock.Lock() + singleton.UserInfoMap = originalUserInfo + singleton.UserLock.Unlock() }) gin.SetMode(gin.TestMode) ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) ctx.Set(model.CtxKeyAuthorizedUser, &model.User{ - Common: model.Common{ID: 200}, - Role: model.RoleMember, + Common: model.Common{ID: userID}, + Role: role, }) return ctx } diff --git a/cmd/dashboard/controller/permission_matrix_test.go b/cmd/dashboard/controller/permission_matrix_test.go new file mode 100644 index 00000000..3b6fa7aa --- /dev/null +++ b/cmd/dashboard/controller/permission_matrix_test.go @@ -0,0 +1,404 @@ +package controller + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/gin-gonic/gin" + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/service/singleton" + "github.com/stretchr/testify/assert" +) + +func TestValidateRuleAcceptsMemberSelfTriggerTasks(t *testing.T) { + ctx := newMemberValidationContext(t) + rule := &model.AlertRule{ + Common: model.Common{UserID: 200}, + Name: "member alert", + Rules: []*model.Rule{{Type: "offline", Duration: 3}}, + FailTriggerTasks: []uint64{43}, + RecoverTriggerTasks: []uint64{43}, + } + assert.NoError(t, validateRule(ctx, rule)) +} + +func TestValidateRuleAcceptsAdminCrossUserTriggerTasks(t *testing.T) { + ctx := newAdminValidationContext(t) + rule := &model.AlertRule{ + Common: model.Common{UserID: 1}, + Name: "admin alert", + Rules: []*model.Rule{{Type: "offline", Duration: 3}}, + FailTriggerTasks: []uint64{43}, + RecoverTriggerTasks: []uint64{43}, + } + assert.NoError(t, validateRule(ctx, rule)) +} + +func TestValidateRuleAcceptsEmptyTriggerTasks(t *testing.T) { + ctx := newMemberValidationContext(t) + rule := &model.AlertRule{ + Common: model.Common{UserID: 200}, + Name: "member alert", + Rules: []*model.Rule{{Type: "offline", Duration: 3}}, + } + assert.NoError(t, validateRule(ctx, rule)) +} + +func TestValidateRuleAcceptsUnknownTriggerTaskID(t *testing.T) { + ctx := newMemberValidationContext(t) + rule := &model.AlertRule{ + Common: model.Common{UserID: 200}, + Name: "member alert", + Rules: []*model.Rule{{Type: "offline", Duration: 3}}, + FailTriggerTasks: []uint64{9999}, + } + assert.NoError(t, validateRule(ctx, rule)) +} + +func TestValidateRuleRejectsForeignNotificationGroup(t *testing.T) { + ctx := newMemberValidationContext(t) + rule := &model.AlertRule{ + Common: model.Common{UserID: 200}, + Name: "member alert", + Rules: []*model.Rule{{Type: "offline", Duration: 3}}, + NotificationGroupID: 7, + } + assert.Error(t, validateRule(ctx, rule)) +} + +func TestValidateRuleAcceptsMemberOwnedNotificationGroup(t *testing.T) { + ctx := newMemberValidationContext(t) + rule := &model.AlertRule{ + Common: model.Common{UserID: 200}, + Name: "member alert", + Rules: []*model.Rule{{Type: "offline", Duration: 3}}, + NotificationGroupID: 8, + } + assert.NoError(t, validateRule(ctx, rule)) +} + +func TestValidateRuleAdminCanReferenceAnyNotificationGroup(t *testing.T) { + ctx := newAdminValidationContext(t) + rule := &model.AlertRule{ + Common: model.Common{UserID: 1}, + Name: "admin alert", + Rules: []*model.Rule{{Type: "offline", Duration: 3}}, + NotificationGroupID: 8, + } + assert.NoError(t, validateRule(ctx, rule)) +} + +func TestValidateServersRejectsForeignNotificationGroup(t *testing.T) { + ctx := newMemberValidationContext(t) + service := &model.Service{ + Common: model.Common{UserID: 200}, + Name: "member service", + SkipServers: map[uint64]bool{}, + NotificationGroupID: 7, + } + assert.Error(t, validateServers(ctx, service)) +} + +func TestValidateServersAcceptsMemberOwnedNotificationGroup(t *testing.T) { + ctx := newMemberValidationContext(t) + service := &model.Service{ + Common: model.Common{UserID: 200}, + Name: "member service", + SkipServers: map[uint64]bool{}, + NotificationGroupID: 8, + } + assert.NoError(t, validateServers(ctx, service)) +} + +func TestUserCanViewServer(t *testing.T) { + memberServer := &model.Server{Common: model.Common{ID: 1, UserID: 200}} + adminServer := &model.Server{Common: model.Common{ID: 2, UserID: 1}} + hiddenAdminServer := &model.Server{Common: model.Common{ID: 3, UserID: 1}, HideForGuest: true} + publicAdminServer := &model.Server{Common: model.Common{ID: 4, UserID: 1}, HideForGuest: false} + + cases := []struct { + name string + setup func(c *gin.Context) + server *model.Server + wantAllow bool + }{ + { + name: "guest sees public", + setup: func(c *gin.Context) {}, + server: publicAdminServer, + wantAllow: true, + }, + { + name: "guest blocked by HideForGuest", + setup: func(c *gin.Context) {}, + server: hiddenAdminServer, + wantAllow: false, + }, + { + name: "member sees own", + setup: func(c *gin.Context) { + c.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: 200}, Role: model.RoleMember}) + }, + server: memberServer, + wantAllow: true, + }, + { + name: "member can still see public foreign", + setup: func(c *gin.Context) { + c.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: 200}, Role: model.RoleMember}) + }, + server: publicAdminServer, + wantAllow: true, + }, + { + name: "member blocked by HideForGuest foreign", + setup: func(c *gin.Context) { + c.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: 200}, Role: model.RoleMember}) + }, + server: hiddenAdminServer, + wantAllow: false, + }, + { + name: "admin sees hidden foreign", + setup: func(c *gin.Context) { + c.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: 1}, Role: model.RoleAdmin}) + }, + server: hiddenAdminServer, + wantAllow: true, + }, + { + name: "admin sees other admin server", + setup: func(c *gin.Context) { + c.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: 1}, Role: model.RoleAdmin}) + }, + server: adminServer, + wantAllow: true, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + gin.SetMode(gin.TestMode) + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + tc.setup(ctx) + if got := userCanViewServer(ctx, tc.server); got != tc.wantAllow { + t.Fatalf("userCanViewServer = %v, want %v", got, tc.wantAllow) + } + }) + } +} + +type permissionTestResource struct { + ID uint64 `json:"id"` + UserID uint64 `json:"user_id"` +} + +func (r *permissionTestResource) GetID() uint64 { return r.ID } +func (r *permissionTestResource) GetUserID() uint64 { return r.UserID } +func (r *permissionTestResource) HasPermission(c *gin.Context) bool { + auth, ok := c.Get(model.CtxKeyAuthorizedUser) + if !ok { + return false + } + user := *auth.(*model.User) + if user.Role == model.RoleAdmin { + return true + } + return user.ID == r.UserID +} + +func TestListHandlerFiltersByOwnership(t *testing.T) { + gin.SetMode(gin.TestMode) + data := []*permissionTestResource{ + {ID: 1, UserID: 100}, + {ID: 2, UserID: 200}, + {ID: 3, UserID: 200}, + } + handler := listHandler(func(c *gin.Context) ([]*permissionTestResource, error) { + return append([]*permissionTestResource{}, data...), nil + }) + + t.Run("member only sees own", func(t *testing.T) { + r := gin.New() + r.Use(func(c *gin.Context) { + c.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: 200}, Role: model.RoleMember}) + c.Next() + }) + r.GET("/test", handler) + w := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/test", nil) + r.ServeHTTP(w, req) + + ids := decodeIDs[uint64](t, w.Body.Bytes()) + assert.ElementsMatch(t, []uint64{2, 3}, ids) + }) + + t.Run("admin sees all", func(t *testing.T) { + r := gin.New() + r.Use(func(c *gin.Context) { + c.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: 1}, Role: model.RoleAdmin}) + c.Next() + }) + r.GET("/test", handler) + w := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/test", nil) + r.ServeHTTP(w, req) + + ids := decodeIDs[uint64](t, w.Body.Bytes()) + assert.ElementsMatch(t, []uint64{1, 2, 3}, ids) + }) +} + +func decodeIDs[T ~uint64](t *testing.T, body []byte) []T { + t.Helper() + var resp struct { + Data []map[string]any `json:"data"` + } + if err := json.Unmarshal(body, &resp); err != nil { + t.Fatalf("decode response: %v body=%s", err, string(body)) + } + ids := make([]T, 0, len(resp.Data)) + for _, item := range resp.Data { + switch v := item["id"].(type) { + case float64: + ids = append(ids, T(v)) + case json.Number: + n, _ := v.Int64() + ids = append(ids, T(n)) + } + } + return ids +} + +func TestCallerIsAdmin(t *testing.T) { + gin.SetMode(gin.TestMode) + + tests := []struct { + name string + setup func(c *gin.Context) + want bool + }{ + {name: "unauth", setup: func(c *gin.Context) {}, want: false}, + { + name: "member", + setup: func(c *gin.Context) { + c.Set(model.CtxKeyAuthorizedUser, &model.User{Role: model.RoleMember}) + }, + want: false, + }, + { + name: "admin", + setup: func(c *gin.Context) { + c.Set(model.CtxKeyAuthorizedUser, &model.User{Role: model.RoleAdmin}) + }, + want: true, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + tc.setup(ctx) + if got := callerIsAdmin(ctx); got != tc.want { + t.Fatalf("callerIsAdmin = %v, want %v", got, tc.want) + } + }) + } +} + +func TestAssertOwnsNotificationGroup(t *testing.T) { + memberCtx := newMemberValidationContext(t) + assert.NoError(t, assertOwnsNotificationGroup(memberCtx, 0)) + assert.NoError(t, assertOwnsNotificationGroup(memberCtx, 8)) + assert.Error(t, assertOwnsNotificationGroup(memberCtx, 7)) + assert.ErrorContains(t, assertOwnsNotificationGroup(memberCtx, 9999), "does not exist") + + adminCtx := newAdminValidationContext(t) + assert.NoError(t, assertOwnsNotificationGroup(adminCtx, 7)) + assert.NoError(t, assertOwnsNotificationGroup(adminCtx, 8)) +} + +func TestListServerGroupFiltersByOwnership(t *testing.T) { + ctx := newMemberValidationContext(t) + assert.NoError(t, singleton.DB.Create(&model.ServerGroup{Common: model.Common{ID: 1, UserID: 200}, Name: "member group"}).Error) + assert.NoError(t, singleton.DB.Create(&model.ServerGroup{Common: model.Common{ID: 2, UserID: 1}, Name: "admin group"}).Error) + + got, err := listServerGroup(ctx) + assert.NoError(t, err) + var names []string + for _, g := range got { + names = append(names, g.Group.Name) + } + assert.ElementsMatch(t, []string{"member group"}, names) +} + +func TestListServerGroupAdminSeesAll(t *testing.T) { + ctx := newAdminValidationContext(t) + assert.NoError(t, singleton.DB.Create(&model.ServerGroup{Common: model.Common{ID: 1, UserID: 200}, Name: "member group"}).Error) + assert.NoError(t, singleton.DB.Create(&model.ServerGroup{Common: model.Common{ID: 2, UserID: 1}, Name: "admin group"}).Error) + + got, err := listServerGroup(ctx) + assert.NoError(t, err) + var names []string + for _, g := range got { + names = append(names, g.Group.Name) + } + assert.ElementsMatch(t, []string{"member group", "admin group"}, names) +} + +func TestListNotificationGroupFiltersByOwnership(t *testing.T) { + ctx := newMemberValidationContext(t) + got, err := listNotificationGroup(ctx) + assert.NoError(t, err) + var names []string + for _, g := range got { + names = append(names, g.Group.Name) + } + assert.ElementsMatch(t, []string{"member group"}, names) +} + +func TestListNotificationGroupAdminSeesAll(t *testing.T) { + ctx := newAdminValidationContext(t) + got, err := listNotificationGroup(ctx) + assert.NoError(t, err) + var names []string + for _, g := range got { + names = append(names, g.Group.Name) + } + assert.ElementsMatch(t, []string{"admin group", "member group"}, names) +} + +func TestBatchMoveServerRejectsNonAdminCrossUser(t *testing.T) { + ctx := newMemberValidationContext(t) + ctx.Request = httptest.NewRequest(http.MethodPost, "/batch-move/server", strings.NewReader(`{"ids":[],"to_user":1}`)) + ctx.Request.Header.Set("Content-Type", "application/json") + _, err := batchMoveServer(ctx) + assert.Error(t, err) +} + +func TestBatchMoveServerAllowsMemberSelfMove(t *testing.T) { + ctx := newMemberValidationContext(t) + ctx.Request = httptest.NewRequest(http.MethodPost, "/batch-move/server", strings.NewReader(`{"ids":[],"to_user":200}`)) + ctx.Request.Header.Set("Content-Type", "application/json") + _, err := batchMoveServer(ctx) + assert.NoError(t, err) +} + +func TestBatchMoveServerAllowsAdminCrossUser(t *testing.T) { + ctx := newAdminValidationContext(t) + ctx.Request = httptest.NewRequest(http.MethodPost, "/batch-move/server", strings.NewReader(`{"ids":[],"to_user":200}`)) + ctx.Request.Header.Set("Content-Type", "application/json") + _, err := batchMoveServer(ctx) + assert.NoError(t, err) +} + +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}`)) + ctx.Request.Header.Set("Content-Type", "application/json") + _, err := createNAT(ctx) + assert.Error(t, err) +} diff --git a/model/common_test.go b/model/common_test.go index ed4d94c8..0d773d50 100644 --- a/model/common_test.go +++ b/model/common_test.go @@ -1,11 +1,49 @@ package model import ( + "net/http/httptest" "reflect" "slices" "testing" + + "github.com/gin-gonic/gin" ) +func TestCommonHasPermission(t *testing.T) { + resource := &Common{ID: 10, UserID: 100} + + t.Run("unauthenticated denied", func(t *testing.T) { + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + if resource.HasPermission(ctx) { + t.Fatal("expected unauthenticated request to be denied") + } + }) + + t.Run("owner allowed", func(t *testing.T) { + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 100}, Role: RoleMember}) + if !resource.HasPermission(ctx) { + t.Fatal("expected owner to be allowed") + } + }) + + t.Run("foreign member denied", func(t *testing.T) { + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 200}, Role: RoleMember}) + if resource.HasPermission(ctx) { + t.Fatal("expected non-owner member to be denied") + } + }) + + t.Run("admin allowed", func(t *testing.T) { + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 1}, Role: RoleAdmin}) + if !resource.HasPermission(ctx) { + t.Fatal("expected admin to be allowed") + } + }) +} + func TestSearchByID(t *testing.T) { t.Run("WithoutPriorityList", func(t *testing.T) { list, exp := []*DDNSProfile{ diff --git a/model/notification_test.go b/model/notification_test.go index aee11446..28dd4873 100644 --- a/model/notification_test.go +++ b/model/notification_test.go @@ -273,24 +273,69 @@ func TestNotificationSendRejectsLoopbackTarget(t *testing.T) { } } -func TestNotificationTargetRejectsSpecialUseAddresses(t *testing.T) { +func TestNotificationTargetRejectsBlockedRanges(t *testing.T) { cases := []string{ - "http://100.64.0.1/", // CGNAT - "http://192.0.2.1/", // documentation range - "http://[fc00::1]/", // IPv6 unique local - "http://[2001:db8::1]/", // IPv6 documentation range - "http://[::ffff:127.0.0.1]/", // IPv4-mapped loopback + "http://0.0.0.0/", + "http://10.1.2.3/", + "http://100.64.0.1/", + "http://127.0.0.1/", + "http://127.255.255.254/", + "http://169.254.169.254/", + "http://172.16.0.1/", + "http://192.0.0.1/", + "http://192.0.2.1/", + "http://192.168.1.1/", + "http://198.18.0.1/", + "http://198.51.100.1/", + "http://203.0.113.1/", + "http://224.0.0.1/", + "http://240.0.0.1/", + "http://[::]/", + "http://[::1]/", + "http://[::ffff:127.0.0.1]/", + "http://[64:ff9b::1]/", + "http://[100::1]/", + "http://[2001:0:0:0:0:0:0:1]/", + "http://[2001:db8::1]/", + "http://[fc00::1]/", + "http://[fe80::1]/", + "http://[ff00::1]/", + "ftp://example.com/", + "file:///etc/passwd", + "http:///path", } for _, rawURL := range cases { - if _, _, err := resolveNotificationTarget(rawURL); err == nil { - t.Fatalf("expected %s to be rejected", rawURL) - } + t.Run(rawURL, func(t *testing.T) { + if _, _, err := resolveNotificationTarget(rawURL); err == nil { + t.Fatalf("expected %s to be rejected", rawURL) + } + }) + } +} + +func TestNotificationTargetAllowsPublicAddresses(t *testing.T) { + cases := []string{ + "http://1.1.1.1/path", + "https://8.8.8.8/", + "https://[2606:4700:4700::1111]/", + } + + for _, rawURL := range cases { + t.Run(rawURL, func(t *testing.T) { + parsedURL, _, err := resolveNotificationTarget(rawURL) + if err != nil { + t.Fatalf("expected %s to be allowed, got %v", rawURL, err) + } + if parsedURL == nil { + t.Fatalf("expected parsed url for %s", rawURL) + } + }) } } func TestNotificationHTTPClientPreservesTLSServerName(t *testing.T) { - client, err := newNotificationHTTPClient("https://example.com/webhook", true) + client, err := newNotificationHTTPClient("https://1.1.1.1/webhook", true) if err != nil { t.Fatalf("expected public HTTPS URL to create client: %v", err) } @@ -298,7 +343,22 @@ func TestNotificationHTTPClientPreservesTLSServerName(t *testing.T) { if !ok { t.Fatalf("expected http.Transport, got %T", client.Transport) } - if transport.TLSClientConfig == nil || transport.TLSClientConfig.ServerName != "example.com" { - t.Fatalf("expected TLS ServerName example.com, got %#v", transport.TLSClientConfig) + if transport.TLSClientConfig == nil || transport.TLSClientConfig.ServerName != "1.1.1.1" { + t.Fatalf("expected TLS ServerName 1.1.1.1, got %#v", transport.TLSClientConfig) + } +} + +func TestNotificationHTTPClientRejectsRedirects(t *testing.T) { + client, err := newNotificationHTTPClient("https://1.1.1.1/webhook", true) + if err != nil { + t.Fatalf("expected client construction: %v", err) + } + req, err := http.NewRequest(http.MethodGet, "https://1.1.1.1/start", nil) + if err != nil { + t.Fatalf("new request: %v", err) + } + via := []*http.Request{req} + if err := client.CheckRedirect(req, via); err != http.ErrUseLastResponse { + t.Fatalf("expected ErrUseLastResponse, got %v", err) } } diff --git a/service/singleton/security_regression_test.go b/service/singleton/security_regression_test.go index f88b33f3..cd3c3fc1 100644 --- a/service/singleton/security_regression_test.go +++ b/service/singleton/security_regression_test.go @@ -2,9 +2,12 @@ package singleton import ( "context" + "net/http/httptest" + "slices" "testing" "time" + "github.com/gin-gonic/gin" "github.com/nezhahq/nezha/model" pb "github.com/nezhahq/nezha/proto" "google.golang.org/grpc/metadata" @@ -135,3 +138,242 @@ func assertNoTask(t *testing.T, stream *capturedTaskStream) { case <-time.After(50 * time.Millisecond): } } + +func TestCronTriggerSendsToMemberOwnedServer(t *testing.T) { + memberStream := newCapturedTaskStream() + replaceServerSharedForSecurityTest(t, + &model.Server{Common: model.Common{ID: 1, UserID: 100}, Name: "member-server", TaskStream: memberStream}, + ) + replaceUserInfoMapForSecurityTest(t, map[uint64]model.UserInfo{ + 100: {Role: model.RoleMember}, + }) + + cronTask := &model.Cron{ + Common: model.Common{ID: 99, UserID: 100}, + Command: "id", + Cover: model.CronCoverAll, + } + + CronTrigger(cronTask)() + + assertTaskCommand(t, memberStream, "id") +} + +func TestCronTriggerAdminCronFansOutAcrossOwners(t *testing.T) { + first := newCapturedTaskStream() + second := newCapturedTaskStream() + replaceServerSharedForSecurityTest(t, + &model.Server{Common: model.Common{ID: 1, UserID: 100}, Name: "member-server", TaskStream: first}, + &model.Server{Common: model.Common{ID: 2, UserID: 200}, Name: "admin-server", TaskStream: second}, + ) + replaceUserInfoMapForSecurityTest(t, map[uint64]model.UserInfo{ + 1: {Role: model.RoleAdmin}, + 100: {Role: model.RoleMember}, + 200: {Role: model.RoleAdmin}, + }) + + cronTask := &model.Cron{ + Common: model.Common{ID: 99, UserID: 1}, + Command: "maintenance", + Cover: model.CronCoverAll, + } + + CronTrigger(cronTask)() + + assertTaskCommand(t, first, "maintenance") + assertTaskCommand(t, second, "maintenance") +} + +func TestCronTriggerLegacyZeroOwnerFansOut(t *testing.T) { + first := newCapturedTaskStream() + replaceServerSharedForSecurityTest(t, + &model.Server{Common: model.Common{ID: 1, UserID: 100}, Name: "member-server", TaskStream: first}, + ) + replaceUserInfoMapForSecurityTest(t, map[uint64]model.UserInfo{ + 100: {Role: model.RoleMember}, + }) + + cronTask := &model.Cron{ + Common: model.Common{ID: 99, UserID: 0}, + Command: "legacy", + Cover: model.CronCoverAll, + } + + CronTrigger(cronTask)() + + assertTaskCommand(t, first, "legacy") +} + +func TestCronTriggerSkipsServersWhenOwnerNotKnown(t *testing.T) { + stream := newCapturedTaskStream() + replaceServerSharedForSecurityTest(t, + &model.Server{Common: model.Common{ID: 1, UserID: 100}, Name: "member-server", TaskStream: stream}, + ) + replaceUserInfoMapForSecurityTest(t, map[uint64]model.UserInfo{ + 100: {Role: model.RoleMember}, + }) + + cronTask := &model.Cron{ + Common: model.Common{ID: 99, UserID: 999}, + Command: "ghost", + Cover: model.CronCoverAll, + } + + CronTrigger(cronTask)() + + assertNoTask(t, stream) +} + +func TestSendTriggerTasksAllowsSelfOwnedCron(t *testing.T) { + stream := newCapturedTaskStream() + replaceServerSharedForSecurityTest(t, + &model.Server{Common: model.Common{ID: 7, UserID: 200}, Name: "member-server", TaskStream: stream}, + ) + replaceUserInfoMapForSecurityTest(t, map[uint64]model.UserInfo{ + 200: {Role: model.RoleMember}, + }) + + memberCron := &model.Cron{ + Common: model.Common{ID: 42, UserID: 200}, + Command: "member-task", + Cover: model.CronCoverAlertTrigger, + } + cronClass := &CronClass{ + class: class[uint64, *model.Cron]{ + list: map[uint64]*model.Cron{memberCron.ID: memberCron}, + }, + } + + cronClass.SendTriggerTasks([]uint64{memberCron.ID}, 7, 200) + + assertTaskCommand(t, stream, "member-task") +} + +func TestSendTriggerTasksAllowsAdminCallerToTriggerAny(t *testing.T) { + stream := newCapturedTaskStream() + replaceServerSharedForSecurityTest(t, + &model.Server{Common: model.Common{ID: 9, UserID: 100}, Name: "any-server", TaskStream: stream}, + ) + replaceUserInfoMapForSecurityTest(t, map[uint64]model.UserInfo{ + 1: {Role: model.RoleAdmin}, + 100: {Role: model.RoleMember}, + }) + + memberCron := &model.Cron{ + Common: model.Common{ID: 42, UserID: 100}, + Command: "member-task", + Cover: model.CronCoverAlertTrigger, + } + cronClass := &CronClass{ + class: class[uint64, *model.Cron]{ + list: map[uint64]*model.Cron{memberCron.ID: memberCron}, + }, + } + + cronClass.SendTriggerTasks([]uint64{memberCron.ID}, 9, 1) + + assertTaskCommand(t, stream, "member-task") +} + +func TestSendTriggerTasksIgnoresUnknownTaskIDs(t *testing.T) { + stream := newCapturedTaskStream() + replaceServerSharedForSecurityTest(t, + &model.Server{Common: model.Common{ID: 7, UserID: 200}, Name: "member-server", TaskStream: stream}, + ) + replaceUserInfoMapForSecurityTest(t, map[uint64]model.UserInfo{ + 200: {Role: model.RoleMember}, + }) + + cronClass := &CronClass{ + class: class[uint64, *model.Cron]{ + list: map[uint64]*model.Cron{}, + }, + } + + cronClass.SendTriggerTasks([]uint64{12345}, 7, 200) + cronClass.SendTriggerTasks(nil, 7, 200) + + assertNoTask(t, stream) +} + +func TestSendTriggerTasksMixedCronIDsOnlyFiresAllowed(t *testing.T) { + stream := newCapturedTaskStream() + replaceServerSharedForSecurityTest(t, + &model.Server{Common: model.Common{ID: 7, UserID: 200}, Name: "member-server", TaskStream: stream}, + ) + replaceUserInfoMapForSecurityTest(t, map[uint64]model.UserInfo{ + 1: {Role: model.RoleAdmin}, + 200: {Role: model.RoleMember}, + }) + + memberCron := &model.Cron{ + Common: model.Common{ID: 7, UserID: 200}, + Command: "member-task", + Cover: model.CronCoverAlertTrigger, + } + adminCron := &model.Cron{ + Common: model.Common{ID: 8, UserID: 1}, + Command: "admin-task", + Cover: model.CronCoverAlertTrigger, + } + cronClass := &CronClass{ + class: class[uint64, *model.Cron]{ + list: map[uint64]*model.Cron{memberCron.ID: memberCron, adminCron.ID: adminCron}, + }, + } + + cronClass.SendTriggerTasks([]uint64{memberCron.ID, adminCron.ID}, 7, 200) + + select { + case task := <-stream.tasks: + if task.GetData() != "member-task" { + t.Fatalf("expected member-task, got %q", task.GetData()) + } + case <-time.After(time.Second): + t.Fatalf("expected member-task to be sent") + } + assertNoTask(t, stream) +} + +func TestClassCheckPermission(t *testing.T) { + replaceUserInfoMapForSecurityTest(t, map[uint64]model.UserInfo{ + 1: {Role: model.RoleAdmin}, + 200: {Role: model.RoleMember}, + }) + sharedClass := &ServerClass{ + class: class[uint64, *model.Server]{ + list: map[uint64]*model.Server{ + 1: {Common: model.Common{ID: 1, UserID: 200}}, + 2: {Common: model.Common{ID: 2, UserID: 1}}, + }, + }, + uuidToID: map[string]uint64{}, + } + + memberCtx, _ := gin.CreateTestContext(httptest.NewRecorder()) + memberCtx.Set(model.CtxKeyAuthorizedUser, &model.User{ + Common: model.Common{ID: 200}, + Role: model.RoleMember, + }) + adminCtx, _ := gin.CreateTestContext(httptest.NewRecorder()) + adminCtx.Set(model.CtxKeyAuthorizedUser, &model.User{ + Common: model.Common{ID: 1}, + Role: model.RoleAdmin, + }) + + if !sharedClass.CheckPermission(memberCtx, slices.Values([]uint64{1})) { + t.Fatal("expected member to access own resource") + } + if sharedClass.CheckPermission(memberCtx, slices.Values([]uint64{2})) { + t.Fatal("expected member to be denied foreign resource") + } + if !sharedClass.CheckPermission(memberCtx, slices.Values([]uint64{})) { + t.Fatal("expected empty iterator to be allowed") + } + if !sharedClass.CheckPermission(memberCtx, slices.Values([]uint64{999})) { + t.Fatal("expected unknown id to be ignored (vacuous true)") + } + if !sharedClass.CheckPermission(adminCtx, slices.Values([]uint64{1, 2})) { + t.Fatal("expected admin to access any resource") + } +}