diff --git a/cmd/ateapi/internal/controlapi/actor.go b/cmd/ateapi/internal/controlapi/actor.go index fd18e3ba88..6af07b2509 100644 --- a/cmd/ateapi/internal/controlapi/actor.go +++ b/cmd/ateapi/internal/controlapi/actor.go @@ -105,6 +105,9 @@ func (s *Service) CreateActor(ctx context.Context, req *ateapipb.CreateActorRequ if errors.Is(err, store.ErrAlreadyExists) { return nil, status.Errorf(codes.AlreadyExists, "Actor %s already exists", name) } + if errors.Is(err, store.ErrFailedPrecondition) { + return nil, status.Errorf(codes.FailedPrecondition, "Atespace %s not found", atespace) + } return nil, fmt.Errorf("while recording actor: %w", err) } @@ -243,7 +246,7 @@ func (s *Service) ListActors(ctx context.Context, req *ateapipb.ListActorsReques page, err := s.persistence.ListActors(ctx, req.GetAtespace(), store.ListOptions{PageSize: effectivePageSize(req.GetPageSize()), PageToken: req.GetPageToken()}) if err != nil { - return nil, fmt.Errorf("while listing actors in db: %w", err) + return nil, mapListError(fmt.Errorf("while listing actors in db: %w", err)) } return &ateapipb.ListActorsResponse{ Actors: page.Items, diff --git a/cmd/ateapi/internal/controlapi/actor_snapshot.go b/cmd/ateapi/internal/controlapi/actor_snapshot.go index 7e34d8de06..9496b668df 100644 --- a/cmd/ateapi/internal/controlapi/actor_snapshot.go +++ b/cmd/ateapi/internal/controlapi/actor_snapshot.go @@ -107,7 +107,7 @@ func (s *Service) ListActorSnapshots(ctx context.Context, req *ateapipb.ListActo } page, err := s.persistence.ListActorSnapshots(ctx, req.GetAtespace(), store.ListOptions{PageSize: effectivePageSize(req.GetPageSize()), PageToken: req.GetPageToken()}) if err != nil { - return nil, fmt.Errorf("while listing actor snapshots: %w", err) + return nil, mapListError(fmt.Errorf("while listing actor snapshots: %w", err)) } return &ateapipb.ListActorSnapshotsResponse{Snapshots: page.Items, NextPageToken: page.NextPageToken}, nil } diff --git a/cmd/ateapi/internal/controlapi/actor_snapshot_test.go b/cmd/ateapi/internal/controlapi/actor_snapshot_test.go index 0193adeed6..9af3fdce26 100644 --- a/cmd/ateapi/internal/controlapi/actor_snapshot_test.go +++ b/cmd/ateapi/internal/controlapi/actor_snapshot_test.go @@ -252,6 +252,23 @@ func TestValidateUpdateActorSnapshotTagRequest(t *testing.T) { } } +func TestCreateActorSnapshotTag_MissingSnapshotIsNotFound(t *testing.T) { + persistence, cleanup := storetest.SetupTestStore(t) + t.Cleanup(cleanup) + s := &Service{persistence: persistence} + + _, err := s.CreateActorSnapshotTag(context.Background(), &ateapipb.CreateActorSnapshotTagRequest{ + ActorSnapshotTag: &ateapipb.ActorSnapshotTag{ + Metadata: &ateapipb.ResourceMetadata{Atespace: "team-a", Name: "latest"}, + Snapshot: &ateapipb.ObjectRef{Atespace: "team-a", Name: "missing"}, + Scope: ateapipb.ActorSnapshotTagScope_ACTOR_SNAPSHOT_TAG_SCOPE_ATESPACE, + }, + }) + if status.Code(err) != codes.NotFound { + t.Fatalf("CreateActorSnapshotTag status = %v, want NotFound (error: %v)", status.Code(err), err) + } +} + func TestUpdateActorSnapshotTag_FieldMasks(t *testing.T) { tests := []struct { name string diff --git a/cmd/ateapi/internal/controlapi/actor_test.go b/cmd/ateapi/internal/controlapi/actor_test.go index c04c86953c..a023acd1ca 100644 --- a/cmd/ateapi/internal/controlapi/actor_test.go +++ b/cmd/ateapi/internal/controlapi/actor_test.go @@ -21,6 +21,7 @@ import ( "testing" "time" + "github.com/agent-substrate/substrate/cmd/ateapi/internal/store" "github.com/agent-substrate/substrate/cmd/ateapi/internal/store/storetest" "github.com/agent-substrate/substrate/internal/ateattr" "github.com/agent-substrate/substrate/internal/resources" @@ -39,6 +40,32 @@ import ( "k8s.io/apimachinery/pkg/util/wait" ) +type createActorErrorStore struct { + serviceStore + err error +} + +func (s *createActorErrorStore) CreateActor(context.Context, *ateapipb.Actor) (*ateapipb.Actor, error) { + return nil, s.err +} + +func TestCreateActor_AtespaceDeletedAfterPrecheck(t *testing.T) { + ns := namespaceForTest("ns-create-atespace-race") + tc := setupTest(t, ns) + defer tc.cleanup() + createTemplate(t, tc, ns) + tc.service.persistence = &createActorErrorStore{serviceStore: tc.service.persistence, err: store.ErrFailedPrecondition} + + _, err := tc.service.CreateActor(context.Background(), &ateapipb.CreateActorRequest{Actor: &ateapipb.Actor{ + Metadata: &ateapipb.ResourceMetadata{Atespace: testAtespace, Name: "racing-create"}, + ActorTemplateNamespace: ns, + ActorTemplateName: "tmpl1", + }}) + if status.Code(err) != codes.FailedPrecondition { + t.Fatalf("CreateActor status = %v, want FailedPrecondition (error: %v)", status.Code(err), err) + } +} + // CreateActor is the only lifecycle op with the full identity (incl. version) // available in the request, so the whole ate.* set should land on its span. func TestCreateActor_StampsFullSpanIdentity(t *testing.T) { diff --git a/cmd/ateapi/internal/controlapi/atespace.go b/cmd/ateapi/internal/controlapi/atespace.go index 220275791e..8d7fd18913 100644 --- a/cmd/ateapi/internal/controlapi/atespace.go +++ b/cmd/ateapi/internal/controlapi/atespace.go @@ -110,7 +110,7 @@ func (s *Service) ListAtespaces(ctx context.Context, req *ateapipb.ListAtespaces page, err := s.persistence.ListAtespaces(ctx, store.ListOptions{PageSize: effectivePageSize(req.GetPageSize()), PageToken: req.GetPageToken()}) if err != nil { - return nil, fmt.Errorf("while listing atespaces in db: %w", err) + return nil, mapListError(fmt.Errorf("while listing atespaces in db: %w", err)) } return &ateapipb.ListAtespacesResponse{ Atespaces: page.Items, diff --git a/cmd/ateapi/internal/controlapi/functional_test.go b/cmd/ateapi/internal/controlapi/functional_test.go index 7d761f0fb0..95c2bed3e7 100644 --- a/cmd/ateapi/internal/controlapi/functional_test.go +++ b/cmd/ateapi/internal/controlapi/functional_test.go @@ -2970,16 +2970,36 @@ func TestValidation(t *testing.T) { assertGrpcErrorRegex(t, err, codes.InvalidArgument, "page_size: Invalid value") }) + t.Run("ListActors invalid token", func(t *testing.T) { + _, err := tc.client.ListActors(context.Background(), &ateapipb.ListActorsRequest{PageToken: "%%%"}) + assertGrpcError(t, err, codes.InvalidArgument, "invalid page_token") + }) + t.Run("ListWorkers", func(t *testing.T) { _, err := tc.client.ListWorkers(context.Background(), &ateapipb.ListWorkersRequest{PageSize: -1}) assertGrpcErrorRegex(t, err, codes.InvalidArgument, "page_size: Invalid value") }) + t.Run("ListWorkers invalid token", func(t *testing.T) { + _, err := tc.client.ListWorkers(context.Background(), &ateapipb.ListWorkersRequest{PageToken: "%%%"}) + assertGrpcError(t, err, codes.InvalidArgument, "invalid page_token") + }) + t.Run("ListAtespaces", func(t *testing.T) { _, err := tc.client.ListAtespaces(context.Background(), &ateapipb.ListAtespacesRequest{PageSize: -1}) assertGrpcErrorRegex(t, err, codes.InvalidArgument, "page_size: Invalid value") }) + t.Run("ListAtespaces invalid token", func(t *testing.T) { + _, err := tc.client.ListAtespaces(context.Background(), &ateapipb.ListAtespacesRequest{PageToken: "%%%"}) + assertGrpcError(t, err, codes.InvalidArgument, "invalid page_token") + }) + + t.Run("ListActorSnapshots invalid token", func(t *testing.T) { + _, err := tc.client.ListActorSnapshots(context.Background(), &ateapipb.ListActorSnapshotsRequest{PageToken: "%%%"}) + assertGrpcError(t, err, codes.InvalidArgument, "invalid page_token") + }) + t.Run("CreateAtespace", func(t *testing.T) { _, err := tc.client.CreateAtespace(context.Background(), &ateapipb.CreateAtespaceRequest{}) assertGrpcErrorRegex(t, err, codes.InvalidArgument, "atespace: Required value") diff --git a/cmd/ateapi/internal/controlapi/pagination.go b/cmd/ateapi/internal/controlapi/pagination.go index 1d9d591f2f..9a458be401 100644 --- a/cmd/ateapi/internal/controlapi/pagination.go +++ b/cmd/ateapi/internal/controlapi/pagination.go @@ -14,6 +14,14 @@ package controlapi +import ( + "errors" + + "github.com/agent-substrate/substrate/cmd/ateapi/internal/store" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" +) + const maxPageSize = 1000 // effectivePageSize applies the server-chosen default for an unset page_size @@ -24,3 +32,13 @@ func effectivePageSize(requested int32) int32 { } return requested } + +func mapListError(err error) error { + if errors.Is(err, store.ErrInvalidPageToken) { + return status.Error(codes.InvalidArgument, "invalid page_token") + } + if errors.Is(err, store.ErrInvalidPageSize) { + return status.Error(codes.InvalidArgument, "invalid page_size") + } + return err +} diff --git a/cmd/ateapi/internal/controlapi/worker.go b/cmd/ateapi/internal/controlapi/worker.go index 5c394742bd..0e7cc3ef12 100644 --- a/cmd/ateapi/internal/controlapi/worker.go +++ b/cmd/ateapi/internal/controlapi/worker.go @@ -30,7 +30,7 @@ func (s *Service) ListWorkers(ctx context.Context, req *ateapipb.ListWorkersRequ page, err := s.persistence.ListWorkers(ctx, store.ListOptions{PageSize: effectivePageSize(req.GetPageSize()), PageToken: req.GetPageToken()}) if err != nil { - return nil, fmt.Errorf("while listing workers in db: %w", err) + return nil, mapListError(fmt.Errorf("while listing workers in db: %w", err)) } return &ateapipb.ListWorkersResponse{ Workers: page.Items, diff --git a/cmd/ateapi/internal/controlapi/workflow_resume.go b/cmd/ateapi/internal/controlapi/workflow_resume.go index 43d854237a..e0661bfa29 100644 --- a/cmd/ateapi/internal/controlapi/workflow_resume.go +++ b/cmd/ateapi/internal/controlapi/workflow_resume.go @@ -83,6 +83,18 @@ func (w *ActorWorkflow) ResumeActor(ctx context.Context, actorRef resources.Acto lifecycleOpAttrs(actor, actorTemplate, tele.SnapshotKind, tele.WireSnapshotScope)...) }() + // Routed requests call ResumeActor even when the actor is already running. + // Read before taking the distributed lease so that hot-path checks do not + // upsert and delete a PostgreSQL lease row. Any state that needs work is read + // again under the lock below. + actor, err = w.store.GetActor(ctx, actorRef) + if err != nil { + return nil, false, err + } + if wasRunning = actor.GetStatus() == ateapipb.Actor_STATUS_RUNNING; wasRunning { + return actor, false, nil + } + lockCtx, lock, err := w.acquireActorLock(ctx, actorRef) if err != nil { return nil, false, err @@ -329,6 +341,9 @@ func (w *ActorWorkflow) ensureWorkerAssigned(ctx context.Context, actorRef resou return false, attemptErr }) if err != nil { + if wait.Interrupted(err) && !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) { + return nil, nil, store.ErrVersionConflict + } return nil, nil, err } return assignedActor, assignedWorker, nil @@ -505,6 +520,10 @@ func (w *ActorWorkflow) assignWorkerAttempt(ctx context.Context, actorRef resour } if err := w.store.UpdateWorker(ctx, assignedWorker, assignedWorker.Version); err != nil { + if errors.Is(err, store.ErrNotFound) { + w.workerCache.Forget(assignedWorker.GetWorkerNamespace(), assignedWorker.GetWorkerPod()) + return nil, nil, fmt.Errorf("selected worker disappeared before claim: %w", store.ErrVersionConflict) + } return nil, nil, err } diff --git a/cmd/ateapi/internal/controlapi/workflow_resume_test.go b/cmd/ateapi/internal/controlapi/workflow_resume_test.go index 7af031ca14..56a6873291 100644 --- a/cmd/ateapi/internal/controlapi/workflow_resume_test.go +++ b/cmd/ateapi/internal/controlapi/workflow_resume_test.go @@ -62,6 +62,88 @@ func TestSchedulerRecordable(t *testing.T) { } } +type lockCountingStore struct { + store.Interface + acquireCalls int +} + +func (s *lockCountingStore) AcquireLock(ctx context.Context, key string) (*store.Lock, error) { + s.acquireCalls++ + return s.Interface.AcquireLock(ctx, key) +} + +func TestResumeActor_RunningFastPathDoesNotAcquireLock(t *testing.T) { + ctx := context.Background() + persistence := newTestPersistence(t) + created, err := persistence.CreateActor(ctx, &ateapipb.Actor{ + Metadata: &ateapipb.ResourceMetadata{Atespace: "team-a", Name: "id1"}, + Status: ateapipb.Actor_STATUS_RUNNING, + }) + if err != nil { + t.Fatalf("CreateActor: %v", err) + } + st := &lockCountingStore{Interface: persistence} + w := &ActorWorkflow{store: st} + + got, resumed, err := w.ResumeActor(ctx, resources.ActorRef{Atespace: "team-a", Name: "id1"}, false) + if err != nil { + t.Fatalf("ResumeActor: %v", err) + } + if resumed { + t.Error("ResumeActor resumed = true, want false") + } + if !proto.Equal(got, created) { + t.Errorf("ResumeActor actor = %v, want %v", got, created) + } + if st.acquireCalls != 0 { + t.Errorf("AcquireLock calls = %d, want 0", st.acquireCalls) + } +} + +type updateWorkerErrorStore struct { + store.Interface + err error +} + +func (s *updateWorkerErrorStore) UpdateWorker(context.Context, *ateapipb.Worker, int64) error { + return s.err +} + +func TestAssignWorkerAttempt_MissingSelectedWorkerIsRetried(t *testing.T) { + ctx := context.Background() + persistence := newTestPersistence(t) + actor, wc := seedAssignFixture(t, ctx, persistence) + st := &updateWorkerErrorStore{Interface: persistence, err: store.ErrNotFound} + w := &ActorWorkflow{store: st, workerCache: wc, scheduler: scheduling.New(wc)} + tmpl := &atev1alpha1.ActorTemplate{Spec: atev1alpha1.ActorTemplateSpec{SandboxClass: atev1alpha1.SandboxClassGvisor}} + + _, _, err := w.assignWorkerAttempt(ctx, resources.ActorRef{Atespace: "team-a", Name: "id1"}, actor, tmpl) + if !errors.Is(err, store.ErrVersionConflict) { + t.Fatalf("assignWorkerAttempt error = %v, want ErrVersionConflict", err) + } + workers, err := wc.Workers() + if err != nil { + t.Fatalf("Workers: %v", err) + } + if len(workers) != 0 { + t.Errorf("cached workers after missing claim = %d, want 0", len(workers)) + } +} + +func TestEnsureWorkerAssigned_ConflictExhaustionIsRetryable(t *testing.T) { + ctx := context.Background() + persistence := newTestPersistence(t) + actor, wc := seedAssignFixture(t, ctx, persistence) + st := &updateWorkerErrorStore{Interface: persistence, err: store.ErrVersionConflict} + w := &ActorWorkflow{store: st, workerCache: wc, scheduler: scheduling.New(wc)} + tmpl := &atev1alpha1.ActorTemplate{Spec: atev1alpha1.ActorTemplateSpec{SandboxClass: atev1alpha1.SandboxClassGvisor}} + + _, _, err := w.ensureWorkerAssigned(ctx, resources.ActorRef{Atespace: "team-a", Name: "id1"}, actor, tmpl) + if !errors.Is(err, store.ErrVersionConflict) { + t.Fatalf("ensureWorkerAssigned error = %v, want ErrVersionConflict", err) + } +} + func TestAssignWorkerAttempt_SkipsWorkerAssignedInOtherAtespace(t *testing.T) { ctx := context.Background() persistence := newTestPersistence(t) diff --git a/cmd/ateapi/internal/store/atepg/atepg.go b/cmd/ateapi/internal/store/atepg/atepg.go index ead0b22f8d..1b9d08b58e 100644 --- a/cmd/ateapi/internal/store/atepg/atepg.go +++ b/cmd/ateapi/internal/store/atepg/atepg.go @@ -105,6 +105,20 @@ func newUpdateMetadata(current *ateapipb.ResourceMetadata) *ateapipb.ResourceMet return metadata } +// updateMaxAttempts bounds how many times a read-modify-write is retried after +// its optimistic uid/version check loses to a concurrent writer. +const updateMaxAttempts = 5 + +func validateMetadataProjection(resource string, metadata *ateapipb.ResourceMetadata, uid string, version int64) error { + if metadata.GetUid() != uid { + return fmt.Errorf("%s uid projection %q does not match proto metadata uid %q", resource, uid, metadata.GetUid()) + } + if metadata.GetVersion() != version { + return fmt.Errorf("%s version projection %d does not match proto metadata version %d", resource, version, metadata.GetVersion()) + } + return nil +} + func isUniqueViolation(err error) bool { return pgErrCode(err) == "23505" } // isForeignKeyViolation matches both the insert/update-side violation @@ -128,6 +142,14 @@ func pgErrCode(err error) string { return "" } +func pgErrConstraint(err error) string { + var pgErr *pgconn.PgError + if errors.As(err, &pgErr) { + return pgErr.ConstraintName + } + return "" +} + // --- Atespaces --- func (p *Persistence) CreateAtespace(ctx context.Context, atespace *ateapipb.Atespace) (*ateapipb.Atespace, error) { @@ -183,6 +205,10 @@ func (p *Persistence) AtespaceExists(ctx context.Context, name string) (bool, er } func (p *Persistence) ListAtespaces(ctx context.Context, opts store.ListOptions) (store.ListResponse[*ateapipb.Atespace], error) { + opts, err := store.NormalizeListOptions(opts) + if err != nil { + return store.ListResponse[*ateapipb.Atespace]{}, err + } pageSize, pageTokenStr := opts.PageSize, opts.PageToken token, err := decodePageToken(pageTokenStr, kindAtespace, "", 1) if err != nil { @@ -314,52 +340,60 @@ func validateUpdateActorTemplateMutation(storedTemplate, mutatedTemplate *ateapi } func (p *Persistence) UpdateActorTemplate(ctx context.Context, templateRef resources.ActorTemplateRef, mutate func(*ateapipb.ActorTemplate) error) (*ateapipb.ActorTemplate, error) { - tx, err := p.pool.Begin(ctx) - if err != nil { - return nil, fmt.Errorf("beginning actor template update: %w", err) - } - defer tx.Rollback(ctx) //nolint:errcheck // no-op once committed - - var currentBytes []byte - if err := tx.QueryRow(ctx, ` - SELECT proto FROM actor_templates - WHERE atespace = $1 AND name = $2 - FOR UPDATE`, templateRef.Atespace, templateRef.Name).Scan(¤tBytes); err != nil { - if errors.Is(err, pgx.ErrNoRows) { - return nil, store.ErrNotFound + for range updateMaxAttempts { + var currentUID string + var currentVersion int64 + var currentBytes []byte + if err := p.pool.QueryRow(ctx, ` + SELECT uid, version, proto FROM actor_templates + WHERE atespace = $1 AND name = $2`, templateRef.Atespace, templateRef.Name).Scan(¤tUID, ¤tVersion, ¤tBytes); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return nil, store.ErrNotFound + } + return nil, fmt.Errorf("getting actor template %s for update: %w", templateRef, err) } - return nil, fmt.Errorf("locking actor template %s for update: %w", templateRef, err) - } - dbTemplate := &ateapipb.ActorTemplate{} - if err := proto.Unmarshal(currentBytes, dbTemplate); err != nil { - return nil, fmt.Errorf("unmarshaling actor template for update: %w", err) - } - templateBeforeMutation := proto.Clone(dbTemplate).(*ateapipb.ActorTemplate) - if err := mutate(dbTemplate); err != nil { - return nil, err - } - if err := validateUpdateActorTemplateMutation(templateBeforeMutation, dbTemplate); err != nil { - return nil, err - } - dbTemplate.Metadata = newUpdateMetadata(templateBeforeMutation.GetMetadata()) - updatedBytes, err := proto.Marshal(dbTemplate) - if err != nil { - return nil, fmt.Errorf("marshaling actor template: %w", err) - } - if _, err := tx.Exec(ctx, ` - UPDATE actor_templates SET version = $1, proto = $2 - WHERE atespace = $3 AND name = $4`, - dbTemplate.GetMetadata().GetVersion(), updatedBytes, templateRef.Atespace, templateRef.Name); err != nil { - return nil, fmt.Errorf("updating actor template %s: %w", templateRef, err) - } - if err := tx.Commit(ctx); err != nil { - return nil, fmt.Errorf("committing actor template update: %w", err) + dbTemplate := &ateapipb.ActorTemplate{} + if err := proto.Unmarshal(currentBytes, dbTemplate); err != nil { + return nil, fmt.Errorf("unmarshaling actor template for update: %w", err) + } + if err := validateMetadataProjection("actor template "+templateRef.String(), dbTemplate.GetMetadata(), currentUID, currentVersion); err != nil { + return nil, err + } + templateBeforeMutation := proto.Clone(dbTemplate).(*ateapipb.ActorTemplate) + if err := mutate(dbTemplate); err != nil { + return nil, err + } + if err := validateUpdateActorTemplateMutation(templateBeforeMutation, dbTemplate); err != nil { + return nil, err + } + dbTemplate.Metadata = newUpdateMetadata(templateBeforeMutation.GetMetadata()) + updatedBytes, err := proto.Marshal(dbTemplate) + if err != nil { + return nil, fmt.Errorf("marshaling actor template: %w", err) + } + commandTag, err := p.pool.Exec(ctx, ` + UPDATE actor_templates SET version = $1, proto = $2 + WHERE atespace = $3 AND name = $4 AND uid = $5 AND version = $6`, + dbTemplate.GetMetadata().GetVersion(), updatedBytes, templateRef.Atespace, templateRef.Name, currentUID, currentVersion) + if err != nil { + return nil, fmt.Errorf("updating actor template %s: %w", templateRef, err) + } + if commandTag.RowsAffected() == 1 { + return dbTemplate, nil + } + if commandTag.RowsAffected() != 0 { + return nil, fmt.Errorf("updating actor template %s affected %d rows, want at most 1", templateRef, commandTag.RowsAffected()) + } } - return dbTemplate, nil + return nil, store.ErrVersionConflict } func (p *Persistence) ListActorTemplates(ctx context.Context, atespace string, opts store.ListOptions) (store.ListResponse[*ateapipb.ActorTemplate], error) { + opts, err := store.NormalizeListOptions(opts) + if err != nil { + return store.ListResponse[*ateapipb.ActorTemplate]{}, err + } pageSize, pageTokenStr := opts.PageSize, opts.PageToken keyParts := 2 if atespace != "" { @@ -510,6 +544,10 @@ func actorTemplateVersionListScope(atespace string, parent resources.ActorTempla } func (p *Persistence) ListActorTemplateVersions(ctx context.Context, atespace string, parent resources.ActorTemplateRef, opts store.ListOptions) (store.ListResponse[*ateapipb.ActorTemplateVersion], error) { + opts, err := store.NormalizeListOptions(opts) + if err != nil { + return store.ListResponse[*ateapipb.ActorTemplateVersion]{}, err + } pageSize, pageTokenStr := opts.PageSize, opts.PageToken scope := actorTemplateVersionListScope(atespace, parent) keyParts := 2 @@ -710,60 +748,57 @@ func validateUpdateActorMutation(storedActor, mutatedActor *ateapipb.Actor) erro func (p *Persistence) UpdateActor(ctx context.Context, actorRef resources.ActorRef, mutate func(*ateapipb.Actor) error) (*ateapipb.Actor, error) { atespace, name := actorRef.Atespace, actorRef.Name - - tx, err := p.pool.Begin(ctx) - if err != nil { - return nil, fmt.Errorf("beginning actor update: %w", err) - } - defer tx.Rollback(ctx) //nolint:errcheck // no-op once committed - - var protoBytes []byte - if err := tx.QueryRow(ctx, ` - SELECT proto FROM actors - WHERE atespace = $1 AND name = $2 - FOR UPDATE`, atespace, name).Scan(&protoBytes); err != nil { - if errors.Is(err, pgx.ErrNoRows) { - return nil, store.ErrNotFound + for range updateMaxAttempts { + var currentUID string + var currentVersion int64 + var currentBytes []byte + if err := p.pool.QueryRow(ctx, ` + SELECT uid, version, proto FROM actors + WHERE atespace = $1 AND name = $2`, atespace, name).Scan(¤tUID, ¤tVersion, ¤tBytes); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return nil, store.ErrNotFound + } + return nil, fmt.Errorf("getting actor %s/%s for update: %w", atespace, name, err) } - return nil, fmt.Errorf("locking actor %s/%s for update: %w", atespace, name, err) - } - - dbActor := &ateapipb.Actor{} - if err := proto.Unmarshal(protoBytes, dbActor); err != nil { - return nil, fmt.Errorf("unmarshaling actor for update: %w", err) - } - actorBeforeMutation := proto.Clone(dbActor).(*ateapipb.Actor) - if err := mutate(dbActor); err != nil { - return nil, err - } - if err := validateUpdateActorMutation(actorBeforeMutation, dbActor); err != nil { - return nil, err - } - // Stored metadata is authoritative; discard any metadata edits made by the - // closure and derive the next revision from the transactionally read actor. - dbActor.Metadata = newUpdateMetadata(actorBeforeMutation.GetMetadata()) - - protoBytes, err = proto.Marshal(dbActor) - if err != nil { - return nil, fmt.Errorf("marshaling actor: %w", err) - } - commandTag, err := tx.Exec(ctx, ` - UPDATE actors - SET version = $1, proto = $2 - WHERE atespace = $3 AND name = $4`, - dbActor.GetMetadata().GetVersion(), protoBytes, atespace, name) - if err != nil { - return nil, fmt.Errorf("updating actor %s/%s: %w", atespace, name, err) - } - if commandTag.RowsAffected() != 1 { - return nil, fmt.Errorf("updating actor %s/%s affected %d rows, want 1", atespace, name, commandTag.RowsAffected()) - } + dbActor := &ateapipb.Actor{} + if err := proto.Unmarshal(currentBytes, dbActor); err != nil { + return nil, fmt.Errorf("unmarshaling actor for update: %w", err) + } + if err := validateMetadataProjection("actor "+actorRef.String(), dbActor.GetMetadata(), currentUID, currentVersion); err != nil { + return nil, err + } + actorBeforeMutation := proto.Clone(dbActor).(*ateapipb.Actor) + if err := mutate(dbActor); err != nil { + return nil, err + } + if err := validateUpdateActorMutation(actorBeforeMutation, dbActor); err != nil { + return nil, err + } + // Stored metadata is authoritative; discard any metadata edits made by the + // closure and derive the next revision from the state this attempt read. + dbActor.Metadata = newUpdateMetadata(actorBeforeMutation.GetMetadata()) - if err := tx.Commit(ctx); err != nil { - return nil, fmt.Errorf("committing actor update: %w", err) + updatedBytes, err := proto.Marshal(dbActor) + if err != nil { + return nil, fmt.Errorf("marshaling actor: %w", err) + } + commandTag, err := p.pool.Exec(ctx, ` + UPDATE actors + SET version = $1, proto = $2 + WHERE atespace = $3 AND name = $4 AND uid = $5 AND version = $6`, + dbActor.GetMetadata().GetVersion(), updatedBytes, atespace, name, currentUID, currentVersion) + if err != nil { + return nil, fmt.Errorf("updating actor %s/%s: %w", atespace, name, err) + } + if commandTag.RowsAffected() == 1 { + return dbActor, nil + } + if commandTag.RowsAffected() != 0 { + return nil, fmt.Errorf("updating actor %s/%s affected %d rows, want at most 1", atespace, name, commandTag.RowsAffected()) + } } - return dbActor, nil + return nil, store.ErrVersionConflict } func (p *Persistence) DeleteActor(ctx context.Context, actorRef resources.ActorRef) (*ateapipb.Actor, error) { @@ -805,9 +840,12 @@ func (p *Persistence) DeleteActor(ctx context.Context, actorRef resources.ActorR } func (p *Persistence) ListActors(ctx context.Context, atespace string, opts store.ListOptions) (store.ListResponse[*ateapipb.Actor], error) { + opts, err := store.NormalizeListOptions(opts) + if err != nil { + return store.ListResponse[*ateapipb.Actor]{}, err + } var items []*ateapipb.Actor var nextToken string - var err error if atespace != "" { items, nextToken, err = p.listActorsScoped(ctx, atespace, opts.PageSize, opts.PageToken) } else { @@ -978,9 +1016,12 @@ func (p *Persistence) GetActorSnapshotTag(ctx context.Context, atespace, name st } func (p *Persistence) ListActorSnapshots(ctx context.Context, atespace string, opts store.ListOptions) (store.ListResponse[*ateapipb.ActorSnapshot], error) { + opts, err := store.NormalizeListOptions(opts) + if err != nil { + return store.ListResponse[*ateapipb.ActorSnapshot]{}, err + } var items []*ateapipb.ActorSnapshot var nextToken string - var err error if atespace != "" { items, nextToken, err = p.listActorSnapshotsScoped(ctx, atespace, opts.PageSize, opts.PageToken) } else { @@ -1100,18 +1141,15 @@ func (p *Persistence) CreateActorSnapshotTag(ctx context.Context, snapshotAtespa return nil, fmt.Errorf("beginning actor snapshot tag create: %w", err) } defer tx.Rollback(ctx) //nolint:errcheck // no-op once committed - if _, err := getActorSnapshotRow(ctx, tx, snapshotAtespace, snapshotName); err != nil { - return nil, err - } var inserted []byte err = tx.QueryRow(ctx, ` INSERT INTO actor_snapshot_tags - (atespace, name, snapshot_atespace, snapshot_name, version, proto) - VALUES ($1, $2, $3, $4, $5, $6) + (atespace, name, snapshot_atespace, snapshot_name, uid, version, proto) + VALUES ($1, $2, $3, $4, $5, $6, $7) ON CONFLICT (atespace, name) DO NOTHING RETURNING proto`, tagAtespace, tagName, snapshotAtespace, snapshotName, - dbTag.GetMetadata().GetVersion(), protoBytes).Scan(&inserted) + dbTag.GetMetadata().GetUid(), dbTag.GetMetadata().GetVersion(), protoBytes).Scan(&inserted) if err == nil { if err := tx.Commit(ctx); err != nil { return nil, fmt.Errorf("committing actor snapshot tag create: %w", err) @@ -1119,7 +1157,14 @@ func (p *Persistence) CreateActorSnapshotTag(ctx context.Context, snapshotAtespa return dbTag, nil } if isForeignKeyViolation(err) { - return nil, store.ErrFailedPrecondition + switch pgErrConstraint(err) { + case "actor_snapshot_tags_snapshot_fk": + return nil, store.ErrNotFound + case "actor_snapshot_tags_atespace_fk": + return nil, store.ErrFailedPrecondition + default: + return nil, fmt.Errorf("inserting actor snapshot tag %s/%s violated unknown foreign key %q: %w", tagAtespace, tagName, pgErrConstraint(err), err) + } } if !errors.Is(err, pgx.ErrNoRows) { return nil, fmt.Errorf("inserting actor snapshot tag %s/%s: %w", tagAtespace, tagName, err) @@ -1161,57 +1206,57 @@ func validateUpdateActorSnapshotTagMutation(storedTag, mutatedTag *ateapipb.Acto } func (p *Persistence) UpdateActorSnapshotTag(ctx context.Context, atespace, name string, mutate func(*ateapipb.ActorSnapshotTag) error) (*ateapipb.ActorSnapshotTag, error) { - tx, err := p.pool.Begin(ctx) - if err != nil { - return nil, fmt.Errorf("beginning actor snapshot tag update: %w", err) - } - defer tx.Rollback(ctx) //nolint:errcheck // no-op once committed - - var currentBytes []byte - if err := tx.QueryRow(ctx, ` - SELECT proto FROM actor_snapshot_tags - WHERE atespace = $1 AND name = $2 - FOR UPDATE`, atespace, name).Scan(¤tBytes); err != nil { - if errors.Is(err, pgx.ErrNoRows) { - return nil, store.ErrNotFound + for range updateMaxAttempts { + var currentUID string + var currentVersion int64 + var currentBytes []byte + if err := p.pool.QueryRow(ctx, ` + SELECT uid, version, proto FROM actor_snapshot_tags + WHERE atespace = $1 AND name = $2`, atespace, name).Scan(¤tUID, ¤tVersion, ¤tBytes); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return nil, store.ErrNotFound + } + return nil, fmt.Errorf("getting actor snapshot tag %s/%s for update: %w", atespace, name, err) } - return nil, fmt.Errorf("locking actor snapshot tag %s/%s for update: %w", atespace, name, err) - } - dbTag := &ateapipb.ActorSnapshotTag{} - if err := proto.Unmarshal(currentBytes, dbTag); err != nil { - return nil, fmt.Errorf("unmarshaling actor snapshot tag: %w", err) - } - tagBeforeMutation := proto.Clone(dbTag).(*ateapipb.ActorSnapshotTag) - if err := mutate(dbTag); err != nil { - return nil, err - } - if err := validateUpdateActorSnapshotTagMutation(tagBeforeMutation, dbTag); err != nil { - return nil, err - } - // Stored metadata is authoritative; discard any metadata edits made by the - // closure and derive the next revision from the transactionally read tag. - dbTag.Metadata = newUpdateMetadata(tagBeforeMutation.GetMetadata()) + dbTag := &ateapipb.ActorSnapshotTag{} + if err := proto.Unmarshal(currentBytes, dbTag); err != nil { + return nil, fmt.Errorf("unmarshaling actor snapshot tag: %w", err) + } + if err := validateMetadataProjection(fmt.Sprintf("actor snapshot tag %s/%s", atespace, name), dbTag.GetMetadata(), currentUID, currentVersion); err != nil { + return nil, err + } + tagBeforeMutation := proto.Clone(dbTag).(*ateapipb.ActorSnapshotTag) + if err := mutate(dbTag); err != nil { + return nil, err + } + if err := validateUpdateActorSnapshotTagMutation(tagBeforeMutation, dbTag); err != nil { + return nil, err + } + // Stored metadata is authoritative; discard any metadata edits made by the + // closure and derive the next revision from the state this attempt read. + dbTag.Metadata = newUpdateMetadata(tagBeforeMutation.GetMetadata()) - updatedBytes, err := proto.Marshal(dbTag) - if err != nil { - return nil, fmt.Errorf("marshaling actor snapshot tag: %w", err) - } - commandTag, err := tx.Exec(ctx, ` - UPDATE actor_snapshot_tags - SET version = $1, proto = $2 - WHERE atespace = $3 AND name = $4`, - dbTag.GetMetadata().GetVersion(), updatedBytes, atespace, name) - if err != nil { - return nil, fmt.Errorf("updating actor snapshot tag %s/%s: %w", atespace, name, err) - } - if commandTag.RowsAffected() != 1 { - return nil, fmt.Errorf("updating actor snapshot tag %s/%s affected %d rows, want 1", atespace, name, commandTag.RowsAffected()) - } - if err := tx.Commit(ctx); err != nil { - return nil, fmt.Errorf("committing actor snapshot tag update: %w", err) + updatedBytes, err := proto.Marshal(dbTag) + if err != nil { + return nil, fmt.Errorf("marshaling actor snapshot tag: %w", err) + } + commandTag, err := p.pool.Exec(ctx, ` + UPDATE actor_snapshot_tags + SET version = $1, proto = $2 + WHERE atespace = $3 AND name = $4 AND uid = $5 AND version = $6`, + dbTag.GetMetadata().GetVersion(), updatedBytes, atespace, name, currentUID, currentVersion) + if err != nil { + return nil, fmt.Errorf("updating actor snapshot tag %s/%s: %w", atespace, name, err) + } + if commandTag.RowsAffected() == 1 { + return dbTag, nil + } + if commandTag.RowsAffected() != 0 { + return nil, fmt.Errorf("updating actor snapshot tag %s/%s affected %d rows, want at most 1", atespace, name, commandTag.RowsAffected()) + } } - return dbTag, nil + return nil, store.ErrVersionConflict } func (p *Persistence) DeleteActorSnapshotTag(ctx context.Context, atespace, name string) (*ateapipb.ActorSnapshotTag, error) { @@ -1413,6 +1458,10 @@ 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) { + opts, err := store.NormalizeListOptions(opts) + if err != nil { + return store.ListResponse[*ateapipb.Worker]{}, err + } pageSize, pageTokenStr := opts.PageSize, opts.PageToken token, err := decodePageToken(pageTokenStr, kindWorker, "", 3) if err != nil { @@ -1496,8 +1545,8 @@ func (p *Persistence) WatchWorkers(ctx context.Context) (*store.WorkerWatch, err } event, err := unmarshalWorkerEvent(notification.Payload) if err != nil { - slog.ErrorContext(ctx, "worker event unmarshal failed", slog.Any("err", err)) - continue + slog.ErrorContext(ctx, "worker event unmarshal failed; closing watch", slog.Any("err", err)) + return } select { case ch <- event: @@ -1518,6 +1567,9 @@ const defaultLockTTL = 30 * time.Second func (p *Persistence) AcquireLock(ctx context.Context, key string) (*store.Lock, error) { ttl := p.lockTTL token := uuid.NewString() + if err := p.cleanupExpiredLeases(ctx); err != nil { + slog.WarnContext(ctx, "failed to clean up expired PostgreSQL leases", "error", err) + } acquired, err := p.acquireLease(ctx, key, token, ttl) if err != nil { @@ -1548,6 +1600,13 @@ func (p *Persistence) AcquireLock(ctx context.Context, key string) (*store.Lock, return store.NewLock(leaseCtx, closeFn), nil } +func (p *Persistence) cleanupExpiredLeases(ctx context.Context) error { + if _, err := p.pool.Exec(ctx, `DELETE FROM leases WHERE expires_at <= clock_timestamp()`); err != nil { + return fmt.Errorf("deleting expired leases: %w", err) + } + return nil +} + func (p *Persistence) acquireLease(ctx context.Context, key, token string, ttl time.Duration) (bool, error) { var returnedKey string err := p.pool.QueryRow(ctx, ` diff --git a/cmd/ateapi/internal/store/atepg/atepg_test.go b/cmd/ateapi/internal/store/atepg/atepg_test.go index e3f1c3b1d1..e69d1ccac4 100644 --- a/cmd/ateapi/internal/store/atepg/atepg_test.go +++ b/cmd/ateapi/internal/store/atepg/atepg_test.go @@ -145,10 +145,181 @@ func createTestAtespace(t *testing.T, s *Persistence, name string) { } } +func createTestActorTemplate(t *testing.T, s *Persistence, atespace, name string) { + t.Helper() + if _, err := s.CreateActorTemplate(context.Background(), &ateapipb.ActorTemplate{ + Metadata: &ateapipb.ResourceMetadata{Atespace: atespace, Name: name}, + }); err != nil { + t.Fatalf("CreateActorTemplate(%q/%q) failed: %v", atespace, name, err) + } +} + +func TestUpdateActor_RetriesConcurrentWrite(t *testing.T) { + s := setupPostgresPersistence(t) + ctx := context.Background() + createTestAtespace(t, s, "team-a") + created, err := s.CreateActor(ctx, &ateapipb.Actor{ + Metadata: &ateapipb.ResourceMetadata{Atespace: "team-a", Name: "actor-a"}, + ActorTemplateNamespace: "default", + ActorTemplateName: "template-a", + Status: ateapipb.Actor_STATUS_SUSPENDED, + }) + if err != nil { + t.Fatalf("CreateActor failed: %v", err) + } + actorRef := resources.ActorRefFromActor(created) + + attempts := 0 + updated, err := s.UpdateActor(ctx, actorRef, func(toUpdate *ateapipb.Actor) error { + attempts++ + if attempts == 1 { + if _, err := s.UpdateActor(ctx, actorRef, func(concurrent *ateapipb.Actor) error { + concurrent.WorkerSelector = &ateapipb.Selector{MatchLabels: map[string]string{"tier": "paid"}} + return nil + }); err != nil { + return fmt.Errorf("concurrent actor update: %w", err) + } + } + toUpdate.Status = ateapipb.Actor_STATUS_RUNNING + return nil + }) + if err != nil { + t.Fatalf("UpdateActor failed: %v", err) + } + if attempts != 2 { + t.Errorf("mutate ran %d times, want 2", attempts) + } + if updated.GetStatus() != ateapipb.Actor_STATUS_RUNNING { + t.Errorf("status = %v, want RUNNING", updated.GetStatus()) + } + if got := updated.GetWorkerSelector().GetMatchLabels()["tier"]; got != "paid" { + t.Errorf("worker selector tier = %q, want paid: concurrent update was lost", got) + } + if got, want := updated.GetMetadata().GetVersion(), created.GetMetadata().GetVersion()+2; got != want { + t.Errorf("version = %d, want %d", got, want) + } +} + +func TestUpdateActor_ExhaustsOptimisticRetries(t *testing.T) { + s := setupPostgresPersistence(t) + ctx := context.Background() + createTestAtespace(t, s, "team-a") + created, err := s.CreateActor(ctx, &ateapipb.Actor{ + Metadata: &ateapipb.ResourceMetadata{Atespace: "team-a", Name: "actor-a"}, + ActorTemplateNamespace: "default", + ActorTemplateName: "template-a", + Status: ateapipb.Actor_STATUS_SUSPENDED, + }) + if err != nil { + t.Fatalf("CreateActor failed: %v", err) + } + actorRef := resources.ActorRefFromActor(created) + + attempts := 0 + _, err = s.UpdateActor(ctx, actorRef, func(toUpdate *ateapipb.Actor) error { + attempts++ + _, err := s.UpdateActor(ctx, actorRef, func(concurrent *ateapipb.Actor) error { + concurrent.Status = ateapipb.Actor_STATUS_RUNNING + return nil + }) + return err + }) + if !errors.Is(err, store.ErrVersionConflict) { + t.Fatalf("UpdateActor error = %v, want ErrVersionConflict", err) + } + if attempts != updateMaxAttempts { + t.Errorf("mutate ran %d times, want %d", attempts, updateMaxAttempts) + } +} + +func TestUpdateActorTemplate_RetriesConcurrentWrite(t *testing.T) { + s := setupPostgresPersistence(t) + ctx := context.Background() + createTestAtespace(t, s, "team-a") + createTestActorTemplate(t, s, "team-a", "template-a") + templateRef := resources.ActorTemplateRef{Atespace: "team-a", Name: "template-a"} + + attempts := 0 + updated, err := s.UpdateActorTemplate(ctx, templateRef, func(toUpdate *ateapipb.ActorTemplate) error { + attempts++ + if attempts == 1 { + if _, err := s.UpdateActorTemplate(ctx, templateRef, func(concurrent *ateapipb.ActorTemplate) error { + concurrent.DefaultVersionOnCreate = &ateapipb.ObjectRef{Atespace: "team-a", Name: "version-a"} + return nil + }); err != nil { + return fmt.Errorf("concurrent actor template update: %w", err) + } + } + return nil + }) + if err != nil { + t.Fatalf("UpdateActorTemplate failed: %v", err) + } + if attempts != 2 { + t.Errorf("mutate ran %d times, want 2", attempts) + } + if got := updated.GetDefaultVersionOnCreate().GetName(); got != "version-a" { + t.Errorf("default version = %q, want version-a: concurrent update was lost", got) + } +} + +func TestUpdateActorSnapshotTag_UIDPreventsDeleteRecreateABA(t *testing.T) { + s := setupPostgresPersistence(t) + ctx := context.Background() + createTestAtespace(t, s, "team-a") + for _, name := range []string{"snapshot-a", "snapshot-b"} { + if _, err := s.CreateActorSnapshot(ctx, &ateapipb.ActorSnapshot{ + Metadata: &ateapipb.ResourceMetadata{Atespace: "team-a", Name: name}, + SnapshotUri: "gs://bucket/" + name, + }); err != nil { + t.Fatalf("CreateActorSnapshot(%q) failed: %v", name, err) + } + } + original, err := s.CreateActorSnapshotTag(ctx, "team-a", "snapshot-a", &ateapipb.ActorSnapshotTag{ + Metadata: &ateapipb.ResourceMetadata{Atespace: "team-a", Name: "tag-a"}, + Scope: ateapipb.ActorSnapshotTagScope_ACTOR_SNAPSHOT_TAG_SCOPE_ATESPACE, + }) + if err != nil { + t.Fatalf("CreateActorSnapshotTag failed: %v", err) + } + + mutations := 0 + var recreated *ateapipb.ActorSnapshotTag + _, err = s.UpdateActorSnapshotTag(ctx, "team-a", "tag-a", store.WithPrecondition(original, func(toUpdate *ateapipb.ActorSnapshotTag) error { + mutations++ + if _, err := s.DeleteActorSnapshotTag(ctx, "team-a", "tag-a"); err != nil { + return fmt.Errorf("deleting original tag: %w", err) + } + recreated, err = s.CreateActorSnapshotTag(ctx, "team-a", "snapshot-b", &ateapipb.ActorSnapshotTag{ + Metadata: &ateapipb.ResourceMetadata{Atespace: "team-a", Name: "tag-a"}, + Scope: ateapipb.ActorSnapshotTagScope_ACTOR_SNAPSHOT_TAG_SCOPE_ATESPACE, + }) + if err != nil { + return fmt.Errorf("recreating tag: %w", err) + } + toUpdate.Scope = ateapipb.ActorSnapshotTagScope_ACTOR_SNAPSHOT_TAG_SCOPE_PUBLISHED + return nil + })) + if !errors.Is(err, store.ErrUIDConflict) { + t.Fatalf("UpdateActorSnapshotTag error = %v, want ErrUIDConflict", err) + } + if mutations != 1 { + t.Errorf("guarded mutation ran %d times, want 1", mutations) + } + stored, err := s.GetActorSnapshotTag(ctx, "team-a", "tag-a") + if err != nil { + t.Fatalf("GetActorSnapshotTag failed: %v", err) + } + if diff := cmp.Diff(recreated, stored, protocmp.Transform()); diff != "" { + t.Errorf("recreated tag was overwritten (-want +got):\n%s", diff) + } +} + func TestDeleteActorTemplateVersion_TaggedGoldenSnapshotRollsBack(t *testing.T) { s := setupPostgresPersistence(t) ctx := context.Background() createTestAtespace(t, s, "team-a") + createTestActorTemplate(t, s, "team-a", "tmpl-a") if _, err := s.CreateActorSnapshot(ctx, &ateapipb.ActorSnapshot{ Metadata: &ateapipb.ResourceMetadata{Atespace: "ate-golden", Name: "golden-1"}, SnapshotUri: "gs://bucket/golden-1", @@ -182,6 +353,7 @@ func TestListActorTemplateVersions_PageTokenRejectsDifferentFilter(t *testing.T) s := setupPostgresPersistence(t) ctx := context.Background() createTestAtespace(t, s, "team-a") + createTestActorTemplate(t, s, "team-a", "tmpl-a") for _, name := range []string{"a-1", "a-2"} { if _, err := s.CreateActorTemplateVersion(ctx, newTestActorTemplateVersion("team-a", name, "tmpl-a")); err != nil { t.Fatalf("CreateActorTemplateVersion(%q) failed: %v", name, err) @@ -201,6 +373,64 @@ func TestListActorTemplateVersions_PageTokenRejectsDifferentFilter(t *testing.T) } } +func TestCreateActorTemplateVersion_MissingParent(t *testing.T) { + s := setupPostgresPersistence(t) + ctx := context.Background() + createTestAtespace(t, s, "team-a") + + if _, err := s.CreateActorTemplateVersion(ctx, newTestActorTemplateVersion("team-a", "v1", "missing")); !errors.Is(err, store.ErrFailedPrecondition) { + t.Fatalf("CreateActorTemplateVersion with missing parent = %v, want ErrFailedPrecondition", err) + } +} + +func TestCreateActorSnapshotTag_ForeignKeyErrors(t *testing.T) { + s := setupPostgresPersistence(t) + ctx := context.Background() + createTestAtespace(t, s, "team-a") + tag := func() *ateapipb.ActorSnapshotTag { + return &ateapipb.ActorSnapshotTag{Metadata: &ateapipb.ResourceMetadata{Atespace: "team-a", Name: "latest"}} + } + + if _, err := s.CreateActorSnapshotTag(ctx, "team-a", "missing", tag()); !errors.Is(err, store.ErrNotFound) { + t.Errorf("missing snapshot error = %v, want ErrNotFound", err) + } + if _, err := s.CreateActorSnapshot(ctx, &ateapipb.ActorSnapshot{Metadata: &ateapipb.ResourceMetadata{Atespace: "gone", Name: "snapshot"}}); err != nil { + t.Fatalf("CreateActorSnapshot: %v", err) + } + tagWithoutAtespace := tag() + tagWithoutAtespace.Metadata.Atespace = "gone" + if _, err := s.CreateActorSnapshotTag(ctx, "gone", "snapshot", tagWithoutAtespace); !errors.Is(err, store.ErrFailedPrecondition) { + t.Errorf("missing tag atespace error = %v, want ErrFailedPrecondition", err) + } +} + +func TestAcquireLock_CleansExpiredLeases(t *testing.T) { + s := setupPostgresPersistence(t) + ctx := context.Background() + if _, err := s.pool.Exec(ctx, ` + INSERT INTO leases (key, token, expires_at) VALUES + ('expired', 'old', clock_timestamp() - interval '1 minute'), + ('active', 'live', clock_timestamp() + interval '1 hour')`); err != nil { + t.Fatalf("seeding leases: %v", err) + } + lock, err := s.AcquireLock(ctx, "new") + if err != nil { + t.Fatalf("AcquireLock: %v", err) + } + defer lock.Close() + + var expired, active int + if err := s.pool.QueryRow(ctx, `SELECT count(*) FROM leases WHERE key = 'expired'`).Scan(&expired); err != nil { + t.Fatalf("counting expired lease: %v", err) + } + if err := s.pool.QueryRow(ctx, `SELECT count(*) FROM leases WHERE key = 'active'`).Scan(&active); err != nil { + t.Fatalf("counting active lease: %v", err) + } + if expired != 0 || active != 1 { + t.Errorf("lease counts = expired:%d active:%d, want 0 and 1", expired, active) + } +} + // TestCreateActor_MissingAtespace_FailedPrecondition exercises the // foreign-key race the doc calls out: CreateActor rejects an actor whose // atespace doesn't exist (including a concurrently-deleted one), closing the @@ -284,6 +514,30 @@ func TestWorkerNotification_OnlyAfterCommit(t *testing.T) { } } +func TestWatchWorkers_MalformedNotificationClosesWatch(t *testing.T) { + s := setupPostgresStore(t).(*Persistence) + ctx := context.Background() + + watch, err := s.WatchWorkers(ctx) + if err != nil { + t.Fatalf("WatchWorkers failed: %v", err) + } + defer watch.Close() + + if _, err := s.pool.Exec(ctx, `SELECT pg_notify($1, $2)`, workerChangeChannel, "not-json"); err != nil { + t.Fatalf("pg_notify failed: %v", err) + } + + select { + case event, ok := <-watch.Events: + if ok { + t.Fatalf("received event %+v from malformed notification; want closed watch", event) + } + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for malformed notification to close watch") + } +} + func TestListActors_InvalidPageToken(t *testing.T) { s := setupPostgresStore(t).(*Persistence) ctx := context.Background() diff --git a/cmd/ateapi/internal/store/atepg/pagetoken.go b/cmd/ateapi/internal/store/atepg/pagetoken.go index 46e395e700..28e818a27e 100644 --- a/cmd/ateapi/internal/store/atepg/pagetoken.go +++ b/cmd/ateapi/internal/store/atepg/pagetoken.go @@ -18,6 +18,8 @@ import ( "encoding/base64" "encoding/json" "fmt" + + "github.com/agent-substrate/substrate/cmd/ateapi/internal/store" ) // pageTokenVersion guards against decoding a token produced by an incompatible @@ -61,23 +63,23 @@ func decodePageToken(tokenStr string, wantKind resourceKind, wantScope string, w } b, err := base64.StdEncoding.DecodeString(tokenStr) if err != nil { - return pageToken{}, fmt.Errorf("invalid page token: %w", err) + return pageToken{}, fmt.Errorf("%w: %v", store.ErrInvalidPageToken, err) } var token pageToken if err := json.Unmarshal(b, &token); err != nil { - return pageToken{}, fmt.Errorf("invalid page token: %w", err) + return pageToken{}, fmt.Errorf("%w: %v", store.ErrInvalidPageToken, err) } if token.Version != pageTokenVersion { - return pageToken{}, fmt.Errorf("invalid page token: unsupported version %d", token.Version) + return pageToken{}, fmt.Errorf("%w: unsupported version %d", store.ErrInvalidPageToken, token.Version) } if token.Kind != wantKind { - return pageToken{}, fmt.Errorf("invalid page token: for %q, used with %q", token.Kind, wantKind) + return pageToken{}, fmt.Errorf("%w: for %q, used with %q", store.ErrInvalidPageToken, token.Kind, wantKind) } if token.Scope != wantScope { - return pageToken{}, fmt.Errorf("invalid page token: for scope %q, used with scope %q", token.Scope, wantScope) + return pageToken{}, fmt.Errorf("%w: for scope %q, used with scope %q", store.ErrInvalidPageToken, token.Scope, wantScope) } if len(token.Last) != wantKeyParts { - return pageToken{}, fmt.Errorf("invalid page token: got %d key parts, want %d", len(token.Last), wantKeyParts) + return pageToken{}, fmt.Errorf("%w: got %d key parts, want %d", store.ErrInvalidPageToken, len(token.Last), wantKeyParts) } return token, nil } diff --git a/cmd/ateapi/internal/store/atepg/schema.go b/cmd/ateapi/internal/store/atepg/schema.go index f4f0ce0ad1..93a5f26ba4 100644 --- a/cmd/ateapi/internal/store/atepg/schema.go +++ b/cmd/ateapi/internal/store/atepg/schema.go @@ -36,7 +36,7 @@ CREATE TABLE IF NOT EXISTS actors ( atespace text NOT NULL REFERENCES atespaces(name) ON DELETE RESTRICT, name text NOT NULL, - uid text NOT NULL UNIQUE, + uid text NOT NULL, version bigint NOT NULL, proto bytea NOT NULL, PRIMARY KEY (atespace, name) @@ -46,7 +46,7 @@ CREATE TABLE IF NOT EXISTS actor_templates ( atespace text NOT NULL REFERENCES atespaces(name) ON DELETE RESTRICT, name text NOT NULL, - uid text NOT NULL UNIQUE, + uid text NOT NULL, version bigint NOT NULL, proto bytea NOT NULL, PRIMARY KEY (atespace, name) @@ -58,9 +58,12 @@ CREATE TABLE IF NOT EXISTS actor_template_versions ( name text NOT NULL, actor_template_atespace text NOT NULL, actor_template_name text NOT NULL, - uid text NOT NULL UNIQUE, + uid text NOT NULL, proto bytea NOT NULL, - PRIMARY KEY (atespace, name) + PRIMARY KEY (atespace, name), + CONSTRAINT actor_template_versions_parent_fk + FOREIGN KEY (actor_template_atespace, actor_template_name) + REFERENCES actor_templates(atespace, name) ON DELETE RESTRICT ); CREATE INDEX IF NOT EXISTS actor_template_versions_parent_idx @@ -74,18 +77,24 @@ CREATE TABLE IF NOT EXISTS actor_snapshots ( ); CREATE TABLE IF NOT EXISTS actor_snapshot_tags ( - atespace text NOT NULL - REFERENCES atespaces(name) ON DELETE RESTRICT, + atespace text NOT NULL, name text NOT NULL, snapshot_atespace text NOT NULL, snapshot_name text NOT NULL, + uid text NOT NULL, version bigint NOT NULL, proto bytea NOT NULL, PRIMARY KEY (atespace, name), - FOREIGN KEY (snapshot_atespace, snapshot_name) + CONSTRAINT actor_snapshot_tags_atespace_fk + FOREIGN KEY (atespace) REFERENCES atespaces(name) ON DELETE RESTRICT, + CONSTRAINT actor_snapshot_tags_snapshot_fk + FOREIGN KEY (snapshot_atespace, snapshot_name) REFERENCES actor_snapshots(atespace, name) ON DELETE RESTRICT ); +CREATE INDEX IF NOT EXISTS actor_snapshot_tags_snapshot_idx + ON actor_snapshot_tags (snapshot_atespace, snapshot_name); + CREATE TABLE IF NOT EXISTS workers ( worker_namespace text NOT NULL, worker_pool text NOT NULL, @@ -100,6 +109,8 @@ CREATE TABLE IF NOT EXISTS leases ( token text NOT NULL, expires_at timestamptz NOT NULL ); + +CREATE INDEX IF NOT EXISTS leases_expires_at_idx ON leases (expires_at); ` // applySchema idempotently creates atepg's tables. diff --git a/cmd/ateapi/internal/store/ateredis/ateredis.go b/cmd/ateapi/internal/store/ateredis/ateredis.go index a279627543..e61f596413 100644 --- a/cmd/ateapi/internal/store/ateredis/ateredis.go +++ b/cmd/ateapi/internal/store/ateredis/ateredis.go @@ -487,6 +487,14 @@ func (s *Persistence) DeleteActorTemplate(ctx context.Context, templateRef resou } func (s *Persistence) CreateActorTemplateVersion(ctx context.Context, atv *ateapipb.ActorTemplateVersion) (*ateapipb.ActorTemplateVersion, error) { + parent := resources.ActorTemplateRefFromObjectRef(atv.GetActorTemplate()) + exists, err := s.ActorTemplateExists(ctx, parent) + if err != nil { + return nil, fmt.Errorf("while checking actor template parent %s: %w", parent, err) + } + if !exists { + return nil, store.ErrFailedPrecondition + } dbKey := actorTemplateVersionDBKey(resources.ActorTemplateVersionRefFromActorTemplateVersion(atv)) dbVersion := proto.Clone(atv).(*ateapipb.ActorTemplateVersion) @@ -1294,9 +1302,14 @@ func (s *Persistence) ListActors(ctx context.Context, atespace string, opts stor // listPage SCANs pattern across the redis masters from the page token, feeding key batches to collect and returns the next-page token. func (s *Persistence) listPage(ctx context.Context, pattern string, pageSize int32, pageTokenStr string, collect func(ctx context.Context, master *redis.Client, keys []string) (int, error)) (string, error) { + normalized, err := store.NormalizeListOptions(store.ListOptions{PageSize: pageSize}) + if err != nil { + return "", err + } + pageSize = normalized.PageSize token, err := decodePageToken(pageTokenStr) if err != nil { - return "", fmt.Errorf("invalid page token: %w", err) + return "", fmt.Errorf("%w: %v", store.ErrInvalidPageToken, err) } masters, err := s.getSortedMasters(ctx) diff --git a/cmd/ateapi/internal/store/ateredis/ateredis_test.go b/cmd/ateapi/internal/store/ateredis/ateredis_test.go index 23b9a25366..97dfaf755b 100644 --- a/cmd/ateapi/internal/store/ateredis/ateredis_test.go +++ b/cmd/ateapi/internal/store/ateredis/ateredis_test.go @@ -2018,6 +2018,9 @@ func TestDeleteAtespace_WithActorTemplateVersions_Rejected(t *testing.T) { if _, err := s.CreateAtespace(ctx, newTestAtespace("team-a")); err != nil { t.Fatalf("CreateAtespace: %v", err) } + if _, err := s.CreateActorTemplate(ctx, newTestActorTemplate("team-a", "tmpl-a")); err != nil { + t.Fatalf("CreateActorTemplate: %v", err) + } if _, err := s.CreateActorTemplateVersion(ctx, newTestActorTemplateVersion("team-a", "tmpl-a-v1", "tmpl-a")); err != nil { t.Fatalf("CreateActorTemplateVersion: %v", err) } @@ -2028,6 +2031,9 @@ func TestDeleteAtespace_WithActorTemplateVersions_Rejected(t *testing.T) { if _, err := s.DeleteActorTemplateVersion(ctx, resources.ActorTemplateVersionRef{Atespace: "team-a", Name: "tmpl-a-v1"}); err != nil { t.Fatalf("DeleteActorTemplateVersion: %v", err) } + if _, err := s.DeleteActorTemplate(ctx, resources.ActorTemplateRef{Atespace: "team-a", Name: "tmpl-a"}); err != nil { + t.Fatalf("DeleteActorTemplate: %v", err) + } if _, err := s.DeleteAtespace(ctx, "team-a"); err != nil { t.Errorf("DeleteAtespace after version removed = %v, want nil", err) } @@ -2852,6 +2858,9 @@ func TestDeleteActorTemplate_HasVersions_Rejected(t *testing.T) { if _, err := s.CreateActorTemplate(ctx, newTestActorTemplate("team-a", "tmpl-a")); err != nil { t.Fatalf("CreateActorTemplate failed: %v", err) } + if _, err := s.CreateActorTemplate(ctx, newTestActorTemplate("team-a", "tmpl-b")); err != nil { + t.Fatalf("CreateActorTemplate failed: %v", err) + } // A version parented to a DIFFERENT template must not block the delete. if _, err := s.CreateActorTemplateVersion(ctx, newTestActorTemplateVersion("team-a", "tmpl-b-v1", "tmpl-b")); err != nil { t.Fatalf("CreateActorTemplateVersion failed: %v", err) @@ -2878,6 +2887,9 @@ func TestDeleteActorTemplate_HasVersions_Rejected(t *testing.T) { func TestActorTemplateVersionLifecycle(t *testing.T) { _, s, ctx := setupTest(t) + if _, err := s.CreateActorTemplate(ctx, newTestActorTemplate("team-a", "tmpl-a")); err != nil { + t.Fatalf("CreateActorTemplate failed: %v", err) + } want := newTestActorTemplateVersion("team-a", "tmpl-a-v1", "tmpl-a") created, err := s.CreateActorTemplateVersion(ctx, want) @@ -2917,6 +2929,9 @@ func TestActorTemplateVersionLifecycle(t *testing.T) { func TestCreateActorTemplateVersion_AlreadyExists(t *testing.T) { _, s, ctx := setupTest(t) + if _, err := s.CreateActorTemplate(ctx, newTestActorTemplate("team-a", "tmpl-a")); err != nil { + t.Fatalf("CreateActorTemplate failed: %v", err) + } if _, err := s.CreateActorTemplateVersion(ctx, newTestActorTemplateVersion("team-a", "v1", "tmpl-a")); err != nil { t.Fatalf("first CreateActorTemplateVersion failed: %v", err) @@ -2962,19 +2977,19 @@ func TestDeleteActorTemplateVersion_IsParentDefault_Rejected(t *testing.T) { } } -func TestDeleteActorTemplateVersion_MissingParent_Allowed(t *testing.T) { +func TestCreateActorTemplateVersion_MissingParent_Rejected(t *testing.T) { _, s, ctx := setupTest(t) - if _, err := s.CreateActorTemplateVersion(ctx, newTestActorTemplateVersion("team-a", "orphan-v1", "gone")); err != nil { - t.Fatalf("CreateActorTemplateVersion failed: %v", err) - } - if _, err := s.DeleteActorTemplateVersion(ctx, resources.ActorTemplateVersionRef{Atespace: "team-a", Name: "orphan-v1"}); err != nil { - t.Errorf("DeleteActorTemplateVersion with missing parent = %v, want nil", err) + if _, err := s.CreateActorTemplateVersion(ctx, newTestActorTemplateVersion("team-a", "orphan-v1", "gone")); !errors.Is(err, store.ErrFailedPrecondition) { + t.Fatalf("CreateActorTemplateVersion with missing parent = %v, want ErrFailedPrecondition", err) } } func TestDeleteActorTemplateVersion_DeletesGoldenSnapshot(t *testing.T) { _, s, ctx := setupTest(t) + if _, err := s.CreateActorTemplate(ctx, newTestActorTemplate("team-a", "tmpl-a")); err != nil { + t.Fatalf("CreateActorTemplate failed: %v", err) + } if _, err := s.CreateActorSnapshot(ctx, &ateapipb.ActorSnapshot{ Metadata: &ateapipb.ResourceMetadata{Atespace: "ate-golden", Name: "golden-1"}, @@ -2998,6 +3013,9 @@ func TestDeleteActorTemplateVersion_DeletesGoldenSnapshot(t *testing.T) { func TestDeleteActorTemplateVersion_GoldenSnapshotAlreadyGone(t *testing.T) { _, s, ctx := setupTest(t) + if _, err := s.CreateActorTemplate(ctx, newTestActorTemplate("team-a", "tmpl-a")); err != nil { + t.Fatalf("CreateActorTemplate failed: %v", err) + } version := newTestActorTemplateVersion("team-a", "tmpl-a-v1", "tmpl-a") version.GoldenSnapshot = &ateapipb.ObjectRef{Atespace: "ate-golden", Name: "never-created"} @@ -3046,6 +3064,11 @@ func TestListActorTemplates_Pagination(t *testing.T) { func TestListActorTemplateVersions_ParentFilter(t *testing.T) { _, s, ctx := setupTest(t) + for _, name := range []string{"tmpl-a", "tmpl-b"} { + if _, err := s.CreateActorTemplate(ctx, newTestActorTemplate("team-a", name)); err != nil { + t.Fatalf("CreateActorTemplate(%s) failed: %v", name, err) + } + } // Interleave versions of two templates. for i := 0; i < 3; i++ { @@ -3100,6 +3123,9 @@ func TestListActorTemplateVersions_ParentFilter(t *testing.T) { // The filter matches the parent's atespace too: scanning all atespaces // with team-a's tmpl-a must not pick up team-b versions whose parent // merely shares the name. + if _, err := s.CreateActorTemplate(ctx, newTestActorTemplate("team-b", "tmpl-a")); err != nil { + t.Fatalf("failed to create team-b template: %v", err) + } if _, err := s.CreateActorTemplateVersion(ctx, newTestActorTemplateVersion("team-b", "tmpl-a-v0", "tmpl-a")); err != nil { t.Fatalf("failed to create team-b version: %v", err) } @@ -3189,6 +3215,11 @@ func TestListActorTemplates_AtespaceFilter(t *testing.T) { func TestActorTemplateVersions_AtespaceIsolation(t *testing.T) { _, s, ctx := setupTest(t) + for _, atespace := range []string{"team-a", "team-b"} { + if _, err := s.CreateActorTemplate(ctx, newTestActorTemplate(atespace, "tmpl")); err != nil { + t.Fatalf("CreateActorTemplate in %s failed: %v", atespace, err) + } + } if _, err := s.CreateActorTemplateVersion(ctx, newTestActorTemplateVersion("team-a", "tmpl-v1", "tmpl")); err != nil { t.Fatalf("CreateActorTemplateVersion in team-a failed: %v", err) @@ -3227,6 +3258,9 @@ func TestDeleteActorTemplate_VersionInOtherAtespace_NotBlocking(t *testing.T) { if _, err := s.CreateActorTemplate(ctx, newTestActorTemplate("team-a", "tmpl-a")); err != nil { t.Fatalf("CreateActorTemplate failed: %v", err) } + if _, err := s.CreateActorTemplate(ctx, newTestActorTemplate("team-b", "tmpl-a")); err != nil { + t.Fatalf("CreateActorTemplate failed: %v", err) + } // A version of a same-named template in ANOTHER atespace must not block // the delete. if _, err := s.CreateActorTemplateVersion(ctx, newTestActorTemplateVersion("team-b", "tmpl-a-v1", "tmpl-a")); err != nil { diff --git a/cmd/ateapi/internal/store/store.go b/cmd/ateapi/internal/store/store.go index 54157bc140..50f7073d78 100644 --- a/cmd/ateapi/internal/store/store.go +++ b/cmd/ateapi/internal/store/store.go @@ -42,6 +42,13 @@ var ( // ErrLockConflict indicates that a distributed lock is already held by another client. ErrLockConflict = errors.New("persistence: lock conflict") + // ErrInvalidPageToken indicates that a list page token is malformed or was + // issued for a different list operation or scope. + ErrInvalidPageToken = errors.New("persistence: invalid page token") + + // ErrInvalidPageSize indicates that a negative page size was supplied. + ErrInvalidPageSize = errors.New("persistence: invalid page size") + // ErrUIDConflict indicates a precondition pinned a uid the stored object does // not carry, meaning the name now addresses a different incarnation. Retrying // can never resolve it. @@ -52,7 +59,8 @@ var ( type Interface interface { // Stores a new actor in suspended state and returns the stored resource with // server-assigned metadata (uid, version, timestamps). The input is not - // mutated. Returns ErrAlreadyExists if key is taken. + // mutated. Returns ErrAlreadyExists if key is taken, or + // ErrFailedPrecondition if the actor's atespace does not exist. CreateActor(ctx context.Context, actor *ateapipb.Actor) (*ateapipb.Actor, error) // Fetches an actor by reference. Returns ErrNotFound if missing. @@ -90,7 +98,9 @@ type Interface interface { // Lists ActorSnapshots in one atespace, or all atespaces when empty. ListActorSnapshots(ctx context.Context, atespace string, opts ListOptions) (ListResponse[*ateapipb.ActorSnapshot], error) - // Adds an immutable Atespace-owned tag to an ActorSnapshot. + // Adds an immutable Atespace-owned tag to an ActorSnapshot. Returns + // ErrNotFound if the snapshot does not exist, or ErrFailedPrecondition if + // the tag's atespace does not exist. CreateActorSnapshotTag(ctx context.Context, atespace, name string, tag *ateapipb.ActorSnapshotTag) (*ateapipb.ActorSnapshotTag, error) // Fetches an Atespace-owned tag. Returns ErrNotFound if missing. The tag's @@ -160,9 +170,9 @@ type Interface interface { DeleteActorTemplate(ctx context.Context, templateRef resources.ActorTemplateRef) (*ateapipb.ActorTemplate, error) // Stores a new ActorTemplateVersion and returns the stored resource with - // server-assigned metadata. The caller is responsible for the - // parent-exists check and for initializing the status fields. The input is not - // mutated. Returns ErrAlreadyExists if the (atespace, name) is taken. + // server-assigned metadata. The caller initializes the status fields. The + // input is not mutated. Returns ErrAlreadyExists if the (atespace, name) is + // taken, or ErrFailedPrecondition if the parent ActorTemplate is missing. CreateActorTemplateVersion(ctx context.Context, version *ateapipb.ActorTemplateVersion) (*ateapipb.ActorTemplateVersion, error) // Fetches an ActorTemplateVersion by reference. Returns ErrNotFound if @@ -171,7 +181,7 @@ type Interface interface { // Lists ActorTemplateVersions in an atespace (all atespaces when atespace // is empty), filtered to one parent template when actorTemplateRef is - // non-zero. The parent lives in the same atespace as its versions. + // non-zero. The parent filter is a fully qualified reference. ListActorTemplateVersions(ctx context.Context, atespace string, actorTemplateRef resources.ActorTemplateRef, opts ListOptions) (ListResponse[*ateapipb.ActorTemplateVersion], error) // Removes an ActorTemplateVersion and returns the deleted resource, also @@ -328,6 +338,22 @@ type ListOptions struct { PageToken string } +// DefaultPageSize is used by store implementations when PageSize is unset. +const DefaultPageSize int32 = 1000 + +// NormalizeListOptions applies the store default and rejects invalid sizes. +// RPC handlers validate user input separately, but store callers also need a +// safe contract because list implementations use PageSize in slice indexes. +func NormalizeListOptions(opts ListOptions) (ListOptions, error) { + if opts.PageSize < 0 { + return ListOptions{}, ErrInvalidPageSize + } + if opts.PageSize == 0 { + opts.PageSize = DefaultPageSize + } + return opts, nil +} + // ListResponse is the return value of a List method: the page of items it // addressed, plus the token to fetch the next page. NextPageToken is empty // once the listing has reached its last page. diff --git a/cmd/ateapi/internal/store/store_test.go b/cmd/ateapi/internal/store/store_test.go index 38609b1275..7673362f5f 100644 --- a/cmd/ateapi/internal/store/store_test.go +++ b/cmd/ateapi/internal/store/store_test.go @@ -145,3 +145,16 @@ func TestWithPrecondition(t *testing.T) { }) } } + +func TestNormalizeListOptions(t *testing.T) { + got, err := NormalizeListOptions(ListOptions{}) + if err != nil { + t.Fatalf("NormalizeListOptions(zero) error = %v", err) + } + if got.PageSize != DefaultPageSize { + t.Errorf("NormalizeListOptions(zero).PageSize = %d, want %d", got.PageSize, DefaultPageSize) + } + if _, err := NormalizeListOptions(ListOptions{PageSize: -1}); !errors.Is(err, ErrInvalidPageSize) { + t.Errorf("NormalizeListOptions(negative) error = %v, want ErrInvalidPageSize", err) + } +} diff --git a/cmd/ateapi/internal/store/storecontract/contract.go b/cmd/ateapi/internal/store/storecontract/contract.go index fe3855a8f2..d0da87e6ae 100644 --- a/cmd/ateapi/internal/store/storecontract/contract.go +++ b/cmd/ateapi/internal/store/storecontract/contract.go @@ -107,9 +107,69 @@ func RunContractTests(t *testing.T, setup func(t *testing.T) store.Interface) { runActorTemplateContractTests(t, setup) runActorSnapshotContractTests(t, setup) runLockContractTests(t, setup) + runListOptionsContractTests(t, setup) runDebugContractTests(t, setup) } +func runListOptionsContractTests(t *testing.T, setup func(t *testing.T) store.Interface) { + t.Helper() + + t.Run("ListOptions_InvalidPageSize", func(t *testing.T) { + s := setup(t) + ctx := context.Background() + calls := []struct { + name string + call func(store.ListOptions) error + }{ + {"atespaces", func(opts store.ListOptions) error { _, err := s.ListAtespaces(ctx, opts); return err }}, + {"actors", func(opts store.ListOptions) error { _, err := s.ListActors(ctx, "", opts); return err }}, + {"actor templates", func(opts store.ListOptions) error { _, err := s.ListActorTemplates(ctx, "", opts); return err }}, + {"actor template versions", func(opts store.ListOptions) error { + _, err := s.ListActorTemplateVersions(ctx, "", resources.ActorTemplateRef{}, opts) + return err + }}, + {"actor snapshots", func(opts store.ListOptions) error { _, err := s.ListActorSnapshots(ctx, "", opts); return err }}, + {"workers", func(opts store.ListOptions) error { _, err := s.ListWorkers(ctx, opts); return err }}, + } + for _, call := range calls { + t.Run(call.name, func(t *testing.T) { + if err := call.call(store.ListOptions{PageSize: -1}); !errors.Is(err, store.ErrInvalidPageSize) { + t.Errorf("negative PageSize error = %v, want ErrInvalidPageSize", err) + } + if err := call.call(store.ListOptions{}); err != nil { + t.Errorf("zero PageSize error = %v, want nil", err) + } + }) + } + }) + + t.Run("ListOptions_InvalidPageToken", func(t *testing.T) { + s := setup(t) + ctx := context.Background() + calls := []struct { + name string + call func(store.ListOptions) error + }{ + {"atespaces", func(opts store.ListOptions) error { _, err := s.ListAtespaces(ctx, opts); return err }}, + {"actors", func(opts store.ListOptions) error { _, err := s.ListActors(ctx, "", opts); return err }}, + {"actor templates", func(opts store.ListOptions) error { _, err := s.ListActorTemplates(ctx, "", opts); return err }}, + {"actor template versions", func(opts store.ListOptions) error { + _, err := s.ListActorTemplateVersions(ctx, "", resources.ActorTemplateRef{}, opts) + return err + }}, + {"actor snapshots", func(opts store.ListOptions) error { _, err := s.ListActorSnapshots(ctx, "", opts); return err }}, + {"workers", func(opts store.ListOptions) error { _, err := s.ListWorkers(ctx, opts); return err }}, + } + for _, call := range calls { + t.Run(call.name, func(t *testing.T) { + if err := call.call(store.ListOptions{PageSize: 1, PageToken: "%%%"}); !errors.Is(err, store.ErrInvalidPageToken) { + t.Errorf("malformed PageToken error = %v, want ErrInvalidPageToken", err) + } + }) + } + }) +} + func runActorContractTests(t *testing.T, setup func(t *testing.T) store.Interface) { t.Helper() @@ -652,6 +712,15 @@ func runActorTemplateContractTests(t *testing.T, setup func(t *testing.T) store. for _, atespace := range []string{"team-a", "team-b"} { mustCreateAtespace(t, s, atespace) } + for _, template := range []struct{ atespace, name string }{ + {"team-a", "tmpl-a"}, + {"team-a", "tmpl-b"}, + {"team-b", "tmpl-a"}, + } { + if _, err := s.CreateActorTemplate(ctx, newTestActorTemplate(template.atespace, template.name)); err != nil { + t.Fatalf("CreateActorTemplate(%s/%s) failed: %v", template.atespace, template.name, err) + } + } for _, item := range []struct{ atespace, name, parent string }{ {"team-a", "a-1", "tmpl-a"}, {"team-a", "a-2", "tmpl-a"}, @@ -696,11 +765,11 @@ func runActorTemplateContractTests(t *testing.T, setup func(t *testing.T) store. if _, err := s.DeleteActorTemplate(ctx, resources.ActorTemplateRef{Atespace: "team-a", Name: "tmpl-a"}); err != nil { t.Fatalf("DeleteActorTemplate failed: %v", err) } - if _, err := s.CreateActorTemplateVersion(ctx, newTestActorTemplateVersion("team-a", "orphan-v1", "gone")); err != nil { - t.Fatalf("CreateActorTemplateVersion failed: %v", err) + if _, err := s.CreateActorTemplateVersion(ctx, newTestActorTemplateVersion("team-a", "orphan-v1", "gone")); !errors.Is(err, store.ErrFailedPrecondition) { + t.Fatalf("CreateActorTemplateVersion without parent = %v, want ErrFailedPrecondition", err) } - if _, err := s.DeleteAtespace(ctx, "team-a"); !errors.Is(err, store.ErrFailedPrecondition) { - t.Errorf("DeleteAtespace with version = %v, want ErrFailedPrecondition", err) + if _, err := s.DeleteAtespace(ctx, "team-a"); err != nil { + t.Errorf("DeleteAtespace after rejected orphan version = %v, want nil", err) } }) @@ -708,6 +777,9 @@ func runActorTemplateContractTests(t *testing.T, setup func(t *testing.T) store. s := setup(t) ctx := context.Background() mustCreateAtespace(t, s, "team-a") + if _, err := s.CreateActorTemplate(ctx, newTestActorTemplate("team-a", "tmpl-a")); err != nil { + t.Fatalf("CreateActorTemplate failed: %v", err) + } if _, err := s.CreateActorSnapshot(ctx, &ateapipb.ActorSnapshot{ Metadata: &ateapipb.ResourceMetadata{Atespace: "ate-golden", Name: "golden-1"}, SnapshotUri: "gs://bucket/golden-1", diff --git a/cmd/ateapi/internal/workercache/workercache.go b/cmd/ateapi/internal/workercache/workercache.go index 56b1583b47..1408f69b3a 100644 --- a/cmd/ateapi/internal/workercache/workercache.go +++ b/cmd/ateapi/internal/workercache/workercache.go @@ -110,6 +110,16 @@ func (c *Cache) Worker(namespace, pod string) (*ateapipb.Worker, error) { return worker, nil } +// Forget removes a worker that a store write proved no longer exists. The +// normal delete watch remains authoritative, but this closes the short race in +// which scheduling selected a worker just after its row was deleted and before +// the watch event reached the cache. +func (c *Cache) Forget(namespace, pod string) { + c.mu.Lock() + defer c.mu.Unlock() + delete(c.workers, namespace+":"+pod) +} + func (c *Cache) sync(ctx context.Context) (*store.WorkerWatch, error) { watch, err := c.store.WatchWorkers(ctx) if err != nil { diff --git a/cmd/ateapi/main.go b/cmd/ateapi/main.go index 5d1d3e0d5a..345ea631f0 100644 --- a/cmd/ateapi/main.go +++ b/cmd/ateapi/main.go @@ -43,6 +43,7 @@ import ( "github.com/agent-substrate/substrate/pkg/client/clientset/versioned" "github.com/agent-substrate/substrate/pkg/client/informers/externalversions" "github.com/agent-substrate/substrate/pkg/proto/ateapipb" + "github.com/jackc/pgx/v5/pgxpool" "github.com/redis/go-redis/v9" "github.com/spf13/pflag" "go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc" @@ -133,7 +134,7 @@ func main() { serverboot.Fatal(ctx, "Failed to initialize JWT providers", err) } - persistence, err := connectStore(ctx) + persistence, err := connectStore(shutdownCtx) if err != nil { serverboot.Fatal(ctx, "Failed to set up persistence backend", err) } @@ -331,7 +332,10 @@ func connectStore(ctx context.Context) (store.Interface, error) { if *postgresConnectionString == "" { return nil, fmt.Errorf("--store-backend=postgres requires --postgres-connection-string") } - persistence, err := atepg.Connect(ctx, *postgresConnectionString) + if _, err := pgxpool.ParseConfig(*postgresConnectionString); err != nil { + return nil, fmt.Errorf("parsing PostgreSQL connection string: %w", err) + } + persistence, err := connectPostgresWithRetries(ctx) if err != nil { return nil, fmt.Errorf("setting up PostgreSQL: %w", err) } @@ -341,6 +345,32 @@ func connectStore(ctx context.Context) (store.Interface, error) { } } +var ( + postgresConnectTries = 30 + postgresConnectPeriod = 2 * time.Second +) + +func connectPostgresWithRetries(ctx context.Context) (*atepg.Persistence, error) { + var connectErr error + for attempt := 1; attempt <= postgresConnectTries; attempt++ { + persistence, err := atepg.Connect(ctx, *postgresConnectionString) + if err == nil { + return persistence, nil + } + connectErr = err + slog.WarnContext(ctx, "Failed to connect to PostgreSQL, retrying...", slog.Int("attempt", attempt), slog.Any("err", err)) + if attempt == postgresConnectTries { + break + } + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-time.After(postgresConnectPeriod): + } + } + return nil, fmt.Errorf("connect to PostgreSQL after %d attempts: %w", postgresConnectTries, connectErr) +} + // connectRedis builds the Redis/Valkey TLS config, plumbs IAM auth if // requested, opens the cluster client, and pings with retries. func connectRedis(ctx context.Context) (*redis.ClusterClient, error) {