diff --git a/tms/task-manager/internal/service/task_svc.go b/tms/task-manager/internal/service/task_svc.go index 732de2e7b54f278b1ebc3b95200ca622a404d02b..71ecad9e5e71a768f1fbedb35ca92bd08de808eb 100644 --- a/tms/task-manager/internal/service/task_svc.go +++ b/tms/task-manager/internal/service/task_svc.go @@ -192,11 +192,21 @@ func (s *TaskService) CreateOpsTask(ctx context.Context, req *pb.CreateOpsTaskRe } // 持久化 target 行。 + acceptedAgentIDs := make([]uint64, 0, len(agentIDs)) for _, agentID := range agentIDs { if err := s.store.AddTarget(taskID, agentID); err != nil { log.ErrorContextf(ctx, "add target failed, task_id: %s, agent_id: %d, err: %v", taskID, agentID, err) + continue } + acceptedAgentIDs = append(acceptedAgentIDs, agentID) + } + if len(acceptedAgentIDs) == 0 { + _ = s.store.UpdateOpsTaskStatus(taskID, "FAILED") + return &pb.CreateOpsTaskResponse{ + ErrorCode: pb.TaskErrorCode_ERR_INTERNAL, + ErrorMessage: "add target failed", + }, nil } s.metrics.OpsTasksCreated.Add(1) _ = s.store.UpdateOpsTaskStatus(taskID, "RUNNING") @@ -220,11 +230,11 @@ func (s *TaskService) CreateOpsTask(ctx context.Context, req *pb.CreateOpsTaskRe ErrorMessage: werr.Error(), }, nil } - go s.deployTargets(context.Background(), taskID, req.Action, wireBody, agentIDs) + go s.deployTargets(context.Background(), taskID, req.Action, wireBody, acceptedAgentIDs) return &pb.CreateOpsTaskResponse{ TaskId: taskID, - AcceptedTargetCount: int32(len(agentIDs)), + AcceptedTargetCount: int32(len(acceptedAgentIDs)), }, nil } diff --git a/tms/task-manager/internal/service/task_svc_addtarget_test.go b/tms/task-manager/internal/service/task_svc_addtarget_test.go new file mode 100644 index 0000000000000000000000000000000000000000..e36ab52d78c49c190deaba62e08ea3e2e73dd15e --- /dev/null +++ b/tms/task-manager/internal/service/task_svc_addtarget_test.go @@ -0,0 +1,43 @@ +// Copyright (C) 2024 OpenCloudOS +// License: GPL-3.0-or-later + +package service + +import ( + "context" + "testing" + + pb "gitee.com/OpenCloudOS/ocmanager/tms/proto/task" +) + +func TestCreateOpsTaskCountsOnlyAcceptedTargets(t *testing.T) { + svc := newIsolatedTaskService(t) + unsupportedSQLiteUint := uint64(1) << 63 + + resp, err := svc.CreateOpsTask(context.Background(), &pb.CreateOpsTaskRequest{ + Creator: "alice", + Action: pb.OpsTaskAction_ACTION_RESTART, + Restart: &pb.RestartPayload{CountdownSeconds: 1}, + Targets: &pb.TargetSelector{AgentIds: []uint64{1, unsupportedSQLiteUint}}, + }) + if err != nil { + t.Fatalf("CreateOpsTask returned error: %v", err) + } + if resp.GetErrorCode() != pb.TaskErrorCode_ERR_OK { + t.Fatalf("CreateOpsTask error_code = %v, message = %q", resp.GetErrorCode(), resp.GetErrorMessage()) + } + if resp.GetAcceptedTargetCount() != 1 { + t.Fatalf("AcceptedTargetCount = %d, want 1", resp.GetAcceptedTargetCount()) + } + + targets, err := svc.store.ListTargets(resp.GetTaskId()) + if err != nil { + t.Fatalf("ListTargets: %v", err) + } + if len(targets) != 1 { + t.Fatalf("stored targets = %d, want 1", len(targets)) + } + if targets[0].AgentID != 1 { + t.Fatalf("stored agent_id = %d, want 1", targets[0].AgentID) + } +}