Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 1 addition & 2 deletions kubernetes/internal/scheduler/default_scheduler.go
Original file line number Diff line number Diff line change
Expand Up @@ -176,7 +176,6 @@ type defaultTaskScheduler struct {
taskNodeByNameIndex map[string]*taskNode

maxConcurrency int
once sync.Once

taskStatusCollector taskStatusCollector
taskClientCreator taskClientCreator
Expand Down Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion kubernetes/internal/scheduler/default_scheduler_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
22 changes: 15 additions & 7 deletions kubernetes/internal/scheduler/recovery.go
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@ package scheduler

import (
"context"
"fmt"

"github.com/go-logr/logr"
v1 "k8s.io/api/core/v1"
Expand All @@ -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 {
Expand All @@ -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]
Expand Down
113 changes: 109 additions & 4 deletions kubernetes/internal/scheduler/recovery_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -15,8 +15,8 @@
package scheduler

import (
"errors"
"reflect"
"sync"
"testing"
"time"

Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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
}(),
},
Expand Down Expand Up @@ -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
}(),
},
Expand Down Expand Up @@ -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)
}
}
24 changes: 19 additions & 5 deletions kubernetes/internal/scheduler/status_collector.go
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,8 @@ package scheduler

import (
"context"
"errors"
"fmt"
"sync"

"github.com/go-logr/logr"
Expand All @@ -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
Expand All @@ -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{}{}
Expand All @@ -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...)
}
5 changes: 3 additions & 2 deletions kubernetes/internal/scheduler/status_collector_mock.go

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

89 changes: 89 additions & 0 deletions kubernetes/internal/scheduler/status_collector_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
}
Loading