Skip to content
This repository was archived by the owner on Oct 9, 2023. It is now read-only.

Commit 5ac7181

Browse files
author
Nick Müller
committed
Refactored listing of node and task executions to shared util
Allows for re-use by cache manager Signed-off-by: Nick Müller <nmueller@blackshark.ai>
1 parent 6d5549c commit 5ac7181

3 files changed

Lines changed: 169 additions & 125 deletions

File tree

pkg/manager/impl/node_execution_manager.go

Lines changed: 20 additions & 88 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,6 @@ package impl
22

33
import (
44
"context"
5-
"strconv"
65

76
cloudeventInterfaces "github.com/flyteorg/flyteadmin/pkg/async/cloudevent/interfaces"
87

@@ -17,7 +16,6 @@ import (
1716

1817
"github.com/flyteorg/flytestdlib/contextutils"
1918

20-
"github.com/flyteorg/flyteadmin/pkg/manager/impl/shared"
2119
"github.com/flyteorg/flytestdlib/promutils"
2220
"github.com/prometheus/client_golang/prometheus"
2321

@@ -74,11 +72,6 @@ const (
7472
alreadyInTerminalStatus
7573
)
7674

77-
var isParent = common.NewMapFilter(map[string]interface{}{
78-
shared.ParentTaskExecutionID: nil,
79-
shared.ParentID: nil,
80-
})
81-
8275
func getNodeExecutionContext(ctx context.Context, identifier *core.NodeExecutionIdentifier) context.Context {
8376
ctx = contextutils.WithProjectDomain(ctx, identifier.ExecutionId.Project, identifier.ExecutionId.Domain)
8477
ctx = contextutils.WithExecutionID(ctx, identifier.ExecutionId.Name)
@@ -369,48 +362,23 @@ func (m *NodeExecutionManager) GetNodeExecution(
369362
return nodeExecution, nil
370363
}
371364

372-
func (m *NodeExecutionManager) listNodeExecutions(
373-
ctx context.Context, identifierFilters []common.InlineFilter,
374-
requestFilters string, limit uint32, requestToken string, sortBy *admin.Sort, mapFilters []common.MapFilter) (
375-
*admin.NodeExecutionList, error) {
376-
377-
filters, err := util.AddRequestFilters(requestFilters, common.NodeExecution, identifierFilters)
378-
if err != nil {
365+
func (m *NodeExecutionManager) ListNodeExecutions(
366+
ctx context.Context, request admin.NodeExecutionListRequest) (*admin.NodeExecutionList, error) {
367+
// Check required fields
368+
if err := validation.ValidateNodeExecutionListRequest(request); err != nil {
379369
return nil, err
380370
}
381-
var sortParameter common.SortParameter
382-
if sortBy != nil {
383-
sortParameter, err = common.NewSortParameter(*sortBy)
384-
if err != nil {
385-
return nil, err
386-
}
387-
}
388-
offset, err := validation.ValidateToken(requestToken)
389-
if err != nil {
390-
return nil, errors.NewFlyteAdminErrorf(codes.InvalidArgument,
391-
"invalid pagination token %s for ListNodeExecutions", requestToken)
392-
}
393-
listInput := repoInterfaces.ListResourceInput{
394-
Limit: int(limit),
395-
Offset: offset,
396-
InlineFilters: filters,
397-
SortParameter: sortParameter,
398-
}
371+
ctx = getExecutionContext(ctx, request.WorkflowExecutionId)
399372

400-
listInput.MapFilters = mapFilters
401-
output, err := m.db.NodeExecutionRepo().List(ctx, listInput)
373+
nodeExecutions, token, err := util.ListNodeExecutionsForWorkflow(ctx, m.db, request.WorkflowExecutionId,
374+
request.UniqueParentId, request.Filters, request.Limit, request.Token, request.SortBy)
402375
if err != nil {
403-
logger.Debugf(ctx, "Failed to list node executions for request with err %v", err)
404376
return nil, err
405377
}
406378

407-
var token string
408-
if len(output.NodeExecutions) == int(limit) {
409-
token = strconv.Itoa(offset + len(output.NodeExecutions))
410-
}
411-
nodeExecutionList, err := m.transformNodeExecutionModelList(ctx, output.NodeExecutions)
379+
nodeExecutionList, err := m.transformNodeExecutionModelList(ctx, nodeExecutions)
412380
if err != nil {
413-
logger.Debugf(ctx, "failed to transform node execution models for request with err: %v", err)
381+
logger.Debugf(ctx, "failed to transform node execution models for request [%+v] with err: %v", request, err)
414382
return nil, err
415383
}
416384

@@ -420,42 +388,6 @@ func (m *NodeExecutionManager) listNodeExecutions(
420388
}, nil
421389
}
422390

423-
func (m *NodeExecutionManager) ListNodeExecutions(
424-
ctx context.Context, request admin.NodeExecutionListRequest) (*admin.NodeExecutionList, error) {
425-
// Check required fields
426-
if err := validation.ValidateNodeExecutionListRequest(request); err != nil {
427-
return nil, err
428-
}
429-
ctx = getExecutionContext(ctx, request.WorkflowExecutionId)
430-
431-
identifierFilters, err := util.GetWorkflowExecutionIdentifierFilters(ctx, *request.WorkflowExecutionId)
432-
if err != nil {
433-
return nil, err
434-
}
435-
var mapFilters []common.MapFilter
436-
if request.UniqueParentId != "" {
437-
parentNodeExecution, err := util.GetNodeExecutionModel(ctx, m.db, &core.NodeExecutionIdentifier{
438-
ExecutionId: request.WorkflowExecutionId,
439-
NodeId: request.UniqueParentId,
440-
})
441-
if err != nil {
442-
return nil, err
443-
}
444-
parentIDFilter, err := common.NewSingleValueFilter(
445-
common.NodeExecution, common.Equal, shared.ParentID, parentNodeExecution.ID)
446-
if err != nil {
447-
return nil, err
448-
}
449-
identifierFilters = append(identifierFilters, parentIDFilter)
450-
} else {
451-
mapFilters = []common.MapFilter{
452-
isParent,
453-
}
454-
}
455-
return m.listNodeExecutions(
456-
ctx, identifierFilters, request.Filters, request.Limit, request.Token, request.SortBy, mapFilters)
457-
}
458-
459391
// Filters on node executions matching the execution parameters (execution project, domain, and name) as well as the
460392
// parent task execution id corresponding to the task execution identified in the request params.
461393
func (m *NodeExecutionManager) ListNodeExecutionsForTask(
@@ -465,23 +397,23 @@ func (m *NodeExecutionManager) ListNodeExecutionsForTask(
465397
return nil, err
466398
}
467399
ctx = getTaskExecutionContext(ctx, request.TaskExecutionId)
468-
identifierFilters, err := util.GetWorkflowExecutionIdentifierFilters(
469-
ctx, *request.TaskExecutionId.NodeExecutionId.ExecutionId)
470-
if err != nil {
471-
return nil, err
472-
}
473-
parentTaskExecutionModel, err := util.GetTaskExecutionModel(ctx, m.db, request.TaskExecutionId)
400+
401+
nodeExecutions, token, err := util.ListNodeExecutionsForTask(ctx, m.db, request.TaskExecutionId,
402+
request.TaskExecutionId.NodeExecutionId.ExecutionId, request.Filters, request.Limit, request.Token, request.SortBy)
474403
if err != nil {
475404
return nil, err
476405
}
477-
nodeIDFilter, err := common.NewSingleValueFilter(
478-
common.NodeExecution, common.Equal, shared.ParentTaskExecutionID, parentTaskExecutionModel.ID)
406+
407+
nodeExecutionList, err := m.transformNodeExecutionModelList(ctx, nodeExecutions)
479408
if err != nil {
409+
logger.Debugf(ctx, "failed to transform node execution models for request [%+v] with err: %v", request, err)
480410
return nil, err
481411
}
482-
identifierFilters = append(identifierFilters, nodeIDFilter)
483-
return m.listNodeExecutions(
484-
ctx, identifierFilters, request.Filters, request.Limit, request.Token, request.SortBy, nil)
412+
413+
return &admin.NodeExecutionList{
414+
NodeExecutions: nodeExecutionList,
415+
Token: token,
416+
}, nil
485417
}
486418

487419
func (m *NodeExecutionManager) GetNodeExecutionData(

pkg/manager/impl/task_execution_manager.go

Lines changed: 4 additions & 37 deletions
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,6 @@ package impl
33
import (
44
"context"
55
"fmt"
6-
"strconv"
76

87
cloudeventInterfaces "github.com/flyteorg/flyteadmin/pkg/async/cloudevent/interfaces"
98

@@ -249,50 +248,18 @@ func (m *TaskExecutionManager) ListTaskExecutions(
249248
}
250249
ctx = getNodeExecutionContext(ctx, request.NodeExecutionId)
251250

252-
identifierFilters, err := util.GetNodeExecutionIdentifierFilters(ctx, *request.NodeExecutionId)
251+
taskExecutions, token, err := util.ListTaskExecutions(ctx, m.db, request.NodeExecutionId, request.Filters,
252+
request.Limit, request.Token, request.SortBy)
253253
if err != nil {
254254
return nil, err
255255
}
256256

257-
filters, err := util.AddRequestFilters(request.Filters, common.TaskExecution, identifierFilters)
258-
if err != nil {
259-
return nil, err
260-
}
261-
var sortParameter common.SortParameter
262-
if request.SortBy != nil {
263-
sortParameter, err = common.NewSortParameter(*request.SortBy)
264-
if err != nil {
265-
return nil, err
266-
}
267-
}
268-
269-
offset, err := validation.ValidateToken(request.Token)
270-
if err != nil {
271-
return nil, errors.NewFlyteAdminErrorf(codes.InvalidArgument,
272-
"invalid pagination token %s for ListTaskExecutions", request.Token)
273-
}
274-
275-
output, err := m.db.TaskExecutionRepo().List(ctx, repoInterfaces.ListResourceInput{
276-
InlineFilters: filters,
277-
Offset: offset,
278-
Limit: int(request.Limit),
279-
SortParameter: sortParameter,
280-
})
281-
if err != nil {
282-
logger.Debugf(ctx, "Failed to list task executions with request [%+v] with err %v",
283-
request, err)
284-
return nil, err
285-
}
286-
287-
taskExecutionList, err := transformers.FromTaskExecutionModels(output.TaskExecutions)
257+
taskExecutionList, err := transformers.FromTaskExecutionModels(taskExecutions)
288258
if err != nil {
289259
logger.Debugf(ctx, "failed to transform task execution models for request [%+v] with err: %v", request, err)
290260
return nil, err
291261
}
292-
var token string
293-
if len(taskExecutionList) == int(request.Limit) {
294-
token = strconv.Itoa(offset + len(taskExecutionList))
295-
}
262+
296263
return &admin.TaskExecutionList{
297264
TaskExecutions: taskExecutionList,
298265
Token: token,

pkg/manager/impl/util/shared.go

Lines changed: 145 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@ package util
33

44
import (
55
"context"
6+
"strconv"
67
"time"
78

89
"github.com/flyteorg/flyteadmin/pkg/common"
@@ -327,3 +328,147 @@ func MergeIntoExecConfig(workflowExecConfig admin.WorkflowExecutionConfig, spec
327328

328329
return workflowExecConfig
329330
}
331+
332+
func ListNodeExecutions(ctx context.Context, repo repoInterfaces.Repository, identifierFilters []common.InlineFilter,
333+
requestFilters string, limit uint32, requestToken string, sortBy *admin.Sort,
334+
mapFilters []common.MapFilter) ([]models.NodeExecution, string, error) {
335+
filters, err := AddRequestFilters(requestFilters, common.NodeExecution, identifierFilters)
336+
if err != nil {
337+
return nil, "", err
338+
}
339+
var sortParameter common.SortParameter
340+
if sortBy != nil {
341+
sortParameter, err = common.NewSortParameter(*sortBy)
342+
if err != nil {
343+
return nil, "", err
344+
}
345+
}
346+
offset, err := validation.ValidateToken(requestToken)
347+
if err != nil {
348+
return nil, "", errors.NewFlyteAdminErrorf(codes.InvalidArgument,
349+
"invalid pagination token %s for ListNodeExecutions", requestToken)
350+
}
351+
listInput := repoInterfaces.ListResourceInput{
352+
Limit: int(limit),
353+
Offset: offset,
354+
InlineFilters: filters,
355+
SortParameter: sortParameter,
356+
MapFilters: mapFilters,
357+
}
358+
359+
output, err := repo.NodeExecutionRepo().List(ctx, listInput)
360+
if err != nil {
361+
logger.Debugf(ctx, "Failed to list node executions: %v", err)
362+
return nil, "", err
363+
}
364+
365+
var token string
366+
if len(output.NodeExecutions) == int(limit) {
367+
token = strconv.Itoa(offset + len(output.NodeExecutions))
368+
}
369+
370+
return output.NodeExecutions, token, nil
371+
}
372+
373+
func ListNodeExecutionsForWorkflow(ctx context.Context, repo repoInterfaces.Repository,
374+
workflowExecutionID *core.WorkflowExecutionIdentifier, uniqueParentID string, requestFilters string,
375+
limit uint32, requestToken string, sortBy *admin.Sort) ([]models.NodeExecution, string, error) {
376+
identifierFilters, err := GetWorkflowExecutionIdentifierFilters(ctx, *workflowExecutionID)
377+
if err != nil {
378+
return nil, "", err
379+
}
380+
381+
var mapFilters []common.MapFilter
382+
if len(uniqueParentID) > 0 {
383+
parentNodeExecution, err := GetNodeExecutionModel(ctx, repo, &core.NodeExecutionIdentifier{
384+
ExecutionId: workflowExecutionID,
385+
NodeId: uniqueParentID,
386+
})
387+
if err != nil {
388+
return nil, "", err
389+
}
390+
parentIDFilter, err := common.NewSingleValueFilter(
391+
common.NodeExecution, common.Equal, shared.ParentID, parentNodeExecution.ID)
392+
if err != nil {
393+
return nil, "", err
394+
}
395+
identifierFilters = append(identifierFilters, parentIDFilter)
396+
} else {
397+
mapFilters = []common.MapFilter{
398+
common.NewMapFilter(map[string]interface{}{
399+
shared.ParentTaskExecutionID: nil,
400+
shared.ParentID: nil,
401+
}),
402+
}
403+
}
404+
405+
return ListNodeExecutions(ctx, repo, identifierFilters, requestFilters, limit, requestToken, sortBy, mapFilters)
406+
}
407+
408+
func ListNodeExecutionsForTask(ctx context.Context, repo repoInterfaces.Repository,
409+
taskExecutionID *core.TaskExecutionIdentifier, workflowExecutionID *core.WorkflowExecutionIdentifier,
410+
requestFilters string, limit uint32, requestToken string, sortBy *admin.Sort) ([]models.NodeExecution, string, error) {
411+
identifierFilters, err := GetWorkflowExecutionIdentifierFilters(ctx, *workflowExecutionID)
412+
if err != nil {
413+
return nil, "", err
414+
}
415+
416+
parentTaskExecutionModel, err := GetTaskExecutionModel(ctx, repo, taskExecutionID)
417+
if err != nil {
418+
return nil, "", err
419+
}
420+
421+
nodeIDFilter, err := common.NewSingleValueFilter(
422+
common.NodeExecution, common.Equal, shared.ParentTaskExecutionID, parentTaskExecutionModel.ID)
423+
if err != nil {
424+
return nil, "", err
425+
}
426+
identifierFilters = append(identifierFilters, nodeIDFilter)
427+
428+
return ListNodeExecutions(ctx, repo, identifierFilters, requestFilters, limit, requestToken, sortBy, nil)
429+
}
430+
431+
func ListTaskExecutions(ctx context.Context, repo repoInterfaces.Repository,
432+
nodeExecutionID *core.NodeExecutionIdentifier, requestFilters string, limit uint32, requestToken string,
433+
sortBy *admin.Sort) ([]models.TaskExecution, string, error) {
434+
identifierFilters, err := GetNodeExecutionIdentifierFilters(ctx, *nodeExecutionID)
435+
if err != nil {
436+
return nil, "", err
437+
}
438+
439+
filters, err := AddRequestFilters(requestFilters, common.TaskExecution, identifierFilters)
440+
if err != nil {
441+
return nil, "", err
442+
}
443+
var sortParameter common.SortParameter
444+
if sortBy != nil {
445+
sortParameter, err = common.NewSortParameter(*sortBy)
446+
if err != nil {
447+
return nil, "", err
448+
}
449+
}
450+
451+
offset, err := validation.ValidateToken(requestToken)
452+
if err != nil {
453+
return nil, "", errors.NewFlyteAdminErrorf(codes.InvalidArgument,
454+
"invalid pagination token %s for ListTaskExecutions", requestToken)
455+
}
456+
457+
output, err := repo.TaskExecutionRepo().List(ctx, repoInterfaces.ListResourceInput{
458+
InlineFilters: filters,
459+
Offset: offset,
460+
Limit: int(limit),
461+
SortParameter: sortParameter,
462+
})
463+
if err != nil {
464+
logger.Debugf(ctx, "Failed to list task executions: %v", err)
465+
return nil, "", err
466+
}
467+
468+
var token string
469+
if len(output.TaskExecutions) == int(limit) {
470+
token = strconv.Itoa(offset + len(output.TaskExecutions))
471+
}
472+
473+
return output.TaskExecutions, token, nil
474+
}

0 commit comments

Comments
 (0)