test: cover multi-user permission boundaries

Co-authored-by: naiba/CloudCode <hi+cloudcode@nai.ba>
This commit is contained in:
naiba
2026-05-17 10:24:19 +08:00
co-authored by naiba/CloudCode
parent 2573ba7522
commit 7061651587
5 changed files with 801 additions and 15 deletions
+45 -3
View File
@@ -115,16 +115,34 @@ func TestValidateServersRejectsForeignTriggerTasks(t *testing.T) {
func newMemberValidationContext(t *testing.T) *gin.Context { func newMemberValidationContext(t *testing.T) *gin.Context {
t.Helper() 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 originalDB := singleton.DB
originalLoc := singleton.Loc originalLoc := singleton.Loc
originalLocalizer := singleton.Localizer originalLocalizer := singleton.Localizer
originalCronShared := singleton.CronShared originalCronShared := singleton.CronShared
originalServerShared := singleton.ServerShared originalServerShared := singleton.ServerShared
originalUserInfo := singleton.UserInfoMap
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
assert.NoError(t, err) 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{ assert.NoError(t, db.Create(&model.Cron{
Common: model.Common{ID: 42, UserID: 1}, Common: model.Common{ID: 42, UserID: 1},
Name: "foreign trigger task", Name: "foreign trigger task",
@@ -132,25 +150,49 @@ func newMemberValidationContext(t *testing.T) *gin.Context {
TaskType: model.CronTypeTriggerTask, TaskType: model.CronTypeTriggerTask,
Cover: model.CronCoverAlertTrigger, Cover: model.CronCoverAlertTrigger,
}).Error) }).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.DB = db
singleton.Loc = time.Local singleton.Loc = time.Local
singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations) singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations)
singleton.CronShared = singleton.NewCronClass() singleton.CronShared = singleton.NewCronClass()
singleton.ServerShared = singleton.NewServerClass() 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() { t.Cleanup(func() {
singleton.DB = originalDB singleton.DB = originalDB
singleton.Loc = originalLoc singleton.Loc = originalLoc
singleton.Localizer = originalLocalizer singleton.Localizer = originalLocalizer
singleton.CronShared = originalCronShared singleton.CronShared = originalCronShared
singleton.ServerShared = originalServerShared singleton.ServerShared = originalServerShared
singleton.UserLock.Lock()
singleton.UserInfoMap = originalUserInfo
singleton.UserLock.Unlock()
}) })
gin.SetMode(gin.TestMode) gin.SetMode(gin.TestMode)
ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) ctx, _ := gin.CreateTestContext(httptest.NewRecorder())
ctx.Set(model.CtxKeyAuthorizedUser, &model.User{ ctx.Set(model.CtxKeyAuthorizedUser, &model.User{
Common: model.Common{ID: 200}, Common: model.Common{ID: userID},
Role: model.RoleMember, Role: role,
}) })
return ctx return ctx
} }
@@ -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)
}
+38
View File
@@ -1,11 +1,49 @@
package model package model
import ( import (
"net/http/httptest"
"reflect" "reflect"
"slices" "slices"
"testing" "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) { func TestSearchByID(t *testing.T) {
t.Run("WithoutPriorityList", func(t *testing.T) { t.Run("WithoutPriorityList", func(t *testing.T) {
list, exp := []*DDNSProfile{ list, exp := []*DDNSProfile{
+72 -12
View File
@@ -273,24 +273,69 @@ func TestNotificationSendRejectsLoopbackTarget(t *testing.T) {
} }
} }
func TestNotificationTargetRejectsSpecialUseAddresses(t *testing.T) { func TestNotificationTargetRejectsBlockedRanges(t *testing.T) {
cases := []string{ cases := []string{
"http://100.64.0.1/", // CGNAT "http://0.0.0.0/",
"http://192.0.2.1/", // documentation range "http://10.1.2.3/",
"http://[fc00::1]/", // IPv6 unique local "http://100.64.0.1/",
"http://[2001:db8::1]/", // IPv6 documentation range "http://127.0.0.1/",
"http://[::ffff:127.0.0.1]/", // IPv4-mapped loopback "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 { for _, rawURL := range cases {
if _, _, err := resolveNotificationTarget(rawURL); err == nil { t.Run(rawURL, func(t *testing.T) {
t.Fatalf("expected %s to be rejected", rawURL) 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) { 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 { if err != nil {
t.Fatalf("expected public HTTPS URL to create client: %v", err) t.Fatalf("expected public HTTPS URL to create client: %v", err)
} }
@@ -298,7 +343,22 @@ func TestNotificationHTTPClientPreservesTLSServerName(t *testing.T) {
if !ok { if !ok {
t.Fatalf("expected http.Transport, got %T", client.Transport) t.Fatalf("expected http.Transport, got %T", client.Transport)
} }
if transport.TLSClientConfig == nil || transport.TLSClientConfig.ServerName != "example.com" { if transport.TLSClientConfig == nil || transport.TLSClientConfig.ServerName != "1.1.1.1" {
t.Fatalf("expected TLS ServerName example.com, got %#v", transport.TLSClientConfig) 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)
} }
} }
@@ -2,9 +2,12 @@ package singleton
import ( import (
"context" "context"
"net/http/httptest"
"slices"
"testing" "testing"
"time" "time"
"github.com/gin-gonic/gin"
"github.com/nezhahq/nezha/model" "github.com/nezhahq/nezha/model"
pb "github.com/nezhahq/nezha/proto" pb "github.com/nezhahq/nezha/proto"
"google.golang.org/grpc/metadata" "google.golang.org/grpc/metadata"
@@ -135,3 +138,242 @@ func assertNoTask(t *testing.T, stream *capturedTaskStream) {
case <-time.After(50 * time.Millisecond): 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")
}
}