diff --git a/apps/api/internal/runner/admission.go b/apps/api/internal/runner/admission.go index 4c7d6f1..fd91566 100644 --- a/apps/api/internal/runner/admission.go +++ b/apps/api/internal/runner/admission.go @@ -187,6 +187,45 @@ func acceptanceAdmissionScopes(task store.GatewayTask, scopes []store.AdmissionS return out } +func (s *Service) loadAsyncTaskAdmission(ctx context.Context, task store.GatewayTask) (*store.TaskAdmission, error) { + if !task.AsyncMode { + return nil, nil + } + admission, err := s.store.GetTaskAdmission(ctx, task.ID) + if errors.Is(err, pgx.ErrNoRows) { + return nil, nil + } + if err != nil { + return nil, err + } + return &admission, nil +} + +func pinCandidatesToTaskAdmission( + candidates []store.RuntimeModelCandidate, + admission *store.TaskAdmission, +) ([]store.RuntimeModelCandidate, bool) { + if admission == nil || admission.Status != "admitted" || len(candidates) < 2 { + return candidates, false + } + pinnedIndex := -1 + for index := range candidates { + if candidates[index].PlatformID == admission.PlatformID && + candidates[index].PlatformModelID == admission.PlatformModelID { + pinnedIndex = index + break + } + } + if pinnedIndex <= 0 { + return candidates, false + } + out := append([]store.RuntimeModelCandidate(nil), candidates...) + pinned := out[pinnedIndex] + copy(out[1:pinnedIndex+1], out[:pinnedIndex]) + out[0] = pinned + return out, true +} + func (s *Service) tryTaskAdmission(ctx context.Context, task store.GatewayTask, plan taskAdmissionPlan, waiterID string) (store.TaskAdmissionResult, error) { return s.tryTaskAdmissionWithAdmittedHook(ctx, task, plan, waiterID, nil) } @@ -195,17 +234,11 @@ func (s *Service) activeAsyncTaskAdmission( ctx context.Context, task store.GatewayTask, plan taskAdmissionPlan, + admission *store.TaskAdmission, ) (store.TaskAdmissionResult, bool, error) { - if !task.AsyncMode { + if !task.AsyncMode || admission == nil { return store.TaskAdmissionResult{}, false, nil } - admission, err := s.store.GetTaskAdmission(ctx, task.ID) - if errors.Is(err, pgx.ErrNoRows) { - return store.TaskAdmissionResult{}, false, nil - } - if err != nil { - return store.TaskAdmissionResult{}, false, err - } if admission.Status != "admitted" || admission.PlatformID != plan.Candidate.PlatformID || admission.PlatformModelID != plan.Candidate.PlatformModelID || @@ -220,7 +253,7 @@ func (s *Service) activeAsyncTaskAdmission( return store.TaskAdmissionResult{}, false, nil } return store.TaskAdmissionResult{ - Admission: admission, + Admission: *admission, Admitted: true, Leases: leases, }, true, nil diff --git a/apps/api/internal/runner/admission_test.go b/apps/api/internal/runner/admission_test.go index e8ee007..4d8aa43 100644 --- a/apps/api/internal/runner/admission_test.go +++ b/apps/api/internal/runner/admission_test.go @@ -61,3 +61,57 @@ func TestAcceptanceAdmissionScopesLeaveProductionPolicyUnchanged(t *testing.T) { t.Fatalf("production admission policy changed: %+v", got[0]) } } + +func TestPinCandidatesToTaskAdmissionPreservesAdmittedCandidate(t *testing.T) { + input := []store.RuntimeModelCandidate{ + {PlatformID: "platform-a", PlatformModelID: "model-a"}, + {PlatformID: "platform-b", PlatformModelID: "model-b"}, + {PlatformID: "platform-c", PlatformModelID: "model-c"}, + } + admission := &store.TaskAdmission{ + Status: "admitted", + PlatformID: "platform-b", + PlatformModelID: "model-b", + } + + got, pinned := pinCandidatesToTaskAdmission(input, admission) + + if !pinned { + t.Fatal("expected admitted candidate to be pinned") + } + if got[0].PlatformModelID != "model-b" || got[1].PlatformModelID != "model-a" || got[2].PlatformModelID != "model-c" { + t.Fatalf("unexpected candidate order: %+v", got) + } + if input[0].PlatformModelID != "model-a" { + t.Fatalf("candidate pinning mutated the caller slice: %+v", input) + } +} + +func TestPinCandidatesToTaskAdmissionIgnoresWaitingOrMissingCandidate(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) + } + }) + } +} diff --git a/apps/api/internal/runner/service.go b/apps/api/internal/runner/service.go index a5eb450..970698d 100644 --- a/apps/api/internal/runner/service.go +++ b/apps/api/internal/runner/service.go @@ -422,6 +422,18 @@ func (s *Service) executeWithToken(ctx context.Context, task store.GatewayTask, return Result{Task: failed, Output: failed.Result}, err } } + var asyncAdmission *store.TaskAdmission + if distributedAdmission && task.AsyncMode { + asyncAdmission, err = s.loadAsyncTaskAdmission(ctx, task) + if err != nil { + return Result{}, err + } + var pinned bool + candidates, pinned = pinCandidatesToTaskAdmission(candidates, asyncAdmission) + if pinned { + s.observeTaskAdmission("candidate_pinned") + } + } pricingByCandidate := map[string]resolvedPricing{} preprocessingByCandidate := map[string]parameterPreprocessResult{} reservationBillings := []any(nil) @@ -561,7 +573,7 @@ func (s *Service) executeWithToken(ctx context.Context, task store.GatewayTask, var admissionErr error if task.AsyncMode { var alreadyAdmitted bool - admissionResult, alreadyAdmitted, admissionErr = s.activeAsyncTaskAdmission(ctx, task, plan) + admissionResult, alreadyAdmitted, admissionErr = s.activeAsyncTaskAdmission(ctx, task, plan, asyncAdmission) if admissionErr == nil && !alreadyAdmitted { admissionResult, admissionErr = s.tryTaskAdmission(ctx, task, plan, "") }