diff --git a/kubernetes/internal/scheduler/default_scheduler.go b/kubernetes/internal/scheduler/default_scheduler.go index 7e1f422ec..00f999956 100644 --- a/kubernetes/internal/scheduler/default_scheduler.go +++ b/kubernetes/internal/scheduler/default_scheduler.go @@ -176,7 +176,6 @@ type defaultTaskScheduler struct { taskNodeByNameIndex map[string]*taskNode maxConcurrency int - once sync.Once taskStatusCollector taskStatusCollector taskClientCreator taskClientCreator @@ -298,7 +297,7 @@ func (sch *defaultTaskScheduler) collectTaskStatus(taskNodes []*taskNode) { if len(ips) == 0 { return } - tasks := sch.taskStatusCollector.Collect(context.Background(), ips) + tasks, _ := sch.taskStatusCollector.Collect(context.Background(), ips) for _, tNode := range taskNodes { task, ok := tasks[tNode.IP] tNode.Status = task diff --git a/kubernetes/internal/scheduler/default_scheduler_test.go b/kubernetes/internal/scheduler/default_scheduler_test.go index 274e0bc1e..8ac8ad29d 100644 --- a/kubernetes/internal/scheduler/default_scheduler_test.go +++ b/kubernetes/internal/scheduler/default_scheduler_test.go @@ -704,7 +704,7 @@ func Test_collectTaskStatus(t *testing.T) { // Create mock task status collector mockCollector := NewMocktaskStatusCollector(ctl) if len(tt.expectedCollectIPs) > 0 { - mockCollector.EXPECT().Collect(gomock.Any(), tt.expectedCollectIPs).Return(tt.mockReturnTasks).Times(1) + mockCollector.EXPECT().Collect(gomock.Any(), tt.expectedCollectIPs).Return(tt.mockReturnTasks, nil).Times(1) } // Create scheduler with mock collector diff --git a/kubernetes/internal/scheduler/recovery.go b/kubernetes/internal/scheduler/recovery.go index a1491011a..250cb4b52 100644 --- a/kubernetes/internal/scheduler/recovery.go +++ b/kubernetes/internal/scheduler/recovery.go @@ -16,6 +16,7 @@ package scheduler import ( "context" + "fmt" "github.com/go-logr/logr" v1 "k8s.io/api/core/v1" @@ -26,12 +27,11 @@ import ( // recover reconstructs the task scheduler state from existing pods and their endpoints // This function is used to restore the scheduler state after a restart func (sch *defaultTaskScheduler) recover() error { - var err error - sch.once.Do(func() { - sch.recoverTaskNodesStatus() - sch.logger.Info("task scheduler recovered", "scheduler", sch.name, "task_nodes", len(sch.taskNodes), "all_pods", len(sch.allPods)) - }) - return err + if err := sch.recoverTaskNodesStatus(); err != nil { + return err + } + sch.logger.Info("task scheduler recovered", "scheduler", sch.name, "task_nodes", len(sch.taskNodes), "all_pods", len(sch.allPods)) + return nil } func (sch *defaultTaskScheduler) recoverTaskNodesStatus() error { @@ -52,7 +52,15 @@ func (sch *defaultTaskScheduler) recoverTaskNodesStatus() error { // the recovery may complete after the agent has already finished stopping the task and returned an empty task list. // This could cause the scheduler to be unable to determine whether the task was never executed or has already completed. // It might lead to duplicate execution, but it ensures at-least-once delivery semantics. - tasks := sch.taskStatusCollector.Collect(context.Background(), ips) + tasks, err := sch.taskStatusCollector.Collect(context.Background(), ips) + if err != nil { + return fmt.Errorf("collect task status during recovery: %w", err) + } + for _, ip := range ips { + if _, ok := tasks[ip]; !ok { + return fmt.Errorf("collect task status during recovery: missing result for pod IP %s", ip) + } + } for i := range ips { ip := ips[i] pod := pods[i] diff --git a/kubernetes/internal/scheduler/recovery_test.go b/kubernetes/internal/scheduler/recovery_test.go index 4a37b1840..77a939a84 100644 --- a/kubernetes/internal/scheduler/recovery_test.go +++ b/kubernetes/internal/scheduler/recovery_test.go @@ -15,8 +15,8 @@ package scheduler import ( + "errors" "reflect" - "sync" "testing" "time" @@ -173,7 +173,6 @@ func Test_defaultTaskScheduler_recoverTaskNodesStatus(t *testing.T) { taskNodes []*taskNode taskNodeByNameIndex map[string]*taskNode maxConcurrency int - once sync.Once taskStatusCollector taskStatusCollector } tests := []struct { @@ -224,7 +223,7 @@ func Test_defaultTaskScheduler_recoverTaskNodesStatus(t *testing.T) { }, taskStatusCollector: func() taskStatusCollector { mock := NewMocktaskStatusCollector(ctl) - mock.EXPECT().Collect(gomock.Any(), []string{"1.2.3.4"}).Return(map[string]*api.Task{"1.2.3.4": nil}).Times(1) + mock.EXPECT().Collect(gomock.Any(), []string{"1.2.3.4"}).Return(map[string]*api.Task{"1.2.3.4": nil}, nil).Times(1) return mock }(), }, @@ -253,7 +252,7 @@ func Test_defaultTaskScheduler_recoverTaskNodesStatus(t *testing.T) { }, taskStatusCollector: func() taskStatusCollector { mock := NewMocktaskStatusCollector(ctl) - mock.EXPECT().Collect(gomock.Any(), []string{"1.2.3.4"}).Return(map[string]*api.Task{"1.2.3.4": testTask}).Times(1) + mock.EXPECT().Collect(gomock.Any(), []string{"1.2.3.4"}).Return(map[string]*api.Task{"1.2.3.4": testTask}, nil).Times(1) return mock }(), }, @@ -285,3 +284,109 @@ func Test_defaultTaskScheduler_recoverTaskNodesStatus(t *testing.T) { }) } } + +func Test_defaultTaskScheduler_recoverTaskNodesStatusIsAtomicOnCollectionError(t *testing.T) { + ctl := gomock.NewController(t) + defer ctl.Finish() + + firstNode := &taskNode{ + ObjectMeta: v1.ObjectMeta{Name: "task-1"}, + Status: &api.Task{Name: "task-1"}, + IP: "192.0.2.1", + PodName: "pod-1", + tState: RunningTaskState, + } + secondNode := &taskNode{ + ObjectMeta: v1.ObjectMeta{Name: "task-2"}, + Status: &api.Task{Name: "task-2"}, + IP: "192.0.2.2", + PodName: "pod-2", + tState: RunningTaskState, + } + firstNodeBefore := *firstNode + secondNodeBefore := *secondNode + queryErr := errors.New("executor unavailable") + + collector := NewMocktaskStatusCollector(ctl) + collector.EXPECT().Collect(gomock.Any(), []string{"192.0.2.1", "192.0.2.2"}).Return(map[string]*api.Task{ + "192.0.2.1": {Name: "task-1"}, + }, queryErr).Times(1) + + sch := &defaultTaskScheduler{ + allPods: []*corev1.Pod{ + {ObjectMeta: v1.ObjectMeta{Name: "pod-1"}, Status: corev1.PodStatus{PodIP: "192.0.2.1"}}, + {ObjectMeta: v1.ObjectMeta{Name: "pod-2"}, Status: corev1.PodStatus{PodIP: "192.0.2.2"}}, + }, + taskNodes: []*taskNode{firstNode, secondNode}, + taskNodeByNameIndex: map[string]*taskNode{"task-1": firstNode, "task-2": secondNode}, + taskStatusCollector: collector, + logger: testLogger, + } + + err := sch.recoverTaskNodesStatus() + if !errors.Is(err, queryErr) { + t.Fatalf("recoverTaskNodesStatus() error = %v, want wrapped collection error", err) + } + if !reflect.DeepEqual(firstNodeBefore, *firstNode) { + t.Fatalf("first task node changed after failed recovery: before=%+v after=%+v", firstNodeBefore, *firstNode) + } + if !reflect.DeepEqual(secondNodeBefore, *secondNode) { + t.Fatalf("second task node changed after failed recovery: before=%+v after=%+v", secondNodeBefore, *secondNode) + } +} + +func Test_defaultTaskScheduler_recoverPropagatesCollectionError(t *testing.T) { + ctl := gomock.NewController(t) + defer ctl.Finish() + + queryErr := errors.New("executor unavailable") + collector := NewMocktaskStatusCollector(ctl) + collector.EXPECT().Collect(gomock.Any(), []string{"192.0.2.1"}).Return(nil, queryErr).Times(1) + + sch := &defaultTaskScheduler{ + allPods: []*corev1.Pod{{ + ObjectMeta: v1.ObjectMeta{Name: "pod-1"}, + Status: corev1.PodStatus{PodIP: "192.0.2.1"}, + }}, + taskStatusCollector: collector, + logger: testLogger, + } + if err := sch.recover(); !errors.Is(err, queryErr) { + t.Fatalf("recover() error = %v, want wrapped collection error", err) + } +} + +func Test_defaultTaskScheduler_recoverTaskNodesStatusRejectsIncompleteCollection(t *testing.T) { + ctl := gomock.NewController(t) + defer ctl.Finish() + + firstNode := &taskNode{ObjectMeta: v1.ObjectMeta{Name: "task-1"}} + secondNode := &taskNode{ObjectMeta: v1.ObjectMeta{Name: "task-2"}} + firstNodeBefore := *firstNode + secondNodeBefore := *secondNode + collector := NewMocktaskStatusCollector(ctl) + collector.EXPECT().Collect(gomock.Any(), []string{"192.0.2.1", "192.0.2.2"}).Return(map[string]*api.Task{ + "192.0.2.1": {Name: "task-1"}, + }, nil).Times(1) + + sch := &defaultTaskScheduler{ + allPods: []*corev1.Pod{ + {ObjectMeta: v1.ObjectMeta{Name: "pod-1"}, Status: corev1.PodStatus{PodIP: "192.0.2.1"}}, + {ObjectMeta: v1.ObjectMeta{Name: "pod-2"}, Status: corev1.PodStatus{PodIP: "192.0.2.2"}}, + }, + taskNodes: []*taskNode{firstNode, secondNode}, + taskNodeByNameIndex: map[string]*taskNode{"task-1": firstNode, "task-2": secondNode}, + taskStatusCollector: collector, + logger: testLogger, + } + + if err := sch.recoverTaskNodesStatus(); err == nil { + t.Fatal("recoverTaskNodesStatus() error = nil, want incomplete collection error") + } + if !reflect.DeepEqual(firstNodeBefore, *firstNode) { + t.Fatalf("first task node changed after incomplete recovery: before=%+v after=%+v", firstNodeBefore, *firstNode) + } + if !reflect.DeepEqual(secondNodeBefore, *secondNode) { + t.Fatalf("second task node changed after incomplete recovery: before=%+v after=%+v", secondNodeBefore, *secondNode) + } +} diff --git a/kubernetes/internal/scheduler/status_collector.go b/kubernetes/internal/scheduler/status_collector.go index adc8db364..670a71db6 100644 --- a/kubernetes/internal/scheduler/status_collector.go +++ b/kubernetes/internal/scheduler/status_collector.go @@ -16,6 +16,8 @@ package scheduler import ( "context" + "errors" + "fmt" "sync" "github.com/go-logr/logr" @@ -30,9 +32,8 @@ func newTaskStatusCollector(creator taskClientCreator, logger logr.Logger) taskS return &defaultTaskStatusCollector{creator: creator, logger: logger} } -// TODO error type taskStatusCollector interface { - Collect(ctx context.Context, ipList []string) map[string]*api.Task /*ip<->task*/ + Collect(ctx context.Context, ipList []string) (map[string]*api.Task, error) /*ip<->task*/ } // TODO maybe cache @@ -41,11 +42,12 @@ type defaultTaskStatusCollector struct { logger logr.Logger } -func (s *defaultTaskStatusCollector) Collect(ctx context.Context, ipList []string) map[string]*api.Task { +func (s *defaultTaskStatusCollector) Collect(ctx context.Context, ipList []string) (map[string]*api.Task, error) { semaphore := make(chan struct{}, len(ipList)) var wg sync.WaitGroup var mu sync.Mutex ret := make(map[string]*api.Task, len(ipList)) + var errs []error for idx := range ipList { ip := ipList[idx] semaphore <- struct{}{} @@ -61,14 +63,26 @@ func (s *defaultTaskStatusCollector) Collect(ctx context.Context, ipList []strin task, err := client.Get(ctx) if err != nil { s.logger.Error(err, "failed to GetTask", "ip", ip) + mu.Lock() + errs = append(errs, fmt.Errorf("get task status for pod IP %s: %w", ip, err)) + mu.Unlock() } else if task != nil { mu.Lock() ret[ip] = task mu.Unlock() + } else { + // Keep an explicit nil entry so callers can distinguish a + // confirmed empty task list from a failed query. + mu.Lock() + ret[ip] = nil + mu.Unlock() } }(ip) } wg.Wait() - s.logger.Info("Collect task status", "result", utils.DumpJSON(ret)) - return ret + verboseLog := s.logger.V(3) + if verboseLog.Enabled() { + verboseLog.Info("Collect task status", "result", utils.DumpJSON(ret)) + } + return ret, errors.Join(errs...) } diff --git a/kubernetes/internal/scheduler/status_collector_mock.go b/kubernetes/internal/scheduler/status_collector_mock.go index d8aa1e2a3..5cbf5c301 100644 --- a/kubernetes/internal/scheduler/status_collector_mock.go +++ b/kubernetes/internal/scheduler/status_collector_mock.go @@ -37,11 +37,12 @@ func (m *MocktaskStatusCollector) EXPECT() *MocktaskStatusCollectorMockRecorder } // Collect mocks base method. -func (m *MocktaskStatusCollector) Collect(ctx context.Context, ipList []string) map[string]*api.Task { +func (m *MocktaskStatusCollector) Collect(ctx context.Context, ipList []string) (map[string]*api.Task, error) { m.ctrl.T.Helper() ret := m.ctrl.Call(m, "Collect", ctx, ipList) ret0, _ := ret[0].(map[string]*api.Task) - return ret0 + ret1, _ := ret[1].(error) + return ret0, ret1 } // Collect indicates an expected call of Collect. diff --git a/kubernetes/internal/scheduler/status_collector_test.go b/kubernetes/internal/scheduler/status_collector_test.go new file mode 100644 index 000000000..c0e6dd47d --- /dev/null +++ b/kubernetes/internal/scheduler/status_collector_test.go @@ -0,0 +1,89 @@ +// Copyright 2025 Alibaba Group Holding Ltd. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package scheduler + +import ( + "context" + "errors" + "strings" + "testing" + + api "github.com/alibaba/OpenSandbox/sandbox-k8s/pkg/task-executor" +) + +type testTaskStatusClient struct { + task *api.Task + err error +} + +func (c *testTaskStatusClient) Get(context.Context) (*api.Task, error) { + return c.task, c.err +} + +func (c *testTaskStatusClient) Set(context.Context, *api.Task) (*api.Task, error) { + return nil, nil +} + +func TestDefaultTaskStatusCollectorCollectDistinguishesEmptyAndError(t *testing.T) { + queryErr := errors.New("executor unavailable") + activeTask := &api.Task{Name: "task-1"} + collector := &defaultTaskStatusCollector{ + creator: func(ip string) taskClient { + switch ip { + case "10.0.0.1": + return &testTaskStatusClient{} + case "10.0.0.2": + return &testTaskStatusClient{task: activeTask} + default: + return &testTaskStatusClient{err: queryErr} + } + }, + logger: testLogger, + } + + tasks, err := collector.Collect(context.Background(), []string{"10.0.0.1", "10.0.0.2", "10.0.0.3"}) + if !errors.Is(err, queryErr) { + t.Fatalf("Collect() error = %v, want wrapped query error", err) + } + if !strings.Contains(err.Error(), "10.0.0.3") { + t.Fatalf("Collect() error = %v, want IP context", err) + } + if task, ok := tasks["10.0.0.1"]; !ok || task != nil { + t.Fatalf("Collect() empty result = (%v, %t), want explicit nil entry", task, ok) + } + if task, ok := tasks["10.0.0.2"]; !ok || task != activeTask { + t.Fatalf("Collect() active result = (%v, %t), want active task", task, ok) + } + if _, ok := tasks["10.0.0.3"]; ok { + t.Fatal("Collect() should omit the result for a failed query") + } +} + +func TestDefaultTaskStatusCollectorCollectWithoutErrors(t *testing.T) { + collector := &defaultTaskStatusCollector{ + creator: func(string) taskClient { + return &testTaskStatusClient{} + }, + logger: testLogger, + } + + tasks, err := collector.Collect(context.Background(), []string{"10.0.0.1"}) + if err != nil { + t.Fatalf("Collect() error = %v, want nil", err) + } + if task, ok := tasks["10.0.0.1"]; !ok || task != nil { + t.Fatalf("Collect() empty result = (%v, %t), want explicit nil entry", task, ok) + } +}