From 74d149d26ae182665e0c3d23b5df87db51b65298 Mon Sep 17 00:00:00 2001 From: NekoPunch Date: Fri, 14 Aug 2026 16:20:39 -0700 Subject: [PATCH] fix(ateapi): authorize MintCert against the store A worker's store key included its pool, which a mint request cannot carry, so authorization resolved it through the watch-fed worker cache and denied assignments the store had already committed. Keying workers by the Pod backing them lets the gate read the authoritative record. Removing the pool from the key also removed the ownership check it implicitly provided, so the syncer now carries the Pod identity in its queue key and conditions worker deletes and actor releases on the exact incarnation it inspected. --- .../internal/actoridentity/actoridentity.go | 9 +- .../actoridentity/actoridentity_test.go | 67 ++++-- cmd/ateapi/internal/controlapi/crash.go | 2 +- cmd/ateapi/internal/controlapi/crash_test.go | 8 +- .../internal/controlapi/functional_test.go | 2 +- cmd/ateapi/internal/controlapi/syncer.go | 103 ++++----- cmd/ateapi/internal/controlapi/syncer_test.go | 197 ++++++++++++++---- cmd/ateapi/internal/controlapi/volumes.go | 2 +- .../internal/controlapi/workflow_pause.go | 2 +- .../internal/controlapi/workflow_resume.go | 2 +- .../controlapi/workflow_resume_test.go | 14 +- .../internal/controlapi/workflow_suspend.go | 2 +- .../controlapi/workflow_suspend_test.go | 2 +- cmd/ateapi/internal/store/atepg/atepg.go | 77 ++++--- cmd/ateapi/internal/store/atepg/schema.go | 2 +- .../internal/store/ateredis/ateredis.go | 55 +++-- .../internal/store/ateredis/ateredis_test.go | 18 +- cmd/ateapi/internal/store/store.go | 21 +- .../internal/store/storecontract/contract.go | 138 +++++++++++- cmd/ateapi/main.go | 2 +- 20 files changed, 531 insertions(+), 194 deletions(-) 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{}