mirror of
https://github.com/Buriburizaem0n/nezha_domains.git
synced 2026-09-19 17:50:12 +00:00
fix(controller): filter service transfer stats by viewer
Co-authored-by: naiba/CloudCode <hi+cloudcode@nai.ba>
This commit is contained in:
@@ -6,6 +6,7 @@ import (
|
|||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/nezhahq/nezha/model"
|
"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 {
|
func decodeIDs[T ~uint64](t *testing.T, body []byte) []T {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
var resp struct {
|
var resp struct {
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
package controller
|
package controller
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"fmt"
|
||||||
"maps"
|
"maps"
|
||||||
"slices"
|
"slices"
|
||||||
"strconv"
|
"strconv"
|
||||||
@@ -26,14 +27,14 @@ import (
|
|||||||
// @Success 200 {object} model.CommonResponse[model.ServiceResponse]
|
// @Success 200 {object} model.CommonResponse[model.ServiceResponse]
|
||||||
// @Router /service [get]
|
// @Router /service [get]
|
||||||
func showService(c *gin.Context) (*model.ServiceResponse, error) {
|
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()
|
singleton.AlertsLock.RLock()
|
||||||
defer singleton.AlertsLock.RUnlock()
|
defer singleton.AlertsLock.RUnlock()
|
||||||
stats := singleton.ServiceSentinelShared.CopyStats()
|
stats := singleton.ServiceSentinelShared.CopyStats()
|
||||||
var cycleTransferStats map[uint64]model.CycleTransferStats
|
var cycleTransferStats map[uint64]model.CycleTransferStats
|
||||||
copier.Copy(&cycleTransferStats, singleton.AlertsCycleTransferStatsStore)
|
copier.Copy(&cycleTransferStats, singleton.AlertsCycleTransferStatsStore)
|
||||||
return []any{
|
return []any{
|
||||||
stats, cycleTransferStats,
|
stats, filterCycleTransferStatsForViewer(c, cycleTransferStats),
|
||||||
}, nil
|
}, nil
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -46,6 +47,51 @@ func showService(c *gin.Context) (*model.ServiceResponse, error) {
|
|||||||
}, nil
|
}, 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
|
// List service
|
||||||
// @Summary List service
|
// @Summary List service
|
||||||
// @Security BearerAuth
|
// @Security BearerAuth
|
||||||
|
|||||||
Reference in New Issue
Block a user