fix(security): restrict service monitors to probe tasks

This commit is contained in:
naiba
2026-08-15 05:00:28 +00:00
parent 42d9e4c8c3
commit 38824dbc11
16 changed files with 338 additions and 20 deletions
+13 -2
View File
@@ -14,6 +14,17 @@ func DispatchTask(serviceSentinelDispatchBus <-chan *model.Service) {
if task == nil {
continue
}
if err := model.ValidateServiceMonitorType(uint64(task.Type)); err != nil {
// Defense in depth for stale database rows and future internal callers:
// Service.Type shares its integer namespace with command/config tasks.
log.Printf("NEZHA>> DispatchTask rejected service %d: %v", task.ID, err)
continue
}
probe := task.PB()
if probe == nil {
log.Printf("NEZHA>> DispatchTask rejected service %d: invalid probe", task.ID)
continue
}
switch task.Cover {
case model.ServiceCoverIgnoreAll:
@@ -29,7 +40,7 @@ func DispatchTask(serviceSentinelDispatchBus <-chan *model.Service) {
if !canSendTaskToServer(task, server) {
continue
}
if err := server.SendTask(task.PB()); err != nil && !errors.Is(err, model.ErrTaskStreamOffline) {
if err := server.SendTask(probe); err != nil && !errors.Is(err, model.ErrTaskStreamOffline) {
log.Printf("NEZHA>> DispatchTask send error (server=%d): %v", id, err)
}
}
@@ -41,7 +52,7 @@ func DispatchTask(serviceSentinelDispatchBus <-chan *model.Service) {
if !canSendTaskToServer(task, server) {
continue
}
if err := server.SendTask(task.PB()); err != nil && !errors.Is(err, model.ErrTaskStreamOffline) {
if err := server.SendTask(probe); err != nil && !errors.Is(err, model.ErrTaskStreamOffline) {
log.Printf("NEZHA>> DispatchTask send error (server=%d): %v", id, err)
}
}
@@ -0,0 +1,59 @@
package rpc
import (
"testing"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/model"
"github.com/nezhahq/nezha/service/singleton"
)
func TestDispatchTaskSendsOnlyProbeTypes(t *testing.T) {
originalServerShared := singleton.ServerShared
originalUserInfo := singleton.UserInfoMap
t.Cleanup(func() {
singleton.ServerShared = originalServerShared
singleton.UserLock.Lock()
singleton.UserInfoMap = originalUserInfo
singleton.UserLock.Unlock()
})
server := &model.Server{Common: model.Common{ID: 1, UserID: 100}}
stream := &serveNATTaskStream{}
server.SetTaskStream(stream)
serverShared := singleton.NewEmptyServerClassForTest()
serverShared.InsertForTest(server)
singleton.ServerShared = serverShared
singleton.UserLock.Lock()
singleton.UserInfoMap = map[uint64]model.UserInfo{100: {Role: model.RoleMember}}
singleton.UserLock.Unlock()
bus := make(chan *model.Service, 8)
done := make(chan struct{})
go func() {
DispatchTask(bus)
close(done)
}()
for _, taskType := range []uint8{model.TaskTypeCommand, model.TaskTypeApplyConfig, model.TaskTypeExec, 255} {
bus <- &model.Service{
Common: model.Common{ID: uint64(taskType), UserID: 100},
Type: taskType,
Cover: model.ServiceCoverIgnoreAll,
SkipServers: map[uint64]bool{1: true},
}
}
bus <- &model.Service{
Common: model.Common{ID: 1000, UserID: 100},
Type: model.TaskTypeTCPPing,
Target: "example.invalid:443",
Cover: model.ServiceCoverIgnoreAll,
SkipServers: map[uint64]bool{1: true},
}
close(bus)
<-done
require.Len(t, stream.sent, 1)
require.Equal(t, uint64(model.TaskTypeTCPPing), stream.sent[0].GetType())
require.Equal(t, uint64(1000), stream.sent[0].GetId())
}