diff --git a/apps/api/internal/runner/admission.go b/apps/api/internal/runner/admission.go index d198ee6..036629e 100644 --- a/apps/api/internal/runner/admission.go +++ b/apps/api/internal/runner/admission.go @@ -66,6 +66,15 @@ func distributedAdmissionModelType(modelType string) bool { } func (s *Service) buildTaskAdmissionPlan(ctx context.Context, task store.GatewayTask, user *auth.User) (taskAdmissionPlan, error) { + return s.buildTaskAdmissionPlanForCurrentBinding(ctx, task, user, nil) +} + +func (s *Service) buildTaskAdmissionPlanForCurrentBinding( + ctx context.Context, + task store.GatewayTask, + user *auth.User, + admission *store.TaskAdmission, +) (taskAdmissionPlan, error) { restoredRequest, err := s.restoreTaskRequestReferences(ctx, task) if err != nil { return taskAdmissionPlan{}, err @@ -107,6 +116,7 @@ func (s *Service) buildTaskAdmissionPlan(ctx context.Context, task store.Gateway if err != nil { return taskAdmissionPlan{}, err } + candidates, _ = pinCandidatesToTaskAdmission(candidates, admission) for _, candidate := range candidates { available, availabilityErr := s.store.RuntimeCandidateAvailable(ctx, candidate.PlatformID, candidate.PlatformModelID) if availabilityErr != nil { @@ -220,7 +230,9 @@ func pinCandidatesToTaskAdmission( candidates []store.RuntimeModelCandidate, admission *store.TaskAdmission, ) ([]store.RuntimeModelCandidate, bool) { - if admission == nil || admission.Status != "admitted" || len(candidates) < 2 { + if admission == nil || + (admission.Status != "waiting" && admission.Status != "admitted") || + len(candidates) < 2 { return candidates, false } pinnedIndex := -1 @@ -629,7 +641,14 @@ func (s *Service) dispatchWaitingAsyncTasks(ctx context.Context, tasks []store.G tasksByID := make(map[string]store.GatewayTask, len(tasks)) for _, task := range tasks { user := authUserFromTask(task) - plan, err := s.buildTaskAdmissionPlan(ctx, task, user) + current, err := s.store.GetTaskAdmission(ctx, task.ID) + if err != nil { + if errors.Is(err, pgx.ErrNoRows) { + continue + } + return false, err + } + plan, err := s.buildTaskAdmissionPlanForCurrentBinding(ctx, task, user, ¤t) if err != nil { return false, err } diff --git a/apps/api/internal/runner/admission_test.go b/apps/api/internal/runner/admission_test.go index 4d8aa43..f2d64e0 100644 --- a/apps/api/internal/runner/admission_test.go +++ b/apps/api/internal/runner/admission_test.go @@ -87,31 +87,41 @@ func TestPinCandidatesToTaskAdmissionPreservesAdmittedCandidate(t *testing.T) { } } -func TestPinCandidatesToTaskAdmissionIgnoresWaitingOrMissingCandidate(t *testing.T) { +func TestPinCandidatesToTaskAdmissionPreservesWaitingCandidate(t *testing.T) { input := []store.RuntimeModelCandidate{ {PlatformID: "platform-a", PlatformModelID: "model-a"}, {PlatformID: "platform-b", PlatformModelID: "model-b"}, } - for name, admission := range map[string]*store.TaskAdmission{ - "waiting": { - Status: "waiting", - PlatformID: "platform-b", - PlatformModelID: "model-b", - }, - "missing": { - Status: "admitted", - PlatformID: "platform-c", - PlatformModelID: "model-c", - }, - } { - t.Run(name, func(t *testing.T) { - got, pinned := pinCandidatesToTaskAdmission(input, admission) - if pinned { - t.Fatalf("unexpected candidate pin for %s: %+v", name, got) - } - if got[0].PlatformModelID != "model-a" { - t.Fatalf("candidate order changed for %s: %+v", name, got) - } - }) + admission := &store.TaskAdmission{ + Status: "waiting", + PlatformID: "platform-b", + PlatformModelID: "model-b", + } + + got, pinned := pinCandidatesToTaskAdmission(input, admission) + + if !pinned || got[0].PlatformModelID != "model-b" { + t.Fatalf("waiting candidate was not pinned: pinned=%v candidates=%+v", pinned, got) + } +} + +func TestPinCandidatesToTaskAdmissionIgnoresMissingCandidate(t *testing.T) { + input := []store.RuntimeModelCandidate{ + {PlatformID: "platform-a", PlatformModelID: "model-a"}, + {PlatformID: "platform-b", PlatformModelID: "model-b"}, + } + admission := &store.TaskAdmission{ + Status: "admitted", + PlatformID: "platform-c", + PlatformModelID: "model-c", + } + + got, pinned := pinCandidatesToTaskAdmission(input, admission) + + if pinned { + t.Fatalf("unexpected candidate pin for missing binding: %+v", got) + } + if got[0].PlatformModelID != "model-a" { + t.Fatalf("candidate order changed for missing binding: %+v", got) } }