diff --git a/.github/workflows/golangci-lint.yml b/.github/workflows/golangci-lint.yml index 458931293..e0cbccc4e 100644 --- a/.github/workflows/golangci-lint.yml +++ b/.github/workflows/golangci-lint.yml @@ -18,7 +18,18 @@ jobs: steps: - name: Check out code into the Go module directory uses: actions/checkout@v7 + - name: Set up Go + uses: actions/setup-go@v6 + with: + go-version: "1.x" + - name: Install golangci-lint + run: | + curl -sSfL https://raw.githubusercontent.com/golangci/golangci-lint/master/install.sh | sh -s -- -b "$(go env GOPATH)/bin" v2.9.0 + echo "$(go env GOPATH)/bin" >> "$GITHUB_PATH" + - name: Run golangci-lint + run: golangci-lint run --config=.golangci.yaml - name: golangci-lint + continue-on-error: true uses: reviewdog/action-golangci-lint@v2 with: github_token: ${{ secrets.GITHUB_TOKEN }} @@ -28,4 +39,4 @@ jobs: filter_mode: nofilter reporter: github-pr-review cache: false - fail_on_error: true + fail_on_error: false diff --git a/.github/workflows/redis-proxy-docker.yml b/.github/workflows/redis-proxy-docker.yml index 6e79ebe25..b79ad0454 100644 --- a/.github/workflows/redis-proxy-docker.yml +++ b/.github/workflows/redis-proxy-docker.yml @@ -45,13 +45,29 @@ jobs: - name: Docker metadata id: meta - uses: docker/metadata-action@v6 - with: - images: ghcr.io/${{ github.repository }}/redis-proxy - tags: | - type=sha - type=ref,event=branch - type=raw,value=latest,enable={{is_default_branch}} + shell: bash + env: + IMAGE: ghcr.io/${{ github.repository }}/redis-proxy + DEFAULT_BRANCH: ${{ github.event.repository.default_branch }} + run: | + short_sha="${GITHUB_SHA::7}" + ref_tag="$(printf '%s' "$GITHUB_REF_NAME" | sed -E 's/[^A-Za-z0-9_.-]+/-/g; s/^-+//; s/-+$//')" + ref_tag="${ref_tag:0:128}" + { + echo "tags<> "$GITHUB_OUTPUT" - name: Build and push uses: docker/build-push-action@v7 diff --git a/adapter/distribution_server.go b/adapter/distribution_server.go index 7e4110a99..3598c2101 100644 --- a/adapter/distribution_server.go +++ b/adapter/distribution_server.go @@ -29,6 +29,7 @@ type DistributionServer struct { watchInterval time.Duration watchLeader func() bool fsObserver DistributionFilesystemObserver + readBlocked func() bool reloadRetry struct { attempts int interval time.Duration @@ -65,6 +66,12 @@ func WithDistributionFilesystemObserver(observer DistributionFilesystemObserver) } } +func WithDistributionReadGate(blocked func() bool) DistributionServerOption { + return func(s *DistributionServer) { + s.readBlocked = blocked + } +} + // WithCatalogReloadRetryPolicy configures the retry policy used after split // commit when waiting for the local catalog snapshot to become visible. func WithCatalogReloadRetryPolicy(attempts int, interval time.Duration) DistributionServerOption { @@ -136,6 +143,20 @@ func NewDistributionServer(e *distribution.Engine, catalog *distribution.Catalog return s } +func (s *DistributionServer) SetReadGate(blocked func() bool) { + if s != nil { + s.readBlocked = blocked + } +} + +func (s *DistributionServer) requireReadReady() error { + if s != nil && s.readBlocked != nil && s.readBlocked() { + //nolint:wrapcheck // Preserve the gRPC status code for startup readers. + return status.Error(codes.Unavailable, "distribution startup has not completed") + } + return nil +} + // UpdateRoute allows updating route information. func (s *DistributionServer) UpdateRoute(start, end []byte, group uint64) { s.engine.UpdateRoute(start, end, group) @@ -143,6 +164,9 @@ func (s *DistributionServer) UpdateRoute(start, end []byte, group uint64) { // GetRoute returns route for a key. func (s *DistributionServer) GetRoute(ctx context.Context, req *pb.GetRouteRequest) (*pb.GetRouteResponse, error) { + if err := s.requireReadReady(); err != nil { + return nil, err + } r, ok := s.engine.GetRoute(kv.RouteKey(req.Key)) if !ok { return &pb.GetRouteResponse{}, nil @@ -162,6 +186,9 @@ func (s *DistributionServer) GetTimestamp(ctx context.Context, req *pb.GetTimest // ListRoutes returns all durable routes from catalog storage. func (s *DistributionServer) ListRoutes(ctx context.Context, req *pb.ListRoutesRequest) (*pb.ListRoutesResponse, error) { + if err := s.requireReadReady(); err != nil { + return nil, err + } snapshot, err := s.loadCatalogSnapshot(ctx) if err != nil { return nil, err @@ -352,32 +379,19 @@ func (s *DistributionServer) SplitRange(ctx context.Context, req *pb.SplitRangeR return nil, err } - parent, found := findRouteByID(snapshot.Routes, req.GetRouteId()) - if !found { - return nil, grpcStatusError(codes.NotFound, errDistributionUnknownRoute.Error()) - } - - rawSplitKey := req.GetSplitKey() - splitKey := distribution.CloneBytes(fskeys.NormalizeSplitBoundary(kv.RouteKey(rawSplitKey))) - if err := validateSplitKey(parent, splitKey); err != nil { - s.observeFilePinnedHotspotIfNeeded(rawSplitKey, splitKey, err) - return nil, err - } - - leftID, rightID, err := s.allocateChildRouteIDs(ctx, snapshot.ReadTS, snapshot.Routes) + plan, err := s.planSplitRange(ctx, snapshot, req) if err != nil { return nil, err } - left, right := splitCatalogRoutes(parent, splitKey, leftID, rightID, 0) - saved, err := s.saveSplitResultViaCoordinator(ctx, snapshot.ReadTS, req.GetExpectedCatalogVersion(), parent.RouteID, left, right) + saved, err := s.saveSplitResultViaCoordinator(ctx, snapshot.ReadTS, req.GetExpectedCatalogVersion(), plan.parentID, plan.readKeys, plan.left, plan.right) if err != nil { return nil, err } if err := s.applyEngineSnapshot(saved); err != nil { return nil, err } - savedLeft, savedRight, err := splitChildrenFromSnapshot(saved, left.RouteID, right.RouteID) + savedLeft, savedRight, err := splitChildrenFromSnapshot(saved, plan.left.RouteID, plan.right.RouteID) if err != nil { return nil, err } @@ -389,6 +403,37 @@ func (s *DistributionServer) SplitRange(ctx context.Context, req *pb.SplitRangeR }, nil } +type splitRangePlan struct { + parentID uint64 + readKeys [][]byte + left distribution.RouteDescriptor + right distribution.RouteDescriptor +} + +func (s *DistributionServer) planSplitRange(ctx context.Context, snapshot distribution.CatalogSnapshot, req *pb.SplitRangeRequest) (splitRangePlan, error) { + parent, found := findRouteByID(snapshot.Routes, req.GetRouteId()) + if !found { + return splitRangePlan{}, grpcStatusError(codes.NotFound, errDistributionUnknownRoute.Error()) + } + + rawSplitKey := req.GetSplitKey() + splitKey := distribution.CloneBytes(fskeys.NormalizeSplitBoundary(kv.RouteKey(rawSplitKey))) + if err := validateSplitKey(parent, splitKey); err != nil { + s.observeFilePinnedHotspotIfNeeded(rawSplitKey, splitKey, err) + return splitRangePlan{}, err + } + readKeys, err := s.splitJobOverlapReadKeys(ctx, snapshot, parent) + if err != nil { + return splitRangePlan{}, err + } + leftID, rightID, err := s.allocateChildRouteIDs(ctx, snapshot.ReadTS, snapshot.Routes) + if err != nil { + return splitRangePlan{}, err + } + left, right := splitCatalogRoutes(parent, splitKey, leftID, rightID, 0) + return splitRangePlan{parentID: parent.RouteID, readKeys: readKeys, left: left, right: right}, nil +} + func (s *DistributionServer) pinReadTS(ts uint64) *kv.ActiveTimestampToken { if s == nil || s.readTracker == nil { return nil @@ -437,17 +482,14 @@ func (s *DistributionServer) saveSplitResultViaCoordinator( readTS uint64, expectedVersion uint64, parentID uint64, + readKeys [][]byte, left distribution.RouteDescriptor, right distribution.RouteDescriptor, ) (distribution.CatalogSnapshot, error) { - if expectedVersion == math.MaxUint64 { - return distribution.CatalogSnapshot{}, grpcStatusError(codes.Internal, "catalog version overflow") - } - nextVersion := expectedVersion + 1 - if right.RouteID == math.MaxUint64 { - return distribution.CatalogSnapshot{}, grpcStatusError(codes.Internal, errDistributionRouteIDOverflow.Error()) + nextVersion, nextRouteID, err := nextCatalogSplitIDs(expectedVersion, right.RouteID) + if err != nil { + return distribution.CatalogSnapshot{}, err } - nextRouteID := right.RouteID + 1 commitTS, err := kv.NextTimestampAfterThrough(ctx, s.coordinator, readTS, "split range: allocate commitTS") if err != nil { return distribution.CatalogSnapshot{}, grpcStatusErrorf(codes.Internal, "allocate split commit timestamp: %v", err) @@ -482,8 +524,12 @@ func (s *DistributionServer) saveSplitResultViaCoordinator( IsTxn: true, StartTS: readTS, CommitTS: commitTS, + ReadKeys: readKeys, }) if err != nil { + if errors.Is(err, store.ErrWriteConflict) { + return distribution.CatalogSnapshot{}, grpcStatusError(codes.Aborted, errDistributionCatalogConflict.Error()) + } return distribution.CatalogSnapshot{}, grpcStatusErrorf(codes.Internal, "commit split mutations: %v", err) } if resp == nil || resp.CommitTS == 0 { @@ -492,6 +538,16 @@ func (s *DistributionServer) saveSplitResultViaCoordinator( return s.loadCatalogSnapshotAtVersion(ctx, resp.CommitTS, nextVersion) } +func nextCatalogSplitIDs(expectedVersion uint64, rightRouteID uint64) (uint64, uint64, error) { + if expectedVersion == math.MaxUint64 { + return 0, 0, grpcStatusError(codes.Internal, "catalog version overflow") + } + if rightRouteID == math.MaxUint64 { + return 0, 0, grpcStatusError(codes.Internal, errDistributionRouteIDOverflow.Error()) + } + return expectedVersion + 1, rightRouteID + 1, nil +} + func catalogStoreMutationsToOps(mutations []*store.KVPairMutation) ([]*kv.Elem[kv.OP], error) { ops := make([]*kv.Elem[kv.OP], 0, len(mutations)) for _, mutation := range mutations { @@ -672,6 +728,81 @@ func validateSplitKey(parent distribution.RouteDescriptor, splitKey []byte) erro return nil } +func (s *DistributionServer) splitJobOverlapReadKeys(ctx context.Context, snapshot distribution.CatalogSnapshot, parent distribution.RouteDescriptor) ([][]byte, error) { + jobs, err := s.catalog.ListSplitJobsAt(ctx, snapshot.ReadTS) + if err != nil { + return nil, grpcStatusErrorf(codes.Internal, "load split jobs: %v", err) + } + readKeys := splitJobReadFenceKeys(jobs) + for _, job := range jobs { + if !splitJobIsLive(job) { + continue + } + for _, interval := range liveSplitJobIntervals(job, snapshot.Routes) { + if routeRangeIntersects(parent.Start, parent.End, interval.start, interval.end) { + return nil, grpcStatusError(codes.Aborted, distribution.ErrSplitJobOverlap.Error()) + } + } + } + return readKeys, nil +} + +func splitJobReadFenceKeys(jobs []distribution.SplitJob) [][]byte { + readKeys := make([][]byte, 0, len(jobs)+1) + readKeys = append(readKeys, distribution.CatalogNextSplitJobIDKey()) + for _, job := range jobs { + if splitJobIsLive(job) { + readKeys = append(readKeys, distribution.CatalogSplitJobKey(job.JobID)) + } + } + return readKeys +} + +func splitJobIsLive(job distribution.SplitJob) bool { + return job.Phase != distribution.SplitJobPhaseDone && job.Phase != distribution.SplitJobPhaseAbandoned +} + +type routeInterval struct { + start []byte + end []byte +} + +const initialLiveSplitJobIntervalCapacity = 2 + +func liveSplitJobIntervals(job distribution.SplitJob, routes []distribution.RouteDescriptor) []routeInterval { + out := make([]routeInterval, 0, initialLiveSplitJobIntervalCapacity) + for _, route := range routes { + switch { + case route.RouteID == job.SourceRouteID: + out = append(out, routeInterval{ + start: distribution.CloneBytes(job.SplitKey), + end: distribution.CloneBytes(route.End), + }) + case route.ParentRouteID == job.SourceRouteID && routeRangeIntersects(route.Start, route.End, job.SplitKey, nil): + out = append(out, routeInterval{ + start: distribution.CloneBytes(route.Start), + end: distribution.CloneBytes(route.End), + }) + case job.JobID != 0 && route.MigrationJobID == job.JobID: + out = append(out, routeInterval{ + start: distribution.CloneBytes(route.Start), + end: distribution.CloneBytes(route.End), + }) + } + } + return out +} + +func routeRangeIntersects(aStart, aEnd, bStart, bEnd []byte) bool { + if aEnd != nil && bytes.Compare(aEnd, bStart) <= 0 { + return false + } + if bEnd != nil && bytes.Compare(bEnd, aStart) <= 0 { + return false + } + return true +} + func splitCatalogRoutes( parent distribution.RouteDescriptor, splitKey []byte, @@ -687,10 +818,10 @@ func splitCatalogRoutes( GroupID: parent.GroupID, State: parent.State, ParentRouteID: parent.RouteID, - SplitAtHLC: splitAtHLC, StagedVisibilityActive: parent.StagedVisibilityActive, MigrationJobID: parent.MigrationJobID, MinWriteTSExclusive: parent.MinWriteTSExclusive, + SplitAtHLC: splitAtHLC, } right := distribution.RouteDescriptor{ RouteID: rightID, @@ -699,10 +830,10 @@ func splitCatalogRoutes( GroupID: parent.GroupID, State: parent.State, ParentRouteID: parent.RouteID, - SplitAtHLC: splitAtHLC, StagedVisibilityActive: parent.StagedVisibilityActive, MigrationJobID: parent.MigrationJobID, MinWriteTSExclusive: parent.MinWriteTSExclusive, + SplitAtHLC: splitAtHLC, } return left, right } @@ -758,10 +889,10 @@ func toProtoRouteDescriptor(route distribution.RouteDescriptor) *pb.RouteDescrip RaftGroupId: route.GroupID, State: toProtoRouteState(route.State), ParentRouteId: route.ParentRouteID, - SplitAtHlc: route.SplitAtHLC, StagedVisibilityActive: route.StagedVisibilityActive, MigrationJobId: route.MigrationJobID, MinWriteTsExclusive: route.MinWriteTSExclusive, + SplitAtHlc: route.SplitAtHLC, } } diff --git a/adapter/distribution_server_test.go b/adapter/distribution_server_test.go index 514a9608a..dc06696b0 100644 --- a/adapter/distribution_server_test.go +++ b/adapter/distribution_server_test.go @@ -1,6 +1,7 @@ package adapter import ( + "bytes" "context" "encoding/binary" "testing" @@ -58,6 +59,35 @@ func TestDistributionServerGetRoute_NormalizesFilesystemChunkKeys(t *testing.T) require.Equal(t, uint64(2), resp.RaftGroupId) } +func TestDistributionServerRouteReadsHonorStartupGate(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + engine.UpdateRoute([]byte("a"), nil, 1) + catalog := distribution.NewCatalogStore(store.NewMVCCStore()) + _, err := catalog.Save(context.Background(), 0, []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte("a"), End: nil, GroupID: 1, State: distribution.RouteStateActive}, + }) + require.NoError(t, err) + + blocked := true + s := NewDistributionServer(engine, catalog, WithDistributionReadGate(func() bool { return blocked })) + + _, err = s.GetRoute(context.Background(), &pb.GetRouteRequest{Key: []byte("a")}) + require.Error(t, err) + require.Equal(t, codes.Unavailable, status.Code(err)) + + _, err = s.ListRoutes(context.Background(), &pb.ListRoutesRequest{}) + require.Error(t, err) + require.Equal(t, codes.Unavailable, status.Code(err)) + + blocked = false + _, err = s.GetRoute(context.Background(), &pb.GetRouteRequest{Key: []byte("a")}) + require.NoError(t, err) + _, err = s.ListRoutes(context.Background(), &pb.ListRoutesRequest{}) + require.NoError(t, err) +} + func TestDistributionServerGetTimestamp_IsMonotonic(t *testing.T) { t.Parallel() @@ -106,10 +136,10 @@ func TestDistributionServerListRoutes_ReadsDurableCatalog(t *testing.T) { GroupID: 2, State: distribution.RouteStateWriteFenced, ParentRouteID: 1, - SplitAtHLC: 77, StagedVisibilityActive: true, MigrationJobID: 42, MinWriteTSExclusive: 99, + SplitAtHLC: 100, }, { RouteID: 1, @@ -136,10 +166,10 @@ func TestDistributionServerListRoutes_ReadsDurableCatalog(t *testing.T) { require.Equal(t, uint64(2), resp.Routes[1].RouteId) require.Nil(t, resp.Routes[1].End) require.Equal(t, pb.RouteState_ROUTE_STATE_WRITE_FENCED, resp.Routes[1].State) - require.Equal(t, uint64(77), resp.Routes[1].SplitAtHlc) require.True(t, resp.Routes[1].StagedVisibilityActive) require.Equal(t, uint64(42), resp.Routes[1].MigrationJobId) require.Equal(t, uint64(99), resp.Routes[1].MinWriteTsExclusive) + require.Equal(t, uint64(100), resp.Routes[1].SplitAtHlc) } func TestDistributionServerListRoutes_RequiresCatalog(t *testing.T) { @@ -205,6 +235,7 @@ func TestDistributionServerSplitRange_Success(t *testing.T) { require.Equal(t, []byte("m"), resp.Right.End) require.Equal(t, uint64(1), resp.Right.RaftGroupId) require.Equal(t, uint64(1), resp.Right.ParentRouteId) + require.Equal(t, uint64(99), resp.Right.MinWriteTsExclusive) require.NotZero(t, resp.Left.SplitAtHlc) require.Equal(t, resp.Left.SplitAtHlc, resp.Right.SplitAtHlc) require.Equal(t, uint64(99), resp.Right.MinWriteTsExclusive) @@ -471,6 +502,138 @@ func TestDistributionServerSplitRange_VersionConflict(t *testing.T) { require.ErrorContains(t, err, errDistributionCatalogConflict.Error()) } +func TestDistributionServerSplitRange_RejectsLiveSplitJobOverlap(t *testing.T) { + t.Parallel() + + ctx := context.Background() + baseStore := store.NewMVCCStore() + catalog := distribution.NewCatalogStore(baseStore) + saved, err := catalog.Save(ctx, 0, []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte("a"), End: []byte("m"), GroupID: 1, State: distribution.RouteStateActive}, + {RouteID: 2, Start: []byte("m"), End: nil, GroupID: 2, State: distribution.RouteStateActive}, + }) + require.NoError(t, err) + require.NoError(t, catalog.CreateSplitJob(ctx, distribution.SplitJob{ + JobID: 10, + SourceRouteID: 1, + SplitKey: []byte("g"), + TargetGroupID: 8, + Phase: distribution.SplitJobPhaseBackfill, + })) + + coordinator := newDistributionCoordinatorStub(baseStore, true) + s := NewDistributionServer(distribution.NewEngine(), catalog, WithDistributionCoordinator(coordinator)) + _, err = s.SplitRange(ctx, &pb.SplitRangeRequest{ + ExpectedCatalogVersion: saved.Version, + RouteId: 1, + SplitKey: []byte("c"), + }) + require.Error(t, err) + require.Equal(t, codes.Aborted, status.Code(err)) + require.ErrorContains(t, err, distribution.ErrSplitJobOverlap.Error()) + require.Zero(t, coordinator.dispatchCalls) +} + +func TestDistributionServerSplitRange_AllowsDisjointRouteWhileSplitJobLive(t *testing.T) { + t.Parallel() + + ctx := context.Background() + baseStore := store.NewMVCCStore() + catalog := distribution.NewCatalogStore(baseStore) + saved, err := catalog.Save(ctx, 0, []distribution.RouteDescriptor{ + {RouteID: 3, Start: []byte("a"), End: []byte("g"), GroupID: 1, State: distribution.RouteStateActive, ParentRouteID: 1}, + {RouteID: 4, Start: []byte("g"), End: []byte("m"), GroupID: 1, State: distribution.RouteStateWriteFenced, ParentRouteID: 1}, + {RouteID: 2, Start: []byte("m"), End: nil, GroupID: 2, State: distribution.RouteStateActive}, + }) + require.NoError(t, err) + require.NoError(t, catalog.CreateSplitJob(ctx, distribution.SplitJob{ + JobID: 10, + SourceRouteID: 1, + SplitKey: []byte("g"), + TargetGroupID: 8, + Phase: distribution.SplitJobPhaseFence, + })) + + coordinator := newDistributionCoordinatorStub(baseStore, true) + s := NewDistributionServer(distribution.NewEngine(), catalog, WithDistributionCoordinator(coordinator)) + resp, err := s.SplitRange(ctx, &pb.SplitRangeRequest{ + ExpectedCatalogVersion: saved.Version, + RouteId: 3, + SplitKey: []byte("c"), + }) + require.NoError(t, err) + require.Equal(t, uint64(2), resp.CatalogVersion) + require.Equal(t, 1, coordinator.dispatchCalls) + requireReadKeysContain(t, coordinator.lastReadKeys, distribution.CatalogNextSplitJobIDKey()) + requireReadKeysContain(t, coordinator.lastReadKeys, distribution.CatalogSplitJobKey(10)) +} + +func TestSplitJobReadFenceKeysExcludesTerminalHistory(t *testing.T) { + t.Parallel() + + readKeys := splitJobReadFenceKeys([]distribution.SplitJob{ + { + JobID: 10, + Phase: distribution.SplitJobPhaseBackfill, + }, + { + JobID: 11, + Phase: distribution.SplitJobPhaseDone, + TerminalAtMs: 1000, + }, + { + JobID: 12, + Phase: distribution.SplitJobPhaseAbandoned, + TerminalAtMs: 1001, + }, + }) + + require.Len(t, readKeys, 2) + requireReadKeysContain(t, readKeys, distribution.CatalogNextSplitJobIDKey()) + requireReadKeysContain(t, readKeys, distribution.CatalogSplitJobKey(10)) + requireReadKeysNotContain(t, readKeys, distribution.CatalogSplitJobHistoryKey(1000, 11)) + requireReadKeysNotContain(t, readKeys, distribution.CatalogSplitJobHistoryKey(1001, 12)) +} + +func TestDistributionServerSplitRange_ConflictsWhenSplitJobCreatedAfterOverlapScan(t *testing.T) { + t.Parallel() + + ctx := context.Background() + baseStore := store.NewMVCCStore() + catalog := distribution.NewCatalogStore(baseStore) + saved, err := catalog.Save(ctx, 0, []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte("a"), End: []byte("m"), GroupID: 1, State: distribution.RouteStateActive}, + {RouteID: 2, Start: []byte("m"), End: nil, GroupID: 2, State: distribution.RouteStateActive}, + }) + require.NoError(t, err) + + coordinator := newDistributionCoordinatorStub(baseStore, true) + coordinator.beforeApply = func(ctx context.Context, _ store.MVCCStore) error { + return catalog.CreateSplitJob(ctx, distribution.SplitJob{ + JobID: 10, + SourceRouteID: 1, + SplitKey: []byte("g"), + TargetGroupID: 8, + Phase: distribution.SplitJobPhaseBackfill, + }) + } + s := NewDistributionServer(distribution.NewEngine(), catalog, WithDistributionCoordinator(coordinator)) + + _, err = s.SplitRange(ctx, &pb.SplitRangeRequest{ + ExpectedCatalogVersion: saved.Version, + RouteId: 1, + SplitKey: []byte("c"), + }) + require.Error(t, err) + require.Equal(t, codes.Aborted, status.Code(err)) + require.ErrorContains(t, err, errDistributionCatalogConflict.Error()) + require.Equal(t, 1, coordinator.dispatchCalls) + + snapshot, err := catalog.Snapshot(ctx) + require.NoError(t, err) + require.Equal(t, saved.Version, snapshot.Version) +} + func TestDistributionServerSplitRange_UsesCoordinatorForCatalogWrites(t *testing.T) { t.Parallel() @@ -838,6 +1001,8 @@ type distributionCoordinatorStub struct { lastStartTS uint64 lastCommitTS uint64 lastRequestedCommitTS uint64 + lastReadKeys [][]byte + beforeApply func(context.Context, store.MVCCStore) error afterDispatch func(context.Context, store.MVCCStore, uint64) error asyncApplyDone chan error asyncApplyDelay time.Duration @@ -861,6 +1026,8 @@ func (s *distributionCoordinatorStub) Dispatch(ctx context.Context, reqs *kv.Ope startTS, commitTS := s.nextTimestamps(reqs.StartTS, reqs.CommitTS) s.lastStartTS = startTS s.lastCommitTS = commitTS + readKeys := cloneDistributionReadKeys(reqs.ReadKeys) + s.lastReadKeys = readKeys if err := kv.ValidateElemCommitTSPatches(reqs.Elems, commitTS); err != nil { return nil, err @@ -874,14 +1041,14 @@ func (s *distributionCoordinatorStub) Dispatch(ctx context.Context, reqs *kv.Ope delay := s.asyncApplyDelay go func() { time.Sleep(delay) - err := s.applyDispatch(ctx, mutations, startTS, commitTS) + err := s.applyDispatch(ctx, mutations, readKeys, startTS, commitTS) if done != nil { done <- err } }() return &kv.CoordinateResponse{CommitIndex: commitTS, CommitTS: commitTS}, nil } - if err := s.applyDispatch(ctx, mutations, startTS, commitTS); err != nil { + if err := s.applyDispatch(ctx, mutations, readKeys, startTS, commitTS); err != nil { return nil, err } return &kv.CoordinateResponse{CommitIndex: commitTS, CommitTS: commitTS}, nil @@ -918,10 +1085,16 @@ func (s *distributionCoordinatorStub) nextTimestamps(startTS uint64, requestedCo func (s *distributionCoordinatorStub) applyDispatch( ctx context.Context, mutations []*store.KVPairMutation, + readKeys [][]byte, startTS uint64, commitTS uint64, ) error { - if err := s.store.ApplyMutations(ctx, mutations, nil, startTS, commitTS); err != nil { + if s.beforeApply != nil { + if err := s.beforeApply(ctx, s.store); err != nil { + return err + } + } + if err := s.store.ApplyMutations(ctx, mutations, readKeys, startTS, commitTS); err != nil { return err } if s.afterDispatch != nil { @@ -932,6 +1105,36 @@ func (s *distributionCoordinatorStub) applyDispatch( return nil } +func cloneDistributionReadKeys(in [][]byte) [][]byte { + if len(in) == 0 { + return nil + } + out := make([][]byte, len(in)) + for i := range in { + out[i] = distribution.CloneBytes(in[i]) + } + return out +} + +func requireReadKeysContain(t *testing.T, readKeys [][]byte, want []byte) { + t.Helper() + for _, key := range readKeys { + if bytes.Equal(key, want) { + return + } + } + t.Fatalf("expected read keys to contain %q, got %q", want, readKeys) +} + +func requireReadKeysNotContain(t *testing.T, readKeys [][]byte, want []byte) { + t.Helper() + for _, key := range readKeys { + if bytes.Equal(key, want) { + t.Fatalf("expected read keys not to contain %q, got %q", want, readKeys) + } + } +} + func coordinatorStubMutations(elems []*kv.Elem[kv.OP], commitTS uint64) ([]*store.KVPairMutation, error) { mutations := make([]*store.KVPairMutation, 0, len(elems)) for _, elem := range elems { diff --git a/adapter/dynamodb_cleanup_retry_test.go b/adapter/dynamodb_cleanup_retry_test.go new file mode 100644 index 000000000..71b535e85 --- /dev/null +++ b/adapter/dynamodb_cleanup_retry_test.go @@ -0,0 +1,42 @@ +package adapter + +import ( + "context" + "testing" + + "github.com/bootjp/elastickv/kv" + "github.com/bootjp/elastickv/store" + "github.com/stretchr/testify/require" +) + +func TestDynamoDBCleanupDeletedTableGeneration_RetriesRouteFence(t *testing.T) { + t.Parallel() + + st := store.NewMVCCStore() + coord := &cleanupRouteFenceCoordinator{ + localAdapterCoordinator: newLocalAdapterCoordinator(st), + failuresRemaining: 1, + } + server := NewDynamoDBServer(nil, st, coord) + + err := server.cleanupDeletedTableGeneration(context.Background(), "table-a", 7) + require.NoError(t, err) + require.Equal(t, 3, coord.calls, "first prefix is retried once, second prefix succeeds once") +} + +type cleanupRouteFenceCoordinator struct { + *localAdapterCoordinator + failuresRemaining int + calls int +} + +func (c *cleanupRouteFenceCoordinator) Dispatch(ctx context.Context, req *kv.OperationGroup[kv.OP]) (*kv.CoordinateResponse, error) { + if req != nil && operationGroupHasDelPrefix(req.Elems) { + c.calls++ + if c.failuresRemaining > 0 { + c.failuresRemaining-- + return nil, kv.ErrRouteWriteFenced + } + } + return c.localAdapterCoordinator.Dispatch(ctx, req) +} diff --git a/adapter/dynamodb_item_write.go b/adapter/dynamodb_item_write.go index a5aecda19..b3ad5678f 100644 --- a/adapter/dynamodb_item_write.go +++ b/adapter/dynamodb_item_write.go @@ -285,7 +285,7 @@ func (d *DynamoDBServer) itemWriteFirstAttempt( plan.req.CommitTS = commitTS if dispErr := d.commitItemWrite(ctx, plan.req); dispErr != nil { // dispErr is already wrapped by commitItemWrite; return it raw. - if isRetryableTransactWriteError(dispErr) { + if shouldPreserveTransactWriteAttempt(dispErr) { return nil, &reusableItemWrite{ plan: plan, commitTS: commitTS, @@ -319,10 +319,13 @@ func (d *DynamoDBServer) itemWriteReuseAttempt( return d.resolveReuseWriteConflict(ctx, tableName, pending, commitTS, dispErr) } if isRetryableTransactWriteError(dispErr) { - // Still ambiguous (e.g. TxnLocked): this reuse may itself have landed, - // so the next retry must probe THIS commit_ts. dispErr is already - // wrapped by commitItemWrite; return it raw. - pending.commitTS = commitTS + // Still ambiguous (e.g. TxnLocked): this reuse may itself have + // landed, so the next retry must probe THIS commit_ts. Route-fence + // rejections are retryable but happen before this write set can apply, + // so keep the older witness. + if shouldPreserveTransactWriteAttempt(dispErr) { + pending.commitTS = commitTS + } return nil, pending, dispErr } return nil, nil, dispErr diff --git a/adapter/dynamodb_onephase_dedup_test.go b/adapter/dynamodb_onephase_dedup_test.go index ce7b9c7d4..b28afe3bd 100644 --- a/adapter/dynamodb_onephase_dedup_test.go +++ b/adapter/dynamodb_onephase_dedup_test.go @@ -113,6 +113,24 @@ func TestItemWriteDedup_LandedPriorAttempt_NoDuplicate(t *testing.T) { require.Equal(t, 1, coord.probeNoOps, "the reuse must dedup via the FSM exact-ts probe") } +func TestItemWriteDedup_RouteFenceRetryPreservesPriorProbe(t *testing.T) { + t.Parallel() + ctx := context.Background() + st := store.NewMVCCStore() + coord := newDedupTestCoordinator(st, 1, true) // dispatch 1 lands then errors + coord.routeFenceAtDispatch = 2 + schema, server := newDedupItemWriteServer(st, coord, true) + seedDedupItem(t, st, schema, "1", "2") + + plan, err := server.updateItemWithRetry(ctx, appendListInput()) + require.NoError(t, err) + require.NotNil(t, plan) + + require.Equal(t, []string{"1", "2", "3"}, readListValues(t, server, schema)) + require.Equal(t, 3, coord.dispatches, "attempt 1 landed, route-fenced reuse, then dedup probe retry") + require.Equal(t, 1, coord.probeNoOps, "route-fenced reuse must not replace the prior landed probe") +} + // TestItemWriteDedup_PriorAttemptDidNotLand_Applies: attempt 1 pre-rejects // (definitely did not commit); the reuse's probe misses, so it applies the // reused write set at a fresh commit_ts. One element, no duplicate. diff --git a/adapter/dynamodb_schema.go b/adapter/dynamodb_schema.go index 0ffa0f917..b3843b8d5 100644 --- a/adapter/dynamodb_schema.go +++ b/adapter/dynamodb_schema.go @@ -349,16 +349,36 @@ func (d *DynamoDBServer) cleanupDeletedTableGeneration(ctx context.Context, tabl // scans and writes tombstones locally, avoiding the enumerate-then-batch- // delete loop that previously required many Raft proposals. for _, prefix := range prefixes { + if err := d.dispatchDeletedTableCleanupPrefix(ctx, prefix); err != nil { + return err + } + } + return nil +} + +func (d *DynamoDBServer) dispatchDeletedTableCleanupPrefix(ctx context.Context, prefix []byte) error { + backoff := transactRetryInitialBackoff + deadline := time.Now().Add(transactRetryMaxDuration) + var lastErr error + for range transactRetryMaxAttempts { _, err := d.coordinator.Dispatch(ctx, &kv.OperationGroup[kv.OP]{ Elems: []*kv.Elem[kv.OP]{ {Op: kv.DelPrefix, Key: prefix}, }, }) - if err != nil { + if err == nil { + return nil + } + if !isRetryableTransactWriteError(err) { return errors.WithStack(err) } + lastErr = err + if waitErr := waitRetryWithDeadline(ctx, deadline, backoff); waitErr != nil { + return errors.Wrap(errors.Join(waitErr, lastErr), "dynamodb delete table cleanup retry canceled") + } + backoff = nextTransactRetryBackoff(backoff) } - return nil + return errors.WithStack(lastErr) } func (d *DynamoDBServer) dispatchDeleteBatch(ctx context.Context, keys [][]byte) error { diff --git a/adapter/dynamodb_transact.go b/adapter/dynamodb_transact.go index 3f58a688a..597c76472 100644 --- a/adapter/dynamodb_transact.go +++ b/adapter/dynamodb_transact.go @@ -1105,9 +1105,19 @@ func (d *DynamoDBServer) resolveTransactTableSchema(ctx context.Context, cache m } func isRetryableTransactWriteError(err error) bool { + return errors.Is(err, store.ErrWriteConflict) || errors.Is(err, kv.ErrTxnLocked) || isRouteWriteFencedError(err) +} + +func isIgnorableTransactRaceError(err error) bool { + // Route fences reject before applying the write. They must be retried or + // propagated, not swallowed as "another worker already completed it". return errors.Is(err, store.ErrWriteConflict) || errors.Is(err, kv.ErrTxnLocked) } +func shouldPreserveTransactWriteAttempt(err error) bool { + return isRetryableTransactWriteError(err) && !isRouteWriteFencedError(err) +} + func waitTransactRetryBackoff(ctx context.Context, delay time.Duration) error { timer := time.NewTimer(delay) defer timer.Stop() diff --git a/adapter/grpc.go b/adapter/grpc.go index 602250bf2..f2be4382a 100644 --- a/adapter/grpc.go +++ b/adapter/grpc.go @@ -25,6 +25,7 @@ type GRPCServer struct { grpcTranscoder *grpcTranscoder coordinator kv.Coordinator store store.MVCCStore + readBlocked func() bool closeStore bool closeOnce sync.Once @@ -78,6 +79,12 @@ func WithCloseStore() GRPCServerOption { } } +func WithGRPCReadGate(blocked func() bool) GRPCServerOption { + return func(s *GRPCServer) { + s.readBlocked = blocked + } +} + func NewGRPCServer(store store.MVCCStore, coordinate kv.Coordinator, opts ...GRPCServerOption) *GRPCServer { s := &GRPCServer{ log: slog.New(slog.NewJSONHandler(os.Stdout, &slog.HandlerOptions{ @@ -96,6 +103,14 @@ func NewGRPCServer(store store.MVCCStore, coordinate kv.Coordinator, opts ...GRP return s } +func (r *GRPCServer) requireReadReady() error { + if r != nil && r.readBlocked != nil && r.readBlocked() { + //nolint:wrapcheck // Preserve the gRPC status code for startup readers. + return status.Error(codes.Unavailable, "startup rotation has not completed") + } + return nil +} + func (r *GRPCServer) Close() error { if r == nil { return nil @@ -119,6 +134,9 @@ func (r *GRPCServer) clock() *kv.HLC { } func (r *GRPCServer) RawGet(ctx context.Context, req *pb.RawGetRequest) (*pb.RawGetResponse, error) { + if err := r.requireReadReady(); err != nil { + return nil, err + } readTS := req.GetTs() if readTS == 0 { readTS = globalSnapshotTS(ctx, r.clock(), r.store) @@ -157,6 +175,9 @@ func (r *GRPCServer) RawGet(ctx context.Context, req *pb.RawGetRequest) (*pb.Raw } func (r *GRPCServer) RawLatestCommitTS(ctx context.Context, req *pb.RawLatestCommitTSRequest) (*pb.RawLatestCommitTSResponse, error) { + if err := r.requireReadReady(); err != nil { + return nil, err + } key := req.GetKey() if len(key) == 0 { // No key: return the store's global last-committed watermark. @@ -189,6 +210,9 @@ func (r *GRPCServer) RawLatestCommitTS(ctx context.Context, req *pb.RawLatestCom } func (r *GRPCServer) RawScanAt(ctx context.Context, req *pb.RawScanAtRequest) (*pb.RawScanAtResponse, error) { + if err := r.requireReadReady(); err != nil { + return nil, err + } limit64 := req.GetLimit() limit, err := rawScanLimit(limit64) if err != nil { @@ -525,6 +549,9 @@ func (r *GRPCServer) Put(ctx context.Context, req *pb.PutRequest) (*pb.PutRespon } func (r *GRPCServer) Get(ctx context.Context, req *pb.GetRequest) (*pb.GetResponse, error) { + if err := r.requireReadReady(); err != nil { + return nil, err + } h := murmur3.New64() if _, err := h.Write(req.Key); err != nil { return nil, errors.WithStack(err) @@ -572,6 +599,9 @@ func (r *GRPCServer) Delete(ctx context.Context, req *pb.DeleteRequest) (*pb.Del } func (r *GRPCServer) Scan(ctx context.Context, req *pb.ScanRequest) (*pb.ScanResponse, error) { + if err := r.requireReadReady(); err != nil { + return nil, err + } limit, err := internal.Uint64ToInt(req.Limit) if err != nil { return &pb.ScanResponse{ diff --git a/adapter/grpc_test.go b/adapter/grpc_test.go index 4c080b4f0..00b8de807 100644 --- a/adapter/grpc_test.go +++ b/adapter/grpc_test.go @@ -355,6 +355,36 @@ func TestGRPCServer_RawReadFenceHelpersStampCurrentRouteVersion(t *testing.T) { require.Equal(t, uint64(55), st.scanReadRouteVersion) } +func TestGRPCServer_ReadsHonorStartupGate(t *testing.T) { + t.Parallel() + + blocked := true + s := NewGRPCServer(store.NewMVCCStore(), nil, WithGRPCReadGate(func() bool { return blocked })) + ctx := context.Background() + + _, err := s.RawGet(ctx, &pb.RawGetRequest{Key: []byte("k"), Ts: 10}) + require.Error(t, err) + require.Equal(t, codes.Unavailable, status.Code(err)) + _, err = s.RawLatestCommitTS(ctx, &pb.RawLatestCommitTSRequest{}) + require.Error(t, err) + require.Equal(t, codes.Unavailable, status.Code(err)) + _, err = s.RawScanAt(ctx, &pb.RawScanAtRequest{Limit: 10, Ts: 10}) + require.Error(t, err) + require.Equal(t, codes.Unavailable, status.Code(err)) + _, err = s.Get(ctx, &pb.GetRequest{Key: []byte("k")}) + require.Error(t, err) + require.Equal(t, codes.Unavailable, status.Code(err)) + _, err = s.Scan(ctx, &pb.ScanRequest{Limit: 10}) + require.Error(t, err) + require.Equal(t, codes.Unavailable, status.Code(err)) + + blocked = false + _, err = s.RawGet(ctx, &pb.RawGetRequest{Key: []byte("k"), Ts: 10}) + require.NoError(t, err) + _, err = s.Scan(ctx, &pb.ScanRequest{Limit: 10}) + require.NoError(t, err) +} + func TestGRPCServer_RawReadFenceHelpersKeepCallerRouteVersion(t *testing.T) { t.Parallel() diff --git a/adapter/redis_collection_ttl.go b/adapter/redis_collection_ttl.go index 7aadfa945..60ad8ca59 100644 --- a/adapter/redis_collection_ttl.go +++ b/adapter/redis_collection_ttl.go @@ -344,8 +344,7 @@ func (r *RedisServer) listMetaExpireElems(ctx context.Context, key []byte, readT if err != nil { return nil, false, err } - prefix := store.ListMetaDeltaScanPrefix(key) - deltas, err := r.scanDeltaKVs(ctx, prefix, readTS) + deltas, err := r.scanListMetaDeltaKVs(ctx, key, readTS) if err != nil { return r.listMetaExpireScanErr(key, meta, exists, expireAtMs, err) } @@ -383,6 +382,7 @@ func (r *RedisServer) listMetaExpireScanErr( return nil, false, err } r.triggerUrgentCompaction("list", key) + r.triggerUrgentCompaction("list-legacy", key) if exists { return listMetaTTLUpdateElem(key, meta, expireAtMs) } @@ -520,6 +520,28 @@ func (r *RedisServer) scanDeltaKVs(ctx context.Context, prefix []byte, readTS ui return deltas, nil } +func (r *RedisServer) scanListMetaDeltaKVs(ctx context.Context, userKey []byte, readTS uint64) ([]*store.KVPair, error) { + deltas := make([]*store.KVPair, 0, store.MaxDeltaScanLimit) + for _, prefix := range store.ListMetaDeltaScanPrefixes(userKey) { + accept := func(*store.KVPair) bool { return true } + if isLegacyListMetaDeltaPrefix(prefix) { + accept = func(pair *store.KVPair) bool { + return legacyListDeltaPairForUserKey(pair, userKey) + } + } + remaining := store.MaxDeltaScanLimit - len(deltas) + found, truncated, err := scanAcceptedDeltaKVsAt(ctx, r.store, prefix, remaining, readTS, accept) + if err != nil { + return nil, err + } + if truncated { + return nil, ErrDeltaScanTruncated + } + deltas = append(deltas, found...) + } + return deltas, nil +} + type redisDeltaKVScanner interface { scanDeltaKVs(context.Context, []byte, uint64) ([]*store.KVPair, error) } diff --git a/adapter/redis_collection_ttl_test.go b/adapter/redis_collection_ttl_test.go index d7f3ce1b8..56ac8a954 100644 --- a/adapter/redis_collection_ttl_test.go +++ b/adapter/redis_collection_ttl_test.go @@ -142,6 +142,41 @@ func TestRedisCollectionExpireWritesInlineMetaTTL(t *testing.T) { } } +func TestListMetaExpireElemsIncludesLegacyDeltas(t *testing.T) { + t.Parallel() + ctx := context.Background() + st := store.NewMVCCStore() + server := &RedisServer{store: st} + key := []byte("ttl:list:legacy-delta") + deltaKey := legacyListMetaDeltaKey(key, 2) + require.NoError(t, st.PutAt(ctx, deltaKey, store.MarshalListMetaDelta(store.ListMetaDelta{LenDelta: 1}), 2, 0)) + + elems, exists, err := server.listMetaExpireElems(ctx, key, 2, 1234) + + require.NoError(t, err) + require.True(t, exists) + meta, err := store.UnmarshalListMeta(elemValueForKey(t, elems, store.ListMetaKey(key))) + require.NoError(t, err) + require.Equal(t, int64(1), meta.Len) + require.Equal(t, uint64(1234), meta.ExpireAt) + require.True(t, elemDelKeysContain(elems, deltaKey)) +} + +func TestExpiredTTLIndexPrecedesLegacyListDeltaOnlyCollection(t *testing.T) { + t.Parallel() + ctx := context.Background() + st := store.NewMVCCStore() + server := &RedisServer{store: st} + key := []byte("ttl:list:legacy-delta-recreate") + require.NoError(t, st.PutAt(ctx, redisTTLKey(key), encodeRedisTTL(time.Now().Add(-time.Hour)), 1, 0)) + require.NoError(t, st.PutAt(ctx, legacyListMetaDeltaKey(key, 3), store.MarshalListMetaDelta(store.ListMetaDelta{LenDelta: 1}), 3, 0)) + + stale, err := expiredTTLIndexPrecedesDeltaOnlyCollection(ctx, st, server, key, redisTypeList, 3) + + require.NoError(t, err) + require.True(t, stale) +} + func TestRedisCollectionExpireHandlesLegacyBlobs(t *testing.T) { t.Parallel() @@ -307,6 +342,7 @@ func TestRedisCollectionExpireAllowsDeltaHeavyCollections(t *testing.T) { deltaKey func([]byte, uint64) []byte deltaValue []byte metaExpireAt func([]byte) (uint64, error) + urgentTypes []string }{ { name: "list", @@ -326,6 +362,7 @@ func TestRedisCollectionExpireAllowsDeltaHeavyCollections(t *testing.T) { meta, err := store.UnmarshalListMeta(raw) return meta.ExpireAt, err }, + urgentTypes: []string{"list", "list-legacy"}, }, { name: "hash", @@ -343,6 +380,7 @@ func TestRedisCollectionExpireAllowsDeltaHeavyCollections(t *testing.T) { meta, err := store.UnmarshalHashMeta(raw) return meta.ExpireAt, err }, + urgentTypes: []string{"hash"}, }, { name: "set", @@ -360,6 +398,7 @@ func TestRedisCollectionExpireAllowsDeltaHeavyCollections(t *testing.T) { meta, err := store.UnmarshalSetMeta(raw) return meta.ExpireAt, err }, + urgentTypes: []string{"set"}, }, { name: "zset", @@ -377,6 +416,7 @@ func TestRedisCollectionExpireAllowsDeltaHeavyCollections(t *testing.T) { meta, err := store.UnmarshalZSetMeta(raw) return meta.ExpireAt, err }, + urgentTypes: []string{"zset"}, }, } @@ -412,13 +452,7 @@ func TestRedisCollectionExpireAllowsDeltaHeavyCollections(t *testing.T) { ttl, err := decodeRedisTTL(rawTTL) require.NoError(t, err) require.Equal(t, redisExpireAtMillis(expireAt), redisExpireAtMillis(ttl)) - select { - case req := <-compactor.urgentCh: - require.Equal(t, tc.name, req.typeName) - require.Equal(t, key, req.userKey) - default: - t.Fatalf("expected urgent compaction request for %s", tc.name) - } + requireUrgentCompactionRequests(t, compactor, key, tc.urgentTypes...) }) } } @@ -429,34 +463,39 @@ func TestRedisCollectionExpireTriggersCompactionForDeltaOnlyTruncatedCollections ctx := context.Background() expireAt := time.Now().Add(time.Hour) cases := []struct { - name string - typ redisValueType - deltaKey func([]byte, uint64) []byte - deltaValue []byte + name string + typ redisValueType + deltaKey func([]byte, uint64) []byte + deltaValue []byte + urgentTypes []string }{ { - name: "list", - typ: redisTypeList, - deltaKey: func(key []byte, ts uint64) []byte { return store.ListMetaDeltaKey(key, ts, 0) }, - deltaValue: store.MarshalListMetaDelta(store.ListMetaDelta{LenDelta: 1}), + name: "list", + typ: redisTypeList, + deltaKey: func(key []byte, ts uint64) []byte { return store.ListMetaDeltaKey(key, ts, 0) }, + deltaValue: store.MarshalListMetaDelta(store.ListMetaDelta{LenDelta: 1}), + urgentTypes: []string{"list", "list-legacy"}, }, { - name: "hash", - typ: redisTypeHash, - deltaKey: func(key []byte, ts uint64) []byte { return store.HashMetaDeltaKey(key, ts, 0) }, - deltaValue: store.MarshalHashMetaDelta(store.HashMetaDelta{LenDelta: 1}), + name: "hash", + typ: redisTypeHash, + deltaKey: func(key []byte, ts uint64) []byte { return store.HashMetaDeltaKey(key, ts, 0) }, + deltaValue: store.MarshalHashMetaDelta(store.HashMetaDelta{LenDelta: 1}), + urgentTypes: []string{"hash"}, }, { - name: "set", - typ: redisTypeSet, - deltaKey: func(key []byte, ts uint64) []byte { return store.SetMetaDeltaKey(key, ts, 0) }, - deltaValue: store.MarshalSetMetaDelta(store.SetMetaDelta{LenDelta: 1}), + name: "set", + typ: redisTypeSet, + deltaKey: func(key []byte, ts uint64) []byte { return store.SetMetaDeltaKey(key, ts, 0) }, + deltaValue: store.MarshalSetMetaDelta(store.SetMetaDelta{LenDelta: 1}), + urgentTypes: []string{"set"}, }, { - name: "zset", - typ: redisTypeZSet, - deltaKey: func(key []byte, ts uint64) []byte { return store.ZSetMetaDeltaKey(key, ts, 0) }, - deltaValue: store.MarshalZSetMetaDelta(store.ZSetMetaDelta{LenDelta: 1}), + name: "zset", + typ: redisTypeZSet, + deltaKey: func(key []byte, ts uint64) []byte { return store.ZSetMetaDeltaKey(key, ts, 0) }, + deltaValue: store.MarshalZSetMetaDelta(store.ZSetMetaDelta{LenDelta: 1}), + urgentTypes: []string{"zset"}, }, } @@ -479,17 +518,24 @@ func TestRedisCollectionExpireTriggersCompactionForDeltaOnlyTruncatedCollections _, err = st.GetAt(ctx, redisTTLKey(key), st.LastCommitTS()) require.ErrorIs(t, err, store.ErrKeyNotFound) - select { - case req := <-compactor.urgentCh: - require.Equal(t, tc.name, req.typeName) - require.Equal(t, key, req.userKey) - default: - t.Fatalf("expected urgent compaction request for %s", tc.name) - } + requireUrgentCompactionRequests(t, compactor, key, tc.urgentTypes...) }) } } +func requireUrgentCompactionRequests(t *testing.T, compactor *DeltaCompactor, key []byte, typeNames ...string) { + t.Helper() + for _, typeName := range typeNames { + select { + case req := <-compactor.urgentCh: + require.Equal(t, typeName, req.typeName) + require.Equal(t, key, req.userKey) + default: + t.Fatalf("expected urgent compaction request for %s", typeName) + } + } +} + func TestRedisCollectionExpireWritesInlineTTLForZeroLengthSimpleMeta(t *testing.T) { t.Parallel() diff --git a/adapter/redis_compat_helpers.go b/adapter/redis_compat_helpers.go index df4c0caa8..b0a6289c7 100644 --- a/adapter/redis_compat_helpers.go +++ b/adapter/redis_compat_helpers.go @@ -267,16 +267,43 @@ func (r *RedisServer) probeListType(ctx context.Context, key []byte, readTS uint if metaExists { return redisTypeList, true, nil } - deltaPrefix := store.ListMetaDeltaScanPrefix(key) + for _, deltaPrefix := range store.ListMetaDeltaScanPrefixes(key) { + found, err := r.listMetaDeltaExistsAt(ctx, key, deltaPrefix, readTS) + if err != nil { + return redisTypeNone, false, err + } + if found { + return redisTypeList, true, nil + } + } + return redisTypeNone, false, nil +} + +func (r *RedisServer) listMetaDeltaExistsAt(ctx context.Context, key []byte, deltaPrefix []byte, readTS uint64) (bool, error) { deltaEnd := store.PrefixScanEnd(deltaPrefix) - deltaKVs, err := r.store.ScanAt(ctx, deltaPrefix, deltaEnd, 1, readTS) - if err != nil { - return redisTypeNone, false, errors.WithStack(err) + if !isLegacyListMetaDeltaPrefix(deltaPrefix) { + deltaKVs, err := r.store.ScanAt(ctx, deltaPrefix, deltaEnd, 1, readTS) + if err != nil { + return false, errors.WithStack(err) + } + return len(deltaKVs) > 0, nil } - if len(deltaKVs) > 0 { - return redisTypeList, true, nil + cursor := deltaPrefix + for { + deltaKVs, err := r.store.ScanAt(ctx, cursor, deltaEnd, store.MaxDeltaScanLimit, readTS) + if err != nil { + return false, errors.WithStack(err) + } + for _, pair := range deltaKVs { + if legacyListDeltaPairForUserKey(pair, key) { + return true, nil + } + } + if len(deltaKVs) < store.MaxDeltaScanLimit { + return false, nil + } + cursor = append(bytes.Clone(deltaKVs[len(deltaKVs)-1].Key), 0) } - return redisTypeNone, false, nil } // probeLegacyCollectionTypes checks for single-blob hash/set/zset/stream @@ -931,12 +958,14 @@ func (r *RedisServer) deleteListElems(ctx context.Context, key []byte, readTS ui } // Always delete the base meta key (no-op tombstone if it doesn't exist). elems = append(elems, &kv.Elem[kv.OP]{Op: kv.Del, Key: listMetaKey(key)}) - // Delete all delta keys (paginated). - deltaElems, err := r.scanAllDeltaElems(ctx, store.ListMetaDeltaScanPrefix(key), readTS) - if err != nil { - return nil, err + // Delete all delta keys (paginated), including the pre-upgrade prefix. + for _, deltaPrefix := range store.ListMetaDeltaScanPrefixes(key) { + deltaElems, err := r.scanListDeltaDelElems(ctx, key, deltaPrefix, readTS) + if err != nil { + return nil, err + } + elems = append(elems, deltaElems...) } - elems = append(elems, deltaElems...) // Delete all claim keys (paginated). claimElems, err := r.scanAllDeltaElems(ctx, store.ListClaimScanPrefix(key), readTS) if err != nil { @@ -968,7 +997,7 @@ func (r *RedisServer) scanListItemDelElems(ctx context.Context, key []byte, read if len(elems)+1 > maxWideColumnItems { return nil, errors.Wrapf(ErrCollectionTooLarge, "list %q exceeds %d items", key, maxWideColumnItems) } - elems = append(elems, &kv.Elem[kv.OP]{Op: kv.Del, Key: pair.Key}) + elems = append(elems, &kv.Elem[kv.OP]{Op: kv.Del, Key: bytes.Clone(pair.Key), GroupID: pair.RouteGroupID}) } if len(itemKVs) < store.MaxDeltaScanLimit { break @@ -983,6 +1012,24 @@ func (r *RedisServer) scanListItemDelElems(ctx context.Context, key []byte, read // MaxDeltaScanLimit entries. Total results are capped at maxWideColumnItems // to prevent unbounded memory growth if the compactor falls behind. func (r *RedisServer) scanAllDeltaElems(ctx context.Context, deltaPrefix []byte, readTS uint64) ([]*kv.Elem[kv.OP], error) { + return r.scanAllDeltaElemsFiltered(ctx, deltaPrefix, readTS, nil) +} + +func (r *RedisServer) scanListDeltaDelElems(ctx context.Context, key []byte, deltaPrefix []byte, readTS uint64) ([]*kv.Elem[kv.OP], error) { + if !isLegacyListMetaDeltaPrefix(deltaPrefix) { + return r.scanAllDeltaElems(ctx, deltaPrefix, readTS) + } + return r.scanAllDeltaElemsFiltered(ctx, deltaPrefix, readTS, func(pair *store.KVPair) bool { + return legacyListDeltaPairForUserKey(pair, key) + }) +} + +func (r *RedisServer) scanAllDeltaElemsFiltered( + ctx context.Context, + deltaPrefix []byte, + readTS uint64, + include func(*store.KVPair) bool, +) ([]*kv.Elem[kv.OP], error) { const cursorAdv = byte(0x00) // appended to advance past the last scanned key var elems []*kv.Elem[kv.OP] deltaEnd := store.PrefixScanEnd(deltaPrefix) @@ -992,13 +1039,15 @@ func (r *RedisServer) scanAllDeltaElems(ctx context.Context, deltaPrefix []byte, if scanErr != nil { return nil, errors.WithStack(scanErr) } - // Check before appending so len(elems) never exceeds maxWideColumnItems - // by more than one scan page. - if len(elems)+len(deltaKVs) > maxWideColumnItems { - return nil, errors.Wrapf(ErrCollectionTooLarge, "delta key count exceeds %d", maxWideColumnItems) - } for _, pair := range deltaKVs { - elems = append(elems, &kv.Elem[kv.OP]{Op: kv.Del, Key: pair.Key}) + if include != nil && !include(pair) { + continue + } + // Check before appending so len(elems) never exceeds maxWideColumnItems. + if len(elems)+1 > maxWideColumnItems { + return nil, errors.Wrapf(ErrCollectionTooLarge, "delta key count exceeds %d", maxWideColumnItems) + } + elems = append(elems, &kv.Elem[kv.OP]{Op: kv.Del, Key: bytes.Clone(pair.Key), GroupID: pair.RouteGroupID}) } if len(deltaKVs) < store.MaxDeltaScanLimit { break @@ -1008,6 +1057,50 @@ func (r *RedisServer) scanAllDeltaElems(ctx context.Context, deltaPrefix []byte, return elems, nil } +func scanAcceptedDeltaKVsAt( + ctx context.Context, + st store.MVCCStore, + prefix []byte, + limit int, + readTS uint64, + accept func(*store.KVPair) bool, +) ([]*store.KVPair, bool, error) { + const cursorAdv = byte(0x00) + end := store.PrefixScanEnd(prefix) + cursor := prefix + out := make([]*store.KVPair, 0, limit) + for { + rawKVs, err := st.ScanAt(ctx, cursor, end, limit+1, readTS) + if err != nil { + return nil, false, errors.WithStack(err) + } + for _, pair := range rawKVs { + if accept != nil && !accept(pair) { + continue + } + out = append(out, pair) + if len(out) > limit { + return out, true, nil + } + } + if len(rawKVs) < limit+1 { + return out, false, nil + } + cursor = append(bytes.Clone(rawKVs[len(rawKVs)-1].Key), cursorAdv) + } +} + +func isLegacyListMetaDeltaPrefix(prefix []byte) bool { + return bytes.HasPrefix(prefix, []byte(store.LegacyListMetaDeltaPrefix)) +} + +func legacyListDeltaPairForUserKey(pair *store.KVPair, userKey []byte) bool { + if pair == nil || !store.IsListMetaDeltaValue(pair.Value) { + return false + } + return bytes.Equal(store.ExtractLegacyListUserKeyFromDelta(pair.Key), userKey) +} + // deleteWideColumnElems returns delete operations for all wide-column field/member keys, // the base meta key, and all delta keys for a collection identified by the given scan prefix, // meta key, and delta prefix. @@ -1265,29 +1358,82 @@ func minRedisInt(a, b int) int { // aggregateLenDeltas scans delta keys under prefix and sums the LenDelta values // via unmarshalDelta. Returns (sum, hasDeltas, error). -// ErrDeltaScanTruncated is returned when the scan hits MaxDeltaScanLimit. -func (r *RedisServer) aggregateLenDeltas(ctx context.Context, prefix []byte, readTS uint64, unmarshalDelta func([]byte) (int64, error)) (int64, bool, error) { +// ErrDeltaScanTruncated is returned when accepted deltas exceed MaxDeltaScanLimit. +func (r *RedisServer) aggregateLenDeltas( + ctx context.Context, + prefix []byte, + readTS uint64, + scanCap *deltaScanCap, + unmarshalDelta func(key []byte, value []byte) (int64, bool, error), +) (int64, bool, error) { + const cursorAdv = byte(0x00) end := store.PrefixScanEnd(prefix) - // Scan one extra key beyond the limit so we can distinguish "exactly - // MaxDeltaScanLimit results" (no truncation) from "more than MaxDeltaScanLimit - // results" (truncated). Without the +1, a collection with exactly - // MaxDeltaScanLimit deltas would incorrectly trigger ErrDeltaScanTruncated. - deltas, err := r.store.ScanAt(ctx, prefix, end, store.MaxDeltaScanLimit+1, readTS) - if err != nil { - return 0, false, errors.WithStack(err) - } - if len(deltas) > store.MaxDeltaScanLimit { - return 0, false, ErrDeltaScanTruncated - } var sum int64 - for _, d := range deltas { - delta, err := unmarshalDelta(d.Value) + var any bool + cursor := prefix + for { + deltas, err := r.store.ScanAt(ctx, cursor, end, store.MaxDeltaScanLimit+1, readTS) if err != nil { return 0, false, errors.WithStack(err) } - sum += delta + for _, d := range deltas { + delta, include, err := unmarshalDelta(d.Key, d.Value) + if err != nil { + return 0, false, errors.WithStack(err) + } + if !include { + continue + } + if scanCap == nil { + scanCap = &deltaScanCap{} + } + if !scanCap.accept() { + return 0, false, ErrDeltaScanTruncated + } + any = true + sum += delta + } + if len(deltas) < store.MaxDeltaScanLimit+1 { + break + } + cursor = append(bytes.Clone(deltas[len(deltas)-1].Key), cursorAdv) } - return sum, len(deltas) > 0, nil + return sum, any, nil +} + +type deltaScanCap struct { + accepted int +} + +func (c *deltaScanCap) accept() bool { + c.accepted++ + return c.accepted <= store.MaxDeltaScanLimit +} + +func (r *RedisServer) aggregateListMetaDeltas(ctx context.Context, key []byte, readTS uint64, applyDelta func(store.ListMetaDelta)) (int64, bool, error) { + var total int64 + var any bool + scanCap := &deltaScanCap{} + for _, prefix := range store.ListMetaDeltaScanPrefixes(key) { + legacy := isLegacyListMetaDeltaPrefix(prefix) + lenSum, hasDeltas, err := r.aggregateLenDeltas(ctx, prefix, readTS, scanCap, func(deltaKey []byte, b []byte) (int64, bool, error) { + if legacy && (!store.IsListMetaDeltaValue(b) || !bytes.Equal(store.ExtractLegacyListUserKeyFromDelta(deltaKey), key)) { + return 0, false, nil + } + d, unmarshalErr := store.UnmarshalListMetaDelta(b) + if unmarshalErr != nil { + return 0, false, errors.WithStack(unmarshalErr) + } + applyDelta(d) + return d.LenDelta, true, nil + }) + if err != nil { + return 0, false, err + } + total += lenSum + any = any || hasDeltas + } + return total, any, nil } // resolveListMeta aggregates the base list metadata with all uncompacted Delta keys @@ -1301,15 +1447,13 @@ func (r *RedisServer) resolveListMeta(ctx context.Context, key []byte, readTS ui // 2. Scan and aggregate delta keys. // The closure also captures baseMeta to accumulate the list-specific HeadDelta. - prefix := store.ListMetaDeltaScanPrefix(key) - lenSum, hasDeltas, err := r.aggregateLenDeltas(ctx, prefix, readTS, func(b []byte) (int64, error) { - d, unmarshalErr := store.UnmarshalListMetaDelta(b) + lenSum, hasDeltas, err := r.aggregateListMetaDeltas(ctx, key, readTS, func(d store.ListMetaDelta) { baseMeta.Head += d.HeadDelta - return d.LenDelta, errors.WithStack(unmarshalErr) }) if err != nil { if errors.Is(err, ErrDeltaScanTruncated) { r.triggerUrgentCompaction("list", key) + r.triggerUrgentCompaction("list-legacy", key) } return store.ListMeta{}, false, err } @@ -1351,7 +1495,10 @@ func (r *RedisServer) resolveCollectionLen( } } - deltaSum, hasDeltas, err := r.aggregateLenDeltas(ctx, deltaPrefix, readTS, unmarshalDelta) + deltaSum, hasDeltas, err := r.aggregateLenDeltas(ctx, deltaPrefix, readTS, nil, func(_ []byte, value []byte) (int64, bool, error) { + delta, err := unmarshalDelta(value) + return delta, true, err + }) if err != nil { return 0, false, err } diff --git a/adapter/redis_delta_compactor.go b/adapter/redis_delta_compactor.go index 424b71a7c..0e01d55b7 100644 --- a/adapter/redis_delta_compactor.go +++ b/adapter/redis_delta_compactor.go @@ -225,11 +225,10 @@ func (c *DeltaCompactor) compactUrgentKey(ctx context.Context, req urgentCompact defer cancel() prefix := h.deltaKeyPrefixFn(req.userKey) - end := store.PrefixScanEnd(prefix) totalCompacted, done := 0, false for !done { var n int - n, done = c.compactUrgentKeyBatch(tickCtx, req, h, prefix, end) + n, done = c.compactUrgentKeyBatch(tickCtx, req, h, prefix) totalCompacted += n if n == 0 { break @@ -246,13 +245,12 @@ func (c *DeltaCompactor) compactUrgentKey(ctx context.Context, req urgentCompact // Returns (n, done): n is the number of delta keys processed in this batch; // done is true when the remaining delta count is at or below MaxDeltaScanLimit // (the key is now readable). -func (c *DeltaCompactor) compactUrgentKeyBatch(ctx context.Context, req urgentCompactionRequest, h *collectionDeltaHandler, prefix, end []byte) (int, bool) { +func (c *DeltaCompactor) compactUrgentKeyBatch(ctx context.Context, req urgentCompactionRequest, h *collectionDeltaHandler, prefix []byte) (int, bool) { // Use a fresh readTS each iteration so we observe the committed state from // the previous compaction pass and do not re-scan already-deleted delta keys. readTS := snapshotTS(c.coord.Clock(), c.st) - // Scan one extra beyond MaxDeltaScanLimit to detect whether more remain. - kvs, err := c.st.ScanAt(ctx, prefix, end, store.MaxDeltaScanLimit+1, readTS) + kvs, truncated, err := scanAcceptedDeltaKVsAt(ctx, c.st, prefix, store.MaxDeltaScanLimit, readTS, h.acceptDeltaKV) if err != nil { c.logger.WarnContext(ctx, "delta compactor urgent: scan failed", "type", req.typeName, "key", string(req.userKey), "error", err) @@ -273,8 +271,7 @@ func (c *DeltaCompactor) compactUrgentKeyBatch(ctx context.Context, req urgentCo "type", req.typeName, "key", string(req.userKey), "error", err) return 0, true } - // Fewer than MaxDeltaScanLimit+1 results means no more deltas remain; key is readable. - return len(kvs), len(kvs) <= store.MaxDeltaScanLimit + return len(kvs), !truncated } // SyncOnce runs one compaction pass. The IsLeader() guard avoids the @@ -416,6 +413,7 @@ type collectionDeltaHandler struct { typeName string prefix []byte extractUserKey func(key []byte) []byte + acceptDeltaKV func(*store.KVPair) bool // deltaKeyPrefixFn returns the prefix that covers all delta keys for a single // user key. Used by compactUrgentKey to perform a targeted single-key scan. deltaKeyPrefixFn func(userKey []byte) []byte @@ -426,6 +424,7 @@ type collectionDeltaHandler struct { func (c *DeltaCompactor) allHandlers() []collectionDeltaHandler { return []collectionDeltaHandler{ c.listHandler(), + c.legacyListHandler(), c.hashHandler(), c.setHandler(), c.zsetHandler(), @@ -485,8 +484,11 @@ func (c *DeltaCompactor) compactHandler(ctx context.Context, h collectionDeltaHa return err } + rawKVs := kvs + kvs = filterDeltaKVs(rawKVs, h.acceptDeltaKV) byKey, ukOrder := c.groupByUserKey(kvs, h.extractUserKey) lastScannedKey := c.splitGuardCursor(byKey, ukOrder, kvs, truncated) + lastScannedKey = deltaCompactorCursorAfterFiltering(rawKVs, kvs, lastScannedKey, truncated, h.extractUserKey) allElems := c.buildBatchElems(ctx, h, byKey, ukOrder, readTS) @@ -500,8 +502,114 @@ func (c *DeltaCompactor) compactHandler(ctx context.Context, h collectionDeltaHa return nil } +func deltaCompactorCursorAfterFiltering( + rawKVs, acceptedKVs []*store.KVPair, + lastScannedKey []byte, + truncated bool, + extractUserKey func([]byte) []byte, +) []byte { + if !truncated || len(rawKVs) == 0 { + return lastScannedKey + } + if len(acceptedKVs) == 0 { + return rawKVs[len(rawKVs)-1].Key + } + lastAcceptedKey := acceptedKVs[len(acceptedKVs)-1].Key + switch { + case len(lastScannedKey) == 0: + if prev := rawKeyBeforeAcceptedTail(rawKVs, acceptedKVs, lastAcceptedKey, extractUserKey); len(prev) > 0 { + return prev + } + if rawPageHasRejectedTailAfter(rawKVs, lastAcceptedKey) { + return rawKVs[len(rawKVs)-1].Key + } + case !bytes.Equal(lastScannedKey, lastAcceptedKey) && rawPageHasRejectedTailAfter(rawKVs, lastAcceptedKey): + return rawKVs[len(rawKVs)-1].Key + } + return lastScannedKey +} + +func rawPageHasRejectedTailAfter(rawKVs []*store.KVPair, acceptedKey []byte) bool { + if len(rawKVs) == 0 || len(acceptedKey) == 0 { + return false + } + for i := len(rawKVs) - 1; i >= 0; i-- { + if rawKVs[i] != nil && bytes.Equal(rawKVs[i].Key, acceptedKey) { + return i < len(rawKVs)-1 + } + } + return false +} + +func rawKeyBeforeAcceptedTail( + rawKVs, acceptedKVs []*store.KVPair, + acceptedKey []byte, + extractUserKey func([]byte) []byte, +) []byte { + if len(rawKVs) < 2 || len(acceptedKey) == 0 { + return nil + } + last := len(rawKVs) - 1 + if rawKVs[last] == nil || !bytes.Equal(rawKVs[last].Key, acceptedKey) { + return nil + } + acceptedUserKey := []byte(nil) + if extractUserKey != nil { + acceptedUserKey = extractUserKey(acceptedKey) + } + acceptedKeys := deltaCompactorAcceptedKeySet(acceptedKVs) + for i := last - 1; i >= 0; i-- { + if rawKVs[i] == nil { + continue + } + if deltaCompactorSameAcceptedTailUser(rawKVs[i].Key, acceptedKeys, acceptedUserKey, extractUserKey) { + continue + } + return rawKVs[i].Key + } + return nil +} + +func deltaCompactorAcceptedKeySet(acceptedKVs []*store.KVPair) map[string]struct{} { + acceptedKeys := make(map[string]struct{}, len(acceptedKVs)) + for _, kvp := range acceptedKVs { + if kvp != nil { + acceptedKeys[string(kvp.Key)] = struct{}{} + } + } + return acceptedKeys +} + +func deltaCompactorSameAcceptedTailUser( + rawKey []byte, + acceptedKeys map[string]struct{}, + acceptedUserKey []byte, + extractUserKey func([]byte) []byte, +) bool { + if len(acceptedUserKey) == 0 || extractUserKey == nil { + return false + } + if _, accepted := acceptedKeys[string(rawKey)]; !accepted { + return false + } + return bytes.Equal(extractUserKey(rawKey), acceptedUserKey) +} + // groupByUserKey groups KVPairs by their user key, returning both the map and // the unique user keys in lexicographic (scan) order. +func filterDeltaKVs(kvs []*store.KVPair, accept func(*store.KVPair) bool) []*store.KVPair { + if accept == nil || len(kvs) == 0 { + return kvs + } + out := kvs[:0] + for _, pair := range kvs { + if accept(pair) { + out = append(out, pair) + } + } + return out +} + func (c *DeltaCompactor) groupByUserKey(kvs []*store.KVPair, extractUserKey func([]byte) []byte) (map[string][]*store.KVPair, []string) { byKey := make(map[string][]*store.KVPair) var ukOrder []string @@ -716,12 +824,36 @@ func (c *DeltaCompactor) listHandler() collectionDeltaHandler { typeName: "list", prefix: []byte(store.ListMetaDeltaPrefix), extractUserKey: store.ExtractListUserKeyFromDelta, + acceptDeltaKV: isListMetaDeltaKV, deltaKeyPrefixFn: store.ListMetaDeltaScanPrefix, buildElems: c.buildListCompactElems, } } +func (c *DeltaCompactor) legacyListHandler() collectionDeltaHandler { + return collectionDeltaHandler{ + typeName: "list-legacy", + prefix: []byte(store.LegacyListMetaDeltaPrefix), + extractUserKey: store.ExtractLegacyListUserKeyFromDelta, + acceptDeltaKV: isListMetaDeltaKV, + deltaKeyPrefixFn: store.LegacyListMetaDeltaScanPrefix, + buildElems: c.buildLegacyListCompactElems, + } +} + +func isListMetaDeltaKV(pair *store.KVPair) bool { + return pair != nil && store.IsListMetaDeltaValue(pair.Value) +} + func (c *DeltaCompactor) buildListCompactElems(ctx context.Context, userKey []byte, deltaKVs []*store.KVPair, readTS uint64) ([]*kv.Elem[kv.OP], error) { + return c.buildListCompactElemsWithDeltaPrefix(ctx, userKey, deltaKVs, readTS, store.ListMetaDeltaScanPrefix(userKey)) +} + +func (c *DeltaCompactor) buildLegacyListCompactElems(ctx context.Context, userKey []byte, deltaKVs []*store.KVPair, readTS uint64) ([]*kv.Elem[kv.OP], error) { + return c.buildListCompactElemsWithDeltaPrefix(ctx, userKey, deltaKVs, readTS, store.LegacyListMetaDeltaScanPrefix(userKey)) +} + +func (c *DeltaCompactor) buildListCompactElemsWithDeltaPrefix(ctx context.Context, userKey []byte, deltaKVs []*store.KVPair, readTS uint64, deltaScanPrefix []byte) ([]*kv.Elem[kv.OP], error) { // Read base metadata (may not exist if all state is in deltas). baseMeta, rawBaseMeta, err := c.loadListBaseMeta(ctx, userKey, readTS) if err != nil { @@ -729,7 +861,7 @@ func (c *DeltaCompactor) buildListCompactElems(ctx context.Context, userKey []by } expireAt, err := compactedMetaExpireAt( ctx, c.st, userKey, readTS, rawBaseMeta, redisWideMetaInlineSizeBytes, baseMeta.ExpireAt, - deltaKVs, store.ListMetaDeltaScanPrefix(userKey), + deltaKVs, deltaScanPrefix, ) if err != nil { return nil, err @@ -773,7 +905,7 @@ func (c *DeltaCompactor) buildListCompactElems(ctx context.Context, userKey []by elems := make([]*kv.Elem[kv.OP], 0, 1+len(deltaKVs)+len(claimElems)) elems = append(elems, metaElem) for _, d := range deltaKVs { - elems = append(elems, &kv.Elem[kv.OP]{Op: kv.Del, Key: bytes.Clone(d.Key)}) + elems = append(elems, &kv.Elem[kv.OP]{Op: kv.Del, Key: bytes.Clone(d.Key), GroupID: d.RouteGroupID}) } return append(elems, claimElems...), nil } @@ -946,7 +1078,7 @@ func foldSimpleLenDeltas( elems := make([]*kv.Elem[kv.OP], 0, 1+len(deltaKVs)) elems = append(elems, metaElem) for _, d := range deltaKVs { - elems = append(elems, &kv.Elem[kv.OP]{Op: kv.Del, Key: bytes.Clone(d.Key)}) + elems = append(elems, &kv.Elem[kv.OP]{Op: kv.Del, Key: bytes.Clone(d.Key), GroupID: d.RouteGroupID}) } return elems, nil } diff --git a/adapter/redis_delta_compactor_test.go b/adapter/redis_delta_compactor_test.go index b63e23451..956b4ecae 100644 --- a/adapter/redis_delta_compactor_test.go +++ b/adapter/redis_delta_compactor_test.go @@ -130,6 +130,24 @@ func TestDeltaCompactor_TTLInlineMigratesListUserKeyStartingWithDeltaPrefix(t *t require.Equal(t, delta, rawDelta) } +func TestIsListMetaMigrationDeltaRecognizesLegacyRowsByValue(t *testing.T) { + t.Parallel() + + userKey := []byte("list") + delta := store.MarshalListMetaDelta(store.ListMetaDelta{LenDelta: 1}) + legacyKey := legacyListMetaDeltaKey(userKey, 2) + + require.True(t, isListMetaMigrationDelta(&store.KVPair{Key: legacyKey, Value: delta})) + require.True(t, isListMetaMigrationDelta(&store.KVPair{ + Key: store.ListMetaDeltaKey(userKey, 2, 0), + Value: delta, + })) + require.False(t, isListMetaMigrationDelta(&store.KVPair{ + Key: store.ListMetaKey(append([]byte("d|"), userKey...)), + Value: make([]byte, redisWideMetaLegacySizeBytes), + })) +} + func TestDeltaCompactor_TTLInlineMigratesLegacyStreamTTL(t *testing.T) { t.Parallel() @@ -960,6 +978,32 @@ func TestDeltaCompactor_DoesNotInheritStaleTTLBeforeDeltaOnlyRecreate(t *testing return meta.Len, meta.ExpireAt }, }, + { + name: "list-legacy", + key: []byte("ttl:compact:recreate:list-legacy"), + metaKey: store.ListMetaKey, + seed: func(t *testing.T, st store.MVCCStore, key []byte) { + t.Helper() + require.NoError(t, st.PutAt(ctx, store.ListItemKey(key, 0), []byte("a"), 3, 0)) + deltaKey := legacyListMetaDeltaKey(key, 3) + require.NoError(t, st.PutAt(ctx, deltaKey, store.MarshalListMetaDelta(store.ListMetaDelta{LenDelta: 1}), 3, 0)) + }, + build: func(t *testing.T, c *DeltaCompactor, st store.MVCCStore, key []byte, readTS uint64) []*kv.Elem[kv.OP] { + t.Helper() + prefix := store.LegacyListMetaDeltaScanPrefix(key) + deltas, err := st.ScanAt(ctx, prefix, store.PrefixScanEnd(prefix), 10, readTS) + require.NoError(t, err) + elems, err := c.buildLegacyListCompactElems(ctx, key, deltas, readTS) + require.NoError(t, err) + return elems + }, + readMeta: func(t *testing.T, raw []byte) (int64, uint64) { + t.Helper() + meta, err := store.UnmarshalListMeta(raw) + require.NoError(t, err) + return meta.Len, meta.ExpireAt + }, + }, { name: "hash", key: []byte("ttl:compact:recreate:hash"), @@ -1249,6 +1293,7 @@ func TestDeltaCompactor_RotatesHandlerAfterTimeout(t *testing.T) { wantPrefixes := []string{ store.ListMetaDeltaPrefix, + store.LegacyListMetaDeltaPrefix, store.HashMetaDeltaPrefix, store.SetMetaDeltaPrefix, store.ZSetMetaDeltaPrefix, @@ -1304,6 +1349,42 @@ func TestDeltaCompactor_ListDeltaFoldedIntoBaseMeta(t *testing.T) { } } +func TestDeltaCompactor_LegacyListDeltaFoldedIntoBaseMeta(t *testing.T) { + t.Parallel() + + st, c := newDeltaCompactorTestFixture(t) + ctx := context.Background() + userKey := []byte("legacy-list") + + baseMeta := store.ListMeta{Head: 5, Len: 2} + baseMeta.Tail = baseMeta.Head + baseMeta.Len + metaBytes, err := store.MarshalListMeta(baseMeta) + require.NoError(t, err) + require.NoError(t, st.PutAt(ctx, store.ListMetaKey(userKey), metaBytes, 1, 0)) + + delta := store.MarshalListMetaDelta(store.ListMetaDelta{HeadDelta: -1, LenDelta: 2}) + d1Key := legacyListMetaDeltaKey(userKey, 10) + d2Key := legacyListMetaDeltaKey(userKey, 11) + require.NoError(t, st.PutAt(ctx, d1Key, delta, 10, 0)) + require.NoError(t, st.PutAt(ctx, d2Key, delta, 11, 0)) + + require.NoError(t, c.SyncOnce(ctx)) + + readTS := st.LastCommitTS() + raw, err := st.GetAt(ctx, store.ListMetaKey(userKey), readTS) + require.NoError(t, err) + got, err := store.UnmarshalListMeta(raw) + require.NoError(t, err) + require.Equal(t, int64(3), got.Head) + require.Equal(t, int64(6), got.Len) + require.Equal(t, int64(9), got.Tail) + + for _, dk := range [][]byte{d1Key, d2Key} { + _, getErr := st.GetAt(ctx, dk, readTS) + require.ErrorIs(t, getErr, store.ErrKeyNotFound, "legacy delta key should be deleted after compaction: %s", dk) + } +} + func TestDeltaCompactor_ListBelowThresholdNotCompacted(t *testing.T) { t.Parallel() @@ -1492,6 +1573,121 @@ func TestDeltaCompactor_ListNoBaseMeta(t *testing.T) { require.ErrorIs(t, err, store.ErrKeyNotFound) } +func TestDeltaCompactor_LegacyListCursorAdvancesPastFilteredRawRows(t *testing.T) { + t.Parallel() + + st, c := newDeltaCompactorTestFixture(t) + ctx := context.Background() + meta, err := store.MarshalListMeta(store.ListMeta{Head: 4, Tail: 6, Len: 2}) + require.NoError(t, err) + for i := uint64(1); i <= deltaCompactorTickScanLimit; i++ { + userKey := deltaLookingListMetaUserKeyAt([]byte("compactor-collision"), i, 0) + require.NoError(t, st.PutAt(ctx, store.ListMetaKey(userKey), meta, i, 0)) + } + + h := c.legacyListHandler() + readTS := st.LastCommitTS() + require.NoError(t, c.compactHandler(ctx, h, readTS)) + + c.cursorMu.Lock() + got := bytes.Clone(c.cursors[h.typeName]) + c.cursorMu.Unlock() + require.Equal(t, store.ListMetaKey(deltaLookingListMetaUserKeyAt([]byte("compactor-collision"), deltaCompactorTickScanLimit, 0)), got) +} + +func TestDeltaCompactor_LegacyListCursorAdvancesPastFilteredTailAfterBacktrack(t *testing.T) { + t.Parallel() + + st, c := newDeltaCompactorTestFixture(t) + ctx := context.Background() + userKey := []byte("compactor-tail") + delta := store.MarshalListMetaDelta(store.ListMetaDelta{HeadDelta: 0, LenDelta: 1}) + acceptedKey := legacyListMetaDeltaKey(userKey, 1) + require.NoError(t, st.PutAt(ctx, acceptedKey, delta, 1, 0)) + + meta, err := store.MarshalListMeta(store.ListMeta{Head: 4, Tail: 6, Len: 2}) + require.NoError(t, err) + for i := uint64(2); i <= deltaCompactorTickScanLimit; i++ { + collidingUserKey := deltaLookingListMetaUserKeyAt(userKey, i, 0) + require.NoError(t, st.PutAt(ctx, store.ListMetaKey(collidingUserKey), meta, i, 0)) + } + + h := c.legacyListHandler() + readTS := st.LastCommitTS() + require.NoError(t, c.compactHandler(ctx, h, readTS)) + + c.cursorMu.Lock() + got := bytes.Clone(c.cursors[h.typeName]) + c.cursorMu.Unlock() + require.Equal(t, store.ListMetaKey(deltaLookingListMetaUserKeyAt(userKey, deltaCompactorTickScanLimit, 0)), got) + _, err = st.GetAt(ctx, acceptedKey, readTS) + require.NoError(t, err) +} + +func TestDeltaCompactor_LegacyListCursorAdvancesPastRejectedPrefixBeforeAcceptedTail(t *testing.T) { + t.Parallel() + + st, c := newDeltaCompactorTestFixture(t) + ctx := context.Background() + userKey := []byte("compactor-tail-after-prefix") + meta, err := store.MarshalListMeta(store.ListMeta{Head: 4, Tail: 6, Len: 2}) + require.NoError(t, err) + for i := uint64(1); i < deltaCompactorTickScanLimit; i++ { + collidingUserKey := deltaLookingListMetaUserKeyAt(userKey, i, 0) + require.NoError(t, st.PutAt(ctx, store.ListMetaKey(collidingUserKey), meta, i, 0)) + } + delta := store.MarshalListMetaDelta(store.ListMetaDelta{HeadDelta: 0, LenDelta: 1}) + acceptedKey := legacyListMetaDeltaKey(userKey, deltaCompactorTickScanLimit) + require.NoError(t, st.PutAt(ctx, acceptedKey, delta, deltaCompactorTickScanLimit, 0)) + + h := c.legacyListHandler() + readTS := st.LastCommitTS() + require.NoError(t, c.compactHandler(ctx, h, readTS)) + + c.cursorMu.Lock() + got := bytes.Clone(c.cursors[h.typeName]) + c.cursorMu.Unlock() + require.Equal(t, store.ListMetaKey(deltaLookingListMetaUserKeyAt(userKey, deltaCompactorTickScanLimit-1, 0)), got) + _, err = st.GetAt(ctx, acceptedKey, readTS) + require.NoError(t, err) +} + +func TestDeltaCompactor_CursorBacktracksBeforeWholeAcceptedTail(t *testing.T) { + t.Parallel() + + userKey := []byte("compactor-accepted-tail") + delta := store.MarshalListMetaDelta(store.ListMetaDelta{HeadDelta: 0, LenDelta: 1}) + d1 := legacyListMetaDeltaKey(userKey, 1) + d2 := legacyListMetaDeltaKey(userKey, 2) + rawKVs := []*store.KVPair{ + {Key: d1, Value: delta}, + {Key: d2, Value: delta}, + } + + got := deltaCompactorCursorAfterFiltering(rawKVs, rawKVs, nil, true, store.ExtractLegacyListUserKeyFromDelta) + require.Nil(t, got) +} + +func TestDeltaCompactor_ListDeltaDeletesPreserveScanRouteGroup(t *testing.T) { + t.Parallel() + + _, c := newDeltaCompactorTestFixture(t) + ctx := context.Background() + userKey := []byte("compactor-route-group") + delta := store.MarshalListMetaDelta(store.ListMetaDelta{LenDelta: 1}) + d1 := legacyListMetaDeltaKey(userKey, 1) + d2 := legacyListMetaDeltaKey(userKey, 2) + + elems, err := c.buildListCompactElems(ctx, userKey, []*store.KVPair{ + {Key: d1, Value: delta, RouteGroupID: 7}, + {Key: d2, Value: delta, RouteGroupID: 7}, + }, 3) + require.NoError(t, err) + require.Zero(t, elems[0].GroupID) + require.Equal(t, uint64(7), requireElemByKey(t, elems, d1).GroupID) + require.Equal(t, uint64(7), requireElemByKey(t, elems, d2).GroupID) +} + // TestDeltaCompactor_UrgentCompactionTriggeredByChannel verifies that a request // queued via TriggerUrgentCompaction is processed by the Run loop, compacting // the targeted key without waiting for the next regular tick. diff --git a/adapter/redis_exec_dedup_test.go b/adapter/redis_exec_dedup_test.go index 2edfaa5fb..fb643663b 100644 --- a/adapter/redis_exec_dedup_test.go +++ b/adapter/redis_exec_dedup_test.go @@ -47,6 +47,31 @@ func TestExecDedup_LandedPriorAttempt_ReturnsCachedResults(t *testing.T) { require.Equal(t, []byte("v1"), val) } +func TestExecDedup_RouteFenceRetryPreservesPriorProbe(t *testing.T) { + t.Parallel() + ctx := context.Background() + st := store.NewMVCCStore() + coord := newDedupTestCoordinator(st, 1, true) + coord.routeFenceAtDispatch = 2 + srv := &RedisServer{store: st, coordinator: coord, scriptCache: map[string]string{}, onePhaseTxnDedup: true} + + queue := []redcon.Command{ + {Args: [][]byte{[]byte(cmdSet), []byte("k"), []byte("v1")}}, + } + results, err := srv.runTransaction(queue) + require.NoError(t, err) + require.Len(t, results, 1) + require.Equal(t, "OK", results[0].str) + require.Equal(t, 3, coord.dispatches) + require.Equal(t, 1, coord.probeNoOps, "route-fenced reuse must not replace the prior landed probe") + + rawVal, err := st.GetAt(ctx, redisStrKey([]byte("k")), snapshotTS(coord.Clock(), st)) + require.NoError(t, err) + val, _, err := decodeRedisStr(rawVal) + require.NoError(t, err) + require.Equal(t, []byte("v1"), val) +} + // TestExecDedup_PriorAttemptDidNotLand_Applies covers the truncated case for // MULTI/EXEC: attempt 1 errored without committing (OCC-style pre-reject), // so the probe misses and the reuse applies the same write set at a fresh diff --git a/adapter/redis_list_dedup_test.go b/adapter/redis_list_dedup_test.go index 428bbe110..93ab303d0 100644 --- a/adapter/redis_list_dedup_test.go +++ b/adapter/redis_list_dedup_test.go @@ -3,6 +3,7 @@ package adapter import ( "bytes" "context" + "encoding/binary" "errors" "testing" @@ -46,8 +47,12 @@ type dedupTestCoordinator struct { // forwarding. It exercises adapter-side type restoration without changing // the underlying landed-vs-not-landed scenario. wireWriteConflicts bool - dispatches int - probeNoOps int + // routeFenceAtDispatch makes the named dispatch return ErrRouteWriteFenced + // before the FSM dedup probe or apply. It is retryable but cannot be + // treated as an ambiguous landing. + routeFenceAtDispatch int + dispatches int + probeNoOps int // beforeDispatch, if set, runs at the start of each Dispatch with the // 1-based dispatch number — lets a test inject a concurrent commit // between the adapter's attempts. @@ -62,22 +67,235 @@ func newDedupTestCoordinator(st store.MVCCStore, ambiguousDispatch int, lands bo } } +func TestResolveListMetaReadsLegacyDeltaPrefix(t *testing.T) { + t.Parallel() + + ctx := context.Background() + st := store.NewMVCCStore() + srv := &RedisServer{store: st} + key := []byte("legacy-list") + base, err := store.MarshalListMeta(store.ListMeta{Head: 10, Tail: 12, Len: 2}) + require.NoError(t, err) + require.NoError(t, st.PutAt(ctx, store.ListMetaKey(key), base, 1, 0)) + delta := store.MarshalListMetaDelta(store.ListMetaDelta{HeadDelta: -1, LenDelta: 3}) + require.NoError(t, st.PutAt(ctx, legacyListMetaDeltaKey(key, 2), delta, 2, 0)) + + meta, exists, err := srv.resolveListMeta(ctx, key, 3) + require.NoError(t, err) + require.True(t, exists) + require.Equal(t, int64(9), meta.Head) + require.Equal(t, int64(5), meta.Len) + require.Equal(t, int64(14), meta.Tail) +} + +func TestResolveListMetaEnforcesDeltaLimitAcrossCurrentAndLegacyPrefixes(t *testing.T) { + t.Parallel() + + ctx := context.Background() + st := store.NewMVCCStore() + srv := &RedisServer{store: st} + key := []byte("list-delta-cap") + delta := store.MarshalListMetaDelta(store.ListMetaDelta{LenDelta: 1}) + for i := uint64(1); i <= uint64(store.MaxDeltaScanLimit); i++ { + require.NoError(t, st.PutAt(ctx, store.ListMetaDeltaKey(key, i, 0), delta, i, 0)) + } + legacyTS := uint64(store.MaxDeltaScanLimit + 1) + require.NoError(t, st.PutAt(ctx, legacyListMetaDeltaKey(key, legacyTS), delta, legacyTS, 0)) + + _, _, err := srv.resolveListMeta(ctx, key, legacyTS) + require.ErrorIs(t, err, ErrDeltaScanTruncated) +} + +func TestResolveListMetaDoesNotTreatDeltaLookingMetaValueAsLegacyDelta(t *testing.T) { + t.Parallel() + + ctx := context.Background() + st := store.NewMVCCStore() + srv := &RedisServer{store: st} + key := deltaLookingListMetaUserKey([]byte("embedded")) + base, err := store.MarshalListMeta(store.ListMeta{Head: 4, Tail: 6, Len: 2}) + require.NoError(t, err) + require.NoError(t, st.PutAt(ctx, store.ListMetaKey(key), base, 1, 0)) + + meta, exists, err := srv.resolveListMeta(ctx, key, 2) + require.NoError(t, err) + require.True(t, exists) + require.Equal(t, int64(4), meta.Head) + require.Equal(t, int64(2), meta.Len) + require.Equal(t, int64(6), meta.Tail) +} + +func TestResolveListMetaIgnoresLegacyDeltaPrefixCollisionForMissingKey(t *testing.T) { + t.Parallel() + + ctx := context.Background() + st := store.NewMVCCStore() + srv := &RedisServer{store: st} + key := []byte("legacy-collision") + collidingUserKey := deltaLookingListMetaUserKey(key) + base, err := store.MarshalListMeta(store.ListMeta{Head: 4, Tail: 6, Len: 2}) + require.NoError(t, err) + require.NoError(t, st.PutAt(ctx, store.ListMetaKey(collidingUserKey), base, 1, 0)) + + _, exists, err := srv.resolveListMeta(ctx, key, 2) + require.NoError(t, err) + require.False(t, exists) +} + +func TestResolveListMetaCountsOnlyAcceptedLegacyDeltasForTruncation(t *testing.T) { + t.Parallel() + + ctx := context.Background() + st := store.NewMVCCStore() + srv := &RedisServer{store: st} + key := []byte("legacy-collision-window") + base, err := store.MarshalListMeta(store.ListMeta{Head: 0, Tail: 1, Len: 1}) + require.NoError(t, err) + require.NoError(t, st.PutAt(ctx, store.ListMetaKey(key), base, 1, 0)) + + collidingMeta, err := store.MarshalListMeta(store.ListMeta{Head: 4, Tail: 6, Len: 2}) + require.NoError(t, err) + for i := uint64(2); i < uint64(store.MaxDeltaScanLimit+2); i++ { + collidingUserKey := deltaLookingListMetaUserKeyAt(key, i, 0) + require.NoError(t, st.PutAt(ctx, store.ListMetaKey(collidingUserKey), collidingMeta, i, 0)) + } + delta := store.MarshalListMetaDelta(store.ListMetaDelta{HeadDelta: 0, LenDelta: 1}) + deltaTS := uint64(store.MaxDeltaScanLimit + 2) + require.NoError(t, st.PutAt(ctx, legacyListMetaDeltaKey(key, deltaTS), delta, deltaTS, 0)) + + meta, exists, err := srv.resolveListMeta(ctx, key, deltaTS+1) + require.NoError(t, err) + require.True(t, exists) + require.Equal(t, int64(2), meta.Len) +} + +func TestProbeListTypeReadsLegacyDeltaPrefix(t *testing.T) { + t.Parallel() + + ctx := context.Background() + st := store.NewMVCCStore() + srv := &RedisServer{store: st} + key := []byte("legacy-list-type") + delta := store.MarshalListMetaDelta(store.ListMetaDelta{HeadDelta: 0, LenDelta: 1}) + require.NoError(t, st.PutAt(ctx, legacyListMetaDeltaKey(key, 2), delta, 2, 0)) + + typ, found, err := srv.probeListType(ctx, key, 3) + require.NoError(t, err) + require.True(t, found) + require.Equal(t, redisTypeList, typ) +} + +func TestProbeListTypeIgnoresLegacyDeltaPrefixCollision(t *testing.T) { + t.Parallel() + + ctx := context.Background() + st := store.NewMVCCStore() + srv := &RedisServer{store: st} + key := []byte("legacy-list-type-collision") + collidingUserKey := deltaLookingListMetaUserKey(key) + base, err := store.MarshalListMeta(store.ListMeta{Head: 4, Tail: 6, Len: 2}) + require.NoError(t, err) + require.NoError(t, st.PutAt(ctx, store.ListMetaKey(collidingUserKey), base, 1, 0)) + + _, found, err := srv.probeListType(ctx, key, 2) + require.NoError(t, err) + require.False(t, found) +} + +func TestProbeListTypePagesPastLegacyDeltaPrefixCollisions(t *testing.T) { + t.Parallel() + + ctx := context.Background() + st := store.NewMVCCStore() + srv := &RedisServer{store: st} + key := []byte("legacy-list-type-page") + collidingMeta, err := store.MarshalListMeta(store.ListMeta{Head: 4, Tail: 6, Len: 2}) + require.NoError(t, err) + for i := uint64(1); i <= uint64(store.MaxDeltaScanLimit); i++ { + collidingUserKey := deltaLookingListMetaUserKeyAt(key, i, 0) + require.NoError(t, st.PutAt(ctx, store.ListMetaKey(collidingUserKey), collidingMeta, i, 0)) + } + deltaTS := uint64(store.MaxDeltaScanLimit + 1) + delta := store.MarshalListMetaDelta(store.ListMetaDelta{HeadDelta: 0, LenDelta: 1}) + require.NoError(t, st.PutAt(ctx, legacyListMetaDeltaKey(key, deltaTS), delta, deltaTS, 0)) + + typ, found, err := srv.probeListType(ctx, key, deltaTS+1) + require.NoError(t, err) + require.True(t, found) + require.Equal(t, redisTypeList, typ) +} + +func TestDeleteListElemsFiltersLegacyDeltaPrefixCollisions(t *testing.T) { + t.Parallel() + + ctx := context.Background() + st := store.NewMVCCStore() + srv := &RedisServer{store: st} + key := []byte("legacy-list-delete") + collidingUserKey := deltaLookingListMetaUserKey(key) + base, err := store.MarshalListMeta(store.ListMeta{Head: 4, Tail: 6, Len: 2}) + require.NoError(t, err) + collidingMetaKey := store.ListMetaKey(collidingUserKey) + require.NoError(t, st.PutAt(ctx, collidingMetaKey, base, 1, 0)) + deltaKey := legacyListMetaDeltaKey(key, 2) + delta := store.MarshalListMetaDelta(store.ListMetaDelta{HeadDelta: 0, LenDelta: 1}) + require.NoError(t, st.PutAt(ctx, deltaKey, delta, 2, 0)) + + elems, err := srv.deleteListElems(ctx, key, 3) + require.NoError(t, err) + deleted := make(map[string]struct{}, len(elems)) + for _, elem := range elems { + if elem.Op == kv.Del { + deleted[string(elem.Key)] = struct{}{} + } + } + require.Contains(t, deleted, string(deltaKey)) + require.NotContains(t, deleted, string(collidingMetaKey)) +} + +func legacyListMetaDeltaKey(userKey []byte, commitTS uint64) []byte { + key := store.LegacyListMetaDeltaScanPrefix(userKey) + var ts [8]byte + binary.BigEndian.PutUint64(ts[:], commitTS) + key = append(key, ts[:]...) + var seq [4]byte + binary.BigEndian.PutUint32(seq[:], 0) + return append(key, seq[:]...) +} + +func deltaLookingListMetaUserKey(fakeUserKey []byte) []byte { + return deltaLookingListMetaUserKeyAt(fakeUserKey, 7, 1) +} + +func deltaLookingListMetaUserKeyAt(fakeUserKey []byte, commitTS uint64, seqInTxn uint32) []byte { + key := make([]byte, 0, len("d|")+4+len(fakeUserKey)+8+4) + key = append(key, "d|"...) + var lenPrefix [4]byte + binary.BigEndian.PutUint32(lenPrefix[:], uint32(len(fakeUserKey))) //nolint:gosec // test data is small. + key = append(key, lenPrefix[:]...) + key = append(key, fakeUserKey...) + var ts [8]byte + binary.BigEndian.PutUint64(ts[:], commitTS) + key = append(key, ts[:]...) + var seq [4]byte + binary.BigEndian.PutUint32(seq[:], seqInTxn) + return append(key, seq[:]...) +} + func (c *dedupTestCoordinator) Dispatch(ctx context.Context, req *kv.OperationGroup[kv.OP]) (*kv.CoordinateResponse, error) { c.dispatches++ n := c.dispatches if c.beforeDispatch != nil { c.beforeDispatch(n) } + if c.shouldRouteFence(n) { + return nil, kv.ErrRouteWriteFenced + } if handled, resp, err := c.maybeProbe(ctx, req); handled { return resp, err } - if n == c.ambiguousDispatch && !c.ambiguousLands { - // OCC-style pre-reject: nothing is written, definitely did not land. - return nil, c.maybeWireWriteConflict(req, store.ErrWriteConflict) - } - if n == c.txnLockedAtDispatch { - // Ambiguous lock error, nothing written: definitely did not land. - return nil, kv.ErrTxnLocked + if err := c.preApplyError(n); err != nil { + return nil, c.maybeWireWriteConflict(req, err) } resp, err := c.occAdapterCoordinator.Dispatch(ctx, req) if err != nil { @@ -108,6 +326,21 @@ func (c *dedupTestCoordinator) maybeWireWriteConflict(req *kv.OperationGroup[kv. return status.Error(codes.Unknown, store.NewWriteConflictError(key).Error()) } +func (c *dedupTestCoordinator) shouldRouteFence(dispatch int) bool { + return dispatch == c.routeFenceAtDispatch +} + +func (c *dedupTestCoordinator) preApplyError(dispatch int) error { + switch { + case dispatch == c.ambiguousDispatch && !c.ambiguousLands: + return store.ErrWriteConflict + case dispatch == c.txnLockedAtDispatch: + return kv.ErrTxnLocked + default: + return nil + } +} + // maybeProbe mimics handleOnePhaseTxnRequest's exact-ts dedup check. It // returns handled=true when the probe owns the response (hit → no-op success // or probe error), and handled=false to fall through to the normal apply. @@ -186,6 +419,27 @@ func TestListPushDedup_LandedPriorAttempt_NoDuplicate(t *testing.T) { require.Equal(t, []byte("v"), val) } +func TestListPushDedup_RouteFenceRetryPreservesPriorProbe(t *testing.T) { + t.Parallel() + ctx := context.Background() + st := store.NewMVCCStore() + coord := newDedupTestCoordinator(st, 1, true) + coord.routeFenceAtDispatch = 2 + srv := &RedisServer{store: st, coordinator: coord, scriptCache: map[string]string{}, onePhaseTxnDedup: true} + + key := []byte("mylist") + n, err := srv.listRPush(ctx, key, [][]byte{[]byte("v")}) + require.NoError(t, err) + require.Equal(t, int64(1), n) + require.Equal(t, 3, coord.dispatches, "attempt 1 landed, route-fenced reuse, then dedup probe retry") + require.Equal(t, 1, coord.probeNoOps, "route-fenced reuse must not replace the prior landed probe") + + readTS := snapshotTS(coord.Clock(), st) + meta, _, err := srv.resolveListMeta(ctx, key, readTS) + require.NoError(t, err) + require.Equal(t, int64(1), meta.Len, "route-fence retry must not append a duplicate") +} + // TestListPushDedup_PriorAttemptDidNotLand_Applies covers the truncated case: // attempt 1 errored without committing, so the probe misses and the reuse // applies the same write set at a fresh commit_ts. The element lands exactly diff --git a/adapter/redis_lists.go b/adapter/redis_lists.go index 1f7a6fcf2..f80780fbb 100644 --- a/adapter/redis_lists.go +++ b/adapter/redis_lists.go @@ -259,13 +259,10 @@ func (r *RedisServer) dispatchListPushReuse(ctx context.Context, key []byte, pen // iteration recomputes from a fresh meta read. return 0, true, errors.WithStack(dispErr) } - // Still ambiguous (lock / other retryable): this reuse may itself - // have landed, so the next retry must probe THIS commit_ts. Only - // advance pending.commitTS if retryRedisWrite will actually loop - // (non-retryable errors escape to the client; pending is then - // discarded with the goroutine, so the update is wasted and the - // stale value would be misleading if some future caller reads it). - if isRetryableRedisTxnErr(dispErr) { + // Still ambiguous (lock / other retryable): this reuse may itself have + // landed, so the next retry must probe THIS commit_ts. Route-fence + // rejections are retryable but pre-apply, so keep the older witness. + if shouldPreserveRedisTxnAttempt(dispErr) { pending.commitTS = commitTS } return 0, false, errors.WithStack(dispErr) @@ -432,7 +429,9 @@ func (r *RedisServer) listPushCoreWithDedup(ctx context.Context, key []byte, val // operations instead of recomputing a second list append. dispErr = normalizeRetryableRedisTxnErr(dispErr) // Only remember the attempt for reuse if retryRedisWrite will actually - // loop — i.e. the error is one of WriteConflict / TxnLocked. For + // loop and the attempt may have landed. Route-fence rejections are + // retryable but happen before this write set can apply, so preserving + // that commitTS would overwrite an older ambiguous witness. For // errors that escape the loop (transient-leader, context deadline, // FSM apply error, etc.), `pending` would be discarded with the // goroutine, and recording it would mislead a future reader about @@ -440,7 +439,7 @@ func (r *RedisServer) listPushCoreWithDedup(ctx context.Context, key []byte, val // retryRedisWrite's retry predicate; ambiguous errors that escape // to the client are a separate problem space (cross-request // idempotency cache) and out of scope for this design. - if isRetryableRedisTxnErr(dispErr) { + if shouldPreserveRedisTxnAttempt(dispErr) { pending = &reusableListPush{ ops: ops, startTS: startTS, diff --git a/adapter/redis_lua_compat_test.go b/adapter/redis_lua_compat_test.go index a77bc631f..85abb76d2 100644 --- a/adapter/redis_lua_compat_test.go +++ b/adapter/redis_lua_compat_test.go @@ -2,6 +2,8 @@ package adapter import ( "context" + "strconv" + "strings" "testing" "time" @@ -97,6 +99,327 @@ return {eventId, redis.call("HGET", KEYS[1], "name"), cjson.encode(event), redis require.Equal(t, map[string]any{"event": "waiting", "jobId": "job-1"}, events[0].Values) } +func TestRedis_LuaXAddMaxLenExistingStream(t *testing.T) { + nodes, _, _ := createNode(t, 3) + defer shutdown(nodes) + + ctx := context.Background() + rdb := redis.NewClient(&redis.Options{Addr: nodes[0].redisAddress}) + defer func() { _ = rdb.Close() }() + + const stream = "bull:test:events-existing" + for i := range 128 { + _, err := rdb.XAdd(ctx, &redis.XAddArgs{ + Stream: stream, + ID: "*", + Values: []string{"i", strconv.Itoa(i)}, + }).Result() + require.NoError(t, err) + } + + id, err := rdb.Eval(ctx, ` +return redis.call("XADD", KEYS[1], "MAXLEN", "~", 10, "*", "event", "waiting", "jobId", "job-1") +`, []string{stream}).Text() + require.NoError(t, err) + require.NotEmpty(t, id) + + xlen, err := rdb.XLen(ctx, stream).Result() + require.NoError(t, err) + require.Equal(t, int64(10), xlen) + + events, err := rdb.XRange(ctx, stream, "-", "+").Result() + require.NoError(t, err) + require.Len(t, events, 10) + require.Equal(t, id, events[len(events)-1].ID) + require.Equal(t, map[string]any{"event": "waiting", "jobId": "job-1"}, events[len(events)-1].Values) +} + +func TestRedis_LuaXAddMaxLenZero(t *testing.T) { + nodes, _, _ := createNode(t, 3) + defer shutdown(nodes) + + ctx := context.Background() + rdb := redis.NewClient(&redis.Options{Addr: nodes[0].redisAddress}) + defer func() { _ = rdb.Close() }() + + const stream = "bull:test:events-maxlen0" + for i := range 3 { + _, err := rdb.XAdd(ctx, &redis.XAddArgs{ + Stream: stream, + ID: "*", + Values: []string{"i", strconv.Itoa(i)}, + }).Result() + require.NoError(t, err) + } + + id, err := rdb.Eval(ctx, ` +return redis.call("XADD", KEYS[1], "MAXLEN", "0", "*", "event", "trimmed") +`, []string{stream}).Text() + require.NoError(t, err) + require.NotEmpty(t, id) + + xlen, err := rdb.XLen(ctx, stream).Result() + require.NoError(t, err) + require.Equal(t, int64(0), xlen) + + events, err := rdb.XRange(ctx, stream, "-", "+").Result() + require.NoError(t, err) + require.Empty(t, events) +} + +func TestRedis_LuaXAddMaxLenZeroThenAppendInScript(t *testing.T) { + nodes, _, _ := createNode(t, 3) + defer shutdown(nodes) + + ctx := context.Background() + rdb := redis.NewClient(&redis.Options{Addr: nodes[0].redisAddress}) + defer func() { _ = rdb.Close() }() + + const stream = "bull:test:events-maxlen0-then-append" + result, err := rdb.Eval(ctx, ` +local trimmed = redis.call("XADD", KEYS[1], "MAXLEN", "0", "*", "event", "trimmed") +local kept = redis.call("XADD", KEYS[1], "MAXLEN", "1", "*", "event", "kept") +return {trimmed, kept} +`, []string{stream}).Result() + require.NoError(t, err) + ids, ok := result.([]any) + require.True(t, ok) + require.Len(t, ids, 2) + + xlen, err := rdb.XLen(ctx, stream).Result() + require.NoError(t, err) + require.Equal(t, int64(1), xlen) + + events, err := rdb.XRange(ctx, stream, "-", "+").Result() + require.NoError(t, err) + require.Len(t, events, 1) + require.Equal(t, ids[1], events[0].ID) + require.Equal(t, map[string]any{"event": "kept"}, events[0].Values) +} + +func TestRedis_LuaXAddMultipleMaxLenTrimsInScript(t *testing.T) { + nodes, _, _ := createNode(t, 3) + defer shutdown(nodes) + + ctx := context.Background() + rdb := redis.NewClient(&redis.Options{Addr: nodes[0].redisAddress}) + defer func() { _ = rdb.Close() }() + + const stream = "bull:test:events-multi-xadd" + result, err := rdb.Eval(ctx, ` +local first = redis.call("XADD", KEYS[1], "MAXLEN", "1", "*", "event", "first") +local second = redis.call("XADD", KEYS[1], "MAXLEN", "1", "*", "event", "second") +return {first, second} +`, []string{stream}).Result() + require.NoError(t, err) + ids, ok := result.([]any) + require.True(t, ok) + require.Len(t, ids, 2) + + xlen, err := rdb.XLen(ctx, stream).Result() + require.NoError(t, err) + require.Equal(t, int64(1), xlen) + + events, err := rdb.XRange(ctx, stream, "-", "+").Result() + require.NoError(t, err) + require.Len(t, events, 1) + require.Equal(t, ids[1], events[0].ID) + require.Equal(t, map[string]any{"event": "second"}, events[0].Values) +} + +func TestRedis_LuaXAddAndExpirePreservesUpdatedStreamMeta(t *testing.T) { + nodes, _, _ := createNode(t, 3) + defer shutdown(nodes) + + ctx := context.Background() + rdb := redis.NewClient(&redis.Options{Addr: nodes[0].redisAddress}) + defer func() { _ = rdb.Close() }() + + const stream = "bull:test:events-xadd-expire" + firstID, err := rdb.XAdd(ctx, &redis.XAddArgs{ + Stream: stream, + ID: "*", + Values: []string{"event", "first"}, + }).Result() + require.NoError(t, err) + + result, err := rdb.Eval(ctx, ` +local id = redis.call("XADD", KEYS[1], "*", "event", "second") +local applied = redis.call("PEXPIRE", KEYS[1], 60000) +return {id, applied} +`, []string{stream}).Result() + require.NoError(t, err) + values, ok := result.([]any) + require.True(t, ok) + require.Len(t, values, 2) + secondID, ok := values[0].(string) + require.True(t, ok) + require.Equal(t, int64(1), values[1]) + require.Greater(t, secondID, firstID) + + xlen, err := rdb.XLen(ctx, stream).Result() + require.NoError(t, err) + require.Equal(t, int64(2), xlen) + events, err := rdb.XRange(ctx, stream, "-", "+").Result() + require.NoError(t, err) + require.Len(t, events, 2) + require.Equal(t, secondID, events[1].ID) + require.Equal(t, map[string]any{"event": "second"}, events[1].Values) + + ttl, err := rdb.PTTL(ctx, stream).Result() + require.NoError(t, err) + require.Positive(t, ttl) + + thirdID, err := rdb.XAdd(ctx, &redis.XAddArgs{ + Stream: stream, + ID: "*", + Values: []string{"event", "third"}, + }).Result() + require.NoError(t, err) + require.Greater(t, thirdID, secondID) +} + +func TestRedis_LuaXAddRecreatesTTLExpiredStream(t *testing.T) { + nodes, _, _ := createNode(t, 3) + defer shutdown(nodes) + + ctx := context.Background() + rdb := redis.NewClient(&redis.Options{Addr: nodes[0].redisAddress}) + defer func() { _ = rdb.Close() }() + + const ( + stream = "bull:test:events-expired" + ttl = 80 * time.Millisecond + ) + _, err := rdb.XAdd(ctx, &redis.XAddArgs{ + Stream: stream, + ID: "*", + Values: []string{"event", "old"}, + }).Result() + require.NoError(t, err) + require.NoError(t, rdb.PExpire(ctx, stream, ttl).Err()) + eventuallyExpired(t, ttl, func() bool { + xlen, err := rdb.XLen(ctx, stream).Result() + return err == nil && xlen == 0 + }, "stream must be logically expired before Lua XADD recreates it") + + id, err := rdb.Eval(ctx, ` +return redis.call("XADD", KEYS[1], "MAXLEN", "~", 10, "*", "event", "new") +`, []string{stream}).Text() + require.NoError(t, err) + require.NotEmpty(t, id) + + xlen, err := rdb.XLen(ctx, stream).Result() + require.NoError(t, err) + require.Equal(t, int64(1), xlen) + ttlAfter, err := rdb.TTL(ctx, stream).Result() + require.NoError(t, err) + require.Equal(t, time.Duration(-1), ttlAfter) + + events, err := rdb.XRange(ctx, stream, "-", "+").Result() + require.NoError(t, err) + require.Len(t, events, 1) + require.Equal(t, id, events[0].ID) + require.Equal(t, map[string]any{"event": "new"}, events[0].Values) +} + +func TestRedis_LuaXAddRecreatesTTLExpiredHash(t *testing.T) { + nodes, _, _ := createNode(t, 3) + defer shutdown(nodes) + + ctx := context.Background() + rdb := redis.NewClient(&redis.Options{Addr: nodes[0].redisAddress}) + defer func() { _ = rdb.Close() }() + + const ( + stream = "bull:test:events-expired-hash" + ttl = 80 * time.Millisecond + ) + require.NoError(t, rdb.HSet(ctx, stream, "event", "old").Err()) + require.NoError(t, rdb.PExpire(ctx, stream, ttl).Err()) + eventuallyExpired(t, ttl, func() bool { + hlen, err := rdb.HLen(ctx, stream).Result() + return err == nil && hlen == 0 + }, "hash must be logically expired before Lua XADD recreates it") + + id, err := rdb.Eval(ctx, ` +return redis.call("XADD", KEYS[1], "MAXLEN", "~", 10, "*", "event", "new") +`, []string{stream}).Text() + require.NoError(t, err) + require.NotEmpty(t, id) + + xlen, err := rdb.XLen(ctx, stream).Result() + require.NoError(t, err) + require.Equal(t, int64(1), xlen) + ttlAfter, err := rdb.TTL(ctx, stream).Result() + require.NoError(t, err) + require.Equal(t, time.Duration(-1), ttlAfter) + + events, err := rdb.XRange(ctx, stream, "-", "+").Result() + require.NoError(t, err) + require.Len(t, events, 1) + require.Equal(t, id, events[0].ID) + require.Equal(t, map[string]any{"event": "new"}, events[0].Values) +} + +func TestRedis_LuaXAddRecreatesTTLExpiredString(t *testing.T) { + nodes, _, _ := createNode(t, 3) + defer shutdown(nodes) + + ctx := context.Background() + rdb := redis.NewClient(&redis.Options{Addr: nodes[0].redisAddress}) + defer func() { _ = rdb.Close() }() + + const ( + stream = "bull:test:events-expired-string" + ttl = 80 * time.Millisecond + ) + require.NoError(t, rdb.Set(ctx, stream, "old", ttl).Err()) + eventuallyExpired(t, ttl, func() bool { + exists, err := rdb.Exists(ctx, stream).Result() + return err == nil && exists == 0 + }, "string must be logically expired before Lua XADD recreates it") + + id, err := rdb.Eval(ctx, ` +return redis.call("XADD", KEYS[1], "MAXLEN", "~", 10, "*", "event", "new") +`, []string{stream}).Text() + require.NoError(t, err) + require.NotEmpty(t, id) + + typ, err := rdb.Type(ctx, stream).Result() + require.NoError(t, err) + require.Equal(t, "stream", typ) + xlen, err := rdb.XLen(ctx, stream).Result() + require.NoError(t, err) + require.Equal(t, int64(1), xlen) + ttlAfter, err := rdb.TTL(ctx, stream).Result() + require.NoError(t, err) + require.Equal(t, time.Duration(-1), ttlAfter) + + events, err := rdb.XRange(ctx, stream, "-", "+").Result() + require.NoError(t, err) + require.Len(t, events, 1) + require.Equal(t, id, events[0].ID) + require.Equal(t, map[string]any{"event": "new"}, events[0].Values) +} + +func TestRedis_LuaXAddHonorsScriptLocalStringType(t *testing.T) { + nodes, _, _ := createNode(t, 3) + defer shutdown(nodes) + + ctx := context.Background() + rdb := redis.NewClient(&redis.Options{Addr: nodes[0].redisAddress}) + defer func() { _ = rdb.Close() }() + + const key = "bull:test:events-local-string" + _, err := rdb.Eval(ctx, ` +redis.call("SET", KEYS[1], "string-value") +return redis.call("XADD", KEYS[1], "*", "event", "bad") +`, []string{key}).Result() + require.Error(t, err) + require.Contains(t, strings.ToUpper(err.Error()), "WRONGTYPE") +} + func TestRedis_LuaReplyHelpers(t *testing.T) { nodes, _, _ := createNode(t, 3) defer shutdown(nodes) diff --git a/adapter/redis_lua_context.go b/adapter/redis_lua_context.go index 94cd49627..5427a3b8c 100644 --- a/adapter/redis_lua_context.go +++ b/adapter/redis_lua_context.go @@ -194,6 +194,14 @@ type luaPhysicalLimitedScanStore interface { var errLuaStreamDeltaNeedsFullRewrite = errors.New("lua stream delta requires full rewrite") +type luaExactScanFallbackDecider interface { + AllowExactScanFallbackAfterPhysicalLimit(ctx context.Context, start []byte, end []byte, visibleLimit, physicalLimit int, ts uint64, reverse bool) bool +} + +type luaRoutePinnedScanStore interface { + ScanAtWithReadFence(ctx context.Context, start []byte, end []byte, limit int, ts uint64, reverse bool, groupID uint64, readRouteVersion uint64, routeStart []byte, routeEnd []byte) ([]*store.KVPair, error) +} + type luaCommandHandler func(*luaScriptContext, []string) (luaReply, error) type luaRenameHandler func(*luaScriptContext, []byte, []byte) error @@ -1019,7 +1027,10 @@ func (c *luaScriptContext) streamType(key []byte) (redisValueType, error) { func (c *luaScriptContext) streamState(key []byte) (*luaStreamState, error) { k := string(key) if st, ok := c.streams[k]; ok { - return st, nil + if st.loaded { + return st, nil + } + return c.materializeStreamState(key, st) } st := &luaStreamState{} c.streams[k] = st @@ -1070,6 +1081,14 @@ func (c *luaScriptContext) streamState(key []byte) (*luaStreamState, error) { return st, nil } +func (c *luaScriptContext) materializeStreamState(key []byte, st *luaStreamState) (*luaStreamState, error) { + if err := c.materializeStream(key, st); err != nil { + return nil, err + } + st.loaded = true + return st, nil +} + // materializeStream reconstructs the script-visible stream only for commands // that require arbitrary entry access. XADD, XLEN, and head trims stay lazy. func (c *luaScriptContext) materializeStream(key []byte, st *luaStreamState) error { @@ -2324,6 +2343,81 @@ func (c *luaScriptContext) scanListItemWindow(key []byte, meta store.ListMeta, s if err != nil { return luaLazyListBoundaryItem{}, false, errors.WithStack(err) } + if item, ok := luaListBoundaryItemFromKVs(key, meta, kvs); ok { + return item, true, nil + } + if physicalLimitReached { + if !c.allowExactListScanFallback(startKey, endKey, scanLimit, !left) { + return luaLazyListBoundaryItem{}, false, errors.Wrapf(ErrCollectionTooLarge, + "list %q sparse pop scanned %d physical item rows", string(key), scanLimit) + } + return c.scanListItemWindowExact(key, meta, startKey, endKey, left) + } + if len(kvs) >= scanLimit { + return luaLazyListBoundaryItem{}, false, errors.Wrapf(ErrCollectionTooLarge, + "list %q sparse pop scanned %d non-matching item rows", string(key), scanLimit) + } + return luaLazyListBoundaryItem{}, false, nil +} + +func (c *luaScriptContext) allowExactListScanFallback(startKey, endKey []byte, scanLimit int, reverse bool) bool { + decider, ok := c.server.store.(luaExactScanFallbackDecider) + if !ok { + return true + } + return decider.AllowExactScanFallbackAfterPhysicalLimit(c.ctx, startKey, endKey, scanLimit, scanLimit, c.startTS, reverse) +} + +func (c *luaScriptContext) scanListItemWindowExact(key []byte, meta store.ListMeta, startKey, endKey []byte, left bool) (luaLazyListBoundaryItem, bool, error) { + routeStart, routeEnd := luaListRouteScanBounds(key) + for { + kvs, err := c.scanExactListPage(startKey, endKey, routeStart, routeEnd, left) + if err != nil { + return luaLazyListBoundaryItem{}, false, err + } + if len(kvs) == 0 { + return luaLazyListBoundaryItem{}, false, nil + } + if item, ok := luaListBoundaryItemFromKVs(key, meta, kvs); ok { + return item, true, nil + } + if left { + startKey = scanStartAfterKey(kvs[0].Key) + } else { + endKey = append([]byte(nil), kvs[0].Key...) + } + } +} + +func luaListRouteScanBounds(key []byte) ([]byte, []byte) { + return append([]byte(nil), key...), prefixScanEnd(key) +} + +func (c *luaScriptContext) scanExactListPage(startKey, endKey, routeStart, routeEnd []byte, left bool) ([]*store.KVPair, error) { + var ( + kvs []*store.KVPair + err error + ) + if scanner, ok := c.server.store.(luaRoutePinnedScanStore); ok { + kvs, err = scanner.ScanAtWithReadFence(c.ctx, startKey, endKey, 1, c.startTS, !left, 0, 0, routeStart, routeEnd) + } else if left { + kvs, err = c.server.store.ScanAt(c.ctx, startKey, endKey, 1, c.startTS) + } else { + kvs, err = c.server.store.ReverseScanAt(c.ctx, startKey, endKey, 1, c.startTS) + } + if err != nil { + return nil, errors.WithStack(err) + } + return kvs, nil +} + +func scanStartAfterKey(key []byte) []byte { + next := make([]byte, len(key)+1) + copy(next, key) + return next +} + +func luaListBoundaryItemFromKVs(key []byte, meta store.ListMeta, kvs []*store.KVPair) (luaLazyListBoundaryItem, bool) { for _, kvp := range kvs { seq, ok := store.ExtractListItemSeq(kvp.Key, key) if !ok { @@ -2332,17 +2426,9 @@ func (c *luaScriptContext) scanListItemWindow(key []byte, meta store.ListMeta, s return luaLazyListBoundaryItem{ value: string(kvp.Value), index: seq - meta.Head, - }, true, nil + }, true } - if physicalLimitReached { - return luaLazyListBoundaryItem{}, false, errors.Wrapf(ErrCollectionTooLarge, - "list %q sparse pop scanned %d physical item rows", string(key), scanLimit) - } - if len(kvs) >= scanLimit { - return luaLazyListBoundaryItem{}, false, errors.Wrapf(ErrCollectionTooLarge, - "list %q sparse pop scanned %d non-matching item rows", string(key), scanLimit) - } - return luaLazyListBoundaryItem{}, false, nil + return luaLazyListBoundaryItem{}, false } func (c *luaScriptContext) scanAtPhysicalLimit(startKey, endKey []byte, scanLimit int) ([]*store.KVPair, bool, error) { @@ -3285,6 +3371,10 @@ func (c *luaScriptContext) cmdXAdd(args []string) (luaReply, error) { if err != nil { return luaReply{}, err } + return c.cmdXAddMaterialized(parsed) +} + +func (c *luaScriptContext) cmdXAddMaterialized(parsed luaXAddArgs) (luaReply, error) { st, err := c.streamState(parsed.key) if err != nil { return luaReply{}, err @@ -4186,7 +4276,7 @@ func (c *luaScriptContext) streamCommitPlan(ctx context.Context, key string) (lu elems, err := c.streamCommitElems(ctx, key) return luaCommitPlan{elems: elems, inlineMetaRewritten: true}, err } - elems, err := c.streamDeltaCommitElems(ctx, key, st) + elems, err := c.streamStateDeltaCommitElems(ctx, key, st) if errors.Is(err, errLuaStreamDeltaNeedsFullRewrite) { if materializeErr := c.materializeStream([]byte(key), st); materializeErr != nil { return luaCommitPlan{}, materializeErr @@ -4197,6 +4287,62 @@ func (c *luaScriptContext) streamCommitPlan(ctx context.Context, key string) (lu return luaCommitPlan{preserveExisting: true, inlineMetaRewritten: true, elems: elems}, err } +// streamStateDeltaCommitElems emits only newly appended entries, bounded head +// tombstones, and the authoritative metadata record. Untouched entries remain +// in place, avoiding an O(stream length) Lua XADD rewrite. +func (c *luaScriptContext) streamStateDeltaCommitElems( + ctx context.Context, + key string, + st *luaStreamState, +) ([]*kv.Elem[kv.OP], error) { + keyBytes := []byte(key) + ttl, err := c.finalTTL(ctx, keyBytes) + if err != nil { + return nil, err + } + + trimCount := int(st.baseTrim) + trimElems, trimmedThrough, err := c.server.buildXTrimHeadElems(ctx, keyBytes, c.startTS, trimCount, st.baseMeta) + if err != nil { + return nil, err + } + if len(trimElems) != trimCount { + return nil, errLuaStreamDeltaNeedsFullRewrite + } + elems := make([]*kv.Elem[kv.OP], 0, len(trimElems)+len(st.appended)+1) + elems = append(elems, trimElems...) + for _, entry := range st.appended { + parsed, err := parseRedisStreamID(entry.ID) + if err != nil { + return nil, errors.WithStack(err) + } + entryValue, err := marshalStreamEntry(entry) + if err != nil { + return nil, err + } + elems = append(elems, &kv.Elem[kv.OP]{ + Op: kv.Put, + Key: store.StreamEntryKey(keyBytes, parsed.ms, parsed.seq), + Value: entryValue, + }) + } + meta := st.meta + meta.ExpireAt = ttlMillis(ttl) + if trimmedThrough.ok { + meta.TrimmedMs, meta.TrimmedSeq = trimmedThrough.ms, trimmedThrough.seq + } + metaBytes, err := store.MarshalStreamMeta(meta) + if err != nil { + return nil, errors.WithStack(err) + } + elems = append(elems, &kv.Elem[kv.OP]{ + Op: kv.Put, + Key: store.StreamMetaKey(keyBytes), + Value: metaBytes, + }) + return elems, nil +} + // streamCommitElems writes the script's final stream state in the // wide-column layout — one StreamEntryKey Put per entry plus a StreamMetaKey // Put for the aggregate meta. The legacy single-blob path is no longer @@ -4268,62 +4414,6 @@ func marshalLuaStreamEntries( return elems, meta, nil } -// streamDeltaCommitElems emits only the newly appended entries, bounded head -// tombstones, and the authoritative metadata record. Untouched entries remain -// in place, avoiding the O(stream length) Lua XADD rewrite. -func (c *luaScriptContext) streamDeltaCommitElems( - ctx context.Context, - key string, - st *luaStreamState, -) ([]*kv.Elem[kv.OP], error) { - keyBytes := []byte(key) - ttl, err := c.finalTTL(ctx, keyBytes) - if err != nil { - return nil, err - } - - trimCount := int(st.baseTrim) - trimElems, trimmedThrough, err := c.server.buildXTrimHeadElems(ctx, keyBytes, c.startTS, trimCount, st.baseMeta) - if err != nil { - return nil, err - } - if len(trimElems) != trimCount { - return nil, errLuaStreamDeltaNeedsFullRewrite - } - elems := make([]*kv.Elem[kv.OP], 0, len(trimElems)+len(st.appended)+1) - elems = append(elems, trimElems...) - for _, entry := range st.appended { - parsed, err := parseRedisStreamID(entry.ID) - if err != nil { - return nil, errors.WithStack(err) - } - entryValue, err := marshalStreamEntry(entry) - if err != nil { - return nil, err - } - elems = append(elems, &kv.Elem[kv.OP]{ - Op: kv.Put, - Key: store.StreamEntryKey(keyBytes, parsed.ms, parsed.seq), - Value: entryValue, - }) - } - meta := st.meta - meta.ExpireAt = ttlMillis(ttl) - if trimmedThrough.ok { - meta.TrimmedMs, meta.TrimmedSeq = trimmedThrough.ms, trimmedThrough.seq - } - metaBytes, err := store.MarshalStreamMeta(meta) - if err != nil { - return nil, errors.WithStack(err) - } - elems = append(elems, &kv.Elem[kv.OP]{ - Op: kv.Put, - Key: store.StreamMetaKey(keyBytes), - Value: metaBytes, - }) - return elems, nil -} - func (c *luaScriptContext) finalType(ctx context.Context, key []byte) (redisValueType, error) { if typ, ok := c.cachedType(key); ok { return typ, nil diff --git a/adapter/redis_lua_list_holes_test.go b/adapter/redis_lua_list_holes_test.go index f3211385e..fd4f215d2 100644 --- a/adapter/redis_lua_list_holes_test.go +++ b/adapter/redis_lua_list_holes_test.go @@ -217,7 +217,7 @@ func TestRedisLua_RPopLPushKeepsRemainingSparseHeadItem(t *testing.T) { require.Equal(t, []string{"job-2"}, dstValues) } -func TestRedisLua_RPopLPushFailsOnTooManyPhysicalTailTombstones(t *testing.T) { +func TestRedisLua_RPopLPushScansPastPhysicalTailTombstones(t *testing.T) { t.Parallel() ctx := context.Background() @@ -226,6 +226,8 @@ func TestRedisLua_RPopLPushFailsOnTooManyPhysicalTailTombstones(t *testing.T) { dst := []byte("bull:test:active:tombstone-tail") largeLen := int64(luaSparseListPopScanLimit) + 2 seedListMeta(t, r, src, store.ListMeta{Head: 0, Len: largeLen, Tail: largeLen}) + require.NoError(t, r.store.PutAt(ctx, listItemKey(src, 0), []byte("job-1"), 2, 0)) + seedListPrefixCollider(t, r, ctx, src, 0) commitTS := uint64(3) for seq := int64(1); seq <= int64(luaSparseListPopScanLimit)+1; seq++ { require.NoError(t, r.store.DeleteAt(ctx, listItemKey(src, seq), commitTS)) @@ -236,10 +238,142 @@ func TestRedisLua_RPopLPushFailsOnTooManyPhysicalTailTombstones(t *testing.T) { require.NoError(t, err) defer scriptCtx.Close() + reply, err := scriptCtx.cmdRPopLPush([]string{string(src), string(dst)}) + require.NoError(t, err) + require.Equal(t, luaReplyString, reply.kind) + require.Equal(t, "job-1", reply.text) + require.NoError(t, scriptCtx.commit()) + + readTS := r.readTS() + typ, err := r.keyTypeAt(ctx, src, readTS) + require.NoError(t, err) + require.Equal(t, redisTypeNone, typ) + dstValues, err := r.listValuesAt(ctx, dst, readTS) + require.NoError(t, err) + require.Equal(t, []string{"job-1"}, dstValues) +} + +func TestRedisLua_LPopScansPastPhysicalHeadTombstones(t *testing.T) { + t.Parallel() + + ctx := context.Background() + r := newListPopTestServer(t) + src := []byte("bull:test:wait:tombstone-head") + largeLen := int64(luaSparseListPopScanLimit) + 2 + seedListMeta(t, r, src, store.ListMeta{Head: 0, Len: largeLen, Tail: largeLen}) + require.NoError(t, r.store.PutAt(ctx, listItemKey(src, largeLen-1), []byte("job-1"), 2, 0)) + seedListPrefixCollider(t, r, ctx, src, 0) + commitTS := uint64(3) + for seq := int64(0); seq <= int64(luaSparseListPopScanLimit); seq++ { + require.NoError(t, r.store.DeleteAt(ctx, listItemKey(src, seq), commitTS)) + commitTS++ + } + + scriptCtx, err := newLuaScriptContext(ctx, r) + require.NoError(t, err) + defer scriptCtx.Close() + + reply, err := scriptCtx.cmdLPop([]string{string(src)}) + require.NoError(t, err) + require.Equal(t, luaReplyString, reply.kind) + require.Equal(t, "job-1", reply.text) + require.NoError(t, scriptCtx.commit()) + + readTS := r.readTS() + typ, err := r.keyTypeAt(ctx, src, readTS) + require.NoError(t, err) + require.Equal(t, redisTypeNone, typ) +} + +func TestRedisLua_RPopLPushDeletesPhysicalTombstoneOnlyList(t *testing.T) { + t.Parallel() + + ctx := context.Background() + r := newListPopTestServer(t) + src := []byte("bull:test:wait:tombstone-only") + dst := []byte("bull:test:active:tombstone-only") + largeLen := int64(luaSparseListPopScanLimit) + 2 + seedListMeta(t, r, src, store.ListMeta{Head: 0, Len: largeLen, Tail: largeLen}) + commitTS := uint64(2) + for seq := int64(0); seq <= int64(luaSparseListPopScanLimit)+1; seq++ { + require.NoError(t, r.store.DeleteAt(ctx, listItemKey(src, seq), commitTS)) + commitTS++ + } + + scriptCtx, err := newLuaScriptContext(ctx, r) + require.NoError(t, err) + defer scriptCtx.Close() + + reply, err := scriptCtx.cmdRPopLPush([]string{string(src), string(dst)}) + require.NoError(t, err) + require.Equal(t, luaReplyNil, reply.kind) + require.NoError(t, scriptCtx.commit()) + + readTS := r.readTS() + typ, err := r.keyTypeAt(ctx, src, readTS) + require.NoError(t, err) + require.Equal(t, redisTypeNone, typ) + dstValues, err := r.listValuesAt(ctx, dst, readTS) + require.NoError(t, err) + require.Empty(t, dstValues) +} + +func TestRedisLua_RPopLPushSyntheticPhysicalLimitFailsClosed(t *testing.T) { + t.Parallel() + + ctx := context.Background() + base := store.NewMVCCStore() + st := noExactFallbackPhysicalLimitStore{MVCCStore: base} + coord := newLocalAdapterCoordinator(st) + r := NewRedisServer(nil, "", st, coord, nil, nil) + src := []byte("bull:test:wait:synthetic-limit") + dst := []byte("bull:test:active:synthetic-limit") + seedListMeta(t, r, src, store.ListMeta{Head: 0, Len: 1, Tail: 1}) + require.NoError(t, r.store.PutAt(ctx, listItemKey(src, 0), []byte("job-1"), 2, 0)) + + scriptCtx, err := newLuaScriptContext(ctx, r) + require.NoError(t, err) + defer scriptCtx.Close() + _, err = scriptCtx.cmdRPopLPush([]string{string(src), string(dst)}) require.ErrorIs(t, err, ErrCollectionTooLarge) } +func TestRedisLua_ListExactFallbackPinsRouteBounds(t *testing.T) { + t.Parallel() + + ctx := context.Background() + src := []byte("bull:test:wait:route-pin") + meta := store.ListMeta{Head: 0, Len: 2, Tail: 2} + startKey := listItemKey(src, 0) + endKey := listItemKey(src, 2) + + collider := append([]byte{}, src...) + collider = append(collider, listItemKey(src, 0)[len(store.ListItemPrefix)+len(src):]...) + collider = append(collider, 'x') + st := &routePinnedExactScanStore{ + MVCCStore: store.NewMVCCStore(), + t: t, + wantRouteStart: src, + wantRouteEnd: prefixScanEnd(src), + pages: [][]*store.KVPair{ + {{Key: listItemKey(collider, 0), Value: []byte("other-list-job")}}, + {{Key: listItemKey(src, 1), Value: []byte("job-1")}}, + }, + } + r := NewRedisServer(nil, "", st, newLocalAdapterCoordinator(st), nil, nil) + scriptCtx, err := newLuaScriptContext(ctx, r) + require.NoError(t, err) + defer scriptCtx.Close() + + item, ok, err := scriptCtx.scanListItemWindowExact(src, meta, startKey, endKey, true) + require.NoError(t, err) + require.True(t, ok) + require.Equal(t, int64(1), item.index) + require.Equal(t, "job-1", item.value) + require.Equal(t, 2, st.calls) +} + func TestRedisLua_RPopLPushDeletesLargeSparseListWithoutItems(t *testing.T) { t.Parallel() @@ -275,3 +409,61 @@ func seedListMeta(t *testing.T, r *RedisServer, key []byte, meta store.ListMeta) require.NoError(t, err) require.NoError(t, r.store.PutAt(context.Background(), store.ListMetaKey(key), raw, 1, 0)) } + +func seedListPrefixCollider(t *testing.T, r *RedisServer, ctx context.Context, key []byte, seq int64) { + t.Helper() + + collider := append([]byte{}, key...) + collider = append(collider, listItemKey(key, seq)[len(store.ListItemPrefix)+len(key):]...) + collider = append(collider, 'x') + require.NoError(t, r.store.PutAt(ctx, listItemKey(collider, 0), []byte("other-list-job"), 2, 0)) +} + +type noExactFallbackPhysicalLimitStore struct { + store.MVCCStore +} + +func (s noExactFallbackPhysicalLimitStore) ScanAtPhysicalLimit(context.Context, []byte, []byte, int, int, uint64) ([]*store.KVPair, bool, error) { + return nil, true, nil +} + +func (s noExactFallbackPhysicalLimitStore) ReverseScanAtPhysicalLimit(context.Context, []byte, []byte, int, int, uint64) ([]*store.KVPair, bool, error) { + return nil, true, nil +} + +func (s noExactFallbackPhysicalLimitStore) AllowExactScanFallbackAfterPhysicalLimit(context.Context, []byte, []byte, int, int, uint64, bool) bool { + return false +} + +type routePinnedExactScanStore struct { + store.MVCCStore + t *testing.T + wantRouteStart []byte + wantRouteEnd []byte + pages [][]*store.KVPair + calls int +} + +func (s *routePinnedExactScanStore) ScanAtWithReadFence( + _ context.Context, + _ []byte, + _ []byte, + _ int, + _ uint64, + _ bool, + _ uint64, + _ uint64, + routeStart []byte, + routeEnd []byte, +) ([]*store.KVPair, error) { + s.t.Helper() + require.Equal(s.t, s.wantRouteStart, routeStart) + require.Equal(s.t, s.wantRouteEnd, routeEnd) + if s.calls >= len(s.pages) { + s.calls++ + return nil, nil + } + page := s.pages[s.calls] + s.calls++ + return page, nil +} diff --git a/adapter/redis_retry.go b/adapter/redis_retry.go index 77233cbbb..41ef3ae19 100644 --- a/adapter/redis_retry.go +++ b/adapter/redis_retry.go @@ -46,7 +46,12 @@ var ( func isRetryableRedisTxnErr(err error) bool { return errors.Is(err, store.ErrWriteConflict) || errors.Is(err, kv.ErrTxnLocked) || - wireRedisTxnErrKind(err) == redisTxnWireErrLocked + wireRedisTxnErrKind(err) == redisTxnWireErrLocked || + isRouteWriteFencedError(err) +} + +func shouldPreserveRedisTxnAttempt(err error) bool { + return isRetryableRedisTxnErr(err) && !isRouteWriteFencedError(err) } func retryPolicyForRedisTxnErr(err error) redisTxnRetryPolicy { diff --git a/adapter/redis_retry_test.go b/adapter/redis_retry_test.go index dd5eab7f5..4b893d50f 100644 --- a/adapter/redis_retry_test.go +++ b/adapter/redis_retry_test.go @@ -550,6 +550,7 @@ func TestRetryPolicyForRedisTxnErr(t *testing.T) { require.Equal(t, redisWriteConflictRetryPolicy, retryPolicyForRedisTxnErr(store.ErrWriteConflict)) require.Equal(t, redisTxnLockedRetryPolicy, retryPolicyForRedisTxnErr(kv.ErrTxnLocked)) + require.Equal(t, redisWriteConflictRetryPolicy, retryPolicyForRedisTxnErr(kv.ErrRouteWriteFenced)) } // TestZCard_LegacyBlobZSet verifies that ZCARD inside a Lua script returns the diff --git a/adapter/redis_ttl_inline_migrator.go b/adapter/redis_ttl_inline_migrator.go index 4274f0e62..cc02aed86 100644 --- a/adapter/redis_ttl_inline_migrator.go +++ b/adapter/redis_ttl_inline_migrator.go @@ -423,7 +423,7 @@ func expiredTTLIndexPrecedesDeltaOnlyCollection( if err != nil || baseExists { return false, err } - prefix, ok := collectionMetaDeltaScanPrefix(userKey, typ) + prefixes, ok := collectionMetaDeltaScanPrefixes(userKey, typ) if !ok { return false, nil } @@ -431,26 +431,61 @@ func expiredTTLIndexPrecedesDeltaOnlyCollection( if err != nil || !found { return false, err } - deltas, err := scanner.scanDeltaKVs(ctx, prefix, readTS) - if err != nil { - if errors.Is(err, ErrDeltaScanTruncated) { + return legacyTTLPrecedesCollectionMetaDeltas(ctx, scanner, userKey, prefixes, ttlCommitTS, readTS) +} + +func legacyTTLPrecedesCollectionMetaDeltas( + ctx context.Context, + scanner redisDeltaKVScanner, + userKey []byte, + prefixes [][]byte, + ttlCommitTS uint64, + readTS uint64, +) (bool, error) { + found := false + for _, prefix := range prefixes { + deltas, err := scanner.scanDeltaKVs(ctx, prefix, readTS) + if err != nil { + if errors.Is(err, ErrDeltaScanTruncated) { + return false, nil + } + return false, err + } + if isLegacyListMetaDeltaPrefix(prefix) { + deltas = filterLegacyListMetaDeltas(deltas, userKey) + } + minDeltaTS, ok := minMetaDeltaCommitTS(deltas, prefix) + if !ok { + continue + } + found = true + if ttlCommitTS >= minDeltaTS { return false, nil } - return false, err } - return legacyTTLPrecedesAllMetaDeltas(ttlCommitTS, deltas, prefix), nil + return found, nil +} + +func filterLegacyListMetaDeltas(deltas []*store.KVPair, userKey []byte) []*store.KVPair { + filtered := deltas[:0] + for _, pair := range deltas { + if legacyListDeltaPairForUserKey(pair, userKey) { + filtered = append(filtered, pair) + } + } + return filtered } -func collectionMetaDeltaScanPrefix(userKey []byte, typ redisValueType) ([]byte, bool) { +func collectionMetaDeltaScanPrefixes(userKey []byte, typ redisValueType) ([][]byte, bool) { switch typ { case redisTypeList: - return store.ListMetaDeltaScanPrefix(userKey), true + return store.ListMetaDeltaScanPrefixes(userKey), true case redisTypeHash: - return store.HashMetaDeltaScanPrefix(userKey), true + return [][]byte{store.HashMetaDeltaScanPrefix(userKey)}, true case redisTypeSet: - return store.SetMetaDeltaScanPrefix(userKey), true + return [][]byte{store.SetMetaDeltaScanPrefix(userKey)}, true case redisTypeZSet: - return store.ZSetMetaDeltaScanPrefix(userKey), true + return [][]byte{store.ZSetMetaDeltaScanPrefix(userKey)}, true case redisTypeNone, redisTypeString, redisTypeStream: return nil, false } @@ -657,10 +692,10 @@ func (c *DeltaCompactor) migrateListTTLInlineElems(ctx context.Context, pair *st } func isListMetaMigrationDelta(pair *store.KVPair) bool { - if pair == nil || !store.IsListMetaDeltaKey(pair.Key) { + if pair == nil || !store.IsListMetaDeltaValue(pair.Value) { return false } - return len(pair.Value) != redisWideMetaLegacySizeBytes && len(pair.Value) != redisWideMetaInlineSizeBytes + return store.IsListMetaDeltaKey(pair.Key) || store.ExtractLegacyListUserKeyFromDelta(pair.Key) != nil } func (c *DeltaCompactor) migrateStreamTTLInlineElems(ctx context.Context, pair *store.KVPair, readTS uint64) ([]*kv.Elem[kv.OP], error) { diff --git a/adapter/redis_txn.go b/adapter/redis_txn.go index a38001517..907037582 100644 --- a/adapter/redis_txn.go +++ b/adapter/redis_txn.go @@ -98,9 +98,12 @@ type txnValue struct { } type stringReplacement struct { - key []byte - value []byte - ttl *time.Time + key []byte + value []byte + ttl *time.Time + rawTyp redisValueType + rawTypKnown bool + rawPrefixedString bool } type txnContext struct { @@ -142,7 +145,12 @@ type listTxnState struct { deleted bool purge bool purgeMeta store.ListMeta - existingDeltas [][]byte // delta key bytes present at load time; deleted on purge/delete + existingDeltas []listDeltaRef // delta keys present at load time; deleted on purge/delete +} + +type listDeltaRef struct { + key []byte + groupID uint64 } type hashTxnState struct { @@ -356,18 +364,33 @@ func (t *txnContext) loadListState(key []byte) (*listTxnState, error) { // truncation: if >MaxDeltaScanLimit deltas exist the transaction cannot // safely enumerate all of them for deletion, so we return ErrDeltaScanTruncated // and let the caller retry after the background compactor has caught up. - deltaPrefix := store.ListMetaDeltaScanPrefix(key) - deltaEnd := store.PrefixScanEnd(deltaPrefix) - deltaKVs, err := t.server.store.ScanAt(ctx, deltaPrefix, deltaEnd, store.MaxDeltaScanLimit+1, t.startTS) - if err != nil { - return nil, errors.WithStack(err) - } - if len(deltaKVs) > store.MaxDeltaScanLimit { - return nil, ErrDeltaScanTruncated - } - existingDeltas := make([][]byte, 0, len(deltaKVs)) - for _, kv := range deltaKVs { - existingDeltas = append(existingDeltas, kv.Key) + var existingDeltas []listDeltaRef + acceptedDeltas := 0 + for _, deltaPrefix := range store.ListMetaDeltaScanPrefixes(key) { + deltaKVs, truncated, err := scanAcceptedDeltaKVsAt( + ctx, + t.server.store, + deltaPrefix, + store.MaxDeltaScanLimit, + t.startTS, + listTxnDeltaFilter(key, deltaPrefix), + ) + if err != nil { + return nil, err + } + if truncated { + return nil, ErrDeltaScanTruncated + } + for _, kv := range deltaKVs { + acceptedDeltas++ + if acceptedDeltas > store.MaxDeltaScanLimit { + return nil, ErrDeltaScanTruncated + } + existingDeltas = append(existingDeltas, listDeltaRef{ + key: bytes.Clone(kv.Key), + groupID: kv.RouteGroupID, + }) + } } st := &listTxnState{ @@ -395,6 +418,15 @@ func (t *txnContext) loadListState(key []byte) (*listTxnState, error) { return st, nil } +func listTxnDeltaFilter(key []byte, deltaPrefix []byte) func(*store.KVPair) bool { + if !isLegacyListMetaDeltaPrefix(deltaPrefix) { + return nil + } + return func(pair *store.KVPair) bool { + return legacyListDeltaPairForUserKey(pair, key) + } +} + func (t *txnContext) loadHashStateForFields(key []byte, fields [][]byte) (*hashTxnState, error) { k := string(key) if t.hashStates == nil { @@ -615,18 +647,48 @@ func (t *txnContext) loadTTLState(key []byte) (*ttlTxnState, error) { } func (t *txnContext) stagedKeyType(key []byte) (redisValueType, error) { + view, err := t.stagedKeyTypeView(key) + if err != nil { + return redisTypeNone, err + } + return view.typ, nil +} + +type txnKeyTypeView struct { + typ redisValueType + rawTyp redisValueType + rawTypKnown bool +} + +func (t *txnContext) stagedKeyTypeView(key []byte) (txnKeyTypeView, error) { k := string(key) - if _, ok := t.replacers[k]; ok { - return redisTypeString, nil + if repl, ok := t.replacers[k]; ok { + return txnKeyTypeView{ + typ: redisTypeString, + rawTyp: repl.rawTyp, + rawTypKnown: repl.rawTypKnown, + }, nil } if typ, ok := t.stagedPositiveKeyType(k); ok { - return typ, nil + return txnKeyTypeView{typ: typ}, nil } if t.hasStagedTypeDeletion(k) { - return redisTypeNone, nil + return txnKeyTypeView{typ: redisTypeNone}, nil } t.trackTypeReadKeys(key) - return t.server.keyTypeAt(t.ctxOrBackground(), key, t.startTS) + rawTyp, err := t.server.rawKeyTypeAt(t.ctxOrBackground(), key, t.startTS) + if err != nil { + return txnKeyTypeView{}, err + } + typ, err := t.server.applyTTLFilter(t.ctxOrBackground(), key, t.startTS, rawTyp) + if err != nil { + return txnKeyTypeView{}, err + } + return txnKeyTypeView{ + typ: typ, + rawTyp: rawTyp, + rawTypKnown: true, + }, nil } func (t *txnContext) stagedPositiveKeyType(key string) (redisValueType, bool) { @@ -719,10 +781,11 @@ func (t *txnContext) applySet(cmd redcon.Command) (redisResult, error) { if err != nil { return redisResult{}, err } - typ, err := t.stagedKeyType(cmd.Args[1]) + typeView, err := t.stagedKeyTypeView(cmd.Args[1]) if err != nil { return redisResult{}, err } + typ := typeView.typ // NX/XX: skip the write if the key-existence condition is not met. exists := typ != redisTypeNone @@ -737,7 +800,12 @@ func (t *txnContext) applySet(cmd redcon.Command) (redisResult, error) { if err != nil { return redisResult{}, err } - t.stageStringReplacement(cmd.Args[1], cmd.Args[2], opts.ttl) + t.trackWideCollectionFenceReads(cmd.Args[1]) + rawPrefixedString, err := t.rawPrefixedStringAtStart(cmd.Args[1], typeView) + if err != nil { + return redisResult{}, err + } + t.stageStringReplacementWithRawType(cmd.Args[1], cmd.Args[2], opts.ttl, typeView.rawTyp, typeView.rawTypKnown, rawPrefixedString) return applySetResult(opts, oldValue), nil } @@ -775,19 +843,49 @@ func cloneTimePtr(in *time.Time) *time.Time { } func (t *txnContext) stageStringReplacement(key, value []byte, ttl *time.Time) { + t.stageStringReplacementWithRawType(key, value, ttl, redisTypeNone, false, false) +} + +func (t *txnContext) stageStringReplacementWithRawType(key, value []byte, ttl *time.Time, rawTyp redisValueType, rawTypKnown bool, rawPrefixedString bool) { if t.replacers == nil { t.replacers = map[string]*stringReplacement{} } k := string(key) - t.replacers[k] = &stringReplacement{ - key: bytes.Clone(key), - value: bytes.Clone(value), - ttl: cloneTimePtr(ttl), + if repl, ok := t.replacers[k]; ok { + repl.value = bytes.Clone(value) + repl.ttl = cloneTimePtr(ttl) + if rawTypKnown { + repl.rawTyp = rawTyp + repl.rawTypKnown = true + repl.rawPrefixedString = rawPrefixedString + } + delete(t.deletedKeys, k) + return } + repl := &stringReplacement{ + key: bytes.Clone(key), + value: bytes.Clone(value), + ttl: cloneTimePtr(ttl), + rawTyp: rawTyp, + rawTypKnown: rawTypKnown, + rawPrefixedString: rawPrefixedString, + } + t.replacers[k] = repl delete(t.deletedKeys, k) delete(t.collectionExpireTypes, k) } +func (t *txnContext) rawPrefixedStringAtStart(key []byte, view txnKeyTypeView) (bool, error) { + if !view.rawTypKnown || view.rawTyp != redisTypeString { + return false, nil + } + exists, err := t.server.store.ExistsAt(t.ctxOrBackground(), redisStrKey(key), t.startTS) + if err != nil { + return false, errors.WithStack(err) + } + return exists, nil +} + func (t *txnContext) updateStringReplacementTTL(key []byte, ttl *time.Time) bool { repl, ok := t.replacers[string(key)] if !ok { @@ -1625,11 +1723,13 @@ func (t *txnContext) buildReplacementElems(ctx context.Context) ([]*kv.Elem[kv.O for _, k := range keys { repl := t.replacers[k] t.trackWideCollectionFenceReads(repl.key) - deleteElems, _, err := t.server.deleteLogicalKeyElems(ctx, repl.key, t.startTS) - if err != nil { - return nil, err + if repl.needsFullLogicalDelete() { + deleteElems, _, err := t.server.deleteLogicalKeyElems(ctx, repl.key, t.startTS) + if err != nil { + return nil, err + } + elems = append(elems, deleteElems...) } - elems = append(elems, deleteElems...) elems = append(elems, redisTxnWideCollectionFenceElems(repl.key)...) elems = append(elems, &kv.Elem[kv.OP]{ Op: kv.Put, @@ -1645,6 +1745,16 @@ func (t *txnContext) buildReplacementElems(ctx context.Context) ([]*kv.Elem[kv.O return elems, nil } +func (r *stringReplacement) needsFullLogicalDelete() bool { + if !r.rawTypKnown { + return true + } + if r.rawTyp == redisTypeString { + return !r.rawPrefixedString + } + return isNonStringCollectionType(r.rawTyp) +} + func (t *txnContext) buildLogicalDeletionElems(ctx context.Context) ([]*kv.Elem[kv.OP], error) { if len(t.logicalDeletes) == 0 { return nil, nil @@ -1840,7 +1950,7 @@ func appendListDeletionElems(elems []*kv.Elem[kv.OP], userKey []byte, st *listTx } // Delete existing delta keys so they do not survive logical delete/purge. for _, dk := range st.existingDeltas { - elems = append(elems, &kv.Elem[kv.OP]{Op: kv.Del, Key: dk}) + elems = append(elems, &kv.Elem[kv.OP]{Op: kv.Del, Key: dk.key, GroupID: dk.groupID}) } elems = append(elems, redisTxnWideListFenceElem(userKey)) return elems @@ -2516,12 +2626,10 @@ func (r *RedisServer) dispatchExecReuse(ctx context.Context, pending *reusableEx // iteration rebuilds from a fresh snapshot. return nil, true, errors.WithStack(dispErr) } - // Still ambiguous (lock / other retryable): the reuse may itself - // have landed, so the next retry must probe THIS commit_ts. Only - // advance pending.commitTS if retryRedisWrite will actually loop - // (non-retryable errors escape to the client; pending is then - // discarded with the goroutine). - if isRetryableRedisTxnErr(dispErr) { + // Still ambiguous (lock / other retryable): the reuse may itself have + // landed, so the next retry must probe THIS commit_ts. Route-fence + // rejections are retryable but pre-apply, so keep the older witness. + if shouldPreserveRedisTxnAttempt(dispErr) { pending.commitTS = commitTS } return nil, false, errors.WithStack(dispErr) @@ -2653,12 +2761,10 @@ func (r *RedisServer) firstExecAttempt(dispatchCtx context.Context, queue []redc // write set instead of replaying the EXEC body from a new snapshot. dispErr = normalizeRetryableRedisTxnErr(dispErr) // Only remember the attempt for reuse if retryRedisWrite will - // actually loop. Mirrors listPushCoreWithDedup's gating - // rationale — errors that escape the loop (transient-leader, - // context deadline, FSM apply error) leave pending pointing at - // state wasted with the goroutine; ambiguous errors that - // escape to the client are out of scope for this loop. - if isRetryableRedisTxnErr(dispErr) { + // actually loop and the attempt may have landed. Mirrors + // listPushCoreWithDedup's gating rationale; route-fence + // rejections are retryable but pre-apply. + if shouldPreserveRedisTxnAttempt(dispErr) { return nil, &reusableExecTxn{ elems: prepared.elems, startTS: txn.startTS, diff --git a/adapter/redis_txn_test.go b/adapter/redis_txn_test.go index a44bad0b0..872d728a1 100644 --- a/adapter/redis_txn_test.go +++ b/adapter/redis_txn_test.go @@ -42,6 +42,108 @@ func newRedisTxnTestContext(server *RedisServer) *txnContext { } } +type routeGroupScanStore struct { + store.MVCCStore + groupID uint64 +} + +func (s *routeGroupScanStore) ScanAt(ctx context.Context, start []byte, end []byte, limit int, ts uint64) ([]*store.KVPair, error) { + kvs, err := s.MVCCStore.ScanAt(ctx, start, end, limit, ts) + if err != nil { + return nil, err + } + for _, kvp := range kvs { + kvp.RouteGroupID = s.groupID + } + return kvs, nil +} + +func TestRedisTxnLoadListStateFiltersLegacyDeltaPrefixCollisions(t *testing.T) { + t.Parallel() + + server, st := newRedisStorageMigrationTestServer(t) + ctx := context.Background() + key := []byte("txn-legacy-list-delete") + base, err := store.MarshalListMeta(store.ListMeta{Head: 0, Tail: 1, Len: 1}) + require.NoError(t, err) + require.NoError(t, st.PutAt(ctx, store.ListMetaKey(key), base, 1, 0)) + + collidingMeta, err := store.MarshalListMeta(store.ListMeta{Head: 4, Tail: 6, Len: 2}) + require.NoError(t, err) + collidingMetaKey := store.ListMetaKey(deltaLookingListMetaUserKey(key)) + require.NoError(t, st.PutAt(ctx, collidingMetaKey, collidingMeta, 2, 0)) + + deltaKey := legacyListMetaDeltaKey(key, 3) + delta := store.MarshalListMetaDelta(store.ListMetaDelta{HeadDelta: 0, LenDelta: 1}) + require.NoError(t, st.PutAt(ctx, deltaKey, delta, 3, 0)) + + txn := newRedisTxnTestContext(server) + txn.startTS = 4 + stState, err := txn.loadListState(key) + require.NoError(t, err) + require.Equal(t, []listDeltaRef{{key: deltaKey}}, stState.existingDeltas) +} + +func TestRedisTxnLoadListStatePreservesDeltaRouteGroupID(t *testing.T) { + t.Parallel() + + baseStore := store.NewMVCCStore() + scanStore := &routeGroupScanStore{MVCCStore: baseStore, groupID: 7} + server := NewRedisServer(nil, "", scanStore, newLocalAdapterCoordinator(baseStore), nil, nil) + ctx := context.Background() + key := []byte("txn-list-route-group") + base, err := store.MarshalListMeta(store.ListMeta{Head: 0, Tail: 1, Len: 1}) + require.NoError(t, err) + require.NoError(t, baseStore.PutAt(ctx, store.ListMetaKey(key), base, 1, 0)) + deltaKey := store.ListMetaDeltaKey(key, 2, 0) + delta := store.MarshalListMetaDelta(store.ListMetaDelta{HeadDelta: 0, LenDelta: 1}) + require.NoError(t, baseStore.PutAt(ctx, deltaKey, delta, 2, 0)) + + txn := newRedisTxnTestContext(server) + txn.startTS = 3 + stState, err := txn.loadListState(key) + require.NoError(t, err) + require.Equal(t, []listDeltaRef{{key: deltaKey, groupID: 7}}, stState.existingDeltas) +} + +func TestRedisScanAllDeltaElemsFilteredPreservesRouteGroupID(t *testing.T) { + t.Parallel() + + baseStore := store.NewMVCCStore() + scanStore := &routeGroupScanStore{MVCCStore: baseStore, groupID: 11} + server := NewRedisServer(nil, "", scanStore, newLocalAdapterCoordinator(baseStore), nil, nil) + ctx := context.Background() + key := []byte("scan-list-route-group") + deltaKey := store.ListMetaDeltaKey(key, 2, 0) + delta := store.MarshalListMetaDelta(store.ListMetaDelta{LenDelta: 1}) + require.NoError(t, baseStore.PutAt(ctx, deltaKey, delta, 2, 0)) + + elems, err := server.scanAllDeltaElemsFiltered(ctx, store.ListMetaDeltaScanPrefix(key), 3, nil) + require.NoError(t, err) + require.Len(t, elems, 1) + require.Equal(t, deltaKey, elems[0].Key) + require.Equal(t, uint64(11), elems[0].GroupID) +} + +func TestRedisTxnLoadListStateEnforcesDeltaLimitAcrossCurrentAndLegacyPrefixes(t *testing.T) { + t.Parallel() + + server, st := newRedisStorageMigrationTestServer(t) + ctx := context.Background() + key := []byte("txn-list-delta-cap") + delta := store.MarshalListMetaDelta(store.ListMetaDelta{LenDelta: 1}) + for i := uint64(1); i <= uint64(store.MaxDeltaScanLimit); i++ { + require.NoError(t, st.PutAt(ctx, store.ListMetaDeltaKey(key, i, 0), delta, i, 0)) + } + legacyTS := uint64(store.MaxDeltaScanLimit + 1) + require.NoError(t, st.PutAt(ctx, legacyListMetaDeltaKey(key, legacyTS), delta, legacyTS, 0)) + + txn := newRedisTxnTestContext(server) + txn.startTS = legacyTS + _, err := txn.loadListState(key) + require.ErrorIs(t, err, ErrDeltaScanTruncated) +} + func elemKeysContain(elems []*kv.Elem[kv.OP], want []byte) bool { for _, elem := range elems { if elem != nil && string(elem.Key) == string(want) { @@ -668,6 +770,139 @@ func TestRedisTxnSetReplacementConflictsWithConcurrentWideHashWrite(t *testing.T "SET replacement in MULTI must conflict with concurrent HSET of a new field") } +func TestRedisTxnSetReplacementTracksWideFencesBeforeBuild(t *testing.T) { + t.Parallel() + + ctx := context.Background() + server, st := newRedisStorageMigrationTestServer(t) + key := []byte("set-replace:fence-read") + + txn := newRedisTxnTestContext(server) + res, err := txn.applySet(redcon.Command{Args: [][]byte{[]byte(cmdSet), key, []byte("string")}}) + require.NoError(t, err) + require.Equal(t, "OK", res.str) + for _, fenceKey := range redisTxnWideCollectionFenceKeys(key) { + require.Contains(t, txn.readKeys, string(fenceKey)) + } + + require.NoError(t, st.PutAt(ctx, redisTxnWideHashFenceKey(key), []byte{}, redisTxnTestStartTS+1, 0)) + require.ErrorIs(t, txn.validateReadSet(ctx), store.ErrWriteConflict) +} + +func TestRedisTxnSetReplacementSkipsWideCleanupForRawStringOrMissing(t *testing.T) { + t.Parallel() + + ctx := context.Background() + cases := []struct { + name string + seed func(store.MVCCStore, []byte) + }{ + {name: "missing"}, + { + name: "string", + seed: func(st store.MVCCStore, key []byte) { + require.NoError(t, st.PutAt(ctx, redisStrKey(key), encodeRedisStr([]byte("old"), nil), redisTxnTestStartTS, 0)) + }, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + server, st := newRedisStorageMigrationTestServer(t) + key := []byte("set-replace:no-wide-cleanup:" + tc.name) + if tc.seed != nil { + tc.seed(st, key) + } + + txn := newRedisTxnTestContext(server) + res, err := txn.applySet(redcon.Command{Args: [][]byte{[]byte(cmdSet), key, []byte("string")}}) + require.NoError(t, err) + require.Equal(t, "OK", res.str) + + elems, err := txn.buildReplacementElems(ctx) + require.NoError(t, err) + require.False(t, elemKeysContain(elems, store.HashMetaKey(key))) + require.False(t, elemKeysContain(elems, store.SetMetaKey(key))) + require.False(t, elemKeysContain(elems, store.ZSetMetaKey(key))) + require.False(t, elemKeysContain(elems, store.ListMetaKey(key))) + require.False(t, elemKeysContain(elems, store.StreamMetaKey(key))) + require.True(t, elemKeysContain(elems, redisStrKey(key))) + }) + } +} + +func TestRedisTxnSetReplacementDeletesNonPrefixedStringEncodings(t *testing.T) { + t.Parallel() + + ctx := context.Background() + hllValue, err := encodeRedisHLL(redisSetValue{Members: []string{"member"}}, nil) + require.NoError(t, err) + cases := []struct { + name string + seedKey func([]byte) []byte + seedValue []byte + expectedDel func([]byte) []byte + }{ + { + name: "hll", + seedKey: redisHLLKey, + seedValue: hllValue, + expectedDel: redisHLLKey, + }, + { + name: "legacy-bare-string", + seedKey: func(key []byte) []byte { return key }, + seedValue: []byte("old"), + expectedDel: func(key []byte) []byte { return key }, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + server, st := newRedisStorageMigrationTestServer(t) + key := []byte("set-replace:non-prefixed-string:" + tc.name) + require.NoError(t, st.PutAt(ctx, tc.seedKey(key), tc.seedValue, redisTxnTestStartTS, 0)) + + txn := newRedisTxnTestContext(server) + res, err := txn.applySet(redcon.Command{Args: [][]byte{[]byte(cmdSet), key, []byte("string")}}) + require.NoError(t, err) + require.Equal(t, "OK", res.str) + + elems, err := txn.buildReplacementElems(ctx) + require.NoError(t, err) + require.True(t, elemKeysContain(elems, tc.expectedDel(key))) + require.True(t, elemKeysContain(elems, redisStrKey(key))) + }) + } +} + +func TestRedisTxnSetReplacementDeletesExpiredRawHash(t *testing.T) { + t.Parallel() + + ctx := context.Background() + server, st := newRedisStorageMigrationTestServer(t) + key := []byte("set-replace:expired-hash") + expired := time.Now().Add(-time.Hour) + require.NoError(t, st.PutAt(ctx, store.HashFieldKey(key, []byte("old")), []byte("v"), redisTxnTestStartTS, 0)) + require.NoError(t, st.PutAt(ctx, store.HashMetaKey(key), store.MarshalHashMeta(store.HashMeta{Len: 1}), redisTxnTestStartTS, 0)) + require.NoError(t, st.PutAt(ctx, redisTTLKey(key), encodeRedisTTL(expired), redisTxnTestStartTS, 0)) + + txn := newRedisTxnTestContext(server) + res, err := txn.applySet(redcon.Command{Args: [][]byte{[]byte(cmdSet), key, []byte("string")}}) + require.NoError(t, err) + require.Equal(t, "OK", res.str) + + elems, err := txn.buildReplacementElems(ctx) + require.NoError(t, err) + require.True(t, elemKeysContain(elems, store.HashFieldKey(key, []byte("old")))) + require.True(t, elemKeysContain(elems, store.HashMetaKey(key))) + require.True(t, elemKeysContain(elems, redisStrKey(key))) +} + func TestRedisTxnSetReplacementConflictsWithConcurrentListPush(t *testing.T) { t.Parallel() @@ -1180,6 +1415,19 @@ func TestRedisTxnListDeletionElemsWriteFence(t *testing.T) { require.True(t, elemKeysContain(elems, redisTxnWideListFenceKey(key))) } +func TestRedisTxnListDeletionElemsPreserveDeltaRouteGroupID(t *testing.T) { + t.Parallel() + + key := []byte("delete-route-group:list") + deltaKey := store.ListMetaDeltaKey(key, 9, 0) + elems := appendListDeletionElems(nil, key, &listTxnState{ + existingDeltas: []listDeltaRef{{key: deltaKey, groupID: 17}}, + deleted: true, + }) + require.Equal(t, uint64(17), requireElemByKey(t, elems, deltaKey).GroupID) + require.True(t, elemKeysContain(elems, redisTxnWideListFenceKey(key))) +} + func TestRedisTxnHashLegacyRewriteWritesFence(t *testing.T) { t.Parallel() diff --git a/adapter/retryable_write_fence_test.go b/adapter/retryable_write_fence_test.go new file mode 100644 index 000000000..94cd71811 --- /dev/null +++ b/adapter/retryable_write_fence_test.go @@ -0,0 +1,49 @@ +package adapter + +import ( + "testing" + + "github.com/bootjp/elastickv/kv" + "github.com/bootjp/elastickv/store" + "github.com/cockroachdb/errors" + "github.com/stretchr/testify/require" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" +) + +func TestWriteFenceErrorsAreAdapterRetryable(t *testing.T) { + t.Parallel() + + require.True(t, isRetryableRedisTxnErr(kv.ErrRouteWriteFenced)) + require.True(t, isRetryableS3MutationErr(kv.ErrRouteWriteFenced)) + require.True(t, isRetryableTransactWriteError(kv.ErrRouteWriteFenced)) + require.False(t, shouldPreserveTransactWriteAttempt(kv.ErrRouteWriteFenced)) + require.False(t, isIgnorableTransactRaceError(kv.ErrRouteWriteFenced)) + require.True(t, isIgnorableTransactRaceError(store.ErrWriteConflict)) + require.True(t, isIgnorableTransactRaceError(kv.ErrTxnLocked)) +} + +func TestWireWriteFenceErrorsAreAdapterRetryable(t *testing.T) { + t.Parallel() + + err := errors.WithStack(status.Error( + codes.Unknown, + "commit-version v=12: key \"k\" routeKey \"k\": "+kv.ErrRouteWriteFenced.Error(), + )) + + require.True(t, isRetryableRedisTxnErr(err)) + require.True(t, isRetryableS3MutationErr(err)) + require.True(t, isRetryableTransactWriteError(err)) + require.False(t, shouldPreserveRedisTxnAttempt(err)) + require.False(t, shouldPreserveTransactWriteAttempt(err)) + require.False(t, isIgnorableTransactRaceError(err)) +} + +func TestWireWriteFenceMatcherRequiresSentinelSuffix(t *testing.T) { + t.Parallel() + + err := status.Error(codes.Unknown, kv.ErrRouteWriteFenced.Error()+": "+store.ErrWriteConflict.Error()) + + require.False(t, isRouteWriteFencedError(err)) + require.False(t, isRetryableTransactWriteError(err)) +} diff --git a/adapter/route_write_fence.go b/adapter/route_write_fence.go new file mode 100644 index 000000000..3f7802713 --- /dev/null +++ b/adapter/route_write_fence.go @@ -0,0 +1,37 @@ +package adapter + +import ( + "strings" + + "github.com/bootjp/elastickv/kv" + "github.com/cockroachdb/errors" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" +) + +func isRouteWriteFencedError(err error) bool { + if errors.Is(err, kv.ErrRouteWriteFenced) { + return true + } + st, ok := grpcStatusFromError(err) + if !ok { + return false + } + code := st.Code() + if code != codes.Unknown && code != codes.Aborted && code != codes.FailedPrecondition { + return false + } + return strings.HasSuffix(st.Message(), kv.ErrRouteWriteFenced.Error()) +} + +func grpcStatusFromError(err error) (*status.Status, bool) { + type grpcStatusCarrier interface { + GRPCStatus() *status.Status + } + var carrier grpcStatusCarrier + if !errors.As(err, &carrier) { + return nil, false + } + st := carrier.GRPCStatus() + return st, st != nil +} diff --git a/adapter/s3.go b/adapter/s3.go index 1893af37e..3f52063d4 100644 --- a/adapter/s3.go +++ b/adapter/s3.go @@ -762,9 +762,12 @@ func (s *S3Server) deleteBucket(w http.ResponseWriter, r *http.Request, bucket s writeS3MutationError(w, err, bucket, "") return } - // Phase 2: best-effort DEL_PREFIX safety net. See - // AdminDeleteBucket / runBucketDeleteSafetyNet for the contract. - s.runBucketDeleteSafetyNet(r.Context(), bucket, deletedGeneration) + // Phase 2: DEL_PREFIX safety net. See AdminDeleteBucket / + // runBucketDeleteSafetyNet for the contract. + if err := s.runBucketDeleteSafetyNet(r.Context(), bucket, deletedGeneration); err != nil { + writeS3MutationError(w, err, bucket, "") + return + } w.WriteHeader(http.StatusNoContent) } @@ -1509,7 +1512,7 @@ func (s *S3Server) cleanupPartBlobsAsync( if len(pending) == 0 { return } - if _, err := s.coordinator.Dispatch(ctx, &kv.OperationGroup[kv.OP]{Elems: pending}); err != nil { + if err := s.dispatchS3CleanupBatch(ctx, pending); err != nil { slog.ErrorContext(ctx, "cleanupPartBlobsAsync: coordinator dispatch failed", "bucket", bucket, "object_key", objectKey, @@ -1580,7 +1583,7 @@ func (s *S3Server) deleteByPrefix(ctx context.Context, prefix []byte, bucket str for _, kvp := range kvs { pending = append(pending, &kv.Elem[kv.OP]{Op: kv.Del, Key: kvp.Key}) } - if _, err := s.coordinator.Dispatch(ctx, &kv.OperationGroup[kv.OP]{Elems: pending}); err != nil { + if err := s.dispatchS3CleanupBatch(ctx, pending); err != nil { slog.ErrorContext(ctx, "deleteByPrefix: dispatch failed", "bucket", bucket, "generation", generation, "object_key", objectKey, "upload_id", uploadID, "err", err) @@ -1590,6 +1593,13 @@ func (s *S3Server) deleteByPrefix(ctx context.Context, prefix []byte, bucket str } } +func (s *S3Server) dispatchS3CleanupBatch(ctx context.Context, elems []*kv.Elem[kv.OP]) error { + return s.retryS3Mutation(ctx, func() error { + _, err := s.coordinator.Dispatch(ctx, &kv.OperationGroup[kv.OP]{Elems: elems}) + return errors.WithStack(err) + }) +} + func parseS3MaxParts(raw string) int { if strings.TrimSpace(raw) == "" { return s3ListPartsMaxParts @@ -2527,7 +2537,7 @@ func (s *S3Server) nextTxnCommitTS(ctx context.Context, startTS uint64) (uint64, } func isRetryableS3MutationErr(err error) bool { - return errors.Is(err, store.ErrWriteConflict) || errors.Is(err, kv.ErrTxnLocked) + return errors.Is(err, store.ErrWriteConflict) || errors.Is(err, kv.ErrTxnLocked) || isRouteWriteFencedError(err) } func waitS3RetryBackoff(ctx context.Context, delay time.Duration) bool { diff --git a/adapter/s3_admin.go b/adapter/s3_admin.go index 0d8c58d28..50ebb0ca1 100644 --- a/adapter/s3_admin.go +++ b/adapter/s3_admin.go @@ -423,12 +423,7 @@ func (s *S3Server) AdminDeleteBucket(ctx context.Context, principal AdminPrincip if err != nil { return err //nolint:wrapcheck // sentinel errors propagate as-is. } - // Phase 2: best-effort safety-net DEL_PREFIX. Outside the - // retryS3Mutation closure because retrying after Phase 1 - // committed would 404 at loadBucketMetaAt; we want the error - // (if any) logged but not propagated to the operator. - s.runBucketDeleteSafetyNet(ctx, name, deletedGeneration) - return nil + return s.runBucketDeleteSafetyNet(ctx, name, deletedGeneration) } // adminDeleteBucketTxnBody is the per-attempt body retryS3Mutation @@ -502,22 +497,26 @@ func bucketDeleteSafetyNetElems(bucket string, generation uint64) []*kv.Elem[kv. } } -// runBucketDeleteSafetyNet runs the Phase-2 DEL_PREFIX dispatch -// and swallows transport / cluster errors after logging — the -// caller has already deleted the bucket meta and the operator- -// visible state is consistent with that. Shared between admin and -// SigV4 paths. -func (s *S3Server) runBucketDeleteSafetyNet(ctx context.Context, bucket string, generation uint64) { - if _, err := s.coordinator.Dispatch(ctx, &kv.OperationGroup[kv.OP]{ - Elems: bucketDeleteSafetyNetElems(bucket, generation), - }); err != nil { +// runBucketDeleteSafetyNet runs the Phase-2 DEL_PREFIX dispatch. Phase 1 has +// already deleted the bucket meta, so this helper retries transient fencing in +// place instead of re-entering the full delete transaction. +func (s *S3Server) runBucketDeleteSafetyNet(ctx context.Context, bucket string, generation uint64) error { + err := s.retryS3Mutation(ctx, func() error { + _, err := s.coordinator.Dispatch(ctx, &kv.OperationGroup[kv.OP]{ + Elems: bucketDeleteSafetyNetElems(bucket, generation), + }) + return errors.WithStack(err) + }) + if err != nil { slog.WarnContext(ctx, "bucket delete safety-net DEL_PREFIX failed; bucket meta is gone but orphan sweep incomplete", slog.String("bucket", bucket), slog.Uint64("generation", generation), slog.String("error", err.Error()), ) + return nil } + return nil } // adminCanonicalACL normalises an empty input to the canned diff --git a/adapter/s3_admin_test.go b/adapter/s3_admin_test.go index 5caf92e26..ba32489ec 100644 --- a/adapter/s3_admin_test.go +++ b/adapter/s3_admin_test.go @@ -264,6 +264,59 @@ func TestS3Server_AdminDeleteBucket_HappyPath(t *testing.T) { require.False(t, exists) } +func TestS3Server_AdminDeleteBucket_RetriesSafetyNetRouteFence(t *testing.T) { + t.Parallel() + + st := store.NewMVCCStore() + coord := &bucketDeleteSafetyNetFenceCoordinator{ + localAdapterCoordinator: newLocalAdapterCoordinator(st), + failuresRemaining: 1, + } + server := NewS3Server(nil, "", st, coord, nil) + ctx := context.Background() + + summary, err := server.AdminCreateBucket(ctx, + fullAdminBucketsPrincipal(), "to-delete", s3AclPrivate) + require.NoError(t, err) + orphan := s3keys.RouteKey("to-delete", summary.Generation, "orphan") + _, err = coord.localAdapterCoordinator.Dispatch(ctx, &kv.OperationGroup[kv.OP]{ + Elems: []*kv.Elem[kv.OP]{{Op: kv.Put, Key: orphan, Value: []byte("orphan")}}, + }) + require.NoError(t, err) + + err = server.AdminDeleteBucket(ctx, + fullAdminBucketsPrincipal(), "to-delete") + require.NoError(t, err) + require.Equal(t, 2, coord.safetyNetCalls) + + _, err = st.GetAt(ctx, orphan, snapshotTS(coord.Clock(), st)) + require.ErrorIs(t, err, store.ErrKeyNotFound, "safety-net retry must sweep the orphan before acknowledging delete") +} + +func TestS3Server_AdminDeleteBucket_SwallowsPersistentSafetyNetRouteFence(t *testing.T) { + t.Parallel() + + st := store.NewMVCCStore() + coord := &bucketDeleteSafetyNetFenceCoordinator{ + localAdapterCoordinator: newLocalAdapterCoordinator(st), + failuresRemaining: s3TxnRetryMaxAttempts, + } + server := NewS3Server(nil, "", st, coord, nil) + ctx := context.Background() + + _, err := server.AdminCreateBucket(ctx, + fullAdminBucketsPrincipal(), "to-delete", s3AclPrivate) + require.NoError(t, err) + + err = server.AdminDeleteBucket(ctx, + fullAdminBucketsPrincipal(), "to-delete") + require.NoError(t, err) + require.Equal(t, s3TxnRetryMaxAttempts, coord.safetyNetCalls) + _, exists, err := server.AdminDescribeBucket(ctx, "to-delete") + require.NoError(t, err) + require.False(t, exists) +} + func TestS3Server_AdminDeleteBucket_MissingBucket(t *testing.T) { t.Parallel() @@ -275,6 +328,32 @@ func TestS3Server_AdminDeleteBucket_MissingBucket(t *testing.T) { require.ErrorIs(t, err, ErrAdminBucketNotFound) } +type bucketDeleteSafetyNetFenceCoordinator struct { + *localAdapterCoordinator + failuresRemaining int + safetyNetCalls int +} + +func (c *bucketDeleteSafetyNetFenceCoordinator) Dispatch(ctx context.Context, req *kv.OperationGroup[kv.OP]) (*kv.CoordinateResponse, error) { + if req != nil && operationGroupHasDelPrefix(req.Elems) { + c.safetyNetCalls++ + if c.failuresRemaining > 0 { + c.failuresRemaining-- + return nil, kv.ErrRouteWriteFenced + } + } + return c.localAdapterCoordinator.Dispatch(ctx, req) +} + +func operationGroupHasDelPrefix(elems []*kv.Elem[kv.OP]) bool { + for _, elem := range elems { + if elem != nil && elem.Op == kv.DelPrefix { + return true + } + } + return false +} + func TestS3Server_AdminDeleteBucket_RejectsReadOnly(t *testing.T) { t.Parallel() diff --git a/adapter/s3_cleanup_retry_test.go b/adapter/s3_cleanup_retry_test.go new file mode 100644 index 000000000..c7591a90b --- /dev/null +++ b/adapter/s3_cleanup_retry_test.go @@ -0,0 +1,57 @@ +package adapter + +import ( + "context" + "testing" + + "github.com/bootjp/elastickv/internal/s3keys" + "github.com/bootjp/elastickv/kv" + "github.com/bootjp/elastickv/store" + "github.com/stretchr/testify/require" +) + +func TestS3DeleteByPrefix_RetriesRouteFence(t *testing.T) { + t.Parallel() + + ctx := context.Background() + st := store.NewMVCCStore() + coord := &s3CleanupRouteFenceCoordinator{ + localAdapterCoordinator: newLocalAdapterCoordinator(st), + failuresRemaining: 1, + } + server := NewS3Server(nil, "", st, coord, nil) + key := s3keys.BlobKey("bucket-a", 3, "object-a", "upload-a", 1, 0) + require.NoError(t, st.PutAt(ctx, key, []byte("blob"), 1, 0)) + + server.deleteByPrefix(ctx, s3keys.BlobPrefixForUpload("bucket-a", 3, "object-a", "upload-a"), "bucket-a", 3, "object-a", "upload-a") + + require.Equal(t, 2, coord.calls, "route-fenced cleanup batch should retry") + _, err := st.GetAt(ctx, key, snapshotTS(coord.Clock(), st)) + require.ErrorIs(t, err, store.ErrKeyNotFound) +} + +type s3CleanupRouteFenceCoordinator struct { + *localAdapterCoordinator + failuresRemaining int + calls int +} + +func (c *s3CleanupRouteFenceCoordinator) Dispatch(ctx context.Context, req *kv.OperationGroup[kv.OP]) (*kv.CoordinateResponse, error) { + if req != nil && operationGroupHasDel(req.Elems) { + c.calls++ + if c.failuresRemaining > 0 { + c.failuresRemaining-- + return nil, kv.ErrRouteWriteFenced + } + } + return c.localAdapterCoordinator.Dispatch(ctx, req) +} + +func operationGroupHasDel(elems []*kv.Elem[kv.OP]) bool { + for _, elem := range elems { + if elem != nil && elem.Op == kv.Del { + return true + } + } + return false +} diff --git a/adapter/sqs_messages.go b/adapter/sqs_messages.go index 516d53d87..a75e0c96d 100644 --- a/adapter/sqs_messages.go +++ b/adapter/sqs_messages.go @@ -1292,7 +1292,7 @@ func (s *SQSServer) expireMessage(ctx context.Context, queueName string, meta *s Elems: elems, } if _, err := s.coordinator.Dispatch(ctx, req); err != nil { - if isRetryableTransactWriteError(err) { + if isIgnorableTransactRaceError(err) { return nil } return errors.WithStack(err) @@ -1340,8 +1340,9 @@ func (s *SQSServer) rotateMessagesForDelivery( // - (msg, false, nil) → delivered, caller appends. // - (nil, true, nil) → expected race; skip this candidate only. // Covers ErrKeyNotFound (someone deleted the record between the -// vis-index scan and our GetAt) and ErrWriteConflict on dispatch -// (another receive rotated the same record). +// vis-index scan and our GetAt), plus non-fence dispatch races +// like ErrWriteConflict / ErrTxnLocked (another receive rotated +// the same record). // - (nil, false, err) → non-retryable failure; propagate up the // stack so ReceiveMessage returns an actionable 5xx instead of // a false-empty 200. @@ -1441,7 +1442,7 @@ func (s *SQSServer) commitReceiveRotation(ctx context.Context, queueName string, return nil, false, err } if _, err := s.coordinator.Dispatch(ctx, req); err != nil { - if isRetryableTransactWriteError(err) { + if isIgnorableTransactRaceError(err) { return nil, true, nil } return nil, false, errors.WithStack(err) diff --git a/adapter/sqs_reaper.go b/adapter/sqs_reaper.go index ae7f6c20d..a25b83ca6 100644 --- a/adapter/sqs_reaper.go +++ b/adapter/sqs_reaper.go @@ -664,7 +664,7 @@ func (s *SQSServer) reapOneRecord(ctx context.Context, queueName string, meta *s return err } if _, err := s.coordinator.Dispatch(ctx, req); err != nil { - if isRetryableTransactWriteError(err) { + if isIgnorableTransactRaceError(err) { return nil } return errors.WithStack(err) @@ -873,7 +873,7 @@ func (s *SQSServer) dispatchDedupDelete(ctx context.Context, key []byte, readTS }, } if _, err := s.coordinator.Dispatch(ctx, req); err != nil { - if isRetryableTransactWriteError(err) { + if isIgnorableTransactRaceError(err) { return nil } return errors.WithStack(err) diff --git a/adapter/sqs_receive_route_fence_test.go b/adapter/sqs_receive_route_fence_test.go new file mode 100644 index 000000000..9c584f791 --- /dev/null +++ b/adapter/sqs_receive_route_fence_test.go @@ -0,0 +1,86 @@ +package adapter + +import ( + "context" + "testing" + + "github.com/bootjp/elastickv/kv" + "github.com/bootjp/elastickv/store" + "github.com/stretchr/testify/require" +) + +type receiveRotationErrorCoordinator struct { + stubAdapterCoordinator + err error +} + +func (c *receiveRotationErrorCoordinator) Dispatch(context.Context, *kv.OperationGroup[kv.OP]) (*kv.CoordinateResponse, error) { + return nil, c.err +} + +func TestSQSReceiveRotationClassifiesDispatchErrors(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + err error + wantSkip bool + wantErr error + }{ + {name: "write conflict", err: store.ErrWriteConflict, wantSkip: true}, + {name: "txn locked", err: kv.ErrTxnLocked, wantSkip: true}, + {name: "route write fenced", err: kv.ErrRouteWriteFenced, wantErr: kv.ErrRouteWriteFenced}, + } + + for _, tt := range tests { + + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + const queueName = "receive-route-fence" + const messageID = "msg-1" + const gen = uint64(1) + const visibleAt = int64(1000) + + meta := &sqsQueueMeta{Name: queueName, Generation: gen, PartitionCount: 1} + rec := &sqsMessageRecord{ + MessageID: messageID, + Body: []byte("body"), + MD5OfBody: sqsMD5Hex([]byte("body")), + SendTimestampMillis: visibleAt, + VisibleAtMillis: visibleAt, + QueueGeneration: gen, + } + cand := sqsMsgCandidate{ + visKey: sqsMsgVisKeyDispatch(meta, queueName, 0, gen, visibleAt, messageID), + messageID: messageID, + partition: 0, + } + dataKey := sqsMsgDataKeyDispatch(meta, queueName, 0, gen, messageID) + srv := &SQSServer{ + coordinator: &receiveRotationErrorCoordinator{err: tt.err}, + } + + msg, skip, err := srv.commitReceiveRotation( + context.Background(), + queueName, + meta, + cand, + dataKey, + rec, + uint64(visibleAt), + sqsReceiveOptions{VisibilityTimeout: 30}, + nil, + fifoLockAcquire, + ) + + require.Nil(t, msg) + require.Equal(t, tt.wantSkip, skip) + if tt.wantErr != nil { + require.ErrorIs(t, err, tt.wantErr) + } else { + require.NoError(t, err) + } + }) + } +} diff --git a/adapter/sqs_redrive.go b/adapter/sqs_redrive.go index 3fb13b8d3..cfc4437d0 100644 --- a/adapter/sqs_redrive.go +++ b/adapter/sqs_redrive.go @@ -237,7 +237,7 @@ func (s *SQSServer) redriveCandidateToDLQ( return false, err } if _, err := s.coordinator.Dispatch(ctx, req); err != nil { - if isRetryableTransactWriteError(err) { + if isIgnorableTransactRaceError(err) { return true, nil } return false, errors.WithStack(err) diff --git a/cmd/redis-proxy/main.go b/cmd/redis-proxy/main.go index f7f2358a7..cd353f64b 100644 --- a/cmd/redis-proxy/main.go +++ b/cmd/redis-proxy/main.go @@ -7,6 +7,7 @@ import ( "log/slog" "net" "net/http" + "net/http/pprof" "os" "os/signal" "strings" @@ -22,6 +23,8 @@ const ( sentryFlushTimeout = 2 * time.Second metricsShutdownTimeout = 5 * time.Second secondaryConcurrencyDivisor = 2 + elasticKVDispatchTimeout = 10 * time.Second + backendTimeoutGrace = time.Second ) func main() { @@ -65,6 +68,8 @@ func run() error { flag.StringVar(&cfg.SentryEnv, "sentry-env", cfg.SentryEnv, "Sentry environment") flag.Float64Var(&cfg.SentrySampleRate, "sentry-sample", cfg.SentrySampleRate, "Sentry sample rate") flag.StringVar(&cfg.MetricsAddr, "metrics", cfg.MetricsAddr, "Prometheus metrics address") + flag.StringVar(&cfg.PProfAddr, "pprof", cfg.PProfAddr, "pprof listen address (empty = disabled)") + flag.BoolVar(&cfg.RedisOnlyRaw, "redis-only-raw", cfg.RedisOnlyRaw, "Use raw TCP bridging in redis-only mode") flag.Parse() mode, resolvedWriteConcurrency, resolvedScriptConcurrency, resolvedBlockingReplayConcurrency, err := resolveRuntimeOptions( @@ -95,11 +100,38 @@ func run() error { sentryReporter := proxy.NewSentryReporter(cfg.SentryDSN, cfg.SentryEnv, cfg.SentrySampleRate, logger) defer sentryReporter.Flush(sentryFlushTimeout) - // Prometheus reg := prometheus.NewRegistry() metrics := proxy.NewProxyMetrics(reg) - // Backends + primary, secondary, err := newBackends(cfg, primaryPoolSize, elasticKVPoolSize, logger) + if err != nil { + return err + } + defer primary.Close() + defer secondary.Close() + + dual := proxy.NewDualWriter(primary, secondary, cfg, metrics, sentryReporter, logger) + defer dual.Close() // wait for in-flight async goroutines + srv := proxy.NewProxyServer(cfg, dual, metrics, sentryReporter, logger) + + // Context for graceful shutdown + ctx, cancel := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) + defer cancel() + + go serveMetrics(ctx, cfg.MetricsAddr, reg, logger) + + if cfg.PProfAddr != "" { + go servePProf(ctx, cfg.PProfAddr, logger) + } + + // Start proxy + if err := srv.ListenAndServe(ctx); err != nil { + return fmt.Errorf("proxy server: %w", err) + } + return nil +} + +func newBackends(cfg proxy.ProxyConfig, primaryPoolSize, elasticKVPoolSize int, logger *slog.Logger) (proxy.Backend, proxy.Backend, error) { primaryOpts := proxy.DefaultBackendOptions() primaryOpts.DB = cfg.PrimaryDB primaryOpts.Password = cfg.PrimaryPassword @@ -108,62 +140,90 @@ func run() error { secondaryOpts.DB = cfg.SecondaryDB secondaryOpts.Password = cfg.SecondaryPassword secondaryOpts.PoolSize = elasticKVPoolSize + alignElasticKVBackendTimeouts(&secondaryOpts, cfg.SecondaryTimeout) secondarySeeds := parseAddrList(cfg.SecondaryAddr) - if len(secondarySeeds) == 0 { - return fmt.Errorf("at least one secondary address is required") - } - var primary, secondary proxy.Backend switch cfg.Mode { - case proxy.ModeElasticKVPrimary, proxy.ModeElasticKVOnly: - primary = proxy.NewLeaderAwareRedisBackend(secondarySeeds, "elastickv", secondaryOpts, logger) - secondary = proxy.NewRedisBackendWithOptions(cfg.PrimaryAddr, "redis", primaryOpts) - case proxy.ModeRedisOnly, proxy.ModeDualWrite, proxy.ModeDualWriteShadow: - primary = proxy.NewRedisBackendWithOptions(cfg.PrimaryAddr, "redis", primaryOpts) - secondary = proxy.NewLeaderAwareRedisBackend(secondarySeeds, "elastickv", secondaryOpts, logger) + case proxy.ModeElasticKVPrimary: + if len(secondarySeeds) == 0 { + return nil, nil, fmt.Errorf("at least one secondary address is required") + } + return proxy.NewLeaderAwareRedisBackend(secondarySeeds, "elastickv", secondaryOpts, logger), + proxy.NewRedisBackendWithOptions(cfg.PrimaryAddr, "redis", primaryOpts), nil + case proxy.ModeElasticKVOnly: + if len(secondarySeeds) == 0 { + return nil, nil, fmt.Errorf("at least one secondary address is required") + } + return proxy.NewLeaderAwareRedisBackend(secondarySeeds, "elastickv", secondaryOpts, logger), + proxy.NewNoopBackend("redis"), nil + case proxy.ModeRedisOnly: + return proxy.NewRedisBackendWithOptions(cfg.PrimaryAddr, "redis", primaryOpts), + proxy.NewNoopBackend("elastickv"), nil + case proxy.ModeDualWrite, proxy.ModeDualWriteShadow: + if len(secondarySeeds) == 0 { + return nil, nil, fmt.Errorf("at least one secondary address is required") + } + return proxy.NewRedisBackendWithOptions(cfg.PrimaryAddr, "redis", primaryOpts), + proxy.NewLeaderAwareRedisBackend(secondarySeeds, "elastickv", secondaryOpts, logger), nil + default: + return nil, nil, fmt.Errorf("unsupported mode: %s", cfg.Mode.String()) } - defer primary.Close() - defer secondary.Close() +} - dual := proxy.NewDualWriter(primary, secondary, cfg, metrics, sentryReporter, logger) - defer dual.Close() // wait for in-flight async goroutines - srv := proxy.NewProxyServer(cfg, dual, metrics, sentryReporter, logger) +func alignElasticKVBackendTimeouts(opts *proxy.BackendOptions, operationTimeout time.Duration) { + if opts == nil { + return + } + floor := elasticKVDispatchTimeout + if operationTimeout > floor { + floor = operationTimeout + } + floor += backendTimeoutGrace + if opts.ReadTimeout > 0 && opts.ReadTimeout < floor { + opts.ReadTimeout = floor + } + if opts.WriteTimeout > 0 && opts.WriteTimeout < floor { + opts.WriteTimeout = floor + } +} - // Context for graceful shutdown - ctx, cancel := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) - defer cancel() +func serveMetrics(ctx context.Context, addr string, reg *prometheus.Registry, logger *slog.Logger) { + mux := http.NewServeMux() + mux.Handle("/metrics", promhttp.HandlerFor(reg, promhttp.HandlerOpts{})) + serveHTTP(ctx, addr, mux, "metrics", logger) +} - // Start metrics server +func servePProf(ctx context.Context, addr string, logger *slog.Logger) { + mux := http.NewServeMux() + mux.HandleFunc("/debug/pprof/", pprof.Index) + mux.HandleFunc("/debug/pprof/cmdline", pprof.Cmdline) + mux.HandleFunc("/debug/pprof/profile", pprof.Profile) + mux.HandleFunc("/debug/pprof/symbol", pprof.Symbol) + mux.HandleFunc("/debug/pprof/trace", pprof.Trace) + serveHTTP(ctx, addr, mux, "pprof", logger) +} + +func serveHTTP(ctx context.Context, addr string, handler http.Handler, name string, logger *slog.Logger) { + var lc net.ListenConfig + ln, err := lc.Listen(ctx, "tcp", addr) + if err != nil { + logger.Error(name+" listen failed", "addr", addr, "err", err) + return + } + srv := &http.Server{Handler: handler, ReadHeaderTimeout: time.Second} go func() { - mux := http.NewServeMux() - mux.Handle("/metrics", promhttp.HandlerFor(reg, promhttp.HandlerOpts{})) - var lc net.ListenConfig - ln, err := lc.Listen(ctx, "tcp", cfg.MetricsAddr) - if err != nil { - logger.Error("metrics listen failed", "addr", cfg.MetricsAddr, "err", err) - return - } - metricsSrv := &http.Server{Handler: mux, ReadHeaderTimeout: time.Second} - go func() { - <-ctx.Done() - shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), metricsShutdownTimeout) - defer shutdownCancel() - if err := metricsSrv.Shutdown(shutdownCtx); err != nil { - logger.Warn("metrics server shutdown error", "err", err) - } - }() - logger.Info("metrics server starting", "addr", cfg.MetricsAddr) - if err := metricsSrv.Serve(ln); err != nil && err != http.ErrServerClosed { - logger.Error("metrics server error", "err", err) + <-ctx.Done() + shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), metricsShutdownTimeout) + defer shutdownCancel() + if err := srv.Shutdown(shutdownCtx); err != nil { + logger.Warn(name+" server shutdown error", "err", err) } }() - - // Start proxy - if err := srv.ListenAndServe(ctx); err != nil { - return fmt.Errorf("proxy server: %w", err) + logger.Info(name+" server starting", "addr", addr) + if err := srv.Serve(ln); err != nil && err != http.ErrServerClosed { + logger.Error(name+" server error", "err", err) } - return nil } func resolveRuntimeOptions( diff --git a/cmd/redis-proxy/main_test.go b/cmd/redis-proxy/main_test.go index 7188637f9..1730e9ed5 100644 --- a/cmd/redis-proxy/main_test.go +++ b/cmd/redis-proxy/main_test.go @@ -2,6 +2,7 @@ package main import ( "testing" + "time" "github.com/bootjp/elastickv/proxy" "github.com/stretchr/testify/assert" @@ -55,6 +56,33 @@ func TestValidateSecondaryConcurrency(t *testing.T) { assert.Contains(t, err.Error(), "secondary-blocking-replay-concurrency") } +func TestNewBackendsAllowsRedisOnlyWithoutSecondarySeeds(t *testing.T) { + cfg := proxy.DefaultConfig() + cfg.Mode = proxy.ModeRedisOnly + cfg.PrimaryAddr = "127.0.0.1:6379" + cfg.SecondaryAddr = "" + + primary, secondary, err := newBackends(cfg, 2, 2, nil) + require.NoError(t, err) + t.Cleanup(func() { + require.NoError(t, primary.Close()) + require.NoError(t, secondary.Close()) + }) + assert.Equal(t, "redis", primary.Name()) + assert.Equal(t, "elastickv", secondary.Name()) +} + +func TestNewBackendsRejectsDualWriteWithoutSecondarySeeds(t *testing.T) { + cfg := proxy.DefaultConfig() + cfg.Mode = proxy.ModeDualWrite + cfg.PrimaryAddr = "127.0.0.1:6379" + cfg.SecondaryAddr = "" + + _, _, err := newBackends(cfg, 2, 2, nil) + require.Error(t, err) + assert.Contains(t, err.Error(), "secondary address") +} + func TestDeriveSecondaryConcurrency(t *testing.T) { tests := []struct { name string @@ -154,3 +182,34 @@ func TestDeriveSecondaryConcurrency(t *testing.T) { }) } } + +func TestAlignElasticKVBackendTimeouts(t *testing.T) { + t.Run("uses ElasticKV dispatch floor by default", func(t *testing.T) { + opts := proxy.DefaultElasticKVBackendOptions() + + alignElasticKVBackendTimeouts(&opts, 5*time.Second) + + assert.Equal(t, 11*time.Second, opts.ReadTimeout) + assert.Equal(t, 11*time.Second, opts.WriteTimeout) + }) + + t.Run("follows longer secondary timeout", func(t *testing.T) { + opts := proxy.DefaultElasticKVBackendOptions() + + alignElasticKVBackendTimeouts(&opts, 15*time.Second) + + assert.Equal(t, 16*time.Second, opts.ReadTimeout) + assert.Equal(t, 16*time.Second, opts.WriteTimeout) + }) + + t.Run("keeps explicit larger timeout", func(t *testing.T) { + opts := proxy.DefaultElasticKVBackendOptions() + opts.ReadTimeout = 30 * time.Second + opts.WriteTimeout = 31 * time.Second + + alignElasticKVBackendTimeouts(&opts, 15*time.Second) + + assert.Equal(t, 30*time.Second, opts.ReadTimeout) + assert.Equal(t, 31*time.Second, opts.WriteTimeout) + }) +} diff --git a/distribution/catalog.go b/distribution/catalog.go index 9447a018a..4d51093a7 100644 --- a/distribution/catalog.go +++ b/distribution/catalog.go @@ -206,11 +206,8 @@ func encodeRouteDescriptorWithSplitAtHLCOffset(route RouteDescriptor) ([]byte, u } out := make([]byte, 0, routeDescriptorEncodedSize(route)) - version := catalogRouteCodecVersionV1 - if route.SplitAtHLC != 0 { - version = catalogRouteCodecVersionV2 - } - if routeDescriptorRequiresV3(route) { + version := catalogRouteCodecVersionV2 + if routeDescriptorRequiresV2(route) { version = catalogRouteCodecVersionV3 } out = append(out, version) @@ -229,13 +226,10 @@ func encodeRouteDescriptorWithSplitAtHLCOffset(route RouteDescriptor) ([]byte, u out = append(out, route.End...) } - splitAtHLCOffset := uint64(0) - switch version { - case catalogRouteCodecVersionV3: - splitAtHLCOffset = uint64(len(out)) + splitAtHLCOffset := uint64(len(out)) + if version == catalogRouteCodecVersionV3 { out = appendRouteDescriptorV3Tail(out, route) - case catalogRouteCodecVersionV2: - splitAtHLCOffset = uint64(len(out)) + } else { out = appendU64(out, route.SplitAtHLC) } return out, splitAtHLCOffset, nil @@ -832,19 +826,14 @@ func routeDescriptorEncodedSize(route RouteDescriptor) int { if route.End != nil { size += catalogUint64Bytes + len(route.End) } - if routeDescriptorRequiresV3(route) { - size += catalogRouteV3TailSize - } else if route.SplitAtHLC != 0 { - size += catalogUint64Bytes + size += catalogUint64Bytes + if routeDescriptorRequiresV2(route) { + size += 1 + catalogUint64Bytes + catalogUint64Bytes } return size } func routeDescriptorRequiresV2(route RouteDescriptor) bool { - return routeDescriptorRequiresV3(route) -} - -func routeDescriptorRequiresV3(route RouteDescriptor) bool { return route.StagedVisibilityActive || route.MigrationJobID != 0 || route.MinWriteTSExclusive != 0 } diff --git a/distribution/catalog_test.go b/distribution/catalog_test.go index 162a2edf2..d7e817a7d 100644 --- a/distribution/catalog_test.go +++ b/distribution/catalog_test.go @@ -66,7 +66,7 @@ func TestRouteDescriptorCodecRoundTrip(t *testing.T) { t.Fatalf("encode route: %v", err) } if raw[0] != catalogRouteCodecVersionV2 { - t.Fatalf("split route encoded version = %d, want v2", raw[0]) + t.Fatalf("zero-M2 route encoded version = %d, want v2", raw[0]) } got, err := DecodeRouteDescriptor(raw) if err != nil { @@ -90,7 +90,7 @@ func TestRouteDescriptorCodecRoundTripNilEnd(t *testing.T) { t.Fatalf("encode route: %v", err) } if raw[0] != catalogRouteCodecVersionV2 { - t.Fatalf("split nil-end route encoded version = %d, want v2", raw[0]) + t.Fatalf("zero-M2 nil-end route encoded version = %d, want v2", raw[0]) } got, err := DecodeRouteDescriptor(raw) if err != nil { @@ -912,8 +912,16 @@ func TestCatalogStoreApplySaveMutations_UsesMonotonicCommitTS(t *testing.T) { func assertRouteEqual(t *testing.T, want, got RouteDescriptor) { t.Helper() assertRouteIdentityEqual(t, want, got) - assertRouteMetadataEqual(t, want, got) - assertRouteBoundsEqual(t, want, got) + assertRouteMigrationEqual(t, want, got) + if want.State != got.State { + t.Fatalf("state mismatch: want %d, got %d", want.State, got.State) + } + if !bytes.Equal(want.Start, got.Start) { + t.Fatalf("start mismatch: want %q, got %q", want.Start, got.Start) + } + if !bytes.Equal(want.End, got.End) { + t.Fatalf("end mismatch: want %q, got %q", want.End, got.End) + } } func assertRouteIdentityEqual(t *testing.T, want, got RouteDescriptor) { @@ -927,16 +935,10 @@ func assertRouteIdentityEqual(t *testing.T, want, got RouteDescriptor) { if want.ParentRouteID != got.ParentRouteID { t.Fatalf("parent route id mismatch: want %d, got %d", want.ParentRouteID, got.ParentRouteID) } - if want.State != got.State { - t.Fatalf("state mismatch: want %d, got %d", want.State, got.State) - } } -func assertRouteMetadataEqual(t *testing.T, want, got RouteDescriptor) { +func assertRouteMigrationEqual(t *testing.T, want, got RouteDescriptor) { t.Helper() - if want.SplitAtHLC != got.SplitAtHLC { - t.Fatalf("split at HLC mismatch: want %d, got %d", want.SplitAtHLC, got.SplitAtHLC) - } if want.StagedVisibilityActive != got.StagedVisibilityActive { t.Fatalf("staged visibility mismatch: want %v, got %v", want.StagedVisibilityActive, got.StagedVisibilityActive) } @@ -946,15 +948,8 @@ func assertRouteMetadataEqual(t *testing.T, want, got RouteDescriptor) { if want.MinWriteTSExclusive != got.MinWriteTSExclusive { t.Fatalf("min write ts mismatch: want %d, got %d", want.MinWriteTSExclusive, got.MinWriteTSExclusive) } -} - -func assertRouteBoundsEqual(t *testing.T, want, got RouteDescriptor) { - t.Helper() - if !bytes.Equal(want.Start, got.Start) { - t.Fatalf("start mismatch: want %q, got %q", want.Start, got.Start) - } - if !bytes.Equal(want.End, got.End) { - t.Fatalf("end mismatch: want %q, got %q", want.End, got.End) + if want.SplitAtHLC != got.SplitAtHLC { + t.Fatalf("split at HLC mismatch: want %d, got %d", want.SplitAtHLC, got.SplitAtHLC) } } diff --git a/distribution/engine.go b/distribution/engine.go index 726c83b16..484854a54 100644 --- a/distribution/engine.go +++ b/distribution/engine.go @@ -266,6 +266,36 @@ func (s RouteHistorySnapshot) OwnerOf(key []byte) (uint64, bool) { return 0, false } +// RouteOf returns the route that covered key at this snapshot's version. +func (s RouteHistorySnapshot) RouteOf(key []byte) (Route, bool) { + for _, r := range s.routes { + if bytes.Compare(key, r.Start) < 0 { + break + } + if r.End != nil && bytes.Compare(key, r.End) >= 0 { + continue + } + return cloneRoute(r), true + } + return Route{}, false +} + +// IntersectingRoutes returns every route whose range intersects [start, end) +// in this snapshot. A nil end denotes +infinity. +func (s RouteHistorySnapshot) IntersectingRoutes(start, end []byte) []Route { + out := make([]Route, 0) + for _, r := range s.routes { + if r.End != nil && bytes.Compare(r.End, start) <= 0 { + continue + } + if end != nil && bytes.Compare(r.Start, end) >= 0 { + break + } + out = append(out, cloneRoute(r)) + } + return out +} + // Current returns the route catalog snapshot at the engine's current // catalogVersion. Returns (zero, false) when the history ring has // not been initialised (bare-struct Engine). Used by the M3 @@ -486,7 +516,7 @@ func (e *Engine) GetIntersectingRoutesWithVersion(start, end []byte) ([]Route, u } // Route starts at or after scan ends: end != nil && rStart >= end if end != nil && bytes.Compare(r.Start, end) >= 0 { - continue + break } // Route intersects with scan range result = append(result, Route{ @@ -504,6 +534,20 @@ func (e *Engine) GetIntersectingRoutesWithVersion(start, end []byte) ([]Route, u return result, e.catalogVersion } +func cloneRoute(r Route) Route { + return Route{ + RouteID: r.RouteID, + Start: CloneBytes(r.Start), + End: CloneBytes(r.End), + GroupID: r.GroupID, + State: r.State, + StagedVisibilityActive: r.StagedVisibilityActive, + MigrationJobID: r.MigrationJobID, + MinWriteTSExclusive: r.MinWriteTSExclusive, + Load: r.Load, + } +} + func (e *Engine) routeIndex(key []byte) int { if len(e.routes) == 0 { return -1 diff --git a/distribution/engine_test.go b/distribution/engine_test.go index a54ef5863..67a7c0053 100644 --- a/distribution/engine_test.go +++ b/distribution/engine_test.go @@ -126,69 +126,21 @@ func TestEngineApplySnapshot_PreservesMigrationRouteFields(t *testing.T) { t.Fatalf("expected 1 intersecting route, got %d", len(intersections)) } requireMigrationRouteFields(t, "GetIntersectingRoutes", intersections[0]) -} - -func TestEngineApplyDelta_PreservesMigrationRouteFields(t *testing.T) { - t.Parallel() - - e := NewEngine() - if err := e.ApplySnapshot(CatalogSnapshot{ - Version: 1, - Routes: []RouteDescriptor{ - { - RouteID: 1, - Start: []byte(""), - End: nil, - GroupID: 1, - State: RouteStateActive, - }, - }, - }); err != nil { - t.Fatalf("ApplySnapshot: %v", err) - } - - err := e.ApplyDelta(CatalogDelta{ - PreviousVersion: 1, - Version: 2, - Mutations: []CatalogRouteMutation{ - {Op: CatalogMutationDelete, RouteID: 1}, - { - Op: CatalogMutationUpsert, - RouteID: 7, - Route: RouteDescriptor{ - RouteID: 7, - Start: []byte("a"), - End: []byte("z"), - GroupID: 2, - State: RouteStateMigratingTarget, - StagedVisibilityActive: true, - MigrationJobID: 42, - MinWriteTSExclusive: 99, - }, - }, - }, - }) - if err != nil { - t.Fatalf("ApplyDelta: %v", err) - } - route, ok := e.GetRoute([]byte("m")) + snapshot, ok := e.Current() if !ok { - t.Fatal("expected route") + t.Fatal("expected current history snapshot") } - requireMigrationRouteFields(t, "GetRoute", route) - - stats := e.Stats() - if len(stats) != 1 { - t.Fatalf("expected 1 stat route, got %d", len(stats)) + historyRoute, ok := snapshot.RouteOf([]byte("m")) + if !ok { + t.Fatal("expected history route") } - requireMigrationRouteFields(t, "Stats", stats[0]) - - intersections := e.GetIntersectingRoutes([]byte("b"), []byte("c")) - if len(intersections) != 1 { - t.Fatalf("expected 1 intersecting route, got %d", len(intersections)) + requireMigrationRouteFields(t, "RouteHistorySnapshot.RouteOf", historyRoute) + historyIntersections := snapshot.IntersectingRoutes([]byte("b"), []byte("c")) + if len(historyIntersections) != 1 { + t.Fatalf("expected 1 history intersecting route, got %d", len(historyIntersections)) } - requireMigrationRouteFields(t, "GetIntersectingRoutes", intersections[0]) + requireMigrationRouteFields(t, "RouteHistorySnapshot.IntersectingRoutes", historyIntersections[0]) } func requireMigrationRouteFields(t *testing.T, label string, route Route) { diff --git a/distribution/migrator.go b/distribution/migrator.go new file mode 100644 index 000000000..db46b6cfe --- /dev/null +++ b/distribution/migrator.go @@ -0,0 +1,546 @@ +package distribution + +import ( + "bytes" + + "github.com/bootjp/elastickv/internal/fskeys" + "github.com/bootjp/elastickv/internal/s3keys" + "github.com/bootjp/elastickv/store" + "github.com/cockroachdb/errors" +) + +const ( + MigrationFamilyUser uint32 = iota + 1 + MigrationFamilyTxnIntent + MigrationFamilyTxnCommit + MigrationFamilyTxnRollback + MigrationFamilyTxnSuccess + MigrationFamilyTxnMeta + MigrationFamilyTxnLock + MigrationFamilyListMeta + MigrationFamilyListItem + MigrationFamilyListMetaDelta + MigrationFamilyListClaim + MigrationFamilyRedisLegacy + MigrationFamilyHash + MigrationFamilySet + MigrationFamilyZSet + MigrationFamilyStreamMeta + MigrationFamilyStreamEntry + MigrationFamilyDynamoTableMeta + MigrationFamilyDynamoTableGeneration + MigrationFamilyDynamoItem + MigrationFamilyDynamoGSI + MigrationFamilySQSQueueMeta + MigrationFamilySQSQueueGeneration + MigrationFamilySQSQueueSequence + MigrationFamilySQSQueueTombstone + MigrationFamilySQSMessageData + MigrationFamilySQSMessageVisibility + MigrationFamilySQSMessageDedup + MigrationFamilySQSMessageGroup + MigrationFamilySQSMessageByAge + MigrationFamilySQSPartitionedMessageData + MigrationFamilySQSPartitionedMessageVisibility + MigrationFamilySQSPartitionedMessageDedup + MigrationFamilySQSPartitionedMessageGroup + MigrationFamilySQSPartitionedMessageByAge + MigrationFamilyS3BucketMeta + MigrationFamilyS3BucketGeneration + MigrationFamilyS3ObjectManifest + MigrationFamilyS3UploadMeta + MigrationFamilyS3UploadPart + MigrationFamilyS3Blob + MigrationFamilyS3GCUpload + MigrationFamilyLegacyListMetaDelta + MigrationFamilyS3ChunkRef + MigrationFamilyFilesystemChunk +) + +const ( + migrationTxnIntentPrefix = "!txn|int|" + migrationTxnCommitPrefix = "!txn|cmt|" + migrationTxnRollbackPrefix = "!txn|rb|" + migrationTxnSuccessPrefix = "!txn|ok|" + migrationTxnMetaPrefix = "!txn|meta|" + migrationTxnLockPrefix = "!txn|lock|" + migrationRedisPrefix = "!redis|" + migrationHashPrefix = "!hs|" + migrationSetPrefix = "!st|" + migrationZSetPrefix = "!zs|" + migrationDynamoMetaPrefix = "!ddb|meta|table|" + migrationDynamoGenPrefix = "!ddb|meta|gen|" + migrationDynamoItemPrefix = "!ddb|item|" + migrationDynamoGSIPrefix = "!ddb|gsi|" +) + +const ( + migrationSQSQueueMetaPrefix = "!sqs|queue|meta|" + migrationSQSQueueGenPrefix = "!sqs|queue|gen|" + migrationSQSQueueSeqPrefix = "!sqs|queue|seq|" + migrationSQSQueueTombstonePrefix = "!sqs|queue|tombstone|" + migrationSQSMsgDataPrefix = "!sqs|msg|data|" + migrationSQSMsgVisPrefix = "!sqs|msg|vis|" + migrationSQSMsgDedupPrefix = "!sqs|msg|dedup|" + migrationSQSMsgGroupPrefix = "!sqs|msg|group|" + migrationSQSMsgByAgePrefix = "!sqs|msg|byage|" + migrationSQSPartitionedSuffix = "p|" +) + +var ( + ErrMigrationReservedRange = errors.New("migration range intersects reserved control prefix") + ErrMigrationInvalidRoute = errors.New("migration route is invalid") + ErrMigrationDataMoveRequired = errors.New("migration data move is not implemented") + ErrMigrationSourceRouteChanged = errors.New("migration source route does not match split job") +) + +var migrationReservedControlPrefixes = [][]byte{ + []byte("!dist|"), + []byte("!migstage|"), + []byte("!migwrite|"), + []byte("!migfence|"), +} + +var migrationInternalFamilyPrefixes = [][]byte{ + []byte(migrationTxnLockPrefix), + []byte(migrationTxnIntentPrefix), + []byte(migrationTxnCommitPrefix), + []byte(migrationTxnRollbackPrefix), + []byte(migrationTxnSuccessPrefix), + []byte(migrationTxnMetaPrefix), + []byte(store.ListMetaDeltaPrefix), + []byte(store.LegacyListMetaDeltaPrefix), + []byte(store.ListClaimPrefix), + []byte(store.ListMetaPrefix), + []byte(store.ListItemPrefix), + []byte(migrationRedisPrefix), + []byte(migrationHashPrefix), + []byte(migrationSetPrefix), + []byte(migrationZSetPrefix), + []byte(store.StreamMetaPrefix), + []byte(store.StreamEntryPrefix), + []byte(migrationDynamoMetaPrefix), + []byte(migrationDynamoGenPrefix), + []byte(migrationDynamoItemPrefix), + []byte(migrationDynamoGSIPrefix), + []byte(migrationSQSQueueMetaPrefix), + []byte(migrationSQSQueueGenPrefix), + []byte(migrationSQSQueueSeqPrefix), + []byte(migrationSQSQueueTombstonePrefix), + []byte(migrationSQSMsgDataPrefix), + []byte(migrationSQSMsgVisPrefix), + []byte(migrationSQSMsgDedupPrefix), + []byte(migrationSQSMsgGroupPrefix), + []byte(migrationSQSMsgByAgePrefix), + []byte(s3keys.BucketMetaPrefix), + []byte(s3keys.BucketGenerationPrefix), + []byte(s3keys.ObjectManifestPrefix), + []byte(s3keys.UploadMetaPrefix), + []byte(s3keys.UploadPartPrefix), + []byte(s3keys.BlobPrefix), + []byte(s3keys.ChunkRefPrefix), + []byte(s3keys.GCUploadPrefix), + fskeys.ChunkAllPrefix(), +} + +// MigrationBracket is a raw MVCC export or drain slice used by the migrator. +type MigrationBracket struct { + BracketID uint64 + Family uint32 + Start []byte + End []byte + ExcludePrefixes [][]byte + ExcludeKnownInternal bool + DrainOnly bool + RequiresRouteKeyCheck bool + RequiresDecodedS3 bool +} + +// PlanMigrationBrackets returns the full M2 migration plan, including the +// drain-only transaction lock bracket. Data export callers should use +// PlanExportBrackets, which omits drain-only control state. +func PlanMigrationBrackets(routeStart, routeEnd []byte) ([]MigrationBracket, error) { + routeEnd = normalizeMigrationRouteEnd(routeEnd) + if err := ValidateMigrationRouteRange(routeStart, routeEnd); err != nil { + return nil, err + } + + brackets := []MigrationBracket{ + { + BracketID: uint64(MigrationFamilyUser), + Family: MigrationFamilyUser, + Start: CloneBytes(routeStart), + End: CloneBytes(routeEnd), + ExcludeKnownInternal: true, + RequiresRouteKeyCheck: true, + }, + } + brackets = append(brackets, migrationFamilyBrackets()...) + return brackets, nil +} + +// PlanExportBrackets returns the data-copy bracket plan. Intent locks are +// deliberately absent because the source drains them route-faithfully before +// cutover and the target must not materialize in-flight intents as data. +func PlanExportBrackets(routeStart, routeEnd []byte) ([]MigrationBracket, error) { + brackets, err := PlanMigrationBrackets(routeStart, routeEnd) + if err != nil { + return nil, err + } + out := make([]MigrationBracket, 0, len(brackets)) + for _, bracket := range brackets { + if bracket.DrainOnly { + continue + } + out = append(out, bracket) + } + return out, nil +} + +// SplitJobBracketProgressForPlan creates durable per-bracket resume state for +// a SplitJob phase. +func SplitJobBracketProgressForPlan(brackets []MigrationBracket, phase SplitJobExportPhase) []SplitJobBracketProgress { + out := make([]SplitJobBracketProgress, 0, len(brackets)) + for _, bracket := range brackets { + if bracket.DrainOnly { + continue + } + out = append(out, SplitJobBracketProgress{ + BracketID: bracket.BracketID, + Family: bracket.Family, + ExportPhase: phase, + }) + } + return out +} + +// ContainsRawKey reports whether rawKey is inside the bracket's raw scan +// interval after applying bracket-local exclusions. Route ownership still +// requires the caller's RouteKeyFilter for every bracket. +func (b MigrationBracket) ContainsRawKey(rawKey []byte) bool { + if bytes.Compare(rawKey, b.Start) < 0 { + return false + } + if len(b.End) > 0 && bytes.Compare(rawKey, b.End) >= 0 { + return false + } + if !b.containsFamilyShape(rawKey) { + return false + } + if b.ExcludeKnownInternal && IsMigrationKnownInternalKey(rawKey) { + return false + } + return !hasAnyPrefix(rawKey, b.ExcludePrefixes) +} + +// ContainsRoutedKey applies both the bracket's raw family interval and its +// route ownership predicate. S3 bucket-level auxiliary rows do not encode an +// object route key, so they are matched by bucket route-prefix intersection. +func (b MigrationBracket) ContainsRoutedKey(rawKey, routeStart, routeEnd []byte, routeKey func([]byte) []byte) bool { + return b.ContainsRoutedVersion(rawKey, nil, routeStart, routeEnd, routeKey) +} + +// ContainsRoutedVersion is the value-aware variant of ContainsRoutedKey. It is +// needed for legacy list metadata because old delta keys overlap byte-for-byte +// with base metadata keys whose user key begins with "d|". +func (b MigrationBracket) ContainsRoutedVersion(rawKey, value, routeStart, routeEnd []byte, routeKey func([]byte) []byte) bool { + if !b.ContainsRawKey(rawKey) { + return false + } + routeEnd = normalizeMigrationRouteEnd(routeEnd) + if b.RequiresDecodedS3 { + return b.containsDecodedS3Route(rawKey, routeStart, routeEnd) + } + if b.Family == MigrationFamilyLegacyListMetaDelta { + return b.containsLegacyListMetaDeltaRoute(rawKey, value, routeStart, routeEnd) + } + if !b.RequiresRouteKeyCheck { + return true + } + if routeKey == nil { + return false + } + return routeKeyInRange(routeKey(rawKey), routeStart, routeEnd) +} + +func (b MigrationBracket) containsLegacyListMetaDeltaRoute(rawKey, value, routeStart, routeEnd []byte) bool { + if value == nil { + return routeKeyInRange(store.ExtractListUserKey(rawKey), routeStart, routeEnd) || + routeKeyInRange(store.ExtractLegacyListUserKeyFromDelta(rawKey), routeStart, routeEnd) + } + if value != nil && !store.IsListMetaDeltaValue(value) { + return routeKeyInRange(store.ExtractListUserKey(rawKey), routeStart, routeEnd) + } + return routeKeyInRange(store.ExtractLegacyListUserKeyFromDelta(rawKey), routeStart, routeEnd) +} + +func (b MigrationBracket) containsFamilyShape(rawKey []byte) bool { + switch b.Family { + case MigrationFamilyListMetaDelta: + return store.ExtractListUserKeyFromDelta(rawKey) != nil + case MigrationFamilyLegacyListMetaDelta: + return store.ExtractLegacyListUserKeyFromDelta(rawKey) != nil + case MigrationFamilyListMeta: + return store.ExtractListUserKeyFromDelta(rawKey) == nil && + store.ExtractLegacyListUserKeyFromDelta(rawKey) == nil + default: + return true + } +} + +func (b MigrationBracket) containsDecodedS3Route(rawKey, routeStart, routeEnd []byte) bool { + bucket, ok := b.decodedS3Bucket(rawKey) + if !ok { + return false + } + if routeKeyInRange(rawKey, routeStart, routeEnd) { + return true + } + bucketRouteStart := s3keys.RoutePrefixForBucketAnyGeneration(bucket) + return rangesIntersect(routeStart, routeEnd, bucketRouteStart, prefixScanEnd(bucketRouteStart)) +} + +func (b MigrationBracket) decodedS3Bucket(rawKey []byte) (string, bool) { + switch b.Family { + case MigrationFamilyS3BucketMeta: + return s3keys.ParseBucketMetaKey(rawKey) + case MigrationFamilyS3BucketGeneration: + return s3keys.ParseBucketGenerationKey(rawKey) + default: + return "", false + } +} + +func routeKeyInRange(routeKey, routeStart, routeEnd []byte) bool { + if routeKey == nil { + return false + } + if bytes.Compare(routeKey, routeStart) < 0 { + return false + } + return len(routeEnd) == 0 || bytes.Compare(routeKey, routeEnd) < 0 +} + +// InitializeSplitJobPlan validates the source route and seeds the job's +// bracket progress for the moving right child [SplitKey, source.End). +func InitializeSplitJobPlan(job SplitJob, source RouteDescriptor, nowMs int64) (SplitJob, error) { + if err := validateSplitJobSource(job, source); err != nil { + return SplitJob{}, err + } + routeStart := CloneBytes(job.SplitKey) + routeEnd := CloneBytes(normalizeMigrationRouteEnd(source.End)) + brackets, err := PlanExportBrackets(routeStart, routeEnd) + if err != nil { + return SplitJob{}, err + } + + out := CloneSplitJob(job) + if out.Phase == SplitJobPhaseNone { + out.Phase = SplitJobPhasePlanned + } + if len(out.BracketProgress) == 0 { + out.BracketProgress = SplitJobBracketProgressForPlan(brackets, SplitJobExportPhaseBackfill) + } + if out.StartedAtMs == 0 { + out.StartedAtMs = nowMs + } + out.UpdatedAtMs = nowMs + return out, nil +} + +// AdvanceSameGroupNoop completes the PR4 same-group path without attempting +// data movement. Cross-group data copy is added by later M2 PRs. +func AdvanceSameGroupNoop(job SplitJob, source RouteDescriptor, nowMs int64) (SplitJob, error) { + planned, err := InitializeSplitJobPlan(job, source, nowMs) + if err != nil { + return SplitJob{}, err + } + if planned.TargetGroupID != source.GroupID { + return SplitJob{}, errors.WithStack(ErrMigrationDataMoveRequired) + } + + planned.Phase = SplitJobPhaseDone + planned.TargetPromotionDone = true + planned.PromotionCompletedTS = migrationWallMillisToUint64(nowMs) + planned.UpdatedAtMs = nowMs + planned.TerminalAtMs = nowMs + for i := range planned.BracketProgress { + planned.BracketProgress[i].Done = true + } + return planned, nil +} + +// MigrationKnownInternalPrefixes returns the concrete internal data/control +// prefixes that the user bracket must exclude. It intentionally does not +// include broad umbrellas such as !txn|, !ddb|, !sqs|, !s3|, or !stream|. +func MigrationKnownInternalPrefixes() [][]byte { + return cloneByteSlices(migrationInternalFamilyPrefixes) +} + +// IsMigrationKnownInternalKey reports whether a raw key belongs to a concrete +// internal family owned by an explicit export bracket or by the txn-lock drain. +func IsMigrationKnownInternalKey(key []byte) bool { + return hasAnyPrefix(key, migrationInternalFamilyPrefixes) +} + +// ValidateMigrationRouteRange rejects route intervals that intersect reserved +// distribution/migration control namespaces. +func ValidateMigrationRouteRange(routeStart, routeEnd []byte) error { + if len(routeEnd) > 0 && bytes.Compare(routeStart, routeEnd) >= 0 { + return errors.WithStack(ErrMigrationInvalidRoute) + } + for _, prefix := range migrationReservedControlPrefixes { + if rangesIntersect(routeStart, routeEnd, prefix, prefixScanEnd(prefix)) { + return errors.WithStack(ErrMigrationReservedRange) + } + } + return nil +} + +func migrationFamilyBrackets() []MigrationBracket { + defs := []struct { + family uint32 + prefix string + drainOnly bool + excludePrefixes []string + }{ + {family: MigrationFamilyTxnLock, prefix: migrationTxnLockPrefix, drainOnly: true}, + {family: MigrationFamilyTxnIntent, prefix: migrationTxnIntentPrefix}, + {family: MigrationFamilyTxnCommit, prefix: migrationTxnCommitPrefix}, + {family: MigrationFamilyTxnRollback, prefix: migrationTxnRollbackPrefix}, + {family: MigrationFamilyTxnSuccess, prefix: migrationTxnSuccessPrefix}, + {family: MigrationFamilyTxnMeta, prefix: migrationTxnMetaPrefix}, + {family: MigrationFamilyListMetaDelta, prefix: store.ListMetaDeltaPrefix}, + {family: MigrationFamilyLegacyListMetaDelta, prefix: store.LegacyListMetaDeltaPrefix}, + {family: MigrationFamilyListClaim, prefix: store.ListClaimPrefix}, + {family: MigrationFamilyListMeta, prefix: store.ListMetaPrefix}, + {family: MigrationFamilyListItem, prefix: store.ListItemPrefix}, + {family: MigrationFamilyRedisLegacy, prefix: migrationRedisPrefix}, + {family: MigrationFamilyHash, prefix: migrationHashPrefix}, + {family: MigrationFamilySet, prefix: migrationSetPrefix}, + {family: MigrationFamilyZSet, prefix: migrationZSetPrefix}, + {family: MigrationFamilyStreamMeta, prefix: store.StreamMetaPrefix}, + {family: MigrationFamilyStreamEntry, prefix: store.StreamEntryPrefix}, + {family: MigrationFamilyDynamoTableMeta, prefix: migrationDynamoMetaPrefix}, + {family: MigrationFamilyDynamoTableGeneration, prefix: migrationDynamoGenPrefix}, + {family: MigrationFamilyDynamoItem, prefix: migrationDynamoItemPrefix}, + {family: MigrationFamilyDynamoGSI, prefix: migrationDynamoGSIPrefix}, + {family: MigrationFamilySQSQueueMeta, prefix: migrationSQSQueueMetaPrefix}, + {family: MigrationFamilySQSQueueGeneration, prefix: migrationSQSQueueGenPrefix}, + {family: MigrationFamilySQSQueueSequence, prefix: migrationSQSQueueSeqPrefix}, + {family: MigrationFamilySQSQueueTombstone, prefix: migrationSQSQueueTombstonePrefix}, + {family: MigrationFamilySQSMessageData, prefix: migrationSQSMsgDataPrefix, excludePrefixes: []string{migrationSQSMsgDataPrefix + migrationSQSPartitionedSuffix}}, + {family: MigrationFamilySQSMessageVisibility, prefix: migrationSQSMsgVisPrefix, excludePrefixes: []string{migrationSQSMsgVisPrefix + migrationSQSPartitionedSuffix}}, + {family: MigrationFamilySQSMessageDedup, prefix: migrationSQSMsgDedupPrefix, excludePrefixes: []string{migrationSQSMsgDedupPrefix + migrationSQSPartitionedSuffix}}, + {family: MigrationFamilySQSMessageGroup, prefix: migrationSQSMsgGroupPrefix, excludePrefixes: []string{migrationSQSMsgGroupPrefix + migrationSQSPartitionedSuffix}}, + {family: MigrationFamilySQSMessageByAge, prefix: migrationSQSMsgByAgePrefix, excludePrefixes: []string{migrationSQSMsgByAgePrefix + migrationSQSPartitionedSuffix}}, + {family: MigrationFamilySQSPartitionedMessageData, prefix: migrationSQSMsgDataPrefix + migrationSQSPartitionedSuffix}, + {family: MigrationFamilySQSPartitionedMessageVisibility, prefix: migrationSQSMsgVisPrefix + migrationSQSPartitionedSuffix}, + {family: MigrationFamilySQSPartitionedMessageDedup, prefix: migrationSQSMsgDedupPrefix + migrationSQSPartitionedSuffix}, + {family: MigrationFamilySQSPartitionedMessageGroup, prefix: migrationSQSMsgGroupPrefix + migrationSQSPartitionedSuffix}, + {family: MigrationFamilySQSPartitionedMessageByAge, prefix: migrationSQSMsgByAgePrefix + migrationSQSPartitionedSuffix}, + {family: MigrationFamilyS3BucketMeta, prefix: s3keys.BucketMetaPrefix}, + {family: MigrationFamilyS3BucketGeneration, prefix: s3keys.BucketGenerationPrefix}, + {family: MigrationFamilyS3ObjectManifest, prefix: s3keys.ObjectManifestPrefix}, + {family: MigrationFamilyS3UploadMeta, prefix: s3keys.UploadMetaPrefix}, + {family: MigrationFamilyS3UploadPart, prefix: s3keys.UploadPartPrefix}, + {family: MigrationFamilyS3Blob, prefix: s3keys.BlobPrefix}, + {family: MigrationFamilyS3ChunkRef, prefix: s3keys.ChunkRefPrefix}, + {family: MigrationFamilyS3GCUpload, prefix: s3keys.GCUploadPrefix}, + {family: MigrationFamilyFilesystemChunk, prefix: string(fskeys.ChunkAllPrefix())}, + } + + out := make([]MigrationBracket, 0, len(defs)) + for _, def := range defs { + start := []byte(def.prefix) + requiresRouteKeyCheck, requiresDecodedS3 := migrationBracketRouteCheck(def.family) + out = append(out, MigrationBracket{ + BracketID: uint64(def.family), + Family: def.family, + Start: CloneBytes(start), + End: prefixScanEnd(start), + ExcludePrefixes: stringsToByteSlices(def.excludePrefixes), + DrainOnly: def.drainOnly, + RequiresRouteKeyCheck: requiresRouteKeyCheck, + RequiresDecodedS3: requiresDecodedS3, + }) + } + return out +} + +func migrationBracketRouteCheck(family uint32) (requiresRouteKeyCheck bool, requiresDecodedS3 bool) { + switch family { + case MigrationFamilyS3BucketMeta, MigrationFamilyS3BucketGeneration: + return false, true + default: + return true, false + } +} + +func normalizeMigrationRouteEnd(routeEnd []byte) []byte { + if len(routeEnd) == 0 { + return nil + } + return routeEnd +} + +func stringsToByteSlices(in []string) [][]byte { + if len(in) == 0 { + return nil + } + out := make([][]byte, len(in)) + for i := range in { + out[i] = []byte(in[i]) + } + return out +} + +func migrationWallMillisToUint64(ms int64) uint64 { + if ms <= 0 { + return 0 + } + return uint64(ms) +} + +func validateSplitJobSource(job SplitJob, source RouteDescriptor) error { + if job.SourceRouteID != source.RouteID { + return errors.WithStack(ErrMigrationSourceRouteChanged) + } + if len(job.SplitKey) == 0 { + return errors.WithStack(ErrMigrationInvalidRoute) + } + if bytes.Compare(job.SplitKey, source.Start) <= 0 { + return errors.WithStack(ErrMigrationInvalidRoute) + } + if len(source.End) > 0 && bytes.Compare(job.SplitKey, source.End) >= 0 { + return errors.WithStack(ErrMigrationInvalidRoute) + } + return nil +} + +func rangesIntersect(aStart, aEnd, bStart, bEnd []byte) bool { + if len(aEnd) > 0 && bytes.Compare(aEnd, bStart) <= 0 { + return false + } + if len(bEnd) > 0 && bytes.Compare(bEnd, aStart) <= 0 { + return false + } + return true +} + +func hasAnyPrefix(key []byte, prefixes [][]byte) bool { + for _, prefix := range prefixes { + if bytes.HasPrefix(key, prefix) { + return true + } + } + return false +} + +func cloneByteSlices(in [][]byte) [][]byte { + out := make([][]byte, len(in)) + for i := range in { + out[i] = CloneBytes(in[i]) + } + return out +} diff --git a/distribution/migrator_export_plan_test.go b/distribution/migrator_export_plan_test.go new file mode 100644 index 000000000..c9bd6d735 --- /dev/null +++ b/distribution/migrator_export_plan_test.go @@ -0,0 +1,544 @@ +package distribution + +import ( + "bytes" + "encoding/binary" + "testing" + + "github.com/bootjp/elastickv/internal/fskeys" + "github.com/bootjp/elastickv/internal/s3keys" + "github.com/bootjp/elastickv/store" + "github.com/cockroachdb/errors" + "github.com/stretchr/testify/require" +) + +func TestPlanMigrationBracketsIncludesRequiredFamilies(t *testing.T) { + t.Parallel() + + brackets, err := PlanMigrationBrackets([]byte("m"), []byte("z")) + require.NoError(t, err) + + byFamily := bracketsByFamily(brackets) + required := map[uint32]string{ + MigrationFamilyUser: "user", + MigrationFamilyTxnIntent: migrationTxnIntentPrefix, + MigrationFamilyTxnCommit: migrationTxnCommitPrefix, + MigrationFamilyTxnRollback: migrationTxnRollbackPrefix, + MigrationFamilyTxnSuccess: migrationTxnSuccessPrefix, + MigrationFamilyTxnMeta: migrationTxnMetaPrefix, + MigrationFamilyTxnLock: migrationTxnLockPrefix, + MigrationFamilyListMeta: store.ListMetaPrefix, + MigrationFamilyListItem: store.ListItemPrefix, + MigrationFamilyListMetaDelta: store.ListMetaDeltaPrefix, + MigrationFamilyLegacyListMetaDelta: store.LegacyListMetaDeltaPrefix, + MigrationFamilyListClaim: store.ListClaimPrefix, + MigrationFamilyRedisLegacy: migrationRedisPrefix, + MigrationFamilyHash: migrationHashPrefix, + MigrationFamilySet: migrationSetPrefix, + MigrationFamilyZSet: migrationZSetPrefix, + MigrationFamilyStreamMeta: store.StreamMetaPrefix, + MigrationFamilyStreamEntry: store.StreamEntryPrefix, + MigrationFamilyDynamoTableMeta: migrationDynamoMetaPrefix, + MigrationFamilyDynamoTableGeneration: migrationDynamoGenPrefix, + MigrationFamilyDynamoItem: migrationDynamoItemPrefix, + MigrationFamilyDynamoGSI: migrationDynamoGSIPrefix, + MigrationFamilySQSQueueMeta: migrationSQSQueueMetaPrefix, + MigrationFamilySQSQueueGeneration: migrationSQSQueueGenPrefix, + MigrationFamilySQSQueueSequence: migrationSQSQueueSeqPrefix, + MigrationFamilySQSQueueTombstone: migrationSQSQueueTombstonePrefix, + MigrationFamilySQSMessageData: migrationSQSMsgDataPrefix, + MigrationFamilySQSMessageVisibility: migrationSQSMsgVisPrefix, + MigrationFamilySQSMessageDedup: migrationSQSMsgDedupPrefix, + MigrationFamilySQSMessageGroup: migrationSQSMsgGroupPrefix, + MigrationFamilySQSMessageByAge: migrationSQSMsgByAgePrefix, + MigrationFamilySQSPartitionedMessageData: migrationSQSMsgDataPrefix + migrationSQSPartitionedSuffix, + MigrationFamilySQSPartitionedMessageVisibility: migrationSQSMsgVisPrefix + migrationSQSPartitionedSuffix, + MigrationFamilySQSPartitionedMessageDedup: migrationSQSMsgDedupPrefix + migrationSQSPartitionedSuffix, + MigrationFamilySQSPartitionedMessageGroup: migrationSQSMsgGroupPrefix + migrationSQSPartitionedSuffix, + MigrationFamilySQSPartitionedMessageByAge: migrationSQSMsgByAgePrefix + migrationSQSPartitionedSuffix, + MigrationFamilyS3BucketMeta: s3keys.BucketMetaPrefix, + MigrationFamilyS3BucketGeneration: s3keys.BucketGenerationPrefix, + MigrationFamilyS3ObjectManifest: s3keys.ObjectManifestPrefix, + MigrationFamilyS3UploadMeta: s3keys.UploadMetaPrefix, + MigrationFamilyS3UploadPart: s3keys.UploadPartPrefix, + MigrationFamilyS3Blob: s3keys.BlobPrefix, + MigrationFamilyS3ChunkRef: s3keys.ChunkRefPrefix, + MigrationFamilyS3GCUpload: s3keys.GCUploadPrefix, + MigrationFamilyFilesystemChunk: string(fskeys.ChunkAllPrefix()), + } + + for family, prefix := range required { + bracket, ok := byFamily[family] + require.True(t, ok, "missing family %d", family) + require.Equal(t, uint64(family), bracket.BracketID) + if family == MigrationFamilyS3BucketMeta || family == MigrationFamilyS3BucketGeneration { + require.False(t, bracket.RequiresRouteKeyCheck) + require.True(t, bracket.RequiresDecodedS3) + } else { + require.True(t, bracket.RequiresRouteKeyCheck) + require.False(t, bracket.RequiresDecodedS3) + } + if family == MigrationFamilyUser { + require.Equal(t, []byte("m"), bracket.Start) + require.Equal(t, []byte("z"), bracket.End) + require.True(t, bracket.ExcludeKnownInternal) + continue + } + require.Equal(t, []byte(prefix), bracket.Start, "family %d start", family) + require.Equal(t, prefixScanEnd([]byte(prefix)), bracket.End, "family %d end", family) + } + + require.True(t, byFamily[MigrationFamilyTxnLock].DrainOnly) + export, err := PlanExportBrackets([]byte("m"), []byte("z")) + require.NoError(t, err) + _, exportedLock := bracketsByFamily(export)[MigrationFamilyTxnLock] + require.False(t, exportedLock, "txn locks are drain-only and must not be exported as data") +} + +func TestPlanMigrationBracketsDisjointPrefixContainment(t *testing.T) { + t.Parallel() + + brackets, err := PlanMigrationBrackets([]byte("m"), []byte("z")) + require.NoError(t, err) + byFamily := bracketsByFamily(brackets) + + listDelta := store.ListMetaDeltaKey([]byte("list"), 1, 0) + require.True(t, byFamily[MigrationFamilyListMetaDelta].ContainsRawKey(listDelta)) + require.False(t, byFamily[MigrationFamilyListMeta].ContainsRawKey(listDelta)) + require.False(t, byFamily[MigrationFamilyLegacyListMetaDelta].ContainsRawKey(listDelta)) + + legacyListDelta := legacyListMetaDeltaKey([]byte("legacy-list"), 2, 0) + require.True(t, byFamily[MigrationFamilyLegacyListMetaDelta].ContainsRawKey(legacyListDelta)) + require.False(t, byFamily[MigrationFamilyListMeta].ContainsRawKey(legacyListDelta)) + require.False(t, byFamily[MigrationFamilyListMetaDelta].ContainsRawKey(legacyListDelta)) + + listMetaWithDeltaLookingUserKey := store.ListMetaKey(deltaLookingListMetaUserKey([]byte("list"), 2, 0)) + require.False(t, byFamily[MigrationFamilyListMeta].ContainsRawKey(listMetaWithDeltaLookingUserKey)) + require.False(t, byFamily[MigrationFamilyListMetaDelta].ContainsRawKey(listMetaWithDeltaLookingUserKey)) + require.True(t, byFamily[MigrationFamilyLegacyListMetaDelta].ContainsRawKey(listMetaWithDeltaLookingUserKey)) + + partitionedSQS := []byte(migrationSQSMsgDataPrefix + migrationSQSPartitionedSuffix + "queue|0|1|msg") + require.True(t, byFamily[MigrationFamilySQSPartitionedMessageData].ContainsRawKey(partitionedSQS)) + require.False(t, byFamily[MigrationFamilySQSMessageData].ContainsRawKey(partitionedSQS)) + + user := byFamily[MigrationFamilyUser] + user.Start = nil + user.End = nil + require.False(t, user.ContainsRawKey(s3keys.ChunkRefKey("bucket", 1, "object", "upload", 1, 0))) + require.False(t, user.ContainsRawKey(fskeys.ChunkKey(1, 2, 3))) + for _, raw := range [][]byte{ + []byte("!txn|foo"), + []byte("!stream|foo"), + []byte("!ddb|foo"), + []byte("!sqs|foo"), + []byte("!s3|foo"), + []byte("ordinary-user-key"), + } { + require.True(t, user.ContainsRawKey(raw), "raw user key %q must stay in familyUser", raw) + } + for _, raw := range [][]byte{ + []byte(migrationTxnSuccessPrefix + "x"), + []byte(store.StreamMetaPrefix + "x"), + []byte(migrationDynamoItemPrefix + "x"), + []byte(migrationSQSMsgVisPrefix + "x"), + []byte(s3keys.ObjectManifestPrefix + "x"), + []byte(migrationRedisPrefix + "string|k"), + []byte(migrationHashPrefix + "meta|x"), + } { + require.False(t, user.ContainsRawKey(raw), "concrete internal key %q must be excluded from familyUser", raw) + } +} + +func TestPlanMigrationBracketsNormalizesEmptyRouteEnd(t *testing.T) { + t.Parallel() + + brackets, err := PlanMigrationBrackets([]byte("m"), []byte{}) + require.NoError(t, err) + user := bracketsByFamily(brackets)[MigrationFamilyUser] + require.Nil(t, user.End) + require.True(t, user.ContainsRawKey([]byte("z"))) +} + +func TestSplitJobPlanNormalizesEmptySourceRouteEnd(t *testing.T) { + t.Parallel() + + source := RouteDescriptor{ + RouteID: 9, + Start: []byte("a"), + End: []byte{}, + GroupID: 3, + State: RouteStateActive, + } + job := SplitJob{ + JobID: 1, + SourceRouteID: source.RouteID, + SplitKey: []byte("m"), + TargetGroupID: source.GroupID, + Phase: SplitJobPhasePlanned, + } + + planned, err := InitializeSplitJobPlan(job, source, 1000) + require.NoError(t, err) + for _, progress := range planned.BracketProgress { + if progress.Family != MigrationFamilyUser { + continue + } + require.False(t, progress.Done) + return + } + require.Fail(t, "missing user bracket progress") +} + +func TestMigrationBracketContainsRoutedKeyForS3BucketAuxiliaryState(t *testing.T) { + t.Parallel() + + brackets, err := PlanMigrationBrackets([]byte("m"), []byte("z")) + require.NoError(t, err) + byFamily := bracketsByFamily(brackets) + + routeStart := s3keys.RouteKey("bucket-b", 7, "a") + routeEnd := s3keys.RouteKey("bucket-b", 7, "z") + for _, tc := range []struct { + name string + family uint32 + key []byte + want bool + }{ + {name: "meta same bucket", family: MigrationFamilyS3BucketMeta, key: s3keys.BucketMetaKey("bucket-b"), want: true}, + {name: "generation same bucket", family: MigrationFamilyS3BucketGeneration, key: s3keys.BucketGenerationKey("bucket-b"), want: true}, + {name: "meta different bucket", family: MigrationFamilyS3BucketMeta, key: s3keys.BucketMetaKey("bucket-c"), want: false}, + {name: "generation different bucket", family: MigrationFamilyS3BucketGeneration, key: s3keys.BucketGenerationKey("bucket-c"), want: false}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + got := byFamily[tc.family].ContainsRoutedKey(tc.key, routeStart, routeEnd, s3keys.ExtractRouteKey) + require.Equal(t, tc.want, got) + }) + } +} + +func TestMigrationBracketContainsRoutedKeyForS3BucketRawRoute(t *testing.T) { + t.Parallel() + + brackets, err := PlanMigrationBrackets([]byte("m"), []byte("z")) + require.NoError(t, err) + byFamily := bracketsByFamily(brackets) + routeStart := []byte("!s3|") + + require.True(t, byFamily[MigrationFamilyS3BucketMeta].ContainsRoutedKey( + s3keys.BucketMetaKey("bucket-b"), routeStart, nil, s3keys.ExtractRouteKey, + )) + require.True(t, byFamily[MigrationFamilyS3BucketGeneration].ContainsRoutedKey( + s3keys.BucketGenerationKey("bucket-b"), routeStart, nil, s3keys.ExtractRouteKey, + )) +} + +func TestMigrationBracketContainsRoutedKeyUsesObjectRoutes(t *testing.T) { + t.Parallel() + + brackets, err := PlanMigrationBrackets([]byte("m"), []byte("z")) + require.NoError(t, err) + manifest := bracketsByFamily(brackets)[MigrationFamilyS3ObjectManifest] + + key := s3keys.ObjectManifestKey("bucket-b", 7, "m") + require.True(t, manifest.ContainsRoutedKey( + key, + s3keys.RouteKey("bucket-b", 7, "a"), + s3keys.RouteKey("bucket-b", 7, "z"), + s3keys.ExtractRouteKey, + )) + require.False(t, manifest.ContainsRoutedKey( + key, + s3keys.RouteKey("bucket-c", 1, "a"), + nil, + s3keys.ExtractRouteKey, + )) +} + +func TestMigrationBracketContainsRoutedKeyUsesS3ChunkRefRoutes(t *testing.T) { + t.Parallel() + + brackets, err := PlanMigrationBrackets([]byte("m"), []byte("z")) + require.NoError(t, err) + chunkRef := bracketsByFamily(brackets)[MigrationFamilyS3ChunkRef] + key := s3keys.ChunkRefKey("bucket-b", 7, "m", "upload", 1, 0) + + require.True(t, chunkRef.ContainsRoutedKey( + key, + s3keys.RouteKey("bucket-b", 7, "a"), + s3keys.RouteKey("bucket-b", 7, "z"), + s3keys.ExtractRouteKey, + )) + require.False(t, chunkRef.ContainsRoutedKey( + key, + s3keys.RouteKey("bucket-c", 1, "a"), + nil, + s3keys.ExtractRouteKey, + )) +} + +func TestMigrationBracketContainsRoutedKeyUsesFilesystemChunkRoutes(t *testing.T) { + t.Parallel() + + brackets, err := PlanMigrationBrackets([]byte("m"), []byte("z")) + require.NoError(t, err) + chunk := bracketsByFamily(brackets)[MigrationFamilyFilesystemChunk] + key := fskeys.ChunkKey(10, 20, 3) + routeKey := fskeys.ChunkRouteKey(10, 20) + + require.True(t, chunk.ContainsRoutedKey(key, routeKey, prefixScanEnd(routeKey), fskeys.ExtractRouteKey)) + require.False(t, chunk.ContainsRoutedKey( + key, + fskeys.ChunkRouteKey(11, 20), + nil, + fskeys.ExtractRouteKey, + )) +} + +func TestMigrationBracketContainsRoutedKeyAcceptsEmptyLogicalRouteKey(t *testing.T) { + t.Parallel() + + routeEnd := []byte{0x01} + brackets, err := PlanMigrationBrackets(nil, routeEnd) + require.NoError(t, err) + hash := bracketsByFamily(brackets)[MigrationFamilyHash] + rawKey := store.HashMetaKey(nil) + + require.True(t, hash.ContainsRoutedKey( + rawKey, + nil, + routeEnd, + store.ExtractHashUserKeyFromMeta, + )) + require.False(t, hash.ContainsRoutedKey( + rawKey, + nil, + routeEnd, + func([]byte) []byte { return nil }, + )) +} + +func TestMigrationBracketContainsRoutedKeyUsesLegacyListDeltaUserKey(t *testing.T) { + t.Parallel() + + brackets, err := PlanMigrationBrackets([]byte("a"), []byte("z")) + require.NoError(t, err) + legacy := bracketsByFamily(brackets)[MigrationFamilyLegacyListMetaDelta] + raw := legacyListMetaDeltaKey([]byte("target-list"), 10, 0) + value := store.MarshalListMetaDelta(store.ListMetaDelta{LenDelta: 1}) + + require.True(t, legacy.ContainsRoutedVersion( + raw, + value, + []byte("target"), + []byte("target-list\x00"), + store.ExtractListUserKey, + )) + require.False(t, legacy.ContainsRoutedVersion( + raw, + value, + []byte("d|"), + []byte("d}"), + store.ExtractListUserKey, + )) + require.False(t, legacy.ContainsRoutedVersion( + raw, + value, + []byte("zzz"), + nil, + store.ExtractListUserKey, + )) +} + +func TestMigrationBracketContainsRoutedKeyRoutesAmbiguousLegacyListMetaByBaseKey(t *testing.T) { + t.Parallel() + + brackets, err := PlanMigrationBrackets([]byte("a"), []byte("z")) + require.NoError(t, err) + legacy := bracketsByFamily(brackets)[MigrationFamilyLegacyListMetaDelta] + userKey := deltaLookingListMetaUserKey([]byte("target-list"), 10, 0) + raw := store.ListMetaKey(userKey) + value, err := store.MarshalListMeta(store.ListMeta{Head: 1, Tail: 2, Len: 1}) + require.NoError(t, err) + + require.True(t, legacy.ContainsRoutedVersion( + raw, + value, + []byte("d|"), + []byte("d}"), + store.ExtractListUserKey, + )) + require.False(t, legacy.ContainsRoutedVersion( + raw, + value, + []byte("zzz"), + nil, + store.ExtractListUserKey, + )) +} + +func TestMigrationBracketContainsRoutedKeyRoutesAmbiguousListMetaTombstoneConservatively(t *testing.T) { + t.Parallel() + + brackets, err := PlanMigrationBrackets([]byte("a"), []byte("z")) + require.NoError(t, err) + legacy := bracketsByFamily(brackets)[MigrationFamilyLegacyListMetaDelta] + baseUserKey := deltaLookingListMetaUserKey([]byte("target-list"), 10, 0) + raw := store.ListMetaKey(baseUserKey) + + require.True(t, legacy.ContainsRoutedVersion( + raw, + nil, + []byte("d|"), + []byte("d}"), + store.ExtractListUserKey, + )) + require.True(t, legacy.ContainsRoutedVersion( + raw, + nil, + []byte("target"), + []byte("target-list\x00"), + store.ExtractListUserKey, + )) + require.False(t, legacy.ContainsRoutedVersion( + raw, + nil, + []byte("zzz"), + nil, + store.ExtractListUserKey, + )) +} + +func TestMigrationKnownInternalPrefixesAreConcreteOnly(t *testing.T) { + t.Parallel() + + for _, raw := range [][]byte{ + []byte(migrationTxnIntentPrefix + "k"), + []byte(migrationTxnSuccessPrefix + "k"), + []byte(store.ListClaimPrefix + "k"), + []byte(store.HashFieldPrefix + "k"), + []byte(store.StreamEntryPrefix + "k"), + []byte(migrationDynamoMetaPrefix + "t"), + []byte(migrationSQSQueueMetaPrefix + "q"), + []byte(s3keys.BlobPrefix + "b"), + } { + require.True(t, IsMigrationKnownInternalKey(raw), "concrete internal key %q", raw) + } + + for _, raw := range [][]byte{ + []byte("!txn|foo"), + []byte("!stream|foo"), + []byte("!ddb|foo"), + []byte("!sqs|foo"), + []byte("!s3|foo"), + } { + require.False(t, IsMigrationKnownInternalKey(raw), "umbrella-looking user key %q", raw) + } + + prefixes := MigrationKnownInternalPrefixes() + require.NotEmpty(t, prefixes) + prefixes[0][0] ^= 0xff + require.False(t, bytes.Equal(prefixes[0], MigrationKnownInternalPrefixes()[0]), "prefix list must be cloned") +} + +func TestValidateMigrationRouteRangeRejectsReservedControlPrefixes(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name string + start []byte + end []byte + }{ + {name: "exact dist", start: []byte("!dist|"), end: prefixScanEnd([]byte("!dist|"))}, + {name: "migstage", start: []byte("!migstage|"), end: prefixScanEnd([]byte("!migstage|"))}, + {name: "broad intersection", start: []byte("!"), end: []byte("~")}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + err := ValidateMigrationRouteRange(tc.start, tc.end) + require.True(t, errors.Is(err, ErrMigrationReservedRange), "got %v", err) + }) + } + + require.NoError(t, ValidateMigrationRouteRange([]byte("m"), []byte("z"))) + err := ValidateMigrationRouteRange([]byte("z"), []byte("m")) + require.True(t, errors.Is(err, ErrMigrationInvalidRoute), "got %v", err) +} + +func TestSplitJobPlannerAndSameGroupNoop(t *testing.T) { + t.Parallel() + + source := RouteDescriptor{ + RouteID: 9, + Start: []byte("a"), + End: []byte("z"), + GroupID: 3, + State: RouteStateActive, + } + job := SplitJob{ + JobID: 1, + SourceRouteID: source.RouteID, + SplitKey: []byte("m"), + TargetGroupID: source.GroupID, + Phase: SplitJobPhasePlanned, + } + + planned, err := InitializeSplitJobPlan(job, source, 1000) + require.NoError(t, err) + require.Equal(t, SplitJobPhasePlanned, planned.Phase) + require.NotEmpty(t, planned.BracketProgress) + require.Equal(t, int64(1000), planned.StartedAtMs) + require.Equal(t, int64(1000), planned.UpdatedAtMs) + for _, progress := range planned.BracketProgress { + require.Equal(t, SplitJobExportPhaseBackfill, progress.ExportPhase) + require.NotEqual(t, MigrationFamilyTxnLock, progress.Family) + } + + done, err := AdvanceSameGroupNoop(job, source, 2000) + require.NoError(t, err) + require.Equal(t, SplitJobPhaseDone, done.Phase) + require.True(t, done.TargetPromotionDone) + require.Equal(t, uint64(2000), done.PromotionCompletedTS) + require.Equal(t, int64(2000), done.TerminalAtMs) + for _, progress := range done.BracketProgress { + require.True(t, progress.Done) + } + + crossGroup := job + crossGroup.TargetGroupID = source.GroupID + 1 + _, err = AdvanceSameGroupNoop(crossGroup, source, 3000) + require.True(t, errors.Is(err, ErrMigrationDataMoveRequired), "got %v", err) +} + +func bracketsByFamily(brackets []MigrationBracket) map[uint32]MigrationBracket { + out := make(map[uint32]MigrationBracket, len(brackets)) + for _, bracket := range brackets { + out[bracket.Family] = bracket + } + return out +} + +func deltaLookingListMetaUserKey(fakeUserKey []byte, commitTS uint64, seqInTxn uint32) []byte { + key := make([]byte, 0, len("d|")+4+len(fakeUserKey)+8+4) + key = append(key, "d|"...) + var lenPrefix [4]byte + binary.BigEndian.PutUint32(lenPrefix[:], uint32(len(fakeUserKey))) //nolint:gosec // test data is small. + key = append(key, lenPrefix[:]...) + key = append(key, fakeUserKey...) + var ts [8]byte + binary.BigEndian.PutUint64(ts[:], commitTS) + key = append(key, ts[:]...) + var seq [4]byte + binary.BigEndian.PutUint32(seq[:], seqInTxn) + return append(key, seq[:]...) +} + +func legacyListMetaDeltaKey(userKey []byte, commitTS uint64, seqInTxn uint32) []byte { + key := store.LegacyListMetaDeltaScanPrefix(userKey) + var ts [8]byte + binary.BigEndian.PutUint64(ts[:], commitTS) + key = append(key, ts[:]...) + var seq [4]byte + binary.BigEndian.PutUint32(seq[:], seqInTxn) + return append(key, seq[:]...) +} diff --git a/distribution/split_job_catalog.go b/distribution/split_job_catalog.go index 27be76834..2470a1252 100644 --- a/distribution/split_job_catalog.go +++ b/distribution/split_job_catalog.go @@ -35,6 +35,7 @@ var ( ErrCatalogSplitJobKeyIDMismatch = errors.New("catalog split job key and record job id mismatch") ErrCatalogSplitJobConflict = errors.New("catalog split job conflict") ErrCatalogSplitJobTerminalRequired = errors.New("catalog split job terminal state is required") + ErrSplitJobOverlap = errors.New("split job overlaps requested route") ) // SplitJobPhase is the durable phase of a split migration job. diff --git a/docs/design/2026_04_14_implemented_etcd_snapshot_disk_offload.md b/docs/design/2026_04_14_implemented_etcd_snapshot_disk_offload.md index 5312c3784..b800f2906 100644 --- a/docs/design/2026_04_14_implemented_etcd_snapshot_disk_offload.md +++ b/docs/design/2026_04_14_implemented_etcd_snapshot_disk_offload.md @@ -5,7 +5,7 @@ ### Observed Symptoms Clusters using the etcd engine exhibit large memory spikes across all nodes at snapshot -creation intervals (`defaultSnapshotEvery = 10,000` entries). The spike magnitude scales +creation intervals (`defaultSnapshotEvery = 100,000` entries). The spike magnitude scales with FSM data size, and simultaneous spikes across multiple nodes compound the pressure. ### Root Cause @@ -34,7 +34,7 @@ storage.CreateSnapshot(applied, &confState, payload) ``` `MemoryStorage` holds `raftpb.Snapshot.Data = payload` until the next snapshot is created -(i.e., until another 10,000 entries are processed), keeping the full FSM export resident +(i.e., until another 100,000 entries are processed), keeping the full FSM export resident in memory the entire time. #### Problem 3: Re-allocation when sending to followers diff --git a/docs/design/2026_04_27_implemented_keyviz_cluster_fanout.md b/docs/design/2026_04_27_implemented_keyviz_cluster_fanout.md index 2c307da2a..4cb913cce 100644 --- a/docs/design/2026_04_27_implemented_keyviz_cluster_fanout.md +++ b/docs/design/2026_04_27_implemented_keyviz_cluster_fanout.md @@ -46,6 +46,15 @@ This implemented design records both the static-node-list Phase 2-C fan-out and the Phase 2-C+ wire/merge extension that now carries per-cell Raft group and leader-term identity. +## 1.1 Implementation status + +Implemented in `internal/admin/keyviz_fanout.go`, +`internal/admin/keyviz_handler.go`, `web/admin/src/pages/KeyViz.tsx`, +and `main.go`'s `--keyvizFanoutNodes` / `--keyvizFanoutTimeout` +flags. Tests cover fan-out merge and degraded-node handling in +`internal/admin/keyviz_fanout_test.go` and the admin KeyViz UI +suite. + ## 2. Scope ### 2.1 In scope diff --git a/docs/design/2026_04_29_implemented_snapshot_logical_decoder.md b/docs/design/2026_04_29_implemented_snapshot_logical_decoder.md index 07205fc7c..f4f721ba4 100644 --- a/docs/design/2026_04_29_implemented_snapshot_logical_decoder.md +++ b/docs/design/2026_04_29_implemented_snapshot_logical_decoder.md @@ -37,7 +37,7 @@ which is a different framing the decoder/encoder do not touch. See `2026_05_25_implemented_snapshot_logical_encoder.md` §"Why a separate design doc" item 3. -Snapshots are taken automatically every `defaultSnapshotEvery = 10000` +Snapshots are taken automatically every `defaultSnapshotEvery = 100000` log entries (`internal/raftengine/etcd/engine.go:92`) and stored under `{dataDir}/fsm-snap/.fsm`. They are crash-consistent by construction — the writer takes a Pebble snapshot at the FSM's @@ -608,7 +608,7 @@ bespoke parser, the format has failed its goal. - **Staleness.** Whatever was written between the snapshot's `applied_index` and "now" is not in the snapshot, so it is not in the decoded output. The gap is bounded by `SnapshotEvery × write_rate` - (default 10000 entries; for a write-heavy cluster, seconds; for a + (default 100000 entries; for a write-heavy cluster, minutes; for a quiet one, hours). - **Cadence is not user-controlled.** "Snapshot now" requires a Raft trigger; the decoder cannot create a fresh snapshot, only consume diff --git a/docs/design/2026_04_29_proposed_logical_backup.md b/docs/design/2026_04_29_proposed_logical_backup.md index 39350d188..1e307c832 100644 --- a/docs/design/2026_04_29_proposed_logical_backup.md +++ b/docs/design/2026_04_29_proposed_logical_backup.md @@ -884,7 +884,7 @@ refuse if: remaining_headroom < --snapshot-headroom-entries ``` `SnapshotEvery` is the per-engine snapshot trigger (default -`defaultSnapshotEvery = 10000` from `internal/raftengine/etcd/engine.go:92`, +`defaultSnapshotEvery = 100000` from `internal/raftengine/etcd/engine.go:92`, overridden via the `ELASTICKV_RAFT_SNAPSHOT_COUNT` env var — there is no `--raftSnapshotEvery` CLI flag). The value is currently a private field on the etcd `Engine` struct (`internal/raftengine/etcd/engine.go:224`) @@ -906,9 +906,9 @@ implemented on the etcd backend by returning the place, `BeginBackup` reads each group's `SnapshotEvery` rather than hardcoding `defaultSnapshotEvery`, so an operator who tuned `ELASTICKV_RAFT_SNAPSHOT_COUNT` sees consistent behavior. With -`SnapshotEvery = 10000` and -`--snapshot-headroom-entries = 1000` (default; one-tenth of -SnapshotEvery), the check refuses backups when fewer than 1000 entries +`SnapshotEvery = 100000` and +`--snapshot-headroom-entries = 10000` (default; one-tenth of +SnapshotEvery), the check refuses backups when fewer than 10000 entries remain before the next snapshot fires — i.e. when an in-flight backup is at risk of triggering the snapshot-installation corner case. A freshly-snapshotted cluster has the *largest* remaining headroom and @@ -1341,7 +1341,7 @@ written. `s.groups[id]` map without a typed assertion fallback. Tests that mock `AdminGroup` (e.g. `adapter/admin_grpc_test.go`) gain one extra method to implement; they can return - `defaultSnapshotEvery = 10000` for parity with production. + `defaultSnapshotEvery = 100000` for parity with production. - Extend `kv/active_timestamp_tracker.go` with `PinWithDeadline`, `Extend`, and the per-second sweeper goroutine that reaps expired pins and emits the `backup_pin_expired` structured warning. @@ -1601,7 +1601,7 @@ Scope: out of this proposal; mentioned only to draw the boundary. | `TestAdminGRPCConnCacheReuse` | Two consecutive `BeginBackup` calls dialing the same peer share one underlying `*grpc.ClientConn` (verified via the cache size); shutdown of a peer evicts only that entry, not the whole cache; admin cache is independent of `ShardStore.connCache` (no cross-cache eviction) | | `TestBeginBackupPropagatesAdminAuthToken` | A cluster booted with `--adminToken` accepts a `BeginBackup` call carrying `authorization: Bearer `; the handler propagates the same metadata via `metadata.NewOutgoingContext` to every `GetNodeVersion` fan-out dial, so peers return their version (not `Unauthenticated`). A `BeginBackup` call without the token is itself rejected before any fan-out happens | | `TestVersionCacheRaceUnderLoad` | `go test -race` with 50 concurrent `GetRaftGroups` callers and async `GetNodeVersion` probe goroutines writing the cache simultaneously emits no data-race report; the `sync.Map` choice is enforced by the lack of a separate `versionCacheMu` field on `AdminServer` | -| `TestSnapshotEveryReadsFromEngine` | A node started with `ELASTICKV_RAFT_SNAPSHOT_COUNT=5000` reports `Engine.SnapshotEvery() == 5000`; `BeginBackup` uses 5000 (not the default 10000) when computing remaining headroom | +| `TestSnapshotEveryReadsFromEngine` | A node started with `ELASTICKV_RAFT_SNAPSHOT_COUNT=5000` reports `Engine.SnapshotEvery() == 5000`; `BeginBackup` uses 5000 (not the default 100000) when computing remaining headroom | | `TestRenewBackupRetriesLeaderElection` | Force a leader election mid-`RenewBackup`; the admin server retries `BackupExtend` up to 3 times with 500ms backoff and succeeds once the new leader is established, without aborting the dump | | `TestPinWithDeadlineExpiry` | `PinWithDeadline(ts, now+100ms)` is auto-released by the sweeper after the deadline; compactor unblocked; `backup_pin_expired` log emitted | | `TestBeginBackupWaitsForLaggingShard` | Force shard B's `applied_index` to lag; `BeginBackup` polls until it catches up or times out with `FailedPrecondition`; no scan starts in the timeout case | diff --git a/docs/design/2026_05_28_implemented_tla_safety_spec.md b/docs/design/2026_05_28_implemented_tla_safety_spec.md index cd94bb5d0..463fa16ba 100644 --- a/docs/design/2026_05_28_implemented_tla_safety_spec.md +++ b/docs/design/2026_05_28_implemented_tla_safety_spec.md @@ -199,8 +199,8 @@ cannot check them as state invariants. independently reach `ceiling + 1` before a fresh ceiling is renewed, the new leader's first `Next()` can tie or undercut the old leader's last commit. Bounding inter-node skew to less than - one ceiling window (< 3s with the current `hlcPhysicalWindowMs = - 3000ms`) keeps the window wide enough that the new leader cannot + one ceiling window (< 15s with the current `hlcPhysicalWindowMs = + 15000ms`) keeps the window wide enough that the new leader cannot independently reach the overflow value before a renewal applies. Should be surfaced in operator docs as a cluster prerequisite. - **(ii) Logical-counter handoff.** The 16-bit logical half of the HLC @@ -270,7 +270,7 @@ cannot check them as state invariants. `hlcPhysicalWindowMs` cannot serve any persistence timestamp — every client commit is rejected until renewal succeeds. This is a CP, not AP, trade-off and operators must size - `hlcPhysicalWindowMs` (currently 3s) relative to expected + `hlcPhysicalWindowMs` (currently 15s) relative to expected partition duration; see §9 risk 7. ### 5.2 OCC @@ -676,7 +676,7 @@ does not keep this document in `partial`. 7. **Fail-closed availability under partition.** HLC-4 precondition (iii) makes the ceiling-fence behaviour normative: a leader partitioned from the default group's quorum for longer than - `hlcPhysicalWindowMs` (currently 3s) cannot serve any persistence + `hlcPhysicalWindowMs` (currently 15s) cannot serve any persistence timestamp, so client commits are rejected until renewal succeeds. This is a CP, not AP, trade-off and is a stricter regime than the current implementation (which silently keeps issuing). Mitigation: diff --git a/docs/design/2026_06_02_implemented_idempotent_snapshot_restore.md b/docs/design/2026_06_02_implemented_idempotent_snapshot_restore.md index 2317fb9db..a3db4b2f2 100644 --- a/docs/design/2026_06_02_implemented_idempotent_snapshot_restore.md +++ b/docs/design/2026_06_02_implemented_idempotent_snapshot_restore.md @@ -297,9 +297,9 @@ func restoreSnapshotState(fsm StateMachine, snapshot raftpb.Snapshot, fsmSnapDir // The body restore is skipped, but we MUST still consume // the v1/v2 snapshot header so the FSM picks up the HLC // ceiling AND the Stage 8a cutover (see §5). Thread - // tok.CRC32C through so the skip path verifies the file - // before mutating FSM state, matching the existing - // openAndRestoreFSMSnapshot safety contract. + // tok.CRC32C through so the skip path verifies the cheap + // snapshot envelope before mutating FSM state. Full-body CRC + // remains on the restore path where body bytes are consumed. return applyHeaderStateOnSkip(fsm, fsmSnapPath(fsmSnapDir, tok.Index), tok.CRC32C) } return openAndRestoreFSMSnapshot(fsm, fsmSnapPath(fsmSnapDir, tok.Index), tok.CRC32C) @@ -423,42 +423,35 @@ or `internal/raftengine/etcd → kv` edges (test-only on both sides), and adding either to satisfy the round-6 design either duplicates the CRC verifier in `kv` or breaks the layering. -Round-7 keeps the CRC verifier in its existing package and splits the -seam into two phases — a **parse** phase that reads the header from a -caller-supplied reader (and drains the rest, for CRC coverage), and an -**apply** phase that is pure assignment. The engine orchestrates -size + footer + tee'd CRC computation around the parse phase, then -calls apply only after all three pass: +Round-7 keeps the envelope checks in the engine package and splits the +seam into two phases — a **parse** phase that reads only the header from +a caller-supplied reader, and an **apply** phase that is pure assignment. +The engine validates size + footer-vs-tokenCRC before parsing and calls +apply only after those checks and the header parse pass: ```go // internal/raftengine/statemachine.go (new, sibling to ApplyIndexAware) type SnapshotHeaderApplier interface { - // ParseSnapshotHeader reads the v1/v2 header from r, drains the - // remaining bytes (so a wrapping crc32 TeeReader covers the full - // payload), and returns the parsed (ceiling, cutover) pair WITHOUT - // mutating FSM state. Implementations MUST NOT touch any FSM - // fields here; the engine calls ApplySnapshotHeader separately - // only after the wrapping CRC verification passes. + // ParseSnapshotHeader reads the v1/v2 header from r and returns the + // parsed (ceiling, cutover) pair WITHOUT mutating FSM state. + // Implementations MUST NOT touch any FSM fields here; the engine + // calls ApplySnapshotHeader separately only after the snapshot file + // footer matches the raft token. // // Errors propagate from the underlying header parser - // (ErrSnapshotHeaderUnknownMagic / InvalidLength) or from the - // drain pass (I/O errors). FSM state stays untouched on error. + // (ErrSnapshotHeaderUnknownMagic / InvalidLength). FSM state stays + // untouched on error. ParseSnapshotHeader(r io.Reader) (ceiling, cutover uint64, err error) // ApplySnapshotHeader is pure assignment of the verified header // state. The engine calls this only after ParseSnapshotHeader - // returned and the wrapping crc32 hash matched the file footer. + // returned and the snapshot file's footer matched the raft token. ApplySnapshotHeader(ceiling, cutover uint64) } // internal/raftengine/etcd/wal_store.go -- never imports kv; -// CRC verification stays here where the helpers live. +// cheap snapshot envelope checks stay here where the helpers live. func applyHeaderStateOnSkip(fsm StateMachine, snapPath string, tokenCRC uint32) error { - setter, ok := fsm.(SnapshotHeaderApplier) - if !ok { - return nil // FSM has no header state; skip is harmless. - } - file, err := os.Open(snapPath) if err != nil { return statFSMFileError(err) } defer file.Close() @@ -480,29 +473,24 @@ func applyHeaderStateOnSkip(fsm StateMachine, snapPath string, tokenCRC uint32) "path=%s footer=%08x token=%08x", snapPath, footer, tokenCRC) } - // Step 3: full-body CRC. Wrap the payload in a crc32 TeeReader and - // hand it to the FSM's ParseSnapshotHeader for header parse + drain. - // The header bytes are included in the computed CRC because the - // FSM reads them from the tee'd reader. + setter, ok := fsm.(SnapshotHeaderApplier) + if !ok { + return nil // FSM has no header state; skip is harmless. + } + if _, err := file.Seek(0, io.SeekStart); err != nil { return errors.WithStack(err) } payloadSize := info.Size() - fsmFooterSize - h := crc32.New(crc32cTable) - tee := io.TeeReader(io.LimitReader(file, payloadSize), h) - ceiling, cutover, perr := setter.ParseSnapshotHeader(tee) + ceiling, cutover, perr := setter.ParseSnapshotHeader(io.LimitReader(file, payloadSize)) if perr != nil { - // ErrSnapshotHeaderUnknownMagic / InvalidLength / I/O error - // surfaced from the FSM's parse pass. State unchanged. + // ErrSnapshotHeaderUnknownMagic / InvalidLength surfaced from + // the FSM's parse pass. State unchanged. return errors.WithStack(perr) } - if h.Sum32() != footer { - return errors.Wrapf(ErrFSMSnapshotFileCRC, - "path=%s footer=%08x computed=%08x", snapPath, footer, h.Sum32()) - } - // All three checks passed; apply side-effects. + // Envelope checks and header parse passed; apply side-effects. setter.ApplySnapshotHeader(ceiling, cutover) return nil } @@ -510,16 +498,10 @@ func applyHeaderStateOnSkip(fsm StateMachine, snapPath string, tokenCRC uint32) // kv/fsm.go (new methods on kvFSM) -- kv.ReadSnapshotHeader stays inside kv; // no imports of internal/raftengine/etcd or its private helpers. func (f *kvFSM) ParseSnapshotHeader(r io.Reader) (uint64, uint64, error) { - // The engine has already wrapped r in a crc32 TeeReader sized at - // the body payload (file size minus 4-byte footer). We read the - // header, then drain the rest of the body so the engine's CRC - // covers every byte (matching restoreAndComputeCRC's behaviour). - br := bufio.NewReaderSize(r, 1<<20) //nolint:mnd // 1 MiB, local to kv - ceiling, cutover, err := ReadSnapshotHeader(br) + // The skip path already has the FSM body locally, so read only + // the snapshot header needed for HLC/cutover state. + ceiling, cutover, err := ReadSnapshotHeader(bufio.NewReaderSize(r, 4<<10)) if err != nil { return 0, 0, errors.WithStack(err) } - if _, err := io.Copy(io.Discard, br); err != nil { - return 0, 0, errors.WithStack(err) - } return ceiling, cutover, nil } @@ -531,27 +513,16 @@ func (f *kvFSM) ApplySnapshotHeader(ceiling, cutover uint64) { } ``` -**Cost note**. Step 3 reads the full snapshot file once (through the -crc32 TeeReader). For multi-GiB FSMs this is a non-trivial I/O cost -— but it is **strictly cheaper** than the restore path it replaces -(which also reads the file once via `restoreAndComputeCRC` AND -additionally writes a temp Pebble database with sstable / WAL output). -Observed restore wall-clock is dominated by Pebble writes, not reads; -eliding the writes preserves the bulk of the win. A future -optimisation could persist the HLC ceiling + cutover durably -(analogous to `metaAppliedIndex`) and elide the file read entirely — -out of scope here, flagged under Open Questions. - -**Why this seam shape**. The two-phase split lets the CRC verifier -stay co-located with its private helpers in -`internal/raftengine/etcd/fsm_snapshot_file.go`'s package, **and** -keeps the v1/v2 header parser inside `kv` where it already lives. -Neither package imports the other in production. The "do CRC on -engine side, side-effects on FSM side after verify" contract is -exactly the inversion of `openAndRestoreFSMSnapshot` (which inlines -`fsm.Restore` inside the CRC tee for performance reasons): for the -skip path we don't need a single-pass restore, so splitting the -phases costs nothing and buys layer hygiene. +**Cost note**. The skip path does not read the multi-GiB body. Startup +cost is bounded by opening the snapshot file, reading the footer, and +parsing the fixed-size header. Full-body CRC verification remains on +the execute/full-restore path, where the body bytes are actually used. + +**Why this seam shape**. The two-phase split keeps the snapshot +envelope checks in `internal/raftengine/etcd`, and keeps the v1/v2 +header parser inside `kv` where it already lives. Neither package +imports the other in production. The skip path deliberately avoids the +restore path's body CRC work because it does not consume body bytes. ### 6. Crash-safety argument @@ -617,7 +588,7 @@ window. By forcing `pebble.Sync` on `SetDurableAppliedIndex` we make the checkpoint at least as durable as the snapshot pointer that follows. Cost: +1 extra fsync per snapshot persist (rare; default -`SnapshotCount=10000`). Negligible vs. the savings, and the only +`SnapshotCount=100000`). Negligible vs. the savings, and the only way to keep the round-4 ordering proof intact under nosync mode. #### Encryption opcodes (`OpRegistration`/`OpBootstrap`/`OpRotation`) @@ -768,7 +739,7 @@ func (s *pebbleStore) SetDurableAppliedIndex(idx uint64) error { // because WAL compaction starts at the snapshot index, no future // replay can re-bump metaAppliedIndex from the lost lease/data // applies, and the skip permanently falls back. The +1 extra fsync - // per snapshot persist (rare; default SnapshotCount=10000) is the + // per snapshot persist (rare; default SnapshotCount=100000) is the // right price. return errors.WithStack(b.Commit(pebble.Sync)) } @@ -801,12 +772,12 @@ permanent fallback case is closed. **Cost**: one extra pebble `Batch.Commit` (Sync per `ELASTICKV_FSM_SYNC_MODE`) per snapshot persist. Snapshots fire on the -etcd raft `SnapshotCount` cadence (default 10000 entries), so this is -~one extra fsync per ~10000 entries — negligible. +etcd raft `SnapshotCount` cadence (default 100000 entries), so this is +~one extra fsync per ~100000 entries — negligible. **Why not bump on every HLC lease apply**. Option A (1 pebble batch per lease tick) costs ~1 fsync/sec/group continuously. Option B (the -snapshot-persist hook) costs ~1 fsync per 10000 entries. Both close +snapshot-persist hook) costs ~1 fsync per 100000 entries. Both close the skip gap; B costs ~10⁴× less and aligns with the natural durability boundary the engine already maintains. @@ -903,7 +874,7 @@ restoreSnapshotState skipped (FSM at index %d, snapshot at %d, ceiling=%d, cutov |---|---|---| | **B1** (this PR) | Design doc | None | | **B2** | `ApplyMutationsRaftAt` / `DeletePrefixAtRaftAt` overloads + meta-key bundling in both leaves + `pebbleStore.LastAppliedIndex()` (under `dbMu.RLock()`) + `pebbleStore.SetDurableAppliedIndex()` (under `dbMu.RLock()` + `applyMu.Lock()` RMW monotonic guard, **`pebble.Sync` unconditionally**) + `kvFSM.LastAppliedIndex()` directly satisfies `raftengine.AppliedIndexReader` (compile-time guard in `kv/fsm_applied_index_iface_check.go`) + `kvFSM.SetDurableAppliedIndex` forwarding + thread `f.pendingApplyIdx` into the data-Apply leaves + BOTH `persistCreatedSnapshot` (`engine.go:2679`) AND `e.persistLocalSnapshotPayload` (`engine.go:4032`, the SnapshotCount-triggered hot path) call `SetDurableAppliedIndex` BEFORE the corresponding `persist.SaveSnap` | Meta key starts being written on every data Apply AND at every snapshot persist (both config-snapshot and steady-state local-snapshot paths). Skip is still disabled. Soak in production for one release. | -| **B3** | `restoreSnapshotState` skip gate + `applyHeaderStateOnSkip(snapPath, tok.CRC32C)` orchestrating size + footer-vs-tokenCRC + full-body-CRC verification using `internal/raftengine/etcd`'s existing helpers (matching `openAndRestoreFSMSnapshot`'s safety contract) + two-phase `SnapshotHeaderApplier` seam on `kvFSM` (`ParseSnapshotHeader(r io.Reader) (ceiling, cutover, err)` + pure `ApplySnapshotHeader(ceiling, cutover)`) + metrics + INFO log | **User-visible cold-start win.** | +| **B3** | `restoreSnapshotState` skip gate + `applyHeaderStateOnSkip(snapPath, tok.CRC32C)` orchestrating size + footer-vs-tokenCRC checks using `internal/raftengine/etcd`'s existing helpers + two-phase `SnapshotHeaderApplier` seam on `kvFSM` (`ParseSnapshotHeader(r io.Reader) (ceiling, cutover, err)` + pure `ApplySnapshotHeader(ceiling, cutover)`) + metrics + INFO log. Full-body CRC remains on the execute/full-restore path. | **User-visible cold-start win.** | | **B4** | Lower `HEALTH_TIMEOUT_SECONDS` default once production data shows steady-state skip rate ≥ 90 % | Tighter ceiling; the env override remains honoured. | Each of B2–B3 ships behind tests: @@ -927,8 +898,8 @@ Each of B2–B3 ships behind tests: test asserts that `applyHeaderStateOnSkip` sets `f.hlc.PhysicalCeiling()` **and** `f.restoredCutover` for both v1 and v2 snapshot headers — the ceiling+cutover are invariant under - the optimisation. Three additional CRC-corruption tests (round-6, - one per failure mode) inject the corruption and drive the skip path + the optimisation. Additional envelope/header tests inject corruption + and drive the skip path through `applyHeaderStateOnSkip` (either directly or via `restoreSnapshotState` with `fsmAlreadyAtIndex` returning true), asserting the specific typed error surfaces and that the FSM did @@ -937,16 +908,15 @@ Each of B2–B3 ships behind tests: `ErrFSMSnapshotTooSmall`. - Pair the file with a wrong-token CRC → `ErrFSMSnapshotTokenCRC`. - - Flip one body byte (post-header) → - `ErrFSMSnapshotFileCRC` (round-6's full-body CRC pass catches - this; without round-6 the skip would silently install state - from a corrupt file). + - Flip one body byte (post-header) and assert the skip path still + succeeds without scanning the body. Body corruption is caught by + the full restore path if that path is needed. An idle-cluster integration test runs a 3-node cluster with `ELASTICKV_RAFT_SNAPSHOT_COUNT=10` (overriding the default - `defaultSnapshotEvery = 10000` at `engine.go:93` so the scenario is + `defaultSnapshotEvery = 100000` at `engine.go:93` so the scenario is tractable — at default + `hlcRenewalInterval = 1 s` an idle period - would need ≥ 20 000 s). With the override, the test issues no data + would need ≥ 200 000 s). With the override, the test issues no data writes for `2 × 10 × hlcRenewalInterval = 20 s`, takes a snapshot, restarts a node, and asserts the skip fires — proving the codex round-3 P2 scenario is closed end-to-end through the @@ -1061,7 +1031,7 @@ subsection + B2 row + B2 test list) closes the gap by bumping successful snapshot persist, `LastAppliedIndex >= snapshot.Index` holds unconditionally, so the skip fires reliably on the next restart. Cost: one extra pebble `Batch.Commit` per snapshot persist -(~one extra fsync per `SnapshotCount` entries, default 10000) versus +(~one extra fsync per `SnapshotCount` entries, default 100000) versus Option A's continuous ~1 fsync/sec/group. Lesson: "rare" should be a quantitative claim against the actual @@ -1138,16 +1108,13 @@ and silently apply `ceiling=0, cutover=0` for a too-short file. Round 6 threads `tok.CRC32C` through `applyHeaderStateOnSkip` and `SnapshotHeaderApplier.ApplySnapshotHeaderFromFile`, and the §5 -pseudocode runs the same three-step verification before applying any -side-effect. Cost: one extra full-file read on the skip path — still -strictly cheaper than the restore path it replaces (the same read -happens during `restoreAndComputeCRC`, plus restore additionally -writes a temp Pebble database via `restoreBatchLoopInto`). - -A follow-up optimisation that persists the HLC ceiling + cutover as -durable meta keys (analogous to `metaAppliedIndex`) would let the -skip path elide the file read entirely; flagged under Open Questions -but not part of this proposal. +pseudocode originally ran the same three-step verification before +applying any side-effect. Production profiling later showed that the +extra full-file read is too expensive for multi-GiB snapshots during +restart. The final implementation keeps the cheap fail-closed checks +that matter for the skipped header side-effect (size, footer-vs-token, +header parse), and leaves full-body CRC on the full restore path where +body bytes are consumed. Lesson: when a §X seam takes over a side-effect previously gated by existing fail-closed checks, **inventory those checks first and @@ -1176,12 +1143,11 @@ or force a layering change. Round 7 splits `SnapshotHeaderApplier` into two methods: `ParseSnapshotHeader(r io.Reader) (ceiling, cutover, err)` and the -pure-assignment `ApplySnapshotHeader(ceiling, cutover)`. The CRC -verification orchestration stays in `internal/raftengine/etcd/wal_store.go` -where the helpers already live; the engine wraps the file in a crc32 -TeeReader and hands the reader to `setter.ParseSnapshotHeader`, which -calls the still-in-`kv`-package `ReadSnapshotHeader(*bufio.Reader)` -and drains. No package imports change; no helpers need to be +pure-assignment `ApplySnapshotHeader(ceiling, cutover)`. The envelope +checks stay in `internal/raftengine/etcd/wal_store.go` where the helpers +already live; the engine hands a payload-limited reader to +`setter.ParseSnapshotHeader`, which calls the still-in-`kv`-package +`ReadSnapshotHeader`. No package imports change; no helpers need to be exported. **P2 line 660 — `SetDurableAppliedIndex` honoured nosync mode.** @@ -1201,7 +1167,7 @@ Round 7 pins `pebbleStore.SetDurableAppliedIndex` to `pebble.Sync` unconditionally. The checkpoint must be at least as durable as the snapshot pointer that immediately follows; there is no raft log entry to replay it from after WAL compaction. Cost: +1 extra fsync per -snapshot persist (rare, default `SnapshotCount=10000`). +snapshot persist (rare, default `SnapshotCount=100000`). Lesson: - (P2 line 410) **Before pseudocoding cross-package helper use, diff --git a/docs/design/2026_06_12_proposed_scaling_roadmap.md b/docs/design/2026_06_12_proposed_scaling_roadmap.md index 76a179915..c12a194dc 100644 --- a/docs/design/2026_06_12_proposed_scaling_roadmap.md +++ b/docs/design/2026_06_12_proposed_scaling_roadmap.md @@ -135,7 +135,7 @@ Storage breakage at 1–10 TB/shard: - Snapshot transfer is a single-stream full-iter scan; at 1 TB it is hours, pins SSTs (blocks compaction reclamation), and any raft snapshot transfer serializes through it. -- WAL replay between `defaultSnapshotEvery = 10 000` raft entries +- WAL replay between `defaultSnapshotEvery = 100 000` raft entries can be multi-GiB; with `pebble.Sync` durability path the replay is tens of minutes. - Pebble L0CompactionThreshold / LBaseMaxBytes / compaction diff --git a/internal/backup/redis_list.go b/internal/backup/redis_list.go index ba9cc6244..77d09243c 100644 --- a/internal/backup/redis_list.go +++ b/internal/backup/redis_list.go @@ -34,7 +34,7 @@ import ( // item record for the popped // seq. The encoder therefore // skips claim keys entirely. -// - !lst|meta|d|... -> meta delta. The hash encoder +// - !lst|delta|... -> meta delta. The hash encoder // skips its analogous deltas // and treats !hs|fld| as the // source of truth; the list @@ -43,11 +43,16 @@ import ( // source of truth and the // delta arithmetic is not // replayed at backup time. +// - !lst|meta|d|... -> legacy meta delta. This overlaps +// with base !lst|meta| keys whose user key begins with d|, so routing +// checks both the key prefix and the 16-byte delta value shape before +// dropping it as a delta. const ( - ListMetaPrefix = "!lst|meta|" - ListItemPrefix = "!lst|itm|" - ListMetaDeltaPrefix = "!lst|meta|d|" - ListClaimPrefix = "!lst|claim|" + ListMetaPrefix = "!lst|meta|" + ListItemPrefix = "!lst|itm|" + ListMetaDeltaPrefix = "!lst|delta|" + LegacyListMetaDeltaPrefix = "!lst|meta|d|" + ListClaimPrefix = "!lst|claim|" // listMetaBinarySize is the legacy Head(8) + Tail(8) + Len(8) shape. listMetaBinarySize = 24 @@ -57,6 +62,8 @@ const ( // listSeqBytes is the fixed width of the trailing sortable-int64 // sequence number in an !lst|itm| key. listSeqBytes = 8 + + listMetaDeltaBinarySize = 16 ) // ErrRedisInvalidListMeta is returned when an !lst|meta| value is not @@ -86,15 +93,12 @@ type redisListState struct { // mismatch with the observed item count and register the user key so a // later !redis|ttl| record routes back to this list state. // -// !lst|meta|d|... delta keys share the !lst|meta| string -// prefix, so a snapshot dispatcher that routes by "starts with -// ListMetaPrefix" lands delta records here too. The hash encoder -// solved the analogous problem (Codex P1 round 14 PR #725) by silently -// skipping the delta family; we mirror that policy because !lst|itm| +// List deltas normally dispatch through HandleListMetaDelta, but keep +// the guard here so direct callers also skip the delta family. !lst|itm| // records are the source of truth for the restored list contents and // the delta arithmetic does not need to be replayed at backup time. func (r *RedisDB) HandleListMeta(key, value []byte) error { - if bytes.HasPrefix(key, []byte(ListMetaDeltaPrefix)) { + if isListMetaDeltaRecord(key, value) { return nil } userKey, ok := parseListMetaKey(key) @@ -147,7 +151,7 @@ func (r *RedisDB) HandleListItem(key, value []byte) error { // therefore reflect the post-POP state without any claim replay. func (r *RedisDB) HandleListClaim(_, _ []byte) error { return nil } -// HandleListMetaDelta accepts and discards one !lst|meta|d|... record. +// HandleListMetaDelta accepts and discards one !lst|delta|... record. // See HandleListMeta's docstring for the rationale; !lst|itm| is the // source of truth at backup time. func (r *RedisDB) HandleListMetaDelta(_, _ []byte) error { return nil } @@ -169,10 +173,9 @@ func (r *RedisDB) listState(userKey []byte) *redisListState { // parseListMetaKey strips !lst|meta| from a meta key and returns // (userKey, true). The list meta key shape is `prefix + userKey` with // no length prefix (mirror of store.ListMetaKey), so the trimmed -// remainder is the userKey verbatim. Delta keys (!lst|meta|d|...) -// share the meta string prefix and must be rejected here so a -// misrouted delta surfaces a parse failure rather than silent state -// corruption — analogous to parseHashMetaKey's delta guard. +// remainder is the userKey verbatim. Delta keys are rejected here so +// a misrouted delta surfaces a parse failure rather than silent state +// corruption. func parseListMetaKey(key []byte) ([]byte, bool) { if bytes.HasPrefix(key, []byte(ListMetaDeltaPrefix)) { return nil, false @@ -184,6 +187,14 @@ func parseListMetaKey(key []byte) ([]byte, bool) { return rest, true } +func isListMetaDeltaRecord(key, value []byte) bool { + if bytes.HasPrefix(key, []byte(ListMetaDeltaPrefix)) { + return true + } + return bytes.HasPrefix(key, []byte(LegacyListMetaDeltaPrefix)) && + len(value) == listMetaDeltaBinarySize +} + // parseListItemKey strips !lst|itm| and extracts (userKey, seq). The // list item key shape (mirror of store.ListItemKey) is // `prefix + userKey + sortableInt64(seq)`, with no userKey length diff --git a/internal/backup/redis_list_test.go b/internal/backup/redis_list_test.go index baa691c2d..1417400b6 100644 --- a/internal/backup/redis_list_test.go +++ b/internal/backup/redis_list_test.go @@ -40,7 +40,7 @@ func listItemKey(userKey string, seq int64) []byte { } // listMetaDeltaKey mirrors store.ListMetaDeltaKey: -// !lst|meta|d|. +// !lst|delta|. // The shape is irrelevant to the encoder (it skips deltas), but we // build a well-formed key here so the dispatcher integration test // would exercise the same byte sequence the live store emits. @@ -58,6 +58,20 @@ func listMetaDeltaKey(userKey string, commitTS uint64, seqInTxn uint32) []byte { return append(out, seq[:]...) } +func legacyListMetaDeltaKey(userKey string, commitTS uint64, seqInTxn uint32) []byte { + out := []byte(LegacyListMetaDeltaPrefix) + var l [4]byte + binary.BigEndian.PutUint32(l[:], uint32(len(userKey))) //nolint:gosec + out = append(out, l[:]...) + out = append(out, userKey...) + var ts [8]byte + binary.BigEndian.PutUint64(ts[:], commitTS) + out = append(out, ts[:]...) + var seq [4]byte + binary.BigEndian.PutUint32(seq[:], seqInTxn) + return append(out, seq[:]...) +} + // listClaimKey mirrors store.ListClaimKey: // !lst|claim|. func listClaimKey(userKey string, seq int64) []byte { @@ -187,6 +201,45 @@ func TestRedisDB_ListEmptyListStillEmitsFile(t *testing.T) { } } +func TestRedisDB_ListLegacyDeltaIsSkippedButDeltaLookingMetaIsPreserved(t *testing.T) { + t.Parallel() + db, root := newRedisDB(t) + if err := db.HandleListMeta(legacyListMetaDeltaKey("q", 10, 0), make([]byte, listMetaDeltaBinarySize)); err != nil { + t.Fatal(err) + } + userKey := string(deltaLookingListMetaUserKeyForBackup([]byte("real"), 11, 1)) + if err := db.HandleListMeta(listMetaKey(userKey), listMetaValue(0, 1)); err != nil { + t.Fatal(err) + } + if err := db.HandleListItem(listItemKey(userKey, 0), []byte("v")); err != nil { + t.Fatal(err) + } + if err := db.Finalize(); err != nil { + t.Fatal(err) + } + + got := readListJSON(t, filepath.Join(root, "redis", "db_0", "lists", EncodeSegment([]byte(userKey))+".json")) + assertListItems(t, got, []any{"v"}) + if _, err := os.Stat(filepath.Join(root, "redis", "db_0", "lists", "q.json")); !os.IsNotExist(err) { + t.Fatalf("legacy delta must not emit q.json: stat err=%v", err) + } +} + +func deltaLookingListMetaUserKeyForBackup(fakeUserKey []byte, commitTS uint64, seqInTxn uint32) []byte { + key := make([]byte, 0, len("d|")+4+len(fakeUserKey)+8+4) + key = append(key, "d|"...) + var lenPrefix [4]byte + binary.BigEndian.PutUint32(lenPrefix[:], uint32(len(fakeUserKey))) //nolint:gosec // test data is small. + key = append(key, lenPrefix[:]...) + key = append(key, fakeUserKey...) + var ts [8]byte + binary.BigEndian.PutUint64(ts[:], commitTS) + key = append(key, ts[:]...) + var seq [4]byte + binary.BigEndian.PutUint32(seq[:], seqInTxn) + return append(key, seq[:]...) +} + // TestRedisDB_ListTTLInlinedFromScanIndex pins that !redis|ttl| records // for a list user key fold into the list's JSON `expire_at_ms` rather // than landing in a separate sidecar (the strings/HLL pattern). A @@ -275,10 +328,9 @@ func TestRedisDB_ListBinaryItemUsesBase64Envelope(t *testing.T) { } // TestRedisDB_ListHandleListMetaSkipsDeltaKey pins that the -// !lst|meta|d|... family is silently skipped by HandleListMeta. Without +// !lst|delta|... family is silently skipped by HandleListMeta. Without // this, parsing the delta's userKeyLen prefix as the start of a -// userKey would corrupt the lists map. Mirrors the hash delta-key -// guard (Codex P1 round 14 PR #725). +// userKey would corrupt the lists map. Mirrors the hash delta-key guard. func TestRedisDB_ListHandleListMetaSkipsDeltaKey(t *testing.T) { t.Parallel() db, _ := newRedisDB(t) @@ -353,7 +405,7 @@ func TestRedisDB_ListRejectsMalformedMetaValueLength(t *testing.T) { // firing the declared-vs-observed length mismatch warning (because // metaSeen=false means we have no "declared" baseline to compare // against). Mirrors the items-as-source-of-truth contract that -// makes the !lst|meta|d| delta family safe to skip. +// makes the !lst|delta| delta family safe to skip. func TestRedisDB_ListItemsWithoutMetaStillEmitsFile(t *testing.T) { t.Parallel() db, root := newRedisDB(t) diff --git a/internal/raftengine/engine.go b/internal/raftengine/engine.go index 7b010a3b4..7e51ceee7 100644 --- a/internal/raftengine/engine.go +++ b/internal/raftengine/engine.go @@ -251,6 +251,12 @@ type Lifecycle interface { Err() error } +// StartupBarrier is an optional capability for engines that can report when +// their local startup replay has drained far enough for user traffic. +type StartupBarrier interface { + WaitStarted(ctx context.Context) error +} + type Admin interface { LeaderView StatusReader diff --git a/internal/raftengine/etcd/dispatch_report_test.go b/internal/raftengine/etcd/dispatch_report_test.go index d0ba97049..63c7fb631 100644 --- a/internal/raftengine/etcd/dispatch_report_test.go +++ b/internal/raftengine/etcd/dispatch_report_test.go @@ -33,6 +33,96 @@ func TestPostDispatchReport_DeliversWhenChannelHasSpace(t *testing.T) { } } +func TestReportSuccessfulDispatchReportsSnapshotFinish(t *testing.T) { + t.Parallel() + e := &Engine{ + dispatchReportCh: make(chan dispatchReport, 1), + closeCh: make(chan struct{}), + } + + e.reportSuccessfulDispatch(raftpb.Message{Type: messageTypePtr(raftpb.MsgSnap), To: uint64Ptr(2)}) + + select { + case got := <-e.dispatchReportCh: + require.Equal(t, dispatchReport{to: 2, msgType: raftpb.MsgSnap, snapshotFinish: true}, got) + default: + t.Fatal("expected successful MsgSnap dispatch to report SnapshotFinish input") + } +} + +func TestReportSuccessfulDispatchWaitsForSnapshotFinishSlot(t *testing.T) { + t.Parallel() + e := &Engine{ + dispatchReportCh: make(chan dispatchReport, 1), + closeCh: make(chan struct{}), + } + blockingReport := dispatchReport{to: 1, msgType: raftpb.MsgApp} + e.dispatchReportCh <- blockingReport + + done := make(chan struct{}) + go func() { + e.reportSuccessfulDispatch(raftpb.Message{Type: messageTypePtr(raftpb.MsgSnap), To: uint64Ptr(2)}) + close(done) + }() + + select { + case <-done: + t.Fatal("snapshot finish report returned while dispatchReportCh was full") + case <-time.After(50 * time.Millisecond): + } + + require.Equal(t, blockingReport, <-e.dispatchReportCh) + select { + case <-done: + case <-time.After(testDispatchReportTimeout): + t.Fatal("snapshot finish report did not complete after dispatchReportCh had space") + } + select { + case got := <-e.dispatchReportCh: + require.Equal(t, dispatchReport{to: 2, msgType: raftpb.MsgSnap, snapshotFinish: true}, got) + default: + t.Fatal("expected reliable successful MsgSnap dispatch report") + } +} + +func TestReportSuccessfulDispatchAbortsOnCloseWhenReportFull(t *testing.T) { + t.Parallel() + e := &Engine{ + dispatchReportCh: make(chan dispatchReport, 1), + closeCh: make(chan struct{}), + } + e.dispatchReportCh <- dispatchReport{to: 1, msgType: raftpb.MsgApp} + close(e.closeCh) + + done := make(chan struct{}) + go func() { + e.reportSuccessfulDispatch(raftpb.Message{Type: messageTypePtr(raftpb.MsgSnap), To: uint64Ptr(2)}) + close(done) + }() + + select { + case <-done: + case <-time.After(testDispatchReportTimeout): + t.Fatal("snapshot finish report did not abort when closeCh was signalled") + } +} + +func TestReportSuccessfulDispatchIgnoresRegularMessage(t *testing.T) { + t.Parallel() + e := &Engine{ + dispatchReportCh: make(chan dispatchReport, 1), + closeCh: make(chan struct{}), + } + + e.reportSuccessfulDispatch(raftpb.Message{Type: messageTypePtr(raftpb.MsgApp), To: uint64Ptr(2)}) + + select { + case got := <-e.dispatchReportCh: + t.Fatalf("unexpected dispatch report for regular message: %+v", got) + default: + } +} + // TestPostDispatchReport_DropsWhenChannelFull asserts the non-blocking // contract: dispatch workers must not stall because the event loop is busy. // The worst case is an eventually-consistent gap that raft will fix on the @@ -86,7 +176,7 @@ func TestPostDispatchReport_AbortsOnClose(t *testing.T) { } } -func TestEnqueueDispatchReportsDroppedSnapshotWhenLaneFull(t *testing.T) { +func TestEnqueueDispatchDefersDroppedSnapshotReportWhenLaneFull(t *testing.T) { t.Parallel() snapshotCh := make(chan dispatchRequest, 1) snapshotCh <- dispatchRequest{msg: raftpb.Message{Type: messageTypePtr(raftpb.MsgSnap), To: uint64Ptr(2)}} @@ -105,10 +195,10 @@ func TestEnqueueDispatchReportsDroppedSnapshotWhenLaneFull(t *testing.T) { require.Equal(t, uint64(1), e.DispatchDropCount()) select { case got := <-e.dispatchReportCh: - require.Equal(t, dispatchReport{to: 2, msgType: raftpb.MsgSnap}, got) + t.Fatalf("dropped MsgSnap should defer without queueing: %+v", got) default: - t.Fatal("expected dropped MsgSnap to report SnapshotFailure input") } + require.Equal(t, []dispatchReport{{to: 2, msgType: raftpb.MsgSnap}}, e.deferredReadyDispatchReports) } func TestEnqueueDispatchReportsDroppedRegularMessageWhenLaneFull(t *testing.T) { diff --git a/internal/raftengine/etcd/engine.go b/internal/raftengine/etcd/engine.go index 7a8680909..461110ea3 100644 --- a/internal/raftengine/etcd/engine.go +++ b/internal/raftengine/etcd/engine.go @@ -60,9 +60,10 @@ const ( // from shrinking the shared inbound queue enough to drop heartbeats. minInboundQueueCapacity = 128 // priorityStepQueueCapacity is the inbound control-plane queue size. - // Heartbeats, votes, read-index responses, and timeout-now messages are - // tiny but time-sensitive; keeping them off the bulk stepCh prevents a - // MsgApp burst from forcing followers into avoidable elections. + // Heartbeats, votes, read-index responses, timeout-now messages, and + // received snapshot tokens are tiny but time-sensitive; keeping them off the + // bulk stepCh prevents a MsgApp burst from forcing followers into avoidable + // elections or rejecting a catch-up snapshot after its payload was streamed. priorityStepQueueCapacity = 1024 // priorityStepBurstLimit bounds consecutive non-blocking priority drains // so a sustained control-message stream cannot starve Tick, proposals, or @@ -71,8 +72,8 @@ const ( priorityStepBurstLimit = 64 // defaultHeartbeatBufPerPeer is the capacity of the priority dispatch channel. // It carries low-frequency control traffic: heartbeats, votes, read-index, - // leader-transfer, and their corresponding response messages - // (MsgHeartbeatResp, MsgReadIndexResp, MsgVoteResp, MsgPreVoteResp). + // leader-transfer, and their corresponding response messages except for + // MsgHeartbeatResp, which uses its own coalescing response lane. // MsgAppResp is intentionally kept in the normal channel: followers — the // only senders of MsgAppResp — do not send MsgApp, so there is no // head-of-line blocking risk there. @@ -84,8 +85,21 @@ const ( // upside is that a ~5 s transient pause (election-timeout scale) // no longer drops heartbeats and forces the peers' lease to expire. defaultHeartbeatBufPerPeer = 512 + // defaultReadIndexRespBufPerPeer sizes the dedicated follower-to-leader + // ReadIndex heartbeat response lane. etcd/raft encodes ReadIndex + // completions as MsgHeartbeatResp messages with Context set; unlike plain + // heartbeat acks, each context is tied to a caller waiting in handleRead + // and must not be coalesced or dropped behind superseded empty acks. + defaultReadIndexRespBufPerPeer = 512 + // defaultHeartbeatRespBufPerPeer sizes the dedicated follower-to-leader + // heartbeat response lane. When a follower is receiving a large snapshot, + // the leader may continue to send heartbeats while the follower's outbound + // transport is slow. Heartbeat responses are superseded by newer heartbeat + // responses for the same peer, so enqueueDispatchMessage coalesces this lane + // instead of reporting a dropped raft message and starving the leader lease. + defaultHeartbeatRespBufPerPeer = 128 // defaultSnapshotLaneBufPerPeer sizes the per-peer MsgSnap lane when the - // 4-lane dispatcher mode is enabled (see ELASTICKV_RAFT_DISPATCHER_LANES). + // opt-in multi-lane dispatcher is enabled (see ELASTICKV_RAFT_DISPATCHER_LANES). // MsgSnap is rare and bulky; 4 is enough to absorb a retry or two without // holding up MsgApp replication behind a multi-MiB payload. defaultSnapshotLaneBufPerPeer = 4 @@ -93,11 +107,11 @@ const ( // types not classified as heartbeat/replication/snapshot (e.g. surprise // locally-addressed control types). Small buffer: traffic volume is tiny. defaultOtherLaneBufPerPeer = 16 - // dispatcherLanesEnvVar toggles the 4-lane dispatcher (heartbeat / - // replication / snapshot / other). When unset or "0", the legacy - // 2-lane layout (heartbeat + normal) is used. Opt-in by design: the - // raft hot path is high blast radius and a regression here can cause - // cluster-wide elections. + // dispatcherLanesEnvVar toggles the multi-lane dispatcher (heartbeat / + // heartbeatResp / replication / snapshot / other). When unset or "0", the + // legacy 3-lane layout (heartbeat + heartbeatResp + normal) is used. + // Opt-in by design: the raft hot path is high blast radius and a regression + // here can cause cluster-wide elections. dispatcherLanesEnvVar = "ELASTICKV_RAFT_DISPATCHER_LANES" // preVoteEnvVar permits an operator to temporarily disable raft // pre-vote during manual quorum recovery. It defaults to enabled. @@ -105,10 +119,10 @@ const ( // defaultSnapshotEvery is the fallback trigger threshold: take an FSM // snapshot once the applied index has advanced this many entries past // the last snapshot's index. etcd/raft itself uses 10_000 as a default, - // but with fat proposal payloads (e.g. Lua scripts) this can produce a - // multi-GiB WAL between snapshots. Operators can lower via + // but multi-GiB FSM snapshots can take tens of seconds to persist and + // contend with the Raft hot path. Operators can lower or raise via // ELASTICKV_RAFT_SNAPSHOT_COUNT without a rebuild. - defaultSnapshotEvery = 10_000 + defaultSnapshotEvery = 100_000 snapshotEveryEnvVar = "ELASTICKV_RAFT_SNAPSHOT_COUNT" defaultSnapshotQueueSize = 1 defaultAdminPollInterval = 10 * time.Millisecond @@ -323,16 +337,17 @@ type Engine struct { nextRequestID atomic.Uint64 - proposeCh chan proposalRequest - readCh chan readRequest - adminCh chan adminRequest - stepCh chan raftpb.Message - priorityStepCh chan raftpb.Message - priorityStepBurst int - dispatchReportCh chan dispatchReport - peerDispatchers map[uint64]*peerQueues - perPeerQueueSize int - // dispatcherLanesEnabled toggles the 4-lane dispatcher layout. Captured + proposeCh chan proposalRequest + readCh chan readRequest + adminCh chan adminRequest + stepCh chan raftpb.Message + priorityStepCh chan raftpb.Message + priorityStepBurst int + dispatchReportCh chan dispatchReport + deferredReadyDispatchReports []dispatchReport + peerDispatchers map[uint64]*peerQueues + perPeerQueueSize int + // dispatcherLanesEnabled toggles the opt-in multi-lane dispatcher layout. Captured // once at Open from ELASTICKV_RAFT_DISPATCHER_LANES so the run-time code // path is branch-free per message and does not need to re-read env vars. dispatcherLanesEnabled bool @@ -355,7 +370,10 @@ type Engine struct { snapshotStopCh chan struct{} closeCh chan struct{} doneCh chan struct{} - startedCh chan struct{} + // startedCh is closed after startup has drained committed Ready entries. + // Multi-node Open must return before this so callers can register the + // transport listener; service startup waits through WaitStarted instead. + startedCh chan struct{} leaderReady chan struct{} leaderOnce sync.Once @@ -571,22 +589,25 @@ type dispatchRequest struct { // peerQueues holds separate dispatch channels per peer so that heartbeats // are never blocked behind large log-entry RPCs. // -// Legacy 2-lane layout (default): heartbeat + normal. +// Legacy 4-lane layout (default): heartbeat + heartbeatResp + +// readIndexResp + normal. // -// 4-lane layout (opt-in via ELASTICKV_RAFT_DISPATCHER_LANES=1): heartbeat + -// replication (MsgApp/MsgAppResp) + snapshot (MsgSnap) + other. Each lane -// gets its own goroutine so a bulky MsgSnap transfer cannot stall MsgApp -// replication and vice versa. Per-peer ordering within a given message type -// is preserved because a single peer's MsgApp stream all share one lane and -// one worker. +// 6-lane layout (opt-in via ELASTICKV_RAFT_DISPATCHER_LANES=1): heartbeat + +// heartbeatResp + readIndexResp + replication (MsgApp/MsgAppResp) + snapshot +// (MsgSnap) + other. Each lane gets its own goroutine so a bulky MsgSnap +// transfer cannot stall MsgApp replication and vice versa. Per-peer ordering +// within a given message type is preserved because a single peer's MsgApp +// stream all share one lane and one worker. type peerQueues struct { - normal chan dispatchRequest - heartbeat chan dispatchRequest - replication chan dispatchRequest // 4-lane mode only; nil otherwise - snapshot chan dispatchRequest // 4-lane mode only; nil otherwise - other chan dispatchRequest // 4-lane mode only; nil otherwise - ctx context.Context - cancel context.CancelFunc + normal chan dispatchRequest + heartbeat chan dispatchRequest + heartbeatResp chan dispatchRequest + readIndexResp chan dispatchRequest + replication chan dispatchRequest // 6-lane mode only; nil otherwise + snapshot chan dispatchRequest // 6-lane mode only; nil otherwise + other chan dispatchRequest // 6-lane mode only; nil otherwise + ctx context.Context + cancel context.CancelFunc } type preparedOpenState struct { @@ -960,17 +981,43 @@ func newRawNode(cfg OpenConfig, storage *etcdraft.MemoryStorage, applied uint64) } func waitForOpen(ctx context.Context, engine *Engine, waitForLeader bool) (*Engine, error) { + if err := waitForOpenSignal(ctx, engine, engine.startedCh); err != nil { + return nil, err + } + if !waitForLeader { + return engine, nil + } + if err := waitForOpenSignal(ctx, engine, engine.leaderReady); err != nil { + return nil, err + } + return engine, nil +} + +func (e *Engine) WaitStarted(ctx context.Context) error { + if e == nil { + return errors.WithStack(errNilEngine) + } + return waitForEngineSignal(ctx, e, e.startedCh, false) +} + +func waitForOpenSignal(ctx context.Context, engine *Engine, ready <-chan struct{}) error { + return waitForEngineSignal(ctx, engine, ready, true) +} + +func waitForEngineSignal(ctx context.Context, engine *Engine, ready <-chan struct{}, closeOnCancel bool) error { select { case <-ctx.Done(): - _ = engine.Close() - return nil, errors.WithStack(ctx.Err()) - case <-engine.openReady(waitForLeader): - return engine, nil + if closeOnCancel { + _ = engine.Close() + } + return errors.WithStack(ctx.Err()) + case <-ready: + return nil case <-engine.doneCh: if err := engine.currentError(); err != nil { - return nil, err + return err } - return nil, errors.WithStack(errClosed) + return errors.WithStack(errClosed) } } @@ -1283,12 +1330,11 @@ func (e *Engine) recordDispatchErrorCode(code string) uint64 { } // StepQueueFullCount returns the total number of inbound raft messages -// that could not be enqueued into the selected inbound step queue -// because the channel was at capacity. This is the "etcd raft inbound -// step queue is full" signal from the task description: a spike -// indicates the local raft loop is starved, usually by something -// blocking the apply path such as -// the pre-#560 rawKeyTypeAt seek storm. +// that found the selected inbound step queue at capacity. Blocking inbound +// message classes wait for space after incrementing this counter; best-effort +// classes still return errStepQueueFull. A spike indicates the local raft loop +// is starved, usually by something blocking the apply path such as the +// pre-#560 rawKeyTypeAt seek storm. func (e *Engine) StepQueueFullCount() uint64 { if e == nil { return 0 @@ -1780,6 +1826,10 @@ func (e *Engine) run() { e.fail(err) return } + if err := e.persistStartupAppliedIndex(); err != nil { + e.fail(err) + return + } e.markStarted() for { @@ -1892,13 +1942,14 @@ func (e *Engine) tryReceivePriorityStep() (raftpb.Message, bool) { } } -// dispatchReport is posted by the dispatch workers when a transport send -// to a peer fails; the engine goroutine drains these and informs etcd/raft -// via rawNode so follower Progress leaves StateReplicate / StateSnapshot on -// unreachable peers and does not silently stall. +// dispatchReport is posted by the dispatch workers after a transport send +// completes or fails; the engine goroutine drains these and informs etcd/raft +// via rawNode so follower Progress leaves StateReplicate / StateSnapshot and +// does not silently stall. type dispatchReport struct { - to uint64 - msgType raftpb.MessageType + to uint64 + msgType raftpb.MessageType + snapshotFinish bool } func (e *Engine) handleDispatchReport(report dispatchReport) { @@ -1911,19 +1962,24 @@ func (e *Engine) handleDispatchReport(report dispatchReport) { // peer from StateReplicate to StateProbe so the next heartbeat response // drives a fresh sendAppend attempt. if report.msgType == raftpb.MsgSnap { - e.rawNode.ReportSnapshot(report.to, etcdraft.SnapshotFailure) + status := etcdraft.SnapshotFailure + if report.snapshotFinish { + status = etcdraft.SnapshotFinish + } + e.rawNode.ReportSnapshot(report.to, status) return } e.rawNode.ReportUnreachable(report.to) } -// postDispatchReport delivers a dispatch failure to the event loop without -// blocking the caller. Dispatch workers use it for transport failures, and the -// event loop uses it for local queue drops before transport. If the channel is -// full (unlikely — the buffer is sized to MaxInflightMsg), the report is -// dropped and logged; this is acceptable because raft will retry on the next -// tick and we only need eventual consistency between transport state and -// Progress state. +// postDispatchReport delivers a dispatch outcome to the event loop without +// blocking the caller. Dispatch workers use it for transport completion or +// failures, and the event loop uses it for local queue drops before transport. +// If the channel is full (unlikely — the buffer is sized to MaxInflightMsg), +// the report is dropped and logged; this is acceptable because raft will retry +// on the next tick and we only need eventual consistency between transport +// state and Progress state. MsgSnap reports are the exception; see +// postReliableDispatchReport. func (e *Engine) postDispatchReport(report dispatchReport) { select { case e.dispatchReportCh <- report: @@ -1936,6 +1992,21 @@ func (e *Engine) postDispatchReport(report dispatchReport) { } } +// postReliableDispatchReport delivers a dispatch outcome that must not be +// dropped. SnapshotFinish and SnapshotFailure are not eventually consistent +// with ordinary unreachable reports: if either is lost, raft can keep the +// follower's Progress in StateSnapshot after the follower accepted or missed +// the snapshot. +func (e *Engine) postReliableDispatchReport(report dispatchReport) { + if e.dispatchReportCh == nil { + return + } + select { + case e.dispatchReportCh <- report: + case <-e.closeCh: + } +} + func (e *Engine) handleProposal(req proposalRequest) { if err := contextErr(req.ctx); err != nil { req.done <- proposalResult{err: err} @@ -2220,6 +2291,7 @@ func (e *Engine) drainReady() error { e.releaseProtectedReceivedFSMSnapshotsUpTo(e.appliedIndex.Load()) e.handleReadStates(rd.ReadStates) e.rawNode.Advance(rd) + e.flushDeferredDispatchReports() if err := e.maybePersistLocalSnapshot(); err != nil { return err } @@ -2260,6 +2332,9 @@ func (e *Engine) persistReadyWithSnapshotLocked(rd etcdraft.Ready) error { if err := persistReadyToWAL(e.persist, rd); err != nil { return err } + if err := e.persistReceivedSnapshotAppliedIndex(rd.Snapshot); err != nil { + return err + } e.releaseProtectedReceivedFSMSnapshotsUpToLocked(snapshotIndex(rd.Snapshot)) return nil } @@ -2293,6 +2368,7 @@ func (e *Engine) handleStep(msg raftpb.Message) { commitBeforeStep := e.rawNode.Status().GetCommit() if err := e.rawNode.Step(&msg); err != nil { if errors.Is(err, etcdraft.ErrStepPeerNotFound) { + e.removeReceivedFSMSnapshotToken(msg) e.unprotectReceivedFSMSnapshotToken(msg) return } @@ -2300,9 +2376,11 @@ func (e *Engine) handleStep(msg raftpb.Message) { return } if e.unprotectReceivedFSMSnapshotTokenIfCommitted(msg, commitBeforeStep) { + e.removeReceivedFSMSnapshotToken(msg) return } if !e.rawNode.HasReady() { + e.removeReceivedFSMSnapshotToken(msg) e.unprotectReceivedFSMSnapshotToken(msg) return } @@ -2419,11 +2497,14 @@ func (e *Engine) enqueueDispatchMessage(msg raftpb.Message) error { e.recordDroppedDispatch(msg) return nil } - ch := e.selectDispatchLane(pd, msg.GetType()) + ch := e.selectDispatchLaneForMessage(pd, msg) // Avoid the expensive deep-clone in prepareDispatchRequest when the channel // is already full. The len/cap check is safe here because this function is // only ever called from the single engine event-loop goroutine. if len(ch) >= cap(ch) { + if msg.GetType() == raftpb.MsgHeartbeatResp && coalesceHeartbeatResp(ch, msg) { + return nil + } e.recordDroppedDispatch(msg) return nil } @@ -2438,12 +2519,65 @@ func (e *Engine) enqueueDispatchMessage(msg raftpb.Message) error { } } +func coalesceHeartbeatResp(ch chan dispatchRequest, msg raftpb.Message) bool { + if ch == nil || cap(ch) == 0 { + return false + } + n := len(ch) + if n == 0 { + return false + } + + drained := make([]dispatchRequest, 0, n) + replaced := false + for i := 0; i < n; i++ { + req, ok := tryReceiveDispatchRequest(ch) + if !ok { + break + } + if !replaced && req.msg.GetType() == raftpb.MsgHeartbeatResp && len(req.msg.GetContext()) == 0 { + closeDispatchRequest(req) + drained = append(drained, prepareDispatchRequest(msg)) + replaced = true + continue + } + drained = append(drained, req) + } + for _, req := range drained { + restoreDispatchRequest(ch, req) + } + return replaced +} + +func tryReceiveDispatchRequest(ch chan dispatchRequest) (dispatchRequest, bool) { + select { + case req := <-ch: + return req, true + default: + return dispatchRequest{}, false + } +} + +func restoreDispatchRequest(ch chan dispatchRequest, req dispatchRequest) { + select { + case ch <- req: + default: + closeDispatchRequest(req) + } +} + +func closeDispatchRequest(req dispatchRequest) { + if err := req.Close(); err != nil { + slog.Error("etcd raft dispatch: failed to close request", "err", err) + } +} + // isPriorityMsg returns true for small, low-frequency control messages that // must not be queued behind large MsgApp payloads in the normal channel. // MsgAppResp is intentionally excluded: it is sent by followers, which never // send MsgApp, so it faces no head-of-line blocking in the normal channel. -// Keeping it out of the priority queue preserves the low-frequency invariant -// that justifies defaultHeartbeatBufPerPeer = 64. +// Keeping MsgHeartbeatResp on its own coalescing lane preserves the +// low-frequency invariant that justifies defaultHeartbeatBufPerPeer. func isPriorityMsg(t raftpb.MessageType) bool { return t == raftpb.MsgHeartbeat || t == raftpb.MsgHeartbeatResp || t == raftpb.MsgReadIndex || t == raftpb.MsgReadIndexResp || @@ -2452,12 +2586,32 @@ func isPriorityMsg(t raftpb.MessageType) bool { t == raftpb.MsgTimeoutNow } +func isInboundPriorityMsg(t raftpb.MessageType) bool { + return isPriorityMsg(t) || t == raftpb.MsgSnap +} + +func isBlockingInboundStepMsg(t raftpb.MessageType) bool { + return t == raftpb.MsgSnap || t == raftpb.MsgApp || t == raftpb.MsgAppResp +} + // selectDispatchLane picks the per-peer channel for msgType. In the legacy -// 2-lane layout it returns pd.heartbeat for priority control traffic and -// pd.normal for everything else. In the 4-lane layout it additionally -// partitions the non-heartbeat traffic so that MsgApp/MsgAppResp and MsgSnap -// do not share a goroutine and cannot block each other. +// layout it returns pd.heartbeatResp for MsgHeartbeatResp, pd.heartbeat for +// other priority control traffic, and pd.normal for everything else. In the +// opt-in multi-lane layout it additionally partitions the non-heartbeat traffic +// so that MsgApp/MsgAppResp and MsgSnap do not share a goroutine and cannot +// block each other. +func (e *Engine) selectDispatchLaneForMessage(pd *peerQueues, msg raftpb.Message) chan dispatchRequest { + msgType := msg.GetType() + if msgType == raftpb.MsgHeartbeatResp && len(msg.GetContext()) > 0 && pd.readIndexResp != nil { + return pd.readIndexResp + } + return e.selectDispatchLane(pd, msgType) +} + func (e *Engine) selectDispatchLane(pd *peerQueues, msgType raftpb.MessageType) chan dispatchRequest { + if msgType == raftpb.MsgHeartbeatResp && pd.heartbeatResp != nil { + return pd.heartbeatResp + } // Priority control traffic (heartbeats, votes, read-index, timeout-now) // always rides the heartbeat lane in both layouts so it keeps its // low-latency treatment and is never stuck behind MsgApp payloads. @@ -2525,15 +2679,6 @@ func (e *Engine) applyReadySnapshotLocked(snapshot *raftpb.Snapshot) error { if err != nil { return errors.Wrapf(err, "decode snapshot token index=%d", snapshot.GetMetadata().GetIndex()) } - // B3/follow-up: also call SetDurableAppliedIndex(tok.Index) here - // after Restore so peer-after-InstallSnapshot populates the meta - // key. The local-snapshot persist path already bumps the live - // store (engine.persistLocalSnapshotPayload), but the receiving - // node's restored store inherits the pre-bump value embedded in - // the snapshot artifact. Design Non-Goals § - // docs/design/2026_06_02_implemented_idempotent_snapshot_restore.md#non-goals - // scopes this out of Branch 2; see PR #915 round-4/5 codex P2 on - // engine.go:4077 for the rationale. if err := openAndRestoreFSMSnapshot(e.fsm, fsmSnapPath(e.fsmSnapDir, tok.Index), tok.CRC32C); err != nil { return errors.Wrapf(err, "restore fsm snapshot file index=%d crc=%08x", tok.Index, tok.CRC32C) } @@ -2543,7 +2688,6 @@ func (e *Engine) applyReadySnapshotLocked(snapshot *raftpb.Snapshot) error { return errors.Wrapf(err, "restore fsm from legacy snapshot payload index=%d", snapshot.GetMetadata().GetIndex()) } } - if err := e.storage.ApplySnapshot(snapshot); err != nil { return errors.Wrapf(err, "apply snapshot to raft storage index=%d term=%d", snapshot.GetMetadata().GetIndex(), snapshot.GetMetadata().GetTerm()) @@ -3327,6 +3471,9 @@ func (e *Engine) protectReceivedFSMSnapshot(index uint64) bool { if e.protectedReceivedFSMSnaps == nil { e.protectedReceivedFSMSnaps = make(map[uint64]int, 1) } + if e.protectedReceivedFSMSnaps[index] > 0 { + return false + } e.protectedReceivedFSMSnaps[index]++ return true } @@ -3450,6 +3597,7 @@ func (e *Engine) releaseIgnoredReceivedFSMSnapshotSteps(rd etcdraft.Ready) { if index == readySnapshotIndex { continue } + e.removeReceivedFSMSnapshotIndex(index) for i := 0; i < count; i++ { e.unprotectReceivedFSMSnapshot(index) } @@ -3464,6 +3612,23 @@ func (e *Engine) unprotectReceivedFSMSnapshotToken(msg raftpb.Message) { e.unprotectReceivedFSMSnapshot(index) } +func (e *Engine) removeReceivedFSMSnapshotToken(msg raftpb.Message) { + index, ok := receivedFSMSnapshotTokenIndex(msg) + if !ok { + return + } + e.removeReceivedFSMSnapshotIndex(index) +} + +func (e *Engine) removeReceivedFSMSnapshotIndex(index uint64) { + if e == nil || e.fsmSnapDir == "" || index == 0 { + return + } + e.snapshotMu.Lock() + defer e.snapshotMu.Unlock() + removeWithWarn(fsmSnapPath(e.fsmSnapDir, index), "ignored received fsm snapshot") +} + func receivedFSMSnapshotTokenIndex(msg raftpb.Message) (uint64, bool) { if msg.GetType() != raftpb.MsgSnap || msg.GetSnapshot() == nil || !isSnapshotToken(msg.GetSnapshot().GetData()) { return 0, false @@ -3555,29 +3720,60 @@ func (e *Engine) createConfigSnapshot(index uint64, confState raftpb.ConfState, } } +// setDurableAppliedIndex pins the FSM's durable applied index to a +// known raft log index. FSMs that do not expose +// raftengine.AppliedIndexWriter silently no-op; the skip optimisation +// falls back to full restore for them (legacy test fakes, in-memory +// backends). +func (e *Engine) setDurableAppliedIndex(index uint64) error { + w, ok := e.fsm.(raftengine.AppliedIndexWriter) + if !ok { + return nil + } + return errors.WithStack(w.SetDurableAppliedIndex(index)) +} + // bumpDurableAppliedIndexBeforeSave pins the FSM's durable applied // index to `index` BEFORE the engine calls persist.SaveSnap, so a // successful snapshot persist always implies LastAppliedIndex >= // snap.Metadata.Index — closes the HLC-lease-only / encryption-only // fallback (PR #910 design §6). // -// FSMs that do not expose raftengine.AppliedIndexWriter silently -// no-op; the skip optimisation falls back to full restore for them -// (legacy test fakes, in-memory backends). pebble.Sync is forced on -// the writer side regardless of ELASTICKV_FSM_SYNC_MODE — once -// persist.SaveSnap returns, WAL compaction discards every log entry -// at or before snap.Metadata.Index, so there is no source to replay -// the meta key bump from. +// pebble.Sync is forced on the writer side regardless of +// ELASTICKV_FSM_SYNC_MODE — once persist.SaveSnap returns, WAL +// compaction discards every log entry at or before snap.Metadata.Index, +// so there is no source to replay the meta key bump from. // // Used by BOTH snapshot persist sites: persistCreatedSnapshot (this // file) and e.persistLocalSnapshotPayload (the steady-state // SnapshotCount-triggered hot path). func (e *Engine) bumpDurableAppliedIndexBeforeSave(index uint64) error { - w, ok := e.fsm.(raftengine.AppliedIndexWriter) - if !ok { + return e.setDurableAppliedIndex(index) +} + +func (e *Engine) persistStartupAppliedIndex() error { + if e.applied == 0 { return nil } - return errors.WithStack(w.SetDurableAppliedIndex(index)) + return e.setDurableAppliedIndex(e.applied) +} + +// persistReceivedSnapshotAppliedIndex runs only after persistReadyToWAL has +// saved the incoming raft snapshot. A crash before this point must fall back to +// a full FSM restore instead of letting the FSM meta index get ahead of the +// local durable raft snapshot. +func (e *Engine) persistReceivedSnapshotAppliedIndex(snapshot *raftpb.Snapshot) error { + if etcdraft.IsEmptySnap(snapshot) { + return nil + } + index := snapshot.GetMetadata().GetIndex() + if index == 0 { + return nil + } + if err := e.setDurableAppliedIndex(index); err != nil { + return errors.Wrapf(err, "persist durable applied index from snapshot index=%d", index) + } + return nil } func (e *Engine) persistCreatedSnapshot(snap raftpb.Snapshot) error { @@ -3714,13 +3910,6 @@ func (e *Engine) markStarted() { e.startOnce.Do(func() { close(e.startedCh) }) } -func (e *Engine) openReady(waitForLeader bool) <-chan struct{} { - if waitForLeader { - return e.leaderReady - } - return e.startedCh -} - func (e *Engine) requestShutdown() { e.closeOnce.Do(func() { close(e.closeCh) @@ -4395,9 +4584,10 @@ func maxAppliedIndex(snapshot raftpb.Snapshot) uint64 { } func (e *Engine) enqueueStep(ctx context.Context, msg raftpb.Message) error { - ch := e.stepCh - if isPriorityMsg(msg.GetType()) && e.priorityStepCh != nil { - ch = e.priorityStepCh + ch := e.stepChannelFor(msg.GetType()) + + if isBlockingInboundStepMsg(msg.GetType()) { + return e.enqueueBlockingStep(ctx, ch, msg) } select { @@ -4413,6 +4603,31 @@ func (e *Engine) enqueueStep(ctx context.Context, msg raftpb.Message) error { } } +func (e *Engine) stepChannelFor(msgType raftpb.MessageType) chan raftpb.Message { + if isInboundPriorityMsg(msgType) && e.priorityStepCh != nil { + return e.priorityStepCh + } + return e.stepCh +} + +func (e *Engine) enqueueBlockingStep(ctx context.Context, ch chan raftpb.Message, msg raftpb.Message) error { + select { + case ch <- msg: + return nil + default: + e.stepQueueFullCount.Add(1) + } + + select { + case <-ctx.Done(): + return errors.WithStack(ctx.Err()) + case <-e.doneCh: + return e.currentErrorOrClosed() + case ch <- msg: + return nil + } +} + func (e *Engine) handleTransportMessage(ctx context.Context, msg raftpb.Message) error { select { case <-ctx.Done(): @@ -4465,24 +4680,27 @@ func (e *Engine) startPeerDispatcher(nodeID uint64) { } ctx, cancel := context.WithCancel(baseCtx) pd := &peerQueues{ - heartbeat: make(chan dispatchRequest, defaultHeartbeatBufPerPeer), - ctx: ctx, - cancel: cancel, + heartbeat: make(chan dispatchRequest, defaultHeartbeatBufPerPeer), + heartbeatResp: make(chan dispatchRequest, defaultHeartbeatRespBufPerPeer), + readIndexResp: make(chan dispatchRequest, defaultReadIndexRespBufPerPeer), + ctx: ctx, + cancel: cancel, } var workers []chan dispatchRequest if e.dispatcherLanesEnabled { - // 4-lane layout: split MsgApp/MsgAppResp (replication), MsgSnap - // (snapshot), and misc (other) onto independent goroutines so a - // bulky snapshot transfer cannot stall replication. Each channel - // still serves a single peer, so within-type ordering (the raft - // invariant we care about for MsgApp) is preserved. + // 6-lane layout: split MsgHeartbeatResp, ReadIndex heartbeat + // responses, MsgApp/MsgAppResp (replication), MsgSnap (snapshot), and + // misc (other) onto independent goroutines so a bulky snapshot transfer + // cannot stall replication or follower heartbeat responses. Each channel + // still serves a single peer, so within-type ordering (the raft invariant + // we care about for MsgApp) is preserved. pd.replication = make(chan dispatchRequest, size) pd.snapshot = make(chan dispatchRequest, defaultSnapshotLaneBufPerPeer) pd.other = make(chan dispatchRequest, defaultOtherLaneBufPerPeer) - workers = []chan dispatchRequest{pd.heartbeat, pd.replication, pd.snapshot, pd.other} + workers = []chan dispatchRequest{pd.heartbeat, pd.heartbeatResp, pd.readIndexResp, pd.replication, pd.snapshot, pd.other} } else { pd.normal = make(chan dispatchRequest, size) - workers = []chan dispatchRequest{pd.normal, pd.heartbeat} + workers = []chan dispatchRequest{pd.normal, pd.heartbeat, pd.heartbeatResp, pd.readIndexResp} } e.peerDispatchers[nodeID] = pd e.dispatchWG.Add(len(workers)) @@ -4539,7 +4757,7 @@ func snapshotEveryFromEnv() uint64 { return n } -// dispatcherLanesEnabledFromEnv returns true when the 4-lane dispatcher has +// dispatcherLanesEnabledFromEnv returns true when the multi-lane dispatcher has // been explicitly opted into via ELASTICKV_RAFT_DISPATCHER_LANES. The value // is parsed with strconv.ParseBool, which accepts the standard tokens // (1, t, T, TRUE, true, True enable; 0, f, F, FALSE, false, False disable). @@ -4568,10 +4786,10 @@ func preVoteEnabledFromEnv() bool { } // closePeerLanes closes every non-nil dispatch channel on pd so that the -// drain loops in runDispatchWorker exit. It is safe to call with either the -// 2-lane or 4-lane layout because unused lanes are nil. +// drain loops in runDispatchWorker exit. It is safe to call with any dispatch +// layout because unused lanes are nil. func closePeerLanes(pd *peerQueues) { - for _, ch := range []chan dispatchRequest{pd.heartbeat, pd.normal, pd.replication, pd.snapshot, pd.other} { + for _, ch := range []chan dispatchRequest{pd.heartbeat, pd.heartbeatResp, pd.readIndexResp, pd.normal, pd.replication, pd.snapshot, pd.other} { if ch != nil { close(ch) } @@ -4610,7 +4828,11 @@ func (e *Engine) handleDispatchRequest(ctx context.Context, req dispatchRequest) if err := req.Close(); err != nil { slog.Error("etcd raft dispatch: failed to close request", "err", err) } - if dispatchErr == nil || errors.Is(dispatchErr, ctx.Err()) { + if dispatchErr == nil { + e.reportSuccessfulDispatch(req.msg) + return + } + if errors.Is(dispatchErr, ctx.Err()) { return } code := dispatchErrorCodeOf(dispatchErr) @@ -4629,7 +4851,18 @@ func (e *Engine) handleDispatchRequest(ctx context.Context, req dispatchRequest) // out of StateReplicate / StateSnapshot. Without this the leader keeps // Progress stuck and never retries sendAppend/sendSnap for the peer, // leaving the follower indefinitely stale even after heartbeats resume. - e.postDispatchReport(dispatchReport{to: req.msg.GetTo(), msgType: req.msg.GetType()}) + e.reportFailedDispatch(req.msg) +} + +func (e *Engine) reportSuccessfulDispatch(msg raftpb.Message) { + if msg.GetType() != raftpb.MsgSnap { + return + } + e.postReliableDispatchReport(dispatchReport{ + to: msg.GetTo(), + msgType: msg.GetType(), + snapshotFinish: true, + }) } func (e *Engine) stopDispatchWorkers() { @@ -4981,10 +5214,39 @@ func (e *Engine) recordDroppedDispatch(msg raftpb.Message) { } func (e *Engine) reportDroppedDispatch(msg raftpb.Message) { + report := dispatchReport{to: msg.GetTo(), msgType: msg.GetType()} + if msg.GetType() == raftpb.MsgSnap { + e.deferReadyDispatchReport(report) + return + } if e.dispatchReportCh == nil { return } - e.postDispatchReport(dispatchReport{to: msg.GetTo(), msgType: msg.GetType()}) + e.postDispatchReport(report) +} + +func (e *Engine) reportFailedDispatch(msg raftpb.Message) { + report := dispatchReport{to: msg.GetTo(), msgType: msg.GetType()} + if msg.GetType() == raftpb.MsgSnap { + e.postReliableDispatchReport(report) + return + } + e.postDispatchReport(report) +} + +func (e *Engine) deferReadyDispatchReport(report dispatchReport) { + e.deferredReadyDispatchReports = append(e.deferredReadyDispatchReports, report) +} + +func (e *Engine) flushDeferredDispatchReports() { + if len(e.deferredReadyDispatchReports) == 0 { + return + } + reports := e.deferredReadyDispatchReports + e.deferredReadyDispatchReports = nil + for _, report := range reports { + e.handleDispatchReport(report) + } } // dispatchErrorCodeOf extracts the grpc status code name from err, or diff --git a/internal/raftengine/etcd/engine_applied_index_test.go b/internal/raftengine/etcd/engine_applied_index_test.go index 6a47d431e..107dcaefa 100644 --- a/internal/raftengine/etcd/engine_applied_index_test.go +++ b/internal/raftengine/etcd/engine_applied_index_test.go @@ -60,7 +60,9 @@ func (f *recordingAppliedIndexFSM) SetDurableAppliedIndex(idx uint64) error { f.failNext = false return f.failErr } - f.rec.record("bump", idx) + if f.rec != nil { + f.rec.record("bump", idx) + } return nil } @@ -236,6 +238,18 @@ func TestProtectReceivedFSMSnapshotWaitsOnSnapshotMuForAlreadyAppliedIndex(t *te require.Empty(t, e.protectedReceivedFSMSnaps) } +func TestProtectReceivedFSMSnapshotRejectsDuplicateInFlightIndex(t *testing.T) { + e := &Engine{} + + require.True(t, e.protectReceivedFSMSnapshot(9)) + require.False(t, e.protectReceivedFSMSnapshot(9)) + require.Equal(t, map[uint64]int{9: 1}, e.protectedReceivedFSMSnaps) + + e.unprotectReceivedFSMSnapshot(9) + require.Empty(t, e.protectedReceivedFSMSnaps) + require.True(t, e.protectReceivedFSMSnapshot(9)) +} + func TestUnprotectReceivedFSMSnapshotTokenIfApplied(t *testing.T) { e := &Engine{ protectedReceivedFSMSnaps: map[uint64]int{9: 1}, @@ -284,8 +298,29 @@ func TestReleaseIgnoredReceivedFSMSnapshotStepsUnprotectsNonSnapshotReady(t *tes require.Empty(t, e.pendingReceivedFSMSnapshotStep) } +func TestReleaseIgnoredReceivedFSMSnapshotStepsRemovesIgnoredFSMFile(t *testing.T) { + fsmSnapDir := t.TempDir() + writeFSMFileForTest(t, fsmSnapDir, 10, []byte("ignored snapshot")) + e := &Engine{ + fsmSnapDir: fsmSnapDir, + protectedReceivedFSMSnaps: map[uint64]int{10: 1}, + pendingReceivedFSMSnapshotStep: map[uint64]int{ + 10: 1, + }, + } + + e.releaseIgnoredReceivedFSMSnapshotSteps(etcdraft.Ready{}) + + require.NoFileExists(t, fsmSnapPath(fsmSnapDir, 10)) + require.Empty(t, e.protectedReceivedFSMSnaps) + require.Empty(t, e.pendingReceivedFSMSnapshotStep) +} + func TestReleaseIgnoredReceivedFSMSnapshotStepsKeepsSnapshotReadyProtected(t *testing.T) { + fsmSnapDir := t.TempDir() + writeFSMFileForTest(t, fsmSnapDir, 10, []byte("accepted snapshot")) e := &Engine{ + fsmSnapDir: fsmSnapDir, protectedReceivedFSMSnaps: map[uint64]int{10: 1}, pendingReceivedFSMSnapshotStep: map[uint64]int{ 10: 1, @@ -299,6 +334,7 @@ func TestReleaseIgnoredReceivedFSMSnapshotStepsKeepsSnapshotReadyProtected(t *te require.Equal(t, map[uint64]int{10: 1}, e.protectedReceivedFSMSnaps) require.Empty(t, e.pendingReceivedFSMSnapshotStep) + require.FileExists(t, fsmSnapPath(fsmSnapDir, 10)) } // TestRecordingFSM_SatisfiesAppliedIndexWriter is a compile-time- @@ -368,6 +404,60 @@ func TestPersistCreatedSnapshot_BumpErrorAborts(t *testing.T) { "failed bump MUST NOT have recorded; SaveSnap MUST NOT have run") } +func TestPersistReadyWithSnapshot_BumpsAppliedIndexAfterWALSnapshot(t *testing.T) { + rec := &applyIndexOrderRecorder{} + fsm := &recordingAppliedIndexFSM{rec: rec} + persist := &recordingPersistStorage{rec: rec} + fsmSnapDir := t.TempDir() + const index uint64 = 77 + crc, _ := writeFSMFileForTest(t, fsmSnapDir, index, []byte("snapshot-payload")) + e := &Engine{ + storage: etcdraft.NewMemoryStorage(), + fsm: fsm, + persist: persist, + dataDir: t.TempDir(), + fsmSnapDir: fsmSnapDir, + peers: map[uint64]Peer{ + 1: {NodeID: 1, ID: "n1", Address: "127.0.0.1:7001"}, + }, + } + + snap := appliedIndexTestSnapshot(index, encodeSnapshotToken(index, crc)) + require.NoError(t, e.persistReadyWithSnapshotLocked(etcdraft.Ready{Snapshot: &snap})) + + require.Equal(t, []orderEvent{ + {kind: "save", index: index}, + {kind: "bump", index: index}, + }, rec.snapshot(), + "received snapshot restore MUST publish its applied index only after the raft snapshot is durable") + require.Equal(t, index, e.applied) +} + +func TestPersistReadyWithSnapshot_BumpErrorAfterWALSnapshotSurfaces(t *testing.T) { + rec := &applyIndexOrderRecorder{} + fsm := &recordingAppliedIndexFSM{rec: rec, failNext: true, failErr: io.ErrShortBuffer} + persist := &recordingPersistStorage{rec: rec} + fsmSnapDir := t.TempDir() + const index uint64 = 78 + crc, _ := writeFSMFileForTest(t, fsmSnapDir, index, []byte("snapshot-payload")) + e := &Engine{ + storage: etcdraft.NewMemoryStorage(), + fsm: fsm, + persist: persist, + dataDir: t.TempDir(), + fsmSnapDir: fsmSnapDir, + peers: map[uint64]Peer{ + 1: {NodeID: 1, ID: "n1", Address: "127.0.0.1:7001"}, + }, + } + + snap := appliedIndexTestSnapshot(index, encodeSnapshotToken(index, crc)) + err := e.persistReadyWithSnapshotLocked(etcdraft.Ready{Snapshot: &snap}) + require.Error(t, err, "received snapshot bump failure MUST be surfaced") + require.Equal(t, []orderEvent{{kind: "save", index: index}}, rec.snapshot(), + "received snapshot path must not bump before the WAL snapshot is durable") +} + // --- Site 2: persistLocalSnapshotPayload (steady-state hot path) --- // // These mirror the Site 1 tests above but exercise the engine's diff --git a/internal/raftengine/etcd/engine_test.go b/internal/raftengine/etcd/engine_test.go index 22beaf800..7c5988461 100644 --- a/internal/raftengine/etcd/engine_test.go +++ b/internal/raftengine/etcd/engine_test.go @@ -108,6 +108,24 @@ type blockingSnapshotStateMachine struct { release chan struct{} } +type blockingApplyStateMachine struct { + started chan struct{} + release chan struct{} + startOnce sync.Once +} + +type recordingStartupApplyStateMachine struct { + *blockingApplyStateMachine + rec *applyIndexOrderRecorder +} + +func (s *recordingStartupApplyStateMachine) SetDurableAppliedIndex(idx uint64) error { + if s.rec != nil { + s.rec.record("bump", idx) + } + return nil +} + type blockingSnapshot struct { started chan struct{} release chan struct{} @@ -141,6 +159,21 @@ func (s *blockingSnapshotStateMachine) Apply(data []byte) any { return string(data) } +func (s *blockingApplyStateMachine) Apply(data []byte) any { + s.startOnce.Do(func() { close(s.started) }) + <-s.release + return string(data) +} + +func (s *blockingApplyStateMachine) Snapshot() (Snapshot, error) { + return &testSnapshot{}, nil +} + +func (s *blockingApplyStateMachine) Restore(r io.Reader) error { + _, err := io.Copy(io.Discard, r) + return err +} + func (s *blockingSnapshotStateMachine) Snapshot() (Snapshot, error) { return &blockingSnapshot{ started: s.started, @@ -570,7 +603,7 @@ func TestHandleTransportMessageWaitsForStartup(t *testing.T) { require.NoError(t, <-errCh) } -func TestEnqueueStepReturnsQueueFull(t *testing.T) { +func TestEnqueueStepBestEffortReturnsQueueFull(t *testing.T) { engine := &Engine{ doneCh: make(chan struct{}), stepCh: make(chan raftpb.Message, 1), @@ -579,20 +612,66 @@ func TestEnqueueStepReturnsQueueFull(t *testing.T) { require.Equal(t, uint64(0), engine.StepQueueFullCount()) - err := engine.enqueueStep(context.Background(), raftpb.Message{Type: messageTypePtr(raftpb.MsgApp)}) + err := engine.enqueueStep(context.Background(), raftpb.Message{Type: messageTypePtr(raftpb.MsgStorageAppend)}) require.Error(t, err) require.True(t, errors.Is(err, errStepQueueFull)) // The Prometheus hot-path dashboard relies on StepQueueFullCount - // advancing exactly once per rejected enqueue so the scraped rate - // equals the true drop rate, not a multiple of it. + // advancing exactly once per full enqueue attempt so the scraped + // rate equals the true congestion rate, not a multiple of it. require.Equal(t, uint64(1), engine.StepQueueFullCount()) - err = engine.enqueueStep(context.Background(), raftpb.Message{Type: messageTypePtr(raftpb.MsgApp)}) + err = engine.enqueueStep(context.Background(), raftpb.Message{Type: messageTypePtr(raftpb.MsgStorageAppend)}) require.Error(t, err) require.Equal(t, uint64(2), engine.StepQueueFullCount()) } +func TestEnqueueStepMsgAppWaitsForQueueSlot(t *testing.T) { + engine := &Engine{ + doneCh: make(chan struct{}), + stepCh: make(chan raftpb.Message, 1), + } + engine.stepCh <- raftpb.Message{Type: messageTypePtr(raftpb.MsgHeartbeat)} + + errCh := make(chan error, 1) + go func() { + errCh <- engine.enqueueStep(context.Background(), raftpb.Message{Type: messageTypePtr(raftpb.MsgApp)}) + }() + + select { + case err := <-errCh: + t.Fatalf("MsgApp enqueue returned before the queue had room: %v", err) + case <-time.After(20 * time.Millisecond): + } + require.Equal(t, uint64(1), engine.StepQueueFullCount()) + + firstMsg := <-engine.stepCh + require.Equal(t, raftpb.MsgHeartbeat, firstMsg.GetType()) + select { + case err := <-errCh: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("MsgApp enqueue did not resume after the queue had room") + } + nextMsg := <-engine.stepCh + require.Equal(t, raftpb.MsgApp, nextMsg.GetType()) +} + +func TestEnqueueStepMsgAppReturnsContextDeadlineWhenQueueStaysFull(t *testing.T) { + engine := &Engine{ + doneCh: make(chan struct{}), + stepCh: make(chan raftpb.Message, 1), + } + engine.stepCh <- raftpb.Message{Type: messageTypePtr(raftpb.MsgHeartbeat)} + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Millisecond) + defer cancel() + err := engine.enqueueStep(ctx, raftpb.Message{Type: messageTypePtr(raftpb.MsgApp)}) + require.Error(t, err) + require.True(t, errors.Is(err, context.DeadlineExceeded)) + require.Equal(t, uint64(1), engine.StepQueueFullCount()) +} + func TestEnqueueStepPriorityBypassesFullBulkQueue(t *testing.T) { engine := &Engine{ doneCh: make(chan struct{}), @@ -612,12 +691,36 @@ func TestEnqueueStepPriorityBypassesFullBulkQueue(t *testing.T) { t.Fatal("priority heartbeat was not enqueued") } - err = engine.enqueueStep(context.Background(), raftpb.Message{Type: messageTypePtr(raftpb.MsgApp)}) + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Millisecond) + defer cancel() + err = engine.enqueueStep(ctx, raftpb.Message{Type: messageTypePtr(raftpb.MsgApp)}) require.Error(t, err) - require.True(t, errors.Is(err, errStepQueueFull)) + require.True(t, errors.Is(err, context.DeadlineExceeded)) require.Equal(t, uint64(1), engine.StepQueueFullCount()) } +func TestEnqueueStepSnapshotBypassesFullBulkQueue(t *testing.T) { + engine := &Engine{ + doneCh: make(chan struct{}), + stepCh: make(chan raftpb.Message, 1), + priorityStepCh: make(chan raftpb.Message, 1), + } + engine.stepCh <- raftpb.Message{Type: messageTypePtr(raftpb.MsgApp)} + + err := engine.enqueueStep(context.Background(), raftpb.Message{Type: messageTypePtr(raftpb.MsgSnap)}) + require.NoError(t, err) + require.Equal(t, uint64(0), engine.StepQueueFullCount()) + + select { + case msg := <-engine.priorityStepCh: + require.Equal(t, raftpb.MsgSnap, msg.GetType()) + default: + t.Fatal("snapshot step was not enqueued on the priority queue") + } + + require.Len(t, engine.stepCh, 1) +} + func TestHandleEventDrainsPriorityStepBeforeBulkStep(t *testing.T) { engine := &Engine{ closeCh: make(chan struct{}), @@ -800,6 +903,58 @@ func TestSendMessagesDoesNotBlockWhenDispatchQueueIsFull(t *testing.T) { } } +func TestReportFailedDispatchReliablyQueuesMsgSnapFailure(t *testing.T) { + engine := &Engine{ + dispatchReportCh: make(chan dispatchReport, 1), + closeCh: make(chan struct{}), + } + engine.dispatchReportCh <- dispatchReport{to: 2, msgType: raftpb.MsgHeartbeat} + + done := make(chan struct{}) + go func() { + engine.reportFailedDispatch(raftpb.Message{ + Type: messageTypePtr(raftpb.MsgSnap), + To: uint64Ptr(3), + }) + close(done) + }() + + select { + case <-done: + t.Fatal("MsgSnap failure report must wait instead of dropping when report channel is full") + case <-time.After(20 * time.Millisecond): + } + + <-engine.dispatchReportCh + requireSignal(t, done, time.Second, "MsgSnap failure report was not delivered after report channel drained") + report := <-engine.dispatchReportCh + require.Equal(t, uint64(3), report.to) + require.Equal(t, raftpb.MsgSnap, report.msgType) + require.False(t, report.snapshotFinish) +} + +func TestReportDroppedDispatchDefersMsgSnapUntilReadyAdvanced(t *testing.T) { + engine := &Engine{ + dispatchReportCh: make(chan dispatchReport, 1), + closeCh: make(chan struct{}), + } + queued := dispatchReport{to: 2, msgType: raftpb.MsgHeartbeat} + engine.dispatchReportCh <- queued + + engine.reportDroppedDispatch(raftpb.Message{ + Type: messageTypePtr(raftpb.MsgSnap), + To: uint64Ptr(3), + }) + + require.Equal(t, queued, <-engine.dispatchReportCh) + select { + case report := <-engine.dispatchReportCh: + t.Fatalf("dropped MsgSnap report should not enqueue from the event loop: %+v", report) + default: + } + require.Equal(t, []dispatchReport{{to: 3, msgType: raftpb.MsgSnap}}, engine.deferredReadyDispatchReports) +} + func TestStopDispatchWorkersCancelsInflightDispatch(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) t.Cleanup(cancel) @@ -864,6 +1019,8 @@ func TestUpsertPeerStartsDispatcherAndAcceptsMessages(t *testing.T) { pd, ok := engine.peerDispatchers[2] require.True(t, ok, "dispatcher must be created on upsert") require.Equal(t, defaultHeartbeatBufPerPeer, cap(pd.heartbeat)) + require.Equal(t, defaultHeartbeatRespBufPerPeer, cap(pd.heartbeatResp)) + require.Equal(t, defaultReadIndexRespBufPerPeer, cap(pd.readIndexResp)) require.Equal(t, 4, cap(pd.normal)) require.NoError(t, engine.enqueueDispatchMessage(raftpb.Message{Type: messageTypePtr(raftpb.MsgHeartbeat), To: uint64Ptr(2)})) @@ -883,10 +1040,12 @@ func TestRemovePeerClosesDispatcherAndDropsSubsequentMessages(t *testing.T) { stopCh := make(chan struct{}) ctx, cancel := context.WithCancel(context.Background()) pd := &peerQueues{ - normal: make(chan dispatchRequest, 4), - heartbeat: make(chan dispatchRequest, 4), - ctx: ctx, - cancel: cancel, + normal: make(chan dispatchRequest, 4), + heartbeat: make(chan dispatchRequest, 4), + heartbeatResp: make(chan dispatchRequest, 4), + readIndexResp: make(chan dispatchRequest, 4), + ctx: ctx, + cancel: cancel, } engine := &Engine{ nodeID: 1, @@ -894,9 +1053,11 @@ func TestRemovePeerClosesDispatcherAndDropsSubsequentMessages(t *testing.T) { peerDispatchers: map[uint64]*peerQueues{2: pd}, dispatchStopCh: stopCh, } - engine.dispatchWG.Add(2) + engine.dispatchWG.Add(4) go engine.runDispatchWorker(ctx, pd.normal) go engine.runDispatchWorker(ctx, pd.heartbeat) + go engine.runDispatchWorker(ctx, pd.heartbeatResp) + go engine.runDispatchWorker(ctx, pd.readIndexResp) engine.removePeer(2) @@ -1404,6 +1565,148 @@ func TestPrepareDispatchRequestClonesSnapshotPayload(t *testing.T) { require.Equal(t, []uint64{1, 2}, req.msg.Snapshot.GetMetadata().GetConfState().GetVoters()) } +func TestEnqueueDispatchMessageCoalescesPlainHeartbeatResponses(t *testing.T) { + t.Parallel() + pd := &peerQueues{ + heartbeat: make(chan dispatchRequest, 1), + heartbeatResp: make(chan dispatchRequest, 1), + } + engine := &Engine{ + nodeID: 1, + peerDispatchers: map[uint64]*peerQueues{ + 2: pd, + }, + } + pd.heartbeatResp <- prepareDispatchRequest(raftpb.Message{ + Type: messageTypePtr(raftpb.MsgHeartbeatResp), + To: uint64Ptr(2), + }) + + require.NoError(t, engine.enqueueDispatchMessage(raftpb.Message{ + Type: messageTypePtr(raftpb.MsgHeartbeatResp), + To: uint64Ptr(2), + Context: []byte("new"), + })) + + require.Zero(t, engine.DispatchDropCount()) + require.Len(t, pd.heartbeatResp, 1) + req := <-pd.heartbeatResp + require.Equal(t, raftpb.MsgHeartbeatResp, req.msg.GetType()) + require.Equal(t, []byte("new"), req.msg.Context) +} + +func TestEnqueueDispatchMessagePreservesReadIndexHeartbeatResponses(t *testing.T) { + t.Parallel() + pd := &peerQueues{ + heartbeat: make(chan dispatchRequest, 1), + heartbeatResp: make(chan dispatchRequest, 1), + } + engine := &Engine{ + nodeID: 1, + peerDispatchers: map[uint64]*peerQueues{ + 2: pd, + }, + } + pd.heartbeatResp <- prepareDispatchRequest(raftpb.Message{ + Type: messageTypePtr(raftpb.MsgHeartbeatResp), + To: uint64Ptr(2), + Context: []byte("read-index"), + }) + + require.NoError(t, engine.enqueueDispatchMessage(raftpb.Message{ + Type: messageTypePtr(raftpb.MsgHeartbeatResp), + To: uint64Ptr(2), + })) + + require.Equal(t, uint64(1), engine.DispatchDropCount()) + require.Len(t, pd.heartbeatResp, 1) + req := <-pd.heartbeatResp + require.Equal(t, raftpb.MsgHeartbeatResp, req.msg.GetType()) + require.Equal(t, []byte("read-index"), req.msg.Context) +} + +func TestEnqueueDispatchMessageCoalescesPlainHeartbeatBehindReadIndex(t *testing.T) { + t.Parallel() + pd := &peerQueues{ + heartbeat: make(chan dispatchRequest, 1), + heartbeatResp: make(chan dispatchRequest, 3), + } + engine := &Engine{ + nodeID: 1, + peerDispatchers: map[uint64]*peerQueues{ + 2: pd, + }, + } + pd.heartbeatResp <- prepareDispatchRequest(raftpb.Message{ + Type: messageTypePtr(raftpb.MsgHeartbeatResp), + To: uint64Ptr(2), + Context: []byte("read-index"), + }) + pd.heartbeatResp <- prepareDispatchRequest(raftpb.Message{ + Type: messageTypePtr(raftpb.MsgHeartbeatResp), + To: uint64Ptr(2), + Index: uint64Ptr(10), + }) + pd.heartbeatResp <- prepareDispatchRequest(raftpb.Message{ + Type: messageTypePtr(raftpb.MsgHeartbeatResp), + To: uint64Ptr(2), + Index: uint64Ptr(11), + }) + + require.NoError(t, engine.enqueueDispatchMessage(raftpb.Message{ + Type: messageTypePtr(raftpb.MsgHeartbeatResp), + To: uint64Ptr(2), + Index: uint64Ptr(99), + })) + + require.Zero(t, engine.DispatchDropCount()) + require.Len(t, pd.heartbeatResp, 3) + req := <-pd.heartbeatResp + require.Equal(t, []byte("read-index"), req.msg.Context) + req = <-pd.heartbeatResp + require.Equal(t, uint64(99), req.msg.GetIndex()) + req = <-pd.heartbeatResp + require.Equal(t, uint64(11), req.msg.GetIndex()) +} + +func TestEnqueueDispatchMessagePreservesIncomingReadIndexHeartbeatWhenRespLaneFull(t *testing.T) { + t.Parallel() + pd := &peerQueues{ + heartbeat: make(chan dispatchRequest, 1), + heartbeatResp: make(chan dispatchRequest, 2), + readIndexResp: make(chan dispatchRequest, 1), + } + engine := &Engine{ + nodeID: 1, + peerDispatchers: map[uint64]*peerQueues{ + 2: pd, + }, + } + pd.heartbeatResp <- prepareDispatchRequest(raftpb.Message{ + Type: messageTypePtr(raftpb.MsgHeartbeatResp), + To: uint64Ptr(2), + Context: []byte("read-index-a"), + }) + pd.heartbeatResp <- prepareDispatchRequest(raftpb.Message{ + Type: messageTypePtr(raftpb.MsgHeartbeatResp), + To: uint64Ptr(2), + Context: []byte("read-index-b"), + }) + + require.NoError(t, engine.enqueueDispatchMessage(raftpb.Message{ + Type: messageTypePtr(raftpb.MsgHeartbeatResp), + To: uint64Ptr(2), + Context: []byte("read-index-c"), + })) + + require.Zero(t, engine.DispatchDropCount()) + require.Len(t, pd.heartbeatResp, 2) + require.Len(t, pd.readIndexResp, 1) + req := <-pd.readIndexResp + require.Equal(t, raftpb.MsgHeartbeatResp, req.msg.GetType()) + require.Equal(t, []byte("read-index-c"), req.msg.Context) +} + func TestMaxAppliedIndexStartsFromSnapshotIndex(t *testing.T) { storage := etcdraft.NewMemoryStorage() snap := raftTestSnapshot(5, 2, []uint64{1}, nil) @@ -1444,6 +1747,154 @@ func TestOpenRestoresLegacySnapshotState(t *testing.T) { require.Equal(t, [][]byte{[]byte("snap"), []byte("tail")}, fsm.Applied()) } +func TestOpenMultiNodeWaitsForCommittedTailDrain(t *testing.T) { + dir := t.TempDir() + peers := []Peer{ + {NodeID: 1, ID: "n1", Address: "127.0.0.1:7001"}, + {NodeID: 2, ID: "n2", Address: "127.0.0.1:7002"}, + } + require.NoError(t, saveStateFile(stateFilePath(dir), persistedState{ + HardState: testHardState(2, 2), + Snapshot: raftTestSnapshot(1, 1, []uint64{1, 2}, mustEncodeSnapshotData(t, nil)), + Entries: []raftpb.Entry{{ + Type: entryTypePtr(raftpb.EntryNormal), + Term: uint64Ptr(2), + Index: uint64Ptr(2), + Data: encodeProposalEnvelope(1, []byte("tail")), + }}, + })) + + rec := &applyIndexOrderRecorder{} + fsm := &recordingStartupApplyStateMachine{ + blockingApplyStateMachine: &blockingApplyStateMachine{ + started: make(chan struct{}), + release: make(chan struct{}), + }, + rec: rec, + } + done := make(chan openResult, 1) + go func() { + engine, err := Open(context.Background(), OpenConfig{ + NodeID: 1, + LocalID: "n1", + LocalAddress: "127.0.0.1:7001", + DataDir: dir, + Peers: peers, + StateMachine: fsm, + }) + done <- openResult{engine: engine, err: err} + }() + + requireSignal(t, fsm.started, time.Second, "startup committed tail was not being applied") + requireNoOpenResult(t, done, 20*time.Millisecond, "multi-node Open returned before committed tail drain completed") + + close(fsm.release) + + result := requireOpenResult(t, done, time.Second, "multi-node Open did not return after committed tail drain completed") + require.NoError(t, result.err) + require.NotNil(t, result.engine) + defer func() { + require.NoError(t, result.engine.Close()) + }() + + require.NoError(t, result.engine.WaitStarted(context.Background())) + requireSignal(t, result.engine.startedCh, time.Second, "engine did not mark started after committed tail drain completed") + require.Equal(t, []orderEvent{{kind: "bump", index: 2}}, rec.snapshot(), + "startup must persist the committed tail applied index before marking the engine started") +} + +type openResult struct { + engine *Engine + err error +} + +func requireOpenResult(t *testing.T, ch <-chan openResult, timeout time.Duration, msg string) openResult { + t.Helper() + select { + case result := <-ch: + return result + case <-time.After(timeout): + t.Fatal(msg) + return openResult{} + } +} + +func requireNoOpenResult(t *testing.T, ch <-chan openResult, timeout time.Duration, msg string) { + t.Helper() + select { + case result := <-ch: + if result.engine != nil { + _ = result.engine.Close() + } + if result.err != nil { + t.Fatalf("%s: Open returned error: %v", msg, result.err) + } + t.Fatal(msg) + case <-time.After(timeout): + } +} + +func requireSignal(t *testing.T, ch <-chan struct{}, timeout time.Duration, msg string) { + t.Helper() + select { + case <-ch: + case <-time.After(timeout): + t.Fatal(msg) + } +} + +func TestWaitStartedTimeoutDoesNotCloseEngine(t *testing.T) { + engine := &Engine{ + startedCh: make(chan struct{}), + doneCh: make(chan struct{}), + closeCh: make(chan struct{}), + } + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + err := engine.WaitStarted(ctx) + require.ErrorIs(t, err, context.Canceled) + select { + case <-engine.closeCh: + t.Fatal("WaitStarted timeout must not close an already-returned engine") + default: + } +} + +func TestWaitForOpenCanceledContextClosesMultiNodeEngine(t *testing.T) { + engine := &Engine{ + doneCh: make(chan struct{}), + closeCh: make(chan struct{}), + } + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + type result struct { + engine *Engine + err error + } + resultCh := make(chan result, 1) + go func() { + opened, err := waitForOpen(ctx, engine, false) + resultCh <- result{engine: opened, err: err} + }() + + select { + case <-engine.closeCh: + case <-time.After(time.Second): + t.Fatal("canceled multi-node Open must close the engine before returning") + } + close(engine.doneCh) + + select { + case got := <-resultCh: + require.Nil(t, got.engine) + require.ErrorIs(t, got.err, context.Canceled) + case <-time.After(time.Second): + t.Fatal("waitForOpen did not return after Close completed") + } +} + func TestOpenMultiNodeReplicatesOverGRPCTransport(t *testing.T) { nodes, peers := newTransportTestNodes(t, 3) startTransportTestServers(nodes, peers) @@ -2038,20 +2489,23 @@ func TestErrNotLeaderMatchesRaftEngineSentinel(t *testing.T) { require.True(t, errors.Is(errors.WithStack(errLeadershipTransferConfChangePending), raftengine.ErrLeadershipTransferConfChangePending)) } -// TestSelectDispatchLane_LegacyTwoLane verifies that, when the 4-lane -// dispatcher is disabled (default), messages are routed exactly as before: -// priority control traffic → heartbeat lane, everything else → normal lane. -func TestSelectDispatchLane_LegacyTwoLane(t *testing.T) { +// TestSelectDispatchLane_LegacyThreeLane verifies that, when the opt-in +// multi-lane dispatcher is disabled (default), priority control traffic uses +// the heartbeat lane, heartbeat responses use their coalescing response lane, +// and everything else uses the normal lane. +func TestSelectDispatchLane_LegacyThreeLane(t *testing.T) { t.Parallel() engine := &Engine{dispatcherLanesEnabled: false} pd := &peerQueues{ - normal: make(chan dispatchRequest, 1), - heartbeat: make(chan dispatchRequest, 1), + normal: make(chan dispatchRequest, 1), + heartbeat: make(chan dispatchRequest, 1), + heartbeatResp: make(chan dispatchRequest, 1), + readIndexResp: make(chan dispatchRequest, 1), } cases := map[raftpb.MessageType]chan dispatchRequest{ raftpb.MsgHeartbeat: pd.heartbeat, - raftpb.MsgHeartbeatResp: pd.heartbeat, + raftpb.MsgHeartbeatResp: pd.heartbeatResp, raftpb.MsgReadIndex: pd.heartbeat, raftpb.MsgReadIndexResp: pd.heartbeat, raftpb.MsgVote: pd.heartbeat, @@ -2067,24 +2521,32 @@ func TestSelectDispatchLane_LegacyTwoLane(t *testing.T) { got := engine.selectDispatchLane(pd, mt) require.Equalf(t, want, got, "legacy mode routing for %s", mt) } + got := engine.selectDispatchLaneForMessage(pd, raftpb.Message{ + Type: messageTypePtr(raftpb.MsgHeartbeatResp), + Context: []byte("read-index"), + }) + require.Equal(t, pd.readIndexResp, got) } -// TestSelectDispatchLane_FourLane verifies that, when ELASTICKV_RAFT_DISPATCHER_LANES +// TestSelectDispatchLane_FiveLane verifies that, when ELASTICKV_RAFT_DISPATCHER_LANES // is enabled, MsgApp/MsgAppResp goes to the replication lane, MsgSnap goes to -// the snapshot lane, and heartbeats/votes/read-index share the priority lane. -func TestSelectDispatchLane_FourLane(t *testing.T) { +// the snapshot lane, heartbeat responses get their own lane, and +// heartbeats/votes/read-index share the priority lane. +func TestSelectDispatchLane_FiveLane(t *testing.T) { t.Parallel() engine := &Engine{dispatcherLanesEnabled: true} pd := &peerQueues{ - heartbeat: make(chan dispatchRequest, 1), - replication: make(chan dispatchRequest, 1), - snapshot: make(chan dispatchRequest, 1), - other: make(chan dispatchRequest, 1), + heartbeat: make(chan dispatchRequest, 1), + heartbeatResp: make(chan dispatchRequest, 1), + readIndexResp: make(chan dispatchRequest, 1), + replication: make(chan dispatchRequest, 1), + snapshot: make(chan dispatchRequest, 1), + other: make(chan dispatchRequest, 1), } cases := map[raftpb.MessageType]chan dispatchRequest{ raftpb.MsgHeartbeat: pd.heartbeat, - raftpb.MsgHeartbeatResp: pd.heartbeat, + raftpb.MsgHeartbeatResp: pd.heartbeatResp, raftpb.MsgVote: pd.heartbeat, raftpb.MsgVoteResp: pd.heartbeat, raftpb.MsgPreVote: pd.heartbeat, @@ -2098,8 +2560,13 @@ func TestSelectDispatchLane_FourLane(t *testing.T) { } for mt, want := range cases { got := engine.selectDispatchLane(pd, mt) - require.Equalf(t, want, got, "4-lane mode routing for %s", mt) + require.Equalf(t, want, got, "multi-lane mode routing for %s", mt) } + got := engine.selectDispatchLaneForMessage(pd, raftpb.Message{ + Type: messageTypePtr(raftpb.MsgHeartbeatResp), + Context: []byte("read-index"), + }) + require.Equal(t, pd.readIndexResp, got) } // TestSelectDispatchLane_MsgPropReachesDefaultFallback verifies that MsgProp, @@ -2111,10 +2578,12 @@ func TestSelectDispatchLane_MsgPropReachesDefaultFallback(t *testing.T) { t.Parallel() engine := &Engine{nodeID: 1, dispatcherLanesEnabled: true} pd := &peerQueues{ - heartbeat: make(chan dispatchRequest, 1), - replication: make(chan dispatchRequest, 1), - snapshot: make(chan dispatchRequest, 1), - other: make(chan dispatchRequest, 1), + heartbeat: make(chan dispatchRequest, 1), + heartbeatResp: make(chan dispatchRequest, 1), + readIndexResp: make(chan dispatchRequest, 1), + replication: make(chan dispatchRequest, 1), + snapshot: make(chan dispatchRequest, 1), + other: make(chan dispatchRequest, 1), } require.NotPanics(t, func() { got := engine.selectDispatchLane(pd, raftpb.MsgProp) @@ -2122,11 +2591,11 @@ func TestSelectDispatchLane_MsgPropReachesDefaultFallback(t *testing.T) { }) } -// TestFourLaneDispatcher_SnapshotDoesNotBlockReplication exercises the key -// correctness invariant for the 4-lane layout: a stuck MsgSnap transfer must +// TestMultiLaneDispatcher_SnapshotDoesNotBlockReplication exercises the key +// correctness invariant for the multi-lane layout: a stuck MsgSnap transfer must // not prevent MsgApp from being dispatched, because they now run on // independent goroutines. -func TestFourLaneDispatcher_SnapshotDoesNotBlockReplication(t *testing.T) { +func TestMultiLaneDispatcher_SnapshotDoesNotBlockReplication(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) t.Cleanup(cancel) @@ -2179,19 +2648,20 @@ func TestFourLaneDispatcher_SnapshotDoesNotBlockReplication(t *testing.T) { engine.dispatchWG.Wait() } -// TestFourLaneDispatcher_RemovePeerClosesAllLanes confirms removePeer closes +// TestMultiLaneDispatcher_RemovePeerClosesAllLanes confirms removePeer closes // every lane (not just normal/heartbeat) so no worker goroutine leaks under -// the opt-in 4-lane layout. -func TestFourLaneDispatcher_RemovePeerClosesAllLanes(t *testing.T) { +// the opt-in multi-lane layout. +func TestMultiLaneDispatcher_RemovePeerClosesAllLanes(t *testing.T) { stopCh := make(chan struct{}) ctx, cancel := context.WithCancel(context.Background()) pd := &peerQueues{ - heartbeat: make(chan dispatchRequest, 4), - replication: make(chan dispatchRequest, 4), - snapshot: make(chan dispatchRequest, 4), - other: make(chan dispatchRequest, 4), - ctx: ctx, - cancel: cancel, + heartbeat: make(chan dispatchRequest, 4), + heartbeatResp: make(chan dispatchRequest, 4), + replication: make(chan dispatchRequest, 4), + snapshot: make(chan dispatchRequest, 4), + other: make(chan dispatchRequest, 4), + ctx: ctx, + cancel: cancel, } engine := &Engine{ nodeID: 1, @@ -2200,8 +2670,9 @@ func TestFourLaneDispatcher_RemovePeerClosesAllLanes(t *testing.T) { dispatchStopCh: stopCh, dispatcherLanesEnabled: true, } - engine.dispatchWG.Add(4) + engine.dispatchWG.Add(5) go engine.runDispatchWorker(ctx, pd.heartbeat) + go engine.runDispatchWorker(ctx, pd.heartbeatResp) go engine.runDispatchWorker(ctx, pd.replication) go engine.runDispatchWorker(ctx, pd.snapshot) go engine.runDispatchWorker(ctx, pd.other) @@ -2216,7 +2687,7 @@ func TestFourLaneDispatcher_RemovePeerClosesAllLanes(t *testing.T) { select { case <-done: case <-time.After(time.Second): - t.Fatal("4-lane dispatch workers did not exit after peer removal") + t.Fatal("multi-lane dispatch workers did not exit after peer removal") } // Subsequent sends to the removed peer must be dropped without panic. diff --git a/internal/raftengine/etcd/fsm_snapshot_file.go b/internal/raftengine/etcd/fsm_snapshot_file.go index 574724cae..13f3d3b0c 100644 --- a/internal/raftengine/etcd/fsm_snapshot_file.go +++ b/internal/raftengine/etcd/fsm_snapshot_file.go @@ -46,6 +46,12 @@ const ( // fsmWriteBufSize is the bufio.Writer buffer size used when writing .fsm files. fsmWriteBufSize = 1 << 20 // 1 MiB + // defaultMaxRetainedFSMSnapshotBytes bounds retained .fsm payload bytes + // after successful snapshot publication. Large production FSM snapshots can + // be tens of GiB; keeping defaultMaxSnapFiles full copies leaves too little + // headroom for the next receive-side spool. + defaultMaxRetainedFSMSnapshotBytes = int64(16 << 30) // 16 GiB + // fsmMaxInMemPayload is the maximum payload size that readFSMSnapshotPayload // will materialise into memory. Larger snapshots must use the streaming path // (openFSMSnapshotPayloadReader) to avoid OOM. 1 GiB is chosen as a generous @@ -53,6 +59,8 @@ const ( fsmMaxInMemPayload = int64(1 << 30) // 1 GiB ) +const maxRetainedFSMSnapshotBytesEnvVar = "ELASTICKV_RAFT_MAX_RETAINED_FSM_SNAPSHOT_BYTES" + var ( snapshotTokenMagic = [snapshotTokenMagicLen]byte{'E', 'K', 'V', 'T'} crc32cTable = crc32.MakeTable(crc32.Castagnoli) @@ -353,6 +361,18 @@ func restoreAndComputeCRC(f *os.File, fileSize int64, fsm StateMachine) (uint32, return h.Sum32(), nil } +func computeFSMSnapshotPayloadCRC(f *os.File, fileSize int64) (uint32, error) { + if _, err := f.Seek(0, io.SeekStart); err != nil { + return 0, errors.WithStack(err) + } + payloadSize := fileSize - fsmFooterSize + h := crc32.New(crc32cTable) + if _, err := io.CopyN(h, f, payloadSize); err != nil { + return 0, errors.WithStack(err) + } + return h.Sum32(), nil +} + // verifyFSMSnapshotFile performs a read-only CRC check without restoring the FSM. // Used for startup orphan detection. Pass tokenCRC=0 to skip the token comparison. func verifyFSMSnapshotFile(path string, tokenCRC uint32) error { @@ -619,6 +639,7 @@ type snapFileCandidate struct { index uint64 restorable bool walValid bool + tokenCRC uint32 } type prewriteSnapshotRetention struct { @@ -655,7 +676,7 @@ func purgeOlderSnapshotPairsBeforeWrite( return candidates[i].index < candidates[j].index }) - retention := keepRestorablePrewriteSnapshots(candidates) + retention := keepVerifiedPrewriteSnapshots(fsmSnapDir, candidates) var combined error combined = errors.CombineErrors(combined, purgeUnretainedPrewriteSnapshots(snapDir, fsmSnapDir, candidates, retention)) combined = errors.CombineErrors(combined, removePrewriteFSMOrphansBeforeIndex( @@ -746,6 +767,69 @@ func keepRestorablePrewriteSnapshots(candidates []snapFileCandidate) prewriteSna return retention } +func keepVerifiedPrewriteSnapshots(fsmSnapDir string, candidates []snapFileCandidate) prewriteSnapshotRetention { + retention := keepRestorablePrewriteSnapshots(candidates) + if fsmSnapDir == "" || len(candidates) <= prewriteSnapKeep { + return retention + } + for { + if !invalidateUnverifiedRetainedPrewriteSnapshots(fsmSnapDir, candidates, retention) { + return retention + } + retention = keepRestorablePrewriteSnapshots(candidates) + if !retentionHasRestorable(candidates, retention) { + return keepAllPrewriteSnapshots(candidates) + } + } +} + +func invalidateUnverifiedRetainedPrewriteSnapshots( + fsmSnapDir string, + candidates []snapFileCandidate, + retention prewriteSnapshotRetention, +) bool { + invalidated := false + for i := range candidates { + candidate := &candidates[i] + if !candidate.restorable || !retention.keep[candidate.name] || !retainedPrewriteSnapshotPrunesOlder(candidates, retention, *candidate) { + continue + } + if verifyFSMSnapshotFileWithToken(fsmSnapPath(fsmSnapDir, candidate.index), candidate.tokenCRC, true) == nil { + continue + } + candidate.restorable = false + invalidated = true + } + return invalidated +} + +func retainedPrewriteSnapshotPrunesOlder(candidates []snapFileCandidate, retention prewriteSnapshotRetention, retained snapFileCandidate) bool { + for _, candidate := range candidates { + if candidate.index >= retained.index || retention.keep[candidate.name] { + continue + } + return true + } + return false +} + +func retentionHasRestorable(candidates []snapFileCandidate, retention prewriteSnapshotRetention) bool { + for _, candidate := range candidates { + if candidate.restorable && retention.keep[candidate.name] { + return true + } + } + return false +} + +func keepAllPrewriteSnapshots(candidates []snapFileCandidate) prewriteSnapshotRetention { + retention := prewriteSnapshotRetention{keep: make(map[string]bool, len(candidates))} + for _, candidate := range candidates { + retention.keep[candidate.name] = true + } + return retention +} + func keepNewestMatchingPrewriteSnapshots( candidates []snapFileCandidate, retention *prewriteSnapshotRetention, @@ -813,25 +897,56 @@ func collectPrewriteSnapCandidates( if index == 0 || index >= nextIndex { continue } + restorable, tokenCRC := fsmSnapshotPairRestorable(snapDir, fsmSnapDir, e.Name(), term, index) candidates = append(candidates, snapFileCandidate{ name: e.Name(), index: index, - restorable: fsmSnapshotPairRestorable(snapDir, fsmSnapDir, e.Name(), term, index), + restorable: restorable, walValid: walValidIndexes == nil || walValidIndexes[walSnapshotKey{term: term, index: index}], + tokenCRC: tokenCRC, }) } return candidates } -func fsmSnapshotPairRestorable(snapDir, fsmSnapDir, snapName string, term, index uint64) bool { +func fsmSnapshotPairRestorable(snapDir, fsmSnapDir, snapName string, term, index uint64) (bool, uint32) { if fsmSnapDir == "" { - return false + return false, 0 } tok, ok := snapshotTokenFromSnapFile(snapDir, snapName, term, index) if !ok { - return false + return false, 0 } - return verifyFSMSnapshotFileWithToken(fsmSnapPath(fsmSnapDir, index), tok.CRC32C, true) == nil + // Prewrite cleanup runs on the snapshot receive hot path before gRPC starts + // draining payload chunks. Only do a footer/token check here; full-payload + // CRC remains in the actual restore/open paths. + return fsmSnapshotFooterMatchesToken(fsmSnapPath(fsmSnapDir, index), tok.CRC32C) == nil, tok.CRC32C +} + +func fsmSnapshotFooterMatchesToken(path string, tokenCRC uint32) error { + f, err := os.Open(path) + if err != nil { + return statFSMFileError(err) + } + defer f.Close() + + info, err := f.Stat() + if err != nil { + return errors.WithStack(err) + } + if info.Size() < fsmMinFileSize { + return errors.Wrapf(ErrFSMSnapshotTooSmall, + "file too small: %d bytes (minimum %d)", info.Size(), fsmMinFileSize) + } + footer, err := readFSMFooter(f, info.Size()) + if err != nil { + return err + } + if footer != tokenCRC { + return errors.Wrapf(ErrFSMSnapshotTokenCRC, + "path=%s footer=%08x token=%08x", path, footer, tokenCRC) + } + return nil } func snapshotTokenFromSnapFile(snapDir, snapName string, term, index uint64) (snapshotToken, bool) { @@ -986,7 +1101,7 @@ func purgeOldSnapshotFiles(snapDir, fsmSnapDir string) error { } snaps := collectSnapNames(entries) - if len(snaps) <= defaultMaxSnapFiles { + if len(snaps) == 0 { return nil } // Sort explicitly: os.ReadDir returns lexicographic order on most systems, @@ -994,8 +1109,12 @@ func purgeOldSnapshotFiles(snapDir, fsmSnapDir string) error { // hex, so lexicographic == chronological order (oldest first). sort.Strings(snaps) + maxKeep := retainedSnapshotFileLimit(snaps, fsmSnapDir) + if len(snaps) <= maxKeep { + return nil + } var combined error - for _, name := range snaps[:len(snaps)-defaultMaxSnapFiles] { + for _, name := range snaps[:len(snaps)-maxKeep] { if err := purgeSnapPair(snapDir, fsmSnapDir, name); err != nil { combined = errors.CombineErrors(combined, err) } @@ -1006,6 +1125,60 @@ func purgeOldSnapshotFiles(snapDir, fsmSnapDir string) error { return errors.WithStack(combined) } +func retainedSnapshotFileLimit(snaps []string, fsmSnapDir string) int { + maxKeep := defaultMaxSnapFiles + if maxKeep < 1 { + maxKeep = 1 + } + if maxKeep > len(snaps) { + maxKeep = len(snaps) + } + if fsmSnapDir == "" { + return maxKeep + } + budget := maxRetainedFSMSnapshotBytes() + if budget <= 0 { + return maxKeep + } + for maxKeep > 1 && retainedFSMSnapshotBytes(snaps[len(snaps)-maxKeep:], fsmSnapDir) > budget { + maxKeep-- + } + return maxKeep +} + +func maxRetainedFSMSnapshotBytes() int64 { + raw := strings.TrimSpace(os.Getenv(maxRetainedFSMSnapshotBytesEnvVar)) + if raw == "" { + return defaultMaxRetainedFSMSnapshotBytes + } + n, err := strconv.ParseInt(raw, 10, 64) + if err != nil { + slog.Warn("invalid max retained FSM snapshot bytes; using default", + "env", maxRetainedFSMSnapshotBytesEnvVar, + "value", raw, + "error", err, + ) + return defaultMaxRetainedFSMSnapshotBytes + } + return n +} + +func retainedFSMSnapshotBytes(snaps []string, fsmSnapDir string) int64 { + var total int64 + for _, name := range snaps { + idx := parseSnapFileIndex(name) + if idx == 0 { + continue + } + info, err := os.Stat(fsmSnapPath(fsmSnapDir, idx)) + if err != nil { + continue + } + total += info.Size() + } + return total +} + func collectSnapNames(entries []os.DirEntry) []string { var snaps []string for _, e := range entries { diff --git a/internal/raftengine/etcd/fsm_snapshot_file_test.go b/internal/raftengine/etcd/fsm_snapshot_file_test.go index a7d509b44..60cc33344 100644 --- a/internal/raftengine/etcd/fsm_snapshot_file_test.go +++ b/internal/raftengine/etcd/fsm_snapshot_file_test.go @@ -446,6 +446,54 @@ func TestPrepareFSMSnapshotWriteKeepsTokenMatchingFallbackPair(t *testing.T) { require.FileExists(t, fsmSnapPath(fsmSnapDir, 100)) } +func TestPrewriteRestorableCheckUsesSnapshotFooterOnly(t *testing.T) { + snapDir := t.TempDir() + fsmSnapDir := t.TempDir() + payload := []byte("payload") + + crc, path := writeFSMFileForTest(t, fsmSnapDir, 100, payload) + createTokenSnapFileWithTerm(t, snapDir, 1, 100, crc) + + f, err := os.OpenFile(path, os.O_WRONLY, 0) + require.NoError(t, err) + _, err = f.WriteAt([]byte("X"), 0) + require.NoError(t, err) + require.NoError(t, f.Close()) + + restorable, _ := fsmSnapshotPairRestorable( + snapDir, + fsmSnapDir, + "0000000000000001-0000000000000064.snap", + 1, + 100, + ) + require.True(t, restorable) + require.ErrorIs(t, verifyFSMSnapshotFile(path, crc), ErrFSMSnapshotFileCRC) +} + +func TestPrepareFSMSnapshotWriteVerifiesRetainedFooterOnlyFallbackBeforePruning(t *testing.T) { + snapDir := t.TempDir() + fsmSnapDir := t.TempDir() + + crc100, _ := writeFSMFileForTest(t, fsmSnapDir, 100, []byte("valid fallback")) + createTokenSnapFileWithTerm(t, snapDir, 1, 100, crc100) + crc200, path200 := writeFSMFileForTest(t, fsmSnapDir, 200, []byte("corrupt newest fallback")) + createTokenSnapFileWithTerm(t, snapDir, 1, 200, crc200) + + f, err := os.OpenFile(path200, os.O_WRONLY, 0) + require.NoError(t, err) + _, err = f.WriteAt([]byte("X"), 0) + require.NoError(t, err) + require.NoError(t, f.Close()) + + require.NoError(t, prepareFSMSnapshotWrite(snapDir, fsmSnapDir, 300)) + + require.FileExists(t, filepath.Join(snapDir, "0000000000000001-0000000000000064.snap")) + require.FileExists(t, fsmSnapPath(fsmSnapDir, 100)) + require.NoFileExists(t, filepath.Join(snapDir, "0000000000000001-00000000000000c8.snap")) + require.NoFileExists(t, fsmSnapPath(fsmSnapDir, 200)) +} + func TestPrepareFSMSnapshotWriteKeepsWALValidAndRestorableFallbacksWhenWALValidCandidateIsBroken(t *testing.T) { snapDir := t.TempDir() fsmSnapDir := t.TempDir() diff --git a/internal/raftengine/etcd/grpc_transport.go b/internal/raftengine/etcd/grpc_transport.go index 236c3989d..fdf5cc0b3 100644 --- a/internal/raftengine/etcd/grpc_transport.go +++ b/internal/raftengine/etcd/grpc_transport.go @@ -42,6 +42,7 @@ var ( errSnapshotMetadataDuplicate = errors.New("etcd raft snapshot metadata was sent more than once") errSnapshotMessageNil = errors.New("etcd raft snapshot message is required") errSnapshotStreamShort = errors.New("etcd raft snapshot stream closed before final chunk") + errSnapshotDispatchBusy = errors.New("etcd raft snapshot dispatch already in progress") errReceivedFSMSnapshotStale = errors.New("etcd raft received fsm snapshot is stale") errPeerStreamClosed = errors.New("etcd raft SendStream closed") errSendStreamDisabled = errors.New("etcd raft SendStream is disabled") @@ -145,7 +146,8 @@ type GRPCTransport struct { // bridgeSem limits concurrent bridge-mode snapshot materializations so // that aggregate in-memory allocation stays bounded even when multiple // dispatch workers run simultaneously. - bridgeSem chan struct{} + bridgeSem chan struct{} + snapshotSendSem chan struct{} sendStreamOpenCount atomic.Uint64 sendStreamReconnectCount atomic.Uint64 @@ -187,6 +189,7 @@ func NewGRPCTransport(peers []Peer) *GRPCTransport { sendStreamCancel: sendStreamCancel, snapshotChunkSize: defaultSnapshotChunkSize, bridgeSem: make(chan struct{}, defaultBridgeMaterializeLimit), + snapshotSendSem: make(chan struct{}, 1), } } @@ -436,6 +439,10 @@ func isSnapshotMsg(msg raftpb.Message) bool { func (t *GRPCTransport) dispatchSnapshot(ctx context.Context, msg raftpb.Message) error { ctx, cancel := transportContext(ctx, defaultSnapshotDispatchTimeout) defer cancel() + if !t.tryAcquireSnapshotSend() { + return errors.WithStack(errors.Mark(status.Error(codes.ResourceExhausted, errSnapshotDispatchBusy.Error()), errSnapshotDispatchBusy)) + } + defer t.releaseSnapshotSend() // Prefer streaming when the snapshot holds a token and an opener is wired. // This avoids materialising the full FSM payload in memory on the sender. @@ -460,6 +467,28 @@ func (t *GRPCTransport) dispatchSnapshot(ctx context.Context, msg raftpb.Message return t.sendSnapshot(ctx, patched) } +func (t *GRPCTransport) tryAcquireSnapshotSend() bool { + if t == nil || t.snapshotSendSem == nil { + return true + } + select { + case t.snapshotSendSem <- struct{}{}: + return true + default: + return false + } +} + +func (t *GRPCTransport) releaseSnapshotSend() { + if t == nil || t.snapshotSendSem == nil { + return + } + select { + case <-t.snapshotSendSem: + default: + } +} + // streamFSMSnapshot streams the .fsm payload file directly to the peer using // chunked gRPC without loading the full content into memory. func (t *GRPCTransport) streamFSMSnapshot(ctx context.Context, msg raftpb.Message, index uint64, openFn func(uint64) (io.ReadCloser, error)) error { @@ -581,7 +610,9 @@ func (t *GRPCTransport) dispatchRegular(ctx context.Context, msg raftpb.Message) } req := &pb.EtcdRaftMessage{Message: raw} if isPriorityMsg(msg.GetType()) || !t.sendStreamEnabledNow() || !t.allowPeerStreamProbe(peer.Address, time.Now()) { - return t.dispatchRegularUnary(ctx, client, req) + err := t.dispatchRegularUnary(ctx, client, req) + t.closePeerConnOnRetryableDialError(peer.Address, err) + return err } err = t.dispatchRegularStream(ctx, peer.Address, client, req) if err == nil { @@ -589,11 +620,16 @@ func (t *GRPCTransport) dispatchRegular(ctx context.Context, msg raftpb.Message) } if grpcStatusCode(err) == codes.Unimplemented { t.markPeerStreamUnsupported(peer.Address) - return t.dispatchRegularUnary(ctx, client, req) + err := t.dispatchRegularUnary(ctx, client, req) + t.closePeerConnOnRetryableDialError(peer.Address, err) + return err } if isSendStreamDisabled(err) { - return t.dispatchRegularUnary(ctx, client, req) + err := t.dispatchRegularUnary(ctx, client, req) + t.closePeerConnOnRetryableDialError(peer.Address, err) + return err } + t.closePeerConnOnRetryableDialError(peer.Address, err) return errors.WithStack(err) } @@ -602,6 +638,15 @@ func (t *GRPCTransport) dispatchRegularUnary(ctx context.Context, client pb.Etcd return errors.WithStack(err) } +func (t *GRPCTransport) closePeerConnOnRetryableDialError(address string, err error) { + if grpcStatusCode(err) != codes.Unavailable { + return + } + t.mu.Lock() + defer t.mu.Unlock() + t.closePeerConnLocked(address) +} + func (t *GRPCTransport) dispatchRegularStream(ctx context.Context, address string, client pb.EtcdRaftClient, req *pb.EtcdRaftMessage) error { stream, err := t.streamFor(ctx, address, client) if err != nil { @@ -1103,8 +1148,11 @@ func snapshotMessageHeader(msg raftpb.Message) ([]byte, error) { } func sendSnapshotChunks(stream pb.EtcdRaft_SendSnapshotClient, header []byte, payload []byte, chunkSize int) error { + if err := sendSnapshotChunk(stream, &pb.EtcdRaftSnapshotChunk{Metadata: header}); err != nil { + return err + } if len(payload) == 0 { - return sendSnapshotChunk(stream, &pb.EtcdRaftSnapshotChunk{Metadata: header, Final: true}) + return sendSnapshotChunk(stream, &pb.EtcdRaftSnapshotChunk{Final: true}) } for offset := 0; offset < len(payload); offset += chunkSize { end := offset + chunkSize @@ -1115,9 +1163,6 @@ func sendSnapshotChunks(stream pb.EtcdRaft_SendSnapshotClient, header []byte, pa Chunk: payload[offset:end], Final: end == len(payload), } - if offset == 0 { - chunk.Metadata = header - } if err := sendSnapshotChunk(stream, chunk); err != nil { return err } @@ -1136,6 +1181,9 @@ func sendSnapshotReaderChunks(stream pb.EtcdRaft_SendSnapshotClient, header []by if chunkSize <= 0 { chunkSize = defaultSnapshotChunkSize } + if err := sendSnapshotChunk(stream, &pb.EtcdRaftSnapshotChunk{Metadata: header}); err != nil { + return err + } buffered := bufio.NewReaderSize(reader, chunkSize) current, err := readSnapshotChunk(buffered, chunkSize) if err != nil { @@ -1145,14 +1193,13 @@ func sendSnapshotReaderChunks(stream pb.EtcdRaft_SendSnapshotClient, header []by // a small snapshot. Include it so the receiver does not get an // empty snapshot.Data. return sendSnapshotChunk(stream, &pb.EtcdRaftSnapshotChunk{ - Metadata: header, - Chunk: current, - Final: true, + Chunk: current, + Final: true, }) } return errors.WithStack(err) } - return streamReaderChunks(stream, header, buffered, current, chunkSize) + return streamReaderChunks(stream, nil, buffered, current, chunkSize) } // streamReaderChunks drains buffered starting from `current` (the first full @@ -1292,7 +1339,12 @@ func (t *GRPCTransport) receiveSnapshotStream(stream pb.EtcdRaft_SendSnapshotSer if err != nil { return raftpb.Message{}, err } - spool, err := newSnapshotSpool(spoolPlacement) + protection, err := protectReceivedSnapshotMetadata(metadata, fsmSnapDir, protectFn, unprotectFn) + if err != nil { + return raftpb.Message{}, err + } + defer protection.releaseUnlessHandedOff() + spool, err := newReceiveSnapshotSpool(spoolPlacement) if err != nil { return raftpb.Message{}, err } @@ -1320,10 +1372,12 @@ func (t *GRPCTransport) receiveSnapshotStream(stream pb.EtcdRaft_SendSnapshotSer metadata, firstPayloadChunk, preparedFSMWrite, + protection.protected, ) if err != nil { return raftpb.Message{}, err } + protection.handoff() index := uint64(0) if msg.Snapshot != nil { index = msg.Snapshot.GetMetadata().GetIndex() @@ -1337,6 +1391,50 @@ func (t *GRPCTransport) receiveSnapshotStream(stream pb.EtcdRaft_SendSnapshotSer return msg, nil } +type receivedSnapshotProtection struct { + index uint64 + protected bool + handedOff bool + unprotectFn func(uint64) +} + +func protectReceivedSnapshotMetadata( + metadata raftpb.Message, + fsmSnapDir string, + protectFn func(uint64) bool, + unprotectFn func(uint64), +) (receivedSnapshotProtection, error) { + if fsmSnapDir == "" || metadata.Snapshot == nil { + return receivedSnapshotProtection{}, nil + } + index := metadata.Snapshot.GetMetadata().GetIndex() + if index == 0 || protectFn == nil { + return receivedSnapshotProtection{}, nil + } + if !protectFn(index) { + return receivedSnapshotProtection{}, errors.WithStack(errReceivedFSMSnapshotStale) + } + return receivedSnapshotProtection{ + index: index, + protected: true, + unprotectFn: unprotectFn, + }, nil +} + +func (p *receivedSnapshotProtection) handoff() { + if p == nil { + return + } + p.handedOff = true +} + +func (p *receivedSnapshotProtection) releaseUnlessHandedOff() { + if p == nil || !p.protected || p.handedOff || p.unprotectFn == nil { + return + } + p.unprotectFn(p.index) +} + // drainSnapshotChunks consumes the SendSnapshot stream into spool, computes // CRC32C over the payload bytes as they hit disk, and on the final chunk // hands off to finalizeReceivedSnapshot — which decides between the @@ -1353,7 +1451,7 @@ func drainSnapshotChunks( unprotectFn func(uint64), ) (raftpb.Message, int64, error) { var metadata raftpb.Message - return drainSnapshotChunksFrom(stream, spool, fsmSnapDir, prepareFn, protectFn, unprotectFn, metadata, nil, false) + return drainSnapshotChunksFrom(stream, spool, fsmSnapDir, prepareFn, protectFn, unprotectFn, metadata, nil, false, false) } func drainSnapshotChunksFrom( @@ -1366,6 +1464,7 @@ func drainSnapshotChunksFrom( metadata raftpb.Message, firstPayloadChunk *pb.EtcdRaftSnapshotChunk, preparedFSMWrite bool, + preprotected bool, ) (raftpb.Message, int64, error) { seenMetadata := metadata.Snapshot != nil // Wrap spool with crc32CWriter so the CRC accumulates as bytes hit @@ -1398,7 +1497,7 @@ func drainSnapshotChunksFrom( } payloadBytes += int64(len(chunk.Chunk)) if chunk.Final { - msg, err := finalizeReceivedSnapshot(metadata, spool, crcWriter.Sum32(), fsmSnapDir, protectFn, unprotectFn, seenMetadata) + msg, err := finalizeReceivedSnapshot(metadata, spool, crcWriter.Sum32(), fsmSnapDir, protectFn, unprotectFn, seenMetadata, preprotected) if err != nil { return raftpb.Message{}, 0, err } @@ -1481,6 +1580,7 @@ func finalizeReceivedSnapshot( protectFn func(uint64) bool, unprotectFn func(uint64), seenMetadata bool, + preprotected bool, ) (raftpb.Message, error) { if !seenMetadata || metadata.Snapshot == nil { return raftpb.Message{}, errors.WithStack(errSnapshotMetadataNil) @@ -1492,23 +1592,38 @@ func finalizeReceivedSnapshot( // rename to). return buildSnapshotMessage(metadata, spool, seenMetadata) } - protected := false - if protectFn != nil { - if !protectFn(index) { - return raftpb.Message{}, errors.WithStack(errReceivedFSMSnapshotStale) - } - protected = true + protected, err := protectReceivedSnapshotForFinalize(index, preprotected, protectFn) + if err != nil { + return raftpb.Message{}, err } if err := spool.FinalizeAsFSMFile(fsmSnapDir, index, crc32c); err != nil { - if protected && unprotectFn != nil { - unprotectFn(index) - } + releaseFinalizeSnapshotProtection(index, protected, preprotected, unprotectFn) return raftpb.Message{}, err } metadata.Snapshot.Data = encodeSnapshotToken(index, crc32c) return metadata, nil } +func protectReceivedSnapshotForFinalize(index uint64, preprotected bool, protectFn func(uint64) bool) (bool, error) { + if preprotected { + return true, nil + } + if protectFn == nil { + return false, nil + } + if !protectFn(index) { + return false, errors.WithStack(errReceivedFSMSnapshotStale) + } + return true, nil +} + +func releaseFinalizeSnapshotProtection(index uint64, protected bool, preprotected bool, unprotectFn func(uint64)) { + if !protected || preprotected || unprotectFn == nil { + return + } + unprotectFn(index) +} + func maybePrepareReceivedFSMSnapshotWrite( metadata raftpb.Message, fsmSnapDir string, diff --git a/internal/raftengine/etcd/grpc_transport_test.go b/internal/raftengine/etcd/grpc_transport_test.go index 02ffcb97f..a57fc43bb 100644 --- a/internal/raftengine/etcd/grpc_transport_test.go +++ b/internal/raftengine/etcd/grpc_transport_test.go @@ -19,6 +19,7 @@ import ( "google.golang.org/grpc" "google.golang.org/grpc/codes" "google.golang.org/grpc/metadata" + "google.golang.org/grpc/status" "google.golang.org/protobuf/proto" ) @@ -124,6 +125,58 @@ func TestReceiveSnapshotStreamRejectsDuplicateMetadata(t *testing.T) { require.True(t, errors.Is(err, errSnapshotMetadataDuplicate)) } +func TestSnapshotSendGateAllowsOnlyOneInFlightSnapshot(t *testing.T) { + transport := NewGRPCTransport(nil) + + require.True(t, transport.tryAcquireSnapshotSend()) + require.False(t, transport.tryAcquireSnapshotSend()) + + transport.releaseSnapshotSend() + require.True(t, transport.tryAcquireSnapshotSend()) + transport.releaseSnapshotSend() +} + +func TestReceiveSnapshotStreamRejectsProtectedIndexBeforePayload(t *testing.T) { + const index = uint64(127) + metadata := raftpb.Message{ + Type: messageTypePtr(raftpb.MsgSnap), + From: uint64Ptr(1), + To: uint64Ptr(2), + Snapshot: &raftpb.Snapshot{ + Metadata: testSnapshotMetadata(index, 1, nil), + }, + } + raw, err := proto.Marshal(&metadata) + require.NoError(t, err) + + transport := NewGRPCTransport(nil) + fsmSnapDir := t.TempDir() + transport.SetFSMSnapDir(fsmSnapDir) + var protected []uint64 + transport.SetFSMSnapshotProtection( + func(got uint64) bool { + protected = append(protected, got) + return false + }, + func(uint64) { + t.Fatal("stale snapshot rejection must not unprotect an index it did not protect") + }, + ) + stream := &testSendSnapshotServer{ + chunks: []*pb.EtcdRaftSnapshotChunk{ + {Metadata: raw}, + {Chunk: []byte("payload that must not be read"), Final: true}, + }, + } + + _, err = transport.receiveSnapshotStream(stream) + require.Error(t, err) + require.True(t, errors.Is(err, errReceivedFSMSnapshotStale)) + require.Equal(t, []uint64{index}, protected) + require.Equal(t, 1, stream.index, "receiver must reject after metadata before reading payload chunks") + require.NoFileExists(t, fsmSnapPath(fsmSnapDir, index)) +} + // TestReceiveSnapshotStream_StreamingTokenWhenFSMSnapDirSet pins the // memory-safety win: when the receive transport has fsmSnapDir wired, // the spool file is renamed to fsmSnapPath(...) and Snapshot.Data is @@ -908,6 +961,30 @@ func TestDispatchRegularUsesUnaryForPriorityMessages(t *testing.T) { require.Zero(t, client.sendStreamCalls.Load()) } +func TestDispatchRegularDropsCachedPeerConnAfterUnaryUnavailable(t *testing.T) { + const addr = "host:2" + transport := NewGRPCTransport([]Peer{{NodeID: 2, Address: addr}}) + t.Cleanup(func() { require.NoError(t, transport.Close()) }) + client := &testEtcdRaftClient{sendErr: status.Error(codes.Unavailable, "connection refused")} + injectClient(t, transport, addr, client) + + err := transport.dispatchRegular(context.Background(), raftpb.Message{ + Type: messageTypePtr(raftpb.MsgHeartbeat), + From: uint64Ptr(1), + To: uint64Ptr(2), + Term: uint64Ptr(4), + Commit: uint64Ptr(22), + }) + require.Error(t, err) + require.Equal(t, codes.Unavailable, grpcStatusCode(err)) + + transport.mu.RLock() + _, cached := transport.clients[addr] + transport.mu.RUnlock() + require.False(t, cached) + require.Equal(t, int32(1), client.sendCalls.Load()) +} + func TestDispatchRegularUsesUnaryWhenSendStreamDisabled(t *testing.T) { t.Setenv(sendStreamEnabledEnvVar, "false") const addr = "host:2" @@ -1472,10 +1549,13 @@ func TestSendSnapshotReaderChunksSmallPayloadPreservesData(t *testing.T) { err := sendSnapshotReaderChunks(client, header, bytes.NewReader(payload), defaultSnapshotChunkSize) require.NoError(t, err) - require.Len(t, client.chunks, 1) + require.Len(t, client.chunks, 2) require.Equal(t, header, client.chunks[0].Metadata) - require.Equal(t, payload, client.chunks[0].Chunk) - require.True(t, client.chunks[0].Final) + require.Empty(t, client.chunks[0].Chunk) + require.False(t, client.chunks[0].Final) + require.Empty(t, client.chunks[1].Metadata) + require.Equal(t, payload, client.chunks[1].Chunk) + require.True(t, client.chunks[1].Final) } func TestSendSnapshotReaderChunksEmptyPayloadSendsHeaderOnly(t *testing.T) { @@ -1485,10 +1565,13 @@ func TestSendSnapshotReaderChunksEmptyPayloadSendsHeaderOnly(t *testing.T) { err := sendSnapshotReaderChunks(client, header, bytes.NewReader(nil), defaultSnapshotChunkSize) require.NoError(t, err) - require.Len(t, client.chunks, 1) + require.Len(t, client.chunks, 2) require.Equal(t, header, client.chunks[0].Metadata) require.Empty(t, client.chunks[0].Chunk) - require.True(t, client.chunks[0].Final) + require.False(t, client.chunks[0].Final) + require.Empty(t, client.chunks[1].Metadata) + require.Empty(t, client.chunks[1].Chunk) + require.True(t, client.chunks[1].Final) } // testSnapshotSendClient captures chunks sent via sendSnapshotReaderChunks / sendSnapshotChunks. @@ -1522,12 +1605,16 @@ type testEtcdRaftClient struct { blockSendStreamUntilContext bool sendStreamStarted chan struct{} releaseSendStream chan struct{} + sendErr error sendCalls atomic.Int32 sendStreamCalls atomic.Int32 } func (c *testEtcdRaftClient) Send(_ context.Context, _ *pb.EtcdRaftMessage, _ ...grpc.CallOption) (*pb.EtcdRaftAck, error) { c.sendCalls.Add(1) + if c.sendErr != nil { + return nil, c.sendErr + } return &pb.EtcdRaftAck{}, nil } @@ -1637,19 +1724,23 @@ func TestSendSnapshotReaderChunksMultiChunk(t *testing.T) { err := sendSnapshotReaderChunks(client, header, bytes.NewReader(payload), chunkSize) require.NoError(t, err) - require.Len(t, client.chunks, 3) + require.Len(t, client.chunks, 4) require.Equal(t, header, client.chunks[0].Metadata) - require.Equal(t, []byte("1234"), client.chunks[0].Chunk) + require.Empty(t, client.chunks[0].Chunk) require.False(t, client.chunks[0].Final) require.Empty(t, client.chunks[1].Metadata) - require.Equal(t, []byte("abcd"), client.chunks[1].Chunk) + require.Equal(t, []byte("1234"), client.chunks[1].Chunk) require.False(t, client.chunks[1].Final) require.Empty(t, client.chunks[2].Metadata) - require.Equal(t, []byte("5678"), client.chunks[2].Chunk) - require.True(t, client.chunks[2].Final) + require.Equal(t, []byte("abcd"), client.chunks[2].Chunk) + require.False(t, client.chunks[2].Final) + + require.Empty(t, client.chunks[3].Metadata) + require.Equal(t, []byte("5678"), client.chunks[3].Chunk) + require.True(t, client.chunks[3].Final) } func TestSendSnapshotReaderChunksExactBoundary(t *testing.T) { @@ -1662,13 +1753,16 @@ func TestSendSnapshotReaderChunksExactBoundary(t *testing.T) { err := sendSnapshotReaderChunks(client, header, bytes.NewReader(payload), chunkSize) require.NoError(t, err) - require.Len(t, client.chunks, 2) + require.Len(t, client.chunks, 3) require.Equal(t, header, client.chunks[0].Metadata) - require.Equal(t, []byte("1234"), client.chunks[0].Chunk) + require.Empty(t, client.chunks[0].Chunk) require.False(t, client.chunks[0].Final) - require.Equal(t, []byte("5678"), client.chunks[1].Chunk) - require.True(t, client.chunks[1].Final) + require.Equal(t, []byte("1234"), client.chunks[1].Chunk) + require.False(t, client.chunks[1].Final) + + require.Equal(t, []byte("5678"), client.chunks[2].Chunk) + require.True(t, client.chunks[2].Final) } // TestSendSnapshotReaderChunksTrailingPartialChunk regressions a production @@ -1691,17 +1785,20 @@ func TestSendSnapshotReaderChunksTrailingPartialChunk(t *testing.T) { err := sendSnapshotReaderChunks(client, header, bytes.NewReader(payload), chunkSize) require.NoError(t, err) - require.Len(t, client.chunks, 3, "expected two full chunks plus a trailing partial") + require.Len(t, client.chunks, 4, "expected metadata, two full chunks, and a trailing partial") require.Equal(t, header, client.chunks[0].Metadata) - require.Equal(t, []byte("1234"), client.chunks[0].Chunk) + require.Empty(t, client.chunks[0].Chunk) require.False(t, client.chunks[0].Final) - require.Equal(t, []byte("5678"), client.chunks[1].Chunk) + require.Equal(t, []byte("1234"), client.chunks[1].Chunk) require.False(t, client.chunks[1].Final) - require.Equal(t, []byte("X"), client.chunks[2].Chunk) - require.True(t, client.chunks[2].Final) + require.Equal(t, []byte("5678"), client.chunks[2].Chunk) + require.False(t, client.chunks[2].Final) + + require.Equal(t, []byte("X"), client.chunks[3].Chunk) + require.True(t, client.chunks[3].Final) var delivered []byte for _, c := range client.chunks { diff --git a/internal/raftengine/etcd/snapshot_spool.go b/internal/raftengine/etcd/snapshot_spool.go index 4065c4526..678d7857a 100644 --- a/internal/raftengine/etcd/snapshot_spool.go +++ b/internal/raftengine/etcd/snapshot_spool.go @@ -4,6 +4,7 @@ import ( "encoding/binary" "io" "log/slog" + "math" "os" "path/filepath" "strconv" @@ -30,6 +31,10 @@ const defaultMaxSnapshotPayloadBytes int64 = 16 << 30 // 16 GiB const maxSnapshotPayloadBytesEnvVar = "ELASTICKV_RAFT_MAX_SNAPSHOT_PAYLOAD_BYTES" +const snapshotSpoolMinFreeBytesEnvVar = "ELASTICKV_RAFT_SNAPSHOT_SPOOL_MIN_FREE_BYTES" + +const defaultReceiveSnapshotSpoolMinFreeBytes int64 = 1 << 30 // 1 GiB + // resolveMaxSnapshotPayloadBytes evaluates the env override once per spool // creation. Snapshots are infrequent enough that one Getenv + ParseInt per // spool is invisible in profiles, and resolving at construction means tests @@ -48,8 +53,24 @@ func resolveMaxSnapshotPayloadBytes() int64 { return n } +func resolveSnapshotSpoolMinFreeBytes(defaultBytes int64) int64 { + v := strings.TrimSpace(os.Getenv(snapshotSpoolMinFreeBytesEnvVar)) + if v == "" { + return defaultBytes + } + n, err := strconv.ParseInt(v, 10, 64) + if err != nil || n < 0 { + slog.Warn("invalid ELASTICKV_RAFT_SNAPSHOT_SPOOL_MIN_FREE_BYTES; using default", + "value", v, "default_bytes", defaultBytes) + return defaultBytes + } + return n +} + var errSnapshotPayloadTooLarge = errors.New("etcd raft snapshot payload exceeds limit") +var errSnapshotSpoolDiskHeadroom = errors.New("etcd raft snapshot spool insufficient disk headroom") + // snapshotSyncDir indirects fsync-on-directory through a package var so a // fault-injection test can simulate "rename succeeded but the fsync that // would persist the directory entry failed" — that's the partial-failure @@ -57,24 +78,43 @@ var errSnapshotPayloadTooLarge = errors.New("etcd raft snapshot payload exceeds // and without an injection seam there's no portable way to reproduce it. var snapshotSyncDir = syncDir +var snapshotSpoolAvailableBytes = snapshotSpoolAvailableBytesFS + const snapshotSpoolPattern = "elastickv-etcd-snapshot-*" type snapshotSpool struct { - file *os.File - path string - size int64 - maxSize int64 + file *os.File + path string + size int64 + maxSize int64 + minFreeBytes int64 } func newSnapshotSpool(dir string) (*snapshotSpool, error) { + return newSnapshotSpoolWithLimits(dir, resolveMaxSnapshotPayloadBytes(), 0) +} + +func newReceiveSnapshotSpool(dir string) (*snapshotSpool, error) { + maxSize := resolveMaxSnapshotPayloadBytes() + // Keep a fixed emergency reserve after receive-side spooling. Tying this + // reserve to the configured maximum snapshot size made the default 16 GiB + // cap also require 16 GiB of free space after every chunk; production nodes + // with enough space for a real 13 GiB FSM snapshot were rejecting the stream + // around 4 GiB and retrying forever. Operators that restore into a layout + // needing a larger reserve can still raise it with the env knob. + return newSnapshotSpoolWithLimits(dir, maxSize, resolveSnapshotSpoolMinFreeBytes(defaultReceiveSnapshotSpoolMinFreeBytes)) +} + +func newSnapshotSpoolWithLimits(dir string, maxSize, minFreeBytes int64) (*snapshotSpool, error) { file, err := os.CreateTemp(dir, snapshotSpoolPattern) if err != nil { return nil, errors.WithStack(err) } return &snapshotSpool{ - file: file, - path: file.Name(), - maxSize: resolveMaxSnapshotPayloadBytes(), + file: file, + path: file.Name(), + maxSize: maxSize, + minFreeBytes: minFreeBytes, }, nil } @@ -87,6 +127,9 @@ func (s *snapshotSpool) Write(p []byte) (int, error) { if int64(len(p)) > s.maxSize-s.size { return 0, errors.Wrapf(errSnapshotPayloadTooLarge, "adding %d bytes to current %d would exceed limit %d", len(p), s.size, s.maxSize) } + if err := s.checkDiskHeadroom(len(p)); err != nil { + return 0, err + } n, err := s.file.Write(p) s.size += int64(n) if err != nil { @@ -95,6 +138,31 @@ func (s *snapshotSpool) Write(p []byte) (int, error) { return n, nil } +func (s *snapshotSpool) checkDiskHeadroom(writeBytes int) error { + if writeBytes <= 0 || s.minFreeBytes <= 0 { + return nil + } + available, err := snapshotSpoolAvailableBytes(filepath.Dir(s.path)) + if err != nil { + return errors.WithStack(err) + } + if available < 0 { + available = 0 + } + if s.minFreeBytes > math.MaxInt64-int64(writeBytes) { + return errors.Wrapf(errSnapshotSpoolDiskHeadroom, + "write_bytes=%d min_free_bytes=%d available_bytes=%d", + writeBytes, s.minFreeBytes, available) + } + required := s.minFreeBytes + int64(writeBytes) + if available < required { + return errors.Wrapf(errSnapshotSpoolDiskHeadroom, + "write_bytes=%d min_free_bytes=%d available_bytes=%d", + writeBytes, s.minFreeBytes, available) + } + return nil +} + func (s *snapshotSpool) Bytes() ([]byte, error) { if _, err := s.file.Seek(0, io.SeekStart); err != nil { return nil, errors.WithStack(err) @@ -160,6 +228,9 @@ func (s *snapshotSpool) FinalizeAsFSMFile(fsmSnapDir string, index uint64, crc32 // so Close() is a no-op. Without the per-step clear, Close() // would attempt os.Remove(s.path) and surface a misleading // ErrNotExist that buries the real syncDir error. + if err := s.checkDiskHeadroom(fsmFooterSize); err != nil { + return err + } if err := binary.Write(s.file, binary.BigEndian, crc32c); err != nil { return errors.WithStack(err) } diff --git a/internal/raftengine/etcd/snapshot_spool_space_other.go b/internal/raftengine/etcd/snapshot_spool_space_other.go new file mode 100644 index 000000000..f9d733c4a --- /dev/null +++ b/internal/raftengine/etcd/snapshot_spool_space_other.go @@ -0,0 +1,9 @@ +//go:build !(aix || darwin || linux) + +package etcd + +import "math" + +func snapshotSpoolAvailableBytesFS(string) (int64, error) { + return math.MaxInt64, nil +} diff --git a/internal/raftengine/etcd/snapshot_spool_space_unix.go b/internal/raftengine/etcd/snapshot_spool_space_unix.go new file mode 100644 index 000000000..a2c5db370 --- /dev/null +++ b/internal/raftengine/etcd/snapshot_spool_space_unix.go @@ -0,0 +1,30 @@ +//go:build aix || darwin || linux + +package etcd + +import ( + "math" + + "github.com/cockroachdb/errors" + "golang.org/x/sys/unix" +) + +func snapshotSpoolAvailableBytesFS(path string) (int64, error) { + var st unix.Statfs_t + if err := unix.Statfs(path, &st); err != nil { + return 0, errors.WithStack(err) + } + if st.Bsize <= 0 { + return 0, nil + } + blockSize := uint64(st.Bsize) + availableBlocks := st.Bavail + if availableBlocks > uint64(math.MaxInt64)/blockSize { + return math.MaxInt64, nil + } + availableBytes := availableBlocks * blockSize + if availableBytes > uint64(math.MaxInt64) { + return math.MaxInt64, nil + } + return int64(availableBytes), nil +} diff --git a/internal/raftengine/etcd/snapshot_spool_test.go b/internal/raftengine/etcd/snapshot_spool_test.go index c8ed62474..211077312 100644 --- a/internal/raftengine/etcd/snapshot_spool_test.go +++ b/internal/raftengine/etcd/snapshot_spool_test.go @@ -22,6 +22,7 @@ func TestSnapshotSpool_DefaultCapAcceptsRealisticFSM(t *testing.T) { if testing.Short() { t.Skip("skipping: writes 1.5 GiB to a temp file") } + t.Setenv(snapshotSpoolMinFreeBytesEnvVar, "0") dir := t.TempDir() spool, err := newSnapshotSpool(dir) require.NoError(t, err) @@ -66,6 +67,7 @@ func TestSnapshotSpool_DefaultCapAcceptsRealisticFSM(t *testing.T) { func TestSnapshotSpool_OverrideViaEnv(t *testing.T) { const spoolCap = int64(4096) t.Setenv(maxSnapshotPayloadBytesEnvVar, strconv.FormatInt(spoolCap, 10)) + t.Setenv(snapshotSpoolMinFreeBytesEnvVar, "0") spool, err := newSnapshotSpool(t.TempDir()) require.NoError(t, err) @@ -95,6 +97,59 @@ func TestSnapshotSpool_OverrideInvalidFallsBack(t *testing.T) { require.Equal(t, defaultMaxSnapshotPayloadBytes, spool.maxSize) } +func TestSnapshotSpoolDefaultReserveUsesFixedEmergencyHeadroom(t *testing.T) { + const spoolCap = int64(4096) + t.Setenv(maxSnapshotPayloadBytesEnvVar, strconv.FormatInt(spoolCap, 10)) + + spool, err := newReceiveSnapshotSpool(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { _ = spool.Close() }) + + require.Equal(t, spoolCap, spool.maxSize) + require.Equal(t, defaultReceiveSnapshotSpoolMinFreeBytes, spool.minFreeBytes) +} + +func TestSnapshotSpoolMaterializeDoesNotReserveDiskHeadroom(t *testing.T) { + t.Setenv(snapshotSpoolMinFreeBytesEnvVar, "1024") + originalAvailable := snapshotSpoolAvailableBytes + snapshotSpoolAvailableBytes = func(string) (int64, error) { + return 0, nil + } + t.Cleanup(func() { snapshotSpoolAvailableBytes = originalAvailable }) + + spool, err := newSnapshotSpool(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { _ = spool.Close() }) + + n, err := spool.Write([]byte("x")) + require.NoError(t, err) + require.Equal(t, 1, n) + require.Equal(t, int64(1), spool.size) +} + +func TestSnapshotSpoolRejectsWhenReserveWouldBeConsumed(t *testing.T) { + t.Setenv(snapshotSpoolMinFreeBytesEnvVar, "1024") + originalAvailable := snapshotSpoolAvailableBytes + snapshotSpoolAvailableBytes = func(string) (int64, error) { + return 1024, nil + } + t.Cleanup(func() { snapshotSpoolAvailableBytes = originalAvailable }) + + spool, err := newReceiveSnapshotSpool(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { _ = spool.Close() }) + + n, err := spool.Write([]byte("x")) + require.Zero(t, n) + require.Error(t, err) + require.True(t, errors.Is(err, errSnapshotSpoolDiskHeadroom), "got %v", err) + require.Zero(t, spool.size) + + info, statErr := spool.file.Stat() + require.NoError(t, statErr) + require.Zero(t, info.Size(), "headroom rejection must happen before writing bytes") +} + // TestFinalizeAsFSMFile_PostFinalizeCloseIsNoop pins the gemini-medium // review on PR #747: after a successful FinalizeAsFSMFile, the deferred // caller-side spool.Close() must NOT attempt to remove the renamed file @@ -319,6 +374,48 @@ func TestPurgeOldSnapFiles(t *testing.T) { require.NoError(t, err) } +func TestPurgeOldSnapFilesHonorsFSMByteBudget(t *testing.T) { + t.Setenv(maxRetainedFSMSnapshotBytesEnvVar, "100") + + snapDir := t.TempDir() + fsmSnapDir := t.TempDir() + + for i := uint64(1); i <= 3; i++ { + index := i * 10000 + createSnapFile(t, snapDir, index) + require.NoError(t, os.WriteFile(fsmSnapPath(fsmSnapDir, index), bytes.Repeat([]byte("x"), 60), 0o600)) + } + + require.NoError(t, purgeOldSnapshotFiles(snapDir, fsmSnapDir)) + + require.NoFileExists(t, filepath.Join(snapDir, fmt.Sprintf("%016x-%016x.snap", 1, uint64(10000)))) + require.NoFileExists(t, filepath.Join(snapDir, fmt.Sprintf("%016x-%016x.snap", 1, uint64(20000)))) + require.FileExists(t, filepath.Join(snapDir, fmt.Sprintf("%016x-%016x.snap", 1, uint64(30000)))) + require.NoFileExists(t, fsmSnapPath(fsmSnapDir, 10000)) + require.NoFileExists(t, fsmSnapPath(fsmSnapDir, 20000)) + require.FileExists(t, fsmSnapPath(fsmSnapDir, 30000)) +} + +func TestPurgeOldSnapFilesBudgetCanBeDisabled(t *testing.T) { + t.Setenv(maxRetainedFSMSnapshotBytesEnvVar, "0") + + snapDir := t.TempDir() + fsmSnapDir := t.TempDir() + + for i := uint64(1); i <= 4; i++ { + index := i * 10000 + createSnapFile(t, snapDir, index) + require.NoError(t, os.WriteFile(fsmSnapPath(fsmSnapDir, index), bytes.Repeat([]byte("x"), 60), 0o600)) + } + + require.NoError(t, purgeOldSnapshotFiles(snapDir, fsmSnapDir)) + + require.NoFileExists(t, filepath.Join(snapDir, fmt.Sprintf("%016x-%016x.snap", 1, uint64(10000)))) + require.FileExists(t, filepath.Join(snapDir, fmt.Sprintf("%016x-%016x.snap", 1, uint64(20000)))) + require.FileExists(t, filepath.Join(snapDir, fmt.Sprintf("%016x-%016x.snap", 1, uint64(30000)))) + require.FileExists(t, filepath.Join(snapDir, fmt.Sprintf("%016x-%016x.snap", 1, uint64(40000)))) +} + func TestPurgeOldSnapFilesUnderLimit(t *testing.T) { snapDir := t.TempDir() fsmSnapDir := t.TempDir() diff --git a/internal/raftengine/etcd/wal_store.go b/internal/raftengine/etcd/wal_store.go index 0612ba5f0..32c7cc88f 100644 --- a/internal/raftengine/etcd/wal_store.go +++ b/internal/raftengine/etcd/wal_store.go @@ -2,7 +2,6 @@ package etcd import ( "bytes" - "hash/crc32" "io" "os" "path/filepath" @@ -142,17 +141,20 @@ func loadWalState(logger *zap.Logger, walDir, snapDir, fsmSnapDir string, fsm St return nil, err } - // Open the WAL before the restore decision so we know the committed - // replay target. The FSM restore can be skipped as soon as the - // durable FSM is at least at the snapshot pointer: Open seeds the - // engine's applied counter from EffectiveApplied, and WAL replay - // then delivers only committed entries above that durable point. + // Open the WAL before the skip-gate decision so we know the committed + // replay target for metrics and raw-node initialization. The skip + // gate itself only requires the FSM to be at least as fresh as the + // persisted snapshot index: Engine.Open seeds e.applied with the + // FSM's durable applied index, rawNodeAppliedForOpen trims the RawNode + // applied pointer when volatile/conf-change entries must replay, and + // applyNormalCommitted drops data-mutating duplicates while applying + // the remaining committed tail. w, hardState, entries, err := openAndReadWALWithRepair(logger, walDir, walSnapshotFor(snapshot)) if err != nil { return nil, err } - lastCommittedIndex := coldStartSkipThreshold(snapshot, hardState) - effectiveApplied, err := restoreSnapshotState(fsm, snapshot, lastCommittedIndex, fsmSnapDir, obs, logger) + committedReplayTarget := coldStartReplayTarget(snapshot, hardState) + effectiveApplied, err := restoreSnapshotState(fsm, snapshot, committedReplayTarget, fsmSnapDir, obs, logger) if err != nil { if closeErr := w.Close(); closeErr != nil { logger.Warn("WAL close failed after restoreSnapshotState error", @@ -197,50 +199,45 @@ func reportColdStartExecute(obs raftengine.ColdStartObserver, logger *zap.Logger if logger == nil { return } - // gap_to_snapshot uses absolute value because stale metadata can - // still report an FSM ahead of the snapshot while the execute path - // runs due to an earlier fallback. gap_behind_committed is clamped - // to avoid underflow when tests call this helper with a lower - // synthetic target. + // gap_to_snapshot uses absolute value because the gate now + // permits have > snapIndex (FSM ahead of snapshot but behind + // committed tail). gap_behind_committed is target-have; can be + // 0 when have==target. var gapToSnapshot uint64 if have >= snapIndex { gapToSnapshot = have - snapIndex } else { gapToSnapshot = snapIndex - have } - var gapBehindCommitted uint64 - if target > have { - gapBehindCommitted = target - have - } logger.Info("restoreSnapshotState executed (FSM behind WAL committed tail)", zap.Uint64("fsm_applied", have), zap.Uint64("snapshot_index", snapIndex), zap.Uint64("last_committed_index", target), zap.Uint64("gap_to_snapshot", gapToSnapshot), - zap.Uint64("gap_behind_committed", gapBehindCommitted), + zap.Uint64("gap_behind_committed", target-have), ) } -// coldStartSkipThreshold returns the maximum log index the cold- -// start replay can deliver via Ready.CommittedEntries on this -// node: max(snapshot.Metadata.Index, hardState.Commit). This is the -// committed replay target used for observability; the snapshot body -// restore decision itself only requires the durable FSM to be at -// least at snapshot.Metadata.Index. +// coldStartReplayTarget returns the maximum log index cold-start replay +// can deliver via Ready.CommittedEntries on this node: +// max(snapshot.Metadata.Index, hardState.Commit). The skip gate uses the +// snapshot index for safety; this target is still reported in metrics and +// used by RawNode initialization to replay only the entries still needed +// after e.applied is seeded from the FSM's durable applied index. // // Followers can carry an UNCOMMITTED WAL suffix // (entries[n-1].Index > hardState.Commit). Raft does NOT surface // those entries in CommittedEntries until the leader confirms -// them. The previous gate used the WAL tail (entries[n-1].Index) +// them. The previous target used the WAL tail (entries[n-1].Index) // which forced a multi-GiB restore on every restart of any // follower with an uncommitted suffix, defeating the cold-start -// optimization. Codex P2 #934 round 3. +// optimization. // // The lower bound stays at the snapshot pointer because an empty // WAL still requires the FSM to be at least at the snapshot // index. Raft's invariant guarantees hardState.Commit >= 0; we do // not need to bound from below explicitly beyond snap.Index. -func coldStartSkipThreshold(snapshot raftpb.Snapshot, hardState raftpb.HardState) uint64 { +func coldStartReplayTarget(snapshot raftpb.Snapshot, hardState raftpb.HardState) uint64 { threshold := snapshot.GetMetadata().GetIndex() if hardState.GetCommit() > threshold { threshold = hardState.GetCommit() @@ -253,7 +250,8 @@ func coldStartSkipThreshold(snapshot raftpb.Snapshot, hardState raftpb.HardState // skip gate fires with the FSM at `have > snapshot.Metadata.Index`, // EffectiveApplied carries `have`; without this seed the engine // would deliver entries snapshot.Index+1..have to applyCommitted -// and re-apply them onto a Pebble store already containing them. +// and re-apply them onto a Pebble store already containing them +// (codex P1 #934 root cause). func coldStartApplied(disk *diskState) uint64 { base := maxAppliedIndex(disk.LocalSnap) if disk.EffectiveApplied > base { @@ -356,12 +354,12 @@ func loadPersistedSnapshot(logger *zap.Logger, walDir string, snapshotter *snap. // <= have or the Pebble store would observe them twice (OCC // conflicts; HLC ceiling inversion). The execute path returns // snapshot.Metadata.Index to leave engine behaviour unchanged. -func restoreSnapshotState(fsm StateMachine, snapshot raftpb.Snapshot, committedTailIndex uint64, fsmSnapDir string, obs raftengine.ColdStartObserver, logger *zap.Logger) (uint64, error) { +func restoreSnapshotState(fsm StateMachine, snapshot raftpb.Snapshot, replayTarget uint64, fsmSnapDir string, obs raftengine.ColdStartObserver, logger *zap.Logger) (uint64, error) { if etcdraft.IsEmptySnap(&snapshot) || len(snapshot.Data) == 0 || fsm == nil { return 0, nil } if isSnapshotToken(snapshot.Data) { - return restoreSnapshotStateFromToken(fsm, snapshot, committedTailIndex, fsmSnapDir, obs, logger) + return restoreSnapshotStateFromToken(fsm, snapshot, replayTarget, fsmSnapDir, obs, logger) } // Legacy format: full FSM payload embedded in snapshot.Data. if err := fsm.Restore(bytes.NewReader(snapshot.Data)); err != nil { @@ -376,33 +374,33 @@ func restoreSnapshotState(fsm StateMachine, snapshot raftpb.Snapshot, committedT // - skip path: `have` (FSM is already past snapshot.Metadata.Index) // - execute path: snapshot.Metadata.Index (restored from snapshot) // -// The skip threshold is snapshot.Metadata.Index: once the FSM is at -// the snapshot pointer, the engine can seed its applied counter from -// the durable FSM index and replay only the committed WAL suffix above -// that point. +// The skip threshold is tok.Index: if the FSM already contains the +// snapshot state, the engine can seed e.applied with the FSM's durable +// applied index and let normal replay apply any committed tail after +// that point. Entries at or below e.applied are handled by the existing +// duplicate-replay seam in applyCommitted. // // Metrics + log fire AFTER the restore-side work succeeds (coderabbit // Major #934): a header/CRC failure must not register a "successful" // outcome in the soak metrics. -func restoreSnapshotStateFromToken(fsm StateMachine, snapshot raftpb.Snapshot, committedTailIndex uint64, fsmSnapDir string, obs raftengine.ColdStartObserver, logger *zap.Logger) (uint64, error) { +func restoreSnapshotStateFromToken(fsm StateMachine, snapshot raftpb.Snapshot, replayTarget uint64, fsmSnapDir string, obs raftengine.ColdStartObserver, logger *zap.Logger) (uint64, error) { tok, err := decodeSnapshotToken(snapshot.Data) if err != nil { return 0, err } - snapIndex := snapshot.GetMetadata().GetIndex() - decision, have := decideSkipOutcome(fsm, snapIndex) + decision, have := decideSkipOutcome(fsm, tok.Index) snapPath := fsmSnapPath(fsmSnapDir, tok.Index) if decision == coldStartSkip { if err := applyHeaderStateOnSkip(fsm, snapPath, tok.CRC32C); err != nil { return 0, err } - reportColdStart(obs, logger, decision, snapIndex, committedTailIndex, have) + reportColdStart(obs, logger, decision, tok.Index, replayTarget, have) return have, nil } if err := openAndRestoreFSMSnapshot(fsm, snapPath, tok.CRC32C); err != nil { return 0, err } - reportColdStart(obs, logger, decision, snapIndex, committedTailIndex, have) + reportColdStart(obs, logger, decision, tok.Index, replayTarget, have) return snapshot.GetMetadata().GetIndex(), nil } @@ -462,27 +460,21 @@ func reportColdStart(obs raftengine.ColdStartObserver, logger *zap.Logger, d col case coldStartSkip: // Observer contract (cold_start.go + monitoring/cold_start.go // Prometheus impl): args are (snapshotIndex, haveAppliedIndex); - // gauges compute have-snapIndex. Codex P2 + coderabbit Major - // #934: do NOT pass target/lastWalIndex here or the exported + // gauges compute have-snapIndex. Do NOT pass the replay target + // here or the exported // gauge measures the wrong baseline. if obs != nil { obs.RestoreSkipped(snapIndex, have) } if logger != nil { - var gapAheadCommitted uint64 - var gapBehindCommitted uint64 - if have >= target { - gapAheadCommitted = have - target - } else { - gapBehindCommitted = target - have - } - // Two named gap fields so an operator correlating the - // log against the Prometheus gauge sees consistent - // magnitudes (claude #934 round 5): + gapAheadCommitted, gapBehindCommitted := coldStartCommittedGaps(target, have) + // Separate gap fields let operators correlate the log with + // the Prometheus gauge without unsigned underflow: // - gap_ahead_snapshot mirrors monitoring.ColdStartObserver // (have - snapIndex), the metric baseline. // - gap_ahead_committed measures how far past the WAL // committed tail (target) the FSM is. + // - gap_behind_committed measures the replay tail still pending. logger.Info("restoreSnapshotState skipped", zap.Uint64("fsm_applied", have), zap.Uint64("snapshot_index", snapIndex), @@ -510,21 +502,23 @@ func reportColdStart(obs raftengine.ColdStartObserver, logger *zap.Logger, d col } } -// applyHeaderStateOnSkip mirrors openAndRestoreFSMSnapshot's safety -// contract (size + footer-vs-tokenCRC + full-body-CRC) but applies -// only the header side-effects (HLC ceiling + Stage 8a cutover) -// instead of running the body restore. The body bytes are read for -// CRC coverage but discarded -- fsm.db already holds equivalent -// state, which is precisely the reason we're skipping the restore. +func coldStartCommittedGaps(target, have uint64) (ahead, behind uint64) { + if have >= target { + return have - target, 0 + } + return 0, target - have +} + +// applyHeaderStateOnSkip validates the snapshot envelope and full payload CRC, +// then applies only the header side-effects (HLC ceiling + Stage 8a cutover) +// instead of running the body restore. The FSM body is already durable enough +// to skip restoring, but header side-effects still come from this file and +// must not be applied from corrupt bytes. // // FSMs that do not implement raftengine.SnapshotHeaderApplier // silently no-op the apply phase -- the FSM has no header state to -// carry forward, and the CRC verification still runs (with no -// observable side-effect on success). On any verification failure -// the typed error propagates and FSM state stays untouched. -// -// See PR #910 design §5 round-7 (two-phase seam) + round-6 -// (three-step CRC mirroring openAndRestoreFSMSnapshot). +// carry forward. On any envelope or header parse failure the typed error +// propagates and FSM state stays untouched. func applyHeaderStateOnSkip(fsm StateMachine, snapPath string, tokenCRC uint32) error { file, err := os.Open(snapPath) if err != nil { @@ -540,57 +534,35 @@ func applyHeaderStateOnSkip(fsm StateMachine, snapPath string, tokenCRC uint32) if err != nil { return err } - - // Step 3: full-body CRC. Wrap the payload in a crc32 TeeReader - // and hand it to the FSM's ParseSnapshotHeader for header parse - // + drain. Every payload byte flows through h, matching - // restoreAndComputeCRC's boundary in openAndRestoreFSMSnapshot. - // - // Error-ordering contract (claude #934 R1-F1): header parse - // errors surface BEFORE the body-CRC compare runs, so callers - // (the skip-gate fallback in restoreSnapshotState) may observe - // either an ErrSnapshotHeaderUnknownMagic / InvalidLength chain or - // an ErrFSMSnapshotFileCRC chain depending on which check fails - // first. This is the same ordering openAndRestoreFSMSnapshot has - // — both errors are equally fatal for the skip path (they signal - // snapshot file corruption) and both must propagate without ever - // calling ApplySnapshotHeader. The CRC check stays AFTER the - // header parse so the TeeReader has actually been drained before - // we read h.Sum32(); inverting the order would let a CRC mismatch - // surface on a truncated body even when the header was valid, - // muddying the operator-facing diagnostic. - if _, err := file.Seek(0, io.SeekStart); err != nil { - return errors.WithStack(err) - } - payloadSize := info.Size() - fsmFooterSize - h := crc32.New(crc32cTable) - tee := io.TeeReader(io.LimitReader(file, payloadSize), h) - - setter, hasSetter := fsm.(raftengine.SnapshotHeaderApplier) - ceiling, cutover, err := readSnapshotHeaderOrDrain(setter, hasSetter, tee) + computed, err := computeFSMSnapshotPayloadCRC(file, info.Size()) if err != nil { return err } - - if h.Sum32() != footer { + if computed != footer { return errors.Wrapf(ErrFSMSnapshotFileCRC, - "path=%s footer=%08x computed=%08x", snapPath, footer, h.Sum32()) + "path=%s footer=%08x computed=%08x", snapPath, footer, computed) + } + + setter, hasSetter := fsm.(raftengine.SnapshotHeaderApplier) + if !hasSetter { + return nil } - // All three checks passed; apply side-effects (pure assignment - // in the FSM). Skipped silently when the FSM does not expose - // the seam. - if hasSetter { - setter.ApplySnapshotHeader(ceiling, cutover) + if _, err := file.Seek(0, io.SeekStart); err != nil { + return errors.WithStack(err) } + payloadSize := info.Size() - fsmFooterSize + ceiling, cutover, err := setter.ParseSnapshotHeader(io.LimitReader(file, payloadSize)) + if err != nil { + return errors.WithStack(err) + } + setter.ApplySnapshotHeader(ceiling, cutover) return nil } -// verifyFSMSnapshotPrefix runs the first two cheap checks of -// openAndRestoreFSMSnapshot's three-step contract: size and -// footer-vs-tokenCRC. Returns the on-disk footer value (caller -// reuses it for the step-3 full-body CRC compare). Typed errors -// surface unchanged. +// verifyFSMSnapshotPrefix runs the cheap checks shared with +// openAndRestoreFSMSnapshot: size and footer-vs-tokenCRC. Returns the +// on-disk footer value. Typed errors surface unchanged. func verifyFSMSnapshotPrefix(file *os.File, fileSize int64, snapPath string, tokenCRC uint32) (uint32, error) { if fileSize < fsmMinFileSize { return 0, errors.Wrapf(ErrFSMSnapshotTooSmall, @@ -607,27 +579,6 @@ func verifyFSMSnapshotPrefix(file *os.File, fileSize int64, snapPath string, tok return footer, nil } -// readSnapshotHeaderOrDrain branches on whether the FSM exposes the -// SnapshotHeaderApplier seam: when present, delegate to -// ParseSnapshotHeader (which parses the header AND drains the rest); -// otherwise drain the entire payload through the tee'd reader so the -// CRC pass covers every byte. The (ceiling, cutover) tuple is zero -// in the no-seam case -- the caller's ApplySnapshotHeader branch -// short-circuits on hasSetter, so the zero values are inert. -func readSnapshotHeaderOrDrain(setter raftengine.SnapshotHeaderApplier, hasSetter bool, tee io.Reader) (uint64, uint64, error) { - if hasSetter { - ceiling, cutover, err := setter.ParseSnapshotHeader(tee) - if err != nil { - return 0, 0, errors.WithStack(err) - } - return ceiling, cutover, nil - } - if _, err := io.Copy(io.Discard, tee); err != nil { - return 0, 0, errors.WithStack(err) - } - return 0, 0, nil -} - func walSnapshotFor(snapshot raftpb.Snapshot) walpb.Snapshot { return walpb.Snapshot{ Index: proto.Uint64(snapshot.GetMetadata().GetIndex()), diff --git a/internal/raftengine/etcd/wal_store_skip_gate_test.go b/internal/raftengine/etcd/wal_store_skip_gate_test.go index 6348a46a0..5f331156f 100644 --- a/internal/raftengine/etcd/wal_store_skip_gate_test.go +++ b/internal/raftengine/etcd/wal_store_skip_gate_test.go @@ -53,17 +53,13 @@ func (f *skipGateFSM) ParseSnapshotHeader(r io.Reader) (uint64, uint64, error) { if f.parseErr != nil { return 0, 0, f.parseErr } - // Mimic the real kvFSM contract: parse + drain. We don't actually - // parse a header here; the test fixtures embed magic+ceiling but - // for the gate-level tests we just drain so the CRC matches. + // Mimic the real kvFSM contract: parse only the header. The skip path + // must not drain multi-GiB snapshot bodies just to seed header state. hdrLen := 16 hdr := make([]byte, hdrLen) if n, _ := io.ReadFull(r, hdr); n == hdrLen && bytes.HasPrefix(hdr, []byte("EKVTHLC1")) { f.parsedCeiling = binary.BigEndian.Uint64(hdr[8:16]) } - if _, err := io.Copy(io.Discard, r); err != nil { - return 0, 0, err - } return f.parsedCeiling, 0, nil } @@ -177,10 +173,10 @@ func TestSkipGate_ExecutesWhenFSMStale(t *testing.T) { require.Empty(t, obs.fallbacks) } -// TestSkipGate_ReturnsEffectiveAppliedOnSkip verifies that the skip -// path returns `have` so Engine.Open can seed e.applied above -// snapshot.Index, preventing applyCommitted from re-delivering the -// snapshot..have tail. +// TestSkipGate_ReturnsEffectiveAppliedOnSkip pins codex P1 #934 +// round 2. The skip path MUST return `have` so Engine.Open can seed +// e.applied above snapshot.Index, preventing applyCommitted from +// re-delivering the snapshot..have tail. func TestSkipGate_ReturnsEffectiveAppliedOnSkip(t *testing.T) { dir := t.TempDir() const ( @@ -230,13 +226,10 @@ func TestSkipGate_EmitsAfterSuccess(t *testing.T) { require.Empty(t, obs.fallbacks) } -// TestColdStartSkipThreshold verifies that the threshold caps at -// hardState.Commit so a follower carrying an -// uncommitted WAL suffix is NOT forced to run the full restore -// every restart (the original gate used the WAL tail, which can -// exceed Commit, and raft would not deliver those entries until -// the leader confirmed them). -func TestColdStartSkipThreshold(t *testing.T) { +// TestColdStartReplayTarget verifies the committed replay target caps at +// hardState.Commit so a follower carrying an uncommitted WAL suffix does not +// report or initialize from entries raft cannot deliver yet. +func TestColdStartReplayTarget(t *testing.T) { t.Parallel() mkSnap := func(idx uint64) raftpb.Snapshot { return raftTestSnapshot(idx, 0, nil, nil) @@ -253,27 +246,53 @@ func TestColdStartSkipThreshold(t *testing.T) { {"hardState.Commit zero", mkSnap(100), testHardState(0, 0), 100}, } for _, c := range cases { - got := coldStartSkipThreshold(c.snap, c.hs) + got := coldStartReplayTarget(c.snap, c.hs) if got != c.expected { t.Errorf("%s: got %d, want %d", c.name, got, c.expected) } } } -// TestSkipGate_SkipsWhenFSMAtSnapshotButBehindCommitTail verifies -// the cold-start fast path for the normal crash window: the FSM is -// already at or beyond the snapshot pointer, but durable Raft commit -// is ahead. The snapshot body must still be skipped; Engine.Open uses -// the returned EffectiveApplied to replay only the committed WAL suffix -// above the FSM's durable applied index. -func TestSkipGate_SkipsWhenFSMAtSnapshotButBehindCommitTail(t *testing.T) { +func TestColdStartCommittedGaps(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name string + target uint64 + have uint64 + wantAhead uint64 + wantBehind uint64 + }{ + {name: "ahead", target: 100, have: 150, wantAhead: 50}, + {name: "equal", target: 100, have: 100}, + {name: "replay tail remains", target: 150, have: 100, wantBehind: 50}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + ahead, behind := coldStartCommittedGaps(tc.target, tc.have) + require.Equal(t, tc.wantAhead, ahead) + require.Equal(t, tc.wantBehind, behind) + }) + } +} + +// TestSkipGate_SkipsWhenWALCarriesPostSnapshotTail verifies a node can skip +// the multi-GiB snapshot restore even when the committed WAL has entries past +// the FSM's durable applied index. Engine.Open seeds e.applied with `have`, +// then RawNode/applyCommitted replays only the needed tail and drops +// data-mutating duplicates at or below `have`. +func TestSkipGate_SkipsWhenWALCarriesPostSnapshotTail(t *testing.T) { dir := t.TempDir() const ( - snapIndex uint64 = 100 - appliedIdx uint64 = 150 - committedTailIndex uint64 = 200 + snapIndex uint64 = 100 + appliedIdx uint64 = 150 + replayTarget uint64 = 200 + ceilingMs uint64 = 1700_000_000_002 ) - payload := []byte("body-bytes-for-skip") + payload := make([]byte, 16, 16+len("body-bytes-after-header")) + copy(payload[:8], "EKVTHLC1") + binary.BigEndian.PutUint64(payload[8:], ceilingMs) + payload = append(payload, []byte("body-bytes-after-header")...) crc, _ := writeFSMFileForTest(t, dir, snapIndex, payload) fsm := &skipGateFSM{applied: appliedIdx, appliedPresent: true} @@ -282,13 +301,13 @@ func TestSkipGate_SkipsWhenFSMAtSnapshotButBehindCommitTail(t *testing.T) { Metadata: testSnapshotMetadata(snapIndex, 0, nil), } obs := &recordingObs{} - effective, gateErr := restoreSnapshotState(fsm, snap, committedTailIndex, dir, obs, nil) + effective, gateErr := restoreSnapshotState(fsm, snap, replayTarget, dir, obs, nil) require.NoError(t, gateErr) - require.Equal(t, appliedIdx, effective, - "skip path MUST return the durable FSM applied index so WAL replay starts above it") require.Empty(t, fsm.bodyBytes, "skip path MUST NOT call fsm.Restore") - require.True(t, fsm.restoredHeader, "skip path MUST still apply snapshot header state") + require.True(t, fsm.restoredHeader) + require.Equal(t, ceilingMs, fsm.appliedCeiling) + require.Equal(t, appliedIdx, effective) require.Equal(t, []uint64{appliedIdx - snapIndex}, obs.skipped) require.Empty(t, obs.executed) require.Empty(t, obs.fallbacks) @@ -379,30 +398,33 @@ func TestApplyHeaderStateOnSkip_WrongTokenCRC(t *testing.T) { require.False(t, fsm.restoredHeader, "FSM state MUST NOT mutate on verification failure") } -// TestApplyHeaderStateOnSkip_BodyCorruption asserts step 3 catches a -// flipped body byte (CRC mismatch). -func TestApplyHeaderStateOnSkip_BodyCorruption(t *testing.T) { +// TestApplyHeaderStateOnSkip_VerifiesPayloadBeforeHeaderApply asserts the skip +// path checks payload integrity before applying header side-effects. +func TestApplyHeaderStateOnSkip_VerifiesPayloadBeforeHeaderApply(t *testing.T) { dir := t.TempDir() - crc, path := writeFSMFileForTest(t, dir, 1, []byte("payload-bytes")) + const ceilingMs uint64 = 1700_000_000_001 + payload := make([]byte, 16, 16+len("payload-bytes")) + copy(payload[:8], "EKVTHLC1") + binary.BigEndian.PutUint64(payload[8:], ceilingMs) + payload = append(payload, []byte("payload-bytes")...) + crc, path := writeFSMFileForTest(t, dir, 1, payload) - // Flip the first byte of the body in-place. The footer still - // reads as `crc`, but the on-the-wire content no longer matches - // it. Step 2 (footer-vs-token) passes (we pass the same `crc` - // as tokenCRC), step 3 (full-body CRC) fails. + // Flip a byte in the fixed 16-byte header. The footer still reads as `crc`, + // but the bytes used to derive header side-effects no longer match it. f, err := os.OpenFile(path, os.O_RDWR, 0) require.NoError(t, err) defer f.Close() var b [1]byte - _, err = f.ReadAt(b[:], 0) + _, err = f.ReadAt(b[:], 8) require.NoError(t, err) b[0] ^= 0x01 - _, err = f.WriteAt(b[:], 0) + _, err = f.WriteAt(b[:], 8) require.NoError(t, err) fsm := &skipGateFSM{} err = applyHeaderStateOnSkip(fsm, path, crc) require.ErrorIs(t, err, ErrFSMSnapshotFileCRC) - require.False(t, fsm.restoredHeader, "FSM state MUST NOT mutate on verification failure") + require.False(t, fsm.restoredHeader) } // --- kvFSM header preservation contract --- diff --git a/internal/raftengine/statemachine.go b/internal/raftengine/statemachine.go index d834ad318..ec14dc517 100644 --- a/internal/raftengine/statemachine.go +++ b/internal/raftengine/statemachine.go @@ -69,24 +69,25 @@ type AppliedIndexReader interface { } // AppliedIndexWriter is an OPTIONAL extension that lets the engine -// pin the FSM's durable applied-index to a known value at snapshot -// persist time. See docs/design/2026_06_02_implemented_idempotent_snapshot_restore.md +// pin the FSM's durable applied-index to a known value at raft +// durability boundaries. See docs/design/2026_06_02_implemented_idempotent_snapshot_restore.md // §6 "HLC lease entries — checkpoint at snapshot persist". // -// The engine calls SetDurableAppliedIndex(snap.Metadata.Index) -// before it calls persist.SaveSnap, so that on every successful -// snapshot persist the invariant `LastAppliedIndex >= -// snapshot.Metadata.Index` holds unconditionally — closing the -// HLC-lease-only / encryption-only fallback that would otherwise -// leave LastAppliedIndex stuck at the last data-Apply index. +// The engine calls SetDurableAppliedIndex at local snapshot persist, +// after received-snapshot WAL persistence, and at startup +// committed-tail drain boundaries, so the invariant +// `LastAppliedIndex >= the locally durable raft state` holds once the +// engine is store-ready. This closes the HLC-lease-only / +// encryption-only fallback that would otherwise leave LastAppliedIndex +// stuck at the last data-Apply index. // // Implementations MUST persist the value with pebble.Sync (or the // equivalent strong-durability flag for the backing store) // regardless of ELASTICKV_FSM_SYNC_MODE. The checkpoint is the only -// durable carrier of metaAppliedIndex at this point — once -// persist.SaveSnap returns, WAL compaction discards every log entry -// at or before snap.Metadata.Index, so there is no source to replay -// the meta key bump from. +// durable carrier of metaAppliedIndex at local snapshot persist time +// — once persist.SaveSnap returns, WAL compaction discards every log +// entry at or before snap.Metadata.Index, so there is no source to +// replay the meta key bump from. type AppliedIndexWriter interface { SetDurableAppliedIndex(idx uint64) error } @@ -99,24 +100,22 @@ type AppliedIndexWriter interface { // // The interface is two-phase by design: // -// - ParseSnapshotHeader reads the v1/v2 header from a caller- -// supplied io.Reader (wrapped in a crc32 TeeReader by the -// engine) and drains the remaining bytes so the wrapping CRC -// covers the full payload. It returns the parsed (ceiling, -// cutover) pair WITHOUT mutating FSM state. Errors propagate -// from the underlying header parser -// (ErrSnapshotHeaderUnknownMagic / InvalidLength) or from the -// drain pass (I/O errors); FSM state stays untouched on error. +// - ParseSnapshotHeader reads only the v1/v2 header from a caller- +// supplied io.Reader. It returns the parsed (ceiling, cutover) pair +// WITHOUT mutating FSM state. Errors propagate from the underlying +// header parser (ErrSnapshotHeaderUnknownMagic / InvalidLength); +// FSM state stays untouched on error. // // - ApplySnapshotHeader is pure assignment of the verified header // state. The engine calls this only after ParseSnapshotHeader -// returned successfully AND the wrapping crc32 hash matched -// the file footer. +// returned successfully and the snapshot file's footer matches the +// raft token. // -// Splitting parse from apply lets the CRC verifier stay co-located -// with its private helpers in internal/raftengine/etcd (matching -// the openAndRestoreFSMSnapshot safety contract) while the v1/v2 -// header parser stays inside the kv package where it already lives. +// Splitting parse from apply keeps the v1/v2 header parser inside the +// kv package where it already lives, while the engine remains responsible +// for deciding whether the snapshot body is actually needed. Full-body CRC +// verification still happens on the restore path; the skip path only reads +// the header because the FSM body state is already present locally. // Neither package imports the other in production. type SnapshotHeaderApplier interface { ParseSnapshotHeader(r io.Reader) (ceiling, cutover uint64, err error) diff --git a/internal/s3keys/keys.go b/internal/s3keys/keys.go index 1ee2a09e6..829577f3f 100644 --- a/internal/s3keys/keys.go +++ b/internal/s3keys/keys.go @@ -84,6 +84,17 @@ func ParseBucketMetaKey(key []byte) (string, bool) { return string(segment), true } +func ParseBucketGenerationKey(key []byte) (string, bool) { + if !bytes.HasPrefix(key, bucketGenerationPrefixBytes) { + return "", false + } + segment, next, ok := decodeSegment(key, len(bucketGenerationPrefixBytes)) + if !ok || next != len(key) { + return "", false + } + return string(segment), true +} + func ObjectManifestKey(bucket string, generation uint64, object string) []byte { return buildObjectKey(objectManifestPrefixBytes, bucket, generation, object, "", 0, 0) } @@ -270,6 +281,50 @@ func RoutePrefixForBucket(bucket string, generation uint64) []byte { return bucketScopedPrefix(routePrefixBytes, bucket, generation) } +func RoutePrefixForBucketAnyGeneration(bucket string) []byte { + out := make([]byte, 0, len(RoutePrefix)+len(bucket)+segmentEscapeOverhead) + out = append(out, routePrefixBytes...) + out = append(out, EncodeSegment([]byte(bucket))...) + return out +} + +func BucketGenerationRoutePrefixForCleanupPrefix(prefix []byte) ([]byte, bool) { + familyPrefix := bucketGenerationFamilyPrefix(prefix) + if familyPrefix == nil { + return nil, false + } + bucketRaw, next, ok := decodeSegment(prefix, len(familyPrefix)) + if !ok { + return nil, false + } + generation, next, ok := readU64(prefix, next) + if !ok || next != len(prefix) { + return nil, false + } + return RoutePrefixForBucket(string(bucketRaw), generation), true +} + +func bucketGenerationFamilyPrefix(key []byte) []byte { + switch { + case bytes.HasPrefix(key, objectManifestPrefixBytes): + return objectManifestPrefixBytes + case bytes.HasPrefix(key, uploadMetaPrefixBytes): + return uploadMetaPrefixBytes + case bytes.HasPrefix(key, uploadPartPrefixBytes): + return uploadPartPrefixBytes + case bytes.HasPrefix(key, blobPrefixBytes): + return blobPrefixBytes + case bytes.HasPrefix(key, chunkRefPrefixBytes): + return chunkRefPrefixBytes + case bytes.HasPrefix(key, gcUploadPrefixBytes): + return gcUploadPrefixBytes + case bytes.HasPrefix(key, routePrefixBytes): + return routePrefixBytes + default: + return nil + } +} + func bucketScopedPrefix(prefix []byte, bucket string, generation uint64) []byte { out := make([]byte, 0, len(prefix)+len(bucket)+u64Bytes+segmentEscapeOverhead) out = append(out, prefix...) diff --git a/internal/s3keys/keys_test.go b/internal/s3keys/keys_test.go index 97431d286..ff602c849 100644 --- a/internal/s3keys/keys_test.go +++ b/internal/s3keys/keys_test.go @@ -19,6 +19,17 @@ func TestBucketMetaKey_RoundTripsZeroByteSegments(t *testing.T) { require.Equal(t, bucket, parsed) } +func TestBucketGenerationKey_RoundTripsZeroByteSegments(t *testing.T) { + t.Parallel() + + bucket := string([]byte{'b', 'u', 0x00, 'c', 'k', 'e', 't'}) + key := BucketGenerationKey(bucket) + + parsed, ok := ParseBucketGenerationKey(key) + require.True(t, ok) + require.Equal(t, bucket, parsed) +} + func TestObjectManifestKey_RoundTripsZeroByteSegments(t *testing.T) { t.Parallel() @@ -53,6 +64,17 @@ func TestExtractRouteKey_ObjectScopedKeys(t *testing.T) { } } +func TestRoutePrefixForBucketAnyGeneration(t *testing.T) { + t.Parallel() + + bucket := string([]byte{'b', 0x00, 'k'}) + prefix := RoutePrefixForBucketAnyGeneration(bucket) + + require.True(t, bytes.HasPrefix(RouteKey(bucket, 1, "a"), prefix)) + require.True(t, bytes.HasPrefix(RouteKey(bucket, 2, "b"), prefix)) + require.False(t, bytes.HasPrefix(RouteKey("other", 1, "a"), prefix)) +} + func TestManifestScanRouteBounds(t *testing.T) { t.Parallel() @@ -504,3 +526,13 @@ func TestPerBucketPrefixes_IsolateByBucketAndGeneration(t *testing.T) { }) } } + +func TestBucketGenerationRoutePrefixForCleanupPrefixIncludesChunkRefs(t *testing.T) { + t.Parallel() + + prefix := ChunkRefPrefixForBucket("bucket-a", 7) + got, ok := BucketGenerationRoutePrefixForCleanupPrefix(prefix) + + require.True(t, ok) + require.Equal(t, RoutePrefixForBucket("bucket-a", 7), got) +} diff --git a/kv/coordinator.go b/kv/coordinator.go index 093d77686..a9b392890 100644 --- a/kv/coordinator.go +++ b/kv/coordinator.go @@ -37,15 +37,20 @@ const dispatchLeaderRetryInterval = 25 * time.Millisecond // hlcPhysicalWindowMs is the duration in milliseconds that the Raft-agreed // physical ceiling extends ahead of the current wall clock. Modelled after -// TiDB's TSO 3-second window: the leader commits ceiling = now + window, and +// TiDB's TSO window strategy: the leader commits ceiling = now + window, and // renews before the window expires. A new leader inherits the committed ceiling // so it never issues timestamps that collide with the previous leader's window. -const hlcPhysicalWindowMs int64 = 3_000 +const hlcPhysicalWindowMs int64 = 15_000 // hlcRenewalInterval controls how often the leader proposes a new ceiling. // Must be less than hlcPhysicalWindowMs to guarantee the window never expires. const hlcRenewalInterval = 1 * time.Second +// hlcRenewalTimeout bounds a single renewal proposal. It is intentionally +// longer than hlcRenewalInterval so transient Raft write backlog does not +// cancel the renewal before the physical ceiling has real risk of expiring. +const hlcRenewalTimeout = 5 * time.Second + // CoordinatorOption is a functional option for Coordinate constructors. type CoordinatorOption func(*Coordinate) @@ -855,10 +860,10 @@ func (c *Coordinate) extendLeaseAfterRenewal(dispatchStart monoclock.Instant, ex // RunHLCLeaseRenewal runs a background loop that periodically proposes a new // physical ceiling to the Raft cluster while this node is the leader. // -// The ceiling is set to now + hlcPhysicalWindowMs (3 s) and is renewed every -// hlcRenewalInterval (1 s), mirroring TiDB's TSO window strategy. Because the -// window is always at least 2 s ahead of any real timestamp, a new leader will -// never issue timestamps that overlap with the previous leader's window. +// The ceiling is set to now + hlcPhysicalWindowMs and is renewed every +// hlcRenewalInterval, mirroring TiDB's TSO window strategy. Because the window +// stays ahead of real timestamps, a new leader will never issue timestamps that +// overlap with the previous leader's window. // // RunHLCLeaseRenewal blocks until ctx is cancelled; call it in a goroutine. func (c *Coordinate) RunHLCLeaseRenewal(ctx context.Context) { @@ -877,7 +882,10 @@ func (c *Coordinate) RunHLCLeaseRenewal(ctx context.Context) { } if c.IsLeaderAcceptingWrites() { ceilingMs := time.Now().UnixMilli() + hlcPhysicalWindowMs - if err := c.ProposeHLCLease(ctx, ceilingMs); err != nil { + pctx, cancel := context.WithTimeout(ctx, hlcRenewalTimeout) + err := c.ProposeHLCLease(pctx, ceilingMs) + cancel() + if err != nil { c.log.WarnContext(ctx, "hlc lease renewal failed", slog.Int64("ceiling_ms", ceilingMs), slog.Any("err", err), @@ -1129,7 +1137,7 @@ func (c *Coordinate) dispatchTxn(ctx context.Context, reqs []*Elem[OP], readKeys // carries the option-2 one-phase dedup probe key for a retry that reuses // a failed attempt's write set. r, err := c.transactionManager.Commit(ctx, []*pb.Request{ - onePhaseTxnRequestWithPrevCommit(startTS, commitTS, prevCommitTS, primary, reqs, readKeys, observedRouteVersion), + onePhaseTxnRequestWithPrevCommit(startTS, commitTS, prevCommitTS, primary, reqs, readKeys, observedRouteVersion, nil), }) if err != nil { return nil, errors.WithStack(err) @@ -1188,12 +1196,13 @@ func (c *Coordinate) dispatchRaw(ctx context.Context, req []*Elem[OP]) (*Coordin // The returned Request is structurally identical to the pre-stamping // shape the leader's stampRawTimestamps already handles for Ts == 0 // (see adapter/internal.go). -func (c *Coordinate) toRawRequest(req *Elem[OP]) *pb.Request { +func (c *Coordinate) toRawRequest(req *Elem[OP], observedRouteVersion uint64) *pb.Request { switch req.Op { case Put: return &pb.Request{ - IsTxn: false, - Phase: pb.Phase_NONE, + IsTxn: false, + Phase: pb.Phase_NONE, + ObservedRouteVersion: observedRouteVersion, Mutations: []*pb.Mutation{ { Op: pb.Op_PUT, @@ -1205,8 +1214,9 @@ func (c *Coordinate) toRawRequest(req *Elem[OP]) *pb.Request { case Del: return &pb.Request{ - IsTxn: false, - Phase: pb.Phase_NONE, + IsTxn: false, + Phase: pb.Phase_NONE, + ObservedRouteVersion: observedRouteVersion, Mutations: []*pb.Mutation{ { Op: pb.Op_DEL, @@ -1217,8 +1227,9 @@ func (c *Coordinate) toRawRequest(req *Elem[OP]) *pb.Request { case DelPrefix: return &pb.Request{ - IsTxn: false, - Phase: pb.Phase_NONE, + IsTxn: false, + Phase: pb.Phase_NONE, + ObservedRouteVersion: observedRouteVersion, Mutations: []*pb.Mutation{ { Op: pb.Op_DEL_PREFIX, @@ -1282,7 +1293,7 @@ func (c *Coordinate) buildRedirectRequests(reqs *OperationGroup[OP]) ([]*pb.Requ if !reqs.IsTxn { requests := make([]*pb.Request, 0, len(reqs.Elems)) for _, req := range reqs.Elems { - requests = append(requests, c.toRawRequest(req)) + requests = append(requests, c.toRawRequest(req, reqs.ObservedRouteVersion)) } return requests, nil } @@ -1306,7 +1317,7 @@ func (c *Coordinate) buildRedirectRequests(reqs *OperationGroup[OP]) ([]*pb.Requ commitTS = 0 } return []*pb.Request{ - onePhaseTxnRequestWithPrevCommit(reqs.StartTS, commitTS, reqs.PrevCommitTS, primary, reqs.Elems, reqs.ReadKeys, reqs.ObservedRouteVersion), + onePhaseTxnRequestWithPrevCommit(reqs.StartTS, commitTS, reqs.PrevCommitTS, primary, reqs.Elems, reqs.ReadKeys, reqs.ObservedRouteVersion, nil), }, nil } @@ -1365,7 +1376,7 @@ func elemToMutation(req *Elem[OP]) *pb.Mutation { // route catalog snapshot at txn-begin (M1 plumbing, see // docs/design/2026_05_29_implemented_composed1_cross_group_commit_guard.md). // Zero is the legacy "unpinned" sentinel. -func onePhaseTxnRequestWithPrevCommit(startTS, commitTS, prevCommitTS uint64, primaryKey []byte, reqs []*Elem[OP], readKeys [][]byte, observedRouteVersion uint64) *pb.Request { +func onePhaseTxnRequestWithPrevCommit(startTS, commitTS, prevCommitTS uint64, primaryKey []byte, reqs []*Elem[OP], readKeys [][]byte, observedRouteVersion uint64, writeFenceBypassKeys [][]byte) *pb.Request { muts := make([]*pb.Mutation, 0, len(reqs)+1) muts = append(muts, &pb.Mutation{ Op: pb.Op_PUT, @@ -1388,6 +1399,7 @@ func onePhaseTxnRequestWithPrevCommit(startTS, commitTS, prevCommitTS uint64, pr Mutations: muts, ReadKeys: readKeys, ObservedRouteVersion: observedRouteVersion, + WriteFenceBypassKeys: writeFenceBypassKeys, } } diff --git a/kv/coordinator_dispatch_test.go b/kv/coordinator_dispatch_test.go index a559c6405..ae1ffdf30 100644 --- a/kv/coordinator_dispatch_test.go +++ b/kv/coordinator_dispatch_test.go @@ -213,7 +213,7 @@ func TestToRawRequestLeavesTsForLeaderStamping(t *testing.T) { } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { - r := c.toRawRequest(tc.req) + r := c.toRawRequest(tc.req, 0) require.NotNil(t, r) require.Equal(t, uint64(0), r.Ts, "forwarded raw requests must arrive with Ts==0 so the leader's stampRawTimestamps assigns the canonical ts (HLC leader-only invariant + HLC-4 (iii) fence)") @@ -221,6 +221,21 @@ func TestToRawRequestLeavesTsForLeaderStamping(t *testing.T) { } } +func TestBuildRedirectRequests_PreservesRawObservedRouteVersion(t *testing.T) { + t.Parallel() + + c := &Coordinate{} + got, err := c.buildRedirectRequests(&OperationGroup[OP]{ + ObservedRouteVersion: 17, + Elems: []*Elem[OP]{ + {Op: Put, Key: []byte("k"), Value: []byte("v")}, + }, + }) + require.NoError(t, err) + require.Len(t, got, 1) + require.Equal(t, uint64(17), got[0].GetObservedRouteVersion()) +} + // TestBuildRedirectRequestsSurvivesStaleFollowerCeiling exercises the // follower's redirect path end-to-end with an expired ceiling: it // confirms the follower hands off a Ts==0 request to the leader diff --git a/kv/fsm.go b/kv/fsm.go index 6ed81f076..09846821b 100644 --- a/kv/fsm.go +++ b/kv/fsm.go @@ -11,6 +11,7 @@ import ( "github.com/bootjp/elastickv/internal/encryption/fsmwire" "github.com/bootjp/elastickv/internal/raftengine" + "github.com/bootjp/elastickv/internal/s3keys" pb "github.com/bootjp/elastickv/proto" "github.com/bootjp/elastickv/store" "github.com/cockroachdb/errors" @@ -127,6 +128,12 @@ type RouteSnapshot interface { // OwnerOf returns the Raft group ID that owned key at this // snapshot's version. (0, false) when no route covered key. OwnerOf(key []byte) (uint64, bool) + // WriteFencedForKey reports whether key is currently inside a + // WriteFenced route in this snapshot. + WriteFencedForKey(key []byte) bool + // WriteFencedIntersects reports whether [start, end) intersects + // any WriteFenced route in this snapshot. + WriteFencedIntersects(start, end []byte) bool } // SetApplyIndex implements raftengine.ApplyIndexAware. The engine @@ -269,6 +276,11 @@ var _ raftengine.StateMachine = (*kvFSM)(nil) var ErrUnknownRequestType = errors.New("unknown request type") +// ErrRouteWriteFenced is returned when a mutation targets a route that is in +// WriteFenced state during split migration. Callers should retry after routing +// catches up to the promoted owner. +var ErrRouteWriteFenced = errors.New("route is write-fenced; retry after route migration") + // ErrComposed1Violation is returned by verifyComposed1 when the // transaction's commit cannot proceed on this Raft group because the // txn's read-set or write-set keys are not owned by this group at @@ -478,6 +490,9 @@ func (f *kvFSM) handleRequest(ctx context.Context, r *pb.Request, commitTS uint6 } func (f *kvFSM) handleRawRequest(ctx context.Context, r *pb.Request, commitTS uint64) error { + if err := f.verifyWriteFence(r); err != nil { + return err + } // DEL_PREFIX mutations are handled by the store's DeletePrefixAt which // scans and writes tombstones locally. A DEL_PREFIX request must be the // sole mutation in a request (enforced by the coordinator's toRawRequest). @@ -532,6 +547,94 @@ func (f *kvFSM) handleDelPrefix(ctx context.Context, prefix []byte, commitTS uin return nil } +func routePrefixRange(prefix []byte) ([]byte, []byte) { + if len(prefix) == 0 { + return []byte(""), nil + } + if start, ok := s3keys.BucketGenerationRoutePrefixForCleanupPrefix(prefix); ok { + return start, prefixScanEnd(start) + } + if start, ok := dynamoExactCleanupRouteKey(prefix); ok { + return start, routePointRangeEnd(start) + } + if routeKeyspaceWideRawPrefix(prefix) { + return []byte(""), nil + } + start := routeKey(prefix) + return start, prefixScanEnd(start) +} + +func dynamoExactCleanupRouteKey(prefix []byte) ([]byte, bool) { + switch { + case bytes.HasPrefix(prefix, dynamoTableMetaPrefixBytes), + bytes.HasPrefix(prefix, dynamoTableGenerationPrefixBytes), + bytes.HasPrefix(prefix, dynamoItemPrefixBytes), + bytes.HasPrefix(prefix, dynamoGSIPrefixBytes): + default: + return nil, false + } + start := routeKey(prefix) + if len(start) == 0 || bytes.Equal(start, prefix) || !bytes.HasPrefix(start, dynamoRoutePrefixBytes) { + return nil, false + } + return start, true +} + +func routePointRangeEnd(start []byte) []byte { + end := make([]byte, 0, len(start)+1) + end = append(end, start...) + end = append(end, 0) + return end +} + +func routeKeyspaceWideRawPrefix(prefix []byte) bool { + if !rawPrefixMayContainRouteMappedKeys(prefix) { + return false + } + return bytes.Equal(routeKey(prefix), prefix) +} + +func rawPrefixMayContainRouteMappedKeys(prefix []byte) bool { + for _, mappedPrefix := range routeMappedRawPrefixes { + if bytes.HasPrefix(prefix, mappedPrefix) || bytes.HasPrefix(mappedPrefix, prefix) { + return true + } + } + return false +} + +var routeMappedRawPrefixes = append([][]byte{ + []byte(redisInternalRoutePrefix), + []byte(DynamoTableMetaPrefix), + []byte(DynamoTableGenerationPrefix), + []byte(DynamoItemPrefix), + []byte(DynamoGSIPrefix), + []byte(store.ListMetaPrefix), + []byte(store.ListItemPrefix), + []byte(store.ListMetaDeltaPrefix), + []byte(store.ListClaimPrefix), + []byte(store.HashMetaPrefix), + []byte(store.HashFieldPrefix), + []byte(store.HashMetaDeltaPrefix), + []byte(store.SetMetaPrefix), + []byte(store.SetMemberPrefix), + []byte(store.SetMetaDeltaPrefix), + []byte(store.ZSetMetaPrefix), + []byte(store.ZSetMemberPrefix), + []byte(store.ZSetScorePrefix), + []byte(store.ZSetMetaDeltaPrefix), + []byte(store.StreamMetaPrefix), + []byte(store.StreamEntryPrefix), + []byte(s3keys.BucketMetaPrefix), + []byte(s3keys.BucketGenerationPrefix), + []byte(s3keys.ObjectManifestPrefix), + []byte(s3keys.UploadMetaPrefix), + []byte(s3keys.UploadPartPrefix), + []byte(s3keys.BlobPrefix), + []byte(s3keys.GCUploadPrefix), + []byte(s3keys.RoutePrefix), +}, sqsConcreteInternalPrefixBytes...) + var ErrNotImplemented = errors.New("not implemented") func (f *kvFSM) Snapshot() (raftengine.Snapshot, error) { @@ -595,37 +698,30 @@ func (f *kvFSM) RestoredCutover() uint64 { // ParseSnapshotHeader implements raftengine.SnapshotHeaderApplier // phase 1 — the cold-start skip path's parse-without-side-effect -// step. The engine has wrapped `r` in a crc32 TeeReader sized at -// the body payload (file size minus 4-byte footer), so every byte -// pulled from `r` flows through the engine's hash. We read the -// v1/v2 header via ReadSnapshotHeader, then drain the rest of the -// body so the wrapping hash covers every payload byte — matching -// restoreAndComputeCRC's behaviour in openAndRestoreFSMSnapshot. +// step. The skip path only needs the header state because the FSM body +// is already present locally, so this reads the v1/v2 header and leaves +// the remainder untouched. Full-body CRC verification still happens on +// the restore path where the body bytes are consumed. // // IMPORTANT: this method MUST NOT touch f.hlc or f.restoredCutover. // The engine calls ApplySnapshotHeader separately, only after the -// wrapping CRC verification passes. Mutating FSM state here would -// defeat the "no side-effect on CRC failure" contract that the -// PR #910 design §5 round-7 split is designed to preserve. +// snapshot envelope checks pass. Mutating FSM state here would defeat +// the "no side-effect on parse failure" contract that the PR #910 +// design §5 round-7 split is designed to preserve. func (f *kvFSM) ParseSnapshotHeader(r io.Reader) (uint64, uint64, error) { - br := bufio.NewReaderSize(r, 1<<20) //nolint:mnd // 1 MiB, local to kv - ceiling, cutover, err := ReadSnapshotHeader(br) + const headerReadBufferSize = 4 << 10 + + ceiling, cutover, err := ReadSnapshotHeader(bufio.NewReaderSize(r, headerReadBufferSize)) if err != nil { return 0, 0, errors.WithStack(err) } - // Drain the remainder so the engine's TeeReader-wrapped CRC - // covers every byte of the body (LimitReader exhaustion - // signals "full payload consumed" to the caller). - if _, err := io.Copy(io.Discard, br); err != nil { - return 0, 0, errors.WithStack(err) - } return ceiling, cutover, nil } // ApplySnapshotHeader implements raftengine.SnapshotHeaderApplier // phase 2 — pure assignment of the verified header state. Called -// only after ParseSnapshotHeader returned successfully AND the -// engine's wrapping crc32 hash matched the file footer. Mirrors +// only after ParseSnapshotHeader returned successfully and the +// snapshot file's footer matched the raft token. Mirrors // the two side-effects Restore would have applied for the header // portion (HLC physical ceiling + restoredCutover). See PR #910 // design §5 round-7. @@ -657,6 +753,9 @@ func (f *kvFSM) IsVolatileOnlyPayload(payload []byte) bool { } func (f *kvFSM) handleTxnRequest(ctx context.Context, r *pb.Request, commitTS uint64) error { + if err := f.verifyWriteFence(r); err != nil { + return err + } if err := f.verifyComposed1(r); err != nil { return err } @@ -733,13 +832,14 @@ func (f *kvFSM) verifyComposed1(r *pb.Request) error { if observedVer == 0 { return nil } + bypassKeys := writeFenceBypassKeySet(r.GetWriteFenceBypassKeys()) // (a) Observed-version check. observedSnap, ok := f.routes.SnapshotAt(observedVer) if !ok { return errors.WithStack(ErrComposed1VersionGCd) } - if err := f.verifyOwnerFromSnapshot(r.GetMutations(), observedSnap, observedVer, "observed"); err != nil { + if err := f.verifyOwnerFromSnapshot(r.GetMutations(), bypassKeys, observedSnap, observedVer, "observed"); err != nil { return err } @@ -751,7 +851,104 @@ func (f *kvFSM) verifyComposed1(r *pb.Request) error { // short-circuit posture of an unwired FSM). return nil } - return f.verifyOwnerFromSnapshot(r.GetMutations(), currentSnap, currentSnap.Version(), "current") + return f.verifyOwnerFromSnapshot(r.GetMutations(), bypassKeys, currentSnap, currentSnap.Version(), "current") +} + +func (f *kvFSM) verifyWriteFence(r *pb.Request) error { + if requestBypassesWriteFence(r) { + return nil + } + observedVer := r.GetObservedRouteVersion() + if !f.writeFenceHistoryReady() { + return nil + } + currentSnap, ok := f.routes.Current() + if !ok { + return nil + } + + if observedVer != 0 { + observedSnap, ok := f.routes.SnapshotAt(observedVer) + if !ok { + return errors.WithStack(ErrComposed1VersionGCd) + } + if err := verifyWriteFenceFromSnapshot(r.GetMutations(), r.GetWriteFenceBypassKeys(), observedSnap, observedVer, "observed"); err != nil { + return err + } + if currentSnap.Version() == observedSnap.Version() { + return nil + } + } + + return verifyWriteFenceFromSnapshot(r.GetMutations(), r.GetWriteFenceBypassKeys(), currentSnap, currentSnap.Version(), "current") +} + +func requestBypassesWriteFence(r *pb.Request) bool { + if !r.GetIsTxn() { + return false + } + switch r.GetPhase() { + case pb.Phase_COMMIT, pb.Phase_ABORT: + return true + case pb.Phase_NONE, pb.Phase_PREPARE: + return false + } + return false +} + +func (f *kvFSM) writeFenceHistoryReady() bool { + return f.routes != nil && f.shardGroupID != 0 +} + +func verifyWriteFenceFromSnapshot(mutations []*pb.Mutation, writeFenceBypassKeys [][]byte, snap RouteSnapshot, snapVer uint64, phase string) error { + bypassKeys := writeFenceBypassKeySet(writeFenceBypassKeys) + for _, mut := range mutations { + if mut == nil { + continue + } + if isTxnInternalKey(mut.Key) { + continue + } + if mut.GetOp() == pb.Op_DEL_PREFIX { + start, end := routePrefixRange(mut.Key) + if snap.WriteFencedIntersects(start, end) { + return errors.Wrapf(ErrRouteWriteFenced, + "%s-version v=%d: prefix %q route range [%q,%q)", + phase, snapVer, mut.Key, start, end) + } + continue + } + if _, ok := bypassKeys[string(mut.Key)]; ok { + continue + } + rKey := routeKey(mut.Key) + if snap.WriteFencedForKey(rKey) { + return errors.Wrapf(ErrRouteWriteFenced, + "%s-version v=%d: key %q routeKey %q", + phase, snapVer, mut.Key, rKey) + } + start, end, ok := s3BucketAuxiliaryRouteRange(mut.Key) + if ok && snap.WriteFencedIntersects(start, end) { + return errors.Wrapf(ErrRouteWriteFenced, + "%s-version v=%d: key %q route range [%q,%q)", + phase, snapVer, mut.Key, start, end) + } + } + return nil +} + +func writeFenceBypassKeySet(keys [][]byte) map[string]struct{} { + if len(keys) == 0 { + return nil + } + out := make(map[string]struct{}, len(keys)) + for _, key := range keys { + if len(key) == 0 { + continue + } + out[string(key)] = struct{}{} + } + return out } // verifyOwnerFromSnapshot is the shared per-mutation owner-check @@ -760,7 +957,7 @@ func (f *kvFSM) verifyComposed1(r *pb.Request) error { // "current") that ends up in the wrapped error. isTxnInternalKey // mutations (the TxnMeta marker prefix) are skipped — they are // always on every shard and have no Composed-1 ownership. -func (f *kvFSM) verifyOwnerFromSnapshot(mutations []*pb.Mutation, snap RouteSnapshot, snapVer uint64, phase string) error { +func (f *kvFSM) verifyOwnerFromSnapshot(mutations []*pb.Mutation, bypassKeys map[string]struct{}, snap RouteSnapshot, snapVer uint64, phase string) error { for _, mut := range mutations { if mut == nil || len(mut.Key) == 0 { continue @@ -768,6 +965,9 @@ func (f *kvFSM) verifyOwnerFromSnapshot(mutations []*pb.Mutation, snap RouteSnap if isTxnInternalKey(mut.Key) { continue } + if _, ok := bypassKeys[string(mut.Key)]; ok { + continue + } // routeKey-normalize before OwnerOf so the gate routes the // same way as ShardRouter.ResolveGroup — raw adapter keys // and route catalog ranges live in different lex bands diff --git a/kv/fsm_migration_fence_test.go b/kv/fsm_migration_fence_test.go new file mode 100644 index 000000000..807a5f551 --- /dev/null +++ b/kv/fsm_migration_fence_test.go @@ -0,0 +1,321 @@ +package kv + +import ( + "context" + "testing" + + "github.com/bootjp/elastickv/distribution" + "github.com/bootjp/elastickv/internal/s3keys" + pb "github.com/bootjp/elastickv/proto" + "github.com/stretchr/testify/require" +) + +func newWriteFencedFSM(t *testing.T) *kvFSM { + t.Helper() + + engine := distribution.NewEngine() + applyComposed1Snapshot(t, engine, 1, []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte(""), End: []byte("m"), GroupID: 1, State: distribution.RouteStateActive}, + {RouteID: 2, Start: []byte("m"), End: nil, GroupID: 1, State: distribution.RouteStateWriteFenced}, + }) + return newComposed1FSM(t, engine, 1) +} + +func newFirstRouteWriteFencedFSM(t *testing.T) *kvFSM { + t.Helper() + + engine := distribution.NewEngine() + applyComposed1Snapshot(t, engine, 1, []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte(""), End: []byte("m"), GroupID: 1, State: distribution.RouteStateWriteFenced}, + {RouteID: 2, Start: []byte("m"), End: nil, GroupID: 1, State: distribution.RouteStateActive}, + }) + return newComposed1FSM(t, engine, 1) +} + +func s3BucketAuxiliaryFenceRoutes(bucket string, rawGroupID, fencedGroupID uint64) []distribution.RouteDescriptor { + start := s3keys.RoutePrefixForBucketAnyGeneration(bucket) + end := prefixScanEnd(start) + return []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte(""), End: start, GroupID: rawGroupID, State: distribution.RouteStateActive}, + {RouteID: 2, Start: start, End: end, GroupID: fencedGroupID, State: distribution.RouteStateWriteFenced}, + {RouteID: 3, Start: end, End: nil, GroupID: rawGroupID, State: distribution.RouteStateActive}, + } +} + +func newS3BucketAuxiliaryWriteFencedFSM(t *testing.T, bucket string) *kvFSM { + t.Helper() + + engine := distribution.NewEngine() + applyComposed1Snapshot(t, engine, 1, s3BucketAuxiliaryFenceRoutes(bucket, 1, 1)) + return newComposed1FSM(t, engine, 1) +} + +func TestFSMRejectsCurrentWriteFencedRawPointWrite(t *testing.T) { + t.Parallel() + + fsm := newWriteFencedFSM(t) + err := fsm.handleRawRequest(context.Background(), &pb.Request{ + Mutations: []*pb.Mutation{{Op: pb.Op_PUT, Key: []byte("z"), Value: []byte("v")}}, + }, 10) + require.ErrorIs(t, err, ErrRouteWriteFenced) +} + +func TestFSMRejectsCurrentWriteFencedEmptyRawPointWrite(t *testing.T) { + t.Parallel() + + fsm := newFirstRouteWriteFencedFSM(t) + err := fsm.handleRawRequest(context.Background(), &pb.Request{ + Mutations: []*pb.Mutation{{Op: pb.Op_PUT, Key: []byte(""), Value: []byte("v")}}, + }, 10) + require.ErrorIs(t, err, ErrRouteWriteFenced) +} + +func TestFSMRejectsObservedWriteFencedRawPointWrite(t *testing.T) { + t.Parallel() + + fsm := newWriteFencedFSM(t) + err := fsm.handleRawRequest(context.Background(), &pb.Request{ + ObservedRouteVersion: 1, + Mutations: []*pb.Mutation{{Op: pb.Op_PUT, Key: []byte("z"), Value: []byte("v")}}, + }, 10) + require.ErrorIs(t, err, ErrRouteWriteFenced) +} + +func TestFSMWriteFenceBypassAllowsMarkedRawPointWrite(t *testing.T) { + t.Parallel() + + fsm := newFirstRouteWriteFencedFSM(t) + key := []byte("!sqs|msg|data|p|partitioned-key") + err := fsm.handleRawRequest(context.Background(), &pb.Request{ + WriteFenceBypassKeys: [][]byte{key}, + Mutations: []*pb.Mutation{{Op: pb.Op_PUT, Key: key, Value: []byte("v")}}, + }, 10) + require.NoError(t, err) + + got, err := fsm.store.GetAt(context.Background(), key, 10) + require.NoError(t, err) + require.Equal(t, []byte("v"), got) +} + +func TestFSMWriteFenceBypassAllowsPinnedTxnOnNonOwningGroup(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + applyComposed1Snapshot(t, engine, 1, []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte(""), End: []byte("m"), GroupID: 1, State: distribution.RouteStateActive}, + {RouteID: 2, Start: []byte("m"), End: nil, GroupID: 2, State: distribution.RouteStateWriteFenced}, + }) + fsm := newComposed1FSM(t, engine, 1) + key := []byte("z") + err := fsm.handleTxnRequest(context.Background(), &pb.Request{ + IsTxn: true, + Phase: pb.Phase_PREPARE, + Ts: 10, + ObservedRouteVersion: 1, + WriteFenceBypassKeys: [][]byte{key}, + Mutations: []*pb.Mutation{ + {Op: pb.Op_PUT, Key: []byte(txnMetaPrefix), Value: EncodeTxnMeta(TxnMeta{PrimaryKey: key, LockTTLms: defaultTxnLockTTLms})}, + {Op: pb.Op_DEL, Key: key}, + }, + }, 10) + require.NoError(t, err) +} + +func TestFSMWriteFenceBypassDoesNotAllowDelPrefix(t *testing.T) { + t.Parallel() + + fsm := newFirstRouteWriteFencedFSM(t) + prefix := []byte("!sqs|msg|data|p|") + err := fsm.handleRawRequest(context.Background(), &pb.Request{ + WriteFenceBypassKeys: [][]byte{prefix}, + Mutations: []*pb.Mutation{{Op: pb.Op_DEL_PREFIX, Key: prefix}}, + }, 10) + require.ErrorIs(t, err, ErrRouteWriteFenced) +} + +func TestFSMRejectsCurrentWriteFenceAfterObservedActiveRawPointWrite(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + applyComposed1Snapshot(t, engine, 1, []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte(""), End: nil, GroupID: 1, State: distribution.RouteStateActive}, + }) + fsm := newComposed1FSM(t, engine, 1) + applyComposed1Snapshot(t, engine, 2, []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte(""), End: []byte("m"), GroupID: 1, State: distribution.RouteStateActive}, + {RouteID: 2, Start: []byte("m"), End: nil, GroupID: 1, State: distribution.RouteStateWriteFenced}, + }) + + err := fsm.handleRawRequest(context.Background(), &pb.Request{ + ObservedRouteVersion: 1, + Mutations: []*pb.Mutation{{Op: pb.Op_PUT, Key: []byte("z"), Value: []byte("v")}}, + }, 10) + require.ErrorIs(t, err, ErrRouteWriteFenced) +} + +func TestFSMRejectsCurrentWriteFencedUnpinnedPrepare(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + applyComposed1Snapshot(t, engine, 1, []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte(""), End: nil, GroupID: 1, State: distribution.RouteStateActive}, + }) + fsm := newComposed1FSM(t, engine, 1) + applyComposed1Snapshot(t, engine, 2, []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte(""), End: []byte("m"), GroupID: 1, State: distribution.RouteStateActive}, + {RouteID: 2, Start: []byte("m"), End: nil, GroupID: 1, State: distribution.RouteStateWriteFenced}, + }) + + err := fsm.handleTxnRequest(context.Background(), &pb.Request{ + IsTxn: true, + Phase: pb.Phase_PREPARE, + Ts: 10, + Mutations: []*pb.Mutation{ + {Op: pb.Op_PUT, Key: []byte(txnMetaPrefix), Value: EncodeTxnMeta(TxnMeta{PrimaryKey: []byte("z"), LockTTLms: defaultTxnLockTTLms})}, + {Op: pb.Op_PUT, Key: []byte("z"), Value: []byte("v")}, + }, + }, 10) + require.ErrorIs(t, err, ErrRouteWriteFenced) +} + +func TestFSMRejectsCurrentWriteFencedS3BucketAuxiliaryPointWrite(t *testing.T) { + t.Parallel() + + ctx := context.Background() + const bucket = "bucket-a" + fsm := newS3BucketAuxiliaryWriteFencedFSM(t, bucket) + + for _, key := range [][]byte{ + s3keys.BucketMetaKey(bucket), + s3keys.BucketGenerationKey(bucket), + } { + err := fsm.handleRawRequest(ctx, &pb.Request{ + Mutations: []*pb.Mutation{{Op: pb.Op_PUT, Key: key, Value: []byte("v")}}, + }, 10) + require.ErrorIs(t, err, ErrRouteWriteFenced) + } +} + +func TestFSMRejectsObservedWriteFencedS3BucketAuxiliaryPointWrite(t *testing.T) { + t.Parallel() + + const bucket = "bucket-a" + fsm := newS3BucketAuxiliaryWriteFencedFSM(t, bucket) + + err := fsm.handleRawRequest(context.Background(), &pb.Request{ + ObservedRouteVersion: 1, + Mutations: []*pb.Mutation{ + {Op: pb.Op_PUT, Key: s3keys.BucketGenerationKey(bucket), Value: []byte("v")}, + }, + }, 10) + require.ErrorIs(t, err, ErrRouteWriteFenced) +} + +func TestFSMRejectsCurrentWriteFencedDelPrefix(t *testing.T) { + t.Parallel() + + fsm := newWriteFencedFSM(t) + require.NoError(t, fsm.store.PutAt(context.Background(), []byte("z"), []byte("v"), 1, 0)) + + err := fsm.handleRawRequest(context.Background(), &pb.Request{ + Mutations: []*pb.Mutation{{Op: pb.Op_DEL_PREFIX, Key: []byte("z")}}, + }, 10) + require.ErrorIs(t, err, ErrRouteWriteFenced) +} + +func TestFSMRejectsObservedWriteFencedDelPrefix(t *testing.T) { + t.Parallel() + + fsm := newWriteFencedFSM(t) + require.NoError(t, fsm.store.PutAt(context.Background(), []byte("z"), []byte("v"), 1, 0)) + + err := fsm.handleRawRequest(context.Background(), &pb.Request{ + ObservedRouteVersion: 1, + Mutations: []*pb.Mutation{{Op: pb.Op_DEL_PREFIX, Key: []byte("z")}}, + }, 10) + require.ErrorIs(t, err, ErrRouteWriteFenced) +} + +func TestFSMRejectsCurrentWriteFencedFullRangeDelPrefix(t *testing.T) { + t.Parallel() + + fsm := newWriteFencedFSM(t) + require.NoError(t, fsm.store.PutAt(context.Background(), []byte("z"), []byte("v"), 1, 0)) + + err := fsm.handleRawRequest(context.Background(), &pb.Request{ + Mutations: []*pb.Mutation{{Op: pb.Op_DEL_PREFIX, Key: nil}}, + }, 10) + require.ErrorIs(t, err, ErrRouteWriteFenced) +} + +func TestFSMRejectsCurrentWriteFencedBroadInternalDelPrefix(t *testing.T) { + t.Parallel() + + fsm := newWriteFencedFSM(t) + key := []byte("!redis|string|z") + require.NoError(t, fsm.store.PutAt(context.Background(), key, []byte("v"), 1, 0)) + + err := fsm.handleRawRequest(context.Background(), &pb.Request{ + Mutations: []*pb.Mutation{{Op: pb.Op_DEL_PREFIX, Key: []byte("!redis|")}}, + }, 10) + require.ErrorIs(t, err, ErrRouteWriteFenced) +} + +func TestFSMRejectsCurrentWriteFencedPrepareButAllowsAbort(t *testing.T) { + t.Parallel() + + ctx := context.Background() + fsm := newWriteFencedFSM(t) + prepare := &pb.Request{ + IsTxn: true, + Phase: pb.Phase_PREPARE, + Ts: 10, + Mutations: []*pb.Mutation{ + {Op: pb.Op_PUT, Key: []byte(txnMetaPrefix), Value: EncodeTxnMeta(TxnMeta{PrimaryKey: []byte("z"), LockTTLms: defaultTxnLockTTLms})}, + {Op: pb.Op_PUT, Key: []byte("z"), Value: []byte("v")}, + }, + } + require.ErrorIs(t, fsm.handleTxnRequest(ctx, prepare, 10), ErrRouteWriteFenced) + + abort := &pb.Request{ + IsTxn: true, + Phase: pb.Phase_ABORT, + Ts: 11, + Mutations: []*pb.Mutation{ + {Op: pb.Op_PUT, Key: []byte(txnMetaPrefix), Value: EncodeTxnMeta(TxnMeta{PrimaryKey: []byte("z"), CommitTS: 11})}, + {Op: pb.Op_PUT, Key: []byte("z"), Value: []byte("v")}, + }, + } + err := fsm.handleTxnRequest(ctx, abort, 11) + require.NotErrorIs(t, err, ErrRouteWriteFenced, "ABORT must keep the narrow cleanup lane open") +} + +func TestFSMRejectsObservedWriteFencedPrepareButAllowsAbort(t *testing.T) { + t.Parallel() + + ctx := context.Background() + fsm := newWriteFencedFSM(t) + prepare := &pb.Request{ + IsTxn: true, + Phase: pb.Phase_PREPARE, + Ts: 10, + ObservedRouteVersion: 1, + Mutations: []*pb.Mutation{ + {Op: pb.Op_PUT, Key: []byte(txnMetaPrefix), Value: EncodeTxnMeta(TxnMeta{PrimaryKey: []byte("z"), LockTTLms: defaultTxnLockTTLms})}, + {Op: pb.Op_PUT, Key: []byte("z"), Value: []byte("v")}, + }, + } + require.ErrorIs(t, fsm.handleTxnRequest(ctx, prepare, 10), ErrRouteWriteFenced) + + abort := &pb.Request{ + IsTxn: true, + Phase: pb.Phase_ABORT, + Ts: 11, + ObservedRouteVersion: 1, + Mutations: []*pb.Mutation{ + {Op: pb.Op_PUT, Key: []byte(txnMetaPrefix), Value: EncodeTxnMeta(TxnMeta{PrimaryKey: []byte("z"), CommitTS: 11})}, + {Op: pb.Op_PUT, Key: []byte("z"), Value: []byte("v")}, + }, + } + require.NotErrorIs(t, fsm.handleTxnRequest(ctx, abort, 11), ErrRouteWriteFenced) +} diff --git a/kv/leader_routed_store.go b/kv/leader_routed_store.go index a0b06c6be..8b8d570b7 100644 --- a/kv/leader_routed_store.go +++ b/kv/leader_routed_store.go @@ -109,9 +109,6 @@ func (s *LeaderRoutedStore) GetAtWithReadFence(ctx context.Context, key []byte, } ok, fenceTS := s.leaderFenceTS(ctx, key) if ok { - if readRouteVersion != 0 { - return nil, errors.WithStack(store.ErrNotSupported) - } val, err := s.local.GetAt(ctx, key, max(ts, fenceTS)) return val, errors.WithStack(err) } @@ -148,9 +145,6 @@ func (s *LeaderRoutedStore) LatestCommitTSWithReadFence(ctx context.Context, key return 0, false, nil } if s.leaderOKForKey(ctx, key) { - if readRouteVersion != 0 { - return 0, false, errors.WithStack(store.ErrNotSupported) - } ts, exists, err := s.local.LatestCommitTS(ctx, key) return ts, exists, errors.WithStack(err) } @@ -274,9 +268,6 @@ func (s *LeaderRoutedStore) ScanAtWithReadFence(ctx context.Context, start []byt if !ok { return s.proxyRawScanAtWithReadFence(ctx, start, end, limit, ts, reverse, readRouteVersion, routeStart, routeEnd) } - if readRouteVersion != 0 { - return nil, errors.WithStack(store.ErrNotSupported) - } readTS := max(ts, fenceTS) if routeScanBoundsPresent(routeStart, routeEnd) { return s.scanLocalRouteFilteredAt(ctx, start, end, limit, readTS, reverse, routeStart, routeEnd) @@ -475,6 +466,20 @@ func (s *LeaderRoutedStore) ReverseScanAtPhysicalLimit(ctx context.Context, star return s.scanAtPhysicalLimit(ctx, start, end, visibleLimit, physicalLimit, ts, true) } +func (s *LeaderRoutedStore) AllowExactScanFallbackAfterPhysicalLimit(ctx context.Context, start []byte, _ []byte, visibleLimit, physicalLimit int, _ uint64, _ bool) bool { + if s == nil || s.local == nil { + return false + } + if visibleLimit <= 0 || physicalLimit <= 0 { + return false + } + if ok, _ := s.leaderFenceTS(ctx, start); !ok { + return false + } + _, ok := s.local.(physicalLimitedStore) + return ok +} + func (s *LeaderRoutedStore) scanAtPhysicalLimit(ctx context.Context, start []byte, end []byte, visibleLimit, physicalLimit int, ts uint64, reverse bool) ([]*store.KVPair, bool, error) { if s == nil || s.local == nil { return []*store.KVPair{}, false, nil @@ -695,6 +700,35 @@ func (s *LeaderRoutedStore) Compact(ctx context.Context, minTS uint64) error { return errors.WithStack(s.local.Compact(ctx, minTS)) } +func (s *LeaderRoutedStore) ExportVersions(ctx context.Context, opts store.ExportVersionsOptions) (store.ExportVersionsResult, error) { + if s == nil || s.local == nil { + return store.ExportVersionsResult{}, errors.WithStack(store.ErrNotSupported) + } + if s.coordinator != nil { + if _, err := s.coordinator.LinearizableRead(ctx); err != nil { + return store.ExportVersionsResult{}, errors.WithStack(err) + } + } + result, err := s.local.ExportVersions(ctx, opts) + return result, errors.WithStack(err) +} + +func (s *LeaderRoutedStore) ImportVersions(ctx context.Context, opts store.ImportVersionsOptions) (store.ImportVersionsResult, error) { + if s == nil || s.local == nil { + return store.ImportVersionsResult{}, errors.WithStack(store.ErrNotSupported) + } + result, err := s.local.ImportVersions(ctx, opts) + return result, errors.WithStack(err) +} + +func (s *LeaderRoutedStore) MigrationHLCFloor(ctx context.Context, jobID uint64) (uint64, error) { + if s == nil || s.local == nil { + return 0, errors.WithStack(store.ErrNotSupported) + } + floor, err := s.local.MigrationHLCFloor(ctx, jobID) + return floor, errors.WithStack(err) +} + func (s *LeaderRoutedStore) Snapshot() (store.Snapshot, error) { if s == nil || s.local == nil { return nil, errors.WithStack(store.ErrNotSupported) diff --git a/kv/leader_routed_store_test.go b/kv/leader_routed_store_test.go index b2b9f8304..650cb0096 100644 --- a/kv/leader_routed_store_test.go +++ b/kv/leader_routed_store_test.go @@ -199,47 +199,19 @@ func TestLeaderRoutedStore_ScanAtWithReadFenceFiltersRouteBoundsLocally(t *testi s := NewLeaderRoutedStore(local, coord) t.Cleanup(func() { _ = s.Close() }) - kvs, err := s.ScanAtWithReadFence(ctx, rawPrefix, prefixScanEnd(rawPrefix), 1, 2, false, 0, 0, []byte("m"), nil) + kvs, err := s.ScanAtWithReadFence(ctx, rawPrefix, prefixScanEnd(rawPrefix), 1, 2, false, 0, 7, []byte("m"), nil) require.NoError(t, err) require.Len(t, kvs, 1) require.Equal(t, right, kvs[0].Key) require.Equal(t, []byte("right"), kvs[0].Value) - kvs, err = s.ScanAtWithReadFence(ctx, rawPrefix, prefixScanEnd(rawPrefix), 1, 2, true, 0, 0, []byte{}, []byte("m")) + kvs, err = s.ScanAtWithReadFence(ctx, rawPrefix, prefixScanEnd(rawPrefix), 1, 2, true, 0, 7, []byte{}, []byte("m")) require.NoError(t, err) require.Len(t, kvs, 1) require.Equal(t, left, kvs[0].Key) require.Equal(t, []byte("left"), kvs[0].Value) } -func TestLeaderRoutedStore_RejectsLocalReadRouteVersion(t *testing.T) { - t.Parallel() - - ctx := context.Background() - local := store.NewMVCCStore() - require.NoError(t, local.PutAt(ctx, []byte("k"), []byte("v"), 10, 0)) - require.NoError(t, local.PutAt(ctx, []byte("a"), []byte("va"), 10, 0)) - - coord := &stubLeaderCoordinator{ - isLeader: true, - clock: NewHLC(), - } - s := NewLeaderRoutedStore(local, coord) - t.Cleanup(func() { _ = s.Close() }) - - _, err := s.GetAtWithReadFence(ctx, []byte("k"), 10, 0, 7) - require.ErrorIs(t, err, store.ErrNotSupported) - - _, _, err = s.LatestCommitTSWithReadFence(ctx, []byte("k"), 7) - require.ErrorIs(t, err, store.ErrNotSupported) - - _, err = s.ScanAtWithReadFence(ctx, []byte("a"), []byte("z"), 10, 10, false, 0, 7, nil, nil) - require.ErrorIs(t, err, store.ErrNotSupported) - - _, err = s.ScanKeysAtWithReadFence(ctx, []byte("a"), []byte("z"), 10, 10, 0, 7) - require.ErrorIs(t, err, store.ErrNotSupported) -} - func TestLeaderRoutedStore_PrefersLinearizableReadFence(t *testing.T) { t.Parallel() @@ -460,3 +432,39 @@ func TestLeaderRoutedStore_GlobalLastCommitTS_FallsBackWhenNoLeader(t *testing.T ts := s.GlobalLastCommitTS(ctx) require.Equal(t, uint64(7), ts) } + +func TestLeaderRoutedStore_ExportVersionsUsesLinearizableFence(t *testing.T) { + t.Parallel() + + ctx := context.Background() + local := store.NewMVCCStore() + require.NoError(t, local.PutAt(ctx, []byte("k"), []byte("v"), 10, 0)) + + coord := &stubLeaderCoordinator{isLeader: true, clock: NewHLC()} + s := NewLeaderRoutedStore(local, coord) + t.Cleanup(func() { _ = s.Close() }) + + result, err := s.ExportVersions(ctx, store.ExportVersionsOptions{MaxVersions: 10}) + require.NoError(t, err) + require.Len(t, result.Versions, 1) + require.Equal(t, []byte("k"), result.Versions[0].Key) + require.Equal(t, uint64(10), result.Versions[0].CommitTS) + require.Equal(t, 1, coord.linearizableCalls) +} + +func TestLeaderRoutedStore_ExportVersionsFailsClosedWithoutFence(t *testing.T) { + t.Parallel() + + ctx := context.Background() + local := store.NewMVCCStore() + require.NoError(t, local.PutAt(ctx, []byte("k"), []byte("v"), 10, 0)) + + coord := &stubLeaderCoordinator{isLeader: false, clock: NewHLC()} + s := NewLeaderRoutedStore(local, coord) + t.Cleanup(func() { _ = s.Close() }) + + result, err := s.ExportVersions(ctx, store.ExportVersionsOptions{MaxVersions: 10}) + require.ErrorIs(t, err, ErrLeaderNotFound) + require.Empty(t, result.Versions) + require.Equal(t, 1, coord.linearizableCalls) +} diff --git a/kv/lease_read_test.go b/kv/lease_read_test.go index 19bb598f2..6fab2ead0 100644 --- a/kv/lease_read_test.go +++ b/kv/lease_read_test.go @@ -26,6 +26,7 @@ type fakeLeaseEngine struct { proposeCalls atomic.Int32 proposeHook func() // invoked inside Propose before returning (race injection) proposeApply func([]byte) // invoked after a successful propose (FSM apply simulation) + proposeCtxHook func(context.Context) state atomic.Value // stores raftengine.State; default Leader lastQuorumAckMonoNs atomic.Int64 // 0 = no ack yet. Updated by setQuorumAck(). leaderLossCallbacksMu sync.Mutex @@ -64,8 +65,11 @@ func (e *fakeLeaseEngine) Status() raftengine.Status { func (e *fakeLeaseEngine) Configuration(context.Context) (raftengine.Configuration, error) { return raftengine.Configuration{}, nil } -func (e *fakeLeaseEngine) Propose(_ context.Context, data []byte) (*raftengine.ProposalResult, error) { +func (e *fakeLeaseEngine) Propose(ctx context.Context, data []byte) (*raftengine.ProposalResult, error) { e.proposeCalls.Add(1) + if e.proposeCtxHook != nil { + e.proposeCtxHook(ctx) + } if e.proposeHook != nil { e.proposeHook() } diff --git a/kv/lease_warmup_test.go b/kv/lease_warmup_test.go index 22385c7ea..18d64d2d6 100644 --- a/kv/lease_warmup_test.go +++ b/kv/lease_warmup_test.go @@ -169,6 +169,41 @@ func TestCoordinate_RunHLCLeaseRenewal_BlockerSuppressesProposals(t *testing.T) "HLC renewal should resume after startup rotation blocker clears") } +func TestCoordinate_RunHLCLeaseRenewal_UsesRenewalTimeout(t *testing.T) { + eng := &fakeLeaseEngine{applied: 11, leaseDur: time.Hour} + c := NewCoordinatorWithEngine(nil, eng) + deadline := make(chan time.Duration, 1) + eng.proposeCtxHook = func(ctx context.Context) { + d, ok := ctx.Deadline() + if !ok { + deadline <- 0 + return + } + deadline <- time.Until(d) + } + + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan struct{}) + go func() { + c.RunHLCLeaseRenewal(ctx) + close(done) + }() + t.Cleanup(func() { + cancel() + <-done + }) + + select { + case got := <-deadline: + require.Greater(t, got, hlcRenewalInterval, + "HLC renewal proposal deadline must outlive the renewal cadence") + require.LessOrEqual(t, got, hlcRenewalTimeout, + "HLC renewal proposal deadline must remain bounded") + case <-time.After(2 * hlcRenewalInterval): + t.Fatal("timed out waiting for HLC renewal proposal") + } +} + // TestShardedCoordinator_RenewHLCLease_WarmsGroupLease proves the // sharded renewal path warms the target group's lease on a successful // propose, so LeaseReadForKey on a key owned by that group serves from the @@ -294,61 +329,32 @@ func TestShardedCoordinator_RenewHLCLeases_ProposesToEveryLedGroup(t *testing.T) "the non-default group lease must be warmed by all-group renewal") } -func TestShardedCoordinator_RecoverHLCLease_ProposesToEveryLedGroup(t *testing.T) { +func TestShardedCoordinator_RenewHLCLeases_UsesRenewalTimeout(t *testing.T) { t.Parallel() - clock := NewHLC() - clock.SetPhysicalCeiling(time.Now().Add(-time.Millisecond).UnixMilli()) eng1 := newShardedLeaseEngine(100) eng2 := newShardedLeaseEngine(200) - eng1.proposeApply = applyHLCLeaseEntryToClock(t, clock) - eng2.proposeApply = applyHLCLeaseEntryToClock(t, clock) + deadline := make(chan time.Duration, 1) + eng1.proposeCtxHook = func(ctx context.Context) { + d, ok := ctx.Deadline() + if !ok { + deadline <- 0 + return + } + deadline <- time.Until(d) + } coord := mustShardedLeaseCoord(t, eng1, eng2) - coord.clock = clock - - require.NoError(t, coord.RecoverHLCLease(context.Background())) - require.Equal(t, int32(1), eng1.proposeCalls.Load()) - require.Equal(t, int32(1), eng2.proposeCalls.Load()) - got, err := clock.NextFenced() - require.NoError(t, err) - require.NotZero(t, got) -} + done := coord.renewHLCLeases(context.Background()) + requireRenewalDone(t, done) -func TestShardedCoordinator_RecoverHLCLease_SucceedsWhenAnyTargetAdvancesCeiling(t *testing.T) { - t.Parallel() - for _, tc := range []struct { - name string - failFirst bool - failSecond bool - }{ - {name: "first fails then second advances", failFirst: true}, - {name: "first advances then second fails", failSecond: true}, - } { - t.Run(tc.name, func(t *testing.T) { - t.Parallel() - clock := NewHLC() - clock.SetPhysicalCeiling(time.Now().Add(-time.Millisecond).UnixMilli()) - eng1 := newShardedLeaseEngine(100) - eng2 := newShardedLeaseEngine(200) - eng1.proposeApply = applyHLCLeaseEntryToClock(t, clock) - eng2.proposeApply = applyHLCLeaseEntryToClock(t, clock) - if tc.failFirst { - eng1.proposeErr = errors.New("group 1 unavailable") - } - if tc.failSecond { - eng2.proposeErr = errors.New("group 2 unavailable") - } - coord := mustShardedLeaseCoord(t, eng1, eng2) - coord.clock = clock - - require.NoError(t, coord.RecoverHLCLease(context.Background())) - require.Equal(t, int32(1), eng1.proposeCalls.Load()) - require.Equal(t, int32(1), eng2.proposeCalls.Load()) - - got, err := clock.NextFenced() - require.NoError(t, err) - require.NotZero(t, got) - }) + select { + case got := <-deadline: + require.Greater(t, got, hlcRenewalInterval, + "HLC renewal proposal deadline must outlive the renewal cadence") + require.LessOrEqual(t, got, hlcRenewalTimeout, + "HLC renewal proposal deadline must remain bounded") + default: + t.Fatal("missing HLC renewal proposal deadline sample") } } diff --git a/kv/migrator_filter.go b/kv/migrator_filter.go new file mode 100644 index 000000000..acbf1a796 --- /dev/null +++ b/kv/migrator_filter.go @@ -0,0 +1,80 @@ +package kv + +import ( + "bytes" + + "github.com/bootjp/elastickv/internal/s3keys" +) + +// RouteKeyFilter returns the migration export predicate for raw MVCC keys. +// rangeEnd nil or empty means +infinity, matching the route descriptor wire +// convention. +func RouteKeyFilter(rangeStart, rangeEnd []byte) func([]byte) bool { + return RouteKeyFilterForGroup(rangeStart, rangeEnd, 0, nil) +} + +// RouteKeyFilterForGroup returns the migration export predicate for a source +// route and group. Partition-resolved keyspaces such as HT-FIFO SQS are matched +// by resolver group instead of the byte-range route key. +func RouteKeyFilterForGroup(rangeStart, rangeEnd []byte, sourceGroupID uint64, resolver PartitionResolver) func([]byte) bool { + start := bytes.Clone(rangeStart) + end := bytes.Clone(rangeEnd) + return func(rawKey []byte) bool { + if resolver != nil { + if gid, ok := resolver.ResolveGroup(rawKey); ok { + return gid == sourceGroupID + } + if resolver.RecognisesPartitionedKey(rawKey) { + return false + } + } + if s3BucketAuxiliaryRouteInRange(rawKey, start, end) { + return true + } + rkey := routeKey(rawKey) + return keyInMigrationRouteRange(rkey, start, end) + } +} + +func s3BucketAuxiliaryRouteInRange(rawKey, routeStart, routeEnd []byte) bool { + bucketRouteStart, bucketRouteEnd, ok := s3BucketAuxiliaryRouteRange(rawKey) + if !ok { + return false + } + if keyInMigrationRouteRange(rawKey, routeStart, routeEnd) { + return true + } + return migrationRouteRangesIntersect(routeStart, routeEnd, bucketRouteStart, bucketRouteEnd) +} + +func s3BucketAuxiliaryRouteRange(rawKey []byte) ([]byte, []byte, bool) { + bucket, ok := s3keys.ParseBucketMetaKey(rawKey) + if !ok { + bucket, ok = s3keys.ParseBucketGenerationKey(rawKey) + } + if !ok { + return nil, nil, false + } + bucketRouteStart := s3keys.RoutePrefixForBucketAnyGeneration(bucket) + return bucketRouteStart, prefixScanEnd(bucketRouteStart), true +} + +func keyInMigrationRouteRange(key, routeStart, routeEnd []byte) bool { + if key == nil { + return false + } + if bytes.Compare(key, routeStart) < 0 { + return false + } + return len(routeEnd) == 0 || bytes.Compare(key, routeEnd) < 0 +} + +func migrationRouteRangesIntersect(aStart, aEnd, bStart, bEnd []byte) bool { + if len(aEnd) > 0 && bytes.Compare(aEnd, bStart) <= 0 { + return false + } + if len(bEnd) > 0 && bytes.Compare(bEnd, aStart) <= 0 { + return false + } + return true +} diff --git a/kv/migrator_lock_drain.go b/kv/migrator_lock_drain.go new file mode 100644 index 000000000..ef7a7e531 --- /dev/null +++ b/kv/migrator_lock_drain.go @@ -0,0 +1,92 @@ +package kv + +import ( + "bytes" + "context" + + "github.com/bootjp/elastickv/store" + "github.com/cockroachdb/errors" +) + +// TxnLockDrainEntry describes one prepared transaction lock that still belongs +// to a migration route. The lock key is cloned so callers can safely retain it +// across drain ticks. +type TxnLockDrainEntry struct { + LockKey []byte + UserKey []byte + StartTS uint64 + TTLExpireAt uint64 + PrimaryKey []byte + IsPrimaryKey bool +} + +// PendingTxnLocksInRoute scans the txn-lock namespace and filters each lock by +// routeKey(userKey). It intentionally does not bracket the scan by +// txnLockKey(routeStart)/txnLockKey(routeEnd): txn locks are sorted by raw user +// key while migration routes are defined in route-key space. +func PendingTxnLocksInRoute(ctx context.Context, st store.MVCCStore, routeStart, routeEnd []byte, ts uint64, limit int) ([]TxnLockDrainEntry, error) { + if st == nil { + return nil, nil + } + if limit <= 0 { + limit = maxTxnLockScanResults + } + start := txnLockKey(nil) + end := prefixScanEnd(start) + filter := RouteKeyFilter(routeStart, routeEnd) + cursor := start + out := make([]TxnLockDrainEntry, 0, min(limit, lockPageLimit)) + + for { + lockKVs, nextCursor, done, err := scanTxnLockDrainPage(ctx, st, cursor, end, ts) + if err != nil { + return nil, err + } + for _, kvp := range lockKVs { + entry, ok, err := txnLockDrainEntry(kvp, filter) + if err != nil { + return nil, err + } + if !ok { + continue + } + out = append(out, entry) + if len(out) >= limit { + return out, nil + } + } + if done { + return out, nil + } + cursor = nextCursor + } +} + +func scanTxnLockDrainPage(ctx context.Context, st store.MVCCStore, cursor, end []byte, ts uint64) ([]*store.KVPair, []byte, bool, error) { + if err := ctx.Err(); err != nil { + return nil, nil, false, errors.WithStack(err) + } + return scanTxnLockPageAt(ctx, st, cursor, end, ts) +} + +func txnLockDrainEntry(kvp *store.KVPair, filter func([]byte) bool) (TxnLockDrainEntry, bool, error) { + if kvp == nil || !bytes.HasPrefix(kvp.Key, txnLockPrefixBytes) { + return TxnLockDrainEntry{}, false, nil + } + userKey := kvp.Key[len(txnLockPrefixBytes):] + if !filter(userKey) { + return TxnLockDrainEntry{}, false, nil + } + lock, err := decodeTxnLock(kvp.Value) + if err != nil { + return TxnLockDrainEntry{}, false, errors.Wrap(err, "decode txn lock during migration drain") + } + return TxnLockDrainEntry{ + LockKey: bytes.Clone(kvp.Key), + UserKey: bytes.Clone(userKey), + StartTS: lock.StartTS, + TTLExpireAt: lock.TTLExpireAt, + PrimaryKey: bytes.Clone(lock.PrimaryKey), + IsPrimaryKey: lock.IsPrimaryKey, + }, true, nil +} diff --git a/kv/migrator_lock_drain_test.go b/kv/migrator_lock_drain_test.go new file mode 100644 index 000000000..d790e85c0 --- /dev/null +++ b/kv/migrator_lock_drain_test.go @@ -0,0 +1,56 @@ +package kv + +import ( + "bytes" + "context" + "testing" + + "github.com/bootjp/elastickv/store" + "github.com/stretchr/testify/require" +) + +func TestPendingTxnLocksInRouteFiltersLocksByRouteKey(t *testing.T) { + t.Parallel() + + ctx := context.Background() + st := store.NewMVCCStore() + tableSegment := "tenant-a" + routeStart := []byte(dynamoRoutePrefix + tableSegment) + routeEnd := append(bytes.Clone(routeStart), 0xff) + itemKey := append([]byte(DynamoItemPrefix+tableSegment+"|7|"), []byte("pk\x00\x01")...) + outsideKey := []byte("outside") + + require.Less(t, bytes.Compare(itemKey, routeStart), 0, + "the red-control key must sort outside a txnLockKey(routeStart)/txnLockKey(routeEnd) bracket") + lock := encodeTxnLock(txnLock{ + StartTS: 11, + TTLExpireAt: 99, + PrimaryKey: itemKey, + IsPrimaryKey: true, + }) + require.NoError(t, st.PutAt(ctx, txnLockKey(itemKey), lock, 1, 0)) + require.NoError(t, st.PutAt(ctx, txnLockKey(outsideKey), encodeTxnLock(txnLock{ + StartTS: 12, + PrimaryKey: outsideKey, + }), 2, 0)) + + pending, err := PendingTxnLocksInRoute(ctx, st, routeStart, routeEnd, ^uint64(0), 10) + require.NoError(t, err) + require.Len(t, pending, 1) + require.Equal(t, itemKey, pending[0].UserKey) + require.Equal(t, uint64(11), pending[0].StartTS) + require.Equal(t, uint64(99), pending[0].TTLExpireAt) + require.True(t, pending[0].IsPrimaryKey) + require.Equal(t, txnLockKey(itemKey), pending[0].LockKey) +} + +func TestPendingTxnLocksInRouteHonorsCanceledContext(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + pending, err := PendingTxnLocksInRoute(ctx, store.NewMVCCStore(), nil, nil, ^uint64(0), 10) + require.ErrorIs(t, err, context.Canceled) + require.Nil(t, pending) +} diff --git a/kv/route_history.go b/kv/route_history.go index cb768b30b..c6ac85608 100644 --- a/kv/route_history.go +++ b/kv/route_history.go @@ -60,3 +60,17 @@ func (s distributionRouteSnapshot) Version() uint64 { func (s distributionRouteSnapshot) OwnerOf(key []byte) (uint64, bool) { return s.snap.OwnerOf(key) } + +func (s distributionRouteSnapshot) WriteFencedForKey(key []byte) bool { + route, ok := s.snap.RouteOf(key) + return ok && route.State == distribution.RouteStateWriteFenced +} + +func (s distributionRouteSnapshot) WriteFencedIntersects(start, end []byte) bool { + for _, route := range s.snap.IntersectingRoutes(start, end) { + if route.State == distribution.RouteStateWriteFenced { + return true + } + } + return false +} diff --git a/kv/shard_key.go b/kv/shard_key.go index d4132c8c1..369c456f9 100644 --- a/kv/shard_key.go +++ b/kv/shard_key.go @@ -40,6 +40,17 @@ const ( // family (!sqs|queue|meta|, !sqs|msg|vis|, etc.). Used by // sqsRouteKey to dispatch the routing decision. sqsInternalPrefix = "!sqs|" + + sqsQueueMetaPrefix = "!sqs|queue|meta|" + sqsQueueGenPrefix = "!sqs|queue|gen|" + sqsQueueSeqPrefix = "!sqs|queue|seq|" + sqsQueueTombstonePrefix = "!sqs|queue|tombstone|" + sqsMsgDataPrefix = "!sqs|msg|data|" + sqsMsgVisPrefix = "!sqs|msg|vis|" + sqsMsgDedupPrefix = "!sqs|msg|dedup|" + sqsMsgGroupPrefix = "!sqs|msg|group|" + sqsMsgByAgePrefix = "!sqs|msg|byage|" + sqsPartitionMarker = "p|" ) var ( @@ -48,23 +59,31 @@ var ( dynamoTableGenerationPrefixBytes = []byte(DynamoTableGenerationPrefix) dynamoItemPrefixBytes = []byte(DynamoItemPrefix) dynamoGSIPrefixBytes = []byte(DynamoGSIPrefix) - sqsRoutePrefixBytes = []byte(sqsRoutePrefix) sqsInternalPrefixBytes = []byte(sqsInternalPrefix) - redisWideColumnScanPrefixes = [][]byte{ - []byte(store.HashMetaDeltaPrefix), - []byte(store.HashMetaPrefix), - []byte(store.HashFieldPrefix), - []byte(store.SetMetaDeltaPrefix), - []byte(store.SetMetaPrefix), - []byte(store.SetMemberPrefix), - []byte(store.ZSetMetaDeltaPrefix), - []byte(store.ZSetMetaPrefix), - []byte(store.ZSetMemberPrefix), - []byte(store.ZSetScorePrefix), - } - redisListAuxiliaryScanPrefixes = [][]byte{ - []byte(store.ListMetaDeltaPrefix), - []byte(store.ListClaimPrefix), + sqsGlobalRouteKey = []byte(sqsRoutePrefix + "global") + sqsConcreteInternalPrefixBytes = [][]byte{ + []byte(sqsQueueMetaPrefix), + []byte(sqsQueueGenPrefix), + []byte(sqsQueueSeqPrefix), + []byte(sqsQueueTombstonePrefix), + []byte(sqsMsgDataPrefix), + []byte(sqsMsgVisPrefix), + []byte(sqsMsgDedupPrefix), + []byte(sqsMsgGroupPrefix), + []byte(sqsMsgByAgePrefix), + } + routeKeyExtractors = []func([]byte) []byte{ + redisRouteKey, + dynamoRouteKey, + sqsRouteKey, + s3keys.ExtractRouteKey, + fskeys.ExtractRouteKey, + listRouteKey, + hashRouteKey, + setRouteKey, + zsetRouteKey, + streamRouteKey, + store.ExtractListUserKey, } ) @@ -95,35 +114,19 @@ func routeFilterKey(key []byte) []byte { } func normalizeRouteKey(key []byte) []byte { - if user := redisRouteKey(key); user != nil { - return user - } - if user := redisWideColumnRouteKey(key); user != nil { - return user - } - if table := dynamoRouteKey(key); table != nil { - return table - } - if route := sqsRouteKey(key); route != nil { - return route - } - if user := s3keys.ExtractRouteKey(key); user != nil { - return user - } - if user := fskeys.ExtractRouteKey(key); user != nil { - return user - } - if user := store.ExtractListUserKey(key); user != nil { - return user + for _, extract := range routeKeyExtractors { + if user := extract(key); user != nil { + return user + } } return key } func normalizeRouteFilterKey(key []byte) []byte { - if user := redisListAuxiliaryRouteKey(key); user != nil { + if user := listRouteKey(key); user != nil { return user } - if user := redisStreamRouteKey(key); user != nil { + if user := streamRouteKey(key); user != nil { return user } return normalizeRouteKey(key) @@ -140,17 +143,28 @@ func redisWideColumnLegacyPointRouteKey(key []byte) []byte { } func redisWideColumnRouteKey(key []byte) []byte { - if user := redisHashRouteKey(key); user != nil { + if user := hashRouteKey(key); user != nil { return user } - if user := redisSetRouteKey(key); user != nil { + if user := setRouteKey(key); user != nil { return user } - return redisZSetRouteKey(key) + return zsetRouteKey(key) } func redisWideColumnScanRouteParts(key []byte) (prefix []byte, userKey []byte, userPrefix []byte, owned bool, parsed bool) { - for _, prefix := range redisWideColumnScanPrefixes { + for _, prefix := range [][]byte{ + []byte(store.HashMetaDeltaPrefix), + []byte(store.HashMetaPrefix), + []byte(store.HashFieldPrefix), + []byte(store.SetMetaDeltaPrefix), + []byte(store.SetMetaPrefix), + []byte(store.SetMemberPrefix), + []byte(store.ZSetMetaDeltaPrefix), + []byte(store.ZSetMetaPrefix), + []byte(store.ZSetMemberPrefix), + []byte(store.ZSetScorePrefix), + } { if !bytes.HasPrefix(key, prefix) { continue } @@ -164,14 +178,6 @@ func redisWideColumnScanRouteParts(key []byte) (prefix []byte, userKey []byte, u return nil, nil, nil, false, false } -func redisWideColumnLegacyScanRouteRange(start []byte, end []byte) ([]byte, []byte, bool) { - _, _, _, owned, parsed := redisWideColumnScanRouteParts(start) - if !owned || !parsed { - return nil, nil, false - } - return start, end, true -} - func redisWideColumnScanRouteRange(start []byte, end []byte) (routeStart []byte, routeEnd []byte, exact bool, ok bool) { prefix, userKey, userPrefix, owned, parsed := redisWideColumnScanRouteParts(start) if !owned { @@ -193,7 +199,10 @@ func redisWideColumnScanRouteRange(start []byte, end []byte) (routeStart []byte, } func listAuxiliaryScanRouteRange(start []byte, end []byte) (routeStart []byte, exact bool, ok bool) { - for _, prefix := range redisListAuxiliaryScanPrefixes { + for _, prefix := range [][]byte{ + []byte(store.ListMetaDeltaPrefix), + []byte(store.ListClaimPrefix), + } { if !bytes.HasPrefix(start, prefix) { continue } @@ -227,7 +236,62 @@ func wideColumnScanUserKey(key []byte, prefix []byte) []byte { return rest[:keyLen] } -func redisHashRouteKey(key []byte) []byte { +func redisRouteKey(key []byte) []byte { + if !bytes.HasPrefix(key, redisInternalRoutePrefixBytes) { + return nil + } + rest := key[len(redisInternalRoutePrefix):] + sep := bytes.IndexByte(rest, '|') + if sep < 0 || sep+1 >= len(rest) { + return nil + } + return rest[sep+1:] +} + +func dynamoRouteKey(key []byte) []byte { + switch { + case bytes.HasPrefix(key, dynamoTableMetaPrefixBytes): + return dynamoRouteTableKey(key[len(dynamoTableMetaPrefixBytes):]) + case bytes.HasPrefix(key, dynamoTableGenerationPrefixBytes): + return dynamoRouteTableKey(key[len(dynamoTableGenerationPrefixBytes):]) + case bytes.HasPrefix(key, dynamoItemPrefixBytes): + return dynamoRouteFromTablePrefixedKey(key[len(dynamoItemPrefixBytes):]) + case bytes.HasPrefix(key, dynamoGSIPrefixBytes): + return dynamoRouteFromTablePrefixedKey(key[len(dynamoGSIPrefixBytes):]) + default: + return nil + } +} + +func dynamoRouteFromTablePrefixedKey(rest []byte) []byte { + sep := bytes.IndexByte(rest, '|') + if sep <= 0 { + return nil + } + return dynamoRouteTableKey(rest[:sep]) +} + +func dynamoRouteTableKey(tableSegment []byte) []byte { + if len(tableSegment) == 0 { + return nil + } + out := make([]byte, 0, len(dynamoRoutePrefixBytes)+len(tableSegment)) + out = append(out, dynamoRoutePrefixBytes...) + out = append(out, tableSegment...) + return out +} + +func listRouteKey(key []byte) []byte { + if userKey := store.ExtractListUserKeyFromDelta(key); userKey != nil { + return userKey + } + if userKey := store.ExtractListUserKeyFromClaim(key); userKey != nil { + return userKey + } + return nil +} + +func hashRouteKey(key []byte) []byte { switch { case store.IsHashMetaDeltaKey(key): return store.ExtractHashUserKeyFromDelta(key) @@ -240,7 +304,7 @@ func redisHashRouteKey(key []byte) []byte { } } -func redisSetRouteKey(key []byte) []byte { +func setRouteKey(key []byte) []byte { switch { case store.IsSetMetaDeltaKey(key): return store.ExtractSetUserKeyFromDelta(key) @@ -253,7 +317,7 @@ func redisSetRouteKey(key []byte) []byte { } } -func redisZSetRouteKey(key []byte) []byte { +func zsetRouteKey(key []byte) []byte { switch { case store.IsZSetMetaDeltaKey(key): return store.ExtractZSetUserKeyFromDelta(key) @@ -268,18 +332,7 @@ func redisZSetRouteKey(key []byte) []byte { } } -func redisListAuxiliaryRouteKey(key []byte) []byte { - switch { - case store.IsListMetaDeltaKey(key): - return store.ExtractListUserKeyFromDelta(key) - case store.IsListClaimKey(key): - return store.ExtractListUserKeyFromClaim(key) - default: - return nil - } -} - -func redisStreamRouteKey(key []byte) []byte { +func streamRouteKey(key []byte) []byte { switch { case store.IsStreamMetaKey(key): return store.ExtractStreamUserKeyFromMeta(key) @@ -290,64 +343,24 @@ func redisStreamRouteKey(key []byte) []byte { } } -func redisRouteKey(key []byte) []byte { - if !bytes.HasPrefix(key, redisInternalRoutePrefixBytes) { - return nil - } - rest := key[len(redisInternalRoutePrefix):] - sep := bytes.IndexByte(rest, '|') - if sep < 0 || sep+1 >= len(rest) { - return nil - } - return rest[sep+1:] -} - -func dynamoRouteKey(key []byte) []byte { - switch { - case bytes.HasPrefix(key, dynamoTableMetaPrefixBytes): - return dynamoRouteTableKey(key[len(dynamoTableMetaPrefixBytes):]) - case bytes.HasPrefix(key, dynamoTableGenerationPrefixBytes): - return dynamoRouteTableKey(key[len(dynamoTableGenerationPrefixBytes):]) - case bytes.HasPrefix(key, dynamoItemPrefixBytes): - return dynamoRouteFromTablePrefixedKey(key[len(dynamoItemPrefixBytes):]) - case bytes.HasPrefix(key, dynamoGSIPrefixBytes): - return dynamoRouteFromTablePrefixedKey(key[len(dynamoGSIPrefixBytes):]) - default: - return nil - } -} - -func dynamoRouteFromTablePrefixedKey(rest []byte) []byte { - sep := bytes.IndexByte(rest, '|') - if sep <= 0 { +// sqsRouteKey maps concrete persisted !sqs|... storage prefixes to a stable +// route key. Adapter-looking raw user keys such as !sqs|foo intentionally stay +// on their raw route and migrate through the user-key bracket. +func sqsRouteKey(key []byte) []byte { + if !bytes.HasPrefix(key, sqsInternalPrefixBytes) { return nil } - return dynamoRouteTableKey(rest[:sep]) -} - -func dynamoRouteTableKey(tableSegment []byte) []byte { - if len(tableSegment) == 0 { + if !hasSQSConcreteInternalPrefix(key) { return nil } - out := make([]byte, 0, len(dynamoRoutePrefixBytes)+len(tableSegment)) - out = append(out, dynamoRoutePrefixBytes...) - out = append(out, tableSegment...) - return out + return sqsGlobalRouteKey } -// sqsRouteKey maps any !sqs|... internal key to a stable route key so -// multi-shard deployments that partition by user-key range still land -// every SQS mutation on a configured group. Milestone 1 collapses all -// SQS keys to a single !sqs|route|global route — this keeps the -// catalog and every queue's message keyspace on the same group, which -// is the minimum needed for FIFO group-lock semantics (landing later) -// to work. When per-queue sharding is implemented it will live here. -func sqsRouteKey(key []byte) []byte { - if !bytes.HasPrefix(key, sqsInternalPrefixBytes) { - return nil +func hasSQSConcreteInternalPrefix(key []byte) bool { + for _, prefix := range sqsConcreteInternalPrefixBytes { + if bytes.HasPrefix(key, prefix) { + return true + } } - out := make([]byte, 0, len(sqsRoutePrefixBytes)+len("global")) - out = append(out, sqsRoutePrefixBytes...) - out = append(out, []byte("global")...) - return out + return false } diff --git a/kv/shard_key_test.go b/kv/shard_key_test.go index de560d250..d8107cc33 100644 --- a/kv/shard_key_test.go +++ b/kv/shard_key_test.go @@ -1,7 +1,9 @@ package kv import ( + "bytes" "encoding/base64" + "encoding/binary" "testing" "github.com/bootjp/elastickv/internal/fskeys" @@ -76,21 +78,6 @@ func TestRouteKey_NormalizesRedisWideColumnKeys(t *testing.T) { } } -func TestRouteFilterKey_NormalizesRedisAuxiliaryKeys(t *testing.T) { - t.Parallel() - - userKey := []byte("user:key") - for _, raw := range [][]byte{ - store.ListMetaDeltaKey(userKey, 10, 0), - store.ListClaimKey(userKey, 1), - store.StreamMetaKey(userKey), - store.StreamEntryKey(userKey, 123, 4), - } { - require.Equal(t, userKey, routeFilterKey(raw)) - require.Equal(t, userKey, routeFilterKey(txnLockKey(raw))) - } -} - func TestRedisWideColumnScanRouteRangeFansOutBareFamilyAndCursor(t *testing.T) { t.Parallel() @@ -122,35 +109,6 @@ func TestRedisWideColumnScanRouteRangeFansOutBareFamilyAndCursor(t *testing.T) { require.Nil(t, routeEnd) } -func TestListAuxiliaryScanRouteRangeFansOutBareFamilyAndCursor(t *testing.T) { - t.Parallel() - - prefix := []byte(store.ListMetaDeltaPrefix) - familyEnd := prefixScanEnd(prefix) - userPrefix := store.ListMetaDeltaScanPrefix([]byte("alice")) - cursor := store.ListMetaDeltaKey([]byte("alice"), 10, 0) - - for _, tc := range []struct { - name string - start []byte - }{ - {name: "bare family", start: prefix}, - {name: "physical cursor", start: cursor}, - } { - t.Run(tc.name, func(t *testing.T) { - routeStart, exact, ok := listAuxiliaryScanRouteRange(tc.start, familyEnd) - require.True(t, ok) - require.False(t, exact) - require.Nil(t, routeStart) - }) - } - - routeStart, exact, ok := listAuxiliaryScanRouteRange(userPrefix, prefixScanEnd(userPrefix)) - require.True(t, ok) - require.True(t, exact) - require.Equal(t, []byte("alice"), routeStart) -} - func TestRouteKey_NormalizesDynamoKeysToTable(t *testing.T) { t.Parallel() @@ -203,3 +161,337 @@ func TestRouteKey_CollapsesDynamoGenerationsToSameTableRoute(t *testing.T) { require.Equal(t, want, routeKey(sourceGSIKey), "migration source generation GSI key must route to the same table group as the current generation") } + +func TestRouteKey_NormalizesCollectionMigrationFamilies(t *testing.T) { + t.Parallel() + + userKey := []byte("redis:user|with|separators") + cases := [][]byte{ + store.ListMetaDeltaKey(userKey, 10, 1), + store.ListClaimKey(userKey, -2), + store.HashMetaKey(userKey), + store.HashFieldKey(userKey, []byte("field")), + store.HashMetaDeltaKey(userKey, 11, 2), + store.SetMetaKey(userKey), + store.SetMemberKey(userKey, []byte("member")), + store.SetMetaDeltaKey(userKey, 12, 3), + store.ZSetMetaKey(userKey), + store.ZSetMemberKey(userKey, []byte("member")), + store.ZSetScoreKey(userKey, 1.25, []byte("member")), + store.ZSetMetaDeltaKey(userKey, 13, 4), + store.StreamMetaKey(userKey), + store.StreamEntryKey(userKey, 14, 5), + } + + for _, raw := range cases { + require.Equal(t, userKey, routeKey(raw), "raw key %q must route by its logical user key", raw) + } +} + +func TestRouteKey_ListMetaKeyThatLooksLikeNewDeltaRoutesByRealListKey(t *testing.T) { + t.Parallel() + + userKey := []byte(store.ListMetaDeltaPrefix + "fake:user") + baseMeta := store.ListMetaKey(userKey) + deltaKey := store.ListMetaDeltaKey(userKey, 10, 1) + + require.Equal(t, userKey, routeKey(baseMeta), "base list metadata must not decode as a new delta") + require.Equal(t, userKey, routeKey(deltaKey), "real list deltas must still route by the logical list key") +} + +func TestRouteKey_LegacyListDeltaKeyOnlyUsesBaseMetaRoute(t *testing.T) { + t.Parallel() + + userKey := []byte("legacy:list") + raw := legacyListMetaDeltaKey(userKey, 42, 7) + require.Equal(t, store.ExtractListUserKey(raw), routeKey(raw)) + require.NotEqual(t, userKey, routeKey(raw), "key-only routing must not decode ambiguous legacy deltas") + + collidingUserKey := deltaLookingListMetaUserKey(userKey, 42, 7) + collidingMeta := store.ListMetaKey(collidingUserKey) + require.Equal(t, collidingUserKey, routeKey(collidingMeta)) +} + +func TestRouteKey_MalformedWideColumnKeysFallBackToRaw(t *testing.T) { + t.Parallel() + + for _, raw := range [][]byte{ + malformedWideColumnKey(store.ListClaimPrefix, 8), + malformedWideColumnKey(store.HashMetaPrefix, 0), + malformedWideColumnKey(store.HashFieldPrefix, 0), + malformedWideColumnKey(store.HashMetaDeltaPrefix, 12), + malformedWideColumnKey(store.SetMetaPrefix, 0), + malformedWideColumnKey(store.SetMemberPrefix, 0), + malformedWideColumnKey(store.SetMetaDeltaPrefix, 12), + malformedWideColumnKey(store.ZSetMetaPrefix, 0), + malformedWideColumnKey(store.ZSetMemberPrefix, 0), + malformedWideColumnKey(store.ZSetScorePrefix, 8), + malformedWideColumnKey(store.ZSetMetaDeltaPrefix, 12), + } { + require.NotPanics(t, func() { + require.Equal(t, raw, routeKey(raw), "malformed key %q must not decode to a logical route", raw) + }) + } +} + +func malformedWideColumnKey(prefix string, suffixLen int) []byte { + key := make([]byte, 0, len(prefix)+4+suffixLen) + key = append(key, prefix...) + var lenPrefix [4]byte + binary.BigEndian.PutUint32(lenPrefix[:], ^uint32(0)) + key = append(key, lenPrefix[:]...) + key = append(key, make([]byte, suffixLen)...) + return key +} + +func legacyListMetaDeltaKey(userKey []byte, commitTS uint64, seqInTxn uint32) []byte { + key := store.LegacyListMetaDeltaScanPrefix(userKey) + var ts [8]byte + binary.BigEndian.PutUint64(ts[:], commitTS) + key = append(key, ts[:]...) + var seq [4]byte + binary.BigEndian.PutUint32(seq[:], seqInTxn) + return append(key, seq[:]...) +} + +func deltaLookingListMetaUserKey(fakeUserKey []byte, commitTS uint64, seqInTxn uint32) []byte { + key := make([]byte, 0, len("d|")+4+len(fakeUserKey)+8+4) + key = append(key, "d|"...) + var lenPrefix [4]byte + binary.BigEndian.PutUint32(lenPrefix[:], uint32(len(fakeUserKey))) //nolint:gosec // test data is small. + key = append(key, lenPrefix[:]...) + key = append(key, fakeUserKey...) + var ts [8]byte + binary.BigEndian.PutUint64(ts[:], commitTS) + key = append(key, ts[:]...) + var seq [4]byte + binary.BigEndian.PutUint32(seq[:], seqInTxn) + return append(key, seq[:]...) +} + +func TestRouteKey_NormalizesTxnSuccessMarkerByLockedKey(t *testing.T) { + t.Parallel() + + lockedKey := []byte("secondary|key\x00with|separators") + primaryKey := []byte("primary|key\x00with|separators") + marker := TxnSuccessMarkerKey(lockedKey, 100, 200, primaryKey) + + require.Equal(t, lockedKey, routeKey(marker)) + + malformed := append([]byte(nil), marker...) + malformed[len(txnSuccessPrefixBytes)] = 2 + require.Equal(t, malformed, routeKey(malformed), "malformed success markers must fall back to their raw key") +} + +func TestRouteKey_SQSDecoderIsConcreteOnly(t *testing.T) { + t.Parallel() + + want := []byte(sqsRoutePrefix + "global") + for _, raw := range [][]byte{ + []byte(sqsQueueMetaPrefix + "queue"), + []byte(sqsQueueGenPrefix + "queue"), + []byte(sqsQueueSeqPrefix + "queue"), + []byte(sqsQueueTombstonePrefix + "queue"), + []byte(sqsMsgDataPrefix + "queue|1|msg"), + []byte(sqsMsgVisPrefix + "queue|1|msg"), + []byte(sqsMsgDedupPrefix + "queue|1|dedup"), + []byte(sqsMsgGroupPrefix + "queue|1|group"), + []byte(sqsMsgByAgePrefix + "queue|1|ts"), + []byte(sqsMsgDataPrefix + sqsPartitionMarker + "queue|0|1|msg"), + []byte(sqsMsgVisPrefix + sqsPartitionMarker + "queue|0|1|msg"), + []byte(sqsMsgDedupPrefix + sqsPartitionMarker + "queue|0|1|dedup"), + []byte(sqsMsgGroupPrefix + sqsPartitionMarker + "queue|0|1|group"), + []byte(sqsMsgByAgePrefix + sqsPartitionMarker + "queue|0|1|ts"), + } { + require.Equal(t, want, routeKey(raw), "concrete SQS key %q must use the SQS route", raw) + } + + rawUser := []byte("!sqs|foo") + require.Equal(t, rawUser, routeKey(rawUser), "adapter-looking raw user key must stay on its raw route") +} + +func TestRouteKey_S3DecoderIsConcreteOnly(t *testing.T) { + t.Parallel() + + manifest := s3keys.ObjectManifestKey("bucket", 2, "obj") + require.Equal(t, s3keys.RouteKey("bucket", 2, "obj"), routeKey(manifest)) + + rawUser := []byte("!s3|foo") + require.Equal(t, rawUser, routeKey(rawUser), "adapter-looking raw user key must stay on its raw route") +} + +func TestRoutePrefixRangeTreatsBroadMappedPrefixesAsFullKeyspace(t *testing.T) { + t.Parallel() + + tableSegment := base64.RawURLEncoding.EncodeToString([]byte("users")) + dynamoTableRoute := dynamoRouteTableKey([]byte(tableSegment)) + + for _, tc := range []struct { + name string + prefix []byte + wantStart []byte + wantEnd []byte + }{ + { + name: "raw user prefix", + prefix: []byte("ab"), + wantStart: []byte("ab"), + wantEnd: prefixScanEnd([]byte("ab")), + }, + { + name: "concrete redis prefix", + prefix: []byte("!redis|string|ab"), + wantStart: []byte("ab"), + wantEnd: prefixScanEnd([]byte("ab")), + }, + { + name: "dynamo table cleanup prefix", + prefix: []byte(DynamoItemPrefix + tableSegment + "|7|"), + wantStart: dynamoTableRoute, + wantEnd: routePointRangeEnd(dynamoTableRoute), + }, + { + name: "broad redis namespace", + prefix: []byte("!redis|"), + wantStart: []byte(""), + wantEnd: nil, + }, + { + name: "broad wide-column namespace", + prefix: []byte("!lst|"), + wantStart: []byte(""), + wantEnd: nil, + }, + { + name: "raw sqs-looking user prefix", + prefix: []byte("!sqs|foo"), + wantStart: []byte("!sqs|foo"), + wantEnd: prefixScanEnd([]byte("!sqs|foo")), + }, + { + name: "concrete sqs storage prefix", + prefix: []byte(sqsMsgDataPrefix), + wantStart: sqsGlobalRouteKey, + wantEnd: prefixScanEnd(sqsGlobalRouteKey), + }, + { + name: "s3 bucket cleanup prefix", + prefix: s3keys.ObjectManifestPrefixForBucket("bucket", 2), + wantStart: s3keys.RoutePrefixForBucket("bucket", 2), + wantEnd: prefixScanEnd(s3keys.RoutePrefixForBucket("bucket", 2)), + }, + { + name: "s3 bucket cleanup route prefix", + prefix: s3keys.RoutePrefixForBucket("bucket", 2), + wantStart: s3keys.RoutePrefixForBucket("bucket", 2), + wantEnd: prefixScanEnd(s3keys.RoutePrefixForBucket("bucket", 2)), + }, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + start, end := routePrefixRange(tc.prefix) + require.Equal(t, tc.wantStart, start) + require.Equal(t, tc.wantEnd, end) + }) + } +} + +func TestRoutePrefixRangeTreatsDynamoCleanupAsExactRouteKey(t *testing.T) { + t.Parallel() + + fooSegment := base64.RawURLEncoding.EncodeToString([]byte("foo")) + foobarSegment := base64.RawURLEncoding.EncodeToString([]byte("foobar")) + fooRoute := dynamoRouteTableKey([]byte(fooSegment)) + foobarRoute := dynamoRouteTableKey([]byte(foobarSegment)) + + start, end := routePrefixRange([]byte(DynamoItemPrefix + fooSegment + "|7|")) + + require.Equal(t, fooRoute, start) + require.Equal(t, routePointRangeEnd(fooRoute), end) + require.False(t, rangesIntersectForTest(start, end, foobarRoute, prefixScanEnd(foobarRoute)), + "cleanup for table foo must not intersect table foobar's route") +} + +func rangesIntersectForTest(aStart, aEnd, bStart, bEnd []byte) bool { + if aEnd != nil && bytes.Compare(aEnd, bStart) <= 0 { + return false + } + if bEnd != nil && bytes.Compare(bEnd, aStart) <= 0 { + return false + } + return true +} + +func TestRouteKeyFilterTreatsNilAndEmptyEndAsInfinity(t *testing.T) { + t.Parallel() + + start := []byte("m") + for _, tc := range []struct { + name string + end []byte + }{ + {name: "nil"}, + {name: "empty", end: []byte{}}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + filter := RouteKeyFilter(start, tc.end) + require.False(t, filter([]byte("a"))) + require.True(t, filter([]byte("m"))) + require.True(t, filter([]byte("z"))) + }) + } +} + +func TestRouteKeyFilterIncludesS3BucketAuxiliaryKeys(t *testing.T) { + t.Parallel() + + filter := RouteKeyFilter( + s3keys.RouteKey("bucket-b", 7, "a"), + s3keys.RouteKey("bucket-b", 7, "z"), + ) + + require.True(t, filter(s3keys.BucketMetaKey("bucket-b"))) + require.True(t, filter(s3keys.BucketGenerationKey("bucket-b"))) + require.True(t, filter(s3keys.ObjectManifestKey("bucket-b", 7, "m"))) + require.False(t, filter(s3keys.BucketMetaKey("bucket-c"))) + require.False(t, filter(s3keys.BucketGenerationKey("bucket-c"))) +} + +func TestRouteKeyFilterIncludesS3BucketAuxiliaryRawRoute(t *testing.T) { + t.Parallel() + + filter := RouteKeyFilter([]byte("!s3|"), nil) + + require.True(t, filter(s3keys.BucketMetaKey("bucket-b"))) + require.True(t, filter(s3keys.BucketGenerationKey("bucket-b"))) +} + +func TestRouteKeyFilterForGroupUsesPartitionResolver(t *testing.T) { + t.Parallel() + + partitionedKey := []byte(sqsMsgDataPrefix + sqsPartitionMarker + "orders|partition-0|message") + resolver := &migrationFilterPartitionResolver{ + groups: map[string]uint64{string(partitionedKey): 42}, + } + + require.True(t, RouteKeyFilterForGroup(nil, nil, 42, resolver)(partitionedKey)) + require.False(t, RouteKeyFilterForGroup(nil, nil, 7, resolver)(partitionedKey)) + require.False(t, RouteKeyFilterForGroup(nil, nil, 42, resolver)( + []byte(sqsMsgDataPrefix+sqsPartitionMarker+"orders|unknown-partition"), + )) +} + +type migrationFilterPartitionResolver struct { + groups map[string]uint64 +} + +func (r *migrationFilterPartitionResolver) ResolveGroup(key []byte) (uint64, bool) { + gid, ok := r.groups[string(key)] + return gid, ok +} + +func (r *migrationFilterPartitionResolver) RecognisesPartitionedKey(key []byte) bool { + return bytes.HasPrefix(key, []byte(sqsMsgDataPrefix+sqsPartitionMarker)) +} diff --git a/kv/shard_store.go b/kv/shard_store.go index 0435707de..f4eb9e010 100644 --- a/kv/shard_store.go +++ b/kv/shard_store.go @@ -422,7 +422,7 @@ func (s *ShardStore) ScanAtWithReadFence(ctx context.Context, start []byte, end if reverse { if groupID != 0 { if routeScanBoundsPresent(routeStart, routeEnd) { - return s.scanRouteAtDirectionWithReadFence(ctx, distribution.Route{GroupID: groupID}, start, end, limit, ts, true, readRouteVersion, routeStart, routeEnd) + return s.scanRouteAtDirectionWithReadFence(ctx, distribution.Route{GroupID: groupID}, start, end, limit, ts, true, true, readRouteVersion, routeStart, routeEnd) } return nil, errors.WithStack(store.ErrNotSupported) } @@ -437,7 +437,7 @@ func (s *ShardStore) scanAtWithReadFence(ctx context.Context, start []byte, end } if groupID != 0 { - return s.scanRouteAtDirectionWithReadFence(ctx, distribution.Route{GroupID: groupID}, start, end, limit, ts, false, readRouteVersion, routeStart, routeEnd) + return s.scanRouteAtDirectionWithReadFence(ctx, distribution.Route{GroupID: groupID}, start, end, limit, ts, false, true, readRouteVersion, routeStart, routeEnd) } routes, clampToRoutes, routeVersion := s.routesForFencedScanWithVersion(start, end, routeStart, routeEnd) @@ -446,7 +446,7 @@ func (s *ShardStore) scanAtWithReadFence(ctx context.Context, start []byte, end if err != nil { return nil, err } - sort.SliceStable(out, func(i, j int) bool { + sort.Slice(out, func(i, j int) bool { return bytes.Compare(out[i].Key, out[j].Key) < 0 }) out = dedupeSortedScanResults(out) @@ -472,7 +472,7 @@ func (s *ShardStore) ScanKeysAtWithReadFence(ctx context.Context, start []byte, return nil, err } if groupID != 0 { - return s.scanKeyRouteAtWithReadFence(ctx, distribution.Route{GroupID: groupID}, start, end, limit, ts, readRouteVersion) + return s.scanKeyRouteAtWithReadFence(ctx, distribution.Route{GroupID: groupID}, start, end, limit, ts, true, readRouteVersion) } routes, clampToRoutes, routeVersion := s.routesForScanWithVersion(start, end) @@ -576,7 +576,11 @@ func (s *ShardStore) canonicalizeRedisWideColumnScanResults(ctx context.Context, } return nil, err } - out = append(out, &store.KVPair{Key: bytes.Clone(kvp.Key), Value: value}) + out = append(out, &store.KVPair{ + Key: bytes.Clone(kvp.Key), + Value: value, + RouteGroupID: kvp.RouteGroupID, + }) } return out, nil } @@ -601,6 +605,36 @@ func (s *ShardStore) routesForReverseScan(start []byte, end []byte) ([]distribut return s.routesForScan(start, end) } +func (s *ShardStore) AllowExactScanFallbackAfterPhysicalLimit(ctx context.Context, start []byte, end []byte, visibleLimit, physicalLimit int, _ uint64, _ bool) bool { + if visibleLimit <= 0 || physicalLimit <= 0 { + return false + } + g := s.exactFallbackPhysicalLimitGroup(start, end) + if g == nil { + return false + } + if _, ok := g.Store.(physicalLimitedStore); !ok { + return false + } + engine := engineForGroup(g) + return engine == nil || isLinearizableRaftLeader(ctx, engine) +} + +func (s *ShardStore) exactFallbackPhysicalLimitGroup(start []byte, end []byte) *ShardGroup { + if s == nil || s.engine == nil { + return nil + } + routes, clampToRoutes := s.routesForScan(start, end) + if len(routes) != 1 || clampToRoutes { + return nil + } + g, ok := s.groupForID(routes[0].GroupID) + if !ok || g == nil || g.Store == nil { + return nil + } + return g +} + func (s *ShardStore) routesForScan(start []byte, end []byte) ([]distribution.Route, bool) { routes, clampToRoutes, _ := s.routesForScanWithVersion(start, end) return routes, clampToRoutes @@ -611,13 +645,15 @@ func (s *ShardStore) routesForScanWithVersion(start []byte, end []byte) ([]distr routes, version := s.engine.GetIntersectingRoutesWithVersion(routeStart, routeEnd) return routes, false, version } - if routes, version, ok := s.routesForEncodedScanWithVersion(start, end); ok { + if routes, version, ok := s.routesForFilesystemUsageScanWithVersion(start, end); ok { return routes, false, version } - if routes, version, ok := s.routesForRedisWideColumnScanWithVersion(start, end); ok { + if routes, version, ok := s.routesForFilesystemChunkScanWithVersion(start, end); ok { return routes, false, version } - + if selected, ok := s.routesForInternalScanWithVersion(start, end); ok { + return selected.routes, false, selected.version + } routes, version := s.engine.GetIntersectingRoutesWithVersion(start, end) // If the scan can include internal list keys (which use a fixed prefix), // avoid clamping to shard range bounds because those keys may be ordered @@ -629,89 +665,148 @@ func (s *ShardStore) routesForScanWithVersion(start []byte, end []byte) ([]distr return routes, true, version } -func (s *ShardStore) routesForRedisWideColumnScanWithVersion(start []byte, end []byte) ([]distribution.Route, uint64, bool) { - routeStart, routeEnd, exact, ok := redisWideColumnScanRouteRange(start, end) - if !ok { - return nil, 0, false +type internalScanRouteSelection struct { + routes []distribution.Route + version uint64 +} + +func (s *ShardStore) routesForInternalScanWithVersion(start []byte, end []byte) (internalScanRouteSelection, bool) { + if selected, ok := s.routesForListAuxiliaryScanWithVersion(start, end); ok { + return selected, true } - if !exact { - routes, version := s.engine.GetIntersectingRoutesWithVersion(routeStart, routeEnd) - routes, version = s.appendRedisWideColumnLegacyScanRoutesWithVersion(routes, version, start, end) - return routes, version, true + if isBroadLegacyListDeltaScan(start) { + routes, version := s.engine.GetIntersectingRoutesWithVersion(nil, nil) + return internalScanRouteSelection{routes: routes, version: version}, true } - route, version, ok := s.engine.GetRouteWithVersion(routeStart) - if !ok { - return []distribution.Route{}, version, true + if store.ExtractLegacyListUserKeyFromDeltaScanPrefix(start) != nil { + catalogRoutes, version := s.engine.GetIntersectingRoutesWithVersion(nil, nil) + return internalScanRouteSelection{routes: routesForLegacyListDeltaScan(catalogRoutes, start, end), version: version}, true } - routes := []distribution.Route{route} - routes, version = s.appendRedisWideColumnLegacyScanRoutesWithVersion(routes, version, start, end) - return routes, version, true + if routeStart, routeEnd, exact, ok := redisWideColumnScanRouteRange(start, end); ok { + routes, version := s.redisWideColumnScanRoutesWithVersion(start, end, routeStart, routeEnd, exact) + return internalScanRouteSelection{routes: routes, version: version}, true + } + // Remaining internal collection scans route by their logical user key. + if userKey := scanRouteUserKey(start); userKey != nil { + route, version, ok := s.engine.GetRouteWithVersion(userKey) + if !ok { + return internalScanRouteSelection{routes: []distribution.Route{}, version: version}, true + } + return internalScanRouteSelection{routes: []distribution.Route{route}, version: version}, true + } + return internalScanRouteSelection{}, false } -func (s *ShardStore) appendRedisWideColumnLegacyScanRoutesWithVersion(routes []distribution.Route, version uint64, start []byte, end []byte) ([]distribution.Route, uint64) { - legacyStart, legacyEnd, ok := redisWideColumnLegacyScanRouteRange(start, end) - if !ok { - return routes, version +func (s *ShardStore) redisWideColumnScanRoutesWithVersion(start []byte, end []byte, routeStart []byte, routeEnd []byte, exact bool) ([]distribution.Route, uint64) { + var routes []distribution.Route + var version uint64 + if exact { + route, routeVersion, ok := s.engine.GetRouteWithVersion(routeStart) + version = routeVersion + if ok { + routes = append(routes, route) + } + } else { + routes, version = s.engine.GetIntersectingRoutesWithVersion(routeStart, routeEnd) } - legacyRoutes, legacyVersion := s.engine.GetIntersectingRoutesWithVersion(legacyStart, legacyEnd) + + legacyRoutes, legacyVersion := s.engine.GetIntersectingRoutesWithVersion(start, end) version = max(version, legacyVersion) - return appendDistinctRoutesByGroup(routes, legacyRoutes), version + return appendUniqueRouteGroups(routes, legacyRoutes...), version } -func appendDistinctRoutesByGroup(routes []distribution.Route, candidates []distribution.Route) []distribution.Route { - seen := make(map[uint64]struct{}, len(routes)) +func appendUniqueRouteGroups(routes []distribution.Route, extra ...distribution.Route) []distribution.Route { + seen := make(map[uint64]struct{}, len(routes)+len(extra)) + out := make([]distribution.Route, 0, len(routes)+len(extra)) for _, route := range routes { + if _, ok := seen[route.GroupID]; ok { + continue + } seen[route.GroupID] = struct{}{} + out = append(out, route) } - for _, route := range candidates { + for _, route := range extra { if _, ok := seen[route.GroupID]; ok { continue } seen[route.GroupID] = struct{}{} - routes = append(routes, route) + out = append(out, route) } - return routes + return out } -func (s *ShardStore) routesForEncodedScanWithVersion(start []byte, end []byte) ([]distribution.Route, uint64, bool) { - if routes, version, ok := s.routesForFilesystemUsageScanWithVersion(start, end); ok { - return routes, version, true +func (s *ShardStore) routesForListAuxiliaryScanWithVersion(start []byte, end []byte) (internalScanRouteSelection, bool) { + routeStart, exact, ok := listAuxiliaryScanRouteRange(start, end) + if !ok { + return internalScanRouteSelection{}, false } - if routes, version, ok := s.routesForFilesystemChunkScanWithVersion(start, end); ok { - return routes, version, true + if !exact { + routes, version := s.engine.GetIntersectingRoutesWithVersion(routeStart, nil) + return internalScanRouteSelection{routes: routes, version: version}, true } - if routeStart, exact, ok := listAuxiliaryScanRouteRange(start, end); ok { - if !exact { - routes, version := s.engine.GetIntersectingRoutesWithVersion(routeStart, nil) - return routes, version, true - } - route, version, ok := s.engine.GetRouteWithVersion(routeStart) - if !ok { - return []distribution.Route{}, version, true - } - return []distribution.Route{route}, version, true + route, version, ok := s.engine.GetRouteWithVersion(routeStart) + if !ok { + return internalScanRouteSelection{routes: []distribution.Route{}, version: version}, true } - userKey := listScanUserKey(start) - if userKey == nil { - return nil, 0, false + return internalScanRouteSelection{routes: []distribution.Route{route}, version: version}, true +} + +func routesForLegacyListDeltaScan(catalogRoutes []distribution.Route, start []byte, end []byte) []distribution.Route { + logicalUserKey := store.ExtractLegacyListUserKeyFromDeltaScanPrefix(start) + routes := make([]distribution.Route, 0) + for _, route := range catalogRoutes { + if routeContainsKey(route, logicalUserKey) { + routes = append(routes, route) + break + } } - route, version, ok := s.engine.GetRouteWithVersion(userKey) - if !ok { - return []distribution.Route{}, version, true + storedStart := store.ExtractListUserKey(start) + if storedStart != nil { + var storedEnd []byte + if len(end) > 0 { + storedEnd = store.ExtractListUserKey(end) + } + for _, route := range catalogRoutes { + if migrationRouteRangesIntersect(route.Start, route.End, storedStart, storedEnd) { + routes = append(routes, route) + } + } } - return []distribution.Route{route}, version, true + return routes } -func listScanUserKey(start []byte) []byte { - if userKey := store.ExtractListUserKeyFromDeltaScanKey(start); userKey != nil { - return userKey +func isBroadLegacyListDeltaScan(start []byte) bool { + prefix := []byte(store.LegacyListMetaDeltaPrefix) + if !bytes.HasPrefix(start, prefix) { + return false } - if userKey := store.ExtractListUserKeyFromClaimScanKey(start); userKey != nil { - return userKey + logicalUserKey := store.ExtractLegacyListUserKeyFromDeltaScanPrefix(start) + return logicalUserKey == nil || !bytes.Equal(start, store.LegacyListMetaDeltaScanPrefix(logicalUserKey)) +} + +func scanRouteUserKey(start []byte) []byte { + for _, extract := range scanRouteUserKeyExtractors { + if userKey := extract(start); userKey != nil { + return userKey + } } - // Internal list keys route by their logical user key rather than their raw - // storage prefix. - return store.ExtractListUserKey(start) + return nil +} + +var scanRouteUserKeyExtractors = []func([]byte) []byte{ + store.ExtractListUserKeyFromDeltaScanKey, + store.ExtractListUserKeyFromClaimScanKey, + store.ExtractListUserKey, + store.ExtractHashUserKeyFromField, + store.ExtractHashUserKeyFromDeltaScanPrefix, + store.ExtractSetUserKeyFromMember, + store.ExtractSetUserKeyFromDeltaScanPrefix, + store.ExtractZSetUserKeyFromMember, + store.ExtractZSetUserKeyFromScore, + store.ExtractZSetUserKeyFromScoreScanPrefix, + store.ExtractZSetUserKeyFromDeltaScanPrefix, + store.ExtractStreamUserKeyFromMeta, + store.ExtractStreamUserKeyFromEntryScanPrefix, } func (s *ShardStore) routesForFencedScanWithVersion(start []byte, end []byte, routeStart []byte, routeEnd []byte) ([]distribution.Route, bool, uint64) { @@ -752,7 +847,7 @@ func (s *ShardStore) scanRoutesAtWithReadFence(ctx context.Context, routes []dis } kvs, err := s.scanRouteAtWithOptionalFilesystemUsageOwnerFilter( - ctx, route, scanStart, scanEnd, limit, ts, false, + ctx, route, scanStart, scanEnd, limit, ts, false, !clampToRoutes, readRouteVersion, routeStart, routeEnd, filterUsageOwners, ) if err != nil { @@ -779,6 +874,7 @@ func (s *ShardStore) scanRouteAtWithOptionalFilesystemUsageOwnerFilter( limit int, ts uint64, reverse bool, + explicitGroup bool, readRouteVersion uint64, routeStart []byte, routeEnd []byte, @@ -786,12 +882,12 @@ func (s *ShardStore) scanRouteAtWithOptionalFilesystemUsageOwnerFilter( ) ([]*store.KVPair, error) { if filterUsageOwners { return s.scanRouteAtWithFilesystemUsageOwnerFilter( - ctx, route, start, end, limit, ts, reverse, + ctx, route, start, end, limit, ts, reverse, explicitGroup, readRouteVersion, routeStart, routeEnd, ) } return s.scanRouteAtDirectionWithReadFence( - ctx, route, start, end, limit, ts, reverse, + ctx, route, start, end, limit, ts, reverse, explicitGroup, readRouteVersion, routeStart, routeEnd, ) } @@ -807,19 +903,18 @@ func (s *ShardStore) routesForFilesystemUsageScanWithVersion(start []byte, end [ } func (s *ShardStore) routesForFilesystemChunkScanWithVersion(start []byte, end []byte) ([]distribution.Route, uint64, bool) { + allRoutes, version := s.engine.GetIntersectingRoutesWithVersion(nil, nil) if routeStart, routeEnd, ok := fskeys.ChunkScanRouteBounds(start, end); ok { - allRoutes, version := s.engine.GetIntersectingRoutesWithVersion(nil, nil) return intersectingRoutes(allRoutes, routeStart, routeEnd), version, true } chunkStart, chunkEnd, ok := filesystemChunkScanOverlap(start, end) if !ok { - return nil, 0, false + return nil, version, false } routeStart, routeEnd, ok := fskeys.ChunkScanRouteBounds(chunkStart, chunkEnd) if !ok { - return nil, 0, false + return nil, version, false } - allRoutes, version := s.engine.GetIntersectingRoutesWithVersion(nil, nil) // Raw scans can continue from the chunk keyspace into later filesystem // families, so include both raw and virtual chunk route groups rather than // narrowing the scan to chunks only. @@ -901,7 +996,7 @@ func (s *ShardStore) scanKeyRoutesAtWithReadFence(ctx context.Context, routes [] ctx, route, scanStart, scanEnd, limit, ts, readRouteVersion, ) } else { - keys, err = s.scanKeyRouteAtWithReadFence(ctx, route, scanStart, scanEnd, limit, ts, readRouteVersion) + keys, err = s.scanKeyRouteAtWithReadFence(ctx, route, scanStart, scanEnd, limit, ts, !clampToRoutes, readRouteVersion) } if err != nil { return nil, err @@ -956,6 +1051,7 @@ func (s *ShardStore) scanRouteAtWithFilesystemUsageOwnerFilter( limit int, ts uint64, reverse bool, + explicitGroup bool, readRouteVersion uint64, routeStart []byte, routeEnd []byte, @@ -969,7 +1065,7 @@ func (s *ShardStore) scanRouteAtWithFilesystemUsageOwnerFilter( cursorEnd := end for len(out) < limit { page, err := s.scanRouteAtDirectionWithReadFence( - ctx, route, cursorStart, cursorEnd, limit, ts, reverse, + ctx, route, cursorStart, cursorEnd, limit, ts, reverse, explicitGroup, readRouteVersion, routeStart, routeEnd, ) if err != nil { @@ -1017,7 +1113,7 @@ func (s *ShardStore) scanKeyRouteAtWithFilesystemUsageOwnerFilter( out := make([][]byte, 0, limit) cursor := start for len(out) < limit { - page, err := s.scanKeyRouteAtWithReadFence(ctx, route, cursor, end, limit, ts, readRouteVersion) + page, err := s.scanKeyRouteAtWithReadFence(ctx, route, cursor, end, limit, ts, true, readRouteVersion) if err != nil { return nil, err } @@ -1052,10 +1148,9 @@ func (s *ShardStore) reverseScanRoutesAtWithReadFence( seenGroups := make(map[uint64]struct{}) routeFilterPresent := routeScanBoundsPresent(routeStart, routeEnd) filterUsageOwners := !clampToRoutes && !routeFilterPresent && filesystemUsageScanOverlap(start, end) - for i := range routes { + for i := len(routes) - 1; i >= 0; i-- { route := routes[i] if clampToRoutes { - route = routes[len(routes)-1-i] kvs, done, err := s.clampedReverseScanRouteAtWithReadFence(ctx, route, start, end, limit, len(out), ts, readRouteVersion, routeStart, routeEnd) if err != nil { return nil, err @@ -1081,7 +1176,7 @@ func (s *ShardStore) reverseScanRoutesAtWithReadFence( seenGroups[route.GroupID] = struct{}{} } kvs, err := s.scanRouteAtWithOptionalFilesystemUsageOwnerFilter( - ctx, route, start, end, limit, ts, true, + ctx, route, start, end, limit, ts, true, true, readRouteVersion, routeStart, routeEnd, filterUsageOwners, ) if err != nil { @@ -1100,7 +1195,7 @@ func (s *ShardStore) scanKeyRouteAt( limit int, ts uint64, ) ([][]byte, error) { - return s.scanKeyRouteAtWithReadFence(ctx, route, start, end, limit, ts, 0) + return s.scanKeyRouteAtWithReadFence(ctx, route, start, end, limit, ts, false, 0) } func (s *ShardStore) scanKeyRouteAtWithReadFence( @@ -1110,6 +1205,7 @@ func (s *ShardStore) scanKeyRouteAtWithReadFence( end []byte, limit int, ts uint64, + explicitGroup bool, readRouteVersion uint64, ) ([][]byte, error) { g, ok := s.groupForID(route.GroupID) @@ -1125,7 +1221,8 @@ func (s *ShardStore) scanKeyRouteAtWithReadFence( return s.scanKeysRouteAtLeader(ctx, g, start, end, limit, ts) } - return s.proxyScanKeysAt(ctx, g, start, end, limit, ts, route.GroupID, readRouteVersion) + groupID := proxyScanGroupID(route, explicitGroup, readRouteVersion, nil, nil) + return s.proxyScanKeysAt(ctx, g, start, end, limit, ts, groupID, readRouteVersion) } func (s *ShardStore) scanKeysRouteLocal( @@ -1302,7 +1399,7 @@ func (s *ShardStore) clampedReverseScanRouteAtWithReadFence( scanStart := clampScanStart(start, route.Start) scanEnd := clampScanEnd(end, route.End) - kvs, err := s.scanRouteAtDirectionWithReadFence(ctx, route, scanStart, scanEnd, limit-currentLen, ts, true, readRouteVersion, routeStart, routeEnd) + kvs, err := s.scanRouteAtDirectionWithReadFence(ctx, route, scanStart, scanEnd, limit-currentLen, ts, true, false, readRouteVersion, routeStart, routeEnd) if err != nil { return nil, false, err } @@ -1318,7 +1415,7 @@ func (s *ShardStore) scanRouteAtDirection( ts uint64, reverse bool, ) ([]*store.KVPair, error) { - return s.scanRouteAtDirectionWithReadFence(ctx, route, start, end, limit, ts, reverse, 0, nil, nil) + return s.scanRouteAtDirectionWithReadFence(ctx, route, start, end, limit, ts, reverse, false, 0, nil, nil) } func (s *ShardStore) scanRouteAtDirectionWithReadFence( @@ -1329,14 +1426,15 @@ func (s *ShardStore) scanRouteAtDirectionWithReadFence( limit int, ts uint64, reverse bool, + explicitGroup bool, readRouteVersion uint64, routeStart []byte, routeEnd []byte, ) ([]*store.KVPair, error) { if routeScanBoundsPresent(routeStart, routeEnd) { - return s.scanRouteAtDirectionWithReadFenceRouteFilter(ctx, route, start, end, limit, ts, reverse, readRouteVersion, routeStart, routeEnd) + return s.scanRouteAtDirectionWithReadFenceRouteFilter(ctx, route, start, end, limit, ts, reverse, explicitGroup, readRouteVersion, routeStart, routeEnd) } - return s.scanRouteAtDirectionWithReadFenceOnce(ctx, route, start, end, limit, ts, reverse, readRouteVersion, routeStart, routeEnd) + return s.scanRouteAtDirectionWithReadFenceOnce(ctx, route, start, end, limit, ts, reverse, explicitGroup, readRouteVersion, routeStart, routeEnd) } func (s *ShardStore) scanRouteAtDirectionWithReadFenceRouteFilter( @@ -1347,6 +1445,7 @@ func (s *ShardStore) scanRouteAtDirectionWithReadFenceRouteFilter( limit int, ts uint64, reverse bool, + explicitGroup bool, readRouteVersion uint64, routeStart []byte, routeEnd []byte, @@ -1361,7 +1460,7 @@ func (s *ShardStore) scanRouteAtDirectionWithReadFenceRouteFilter( for len(out) < limit { remaining := limit - len(out) batchLimit := routeFilteredScanBatchLimit(remaining) - kvs, cursorKVs, err := s.scanRouteAtDirectionWithReadFenceRouteFilterPage(ctx, route, scanStart, scanEnd, batchLimit, remaining, ts, reverse, readRouteVersion, filterStart, filterEnd) + kvs, cursorKVs, err := s.scanRouteAtDirectionWithReadFenceRouteFilterPage(ctx, route, scanStart, scanEnd, batchLimit, remaining, ts, reverse, explicitGroup, readRouteVersion, filterStart, filterEnd) if err != nil { return nil, err } @@ -1402,6 +1501,7 @@ func (s *ShardStore) scanRouteAtDirectionWithReadFenceRouteFilterPage( visibleLimit int, ts uint64, reverse bool, + explicitGroup bool, readRouteVersion uint64, routeStart []byte, routeEnd []byte, @@ -1416,19 +1516,28 @@ func (s *ShardStore) scanRouteAtDirectionWithReadFenceRouteFilterPage( if err != nil { return nil, nil, errors.WithStack(err) } - return filterTxnInternalKVs(kvs), kvs, nil + return markScanRouteGroup(filterTxnInternalKVs(kvs), route.GroupID), markScanRouteGroup(kvs, route.GroupID), nil } if isLinearizableRaftLeader(ctx, engineForGroup(g)) { - return s.scanRouteAtLeaderRouteFilter(ctx, g, start, end, limit, visibleLimit, ts, reverse, routeStart, routeEnd) + kvs, cursorKVs, err := s.scanRouteAtLeaderRouteFilter(ctx, g, start, end, limit, visibleLimit, ts, reverse, routeStart, routeEnd) + return markScanRouteGroup(kvs, route.GroupID), markScanRouteGroup(cursorKVs, route.GroupID), err } - kvs, err := s.proxyRawScanAt(ctx, g, start, end, limit, ts, reverse, route.GroupID, readRouteVersion, routeStart, routeEnd) + groupID := proxyScanGroupID(route, explicitGroup, readRouteVersion, routeStart, routeEnd) + kvs, err := s.proxyRawScanAt(ctx, g, start, end, limit, ts, reverse, groupID, readRouteVersion, routeStart, routeEnd) if err != nil { return nil, nil, err } filtered := filterTxnInternalKVs(kvs) - return filtered, kvs, nil + return markScanRouteGroup(filtered, route.GroupID), markScanRouteGroup(kvs, route.GroupID), nil +} + +func proxyScanGroupID(route distribution.Route, explicitGroup bool, readRouteVersion uint64, routeStart []byte, routeEnd []byte) uint64 { + if explicitGroup || readRouteVersion == 0 || routeScanBoundsPresent(routeStart, routeEnd) { + return route.GroupID + } + return 0 } func (s *ShardStore) scanRouteAtDirectionWithReadFenceOnce( @@ -1439,6 +1548,7 @@ func (s *ShardStore) scanRouteAtDirectionWithReadFenceOnce( limit int, ts uint64, reverse bool, + explicitGroup bool, readRouteVersion uint64, routeStart []byte, routeEnd []byte, @@ -1449,7 +1559,7 @@ func (s *ShardStore) scanRouteAtDirectionWithReadFenceOnce( } if !reverse { - return s.scanRouteAtForward(ctx, route, g, start, end, limit, ts, readRouteVersion, routeStart, routeEnd) + return s.scanRouteAtForward(ctx, route, g, start, end, limit, ts, explicitGroup, readRouteVersion, routeStart, routeEnd) } if engineForGroup(g) == nil { @@ -1457,20 +1567,22 @@ func (s *ShardStore) scanRouteAtDirectionWithReadFenceOnce( if err != nil { return nil, errors.WithStack(err) } - return filterTxnInternalKVs(kvs), nil + return markScanRouteGroup(filterTxnInternalKVs(kvs), route.GroupID), nil } if isLinearizableRaftLeader(ctx, engineForGroup(g)) { - return s.scanRouteAtLeader(ctx, g, start, end, limit, ts, reverse) + kvs, err := s.scanRouteAtLeader(ctx, g, start, end, limit, ts, reverse) + return markScanRouteGroup(kvs, route.GroupID), err } - kvs, err := s.proxyRawScanAt(ctx, g, start, end, limit, ts, reverse, route.GroupID, readRouteVersion, routeStart, routeEnd) + groupID := proxyScanGroupID(route, explicitGroup, readRouteVersion, routeStart, routeEnd) + kvs, err := s.proxyRawScanAt(ctx, g, start, end, limit, ts, reverse, groupID, readRouteVersion, routeStart, routeEnd) if err != nil { return nil, err } // The leader's RawScanAt is expected to perform lock resolution and filtering // via ShardStore.ScanAt, so avoid N+1 proxy gets here. - return filterTxnInternalKVs(kvs), nil + return markScanRouteGroup(filterTxnInternalKVs(kvs), route.GroupID), nil } const routeFilteredScanBatchMin = 128 @@ -1579,6 +1691,7 @@ func (s *ShardStore) scanRouteAtForward( end []byte, limit int, ts uint64, + explicitGroup bool, readRouteVersion uint64, routeStart []byte, routeEnd []byte, @@ -1590,7 +1703,7 @@ func (s *ShardStore) scanRouteAtForward( out := make([]*store.KVPair, 0, limit) cursor := start for len(out) < limit { - page, err := s.scanRouteAtForwardPage(ctx, route, g, cursor, end, limit, ts, readRouteVersion, routeStart, routeEnd) + page, err := s.scanRouteAtForwardPage(ctx, route, g, cursor, end, limit, ts, explicitGroup, readRouteVersion, routeStart, routeEnd) if err != nil { return nil, err } @@ -1614,7 +1727,7 @@ func (s *ShardStore) scanRouteAtForward( if len(out) > limit { out = out[:limit] } - return out, nil + return markScanRouteGroup(out, route.GroupID), nil } func (s *ShardStore) scanRouteAtForwardPage( @@ -1625,6 +1738,7 @@ func (s *ShardStore) scanRouteAtForwardPage( end []byte, limit int, ts uint64, + explicitGroup bool, readRouteVersion uint64, routeStart []byte, routeEnd []byte, @@ -1663,7 +1777,8 @@ func (s *ShardStore) scanRouteAtForwardPage( }, nil } - raw, err := s.proxyRawScanAt(ctx, g, start, end, limit, ts, false, route.GroupID, readRouteVersion, routeStart, routeEnd) + groupID := proxyScanGroupID(route, explicitGroup, readRouteVersion, routeStart, routeEnd) + raw, err := s.proxyRawScanAt(ctx, g, start, end, limit, ts, false, groupID, readRouteVersion, routeStart, routeEnd) if err != nil { return scanRoutePage{}, err } @@ -1699,11 +1814,12 @@ func (s *ShardStore) scanRouteAtDirectionPhysicalLimit( if err != nil { return nil, limitReached, errors.WithStack(err) } - return filterTxnInternalKVs(kvs), limitReached, nil + return markScanRouteGroup(filterTxnInternalKVs(kvs), route.GroupID), limitReached, nil } if isLinearizableRaftLeader(ctx, engineForGroup(g)) { - return s.scanRouteAtLeaderPhysicalLimit(ctx, g, start, end, visibleLimit, physicalLimit, ts, reverse) + kvs, limitReached, err := s.scanRouteAtLeaderPhysicalLimit(ctx, g, start, end, visibleLimit, physicalLimit, ts, reverse) + return markScanRouteGroup(kvs, route.GroupID), limitReached, err } // RawScanAt cannot enforce physicalLimit, so report truncation and let @@ -2139,18 +2255,18 @@ func (s *ShardStore) LatestCommitTSWithReadFence(ctx context.Context, key []byte return 0, false, nil } var latest uint64 - found := false + var exists bool for _, route := range routes { - ts, exists, err := s.latestCommitTSForRoute(ctx, route, key, readRouteVersion) + ts, ok, err := s.latestCommitTSForRoute(ctx, route, key, readRouteVersion) if err != nil { return 0, false, err } - if exists && (!found || ts > latest) { + if ok && (!exists || ts > latest) { latest = ts - found = true + exists = true } } - return latest, found, nil + return latest, exists, nil } func (s *ShardStore) latestCommitTSForRoute(ctx context.Context, route distribution.Route, key []byte, readRouteVersion uint64) (uint64, bool, error) { @@ -2660,6 +2776,18 @@ func filterTxnInternalKVs(kvs []*store.KVPair) []*store.KVPair { return out } +func markScanRouteGroup(kvs []*store.KVPair, groupID uint64) []*store.KVPair { + if groupID == 0 { + return kvs + } + for _, kvp := range kvs { + if kvp != nil { + kvp.RouteGroupID = groupID + } + } + return kvs +} + type txnStatus int const ( @@ -3135,6 +3263,18 @@ func (s *ShardStore) Snapshot() (store.Snapshot, error) { return nil, store.ErrNotSupported } +func (s *ShardStore) ExportVersions(context.Context, store.ExportVersionsOptions) (store.ExportVersionsResult, error) { + return store.ExportVersionsResult{}, store.ErrNotSupported +} + +func (s *ShardStore) ImportVersions(context.Context, store.ImportVersionsOptions) (store.ImportVersionsResult, error) { + return store.ImportVersionsResult{}, store.ErrNotSupported +} + +func (s *ShardStore) MigrationHLCFloor(context.Context, uint64) (uint64, error) { + return 0, store.ErrNotSupported +} + func (s *ShardStore) Restore(_ io.Reader) error { return store.ErrNotSupported } @@ -3192,21 +3332,25 @@ func (s *ShardStore) LocalStoreForKey(key []byte) (store.MVCCStore, bool) { return g.Store, true } -// LocalStores returns every process-local shard store in stable group order. -// It is used by node-local auxiliary maintenance that must recover state after -// snapshot restore without leader routing. +// LocalStores returns every local shard store in stable group-ID order. It is +// reserved for node-local maintenance workers that need to scan auxiliary +// state present on any local shard, not for replicated reads. func (s *ShardStore) LocalStores() []store.MVCCStore { + if s == nil { + return nil + } groupIDs := make([]uint64, 0, len(s.groups)) for groupID := range s.groups { groupIDs = append(groupIDs, groupID) } - sort.Slice(groupIDs, func(i, j int) bool { return groupIDs[i] < groupIDs[j] }) + slices.Sort(groupIDs) stores := make([]store.MVCCStore, 0, len(groupIDs)) for _, groupID := range groupIDs { - group := s.groups[groupID] - if group != nil && group.Store != nil { - stores = append(stores, group.Store) + g := s.groups[groupID] + if g == nil || g.Store == nil { + continue } + stores = append(stores, g.Store) } return stores } diff --git a/kv/shard_store_test.go b/kv/shard_store_test.go index e14c8f66a..04e6aa031 100644 --- a/kv/shard_store_test.go +++ b/kv/shard_store_test.go @@ -58,6 +58,20 @@ func (e *followerProxyEngine) Close() error { return nil } +func TestShardStoreLocalStoresReturnsLocalShardStoresInStableOrder(t *testing.T) { + t.Parallel() + + first := store.NewMVCCStore() + second := store.NewMVCCStore() + st := NewShardStore(distribution.NewEngine(), map[uint64]*ShardGroup{ + 3: {Store: second}, + 1: {Store: first}, + 2: {}, + }) + + require.Equal(t, []store.MVCCStore{first, second}, st.LocalStores()) +} + func TestShardStoreScanAt_IncludesListKeysAcrossShards(t *testing.T) { t.Parallel() @@ -119,6 +133,101 @@ func TestShardStoreScanAt_RoutesListItemScansByUserKey(t *testing.T) { require.Equal(t, k2, kvs[2].Key) } +func TestShardStoreScanAt_RoutesListDeltaScansByUserKey(t *testing.T) { + t.Parallel() + + ctx := context.Background() + userKey := []byte("x") // routes to group 2; raw !lst|* prefixes route to group 1. + for _, tc := range []struct { + name string + key []byte + scanStart []byte + legacyRouting bool + }{ + {name: "current", key: store.ListMetaDeltaKey(userKey, 10, 1), scanStart: store.ListMetaDeltaScanPrefix(userKey)}, + {name: "legacy", key: legacyListMetaDeltaKey(userKey, 10, 1), scanStart: store.LegacyListMetaDeltaScanPrefix(userKey), legacyRouting: true}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + st := newTwoRouteShardStoreForScanTest() + deltaValue := store.MarshalListMetaDelta(store.ListMetaDelta{LenDelta: 1}) + if tc.legacyRouting { + require.NoError(t, st.groups[1].Store.PutAt(ctx, tc.key, deltaValue, 1, 0)) + } else { + require.NoError(t, st.PutAt(ctx, tc.key, deltaValue, 1, 0)) + } + + kvs, err := st.ScanAt(ctx, tc.scanStart, store.PrefixScanEnd(tc.scanStart), 10, ^uint64(0)) + require.NoError(t, err) + require.Len(t, kvs, 1) + require.Equal(t, tc.key, kvs[0].Key) + }) + } +} + +func TestShardStoreScanAt_BroadLegacyListDeltaScansAllRoutes(t *testing.T) { + t.Parallel() + + ctx := context.Background() + st := newTwoRouteShardStoreForScanTest() + deltaValue := store.MarshalListMetaDelta(store.ListMetaDelta{LenDelta: 1}) + leftKey := legacyListMetaDeltaKey([]byte("left-list"), 10, 1) + rightKey := legacyListMetaDeltaKey([]byte("right-list"), 11, 1) + require.NoError(t, st.groups[1].Store.PutAt(ctx, leftKey, deltaValue, 1, 0)) + require.NoError(t, st.groups[2].Store.PutAt(ctx, rightKey, deltaValue, 1, 0)) + + kvs, err := st.ScanAt(ctx, []byte(store.LegacyListMetaDeltaPrefix), store.PrefixScanEnd([]byte(store.LegacyListMetaDeltaPrefix)), 10, ^uint64(0)) + require.NoError(t, err) + require.Len(t, kvs, 2) + require.Equal(t, leftKey, kvs[0].Key) + require.Equal(t, uint64(1), kvs[0].RouteGroupID) + require.Equal(t, rightKey, kvs[1].Key) + require.Equal(t, uint64(2), kvs[1].RouteGroupID) +} + +func TestShardStoreScanAt_RoutesWideColumnScansByUserKey(t *testing.T) { + t.Parallel() + + ctx := context.Background() + for _, tc := range []struct { + name string + key []byte + scanStart []byte + }{ + {name: "hash field", key: store.HashFieldKey([]byte("x"), []byte("f")), scanStart: store.HashFieldScanPrefix([]byte("x"))}, + {name: "hash delta", key: store.HashMetaDeltaKey([]byte("x"), 10, 0), scanStart: store.HashMetaDeltaScanPrefix([]byte("x"))}, + {name: "set member", key: store.SetMemberKey([]byte("x"), []byte("m")), scanStart: store.SetMemberScanPrefix([]byte("x"))}, + {name: "set delta", key: store.SetMetaDeltaKey([]byte("x"), 10, 0), scanStart: store.SetMetaDeltaScanPrefix([]byte("x"))}, + {name: "zset member", key: store.ZSetMemberKey([]byte("x"), []byte("m")), scanStart: store.ZSetMemberScanPrefix([]byte("x"))}, + {name: "zset score", key: store.ZSetScoreKey([]byte("x"), 1.5, []byte("m")), scanStart: store.ZSetScoreScanPrefix([]byte("x"))}, + {name: "zset delta", key: store.ZSetMetaDeltaKey([]byte("x"), 10, 0), scanStart: store.ZSetMetaDeltaScanPrefix([]byte("x"))}, + {name: "stream meta", key: store.StreamMetaKey([]byte("x")), scanStart: store.StreamMetaKey([]byte("x"))}, + {name: "stream entry", key: store.StreamEntryKey([]byte("x"), 10, 0), scanStart: store.StreamEntryScanPrefix([]byte("x"))}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + st := newTwoRouteShardStoreForScanTest() + require.NoError(t, st.PutAt(ctx, tc.key, []byte("v"), 1, 0)) + + kvs, err := st.ScanAt(ctx, tc.scanStart, store.PrefixScanEnd(tc.scanStart), 10, ^uint64(0)) + require.NoError(t, err) + require.Len(t, kvs, 1) + require.Equal(t, tc.key, kvs[0].Key) + }) + } +} + +func newTwoRouteShardStoreForScanTest() *ShardStore { + engine := distribution.NewEngine() + engine.UpdateRoute([]byte(""), []byte("m"), 1) + engine.UpdateRoute([]byte("m"), nil, 2) + groups := map[uint64]*ShardGroup{ + 1: {Store: store.NewMVCCStore()}, + 2: {Store: store.NewMVCCStore()}, + } + return NewShardStore(engine, groups) +} + func TestShardStoreScanAtWithReadFence_RoutesListAuxiliaryScansByUserKey(t *testing.T) { t.Parallel() @@ -160,35 +269,6 @@ func TestShardStoreScanAtWithReadFence_RoutesListAuxiliaryScansByUserKey(t *test } } -func TestShardStoreScanAt_RoutesBareListAuxiliaryScansAcrossShards(t *testing.T) { - t.Parallel() - - ctx := context.Background() - engine := distribution.NewEngine() - engine.UpdateRoute([]byte(""), []byte("m"), 1) - engine.UpdateRoute([]byte("m"), nil, 2) - groups := map[uint64]*ShardGroup{ - 1: {Store: store.NewMVCCStore()}, - 2: {Store: store.NewMVCCStore()}, - } - t.Cleanup(func() { - _ = groups[1].Store.Close() - _ = groups[2].Store.Close() - }) - st := NewShardStore(engine, groups) - - left := store.ListMetaDeltaKey([]byte("anna"), 10, 0) - right := store.ListMetaDeltaKey([]byte("zoey"), 11, 0) - require.NoError(t, groups[1].Store.PutAt(ctx, left, []byte("left"), 10, 0)) - require.NoError(t, groups[2].Store.PutAt(ctx, right, []byte("right"), 11, 0)) - - prefix := []byte(store.ListMetaDeltaPrefix) - kvs, err := st.ScanAt(ctx, prefix, prefixScanEnd(prefix), 10, ^uint64(0)) - require.NoError(t, err) - require.Len(t, kvs, 2) - require.Equal(t, [][]byte{left, right}, [][]byte{kvs[0].Key, kvs[1].Key}) -} - func TestShardStoreScanGroupAt_UsesExplicitGroup(t *testing.T) { t.Parallel() @@ -234,76 +314,6 @@ func TestShardStoreGetGroupAt_UsesExplicitGroup(t *testing.T) { require.ErrorIs(t, err, store.ErrKeyNotFound) } -func TestShardStoreWritePathsRejectRouteWriteTimestampFloor(t *testing.T) { - t.Parallel() - - ctx := context.Background() - engine := distribution.NewEngine() - require.NoError(t, engine.ApplySnapshot(distribution.CatalogSnapshot{ - Version: 1, - Routes: []distribution.RouteDescriptor{ - { - RouteID: 1, - Start: []byte(""), - End: nil, - GroupID: 1, - State: distribution.RouteStateActive, - MinWriteTSExclusive: 10, - }, - }, - })) - groups := map[uint64]*ShardGroup{ - 1: {Store: store.NewMVCCStore()}, - } - t.Cleanup(func() { - _ = groups[1].Store.Close() - }) - st := NewShardStore(engine, groups) - - require.ErrorIs(t, st.PutAt(ctx, []byte("put-stale"), []byte("v"), 10, 0), store.ErrWriteConflict) - require.NoError(t, st.PutAt(ctx, []byte("put-fresh"), []byte("v"), 11, 0)) - - require.ErrorIs(t, st.DeleteAt(ctx, []byte("delete-stale"), 10), store.ErrWriteConflict) - require.NoError(t, st.DeleteAt(ctx, []byte("delete-fresh"), 11)) - - require.ErrorIs(t, st.PutWithTTLAt(ctx, []byte("ttl-stale"), []byte("v"), 10, 99), store.ErrWriteConflict) - require.NoError(t, st.PutWithTTLAt(ctx, []byte("ttl-fresh"), []byte("v"), 11, 99)) - - require.ErrorIs(t, st.ExpireAt(ctx, []byte("expire-stale"), 99, 10), store.ErrWriteConflict) - require.NoError(t, st.PutAt(ctx, []byte("expire-fresh"), []byte("v"), 11, 0)) - require.NoError(t, st.ExpireAt(ctx, []byte("expire-fresh"), 99, 12)) - - require.ErrorIs(t, st.ApplyMutations(ctx, []*store.KVPairMutation{ - {Op: store.OpTypePut, Key: []byte("apply-stale"), Value: []byte("v")}, - }, nil, 0, 10), store.ErrWriteConflict) - require.NoError(t, st.ApplyMutations(ctx, []*store.KVPairMutation{ - {Op: store.OpTypePut, Key: []byte("apply-fresh"), Value: []byte("v")}, - }, nil, 0, 11)) - - require.ErrorIs(t, st.ApplyMutationsRaft(ctx, []*store.KVPairMutation{ - {Op: store.OpTypePut, Key: []byte("raft-stale"), Value: []byte("v")}, - }, nil, 0, 10), store.ErrWriteConflict) - require.NoError(t, st.ApplyMutationsRaft(ctx, []*store.KVPairMutation{ - {Op: store.OpTypePut, Key: []byte("raft-fresh"), Value: []byte("v")}, - }, nil, 0, 11)) - - require.ErrorIs(t, st.ApplyMutationsRaftAt(ctx, []*store.KVPairMutation{ - {Op: store.OpTypePut, Key: []byte("raft-at-stale"), Value: []byte("v")}, - }, nil, 0, 10, 1), store.ErrWriteConflict) - require.NoError(t, st.ApplyMutationsRaftAt(ctx, []*store.KVPairMutation{ - {Op: store.OpTypePut, Key: []byte("raft-at-fresh"), Value: []byte("v")}, - }, nil, 0, 11, 2)) - - require.ErrorIs(t, st.DeletePrefixAt(ctx, []byte("prefix-stale"), nil, 10), store.ErrWriteConflict) - require.NoError(t, st.DeletePrefixAt(ctx, []byte("prefix-fresh"), nil, 11)) - - require.ErrorIs(t, st.DeletePrefixAtRaft(ctx, []byte("raft-prefix-stale"), nil, 10), store.ErrWriteConflict) - require.NoError(t, st.DeletePrefixAtRaft(ctx, []byte("raft-prefix-fresh"), nil, 11)) - - require.ErrorIs(t, st.DeletePrefixAtRaftAt(ctx, []byte("raft-at-prefix-stale"), nil, 10, 3), store.ErrWriteConflict) - require.NoError(t, st.DeletePrefixAtRaftAt(ctx, []byte("raft-at-prefix-fresh"), nil, 11, 4)) -} - func TestShardStore_ForwardsReadFenceStamps(t *testing.T) { t.Parallel() @@ -358,7 +368,7 @@ func TestShardStore_ForwardsReadFenceStamps(t *testing.T) { require.NoError(t, err) fake.mu.Lock() - require.Equal(t, uint64(1), fake.lastScanReq.GetGroupId()) + require.Equal(t, uint64(0), fake.lastScanReq.GetGroupId()) require.Equal(t, uint64(100), fake.lastScanReq.GetReadRouteVersion()) fake.mu.Unlock() @@ -366,7 +376,7 @@ func TestShardStore_ForwardsReadFenceStamps(t *testing.T) { require.NoError(t, err) fake.mu.Lock() - require.Equal(t, uint64(1), fake.lastScanReq.GetGroupId()) + require.Equal(t, uint64(0), fake.lastScanReq.GetGroupId()) require.Equal(t, uint64(100), fake.lastScanReq.GetReadRouteVersion()) require.True(t, fake.lastScanReq.GetKeysOnly()) fake.mu.Unlock() @@ -384,7 +394,7 @@ func TestShardStore_ForwardsReadFenceStamps(t *testing.T) { require.NoError(t, err) fake.mu.Lock() - require.Equal(t, uint64(1), fake.lastScanReq.GetGroupId()) + require.Equal(t, uint64(0), fake.lastScanReq.GetGroupId()) require.Equal(t, uint64(100), fake.lastScanReq.GetReadRouteVersion()) require.False(t, fake.lastScanReq.GetRouteBoundsPresent()) fake.mu.Unlock() @@ -398,7 +408,7 @@ func TestShardStore_ForwardsReadFenceStamps(t *testing.T) { require.Equal(t, uint64(100), fake.lastScanReq.GetReadRouteVersion()) } -func TestShardStoreRoutesForScanUsesWideColumnUserKey(t *testing.T) { +func TestShardStoreRoutesForScanUsesWideColumnUserKeyAndLegacyRoute(t *testing.T) { t.Parallel() engine := distribution.NewEngine() @@ -423,8 +433,7 @@ func TestShardStoreRoutesForScanUsesWideColumnUserKey(t *testing.T) { routes, clamp := st.routesForScan(tc.prefix, prefixScanEnd(tc.prefix)) require.False(t, clamp) require.Len(t, routes, 2) - require.Equal(t, uint64(2), routes[0].GroupID) - require.Equal(t, uint64(1), routes[1].GroupID) + require.ElementsMatch(t, []uint64{1, 2}, []uint64{routes[0].GroupID, routes[1].GroupID}) }) } } @@ -460,7 +469,7 @@ func TestShardStoreScanAtRoutesWideColumnPrefixesByUserKey(t *testing.T) { require.NoError(t, st.PutAt(ctx, tc.key, []byte("value"), 20, 0)) kvs, err := st.ScanAt(ctx, tc.prefix, prefixScanEnd(tc.prefix), 10, 20) require.NoError(t, err) - require.Equal(t, []*store.KVPair{{Key: tc.key, Value: []byte("value")}}, kvs) + require.Equal(t, []*store.KVPair{{Key: tc.key, Value: []byte("value"), RouteGroupID: 2}}, kvs) _, err = groups[1].Store.GetAt(ctx, tc.key, 20) require.ErrorIs(t, err, store.ErrKeyNotFound) }) @@ -707,64 +716,6 @@ func TestShardStoreScanAtWithReadFence_FiltersSuppliedBoundsByRouteKey(t *testin require.Equal(t, left, kvs[0].Key) } -func TestShardStoreScanAtWithReadFence_FiltersRedisAuxiliaryBoundsByRouteKey(t *testing.T) { - t.Parallel() - - ctx := context.Background() - - engine := distribution.NewEngine() - engine.UpdateRoute([]byte(""), []byte("m"), 1) - engine.UpdateRoute([]byte("m"), nil, 1) - - groups := map[uint64]*ShardGroup{ - 1: {Store: store.NewMVCCStore()}, - } - t.Cleanup(func() { _ = groups[1].Store.Close() }) - st := NewShardStore(engine, groups) - - for _, tc := range []struct { - name string - prefix []byte - left []byte - right []byte - }{ - { - name: "list delta", - prefix: []byte(store.ListMetaDeltaPrefix), - left: store.ListMetaDeltaKey([]byte("alpha"), 10, 0), - right: store.ListMetaDeltaKey([]byte("zulu"), 11, 0), - }, - { - name: "list claim", - prefix: []byte(store.ListClaimPrefix), - left: store.ListClaimKey([]byte("alpha"), 1), - right: store.ListClaimKey([]byte("zulu"), 1), - }, - { - name: "stream meta", - prefix: []byte(store.StreamMetaPrefix), - left: store.StreamMetaKey([]byte("alpha")), - right: store.StreamMetaKey([]byte("zulu")), - }, - { - name: "stream entry", - prefix: []byte(store.StreamEntryPrefix), - left: store.StreamEntryKey([]byte("alpha"), 1, 0), - right: store.StreamEntryKey([]byte("zulu"), 1, 0), - }, - } { - t.Run(tc.name, func(t *testing.T) { - require.NoError(t, groups[1].Store.PutAt(ctx, tc.left, []byte("left"), 1, 0)) - require.NoError(t, groups[1].Store.PutAt(ctx, tc.right, []byte("right"), 2, 0)) - - kvs, err := st.ScanAtWithReadFence(ctx, tc.prefix, prefixScanEnd(tc.prefix), 1, 2, false, 0, st.ReadRouteVersion(), []byte("m"), nil) - require.NoError(t, err) - require.Len(t, kvs, 1) - require.Equal(t, tc.right, kvs[0].Key) - }) - } -} - func TestShardStoreScanAtWithReadFence_FiltersByEachRouteBounds(t *testing.T) { t.Parallel() @@ -1192,7 +1143,7 @@ func TestShardStoreProxyForwardPageAdvancesFromRawPage(t *testing.T) { } st := NewShardStore(distribution.NewEngine(), map[uint64]*ShardGroup{42: g}) - page, err := st.scanRouteAtForwardPage(ctx, distribution.Route{GroupID: 42}, g, []byte(""), nil, 2, ^uint64(0), 0, nil, nil) + page, err := st.scanRouteAtForwardPage(ctx, distribution.Route{GroupID: 42}, g, []byte(""), nil, 2, ^uint64(0), true, 0, nil, nil) require.NoError(t, err) require.True(t, page.full) require.Equal(t, internalKey, page.advanceKey) @@ -1654,31 +1605,7 @@ func TestShardStoreScanAt_RoutesExactRedisWideColumnScanToOneShard(t *testing.T) require.Equal(t, uint64(1), routes[0].GroupID) } -func TestShardStoreRoutesForWideColumnBoundedPatternIncludesLegacyRawRoute(t *testing.T) { - t.Parallel() - - engine := distribution.NewEngine() - engine.UpdateRoute([]byte(""), []byte("m"), 1) - engine.UpdateRoute([]byte("m"), nil, 2) - groups := map[uint64]*ShardGroup{ - 1: {Store: store.NewMVCCStore()}, - 2: {Store: store.NewMVCCStore()}, - } - t.Cleanup(func() { - _ = groups[1].Store.Close() - _ = groups[2].Store.Close() - }) - st := NewShardStore(engine, groups) - - start := store.HashFieldScanPrefix([]byte("m")) - routes, clamp, _ := st.routesForScanWithVersion(start, prefixScanEnd([]byte(store.HashFieldPrefix))) - require.False(t, clamp) - require.Len(t, routes, 2) - require.Equal(t, uint64(2), routes[0].GroupID) - require.Equal(t, uint64(1), routes[1].GroupID) -} - -func TestShardStoreRedisWideColumnReadsLegacyRawRoute(t *testing.T) { +func TestShardStoreScanAt_RoutesExactRedisWideColumnScanIncludesLegacyRoute(t *testing.T) { t.Parallel() ctx := context.Background() @@ -1696,69 +1623,17 @@ func TestShardStoreRedisWideColumnReadsLegacyRawRoute(t *testing.T) { st := NewShardStore(engine, groups) key := store.HashFieldKey([]byte("zulu"), []byte("field")) - require.NoError(t, groups[1].Store.PutAt(ctx, key, []byte("legacy"), 5, 0)) - - value, err := st.GetAt(ctx, key, 5) - require.NoError(t, err) - require.Equal(t, []byte("legacy"), value) - - ts, exists, err := st.LatestCommitTS(ctx, key) - require.NoError(t, err) - require.True(t, exists) - require.Equal(t, uint64(5), ts) + require.NoError(t, groups[1].Store.PutAt(ctx, key, []byte("legacy"), 10, 0)) - prefix := store.HashFieldScanPrefix([]byte("zulu")) - kvs, err := st.ScanAt(ctx, prefix, prefixScanEnd(prefix), 10, 5) + start := store.HashFieldScanPrefix([]byte("zulu")) + kvs, err := st.ScanAt(ctx, start, prefixScanEnd(start), 10, ^uint64(0)) require.NoError(t, err) require.Len(t, kvs, 1) + require.Equal(t, key, kvs[0].Key) require.Equal(t, []byte("legacy"), kvs[0].Value) - - require.NoError(t, st.PutAt(ctx, key, []byte("current"), 6, 0)) - value, err = st.GetAt(ctx, key, 6) - require.NoError(t, err) - require.Equal(t, []byte("current"), value) - - ts, exists, err = st.LatestCommitTS(ctx, key) - require.NoError(t, err) - require.True(t, exists) - require.Equal(t, uint64(6), ts) - - kvs, err = st.ScanAt(ctx, prefix, prefixScanEnd(prefix), 10, 6) - require.NoError(t, err) - require.Len(t, kvs, 1) - require.Equal(t, []byte("current"), kvs[0].Value) - - kvs, err = st.ReverseScanAt(ctx, prefix, prefixScanEnd(prefix), 10, 6) - require.NoError(t, err) - require.Len(t, kvs, 1) - require.Equal(t, []byte("current"), kvs[0].Value) - - require.NoError(t, st.DeleteAt(ctx, key, 7)) - _, err = st.GetAt(ctx, key, 7) - require.ErrorIs(t, err, store.ErrKeyNotFound) - - kvs, err = st.ScanAt(ctx, prefix, prefixScanEnd(prefix), 10, 7) - require.NoError(t, err) - require.Empty(t, kvs) - - kvs, err = st.ReverseScanAt(ctx, prefix, prefixScanEnd(prefix), 10, 7) - require.NoError(t, err) - require.Empty(t, kvs) - - require.NoError(t, st.PutAt(ctx, key, []byte("future"), 9, 0)) - _, err = st.GetAt(ctx, key, 8) - require.ErrorIs(t, err, store.ErrKeyNotFound) - - kvs, err = st.ScanAt(ctx, prefix, prefixScanEnd(prefix), 10, 8) - require.NoError(t, err) - require.Empty(t, kvs) - - value, err = st.GetAt(ctx, key, 9) - require.NoError(t, err) - require.Equal(t, []byte("future"), value) } -func TestShardStoreReverseRedisWideColumnScanPrefersLogicalRoute(t *testing.T) { +func TestShardStoreLatestCommitTS_IncludesLegacyRedisWideColumnRoute(t *testing.T) { t.Parallel() ctx := context.Background() @@ -1776,14 +1651,13 @@ func TestShardStoreReverseRedisWideColumnScanPrefersLogicalRoute(t *testing.T) { st := NewShardStore(engine, groups) key := store.HashFieldKey([]byte("zulu"), []byte("field")) - require.NoError(t, groups[1].Store.PutAt(ctx, key, []byte("legacy"), 5, 0)) - require.NoError(t, st.PutAt(ctx, key, []byte("current"), 6, 0)) + require.NoError(t, groups[2].Store.PutAt(ctx, key, []byte("normalized"), 10, 0)) + require.NoError(t, groups[1].Store.PutAt(ctx, key, []byte("legacy"), 20, 0)) - prefix := store.HashFieldScanPrefix([]byte("zulu")) - kvs, err := st.ReverseScanAt(ctx, prefix, prefixScanEnd(prefix), 10, 6) + latest, ok, err := st.LatestCommitTS(ctx, key) require.NoError(t, err) - require.Len(t, kvs, 1) - require.Equal(t, []byte("current"), kvs[0].Value) + require.True(t, ok) + require.Equal(t, uint64(20), latest) } func TestShardStoreScanAt_RoutesFilesystemChunkScansByChunkRouteKey(t *testing.T) { diff --git a/kv/sharded_coordinator.go b/kv/sharded_coordinator.go index 8c6a81ad4..f7f80d024 100644 --- a/kv/sharded_coordinator.go +++ b/kv/sharded_coordinator.go @@ -734,13 +734,6 @@ func (c *ShardedCoordinator) Dispatch(ctx context.Context, reqs *OperationGroup[ return nil, err } - // DEL_PREFIX cannot be routed to a single shard because the prefix may - // span multiple shards (or be nil, meaning "all keys"). Broadcast the - // operation to every shard group so each FSM scans locally. - if hasDelPrefixElem(reqs.Elems) { - return c.dispatchDelPrefixBroadcast(ctx, reqs.IsTxn, reqs.Elems) - } - // Capture whether the caller supplied a non-zero StartTS BEFORE // the coordinator-allocates-on-zero branch below mutates the // field. A caller-supplied StartTS names a specific snapshot @@ -772,10 +765,27 @@ func (c *ShardedCoordinator) Dispatch(ctx context.Context, reqs *OperationGroup[ } if reqs.IsTxn { + if resp, handled, err := c.dispatchBeforeShardRouting(ctx, reqs); handled { + return resp, err + } return c.dispatchTxnWithComposed1Retry(ctx, reqs, callerSuppliedStartTS) } - return c.dispatchNonTxn(ctx, reqs) + return c.dispatchRawWithComposed1Retry(ctx, reqs) +} + +func (c *ShardedCoordinator) dispatchBeforeShardRouting(ctx context.Context, reqs *OperationGroup[OP]) (*CoordinateResponse, bool, error) { + // DEL_PREFIX cannot be routed to a single shard because the prefix may + // span multiple shards (or be nil, meaning "all keys"). Broadcast the + // operation to every shard group so each FSM scans locally. + if hasDelPrefixElem(reqs.Elems) { + resp, err := c.dispatchDelPrefixBroadcast(ctx, reqs.IsTxn, reqs.Elems, reqs.ObservedRouteVersion) + return resp, true, err + } + if err := c.rejectWriteFencedPointElems(reqs.Elems); err != nil { + return nil, true, err + } + return nil, false, nil } // dispatchTxnWithComposed1Retry runs the M4 Composed-1 retry loop @@ -837,16 +847,17 @@ func (c *ShardedCoordinator) Dispatch(ctx context.Context, reqs *OperationGroup[ // gate spuriously rejects resolver-routed commits even when // the resolver picked the correct gid. Skip the auto-pin // for any resolver-recognised key; the request flows with -// ObservedRouteVersion=0 and the M3 gate short-circuits — -// restoring the pre-auto-pin behaviour for resolver-routed -// txns. Resolver-aware M3 is M5+ work (codex P1 on -// 6a458a28, PR #900). +// ObservedRouteVersion=0 and the Composed-1 owner gate +// short-circuits — restoring the pre-auto-pin behaviour for +// resolver-routed txns. Resolver-aware M3 is M5+ work (PR #900). // // The non-auto-pin case (request flows with ObservedRouteVersion=0, -// M3 gate short-circuits) is the safe non-regressing posture for -// non-migrated callers — the gate cannot retroactively pin reads -// it was not present for. Adapters that want M3 protection must -// migrate to pin at BeginTxn per §4.1. +// Composed-1 owner gate short-circuits) is the safe non-regressing +// posture for non-migrated callers — the owner gate cannot +// retroactively pin reads it was not present for. The write-fence +// gate still checks the current route snapshot at apply time. +// Adapters that want full M3 owner protection must migrate to pin at +// BeginTxn per §4.1. // // Extracted from dispatchTxnWithComposed1Retry to keep its // cyclomatic complexity in the cyclop budget. @@ -1007,6 +1018,9 @@ func isComposed1RetryableError(err error) bool { // shard router. Extracted from Dispatch to keep that method's branch // count within the cyclop budget after the 7a registration gate landed. func (c *ShardedCoordinator) dispatchNonTxn(ctx context.Context, reqs *OperationGroup[OP]) (*CoordinateResponse, error) { + if hasExplicitGroupElem(reqs.Elems) { + return c.dispatchExplicitGroupNonTxn(ctx, reqs) + } logs, err := c.requestLogs(ctx, reqs) if err != nil { return nil, err @@ -1018,6 +1032,92 @@ func (c *ShardedCoordinator) dispatchNonTxn(ctx context.Context, reqs *Operation return &CoordinateResponse{CommitIndex: r.CommitIndex}, nil } +func hasExplicitGroupElem(elems []*Elem[OP]) bool { + for _, elem := range elems { + if elem != nil && elem.GroupID != 0 { + return true + } + } + return false +} + +func (c *ShardedCoordinator) dispatchExplicitGroupNonTxn(ctx context.Context, reqs *OperationGroup[OP]) (*CoordinateResponse, error) { + logs, gids, err := c.rawLogsWithGroups(ctx, reqs) + if err != nil { + return nil, err + } + var maxIndex uint64 + for i, gid := range gids { + g, err := c.txnGroupForID(gid) + if err != nil { + return nil, err + } + resp, err := g.Txn.Commit(ctx, []*pb.Request{logs[i]}) + if err != nil { + return nil, errors.WithStack(err) + } + if resp != nil && resp.CommitIndex > maxIndex { + maxIndex = resp.CommitIndex + } + } + return &CoordinateResponse{CommitIndex: maxIndex}, nil +} + +func (c *ShardedCoordinator) dispatchRawWithComposed1Retry(ctx context.Context, reqs *OperationGroup[OP]) (*CoordinateResponse, error) { + for attempt := 0; attempt <= composed1RetryAttempts; attempt++ { + resp, handled, err := c.dispatchBeforeShardRouting(ctx, reqs) + if !handled { + resp, err = c.dispatchNonTxn(ctx, reqs) + } + if err == nil { + return resp, nil + } + if !errors.Is(err, ErrComposed1VersionGCd) || attempt == composed1RetryAttempts || c.engine == nil { + return resp, err + } + if !c.canRetryRawVersionGC(reqs) { + return resp, err + } + reqs.ObservedRouteVersion = c.engine.Version() + } + return nil, errors.WithStack(ErrInvalidRequest) +} + +func (c *ShardedCoordinator) canRetryRawVersionGC(reqs *OperationGroup[OP]) bool { + if c == nil || c.router == nil || reqs == nil || hasDelPrefixElem(reqs.Elems) { + return false + } + var ( + firstGID uint64 + seen bool + ) + for _, elem := range reqs.Elems { + gid, ok := c.rawElemGroupID(elem) + if !ok { + return false + } + if !seen { + firstGID = gid + seen = true + continue + } + if gid != firstGID { + return false + } + } + return seen +} + +func (c *ShardedCoordinator) rawElemGroupID(elem *Elem[OP]) (uint64, bool) { + if elem == nil { + return 0, false + } + if elem.GroupID != 0 { + return elem.GroupID, true + } + return c.router.ResolveGroup(elem.Key) +} + // hasDelPrefixElem returns true if any element is a DelPrefix operation. func hasDelPrefixElem(elems []*Elem[OP]) bool { for _, e := range elems { @@ -1045,13 +1145,16 @@ func validateDelPrefixOnly(elems []*Elem[OP]) error { // pb.Request (the FSM's extractDelPrefix processes only the first DEL_PREFIX // mutation per request). All requests are batched into a single Commit call // per shard group. -func (c *ShardedCoordinator) dispatchDelPrefixBroadcast(ctx context.Context, isTxn bool, elems []*Elem[OP]) (*CoordinateResponse, error) { +func (c *ShardedCoordinator) dispatchDelPrefixBroadcast(ctx context.Context, isTxn bool, elems []*Elem[OP], observedRouteVersion uint64) (*CoordinateResponse, error) { if isTxn { return nil, errors.Wrap(ErrInvalidRequest, "DEL_PREFIX not supported in transactions") } if err := validateDelPrefixOnly(elems); err != nil { return nil, err } + if err := c.rejectWriteFencedDelPrefixes(elems); err != nil { + return nil, err + } ts, err := c.allocateTimestamp(ctx, "allocate DEL_PREFIX broadcast ts") if err != nil { @@ -1060,16 +1163,121 @@ func (c *ShardedCoordinator) dispatchDelPrefixBroadcast(ctx context.Context, isT requests := make([]*pb.Request, 0, len(elems)) for _, elem := range elems { requests = append(requests, &pb.Request{ - IsTxn: false, - Phase: pb.Phase_NONE, - Ts: ts, - Mutations: []*pb.Mutation{elemToMutation(elem)}, + IsTxn: false, + Phase: pb.Phase_NONE, + Ts: ts, + Mutations: []*pb.Mutation{elemToMutation(elem)}, + ObservedRouteVersion: observedRouteVersion, }) } return c.broadcastToAllGroups(ctx, requests) } +func (c *ShardedCoordinator) rejectWriteFencedPointElems(elems []*Elem[OP]) error { + if c == nil || c.engine == nil { + return nil + } + for _, elem := range elems { + if elem == nil || elem.GroupID != 0 { + continue + } + if err := c.rejectWriteFencedPointKey(elem.Key); err != nil { + return err + } + } + return nil +} + +func (c *ShardedCoordinator) rejectWriteFencedPointKey(key []byte) error { + if c.partitionResolverRecognisesPointKey(key) { + return nil + } + rkey := routeKey(key) + if route, ok := c.engine.GetRoute(rkey); ok && route.State == distribution.RouteStateWriteFenced { + return errors.Wrapf(ErrRouteWriteFenced, "key %q routeKey %q", key, rkey) + } + start, end, ok := s3BucketAuxiliaryRouteRange(key) + if !ok { + return nil + } + for _, route := range c.engine.GetIntersectingRoutes(start, end) { + if route.State == distribution.RouteStateWriteFenced { + return errors.Wrapf(ErrRouteWriteFenced, "key %q route range [%q,%q)", key, start, end) + } + } + return nil +} + +func (c *ShardedCoordinator) partitionResolverRecognisesPointKey(key []byte) bool { + if c == nil || c.router == nil || c.router.partitionResolver == nil || len(key) == 0 { + return false + } + if _, ok := c.router.partitionResolver.ResolveGroup(key); ok { + return true + } + return c.router.partitionResolver.RecognisesPartitionedKey(key) +} + +func (c *ShardedCoordinator) writeFenceBypassKeysForElems(elems []*Elem[OP]) [][]byte { + if len(elems) == 0 { + return nil + } + out := make([][]byte, 0, len(elems)) + for _, elem := range elems { + if elem == nil || elem.Op == DelPrefix || (elem.GroupID == 0 && !c.partitionResolverClaimsPointKey(elem.Key)) { + continue + } + out = append(out, bytes.Clone(elem.Key)) + } + return out +} + +func (c *ShardedCoordinator) writeFenceBypassKeysByGroup(elems []*Elem[OP]) map[uint64][][]byte { + out := make(map[uint64][][]byte) + for _, elem := range elems { + if elem == nil || elem.Op == DelPrefix { + continue + } + gid := elem.GroupID + if gid == 0 { + var ok bool + gid, ok = c.router.ResolveGroup(elem.Key) + if !ok || !c.partitionResolverClaimsPointKey(elem.Key) { + continue + } + } + out[gid] = append(out[gid], bytes.Clone(elem.Key)) + } + return out +} + +func (c *ShardedCoordinator) partitionResolverClaimsPointKey(key []byte) bool { + if c == nil || c.router == nil || c.router.partitionResolver == nil || len(key) == 0 { + return false + } + _, ok := c.router.partitionResolver.ResolveGroup(key) + return ok +} + +func (c *ShardedCoordinator) rejectWriteFencedDelPrefixes(elems []*Elem[OP]) error { + if c == nil || c.engine == nil { + return nil + } + for _, elem := range elems { + if elem == nil { + continue + } + start, end := routePrefixRange(elem.Key) + for _, route := range c.engine.GetIntersectingRoutes(start, end) { + if route.State == distribution.RouteStateWriteFenced { + return errors.Wrapf(ErrRouteWriteFenced, "prefix %q route range [%q,%q)", elem.Key, start, end) + } + } + } + return nil +} + // broadcastToAllGroups sends the same set of requests to every configured // all-shard data group in parallel and returns the maximum commit index. func (c *ShardedCoordinator) broadcastToAllGroups(ctx context.Context, requests []*pb.Request) (*CoordinateResponse, error) { @@ -1122,6 +1330,7 @@ func (c *ShardedCoordinator) dispatchTxn(ctx context.Context, startTS uint64, co if err != nil { return nil, err } + bypassKeysByGroup := c.writeFenceBypassKeysByGroup(elems) primaryKey := primaryKeyForElems(elems) if len(primaryKey) == 0 { return nil, errors.WithStack(ErrTxnPrimaryKeyRequired) @@ -1146,7 +1355,7 @@ func (c *ShardedCoordinator) dispatchTxn(ctx context.Context, startTS uint64, co if err := StampGroupedMutationCommitTS(grouped, commitTS); err != nil { return nil, err } - return c.dispatchMultiShardTxn(ctx, startTS, commitTS, prevCommitTS, primaryKey, grouped, gids, readKeys, observedRouteVersion) + return c.dispatchMultiShardTxn(ctx, startTS, commitTS, prevCommitTS, primaryKey, grouped, gids, readKeys, observedRouteVersion, bypassKeysByGroup) } // dispatchMultiShardTxn runs the 2PC path. Extracted from dispatchTxn to keep @@ -1154,7 +1363,7 @@ func (c *ShardedCoordinator) dispatchTxn(ctx context.Context, startTS uint64, co // P2 round-10) was added; the multi-shard branch already carries five linear // error checks (groupReadKeys, prewrite, commitPrimary, abortCleanup, // commitSecondaries) that pushed the parent over the 10-edge limit. -func (c *ShardedCoordinator) dispatchMultiShardTxn(ctx context.Context, startTS, commitTS, prevCommitTS uint64, primaryKey []byte, grouped map[uint64][]*pb.Mutation, gids []uint64, readKeys [][]byte, observedRouteVersion uint64) (*CoordinateResponse, error) { +func (c *ShardedCoordinator) dispatchMultiShardTxn(ctx context.Context, startTS, commitTS, prevCommitTS uint64, primaryKey []byte, grouped map[uint64][]*pb.Mutation, gids []uint64, readKeys [][]byte, observedRouteVersion uint64, bypassKeysByGroup map[uint64][][]byte) (*CoordinateResponse, error) { // Fail-closed when a retry carries the option-2 dedup probe key but its // write set / read set spans shards (codex P2 round-10 "reject retries // that leave the one-phase path"). The 2PC log builders only encode @@ -1177,12 +1386,12 @@ func (c *ShardedCoordinator) dispatchMultiShardTxn(ctx context.Context, startTS, if err != nil { return nil, err } - prepared, err := c.prewriteTxn(ctx, startTS, commitTS, primaryKey, grouped, gids, groupedReadKeys, observedRouteVersion) + prepared, err := c.prewriteTxn(ctx, startTS, commitTS, primaryKey, grouped, gids, groupedReadKeys, observedRouteVersion, bypassKeysByGroup) if err != nil { return nil, err } - primaryGid, maxIndex, err := c.commitPrimaryTxn(ctx, startTS, primaryKey, grouped, commitTS, observedRouteVersion) + primaryGid, maxIndex, err := c.commitPrimaryTxn(ctx, startTS, primaryKey, grouped, gids, commitTS, observedRouteVersion, bypassKeysByGroup) if err != nil { // abortPreparedTxn must run even when ctx was the reason // commitPrimaryTxn failed — otherwise prewrite intents on @@ -1209,7 +1418,7 @@ func (c *ShardedCoordinator) dispatchMultiShardTxn(ctx context.Context, startTS, // reachable by readers on the new owner. Both directions are // silent partial commits; surfacing the error is the only honest // posture (codex P1 on d8487672 + 6202b964, PR #900). - maxIndex, err = c.commitSecondaryTxns(ctx, startTS, primaryGid, primaryKey, grouped, gids, commitTS, maxIndex, observedRouteVersion) + maxIndex, err = c.commitSecondaryTxns(ctx, startTS, primaryGid, primaryKey, grouped, gids, commitTS, maxIndex, observedRouteVersion, bypassKeysByGroup) if err != nil { return nil, errors.WithStack(err) } @@ -1256,7 +1465,7 @@ func (c *ShardedCoordinator) dispatchSingleShardTxn(ctx context.Context, startTS // carries the one-phase dedup probe key for a retry that reuses a failed // attempt's write set. resp, err := g.Txn.Commit(ctx, []*pb.Request{ - onePhaseTxnRequestWithPrevCommit(startTS, commitTS, prevCommitTS, primaryKey, elems, readKeys, observedRouteVersion), + onePhaseTxnRequestWithPrevCommit(startTS, commitTS, prevCommitTS, primaryKey, elems, readKeys, observedRouteVersion, c.writeFenceBypassKeysForElems(elems)), }) if err != nil { return nil, errors.WithStack(err) @@ -1272,7 +1481,7 @@ type preparedGroup struct { keys []*pb.Mutation } -func (c *ShardedCoordinator) prewriteTxn(ctx context.Context, startTS, commitTS uint64, primaryKey []byte, grouped map[uint64][]*pb.Mutation, gids []uint64, groupedReadKeys map[uint64][][]byte, observedRouteVersion uint64) ([]preparedGroup, error) { +func (c *ShardedCoordinator) prewriteTxn(ctx context.Context, startTS, commitTS uint64, primaryKey []byte, grouped map[uint64][]*pb.Mutation, gids []uint64, groupedReadKeys map[uint64][][]byte, observedRouteVersion uint64, bypassKeysByGroup map[uint64][][]byte) ([]preparedGroup, error) { prepareMeta := txnMetaMutation(primaryKey, defaultTxnLockTTLms, 0) prepared := make([]preparedGroup, 0, len(gids)) @@ -1288,6 +1497,7 @@ func (c *ShardedCoordinator) prewriteTxn(ctx context.Context, startTS, commitTS Mutations: append([]*pb.Mutation{prepareMeta}, grouped[gid]...), ReadKeys: groupedReadKeys[gid], ObservedRouteVersion: observedRouteVersion, + WriteFenceBypassKeys: bypassKeysByGroup[gid], } if _, err := g.Txn.Commit(ctx, []*pb.Request{req}); err != nil { // Same WithoutCancel pattern as dispatchTxn's @@ -1321,9 +1531,9 @@ func (c *ShardedCoordinator) prewriteTxn(ctx context.Context, startTS, commitTS return prepared, nil } -func (c *ShardedCoordinator) commitPrimaryTxn(ctx context.Context, startTS uint64, primaryKey []byte, grouped map[uint64][]*pb.Mutation, commitTS uint64, observedRouteVersion uint64) (uint64, uint64, error) { - primaryGid := c.engineGroupIDForKey(primaryKey) - if primaryGid == 0 { +func (c *ShardedCoordinator) commitPrimaryTxn(ctx context.Context, startTS uint64, primaryKey []byte, grouped map[uint64][]*pb.Mutation, gids []uint64, commitTS uint64, observedRouteVersion uint64, bypassKeysByGroup map[uint64][][]byte) (uint64, uint64, error) { + primaryGid, ok := primaryGroupIDForKey(primaryKey, grouped, gids) + if !ok { return 0, 0, errors.WithStack(ErrInvalidRequest) } @@ -1340,6 +1550,7 @@ func (c *ShardedCoordinator) commitPrimaryTxn(ctx context.Context, startTS uint6 Ts: startTS, Mutations: append([]*pb.Mutation{meta}, keys...), ObservedRouteVersion: observedRouteVersion, + WriteFenceBypassKeys: bypassKeysByGroup[primaryGid], } r, err := g.Txn.Commit(ctx, []*pb.Request{req}) @@ -1352,7 +1563,18 @@ func (c *ShardedCoordinator) commitPrimaryTxn(ctx context.Context, startTS uint6 return primaryGid, r.CommitIndex, nil } -func (c *ShardedCoordinator) commitSecondaryTxns(ctx context.Context, startTS uint64, primaryGid uint64, primaryKey []byte, grouped map[uint64][]*pb.Mutation, gids []uint64, commitTS uint64, maxIndex uint64, observedRouteVersion uint64) (uint64, error) { +func primaryGroupIDForKey(primaryKey []byte, grouped map[uint64][]*pb.Mutation, gids []uint64) (uint64, bool) { + for _, gid := range gids { + for _, mut := range grouped[gid] { + if mut != nil && bytes.Equal(mut.Key, primaryKey) { + return gid, true + } + } + } + return 0, false +} + +func (c *ShardedCoordinator) commitSecondaryTxns(ctx context.Context, startTS uint64, primaryGid uint64, primaryKey []byte, grouped map[uint64][]*pb.Mutation, gids []uint64, commitTS uint64, maxIndex uint64, observedRouteVersion uint64, bypassKeysByGroup map[uint64][][]byte) (uint64, error) { // Secondary commits are best-effort for non-Composed-1 errors: // if a shard is unavailable after the primary commits, read-time // lock resolution will commit the remaining secondaries based on @@ -1389,6 +1611,7 @@ func (c *ShardedCoordinator) commitSecondaryTxns(ctx context.Context, startTS ui Ts: startTS, Mutations: append([]*pb.Mutation{meta}, keyMutations(grouped[gid])...), ObservedRouteVersion: observedRouteVersion, + WriteFenceBypassKeys: bypassKeysByGroup[gid], } r, err := commitSecondaryWithRetry(ctx, g, req) if err != nil { @@ -2111,25 +2334,33 @@ func (c *ShardedCoordinator) requestLogs(ctx context.Context, reqs *OperationGro } func (c *ShardedCoordinator) rawLogs(ctx context.Context, reqs *OperationGroup[OP]) ([]*pb.Request, error) { + logs, _, err := c.rawLogsWithGroups(ctx, reqs) + return logs, err +} + +func (c *ShardedCoordinator) rawLogsWithGroups(ctx context.Context, reqs *OperationGroup[OP]) ([]*pb.Request, []uint64, error) { grouped, gids, err := c.groupMutations(reqs.Elems, reqs.KeyVizLabel) if err != nil { - return nil, err + return nil, nil, err } + bypassKeysByGroup := c.writeFenceBypassKeysByGroup(reqs.Elems) logs := make([]*pb.Request, 0, len(gids)) for _, gid := range gids { ts, err := c.rawLogTimestamp(ctx) if err != nil { - return nil, err + return nil, nil, err } logs = append(logs, &pb.Request{ - IsTxn: false, - Phase: pb.Phase_NONE, - Ts: ts, - Mutations: grouped[gid], + IsTxn: false, + Phase: pb.Phase_NONE, + Ts: ts, + Mutations: grouped[gid], + ObservedRouteVersion: reqs.ObservedRouteVersion, + WriteFenceBypassKeys: bypassKeysByGroup[gid], }) } - return logs, nil + return logs, gids, nil } func (c *ShardedCoordinator) rawLogTimestamp(ctx context.Context) (uint64, error) { @@ -2170,7 +2401,8 @@ func (c *ShardedCoordinator) txnLogs(ctx context.Context, reqs *OperationGroup[O if err := StampGroupedMutationCommitTS(grouped, commitTS); err != nil { return nil, err } - return buildTxnLogs(reqs.StartTS, commitTS, grouped, gids, reqs.ObservedRouteVersion) + bypassKeysByGroup := c.writeFenceBypassKeysByGroup(reqs.Elems) + return buildTxnLogs(reqs.StartTS, commitTS, grouped, gids, reqs.ObservedRouteVersion, bypassKeysByGroup) } // observeMutation: counted pre-commit, so a mutation that subsequently @@ -2224,9 +2456,15 @@ func (c *ShardedCoordinator) groupMutations(reqs []*Elem[OP], label keyviz.Label return nil, nil, ErrInvalidRequest } mut := elemToMutation(req) - gid, ok := c.router.ResolveGroup(mut.Key) - if !ok { - return nil, nil, errors.Wrapf(ErrInvalidRequest, "no route for key %q", mut.Key) + gid := req.GroupID + if gid == 0 { + var ok bool + gid, ok = c.router.ResolveGroup(mut.Key) + if !ok { + return nil, nil, errors.Wrapf(ErrInvalidRequest, "no route for key %q", mut.Key) + } + } else if _, ok := c.groups[gid]; !ok { + return nil, nil, errors.Wrapf(ErrInvalidRequest, "no shard group %d for key %q", gid, mut.Key) } // Engine RouteID for keyviz observation; partition-resolved // keys observe under the !sqs|route|global RouteID until @@ -2246,7 +2484,7 @@ func (c *ShardedCoordinator) groupMutations(reqs []*Elem[OP], label keyviz.Label return grouped, gids, nil } -func buildTxnLogs(startTS uint64, commitTS uint64, grouped map[uint64][]*pb.Mutation, gids []uint64, observedRouteVersion uint64) ([]*pb.Request, error) { +func buildTxnLogs(startTS uint64, commitTS uint64, grouped map[uint64][]*pb.Mutation, gids []uint64, observedRouteVersion uint64, writeFenceBypassKeysByGroup map[uint64][][]byte) ([]*pb.Request, error) { logs := make([]*pb.Request, 0, len(gids)*txnPhaseCount) for _, gid := range gids { muts := grouped[gid] @@ -2263,6 +2501,7 @@ func buildTxnLogs(startTS uint64, commitTS uint64, grouped map[uint64][]*pb.Muta {Op: pb.Op_PUT, Key: []byte(txnMetaPrefix), Value: EncodeTxnMeta(TxnMeta{PrimaryKey: primaryKey, LockTTLms: defaultTxnLockTTLms, CommitTS: 0})}, }, muts...), ObservedRouteVersion: observedRouteVersion, + WriteFenceBypassKeys: writeFenceBypassKeysByGroup[gid], }, &pb.Request{ IsTxn: true, @@ -2272,6 +2511,7 @@ func buildTxnLogs(startTS uint64, commitTS uint64, grouped map[uint64][]*pb.Muta {Op: pb.Op_PUT, Key: []byte(txnMetaPrefix), Value: EncodeTxnMeta(TxnMeta{PrimaryKey: primaryKey, LockTTLms: 0, CommitTS: commitTS})}, }, keys...), ObservedRouteVersion: observedRouteVersion, + WriteFenceBypassKeys: writeFenceBypassKeysByGroup[gid], }, ) } @@ -2359,7 +2599,7 @@ func (c *ShardedCoordinator) renewHLCLeases(ctx context.Context) <-chan struct{} go func(gid uint64, group *ShardGroup) { defer wg.Done() defer c.finishHLCLeaseRenewal(gid) - pctx, cancel := context.WithTimeout(ctx, hlcRenewalInterval) + pctx, cancel := context.WithTimeout(ctx, hlcRenewalTimeout) defer cancel() c.renewHLCLease(pctx, gid, group) }(gid, group) diff --git a/kv/sharded_coordinator_del_prefix_test.go b/kv/sharded_coordinator_del_prefix_test.go index 11de9de78..4c4e6d8dc 100644 --- a/kv/sharded_coordinator_del_prefix_test.go +++ b/kv/sharded_coordinator_del_prefix_test.go @@ -6,6 +6,7 @@ import ( "testing" "github.com/bootjp/elastickv/distribution" + "github.com/bootjp/elastickv/internal/s3keys" pb "github.com/bootjp/elastickv/proto" "github.com/bootjp/elastickv/store" "github.com/stretchr/testify/require" @@ -121,6 +122,420 @@ func TestShardedCoordinator_DelPrefixBroadcastsToAllGroups(t *testing.T) { "same DEL_PREFIX element must use the same timestamp across shards") } +func TestShardedCoordinator_DelPrefixDoesNotAutoPinObservedRouteVersion(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + require.NoError(t, engine.ApplySnapshot(distribution.CatalogSnapshot{ + Version: 7, + Routes: []distribution.RouteDescriptor{ + {RouteID: 1, Start: nil, End: nil, GroupID: 1, State: distribution.RouteStateActive}, + }, + })) + txn := &recordingTransactional{} + coord := NewShardedCoordinator(engine, map[uint64]*ShardGroup{ + 1: {Txn: txn}, + }, 1, NewHLC(), nil) + + _, err := coord.Dispatch(context.Background(), &OperationGroup[OP]{ + Elems: []*Elem[OP]{{Op: DelPrefix, Key: []byte("user:")}}, + }) + require.NoError(t, err) + require.Len(t, txn.requests, 1) + require.Zero(t, txn.requests[0].GetObservedRouteVersion()) +} + +func TestShardedCoordinator_RawWriteDoesNotAutoPinObservedRouteVersion(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + require.NoError(t, engine.ApplySnapshot(distribution.CatalogSnapshot{ + Version: 9, + Routes: []distribution.RouteDescriptor{ + {RouteID: 1, Start: nil, End: nil, GroupID: 1, State: distribution.RouteStateActive}, + }, + })) + txn := &recordingTransactional{} + coord := NewShardedCoordinator(engine, map[uint64]*ShardGroup{ + 1: {Txn: txn}, + }, 1, NewHLC(), nil) + + _, err := coord.Dispatch(context.Background(), &OperationGroup[OP]{ + Elems: []*Elem[OP]{{Op: Put, Key: []byte("k"), Value: []byte("v")}}, + }) + require.NoError(t, err) + require.Len(t, txn.requests, 1) + require.Zero(t, txn.requests[0].GetObservedRouteVersion()) +} + +func TestShardedCoordinator_RetriesRawWriteWhenObservedRouteVersionGCd(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + require.NoError(t, engine.ApplySnapshot(distribution.CatalogSnapshot{ + Version: 9, + Routes: []distribution.RouteDescriptor{ + {RouteID: 1, Start: nil, End: nil, GroupID: 1, State: distribution.RouteStateActive}, + }, + })) + txn := &recordingTransactional{ + errs: []error{ErrComposed1VersionGCd}, + } + coord := NewShardedCoordinator(engine, map[uint64]*ShardGroup{ + 1: {Txn: txn}, + }, 1, NewHLC(), nil) + + _, err := coord.Dispatch(context.Background(), &OperationGroup[OP]{ + ObservedRouteVersion: 3, + Elems: []*Elem[OP]{{Op: Put, Key: []byte("k"), Value: []byte("v")}}, + }) + require.NoError(t, err) + require.Len(t, txn.requests, 2) + require.Equal(t, uint64(3), txn.requests[0].GetObservedRouteVersion()) + require.Equal(t, uint64(9), txn.requests[1].GetObservedRouteVersion()) +} + +func TestShardedCoordinator_DoesNotRetryDelPrefixWhenObservedRouteVersionGCd(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + require.NoError(t, engine.ApplySnapshot(distribution.CatalogSnapshot{ + Version: 7, + Routes: []distribution.RouteDescriptor{ + {RouteID: 1, Start: nil, End: nil, GroupID: 1, State: distribution.RouteStateActive}, + }, + })) + txn := &recordingTransactional{ + errs: []error{ErrComposed1VersionGCd}, + } + coord := NewShardedCoordinator(engine, map[uint64]*ShardGroup{ + 1: {Txn: txn}, + }, 1, NewHLC(), nil) + + _, err := coord.Dispatch(context.Background(), &OperationGroup[OP]{ + ObservedRouteVersion: 2, + Elems: []*Elem[OP]{{Op: DelPrefix, Key: []byte("user:")}}, + }) + require.ErrorIs(t, err, ErrComposed1VersionGCd) + require.Len(t, txn.requests, 1) + require.Equal(t, uint64(2), txn.requests[0].GetObservedRouteVersion()) +} + +func TestShardedCoordinator_DoesNotRetryMultiShardRawWriteWhenObservedRouteVersionGCd(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + require.NoError(t, engine.ApplySnapshot(distribution.CatalogSnapshot{ + Version: 11, + Routes: []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte(""), End: []byte("m"), GroupID: 1, State: distribution.RouteStateActive}, + {RouteID: 2, Start: []byte("m"), End: nil, GroupID: 2, State: distribution.RouteStateActive}, + }, + })) + g1Txn := &recordingTransactional{} + g2Txn := &recordingTransactional{errs: []error{ErrComposed1VersionGCd}} + coord := NewShardedCoordinator(engine, map[uint64]*ShardGroup{ + 1: {Txn: g1Txn}, + 2: {Txn: g2Txn}, + }, 1, NewHLC(), nil) + + _, err := coord.Dispatch(context.Background(), &OperationGroup[OP]{ + ObservedRouteVersion: 3, + Elems: []*Elem[OP]{ + {Op: Put, Key: []byte("a"), Value: []byte("v1")}, + {Op: Put, Key: []byte("z"), Value: []byte("v2")}, + }, + }) + + require.ErrorIs(t, err, ErrComposed1VersionGCd) + require.Len(t, g1Txn.requests, 1) + require.Len(t, g2Txn.requests, 1) + require.Equal(t, uint64(3), g1Txn.requests[0].GetObservedRouteVersion()) + require.Equal(t, uint64(3), g2Txn.requests[0].GetObservedRouteVersion()) +} + +func TestShardedCoordinator_DoesNotRetryRawWriteWhenRetryRouteBecomesMultiShard(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + require.NoError(t, engine.ApplySnapshot(distribution.CatalogSnapshot{ + Version: 11, + Routes: []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte(""), End: nil, GroupID: 1, State: distribution.RouteStateActive}, + }, + })) + g1Txn := &recordingTransactional{ + errs: []error{ErrComposed1VersionGCd}, + onCommit: func(call int, _ *pb.Request) { + if call != 0 { + return + } + require.NoError(t, engine.ApplySnapshot(distribution.CatalogSnapshot{ + Version: 12, + Routes: []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte(""), End: []byte("m"), GroupID: 1, State: distribution.RouteStateActive}, + {RouteID: 2, Start: []byte("m"), End: nil, GroupID: 2, State: distribution.RouteStateActive}, + }, + })) + }, + } + g2Txn := &recordingTransactional{} + coord := NewShardedCoordinator(engine, map[uint64]*ShardGroup{ + 1: {Txn: g1Txn}, + 2: {Txn: g2Txn}, + }, 1, NewHLC(), nil) + + _, err := coord.Dispatch(context.Background(), &OperationGroup[OP]{ + ObservedRouteVersion: 3, + Elems: []*Elem[OP]{ + {Op: Put, Key: []byte("a"), Value: []byte("v1")}, + {Op: Put, Key: []byte("z"), Value: []byte("v2")}, + }, + }) + + require.ErrorIs(t, err, ErrComposed1VersionGCd) + require.Len(t, g1Txn.requests, 1) + require.Len(t, g2Txn.requests, 0) + require.Equal(t, uint64(3), g1Txn.requests[0].GetObservedRouteVersion()) +} + +func TestShardedCoordinatorRejectsPointWriteOnWriteFencedRoute(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + require.NoError(t, engine.ApplySnapshot(distribution.CatalogSnapshot{ + Version: 1, + Routes: []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte(""), End: []byte("m"), GroupID: 1, State: distribution.RouteStateActive}, + {RouteID: 2, Start: []byte("m"), End: nil, GroupID: 2, State: distribution.RouteStateWriteFenced}, + }, + })) + + g2Txn := &recordingTransactional{} + coord := NewShardedCoordinator(engine, map[uint64]*ShardGroup{ + 1: {Txn: &recordingTransactional{}}, + 2: {Txn: g2Txn}, + }, 1, NewHLC(), nil) + + _, err := coord.Dispatch(context.Background(), &OperationGroup[OP]{ + Elems: []*Elem[OP]{{Op: Put, Key: []byte("z"), Value: []byte("v")}}, + }) + require.ErrorIs(t, err, ErrRouteWriteFenced) + require.Empty(t, g2Txn.requests, "coordinator must reject before proposing to the fenced shard") +} + +func TestShardedCoordinatorRejectsEmptyKeyWriteOnLeadingWriteFencedRoute(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + require.NoError(t, engine.ApplySnapshot(distribution.CatalogSnapshot{ + Version: 1, + Routes: []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte(""), End: []byte("m"), GroupID: 1, State: distribution.RouteStateWriteFenced}, + {RouteID: 2, Start: []byte("m"), End: nil, GroupID: 2, State: distribution.RouteStateActive}, + }, + })) + + g1Txn := &recordingTransactional{} + coord := NewShardedCoordinator(engine, map[uint64]*ShardGroup{ + 1: {Txn: g1Txn}, + 2: {Txn: &recordingTransactional{}}, + }, 1, NewHLC(), nil) + + _, err := coord.Dispatch(context.Background(), &OperationGroup[OP]{ + Elems: []*Elem[OP]{{Op: Put, Key: []byte{}, Value: []byte("v")}}, + }) + require.ErrorIs(t, err, ErrRouteWriteFenced) + require.Empty(t, g1Txn.requests, "coordinator must reject the empty key before proposing to the fenced shard") +} + +func TestShardedCoordinatorRejectsS3BucketAuxiliaryPointWriteOnWriteFencedRoute(t *testing.T) { + t.Parallel() + + const bucket = "bucket-a" + engine := distribution.NewEngine() + require.NoError(t, engine.ApplySnapshot(distribution.CatalogSnapshot{ + Version: 1, + Routes: s3BucketAuxiliaryFenceRoutes(bucket, 1, 2), + })) + + g1Txn := &recordingTransactional{} + g2Txn := &recordingTransactional{} + coord := NewShardedCoordinator(engine, map[uint64]*ShardGroup{ + 1: {Txn: g1Txn}, + 2: {Txn: g2Txn}, + }, 1, NewHLC(), nil) + + for _, key := range [][]byte{ + s3keys.BucketMetaKey(bucket), + s3keys.BucketGenerationKey(bucket), + } { + _, err := coord.Dispatch(context.Background(), &OperationGroup[OP]{ + Elems: []*Elem[OP]{{Op: Put, Key: key, Value: []byte("v")}}, + }) + require.ErrorIs(t, err, ErrRouteWriteFenced) + require.Empty(t, g1Txn.requests, "coordinator must reject before proposing to the raw-key shard") + require.Empty(t, g2Txn.requests, "coordinator must reject before proposing to the fenced shard") + } +} + +func TestShardedCoordinatorRejectsDelPrefixIntersectingWriteFencedRoute(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + require.NoError(t, engine.ApplySnapshot(distribution.CatalogSnapshot{ + Version: 1, + Routes: []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte(""), End: []byte("m"), GroupID: 1, State: distribution.RouteStateActive}, + {RouteID: 2, Start: []byte("m"), End: nil, GroupID: 2, State: distribution.RouteStateWriteFenced}, + }, + })) + + g1Txn := &recordingTransactional{} + g2Txn := &recordingTransactional{} + coord := NewShardedCoordinator(engine, map[uint64]*ShardGroup{ + 1: {Txn: g1Txn}, + 2: {Txn: g2Txn}, + }, 1, NewHLC(), nil) + + _, err := coord.Dispatch(context.Background(), &OperationGroup[OP]{ + Elems: []*Elem[OP]{{Op: DelPrefix, Key: []byte("z")}}, + }) + require.ErrorIs(t, err, ErrRouteWriteFenced) + require.Empty(t, g1Txn.requests) + require.Empty(t, g2Txn.requests) +} + +func TestShardedCoordinatorRejectsFullRangeDelPrefixWhenRouteIsWriteFenced(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + require.NoError(t, engine.ApplySnapshot(distribution.CatalogSnapshot{ + Version: 1, + Routes: []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte(""), End: []byte("m"), GroupID: 1, State: distribution.RouteStateActive}, + {RouteID: 2, Start: []byte("m"), End: nil, GroupID: 2, State: distribution.RouteStateWriteFenced}, + }, + })) + + g1Txn := &recordingTransactional{} + g2Txn := &recordingTransactional{} + coord := NewShardedCoordinator(engine, map[uint64]*ShardGroup{ + 1: {Txn: g1Txn}, + 2: {Txn: g2Txn}, + }, 1, NewHLC(), nil) + + _, err := coord.Dispatch(context.Background(), &OperationGroup[OP]{ + Elems: []*Elem[OP]{{Op: DelPrefix, Key: nil}}, + }) + require.ErrorIs(t, err, ErrRouteWriteFenced) + require.Empty(t, g1Txn.requests) + require.Empty(t, g2Txn.requests) +} + +func TestShardedCoordinatorRejectsBroadInternalDelPrefixWhenRouteIsWriteFenced(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + require.NoError(t, engine.ApplySnapshot(distribution.CatalogSnapshot{ + Version: 1, + Routes: []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte(""), End: []byte("m"), GroupID: 1, State: distribution.RouteStateActive}, + {RouteID: 2, Start: []byte("m"), End: nil, GroupID: 2, State: distribution.RouteStateWriteFenced}, + }, + })) + + for _, prefix := range [][]byte{ + []byte("!redis|"), + []byte("!lst|"), + } { + t.Run(string(prefix), func(t *testing.T) { + t.Parallel() + + g1Txn := &recordingTransactional{} + g2Txn := &recordingTransactional{} + coord := NewShardedCoordinator(engine, map[uint64]*ShardGroup{ + 1: {Txn: g1Txn}, + 2: {Txn: g2Txn}, + }, 1, NewHLC(), nil) + + _, err := coord.Dispatch(context.Background(), &OperationGroup[OP]{ + Elems: []*Elem[OP]{{Op: DelPrefix, Key: prefix}}, + }) + require.ErrorIs(t, err, ErrRouteWriteFenced) + require.Empty(t, g1Txn.requests) + require.Empty(t, g2Txn.requests) + }) + } +} + +func TestShardedCoordinatorAllowsRawSQSLookingDelPrefixWhenUnrelatedRouteIsWriteFenced(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + require.NoError(t, engine.ApplySnapshot(distribution.CatalogSnapshot{ + Version: 1, + Routes: []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte(""), End: []byte("m"), GroupID: 1, State: distribution.RouteStateActive}, + {RouteID: 2, Start: []byte("m"), End: nil, GroupID: 2, State: distribution.RouteStateWriteFenced}, + }, + })) + + g1Txn := &recordingTransactional{} + g2Txn := &recordingTransactional{} + coord := NewShardedCoordinator(engine, map[uint64]*ShardGroup{ + 1: {Txn: g1Txn}, + 2: {Txn: g2Txn}, + }, 1, NewHLC(), nil) + + _, err := coord.Dispatch(context.Background(), &OperationGroup[OP]{ + Elems: []*Elem[OP]{{Op: DelPrefix, Key: []byte("!sqs|foo")}}, + }) + require.NoError(t, err) + require.NotEmpty(t, g1Txn.requests) + require.NotEmpty(t, g2Txn.requests) +} + +func TestShardedCoordinatorAllowsS3BucketDelPrefixWhenUnrelatedRouteIsWriteFenced(t *testing.T) { + t.Parallel() + + const ( + activeBucket = "bucket-a" + fencedBucket = "bucket-b" + generation = uint64(7) + ) + activeStart := s3keys.RoutePrefixForBucketAnyGeneration(activeBucket) + activeEnd := prefixScanEnd(activeStart) + fencedStart := s3keys.RoutePrefixForBucketAnyGeneration(fencedBucket) + fencedEnd := prefixScanEnd(fencedStart) + + engine := distribution.NewEngine() + require.NoError(t, engine.ApplySnapshot(distribution.CatalogSnapshot{ + Version: 1, + Routes: []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte(""), End: activeStart, GroupID: 1, State: distribution.RouteStateActive}, + {RouteID: 2, Start: activeStart, End: activeEnd, GroupID: 1, State: distribution.RouteStateActive}, + {RouteID: 3, Start: activeEnd, End: fencedStart, GroupID: 1, State: distribution.RouteStateActive}, + {RouteID: 4, Start: fencedStart, End: fencedEnd, GroupID: 2, State: distribution.RouteStateWriteFenced}, + {RouteID: 5, Start: fencedEnd, End: nil, GroupID: 1, State: distribution.RouteStateActive}, + }, + })) + + g1Txn := &recordingTransactional{} + g2Txn := &recordingTransactional{} + coord := NewShardedCoordinator(engine, map[uint64]*ShardGroup{ + 1: {Txn: g1Txn}, + 2: {Txn: g2Txn}, + }, 1, NewHLC(), nil) + + _, err := coord.Dispatch(context.Background(), &OperationGroup[OP]{ + Elems: []*Elem[OP]{{Op: DelPrefix, Key: s3keys.ObjectManifestPrefixForBucket(activeBucket, generation)}}, + }) + require.NoError(t, err) + require.NotEmpty(t, g1Txn.requests) + require.NotEmpty(t, g2Txn.requests, "DEL_PREFIX still broadcasts to every group after the narrow fence precheck passes") +} + // TestShardedCoordinator_DelPrefixRejectsTxn verifies that DEL_PREFIX inside // a transactional group is rejected. func TestShardedCoordinator_DelPrefixRejectsTxn(t *testing.T) { diff --git a/kv/sharded_coordinator_partition_test.go b/kv/sharded_coordinator_partition_test.go index ae46301a1..7b224db8a 100644 --- a/kv/sharded_coordinator_partition_test.go +++ b/kv/sharded_coordinator_partition_test.go @@ -2,6 +2,7 @@ package kv import ( "context" + "errors" "sync" "testing" @@ -115,6 +116,118 @@ func TestShardedCoordinator_DispatchHonoursPartitionResolver(t *testing.T) { require.Equal(t, []byte("!sqs|msg|data|p|partitioned-key"), calls[0]) } +func TestShardedCoordinatorWriteFencePrecheckHonoursPartitionResolver(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + require.NoError(t, engine.ApplySnapshot(distribution.CatalogSnapshot{ + Version: 1, + Routes: []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte(""), End: nil, GroupID: 1, State: distribution.RouteStateWriteFenced}, + }, + })) + + g1 := &recordingTransactional{ + responses: []*TransactionResponse{{CommitIndex: 1}}, + } + g42 := &recordingTransactional{ + responses: []*TransactionResponse{{CommitIndex: 42}}, + } + coord := NewShardedCoordinator(engine, map[uint64]*ShardGroup{ + 1: {Txn: g1, Store: store.NewMVCCStore()}, + 42: {Txn: g42, Store: store.NewMVCCStore()}, + }, 1, NewHLC(), nil) + + key := []byte("!sqs|msg|data|p|partitioned-key") + coord.WithPartitionResolver(&stubResolver{claim: map[string]uint64{ + string(key): 42, + }}) + + resp, err := coord.Dispatch(context.Background(), &OperationGroup[OP]{ + Elems: []*Elem[OP]{{Op: Put, Key: key, Value: []byte("v")}}, + }) + require.NoError(t, err) + require.NotNil(t, resp) + require.Equal(t, uint64(42), resp.CommitIndex) + require.Empty(t, g1.requests, "engine route fence must not preempt resolver-owned keys") + require.Len(t, g42.requests, 1) + require.Equal(t, [][]byte{key}, g42.requests[0].GetWriteFenceBypassKeys(), + "resolver-owned point writes must carry the FSM write-fence bypass marker") +} + +func TestShardedCoordinatorWriteFencePrecheckSkipsUnresolvedPartitionKey(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + require.NoError(t, engine.ApplySnapshot(distribution.CatalogSnapshot{ + Version: 1, + Routes: []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte(""), End: nil, GroupID: 1, State: distribution.RouteStateWriteFenced}, + }, + })) + + g1 := &recordingTransactional{ + responses: []*TransactionResponse{{CommitIndex: 1}}, + } + coord := NewShardedCoordinator(engine, map[uint64]*ShardGroup{ + 1: {Txn: g1, Store: store.NewMVCCStore()}, + }, 1, NewHLC(), nil) + + key := []byte("!sqs|msg|data|p|unknown-partition-key") + coord.WithPartitionResolver(&stubResolver{ + recognisedPrefix: []byte("!sqs|msg|data|p|"), + }) + + _, err := coord.Dispatch(context.Background(), &OperationGroup[OP]{ + Elems: []*Elem[OP]{{Op: Put, Key: key, Value: []byte("v")}}, + }) + require.ErrorIs(t, err, ErrInvalidRequest) + require.False(t, errors.Is(err, ErrRouteWriteFenced), + "resolver-recognised keys must fail through resolver routing, not engine route fences") + require.Empty(t, g1.requests) +} + +func TestShardedCoordinatorTxnCarriesResolverWriteFenceBypassKeys(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + require.NoError(t, engine.ApplySnapshot(distribution.CatalogSnapshot{ + Version: 1, + Routes: []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte(""), End: nil, GroupID: 1, State: distribution.RouteStateWriteFenced}, + }, + })) + + g1 := &recordingTransactional{ + responses: []*TransactionResponse{{CommitIndex: 1}}, + } + g42 := &recordingTransactional{ + responses: []*TransactionResponse{{CommitIndex: 42}}, + } + coord := NewShardedCoordinator(engine, map[uint64]*ShardGroup{ + 1: {Txn: g1, Store: store.NewMVCCStore()}, + 42: {Txn: g42, Store: store.NewMVCCStore()}, + }, 1, NewHLC(), nil) + + key := []byte("!sqs|msg|data|p|txn-partitioned-key") + coord.WithPartitionResolver(&stubResolver{claim: map[string]uint64{ + string(key): 42, + }}) + + resp, err := coord.Dispatch(context.Background(), &OperationGroup[OP]{ + IsTxn: true, + StartTS: 100, + CommitTS: 200, + Elems: []*Elem[OP]{{Op: Put, Key: key, Value: []byte("v")}}, + }) + require.NoError(t, err) + require.NotNil(t, resp) + require.Empty(t, g1.requests) + require.Len(t, g42.requests, 1) + require.Equal(t, [][]byte{key}, g42.requests[0].GetWriteFenceBypassKeys(), + "single-shard txn prepare/one-phase request must preserve the resolver-owned point key for the FSM") +} + // TestShardedCoordinator_DispatchSplitsMutationsByResolverGroup is // the genuine regression for the Gemini-HIGH groupMutations // bypass: a Dispatch with mutations belonging to TWO different diff --git a/kv/sharded_coordinator_txn_test.go b/kv/sharded_coordinator_txn_test.go index 42b65f328..ec6c651ff 100644 --- a/kv/sharded_coordinator_txn_test.go +++ b/kv/sharded_coordinator_txn_test.go @@ -9,6 +9,7 @@ import ( "github.com/bootjp/elastickv/distribution" "github.com/bootjp/elastickv/internal/raftengine" + "github.com/bootjp/elastickv/keyviz" pb "github.com/bootjp/elastickv/proto" "github.com/bootjp/elastickv/store" "github.com/stretchr/testify/require" @@ -20,6 +21,7 @@ type recordingTransactional struct { requests []*pb.Request responses []*TransactionResponse errs []error + onCommit func(call int, req *pb.Request) } func (s *recordingTransactional) Commit(_ context.Context, reqs []*pb.Request) (*TransactionResponse, error) { @@ -31,6 +33,9 @@ func (s *recordingTransactional) Commit(_ context.Context, reqs []*pb.Request) ( } s.requests = append(s.requests, cloneTxnRequest(reqs[0])) call := len(s.requests) - 1 + if s.onCommit != nil { + s.onCommit(call, s.requests[call]) + } if call < len(s.errs) && s.errs[call] != nil { return nil, s.errs[call] } @@ -56,6 +61,120 @@ func cloneTxnRequest(req *pb.Request) *pb.Request { return request } +func TestShardedCoordinatorGroupMutationsUsesExplicitElemGroup(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + engine.UpdateRoute([]byte(""), []byte("m"), 1) + engine.UpdateRoute([]byte("m"), nil, 2) + coord := NewShardedCoordinator(engine, map[uint64]*ShardGroup{ + 1: {}, + 2: {}, + }, 1, NewHLC(), nil) + + grouped, gids, err := coord.groupMutations([]*Elem[OP]{ + {Op: Del, Key: []byte("a-key"), GroupID: 2}, + }, keyviz.Label("")) + require.NoError(t, err) + require.Equal(t, []uint64{2}, gids) + require.Len(t, grouped[2], 1) + require.Equal(t, []byte("a-key"), grouped[2][0].Key) +} + +func TestShardedCoordinatorDispatchTxn_CommitPrimaryUsesPinnedGroup(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + engine.UpdateRoute([]byte(""), nil, 1) + + g1Txn := &recordingTransactional{ + responses: []*TransactionResponse{ + {CommitIndex: 3}, + {CommitIndex: 11}, + }, + } + g2Txn := &recordingTransactional{ + responses: []*TransactionResponse{ + {CommitIndex: 5}, + {CommitIndex: 27}, + }, + } + + coord := NewShardedCoordinator(engine, map[uint64]*ShardGroup{ + 1: {Txn: g1Txn}, + 2: {Txn: g2Txn}, + }, 1, NewHLC(), nil) + + startTS := uint64(10) + _, err := coord.Dispatch(context.Background(), &OperationGroup[OP]{ + IsTxn: true, + StartTS: startTS, + Elems: []*Elem[OP]{ + {Op: Del, Key: []byte("a-key"), GroupID: 2}, + {Op: Put, Key: []byte("z-key"), Value: []byte("v"), GroupID: 1}, + }, + }) + require.NoError(t, err) + require.Len(t, g1Txn.requests, 2) + require.Len(t, g2Txn.requests, 2) + + g1Commit := g1Txn.requests[1] + g2Commit := g2Txn.requests[1] + require.Equal(t, [][]byte{[]byte("z-key")}, g1Txn.requests[0].WriteFenceBypassKeys) + require.Equal(t, [][]byte{[]byte("z-key")}, g1Commit.WriteFenceBypassKeys) + require.Equal(t, [][]byte{[]byte("a-key")}, g2Txn.requests[0].WriteFenceBypassKeys) + require.Equal(t, [][]byte{[]byte("a-key")}, g2Commit.WriteFenceBypassKeys) + require.Equal(t, pb.Phase_COMMIT, g1Commit.Phase) + require.Equal(t, pb.Phase_COMMIT, g2Commit.Phase) + require.Equal(t, []byte("z-key"), g1Commit.Mutations[1].Key) + require.Equal(t, pb.Op_PUT, g1Commit.Mutations[1].Op) + require.Equal(t, []byte("a-key"), g2Commit.Mutations[1].Key) + require.Equal(t, pb.Op_PUT, g2Commit.Mutations[1].Op) + + primaryCommitMeta := requestTxnMeta(t, g2Commit) + require.Equal(t, []byte("a-key"), primaryCommitMeta.PrimaryKey) + require.Greater(t, primaryCommitMeta.CommitTS, startTS) +} + +func TestShardedCoordinatorPinnedWritesBypassLogicalRouteFence(t *testing.T) { + t.Parallel() + + for _, isTxn := range []bool{false, true} { + t.Run(map[bool]string{false: "raw", true: "txn"}[isTxn], func(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + require.NoError(t, engine.ApplySnapshot(distribution.CatalogSnapshot{ + Version: 1, + Routes: []distribution.RouteDescriptor{ + {RouteID: 1, Start: []byte(""), End: []byte("m"), GroupID: 1, State: distribution.RouteStateActive}, + {RouteID: 2, Start: []byte("m"), End: nil, GroupID: 2, State: distribution.RouteStateWriteFenced}, + }, + })) + + g1Txn := &recordingTransactional{} + coord := NewShardedCoordinator(engine, map[uint64]*ShardGroup{ + 1: {Txn: g1Txn}, + 2: {Txn: &recordingTransactional{}}, + }, 1, NewHLC(), nil) + key := []byte("z-key") + reqs := &OperationGroup[OP]{ + IsTxn: isTxn, + Elems: []*Elem[OP]{{Op: Del, Key: key, GroupID: 1}}, + } + if isTxn { + reqs.StartTS = 10 + } + + _, err := coord.Dispatch(context.Background(), reqs) + require.NoError(t, err) + require.Len(t, g1Txn.requests, 1) + require.Equal(t, [][]byte{key}, g1Txn.requests[0].WriteFenceBypassKeys) + require.Equal(t, key, g1Txn.requests[0].Mutations[len(g1Txn.requests[0].Mutations)-1].Key) + }) + } +} + func requestTxnMeta(t *testing.T, req *pb.Request) TxnMeta { t.Helper() require.NotNil(t, req) @@ -614,6 +733,35 @@ func TestShardedCoordinatorDispatchTxn_SingleShardIncludesReadKeysInRaftEntry(t require.Equal(t, [][]byte{[]byte("rk1"), []byte("rk2")}, g1Txn.requests[0].ReadKeys) } +func TestShardedCoordinatorCommitPrimaryUsesPinnedMutationGroup(t *testing.T) { + t.Parallel() + + engine := distribution.NewEngine() + engine.UpdateRoute([]byte(""), nil, 2) + + g1Txn := &recordingTransactional{responses: []*TransactionResponse{{CommitIndex: 7}}} + g2Txn := &recordingTransactional{} + coord := NewShardedCoordinator(engine, map[uint64]*ShardGroup{ + 1: {Txn: g1Txn}, + 2: {Txn: g2Txn}, + }, 2, NewHLC(), nil) + + primaryKey := []byte("!lst|meta|d|pinned") + grouped := map[uint64][]*pb.Mutation{ + 1: {{Op: pb.Op_DEL, Key: primaryKey}}, + 2: {{Op: pb.Op_PUT, Key: []byte("z"), Value: []byte("v")}}, + } + primaryGid, commitIndex, err := coord.commitPrimaryTxn(context.Background(), 10, primaryKey, grouped, []uint64{1, 2}, 20, 0, nil) + require.NoError(t, err) + require.Equal(t, uint64(1), primaryGid) + require.Equal(t, uint64(7), commitIndex) + require.Len(t, g1Txn.requests, 1) + require.Empty(t, g2Txn.requests) + require.Equal(t, pb.Phase_COMMIT, g1Txn.requests[0].Phase) + require.Len(t, g1Txn.requests[0].Mutations, 2) + require.Equal(t, primaryKey, g1Txn.requests[0].Mutations[1].Key) +} + // TestShardedCoordinatorDispatchTxn_CrossShardPropagatesObservedRouteVersion // is the gemini-critical regression from PR #881. Contract: // every PREPARE and COMMIT envelope across the 2PC paths diff --git a/kv/transcoder.go b/kv/transcoder.go index 1bb4ac455..0b92491aa 100644 --- a/kv/transcoder.go +++ b/kv/transcoder.go @@ -23,6 +23,9 @@ type Elem[T OP] struct { // leader to stamp the resolved transaction commit timestamp into Value at // this byte offset before committing the mutation. CommitTSValueOffset uint64 + // GroupID optionally pins this mutation to a shard group. Zero preserves + // normal key-based routing. + GroupID uint64 } // OperationGroup is a group of operations that should be executed atomically. diff --git a/kv/txn_keys.go b/kv/txn_keys.go index e997722ea..9ddfb3c9f 100644 --- a/kv/txn_keys.go +++ b/kv/txn_keys.go @@ -16,6 +16,7 @@ const ( txnIntentPrefix = TxnKeyPrefix + "int|" txnCommitPrefix = TxnKeyPrefix + "cmt|" txnRollbackPrefix = TxnKeyPrefix + "rb|" + txnSuccessPrefix = TxnKeyPrefix + "ok|" txnMetaPrefix = TxnKeyPrefix + "meta|" ) @@ -27,11 +28,14 @@ var ( txnIntentPrefixBytes = []byte(txnIntentPrefix) txnCommitPrefixBytes = []byte(txnCommitPrefix) txnRollbackPrefixBytes = []byte(txnRollbackPrefix) + txnSuccessPrefixBytes = []byte(txnSuccessPrefix) txnMetaPrefixBytes = []byte(txnMetaPrefix) txnCommonPrefix = []byte(TxnKeyPrefix) ) const txnStartTSSuffixLen = 8 +const txnSuccessMarkerVersion = byte(1) +const maxIntValue = int(^uint(0) >> 1) func txnLockKey(userKey []byte) []byte { k := make([]byte, 0, len(txnLockPrefixBytes)+len(userKey)) @@ -75,6 +79,7 @@ func isTxnInternalKey(key []byte) bool { bytes.HasPrefix(key, txnIntentPrefixBytes) || bytes.HasPrefix(key, txnCommitPrefixBytes) || bytes.HasPrefix(key, txnRollbackPrefixBytes) || + bytes.HasPrefix(key, txnSuccessPrefixBytes) || bytes.HasPrefix(key, txnMetaPrefixBytes) } @@ -104,6 +109,8 @@ func txnRouteKey(key []byte) ([]byte, bool) { return nil, false } return rest[:len(rest)-txnStartTSSuffixLen], true + case bytes.HasPrefix(key, txnSuccessPrefixBytes): + return txnSuccessLockedKey(key) default: return nil, false } @@ -119,3 +126,55 @@ func ExtractTxnUserKey(key []byte) []byte { } return userKey } + +// TxnSuccessMarkerKey builds the route-local transaction-success marker key +// used by the migration planner. Normal transaction traffic only writes this +// after the migration capability gate opens in a later PR. +func TxnSuccessMarkerKey(lockedKey []byte, startTS, commitTS uint64, primaryKey []byte) []byte { + out := make([]byte, 0, len(txnSuccessPrefixBytes)+1+binary.MaxVarintLen64+len(lockedKey)+2*txnStartTSSuffixLen+binary.MaxVarintLen64+len(primaryKey)) + out = append(out, txnSuccessPrefixBytes...) + out = append(out, txnSuccessMarkerVersion) + out = binary.AppendUvarint(out, uint64(len(lockedKey))) + out = append(out, lockedKey...) + var raw [txnStartTSSuffixLen]byte + binary.BigEndian.PutUint64(raw[:], startTS) + out = append(out, raw[:]...) + binary.BigEndian.PutUint64(raw[:], commitTS) + out = append(out, raw[:]...) + out = binary.AppendUvarint(out, uint64(len(primaryKey))) + out = append(out, primaryKey...) + return out +} + +func txnSuccessLockedKey(key []byte) ([]byte, bool) { + rest := key[len(txnSuccessPrefixBytes):] + if len(rest) == 0 || rest[0] != txnSuccessMarkerVersion { + return nil, false + } + rest = rest[1:] + lockedLenRaw, n := binary.Uvarint(rest) + lockedLen, ok := uvarintToInt(lockedLenRaw) + if n <= 0 || !ok || lockedLen > len(rest)-n { + return nil, false + } + lockedStart := n + lockedEnd := lockedStart + lockedLen + rest = rest[lockedEnd:] + if len(rest) < 2*txnStartTSSuffixLen { + return nil, false + } + rest = rest[2*txnStartTSSuffixLen:] + primaryLenRaw, n := binary.Uvarint(rest) + primaryLen, ok := uvarintToInt(primaryLenRaw) + if n <= 0 || !ok || primaryLen != len(rest)-n { + return nil, false + } + return key[len(txnSuccessPrefixBytes)+1+lockedStart : len(txnSuccessPrefixBytes)+1+lockedEnd], true +} + +func uvarintToInt(v uint64) (int, bool) { + if v > uint64(maxIntValue) { + return 0, false + } + return int(v), true +} diff --git a/main.go b/main.go index c058f8651..d3d4412af 100644 --- a/main.go +++ b/main.go @@ -1859,6 +1859,9 @@ func startServersAfterStartupRotation(waitRotateOnStartup startupRotationWaiter, }) } publicKVGate := &startupPublicKVGate{} + if in.distServer != nil { + in.distServer.SetReadGate(publicKVGate.blocked) + } installHLCLeaseRenewalBlocker(in.coordinate, waitRotateOnStartup.BlockMutators) adapterCoordinate := startupGatedCoordinator{ inner: in.coordinate, diff --git a/monitoring/grafana/dashboards/elastickv-redis-summary.json b/monitoring/grafana/dashboards/elastickv-redis-summary.json index 3d05d6176..ddfce4f7c 100644 --- a/monitoring/grafana/dashboards/elastickv-redis-summary.json +++ b/monitoring/grafana/dashboards/elastickv-redis-summary.json @@ -1683,7 +1683,7 @@ ], "title": "Raft Queue Saturation (stepCh full / outbound drops / errors)", "type": "timeseries", - "description": "Counter rates from the etcd raft engine. step-queue-full means inbound messages from remote peers were dropped because the local raft loop was too slow to consume the selected step queue (the 'etcd raft inbound step queue is full' log line). dispatch-dropped means outbound messages were discarded before transport because the per-peer channel was full. dispatch-errors means transport delivery failed. The pre-#560 seek storm caused all three to spike together; watch for them to fall after the rollout and stay flat." + "description": "Counter rates from the etcd raft engine. step-queue-full means inbound messages from remote peers found the selected step queue full; blocking replication messages now wait for space, while best-effort message classes may still be rejected. dispatch-dropped means outbound messages were discarded before transport because the per-peer channel was full. dispatch-errors means transport delivery failed. The pre-#560 seek storm caused all three to spike together; watch for them to fall after the rollout and stay flat." }, { "datasource": "$datasource", diff --git a/monitoring/hotpath.go b/monitoring/hotpath.go index cd56a3587..04e5a7776 100644 --- a/monitoring/hotpath.go +++ b/monitoring/hotpath.go @@ -103,7 +103,7 @@ func newHotPathMetrics(registerer prometheus.Registerer) *HotPathMetrics { stepQueueFullTotal: prometheus.NewCounterVec( prometheus.CounterOpts{ Name: "elastickv_raft_step_queue_full_total", - Help: "Inbound raft messages that could not be enqueued because the selected step queue was full; indicates the raft loop is starved (classic pre-#560 seek-storm symptom).", + Help: "Inbound raft messages that found the selected step queue full. Blocking replication messages wait for space; best-effort messages may still be rejected. Indicates the raft loop is starved.", }, []string{"group"}, ), diff --git a/monitoring/hotpath_test.go b/monitoring/hotpath_test.go index ff843033e..fb8dc9904 100644 --- a/monitoring/hotpath_test.go +++ b/monitoring/hotpath_test.go @@ -168,7 +168,7 @@ elastickv_raft_dispatch_dropped_total{group="1",node_address="10.0.0.1:50051",no # HELP elastickv_raft_dispatch_errors_total Outbound raft dispatches that reached the transport but failed. Mirrors etcd raft Engine.dispatchErrorCount. # TYPE elastickv_raft_dispatch_errors_total counter elastickv_raft_dispatch_errors_total{group="1",node_address="10.0.0.1:50051",node_id="n1"} 2 -# HELP elastickv_raft_step_queue_full_total Inbound raft messages that could not be enqueued because the selected step queue was full; indicates the raft loop is starved (classic pre-#560 seek-storm symptom). +# HELP elastickv_raft_step_queue_full_total Inbound raft messages that found the selected step queue full. Blocking replication messages wait for space; best-effort messages may still be rejected. Indicates the raft loop is starved. # TYPE elastickv_raft_step_queue_full_total counter elastickv_raft_step_queue_full_total{group="1",node_address="10.0.0.1:50051",node_id="n1"} 1 # HELP elastickv_raft_send_stream_opens_total Successful outbound Raft SendStream opens, including reconnects. diff --git a/proto/internal.pb.go b/proto/internal.pb.go index ef52628ca..dabe9d755 100644 --- a/proto/internal.pb.go +++ b/proto/internal.pb.go @@ -221,6 +221,12 @@ type Request struct { // is plumbing only — the FSM ignores the value, so all existing // callers see no behaviour change. ObservedRouteVersion uint64 `protobuf:"varint,6,opt,name=observed_route_version,json=observedRouteVersion,proto3" json:"observed_route_version,omitempty"` + // write_fence_bypass_keys carries point keys whose target group was resolved + // outside the byte-route catalog, either by a partition resolver or an + // explicit maintenance pin. The FSM skips byte-route write-fence and owner + // checks only for those exact point mutations; prefixes and unmarked keys + // still fail closed. + WriteFenceBypassKeys [][]byte `protobuf:"bytes,7,rep,name=write_fence_bypass_keys,json=writeFenceBypassKeys,proto3" json:"write_fence_bypass_keys,omitempty"` unknownFields protoimpl.UnknownFields sizeCache protoimpl.SizeCache } @@ -297,6 +303,13 @@ func (x *Request) GetObservedRouteVersion() uint64 { return 0 } +func (x *Request) GetWriteFenceBypassKeys() [][]byte { + if x != nil { + return x.WriteFenceBypassKeys + } + return nil +} + type RaftCommand struct { state protoimpl.MessageState `protogen:"open.v1"` Requests []*Request `protobuf:"bytes,1,rep,name=requests,proto3" json:"requests,omitempty"` @@ -931,14 +944,15 @@ const file_internal_proto_rawDesc = "" + "\x02op\x18\x01 \x01(\x0e2\x03.OpR\x02op\x12\x10\n" + "\x03key\x18\x02 \x01(\fR\x03key\x12\x14\n" + "\x05value\x18\x03 \x01(\fR\x05value\x123\n" + - "\x16commit_ts_value_offset\x18\x04 \x01(\x04R\x13commitTsValueOffset\"\xca\x01\n" + + "\x16commit_ts_value_offset\x18\x04 \x01(\x04R\x13commitTsValueOffset\"\x81\x02\n" + "\aRequest\x12\x15\n" + "\x06is_txn\x18\x01 \x01(\bR\x05isTxn\x12\x1c\n" + "\x05phase\x18\x02 \x01(\x0e2\x06.PhaseR\x05phase\x12\x0e\n" + "\x02ts\x18\x03 \x01(\x04R\x02ts\x12'\n" + "\tmutations\x18\x04 \x03(\v2\t.MutationR\tmutations\x12\x1b\n" + "\tread_keys\x18\x05 \x03(\fR\breadKeys\x124\n" + - "\x16observed_route_version\x18\x06 \x01(\x04R\x14observedRouteVersion\"3\n" + + "\x16observed_route_version\x18\x06 \x01(\x04R\x14observedRouteVersion\x125\n" + + "\x17write_fence_bypass_keys\x18\a \x03(\fR\x14writeFenceBypassKeys\"3\n" + "\vRaftCommand\x12$\n" + "\brequests\x18\x01 \x03(\v2\b.RequestR\brequests\"M\n" + "\x0eForwardRequest\x12\x15\n" + diff --git a/proto/internal.proto b/proto/internal.proto index 7af3e5fd7..27b45bd0d 100644 --- a/proto/internal.proto +++ b/proto/internal.proto @@ -61,6 +61,12 @@ message Request { // is plumbing only — the FSM ignores the value, so all existing // callers see no behaviour change. uint64 observed_route_version = 6; + // write_fence_bypass_keys carries point keys whose target group was resolved + // outside the byte-route catalog, either by a partition resolver or an + // explicit maintenance pin. The FSM skips byte-route write-fence and owner + // checks only for those exact point mutations; prefixes and unmarked keys + // still fail closed. + repeated bytes write_fence_bypass_keys = 7; } message RaftCommand { diff --git a/proxy/blocking.go b/proxy/blocking.go index be9d7e5d2..973240f96 100644 --- a/proxy/blocking.go +++ b/proxy/blocking.go @@ -1,7 +1,9 @@ package proxy import ( + "bytes" "context" + "fmt" "strconv" "strings" "time" @@ -9,7 +11,43 @@ import ( "github.com/redis/go-redis/v9" ) -const blockingMultiPopMinArgs = 2 +const ( + blockingMultiPopMinArgs = 2 + blockingBLMPopMinArgs = 5 + blockingBLMPopNumKeysArgIndex = 2 + blockingBLMPopFirstKeyArgIndex = 3 + blockingListMoveMinArgs = 3 + blockingBLMoveArgs = 6 + blockingListPopReplayKeyCount = int64(1) + blockingListMoveReplayKeyCount = int64(2) +) + +const ( + blockingListSideLeft = "LEFT" + blockingListSideRight = "RIGHT" +) + +const blockingListMoveReplayScript = ` +local removed = redis.call("LREM", KEYS[1], tonumber(ARGV[1]), ARGV[3]) +if removed == 0 then + return 0 +end +if ARGV[2] == "LEFT" then + return redis.call("LPUSH", KEYS[2], ARGV[3]) +end +return redis.call("RPUSH", KEYS[2], ARGV[3]) +` + +const blockingListMultiPopReplayScript = ` +local removed = 0 +local count = tonumber(ARGV[1]) +for i = 2, #ARGV do + if redis.call("LREM", KEYS[1], count, ARGV[i]) > 0 then + removed = removed + 1 + end +end +return removed +` type blockingTimeoutBackend interface { DoWithTimeout(ctx context.Context, timeout time.Duration, args ...any) *redis.Cmd @@ -27,7 +65,7 @@ func blockingCommandTimeout(cmd string, args [][]byte) time.Duration { return 0 } return parseBlockingSecondsArg(args[1]) - case "XREAD", "XREADGROUP": + case "XREAD", cmdNameXREADGROUP: for i := 1; i+1 < len(args); i++ { if strings.EqualFold(string(args[i]), "BLOCK") { return parseBlockingMillisecondsArg(args[i+1]) @@ -53,27 +91,233 @@ func parseBlockingMillisecondsArg(raw []byte) time.Duration { return time.Duration(millis) * time.Millisecond } -func shouldReplayBlockingToSecondary(cmd string) bool { - return !strings.EqualFold(cmd, "XREAD") -} - -func secondaryBlockingReplay(cmd string, resp any) (string, []any, bool) { +func blockingReplayCommand(cmd string, args [][]byte, resp any) (string, []any, bool) { switch strings.ToUpper(cmd) { + case "BLPOP": + return blockingListPopReplay(1, resp) + case "BRPOP": + return blockingListPopReplay(-1, resp) + case "BRPOPLPUSH": + return blockingListMoveReplay(args, resp, -1, blockingListSideLeft) + case "BLMOVE": + return blockingBLMoveReplay(args, resp) + case "BLMPOP": + return blockingBLMPopReplay(args, resp) case "BZPOPMIN", "BZPOPMAX": - key, member, ok := zsetPopKeyMember(resp) + return blockingZSetPopReplay(resp) + case cmdNameXREADGROUP: + return blockingXReadGroupReplay(args, resp) + default: + return "", nil, false + } +} + +func blockingXReadGroupReplay(args [][]byte, resp any) (string, []any, bool) { + if !xreadGroupResponseHasEntries(resp) { + return "", nil, false + } + return cmdNameXREADGROUP, bytesArgsToInterfaces(args), true +} + +func xreadGroupResponseHasEntries(resp any) bool { + streams, ok := redisArray(resp) + if !ok { + return false + } + for _, stream := range streams { + parts, ok := redisArray(stream) + if !ok || len(parts) < 2 { + continue + } + entries, ok := redisArray(parts[1]) + if ok && len(entries) > 0 { + return true + } + } + return false +} + +func blockingListPopReplay(count int64, resp any) (string, []any, bool) { + parts, ok := redisArray(resp) + if !ok || len(parts) < 2 { + return "", nil, false + } + key, keyOK := redisArg(parts[0]) + value, valueOK := redisArg(parts[1]) + if !keyOK || !valueOK { + return "", nil, false + } + return "LREM", []any{[]byte("LREM"), key, count, value}, true +} + +func blockingZSetPopReplay(resp any) (string, []any, bool) { + parts, ok := redisArray(resp) + if !ok || len(parts) < blockingMultiPopMinArgs { + return "", nil, false + } + key, keyOK := redisArg(parts[0]) + member, memberOK := redisArg(parts[1]) + if !keyOK || !memberOK { + return "", nil, false + } + return "ZREM", []any{[]byte("ZREM"), key, member}, true +} + +func blockingBLMoveReplay(args [][]byte, resp any) (string, []any, bool) { + if len(args) < blockingBLMoveArgs { + return "", nil, false + } + count, ok := blockingListPopCount(args[3]) + if !ok { + return "", nil, false + } + to := strings.ToUpper(string(args[4])) + if to != blockingListSideLeft && to != blockingListSideRight { + return "", nil, false + } + return blockingListMoveReplay(args, resp, count, to) +} + +func blockingBLMPopReplay(args [][]byte, resp any) (string, []any, bool) { + numKeys, count, ok := parseBlockingBLMPopArgs(args) + if !ok { + return "", nil, false + } + key, values, ok := blockingBLMPopResponse(resp) + if !ok || !blockingBLMPopKeyListed(args, numKeys, key) { + return "", nil, false + } + + replay := []any{ + []byte(cmdEval), + blockingListMultiPopReplayScript, + blockingListPopReplayKeyCount, + key, + count, + } + for _, value := range values { + arg, ok := redisArg(value) if !ok { return "", nil, false } - return "ZREM", []any{"ZREM", key, member}, true + replay = append(replay, arg) + } + return cmdEval, replay, true +} + +func parseBlockingBLMPopArgs(args [][]byte) (int, int64, bool) { + if len(args) < blockingBLMPopMinArgs { + return 0, 0, false + } + numKeys, err := strconv.Atoi(string(args[blockingBLMPopNumKeysArgIndex])) + if err != nil || numKeys <= 0 { + return 0, 0, false + } + directionArgIndex := blockingBLMPopFirstKeyArgIndex + numKeys + if directionArgIndex >= len(args) { + return 0, 0, false + } + count, ok := blockingListPopCount(args[directionArgIndex]) + return numKeys, count, ok +} + +func blockingBLMPopResponse(resp any) ([]byte, []any, bool) { + parts, ok := redisArray(resp) + if !ok || len(parts) != blockingMultiPopMinArgs { + return nil, nil, false + } + key, ok := redisArgBytes(parts[0]) + if !ok { + return nil, nil, false + } + values, ok := redisArray(parts[1]) + if !ok || len(values) == 0 { + return nil, nil, false + } + return key, values, true +} + +func blockingBLMPopKeyListed(args [][]byte, numKeys int, key []byte) bool { + for i := 0; i < numKeys; i++ { + if bytes.Equal(args[blockingBLMPopFirstKeyArgIndex+i], key) { + return true + } + } + return false +} + +func blockingListPopCount(side []byte) (int64, bool) { + switch strings.ToUpper(string(side)) { + case blockingListSideLeft: + return 1, true + case blockingListSideRight: + return -1, true default: + return 0, false + } +} + +func blockingListMoveReplay(args [][]byte, resp any, count int64, to string) (string, []any, bool) { + if len(args) < blockingListMoveMinArgs { return "", nil, false } + value, ok := redisArg(resp) + if !ok { + return "", nil, false + } + source := append([]byte(nil), args[1]...) + destination := append([]byte(nil), args[2]...) + return cmdEval, []any{ + []byte(cmdEval), + blockingListMoveReplayScript, + blockingListMoveReplayKeyCount, + source, + destination, + count, + []byte(to), + value, + }, true } -func zsetPopKeyMember(resp any) (any, any, bool) { - arr, ok := resp.([]any) - if !ok || len(arr) < 2 || arr[0] == nil || arr[1] == nil { - return nil, nil, false +func redisArray(v any) ([]any, bool) { + switch x := v.(type) { + case []any: + return x, true + case []string: + out := make([]any, len(x)) + for i := range x { + out[i] = x[i] + } + return out, true + case [][]byte: + out := make([]any, len(x)) + for i := range x { + out[i] = x[i] + } + return out, true + default: + return nil, false + } +} + +func redisArgBytes(v any) ([]byte, bool) { + arg, ok := redisArg(v) + if !ok { + return nil, false + } + b, ok := arg.([]byte) + return b, ok +} + +func redisArg(v any) (any, bool) { + switch x := v.(type) { + case nil: + return nil, false + case []byte: + return append([]byte(nil), x...), true + case string: + return []byte(x), true + default: + return []byte(fmt.Sprint(x)), true } - return arr[0], arr[1], true } diff --git a/proxy/command.go b/proxy/command.go index d3f7abb85..4493ca68c 100644 --- a/proxy/command.go +++ b/proxy/command.go @@ -16,8 +16,9 @@ const ( ) const ( - cmdNameAUTH = "AUTH" - cmdNameSELECT = "SELECT" + cmdNameAUTH = "AUTH" + cmdNameSELECT = "SELECT" + cmdNameXREADGROUP = "XREADGROUP" ) var commandTable = map[string]CommandCategory{ @@ -243,7 +244,7 @@ func ClassifyCommand(name string, args [][]byte) CommandCategory { upper := strings.ToUpper(name) // Special case: XREAD/XREADGROUP with BLOCK - if upper == "XREAD" || upper == "XREADGROUP" { + if upper == "XREAD" || upper == cmdNameXREADGROUP { for _, arg := range args { if strings.ToUpper(string(arg)) == "BLOCK" { return CmdBlocking diff --git a/proxy/config.go b/proxy/config.go index e99229ff9..771dfa7e2 100644 --- a/proxy/config.go +++ b/proxy/config.go @@ -83,7 +83,9 @@ type ProxyConfig struct { SentryEnv string SentrySampleRate float64 MetricsAddr string + PProfAddr string PubSubCompareWindow time.Duration + RedisOnlyRaw bool } // DefaultConfig returns a ProxyConfig with sensible defaults. @@ -99,5 +101,6 @@ func DefaultConfig() ProxyConfig { SentrySampleRate: 1.0, MetricsAddr: ":9191", PubSubCompareWindow: defaultPubSubCompareWindow, + RedisOnlyRaw: true, } } diff --git a/proxy/dualwrite.go b/proxy/dualwrite.go index dc1e16bf7..88f0cd0f9 100644 --- a/proxy/dualwrite.go +++ b/proxy/dualwrite.go @@ -25,9 +25,8 @@ const ( // (EVAL / EVALSHA). Lua scripts under high load cause write conflicts in the Raft // layer, and each conflict triggers a full script re-execution. Capping the // concurrency reduces contention so individual scripts complete within - // SecondaryTimeout. Excess secondary script writes may be dropped to keep - // contention bounded; this is only tolerable in modes where the script write - // is targeting the non-authoritative backend. + // SecondaryTimeout. Strict dual-write script replays wait for capacity instead + // of being dropped; best-effort users of goScript may still drop. maxScriptWriteGoroutines = 64 // maxBlockingReplayGoroutines isolates mutating blocking command replays // from normal secondary writes. Blocking replays may wait for the secondary @@ -40,11 +39,15 @@ const ( maxAsyncQueueCapacity = 8192 asyncQueueConcurrencyFactor = 64 + // maxSecondaryTransientRetries caps proxy-level retries when the secondary + // returns a transient OCC/read-snapshot error after exhausting its own retry + // loop. SecondaryTimeout still bounds the whole replay. + maxSecondaryTransientRetries = 3 // maxCompactedRetries caps retries when the secondary returns // "read timestamp has been compacted". Each attempt re-sends the command so // the secondary re-selects a fresh read snapshot; a small bound is enough // because the compaction waterline advances slowly relative to SecondaryTimeout. - maxCompactedRetries = 3 + maxCompactedRetries = maxSecondaryTransientRetries // maxServerOverloadedRetries caps retries when ElasticKV rejects a secondary // replay before execution because the heavy-command worker pool is full. // Retrying holds the proxy's secondary budget, so a bounded loop turns @@ -88,6 +91,10 @@ const serverOverloadedMarker = "BUSY server overloaded" var errSecondaryReplayNoEffect = errors.New("secondary replay produced no effect") +type leaderRefreshingBackend interface { + RefreshLeaderNow(context.Context) +} + func isReadTSCompactedError(err error) bool { if err == nil { return false @@ -95,6 +102,29 @@ func isReadTSCompactedError(err error) bool { return strings.Contains(err.Error(), readTSCompactedMarker) } +func isElasticKVNotLeaderError(err error) bool { + if err == nil { + return false + } + msg := strings.TrimSpace(err.Error()) + upper := strings.ToUpper(msg) + if upper == "NOTLEADER" || strings.HasPrefix(upper, "NOTLEADER ") { + return true + } + var redisErr redis.Error + if errors.As(err, &redisErr) { + return false + } + if msg == "etcd raft engine is not leader" || + msg == "raft engine: not leader" { + return true + } + return msg == "leader not found" || + strings.HasSuffix(msg, "desc = leader not found") || + strings.HasSuffix(msg, "desc = raft engine: not leader") || + strings.HasSuffix(msg, "desc = etcd raft engine is not leader") +} + func isServerOverloadedError(err error) bool { if err == nil { return false @@ -102,6 +132,15 @@ func isServerOverloadedError(err error) bool { return strings.Contains(err.Error(), serverOverloadedMarker) } +func refreshSecondaryLeader(ctx context.Context, backend Backend, err error) { + if !isElasticKVNotLeaderError(err) { + return + } + if refresher, ok := backend.(leaderRefreshingBackend); ok { + refresher.RefreshLeaderNow(ctx) + } +} + // DualWriter routes commands to primary and secondary backends based on mode. type DualWriter struct { primary Backend @@ -298,7 +337,6 @@ func (d *DualWriter) Write(ctx context.Context, cmd string, args [][]byte) (any, } d.metrics.CommandTotal.WithLabelValues(cmd, d.primary.Name(), "ok").Inc() - // Secondary: async fire-and-forget (bounded) if d.hasSecondaryWrite() { d.goWrite(func(ctx context.Context) { d.writeSecondary(ctx, cmd, iArgs) }) } @@ -361,20 +399,18 @@ func (d *DualWriter) Blocking(ctx context.Context, cmd string, args [][]byte) (a d.metrics.CommandTotal.WithLabelValues(cmd, d.primary.Name(), "ok").Inc() if d.hasSecondaryWrite() { - if replayCmd, replayArgs, ok := secondaryBlockingReplay(cmd, resp); ok { + if replayCmd, replayArgs, ok := blockingReplayCommand(cmd, args, resp); ok { d.goBlockingReplay(func(ctx context.Context) { + if strings.EqualFold(replayCmd, cmdNameXREADGROUP) { + d.writeSecondary(ctx, replayCmd, replayArgs) + return + } d.writeSecondaryPositiveIntWithOptions(ctx, replayCmd, replayArgs, positiveIntReplayOptions{ initialDelay: blockingReplayInitialDelay, noEffectRetryWindow: blockingReplayNoEffectRetryWindow, deadlineAsMiss: true, }) }) - } else if shouldReplayBlockingToSecondary(cmd) { - d.goBlockingReplay(func(ctx context.Context) { - sCtx, cancel := context.WithTimeout(ctx, time.Second) - defer cancel() - d.secondary.Do(sCtx, iArgs...) - }) } } @@ -399,7 +435,9 @@ func (d *DualWriter) Admin(ctx context.Context, cmd string, args [][]byte) (any, return resp, err //nolint:wrapcheck // redis.Nil must pass through unwrapped for callers to detect nil replies } -// Script forwards EVAL/EVALSHA to the primary, and async replays to secondary. +// Script forwards EVAL/EVALSHA to the primary, and replays to secondary. +// Secondary script replays are concurrency-limited and apply caller +// backpressure rather than dropping when the script slot pool is saturated. // cmd must be the pre-uppercased command name. func (d *DualWriter) Script(ctx context.Context, cmd string, args [][]byte) (any, error) { iArgs := bytesArgsToInterfaces(args) @@ -424,10 +462,9 @@ func (d *DualWriter) Script(ctx context.Context, cmd string, args [][]byte) (any } // writeSecondary sends the command to the secondary, handling the NOSCRIPT -// → EVAL fallback and transparently retrying when the secondary reports that -// the read snapshot has been compacted. A re-sent command causes the backend -// to re-select a fresh read timestamp, which is the only way to recover once -// the original startTS has fallen behind MinRetainedTS on a peer node. +// → EVAL fallback and transparently retrying transient secondary errors. A +// re-sent command causes the backend to re-select a fresh timestamp and can +// also recover from hot-key OCC retry exhaustion in the secondary Redis adapter. // // The secondary's raw redis error is kept in sErr (not wrapped) so that // writeSecondary can classify it via errors.Is(sErr, redis.Nil), attach the @@ -463,6 +500,7 @@ func (d *DualWriter) writeSecondary(sCtx context.Context, cmd string, iArgs []an if attempt >= retryLimit { break } + refreshSecondaryLeader(sCtx, d.secondary, sErr) if shouldLogSecondaryRetry(sErr) { d.logger.Debug("retrying secondary write", "cmd", cmd, "reason", retryReason, "attempt", attempt+1, "backoff", backoff, "err", sErr) @@ -755,6 +793,12 @@ func (d *DualWriter) recordSecondaryWriteFailure(cmd string, iArgs []any, elapse d.logger.Warn("secondary write failed", warnArgs...) } +func (d *DualWriter) replaySecondaryPipeline(cmds [][]any) { + d.goWrite(func(ctx context.Context) { + d.writeSecondaryPipeline(ctx, cmds) + }) +} + // waitCompactedRetryBackoff sleeps for a jittered interval or returns early // when the context is cancelled. Returns false if the caller should abort // the retry loop (context done). @@ -807,6 +851,10 @@ func secondaryRetryReasonAndLimit(cmd string, err error) (string, int) { switch { case isReadTSCompactedError(err): return "compacted_snapshot", maxCompactedRetries + case isElasticKVNotLeaderError(err): + return "not_leader", maxSecondaryTransientRetries + case isRetryableSecondaryConflictError(err): + return classifySecondaryWriteError(err), maxSecondaryTransientRetries case isServerOverloadedError(err) && !isRedisScriptCommandName(cmd): return "server_overloaded", maxServerOverloadedRetries default: @@ -814,15 +862,24 @@ func secondaryRetryReasonAndLimit(cmd string, err error) (string, int) { } } +func isRetryableSecondaryConflictError(err error) bool { + switch classifySecondaryWriteError(err) { + case "retry_limit", "write_conflict", "txn_locked": + return true + default: + return false + } +} + // goWrite queues fn for bounded secondary execution. -func (d *DualWriter) goWrite(fn func(context.Context)) { - d.enqueueAsync(d.writeQueue, d.writeQueueSlots, asyncQueueWrite, fn) +func (d *DualWriter) goWrite(fn any) { + d.enqueueAsync(d.writeQueue, d.writeQueueSlots, asyncQueueWrite, normalizeAsyncFunc(fn)) } // goScript launches fn in a bounded Lua-script write goroutine. // It uses a smaller class limit while also consuming the shared write limit. -func (d *DualWriter) goScript(fn func(context.Context)) { - d.enqueueAsync(d.scriptQueue, d.scriptQueueSlots, asyncQueueScript, fn) +func (d *DualWriter) goScript(fn any) { + d.enqueueAsync(d.scriptQueue, d.scriptQueueSlots, asyncQueueScript, normalizeAsyncFunc(fn)) } // goBlockingReplay queues fn for bounded secondary replay of mutating blocking @@ -834,6 +891,17 @@ func (d *DualWriter) goBlockingReplay(fn func(context.Context)) { d.enqueueAsync(d.blockingReplayQueue, d.blockingReplayQueueSlots, asyncQueueBlocking, fn) } +func normalizeAsyncFunc(fn any) func(context.Context) { + switch f := fn.(type) { + case func(context.Context): + return f + case func(): + return func(context.Context) { f() } + default: + return func(context.Context) {} + } +} + // goShadow launches fn in a bounded shadow-read goroutine. func (d *DualWriter) goShadow(fn func()) { d.goShadowWithSem(fn) diff --git a/proxy/leader_aware_backend.go b/proxy/leader_aware_backend.go index 05bd6e6b0..db9e82f93 100644 --- a/proxy/leader_aware_backend.go +++ b/proxy/leader_aware_backend.go @@ -180,6 +180,13 @@ func (b *LeaderAwareRedisBackend) TriggerRefresh() { } } +// RefreshLeaderNow re-probes the cluster before returning. It is used by +// callers that already observed a not-leader response and need the next retry +// to use a fresh target instead of waiting for the background loop. +func (b *LeaderAwareRedisBackend) RefreshLeaderNow(ctx context.Context) { + b.refreshLeader(ctx) +} + // refreshLeader probes INFO replication on the current leader first, then on // each seed, and adopts the first advertised leader address. The current // leader's Redis address is returned by the leader node itself when it's @@ -238,12 +245,6 @@ func (b *LeaderAwareRedisBackend) refreshLeaderOnce(ctx context.Context) { } } -// RefreshLeaderNow synchronously re-probes the cluster. Callers use this only -// after an explicit not-leader rejection, where retrying is known to be safe. -func (b *LeaderAwareRedisBackend) RefreshLeaderNow(ctx context.Context) { - b.refreshLeader(ctx) -} - func (b *LeaderAwareRedisBackend) probeLeader(ctx context.Context, addr string) (string, error) { cli := b.getOrCreateClient(addr) if cli == nil { @@ -414,6 +415,8 @@ func (b *LeaderAwareRedisBackend) doOnce(ctx context.Context, args ...any) *redi } // DoWithTimeout forwards a blocking command with a per-call socket timeout. +// Like Do, a not-leader rejection refreshes the cached leader for the next +// command without replaying the current command. func (b *LeaderAwareRedisBackend) DoWithTimeout(ctx context.Context, timeout time.Duration, args ...any) *redis.Cmd { cmd := b.doWithTimeoutOnce(ctx, timeout, args...) switch { @@ -435,8 +438,13 @@ func (b *LeaderAwareRedisBackend) doWithTimeoutOnce(ctx context.Context, timeout return cli.WithTimeout(effectiveBlockingReadTimeout(timeout)).Do(ctx, args...) } -// Pipeline forwards a batch to the current leader. +// Pipeline forwards a batch to the current leader. NOTLEADER refreshes the +// cached leader for the next command without replaying the current batch. func (b *LeaderAwareRedisBackend) Pipeline(ctx context.Context, cmds [][]any) ([]*redis.Cmd, error) { + return b.pipelineOnce(ctx, cmds) +} + +func (b *LeaderAwareRedisBackend) pipelineOnce(ctx context.Context, cmds [][]any) ([]*redis.Cmd, error) { cli := b.currentClient() if cli == nil { return nil, ErrNoLeaderBackend @@ -452,7 +460,7 @@ func (b *LeaderAwareRedisBackend) Pipeline(ctx context.Context, cmds [][]any) ([ if errors.As(err, &redisErr) || errors.Is(err, redis.Nil) { for _, result := range results { if isElasticKVNotLeaderError(result.Err()) { - b.TriggerRefresh() + b.RefreshLeaderNow(ctx) break } } @@ -475,35 +483,6 @@ func isRedisScriptCommandName(name string) bool { } } -func isElasticKVNotLeaderError(err error) bool { - if err == nil { - return false - } - - msg := strings.TrimSpace(err.Error()) - upper := strings.ToUpper(msg) - if upper == "NOTLEADER" || strings.HasPrefix(upper, "NOTLEADER ") { - return true - } - - // Redis application errors without the NOTLEADER code are not safe to - // replay. Their text may contain a leader phrase supplied by a script or - // command, but the command may already have changed state. - var redisErr redis.Error - if errors.As(err, &redisErr) { - return false - } - - // Keep a closed set for leadership errors that may reach this backend - // before Redis protocol framing (for example through a gRPC wrapper). - return msg == "etcd raft engine is not leader" || - msg == "raft engine: not leader" || - msg == "leader not found" || - strings.HasSuffix(msg, "desc = leader not found") || - strings.HasSuffix(msg, "desc = raft engine: not leader") || - strings.HasSuffix(msg, "desc = etcd raft engine is not leader") -} - func isLeaderRefreshTransportError(err error) bool { if err == nil || errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { return false diff --git a/proxy/leader_aware_backend_test.go b/proxy/leader_aware_backend_test.go index c42cfbeed..0daa899af 100644 --- a/proxy/leader_aware_backend_test.go +++ b/proxy/leader_aware_backend_test.go @@ -3,7 +3,6 @@ package proxy import ( "bufio" "context" - "errors" "fmt" "io" "net" @@ -63,11 +62,9 @@ type fakeElasticKVNode struct { addr string ln net.Listener leaderAddr atomic.Pointer[string] + commandErr atomic.Pointer[string] commands atomic.Int64 infoCalls atomic.Int64 - scriptErr atomic.Bool - infoGateMu sync.RWMutex - infoGate <-chan struct{} } func newFakeElasticKVNode(t *testing.T) *fakeElasticKVNode { @@ -93,23 +90,8 @@ func (n *fakeElasticKVNode) Leader() string { return "" } -func (n *fakeElasticKVNode) SetScriptNotLeaderError(enabled bool) { - n.scriptErr.Store(enabled) -} - -func (n *fakeElasticKVNode) SetInfoGate(gate <-chan struct{}) { - n.infoGateMu.Lock() - n.infoGate = gate - n.infoGateMu.Unlock() -} - -func (n *fakeElasticKVNode) waitInfoGate() { - n.infoGateMu.RLock() - gate := n.infoGate - n.infoGateMu.RUnlock() - if gate != nil { - <-gate - } +func (n *fakeElasticKVNode) SetCommandError(err string) { + n.commandErr.Store(&err) } func (n *fakeElasticKVNode) serve() { @@ -133,36 +115,38 @@ func (n *fakeElasticKVNode) handleConn(conn net.Conn) { if len(args) == 0 { return } - cmd := strings.ToUpper(args[0]) - switch cmd { - case "HELLO": - // Force go-redis to fall back to RESP2 without HELLO: return an - // error the client treats as "HELLO not supported". - _, _ = conn.Write([]byte("-ERR unknown command 'HELLO'\r\n")) - case "CLIENT", "AUTH", "SELECT", "PING": - _, _ = conn.Write([]byte("+OK\r\n")) - case "INFO": - n.infoCalls.Add(1) - n.waitInfoGate() - body := fmt.Sprintf( - "# Replication\r\nrole:slave\r\nraft_leader_redis:%s\r\n", - n.Leader(), - ) - _, _ = fmt.Fprintf(conn, "$%d\r\n%s\r\n", len(body), body) - default: - n.writeCommandResponse(conn, cmd) - } + n.handleCommand(conn, args) } } -func (n *fakeElasticKVNode) writeCommandResponse(conn net.Conn, cmd string) { +func (n *fakeElasticKVNode) handleCommand(conn net.Conn, args []string) { + switch strings.ToUpper(args[0]) { + case "HELLO": + // Force go-redis to fall back to RESP2 without HELLO: return an + // error the client treats as "HELLO not supported". + _, _ = conn.Write([]byte("-ERR unknown command 'HELLO'\r\n")) + case "CLIENT", "AUTH", "SELECT", "PING": + _, _ = conn.Write([]byte("+OK\r\n")) + case "INFO": + n.infoCalls.Add(1) + body := fmt.Sprintf( + "# Replication\r\nrole:slave\r\nraft_leader_redis:%s\r\n", + n.Leader(), + ) + _, _ = fmt.Fprintf(conn, "$%d\r\n%s\r\n", len(body), body) + default: + n.writeCommandReply(conn) + } +} + +func (n *fakeElasticKVNode) writeCommandReply(conn net.Conn) { n.commands.Add(1) - if leader := n.Leader(); leader != "" && leader != n.addr { - _, _ = conn.Write([]byte("-NOTLEADER etcd raft engine is not leader\r\n")) + if p := n.commandErr.Load(); p != nil { + _, _ = fmt.Fprintf(conn, "-%s\r\n", *p) return } - if n.scriptErr.Load() && isRedisScriptCommandName(cmd) { - _, _ = conn.Write([]byte("-NOTLEADER user script\r\n")) + if leader := n.Leader(); leader != "" && leader != n.addr { + _, _ = conn.Write([]byte("-NOTLEADER etcd raft engine is not leader\r\n")) return } _, _ = conn.Write([]byte("+OK\r\n")) @@ -253,6 +237,7 @@ func TestLeaderAwareRedisBackend_FollowsLeaderChange(t *testing.T) { func TestLeaderAwareRedisBackend_NotLeaderRefreshesWithoutReplay(t *testing.T) { nodeA := newFakeElasticKVNode(t) nodeB := newFakeElasticKVNode(t) + nodeA.SetLeader(nodeA.addr) nodeB.SetLeader(nodeA.addr) @@ -264,9 +249,10 @@ func TestLeaderAwareRedisBackend_NotLeaderRefreshesWithoutReplay(t *testing.T) { testLogger, ) t.Cleanup(func() { _ = backend.Close() }) + require.Eventually(t, func() bool { return backend.CurrentLeader() == nodeA.addr - }, 2*time.Second, 10*time.Millisecond) + }, 2*time.Second, 10*time.Millisecond, "initial leader must be A") nodeA.SetLeader(nodeB.addr) nodeB.SetLeader(nodeB.addr) @@ -274,44 +260,19 @@ func TestLeaderAwareRedisBackend_NotLeaderRefreshesWithoutReplay(t *testing.T) { res := backend.Do(context.Background(), "SET", "k", "v") require.Error(t, res.Err()) require.Contains(t, res.Err().Error(), "NOTLEADER") - require.Equal(t, nodeB.addr, backend.CurrentLeader()) - require.Equal(t, int64(1), nodeA.commands.Load(), "first attempt must be rejected by the former leader") + require.Equal(t, nodeB.addr, backend.CurrentLeader(), "not-leader must synchronously adopt the advertised leader") + require.Equal(t, int64(1), nodeA.commands.Load(), "first attempt should hit the former leader and be rejected") require.Equal(t, int64(0), nodeB.commands.Load(), "ambiguous writes must not be replayed") -} - -func TestLeaderAwareRedisBackend_ScriptNotLeaderRefreshesWithoutReplay(t *testing.T) { - nodeA := newFakeElasticKVNode(t) - nodeB := newFakeElasticKVNode(t) - nodeA.SetLeader(nodeA.addr) - nodeB.SetLeader(nodeA.addr) - - backend := NewLeaderAwareRedisBackendWithInterval( - []string{nodeA.addr, nodeB.addr}, - "elastickv", - DefaultBackendOptions(), - time.Hour, 500*time.Millisecond, - testLogger, - ) - t.Cleanup(func() { _ = backend.Close() }) - require.Eventually(t, func() bool { - return backend.CurrentLeader() == nodeA.addr - }, 2*time.Second, 10*time.Millisecond) - - nodeA.SetLeader(nodeB.addr) - nodeB.SetLeader(nodeB.addr) - res := backend.Do(context.Background(), "EVALSHA", "deadbeef", "0") - require.Error(t, res.Err()) - require.Contains(t, res.Err().Error(), "NOTLEADER") - require.Equal(t, nodeB.addr, backend.CurrentLeader()) - require.Equal(t, int64(1), nodeA.commands.Load(), "script reaches the stale leader once") - require.Equal(t, int64(0), nodeB.commands.Load(), "script must not be replayed after refresh") + res = backend.Do(context.Background(), "SET", "k", "v") + require.NoError(t, res.Err()) + require.Equal(t, int64(1), nodeB.commands.Load(), "next command should use the refreshed leader") } -func TestLeaderAwareRedisBackend_DoesNotRetryScriptNotLeaderRedisError(t *testing.T) { +func TestLeaderAwareRedisBackend_DoesNotRetryUserNotLeaderError(t *testing.T) { node := newFakeElasticKVNode(t) node.SetLeader(node.addr) - node.SetScriptNotLeaderError(true) + node.SetCommandError("ERR raft engine: not leader") backend := NewLeaderAwareRedisBackendWithInterval( []string{node.addr}, @@ -321,148 +282,66 @@ func TestLeaderAwareRedisBackend_DoesNotRetryScriptNotLeaderRedisError(t *testin testLogger, ) t.Cleanup(func() { _ = backend.Close() }) + require.Eventually(t, func() bool { return backend.CurrentLeader() == node.addr - }, 2*time.Second, 10*time.Millisecond) - - res := backend.Do(context.Background(), "EVAL", "redis.call('SET', KEYS[1], ARGV[1]); return {err='NOTLEADER user script'}", "1", "k", "v") + }, 2*time.Second, 10*time.Millisecond, "initial leader must be set") + res := backend.Do(context.Background(), "EVAL", "return redis.error_reply('raft engine: not leader')", 0) require.Error(t, res.Err()) - assert.Contains(t, res.Err().Error(), "NOTLEADER user script") - assert.Equal(t, int64(1), node.commands.Load(), "script errors must not be retried because the script may already have mutated state") + require.Contains(t, res.Err().Error(), "raft engine: not leader") + require.Equal(t, int64(1), node.commands.Load(), "user Redis error must not be retried") } -func TestLeaderAwareRedisBackend_CoalescesConcurrentRefreshes(t *testing.T) { - node := newFakeElasticKVNode(t) - node.SetLeader(node.addr) - - backend := NewLeaderAwareRedisBackendWithInterval( - []string{node.addr}, - "elastickv", - DefaultBackendOptions(), - time.Hour, time.Second, - testLogger, - ) - t.Cleanup(func() { _ = backend.Close() }) - require.Eventually(t, func() bool { - return backend.CurrentLeader() == node.addr && node.infoCalls.Load() > 0 - }, 2*time.Second, 10*time.Millisecond) - - gate := make(chan struct{}) - node.SetInfoGate(gate) - before := node.infoCalls.Load() - - ownerDone := make(chan struct{}) - go func() { - backend.RefreshLeaderNow(context.Background()) - close(ownerDone) - }() - require.Eventually(t, func() bool { - return node.infoCalls.Load() == before+1 - }, time.Second, 10*time.Millisecond) - - const callers = 32 - var wg sync.WaitGroup - for range callers { - wg.Add(1) - go func() { - defer wg.Done() - ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) - defer cancel() - backend.RefreshLeaderNow(ctx) - }() - } - waitersDone := make(chan struct{}) - go func() { - wg.Wait() - close(waitersDone) - }() - select { - case <-waitersDone: - case <-time.After(time.Second): - close(gate) - <-ownerDone - t.Fatal("refresh waiters did not respect their contexts") - } - - assert.Equal(t, before+1, node.infoCalls.Load()) - close(gate) - <-ownerDone - assert.Equal(t, before+1, node.infoCalls.Load()) -} +func TestLeaderAwareRedisBackend_PipelineNotLeaderRefreshesWithoutReplay(t *testing.T) { + nodeA := newFakeElasticKVNode(t) + nodeB := newFakeElasticKVNode(t) -func TestLeaderAwareRedisBackend_RefreshOutlivesCallerDeadline(t *testing.T) { - node := newFakeElasticKVNode(t) - node.SetLeader(node.addr) + nodeA.SetLeader(nodeA.addr) + nodeB.SetLeader(nodeA.addr) backend := NewLeaderAwareRedisBackendWithInterval( - []string{node.addr}, + []string{nodeA.addr, nodeB.addr}, "elastickv", DefaultBackendOptions(), - time.Hour, time.Second, + time.Hour, 500*time.Millisecond, testLogger, ) t.Cleanup(func() { _ = backend.Close() }) - require.Eventually(t, func() bool { - return backend.CurrentLeader() == node.addr && node.infoCalls.Load() > 0 - }, 2*time.Second, 10*time.Millisecond) - - gate := make(chan struct{}) - node.SetInfoGate(gate) - before := node.infoCalls.Load() - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Millisecond) - defer cancel() - backend.RefreshLeaderNow(ctx) - require.ErrorIs(t, ctx.Err(), context.DeadlineExceeded) - require.Equal(t, before+1, node.infoCalls.Load()) - - close(gate) require.Eventually(t, func() bool { - backend.refreshMu.Lock() - defer backend.refreshMu.Unlock() - return backend.refreshDone == nil - }, time.Second, 10*time.Millisecond) - assert.Equal(t, before+1, node.infoCalls.Load(), "caller cancellation must not start a replacement probe") -} + return backend.CurrentLeader() == nodeA.addr + }, 2*time.Second, 10*time.Millisecond, "initial leader must be A") -func TestLeaderRefreshTransportErrorClassification(t *testing.T) { - require.True(t, isLeaderRefreshTransportError(io.EOF)) - require.True(t, isLeaderRefreshTransportError(&net.OpError{Op: "read", Err: errors.New("connection reset by peer")})) - require.False(t, isLeaderRefreshTransportError(context.Canceled)) - require.False(t, isLeaderRefreshTransportError(context.DeadlineExceeded)) - require.False(t, isLeaderRefreshTransportError(errors.New("write conflict"))) -} + nodeA.SetLeader(nodeB.addr) + nodeB.SetLeader(nodeB.addr) -func TestElasticKVNotLeaderErrorClassification(t *testing.T) { - tests := []struct { - name string - err error - want bool - }{ - {name: "nil", err: nil, want: false}, - {name: "canonical redis code", err: redisError("NOTLEADER leader not found"), want: true}, - {name: "bare canonical code", err: errors.New("NOTLEADER"), want: true}, - {name: "internal sentinel text", err: errors.New("raft engine: not leader"), want: true}, - {name: "grpc wrapped text", err: errors.New("rpc error: code = Unknown desc = leader not found"), want: true}, - {name: "unrelated error", err: errors.New("write conflict"), want: false}, - {name: "redis user error containing phrase", err: redisError("ERR script says raft engine: not leader"), want: false}, - {name: "redis user error with bare phrase", err: redisError("leader not found"), want: false}, - {name: "free form suffix", err: errors.New("key value says leader not found"), want: false}, + results, err := backend.Pipeline(context.Background(), [][]any{ + {"MULTI"}, + {"SET", "k", "v"}, + {"EXEC"}, + }) + require.NoError(t, err) + for _, result := range results { + require.Error(t, result.Err()) + require.Contains(t, result.Err().Error(), "NOTLEADER") } - - for _, tc := range tests { - t.Run(tc.name, func(t *testing.T) { - assert.Equal(t, tc.want, isElasticKVNotLeaderError(tc.err)) - }) + require.Equal(t, nodeB.addr, backend.CurrentLeader(), "pipeline not-leader must synchronously adopt the advertised leader") + require.Equal(t, int64(3), nodeA.commands.Load(), "first pipeline attempt should hit the former leader") + require.Equal(t, int64(0), nodeB.commands.Load(), "ambiguous pipelines must not be replayed") + + results, err = backend.Pipeline(context.Background(), [][]any{ + {"MULTI"}, + {"SET", "k", "v"}, + {"EXEC"}, + }) + require.NoError(t, err) + for _, result := range results { + require.NoError(t, result.Err()) } + require.Equal(t, int64(3), nodeB.commands.Load(), "next pipeline should use the refreshed leader") } -type redisError string - -func (e redisError) Error() string { return string(e) } -func (redisError) RedisError() {} - func TestLeaderAwareRedisBackend_ConcurrentCloseIsRaceFree(t *testing.T) { // Regression guard: Close() must not race concurrent Do() callers — // currentClient and ensureClientLocked hold the lock consistently with diff --git a/proxy/metrics.go b/proxy/metrics.go index df4f013f5..121bec3d8 100644 --- a/proxy/metrics.go +++ b/proxy/metrics.go @@ -18,6 +18,7 @@ type ProxyMetrics struct { ActiveConnections prometheus.Gauge AsyncDrops prometheus.Counter + AsyncBackpressure prometheus.Counter AsyncDropsByQueue *prometheus.CounterVec AsyncQueueDepth *prometheus.GaugeVec AsyncQueueCapacity *prometheus.GaugeVec @@ -89,7 +90,12 @@ func NewProxyMetrics(reg prometheus.Registerer) *ProxyMetrics { AsyncDrops: prometheus.NewCounter(prometheus.CounterOpts{ Namespace: "proxy", Name: "async_drops_total", - Help: "Total async operations dropped due to queue backpressure or expiry.", + Help: "Total best-effort async operations dropped due to semaphore backpressure.", + }), + AsyncBackpressure: prometheus.NewCounter(prometheus.CounterOpts{ + Namespace: "proxy", + Name: "async_backpressure_total", + Help: "Total strict secondary writes that waited for an async semaphore slot instead of being dropped.", }), AsyncDropsByQueue: prometheus.NewCounterVec(prometheus.CounterOpts{ Namespace: "proxy", @@ -173,6 +179,7 @@ func NewProxyMetrics(reg prometheus.Registerer) *ProxyMetrics { m.Divergences, m.MigrationGaps, m.AsyncDrops, + m.AsyncBackpressure, m.AsyncDropsByQueue, m.AsyncQueueDepth, m.AsyncQueueCapacity, diff --git a/proxy/noop_backend.go b/proxy/noop_backend.go new file mode 100644 index 000000000..fbe5dc811 --- /dev/null +++ b/proxy/noop_backend.go @@ -0,0 +1,38 @@ +package proxy + +import ( + "context" + + "github.com/redis/go-redis/v9" +) + +// noopBackend is used for modes with no secondary backend. Keeping a concrete +// Backend avoids starting unnecessary leader-discovery goroutines in redis-only +// and elastickv-only modes while preserving DualWriter's simple shape. +type noopBackend struct { + name string +} + +// NewNoopBackend returns a Backend placeholder for an intentionally unused +// side of the proxy. +func NewNoopBackend(name string) Backend { + return noopBackend{name: name} +} + +func (n noopBackend) Do(ctx context.Context, args ...any) *redis.Cmd { + cmd := redis.NewCmd(ctx, args...) + cmd.SetErr(ErrNoLeaderBackend) + return cmd +} + +func (n noopBackend) Pipeline(context.Context, [][]any) ([]*redis.Cmd, error) { + return nil, ErrNoLeaderBackend +} + +func (n noopBackend) Close() error { + return nil +} + +func (n noopBackend) Name() string { + return n.name +} diff --git a/proxy/proxy.go b/proxy/proxy.go index b586ff090..9b2add2b1 100644 --- a/proxy/proxy.go +++ b/proxy/proxy.go @@ -59,6 +59,10 @@ func NewProxyServer(cfg ProxyConfig, dual *DualWriter, metrics *ProxyMetrics, se // ListenAndServe starts the redcon proxy server. func (p *ProxyServer) ListenAndServe(ctx context.Context) error { + if p.cfg.Mode == ModeRedisOnly && p.cfg.RedisOnlyRaw { + return p.listenAndServeRawRedis(ctx) + } + p.shutdownCtx = ctx var lc net.ListenConfig @@ -430,11 +434,8 @@ func (p *ProxyServer) execTxn(conn redcon.Conn, state *proxyConnState) { conn.WriteError("ERR empty transaction response") } - // Async replay to secondary (bounded) if p.dual.hasSecondaryWrite() { - p.dual.goAsync(func(ctx context.Context) { - p.dual.writeSecondaryPipeline(ctx, cmds) - }) + p.dual.replaySecondaryPipeline(cmds) } } diff --git a/proxy/proxy_test.go b/proxy/proxy_test.go index 94cac2cf4..c92815b62 100644 --- a/proxy/proxy_test.go +++ b/proxy/proxy_test.go @@ -67,6 +67,34 @@ func (b *mockBackend) CallCount() int { return len(b.calls) } +func (b *mockBackend) Calls() [][]any { + b.mu.Lock() + defer b.mu.Unlock() + out := make([][]any, len(b.calls)) + for i := range b.calls { + out[i] = append([]any(nil), b.calls[i]...) + } + return out +} + +type refreshableMockBackend struct { + *mockBackend + refreshMu sync.Mutex + refreshes int +} + +func (b *refreshableMockBackend) RefreshLeaderNow(ctx context.Context) { + b.refreshMu.Lock() + defer b.refreshMu.Unlock() + b.refreshes++ +} + +func (b *refreshableMockBackend) RefreshCount() int { + b.refreshMu.Lock() + defer b.refreshMu.Unlock() + return b.refreshes +} + // Helper to create a doFunc that returns a specific value. func makeCmd(val any, err error) func(ctx context.Context, args ...any) *redis.Cmd { return func(ctx context.Context, args ...any) *redis.Cmd { @@ -636,301 +664,290 @@ func TestDualWriter_Blocking_UsesTimeoutAwareBackend(t *testing.T) { assert.Equal(t, []any{[]byte("BZPOPMIN"), []byte("queue"), []byte("5")}, primary.args) } -func TestDualWriter_Blocking_ReplaysBZPopAsZRem(t *testing.T) { - for _, tc := range []struct { - cmd string +func TestDualWriter_Blocking_ReplaysBZPopMinAsZRem(t *testing.T) { + primary := &timeoutCapturingBackend{ + name: "primary", + returnValue: []any{[]byte("queue"), []byte("job-1"), []byte("42")}, + } + secondary := newMockBackend("secondary") + secondary.doFunc = makeCmd(int64(1), nil) + + metrics := newTestMetrics() + cfg := ProxyConfig{ + Mode: ModeDualWrite, + SecondaryTimeout: 10 * time.Second, + SecondaryBlockingReplayConcurrency: 1, + } + d := NewDualWriter(primary, secondary, cfg, metrics, newTestSentry(), testLogger) + + resp, err := d.Blocking(context.Background(), "BZPOPMIN", [][]byte{[]byte("BZPOPMIN"), []byte("queue"), []byte("5")}) + assert.NoError(t, err) + assert.Equal(t, []any{[]byte("queue"), []byte("job-1"), []byte("42")}, resp) + d.Close() + + assert.Equal(t, [][]any{{[]byte("ZREM"), []byte("queue"), []byte("job-1")}}, secondary.Calls()) + assert.InDelta(t, 0, testutil.ToFloat64(metrics.AsyncDrops), 0.001) + assert.InDelta(t, 1, testutil.ToFloat64(metrics.CommandTotal.WithLabelValues("ZREM", "secondary", "ok")), 0.001) +} + +func TestDualWriter_Blocking_ReplaysListPopAsLRem(t *testing.T) { + tests := []struct { + name string + cmd string + resp []any + want [][]any }{ - {cmd: "BZPOPMIN"}, - {cmd: "BZPOPMAX"}, - } { - t.Run(tc.cmd, func(t *testing.T) { + { + name: "BLPOP removes first matching element", + cmd: "BLPOP", + resp: []any{[]byte("queue"), []byte("job-1")}, + want: [][]any{{[]byte("LREM"), []byte("queue"), int64(1), []byte("job-1")}}, + }, + { + name: "BRPOP removes last matching element", + cmd: "BRPOP", + resp: []any{[]byte("queue"), []byte("job-2")}, + want: [][]any{{[]byte("LREM"), []byte("queue"), int64(-1), []byte("job-2")}}, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { primary := &timeoutCapturingBackend{ name: "primary", - returnValue: []any{"queue", "job-1", "12.5"}, + returnValue: tc.resp, } secondary := newMockBackend("secondary") secondary.doFunc = makeCmd(int64(1), nil) metrics := newTestMetrics() - d := NewDualWriter( - primary, - secondary, - ProxyConfig{ - Mode: ModeDualWrite, - SecondaryTimeout: time.Second, - SecondaryBlockingReplayConcurrency: 1, - }, - metrics, - newTestSentry(), - testLogger, - ) + cfg := ProxyConfig{ + Mode: ModeDualWrite, + SecondaryTimeout: 10 * time.Second, + SecondaryBlockingReplayConcurrency: 1, + } + d := NewDualWriter(primary, secondary, cfg, metrics, newTestSentry(), testLogger) resp, err := d.Blocking(context.Background(), tc.cmd, [][]byte{[]byte(tc.cmd), []byte("queue"), []byte("5")}) assert.NoError(t, err) - assert.Equal(t, []any{"queue", "job-1", "12.5"}, resp) + assert.Equal(t, tc.resp, resp) d.Close() - assert.Equal(t, 1, secondary.CallCount()) - secondary.mu.Lock() - got := append([]any(nil), secondary.calls[0]...) - secondary.mu.Unlock() - assert.Equal(t, []any{"ZREM", "queue", "job-1"}, got) + assert.Equal(t, tc.want, secondary.Calls()) + assert.InDelta(t, 0, testutil.ToFloat64(metrics.AsyncDrops), 0.001) + assert.InDelta(t, 1, testutil.ToFloat64(metrics.CommandTotal.WithLabelValues("LREM", "secondary", "ok")), 0.001) }) } } -func TestDualWriter_Blocking_RetriesBZPopReplayUntilRemoved(t *testing.T) { - primary := &timeoutCapturingBackend{ - name: "primary", - returnValue: []any{"queue", "job-1", "12.5"}, - } - secondary := newMockBackend("secondary") - var attempts atomic.Int32 - secondary.doFunc = func(ctx context.Context, args ...any) *redis.Cmd { - cmd := redis.NewCmd(ctx, args...) - if attempts.Add(1) == 1 { - cmd.SetVal(int64(0)) - return cmd - } - cmd.SetVal(int64(1)) - return cmd - } - - metrics := newTestMetrics() - d := NewDualWriter( - primary, - secondary, - ProxyConfig{ - Mode: ModeDualWrite, - SecondaryTimeout: time.Second, - SecondaryBlockingReplayConcurrency: 1, +func TestDualWriter_Blocking_ReplaysListMoveAsEval(t *testing.T) { + tests := []struct { + name string + cmd string + args [][]byte + resp any + want []any + }{ + { + name: "BRPOPLPUSH removes source tail and pushes destination head", + cmd: "BRPOPLPUSH", + args: [][]byte{[]byte("BRPOPLPUSH"), []byte("source"), []byte("dest"), []byte("5")}, + resp: []byte("job-1"), + want: []any{ + []byte("EVAL"), + blockingListMoveReplayScript, + int64(2), + []byte("source"), + []byte("dest"), + int64(-1), + []byte("LEFT"), + []byte("job-1"), + }, + }, + { + name: "BLMOVE mirrors requested source and destination sides", + cmd: "BLMOVE", + args: [][]byte{ + []byte("BLMOVE"), []byte("source"), []byte("dest"), + []byte("LEFT"), []byte("RIGHT"), []byte("5"), + }, + resp: "job-2", + want: []any{ + []byte("EVAL"), + blockingListMoveReplayScript, + int64(2), + []byte("source"), + []byte("dest"), + int64(1), + []byte("RIGHT"), + []byte("job-2"), + }, }, - metrics, - newTestSentry(), - testLogger, - ) - - _, err := d.Blocking(context.Background(), "BZPOPMIN", [][]byte{[]byte("BZPOPMIN"), []byte("queue"), []byte("5")}) - assert.NoError(t, err) - d.Close() - - assert.Equal(t, int32(2), attempts.Load()) - assert.Equal(t, 2, secondary.CallCount()) -} - -func TestDualWriter_Blocking_BZPopReplayMissIsNotSecondaryWriteError(t *testing.T) { - primary := &timeoutCapturingBackend{ - name: "primary", - returnValue: []any{"queue", "job-1", "12.5"}, } - secondary := newMockBackend("secondary") - secondary.doFunc = makeCmd(int64(0), nil) - metrics := newTestMetrics() - d := NewDualWriter( - primary, - secondary, - ProxyConfig{ - Mode: ModeDualWrite, - SecondaryTimeout: 2 * time.Second, - SecondaryBlockingReplayConcurrency: 1, - }, - metrics, - newTestSentry(), - testLogger, - ) + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + primary := &timeoutCapturingBackend{ + name: "primary", + returnValue: tc.resp, + } + secondary := newMockBackend("secondary") + secondary.doFunc = makeCmd(int64(1), nil) - _, err := d.Blocking(context.Background(), "BZPOPMIN", [][]byte{[]byte("BZPOPMIN"), []byte("queue"), []byte("5")}) - assert.NoError(t, err) - d.Close() + metrics := newTestMetrics() + cfg := ProxyConfig{ + Mode: ModeDualWrite, + SecondaryTimeout: 10 * time.Second, + SecondaryBlockingReplayConcurrency: 1, + } + d := NewDualWriter(primary, secondary, cfg, metrics, newTestSentry(), testLogger) - assert.Equal(t, noEffectReplayRetryLimit(context.Background(), blockingReplayNoEffectRetryWindow)+1, secondary.CallCount()) - assert.InDelta(t, 0, testutil.ToFloat64(metrics.SecondaryWriteErrors), 0.001) - assert.InDelta(t, 1, testutil.ToFloat64( - metrics.CommandTotal.WithLabelValues("ZREM", "secondary", "miss")), 0.001) -} + resp, err := d.Blocking(context.Background(), tc.cmd, tc.args) + assert.NoError(t, err) + assert.Equal(t, tc.resp, resp) + d.Close() -func TestDualWriter_Blocking_BZPopReplayShortTimeoutStillAttemptsZRem(t *testing.T) { - primary := &timeoutCapturingBackend{ - name: "primary", - returnValue: []any{"queue", "job-1", "12.5"}, + assert.Equal(t, [][]any{tc.want}, secondary.Calls()) + assert.InDelta(t, 0, testutil.ToFloat64(metrics.AsyncDrops), 0.001) + assert.InDelta(t, 1, testutil.ToFloat64(metrics.CommandTotal.WithLabelValues("EVAL", "secondary", "ok")), 0.001) + }) } - secondary := newMockBackend("secondary") - secondary.doFunc = makeCmd(int64(0), nil) +} - metrics := newTestMetrics() - d := NewDualWriter( - primary, - secondary, - ProxyConfig{ - Mode: ModeDualWrite, - SecondaryTimeout: 30 * time.Millisecond, - SecondaryBlockingReplayConcurrency: 1, +func TestDualWriter_Blocking_ReplaysBLMPopAsEval(t *testing.T) { + tests := []struct { + name string + args [][]byte + resp any + want []any + }{ + { + name: "left pop removes returned values from head side", + args: [][]byte{ + []byte("BLMPOP"), []byte("5"), []byte("2"), + []byte("queue-a"), []byte("queue-b"), []byte("LEFT"), + []byte("COUNT"), []byte("2"), + }, + resp: []any{[]byte("queue-b"), []any{[]byte("job-1"), []byte("job-2")}}, + want: []any{ + []byte("EVAL"), + blockingListMultiPopReplayScript, + int64(1), + []byte("queue-b"), + int64(1), + []byte("job-1"), + []byte("job-2"), + }, }, - metrics, - newTestSentry(), - testLogger, - ) + { + name: "right pop removes returned values from tail side", + args: [][]byte{ + []byte("BLMPOP"), []byte("5"), []byte("1"), + []byte("queue"), []byte("RIGHT"), + }, + resp: []any{"queue", []string{"job-3"}}, + want: []any{ + []byte("EVAL"), + blockingListMultiPopReplayScript, + int64(1), + []byte("queue"), + int64(-1), + []byte("job-3"), + }, + }, + } - _, err := d.Blocking(context.Background(), "BZPOPMIN", [][]byte{[]byte("BZPOPMIN"), []byte("queue"), []byte("5")}) - assert.NoError(t, err) - d.Close() + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + primary := &timeoutCapturingBackend{ + name: "primary", + returnValue: tc.resp, + } + secondary := newMockBackend("secondary") + secondary.doFunc = makeCmd(int64(1), nil) - assert.GreaterOrEqual(t, secondary.CallCount(), 1) - assert.InDelta(t, 0, testutil.ToFloat64(metrics.SecondaryWriteErrors), 0.001) - assert.InDelta(t, 1, testutil.ToFloat64( - metrics.CommandTotal.WithLabelValues("ZREM", "secondary", "miss")), 0.001) -} + metrics := newTestMetrics() + cfg := ProxyConfig{ + Mode: ModeDualWrite, + SecondaryTimeout: 10 * time.Second, + SecondaryBlockingReplayConcurrency: 1, + } + d := NewDualWriter(primary, secondary, cfg, metrics, newTestSentry(), testLogger) -func TestNoEffectReplayRetryLimitIncludesJitterBudget(t *testing.T) { - limit := noEffectReplayRetryLimit(context.Background(), blockingReplayNoEffectRetryWindow) - assert.Positive(t, limit) + resp, err := d.Blocking(context.Background(), "BLMPOP", tc.args) + assert.NoError(t, err) + assert.Equal(t, tc.resp, resp) + d.Close() - var spent time.Duration - backoff := compactedRetryInitialBackoff - for range limit { - spent += retryBackoffWithMaxJitter(backoff) - backoff = nextCompactedRetryBackoff(backoff) + assert.Equal(t, [][]any{tc.want}, secondary.Calls()) + assert.InDelta(t, 0, testutil.ToFloat64(metrics.AsyncDrops), 0.001) + assert.InDelta(t, 1, testutil.ToFloat64(metrics.CommandTotal.WithLabelValues("EVAL", "secondary", "ok")), 0.001) + }) } - - assert.LessOrEqual(t, spent, blockingReplayNoEffectRetryWindow) - assert.Greater(t, spent+retryBackoffWithMaxJitter(backoff), blockingReplayNoEffectRetryWindow) } -func TestDualWriter_BlockingReplayDoesNotConsumeWriteWorkers(t *testing.T) { - primary := &timeoutCapturingBackend{ - name: "primary", - returnValue: []any{"queue", "job-1", "12.5"}, - } +func TestDualWriter_Blocking_XReadDoesNotUseWriteSemaphore(t *testing.T) { + primary := &timeoutCapturingBackend{name: "primary", returnValue: []any{}} secondary := newMockBackend("secondary") - secondary.doFunc = makeCmd(int64(1), nil) metrics := newTestMetrics() - d := NewDualWriter( - primary, - secondary, - ProxyConfig{ - Mode: ModeDualWrite, - SecondaryTimeout: time.Second, - SecondaryWriteConcurrency: 1, - SecondaryBlockingReplayConcurrency: 1, - SecondaryWriteQueueCapacity: 1, - SecondaryBlockingReplayQueueCapacity: 1, - }, - metrics, - newTestSentry(), - testLogger, - ) + cfg := ProxyConfig{Mode: ModeDualWrite, SecondaryWriteConcurrency: 1, SecondaryTimeout: 10 * time.Second} + d := NewDualWriter(primary, secondary, cfg, metrics, newTestSentry(), testLogger) blocker := make(chan struct{}) - started := make(chan struct{}) - d.goWrite(func(context.Context) { - close(started) + d.goWrite(func() { <-blocker }) - <-started - _, err := d.Blocking(context.Background(), "BZPOPMIN", [][]byte{[]byte("BZPOPMIN"), []byte("queue"), []byte("5")}) + resp, err := d.Blocking(context.Background(), "XREAD", [][]byte{ + []byte("XREAD"), []byte("BLOCK"), []byte("1"), []byte("STREAMS"), []byte("jobs"), []byte("0"), + }) assert.NoError(t, err) - assert.Eventually(t, func() bool { return secondary.CallCount() == 1 }, - time.Second, 10*time.Millisecond) - assert.InDelta(t, 1, testutil.ToFloat64(metrics.AsyncWorkersActive.WithLabelValues(asyncQueueWrite)), 0.001) + assert.Equal(t, []any{}, resp) + assert.Empty(t, secondary.Calls()) + assert.InDelta(t, 0, testutil.ToFloat64(metrics.AsyncDrops), 0.001) + assert.InDelta(t, 0, testutil.ToFloat64(metrics.AsyncBackpressure), 0.001) close(blocker) d.Close() } -func TestDualWriter_Blocking_ReplaysXReadGroup(t *testing.T) { - primary := &timeoutCapturingBackend{ - name: "primary", - returnValue: []any{"stream-result"}, - } - secondary := newMockBackend("secondary") - - metrics := newTestMetrics() - d := NewDualWriter( - primary, - secondary, - ProxyConfig{ - Mode: ModeDualWrite, - SecondaryTimeout: time.Second, - SecondaryBlockingReplayConcurrency: 1, +func TestDualWriter_Blocking_ReplaysXReadGroupToSecondary(t *testing.T) { + resp := []any{ + []any{ + []byte("jobs"), + []any{ + []any{[]byte("1-0"), []any{[]byte("field"), []byte("value")}}, + }, }, - metrics, - newTestSentry(), - testLogger, - ) - - args := [][]byte{ - []byte("XREADGROUP"), []byte("GROUP"), []byte("g"), []byte("c"), - []byte("BLOCK"), []byte("1000"), []byte("STREAMS"), []byte("jobs"), []byte(">"), - } - resp, err := d.Blocking(context.Background(), "XREADGROUP", args) - assert.NoError(t, err) - assert.Equal(t, []any{"stream-result"}, resp) - d.Close() - - assert.Equal(t, 1, secondary.CallCount()) - secondary.mu.Lock() - got := append([]any(nil), secondary.calls[0]...) - secondary.mu.Unlock() - assert.Equal(t, bytesArgsToInterfaces(args), got) -} - -func TestDualWriter_Blocking_DoesNotReplayXRead(t *testing.T) { - primary := &timeoutCapturingBackend{ - name: "primary", - returnValue: []any{"stream-result"}, } + primary := &timeoutCapturingBackend{name: "primary", returnValue: resp} secondary := newMockBackend("secondary") + secondary.doFunc = makeCmd(resp, nil) metrics := newTestMetrics() - d := NewDualWriter( - primary, - secondary, - ProxyConfig{Mode: ModeDualWrite, SecondaryTimeout: time.Second}, - metrics, - newTestSentry(), - testLogger, - ) - - resp, err := d.Blocking(context.Background(), "XREAD", [][]byte{ - []byte("XREAD"), []byte("BLOCK"), []byte("1000"), []byte("STREAMS"), []byte("jobs"), []byte("0"), - }) - assert.NoError(t, err) - assert.Equal(t, []any{"stream-result"}, resp) - d.Close() - - assert.Equal(t, 0, secondary.CallCount()) -} - -func TestDualWriter_Blocking_DoesNotReplayWhenBlockingReplayDisabled(t *testing.T) { - primary := &timeoutCapturingBackend{ - name: "primary", - returnValue: []any{"queue", "job-1", "12.5"}, + cfg := ProxyConfig{ + Mode: ModeDualWrite, + SecondaryTimeout: 10 * time.Second, + SecondaryBlockingReplayConcurrency: 1, } - secondary := newMockBackend("secondary") - - metrics := newTestMetrics() - d := NewDualWriter( - primary, - secondary, - ProxyConfig{Mode: ModeDualWrite, SecondaryTimeout: time.Second}, - metrics, - newTestSentry(), - testLogger, - ) + d := NewDualWriter(primary, secondary, cfg, metrics, newTestSentry(), testLogger) - resp, err := d.Blocking(context.Background(), "BZPOPMIN", [][]byte{[]byte("BZPOPMIN"), []byte("queue"), []byte("5")}) + args := [][]byte{ + []byte(cmdNameXREADGROUP), []byte("GROUP"), []byte("g"), []byte("c"), + []byte("BLOCK"), []byte("1"), []byte("STREAMS"), []byte("jobs"), []byte(">"), + } + got, err := d.Blocking(context.Background(), cmdNameXREADGROUP, args) assert.NoError(t, err) - assert.Equal(t, []any{"queue", "job-1", "12.5"}, resp) + assert.Equal(t, resp, got) d.Close() - assert.Equal(t, 0, secondary.CallCount()) + assert.Equal(t, [][]any{bytesArgsToInterfaces(args)}, secondary.Calls()) assert.InDelta(t, 0, testutil.ToFloat64(metrics.AsyncDrops), 0.001) + assert.InDelta(t, 1, testutil.ToFloat64(metrics.CommandTotal.WithLabelValues(cmdNameXREADGROUP, "secondary", "ok")), 0.001) } -func TestDualWriter_GoAsync_QueuesBurstBeforeDropping(t *testing.T) { +func TestDualWriter_GoAsync_Bounded(t *testing.T) { primary := newMockBackend("primary") primary.doFunc = makeCmd("OK", nil) secondary := newMockBackend("secondary") @@ -1019,6 +1036,54 @@ func TestDualWriter_GoAsync_DropLogsAreRateLimited(t *testing.T) { d.Close() } +func TestDualWriter_Write_QueuesWhenWriteLimitIsBusy(t *testing.T) { + primary := newMockBackend("primary") + primary.doFunc = makeCmd("OK", nil) + secondary := newMockBackend("secondary") + secondary.doFunc = makeCmd("OK", nil) + + metrics := newTestMetrics() + cfg := ProxyConfig{ + Mode: ModeDualWrite, + SecondaryWriteConcurrency: 1, + SecondaryWriteQueueCapacity: 1, + SecondaryTimeout: 10 * time.Second, + } + d := NewDualWriter(primary, secondary, cfg, metrics, newTestSentry(), testLogger) + + blocker := make(chan struct{}) + started := make(chan struct{}) + d.goAsync(func(context.Context) { + close(started) + <-blocker + }) + <-started + + done := make(chan struct{}) + go func() { + _, err := d.Write(context.Background(), "SET", [][]byte{ + []byte("SET"), []byte("key"), []byte("value"), + }) + assert.NoError(t, err) + close(done) + }() + + select { + case <-done: + // good: the primary succeeded while secondary replay stayed queued. + case <-time.After(time.Second): + t.Fatal("Write blocked while its replay was queued") + } + + assert.Equal(t, 0, secondary.CallCount()) + assert.InDelta(t, 0, testutil.ToFloat64(metrics.AsyncDrops), 0.001) + assert.InDelta(t, 1, testutil.ToFloat64(metrics.AsyncQueueDepth.WithLabelValues(asyncQueueWrite)), 0.001) + + close(blocker) + d.Close() + assert.Equal(t, 1, secondary.CallCount()) +} + func TestDualWriter_Script_QueuesWhenScriptLimitIsBusy(t *testing.T) { primary := newMockBackend("primary") primary.doFunc = makeCmd("OK", nil) @@ -1043,7 +1108,6 @@ func TestDualWriter_Script_QueuesWhenScriptLimitIsBusy(t *testing.T) { }) <-started - // Script returns after the authoritative primary write while replay is queued. done := make(chan struct{}) go func() { _, err := d.Script(context.Background(), "EVALSHA", [][]byte{ @@ -1055,7 +1119,7 @@ func TestDualWriter_Script_QueuesWhenScriptLimitIsBusy(t *testing.T) { select { case <-done: - // good + // good: the primary succeeded while secondary replay stayed queued. case <-time.After(time.Second): t.Fatal("Script blocked while its replay was queued") } @@ -1480,8 +1544,8 @@ func TestDualWriter_writeSecondary_RetriesReadTSCompacted(t *testing.T) { } func TestDualWriter_writeSecondary_ReadTSCompactedRetriesAreBounded(t *testing.T) { - // When the compacted error is persistent, the retry loop must stop after - // maxCompactedRetries+1 attempts so the secondary goroutine returns + // When the transient error is persistent, the retry loop must stop after + // maxSecondaryTransientRetries+1 attempts so the secondary goroutine returns // instead of burning a scriptSem slot indefinitely. primary := newMockBackend("primary") primary.doFunc = makeCmd("OK", nil) @@ -1495,24 +1559,27 @@ func TestDualWriter_writeSecondary_ReadTSCompactedRetriesAreBounded(t *testing.T d.writeSecondary(context.Background(), "EVALSHA", []any{[]byte("EVALSHA"), []byte("deadbeef"), []byte("0")}) - assert.Equal(t, maxCompactedRetries+1, secondary.CallCount(), - "secondary must stop after maxCompactedRetries+1 attempts") + assert.Equal(t, maxSecondaryTransientRetries+1, secondary.CallCount(), + "secondary must stop after maxSecondaryTransientRetries+1 attempts") assert.InDelta(t, 1, testutil.ToFloat64(metrics.SecondaryWriteErrors), 0.001, "a persistent compacted error must still be reported as a secondary write error") } -func TestDualWriter_writeSecondary_RetriesServerOverloadedNonScript(t *testing.T) { +func TestDualWriter_writeSecondary_RetriesRetryLimitWriteConflict(t *testing.T) { + // A hot key can exhaust the secondary Redis adapter's internal OCC retry loop. + // Re-sending the command lets the backend choose a fresh timestamp and keeps + // dual-write traffic from permanently diverging on a transient conflict. primary := newMockBackend("primary") primary.doFunc = makeCmd("OK", nil) secondary := newMockBackend("secondary") - overloadedErr := testRedisErr(serverOverloadedMarker) + retryLimitErr := testRedisErr("redis txn retry limit exceeded: key: myzset: write conflict") var calls int secondary.doFunc = func(ctx context.Context, args ...any) *redis.Cmd { calls++ cmd := redis.NewCmd(ctx, args...) - if calls < 4 { - cmd.SetErr(overloadedErr) + if calls < 3 { + cmd.SetErr(retryLimitErr) return cmd } cmd.SetVal("OK") @@ -1522,121 +1589,106 @@ func TestDualWriter_writeSecondary_RetriesServerOverloadedNonScript(t *testing.T metrics := newTestMetrics() d := NewDualWriter(primary, secondary, ProxyConfig{Mode: ModeDualWrite, SecondaryTimeout: time.Second}, metrics, newTestSentry(), testLogger) - d.writeSecondary(context.Background(), "SET", []any{[]byte("SET"), []byte("k"), []byte("v")}) + d.writeSecondary(context.Background(), "ZADD", []any{[]byte("ZADD"), []byte("myzset"), []byte("1"), []byte("member")}) - assert.Equal(t, 4, calls, "secondary must retry transient admission failures") + assert.Equal(t, 3, calls, "secondary must retry transient retry-limit write conflicts") assert.InDelta(t, 0, testutil.ToFloat64(metrics.SecondaryWriteErrors), 0.001, - "a retried server-overloaded response must not count as a secondary write error") -} - -func TestDualWriter_writeSecondary_ServerOverloadedNonScriptRetriesAreBounded(t *testing.T) { - primary := newMockBackend("primary") - primary.doFunc = makeCmd("OK", nil) - - secondary := newMockBackend("secondary") - secondary.doFunc = makeCmd(nil, testRedisErr(serverOverloadedMarker)) - - metrics := newTestMetrics() - d := NewDualWriter(primary, secondary, ProxyConfig{Mode: ModeDualWrite, SecondaryTimeout: 2 * time.Second}, metrics, newTestSentry(), testLogger) - - d.writeSecondary(context.Background(), "SET", []any{[]byte("SET"), []byte("k"), []byte("v")}) - - assert.Equal(t, maxServerOverloadedRetries+1, secondary.CallCount(), - "secondary must stop after maxServerOverloadedRetries+1 attempts") - assert.InDelta(t, 1, testutil.ToFloat64(metrics.SecondaryWriteErrors), 0.001, - "a persistent server-overloaded response must still be reported as a secondary write error") - assert.InDelta(t, 1, testutil.ToFloat64( - metrics.SecondaryWriteErrorsByReason.WithLabelValues("SET", "busy")), 0.001) -} - -func TestDualWriter_writeSecondary_DoesNotRetryScriptReturnedServerOverloaded(t *testing.T) { - primary := newMockBackend("primary") - primary.doFunc = makeCmd("OK", nil) - - secondary := newMockBackend("secondary") - secondary.doFunc = makeCmd(nil, testRedisErr(serverOverloadedMarker)) - - metrics := newTestMetrics() - d := NewDualWriter(primary, secondary, ProxyConfig{Mode: ModeDualWrite, SecondaryTimeout: time.Second}, metrics, newTestSentry(), testLogger) - - d.writeSecondary(context.Background(), "EVALSHA", []any{[]byte("EVALSHA"), []byte("deadbeef"), []byte("0")}) - - assert.Equal(t, 1, secondary.CallCount(), "script-returned BUSY is ambiguous and must not be replayed") - assert.InDelta(t, 1, testutil.ToFloat64(metrics.SecondaryWriteErrors), 0.001) - assert.InDelta(t, 1, testutil.ToFloat64( - metrics.SecondaryWriteErrorsByReason.WithLabelValues("EVALSHA", "busy")), 0.001) + "a retried success must not count as a secondary write error") } -func TestDualWriter_writeSecondaryPipeline_RetriesServerOverloadedExec(t *testing.T) { +func TestDualWriter_writeSecondary_RetriesNotLeaderAfterRefresh(t *testing.T) { primary := newMockBackend("primary") primary.doFunc = makeCmd("OK", nil) - overloadedErr := testRedisErr(serverOverloadedMarker) - secondary := newMockBackend("secondary") + secondary := &refreshableMockBackend{mockBackend: newMockBackend("secondary")} + notLeaderErr := testRedisErr("NOTLEADER etcd raft engine is not leader") var calls int - secondary.pipelineFunc = func(ctx context.Context, cmds [][]any) ([]*redis.Cmd, error) { + secondary.doFunc = func(ctx context.Context, args ...any) *redis.Cmd { calls++ - var execErr error - if calls < 4 { - execErr = overloadedErr + cmd := redis.NewCmd(ctx, args...) + if calls == 1 { + cmd.SetErr(notLeaderErr) + return cmd } - return pipelineResults(ctx, cmds, execErr), nil + cmd.SetVal("OK") + return cmd } metrics := newTestMetrics() d := NewDualWriter(primary, secondary, ProxyConfig{Mode: ModeDualWrite, SecondaryTimeout: time.Second}, metrics, newTestSentry(), testLogger) - d.writeSecondaryPipeline(context.Background(), [][]any{{"MULTI"}, {"EVALSHA", "deadbeef", "0"}, {"EXEC"}}) + d.writeSecondary(context.Background(), "EVALSHA", []any{[]byte("EVALSHA"), []byte("deadbeef"), []byte("0")}) - assert.Equal(t, 4, calls, "secondary transaction replay must retry transient admission failures") + assert.Equal(t, 2, calls, "secondary must retry after a not-leader rejection") + assert.Equal(t, 1, secondary.RefreshCount(), "not-leader retries must force leader rediscovery before retrying") assert.InDelta(t, 0, testutil.ToFloat64(metrics.SecondaryWriteErrors), 0.001, - "a retried server-overloaded EXEC must not count as a secondary write error") + "a refreshed retry success must not count as a secondary write error") } -func TestDualWriter_writeSecondaryPipeline_ServerOverloadedRetriesAreBounded(t *testing.T) { +func TestDualWriter_writeSecondary_DoesNotRetryUserNotLeaderError(t *testing.T) { primary := newMockBackend("primary") primary.doFunc = makeCmd("OK", nil) - overloadedErr := testRedisErr(serverOverloadedMarker) - secondary := newMockBackend("secondary") + secondary := &refreshableMockBackend{mockBackend: newMockBackend("secondary")} + userErr := testRedisErr("ERR raft engine: not leader") var calls int - secondary.pipelineFunc = func(ctx context.Context, cmds [][]any) ([]*redis.Cmd, error) { + secondary.doFunc = func(ctx context.Context, args ...any) *redis.Cmd { calls++ - return pipelineResults(ctx, cmds, overloadedErr), nil + cmd := redis.NewCmd(ctx, args...) + cmd.SetErr(userErr) + return cmd } metrics := newTestMetrics() - d := NewDualWriter(primary, secondary, ProxyConfig{Mode: ModeDualWrite, SecondaryTimeout: 2 * time.Second}, metrics, newTestSentry(), testLogger) + d := NewDualWriter(primary, secondary, ProxyConfig{Mode: ModeDualWrite, SecondaryTimeout: time.Second}, metrics, newTestSentry(), testLogger) - d.writeSecondaryPipeline(context.Background(), [][]any{{"MULTI"}, {"SET", "k", "v"}, {"EXEC"}}) + d.writeSecondary(context.Background(), "EVALSHA", []any{[]byte("EVALSHA"), []byte("deadbeef"), []byte("0")}) - assert.Equal(t, maxServerOverloadedRetries+1, calls, - "secondary transaction replay must stop after maxServerOverloadedRetries+1 attempts") - assert.InDelta(t, 1, testutil.ToFloat64(metrics.SecondaryWriteErrors), 0.001, - "a persistent server-overloaded EXEC must still be reported as a secondary write error") - assert.InDelta(t, 1, testutil.ToFloat64( - metrics.SecondaryWriteErrorsByReason.WithLabelValues("EXEC", "busy")), 0.001) + assert.Equal(t, 1, calls, "user Redis not-leader text must not be retried") + assert.Equal(t, 0, secondary.RefreshCount(), "user Redis not-leader text must not force rediscovery") + assert.InDelta(t, 1, testutil.ToFloat64(metrics.SecondaryWriteErrors), 0.001) } -func TestDualWriter_writeSecondaryPipeline_DoesNotRetryCompactedExec(t *testing.T) { - primary := newMockBackend("primary") - primary.doFunc = makeCmd("OK", nil) +func TestIsElasticKVNotLeaderErrorClassifiesWireNotLeaderOnly(t *testing.T) { + tests := []struct { + name string + err error + want bool + }{ + {name: "redis wire notleader", err: testRedisErr("NOTLEADER raft engine: not leader"), want: true}, + {name: "grpc wrapped no redis err prefix", err: errors.New("rpc error: code = FailedPrecondition desc = raft engine: not leader"), want: true}, + {name: "grpc wrapped etcd sentinel no redis err prefix", err: errors.New("rpc error: code = FailedPrecondition desc = etcd raft engine is not leader"), want: true}, + {name: "bare leader not found", err: errors.New("leader not found"), want: true}, + {name: "lua user err bare raft phrase", err: testRedisErr("raft engine: not leader"), want: false}, + {name: "lua user err bare etcd phrase", err: testRedisErr("etcd raft engine is not leader"), want: false}, + {name: "lua user err exact raft phrase", err: testRedisErr("ERR raft engine: not leader"), want: false}, + {name: "lua user err etcd phrase", err: testRedisErr("ERR etcd raft engine is not leader"), want: false}, + {name: "lua user err leader not found", err: testRedisErr("ERR leader not found"), want: false}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + assert.Equal(t, tc.want, isElasticKVNotLeaderError(tc.err)) + }) + } +} - compactedErr := testRedisErr("rpc error: code = FailedPrecondition desc = read timestamp has been compacted") +func TestDualWriter_ReplaySecondaryPipeline_RecordsReplyErrors(t *testing.T) { + primary := newMockBackend("primary") secondary := newMockBackend("secondary") - var calls int - secondary.pipelineFunc = func(ctx context.Context, cmds [][]any) ([]*redis.Cmd, error) { - calls++ - return pipelineResults(ctx, cmds, compactedErr), nil + secondary.doFunc = func(ctx context.Context, args ...any) *redis.Cmd { + cmd := redis.NewCmd(ctx, args...) + cmd.SetErr(testRedisErr("NOTLEADER etcd raft engine is not leader")) + return cmd } metrics := newTestMetrics() d := NewDualWriter(primary, secondary, ProxyConfig{Mode: ModeDualWrite, SecondaryTimeout: time.Second}, metrics, newTestSentry(), testLogger) - d.writeSecondaryPipeline(context.Background(), [][]any{{"MULTI"}, {"SET", "k", "v"}, {"EXEC"}}) + d.replaySecondaryPipeline([][]any{{[]byte("MULTI")}, {[]byte("SET"), []byte("k"), []byte("v")}, {[]byte("EXEC")}}) + d.Close() - assert.Equal(t, 1, calls, "transaction replay must not retry errors that may occur after partial execution") assert.InDelta(t, 1, testutil.ToFloat64(metrics.SecondaryWriteErrors), 0.001) + assert.InDelta(t, 1, testutil.ToFloat64(metrics.SecondaryWriteErrorsByReason.WithLabelValues(cmdExec, "not_leader")), 0.001) } func TestDualWriter_writeSecondary_RetriesDoNotRepeatNoScriptProbe(t *testing.T) { @@ -1757,20 +1809,6 @@ type testRedisErr string func (e testRedisErr) Error() string { return string(e) } func (e testRedisErr) RedisError() {} -func pipelineResults(ctx context.Context, cmds [][]any, execErr error) []*redis.Cmd { - results := make([]*redis.Cmd, len(cmds)) - for i, args := range cmds { - cmd := redis.NewCmd(ctx, args...) - if i == len(cmds)-1 && execErr != nil { - cmd.SetErr(execErr) - } else { - cmd.SetVal("OK") - } - results[i] = cmd - } - return results -} - type mockRespWriter struct { writes []any } diff --git a/proxy/pubsub.go b/proxy/pubsub.go index 54ab079f0..221d085c1 100644 --- a/proxy/pubsub.go +++ b/proxy/pubsub.go @@ -470,9 +470,7 @@ func (s *pubsubSession) execTxn() { s.writeMu.Unlock() if s.proxy.dual.hasSecondaryWrite() { - s.proxy.dual.goAsync(func(ctx context.Context) { - s.proxy.dual.writeSecondaryPipeline(ctx, cmds) - }) + s.proxy.dual.replaySecondaryPipeline(cmds) } } diff --git a/proxy/raw_redis_proxy.go b/proxy/raw_redis_proxy.go new file mode 100644 index 000000000..f478d5480 --- /dev/null +++ b/proxy/raw_redis_proxy.go @@ -0,0 +1,156 @@ +package proxy + +import ( + "bufio" + "context" + "fmt" + "io" + "net" + "strconv" + "strings" + "sync" + "time" +) + +const ( + rawRedisCopyDirections = 2 + rawRedisHandshakeTimeout = 5 * time.Second +) + +func (p *ProxyServer) listenAndServeRawRedis(ctx context.Context) error { + p.shutdownCtx = ctx + + var lc net.ListenConfig + ln, err := lc.Listen(ctx, "tcp", p.cfg.ListenAddr) + if err != nil { + return fmt.Errorf("raw redis proxy listen: %w", err) + } + defer ln.Close() + + p.logger.Info("raw redis proxy starting", + "addr", p.cfg.ListenAddr, + "mode", p.cfg.Mode.String(), + "primary", p.cfg.PrimaryAddr, + "primary_db", p.cfg.PrimaryDB, + ) + + var wg sync.WaitGroup + defer wg.Wait() + + go func() { + <-ctx.Done() + p.logger.Info("shutting down raw redis proxy") + _ = ln.Close() + }() + + for { + client, err := ln.Accept() + if err != nil { + if ctx.Err() != nil { + return nil + } + return fmt.Errorf("raw redis proxy accept: %w", err) + } + wg.Add(1) + go func() { + defer wg.Done() + p.handleRawRedisConn(ctx, client) + }() + } +} + +func (p *ProxyServer) handleRawRedisConn(ctx context.Context, client net.Conn) { + p.metrics.ActiveConnections.Inc() + defer p.metrics.ActiveConnections.Dec() + defer client.Close() + + var dialer net.Dialer + upstream, err := dialer.DialContext(ctx, "tcp", p.cfg.PrimaryAddr) + if err != nil { + p.logger.Warn("raw redis upstream dial failed", "addr", p.cfg.PrimaryAddr, "err", err) + return + } + defer upstream.Close() + + upstreamReader := bufio.NewReader(upstream) + if err := p.prepareRawRedisUpstream(upstream, upstreamReader); err != nil { + p.logger.Warn("raw redis upstream prepare failed", "addr", p.cfg.PrimaryAddr, "err", err) + return + } + + stop := make(chan struct{}) + defer close(stop) + go func() { + select { + case <-ctx.Done(): + _ = client.Close() + _ = upstream.Close() + case <-stop: + } + }() + + errCh := make(chan struct{}, rawRedisCopyDirections) + go rawCopy(upstream, client, errCh) + go rawCopy(client, upstreamReader, errCh) + <-errCh +} + +func rawCopy(dst net.Conn, src io.Reader, done chan<- struct{}) { + _, _ = io.Copy(dst, src) + _ = dst.Close() + done <- struct{}{} +} + +func (p *ProxyServer) prepareRawRedisUpstream(upstream net.Conn, upstreamReader *bufio.Reader) error { + if err := upstream.SetDeadline(time.Now().Add(rawRedisHandshakeTimeout)); err != nil { + return fmt.Errorf("set handshake deadline: %w", err) + } + if p.cfg.PrimaryPassword != "" { + if err := rawRedisRoundTrip(upstream, upstreamReader, "AUTH", p.cfg.PrimaryPassword); err != nil { + return fmt.Errorf("auth: %w", err) + } + } + if p.cfg.PrimaryDB != 0 { + if err := rawRedisRoundTrip(upstream, upstreamReader, "SELECT", strconv.Itoa(p.cfg.PrimaryDB)); err != nil { + return fmt.Errorf("select db %d: %w", p.cfg.PrimaryDB, err) + } + } + if err := upstream.SetDeadline(time.Time{}); err != nil { + return fmt.Errorf("clear handshake deadline: %w", err) + } + return nil +} + +func rawRedisRoundTrip(w io.Writer, r *bufio.Reader, args ...string) error { + if _, err := w.Write(rawRedisCommand(args...)); err != nil { + return fmt.Errorf("write command: %w", err) + } + prefix, err := r.ReadByte() + if err != nil { + return fmt.Errorf("read reply prefix: %w", err) + } + line, err := r.ReadString('\n') + if err != nil { + return fmt.Errorf("read reply line: %w", err) + } + line = strings.TrimRight(line, "\r\n") + if prefix == '-' { + return fmt.Errorf("%s", line) + } + return nil +} + +func rawRedisCommand(args ...string) []byte { + var b strings.Builder + b.WriteByte('*') + b.WriteString(strconv.Itoa(len(args))) + b.WriteString("\r\n") + for _, arg := range args { + b.WriteByte('$') + b.WriteString(strconv.Itoa(len(arg))) + b.WriteString("\r\n") + b.WriteString(arg) + b.WriteString("\r\n") + } + return []byte(b.String()) +} diff --git a/proxy/raw_redis_proxy_test.go b/proxy/raw_redis_proxy_test.go new file mode 100644 index 000000000..285c1bd3d --- /dev/null +++ b/proxy/raw_redis_proxy_test.go @@ -0,0 +1,77 @@ +package proxy + +import ( + "bufio" + "io" + "net" + "strconv" + "strings" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestRawRedisCommand(t *testing.T) { + t.Parallel() + + require.Equal(t, + []byte("*2\r\n$4\r\nAUTH\r\n$6\r\nsecret\r\n"), + rawRedisCommand("AUTH", "secret"), + ) +} + +func TestPrepareRawRedisUpstreamAuthenticatesAndSelectsDB(t *testing.T) { + t.Parallel() + + client, server := net.Pipe() + defer client.Close() + defer server.Close() + + seen := make(chan []string, 2) + go func() { + reader := bufio.NewReader(server) + for i := 0; i < 2; i++ { + cmd, err := readRawRedisTestCommand(reader) + if err != nil { + return + } + seen <- cmd + _, _ = server.Write([]byte("+OK\r\n")) + } + }() + + p := &ProxyServer{cfg: ProxyConfig{PrimaryPassword: "secret", PrimaryDB: 1}} + err := p.prepareRawRedisUpstream(client, bufio.NewReader(client)) + require.NoError(t, err) + require.Equal(t, []string{"AUTH", "secret"}, <-seen) + require.Equal(t, []string{"SELECT", "1"}, <-seen) +} + +func readRawRedisTestCommand(r *bufio.Reader) ([]string, error) { + line, err := r.ReadString('\n') + if err != nil { + return nil, err + } + line = strings.TrimRight(line, "\r\n") + n, err := strconv.Atoi(strings.TrimPrefix(line, "*")) + if err != nil { + return nil, err + } + out := make([]string, 0, n) + for range n { + line, err = r.ReadString('\n') + if err != nil { + return nil, err + } + size, err := strconv.Atoi(strings.TrimPrefix(strings.TrimRight(line, "\r\n"), "$")) + if err != nil { + return nil, err + } + arg := make([]byte, size+2) + if _, err := io.ReadFull(r, arg); err != nil { + return nil, err + } + out = append(out, string(arg[:size])) + } + return out, nil +} diff --git a/store/hash_helpers.go b/store/hash_helpers.go index 0ee863a24..653e54a3b 100644 --- a/store/hash_helpers.go +++ b/store/hash_helpers.go @@ -160,40 +160,21 @@ func IsHashMetaDeltaKey(key []byte) bool { // ExtractHashUserKeyFromMeta extracts the logical user key from a hash meta key. func ExtractHashUserKeyFromMeta(key []byte) []byte { - trimmed := bytes.TrimPrefix(key, []byte(HashMetaPrefix)) - if len(trimmed) < wideColKeyLenSize { - return nil - } - ukLen := binary.BigEndian.Uint32(trimmed[:wideColKeyLenSize]) - if uint32(len(trimmed)) < uint32(wideColKeyLenSize)+ukLen { //nolint:gosec // wideColKeyLenSize fits in uint32 - return nil - } - return trimmed[wideColKeyLenSize : wideColKeyLenSize+ukLen] + return extractWideColumnUserKey(key, []byte(HashMetaPrefix), 0, true) } // ExtractHashUserKeyFromField extracts the logical user key from a hash field key. func ExtractHashUserKeyFromField(key []byte) []byte { - trimmed := bytes.TrimPrefix(key, []byte(HashFieldPrefix)) - if len(trimmed) < wideColKeyLenSize { - return nil - } - ukLen := binary.BigEndian.Uint32(trimmed[:wideColKeyLenSize]) - if uint32(len(trimmed)) < uint32(wideColKeyLenSize)+ukLen { //nolint:gosec // wideColKeyLenSize fits in uint32 - return nil - } - return trimmed[wideColKeyLenSize : wideColKeyLenSize+ukLen] + return extractWideColumnUserKey(key, []byte(HashFieldPrefix), 0, false) } // ExtractHashUserKeyFromDelta extracts the logical user key from a hash delta key. func ExtractHashUserKeyFromDelta(key []byte) []byte { - trimmed := bytes.TrimPrefix(key, []byte(HashMetaDeltaPrefix)) - minLen := wideColKeyLenSize + deltaKeyTSSize + deltaKeySeqSize - if len(trimmed) < minLen { - return nil - } - ukLen := binary.BigEndian.Uint32(trimmed[:wideColKeyLenSize]) - if uint32(len(trimmed)) < uint32(wideColKeyLenSize)+ukLen+uint32(deltaKeyTSSize+deltaKeySeqSize) { //nolint:gosec // constants fit in uint32 - return nil - } - return trimmed[wideColKeyLenSize : wideColKeyLenSize+ukLen] + return extractWideColumnUserKey(key, []byte(HashMetaDeltaPrefix), deltaKeyTSSize+deltaKeySeqSize, true) +} + +// ExtractHashUserKeyFromDeltaScanPrefix extracts the user key from a hash +// metadata delta scan start/prefix. +func ExtractHashUserKeyFromDeltaScanPrefix(key []byte) []byte { + return extractWideColumnUserKey(key, []byte(HashMetaDeltaPrefix), 0, false) } diff --git a/store/list_helpers.go b/store/list_helpers.go index e6c4ed964..971b47f6f 100644 --- a/store/list_helpers.go +++ b/store/list_helpers.go @@ -11,8 +11,13 @@ import ( // Delta/Claim key constants. const ( // ListMetaDeltaPrefix is the prefix for all list metadata delta keys. - // Layout: !lst|meta|d| - ListMetaDeltaPrefix = "!lst|meta|d|" + // Layout: !lst|delta| + ListMetaDeltaPrefix = "!lst|delta|" + + // LegacyListMetaDeltaPrefix is the pre-upgrade list metadata delta prefix. + // Writers use ListMetaDeltaPrefix, but readers and compactors keep scanning + // this prefix until old uncompacted deltas have been drained. + LegacyListMetaDeltaPrefix = "!lst|meta|d|" // ListClaimPrefix is the prefix for list claim keys used by POP operations. // Layout: !lst|claim| @@ -85,8 +90,27 @@ func ListMetaDeltaKey(userKey []byte, commitTS uint64, seqInTxn uint32) []byte { // ListMetaDeltaScanPrefix returns the prefix used to scan all delta keys for a userKey. func ListMetaDeltaScanPrefix(userKey []byte) []byte { - buf := make([]byte, 0, len(ListMetaDeltaPrefix)+wideColKeyLenSize+len(userKey)) - buf = append(buf, ListMetaDeltaPrefix...) + return listMetaDeltaScanPrefixFor(ListMetaDeltaPrefix, userKey) +} + +// LegacyListMetaDeltaScanPrefix returns the pre-upgrade delta scan prefix for a userKey. +func LegacyListMetaDeltaScanPrefix(userKey []byte) []byte { + return listMetaDeltaScanPrefixFor(LegacyListMetaDeltaPrefix, userKey) +} + +// ListMetaDeltaScanPrefixes returns all prefixes that may contain visible list +// metadata deltas for userKey. New writers emit only ListMetaDeltaPrefix, but +// upgrade reads must include LegacyListMetaDeltaPrefix until compaction drains it. +func ListMetaDeltaScanPrefixes(userKey []byte) [][]byte { + return [][]byte{ + ListMetaDeltaScanPrefix(userKey), + LegacyListMetaDeltaScanPrefix(userKey), + } +} + +func listMetaDeltaScanPrefixFor(prefix string, userKey []byte) []byte { + buf := make([]byte, 0, len(prefix)+wideColKeyLenSize+len(userKey)) + buf = append(buf, prefix...) var kl [4]byte binary.BigEndian.PutUint32(kl[:], uint32(len(userKey))) //nolint:gosec // len is bounded by max slice size buf = append(buf, kl[:]...) @@ -121,7 +145,7 @@ func ListClaimScanPrefix(userKey []byte) []byte { // IsListMetaDeltaKey reports whether the key is a list metadata delta key. func IsListMetaDeltaKey(key []byte) bool { - return bytes.HasPrefix(key, []byte(ListMetaDeltaPrefix)) + return ExtractListUserKeyFromDelta(key) != nil } // IsListClaimKey reports whether the key is a list claim key. @@ -131,22 +155,63 @@ func IsListClaimKey(key []byte) bool { // ExtractListUserKeyFromDelta extracts the logical user key from a list delta key. func ExtractListUserKeyFromDelta(key []byte) []byte { - trimmed := bytes.TrimPrefix(key, []byte(ListMetaDeltaPrefix)) - end, ok := listUserKeyEnd(trimmed, deltaKeyTSSize+deltaKeySeqSize) - if !ok { - return nil - } - return trimmed[wideColKeyLenSize:end] + return extractWideColumnUserKey(key, []byte(ListMetaDeltaPrefix), deltaKeyTSSize+deltaKeySeqSize, true) +} + +// ExtractLegacyListUserKeyFromDelta extracts the user key from an old +// !lst|meta|d| delta key. Callers that only have a key must be careful: this +// legacy layout overlaps with base !lst|meta| keys whose user key begins with +// d|. Prefer using it when the associated value is known to be a delta value. +func ExtractLegacyListUserKeyFromDelta(key []byte) []byte { + return extractWideColumnUserKey(key, []byte(LegacyListMetaDeltaPrefix), deltaKeyTSSize+deltaKeySeqSize, true) +} + +// ExtractListUserKeyFromDeltaScanPrefix extracts the user key from a new-list +// delta scan start/prefix. +func ExtractListUserKeyFromDeltaScanPrefix(key []byte) []byte { + return extractWideColumnUserKey(key, []byte(ListMetaDeltaPrefix), 0, false) +} + +// ExtractLegacyListUserKeyFromDeltaScanPrefix extracts the user key from a +// legacy-list delta scan start/prefix. +func ExtractLegacyListUserKeyFromDeltaScanPrefix(key []byte) []byte { + return extractWideColumnUserKey(key, []byte(LegacyListMetaDeltaPrefix), 0, false) } // ExtractListUserKeyFromClaim extracts the logical user key from a list claim key. func ExtractListUserKeyFromClaim(key []byte) []byte { - trimmed := bytes.TrimPrefix(key, []byte(ListClaimPrefix)) - end, ok := listUserKeyEnd(trimmed, sortableInt64Bytes) - if !ok { + return extractWideColumnUserKey(key, []byte(ListClaimPrefix), sortableInt64Bytes, true) +} + +// ExtractListUserKeyFromClaimScanPrefix extracts the user key from a list claim +// scan start/prefix. +func ExtractListUserKeyFromClaimScanPrefix(key []byte) []byte { + return extractWideColumnUserKey(key, []byte(ListClaimPrefix), 0, false) +} + +// IsListMetaDeltaValue reports whether value has the fixed delta encoding. +func IsListMetaDeltaValue(value []byte) bool { + return len(value) == listDeltaSizeBytes +} + +func extractWideColumnUserKey(key, prefix []byte, suffixLen uint64, exactLen bool) []byte { + if !bytes.HasPrefix(key, prefix) { + return nil + } + trimmed := key[len(prefix):] + if uint64(len(trimmed)) < uint64(wideColKeyLenSize)+suffixLen { + return nil + } + ukLen := binary.BigEndian.Uint32(trimmed[:wideColKeyLenSize]) + userEnd := uint64(wideColKeyLenSize) + uint64(ukLen) + minLen := userEnd + suffixLen + if minLen > uint64(len(trimmed)) { + return nil + } + if exactLen && minLen != uint64(len(trimmed)) { return nil } - return trimmed[wideColKeyLenSize:end] + return trimmed[wideColKeyLenSize:int(userEnd)] //nolint:gosec // userEnd is bounded by len(trimmed) above. } // ExtractListUserKeyFromDeltaScanKey extracts the logical user key from a @@ -166,24 +231,14 @@ func extractListUserKeyFromScanKey(key []byte, prefix []byte) []byte { return nil } trimmed := key[len(prefix):] - end, ok := listUserKeyEnd(trimmed, 0) - if !ok { + if len(trimmed) < wideColKeyLenSize { return nil } - return trimmed[wideColKeyLenSize:end] -} - -func listUserKeyEnd(trimmed []byte, suffixLen int) (int, bool) { - if len(trimmed) < wideColKeyLenSize+suffixLen { - return 0, false - } userKeyLen := binary.BigEndian.Uint32(trimmed[:wideColKeyLenSize]) - requiredTail := uint64(wideColKeyLenSize) + uint64(suffixLen) //nolint:gosec // suffixLen is one of this file's fixed encoded suffix widths. - available := uint64(len(trimmed)) - if requiredTail > available || uint64(userKeyLen) > available-requiredTail { - return 0, false + if uint32(len(trimmed)) < uint32(wideColKeyLenSize)+userKeyLen { //nolint:gosec // wideColKeyLenSize and encoded lengths fit in uint32 + return nil } - return wideColKeyLenSize + int(userKeyLen), true //nolint:gosec // userKeyLen is bounded by len(trimmed) above. + return trimmed[wideColKeyLenSize : wideColKeyLenSize+userKeyLen] } // PrefixScanEnd returns the exclusive end key for a prefix scan. diff --git a/store/list_helpers_test.go b/store/list_helpers_test.go index ea38a0fc3..dc2916f38 100644 --- a/store/list_helpers_test.go +++ b/store/list_helpers_test.go @@ -1,61 +1,79 @@ package store import ( - "bytes" "encoding/binary" - "math" "testing" + + "github.com/stretchr/testify/require" ) -func TestExtractListUserKeyFromScanKeyBoundsOverflow(t *testing.T) { +func TestExtractListUserKeyFromDeltaRequiresExactDeltaShape(t *testing.T) { + t.Parallel() + + userKey := []byte("d|list") + deltaKey := ListMetaDeltaKey(userKey, 42, 7) + require.True(t, IsListMetaDeltaKey(deltaKey)) + require.Equal(t, userKey, ExtractListUserKeyFromDelta(deltaKey)) + + baseMetaWithDeltaLookingUserKey := ListMetaKey(userKey) + require.False(t, IsListMetaDeltaKey(baseMetaWithDeltaLookingUserKey)) + require.Nil(t, ExtractListUserKeyFromDelta(baseMetaWithDeltaLookingUserKey)) + + trailingGarbage := append([]byte{}, deltaKey...) + trailingGarbage = append(trailingGarbage, 0) + require.False(t, IsListMetaDeltaKey(trailingGarbage)) + require.Nil(t, ExtractListUserKeyFromDelta(trailingGarbage)) + + require.False(t, IsListMetaDeltaKey([]byte("not-a-delta-key"))) + require.Nil(t, ExtractListUserKeyFromDelta([]byte("not-a-delta-key"))) +} + +func TestListMetaDeltaPrefixDoesNotOverlapBaseMetaKeys(t *testing.T) { t.Parallel() - var lenPrefix [wideColKeyLenSize]byte - binary.BigEndian.PutUint32(lenPrefix[:], math.MaxUint32) - - for _, tc := range []struct { - name string - key []byte - extract func([]byte) []byte - }{ - { - name: "delta scan", - key: append(append([]byte(nil), []byte(ListMetaDeltaPrefix)...), lenPrefix[:]...), - extract: ExtractListUserKeyFromDeltaScanKey, - }, - { - name: "claim scan", - key: append(append([]byte(nil), []byte(ListClaimPrefix)...), lenPrefix[:]...), - extract: ExtractListUserKeyFromClaimScanKey, - }, - { - name: "full delta", - key: append(append(append([]byte(nil), []byte(ListMetaDeltaPrefix)...), lenPrefix[:]...), make([]byte, deltaKeyTSSize+deltaKeySeqSize)...), - extract: ExtractListUserKeyFromDelta, - }, - { - name: "full claim", - key: append(append(append([]byte(nil), []byte(ListClaimPrefix)...), lenPrefix[:]...), make([]byte, sortableInt64Bytes)...), - extract: ExtractListUserKeyFromClaim, - }, - } { - t.Run(tc.name, func(t *testing.T) { - t.Parallel() - if got := tc.extract(tc.key); got != nil { - t.Fatalf("overflow user-key length: want nil, got %q", got) - } - }) - } + fakeUserKey := []byte("fake-user") + userKey := deltaLookingListMetaUserKey(fakeUserKey, 42, 7) + baseMeta := ListMetaKey(userKey) + deltaKey := ListMetaDeltaKey(userKey, 42, 7) + + require.True(t, IsListMetaKey(baseMeta)) + require.False(t, IsListMetaDeltaKey(baseMeta)) + require.Equal(t, userKey, ExtractListUserKey(baseMeta)) + require.Nil(t, ExtractListUserKeyFromDelta(baseMeta)) + + require.False(t, IsListMetaKey(deltaKey)) + require.True(t, IsListMetaDeltaKey(deltaKey)) + require.Equal(t, userKey, ExtractListUserKeyFromDelta(deltaKey)) } -func TestExtractListUserKeyFromScanKeyRoundTrip(t *testing.T) { +func TestLegacyListMetaDeltaHelpersScanOldPrefixWithoutReclassifyingNewDelta(t *testing.T) { t.Parallel() - userKey := []byte("list-user") - if got := ExtractListUserKeyFromDeltaScanKey(ListMetaDeltaScanPrefix(userKey)); !bytes.Equal(got, userKey) { - t.Fatalf("delta scan round trip: want %q, got %q", userKey, got) - } - if got := ExtractListUserKeyFromClaimScanKey(ListClaimScanPrefix(userKey)); !bytes.Equal(got, userKey) { - t.Fatalf("claim scan round trip: want %q, got %q", userKey, got) - } + userKey := []byte("list") + legacyPrefix := LegacyListMetaDeltaScanPrefix(userKey) + legacyKey := append(append([]byte(nil), legacyPrefix...), make([]byte, deltaKeyTSSize+deltaKeySeqSize)...) + + require.Equal(t, userKey, ExtractLegacyListUserKeyFromDeltaScanPrefix(legacyPrefix)) + require.Equal(t, userKey, ExtractLegacyListUserKeyFromDelta(legacyKey)) + require.False(t, IsListMetaDeltaKey(legacyKey), "legacy prefix is ambiguous with base !lst|meta| keys") + require.True(t, IsListMetaDeltaValue(MarshalListMetaDelta(ListMetaDelta{HeadDelta: 1, LenDelta: 2}))) + + metaValue, err := MarshalListMeta(ListMeta{Head: 1, Tail: 2, Len: 1}) + require.NoError(t, err) + require.False(t, IsListMetaDeltaValue(metaValue)) +} + +func deltaLookingListMetaUserKey(fakeUserKey []byte, commitTS uint64, seqInTxn uint32) []byte { + buf := make([]byte, 0, len("d|")+4+len(fakeUserKey)+8+4) + buf = append(buf, "d|"...) + var keyLen [4]byte + binary.BigEndian.PutUint32(keyLen[:], uint32(len(fakeUserKey))) //nolint:gosec // test data is small. + buf = append(buf, keyLen[:]...) + buf = append(buf, fakeUserKey...) + var ts [8]byte + binary.BigEndian.PutUint64(ts[:], commitTS) + buf = append(buf, ts[:]...) + var seq [4]byte + binary.BigEndian.PutUint32(seq[:], seqInTxn) + return append(buf, seq[:]...) } diff --git a/store/lsm_migration.go b/store/lsm_migration.go new file mode 100644 index 000000000..6b50651c5 --- /dev/null +++ b/store/lsm_migration.go @@ -0,0 +1,517 @@ +package store + +import ( + "bytes" + "context" + "math" + + "github.com/cockroachdb/errors" + "github.com/cockroachdb/pebble/v2" +) + +func (s *pebbleStore) ExportVersions(ctx context.Context, opts ExportVersionsOptions) (ExportVersionsResult, error) { + opts = normalizeExportVersionsOptions(opts) + pos, err := decodeExportCursorForOptions(opts) + if err != nil { + return ExportVersionsResult{}, err + } + if opts.MaxVersions <= 0 { + return ExportVersionsResult{Done: true}, nil + } + + s.dbMu.RLock() + defer s.dbMu.RUnlock() + + iter, err := s.db.NewIter(pebbleExportIterOptions(opts)) + if err != nil { + return ExportVersionsResult{}, errors.WithStack(err) + } + defer iter.Close() + + seek, err := pebbleExportSeekKey(opts, pos) + if err != nil { + return ExportVersionsResult{}, err + } + result := newExportVersionsResult(opts.MaxVersions) + if err := s.runPebbleExportLoop(ctx, iter, seek, opts, pos, &result); err != nil { + if errors.Is(err, errExportChunkFull) { + return result, nil + } + if errors.Is(err, errExportReachedEnd) { + result.Done = true + result.NextCursor = nil + return result, nil + } + return result, err + } + if err := iter.Error(); err != nil { + return ExportVersionsResult{}, errors.WithStack(err) + } + result.Done = true + result.NextCursor = nil + return result, nil +} + +func (s *pebbleStore) runPebbleExportLoop( + ctx context.Context, + iter *pebble.Iterator, + seek []byte, + opts ExportVersionsOptions, + pos exportCursorPosition, + result *ExportVersionsResult, +) error { + for ok := iter.SeekGE(seek); ok; { + advance, done, err := s.exportPebbleIteratorPosition(ctx, iter, opts, pos, result) + if err != nil { + return err + } + if !done { + return errExportChunkFull + } + if advance { + ok = iter.Next() + continue + } + ok = iter.Valid() + } + return nil +} + +func pebbleExportIterOptions(opts ExportVersionsOptions) *pebble.IterOptions { + return &pebble.IterOptions{} +} + +func pebbleExportSeekKey(opts ExportVersionsOptions, pos exportCursorPosition) ([]byte, error) { + if !pos.hasKey { + return encodeKey(opts.StartKey, math.MaxUint64), nil + } + if err := validateExportCursorRange(opts, pos); err != nil { + return nil, err + } + return encodeKey(pos.key, pos.commitTS), nil +} + +func (s *pebbleStore) exportPebbleIteratorPosition( + ctx context.Context, + iter *pebble.Iterator, + opts ExportVersionsOptions, + pos exportCursorPosition, + result *ExportVersionsResult, +) (advance bool, done bool, err error) { + if err := ctx.Err(); err != nil { + return false, false, errors.WithStack(err) + } + rawKey := iter.Key() + if isPebbleExportMetadataKey(rawKey) { + return true, true, nil + } + userKey, commitTS := decodeKeyView(rawKey) + if userKey == nil { + return true, true, nil + } + if pebbleExportCursorWholeKeySkipped(pos, userKey, commitTS) { + done := advancePebbleExportPastCurrentUserKey(iter, opts, userKey, pos.tag, result) + return false, done, nil + } + if pebbleExportCursorEqual(pos, userKey, commitTS) { + return true, true, nil + } + if skipped, err := s.skipPebbleExportKeyOutsideRange(iter, opts, userKey, commitTS, result); skipped || err != nil { + return false, true, err + } + if commitTS <= opts.MinCommitTSExclusive { + return s.skipPebbleExportVersionBelowMinTS(iter, opts, userKey, commitTS, result) + } + done, err = s.exportPebbleVersion(iter, opts, userKey, commitTS, result) + return true, done, err +} + +func isPebbleExportMetadataKey(rawKey []byte) bool { + return isPebbleMetaKey(rawKey) +} + +func (s *pebbleStore) skipPebbleExportKeyOutsideRange( + iter *pebble.Iterator, + opts ExportVersionsOptions, + userKey []byte, + commitTS uint64, + result *ExportVersionsResult, +) (bool, error) { + if opts.StartKey != nil && bytes.Compare(userKey, opts.StartKey) < 0 { + return true, skipPebbleExportWholeKey(iter, opts, userKey, commitTS, exportCursorTagSkippedKey, result) + } + if opts.EndKey == nil || bytes.Compare(userKey, opts.EndKey) < 0 { + return false, nil + } + if pebbleExportCanStopAtEndKey(opts.StartKey, opts.EndKey, userKey) { + return true, errExportReachedEnd + } + return true, skipPebbleExportWholeKey(iter, opts, userKey, commitTS, exportCursorTagSkippedKey, result) +} + +func skipPebbleExportWholeKey( + iter *pebble.Iterator, + opts ExportVersionsOptions, + userKey []byte, + commitTS uint64, + tag byte, + result *ExportVersionsResult, +) error { + rawValue := iter.Value() + result.ScannedBytes += versionExportSize(userKey, len(rawValue)) + result.NextCursor = encodeExportCursor(userKey, commitTS, tag) + if finishExportIfLimited(opts, result) { + result.Done = false + return errExportChunkFull + } + if !advancePebbleExportPastCurrentUserKey(iter, opts, userKey, tag, result) { + return errExportChunkFull + } + return nil +} + +func (s *pebbleStore) skipPebbleExportVersionBelowMinTS( + iter *pebble.Iterator, + opts ExportVersionsOptions, + userKey []byte, + commitTS uint64, + result *ExportVersionsResult, +) (advance bool, done bool, err error) { + rawValue := iter.Value() + result.ScannedBytes += versionExportSize(userKey, len(rawValue)) + result.NextCursor = encodeExportCursor(userKey, commitTS, exportCursorTagPrunedKey) + if finishExportIfLimited(opts, result) { + result.Done = false + return false, false, nil + } + return false, advancePebbleExportPastCurrentUserKey(iter, opts, userKey, exportCursorTagPrunedKey, result), nil +} + +func advancePebbleExportPastCurrentUserKey( + iter *pebble.Iterator, + opts ExportVersionsOptions, + userKey []byte, + tag byte, + result *ExportVersionsResult, +) bool { + userKey = bytes.Clone(userKey) + for iter.Next() { + currentUserKey, commitTS := decodeKeyView(iter.Key()) + if !bytes.Equal(currentUserKey, userKey) { + return true + } + rawValue := iter.Value() + result.ScannedBytes += versionExportSize(currentUserKey, len(rawValue)) + result.NextCursor = encodeExportCursor(currentUserKey, commitTS, tag) + if finishExportIfLimited(opts, result) { + result.Done = false + return false + } + } + return true +} + +func pebbleExportCanStopAtEndKey(startKey, endKey, userKey []byte) bool { + // Pebble orders userKey||invertedCommitTS physically. The empty logical key + // can therefore appear after non-empty keys; a leading range must keep + // scanning or it can silently omit that key. + if len(startKey) == 0 { + return false + } + for prefixLen := 1; prefixLen <= len(userKey); prefixLen++ { + prefix := userKey[:prefixLen] + if bytes.Compare(prefix, startKey) < 0 { + continue + } + if bytes.Compare(prefix, endKey) < 0 { + return false + } + } + return true +} + +func pebbleExportCursorEqual(pos exportCursorPosition, userKey []byte, commitTS uint64) bool { + return pos.hasKey && bytes.Equal(userKey, pos.key) && commitTS == pos.commitTS +} + +func pebbleExportCursorWholeKeySkipped(pos exportCursorPosition, userKey []byte, commitTS uint64) bool { + return pos.hasKey && + (pos.tag == exportCursorTagPrunedKey || pos.tag == exportCursorTagSkippedKey) && + bytes.Equal(userKey, pos.key) && + commitTS == pos.commitTS +} + +func (s *pebbleStore) exportPebbleVersion( + iter *pebble.Iterator, + opts ExportVersionsOptions, + userKey []byte, + commitTS uint64, + result *ExportVersionsResult, +) (bool, error) { + tag := exportCursorTagScanned + rawValue := iter.Value() + result.ScannedBytes += versionExportSize(userKey, len(rawValue)) + if shouldExportPebbleVersion(opts, userKey, commitTS) { + version, err := s.decodeExportedPebbleVersion(iter, userKey, commitTS, opts.KeyFamily) + if err != nil { + return false, err + } + if opts.AcceptVersion != nil && !opts.AcceptVersion(version.Key, version.Value) { + result.NextCursor = encodeExportCursor(userKey, commitTS, exportCursorTagScanned) + if finishExportIfLimited(opts, result) { + result.Done = false + return false, nil + } + return true, nil + } + result.Versions = append(result.Versions, version) + result.ExportedBytes += versionExportSize(userKey, len(version.Value)) + result.AcceptedRows++ + tag = exportCursorTagEmitted + } + result.NextCursor = encodeExportCursor(userKey, commitTS, tag) + if finishExportIfLimited(opts, result) { + result.Done = false + return false, nil + } + return true, nil +} + +func shouldExportPebbleVersion(opts ExportVersionsOptions, userKey []byte, commitTS uint64) bool { + if shouldSkipMigrationExportKey(userKey) { + return false + } + if opts.AcceptKey != nil && !opts.AcceptKey(userKey) { + return false + } + return opts.MaxCommitTSInclusive == 0 || commitTS <= opts.MaxCommitTSInclusive +} + +func (s *pebbleStore) decodeExportedPebbleVersion(iter *pebble.Iterator, userKey []byte, commitTS uint64, keyFamily uint32) (MVCCVersion, error) { + sv, err := decodeValue(iter.Value()) + if err != nil { + return MVCCVersion{}, errors.WithStack(err) + } + value, err := s.decryptForKey(iter.Key(), sv, sv.Value) + if err != nil { + return MVCCVersion{}, err + } + if sv.Tombstone { + value = nil + } + return MVCCVersion{ + Key: bytes.Clone(userKey), + CommitTS: commitTS, + Tombstone: sv.Tombstone, + Value: bytes.Clone(value), + KeyFamily: keyFamily, + ExpireAt: sv.ExpireAt, + }, nil +} + +func (s *pebbleStore) ImportVersions(ctx context.Context, opts ImportVersionsOptions) (ImportVersionsResult, error) { + s.dbMu.RLock() + defer s.dbMu.RUnlock() + + s.applyMu.Lock() + defer s.applyMu.Unlock() + + duplicate, ackedCursor, err := s.validatePebbleImportBatch(opts) + if err != nil { + return ImportVersionsResult{}, err + } + if duplicate { + return ImportVersionsResult{AckedCursor: ackedCursor, Duplicate: true}, nil + } + + batchMax := importBatchMaxTS(opts.Versions) + if err := s.commitPebbleImportBatch(opts, batchMax); err != nil { + return ImportVersionsResult{}, errors.WithStack(err) + } + s.log.InfoContext(ctx, "import_versions", + "job_id", opts.JobID, + "bracket_id", opts.BracketID, + "batch_seq", opts.BatchSeq, + "versions", len(opts.Versions), + "max_imported_ts", batchMax, + ) + return ImportVersionsResult{AckedCursor: bytes.Clone(opts.Cursor), MaxImportedTS: batchMax}, nil +} + +func (s *pebbleStore) validatePebbleImportBatch(opts ImportVersionsOptions) (bool, []byte, error) { + existing, hasExisting, err := s.readMigrationImportAck(opts.JobID, opts.BracketID) + if err != nil { + return false, nil, err + } + duplicate, err := validateNextImportBatch(existing, hasExisting, opts.BatchSeq) + if err != nil { + return false, nil, err + } + if duplicate { + return true, bytes.Clone(existing.cursor), nil + } + for _, version := range opts.Versions { + if err := validateImportVersion(version); err != nil { + return false, nil, err + } + } + return false, nil, nil +} + +func (s *pebbleStore) commitPebbleImportBatch(opts ImportVersionsOptions, batchMax uint64) error { + batch := s.db.NewBatch() + defer batch.Close() + if err := s.applyImportVersionsBatch(batch, opts.Versions); err != nil { + return err + } + if err := s.stageMigrationImportAck(batch, opts.JobID, opts.BracketID, migrationImportAck{ + batchSeq: opts.BatchSeq, + cursor: opts.Cursor, + }); err != nil { + return err + } + unlock, newLastTS, err := s.stageMigrationClockMetadataIfNeeded(batch, opts.JobID, batchMax) + if err != nil { + return err + } + defer unlock() + if err := batch.Commit(s.directApplyWriteOpts()); err != nil { + return errors.WithStack(err) + } + if batchMax > 0 { + s.lastCommitTS = newLastTS + } + return nil +} + +func (s *pebbleStore) stageMigrationImportAck(batch *pebble.Batch, jobID, bracketID uint64, ack migrationImportAck) error { + acks, err := s.readMigrationImportAcks() + if err != nil { + return err + } + acks[migrationAckID{jobID: jobID, bracketID: bracketID}] = migrationImportAck{ + batchSeq: ack.batchSeq, + cursor: bytes.Clone(ack.cursor), + } + return errors.WithStack(batch.Set(migrationAckMetaKeyBytes, encodeMigrationImportAcks(acks), nil)) +} + +func (s *pebbleStore) stageMigrationClockMetadataIfNeeded(batch *pebble.Batch, jobID, batchMax uint64) (func(), uint64, error) { + if batchMax == 0 { + return func() {}, 0, nil + } + return s.stageMigrationClockMetadata(batch, jobID, batchMax) +} + +func (s *pebbleStore) applyImportVersionsBatch(batch *pebble.Batch, versions []MVCCVersion) error { + for _, version := range versions { + k, err := encodePebbleUserVersionKey(version.Key, version.CommitTS) + if err != nil { + return err + } + var encoded []byte + if version.Tombstone { + encoded = encodeValue(nil, true, 0, encStateCleartext) + } else { + body, encState, err := s.encryptForKey(k, version.Value, version.ExpireAt, true) + if err != nil { + return err + } + encoded = encodeValue(body, false, version.ExpireAt, encState) + } + if err := batch.Set(k, encoded, nil); err != nil { + return errors.WithStack(err) + } + } + return nil +} + +func (s *pebbleStore) stageMigrationClockMetadata(batch *pebble.Batch, jobID, batchMax uint64) (func(), uint64, error) { + s.mtx.Lock() + unlock := func() { s.mtx.Unlock() } + newLastTS := s.lastCommitTS + if batchMax > newLastTS { + newLastTS = batchMax + } + if err := setPebbleUint64InBatch(batch, metaLastCommitTSBytes, newLastTS); err != nil { + unlock() + return nil, 0, err + } + floor, err := s.readMigrationHLCFloorLocked(jobID) + if err != nil { + unlock() + return nil, 0, err + } + if batchMax > floor { + floors, err := s.readMigrationHLCFloors() + if err != nil { + unlock() + return nil, 0, err + } + floors[jobID] = batchMax + if err := batch.Set(migrationHLCFloorMetaKeyBytes, encodeMigrationHLCFloors(floors), nil); err != nil { + unlock() + return nil, 0, errors.WithStack(err) + } + } + return unlock, newLastTS, nil +} + +func (s *pebbleStore) readMigrationImportAck(jobID, bracketID uint64) (migrationImportAck, bool, error) { + acks, err := s.readMigrationImportAcks() + if err != nil { + return migrationImportAck{}, false, err + } + ack, ok := acks[migrationAckID{jobID: jobID, bracketID: bracketID}] + return ack, ok, nil +} + +func (s *pebbleStore) readMigrationImportAcks() (map[migrationAckID]migrationImportAck, error) { + val, closer, err := s.db.Get(migrationAckMetaKeyBytes) + if err != nil { + if errors.Is(err, pebble.ErrNotFound) { + return make(map[migrationAckID]migrationImportAck), nil + } + return nil, errors.WithStack(err) + } + defer func() { _ = closer.Close() }() + acks, ok := decodeMigrationImportAcks(val) + if !ok { + return nil, errors.New("corrupt migration import ack metadata") + } + return acks, nil +} + +func (s *pebbleStore) readMigrationHLCFloorLocked(jobID uint64) (uint64, error) { + floors, err := s.readMigrationHLCFloors() + if err != nil { + return 0, err + } + return floors[jobID], nil +} + +func (s *pebbleStore) readMigrationHLCFloors() (map[uint64]uint64, error) { + val, closer, err := s.db.Get(migrationHLCFloorMetaKeyBytes) + if err != nil { + if errors.Is(err, pebble.ErrNotFound) { + return make(map[uint64]uint64), nil + } + return nil, errors.WithStack(err) + } + defer func() { _ = closer.Close() }() + floors, ok := decodeMigrationHLCFloors(val) + if !ok { + return nil, errors.New("corrupt migration HLC floor metadata") + } + return floors, nil +} + +func (s *pebbleStore) MigrationHLCFloor(_ context.Context, jobID uint64) (uint64, error) { + s.dbMu.RLock() + defer s.dbMu.RUnlock() + floor, err := s.readMigrationHLCFloorLocked(jobID) + if err != nil { + return 0, err + } + return floor, nil +} diff --git a/store/lsm_store.go b/store/lsm_store.go index 0ce5af988..1fc84daf9 100644 --- a/store/lsm_store.go +++ b/store/lsm_store.go @@ -665,24 +665,30 @@ func writePebbleUint64(db *pebble.DB, key []byte, value uint64, opts *pebble.Wri return errors.WithStack(db.Set(key, buf[:], opts)) } -// writeTempDBMetadata writes lastCommitTS and minRetainedTS atomically in a -// single synced batch so that both values are either fully durable or fully -// absent after a crash. This is critical for restore paths that swap a -// temporary Pebble directory into place: losing lastCommitTS could allow -// future commits to reuse timestamps, violating monotonic ordering. -func writeTempDBMetadata(db *pebble.DB, lastCommitTS, minRetainedTS uint64) error { +// writeTempDBMetadata writes restore metadata atomically in a single synced +// batch so every field is either fully durable or fully absent after a crash. +// This is critical for restore paths that swap a temporary Pebble directory +// into place: losing lastCommitTS could allow future commits to reuse +// timestamps, violating monotonic ordering. +func writeTempDBMetadata(db *pebble.DB, meta streamingMVCCRestoreMetadata) error { batch := db.NewBatch() defer func() { _ = batch.Close() }() var buf [timestampSize]byte - binary.LittleEndian.PutUint64(buf[:], lastCommitTS) + binary.LittleEndian.PutUint64(buf[:], meta.lastCommitTS) if err := batch.Set(metaLastCommitTSBytes, buf[:], nil); err != nil { return errors.WithStack(err) } - binary.LittleEndian.PutUint64(buf[:], minRetainedTS) + binary.LittleEndian.PutUint64(buf[:], meta.minRetainedTS) if err := batch.Set(metaMinRetainedTSBytes, buf[:], nil); err != nil { return errors.WithStack(err) } + if err := batch.Set(migrationAckMetaKeyBytes, encodeMigrationImportAcks(meta.migrationAcks), nil); err != nil { + return errors.WithStack(err) + } + if err := batch.Set(migrationHLCFloorMetaKeyBytes, encodeMigrationHLCFloors(meta.migrationHLCFloors), nil); err != nil { + return errors.WithStack(err) + } return errors.WithStack(batch.Commit(pebble.Sync)) } @@ -690,16 +696,39 @@ func isPebbleMetaKey(rawKey []byte) bool { return bytes.Equal(rawKey, metaLastCommitTSBytes) || bytes.Equal(rawKey, metaMinRetainedTSBytes) || bytes.Equal(rawKey, metaPendingMinRetainedTSBytes) || - bytes.Equal(rawKey, metaAppliedIndexBytes) -} - -func isPebbleOperationalKey(rawKey []byte) bool { - return isPebbleMetaKey(rawKey) || + bytes.Equal(rawKey, metaAppliedIndexBytes) || + isMigrationMetadataKey(rawKey) || isPebbleWriterRegistryKey(rawKey) } func isPebbleWriterRegistryKey(rawKey []byte) bool { - return encryption.IsRegistryKey(rawKey) + if !couldBePebbleWriterRegistryKey(rawKey) { + return false + } + _, _, err := encryption.DecodeRegistryKey(rawKey) + return err == nil +} + +func couldBePebbleWriterRegistryKey(rawKey []byte) bool { + const writerRegistrySuffixSize = 4 + 1 + 2 + prefix := encryption.WriterRegistryPrefix + return len(rawKey) == len(prefix)+writerRegistrySuffixSize && + bytes.HasPrefix(rawKey, prefix) && + rawKey[len(prefix)+4] == '|' +} + +var errMVCCMetadataKeyCollision = errors.New("store: mvcc encoded key collides with reserved pebble metadata key") + +func encodePebbleUserVersionKey(key []byte, commitTS uint64) ([]byte, error) { + encoded := encodeKey(key, commitTS) + if isPebbleMetaKey(encoded) { + return nil, errors.WithStack(errMVCCMetadataKeyCollision) + } + return encoded, nil +} + +func isPebbleOperationalKey(rawKey []byte) bool { + return isPebbleMetaKey(rawKey) } func (s *pebbleStore) findMaxCommitTS() (uint64, error) { @@ -1076,26 +1105,32 @@ func (s *pebbleStore) getAt(_ context.Context, key []byte, ts uint64) ([]byte, e // values, in which case the visibility checks below are operating // on authenticated bytes. func (s *pebbleStore) readVisibleVersion(iter *pebble.Iterator, key []byte, ts uint64) ([]byte, error) { - k := iter.Key() - userKey, _ := decodeKeyView(k) - if !bytes.Equal(userKey, key) { - return nil, ErrKeyNotFound - } - sv, err := decodeValue(iter.Value()) - if err != nil { - return nil, errors.WithStack(err) - } - plain, err := s.decryptForKey(k, sv, sv.Value) - if err != nil { - return nil, err - } - if sv.Tombstone { - return nil, ErrKeyNotFound - } - if sv.ExpireAt != 0 && sv.ExpireAt <= ts { - return nil, ErrKeyNotFound + for ; iter.Valid(); iter.Next() { + k := iter.Key() + if isPebbleMetaKey(k) { + continue + } + userKey, _ := decodeKeyView(k) + if !bytes.Equal(userKey, key) { + return nil, ErrKeyNotFound + } + sv, err := decodeValue(iter.Value()) + if err != nil { + return nil, errors.WithStack(err) + } + plain, err := s.decryptForKey(k, sv, sv.Value) + if err != nil { + return nil, err + } + if sv.Tombstone { + return nil, ErrKeyNotFound + } + if sv.ExpireAt != 0 && sv.ExpireAt <= ts { + return nil, ErrKeyNotFound + } + return plain, nil } - return plain, nil + return nil, ErrKeyNotFound } func (s *pebbleStore) GetAt(ctx context.Context, key []byte, ts uint64) ([]byte, error) { @@ -1248,7 +1283,11 @@ func (s *pebbleStore) CommittedVersionAt(_ context.Context, key []byte, commitTS } func (s *pebbleStore) committedVersionAtLocked(key []byte, commitTS uint64) (bool, error) { - _, closer, err := s.db.Get(encodeKey(key, commitTS)) + encoded := encodeKey(key, commitTS) + if isPebbleOperationalKey(encoded) { + return false, nil + } + _, closer, err := s.db.Get(encoded) if err != nil { if errors.Is(err, pebble.ErrNotFound) { return false, nil @@ -1319,7 +1358,7 @@ func (s *pebbleStore) seekToVisibleVersion(iter *pebble.Iterator, userKey []byte func (s *pebbleStore) skipToNextUserKey(iter *pebble.Iterator, userKey []byte) bool { for iter.Next() { rawKey := iter.Key() - if isPebbleMetaKey(rawKey) { + if isPebbleOperationalKey(rawKey) { continue } nextUserKey, _ := decodeKeyView(rawKey) @@ -1917,11 +1956,14 @@ func (s *pebbleStore) PutAt(ctx context.Context, key []byte, value []byte, commi if err := validateValueSize(value); err != nil { return err } + k, err := encodePebbleUserVersionKey(key, commitTS) + if err != nil { + return err + } s.dbMu.RLock() defer s.dbMu.RUnlock() commitTS = s.alignCommitTS(commitTS) - k := encodeKey(key, commitTS) // gateRegistration=true: PutAt is a direct (non-raft) write path. body, encState, err := s.encryptForKey(k, value, expireAt, true) if err != nil { @@ -1937,11 +1979,14 @@ func (s *pebbleStore) PutAt(ctx context.Context, key []byte, value []byte, commi } func (s *pebbleStore) DeleteAt(ctx context.Context, key []byte, commitTS uint64) error { + k, err := encodePebbleUserVersionKey(key, commitTS) + if err != nil { + return err + } s.dbMu.RLock() defer s.dbMu.RUnlock() commitTS = s.alignCommitTS(commitTS) - k := encodeKey(key, commitTS) v := encodeValue(nil, true, 0, encStateCleartext) if err := s.db.Set(k, v, pebble.NoSync); err != nil { @@ -1956,6 +2001,10 @@ func (s *pebbleStore) PutWithTTLAt(ctx context.Context, key []byte, value []byte } func (s *pebbleStore) ExpireAt(ctx context.Context, key []byte, expireAt uint64, commitTS uint64) error { + k, err := encodePebbleUserVersionKey(key, commitTS) + if err != nil { + return err + } s.dbMu.RLock() defer s.dbMu.RUnlock() @@ -1966,8 +2015,7 @@ func (s *pebbleStore) ExpireAt(ctx context.Context, key []byte, expireAt uint64, return err } - commitTS = s.alignCommitTS(commitTS) - k := encodeKey(key, commitTS) + s.alignCommitTS(commitTS) // gateRegistration=true: ExpireAt is a direct (non-raft) write path // that calls encryptForKey directly (it does not delegate to PutAt). body, encState, err := s.encryptForKey(k, val, expireAt, true) @@ -1994,12 +2042,16 @@ func (s *pebbleStore) latestCommitTS(_ context.Context, key []byte) (uint64, boo } defer iter.Close() - if iter.First() { + for ok := iter.First(); ok; ok = iter.Next() { k := iter.Key() + if isPebbleMetaKey(k) { + continue + } userKey, version := decodeKeyView(k) if bytes.Equal(userKey, key) { return version, true, nil } + return 0, false, nil } return 0, false, nil } @@ -2190,7 +2242,10 @@ func (s *pebbleStore) WriteConflictCount() uint64 { func (s *pebbleStore) applyMutationsBatch(b *pebble.Batch, mutations []*KVPairMutation, commitTS uint64, gateRegistration bool) error { for _, mut := range mutations { - k := encodeKey(mut.Key, commitTS) + k, err := encodePebbleUserVersionKey(mut.Key, commitTS) + if err != nil { + return err + } var v []byte switch mut.Op { @@ -2585,8 +2640,8 @@ func (s *pebbleStore) scanDeletePrefix(iter *pebble.Iterator, batch *pebble.Batc return err } if needsTombstone { - if err := batch.Set(encodeKey(userKey, commitTS), tombstoneVal, nil); err != nil { - return errors.WithStack(err) + if err := setDeletePrefixTombstone(batch, userKey, commitTS, tombstoneVal); err != nil { + return err } } if !s.skipToNextUserKey(iter, userKey) { @@ -2596,6 +2651,14 @@ func (s *pebbleStore) scanDeletePrefix(iter *pebble.Iterator, batch *pebble.Batc return nil } +func setDeletePrefixTombstone(batch *pebble.Batch, userKey []byte, commitTS uint64, tombstoneVal []byte) error { + k, err := encodePebbleUserVersionKey(userKey, commitTS) + if err != nil { + return err + } + return errors.WithStack(batch.Set(k, tombstoneVal, nil)) +} + type deletePrefixAction int const ( @@ -2995,8 +3058,12 @@ func flushSnapshotBatch(db *pebble.DB, batch **pebble.Batch, opts *pebble.WriteO } func setEncodedVersionInBatch(batch *pebble.Batch, key []byte, version VersionedValue) error { - deferred := batch.SetDeferred(encodedKeyLen(key), encodedValueLen(len(version.Value))) - fillEncodedKey(deferred.Key, key, version.TS) + encodedKey, err := encodePebbleUserVersionKey(key, version.TS) + if err != nil { + return err + } + deferred := batch.SetDeferred(len(encodedKey), encodedValueLen(len(version.Value))) + copy(deferred.Key, encodedKey) // MVCC snapshot format v2 does not carry encryption_state — Stage 8 of // the encryption rollout (per docs/design/2026_04_29_proposed...) bumps // the format to v3 to round-trip encrypted entries through this path. @@ -3020,6 +3087,33 @@ func writeRestoreEntry(r io.Reader, batch *pebble.Batch, keyBuf []byte, kLen, vL return errors.WithStack(deferred.Finish()) } +func restoreBatchLoopStep(r io.Reader, db *pebble.DB, batch **pebble.Batch, keyBuf *[]byte) (bool, error) { + kLen, vLen, eof, err := readRestoreEntry(r, keyBuf) + if err != nil { + return false, err + } + if eof { + return true, nil + } + if err := flushSnapshotBatchIfNeeded(db, batch, kLen, vLen); err != nil { + return false, err + } + if err := writeRestoreEntry(r, *batch, *keyBuf, kLen, vLen); err != nil { + return false, err + } + if snapshotBatchShouldFlush(*batch) { + return false, flushSnapshotBatch(db, batch, pebble.NoSync) + } + return false, nil +} + +func flushSnapshotBatchIfNeeded(db *pebble.DB, batch **pebble.Batch, kLen, vLen int) error { + if !(*batch).Empty() && (*batch).Len()+kLen+vLen >= snapshotBatchByteLimit { + return flushSnapshotBatch(db, batch, pebble.NoSync) + } + return nil +} + // restoreBatchLoopInto reads raw Pebble key-value entries from r and writes // them into db using batched commits. It is used for both the direct and the // temp-dir atomic native Pebble restore paths. @@ -3028,33 +3122,14 @@ func restoreBatchLoopInto(r io.Reader, db *pebble.DB) error { var keyBuf []byte // reused across entries to reduce per-entry allocations for { - kLen, vLen, eof, err := readRestoreEntry(r, &keyBuf) - if err != nil { - _ = batch.Close() - return err - } - if eof { + done, err := restoreBatchLoopStep(r, db, &batch, &keyBuf) + if done { break } - - // Flush before adding when the batch is non-empty and the anticipated - // entry size would push the batch over the byte limit. - if !batch.Empty() && batch.Len()+kLen+vLen >= snapshotBatchByteLimit { - if err := flushSnapshotBatch(db, &batch, pebble.NoSync); err != nil { - return err - } - } - - if err := writeRestoreEntry(r, batch, keyBuf, kLen, vLen); err != nil { + if err != nil { _ = batch.Close() return err } - - if snapshotBatchShouldFlush(batch) { - if err := flushSnapshotBatch(db, &batch, pebble.NoSync); err != nil { - return err - } - } } return commitSnapshotBatch(batch, pebble.Sync) } @@ -3282,22 +3357,35 @@ func writeNativeSnapshotToTempDir(r io.Reader, tmpDir string, ts uint64) error { // Entries are written to a temporary Pebble directory and only swapped into // place after the CRC32 checksum is verified, preserving the existing store // on failure. -func readStreamingMVCCRestoreHeader(r io.Reader) (io.Reader, hash.Hash32, uint32, uint64, uint64, error) { - expectedChecksum, err := readMVCCSnapshotHeader(r) +type streamingMVCCRestoreMetadata struct { + lastCommitTS uint64 + minRetainedTS uint64 + migrationAcks map[migrationAckID]migrationImportAck + migrationHLCFloors map[uint64]uint64 +} + +func readStreamingMVCCRestoreHeader(r io.Reader) (io.Reader, hash.Hash32, uint32, streamingMVCCRestoreMetadata, error) { + version, expectedChecksum, err := readMVCCSnapshotHeader(r) if err != nil { - return nil, nil, 0, 0, 0, err + return nil, nil, 0, streamingMVCCRestoreMetadata{}, err } hash := crc32.NewIEEE() body := io.TeeReader(r, hash) - lastCommitTS, minRetainedTS, err := readMVCCSnapshotMetadata(body) + lastCommitTS, minRetainedTS, migrationAcks, migrationHLCFloors, err := readMVCCSnapshotMetadata(body, version) if err != nil { - return nil, nil, 0, 0, 0, err + return nil, nil, 0, streamingMVCCRestoreMetadata{}, err + } + meta := streamingMVCCRestoreMetadata{ + lastCommitTS: lastCommitTS, + minRetainedTS: minRetainedTS, + migrationAcks: migrationAcks, + migrationHLCFloors: migrationHLCFloors, } - return body, hash, expectedChecksum, lastCommitTS, minRetainedTS, nil + return body, hash, expectedChecksum, meta, nil } -func writeStreamingMVCCRestoreTempDB(dir string, body io.Reader, hash hash.Hash32, expectedChecksum uint32, lastCommitTS uint64, minRetainedTS uint64) (string, error) { +func writeStreamingMVCCRestoreTempDB(dir string, body io.Reader, hash hash.Hash32, expectedChecksum uint32, meta streamingMVCCRestoreMetadata) (string, error) { tmpDir := filepath.Clean(dir) + ".restore-tmp" if err := os.RemoveAll(tmpDir); err != nil { return "", errors.WithStack(err) @@ -3325,7 +3413,7 @@ func writeStreamingMVCCRestoreTempDB(dir string, body io.Reader, hash hash.Hash3 cleanupTmp() return "", errors.WithStack(ErrInvalidChecksum) } - if err := writeTempDBMetadata(tmpDB, lastCommitTS, minRetainedTS); err != nil { + if err := writeTempDBMetadata(tmpDB, meta); err != nil { cleanupTmp() return "", err } @@ -3337,12 +3425,12 @@ func writeStreamingMVCCRestoreTempDB(dir string, body io.Reader, hash hash.Hash3 } func (s *pebbleStore) restoreFromStreamingMVCC(r io.Reader) error { - body, hash, expectedChecksum, lastCommitTS, minRetainedTS, err := readStreamingMVCCRestoreHeader(r) + body, hash, expectedChecksum, meta, err := readStreamingMVCCRestoreHeader(r) if err != nil { return err } - tmpDir, err := writeStreamingMVCCRestoreTempDB(s.dir, body, hash, expectedChecksum, lastCommitTS, minRetainedTS) + tmpDir, err := writeStreamingMVCCRestoreTempDB(s.dir, body, hash, expectedChecksum, meta) if err != nil { return err } diff --git a/store/lsm_store_test.go b/store/lsm_store_test.go index ecea4526a..dbd6605e1 100644 --- a/store/lsm_store_test.go +++ b/store/lsm_store_test.go @@ -890,6 +890,14 @@ func TestPebbleStore_RestoreFromStreamingMVCC(t *testing.T) { require.NoError(t, src.PutAt(ctx, []byte("key1"), []byte("val1-updated"), 20, 0)) require.NoError(t, src.DeleteAt(ctx, []byte("key2"), 15)) require.NoError(t, src.PutWithTTLAt(ctx, []byte("key3"), []byte("val3"), 30, 9999)) + _, err := src.ImportVersions(ctx, ImportVersionsOptions{ + JobID: 7, + BracketID: 3, + BatchSeq: 1, + Cursor: []byte("streaming-metadata"), + Versions: []MVCCVersion{{Key: []byte("imported"), CommitTS: 50, Value: []byte("v50")}}, + }) + require.NoError(t, err) snap, err := src.Snapshot() require.NoError(t, err) @@ -926,6 +934,22 @@ func TestPebbleStore_RestoreFromStreamingMVCC(t *testing.T) { require.NoError(t, err) assert.Equal(t, []byte("val3"), val) + val, err = dst.GetAt(ctx, []byte("imported"), 50) + require.NoError(t, err) + assert.Equal(t, []byte("v50"), val) + floor, err := dst.MigrationHLCFloor(ctx, 7) + require.NoError(t, err) + assert.Equal(t, uint64(50), floor) + res, err := dst.ImportVersions(ctx, ImportVersionsOptions{ + JobID: 7, + BracketID: 3, + BatchSeq: 1, + Cursor: []byte("different"), + }) + require.NoError(t, err) + assert.True(t, res.Duplicate) + assert.Equal(t, []byte("streaming-metadata"), res.AckedCursor) + assert.Equal(t, src.LastCommitTS(), dst.LastCommitTS()) } diff --git a/store/migration_versions.go b/store/migration_versions.go new file mode 100644 index 000000000..4e42af1db --- /dev/null +++ b/store/migration_versions.go @@ -0,0 +1,528 @@ +package store + +import ( + "bytes" + "context" + "encoding/binary" + "sort" + + "github.com/cockroachdb/errors" + "github.com/emirpasic/gods/maps/treemap" +) + +const ( + exportCursorTagEmitted byte = iota + exportCursorTagScanned + exportCursorTagPrunedKey + exportCursorTagSkippedKey + + migrationAckMetaKey = "_migack" + migrationHLCFloorMetaKey = "_mighlc" + migrationMetadataVersion = 1 + + migrationAckPrefix = "!migstage|ack|" + migrationUint64Bytes = 8 + exportVersionSizeOverhead = 24 + defaultSparseExportMaxScannedBytes = 1 << 20 +) + +var ( + migrationAckMetaKeyBytes = []byte(migrationAckMetaKey) + migrationHLCFloorMetaKeyBytes = []byte(migrationHLCFloorMetaKey) +) + +type exportCursorPosition struct { + key []byte + commitTS uint64 + tag byte + hasKey bool +} + +type migrationAckID struct { + jobID uint64 + bracketID uint64 +} + +type migrationImportAck struct { + batchSeq uint64 + cursor []byte +} + +func encodeExportCursor(key []byte, commitTS uint64, tag byte) []byte { + var buf []byte + buf = binary.AppendUvarint(buf, lenAsUint64(len(key))) + buf = append(buf, key...) + buf = binary.AppendUvarint(buf, commitTS) + buf = append(buf, tag) + return buf +} + +func decodeExportCursor(cursor []byte) (exportCursorPosition, error) { + if len(cursor) == 0 { + return exportCursorPosition{}, nil + } + keyLen, n := binary.Uvarint(cursor) + if n <= 0 { + return exportCursorPosition{}, errors.WithStack(ErrInvalidExportCursor) + } + rest := cursor[n:] + if keyLen > lenAsUint64(len(rest)) { + return exportCursorPosition{}, errors.WithStack(ErrInvalidExportCursor) + } + key := bytes.Clone(rest[:keyLen]) + rest = rest[keyLen:] + commitTS, n := binary.Uvarint(rest) + if n <= 0 { + return exportCursorPosition{}, errors.WithStack(ErrInvalidExportCursor) + } + rest = rest[n:] + if len(rest) != 1 { + return exportCursorPosition{}, errors.WithStack(ErrInvalidExportCursor) + } + tag := rest[0] + if tag != exportCursorTagEmitted && + tag != exportCursorTagScanned && + tag != exportCursorTagPrunedKey && + tag != exportCursorTagSkippedKey { + return exportCursorPosition{}, errors.WithStack(ErrInvalidExportCursor) + } + return exportCursorPosition{key: key, commitTS: commitTS, tag: tag, hasKey: true}, nil +} + +func decodeExportCursorForOptions(opts ExportVersionsOptions) (exportCursorPosition, error) { + pos, err := decodeExportCursor(opts.Cursor) + if err != nil { + return exportCursorPosition{}, err + } + if err := validateExportCursorRange(opts, pos); err != nil { + return exportCursorPosition{}, err + } + return pos, nil +} + +func validateExportCursorRange(opts ExportVersionsOptions, pos exportCursorPosition) error { + if !pos.hasKey { + return nil + } + if pos.tag == exportCursorTagSkippedKey { + if !exportSkippedCursorOutsideRange(opts, pos.key) { + return errors.WithStack(ErrInvalidExportCursor) + } + return nil + } + if opts.StartKey != nil && bytes.Compare(pos.key, opts.StartKey) < 0 { + return errors.WithStack(ErrInvalidExportCursor) + } + if opts.EndKey != nil && bytes.Compare(pos.key, opts.EndKey) >= 0 { + return errors.WithStack(ErrInvalidExportCursor) + } + return nil +} + +func exportSkippedCursorOutsideRange(opts ExportVersionsOptions, key []byte) bool { + return (opts.StartKey != nil && bytes.Compare(key, opts.StartKey) < 0) || + (opts.EndKey != nil && bytes.Compare(key, opts.EndKey) >= 0) +} + +func normalizeExportVersionsOptions(opts ExportVersionsOptions) ExportVersionsOptions { + if opts.EndKey != nil && len(opts.EndKey) == 0 { + opts.EndKey = nil + } + if exportUsesSparseScanBudget(opts) && opts.MaxScannedBytes == 0 { + opts.MaxScannedBytes = defaultSparseExportMaxScannedBytes + } + return opts +} + +func exportUsesSparseScanBudget(opts ExportVersionsOptions) bool { + return opts.AcceptKey != nil || + opts.AcceptVersion != nil || + opts.MaxCommitTSInclusive != 0 || + opts.MinCommitTSExclusive != 0 || + opts.StartKey != nil || + opts.EndKey != nil +} + +func isMigrationMetadataKey(rawKey []byte) bool { + return bytes.Equal(rawKey, migrationAckMetaKeyBytes) || + bytes.Equal(rawKey, migrationHLCFloorMetaKeyBytes) +} + +func encodeMigrationImportAcks(acks map[migrationAckID]migrationImportAck) []byte { + ids := make([]migrationAckID, 0, len(acks)) + for id := range acks { + ids = append(ids, id) + } + sort.Slice(ids, func(i, j int) bool { + if ids[i].jobID != ids[j].jobID { + return ids[i].jobID < ids[j].jobID + } + return ids[i].bracketID < ids[j].bracketID + }) + + buf := make([]byte, 0, 1+binary.MaxVarintLen64+len(ids)*(3*migrationUint64Bytes+binary.MaxVarintLen64)) + buf = append(buf, migrationMetadataVersion) + buf = binary.AppendUvarint(buf, lenAsUint64(len(ids))) + for _, id := range ids { + ack := acks[id] + buf = binary.BigEndian.AppendUint64(buf, id.jobID) + buf = binary.BigEndian.AppendUint64(buf, id.bracketID) + buf = binary.BigEndian.AppendUint64(buf, ack.batchSeq) + buf = binary.AppendUvarint(buf, lenAsUint64(len(ack.cursor))) + buf = append(buf, ack.cursor...) + } + return buf +} + +func decodeMigrationImportAcks(data []byte) (map[migrationAckID]migrationImportAck, bool) { + if len(data) == 0 || data[0] != migrationMetadataVersion { + return nil, false + } + rest := data[1:] + count, n := binary.Uvarint(rest) + if n <= 0 { + return nil, false + } + rest = rest[n:] + acks := make(map[migrationAckID]migrationImportAck) + for i := uint64(0); i < count; i++ { + if len(rest) < 3*migrationUint64Bytes { + return nil, false + } + id := migrationAckID{ + jobID: binary.BigEndian.Uint64(rest[:migrationUint64Bytes]), + bracketID: binary.BigEndian.Uint64(rest[migrationUint64Bytes : 2*migrationUint64Bytes]), + } + ack := migrationImportAck{batchSeq: binary.BigEndian.Uint64(rest[2*migrationUint64Bytes : 3*migrationUint64Bytes])} + rest = rest[3*migrationUint64Bytes:] + cursorLen, n := binary.Uvarint(rest) + if n <= 0 { + return nil, false + } + rest = rest[n:] + if cursorLen > lenAsUint64(len(rest)) { + return nil, false + } + ack.cursor = bytes.Clone(rest[:cursorLen]) + rest = rest[cursorLen:] + acks[id] = ack + } + return acks, len(rest) == 0 +} + +func encodeMigrationHLCFloors(floors map[uint64]uint64) []byte { + jobIDs := make([]uint64, 0, len(floors)) + for jobID := range floors { + jobIDs = append(jobIDs, jobID) + } + sort.Slice(jobIDs, func(i, j int) bool { return jobIDs[i] < jobIDs[j] }) + + buf := make([]byte, 0, 1+binary.MaxVarintLen64+len(jobIDs)*2*migrationUint64Bytes) + buf = append(buf, migrationMetadataVersion) + buf = binary.AppendUvarint(buf, lenAsUint64(len(jobIDs))) + for _, jobID := range jobIDs { + buf = binary.BigEndian.AppendUint64(buf, jobID) + buf = binary.BigEndian.AppendUint64(buf, floors[jobID]) + } + return buf +} + +func decodeMigrationHLCFloors(data []byte) (map[uint64]uint64, bool) { + if len(data) == 0 || data[0] != migrationMetadataVersion { + return nil, false + } + rest := data[1:] + count, n := binary.Uvarint(rest) + if n <= 0 { + return nil, false + } + rest = rest[n:] + floors := make(map[uint64]uint64) + for i := uint64(0); i < count; i++ { + if len(rest) < 2*migrationUint64Bytes { + return nil, false + } + jobID := binary.BigEndian.Uint64(rest[:migrationUint64Bytes]) + floor := binary.BigEndian.Uint64(rest[migrationUint64Bytes : 2*migrationUint64Bytes]) + rest = rest[2*migrationUint64Bytes:] + floors[jobID] = floor + } + return floors, len(rest) == 0 +} + +func validateImportVersion(version MVCCVersion) error { + if version.CommitTS == 0 { + return errors.New("migration import version has zero commit_ts") + } + if version.Tombstone { + if version.ExpireAt != 0 { + return errors.New("migration import tombstone carries expire_at") + } + if len(version.Value) != 0 { + return errors.New("migration import tombstone carries value") + } + return nil + } + return validateValueSize(version.Value) +} + +func versionExportSize(key []byte, valueLen int) uint64 { + return lenAsUint64(len(key)) + lenAsUint64(valueLen) + exportVersionSizeOverhead +} + +func lenAsUint64(n int) uint64 { + if n <= 0 { + return 0 + } + return uint64(n) //nolint:gosec // slice lengths are non-negative and bounded by addressable memory. +} + +func importBatchMaxTS(versions []MVCCVersion) uint64 { + var maxTS uint64 + for _, version := range versions { + if version.CommitTS > maxTS { + maxTS = version.CommitTS + } + } + return maxTS +} + +func validateNextImportBatch(existing migrationImportAck, hasExisting bool, batchSeq uint64) (duplicate bool, err error) { + if hasExisting { + if batchSeq <= existing.batchSeq { + return true, nil + } + if batchSeq != existing.batchSeq+1 { + return false, errors.WithStack(ErrImportBatchGap) + } + return false, nil + } + if batchSeq != 1 { + return false, errors.WithStack(ErrImportBatchGap) + } + return false, nil +} + +func (s *mvccStore) ExportVersions(ctx context.Context, opts ExportVersionsOptions) (ExportVersionsResult, error) { + opts = normalizeExportVersionsOptions(opts) + pos, err := decodeExportCursorForOptions(opts) + if err != nil { + return ExportVersionsResult{}, err + } + if opts.MaxVersions <= 0 { + return ExportVersionsResult{Done: true}, nil + } + + s.mtx.RLock() + defer s.mtx.RUnlock() + + result := newExportVersionsResult(opts.MaxVersions) + it := s.tree.Iterator() + if !s.seekMemoryExportStart(&it, opts.StartKey, pos) { + result.Done = true + return result, nil + } + + for ok := true; ok; ok = it.Next() { + key, ok := it.Key().([]byte) + if err := checkExportKey(ctx, key, ok, opts.EndKey); err != nil { + if errors.Is(err, errExportReachedEnd) { + result.Done = true + result.NextCursor = nil + return result, nil + } + return ExportVersionsResult{}, err + } + if !ok { + continue + } + done, err := exportMemoryIteratorKey(ctx, opts, pos, key, it.Value(), &result) + if err != nil || !done { + return result, err + } + } + result.Done = true + result.NextCursor = nil + return result, nil +} + +var errExportReachedEnd = errors.New("export reached end") +var errExportChunkFull = errors.New("export chunk full") + +func checkExportKey(ctx context.Context, key []byte, keyOK bool, end []byte) error { + if err := ctx.Err(); err != nil { + return errors.WithStack(err) + } + if !keyOK { + return nil + } + if end != nil && bytes.Compare(key, end) >= 0 { + return errExportReachedEnd + } + return nil +} + +func newExportVersionsResult(maxVersions int) ExportVersionsResult { + return ExportVersionsResult{ + Versions: make([]MVCCVersion, 0, min(maxVersions, scanResultCapacityLimit)), + } +} + +func (s *mvccStore) seekMemoryExportStart(it *treemap.Iterator, startKey []byte, pos exportCursorPosition) bool { + if pos.hasKey { + return seekForwardIteratorStart(s.tree, it, pos.key) + } + return seekForwardIteratorStart(s.tree, it, startKey) +} + +func exportMemoryIteratorKey( + ctx context.Context, + opts ExportVersionsOptions, + pos exportCursorPosition, + key []byte, + value any, + result *ExportVersionsResult, +) (bool, error) { + versions, _ := value.([]VersionedValue) + if pos.hasKey && pos.tag == exportCursorTagPrunedKey && bytes.Equal(key, pos.key) { + return true, nil + } + cursorCommitTS := uint64(0) + if pos.hasKey && bytes.Equal(key, pos.key) { + cursorCommitTS = pos.commitTS + } + return exportMemoryVersionsForKey(ctx, opts, cursorCommitTS, key, versions, result) +} + +func finishExportIfLimited(opts ExportVersionsOptions, result *ExportVersionsResult) bool { + return len(result.Versions) >= opts.MaxVersions || + (opts.MaxBytes > 0 && result.ExportedBytes >= opts.MaxBytes) || + (opts.MaxScannedBytes > 0 && result.ScannedBytes >= opts.MaxScannedBytes) +} + +func appendMemoryExportVersion(opts ExportVersionsOptions, key []byte, version VersionedValue, result *ExportVersionsResult) byte { + if shouldSkipMigrationExportKey(key) { + return exportCursorTagScanned + } + if opts.AcceptKey != nil && !opts.AcceptKey(key) { + return exportCursorTagScanned + } + if opts.AcceptVersion != nil && !opts.AcceptVersion(key, version.Value) { + return exportCursorTagScanned + } + if opts.MaxCommitTSInclusive != 0 && version.TS > opts.MaxCommitTSInclusive { + return exportCursorTagScanned + } + result.Versions = append(result.Versions, MVCCVersion{ + Key: bytes.Clone(key), + CommitTS: version.TS, + Tombstone: version.Tombstone, + Value: bytes.Clone(version.Value), + KeyFamily: opts.KeyFamily, + ExpireAt: version.ExpireAt, + }) + result.ExportedBytes += versionExportSize(key, len(version.Value)) + result.AcceptedRows++ + return exportCursorTagEmitted +} + +func shouldSkipMigrationExportKey(key []byte) bool { + return bytes.HasPrefix(key, txnLockKeyPrefix) +} + +func finishMemoryExportPosition(opts ExportVersionsOptions, key []byte, version VersionedValue, tag byte, result *ExportVersionsResult) bool { + result.ScannedBytes += versionExportSize(key, len(version.Value)) + result.NextCursor = encodeExportCursor(key, version.TS, tag) + if finishExportIfLimited(opts, result) { + result.Done = false + return false + } + return true +} + +func shouldSkipMemoryVersion(cursorCommitTS uint64, version VersionedValue) bool { + return cursorCommitTS != 0 && version.TS >= cursorCommitTS +} + +func exportMemoryVersion(opts ExportVersionsOptions, cursorCommitTS uint64, key []byte, version VersionedValue, result *ExportVersionsResult) bool { + if shouldSkipMemoryVersion(cursorCommitTS, version) { + return true + } + tag := appendMemoryExportVersion(opts, key, version, result) + return finishMemoryExportPosition(opts, key, version, tag, result) +} + +func exportMemoryVersionsForKey( + ctx context.Context, + opts ExportVersionsOptions, + cursorCommitTS uint64, + key []byte, + versions []VersionedValue, + result *ExportVersionsResult, +) (bool, error) { + for i := len(versions) - 1; i >= 0; i-- { + if err := ctx.Err(); err != nil { + return false, errors.WithStack(err) + } + if shouldSkipMemoryVersion(cursorCommitTS, versions[i]) { + continue + } + if versions[i].TS <= opts.MinCommitTSExclusive { + if !finishMemoryExportPosition(opts, key, versions[i], exportCursorTagPrunedKey, result) { + return false, nil + } + return true, nil + } + if !exportMemoryVersion(opts, cursorCommitTS, key, versions[i], result) { + return !finishExportIfLimited(opts, result), nil + } + if !result.Done && finishExportIfLimited(opts, result) { + return false, nil + } + } + return true, nil +} + +func (s *mvccStore) ImportVersions(_ context.Context, opts ImportVersionsOptions) (ImportVersionsResult, error) { + s.mtx.Lock() + defer s.mtx.Unlock() + + id := migrationAckID{jobID: opts.JobID, bracketID: opts.BracketID} + existing, hasExisting := s.migrationAcks[id] + duplicate, err := validateNextImportBatch(existing, hasExisting, opts.BatchSeq) + if err != nil { + return ImportVersionsResult{}, err + } + if duplicate { + return ImportVersionsResult{AckedCursor: bytes.Clone(existing.cursor), Duplicate: true}, nil + } + + for _, version := range opts.Versions { + if err := validateImportVersion(version); err != nil { + return ImportVersionsResult{}, err + } + } + for _, version := range opts.Versions { + if version.Tombstone { + s.deleteVersionLocked(version.Key, version.CommitTS) + continue + } + s.putVersionLocked(version.Key, version.Value, version.CommitTS, version.ExpireAt) + } + + batchMax := importBatchMaxTS(opts.Versions) + if batchMax > s.lastCommitTS { + s.lastCommitTS = batchMax + } + if batchMax > s.migrationHLCFloors[opts.JobID] { + s.migrationHLCFloors[opts.JobID] = batchMax + } + s.migrationAcks[id] = migrationImportAck{batchSeq: opts.BatchSeq, cursor: bytes.Clone(opts.Cursor)} + return ImportVersionsResult{AckedCursor: bytes.Clone(opts.Cursor), MaxImportedTS: batchMax}, nil +} + +func (s *mvccStore) MigrationHLCFloor(_ context.Context, jobID uint64) (uint64, error) { + s.mtx.RLock() + defer s.mtx.RUnlock() + return s.migrationHLCFloors[jobID], nil +} diff --git a/store/migration_versions_test.go b/store/migration_versions_test.go new file mode 100644 index 000000000..3dc418888 --- /dev/null +++ b/store/migration_versions_test.go @@ -0,0 +1,918 @@ +package store + +import ( + "bytes" + "context" + "encoding/binary" + "math" + "os" + "testing" + + "github.com/bootjp/elastickv/internal/encryption" + "github.com/stretchr/testify/require" +) + +func runMigrationStoreSuite(t *testing.T, test func(t *testing.T, st MVCCStore)) { + t.Helper() + t.Run("memory", func(t *testing.T) { + test(t, NewMVCCStore()) + }) + t.Run("pebble", func(t *testing.T) { + dir, err := os.MkdirTemp("", "migration-versions-*") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, os.RemoveAll(dir)) }) + st, err := NewPebbleStore(dir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, st.Close()) }) + test(t, st) + }) +} + +func TestExportVersionsPreservesRawVersionMetadata(t *testing.T) { + runMigrationStoreSuite(t, func(t *testing.T, st MVCCStore) { + ctx := context.Background() + require.NoError(t, st.PutAt(ctx, []byte("k1"), []byte("v10"), 10, 0)) + require.NoError(t, st.PutWithTTLAt(ctx, []byte("k1"), []byte("v20"), 20, 55)) + require.NoError(t, st.DeleteAt(ctx, []byte("k1"), 30)) + require.NoError(t, st.PutAt(ctx, []byte("k2"), []byte("v15"), 15, 0)) + + result, err := st.ExportVersions(ctx, ExportVersionsOptions{ + StartKey: []byte("k1"), + EndKey: []byte("k3"), + MinCommitTSExclusive: 9, + MaxCommitTSInclusive: 30, + MaxVersions: 10, + KeyFamily: 7, + }) + require.NoError(t, err) + require.True(t, result.Done) + require.Empty(t, result.NextCursor) + require.Equal(t, []MVCCVersion{ + {Key: []byte("k1"), CommitTS: 30, Tombstone: true, KeyFamily: 7}, + {Key: []byte("k1"), CommitTS: 20, Value: []byte("v20"), KeyFamily: 7, ExpireAt: 55}, + {Key: []byte("k1"), CommitTS: 10, Value: []byte("v10"), KeyFamily: 7}, + {Key: []byte("k2"), CommitTS: 15, Value: []byte("v15"), KeyFamily: 7}, + }, result.Versions) + }) +} + +func TestExportVersionsExcludesTxnLocks(t *testing.T) { + runMigrationStoreSuite(t, func(t *testing.T, st MVCCStore) { + ctx := context.Background() + lockKey := append(bytes.Clone(txnLockKeyPrefix), []byte("user")...) + require.NoError(t, st.PutAt(ctx, lockKey, []byte("lock"), 10, 0)) + require.NoError(t, st.PutAt(ctx, []byte("user"), []byte("value"), 20, 0)) + + result, err := st.ExportVersions(ctx, ExportVersionsOptions{MaxVersions: 10}) + require.NoError(t, err) + require.True(t, result.Done) + require.Equal(t, []MVCCVersion{{Key: []byte("user"), CommitTS: 20, Value: []byte("value")}}, result.Versions) + }) +} + +func TestExportVersionsAcceptVersionFiltersByValue(t *testing.T) { + runMigrationStoreSuite(t, func(t *testing.T, st MVCCStore) { + ctx := context.Background() + require.NoError(t, st.PutAt(ctx, []byte("drop"), []byte("legacy-meta"), 10, 0)) + require.NoError(t, st.PutAt(ctx, []byte("keep"), []byte("legacy-delta"), 20, 0)) + + result, err := st.ExportVersions(ctx, ExportVersionsOptions{ + MaxVersions: 10, + AcceptVersion: func(_ []byte, value []byte) bool { + return bytes.Equal(value, []byte("legacy-delta")) + }, + }) + require.NoError(t, err) + require.True(t, result.Done) + require.Equal(t, []MVCCVersion{{Key: []byte("keep"), CommitTS: 20, Value: []byte("legacy-delta")}}, result.Versions) + }) +} + +func TestExportVersionsCursorResumesWithinHotKey(t *testing.T) { + runMigrationStoreSuite(t, func(t *testing.T, st MVCCStore) { + ctx := context.Background() + require.NoError(t, st.PutAt(ctx, []byte("hot"), []byte("v10"), 10, 0)) + require.NoError(t, st.PutAt(ctx, []byte("hot"), []byte("v20"), 20, 0)) + require.NoError(t, st.PutAt(ctx, []byte("hot"), []byte("v30"), 30, 0)) + require.NoError(t, st.PutAt(ctx, []byte("tail"), []byte("v15"), 15, 0)) + + first, err := st.ExportVersions(ctx, ExportVersionsOptions{MaxVersions: 2}) + require.NoError(t, err) + require.False(t, first.Done) + require.Len(t, first.Versions, 2) + require.Equal(t, uint64(30), first.Versions[0].CommitTS) + require.Equal(t, uint64(20), first.Versions[1].CommitTS) + require.NotEmpty(t, first.NextCursor) + + second, err := st.ExportVersions(ctx, ExportVersionsOptions{ + Cursor: first.NextCursor, + MaxVersions: 10, + }) + require.NoError(t, err) + require.True(t, second.Done) + require.Equal(t, []MVCCVersion{ + {Key: []byte("hot"), CommitTS: 10, Value: []byte("v10")}, + {Key: []byte("tail"), CommitTS: 15, Value: []byte("v15")}, + }, second.Versions) + }) +} + +func TestExportVersionsSparseScanBudgetAdvancesRejectedRows(t *testing.T) { + runMigrationStoreSuite(t, func(t *testing.T, st MVCCStore) { + ctx := context.Background() + require.NoError(t, st.PutAt(ctx, []byte("drop-a"), []byte("a"), 10, 0)) + require.NoError(t, st.PutAt(ctx, []byte("keep"), []byte("b"), 20, 0)) + + first, err := st.ExportVersions(ctx, ExportVersionsOptions{ + MaxVersions: 10, + MaxScannedBytes: 1, + AcceptKey: func(key []byte) bool { + return string(key) == "keep" + }, + }) + require.NoError(t, err) + require.False(t, first.Done) + require.Empty(t, first.Versions) + require.NotEmpty(t, first.NextCursor) + + second, err := st.ExportVersions(ctx, ExportVersionsOptions{ + Cursor: first.NextCursor, + MaxVersions: 10, + AcceptKey: func(key []byte) bool { + return string(key) == "keep" + }, + }) + require.NoError(t, err) + require.True(t, second.Done) + require.Equal(t, []MVCCVersion{{Key: []byte("keep"), CommitTS: 20, Value: []byte("b")}}, second.Versions) + }) +} + +func TestExportVersionsPreservesEmptyKeyCursor(t *testing.T) { + runMigrationStoreSuite(t, func(t *testing.T, st MVCCStore) { + ctx := context.Background() + require.NoError(t, st.PutAt(ctx, nil, []byte("v10"), 10, 0)) + require.NoError(t, st.PutAt(ctx, nil, []byte("v20"), 20, 0)) + + first, err := st.ExportVersions(ctx, ExportVersionsOptions{MaxVersions: 1}) + require.NoError(t, err) + require.False(t, first.Done) + require.Len(t, first.Versions, 1) + require.Empty(t, first.Versions[0].Key) + require.Equal(t, uint64(20), first.Versions[0].CommitTS) + require.Equal(t, []byte("v20"), first.Versions[0].Value) + require.NotEmpty(t, first.NextCursor) + + second, err := st.ExportVersions(ctx, ExportVersionsOptions{ + Cursor: first.NextCursor, + MaxVersions: 10, + }) + require.NoError(t, err) + require.True(t, second.Done) + require.Len(t, second.Versions, 1) + require.Empty(t, second.Versions[0].Key) + require.Equal(t, uint64(10), second.Versions[0].CommitTS) + require.Equal(t, []byte("v10"), second.Versions[0].Value) + }) +} + +func TestExportVersionsAppliesDefaultSparseScanBudget(t *testing.T) { + runMigrationStoreSuite(t, func(t *testing.T, st MVCCStore) { + ctx := context.Background() + value := bytes.Repeat([]byte("x"), defaultSparseExportMaxScannedBytes) + require.NoError(t, st.PutAt(ctx, []byte("drop"), value, 10, 0)) + require.NoError(t, st.PutAt(ctx, []byte("keep"), []byte("v"), 20, 0)) + + first, err := st.ExportVersions(ctx, ExportVersionsOptions{ + MaxVersions: 10, + AcceptKey: func(key []byte) bool { + return bytes.Equal(key, []byte("keep")) + }, + }) + require.NoError(t, err) + require.False(t, first.Done) + require.Empty(t, first.Versions) + require.GreaterOrEqual(t, first.ScannedBytes, uint64(defaultSparseExportMaxScannedBytes)) + require.NotEmpty(t, first.NextCursor) + + second, err := st.ExportVersions(ctx, ExportVersionsOptions{ + Cursor: first.NextCursor, + MaxVersions: 10, + AcceptKey: func(key []byte) bool { + return bytes.Equal(key, []byte("keep")) + }, + }) + require.NoError(t, err) + require.True(t, second.Done) + require.Equal(t, []MVCCVersion{{Key: []byte("keep"), CommitTS: 20, Value: []byte("v")}}, second.Versions) + }) +} + +func TestExportVersionsAppliesDefaultScanBudgetForTimestampFilter(t *testing.T) { + runMigrationStoreSuite(t, func(t *testing.T, st MVCCStore) { + ctx := context.Background() + value := bytes.Repeat([]byte("x"), defaultSparseExportMaxScannedBytes) + require.NoError(t, st.PutAt(ctx, []byte("drop-new"), value, 30, 0)) + require.NoError(t, st.PutAt(ctx, []byte("keep-old"), []byte("v"), 10, 0)) + + first, err := st.ExportVersions(ctx, ExportVersionsOptions{ + MaxCommitTSInclusive: 20, + MaxVersions: 10, + }) + require.NoError(t, err) + require.False(t, first.Done) + require.Empty(t, first.Versions) + require.GreaterOrEqual(t, first.ScannedBytes, uint64(defaultSparseExportMaxScannedBytes)) + require.NotEmpty(t, first.NextCursor) + + second, err := st.ExportVersions(ctx, ExportVersionsOptions{ + Cursor: first.NextCursor, + MaxCommitTSInclusive: 20, + MaxVersions: 10, + }) + require.NoError(t, err) + require.True(t, second.Done) + require.Equal(t, []MVCCVersion{{Key: []byte("keep-old"), CommitTS: 10, Value: []byte("v")}}, second.Versions) + }) +} + +func TestExportVersionsMinTSPruneDoesNotSkipPrefixedKeys(t *testing.T) { + runMigrationStoreSuite(t, func(t *testing.T, st MVCCStore) { + ctx := context.Background() + prefixed := []byte{'a', 0xff, 0x00} + require.NoError(t, st.PutAt(ctx, []byte("a"), []byte("old"), 5, 0)) + require.NoError(t, st.PutAt(ctx, prefixed, []byte("new"), 20, 0)) + + res, err := st.ExportVersions(ctx, ExportVersionsOptions{ + StartKey: []byte("a"), + EndKey: []byte("b"), + MinCommitTSExclusive: 10, + MaxVersions: 10, + }) + require.NoError(t, err) + require.True(t, res.Done) + require.Equal(t, []MVCCVersion{{Key: prefixed, CommitTS: 20, Value: []byte("new")}}, res.Versions) + }) +} + +func TestExportVersionsMinTSSkipHonorsScanBudget(t *testing.T) { + runMigrationStoreSuite(t, func(t *testing.T, st MVCCStore) { + ctx := context.Background() + value := bytes.Repeat([]byte("x"), defaultSparseExportMaxScannedBytes) + require.NoError(t, st.PutAt(ctx, []byte("old"), value, 5, 0)) + require.NoError(t, st.PutAt(ctx, []byte("tail"), []byte("v"), 20, 0)) + + first, err := st.ExportVersions(ctx, ExportVersionsOptions{ + MinCommitTSExclusive: 10, + MaxVersions: 10, + }) + require.NoError(t, err) + require.False(t, first.Done) + require.Empty(t, first.Versions) + require.GreaterOrEqual(t, first.ScannedBytes, uint64(defaultSparseExportMaxScannedBytes)) + require.NotEmpty(t, first.NextCursor) + + second, err := st.ExportVersions(ctx, ExportVersionsOptions{ + Cursor: first.NextCursor, + MinCommitTSExclusive: 10, + MaxVersions: 10, + }) + require.NoError(t, err) + require.True(t, second.Done) + require.Equal(t, []MVCCVersion{{Key: []byte("tail"), CommitTS: 20, Value: []byte("v")}}, second.Versions) + }) +} + +func TestExportVersionsMinTSPruneCursorSkipsWholeKey(t *testing.T) { + runMigrationStoreSuite(t, func(t *testing.T, st MVCCStore) { + ctx := context.Background() + value := bytes.Repeat([]byte("x"), defaultSparseExportMaxScannedBytes) + require.NoError(t, st.PutAt(ctx, []byte("old"), value, 1, 0)) + require.NoError(t, st.PutAt(ctx, []byte("old"), value, 2, 0)) + require.NoError(t, st.PutAt(ctx, []byte("old"), value, 3, 0)) + require.NoError(t, st.PutAt(ctx, []byte("tail"), []byte("v20"), 20, 0)) + + first, err := st.ExportVersions(ctx, ExportVersionsOptions{ + MinCommitTSExclusive: 10, + MaxVersions: 10, + }) + require.NoError(t, err) + require.False(t, first.Done) + require.Empty(t, first.Versions) + require.NotEmpty(t, first.NextCursor) + + cursor := first.NextCursor + for attempts := 0; attempts < 4; attempts++ { + next, err := st.ExportVersions(ctx, ExportVersionsOptions{ + Cursor: cursor, + MinCommitTSExclusive: 10, + MaxVersions: 10, + }) + require.NoError(t, err) + if next.Done { + require.Equal(t, []MVCCVersion{{Key: []byte("tail"), CommitTS: 20, Value: []byte("v20")}}, next.Versions) + return + } + require.Empty(t, next.Versions) + require.NotEmpty(t, next.NextCursor) + cursor = next.NextCursor + } + t.Fatal("export did not finish after bounded pruned-key cursor resumes") + }) +} + +func TestExportVersionsUsesUserKeyRangeBounds(t *testing.T) { + runMigrationStoreSuite(t, func(t *testing.T, st MVCCStore) { + ctx := context.Background() + require.NoError(t, st.PutAt(ctx, []byte("a"), []byte("a-value"), 10, 0)) + require.NoError(t, st.PutAt(ctx, []byte("aa"), []byte("aa-value"), 20, 0)) + require.NoError(t, st.PutAt(ctx, []byte("b"), []byte("b-value"), 30, 0)) + + mid, err := st.ExportVersions(ctx, ExportVersionsOptions{ + StartKey: []byte("aa"), + EndKey: []byte("b"), + MaxVersions: 10, + }) + require.NoError(t, err) + require.True(t, mid.Done) + require.Equal(t, []MVCCVersion{{Key: []byte("aa"), CommitTS: 20, Value: []byte("aa-value")}}, mid.Versions) + + before, err := st.ExportVersions(ctx, ExportVersionsOptions{ + EndKey: []byte("aa"), + MaxVersions: 10, + }) + require.NoError(t, err) + require.True(t, before.Done) + require.Equal(t, []MVCCVersion{{Key: []byte("a"), CommitTS: 10, Value: []byte("a-value")}}, before.Versions) + }) +} + +func TestExportVersionsEmptyEndKeyIsUnbounded(t *testing.T) { + runMigrationStoreSuite(t, func(t *testing.T, st MVCCStore) { + ctx := context.Background() + require.NoError(t, st.PutAt(ctx, []byte("m"), []byte("m-value"), 10, 0)) + require.NoError(t, st.PutAt(ctx, []byte("z"), []byte("z-value"), 20, 0)) + + first, err := st.ExportVersions(ctx, ExportVersionsOptions{ + StartKey: []byte("m"), + EndKey: []byte{}, + MaxVersions: 1, + }) + require.NoError(t, err) + require.False(t, first.Done) + require.Equal(t, []MVCCVersion{{Key: []byte("m"), CommitTS: 10, Value: []byte("m-value")}}, first.Versions) + require.NotEmpty(t, first.NextCursor) + + second, err := st.ExportVersions(ctx, ExportVersionsOptions{ + StartKey: []byte("m"), + EndKey: []byte{}, + Cursor: first.NextCursor, + MaxVersions: 10, + }) + require.NoError(t, err) + require.True(t, second.Done) + require.Empty(t, second.NextCursor) + require.Equal(t, []MVCCVersion{{Key: []byte("z"), CommitTS: 20, Value: []byte("z-value")}}, second.Versions) + }) +} + +func TestPebbleExportDoesNotStopBeforeTrailingEmptyKey(t *testing.T) { + ctx := context.Background() + dir, err := os.MkdirTemp("", "migration-trailing-empty-key-*") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, os.RemoveAll(dir)) }) + st, err := NewPebbleStore(dir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, st.Close()) }) + + require.NoError(t, st.PutAt(ctx, []byte("a"), []byte("a-value"), 10, 0)) + require.NoError(t, st.PutAt(ctx, []byte("b"), []byte("b-value"), 20, 0)) + require.NoError(t, st.PutAt(ctx, nil, []byte("empty-value"), 30, 0)) + + res, err := st.ExportVersions(ctx, ExportVersionsOptions{ + EndKey: []byte("aa"), + MaxVersions: 10, + }) + require.NoError(t, err) + require.True(t, res.Done) + require.Len(t, res.Versions, 2) + require.Equal(t, []byte("a"), res.Versions[0].Key) + require.Equal(t, uint64(10), res.Versions[0].CommitTS) + require.Equal(t, []byte("a-value"), res.Versions[0].Value) + require.Empty(t, res.Versions[1].Key) + require.Equal(t, uint64(30), res.Versions[1].CommitTS) + require.Equal(t, []byte("empty-value"), res.Versions[1].Value) +} + +func TestExportVersionsRejectsCursorOutsideRequestedRange(t *testing.T) { + runMigrationStoreSuite(t, func(t *testing.T, st MVCCStore) { + ctx := context.Background() + require.NoError(t, st.PutAt(ctx, []byte("a"), []byte("a10"), 10, 0)) + require.NoError(t, st.PutAt(ctx, []byte("a"), []byte("a20"), 20, 0)) + require.NoError(t, st.PutAt(ctx, []byte("b"), []byte("b30"), 30, 0)) + + res, err := st.ExportVersions(ctx, ExportVersionsOptions{ + StartKey: []byte("b"), + EndKey: []byte("c"), + Cursor: encodeExportCursor([]byte("a"), 20, exportCursorTagEmitted), + MaxVersions: 10, + }) + require.ErrorIs(t, err, ErrInvalidExportCursor) + require.Empty(t, res.Versions) + }) +} + +func TestExportVersionsDoesNotTreatMigrationPrefixUserKeyAsMetadata(t *testing.T) { + runMigrationStoreSuite(t, func(t *testing.T, st MVCCStore) { + ctx := context.Background() + key := append([]byte(migrationAckPrefix), bytes.Repeat([]byte{0x7}, migrationUint64Bytes)...) + require.NoError(t, st.PutAt(ctx, key, []byte("value"), 10, 0)) + + res, err := st.ExportVersions(ctx, ExportVersionsOptions{MaxVersions: 10}) + require.NoError(t, err) + require.True(t, res.Done) + require.Equal(t, []MVCCVersion{{Key: key, CommitTS: 10, Value: []byte("value")}}, res.Versions) + }) +} + +func TestPebbleExportSkipsWriterRegistryRows(t *testing.T) { + ctx := context.Background() + dir, err := os.MkdirTemp("", "migration-writer-registry-*") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, os.RemoveAll(dir)) }) + st, err := NewPebbleStore(dir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, st.Close()) }) + + registry, err := WriterRegistryFor(st) + require.NoError(t, err) + require.NoError(t, registry.SetRegistryRow( + encryption.RegistryKey(1, 2), + encryption.EncodeRegistryValue(encryption.RegistryValue{ + FullNodeID: 2, + FirstSeenLocalEpoch: 1, + LastSeenLocalEpoch: 1, + }), + )) + require.NoError(t, st.PutAt(ctx, []byte("user"), []byte("value"), 10, 0)) + + res, err := st.ExportVersions(ctx, ExportVersionsOptions{MaxVersions: 10}) + require.NoError(t, err) + require.True(t, res.Done) + require.Equal(t, []MVCCVersion{{Key: []byte("user"), CommitTS: 10, Value: []byte("value")}}, res.Versions) +} + +func TestPebbleExportSkipsWriterRegistryRowsWithoutUserVersions(t *testing.T) { + ctx := context.Background() + dir, err := os.MkdirTemp("", "migration-writer-registry-only-*") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, os.RemoveAll(dir)) }) + st, err := NewPebbleStore(dir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, st.Close()) }) + + registry, err := WriterRegistryFor(st) + require.NoError(t, err) + require.NoError(t, registry.SetRegistryRow( + encryption.RegistryKey(1, 2), + encryption.EncodeRegistryValue(encryption.RegistryValue{ + FullNodeID: 2, + FirstSeenLocalEpoch: 1, + LastSeenLocalEpoch: 1, + }), + )) + + res, err := st.ExportVersions(ctx, ExportVersionsOptions{MaxVersions: 10}) + require.NoError(t, err) + require.True(t, res.Done) + require.Empty(t, res.NextCursor) + require.Empty(t, res.Versions) + require.Zero(t, res.ScannedBytes) + require.Zero(t, res.AcceptedRows) +} + +func TestPebbleExportDoesNotDropWriterRegistryPrefixUserKey(t *testing.T) { + ctx := context.Background() + dir, err := os.MkdirTemp("", "migration-writer-registry-user-key-*") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, os.RemoveAll(dir)) }) + st, err := NewPebbleStore(dir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, st.Close()) }) + + registryKey := encryption.RegistryKey(1, 2) + registry, err := WriterRegistryFor(st) + require.NoError(t, err) + require.NoError(t, registry.SetRegistryRow( + registryKey, + encryption.EncodeRegistryValue(encryption.RegistryValue{ + FullNodeID: 2, + FirstSeenLocalEpoch: 1, + LastSeenLocalEpoch: 1, + }), + )) + require.NoError(t, st.PutAt(ctx, registryKey, []byte("user-value"), 10, 0)) + + res, err := st.ExportVersions(ctx, ExportVersionsOptions{MaxVersions: 10}) + require.NoError(t, err) + require.True(t, res.Done) + require.Equal(t, []MVCCVersion{{Key: registryKey, CommitTS: 10, Value: []byte("user-value")}}, res.Versions) +} + +func TestPebbleRejectsMVCCKeyThatEncodesAsWriterRegistryRow(t *testing.T) { + ctx := context.Background() + dir, err := os.MkdirTemp("", "migration-writer-registry-collision-*") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, os.RemoveAll(dir)) }) + st, err := NewPebbleStore(dir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, st.Close()) }) + + rawRegistryKey := encryption.RegistryKey(1, 2) + userKey := bytes.Clone(rawRegistryKey[:len(rawRegistryKey)-timestampSize]) + commitTS := ^binary.BigEndian.Uint64(rawRegistryKey[len(rawRegistryKey)-timestampSize:]) + require.True(t, isPebbleWriterRegistryKey(encodeKey(userKey, commitTS))) + + require.ErrorIs(t, st.PutAt(ctx, userKey, []byte("value"), commitTS, 0), errMVCCMetadataKeyCollision) + _, err = st.ImportVersions(ctx, ImportVersionsOptions{ + JobID: 1, + BracketID: 1, + BatchSeq: 1, + Versions: []MVCCVersion{{ + Key: userKey, + CommitTS: commitTS, + Value: []byte("value"), + }}, + }) + require.ErrorIs(t, err, errMVCCMetadataKeyCollision) +} + +func TestPebbleWriterRegistryRowIsNotVisibleAsMVCCCollision(t *testing.T) { + ctx := context.Background() + dir, err := os.MkdirTemp("", "migration-writer-registry-read-collision-*") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, os.RemoveAll(dir)) }) + st, err := NewPebbleStore(dir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, st.Close()) }) + + registryKey := encryption.RegistryKey(1, 2) + registry, err := WriterRegistryFor(st) + require.NoError(t, err) + require.NoError(t, registry.SetRegistryRow( + registryKey, + encryption.EncodeRegistryValue(encryption.RegistryValue{ + FullNodeID: 2, + FirstSeenLocalEpoch: 1, + LastSeenLocalEpoch: 1, + }), + )) + + userKey := bytes.Clone(registryKey[:len(registryKey)-timestampSize]) + commitTS := ^binary.BigEndian.Uint64(registryKey[len(registryKey)-timestampSize:]) + _, err = st.GetAt(ctx, userKey, commitTS) + require.ErrorIs(t, err, ErrKeyNotFound) + ok, err := st.CommittedVersionAt(ctx, userKey, commitTS) + require.NoError(t, err) + require.False(t, ok) + latest, ok, err := st.LatestCommitTS(ctx, userKey) + require.NoError(t, err) + require.False(t, ok) + require.Zero(t, latest) +} + +func TestPebbleExportAuthenticatesEncryptedTombstoneHeader(t *testing.T) { + ctx := context.Background() + f := newEncryptedStoreFixture(t, 81) + require.NoError(t, f.mvcc.PutAt(ctx, []byte("tampered-export"), []byte("payload"), 100, 0)) + f.tamperPebbleValue(t, []byte("tampered-export"), 100, func(raw []byte) []byte { + raw[0] |= tombstoneMask + return raw + }) + + res, err := f.mvcc.ExportVersions(ctx, ExportVersionsOptions{MaxVersions: 10}) + require.ErrorIs(t, err, ErrEncryptedReadIntegrity) + require.Empty(t, res.Versions) +} + +func TestPebbleExportStopsAtEndKeyWhenNoLaterInRangeKeyCanTrail(t *testing.T) { + ctx := context.Background() + dir, err := os.MkdirTemp("", "migration-end-key-*") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, os.RemoveAll(dir)) }) + st, err := NewPebbleStore(dir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, st.Close()) }) + ps, ok := st.(*pebbleStore) + require.True(t, ok) + require.NoError(t, ps.PutAt(ctx, []byte("b"), []byte("later"), 10, 0)) + + iter, err := ps.db.NewIter(nil) + require.NoError(t, err) + defer func() { require.NoError(t, iter.Close()) }() + require.True(t, iter.SeekGE(encodeKey([]byte("b"), math.MaxUint64))) + + result := newExportVersionsResult(10) + advance, done, err := ps.exportPebbleIteratorPosition(ctx, iter, ExportVersionsOptions{ + StartKey: []byte("a"), + EndKey: []byte("b"), + }, exportCursorPosition{}, &result) + require.ErrorIs(t, err, errExportReachedEnd) + require.False(t, advance) + require.True(t, done) + require.Empty(t, result.Versions) + require.Zero(t, result.ScannedBytes) +} + +func TestPebbleExportDoesNotStopLeadingRangeWithExplicitEmptyStartAtEndKey(t *testing.T) { + ctx := context.Background() + dir, err := os.MkdirTemp("", "migration-leading-end-key-*") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, os.RemoveAll(dir)) }) + st, err := NewPebbleStore(dir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, st.Close()) }) + require.NoError(t, st.PutAt(ctx, []byte("b"), []byte("later"), 10, 0)) + require.NoError(t, st.PutAt(ctx, nil, []byte("empty"), 20, 0)) + + res, err := st.ExportVersions(ctx, ExportVersionsOptions{ + StartKey: []byte{}, + EndKey: []byte("b"), + MaxVersions: 10, + }) + require.NoError(t, err) + require.True(t, res.Done) + require.Equal(t, []MVCCVersion{{ + Key: []byte{}, + CommitTS: 20, + Value: []byte("empty"), + }}, res.Versions) +} + +func TestPebbleExportOutOfRangeEndSkipReturnsResumableCursor(t *testing.T) { + ctx := context.Background() + dir, err := os.MkdirTemp("", "migration-out-of-range-end-skip-*") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, os.RemoveAll(dir)) }) + st, err := NewPebbleStore(dir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, st.Close()) }) + + require.NoError(t, st.PutAt(ctx, []byte("b"), []byte("later"), 10, 0)) + require.NoError(t, st.PutAt(ctx, nil, []byte("empty"), 20, 0)) + + first, err := st.ExportVersions(ctx, ExportVersionsOptions{ + EndKey: []byte("a"), + MaxVersions: 10, + MaxScannedBytes: 1, + }) + require.NoError(t, err) + require.False(t, first.Done) + require.Empty(t, first.Versions) + require.NotEmpty(t, first.NextCursor) + require.Greater(t, first.ScannedBytes, uint64(0)) + + second, err := st.ExportVersions(ctx, ExportVersionsOptions{ + EndKey: []byte("a"), + Cursor: first.NextCursor, + MaxVersions: 10, + }) + require.NoError(t, err) + require.True(t, second.Done) + require.Equal(t, []MVCCVersion{{Key: []byte{}, CommitTS: 20, Value: []byte("empty")}}, second.Versions) +} + +func TestPebbleExportOutOfRangeEndSkipChargesAllVersions(t *testing.T) { + ctx := context.Background() + dir, err := os.MkdirTemp("", "migration-out-of-range-end-skip-versions-*") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, os.RemoveAll(dir)) }) + st, err := NewPebbleStore(dir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, st.Close()) }) + + require.NoError(t, st.PutAt(ctx, []byte("b"), []byte("v10"), 10, 0)) + require.NoError(t, st.PutAt(ctx, []byte("b"), []byte("v20"), 20, 0)) + require.NoError(t, st.PutAt(ctx, []byte("b"), []byte("v30"), 30, 0)) + require.NoError(t, st.PutAt(ctx, nil, []byte("empty"), 40, 0)) + + first, err := st.ExportVersions(ctx, ExportVersionsOptions{ + EndKey: []byte("a"), + MaxVersions: 10, + MaxScannedBytes: versionExportSize([]byte("b"), len("v30")) + 1, + }) + require.NoError(t, err) + require.False(t, first.Done) + require.Empty(t, first.Versions) + require.NotEmpty(t, first.NextCursor) + require.GreaterOrEqual(t, first.ScannedBytes, versionExportSize([]byte("b"), len("v30"))+1) + + second, err := st.ExportVersions(ctx, ExportVersionsOptions{ + EndKey: []byte("a"), + Cursor: first.NextCursor, + MaxVersions: 10, + }) + require.NoError(t, err) + require.True(t, second.Done) + require.Equal(t, []MVCCVersion{{Key: []byte{}, CommitTS: 40, Value: []byte("empty")}}, second.Versions) +} + +func TestPebbleExportAppliesDefaultScanBudgetForRangeBoundSkip(t *testing.T) { + ctx := context.Background() + dir, err := os.MkdirTemp("", "migration-range-bound-scan-budget-*") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, os.RemoveAll(dir)) }) + st, err := NewPebbleStore(dir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, st.Close()) }) + + value := bytes.Repeat([]byte("x"), defaultSparseExportMaxScannedBytes) + require.NoError(t, st.PutAt(ctx, []byte("b"), value, 10, 0)) + require.NoError(t, st.PutAt(ctx, nil, []byte("empty"), 20, 0)) + + first, err := st.ExportVersions(ctx, ExportVersionsOptions{ + EndKey: []byte("a"), + MaxVersions: 10, + }) + require.NoError(t, err) + require.False(t, first.Done) + require.Empty(t, first.Versions) + require.NotEmpty(t, first.NextCursor) + require.GreaterOrEqual(t, first.ScannedBytes, uint64(defaultSparseExportMaxScannedBytes)) + + second, err := st.ExportVersions(ctx, ExportVersionsOptions{ + EndKey: []byte("a"), + Cursor: first.NextCursor, + MaxVersions: 10, + }) + require.NoError(t, err) + require.True(t, second.Done) + require.Equal(t, []MVCCVersion{{Key: []byte{}, CommitTS: 20, Value: []byte("empty")}}, second.Versions) +} + +func TestPebbleWriterRegistryKeyShapeRejectsOrdinaryRows(t *testing.T) { + registryKey := encryption.RegistryKey(1, 2) + require.True(t, isPebbleWriterRegistryKey(registryKey)) + + ordinary := encodeKey([]byte("ordinary"), 10) + require.False(t, isPebbleWriterRegistryKey(ordinary)) + + malformed := bytes.Clone(registryKey) + malformed[len(encryption.WriterRegistryPrefix)+4] = '#' + require.False(t, isPebbleWriterRegistryKey(malformed)) +} + +func TestImportVersionsIdempotencyAndMetadata(t *testing.T) { + runMigrationStoreSuite(t, func(t *testing.T, st MVCCStore) { + ctx := context.Background() + first := ImportVersionsOptions{ + JobID: 1, + BracketID: 2, + BatchSeq: 1, + Cursor: []byte("c1"), + Versions: []MVCCVersion{ + {Key: []byte("ttl"), CommitTS: 20, Value: []byte("v20"), ExpireAt: 50}, + {Key: []byte("gone"), CommitTS: 30, Tombstone: true}, + }, + } + res, err := st.ImportVersions(ctx, first) + require.NoError(t, err) + require.Equal(t, []byte("c1"), res.AckedCursor) + require.Equal(t, uint64(30), res.MaxImportedTS) + require.False(t, res.Duplicate) + require.Equal(t, uint64(30), st.LastCommitTS()) + floor, err := st.MigrationHLCFloor(ctx, 1) + require.NoError(t, err) + require.Equal(t, uint64(30), floor) + + val, err := st.GetAt(ctx, []byte("ttl"), 25) + require.NoError(t, err) + require.Equal(t, []byte("v20"), val) + _, err = st.GetAt(ctx, []byte("ttl"), 55) + require.ErrorIs(t, err, ErrKeyNotFound) + _, err = st.GetAt(ctx, []byte("gone"), 35) + require.ErrorIs(t, err, ErrKeyNotFound) + + dup := first + dup.Cursor = []byte("changed") + dup.Versions = []MVCCVersion{{Key: []byte("ttl"), CommitTS: 40, Value: []byte("bad")}} + res, err = st.ImportVersions(ctx, dup) + require.NoError(t, err) + require.True(t, res.Duplicate) + require.Equal(t, []byte("c1"), res.AckedCursor) + val, err = st.GetAt(ctx, []byte("ttl"), 45) + require.NoError(t, err) + require.Equal(t, []byte("v20"), val) + + _, err = st.ImportVersions(ctx, ImportVersionsOptions{ + JobID: 1, + BracketID: 2, + BatchSeq: 3, + Cursor: []byte("gap"), + Versions: []MVCCVersion{{Key: []byte("gap"), CommitTS: 60, Value: []byte("bad")}}, + }) + require.ErrorIs(t, err, ErrImportBatchGap) + _, err = st.GetAt(ctx, []byte("gap"), 60) + require.ErrorIs(t, err, ErrKeyNotFound) + + res, err = st.ImportVersions(ctx, ImportVersionsOptions{ + JobID: 1, + BracketID: 2, + BatchSeq: 2, + Cursor: []byte("c2"), + }) + require.NoError(t, err) + require.Equal(t, []byte("c2"), res.AckedCursor) + require.Zero(t, res.MaxImportedTS) + require.Equal(t, uint64(30), st.LastCommitTS()) + floor, err = st.MigrationHLCFloor(ctx, 1) + require.NoError(t, err) + require.Equal(t, uint64(30), floor) + }) +} + +func TestPebbleImportMetadataPersistsAcrossReopen(t *testing.T) { + ctx := context.Background() + dir, err := os.MkdirTemp("", "migration-import-persist-*") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, os.RemoveAll(dir)) }) + + st, err := NewPebbleStore(dir) + require.NoError(t, err) + _, err = st.ImportVersions(ctx, ImportVersionsOptions{ + JobID: 9, + BracketID: 4, + BatchSeq: 1, + Cursor: []byte("persisted"), + Versions: []MVCCVersion{{Key: []byte("k"), CommitTS: 99, Value: []byte("v")}}, + }) + require.NoError(t, err) + require.NoError(t, st.Close()) + + reopened, err := NewPebbleStore(dir) + require.NoError(t, err) + defer func() { require.NoError(t, reopened.Close()) }() + floor, err := reopened.MigrationHLCFloor(ctx, 9) + require.NoError(t, err) + require.Equal(t, uint64(99), floor) + res, err := reopened.ImportVersions(ctx, ImportVersionsOptions{ + JobID: 9, + BracketID: 4, + BatchSeq: 1, + Cursor: []byte("different"), + }) + require.NoError(t, err) + require.True(t, res.Duplicate) + require.Equal(t, []byte("persisted"), res.AckedCursor) +} + +func TestPebbleSnapshotPreservesMigrationMetadata(t *testing.T) { + ctx := context.Background() + srcDir, err := os.MkdirTemp("", "migration-snapshot-src-*") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, os.RemoveAll(srcDir)) }) + src, err := NewPebbleStore(srcDir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, src.Close()) }) + + _, err = src.ImportVersions(ctx, ImportVersionsOptions{ + JobID: 7, + BracketID: 3, + BatchSeq: 1, + Cursor: []byte("stale"), + Versions: []MVCCVersion{{Key: []byte("snapshotted"), CommitTS: 50, Value: []byte("v50")}}, + }) + require.NoError(t, err) + snap, err := src.Snapshot() + require.NoError(t, err) + raw := snapshotBytes(t, snap) + require.NoError(t, snap.Close()) + + dstDir, err := os.MkdirTemp("", "migration-snapshot-dst-*") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, os.RemoveAll(dstDir)) }) + dst, err := NewPebbleStore(dstDir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, dst.Close()) }) + require.NoError(t, dst.Restore(bytes.NewReader(raw))) + + val, err := dst.GetAt(ctx, []byte("snapshotted"), 50) + require.NoError(t, err) + require.Equal(t, []byte("v50"), val) + floor, err := dst.MigrationHLCFloor(ctx, 7) + require.NoError(t, err) + require.Equal(t, uint64(50), floor) + + res, err := dst.ImportVersions(ctx, ImportVersionsOptions{ + JobID: 7, + BracketID: 3, + BatchSeq: 1, + Cursor: []byte("fresh"), + Versions: []MVCCVersion{{Key: []byte("fresh"), CommitTS: 60, Value: []byte("v60")}}, + }) + require.NoError(t, err) + require.True(t, res.Duplicate) + require.Equal(t, []byte("stale"), res.AckedCursor) + _, err = dst.GetAt(ctx, []byte("fresh"), 60) + require.ErrorIs(t, err, ErrKeyNotFound) +} diff --git a/store/mvcc_store.go b/store/mvcc_store.go index 9776e2389..cc173c9c0 100644 --- a/store/mvcc_store.go +++ b/store/mvcc_store.go @@ -25,7 +25,8 @@ type VersionedValue struct { } const ( - mvccSnapshotVersion = uint32(1) + mvccSnapshotVersionV1 = uint32(1) + mvccSnapshotVersion = uint32(2) maxSnapshotKeySize = 1 << 20 // 1 MiB per key maxSnapshotVersionCount = 1 << 20 // 1M versions per key ) @@ -60,11 +61,13 @@ func byteSliceComparator(a, b any) int { // mvccStore is an in-memory MVCC implementation backed by a treemap for // deterministic iteration order and range scans. type mvccStore struct { - tree *treemap.Map // key []byte -> []VersionedValue - mtx sync.RWMutex - log *slog.Logger - lastCommitTS uint64 - minRetainedTS uint64 + tree *treemap.Map // key []byte -> []VersionedValue + mtx sync.RWMutex + log *slog.Logger + lastCommitTS uint64 + minRetainedTS uint64 + migrationAcks map[migrationAckID]migrationImportAck + migrationHLCFloors map[uint64]uint64 // writeConflicts mirrors the per-(kind, key_prefix) counter from // the pebble-backed store so the in-memory implementation shows up // in the same Prometheus series (even if the counts are usually @@ -111,7 +114,9 @@ func NewMVCCStore(opts ...MVCCStoreOption) MVCCStore { log: slog.New(slog.NewJSONHandler(os.Stdout, &slog.HandlerOptions{ Level: slog.LevelWarn, })), - writeConflicts: newWriteConflictCounter(), + migrationAcks: make(map[migrationAckID]migrationImportAck), + migrationHLCFloors: make(map[uint64]uint64), + writeConflicts: newWriteConflictCounter(), } for _, opt := range opts { opt(s) @@ -942,6 +947,12 @@ func (s *mvccStore) writeSnapshotBody(f *os.File) (uint32, error) { if err := binary.Write(w, binary.LittleEndian, s.minRetainedTS); err != nil { return 0, errors.WithStack(err) } + if err := writeMVCCSnapshotBytes(w, encodeMigrationImportAcks(s.migrationAcks)); err != nil { + return 0, err + } + if err := writeMVCCSnapshotBytes(w, encodeMigrationHLCFloors(s.migrationHLCFloors)); err != nil { + return 0, err + } iter := s.tree.Iterator() for iter.Next() { key, ok := iter.Key().([]byte) @@ -975,6 +986,16 @@ func finalizeMVCCSnapshotFile(f *os.File, checksumOffset int64, sum uint32) erro return nil } +func writeMVCCSnapshotBytes(w io.Writer, data []byte) error { + if err := binary.Write(w, binary.LittleEndian, uint64(len(data))); err != nil { + return errors.WithStack(err) + } + if _, err := w.Write(data); err != nil { + return errors.WithStack(err) + } + return nil +} + func writeMVCCSnapshotEntry(w io.Writer, key []byte, versions []VersionedValue) error { if err := binary.Write(w, binary.LittleEndian, uint64(len(key))); err != nil { return errors.WithStack(err) @@ -1020,12 +1041,12 @@ func mvccSnapshotTombstoneByte(tombstone bool) byte { } func (s *mvccStore) restoreStreamingSnapshot(r io.Reader) error { - expected, err := readMVCCSnapshotHeader(r) + version, expected, err := readMVCCSnapshotHeader(r) if err != nil { return err } - tree, lastCommitTS, minRetainedTS, actual, err := restoreStreamingMVCCSnapshotBody(r) + tree, lastCommitTS, minRetainedTS, migrationAcks, migrationHLCFloors, actual, err := restoreStreamingMVCCSnapshotBody(r, version) if err != nil { return err } @@ -1038,60 +1059,100 @@ func (s *mvccStore) restoreStreamingSnapshot(r io.Reader) error { s.tree = tree s.lastCommitTS = lastCommitTS s.minRetainedTS = minRetainedTS + s.migrationAcks = migrationAcks + s.migrationHLCFloors = migrationHLCFloors return nil } -func readMVCCSnapshotHeader(r io.Reader) (uint32, error) { +func readMVCCSnapshotHeader(r io.Reader) (uint32, uint32, error) { var magic [8]byte if _, err := io.ReadFull(r, magic[:]); err != nil { - return 0, errors.WithStack(err) + return 0, 0, errors.WithStack(err) } if magic != mvccSnapshotMagic { - return 0, errors.WithStack(ErrInvalidChecksum) + return 0, 0, errors.WithStack(ErrInvalidChecksum) } var version uint32 if err := binary.Read(r, binary.LittleEndian, &version); err != nil { - return 0, errors.WithStack(err) + return 0, 0, errors.WithStack(err) } - if version != mvccSnapshotVersion { - return 0, errors.WithStack(errors.Newf("unsupported mvcc snapshot version %d", version)) + if version != mvccSnapshotVersionV1 && version != mvccSnapshotVersion { + return 0, 0, errors.WithStack(errors.Newf("unsupported mvcc snapshot version %d", version)) } var expected uint32 if err := binary.Read(r, binary.LittleEndian, &expected); err != nil { - return 0, errors.WithStack(err) + return 0, 0, errors.WithStack(err) } - return expected, nil + return version, expected, nil } -func restoreStreamingMVCCSnapshotBody(r io.Reader) (*treemap.Map, uint64, uint64, uint32, error) { +func restoreStreamingMVCCSnapshotBody(r io.Reader, version uint32) (*treemap.Map, uint64, uint64, map[migrationAckID]migrationImportAck, map[uint64]uint64, uint32, error) { hash := crc32.NewIEEE() body := io.TeeReader(r, hash) - lastCommitTS, minRetainedTS, err := readMVCCSnapshotMetadata(body) + lastCommitTS, minRetainedTS, migrationAcks, migrationHLCFloors, err := readMVCCSnapshotMetadata(body, version) if err != nil { - return nil, 0, 0, 0, err + return nil, 0, 0, nil, nil, 0, err } tree, err := readMVCCSnapshotTree(body) if err != nil { - return nil, 0, 0, 0, err + return nil, 0, 0, nil, nil, 0, err } - return tree, lastCommitTS, minRetainedTS, hash.Sum32(), nil + return tree, lastCommitTS, minRetainedTS, migrationAcks, migrationHLCFloors, hash.Sum32(), nil } -func readMVCCSnapshotMetadata(r io.Reader) (uint64, uint64, error) { +func readMVCCSnapshotMetadata(r io.Reader, version uint32) (uint64, uint64, map[migrationAckID]migrationImportAck, map[uint64]uint64, error) { var lastCommitTS uint64 if err := binary.Read(r, binary.LittleEndian, &lastCommitTS); err != nil { - return 0, 0, errors.WithStack(err) + return 0, 0, nil, nil, errors.WithStack(err) } var minRetainedTS uint64 if err := binary.Read(r, binary.LittleEndian, &minRetainedTS); err != nil { - return 0, 0, errors.WithStack(err) + return 0, 0, nil, nil, errors.WithStack(err) + } + if version == mvccSnapshotVersionV1 { + return lastCommitTS, minRetainedTS, make(map[migrationAckID]migrationImportAck), make(map[uint64]uint64), nil + } + + ackData, err := readMVCCSnapshotBytes(r, "snapshot migration acks") + if err != nil { + return 0, 0, nil, nil, err + } + migrationAcks, ok := decodeMigrationImportAcks(ackData) + if !ok { + return 0, 0, nil, nil, errors.New("invalid snapshot migration acks") + } + + floorData, err := readMVCCSnapshotBytes(r, "snapshot migration hlc floors") + if err != nil { + return 0, 0, nil, nil, err + } + migrationHLCFloors, ok := decodeMigrationHLCFloors(floorData) + if !ok { + return 0, 0, nil, nil, errors.New("invalid snapshot migration hlc floors") + } + + return lastCommitTS, minRetainedTS, migrationAcks, migrationHLCFloors, nil +} + +func readMVCCSnapshotBytes(r io.Reader, field string) ([]byte, error) { + var dataLen uint64 + if err := binary.Read(r, binary.LittleEndian, &dataLen); err != nil { + return nil, errors.WithStack(err) + } + decodedLen, err := restoreFieldLenInt(dataLen, field, maxSnapshotValueSize) + if err != nil { + return nil, err + } + data := make([]byte, decodedLen) + if _, err := io.ReadFull(r, data); err != nil { + return nil, errors.WithStack(err) } - return lastCommitTS, minRetainedTS, nil + return data, nil } func readMVCCSnapshotTree(r io.Reader) (*treemap.Map, error) { diff --git a/store/mvcc_store_snapshot_test.go b/store/mvcc_store_snapshot_test.go index 3ebdbcac8..7b0bd39a3 100644 --- a/store/mvcc_store_snapshot_test.go +++ b/store/mvcc_store_snapshot_test.go @@ -56,6 +56,101 @@ func TestMVCCStore_RestoreRejectsInvalidChecksum(t *testing.T) { require.ErrorIs(t, st.Restore(bytes.NewReader(raw)), ErrInvalidChecksum) } +func TestMVCCStore_RestoreClearsMigrationMetadata(t *testing.T) { + t.Parallel() + + ctx := context.Background() + st := newTestMVCCStore(t) + require.NoError(t, st.PutAt(ctx, []byte("base"), []byte("v1"), 10, 0)) + + snap, err := st.Snapshot() + require.NoError(t, err) + defer snap.Close() + raw := snapshotBytes(t, snap) + + _, err = st.ImportVersions(ctx, ImportVersionsOptions{ + JobID: 7, + BracketID: 3, + BatchSeq: 1, + Cursor: []byte("stale"), + Versions: []MVCCVersion{{Key: []byte("imported"), CommitTS: 50, Value: []byte("v50")}}, + }) + require.NoError(t, err) + floor, err := st.MigrationHLCFloor(ctx, 7) + require.NoError(t, err) + require.Equal(t, uint64(50), floor) + + require.NoError(t, st.Restore(bytes.NewReader(raw))) + floor, err = st.MigrationHLCFloor(ctx, 7) + require.NoError(t, err) + require.Zero(t, floor) + _, err = st.GetAt(ctx, []byte("imported"), 50) + require.ErrorIs(t, err, ErrKeyNotFound) + + res, err := st.ImportVersions(ctx, ImportVersionsOptions{ + JobID: 7, + BracketID: 3, + BatchSeq: 1, + Cursor: []byte("fresh"), + Versions: []MVCCVersion{{Key: []byte("imported"), CommitTS: 60, Value: []byte("v60")}}, + }) + require.NoError(t, err) + require.False(t, res.Duplicate) + require.Equal(t, []byte("fresh"), res.AckedCursor) +} + +func TestMVCCStore_SnapshotRestorePreservesMigrationMetadata(t *testing.T) { + t.Parallel() + + ctx := context.Background() + st := newTestMVCCStore(t) + res, err := st.ImportVersions(ctx, ImportVersionsOptions{ + JobID: 7, + BracketID: 3, + BatchSeq: 1, + Cursor: []byte("snap-cursor"), + Versions: []MVCCVersion{{Key: []byte("imported"), CommitTS: 50, Value: []byte("v50")}}, + }) + require.NoError(t, err) + require.False(t, res.Duplicate) + require.Equal(t, uint64(50), res.MaxImportedTS) + + snap, err := st.Snapshot() + require.NoError(t, err) + defer snap.Close() + raw := snapshotBytes(t, snap) + + _, err = st.ImportVersions(ctx, ImportVersionsOptions{ + JobID: 7, + BracketID: 3, + BatchSeq: 2, + Cursor: []byte("post-snapshot-cursor"), + Versions: []MVCCVersion{{Key: []byte("post-snapshot"), CommitTS: 60, Value: []byte("v60")}}, + }) + require.NoError(t, err) + + require.NoError(t, st.Restore(bytes.NewReader(raw))) + floor, err := st.MigrationHLCFloor(ctx, 7) + require.NoError(t, err) + require.Equal(t, uint64(50), floor) + + _, err = st.GetAt(ctx, []byte("post-snapshot"), 60) + require.ErrorIs(t, err, ErrKeyNotFound) + + res, err = st.ImportVersions(ctx, ImportVersionsOptions{ + JobID: 7, + BracketID: 3, + BatchSeq: 1, + Cursor: []byte("ignored"), + Versions: []MVCCVersion{{Key: []byte("duplicate"), CommitTS: 70, Value: []byte("v70")}}, + }) + require.NoError(t, err) + require.True(t, res.Duplicate) + require.Equal(t, []byte("snap-cursor"), res.AckedCursor) + _, err = st.GetAt(ctx, []byte("duplicate"), 70) + require.ErrorIs(t, err, ErrKeyNotFound) +} + func TestMVCCStore_ApplyMutations_WriteConflict(t *testing.T) { t.Parallel() diff --git a/store/set_helpers.go b/store/set_helpers.go index 318581ab7..2e0d106d6 100644 --- a/store/set_helpers.go +++ b/store/set_helpers.go @@ -160,40 +160,21 @@ func IsSetMetaDeltaKey(key []byte) bool { // ExtractSetUserKeyFromMeta extracts the logical user key from a set meta key. func ExtractSetUserKeyFromMeta(key []byte) []byte { - trimmed := bytes.TrimPrefix(key, []byte(SetMetaPrefix)) - if len(trimmed) < wideColKeyLenSize { - return nil - } - ukLen := binary.BigEndian.Uint32(trimmed[:wideColKeyLenSize]) - if uint32(len(trimmed)) < uint32(wideColKeyLenSize)+ukLen { //nolint:gosec // wideColKeyLenSize fits in uint32 - return nil - } - return trimmed[wideColKeyLenSize : wideColKeyLenSize+ukLen] + return extractWideColumnUserKey(key, []byte(SetMetaPrefix), 0, true) } // ExtractSetUserKeyFromMember extracts the logical user key from a set member key. func ExtractSetUserKeyFromMember(key []byte) []byte { - trimmed := bytes.TrimPrefix(key, []byte(SetMemberPrefix)) - if len(trimmed) < wideColKeyLenSize { - return nil - } - ukLen := binary.BigEndian.Uint32(trimmed[:wideColKeyLenSize]) - if uint32(len(trimmed)) < uint32(wideColKeyLenSize)+ukLen { //nolint:gosec // wideColKeyLenSize fits in uint32 - return nil - } - return trimmed[wideColKeyLenSize : wideColKeyLenSize+ukLen] + return extractWideColumnUserKey(key, []byte(SetMemberPrefix), 0, false) } // ExtractSetUserKeyFromDelta extracts the logical user key from a set delta key. func ExtractSetUserKeyFromDelta(key []byte) []byte { - trimmed := bytes.TrimPrefix(key, []byte(SetMetaDeltaPrefix)) - minLen := wideColKeyLenSize + deltaKeyTSSize + deltaKeySeqSize - if len(trimmed) < minLen { - return nil - } - ukLen := binary.BigEndian.Uint32(trimmed[:wideColKeyLenSize]) - if uint32(len(trimmed)) < uint32(wideColKeyLenSize)+ukLen+uint32(deltaKeyTSSize+deltaKeySeqSize) { //nolint:gosec // constants fit in uint32 - return nil - } - return trimmed[wideColKeyLenSize : wideColKeyLenSize+ukLen] + return extractWideColumnUserKey(key, []byte(SetMetaDeltaPrefix), deltaKeyTSSize+deltaKeySeqSize, true) +} + +// ExtractSetUserKeyFromDeltaScanPrefix extracts the user key from a set +// metadata delta scan start/prefix. +func ExtractSetUserKeyFromDeltaScanPrefix(key []byte) []byte { + return extractWideColumnUserKey(key, []byte(SetMetaDeltaPrefix), 0, false) } diff --git a/store/store.go b/store/store.go index 362c3d505..defc5241d 100644 --- a/store/store.go +++ b/store/store.go @@ -15,6 +15,10 @@ import ( // single definition due to the store→kv import cycle. var txnInternalKeyPrefix = []byte("!txn|") +// txnLockKeyPrefix must match kv's txn lock namespace. Migration exports skip +// these in-flight lock records so a target range never imports a stale lock. +var txnLockKeyPrefix = []byte("!txn|lock|") + var ErrKeyNotFound = errors.New("not found") var ErrUnknownOp = errors.New("unknown op") var ErrNotSupported = errors.New("not supported") @@ -25,6 +29,8 @@ var ErrReadTSCompacted = errors.New("read timestamp has been compacted") var ErrSnapshotKeyTooLarge = errors.New("mvcc snapshot key too large") var ErrSnapshotVersionCountTooLarge = errors.New("mvcc snapshot version count too large") var ErrValueTooLarge = errors.New("value too large") +var ErrInvalidExportCursor = errors.New("invalid export cursor") +var ErrImportBatchGap = errors.New("migration import batch gap") // validateValueSize returns ErrValueTooLarge when the value exceeds maxSnapshotValueSize. func validateValueSize(value []byte) error { @@ -59,8 +65,61 @@ func (e *WriteConflictError) Unwrap() error { } type KVPair struct { - Key []byte - Value []byte + Key []byte + Value []byte + RouteGroupID uint64 +} + +// MVCCVersion is a raw committed MVCC version for range migration. +// Unlike scan results, it preserves tombstones and TTL expiry metadata. +type MVCCVersion struct { + Key []byte + CommitTS uint64 + Tombstone bool + Value []byte + KeyFamily uint32 + ExpireAt uint64 +} + +// ExportVersionsOptions selects a raw MVCC-version export window. +type ExportVersionsOptions struct { + StartKey []byte + EndKey []byte + MinCommitTSExclusive uint64 + MaxCommitTSInclusive uint64 + Cursor []byte + MaxVersions int + MaxBytes uint64 + MaxScannedBytes uint64 + KeyFamily uint32 + AcceptKey func([]byte) bool + AcceptVersion func(key []byte, value []byte) bool +} + +// ExportVersionsResult is one resumable chunk of raw MVCC versions. +type ExportVersionsResult struct { + Versions []MVCCVersion + NextCursor []byte + Done bool + ScannedBytes uint64 + ExportedBytes uint64 + AcceptedRows uint64 +} + +// ImportVersionsOptions applies one idempotent migration-import batch. +type ImportVersionsOptions struct { + JobID uint64 + BracketID uint64 + BatchSeq uint64 + Versions []MVCCVersion + Cursor []byte +} + +// ImportVersionsResult reports the cursor durably acknowledged by the target. +type ImportVersionsResult struct { + AckedCursor []byte + MaxImportedTS uint64 + Duplicate bool } // OpType describes a mutation kind. @@ -215,6 +274,14 @@ type MVCCStore interface { WriteConflictCountsByPrefix() map[string]uint64 // Compact removes versions older than minTS that are no longer needed. Compact(ctx context.Context, minTS uint64) error + // ExportVersions exports raw committed MVCC versions for range migration. + ExportVersions(ctx context.Context, opts ExportVersionsOptions) (ExportVersionsResult, error) + // ImportVersions applies a migration import batch idempotently by + // (jobID, bracketID, batchSeq), preserving tombstones and expireAt. + ImportVersions(ctx context.Context, opts ImportVersionsOptions) (ImportVersionsResult, error) + // MigrationHLCFloor returns the full-HLC target-local migration floor + // persisted by ImportVersions for jobID. + MigrationHLCFloor(ctx context.Context, jobID uint64) (uint64, error) Snapshot() (Snapshot, error) Restore(buf io.Reader) error Close() error diff --git a/store/stream_helpers.go b/store/stream_helpers.go index 35ca5fbb9..88d303d43 100644 --- a/store/stream_helpers.go +++ b/store/stream_helpers.go @@ -183,6 +183,9 @@ func IsStreamEntryKey(key []byte) bool { // math.MaxUint32 cannot wrap (uint32(wideColKeyLenSize)+ukLen) and pass a // false negative, which would then panic on the trimmed[lo:hi] slice below. func ExtractStreamUserKeyFromMeta(key []byte) []byte { + if !bytes.HasPrefix(key, streamMetaPrefixBytes) { + return nil + } trimmed := bytes.TrimPrefix(key, streamMetaPrefixBytes) if len(trimmed) < wideColKeyLenSize { return nil @@ -194,12 +197,21 @@ func ExtractStreamUserKeyFromMeta(key []byte) []byte { return trimmed[wideColKeyLenSize : wideColKeyLenSize+ukLen] } +// ExtractStreamUserKeyFromEntryScanPrefix extracts the logical user key from +// the prefix used to scan all entries for a stream. +func ExtractStreamUserKeyFromEntryScanPrefix(key []byte) []byte { + return extractWideColumnUserKey(key, streamEntryPrefixBytes, 0, false) +} + // ExtractStreamUserKeyFromEntry extracts the logical user key from a stream entry key. // // See ExtractStreamUserKeyFromMeta for the rationale of the uint64 bounds // check; the entry variant additionally has to account for the trailing // StreamIDBytes (16 bytes) suffix. func ExtractStreamUserKeyFromEntry(key []byte) []byte { + if !bytes.HasPrefix(key, streamEntryPrefixBytes) { + return nil + } trimmed := bytes.TrimPrefix(key, streamEntryPrefixBytes) if len(trimmed) < wideColKeyLenSize+StreamIDBytes { return nil diff --git a/store/stream_helpers_test.go b/store/stream_helpers_test.go index 78c66eb78..fd2285b7d 100644 --- a/store/stream_helpers_test.go +++ b/store/stream_helpers_test.go @@ -76,71 +76,25 @@ func TestExtractStreamUserKeyFromEntry_RoundTrip(t *testing.T) { } } -func TestStreamMetaMarshalRoundTripTrimCursor(t *testing.T) { +func TestExtractStreamUserKeyFromEntryScanPrefix_RoundTrip(t *testing.T) { t.Parallel() - - want := StreamMeta{ - Length: 42, - LastMs: 1000, - LastSeq: 7, - ExpireAt: 123456, - TrimmedMs: 999, - TrimmedSeq: 6, - } - raw, err := MarshalStreamMeta(want) - if err != nil { - t.Fatalf("MarshalStreamMeta: %v", err) - } - if len(raw) != streamMetaTrimBinarySize { - t.Fatalf("encoded size: want %d, got %d", streamMetaTrimBinarySize, len(raw)) - } - got, err := UnmarshalStreamMeta(raw) - if err != nil { - t.Fatalf("UnmarshalStreamMeta: %v", err) - } - if got != want { - t.Fatalf("round trip: want %+v, got %+v", want, got) - } -} - -func TestStreamMetaUnmarshalLegacySizes(t *testing.T) { - t.Parallel() - - legacy := make([]byte, streamMetaLegacyBinarySize) - binary.BigEndian.PutUint64(legacy[0:8], 3) - binary.BigEndian.PutUint64(legacy[8:16], 10) - binary.BigEndian.PutUint64(legacy[16:24], 2) - got, err := UnmarshalStreamMeta(legacy) - if err != nil { - t.Fatalf("legacy UnmarshalStreamMeta: %v", err) - } - if got != (StreamMeta{Length: 3, LastMs: 10, LastSeq: 2}) { - t.Fatalf("legacy meta: got %+v", got) - } - - current := make([]byte, streamMetaBinarySize) - copy(current, legacy) - binary.BigEndian.PutUint64(current[24:32], 99) - got, err = UnmarshalStreamMeta(current) - if err != nil { - t.Fatalf("current UnmarshalStreamMeta: %v", err) - } - if got != (StreamMeta{Length: 3, LastMs: 10, LastSeq: 2, ExpireAt: 99}) { - t.Fatalf("current meta: got %+v", got) + want := []byte("entry-scan-user-key") + if got := ExtractStreamUserKeyFromEntryScanPrefix(StreamEntryScanPrefix(want)); !bytes.Equal(got, want) { + t.Fatalf("scan prefix round trip: want %q, got %q", want, got) } } -func TestStreamEntryScanStartUsesTrimCursor(t *testing.T) { +func TestExtractStreamUserKeyRejectsForeignPrefixes(t *testing.T) { t.Parallel() - key := []byte("trim-cursor") - if got := StreamEntryScanStart(key, StreamMeta{}); !bytes.Equal(got, StreamEntryScanPrefix(key)) { - t.Fatalf("no trim cursor: got %q", got) + key := append([]byte{0, 0, 0, 1}, 'x') + if got := ExtractStreamUserKeyFromMeta(key); got != nil { + t.Fatalf("meta extractor accepted non-stream key: %q", got) } - if got := StreamEntryScanStart(key, StreamMeta{TrimmedMs: 9, TrimmedSeq: 4}); !bytes.Equal(got, StreamEntryKey(key, 9, 5)) { - t.Fatalf("trim cursor: want start after 9-4, got %q", got) + if got := ExtractStreamUserKeyFromEntry(key); got != nil { + t.Fatalf("entry extractor accepted non-stream key: %q", got) } - if got := StreamEntryScanStart(key, StreamMeta{TrimmedMs: 9, TrimmedSeq: ^uint64(0)}); !bytes.Equal(got, StreamEntryKey(key, 10, 0)) { - t.Fatalf("seq overflow cursor: want start at 10-0, got %q", got) + if got := ExtractStreamUserKeyFromEntryScanPrefix(key); got != nil { + t.Fatalf("entry scan extractor accepted non-stream key: %q", got) } } diff --git a/store/wide_column_helpers_test.go b/store/wide_column_helpers_test.go new file mode 100644 index 000000000..243cc8a5f --- /dev/null +++ b/store/wide_column_helpers_test.go @@ -0,0 +1,79 @@ +package store + +import ( + "encoding/binary" + "testing" +) + +func TestWideColumnExtractorsRejectOverflowingUserKeyLength(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name string + prefix string + suffixLen int + extract func([]byte) []byte + }{ + {name: "list claim", prefix: ListClaimPrefix, suffixLen: sortableInt64Bytes, extract: ExtractListUserKeyFromClaim}, + {name: "hash meta", prefix: HashMetaPrefix, extract: ExtractHashUserKeyFromMeta}, + {name: "hash field", prefix: HashFieldPrefix, extract: ExtractHashUserKeyFromField}, + {name: "hash delta", prefix: HashMetaDeltaPrefix, suffixLen: deltaKeyTSSize + deltaKeySeqSize, extract: ExtractHashUserKeyFromDelta}, + {name: "set meta", prefix: SetMetaPrefix, extract: ExtractSetUserKeyFromMeta}, + {name: "set member", prefix: SetMemberPrefix, extract: ExtractSetUserKeyFromMember}, + {name: "set delta", prefix: SetMetaDeltaPrefix, suffixLen: deltaKeyTSSize + deltaKeySeqSize, extract: ExtractSetUserKeyFromDelta}, + {name: "zset meta", prefix: ZSetMetaPrefix, extract: ExtractZSetUserKeyFromMeta}, + {name: "zset member", prefix: ZSetMemberPrefix, extract: ExtractZSetUserKeyFromMember}, + {name: "zset score", prefix: ZSetScorePrefix, suffixLen: zsetScalarSizeBytes, extract: ExtractZSetUserKeyFromScore}, + {name: "zset delta", prefix: ZSetMetaDeltaPrefix, suffixLen: deltaKeyTSSize + deltaKeySeqSize, extract: ExtractZSetUserKeyFromDelta}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + if got := tc.extract(malformedWideColumnStorageKey(tc.prefix, tc.suffixLen)); got != nil { + t.Fatalf("overflowing user-key length: want nil, got %q", got) + } + }) + } +} + +func TestWideColumnExtractorsRoundTripEmptyUserKey(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name string + key []byte + extract func([]byte) []byte + }{ + {name: "list claim", key: ListClaimKey(nil, 1), extract: ExtractListUserKeyFromClaim}, + {name: "hash meta", key: HashMetaKey(nil), extract: ExtractHashUserKeyFromMeta}, + {name: "hash field", key: HashFieldKey(nil, []byte("field")), extract: ExtractHashUserKeyFromField}, + {name: "hash delta", key: HashMetaDeltaKey(nil, 2, 3), extract: ExtractHashUserKeyFromDelta}, + {name: "set meta", key: SetMetaKey(nil), extract: ExtractSetUserKeyFromMeta}, + {name: "set member", key: SetMemberKey(nil, []byte("member")), extract: ExtractSetUserKeyFromMember}, + {name: "set delta", key: SetMetaDeltaKey(nil, 2, 3), extract: ExtractSetUserKeyFromDelta}, + {name: "zset meta", key: ZSetMetaKey(nil), extract: ExtractZSetUserKeyFromMeta}, + {name: "zset member", key: ZSetMemberKey(nil, []byte("member")), extract: ExtractZSetUserKeyFromMember}, + {name: "zset score", key: ZSetScoreKey(nil, 1.5, []byte("member")), extract: ExtractZSetUserKeyFromScore}, + {name: "zset delta", key: ZSetMetaDeltaKey(nil, 2, 3), extract: ExtractZSetUserKeyFromDelta}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + got := tc.extract(tc.key) + if got == nil { + t.Fatal("empty user-key extraction returned nil") + } + if len(got) != 0 { + t.Fatalf("empty user-key extraction: want empty, got %q", got) + } + }) + } +} + +func malformedWideColumnStorageKey(prefix string, suffixLen int) []byte { + key := make([]byte, 0, len(prefix)+wideColKeyLenSize+suffixLen) + key = append(key, prefix...) + var lenPrefix [wideColKeyLenSize]byte + binary.BigEndian.PutUint32(lenPrefix[:], ^uint32(0)) + key = append(key, lenPrefix[:]...) + key = append(key, make([]byte, suffixLen)...) + return key +} diff --git a/store/zset_helpers.go b/store/zset_helpers.go index 5424ee8fb..fe73fc007 100644 --- a/store/zset_helpers.go +++ b/store/zset_helpers.go @@ -274,53 +274,32 @@ func IsZSetMetaDeltaKey(key []byte) bool { // ExtractZSetUserKeyFromDelta extracts the logical user key from a zset delta key. func ExtractZSetUserKeyFromDelta(key []byte) []byte { - trimmed := bytes.TrimPrefix(key, []byte(ZSetMetaDeltaPrefix)) - minLen := wideColKeyLenSize + deltaKeyTSSize + deltaKeySeqSize - if len(trimmed) < minLen { - return nil - } - ukLen := binary.BigEndian.Uint32(trimmed[:wideColKeyLenSize]) - if uint32(len(trimmed)) < uint32(wideColKeyLenSize)+ukLen+uint32(deltaKeyTSSize+deltaKeySeqSize) { //nolint:gosec // constants fit in uint32 - return nil - } - return trimmed[wideColKeyLenSize : wideColKeyLenSize+ukLen] + return extractWideColumnUserKey(key, []byte(ZSetMetaDeltaPrefix), deltaKeyTSSize+deltaKeySeqSize, true) } // ExtractZSetUserKeyFromMeta extracts the logical user key from a zset meta key. func ExtractZSetUserKeyFromMeta(key []byte) []byte { - trimmed := bytes.TrimPrefix(key, []byte(ZSetMetaPrefix)) - if len(trimmed) < wideColKeyLenSize { - return nil - } - ukLen := binary.BigEndian.Uint32(trimmed[:wideColKeyLenSize]) - if uint32(len(trimmed)) < uint32(wideColKeyLenSize)+ukLen { //nolint:gosec // wideColKeyLenSize fits in uint32 - return nil - } - return trimmed[wideColKeyLenSize : wideColKeyLenSize+ukLen] + return extractWideColumnUserKey(key, []byte(ZSetMetaPrefix), 0, true) } // ExtractZSetUserKeyFromMember extracts the logical user key from a zset member key. func ExtractZSetUserKeyFromMember(key []byte) []byte { - trimmed := bytes.TrimPrefix(key, []byte(ZSetMemberPrefix)) - if len(trimmed) < wideColKeyLenSize { - return nil - } - ukLen := binary.BigEndian.Uint32(trimmed[:wideColKeyLenSize]) - if uint32(len(trimmed)) < uint32(wideColKeyLenSize)+ukLen { //nolint:gosec // wideColKeyLenSize fits in uint32 - return nil - } - return trimmed[wideColKeyLenSize : wideColKeyLenSize+ukLen] + return extractWideColumnUserKey(key, []byte(ZSetMemberPrefix), 0, false) } // ExtractZSetUserKeyFromScore extracts the logical user key from a zset score index key. func ExtractZSetUserKeyFromScore(key []byte) []byte { - trimmed := bytes.TrimPrefix(key, []byte(ZSetScorePrefix)) - if len(trimmed) < wideColKeyLenSize { - return nil - } - ukLen := binary.BigEndian.Uint32(trimmed[:wideColKeyLenSize]) - if uint32(len(trimmed)) < uint32(wideColKeyLenSize)+ukLen { //nolint:gosec // wideColKeyLenSize fits in uint32 - return nil - } - return trimmed[wideColKeyLenSize : wideColKeyLenSize+ukLen] + return extractWideColumnUserKey(key, []byte(ZSetScorePrefix), zsetScalarSizeBytes, false) +} + +// ExtractZSetUserKeyFromScoreScanPrefix extracts the user key from a zset +// score-index scan start/prefix that does not yet include the sortable score. +func ExtractZSetUserKeyFromScoreScanPrefix(key []byte) []byte { + return extractWideColumnUserKey(key, []byte(ZSetScorePrefix), 0, false) +} + +// ExtractZSetUserKeyFromDeltaScanPrefix extracts the user key from a zset +// metadata delta scan start/prefix. +func ExtractZSetUserKeyFromDeltaScanPrefix(key []byte) []byte { + return extractWideColumnUserKey(key, []byte(ZSetMetaDeltaPrefix), 0, false) }