diff --git a/cmd/dashboard/controller/permission_matrix_test.go b/cmd/dashboard/controller/permission_matrix_test.go index 3b6fa7aa..a3c180e4 100644 --- a/cmd/dashboard/controller/permission_matrix_test.go +++ b/cmd/dashboard/controller/permission_matrix_test.go @@ -6,6 +6,7 @@ import ( "net/http/httptest" "strings" "testing" + "time" "github.com/gin-gonic/gin" "github.com/nezhahq/nezha/model" @@ -252,6 +253,88 @@ func TestListHandlerFiltersByOwnership(t *testing.T) { }) } +func TestShowServiceFiltersCycleTransferStatsLikeServerList(t *testing.T) { + newMemberValidationContext(t) + assert.NoError(t, singleton.DB.AutoMigrate(&model.Service{}, &model.ServiceHistory{})) + assert.NoError(t, singleton.DB.Create(&model.Server{Common: model.Common{ID: 1, UserID: 1}, Name: "public server", UUID: "public-server"}).Error) + assert.NoError(t, singleton.DB.Create(&model.Server{Common: model.Common{ID: 2, UserID: 1}, Name: "hidden admin server", UUID: "hidden-admin-server", HideForGuest: true}).Error) + assert.NoError(t, singleton.DB.Create(&model.Server{Common: model.Common{ID: 3, UserID: 200}, Name: "hidden member server", UUID: "hidden-member-server", HideForGuest: true}).Error) + singleton.ServerShared = singleton.NewServerClass() + + assert.NoError(t, singleton.DB.Create(&model.Service{Common: model.Common{ID: 10, UserID: 1}, Name: "shown service", EnableShowInService: true}).Error) + assert.NoError(t, singleton.DB.Create(&model.Service{Common: model.Common{ID: 11, UserID: 1}, Name: "hidden service"}).Error) + + originalServiceSentinel := singleton.ServiceSentinelShared + serviceSentinel, err := singleton.NewServiceSentinel(make(chan *model.Service, 2)) + assert.NoError(t, err) + singleton.ServiceSentinelShared = serviceSentinel + t.Cleanup(func() { singleton.ServiceSentinelShared = originalServiceSentinel }) + + singleton.AlertsLock.Lock() + originalCycleTransferStats := singleton.AlertsCycleTransferStatsStore + singleton.AlertsCycleTransferStatsStore = map[uint64]*model.CycleTransferStats{ + 7: { + Name: "transfer alert", + ServerName: map[uint64]string{1: "public server", 2: "hidden admin server", 3: "hidden member server"}, + Transfer: map[uint64]uint64{1: 100, 2: 200, 3: 300}, + NextUpdate: map[uint64]time.Time{1: time.Unix(1, 0), 2: time.Unix(2, 0), 3: time.Unix(3, 0)}, + }, + } + singleton.AlertsLock.Unlock() + t.Cleanup(func() { + singleton.AlertsLock.Lock() + singleton.AlertsCycleTransferStatsStore = originalCycleTransferStats + singleton.AlertsLock.Unlock() + }) + + tests := []struct { + name string + viewer *model.User + wantNames map[uint64]string + }{ + { + name: "guest sees public servers only", + wantNames: map[uint64]string{1: "public server"}, + }, + { + name: "member sees public and owned hidden servers", + viewer: &model.User{Common: model.Common{ID: 200}, Role: model.RoleMember}, + wantNames: map[uint64]string{1: "public server", 3: "hidden member server"}, + }, + { + name: "admin sees every server", + viewer: &model.User{Common: model.Common{ID: 1}, Role: model.RoleAdmin}, + wantNames: map[uint64]string{1: "public server", 2: "hidden admin server", 3: "hidden member server"}, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + if tc.viewer != nil { + ctx.Set(model.CtxKeyAuthorizedUser, tc.viewer) + } + + got, err := showService(ctx) + assert.NoError(t, err) + assert.Contains(t, got.Services, uint64(10)) + assert.NotContains(t, got.Services, uint64(11)) + if assert.Contains(t, got.CycleTransferStats, uint64(7)) { + cycleStats := got.CycleTransferStats[7] + assert.Equal(t, tc.wantNames, cycleStats.ServerName) + assert.Len(t, cycleStats.Transfer, len(tc.wantNames)) + assert.Len(t, cycleStats.NextUpdate, len(tc.wantNames)) + for serverID := range cycleStats.Transfer { + assert.Contains(t, tc.wantNames, serverID) + } + for serverID := range cycleStats.NextUpdate { + assert.Contains(t, tc.wantNames, serverID) + } + } + }) + } +} + func decodeIDs[T ~uint64](t *testing.T, body []byte) []T { t.Helper() var resp struct { diff --git a/cmd/dashboard/controller/service.go b/cmd/dashboard/controller/service.go index cf7d03e3..b46cab11 100644 --- a/cmd/dashboard/controller/service.go +++ b/cmd/dashboard/controller/service.go @@ -1,6 +1,7 @@ package controller import ( + "fmt" "maps" "slices" "strconv" @@ -26,14 +27,14 @@ import ( // @Success 200 {object} model.CommonResponse[model.ServiceResponse] // @Router /service [get] func showService(c *gin.Context) (*model.ServiceResponse, error) { - res, err, _ := requestGroup.Do("list-service", func() (any, error) { + res, err, _ := requestGroup.Do(serviceResponseCacheKey(c), func() (any, error) { singleton.AlertsLock.RLock() defer singleton.AlertsLock.RUnlock() stats := singleton.ServiceSentinelShared.CopyStats() var cycleTransferStats map[uint64]model.CycleTransferStats copier.Copy(&cycleTransferStats, singleton.AlertsCycleTransferStatsStore) return []any{ - stats, cycleTransferStats, + stats, filterCycleTransferStatsForViewer(c, cycleTransferStats), }, nil }) if err != nil { @@ -46,6 +47,51 @@ func showService(c *gin.Context) (*model.ServiceResponse, error) { }, nil } +func serviceResponseCacheKey(c *gin.Context) string { + auth, ok := c.Get(model.CtxKeyAuthorizedUser) + if !ok { + return "list-service::guest" + } + user, ok := auth.(*model.User) + if !ok || user == nil { + return "list-service::guest" + } + return fmt.Sprintf("list-service::%t::%d", user.Role.IsAdmin(), user.ID) +} + +func filterCycleTransferStatsForViewer(c *gin.Context, stats map[uint64]model.CycleTransferStats) map[uint64]model.CycleTransferStats { + if len(stats) == 0 { + return stats + } + servers := singleton.ServerShared.GetList() + filteredStats := make(map[uint64]model.CycleTransferStats, len(stats)) + for id, cycleStats := range stats { + cycleStats.ServerName = filterServerMapForViewer(c, cycleStats.ServerName, servers) + cycleStats.Transfer = filterServerMapForViewer(c, cycleStats.Transfer, servers) + cycleStats.NextUpdate = filterServerMapForViewer(c, cycleStats.NextUpdate, servers) + if len(cycleStats.ServerName) == 0 && len(cycleStats.Transfer) == 0 && len(cycleStats.NextUpdate) == 0 { + continue + } + filteredStats[id] = cycleStats + } + return filteredStats +} + +func filterServerMapForViewer[T any](c *gin.Context, values map[uint64]T, servers map[uint64]*model.Server) map[uint64]T { + if len(values) == 0 { + return values + } + filteredValues := make(map[uint64]T, len(values)) + for serverID, value := range values { + server, ok := servers[serverID] + if !ok || !userCanViewServer(c, server) { + continue + } + filteredValues[serverID] = value + } + return filteredValues +} + // List service // @Summary List service // @Security BearerAuth