diff --git a/service/singleton/alertsentinel.go b/service/singleton/alertsentinel.go index c31873c0..cb3cbe91 100644 --- a/service/singleton/alertsentinel.go +++ b/service/singleton/alertsentinel.go @@ -167,7 +167,7 @@ func checkStatus() { alertsPrevState[alert.ID][server.ID] = _RuleCheckFail message := fmt.Sprintf("[%s] %s(%s) %s", Localizer.T("Incident"), server.Name, IPDesensitize(server.GeoIP.IP.Join()), alert.Name) - go CronShared.SendTriggerTasks(alert.FailTriggerTasks, curServer.ID) + go CronShared.SendTriggerTasks(alert.FailTriggerTasks, curServer.ID, alert.UserID) go NotificationShared.SendNotification(alert.NotificationGroupID, message, NotificationMuteLabel.ServerIncident(server.ID, alert.ID), &curServer) // 清除恢复通知的静音缓存 NotificationShared.UnMuteNotification(alert.NotificationGroupID, NotificationMuteLabel.ServerIncidentResolved(server.ID, alert.ID)) @@ -177,7 +177,7 @@ func checkStatus() { if alertsPrevState[alert.ID][server.ID] == _RuleCheckFail { message := fmt.Sprintf("[%s] %s(%s) %s", Localizer.T("Resolved"), server.Name, IPDesensitize(server.GeoIP.IP.Join()), alert.Name) - go CronShared.SendTriggerTasks(alert.RecoverTriggerTasks, curServer.ID) + go CronShared.SendTriggerTasks(alert.RecoverTriggerTasks, curServer.ID, alert.UserID) go NotificationShared.SendNotification(alert.NotificationGroupID, message, NotificationMuteLabel.ServerIncidentResolved(server.ID, alert.ID), &curServer) // 清除失败通知的静音缓存 NotificationShared.UnMuteNotification(alert.NotificationGroupID, NotificationMuteLabel.ServerIncident(server.ID, alert.ID)) diff --git a/service/singleton/crontask.go b/service/singleton/crontask.go index 7928c3d4..ce73c081 100644 --- a/service/singleton/crontask.go +++ b/service/singleton/crontask.go @@ -110,11 +110,11 @@ func (c *CronClass) sortList() { c.sortedList = sortedList } -func (c *CronClass) SendTriggerTasks(taskIDs []uint64, triggerServer uint64) { +func (c *CronClass) SendTriggerTasks(taskIDs []uint64, triggerServer uint64, triggerOwner uint64) { c.listMu.RLock() var cronLists []*model.Cron for _, taskID := range taskIDs { - if c, ok := c.list[taskID]; ok { + if c, ok := c.list[taskID]; ok && cronCanBeTriggeredByOwner(c, triggerOwner) { cronLists = append(cronLists, c) } } @@ -126,6 +126,10 @@ func (c *CronClass) SendTriggerTasks(taskIDs []uint64, triggerServer uint64) { } } +func cronCanBeTriggeredByOwner(cr *model.Cron, triggerOwner uint64) bool { + return cr.UserID == triggerOwner || userIsAdmin(triggerOwner) +} + func ManualTrigger(cr *model.Cron) { CronTrigger(cr)() } @@ -141,6 +145,9 @@ func CronTrigger(cr *model.Cron, triggerServer ...uint64) func() { return } if s, ok := ServerShared.Get(triggerServer[0]); ok { + if !cronCanSendToServer(cr, s) { + return + } if s.TaskStream != nil { s.TaskStream.Send(&pb.Task{ Id: cr.ID, @@ -158,6 +165,9 @@ func CronTrigger(cr *model.Cron, triggerServer ...uint64) func() { } for _, s := range ServerShared.Range { + if !cronCanSendToServer(cr, s) { + continue + } if cr.Cover == model.CronCoverAll && crIgnoreMap[s.ID] { continue } @@ -179,3 +189,19 @@ func CronTrigger(cr *model.Cron, triggerServer ...uint64) func() { } } } + +func cronCanSendToServer(cr *model.Cron, server *model.Server) bool { + return cr.UserID == server.UserID || userIsAdmin(cr.UserID) +} + +func userIsAdmin(userID uint64) bool { + if userID == 0 { + return true + } + + UserLock.RLock() + defer UserLock.RUnlock() + + userInfo, ok := UserInfoMap[userID] + return ok && userInfo.Role.IsAdmin() +} diff --git a/service/singleton/security_regression_test.go b/service/singleton/security_regression_test.go new file mode 100644 index 00000000..f88b33f3 --- /dev/null +++ b/service/singleton/security_regression_test.go @@ -0,0 +1,137 @@ +package singleton + +import ( + "context" + "testing" + "time" + + "github.com/nezhahq/nezha/model" + pb "github.com/nezhahq/nezha/proto" + "google.golang.org/grpc/metadata" +) + +type capturedTaskStream struct { + tasks chan *pb.Task +} + +func newCapturedTaskStream() *capturedTaskStream { + return &capturedTaskStream{tasks: make(chan *pb.Task, 4)} +} + +func (s *capturedTaskStream) Send(task *pb.Task) error { + s.tasks <- task + return nil +} + +func (s *capturedTaskStream) Recv() (*pb.TaskResult, error) { return nil, context.Canceled } +func (s *capturedTaskStream) SetHeader(metadata.MD) error { return nil } +func (s *capturedTaskStream) SendHeader(metadata.MD) error { return nil } +func (s *capturedTaskStream) SetTrailer(metadata.MD) {} +func (s *capturedTaskStream) Context() context.Context { return context.Background() } +func (s *capturedTaskStream) SendMsg(any) error { return nil } +func (s *capturedTaskStream) RecvMsg(any) error { return context.Canceled } + +func replaceServerSharedForSecurityTest(t *testing.T, servers ...*model.Server) { + t.Helper() + + original := ServerShared + serverClass := &ServerClass{ + class: class[uint64, *model.Server]{ + list: make(map[uint64]*model.Server), + }, + uuidToID: make(map[string]uint64), + } + for _, server := range servers { + serverClass.list[server.ID] = server + } + ServerShared = serverClass + t.Cleanup(func() { ServerShared = original }) +} + +func replaceUserInfoMapForSecurityTest(t *testing.T, users map[uint64]model.UserInfo) { + t.Helper() + + UserLock.Lock() + original := UserInfoMap + UserInfoMap = users + UserLock.Unlock() + + t.Cleanup(func() { + UserLock.Lock() + UserInfoMap = original + UserLock.Unlock() + }) +} + +func TestCronTriggerSkipsServersOwnedByOtherUsers(t *testing.T) { + firstStream := newCapturedTaskStream() + secondStream := newCapturedTaskStream() + replaceServerSharedForSecurityTest(t, + &model.Server{Common: model.Common{ID: 1, UserID: 100}, Name: "member-server", TaskStream: firstStream}, + &model.Server{Common: model.Common{ID: 2, UserID: 200}, Name: "admin-server", TaskStream: secondStream}, + ) + + cronTask := &model.Cron{ + Common: model.Common{ID: 99, UserID: 100}, + Command: "id", + Cover: model.CronCoverAll, + Servers: []uint64{}, + } + + CronTrigger(cronTask)() + + assertTaskCommand(t, firstStream, "id") + assertNoTask(t, secondStream) +} + +func TestSendTriggerTasksSkipsCronOwnedByAnotherUser(t *testing.T) { + attackerStream := newCapturedTaskStream() + replaceServerSharedForSecurityTest(t, + &model.Server{Common: model.Common{ID: 7, UserID: 200}, Name: "attacker-server", TaskStream: attackerStream}, + ) + replaceUserInfoMapForSecurityTest(t, map[uint64]model.UserInfo{ + 1: {Role: model.RoleAdmin}, + 200: {Role: model.RoleMember}, + }) + + adminCron := &model.Cron{ + Common: model.Common{ID: 42, UserID: 1}, + Command: "admin-maintenance", + Cover: model.CronCoverAlertTrigger, + } + cronClass := &CronClass{ + class: class[uint64, *model.Cron]{ + list: map[uint64]*model.Cron{adminCron.ID: adminCron}, + }, + } + + cronClass.SendTriggerTasks([]uint64{adminCron.ID}, 7, 200) + + assertNoTask(t, attackerStream) +} + +func assertTaskCommand(t *testing.T, stream *capturedTaskStream, expectedCommand string) { + t.Helper() + + select { + case task := <-stream.tasks: + if task.GetType() != model.TaskTypeCommand { + t.Fatalf("expected command task type, got %v", task.GetType()) + } + if task.GetData() != expectedCommand { + t.Fatalf("expected command %q, got %q", expectedCommand, task.GetData()) + } + case <-time.After(time.Second): + t.Fatalf("expected command %q to be sent", expectedCommand) + } +} + +func assertNoTask(t *testing.T, stream *capturedTaskStream) { + t.Helper() + + select { + case task := <-stream.tasks: + t.Fatalf("expected no task to be sent, got command %q", task.GetData()) + case <-time.After(50 * time.Millisecond): + } +} diff --git a/service/singleton/servicesentinel.go b/service/singleton/servicesentinel.go index 9b7aef49..205f1d41 100644 --- a/service/singleton/servicesentinel.go +++ b/service/singleton/servicesentinel.go @@ -744,10 +744,10 @@ func notifyCheck(r *ReportData, m map[uint64]*model.Server, reporterServer := m[r.Reporter] if stateCode == StatusGood && lastStatus != stateCode { // 当前状态正常 前序状态非正常时 触发恢复任务 - go CronShared.SendTriggerTasks(ss.RecoverTriggerTasks, reporterServer.ID) + go CronShared.SendTriggerTasks(ss.RecoverTriggerTasks, reporterServer.ID, ss.UserID) } else if lastStatus == StatusGood && lastStatus != stateCode { // 前序状态正常 当前状态非正常时 触发失败任务 - go CronShared.SendTriggerTasks(ss.FailTriggerTasks, reporterServer.ID) + go CronShared.SendTriggerTasks(ss.FailTriggerTasks, reporterServer.ID, ss.UserID) } } }