diff --git a/cmd/ateapi/internal/actoridentity/actoridentity.go b/cmd/ateapi/internal/actoridentity/actoridentity.go index 8de6dc0b3..c48c25877 100644 --- a/cmd/ateapi/internal/actoridentity/actoridentity.go +++ b/cmd/ateapi/internal/actoridentity/actoridentity.go @@ -29,7 +29,6 @@ import ( "github.com/agent-substrate/substrate/cmd/ateapi/internal/actoridjwt" "github.com/agent-substrate/substrate/cmd/ateapi/internal/store" - "github.com/agent-substrate/substrate/cmd/ateapi/internal/workercache" "github.com/agent-substrate/substrate/internal/localca" "github.com/agent-substrate/substrate/internal/localjwtauthority" "github.com/agent-substrate/substrate/internal/principal" @@ -54,19 +53,17 @@ type Server struct { // store is the actor database. MintCert consults it to confirm the caller // is entitled to the actor it is asking for a credential for. - store store.Interface - workers *workercache.Cache + store store.Interface } var _ ateapipb.ActorIdentityServer = (*Server)(nil) -func New(actorIdentityJWTIssuer, actorIDJWTPoolFile, actorIDCAPoolFile string, store store.Interface, workers *workercache.Cache) *Server { +func New(actorIdentityJWTIssuer, actorIDJWTPoolFile, actorIDCAPoolFile string, store store.Interface) *Server { return &Server{ actorIdentityJWTIssuer: actorIdentityJWTIssuer, actorIDJWTPoolFile: actorIDJWTPoolFile, actorIDCAPoolFile: actorIDCAPoolFile, store: store, - workers: workers, } } @@ -312,7 +309,7 @@ func (s *Server) authorizeActor(ctx context.Context, caller *ateletCaller, req * return status.Errorf(codes.PermissionDenied, "caller is not permitted to mint credentials for this actor") } - worker, err := s.workers.Worker(req.GetWorkerNamespace(), req.GetWorkerPod()) + worker, err := s.store.GetWorker(ctx, req.GetWorkerNamespace(), req.GetWorkerPod()) if err != nil { if errors.Is(err, store.ErrNotFound) { return nil, resources.ActorRef{}, deny("worker not found") diff --git a/cmd/ateapi/internal/actoridentity/actoridentity_test.go b/cmd/ateapi/internal/actoridentity/actoridentity_test.go index e91420d38..a351d3e7e 100644 --- a/cmd/ateapi/internal/actoridentity/actoridentity_test.go +++ b/cmd/ateapi/internal/actoridentity/actoridentity_test.go @@ -31,7 +31,6 @@ import ( "github.com/agent-substrate/substrate/cmd/ateapi/internal/store" "github.com/agent-substrate/substrate/cmd/ateapi/internal/store/storetest" - "github.com/agent-substrate/substrate/cmd/ateapi/internal/workercache" "github.com/agent-substrate/substrate/internal/localca" "github.com/agent-substrate/substrate/internal/principal" "github.com/agent-substrate/substrate/internal/resources" @@ -131,9 +130,8 @@ func ctxWithCert(cert *x509.Certificate) context.Context { }) } -// newTestServer returns a Server backed by st, with a freshly generated actor -// CA pool written to a temp file. -func newTestServer(t *testing.T, st store.Interface) *Server { +// newActorCAPoolFile writes a freshly generated actor CA pool to a temp file. +func newActorCAPoolFile(t *testing.T) string { t.Helper() ca, err := localca.GenerateED25519CA("test-actor-ca") @@ -148,17 +146,15 @@ func newTestServer(t *testing.T, st store.Interface) *Server { if err := os.WriteFile(poolFile, poolBytes, 0o600); err != nil { t.Fatalf("write CA pool: %v", err) } + return poolFile +} - var workers *workercache.Cache - if st != nil { - workers = workercache.New(st, time.Hour) - ctx, cancel := context.WithCancel(context.Background()) - t.Cleanup(cancel) - if err := workers.Start(ctx); err != nil { - t.Fatalf("start worker cache: %v", err) - } - } - return New("issuer", "", poolFile, st, workers) +// newTestServer returns a Server backed by st, with a freshly generated actor +// CA pool written to a temp file. +func newTestServer(t *testing.T, st store.Interface) *Server { + t.Helper() + + return New("issuer", "", newActorCAPoolFile(t), st) } func TestMintJWTRequiresConfiguredJWTProvider(t *testing.T) { @@ -707,6 +703,41 @@ func TestMintCertDeniesUnassignedActorWhateverItsStatus(t *testing.T) { } } +// TestMintCertUsesAnAssignmentAsSoonAsItIsWritten guards the invariant the gate +// rests on: it reads the store, so an assignment is usable the instant resume +// writes it. Behind a replica-local cache, atelet mints inside the lag window. +func TestMintCertUsesAnAssignmentAsSoonAsItIsWritten(t *testing.T) { + ctx := context.Background() + st, cleanup := storetest.SetupTestStore(t) + defer cleanup() + + // The state pause leaves behind: the worker exists but holds no assignment. + seedActor(t, ctx, st, actorFixture{ + status: ateapipb.Actor_STATUS_RESUMING, + workerNode: testNode, + unassigned: true, + }) + srv := newTestServer(t, st) + + actorRef := resources.ActorRef{Atespace: testAtespace, Name: testActorName} + actor, err := st.GetActor(ctx, actorRef) + if err != nil { + t.Fatal(err) + } + worker, err := st.GetWorker(ctx, testPodNS, testWorkerPod) + if err != nil { + t.Fatal(err) + } + worker.Assignment = &ateapipb.Assignment{Actor: actorRef.ToObjectRef(), ActorUid: actor.GetMetadata().GetUid()} + if err := st.UpdateWorker(ctx, worker, worker.GetVersion()); err != nil { + t.Fatalf("assign worker: %v", err) + } + + if _, err := srv.MintCert(ctxWithCert(ateletCertOn(t, testNode)), mintCertRequest(t, actor.GetMetadata().GetUid())); err != nil { + t.Errorf("MintCert() error = %v, want success from the freshly written assignment", err) + } +} + // TestMintCertAuthorizesBeforeSigning checks that the gate runs before any CSR // parsing or CA material is touched. An unauthorized caller must be rejected // with PermissionDenied even when the rest of the request is unusable, so that @@ -720,13 +751,7 @@ func TestMintCertAuthorizesBeforeSigning(t *testing.T) { // A server whose CA pool file does not exist: reaching the signing path at // all would surface as Internal rather than PermissionDenied. - workers := workercache.New(st, time.Hour) - cacheCtx, cancel := context.WithCancel(ctx) - defer cancel() - if err := workers.Start(cacheCtx); err != nil { - t.Fatal(err) - } - srv := New("issuer", "", filepath.Join(t.TempDir(), "missing.json"), st, workers) + srv := New("issuer", "", filepath.Join(t.TempDir(), "missing.json"), st) actor, err := st.GetActor(ctx, resources.ActorRef{Atespace: testAtespace, Name: testActorName}) if err != nil { diff --git a/cmd/ateapi/internal/controlapi/crash.go b/cmd/ateapi/internal/controlapi/crash.go index e12593c4b..74a73883b 100644 --- a/cmd/ateapi/internal/controlapi/crash.go +++ b/cmd/ateapi/internal/controlapi/crash.go @@ -108,7 +108,7 @@ func releaseWorker(ctx context.Context, st store.Interface, actor *ateapipb.Acto } podUid := assignment.GetWorkerPodUid() - worker, err := st.GetWorker(ctx, assignment.GetWorkerNamespace(), assignment.GetWorkerPool(), assignment.GetWorkerPod()) + worker, err := st.GetWorker(ctx, assignment.GetWorkerNamespace(), assignment.GetWorkerPod()) if errors.Is(err, store.ErrNotFound) { // No need to release if the worker is not found. slog.WarnContext(ctx, "Worker already gone while crashing actor, skipping release", slog.String("worker", podUid)) diff --git a/cmd/ateapi/internal/controlapi/crash_test.go b/cmd/ateapi/internal/controlapi/crash_test.go index 620981e87..abf231b4b 100644 --- a/cmd/ateapi/internal/controlapi/crash_test.go +++ b/cmd/ateapi/internal/controlapi/crash_test.go @@ -145,7 +145,7 @@ func TestCrashActor(t *testing.T) { t.Fatalf("crashActor() = %v, want nil", err) } assertCrashed(t, ctx, st, actorRef) - worker, gerr := st.GetWorker(ctx, "ns", "pool", "pod") + worker, gerr := st.GetWorker(ctx, "ns", "pod") if gerr != nil { t.Fatalf("GetWorker() = %v, want nil", gerr) } @@ -165,7 +165,7 @@ func TestCrashActor(t *testing.T) { t.Fatalf("crashActor() = %v, want nil", err) } assertCrashed(t, ctx, st, actorRef) - worker, gerr := st.GetWorker(ctx, "ns", "pool", "pod") + worker, gerr := st.GetWorker(ctx, "ns", "pod") if gerr != nil { t.Fatalf("GetWorker() = %v, want nil", gerr) } @@ -200,7 +200,7 @@ func TestCrashActor(t *testing.T) { t.Fatalf("crashActor() = %v, want nil", err) } assertCrashed(t, ctx, st, actorRef) - worker, gerr := st.GetWorker(ctx, "ns", "pool", "pod") + worker, gerr := st.GetWorker(ctx, "ns", "pod") if gerr != nil { t.Fatalf("GetWorker() = %v, want nil", gerr) } @@ -227,7 +227,7 @@ func TestCrashActor(t *testing.T) { // Without a binding the worker cannot be looked up, so its // assignment must be left untouched even though it names // the crashed actor. - worker, gerr := st.GetWorker(ctx, "ns", "pool", "pod") + worker, gerr := st.GetWorker(ctx, "ns", "pod") if gerr != nil { t.Fatalf("GetWorker() = %v, want nil", gerr) } diff --git a/cmd/ateapi/internal/controlapi/functional_test.go b/cmd/ateapi/internal/controlapi/functional_test.go index 0087b0da1..75eb086ad 100644 --- a/cmd/ateapi/internal/controlapi/functional_test.go +++ b/cmd/ateapi/internal/controlapi/functional_test.go @@ -2725,7 +2725,7 @@ func TestResumeActor_CrashesIfAssignedWorkerIsDraining(t *testing.T) { t.Fatalf("expected actor to be bound to a worker after the failed attempt") } - assigned, err := tc.persistence.GetWorker(context.Background(), ns, "pool1", assignedPod) + assigned, err := tc.persistence.GetWorker(context.Background(), ns, assignedPod) if err != nil { t.Fatalf("GetWorker(%s) failed: %v", assignedPod, err) } diff --git a/cmd/ateapi/internal/controlapi/syncer.go b/cmd/ateapi/internal/controlapi/syncer.go index 93f3d0507..419748632 100644 --- a/cmd/ateapi/internal/controlapi/syncer.go +++ b/cmd/ateapi/internal/controlapi/syncer.go @@ -38,12 +38,10 @@ import ( // ordering is preserved. const syncerWorkerCount = 2 -// workerKey identifies a worker row in the store. pool is captured at enqueue -// time because once the pod is gone from the informer cache, the -// ate.dev/worker-pool label cannot be recovered from a namespace/name key. +// workerKey is the Pod a reconcile is about. Keying on the pool too would let a +// task queued for the previous owner run alongside, and undo, its successor. type workerKey struct { namespace string - pool string name string } @@ -79,15 +77,8 @@ func (s *WorkerPoolSyncer) Start(ctx context.Context) { AddFunc: func(obj interface{}) { s.enqueuePod(obj.(*corev1.Pod)) }, - UpdateFunc: func(oldObj, newObj interface{}) { - oldPod := oldObj.(*corev1.Pod) - newPod := newObj.(*corev1.Pod) - // If the pool label changed, enqueue the old key too so its - // now-stale store row gets cleaned up. - if oldPod.Labels[workerPodLabel] != newPod.Labels[workerPodLabel] { - s.enqueuePod(oldPod) - } - s.enqueuePod(newPod) + UpdateFunc: func(_, newObj interface{}) { + s.enqueuePod(newObj.(*corev1.Pod)) }, DeleteFunc: func(obj interface{}) { var pod *corev1.Pod @@ -132,7 +123,7 @@ func (s *WorkerPoolSyncer) Start(ctx context.Context) { } func (s *WorkerPoolSyncer) enqueuePod(pod *corev1.Pod) { - s.queue.Add(workerKey{namespace: pod.Namespace, pool: pod.Labels[workerPodLabel], name: pod.Name}) + s.queue.Add(workerKey{namespace: pod.Namespace, name: pod.Name}) } func (s *WorkerPoolSyncer) runWorker(ctx context.Context) { @@ -150,7 +141,6 @@ func (s *WorkerPoolSyncer) processNextWorkItem(ctx context.Context) bool { if err := s.reconcile(ctx, key); err != nil { slog.ErrorContext(ctx, "Syncer: reconcile failed, requeueing", slog.String("worker", key.namespace+"/"+key.name), - slog.String("pool", key.pool), slog.Any("err", err)) s.queue.AddRateLimited(key) return true @@ -168,13 +158,9 @@ func (s *WorkerPoolSyncer) reconcile(ctx context.Context, key workerKey) error { } if !exists { slog.InfoContext(ctx, "Syncer: removing worker from store (pod deleted)", slog.String("worker", key.namespace+"/"+key.name)) - return s.reconcileDeadWorker(ctx, key.namespace, key.pool, key.name) + return s.reconcileDeadWorker(ctx, key) } pod := obj.(*corev1.Pod) - if pod.Labels[workerPodLabel] != key.pool { - // The pod moved to a different pool; this key's store row is stale. - return s.reconcileDeadWorker(ctx, key.namespace, key.pool, key.name) - } // Checked before eligibility: draining works off the stored record by name and // never reads the pod IP, while a Terminating pod can legitimately report no // IP once its sandbox is torn down. Gating on the IP first would drop the @@ -185,7 +171,7 @@ func (s *WorkerPoolSyncer) reconcile(ctx context.Context, key workerKey) error { // the bound actor here — inside the pod ateom has received SIGTERM and is // gracefully shutting the actor down. Actor cleanup happens on the Pod // Deleted event. - return s.markWorkerDraining(ctx, key.namespace, key.pool, key.name) + return s.markWorkerDraining(ctx, key) } if !isWorkerEligible(pod) { // No IP yet; a later update event re-enqueues the pod. @@ -195,12 +181,13 @@ func (s *WorkerPoolSyncer) reconcile(ctx context.Context, key workerKey) error { } func (s *WorkerPoolSyncer) createOrUpdateWorker(ctx context.Context, key workerKey, pod *corev1.Pod) error { - pool, err := s.workerPoolLister.WorkerPools(key.namespace).Get(key.pool) + poolName := pod.Labels[workerPodLabel] + pool, err := s.workerPoolLister.WorkerPools(key.namespace).Get(poolName) if err != nil { - return fmt.Errorf("getting WorkerPool %s/%s: %w", key.namespace, key.pool, err) + return fmt.Errorf("getting WorkerPool %s/%s: %w", key.namespace, poolName, err) } - w, err := s.persistence.GetWorker(ctx, key.namespace, key.pool, key.name) + w, err := s.persistence.GetWorker(ctx, key.namespace, key.name) if err != nil { if !errors.Is(err, store.ErrNotFound) { return fmt.Errorf("getting worker from store: %w", err) @@ -208,7 +195,7 @@ func (s *WorkerPoolSyncer) createOrUpdateWorker(ctx context.Context, key workerK slog.InfoContext(ctx, "Syncer: creating worker in store", slog.String("worker", key.namespace+"/"+key.name)) worker := &ateapipb.Worker{ WorkerNamespace: pod.Namespace, - WorkerPool: key.pool, + WorkerPool: poolName, WorkerPod: pod.Name, Ip: pod.Status.PodIP, WorkerPodUid: string(pod.UID), @@ -238,10 +225,18 @@ func (s *WorkerPoolSyncer) createOrUpdateWorker(ctx context.Context, key workerK // clean it up (releasing any actor bound to the old incarnation) and // requeue to create a fresh row. slog.InfoContext(ctx, "Syncer: worker in store belongs to a replaced pod, deleting", slog.String("worker", key.namespace+"/"+key.name)) - if err := s.reconcileDeadWorker(ctx, key.namespace, key.pool, key.name); err != nil { + if err := s.deleteWorkerAndReleaseActor(ctx, w); err != nil { + return err + } + return fmt.Errorf("worker %s/%s belonged to replaced pod UID %s; requeueing to recreate", key.namespace, key.name, w.WorkerPodUid) + } + if w.WorkerPool != poolName { + // The pool is immutable on the record, so a relabel rebuilds the row. + slog.InfoContext(ctx, "Syncer: worker in store belongs to a previous pool, deleting", slog.String("worker", key.namespace+"/"+key.name)) + if err := s.deleteWorkerAndReleaseActor(ctx, w); err != nil { return err } - return fmt.Errorf("worker %s/%s/%s belonged to replaced pod UID %s; requeueing to recreate", key.namespace, key.pool, key.name, w.WorkerPodUid) + return fmt.Errorf("worker %s/%s belonged to pool %s; requeueing to recreate under %s", key.namespace, key.name, w.WorkerPool, poolName) } changed := false @@ -311,8 +306,8 @@ func workerCapacity(pod *corev1.Pod) *ateapipb.WorkerCapacity { // already gone or already draining there is nothing more to do — the Pod // Deleted event will clean up the record. A version conflict is returned so the // caller requeues and retries against the updated record. -func (s *WorkerPoolSyncer) markWorkerDraining(ctx context.Context, namespace, pool, podName string) error { - worker, err := s.persistence.GetWorker(ctx, namespace, pool, podName) +func (s *WorkerPoolSyncer) markWorkerDraining(ctx context.Context, key workerKey) error { + worker, err := s.persistence.GetWorker(ctx, key.namespace, key.name) if err != nil { if errors.Is(err, store.ErrNotFound) { return nil @@ -322,22 +317,31 @@ func (s *WorkerPoolSyncer) markWorkerDraining(ctx context.Context, namespace, po if worker.GetState() == ateapipb.Worker_STATE_DRAINING { return nil } - slog.InfoContext(ctx, "Syncer: marking worker draining (pod deleting)", slog.String("worker", namespace+"/"+podName)) + slog.InfoContext(ctx, "Syncer: marking worker draining (pod deleting)", slog.String("worker", key.namespace+"/"+key.name)) worker.State = ateapipb.Worker_STATE_DRAINING return s.persistence.UpdateWorker(ctx, worker, worker.GetVersion()) } -// reconcileDeadWorker cleans up a worker whose pod is gone. It releases the -// bound actor first and only deletes the worker record if that succeeds: -// deleting the record is what erases the actor->pod pointer, so on a release -// failure we intentionally leave the record in place (and return the error) so a -// later reconcile can retry. Returns nil once the actor is released and the -// worker record deleted. -func (s *WorkerPoolSyncer) reconcileDeadWorker(ctx context.Context, namespace, pool, podName string) error { - if err := s.releaseActorOnDeadWorker(ctx, namespace, pool, podName); err != nil { +// reconcileDeadWorker cleans up the stored worker for a pod that is gone. +func (s *WorkerPoolSyncer) reconcileDeadWorker(ctx context.Context, key workerKey) error { + worker, err := s.persistence.GetWorker(ctx, key.namespace, key.name) + if err != nil { + if errors.Is(err, store.ErrNotFound) { + return nil + } return err } - return s.persistence.DeleteWorker(ctx, namespace, pool, podName) + return s.deleteWorkerAndReleaseActor(ctx, worker) +} + +// deleteWorkerAndReleaseActor releases the actor before deleting the record: +// deleting it is what erases the actor->pod pointer, so a failed release leaves +// the record for a later retry. worker is also the delete's precondition. +func (s *WorkerPoolSyncer) deleteWorkerAndReleaseActor(ctx context.Context, worker *ateapipb.Worker) error { + if err := s.releaseActorOnDeadWorker(ctx, worker); err != nil { + return err + } + return s.persistence.DeleteWorker(ctx, worker, worker.GetVersion()) } // enqueueStoredWorkers enqueues a key for every worker record in the store. @@ -352,7 +356,7 @@ func (s *WorkerPoolSyncer) enqueueStoredWorkers(ctx context.Context) { return } for _, w := range page.Items { - s.queue.Add(workerKey{namespace: w.GetWorkerNamespace(), pool: w.GetWorkerPool(), name: w.GetWorkerPod()}) + s.queue.Add(workerKey{namespace: w.GetWorkerNamespace(), name: w.GetWorkerPod()}) } if !page.HasNextPage() { return @@ -370,14 +374,7 @@ func (s *WorkerPoolSyncer) enqueueStoredWorkers(ctx context.Context) { // UpdateActor uses optimistic version checking. A concurrent SuspendActor // or ResumeActor wins; we fail this attempt so it can be retried with the // updated state. -func (s *WorkerPoolSyncer) releaseActorOnDeadWorker(ctx context.Context, namespace, pool, podName string) error { - worker, err := s.persistence.GetWorker(ctx, namespace, pool, podName) - if err != nil { - if errors.Is(err, store.ErrNotFound) { - return nil - } - return err - } +func (s *WorkerPoolSyncer) releaseActorOnDeadWorker(ctx context.Context, worker *ateapipb.Worker) error { if worker.Assignment == nil || worker.Assignment.GetActor() == nil { return nil } @@ -392,9 +389,15 @@ func (s *WorkerPoolSyncer) releaseActorOnDeadWorker(ctx context.Context, namespa if actor.GetMetadata().GetUid() != worker.Assignment.GetActorUid() { return nil } - // Skip if a concurrent SuspendActor already cleared the pointer. + // Only the actor still placed on this exact worker incarnation may be + // released. Both ateapi replicas run a syncer, so a snapshot read before a + // replacement would otherwise crash the actor that has since moved on, and + // a concurrent SuspendActor clears the placement outright. assignment := actor.GetWorkerAssignment() - if assignment.GetWorkerNamespace() != namespace || assignment.GetWorkerPod() != podName { + if assignment.GetWorkerNamespace() != worker.GetWorkerNamespace() || + assignment.GetWorkerPool() != worker.GetWorkerPool() || + assignment.GetWorkerPod() != worker.GetWorkerPod() || + assignment.GetWorkerPodUid() != worker.GetWorkerPodUid() { return nil } // If the actor is suspended, it's already been released. diff --git a/cmd/ateapi/internal/controlapi/syncer_test.go b/cmd/ateapi/internal/controlapi/syncer_test.go index 917ef0299..88c600cd6 100644 --- a/cmd/ateapi/internal/controlapi/syncer_test.go +++ b/cmd/ateapi/internal/controlapi/syncer_test.go @@ -136,7 +136,7 @@ func TestSyncer_Lifecycle(t *testing.T) { // 3. Check it's not there (polled for 500ms) err = wait.PollUntilContextTimeout(context.Background(), 50*time.Millisecond, 500*time.Millisecond, true, func(ctx context.Context) (bool, error) { - _, err := persistence.GetWorker(ctx, ns, poolName, podName) + _, err := persistence.GetWorker(ctx, ns, podName) if err == nil { return false, fmt.Errorf("worker unexpectedly found in Redis") } @@ -165,7 +165,7 @@ func TestSyncer_Lifecycle(t *testing.T) { // 5. Check that it's added (eventually by polling) err = wait.PollUntilContextTimeout(context.Background(), 100*time.Millisecond, 2*time.Second, true, func(ctx context.Context) (bool, error) { - w, err := persistence.GetWorker(ctx, ns, poolName, podName) + w, err := persistence.GetWorker(ctx, ns, podName) if err != nil { if errors.Is(err, store.ErrNotFound) { return false, nil @@ -195,7 +195,7 @@ func TestSyncer_Lifecycle(t *testing.T) { // 9. Verify it's gone err = wait.PollUntilContextTimeout(context.Background(), 100*time.Millisecond, 2*time.Second, true, func(ctx context.Context) (bool, error) { - _, err := persistence.GetWorker(ctx, ns, poolName, podName) + _, err := persistence.GetWorker(ctx, ns, podName) if err != nil { if errors.Is(err, store.ErrNotFound) { return true, nil @@ -250,7 +250,7 @@ func TestSyncer_DeleteBoundWorker_ClearsActor(t *testing.T) { t.Fatalf("create pod: %v", err) } if err := wait.PollUntilContextTimeout(ctx, 50*time.Millisecond, 2*time.Second, true, func(c context.Context) (bool, error) { - _, gerr := persistence.GetWorker(c, ns, pool, pod) + _, gerr := persistence.GetWorker(c, ns, pod) return gerr == nil, nil }); err != nil { t.Fatalf("worker row not materialised: %v", err) @@ -260,7 +260,7 @@ func TestSyncer_DeleteBoundWorker_ClearsActor(t *testing.T) { Metadata: &ateapipb.ResourceMetadata{Name: actorName, Atespace: "team-orphan"}, ActorTemplateNamespace: ns, ActorTemplateName: "tmpl", Status: ateapipb.Actor_STATUS_RUNNING, WorkerAssignment: &ateapipb.WorkerAssignment{ - WorkerNamespace: ns, WorkerPool: pool, WorkerPod: pod, WorkerPodIp: ip, + WorkerNamespace: ns, WorkerPool: pool, WorkerPod: pod, WorkerPodIp: ip, WorkerPodUid: "11111111-1111-1111-1111-111111111111", }, // Both in-progress checkpoints are set so the assertion below covers the // shared crash path, which cannot know which workflow was in flight. @@ -271,7 +271,7 @@ func TestSyncer_DeleteBoundWorker_ClearsActor(t *testing.T) { if err != nil { t.Fatalf("create actor: %v", err) } - w, _ := persistence.GetWorker(ctx, ns, pool, pod) + w, _ := persistence.GetWorker(ctx, ns, pod) w.Assignment = &ateapipb.Assignment{ ActorTemplate: &ateapipb.KubeNamespacedObjectRef{ Namespace: ns, @@ -361,7 +361,7 @@ func TestSyncer_OmittedFields(t *testing.T) { // Verify that it is created in Redis with empty SandboxClass and empty Labels err = wait.PollUntilContextTimeout(context.Background(), 100*time.Millisecond, 2*time.Second, true, func(ctx context.Context) (bool, error) { - w, err := persistence.GetWorker(ctx, ns, poolName, podName) + w, err := persistence.GetWorker(ctx, ns, podName) if err != nil { if errors.Is(err, store.ErrNotFound) { return false, nil @@ -408,6 +408,125 @@ func setupReconcileTest(t *testing.T, persistence store.Interface, initPools ... return NewWorkerPoolSyncer(persistence, workerInformer, workerPoolLister) } +// TestSyncer_StaleWorkerSnapshotDoesNotCrashReassignedActor covers the window +// both ateapi replicas share: one reads a worker, the other replaces it, the +// actor is placed on the replacement, and the first replica then acts on its +// snapshot. Namespace and pod alone still match, so only the full placement +// separates the incarnations. +func TestSyncer_StaleWorkerSnapshotDoesNotCrashReassignedActor(t *testing.T) { + ctx := context.Background() + persistence, cleanup := storetest.SetupTestStore(t) + defer cleanup() + s := setupReconcileTest(t, persistence) + + ns, pod := "ns-stale", "worker-stale" + actor, err := persistence.CreateActor(ctx, &ateapipb.Actor{ + Metadata: &ateapipb.ResourceMetadata{Atespace: "team", Name: "actor-stale"}, + Status: ateapipb.Actor_STATUS_RUNNING, + WorkerAssignment: &ateapipb.WorkerAssignment{ + WorkerNamespace: ns, WorkerPool: "pool-new", WorkerPod: pod, WorkerPodUid: "uid-new", + }, + }) + if err != nil { + t.Fatalf("create actor: %v", err) + } + + stale := &ateapipb.Worker{ + WorkerNamespace: ns, WorkerPool: "pool-old", WorkerPod: pod, WorkerPodUid: "uid-old", + Assignment: &ateapipb.Assignment{ + Actor: &ateapipb.ObjectRef{Atespace: "team", Name: "actor-stale"}, + ActorUid: actor.GetMetadata().GetUid(), + }, + } + if err := s.releaseActorOnDeadWorker(ctx, stale); err != nil { + t.Fatalf("releaseActorOnDeadWorker: %v", err) + } + + got, err := persistence.GetActor(ctx, resources.ActorRef{Atespace: "team", Name: "actor-stale"}) + if err != nil { + t.Fatalf("get actor: %v", err) + } + if got.GetStatus() != ateapipb.Actor_STATUS_RUNNING { + t.Errorf("actor status = %v, want STATUS_RUNNING: a stale worker snapshot crashed the reassigned actor", got.GetStatus()) + } +} + +// TestSyncer_PoolRelabelEnqueuesOneKey is what removes the old-pool/new-pool +// race rather than guarding against it: both sides of a relabel address the same +// Pod, so the queue coalesces them and never runs two tasks for it at once. +func TestSyncer_PoolRelabelEnqueuesOneKey(t *testing.T) { + persistence, cleanup := storetest.SetupTestStore(t) + defer cleanup() + s := setupReconcileTest(t, persistence) + + meta := func(pool string) *corev1.Pod { + return &corev1.Pod{ObjectMeta: metav1.ObjectMeta{ + Name: "worker-relabel", Namespace: "ns-relabel", + Labels: map[string]string{workerPodLabel: pool}, + }} + } + s.enqueuePod(meta("pool1")) + s.enqueuePod(meta("pool2")) + + if got := s.queue.Len(); got != 1 { + t.Errorf("queue length after a relabel = %d, want 1 coalesced item", got) + } +} + +// TestSyncer_PoolRelabel_RebuildsRecord covers the relabel end to end: the pool +// is immutable on the record, so the row is rebuilt, and the live Pod must end +// up with a record rather than none. +func TestSyncer_PoolRelabel_RebuildsRecord(t *testing.T) { + ctx := context.Background() + persistence, cleanup := storetest.SetupTestStore(t) + defer cleanup() + + ns, pod, ip, uid := "ns-relabel", "worker-relabel", "10.0.0.9", "22222222-2222-2222-2222-222222222222" + pools := []*atev1alpha1.WorkerPool{ + {ObjectMeta: metav1.ObjectMeta{Name: "pool1", Namespace: ns}, Spec: atev1alpha1.WorkerPoolSpec{SandboxClass: "gvisor"}}, + {ObjectMeta: metav1.ObjectMeta{Name: "pool2", Namespace: ns}, Spec: atev1alpha1.WorkerPoolSpec{SandboxClass: "gvisor"}}, + } + s := setupReconcileTest(t, persistence, pools...) + + if err := persistence.CreateWorker(ctx, &ateapipb.Worker{ + WorkerNamespace: ns, WorkerPool: "pool1", WorkerPod: pod, Ip: ip, + WorkerPodUid: uid, NodeName: "node1", State: ateapipb.Worker_STATE_ACTIVE, + }); err != nil { + t.Fatalf("create worker: %v", err) + } + + relabelled := &corev1.Pod{ + ObjectMeta: metav1.ObjectMeta{ + Name: pod, Namespace: ns, UID: types.UID(uid), + Labels: map[string]string{workerPodLabel: "pool2"}, + }, + Spec: corev1.PodSpec{NodeName: "node1"}, + Status: corev1.PodStatus{Phase: corev1.PodRunning, PodIP: ip, PodIPs: []corev1.PodIP{{IP: ip}}}, + } + if err := s.workerInformer.GetIndexer().Add(relabelled); err != nil { + t.Fatalf("seed indexer: %v", err) + } + + key := workerKey{namespace: ns, name: pod} + if err := s.reconcile(ctx, key); err == nil { + t.Fatal("reconcile of a relabelled pod should requeue, got no error") + } + if err := s.reconcile(ctx, key); err != nil { + t.Fatalf("second reconcile: %v", err) + } + + got, err := persistence.GetWorker(ctx, ns, pod) + if err != nil { + t.Fatalf("the live pod must end up with a worker record: %v", err) + } + if got.GetWorkerPool() != "pool2" { + t.Errorf("worker pool = %q, want pool2", got.GetWorkerPool()) + } + if got.GetWorkerPodUid() != uid { + t.Errorf("worker pod UID = %q, want %q", got.GetWorkerPodUid(), uid) + } +} + // TestSyncer_SoftDelete_MarksDraining verifies that a pod entering Terminating // (DeletionTimestamp set) flips its worker to STATE_DRAINING without deleting the // worker record or touching the bound actor — the actor is still gracefully @@ -439,11 +558,11 @@ func TestSyncer_SoftDelete_MarksDraining(t *testing.T) { if err := s.workerInformer.GetIndexer().Add(deleting); err != nil { t.Fatalf("seed indexer: %v", err) } - if err := s.reconcile(ctx, workerKey{namespace: ns, pool: pool, name: pod}); err != nil { + if err := s.reconcile(ctx, workerKey{namespace: ns, name: pod}); err != nil { t.Fatalf("reconcile: %v", err) } - w, err := persistence.GetWorker(ctx, ns, pool, pod) + w, err := persistence.GetWorker(ctx, ns, pod) if err != nil { t.Fatalf("worker should still exist while draining: %v", err) } @@ -483,11 +602,11 @@ func TestSyncer_SoftDelete_NoPodIP(t *testing.T) { if err := s.workerInformer.GetIndexer().Add(deleting); err != nil { t.Fatalf("seed indexer: %v", err) } - if err := s.reconcile(ctx, workerKey{namespace: ns, pool: pool, name: pod}); err != nil { + if err := s.reconcile(ctx, workerKey{namespace: ns, name: pod}); err != nil { t.Fatalf("reconcile: %v", err) } - w, err := persistence.GetWorker(ctx, ns, pool, pod) + w, err := persistence.GetWorker(ctx, ns, pod) if err != nil { t.Fatalf("get worker: %v", err) } @@ -515,7 +634,7 @@ func TestMarkWorkerDraining(t *testing.T) { persistence, cleanup := storetest.SetupTestStore(t) defer cleanup() s := &WorkerPoolSyncer{persistence: persistence} - if err := s.markWorkerDraining(ctx, ns, pool, pod); err != nil { + if err := s.markWorkerDraining(ctx, workerKey{namespace: ns, name: pod}); err != nil { t.Errorf("markWorkerDraining on missing worker = %v, want nil", err) } }) @@ -527,7 +646,7 @@ func TestMarkWorkerDraining(t *testing.T) { t.Fatalf("create worker: %v", err) } s := &WorkerPoolSyncer{persistence: persistence} - if err := s.markWorkerDraining(ctx, ns, pool, pod); err != nil { + if err := s.markWorkerDraining(ctx, workerKey{namespace: ns, name: pod}); err != nil { t.Errorf("markWorkerDraining on already-draining worker = %v, want nil", err) } }) @@ -539,10 +658,10 @@ func TestMarkWorkerDraining(t *testing.T) { t.Fatalf("create worker: %v", err) } s := &WorkerPoolSyncer{persistence: persistence} - if err := s.markWorkerDraining(ctx, ns, pool, pod); err != nil { + if err := s.markWorkerDraining(ctx, workerKey{namespace: ns, name: pod}); err != nil { t.Fatalf("markWorkerDraining = %v, want nil", err) } - w, err := persistence.GetWorker(ctx, ns, pool, pod) + w, err := persistence.GetWorker(ctx, ns, pod) if err != nil { t.Fatalf("get worker: %v", err) } @@ -566,7 +685,7 @@ func TestReconcileDeadWorker(t *testing.T) { Metadata: &ateapipb.ResourceMetadata{Name: actorID, Atespace: atespace}, ActorTemplateNamespace: ns, ActorTemplateName: "tmpl", Status: ateapipb.Actor_STATUS_RUNNING, WorkerAssignment: &ateapipb.WorkerAssignment{ - WorkerNamespace: ns, WorkerPool: pool, WorkerPod: pod, WorkerPodIp: "10.0.0.5", + WorkerNamespace: ns, WorkerPool: pool, WorkerPod: pod, WorkerPodIp: "10.0.0.5", WorkerPodUid: "11111111-1111-1111-1111-111111111111", }, }) if err != nil { @@ -585,10 +704,10 @@ func TestReconcileDeadWorker(t *testing.T) { t.Fatalf("create worker: %v", err) } - if err := s.reconcileDeadWorker(ctx, ns, pool, pod); err != nil { + if err := s.reconcileDeadWorker(ctx, workerKey{namespace: ns, name: pod}); err != nil { t.Fatalf("reconcileDeadWorker = %v, want nil", err) } - if _, err := persistence.GetWorker(ctx, ns, pool, pod); !errors.Is(err, store.ErrNotFound) { + if _, err := persistence.GetWorker(ctx, ns, pod); !errors.Is(err, store.ErrNotFound) { t.Errorf("worker not deleted: err=%v", err) } got, err := persistence.GetActor(ctx, resources.ActorRef{Name: actorID, Atespace: atespace}) @@ -615,7 +734,7 @@ func TestReconcileDeadWorker_IgnoresStaleIncarnationAssignment(t *testing.T) { Metadata: &ateapipb.ResourceMetadata{Name: actorID, Atespace: atespace}, ActorTemplateNamespace: ns, ActorTemplateName: "tmpl", Status: ateapipb.Actor_STATUS_RUNNING, WorkerAssignment: &ateapipb.WorkerAssignment{ - WorkerNamespace: ns, WorkerPool: pool, WorkerPod: pod, WorkerPodIp: "10.0.0.5", + WorkerNamespace: ns, WorkerPool: pool, WorkerPod: pod, WorkerPodIp: "10.0.0.5", WorkerPodUid: "uid-rdw", }, }) if err != nil { @@ -634,11 +753,11 @@ func TestReconcileDeadWorker_IgnoresStaleIncarnationAssignment(t *testing.T) { t.Fatalf("create worker: %v", err) } - if err := s.reconcileDeadWorker(ctx, ns, pool, pod); err != nil { + if err := s.reconcileDeadWorker(ctx, workerKey{namespace: ns, name: pod}); err != nil { t.Fatalf("reconcileDeadWorker = %v, want nil", err) } // The dead worker should be deleted. - if _, err := persistence.GetWorker(ctx, ns, pool, pod); !errors.Is(err, store.ErrNotFound) { + if _, err := persistence.GetWorker(ctx, ns, pod); !errors.Is(err, store.ErrNotFound) { t.Errorf("worker not deleted: err=%v", err) } // Because ActorUid did not match, the new actor must remain RUNNING. @@ -699,7 +818,7 @@ func TestSyncer_ReconcileOrphanedWorkers(t *testing.T) { Metadata: &ateapipb.ResourceMetadata{Name: actorID, Atespace: atespace}, ActorTemplateNamespace: ns, ActorTemplateName: "tmpl", Status: ateapipb.Actor_STATUS_RUNNING, WorkerAssignment: &ateapipb.WorkerAssignment{ - WorkerNamespace: ns, WorkerPool: pool, WorkerPod: "worker-orphan", WorkerPodIp: "10.0.0.10", + WorkerNamespace: ns, WorkerPool: pool, WorkerPod: "worker-orphan", WorkerPodIp: "10.0.0.10", WorkerPodUid: "22222222-2222-2222-2222-222222222222", }, }) if err != nil { @@ -725,10 +844,10 @@ func TestSyncer_ReconcileOrphanedWorkers(t *testing.T) { s.processNextWorkItem(ctx) } - if _, err := persistence.GetWorker(ctx, ns, pool, "worker-orphan"); !errors.Is(err, store.ErrNotFound) { + if _, err := persistence.GetWorker(ctx, ns, "worker-orphan"); !errors.Is(err, store.ErrNotFound) { t.Errorf("orphan worker not removed: err=%v", err) } - if _, err := persistence.GetWorker(ctx, ns, pool, "worker-live"); err != nil { + if _, err := persistence.GetWorker(ctx, ns, "worker-live"); err != nil { t.Errorf("live worker was wrongly removed: %v", err) } got, err := persistence.GetActor(ctx, resources.ActorRef{Name: actorID, Atespace: atespace}) @@ -782,7 +901,7 @@ func TestReleaseActorOnDeadWorker_StatusTransitions(t *testing.T) { ActorTemplateName: "tmpl", Status: tc.start, WorkerAssignment: &ateapipb.WorkerAssignment{ - WorkerNamespace: ns, WorkerPool: pool, WorkerPod: pod, WorkerPodIp: ip, WorkerPodUid: "uid", + WorkerNamespace: ns, WorkerPool: pool, WorkerPod: pod, WorkerPodIp: ip, WorkerPodUid: "08675309-4a65-6e6e-7973-6e756d626572", }, }) if err != nil { @@ -802,7 +921,11 @@ func TestReleaseActorOnDeadWorker_StatusTransitions(t *testing.T) { t.Fatalf("create worker: %v", err) } - if err := s.releaseActorOnDeadWorker(ctx, ns, pool, pod); err != nil { + stored, err := persistence.GetWorker(ctx, ns, pod) + if err != nil { + t.Fatalf("get worker: %v", err) + } + if err := s.releaseActorOnDeadWorker(ctx, stored); err != nil { t.Fatalf("releaseActorOnDeadWorker: %v", err) } @@ -893,7 +1016,7 @@ func TestSyncer_UpdateWorker_RetryOnVersionConflict(t *testing.T) { } err = wait.PollUntilContextTimeout(context.Background(), 100*time.Millisecond, 5*time.Second, true, func(ctx context.Context) (bool, error) { - w, err := persistence.GetWorker(ctx, ns, poolName, podName) + w, err := persistence.GetWorker(ctx, ns, podName) if err != nil { if errors.Is(err, store.ErrNotFound) { return false, nil @@ -930,7 +1053,7 @@ func TestSyncer_UpdateWorker_RetryOnVersionConflict(t *testing.T) { // Configure conflictStore to inject a concurrent version bump in Redis when the syncer calls UpdateWorker. cs.onUpdate = func(c context.Context, w *ateapipb.Worker) { - if cw, err := cs.Interface.GetWorker(c, ns, poolName, podName); err == nil { + if cw, err := cs.Interface.GetWorker(c, ns, podName); err == nil { cw.NodeName = "node2" _ = cs.Interface.UpdateWorker(c, cw, cw.Version) } @@ -953,7 +1076,7 @@ func TestSyncer_UpdateWorker_RetryOnVersionConflict(t *testing.T) { // Verify that the worker in Redis eventually gets updated to the new SandboxClass despite the version conflict. err = wait.PollUntilContextTimeout(context.Background(), 100*time.Millisecond, 5*time.Second, true, func(ctx context.Context) (bool, error) { - w, err := persistence.GetWorker(ctx, ns, poolName, podName) + w, err := persistence.GetWorker(ctx, ns, podName) if err != nil { if errors.Is(err, store.ErrNotFound) { return false, nil @@ -1008,7 +1131,7 @@ func TestSyncer_RequeueOnMissingWorkerPool(t *testing.T) { } // Wait until the syncer has attempted to reconcile the pod and requeued it due to the missing pool. - key := workerKey{namespace: ns, pool: poolName, name: podName} + key := workerKey{namespace: ns, name: podName} err := wait.PollUntilContextTimeout(context.Background(), 10*time.Millisecond, 5*time.Second, true, func(ctx context.Context) (bool, error) { return syncer.queue.NumRequeues(key) > 0, nil }) @@ -1030,7 +1153,7 @@ func TestSyncer_RequeueOnMissingWorkerPool(t *testing.T) { } err = wait.PollUntilContextTimeout(context.Background(), 100*time.Millisecond, 5*time.Second, true, func(ctx context.Context) (bool, error) { - w, err := persistence.GetWorker(ctx, ns, poolName, podName) + w, err := persistence.GetWorker(ctx, ns, podName) if err != nil { if errors.Is(err, store.ErrNotFound) { return false, nil @@ -1095,7 +1218,7 @@ func TestSyncer_SoftDelete_ViaInformer(t *testing.T) { } if err := wait.PollUntilContextTimeout(context.Background(), 50*time.Millisecond, 2*time.Second, true, func(ctx context.Context) (bool, error) { - _, err := persistence.GetWorker(ctx, ns, poolName, podName) + _, err := persistence.GetWorker(ctx, ns, podName) return err == nil, nil }); err != nil { t.Fatalf("worker row not materialised: %v", err) @@ -1111,7 +1234,7 @@ func TestSyncer_SoftDelete_ViaInformer(t *testing.T) { } if err := wait.PollUntilContextTimeout(context.Background(), 100*time.Millisecond, 2*time.Second, true, func(ctx context.Context) (bool, error) { - w, err := persistence.GetWorker(ctx, ns, poolName, podName) + w, err := persistence.GetWorker(ctx, ns, podName) if err != nil { return false, err } @@ -1177,7 +1300,7 @@ func TestSyncer_PodRecreatedWithNewUID(t *testing.T) { t.Fatalf("failed to create pod: %v", err) } if err := wait.PollUntilContextTimeout(context.Background(), 50*time.Millisecond, 2*time.Second, true, func(ctx context.Context) (bool, error) { - w, err := persistence.GetWorker(ctx, ns, poolName, podName) + w, err := persistence.GetWorker(ctx, ns, podName) if err != nil { if errors.Is(err, store.ErrNotFound) { return false, nil @@ -1196,7 +1319,7 @@ func TestSyncer_PodRecreatedWithNewUID(t *testing.T) { } if err := wait.PollUntilContextTimeout(context.Background(), 100*time.Millisecond, 5*time.Second, true, func(ctx context.Context) (bool, error) { - w, err := persistence.GetWorker(ctx, ns, poolName, podName) + w, err := persistence.GetWorker(ctx, ns, podName) if err != nil { if errors.Is(err, store.ErrNotFound) { return false, nil @@ -1335,11 +1458,11 @@ func TestSyncer_InvalidWorkerIsTerminal(t *testing.T) { t.Fatalf("failed to seed indexer: %v", err) } - key := workerKey{namespace: ns, pool: poolName, name: podName} + key := workerKey{namespace: ns, name: podName} if err := s.reconcile(ctx, key); err != nil { t.Fatalf("reconcile returned error for invalid worker (should be terminal): %v", err) } - if _, err := persistence.GetWorker(ctx, ns, poolName, podName); !errors.Is(err, store.ErrNotFound) { + if _, err := persistence.GetWorker(ctx, ns, podName); !errors.Is(err, store.ErrNotFound) { t.Fatalf("expected no worker row for invalid worker, got err=%v", err) } } diff --git a/cmd/ateapi/internal/controlapi/volumes.go b/cmd/ateapi/internal/controlapi/volumes.go index 10146fda5..4213b268d 100644 --- a/cmd/ateapi/internal/controlapi/volumes.go +++ b/cmd/ateapi/internal/controlapi/volumes.go @@ -195,7 +195,7 @@ func detachActorVolumes(ctx context.Context, st store.Interface, registry Volume return nil } - worker, err := st.GetWorker(ctx, assignment.GetWorkerNamespace(), assignment.GetWorkerPool(), assignment.GetWorkerPod()) + worker, err := st.GetWorker(ctx, assignment.GetWorkerNamespace(), assignment.GetWorkerPod()) if err != nil { if errors.Is(err, store.ErrNotFound) { slog.WarnContext(ctx, fmt.Sprintf("Worker not found in store during %s, skipping detach volumes", action), slog.String("actor_id", actor.GetMetadata().GetName())) diff --git a/cmd/ateapi/internal/controlapi/workflow_pause.go b/cmd/ateapi/internal/controlapi/workflow_pause.go index b088233b5..bff1f3c5f 100644 --- a/cmd/ateapi/internal/controlapi/workflow_pause.go +++ b/cmd/ateapi/internal/controlapi/workflow_pause.go @@ -215,7 +215,7 @@ func (w *ActorWorkflow) ensurePausedFinalized(ctx context.Context, actorRef reso // 1. Free the worker (if it hasn't been freed yet) if assignment := latestActor.GetWorkerAssignment(); assignment != nil { - worker, err := w.store.GetWorker(ctx, assignment.GetWorkerNamespace(), assignment.GetWorkerPool(), assignment.GetWorkerPod()) + worker, err := w.store.GetWorker(ctx, assignment.GetWorkerNamespace(), assignment.GetWorkerPod()) nodeName := "" if err != nil { if !errors.Is(err, store.ErrNotFound) { diff --git a/cmd/ateapi/internal/controlapi/workflow_resume.go b/cmd/ateapi/internal/controlapi/workflow_resume.go index bdf999a3f..5caafc84b 100644 --- a/cmd/ateapi/internal/controlapi/workflow_resume.go +++ b/cmd/ateapi/internal/controlapi/workflow_resume.go @@ -350,7 +350,7 @@ func (w *ActorWorkflow) validateAssignedWorker(ctx context.Context, actorRef res return nil, status.Errorf(codes.Aborted, "actor %s crashed", actorRef) } - worker, err := w.store.GetWorker(ctx, assignment.GetWorkerNamespace(), assignment.GetWorkerPool(), assignment.GetWorkerPod()) + worker, err := w.store.GetWorker(ctx, assignment.GetWorkerNamespace(), assignment.GetWorkerPod()) if err != nil { // Crash the actor if it was assigned to a deleted pod. if errors.Is(err, store.ErrNotFound) { diff --git a/cmd/ateapi/internal/controlapi/workflow_resume_test.go b/cmd/ateapi/internal/controlapi/workflow_resume_test.go index b9b1274b4..5380a0779 100644 --- a/cmd/ateapi/internal/controlapi/workflow_resume_test.go +++ b/cmd/ateapi/internal/controlapi/workflow_resume_test.go @@ -102,7 +102,7 @@ func TestAssignWorkerAttempt_SkipsWorkerAssignedInOtherAtespace(t *testing.T) { t.Fatalf("assignWorkerAttempt() error = %v, want FailedPrecondition (no free workers)", err) } - stored, err := persistence.GetWorker(ctx, "worker-ns", "pool", "pod-1") + stored, err := persistence.GetWorker(ctx, "worker-ns", "pod-1") if err != nil { t.Fatalf("GetWorker: %v", err) } @@ -180,7 +180,7 @@ func TestAssignWorkerAttempt_ReleasesIneligibleStaleWorkerInBackground(t *testin // assignment is cleared. deadline := time.Now().Add(5 * time.Second) for { - stored, err := persistence.GetWorker(ctx, "worker-ns", "pool-a", "stale-pod") + stored, err := persistence.GetWorker(ctx, "worker-ns", "stale-pod") if err != nil { t.Fatalf("GetWorker: %v", err) } @@ -224,7 +224,7 @@ func TestAssignWorkerAttempt_RetryAfterConflictPicksFreshWorker(t *testing.T) { } // Snapshot the contested worker at the version the failed attempt saw. - beforeClaim, err := persistence.GetWorker(ctx, "worker-ns", "pool", "contested-pod") + beforeClaim, err := persistence.GetWorker(ctx, "worker-ns", "contested-pod") if err != nil { t.Fatalf("GetWorker: %v", err) } @@ -267,14 +267,14 @@ func TestAssignWorkerAttempt_RetryAfterConflictPicksFreshWorker(t *testing.T) { t.Errorf("assigned worker = %q, want %q", got, "fallback-pod") } - storedContested, err := persistence.GetWorker(ctx, "worker-ns", "pool", "contested-pod") + storedContested, err := persistence.GetWorker(ctx, "worker-ns", "contested-pod") if err != nil { t.Fatalf("GetWorker(contested-pod): %v", err) } if got := storedContested.GetAssignment().GetActorUid(); got != "other-actor-uid" { t.Errorf("contested worker assignment = %v, want to remain with actor %q", storedContested.GetAssignment(), "other-actor-uid") } - storedFallback, err := persistence.GetWorker(ctx, "worker-ns", "pool", "fallback-pod") + storedFallback, err := persistence.GetWorker(ctx, "worker-ns", "fallback-pod") if err != nil { t.Fatalf("GetWorker(fallback-pod): %v", err) } @@ -704,7 +704,7 @@ func TestValidateAssignedWorker_WorkerOwnership(t *testing.T) { } // Fetch the stored version so the no-write assertion below can // detect any optimistic update. - seeded, err := persistence.GetWorker(ctx, "worker-ns", "pool", "pod-1") + seeded, err := persistence.GetWorker(ctx, "worker-ns", "pod-1") if err != nil { t.Fatalf("GetWorker: %v", err) } @@ -735,7 +735,7 @@ func TestValidateAssignedWorker_WorkerOwnership(t *testing.T) { t.Errorf("stored actor status = %v, want %v", actor.GetStatus(), tt.wantActorStatus) } - stored, err := persistence.GetWorker(ctx, "worker-ns", "pool", "pod-1") + stored, err := persistence.GetWorker(ctx, "worker-ns", "pod-1") if err != nil { t.Fatalf("GetWorker: %v", err) } diff --git a/cmd/ateapi/internal/controlapi/workflow_suspend.go b/cmd/ateapi/internal/controlapi/workflow_suspend.go index 942b00473..55bc56727 100644 --- a/cmd/ateapi/internal/controlapi/workflow_suspend.go +++ b/cmd/ateapi/internal/controlapi/workflow_suspend.go @@ -346,7 +346,7 @@ func (w *ActorWorkflow) ensureSuspendedFinalized(ctx context.Context, actorRef r if assignment := latestActor.GetWorkerAssignment(); assignment != nil { workerPod := assignment.GetWorkerPod() - worker, err := w.store.GetWorker(ctx, assignment.GetWorkerNamespace(), assignment.GetWorkerPool(), workerPod) + worker, err := w.store.GetWorker(ctx, assignment.GetWorkerNamespace(), workerPod) if err != nil { if !errors.Is(err, store.ErrNotFound) { return nil, fmt.Errorf("while getting worker for release: %w", err) diff --git a/cmd/ateapi/internal/controlapi/workflow_suspend_test.go b/cmd/ateapi/internal/controlapi/workflow_suspend_test.go index 38c42c318..16b6fa72e 100644 --- a/cmd/ateapi/internal/controlapi/workflow_suspend_test.go +++ b/cmd/ateapi/internal/controlapi/workflow_suspend_test.go @@ -434,7 +434,7 @@ func TestEnsureSuspendedFinalized_ReleasesOnlyOwnWorker(t *testing.T) { t.Fatalf("ensureSuspendedFinalized: %v", err) } - stored, err := persistence.GetWorker(ctx, "worker-ns", "pool", "pod-1") + stored, err := persistence.GetWorker(ctx, "worker-ns", "pod-1") if err != nil { t.Fatalf("GetWorker: %v", err) } diff --git a/cmd/ateapi/internal/store/atepg/atepg.go b/cmd/ateapi/internal/store/atepg/atepg.go index ead0b22f8..1f1154326 100644 --- a/cmd/ateapi/internal/store/atepg/atepg.go +++ b/cmd/ateapi/internal/store/atepg/atepg.go @@ -1333,15 +1333,15 @@ func (p *Persistence) CreateWorker(ctx context.Context, worker *ateapipb.Worker) return nil } -func getWorkerRow(ctx context.Context, q querier, namespace, poolName, pod string) (*ateapipb.Worker, error) { +func getWorkerRow(ctx context.Context, q querier, namespace, pod string) (*ateapipb.Worker, error) { var protoBytes []byte - err := q.QueryRow(ctx, `SELECT proto FROM workers WHERE worker_namespace = $1 AND worker_pool = $2 AND worker_pod = $3`, - namespace, poolName, pod).Scan(&protoBytes) + err := q.QueryRow(ctx, `SELECT proto FROM workers WHERE worker_namespace = $1 AND worker_pod = $2`, + namespace, pod).Scan(&protoBytes) if err != nil { if errors.Is(err, pgx.ErrNoRows) { return nil, store.ErrNotFound } - return nil, fmt.Errorf("getting worker %s/%s/%s: %w", namespace, poolName, pod, err) + return nil, fmt.Errorf("getting worker %s/%s: %w", namespace, pod, err) } out := &ateapipb.Worker{} if err := proto.Unmarshal(protoBytes, out); err != nil { @@ -1350,12 +1350,12 @@ func getWorkerRow(ctx context.Context, q querier, namespace, poolName, pod strin return out, nil } -func (p *Persistence) GetWorker(ctx context.Context, namespace, poolName, pod string) (*ateapipb.Worker, error) { - return getWorkerRow(ctx, p.pool, namespace, poolName, pod) +func (p *Persistence) GetWorker(ctx context.Context, namespace, pod string) (*ateapipb.Worker, error) { + return getWorkerRow(ctx, p.pool, namespace, pod) } func (p *Persistence) UpdateWorker(ctx context.Context, worker *ateapipb.Worker, expectedVersion int64) error { - namespace, poolName, pod := worker.GetWorkerNamespace(), worker.GetWorkerPool(), worker.GetWorkerPod() + namespace, pod := worker.GetWorkerNamespace(), worker.GetWorkerPod() dbWorker := proto.Clone(worker).(*ateapipb.Worker) dbWorker.Version = expectedVersion + 1 @@ -1370,43 +1370,62 @@ func (p *Persistence) UpdateWorker(ctx context.Context, worker *ateapipb.Worker, err := tx.QueryRow(ctx, ` UPDATE workers SET version = $1, proto = $2 - WHERE worker_namespace = $3 AND worker_pool = $4 AND worker_pod = $5 + WHERE worker_namespace = $3 AND worker_pod = $4 + AND worker_pool = $5 AND version = $6 RETURNING proto`, - dbWorker.GetVersion(), protoBytes, namespace, poolName, pod, expectedVersion, + dbWorker.GetVersion(), protoBytes, namespace, pod, dbWorker.GetWorkerPool(), expectedVersion, ).Scan(&returned) if err == nil { return true, nil } if !errors.Is(err, pgx.ErrNoRows) { - return false, fmt.Errorf("updating worker %s/%s/%s: %w", namespace, poolName, pod, err) + return false, fmt.Errorf("updating worker %s/%s: %w", namespace, pod, err) } - current, getErr := getWorkerRow(ctx, tx, namespace, poolName, pod) + current, getErr := getWorkerRow(ctx, tx, namespace, pod) if getErr != nil { return false, getErr } + // The column is never rewritten, so letting the pool change here would + // drift it away from the stored proto. + if current.GetWorkerPool() != dbWorker.GetWorkerPool() { + return false, fmt.Errorf("worker_pool is immutable: mutation changed it from %q to %q", current.GetWorkerPool(), dbWorker.GetWorkerPool()) + } if current.GetVersion() != expectedVersion { return false, store.ErrVersionConflict } - return false, fmt.Errorf("update worker %s/%s/%s: no row matched but current state is otherwise consistent", namespace, poolName, pod) + return false, fmt.Errorf("update worker %s/%s: no row matched but current state is otherwise consistent", namespace, pod) }) } -func (p *Persistence) DeleteWorker(ctx context.Context, namespace, poolName, pod string) error { +func (p *Persistence) DeleteWorker(ctx context.Context, expected *ateapipb.Worker, expectedVersion int64) error { + namespace, pod := expected.GetWorkerNamespace(), expected.GetWorkerPod() deletedEvent := &ateapipb.Worker{WorkerNamespace: namespace, WorkerPod: pod} return p.writeAndNotify(ctx, store.WorkerEventDeleted, deletedEvent, func(ctx context.Context, tx pgx.Tx) (bool, error) { + // The row lock holds the record still between the check and the delete. var protoBytes []byte - err := tx.QueryRow(ctx, ` - DELETE FROM workers - WHERE worker_namespace = $1 AND worker_pool = $2 AND worker_pod = $3 - RETURNING proto`, namespace, poolName, pod).Scan(&protoBytes) + err := tx.QueryRow(ctx, + `SELECT proto FROM workers WHERE worker_namespace = $1 AND worker_pod = $2 FOR UPDATE`, + namespace, pod).Scan(&protoBytes) if err != nil { if errors.Is(err, pgx.ErrNoRows) { // Idempotent: nothing existed, so nothing to notify either. return false, nil } - return false, fmt.Errorf("deleting worker %s/%s/%s: %w", namespace, poolName, pod, err) + return false, fmt.Errorf("locking worker %s/%s: %w", namespace, pod, err) + } + current := &ateapipb.Worker{} + if err := proto.Unmarshal(protoBytes, current); err != nil { + return false, fmt.Errorf("unmarshaling worker: %w", err) + } + if !store.DeletePreconditionMet(current, expected, expectedVersion) { + return false, store.ErrVersionConflict + } + if _, err := tx.Exec(ctx, + `DELETE FROM workers WHERE worker_namespace = $1 AND worker_pod = $2`, + namespace, pod); err != nil { + return false, fmt.Errorf("deleting worker %s/%s: %w", namespace, pod, err) } return true, nil }) @@ -1414,32 +1433,32 @@ func (p *Persistence) DeleteWorker(ctx context.Context, namespace, poolName, pod func (p *Persistence) ListWorkers(ctx context.Context, opts store.ListOptions) (store.ListResponse[*ateapipb.Worker], error) { pageSize, pageTokenStr := opts.PageSize, opts.PageToken - token, err := decodePageToken(pageTokenStr, kindWorker, "", 3) + token, err := decodePageToken(pageTokenStr, kindWorker, "", 2) if err != nil { return store.ListResponse[*ateapipb.Worker]{}, err } - var lastNS, lastPool, lastPod *string - if len(token.Last) == 3 { - lastNS, lastPool, lastPod = &token.Last[0], &token.Last[1], &token.Last[2] + var lastNS, lastPod *string + if len(token.Last) == 2 { + lastNS, lastPod = &token.Last[0], &token.Last[1] } rows, err := p.pool.Query(ctx, ` - SELECT worker_namespace, worker_pool, worker_pod, proto FROM workers - WHERE $1::text IS NULL OR (worker_namespace, worker_pool, worker_pod) > ($1, $2, $3) - ORDER BY worker_namespace, worker_pool, worker_pod - LIMIT $4`, lastNS, lastPool, lastPod, int64(pageSize)+1) + SELECT worker_namespace, worker_pod, proto FROM workers + WHERE $1::text IS NULL OR (worker_namespace, worker_pod) > ($1, $2) + ORDER BY worker_namespace, worker_pod + LIMIT $3`, lastNS, lastPod, int64(pageSize)+1) if err != nil { return store.ListResponse[*ateapipb.Worker]{}, fmt.Errorf("listing workers: %w", err) } defer rows.Close() - type key struct{ namespace, pool, pod string } + type key struct{ namespace, pod string } var keys []key var result []*ateapipb.Worker for rows.Next() { var k key var protoBytes []byte - if err := rows.Scan(&k.namespace, &k.pool, &k.pod, &protoBytes); err != nil { + if err := rows.Scan(&k.namespace, &k.pod, &protoBytes); err != nil { return store.ListResponse[*ateapipb.Worker]{}, fmt.Errorf("scanning worker row: %w", err) } w := &ateapipb.Worker{} @@ -1457,7 +1476,7 @@ func (p *Persistence) ListWorkers(ctx context.Context, opts store.ListOptions) ( if len(result) > int(pageSize) { result = result[:pageSize] last := keys[pageSize-1] - nextToken = encodePageToken(kindWorker, "", []string{last.namespace, last.pool, last.pod}) + nextToken = encodePageToken(kindWorker, "", []string{last.namespace, last.pod}) } return store.ListResponse[*ateapipb.Worker]{Items: result, NextPageToken: nextToken}, nil } diff --git a/cmd/ateapi/internal/store/atepg/schema.go b/cmd/ateapi/internal/store/atepg/schema.go index f4f0ce0ad..9a68abf64 100644 --- a/cmd/ateapi/internal/store/atepg/schema.go +++ b/cmd/ateapi/internal/store/atepg/schema.go @@ -92,7 +92,7 @@ CREATE TABLE IF NOT EXISTS workers ( worker_pod text NOT NULL, version bigint NOT NULL, proto bytea NOT NULL, - PRIMARY KEY (worker_namespace, worker_pool, worker_pod) + PRIMARY KEY (worker_namespace, worker_pod) ); CREATE TABLE IF NOT EXISTS leases ( diff --git a/cmd/ateapi/internal/store/ateredis/ateredis.go b/cmd/ateapi/internal/store/ateredis/ateredis.go index a27962754..9423f4a6a 100644 --- a/cmd/ateapi/internal/store/ateredis/ateredis.go +++ b/cmd/ateapi/internal/store/ateredis/ateredis.go @@ -20,7 +20,7 @@ // Redis lua. // // Workers are stored in keys of the form -// `worker:::`, holding a DBWorker JSON object. +// `worker::`, holding a DBWorker JSON object. // // Note that redis lua scripting has a restriction that informed the data design // here -- a lua script must predeclare all keys it is going to access. It @@ -598,8 +598,8 @@ func (s *Persistence) DeleteActorTemplateVersion(ctx context.Context, versionRef return deleted, nil } -func workerDBKey(namespace, poolName, podName string) string { - return "worker:" + namespace + ":" + poolName + ":" + podName +func workerDBKey(namespace, podName string) string { + return "worker:" + namespace + ":" + podName } func marshalWorkerEvent(eventType store.WorkerEventType, worker *ateapipb.Worker) (string, error) { @@ -973,7 +973,7 @@ func (s *Persistence) DeleteActorSnapshotTag(ctx context.Context, atespace, name } func (s *Persistence) CreateWorker(ctx context.Context, worker *ateapipb.Worker) error { - dbKey := workerDBKey(worker.GetWorkerNamespace(), worker.GetWorkerPool(), worker.GetWorkerPod()) + dbKey := workerDBKey(worker.GetWorkerNamespace(), worker.GetWorkerPod()) // Clone because we will update the version field, and we don't want to // stomp the caller's copy. @@ -997,8 +997,8 @@ func (s *Persistence) CreateWorker(ctx context.Context, worker *ateapipb.Worker) return nil } -func (s *Persistence) GetWorker(ctx context.Context, namespace, pool, pod string) (*ateapipb.Worker, error) { - dbKey := workerDBKey(namespace, pool, pod) +func (s *Persistence) GetWorker(ctx context.Context, namespace, pod string) (*ateapipb.Worker, error) { + dbKey := workerDBKey(namespace, pod) dbWorkerBytes, err := s.rdb.Get(ctx, dbKey).Bytes() if err != nil { @@ -1013,15 +1013,15 @@ func (s *Persistence) GetWorker(ctx context.Context, namespace, pool, pod string return nil, fmt.Errorf("in protojson.Unmarshal: %w", err) } - if worker.GetWorkerNamespace() != namespace || worker.GetWorkerPool() != pool || worker.GetWorkerPod() != pod { - return nil, fmt.Errorf("(impossible) mismatch between stored namespace/pool/pod and key") + if worker.GetWorkerNamespace() != namespace || worker.GetWorkerPod() != pod { + return nil, fmt.Errorf("(impossible) mismatch between stored namespace/pod and key") } return worker, nil } func (s *Persistence) UpdateWorker(ctx context.Context, worker *ateapipb.Worker, expectedVersion int64) error { - dbKey := workerDBKey(worker.GetWorkerNamespace(), worker.GetWorkerPool(), worker.GetWorkerPod()) + dbKey := workerDBKey(worker.GetWorkerNamespace(), worker.GetWorkerPod()) // Clone because we will update the version field, and we don't want to // stomp the caller's copy. @@ -1082,12 +1082,43 @@ func (s *Persistence) UpdateWorker(ctx context.Context, worker *ateapipb.Worker, return nil } -func (s *Persistence) DeleteWorker(ctx context.Context, namespace, pool, pod string) error { - dbKey := workerDBKey(namespace, pool, pod) - err := s.rdb.Del(ctx, dbKey).Err() +func (s *Persistence) DeleteWorker(ctx context.Context, expected *ateapipb.Worker, expectedVersion int64) error { + namespace, pod := expected.GetWorkerNamespace(), expected.GetWorkerPod() + dbKey := workerDBKey(namespace, pod) + deleted := false + err := s.rdb.Watch(ctx, func(tx *redis.Tx) error { + currentVal, err := tx.Get(ctx, dbKey).Bytes() + if err != nil { + if errors.Is(err, redis.Nil) { + return nil + } + return fmt.Errorf("while getting worker: %w", err) + } + currentWorker := &ateapipb.Worker{} + if err := protojson.Unmarshal(currentVal, currentWorker); err != nil { + return fmt.Errorf("in protojson.Unmarshal: %w", err) + } + if !store.DeletePreconditionMet(currentWorker, expected, expectedVersion) { + return store.ErrVersionConflict + } + if _, err := tx.TxPipelined(ctx, func(pipe redis.Pipeliner) error { + pipe.Del(ctx, dbKey) + return nil + }); err != nil { + return err + } + deleted = true + return nil + }, dbKey) if err != nil { + if errors.Is(err, store.ErrVersionConflict) || errors.Is(err, redis.TxFailedErr) { + return store.ErrVersionConflict + } return fmt.Errorf("while deleting worker key %q: %w", dbKey, err) } + if !deleted { + return nil + } s.publishWorkerEvent(ctx, store.WorkerEventDeleted, &ateapipb.Worker{ WorkerNamespace: namespace, WorkerPod: pod, diff --git a/cmd/ateapi/internal/store/ateredis/ateredis_test.go b/cmd/ateapi/internal/store/ateredis/ateredis_test.go index 23b9a2536..59690bab3 100644 --- a/cmd/ateapi/internal/store/ateredis/ateredis_test.go +++ b/cmd/ateapi/internal/store/ateredis/ateredis_test.go @@ -520,7 +520,7 @@ func TestUpdateWorker_NotFound(t *testing.T) { func TestGetWorker_NotFound(t *testing.T) { _, s, ctx := setupTest(t) - _, err := s.GetWorker(ctx, "default", "pool-1", "non-existent") + _, err := s.GetWorker(ctx, "default", "non-existent") if !errors.Is(err, store.ErrNotFound) { t.Errorf("expected ErrNotFound, got %v", err) } @@ -545,7 +545,7 @@ func TestCreateWorker_Success(t *testing.T) { t.Fatalf("CreateWorker failed: %v", err) } - got, err := s.GetWorker(ctx, "default", "pool-1", "pod-1") + got, err := s.GetWorker(ctx, "default", "pod-1") if err != nil { t.Fatalf("GetWorker failed: %v", err) } @@ -602,7 +602,7 @@ func TestUpdateWorker_Success(t *testing.T) { t.Fatalf("UpdateWorker failed: %v", err) } - got, err := s.GetWorker(ctx, "default", "pool-1", "pod-1") + got, err := s.GetWorker(ctx, "default", "pod-1") if err != nil { t.Fatalf("GetWorker failed: %v", err) } @@ -644,11 +644,15 @@ func TestDeleteWorker(t *testing.T) { t.Fatalf("WatchWorkers failed: %v", err) } - if err := s.DeleteWorker(ctx, "default", "pool-1", "pod-1"); err != nil { + stored, err := s.GetWorker(ctx, "default", "pod-1") + if err != nil { + t.Fatalf("GetWorker failed: %v", err) + } + if err := s.DeleteWorker(ctx, stored, stored.GetVersion()); err != nil { t.Fatalf("DeleteWorker failed: %v", err) } - _, err = s.GetWorker(ctx, "default", "pool-1", "pod-1") + _, err = s.GetWorker(ctx, "default", "pod-1") if !errors.Is(err, store.ErrNotFound) { t.Errorf("expected ErrNotFound after delete, got %v", err) } @@ -1209,13 +1213,13 @@ func TestUpdateWorker_Conflict(t *testing.T) { } // Fetch instance 1 - worker1, err := s.GetWorker(ctx, "default", "pool-1", "pod-1") + worker1, err := s.GetWorker(ctx, "default", "pod-1") if err != nil { t.Fatalf("GetWorker failed: %v", err) } // Fetch instance 2 - worker2, err := s.GetWorker(ctx, "default", "pool-1", "pod-1") + worker2, err := s.GetWorker(ctx, "default", "pod-1") if err != nil { t.Fatalf("GetWorker failed: %v", err) } diff --git a/cmd/ateapi/internal/store/store.go b/cmd/ateapi/internal/store/store.go index 54157bc14..1a21d0537 100644 --- a/cmd/ateapi/internal/store/store.go +++ b/cmd/ateapi/internal/store/store.go @@ -183,8 +183,9 @@ type Interface interface { // Registers a new idle worker. Returns ErrAlreadyExists if already registered. CreateWorker(ctx context.Context, worker *ateapipb.Worker) error - // Fetches worker state by namespace, pool, and pod name. Returns ErrNotFound if missing. - GetWorker(ctx context.Context, namespace, pool, pod string) (*ateapipb.Worker, error) + // Fetches worker state by the Pod backing it: the owning pool is an + // attribute, not part of a worker's identity. Returns ErrNotFound if missing. + GetWorker(ctx context.Context, namespace, pod string) (*ateapipb.Worker, error) // Lists workers. ListWorkers(ctx context.Context, opts ListOptions) (ListResponse[*ateapipb.Worker], error) @@ -192,8 +193,9 @@ type Interface interface { // Updates worker state with optimistic concurrency check. Returns ErrNotFound if missing, or ErrVersionConflict on version mismatch. UpdateWorker(ctx context.Context, worker *ateapipb.Worker, expectedVersion int64) error - // Removes a worker. Idempotent: does nothing if worker is not found. - DeleteWorker(ctx context.Context, namespace, pool, pod string) error + // Removes the worker the caller inspected. Returns nil if it is already + // gone and ErrVersionConflict if the stored record no longer matches. + DeleteWorker(ctx context.Context, expected *ateapipb.Worker, expectedVersion int64) error // WatchWorkers returns an active subscription to track worker state changes. // The watch's Events channel is closed when the caller calls Close, the @@ -319,6 +321,17 @@ func (l *Lock) Context() context.Context { return l.ctx } // times. func (l *Lock) Close() { l.once.Do(l.closeFn) } +// DeletePreconditionMet reports whether current is still the record the caller +// decided to delete. The version alone cannot say so: a delete-and-recreate +// restarts it at 1, so the Pod UID is what separates incarnations. +func DeletePreconditionMet(current, expected *ateapipb.Worker, expectedVersion int64) bool { + return current.GetVersion() == expectedVersion && + current.GetWorkerNamespace() == expected.GetWorkerNamespace() && + current.GetWorkerPod() == expected.GetWorkerPod() && + current.GetWorkerPool() == expected.GetWorkerPool() && + current.GetWorkerPodUid() == expected.GetWorkerPodUid() +} + // ListOptions carries the pagination parameters common to every List method. type ListOptions struct { // PageSize caps how many items a single call returns. diff --git a/cmd/ateapi/internal/store/storecontract/contract.go b/cmd/ateapi/internal/store/storecontract/contract.go index fe3855a8f..a7e2d645c 100644 --- a/cmd/ateapi/internal/store/storecontract/contract.go +++ b/cmd/ateapi/internal/store/storecontract/contract.go @@ -899,12 +899,55 @@ func runWorkerContractTests(t *testing.T, setup func(t *testing.T) store.Interfa s := setup(t) ctx := context.Background() - _, err := s.GetWorker(ctx, "default", "pool-1", "non-existent") + _, err := s.GetWorker(ctx, "default", "non-existent") if !errors.Is(err, store.ErrNotFound) { t.Errorf("expected ErrNotFound, got %v", err) } }) + // Identity is the Pod; the owning pool is an attribute, not a key component. + t.Run("WorkerIdentityIsNamespaceAndPod", func(t *testing.T) { + s := setup(t) + ctx := context.Background() + + if err := s.CreateWorker(ctx, &ateapipb.Worker{ + WorkerNamespace: "default", + WorkerPool: "pool-1", + WorkerPod: "pod-1", + }); err != nil { + t.Fatalf("CreateWorker failed: %v", err) + } + + err := s.CreateWorker(ctx, &ateapipb.Worker{ + WorkerNamespace: "default", + WorkerPool: "pool-2", + WorkerPod: "pod-1", + }) + if !errors.Is(err, store.ErrAlreadyExists) { + t.Errorf("CreateWorker for the same Pod under another pool = %v, want ErrAlreadyExists", err) + } + + got, err := s.GetWorker(ctx, "default", "pod-1") + if err != nil { + t.Fatalf("GetWorker failed: %v", err) + } + if got.GetWorkerPool() != "pool-1" { + t.Errorf("GetWorker returned pool %q, want the originally created pool-1", got.GetWorkerPool()) + } + + got.WorkerPool = "pool-2" + if err := s.UpdateWorker(ctx, got, got.GetVersion()); err == nil { + t.Error("UpdateWorker changing the pool succeeded, want it rejected as immutable") + } + after, err := s.GetWorker(ctx, "default", "pod-1") + if err != nil { + t.Fatalf("GetWorker after rejected update failed: %v", err) + } + if after.GetWorkerPool() != "pool-1" { + t.Errorf("pool after rejected update = %q, want pool-1", after.GetWorkerPool()) + } + }) + t.Run("CreateWorker_Success", func(t *testing.T) { s := setup(t) ctx := context.Background() @@ -925,7 +968,7 @@ func runWorkerContractTests(t *testing.T, setup func(t *testing.T) store.Interfa t.Fatalf("CreateWorker failed: %v", err) } - got, err := s.GetWorker(ctx, "default", "pool-1", "pod-1") + got, err := s.GetWorker(ctx, "default", "pod-1") if err != nil { t.Fatalf("GetWorker failed: %v", err) } @@ -984,7 +1027,7 @@ func runWorkerContractTests(t *testing.T, setup func(t *testing.T) store.Interfa t.Fatalf("UpdateWorker failed: %v", err) } - got, err := s.GetWorker(ctx, "default", "pool-1", "pod-1") + got, err := s.GetWorker(ctx, "default", "pod-1") if err != nil { t.Fatalf("GetWorker failed: %v", err) } @@ -1015,11 +1058,11 @@ func runWorkerContractTests(t *testing.T, setup func(t *testing.T) store.Interfa t.Fatalf("CreateWorker failed: %v", err) } - worker1, err := s.GetWorker(ctx, "default", "pool-1", "pod-1") + worker1, err := s.GetWorker(ctx, "default", "pod-1") if err != nil { t.Fatalf("GetWorker failed: %v", err) } - worker2, err := s.GetWorker(ctx, "default", "pool-1", "pod-1") + worker2, err := s.GetWorker(ctx, "default", "pod-1") if err != nil { t.Fatalf("GetWorker failed: %v", err) } @@ -1051,10 +1094,10 @@ func runWorkerContractTests(t *testing.T, setup func(t *testing.T) store.Interfa } defer watch.Close() - if err := s.DeleteWorker(ctx, "default", "pool-1", "pod-1"); err != nil { + if err := s.DeleteWorker(ctx, worker, 1); err != nil { t.Fatalf("DeleteWorker failed: %v", err) } - if _, err := s.GetWorker(ctx, "default", "pool-1", "pod-1"); !errors.Is(err, store.ErrNotFound) { + if _, err := s.GetWorker(ctx, "default", "pod-1"); !errors.Is(err, store.ErrNotFound) { t.Errorf("expected ErrNotFound after delete, got %v", err) } @@ -1068,11 +1111,90 @@ func runWorkerContractTests(t *testing.T, setup func(t *testing.T) store.Interfa s := setup(t) ctx := context.Background() - if err := s.DeleteWorker(ctx, "default", "pool-1", "non-existent"); err != nil { + if err := s.DeleteWorker(ctx, &ateapipb.Worker{WorkerNamespace: "default", WorkerPod: "non-existent"}, 1); err != nil { t.Errorf("DeleteWorker of a missing worker should be a no-op, got %v", err) } }) + // The syncer decides to delete from a record it has already read; the + // version pins that decision to the incarnation it inspected, so a + // concurrent replacement cannot be deleted by the stale decision. + t.Run("DeleteWorker_VersionConflict", func(t *testing.T) { + s := setup(t) + ctx := context.Background() + + worker := &ateapipb.Worker{WorkerNamespace: "default", WorkerPool: "pool-1", WorkerPod: "pod-1"} + if err := s.CreateWorker(ctx, worker); err != nil { + t.Fatalf("CreateWorker failed: %v", err) + } + stored, err := s.GetWorker(ctx, "default", "pod-1") + if err != nil { + t.Fatalf("GetWorker failed: %v", err) + } + staleVersion := stored.GetVersion() + + stored.NodeName = "node-2" + if err := s.UpdateWorker(ctx, stored, staleVersion); err != nil { + t.Fatalf("UpdateWorker failed: %v", err) + } + + if err := s.DeleteWorker(ctx, stored, staleVersion); !errors.Is(err, store.ErrVersionConflict) { + t.Errorf("DeleteWorker at a stale version = %v, want ErrVersionConflict", err) + } + if _, err := s.GetWorker(ctx, "default", "pod-1"); err != nil { + t.Errorf("worker should survive a rejected delete, got %v", err) + } + }) + + // CreateWorker restarts the version at 1, so a decision taken against the + // previous incarnation reads as current unless the Pod UID is checked too. + t.Run("DeleteWorker_RejectsReplacedIncarnation", func(t *testing.T) { + s := setup(t) + ctx := context.Background() + + dead := &ateapipb.Worker{WorkerNamespace: "default", WorkerPool: "pool-1", WorkerPod: "pod-1", WorkerPodUid: "uid-old"} + if err := s.CreateWorker(ctx, dead); err != nil { + t.Fatalf("CreateWorker failed: %v", err) + } + stored, err := s.GetWorker(ctx, "default", "pod-1") + if err != nil { + t.Fatalf("GetWorker failed: %v", err) + } + if err := s.DeleteWorker(ctx, stored, stored.GetVersion()); err != nil { + t.Fatalf("DeleteWorker failed: %v", err) + } + + reborn := &ateapipb.Worker{WorkerNamespace: "default", WorkerPool: "pool-1", WorkerPod: "pod-1", WorkerPodUid: "uid-new"} + if err := s.CreateWorker(ctx, reborn); err != nil { + t.Fatalf("CreateWorker for the replacement failed: %v", err) + } + + if err := s.DeleteWorker(ctx, stored, stored.GetVersion()); !errors.Is(err, store.ErrVersionConflict) { + t.Errorf("DeleteWorker of a replaced incarnation = %v, want ErrVersionConflict", err) + } + if _, err := s.GetWorker(ctx, "default", "pod-1"); err != nil { + t.Errorf("the replacement must survive, got %v", err) + } + }) + + // A task queued while the Pod belonged to another pool must not remove the + // record its successor now owns. + t.Run("DeleteWorker_RejectsOtherPool", func(t *testing.T) { + s := setup(t) + ctx := context.Background() + + if err := s.CreateWorker(ctx, &ateapipb.Worker{WorkerNamespace: "default", WorkerPool: "pool-2", WorkerPod: "pod-1", WorkerPodUid: "uid-1"}); err != nil { + t.Fatalf("CreateWorker failed: %v", err) + } + stale := &ateapipb.Worker{WorkerNamespace: "default", WorkerPool: "pool-1", WorkerPod: "pod-1", WorkerPodUid: "uid-1"} + if err := s.DeleteWorker(ctx, stale, 1); !errors.Is(err, store.ErrVersionConflict) { + t.Errorf("DeleteWorker from another pool = %v, want ErrVersionConflict", err) + } + if _, err := s.GetWorker(ctx, "default", "pod-1"); err != nil { + t.Errorf("the record must survive, got %v", err) + } + }) + t.Run("WatchWorkers_ClosedOnClose", func(t *testing.T) { s := setup(t) ctx := context.Background() diff --git a/cmd/ateapi/main.go b/cmd/ateapi/main.go index 5d1d3e0d5..ae36e763e 100644 --- a/cmd/ateapi/main.go +++ b/cmd/ateapi/main.go @@ -195,7 +195,7 @@ func main() { ateletDialer := controlapi.NewAteletDialer(workerPodInformer.GetIndexer(), ateletPodInformer.GetIndexer(), *ateletClientCredBundle, *podIdentityCACerts) sm := controlapi.NewService(persistence, workerCache, actorTemplateLister, workerPoolLister, sandboxConfigLister, csiDriverConfigLister, storageClassLister, ateletDialer, instruments, *egressGatewayAddress, volPlugins) - actorIdentitySrv := actoridentity.New(actorIdentityJWTIssuer, *actorIDJWTPoolFile, *actorIDCAPoolFile, persistence, workerCache) + actorIdentitySrv := actoridentity.New(actorIdentityJWTIssuer, *actorIDJWTPoolFile, *actorIDCAPoolFile, persistence) debugSrv := debugapi.NewService(persistence) lisCfg := &net.ListenConfig{}